1. resume_chain 回报真 chain_key:从 awaiting job.result 读回(run 时 _chain_result 持久化),不再按 job.kind 推断(kind 恒为通用常量 "chain",从不在 SUPPORTED_CHAINS → 总错回退 draft_volume)。continue_volume 的 resume 回执现报真值。 2. 模板 title/body 拒纯空白:StringConstraints(strip_whitespace, min_length=1), " " strip 后为空 → 422。 3. 续写链 resume 路径 E2E:continue_volume 第 1 章冲突 → interrupt → awaiting → resume 带裁决 → done,断言 resume 回执 / job result chain_key="continue_volume" (回归守卫 #1)。 测试:+resume continue_volume 单测(chain_key 真值)+2 模板空白 422 单测 +1 续写链 resume E2E。重生成 TS 客户端(仅 description 文案变化)。
399 lines
17 KiB
Python
399 lines
17 KiB
Python
"""F2 续写式链 continue_volume 端到端(真 pg + mock 网关零 token)。
|
||
|
||
证明续写式链闭环(不变量 #1/#5):`continue_volume` 每章写作以**上一章 accepted 正文末尾**
|
||
作前文引子(复用 `build_continuation_context`)。两章无冲突全自动:
|
||
- 第二章写作请求的上下文须含第一章 accepted 正文(断前文注入);
|
||
- DB 真源:两章 accepted + 两 digest 行。
|
||
回归守卫:`draft_volume`(非续写)第二章写作请求**不应**含第一章正文(区分两模式)。
|
||
|
||
镜像 `tests/test_chain_workflow_e2e.py`:真 `Gateway` + 假适配器(据 `req.output_schema` 分支
|
||
返 parsed/text,绝不联网)+ 真 `SqlAlchemyLedgerSink`;MemorySaver 检查点。无 pg → skip。
|
||
|
||
关键差异:写章假适配器**逐章返不同正文**(含章号标记),并**记录每次写章请求的 input**——
|
||
故可断言第二章写章请求的 input 含第一章 accepted 正文标记(前文注入),而 draft_volume 不含。
|
||
"""
|
||
|
||
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
|
||
|
||
_PROVIDER = "deepseek"
|
||
_USAGE = ProviderUsage(input_tokens=13, output_tokens=7)
|
||
|
||
# 每章正文含唯一章号标记,供「前文注入」断言(第二章请求须含第一章标记)。
|
||
_CHAPTER_MARK = "【第{n}章正文END】"
|
||
|
||
_FORESHADOW_REVIEW = ForeshadowReview(planted=[], resolved=[])
|
||
_PACE_REVIEW = PaceReview(water=[], hook=True, beat_map=[1, 2, 3])
|
||
_STYLE_DRIFT = StyleDriftReview(score=88, segments=[])
|
||
|
||
|
||
class _RecordingChainAdapter:
|
||
"""实现 `ProviderAdapter`:写章逐章返带章号标记的正文 + 记录每次写章请求 input。
|
||
|
||
- `output_schema is None`(写章 `gateway.run`):自增计数返第 N 章正文,记录 req.input。
|
||
- 四审 schema → 无冲突 parsed(让图全自动跑完)。
|
||
- 其余(ChapterDigestFacts)→ 终稿 digest 提炼。
|
||
"""
|
||
|
||
provider = _PROVIDER
|
||
|
||
def __init__(self, *, continuity_conflicts: list[Conflict] | None = None) -> None:
|
||
self._write_count = 0
|
||
# (request_input, returned_text) 每次写章一条,按写章顺序。
|
||
self.write_calls: list[tuple[str, str]] = []
|
||
self._continuity = ContinuityReview(conflicts=continuity_conflicts or [])
|
||
|
||
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:
|
||
self._write_count += 1
|
||
text = f"灵气如潮{_CHAPTER_MARK.format(n=self._write_count)}"
|
||
self.write_calls.append((str(req.input), text))
|
||
return ProviderResult(text=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
|
||
)
|
||
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(usage=_USAGE)
|
||
|
||
|
||
@pytest.fixture
|
||
async def e2e_sm() -> AsyncIterator[async_sessionmaker[AsyncSession]]:
|
||
"""真 DB session 工厂;无 pg 时跳过。"""
|
||
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: _RecordingChainAdapter, session: AsyncSession) -> Gateway:
|
||
return Gateway(
|
||
adapters={_PROVIDER: adapter},
|
||
ledger=SqlAlchemyLedgerSink(session),
|
||
resolver=resolve_route,
|
||
)
|
||
|
||
|
||
def _chain_gateway_override(
|
||
adapter: _RecordingChainAdapter,
|
||
) -> Callable[[AsyncSession], Gateway]:
|
||
def _override(session: Annotated[AsyncSession, Depends(get_session)]) -> Gateway:
|
||
return _gateway_for(adapter, session)
|
||
|
||
return _override
|
||
|
||
|
||
def _builder_override(
|
||
adapter: _RecordingChainAdapter,
|
||
) -> Callable[[], Callable[[AsyncSession], Awaitable[Gateway]]]:
|
||
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]]]:
|
||
@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:
|
||
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: _RecordingChainAdapter, saver: MemorySaver, e2e_sm: Any) -> Any:
|
||
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_chain_gateway_builder,
|
||
get_digest_gateway_builder,
|
||
get_session_factory,
|
||
)
|
||
|
||
app = create_app()
|
||
app.dependency_overrides[get_chain_gateway] = _chain_gateway_override(adapter)
|
||
app.dependency_overrides[get_chain_gateway_builder] = _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] = _builder_override(adapter)
|
||
return app
|
||
|
||
|
||
async def _run_two_chapter_chain(
|
||
adapter: _RecordingChainAdapter,
|
||
e2e_sm: async_sessionmaker[AsyncSession],
|
||
chain_key: str,
|
||
) -> tuple[uuid.UUID, dict[str, Any]]:
|
||
"""跑两章链(无冲突全自动)→ 返 (project_uuid, job result)。调用方负责 cleanup。"""
|
||
saver = MemorySaver()
|
||
app = _build_app(adapter, saver, e2e_sm)
|
||
transport = httpx.ASGITransport(app=app)
|
||
async with LifespanManager(app):
|
||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
|
||
create_resp = await client.post("/projects", json={"title": f"F2 {chain_key} 作品"})
|
||
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/{chain_key}/run",
|
||
json={"start_chapter_no": 1, "count": 2},
|
||
)
|
||
assert run_resp.status_code == 202, run_resp.text
|
||
assert run_resp.json()["chain_key"] == chain_key
|
||
job_id = run_resp.json()["job_id"]
|
||
|
||
job = (await client.get(f"/jobs/{job_id}")).json()
|
||
assert job["status"] == "done", job
|
||
return project_uuid, job["result"]
|
||
|
||
|
||
async def test_continue_volume_injects_prior_chapter_into_second_write(
|
||
e2e_sm: async_sessionmaker[AsyncSession],
|
||
) -> None:
|
||
"""用例 1:continue_volume 两章 → 第二章写章请求含第一章 accepted 正文(前文注入)。"""
|
||
adapter = _RecordingChainAdapter()
|
||
project_uuid: uuid.UUID | None = None
|
||
try:
|
||
project_uuid, result = await _run_two_chapter_chain(adapter, e2e_sm, "continue_volume")
|
||
assert result["written"] == [1, 2]
|
||
assert result["completed"] is True
|
||
|
||
# 两次写章请求被记录(按章序)。
|
||
assert len(adapter.write_calls) == 2
|
||
first_input, first_text = adapter.write_calls[0]
|
||
second_input, _second_text = adapter.write_calls[1]
|
||
|
||
# 第一章正文标记(accepted 后即此文本)。
|
||
chapter1_mark = _CHAPTER_MARK.format(n=1)
|
||
|
||
# 核心断言:第二章写章请求上下文含第一章 accepted 正文(前文引子注入,不变量 #1/#5)。
|
||
assert chapter1_mark in second_input, second_input
|
||
# 续写上下文结构标记(build_continuation_context)也应在第二章请求中。
|
||
assert "前文正文" in second_input
|
||
# 第一章无前文 → 其请求不含「第1章标记」(首章占位降级,证明非凭空注入)。
|
||
assert chapter1_mark not in first_input
|
||
|
||
# DB 真源:两章 accepted + 两 digest 行。
|
||
async with e2e_sm() as verify:
|
||
accepted = (
|
||
(
|
||
await verify.execute(
|
||
select(Chapter).where(
|
||
Chapter.project_id == project_uuid, Chapter.status == "accepted"
|
||
)
|
||
)
|
||
)
|
||
.scalars()
|
||
.all()
|
||
)
|
||
assert {c.chapter_no for c in accepted} == {1, 2}
|
||
digests = (
|
||
(
|
||
await verify.execute(
|
||
select(ChapterDigest).where(ChapterDigest.project_id == project_uuid)
|
||
)
|
||
)
|
||
.scalars()
|
||
.all()
|
||
)
|
||
assert {d.chapter_no for d in digests} == {1, 2}
|
||
finally:
|
||
if project_uuid is not None:
|
||
await _cleanup(e2e_sm, project_uuid)
|
||
|
||
|
||
async def test_draft_volume_does_not_inject_prior_chapter_into_second_write(
|
||
e2e_sm: async_sessionmaker[AsyncSession],
|
||
) -> None:
|
||
"""用例 2(回归守卫):draft_volume 第二章写章请求**不含**第一章正文(区分两模式)。"""
|
||
adapter = _RecordingChainAdapter()
|
||
project_uuid: uuid.UUID | None = None
|
||
try:
|
||
project_uuid, result = await _run_two_chapter_chain(adapter, e2e_sm, "draft_volume")
|
||
assert result["written"] == [1, 2]
|
||
|
||
assert len(adapter.write_calls) == 2
|
||
_first_input, _first_text = adapter.write_calls[0]
|
||
second_input, _ = adapter.write_calls[1]
|
||
|
||
chapter1_mark = _CHAPTER_MARK.format(n=1)
|
||
# draft_volume 非续写:第二章写章请求不应含第一章正文(仅按记忆/digest 量产)。
|
||
assert chapter1_mark not in second_input
|
||
# 亦无续写上下文结构标记。
|
||
assert "前文正文(续写须无缝承接)" not in second_input
|
||
finally:
|
||
if project_uuid is not None:
|
||
await _cleanup(e2e_sm, project_uuid)
|
||
|
||
|
||
async def test_continue_volume_conflict_interrupts_then_resume_reports_chain_key(
|
||
e2e_sm: async_sessionmaker[AsyncSession],
|
||
) -> None:
|
||
"""用例 3(resume 路径 + 回归守卫 #1):continue_volume 第 1 章冲突 → interrupt → awaiting →
|
||
resume 带裁决 → 续跑 accept → done,且 resume 受理回执 chain_key="continue_volume"。
|
||
|
||
修复前 bug:resume 端点按 job.kind 推断 chain_key(恒为通用常量 "chain",从不在
|
||
SUPPORTED_CHAINS)→ 总回退 "draft_volume"——continue_volume 的 resume 回执报错类型。
|
||
本测试断言 resume 回执 / job result 均反映真值 "continue_volume"。
|
||
interrupt+resume 横跨两次端点调用,靠单进程单 MemorySaver(同一 thread_id 检查点)。
|
||
"""
|
||
conflict = Conflict(type="性格漂移", where="第2段", suggestion="统一主角姓名为「萧寒」")
|
||
adapter = _RecordingChainAdapter(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": "F2 continue_volume resume 作品"}
|
||
)
|
||
assert create_resp.status_code == 201
|
||
project_id = create_resp.json()["id"]
|
||
project_uuid = uuid.UUID(project_id)
|
||
|
||
# 1) 发起 continue_volume:第 1 章四审报冲突 → interrupt → job awaiting_input。
|
||
run_resp = await client.post(
|
||
f"/projects/{project_id}/chains/continue_volume/run",
|
||
json={"start_chapter_no": 1, "count": 1},
|
||
)
|
||
assert run_resp.status_code == 202, run_resp.text
|
||
assert run_resp.json()["chain_key"] == "continue_volume"
|
||
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
|
||
# job result 已持久真 chain_key(resume 端点据此读回,回归守卫 #1)。
|
||
assert job["result"]["chain_key"] == "continue_volume"
|
||
|
||
# 2) resume:作者裁决冲突 → 续跑 accept → done。回执须报真 chain_key。
|
||
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, resume_resp.text
|
||
# 回归守卫 #1:resume 回执报真 chain_key,而非旧 bug 的 "draft_volume"。
|
||
assert resume_resp.json()["chain_key"] == "continue_volume"
|
||
|
||
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"]["chain_key"] == "continue_volume"
|
||
|
||
# DB 真源:resume 后第 1 章 accepted + digest 一行。
|
||
assert project_uuid is not None
|
||
async with e2e_sm() as verify:
|
||
accepted = (
|
||
(
|
||
await verify.execute(
|
||
select(Chapter).where(
|
||
Chapter.project_id == project_uuid, Chapter.status == "accepted"
|
||
)
|
||
)
|
||
)
|
||
.scalars()
|
||
.all()
|
||
)
|
||
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]
|
||
finally:
|
||
if project_uuid is not None:
|
||
await _cleanup(e2e_sm, project_uuid)
|