57 lines
1.7 KiB
Python
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)
|