feat(provider): 完成阶段E协议适配、国内预设与模型路由
This commit is contained in:
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
@@ -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())
|
||||
Reference in New Issue
Block a user