Files
EvoScientist-Multi/EvoScientist/langgraph_dev/worker_exit.py
T

236 lines
9.1 KiB
Python

"""Process-local exit receipts for the existing LangGraph worker.
No receipt survives a restart. Missing receipts never prove execution absent.
The admission gate also covers a queued worker that has not started yet.
"""
from __future__ import annotations
import asyncio
from concurrent.futures import Future
from contextvars import ContextVar
from functools import wraps
from threading import RLock
from typing import Any, cast
_lock = RLock()
_runs: dict[tuple[str, str], dict] = {}
_installed = False
_owned: ContextVar[bool] = ContextVar("recoverable_worker_owned", default=False)
_evidence: ContextVar[dict | None] = ContextVar("worker_exit_evidence", default=None)
class _ExitBoundary:
def __init__(self, context, *, run=False):
self.context = context
self.run = run
async def __aenter__(self):
return await self.context.__aenter__()
async def __aexit__(self, typ, value, tb):
evidence = _evidence.get()
try:
result = await self.context.__aexit__(typ, value, tb)
except BaseException as exc:
# An asynccontextmanager propagating its body exception is not a
# cleanup failure. Replacement exceptions are fail-closed.
if evidence is not None and exc is not value:
evidence["cleanup_failed"] = True
if evidence is not None and self.run and exc is value:
evidence["run_exited"] = True
raise
else:
if evidence is not None and self.run:
evidence["run_exited"] = True
return result
async def _drain(task):
while not task.done():
try:
await asyncio.shield(task)
except asyncio.CancelledError:
continue
except BaseException:
break
if not task.cancelled():
task.exception()
async def _await_remote_future(remote: Future):
"""Do not let caller cancellation detach a coroutine on a worker thread."""
async def wait():
return await asyncio.wrap_future(remote)
task = asyncio.create_task(wait())
try:
return await asyncio.shield(task)
except asyncio.CancelledError:
await _drain(task)
raise
async def _persistent_cancellation_listener(original, queue, run_id, thread_id, done):
"""The in-memory listener's idle timeout is a poll, not end-of-listening."""
while not done.is_set():
await original(queue, run_id, thread_id, done)
def install() -> None:
global _installed
from langgraph_api import worker, stream
from langgraph_runtime_inmem import ops as inmem_ops
from langgraph.pregel._loop import AsyncPregelLoop
from quickjs_rs.threading import ThreadWorker
with _lock:
if _installed:
return
original = worker.worker
original_to_thread = asyncio.to_thread
original_enter = worker.Runs.enter
original_closing = stream.aclosing
original_stack = stream.AsyncExitStack
original_loop_exit = AsyncPregelLoop.__aexit__
original_cancellation_listener = inmem_ops.listen_for_cancellation
class ObservedStack(original_stack):
async def __aexit__(self, typ, value, tb):
return await _ExitBoundary(super()).__aexit__(typ, value, tb)
async def loop_exit(self, exc_type, exc_value, traceback):
evidence = _evidence.get()
if evidence is None:
return await original_loop_exit(self, exc_type, exc_value, traceback)
try:
return await original_loop_exit(self, exc_type, exc_value, traceback)
except asyncio.CancelledError as exc:
# Installed Pregel exposes its outstanding exit task in args.
tasks = [arg for arg in exc.args if isinstance(arg, asyncio.Task)]
for task in tasks:
await _drain(task)
if evidence is not None and (
not tasks or any(task.cancelled() or task.exception() is not None for task in tasks)
):
evidence["cleanup_failed"] = True
raise
except BaseException:
if evidence is not None:
evidence["cleanup_failed"] = True
raise
cast(Any, worker.Runs).enter = staticmethod(
lambda *args, **kw: _ExitBoundary(cast(Any, original_enter)(*args, **kw), run=True)
)
stream.aclosing = lambda iterator: _ExitBoundary(original_closing(iterator))
stream.AsyncExitStack = ObservedStack
cast(Any, AsyncPregelLoop).__aexit__ = loop_exit
async def persistent_cancellation_listener(queue, run_id, thread_id, done):
return await _persistent_cancellation_listener(
original_cancellation_listener, queue, run_id, thread_id, done
)
def quickjs_run_async(self, coro):
self._ensure_started()
remote = asyncio.run_coroutine_threadsafe(coro, self._loop)
return asyncio.create_task(_await_remote_future(remote))
inmem_ops.listen_for_cancellation = persistent_cancellation_listener
cast(Any, ThreadWorker).run_async = quickjs_run_async
@wraps(original_to_thread)
async def owned_to_thread(func, /, *args, **kwargs):
if not _owned.get():
return await original_to_thread(func, *args, **kwargs)
task = asyncio.create_task(original_to_thread(func, *args, **kwargs))
try:
return await asyncio.shield(task)
except asyncio.CancelledError:
await _drain(task)
raise
@wraps(original)
async def observed(run, attempt, main_loop, **kwargs):
key = (str(run["thread_id"]), str(run["run_id"]))
identity = (attempt, object())
with _lock:
entry = _runs.get(key)
if entry is not None:
if entry["cancel_requested"]:
# Cancel closed admission before this queued attempt ran.
return None
entry["active"] += 1
entry["exited"] = False
entry["attempts"][identity] = False
result = None
token = _owned.set(entry is not None)
evidence = {"run_exited": False, "cleanup_failed": False}
evidence_token = _evidence.set(evidence if entry is not None else None)
try:
result = await original(run, attempt, main_loop, **kwargs)
return result
finally:
_owned.reset(token)
_evidence.reset(evidence_token)
if entry is not None:
with _lock:
entry["active"] -= 1
# Only this invocation can discharge its own evidence.
# Business error/timeout/retry is independent of cleanup.
clean = evidence["run_exited"] and not evidence["cleanup_failed"]
entry["attempts"][identity] = clean
entry["uncertain"] = not all(entry["attempts"].values())
if clean:
entry["exited"] = True
entry["status"] = result["status"] if result else "retry"
worker.worker = observed
asyncio.to_thread = owned_to_thread
_installed = True
def reserve(thread_id: str, run_id: str, request_hash: str) -> None:
with _lock:
_runs.setdefault((thread_id, run_id), {
"request_hash": request_hash, "active": 0,
"cancel_requested": False, "exited": False,
"uncertain": False, "status": "pending", "attempts": {},
})
def is_reserved(thread_id: str, run_id: str) -> bool:
with _lock:
return (thread_id, run_id) in _runs
def cancel_and_inspect(thread_id: str, run_id: str) -> dict:
with _lock:
entry = _runs.get((thread_id, run_id))
if entry is None:
return {"run_id": run_id, "thread_id": thread_id, "execution_exited": False}
entry["cancel_requested"] = True
# No active worker plus closed admission is safe even for pending work.
# A worker that escaped abnormally keeps uncertainty latched.
stopped = entry["active"] == 0 and not entry["uncertain"]
return {
"run_id": run_id, "thread_id": thread_id,
"request_hash": entry["request_hash"],
"execution_exited": stopped,
"exit_kind": "worker_returned" if entry["exited"] else "admission_closed",
"status": entry["status"] if entry["exited"] else "interrupted",
}
async def wait_for_exit(thread_id: str, run_id: str, timeout: float = 3.0) -> dict:
deadline = asyncio.get_running_loop().time() + timeout
while True:
receipt = cancel_and_inspect(thread_id, run_id)
if receipt.get("execution_exited") is True:
return receipt
remaining = deadline - asyncio.get_running_loop().time()
if remaining <= 0:
return receipt
await asyncio.sleep(min(0.02, remaining))