feat: M1 — 立项→写章草稿(SSE)→自动保存;连一家 provider

- 薄自建 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 走通闭环
This commit is contained in:
Yaojia Wang
2026-06-18 11:38:28 +02:00
parent d3dc620a71
commit b523b4fd21
70 changed files with 6642 additions and 0 deletions

View File

@@ -0,0 +1,227 @@
"""内存替身:项目/章节 repo + 写章网关(端点测试用,无 DB/无网络)。
绝对导入 `from fakes_projects import ...`(包目录无 __init__.py见 memory/gotchas
全局唯一名(避免跨包同名碰撞)。
"""
from __future__ import annotations
import uuid
from collections.abc import AsyncIterator
from typing import Any
from pydantic import BaseModel
from ww_core.domain.chapter_repo import ChapterDraftView, ChapterView
from ww_core.domain.project_repo import ProjectCreate, ProjectView
from ww_core.domain.repositories import DigestView
from ww_core.domain.review_repo import ReviewView
from ww_llm_gateway.types import Delta, LlmRequest, LlmResponse, ServedBy, Usage
class FakeProjectRepo:
"""实现 `ProjectRepo` Protocol 的内存版(按 owner_id 隔离)。"""
def __init__(self) -> None:
self.rows: dict[uuid.UUID, tuple[uuid.UUID, ProjectView]] = {}
async def create(self, owner_id: uuid.UUID, data: ProjectCreate) -> ProjectView:
pid = uuid.uuid4()
view = ProjectView(
id=pid,
title=data.title,
genre=data.genre,
logline=data.logline,
premise=data.premise,
theme=data.theme,
selling_points=list(data.selling_points),
structure=data.structure,
)
self.rows[pid] = (owner_id, view)
return view
async def list_for_owner(self, owner_id: uuid.UUID) -> list[ProjectView]:
return [v for (o, v) in self.rows.values() if o == owner_id]
async def get(self, owner_id: uuid.UUID, project_id: uuid.UUID) -> ProjectView | None:
entry = self.rows.get(project_id)
if entry is None or entry[0] != owner_id:
return None
return entry[1]
class FakeChapterRepo:
"""实现 `ChapterRepo` Protocol幂等覆盖同 (project_id, chapter_no) 草稿。"""
def __init__(self) -> None:
self.drafts: dict[tuple[uuid.UUID, int], ChapterDraftView] = {}
self.accepted_versions: dict[tuple[uuid.UUID, int], int] = {}
self.accepted_content: dict[tuple[uuid.UUID, int, int], str] = {}
async def save_draft(
self, project_id: uuid.UUID, chapter_no: int, *, text: str, volume: int = 1
) -> ChapterDraftView:
view = ChapterDraftView(
project_id=project_id,
chapter_no=chapter_no,
volume=volume,
content=text,
status="draft",
version=1,
)
self.drafts[(project_id, chapter_no)] = view
return view
async def get_draft(self, project_id: uuid.UUID, chapter_no: int) -> ChapterDraftView | None:
return self.drafts.get((project_id, chapter_no))
async def max_version(self, project_id: uuid.UUID, chapter_no: int) -> int:
versions = [
v for (p, c), v in self.accepted_versions.items() if p == project_id and c == chapter_no
]
draft = 1 if (project_id, chapter_no) in self.drafts else 0
return max([*versions, draft], default=0)
async def promote_to_accepted(
self, project_id: uuid.UUID, chapter_no: int, *, content: str, volume: int = 1
) -> ChapterView:
next_version = (await self.max_version(project_id, chapter_no)) + 1
self.accepted_versions[(project_id, chapter_no)] = next_version
self.accepted_content[(project_id, chapter_no, next_version)] = content
return ChapterView(
project_id=project_id,
chapter_no=chapter_no,
volume=volume,
content=content,
status="accepted",
version=next_version,
)
async def latest_accepted(self, project_id: uuid.UUID, chapter_no: int) -> ChapterView | None:
ver = self.accepted_versions.get((project_id, chapter_no))
if ver is None:
return None
return ChapterView(
project_id=project_id,
chapter_no=chapter_no,
volume=1,
content=self.accepted_content[(project_id, chapter_no, ver)],
status="accepted",
version=ver,
)
class FakeReviewRepo:
"""实现 `ReviewRepo` Protocol内存留痕 + 历史 + 裁决;只 flush 语义无 DB"""
def __init__(self, *, fail_set_decisions: bool = False) -> None:
self.rows: list[ReviewView] = []
self._fail_set_decisions = fail_set_decisions
async def record(
self,
project_id: uuid.UUID,
chapter_no: int,
*,
chapter_version: int | None = None,
conflicts: list[dict[str, Any]],
foreshadow_sug: list[dict[str, Any]] | None = None,
style: dict[str, Any] | None = None,
pace: dict[str, Any] | None = None,
health_score: int | None = None,
) -> ReviewView:
view = ReviewView(
id=uuid.uuid4(),
project_id=project_id,
chapter_no=chapter_no,
chapter_version=chapter_version,
conflicts=[dict(c) for c in conflicts],
foreshadow_sug=[dict(s) for s in (foreshadow_sug or [])],
style=style,
pace=pace,
health_score=health_score,
)
# 新→旧:插到列表头部。
self.rows.insert(0, view)
return view
async def list_for_chapter(self, project_id: uuid.UUID, chapter_no: int) -> list[ReviewView]:
return [r for r in self.rows if r.project_id == project_id and r.chapter_no == chapter_no]
async def set_decisions(self, review_id: uuid.UUID, *, decisions: dict[str, Any]) -> ReviewView:
if self._fail_set_decisions:
raise RuntimeError("boom: set_decisions failed (transaction rollback test)")
for i, r in enumerate(self.rows):
if r.id == review_id:
updated = r.model_copy(update={"decisions": dict(decisions)})
self.rows[i] = updated
return updated
raise LookupError(f"review not found: {review_id}")
class FakeDigestAppendRepo:
"""实现 `DigestAppendRepo` Protocol内存 append-only"""
def __init__(self) -> None:
self.rows: list[tuple[uuid.UUID, int, dict[str, Any]]] = []
async def append(
self, project_id: uuid.UUID, chapter_no: int, *, facts: dict[str, Any]
) -> DigestView:
self.rows.append((project_id, chapter_no, dict(facts)))
return DigestView(chapter_no=chapter_no, facts=dict(facts))
class FakeSession:
"""最小 fake记录 commit 次数(验收事务边界断言)。"""
def __init__(self) -> None:
self.commits = 0
async def commit(self) -> None:
self.commits += 1
class FakeReviewGateway:
"""实现审稿/digest 节点最小依赖(`run`):吐固定结构化 `parsed`,绝不联网。"""
def __init__(self, parsed: BaseModel | None = None, error: Exception | None = None) -> None:
self._parsed = parsed
self._error = error
self.requests: list[LlmRequest] = []
async def run(self, req: LlmRequest) -> LlmResponse:
self.requests.append(req)
if self._error is not None:
raise self._error
return LlmResponse(
text=self._parsed.model_dump_json() if self._parsed else "{}",
parsed=self._parsed,
usage=Usage(
provider="fake",
model="fake",
input_tokens=1,
output_tokens=1,
cost_minor=0,
currency="USD",
),
served_by=ServedBy(provider="fake", model="fake"),
)
class FakeWriterGateway:
"""实现 write 节点最小依赖(`stream`):吐固定 `Delta` 序列,绝不联网。
可注入 `error` 以验证 SSE 归一层把异常归一为 error 事件。
"""
def __init__(self, chunks: list[str] | None = None, error: Exception | None = None) -> None:
self._chunks = chunks if chunks is not None else ["第一段。", "第二段。"]
self._error = error
self.requests: list[LlmRequest] = []
async def stream(self, req: LlmRequest) -> AsyncIterator[Delta]:
self.requests.append(req)
for c in self._chunks:
yield Delta(text=c)
if self._error is not None:
raise self._error

