Infer scalar types with let generalization

This commit is contained in:
sneeker committed 2017-02-02 14:36:00 +00:00
1 parent 270371c3ed
commit 999b9e3262
15 files changed
+980 -13

No files matched your search

+377
View File
@@ -0,0 +1,377 @@
type env = {
ctx : Types.records;
values : (int * Ident.t * Types.scheme) list;
input : (Ident.t * Types.t) option;
}
let empty_env ctx = { ctx = ctx; values = []; input = None }
let bind env ident scheme = { env with values = (Ident.stamp ident, ident, scheme) :: env.values }
let scheme_of ty = { Types.vars = []; body = ty }
let lookup env ident =
let rec search = function
| [] -> None
| (stamp, _, scheme) :: rest -> if stamp = Ident.stamp ident then Some scheme else search rest
in
search env.values
let env_free_vars env =
List.fold_left
(fun acc (_, _, scheme) ->
let body_vars = Types.free_vars_of scheme.Types.body in
List.fold_left
(fun acc var ->
if List.exists (fun other -> other.Types.id = var.Types.id) scheme.Types.vars then acc
else var :: acc)
acc body_vars)
[] env.values
let generalize env ty =
let free = Types.free_vars_of ty in
let env_free = env_free_vars env in
let quantified =
List.filter
(fun var -> not (List.exists (fun other -> other.Types.id = var.Types.id) env_free))
free
in
{ Types.vars = quantified; body = ty }
let instantiate scheme =
let table = Hashtbl.create 16 in
List.iter
(fun var -> Hashtbl.replace table var.Types.id (Types.fresh_var ()))
scheme.Types.vars;
let rec copy ty =
match Types.repr ty with
| Types.TVar var -> (
match Util.hashtbl_find_opt table var.Types.id with
| Some replacement -> replacement
| None -> ty)
| Types.TTuple items -> Types.TTuple (List.map copy items)
| Types.TCollection element -> Types.TCollection (copy element)
| Types.TArrow (domain, codomain) -> Types.TArrow (copy domain, copy codomain)
| other -> other
in
copy scheme.Types.body
let named_type name =
match name with
| "int" -> Some Types.TInt
| "bool" -> Some Types.TBool
| "string" -> Some Types.TString
| "unit" -> Some Types.TUnit
| _ -> None
let rec of_tyexpr names tyexpr =
match tyexpr.Syntax.ty with
| Syntax.TyName name -> (
match named_type name with
| Some ty -> ty
| None ->
if Util.contains name names then Types.TRecord name
else Diagnostic.error tyexpr.Syntax.tyspan "unknown type `%s`" name)
| Syntax.TyTuple items -> Types.TTuple (List.map (of_tyexpr names) items)
| Syntax.TyCollection _ ->
Diagnostic.error tyexpr.Syntax.tyspan
"collection types are only allowed as the input type and as query results"
let rec contains_collection ty =
match Types.repr ty with
| Types.TCollection _ -> true
| Types.TTuple items -> List.exists contains_collection items
| _ -> false
let rec contains_function ty =
match Types.repr ty with
| Types.TArrow _ -> true
| Types.TTuple items -> List.exists contains_function items
| Types.TCollection element -> contains_function element
| _ -> false
let check_element span label ty =
if contains_collection ty then
Diagnostic.error span "nested collections are not supported: %s" label;
if contains_function ty then
Diagnostic.error span "collections of functions are not supported: %s" label
let rec is_collection_tyexpr tyexpr =
match tyexpr.Syntax.ty with
| Syntax.TyCollection _ -> true
| Syntax.TyTuple items -> List.exists is_collection_tyexpr items
| Syntax.TyName _ -> false
let records_of_decls decls =
let names = List.map (fun decl -> decl.Syntax.rd_name) decls in
let declarations =
List.map
(fun decl ->
let fields =
List.map
(fun (label, tyexpr) ->
if is_collection_tyexpr tyexpr then
Diagnostic.error tyexpr.Syntax.tyspan
"record field `%s` has a collection type; collections may not be stored in records" label;
let ty = of_tyexpr names tyexpr in
if contains_function ty then
Diagnostic.error tyexpr.Syntax.tyspan "record field `%s` has a function type" label;
(label, ty))
decl.Syntax.rd_fields
in
{ Types.ri_name = decl.Syntax.rd_name; ri_fields = fields })
decls
in
let labels =
Util.concat_map
(fun info -> List.map (fun (label, ty) -> (label, (info.Types.ri_name, ty))) info.Types.ri_fields)
declarations
in
{ Types.declarations = declarations; labels = labels }
let internal span fmt =
Printf.ksprintf (fun message -> Diagnostic.error span "internal error: %s" message) fmt
let rec infer env syntax =
let span = syntax.Resolve.rspan in
match syntax.Resolve.r with
| Resolve.RInt value -> Typed.make (Typed.TInt value) Types.TInt span
| Resolve.RBool value -> Typed.make (Typed.TBool value) Types.TBool span
| Resolve.RString value -> Typed.make (Typed.TString value) Types.TString span
| Resolve.RUnit -> Typed.make Typed.TUnit Types.TUnit span
| Resolve.RVar ident -> (
match env.input with
| Some (input, ty) when Ident.equal input ident -> Typed.make (Typed.TSource ident) ty span
| _ -> (
match lookup env ident with
| Some scheme -> Typed.make (Typed.TVar ident) (instantiate scheme) span
| None -> internal span "unresolved variable `%s`" (Ident.display ident)))
| Resolve.RLambda (ident, body) ->
let argument = Types.fresh_var () in
let body = infer (bind env ident (scheme_of argument)) body in
Typed.make (Typed.TLambda (ident, body)) (Types.TArrow (argument, body.Typed.ty)) span
| Resolve.RLet (ident, bound, body) ->
let bound = infer env bound in
let scheme = generalize env bound.Typed.ty in
let body = infer (bind env ident scheme) body in
Typed.make (Typed.TLet (ident, bound, body)) body.Typed.ty span
| Resolve.RApp ({ Resolve.r = Resolve.RFilter predicate; _ }, collection) ->
infer_filter env span predicate collection
| Resolve.RApp ({ Resolve.r = Resolve.RMap projection; _ }, collection) ->
infer_map env span projection collection
| Resolve.RApp ({ Resolve.r = Resolve.RSum; _ }, collection) ->
let operand = infer env collection in
let where = collection.Resolve.rspan in
if not (Types.is_collection operand.Typed.ty) then
Diagnostic.error where "sum expects a collection of integers but got %s" (Types.pp operand.Typed.ty);
let element = Types.fresh_var () in
Types.unify where operand.Typed.ty (Types.TCollection element);
(match Types.repr element with
| Types.TInt | Types.TVar _ -> ()
| other ->
Diagnostic.error where "sum expects a collection of integers but got collection %s"
(Types.pp other));
Types.unify where element Types.TInt;
Typed.make (Typed.TSum operand) Types.TInt span
| Resolve.RApp ({ Resolve.r = Resolve.RCount; _ }, collection) ->
let operand = infer env collection in
let where = collection.Resolve.rspan in
if not (Types.is_collection operand.Typed.ty) then
Diagnostic.error where "count expects a collection but got %s" (Types.pp operand.Typed.ty);
let element = Types.fresh_var () in
Types.unify where operand.Typed.ty (Types.TCollection element);
check_element where "count element" element;
Typed.make (Typed.TCount operand) Types.TInt span
| Resolve.RApp (fn, argument) ->
let fn = infer env fn in
let argument = infer env argument in
let result = Types.fresh_var () in
Types.unify span fn.Typed.ty (Types.TArrow (argument.Typed.ty, result));
Typed.make (Typed.TApp (fn, argument)) result span
| Resolve.RFilter _ | Resolve.RMap _ ->
Diagnostic.error span "`filter` and `map` must be applied to a collection"
| Resolve.RSum -> Diagnostic.error span "`sum` must be applied to a collection of integers"
| Resolve.RCount -> Diagnostic.error span "`count` must be applied to a collection"
| Resolve.RIf (condition, then_branch, else_branch) ->
let condition = infer env condition in
Types.unify condition.Typed.tspan condition.Typed.ty Types.TBool;
let then_branch = infer env then_branch in
let else_branch = infer env else_branch in
Types.unify span then_branch.Typed.ty else_branch.Typed.ty;
Typed.make (Typed.TIf (condition, then_branch, else_branch)) then_branch.Typed.ty span
| Resolve.RBinop (operator, left, right) ->
let left = infer env left in
let right = infer env right in
(match operator with
| Syntax.Add | Syntax.Sub | Syntax.Mul | Syntax.Div ->
Types.unify span left.Typed.ty Types.TInt;
Types.unify span right.Typed.ty Types.TInt;
Typed.make (Typed.TBinop (operator, left, right)) Types.TInt span
| Syntax.Lt | Syntax.Le | Syntax.Gt | Syntax.Ge ->
Types.unify span left.Typed.ty Types.TInt;
Types.unify span right.Typed.ty Types.TInt;
Typed.make (Typed.TBinop (operator, left, right)) Types.TBool span
| Syntax.And | Syntax.Or ->
Types.unify span left.Typed.ty Types.TBool;
Types.unify span right.Typed.ty Types.TBool;
Typed.make (Typed.TBinop (operator, left, right)) Types.TBool span
| Syntax.Eq | Syntax.Ne ->
Types.unify span left.Typed.ty right.Typed.ty;
if not (Types.is_equatable left.Typed.ty) then
Diagnostic.error span "`%s` is only supported on integers, booleans, strings, unit, tuples and records of these, but the operands have type %s"
(Syntax.binop_name operator) (Types.pp left.Typed.ty);
Typed.make (Typed.TBinop (operator, left, right)) Types.TBool span)
| Resolve.RNeg operand ->
let operand = infer env operand in
Types.unify span operand.Typed.ty Types.TInt;
Typed.make (Typed.TBinop (Syntax.Sub, Typed.make (Typed.TInt 0) Types.TInt span, operand)) Types.TInt span
| Resolve.RTuple items ->
let items = List.map (infer env) items in
Typed.make (Typed.TTuple items) (Types.TTuple (List.map (fun item -> item.Typed.ty) items)) span
| Resolve.RRecord fields -> infer_record env span fields
| Resolve.RField (record, label) ->
let record = infer env record in
(match Util.assoc_opt label env.ctx.Types.labels with
| None -> Diagnostic.error span "no record type declares a field named `%s`" label
| Some (name, field_ty) ->
Types.unify span record.Typed.ty (Types.TRecord name);
Typed.make (Typed.TField (record, label)) field_ty span)
and infer_filter env span predicate collection =
let collection = infer env collection in
let element = Types.fresh_var () in
Types.unify collection.Typed.tspan collection.Typed.ty (Types.TCollection element);
check_element collection.Typed.tspan "filter element" element;
let predicate = infer env predicate in
Types.unify predicate.Typed.tspan predicate.Typed.ty (Types.TArrow (element, Types.TBool));
Typed.make (Typed.TFilter (collection, predicate)) collection.Typed.ty span
and infer_map env span projection collection =
let collection = infer env collection in
let element = Types.fresh_var () in
Types.unify collection.Typed.tspan collection.Typed.ty (Types.TCollection element);
check_element collection.Typed.tspan "map element" element;
let projection = infer env projection in
let result = Types.fresh_var () in
Types.unify projection.Typed.tspan projection.Typed.ty (Types.TArrow (element, result));
check_element projection.Typed.tspan "mapped element" result;
Typed.make (Typed.TMap (collection, projection)) (Types.TCollection result) span
and infer_record env span fields =
let inferred = List.map (fun (label, value) -> (label, infer env value)) fields in
match inferred with
| [] -> Diagnostic.error span "a record literal must have at least one field"
| (first_label, _) :: _ ->
let name =
match Util.assoc_opt first_label env.ctx.Types.labels with
| Some (name, _) -> name
| None -> Diagnostic.error span "no record type declares a field named `%s`" first_label
in
let info =
match Types.record_info env.ctx name with
| Some info -> info
| None -> internal span "unknown record %s" name
in
let provided =
List.map
(fun (label, value) ->
match Util.assoc_opt label info.Types.ri_fields with
| None ->
Diagnostic.error value.Typed.tspan "record `%s` has no field named `%s`" name label
| Some expected ->
Types.unify value.Typed.tspan value.Typed.ty expected;
(label, value))
inferred
in
List.iter
(fun (label, _) ->
if not (List.exists (fun (provided_label, _) -> provided_label = label) provided) then
Diagnostic.error span "record literal of type `%s` is missing field `%s`" name label)
info.Types.ri_fields;
let ordered =
List.map
(fun (label, _) ->
let value = List.assoc label provided in
(label, value))
info.Types.ri_fields
in
Typed.make (Typed.TRecord (name, ordered)) (Types.TRecord name) span
let helper_order helpers =
let table = Hashtbl.create 16 in
List.iter (fun helper -> Hashtbl.replace table (Ident.stamp helper.Resolve.h_ident) helper) helpers;
let visited = Hashtbl.create 16 in
let order = ref [] in
let rec visit helper =
let key = Ident.stamp helper.Resolve.h_ident in
match Util.hashtbl_find_opt visited key with
| Some () -> ()
| None ->
Hashtbl.replace visited key ();
List.iter
(fun ident ->
match Util.hashtbl_find_opt table (Ident.stamp ident) with
| Some callee -> visit callee
| None -> ())
(Resolve.refs_of_expr helper.Resolve.h_body);
order := helper :: !order
in
List.iter visit helpers;
List.rev !order
let program resolved =
let ctx = records_of_decls resolved.Resolve.rp_records in
let element =
match resolved.Resolve.rp_input_ty.Syntax.ty with
| Syntax.TyCollection element ->
let ty = of_tyexpr (List.map (fun decl -> decl.Syntax.rd_name) resolved.Resolve.rp_records) element in
check_element resolved.Resolve.rp_input_ty.Syntax.tyspan "input element" ty;
ty
| _ ->
Diagnostic.error resolved.Resolve.rp_input_ty.Syntax.tyspan
"the input collection must be declared with a collection type: input NAME : collection T"
in
let base = empty_env ctx in
let helpers = ref [] in
let env = ref base in
List.iter
(fun helper ->
let helper_env = { !env with input = None } in
let params, body_env =
List.fold_left
(fun (params, env) parameter ->
let ty = Types.fresh_var () in
(params @ [ (parameter, ty) ], bind env parameter (scheme_of ty)))
([], helper_env) helper.Resolve.h_params
in
let body = infer body_env helper.Resolve.h_body in
let body =
List.fold_right
(fun (parameter, ty) body ->
Typed.make (Typed.TLambda (parameter, body)) (Types.TArrow (ty, body.Typed.ty))
helper.Resolve.h_span)
params body
in
let scheme = generalize base body.Typed.ty in
helpers := { Typed.th_ident = helper.Resolve.h_ident; th_scheme = scheme; th_body = body } :: !helpers;
env := bind !env helper.Resolve.h_ident scheme)
(helper_order resolved.Resolve.rp_helpers);
let query_env =
{ !env with input = Some (resolved.Resolve.rp_input, Types.TCollection element) }
in
let query_body = infer query_env resolved.Resolve.rp_query_body in
(match Types.repr query_body.Typed.ty with
| Types.TInt -> ()
| Types.TCollection element ->
check_element query_body.Typed.tspan "query result" element
| other ->
Diagnostic.error query_body.Typed.tspan
"a query must produce a collection or an integer, but this query produces %s" (Types.pp other));
{
Typed.tp_records = ctx;
tp_input = resolved.Resolve.rp_input;
tp_input_element = element;
tp_helpers = List.rev !helpers;
tp_query = resolved.Resolve.rp_query;
tp_query_body = query_body;
}