225 lines
7.7 KiB
Python
225 lines
7.7 KiB
Python
from __future__ import annotations
|
|
|
|
import io
|
|
import json
|
|
import urllib.error
|
|
|
|
from agents.providers import DeterministicCodingProvider, DeterministicSolProvider
|
|
from control_plane.resources.models import ModelRequest, Resource, ResourceKind
|
|
from model_router.policy import model_for_purpose
|
|
from model_router.providers import (
|
|
ProviderError,
|
|
QwenProvider,
|
|
SolProvider,
|
|
providers_from_resources,
|
|
)
|
|
from model_router.router import ModelCapability, ModelChunk, ModelRequestContract, ModelRouter
|
|
|
|
|
|
def test_model_router_persists_sanitized_request_metadata() -> None:
|
|
Resource.objects.create(
|
|
name="Qwen Test",
|
|
kind=ResourceKind.MODEL,
|
|
provider="qwen",
|
|
roles=["CODING"],
|
|
)
|
|
router = ModelRouter({"qwen": DeterministicCodingProvider()}, persist_requests=True)
|
|
|
|
router.complete(ModelRequestContract(purpose=ModelCapability.CODING, prompt="secret-looking prompt must not persist"))
|
|
|
|
request = ModelRequest.objects.get()
|
|
assert request.status == "COMPLETE"
|
|
assert request.logical_role == "CODING"
|
|
assert request.request["prompt_chars"] > 0
|
|
assert request.request["contains_raw_prompt"] is False
|
|
assert "secret-looking" not in str(request.request)
|
|
|
|
|
|
def test_model_router_health_is_non_throwing() -> None:
|
|
router = ModelRouter({"qwen": DeterministicCodingProvider()})
|
|
|
|
assert router.health() == {"qwen": "AVAILABLE"}
|
|
|
|
|
|
def test_model_router_stream_falls_back_to_complete() -> None:
|
|
router = ModelRouter({"qwen": DeterministicCodingProvider()})
|
|
|
|
chunks = list(
|
|
router.stream(
|
|
ModelRequestContract(
|
|
purpose=ModelCapability.CODING, model_hint="qwen", prompt="hello"
|
|
)
|
|
)
|
|
)
|
|
|
|
assert len(chunks) == 1
|
|
assert isinstance(chunks[0], ModelChunk)
|
|
assert chunks[0].content
|
|
|
|
|
|
def test_opencode_model_key_resources_load_as_distinct_providers() -> None:
|
|
for key in ["sol", "terra", "luna"]:
|
|
Resource.objects.create(name=key.title(), kind=ResourceKind.MODEL, provider="opencode", roles=["REASONING"], config={"model_key": key})
|
|
|
|
providers = providers_from_resources()
|
|
|
|
assert {"sol", "terra", "luna"}.issubset(providers)
|
|
assert all(isinstance(providers[key], SolProvider) for key in ["sol", "terra", "luna"])
|
|
|
|
|
|
def test_model_router_prefers_exact_model_key_resource_for_persisted_requests() -> None:
|
|
Resource.objects.create(name="Generic Opencode", kind=ResourceKind.MODEL, provider="opencode", roles=["REASONING"], config={})
|
|
Resource.objects.create(name="Terra", kind=ResourceKind.MODEL, provider="opencode", roles=["REASONING"], config={"model_key": "terra"})
|
|
router = ModelRouter({"terra": DeterministicSolProvider("{}")}, persist_requests=True)
|
|
|
|
router.complete(ModelRequestContract(purpose=ModelCapability.REASONING, model_hint="terra", prompt="{}"))
|
|
|
|
request = ModelRequest.objects.get()
|
|
assert request.model == "Terra"
|
|
|
|
|
|
def test_qwen_provider_retries_transient_http_failures(monkeypatch) -> None:
|
|
resource = Resource.objects.create(
|
|
name="Qwen Test",
|
|
kind=ResourceKind.MODEL,
|
|
provider="local_inference",
|
|
roles=["REASONING"],
|
|
config={
|
|
"endpoint_url": "http://qwen.test/v1/chat/completions",
|
|
"model": "qwen38",
|
|
"retry_attempts": 2,
|
|
"retry_backoff_seconds": 0,
|
|
},
|
|
)
|
|
calls = {"count": 0}
|
|
|
|
class Response:
|
|
status = 200
|
|
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
def read(self):
|
|
return b'{"model":"qwen38","choices":[{"message":{"content":"ok"}}],"usage":{"total_tokens":1}}'
|
|
|
|
def fake_urlopen(request, timeout):
|
|
calls["count"] += 1
|
|
if calls["count"] == 1:
|
|
raise urllib.error.HTTPError(request.full_url, 500, "Internal Server Error", {}, io.BytesIO())
|
|
return Response()
|
|
|
|
monkeypatch.setattr("model_router.providers.urllib.request.urlopen", fake_urlopen)
|
|
monkeypatch.setattr("model_router.providers.time.sleep", lambda seconds: None)
|
|
|
|
response = QwenProvider(resource).complete(ModelRequestContract(purpose=ModelCapability.REASONING, prompt="hello"))
|
|
|
|
assert response.content == "ok"
|
|
assert response.metadata["attempts"] == 2
|
|
assert calls["count"] == 2
|
|
|
|
|
|
def test_qwen_provider_reports_retry_exhaustion(monkeypatch) -> None:
|
|
resource = Resource.objects.create(
|
|
name="Qwen Test",
|
|
kind=ResourceKind.MODEL,
|
|
provider="local_inference",
|
|
roles=["REASONING"],
|
|
config={
|
|
"endpoint_url": "http://qwen.test/v1/chat/completions",
|
|
"retry_attempts": 2,
|
|
"retry_backoff_seconds": 0,
|
|
},
|
|
)
|
|
|
|
def fake_urlopen(request, timeout):
|
|
raise urllib.error.URLError("connection refused")
|
|
|
|
monkeypatch.setattr("model_router.providers.urllib.request.urlopen", fake_urlopen)
|
|
monkeypatch.setattr("model_router.providers.time.sleep", lambda seconds: None)
|
|
|
|
try:
|
|
QwenProvider(resource).complete(ModelRequestContract(purpose=ModelCapability.REASONING, prompt="hello"))
|
|
except ProviderError as exc:
|
|
message = str(exc)
|
|
else:
|
|
raise AssertionError("expected ProviderError")
|
|
|
|
assert "Qwen provider failed after retries" in message
|
|
assert "attempt 1/2" in message
|
|
assert "attempt 2/2" in message
|
|
|
|
|
|
def test_qwen_stream_forwards_no_thinking_and_persists_usage(monkeypatch) -> None:
|
|
resource = Resource.objects.create(
|
|
name="Qwen",
|
|
kind=ResourceKind.MODEL,
|
|
provider="local_inference",
|
|
roles=["STORY_PROSE"],
|
|
config={
|
|
"endpoint_url": "http://qwen.test/v1/chat/completions",
|
|
"model": "qwen38",
|
|
"extra_body": {"chat_template_kwargs": {"enable_thinking": False}},
|
|
},
|
|
)
|
|
bodies = []
|
|
|
|
class Response:
|
|
status = 200
|
|
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
def __iter__(self):
|
|
return iter(
|
|
[
|
|
b'data: {"choices":[{"delta":{"content":"draft"}}]}\n',
|
|
b'data: {"choices":[],"usage":{"prompt_tokens":12,"completion_tokens":3}}\n',
|
|
b"data: [DONE]\n",
|
|
]
|
|
)
|
|
|
|
def fake_urlopen(request, timeout):
|
|
bodies.append(json.loads(request.data.decode("utf-8")))
|
|
return Response()
|
|
|
|
monkeypatch.setattr("model_router.providers.urllib.request.urlopen", fake_urlopen)
|
|
router = ModelRouter({"qwen": QwenProvider(resource)}, persist_requests=True)
|
|
|
|
chunks = list(
|
|
router.stream(
|
|
ModelRequestContract(
|
|
purpose=ModelCapability.STORY_PROSE,
|
|
prompt="write",
|
|
model_hint="qwen",
|
|
)
|
|
)
|
|
)
|
|
|
|
request = ModelRequest.objects.get()
|
|
assert "".join(chunk.content for chunk in chunks) == "draft"
|
|
assert bodies[0]["chat_template_kwargs"]["enable_thinking"] is False
|
|
assert request.prompt_tokens == 12
|
|
assert request.completion_tokens == 3
|
|
|
|
|
|
def test_story_defaults_use_terra_luna_qwen_policy(monkeypatch) -> None:
|
|
for name in [
|
|
"ARTIFEX_STORY_PLANNING_MODEL",
|
|
"ARTIFEX_STORY_PROSE_MODEL",
|
|
"ARTIFEX_STORY_CONTINUITY_MODEL",
|
|
"ARTIFEX_STORY_REVIEW_MODEL",
|
|
"ARTIFEX_STORY_REVISION_MODEL",
|
|
]:
|
|
monkeypatch.delenv(name, raising=False)
|
|
|
|
assert model_for_purpose(ModelCapability.STORY_PLANNING) == "terra"
|
|
assert model_for_purpose(ModelCapability.STORY_PROSE) == "terra"
|
|
assert model_for_purpose(ModelCapability.STORY_CONTINUITY) == "luna"
|
|
assert model_for_purpose(ModelCapability.STORY_REVIEW) == "terra"
|
|
assert model_for_purpose(ModelCapability.STORY_REVISION) == "luna"
|