添加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 902ac4a234
commit 02651505e2
16 changed files with 745 additions and 17 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" },
] ]