Emit specialized OCaml initialization functions
This commit is contained in:
6 files changed
+732
-4
No files matched your search
@@ -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)
|
||||
|
||||
+574
@@ -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
|
||||
@@ -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
@@ -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);
|
||||
|
||||
@@ -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) );
|
||||
]
|
||||
@@ -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)
|
||||
Reference in new issue
Block a user