Branch data Line data Source code
1 : : /****************************************************************************** 2 : : * This file is part of the cvc5 project. 3 : : * 4 : : * Copyright (c) 2009-2026 by the authors listed in the file AUTHORS 5 : : * in the top-level source directory and their institutional affiliations. 6 : : * All rights reserved. See the file COPYING in the top-level source 7 : : * directory for licensing information. 8 : : * **************************************************************************** 9 : : * 10 : : * Implementation of bags state object. 11 : : */ 12 : : 13 : : #include "theory/bags/solver_state.h" 14 : : 15 : : #include "expr/attribute.h" 16 : : #include "expr/bound_var_manager.h" 17 : : #include "expr/skolem_manager.h" 18 : : #include "theory/smt_engine_subsolver.h" 19 : : #include "theory/uf/equality_engine.h" 20 : : 21 : : using namespace std; 22 : : using namespace cvc5::internal::kind; 23 : : 24 : : namespace cvc5::internal { 25 : : namespace theory { 26 : : namespace bags { 27 : : 28 : 28698 : SolverState::SolverState(Env& env, Valuation val) 29 : 28698 : : TheoryState(env, val), d_partElementSkolems(env.getUserContext()) 30 : : { 31 : 28698 : d_true = nodeManager()->mkConst(true); 32 : 28698 : d_false = nodeManager()->mkConst(false); 33 : 28698 : d_nm = nodeManager(); 34 : 28698 : } 35 : : 36 : 9445 : void SolverState::registerBag(TNode n) 37 : : { 38 [ - + ][ - + ]: 9445 : Assert(n.getType().isBag()); [ - - ] 39 : 9445 : d_bags.insert(n); 40 : 9445 : } 41 : : 42 : 66923 : void SolverState::registerCountTerm(Node bag, Node element, Node skolem) 43 : : { 44 : 66923 : Assert(bag.getType().isBag() && bag == getRepresentative(bag)); 45 : 267692 : Assert(CVC5_EQUAL(element.getType(), bag.getType().getBagElementType()) 46 : : && element == getRepresentative(element)); 47 : 66923 : Assert(skolem.isVar() && skolem.getType().isInteger()); 48 : 66923 : std::pair<Node, Node> pair = std::make_pair(element, skolem); 49 : 66923 : if (std::find(d_bagElements[bag].begin(), d_bagElements[bag].end(), pair) 50 [ + + ]: 133846 : == d_bagElements[bag].end()) 51 : : { 52 : 20542 : d_bagElements[bag].push_back(pair); 53 : : } 54 : 66923 : } 55 : : 56 : 508 : void SolverState::registerGroupTerm(Node n) 57 : : { 58 : : std::shared_ptr<context::CDHashSet<Node>> set = 59 : 508 : std::make_shared<context::CDHashSet<Node>>(d_env.getUserContext()); 60 : 508 : d_partElementSkolems[n] = set; 61 : 508 : } 62 : : 63 : 0 : void SolverState::registerCardinalityTerm(Node n, Node skolem) 64 : : { 65 : 0 : Assert(n.getKind() == Kind::BAG_CARD); 66 : 0 : Assert(skolem.isVar()); 67 : 0 : d_cardTerms[n] = skolem; 68 : 0 : } 69 : : 70 : 0 : Node SolverState::getCardinalitySkolem(Node n) 71 : : { 72 : 0 : Assert(n.getKind() == Kind::BAG_CARD); 73 : 0 : Node bag = getRepresentative(n[0]); 74 : 0 : Node cardTerm = d_nm->mkNode(Kind::BAG_CARD, bag); 75 : 0 : return d_cardTerms[cardTerm]; 76 : 0 : } 77 : : 78 : 0 : bool SolverState::hasCardinalityTerms() const { return !d_cardTerms.empty(); } 79 : : 80 : 179346 : const std::set<Node>& SolverState::getBags() { return d_bags; } 81 : : 82 : 0 : const std::map<Node, Node>& SolverState::getCardinalityTerms() 83 : : { 84 : 0 : return d_cardTerms; 85 : : } 86 : : 87 : 16059 : std::set<Node> SolverState::getElements(Node B) 88 : : { 89 : 32118 : Node bag = getRepresentative(B); 90 : 16059 : std::set<Node> elements; 91 : 16059 : std::vector<std::pair<Node, Node>> pairs = d_bagElements[bag]; 92 [ + + ]: 51568 : for (std::pair<Node, Node> pair : pairs) 93 : : { 94 : 35509 : elements.insert(pair.first); 95 : 35509 : } 96 : 32118 : return elements; 97 : 16059 : } 98 : : 99 : 452 : const std::vector<std::pair<Node, Node>>& SolverState::getElementCountPairs( 100 : : Node n) 101 : : { 102 : 904 : Node bag = getRepresentative(n); 103 : 904 : return d_bagElements[bag]; 104 : 452 : } 105 : : 106 : : struct BagsDeqAttributeId 107 : : { 108 : : }; 109 : : typedef expr::Attribute<BagsDeqAttributeId, Node> BagsDeqAttribute; 110 : : 111 : 36849 : void SolverState::collectDisequalBagTerms() 112 : : { 113 : 36849 : eq::EqClassIterator it = eq::EqClassIterator(d_false, d_ee); 114 [ + + ]: 100312 : while (!it.isFinished()) 115 : : { 116 : 63463 : Node n = (*it); 117 : 63463 : if (n.getKind() == Kind::EQUAL && n[0].getType().isBag()) 118 : : { 119 [ + - ]: 14251 : Trace("bags-eqc") << "Disequal terms: " << n << std::endl; 120 : 28502 : Node A = getRepresentative(n[0]); 121 : 28502 : Node B = getRepresentative(n[1]); 122 [ + + ]: 14251 : Node equal = A <= B ? A.eqNode(B) : B.eqNode(A); 123 [ + + ]: 14251 : if (d_deq.find(equal) == d_deq.end()) 124 : : { 125 : 4487 : SkolemManager* sm = d_nm->getSkolemManager(); 126 : 17948 : Node skolem = sm->mkSkolemFunction(SkolemId::BAGS_DEQ_DIFF, {A, B}); 127 : 4487 : d_deq[equal] = skolem; 128 : 4487 : } 129 : 14251 : } 130 : 63463 : ++it; 131 : 63463 : } 132 : 36849 : } 133 : : 134 : 36136 : const std::map<Node, Node>& SolverState::getDisequalBagTerms() { return d_deq; } 135 : : 136 : 154 : void SolverState::registerPartElementSkolem(Node group, Node skolemElement) 137 : : { 138 [ - + ][ - + ]: 154 : Assert(group.getKind() == Kind::TABLE_GROUP); [ - - ] 139 [ - + ][ - + ]: 616 : AssertEqual(skolemElement.getType(), group[0].getType().getBagElementType()); [ - - ] 140 : 154 : d_partElementSkolems[group].get()->insert(skolemElement); 141 : 154 : } 142 : : 143 : 286 : std::shared_ptr<context::CDHashSet<Node>> SolverState::getPartElementSkolems( 144 : : Node n) 145 : : { 146 [ - + ][ - + ]: 286 : Assert(n.getKind() == Kind::TABLE_GROUP); [ - - ] 147 : 286 : return d_partElementSkolems[n]; 148 : : } 149 : : 150 : 36849 : void SolverState::reset() 151 : : { 152 : 36849 : d_bagElements.clear(); 153 : 36849 : d_bags.clear(); 154 : 36849 : d_deq.clear(); 155 : 36849 : d_cardTerms.clear(); 156 : 36849 : } 157 : : 158 : 44 : void SolverState::checkInjectivity(Node n) 159 : : { 160 : 44 : SkolemManager* sm = d_nm->getSkolemManager(); 161 : 44 : Node f = sm->getOriginalForm(n); 162 [ + + ]: 44 : if (d_functions.find(f) != d_functions.end()) 163 : : { 164 : : // we already know f 165 : 4 : return; 166 : : } 167 : : 168 [ + + ]: 40 : if (f.isVar()) 169 : : { 170 : : // no need to solve. f can be assigned any non injective function 171 : 6 : d_functions[f] = false; 172 : 6 : return; 173 : : } 174 : : 175 : 68 : TypeNode domainType = f.getType().getArgTypes()[0]; 176 : 68 : Node x = NodeManager::mkDummySkolem("x", domainType); 177 : 68 : Node y = NodeManager::mkDummySkolem("y", domainType); 178 : 68 : Node f_x = d_nm->mkNode(Kind::APPLY_UF, f, x); 179 : 68 : Node f_y = d_nm->mkNode(Kind::APPLY_UF, f, y); 180 : 34 : Node f_x_equals_f_y = f_x.eqNode(f_y); 181 : 34 : Node not_x_equals_y = x.eqNode(y).notNode(); 182 : 34 : Node query = f_x_equals_f_y.andNode(not_x_equals_y); 183 : : 184 : 34 : Options subOptions; 185 : 34 : subOptions.copyValues(d_env.getOptions()); 186 : 34 : SubsolverSetupInfo ssi(d_env, subOptions); 187 : 34 : Result result = checkWithSubsolver(query, ssi); 188 [ + + ]: 34 : if (result.getStatus() == Result::Status::UNSAT) 189 : : { 190 : 21 : d_functions[f] = true; 191 : : } 192 : : else 193 : : { 194 : 13 : d_functions[f] = false; 195 : : } 196 [ + + ]: 44 : } 197 : : 198 : 433 : bool SolverState::isInjective(Node n) const 199 : : { 200 : 433 : Node f = d_nm->getSkolemManager()->getOriginalForm(n); 201 [ + - ]: 433 : if (d_functions.find(f) != d_functions.end()) 202 : : { 203 : 433 : return d_functions.at(f); 204 : : } 205 : 0 : return false; 206 : 433 : } 207 : : 208 : : } // namespace bags 209 : : } // namespace theory 210 : : } // namespace cvc5::internal