View File

@@ -0,0 +1,56 @@
"""内存替身:凭据存储 + 提供商探测(端点测试用,无 DB/无网络)。
绝对导入 `from fakes import ...`(包目录无 __init__.py见 memory/gotchas
"""
from __future__ import annotations
import uuid
from ww_api.services.credentials import StoredCredential, StoredRouting
from ww_llm_gateway.adapters.base import Capabilities
class FakeCredentialStore:
"""实现 `CredentialStore` Protocol 的内存版。"""
def __init__(self) -> None:
# key: (owner_id, provider) → 密文
self.creds: dict[tuple[uuid.UUID, str], bytes] = {}
self.routing: dict[str, StoredRouting] = {}
async def list_credentials(self, owner_id: uuid.UUID) -> list[StoredCredential]:
return [
StoredCredential(provider=p, api_key_enc=blob)
for (o, p), blob in self.creds.items()
if o == owner_id
]
async def list_routing(self) -> list[StoredRouting]:
return list(self.routing.values())
async def get_credential(self, owner_id: uuid.UUID, provider: str) -> StoredCredential | None:
blob = self.creds.get((owner_id, provider))
if blob is None:
return None
return StoredCredential(provider=provider, api_key_enc=blob)
async def upsert_credential(
self, owner_id: uuid.UUID, provider: str, api_key_enc: bytes
) -> None:
self.creds[(owner_id, provider)] = api_key_enc
async def upsert_routing(self, routing: StoredRouting) -> None:
self.routing[routing.tier] = routing
class FakeProviderProbe:
"""实现 `ProviderProbe` Protocol返回固定能力矩阵绝不联网。"""
def __init__(self, caps: Capabilities | None = None) -> None:
self.caps = caps or Capabilities(structured_output=True, prefix_cache=True, thinking=False)
self.calls: list[tuple[uuid.UUID, str]] = []
async def probe(self, owner_id: uuid.UUID, provider: str) -> Capabilities:
self.calls.append((owner_id, provider))
return self.caps

View File

@@ -0,0 +1,59 @@
"""T1.7 凭据加密工具单测——Fernet 往返 + 掩码ARCH §4.7)。
绝不在响应/日志回显明文 Key。
"""
from __future__ import annotations
import pytest
from cryptography.fernet import Fernet
from ww_api.security.credentials import (
CredentialKeyError,
decrypt_api_key,
encrypt_api_key,
mask_api_key,
)
KEY = Fernet.generate_key().decode()
def test_encrypt_decrypt_round_trip() -> None:
# Arrange
plaintext = "sk-abc123def456"
# Act
blob = encrypt_api_key(plaintext, key=KEY)
# Assert
assert isinstance(blob, bytes)
assert plaintext.encode() not in blob # 密文不含明文
assert decrypt_api_key(blob, key=KEY) == plaintext
def test_encrypt_is_non_deterministic() -> None:
# Fernet 含随机 IV——两次加密产物不同但都能解回
a = encrypt_api_key("sk-secret", key=KEY)
b = encrypt_api_key("sk-secret", key=KEY)
assert a != b
assert decrypt_api_key(a, key=KEY) == decrypt_api_key(b, key=KEY) == "sk-secret"
def test_missing_key_fails_fast() -> None:
with pytest.raises(CredentialKeyError):
encrypt_api_key("sk-x", key="")
def test_invalid_key_fails_fast() -> None:
with pytest.raises(CredentialKeyError):
encrypt_api_key("sk-x", key="not-a-valid-fernet-key")
def test_mask_shows_only_last_four() -> None:
assert mask_api_key("sk-abcdefgh1234") == "sk-…1234"
def test_mask_short_key_fully_hidden() -> None:
# 极短 key 不泄露任何字符
assert mask_api_key("ab") == "sk-…••••"
def test_mask_empty_key() -> None:
assert mask_api_key("") == "sk-…••••"

View File

