main
py 39 lines 1.25 KB
Raw
1 from __future__ import annotations
2
3 from plugins._oauth.helpers.providers.base import (
4 CODEX_PROVIDER_ID,
5 OAuthProvider,
6 )
7 from plugins._oauth.helpers.providers.codex import CodexOAuthProvider
8
9
10 def provider_registry() -> dict[str, OAuthProvider]:
11 from plugins._oauth.helpers.providers.gemini_api import GeminiApiOAuthProvider
12 from plugins._oauth.helpers.providers.github_copilot import GitHubCopilotOAuthProvider
13 from plugins._oauth.helpers.providers.xai_grok import XaiGrokOAuthProvider
14
15 providers: list[OAuthProvider] = [
16 CodexOAuthProvider(),
17 GitHubCopilotOAuthProvider(),
18 GeminiApiOAuthProvider(),
19 XaiGrokOAuthProvider(),
20 ]
21 return {provider.provider_id: provider for provider in providers}
22
23
24 def get_provider(provider_id: str | None = None) -> OAuthProvider:
25 if provider_id is None:
26 normalized = CODEX_PROVIDER_ID
27 else:
28 normalized = str(provider_id).strip()
29 if not normalized:
30 normalized = CODEX_PROVIDER_ID
31 registry = provider_registry()
32 try:
33 return registry[normalized]
34 except KeyError as exc:
35 raise KeyError(f"Unknown OAuth provider: {normalized}") from exc
36
37
38 def oauth_provider_ids() -> set[str]:
39 return set(provider_registry())