main
py 151 lines 4.17 KB
Raw
1 import asyncio
2 import threading
3 import uuid
4 import weakref
5
6 import pytest
7
8 from helpers.defer import DeferredTask
9
10
11 class Owner:
12 pass
13
14
15 def make_task() -> DeferredTask:
16 return DeferredTask(f"defer-lifecycle-{uuid.uuid4()}")
17
18
19 def test_completed_task_releases_call_references_and_children():
20 task = make_task()
21 owner = Owner()
22 owner_ref = weakref.ref(owner)
23 child_killed = threading.Event()
24
25 class Child:
26 def kill(self, terminate_thread: bool = False) -> None:
27 assert terminate_thread
28 child_killed.set()
29
30 async def run(captured_owner):
31 return "done"
32
33 try:
34 task.add_child_task(Child(), terminate_thread=True) # type: ignore[arg-type]
35 task.start_task(run, owner)
36 assert task.result_sync(timeout=2) == "done"
37 assert child_killed.wait(2)
38 assert task.func is None
39 assert task.args == ()
40 assert task.kwargs == {}
41
42 del owner
43 assert owner_ref() is None
44 assert task.result_sync(timeout=2) == "done"
45 with pytest.raises(RuntimeError, match="Completed task cannot be restarted"):
46 task.restart()
47 finally:
48 task.kill(terminate_thread=True)
49
50
51 def test_run_task_end_extension_marks_state_dirty_after_completion(monkeypatch):
52 from extensions.python._functions.agent.AgentContext.run_task.end import (
53 _10_mark_state_dirty as task_done_extension,
54 )
55
56 task = make_task()
57 callback_called = threading.Event()
58 observations: list[tuple[str | None, bool]] = []
59
60 def mark_dirty(*, reason=None):
61 observations.append((reason, bool(task.is_alive())))
62 callback_called.set()
63
64 monkeypatch.setattr(
65 task_done_extension,
66 "mark_dirty_all",
67 mark_dirty,
68 )
69
70 async def run():
71 return "done"
72
73 try:
74 with pytest.raises(RuntimeError, match="Task hasn't been started"):
75 task.add_done_callback(lambda _future: None)
76 task.start_task(run)
77 task_done_extension.MarkStateDirty(agent=None).execute(
78 data={"result": task}
79 )
80 assert task.result_sync(timeout=2) == "done"
81 assert callback_called.wait(2)
82 assert observations == [("agent.AgentContext.run_task_done", False)]
83 finally:
84 task.kill(terminate_thread=True)
85
86
87 def test_kill_clears_stored_call_without_clearing_running_arguments():
88 task = make_task()
89 owner = Owner()
90 owner_ref = weakref.ref(owner)
91 started = threading.Event()
92 cancelled = threading.Event()
93 finished = threading.Event()
94 release: list[asyncio.Event] = []
95
96 async def run(captured_owner):
97 release.append(asyncio.Event())
98 started.set()
99 try:
100 await asyncio.Future()
101 except asyncio.CancelledError:
102 cancelled.set()
103 await release[0].wait()
104 finally:
105 finished.set()
106
107 try:
108 task.start_task(run, owner)
109 assert started.wait(2)
110 task.kill()
111 assert cancelled.wait(2)
112 assert task.func is None
113 assert task.args == ()
114 assert task.kwargs == {}
115
116 del owner
117 assert owner_ref() is not None
118 task.event_loop_thread.loop.call_soon_threadsafe(release[0].set)
119 assert finished.wait(2)
120 asyncio.run_coroutine_threadsafe(
121 asyncio.sleep(0), task.event_loop_thread.loop
122 ).result(2)
123 assert owner_ref() is None
124 finally:
125 if release and task.event_loop_thread.loop:
126 task.event_loop_thread.loop.call_soon_threadsafe(release[0].set)
127 task.kill(terminate_thread=True)
128
129
130 def test_active_task_can_restart_from_its_snapshot():
131 task = make_task()
132 starts = [threading.Event(), threading.Event()]
133 run_count = 0
134
135 async def run(value):
136 nonlocal run_count
137 current_run = run_count
138 run_count += 1
139 assert value == "argument"
140 starts[current_run].set()
141 await asyncio.Future()
142
143 try:
144 task.start_task(run, "argument")
145 assert starts[0].wait(2)
146 task.restart()
147 assert starts[1].wait(2)
148 assert task.func is run
149 assert task.args == ("argument",)
150 finally:
151 task.kill(terminate_thread=True)