添加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,
)
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(
+5 -1
View File
@@ -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,
+1
View File
@@ -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):
+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.registry import ProviderRegistry
__all__ = [
"MockProvider",
"ModelProvider",
"ProviderError",
"ProviderFactory",
"ProviderRegistry",
"ProviderToolCall",
"ProviderTurn",
"UnsupportedProviderError",
]
+7
View File
@@ -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
+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:
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()]
+63 -6
View File
@@ -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)
+1
View File
@@ -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",
]
+24
View File
@@ -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
+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" },
]
[[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" },
]
+47 -2
View File
@@ -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": "<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 AdapterOpenAI Responses 和 Anthropic Messages 尚未实现
- Provider 配置暂存内存,后续通过 Repository 接入 SQLite。
- Run/Trace 暂存内存;下一步抽象 Repository 并接入 SQLite。
- Permission 已有核心等待/恢复机制,前端确认 UI 尚未联调。
- 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。
- Notes、Search、Skills、Plugins、Tasks、Media、Index 等尚未接入业务服务的接口继续返回空结果、`idle``501`
- 需要尚未接入的数据库、文件或扩展 Runtime 的操作统一返回 `501`