@@ -0,0 +1,210 @@
"""T1.4 端点测试:立项 CRUD + 写章 SSE + 自动保存(内存替身,无 DB/无网络)。
DoD端点入 OpenAPIPOST draft 回 SSE 流mock 网关验证PUT draft 幂等。
"""
from __future__ import annotations
import uuid
import httpx
import pytest
from fakes_projects import FakeChapterRepo, FakeProjectRepo, FakeWriterGateway
from ww_core.domain.repositories import (
CharacterView,
DigestView,
ForeshadowView,
MemoryRepos,
OutlineView,
RuleView,
StyleView,
WorldEntityView,
)
from ww_shared import ErrorCode
class _EmptyOutlineRepo:
async def get(self, project_id: uuid.UUID, chapter_no: int) -> OutlineView | None:
return None
class _EmptyCharacterRepo:
async def list_for_project(self, project_id: uuid.UUID) -> list[CharacterView]:
return []
class _EmptyWorldEntityRepo:
async def list_for_project(self, project_id: uuid.UUID) -> list[WorldEntityView]:
return []
class _EmptyDigestRepo:
async def recent(self, project_id: uuid.UUID, k: int) -> list[DigestView]:
return []
class _EmptyForeshadowRepo:
async def list_for_codes(self, project_id: uuid.UUID, codes: list[str]) -> list[ForeshadowView]:
return []
class _EmptyStyleRepo:
async def latest(self, project_id: uuid.UUID) -> StyleView | None:
return None
class _EmptyRulesRepo:
async def all_for_project(self, project_id: uuid.UUID) -> list[RuleView]:
return []
def _empty_memory_repos() -> MemoryRepos:
return MemoryRepos(
outline=_EmptyOutlineRepo(),
character=_EmptyCharacterRepo(),
world_entity=_EmptyWorldEntityRepo(),
digest=_EmptyDigestRepo(),
foreshadow=_EmptyForeshadowRepo(),
style=_EmptyStyleRepo(),
rules=_EmptyRulesRepo(),
)
def _make_client(
*,
project_repo: FakeProjectRepo | None = None,
chapter_repo: FakeChapterRepo | None = None,
gateway: FakeWriterGateway | None = None,
) -> tuple[httpx.AsyncClient, FakeProjectRepo, FakeChapterRepo, FakeWriterGateway]:
import os
os.environ.setdefault("CREDENTIAL_ENC_KEY", "x" * 44)
from ww_api.main import create_app
from ww_api.services.project_deps import (
get_chapter_repo,
get_memory_repos,
get_project_repo,
get_writer_gateway,
)
project_repo = project_repo or FakeProjectRepo()
chapter_repo = chapter_repo or FakeChapterRepo()
gateway = gateway or FakeWriterGateway()
app = create_app()
app.dependency_overrides[get_project_repo] = lambda: project_repo
app.dependency_overrides[get_chapter_repo] = lambda: chapter_repo
app.dependency_overrides[get_memory_repos] = _empty_memory_repos
app.dependency_overrides[get_writer_gateway] = lambda: gateway
transport = httpx.ASGITransport(app=app)
client = httpx.AsyncClient(transport=transport, base_url="http://test")
return client, project_repo, chapter_repo, gateway
@pytest.mark.asyncio
async def test_create_project_returns_201() -> None:
client, _, _, _ = _make_client()
async with client:
resp = await client.post(
"/projects",
json={"title": "我的网文", "genre": "玄幻", "selling_points": ["爽点密集"]},
)
assert resp.status_code == 201
body = resp.json()
assert body["title"] == "我的网文"
assert body["genre"] == "玄幻"
assert body["selling_points"] == ["爽点密集"]
assert uuid.UUID(body["id"])
@pytest.mark.asyncio
async def test_list_projects() -> None:
client, repo, _, _ = _make_client()
async with client:
await client.post("/projects", json={"title": ""})
await client.post("/projects", json={"title": ""})
resp = await client.get("/projects")
assert resp.status_code == 200
titles = {p["title"] for p in resp.json()["projects"]}
assert titles == {"", ""}
@pytest.mark.asyncio
async def test_get_project_detail() -> None:
client, _, _, _ = _make_client()
async with client:
created = (await client.post("/projects", json={"title": "详情"})).json()
resp = await client.get(f"/projects/{created['id']}")
assert resp.status_code == 200
assert resp.json()["title"] == "详情"
@pytest.mark.asyncio
async def test_get_unknown_project_404() -> None:
client, _, _, _ = _make_client()
async with client:
resp = await client.get(f"/projects/{uuid.uuid4()}")
assert resp.status_code == 404
assert resp.json()["error"]["code"] == ErrorCode.NOT_FOUND
@pytest.mark.asyncio
async def test_create_project_rejects_blank_title() -> None:
client, _, _, _ = _make_client()
async with client:
resp = await client.post("/projects", json={"title": ""})
assert resp.status_code == 422
@pytest.mark.asyncio
async def test_draft_stream_yields_sse_tokens_and_done() -> None:
gateway = FakeWriterGateway(chunks=["阿福", "走进门。"])
client, _, _, _ = _make_client(gateway=gateway)
pid = uuid.uuid4()
async with client:
resp = await client.post(f"/projects/{pid}/chapters/1/draft")
assert resp.status_code == 200
assert resp.headers["content-type"].startswith("text/event-stream")
text = resp.text
assert "event: token" in text
assert '"text": "阿福"' in text
assert "event: done" in text
# done 带累计长度(不回灌全文)
assert '"length": 6' in text
# 网关确实被调用,且只传 tier不传具体 model
assert len(gateway.requests) == 1
assert gateway.requests[0].tier == "writer"
@pytest.mark.asyncio
async def test_draft_stream_maps_error_to_sse_error_event() -> None:
from ww_shared import AppError
gateway = FakeWriterGateway(chunks=["半段"], error=AppError(ErrorCode.LLM_UNAVAILABLE, "boom"))
client, _, _, _ = _make_client(gateway=gateway)
pid = uuid.uuid4()
async with client:
resp = await client.post(f"/projects/{pid}/chapters/1/draft")
assert resp.status_code == 200
text = resp.text
assert "event: token" in text
assert "event: error" in text
assert ErrorCode.LLM_UNAVAILABLE in text
@pytest.mark.asyncio
async def test_put_draft_is_idempotent() -> None:
chapter_repo = FakeChapterRepo()
client, _, _, _ = _make_client(chapter_repo=chapter_repo)
pid = uuid.uuid4()
async with client:
r1 = await client.put(f"/projects/{pid}/chapters/3/draft", json={"text": "初稿"})
r2 = await client.put(f"/projects/{pid}/chapters/3/draft", json={"text": "改稿更长"})
assert r1.status_code == 200
assert r2.status_code == 200
# 同章节只一条草稿(版次不爆炸)
assert len(chapter_repo.drafts) == 1
saved = chapter_repo.drafts[(pid, 3)]
assert saved.content == "改稿更长"
assert saved.version == 1
assert r2.json()["length"] == len("改稿更长")

View File

@@ -0,0 +1,93 @@
"""T1.7 端点测试:凭据 upsert/列表/探测(内存替身,无 DB/无网络)。
DoD掩码往返、test-connection 经可注入探测回能力矩阵、明文绝不出响应。
"""
from __future__ import annotations
from fakes_providers import FakeCredentialStore, FakeProviderProbe
from fastapi.testclient import TestClient
from ww_api.security.credentials import decrypt_api_key
from ww_api.services.credentials import STUB_OWNER_ID
def test_get_empty_providers(client: TestClient) -> None:
resp = client.get("/settings/providers")
assert resp.status_code == 200
body = resp.json()
assert body == {"providers": [], "tier_routing": []}
def test_put_credential_encrypts_and_masks(
client: TestClient, store: FakeCredentialStore, enc_key: str
) -> None:
# Act
resp = client.put(
"/settings/providers",
json={"credentials": [{"provider": "deepseek", "api_key": "sk-secret-9999"}]},
)
# Assert: 响应含掩码,绝不含明文
assert resp.status_code == 200
body = resp.json()
assert "sk-secret-9999" not in resp.text
assert body["providers"][0]["provider"] == "deepseek"
# 存的是密文且能解回明文(往返)
blob = store.creds[(STUB_OWNER_ID, "deepseek")]
assert b"sk-secret-9999" not in blob
assert decrypt_api_key(blob, key=enc_key) == "sk-secret-9999"
def test_put_is_idempotent(client: TestClient, store: FakeCredentialStore) -> None:
payload = {"credentials": [{"provider": "deepseek", "api_key": "sk-a"}]}
client.put("/settings/providers", json=payload)
client.put(
"/settings/providers", json={"credentials": [{"provider": "deepseek", "api_key": "sk-b"}]}
)
# 同 (owner, provider) 只一条
assert len([k for k in store.creds if k[1] == "deepseek"]) == 1
def test_put_tier_routing(client: TestClient, store: FakeCredentialStore) -> None:
resp = client.put(
"/settings/providers",
json={
"tier_routing": [
{
"tier": "writer",
"provider": "deepseek",
"model": "deepseek-chat",
"fallback": ["kimi"],
}
]
},
)
assert resp.status_code == 200
routing = resp.json()["tier_routing"]
assert routing[0] == {
"tier": "writer",
"provider": "deepseek",
"model": "deepseek-chat",
"fallback": ["kimi"],
}
def test_test_connection_returns_capabilities(client: TestClient, probe: FakeProviderProbe) -> None:
resp = client.post("/settings/providers/test", json={"provider": "deepseek"})
assert resp.status_code == 200
body = resp.json()
assert body["provider"] == "deepseek"
assert body["ok"] is True
assert body["capabilities"] == {
"structured_output": True,
"prefix_cache": True,
"thinking": False,
}
assert probe.calls == [(STUB_OWNER_ID, "deepseek")]
def test_put_validation_rejects_blank_provider(client: TestClient) -> None:
resp = client.put(
"/settings/providers",
json={"credentials": [{"provider": "", "api_key": "sk-x"}]},
)
assert resp.status_code == 422

View File

