- 薄自建 LLM 网关:OpenAI 兼容适配器(DeepSeek) + instructor 结构化输出 + usage_ledger 记账 + 档位路由 - 记忆服务 assemble:确定性选择(显式+主角+近况) + 渲染卡 + 缓存断点(中性文本) - LangGraph 写章节点 + Postgres checkpointer + SSE 归一(token/done/error) - API:立项 + 写章 draft(SSE) + PUT 自动保存 + 提供商凭据(Fernet 加密/测试连接) - 前端:AppShell + 作品库 + 5 步立项向导 + 写作工作台(流式打字机+自动保存) + 设置页 - M1 E2E:真实 DB + mock 网关零 token 走通闭环
257 lines
11 KiB
Python
257 lines
11 KiB
Python
"""M1 端到端:立项 → 写一章草稿(mock 网关) → 自动保存(真实 DB,零 token)。
|
||
|
||
证明 M1 闭环:`POST /projects` → `GET /projects/{id}` → `POST .../draft`(SSE) →
|
||
`PUT .../draft`(自动保存),且 DB 为真源(projects/chapters/usage_ledger 落库)。
|
||
|
||
确定性 & 零成本:网关用**真实 `Gateway`** 但注入一个吐固定 `Delta` 的假适配器
|
||
(不联网、不花 token),其 `ledger` 是**真实 `SqlAlchemyLedgerSink`** —— 闭环含用量
|
||
记账也被走通。无 DB 时跳过(对齐 `tests/test_jobs_integration.py`)。
|
||
|
||
坑(见 memory/gotchas):
|
||
- `ASGITransport` 不跑 lifespan → 用 `asgi-lifespan` 的 `LifespanManager` 触发
|
||
`seed_stub_user`,否则 `projects.owner_id`/`usage_ledger.owner_id` FK 报错。
|
||
- `get_sessionmaker` 缓存的 async engine 绑定首个事件循环;每个 DB 测试清缓存重建、
|
||
结束 dispose。
|
||
- 网关 `SqlAlchemyLedgerSink` 只 `flush()` 不 `commit()`(写库事务由编排层控制,
|
||
见不变量);draft 端点在 SSE 流耗尽后对**请求 session** `commit()`,usage_ledger
|
||
方落库。本测试用请求 session 做 ledger,故无需手动提交——即对该提交的回归校验。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from collections.abc import AsyncIterator
|
||
from typing import Annotated
|
||
|
||
import httpx
|
||
import pytest
|
||
from asgi_lifespan import LifespanManager
|
||
from fastapi import Depends
|
||
from sqlalchemy import delete, func, select
|
||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
||
from ww_db import get_session, get_sessionmaker
|
||
from ww_db.models import Chapter, 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 档位的全局默认 provider(config.tier_defaults["writer"] = "deepseek:...")。
|
||
_WRITER_PROVIDER = "deepseek"
|
||
# 确定性流式 token(不联网、不花 token)。
|
||
_TOKENS = ["第", "一", "章", ":", "开端。"]
|
||
# 假适配器在末尾块回报的 token 数(喂记账,证明用量闭环走通)。
|
||
_FAKE_INPUT_TOKENS = 7
|
||
_FAKE_OUTPUT_TOKENS = 5
|
||
|
||
|
||
class _FakeStreamingAdapter:
|
||
"""实现 `ProviderAdapter` Protocol:吐固定 `StreamChunk`,绝不联网。
|
||
|
||
末尾块带 `ProviderUsage` → 网关据此记账(cost 经 pricing 表,未知 model→0)。
|
||
"""
|
||
|
||
provider = _WRITER_PROVIDER
|
||
|
||
def capabilities(self) -> Capabilities:
|
||
return Capabilities()
|
||
|
||
async def complete(self, req: LlmRequest, model: str) -> ProviderResult:
|
||
# M1 草稿走 stream;complete 不参与本 E2E,仍给确定性实现。
|
||
return ProviderResult(
|
||
text="".join(_TOKENS),
|
||
usage=ProviderUsage(input_tokens=_FAKE_INPUT_TOKENS, output_tokens=_FAKE_OUTPUT_TOKENS),
|
||
)
|
||
|
||
async def stream(self, req: LlmRequest, model: str) -> AsyncIterator[StreamChunk]:
|
||
for token in _TOKENS:
|
||
yield StreamChunk(text=token)
|
||
# 末尾用量块(text=""):网关收尾时落 1 条 usage_ledger。
|
||
yield StreamChunk(
|
||
usage=ProviderUsage(input_tokens=_FAKE_INPUT_TOKENS, output_tokens=_FAKE_OUTPUT_TOKENS)
|
||
)
|
||
|
||
|
||
@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 _parse_sse(raw: str) -> list[tuple[str, str]]:
|
||
"""把 text/event-stream 原文解析为 `(event, data)` 帧列表。"""
|
||
frames: list[tuple[str, str]] = []
|
||
event: str | None = None
|
||
data: str | None = None
|
||
for line in raw.splitlines():
|
||
if line.startswith("event:"):
|
||
event = line[len("event:") :].strip()
|
||
elif line.startswith("data:"):
|
||
data = line[len("data:") :].strip()
|
||
elif line == "":
|
||
if event is not None and data is not None:
|
||
frames.append((event, data))
|
||
event, data = None, None
|
||
if event is not None and data is not None:
|
||
frames.append((event, data))
|
||
return frames
|
||
|
||
|
||
async def test_m1_closed_loop_project_to_draft_to_autosave(
|
||
e2e_sm: async_sessionmaker[AsyncSession],
|
||
) -> None:
|
||
import json
|
||
|
||
from ww_api.main import create_app
|
||
from ww_api.services.project_deps import get_writer_gateway
|
||
|
||
# 注入:真实 Gateway + 假适配器 + 真实 ledger。ledger 用**请求 session**
|
||
# (`Depends(get_session)` 被 FastAPI 按请求缓存,与 draft 端点同一实例)——
|
||
# 故由端点流末的 `session.commit()` 落库,不在测试里手动提交。这正是对端点
|
||
# 提交修复的回归校验:若端点不提交,下面 usage_ledger 断言会失败。
|
||
def _override_gateway(
|
||
session: Annotated[AsyncSession, Depends(get_session)],
|
||
) -> Gateway:
|
||
return Gateway(
|
||
adapters={_WRITER_PROVIDER: _FakeStreamingAdapter()},
|
||
ledger=SqlAlchemyLedgerSink(session),
|
||
resolver=resolve_route,
|
||
)
|
||
|
||
app = create_app()
|
||
app.dependency_overrides[get_writer_gateway] = _override_gateway
|
||
|
||
transport = httpx.ASGITransport(app=app)
|
||
# LifespanManager 触发 lifespan → seed_stub_user(owner_id FK 依赖它)。
|
||
async with LifespanManager(app):
|
||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
|
||
# 1) 立项 → 201。唯一标题避免跨次干扰。
|
||
title = "M1 E2E 闭环验证作品"
|
||
create_resp = await client.post(
|
||
"/projects",
|
||
json={
|
||
"title": title,
|
||
"genre": "玄幻",
|
||
"logline": "少年逆袭",
|
||
"selling_points": ["爽点密集"],
|
||
},
|
||
)
|
||
assert create_resp.status_code == 201
|
||
created = create_resp.json()
|
||
project_id = created["id"]
|
||
assert created["title"] == title
|
||
|
||
# 2) GET 详情 → 200,与立项一致。
|
||
get_resp = await client.get(f"/projects/{project_id}")
|
||
assert get_resp.status_code == 200
|
||
fetched = get_resp.json()
|
||
assert fetched["id"] == project_id
|
||
assert fetched["title"] == title
|
||
assert fetched["genre"] == "玄幻"
|
||
assert fetched["selling_points"] == ["爽点密集"]
|
||
|
||
# 3) 流式写章草稿 → 消费 SSE:≥1 token 帧 + 终结 done 帧,重组文本。
|
||
draft_resp = await client.post(f"/projects/{project_id}/chapters/1/draft")
|
||
assert draft_resp.status_code == 200
|
||
assert draft_resp.headers["content-type"].startswith("text/event-stream")
|
||
frames = _parse_sse(draft_resp.text)
|
||
token_frames = [d for (ev, d) in frames if ev == "token"]
|
||
done_frames = [d for (ev, d) in frames if ev == "done"]
|
||
error_frames = [d for (ev, d) in frames if ev == "error"]
|
||
assert len(token_frames) >= 1
|
||
assert len(done_frames) == 1
|
||
assert error_frames == []
|
||
streamed_text = "".join(json.loads(d)["text"] for d in token_frames)
|
||
assert streamed_text == "".join(_TOKENS)
|
||
# done 帧带累计长度。
|
||
assert json.loads(done_frames[0])["length"] == len(streamed_text)
|
||
# 不手动提交:draft 端点流末 `session.commit()` 已把 usage_ledger 落库。
|
||
|
||
# 4) 自动保存草稿 ← 重组文本 → 200 DraftResponse{status:'draft', version:1}。
|
||
save_resp = await client.put(
|
||
f"/projects/{project_id}/chapters/1/draft",
|
||
json={"text": streamed_text},
|
||
)
|
||
assert save_resp.status_code == 200
|
||
saved = save_resp.json()
|
||
assert saved["project_id"] == project_id
|
||
assert saved["chapter_no"] == 1
|
||
assert saved["status"] == "draft"
|
||
assert saved["version"] == 1
|
||
assert saved["length"] == len(streamed_text)
|
||
|
||
# PUT 再次 → 幂等:不新增章节版本(覆盖同一行)。
|
||
save_resp2 = await client.put(
|
||
f"/projects/{project_id}/chapters/1/draft",
|
||
json={"text": streamed_text},
|
||
)
|
||
assert save_resp2.status_code == 200
|
||
assert save_resp2.json()["version"] == 1
|
||
|
||
# 5) DB 断言(经 e2e session 查询)——DB 是唯一真源。
|
||
project_uuid = created["id"]
|
||
async with e2e_sm() as verify:
|
||
# projects 行存在且字段一致。
|
||
project_row = (
|
||
await verify.execute(select(Project).where(Project.id == project_uuid))
|
||
).scalar_one()
|
||
assert project_row.title == title
|
||
|
||
# chapters 草稿行:status='draft', version=1, content=已保存文本。
|
||
chapter_rows = (
|
||
(await verify.execute(select(Chapter).where(Chapter.project_id == project_uuid)))
|
||
.scalars()
|
||
.all()
|
||
)
|
||
assert len(chapter_rows) == 1 # 幂等:两次 PUT 仍一行
|
||
chapter = chapter_rows[0]
|
||
assert chapter.chapter_no == 1
|
||
assert chapter.status == "draft"
|
||
assert chapter.version == 1
|
||
assert chapter.content == streamed_text
|
||
|
||
# usage_ledger:草稿流至少写 1 条(用量记账闭环走通)。
|
||
ledger_count = (
|
||
await verify.execute(
|
||
select(func.count())
|
||
.select_from(UsageLedger)
|
||
.where(UsageLedger.project_id == project_uuid)
|
||
)
|
||
).scalar_one()
|
||
assert ledger_count >= 1
|
||
ledger_row = (
|
||
(
|
||
await verify.execute(
|
||
select(UsageLedger).where(UsageLedger.project_id == project_uuid)
|
||
)
|
||
)
|
||
.scalars()
|
||
.first()
|
||
)
|
||
assert ledger_row is not None
|
||
assert ledger_row.provider == _WRITER_PROVIDER
|
||
assert ledger_row.input_tokens == _FAKE_INPUT_TOKENS
|
||
assert ledger_row.output_tokens == _FAKE_OUTPUT_TOKENS
|
||
|
||
# 清理:usage_ledger 的 project FK 无级联 → 先删它,再删 project(chapters 经 FK 级联)。
|
||
async with e2e_sm() as cleanup:
|
||
await cleanup.execute(delete(UsageLedger).where(UsageLedger.project_id == project_uuid))
|
||
await cleanup.execute(delete(Project).where(Project.id == project_uuid))
|
||
await cleanup.commit()
|