1077 lines
53 KiB
OCaml
1077 lines
53 KiB
OCaml
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 binders_of 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;
|
|
!binders
|
|
|
|
let names_of_idents idents =
|
|
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))
|
|
idents
|
|
|
|
let binder_names of_expr = names_of_idents (binders_of of_expr)
|
|
|
|
let names_with_suffix suffix names = List.map (fun (stamp, name) -> (stamp, name ^ suffix)) names
|
|
|
|
let atom_expr span ty atom = { Anf.a = Anf.AAtom atom; aty = ty; aspan = span }
|
|
|
|
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)));
|
|
line buffer "";
|
|
line buffer " let output_change_to_string (change : output_change) : string =";
|
|
line buffer
|
|
(Printf.sprintf
|
|
" String.concat \"; \" (List.rev (List.rev_map (fun (key, item) -> match item with Ins value -> Printf.sprintf \"insert %%d %%s\" key (show_%s value) | Rem value -> Printf.sprintf \"remove %%d %%s\" key (show_%s value) | Rep (old_value, new_value) -> Printf.sprintf \"replace %%d %%s %%s\" key (show_%s old_value) (show_%s new_value)) change))"
|
|
(type_name (output_value_type plan)) (type_name (output_value_type plan))
|
|
(type_name (output_value_type plan)) (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 "";
|
|
line buffer " let output_change_to_string (change : output_change) : string = Printf.sprintf \"%+d\" change");
|
|
line buffer ""
|
|
|
|
let emit_wire_type_functions plan buffer info =
|
|
let name = info.ti_name in
|
|
let decode_failure what =
|
|
Printf.sprintf
|
|
" | other -> Delta_runtime.Failure (Printf.sprintf \"expected %s for a value of type %s but got %%s\" (Wire.to_string other))"
|
|
what name
|
|
in
|
|
match Types.repr info.ti_ty with
|
|
| Types.TInt ->
|
|
line buffer (Printf.sprintf " let encode_%s (value : %s) : Wire.value = Wire.WInt value" name name);
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let decode_%s (value : Wire.value) : %s Delta_runtime.outcome = match value with Wire.WInt number -> Delta_runtime.Success number%s"
|
|
name name (decode_failure "an integer"))
|
|
| Types.TBool ->
|
|
line buffer (Printf.sprintf " let encode_%s (value : %s) : Wire.value = Wire.WBool value" name name);
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let decode_%s (value : Wire.value) : %s Delta_runtime.outcome = match value with Wire.WBool truth -> Delta_runtime.Success truth%s"
|
|
name name (decode_failure "a boolean"))
|
|
| Types.TString ->
|
|
line buffer (Printf.sprintf " let encode_%s (value : %s) : Wire.value = Wire.WString value" name name);
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let decode_%s (value : Wire.value) : %s Delta_runtime.outcome = match value with Wire.WString text -> Delta_runtime.Success text%s"
|
|
name name (decode_failure "a string"))
|
|
| Types.TUnit ->
|
|
line buffer (Printf.sprintf " let encode_%s (value : %s) : Wire.value = Wire.WUnit" name name);
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let decode_%s (value : Wire.value) : %s Delta_runtime.outcome = match value with Wire.WUnit -> Delta_runtime.Success ()%s"
|
|
name name (decode_failure "unit"))
|
|
| 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 encode_%s (value : %s) : Wire.value =" name name);
|
|
line buffer (Printf.sprintf " let %s = value in" ("(" ^ Util.join ", " (Util.list_init count (fun index -> indexed "v" (index + 1))) ^ ")"));
|
|
line buffer
|
|
(Printf.sprintf " Wire.WTuple [ %s ]"
|
|
(Util.join "; "
|
|
(Util.list_init count (fun index ->
|
|
Printf.sprintf "encode_%s %s" (List.nth names index) (indexed "v" (index + 1))))));
|
|
line buffer
|
|
(Printf.sprintf " let decode_%s (value : Wire.value) : %s Delta_runtime.outcome =" name name);
|
|
line buffer (Printf.sprintf " match value with");
|
|
line buffer (Printf.sprintf " | Wire.WTuple %s ->" (pattern "v"));
|
|
let rec chain index =
|
|
if index > count then
|
|
Printf.sprintf " Delta_runtime.Success (%s)"
|
|
("(" ^ Util.join ", " (Util.list_init count (fun position -> indexed "x" (position + 1))) ^ ")")
|
|
else
|
|
Printf.sprintf " (match decode_%s %s with Delta_runtime.Failure message -> Delta_runtime.Failure message | Delta_runtime.Success %s ->\n%s)"
|
|
(List.nth names (index - 1))
|
|
(indexed "v" index)
|
|
(indexed "x" index)
|
|
(chain (index + 1))
|
|
in
|
|
line buffer (chain 1);
|
|
line buffer (Printf.sprintf "%s" (decode_failure "a tuple of the right size"))))
|
|
| Types.TRecord record_name -> (
|
|
match Types.record_info plan.Graph.pl_records record_name with
|
|
| None -> ()
|
|
| Some record ->
|
|
let fields = record.Types.ri_fields in
|
|
line buffer (Printf.sprintf " let encode_%s (value : %s) : Wire.value =" name name);
|
|
line buffer
|
|
(Printf.sprintf " Wire.WRecord [ %s ]"
|
|
(Util.join "; "
|
|
(List.map
|
|
(fun (label, ty) ->
|
|
Printf.sprintf "(%S, encode_%s value.%s)" label (type_name ty) label)
|
|
fields)));
|
|
line buffer
|
|
(Printf.sprintf " let decode_%s (value : Wire.value) : %s Delta_runtime.outcome =" name name);
|
|
line buffer " match value with";
|
|
line buffer " | Wire.WRecord fields ->";
|
|
let rec chain remaining =
|
|
match remaining with
|
|
| [] ->
|
|
Printf.sprintf " Delta_runtime.Success { %s }"
|
|
(Util.join "; "
|
|
(List.map (fun (label, _) -> Printf.sprintf "%s = %s" label label) fields))
|
|
| (label, ty) :: rest ->
|
|
Printf.sprintf
|
|
" (match Wire.field_value %S fields with Delta_runtime.Failure message -> Delta_runtime.Failure message | Delta_runtime.Success raw ->\n (match decode_%s raw with Delta_runtime.Failure message -> Delta_runtime.Failure message | Delta_runtime.Success %s ->\n%s))"
|
|
label (type_name ty) label (chain rest)
|
|
in
|
|
line buffer (chain fields);
|
|
line buffer (decode_failure "a record"))
|
|
| Types.TCollection _ | Types.TArrow _ | Types.TVar _ -> ()
|
|
|
|
let emit_wire_functions plan buffer =
|
|
List.iter (fun info -> emit_wire_type_functions plan buffer info) (collect_types plan);
|
|
line buffer "";
|
|
let input = input_type plan in
|
|
line buffer
|
|
(Printf.sprintf " let decode_row (value : Wire.value) : %s Delta_runtime.outcome = decode_%s value"
|
|
(ocaml_ty input) (type_name input));
|
|
line buffer "";
|
|
line buffer " let encode_output (value : output) : Wire.value =";
|
|
(if is_collection_output plan then (
|
|
line buffer
|
|
(Printf.sprintf
|
|
" Wire.WCollection (List.rev (List.rev_map (fun (key, item) -> (key, encode_%s item)) (Delta_runtime.Pure_map.bindings value)))"
|
|
(type_name (output_value_type plan))))
|
|
else line buffer " Wire.WInt value");
|
|
line buffer "";
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let decode_batch_op (op : Wire.op) : batch_op Delta_runtime.outcome =\n match op with\n | Wire.WInsert (key, value) -> (match decode_%s value with Delta_runtime.Failure message -> Delta_runtime.Failure message | Delta_runtime.Success row -> Delta_runtime.Success (Insert (key, row)))\n | Wire.WRemove key -> Delta_runtime.Success (Remove key)\n | Wire.WReplace (key, value) -> (match decode_%s value with Delta_runtime.Failure message -> Delta_runtime.Failure message | Delta_runtime.Success row -> Delta_runtime.Success (Replace (key, row)))"
|
|
(type_name input) (type_name input));
|
|
line buffer ""
|
|
|
|
let emit_driver plan buffer =
|
|
let input = ocaml_ty (input_type plan) in
|
|
line buffer "";
|
|
line buffer "let () =";
|
|
line buffer " let input_file = ref None in";
|
|
line buffer " let updates_file = ref None in";
|
|
line buffer " let print_result = ref false in";
|
|
line buffer " let stats = ref false in";
|
|
line buffer " let trace = ref false in";
|
|
line buffer " let fail message = prerr_endline (\"error: \" ^ message); exit 1 in";
|
|
line buffer " let usage message = prerr_endline (\"usage: \" ^ message); exit 2 in";
|
|
line buffer " let rec parse_arguments arguments =";
|
|
line buffer " match arguments with";
|
|
line buffer " | [] -> ()";
|
|
line buffer " | \"--input\" :: file :: rest -> input_file := Some file; parse_arguments rest";
|
|
line buffer " | \"--updates\" :: file :: rest -> updates_file := Some file; parse_arguments rest";
|
|
line buffer " | \"--print-result\" :: rest -> print_result := true; parse_arguments rest";
|
|
line buffer " | \"--stats\" :: rest -> stats := true; parse_arguments rest";
|
|
line buffer " | \"--trace\" :: rest -> trace := true; parse_arguments rest";
|
|
line buffer " | flag :: _ -> usage (\"unknown option \" ^ flag)";
|
|
line buffer " in";
|
|
line buffer " parse_arguments (List.tl (Array.to_list Sys.argv));";
|
|
line buffer " let started = Sys.time () in";
|
|
line buffer " let rows =";
|
|
line buffer " match !input_file with";
|
|
line buffer " | None -> usage \"the --input FILE option is required\"";
|
|
line buffer " | Some file -> (";
|
|
line buffer " match Wire.read_file file with";
|
|
line buffer " | Delta_runtime.Failure message -> fail message";
|
|
line buffer " | Delta_runtime.Success text -> (";
|
|
line buffer " match Wire.parse_input text with";
|
|
line buffer " | Delta_runtime.Failure message -> fail (file ^ \": \" ^ message)";
|
|
line buffer
|
|
(Printf.sprintf
|
|
" | Delta_runtime.Success entries ->\n let rec decode acc pending =\n match pending with\n | [] -> List.rev acc\n | (key, value) :: rest ->\n (match Query.decode_row value with\n | Delta_runtime.Failure message ->\n fail (Printf.sprintf \"%%s: key %%d: %%s\" file key message)\n | Delta_runtime.Success row ->\n if Delta_runtime.Pure_map.mem key (List.fold_left (fun map (key, _) -> Delta_runtime.Pure_map.add key () map) Delta_runtime.Pure_map.empty acc) then\n fail (Printf.sprintf \"%%s: duplicate key %%d\" file key)\n else decode ((key, row) :: acc) rest)\n in\n decode [] entries));");
|
|
line buffer " in";
|
|
line buffer
|
|
(Printf.sprintf " let state = ref (Query.init (rows : (int * %s) list)) in" input);
|
|
line buffer " let init_seconds = Sys.time () -. started in";
|
|
line buffer " let init_counters = Delta_runtime.copy_counters (Query.counters ()) in";
|
|
line buffer " Query.reset_counters ();";
|
|
line buffer " let minor_words_before = Gc.minor_words () in";
|
|
line buffer " let update_seconds = ref 0.0 in";
|
|
line buffer " (match !updates_file with";
|
|
line buffer " | None -> ()";
|
|
line buffer " | Some file -> (";
|
|
line buffer " match Wire.read_file file with";
|
|
line buffer " | Delta_runtime.Failure message -> fail message";
|
|
line buffer " | Delta_runtime.Success text -> (";
|
|
line buffer " match Wire.parse_updates text with";
|
|
line buffer " | Delta_runtime.Failure message -> fail (file ^ \": \" ^ message)";
|
|
line buffer " | Delta_runtime.Success batches ->";
|
|
line buffer " List.iteri";
|
|
line buffer " (fun index batch ->";
|
|
line buffer " let rec decode acc pending =";
|
|
line buffer " match pending with";
|
|
line buffer " | [] -> List.rev acc";
|
|
line buffer " | op :: rest -> (";
|
|
line buffer " match Query.decode_batch_op op with";
|
|
line buffer " | Delta_runtime.Failure message ->";
|
|
line buffer
|
|
" fail (Printf.sprintf \"%s: batch %d: %s\" file (index + 1) message)";
|
|
line buffer " | Delta_runtime.Success operation -> decode (operation :: acc) rest)";
|
|
line buffer " in";
|
|
line buffer " let operations = decode [] batch in";
|
|
line buffer " let before = Sys.time () in";
|
|
line buffer " (match Query.apply_batch !state operations with";
|
|
line buffer " | Delta_runtime.Failure message ->";
|
|
line buffer
|
|
" fail (Printf.sprintf \"%s: batch %d: %s\" file (index + 1) message)";
|
|
line buffer " | Delta_runtime.Success (next, _) ->";
|
|
line buffer " state := next;";
|
|
line buffer " update_seconds := !update_seconds +. (Sys.time () -. before);";
|
|
line buffer " if !trace then";
|
|
line buffer
|
|
" print_endline (Wire.to_string (Query.encode_output (Query.result next)))))";
|
|
line buffer " batches)));";
|
|
line buffer " let minor_words_after = Gc.minor_words () in";
|
|
line buffer " if !print_result then";
|
|
line buffer " print_endline (Wire.to_string (Query.encode_output (Query.result !state)));";
|
|
line buffer " if !stats then";
|
|
line buffer
|
|
" Printf.eprintf \"init_counters: %s\\nupdate_counters: %s\\ninit_seconds: %.6f\\nupdate_seconds: %.6f\\nminor_words: %.0f\\n\" (Delta_runtime.counters_to_string init_counters) (Delta_runtime.counters_to_string (Query.counters ())) init_seconds !update_seconds (minor_words_after -. minor_words_before);";
|
|
line buffer " exit 0"
|
|
|
|
let emit_signature plan buffer =
|
|
line buffer " type state";
|
|
line buffer " type output";
|
|
line buffer " type output_change";
|
|
line buffer
|
|
(Printf.sprintf " type batch_op = Insert of int * %s | Remove of int | Replace of int * %s"
|
|
(ocaml_ty (input_type plan)) (ocaml_ty (input_type plan)));
|
|
line buffer (Printf.sprintf " val init : (int * %s) list -> state" (ocaml_ty (input_type plan)));
|
|
line buffer
|
|
" val apply_batch : state -> batch_op list -> (state * output_change) Delta_runtime.outcome";
|
|
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 output_change_to_string : output_change -> string";
|
|
line buffer
|
|
(Printf.sprintf " val decode_row : Wire.value -> %s Delta_runtime.outcome" (ocaml_ty (input_type plan)));
|
|
line buffer " val encode_output : output -> Wire.value";
|
|
line buffer " val decode_batch_op : Wire.op -> batch_op Delta_runtime.outcome";
|
|
line buffer " val counters : unit -> Delta_runtime.counters";
|
|
line buffer " val reset_counters : unit -> unit"
|
|
|
|
let emit_batch_type plan buffer =
|
|
line buffer
|
|
(Printf.sprintf " type batch_op = Insert of int * %s | Remove of int | Replace of int * %s"
|
|
(ocaml_ty (input_type plan)) (ocaml_ty (input_type plan)));
|
|
line buffer ""
|
|
|
|
let rec emit_delta_functions plan buffer =
|
|
List.iter
|
|
(fun node ->
|
|
match node.Graph.n_kind with
|
|
| Graph.Map (parameter, body) ->
|
|
let names = names_of_idents (parameter :: binders_of body) in
|
|
let old_names = names_with_suffix "_old" names in
|
|
let new_names = names_with_suffix "_new" names in
|
|
let parameter_name = sanitize (Ident.display parameter) in
|
|
let element = node_input_element plan node in
|
|
line buffer
|
|
(Printf.sprintf " let delta_%d (%s_old : %s) (%s_new : %s) : %s =" node.Graph.n_id
|
|
parameter_name (ocaml_ty element) parameter_name (ocaml_ty element)
|
|
(change_ty body.Anf.aty));
|
|
let rec statements expr =
|
|
match expr.Anf.a with
|
|
| Anf.ALet (ident, bound, rest) ->
|
|
let name = var_name names ident in
|
|
line buffer (Printf.sprintf " let %s_old = %s in" name (render_expr old_names bound));
|
|
line buffer (Printf.sprintf " let %s_new = %s in" name (render_expr new_names bound));
|
|
statements rest
|
|
| _ -> ()
|
|
in
|
|
line buffer " Delta_runtime.count_scalar_delta query_counters;";
|
|
statements body;
|
|
line buffer (Printf.sprintf " %s" (render_delta plan old_names new_names body));
|
|
line buffer ""
|
|
| _ -> ())
|
|
plan.Graph.pl_nodes
|
|
|
|
and render_delta plan old_names new_names expr =
|
|
match expr.Anf.a with
|
|
| Anf.AAtom (Anf.AInt _ | Anf.ABool _ | Anf.AString _ | Anf.AUnit) -> "empty_" ^ type_name expr.Anf.aty
|
|
| Anf.AAtom (Anf.AVar _) -> (
|
|
match Types.repr expr.Anf.aty with
|
|
| Types.TInt -> Printf.sprintf "(%s - %s)" (render_expr new_names expr) (render_expr old_names expr)
|
|
| _ ->
|
|
Printf.sprintf "(chg_%s %s %s)" (type_name expr.Anf.aty) (render_expr old_names expr)
|
|
(render_expr new_names expr))
|
|
| Anf.ABinop (Syntax.Add, left, right) ->
|
|
Printf.sprintf "(%s + %s)"
|
|
(render_delta_atom plan old_names new_names expr.Anf.aspan left)
|
|
(render_delta_atom plan old_names new_names expr.Anf.aspan right)
|
|
| Anf.ABinop (Syntax.Sub, left, right) ->
|
|
Printf.sprintf "(%s - %s)"
|
|
(render_delta_atom plan old_names new_names expr.Anf.aspan left)
|
|
(render_delta_atom plan old_names new_names expr.Anf.aspan right)
|
|
| Anf.ABinop (Syntax.Mul, left, right) ->
|
|
let left_old = render_expr old_names (atom_expr expr.Anf.aspan Types.TInt left) in
|
|
let right_old = render_expr old_names (atom_expr expr.Anf.aspan Types.TInt right) in
|
|
let left_delta = render_delta_atom plan old_names new_names expr.Anf.aspan left in
|
|
let right_delta = render_delta_atom plan old_names new_names expr.Anf.aspan right in
|
|
Printf.sprintf "((%s * %s) + (%s * %s) + (%s * %s))" left_old right_delta right_old left_delta
|
|
left_delta right_delta
|
|
| _ ->
|
|
Printf.sprintf "(chg_%s %s %s)" (type_name expr.Anf.aty) (render_expr old_names expr)
|
|
(render_expr new_names expr)
|
|
|
|
and render_delta_atom plan old_names new_names span atom =
|
|
render_delta plan old_names new_names (atom_expr span Types.TInt atom)
|
|
|
|
let emit_update_functions plan buffer =
|
|
let nodes = plan.Graph.pl_nodes in
|
|
List.iteri
|
|
(fun index node ->
|
|
let keyword = if index = 0 then "let rec" else "and" in
|
|
let input_element = node_input_element plan node in
|
|
let consumers = plan.Graph.pl_consumers.(node.Graph.n_id) in
|
|
let forward old_change new_change =
|
|
Printf.sprintf " (match (%s, %s) with\n | None, None -> state\n | _ ->\n%s)"
|
|
old_change new_change
|
|
(Util.join "\n"
|
|
(List.map
|
|
(fun consumer ->
|
|
Printf.sprintf " update_%d state key %s %s" consumer old_change new_change)
|
|
consumers))
|
|
in
|
|
line buffer
|
|
(Printf.sprintf
|
|
" %s update_%d (state : state) (key : int) (up_old : %s option) (up_new : %s option) : state ="
|
|
keyword node.Graph.n_id (ocaml_ty input_element) (ocaml_ty input_element));
|
|
line buffer " Delta_runtime.count_changed_key query_counters;";
|
|
match node.Graph.n_kind with
|
|
| Graph.Source ->
|
|
line buffer
|
|
" let input = (match up_new with None -> Delta_runtime.Pure_map.remove key state.input | Some value -> Delta_runtime.Pure_map.add key value state.input) in";
|
|
if consumers = [] then line buffer " { state with input }"
|
|
else (
|
|
line buffer " let state = { state with input } in";
|
|
line buffer (forward "up_old" "up_new"))
|
|
| Graph.Filter (_, _) ->
|
|
let cache = cache_field node in
|
|
line buffer
|
|
(Printf.sprintf " let old_member = Delta_runtime.Pure_map.find_opt key state.%s in" cache);
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let new_member = (match up_new with None -> None | Some value -> if predicate_%d value then Some value else None) in"
|
|
node.Graph.n_id);
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let cached = (match (old_member, new_member) with None, None -> state.%s | None, Some value -> Delta_runtime.Pure_map.add key value state.%s | Some _, None -> Delta_runtime.Pure_map.remove key state.%s | Some _, Some value -> Delta_runtime.Pure_map.add key value state.%s) in"
|
|
cache cache cache cache);
|
|
line buffer (Printf.sprintf " let state = { state with %s = cached } in" cache);
|
|
if consumers = [] then line buffer " state"
|
|
else line buffer (forward "old_member" "new_member")
|
|
| Graph.Map (_, _) ->
|
|
let cached_values = node.Graph.n_cache = Graph.Cached_values in
|
|
let element = node.Graph.n_element in
|
|
(if cached_values then
|
|
line buffer
|
|
(Printf.sprintf " let old_value = Delta_runtime.Pure_map.find_opt key state.%s in"
|
|
(cache_field node))
|
|
else line buffer " let old_value = None in");
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let new_value = (match up_new with None -> None | Some value -> %s) in"
|
|
(if cached_values then
|
|
Printf.sprintf
|
|
"(match up_old with None -> Some (mapping_%d value) | Some previous -> (match old_value with Some cached_value -> Some (apply_%s cached_value (delta_%d previous value)) | None -> Some (mapping_%d value)))"
|
|
node.Graph.n_id (type_name element) node.Graph.n_id node.Graph.n_id
|
|
else Printf.sprintf "Some (mapping_%d value)" node.Graph.n_id));
|
|
(if cached_values then
|
|
let cache = cache_field node in
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let cached = (match (old_value, new_value) with None, None -> state.%s | None, Some value -> Delta_runtime.Pure_map.add key value state.%s | Some _, None -> Delta_runtime.Pure_map.remove key state.%s | Some _, Some value -> Delta_runtime.Pure_map.add key value state.%s) in"
|
|
cache cache cache cache);
|
|
line buffer (Printf.sprintf " let state = { state with %s = cached } in" cache)
|
|
else
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let old_value = (match up_old with None -> None | Some previous -> Some (mapping_%d previous)) in"
|
|
node.Graph.n_id));
|
|
if consumers = [] then line buffer " state" else line buffer (forward "old_value" "new_value")
|
|
| Graph.Sum ->
|
|
line buffer
|
|
" let delta = (match up_new with None -> 0 | Some value -> value) - (match up_old with None -> 0 | Some value -> value) in";
|
|
if consumers = [] then
|
|
line buffer
|
|
(Printf.sprintf " { state with %s = state.%s + delta }" (accumulator_field node)
|
|
(accumulator_field node))
|
|
else (
|
|
line buffer
|
|
(Printf.sprintf " let state = { state with %s = state.%s + delta } in"
|
|
(accumulator_field node) (accumulator_field node));
|
|
line buffer (forward "None" "None"))
|
|
| Graph.Count ->
|
|
line buffer
|
|
" let delta = (match (up_old, up_new) with None, None -> 0 | None, Some _ -> 1 | Some _, None -> -1 | Some _, Some _ -> 0) in";
|
|
if consumers = [] then
|
|
line buffer
|
|
(Printf.sprintf " { state with %s = state.%s + delta }" (accumulator_field node)
|
|
(accumulator_field node))
|
|
else (
|
|
line buffer
|
|
(Printf.sprintf " let state = { state with %s = state.%s + delta } in"
|
|
(accumulator_field node) (accumulator_field node));
|
|
line buffer (forward "None" "None"));
|
|
line buffer "")
|
|
nodes
|
|
|
|
let root_observation plan =
|
|
match plan.Graph.pl_result with
|
|
| Graph.Result_collection id -> (
|
|
let node = Graph.node_of_id plan id in
|
|
match node.Graph.n_kind with
|
|
| Graph.Source -> "Delta_runtime.Pure_map.find_opt key state.input"
|
|
| Graph.Filter _ | Graph.Map _ ->
|
|
Printf.sprintf "Delta_runtime.Pure_map.find_opt key state.%s" (cache_field node)
|
|
| Graph.Sum | Graph.Count -> Printf.sprintf "state.%s" (accumulator_field node))
|
|
| Graph.Result_scalar _ -> "0"
|
|
|
|
let emit_apply_key plan buffer =
|
|
let input_element = input_type plan in
|
|
let source = match source_node plan with Some node -> node.Graph.n_id | None -> -1 in
|
|
let scalar_env =
|
|
List.map
|
|
(fun (stamp, node_id) ->
|
|
(stamp, Printf.sprintf "state.%s" (accumulator_field (Graph.node_of_id plan node_id))))
|
|
plan.Graph.pl_scalar_bindings
|
|
in
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let apply_key (state : state) (key : int) (up_old : %s option) (up_new : %s option) : state * output_change ="
|
|
(ocaml_ty input_element) (ocaml_ty input_element));
|
|
line buffer (Printf.sprintf " let before = %s in" (root_observation plan));
|
|
line buffer (Printf.sprintf " let state = update_%d state key up_old up_new in" source);
|
|
(match plan.Graph.pl_result with
|
|
| Graph.Result_collection id -> (
|
|
let node = Graph.node_of_id plan id in
|
|
match node.Graph.n_kind with
|
|
| Graph.Sum | Graph.Count ->
|
|
line buffer
|
|
(Printf.sprintf " (state, state.%s - before)" (accumulator_field node))
|
|
| Graph.Source | Graph.Filter _ | Graph.Map _ ->
|
|
line buffer (Printf.sprintf " let after = %s in" (root_observation plan));
|
|
line buffer
|
|
" let change = (match (before, after) with None, None -> [] | None, Some value -> [ (key, Ins value) ] | Some value, None -> [ (key, Rem value) ] | Some old_value, Some new_value -> if old_value = new_value then [] else [ (key, Rep (old_value, new_value)) ]) in";
|
|
line buffer " (state, change)")
|
|
| Graph.Result_scalar expr ->
|
|
let names = names_of_idents (binders_of expr) in
|
|
let env = names @ scalar_env in
|
|
line buffer (Printf.sprintf " let after = %s in" (render_expr env expr));
|
|
line buffer " (state, after - before)");
|
|
line buffer ""
|
|
|
|
let emit_apply_batch plan buffer =
|
|
let input_element = input_type plan in
|
|
let empty_output = if is_collection_output plan then "[]" else "0" in
|
|
let combine = if is_collection_output plan then "(!output @ change)" else "(!output + change)" in
|
|
line buffer
|
|
(Printf.sprintf
|
|
" let apply_batch (state : state) (ops : batch_op list) : (state * output_change) Delta_runtime.outcome ="
|
|
);
|
|
line buffer " let rec validate validated touched pending =";
|
|
line buffer " match pending with";
|
|
line buffer " | [] -> Delta_runtime.Success (validated, touched)";
|
|
line buffer
|
|
" | Insert (key, value) :: rest -> if Delta_runtime.Pure_map.mem key validated then Delta_runtime.Failure (Printf.sprintf \"cannot insert key %d: it is already present\" key) else validate (Delta_runtime.Pure_map.add key value validated) (key :: touched) rest";
|
|
line buffer
|
|
" | Remove key :: rest -> if Delta_runtime.Pure_map.mem key validated then validate (Delta_runtime.Pure_map.remove key validated) (key :: touched) rest else Delta_runtime.Failure (Printf.sprintf \"cannot remove key %d: it is not present\" key)";
|
|
line buffer
|
|
" | Replace (key, value) :: rest -> if Delta_runtime.Pure_map.mem key validated then validate (Delta_runtime.Pure_map.add key value validated) (key :: touched) rest else Delta_runtime.Failure (Printf.sprintf \"cannot replace key %d: it is not present\" key)";
|
|
line buffer " in";
|
|
line buffer " let original = state.input in";
|
|
line buffer " let rec collect validated keys pending =";
|
|
line buffer " match pending with";
|
|
line buffer " | [] -> List.rev keys";
|
|
line buffer " | key :: rest -> (";
|
|
line buffer
|
|
" match (Delta_runtime.Pure_map.find_opt key original, Delta_runtime.Pure_map.find_opt key validated) with";
|
|
line buffer " | None, None -> collect validated keys rest";
|
|
line buffer " | None, Some value -> collect validated ((key, None, Some value) :: keys) rest";
|
|
line buffer " | Some value, None -> collect validated ((key, Some value, None) :: keys) rest";
|
|
line buffer
|
|
" | Some old_value, Some new_value -> if old_value = new_value then collect validated keys rest else collect validated ((key, Some old_value, Some new_value) :: keys) rest)";
|
|
line buffer " in";
|
|
line buffer " match validate original [] ops with";
|
|
line buffer " | Delta_runtime.Failure message -> Delta_runtime.Failure message";
|
|
line buffer " | Delta_runtime.Success (validated, touched) ->";
|
|
line buffer " let items = collect validated [] (List.sort_uniq compare touched) in";
|
|
line buffer " (try";
|
|
line buffer " let current = ref state in";
|
|
line buffer (Printf.sprintf " let output = ref %s in" empty_output);
|
|
line buffer " List.iter (fun (key, up_old, up_new) ->";
|
|
line buffer " let (next_state, change) = apply_key !current key up_old up_new in";
|
|
line buffer " current := next_state;";
|
|
line buffer (Printf.sprintf " output := %s) items;" combine);
|
|
line buffer " Delta_runtime.Success (!current, !output)";
|
|
line buffer " with";
|
|
line buffer " | Division_by_zero -> Delta_runtime.Failure \"division by zero while applying the batch\")";
|
|
line buffer ""
|
|
|
|
let emit_structure plan buffer =
|
|
emit_output_types plan buffer;
|
|
emit_state plan buffer;
|
|
emit_batch_type plan buffer;
|
|
emit_delta_functions plan buffer;
|
|
emit_node_functions plan buffer;
|
|
emit_update_functions plan buffer;
|
|
emit_apply_key plan buffer;
|
|
emit_apply_batch plan buffer;
|
|
emit_init plan buffer;
|
|
emit_result plan buffer;
|
|
emit_output_functions plan buffer;
|
|
emit_wire_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 module_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
|
|
|
|
let program_to_string plan =
|
|
let module_source = module_to_string plan in
|
|
let buffer = Buffer.create (String.length module_source + 4096) in
|
|
Buffer.add_string buffer module_source;
|
|
emit_driver plan buffer;
|
|
Buffer.contents buffer
|