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 : : * Utilities for function constants
11 : : */
12 : :
13 : : #include "theory/uf/function_const.h"
14 : :
15 : : #include "expr/array_store_all.h"
16 : : #include "expr/attribute.h"
17 : : #include "expr/bound_var_manager.h"
18 : : #include "expr/function_array_const.h"
19 : : #include "expr/skolem_manager.h"
20 : : #include "theory/arrays/theory_arrays_rewriter.h"
21 : : #include "theory/rewriter.h"
22 : : #include "util/rational.h"
23 : :
24 : : namespace cvc5::internal {
25 : : namespace theory {
26 : : namespace uf {
27 : :
28 : : /**
29 : : * An attribute to cache the conversion between array constants and lambdas.
30 : : */
31 : : struct ArrayToLambdaTag
32 : : {
33 : : };
34 : : using ArrayToLambdaAttribute = expr::Attribute<ArrayToLambdaTag, Node>;
35 : :
36 : 1314409 : Node FunctionConst::toLambda(TNode n)
37 : : {
38 : 1314409 : Kind nk = n.getKind();
39 [ + + ]: 1314409 : if (nk == Kind::LAMBDA)
40 : : {
41 : 46737 : return n;
42 : : }
43 [ + + ]: 1267672 : else if (nk == Kind::FUNCTION_ARRAY_CONST)
44 : : {
45 : : ArrayToLambdaAttribute atla;
46 [ + + ]: 224307 : if (n.hasAttribute(atla))
47 : : {
48 : 223164 : return n.getAttribute(atla);
49 : : }
50 : 1143 : const FunctionArrayConst& fc = n.getConst<FunctionArrayConst>();
51 : 1143 : Node avalue = fc.getArrayValue();
52 : 1143 : TypeNode tn = fc.getType();
53 [ - + ][ - + ]: 1143 : Assert(tn.isFunction());
[ - - ]
54 : 1143 : std::vector<TypeNode> argTypes = tn.getArgTypes();
55 : 1143 : std::vector<Node> bvs;
56 : 1143 : NodeManager* nm = n.getNodeManager();
57 : 1143 : BoundVarManager* bvm = nm->getBoundVarManager();
58 : : // associate a unique bound variable list with the value
59 [ + + ]: 2791 : for (size_t i = 0, nargs = argTypes.size(); i < nargs; i++)
60 : : {
61 : : Node cacheVal =
62 : 3296 : BoundVarManager::getCacheValue(n, nm->mkConstInt(Rational(i)));
63 : : Node v = bvm->mkBoundVar(
64 : 3296 : BoundVarId::FUN_BOUND_VAR_LIST, cacheVal, argTypes[i]);
65 : 1648 : bvs.push_back(v);
66 : 1648 : }
67 : 1143 : Node bvl = nm->mkNode(Kind::BOUND_VAR_LIST, bvs);
68 : 2286 : Node lam = getLambdaForArrayRepresentation(avalue, bvl);
69 : 1143 : n.setAttribute(atla, lam);
70 : 1143 : return lam;
71 : 1143 : }
72 : 1043365 : return Node::null();
73 : : }
74 : :
75 : 540 : Node FunctionConst::getDefinition(TNode f)
76 : : {
77 : 540 : return toLambda(SkolemManager::getUnpurifiedForm(f));
78 : : }
79 : :
80 : 0 : TypeNode FunctionConst::getFunctionTypeForArrayType(TypeNode atn, Node bvl)
81 : : {
82 : 0 : std::vector<TypeNode> children;
83 [ - - ]: 0 : for (unsigned i = 0; i < bvl.getNumChildren(); i++)
84 : : {
85 : 0 : Assert(atn.isArray());
86 : 0 : Assert(bvl[i].getType() == atn.getArrayIndexType());
87 : 0 : children.push_back(atn.getArrayIndexType());
88 : 0 : atn = atn.getArrayConstituentType();
89 : : }
90 : 0 : children.push_back(atn);
91 : 0 : return bvl.getNodeManager()->mkFunctionType(children);
92 : 0 : }
93 : :
94 : 9 : TypeNode FunctionConst::getArrayTypeForFunctionType(TypeNode ftn)
95 : : {
96 [ - + ][ - + ]: 9 : Assert(ftn.isFunction());
[ - - ]
97 : : // construct the curried array type
98 : 9 : size_t nchildren = ftn.getNumChildren();
99 : 9 : TypeNode ret = ftn[nchildren - 1];
100 [ + + ]: 18 : for (size_t i = 0; i < nchildren - 1; i++)
101 : : {
102 : 9 : size_t ii = nchildren - i - 2;
103 : 9 : ret = NodeManager::mkArrayType(ftn[ii], ret);
104 : : }
105 : 9 : return ret;
106 : 0 : }
107 : :
108 : 7398 : Node FunctionConst::getLambdaForArrayRepresentationRec(
109 : : TNode a,
110 : : TNode bvl,
111 : : unsigned bvlIndex,
112 : : std::unordered_map<TNode, Node>& visited)
113 : : {
114 : 7398 : std::unordered_map<TNode, Node>::iterator it = visited.find(a);
115 [ + + ]: 7398 : if (it != visited.end())
116 : : {
117 : 989 : return it->second;
118 : : }
119 : 6409 : Node ret;
120 [ + + ]: 6409 : if (bvlIndex < bvl.getNumChildren())
121 : : {
122 [ - + ][ - + ]: 4048 : Assert(a.getType().isArray());
[ - - ]
123 [ + + ]: 4048 : if (a.getKind() == Kind::STORE)
124 : : {
125 : : // convert the array recursively
126 : : Node body =
127 : 4414 : getLambdaForArrayRepresentationRec(a[0], bvl, bvlIndex, visited);
128 [ + - ]: 2207 : if (!body.isNull())
129 : : {
130 : : // convert the value recursively (bounded by the number of arguments
131 : : // in bvl)
132 : : Node val = getLambdaForArrayRepresentationRec(
133 : 4414 : a[2], bvl, bvlIndex + 1, visited);
134 [ + - ]: 2207 : if (!val.isNull())
135 : : {
136 [ - + ][ - + ]: 11035 : AssertEqual(a[1].getType(), bvl[bvlIndex].getType());
[ - - ]
137 [ - + ][ - + ]: 6621 : AssertEqual(val.getType(), body.getType());
[ - - ]
138 : 4414 : Node cond = bvl[bvlIndex].eqNode(a[1]);
139 : 2207 : ret = NodeManager::mkNode(Kind::ITE, cond, val, body);
140 : 2207 : }
141 : 2207 : }
142 : 2207 : }
143 [ + - ]: 1841 : else if (a.getKind() == Kind::STORE_ALL)
144 : : {
145 : 1841 : ArrayStoreAll storeAll = a.getConst<ArrayStoreAll>();
146 : 1841 : Node sa = storeAll.getValue();
147 : : // convert the default value recursively (bounded by the number of
148 : : // arguments in bvl)
149 : 1841 : ret = getLambdaForArrayRepresentationRec(sa, bvl, bvlIndex + 1, visited);
150 : 1841 : }
151 : : }
152 : : else
153 : : {
154 : 2361 : ret = a;
155 : : }
156 : 6409 : visited[a] = ret;
157 : 6409 : return ret;
158 : 6409 : }
159 : :
160 : 1143 : Node FunctionConst::getLambdaForArrayRepresentation(TNode a, TNode bvl)
161 : : {
162 [ - + ][ - + ]: 1143 : Assert(a.getType().isArray());
[ - - ]
163 : 1143 : std::unordered_map<TNode, Node> visited;
164 [ + - ]: 2286 : Trace("builtin-rewrite-debug")
165 : 1143 : << "Get lambda for : " << a << ", with variables " << bvl << std::endl;
166 : 2286 : Node body = getLambdaForArrayRepresentationRec(a, bvl, 0, visited);
167 [ + - ]: 1143 : if (!body.isNull())
168 : : {
169 [ + - ]: 2286 : Trace("builtin-rewrite-debug")
170 : 1143 : << "...got lambda body " << body << std::endl;
171 : 1143 : return NodeManager::mkNode(Kind::LAMBDA, bvl, body);
172 : : }
173 [ - - ]: 0 : Trace("builtin-rewrite-debug") << "...failed to get lambda body" << std::endl;
174 : 0 : return Node::null();
175 : 1143 : }
176 : :
177 : 60106 : Node FunctionConst::getArrayRepresentationForLambdaRec(TNode n,
178 : : TypeNode retType)
179 : : {
180 [ - + ][ - + ]: 60106 : Assert(n.getKind() == Kind::LAMBDA);
[ - - ]
181 : 60106 : NodeManager* nm = n.getNodeManager();
182 [ + - ]: 120212 : Trace("builtin-rewrite-debug")
183 : 60106 : << "Get array representation for : " << n << std::endl;
184 : :
185 : 120212 : Node first_arg = n[0][0];
186 : 60106 : Node rec_bvl;
187 : 60106 : size_t size = n[0].getNumChildren();
188 [ + + ]: 60106 : if (size > 1)
189 : : {
190 : 31225 : std::vector<Node> args;
191 [ + + ]: 248469 : for (size_t i = 1; i < size; i++)
192 : : {
193 : 217244 : args.push_back(n[0][i]);
194 : : }
195 : 31225 : rec_bvl = nm->mkNode(Kind::BOUND_VAR_LIST, args);
196 : 31225 : }
197 : :
198 [ + - ]: 60106 : Trace("builtin-rewrite-debug2") << " process body..." << std::endl;
199 : 60106 : std::vector<Node> conds;
200 : 60106 : std::vector<Node> vals;
201 : 60106 : Node curr = n[1];
202 : 60106 : Kind ck = curr.getKind();
203 [ + + ][ + + ]: 121889 : while (ck == Kind::ITE || ck == Kind::OR || ck == Kind::AND
204 [ + + ][ + + ]: 113030 : || ck == Kind::EQUAL || ck == Kind::NOT || ck == Kind::BOUND_VARIABLE)
[ + + ][ + + ]
205 : : {
206 : 36418 : Node index_eq;
207 : 36418 : Node curr_val;
208 : 36418 : Node next;
209 : : // Each iteration of this loop infers an entry in the function, e.g. it
210 : : // has a value under some condition.
211 : :
212 : : // [1] We infer that the entry has value "curr_val" under condition
213 : : // "index_eq". We set "next" to the node that is the remainder of the
214 : : // function to process.
215 [ + + ]: 36418 : if (ck == Kind::ITE)
216 : : {
217 [ + - ]: 12204 : Trace("builtin-rewrite-debug2")
218 [ - + ][ - - ]: 6102 : << " process condition : " << curr[0] << std::endl;
219 : 6102 : index_eq = curr[0];
220 : 6102 : curr_val = curr[1];
221 : 6102 : next = curr[2];
222 : : }
223 [ + + ][ + + ]: 30316 : else if (ck == Kind::OR || ck == Kind::AND)
224 : : {
225 [ + - ]: 33276 : Trace("builtin-rewrite-debug2")
226 : 16638 : << " process base : " << curr << std::endl;
227 : : // Complex Boolean return cases, in which
228 : : // (1) lambda x. (= x v1) v ... becomes
229 : : // lambda x. (ite (= x v1) true [...])
230 : : //
231 : : // (2) lambda x. (not (= x v1)) ^ ... becomes
232 : : // lambda x. (ite (= x v1) false [...])
233 : : //
234 : : // Note the negated cases of the lhs of the OR/AND operators above are
235 : : // handled by pushing the recursion to the then-branch, with the
236 : : // else-branch being the constant value. For example, the negated (1)
237 : : // would be
238 : : // (1') lambda x. (not (= x v1)) v ... becomes
239 : : // lambda x. (ite (= x v1) [...] true)
240 : : // thus requiring the rest of the disjunction to be further processed in
241 : : // the then-branch as the current value.
242 : 16638 : bool pol = curr[0].getKind() != Kind::NOT;
243 : 16638 : bool inverted = (pol == (ck == Kind::AND));
244 [ + + ][ + + ]: 16638 : index_eq = pol ? curr[0] : curr[0][0];
[ - - ]
245 : : // processed : the value that is determined by the first child of curr
246 : : // remainder : the remaining children of curr
247 : 16638 : Node processed, remainder;
248 : : // the value is the polarity of the first child or its inverse if we are
249 : : // in the inverted case
250 [ + + ]: 20486 : processed = nm->mkConst(!inverted ? pol : !pol);
251 : : // build an OR/AND with the remaining components
252 [ + + ]: 16638 : if (curr.getNumChildren() == 2)
253 : : {
254 : 6464 : remainder = curr[1];
255 : : }
256 : : else
257 : : {
258 : 10174 : std::vector<Node> remainderNodes{curr.begin() + 1, curr.end()};
259 : 10174 : remainder = nm->mkNode(ck, remainderNodes);
260 : 10174 : }
261 [ + + ]: 16638 : if (inverted)
262 : : {
263 : 12790 : curr_val = remainder;
264 : 12790 : next = processed;
265 : : // If the lambda contains more variables than the one being currently
266 : : // processed, the current value can be non-constant, since it'll be
267 : : // processed recursively below. Otherwise we fail.
268 [ + + ][ + - ]: 12790 : if (rec_bvl.isNull() && !curr_val.isConst())
[ + + ]
269 : : {
270 [ + - ]: 4452 : Trace("builtin-rewrite-debug2")
271 : 2226 : << "...non-const curr_val " << curr_val << "\n";
272 : 2226 : return Node::null();
273 : : }
274 : : }
275 : : else
276 : : {
277 : 3848 : curr_val = processed;
278 : 3848 : next = remainder;
279 : : }
280 [ + - ]: 14412 : Trace("builtin-rewrite-debug2") << " index_eq : " << index_eq << "\n";
281 [ + - ]: 14412 : Trace("builtin-rewrite-debug2") << " curr_val : " << curr_val << "\n";
282 [ + - ]: 14412 : Trace("builtin-rewrite-debug2") << " next : " << next << std::endl;
283 [ + + ][ + + ]: 33276 : }
284 : : else
285 : : {
286 [ + - ]: 27356 : Trace("builtin-rewrite-debug2")
287 : 13678 : << " process base : " << curr << std::endl;
288 : : // Simple Boolean return cases, in which
289 : : // (1) lambda x. (= x v) becomes lambda x. (ite (= x v) true false)
290 : : // (2) lambda x. x becomes lambda x. (ite (= x true) true false)
291 : : // Note the negateg cases of the bodies above are also handled.
292 : 13678 : bool pol = ck != Kind::NOT;
293 [ + + ]: 13678 : index_eq = pol ? curr : curr[0];
294 : 13678 : curr_val = nm->mkConst(pol);
295 : 13678 : next = nm->mkConst(!pol);
296 : : }
297 : :
298 : : // [2] We ensure that "index_eq" is an equality, if possible.
299 [ + + ]: 34192 : if (index_eq.getKind() != Kind::EQUAL)
300 : : {
301 : 15214 : bool pol = index_eq.getKind() != Kind::NOT;
302 [ + - ]: 15214 : Node indexEqAtom = pol ? index_eq : index_eq[0];
303 [ + + ]: 15214 : if (indexEqAtom.getKind() == Kind::BOUND_VARIABLE)
304 : : {
305 [ + + ]: 10128 : if (!indexEqAtom.getType().isBoolean())
306 : : {
307 : : // Catches default case of non-Boolean variable, e.g.
308 : : // lambda x : Int. x. In this case, it is not canonical and we fail.
309 [ + - ]: 4364 : Trace("builtin-rewrite-debug2")
310 : 2182 : << " ...non-Boolean variable." << std::endl;
311 : 2182 : return Node::null();
312 : : }
313 : : // Boolean argument case, e.g. lambda x. ite( x, t, s ) is processed as
314 : : // lambda x. (ite (= x true) t s)
315 : 7946 : index_eq = indexEqAtom.eqNode(nm->mkConst(pol));
316 : : }
317 : : else
318 : : {
319 : : // non-equality condition
320 [ + - ]: 10172 : Trace("builtin-rewrite-debug2")
321 : 5086 : << " ...non-equality condition." << std::endl;
322 : 5086 : return Node::null();
323 : : }
324 [ + + ]: 15214 : }
325 : :
326 : : // [3] We ensure that "index_eq" is an equality that is equivalent to
327 : : // "first_arg" = "curr_index", where curr_index is a constant, and
328 : : // "first_arg" is the current argument we are processing, if possible.
329 : 26924 : Node curr_index;
330 [ + + ]: 55682 : for (unsigned r = 0; r < 2; r++)
331 : : {
332 : 42967 : Node arg = index_eq[r];
333 : 42967 : Node val = index_eq[1 - r];
334 [ + + ]: 42967 : if (arg == first_arg)
335 : : {
336 : 14209 : curr_index = val;
337 [ + - ]: 28418 : Trace("builtin-rewrite-debug2")
338 : 14209 : << " arg " << arg << " -> " << val << std::endl;
339 : 14209 : break;
340 : : }
341 [ + + ][ + + ]: 57176 : }
342 [ + + ]: 26924 : if (curr_index.isNull())
343 : : {
344 [ + - ]: 25430 : Trace("builtin-rewrite-debug2")
345 : 12715 : << " ...could not infer index value." << std::endl;
346 : : // it could correspond to the default value that does not involve the
347 : : // current argument, hence we break and take curr as the default value
348 : : // below. For example, if we are processing lambda xy. (not y) for x,
349 : : // we have index_eq is (= y true), which does not match for x, hence
350 : : // (not y) is taken as the default value below.
351 : 12715 : break;
352 : : }
353 : :
354 : : // [4] Recurse to ensure that "curr_val" has been normalized w.r.t. the
355 : : // remaining arguments (rec_bvl).
356 [ + + ]: 14209 : if (!rec_bvl.isNull())
357 : : {
358 : 6945 : curr_val = nm->mkNode(Kind::LAMBDA, rec_bvl, curr_val);
359 [ + - ]: 6945 : Trace("builtin-rewrite-debug") << push;
360 [ + - ]: 6945 : Trace("builtin-rewrite-debug2") << push;
361 : 6945 : curr_val = getArrayRepresentationForLambdaRec(curr_val, retType);
362 [ + - ]: 6945 : Trace("builtin-rewrite-debug") << pop;
363 [ + - ]: 6945 : Trace("builtin-rewrite-debug2") << pop;
364 [ + + ]: 6945 : if (curr_val.isNull())
365 : : {
366 [ + - ]: 5180 : Trace("builtin-rewrite-debug2")
367 : 2590 : << " ...failed to recursively find value." << std::endl;
368 : 2590 : return Node::null();
369 : : }
370 : : }
371 [ + - ]: 23238 : Trace("builtin-rewrite-debug2")
372 : 11619 : << " ...condition is index " << curr_val << std::endl;
373 : :
374 : : // [5] Add the entry
375 [ - + ][ - + ]: 11619 : Assert(!curr_index.isNull());
[ - - ]
376 [ - + ][ - + ]: 11619 : Assert(!curr_val.isNull());
[ - - ]
377 [ + + ][ + + ]: 11619 : if (!curr_index.isConst() || !curr_val.isConst())
[ + + ]
378 : : {
379 : : // non-constant value
380 [ + - ]: 3840 : Trace("builtin-rewrite-debug2") << " ...non-constant value for entry\n.";
381 : 3840 : return Node::null();
382 : : }
383 : 7779 : conds.push_back(curr_index);
384 : 7779 : vals.push_back(curr_val);
385 : :
386 : : // we will now process the remainder
387 : 7779 : curr = next;
388 : 7779 : ck = curr.getKind();
389 [ + - ]: 15558 : Trace("builtin-rewrite-debug2")
390 : 7779 : << " process remainder : " << curr << std::endl;
391 [ + + ][ + + ]: 112841 : }
[ + + ][ + + ]
[ + + ][ + + ]
392 [ + + ]: 44182 : if (!rec_bvl.isNull())
393 : : {
394 : 20847 : curr = nm->mkNode(Kind::LAMBDA, rec_bvl, curr);
395 [ + - ]: 20847 : Trace("builtin-rewrite-debug") << push;
396 [ + - ]: 20847 : Trace("builtin-rewrite-debug2") << push;
397 : 20847 : curr = getArrayRepresentationForLambdaRec(curr, retType);
398 [ + - ]: 20847 : Trace("builtin-rewrite-debug") << pop;
399 [ + - ]: 20847 : Trace("builtin-rewrite-debug2") << pop;
400 : : }
401 [ + + ][ + + ]: 44182 : if (!curr.isNull() && curr.isConst())
[ + + ]
402 : : {
403 : : // compute the return type
404 : 11795 : TypeNode array_type = retType;
405 [ + + ]: 37011 : for (size_t i = 0; i < size; i++)
406 : : {
407 : 25216 : size_t index = (size - 1) - i;
408 : 25216 : array_type = nm->mkArrayType(n[0][index].getType(), array_type);
409 : : }
410 [ + - ]: 23590 : Trace("builtin-rewrite-debug2")
411 [ - + ][ - - ]: 11795 : << " make array store all " << curr.getType()
412 : 11795 : << " annotated : " << array_type << " from " << curr << std::endl;
413 [ - + ][ - + ]: 11795 : Assert(curr.getType() == array_type.getArrayConstituentType());
[ - - ]
414 : 11795 : curr = nm->mkConst(ArrayStoreAll(array_type, curr));
415 [ + - ]: 11795 : Trace("builtin-rewrite-debug2") << " build array..." << std::endl;
416 : : // can only build if default value is constant (since array store all must
417 : : // be constant)
418 [ + - ]: 23590 : Trace("builtin-rewrite-debug2")
419 : 11795 : << " got constant base " << curr << std::endl;
420 [ + - ]: 11795 : Trace("builtin-rewrite-debug2") << " conditions " << conds << std::endl;
421 [ + - ]: 11795 : Trace("builtin-rewrite-debug2") << " values " << vals << std::endl;
422 : : // construct store chain
423 [ + + ]: 19398 : for (size_t i = 0, numCond = conds.size(); i < numCond; i++)
424 : : {
425 : 7603 : size_t ii = (numCond - 1) - i;
426 [ - + ][ - + ]: 22809 : AssertEqual(conds[ii].getType(), first_arg.getType());
[ - - ]
427 : 7603 : curr = nm->mkNode(Kind::STORE, curr, conds[ii], vals[ii]);
428 : : // normalize it using the array rewriter utility, which must be done at
429 : : // each iteration of this loop
430 : 7603 : curr = arrays::TheoryArraysRewriter::normalizeConstant(nm, curr);
431 : : }
432 [ + - ]: 23590 : Trace("builtin-rewrite-debug")
433 : 11795 : << "...got array " << curr << " for " << n << std::endl;
434 : 11795 : return curr;
435 : 11795 : }
436 [ + - ]: 64774 : Trace("builtin-rewrite-debug")
437 : 0 : << "...failed to get array (cannot get constant default value)"
438 : 32387 : << std::endl;
439 : 32387 : return Node::null();
440 : 63954 : }
441 : :
442 : 32314 : Node FunctionConst::toArrayConst(TNode n)
443 : : {
444 : 32314 : Kind nk = n.getKind();
445 [ - + ]: 32314 : if (nk == Kind::FUNCTION_ARRAY_CONST)
446 : : {
447 : 0 : const FunctionArrayConst& fc = n.getConst<FunctionArrayConst>();
448 : 0 : return fc.getArrayValue();
449 : : }
450 [ + - ]: 32314 : else if (nk == Kind::LAMBDA)
451 : : {
452 : : // must carry the overall return type to deal with cases like (lambda ((x
453 : : // Int) (y Int)) (ite (= x _) 0.5 0.0)), where the inner construction for
454 : : // the else case above should be (arraystoreall (Array Int Real) 0.0)
455 : 32314 : return getArrayRepresentationForLambdaRec(n, n[1].getType());
456 : : }
457 : 0 : return Node::null();
458 : : }
459 : :
460 : : } // namespace uf
461 : : } // namespace theory
462 : : } // namespace cvc5::internal
|