From 0f087dc75353386ab026ba761231f2204ac22d18 Mon Sep 17 00:00:00 2001 From: milner Date: Mon, 12 Jun 2017 18:44:00 +0000 Subject: [PATCH] Exercise generated programs with update sequences --- src/emit.ml | 240 ++++++++++++++++++++++++++++++++++++++++++- src/util.ml | 6 +- test/test_codegen.ml | 142 +++++++++++++++++++++++++ test/test_main.ml | 1 + 4 files changed, 385 insertions(+), 4 deletions(-) diff --git a/src/emit.ml b/src/emit.ml index 6264b23..fde27d6 100644 --- a/src/emit.ml +++ b/src/emit.ml @@ -547,13 +547,239 @@ let emit_output_functions plan 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)))) + (type_name (output_value_type plan)) (type_name (output_value_type plan))); + line buffer ""; + line buffer + " let equal_output (left : output) (right : output) = Delta_runtime.Pure_map.equal (fun a b -> a = b) left right") 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 output_change_to_string (change : output_change) : string = Printf.sprintf \"%+d\" change"; + line buffer ""; + line buffer " let equal_output (left : output) (right : output) = left = right"); + line buffer "" + +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))); + line buffer ""; + line buffer + " let equal_output (left : output) (right : output) = Delta_runtime.Pure_map.equal (fun a b -> a = b) left right") + 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 ""; + line buffer " let equal_output (left : output) (right : output) = left = right"); line buffer "" let emit_wire_type_functions plan buffer info = @@ -743,7 +969,14 @@ let emit_driver plan buffer = 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 " | Delta_runtime.Success (next, change) ->"; + line buffer + " let applied = Query.apply_output_change (Query.result !state) change in"; + line buffer " let current = Query.result next in"; + line buffer + " if not (Query.equal_output applied current) then"; + line buffer + " fail (Printf.sprintf \"output change does not reproduce the result after batch %d\" (index + 1));"; line buffer " state := next;"; line buffer " update_seconds := !update_seconds +. (Sys.time () -. before);"; line buffer " if !trace then"; @@ -772,6 +1005,7 @@ let emit_signature plan buffer = 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 " val equal_output : output -> output -> bool"; 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"; diff --git a/src/util.ml b/src/util.ml index 9aa34e3..d379182 100644 --- a/src/util.ml +++ b/src/util.ml @@ -19,7 +19,11 @@ let rec concat_map f = function | item :: rest -> f item @ concat_map f rest let rec list_init count f = - if count <= 0 then [] else f 0 :: list_init (count - 1) (fun index -> f (index + 1)) + if count <= 0 then [] + else + let head = f 0 in + let tail = list_init (count - 1) (fun index -> f (index + 1)) in + head :: tail let rec take count items = if count <= 0 then [] diff --git a/test/test_codegen.ml b/test/test_codegen.ml index 5412b9e..470cb91 100644 --- a/test/test_codegen.ml +++ b/test/test_codegen.ml @@ -546,3 +546,145 @@ let cli_cases = check "the full traversal counter appears" (String.length output > 0); Native.remove_dir dir ); ] + +let rec wire_of_value value = + match value with + | Value.VUnit -> Wire.WUnit + | Value.VInt number -> Wire.WInt number + | Value.VBool truth -> Wire.WBool truth + | Value.VString text -> Wire.WString text + | Value.VTuple items -> Wire.WTuple (List.map wire_of_value items) + | Value.VRecord (_, fields) -> Wire.WRecord (List.map (fun (label, item) -> (label, wire_of_value item)) fields) + | Value.VCollection entries -> + Wire.WCollection (List.map (fun (key, item) -> (key, wire_of_value item)) (Delta_runtime.Pure_map.bindings entries)) + +let rec value_of_wire value = + match value with + | Wire.WUnit -> Value.VUnit + | Wire.WInt number -> Value.VInt number + | Wire.WBool truth -> Value.VBool truth + | Wire.WString text -> Value.VString text + | Wire.WTuple items -> Value.VTuple (List.map value_of_wire items) + | Wire.WRecord fields -> Value.VRecord ("", List.map (fun (label, item) -> (label, value_of_wire item)) fields) + | Wire.WCollection entries -> + Value.VCollection + (Value.collection_of_list (List.map (fun (key, item) -> (key, value_of_wire item)) entries)) + +let write_entries path entries = + write_file path + (String.concat "" + (List.map (fun (key, value) -> Printf.sprintf "(%d %s)\n" key (Wire.to_string (wire_of_value value))) entries)) + +let write_batches path batches = + write_file path + (String.concat "" + (List.map + (fun ops -> + "(batch" + ^ String.concat "" + (List.map + (fun op -> + match op with + | Change.OpInsert (key, value) -> + Printf.sprintf " (insert %d %s)" key (Wire.to_string (wire_of_value value)) + | Change.OpRemove key -> Printf.sprintf " (remove %d)" key + | Change.OpReplace (key, value) -> + Printf.sprintf " (replace %d %s)" key (Wire.to_string (wire_of_value value))) + ops) + ^ ")\n") + batches)) + +let compile_query name text = + let typed = infer text in + let plan = Simplify.simplify (Graph.build (Anf.program (Specialize.program typed))) in + let dir = Filename.temp_file "delta_prog" "" 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); + match compile source_path exe_path with + | Some message -> Error ("compilation failed:\n" ^ message) + | None -> Ok (typed, exe_path, dir) + +let reference_outputs typed entries batches = + let state = ref (Value.collection_of_list entries) in + let results = ref [ Interpret.program typed (Delta_runtime.Pure_map.bindings !state) ] in + List.iter + (fun ops -> + match Change.validate_batch ~existing:!state ops with + | Change.Failure message -> + fail "reference" + (Printf.sprintf "%s with state %s" message + (Value.to_string (Value.VCollection !state))) + | Change.Success (temp, _) -> + state := temp; + results := !results @ [ Interpret.program typed (Delta_runtime.Pure_map.bindings temp) ]) + batches; + !results + +let generated_differential_case seed_count batch_count = + List.iter + (fun (name, text) -> + match compile_query name text with + | Error message -> fail name message + | Ok (typed, executable, dir) -> + let element = typed.Typed.tp_input_element in + List.iter + (fun seed -> + let rng = rng (seed + (77 * String.length name)) in + let entries = + Util.list_init (range rng 5 + 1) (fun index -> + (index + 1, Test_change.value_for rng typed.Typed.tp_records element)) + in + let rec generate count current acc = + if count <= 0 then List.rev acc + else + let ops = Test_incremental.random_batch rng typed.Typed.tp_records element current in + let next = + match Change.validate_batch ~existing:current ops with + | Change.Failure message -> + fail name (Printf.sprintf "seed %d: generated an invalid batch: %s" seed message); + current + | Change.Success (temp, _) -> temp + in + generate (count - 1) next (ops :: acc) + in + let batches = generate batch_count (Value.collection_of_list entries) [] in + let reference = reference_outputs typed entries batches in + let expected = + List.tl reference @ [ List.nth reference (List.length reference - 1) ] + in + let input_path = Filename.concat dir (Printf.sprintf "input_%d.sexp" seed) in + let updates_path = Filename.concat dir (Printf.sprintf "updates_%d.sexp" seed) in + write_entries input_path entries; + write_batches updates_path batches; + (match run [ executable; "--input"; input_path; "--updates"; updates_path; "--trace"; "--print-result" ] with + | Unix.WEXITED 0, output -> + let lines = List.filter (fun line -> line <> "") (String.split_on_char '\n' output) in + if List.length lines <> List.length expected then + fail name + (Printf.sprintf "seed %d: expected %d results but the program printed %d" seed + (List.length expected) (List.length lines)) + else + List.iteri + (fun index line -> + let want = Wire.to_string (wire_of_value (List.nth expected index)) in + if line <> want then + fail name + (Printf.sprintf "seed %d step %d: expected %s but got %s" seed index want line)) + lines + | _, output -> + fail name (Printf.sprintf "seed %d: the program failed:\n%s" seed output)); + Sys.remove input_path; + Sys.remove updates_path) + (Util.list_init seed_count (fun index -> index + 1)); + Native.remove_dir dir) + (List.filter (fun (name, _) -> name <> "strings") Test_incremental.differential_queries) + +let generated_differential_cases = + [ + ("generated programs match full evaluation", fun () -> generated_differential_case 12 25); + ("generated programs match full evaluation on a short run", + fun () -> generated_differential_case 2 3); + ] diff --git a/test/test_main.ml b/test/test_main.ml index b725d54..0dbf493 100644 --- a/test/test_main.ml +++ b/test/test_main.ml @@ -389,5 +389,6 @@ let () = Test_harness.run_suite "wire" Test_codegen.wire_cases; Test_harness.run_suite "cli" Test_codegen.cli_cases; Test_harness.run_suite "differential" Test_incremental.differential_cases; + Test_harness.run_suite "generated differential" Test_codegen.generated_differential_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)