From b4dee2f4f2738a62fa3a53ed97387897faf8a676 Mon Sep 17 00:00:00 2001 From: milner Date: Sun, 7 May 2017 11:39:00 +0000 Subject: [PATCH] Emit specialized OCaml initialization functions --- Makefile | 2 +- src/emit.ml | 574 +++++++++++++++++++++++++++++++++++++++++++ src/emit.mli | 3 + src/main.ml | 7 +- test/test_codegen.ml | 149 +++++++++++ test/test_main.ml | 1 + 6 files changed, 732 insertions(+), 4 deletions(-) create mode 100644 src/emit.ml create mode 100644 src/emit.mli create mode 100644 test/test_codegen.ml diff --git a/Makefile b/Makefile index 120b385..82ea085 100644 --- a/Makefile +++ b/Makefile @@ -150,7 +150,7 @@ $(TEST_EXE): $(TEST_CMX) $(LIB_ARCHIVE) $(RUNTIME_ARCHIVE) endif test: $(TEST_EXE) - DELTA_ROOT=$(CURDIR) DELTA_BUILD_DIR=$(CURDIR)/$(BUILD) ./$(TEST_EXE) + DELTA_ROOT=$(CURDIR) DELTA_BUILD_DIR=$(CURDIR)/$(BUILD) DELTA_OCAMLOPT="$(OCAMLOPT)" ./$(TEST_EXE) bench: $(BENCH_EXE) DELTA_ROOT=$(CURDIR) DELTA_BUILD_DIR=$(CURDIR)/$(BUILD) ./$(BENCH_EXE) diff --git a/src/emit.ml b/src/emit.ml new file mode 100644 index 0000000..cd68ba1 --- /dev/null +++ b/src/emit.ml @@ -0,0 +1,574 @@ +let sanitize name = + let buffer = Buffer.create (String.length name) in + String.iter + (fun character -> + if + (character >= 'a' && character <= 'z') + || (character >= 'A' && character <= 'Z') + || (character >= '0' && character <= '9') + || character = '_' + then Buffer.add_char buffer character + else Buffer.add_char buffer '_') + name; + let text = Buffer.contents buffer in + if text = "" then "v" else if text.[0] >= '0' && text.[0] <= '9' then "v" ^ text else text + +let rec type_name ty = + match Types.repr ty with + | Types.TInt -> "int" + | Types.TBool -> "bool" + | Types.TString -> "string" + | Types.TUnit -> "unit" + | Types.TRecord name -> sanitize name + | Types.TTuple items -> "t_" ^ String.concat "_" (List.map type_name items) + | Types.TCollection element -> "coll_" ^ type_name element + | Types.TArrow _ -> "fun" + | Types.TVar var -> Printf.sprintf "tvar%d" var.Types.id + +let rec ocaml_ty ty = + match Types.repr ty with + | Types.TInt -> "int" + | Types.TBool -> "bool" + | Types.TString -> "string" + | Types.TUnit -> "unit" + | Types.TRecord name -> sanitize name + | Types.TTuple items -> "(" ^ String.concat " * " (List.map ocaml_ty items) ^ ")" + | Types.TCollection element -> ocaml_ty element ^ " Delta_runtime.Pure_map.t" + | Types.TArrow (domain, codomain) -> "(" ^ ocaml_ty domain ^ " -> " ^ ocaml_ty codomain ^ ")" + | Types.TVar var -> Printf.sprintf "tvar%d" var.Types.id + +let rec change_ty ty = type_name ty ^ "_change" + +type tyinfo = { + ti_ty : Types.t; + ti_name : string; +} + +let same_type left right = Types.pp left = Types.pp right + +let rec add_subtypes plan add ty = + match Types.repr ty with + | Types.TTuple items -> + add ty; + List.iter (add_subtypes plan add) items + | Types.TRecord name -> + add ty; + (match Types.record_info plan.Graph.pl_records name with + | Some info -> List.iter (fun (_, field_ty) -> add_subtypes plan add field_ty) info.Types.ri_fields + | None -> ()) + | other -> add other + +let collect_types plan = + let acc = ref [] in + let add ty = + if not (List.exists (fun info -> same_type info.ti_ty ty) !acc) then + acc := !acc @ [ { ti_ty = ty; ti_name = type_name ty } ] + in + let add_with_subtypes ty = add_subtypes plan add ty in + let rec walk expr = + add_with_subtypes expr.Anf.aty; + match expr.Anf.a with + | Anf.ALet (_, bound, body) -> + walk bound; + walk body + | Anf.AIf (_, then_branch, else_branch) -> + walk then_branch; + walk else_branch + | Anf.ALambda (_, body) -> walk body + | Anf.AFilter (_, _, body) | Anf.AMap (_, _, body) -> walk body + | Anf.ABinop _ | Anf.AAtom _ | Anf.ATuple _ | Anf.ARecord _ | Anf.AField _ | Anf.AApp _ + | Anf.ASum _ | Anf.ACount _ -> + () + in + List.iter (fun node -> add_with_subtypes node.Graph.n_element) plan.Graph.pl_nodes; + List.iter + (fun node -> + match node.Graph.n_kind with + | Graph.Filter (_, body) | Graph.Map (_, body) -> walk body + | _ -> ()) + plan.Graph.pl_nodes; + List.iter (fun (_, ty) -> add_with_subtypes ty) plan.Graph.pl_types; + add_with_subtypes plan.Graph.pl_output; + let rec depth ty = + match Types.repr ty with + | Types.TTuple items -> + 1 + + List.fold_left + (fun acc item -> max acc (depth item)) + 0 items + | Types.TRecord name -> ( + match Types.record_info plan.Graph.pl_records name with + | Some info -> + 1 + List.fold_left (fun acc (_, field_ty) -> max acc (depth field_ty)) 0 info.Types.ri_fields + | None -> 0) + | _ -> 0 + in + List.stable_sort (fun a b -> compare (depth a.ti_ty) (depth b.ti_ty)) !acc + +let source_node plan = + match plan.Graph.pl_nodes with + | node :: _ when node.Graph.n_kind = Graph.Source -> Some node + | _ -> None + +let input_type plan = + match source_node plan with Some node -> node.Graph.n_element | None -> Types.TInt + +let is_collection_output plan = Types.is_collection plan.Graph.pl_output + +let output_value_type plan = + match plan.Graph.pl_output_element with Some element -> element | None -> Types.TInt + +let record_fields plan name = + match Types.record_info plan.Graph.pl_records name with + | Some info -> info.Types.ri_fields + | None -> [] + +let cache_field node = Printf.sprintf "cache_%d" node.Graph.n_id + +let accumulator_field node = Printf.sprintf "acc_%d" node.Graph.n_id + +let components ty = match Types.repr ty with Types.TTuple items -> items | _ -> [] + +let line buffer text = Buffer.add_string buffer (text ^ "\n") + +let indexed prefix index = Printf.sprintf "%s%d" prefix index + +let binder_names of_expr = + let binders = ref [] in + let rec walk expr = + match expr.Anf.a with + | Anf.ALet (ident, bound, body) -> + binders := !binders @ [ ident ]; + walk bound; + walk body + | Anf.AIf (_, then_branch, else_branch) -> + walk then_branch; + walk else_branch + | Anf.ALambda (ident, body) -> + binders := !binders @ [ ident ]; + walk body + | Anf.AFilter (_, parameter, body) | Anf.AMap (_, parameter, body) -> + binders := !binders @ [ parameter ]; + walk body + | Anf.ABinop _ | Anf.AAtom _ | Anf.ATuple _ | Anf.ARecord _ | Anf.AField _ | Anf.AApp _ + | Anf.ASum _ | Anf.ACount _ -> + () + in + walk of_expr; + let used = Hashtbl.create 32 in + List.map + (fun ident -> + let base = sanitize (Ident.display ident) in + let candidate = + if Hashtbl.mem used base then Printf.sprintf "%s_%d" base (Ident.stamp ident) else base + in + Hashtbl.replace used candidate (); + (Ident.stamp ident, candidate)) + !binders + +let var_name names ident = + match Util.assoc_opt (Ident.stamp ident) names with + | Some name -> name + | None -> sanitize (Ident.display ident) + +let rec render_expr names expr = + match expr.Anf.a with + | Anf.AAtom atom -> render_atom names atom + | Anf.ABinop (operator, left, right) -> + Printf.sprintf "(%s %s %s)" (render_atom names left) (Syntax.binop_name operator) + (render_atom names right) + | Anf.AIf (condition, then_branch, else_branch) -> + Printf.sprintf "(if %s then %s else %s)" (render_atom names condition) + (render_expr names then_branch) (render_expr names else_branch) + | Anf.ATuple atoms -> "(" ^ Util.join ", " (List.map (render_atom names) atoms) ^ ")" + | Anf.ARecord (name, fields) -> + "{ " + ^ Util.join "; " (List.map (fun (label, atom) -> label ^ " = " ^ render_atom names atom) fields) + ^ " }" + | Anf.AField (record, label) -> Printf.sprintf "%s.%s" (render_atom names record) label + | Anf.ALet (ident, bound, body) -> + Printf.sprintf "(let %s = %s in %s)" (var_name names ident) (render_expr names bound) + (render_expr names body) + | Anf.AApp _ | Anf.ALambda _ | Anf.AFilter _ | Anf.AMap _ | Anf.ASum _ | Anf.ACount _ -> + Diagnostic.error expr.Anf.aspan + "internal error: this expression cannot be compiled into scalar code" + +and render_atom names atom = + match atom with + | Anf.AInt value -> string_of_int value + | Anf.ABool true -> "true" + | Anf.ABool false -> "false" + | Anf.AString value -> Printf.sprintf "%S" value + | Anf.AUnit -> "()" + | Anf.AVar ident -> var_name names ident + +let rec render_statement buffer names env expr = + match expr.Anf.a with + | Anf.ALet (ident, bound, body) -> + let name = var_name names ident in + line buffer (Printf.sprintf " let %s = %s in" name (render_expr names bound)); + render_statement buffer names ((Ident.stamp ident, name) :: env) body + | _ -> render_expr names expr + +let emit_type_functions buffer plan info = + let name = info.ti_name in + match Types.repr info.ti_ty with + | Types.TInt -> + line buffer (Printf.sprintf "let empty_%s : %s_change = 0" name name); + line buffer + (Printf.sprintf "let apply_%s (value : %s) (change : %s_change) : %s = value + change" name name + name name); + line buffer + (Printf.sprintf + "let chg_%s (old_value : %s) (new_value : %s) : %s_change = new_value - old_value" name name + name name); + line buffer (Printf.sprintf "let show_%s (value : %s) : string = string_of_int value" name name) + | Types.TBool -> + line buffer (Printf.sprintf "let empty_%s : %s_change = None" name name); + line buffer + (Printf.sprintf + "let apply_%s (value : %s) (change : %s_change) : %s = match change with None -> value | Some next -> next" + name name name name); + line buffer + (Printf.sprintf + "let chg_%s (old_value : %s) (new_value : %s) : %s_change = if old_value = new_value then None else Some new_value" + name name name name); + line buffer (Printf.sprintf "let show_%s (value : %s) : string = string_of_bool value" name name) + | Types.TString -> + line buffer (Printf.sprintf "let empty_%s : %s_change = None" name name); + line buffer + (Printf.sprintf + "let apply_%s (value : %s) (change : %s_change) : %s = match change with None -> value | Some next -> next" + name name name name); + line buffer + (Printf.sprintf + "let chg_%s (old_value : %s) (new_value : %s) : %s_change = if old_value = new_value then None else Some new_value" + name name name name); + line buffer + (Printf.sprintf "let show_%s (value : %s) : string = Printf.sprintf %S value" name name "%S") + | Types.TUnit -> + line buffer (Printf.sprintf "let empty_%s : %s_change = ()" name name); + line buffer + (Printf.sprintf "let apply_%s (value : %s) (change : %s_change) : %s = ()" name name name name); + line buffer + (Printf.sprintf "let chg_%s (old_value : %s) (new_value : %s) : %s_change = ()" name name name + name); + line buffer (Printf.sprintf "let show_%s (value : %s) : string = %S" name name "()") + | Types.TTuple items -> ( + let count = List.length items in + if count >= 2 then ( + let names = List.map type_name items in + let pattern prefix = + "(" ^ Util.join ", " (Util.list_init count (fun index -> indexed prefix (index + 1))) ^ ")" + in + line buffer + (Printf.sprintf "let empty_%s : %s_change = (%s)" name name + (Util.join ", " (List.map (fun component -> "empty_" ^ component) names))); + line buffer (Printf.sprintf "let apply_%s (value : %s) (change : %s_change) : %s =" name name name name); + line buffer (Printf.sprintf " let %s = value in" (pattern "v")); + line buffer (Printf.sprintf " let %s = change in" (pattern "c")); + line buffer + (Printf.sprintf " (%s)" + (Util.join ", " + (Util.list_init count (fun index -> + Printf.sprintf "apply_%s %s %s" (List.nth names index) (indexed "v" (index + 1)) + (indexed "c" (index + 1)))))); + line buffer (Printf.sprintf "let chg_%s (old_value : %s) (new_value : %s) : %s_change =" name name name name); + line buffer (Printf.sprintf " let %s = old_value in" (pattern "o")); + line buffer (Printf.sprintf " let %s = new_value in" (pattern "n")); + line buffer + (Printf.sprintf " (%s)" + (Util.join ", " + (Util.list_init count (fun index -> + Printf.sprintf "chg_%s %s %s" (List.nth names index) (indexed "o" (index + 1)) + (indexed "n" (index + 1)))))); + line buffer (Printf.sprintf "let show_%s (value : %s) : string =" name name); + line buffer (Printf.sprintf " let %s = value in" (pattern "v")); + line buffer + (Printf.sprintf " Printf.sprintf %S %s" + ("(" ^ Util.join ", " (Util.list_init count (fun _ -> "%s")) ^ ")") + (Util.join " " + (Util.list_init count (fun index -> + Printf.sprintf "(show_%s %s)" (List.nth names index) (indexed "v" (index + 1)))))))) + | Types.TRecord record_name -> ( + match Types.record_info plan.Graph.pl_records record_name with + | None -> () + | Some info -> + let fields = info.Types.ri_fields in + line buffer + (Printf.sprintf "type %s_change = { %s }" name + (Util.join "; " + (List.map (fun (label, ty) -> Printf.sprintf "ch_%s : %s" label (change_ty ty)) fields))); + line buffer + (Printf.sprintf "let empty_%s : %s_change = { %s }" name name + (Util.join "; " + (List.map + (fun (label, ty) -> Printf.sprintf "ch_%s = empty_%s" label (type_name ty)) + fields))); + line buffer + (Printf.sprintf "let apply_%s (value : %s) (change : %s_change) : %s = { %s }" name name name + name + (Util.join "; " + (List.map + (fun (label, ty) -> + Printf.sprintf "%s = apply_%s value.%s change.ch_%s" label (type_name ty) label + label) + fields))); + line buffer + (Printf.sprintf "let chg_%s (old_value : %s) (new_value : %s) : %s_change = { %s }" name name + name name + (Util.join "; " + (List.map + (fun (label, ty) -> + Printf.sprintf "ch_%s = chg_%s old_value.%s new_value.%s" label (type_name ty) + label label) + fields))); + line buffer (Printf.sprintf "let show_%s (value : %s) : string =" name name); + line buffer + (Printf.sprintf " Printf.sprintf %S %s" + ("{ " ^ Util.join "; " (List.map (fun (label, _) -> label ^ " = %s") fields) ^ " }") + (Util.join " " + (List.map (fun (label, ty) -> Printf.sprintf "(show_%s value.%s)" (type_name ty) label) fields)))) + | Types.TCollection _ | Types.TArrow _ | Types.TVar _ -> () + +let emit_declarations buffer plan = + let declarations = + List.map + (fun info -> + Printf.sprintf " %s = { %s }" (sanitize info.Types.ri_name) + (Util.join "; " + (List.map (fun (label, ty) -> Printf.sprintf "%s : %s" label (ocaml_ty ty)) info.Types.ri_fields))) + plan.Graph.pl_records.Types.declarations + @ List.map + (fun info -> + match Types.repr info.ti_ty with + | Types.TTuple items -> + Printf.sprintf " %s = %s" info.ti_name (ocaml_ty info.ti_ty) + | _ -> "") + (collect_types plan) + @ [ " int_change = int"; " bool_change = bool option"; " string_change = string option"; " unit_change = unit" ] + @ List.map + (fun info -> + match Types.repr info.ti_ty with + | Types.TTuple items -> + Printf.sprintf " %s_change = (%s)" info.ti_name + (Util.join " * " (List.map change_ty items)) + | _ -> "") + (collect_types plan) + in + let declarations = List.filter (fun text -> text <> "") declarations in + (match declarations with + | [] -> () + | first :: rest -> + line buffer ("type" ^ String.sub first 1 (String.length first - 1)); + List.iter (fun item -> line buffer ("and" ^ String.sub item 1 (String.length item - 1))) rest); + line buffer "" + +let upstream_id plan node = + match node.Graph.n_input with + | Some id -> id + | None -> + Diagnostic.error node.Graph.n_span "internal error: node %d has no input collection" + node.Graph.n_id + +let node_input_element plan node = + match node.Graph.n_input with + | Some id -> (Graph.node_of_id plan id).Graph.n_element + | None -> node.Graph.n_element + +let emit_node_functions plan buffer = + List.iter + (fun node -> + match node.Graph.n_kind with + | Graph.Filter (parameter, body) -> + let names = binder_names body in + line buffer + (Printf.sprintf " let predicate_%d (%s : %s) : bool =" node.Graph.n_id + (var_name names parameter) (ocaml_ty (node_input_element plan node))); + line buffer + (Printf.sprintf " (Delta_runtime.count_predicate query_counters; %s)" (render_expr names body)); + line buffer "" + | Graph.Map (parameter, body) -> + let names = binder_names body in + line buffer + (Printf.sprintf " let mapping_%d (%s : %s) : %s =" node.Graph.n_id + (var_name names parameter) (ocaml_ty (node_input_element plan node)) + (ocaml_ty body.Anf.aty)); + line buffer + (Printf.sprintf " (Delta_runtime.count_mapping query_counters; %s)" (render_expr names body)); + line buffer "" + | Graph.Source | Graph.Sum | Graph.Count -> ()) + plan.Graph.pl_nodes + +let emit_state plan buffer = + line buffer " type state = {"; + line buffer (Printf.sprintf " input : %s Delta_runtime.Pure_map.t;" (ocaml_ty (input_type plan))); + List.iter + (fun node -> + match (node.Graph.n_kind, node.Graph.n_cache) with + | Graph.Source, _ -> () + | _, Graph.Cached_values -> + line buffer + (Printf.sprintf " %s : %s Delta_runtime.Pure_map.t;" (cache_field node) + (ocaml_ty node.Graph.n_element)) + | _, Graph.Accumulator -> + line buffer (Printf.sprintf " %s : int;" (accumulator_field node)) + | _, Graph.No_cache -> ()) + plan.Graph.pl_nodes; + line buffer " }"; + line buffer "" + +let emit_init plan buffer = + let input = input_type plan in + line buffer + (Printf.sprintf " let init (rows : (int * %s) list) : state =" (ocaml_ty input)); + line buffer " Delta_runtime.count_full_traversal query_counters;"; + line buffer + " let input = List.fold_left (fun acc (key, value) -> Delta_runtime.Pure_map.add key value acc) Delta_runtime.Pure_map.empty rows in"; + List.iter + (fun node -> + match node.Graph.n_kind with + | Graph.Source -> () + | Graph.Filter (_, _) -> + let upstream = Graph.node_of_id plan (upstream_id plan node) in + line buffer " Delta_runtime.count_full_traversal query_counters;"; + line buffer + (Printf.sprintf + " let %s = Delta_runtime.Pure_map.fold (fun key value acc -> if predicate_%d value then Delta_runtime.Pure_map.add key value acc else acc) %s Delta_runtime.Pure_map.empty in" + (cache_field node) node.Graph.n_id + (if upstream.Graph.n_kind = Graph.Source then "input" else cache_field upstream)) + | Graph.Map (_, _) -> + let upstream = Graph.node_of_id plan (upstream_id plan node) in + line buffer " Delta_runtime.count_full_traversal query_counters;"; + line buffer + (Printf.sprintf + " let %s = Delta_runtime.Pure_map.fold (fun key value acc -> Delta_runtime.Pure_map.add key (mapping_%d value) acc) %s Delta_runtime.Pure_map.empty in" + (cache_field node) node.Graph.n_id + (if upstream.Graph.n_kind = Graph.Source then "input" else cache_field upstream)) + | Graph.Sum -> + let upstream = Graph.node_of_id plan (upstream_id plan node) in + line buffer " Delta_runtime.count_full_traversal query_counters;"; + line buffer + (Printf.sprintf + " let %s = Delta_runtime.Pure_map.fold (fun _ value acc -> acc + value) %s 0 in" + (accumulator_field node) + (if upstream.Graph.n_kind = Graph.Source then "input" else cache_field upstream)) + | Graph.Count -> + let upstream = Graph.node_of_id plan (upstream_id plan node) in + line buffer " Delta_runtime.count_full_traversal query_counters;"; + line buffer + (Printf.sprintf " let %s = Delta_runtime.Pure_map.cardinal %s in" + (accumulator_field node) + (if upstream.Graph.n_kind = Graph.Source then "input" else cache_field upstream))) + plan.Graph.pl_nodes; + line buffer + (Printf.sprintf " { input; %s }" + (Util.join "; " + (Util.filter_map + (fun node -> + match (node.Graph.n_kind, node.Graph.n_cache) with + | Graph.Source, _ -> None + | _, Graph.Cached_values -> + Some (Printf.sprintf "%s = %s" (cache_field node) (cache_field node)) + | _, Graph.Accumulator -> + Some (Printf.sprintf "%s = %s" (accumulator_field node) (accumulator_field node)) + | _, Graph.No_cache -> None) + plan.Graph.pl_nodes))); + line buffer "" + +let emit_result plan buffer = + let node_by_id id = Graph.node_of_id plan id in + match plan.Graph.pl_result with + | Graph.Result_collection id -> + let node = node_by_id id in + (match node.Graph.n_kind with + | Graph.Sum | Graph.Count -> + line buffer + (Printf.sprintf " let result (state : state) : output = state.%s" + (accumulator_field node)) + | Graph.Source -> line buffer " let result (state : state) : output = state.input" + | Graph.Filter _ | Graph.Map _ -> + line buffer + (Printf.sprintf " let result (state : state) : output = state.%s" (cache_field node))) + | Graph.Result_scalar expr -> + let names = binder_names expr in + line buffer " let result (state : state) : output ="; + let env = + List.map + (fun (stamp, node_id) -> + (stamp, Printf.sprintf "state.%s" (accumulator_field (node_by_id node_id)))) + plan.Graph.pl_scalar_bindings + in + let body = render_statement buffer names env expr in + line buffer (Printf.sprintf " %s" body); + line buffer "" + +let emit_output_types plan buffer = + line buffer " type output ="; + if is_collection_output plan then ( + line buffer (Printf.sprintf " %s Delta_runtime.Pure_map.t" (ocaml_ty (output_value_type plan))); + line buffer ""; + line buffer " type output_change = (int * outval_map_change) list"; + line buffer ""; + line buffer + (Printf.sprintf " and outval_map_change = Ins of %s | Rem of %s | Rep of %s * %s" + (ocaml_ty (output_value_type plan)) (ocaml_ty (output_value_type plan)) + (ocaml_ty (output_value_type plan)) (ocaml_ty (output_value_type plan)))) + else ( + line buffer " int"; + line buffer ""; + line buffer " type output_change = int"); + line buffer "" + +let emit_output_functions plan buffer = + if is_collection_output plan then ( + line buffer + " let apply_output_change (value : output) (change : output_change) : output ="; + line buffer + " List.fold_left (fun acc (key, item) -> match item with Ins next -> Delta_runtime.Pure_map.add key next acc | Rem _ -> Delta_runtime.Pure_map.remove key acc | Rep (_, next) -> Delta_runtime.Pure_map.add key next acc) value change"; + line buffer ""; + line buffer " let output_to_string (value : output) : string ="; + line buffer + (Printf.sprintf + " \"[ \" ^ String.concat \"; \" (List.rev (List.rev_map (fun (key, item) -> Printf.sprintf \"(%%d, %%s)\" key (show_%s item)) (Delta_runtime.Pure_map.bindings value))) ^ \" ]\"" + (type_name (output_value_type plan)))) + else ( + line buffer " let apply_output_change (value : output) (change : output_change) : output = value + change"; + line buffer ""; + line buffer " let output_to_string (value : output) : string = string_of_int value"); + line buffer "" + +let emit_signature plan buffer = + line buffer " type state"; + line buffer " type output"; + line buffer " type output_change"; + line buffer (Printf.sprintf " val init : (int * %s) list -> state" (ocaml_ty (input_type plan))); + line buffer " val result : state -> output"; + line buffer " val apply_output_change : output -> output_change -> output"; + line buffer " val output_to_string : output -> string"; + line buffer " val counters : unit -> Delta_runtime.counters"; + line buffer " val reset_counters : unit -> unit" + +let emit_structure plan buffer = + emit_output_types plan buffer; + emit_state plan buffer; + emit_node_functions plan buffer; + emit_init plan buffer; + emit_result plan buffer; + emit_output_functions plan buffer; + line buffer " let counters () = Delta_runtime.copy_counters query_counters"; + line buffer ""; + line buffer " let reset_counters () = Delta_runtime.reset_counters query_counters" + +let program_to_string plan = + let buffer = Buffer.create 8192 in + emit_declarations buffer plan; + List.iter (fun info -> emit_type_functions buffer plan info) (collect_types plan); + line buffer "let query_counters = Delta_runtime.new_counters ()"; + line buffer ""; + line buffer "module Query : sig"; + emit_signature plan buffer; + line buffer "end = struct"; + emit_structure plan buffer; + line buffer "end"; + Buffer.contents buffer diff --git a/src/emit.mli b/src/emit.mli new file mode 100644 index 0000000..f9bf7c9 --- /dev/null +++ b/src/emit.mli @@ -0,0 +1,3 @@ +val program_to_string : Graph.plan -> string +val input_type : Graph.plan -> Types.t +val ocaml_ty : Types.t -> string diff --git a/src/main.ml b/src/main.ml index d5ca9ca..cb1f405 100644 --- a/src/main.ml +++ b/src/main.ml @@ -119,9 +119,10 @@ let dump stage path = | _ -> frontend_unavailable () let emit path output = - ignore (plan_source path); - ignore output; - frontend_unavailable () + let program = plan_source path in + let channel = open_out output in + output_string channel (Emit.program_to_string program); + close_out channel let build path output = ignore (plan_source path); diff --git a/test/test_codegen.ml b/test/test_codegen.ml new file mode 100644 index 0000000..6b764fa --- /dev/null +++ b/test/test_codegen.ml @@ -0,0 +1,149 @@ +open Test_harness + +let build_dir () = try Sys.getenv "DELTA_BUILD_DIR" with Not_found -> "_build" + +let ocamlopt () = try Sys.getenv "DELTA_OCAMLOPT" with Not_found -> "ocamlopt" + +let unique name = + Printf.sprintf "%s_%d_%d" name (Unix.getpid ()) (Random.self_init (); Random.int 1000000) + +let write_file path text = + let channel = open_out path in + output_string channel text; + close_out channel + +let run command = + let log = Filename.temp_file "delta_test" ".log" in + let fd = Unix.openfile log [ Unix.O_WRONLY; Unix.O_CREAT; Unix.O_TRUNC ] 0o600 in + let argv = Array.of_list command in + let pid = Unix.create_process argv.(0) argv Unix.stdin fd fd in + let status = snd (Unix.waitpid [] pid) in + Unix.close fd; + let text = Native.read_file log in + (try Sys.remove log with _ -> ()); + (status, text) + +let compile source_path output_path = + let command = + [ ocamlopt (); "-I"; build_dir (); "-I"; "+unix"; "-o"; output_path; + Filename.concat (build_dir ()) "delta_runtime.cmxa"; "unix.cmxa"; source_path ] + in + match run command with + | Unix.WEXITED 0, _ -> None + | _, text -> Some text + +let example_driver = + {| +let () = + let rows = + [ (1, { customer = "Ada"; total = 1500 }); + (2, { customer = "Bo"; total = 900 }); + (3, { customer = "Lin"; total = 2200 }) ] + in + let state = Query.init rows in + print_endline (Query.output_to_string (Query.result state)); + print_endline (Delta_runtime.counters_to_string (Query.counters ())) +|} + +let run_generated name plan driver = + let dir = Filename.temp_file "delta_gen" "" in + Sys.remove dir; + Unix.mkdir dir 0o700; + let source_path = Filename.concat dir (unique name ^ ".ml") in + let exe_path = Filename.concat dir (unique name ^ ".exe") in + write_file source_path (Emit.program_to_string plan ^ driver); + let result = + match compile source_path exe_path with + | Some message -> Error ("the generated program did not compile:\n" ^ message) + | None -> ( + match run [ exe_path ] with + | Unix.WEXITED 0, text -> Ok text + | _, text -> Error ("the generated program failed:\n" ^ text)) + in + Native.remove_dir dir; + result + +let sample_plan fixture = + let typed = infer (read_fixture fixture) in + Simplify.simplify (Graph.build (Anf.program (Specialize.program typed))) + +let codegen_cases = + [ + ( "the emitted module initializes the example query", + fun () -> + let plan = sample_plan "expensive_order.delta" in + match run_generated "expensive_orders" plan example_driver with + | Error message -> fail "emitted program" message + | Ok output -> + let lines = String.split_on_char '\n' output in + check_equal_string "output" + "[ (1, (\"Ada\", 300)); (3, (\"Lin\", 440)) ]" + (List.nth lines 0); + check "counters are reported" (String.length (List.nth lines 1) > 0) ); + ( "the emitted module handles an integer query", + fun () -> + let plan = sample_plan "revenue.delta" in + let driver = + {| +let () = + let rows = [ (1, { customer = "Ada"; total = 1000 }); (2, { customer = "Bo"; total = 250 }) ] in + let state = Query.init rows in + print_endline (Query.output_to_string (Query.result state)) +|} + in + (match run_generated "revenue" plan driver with + | Error message -> fail "emitted program" message + | Ok output -> check_equal_string "revenue" "250" (String.trim output)) ); + ( "the emitted module handles a count query", + fun () -> + let plan = sample_plan "count_large.delta" in + let driver = + {| +let () = + let rows = + [ (1, { customer = "Ada"; total = 1000 }); + (2, { customer = "Bo"; total = 250 }); + (3, { customer = "Lin"; total = 501 }) ] + in + let state = Query.init rows in + print_endline (Query.output_to_string (Query.result state)) +|} + in + (match run_generated "count_large" plan driver with + | Error message -> fail "emitted program" message + | Ok output -> check_equal_string "count" "2" (String.trim output)) ); + ( "the emitted module compiles for a tuple valued query", + fun () -> + let typed = + infer + "type line = { price : int; quantity : int }\ninput lines : collection line\nquery q = lines |> filter (fun l -> l.quantity > 0) |> map (fun l -> (l.price, l.price * l.quantity))\n" + in + let plan = Simplify.simplify (Graph.build (Anf.program (Specialize.program typed))) in + let driver = "let _ = Query.init []\n" in + (match run_generated "tuple_query" plan driver with + | Error message -> fail "emitted program" message + | Ok _ -> check "compiles and runs" true) ); + ( "the emitted module compiles for a boolean valued query", + fun () -> + let typed = + infer + "type line = { price : int; quantity : int }\ninput lines : collection line\nquery q = lines |> map (fun l -> (l.price > 0, l.quantity))\n" + in + let plan = Simplify.simplify (Graph.build (Anf.program (Specialize.program typed))) in + (match run_generated "bool_query" plan "let _ = Query.init []\n" with + | Error message -> fail "emitted program" message + | Ok _ -> check "compiles and runs" true) ); + ( "the generated module keeps generated identifiers distinct", + fun () -> + let typed = + infer + "input rows : collection int\nlet scale n = n * 2\nlet quad n = scale (scale n)\nquery q = rows |> map (fun r -> quad r + quad r)\n" + in + let plan = Simplify.simplify (Graph.build (Anf.program (Specialize.program typed))) in + let text = Emit.program_to_string plan in + check "no duplicated let binding in the emitted mapping function" + (not (Util.starts_with "internal error" text)); + (match run_generated "identifiers" plan "let _ = Query.init []\n" with + | Error message -> fail "emitted program" message + | Ok _ -> check "compiles and runs" true) ); + ] diff --git a/test/test_main.ml b/test/test_main.ml index 88148da..737a970 100644 --- a/test/test_main.ml +++ b/test/test_main.ml @@ -384,5 +384,6 @@ let () = Test_harness.run_suite "filter" Test_incremental.filter_cases; Test_harness.run_suite "aggregates" Test_incremental.aggregate_cases; Test_harness.run_suite "simplify" Test_incremental.simplify_cases; + Test_harness.run_suite "codegen" Test_codegen.codegen_cases; Printf.printf "%d cases, %d failures\n" (Test_harness.case_count ()) (Test_harness.failure_count ()); exit (if Test_harness.failure_count () = 0 then 0 else 1)