3
0
Fork 0
mirror of https://github.com/Z3Prover/z3 synced 2026-08-02 12:13:25 +00:00

term_enumeration: add tuple iterator over a vector of sorts

Expose enum_tuples(sorts) which produces an iterator over vectors of
terms, one term per input sort, dovetailing the per-sort streams so all
combinations are enumerated even when individual streams are infinite.

Each sort is enumerated by a self-contained sort_stream owning its own
grammar and bottom_up_enumerator seeded from the user productions, so
array sorts (with their fresh select ops and bound vars) do not collide
across sorts.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: cb1e958b-f89b-4407-958d-8e7ecf172bbc
This commit is contained in:
Nikolaj Bjorner 2026-07-28 14:47:54 -07:00
parent 5bd9e6a009
commit 1c89937473
3 changed files with 354 additions and 1 deletions

View file

@ -488,6 +488,106 @@ private:
}
};
// ============================================================================
// sort_stream - self-contained enumeration of terms of a single sort
// ============================================================================
/**
* A sort_stream owns a private grammar and bottom_up_enumerator so that
* several streams can be advanced independently and concurrently (as needed
* by the tuple iterator). It is seeded from the user-configured productions
* and, for array sorts, augments the grammar with fresh select operators and
* bound variables, wrapping enumerated bodies into lambdas.
*/
class sort_stream {
public:
sort_stream(ast_manager& m, func_decl_ref_vector const& funcs, expr_ref_vector const& exprs, sort* s)
: m(m), m_grammar(m), m_enum(m_grammar), autil(m), m_sort(s), m_current(m), m_pinned(m) {
for (func_decl* f : funcs)
m_grammar.add_func_decl(f);
for (expr* e : exprs)
m_grammar.add_expr(e);
m_enum.reset();
init_sort();
advance();
}
// Return the current term (already lambda-wrapped) and advance to the next
// one. Returns nullptr once the stream is exhausted. Returned terms remain
// pinned for the lifetime of the stream.
expr* next() {
if (m_end)
return nullptr;
expr* r = m_current.get();
advance();
return r;
}
private:
ast_manager& m;
grammar m_grammar;
bottom_up_enumerator m_enum;
array_util autil;
sort* m_sort;
expr_ref m_current;
expr_ref_vector m_pinned;
bool m_end = false;
vector<expr_ref_vector> m_vars;
vector<ptr_vector<sort>> m_decls;
vector<vector<symbol>> m_names;
void init_sort() {
sort* range = m_sort;
while (autil.is_array(range)) {
m_vars.push_back(expr_ref_vector(m));
m_decls.push_back(ptr_vector<sort>());
m_names.push_back(vector<symbol>());
for (unsigned i = 0; i < get_array_arity(range); ++i) {
m_decls.back().push_back(get_array_domain(range, i));
m_vars.back().push_back(nullptr);
m_names.back().push_back(symbol());
}
expr_ref_vector args(m);
args.push_back(m.mk_const("a", range));
for (unsigned i = 0; i < m_decls.back().size(); ++i)
args.push_back(m.mk_var(i, m_decls.back().get(i)));
app_ref sel(autil.mk_select(args), m);
m_grammar.add_func_decl(sel->get_decl());
range = get_array_range(range);
}
unsigned n = 0;
for (unsigned i = m_decls.size(); i-- > 0;) {
for (unsigned j = m_decls[i].size(); j-- > 0;) {
m_vars[i][j] = m.mk_var(n, m_decls[i][j]);
m_names[i][j] = symbol(n);
m_grammar.add_expr(m_vars[i].get(j));
n++;
}
}
m_sort = range;
m_enum.set_target_sort(range);
}
void mk_lambda() {
if (!m_current)
return;
for (unsigned i = m_decls.size(); i-- > 0;)
m_current = m.mk_lambda(m_decls[i].size(), m_decls[i].data(), m_names[i].data(), m_current);
}
void advance() {
if (m_end)
return;
m_current = m_enum.next();
SASSERT(!m_current || m_current->get_sort() == m_sort);
mk_lambda();
if (!m_current)
m_end = true;
else
m_pinned.push_back(m_current);
}
};
} // namespace term_enum
// ============================================================================
@ -499,16 +599,21 @@ struct term_enumeration::imp {
term_enum::grammar m_grammar;
term_enum::bottom_up_enumerator m_bottom_up_enumerator;
std::function<unsigned(expr*)> m_cost;
func_decl_ref_vector m_user_funcs;
expr_ref_vector m_user_exprs;
imp(ast_manager& m) :
m(m), m_grammar(m), m_bottom_up_enumerator(m_grammar) {}
m(m), m_grammar(m), m_bottom_up_enumerator(m_grammar),
m_user_funcs(m), m_user_exprs(m) {}
void add_production(func_decl* f) {
m_grammar.add_func_decl(f);
m_user_funcs.push_back(f);
}
void add_production(expr* e) {
m_grammar.add_expr(e);
m_user_exprs.push_back(e);
}
void set_cost(std::function<unsigned(expr*)> const& cost) {
@ -650,6 +755,157 @@ term_enumeration::iterator term_enumeration::terms::end() {
return iterator(nullptr);
}
// -- tuple iterator implementation --
struct term_enumeration::tuple_iterator::timp {
imp& m_imp;
ast_manager& m;
unsigned m_n;
scoped_ptr_vector<term_enum::sort_stream> m_streams;
vector<expr_ref_vector> m_terms; // materialized terms per dimension
svector<bool> m_done; // per-dimension exhausted flag
unsigned m_rr = 0; // round-robin cursor
vector<expr_ref_vector> m_buffer; // pending tuples
unsigned m_buf_idx = 0;
expr_ref_vector m_current;
bool m_end = false;
bool m_dead = false;
timp(imp& i, unsigned n, sort* const* sorts) :
m_imp(i), m(i.m), m_n(n), m_current(i.m) {
for (unsigned k = 0; k < n; ++k) {
m_streams.push_back(alloc(term_enum::sort_stream, m, i.m_user_funcs, i.m_user_exprs, sorts[k]));
m_terms.push_back(expr_ref_vector(m));
m_done.push_back(false);
}
if (n == 0) {
m_end = true;
return;
}
advance();
}
// Emit all tuples where dimension d is the fresh term t and every other
// dimension ranges over its currently materialized terms.
void emit(unsigned d, expr* t) {
for (unsigned j = 0; j < m_n; ++j)
if (j != d && m_terms[j].empty())
return;
svector<unsigned> idx;
idx.resize(m_n, 0);
while (true) {
expr_ref_vector tup(m);
for (unsigned j = 0; j < m_n; ++j)
tup.push_back(j == d ? t : m_terms[j].get(idx[j]));
m_buffer.push_back(tup);
bool carried = false;
for (unsigned j = m_n; j-- > 0;) {
if (j == d)
continue;
idx[j]++;
if (idx[j] < m_terms[j].size()) {
carried = true;
break;
}
idx[j] = 0;
}
if (!carried)
break;
}
}
// Advance one dimension. Returns false if every dimension is exhausted.
bool step() {
for (unsigned tries = 0; tries < m_n; ++tries) {
unsigned d = m_rr;
m_rr = (m_rr + 1) % m_n;
if (m_done[d])
continue;
expr* t = m_streams[d]->next();
if (!t) {
m_done[d] = true;
if (m_terms[d].empty())
m_dead = true; // this dimension yields no terms at all
return true;
}
emit(d, t);
m_terms[d].push_back(t);
return true;
}
return false;
}
bool fill() {
while (m_buf_idx >= m_buffer.size()) {
if (m_dead)
return false;
if (!step())
return false;
}
return true;
}
void advance() {
if (m_end)
return;
if (m_buf_idx >= m_buffer.size()) {
m_buffer.reset();
m_buf_idx = 0;
if (!fill()) {
m_end = true;
m_current.reset();
return;
}
}
m_current.reset();
m_current.append(m_buffer[m_buf_idx]);
m_buf_idx++;
}
};
term_enumeration::tuple_iterator::tuple_iterator(imp& i, unsigned n, sort* const* sorts) {
m_imp = alloc(timp, i, n, sorts);
}
term_enumeration::tuple_iterator::tuple_iterator(std::nullptr_t) {
m_imp = nullptr;
}
term_enumeration::tuple_iterator::~tuple_iterator() {
dealloc(m_imp);
}
expr_ref_vector term_enumeration::tuple_iterator::operator*() {
SASSERT(m_imp);
return m_imp->m_current;
}
term_enumeration::tuple_iterator& term_enumeration::tuple_iterator::operator++() {
if (m_imp)
m_imp->advance();
return *this;
}
bool term_enumeration::tuple_iterator::operator==(tuple_iterator const& other) const {
if (!m_imp && !other.m_imp) return true;
if (!m_imp) return other.m_imp->m_end;
if (!other.m_imp) return m_imp->m_end;
return m_imp->m_end == other.m_imp->m_end;
}
// -- tuples implementation --
term_enumeration::tuples::tuples(imp* i, unsigned n, sort* const* sorts) :
m_imp(i), m_sorts(n, sorts) {}
term_enumeration::tuple_iterator term_enumeration::tuples::begin() {
return tuple_iterator(*m_imp, m_sorts.size(), m_sorts.data());
}
term_enumeration::tuple_iterator term_enumeration::tuples::end() {
return tuple_iterator(nullptr);
}
// -- term_enumeration implementation --
term_enumeration::term_enumeration(ast_manager& m) {
@ -676,6 +932,14 @@ term_enumeration::terms term_enumeration::enum_terms(sort* s) {
return terms(m_imp, s);
}
term_enumeration::tuples term_enumeration::enum_tuples(unsigned n, sort* const* sorts) {
return tuples(m_imp, n, sorts);
}
term_enumeration::tuples term_enumeration::enum_tuples(sort_ref_vector const& sorts) {
return tuples(m_imp, sorts.size(), sorts.data());
}
std::ostream& term_enumeration::display(std::ostream& out) const {
return m_imp->display(out);
}

View file

@ -46,5 +46,36 @@ public:
terms enum_terms(sort* s);
// -- tuple enumeration --
// Iterate over vectors of terms, one term per input sort. Produces all
// combinations (dovetailed, since individual streams may be infinite).
class tuple_iterator {
struct timp;
timp* m_imp;
public:
tuple_iterator(imp& i, unsigned n, sort* const* sorts);
tuple_iterator(std::nullptr_t);
~tuple_iterator();
expr_ref_vector operator*();
tuple_iterator& operator++();
bool operator!=(tuple_iterator const& other) const {
return !(*this == other);
}
bool operator==(tuple_iterator const& other) const;
};
class tuples {
imp* m_imp;
ptr_vector<sort> m_sorts;
public:
tuples(imp* i, unsigned n, sort* const* sorts);
tuple_iterator begin();
tuple_iterator end();
};
tuples enum_tuples(unsigned n, sort* const* sorts);
tuples enum_tuples(sort_ref_vector const& sorts);
std::ostream& display(std::ostream& out) const;
};

View file

@ -297,6 +297,63 @@ static void tst_nested_array_enumeration() {
te.display(std::cout);
}
static void tst_tuple_enumeration() {
std::cout << "=== test tuple enumeration ===\n";
ast_manager m;
reg_decl_plugins(m);
arith_util a(m);
array_util arr(m);
term_enumeration te(m);
// Leaves: a boolean constant, integer constants, and an integer array.
expr_ref bt(m.mk_true(), m);
expr_ref bf(m.mk_false(), m);
expr_ref zero(a.mk_int(0), m);
expr_ref one(a.mk_int(1), m);
te.add_production(bt);
te.add_production(bf);
te.add_production(zero);
te.add_production(one);
sort* int_sort = a.mk_int();
sort_ref arr_sort(arr.mk_array_sort(int_sort, int_sort), m);
app_ref ca(arr.mk_const_array(arr_sort, zero), m);
te.add_production(ca.get());
// Operators over integers so that streams are non-trivial.
app_ref tmp_add(a.mk_add(zero, one), m);
te.add_production(tmp_add->get_decl());
sort_ref_vector sorts(m);
sorts.push_back(m.mk_bool_sort());
sorts.push_back(int_sort);
sorts.push_back(arr_sort);
unsigned count = 0;
obj_hashtable<expr> bools, ints, arrays;
for (expr_ref_vector const& tup : te.enum_tuples(sorts)) {
ENSURE(tup.size() == 3);
ENSURE(tup[0]->get_sort() == m.mk_bool_sort());
ENSURE(tup[1]->get_sort() == int_sort);
ENSURE(tup[2]->get_sort() == arr_sort.get());
bools.insert(tup[0]);
ints.insert(tup[1]);
arrays.insert(tup[2]);
if (count < 10)
std::cout << " (" << mk_pp(tup[0], m) << ", " << mk_pp(tup[1], m)
<< ", " << mk_pp(tup[2], m) << ")\n";
count++;
if (count >= 30) break;
}
ENSURE(count >= 4);
// Verify combinations mix distinct terms from each component stream.
ENSURE(bools.size() >= 2);
ENSURE(ints.size() >= 2);
std::cout << "Enumerated " << count << " tuples\n";
}
void tst_term_enumeration() {
tst_basic_enumeration();
tst_enumeration_with_operators();
@ -305,5 +362,6 @@ void tst_term_enumeration() {
tst_bitvector_enumeration();
tst_multiple_sorts();
tst_nested_array_enumeration();
tst_tuple_enumeration();
std::cout << "All term_enumeration tests passed!\n";
}