diff --git a/src/graph.luau b/src/graph.luau index bb42fbb..0c0bf21 100644 --- a/src/graph.luau +++ b/src/graph.luau @@ -181,47 +181,80 @@ local function evaluate_node(node: Node) return cur_value ~= new_value -- node has changed value end -local function update_from(node: StartNode, n0: number) - if not node[1] then return end +-- local function update_from(node: StartNode, n0: number) +-- if not node[1] then return end - local n = n0 +-- local n = n0 - -- unparent all children and queue for eval - do - local child = node[1] - while child do - --assert(child.parents.owner) - unparent(child) - n += 1 - update_queue[n] = child - child = node[1] - end +-- -- unparent all children and queue for eval +-- do +-- local child = node[1] +-- while child do +-- --assert(child.parents.owner) +-- unparent(child) +-- n += 1 +-- update_queue[n] = child +-- child = node[1] +-- end +-- end + +-- update_queue.n = n + +-- -- evaluate all queued children +-- for i = n0 + 1, n do +-- local child = update_queue[i] +-- assert(type(child.effect) == "function") + +-- if evaluate_node(child) then +-- update_from(child, n) +-- end + +-- update_queue[i] = false :: any -- false instead of nil to avoid sparse +-- end + +-- update_queue.n = n0 +-- end + +-- local function update(node: StartNode) +-- update_from(node, update_queue.n) +-- end + +local function queue_children(node: StartNode) + local i = update_queue.n + local child = node[1] + while child do + --assert(child.parents.owner) + unparent(child) + i += 1 + update_queue[i] = child + child = node[1] end + update_queue.n = i +end - update_queue.n = n +local function update(root: StartNode) + local n0 = update_queue.n + queue_children(root) - -- evaluate all queued children - for i = n0 + 1, n do - local child = update_queue[i] - assert(type(child.effect) == "function") + local i = n0 + 1 + while i <= update_queue.n do + local node = update_queue[i] + assert(node.effect) - if evaluate_node(child) then - update_from(child, n) + if evaluate_node(node) then + queue_children(node) end update_queue[i] = false :: any -- false instead of nil to avoid sparse + i += 1 end update_queue.n = n0 end -local function update(node: StartNode) - update_from(node, update_queue.n) -end - local function track(node: StartNode) local scope = get_scope() - if scope and type(scope.effect) == "function" then -- do not track nodes with no effect + if scope and scope.effect then -- do not track nodes with no effect add_child(node, scope) end end diff --git a/test/tests.luau b/test/tests.luau index a82712b..0c72f90 100644 --- a/test/tests.luau +++ b/test/tests.luau @@ -1,5 +1,5 @@ local testkit = require("test/testkit") -local TEST, CASE, CHECK, FINISH = testkit.test() +local TEST, CASE, CHECK, FINISH, SKIP = testkit.test() local mock = require "test/mock" local Instance, Signal = mock.Instance, mock.Signal @@ -33,6 +33,8 @@ local NIL = nil :: any vide.strict = false +--SKIP "graph edge cases" + TEST("graph", function() local create_node = graph.create_node local track = graph.track @@ -1907,7 +1909,85 @@ TEST("nested effects cases", function() root(App) end) -vide.strict = true +TEST("graph edge cases", wrap_root(function() + local source = vide.source + local derive = vide.derive + local effect = vide.effect + + do CASE "diamond A,B,C,D" + --[[ + + a > b > d + > c > + + ]] + + local a = source(0) + + local b = derive(function() return (a() % 2 == 0) and 1 or 0 end) + local c = derive(function() return a() * 2 end) + local d = derive(function() return b() + c() end) + + local count = { b = 0, c = 0, d = 0 } + effect(function() b(); count.b += 1 end) + effect(function() c(); count.c += 1 end) + effect(function() d(); count.d += 1 end) + + a(1) + CHECK(count.b == 2) + CHECK(count.c == 2) + CHECK(count.d == 2) + CHECK(d() == 2) + + a(3) + CHECK(count.b == 2) + CHECK(count.c == 3) + CHECK(count.d == 3) + CHECK(d() == 6) + end + + do CASE "diamond A,B,C,D,E" + --[[ + + a > b > > e + > c > d > + + ]] + + local a = source(0) + + local b = derive(function() print "ran b"; return (a() % 2 == 0) and 1 or 0 end) + local c = derive(function() print "ran c"; return a() * 2 end) + local d = derive(function() print "ran d"; return c() * 2 end) + local e = derive(function() print "ran e"; return b() + c() end) + + local count = { b = 0, c = 0, d = 0, e = 0 } + effect(function() b(); count.b += 1 end) + effect(function() c(); count.c += 1 end) + effect(function() d(); count.d += 1 end) + effect(function() d(); count.e += 1 end) + + a(1) + + print(e()) + CHECK(count.b == 2) + CHECK(count.c == 2) + CHECK(count.d == 2) + CHECK(count.e == 2) + CHECK(e() == 4) -- todo: solve e evaluating before d + + a(3) + CHECK(count.b == 2) + CHECK(count.c == 3) + CHECK(count.d == 3) + CHECK(count.e == 3) + CHECK(e() == 12) + end + + do CASE "repeated read" + + end +end)) TEST("strict", wrap_root(function() vide.strict = true