diff --git a/tests/test_chain_workflow_e2e.py b/tests/test_chain_workflow_e2e.py new file mode 100644 index 0000000..410f570 --- /dev/null +++ b/tests/test_chain_workflow_e2e.py @@ -0,0 +1,429 @@ +"""C4 多章工作流链 端到端(chain-workflow §9 C4 / §10)——真 pg + mock 网关零 token。 + +证明杀手锏闭环(不变量 #1/#3/#4/#5/#9):批量量产,但每章过 `write→四审→decide→accept`, +遇冲突 interrupt 交人裁决再 resume 续跑——星月架构做不到的一致性闸。 + + 用例 1 · 两章无冲突全自动 + `POST /projects` → `POST .../chains/draft_volume/run {start=1,count=2}` → 202 {job_id}; + BackgroundTask `run_chain_job` 自建独立 session 驱动链图(mock 网关零冲突)→ 两章 + `write→四审(无冲突)→accept` 自动跑完。DB 真源逐项断言: + - `chapters` 第 1/2 章各一行 `accepted`(write 落 draft → accept 晋升终稿); + - `chapter_digests` 两行(digest 从终稿提炼,不变量 #4); + - `GET /jobs/{id}` status=done,result.written=[1,2]、completed=True、awaiting_chapter=None。 + 负向:job result/状态绝不含正文/prompt/token(§5)。 + + 用例 2 · 注入冲突 → interrupt → awaiting → resume → accept + 续审第 1 章报冲突 → `decide` 命中 → `accept` 节点 `interrupt()` 暂停 → job=awaiting_input、 + result.awaiting_chapter=1、`chapters` 第 1 章仍 `draft`(未晋升)、`chapter_digests` 空。 + `POST .../chains/runs/{job_id}/resume {decisions:[...]}` → 202 → resume 续跑 → accept → + job done、第 1 章 `accepted`、digest 一行。 + interrupt+resume 横跨**两次端点调用**,靠单进程单 `MemorySaver`(经依赖覆盖注入 + `get_checkpointer_factory` 的同一 saver 实例)持久控制流位置 + pending handle。 + +确定性 & 零成本(同 M3/M4):真 `Gateway` + 假适配器(据 `req.output_schema` 分支返固定 +parsed/text,绝不联网)+ 真 `SqlAlchemyLedgerSink`。检查点用 MemorySaver(绝不连真 Postgres +检查点)。无 pg → skip。 + +坑(见 memory/gotchas): +- `get_checkpointer_factory` 必 override 成产**同一** MemorySaver 的工厂——run 与 resume 两次 + 端点调用须落同一 thread_id 的同一检查点存储,否则 resume 找不到 interrupt 现场。 +- `get_session_factory` 必 override → `e2e_sm`(真 sessionmaker,同测试 engine/loop); + BackgroundTask 在请求 session 关闭后才跑。 +- `get_digest_gateway_builder` 必 override → 返「按 session 建真 Gateway 包假适配器」的 builder + (链 accept 节点终稿提炼用 light 档),否则会从空凭据建真 OpenAI 适配器。 +- ASGITransport 下 `await client.post` 等 ASGI app(含 background task)跑完才返回。 +""" + +from __future__ import annotations + +import uuid +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import AbstractAsyncContextManager, asynccontextmanager +from typing import Annotated, Any + +import httpx +import pytest +from asgi_lifespan import LifespanManager +from fastapi import Depends +from langgraph.checkpoint.memory import MemorySaver +from sqlalchemy import delete, select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker +from ww_agents import ( + Conflict, + ContinuityReview, + ForeshadowReview, + PaceReview, + StyleDriftReview, +) +from ww_api.services.digest_extraction import ChapterDigestFacts +from ww_db import get_session, get_sessionmaker +from ww_db.models import ( + Chapter, + ChapterDigest, + ChapterReview, + Job, + Project, + UsageLedger, +) +from ww_llm_gateway import Gateway, SqlAlchemyLedgerSink, resolve_route +from ww_llm_gateway.adapters.base import ( + Capabilities, + ProviderResult, + ProviderUsage, + StreamChunk, +) +from ww_llm_gateway.types import LlmRequest + +# 三档(writer/analyst/light)默认都路由 deepseek(config.tier_defaults)。 +_PROVIDER = "deepseek" + +# 链写章收集版正文(非 SSE,经 gateway.run → adapter.complete,output_schema is None)。 +_CHAPTER_TEXT = "灵气如潮,少年握剑而立,山门之上风雷大作。" + +# 各档假用量(喂记账;证明各档调用各落一条 usage_ledger)。 +_USAGE = ProviderUsage(input_tokens=13, output_tokens=7) + +# 三审固定无冲突产物(让四审图跑通;continuity 的冲突由 fixture 参数注入)。 +_FORESHADOW_REVIEW = ForeshadowReview(planted=[], resolved=[]) +_PACE_REVIEW = PaceReview(water=[], hook=True, beat_map=[1, 2, 3]) +_STYLE_DRIFT = StyleDriftReview(score=88, segments=[]) + + +class _FakeChainAdapter: + """实现 `ProviderAdapter`:据 `req.output_schema` 分支返 parsed/text,绝不联网。 + + - `output_schema is None` → 链写章收集版纯文本(writer 档,`gateway.run`)。 + - 四审 schema → 各审 parsed(continuity 注入冲突,其余无冲突)。 + - `ChapterDigestFacts` → 终稿 digest 提炼(light 档)。 + `stream()` 实现但链不走(收集版量产非 SSE)。 + """ + + provider = _PROVIDER + + def __init__(self, *, continuity_conflicts: list[Conflict]) -> None: + self._continuity = ContinuityReview(conflicts=continuity_conflicts) + + def capabilities(self) -> Capabilities: + return Capabilities(structured_output=True) + + async def complete(self, req: LlmRequest, model: str) -> ProviderResult: + schema = req.output_schema + if schema is None: + return ProviderResult(text=_CHAPTER_TEXT, parsed=None, usage=_USAGE) + if schema is ContinuityReview: + return ProviderResult( + text=self._continuity.model_dump_json(), parsed=self._continuity, usage=_USAGE + ) + if schema is ForeshadowReview: + return ProviderResult( + text=_FORESHADOW_REVIEW.model_dump_json(), parsed=_FORESHADOW_REVIEW, usage=_USAGE + ) + if schema is PaceReview: + return ProviderResult( + text=_PACE_REVIEW.model_dump_json(), parsed=_PACE_REVIEW, usage=_USAGE + ) + if schema is StyleDriftReview: + return ProviderResult( + text=_STYLE_DRIFT.model_dump_json(), parsed=_STYLE_DRIFT, usage=_USAGE + ) + # 否则视为终稿 digest 提炼(ChapterDigestFacts)。 + facts = ChapterDigestFacts(summary="终稿摘要", events=["少年出山"], locations=["山门"]) + return ProviderResult(text=facts.model_dump_json(), parsed=facts, usage=_USAGE) + + async def stream(self, req: LlmRequest, model: str) -> AsyncIterator[StreamChunk]: + yield StreamChunk(text=_CHAPTER_TEXT) + yield StreamChunk(usage=_USAGE) + + +@pytest.fixture +async def e2e_sm() -> AsyncIterator[async_sessionmaker[AsyncSession]]: + """真 DB session 工厂;无 pg 时跳过(每测试清缓存重建 engine、结束 dispose)。""" + get_sessionmaker.cache_clear() + maker = get_sessionmaker() + try: + async with maker() as probe: + await probe.execute(select(1)) + except Exception: + pytest.skip("postgres not reachable") + yield maker + await maker.kw["bind"].dispose() + get_sessionmaker.cache_clear() + + +def _gateway_for(adapter: _FakeChainAdapter, session: AsyncSession) -> Gateway: + """真 Gateway + 给定假适配器 + 真 ledger(绑给定 session)。""" + return Gateway( + adapters={_PROVIDER: adapter}, + ledger=SqlAlchemyLedgerSink(session), + resolver=resolve_route, + ) + + +def _chain_gateway_override( + adapter: _FakeChainAdapter, +) -> Callable[[AsyncSession], Gateway]: + """`get_chain_gateway` 覆盖:请求阶段凭据探测 → 真 Gateway 包假适配器(绕过 503)。""" + + def _override(session: Annotated[AsyncSession, Depends(get_session)]) -> Gateway: + return _gateway_for(adapter, session) + + return _override + + +def _digest_builder_override( + adapter: _FakeChainAdapter, +) -> Callable[[], Callable[[AsyncSession], Awaitable[Gateway]]]: + """`get_digest_gateway_builder` 覆盖:返「按 session 建真 Gateway 包假适配器」的 builder。 + + 链 accept 节点在自建短事务里据该 session 现建 light 档 digest 网关(终稿提炼)。 + """ + + def _get_builder() -> Callable[[AsyncSession], Awaitable[Gateway]]: + async def _build(session: AsyncSession) -> Gateway: + return _gateway_for(adapter, session) + + return _build + + return _get_builder + + +def _memsaver_override( + saver: MemorySaver, +) -> Callable[[], Callable[[], AbstractAsyncContextManager[MemorySaver]]]: + """`get_checkpointer_factory` 覆盖:返产**同一** saver 的工厂(run/resume 跨调用共用)。""" + + @asynccontextmanager + async def _ctx() -> AsyncIterator[MemorySaver]: + yield saver + + return lambda: lambda: _ctx() + + +async def _cleanup(e2e_sm: async_sessionmaker[AsyncSession], project_uuid: uuid.UUID) -> None: + """按 FK 顺序清理(多数子表 FK CASCADE;显式删保独立)。""" + async with e2e_sm() as cleanup: + await cleanup.execute(delete(UsageLedger).where(UsageLedger.project_id == project_uuid)) + await cleanup.execute(delete(Job).where(Job.project_id == project_uuid)) + await cleanup.execute(delete(ChapterDigest).where(ChapterDigest.project_id == project_uuid)) + await cleanup.execute(delete(ChapterReview).where(ChapterReview.project_id == project_uuid)) + await cleanup.execute(delete(Chapter).where(Chapter.project_id == project_uuid)) + await cleanup.execute(delete(Project).where(Project.id == project_uuid)) + await cleanup.commit() + + +def _build_app(adapter: _FakeChainAdapter, saver: MemorySaver, e2e_sm: Any) -> Any: + """组装注入了全 fake 缝的 app(凭据探测/session 工厂/checkpointer/digest builder)。""" + from ww_api.main import create_app + from ww_api.services.chain_deps import get_checkpointer_factory + from ww_api.services.project_deps import ( + get_chain_gateway, + get_digest_gateway_builder, + get_session_factory, + ) + + app = create_app() + app.dependency_overrides[get_chain_gateway] = _chain_gateway_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) + return app + + +async def test_chain_two_chapters_no_conflict_full_auto( + e2e_sm: async_sessionmaker[AsyncSession], +) -> None: + """用例 1:两章无冲突全自动 → 两 accepted 章 + 两 digest 行 + job done written=[1,2]。""" + adapter = _FakeChainAdapter(continuity_conflicts=[]) # 零冲突 → 全自动 + saver = MemorySaver() + app = _build_app(adapter, saver, e2e_sm) + + transport = httpx.ASGITransport(app=app) + project_uuid: uuid.UUID | None = None + try: + async with LifespanManager(app): + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + create_resp = await client.post("/projects", json={"title": "链 E2E 无冲突作品"}) + assert create_resp.status_code == 201 + project_id = create_resp.json()["id"] + project_uuid = uuid.UUID(project_id) + + run_resp = await client.post( + f"/projects/{project_id}/chains/draft_volume/run", + json={"start_chapter_no": 1, "count": 2}, + ) + assert run_resp.status_code == 202 + run_body = run_resp.json() + assert run_body["chain_key"] == "draft_volume" + job_id = run_body["job_id"] + + # 链经 BackgroundTask 已跑完(ASGITransport 等其完成)→ job done。 + job_resp = await client.get(f"/jobs/{job_id}") + assert job_resp.status_code == 200 + job = job_resp.json() + assert job["status"] == "done", job + assert job["kind"] == "chain" + result = job["result"] + assert result["written"] == [1, 2] + assert result["completed"] is True + assert result["awaiting_chapter"] is None + # 负向:job result 绝不含正文/token(只章号/计数/标志)。 + assert _CHAPTER_TEXT not in str(result) + assert "终稿摘要" not in str(result) + + # DB 真源逐项断言。 + assert project_uuid is not None + async with e2e_sm() as verify: + chapters = ( + ( + await verify.execute( + select(Chapter) + .where(Chapter.project_id == project_uuid) + .order_by(Chapter.chapter_no, Chapter.version) + ) + ) + .scalars() + .all() + ) + # 第 1/2 章各晋升一个 accepted 版次。 + accepted = [c for c in chapters if c.status == "accepted"] + assert {c.chapter_no for c in accepted} == {1, 2} + assert all(c.content == _CHAPTER_TEXT for c in accepted) + + digests = ( + ( + await verify.execute( + select(ChapterDigest) + .where(ChapterDigest.project_id == project_uuid) + .order_by(ChapterDigest.chapter_no) + ) + ) + .scalars() + .all() + ) + assert [d.chapter_no for d in digests] == [1, 2] + + job_rows = ( + (await verify.execute(select(Job).where(Job.project_id == project_uuid))) + .scalars() + .all() + ) + assert len(job_rows) == 1 + assert job_rows[0].status == "done" + assert job_rows[0].kind == "chain" + # 负向(DB 真源):job result 列不含正文/token。 + assert _CHAPTER_TEXT not in str(job_rows[0].result) + finally: + if project_uuid is not None: + await _cleanup(e2e_sm, project_uuid) + + +async def test_chain_conflict_interrupts_then_resume_accepts( + e2e_sm: async_sessionmaker[AsyncSession], +) -> None: + """用例 2:注入冲突 → interrupt → job awaiting → resume 带裁决 → 续跑 → accept。 + + interrupt+resume 横跨两次端点调用,靠单进程单 MemorySaver(同一 thread_id 检查点)。 + """ + conflict = Conflict(type="性格漂移", where="第2段", suggestion="统一主角姓名为「萧寒」") + adapter = _FakeChainAdapter(continuity_conflicts=[conflict]) + saver = MemorySaver() # 单实例横跨 run + resume 两次端点调用 + app = _build_app(adapter, saver, e2e_sm) + + transport = httpx.ASGITransport(app=app) + project_uuid: uuid.UUID | None = None + try: + async with LifespanManager(app): + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + create_resp = await client.post("/projects", json={"title": "链 E2E 冲突作品"}) + assert create_resp.status_code == 201 + project_id = create_resp.json()["id"] + project_uuid = uuid.UUID(project_id) + + # 1) 发起:第 1 章四审报冲突 → interrupt → job awaiting_input。 + run_resp = await client.post( + f"/projects/{project_id}/chains/draft_volume/run", + json={"start_chapter_no": 1, "count": 1}, + ) + assert run_resp.status_code == 202 + job_id = run_resp.json()["job_id"] + + job = (await client.get(f"/jobs/{job_id}")).json() + assert job["status"] == "awaiting_input", job + assert job["result"]["awaiting_chapter"] == 1 + assert job["result"]["completed"] is False + assert job["result"]["written"] == [] + + # awaiting 时第 1 章仍 draft(未经 accept 晋升),无 digest。 + async with e2e_sm() as mid: + mid_chapters = ( + ( + await mid.execute( + select(Chapter).where(Chapter.project_id == project_uuid) + ) + ) + .scalars() + .all() + ) + assert all(c.status == "draft" for c in mid_chapters) + assert mid_chapters != [] + mid_digests = ( + ( + await mid.execute( + select(ChapterDigest).where( + ChapterDigest.project_id == project_uuid + ) + ) + ) + .scalars() + .all() + ) + assert mid_digests == [] + + # 2) resume:作者裁决冲突 → 续跑 accept → done。 + resume_resp = await client.post( + f"/projects/{project_id}/chains/runs/{job_id}/resume", + json={ + "decisions": [ + {"conflict_index": 0, "verdict": "ignore", "note": "笔误,忽略"} + ] + }, + ) + assert resume_resp.status_code == 202 + + job = (await client.get(f"/jobs/{job_id}")).json() + assert job["status"] == "done", job + assert job["result"]["written"] == [1] + assert job["result"]["completed"] is True + assert job["result"]["awaiting_chapter"] is None + + # DB 真源断言:resume 后第 1 章 accepted + digest 一行。 + assert project_uuid is not None + async with e2e_sm() as verify: + chapters = ( + (await verify.execute(select(Chapter).where(Chapter.project_id == project_uuid))) + .scalars() + .all() + ) + accepted = [c for c in chapters if c.status == "accepted"] + assert {c.chapter_no for c in accepted} == {1} + + digests = ( + ( + await verify.execute( + select(ChapterDigest).where(ChapterDigest.project_id == project_uuid) + ) + ) + .scalars() + .all() + ) + assert [d.chapter_no for d in digests] == [1] + + # 负向:job 状态/result 不含正文/裁决正文外的 token/prompt。 + job_row = ( + await verify.execute(select(Job).where(Job.project_id == project_uuid)) + ).scalar_one() + assert _CHAPTER_TEXT not in str(job_row.result) + assert job_row.status == "done" + finally: + if project_uuid is not None: + await _cleanup(e2e_sm, project_uuid)