diff --git a/src/frontends/lean/elaborator.cpp b/src/frontends/lean/elaborator.cpp index 12d1aa485e..27d2c509c8 100644 --- a/src/frontends/lean/elaborator.cpp +++ b/src/frontends/lean/elaborator.cpp @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Author: Leonardo de Moura */ +#include #include "kernel/find_fn.h" #include "kernel/for_each_fn.h" #include "kernel/instantiate.h" @@ -12,6 +13,7 @@ Author: Leonardo de Moura #include "library/constants.h" #include "library/placeholder.h" #include "library/explicit.h" +#include "library/scope_pos_info_provider.h" #include "library/choice.h" #include "library/typed_expr.h" #include "library/compiler/rec_fn_macro.h" @@ -38,7 +40,7 @@ elaborator::elaborator(environment const & env, options const & opts, local_leve format elaborator::pp(expr const & e) { formatter_factory const & factory = get_global_ios().get_formatter_factory(); formatter fmt = factory(m_env, m_opts, m_ctx); - return fmt(e); + return fmt(instantiate_mvars(e)); } format elaborator::pp_indent(expr const & e) { @@ -46,6 +48,15 @@ format elaborator::pp_indent(expr const & e) { return nest(i, line() + pp(e)); } +static std::string pos_string_for(expr const & e) { + pos_info_provider * provider = get_pos_info_provider(); + if (!provider) return "'unknown'"; + pos_info pos = provider->get_pos_info_or_some(e); + sstream s; + s << provider->get_file_name() << ":" << pos.first << ":" << pos.second << ":"; + return s.str(); +} + level elaborator::mk_univ_metavar() { level r = m_ctx.mk_univ_metavar_decl(); m_uvar_stack.push_back(r); @@ -63,18 +74,33 @@ expr elaborator::mk_type_metavar() { return mk_metavar(mk_sort(l)); } +expr elaborator::mk_instance_core(expr const & C) { + optional inst = m_ctx.mk_class_instance(C); + trace_elab(tout() << pos_string_for(C) << " type class resolution\n " << C << "\ninstance\n"; + if (!inst) tout() << "[failed]\n"; + else tout() << " " << *inst << "\n";); + if (!inst) + throw_elaborator_exception(C, format("failed to synthesize type class instance for") + pp_indent(C)); + return *inst; +} + expr elaborator::mk_instance(expr const & C) { if (has_expr_metavar(C)) { expr inst = mk_metavar(C); m_instance_stack.push_back(inst); return inst; - } else if (optional inst = m_ctx.mk_class_instance(C)) { - return *inst; } else { - throw_elaborator_exception(C, format("failed to synthesize type class instance for") + pp_indent(C)); + return mk_instance_core(C); } } +expr elaborator::instantiate_mvars(expr const & e) { + expr r = m_ctx.instantiate_mvars(e); + if (r.get_tag() == nulltag) + r.set_tag(e.get_tag()); + return r; +} + level elaborator::get_level(expr const & A) { return ::lean::get_level(m_ctx, A); } @@ -147,9 +173,9 @@ auto elaborator::use_elim_elab_core(name const & fn) -> optional { buffer C_args; expr const & C = get_app_args(type, C_args); - if (!is_local(C) || !std::all_of(C_args.begin(), C_args.end(), is_local)) { + if (!is_local(C) || C_args.empty() || !std::all_of(C_args.begin(), C_args.end(), is_local)) { trace_elab_detail(tout() << "'eliminator' elaboration is not used for '" << fn << - "' because of resulting type is not of the expected form\n";); + "' because resulting type is not of the expected form\n";); return optional(); } @@ -223,24 +249,6 @@ auto elaborator::use_elim_elab(name const & fn) -> optional { return r; } -void elaborator::resolve_instances_from(unsigned old_sz) { - unsigned j = old_sz; - for (unsigned i = old_sz; i < m_instance_stack.size(); i++) { - expr inst_mvar = m_instance_stack[i]; - if (is_mvar_assigned(inst_mvar)) - continue; - expr C = instantiate_mvars(infer_type(inst_mvar)); - if (!has_expr_metavar(C)) { - if (!assign_mvar(inst_mvar, mk_instance(C))) - throw_elaborator_exception(inst_mvar, format("failed to assign type class instance to placeholder")); - } else { - m_instance_stack[j] = m_instance_stack[i]; - j++; - } - } - m_instance_stack.shrink(j); -} - expr elaborator::ensure_function(expr const & e, expr const & ref) { // TODO(Leo): infer type and try to apply coercion to function if needed // ref is only needed for generating error messages @@ -253,29 +261,42 @@ expr elaborator::ensure_type(expr const & e, expr const & ref) { return e; } -expr elaborator::ensure_has_type(expr const & e, expr const & type, expr const & ref) { - // TODO(Leo): infer e type, and apply coercions to type if needed. +optional elaborator::ensure_has_type(expr const & e, expr const & e_type, expr const & type) { + if (is_def_eq(e_type, type)) + return some_expr(e); + + // TODO(Leo): try coercions // ref is only needed for generating error messages - return e; + + return none_expr(); } expr elaborator::visit_typed_expr(expr const & e) { - expr val = get_typed_expr_expr(e); - expr type = get_typed_expr_type(e); - expr new_type = ensure_type(visit(type, none_expr()), type); - expr new_val = visit(val, some_expr(new_type)); - return ensure_has_type(new_val, new_type, e); + expr val = get_typed_expr_expr(e); + expr type = get_typed_expr_type(e); + expr new_type = ensure_type(visit(type, none_expr()), type); + expr new_val = visit(val, some_expr(new_type)); + expr new_val_type = infer_type(new_val); + if (auto r = ensure_has_type(new_val, new_val_type, new_type)) + return *r; + throw_elaborator_exception(e, + format("invalid type ascription, expression has type") + pp_indent(new_val_type) + + line() + format(", but is expected to have type") + pp_indent(new_type)); } -expr elaborator::visit_prenum_core(expr const & e, optional const & expected_type) { +expr elaborator::visit_prenum(expr const & e, optional const & expected_type) { lean_assert(is_prenum(e)); mpz const & v = prenum_value(e); tag e_tag = e.get_tag(); expr A; - if (expected_type) + if (expected_type) { A = *expected_type; - else + if (is_metavar(*expected_type)) + m_numeral_type_stack.push_back(A); + } else { A = mk_type_metavar(); + m_numeral_type_stack.push_back(A); + } level A_lvl = get_level(A); levels ls(A_lvl); if (v.is_neg()) @@ -302,7 +323,8 @@ expr elaborator::visit_prenum_core(expr const & e, optional const & expect return mk_app(mk_app(mk_app(mk_constant(get_bit0_name(), ls), A, e_tag), S_add, e_tag), r, e_tag); } else { expr r = convert(v / 2); - return mk_app(mk_app(mk_app(mk_app(mk_constant(get_bit1_name(), ls), A, e_tag), S_one, e_tag), S_add, e_tag), r, e_tag); + return mk_app(mk_app(mk_app(mk_app(mk_constant(get_bit1_name(), ls), A, e_tag), S_one, e_tag), + S_add, e_tag), r, e_tag); } }; return convert(v); @@ -310,25 +332,6 @@ expr elaborator::visit_prenum_core(expr const & e, optional const & expect } } -expr elaborator::get_default_numeral_type() { - // TODO(Leo): allow user to modify default? - return mk_constant(get_nat_name()); -} - -expr elaborator::visit_prenum(expr const & e, optional const & expected_type) { - unsigned old_sz = m_instance_stack.size(); - expr r = visit_prenum_core(e, expected_type); - if (!expected_type) { - expr t = infer_type(r); - if (!assign_mvar(t, get_default_numeral_type())) - throw_elaborator_exception("invalid numeral, failed to force numeral to be a nat", e); - resolve_instances_from(old_sz); - } - if (m_instance_stack.size() != old_sz) - throw_elaborator_exception("invalid numeral, failed to synthesize type class instances", e); - return instantiate_mvars(r); -} - expr elaborator::visit_sort(expr const & e) { expr r = update_sort(e, replace_univ_placeholder(sort_level(e))); if (contains_placeholder(sort_level(e))) @@ -405,32 +408,128 @@ void elaborator::validate_overloads(buffer const & fns, expr const & ref) } } -expr elaborator::visit_poly_app(expr const & fn, buffer const & args, optional const & expected_type) { +expr elaborator::visit_poly_app(expr const & fn, buffer const & args, optional const & expected_type, + expr const & ref) { // TODO(Leo) lean_unreachable(); } expr elaborator::visit_elim_app(expr const & fn, elim_info const & info, buffer const & args, - optional const & expected_type) { + optional const & expected_type, expr const & ref) { // TODO(Leo): check if fn is fully applied. If it is not, then invoke visit_default_app if (!expected_type) { throw_elaborator_exception(sstream() << "invalid '" << const_name(fn) << "' application, " - << "elaborator has special support for this kind of application, " - << "but the expected type must be known", fn); + << "elaborator has special support for this kind of application " + << "(it is handled as an \"eliminator\"), " + << "but the expected type must be known", ref); } // TODO(Leo) lean_unreachable(); } -expr elaborator::visit_default_app(buffer const & fns, arg_mask amask, buffer const & args, optional const & expected_type) { - if (fns.size() == 1 && args.size() == 0) - return fns[0]; // Easy case +expr elaborator::visit_overloaded_app(buffer const & fns, arg_mask amask, buffer const & args, + optional const & expected_type, expr const & ref) { + // Remark: fns have been elaborated, but args have not. + + buffer fn_types; + buffer ok; + for (expr const & fn : fns) { + fn_types.push_back(infer_type(fn)); + ok.push_back(true); + } + + buffer new_args; + + // while the fn_types are compatible, we keep processing arguments + unsigned i = 0; + while (true) { + bool all_pi = true; + for (unsigned fn_idx = 0; fn_idx < fns.size(); fn_idx++) { + if (!ok[fn_idx]) + continue; + expr & fn_type = fn_types[fn_idx]; + fn_type = whnf(fn_type); + if (!is_pi(fn_type)) { + if (i < args.size()) { + trace_elab_detail(tout() << "overloaded application '" << fns[i] << "' is not being considered at" << + "position " << pos_string_for(ref) << " because " << + "of the number of explicit arguments provided\n";); + ok[fn_idx] = false; + } else { + all_pi = false; + } + } + } + if (!all_pi) break; + } // TODO(Leo) lean_unreachable(); } +expr elaborator::visit_default_app(expr const & fn, arg_mask amask, buffer const & args, + optional const & expected_type, expr const & ref) { + expr fn_type = try_to_pi(infer_type(fn)); + unsigned i = 0; + buffer> types_to_check; + buffer new_args; + expr type = fn_type; + while (is_pi(type)) { + binder_info const & bi = binding_info(type); + expr const & d = binding_domain(type); + expr new_arg; + if ((amask == arg_mask::Default && !is_explicit(bi)) || + (amask == arg_mask::Simple && !bi.is_inst_implicit() && !is_first_order(d))) { + // implicit argument + if (bi.is_inst_implicit()) + new_arg = mk_instance(d); + else + new_arg = mk_metavar(d); + // we don't types of implicit arguments + types_to_check.push_back(none_expr()); + } else if (i < args.size()) { + // explicit argument + new_arg = visit(args[i], some_expr(d)); + types_to_check.push_back(some_expr(d)); + i++; + } else { + break; + } + new_args.push_back(new_arg); + type = try_to_pi(instantiate(binding_body(type), new_arg)); + } + + lean_assert(new_args.size() == types_to_check.size()); + + if (i != args.size()) { + throw_elaborator_exception(ref, format("invalid function application, too many arguments, function type:") + + pp_indent(fn_type)); + } + + for (unsigned i = 0; i < new_args.size(); i++) { + if (optional expected_type = types_to_check[i]) { + expr & new_arg = new_args[i]; + expr new_arg_type = infer_type(new_arg); + if (optional new_new_arg = ensure_has_type(new_arg, new_arg_type, *expected_type)) { + new_arg = *new_new_arg; + } else { + throw_elaborator_exception(ref, + format("type mismatch at application") + + pp_indent(mk_app(fn, i+1, new_args.data())) + + line() + format("term") + + pp_indent(new_arg) + + line() + format("has type") + + pp_indent(new_arg_type) + + line() + format("but is expected to have type") + + pp_indent(*expected_type)); + } + } + } + + return mk_app(fn, new_args.size(), new_args.data()); +} + expr elaborator::visit_app_core(expr fn, buffer const & args, optional const & expected_type) { arg_mask amask = arg_mask::Default; if (is_explicit(fn)) { @@ -440,8 +539,11 @@ expr elaborator::visit_app_core(expr fn, buffer const & args, optional fns; + + bool has_args = !args.empty(); + if (is_choice(fn)) { + buffer fns; if (amask != arg_mask::Default) throw_elaborator_exception("invalid explicit annotation, symbol is overloaded " "(solution: use fully qualified names)", fn); @@ -450,23 +552,21 @@ expr elaborator::visit_app_core(expr fn, buffer const & args, optional const & expected_type) { @@ -534,8 +634,76 @@ expr elaborator::visit(expr const & e, optional const & expected_type) { lean_unreachable(); // LCOV_EXCL_LINE } +expr elaborator::get_default_numeral_type() { + // TODO(Leo): allow user to modify default? + return mk_constant(get_nat_name()); +} + +void elaborator::ensure_numeral_types_assigned(checkpoint const & C) { + for (unsigned i = C.m_numeral_type_stack_sz; i < m_numeral_type_stack.size(); i++) { + expr A = instantiate_mvars(m_numeral_type_stack[i]); + if (is_metavar(A)) { + if (!assign_mvar(A, get_default_numeral_type())) + throw_elaborator_exception("invalid numeral, failed to force numeral to be a nat", A); + } + } +} + +void elaborator::synthesize_type_class_instances(checkpoint const & C) { + for (unsigned i = C.m_instance_stack_sz; i < m_instance_stack.size(); i++) { + expr inst = instantiate_mvars(m_instance_stack[i]); + if (!has_expr_metavar(inst)) { + trace_elab(tout() << "skipping type class resolution at " << pos_string_for(m_instance_stack[i]) + << ", placeholder instantiated using type inference\n";); + continue; + } + expr inst_type = instantiate_mvars(infer_type(inst)); + if (!has_expr_metavar(inst_type)) { + if (!is_def_eq(inst, mk_instance_core(inst_type))) + throw_elaborator_exception(m_instance_stack[i], + format("failed to assign type class instance to placeholder")); + } else { + throw_elaborator_exception(m_instance_stack[i], + format("type class instance cannot be synthesized, type has metavariables") + + pp_indent(inst_type)); + } + } +} + +void elaborator::invoke_tactics(checkpoint const & C) { + // TODO(Leo) +} + +void elaborator::ensure_no_unassigned_metavars(checkpoint const & C) { +} + +void elaborator::process_checkpoint(checkpoint const & C) { + ensure_numeral_types_assigned(C); + synthesize_type_class_instances(C); + invoke_tactics(C); + ensure_no_unassigned_metavars(C); +} + +elaborator::checkpoint::checkpoint(elaborator & e): + m_elaborator(e), + m_uvar_stack_sz(e.m_uvar_stack.size()), + m_mvar_stack_sz(e.m_mvar_stack.size()), + m_instance_stack_sz(e.m_instance_stack.size()), + m_numeral_type_stack_sz(e.m_numeral_type_stack.size()) { +} + +elaborator::checkpoint::~checkpoint() { + m_elaborator.m_uvar_stack.shrink(m_uvar_stack_sz); + m_elaborator.m_mvar_stack.shrink(m_mvar_stack_sz); + m_elaborator.m_instance_stack.shrink(m_instance_stack_sz); + m_elaborator.m_numeral_type_stack.shrink(m_numeral_type_stack_sz); +} + std::tuple elaborator::operator()(expr const & e) { + checkpoint C(*this); expr r = visit(e, none_expr()); + process_checkpoint(C); + r = instantiate_mvars(r); // TODO(Leo) level_param_names ls; return std::make_tuple(r, ls); diff --git a/src/frontends/lean/elaborator.h b/src/frontends/lean/elaborator.h index a9dc67a402..44f988101a 100644 --- a/src/frontends/lean/elaborator.h +++ b/src/frontends/lean/elaborator.h @@ -27,6 +27,17 @@ class elaborator { buffer m_uvar_stack; buffer m_mvar_stack; buffer m_instance_stack; + buffer m_numeral_type_stack; + + struct checkpoint { + elaborator & m_elaborator; + unsigned m_uvar_stack_sz; + unsigned m_mvar_stack_sz; + unsigned m_instance_stack_sz; + unsigned m_numeral_type_stack_sz; + checkpoint(elaborator & e); + ~checkpoint(); + }; /** \brief Cache for constants that are handled using "polymorphic" elaboration. */ name_map m_poly_cache; @@ -62,9 +73,11 @@ class elaborator { format pp_indent(expr const & e); expr infer_type(expr const & e) { return m_ctx.infer(e); } + expr whnf(expr const & e) { return m_ctx.whnf(e); } + expr try_to_pi(expr const & e) { return m_ctx.try_to_pi(e); } bool is_def_eq(expr const & e1, expr const & e2) { return m_ctx.is_def_eq(e1, e2); } bool assign_mvar(expr const & e1, expr const & e2) { lean_assert(is_metavar(e1)); return is_def_eq(e1, e2); } - expr instantiate_mvars(expr const & e) { return m_ctx.instantiate_mvars(e); } + expr instantiate_mvars(expr const & e); bool is_uvar_assigned(level const & l) const { return m_ctx.is_assigned(l); } bool is_mvar_assigned(expr const & e) const { return m_ctx.is_assigned(e); } void resolve_instances_from(unsigned old_sz); @@ -72,14 +85,15 @@ class elaborator { level mk_univ_metavar(); expr mk_metavar(expr const & A); expr mk_type_metavar(); + expr mk_instance_core(expr const & C); expr mk_instance(expr const & C); level get_level(expr const & A); level replace_univ_placeholder(level const & l); expr ensure_type(expr const & e, expr const & ref); - expr ensure_has_type(expr const & e, expr const & type, expr const & ref); expr ensure_function(expr const & e, expr const & ref); + optional ensure_has_type(expr const & e, expr const & e_type, expr const & type); bool use_poly_elab(name const & fn); optional use_elim_elab_core(name const & fn); @@ -95,9 +109,14 @@ class elaborator { expr visit_function(expr const & fn, bool has_args, expr const & ref); void throw_overload_exception(buffer const & fns, sstream & ss, expr const & ref); void validate_overloads(buffer const & fns, expr const & ref); - expr visit_poly_app(expr const & fn, buffer const & args, optional const & expected_type); - expr visit_elim_app(expr const & fn, elim_info const & info, buffer const & args, optional const & expected_type); - expr visit_default_app(buffer const & fns, arg_mask amask, buffer const & args, optional const & expected_type); + expr visit_poly_app(expr const & fn, buffer const & args, + optional const & expected_type, expr const & ref); + expr visit_elim_app(expr const & fn, elim_info const & info, buffer const & args, + optional const & expected_type, expr const & ref); + expr visit_overloaded_app(buffer const & fns, arg_mask amask, buffer const & args, + optional const & expected_type, expr const & ref); + expr visit_default_app(expr const & fn, arg_mask amask, buffer const & args, + optional const & expected_type, expr const & ref); expr visit_app_core(expr fn, buffer const & args, optional const & expected_type); expr visit_local(expr const & e, optional const & expected_type); expr visit_constant(expr const & e, optional const & expected_type); @@ -108,6 +127,12 @@ class elaborator { expr visit_let(expr const & e, optional const & expected_type); expr visit(expr const & e, optional const & expected_type); + void ensure_numeral_types_assigned(checkpoint const & C); + void synthesize_type_class_instances(checkpoint const & C); + void invoke_tactics(checkpoint const & C); + void ensure_no_unassigned_metavars(checkpoint const & C); + void process_checkpoint(checkpoint const & C); + public: elaborator(environment const & env, options const & opts, local_level_decls const & lls, metavar_context const & mctx, local_context const & lctx); diff --git a/tests/lean/elab1.lean.expected.out b/tests/lean/elab1.lean.expected.out index 0a3060d8d0..bcc3a02279 100644 --- a/tests/lean/elab1.lean.expected.out +++ b/tests/lean/elab1.lean.expected.out @@ -7,4 +7,4 @@ elab1.lean:17:6: error: invalid overloaded application, elaborator has special s >> [choice add boo.add foo.add] elab1.lean:19:6: error: invalid overloaded application, elaborator has special support for 'add' (it is handled as a polymorphic application), but this kind of constant cannot be overloaded (solution: use fully qualified names) (overloads: add, boo.add, foo.add) >> eq.subst H1 H2 -elab1.lean:25:6: error: invalid 'eq.subst' application, elaborator has special support for this kind of application, but the expected type must be known +elab1.lean:25:6: error: invalid 'eq.subst' application, elaborator has special support for this kind of application (it is handled as an "eliminator"), but the expected type must be known diff --git a/tests/lean/elab2.lean b/tests/lean/elab2.lean new file mode 100644 index 0000000000..322ba8a76a --- /dev/null +++ b/tests/lean/elab2.lean @@ -0,0 +1,15 @@ +definition foo {A B : Type} [has_add A] (a : A) (b : B) : A := +a + +set_option trace.elaborator true +set_option trace.elaborator_detail true + +set_option pp.all true +#elab foo 0 1 + +definition bla {A B : Type} (a₁ a₂ : A) (b : B) : A := +a₁ + +#elab bla nat.zero tt 1 + +#elab bla 0 0 tt diff --git a/tests/lean/elab2.lean.expected.out b/tests/lean/elab2.lean.expected.out new file mode 100644 index 0000000000..d0337eaf43 --- /dev/null +++ b/tests/lean/elab2.lean.expected.out @@ -0,0 +1,44 @@ +>> foo [prenum] [prenum] +[elaborator_detail] 'polymorphic' elaboration is not used for 'foo' because it has more than one Type parameter +[elaborator_detail] 'eliminator' elaboration is not used for 'foo' because resulting type is not of the expected form +[elaborator] elab2.lean:1:29: type class resolution + has_add.{1} nat +instance + nat_has_add +[elaborator] elab2.lean:8:10: type class resolution + has_zero.{1} nat +instance + nat_has_zero +[elaborator] elab2.lean:8:12: type class resolution + has_one.{1} nat +instance + nat_has_one +>> @foo.{1 1} nat nat nat_has_add (@zero.{1} nat nat_has_zero) (@one.{1} nat nat_has_one) +>> nat +>> bla nat.zero bool.tt [prenum] +[elaborator_detail] 'polymorphic' elaboration is not used for 'bla' because it has more than one Type parameter +[elaborator_detail] 'eliminator' elaboration is not used for 'bla' because resulting type is not of the expected form +[elaborator_detail] 'eliminator' elaboration is not used for 'nat.zero' because resulting type is not of the expected form +[elaborator_detail] 'eliminator' elaboration is not used for 'bool.tt' because resulting type is not of the expected form +elab2.lean:13:6: error: type mismatch at application + @bla.{1 ?M_1} nat ?M_2 nat.zero bool.tt +term + bool.tt +has type + bool +but is expected to have type + nat +>> bla [prenum] [prenum] bool.tt +[elaborator_detail] 'polymorphic' elaboration is not used for 'bla' because it has more than one Type parameter +[elaborator_detail] 'eliminator' elaboration is not used for 'bla' because resulting type is not of the expected form +[elaborator_detail] 'eliminator' elaboration is not used for 'bool.tt' because resulting type is not of the expected form +[elaborator] elab2.lean:15:10: type class resolution + has_zero.{1} nat +instance + nat_has_zero +[elaborator] elab2.lean:15:12: type class resolution + has_zero.{1} nat +instance + nat_has_zero +>> @bla.{1 1} nat bool (@zero.{1} nat nat_has_zero) (@zero.{1} nat nat_has_zero) bool.tt +>> nat diff --git a/tests/lean/run/elab1.lean b/tests/lean/run/elab1.lean index 8fe2c8f888..6b9783dcaf 100644 --- a/tests/lean/run/elab1.lean +++ b/tests/lean/run/elab1.lean @@ -14,5 +14,3 @@ attribute int_comm_ring [instance] #elab int #elab (2 : int) - -#elab @eq.refl