| 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) |