Files
writer-work-flow/packages/llm_gateway/tests/test_openai_adapter.py

137 lines
4.4 KiB
Python

"""OpenAI 兼容适配器单测:用替身 client 验证消息翻译 + usage 映射(不联网)。"""
from __future__ import annotations
import uuid
from collections.abc import AsyncIterator
from types import SimpleNamespace
from typing import Any, cast
import pytest
from openai import AsyncOpenAI
from ww_llm_gateway.adapters.openai_compat import OpenAICompatAdapter, _messages
from ww_llm_gateway.types import Block, LlmRequest, Scope
def _req(**kw: Any) -> LlmRequest:
kw.setdefault("tier", "writer")
return LlmRequest(scope=Scope(user_id=uuid.UUID(int=1)), **kw)
class _FakeStream:
def __init__(self, chunks: list[Any]) -> None:
self._chunks = chunks
async def __aiter__(self) -> AsyncIterator[Any]:
for c in self._chunks:
yield c
class _FakeCompletions:
def __init__(self, response: Any, stream_chunks: list[Any]) -> None:
self._response = response
self._stream_chunks = stream_chunks
self.last_messages: Any = None
async def create(self, **kw: Any) -> Any:
self.last_messages = kw.get("messages")
if kw.get("stream"):
return _FakeStream(self._stream_chunks)
return self._response
class _FakeClient:
def __init__(self, completions: _FakeCompletions) -> None:
self.chat = SimpleNamespace(completions=completions)
def test_messages_put_system_first_then_user() -> None:
req = _req(input="正文", system=[Block(text="世界观硬规则", cache=True)])
msgs = _messages(req)
assert msgs[0]["role"] == "system"
assert msgs[0]["content"] == "世界观硬规则"
assert msgs[1]["role"] == "user"
assert msgs[1]["content"] == "正文"
async def test_complete_maps_text_and_usage() -> None:
response = SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content="草稿"))],
usage=SimpleNamespace(
prompt_tokens=120,
completion_tokens=80,
prompt_tokens_details=SimpleNamespace(cached_tokens=30),
),
)
completions = _FakeCompletions(response, [])
adapter = OpenAICompatAdapter("deepseek", cast(AsyncOpenAI, _FakeClient(completions)))
result = await adapter.complete(_req(input="x"), "deepseek-chat")
assert result.text == "草稿"
assert result.usage.input_tokens == 120
assert result.usage.output_tokens == 80
assert result.usage.cache_read_tokens == 30
async def test_stream_yields_text_then_final_usage() -> None:
chunks = [
SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(content=""))], usage=None),
SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(content=""))], usage=None),
SimpleNamespace(
choices=[],
usage=SimpleNamespace(
prompt_tokens=10,
completion_tokens=2,
prompt_tokens_details=None,
),
),
]
completions = _FakeCompletions(None, chunks)
adapter = OpenAICompatAdapter("deepseek", cast(AsyncOpenAI, _FakeClient(completions)))
texts: list[str] = []
usage_seen = None
async for ch in adapter.stream(_req(input="x"), "deepseek-chat"):
if ch.text:
texts.append(ch.text)
if ch.usage is not None:
usage_seen = ch.usage
assert "".join(texts) == "他说"
assert usage_seen is not None
assert usage_seen.output_tokens == 2
class _FakeModels:
def __init__(self, *, error: Exception | None = None) -> None:
self._error = error
self.calls = 0
async def list(self) -> Any:
self.calls += 1
if self._error is not None:
raise self._error
return SimpleNamespace(data=[])
async def test_probe_connection_lists_models() -> None:
# CR-M1.4:连通性探测经公开 `probe_connection`(不再让调用方反手私有 `_client`)。
models = _FakeModels()
client = SimpleNamespace(models=models)
adapter = OpenAICompatAdapter("deepseek", cast(AsyncOpenAI, client))
await adapter.probe_connection()
assert models.calls == 1
async def test_probe_connection_propagates_failure() -> None:
# 探测失败原样上抛(调用方据此映射 LLM 不可用)。
models = _FakeModels(error=RuntimeError("bad key"))
client = SimpleNamespace(models=models)
adapter = OpenAICompatAdapter("deepseek", cast(AsyncOpenAI, client))
with pytest.raises(RuntimeError):
await adapter.probe_connection()