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)