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:
227
apps/api/tests/fakes_projects.py
Normal file
227
apps/api/tests/fakes_projects.py
Normal 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
|
||||
56
apps/api/tests/fakes_providers.py
Normal file
56
apps/api/tests/fakes_providers.py
Normal 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
|
||||
59
apps/api/tests/test_credentials_crypto.py
Normal file
59
apps/api/tests/test_credentials_crypto.py
Normal 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-…••••"
|
||||
210
apps/api/tests/test_projects.py
Normal file
210
apps/api/tests/test_projects.py
Normal file
@@ -0,0 +1,210 @@
|
||||
"""T1.4 端点测试:立项 CRUD + 写章 SSE + 自动保存(内存替身,无 DB/无网络)。
|
||||
|
||||
DoD:端点入 OpenAPI;POST 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("改稿更长")
|
||||
93
apps/api/tests/test_settings_providers.py
Normal file
93
apps/api/tests/test_settings_providers.py
Normal 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
|
||||
360
apps/api/ww_api/routers/projects.py
Normal file
360
apps/api/ww_api/routers/projects.py
Normal 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 流式写章草稿(SSE,text/event-stream)。
|
||||
- PUT /projects/:id/chapters/:no/draft 自动保存草稿(幂等 upsert)。
|
||||
|
||||
不变量:writer 只产草稿(不提升 accepted 版本/不抽 digest/不写状态——属 M2 验收);
|
||||
agent 只传 tier;DB 按 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.5):gate(事务前)→ 终稿提炼 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
|
||||
|
||||
# 冲突 gate(R5,事务前拦截,不写库)。
|
||||
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,
|
||||
)
|
||||
127
apps/api/ww_api/routers/settings_providers.py
Normal file
127
apps/api/ww_api/routers/settings_providers.py
Normal file
@@ -0,0 +1,127 @@
|
||||
"""提供商凭据与档位路由端点(C3 / ARCH §4.7, §7.2;UX §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),
|
||||
)
|
||||
126
apps/api/ww_api/schemas/projects.py
Normal file
126
apps/api/ww_api/schemas/projects.py
Normal file
@@ -0,0 +1,126 @@
|
||||
"""项目(立项)与章节草稿的请求/响应 schema(C3 / 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")
|
||||
76
apps/api/ww_api/schemas/providers.py
Normal file
76
apps/api/ww_api/schemas/providers.py
Normal file
@@ -0,0 +1,76 @@
|
||||
"""提供商凭据/档位路由的请求/响应 schema(C3 / 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
|
||||
53
apps/api/ww_api/security/credentials.py
Normal file
53
apps/api/ww_api/security/credentials.py
Normal 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:]}"
|
||||
149
apps/api/ww_api/services/credentials.py
Normal file
149
apps/api/ww_api/services/credentials.py
Normal 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_id,PG 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()
|
||||
162
apps/api/ww_api/services/project_deps.py
Normal file
162
apps/api/ww_api/services/project_deps.py
Normal 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 单用户 stub(owner_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 FK(projects/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:
|
||||
"""审稿留痕 repo(review 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")
|
||||
93
apps/api/ww_api/services/provider_deps.py
Normal file
93
apps/api/ww_api/services/provider_deps.py
Normal 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_url(ARCH §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)
|
||||
Reference in New Issue
Block a user