From 6b83414ee8c95f0f7b2d7eff69c8be059c9a6c47 Mon Sep 17 00:00:00 2001 From: Dana Jansens Date: Tue, 29 Jul 2025 17:34:18 -0400 Subject: [PATCH] Dedupe rewrite constraints without sorting (#5864) Dedupe rewrite constraints by consuming them by their LHS from the map of rewrite values, and dropping any LHS that we see more than once. This essentially uses the map to track which LHS we have seen in place of sorting the rewrite constraints by the LHS. --- toolchain/check/facet_type.cpp | 224 +++++++++++---------------------- 1 file changed, 71 insertions(+), 153 deletions(-) diff --git a/toolchain/check/facet_type.cpp b/toolchain/check/facet_type.cpp index 26b0972466bc..9270c9f8bb31 100644 --- a/toolchain/check/facet_type.cpp +++ b/toolchain/check/facet_type.cpp @@ -263,18 +263,6 @@ struct FacetTypeConstraintValue { SemIR::ElementIndex access_index; SemIR::SpecificInterfaceId specific_interface_id; - friend auto operator<=>(const FacetTypeConstraintValue& lhs, - const FacetTypeConstraintValue& rhs) - -> std::weak_ordering { - if (lhs.entity_name_id != rhs.entity_name_id) { - return lhs.entity_name_id.index <=> rhs.entity_name_id.index; - } - if (lhs.access_index != rhs.access_index) { - return lhs.access_index.index <=> rhs.access_index.index; - } - return lhs.specific_interface_id.index <=> rhs.specific_interface_id.index; - } - friend auto operator==(const FacetTypeConstraintValue& lhs, const FacetTypeConstraintValue& rhs) -> bool = default; }; @@ -298,17 +286,14 @@ static auto GetFacetTypeConstraintValue(Context& context, return std::nullopt; } -// Returns an ordering between two values in a rewrite constraint. Two +// Returns true if two values in a rewrite constraint are equivalent. Two // `ImplWitnessAccess` instructions that refer to the same associated constant -// through the same facet value are treated as equivalent. Otherwise, the -// ordering is somewhat arbitrary with `ImplWitnessAccess` instructions coming -// first. +// through the same facet value are treated as equivalent. static auto CompareFacetTypeConstraintValues(Context& context, SemIR::InstId lhs_id, - SemIR::InstId rhs_id) - -> std::weak_ordering { + SemIR::InstId rhs_id) -> bool { if (lhs_id == rhs_id) { - return std::weak_ordering::equivalent; + return true; } auto lhs_access = context.insts().TryGetAs(lhs_id); @@ -316,54 +301,15 @@ static auto CompareFacetTypeConstraintValues(Context& context, if (lhs_access && rhs_access) { auto lhs_access_value = GetFacetTypeConstraintValue(context, *lhs_access); auto rhs_access_value = GetFacetTypeConstraintValue(context, *rhs_access); - if (lhs_access_value && rhs_access_value) { - return *lhs_access_value <=> *rhs_access_value; - } - // We do *not* want to get the evaluated result of `ImplWitnessAccess` here, // we want to keep them as a reference to an associated constant for the // resolution phase. - return lhs_id.index <=> rhs_id.index; + return lhs_access_value && rhs_access_value && + *lhs_access_value == *rhs_access_value; } - // ImplWitnessAccess sorts before other instructions. - if (lhs_access) { - return std::weak_ordering::less; - } - if (rhs_access) { - return std::weak_ordering::greater; - } - - return context.constant_values().GetConstantInstId(lhs_id).index <=> - context.constant_values().GetConstantInstId(rhs_id).index; -} - -// Sort and dedupe the rewrite constraints, with accesses to the same associated -// constants through the same facet value being treated as equivalent. -static auto SortAndDedupeRewriteConstraints( - Context& context, - llvm::SmallVector& rewrites) { - auto ord = [&](const SemIR::FacetTypeInfo::RewriteConstraint& a, - const SemIR::FacetTypeInfo::RewriteConstraint& b) { - auto lhs = CompareFacetTypeConstraintValues(context, a.lhs_id, b.lhs_id); - if (lhs != std::weak_ordering::equivalent) { - return lhs; - } - auto rhs = CompareFacetTypeConstraintValues(context, a.rhs_id, b.rhs_id); - return rhs; - }; - - auto less = [&](const SemIR::FacetTypeInfo::RewriteConstraint& a, - const SemIR::FacetTypeInfo::RewriteConstraint& b) { - return ord(a, b) == std::weak_ordering::less; - }; - llvm::stable_sort(rewrites, less); - - auto eq = [&](const SemIR::FacetTypeInfo::RewriteConstraint& a, - const SemIR::FacetTypeInfo::RewriteConstraint& b) { - return ord(a, b) == std::weak_ordering::equivalent; - }; - rewrites.erase(llvm::unique(rewrites, eq), rewrites.end()); + return context.constant_values().GetConstantInstId(lhs_id) == + context.constant_values().GetConstantInstId(rhs_id); } // A mapping of each associated constant (represented as `ImplWitnessAccess`) to @@ -403,19 +349,12 @@ class AccessRewriteValues { } } - auto SetFullyRewritten(Value& value, SemIR::InstId rewritten_to_inst_id) + auto SetFullyRewritten(Context& context, Value& value, SemIR::InstId inst_id) -> void { - // TODO: If state == FullyRewrtten and the inst id is different (according - // to `CompareFacetTypeConstraintValues`), we can diagnose writing two - // different values for the same associated constant immediately? - // - // TODO: Once the above is done, we don't need to do the SortAndDedupe - // step in ResolveFacetTypeRewriteConstraints()? We can just convert this - // `map_` into a new set of `RewriteConstraint`s which will already be - // deduped. - if (value.state == BeingRewritten) { - value = {FullyRewritten, rewritten_to_inst_id}; - } + CARBON_CHECK( + value.state == BeingRewritten || + CompareFacetTypeConstraintValues(context, value.inst_id, inst_id)); + value = {FullyRewritten, inst_id}; } private: @@ -573,7 +512,7 @@ class SubstImplWitnessAccessCallbacks : public SubstInstCallbacks { subst_inst_id)) { auto* rewrite_value = rewrite_values_->FindRef(context(), *access); CARBON_CHECK(rewrite_value); - rewrite_values_->SetFullyRewritten(*rewrite_value, inst_id); + rewrite_values_->SetFullyRewritten(context(), *rewrite_value, inst_id); } return inst_id; } @@ -584,7 +523,8 @@ class SubstImplWitnessAccessCallbacks : public SubstInstCallbacks { subst_inst_id)) { auto* rewrite_value = rewrite_values_->FindRef(context(), *access); CARBON_CHECK(rewrite_value); - rewrite_values_->SetFullyRewritten(*rewrite_value, orig_inst_id); + rewrite_values_->SetFullyRewritten(context(), *rewrite_value, + orig_inst_id); } return orig_inst_id; } @@ -624,14 +564,9 @@ auto ResolveFacetTypeRewriteConstraints( return true; } - // Apply rewrite constraints to each other, so that for example: - // `.X = Y and .Y = ()` becomes `.X = () and .Y = ()`. AccessRewriteValues rewrite_values; + for (auto& constraint : rewrites) { - if (constraint.lhs_id == SemIR::ErrorInst::InstId || - constraint.rhs_id == SemIR::ErrorInst::InstId) { - continue; - } auto lhs_access = context.insts().TryGetAs(constraint.lhs_id); if (!lhs_access) { @@ -640,11 +575,8 @@ auto ResolveFacetTypeRewriteConstraints( rewrite_values.InsertNotRewritten(context, *lhs_access, constraint.rhs_id); } + for (auto& constraint : rewrites) { - if (constraint.lhs_id == SemIR::ErrorInst::InstId || - constraint.rhs_id == SemIR::ErrorInst::InstId) { - continue; - } auto lhs_access = context.insts().TryGetAs(constraint.lhs_id); if (!lhs_access) { @@ -652,80 +584,37 @@ auto ResolveFacetTypeRewriteConstraints( } auto* lhs_rewrite_value = rewrite_values.FindRef(context, *lhs_access); + // Every LHS was added with InsertNotRewritten above. CARBON_CHECK(lhs_rewrite_value); rewrite_values.SetBeingRewritten(*lhs_rewrite_value); auto replace_witness_callbacks = SubstImplWitnessAccessCallbacks(&context, loc_id, &rewrite_values); - auto subst_inst_id = + auto rhs_subst_inst_id = SubstInst(context, constraint.rhs_id, replace_witness_callbacks); - constraint.rhs_id = subst_inst_id; - if (constraint.rhs_id == SemIR::ErrorInst::InstId) { + if (rhs_subst_inst_id == SemIR::ErrorInst::InstId) { return false; } - rewrite_values.SetFullyRewritten(*lhs_rewrite_value, subst_inst_id); - } - - // We sort the constraints so that we can find different values being written - // to the same LHS by looking at consecutive rewrite constraints. - // - // It is important to dedupe so that we don't have redundant rewrite - // constraints, as these lead to being diagnosed as a cycle. For example: - // ``` - // (T:! Z where .X = .Y) where .X = .Y - // ``` - // Here we drop one of the `.X = .Y` in the resulting facet type. If we don't, - // then the `.X` in the outer facet type can be evaluated to `.Y` from the - // inner facet type, resulting in `.Y = .Y` which is a cycle. By deduping, we - // avoid any LHS of a rewrite constraint from being evaluated to the RHS of - // a duplicate rewrite constraint. - SortAndDedupeRewriteConstraints(context, rewrites); - - for (size_t i = 0; i < rewrites.size() - 1; ++i) { - auto& constraint = rewrites[i]; - if (constraint.lhs_id == SemIR::ErrorInst::InstId || - constraint.rhs_id == SemIR::ErrorInst::InstId) { - continue; - } - - auto lhs_access = - context.insts().TryGetAs(constraint.lhs_id); - if (!lhs_access) { - continue; - } - - // This loop moves `i` to the last position with the same LHS value, so that - // we don't diagnose more than once within the same contiguous range of - // assignments to a single LHS value. - for (; i < rewrites.size() - 1; ++i) { - auto& next = rewrites[i + 1]; - auto next_lhs_access = - context.insts().TryGetAs(next.lhs_id); - if (!next_lhs_access) { - break; - } - - if (CompareFacetTypeConstraintValues(context, constraint.lhs_id, - next.lhs_id) != - std::weak_ordering::equivalent) { - break; - } - - if (constraint.rhs_id != SemIR::ErrorInst::InstId && - next.rhs_id != SemIR::ErrorInst::InstId) { - CARBON_DIAGNOSTIC( - AssociatedConstantWithDifferentValues, Error, - "associated constant {0} given two different values {1} and {2}", - InstIdAsConstant, InstIdAsConstant, InstIdAsConstant); - // Use inst id ordering as a simple proxy for source ordering, to try - // to name the values in the same order they appear in the facet type. - auto source_order1 = constraint.rhs_id.index < next.rhs_id.index - ? constraint.rhs_id - : next.rhs_id; - auto source_order2 = constraint.rhs_id.index >= next.rhs_id.index - ? constraint.rhs_id - : next.rhs_id; + if (lhs_rewrite_value->state == AccessRewriteValues::FullyRewritten && + !CompareFacetTypeConstraintValues(context, lhs_rewrite_value->inst_id, + rhs_subst_inst_id)) { + if (lhs_rewrite_value->inst_id != SemIR::ErrorInst::InstId) { + CARBON_DIAGNOSTIC(AssociatedConstantWithDifferentValues, Error, + "associated constant {0} given two different " + "values {1} and {2}", + InstIdAsConstant, InstIdAsConstant, InstIdAsConstant); + // Use inst id ordering as a simple proxy for source ordering, to + // try to name the values in the same order they appear in the facet + // type. + auto source_order1 = + lhs_rewrite_value->inst_id.index < rhs_subst_inst_id.index + ? lhs_rewrite_value->inst_id + : rhs_subst_inst_id; + auto source_order2 = + lhs_rewrite_value->inst_id.index >= rhs_subst_inst_id.index + ? lhs_rewrite_value->inst_id + : rhs_subst_inst_id; // TODO: It would be nice to note the places where the values are // assigned but rewrite constraint instructions are from canonical // constant values, and have no locations. We'd need to store a @@ -733,12 +622,41 @@ auto ResolveFacetTypeRewriteConstraints( context.emitter().Emit(loc_id, AssociatedConstantWithDifferentValues, constraint.lhs_id, source_order1, source_order2); } - constraint.rhs_id = SemIR::ErrorInst::InstId; - next.rhs_id = SemIR::ErrorInst::InstId; return false; } + + rewrite_values.SetFullyRewritten(context, *lhs_rewrite_value, + rhs_subst_inst_id); } + // Rebuild the `rewrites` vector with resolved values for the RHS. Drop any + // duplicate rewrites in the `rewrites` vector by walking through the + // `rewrite_values` map and dropping the computed RHS value for each LHS the + // first time we see it, and erasing the constraint from the vector if we see + // the same LHS again. + size_t keep_size = rewrites.size(); + for (size_t i = 0; i < keep_size;) { + auto& constraint = rewrites[i]; + + auto lhs_access = + context.insts().TryGetAs(constraint.lhs_id); + if (!lhs_access) { + ++i; + continue; + } + + auto& rewrite_value = *rewrite_values.FindRef(context, *lhs_access); + auto rhs_id = std::exchange(rewrite_value.inst_id, SemIR::InstId::None); + if (rhs_id == SemIR::InstId::None) { + std::swap(rewrites[i], rewrites[keep_size - 1]); + --keep_size; + } else { + rewrites[i].rhs_id = rhs_id; + ++i; + } + } + rewrites.erase(rewrites.begin() + keep_size, rewrites.end()); + return true; }