添加provider工厂和Ollama支持

- 实现ProviderFactory用于构建不同类型的provider适配器
- 添加EnvironmentCredentialResolver用于解析环境变量中的凭证
- 实现OllamaProvider支持本地模型调用
- 实现OpenAICompatibleProvider支持OpenAI兼容接口
- 在AgentRuntime中添加对ProviderError的处理
- 更新Message结构体添加tool_calls字段
- 实现provider配置的增删改查API端点
- 添加provider注册表的replace方法
- 添加HTTP基础类和工具参数解码功能
- 更新依赖添加httpx库
- 添加相关单元测试验证provider适配器功能
```
This commit is contained in:
2026-08-27 14:16:57 +08:00
parent b71984d951
commit 1741f7b1aa
18 changed files with 794 additions and 20 deletions
+14 -5
View File
@@ -20,6 +20,7 @@ from app.contracts import (
ToolResult, ToolResult,
) )
from app.providers.registry import ProviderRegistry from app.providers.registry import ProviderRegistry
from app.providers.base import ProviderError
class AgentRunNotFoundError(LookupError): class AgentRunNotFoundError(LookupError):
@@ -141,6 +142,8 @@ class AgentRuntime:
self._finish_cancelled(record) self._finish_cancelled(record)
except TimeoutError: except TimeoutError:
self._fail(record, "AGENT_TIMEOUT", "Agent run exceeded its timeout.") 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: except Exception as exc:
self._fail(record, "AGENT_FAILED", str(exc)) self._fail(record, "AGENT_FAILED", str(exc))
@@ -183,12 +186,18 @@ class AgentRuntime:
return return
if turn.tool_calls: if turn.tool_calls:
for provider_call in turn.tool_calls: calls = [
call = ToolCall( ToolCall(
tool_call_id=provider_call.tool_call_id, tool_call_id=item.tool_call_id,
name=provider_call.name, name=item.name,
arguments=provider_call.arguments, 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) result = await self._execute_tool(record, call)
record.run.tool_results.append(result) record.run.tool_results.append(result)
messages.append( messages.append(
+5 -1
View File
@@ -3,18 +3,21 @@ from dataclasses import dataclass
from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry
from app.agent.builtin_tools import register_builtin_tools from app.agent.builtin_tools import register_builtin_tools
from app.contracts import ModelCapability, ProviderConfig, ProviderType 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) @dataclass(frozen=True)
class ApplicationContainer: class ApplicationContainer:
providers: ProviderRegistry providers: ProviderRegistry
provider_factory: ProviderFactory
tools: ToolRegistry tools: ToolRegistry
permissions: PermissionManager permissions: PermissionManager
agent: AgentRuntime agent: AgentRuntime
def build_container() -> ApplicationContainer: def build_container() -> ApplicationContainer:
provider_factory = ProviderFactory(EnvironmentCredentialResolver())
providers = ProviderRegistry() providers = ProviderRegistry()
providers.register( providers.register(
ProviderConfig( ProviderConfig(
@@ -40,6 +43,7 @@ def build_container() -> ApplicationContainer:
agent = AgentRuntime(providers=providers, tools=tools, permissions=permissions) agent = AgentRuntime(providers=providers, tools=tools, permissions=permissions)
return ApplicationContainer( return ApplicationContainer(
providers=providers, providers=providers,
provider_factory=provider_factory,
tools=tools, tools=tools,
permissions=permissions, permissions=permissions,
agent=agent, agent=agent,
+1
View File
@@ -145,6 +145,7 @@ class Message(Contract):
content: str content: str
name: str | None = None name: str | None = None
tool_call_id: str | None = None tool_call_id: str | None = None
tool_calls: list["ToolCall"] = Field(default_factory=list)
class ToolDefinition(Contract): class ToolDefinition(Contract):
+5 -1
View File
@@ -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.mock import MockProvider
from app.providers.registry import ProviderRegistry from app.providers.registry import ProviderRegistry
__all__ = [ __all__ = [
"MockProvider", "MockProvider",
"ModelProvider", "ModelProvider",
"ProviderError",
"ProviderFactory",
"ProviderRegistry", "ProviderRegistry",
"ProviderToolCall", "ProviderToolCall",
"ProviderTurn", "ProviderTurn",
"UnsupportedProviderError",
] ]
+7
View File
@@ -5,6 +5,13 @@ from typing import Protocol
from app.contracts import ModelEvent, ModelInfo, ModelRequest 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) @dataclass(slots=True)
class ProviderToolCall: class ProviderToolCall:
tool_call_id: str tool_call_id: str
+17
View File
@@ -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}")
+48
View File
@@ -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 []
+82
View File
@@ -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
+117
View File
@@ -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
+156
View File
@@ -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
+12 -4
View File
@@ -27,15 +27,23 @@ class ProviderRegistry:
def unregister(self, provider_id: str) -> None: def unregister(self, provider_id: str) -> None:
self._providers.pop(provider_id, 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: def get(self, provider_id: str) -> RegisteredProvider:
try: provider = self.get_any(provider_id)
provider = self._providers[provider_id]
except KeyError as exc:
raise ProviderNotFoundError(provider_id) from exc
if not provider.config.enabled: if not provider.config.enabled:
raise ProviderNotFoundError(provider_id) raise ProviderNotFoundError(provider_id)
return provider 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]: def list_configs(self) -> list[ProviderConfig]:
return [item.config.model_copy(deep=True) for item in self._providers.values()] return [item.config.model_copy(deep=True) for item in self._providers.values()]
+63 -6
View File
@@ -1,5 +1,6 @@
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from datetime import datetime, timezone from datetime import datetime, timezone
from uuid import uuid4
from fastapi import APIRouter, Query from fastapi import APIRouter, Query
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
@@ -49,6 +50,7 @@ from app.agent import AgentRunNotFoundError
from app.container import container from app.container import container
from app.errors import ApiError, not_implemented from app.errors import ApiError, not_implemented
from app.providers.registry import ProviderNotFoundError from app.providers.registry import ProviderNotFoundError
from app.providers.factory import UnsupportedProviderError
router = APIRouter(prefix="/api") router = APIRouter(prefix="/api")
not_implemented_response = {501: {"model": ErrorResponse, "description": "业务服务尚未实现"}} not_implemented_response = {501: {"model": ErrorResponse, "description": "业务服务尚未实现"}}
@@ -86,6 +88,18 @@ def agent_run_or_404(run_id: str) -> AgentRun:
) from exc ) 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 # Notes
@router.get("/notes", response_model=NoteListResponse, tags=["Notes"]) @router.get("/notes", response_model=NoteListResponse, tags=["Notes"])
async def list_notes( async def list_notes(
@@ -398,7 +412,7 @@ async def list_providers() -> ProviderListResponse:
tags=["Providers"], tags=["Providers"],
) )
async def get_provider(provider_id: str) -> ProviderConfig: 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( @router.post(
@@ -407,8 +421,27 @@ async def get_provider(provider_id: str) -> ProviderConfig:
responses=not_implemented_response, responses=not_implemented_response,
tags=["Providers"], tags=["Providers"],
) )
async def create_provider(_: ProviderCreateRequest) -> ProviderConfig: async def create_provider(request: ProviderCreateRequest) -> ProviderConfig:
not_implemented("providers.create") 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( @router.patch(
@@ -417,8 +450,16 @@ async def create_provider(_: ProviderCreateRequest) -> ProviderConfig:
responses=not_implemented_response, responses=not_implemented_response,
tags=["Providers"], tags=["Providers"],
) )
async def update_provider(provider_id: str, _: ProviderUpdateRequest) -> ProviderConfig: async def update_provider(
not_implemented(f"providers.update:{provider_id}") 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( @router.delete(
@@ -428,7 +469,11 @@ async def update_provider(provider_id: str, _: ProviderUpdateRequest) -> Provide
tags=["Providers"], tags=["Providers"],
) )
async def delete_provider(provider_id: str) -> OperationResponse: 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( @router.get(
@@ -452,6 +497,18 @@ async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
tags=["Providers"], tags=["Providers"],
) )
async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse: 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) provider_or_404(request.provider_id)
return await container.providers.test(request.provider_id, request.model) return await container.providers.test(request.provider_id, request.model)
+1
View File
@@ -6,6 +6,7 @@ readme = "README.md"
requires-python = ">=3.11" requires-python = ">=3.11"
dependencies = [ dependencies = [
"fastapi>=0.116,<1.0", "fastapi>=0.116,<1.0",
"httpx>=0.28,<1.0",
"uvicorn[standard]>=0.35,<1.0", "uvicorn[standard]>=0.35,<1.0",
] ]
+24
View File
@@ -2,6 +2,8 @@ import asyncio
from app.main import health, service_status 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 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: def test_health() -> None:
@@ -53,3 +55,25 @@ def test_openapi_contains_documented_frontend_interfaces() -> None:
} }
assert expected_paths <= paths.keys() 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
+154
View File
@@ -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
+39
View File
@@ -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" }, { 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]] [[package]]
name = "click" name = "click"
version = "8.5.0" 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" }, { 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]] [[package]]
name = "httptools" name = "httptools"
version = "0.8.0" 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" }, { 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]] [[package]]
name = "idna" name = "idna"
version = "3.19" version = "3.19"
@@ -143,6 +180,7 @@ version = "0.1.0"
source = { virtual = "." } source = { virtual = "." }
dependencies = [ dependencies = [
{ name = "fastapi" }, { name = "fastapi" },
{ name = "httpx" },
{ name = "uvicorn", extra = ["standard"] }, { name = "uvicorn", extra = ["standard"] },
] ]
@@ -154,6 +192,7 @@ dev = [
[package.metadata] [package.metadata]
requires-dist = [ requires-dist = [
{ name = "fastapi", specifier = ">=0.116,<1.0" }, { name = "fastapi", specifier = ">=0.116,<1.0" },
{ name = "httpx", specifier = ">=0.28,<1.0" },
{ name = "uvicorn", extras = ["standard"], specifier = ">=0.35,<1.0" }, { name = "uvicorn", extras = ["standard"], specifier = ">=0.35,<1.0" },
] ]
+47 -2
View File
@@ -53,7 +53,7 @@ Router 只负责 HTTP/SSE 与错误转换,不实现 Agent、Tool 或 Provider
- 文件系统和 API Key 明文读取:由 Rust Host 提供; - 文件系统和 API Key 明文读取:由 Rust Host 提供;
- Skill、Plugin 生命周期:后续在 Extension Core 中实现。 - Skill、Plugin 生命周期:后续在 Extension Core 中实现。
## 开发 Provider ## Provider
默认注册离线 Provider 默认注册离线 Provider
@@ -84,6 +84,50 @@ streaming
`POST /api/chat` 返回 ModelEvent SSE。 `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": "<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
创建普通 Agent Run 创建普通 Agent Run
@@ -199,7 +243,8 @@ tasks.list
## 当前限制与下一步 ## 当前限制与下一步
- Provider 目前只有完全离线的 Mock 实现;下一步实现 OpenAI-Compatible 与 Ollama Adapter - 已实现 Mock、OpenAI-Compatible Chat Completions 与 Ollama AdapterOpenAI Responses 和 Anthropic Messages 尚未实现
- Provider 配置暂存内存,后续通过 Repository 接入 SQLite。
- Run/Trace 暂存内存;下一步抽象 Repository 并接入 SQLite。 - Run/Trace 暂存内存;下一步抽象 Repository 并接入 SQLite。
- Permission 已有核心等待/恢复机制,前端确认 UI 尚未联调。 - Permission 已有核心等待/恢复机制,前端确认 UI 尚未联调。
- Note/RAG Tool 等待对应模块 Service 接入。 - Note/RAG Tool 等待对应模块 Service 接入。
+2 -1
View File
@@ -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。 - 默认提供 `mock/mock-1` 离线 Provider,以及 `system.echo``math.add` 开发 Tool。
- Notes、Search、Skills、Plugins、Tasks、Media、Index 等尚未接入业务服务的接口继续返回空结果、`idle``501` - Notes、Search、Skills、Plugins、Tasks、Media、Index 等尚未接入业务服务的接口继续返回空结果、`idle``501`
- 需要尚未接入的数据库、文件或扩展 Runtime 的操作统一返回 `501` - 需要尚未接入的数据库、文件或扩展 Runtime 的操作统一返回 `501`