refactor(chain): 链 runner 经公开 aget_state.interrupts 判暂停——弃私有 __interrupt__ 键(CR-M3)

This commit is contained in:
Yaojia Wang
2026-07-08 13:04:28 +02:00
parent c9ffada503
commit ec91c8fdf7
2 changed files with 14 additions and 8 deletions

View File

@@ -593,6 +593,10 @@ async def test_run_chain_job_sets_recursion_limit_scaled_to_count(
captured["config"] = config captured["config"] = config
return await self._inner.ainvoke(inp, config=config, **kw) return await self._inner.ainvoke(inp, config=config, **kw)
async def aget_state(self, config: Any, **kw: Any) -> Any:
# runner 现经公开 aget_state 判 interruptCR-M3——代理透传给真 graph。
return await self._inner.aget_state(config, **kw)
def _spy_build(*args: Any, **kwargs: Any) -> Any: def _spy_build(*args: Any, **kwargs: Any) -> Any:
return _SpyGraph(real_build(*args, **kwargs)) return _SpyGraph(real_build(*args, **kwargs))

View File

@@ -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: def _extract_awaiting_chapter(interrupts: Sequence[Any]) -> int | None:
"""图返回值判 interrupt有 `__interrupt__` → 取暂停章号;否则 None跑完 """ pending interrupts 判暂停章号:空 → None跑完否则取首个 interrupt 载荷的章号
interrupt 载荷 = `{"chapter_no": n}`(见链图 `_accept` 节点。LangGraph 把它放在 interrupts 经**公开** `graph.aget_state(config).interrupts``StateSnapshot`)拿到,不再
`final["__interrupt__"]`Interrupt 对象列表)。容错读 `.value`/dict 两种形。 反手 ainvoke 返回值里的私有 `__interrupt__` 键。载荷 = `{"chapter_no": n}`(见链图 `_accept`
节点)。容错读 `.value`/dict 两种形。
""" """
interrupts = final.get("__interrupt__")
if not interrupts: if not interrupts:
return None return None
first = interrupts[0] first = interrupts[0]
@@ -244,9 +244,11 @@ async def run_chain_job(
resume_value = [d.model_dump() for d in resume_decisions] resume_value = [d.model_dump() for d in resume_decisions]
raw_final = await graph.ainvoke(Command(resume=resume_value), config=config) raw_final = await graph.ainvoke(Command(resume=resume_value), config=config)
final: dict[str, Any] = dict(raw_final) # 经公开 `StateSnapshot.interrupts` 判 pending interrupt不反手私有 `__interrupt__`
awaiting_chapter = _extract_awaiting_chapter(final) # checkpointer 仍在 `checkpointer_ctx()` 上下文内aget_state 可读同一 thread 检查点。
result = _chain_result(final, awaiting_chapter=awaiting_chapter) 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) await _finish(session_factory, job_id, result, awaiting=awaiting_chapter is not None)
log.info( log.info(
"chain_job_settled", "chain_job_settled",