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
+4 -3
View File
@@ -160,10 +160,11 @@ def read_attachment(arguments: AttachmentReadArguments, _: ToolExecutionContext)
return attachment_service.read_attachment(**arguments.model_dump())
def transcribe_audio(arguments: AudioTranscribeArguments, _: ToolExecutionContext) -> dict:
return transcription_service.create_transcription(
async def transcribe_audio(arguments: AudioTranscribeArguments, _: ToolExecutionContext) -> dict:
job = await transcription_service.create_transcription(
arguments.attachment_id, arguments.language
).model_dump(mode="json")
)
return job.model_dump(mode="json")
def _register(
+4 -1
View File
@@ -7,6 +7,7 @@ from app.config import BACKEND_DIR, get_settings
from app.extensions import PluginRuntime, SkillRuntime
from app.extensions.mcp_registry import McpServerRegistry
from app.providers import MockProvider, ProviderFactory, ProviderRegistry
from app.providers.routing import ModelRoutingService
from app.providers.credentials import (
ChainedCredentialResolver,
EncryptedCredentialStore,
@@ -18,6 +19,7 @@ from app.providers.credentials import (
class ApplicationContainer:
providers: ProviderRegistry
provider_factory: ProviderFactory
model_routing: ModelRoutingService
credentials: EncryptedCredentialStore
tools: ToolRegistry
permissions: PermissionManager
@@ -33,7 +35,7 @@ def build_container() -> ApplicationContainer:
provider_factory = ProviderFactory(
ChainedCredentialResolver(credentials, EnvironmentCredentialResolver())
)
providers = ProviderRegistry()
providers = ProviderRegistry(provider_factory)
providers.register(
ProviderConfig(
provider_id="mock",
@@ -86,6 +88,7 @@ def build_container() -> ApplicationContainer:
return ApplicationContainer(
providers=providers,
provider_factory=provider_factory,
model_routing=ModelRoutingService(providers, provider_factory.credentials),
credentials=credentials,
tools=tools,
permissions=permissions,
+100 -4
View File
@@ -2,7 +2,7 @@ from datetime import datetime
from enum import Enum
from typing import Annotated, Any, Literal
from pydantic import BaseModel, ConfigDict, Field, SecretStr
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
class Contract(BaseModel):
@@ -230,6 +230,8 @@ class ModelCapability(str, Enum):
streaming = "streaming"
structured_output = "structured_output"
embedding = "embedding"
transcription = "transcription"
speaker_matching = "speaker_matching"
class ModelRequest(Contract):
@@ -757,7 +759,24 @@ class ProviderType(str, Enum):
ollama = "ollama"
class ProviderConfig(Contract):
class ProviderConnectionFields(Contract):
base_url: str | None = None
credential_id: str | None = None
@field_validator("base_url")
@classmethod
def provider_url(cls, value: str | None) -> str | None:
if value is None:
return value
from urllib.parse import urlsplit
parsed = urlsplit(value)
if (parsed.scheme not in {"http", "https"} or not parsed.hostname or
parsed.username or parsed.password or parsed.query or parsed.fragment):
raise ValueError("Base URL requires HTTP(S), without credentials, query or fragment")
return value.rstrip("/")
class ProviderConfig(ProviderConnectionFields):
provider_id: str
provider_type: ProviderType
name: str
@@ -768,7 +787,7 @@ class ProviderConfig(Contract):
capabilities: list[ModelCapability] = Field(default_factory=list)
class ProviderCreateRequest(Contract):
class ProviderCreateRequest(ProviderConnectionFields):
provider_type: ProviderType
name: str
base_url: str | None = None
@@ -777,7 +796,8 @@ class ProviderCreateRequest(Contract):
enabled: bool = True
class ProviderUpdateRequest(Contract):
class ProviderUpdateRequest(ProviderConnectionFields):
provider_type: ProviderType | None = None
name: str | None = None
base_url: str | None = None
default_model: str | None = None
@@ -796,6 +816,80 @@ class ProviderPreset(Contract):
base_url: str
default_credential_id: str | None = None
requires_credential: bool = True
logo_id: str = "custom"
description: str = ""
capabilities: list[ModelCapability] = Field(default_factory=list)
class ModelBinding(Contract):
provider_id: str = Field(min_length=1, max_length=128)
model: str = Field(min_length=1, max_length=256)
endpoint: str = Field(min_length=1, max_length=256)
dimensions: int | None = Field(default=None, ge=1, le=16384)
@field_validator("endpoint")
@classmethod
def relative_endpoint(cls, value: str) -> str:
# An endpoint is a path on the selected provider, never a second origin.
import re
if not re.fullmatch(r"/[A-Za-z0-9_/-]+", value) or value.startswith("//"):
raise ValueError("endpoint must be an absolute API path on the provider")
return value
@field_validator("model", "provider_id")
@classmethod
def non_blank(cls, value: str) -> str:
if not value.strip():
raise ValueError("value must not be blank")
return value.strip()
class ModelRoutingConfig(Contract):
version: int = Field(default=0, ge=0)
embedding: ModelBinding | None = None
transcription: ModelBinding | None = None
speaker_matching: ModelBinding | None = None
class LocalBackendStatus(Contract):
capability: Literal["embedding", "transcription", "speaker_matching"]
status: Literal["placeholder", "not_installed", "ready"]
message: str
class ModelRoutingResponse(Contract):
config: ModelRoutingConfig
local_backends: list[LocalBackendStatus]
class EmbeddingRequest(Contract):
texts: list[str] = Field(min_length=1, max_length=256)
@field_validator("texts")
@classmethod
def bound_texts(cls, value: list[str]) -> list[str]:
if sum(len(text) for text in value) > 200_000:
raise ValueError("embedding input is too large")
return value
class EmbeddingResult(Contract):
vectors: list[list[float]]
source: Literal["api", "local"]
model_id: str
dimensions: int
fallback_reason: str | None = None
class SpeakerMatchRequest(Contract):
attachment_id: str
reference_attachment_id: str
class SpeakerMatchResult(Contract):
score: float = Field(ge=0, le=1, allow_inf_nan=False)
source: Literal["api", "local"]
fallback_reason: str | None = None
class ProviderPresetListResponse(Contract):
@@ -888,6 +982,8 @@ class TranscriptionJob(Contract):
error_code: str | None = None
error_message: str | None = None
created_at: datetime
source: Literal["api", "local", "sidecar"] | None = None
fallback_reason: str | None = None
class IndexStatus(Contract):
+5 -1
View File
@@ -36,7 +36,11 @@ async def validation_error_handler(_: Request, exc: RequestValidationError) -> J
error=ErrorDetail(
code="VALIDATION_ERROR",
message="Request validation failed.",
details={"errors": exc.errors()},
# Pydantic ctx can contain exception objects; input may contain API keys.
details={"errors": [
{key: error[key] for key in ("type", "loc", "msg") if key in error}
for error in exc.errors()
]},
)
)
return JSONResponse(status_code=422, content=jsonable_encoder(body))
+163
View File
@@ -0,0 +1,163 @@
"""Native Anthropic Messages protocol with incrementally decoded content blocks."""
import json
from contextlib import aclosing
from app.contracts import MessageRole, ModelEventType, ModelRequest
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
from app.providers.http_base import (
UsageTracker, check_error, decode_tool_arguments, invalid_response, list_value,
object_value, string_value, token_count, truncated_stream,
)
from app.providers.openai_compatible import OpenAICompatibleProvider
from app.providers.tool_names import mapped_tool_names
class AnthropicMessagesProvider(OpenAICompatibleProvider):
stream_path = "/messages"
def _headers(self) -> dict[str, str]:
headers = super()._headers()
authorization = headers.pop("Authorization", None)
if authorization:
headers["x-api-key"] = authorization.removeprefix("Bearer ")
headers["anthropic-version"] = "2023-06-01"
return headers
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
systems = [request.system] if request.system else []
messages = []
for message in request.messages:
if message.role == MessageRole.system:
systems.append(message.content)
continue
if message.role == MessageRole.tool:
if not message.tool_call_id:
raise ProviderError("PROVIDER_INVALID_REQUEST", "Tool result requires a call identifier.")
role = "user"
content = [{"type": "tool_result", "tool_use_id": message.tool_call_id, "content": message.content}]
else:
role = message.role.value
content = [{"type": "text", "text": message.content}] if message.content else []
content += [{"type": "tool_use", "id": call.tool_call_id, "name": call.name,
"input": call.arguments} for call in message.tool_calls]
if not content:
continue
if messages and messages[-1]["role"] == role:
messages[-1]["content"].extend(content)
else:
messages.append({"role": role, "content": content})
payload: dict[str, object] = {"model": request.model, "messages": messages,
"max_tokens": request.max_tokens or 4096, "stream": stream}
if systems:
payload["system"] = "\n\n".join(systems)
if request.tools:
payload["tools"] = [{"name": tool.name, "description": tool.description,
"input_schema": tool.parameters} for tool in request.tools]
if request.temperature is not None:
payload["temperature"] = request.temperature
if request.response_format is not None:
format_ = request.response_format
if format_.get("type") != "json_schema":
raise ProviderError("PROVIDER_INVALID_REQUEST", "Messages requires a JSON schema response format.")
schema = object_value(format_.get("json_schema"))
payload["output_config"] = {"format": {"type": "json_schema", "schema": object_value(schema.get("schema"))}}
return payload
@mapped_tool_names
async def complete(self, request: ModelRequest) -> ProviderTurn:
data = await self._request("POST", self.stream_path, json=self._payload(request, stream=False))
texts = []
calls = []
for raw in list_value(data.get("content")):
block = object_value(raw)
if block.get("type") == "text":
texts.append(string_value(block.get("text")))
elif block.get("type") == "tool_use":
calls.append(ProviderToolCall(
tool_call_id=string_value(block.get("id"), nonempty=True),
name=string_value(block.get("name"), nonempty=True),
arguments=decode_tool_arguments(block.get("input")),
))
return ProviderTurn(text="".join(texts) or None, tool_calls=calls,
**UsageTracker(cache_tokens=True).update(data.get("usage") or {}))
async def _events(self, request: ModelRequest):
blocks: dict[int, dict] = {}
usage = UsageTracker(cache_tokens=True)
started = False
async with aclosing(self._stream_json(self._payload(request, stream=True))) as chunks:
async for data in chunks:
kind = string_value(data.get("type"), nonempty=True)
if kind == "message_start":
if started:
raise invalid_response()
started = True
message = object_value(data.get("message"))
check_error(message)
if message.get("usage") is not None:
yield ModelEventType.usage, usage.update(message["usage"])
elif kind == "content_block_start":
index = token_count(data.get("index"))
if not started or index in blocks:
raise invalid_response()
block = dict(object_value(data.get("content_block")))
blocks[index] = block
block["closed"] = False
if block.get("type") == "tool_use":
block["id"] = string_value(block.get("id"), nonempty=True)
block["name"] = string_value(block.get("name"), nonempty=True)
block["arguments"] = ""
block["input"] = object_value(block.get("input", {}))
yield ModelEventType.tool_call_start, {"tool_call_id": block["id"], "name": block["name"]}
elif block.get("type") == "text" and block.get("text"):
yield ModelEventType.text_delta, {"text": string_value(block["text"])}
elif block.get("type") == "thinking" and block.get("thinking"):
yield ModelEventType.thinking_delta, {"text": string_value(block["thinking"])}
elif kind == "content_block_delta":
block = blocks.get(token_count(data.get("index")))
if block is None or block["closed"]:
raise invalid_response()
delta = object_value(data.get("delta"))
delta_type = delta.get("type")
if delta_type == "text_delta":
if block.get("type") != "text":
raise invalid_response()
yield ModelEventType.text_delta, {"text": string_value(delta.get("text"))}
elif delta_type == "thinking_delta":
if block.get("type") != "thinking":
raise invalid_response()
yield ModelEventType.thinking_delta, {"text": string_value(delta.get("thinking"))}
elif delta_type == "input_json_delta" and block.get("type") == "tool_use":
fragment = string_value(delta.get("partial_json"))
block["arguments"] += fragment
yield ModelEventType.tool_call_delta, {"tool_call_id": block["id"], "arguments_delta": fragment}
# Signatures and future delta types have no representation in ModelEvent.
elif kind == "content_block_stop":
block = blocks.get(token_count(data.get("index")))
if block is None or block["closed"]:
raise invalid_response()
block["closed"] = True
if block.get("type") == "tool_use":
if block["arguments"]:
decode_tool_arguments(block["arguments"])
else:
yield ModelEventType.tool_call_delta, {
"tool_call_id": block["id"], "arguments_delta": json.dumps(block["input"]),
}
yield ModelEventType.tool_call_end, {"tool_call_id": block["id"]}
elif kind == "message_delta":
if not started:
raise invalid_response()
object_value(data.get("delta"))
if data.get("usage") is not None:
yield ModelEventType.usage, usage.update(data["usage"])
elif kind == "message_stop":
if not started:
raise invalid_response()
if any(not block["closed"] for block in blocks.values()):
raise truncated_stream()
return
elif kind == "[DONE]":
raise truncated_stream()
raise truncated_stream()
+46 -1
View File
@@ -16,6 +16,18 @@ class ProviderFactory:
self.credentials = ProviderCredentialResolver(credentials)
def build(self, config: ProviderConfig) -> ModelProvider:
if config.provider_type == ProviderType.openai_responses:
from app.providers.openai_responses import OpenAIResponsesProvider
return OpenAIResponsesProvider(
base_url=config.base_url or "https://api.openai.com/v1",
credential_id=config.credential_id, credentials=self.credentials,
)
if config.provider_type == ProviderType.anthropic_messages:
from app.providers.anthropic_messages import AnthropicMessagesProvider
return AnthropicMessagesProvider(
base_url=config.base_url or "https://api.anthropic.com/v1",
credential_id=config.credential_id, credentials=self.credentials,
)
if config.provider_type in {
ProviderType.openai_chat,
ProviderType.openai_compatible,
@@ -31,7 +43,7 @@ class ProviderFactory:
@staticmethod
def presets() -> list[ProviderPreset]:
return [
presets = [
ProviderPreset(
preset_id="openai",
name="OpenAI",
@@ -54,12 +66,45 @@ class ProviderFactory:
requires_credential=False,
),
]
# General API endpoints. Coding-plan endpoints and keys are separate products.
domestic = [
("kimi", "Kimi / 月之暗面", "https://api.moonshot.cn/v1", [], "长上下文对话;模型以账号权限为准。"),
("qwen", "阿里云百炼", "https://dashscope.aliyuncs.com/compatible-mode/v1", [ModelCapability.embedding], "中国内地兼容接口;海外地域需修改地址。"),
("zhipu", "智谱 GLM", "https://open.bigmodel.cn/api/paas/v4", [ModelCapability.embedding], "通用 APICoding Plan 请使用其专用地址。"),
("volcengine", "火山方舟 / 豆包", "https://ark.cn-beijing.volces.com/api/v3", [ModelCapability.embedding], "按账号填写模型 ID 或推理接入点 ID。"),
("siliconflow", "硅基流动", "https://api.siliconflow.cn/v1", [ModelCapability.embedding, ModelCapability.transcription], "支持兼容 Embedding 和音频转写接口。"),
("baidu", "百度千帆", "https://qianfan.baidubce.com/v2", [ModelCapability.embedding], "使用千帆 API Key;模型列表取决于账号。"),
("hunyuan", "腾讯混元", "https://api.hunyuan.cloud.tencent.com/v1", [], "OpenAI 兼容对话接口。"),
("minimax", "MiniMax", "https://api.minimaxi.com/v1", [], "文本对话兼容接口;其他媒体协议需独立适配。"),
("stepfun", "阶跃星辰", "https://api.stepfun.com/v1", [], "通用 API;Step Plan 请使用其专用地址。"),
]
for preset_id, name, url, extra, description in domestic:
presets.append(ProviderPreset(
preset_id=preset_id, name=name, provider_type=ProviderType.openai_compatible,
base_url=url, default_credential_id=preset_id, logo_id=preset_id,
capabilities=[ModelCapability.chat, *extra], description=description,
))
presets.extend([
ProviderPreset(preset_id="openai-responses", name="OpenAI Responses", provider_type=ProviderType.openai_responses,
base_url="https://api.openai.com/v1", default_credential_id="openai", logo_id="openai"),
ProviderPreset(preset_id="anthropic", name="Anthropic / Claude", provider_type=ProviderType.anthropic_messages,
base_url="https://api.anthropic.com/v1", default_credential_id="anthropic", logo_id="anthropic"),
])
for preset in presets:
if preset.logo_id == "custom":
preset.logo_id = preset.preset_id
if not preset.capabilities:
preset.capabilities = [ModelCapability.chat]
presets[0].capabilities += [ModelCapability.embedding, ModelCapability.transcription]
return presets
@staticmethod
def capabilities(provider_type: ProviderType) -> list[ModelCapability]:
if provider_type in {
ProviderType.openai_chat,
ProviderType.openai_compatible,
ProviderType.openai_responses,
ProviderType.anthropic_messages,
}:
return [
ModelCapability.chat,
+230
View File
@@ -1,9 +1,13 @@
import json
from collections.abc import AsyncIterator
from contextlib import aclosing
from datetime import datetime, timezone
import httpx
from app.contracts import ModelEvent, ModelEventType, ModelRequest
from app.providers.base import ProviderError, ProviderTurn
from app.providers.tool_names import prepare_tool_names
class TurnStreamingMixin:
@@ -80,3 +84,229 @@ def decode_tool_arguments(value: object) -> dict[str, object]:
if not isinstance(decoded, dict):
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Tool arguments must be an object.")
return decoded
def invalid_response() -> ProviderError:
return ProviderError("PROVIDER_INVALID_RESPONSE", "Provider returned an invalid response.")
def truncated_stream() -> ProviderError:
return ProviderError("PROVIDER_STREAM_TRUNCATED", "Provider stream ended before completion.")
def object_value(value: object) -> dict:
if not isinstance(value, dict):
raise invalid_response()
return value
def list_value(value: object) -> list:
if not isinstance(value, list):
raise invalid_response()
return value
def string_value(value: object, *, nonempty: bool = False) -> str:
if not isinstance(value, str) or (nonempty and not value):
raise invalid_response()
return value
def token_count(value: object) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise invalid_response()
return value
def remote_error(value: object) -> ProviderError:
# Never reflect upstream messages, URLs, request bodies or credentials.
error = value if isinstance(value, dict) else {}
code = error.get("code") or error.get("type")
mapping = {
"authentication_error": "PROVIDER_AUTH_FAILED",
"invalid_api_key": "PROVIDER_AUTH_FAILED",
"permission_error": "PROVIDER_AUTH_FAILED",
"rate_limit_error": "PROVIDER_RATE_LIMITED",
"rate_limit_exceeded": "PROVIDER_RATE_LIMITED",
"insufficient_quota": "PROVIDER_RATE_LIMITED",
"not_found_error": "MODEL_NOT_FOUND",
"model_not_found": "MODEL_NOT_FOUND",
"invalid_request_error": "PROVIDER_INVALID_REQUEST",
"context_length_exceeded": "PROVIDER_INVALID_REQUEST",
}
mapped = mapping.get(code, "PROVIDER_UNAVAILABLE") if isinstance(code, str) else "PROVIDER_UNAVAILABLE"
return ProviderError(mapped, "Provider could not complete the request.")
def check_error(data: dict) -> None:
if data.get("error") is not None or data.get("type") == "error":
raise remote_error(data.get("error") or data)
class UsageTracker:
"""Merge cumulative snapshots, including partial usage updates."""
def __init__(self, input_key: str = "input_tokens", output_key: str = "output_tokens",
*, cache_tokens: bool = False) -> None:
self.input_key = input_key
self.output_key = output_key
self.cache_tokens = cache_tokens
self.counts: dict[str, int] = {}
def update(self, value: object) -> dict[str, int]:
usage = object_value(value)
keys = [self.input_key, self.output_key]
if self.cache_tokens:
keys += ["cache_creation_input_tokens", "cache_read_input_tokens"]
for key in keys:
if key in usage:
self.counts[key] = max(self.counts.get(key, 0), token_count(usage[key]))
inputs = self.counts.get(self.input_key, 0)
if self.cache_tokens:
inputs += sum(self.counts.get(key, 0) for key in keys[2:])
return {"input_tokens": inputs, "output_tokens": self.counts.get(self.output_key, 0)}
class EventStreamingMixin:
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
sequence = 0
status = "completed"
try:
request, originals = prepare_tool_names(request)
# Closing the public iterator must synchronously close every nested iterator.
async with aclosing(self._events(request)) as events:
async for kind, data in events:
if kind == ModelEventType.tool_call_start and "name" in data:
data = {**data, "name": originals.get(data["name"], data["name"])}
if kind == ModelEventType.usage:
data = {**data, "total_tokens": data["input_tokens"] + data["output_tokens"]}
yield ModelEvent(event=kind, data=data, sequence=sequence,
timestamp=datetime.now(timezone.utc))
sequence += 1
except ProviderError as exc:
status = "failed"
yield ModelEvent(event=ModelEventType.error, sequence=sequence,
data={"code": exc.code, "message": exc.message},
timestamp=datetime.now(timezone.utc))
sequence += 1
except (ValueError, TypeError, KeyError, IndexError, AttributeError, OverflowError):
status = "failed"
error = invalid_response()
yield ModelEvent(event=ModelEventType.error, sequence=sequence,
data={"code": error.code, "message": error.message},
timestamp=datetime.now(timezone.utc))
sequence += 1
# CancelledError and GeneratorExit deliberately propagate without a Done event.
yield ModelEvent(event=ModelEventType.done, sequence=sequence,
data={"status": status},
timestamp=datetime.now(timezone.utc))
async def sse_objects(response: httpx.Response) -> AsyncIterator[dict]:
"""Read SSE frames, accepting the adjacent data lines used by some gateways."""
parts: list[str] = []
event_name = ""
def decode() -> dict:
value = "\n".join(parts)
if value.strip() == "[DONE]":
return {"type": "[DONE]"}
try:
data = object_value(json.loads(value))
except (ValueError, TypeError) as exc:
raise invalid_response() from exc
if event_name and "type" not in data:
data["type"] = event_name
check_error(data)
return data
async for line in response.aiter_lines():
if not line:
if parts:
yield decode()
parts = []
event_name = ""
elif line.startswith(":"):
continue
elif line.startswith("event:"):
if parts:
yield decode()
parts = []
event_name = line[6:].strip()
elif line.startswith("data:"):
if parts:
# Legacy compatible endpoints sometimes omit blank separators.
try:
json.loads("\n".join(parts))
except ValueError:
pass
else:
yield decode()
parts = []
event_name = ""
parts.append(line[5:].removeprefix(" "))
if parts:
yield decode()
class HTTPProviderMixin:
stream_path = "/chat/completions"
stream_format = "sse"
def _headers(self) -> dict[str, str]:
return {"Content-Type": "application/json"}
@staticmethod
def _status_error(exc: httpx.HTTPStatusError) -> ProviderError:
status = exc.response.status_code
code = {400: "PROVIDER_INVALID_REQUEST", 401: "PROVIDER_AUTH_FAILED",
403: "PROVIDER_AUTH_FAILED", 404: "MODEL_NOT_FOUND",
408: "PROVIDER_TIMEOUT", 413: "PROVIDER_INVALID_REQUEST",
422: "PROVIDER_INVALID_REQUEST", 429: "PROVIDER_RATE_LIMITED"}.get(
status, "PROVIDER_UNAVAILABLE")
return ProviderError(code, f"Provider returned HTTP {status}.")
async def _request(self, method: str, path: str, **kwargs) -> dict:
headers = self._headers()
try:
async with httpx.AsyncClient(timeout=self.timeout_seconds, transport=self.transport) as client:
response = await client.request(method, f"{self.base_url}{path}", headers=headers, **kwargs)
response.raise_for_status()
data = object_value(response.json())
check_error(data)
return data
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
except httpx.HTTPStatusError as exc:
raise self._status_error(exc) from exc
except httpx.HTTPError as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
except (ValueError, TypeError) as exc:
raise invalid_response() from exc
async def _stream_json(self, payload: dict[str, object]) -> AsyncIterator[dict]:
headers = self._headers()
headers["Accept"] = "text/event-stream" if self.stream_format == "sse" else "application/x-ndjson"
try:
async with httpx.AsyncClient(timeout=self.timeout_seconds, transport=self.transport) as client:
async with client.stream("POST", f"{self.base_url}{self.stream_path}",
headers=headers, json=payload) as response:
response.raise_for_status()
if self.stream_format == "sse":
async with aclosing(sse_objects(response)) as objects:
async for data in objects:
yield data
else:
async for line in response.aiter_lines():
if line.strip():
data = object_value(json.loads(line))
check_error(data)
yield data
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
except httpx.HTTPStatusError as exc:
raise self._status_error(exc) from exc
except httpx.HTTPError as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
except (ValueError, TypeError) as exc:
raise invalid_response() from exc
+88 -178
View File
@@ -1,16 +1,22 @@
from uuid import uuid4
import json
from collections.abc import AsyncIterator
from datetime import datetime, timezone
from contextlib import aclosing
from uuid import uuid4
import httpx
from app.contracts import ModelCapability, ModelEvent, ModelEventType, ModelInfo, ModelRequest
from app.contracts import MessageRole, ModelCapability, ModelEventType, ModelInfo, ModelRequest
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments
from app.providers.tool_names import mapped_tool_names
from app.providers.http_base import (
EventStreamingMixin, HTTPProviderMixin, UsageTracker, decode_tool_arguments,
invalid_response, list_value, object_value, string_value, truncated_stream,
)
class OllamaProvider(TurnStreamingMixin):
class OllamaProvider(EventStreamingMixin, HTTPProviderMixin):
stream_path = "/api/chat"
stream_format = "jsonl"
def __init__(
self,
base_url: str = "http://127.0.0.1:11434",
@@ -21,126 +27,55 @@ class OllamaProvider(TurnStreamingMixin):
self.timeout_seconds = timeout_seconds
self.transport = transport
@mapped_tool_names
async def complete(self, request: ModelRequest) -> ProviderTurn:
messages = []
if request.system:
messages.append({"role": "system", "content": request.system})
for message in request.messages:
item: dict[str, object] = {
"role": message.role.value,
"content": message.content,
}
if message.tool_calls:
item["tool_calls"] = [
{
"function": {
"name": call.name,
"arguments": call.arguments,
}
}
for call in message.tool_calls
]
messages.append(item)
payload: dict[str, object] = {
"model": request.model,
"messages": messages,
"stream": False,
}
if request.tools:
payload["tools"] = [
{
"type": "function",
"function": {
"name": tool.name,
"description": tool.description,
"parameters": tool.parameters,
},
}
for tool in request.tools
]
data = await self._request("POST", "/api/chat", json=payload)
message = data.get("message") or {}
tool_calls = []
for raw_call in message.get("tool_calls") or []:
function = raw_call.get("function") or {}
tool_calls.append(
ProviderToolCall(
tool_call_id=raw_call.get("id") or f"call_{uuid4().hex}",
name=function.get("name") or "",
arguments=decode_tool_arguments(function.get("arguments", {})),
)
)
return ProviderTurn(
text=message.get("content") or None,
tool_calls=tool_calls,
input_tokens=int(data.get("prompt_eval_count") or 0),
output_tokens=int(data.get("eval_count") or 0),
data = await self._request("POST", self.stream_path, json=self._chat_payload(request, stream=False))
message = object_value(data.get("message"))
calls = [self._tool_call(raw) for raw in list_value(message.get("tool_calls", []))]
content = message.get("content")
if content is not None:
content = string_value(content)
return ProviderTurn(text=content or None, tool_calls=calls,
**UsageTracker("prompt_eval_count", "eval_count").update(data))
@staticmethod
def _tool_call(raw: object) -> ProviderToolCall:
call = object_value(raw)
function = object_value(call.get("function"))
return ProviderToolCall(
tool_call_id=string_value(call.get("id") or f"call_{uuid4().hex}"),
name=string_value(function.get("name"), nonempty=True),
arguments=decode_tool_arguments(function.get("arguments", {})),
)
async def list_models(self) -> list[ModelInfo]:
data = await self._request("GET", "/api/tags")
return [
ModelInfo(
model=item["name"],
display_name=item.get("name", ""),
capabilities=[ModelCapability.chat, ModelCapability.streaming],
)
for item in data.get("models", [])
if isinstance(item, dict) and item.get("name")
]
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
payload = self._chat_payload(request, stream=True)
sequence = 0
def event(kind: ModelEventType, data: dict | None = None) -> ModelEvent:
nonlocal sequence
item = ModelEvent(
event=kind, sequence=sequence, data=data or {},
timestamp=datetime.now(timezone.utc),
)
sequence += 1
return item
try:
async for data in self._stream_json(payload):
message = data.get("message") or {}
async def _events(self, request: ModelRequest):
usage = UsageTracker("prompt_eval_count", "eval_count")
async with aclosing(self._stream_json(self._chat_payload(request, stream=True))) as chunks:
async for data in chunks:
message = object_value(data.get("message", {}))
if message.get("thinking"):
yield event(ModelEventType.thinking_delta, {"text": message["thinking"]})
yield ModelEventType.thinking_delta, {"text": string_value(message["thinking"])}
if message.get("content"):
yield event(ModelEventType.text_delta, {"text": message["content"]})
for raw_call in message.get("tool_calls") or []:
function = raw_call.get("function") or {}
call_id = raw_call.get("id") or f"call_{uuid4().hex}"
yield event(
ModelEventType.tool_call_start,
{"tool_call_id": call_id, "name": function.get("name") or ""},
)
yield event(
ModelEventType.tool_call_delta,
{
"tool_call_id": call_id,
"arguments_delta": json.dumps(
function.get("arguments") or {}, ensure_ascii=False
),
},
)
yield event(ModelEventType.tool_call_end, {"tool_call_id": call_id})
if data.get("done"):
yield event(
ModelEventType.usage,
{
"input_tokens": int(data.get("prompt_eval_count") or 0),
"output_tokens": int(data.get("eval_count") or 0),
},
)
yield event(ModelEventType.done)
except ProviderError as exc:
yield event(ModelEventType.error, {"code": exc.code, "message": exc.message})
yield event(ModelEventType.done)
yield ModelEventType.text_delta, {"text": string_value(message["content"])}
for raw in list_value(message.get("tool_calls", [])):
call = self._tool_call(raw)
yield ModelEventType.tool_call_start, {"tool_call_id": call.tool_call_id, "name": call.name}
yield ModelEventType.tool_call_delta, {
"tool_call_id": call.tool_call_id,
"arguments_delta": json.dumps(call.arguments, ensure_ascii=False),
}
yield ModelEventType.tool_call_end, {"tool_call_id": call.tool_call_id}
if "done" in data and not isinstance(data["done"], bool):
raise invalid_response()
if "prompt_eval_count" in data or "eval_count" in data or data.get("done"):
yield ModelEventType.usage, usage.update(data)
if data.get("done") is True:
return
raise truncated_stream()
def _chat_payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
messages = []
names: dict[str, str] = {}
if request.system:
messages.append({"role": "system", "content": request.system})
for message in request.messages:
@@ -150,53 +85,49 @@ class OllamaProvider(TurnStreamingMixin):
{"function": {"name": call.name, "arguments": call.arguments}}
for call in message.tool_calls
]
names.update({call.tool_call_id: call.name for call in message.tool_calls})
if message.role == MessageRole.tool:
name = message.name or names.get(message.tool_call_id or "")
if name:
item["tool_name"] = name
messages.append(item)
payload: dict[str, object] = {
"model": request.model, "messages": messages, "stream": stream
"model": request.model, "messages": messages, "stream": stream,
}
if request.tools:
payload["tools"] = [
{
"type": "function",
"function": {
"name": tool.name,
"description": tool.description,
"parameters": tool.parameters,
},
}
for tool in request.tools
{"type": "function", "function": {
"name": tool.name, "description": tool.description, "parameters": tool.parameters,
}} for tool in request.tools
]
options = {}
if request.temperature is not None:
options["temperature"] = request.temperature
if request.max_tokens is not None:
options["num_predict"] = request.max_tokens
if options:
payload["options"] = options
if request.response_format:
format_ = request.response_format
if format_.get("type") == "json_object":
payload["format"] = "json"
elif format_.get("type") == "json_schema":
payload["format"] = object_value(object_value(format_.get("json_schema")).get("schema"))
else:
payload["format"] = format_
return payload
async def _stream_json(self, payload: dict[str, object]) -> AsyncIterator[dict]:
try:
async with httpx.AsyncClient(
timeout=self.timeout_seconds, transport=self.transport
) as client:
async with client.stream(
"POST", f"{self.base_url}/api/chat", json=payload
) as response:
response.raise_for_status()
async for line in response.aiter_lines():
if not line.strip():
continue
try:
data = json.loads(line)
except json.JSONDecodeError as exc:
raise ProviderError(
"PROVIDER_INVALID_RESPONSE", "Ollama returned invalid JSONL."
) from exc
if isinstance(data, dict):
yield data
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Ollama request timed out.") from exc
except httpx.HTTPStatusError as exc:
raise ProviderError(
"MODEL_NOT_FOUND" if exc.response.status_code == 404 else "PROVIDER_UNAVAILABLE",
f"Ollama returned HTTP {exc.response.status_code}.",
) from exc
except httpx.HTTPError as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Ollama is unavailable.") from exc
async def list_models(self) -> list[ModelInfo]:
data = await self._request("GET", "/api/tags")
return [
ModelInfo(
model=string_value(item["name"]), display_name=item["name"],
capabilities=([ModelCapability.embedding] if "embed" in item["name"].lower()
else [ModelCapability.chat, ModelCapability.streaming]),
)
for item in list_value(data.get("models"))
if isinstance(item, dict) and isinstance(item.get("name"), str) and item["name"]
]
async def test_connection(self, model: str | None = None) -> tuple[bool, str]:
try:
@@ -206,24 +137,3 @@ class OllamaProvider(TurnStreamingMixin):
if model and model not in {item.model for item in models}:
return False, f"Model is not installed: {model}"
return True, f"Connected; discovered {len(models)} local model(s)."
async def _request(self, method: str, path: str, **kwargs) -> dict:
try:
async with httpx.AsyncClient(
timeout=self.timeout_seconds, transport=self.transport
) as client:
response = await client.request(method, f"{self.base_url}{path}", **kwargs)
response.raise_for_status()
data = response.json()
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Ollama request timed out.") from exc
except httpx.HTTPStatusError as exc:
raise ProviderError(
"MODEL_NOT_FOUND" if exc.response.status_code == 404 else "PROVIDER_UNAVAILABLE",
f"Ollama returned HTTP {exc.response.status_code}.",
) from exc
except (httpx.HTTPError, ValueError) as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Ollama is unavailable.") from exc
if not isinstance(data, dict):
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Ollama returned non-object JSON.")
return data
+106 -209
View File
@@ -1,24 +1,20 @@
import json
from collections.abc import AsyncIterator
from datetime import datetime, timezone
from contextlib import aclosing
from uuid import uuid4
import httpx
from app.contracts import (
MessageRole,
ModelCapability,
ModelEvent,
ModelEventType,
ModelInfo,
ModelRequest,
)
from app.contracts import MessageRole, ModelCapability, ModelEventType, ModelInfo, ModelRequest
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
from app.providers.credentials import CredentialResolver, CredentialStoreError
from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments
from app.providers.tool_names import mapped_tool_names
from app.providers.http_base import (
EventStreamingMixin, HTTPProviderMixin, UsageTracker, decode_tool_arguments,
invalid_response, list_value, object_value, string_value, token_count, truncated_stream,
)
class OpenAICompatibleProvider(TurnStreamingMixin):
class OpenAICompatibleProvider(EventStreamingMixin, HTTPProviderMixin):
def __init__(
self,
base_url: str,
@@ -33,50 +29,37 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
self.timeout_seconds = timeout_seconds
self.transport = transport
@mapped_tool_names
async def complete(self, request: ModelRequest) -> ProviderTurn:
payload = self._payload(request, stream=False)
data = await self._request("POST", "/chat/completions", json=payload)
try:
message = data["choices"][0]["message"]
except (KeyError, IndexError, TypeError) as exc:
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Missing completion message.") from exc
tool_calls = []
for raw_call in message.get("tool_calls") or []:
function = raw_call.get("function") or {}
tool_calls.append(
ProviderToolCall(
tool_call_id=raw_call.get("id") or f"call_{uuid4().hex}",
name=function.get("name") or "",
arguments=decode_tool_arguments(function.get("arguments", "{}")),
)
)
usage = data.get("usage") or {}
return ProviderTurn(
text=message.get("content"),
tool_calls=tool_calls,
input_tokens=int(usage.get("prompt_tokens") or 0),
output_tokens=int(usage.get("completion_tokens") or 0),
)
data = await self._request("POST", self.stream_path, json=self._payload(request, stream=False))
choices = list_value(data.get("choices"))
if not choices:
raise invalid_response()
message = object_value(object_value(choices[0]).get("message"))
calls = []
for raw in list_value(message.get("tool_calls", [])):
raw = object_value(raw)
function = object_value(raw.get("function"))
calls.append(ProviderToolCall(
tool_call_id=string_value(raw.get("id") or f"call_{uuid4().hex}"),
name=string_value(function.get("name"), nonempty=True),
arguments=decode_tool_arguments(function.get("arguments", "{}")),
))
text = message.get("content")
if text is not None:
text = string_value(text)
usage = UsageTracker("prompt_tokens", "completion_tokens").update(data.get("usage") or {})
return ProviderTurn(text=text, tool_calls=calls, **usage)
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
payload: dict[str, object] = {
"model": request.model,
"messages": self._messages(request),
"stream": stream,
"model": request.model, "messages": self._messages(request), "stream": stream,
}
if request.tools:
payload["tools"] = [
{
"type": "function",
"function": {
"name": tool.name,
"description": tool.description,
"parameters": tool.parameters,
},
}
for tool in request.tools
{"type": "function", "function": {
"name": tool.name, "description": tool.description, "parameters": tool.parameters,
}} for tool in request.tools
]
if request.temperature is not None:
payload["temperature"] = request.temperature
@@ -84,124 +67,81 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
payload["max_tokens"] = request.max_tokens
if request.response_format is not None:
payload["response_format"] = request.response_format
if stream:
payload["stream_options"] = {"include_usage": True}
return payload
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
sequence = 0
open_calls: dict[int, str] = {}
def event(kind: ModelEventType, data: dict | None = None) -> ModelEvent:
nonlocal sequence
item = ModelEvent(
event=kind,
sequence=sequence,
data=data or {},
timestamp=datetime.now(timezone.utc),
)
sequence += 1
return item
try:
async for data in self._stream_json(self._payload(request, stream=True)):
usage = data.get("usage") or {}
if usage:
yield event(
ModelEventType.usage,
{
"input_tokens": int(usage.get("prompt_tokens") or 0),
"output_tokens": int(usage.get("completion_tokens") or 0),
},
)
choices = data.get("choices") or []
async def _events(self, request: ModelRequest):
calls: dict[int, dict] = {}
usage = UsageTracker("prompt_tokens", "completion_tokens")
finished = False
seen = False
async with aclosing(self._stream_json(self._payload(request, stream=True))) as chunks:
async for data in chunks:
if data.get("type") == "[DONE]":
if not seen:
raise invalid_response()
finished = True
break
if data.get("usage") is not None:
yield ModelEventType.usage, usage.update(data["usage"])
choices = list_value(data.get("choices", []))
if not choices:
continue
choice = choices[0]
delta = choice.get("delta") or {}
seen = True
choice = object_value(choices[0])
delta = object_value(choice.get("delta") or {})
if delta.get("reasoning_content"):
yield event(
ModelEventType.thinking_delta,
{"text": delta["reasoning_content"]},
)
yield ModelEventType.thinking_delta, {"text": string_value(delta["reasoning_content"])}
if delta.get("content"):
yield event(ModelEventType.text_delta, {"text": delta["content"]})
for raw_call in delta.get("tool_calls") or []:
index = int(raw_call.get("index") or 0)
function = raw_call.get("function") or {}
call_id = raw_call.get("id") or open_calls.get(index) or f"call_{uuid4().hex}"
if index not in open_calls:
open_calls[index] = call_id
yield event(
ModelEventType.tool_call_start,
{"tool_call_id": call_id, "name": function.get("name") or ""},
)
if function.get("arguments"):
yield event(
ModelEventType.tool_call_delta,
{
"tool_call_id": open_calls[index],
"arguments_delta": function["arguments"],
},
)
if choice.get("finish_reason") == "tool_calls":
for call_id in open_calls.values():
yield event(
ModelEventType.tool_call_end, {"tool_call_id": call_id}
)
open_calls.clear()
for call_id in open_calls.values():
yield event(ModelEventType.tool_call_end, {"tool_call_id": call_id})
yield event(ModelEventType.done)
except ProviderError as exc:
yield event(ModelEventType.error, {"code": exc.code, "message": exc.message})
yield event(ModelEventType.done)
async def _stream_json(self, payload: dict[str, object]) -> AsyncIterator[dict]:
headers = self._headers()
try:
async with httpx.AsyncClient(
timeout=self.timeout_seconds, transport=self.transport
) as client:
async with client.stream(
"POST", f"{self.base_url}/chat/completions", headers=headers, json=payload
) as response:
response.raise_for_status()
async for line in response.aiter_lines():
if not line.startswith("data:"):
continue
value = line[5:].strip()
if not value or value == "[DONE]":
continue
try:
data = json.loads(value)
except json.JSONDecodeError as exc:
raise ProviderError(
"PROVIDER_INVALID_RESPONSE", "Provider returned invalid SSE JSON."
) from exc
if isinstance(data, dict):
yield data
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
except httpx.HTTPStatusError as exc:
raise self._status_error(exc) from exc
except httpx.HTTPError as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
yield ModelEventType.text_delta, {"text": string_value(delta["content"])}
for raw in list_value(delta.get("tool_calls", [])):
raw = object_value(raw)
index = token_count(raw.get("index", 0))
function = object_value(raw.get("function") or {})
call = calls.setdefault(index, {"id": "", "name": "", "arguments": "", "started": False})
if raw.get("id"):
call["id"] = string_value(raw["id"])
if function.get("name"):
call["name"] += string_value(function["name"])
fragment = string_value(function.get("arguments", ""))
call["arguments"] += fragment
if not call["started"] and call["name"]:
call["id"] = call["id"] or f"call_{uuid4().hex}"
call["started"] = True
yield ModelEventType.tool_call_start, {"tool_call_id": call["id"], "name": call["name"]}
fragment = call["arguments"]
if call["started"] and fragment:
yield ModelEventType.tool_call_delta, {"tool_call_id": call["id"], "arguments_delta": fragment}
if choice.get("finish_reason"):
finished = True
if not finished:
raise truncated_stream()
for call in calls.values():
if not call["started"]:
raise invalid_response()
decode_tool_arguments(call["arguments"] or "{}")
yield ModelEventType.tool_call_end, {"tool_call_id": call["id"]}
async def list_models(self) -> list[ModelInfo]:
data = await self._request("GET", "/models")
return [
ModelInfo(
model=item["id"],
display_name=item["id"],
capabilities=[
ModelCapability.chat,
ModelCapability.tool_calling,
ModelCapability.streaming,
],
)
for item in data.get("data", [])
if isinstance(item, dict) and item.get("id")
]
return [ModelInfo(model=string_value(item["id"]), display_name=item["id"],
capabilities=self._model_capabilities(string_value(item["id"])))
for item in list_value(data.get("data"))
if isinstance(item, dict) and item.get("id")]
@staticmethod
def _model_capabilities(model: str) -> list[ModelCapability]:
# /models does not advertise capabilities. Avoid known non-chat families;
# these are discovery hints, not a guarantee of support by a gateway.
name = model.lower()
if "embed" in name or name.startswith(("bge-", "bge/")):
return [ModelCapability.embedding]
if any(marker in name for marker in (
"whisper", "tts", "transcri", "audio", "realtime", "dall-e", "image", "moderation", "rerank",
)):
return []
return [ModelCapability.chat]
async def test_connection(self, model: str | None = None) -> tuple[bool, str]:
try:
@@ -217,73 +157,30 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
if request.system:
result.append({"role": "system", "content": request.system})
for message in request.messages:
item: dict[str, object] = {
"role": message.role.value,
"content": message.content,
}
item: dict[str, object] = {"role": message.role.value, "content": message.content}
if message.name:
item["name"] = message.name
if message.role == MessageRole.tool and message.tool_call_id:
item["tool_call_id"] = message.tool_call_id
if message.tool_calls:
item["tool_calls"] = [
{
"id": call.tool_call_id,
"type": "function",
"function": {
"name": call.name,
"arguments": json.dumps(call.arguments),
},
}
for call in message.tool_calls
{"id": call.tool_call_id, "type": "function", "function": {
"name": call.name, "arguments": json.dumps(call.arguments),
}} for call in message.tool_calls
]
result.append(item)
return result
async def _request(self, method: str, path: str, **kwargs) -> dict:
headers = self._headers()
try:
async with httpx.AsyncClient(
timeout=self.timeout_seconds, transport=self.transport
) as client:
response = await client.request(
method, f"{self.base_url}{path}", headers=headers, **kwargs
)
response.raise_for_status()
data = response.json()
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
except httpx.HTTPStatusError as exc:
raise self._status_error(exc) from exc
except (httpx.HTTPError, ValueError) as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
if not isinstance(data, dict):
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Provider returned non-object JSON.")
return data
def _headers(self) -> dict[str, str]:
headers = {"Content-Type": "application/json"}
try:
api_key = self.credentials.resolve(self.credential_id)
except CredentialStoreError as exc:
raise ProviderError(
"PROVIDER_CREDENTIAL_UNAVAILABLE",
"Credential could not be decrypted by the AI Core.",
) from exc
raise ProviderError("PROVIDER_CREDENTIAL_UNAVAILABLE",
"Credential could not be decrypted by the AI Core.") from exc
if self.credential_id and not api_key:
raise ProviderError(
"PROVIDER_CREDENTIAL_MISSING",
f'Credential "{self.credential_id}" is not available in the AI Core process.',
)
raise ProviderError("PROVIDER_CREDENTIAL_MISSING",
"Credential is not available in the AI Core process.")
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
return headers
@staticmethod
def _status_error(exc: httpx.HTTPStatusError) -> ProviderError:
code = {
401: "PROVIDER_AUTH_FAILED",
404: "MODEL_NOT_FOUND",
429: "PROVIDER_RATE_LIMITED",
}.get(exc.response.status_code, "PROVIDER_UNAVAILABLE")
return ProviderError(code, f"Provider returned HTTP {exc.response.status_code}.")
+168
View File
@@ -0,0 +1,168 @@
"""Native /responses adapter; stateless history uses function_call/output items."""
import json
from contextlib import aclosing
from app.contracts import MessageRole, ModelEventType, ModelRequest
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
from app.providers.http_base import (
UsageTracker, check_error, decode_tool_arguments, invalid_response, list_value,
object_value, remote_error, string_value, token_count, truncated_stream,
)
from app.providers.openai_compatible import OpenAICompatibleProvider
from app.providers.tool_names import mapped_tool_names
class OpenAIResponsesProvider(OpenAICompatibleProvider):
stream_path = "/responses"
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
inputs = []
for message in request.messages:
if message.role == MessageRole.tool:
if not message.tool_call_id:
raise ProviderError("PROVIDER_INVALID_REQUEST", "Tool result requires a call identifier.")
inputs.append({"type": "function_call_output", "call_id": message.tool_call_id,
"output": message.content})
continue
if message.content or not message.tool_calls:
inputs.append({"role": message.role.value, "content": message.content})
for call in message.tool_calls:
inputs.append({"type": "function_call", "call_id": call.tool_call_id,
"name": call.name, "arguments": json.dumps(call.arguments)})
payload: dict[str, object] = {"model": request.model, "input": inputs, "stream": stream}
if request.system:
payload["instructions"] = request.system
if request.tools:
payload["tools"] = [{"type": "function", "name": tool.name,
"description": tool.description, "parameters": tool.parameters}
for tool in request.tools]
if request.temperature is not None:
payload["temperature"] = request.temperature
if request.max_tokens is not None:
payload["max_output_tokens"] = request.max_tokens
if request.response_format is not None:
format_ = dict(request.response_format)
if format_.get("type") == "json_schema":
format_ = {"type": "json_schema", **object_value(format_.get("json_schema"))}
payload["text"] = {"format": format_}
return payload
@staticmethod
def _check_response(data: dict) -> None:
check_error(data)
status = data.get("status")
if status == "incomplete":
raise ProviderError("PROVIDER_INCOMPLETE_RESPONSE", "Provider response is incomplete.")
if status == "failed":
raise remote_error(data.get("error"))
if status is not None and status != "completed":
raise invalid_response()
@mapped_tool_names
async def complete(self, request: ModelRequest) -> ProviderTurn:
data = await self._request("POST", self.stream_path, json=self._payload(request, stream=False))
self._check_response(data)
texts = []
calls = []
for raw in list_value(data.get("output")):
item = object_value(raw)
if item.get("type") == "message":
for raw_part in list_value(item.get("content")):
part = object_value(raw_part)
if part.get("type") == "output_text":
texts.append(string_value(part.get("text")))
elif part.get("type") == "refusal":
texts.append(string_value(part.get("refusal")))
elif item.get("type") == "function_call":
calls.append(ProviderToolCall(
tool_call_id=string_value(item.get("call_id"), nonempty=True),
name=string_value(item.get("name"), nonempty=True),
arguments=decode_tool_arguments(item.get("arguments")),
))
return ProviderTurn(text="".join(texts) or None, tool_calls=calls,
**UsageTracker().update(data.get("usage") or {}))
async def _events(self, request: ModelRequest):
calls: dict[int, dict] = {}
usage = UsageTracker()
def finish_call(index: int, final: object = None):
call = calls[index]
if call["ended"]:
return []
events = []
if final is not None:
arguments = string_value(final)
if not arguments.startswith(call["arguments"]):
raise invalid_response()
remainder = arguments[len(call["arguments"]):]
if remainder:
events.append((ModelEventType.tool_call_delta,
{"tool_call_id": call["id"], "arguments_delta": remainder}))
call["arguments"] = arguments
decode_tool_arguments(call["arguments"])
call["ended"] = True
events.append((ModelEventType.tool_call_end, {"tool_call_id": call["id"]}))
return events
async with aclosing(self._stream_json(self._payload(request, stream=True))) as chunks:
async for data in chunks:
kind = string_value(data.get("type"), nonempty=True)
if kind in {"response.failed", "response.incomplete"}:
response = object_value(data.get("response"))
self._check_response({**response, "status": kind.split(".")[1]})
elif kind in {"response.output_text.delta", "response.refusal.delta"}:
yield ModelEventType.text_delta, {"text": string_value(data.get("delta"))}
elif kind in {"response.reasoning_summary_text.delta", "response.reasoning_text.delta"}:
yield ModelEventType.thinking_delta, {"text": string_value(data.get("delta"))}
elif kind in {"response.output_item.added", "response.output_item.done"}:
item = object_value(data.get("item"))
if item.get("type") != "function_call":
continue
index = token_count(data.get("output_index"))
call_id = string_value(item.get("call_id"), nonempty=True)
name = string_value(item.get("name"), nonempty=True)
if index not in calls:
calls[index] = {"id": call_id, "name": name, "arguments": "", "ended": False,
"item_id": item.get("id")}
yield ModelEventType.tool_call_start, {"tool_call_id": call_id, "name": name}
elif calls[index]["id"] != call_id or calls[index]["name"] != name:
raise invalid_response()
if kind == "response.output_item.done":
for event in finish_call(index, item.get("arguments")):
yield event
elif item.get("arguments"):
arguments = string_value(item["arguments"])
calls[index]["arguments"] += arguments
yield ModelEventType.tool_call_delta, {"tool_call_id": call_id, "arguments_delta": arguments}
elif kind in {"response.function_call_arguments.delta", "response.function_call_arguments.done"}:
index = token_count(data.get("output_index"))
call = calls.get(index)
if call is None or (data.get("item_id") and call["item_id"] != data["item_id"]):
raise invalid_response()
if kind.endswith(".done"):
for event in finish_call(index, data.get("arguments")):
yield event
else:
if call["ended"]:
raise invalid_response()
fragment = string_value(data.get("delta"))
call["arguments"] += fragment
yield ModelEventType.tool_call_delta, {"tool_call_id": call["id"], "arguments_delta": fragment}
elif kind == "response.completed":
response = object_value(data.get("response"))
self._check_response(response)
if any(not call["ended"] for call in calls.values()):
raise truncated_stream()
if response.get("usage") is not None:
yield ModelEventType.usage, usage.update(response["usage"])
return
elif kind == "[DONE]":
raise truncated_stream()
elif kind in {"response.created", "response.in_progress"}:
response = object_value(data.get("response"))
check_error(response)
if response.get("usage") is not None:
yield ModelEventType.usage, usage.update(response["usage"])
raise truncated_stream()
+52 -1
View File
@@ -1,5 +1,10 @@
from dataclasses import dataclass
from time import perf_counter
from pathlib import Path
from app.config import get_settings
from app.database.db import connect
from app.errors import ApiError
from app.contracts import ModelInfo, ProviderConfig, ProviderTestResponse
from app.providers.base import ModelProvider
@@ -16,20 +21,64 @@ class RegisteredProvider:
class ProviderRegistry:
def __init__(self) -> None:
def __init__(self, factory=None) -> None:
self._providers: dict[str, RegisteredProvider] = {}
self._factory = factory
self._loaded_path: Path | None = None
def _restore(self) -> None:
if self._factory is None or self._loaded_path == get_settings().db_path:
return
conn = connect()
try:
conn.execute("CREATE TABLE IF NOT EXISTS provider_configs (provider_id TEXT PRIMARY KEY, config_json TEXT NOT NULL)")
restored = {}
for row in conn.execute("SELECT config_json FROM provider_configs"):
config = ProviderConfig.model_validate_json(row["config_json"])
if config.provider_id == "mock":
raise ValueError("reserved provider")
restored[config.provider_id] = RegisteredProvider(config, self._factory.build(config))
if "mock" in self._providers:
restored["mock"] = self._providers["mock"]
self._providers = restored
self._loaded_path = get_settings().db_path
except (ValueError, TypeError) as exc:
raise ApiError(500, "PROVIDER_STORAGE_INVALID", "Saved provider configuration could not be loaded.") from exc
finally:
conn.close()
def _save(self, config: ProviderConfig) -> None:
if self._factory is None or config.provider_id == "mock":
return
conn = connect()
try:
conn.execute("INSERT OR REPLACE INTO provider_configs VALUES (?, ?)", (config.provider_id, config.model_dump_json()))
finally:
conn.close()
def register(self, config: ProviderConfig, adapter: ModelProvider) -> None:
if config.provider_id != "mock":
self._restore()
if config.provider_id in self._providers:
raise ValueError(f"Provider already registered: {config.provider_id}")
self._save(config)
self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter)
def unregister(self, provider_id: str) -> None:
self._restore()
if self._factory is not None:
conn = connect()
try:
conn.execute("DELETE FROM provider_configs WHERE provider_id = ?", (provider_id,))
finally:
conn.close()
self._providers.pop(provider_id, None)
def replace(self, config: ProviderConfig, adapter: ModelProvider) -> None:
self._restore()
if config.provider_id not in self._providers:
raise ProviderNotFoundError(config.provider_id)
self._save(config)
self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter)
def get(self, provider_id: str) -> RegisteredProvider:
@@ -39,12 +88,14 @@ class ProviderRegistry:
return provider
def get_any(self, provider_id: str) -> RegisteredProvider:
self._restore()
try:
return self._providers[provider_id]
except KeyError as exc:
raise ProviderNotFoundError(provider_id) from exc
def list_configs(self) -> list[ProviderConfig]:
self._restore()
return [item.config.model_copy(deep=True) for item in self._providers.values()]
async def list_models(self, provider_id: str) -> list[ModelInfo]:
+287
View File
@@ -0,0 +1,287 @@
"""Capability routing: validated remote results, then an explicit local backend.
Phase E supplies HTTP adapters and injectable local contracts. Hash embeddings are
still a development placeholder; speech models are installed in phase F.
"""
from __future__ import annotations
import hashlib
import json
import math
from dataclasses import dataclass
from pathlib import Path
from typing import Protocol
import httpx
from app.contracts import (
EmbeddingResult, LocalBackendStatus, ModelBinding, ModelRoutingConfig,
ModelRoutingResponse, ProviderType, SpeakerMatchResult,
)
from app.database.db import connect, transaction
from app.errors import ApiError
from app.providers.base import ProviderError
from app.providers.credentials import CredentialResolver, CredentialStoreError
from app.providers.registry import ProviderNotFoundError, ProviderRegistry
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
CAPABILITIES = ("embedding", "transcription", "speaker_matching")
HTTP_TYPES = {ProviderType.openai_chat, ProviderType.openai_compatible}
MAX_MEDIA_BYTES = 25 * 1024 * 1024
MAX_RESPONSE_BYTES = 16 * 1024 * 1024
class LocalSpeechBackend(Protocol):
available: bool
async def transcribe(self, source: Path, language: str | None) -> str: ...
async def match(self, source: Path, reference: Path) -> float: ...
class PendingSpeechBackend:
available = False
async def transcribe(self, source: Path, language: str | None) -> str:
raise ProviderError("LOCAL_MODEL_NOT_INSTALLED", "本地音频转写模型尚未安装,将在阶段 F 接入。")
async def match(self, source: Path, reference: Path) -> float:
raise ProviderError("LOCAL_MODEL_NOT_INSTALLED", "本地声纹模型尚未安装,将在阶段 F 接入。")
@dataclass(frozen=True)
class RoutedTranscript:
text: str
source: str
fallback_reason: str | None = None
def invalid_response() -> ProviderError:
return ProviderError("PROVIDER_INVALID_RESPONSE", "Model API returned an invalid result.")
def finite_number(value: object) -> bool:
if type(value) not in (int, float):
return False
try:
return math.isfinite(value)
except (OverflowError, ValueError):
return False
class ModelRoutingService:
def __init__(self, providers: ProviderRegistry, credentials: CredentialResolver, *,
local_embedding: EmbeddingProvider | None = None,
local_speech: LocalSpeechBackend | None = None,
transport: httpx.AsyncBaseTransport | None = None) -> None:
self.providers = providers
self.credentials = credentials
self.local_embedding = local_embedding or HashEmbeddingProvider()
self.local_speech = local_speech or PendingSpeechBackend()
self.transport = transport
@staticmethod
def _connection():
conn = connect()
conn.execute("CREATE TABLE IF NOT EXISTS model_routing (id INTEGER PRIMARY KEY CHECK(id=1), config_json TEXT NOT NULL)")
return conn
def configuration(self) -> ModelRoutingConfig:
conn = self._connection()
try:
row = conn.execute("SELECT config_json FROM model_routing WHERE id=1").fetchone()
return ModelRoutingConfig.model_validate_json(row[0]) if row else ModelRoutingConfig()
except ValueError as exc:
raise ApiError(500, "MODEL_ROUTING_STORAGE_INVALID", "Saved model routing could not be loaded.") from exc
finally:
conn.close()
def describe(self) -> ModelRoutingResponse:
return ModelRoutingResponse(config=self.configuration(), local_backends=[
LocalBackendStatus(capability="embedding", status="placeholder" if isinstance(self.local_embedding, HashEmbeddingProvider) else "ready",
message="当前为 hash-v1 确定性占位向量,真实本地语义模型尚未集成。" if isinstance(self.local_embedding, HashEmbeddingProvider) else "本地 Embedding 模型已就绪。"),
*[LocalBackendStatus(capability=capability, status="ready" if self.local_speech.available else "not_installed",
message="本地模型已就绪。" if self.local_speech.available else "阶段 F 接入本地模型;当前保留回退接口。")
for capability in ("transcription", "speaker_matching")],
])
def update(self, config: ModelRoutingConfig) -> ModelRoutingResponse:
for capability in CAPABILITIES:
binding = getattr(config, capability)
if binding:
try:
provider = self.providers.get_any(binding.provider_id).config
except ProviderNotFoundError as exc:
raise ApiError(422, "PROVIDER_NOT_FOUND", "请选择已保存的提供商。") from exc
if provider.provider_type not in HTTP_TYPES:
raise ApiError(422, "MODEL_ROUTING_PROTOCOL_UNSUPPORTED", "该能力当前需要 OpenAI Compatible HTTP 接口。")
conn = self._connection()
try:
with transaction(conn):
row = conn.execute("SELECT config_json FROM model_routing WHERE id=1").fetchone()
current = ModelRoutingConfig.model_validate_json(row[0]) if row else ModelRoutingConfig()
if current.version != config.version:
raise ApiError(409, "MODEL_ROUTING_VERSION_CONFLICT", "配置已更新,请重新加载后再保存。")
saved = config.model_copy(update={"version": config.version + 1})
conn.execute("INSERT OR REPLACE INTO model_routing VALUES (1, ?)", (saved.model_dump_json(),))
finally:
conn.close()
return self.describe()
def uses_provider(self, provider_id: str) -> bool:
config = self.configuration()
return any(binding and binding.provider_id == provider_id for binding in
(getattr(config, name) for name in CAPABILITIES))
def _remote(self, binding: ModelBinding) -> tuple[str, dict[str, str]]:
try:
provider = self.providers.get(binding.provider_id).config
except ProviderNotFoundError as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Configured provider is unavailable.") from exc
if provider.provider_type not in HTTP_TYPES:
raise ProviderError("PROVIDER_CAPABILITY_UNSUPPORTED", "Provider does not support this HTTP capability.")
try:
key = self.credentials.resolve(provider.credential_id)
except CredentialStoreError as exc:
raise ProviderError("PROVIDER_CREDENTIAL_UNAVAILABLE", "Provider credential is unavailable.") from exc
if provider.credential_id and not key:
raise ProviderError("PROVIDER_CREDENTIAL_MISSING", "Provider credential is not configured.")
url = (provider.base_url or "https://api.openai.com/v1").rstrip("/") + binding.endpoint
return url, {"Authorization": f"Bearer {key}"} if key else {}
async def _request(self, binding: ModelBinding, *, remote: tuple[str, dict[str, str]] | None = None, **kwargs) -> tuple[dict, str]:
url, headers = remote or self._remote(binding)
try:
async with httpx.AsyncClient(timeout=30, transport=self.transport) as client:
async with client.stream("POST", url, headers=headers, **kwargs) as response:
response.raise_for_status()
body = bytearray()
async for chunk in response.aiter_bytes():
body.extend(chunk)
if len(body) > MAX_RESPONSE_BYTES:
raise invalid_response()
data = json.loads(body)
except httpx.TimeoutException as exc:
raise ProviderError("PROVIDER_TIMEOUT", "Model API timed out.") from exc
except httpx.HTTPStatusError as exc:
code = {401: "PROVIDER_AUTH_FAILED", 403: "PROVIDER_AUTH_FAILED", 404: "MODEL_NOT_FOUND", 429: "PROVIDER_RATE_LIMITED"}.get(exc.response.status_code, "PROVIDER_UNAVAILABLE")
raise ProviderError(code, f"Model API returned HTTP {exc.response.status_code}.") from exc
except (httpx.HTTPError, httpx.InvalidURL) as exc:
raise ProviderError("PROVIDER_UNAVAILABLE", "Model API is unavailable.") from exc
except (ValueError, UnicodeError) as exc:
raise invalid_response() from exc
if not isinstance(data, dict) or data.get("error"):
raise invalid_response()
return data, url
async def embed(self, texts: list[str]) -> EmbeddingResult:
binding = self.configuration().embedding
reason = None
if binding and texts:
try:
vectors = []
dimension = binding.dimensions
# Freeze the origin across batches, even if the user edits the provider.
remote = self._remote(binding)
for start in range(0, len(texts), 32):
batch = texts[start:start + 32]
payload = {"model": binding.model, "input": batch, "encoding_format": "float"}
if binding.dimensions is not None:
payload["dimensions"] = binding.dimensions
data, url = await self._request(binding, remote=remote, json=payload)
items = data.get("data")
if not isinstance(items, list) or len(items) != len(batch):
raise invalid_response()
indexed = {}
for item in items:
if not isinstance(item, dict):
raise invalid_response()
index, vector = item.get("index"), item.get("embedding")
if type(index) is not int or index in indexed or not 0 <= index < len(batch):
raise invalid_response()
if not isinstance(vector, list) or not 1 <= len(vector) <= 16384:
raise invalid_response()
if any(not finite_number(value) for value in vector):
raise invalid_response()
dimension = dimension or len(vector)
norm = math.hypot(*vector)
if len(vector) != dimension or not norm or not math.isfinite(norm):
raise invalid_response()
indexed[index] = [value / norm for value in vector]
vectors.extend(indexed[index] for index in range(len(batch)))
identity = json.dumps([url, binding.model, dimension], separators=(",", ":"))
return EmbeddingResult(vectors=vectors, source="api", dimensions=dimension,
model_id="api-" + hashlib.sha256(identity.encode()).hexdigest())
except ProviderError as exc:
reason = exc.code
vectors = await self.local_embedding.embed_documents(texts)
return EmbeddingResult(vectors=vectors, source="local", model_id=self.local_embedding.model_id,
dimensions=self.local_embedding.dim, fallback_reason=reason)
@staticmethod
def _media_file(path: Path):
try:
handle = path.open("rb")
except OSError as exc:
raise ApiError(404, "ATTACHMENT_NOT_FOUND", "Audio attachment was not found.") from exc
import os
if not 0 < os.fstat(handle.fileno()).st_size <= MAX_MEDIA_BYTES:
handle.close()
raise ApiError(413, "ATTACHMENT_TOO_LARGE", "Audio attachment must be between 1 byte and 25 MiB.")
return handle
async def transcribe(self, source: Path, language: str | None) -> RoutedTranscript:
binding = self.configuration().transcription
if binding is None:
with self._media_file(source):
pass
reason = None
if binding:
try:
fields = {"model": binding.model}
if language:
fields["language"] = language
with self._media_file(source) as handle:
data, _ = await self._request(binding, data=fields,
files={"file": (source.name, handle, "application/octet-stream")})
text = data.get("text")
if not isinstance(text, str) or not text.strip():
raise invalid_response()
return RoutedTranscript(text=text, source="api")
except ProviderError as exc:
reason = exc.code
try:
text = await self.local_speech.transcribe(source, language)
if not isinstance(text, str) or not text.strip():
raise ProviderError("LOCAL_MODEL_INVALID_RESPONSE", "Local transcription was empty.")
return RoutedTranscript(text=text, source="local", fallback_reason=reason)
except ProviderError as exc:
raise ApiError(503, exc.code, exc.message, {"fallback_reason": reason}) from exc
async def match_speakers(self, source: Path, reference: Path) -> SpeakerMatchResult:
binding = self.configuration().speaker_matching
if binding is None:
with self._media_file(source), self._media_file(reference):
pass
reason = None
if binding:
try:
# Explicit application contract, not an OpenAI-standard endpoint.
with self._media_file(source) as audio, self._media_file(reference) as sample:
data, _ = await self._request(binding, data={"model": binding.model}, files={
"file": (source.name, audio, "application/octet-stream"),
"reference_file": (reference.name, sample, "application/octet-stream"),
})
score = data.get("score")
if not finite_number(score) or not 0 <= score <= 1:
raise invalid_response()
return SpeakerMatchResult(score=score, source="api")
except ProviderError as exc:
reason = exc.code
try:
score = await self.local_speech.match(source, reference)
if not finite_number(score) or not 0 <= score <= 1:
raise ProviderError("LOCAL_MODEL_INVALID_RESPONSE", "Local speaker matching was invalid.")
return SpeakerMatchResult(score=score, source="local", fallback_reason=reason)
except ProviderError as exc:
raise ApiError(503, exc.code, exc.message, {"fallback_reason": reason}) from exc
+47
View File
@@ -0,0 +1,47 @@
"""Keep internal namespaced tools compatible with providers' 64-character names."""
import hashlib
import re
from functools import wraps
from app.contracts import MessageRole, ModelRequest
def prepare_tool_names(request: ModelRequest) -> tuple[ModelRequest, dict[str, str]]:
names = {tool.name for tool in request.tools}
for message in request.messages:
names.update(call.name for call in message.tool_calls)
if message.role == MessageRole.tool and message.name:
names.add(message.name)
mapping = {name: name for name in names if re.fullmatch(r"[A-Za-z0-9_-]{1,64}", name)}
used = set(mapping)
for name in sorted(names - mapping.keys()):
salt = 0
while True:
alias = "tool_" + hashlib.sha256(f"{name}:{salt}".encode()).hexdigest()[:56]
if alias not in used:
break
salt += 1
mapping[name] = alias
used.add(alias)
if all(name == alias for name, alias in mapping.items()):
return request, {}
wire = request.model_copy(deep=True)
for tool in wire.tools:
tool.name = mapping[tool.name]
for message in wire.messages:
for call in message.tool_calls:
call.name = mapping[call.name]
if message.role == MessageRole.tool and message.name:
message.name = mapping[message.name]
return wire, {alias: name for name, alias in mapping.items()}
def mapped_tool_names(complete):
@wraps(complete)
async def wrapped(self, request: ModelRequest):
wire, originals = prepare_tool_names(request)
turn = await complete(self, wire)
for call in turn.tool_calls:
call.name = originals.get(call.name, call.name)
return turn
return wrapped
+19 -3
View File
@@ -22,6 +22,7 @@ from app.repository import BlockHit
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
from app.retrieval.hybrid import normalize_scores, rrf_fuse
from app.retrieval.reranker import LexicalReranker, RankedCandidate, RerankerProvider
from app.retrieval import routed_vectors
from app.retrieval.vectorstore import SqliteVecStore, VectorStore
from app.textutils import make_snippet, match_query
@@ -39,10 +40,15 @@ class RetrievalEngine:
embedding: EmbeddingProvider,
reranker: RerankerProvider,
vector_store: VectorStore,
*,
route_embeddings: bool = False,
) -> None:
self.embedding = embedding
self.reranker = reranker
self.vector_store = vector_store
# Only the production instance opts in. Replaced test dependencies must
# remain authoritative, including monkeypatches on the singleton.
self._routed_defaults = (embedding, vector_store) if route_embeddings else None
async def search(self, request: SearchRequest) -> SearchResponse:
if request.mode == SearchMode.fts:
@@ -74,8 +80,16 @@ class RetrievalEngine:
fts_scores = {h.block_id: -h.bm25 for h in fts_hits}
if request.mode in (SearchMode.vector, SearchMode.hybrid):
query_vec = await self.embedding.embed_query(request.query)
vec_hits = await self.vector_store.search(query_vec, top_k=recall)
vec_hits = None
if (
self._routed_defaults is not None
and self.embedding is self._routed_defaults[0]
and self.vector_store is self._routed_defaults[1]
):
vec_hits = await routed_vectors.search_remote(request.query, top_k=recall)
if vec_hits is None:
query_vec = await self.embedding.embed_query(request.query)
vec_hits = await self.vector_store.search(query_vec, top_k=recall)
vec_ranked = [v.id for v in vec_hits]
vec_scores = {v.id: v.score for v in vec_hits}
@@ -217,4 +231,6 @@ def _utc(dt: datetime) -> datetime:
# 默认引擎实例:轻量实现跑通链路,后续可替换真实模型实现
engine = RetrievalEngine(HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore())
engine = RetrievalEngine(
HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore(), route_embeddings=True,
)
+192
View File
@@ -0,0 +1,192 @@
"""Optional API embeddings, isolated from the stable hash/sqlite-vec index.
The runtime's model_id is the authoritative space ID (including provider URL,
endpoint, model and dimensions); equal dimensions alone never imply compatibility.
This phase uses a lazy, rebuildable SQLite side table instead of a schema migration.
Search scans only current blocks in one database snapshot and requires complete
coverage. Cosine ranking costs O(blocks * dimensions) with an O(top_k) heap; this
small-vault implementation should become a per-space ANN index at larger scale.
"""
from __future__ import annotations
import heapq
import json
import logging
import math
import sqlite3
from dataclasses import dataclass
from typing import Protocol
from app.database.db import connect, transaction
from app.retrieval.vectorstore import VectorHit
logger = logging.getLogger(__name__)
class EmbeddingResult(Protocol):
vectors: list[list[float]]
source: str
model_id: str
dimensions: int
fallback_reason: str | None
class EmbeddingRuntime(Protocol):
async def embed(self, texts: list[str]) -> EmbeddingResult: ...
@dataclass(frozen=True)
class RemoteEmbeddings:
space_id: str
dimensions: int
vectors: list[list[float]]
def get_model_routing() -> EmbeddingRuntime | None:
"""Lazy integration hook; tests can inject a runtime without any network I/O."""
from app.container import container
return getattr(container, "model_routing", None)
def _unit_vector(vector: list[float], dimensions: int) -> list[float]:
if len(vector) != dimensions:
raise ValueError("embedding dimension mismatch")
if any(isinstance(value, bool) or not isinstance(value, (int, float)) for value in vector):
raise ValueError("embedding must be numeric")
if not all(math.isfinite(value) for value in vector):
raise ValueError("embedding must be finite")
scale = max(abs(value) for value in vector)
if scale == 0:
raise ValueError("embedding must be nonzero")
# Scaling first avoids overflow/underflow for finite but extreme API values.
scaled = [value / scale for value in vector]
norm = math.sqrt(math.fsum(value * value for value in scaled))
return [value / norm for value in scaled]
async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
"""Return validated API vectors, or None to use the caller's local baseline.
Do not use the runtime's local result: the caller may have injected its own
embedding/store pair. Exception deliberately excludes cancellation.
"""
if not texts:
return None
try:
runtime = get_model_routing()
if runtime is None:
return None
result = await runtime.embed(texts)
if result.source != "api":
return None
if not isinstance(result.model_id, str) or not result.model_id or result.model_id == "hash-v1":
raise ValueError("API embedding needs a distinct space ID")
if type(result.dimensions) is not int or result.dimensions <= 0:
raise ValueError("invalid embedding dimensions")
if len(result.vectors) != len(texts):
raise ValueError("embedding count mismatch")
return RemoteEmbeddings(
space_id=result.model_id,
dimensions=result.dimensions,
vectors=[_unit_vector(vector, result.dimensions) for vector in result.vectors],
)
except Exception as exc:
# Avoid logging provider exceptions containing credentials or note text.
logger.warning("Remote embedding unavailable (%s); using local index", type(exc).__name__)
return None
def _ensure_table(conn: sqlite3.Connection) -> None:
conn.execute("""
CREATE TABLE IF NOT EXISTS routed_block_vectors (
space_id TEXT NOT NULL,
block_id TEXT NOT NULL REFERENCES blocks(block_id) ON DELETE CASCADE,
dimensions INTEGER NOT NULL CHECK (dimensions > 0),
vector TEXT NOT NULL,
PRIMARY KEY (space_id, block_id)
)
""")
conn.execute("""
CREATE INDEX IF NOT EXISTS routed_block_vectors_block_id
ON routed_block_vectors(block_id)
""")
def store_remote(
conn: sqlite3.Connection, block_ids: list[str], batch: RemoteEmbeddings | None,
) -> None:
"""Best-effort side-index write inside the caller's metadata transaction.
A savepoint prevents partial remote batches and isolates storage failures from
note saving. Replacing/deleting blocks cascades all old spaces automatically.
"""
if batch is None:
return
try:
conn.execute("SAVEPOINT routed_vectors_write")
try:
if len(block_ids) != len(batch.vectors):
raise ValueError("block/vector count mismatch")
_ensure_table(conn)
conn.executemany(
"""INSERT INTO routed_block_vectors (space_id, block_id, dimensions, vector)
VALUES (?, ?, ?, ?)
ON CONFLICT (space_id, block_id) DO UPDATE SET
dimensions = excluded.dimensions, vector = excluded.vector""",
[
(batch.space_id, block_id, batch.dimensions, json.dumps(vector, allow_nan=False))
for block_id, vector in zip(block_ids, batch.vectors)
],
)
except BaseException:
conn.execute("ROLLBACK TO routed_vectors_write")
raise
finally:
conn.execute("RELEASE routed_vectors_write")
except Exception as exc:
logger.warning("Remote vector storage unavailable (%s); local index retained", type(exc).__name__)
async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
"""None means fallback, including any missing/invalid current-block vector.
Read coverage and vectors together so concurrent note updates cannot produce
an apparently complete subset. Never fill missing remote hits with local hits.
"""
batch = await embed_remote([query])
if batch is None:
return None
try:
conn = connect()
try:
with transaction(conn):
exists = conn.execute(
"SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'routed_block_vectors'"
).fetchone()
if exists is None:
return None
rows = conn.execute(
"""SELECT b.block_id, r.vector
FROM blocks AS b
LEFT JOIN routed_block_vectors AS r
ON r.block_id = b.block_id AND r.space_id = ? AND r.dimensions = ?
ORDER BY b.block_id""",
(batch.space_id, batch.dimensions),
)
def hits():
for row in rows:
if row["vector"] is None:
raise ValueError("remote space has incomplete block coverage")
vector = _unit_vector(json.loads(row["vector"]), batch.dimensions)
score = math.fsum(a * b for a, b in zip(batch.vectors[0], vector))
yield VectorHit(id=row["block_id"], score=max(0.0, min(1.0, score)))
return heapq.nlargest(top_k, hits(), key=lambda hit: hit.score)
finally:
conn.close()
except Exception as exc:
logger.debug("Remote vector search unavailable (%s); using local index", type(exc).__name__)
return None
+51 -10
View File
@@ -1,5 +1,6 @@
import asyncio
from collections.abc import AsyncIterator
from contextlib import aclosing
from datetime import datetime, timezone
from uuid import uuid4
@@ -33,6 +34,12 @@ from app.contracts import (
McpToolSummaryListResponse,
ModelEvent,
ModelEventType,
EmbeddingRequest,
EmbeddingResult,
ModelRoutingConfig,
ModelRoutingResponse,
SpeakerMatchRequest,
SpeakerMatchResult,
Note,
NoteCreateRequest,
NoteListResponse,
@@ -96,6 +103,7 @@ from app.services import (
transcription_service,
workspace_service,
)
from app.services.attachment_service import attachment_path
router = APIRouter(prefix="/api")
@@ -294,17 +302,22 @@ async def chat(request: ChatRequest) -> StreamingResponse:
provider = provider_or_404(request.provider_id)
async def stream() -> AsyncIterator[str]:
sequence = 0
try:
async for event in provider.adapter.stream(request):
yield as_sse(event.event.value, event.model_dump_json())
except Exception as exc:
async with aclosing(provider.adapter.stream(request)) as events:
async for event in events:
sequence = event.sequence + 1
yield as_sse(event.event.value, event.model_dump_json())
except Exception:
error = ModelEvent(
event=ModelEventType.error,
data={"code": "PROVIDER_ERROR", "message": str(exc)},
sequence=sequence,
data={"code": "PROVIDER_ERROR", "message": "Provider could not complete the request."},
timestamp=utc_now(),
)
done = ModelEvent(
event=ModelEventType.done, sequence=1, timestamp=utc_now()
event=ModelEventType.done, sequence=sequence + 1,
data={"status": "failed"}, timestamp=utc_now()
)
yield as_sse(error.event.value, error.model_dump_json())
yield as_sse(done.event.value, done.model_dump_json())
@@ -911,13 +924,13 @@ async def update_provider(
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified."
)
fields = request.model_fields_set
if ("name" in fields and request.name is None) or (
if ("provider_type" in fields and request.provider_type is None) or ("name" in fields and request.name is None) or (
"enabled" in fields and request.enabled is None
):
raise ApiError(
422,
"VALIDATION_ERROR",
"name and enabled cannot be null when explicitly provided.",
"provider_type, name and enabled cannot be null when explicitly provided.",
)
updates = {name: getattr(request, name) for name in fields}
if "credential_id" in fields:
@@ -925,7 +938,11 @@ async def update_provider(
config = ProviderConfig.model_validate(
{**current.model_dump(mode="python"), **updates}
)
adapter = container.provider_factory.build(config)
config.capabilities = container.provider_factory.capabilities(config.provider_type)
try:
adapter = container.provider_factory.build(config)
except UnsupportedProviderError as exc:
raise ApiError(422, "PROVIDER_TYPE_UNSUPPORTED", "Provider adapter is not supported.") from exc
container.providers.replace(config, adapter)
return config
@@ -941,6 +958,8 @@ async def delete_provider(provider_id: str) -> OperationResponse:
raise ApiError(
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted."
)
if container.model_routing.uses_provider(provider_id):
raise ApiError(409, "PROVIDER_IN_USE", "请先在索引与模型中解除该提供商的模型绑定。")
container.providers.unregister(provider_id)
return OperationResponse(status="completed", resource_id=provider_id)
@@ -1043,6 +1062,28 @@ async def delete_task(task_id: str) -> OperationResponse:
# Media and index
@router.get("/model-routing", response_model=ModelRoutingResponse, tags=["Providers"])
async def get_model_routing() -> ModelRoutingResponse:
return container.model_routing.describe()
@router.put("/model-routing", response_model=ModelRoutingResponse, tags=["Providers"])
async def update_model_routing(request: ModelRoutingConfig) -> ModelRoutingResponse:
return container.model_routing.update(request)
@router.post("/models/embeddings", response_model=EmbeddingResult, tags=["Providers"])
async def create_embeddings(request: EmbeddingRequest) -> EmbeddingResult:
return await container.model_routing.embed(request.texts)
@router.post("/media/speaker-matches", response_model=SpeakerMatchResult, tags=["Media"])
async def match_speakers(request: SpeakerMatchRequest) -> SpeakerMatchResult:
return await container.model_routing.match_speakers(
attachment_path(request.attachment_id), attachment_path(request.reference_attachment_id),
)
@router.post(
"/media/transcriptions",
response_model=TranscriptionJob,
@@ -1050,8 +1091,8 @@ async def delete_task(task_id: str) -> OperationResponse:
tags=["Media"],
)
async def create_transcription(request: TranscriptionRequest) -> TranscriptionJob:
return transcription_service.create_transcription(
request.attachment_id, request.language
return await transcription_service.create_transcription(
request.attachment_id, request.language, diarization=request.diarization
)
+2
View File
@@ -94,6 +94,8 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
created_at=datetime.now(timezone.utc),
))
try:
# Deleting blocks also cascades every space in routed_block_vectors;
# index_note repopulates only the currently successful API space.
repository.clear_all()
await vector_store.clear()
for rel, folder, markdown, created, updated in docs:
+7 -1
View File
@@ -16,6 +16,7 @@ from app.database.db import connect, transaction
from app.errors import ApiError
from app.knowledge.parser import ParsedNote, parse_note
from app.retrieval.embedding import HashEmbeddingProvider
from app.retrieval import routed_vectors
from app.retrieval.vectorstore import SqliteVecStore, VectorRecord
from app.services.coordination import serialized_vault_mutation
from app.services.vault_paths import (
@@ -78,7 +79,11 @@ async def index_note(parsed: ParsedNote) -> None:
半提交状态替换元数据时拿到旧 block_id清理已删除/内容变化的旧向量只为新增
block 写向量内容未变的 block 其向量仍有效无需重复写入
"""
vectors = await embedding.embed_documents([block.content for block in parsed.blocks])
texts = [block.content for block in parsed.blocks]
vectors = await embedding.embed_documents(texts)
# Network I/O stays outside the write transaction. The hash index remains
# complete even when the optional API route fails or changes vector spaces.
remote = await routed_vectors.embed_remote(texts)
conn = connect()
try:
with transaction(conn):
@@ -105,6 +110,7 @@ async def index_note(parsed: ParsedNote) -> None:
if block.block_id in missing_ids
]
await vector_store.upsert(records, conn=conn)
routed_vectors.store_remote(conn, [block.block_id for block in parsed.blocks], remote)
repository.set_index_meta(
{"embedding_model": embedding.model_id, "embedding_dim": str(embedding.dim)},
conn=conn,
+36 -14
View File
@@ -1,37 +1,59 @@
"""转写适配层;第一阶段消费文本附件或桌面 Host 预生成的旁路文本"""
"""转写作业:API 优先,本地模型回退;保留已有 Host 文本入口"""
from __future__ import annotations
from collections import OrderedDict
from datetime import datetime, timezone
from pathlib import Path
from uuid import uuid4
from app.contracts import TranscriptionJob
from app.errors import ApiError
from app.services.attachment_service import attachment_path
_jobs: OrderedDict[str, TranscriptionJob] = OrderedDict()
MAX_JOBS = 100
def create_transcription(attachment_id: str, language: str | None = None) -> TranscriptionJob:
# TODO(ai-core): 第二阶段接入本地 ASR 队列后,保留相同 Job 契约替换此同步降级实现。
del language # 预生成 transcript 暂不需要语言识别。
async def create_transcription(attachment_id: str, language: str | None = None, *, diarization: bool = False) -> TranscriptionJob:
from app.container import container
source = attachment_path(attachment_id)
transcript = source if source.suffix.lower() in {".txt", ".md"} else Path(f"{source}.txt")
job = TranscriptionJob(
job_id=f"transcription_{uuid4().hex}",
attachment_id=attachment_id,
status="completed" if transcript.is_file() else "failed",
text=transcript.read_text(encoding="utf-8") if transcript.is_file() else None,
error_code=None if transcript.is_file() else "TRANSCRIPTION_BACKEND_UNAVAILABLE",
error_message=(
None
if transcript.is_file()
else "No host-generated transcript is available; local speech models are phase two."
),
status="processing",
created_at=datetime.now(timezone.utc),
)
try:
if diarization:
# Speaker verification and diarization are different capabilities.
raise ApiError(501, "DIARIZATION_NOT_IMPLEMENTED", "说话人分离将在阶段 F 接入,当前不能忽略 diarization 请求。")
transcript = source if source.suffix.lower() in {".txt", ".md"} else attachment_path(f"{attachment_id}.txt")
# A saved transcript remains an explicit import path, never faked ASR.
if transcript.is_file() and (source == transcript or container.model_routing.configuration().transcription is None):
with transcript.open("rb") as handle:
content = handle.read(1024 * 1024 + 1)
if len(content) > 1024 * 1024:
raise ApiError(413, "TRANSCRIPT_TOO_LARGE", "Transcript exceeds 1 MiB.")
job.text = content.decode("utf-8")
if not job.text.strip():
raise ApiError(422, "TRANSCRIPT_EMPTY", "Transcript is empty.")
job.source = "sidecar"
else:
result = await container.model_routing.transcribe(source, language)
job.text = result.text
job.source = result.source
job.fallback_reason = result.fallback_reason
job.status = "completed"
except ApiError as exc:
job.status = "failed"
job.error_code = exc.code
job.error_message = exc.message
job.fallback_reason = exc.details.get("fallback_reason")
except (OSError, UnicodeError):
job.status = "failed"
job.error_code = "TRANSCRIPT_UNREADABLE"
job.error_message = "Transcript could not be read."
_jobs[job.job_id] = job
while len(_jobs) > MAX_JOBS:
_jobs.popitem(last=False)
+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())