Normalize query expressions into ANF

This commit is contained in:
sneeker committed 2017-03-03 15:09:00 +00:00
1 parent 05769abc11
commit a8ef268c9c
5 files changed
+406 -3

No files matched your search

+264
View File
@@ -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
+42
View File
@@ -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
+8 -3
View File
@@ -93,25 +93,30 @@ let infer_source path = Infer.program (resolve_source path)
let specialize_source path = Specialize.program (infer_source path) let specialize_source path = Specialize.program (infer_source path)
let anf_source path = Anf.program (specialize_source path)
let frontend_unavailable () = let frontend_unavailable () =
Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision" Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision"
let check path = let check path =
ignore (specialize_source path) ignore (anf_source path)
let dump stage path = let dump stage path =
match stage with
| "anf" -> print_string (Anf.program_to_string (anf_source path))
| _ ->
let program = infer_source path in let program = infer_source path in
match stage with match stage with
| "typed" -> print_string (Typed.program_to_string program) | "typed" -> print_string (Typed.program_to_string program)
| _ -> frontend_unavailable () | _ -> frontend_unavailable ()
let emit path output = let emit path output =
ignore (specialize_source path); ignore (anf_source path);
ignore output; ignore output;
frontend_unavailable () frontend_unavailable ()
let build path output = let build path output =
ignore (specialize_source path); ignore (anf_source path);
ignore output; ignore output;
frontend_unavailable () frontend_unavailable ()
+1
View File
@@ -375,6 +375,7 @@ let () =
Test_harness.run_suite "resolve" resolve_cases; Test_harness.run_suite "resolve" resolve_cases;
Test_harness.run_suite "types" Test_type.cases; Test_harness.run_suite "types" Test_type.cases;
Test_harness.run_suite "specialize" Test_type.specialize_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; Test_harness.run_suite "interpret" Test_incremental.cases;
Printf.printf "%d cases, %d failures\n" (Test_harness.case_count ()) (Test_harness.failure_count ()); Printf.printf "%d cases, %d failures\n" (Test_harness.case_count ()) (Test_harness.failure_count ());
exit (if Test_harness.failure_count () = 0 then 0 else 1) exit (if Test_harness.failure_count () = 0 then 0 else 1)
+91
View File
@@ -324,3 +324,94 @@ let specialize_cases =
in in
check "no let bindings survive" (not (has_let program.Typed.tp_query_body)) ); 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)) );
]