Files
worldquant-alpha-system/backend/tests/test_ai_provider.py
T

77 lines
3.2 KiB
Python
Raw Normal View History

import json
import httpx
import pytest
from app.ai.provider import model_connection
from app.ai.provider import test_capabilities as check_capabilities
from app.models import AISettings
from app.security import cipher
from tests.model_protocol import model_events
@pytest.mark.parametrize("protocol", ["chat_completions", "responses"])
async def test_actual_provider_protocol_and_tool_roundtrip(app, protocol):
paths = []
def gateway(request):
paths.append(request.url.path)
assert request.headers["authorization"] == "Bearer synthetic-key"
body = json.loads(request.content)
assert body["model"] == "mock-model" and body["stream"] is True
return httpx.Response(
200, headers={"content-type": "text/event-stream"}, content="".join(model_events(body, protocol))
)
config = AISettings(
base_url="http://model.test/v1",
model="mock-model",
protocol=protocol,
api_key_encrypted=cipher(app.state.settings).encrypt(b"synthetic-key").decode(),
)
async with model_connection(config, app.state.settings, httpx.MockTransport(gateway)) as model:
result = await check_capabilities(model)
assert all(item["ok"] for item in result.values()), result
assert paths == ["/v1/" + ("chat/completions" if protocol == "chat_completions" else "responses")] * 3
@pytest.mark.parametrize("protocol", ["chat_completions", "responses"])
@pytest.mark.parametrize("failure", [401, 404, 429, "timeout", "broken", "truncated", "no-tools"])
async def test_provider_failures_are_safe(app, protocol, failure, caplog):
def gateway(request):
if isinstance(failure, int):
return httpx.Response(failure, json={"error": {"message": "synthetic-key private provider body"}})
if failure == "timeout":
raise httpx.ReadTimeout("synthetic-key private provider body")
if failure == "broken":
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
content="data: invalid-private-synthetic-key\n\n",
)
body = json.loads(request.content)
body["tools"] = []
if failure == "truncated":
frames = model_events(body, protocol)
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
content="".join(frames[:-2] if protocol == "chat_completions" else frames[:-1]),
)
return httpx.Response(
200, headers={"content-type": "text/event-stream"}, content="".join(model_events(body, protocol))
)
config = AISettings(
base_url="http://model.test/v1",
model="mock-model",
protocol=protocol,
api_key_encrypted=cipher(app.state.settings).encrypt(b"synthetic-key").decode(),
)
async with model_connection(config, app.state.settings, httpx.MockTransport(gateway)) as model:
result = await check_capabilities(model)
assert not result["tools"]["ok"], result
assert "synthetic-key" not in json.dumps(result) + caplog.text
if failure == "no-tools":
assert result["answer"]["ok"] and result["stream"]["ok"]