From 1fe75e3fd24961e4301422da21ada653193006df Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Fri, 4 Sep 2026 06:19:32 +0800 Subject: [PATCH 1/2] =?UTF-8?q?feat(provider):=20=E5=AE=8C=E6=88=90?= =?UTF-8?q?=E9=98=B6=E6=AE=B5E=E5=8D=8F=E8=AE=AE=E9=80=82=E9=85=8D?= =?UTF-8?q?=E3=80=81=E5=9B=BD=E5=86=85=E9=A2=84=E8=AE=BE=E4=B8=8E=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E8=B7=AF=E7=94=B1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/agent/builtin_tools.py | 7 +- backend/app/container.py | 5 +- backend/app/contracts.py | 104 ++- backend/app/errors.py | 6 +- backend/app/providers/anthropic_messages.py | 163 +++++ backend/app/providers/factory.py | 47 +- backend/app/providers/http_base.py | 230 +++++++ backend/app/providers/ollama.py | 266 +++----- backend/app/providers/openai_compatible.py | 315 +++------ backend/app/providers/openai_responses.py | 168 +++++ backend/app/providers/registry.py | 53 +- backend/app/providers/routing.py | 287 ++++++++ backend/app/providers/tool_names.py | 47 ++ backend/app/retrieval/engine.py | 22 +- backend/app/retrieval/routed_vectors.py | 192 ++++++ backend/app/routes.py | 61 +- backend/app/services/index_service.py | 2 + backend/app/services/note_service.py | 8 +- backend/app/services/transcription_service.py | 50 +- backend/tests/test_model_routing.py | 642 ++++++++++++++++++ backend/tests/test_provider_adapters.py | 2 +- backend/tests/test_provider_protocols.py | 585 ++++++++++++++++ backend/tests/test_routed_retrieval.py | 328 +++++++++ .../AI笔记软件技术栈说明-团队版-v2.3.md | 2 +- docs/contracts/后端接口契约-开发版.md | 2 +- docs/contracts/第二阶段接口契约-开发版.md | 25 +- .../AI-Core与Agent-Core开发说明.md | 2 +- .../模型提供商与模型发现开发说明.md | 146 ++-- frontend/public/provider-icons-LICENSE.txt | 40 ++ frontend/src/assets/providers/ATTRIBUTION.txt | 19 + frontend/src/assets/providers/LICENSE | 21 + frontend/src/assets/providers/anthropic.svg | 1 + frontend/src/assets/providers/baidu.svg | 1 + frontend/src/assets/providers/deepseek.svg | 1 + frontend/src/assets/providers/hunyuan.svg | 1 + frontend/src/assets/providers/kimi.svg | 1 + frontend/src/assets/providers/minimax.svg | 1 + frontend/src/assets/providers/ollama.svg | 1 + frontend/src/assets/providers/openai.svg | 1 + frontend/src/assets/providers/qwen.svg | 1 + frontend/src/assets/providers/siliconflow.svg | 1 + frontend/src/assets/providers/stepfun.svg | 1 + frontend/src/assets/providers/volcengine.svg | 1 + frontend/src/assets/providers/zhipu.svg | 1 + frontend/src/contracts/index.ts | 36 + .../settings/ModelRoutingSettings.spec.ts | 170 +++++ .../settings/ModelRoutingSettings.vue | 160 +++++ .../features/settings/ProviderForm.spec.ts | 157 +++++ .../src/features/settings/ProviderForm.vue | 177 +++++ .../src/features/settings/ProviderLogo.vue | 24 + .../settings/ProviderPresetSelector.vue | 36 + .../src/features/settings/SettingsView.vue | 100 +-- frontend/src/services/index.ts | 1 + .../src/services/modelRoutingService.spec.ts | 28 + frontend/src/services/modelRoutingService.ts | 11 + frontend/src/services/providerService.spec.ts | 35 + frontend/src/services/providerService.ts | 5 +- 57 files changed, 4208 insertions(+), 592 deletions(-) create mode 100644 backend/app/providers/anthropic_messages.py create mode 100644 backend/app/providers/openai_responses.py create mode 100644 backend/app/providers/routing.py create mode 100644 backend/app/providers/tool_names.py create mode 100644 backend/app/retrieval/routed_vectors.py create mode 100644 backend/tests/test_model_routing.py create mode 100644 backend/tests/test_provider_protocols.py create mode 100644 backend/tests/test_routed_retrieval.py create mode 100644 frontend/public/provider-icons-LICENSE.txt create mode 100644 frontend/src/assets/providers/ATTRIBUTION.txt create mode 100644 frontend/src/assets/providers/LICENSE create mode 100644 frontend/src/assets/providers/anthropic.svg create mode 100644 frontend/src/assets/providers/baidu.svg create mode 100644 frontend/src/assets/providers/deepseek.svg create mode 100644 frontend/src/assets/providers/hunyuan.svg create mode 100644 frontend/src/assets/providers/kimi.svg create mode 100644 frontend/src/assets/providers/minimax.svg create mode 100644 frontend/src/assets/providers/ollama.svg create mode 100644 frontend/src/assets/providers/openai.svg create mode 100644 frontend/src/assets/providers/qwen.svg create mode 100644 frontend/src/assets/providers/siliconflow.svg create mode 100644 frontend/src/assets/providers/stepfun.svg create mode 100644 frontend/src/assets/providers/volcengine.svg create mode 100644 frontend/src/assets/providers/zhipu.svg create mode 100644 frontend/src/features/settings/ModelRoutingSettings.spec.ts create mode 100644 frontend/src/features/settings/ModelRoutingSettings.vue create mode 100644 frontend/src/features/settings/ProviderForm.spec.ts create mode 100644 frontend/src/features/settings/ProviderForm.vue create mode 100644 frontend/src/features/settings/ProviderLogo.vue create mode 100644 frontend/src/features/settings/ProviderPresetSelector.vue create mode 100644 frontend/src/services/modelRoutingService.spec.ts create mode 100644 frontend/src/services/modelRoutingService.ts create mode 100644 frontend/src/services/providerService.spec.ts diff --git a/backend/app/agent/builtin_tools.py b/backend/app/agent/builtin_tools.py index 449c6da..a1db632 100644 --- a/backend/app/agent/builtin_tools.py +++ b/backend/app/agent/builtin_tools.py @@ -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( diff --git a/backend/app/container.py b/backend/app/container.py index ee05cd8..2f3686d 100644 --- a/backend/app/container.py +++ b/backend/app/container.py @@ -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, diff --git a/backend/app/contracts.py b/backend/app/contracts.py index 9840f45..a95e458 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -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): diff --git a/backend/app/errors.py b/backend/app/errors.py index 8436205..d0c0f38 100644 --- a/backend/app/errors.py +++ b/backend/app/errors.py @@ -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)) diff --git a/backend/app/providers/anthropic_messages.py b/backend/app/providers/anthropic_messages.py new file mode 100644 index 0000000..fad1f42 --- /dev/null +++ b/backend/app/providers/anthropic_messages.py @@ -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() diff --git a/backend/app/providers/factory.py b/backend/app/providers/factory.py index 84f08a7..bab46c8 100644 --- a/backend/app/providers/factory.py +++ b/backend/app/providers/factory.py @@ -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], "通用 API;Coding 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, diff --git a/backend/app/providers/http_base.py b/backend/app/providers/http_base.py index 5e1a96c..db288ff 100644 --- a/backend/app/providers/http_base.py +++ b/backend/app/providers/http_base.py @@ -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 diff --git a/backend/app/providers/ollama.py b/backend/app/providers/ollama.py index 160edd8..2ecb406 100644 --- a/backend/app/providers/ollama.py +++ b/backend/app/providers/ollama.py @@ -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 diff --git a/backend/app/providers/openai_compatible.py b/backend/app/providers/openai_compatible.py index ff2c962..abdc3d5 100644 --- a/backend/app/providers/openai_compatible.py +++ b/backend/app/providers/openai_compatible.py @@ -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}.") diff --git a/backend/app/providers/openai_responses.py b/backend/app/providers/openai_responses.py new file mode 100644 index 0000000..49d0151 --- /dev/null +++ b/backend/app/providers/openai_responses.py @@ -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() diff --git a/backend/app/providers/registry.py b/backend/app/providers/registry.py index 77ce9d9..774b16f 100644 --- a/backend/app/providers/registry.py +++ b/backend/app/providers/registry.py @@ -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]: diff --git a/backend/app/providers/routing.py b/backend/app/providers/routing.py new file mode 100644 index 0000000..dd44572 --- /dev/null +++ b/backend/app/providers/routing.py @@ -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 diff --git a/backend/app/providers/tool_names.py b/backend/app/providers/tool_names.py new file mode 100644 index 0000000..7671e2f --- /dev/null +++ b/backend/app/providers/tool_names.py @@ -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 diff --git a/backend/app/retrieval/engine.py b/backend/app/retrieval/engine.py index 432edc4..5a79f5b 100644 --- a/backend/app/retrieval/engine.py +++ b/backend/app/retrieval/engine.py @@ -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, +) diff --git a/backend/app/retrieval/routed_vectors.py b/backend/app/retrieval/routed_vectors.py new file mode 100644 index 0000000..bd74bff --- /dev/null +++ b/backend/app/retrieval/routed_vectors.py @@ -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 diff --git a/backend/app/routes.py b/backend/app/routes.py index ec8908f..b0a605f 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -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 ) diff --git a/backend/app/services/index_service.py b/backend/app/services/index_service.py index 788c8ab..fd66e1d 100644 --- a/backend/app/services/index_service.py +++ b/backend/app/services/index_service.py @@ -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: diff --git a/backend/app/services/note_service.py b/backend/app/services/note_service.py index 2d7579a..28117e2 100644 --- a/backend/app/services/note_service.py +++ b/backend/app/services/note_service.py @@ -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, diff --git a/backend/app/services/transcription_service.py b/backend/app/services/transcription_service.py index d4c6f81..1e1fdda 100644 --- a/backend/app/services/transcription_service.py +++ b/backend/app/services/transcription_service.py @@ -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) diff --git a/backend/tests/test_model_routing.py b/backend/tests/test_model_routing.py new file mode 100644 index 0000000..dbcf8e1 --- /dev/null +++ b/backend/tests/test_model_routing.py @@ -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" diff --git a/backend/tests/test_provider_adapters.py b/backend/tests/test_provider_adapters.py index 781e818..94c8ab4 100644 --- a/backend/tests/test_provider_adapters.py +++ b/backend/tests/test_provider_adapters.py @@ -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 diff --git a/backend/tests/test_provider_protocols.py b/backend/tests/test_provider_protocols.py new file mode 100644 index 0000000..46aaeaa --- /dev/null +++ b/backend/tests/test_provider_protocols.py @@ -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()) diff --git a/backend/tests/test_routed_retrieval.py b/backend/tests/test_routed_retrieval.py new file mode 100644 index 0000000..685fc4e --- /dev/null +++ b/backend/tests/test_routed_retrieval.py @@ -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()) diff --git a/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md b/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md index 3b3a9a2..d86f25b 100644 --- a/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md +++ b/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md @@ -5,7 +5,7 @@ > 适用范围:桌面客户端、本地知识库、RAG、Agent、Skill、多模型接入、多模态处理与可选云同步 > 目标读者:前端、Rust 桌面端、Python AI Core、算法、测试与后续接手项目的开发成员 -> 实施状态更新:2026-09-02。本文同时包含目标架构、当前实现和第二阶段接口基线。第一阶段已完成 Vue Web 联调前端、FastAPI、Knowledge/Retrieval、Agent/Tool/Permission、Skill/Plugin 声明式运行时、Mock/OpenAI-Compatible/Ollama Provider、DeepSeek/OpenAI 预设、模型发现及开发阶段 Fernet 凭据存储。Web Workspace 已通过 FastAPI 接入后端配置的真实单 Vault;第二阶段 Agent Trace 持久化、分页快照、可恢复 SSE、stdio MCP Bridge、隔离 Plugin Host、Plugin Command 与 Plugin Settings/Secret Contract 已完成。后续继续接入真实音频处理、Provider 协议增强、Benchmark、文档导出、主题包、Trace 可视化、Mermaid 和函数图像。Tauri/Rust Host、Stronghold、原生多 Vault 文件系统和 Sync Server 仍未实现。 +> 实施状态更新:2026-09-02。本文同时包含目标架构、当前实现和第二阶段接口基线。第一阶段已完成 Vue Web 联调前端、FastAPI、Knowledge/Retrieval、Agent/Tool/Permission、Skill/Plugin 声明式运行时、Mock/OpenAI-Compatible/Ollama Provider、DeepSeek/OpenAI 预设、模型发现及开发阶段 Fernet 凭据存储。Web Workspace 已通过 FastAPI 接入后端配置的真实单 Vault;第二阶段 Agent Trace 持久化、分页快照、可恢复 SSE、stdio MCP Bridge、隔离 Plugin Host、Plugin Command 与 Plugin Settings/Secret Contract 已完成。阶段 E 已完成 Responses/Anthropic 协议、国内 logo 预设、Provider 配置恢复和 Embedding/转写/声纹 API 路由;本地语音模型仍为阶段 F 接口预留。后续继续接入真实音频处理、Benchmark、文档导出、主题包、Trace 可视化、Mermaid 和函数图像。Tauri/Rust Host、Stronghold、原生多 Vault 文件系统和 Sync Server 仍未实现。 --- diff --git a/docs/contracts/后端接口契约-开发版.md b/docs/contracts/后端接口契约-开发版.md index 6cfa1f0..60b401c 100644 --- a/docs/contracts/后端接口契约-开发版.md +++ b/docs/contracts/后端接口契约-开发版.md @@ -180,7 +180,7 @@ RunCancelled - Chat、Agent Run、Agent Events、Tool 列表、Provider 配置生命周期、模型列表和连接测试已经接入 AI Core。 - Agent Run/Event 已持久化到 SQLite;SSE 帧携带 sequence `id`,断线后可以回放缺失事件。Trace API 与 Benchmark 共用同一事件事实,并在入库前执行 Secret 脱敏和结果限长。 -- Provider Adapter 当前包含 Mock、真正增量 SSE 的 OpenAI-Compatible Chat Completions,以及 Ollama JSONL Streaming。 +- Provider Adapter 当前包含 Mock、增量 SSE 的 OpenAI-Compatible Chat Completions、OpenAI Responses、Anthropic Messages,以及 Ollama JSONL Streaming。阶段 E 增加 `/api/model-routing`、`/api/models/embeddings`、`/api/media/speaker-matches`;具体请求和阶段边界见第二阶段契约 §8.5。 - Notes、Search、Index、Skills、Plugins、Tasks 和 Provider 生命周期均已接入业务服务。 - Workspace 已接入后端配置的真实 Vault;文件树、笔记读写、文件/目录新建、重命名和删除不再使用前端 Mock Fallback。 - Note Move 保留 `note_id`;Citation 的字符偏移统一使用 UTF-16 code unit,供浏览器编辑器直接定位。 diff --git a/docs/contracts/第二阶段接口契约-开发版.md b/docs/contracts/第二阶段接口契约-开发版.md index f04a7a2..c64d061 100644 --- a/docs/contracts/第二阶段接口契约-开发版.md +++ b/docs/contracts/第二阶段接口契约-开发版.md @@ -718,6 +718,8 @@ stdio 命令始终以 executable 与 args 数组通过 `shell=False` 启动; ## 8. Provider Adapter 扩展 +> 阶段 E 实施更新(2026-09-04):OpenAI Responses、Anthropic Messages、Chat Completions 与 Ollama Adapter 已接入;国内提供商 logo 预设、独立凭据输入、配置恢复、Embedding / 转写 / 声纹 API 路由已实现。真实本地语音模型仍属于阶段 F。实现细节见 [模型提供商与模型发现开发说明](../development/模型提供商与模型发现开发说明.md)。 + 第二阶段不新增平行 Provider CRUD,继续使用第一阶段接口: ```text @@ -734,7 +736,7 @@ POST /api/chat ### 8.1 ModelInfo 扩展 -`GET /api/providers/{provider_id}/models` 的 item 增加可选字段: +以下为后续计划的可选字段;阶段 E 的 `GET /api/providers/{provider_id}/models` 实际 item 仍只包含 `model`、`display_name`、`capabilities`: ```json { @@ -774,6 +776,8 @@ Done - 浏览器取消 Fetch 或 SSE 后,服务端必须取消上游 Provider 请求。 - 不支持 reasoning 的 Provider 不发送伪造 ThinkingDelta。 +阶段 E 补充:取消或关闭迭代器直接关闭上游连接并传播取消,不向已断开的客户端继续发送 Done。内部带点号、长名称的工具映射为合法的 64 字符以内名称,响应恢复原命名空间,映射在请求内隔离。实际流中断错误码为 `PROVIDER_STREAM_TRUNCATED`;`PROVIDER_INVALID_RESPONSE` 用于无效结构/参数。上面的 `Done.data.status` 适用于真实 HTTP Adapter;开发 Mock 保留原有测试事件。 + ### 8.3 Provider 一致性测试 Contract 每个 Adapter 使用相同 Case 描述: @@ -811,6 +815,25 @@ MODEL_CONTEXT_LENGTH_EXCEEDED --- +### 8.5 阶段 E 模型路由接口(已实现) + +| 方法 | 路径 | 契约 | +| --- | --- | --- | +| GET | `/api/model-routing` | `{config, local_backends}` | +| PUT | `/api/model-routing` | 提交 ModelRoutingConfig,返回递增版本配置 | +| POST | `/api/models/embeddings` | `{texts: string[]}` → `{vectors, source, model_id, dimensions, fallback_reason}` | +| POST | `/api/media/speaker-matches` | `{attachment_id, reference_attachment_id}` → `{score, source, fallback_reason}` | + +`ModelRoutingConfig` 包含 `version`、`embedding`、`transcription`、`speaker_matching`。每个能力为 null 或 `{provider_id, model, endpoint, dimensions?}`。dimensions 仅 Embedding 使用,范围 1–16384;endpoint 是选定 Provider 下不带查询的路径,不能传第二个 URL。PUT 不提交 GET 返回的 local_backends;版本冲突返回 409 `MODEL_ROUTING_VERSION_CONFLICT`。删除仍被引用的 Provider 返回 409 `PROVIDER_IN_USE`。 + +本阶段远程能力仅接受 OpenAI Chat / Compatible HTTP 配置,默认路径分别是 `/embeddings`、`/audio/transcriptions`、`/audio/speaker-matches`。最后一个是本项目自定义 multipart 接口,**不是公共 OpenAI 标准协议**:请求 model、file、reference_file,响应有限 0–1 的 score。转写采用 multipart model、file、可选 language,必须返回非空 text。文件来自受控附件目录,限制 25 MiB。 + +`TranscriptionJob` 新增可选 `source: api|local|sidecar` 和 `fallback_reason`。保留已有转写 Job 路径;`diarization=true` 返回失败 Job,错误为 `DIARIZATION_NOT_IMPLEMENTED`,不能静默忽略。 + +无配置时调用本地接口;有配置时 API 优先,网络/鉴权/限流/结果无效时回退。本地 Embedding 当前为 hash 占位;本地 ASR / 声纹后端尚未安装时返回 `LOCAL_MODEL_NOT_INSTALLED`,而非伪成功。远程 Embedding 独立索引并检查完整覆盖,模型变化后需重建;不与本地向量混算。 + +Provider PATCH 支持 provider_type;普通配置持久化到 SQLite,凭据继续独立加密。预设新增 logo_id、description、capabilities,前端图标随应用打包。 + ## 9. RAG / Agent Benchmark Benchmark Service 同时提供 Python 调用接口和本地 HTTP 接口。CLI、测试和前端报告页调用同一 Service,不各自实现指标。 diff --git a/docs/development/AI-Core与Agent-Core开发说明.md b/docs/development/AI-Core与Agent-Core开发说明.md index dac8187..e1cf844 100644 --- a/docs/development/AI-Core与Agent-Core开发说明.md +++ b/docs/development/AI-Core与Agent-Core开发说明.md @@ -342,7 +342,7 @@ Skill Manifest 前端智能体页面已经完成中文联调:运行状态、Agent Event、内置 Tool、Permission 和常用事件详情字段均通过集中标签映射展示中文;`notes.search` 等技术 ID 继续保留,便于与后端 Trace、日志和接口契约对应。 -- 已实现 Mock、OpenAI-Compatible Chat Completions 与 Ollama Adapter;OpenAI Responses 和 Anthropic Messages 尚未实现。 +- 已实现 Mock、OpenAI-Compatible Chat Completions、Ollama、OpenAI Responses 和 Anthropic Messages Adapter;阶段 E 同时完成国内预设、持久化配置和能力模型路由,详见 [模型提供商与模型发现开发说明](模型提供商与模型发现开发说明.md)。 - Provider 配置暂存内存,后续通过 Repository 接入 SQLite;PATCH 已支持用显式 `null` 清空 base URL、默认模型和凭据引用。 - Run/Trace 已通过 Repository 接入 SQLite;后续增加按保留策略归档和 Benchmark 引用保护。 - Permission 已有核心等待/恢复机制,前端确认 UI 已完成联调和中文展示。 diff --git a/docs/development/模型提供商与模型发现开发说明.md b/docs/development/模型提供商与模型发现开发说明.md index 825e501..226471c 100644 --- a/docs/development/模型提供商与模型发现开发说明.md +++ b/docs/development/模型提供商与模型发现开发说明.md @@ -1,107 +1,105 @@ -# 模型提供商与模型发现开发说明 +# 模型提供商、协议适配与模型路由开发说明 -> 更新日期:2026-09-02。OpenAI、DeepSeek、Ollama 预设、模型自动发现、默认模型选择和开发阶段加密凭据存储均已实现并接入设置页。 +> 更新日期:2026-09-04。阶段 E 实现记录。本地小模型的实际安装与多模态队列属于阶段 F;本阶段保留并测试可注入的本地后端接口。 -## 1. 本次目标 +## 1. 设置与凭据 -本次完善设置页的模型提供商配置,不改变 Agent、Chat 和 Skill 对统一 Model Core 接口的依赖: +设置 → 模型提供商 → 新增 Provider 提供可搜索的 logo 预设网格,包含 DeepSeek、Kimi、阿里云百炼、智谱 GLM、火山方舟、硅基流动、百度千帆、腾讯混元、MiniMax、阶跃星辰,以及 OpenAI Chat / Responses、Anthropic 和 Ollama。图标打包到前端,使用时不请求第三方图片服务;来源和许可见前端 assets/providers 目录。 -- 提供 OpenAI、DeepSeek 和 Ollama 配置预设; -- 保存 Provider 后自动获取该账号或服务当前可用的模型列表; -- 支持手动刷新模型列表和选择默认模型; -- 保留自定义 OpenAI-Compatible 服务入口; -- 不在 Vue、FastAPI 配置或仓库文件中保存、回显 API Key 明文。 +预设返回 `preset_id`、`logo_id`、`name`、`provider_type`、`base_url`、`requires_credential`、`description` 和 `capabilities`。能力标签表示预设接入范围,不保证该账号的每个模型支持全部能力。厂商专用媒体协议、Coding Plan 和海外地域需要使用对应地址,不能仅凭厂商名称推断协议兼容。 -## 2. 接口与实现 +预设和自定义服务都可以直接输入 API Key。每个新配置分配独立 Credential ID,避免同厂商多账号相互覆盖。明文只留在密码输入框和专用请求中,提交、失败、切换预设及关闭时清空;密钥不进入 Pinia、localStorage、Provider 配置响应或模型路由。 -### 2.1 Provider 预设 +凭据继续使用独立的 `PUT /api/credentials/{credential_id}` 和 Fernet 开发存储。`plugin.*`、`mcp.*` 是保留命名空间。桌面端阶段仍需要把主密钥管理迁移到 Stronghold。保存密钥与保存 Provider 是两个请求,Provider 保存失败时可能留下未引用的加密凭据,可通过凭据删除接口清理。 -新增接口: +Provider 配置和 Credential ID 写入 SQLite `provider_configs`,重启后恢复。Mock 为内置 Provider,不能编辑或删除。PATCH 已支持变更 `provider_type` 并重新创建 Adapter;Base URL 限制为不带用户信息、查询或 fragment 的 HTTP(S) 地址。 -```http -GET /api/providers/presets +## 2. 协议适配 + +支持的协议是 OpenAI Chat Completions、OpenAI-Compatible、OpenAI Responses、Anthropic Messages 和 Ollama。Agent、Chat、Skill 仍只依赖内部 `ModelRequest` / `ModelEvent` / `ProviderTurn`,不直接解释厂商协议。 + +Adapter 负责消息及 Tool 历史转换、增量文本、可用的 reasoning delta、工具参数片段、usage、终止与统一错误。外部错误正文不原样返回;HTTP 鉴权、限流、超时、无效数据、流中断分别映射为内部错误。取消继续传播并关闭上游连接,不触发第二次本地推理。 + +`GET /api/providers/{provider_id}/models` 用于发现模型。模型列表不等于每个模型的能力承诺;部分厂商或代理不提供 `/models` 时,允许直接手动输入模型 ID。连接测试验证模型发现接口,不代表每一种媒体模型已完成真实推理验收。 + +## 3. 三类模型路由 + +接口: + +| 方法 | 路径 | 用途 | +| --- | --- | --- | +| GET | `/api/model-routing` | 读取配置和本地后端状态 | +| PUT | `/api/model-routing` | 带版本更新三类模型绑定 | +| POST | `/api/models/embeddings` | 文本向量,返回来源和回退原因 | +| POST | `/api/media/transcriptions` | 附件转写作业 | +| GET | `/api/media/transcriptions/{job_id}` | 获取转写作业 | +| POST | `/api/media/speaker-matches` | 两个音频附件的声纹相似度 | + +设置 → 索引与模型分别选择 Embedding、音频转文本和声纹匹配。三种绑定互相独立,可使用不同提供商、模型、密钥和 API 路径。 + +GET / PUT 响应: + +```json +{ + "config": { + "version": 1, + "embedding": { + "provider_id": "provider_example", + "model": "your-embedding-model", + "endpoint": "/embeddings", + "dimensions": null + }, + "transcription": null, + "speaker_matching": null + }, + "local_backends": [ + {"capability": "embedding", "status": "placeholder", "message": "当前为 hash-v1 占位向量"}, + {"capability": "transcription", "status": "not_installed", "message": "阶段 F 接入"}, + {"capability": "speaker_matching", "status": "not_installed", "message": "阶段 F 接入"} + ] +} ``` -预设由后端 `ProviderFactory` 提供,前端只消费名称、协议类型、Base URL 和是否需要凭据等配置元数据,不直接实现厂商协议。 +PUT body 只提交 `config` 的内容。`version` 为读取时的版本,成功递增;并发更新返回 `MODEL_ROUTING_VERSION_CONFLICT`。绑定为空表示使用本地后端。删除仍被路由引用的 Provider 返回 `PROVIDER_IN_USE`,须先解除绑定。 -当前预设: +本阶段三类远程路由使用 `openai_chat` / `openai_compatible` 的 Bearer HTTP 配置,endpoint 只能是该提供商下的路径。Responses、Anthropic 和 Ollama 原生协议不冒充上述媒体协议;Ollama 用户需要另建兼容 HTTP 配置才能用于当前远程 Embedding 接口。 -| 提供商 | Provider Type | Base URL | 默认 Credential ID | -| --- | --- | --- | --- | -| OpenAI | `openai_chat` | `https://api.openai.com/v1` | `openai` | -| DeepSeek | `openai_compatible` | `https://api.deepseek.com` | `deepseek` | -| Ollama | `ollama` | `http://127.0.0.1:11434` | 无 | +调用规则:无绑定 → 本地接口;有绑定 → API → 校验结果 → 失败或无效时调用本地接口。Provider 停用、密钥缺失、鉴权失败、限流、网络超时及无效结果均可回退;用户取消不会回退。附件不存在、大小非法等输入错误直接返回,不把用户输入错误当成模型故障。 -OpenAI 和 DeepSeek 都通过项目已有的 `OpenAICompatibleProvider` 访问。模型发现分别请求 Base URL 下的 `/models`,不引入厂商 SDK。 +## 4. Embedding 与索引一致性 -### 2.2 自动获取模型 +请求使用 `model`、`input`、`encoding_format: float`;只有明确配置维度时才发送 `dimensions`。按最多 32 条分批请求,全部批次有效才使用 API 结果。校验返回数量、连续唯一 index、维度一致性、有限数值、非零范数,并 L2 归一化。维度可为 1–16384,不截断、补零或混用不同模型的向量。 -模型列表继续使用既有接口: +返回 `vectors`、`source`、`model_id`、`dimensions`、`fallback_reason`。远程空间 ID 由完整 API URL、模型和实际维度生成;即使维度相同,不同模型的空间也不同。 -```http -GET /api/providers/{provider_id}/models -``` +笔记索引始终保留现有 hash/sqlite-vec 本地基线,远程向量写入独立 `routed_block_vectors` 表。远程查询只搜索对应空间,并要求覆盖全部当前 Block。API 失败、索引缺失、不完整或损坏时使用完整本地索引。切换模型、URL、维度后应在设置中重建全部索引。旧空间与当前文本不会混合打分,删除笔记或重建索引会通过外键清理远程向量。 -设置页在以下时机调用该接口: +当前远程侧索引采用 SQLite JSON 向量和精确余弦扫描,复杂度 O(Block 数量 × 维度),适用于当前小型 Vault;后续大规模索引需替换为按空间隔离的 ANN。网络等待发生在数据库写事务之前,当前仍会增加保存或重建延迟,异步索引队列尚未接入。 -- Provider 列表加载完成后,为所有已启用 Provider 自动刷新; -- 新增或编辑 Provider 保存成功后自动刷新; -- 用户点击“刷新模型”时手动刷新; -- 打开已有 Provider 的编辑窗口时刷新可选模型。 +无 API 时使用的 `HashEmbeddingProvider` 是确定性特征哈希占位实现,**不是已集成的小型语义模型**。真实本地 Embedding 可实现既有 `EmbeddingProvider` 接口注入。 -前端按模型名称排序并按 `model_id` 去重。获取结果保存在 `providerStore.modelsByProvider`,加载状态和错误按 Provider 隔离,单个外部服务失败不会阻止其他服务展示。 +## 5. 音频与声纹边界 -获取成功后,Provider 卡片展示模型数量和默认模型下拉框。更换默认模型会调用 Provider PATCH 接口写回配置;编辑窗口仍允许手动输入模型 ID,以兼容未出现在列表中的代理模型或部署别名。 +转写默认请求 `/audio/transcriptions`,multipart 字段 `model`、可选 `language` 和 `file`,响应必须包含非空字符串 `text`。已有纯文本附件和 Host 旁路 `.txt` 导入保留,来源标记 `sidecar`,不伪称 ASR。转写作业新增 `source`、`fallback_reason`;回退失败的作业记录 `LOCAL_MODEL_NOT_INSTALLED` 等明确错误。作业目前同步执行、限量保存在内存中,不是持久化异步队列。 -### 2.3 错误处理 +声纹匹配使用**本项目自定义 HTTP 契约**,默认 `/audio/speaker-matches`,multipart 字段 `model`、`file`、`reference_file`;响应为 `{"score": 0.85}`,score 必须为有限的 0–1 数值。公共入口只接受 `attachment_id` 和 `reference_attachment_id`,不接收任意文件路径。此接口用于一对一声纹比对,不等同于 pyannote 说话人分离,也不声称任意国内厂商原生支持该路径。 -Provider Adapter 的错误在 FastAPI 路由转换为统一 API Error: +媒体文件限制 1 字节至 25 MiB,API 响应限制 16 MiB,单次请求超时 30 秒。文件从后端受控附件目录读取,使用结束或取消时关闭句柄。 -| Provider Error | HTTP 状态 | -| --- | --- | -| `PROVIDER_AUTH_FAILED` | 401 | -| `MODEL_NOT_FOUND` | 404 | -| `PROVIDER_RATE_LIMITED` | 429 | -| `PROVIDER_TIMEOUT` | 504 | -| 其他 Provider 可用性错误 | 502 | +`LocalSpeechBackend` 提供 `transcribe` 和 `match` 接口。阶段 E 默认 `PendingSpeechBackend` 明确报告未安装;阶段 F 接入 faster-whisper、pyannote.audio 及模型资源后替换。当前 `diarization=true` 明确返回失败作业 `DIARIZATION_NOT_IMPLEMENTED`,不会静默忽略。视频解码、TTS、视频生成及厂商专用异步媒体协议不在本次交付内。 -前端在对应 Provider 卡片内展示失败原因,并允许用户修正 Credential ID、Base URL 后重新获取。 +## 6. 官方协议依据与验证 -## 3. 凭据边界 +国内通用地址核对依据:[阿里云百炼兼容接口](https://help.aliyun.com/zh/model-studio/compatibility-of-openai-with-dashscope)、[百度千帆兼容接口](https://cloud.baidu.com/doc/qianfan/s/Hmh4suq26)、[腾讯混元兼容接口](https://cloud.tencent.com/document/product/1729/111007)、[MiniMax 文本接口](https://platform.minimaxi.com/docs/guides/text-generation)、[阶跃星辰通用与套餐地址区别](https://platform.stepfun.com/docs/zh/step-plan/overview)、[火山方舟 API](https://www.volcengine.com/docs/82379/1795150)、[智谱开放接口](https://docs.bigmodel.cn/api-reference/文件-api/文件列表)。模型 ID 以账号实际开通列表为准,不写死“最新模型”。 -设置页选择 OpenAI 或 DeepSeek 预设后展示密码类型的 API Key 输入框,不再要求用户理解 Credential ID。输入值只存在于表单的临时 `ref`,不会写入 Pinia 或 localStorage;请求完成、取消表单或失败后都会清空。 +流式事件依据:[OpenAI Responses streaming](https://platform.openai.com/docs/api-reference/responses-streaming)、[Anthropic streaming](https://platform.claude.com/docs/en/build-with-claude/streaming)。音频请求依据:[SiliconFlow transcription](https://docs.siliconflow.com/en/api-reference/audio/create-audio-transcriptions)。 -API Key 通过独立接口写入: +自动化验证使用虚构凭据、本地附件、httpx.MockTransport 和可注入本地模型,覆盖流式 Tool/Usage/取消、错误映射、回退、索引空间隔离、版本冲突、重启恢复和界面凭据行为。没有使用真实 API Key 或向厂商发送推理请求。最终验证:后端全量 410 项、前端 76 项测试通过,Vue/TypeScript 类型检查和生产构建通过,浅色/深色预设页面与路由保存经过浏览器检查,git diff --check 通过。后端仅保留既有 Starlette 测试客户端弃用提示,前端保留既有大 bundle 提示。 -```http -GET /api/credentials/{credential_id} -PUT /api/credentials/{credential_id} -DELETE /api/credentials/{credential_id} -``` - -PUT 请求使用 Pydantic `SecretStr` 接收密钥,响应仅包含 Credential ID 和 `configured` 状态。后端使用 Fernet 认证加密,将密文保存到 `data/credentials/credentials.json`,主密钥保存到 `data/credentials/master.key`;目录和文件尽可能设置为仅当前用户可访问并整体排除版本控制。写入采用临时文件替换,避免进程中断留下半写文件。Provider 发起请求时按 Credential ID 解密,解密失败转换为统一 Provider Error,任何读取接口均不返回明文。 - -本地开发存储的主密钥与密文仍位于同一用户数据目录,因此它解决的是仓库泄漏、普通配置误提交和静态明文暴露,不等同于操作系统安全硬件或 Stronghold。Tauri 集成后应以 Stronghold 实现替换 `EncryptedCredentialStore`。无界面环境仍兼容 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 和 Host 注入的 `AINOTE_CREDENTIAL_`;设置页保存的本地密钥优先,环境变量仅作为回退。 - -自动化测试仅使用虚构测试值,验证磁盘文件不包含明文、加解密往返、API 响应不泄密,以及 Provider 能用解密后的值构造 Authorization Header。本次没有使用真实 OpenAI 或 DeepSeek Key,也没有向厂商发起真实请求。 - -## 4. 验证 - -后端: - -```bash +```powershell cd backend uv run pytest -q -p no:cacheprovider -``` - -前端: - -```bash -cd frontend +cd ../frontend pnpm test pnpm build ``` - -自动化验证覆盖 Provider 预设、OpenAI-Compatible `/models` 请求与鉴权头、模型映射、前端自动刷新、排序去重及按 Provider 隔离错误。生产构建同时执行 Vue 和 TypeScript 类型检查。 - -当前完整回归基线:后端 136 项测试、前端 29 项测试通过,前端类型检查和生产构建通过。Provider 配置目前仍保存在内存 Registry,AI Core 重启后需要重新创建;凭据密文会保留。`plugin.*` 为 Plugin Secret 保留命名空间,Provider 配置、临时测试凭据和通用凭据 API 均拒绝该前缀。OpenAI Responses 与 Anthropic Messages Adapter 尚未实现,设置页正式预设不会使用这两种协议。 diff --git a/frontend/public/provider-icons-LICENSE.txt b/frontend/public/provider-icons-LICENSE.txt new file mode 100644 index 0000000..3f8b3f8 --- /dev/null +++ b/frontend/public/provider-icons-LICENSE.txt @@ -0,0 +1,40 @@ +Lobe Icons — Copyright (c) 2023 LobeHub. MIT license; see LICENSE. +Source: https://github.com/lobehub/lobe-icons +Revision: 4aaf4ee1fb2678a7f989ea570f0f6ce14a9abf75 +Source directory: packages/static-svg/icons/ +Assets are bundled locally. Brand trademarks belong to their respective owners. +File mapping (local: upstream): +anthropic.svg: anthropic.svg +baidu.svg: baidu-color.svg +deepseek.svg: deepseek-color.svg +hunyuan.svg: hunyuan-color.svg +kimi.svg: kimi-color.svg +minimax.svg: minimax-color.svg +ollama.svg: ollama.svg +openai.svg: openai.svg +qwen.svg: qwen-color.svg +siliconflow.svg: siliconcloud-color.svg +stepfun.svg: stepfun-color.svg +volcengine.svg: volcengine-color.svg +zhipu.svg: zhipu-color.svg +MIT License + +Copyright (c) 2023 LobeHub + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/frontend/src/assets/providers/ATTRIBUTION.txt b/frontend/src/assets/providers/ATTRIBUTION.txt new file mode 100644 index 0000000..790c54a --- /dev/null +++ b/frontend/src/assets/providers/ATTRIBUTION.txt @@ -0,0 +1,19 @@ +Lobe Icons — Copyright (c) 2023 LobeHub. MIT license; see LICENSE. +Source: https://github.com/lobehub/lobe-icons +Revision: 4aaf4ee1fb2678a7f989ea570f0f6ce14a9abf75 +Source directory: packages/static-svg/icons/ +Assets are bundled locally. Brand trademarks belong to their respective owners. +File mapping (local: upstream): +anthropic.svg: anthropic.svg +baidu.svg: baidu-color.svg +deepseek.svg: deepseek-color.svg +hunyuan.svg: hunyuan-color.svg +kimi.svg: kimi-color.svg +minimax.svg: minimax-color.svg +ollama.svg: ollama.svg +openai.svg: openai.svg +qwen.svg: qwen-color.svg +siliconflow.svg: siliconcloud-color.svg +stepfun.svg: stepfun-color.svg +volcengine.svg: volcengine-color.svg +zhipu.svg: zhipu-color.svg diff --git a/frontend/src/assets/providers/LICENSE b/frontend/src/assets/providers/LICENSE new file mode 100644 index 0000000..1dd53d2 --- /dev/null +++ b/frontend/src/assets/providers/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2023 LobeHub + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/frontend/src/assets/providers/anthropic.svg b/frontend/src/assets/providers/anthropic.svg new file mode 100644 index 0000000..5b81844 --- /dev/null +++ b/frontend/src/assets/providers/anthropic.svg @@ -0,0 +1 @@ +Anthropic \ No newline at end of file diff --git a/frontend/src/assets/providers/baidu.svg b/frontend/src/assets/providers/baidu.svg new file mode 100644 index 0000000..ead7f89 --- /dev/null +++ b/frontend/src/assets/providers/baidu.svg @@ -0,0 +1 @@ +Baidu \ No newline at end of file diff --git a/frontend/src/assets/providers/deepseek.svg b/frontend/src/assets/providers/deepseek.svg new file mode 100644 index 0000000..3fc2302 --- /dev/null +++ b/frontend/src/assets/providers/deepseek.svg @@ -0,0 +1 @@ +DeepSeek \ No newline at end of file diff --git a/frontend/src/assets/providers/hunyuan.svg b/frontend/src/assets/providers/hunyuan.svg new file mode 100644 index 0000000..42edd6c --- /dev/null +++ b/frontend/src/assets/providers/hunyuan.svg @@ -0,0 +1 @@ +Hunyuan \ No newline at end of file diff --git a/frontend/src/assets/providers/kimi.svg b/frontend/src/assets/providers/kimi.svg new file mode 100644 index 0000000..83878fa --- /dev/null +++ b/frontend/src/assets/providers/kimi.svg @@ -0,0 +1 @@ +Kimi \ No newline at end of file diff --git a/frontend/src/assets/providers/minimax.svg b/frontend/src/assets/providers/minimax.svg new file mode 100644 index 0000000..beb7adb --- /dev/null +++ b/frontend/src/assets/providers/minimax.svg @@ -0,0 +1 @@ +Minimax \ No newline at end of file diff --git a/frontend/src/assets/providers/ollama.svg b/frontend/src/assets/providers/ollama.svg new file mode 100644 index 0000000..cc887e3 --- /dev/null +++ b/frontend/src/assets/providers/ollama.svg @@ -0,0 +1 @@ +Ollama \ No newline at end of file diff --git a/frontend/src/assets/providers/openai.svg b/frontend/src/assets/providers/openai.svg new file mode 100644 index 0000000..78caf4f --- /dev/null +++ b/frontend/src/assets/providers/openai.svg @@ -0,0 +1 @@ +OpenAI \ No newline at end of file diff --git a/frontend/src/assets/providers/qwen.svg b/frontend/src/assets/providers/qwen.svg new file mode 100644 index 0000000..f2d0ada --- /dev/null +++ b/frontend/src/assets/providers/qwen.svg @@ -0,0 +1 @@ +Qwen \ No newline at end of file diff --git a/frontend/src/assets/providers/siliconflow.svg b/frontend/src/assets/providers/siliconflow.svg new file mode 100644 index 0000000..6b5f6d8 --- /dev/null +++ b/frontend/src/assets/providers/siliconflow.svg @@ -0,0 +1 @@ +SiliconCloud \ No newline at end of file diff --git a/frontend/src/assets/providers/stepfun.svg b/frontend/src/assets/providers/stepfun.svg new file mode 100644 index 0000000..920e8a6 --- /dev/null +++ b/frontend/src/assets/providers/stepfun.svg @@ -0,0 +1 @@ +Stepfun \ No newline at end of file diff --git a/frontend/src/assets/providers/volcengine.svg b/frontend/src/assets/providers/volcengine.svg new file mode 100644 index 0000000..ecf6d75 --- /dev/null +++ b/frontend/src/assets/providers/volcengine.svg @@ -0,0 +1 @@ +Volcengine \ No newline at end of file diff --git a/frontend/src/assets/providers/zhipu.svg b/frontend/src/assets/providers/zhipu.svg new file mode 100644 index 0000000..0c6e61c --- /dev/null +++ b/frontend/src/assets/providers/zhipu.svg @@ -0,0 +1 @@ +Zhipu \ No newline at end of file diff --git a/frontend/src/contracts/index.ts b/frontend/src/contracts/index.ts index 07f6679..8449a6e 100644 --- a/frontend/src/contracts/index.ts +++ b/frontend/src/contracts/index.ts @@ -416,6 +416,39 @@ export interface ProviderPreset { base_url: string default_credential_id?: string | null requires_credential: boolean + logo_id?: string + description?: string + capabilities?: string[] +} + +export type ProviderUpdateRequest = Partial> & { + credential_id?: string | null + base_url?: string | null +} + +export type RoutingCapability = 'embedding' | 'transcription' | 'speaker_matching' + +export interface ModelBinding { + provider_id: string + model: string + endpoint: string + dimensions?: number | null +} + +export interface ModelRoutingConfig { + version: number + embedding: ModelBinding | null + transcription: ModelBinding | null + speaker_matching: ModelBinding | null +} + +export interface ModelRoutingResponse { + config: ModelRoutingConfig + local_backends: Array<{ + capability: RoutingCapability + status: 'placeholder' | 'not_installed' | 'ready' + message: string + }> } // ============ Tasks ============ @@ -721,6 +754,9 @@ export interface ApiProviderPreset { base_url: string default_credential_id?: string | null requires_credential: boolean + logo_id?: string + description?: string + capabilities?: string[] } export interface ApiModelInfo { diff --git a/frontend/src/features/settings/ModelRoutingSettings.spec.ts b/frontend/src/features/settings/ModelRoutingSettings.spec.ts new file mode 100644 index 0000000..56d695a --- /dev/null +++ b/frontend/src/features/settings/ModelRoutingSettings.spec.ts @@ -0,0 +1,170 @@ +// @vitest-environment happy-dom +import { flushPromises, mount } from '@vue/test-utils' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import type { ModelRoutingResponse, ProviderConfig } from '@/contracts' +import { ApiErrorClass } from '@/services/apiClient' +import * as service from '@/services/modelRoutingService' +import { listProviders } from '@/services/providerService' +import ModelRoutingSettings from './ModelRoutingSettings.vue' + +vi.mock('@/services/modelRoutingService', () => ({ getModelRouting: vi.fn(), saveModelRouting: vi.fn() })) +vi.mock('@/services/providerService', () => ({ listProviders: vi.fn() })) +const providers: ProviderConfig[] = [ + { provider_id: 'p1', provider_type: 'openai_compatible', name: 'Custom API', enabled: true, default_model: 'chat-model', capabilities: {}, has_credential: true }, + { provider_id: 'p2', provider_type: 'openai_chat', name: 'OpenAI', enabled: true, default_model: '', capabilities: {}, has_credential: true }, + { provider_id: 'responses', provider_type: 'openai_responses', name: 'Responses', enabled: true, default_model: '', capabilities: {}, has_credential: true }, + { provider_id: 'anthropic', provider_type: 'anthropic_messages', name: 'Anthropic', enabled: true, default_model: '', capabilities: {}, has_credential: true }, + { provider_id: 'ollama', provider_type: 'ollama', name: 'Ollama', enabled: true, default_model: '', capabilities: {}, has_credential: false }, + { provider_id: 'disabled', provider_type: 'openai_chat', name: 'Disabled', enabled: false, default_model: '', capabilities: {}, has_credential: true }, +] +const initial: ModelRoutingResponse = { config: { version: 3, embedding: null, transcription: null, speaker_matching: null }, local_backends: [ + { capability: 'embedding', status: 'placeholder', message: 'hash fallback' }, + { capability: 'transcription', status: 'not_installed', message: 'ASR not installed' }, + { capability: 'speaker_matching', status: 'not_installed', message: 'speaker not installed' }, +] } +const wrappers: ReturnType[] = [] +async function render() { + const wrapper = mount(ModelRoutingSettings) + wrappers.push(wrapper) + await flushPromises() + return wrapper +} +beforeEach(() => { + vi.resetAllMocks() + vi.mocked(service.getModelRouting).mockResolvedValue(structuredClone(initial)) + vi.mocked(listProviders).mockResolvedValue(providers) + vi.mocked(service.saveModelRouting).mockImplementation(async config => ({ ...initial, config: { ...config, version: config.version + 1 } })) +}) +afterEach(() => { wrappers.splice(0).forEach(wrapper => wrapper.unmount()) }) + +describe('ModelRoutingSettings', () => { + it('loads local selections honestly, explains index rebuilds, and disables incompatible providers', async () => { + const wrapper = await render() + expect(wrapper.findAll('select').map(select => (select.element as HTMLSelectElement).value)).toEqual(['', '', '']) + expect(wrapper.text()).toContain('当前为占位实现') + expect(wrapper.text()).toContain('真实本地 ASR 尚未接入') + expect(wrapper.text()).toContain('真实本地说话人匹配尚未接入') + expect(wrapper.text()).toContain('重建全部') + expect(wrapper.text()).toContain('重建完成前继续使用本地检索') + expect(wrapper.text()).toContain('不是 OpenAI 标准接口') + for (const id of ['responses', 'anthropic', 'ollama', 'disabled']) expect(wrapper.get(`option[value="${id}"]`).attributes()).toHaveProperty('disabled') + expect(wrapper.get('option[value="p1"]').attributes()).not.toHaveProperty('disabled') + }) + + it('saves all three independent bindings with expected version and remote dimensions', async () => { + const wrapper = await render() + for (const capability of ['embedding', 'transcription', 'speaker_matching']) { + const card = wrapper.get(`[data-capability="${capability}"]`) + await card.get('select').setValue('p1') + await card.get('[data-field="model"]').setValue(`${capability}-model`) + } + await wrapper.get('[data-field="dimensions"]').setValue('3072') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.saveModelRouting).toHaveBeenCalledWith({ version: 3, + embedding: { provider_id: 'p1', model: 'embedding-model', endpoint: '/embeddings', dimensions: 3072 }, + transcription: { provider_id: 'p1', model: 'transcription-model', endpoint: '/audio/transcriptions' }, + speaker_matching: { provider_id: 'p1', model: 'speaker_matching-model', endpoint: '/audio/speaker-matches' }, + }) + expect(wrapper.text()).toContain('配置版本 4') + expect(wrapper.text()).toContain('模型路由已保存') + await wrapper.get('[data-capability="embedding"] select').setValue('') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.saveModelRouting).toHaveBeenLastCalledWith(expect.objectContaining({ version: 4, embedding: null })) + }) + + it('supports omitted dimensions and clears stale models/endpoints when switching providers', async () => { + const wrapper = await render() + const card = wrapper.get('[data-capability="embedding"]') + await card.get('select').setValue('p1') + await card.get('[data-field="model"]').setValue('custom-embedding') + await card.get('[data-field="endpoint"]').setValue('/custom/embeddings') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.saveModelRouting).toHaveBeenCalledWith(expect.objectContaining({ embedding: { provider_id: 'p1', model: 'custom-embedding', endpoint: '/custom/embeddings', dimensions: null } })) + await card.get('select').setValue('p2') + expect((card.get('[data-field="model"]').element as HTMLInputElement).value).toBe('') + expect((card.get('[data-field="endpoint"]').element as HTMLInputElement).value).toBe('/embeddings') + }) + + it('shows a loading state and does not offer a default local configuration after load failure', async () => { + let fail!: (error: Error) => void + vi.mocked(service.getModelRouting).mockReturnValueOnce(new Promise((_, reject) => { fail = reject })) + const wrapper = await render() + expect(wrapper.text()).toContain('正在加载模型路由') + expect(wrapper.find('form').exists()).toBe(false) + fail(new Error('offline')) + await flushPromises() + expect(wrapper.text()).toContain('加载失败:offline') + expect(wrapper.find('form').exists()).toBe(false) + await wrapper.get('button').trigger('click') + await flushPromises() + expect(wrapper.find('form').exists()).toBe(true) + }) + + it('keeps unsaved input on save failure and retries without inventing a new version', async () => { + const wrapper = await render() + const card = wrapper.get('[data-capability="transcription"]') + await card.get('select').setValue('p1') + await card.get('[data-field="model"]').setValue('asr-model') + vi.mocked(service.saveModelRouting).mockRejectedValueOnce(new Error('disk full')) + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(wrapper.text()).toContain('保存失败:disk full') + expect((card.get('[data-field="model"]').element as HTMLInputElement).value).toBe('asr-model') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.saveModelRouting).toHaveBeenLastCalledWith(expect.objectContaining({ version: 3 })) + }) + + it('blocks overwrite after a conflict until explicitly reloading the latest configuration', async () => { + const wrapper = await render() + vi.mocked(service.saveModelRouting).mockRejectedValueOnce(new ApiErrorClass('MODEL_ROUTING_VERSION_CONFLICT', 'stale')) + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(wrapper.text()).toContain('配置版本冲突') + expect(wrapper.get('button[type="submit"]').attributes()).toHaveProperty('disabled') + await wrapper.get('form').trigger('submit') + expect(service.saveModelRouting).toHaveBeenCalledTimes(1) + vi.mocked(service.getModelRouting).mockResolvedValueOnce({ ...initial, config: { ...initial.config, version: 8 } }) + await wrapper.findAll('button').find(button => button.text().includes('放弃当前输入'))!.trigger('click') + await flushPromises() + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.saveModelRouting).toHaveBeenLastCalledWith(expect.objectContaining({ version: 8 })) + }) + + it('prevents invalid dimensions, endpoints, and missing provider bindings from being saved', async () => { + vi.mocked(service.getModelRouting).mockResolvedValueOnce({ ...initial, config: { ...initial.config, embedding: { provider_id: 'missing', model: 'old-model', endpoint: '/embeddings' } } }) + const wrapper = await render() + expect(wrapper.text()).toContain('原提供商已不可用') + await wrapper.get('form').trigger('submit') + expect(service.saveModelRouting).not.toHaveBeenCalled() + const card = wrapper.get('[data-capability="embedding"]') + await card.get('select').setValue('p1') + await card.get('[data-field="model"]').setValue('embedding-model') + await card.get('[data-field="dimensions"]').setValue('1.5') + await wrapper.get('form').trigger('submit') + expect(wrapper.text()).toContain('嵌入维度必须为 1–16384 的整数') + await card.get('[data-field="dimensions"]').setValue('16385') + await wrapper.get('form').trigger('submit') + expect(service.saveModelRouting).not.toHaveBeenCalled() + expect(card.get('[data-field="dimensions"]').attributes('max')).toBe('16384') + await card.get('[data-field="dimensions"]').setValue('') + await card.get('[data-field="endpoint"]').setValue('https://example.test/embeddings') + await wrapper.get('form').trigger('submit') + expect(wrapper.text()).toContain('Endpoint 必须是以 / 开头的相对路径') + expect(service.saveModelRouting).not.toHaveBeenCalled() + }) + + it('labels injected ready local backends accurately', async () => { + vi.mocked(service.getModelRouting).mockResolvedValueOnce({ ...initial, local_backends: [{ capability: 'transcription', status: 'ready', message: 'Local ASR ready' }] }) + const wrapper = await render() + const card = wrapper.get('[data-capability="transcription"]') + expect(card.get('option[value=""]').text()).toBe('本地 · 已就绪') + expect(card.text()).toContain('本地后端已就绪') + expect(card.text()).toContain('Local ASR ready') + expect(card.text()).not.toContain('真实本地 ASR 尚未接入') + }) +}) diff --git a/frontend/src/features/settings/ModelRoutingSettings.vue b/frontend/src/features/settings/ModelRoutingSettings.vue new file mode 100644 index 0000000..492e105 --- /dev/null +++ b/frontend/src/features/settings/ModelRoutingSettings.vue @@ -0,0 +1,160 @@ + + + + + diff --git a/frontend/src/features/settings/ProviderForm.spec.ts b/frontend/src/features/settings/ProviderForm.spec.ts new file mode 100644 index 0000000..c1a33bb --- /dev/null +++ b/frontend/src/features/settings/ProviderForm.spec.ts @@ -0,0 +1,157 @@ +// @vitest-environment happy-dom +import { flushPromises, mount } from '@vue/test-utils' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import type { ProviderConfig, ProviderPreset } from '@/contracts' +import * as service from '@/services/providerService' +import ProviderForm from './ProviderForm.vue' +import ProviderPresetSelector from './ProviderPresetSelector.vue' + +vi.mock('@/services/providerService', () => ({ listProviderPresets: vi.fn(), getCredentialStatus: vi.fn(), putCredential: vi.fn(), createProvider: vi.fn(), updateProvider: vi.fn() })) +const presets: ProviderPreset[] = [ + { preset_id: 'deepseek', name: 'DeepSeek', provider_type: 'openai_compatible', base_url: 'https://deepseek.example.test', default_credential_id: 'shared-deepseek', requires_credential: true, logo_id: 'deepseek' }, + { preset_id: 'qwen', name: '通义千问', provider_type: 'openai_compatible', base_url: 'https://qwen.example.test', default_credential_id: 'shared-qwen', requires_credential: true, logo_id: 'qwen' }, +] +const existing: ProviderConfig = { provider_id: 'p1', provider_type: 'openai_compatible', name: 'DeepSeek', base_url: presets[0].base_url, default_model: 'old-model', enabled: true, credential_id: 'old-shared-key', has_credential: true, capabilities: {} } +const wrappers: ReturnType[] = [] +async function render(provider?: ProviderConfig) { + const wrapper = mount(ProviderForm, { props: { provider, models: [{ model_id: 'old-model', name: 'Old', capabilities: {} }] } }) + wrappers.push(wrapper) + await flushPromises() + return wrapper +} +beforeEach(() => { + vi.resetAllMocks() + vi.mocked(service.listProviderPresets).mockResolvedValue(presets) + vi.mocked(service.getCredentialStatus).mockResolvedValue(true) + vi.mocked(service.putCredential).mockResolvedValue() + vi.mocked(service.createProvider).mockResolvedValue(existing) + vi.mocked(service.updateProvider).mockResolvedValue(existing) +}) +afterEach(() => { wrappers.splice(0).forEach(wrapper => wrapper.unmount()) }) + +describe('ProviderForm', () => { + it('filters compact preset chips and resolves bundled logos', async () => { + const wrapper = await render() + await wrapper.get('#provider-search').setValue('通义') + expect(wrapper.find('[data-preset="deepseek"]').exists()).toBe(false) + expect(wrapper.find('[data-preset="qwen"]').exists()).toBe(true) + expect(wrapper.get('[data-preset="qwen"] img').attributes('src')).not.toMatch(/^https?:/) + }) + + it('clears the secret, old model and credential on preset and custom selection', async () => { + const wrapper = await render(existing) + await wrapper.get('input[type="password"]').setValue('draft-secret') + await wrapper.get('[data-preset="qwen"]').trigger('click') + expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('') + expect((wrapper.get('[data-field="model"]').element as HTMLInputElement).value).toBe('') + expect(wrapper.findAll('datalist option')).toHaveLength(0) + await wrapper.get('form').trigger('submit') + expect(service.updateProvider).not.toHaveBeenCalled() + expect(wrapper.text()).toContain('请输入 API Key') + await wrapper.get('input[type="password"]').setValue('new-secret') + wrapper.getComponent(ProviderPresetSelector).vm.$emit('update:modelValue', '') + await flushPromises() + expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('') + }) + + it('allocates different credential IDs for two new providers using the same preset', async () => { + for (let i = 0; i < 2; i++) { + const wrapper = await render() + await wrapper.get('[data-preset="deepseek"]').trigger('click') + await wrapper.get('input[type="password"]').setValue(`test-key-${i}`) + await wrapper.get('form').trigger('submit') + await flushPromises() + } + const ids = vi.mocked(service.putCredential).mock.calls.map(call => call[0]) + expect(ids).toHaveLength(2) + expect(new Set(ids).size).toBe(2) + ids.forEach(id => expect(id).toMatch(/^provider-key-[0-9a-f-]{36}$/)) + vi.mocked(service.createProvider).mock.calls.forEach(([data], index) => { + expect(data.credential_id).toBe(ids[index]) + expect(JSON.stringify(data)).not.toContain('test-key') + expect(JSON.stringify(data)).not.toContain('shared-deepseek') + }) + }) + + it('accepts a custom API key and clears it after a failed credential save', async () => { + const wrapper = await render() + await wrapper.get('[data-field="name"]').setValue('Custom') + await wrapper.get('[data-field="base-url"]').setValue('https://custom.example.test/v1') + await wrapper.get('input[type="password"]').setValue('custom-test-key') + vi.mocked(service.putCredential).mockRejectedValueOnce(new Error('credential store unavailable')) + await wrapper.get('form').trigger('submit') + await flushPromises() + expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('') + expect(wrapper.text()).toContain('credential store unavailable') + expect(service.createProvider).not.toHaveBeenCalled() + await wrapper.get('input[type="password"]').setValue('custom-test-key') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.createProvider).toHaveBeenCalledWith(expect.objectContaining({ name: 'Custom', credential_id: expect.stringMatching(/^provider-key-/) })) + }) + + it('preserves its own existing key when untouched and rotates shared legacy references when replacing a key', async () => { + const untouched = await render(existing) + await untouched.get('form').trigger('submit') + await flushPromises() + expect(service.updateProvider).toHaveBeenLastCalledWith('p1', expect.objectContaining({ credential_id: 'old-shared-key', default_model: 'old-model' })) + const rotated = await render(existing) + await rotated.get('input[type="password"]').setValue('replacement-test-key') + await rotated.get('form').trigger('submit') + await flushPromises() + expect(service.putCredential).toHaveBeenCalledWith(expect.stringMatching(/^provider-key-/), 'replacement-test-key') + expect(service.updateProvider).toHaveBeenLastCalledWith('p1', expect.objectContaining({ credential_id: vi.mocked(service.putCredential).mock.calls[0][0] })) + }) + + it('persists edited protocols and unlinks the previous credential and model', async () => { + const wrapper = await render(existing) + await wrapper.get('[data-field="protocol"]').setValue('openai_responses') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.updateProvider).toHaveBeenCalledWith('p1', expect.objectContaining({ provider_type: 'openai_responses', default_model: '', credential_id: null })) + }) + + it('does not let a late credential status reuse a key after switching presets', async () => { + let resolveStatus!: (configured: boolean) => void + vi.mocked(service.getCredentialStatus).mockReturnValue(new Promise(resolve => { resolveStatus = resolve })) + const wrapper = await render(existing) + await wrapper.get('[data-preset="qwen"]').trigger('click') + resolveStatus(true) + await flushPromises() + await wrapper.get('form').trigger('submit') + expect(service.updateProvider).not.toHaveBeenCalled() + expect(wrapper.text()).toContain('请输入 API Key') + }) + + it('clears secrets on close and unmount, and stops a pending credential save from creating a provider', async () => { + let finish!: () => void + vi.mocked(service.putCredential).mockReturnValue(new Promise(resolve => { finish = resolve })) + const wrapper = await render() + await wrapper.get('[data-preset="deepseek"]').trigger('click') + await wrapper.get('input[type="password"]').setValue('pending-test-key') + await wrapper.get('form').trigger('submit') + await wrapper.get('[aria-label="关闭提供商表单"]').trigger('click') + expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('') + wrapper.unmount() + finish() + await flushPromises() + expect(service.createProvider).not.toHaveBeenCalled() + const reopened = await render() + expect((reopened.get('input[type="password"]').element as HTMLInputElement).value).toBe('') + }) + + it('clears a secret after provider save failure and keeps only the successfully saved reference for retry', async () => { + const wrapper = await render() + await wrapper.get('[data-preset="deepseek"]').trigger('click') + await wrapper.get('input[type="password"]').setValue('retry-test-key') + vi.mocked(service.createProvider).mockRejectedValueOnce(new Error('provider save failed')) + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(wrapper.text()).toContain('provider save failed') + expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(service.putCredential).toHaveBeenCalledTimes(1) + expect(service.createProvider).toHaveBeenCalledTimes(2) + }) +}) diff --git a/frontend/src/features/settings/ProviderForm.vue b/frontend/src/features/settings/ProviderForm.vue new file mode 100644 index 0000000..eed4f28 --- /dev/null +++ b/frontend/src/features/settings/ProviderForm.vue @@ -0,0 +1,177 @@ + + + + + diff --git a/frontend/src/features/settings/ProviderLogo.vue b/frontend/src/features/settings/ProviderLogo.vue new file mode 100644 index 0000000..ac59061 --- /dev/null +++ b/frontend/src/features/settings/ProviderLogo.vue @@ -0,0 +1,24 @@ + + + + + diff --git a/frontend/src/features/settings/ProviderPresetSelector.vue b/frontend/src/features/settings/ProviderPresetSelector.vue new file mode 100644 index 0000000..95a361f --- /dev/null +++ b/frontend/src/features/settings/ProviderPresetSelector.vue @@ -0,0 +1,36 @@ + + + + + diff --git a/frontend/src/features/settings/SettingsView.vue b/frontend/src/features/settings/SettingsView.vue index cb6a90f..a9dd27f 100644 --- a/frontend/src/features/settings/SettingsView.vue +++ b/frontend/src/features/settings/SettingsView.vue @@ -1,6 +1,9 @@