@@ -0,0 +1,360 @@
"""项目(立项)+ 章节草稿端点C3 / ARCH §7.2, §7.3;不变量 #7
- POST /projects 立项向导 → projects 行owner_id=stub
- GET /projects 列出项目。
- GET /projects/:id 项目详情404 → NOT_FOUND
- POST /projects/:id/chapters/:no/draft 流式写章草稿SSEtext/event-stream
- PUT /projects/:id/chapters/:no/draft 自动保存草稿(幂等 upsert
不变量writer 只产草稿(不提升 accepted 版本/不抽 digest/不写状态——属 M2 验收);
agent 只传 tierDB 按 project_id/owner_id 过滤;日志脱敏(只记长度,不记正文)。
"""
from __future__ import annotations
import json
import uuid
from collections.abc import AsyncIterator
from typing import Annotated
from fastapi import APIRouter, Depends, Request
from fastapi.responses import StreamingResponse
from sqlalchemy.ext.asyncio import AsyncSession
from ww_core.domain.chapter_repo import ChapterRepo
from ww_core.domain.digest_repo import DigestAppendRepo
from ww_core.domain.project_repo import ProjectCreate, ProjectRepo
from ww_core.domain.repositories import MemoryRepos
from ww_core.domain.review_repo import ReviewRepo
from ww_core.memory import assemble
from ww_core.orchestrator import (
ChapterState,
SseEvent,
build_review_context,
build_review_graph,
normalize_deltas,
normalize_review,
stream_chapter_draft,
)
from ww_db import get_session
from ww_llm_gateway import Gateway
from ww_shared import AppError, ErrorCode
from ww_api.logging_config import get_logger
from ww_api.schemas.projects import (
AcceptRequest,
AcceptResponse,
DraftResponse,
DraftSaveRequest,
ProjectCreateRequest,
ProjectListResponse,
ProjectResponse,
ReviewHistoryItem,
ReviewHistoryResponse,
ReviewRequest,
)
from ww_api.services.accept_service import (
AcceptOutcome,
assert_conflicts_resolved,
run_accept_transaction,
)
from ww_api.services.credentials import STUB_OWNER_ID
from ww_api.services.digest_extraction import extract_digest_facts
from ww_api.services.project_deps import (
get_chapter_repo,
get_digest_append_repo,
get_digest_gateway,
get_memory_repos,
get_project_repo,
get_review_gateway,
get_review_repo,
get_writer_gateway,
)
log = get_logger("ww.api.projects")
router = APIRouter(prefix="/projects", tags=["projects"])
ProjectRepoDep = Annotated[ProjectRepo, Depends(get_project_repo)]
ChapterRepoDep = Annotated[ChapterRepo, Depends(get_chapter_repo)]
GatewayDep = Annotated[Gateway, Depends(get_writer_gateway)]
ReviewGatewayDep = Annotated[Gateway, Depends(get_review_gateway)]
DigestGatewayDep = Annotated[Gateway, Depends(get_digest_gateway)]
MemoryReposDep = Annotated[MemoryRepos, Depends(get_memory_repos)]
ReviewRepoDep = Annotated[ReviewRepo, Depends(get_review_repo)]
DigestRepoDep = Annotated[DigestAppendRepo, Depends(get_digest_append_repo)]
def _to_response(view: object) -> ProjectResponse:
# ProjectView 与 ProjectResponse 字段同名逐字段映射snake_case 契约)。
return ProjectResponse.model_validate(view, from_attributes=True)
@router.post("", status_code=201)
async def create_project(body: ProjectCreateRequest, repo: ProjectRepoDep) -> ProjectResponse:
view = await repo.create(
STUB_OWNER_ID,
ProjectCreate(
title=body.title,
genre=body.genre,
logline=body.logline,
premise=body.premise,
theme=body.theme,
selling_points=body.selling_points,
structure=body.structure,
),
)
log.info("project_created", project_id=str(view.id), title_len=len(view.title))
return _to_response(view)
@router.get("")
async def list_projects(repo: ProjectRepoDep) -> ProjectListResponse:
views = await repo.list_for_owner(STUB_OWNER_ID)
return ProjectListResponse(projects=[_to_response(v) for v in views])
@router.get("/{project_id}")
async def get_project(project_id: uuid.UUID, repo: ProjectRepoDep) -> ProjectResponse:
view = await repo.get(STUB_OWNER_ID, project_id)
if view is None:
raise AppError(ErrorCode.NOT_FOUND, f"project {project_id} not found")
return _to_response(view)
def _encode_sse(event: SseEvent) -> str:
"""把归一事件编码为 text/event-stream 帧:`event: <name>\\ndata: <json>\\n\\n`。"""
payload = json.dumps(event.data, ensure_ascii=False)
return f"event: {event.event}\ndata: {payload}\n\n"
@router.post("/{project_id}/chapters/{chapter_no}/draft")
async def stream_draft(
project_id: uuid.UUID,
chapter_no: int,
request: Request,
repos: MemoryReposDep,
gateway: GatewayDep,
session: Annotated[AsyncSession, Depends(get_session)],
) -> StreamingResponse:
"""流式写章草稿:组装记忆 → 网关流 → 归一为 SSE 事件 → text/event-stream。"""
request_id = getattr(request.state, "request_id", None)
context = await assemble(repos, project_id, chapter_no)
log.info(
"draft_stream_start",
project_id=str(project_id),
chapter_no=chapter_no,
request_id=request_id,
stable_core_len=len(context.stable_core),
volatile_len=len(context.volatile),
)
deltas = stream_chapter_draft(
gateway,
stable_core=context.stable_core,
volatile=context.volatile,
user_id=STUB_OWNER_ID,
project_id=project_id,
)
async def _frames() -> AsyncIterator[str]:
async for event in normalize_deltas(deltas, request_id=request_id):
yield _encode_sse(event)
# 网关在流末经 SqlAlchemyLedgerSink.record 把 usage_ledger 行 flush 进本请求 session
# sink 按设计不提交写库事务由编排层控制见不变量。draft 端点无其他写副作用,
# 故流耗尽后在此提交,确保「每次调用一条 usage_ledger」真正落库T1.9 暴露)。
await session.commit()
return StreamingResponse(
_frames(),
media_type="text/event-stream",
headers={"cache-control": "no-cache", "x-accel-buffering": "no"},
)
@router.put("/{project_id}/chapters/{chapter_no}/draft")
async def save_draft(
project_id: uuid.UUID,
chapter_no: int,
body: DraftSaveRequest,
repo: ChapterRepoDep,
) -> DraftResponse:
"""自动保存:幂等 upsert 草稿(同章节覆盖同一行,版次不爆炸)。"""
view = await repo.save_draft(project_id, chapter_no, text=body.text)
log.info(
"draft_saved",
project_id=str(project_id),
chapter_no=chapter_no,
length=len(view.content),
)
return DraftResponse(
project_id=view.project_id,
chapter_no=view.chapter_no,
volume=view.volume,
status=view.status,
version=view.version,
length=len(view.content),
)
async def _resolve_review_draft(
body: ReviewRequest,
chapter_repo: ChapterRepo,
project_id: uuid.UUID,
chapter_no: int,
) -> str:
"""取待审正文:请求体 `draft` 优先;否则回退到已保存草稿;都无 → NOT_FOUND。"""
if body.draft is not None and body.draft.strip():
return body.draft
saved = await chapter_repo.get_draft(project_id, chapter_no)
if saved is None or not saved.content.strip():
raise AppError(
ErrorCode.NOT_FOUND,
f"chapter {chapter_no} has no draft to review; provide `draft` in body",
)
return saved.content
@router.post("/{project_id}/chapters/{chapter_no}/review")
async def review_chapter(
project_id: uuid.UUID,
chapter_no: int,
body: ReviewRequest,
request: Request,
repos: MemoryReposDep,
chapter_repo: ChapterRepoDep,
review_repo: ReviewRepoDep,
gateway: ReviewGatewayDep,
session: Annotated[AsyncSession, Depends(get_session)],
) -> StreamingResponse:
"""续审SSE组审稿上下文 → 跑审稿子图 → 归一为 section/conflict/done 事件。
提交边界:网关 ledger + collect 经 review_repo.record 均只 flush端点在**流耗尽后**
`await session.commit()`(镜像 draft 端点,否则记账/留痕静默丢失)。
"""
request_id = getattr(request.state, "request_id", None)
draft = await _resolve_review_draft(body, chapter_repo, project_id, chapter_no)
context = await assemble(repos, project_id, chapter_no)
review_context = build_review_context(
draft=draft, stable_core=context.stable_core, volatile=context.volatile
)
log.info(
"review_stream_start",
project_id=str(project_id),
chapter_no=chapter_no,
request_id=request_id,
draft_len=len(draft),
review_context_len=len(review_context),
)
graph = build_review_graph(gateway, review_repo)
initial: ChapterState = {
"project_id": project_id,
"chapter_no": chapter_no,
"user_id": STUB_OWNER_ID,
"review_context": review_context,
}
async def _frames() -> AsyncIterator[str]:
final = await graph.ainvoke(initial)
reviews = final.get("reviews") or {}
async for event in normalize_review(reviews, request_id=request_id):
yield _encode_sse(event)
# collect 经 review_repo.record 落 chapter_reviews只 flush+ 网关 ledger 只 flush
# → 流耗尽后在此提交,确保审稿留痕 + usage_ledger 真正落库(同 draft 端点)。
await session.commit()
return StreamingResponse(
_frames(),
media_type="text/event-stream",
headers={"cache-control": "no-cache", "x-accel-buffering": "no"},
)
@router.get("/{project_id}/chapters/{chapter_no}/reviews")
async def list_reviews(
project_id: uuid.UUID,
chapter_no: int,
review_repo: ReviewRepoDep,
) -> ReviewHistoryResponse:
"""审稿历史(新→旧):供前端审稿页加载既往审稿留痕 + 裁决。"""
views = await review_repo.list_for_chapter(project_id, chapter_no)
items = [
ReviewHistoryItem(
id=v.id,
project_id=v.project_id,
chapter_no=v.chapter_no,
chapter_version=v.chapter_version,
conflicts=v.conflicts,
foreshadow_sug=v.foreshadow_sug,
style=v.style,
pace=v.pace,
health_score=v.health_score,
decisions=v.decisions,
)
for v in views
]
return ReviewHistoryResponse(reviews=items)
@router.post("/{project_id}/chapters/{chapter_no}/accept")
async def accept_chapter(
project_id: uuid.UUID,
chapter_no: int,
body: AcceptRequest,
request: Request,
chapter_repo: ChapterRepoDep,
digest_repo: DigestRepoDep,
review_repo: ReviewRepoDep,
gateway: DigestGatewayDep,
session: Annotated[AsyncSession, Depends(get_session)],
) -> AcceptResponse:
"""验收事务 + 冲突 gate§5.5gate事务前→ 终稿提炼 digest事务外R2
单原子事务(晋升 + digest + 裁决留痕)→ 一次 commit。
"""
request_id = getattr(request.state, "request_id", None)
# R3审稿真相从领域表重读最近一条 chapter_reviews不依赖 checkpoint。
history = await review_repo.list_for_chapter(project_id, chapter_no)
latest_review = history[0] if history else None
# 冲突 gateR5事务前拦截不写库
assert_conflicts_resolved(latest_review, body.decisions)
log.info(
"accept_start",
project_id=str(project_id),
chapter_no=chapter_no,
request_id=request_id,
final_text_len=len(body.final_text),
decision_count=len(body.decisions),
has_review=latest_review is not None,
)
# R2终稿 digest 提炼在**事务外**做(别在持开事务里跨网络调 LLM
digest_facts = await extract_digest_facts(
gateway,
final_text=body.final_text,
user_id=STUB_OWNER_ID,
project_id=project_id,
chapter_no=chapter_no,
)
outcome: AcceptOutcome = await run_accept_transaction(
session=session,
chapter_repo=chapter_repo,
digest_repo=digest_repo,
review_repo=review_repo,
project_id=project_id,
chapter_no=chapter_no,
final_text=body.final_text,
digest_facts=digest_facts,
latest_review=latest_review,
decisions=body.decisions,
)
return AcceptResponse(
project_id=project_id,
chapter_no=chapter_no,
accepted_version=outcome.accepted_version,
digest_added=outcome.digest_added,
decisions_recorded=outcome.decisions_recorded,
review_id=outcome.review_id,
)

