236 lines
9.1 KiB
Python
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)) |