Files
cmhub/apps/ai/providers/registry.py
T
2026-07-02 10:33:15 +08:00

57 lines
1.7 KiB
Python

from __future__ import annotations
from .base import AiProviderConfigError, Provider
from .openai_compatible import (
ChatCompletionsProvider,
GeminiProvider,
ImagesEditsProvider,
ImagesGenerationProvider,
)
from .utils import (
API_AUTO,
API_CHAT,
API_GEMINI,
API_IMAGES,
API_IMAGES_EDITS,
detect_api_type,
)
class ProviderRegistry:
def __init__(self) -> None:
self._providers: dict[str, Provider] = {}
def register(self, api_type: str, provider: Provider) -> None:
self._providers[api_type] = provider
def resolve_api_type(self, api_type: str, url: str = "") -> str:
return detect_api_type(url, api_type) if api_type == API_AUTO else api_type
def get(self, api_type: str, url: str = "") -> Provider:
resolved_api_type = self.resolve_api_type(api_type, url)
try:
return self._providers[resolved_api_type]
except KeyError as exc:
raise AiProviderConfigError(
f"no provider registered for api_type: {resolved_api_type}"
) from exc
default_registry = ProviderRegistry()
default_registry.register(API_CHAT, ChatCompletionsProvider())
default_registry.register(API_GEMINI, GeminiProvider())
default_registry.register(API_IMAGES, ImagesGenerationProvider())
default_registry.register(API_IMAGES_EDITS, ImagesEditsProvider())
def register_provider(api_type: str, provider: Provider) -> None:
default_registry.register(api_type, provider)
def resolve_api_type(api_type: str, url: str = "") -> str:
return default_registry.resolve_api_type(api_type, url)
def get_provider(api_type: str, url: str = "") -> Provider:
return default_registry.get(api_type, url)