From ec91c8fdf78dac3f3501de8370492e94493818fd Mon Sep 17 00:00:00 2001 From: Yaojia Wang Date: Wed, 8 Jul 2026 13:04:28 +0200 Subject: [PATCH] =?UTF-8?q?refactor(chain):=20=E9=93=BE=20runner=20?= =?UTF-8?q?=E7=BB=8F=E5=85=AC=E5=BC=80=20aget=5Fstate.interrupts=20?= =?UTF-8?q?=E5=88=A4=E6=9A=82=E5=81=9C=E2=80=94=E2=80=94=E5=BC=83=E7=A7=81?= =?UTF-8?q?=E6=9C=89=20=5F=5Finterrupt=5F=5F=20=E9=94=AE=EF=BC=88CR-M3?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/api/tests/test_chain.py | 4 ++++ apps/api/ww_api/services/chain_runner.py | 18 ++++++++++-------- 2 files changed, 14 insertions(+), 8 deletions(-) diff --git a/apps/api/tests/test_chain.py b/apps/api/tests/test_chain.py index 405805b..ce32c7f 100644 --- a/apps/api/tests/test_chain.py +++ b/apps/api/tests/test_chain.py @@ -593,6 +593,10 @@ async def test_run_chain_job_sets_recursion_limit_scaled_to_count( captured["config"] = config return await self._inner.ainvoke(inp, config=config, **kw) + async def aget_state(self, config: Any, **kw: Any) -> Any: + # runner 现经公开 aget_state 判 interrupt(CR-M3)——代理透传给真 graph。 + return await self._inner.aget_state(config, **kw) + def _spy_build(*args: Any, **kwargs: Any) -> Any: return _SpyGraph(real_build(*args, **kwargs)) diff --git a/apps/api/ww_api/services/chain_runner.py b/apps/api/ww_api/services/chain_runner.py index 7f8cec2..2552675 100644 --- a/apps/api/ww_api/services/chain_runner.py +++ b/apps/api/ww_api/services/chain_runner.py @@ -161,13 +161,13 @@ def _chain_result(state: dict[str, Any], *, awaiting_chapter: int | None) -> dic } -def _extract_awaiting_chapter(final: dict[str, Any]) -> int | None: - """从图返回值判 interrupt:有 `__interrupt__` → 取暂停章号;否则 None(跑完)。 +def _extract_awaiting_chapter(interrupts: Sequence[Any]) -> int | None: + """从 pending interrupts 判暂停章号:空 → None(跑完);否则取首个 interrupt 载荷的章号。 - interrupt 载荷 = `{"chapter_no": n}`(见链图 `_accept` 节点)。LangGraph 把它放在 - `final["__interrupt__"]`(Interrupt 对象列表)。容错读 `.value`/dict 两种形。 + interrupts 经**公开** `graph.aget_state(config).interrupts`(`StateSnapshot`)拿到,不再 + 反手 ainvoke 返回值里的私有 `__interrupt__` 键。载荷 = `{"chapter_no": n}`(见链图 `_accept` + 节点)。容错读 `.value`/dict 两种形。 """ - interrupts = final.get("__interrupt__") if not interrupts: return None first = interrupts[0] @@ -244,9 +244,11 @@ async def run_chain_job( resume_value = [d.model_dump() for d in resume_decisions] raw_final = await graph.ainvoke(Command(resume=resume_value), config=config) - final: dict[str, Any] = dict(raw_final) - awaiting_chapter = _extract_awaiting_chapter(final) - result = _chain_result(final, awaiting_chapter=awaiting_chapter) + # 经公开 `StateSnapshot.interrupts` 判 pending interrupt(不反手私有 `__interrupt__`); + # checkpointer 仍在 `checkpointer_ctx()` 上下文内,aget_state 可读同一 thread 检查点。 + snapshot = await graph.aget_state(config) + awaiting_chapter = _extract_awaiting_chapter(snapshot.interrupts) + result = _chain_result(dict(raw_final), awaiting_chapter=awaiting_chapter) await _finish(session_factory, job_id, result, awaiting=awaiting_chapter is not None) log.info( "chain_job_settled",