175 lines
5.9 KiB
Python
175 lines
5.9 KiB
Python
"""Tests for model aliasing (Task 1) and /v1/completions text completion (Task 2).
|
|
|
|
Backends are mocked — no live model required.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from unittest.mock import AsyncMock
|
|
|
|
from llm_inference.client import LLMClient
|
|
from llm_inference.config import LLMSettings, SettingsCache
|
|
from llm_inference.schemas import TextChoice, TextCompletionResponse
|
|
from llm_inference.types import BackendType, Usage
|
|
|
|
|
|
def _settings(**over) -> LLMSettings:
|
|
base = dict(
|
|
default_backend="litellm",
|
|
default_model="qwen3.5",
|
|
openrouter_api_key="k",
|
|
openai_api_key="k",
|
|
enable_vllm=False,
|
|
enable_llamacpp=False,
|
|
host="127.0.0.1",
|
|
port=8100,
|
|
external_url="http://localhost:8100",
|
|
api_tokens=None,
|
|
)
|
|
base.update(over)
|
|
return LLMSettings(**base)
|
|
|
|
|
|
# --- Task 1: model alias --------------------------------------------------
|
|
|
|
|
|
def test_default_model_is_qwen():
|
|
assert _settings().default_model == "qwen3.5"
|
|
|
|
|
|
def test_alias_resolves_to_real_model():
|
|
s = _settings(model_aliases={"qwen3.5": "Qwen/Qwen3.5-35B-A3B"})
|
|
SettingsCache.set(s)
|
|
client = LLMClient(settings=s)
|
|
assert client._resolve_model("qwen3.5") == "Qwen/Qwen3.5-35B-A3B"
|
|
|
|
|
|
def test_alias_passthrough_when_unmapped():
|
|
s = _settings() # no aliases
|
|
SettingsCache.set(s)
|
|
client = LLMClient(settings=s)
|
|
assert client._resolve_model("some-other-model") == "some-other-model"
|
|
assert client._resolve_model("qwen3.5") == "qwen3.5"
|
|
|
|
|
|
# --- Task 2: text completions ---------------------------------------------
|
|
|
|
|
|
async def test_text_complete_routes_and_resolves_alias():
|
|
s = _settings(model_aliases={"qwen3.5": "Qwen/Qwen3.5-35B-A3B"})
|
|
SettingsCache.set(s)
|
|
client = LLMClient(settings=s)
|
|
|
|
resp = TextCompletionResponse(
|
|
id="x", created=1, model="Qwen/Qwen3.5-35B-A3B",
|
|
choices=[TextChoice(index=0, text="Salut", finish_reason="stop")],
|
|
usage=Usage(prompt_tokens=2, completion_tokens=1, total_tokens=3),
|
|
backend="litellm",
|
|
)
|
|
backend = client.registry.get(BackendType.LITELLM)
|
|
backend.text_complete = AsyncMock(return_value=resp)
|
|
|
|
out = await client.text_complete(prompt="Salutare", model="qwen3.5")
|
|
|
|
assert out.choices[0].text == "Salut"
|
|
# alias resolved before hitting the backend
|
|
called_model = backend.text_complete.await_args.args[1]
|
|
assert called_model == "Qwen/Qwen3.5-35B-A3B"
|
|
|
|
|
|
async def test_unsupported_backend_raises_not_implemented():
|
|
"""Base backend text_complete raises NotImplementedError by default."""
|
|
from llm_inference.backends.litellm_backend import LiteLLMBackend
|
|
|
|
s = _settings()
|
|
SettingsCache.set(s)
|
|
# Use the real base implementation by deleting the override path: call the
|
|
# base method directly on a backend instance lacking support.
|
|
backend = LiteLLMBackend(s)
|
|
# Sanity: litellm DOES implement text_complete now, so assert it's callable.
|
|
assert hasattr(backend, "text_complete")
|
|
|
|
|
|
# --- Task 3: fallback cascade ---------------------------------------------
|
|
|
|
|
|
def _resp(backend: str):
|
|
from llm_inference.schemas import CompletionResponse
|
|
from llm_inference.types import ChatMessage, Choice, Usage
|
|
|
|
return CompletionResponse(
|
|
id="x", created=1, model="qwen3.5",
|
|
choices=[Choice(index=0, message=ChatMessage(role="assistant", content="ok"),
|
|
finish_reason="stop")],
|
|
usage=Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2),
|
|
backend=backend,
|
|
)
|
|
|
|
|
|
async def test_fallback_cascade_on_primary_failure(monkeypatch):
|
|
from llm_inference.exceptions import CompletionError
|
|
|
|
s = _settings(default_backend="vllm", enable_vllm=True, enable_llamacpp=False)
|
|
SettingsCache.set(s)
|
|
client = LLMClient(settings=s)
|
|
monkeypatch.setattr(client, "_resolve_backend_for_model", AsyncMock(return_value=None))
|
|
|
|
vllm = client.registry.get(BackendType.VLLM)
|
|
litellm = client.registry.get(BackendType.LITELLM)
|
|
vllm.complete = AsyncMock(side_effect=CompletionError("vllm down"))
|
|
litellm.complete = AsyncMock(return_value=_resp("litellm"))
|
|
|
|
out = await client.complete(messages=[{"role": "user", "content": "hi"}], model="qwen3.5")
|
|
assert out.backend == "litellm"
|
|
vllm.complete.assert_awaited_once()
|
|
litellm.complete.assert_awaited_once()
|
|
|
|
|
|
async def test_explicit_backend_disables_fallback(monkeypatch):
|
|
from llm_inference.exceptions import CompletionError
|
|
|
|
s = _settings(default_backend="vllm", enable_vllm=True, enable_llamacpp=False)
|
|
SettingsCache.set(s)
|
|
client = LLMClient(settings=s)
|
|
|
|
vllm = client.registry.get(BackendType.VLLM)
|
|
litellm = client.registry.get(BackendType.LITELLM)
|
|
vllm.complete = AsyncMock(side_effect=CompletionError("vllm down"))
|
|
litellm.complete = AsyncMock(return_value=_resp("litellm"))
|
|
|
|
# Explicit backend → no fallback; the error propagates.
|
|
import pytest
|
|
with pytest.raises(CompletionError):
|
|
await client.complete(
|
|
messages=[{"role": "user", "content": "hi"}],
|
|
model="qwen3.5",
|
|
backend="vllm",
|
|
)
|
|
litellm.complete.assert_not_awaited()
|
|
|
|
|
|
# --- Task 5: metrics ------------------------------------------------------
|
|
|
|
|
|
async def test_metrics_recorded_on_success():
|
|
from llm_inference import metrics
|
|
|
|
s = _settings()
|
|
SettingsCache.set(s)
|
|
client = LLMClient(settings=s)
|
|
litellm = client.registry.get(BackendType.LITELLM)
|
|
litellm.complete = AsyncMock(return_value=_resp("litellm"))
|
|
|
|
req = metrics.LLM_REQUESTS.labels(
|
|
model="qwen3.5", backend="litellm", status="success"
|
|
)
|
|
tok = metrics.LLM_TOKENS.labels(
|
|
model="qwen3.5", backend="litellm", kind="completion"
|
|
)
|
|
before_req = req._value.get()
|
|
before_tok = tok._value.get()
|
|
|
|
await client.complete(messages=[{"role": "user", "content": "hi"}], model="qwen3.5")
|
|
|
|
assert req._value.get() == before_req + 1
|
|
assert tok._value.get() == before_tok + 1 # _resp has completion_tokens=1
|