添加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:
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}")
|
||||
@@ -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 []
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Generated
+39
@@ -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" },
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user