feat(provider): 完成阶段E协议适配、国内预设与模型路由

This commit is contained in:
2026-09-04 06:19:32 +08:00
parent 9b8b10cdb1
commit 1fe75e3fd2
57 changed files with 4208 additions and 592 deletions
+642
View File
@@ -0,0 +1,642 @@
"""Offline model-routing contracts, HTTP validation, media lifetimes and persistence.
All HTTP uses MockTransport (or the in-process API). Credentials, models and
attachments are fakes, and conftest redirects all storage to temporary paths.
"""
from __future__ import annotations
import asyncio
import hashlib
import json
from email import policy
from email.parser import BytesParser
from types import SimpleNamespace
import httpx
import pytest
from fastapi.testclient import TestClient
from app.contracts import ModelBinding, ModelRoutingConfig, ProviderConfig, ProviderType
from app.errors import ApiError
from app.providers import MockProvider
from app.providers.credentials import CredentialStoreError
from app.providers.registry import ProviderRegistry
from app.providers.routing import ModelRoutingService, PendingSpeechBackend
from app.retrieval.embedding import HashEmbeddingProvider
def run(awaitable):
return asyncio.run(awaitable)
def response(data, status=200):
# Raw JSON intentionally permits NaN/Infinity to exercise hostile API output.
return httpx.Response(status, content=json.dumps(data).encode(), headers={"content-type": "application/json"})
class FakeCredentials:
def __init__(self):
self.value = "unit-test-placeholder"
self.error = None
self.calls = []
def resolve(self, credential_id):
self.calls.append(credential_id)
if self.error:
raise self.error
return self.value if credential_id else None
class FakeEmbedding:
model_id = "fake-local-model"
dim = 3
def __init__(self):
self.calls = []
self.error = None
async def embed_documents(self, texts):
self.calls.append(list(texts))
if self.error:
raise self.error
return [[0.6, 0.8, 0.0] for _ in texts]
class FakeSpeech:
available = True
def __init__(self):
self.calls = []
self.text = "local transcript"
self.score = 0.25
self.error = None
async def transcribe(self, source, language):
self.calls.append(("transcribe", source, language))
if self.error:
raise self.error
return self.text
async def match(self, source, reference):
self.calls.append(("match", source, reference))
if self.error:
raise self.error
return self.score
@pytest.fixture(autouse=True)
def no_real_http(monkeypatch):
async def reject_async(*args, **kwargs):
pytest.fail("Real HTTP transport is forbidden in model-routing tests")
def reject_sync(*args, **kwargs):
pytest.fail("Real HTTP transport is forbidden in model-routing tests")
monkeypatch.setattr(httpx.AsyncHTTPTransport, "handle_async_request", reject_async)
monkeypatch.setattr(httpx.HTTPTransport, "handle_request", reject_sync)
@pytest.fixture
def rig():
requests = []
def unexpected(request):
pytest.fail(f"Unexpected model HTTP request: {request.url}")
state = SimpleNamespace(handler=unexpected)
async def dispatch(request):
requests.append(request)
result = state.handler(request)
return await result if hasattr(result, "__await__") else result
providers = ProviderRegistry()
config = ProviderConfig(
provider_id="test-provider", provider_type=ProviderType.openai_compatible,
name="Fake provider", base_url="https://models.invalid/v1/", credential_id="test-credential",
)
providers.register(config, MockProvider())
credentials, embedding, speech = FakeCredentials(), FakeEmbedding(), FakeSpeech()
service = ModelRoutingService(
providers, credentials, local_embedding=embedding, local_speech=speech,
transport=httpx.MockTransport(dispatch),
)
return SimpleNamespace(
service=service, providers=providers, credentials=credentials,
embedding=embedding, speech=speech, requests=requests, http=state,
)
def bind(rig, capability="embedding", **overrides):
endpoints = {
"embedding": "/embeddings", "transcription": "/audio/transcriptions",
"speaker_matching": "/audio/speaker-matches",
}
binding = ModelBinding(**{
"provider_id": "test-provider", "model": "test-model",
"endpoint": endpoints[capability], **overrides,
})
current = rig.service.configuration()
return rig.service.update(current.model_copy(update={capability: binding}))
def assert_local(rig, result, texts, reason):
assert result.source == "local"
assert result.model_id == rig.embedding.model_id
assert result.dimensions == 3
assert result.vectors == [[0.6, 0.8, 0.0] for _ in texts]
assert result.fallback_reason == reason
assert rig.embedding.calls == [texts]
@pytest.fixture
def audio(tmp_path):
source, reference = tmp_path / "audio.wav", tmp_path / "reference.wav"
source.write_bytes(b"fake-audio-content")
reference.write_bytes(b"fake-reference-content")
return source, reference
def media_call(rig, capability, audio):
if capability == "transcription":
return rig.service.transcribe(audio[0], "zh")
return rig.service.match_speakers(*audio)
def track_media_handles(rig, monkeypatch):
handles = []
original = rig.service._media_file
def tracked(path):
handle = original(path)
handles.append(handle)
return handle
monkeypatch.setattr(rig.service, "_media_file", tracked)
return handles
def test_absent_binding_uses_hash_without_network(rig):
rig.service.local_embedding = HashEmbeddingProvider()
texts = ["hello retrieval", "向量检索"]
result = run(rig.service.embed(texts))
assert result.source == "local"
assert result.model_id == "hash-v1"
assert result.dimensions == 128
assert result.vectors == run(HashEmbeddingProvider().embed_documents(texts))
assert result.fallback_reason is None
assert rig.requests == rig.credentials.calls == []
statuses = {item.capability: item.status for item in rig.service.describe().local_backends}
assert statuses == {"embedding": "placeholder", "transcription": "ready", "speaker_matching": "ready"}
def test_empty_embedding_input_does_not_call_remote(rig):
bind(rig)
result = run(rig.service.embed([]))
assert result.vectors == [] and result.source == "local"
assert rig.requests == []
def test_remote_embedding_restores_batch_order_normalizes_and_sends_auth(rig):
bind(rig, dimensions=2)
texts = [str(index) for index in range(35)]
def handler(request):
assert request.method == "POST"
assert str(request.url) == "https://models.invalid/v1/embeddings"
assert request.headers["authorization"] == "Bearer unit-test-placeholder"
payload = json.loads(request.content)
assert payload["model"] == "test-model"
assert payload["dimensions"] == 2
assert payload["encoding_format"] == "float"
return response({"data": [
{"index": index, "embedding": [float(int(text) + 1), 1.0]}
for index, text in reversed(list(enumerate(payload["input"])))
]})
rig.http.handler = handler
result = run(rig.service.embed(texts))
assert result.source == "api" and result.fallback_reason is None
assert result.dimensions == 2 and len(result.vectors) == 35
for index, vector in enumerate(result.vectors):
assert sum(value * value for value in vector) == pytest.approx(1.0)
assert vector[0] / vector[1] == pytest.approx(index + 1)
assert [json.loads(req.content)["input"] for req in rig.requests] == [texts[:32], texts[32:]]
assert rig.embedding.calls == []
def test_space_id_is_stable_and_includes_full_url_model_and_inferred_dimensions(rig):
dimensions = 2
def handler(request):
assert "dimensions" not in json.loads(request.content)
return response({"data": [{"index": 0, "embedding": [1.0] * dimensions}]})
rig.http.handler = handler
bind(rig, model=" trimmed-model ")
def check(url, model, dimension):
result = run(rig.service.embed(["hello"]))
digest = hashlib.sha256(json.dumps([url, model, dimension], separators=(",", ":")).encode()).hexdigest()
assert result.model_id == "api-" + digest
assert result.source == "api"
return result.model_id
first = check("https://models.invalid/v1/embeddings", "trimmed-model", 2)
assert first == check("https://models.invalid/v1/embeddings", "trimmed-model", 2)
config = rig.providers.get_any("test-provider").config.model_copy(update={"base_url": "https://models.invalid/v1"})
rig.providers.replace(config, MockProvider())
assert first == check("https://models.invalid/v1/embeddings", "trimmed-model", 2)
bind(rig, model="trimmed-model", endpoint="/custom/embeddings")
endpoint_id = check("https://models.invalid/v1/custom/embeddings", "trimmed-model", 2)
bind(rig, model="another-model", endpoint="/custom/embeddings")
model_id = check("https://models.invalid/v1/custom/embeddings", "another-model", 2)
dimensions = 3
dim_id = check("https://models.invalid/v1/custom/embeddings", "another-model", 3)
config = config.model_copy(update={"base_url": "https://other.invalid/v1"})
rig.providers.replace(config, MockProvider())
provider_id = check("https://other.invalid/v1/custom/embeddings", "another-model", 3)
assert len({first, endpoint_id, model_id, dim_id, provider_id}) == 5
@pytest.mark.parametrize("data", [
{"data": []},
{"data": [{"index": 0, "embedding": [1, 0]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 0, "embedding": [0, 1]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 2, "embedding": [0, 1]}]},
{"data": [{"index": False, "embedding": [1, 0]}, {"index": 1, "embedding": [0, 1]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [0, 1, 0]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [float("nan"), 1]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [float("inf"), 1]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [True, 1]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [0, 0]}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": []}]},
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": ["1", 0]}]},
{"data": [None, None]},
{"error": {"message": "in-band failure"}, "data": []},
[],
], ids=["empty", "count", "duplicate-index", "out-of-range-index", "bool-index", "dimensions", "nan", "infinity", "bool", "zero", "empty-vector", "string", "invalid-items", "in-band-error", "non-object"])
def test_invalid_remote_embeddings_fall_back_as_a_whole(rig, data):
bind(rig)
rig.http.handler = lambda request: response(data)
texts = ["first", "second"]
assert_local(rig, run(rig.service.embed(texts)), texts, "PROVIDER_INVALID_RESPONSE")
def test_explicit_embedding_dimension_mismatch_falls_back(rig):
bind(rig, dimensions=3)
rig.http.handler = lambda request: response({"data": [{"index": 0, "embedding": [1, 0]}]})
assert_local(rig, run(rig.service.embed(["text"])), ["text"], "PROVIDER_INVALID_RESPONSE")
def test_later_batch_dimension_mismatch_discards_earlier_remote_vectors(rig):
bind(rig)
def handler(request):
batch = json.loads(request.content)["input"]
dimension = 2 if len(rig.requests) == 1 else 3
return response({"data": [{"index": i, "embedding": [1] * dimension} for i in range(len(batch))]})
rig.http.handler = handler
texts = [str(i) for i in range(33)]
assert_local(rig, run(rig.service.embed(texts)), texts, "PROVIDER_INVALID_RESPONSE")
assert len(rig.requests) == 2
@pytest.mark.parametrize("failure, reason", [
(401, "PROVIDER_AUTH_FAILED"), (403, "PROVIDER_AUTH_FAILED"),
(404, "MODEL_NOT_FOUND"), (429, "PROVIDER_RATE_LIMITED"), (500, "PROVIDER_UNAVAILABLE"),
("timeout", "PROVIDER_TIMEOUT"), ("connect", "PROVIDER_UNAVAILABLE"),
("json", "PROVIDER_INVALID_RESPONSE"),
])
def test_embedding_http_failures_use_injected_local(rig, failure, reason):
bind(rig)
def handler(request):
if failure == "timeout":
raise httpx.ReadTimeout("simulated timeout", request=request)
if failure == "connect":
raise httpx.ConnectError("simulated connection failure", request=request)
if failure == "json":
return httpx.Response(200, content=b"not JSON")
return response({"error": "failed"}, failure)
rig.http.handler = handler
assert_local(rig, run(rig.service.embed(["text"])), ["text"], reason)
@pytest.mark.parametrize("failure, reason", [
("missing-key", "PROVIDER_CREDENTIAL_MISSING"),
("unreadable-key", "PROVIDER_CREDENTIAL_UNAVAILABLE"),
("disabled-provider", "PROVIDER_UNAVAILABLE"),
])
def test_unavailable_remote_configuration_falls_back_without_http(rig, failure, reason):
bind(rig)
if failure == "missing-key":
rig.credentials.value = None
elif failure == "unreadable-key":
rig.credentials.error = CredentialStoreError("fake unavailable store")
else:
config = rig.providers.get_any("test-provider").config.model_copy(update={"enabled": False})
rig.providers.replace(config, MockProvider())
assert_local(rig, run(rig.service.embed(["text"])), ["text"], reason)
assert rig.requests == []
@pytest.mark.parametrize("capability", ["transcription", "speaker_matching"])
def test_media_success_sends_expected_multipart_and_closes_files(rig, audio, monkeypatch, capability):
bind(rig, capability)
handles = track_media_handles(rig, monkeypatch)
def handler(request):
assert request.headers["authorization"] == "Bearer unit-test-placeholder"
assert str(request.url).endswith("/audio/transcriptions" if capability == "transcription" else "/audio/speaker-matches")
message = BytesParser(policy=policy.default).parsebytes(
b"Content-Type: " + request.headers["content-type"].encode() + b"\r\nMIME-Version: 1.0\r\n\r\n" + request.content,
)
parts = {part.get_param("name", header="content-disposition"): part for part in message.iter_parts()}
assert parts["model"].get_payload(decode=True) == b"test-model"
assert parts["file"].get_filename() == audio[0].name
assert parts["file"].get_payload(decode=True) == audio[0].read_bytes()
if capability == "transcription":
assert set(parts) == {"model", "language", "file"}
assert parts["language"].get_payload(decode=True) == b"zh"
return response({"text": "remote transcript"})
assert set(parts) == {"model", "file", "reference_file"}
assert parts["reference_file"].get_filename() == audio[1].name
assert parts["reference_file"].get_payload(decode=True) == audio[1].read_bytes()
return response({"score": 0.875})
rig.http.handler = handler
result = run(media_call(rig, capability, audio))
assert result.source == "api" and result.fallback_reason is None
assert result.text == "remote transcript" if capability == "transcription" else result.score == 0.875
assert len(handles) == (1 if capability == "transcription" else 2)
assert all(handle.closed for handle in handles)
assert rig.speech.calls == []
@pytest.mark.parametrize("capability, data", [
("transcription", {}), ("transcription", {"text": " "}), ("transcription", {"text": False}),
("transcription", {"error": "in-band", "text": "must not use"}),
("speaker_matching", {}), ("speaker_matching", {"score": -0.1}),
("speaker_matching", {"score": 1.1}), ("speaker_matching", {"score": True}),
("speaker_matching", {"score": float("nan")}), ("speaker_matching", {"score": "0.5"}),
("speaker_matching", {"error": "in-band", "score": 0.9}),
])
def test_invalid_remote_media_falls_back_to_injected_local(rig, audio, monkeypatch, capability, data):
bind(rig, capability)
handles = track_media_handles(rig, monkeypatch)
rig.http.handler = lambda request: response(data)
result = run(media_call(rig, capability, audio))
assert result.source == "local" and result.fallback_reason == "PROVIDER_INVALID_RESPONSE"
assert result.text == "local transcript" if capability == "transcription" else result.score == 0.25
assert rig.speech.calls == [
("transcribe", audio[0], "zh") if capability == "transcription" else ("match", *audio)
]
assert handles and all(handle.closed for handle in handles)
@pytest.mark.parametrize("capability", ["transcription", "speaker_matching"])
@pytest.mark.parametrize("configured", [False, True])
def test_pending_local_backend_has_explicit_503_and_fallback_details(rig, audio, capability, configured):
rig.service.local_speech = PendingSpeechBackend()
if configured:
bind(rig, capability)
rig.http.handler = lambda request: response({"error": "unauthorized"}, 401)
with pytest.raises(ApiError) as caught:
run(media_call(rig, capability, audio))
assert caught.value.status_code == 503
assert caught.value.code == "LOCAL_MODEL_NOT_INSTALLED"
assert caught.value.details == {"fallback_reason": "PROVIDER_AUTH_FAILED" if configured else None}
statuses = {item.capability: item.status for item in rig.service.describe().local_backends}
assert statuses["transcription"] == statuses["speaker_matching"] == "not_installed"
assert len(rig.requests) == int(configured)
@pytest.mark.parametrize("capability", ["transcription", "speaker_matching"])
def test_invalid_local_speech_returns_explicit_503(rig, audio, capability):
rig.speech.text = ""
rig.speech.score = True
with pytest.raises(ApiError) as caught:
run(media_call(rig, capability, audio))
assert (caught.value.status_code, caught.value.code) == (503, "LOCAL_MODEL_INVALID_RESPONSE")
assert caught.value.details == {"fallback_reason": None}
@pytest.mark.parametrize("capability", ["embedding", "transcription", "speaker_matching"])
@pytest.mark.parametrize("stage", ["remote", "local"])
def test_cancellation_propagates_and_upload_handles_close(rig, audio, monkeypatch, capability, stage):
bind(rig, capability)
handles = track_media_handles(rig, monkeypatch)
async def cancelled(request):
raise asyncio.CancelledError()
if stage == "remote":
rig.http.handler = cancelled
else:
rig.http.handler = lambda request: response({"error": "fallback"}, 500)
rig.embedding.error = rig.speech.error = asyncio.CancelledError()
operation = rig.service.embed(["text"]) if capability == "embedding" else media_call(rig, capability, audio)
with pytest.raises(asyncio.CancelledError):
run(operation)
assert len(handles) == {"embedding": 0, "transcription": 1, "speaker_matching": 2}[capability]
assert all(handle.closed for handle in handles)
if stage == "remote":
assert rig.embedding.calls == rig.speech.calls == []
def test_missing_reference_closes_already_open_source(rig, audio, monkeypatch):
bind(rig, "speaker_matching")
handles = track_media_handles(rig, monkeypatch)
audio[1].unlink()
with pytest.raises(ApiError) as caught:
run(rig.service.match_speakers(*audio))
assert caught.value.status_code == 404
assert len(handles) == 1 and handles[0].closed
assert rig.requests == []
def test_config_optimistic_conflict_preserves_saved_bindings(rig):
assert rig.service.configuration().version == 0
saved = bind(rig).config
assert saved.version == 1
with pytest.raises(ApiError) as caught:
rig.service.update(ModelRoutingConfig(version=0))
assert (caught.value.status_code, caught.value.code) == (409, "MODEL_ROUTING_VERSION_CONFLICT")
assert rig.service.configuration() == saved
assert rig.service.uses_provider("test-provider")
assert not rig.service.uses_provider("not-a-provider")
cleared = rig.service.update(ModelRoutingConfig(version=1)).config
assert cleared.version == 2 and cleared.embedding is None
assert not rig.service.uses_provider("test-provider")
@pytest.mark.parametrize("capability", ["embedding", "transcription", "speaker_matching"])
@pytest.mark.parametrize("provider_id, code", [
("missing", "PROVIDER_NOT_FOUND"), ("unsupported", "MODEL_ROUTING_PROTOCOL_UNSUPPORTED"),
])
def test_config_references_require_existing_supported_providers(rig, capability, provider_id, code):
rig.providers.register(
ProviderConfig(provider_id="unsupported", provider_type=ProviderType.ollama, name="unsupported"), MockProvider(),
)
with pytest.raises(ApiError) as caught:
bind(rig, capability, provider_id=provider_id)
assert (caught.value.status_code, caught.value.code) == (422, code)
assert rig.service.configuration() == ModelRoutingConfig()
assert rig.requests == []
@pytest.fixture
def api(monkeypatch, no_real_http, _isolate_data_dir):
# Import the production container only after temporary storage is configured.
from app import container as container_module, routes
from app.main import app
containers = []
def restart():
container = container_module.build_container()
container.model_routing.credentials = FakeCredentials()
def unexpected(request):
pytest.fail(f"Unexpected API-side provider HTTP: {request.url}")
container.model_routing.transport = httpx.MockTransport(unexpected)
monkeypatch.setattr(container_module, "container", container)
monkeypatch.setattr(routes, "container", container)
containers.append(container)
return container
container = restart()
client = TestClient(app)
yield SimpleNamespace(client=client, container=container, restart=restart)
client.close()
for container in containers:
container.plugins.shutdown()
container.mcp_servers.shutdown()
def create_api_provider(api):
result = api.client.post("/api/providers", json={
"provider_type": "openai_compatible", "name": "Persisted fake",
"base_url": "https://persist.invalid/v1", "default_model": "fake-model",
})
assert result.status_code == 200, result.text
return result.json()
def test_api_config_conflict_reference_delete_and_restart_persistence(api):
provider = create_api_provider(api)
provider_id = provider["provider_id"]
assert api.client.get("/api/model-routing").json()["config"]["version"] == 0
config = {"version": 0, "embedding": {"provider_id": provider_id, "model": "embed-model", "endpoint": "/embeddings"}}
saved = api.client.put("/api/model-routing", json=config)
assert saved.status_code == 200
assert saved.json()["config"]["version"] == 1
conflict = api.client.put("/api/model-routing", json=config)
assert conflict.status_code == 409
assert conflict.json()["error"]["code"] == "MODEL_ROUTING_VERSION_CONFLICT"
blocked = api.client.delete(f"/api/providers/{provider_id}")
assert blocked.status_code == 409 and blocked.json()["error"]["code"] == "PROVIDER_IN_USE"
restarted = api.restart()
assert restarted.providers.get_any(provider_id).config.model_dump(mode="json") == provider
assert api.client.get("/api/model-routing").json()["config"] == saved.json()["config"]
assert {item["provider_id"] for item in api.client.get("/api/providers").json()["items"]} == {"mock", provider_id}
cleared = api.client.put("/api/model-routing", json={"version": 1})
assert cleared.status_code == 200
assert api.client.delete(f"/api/providers/{provider_id}").status_code == 200
api.restart()
assert api.client.get(f"/api/providers/{provider_id}").status_code == 404
assert api.client.get("/api/model-routing").json()["config"]["version"] == 2
def test_api_provider_type_patch_rebuilds_adapter_and_persists(api):
from app.providers.anthropic_messages import AnthropicMessagesProvider
provider = create_api_provider(api)
provider_id = provider["provider_id"]
changed = api.client.patch(f"/api/providers/{provider_id}", json={
"provider_type": "anthropic_messages", "base_url": "https://anthropic.invalid/v1",
})
assert changed.status_code == 200, changed.text
assert changed.json()["provider_type"] == "anthropic_messages"
assert changed.json()["name"] == provider["name"]
assert isinstance(api.container.providers.get_any(provider_id).adapter, AnthropicMessagesProvider)
restarted = api.restart()
assert isinstance(restarted.providers.get_any(provider_id).adapter, AnthropicMessagesProvider)
assert api.client.get(f"/api/providers/{provider_id}").json() == changed.json()
for invalid_type in (None, "mock", "nonexistent-type"):
rejected = api.client.patch(f"/api/providers/{provider_id}", json={"provider_type": invalid_type})
assert rejected.status_code == 422
assert api.client.get(f"/api/providers/{provider_id}").json() == changed.json()
@pytest.mark.parametrize("endpoint", ["https://elsewhere.invalid/embed", "//elsewhere.invalid/embed", "relative", "/../embed", "/embed?key=test"])
def test_api_config_rejects_non_provider_endpoint_paths(api, endpoint):
provider = create_api_provider(api)
result = api.client.put("/api/model-routing", json={
"embedding": {"provider_id": provider["provider_id"], "model": "embed", "endpoint": endpoint},
})
assert result.status_code == 422
assert api.client.get("/api/model-routing").json()["config"]["version"] == 0
def test_api_embedding_reports_remote_and_fallback_sources(api):
provider = create_api_provider(api)
assert api.client.put("/api/model-routing", json={
"embedding": {"provider_id": provider["provider_id"], "model": "embed", "endpoint": "/embeddings"},
}).status_code == 200
api.container.model_routing.transport = httpx.MockTransport(
lambda request: response({"data": [{"index": 0, "embedding": [3, 4]}]}),
)
result = api.client.post("/api/models/embeddings", json={"texts": ["hello"]})
assert result.status_code == 200
assert result.json()["source"] == "api" and result.json()["vectors"][0] == pytest.approx([0.6, 0.8])
api.container.model_routing.transport = httpx.MockTransport(lambda request: response({"error": "denied"}, 401))
result = api.client.post("/api/models/embeddings", json={"texts": ["hello"]})
assert result.status_code == 200
assert result.json()["source"] == "local" and result.json()["model_id"] == "hash-v1"
assert result.json()["fallback_reason"] == "PROVIDER_AUTH_FAILED"
assert api.client.post("/api/models/embeddings", json={"texts": []}).status_code == 422
def test_api_speech_failure_reports_reason_in_503_and_transcription_job(api):
from app.services.attachment_service import attachment_path
source, reference = attachment_path("audio.wav"), attachment_path("reference.wav")
source.parent.mkdir(parents=True, exist_ok=True)
source.write_bytes(b"test audio")
reference.write_bytes(b"test reference")
provider = create_api_provider(api)
assert api.client.put("/api/model-routing", json={
"transcription": {"provider_id": provider["provider_id"], "model": "asr", "endpoint": "/audio/transcriptions"},
"speaker_matching": {"provider_id": provider["provider_id"], "model": "voice", "endpoint": "/audio/speaker-matches"},
}).status_code == 200
api.container.model_routing.transport = httpx.MockTransport(lambda request: response({"error": "offline"}, 500))
match = api.client.post("/api/media/speaker-matches", json={"attachment_id": source.name, "reference_attachment_id": reference.name})
assert match.status_code == 503
assert match.json()["error"]["code"] == "LOCAL_MODEL_NOT_INSTALLED"
assert match.json()["error"]["details"] == {"fallback_reason": "PROVIDER_UNAVAILABLE"}
transcript = api.client.post("/api/media/transcriptions", json={"attachment_id": source.name, "language": "zh"})
assert transcript.status_code == 202
job = transcript.json()
assert job["status"] == "failed" and job["error_code"] == "LOCAL_MODEL_NOT_INSTALLED"
assert job["fallback_reason"] == "PROVIDER_UNAVAILABLE"
assert api.client.get(f"/api/media/transcriptions/{job['job_id']}").json() == job
@pytest.mark.parametrize("capability", ["embedding", "speaker_matching"])
def test_out_of_float_range_json_number_is_invalid_remote_and_falls_back(rig, audio, capability):
"""JSON integers may be finite but too large to convert to a Python float."""
bind(rig, capability)
data = {"data": [{"index": 0, "embedding": [10 ** 400, 1]}]} if capability == "embedding" else {"score": 10 ** 400}
rig.http.handler = lambda request: response(data)
if capability == "embedding":
assert_local(rig, run(rig.service.embed(["text"])), ["text"], "PROVIDER_INVALID_RESPONSE")
else:
result = run(media_call(rig, capability, audio))
assert result.source == "local" and result.score == rig.speech.score
assert result.fallback_reason == "PROVIDER_INVALID_RESPONSE"
+1 -1
View File
@@ -84,7 +84,7 @@ def test_openai_compatible_maps_tool_call_and_credentials() -> None:
)
)
assert captured["tools"][0]["function"]["name"] == "math.add"
assert captured["tools"][0]["function"]["name"].startswith("tool_")
assert turn.tool_calls[0].name == "math.add"
assert turn.tool_calls[0].arguments == {"left": 1, "right": 2}
assert turn.input_tokens == 8
+585
View File
@@ -0,0 +1,585 @@
"""Wire-level provider tests: no credentials, SDKs, clocks, or network services."""
import asyncio
import json
import httpx
import pytest
from app.contracts import Message, MessageRole, ModelCapability, ModelEventType as E, ModelRequest, ToolCall, ToolDefinition
from app.providers.anthropic_messages import AnthropicMessagesProvider
from app.providers.base import ProviderError
from app.providers.ollama import OllamaProvider
from app.providers.openai_compatible import OpenAICompatibleProvider
from app.providers.openai_responses import OpenAIResponsesProvider
NATIVE = ["responses", "anthropic"]
PROTOCOLS = [*NATIVE, "compatible", "ollama"]
SECRET = "test-only-sensitive-upstream-body"
class Credentials:
def resolve(self, credential_id):
return SECRET if credential_id else None
class Bytes(httpx.AsyncByteStream):
def __init__(self, body: bytes, *, fragment: int = 17):
self.body = body
self.fragment = fragment
self.closed = False
async def __aiter__(self):
for offset in range(0, len(self.body), self.fragment):
yield self.body[offset:offset + self.fragment]
async def aclose(self):
self.closed = True
class GatedBytes(Bytes):
def __init__(self, body):
super().__init__(body)
self.waiting = asyncio.Event()
self.release = asyncio.Event()
async def __aiter__(self):
yield self.body
self.waiting.set()
await self.release.wait()
def provider(protocol, handler, *, credential_id="test"):
transport = httpx.MockTransport(handler)
if protocol == "ollama":
return OllamaProvider("https://provider.test", transport=transport)
cls = {"responses": OpenAIResponsesProvider, "anthropic": AnthropicMessagesProvider,
"compatible": OpenAICompatibleProvider}[protocol]
return cls("https://provider.test/v1/", credential_id, Credentials(), transport=transport)
def request(*, history=False):
messages = [Message(role=MessageRole.user, content="查笔记")]
if history:
messages += [
Message(role=MessageRole.system, content="Additional rules"),
Message(role=MessageRole.assistant, content="Checking", tool_calls=[
ToolCall(tool_call_id="old_1", name="lookup", arguments={"query": "a"}),
ToolCall(tool_call_id="old_2", name="lookup", arguments={"query": "b"}),
]),
Message(role=MessageRole.tool, tool_call_id="old_1", content='{"found":1}'),
Message(role=MessageRole.tool, tool_call_id="old_2", content='{"found":2}'),
]
return ModelRequest(
provider_id="test", model="model", system="System rules", messages=messages,
tools=[ToolDefinition(name="lookup", description="Find notes", parameters={"type": "object"})],
max_tokens=512, temperature=0,
)
async def collect(iterator):
return [event async for event in iterator]
def sse(*events):
return "".join(
f"event: {event.get('type', 'message')}\r\ndata: {json.dumps(event, ensure_ascii=False)}\r\n\r\n"
for event in events
).encode()
def wire(protocol, *events):
if protocol == "ollama":
return ("\n".join(json.dumps(event, ensure_ascii=False) for event in events) + "\n").encode()
return sse(*events)
def start(protocol):
if protocol == "responses":
return [{"type": "response.output_text.delta", "delta": "你好"}]
if protocol == "anthropic":
return [{"type": "message_start", "message": {"usage": {"input_tokens": 7, "output_tokens": 0}}},
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "你好"}}]
if protocol == "compatible":
return [{"choices": [{"delta": {"content": "你好"}}]}]
return [{"message": {"content": "你好"}, "done": False}]
def terminal(protocol):
if protocol == "responses":
return [{"type": "response.completed", "response": {"status": "completed", "usage": {"input_tokens": 7, "output_tokens": 2}}}]
if protocol == "anthropic":
return [{"type": "content_block_stop", "index": 0},
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 2}},
{"type": "message_stop"}]
if protocol == "compatible":
return [{"choices": [{"delta": {}, "finish_reason": "stop"}], "usage": {"prompt_tokens": 7, "completion_tokens": 2}}]
return [{"message": {}, "done": True, "prompt_eval_count": 7, "eval_count": 2}]
def assert_events(events):
assert events[-1].event == E.done
assert events[-1].data["status"] == ("failed" if any(event.event == E.error for event in events) else "completed")
assert sum(event.event == E.done for event in events) == 1
assert [event.sequence for event in events] == list(range(len(events)))
assert all(event.timestamp.tzinfo is not None for event in events)
def assert_error(events, code):
assert_events(events)
assert events[-2].event == E.error
assert events[-2].data["code"] == code
assert SECRET not in str(events[-2].data)
@pytest.mark.parametrize("protocol", NATIVE)
def test_native_completion_and_history(protocol):
captured = {}
def handler(req):
captured.update(json.loads(req.content))
assert req.url.path == ("/v1/responses" if protocol == "responses" else "/v1/messages")
if protocol == "responses":
assert req.headers["authorization"] == f"Bearer {SECRET}"
body = {"status": "completed", "output": [
{"type": "reasoning", "summary": [{"type": "summary_text", "text": "thinking"}]},
{"type": "message", "content": [{"type": "output_text", "text": "完成"}]},
{"type": "function_call", "call_id": "next", "name": "lookup", "arguments": '{"query":"c"}'},
], "usage": {"input_tokens": 10, "output_tokens": 3}}
else:
assert "authorization" not in req.headers
assert req.headers["x-api-key"] == SECRET
assert req.headers["anthropic-version"] == "2023-06-01"
body = {"type": "message", "content": [
{"type": "thinking", "thinking": "thinking", "signature": "sig"},
{"type": "text", "text": "完成"},
{"type": "tool_use", "id": "next", "name": "lookup", "input": {"query": "c"}},
], "usage": {"input_tokens": 5, "cache_creation_input_tokens": 2, "cache_read_input_tokens": 3, "output_tokens": 3}}
return httpx.Response(200, json=body)
turn = asyncio.run(provider(protocol, handler).complete(request(history=True)))
assert turn.text == "完成"
assert (turn.input_tokens, turn.output_tokens) == (10, 3)
assert turn.tool_calls[0].tool_call_id == "next"
assert turn.tool_calls[0].arguments == {"query": "c"}
assert captured["stream"] is False
assert captured["temperature"] == 0
if protocol == "responses":
assert captured["instructions"] == "System rules"
assert captured["max_output_tokens"] == 512
assert captured["tools"][0]["parameters"] == {"type": "object"}
calls = [item for item in captured["input"] if item.get("type") == "function_call"]
outputs = [item for item in captured["input"] if item.get("type") == "function_call_output"]
assert [call["call_id"] for call in calls] == ["old_1", "old_2"]
assert json.loads(calls[1]["arguments"]) == {"query": "b"}
assert outputs == [{"type": "function_call_output", "call_id": "old_1", "output": '{"found":1}'},
{"type": "function_call_output", "call_id": "old_2", "output": '{"found":2}'}]
assert {"role": "system", "content": "Additional rules"} in captured["input"]
else:
assert captured["system"] == "System rules\n\nAdditional rules"
assert captured["max_tokens"] == 512
assert captured["tools"][0]["input_schema"] == {"type": "object"}
assert captured["messages"][1]["content"][2] == {
"type": "tool_use", "id": "old_2", "name": "lookup", "input": {"query": "b"},
}
assert captured["messages"][-1] == {"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "old_1", "content": '{"found":1}'},
{"type": "tool_result", "tool_use_id": "old_2", "content": '{"found":2}'},
]}
def responses_tool_events():
events = [
{"type": "response.created", "response": {"usage": {"input_tokens": 10, "output_tokens": 0}}},
{"type": "response.reasoning_summary_text.delta", "delta": "计划"},
{"type": "response.output_text.delta", "delta": ""},
{"type": "response.output_text.delta", "delta": ""},
]
for index in (2, 3):
events.append({"type": "response.output_item.added", "output_index": index, "item": {
"id": f"item_{index}", "type": "function_call", "call_id": f"call_{index}", "name": "lookup", "arguments": "",
}})
for index, fragment in [(2, '{"query":'), (3, '{}'), (2, '"笔记"}')]:
events.append({"type": "response.function_call_arguments.delta", "output_index": index,
"item_id": f"item_{index}", "delta": fragment})
for index, arguments in [(3, '{}'), (2, '{"query":"笔记"}')]:
events += [
{"type": "response.function_call_arguments.done", "output_index": index, "item_id": f"item_{index}", "arguments": arguments},
{"type": "response.output_item.done", "output_index": index, "item": {
"id": f"item_{index}", "type": "function_call", "call_id": f"call_{index}", "name": "lookup", "arguments": arguments,
}},
]
events += [{"type": "future.event"}, {"type": "response.completed", "response": {
"status": "completed", "usage": {"input_tokens": 10, "output_tokens": 9},
}}]
return events
def anthropic_tool_events():
events = [
{"type": "message_start", "message": {"usage": {
"input_tokens": 5, "cache_read_input_tokens": 3, "cache_creation_input_tokens": 2, "output_tokens": 1,
}}},
{"type": "ping"},
{"type": "content_block_start", "index": 0, "content_block": {"type": "thinking", "thinking": ""}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "thinking_delta", "thinking": "计划"}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "signature_delta", "signature": "sig"}},
{"type": "content_block_stop", "index": 0},
{"type": "content_block_start", "index": 1, "content_block": {"type": "text", "text": ""}},
{"type": "content_block_delta", "index": 1, "delta": {"type": "text_delta", "text": ""}},
{"type": "content_block_stop", "index": 1},
]
for index, fragments in [(2, ['{"query":', '"笔记"}']), (3, [])]:
events.append({"type": "content_block_start", "index": index, "content_block": {
"type": "tool_use", "id": f"call_{index}", "name": "lookup", "input": {},
}})
for fragment in fragments:
events.append({"type": "content_block_delta", "index": index,
"delta": {"type": "input_json_delta", "partial_json": fragment}})
events.append({"type": "content_block_stop", "index": index})
events += [
{"type": "message_delta", "delta": {"stop_reason": "tool_use"}, "usage": {"output_tokens": 4}},
{"type": "future.event"},
{"type": "message_delta", "delta": {}, "usage": {"output_tokens": 9}},
{"type": "message_stop"},
]
return events
@pytest.mark.parametrize("protocol", NATIVE)
def test_native_stream_tools_reasoning_usage_and_fragmented_utf8(protocol):
frames = responses_tool_events() if protocol == "responses" else anthropic_tool_events()
body = Bytes(b": comment\r\n\r\n" + sse(*frames) + b"data: malformed after completion\n\n", fragment=1)
def handler(req):
payload = json.loads(req.content)
assert payload["stream"] is True
assert payload["tools"]
assert (payload.get("input") or payload.get("messages"))
return httpx.Response(200, stream=body)
events = asyncio.run(collect(provider(protocol, handler).stream(request(history=True))))
assert_events(events)
assert not any(event.event == E.error for event in events)
assert [event.data["text"] for event in events if event.event == E.text_delta] == ["", ""]
assert [event.data["text"] for event in events if event.event == E.thinking_delta] == ["计划"]
assert [event.data["tool_call_id"] for event in events if event.event == E.tool_call_start] == ["call_2", "call_3"]
assert sorted(event.data["tool_call_id"] for event in events if event.event == E.tool_call_end) == ["call_2", "call_3"]
for call_id, expected in [("call_2", {"query": "笔记"}), ("call_3", {})]:
arguments = "".join(event.data["arguments_delta"] for event in events
if event.event == E.tool_call_delta and event.data["tool_call_id"] == call_id)
assert json.loads(arguments) == expected
usages = [event.data for event in events if event.event == E.usage]
assert usages[-1] == {"input_tokens": 10, "output_tokens": 9, "total_tokens": 19}
assert all(usage["input_tokens"] == 10 for usage in usages)
if protocol == "anthropic":
assert [usage["output_tokens"] for usage in usages] == [1, 4, 9]
assert body.closed
@pytest.mark.parametrize("protocol", PROTOCOLS)
def test_stream_terminal_usage_and_closure(protocol):
body = Bytes(wire(protocol, *start(protocol), *terminal(protocol)))
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
assert_events(events)
assert not any(event.event == E.error for event in events)
assert [event.data["text"] for event in events if event.event == E.text_delta] == ["你好"]
assert [event.data for event in events if event.event == E.usage][-1] == {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}
assert body.closed
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("empty", [False, True])
def test_truncated_stream(protocol, empty):
body = Bytes(b"" if empty else wire(protocol, *start(protocol)))
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
assert_error(events, "PROVIDER_STREAM_TRUNCATED")
assert body.closed
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("bad", [b"not-json", b"[]", b"null", b'{"usage":'])
def test_malformed_stream_is_sanitized(protocol, bad):
suffix = bad + b"\n" if protocol == "ollama" else b"data: " + bad + b"\n\n"
body = Bytes(wire(protocol, *start(protocol)) + suffix)
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
assert_error(events, "PROVIDER_INVALID_RESPONSE")
assert body.closed
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("error_type,code", [("rate_limit_error", "PROVIDER_RATE_LIMITED"),
("authentication_error", "PROVIDER_AUTH_FAILED"),
("overloaded_error", "PROVIDER_UNAVAILABLE")])
def test_in_band_error_after_partial_output(protocol, error_type, code):
body = Bytes(wire(protocol, *start(protocol), {"type": "error", "error": {"type": error_type, "message": SECRET}}))
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
assert any(event.event == E.text_delta for event in events)
assert_error(events, code)
assert body.closed
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("status,code", [(400, "PROVIDER_INVALID_REQUEST"), (401, "PROVIDER_AUTH_FAILED"),
(403, "PROVIDER_AUTH_FAILED"), (404, "MODEL_NOT_FOUND"),
(429, "PROVIDER_RATE_LIMITED"), (500, "PROVIDER_UNAVAILABLE")])
def test_http_errors_completion_and_stream(protocol, status, code):
adapter = provider(protocol, lambda _: httpx.Response(status, text=SECRET))
with pytest.raises(ProviderError) as exc:
asyncio.run(adapter.complete(request()))
assert exc.value.code == code
assert SECRET not in str(exc.value)
assert_error(asyncio.run(collect(adapter.stream(request()))), code)
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("body,code", [(b"broken", "PROVIDER_INVALID_RESPONSE"),
(b"[]", "PROVIDER_INVALID_RESPONSE"),
(b"{}", "PROVIDER_INVALID_RESPONSE"),
(json.dumps({"error": {"code": "invalid_api_key", "message": SECRET}}).encode(), "PROVIDER_AUTH_FAILED")])
def test_bad_completion(protocol, body, code):
adapter = provider(protocol, lambda _: httpx.Response(200, content=body))
with pytest.raises(ProviderError) as exc:
asyncio.run(adapter.complete(request()))
assert exc.value.code == code
assert SECRET not in str(exc.value)
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("error,code", [(httpx.ReadTimeout, "PROVIDER_TIMEOUT"),
(httpx.ConnectError, "PROVIDER_UNAVAILABLE")])
def test_transport_error_mapping(protocol, error, code):
def handler(req):
raise error(SECRET, request=req)
adapter = provider(protocol, handler)
with pytest.raises(ProviderError) as exc:
asyncio.run(adapter.complete(request()))
assert exc.value.code == code
assert SECRET not in str(exc.value)
assert_error(asyncio.run(collect(adapter.stream(request()))), code)
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("cancel", [True, False])
def test_incremental_delivery_cancellation_and_explicit_close(protocol, cancel):
async def scenario():
body = GatedBytes(wire(protocol, *start(protocol)))
adapter = provider(protocol, lambda _: httpx.Response(200, stream=body))
iterator = adapter.stream(request())
seen = []
while True:
event = await asyncio.wait_for(anext(iterator), timeout=1)
seen.append(event)
if event.event == E.text_delta:
break
# The first token arrives while the response is still open and blocked.
assert seen[-1].data["text"] == "你好"
assert not body.closed
if cancel:
pending = asyncio.create_task(anext(iterator))
await asyncio.wait_for(body.waiting.wait(), timeout=1)
pending.cancel()
with pytest.raises(asyncio.CancelledError):
await pending
else:
await iterator.aclose()
assert body.closed
assert not any(event.event in {E.error, E.done} for event in seen)
asyncio.run(scenario())
@pytest.mark.parametrize("protocol", NATIVE)
def test_cancellation_before_response_headers(protocol):
async def scenario():
entered = asyncio.Event()
closed = asyncio.Event()
async def handler(req):
entered.set()
try:
await asyncio.Event().wait()
finally:
closed.set()
adapter = provider(protocol, handler)
pending = asyncio.create_task(adapter.complete(request()))
await asyncio.wait_for(entered.wait(), timeout=1)
pending.cancel()
with pytest.raises(asyncio.CancelledError):
await pending
assert closed.is_set()
asyncio.run(scenario())
@pytest.mark.parametrize("protocol", NATIVE)
def test_native_discovery_does_not_claim_non_chat_capabilities(protocol):
def handler(req):
assert req.url.path == "/v1/models"
return httpx.Response(200, json={"data": [{"id": name} for name in ["chat-model", "text-embedding-3-small", "whisper-1", "gpt-audio"]]})
models = asyncio.run(provider(protocol, handler).list_models())
assert ModelCapability.chat in models[0].capabilities
assert models[1].capabilities == [ModelCapability.embedding]
assert all(ModelCapability.chat not in model.capabilities for model in models[1:])
@pytest.mark.parametrize("protocol", NATIVE)
def test_native_structured_format_mapping(protocol):
adapter = provider(protocol, lambda _: pytest.fail("No network expected"))
req = request()
req.response_format = {"type": "json_schema", "json_schema": {
"name": "answer", "strict": True, "schema": {"type": "object", "properties": {}},
}}
payload = adapter._payload(req, stream=False)
format_ = payload["text"]["format"] if protocol == "responses" else payload["output_config"]["format"]
assert format_["type"] == "json_schema"
assert format_["schema"] == {"type": "object", "properties": {}}
if protocol == "responses":
assert format_["name"] == "answer"
assert format_["strict"] is True
@pytest.mark.parametrize("protocol", NATIVE)
def test_invalid_tool_arguments_and_unclosed_tool(protocol):
frames = responses_tool_events() if protocol == "responses" else anthropic_tool_events()
# A syntactically valid terminal cannot rescue an unfinished tool block.
index = next(i for i, frame in enumerate(frames)
if frame["type"] in {"response.function_call_arguments.delta", "content_block_delta"}
and (frame.get("output_index") == 2 or frame.get("index") == 2))
partial = frames[:index + 1]
final = frames[-1]
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, content=sse(*partial, final))).stream(request())))
assert_error(events, "PROVIDER_STREAM_TRUNCATED")
assert not any(event.event == E.tool_call_end for event in events)
for frame in frames:
if frame["type"] == "response.function_call_arguments.done":
frame["arguments"] = "[]"
break
if frame["type"] == "content_block_delta" and frame.get("index") == 2:
frame["delta"]["partial_json"] = "malformed"
break
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, content=sse(*frames))).stream(request())))
assert_error(events, "PROVIDER_INVALID_RESPONSE")
@pytest.mark.parametrize("kind,code", [("response.failed", "PROVIDER_UNAVAILABLE"),
("response.incomplete", "PROVIDER_INCOMPLETE_RESPONSE")])
def test_responses_failed_and_incomplete(kind, code):
frame = {"type": kind, "response": {"status": kind.split(".")[1], "incomplete_details": {"reason": SECRET}}}
events = asyncio.run(collect(provider("responses", lambda _: httpx.Response(200, content=sse(*start("responses"), frame))).stream(request())))
assert_error(events, code)
def test_sse_multiline_data_and_event_name_without_json_type():
body = (b': keepalive\n\nevent: response.output_text.delta\ndata: {\ndata: "delta": "hello"\ndata: }\n\n'
+ sse({"type": "response.completed", "response": {"status": "completed"}}))
events = asyncio.run(collect(provider("responses", lambda _: httpx.Response(200, content=body)).stream(request())))
assert_events(events)
assert [event.data["text"] for event in events if event.event == E.text_delta] == ["hello"]
assert not any(event.event == E.error for event in events)
def test_ollama_history_options_and_in_band_string_error():
captured = {}
def handler(req):
captured.update(json.loads(req.content))
return httpx.Response(200, json={"error": SECRET})
with pytest.raises(ProviderError) as exc:
asyncio.run(provider("ollama", handler).complete(request(history=True)))
assert exc.value.code == "PROVIDER_UNAVAILABLE"
assert SECRET not in str(exc.value)
assert captured["messages"][-1]["tool_name"] == "lookup"
assert captured["options"] == {"temperature": 0.0, "num_predict": 512}
@pytest.mark.parametrize("protocol", PROTOCOLS)
@pytest.mark.parametrize("streaming", [False, True])
def test_namespaced_tools_roundtrip_without_changing_internal_request(protocol, streaming):
import re
model_request = request(history=True)
original_name = "mcp.my-server.search.notes"
model_request.tools[0].name = original_name
for message in model_request.messages:
for call in message.tool_calls:
call.name = original_name
before = model_request.model_dump()
def handler(req):
payload = json.loads(req.content)
definition = payload["tools"][0]
name = (definition.get("function") or definition)["name"]
assert name != original_name and re.fullmatch(r"[a-zA-Z0-9_-]{1,64}", name)
assert original_name not in req.content.decode()
if protocol == "responses":
item = {"type": "function_call", "id": "item1", "call_id": "call1", "name": name, "arguments": "{}"}
body = {"status": "completed", "output": [item]}
events = [
{"type": "response.output_item.done", "output_index": 0, "item": item},
{"type": "response.completed", "response": {"status": "completed"}},
]
elif protocol == "anthropic":
item = {"type": "tool_use", "id": "call1", "name": name, "input": {}}
body = {"content": [item]}
events = [
{"type": "message_start", "message": {}},
{"type": "content_block_start", "index": 0, "content_block": item},
{"type": "content_block_stop", "index": 0},
{"type": "message_stop"},
]
elif protocol == "compatible":
item = {"id": "call1", "function": {"name": name, "arguments": "{}"}}
body = {"choices": [{"message": {"tool_calls": [item]}}]}
events = [{"choices": [{"delta": {"tool_calls": [{"index": 0, **item}]}, "finish_reason": "tool_calls"}]}]
else:
item = {"function": {"name": name, "arguments": {}}}
body = {"message": {"tool_calls": [item]}, "done": True}
events = [body]
return httpx.Response(200, content=wire(protocol, *events)) if streaming else httpx.Response(200, json=body)
adapter = provider(protocol, handler)
if streaming:
events = asyncio.run(collect(adapter.stream(model_request)))
assert_events(events)
assert [event.data["name"] for event in events if event.event == E.tool_call_start] == [original_name]
else:
assert asyncio.run(adapter.complete(model_request)).tool_calls[0].name == original_name
assert model_request.model_dump() == before
def test_chat_route_closes_upstream_and_sanitizes_unexpected_errors(monkeypatch):
from types import SimpleNamespace
from datetime import datetime, timezone
from app import routes
from app.contracts import ChatRequest, ModelEvent
closed = []
class Adapter:
async def stream(self, request):
try:
yield ModelEvent(event=E.text_delta, sequence=0, data={"text": "first"}, timestamp=datetime.now(timezone.utc))
raise RuntimeError(SECRET)
finally:
closed.append(True)
monkeypatch.setattr(routes, "provider_or_404", lambda _: SimpleNamespace(adapter=Adapter()))
async def scenario():
response = await routes.chat(ChatRequest(provider_id="test", model="test", messages=[]))
iterator = response.body_iterator
await anext(iterator)
await iterator.aclose()
assert len(closed) == 1
response = await routes.chat(ChatRequest(provider_id="test", model="test", messages=[]))
items = [json.loads(chunk.split("data: ")[1].strip()) async for chunk in response.body_iterator]
assert [item["sequence"] for item in items] == [0, 1, 2]
assert items[-1]["data"]["status"] == "failed"
assert SECRET not in str(items)
assert len(closed) == 2
asyncio.run(scenario())
+328
View File
@@ -0,0 +1,328 @@
"""Phase E route integration: deterministic runtimes, isolated DBs, no network."""
from __future__ import annotations
import asyncio
import json
from dataclasses import dataclass, field
from types import SimpleNamespace
import pytest
from app import repository
from app.config import get_settings
from app.contracts import IndexRebuildRequest, SearchMode, SearchRequest
from app.database.db import connect, transaction
from app.retrieval import routed_vectors
from app.retrieval.embedding import HashEmbeddingProvider
from app.retrieval.engine import RetrievalEngine, engine
from app.retrieval.reranker import LexicalReranker
from app.retrieval.vectorstore import SqliteVecStore, VectorHit
from app.services import index_service, note_service
@dataclass
class FakeRuntime:
model_id: str = "space-a"
dimensions: int = 3 # Deliberately differs from sqlite-vec's fixed 128.
source: str = "api"
error: BaseException | None = None
calls: list[list[str]] = field(default_factory=list)
result_override: object | None = None
async def embed(self, texts):
self.calls.append(list(texts))
if self.error is not None:
raise self.error
if self.result_override is not None:
return self.result_override
vectors = []
for text in texts:
# The API associates "apple" with banana; hash retrieval picks apple.
first = text == "apple orchard"
if self.model_id == "space-b":
first = not first
vectors.append(([1.0, 0.0] if first else [0.0, 1.0]) + [0.0] * (self.dimensions - 2))
return SimpleNamespace(
vectors=vectors, source=self.source, model_id=self.model_id,
dimensions=self.dimensions, fallback_reason=None,
)
@pytest.fixture
def runtime(monkeypatch):
runtime = FakeRuntime()
monkeypatch.setattr(routed_vectors, "get_model_routing", lambda: runtime)
return runtime
async def seed():
apple = await note_service.create_note(
title="Apple", markdown="apple orchard", folder=None, tags=[],
)
banana = await note_service.create_note(
title="Banana", markdown="banana grove", folder=None, tags=[],
)
return apple, banana
def local_engine():
return RetrievalEngine(HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore())
def request(mode=SearchMode.vector):
return SearchRequest(query="apple", mode=mode, limit=10)
def rows(sql, parameters=()):
conn = connect()
try:
return conn.execute(sql, parameters).fetchall()
finally:
conn.close()
def test_api_index_and_query_use_matching_space_and_keep_local_metadata(runtime):
async def scenario():
apple, banana = await seed()
result = await engine.search(request())
assert result.items[0].note_id == banana.note_id
baseline = await local_engine().search(request())
assert baseline.items[0].note_id == apple.note_id
assert rows("SELECT DISTINCT space_id, dimensions FROM routed_block_vectors")[0][:] == ("space-a", 3)
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == len(apple.blocks) + len(banana.blocks)
meta = repository.get_index_meta()
assert meta["embedding_model"] == "hash-v1"
assert meta["embedding_dim"] == "128"
assert len(runtime.calls) == 3
asyncio.run(scenario())
@pytest.mark.parametrize("failure", ["exception", "local", "missing", "dimension", "corrupt"])
def test_query_falls_back_to_exact_local_results(runtime, failure):
async def scenario():
await seed()
if failure == "exception":
runtime.error = RuntimeError("offline")
elif failure == "local":
runtime.source = "local"
elif failure == "missing":
rows("DELETE FROM routed_block_vectors WHERE block_id = (SELECT MIN(block_id) FROM blocks)")
elif failure == "dimension":
runtime.dimensions = 4
else:
rows("UPDATE routed_block_vectors SET vector = ?", ("[0, 0, 0]",))
actual = await engine.search(request())
baseline = await local_engine().search(request())
assert actual == baseline
asyncio.run(scenario())
def test_same_dimension_model_switch_never_combines_partial_spaces(runtime):
async def scenario():
apple, banana = await seed()
baseline = await local_engine().search(request())
runtime.model_id = "space-b"
assert await engine.search(request()) == baseline
await note_service.update_note(apple.note_id, markdown="apple orchard")
assert {row[0] for row in rows("SELECT DISTINCT space_id FROM routed_block_vectors")} == {"space-a", "space-b"}
assert await routed_vectors.search_remote("apple", top_k=10) is None
assert await engine.search(request()) == baseline
runtime.model_id = "space-a"
assert await engine.search(request()) == baseline
runtime.model_id = "space-b"
await note_service.update_note(banana.note_id, markdown="banana grove")
hits = await routed_vectors.search_remote("apple", top_k=10)
assert hits is not None and hits[0].id == banana.blocks[0].block_id
assert (await engine.search(request())).items[0].note_id == banana.note_id
asyncio.run(scenario())
def test_complete_spaces_coexist_but_only_requested_space_is_ranked(runtime):
async def scenario():
apple, banana = await seed()
conn = connect()
try:
with transaction(conn):
routed_vectors.store_remote(
conn, [apple.blocks[0].block_id, banana.blocks[0].block_id],
routed_vectors.RemoteEmbeddings("space-b", 3, [[1, 0, 0], [0, 1, 0]]),
)
finally:
conn.close()
assert (await engine.search(request())).items[0].note_id == banana.note_id
runtime.model_id = "space-b"
result = await engine.search(request())
assert len(result.items) == 2
assert result.items[0].note_id == apple.note_id
asyncio.run(scenario())
def test_failed_note_embedding_preserves_save_and_forces_coverage_fallback(runtime):
async def scenario():
apple, banana = await seed()
runtime.error = RuntimeError("offline")
await note_service.update_note(banana.note_id, markdown="banana changed")
assert (await note_service.get_note(banana.note_id)).markdown == "banana changed"
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == len(apple.blocks)
runtime.error = None
assert await engine.search(request()) == await local_engine().search(request())
asyncio.run(scenario())
@pytest.mark.parametrize("vectors, dimensions, space", [
([], 3, "space-a"),
([[1, 0]], 3, "space-a"),
([[0, 0, 0]], 3, "space-a"),
([[float("nan"), 0, 0]], 3, "space-a"),
([[float("inf"), 0, 0]], 3, "space-a"),
([[True, 0, 0]], 3, "space-a"),
([[1, 0, 0]], 0, "space-a"),
([[1, 0, 0]], 3, "hash-v1"),
])
def test_invalid_remote_batch_does_not_break_note_saving(runtime, vectors, dimensions, space):
runtime.result_override = SimpleNamespace(
source="api", vectors=vectors, dimensions=dimensions, model_id=space,
)
async def scenario():
note = await note_service.create_note(title="Apple", markdown="apple orchard", folder=None, tags=[])
assert (await local_engine().search(request())).items[0].note_id == note.note_id
assert await routed_vectors.search_remote("apple", top_k=10) is None
asyncio.run(scenario())
def test_remote_storage_failure_rolls_back_batch_but_keeps_local_index(runtime):
async def scenario():
await seed()
rows("""CREATE TRIGGER reject_remote_vector BEFORE INSERT ON routed_block_vectors
WHEN (SELECT content FROM blocks WHERE block_id = NEW.block_id) = 'second'
BEGIN SELECT RAISE(ABORT, 'simulated storage failure'); END""")
note = await note_service.create_note(
title="Multi", markdown="first\n\nsecond", folder=None, tags=[],
)
assert len(note.blocks) == 2
assert rows(
"SELECT COUNT(*) FROM routed_block_vectors r JOIN blocks b USING(block_id) WHERE b.note_id = ?",
(note.note_id,),
)[0][0] == 0
assert rows("SELECT COUNT(*) FROM vec_blocks")[0][0] == rows("SELECT COUNT(*) FROM blocks")[0][0]
assert (get_settings().vault_path / note.file_path).exists()
asyncio.run(scenario())
def test_rebuild_and_delete_clear_old_remote_rows_through_foreign_keys(runtime):
async def scenario():
apple, _ = await seed()
await note_service.delete_note(apple.note_id)
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == 1
runtime.source = "local"
job = await index_service.rebuild(IndexRebuildRequest())
assert job.status == "completed"
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == 0
assert rows("SELECT COUNT(*) FROM vec_blocks")[0][0] == 1
runtime.source = "api"
runtime.model_id = "space-b"
await index_service.rebuild(IndexRebuildRequest())
assert [row[0] for row in rows("SELECT space_id FROM routed_block_vectors")] == ["space-b"]
asyncio.run(scenario())
@pytest.mark.parametrize("operation", ["save", "query", "rebuild"])
def test_cancellation_propagates_and_mutations_roll_back(runtime, operation):
async def scenario():
apple, _ = await seed()
before = [tuple(row) for row in rows("SELECT * FROM routed_block_vectors ORDER BY block_id")]
runtime.error = asyncio.CancelledError()
with pytest.raises(asyncio.CancelledError):
if operation == "query":
await engine.search(request())
elif operation == "rebuild":
await index_service.rebuild(IndexRebuildRequest())
else:
await note_service.update_note(apple.note_id, markdown="changed")
assert (await note_service.get_note(apple.note_id)).markdown == "apple orchard"
assert [tuple(row) for row in rows("SELECT * FROM routed_block_vectors ORDER BY block_id")] == before
asyncio.run(scenario())
@pytest.mark.parametrize("injected", ["embedding", "vector_store", "constructor"])
def test_injected_engine_dependencies_are_respected(runtime, monkeypatch, injected):
async def scenario():
apple, _ = await seed()
target = engine
if injected == "constructor":
target = local_engine()
elif injected == "embedding":
monkeypatch.setattr(engine, "embedding", HashEmbeddingProvider())
else:
class FakeStore:
async def search(self, vector, *, top_k):
assert len(vector) == 128
return [VectorHit(id=apple.blocks[0].block_id, score=1.0)]
monkeypatch.setattr(engine, "vector_store", FakeStore())
runtime.calls.clear()
assert (await target.search(request())).items[0].note_id == apple.note_id
assert runtime.calls == []
asyncio.run(scenario())
def test_fts_skips_routing_and_hybrid_uses_routed_vector_channel(runtime, monkeypatch):
async def scenario():
_, banana = await seed()
runtime.calls.clear()
await engine.search(request(SearchMode.fts))
assert runtime.calls == []
# Empty lexical channel isolates the vector contribution to hybrid fusion.
monkeypatch.setattr(repository, "fts_search", lambda *_: [])
class PreserveOrder:
async def rerank(self, query, candidates):
return sorted(candidates, key=lambda candidate: -candidate.score)
monkeypatch.setattr(engine, "reranker", PreserveOrder())
result = await engine.search(request(SearchMode.hybrid))
assert result.items[0].note_id == banana.note_id
assert runtime.calls == [["apple"]]
asyncio.run(scenario())
def test_arbitrary_dimensions_and_extreme_finite_values(runtime):
dimensions = 257
runtime.result_override = SimpleNamespace(
source="api", model_id="space-wide", dimensions=dimensions,
vectors=[[1e308, 1e308] + [0.0] * (dimensions - 2)],
)
async def scenario():
note = await note_service.create_note(title="Apple", markdown="apple orchard", folder=None, tags=[])
hits = await routed_vectors.search_remote("apple", top_k=1)
assert hits is not None and hits[0].id == note.blocks[0].block_id
assert hits[0].score == pytest.approx(1.0)
vector = json.loads(rows("SELECT vector FROM routed_block_vectors")[0][0])
assert len(vector) == dimensions
asyncio.run(scenario())
def test_missing_runtime_uses_unchanged_local_retrieval(runtime, monkeypatch):
monkeypatch.setattr(routed_vectors, "get_model_routing", lambda: None)
async def scenario():
await seed()
assert await engine.search(request()) == await local_engine().search(request())
assert runtime.calls == []
asyncio.run(scenario())