fix(chain): 修评审 CRITICAL+HIGH — 链网关按 session 重建/日志脱敏/resume 原子化+所有权/死导入
- CRITICAL #1:链 write/review 节点经 gateway_builder 按节点自建 session 现建网关,
usage_ledger sink 绑活 session,随节点 commit 持久化;run_chain_job 不再转发请求网关
(其 session 在 BackgroundTask 跑时已关闭,记账行被静默丢弃)。新增 get_chain_gateway_builder
缝(仿 digest builder),get_chain_gateway 退化为纯 503 凭据预检。守不变量 #1。
- HIGH #2:chain_runner 失败日志不再记 str(exc)(可能含 key/连接串/LLM 输出),改记
_classify_job_error 脱敏文案 + exc_type(设计 §5)。
- HIGH #3:resume 端点原子抢占 awaiting→running(JobRepo.claim_awaiting_to_running 条件
UPDATE),抢不到 → 409,防并发 resume 双 Command(resume) 损坏图。
- HIGH #4:resume 校验 job.project_id == project_id(JobView 新增 project_id),不匹配 → 404。
- HIGH #5:resume 返回新 ChainResumeAccepted{job_id,chain_key},去掉无意义哨兵 start/count=0。
- HIGH #6:删 nodes.py 死导入 extract_conflicts(import + __all__)。
- 测试 #7:e2e 断言链跑后 usage_ledger 有行(#1 回归守卫)+ chapter_reviews 每章一行;
新增 claim 原子抢占单测 + resume 跨项目 404 / 并发 409 端点测。
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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) ----
|
||||
|
||||
|
||||
|
||||
@@ -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=<dict>(多章链 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=<str>。"""
|
||||
...
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user