View File

@@ -0,0 +1,127 @@
"""提供商凭据与档位路由端点C3 / ARCH §4.7, §7.2UX §6.10)。
- GET /settings/providers 列出已配置提供商(掩码)+ 档位路由。
- PUT /settings/providers 幂等 upsert 凭据/档位路由,回掩码视图。
- POST /settings/providers/test 最小探测验 Key + 拉能力矩阵。
不变量:响应/日志绝不含明文 Key加密入库、仅掩码出站
"""
from __future__ import annotations
from typing import Annotated
from fastapi import APIRouter, Depends
from ww_config import get_settings
from ww_llm_gateway.adapters.base import Capabilities
from ww_api.logging_config import get_logger
from ww_api.schemas.providers import (
CapabilitiesView,
ProvidersResponse,
ProvidersUpsertRequest,
ProviderView,
TestConnectionRequest,
TestConnectionResponse,
TierRoutingView,
)
from ww_api.security.credentials import encrypt_api_key, mask_api_key
from ww_api.services.credentials import (
STUB_OWNER_ID,
CredentialStore,
ProviderProbe,
StoredCredential,
StoredRouting,
)
from ww_api.services.provider_deps import (
get_credential_store,
get_provider_probe,
)
log = get_logger("ww.api.providers")
router = APIRouter(prefix="/settings/providers", tags=["settings"])
StoreDep = Annotated[CredentialStore, Depends(get_credential_store)]
ProbeDep = Annotated[ProviderProbe, Depends(get_provider_probe)]
def _mask_credential(cred: StoredCredential, plaintext: str | None) -> ProviderView:
# 存储层只有密文;掩码需明文末四位——若没有则全隐占位。
masked = mask_api_key(plaintext) if plaintext is not None else mask_api_key("")
return ProviderView(provider=cred.provider, masked_key=masked)
def _routing_view(r: StoredRouting) -> TierRoutingView:
return TierRoutingView(tier=r.tier, provider=r.provider, model=r.model, fallback=r.fallback)
async def _build_response(store: CredentialStore) -> ProvidersResponse:
creds = await store.list_credentials(STUB_OWNER_ID)
routing = await store.list_routing()
# 列表视图无明文,掩码占位(不解密历史密文以免无谓 IO/泄露面)。
providers = [ProviderView(provider=c.provider, masked_key=mask_api_key("")) for c in creds]
return ProvidersResponse(
providers=providers,
tier_routing=[_routing_view(r) for r in routing],
)
@router.get("")
async def list_providers(store: StoreDep) -> ProvidersResponse:
return await _build_response(store)
@router.put("")
async def upsert_providers(
body: ProvidersUpsertRequest,
store: StoreDep,
) -> ProvidersResponse:
enc_key = get_settings().credential_enc_key
for cred in body.credentials:
api_key_enc = encrypt_api_key(cred.api_key, key=enc_key)
await store.upsert_credential(STUB_OWNER_ID, cred.provider, api_key_enc)
# 仅记 provider + 末四位掩码,绝不记明文
log.info(
"provider_credential_upserted",
provider=cred.provider,
masked_key=mask_api_key(cred.api_key),
)
for routing in body.tier_routing:
await store.upsert_routing(
StoredRouting(
tier=routing.tier,
provider=routing.provider,
model=routing.model,
fallback=routing.fallback,
)
)
log.info(
"tier_routing_upserted",
tier=routing.tier,
provider=routing.provider,
model=routing.model,
)
return await _build_response(store)
def _capabilities_view(caps: Capabilities) -> CapabilitiesView:
return CapabilitiesView(
structured_output=caps.structured_output,
prefix_cache=caps.prefix_cache,
thinking=caps.thinking,
)
@router.post("/test")
async def test_connection(
body: TestConnectionRequest,
probe: ProbeDep,
) -> TestConnectionResponse:
caps = await probe.probe(STUB_OWNER_ID, body.provider)
log.info("provider_probe_ok", provider=body.provider)
return TestConnectionResponse(
provider=body.provider,
ok=True,
capabilities=_capabilities_view(caps),
)

