添加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
+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()]