Livrare LOT 1 - Didi
This commit is contained in:
commit
5380c3fc63
990 changed files with 133308 additions and 0 deletions
175
ai_platform/modules/llm-inference/tests/test_text_completions.py
Normal file
175
ai_platform/modules/llm-inference/tests/test_text_completions.py
Normal file
|
|
@ -0,0 +1,175 @@
|
|||
"""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
|
||||
Loading…
Add table
Add a link
Reference in a new issue