添加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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user