diff --git a/apps/api/tests/fakes_projects.py b/apps/api/tests/fakes_projects.py index ac61c17..77e1637 100644 --- a/apps/api/tests/fakes_projects.py +++ b/apps/api/tests/fakes_projects.py @@ -414,7 +414,13 @@ class FakeJobRepo: self.rows: dict[uuid.UUID, JobView] = {} async def create(self, project_id: uuid.UUID | None, kind: str) -> JobView: - view = JobView(id=uuid.uuid4(), kind=kind, status=STATUS_QUEUED, progress=0) + view = JobView( + id=uuid.uuid4(), + project_id=project_id, + kind=kind, + status=STATUS_QUEUED, + progress=0, + ) self.rows[view.id] = view return view @@ -442,6 +448,15 @@ class FakeJobRepo: self.rows[job_id] = view return view + async def claim_awaiting_to_running(self, job_id: uuid.UUID) -> JobView | None: + """原子抢占 awaiting→running(内存替身:仅当现态 awaiting 才抢到,否则 None)。""" + current = self.rows.get(job_id) + if current is None or current.status != STATUS_AWAITING: + return None + view = current.model_copy(update={"status": STATUS_RUNNING}) + self.rows[job_id] = view + return view + async def fail(self, job_id: uuid.UUID, error: str) -> JobView: view = self.rows[job_id].model_copy(update={"status": STATUS_FAILED, "error": error}) self.rows[job_id] = view diff --git a/apps/api/tests/test_chain.py b/apps/api/tests/test_chain.py index 9d334c0..9cb2948 100644 --- a/apps/api/tests/test_chain.py +++ b/apps/api/tests/test_chain.py @@ -151,6 +151,19 @@ async def _memsaver_ctx(saver: MemorySaver) -> AsyncIterator[Any]: yield saver +def _async_returning(gateway: Any) -> Any: + """构造「按 session 建网关」的 builder 替身:忽略 session,恒返给定 mock 网关(绝不联网)。 + + 链节点经 `gateway_builder(session)` 现建网关(不变量 #1);测试里 ledger 不落库(fake + session),故 builder 只需把 mock 网关交回即可。 + """ + + async def _build(_session: Any) -> Any: + return gateway + + return _build + + class _FakeContext: stable_core = "## 世界观硬规则\n灵气可凝丹" volatile = "## 写作指令\n写第 N 章" @@ -169,6 +182,7 @@ def _app(project_repo: FakeProjectRepo, job_repo: FakeJobRepo) -> FastAPI: from ww_api.services.chain_deps import get_checkpointer_factory from ww_api.services.project_deps import ( get_chain_gateway, + get_chain_gateway_builder, get_digest_gateway_builder, get_job_repo, get_project_repo, @@ -182,6 +196,9 @@ def _app(project_repo: FakeProjectRepo, job_repo: FakeJobRepo) -> FastAPI: app.dependency_overrides[get_session] = lambda: _FakeSession() app.dependency_overrides[get_session_factory] = lambda: _session_factory app.dependency_overrides[get_chain_gateway] = lambda: _FakeChainGateway(conflicts=[]) + app.dependency_overrides[get_chain_gateway_builder] = lambda: _async_returning( + _FakeChainGateway(conflicts=[]) + ) saver = MemorySaver() app.dependency_overrides[get_checkpointer_factory] = lambda: lambda: _memsaver_ctx(saver) app.dependency_overrides[get_digest_gateway_builder] = lambda: lambda _s: None @@ -363,6 +380,56 @@ async def test_resume_chain_job_not_found_returns_404( assert resp.status_code == 404 +async def test_resume_chain_wrong_project_returns_404( + _noop_run_chain_job: list[dict[str, Any]], +) -> None: + """所有权校验(审评 #4):job 属项目 A,从项目 B 路径 resume → 404(同案,不枚举)。""" + project_repo = FakeProjectRepo() + job_repo = FakeJobRepo() + pid_a = await _seed_project(project_repo) + pid_b = uuid.UUID( + str((await project_repo.create(PROJECT_OWNER, ProjectCreate(title="作品B"))).id) + ) + job = await job_repo.create(pid_a, "chain") # job 属 A + await job_repo.set_awaiting(job.id, {"awaiting_chapter": 1}) + app = _app(project_repo, job_repo) + + async with _client(app) as client: + resp = await client.post( + f"/projects/{pid_b}/chains/runs/{job.id}/resume", # 从 B 续 A 的 job + json={"decisions": []}, + ) + + assert resp.status_code == 404 + assert len(_noop_run_chain_job) == 0 # 未调度 + + +async def test_resume_chain_concurrent_claim_returns_409( + _noop_run_chain_job: list[dict[str, Any]], +) -> None: + """原子抢占(审评 #3):第一个 resume 已把 awaiting→running,第二个抢不到 → 409。""" + project_repo = FakeProjectRepo() + job_repo = FakeJobRepo() + pid = await _seed_project(project_repo) + job = await job_repo.create(pid, "chain") + await job_repo.set_awaiting(job.id, {"awaiting_chapter": 1}) + # 模拟并发:第一个 resume 已抢占 awaiting→running。 + claimed = await job_repo.claim_awaiting_to_running(job.id) + assert claimed is not None + app = _app(project_repo, job_repo) + + async with _client(app) as client: + resp = await client.post( + f"/projects/{pid}/chains/runs/{job.id}/resume", + json={"decisions": []}, + ) + + # job 现态 running(非 awaiting)→ 守卫先在 status 检查处 409(CONFLICT)。 + assert resp.status_code == 409 + assert resp.json()["error"]["code"] == "CONFLICT" + assert len(_noop_run_chain_job) == 0 + + # ---- 服务层 run_chain_job(直接 await,注入全 fake;job 状态由 fake SqlJobRepo 替身记录)---- @@ -440,7 +507,7 @@ async def test_run_chain_job_no_conflict_completes_done(_patched_runner: None) - chain_key="draft_volume", start_chapter_no=1, count=2, - gateway=gateway, + chain_gateway_builder=_async_returning(gateway), checkpointer_ctx=lambda: _memsaver_ctx(saver), accept_op=fakes["accept_op"], chapter_repo_factory=fakes["chapter_repo_factory"], @@ -476,7 +543,7 @@ async def test_run_chain_job_conflict_sets_awaiting_then_resume_done( chain_key="draft_volume", start_chapter_no=1, count=1, - gateway=gateway, + chain_gateway_builder=_async_returning(gateway), checkpointer_ctx=lambda: _memsaver_ctx(saver), accept_op=fakes["accept_op"], chapter_repo_factory=fakes["chapter_repo_factory"], @@ -498,7 +565,7 @@ async def test_run_chain_job_conflict_sets_awaiting_then_resume_done( chain_key="draft_volume", start_chapter_no=1, count=1, - gateway=gateway, + chain_gateway_builder=_async_returning(gateway), checkpointer_ctx=lambda: _memsaver_ctx(saver), accept_op=fakes["accept_op"], chapter_repo_factory=fakes["chapter_repo_factory"], @@ -532,7 +599,7 @@ async def test_run_chain_job_error_marks_failed_safely(_patched_runner: None) -> chain_key="draft_volume", start_chapter_no=1, count=1, - gateway=_BoomGateway(), + chain_gateway_builder=_async_returning(_BoomGateway()), checkpointer_ctx=lambda: _memsaver_ctx(MemorySaver()), accept_op=fakes["accept_op"], chapter_repo_factory=fakes["chapter_repo_factory"], diff --git a/apps/api/ww_api/routers/chain.py b/apps/api/ww_api/routers/chain.py index 738cdf2..73ceea3 100644 --- a/apps/api/ww_api/routers/chain.py +++ b/apps/api/ww_api/routers/chain.py @@ -28,6 +28,7 @@ from ww_shared import AppError, ErrorCode, ErrorEnvelope from ww_api.logging_config import get_logger from ww_api.schemas.chain import ( + ChainResumeAccepted, ChainResumeRequest, ChainRunAccepted, ChainRunRequest, @@ -44,8 +45,10 @@ from ww_api.services.chain_runner import ( from ww_api.services.credentials import STUB_OWNER_ID from ww_api.services.foreshadow_scan import SessionFactory from ww_api.services.project_deps import ( + GatewayChainBuilder, GatewayDigestBuilder, get_chain_gateway, + get_chain_gateway_builder, get_digest_gateway_builder, get_job_repo, get_project_repo, @@ -66,6 +69,7 @@ SessionDep = Annotated[AsyncSession, Depends(get_session)] SessionFactoryDep = Annotated[SessionFactory, Depends(get_session_factory)] CheckpointerFactoryDep = Annotated[CheckpointerFactory, Depends(get_checkpointer_factory)] DigestBuilderDep = Annotated[GatewayDigestBuilder, Depends(get_digest_gateway_builder)] +ChainBuilderDep = Annotated[GatewayChainBuilder, Depends(get_chain_gateway_builder)] _RUN_ERRORS: dict[int | str, dict[str, Any]] = { 404: {"model": ErrorEnvelope, "description": "项目或链 key 不存在"}, @@ -97,11 +101,12 @@ async def run_chain( background_tasks: BackgroundTasks, project_repo: ProjectRepoDep, job_repo: JobRepoDep, - gateway: ChainGatewayDep, # 凭据探测(无凭据 → dep 解析阶段 503) + gateway: ChainGatewayDep, # 凭据探测(无凭据 → dep 解析阶段 503);不用于实际 LLM 调用 session: SessionDep, session_factory: SessionFactoryDep, checkpointer_factory: CheckpointerFactoryDep, digest_gateway_builder: DigestBuilderDep, + chain_gateway_builder: ChainBuilderDep, ) -> ChainRunAccepted: """发起多章链:写一行 job 返 202,链经 BackgroundTask 异步跑。 @@ -132,7 +137,7 @@ async def run_chain( chain_key=chain_key, start_chapter_no=body.start_chapter_no, count=body.count, - gateway=gateway, + chain_gateway_builder=chain_gateway_builder, checkpointer_ctx=checkpointer_factory, accept_op=accept_op, chapter_repo_factory=_chapter_repo_factory, @@ -172,22 +177,26 @@ async def resume_chain( background_tasks: BackgroundTasks, project_repo: ProjectRepoDep, job_repo: JobRepoDep, - gateway: ChainGatewayDep, + gateway: ChainGatewayDep, # 凭据探测(无凭据 → dep 解析阶段 503);不用于实际 LLM 调用 session: SessionDep, session_factory: SessionFactoryDep, checkpointer_factory: CheckpointerFactoryDep, digest_gateway_builder: DigestBuilderDep, -) -> ChainRunAccepted: + chain_gateway_builder: ChainBuilderDep, +) -> ChainResumeAccepted: """裁决续跑:仅当 job=awaiting_input,带裁决经 BackgroundTask resume。 - 项目不存在 → 404;job 不存在 → 404;非 awaiting 态 → 409(CONFLICT)。 - resume 经 `Command(resume=decisions)` 从 interrupt 续跑(同 thread_id=job_id)。 + 项目不存在 → 404;job 不存在 / 不属于该项目 → 404(同案,不枚举);非 awaiting 态 → + 409(CONFLICT)。`awaiting→running` 在 HTTP 处理器内**原子抢占**(条件 UPDATE),避免两个 + 并发 resume 都过守卫后双 `Command(resume=...)` 损坏图(审评 #3);抢不到(已被并发抢走/非 + awaiting)→ 409。resume 经 `Command(resume=decisions)` 从 interrupt 续跑(同 thread_id)。 """ request_id = getattr(request.state, "request_id", None) await _require_project(project_repo, project_id) job = await job_repo.get(job_id) - if job is None: + # 所有权校验(审评 #4):job 不存在 **或** 不属于该项目 → 404(同案,避免跨项目枚举 job)。 + if job is None or job.project_id != project_id: raise AppError(ErrorCode.NOT_FOUND, f"job not found: {job_id}") if job.status != STATUS_AWAITING: raise AppError( @@ -196,6 +205,17 @@ async def resume_chain( {"job_status": job.status}, ) + # 原子抢占 awaiting→running(审评 #3):抢不到(并发竞态/已非 awaiting)→ 409,绝不调度。 + claimed = await job_repo.claim_awaiting_to_running(job_id) + await session.commit() # 抢占须在 202 返回前持久化(防并发 resume 双调度)。 + if claimed is None: + raise AppError( + ErrorCode.CONFLICT, + f"job {job_id} 已被并发续跑抢占或非 awaiting_input 态,不可重复续跑", + ) + + chain_key = claimed.kind if claimed.kind in SUPPORTED_CHAINS else "draft_volume" + accept_op = build_accept_op( session_factory=session_factory, digest_gateway_builder=digest_gateway_builder, @@ -207,10 +227,10 @@ async def resume_chain( job_id, project_id=project_id, user_id=STUB_OWNER_ID, - chain_key=job.kind if job.kind in SUPPORTED_CHAINS else "draft_volume", + chain_key=chain_key, start_chapter_no=1, # resume 不重置区间:图从检查点续,初值不再使用。 count=1, - gateway=gateway, + chain_gateway_builder=chain_gateway_builder, checkpointer_ctx=checkpointer_factory, accept_op=accept_op, chapter_repo_factory=_chapter_repo_factory, @@ -227,12 +247,7 @@ async def resume_chain( decision_count=len(body.decisions), ) response.status_code = 202 - return ChainRunAccepted( - job_id=job_id, - chain_key="draft_volume", - start_chapter_no=0, - count=0, - ) + return ChainResumeAccepted(job_id=job_id, chain_key=chain_key) def _chapter_repo_factory(session: AsyncSession) -> object: diff --git a/apps/api/ww_api/schemas/chain.py b/apps/api/ww_api/schemas/chain.py index f3d047c..f382fda 100644 --- a/apps/api/ww_api/schemas/chain.py +++ b/apps/api/ww_api/schemas/chain.py @@ -52,3 +52,13 @@ class ChainResumeRequest(BaseModel): default_factory=list, description="对 awaiting 章冲突的裁决清单(采纳/忽略/手改)", ) + + +class ChainResumeAccepted(BaseModel): + """resume 受理回执(202):job_id + chain_key(前端走 `GET /jobs/{id}` 轮询续跑进度)。 + + 续跑无独立区间(图从检查点续),故不回显 start/count(会是无意义哨兵值)——只回 job 标识。 + """ + + job_id: uuid.UUID + chain_key: str diff --git a/apps/api/ww_api/services/chain_runner.py b/apps/api/ww_api/services/chain_runner.py index c04cadc..eda1dce 100644 --- a/apps/api/ww_api/services/chain_runner.py +++ b/apps/api/ww_api/services/chain_runner.py @@ -32,7 +32,7 @@ from ww_core.domain.job_repo import SqlJobRepo from ww_core.domain.review_repo import SqlReviewRepo from ww_core.memory import assemble from ww_core.memory.sql_repositories import sql_memory_repos -from ww_core.orchestrator import REVIEW_SPECS, GatewayRun +from ww_core.orchestrator import REVIEW_SPECS from ww_core.orchestrator.chain import build_chain_graph, initial_chain_state from ww_llm_gateway import Gateway @@ -54,6 +54,10 @@ JOB_KIND_CHAIN = "chain" # 「按 session 建 digest 网关」缝:accept 节点在自建短事务里需 light 档网关跑终稿提炼。 DigestGatewayBuilder = Callable[[Any], Awaitable[Gateway]] +# 「按 session 建链网关」缝:write/review 节点在自建短事务里现建网关(ledger 绑节点 session, +# 不变量 #1)。BackgroundTask 跑时请求 session 已关闭,故必须用节点 session 重建网关。 +GatewayChainBuilder = Callable[[Any], Awaitable[Gateway]] + def build_accept_op( *, @@ -178,7 +182,7 @@ async def run_chain_job( chain_key: str, start_chapter_no: int, count: int, - gateway: GatewayRun, + chain_gateway_builder: GatewayChainBuilder, checkpointer_ctx: CheckpointerCtx, accept_op: AcceptOp, chapter_repo_factory: ChapterRepoFactory, @@ -193,6 +197,10 @@ async def run_chain_job( 否则 resume(`graph.ainvoke(Command(resume=decisions), config)`),从 interrupt 续跑。 `thread_id = str(job_id)`:同 job 的初始/续跑落同一检查点 thread。 + `chain_gateway_builder`(非具体网关实例):write/review 节点在自建短事务的**新鲜** session + 上现建网关,使 usage_ledger 绑节点 session(不变量 #1)。**绝不**把请求阶段建的网关传进来—— + BackgroundTask 跑时请求 session 已关闭,其 ledger.record() 会落到死 session 上被静默丢弃。 + checkpointer 上下文(Postgres 连接 / MemorySaver)横跨整次 invoke;本 runner 在 task 内打开/关闭它。异常一律被吞(后台任务边界),脱敏后置 job failed,不冒泡崩进程。 """ @@ -201,7 +209,7 @@ async def run_chain_job( await _set_running(session_factory, job_id) config: RunnableConfig = {"configurable": {"thread_id": str(job_id)}} graph = build_chain_graph( - gateway, + chain_gateway_builder, session_factory=session_factory, memory_repos_factory=sql_memory_repos, chapter_repo_factory=chapter_repo_factory, @@ -236,8 +244,16 @@ async def run_chain_job( written_count=len(result["written"]), ) except Exception as exc: # noqa: BLE001 — 后台任务边界:记错误 + 置 job failed,不冒泡。 - log.error("chain_job_failed", job_id=str(job_id), request_id=request_id, error=str(exc)) - await _mark_failed(session_factory, job_id, _classify_job_error(exc), request_id) + # 日志只记脱敏文案 + 异常类型;绝不记 str(exc)(可能含 API key/连接串/LLM 输出,设计 §5)。 + stored_error = _classify_job_error(exc) + log.error( + "chain_job_failed", + job_id=str(job_id), + request_id=request_id, + exc_type=type(exc).__name__, + error=stored_error, + ) + await _mark_failed(session_factory, job_id, stored_error, request_id) async def _set_running(session_factory: SessionFactory, job_id: uuid.UUID) -> None: @@ -276,11 +292,12 @@ async def _mark_failed( await SqlJobRepo(session).fail(job_id, error) await session.commit() except Exception as exc: # noqa: BLE001 — 置失败态本身再炸只记日志,不冒泡。 + # 同 §5:不记 str(exc),只记异常类型(避免泄露内部细节)。 log.error( "chain_job_fail_mark_failed", job_id=str(job_id), request_id=request_id, - error=str(exc), + exc_type=type(exc).__name__, ) @@ -289,6 +306,7 @@ __all__ = [ "AcceptOp", "CheckpointerCtx", "DigestGatewayBuilder", + "GatewayChainBuilder", "build_accept_op", "run_chain_job", ] diff --git a/apps/api/ww_api/services/project_deps.py b/apps/api/ww_api/services/project_deps.py index befc0eb..3c8af1e 100644 --- a/apps/api/ww_api/services/project_deps.py +++ b/apps/api/ww_api/services/project_deps.py @@ -530,3 +530,27 @@ def get_digest_gateway_builder() -> GatewayDigestBuilder: `app.dependency_overrides[get_digest_gateway_builder]` 注返回 mock 的 builder(绝不联网)。 """ return _digest_gateway_builder + + +# 「按 session 建链网关」缝类型(链 write/review 节点在 BackgroundTask 自建短事务的 session 建)。 +GatewayChainBuilder = Callable[[AsyncSession], Awaitable[Gateway]] + + +def _chain_gateway_builder(session: AsyncSession) -> Awaitable[Gateway]: + """链 write/review 节点在自建短事务里建按请求 tier 分派的链网关。 + + 关键(守不变量 #1):网关的 `SqlAlchemyLedgerSink` 必须绑**节点当前的** session—— + BackgroundTask 跑时请求 session 已关闭,故不能复用请求阶段的网关,必须按节点新 session 现建, + 否则 `gateway.run()` 的 `usage_ledger` 写会落到已死 session 上被静默丢弃(成本记账断裂)。 + """ + return build_chain_gateway(session, SqlCredentialStore(session)) + + +def get_chain_gateway_builder() -> GatewayChainBuilder: + """返回「按 session 建链网关」的缝(链 write/review 节点在 BackgroundTask 内现建网关用)。 + + 与 `get_chain_gateway`(请求阶段凭据探测,返单实例,仅做 503 拦截)不同:本缝供 + `run_chain_job` 在节点自建的**新鲜** session 上重建网关,使 usage_ledger 落到活 session。 + 测试经 `app.dependency_overrides[get_chain_gateway_builder]` 注返 mock 的 builder(绝不联网)。 + """ + return _chain_gateway_builder diff --git a/packages/core/tests/test_chain_graph.py b/packages/core/tests/test_chain_graph.py index e52d035..f004e31 100644 --- a/packages/core/tests/test_chain_graph.py +++ b/packages/core/tests/test_chain_graph.py @@ -184,6 +184,10 @@ def _make_harness(*, conflicts: list[Conflict]) -> dict[str, Any]: draft_text="第 N 章正文。", conflicts=conflicts, by_schema=_empty_by_schema() ) + async def gateway_builder(session: Any) -> Any: + """按节点 session 建网关的替身(不变量 #1):忽略 session,恒返同一 mock 网关。""" + return gateway + def memory_repos_factory(session: Any) -> Any: return object() # assemble 被 fake 替换,不实际用 repos @@ -207,6 +211,7 @@ def _make_harness(*, conflicts: list[Conflict]) -> dict[str, Any]: return { "gateway": gateway, + "gateway_builder": gateway_builder, "session_factory": _session_factory, "memory_repos_factory": memory_repos_factory, "chapter_repo_factory": chapter_repo_factory, @@ -221,7 +226,7 @@ def _make_harness(*, conflicts: list[Conflict]) -> dict[str, Any]: def _build(h: dict[str, Any], checkpointer: Any) -> Any: return build_chain_graph( - h["gateway"], + h["gateway_builder"], session_factory=h["session_factory"], memory_repos_factory=h["memory_repos_factory"], chapter_repo_factory=h["chapter_repo_factory"], @@ -317,7 +322,7 @@ async def test_write_chapter_saves_collected_draft() -> None: out = await write_chapter( state, - gateway=h["gateway"], + gateway_builder=h["gateway_builder"], session_factory=h["session_factory"], memory_repos_factory=h["memory_repos_factory"], chapter_repo_factory=h["chapter_repo_factory"], diff --git a/packages/core/tests/test_job_repo.py b/packages/core/tests/test_job_repo.py index 1c32866..14a4543 100644 --- a/packages/core/tests/test_job_repo.py +++ b/packages/core/tests/test_job_repo.py @@ -100,6 +100,14 @@ class FakeJobRepo: row.result = dict(result) return _view(row) + async def claim_awaiting_to_running(self, job_id: uuid.UUID) -> JobView | None: + """原子抢占 awaiting→running(内存替身:仅当现态 awaiting 才抢到,否则 None)。""" + row = next((r for r in self.rows if r.id == job_id), None) + if row is None or row.status != STATUS_AWAITING: + return None + row.status = STATUS_RUNNING + return _view(row) + async def fail(self, job_id: uuid.UUID, error: str) -> JobView: row = self._require(job_id) row.status = STATUS_FAILED @@ -173,6 +181,39 @@ async def test_fail_sets_failed_and_error() -> None: assert failed.error == "boom" +# ---- 原子抢占 awaiting → running(防并发 resume 竞态,审评 #3)---- + + +async def test_claim_awaiting_to_running_succeeds_once() -> None: + repo: JobRepo = _repo() + job = await repo.create(PROJECT, KIND) + await repo.set_awaiting(job.id, {"awaiting_chapter": 1}) + claimed = await repo.claim_awaiting_to_running(job.id) + assert claimed is not None + assert claimed.status == STATUS_RUNNING + + +async def test_claim_awaiting_second_call_returns_none() -> None: + repo: JobRepo = _repo() + job = await repo.create(PROJECT, KIND) + await repo.set_awaiting(job.id, {"awaiting_chapter": 1}) + first = await repo.claim_awaiting_to_running(job.id) + second = await repo.claim_awaiting_to_running(job.id) # 已 running → 抢不到 + assert first is not None + assert second is None + + +async def test_claim_non_awaiting_returns_none() -> None: + repo: JobRepo = _repo() + job = await repo.create(PROJECT, KIND) # queued(非 awaiting) + assert await repo.claim_awaiting_to_running(job.id) is None + + +async def test_claim_absent_job_returns_none() -> None: + repo: JobRepo = _repo() + assert await repo.claim_awaiting_to_running(uuid.uuid4()) is None + + # ---- progress (clamped) ---- diff --git a/packages/core/ww_core/domain/job_repo.py b/packages/core/ww_core/domain/job_repo.py index bdfcfb3..34c64c6 100644 --- a/packages/core/ww_core/domain/job_repo.py +++ b/packages/core/ww_core/domain/job_repo.py @@ -39,6 +39,7 @@ class JobView(BaseModel): model_config = {"frozen": True} id: uuid.UUID + project_id: uuid.UUID | None = None kind: str status: str progress: int = 0 @@ -69,6 +70,11 @@ class JobRepo(Protocol): """置 status=awaiting_input, result=(多章链 interrupt 暂停等裁决;chain §5/§7)。""" ... + async def claim_awaiting_to_running(self, job_id: uuid.UUID) -> JobView | None: + """**原子**抢占 awaiting_input → running,返回抢到的 JobView;未抢到(不存在/非 + awaiting/已被并发抢走)返 None。供 resume 端点防并发 resume 竞态(chain §3/审评 #3)。""" + ... + async def fail(self, job_id: uuid.UUID, error: str) -> JobView: """置 status=failed, error=。""" ... @@ -85,6 +91,7 @@ class JobRepo(Protocol): def _to_view(row: Job) -> JobView: return JobView( id=row.id, + project_id=row.project_id, kind=row.kind, status=row.status, progress=row.progress, @@ -158,6 +165,34 @@ class SqlJobRepo: await self._s.refresh(row) return _to_view(row) + async def claim_awaiting_to_running(self, job_id: uuid.UUID) -> JobView | None: + """原子 `awaiting_input → running`(条件 UPDATE,returning 行)。 + + 防并发 resume 竞态(审评 #3):两个并发 resume 仅一个能把 awaiting→running,另一个 + UPDATE 命中 0 行 → 返 None(端点据此返 409)。避免「两个 resume 都过 awaiting 守卫 → + 同 thread_id 双 `Command(resume=...)` → 图损坏」。提交交调用方(端点事务)。 + """ + result = await self._s.execute( + update(Job) + .where(Job.id == job_id, Job.status == STATUS_AWAITING) + .values(status=STATUS_RUNNING) + .returning( + Job.id, Job.project_id, Job.kind, Job.status, Job.progress, Job.result, Job.error + ) + ) + row = result.one_or_none() + if row is None: + return None + return JobView( + id=row.id, + project_id=row.project_id, + kind=row.kind, + status=row.status, + progress=row.progress, + result=row.result, + error=row.error, + ) + async def fail(self, job_id: uuid.UUID, error: str) -> JobView: row = await self._require(job_id) row.status = STATUS_FAILED diff --git a/packages/core/ww_core/orchestrator/chain/__init__.py b/packages/core/ww_core/orchestrator/chain/__init__.py index 1b44738..ca448c0 100644 --- a/packages/core/ww_core/orchestrator/chain/__init__.py +++ b/packages/core/ww_core/orchestrator/chain/__init__.py @@ -20,6 +20,7 @@ from .nodes import ( AssembleContext, AssembleFn, ChapterDraftRepo, + GatewayBuilder, MemoryReposFactory, ReviewRecordRepo, SessionFactory, @@ -41,6 +42,7 @@ __all__ = [ "AssembleFn", "ChainState", "ChapterDraftRepo", + "GatewayBuilder", "MemoryReposFactory", "ReviewRecordRepo", "SessionFactory", diff --git a/packages/core/ww_core/orchestrator/chain/graph.py b/packages/core/ww_core/orchestrator/chain/graph.py index 40b0383..3f911cf 100644 --- a/packages/core/ww_core/orchestrator/chain/graph.py +++ b/packages/core/ww_core/orchestrator/chain/graph.py @@ -22,13 +22,13 @@ from langgraph.graph import END, START, StateGraph from langgraph.types import interrupt from ww_agents import AgentSpec -from .._protocols import GatewayRun from ..graph import REVIEW_SPECS from . import nodes from .nodes import ( AcceptChapterOp, AssembleFn, ChapterDraftRepo, + GatewayBuilder, MemoryReposFactory, ReviewRecordRepo, SessionFactory, @@ -46,7 +46,7 @@ ACCEPT_CHAPTER = "accept_chapter" def build_chain_graph( - gateway: GatewayRun, + gateway_builder: GatewayBuilder, *, session_factory: SessionFactory, memory_repos_factory: MemoryReposFactory, @@ -60,7 +60,9 @@ def build_chain_graph( """构建并编译 `draft_volume` 多章链图。 节点(`write_chapter`/`review_chapter`/`decide`/`accept_chapter`)经默认参绑定把 - `gateway`/`session_factory`/各 repo 工厂/`assemble`/`accept_op` 闭包绑入。 + `gateway_builder`/`session_factory`/各 repo 工厂/`assemble`/`accept_op` 闭包绑入。 + `gateway_builder` 而非具体网关实例:write/review 节点在自建 session 上现建网关,使 + usage_ledger 绑节点 session(不变量 #1,修复请求网关 session 已关闭致记账丢失的缺陷)。 条件边:`decide` 命中冲突 → `accept_chapter` 内 `interrupt()` 暂停交人; `accept_chapter` 后判 `current_chapter_no <= last_chapter_no` 回 `write_chapter` 或 END。 @@ -70,7 +72,7 @@ def build_chain_graph( async def _write(state: ChainState) -> dict[str, Any]: return await nodes.write_chapter( state, - gateway=gateway, + gateway_builder=gateway_builder, session_factory=session_factory, memory_repos_factory=memory_repos_factory, chapter_repo_factory=chapter_repo_factory, @@ -80,7 +82,7 @@ def build_chain_graph( async def _review(state: ChainState) -> dict[str, Any]: return await nodes.review_chapter( state, - gateway=gateway, + gateway_builder=gateway_builder, session_factory=session_factory, memory_repos_factory=memory_repos_factory, chapter_repo_factory=chapter_repo_factory, diff --git a/packages/core/ww_core/orchestrator/chain/nodes.py b/packages/core/ww_core/orchestrator/chain/nodes.py index 2ac69ba..2926c94 100644 --- a/packages/core/ww_core/orchestrator/chain/nodes.py +++ b/packages/core/ww_core/orchestrator/chain/nodes.py @@ -19,7 +19,7 @@ import structlog from ww_agents import AgentSpec from .._protocols import GatewayRun -from ..collect import collect_reviews, extract_conflicts +from ..collect import collect_reviews from ..review_node import build_review_context, run_review from ..state import ChapterState from ..write_node import build_write_request @@ -32,6 +32,11 @@ log = structlog.get_logger(__name__) #: `session -> Repo`:节点自建短事务里从 session 造 repo(仿端点依赖工厂)。 SessionFactory = Callable[[], AbstractAsyncContextManager[Any]] +#: `session -> Gateway`:按**节点当前 session** 建网关(不变量 #1)。 +#: 关键:网关的 ledger sink 绑节点 session,节点末尾 `commit()` 一并持久 usage_ledger 行—— +#: 绝不能复用请求阶段建好的网关(其 session 在 BackgroundTask 跑时已关闭,记账行会被静默丢弃)。 +GatewayBuilder = Callable[[Any], Awaitable[GatewayRun]] + class AssembleContext(Protocol): """`assemble` 产出的最小读形——只需 stable_core/volatile(构造写章/审稿请求)。""" @@ -103,7 +108,7 @@ AcceptChapterOp = Callable[..., Awaitable[None]] async def write_chapter( state: ChainState, *, - gateway: GatewayRun, + gateway_builder: GatewayBuilder, session_factory: SessionFactory, memory_repos_factory: MemoryReposFactory, chapter_repo_factory: Callable[[Any], ChapterDraftRepo], @@ -114,11 +119,15 @@ async def write_chapter( 短事务、独立 session(不跨章共享)。state 不变(正文在 DB)。返回空增量。 收集版写章复用 `build_write_request` 同一 prompt 组装路径(stream=False → `gateway.run`), 守不变量 #9(stable_core 进缓存前缀);不确定性锁在网关后,注 mock 即可单测。 + + 网关经 `gateway_builder(session)` 按**本节点 session** 现建(不变量 #1):其 ledger sink + 绑此 session,末尾 `commit()` 一并持久 usage_ledger——绝不复用请求阶段已关闭 session 的网关。 """ project_id = state["project_id"] chapter_no = state["current_chapter_no"] user_id = state["user_id"] async with session_factory() as session: + gateway = await gateway_builder(session) repos = memory_repos_factory(session) context = await assemble(repos, project_id, chapter_no) req = build_write_request( @@ -145,7 +154,7 @@ async def write_chapter( async def review_chapter( state: ChainState, *, - gateway: GatewayRun, + gateway_builder: GatewayBuilder, session_factory: SessionFactory, memory_repos_factory: MemoryReposFactory, chapter_repo_factory: Callable[[Any], ChapterDraftRepo], @@ -157,11 +166,15 @@ async def review_chapter( 复用现有 `run_review × specs` + `collect_reviews`(直接顺序跑各审,非起子图——链节点本身 已是图节点,无需嵌套图)。从 DB 重读本章草稿构审稿上下文(真相在表,不信 state)。 + + 网关经 `gateway_builder(session)` 按**本节点 session** 现建(不变量 #1):四审各调的 + usage_ledger 行随节点末尾 `commit()` 一并持久——绝不复用请求阶段已关闭 session 的网关。 """ project_id = state["project_id"] chapter_no = state["current_chapter_no"] user_id = state["user_id"] async with session_factory() as session: + gateway = await gateway_builder(session) repos = memory_repos_factory(session) context = await assemble(repos, project_id, chapter_no) chapter_repo = chapter_repo_factory(session) @@ -302,9 +315,9 @@ __all__ = [ "MemoryReposFactory", "ReviewRecordRepo", "SessionFactory", + "GatewayBuilder", "accept_chapter", "decide", - "extract_conflicts", "has_conflicts", "review_chapter", "write_chapter", diff --git a/tests/test_chain_workflow_e2e.py b/tests/test_chain_workflow_e2e.py index 410f570..21ee9c7 100644 --- a/tests/test_chain_workflow_e2e.py +++ b/tests/test_chain_workflow_e2e.py @@ -217,12 +217,15 @@ def _build_app(adapter: _FakeChainAdapter, saver: MemorySaver, e2e_sm: Any) -> A from ww_api.services.chain_deps import get_checkpointer_factory from ww_api.services.project_deps import ( get_chain_gateway, + get_chain_gateway_builder, get_digest_gateway_builder, get_session_factory, ) app = create_app() app.dependency_overrides[get_chain_gateway] = _chain_gateway_override(adapter) + # 链 write/review 节点经 builder 在节点 session 上现建网关(ledger 绑节点 session,不变量 #1)。 + app.dependency_overrides[get_chain_gateway_builder] = _digest_builder_override(adapter) app.dependency_overrides[get_session_factory] = lambda: e2e_sm app.dependency_overrides[get_checkpointer_factory] = _memsaver_override(saver) app.dependency_overrides[get_digest_gateway_builder] = _digest_builder_override(adapter) @@ -302,6 +305,35 @@ async def test_chain_two_chapters_no_conflict_full_auto( ) assert [d.chapter_no for d in digests] == [1, 2] + # 回归守卫(审评 #1):链每次 gateway.run 都应落 usage_ledger 行——证明网关 ledger + # 绑节点活 session、节点 commit 持久化。 + # 修复前:请求网关 session 已关闭,行被静默丢弃。 + ledger_rows = ( + ( + await verify.execute( + select(UsageLedger).where(UsageLedger.project_id == project_uuid) + ) + ) + .scalars() + .all() + ) + # 两章 × (1 写 + 四审 + 1 digest);至少应有若干条,绝不为空。 + assert len(ledger_rows) > 0, "chain usage_ledger 行不应为空(成本记账断裂回归守卫)" + + # 每审过一章应在 chapter_reviews 留痕(不变量 #3 只读留痕;审评 #7b)。 + review_rows = ( + ( + await verify.execute( + select(ChapterReview) + .where(ChapterReview.project_id == project_uuid) + .order_by(ChapterReview.chapter_no) + ) + ) + .scalars() + .all() + ) + assert {r.chapter_no for r in review_rows} == {1, 2} + job_rows = ( (await verify.execute(select(Job).where(Job.project_id == project_uuid))) .scalars()