View File

@@ -0,0 +1,126 @@
"""项目(立项)与章节草稿的请求/响应 schemaC3 / ARCH §7.2)。
snake_case前端经 OpenAPI 生成 TS 类型消费。改字段 → 前端必须 `pnpm gen:api`。
"""
from __future__ import annotations
import uuid
from typing import Any, Literal
from pydantic import BaseModel, Field
class ProjectCreateRequest(BaseModel):
"""POST /projects立项向导字段owner_id 由后端补 stub不入参"""
title: str = Field(min_length=1)
genre: str | None = None
logline: str | None = None
premise: str | None = None
theme: str | None = None
selling_points: list[Any] = Field(default_factory=list)
structure: str | None = None
class ProjectResponse(BaseModel):
"""项目视图(创建/列表/详情共用)。"""
id: uuid.UUID
title: str
genre: str | None = None
logline: str | None = None
premise: str | None = None
theme: str | None = None
selling_points: list[Any] = Field(default_factory=list)
structure: str | None = None
class ProjectListResponse(BaseModel):
"""GET /projects项目列表。"""
projects: list[ProjectResponse] = Field(default_factory=list)
class DraftSaveRequest(BaseModel):
"""PUT /projects/:id/chapters/:no/draft自动保存草稿正文。"""
text: str
class DraftResponse(BaseModel):
"""草稿保存结果(脱敏:只回元信息 + 长度,不回灌正文以外的衍生)。"""
project_id: uuid.UUID
chapter_no: int
volume: int
status: str
version: int
length: int
# ---- 审稿T2.5----
class ReviewRequest(BaseModel):
"""POST /projects/:id/chapters/:no/review可选携带待审草稿正文。
不传 `draft` 时端点回退到已保存的草稿chapter_repo
"""
draft: str | None = None
class ReviewHistoryItem(BaseModel):
"""单条审稿留痕GET .../reviews 历史项snake_case"""
id: uuid.UUID
project_id: uuid.UUID
chapter_no: int
chapter_version: int | None = None
conflicts: list[dict[str, Any]] = Field(default_factory=list)
foreshadow_sug: list[dict[str, Any]] = Field(default_factory=list)
style: dict[str, Any] | None = None
pace: dict[str, Any] | None = None
health_score: int | None = None
decisions: dict[str, Any] | None = None
class ReviewHistoryResponse(BaseModel):
"""GET /projects/:id/chapters/:no/reviews审稿历史新→旧"""
reviews: list[ReviewHistoryItem] = Field(default_factory=list)
# ---- 验收T2.4----
# 每个冲突的裁决:采纳改法 / 忽略 / 手改R5
Verdict = Literal["accept", "ignore", "manual"]
class ConflictDecision(BaseModel):
"""对最近一次审稿留痕里**某个冲突**(按其在 conflicts 列表的下标定位)的裁决。"""
conflict_index: int = Field(ge=0, description="冲突在最近审稿 conflicts 列表中的下标")
verdict: Verdict = Field(description="采纳改法 / 忽略 / 手改")
note: str | None = Field(default=None, description="可选裁决备注(如手改说明)")
class AcceptRequest(BaseModel):
"""POST /projects/:id/chapters/:no/accept裁决清单 + 可能改过的终稿。"""
final_text: str = Field(min_length=1, description="作者裁决/改稿后的最终验收文本")
decisions: list[ConflictDecision] = Field(
default_factory=list, description="对最近审稿每个冲突的裁决每冲突必有其一R5"
)
class AcceptResponse(BaseModel):
"""验收「本次将更新」清单ARCH §7.2 写回结果)。"""
project_id: uuid.UUID
chapter_no: int
accepted_version: int = Field(description="晋升到的 accepted 版次max+1")
digest_added: bool = Field(description="是否新增了一行 chapter_digests")
decisions_recorded: int = Field(description="本次写回的裁决条数")
review_id: uuid.UUID | None = Field(default=None, description="写回裁决的审稿留痕行 id")

View File

@@ -0,0 +1,76 @@
"""提供商凭据/档位路由的请求/响应 schemaC3 / ARCH §4.7, §7.2)。
snake_case响应一律 **掩码 Key**,绝不含明文。前端经 OpenAPI 生成 TS 类型消费。
"""
from __future__ import annotations
from pydantic import BaseModel, Field
class ProviderView(BaseModel):
"""已配置提供商(掩码视图)。"""
provider: str
masked_key: str # 形如 sk-…1234明文永不出现
class TierRoutingView(BaseModel):
"""档位 → provider:model 路由(含回退链)。"""
tier: str
provider: str
model: str
fallback: list[str] = Field(default_factory=list)
class ProvidersResponse(BaseModel):
"""GET/PUT 响应:已配置提供商(掩码)+ 当前档位路由。"""
providers: list[ProviderView] = Field(default_factory=list)
tier_routing: list[TierRoutingView] = Field(default_factory=list)
class ProviderCredentialInput(BaseModel):
"""单条提供商凭据写入。"""
provider: str = Field(min_length=1)
api_key: str = Field(min_length=1) # 明文入站,加密入库,绝不回显
class TierRoutingInput(BaseModel):
"""单条档位路由写入。"""
tier: str = Field(min_length=1)
provider: str = Field(min_length=1)
model: str = Field(min_length=1)
fallback: list[str] = Field(default_factory=list)
class ProvidersUpsertRequest(BaseModel):
"""PUT 请求:可同时 upsert 若干凭据与档位路由(幂等)。"""
credentials: list[ProviderCredentialInput] = Field(default_factory=list)
tier_routing: list[TierRoutingInput] = Field(default_factory=list)
class TestConnectionRequest(BaseModel):
"""POST /test最小探测请求。"""
provider: str = Field(min_length=1)
class CapabilitiesView(BaseModel):
"""探测得到的能力矩阵(镜像网关 `Capabilities`)。"""
structured_output: bool = False
prefix_cache: bool = False
thinking: bool = False
class TestConnectionResponse(BaseModel):
"""POST /test 响应:连通性 + 能力矩阵。"""
provider: str
ok: bool
capabilities: CapabilitiesView

View File

