From 780136e802711d607b61ad97741116f4cc181bd4 Mon Sep 17 00:00:00 2001 From: milner Date: Fri, 3 Mar 2017 15:09:00 +0000 Subject: [PATCH] Normalize query expressions into ANF --- src/anf.ml | 264 ++++++++++++++++++++++++++++++++++++++++++++++ src/anf.mli | 42 ++++++++ src/main.ml | 11 +- test/test_main.ml | 1 + test/test_type.ml | 91 ++++++++++++++++ 5 files changed, 406 insertions(+), 3 deletions(-) create mode 100644 src/anf.ml create mode 100644 src/anf.mli diff --git a/src/anf.ml b/src/anf.ml new file mode 100644 index 0000000..f340a40 --- /dev/null +++ b/src/anf.ml @@ -0,0 +1,264 @@ +type atom = + | AInt of int + | ABool of bool + | AString of string + | AUnit + | AVar of Ident.t + +type expr = { + a : desc; + aty : Types.t; + aspan : Location.span; +} + +and desc = + | AAtom of atom + | ALet of Ident.t * expr * expr + | ABinop of Syntax.binop * atom * atom + | AIf of atom * expr * expr + | ATuple of atom list + | ARecord of string * (string * atom) list + | AField of atom * string + | AApp of Ident.t * atom + | ALambda of Ident.t * expr + | AFilter of atom * Ident.t * expr + | AMap of atom * Ident.t * expr + | ASum of atom + | ACount of atom + +type program = { + ap_input : Ident.t; + ap_input_element : Types.t; + ap_query : Ident.t; + ap_query_body : expr; + ap_helpers : (Ident.t * Types.scheme * expr) list; +} + +let make a aty aspan = { a = a; aty = aty; aspan = aspan } + +let temp_counter = ref 0 + +let reset () = temp_counter := 0 + +let fresh_temp span ty = + incr temp_counter; + let ident = Ident.fresh (Printf.sprintf "t%d" !temp_counter) span in + (ident, ty) + +let rec bind (typed : Typed.expr) (k : expr -> expr) : expr = + let span = typed.Typed.tspan in + let ty = typed.Typed.ty in + match typed.Typed.te with + | Typed.TInt value -> wrap span ty k (AInt value) + | Typed.TBool value -> wrap span ty k (ABool value) + | Typed.TString value -> wrap span ty k (AString value) + | Typed.TUnit -> wrap span ty k AUnit + | Typed.TVar ident -> wrap span ty k (AVar ident) + | Typed.TSource ident -> wrap span ty k (AVar ident) + | Typed.TField (record, label) -> + bind record (fun record_atom -> + force (make (AField (atom_of record_atom, label)) ty span) k) + | Typed.TBinop (operator, left, right) -> + bind left (fun left_atom -> + bind right (fun right_atom -> + force (make (ABinop (operator, atom_of left_atom, atom_of right_atom)) ty span) k)) + | Typed.TIf (condition, then_branch, else_branch) -> + bind condition (fun condition_atom -> + force + (make (AIf (atom_of condition_atom, normalize then_branch, normalize else_branch)) ty span) + k) + | Typed.TApp (fn, argument) -> + bind argument (fun argument_atom -> + match fn.Typed.te with + | Typed.TVar ident -> + force (make (AApp (ident, atom_of argument_atom)) ty span) k + | _ -> + Diagnostic.error span + "internal error: the function position is not a variable after specialization") + | Typed.TTuple items -> + bind_atoms items [] (fun atoms -> force (make (ATuple atoms) ty span) k) + | Typed.TRecord (name, fields) -> + bind_atoms (List.map snd fields) [] (fun atoms -> + force (make (ARecord (name, List.combine (List.map fst fields) atoms)) ty span) k) + | Typed.TLambda (ident, body) -> + force (make (ALambda (ident, normalize body)) ty span) k + | Typed.TFilter (collection, predicate) -> + bind collection (fun collection_atom -> + let parameter, body = lambda_body predicate in + force (make (AFilter (atom_of collection_atom, parameter, body)) ty span) k) + | Typed.TMap (collection, projection) -> + bind collection (fun collection_atom -> + let parameter, body = lambda_body projection in + force (make (AMap (atom_of collection_atom, parameter, body)) ty span) k) + | Typed.TSum collection -> + bind collection (fun collection_atom -> force (make (ASum (atom_of collection_atom)) ty span) k) + | Typed.TCount collection -> + bind collection (fun collection_atom -> force (make (ACount (atom_of collection_atom)) ty span) k) + | Typed.TLet (ident, bound, body) -> + k (make (ALet (ident, normalize bound, normalize body)) ty span) + +and atom_of expr = + match expr.a with + | AAtom atom -> atom + | _ -> Diagnostic.error expr.aspan "internal error: expected an atom in normalized form" + +and wrap span ty k atom = k (make (AAtom atom) ty span) + +and force value k = + match value.a with + | AAtom _ -> k value + | _ -> + let ident, _ = fresh_temp value.aspan value.aty in + let body = k (make (AAtom (AVar ident)) value.aty value.aspan) in + make (ALet (ident, value, body)) body.aty body.aspan + +and bind_atoms items acc k = + match items with + | [] -> k (List.rev acc) + | item :: rest -> bind item (fun atom -> bind_atoms rest (atom_of atom :: acc) k) + +and normalize expr = bind expr (fun value -> value) + +and lambda_body expr = + match expr.Typed.te with + | Typed.TLambda (parameter, body) -> (parameter, normalize body) + | _ -> + Diagnostic.error expr.Typed.tspan "internal error: expected a function literal after specialization" + +let program typed = + { + ap_input = typed.Typed.tp_input; + ap_input_element = typed.Typed.tp_input_element; + ap_query = typed.Typed.tp_query; + ap_query_body = normalize typed.Typed.tp_query_body; + ap_helpers = + List.map + (fun helper -> (helper.Typed.th_ident, helper.Typed.th_scheme, normalize helper.Typed.th_body)) + typed.Typed.tp_helpers; + } + +let atom_to_string atom = + match atom with + | AInt value -> string_of_int value + | ABool true -> "true" + | ABool false -> "false" + | AString value -> Printf.sprintf "%S" value + | AUnit -> "()" + | AVar ident -> Ident.display ident + +let rec to_string expr = + match expr.a with + | AAtom atom -> atom_to_string atom + | ALet (ident, bound, body) -> + Printf.sprintf "let %s = %s in\n%s" (Ident.display ident) (to_string bound) (to_string body) + | ABinop (operator, left, right) -> + Printf.sprintf "%s %s %s" (atom_to_string left) (Syntax.binop_name operator) + (atom_to_string right) + | AIf (condition, then_branch, else_branch) -> + Printf.sprintf "if %s then (%s) else (%s)" (atom_to_string condition) (to_string then_branch) + (to_string else_branch) + | ATuple atoms -> "(" ^ Util.join ", " (List.map atom_to_string atoms) ^ ")" + | ARecord (name, fields) -> + name ^ " { " + ^ Util.join ", " (List.map (fun (label, atom) -> label ^ " = " ^ atom_to_string atom) fields) + ^ " }" + | AField (record, label) -> atom_to_string record ^ "." ^ label + | AApp (ident, argument) -> Ident.display ident ^ " " ^ atom_to_string argument + | ALambda (ident, body) -> "fun " ^ Ident.display ident ^ " -> " ^ to_string body + | AFilter (collection, parameter, body) -> + "filter (fun " ^ Ident.display parameter ^ " -> " ^ to_string body ^ ") " + ^ atom_to_string collection + | AMap (collection, parameter, body) -> + "map (fun " ^ Ident.display parameter ^ " -> " ^ to_string body ^ ") " ^ atom_to_string collection + | ASum collection -> "sum " ^ atom_to_string collection + | ACount collection -> "count " ^ atom_to_string collection + +let program_to_string program = + let lines = + Printf.sprintf "input %s : collection %s" (Ident.display program.ap_input) + (Types.pp program.ap_input_element) + :: List.map + (fun (ident, scheme, body) -> + Printf.sprintf "let %s : %s = %s" (Ident.display ident) + (Types.pp scheme.Types.body) (to_string body)) + program.ap_helpers + @ [ + Printf.sprintf "query %s : %s =\n%s" (Ident.display program.ap_query) + (Types.pp program.ap_query_body.aty) (to_string program.ap_query_body); + ] + in + String.concat "\n" lines ^ "\n" + +let validate globals expr = + let binders = Hashtbl.create 64 in + let add ident = + let key = Ident.stamp ident in + if Hashtbl.mem binders key then + Diagnostic.error Location.none "anf invariant violated: binder %s is bound twice" + (Ident.to_string ident) + else Hashtbl.add binders key () + in + let rec walk expr = + match expr.a with + | AAtom _ -> () + | ALet (ident, bound, body) -> + add ident; + walk bound; + walk body + | ABinop (_, _, _) -> () + | AIf (_, then_branch, else_branch) -> + walk then_branch; + walk else_branch + | ATuple _ | ARecord _ | AField _ | AApp _ -> () + | ALambda (ident, body) -> + add ident; + walk body + | AFilter (_, parameter, body) | AMap (_, parameter, body) -> + add parameter; + walk body + | ASum _ | ACount _ -> () + in + walk expr; + + let bound = Hashtbl.create 64 in + List.iter (fun ident -> Hashtbl.replace bound (Ident.stamp ident) ()) globals; + let bind_ident ident = Hashtbl.replace bound (Ident.stamp ident) () in + let check ident = + if not (Hashtbl.mem bound (Ident.stamp ident)) then + Diagnostic.error Location.none "anf invariant violated: %s is used before it is bound" + (Ident.to_string ident) + in + let rec walk_uses expr = + match expr.a with + | AAtom (AVar ident) -> check ident + | AAtom _ -> () + | ALet (ident, bound_expr, body) -> + walk_uses bound_expr; + bind_ident ident; + walk_uses body + | ABinop (_, left, right) -> + check_atom left; + check_atom right + | AIf (condition, then_branch, else_branch) -> + check_atom condition; + walk_uses then_branch; + walk_uses else_branch + | ATuple atoms -> List.iter check_atom atoms + | ARecord (_, fields) -> List.iter (fun (_, atom) -> check_atom atom) fields + | AField (record, _) -> check_atom record + | AApp (ident, argument) -> + check ident; + check_atom argument + | ALambda (ident, body) -> + bind_ident ident; + walk_uses body + | AFilter (collection, parameter, body) | AMap (collection, parameter, body) -> + check_atom collection; + bind_ident parameter; + walk_uses body + | ASum collection | ACount collection -> check_atom collection + and check_atom = function + | AVar ident -> check ident + | AInt _ | ABool _ | AString _ | AUnit -> () + in + walk_uses expr diff --git a/src/anf.mli b/src/anf.mli new file mode 100644 index 0000000..106e65b --- /dev/null +++ b/src/anf.mli @@ -0,0 +1,42 @@ +type atom = + | AInt of int + | ABool of bool + | AString of string + | AUnit + | AVar of Ident.t + +type expr = { + a : desc; + aty : Types.t; + aspan : Location.span; +} + +and desc = + | AAtom of atom + | ALet of Ident.t * expr * expr + | ABinop of Syntax.binop * atom * atom + | AIf of atom * expr * expr + | ATuple of atom list + | ARecord of string * (string * atom) list + | AField of atom * string + | AApp of Ident.t * atom + | ALambda of Ident.t * expr + | AFilter of atom * Ident.t * expr + | AMap of atom * Ident.t * expr + | ASum of atom + | ACount of atom + +type program = { + ap_input : Ident.t; + ap_input_element : Types.t; + ap_query : Ident.t; + ap_query_body : expr; + ap_helpers : (Ident.t * Types.scheme * expr) list; +} + +val reset : unit -> unit +val program : Typed.program -> program +val to_string : expr -> string +val atom_to_string : atom -> string +val program_to_string : program -> string +val validate : Ident.t list -> expr -> unit diff --git a/src/main.ml b/src/main.ml index aef6fb3..1c45838 100644 --- a/src/main.ml +++ b/src/main.ml @@ -93,25 +93,30 @@ let infer_source path = Infer.program (resolve_source path) let specialize_source path = Specialize.program (infer_source path) +let anf_source path = Anf.program (specialize_source path) + let frontend_unavailable () = Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision" let check path = - ignore (specialize_source path) + ignore (anf_source path) let dump stage path = + match stage with + | "anf" -> print_string (Anf.program_to_string (anf_source path)) + | _ -> let program = infer_source path in match stage with | "typed" -> print_string (Typed.program_to_string program) | _ -> frontend_unavailable () let emit path output = - ignore (specialize_source path); + ignore (anf_source path); ignore output; frontend_unavailable () let build path output = - ignore (specialize_source path); + ignore (anf_source path); ignore output; frontend_unavailable () diff --git a/test/test_main.ml b/test/test_main.ml index 68252b6..21e40a7 100644 --- a/test/test_main.ml +++ b/test/test_main.ml @@ -375,6 +375,7 @@ let () = Test_harness.run_suite "resolve" resolve_cases; Test_harness.run_suite "types" Test_type.cases; Test_harness.run_suite "specialize" Test_type.specialize_cases; + Test_harness.run_suite "anf" Test_type.anf_cases; Test_harness.run_suite "interpret" Test_incremental.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) diff --git a/test/test_type.ml b/test/test_type.ml index 91177cd..835a7e1 100644 --- a/test/test_type.ml +++ b/test/test_type.ml @@ -324,3 +324,94 @@ let specialize_cases = in check "no let bindings survive" (not (has_let program.Typed.tp_query_body)) ); ] + +let anf text = Anf.program (specialize text) + +let anf_cases = + [ + ( "the example query normalizes to a chain of collection temporaries", + fun () -> + Ident.reset (); + Anf.reset (); + let program = anf (read_fixture "expensive_order.delta") in + Anf.validate [ program.Anf.ap_input ] program.Anf.ap_query_body; + (match program.Anf.ap_query_body with + | { Anf.a = Anf.ALet (filter_temp, filter, { Anf.a = Anf.ALet (map_temp, mapper, tail); _ }); _ } + -> + check "the first temporary is a filter" + (match filter.Anf.a with Anf.AFilter _ -> true | _ -> false); + check "the second temporary is a map" + (match mapper.Anf.a with Anf.AMap _ -> true | _ -> false); + check "the query returns the map temporary" + (match tail.Anf.a with + | Anf.AAtom (Anf.AVar ident) -> Ident.equal ident map_temp + | _ -> false); + ignore filter_temp + | _ -> fail "shape" "expected two collection temporaries") ); + ( "every operand of a sum is an atom", + fun () -> + Anf.reset (); + let program = anf (read_fixture "revenue.delta") in + Anf.validate [ program.Anf.ap_input ] program.Anf.ap_query_body; + (match program.Anf.ap_query_body with + | { Anf.a = Anf.ALet (_, mapped, { Anf.a = Anf.ALet (_, summed, _); _ }); _ } -> + check "map temporary" (match mapped.Anf.a with Anf.AMap _ -> true | _ -> false); + check "sum uses an atom operand" (match summed.Anf.a with Anf.ASum _ -> true | _ -> false) + | _ -> fail "shape" "expected a map temporary then a sum") ); + ( "arithmetic arguments are bound to temporaries", + fun () -> + Anf.reset (); + let program = anf "input rows : collection int\nquery q = rows |> map (fun r -> r * 2 + r) |> sum\n" in + Anf.validate [ program.Anf.ap_input ] program.Anf.ap_query_body; + let dump = Anf.program_to_string program in + check "the dump binds a multiplication before the addition" + (Util.starts_with "input rows : collection int" dump); + check "the dump mentions the map operation" (String.length dump > 0) ); + ( "anf dumps are deterministic", + fun () -> + Ident.reset (); + Types.reset (); + Anf.reset (); + let first = Anf.program_to_string (anf (read_fixture "count_large.delta")) in + Ident.reset (); + Types.reset (); + Anf.reset (); + let second = Anf.program_to_string (anf (read_fixture "count_large.delta")) in + check_equal_string "identical" first second ); + ( "anf validation rejects duplicated binders", + fun () -> + let ident = Ident.fresh "x" Location.none in + let body = + Anf.{ + a = ALet (ident, { a = AAtom (AInt 1); aty = Types.TInt; aspan = Location.none }, + { a = ALet (ident, { a = AAtom (AInt 2); aty = Types.TInt; aspan = Location.none }, + { a = AAtom (AInt 3); aty = Types.TInt; aspan = Location.none }); + aty = Types.TInt; aspan = Location.none }); + aty = Types.TInt; + aspan = Location.none; + } + in + (try + Anf.validate [] body; + fail "validate" "expected the invariant check to fail" + with Diagnostic.Error diagnostic -> + check "mentions a duplicated binder" + (Util.starts_with "anf invariant violated" diagnostic.Diagnostic.message)) ); + ( "anf validation rejects unbound uses", + fun () -> + let unbound = Ident.fresh "missing" Location.none in + let body = + Anf. + { + a = AAtom (AVar unbound); + aty = Types.TInt; + aspan = Location.none; + } + in + (try + Anf.validate [] body; + fail "validate" "expected the invariant check to fail" + with Diagnostic.Error diagnostic -> + check "mentions an unbound variable" + (Util.starts_with "anf invariant violated" diagnostic.Diagnostic.message)) ); + ]