From 1741f7b1aa2f6ae5e9f4bfaaeba4ffc7ef028e56 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Thu, 27 Aug 2026 14:16:57 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0provider=E5=B7=A5=E5=8E=82?= =?UTF-8?q?=E5=92=8COllama=E6=94=AF=E6=8C=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 实现ProviderFactory用于构建不同类型的provider适配器 - 添加EnvironmentCredentialResolver用于解析环境变量中的凭证 - 实现OllamaProvider支持本地模型调用 - 实现OpenAICompatibleProvider支持OpenAI兼容接口 - 在AgentRuntime中添加对ProviderError的处理 - 更新Message结构体添加tool_calls字段 - 实现provider配置的增删改查API端点 - 添加provider注册表的replace方法 - 添加HTTP基础类和工具参数解码功能 - 更新依赖添加httpx库 - 添加相关单元测试验证provider适配器功能 ``` --- backend/app/agent/runtime.py | 19 ++- backend/app/container.py | 6 +- backend/app/contracts.py | 1 + backend/app/providers/__init__.py | 6 +- backend/app/providers/base.py | 7 + backend/app/providers/credentials.py | 17 +++ backend/app/providers/factory.py | 48 +++++++ backend/app/providers/http_base.py | 82 +++++++++++ backend/app/providers/ollama.py | 117 ++++++++++++++++ backend/app/providers/openai_compatible.py | 156 +++++++++++++++++++++ backend/app/providers/registry.py | 16 ++- backend/app/routes.py | 69 ++++++++- backend/pyproject.toml | 1 + backend/tests/test_api.py | 24 ++++ backend/tests/test_provider_adapters.py | 154 ++++++++++++++++++++ backend/uv.lock | 39 ++++++ docs/AI-Core与Agent-Core开发说明.md | 49 ++++++- docs/后端接口契约-开发版.md | 3 +- 18 files changed, 794 insertions(+), 20 deletions(-) create mode 100644 backend/app/providers/credentials.py create mode 100644 backend/app/providers/factory.py create mode 100644 backend/app/providers/http_base.py create mode 100644 backend/app/providers/ollama.py create mode 100644 backend/app/providers/openai_compatible.py create mode 100644 backend/tests/test_provider_adapters.py diff --git a/backend/app/agent/runtime.py b/backend/app/agent/runtime.py index 08bf61c..4fe122a 100644 --- a/backend/app/agent/runtime.py +++ b/backend/app/agent/runtime.py @@ -20,6 +20,7 @@ from app.contracts import ( ToolResult, ) from app.providers.registry import ProviderRegistry +from app.providers.base import ProviderError class AgentRunNotFoundError(LookupError): @@ -141,6 +142,8 @@ class AgentRuntime: self._finish_cancelled(record) except TimeoutError: self._fail(record, "AGENT_TIMEOUT", "Agent run exceeded its timeout.") + except ProviderError as exc: + self._fail(record, exc.code, exc.message) except Exception as exc: self._fail(record, "AGENT_FAILED", str(exc)) @@ -183,12 +186,18 @@ class AgentRuntime: return if turn.tool_calls: - for provider_call in turn.tool_calls: - call = ToolCall( - tool_call_id=provider_call.tool_call_id, - name=provider_call.name, - arguments=provider_call.arguments, + calls = [ + ToolCall( + tool_call_id=item.tool_call_id, + name=item.name, + arguments=item.arguments, ) + for item in turn.tool_calls + ] + messages.append( + Message(role=MessageRole.assistant, content=turn.text or "", tool_calls=calls) + ) + for call in calls: result = await self._execute_tool(record, call) record.run.tool_results.append(result) messages.append( diff --git a/backend/app/container.py b/backend/app/container.py index 3377e2d..cdf05aa 100644 --- a/backend/app/container.py +++ b/backend/app/container.py @@ -3,18 +3,21 @@ from dataclasses import dataclass from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry from app.agent.builtin_tools import register_builtin_tools from app.contracts import ModelCapability, ProviderConfig, ProviderType -from app.providers import MockProvider, ProviderRegistry +from app.providers import MockProvider, ProviderFactory, ProviderRegistry +from app.providers.credentials import EnvironmentCredentialResolver @dataclass(frozen=True) class ApplicationContainer: providers: ProviderRegistry + provider_factory: ProviderFactory tools: ToolRegistry permissions: PermissionManager agent: AgentRuntime def build_container() -> ApplicationContainer: + provider_factory = ProviderFactory(EnvironmentCredentialResolver()) providers = ProviderRegistry() providers.register( ProviderConfig( @@ -40,6 +43,7 @@ def build_container() -> ApplicationContainer: agent = AgentRuntime(providers=providers, tools=tools, permissions=permissions) return ApplicationContainer( providers=providers, + provider_factory=provider_factory, tools=tools, permissions=permissions, agent=agent, diff --git a/backend/app/contracts.py b/backend/app/contracts.py index c86eb2d..8cc890e 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -145,6 +145,7 @@ class Message(Contract): content: str name: str | None = None tool_call_id: str | None = None + tool_calls: list["ToolCall"] = Field(default_factory=list) class ToolDefinition(Contract): diff --git a/backend/app/providers/__init__.py b/backend/app/providers/__init__.py index 0555af2..8c08f0d 100644 --- a/backend/app/providers/__init__.py +++ b/backend/app/providers/__init__.py @@ -1,11 +1,15 @@ -from app.providers.base import ModelProvider, ProviderToolCall, ProviderTurn +from app.providers.base import ModelProvider, ProviderError, ProviderToolCall, ProviderTurn +from app.providers.factory import ProviderFactory, UnsupportedProviderError from app.providers.mock import MockProvider from app.providers.registry import ProviderRegistry __all__ = [ "MockProvider", "ModelProvider", + "ProviderError", + "ProviderFactory", "ProviderRegistry", "ProviderToolCall", "ProviderTurn", + "UnsupportedProviderError", ] diff --git a/backend/app/providers/base.py b/backend/app/providers/base.py index 039d3f3..10ff073 100644 --- a/backend/app/providers/base.py +++ b/backend/app/providers/base.py @@ -5,6 +5,13 @@ from typing import Protocol from app.contracts import ModelEvent, ModelInfo, ModelRequest +class ProviderError(RuntimeError): + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.code = code + self.message = message + + @dataclass(slots=True) class ProviderToolCall: tool_call_id: str diff --git a/backend/app/providers/credentials.py b/backend/app/providers/credentials.py new file mode 100644 index 0000000..0340c3f --- /dev/null +++ b/backend/app/providers/credentials.py @@ -0,0 +1,17 @@ +import os +import re +from typing import Protocol + + +class CredentialResolver(Protocol): + def resolve(self, credential_id: str | None) -> str | None: ... + + +class EnvironmentCredentialResolver: + """解析由桌面 Host 注入 Sidecar 进程的临时凭证上下文。""" + + def resolve(self, credential_id: str | None) -> str | None: + if not credential_id: + return None + normalized = re.sub(r"[^A-Za-z0-9]", "_", credential_id).upper() + return os.getenv(f"AINOTE_CREDENTIAL_{normalized}") diff --git a/backend/app/providers/factory.py b/backend/app/providers/factory.py new file mode 100644 index 0000000..8236882 --- /dev/null +++ b/backend/app/providers/factory.py @@ -0,0 +1,48 @@ +from app.contracts import ModelCapability, ProviderConfig, ProviderType +from app.providers.base import ModelProvider +from app.providers.credentials import CredentialResolver +from app.providers.ollama import OllamaProvider +from app.providers.openai_compatible import OpenAICompatibleProvider + + +class UnsupportedProviderError(ValueError): + pass + + +class ProviderFactory: + def __init__(self, credentials: CredentialResolver) -> None: + self.credentials = credentials + + def build(self, config: ProviderConfig) -> ModelProvider: + if config.provider_type in { + ProviderType.openai_chat, + ProviderType.openai_compatible, + }: + return OpenAICompatibleProvider( + base_url=config.base_url or "https://api.openai.com/v1", + credential_id=config.credential_id, + credentials=self.credentials, + ) + if config.provider_type == ProviderType.ollama: + return OllamaProvider(config.base_url or "http://127.0.0.1:11434") + raise UnsupportedProviderError(config.provider_type.value) + + @staticmethod + def capabilities(provider_type: ProviderType) -> list[ModelCapability]: + if provider_type in { + ProviderType.openai_chat, + ProviderType.openai_compatible, + }: + return [ + ModelCapability.chat, + ModelCapability.tool_calling, + ModelCapability.streaming, + ModelCapability.structured_output, + ] + if provider_type == ProviderType.ollama: + return [ + ModelCapability.chat, + ModelCapability.tool_calling, + ModelCapability.streaming, + ] + return [] diff --git a/backend/app/providers/http_base.py b/backend/app/providers/http_base.py new file mode 100644 index 0000000..5e1a96c --- /dev/null +++ b/backend/app/providers/http_base.py @@ -0,0 +1,82 @@ +import json +from collections.abc import AsyncIterator +from datetime import datetime, timezone + +from app.contracts import ModelEvent, ModelEventType, ModelRequest +from app.providers.base import ProviderError, ProviderTurn + + +class TurnStreamingMixin: + async def complete(self, request: ModelRequest) -> ProviderTurn: + raise NotImplementedError + + async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]: + try: + turn = await self.complete(request) + sequence = 0 + if turn.text: + yield ModelEvent( + event=ModelEventType.text_delta, + sequence=sequence, + data={"text": turn.text}, + timestamp=datetime.now(timezone.utc), + ) + sequence += 1 + for call in turn.tool_calls: + yield ModelEvent( + event=ModelEventType.tool_call_start, + sequence=sequence, + data={ + "tool_call_id": call.tool_call_id, + "name": call.name, + "arguments": call.arguments, + }, + timestamp=datetime.now(timezone.utc), + ) + sequence += 1 + yield ModelEvent( + event=ModelEventType.tool_call_end, + sequence=sequence, + data={"tool_call_id": call.tool_call_id}, + timestamp=datetime.now(timezone.utc), + ) + sequence += 1 + yield ModelEvent( + event=ModelEventType.usage, + sequence=sequence, + data={ + "input_tokens": turn.input_tokens, + "output_tokens": turn.output_tokens, + }, + timestamp=datetime.now(timezone.utc), + ) + yield ModelEvent( + event=ModelEventType.done, + sequence=sequence + 1, + timestamp=datetime.now(timezone.utc), + ) + except ProviderError as exc: + yield ModelEvent( + event=ModelEventType.error, + data={"code": exc.code, "message": exc.message}, + timestamp=datetime.now(timezone.utc), + ) + yield ModelEvent( + event=ModelEventType.done, + sequence=1, + timestamp=datetime.now(timezone.utc), + ) + + +def decode_tool_arguments(value: object) -> dict[str, object]: + if isinstance(value, dict): + return value + if not isinstance(value, str): + raise ProviderError("PROVIDER_INVALID_RESPONSE", "Tool arguments are not JSON.") + try: + decoded = json.loads(value) + except json.JSONDecodeError as exc: + raise ProviderError("PROVIDER_INVALID_RESPONSE", "Tool arguments are invalid JSON.") from exc + if not isinstance(decoded, dict): + raise ProviderError("PROVIDER_INVALID_RESPONSE", "Tool arguments must be an object.") + return decoded diff --git a/backend/app/providers/ollama.py b/backend/app/providers/ollama.py new file mode 100644 index 0000000..b796b8c --- /dev/null +++ b/backend/app/providers/ollama.py @@ -0,0 +1,117 @@ +from uuid import uuid4 + +import httpx + +from app.contracts import ModelCapability, ModelInfo, ModelRequest +from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn +from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments + + +class OllamaProvider(TurnStreamingMixin): + def __init__( + self, + base_url: str = "http://127.0.0.1:11434", + timeout_seconds: float = 120, + transport: httpx.AsyncBaseTransport | None = None, + ) -> None: + self.base_url = base_url.rstrip("/") + self.timeout_seconds = timeout_seconds + self.transport = transport + + 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), + ) + + 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 test_connection(self, model: str | None = None) -> tuple[bool, str]: + try: + models = await self.list_models() + except ProviderError as exc: + return False, exc.message + 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 new file mode 100644 index 0000000..3b42a1f --- /dev/null +++ b/backend/app/providers/openai_compatible.py @@ -0,0 +1,156 @@ +import json +from uuid import uuid4 + +import httpx + +from app.contracts import MessageRole, ModelCapability, ModelInfo, ModelRequest +from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn +from app.providers.credentials import CredentialResolver +from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments + + +class OpenAICompatibleProvider(TurnStreamingMixin): + def __init__( + self, + base_url: str, + credential_id: str | None, + credentials: CredentialResolver, + timeout_seconds: float = 60, + transport: httpx.AsyncBaseTransport | None = None, + ) -> None: + self.base_url = base_url.rstrip("/") + self.credential_id = credential_id + self.credentials = credentials + self.timeout_seconds = timeout_seconds + self.transport = transport + + async def complete(self, request: ModelRequest) -> ProviderTurn: + payload: dict[str, object] = { + "model": request.model, + "messages": self._messages(request), + "stream": False, + } + if request.tools: + payload["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 + if request.max_tokens is not None: + payload["max_tokens"] = request.max_tokens + if request.response_format is not None: + payload["response_format"] = request.response_format + + 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), + ) + + 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") + ] + + async def test_connection(self, model: str | None = None) -> tuple[bool, str]: + try: + models = await self.list_models() + except ProviderError as exc: + return False, exc.message + if model and model not in {item.model for item in models}: + return False, f"Model is not available: {model}" + return True, f"Connected; discovered {len(models)} model(s)." + + def _messages(self, request: ModelRequest) -> list[dict[str, object]]: + result: list[dict[str, object]] = [] + 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, + } + 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 + ] + result.append(item) + return result + + async def _request(self, method: str, path: str, **kwargs) -> dict: + headers = {"Content-Type": "application/json"} + api_key = self.credentials.resolve(self.credential_id) + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + 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: + code = { + 401: "PROVIDER_AUTH_FAILED", + 404: "MODEL_NOT_FOUND", + 429: "PROVIDER_RATE_LIMITED", + }.get(exc.response.status_code, "PROVIDER_UNAVAILABLE") + raise ProviderError(code, f"Provider returned HTTP {exc.response.status_code}.") 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 diff --git a/backend/app/providers/registry.py b/backend/app/providers/registry.py index c19e7a0..77ce9d9 100644 --- a/backend/app/providers/registry.py +++ b/backend/app/providers/registry.py @@ -27,15 +27,23 @@ class ProviderRegistry: def unregister(self, provider_id: str) -> None: self._providers.pop(provider_id, None) + def replace(self, config: ProviderConfig, adapter: ModelProvider) -> None: + if config.provider_id not in self._providers: + raise ProviderNotFoundError(config.provider_id) + self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter) + def get(self, provider_id: str) -> RegisteredProvider: - try: - provider = self._providers[provider_id] - except KeyError as exc: - raise ProviderNotFoundError(provider_id) from exc + provider = self.get_any(provider_id) if not provider.config.enabled: raise ProviderNotFoundError(provider_id) return provider + def get_any(self, provider_id: str) -> RegisteredProvider: + try: + return self._providers[provider_id] + except KeyError as exc: + raise ProviderNotFoundError(provider_id) from exc + def list_configs(self) -> list[ProviderConfig]: return [item.config.model_copy(deep=True) for item in self._providers.values()] diff --git a/backend/app/routes.py b/backend/app/routes.py index b847d12..7f5db3f 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -1,5 +1,6 @@ from collections.abc import AsyncIterator from datetime import datetime, timezone +from uuid import uuid4 from fastapi import APIRouter, Query from fastapi.responses import StreamingResponse @@ -49,6 +50,7 @@ from app.agent import AgentRunNotFoundError from app.container import container from app.errors import ApiError, not_implemented from app.providers.registry import ProviderNotFoundError +from app.providers.factory import UnsupportedProviderError router = APIRouter(prefix="/api") not_implemented_response = {501: {"model": ErrorResponse, "description": "业务服务尚未实现"}} @@ -86,6 +88,18 @@ def agent_run_or_404(run_id: str) -> AgentRun: ) from exc +def configurable_provider_or_404(provider_id: str): + try: + return container.providers.get_any(provider_id) + except ProviderNotFoundError as exc: + raise ApiError( + 404, + "PROVIDER_NOT_FOUND", + f"Provider is not registered: {provider_id}", + {"provider_id": provider_id}, + ) from exc + + # Notes @router.get("/notes", response_model=NoteListResponse, tags=["Notes"]) async def list_notes( @@ -398,7 +412,7 @@ async def list_providers() -> ProviderListResponse: tags=["Providers"], ) async def get_provider(provider_id: str) -> ProviderConfig: - return provider_or_404(provider_id).config.model_copy(deep=True) + return configurable_provider_or_404(provider_id).config.model_copy(deep=True) @router.post( @@ -407,8 +421,27 @@ async def get_provider(provider_id: str) -> ProviderConfig: responses=not_implemented_response, tags=["Providers"], ) -async def create_provider(_: ProviderCreateRequest) -> ProviderConfig: - not_implemented("providers.create") +async def create_provider(request: ProviderCreateRequest) -> ProviderConfig: + config = ProviderConfig( + provider_id=f"provider_{uuid4().hex}", + provider_type=request.provider_type, + name=request.name, + base_url=request.base_url, + default_model=request.default_model, + credential_id=request.credential_id, + enabled=request.enabled, + capabilities=container.provider_factory.capabilities(request.provider_type), + ) + try: + adapter = container.provider_factory.build(config) + except UnsupportedProviderError as exc: + raise ApiError( + 422, + "PROVIDER_TYPE_UNSUPPORTED", + f"Provider adapter is not implemented: {request.provider_type.value}", + ) from exc + container.providers.register(config, adapter) + return config @router.patch( @@ -417,8 +450,16 @@ async def create_provider(_: ProviderCreateRequest) -> ProviderConfig: responses=not_implemented_response, tags=["Providers"], ) -async def update_provider(provider_id: str, _: ProviderUpdateRequest) -> ProviderConfig: - not_implemented(f"providers.update:{provider_id}") +async def update_provider( + provider_id: str, request: ProviderUpdateRequest +) -> ProviderConfig: + current = configurable_provider_or_404(provider_id).config + if provider_id == "mock": + raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified.") + config = current.model_copy(update=request.model_dump(exclude_none=True)) + adapter = container.provider_factory.build(config) + container.providers.replace(config, adapter) + return config @router.delete( @@ -428,7 +469,11 @@ async def update_provider(provider_id: str, _: ProviderUpdateRequest) -> Provide tags=["Providers"], ) async def delete_provider(provider_id: str) -> OperationResponse: - not_implemented(f"providers.delete:{provider_id}") + configurable_provider_or_404(provider_id) + if provider_id == "mock": + raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted.") + container.providers.unregister(provider_id) + return OperationResponse(status="completed", resource_id=provider_id) @router.get( @@ -452,6 +497,18 @@ async def list_provider_models(provider_id: str) -> ProviderModelsResponse: tags=["Providers"], ) async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse: + registered = configurable_provider_or_404(request.provider_id) + if request.credential_context_id: + temporary_config = registered.config.model_copy( + update={"credential_id": request.credential_context_id, "enabled": True} + ) + adapter = container.provider_factory.build(temporary_config) + success, message = await adapter.test_connection(request.model) + return ProviderTestResponse( + provider_id=request.provider_id, + success=success, + message=message, + ) provider_or_404(request.provider_id) return await container.providers.test(request.provider_id, request.model) diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 3551f6d..64ca599 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -6,6 +6,7 @@ readme = "README.md" requires-python = ">=3.11" dependencies = [ "fastapi>=0.116,<1.0", + "httpx>=0.28,<1.0", "uvicorn[standard]>=0.35,<1.0", ] diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index d1a17fc..1a374b8 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -2,6 +2,8 @@ import asyncio from app.main import health, service_status from app.routes import get_index_status, list_notes, list_plugins, list_providers, list_skills +from app.routes import create_provider, delete_provider, get_provider, update_provider +from app.contracts import ProviderCreateRequest, ProviderType, ProviderUpdateRequest def test_health() -> None: @@ -53,3 +55,25 @@ def test_openapi_contains_documented_frontend_interfaces() -> None: } assert expected_paths <= paths.keys() + + +def test_provider_configuration_lifecycle() -> None: + created = asyncio.run( + create_provider( + ProviderCreateRequest( + provider_type=ProviderType.ollama, + name="Local Ollama", + base_url="http://127.0.0.1:11434", + default_model="qwen3:latest", + ) + ) + ) + fetched = asyncio.run(get_provider(created.provider_id)) + disabled = asyncio.run( + update_provider(created.provider_id, ProviderUpdateRequest(enabled=False)) + ) + deleted = asyncio.run(delete_provider(created.provider_id)) + + assert fetched.provider_type == ProviderType.ollama + assert disabled.enabled is False + assert deleted.resource_id == created.provider_id diff --git a/backend/tests/test_provider_adapters.py b/backend/tests/test_provider_adapters.py new file mode 100644 index 0000000..91e03b9 --- /dev/null +++ b/backend/tests/test_provider_adapters.py @@ -0,0 +1,154 @@ +import asyncio +import json + +import httpx + +from app.contracts import Message, MessageRole, ModelRequest, ToolCall, ToolDefinition +from app.providers.ollama import OllamaProvider +from app.providers.openai_compatible import OpenAICompatibleProvider + + +class StaticCredentials: + def resolve(self, credential_id: str | None) -> str | None: + return "secret-test-key" if credential_id else None + + +def run(coroutine): + return asyncio.run(coroutine) + + +def test_openai_compatible_maps_tool_call_and_credentials() -> None: + captured: dict = {} + + def handler(request: httpx.Request) -> httpx.Response: + assert request.headers["Authorization"] == "Bearer secret-test-key" + captured.update(json.loads(request.content)) + return httpx.Response( + 200, + json={ + "choices": [ + { + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "math.add", + "arguments": '{"left":1,"right":2}', + }, + } + ], + } + } + ], + "usage": {"prompt_tokens": 8, "completion_tokens": 4}, + }, + ) + + provider = OpenAICompatibleProvider( + base_url="https://provider.test/v1", + credential_id="openai-test", + credentials=StaticCredentials(), + transport=httpx.MockTransport(handler), + ) + turn = run( + provider.complete( + ModelRequest( + provider_id="test", + model="test-model", + messages=[Message(role=MessageRole.user, content="add")], + tools=[ + ToolDefinition( + name="math.add", + description="Add numbers", + parameters={"type": "object"}, + ) + ], + ) + ) + ) + + assert captured["tools"][0]["function"]["name"] == "math.add" + assert turn.tool_calls[0].name == "math.add" + assert turn.tool_calls[0].arguments == {"left": 1, "right": 2} + assert turn.input_tokens == 8 + + +def test_openai_compatible_preserves_tool_call_context() -> None: + captured: dict = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured.update(json.loads(request.content)) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"role": "assistant", "content": "done"}}], + "usage": {}, + }, + ) + + provider = OpenAICompatibleProvider( + base_url="https://provider.test/v1", + credential_id=None, + credentials=StaticCredentials(), + transport=httpx.MockTransport(handler), + ) + call = ToolCall( + tool_call_id="call_1", name="math.add", arguments={"left": 1, "right": 2} + ) + turn = run( + provider.complete( + ModelRequest( + provider_id="test", + model="test-model", + messages=[ + Message(role=MessageRole.user, content="add"), + Message(role=MessageRole.assistant, content="", tool_calls=[call]), + Message( + role=MessageRole.tool, + content='{"value":3}', + name="math.add", + tool_call_id="call_1", + ), + ], + ) + ) + ) + + assert captured["messages"][1]["tool_calls"][0]["id"] == "call_1" + assert captured["messages"][2]["tool_call_id"] == "call_1" + assert turn.text == "done" + + +def test_ollama_maps_models_and_completion() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/api/tags": + return httpx.Response(200, json={"models": [{"name": "qwen3:latest"}]}) + return httpx.Response( + 200, + json={ + "message": {"role": "assistant", "content": "local answer"}, + "prompt_eval_count": 5, + "eval_count": 2, + }, + ) + + provider = OllamaProvider(transport=httpx.MockTransport(handler)) + models = run(provider.list_models()) + turn = run( + provider.complete( + ModelRequest( + provider_id="ollama", + model="qwen3:latest", + messages=[Message(role=MessageRole.user, content="hello")], + ) + ) + ) + + assert models[0].model == "qwen3:latest" + assert turn.text == "local answer" + assert turn.input_tokens == 5 + assert turn.output_tokens == 2 diff --git a/backend/uv.lock b/backend/uv.lock index dc2459b..ecebb4b 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -33,6 +33,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] +[[package]] +name = "certifi" +version = "2026.7.22" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/a3/c2/24167ea9858356b47a87a50d39908bfdb72ceeefe0041586e704e5376b3a/certifi-2026.7.22.tar.gz", hash = "sha256:741e2c3b351ddf169a738da9f2c048608ff7f2c5cc02f1ebc6b118bb090d5d55", size = 138112, upload-time = "2026-07-22T03:35:12.644Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0b/a7/71ac2cff56fec219ed242bb11b8efb69fcc4bec75db06fb7bfe35de520e6/certifi-2026.7.22-py3-none-any.whl", hash = "sha256:62f22742b58a1a33014a2b6b706588a8d7e2a88ae7bd1a6ebe8c992928483775", size = 136983, upload-time = "2026-07-22T03:35:11.276Z" }, +] + [[package]] name = "click" version = "8.5.0" @@ -76,6 +85,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" }, ] +[[package]] +name = "httpcore" +version = "1.0.9" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "certifi" }, + { name = "h11" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/06/94/82699a10bca87a5556c9c59b5963f2d039dbd239f25bc2a63907a05a14cb/httpcore-1.0.9.tar.gz", hash = "sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8", size = 85484, upload-time = "2025-04-24T22:06:22.219Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7e/f5/f66802a942d491edb555dd61e3a9961140fd64c90bce1eafd741609d334d/httpcore-1.0.9-py3-none-any.whl", hash = "sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55", size = 78784, upload-time = "2025-04-24T22:06:20.566Z" }, +] + [[package]] name = "httptools" version = "0.8.0" @@ -119,6 +141,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/48/63/b906c01e53f50d432c0defe43ce52764a111dc1bdd028bafbeb54dcfd008/httptools-0.8.0-cp314-cp314t-win_amd64.whl", hash = "sha256:384c17174464c8e873398b7af24f0b1f44d992c820328413951a625323155d77", size = 108209, upload-time = "2026-05-25T22:17:39.473Z" }, ] +[[package]] +name = "httpx" +version = "0.28.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "certifi" }, + { name = "httpcore" }, + { name = "idna" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/b1/df/48c586a5fe32a0f01324ee087459e112ebb7224f646c0b5023f5e79e9956/httpx-0.28.1.tar.gz", hash = "sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc", size = 141406, upload-time = "2024-12-06T15:37:23.222Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517, upload-time = "2024-12-06T15:37:21.509Z" }, +] + [[package]] name = "idna" version = "3.19" @@ -143,6 +180,7 @@ version = "0.1.0" source = { virtual = "." } dependencies = [ { name = "fastapi" }, + { name = "httpx" }, { name = "uvicorn", extra = ["standard"] }, ] @@ -154,6 +192,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "fastapi", specifier = ">=0.116,<1.0" }, + { name = "httpx", specifier = ">=0.28,<1.0" }, { name = "uvicorn", extras = ["standard"], specifier = ">=0.35,<1.0" }, ] diff --git a/docs/AI-Core与Agent-Core开发说明.md b/docs/AI-Core与Agent-Core开发说明.md index 1c88dfc..4801f38 100644 --- a/docs/AI-Core与Agent-Core开发说明.md +++ b/docs/AI-Core与Agent-Core开发说明.md @@ -53,7 +53,7 @@ Router 只负责 HTTP/SSE 与错误转换,不实现 Agent、Tool 或 Provider - 文件系统和 API Key 明文读取:由 Rust Host 提供; - Skill、Plugin 生命周期:后续在 Extension Core 中实现。 -## 开发 Provider +## Provider 默认注册离线 Provider: @@ -84,6 +84,50 @@ streaming `POST /api/chat` 返回 ModelEvent SSE。 +另外已经实现以下可配置 Adapter: + +```text +openai_chat / openai_compatible +ollama +``` + +创建 Ollama Provider: + +```json +{ + "provider_type": "ollama", + "name": "Local Ollama", + "base_url": "http://127.0.0.1:11434", + "default_model": "qwen3:latest" +} +``` + +创建 OpenAI-Compatible Provider: + +```json +{ + "provider_type": "openai_compatible", + "name": "OpenAI Compatible", + "base_url": "https://api.openai.com/v1", + "default_model": "", + "credential_id": "openai-main" +} +``` + +凭证 ID `openai-main` 对应 Sidecar 进程中的临时环境变量 `AINOTE_CREDENTIAL_OPENAI_MAIN`。环境变量由 Rust Host 从 Stronghold 读取后注入,不写入 Provider Config、日志或前端 Store。 + +Provider 配置生命周期接口已经可用: + +```text +GET /api/providers +POST /api/providers +GET /api/providers/{provider_id} +PATCH /api/providers/{provider_id} +DELETE /api/providers/{provider_id} +GET /api/providers/{provider_id}/models +POST /api/providers/test +``` + ## Agent Run 创建普通 Agent Run: @@ -199,7 +243,8 @@ tasks.list ## 当前限制与下一步 -- Provider 目前只有完全离线的 Mock 实现;下一步实现 OpenAI-Compatible 与 Ollama Adapter。 +- 已实现 Mock、OpenAI-Compatible Chat Completions 与 Ollama Adapter;OpenAI Responses 和 Anthropic Messages 尚未实现。 +- Provider 配置暂存内存,后续通过 Repository 接入 SQLite。 - Run/Trace 暂存内存;下一步抽象 Repository 并接入 SQLite。 - Permission 已有核心等待/恢复机制,前端确认 UI 尚未联调。 - Note/RAG Tool 等待对应模块 Service 接入。 diff --git a/docs/后端接口契约-开发版.md b/docs/后端接口契约-开发版.md index 01a59df..9b05463 100644 --- a/docs/后端接口契约-开发版.md +++ b/docs/后端接口契约-开发版.md @@ -153,7 +153,8 @@ RunCancelled ## 当前实现状态 -- Chat、Agent Run、Agent Events、Tool 列表、Provider 列表、模型列表和连接测试已经接入 AI Core。 +- Chat、Agent Run、Agent Events、Tool 列表、Provider 配置生命周期、模型列表和连接测试已经接入 AI Core。 +- Provider Adapter 当前包含 Mock、OpenAI-Compatible Chat Completions 和 Ollama。 - 默认提供 `mock/mock-1` 离线 Provider,以及 `system.echo`、`math.add` 开发 Tool。 - Notes、Search、Skills、Plugins、Tasks、Media、Index 等尚未接入业务服务的接口继续返回空结果、`idle` 或 `501`。 - 需要尚未接入的数据库、文件或扩展 Runtime 的操作统一返回 `501`。