From 999b9e326258ec275493b2e69b57e1521b300317 Mon Sep 17 00:00:00 2001 From: sneeker Date: Thu, 2 Feb 2017 14:36:00 +0000 Subject: [PATCH] Infer scalar types with let generalization --- Makefile | 3 +- src/infer.ml | 377 ++++++++++++++++++++++++++++++++++++++++++ src/infer.mli | 2 + src/main.ml | 15 +- src/resolve.mli | 1 + src/typed.ml | 114 +++++++++++++ src/typed.mli | 44 +++++ src/types.ml | 128 ++++++++++++++ src/types.mli | 42 +++++ src/util.ml | 48 ++++++ src/util.mli | 11 ++ test/test_harness.ml | 29 ++++ test/test_harness.mli | 12 ++ test/test_main.ml | 7 +- test/test_type.ml | 160 ++++++++++++++++++ 15 files changed, 980 insertions(+), 13 deletions(-) create mode 100644 src/infer.ml create mode 100644 src/infer.mli create mode 100644 src/typed.ml create mode 100644 src/typed.mli create mode 100644 src/types.ml create mode 100644 src/types.mli create mode 100644 src/util.ml create mode 100644 src/util.mli create mode 100644 test/test_type.ml diff --git a/Makefile b/Makefile index d38c75e..3545c7d 100644 --- a/Makefile +++ b/Makefile @@ -125,6 +125,7 @@ $(BUILD)/deps.d: $(MAKEFILE_LIST) $(SRC_ML) $(INTERFACES) $(RT_ML) $(TEST_ML) $( $(BUILD)/order.mk: $(MAKEFILE_LIST) $(SRC_ML) $(RT_ML) $(TEST_ML) $(GEN) | $(BUILD) { echo -n "LIB_ORDER = " ; $(OCAMLDEP) -sort $(LIB_ML) $(LIB_GEN_ML) | sed -e 's|$(SRC_DIR)/\([^ ]*\)\.ml|$(BUILD)/\1.cmx|g' -e 's|$(BUILD)/\([^ ]*\)\.ml|$(BUILD)/\1.cmx|g' ; echo ; } > $@ { echo -n "RT_ORDER = " ; $(OCAMLDEP) -sort $(RT_ML) | sed -e 's|$(RUNTIME_DIR)/\([^ ]*\)\.ml|$(BUILD)/\1.cmx|g' ; echo ; } >> $@ + { echo -n "TEST_ORDER = " ; $(OCAMLDEP) -sort $(TEST_ML) | sed -e 's|$(TEST_DIR)/\([^ ]*\)\.ml|$(BUILD)/\1.cmx|g' ; echo ; } >> $@ -include $(BUILD)/order.mk @@ -145,7 +146,7 @@ ifeq ($(strip $(TEST_ML)),) TEST_EXE = else $(TEST_EXE): $(TEST_CMX) $(LIB_ARCHIVE) $(RUNTIME_ARCHIVE) - $(OCAMLOPT) $(FLAGS) -o $@ $(LIBS) $(LIB_ARCHIVE) $(RUNTIME_ARCHIVE) $(TEST_CMX) + $(OCAMLOPT) $(FLAGS) -o $@ $(LIBS) $(LIB_ARCHIVE) $(RUNTIME_ARCHIVE) $(TEST_ORDER) endif test: $(TEST_EXE) diff --git a/src/infer.ml b/src/infer.ml new file mode 100644 index 0000000..4404a18 --- /dev/null +++ b/src/infer.ml @@ -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; + } diff --git a/src/infer.mli b/src/infer.mli new file mode 100644 index 0000000..9b464cc --- /dev/null +++ b/src/infer.mli @@ -0,0 +1,2 @@ +val program : Resolve.program -> Typed.program +val of_tyexpr : string list -> Syntax.tyexpr -> Types.t diff --git a/src/main.ml b/src/main.ml index a18f7b1..f6a00d5 100644 --- a/src/main.ml +++ b/src/main.ml @@ -89,24 +89,27 @@ let parse_source path = Parse.program (load_source path) let resolve_source path = Resolve.program (parse_source path) +let infer_source path = Infer.program (resolve_source path) + let frontend_unavailable () = Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision" let check path = - ignore (resolve_source path) + ignore (infer_source path) let dump stage path = - ignore (stage); - ignore (resolve_source path); - frontend_unavailable () + let program = infer_source path in + match stage with + | "typed" -> print_string (Typed.program_to_string program) + | _ -> frontend_unavailable () let emit path output = - ignore (resolve_source path); + ignore (infer_source path); ignore output; frontend_unavailable () let build path output = - ignore (resolve_source path); + ignore (infer_source path); ignore output; frontend_unavailable () diff --git a/src/resolve.mli b/src/resolve.mli index 8cf0059..33f5c69 100644 --- a/src/resolve.mli +++ b/src/resolve.mli @@ -40,3 +40,4 @@ type program = { } val program : Syntax.program -> program +val refs_of_expr : expr -> Ident.t list diff --git a/src/typed.ml b/src/typed.ml new file mode 100644 index 0000000..7730b98 --- /dev/null +++ b/src/typed.ml @@ -0,0 +1,114 @@ +type expr = { + te : desc; + ty : Types.t; + tspan : Location.span; +} + +and desc = + | TInt of int + | TBool of bool + | TString of string + | TUnit + | TVar of Ident.t + | TLet of Ident.t * expr * expr + | TLambda of Ident.t * expr + | TApp of expr * expr + | TIf of expr * expr * expr + | TBinop of Syntax.binop * expr * expr + | TTuple of expr list + | TRecord of string * (string * expr) list + | TField of expr * string + | TSource of Ident.t + | TFilter of expr * expr + | TMap of expr * expr + | TSum of expr + | TCount of expr + +type helper = { + th_ident : Ident.t; + th_scheme : Types.scheme; + th_body : expr; +} + +type program = { + tp_records : Types.records; + tp_input : Ident.t; + tp_input_element : Types.t; + tp_helpers : helper list; + tp_query : Ident.t; + tp_query_body : expr; +} + +let make te ty tspan = { te = te; ty = ty; tspan = tspan } + +let precedence expr = + match expr.te with + | TInt _ | TBool _ | TString _ | TUnit | TVar _ | TSource _ | TRecord _ | TField _ | TTuple _ -> 6 + | TApp _ -> 5 + | TBinop ((Syntax.Mul | Syntax.Div), _, _) -> 4 + | TBinop ((Syntax.Add | Syntax.Sub), _, _) -> 3 + | TBinop ((Syntax.Eq | Syntax.Ne | Syntax.Lt | Syntax.Le | Syntax.Gt | Syntax.Ge), _, _) -> 2 + | TBinop (Syntax.And, _, _) -> 1 + | TBinop (Syntax.Or, _, _) -> 0 + | TIf _ | TLet _ | TLambda _ | TFilter _ | TMap _ | TSum _ | TCount _ -> 0 + +let rec render parent expr = + let body = + match expr.te with + | TInt value -> string_of_int value + | TBool true -> "true" + | TBool false -> "false" + | TString value -> Printf.sprintf "%S" value + | TUnit -> "()" + | TVar ident -> Ident.display ident + | TSource ident -> Ident.display ident + | TTuple items -> "(" ^ Util.join ", " (List.map (render 0) items) ^ ")" + | TRecord (name, fields) -> + name ^ " { " + ^ Util.join ", " (List.map (fun (label, value) -> label ^ " = " ^ render 0 value) fields) + ^ " }" + | TField (record, label) -> render 6 record ^ "." ^ label + | TApp (fn, argument) -> render 5 fn ^ " " ^ render 6 argument + | TIf (condition, then_branch, else_branch) -> + "if " ^ render 0 condition ^ " then " ^ render 0 then_branch ^ " else " ^ render 0 else_branch + | TLet (ident, bound, body) -> + "let " ^ Ident.display ident ^ " = " ^ render 0 bound ^ " in " ^ render 0 body + | TLambda (ident, body) -> "fun " ^ Ident.display ident ^ " -> " ^ render 0 body + | TBinop (operator, left, right) -> + render (precedence expr) left ^ " " ^ Syntax.binop_name operator ^ " " + ^ render (precedence expr + 1) right + | TFilter (collection, predicate) -> + "filter " ^ render 6 predicate ^ " " ^ render 6 collection + | TMap (collection, projection) -> "map " ^ render 6 projection ^ " " ^ render 6 collection + | TSum collection -> "sum " ^ render 6 collection + | TCount collection -> "count " ^ render 6 collection + in + if precedence expr < parent then "(" ^ body ^ ")" else body + +let to_string expr = render 0 expr + +let helper_to_string helper = + Printf.sprintf "let %s : %s = %s" (Ident.display helper.th_ident) + (Types.pp helper.th_scheme.Types.body) + (to_string helper.th_body) + +let program_to_string program = + let lines = + List.map + (fun info -> + Printf.sprintf "type %s = { %s }" info.Types.ri_name + (Util.join "; " + (List.map (fun (label, ty) -> label ^ " : " ^ Types.pp ty) info.Types.ri_fields))) + program.tp_records.Types.declarations + @ [ + Printf.sprintf "input %s : collection %s" (Ident.display program.tp_input) + (Types.pp program.tp_input_element); + ] + @ List.map helper_to_string program.tp_helpers + @ [ + Printf.sprintf "query %s : %s = %s" (Ident.display program.tp_query) + (Types.pp program.tp_query_body.ty) + (to_string program.tp_query_body); + ] + in + String.concat "\n" lines ^ "\n" diff --git a/src/typed.mli b/src/typed.mli new file mode 100644 index 0000000..4212688 --- /dev/null +++ b/src/typed.mli @@ -0,0 +1,44 @@ +type expr = { + te : desc; + ty : Types.t; + tspan : Location.span; +} + +and desc = + | TInt of int + | TBool of bool + | TString of string + | TUnit + | TVar of Ident.t + | TLet of Ident.t * expr * expr + | TLambda of Ident.t * expr + | TApp of expr * expr + | TIf of expr * expr * expr + | TBinop of Syntax.binop * expr * expr + | TTuple of expr list + | TRecord of string * (string * expr) list + | TField of expr * string + | TSource of Ident.t + | TFilter of expr * expr + | TMap of expr * expr + | TSum of expr + | TCount of expr + +type helper = { + th_ident : Ident.t; + th_scheme : Types.scheme; + th_body : expr; +} + +type program = { + tp_records : Types.records; + tp_input : Ident.t; + tp_input_element : Types.t; + tp_helpers : helper list; + tp_query : Ident.t; + tp_query_body : expr; +} + +val make : desc -> Types.t -> Location.span -> expr +val to_string : expr -> string +val program_to_string : program -> string diff --git a/src/types.ml b/src/types.ml new file mode 100644 index 0000000..b1d0ecf --- /dev/null +++ b/src/types.ml @@ -0,0 +1,128 @@ +type t = + | TInt + | TBool + | TString + | TUnit + | TTuple of t list + | TRecord of string + | TCollection of t + | TArrow of t * t + | TVar of var + +and var = { + id : int; + mutable link : t option; +} + +type scheme = { + vars : var list; + body : t; +} + +type record_info = { + ri_name : string; + ri_fields : (string * t) list; +} + +type records = { + declarations : record_info list; + labels : (string * (string * t)) list; +} + +let next_var = ref 0 + +let reset () = next_var := 0 + +let fresh_var () = + incr next_var; + TVar { id = !next_var; link = None } + +let rec repr ty = + match ty with + | TVar var -> ( + match var.link with + | Some linked -> + let result = repr linked in + var.link <- Some result; + result + | None -> ty) + | _ -> ty + +let rec occurs var ty = + match repr ty with + | TVar other -> other.id = var.id + | TTuple items -> List.exists (occurs var) items + | TCollection element -> occurs var element + | TArrow (domain, codomain) -> occurs var domain || occurs var codomain + | TInt | TBool | TString | TUnit | TRecord _ -> false + +let rec pp ty = + match repr ty with + | TInt -> "int" + | TBool -> "bool" + | TString -> "string" + | TUnit -> "unit" + | TTuple items -> "(" ^ Util.join ", " (List.map pp items) ^ ")" + | TRecord name -> name + | TCollection element -> "collection " ^ pp element + | TArrow (domain, codomain) -> "(" ^ pp domain ^ " -> " ^ pp codomain ^ ")" + | TVar var -> Printf.sprintf "'t%d" var.id + +let rec unify span left right = + let left = repr left in + let right = repr right in + if left == right then () + else + match (left, right) with + | TVar var, other | other, TVar var -> + if occurs var other then + Diagnostic.error span "cannot construct the infinite type %s = %s" (pp other) (pp left) + else var.link <- Some other + | TInt, TInt | TBool, TBool | TString, TString | TUnit, TUnit -> () + | TTuple left_items, TTuple right_items -> + if List.length left_items <> List.length right_items then + Diagnostic.error span "tuple size mismatch: %s has %d elements but %s has %d" (pp left) + (List.length left_items) (pp right) (List.length right_items) + else List.iter2 (unify span) left_items right_items + | TRecord left_name, TRecord right_name -> + if left_name <> right_name then + Diagnostic.error span "record type mismatch: expected %s but got %s" left_name right_name + | TCollection left_element, TCollection right_element -> unify span left_element right_element + | TArrow (left_domain, left_codomain), TArrow (right_domain, right_codomain) -> + unify span left_domain right_domain; + unify span left_codomain right_codomain + | _ -> Diagnostic.error span "type mismatch: expected %s but got %s" (pp left) (pp right) + +let rec free_vars ty acc = + match repr ty with + | TVar var -> if List.exists (fun other -> other.id = var.id) acc then acc else var :: acc + | TTuple items -> List.fold_left (fun acc item -> free_vars item acc) acc items + | TCollection element -> free_vars element acc + | TArrow (domain, codomain) -> free_vars codomain (free_vars domain acc) + | TInt | TBool | TString | TUnit | TRecord _ -> acc + +let free_vars_of ty = free_vars ty [] + +let is_function ty = match repr ty with TArrow _ -> true | _ -> false + +let is_collection ty = match repr ty with TCollection _ -> true | _ -> false + +let rec is_equatable ty = + match repr ty with + | TInt | TBool | TString | TUnit -> true + | TTuple items -> List.for_all is_equatable items + | TVar _ -> true + | TRecord name -> true + | TCollection _ | TArrow _ -> false + +let choose_variables ty = + match repr ty with + | TVar var -> Some var + | _ -> None + +let record_info records name = + let rec search = function + | [] -> None + | info :: rest -> if info.ri_name = name then Some info else search rest + in + search records.declarations diff --git a/src/types.mli b/src/types.mli new file mode 100644 index 0000000..b62a14b --- /dev/null +++ b/src/types.mli @@ -0,0 +1,42 @@ +type t = + | TInt + | TBool + | TString + | TUnit + | TTuple of t list + | TRecord of string + | TCollection of t + | TArrow of t * t + | TVar of var + +and var = { + id : int; + mutable link : t option; +} + +type scheme = { + vars : var list; + body : t; +} + +type record_info = { + ri_name : string; + ri_fields : (string * t) list; +} + +type records = { + declarations : record_info list; + labels : (string * (string * t)) list; +} + +val reset : unit -> unit +val fresh_var : unit -> t +val repr : t -> t +val occurs : var -> t -> bool +val unify : Location.span -> t -> t -> unit +val free_vars_of : t -> var list +val is_function : t -> bool +val is_collection : t -> bool +val is_equatable : t -> bool +val record_info : records -> string -> record_info option +val pp : t -> string diff --git a/src/util.ml b/src/util.ml new file mode 100644 index 0000000..491d59e --- /dev/null +++ b/src/util.ml @@ -0,0 +1,48 @@ +let assoc_opt key items = + let rec search = function + | [] -> None + | (candidate, value) :: rest -> if candidate = key then Some value else search rest + in + search items + +let hashtbl_find_opt table key = try Some (Hashtbl.find table key) with Not_found -> None + +let rec filter_map f = function + | [] -> [] + | item :: rest -> ( + match f item with + | Some value -> value :: filter_map f rest + | None -> filter_map f rest) + +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)) + +let rec take count items = + if count <= 0 then [] + else + match items with + | [] -> [] + | item :: rest -> item :: take (count - 1) rest + +let rec drop count items = + if count <= 0 then items + else match items with [] -> [] | _ :: rest -> drop (count - 1) rest + +let contains value items = List.exists (fun item -> item = value) items + +let string_before_char text character = + match String.index_opt text character with + | Some index -> String.sub text 0 index + | None -> text + +let starts_with prefix text = + String.length text >= String.length prefix && String.sub text 0 (String.length prefix) = prefix + +let rec join separator = function + | [] -> "" + | [ single ] -> single + | item :: rest -> item ^ separator ^ join separator rest diff --git a/src/util.mli b/src/util.mli new file mode 100644 index 0000000..eb19b5d --- /dev/null +++ b/src/util.mli @@ -0,0 +1,11 @@ +val assoc_opt : 'a -> ('a * 'b) list -> 'b option +val hashtbl_find_opt : ('a, 'b) Hashtbl.t -> 'a -> 'b option +val filter_map : ('a -> 'b option) -> 'a list -> 'b list +val concat_map : ('a -> 'b list) -> 'a list -> 'b list +val list_init : int -> (int -> 'a) -> 'a list +val take : int -> 'a list -> 'a list +val drop : int -> 'a list -> 'a list +val contains : 'a -> 'a list -> bool +val string_before_char : string -> char -> string +val starts_with : string -> string -> bool +val join : string -> string list -> string diff --git a/test/test_harness.ml b/test/test_harness.ml index 08f1ee9..1852774 100644 --- a/test/test_harness.ml +++ b/test/test_harness.ml @@ -63,3 +63,32 @@ let run_suite suite cases = let failure_count () = !failures let case_count () = !cases + +let root () = try Sys.getenv "DELTA_ROOT" with Not_found -> "." + +let fixture name = Filename.concat (root ()) (Filename.concat "test/fixture" name) + +let read_fixture name = Native.read_file (fixture name) + +let error_of thunk = + try + ignore (thunk ()); + None + with Diagnostic.Error diagnostic -> Some diagnostic + +let parse text = Parse.program text + +let resolve text = Resolve.program (parse text) + +let infer text = Infer.program (resolve text) + +let parse_error text = error_of (fun () -> parse text) + +let resolve_error text = error_of (fun () -> resolve text) + +let infer_error text = error_of (fun () -> infer text) + +let check_message name expected thunk = + match error_of thunk with + | None -> fail name "expected a diagnostic" + | Some diagnostic -> check_equal_string name expected diagnostic.Diagnostic.message diff --git a/test/test_harness.mli b/test/test_harness.mli index 1f23d82..3297602 100644 --- a/test/test_harness.mli +++ b/test/test_harness.mli @@ -9,3 +9,15 @@ val expect_diagnostic : string -> (unit -> unit) -> unit val run_suite : string -> case list -> unit val failure_count : unit -> int val case_count : unit -> int + +val root : unit -> string +val fixture : string -> string +val read_fixture : string -> string +val error_of : (unit -> 'a) -> Diagnostic.t option +val parse : string -> Syntax.program +val resolve : string -> Resolve.program +val infer : string -> Typed.program +val parse_error : string -> Diagnostic.t option +val resolve_error : string -> Diagnostic.t option +val infer_error : string -> Diagnostic.t option +val check_message : string -> string -> (unit -> 'a) -> unit diff --git a/test/test_main.ml b/test/test_main.ml index 9871c46..2a98865 100644 --- a/test/test_main.ml +++ b/test/test_main.ml @@ -180,12 +180,6 @@ let parse_cases = (match parsed.Syntax.e with Syntax.EInt 7 -> true | _ -> false) ); ] -let root () = try Sys.getenv "DELTA_ROOT" with Not_found -> "." - -let fixture name = Filename.concat (root ()) (Filename.concat "test/fixtures" name) - -let read_fixture name = Native.read_file (fixture name) - let program_error text = try ignore (Parse.program text); @@ -379,5 +373,6 @@ let () = Test_harness.run_suite "parse" parse_cases; Test_harness.run_suite "program" program_cases; Test_harness.run_suite "resolve" resolve_cases; + Test_harness.run_suite "types" Test_type.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 new file mode 100644 index 0000000..b1abc53 --- /dev/null +++ b/test/test_type.ml @@ -0,0 +1,160 @@ +open Test_harness + +let infer_program text = infer text + +let cases = + [ + ( "the example query infers a collection of tuples of string and int", + fun () -> + let program = infer_program (read_fixture "expensive_order.delta") in + check_equal_string "query type" "collection (string, int)" + (Types.pp program.Typed.tp_query_body.Typed.ty); + check_equal_string "input element" "order" (Types.pp program.Typed.tp_input_element) ); + ( "the example helper infers an integer arrow", + fun () -> + let program = infer_program (read_fixture "expensive_order.delta") in + check_equal_string "helper type" "(int -> int)" + (Types.pp (List.hd program.Typed.tp_helpers).Typed.th_scheme.Types.body) ); + ( "revenue infers an integer query", + fun () -> + let program = infer_program (read_fixture "revenue.delta") in + check_equal_string "query type" "int" (Types.pp program.Typed.tp_query_body.Typed.ty) ); + ( "count_large infers an integer query", + fun () -> + let program = infer_program (read_fixture "count_large.delta") in + check_equal_string "query type" "int" (Types.pp program.Typed.tp_query_body.Typed.ty) ); + ( "a helper generalizes at its let binding", + fun () -> + let program = + infer_program + "type order = { customer : string; total : int }\ninput orders : collection order\nlet pair a b = (a, b)\nquery q = orders |> map (fun o -> pair o.total o.customer)\n" + in + let scheme = (List.hd program.Typed.tp_helpers).Typed.th_scheme in + check_equal_int "two generalized variables" 2 (List.length scheme.Types.vars); + check "the helper is a two argument function" + (match Types.repr scheme.Types.body with + | Types.TArrow (_, Types.TArrow _) -> true + | _ -> false); + check_equal_string "query type" "collection (int, string)" + (Types.pp program.Typed.tp_query_body.Typed.ty) ); + ( "a polymorphic helper is instantiated independently at each use", + fun () -> + let program = + infer_program + "type order = { customer : string; total : int }\ninput orders : collection order\nlet identity x = x\nlet tax n = identity n * 20 / 100\nquery q = orders |> map (fun o -> (identity o.customer, tax o.total))\n" + in + check_equal_string "query type" "collection (string, int)" + (Types.pp program.Typed.tp_query_body.Typed.ty) ); + ( "the occurs check rejects self application", + fun () -> + match infer_error "input rows : collection int\nlet f x = x x\nquery q = rows\n" with + | None -> fail "occurs" "expected a diagnostic" + | Some diagnostic -> + check "message mentions an infinite type" + (Util.starts_with "cannot construct the infinite type" diagnostic.Diagnostic.message) ); + ( "an unannotated arithmetic helper that is applied to a string is rejected", + fun () -> + check_message "arithmetic" "type mismatch: expected int but got string" + (fun () -> + infer_program "input rows : collection int\nlet f n = n + 1\nquery q = rows |> map (fun r -> f \"a\")\n") ); + ( "projection on a non record is rejected", + fun () -> + check_message "projection" "type mismatch: expected order but got int" + (fun () -> + infer_program + "type order = { total : int }\ninput rows : collection int\nquery q = rows |> map (fun r -> r.total)\n") ); + ( "unknown record fields are rejected", + fun () -> + check_message "unknown field" "no record type declares a field named `tota`" + (fun () -> + infer_program + "type order = { total : int }\ninput rows : collection order\nquery q = rows |> map (fun r -> r.tota)\n") ); + ( "record literals must provide every field", + fun () -> + check_message "missing field" "record literal of type `order` is missing field `customer`" + (fun () -> + infer_program + "type order = { customer : string; total : int }\ninput rows : collection order\nquery q = rows |> map (fun r -> { total = r.total })\n") ); + ( "record literals reject fields of another record", + fun () -> + check_message "mixed fields" "record `order` has no field named `other`" + (fun () -> + infer_program + "type order = { total : int }\ninput rows : collection order\nquery q = rows |> map (fun r -> { total = r.total; other = 1 })\n") ); + ( "a record field with a collection type is rejected", + fun () -> + check_message "collection field" + "record field `items` has a collection type; collections may not be stored in records" + (fun () -> + infer_program "type box = { items : collection int }\ninput rows : collection box\nquery q = rows\n") ); + ( "nested collection types are rejected", + fun () -> + check_message "collection type position" + "collection types are only allowed as the input type and as query results" + (fun () -> infer_program "input rows : collection (collection int)\nquery q = rows\n") ); + ( "the input must be declared with a collection type", + fun () -> + check_message "input type" + "the input collection must be declared with a collection type: input NAME : collection T" + (fun () -> infer_program "input rows : int\nquery q = 1\n") ); + ( "sum rejects a non integer collection", + fun () -> + check_message "sum" + "sum expects a collection of integers but got collection string" + (fun () -> + infer_program + "type order = { customer : string }\ninput rows : collection order\nquery q = rows |> map (fun r -> r.customer) |> sum\n") ); + ( "sum rejects a scalar operand", + fun () -> + check_message "sum scalar" "sum expects a collection of integers but got int" + (fun () -> infer_program "input rows : collection int\nquery q = sum 1\n") ); + ( "filter must be applied to a collection", + fun () -> + check_message "filter arity" "`filter` and `map` must be applied to a collection" + (fun () -> infer_program "input rows : collection int\nquery q = filter (fun x -> true)\n") ); + ( "map cannot produce a collection", + fun () -> + check_message "nested collection" "nested collections are not supported: mapped element" + (fun () -> infer_program "input rows : collection int\nquery q = rows |> map (fun r -> rows)\n") ); + ( "count rejects a scalar operand", + fun () -> + check_message "count scalar" "count expects a collection but got int" + (fun () -> infer_program "input rows : collection int\nquery q = count 1\n") ); + ( "a query must produce a collection or an integer", + fun () -> + check_message "query result" + "a query must produce a collection or an integer, but this query produces string" + (fun () -> infer_program "input rows : collection int\nquery q = \"hello\"\n") ); + ( "equality is rejected on functions", + fun () -> + check_message "function equality" + "`=` is only supported on integers, booleans, strings, unit, tuples and records of these, but the operands have type (int -> int)" + (fun () -> + infer_program + "input rows : collection int\nlet bad = (fun x -> x + 1) = (fun x -> x + 1)\nquery q = rows\n" ) ); + ( "conditional branches must have the same type", + fun () -> + check_message "branch mismatch" "type mismatch: expected int but got string" + (fun () -> infer_program "input rows : collection int\nquery q = if true then 1 else \"x\"\n") ); + ( "unknown type names are rejected", + fun () -> + check_message "unknown type" "unknown type `order`" + (fun () -> infer_program "input rows : collection order\nquery q = rows\n") ); + ( "negation is integer only", + fun () -> + check_message "negation" "type mismatch: expected bool but got int" + (fun () -> infer_program "input rows : collection int\nquery q = -true\n") ); + ( "a constant integer query is accepted", + fun () -> + let program = infer_program "input rows : collection int\nquery q = 40 + 2\n" in + check_equal_string "query type" "int" (Types.pp program.Typed.tp_query_body.Typed.ty) ); + ( "the typed dump is deterministic", + fun () -> + Types.reset (); + Ident.reset (); + let first = Typed.program_to_string (infer_program (read_fixture "expensive_order.delta")) in + Types.reset (); + Ident.reset (); + let second = Typed.program_to_string (infer_program (read_fixture "expensive_order.delta")) in + check_equal_string "identical dumps" first second ); + ]