@@ -0,0 +1,53 @@
"""提供商 API Key 的对称加解密与掩码ARCH §4.7)。
- 加密用 `cryptography` 的 Fernet密钥取 `settings.credential_enc_key`。
- 缺失/非法密钥 → 立即失败(`CredentialKeyError`),绝不静默。
- `mask_api_key` 产出可安全展示/落库回显的掩码;明文永不出响应、永不进日志。
"""
from __future__ import annotations
from cryptography.fernet import Fernet, InvalidToken
# 掩码常量:前缀 + 省略号 + 末四位。短 key 用 bullet 占位,绝不泄露字符。
_MASK_PREFIX = "sk-…"
_MASK_PLACEHOLDER = "••••"
_LAST_N = 4
class CredentialKeyError(RuntimeError):
"""加密密钥缺失或非法——快速失败,调用方负责映射为 5xx/配置错误。"""
def _fernet(key: str) -> Fernet:
if not key:
raise CredentialKeyError(
"credential_enc_key 未配置;无法加解密提供商凭据。请设置环境变量。"
)
try:
return Fernet(key.encode())
except (ValueError, TypeError) as exc:
# 不回显 key 内容
raise CredentialKeyError(
"credential_enc_key 非法(需为 32 字节 url-safe base64 Fernet key"
) from exc
def encrypt_api_key(plaintext: str, *, key: str) -> bytes:
"""加密明文 API Key返回可入库的密文字节`api_key_enc`)。"""
return _fernet(key).encrypt(plaintext.encode())
def decrypt_api_key(blob: bytes, *, key: str) -> str:
"""解密密文字节回明文。仅供探测/调用使用,绝不进响应或日志。"""
try:
return _fernet(key).decrypt(blob).decode()
except InvalidToken as exc:
raise CredentialKeyError("凭据密文无法解密(密钥不匹配或数据损坏)。") from exc
def mask_api_key(plaintext: str) -> str:
"""掩码:仅露末四位,如 `sk-…1234`;过短/空则全隐。"""
if len(plaintext) < _LAST_N:
return f"{_MASK_PREFIX}{_MASK_PLACEHOLDER}"
return f"{_MASK_PREFIX}{plaintext[-_LAST_N:]}"

View File

@@ -0,0 +1,149 @@
"""凭据存储与提供商探测的接口 + SQLAlchemy/网关实现ARCH §4.7)。
路由依赖这里的 **接口**Protocol测试注入内存替身运行时用 SQLAlchemy/网关实现。
单用户原型:`owner_id` 用固定 stub见 `STUB_OWNER_ID`),多租户化时改为按认证主体取。
"""
from __future__ import annotations
import uuid
from dataclasses import dataclass
from typing import Protocol
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from ww_db.models import ProviderCredential, TierRouting
from ww_llm_gateway.adapters.base import Capabilities
# 单用户 stub owner与网关 Scope.user_id 的 stub 约定一致UUID(int=1))。
# 多租户化时此常量由认证主体替换(见 ARCH §4.7 隔离)。
STUB_OWNER_ID = uuid.UUID(int=1)
@dataclass(frozen=True)
class StoredCredential:
"""存储层视图:含密文,绝不出 API 边界(路由仅取 provider 并掩码)。"""
provider: str
api_key_enc: bytes
@dataclass(frozen=True)
class StoredRouting:
tier: str
provider: str
model: str
fallback: list[str]
class CredentialStore(Protocol):
"""凭据 + 档位路由的读写接口(按 owner_id 隔离)。"""
async def list_credentials(self, owner_id: uuid.UUID) -> list[StoredCredential]: ...
async def list_routing(self) -> list[StoredRouting]: ...
async def get_credential(
self, owner_id: uuid.UUID, provider: str
) -> StoredCredential | None: ...
async def upsert_credential(
self, owner_id: uuid.UUID, provider: str, api_key_enc: bytes
) -> None: ...
async def upsert_routing(self, routing: StoredRouting) -> None: ...
class ProviderProbe(Protocol):
"""最小连通探测:验证 Key + 返回能力矩阵。测试注入假探测,绝不联网。"""
async def probe(self, owner_id: uuid.UUID, provider: str) -> Capabilities: ...
class SqlCredentialStore:
"""SQLAlchemy 实现:写 `provider_credentials` / `tier_routing`,幂等 upsert。"""
def __init__(self, session: AsyncSession) -> None:
self._session = session
async def list_credentials(self, owner_id: uuid.UUID) -> list[StoredCredential]:
rows = (
await self._session.execute(
select(ProviderCredential).where(ProviderCredential.owner_id == owner_id)
)
).scalars()
return [StoredCredential(provider=r.provider, api_key_enc=r.api_key_enc) for r in rows]
async def list_routing(self) -> list[StoredRouting]:
rows = (await self._session.execute(select(TierRouting))).scalars()
return [
StoredRouting(
tier=r.tier, provider=r.provider, model=r.model, fallback=list(r.fallback)
)
for r in rows
]
async def get_credential(self, owner_id: uuid.UUID, provider: str) -> StoredCredential | None:
row = (
await self._session.execute(
select(ProviderCredential).where(
ProviderCredential.owner_id == owner_id,
ProviderCredential.provider == provider,
)
)
).scalar_one_or_none()
if row is None:
return None
return StoredCredential(provider=row.provider, api_key_enc=row.api_key_enc)
async def upsert_credential(
self, owner_id: uuid.UUID, provider: str, api_key_enc: bytes
) -> None:
# 显式 read-modify-write唯一约束含可空 project_idPG ON CONFLICT
# 在 NULL 上不去重NULLS DISTINCT故不用 on_conflict。
existing = (
await self._session.execute(
select(ProviderCredential).where(
ProviderCredential.owner_id == owner_id,
ProviderCredential.project_id.is_(None),
ProviderCredential.provider == provider,
)
)
).scalar_one_or_none()
if existing is None:
self._session.add(
ProviderCredential(
owner_id=owner_id,
project_id=None,
provider=provider,
api_key_enc=api_key_enc,
)
)
else:
existing.api_key_enc = api_key_enc
await self._session.commit()
async def upsert_routing(self, routing: StoredRouting) -> None:
existing = (
await self._session.execute(
select(TierRouting).where(
TierRouting.project_id.is_(None),
TierRouting.tier == routing.tier,
)
)
).scalar_one_or_none()
if existing is None:
self._session.add(
TierRouting(
project_id=None,
tier=routing.tier,
provider=routing.provider,
model=routing.model,
fallback=routing.fallback,
)
)
else:
existing.provider = routing.provider
existing.model = routing.model
existing.fallback = routing.fallback
await self._session.commit()

View File

