137 lines
4.4 KiB
Python
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()
|