Emit specialized OCaml initialization functions

This commit is contained in:
sneeker committed 2017-05-07 11:39:00 +00:00
1 parent 84f6a9fdb1
commit 5f6bb17cbd
6 files changed
+732 -4

No files matched your search

+1 -1
View File
@@ -150,7 +150,7 @@ $(TEST_EXE): $(TEST_CMX) $(LIB_ARCHIVE) $(RUNTIME_ARCHIVE)
endif endif
test: $(TEST_EXE) 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) bench: $(BENCH_EXE)
DELTA_ROOT=$(CURDIR) DELTA_BUILD_DIR=$(CURDIR)/$(BUILD) ./$(BENCH_EXE) DELTA_ROOT=$(CURDIR) DELTA_BUILD_DIR=$(CURDIR)/$(BUILD) ./$(BENCH_EXE)
+574
View File
@@ -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
+3
View File
@@ -0,0 +1,3 @@
val program_to_string : Graph.plan -> string
val input_type : Graph.plan -> Types.t
val ocaml_ty : Types.t -> string
+4 -3
View File
@@ -119,9 +119,10 @@ let dump stage path =
| _ -> frontend_unavailable () | _ -> frontend_unavailable ()
let emit path output = let emit path output =
ignore (plan_source path); let program = plan_source path in
ignore output; let channel = open_out output in
frontend_unavailable () output_string channel (Emit.program_to_string program);
close_out channel
let build path output = let build path output =
ignore (plan_source path); ignore (plan_source path);
+149
View File
@@ -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) );
]
+1
View File
@@ -384,5 +384,6 @@ let () =
Test_harness.run_suite "filter" Test_incremental.filter_cases; Test_harness.run_suite "filter" Test_incremental.filter_cases;
Test_harness.run_suite "aggregates" Test_incremental.aggregate_cases; Test_harness.run_suite "aggregates" Test_incremental.aggregate_cases;
Test_harness.run_suite "simplify" Test_incremental.simplify_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 ()); 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) exit (if Test_harness.failure_count () = 0 then 0 else 1)