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

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