Files
writer-work-flow/packages/llm_gateway/tests/test_structured_output.py
Yaojia Wang b523b4fd21 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 走通闭环
2026-06-18 11:38:28 +02:00

124 lines
3.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""T2.1 结构化输出接线单测output_schema → instructor → LlmResponse.parsed。
注入 fake instructor client返回 (parsed, raw_completion)),不联网。
覆盖:① 带 schema → parsed 是实例且字段正确 + 仍记 1 条 ledger
② 无 schema → parsed is None 且纯文本路径不变。
"""
from __future__ import annotations
import uuid
from types import SimpleNamespace
from typing import Any, cast
from openai import AsyncOpenAI
from pydantic import BaseModel
from ww_llm_gateway.adapters.openai_compat import OpenAICompatAdapter
from ww_llm_gateway.gateway import Gateway
from ww_llm_gateway.routing import Route
from ww_llm_gateway.types import LlmRequest, Scope, Tier
class _Out(BaseModel):
label: str
score: int
def _req(**kw: Any) -> LlmRequest:
kw.setdefault("tier", "analyst")
return LlmRequest(scope=Scope(user_id=uuid.UUID(int=1)), **kw)
def _route(_tier: Tier) -> Route:
return Route(provider="deepseek", model="deepseek-chat")
class _FakeStructured:
"""模拟 instructor.AsyncInstructorcreate_with_completion 返回 (parsed, raw)。"""
def __init__(self, parsed: BaseModel, raw: Any) -> None:
self._parsed = parsed
self._raw = raw
self.calls = 0
self.last_response_model: Any = None
async def create_with_completion(
self, *, messages: Any, response_model: Any, **kw: Any
) -> tuple[BaseModel, Any]:
self.calls += 1
self.last_response_model = response_model
return self._parsed, self._raw
def _raw_with_usage() -> Any:
return SimpleNamespace(
usage=SimpleNamespace(
prompt_tokens=42,
completion_tokens=7,
prompt_tokens_details=SimpleNamespace(cached_tokens=5),
)
)
def _client() -> AsyncOpenAI:
return cast(AsyncOpenAI, SimpleNamespace(chat=SimpleNamespace(completions=None)))
async def test_complete_with_schema_returns_parsed_instance() -> None:
structured = _FakeStructured(_Out(label="ok", score=9), _raw_with_usage())
adapter = OpenAICompatAdapter("deepseek", _client(), structured_client=structured)
result = await adapter.complete(_req(input="x", output_schema=_Out), "deepseek-chat")
assert isinstance(result.parsed, _Out)
assert result.parsed.label == "ok"
assert result.parsed.score == 9
assert structured.last_response_model is _Out
# usage 从 raw completion 提取
assert result.usage.input_tokens == 42
assert result.usage.output_tokens == 7
assert result.usage.cache_read_tokens == 5
async def test_complete_without_schema_keeps_parsed_none() -> None:
response = SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content="纯文本"))],
usage=SimpleNamespace(prompt_tokens=3, completion_tokens=2, prompt_tokens_details=None),
)
class _Completions:
async def create(self, **kw: Any) -> Any:
return response
client = cast(
AsyncOpenAI,
SimpleNamespace(chat=SimpleNamespace(completions=_Completions())),
)
adapter = OpenAICompatAdapter("deepseek", client)
result = await adapter.complete(_req(input="x"), "deepseek-chat")
assert result.parsed is None
assert result.text == "纯文本"
class _FakeLedger:
def __init__(self) -> None:
self.records: list[Any] = []
async def record(self, scope: Any, usage: Any) -> None:
self.records.append(usage)
async def test_gateway_run_passes_parsed_through_and_records_once() -> None:
structured = _FakeStructured(_Out(label="hit", score=1), _raw_with_usage())
adapter = OpenAICompatAdapter("deepseek", _client(), structured_client=structured)
ledger = _FakeLedger()
gw = Gateway({"deepseek": adapter}, ledger, resolver=_route)
resp = await gw.run(_req(input="x", output_schema=_Out))
assert isinstance(resp.parsed, _Out)
assert resp.parsed.label == "hit"
assert len(ledger.records) == 1