diff --git a/EvoScientist/langgraph_dev/http.py b/EvoScientist/langgraph_dev/http.py index bec48be..fff5710 100644 --- a/EvoScientist/langgraph_dev/http.py +++ b/EvoScientist/langgraph_dev/http.py @@ -275,8 +275,46 @@ async def bind_workspace_run(request: Request) -> JSONResponse: return JSONResponse(_run_payload(run)) +def _interrupt_ids(value: Any) -> set[str]: + """从 __interrupt__ 写入值里取出中断标识(Interrupt 对象或字典都兼容)。""" + + items = value if isinstance(value, (list, tuple)) else [value] + found: set[str] = set() + for item in items: + ident = getattr(item, "id", None) + if ident is None and isinstance(item, dict): + ident = item.get("id") or item.get("interrupt_id") + if ident: + found.add(str(ident)) + return found + + +def _anchor_has_interrupt(checkpoint: Any, expected: str) -> bool: + """锚点处是否仍承载待审批中断(并核对中断标识)。 + + 这是"按锚点恢复"的准入判据:只要承载该中断的检查点写入还在 PG,暂停就仍然 + 有效 —— 与运行时进程是否重启过、距暂停多久都无关。 + """ + + found: set[str] = set() + for write in getattr(checkpoint, "pending_writes", None) or (): + try: + if len(write) < 3 or str(write[1]) != "__interrupt__": + continue + found |= _interrupt_ids(write[2]) + except TypeError: + continue + if not found: + return False + # 网关未提供标识(或回退值 default)时只做存在性判定。 + if not expected or expected == "default": + return True + return expected in found + + async def _compatible_checkpoint(conn, thread_id: str, assistant_id: str, - config: dict) -> tuple[bool, bool]: + config: dict, anchor: dict | None = None + ) -> tuple[bool, bool]: """Read through the API-owned saver and graph factory, never execute here.""" from langgraph_api._checkpointer import get_checkpointer from langgraph_api.graph import get_graph, graph_exists @@ -285,10 +323,15 @@ async def _compatible_checkpoint(conn, thread_id: str, assistant_id: str, saver = await get_checkpointer(conn=conn) read_config = {**config, "configurable": { **config.get("configurable", {}), "thread_id": thread_id, - "checkpoint_ns": "", + "checkpoint_ns": str((anchor or {}).get("checkpoint_ns") or ""), }} - # Admission always checks the current head, never a caller-selected ancestor. - read_config["configurable"].pop("checkpoint_id", None) + if anchor and str(anchor.get("checkpoint_id") or ""): + # 按锚点恢复:调用方(网关)指定了承载该中断的祖先检查点,暂停时就已落库。 + # 只有该锚点处确实还有待审批写入才放行 —— 不依赖运行时当前头部。 + read_config["configurable"]["checkpoint_id"] = str(anchor["checkpoint_id"]) + else: + # Admission always checks the current head, never a caller-selected ancestor. + read_config["configurable"].pop("checkpoint_id", None) checkpoint = await saver.aget_tuple(read_config) if checkpoint is None: return False, False @@ -300,6 +343,10 @@ async def _compatible_checkpoint(conn, thread_id: str, assistant_id: str, graph_id = assistant["graph_id"] if checkpoint.metadata.get("graph_id", graph_id) != graph_id: raise ValueError("checkpoint graph mismatch") + if anchor and str(anchor.get("checkpoint_id") or ""): + return True, _anchor_has_interrupt( + checkpoint, str(anchor.get("interrupt_id") or "") + ) # get_graph enters coroutine/async-context-manager factories and binds the # same API saver used by the worker. aget_state also validates delta seeds. async with get_graph(graph_id, read_config, checkpointer=saver, @@ -337,8 +384,11 @@ async def _has_legacy_history(conn, thread_id: str) -> bool: return True -async def _history_admission(conn, thread_id, assistant_id, config, operation, history): - exists, pending = await _compatible_checkpoint(conn, thread_id, assistant_id, config) +async def _history_admission(conn, thread_id, assistant_id, config, operation, history, + anchor: dict | None = None): + exists, pending = await _compatible_checkpoint( + conn, thread_id, assistant_id, config, anchor + ) if operation == "resume": return "resume" if exists and pending else "CHECKPOINT_RESUME_UNAVAILABLE" if pending: @@ -398,6 +448,11 @@ async def create_recoverable_run(request: Request) -> JSONResponse: return JSONResponse({"code": "INVALID_RESUME_REQUEST"}, status_code=400) elif command is not None: return JSONResponse({"code": "INVALID_START_REQUEST"}, status_code=400) + # 按锚点恢复(可选,仅 resume):网关在暂停时把"承载该中断的检查点"落库, + # 继续时回传。这样续接只依赖该检查点仍在 PG,而不依赖运行时当前头部。 + anchor = value.get("anchor") if operation == "resume" else None + if not isinstance(anchor, dict) or not str(anchor.get("checkpoint_id") or ""): + anchor = None history = value.get("history") # History validation is performed under the create lock, after idempotent # lookup. Only that branch can attest this request did not create a Run. @@ -427,12 +482,18 @@ async def create_recoverable_run(request: Request) -> JSONResponse: "run_request_id": run_request_id, "request_hash": request_hash, } - # The admission read and the worker must address the same current head. - # Forking from an ancestor is not part of the recoverable-run contract. + # The admission read and the worker must address the same checkpoint. + # Without an anchor that is the current head (forking from an ancestor is + # not part of the recoverable-run contract); with an anchor it is the + # ancestor recorded at pause time — which is exactly what makes the + # continuation independent of restarts and elapsed time. configurable = dict(payload["config"].get("configurable", {})) for key in ("checkpoint_id", "checkpoint_map"): configurable.pop(key, None) configurable.update(thread_id=thread_id, checkpoint_ns="") + if anchor is not None: + configurable["checkpoint_id"] = str(anchor["checkpoint_id"]) + configurable["checkpoint_ns"] = str(anchor.get("checkpoint_ns") or "") payload["config"] = {**payload["config"], "configurable": configurable} async with _recoverable_run_lock: async with connect() as conn: @@ -459,6 +520,7 @@ async def create_recoverable_run(request: Request) -> JSONResponse: try: admission = await _history_admission( conn, thread_id, assistant_id, payload["config"], operation, history, + anchor, ) except Exception: return JSONResponse({ diff --git a/EvoScientist/middleware/dynamic_review.py b/EvoScientist/middleware/dynamic_review.py index 7ff83d4..cf2bb8f 100644 --- a/EvoScientist/middleware/dynamic_review.py +++ b/EvoScientist/middleware/dynamic_review.py @@ -9,7 +9,7 @@ from typing import Annotated, Any, NotRequired import httpx from EvoScientist.internal_service import internal_service_headers -from langchain.agents.middleware import HumanInTheLoopMiddleware +from langchain.agents.middleware import HumanInTheLoopMiddleware, hook_config from langchain.agents.middleware.types import AgentState, OmitFromSchema from langgraph.config import get_config @@ -137,6 +137,31 @@ class DynamicReviewMiddleware(HumanInTheLoopMiddleware): state_schema = DynamicReviewState + @staticmethod + def _abort_requested() -> bool: + """本轮是否被要求"终止"(网关放弃某个待审批时注入的标记)。 + + 方案 A:中断仍由库正常消费(该工具不执行),但随后不再发起模型调用, + 直接把本轮跳到 end —— 即用户口径里的"不批准 → 停止这条对话"。 + """ + + try: + config = get_config() + except RuntimeError: + return False + configurable = config.get("configurable") if isinstance(config, Mapping) else None + if not isinstance(configurable, Mapping): + return False + return bool(configurable.get("ai4sci_abort_turn")) + + @classmethod + def _finish_turn(cls, update: dict[str, Any] | None) -> dict[str, Any] | None: + """终止本轮:保留库的决议结果(中断被消费掉),再结束本轮。""" + + if not cls._abort_requested(): + return update + return {**(update or {}), "jump_to": "end"} + def before_agent(self, state: DynamicReviewState, runtime: Any) -> dict[str, Any]: del state, runtime run_id, review = _review_context() @@ -153,6 +178,7 @@ class DynamicReviewMiddleware(HumanInTheLoopMiddleware): return {"_verified_review_mode": _manual_state(run_id, review)} return {"_verified_review_mode": await _resolve_async(run_id, review)} + @hook_config(can_jump_to=["end"]) def after_model( self, state: DynamicReviewState, runtime: Any ) -> dict[str, Any] | None: @@ -177,12 +203,12 @@ class DynamicReviewMiddleware(HumanInTheLoopMiddleware): return None if review.get("requested_mode") != "auto": # The gateway now requires manual approval for this turn. - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) try: _resolve_sync(current_run_id, review) except AutoReviewVerificationError: # Auto approval could not be re-verified; fall back to review. - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) return None if mode == "manual": # A LangGraph resume continues at this interrupted node and does not @@ -197,11 +223,12 @@ class DynamicReviewMiddleware(HumanInTheLoopMiddleware): try: _resolve_sync(current_run_id, review) except AutoReviewVerificationError: - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) return None - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) raise AutoReviewVerificationError("REVIEW_MODE_STATE_INVALID") + @hook_config(can_jump_to=["end"]) async def aafter_model( self, state: DynamicReviewState, runtime: Any ) -> dict[str, Any] | None: @@ -218,11 +245,11 @@ class DynamicReviewMiddleware(HumanInTheLoopMiddleware): if review is None: return None if review.get("requested_mode") != "auto": - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) try: await _resolve_async(current_run_id, review) except AutoReviewVerificationError: - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) return None if mode == "manual": # Mirror the sync path: an injected auto context on a resume child @@ -232,7 +259,7 @@ class DynamicReviewMiddleware(HumanInTheLoopMiddleware): try: await _resolve_async(current_run_id, review) except AutoReviewVerificationError: - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) return None - return super().after_model(state, runtime) + return self._finish_turn(super().after_model(state, runtime)) raise AutoReviewVerificationError("REVIEW_MODE_STATE_INVALID")