206 lines
4.8 KiB
Lua
206 lines
4.8 KiB
Lua
require "functools"
|
|
|
|
-- itertools
|
|
-- iterators API, inspired by python's itertools module.
|
|
|
|
-- The API can used indifferently iterators (implemented as LUA's coroutines) or
|
|
-- table values for iterating.
|
|
|
|
-- Warnings: this code is really unsafe on function returning multiple results.
|
|
-- It is suggest to pack the result using op_pack and unpack them when required.
|
|
|
|
-- local function that try to check if the method is a valid iterator.
|
|
local function checkit(it)
|
|
if not ((it == nil) or
|
|
(type(it) == "table") or
|
|
(type(it) == "function")) then
|
|
error("Invalid iterator: <" .. type(it) .. ">")
|
|
end
|
|
end
|
|
|
|
-- return a iterator for which the method f is apply to all elements of it.
|
|
function map(f, it)
|
|
checkit(it)
|
|
local function next ()
|
|
for v in iter(it) do
|
|
coroutine.yield(f(v))
|
|
end
|
|
end
|
|
return coroutine.wrap(next)
|
|
end
|
|
|
|
-- reduce all elements from it by combining them using f, starting with init as first element.
|
|
function reduce(f, init, it)
|
|
checkit(it)
|
|
local acc = init
|
|
for v in iter(it) do
|
|
acc = f(acc, v)
|
|
end
|
|
return acc
|
|
end
|
|
|
|
-- return all elements for which p(element) return true.
|
|
function filter(p, it)
|
|
checkit(it)
|
|
local function next ()
|
|
for v in iter(it) do
|
|
if p(v) then
|
|
coroutine.yield(v)
|
|
end
|
|
end
|
|
end
|
|
return coroutine.wrap(next)
|
|
end
|
|
|
|
-- Concat all iterators into one, consuming them one after each other.
|
|
function concat(it, ...)
|
|
checkit(it)
|
|
map(checkit, unpack(arg))
|
|
if #arg == 0 then
|
|
return iter(it)
|
|
end
|
|
local function next()
|
|
for v in iter(it) do
|
|
coroutine.yield(v)
|
|
end
|
|
if #arg > 0 then
|
|
for v in concat(unpack(arg)) do
|
|
coroutine.yield(v)
|
|
end
|
|
end
|
|
end
|
|
return coroutine.wrap(next)
|
|
end
|
|
|
|
-- Same as concat, but used an iterator of iterators instead.
|
|
function concat_from_iterable(iterables)
|
|
return reduce(concat, nil, iterables)
|
|
end
|
|
|
|
-- Return all unique elements from it. This consumed as much memory
|
|
-- as the number of elements returned.
|
|
function unique(it)
|
|
checkit(it)
|
|
local function next()
|
|
local t = {}
|
|
for v in iter(it) do
|
|
if not t[v] then
|
|
coroutine.yield(v)
|
|
t[v] = true
|
|
end
|
|
end
|
|
end
|
|
return coroutine.wrap(next)
|
|
end
|
|
|
|
-- collect all elements of iterators into a list and return it.
|
|
function collect(...)
|
|
local t = {}
|
|
for v in concat(unpack(arg)) do
|
|
table.insert(t, v)
|
|
end
|
|
return t
|
|
end
|
|
|
|
-- Return true if table t contains element elem
|
|
-- If elem is a pair {key = value}, returns true if t[key] == value
|
|
function contains(t, elem)
|
|
for k, v in iteritems(t) do
|
|
if type(elem) == "table" then
|
|
for key, value in pairs(elem) do
|
|
return t[key] == value
|
|
end
|
|
else
|
|
if type(k) == "number" and t[k] == elem then
|
|
return true
|
|
end
|
|
end
|
|
end
|
|
return false
|
|
end
|
|
-- Return true if any elements is true (short-circuit)
|
|
function any(it)
|
|
checkit(it)
|
|
for v in iter(it) do
|
|
if v then
|
|
return true
|
|
end
|
|
end
|
|
return false
|
|
end
|
|
|
|
-- Return true if all elements are true (short-circuit at first false)
|
|
function all(it)
|
|
checkit(it)
|
|
return not any(map(op_not, it))
|
|
end
|
|
|
|
-- Return true if no elements are true (short-circuit on the first true)
|
|
function none(it)
|
|
checkit(it)
|
|
return not any(it)
|
|
end
|
|
|
|
-- Iterator that yield a pair of key and value as an arrayy.
|
|
function iteritems(t)
|
|
assert(type(t) == "table")
|
|
local function next ()
|
|
for i, v in pairs(t) do
|
|
coroutine.yield({i, v})
|
|
end
|
|
end
|
|
return coroutine.wrap(next)
|
|
end
|
|
|
|
-- Return the first value of t (a key from iteritems).
|
|
function getkey(t)
|
|
return t[1]
|
|
end
|
|
|
|
-- Return the second value of t (a value from iteritems).
|
|
function getvalue(t)
|
|
return t[2]
|
|
end
|
|
|
|
-- Iterate only over the values of a table (in no specific orders)
|
|
itervalues = compose(partial(map,getvalue),iteritems)
|
|
|
|
-- Iterate only over the keys of a table
|
|
iterkeys = compose(partial(map,getkey),iteritems)
|
|
|
|
-- array are table where a.n is equal to table.getn(a)
|
|
-- it's the case for arg, in particular, but not {}
|
|
function isarray(a)
|
|
return type(a) == "table" and table.getn(a) == a.n
|
|
end
|
|
|
|
-- iterated on digit key ordering (with ipairs).
|
|
function iterarray(a)
|
|
assert(type(a) == "table")
|
|
local function next ()
|
|
for _, v in ipairs(a) do
|
|
coroutine.yield(v)
|
|
end
|
|
end
|
|
return coroutine.wrap(next)
|
|
end
|
|
|
|
-- Return an empty iterator
|
|
function nulliter(f)
|
|
return coroutine.wrap(function () return nil end)
|
|
end
|
|
|
|
-- Return a safe iterator based on the type passed as argument.
|
|
function iter(it)
|
|
if it == nil then
|
|
return nulliter()
|
|
elseif type(it) == "table" then
|
|
if isarray(it) then
|
|
return iterarray(it)
|
|
else
|
|
return itervalues(it)
|
|
end
|
|
else
|
|
return it
|
|
end
|
|
end
|