Skip to content

Commit 690a237

Browse files
committed
feat: implement python wrapper for the iterator.
1 parent 0b80208 commit 690a237

3 files changed

Lines changed: 72 additions & 0 deletions

File tree

apyds/chain_t.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -78,3 +78,20 @@ def execute(self, callback: typing.Callable[[Rule], bool]) -> int:
7878
The number of rules processed.
7979
"""
8080
return self._chain.execute(lambda candidate: callback(Rule(candidate.clone())))
81+
82+
def __iter__(self) -> typing.Iterator[Rule]:
83+
"""Iterate over inferred rules.
84+
85+
Returns:
86+
An iterator over Rule objects.
87+
88+
Example:
89+
>>> for rule in chain:
90+
... print(rule)
91+
"""
92+
iterator = self._chain.iter()
93+
while True:
94+
candidate = iterator.next()
95+
if candidate is None:
96+
break
97+
yield Rule(candidate.clone())

apyds/ds.cc

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,31 @@
66

77
namespace py = pybind11;
88

9+
class Iterator {
10+
public:
11+
explicit Iterator(std::generator<ds::rule_t*> _generator) : generator(std::move(_generator)), initialized(false), iterator(nullptr) { }
12+
13+
ds::rule_t* next() {
14+
if (initialized) {
15+
++*iterator;
16+
} else {
17+
iterator = std::make_unique<iterator_t>(generator.begin());
18+
initialized = true;
19+
}
20+
if (*iterator == generator.end()) {
21+
return nullptr;
22+
}
23+
ds::rule_t* result = **iterator;
24+
return result;
25+
}
26+
27+
private:
28+
std::generator<ds::rule_t*> generator;
29+
bool initialized;
30+
using iterator_t = decltype(generator.begin());
31+
std::unique_ptr<iterator_t> iterator;
32+
};
33+
934
template<typename T>
1035
auto from_string(const std::string_view& string, int buffer_size) -> std::unique_ptr<T> {
1136
auto result = reinterpret_cast<T*>(operator new(buffer_size));
@@ -169,6 +194,11 @@ PYBIND11_MODULE(_ds, m, py::mod_gil_not_used()) {
169194
search_t.def("reset", &ds::search_t::reset);
170195
search_t.def("add", &ds::search_t::add);
171196
search_t.def("execute", &ds::search_t::execute);
197+
search_t.def(
198+
"iter",
199+
[](ds::search_t& self) { return Iterator(std::move(self.iterator())); },
200+
py::return_value_policy::reference_internal
201+
);
172202

173203
auto chain_t = py::class_<ds::chain_t>(m, "Chain");
174204
chain_t.def(py::init<ds::length_t, ds::length_t>());
@@ -177,4 +207,12 @@ PYBIND11_MODULE(_ds, m, py::mod_gil_not_used()) {
177207
chain_t.def("reset", &ds::chain_t::reset);
178208
chain_t.def("add", &ds::chain_t::add);
179209
chain_t.def("execute", &ds::chain_t::execute);
210+
chain_t.def(
211+
"iter",
212+
[](ds::chain_t& self) { return Iterator(std::move(self.iterator())); },
213+
py::return_value_policy::reference_internal
214+
);
215+
216+
auto iterator_t = py::class_<Iterator>(m, "Iterator");
217+
iterator_t.def("next", &Iterator::next, py::return_value_policy::reference_internal);
180218
}

apyds/search_t.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,3 +77,20 @@ def execute(self, callback: typing.Callable[[Rule], bool]) -> int:
7777
The number of rules processed.
7878
"""
7979
return self._search.execute(lambda candidate: callback(Rule(candidate.clone())))
80+
81+
def __iter__(self) -> typing.Iterator[Rule]:
82+
"""Iterate over inferred rules.
83+
84+
Returns:
85+
An iterator over Rule objects.
86+
87+
Example:
88+
>>> for rule in search:
89+
... print(rule)
90+
"""
91+
iterator = self._search.iter()
92+
while True:
93+
candidate = iterator.next()
94+
if candidate is None:
95+
break
96+
yield Rule(candidate.clone())

0 commit comments

Comments
 (0)