Files
NotesAgentic/backend/app/providers/registry.py
T
admin 1741f7b1aa 添加provider工厂和Ollama支持
- 实现ProviderFactory用于构建不同类型的provider适配器
- 添加EnvironmentCredentialResolver用于解析环境变量中的凭证
- 实现OllamaProvider支持本地模型调用
- 实现OpenAICompatibleProvider支持OpenAI兼容接口
- 在AgentRuntime中添加对ProviderError的处理
- 更新Message结构体添加tool_calls字段
- 实现provider配置的增删改查API端点
- 添加provider注册表的replace方法
- 添加HTTP基础类和工具参数解码功能
- 更新依赖添加httpx库
- 添加相关单元测试验证provider适配器功能
```
2026-08-27 14:16:57 +08:00

64 lines
2.3 KiB
Python

from dataclasses import dataclass
from time import perf_counter
from app.contracts import ModelInfo, ProviderConfig, ProviderTestResponse
from app.providers.base import ModelProvider
class ProviderNotFoundError(LookupError):
pass
@dataclass(slots=True)
class RegisteredProvider:
config: ProviderConfig
adapter: ModelProvider
class ProviderRegistry:
def __init__(self) -> None:
self._providers: dict[str, RegisteredProvider] = {}
def register(self, config: ProviderConfig, adapter: ModelProvider) -> None:
if config.provider_id in self._providers:
raise ValueError(f"Provider already registered: {config.provider_id}")
self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter)
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:
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()]
async def list_models(self, provider_id: str) -> list[ModelInfo]:
return await self.get(provider_id).adapter.list_models()
async def test(self, provider_id: str, model: str | None = None) -> ProviderTestResponse:
provider = self.get(provider_id)
started = perf_counter()
success, message = await provider.adapter.test_connection(model)
latency_ms = round((perf_counter() - started) * 1000)
return ProviderTestResponse(
provider_id=provider_id,
success=success,
latency_ms=latency_ms,
message=message,
)