@@ -0,0 +1,162 @@
"""项目/章节端点的依赖装配(运行时实现)。
- `get_project_repo` / `get_chapter_repo`:把请求 session 装配成 SQLAlchemy repo。
- `get_writer_gateway`:据 writer 档位路由从已存凭据解密 → 建 OpenAI 兼容适配器
→ `Gateway`(注入 `SqlAlchemyLedgerSink` + `resolve_route`)。这是 **draft SSE 的可注入缝**——
测试经 `app.dependency_overrides[get_writer_gateway]` 注入 mock 网关(产 `Delta`,绝不联网)。
- `seed_stub_user`:幂等 seed 单用户 stubowner_id FK 依赖它,见 memory/gotchas
"""
from __future__ import annotations
from typing import Annotated
from fastapi import Depends
from openai import AsyncOpenAI
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from ww_config import get_settings
from ww_core.domain.chapter_repo import ChapterRepo, SqlChapterRepo
from ww_core.domain.digest_repo import DigestAppendRepo, SqlDigestAppendRepo
from ww_core.domain.project_repo import ProjectRepo, SqlProjectRepo
from ww_core.domain.repositories import MemoryRepos
from ww_core.domain.review_repo import ReviewRepo, SqlReviewRepo
from ww_core.memory.sql_repositories import sql_memory_repos
from ww_db import get_session
from ww_db.models import User
from ww_llm_gateway import (
Gateway,
OpenAICompatAdapter,
SqlAlchemyLedgerSink,
resolve_route,
)
from ww_llm_gateway.types import Tier
from ww_shared import AppError, ErrorCode
from ww_api.security.credentials import (
CredentialKeyError,
decrypt_api_key,
)
from ww_api.services.credentials import (
STUB_OWNER_ID,
CredentialStore,
SqlCredentialStore,
)
from ww_api.services.provider_deps import _PROVIDER_BASE_URLS
# 单用户 stub 的占位邮箱(多租户化时由真实主体替换)。
_STUB_USER_EMAIL = "stub@local"
async def seed_stub_user(session: AsyncSession) -> None:
"""幂等 seed 单用户 stub 行——所有 owner_id FKprojects/usage_ledger/...)依赖它。
无该行时插入;已存在则跳过。在 app lifespan 启动时调用一次。
"""
existing = (
await session.execute(select(User).where(User.id == STUB_OWNER_ID))
).scalar_one_or_none()
if existing is not None:
return
session.add(User(id=STUB_OWNER_ID, email=_STUB_USER_EMAIL, display_name="stub"))
await session.commit()
def get_project_repo(
session: Annotated[AsyncSession, Depends(get_session)],
) -> ProjectRepo:
return SqlProjectRepo(session)
def get_chapter_repo(
session: Annotated[AsyncSession, Depends(get_session)],
) -> ChapterRepo:
return SqlChapterRepo(session)
def get_memory_repos(
session: Annotated[AsyncSession, Depends(get_session)],
) -> MemoryRepos:
"""记忆组装的 7-repo 捆绑draft SSE 用)。测试覆盖此依赖注入内存 fake。"""
return sql_memory_repos(session)
def get_review_repo(
session: Annotated[AsyncSession, Depends(get_session)],
) -> ReviewRepo:
"""审稿留痕 reporeview SSE collect / 历史 / accept 裁决)。"""
return SqlReviewRepo(session)
def get_digest_append_repo(
session: Annotated[AsyncSession, Depends(get_session)],
) -> DigestAppendRepo:
"""章节摘要写侧 repo验收事务追加终稿 digest"""
return SqlDigestAppendRepo(session)
async def build_gateway_for_tier(
session: AsyncSession, store: CredentialStore, tier: Tier
) -> Gateway:
"""据指定档位路由解密对应 provider 凭据 → 建网关(解析器仍为全局 `resolve_route`)。
无凭据/未知 provider → `LLM_UNAVAILABLE`(友好提示,前端引导去配置)。
解析器用 `resolve_route`(按 tier 路由);这里只决定**要预备哪个 provider 的适配器**。
"""
settings = get_settings()
route = resolve_route(tier)
base_url = _PROVIDER_BASE_URLS.get(route.provider)
if base_url is None:
raise AppError(
ErrorCode.LLM_UNAVAILABLE,
f"{tier} 档位 provider {route.provider} 暂不支持",
{"provider": route.provider, "tier": tier},
)
cred = await store.get_credential(STUB_OWNER_ID, route.provider)
if cred is None:
raise AppError(
ErrorCode.LLM_UNAVAILABLE,
f"{tier} 档位 provider {route.provider} 未配置凭据,请先在设置中配置",
{"provider": route.provider, "tier": tier},
)
try:
api_key = decrypt_api_key(cred.api_key_enc, key=settings.credential_enc_key)
except CredentialKeyError as exc:
raise AppError(ErrorCode.INTERNAL, str(exc)) from exc
client = AsyncOpenAI(api_key=api_key, base_url=base_url)
adapter = OpenAICompatAdapter(provider=route.provider, client=client)
return Gateway(
adapters={route.provider: adapter},
ledger=SqlAlchemyLedgerSink(session),
resolver=resolve_route,
)
async def build_writer_gateway(session: AsyncSession, store: CredentialStore) -> Gateway:
"""据 writer 档位路由解密对应 provider 凭据 → 建网关。"""
return await build_gateway_for_tier(session, store, "writer")
async def get_writer_gateway(
session: Annotated[AsyncSession, Depends(get_session)],
) -> Gateway:
"""draft SSE 的可注入网关缝。测试覆盖此依赖注入 mock产 `Delta`,绝不联网)。"""
store = SqlCredentialStore(session)
return await build_writer_gateway(session, store)
async def get_review_gateway(
session: Annotated[AsyncSession, Depends(get_session)],
) -> Gateway:
"""续审analyst 档位)的可注入网关缝。测试经 override 注 mock产 `parsed`,绝不联网)。"""
store = SqlCredentialStore(session)
return await build_gateway_for_tier(session, store, "analyst")
async def get_digest_gateway(
session: Annotated[AsyncSession, Depends(get_session)],
) -> Gateway:
"""验收终稿 digest 提炼light 档位)的可注入网关缝。测试经 override 注 mock。"""
store = SqlCredentialStore(session)
return await build_gateway_for_tier(session, store, "light")

View File

@@ -0,0 +1,93 @@
"""提供商凭据端点的依赖装配(运行时实现)。
把 `get_session` 装配成 `SqlCredentialStore`;把网关适配器装配成探测器。
测试经 `app.dependency_overrides` 注入内存替身——不联网。
"""
from __future__ import annotations
import uuid
from typing import Annotated
from fastapi import Depends
from openai import AsyncOpenAI
from sqlalchemy.ext.asyncio import AsyncSession
from ww_config import get_settings
from ww_db import get_session
from ww_llm_gateway.adapters.base import Capabilities
from ww_llm_gateway.adapters.openai_compat import OpenAICompatAdapter
from ww_shared import AppError, ErrorCode
from ww_api.security.credentials import (
CredentialKeyError,
decrypt_api_key,
)
from ww_api.services.credentials import (
CredentialStore,
SqlCredentialStore,
)
# 已知 OpenAI 兼容提供商 → base_urlARCH §4.2)。
_PROVIDER_BASE_URLS: dict[str, str] = {
"deepseek": "https://api.deepseek.com",
"kimi": "https://api.moonshot.cn/v1",
"qwen": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"glm": "https://open.bigmodel.cn/api/paas/v4",
"openai": "https://api.openai.com/v1",
}
def get_credential_store(
session: Annotated[AsyncSession, Depends(get_session)],
) -> CredentialStore:
return SqlCredentialStore(session)
class GatewayProviderProbe:
"""运行时探测:解密 Key→建 OpenAI 兼容适配器→最小请求验 Key→回能力矩阵。
依赖 store 取密文 + settings 取加密 key。绝不在日志/响应回显明文。
"""
def __init__(self, store: CredentialStore, enc_key: str) -> None:
self._store = store
self._enc_key = enc_key
async def probe(self, owner_id: uuid.UUID, provider: str) -> Capabilities:
cred = await self._store.get_credential(owner_id, provider)
if cred is None:
raise AppError(
ErrorCode.NOT_FOUND,
f"provider {provider} 未配置凭据",
{"provider": provider},
)
base_url = _PROVIDER_BASE_URLS.get(provider)
if base_url is None:
raise AppError(
ErrorCode.VALIDATION,
f"未知提供商 {provider}",
{"provider": provider},
)
try:
api_key = decrypt_api_key(cred.api_key_enc, key=self._enc_key)
except CredentialKeyError as exc:
raise AppError(ErrorCode.INTERNAL, str(exc)) from exc
client = AsyncOpenAI(api_key=api_key, base_url=base_url)
adapter = OpenAICompatAdapter(provider=provider, client=client)
try:
# 最小探测:列模型即可验证 Key 有效(不消耗生成额度)。
await client.models.list()
except Exception as exc: # noqa: BLE001 — 任一失败都映射为 LLM 不可用
raise AppError(
ErrorCode.LLM_UNAVAILABLE,
f"provider {provider} 连接探测失败",
{"provider": provider},
) from exc
return adapter.capabilities()
def get_provider_probe(
store: Annotated[CredentialStore, Depends(get_credential_store)],
) -> GatewayProviderProbe:
return GatewayProviderProbe(store, get_settings().credential_enc_key)