| 1 | from __future__ import annotations |
| 2 | |
| 3 | from collections.abc import Callable, Mapping |
| 4 | from typing import Any |
| 5 | |
| 6 | from plugins._oauth.helpers.providers import CODEX_PROVIDER_ID |
| 7 | from plugins._oauth.helpers.usage_plans import usage_plan_catalog |
| 8 | |
| 9 | |
| 10 | ProviderRegistryFactory = Callable[[], Mapping[str, Any]] |
| 11 | RoutesInstalledFactory = Callable[[], bool] |
| 12 | |
| 13 | |
| 14 | def build_oauth_status_summary( |
| 15 | *, |
| 16 | provider_registry: ProviderRegistryFactory, |
| 17 | routes_installed: RoutesInstalledFactory, |
| 18 | ) -> dict[str, Any]: |
| 19 | """Return the shared OAuth account status payload for APIs and discovery cards.""" |
| 20 | |
| 21 | providers = [_provider_status(provider) for provider in provider_registry().values()] |
| 22 | provider_map = { |
| 23 | provider["provider_id"]: provider |
| 24 | for provider in providers |
| 25 | if provider.get("provider_id") |
| 26 | } |
| 27 | connected = [provider for provider in providers if provider.get("connected")] |
| 28 | available = [provider for provider in providers if not provider.get("connected")] |
| 29 | catalog = usage_plan_catalog() |
| 30 | |
| 31 | return { |
| 32 | "ok": True, |
| 33 | "routes_installed": routes_installed(), |
| 34 | "providers": providers, |
| 35 | "provider_map": provider_map, |
| 36 | "usage_plan_catalog": catalog, |
| 37 | "connected_count": len(connected), |
| 38 | "available_count": len(available), |
| 39 | "oauth_accounts": { |
| 40 | "connected_count": len(connected), |
| 41 | "available_count": len(available), |
| 42 | "total_count": len(providers), |
| 43 | "connected": [_account_summary(provider, catalog) for provider in connected], |
| 44 | "available": [_account_summary(provider, catalog) for provider in available], |
| 45 | }, |
| 46 | "codex": provider_map.get(CODEX_PROVIDER_ID, {}), |
| 47 | } |
| 48 | |
| 49 | |
| 50 | def _provider_status(provider: Any) -> dict[str, Any]: |
| 51 | try: |
| 52 | status = provider.status() |
| 53 | if not isinstance(status, dict): |
| 54 | status = {} |
| 55 | provider_id = str(status.get("provider_id") or getattr(provider, "provider_id", "")) |
| 56 | status = { |
| 57 | **status, |
| 58 | "provider_id": provider_id, |
| 59 | "connected": bool(status.get("connected")), |
| 60 | } |
| 61 | usage_windows = _usage_windows(status) |
| 62 | if usage_windows: |
| 63 | status["usage_windows"] = usage_windows |
| 64 | return status |
| 65 | except Exception as exc: |
| 66 | metadata = _provider_metadata(provider) |
| 67 | return { |
| 68 | **metadata, |
| 69 | "provider_id": str(getattr(provider, "provider_id", "")), |
| 70 | "connected": False, |
| 71 | "error": str(exc), |
| 72 | } |
| 73 | |
| 74 | |
| 75 | def _provider_metadata(provider: Any) -> dict[str, Any]: |
| 76 | try: |
| 77 | metadata = provider.metadata() |
| 78 | to_dict = getattr(metadata, "to_dict", None) |
| 79 | value = to_dict() if callable(to_dict) else metadata |
| 80 | return value if isinstance(value, dict) else {} |
| 81 | except Exception: |
| 82 | return {} |
| 83 | |
| 84 | |
| 85 | def _account_summary(provider: dict[str, Any], catalog: dict[str, Any]) -> dict[str, Any]: |
| 86 | provider_id = str(provider.get("provider_id") or "") |
| 87 | plan_entry = catalog.get(provider_id) if isinstance(catalog, dict) else None |
| 88 | plans = plan_entry.get("plans", []) if isinstance(plan_entry, dict) else [] |
| 89 | return { |
| 90 | "provider_id": provider_id, |
| 91 | "display_name": provider.get("display_name") or provider_id, |
| 92 | "short_name": provider.get("short_name") or provider.get("display_name") or provider_id, |
| 93 | "connected": bool(provider.get("connected")), |
| 94 | "account_label": provider.get("account_label") or provider.get("email") or "", |
| 95 | "auth_flow": provider.get("auth_flow") or "", |
| 96 | "icon": provider.get("icon") or "", |
| 97 | "warning": provider.get("warning") or provider.get("models_warning") or "", |
| 98 | "usage_windows": list(provider.get("usage_windows") or []), |
| 99 | "plan_count": len(plans) if isinstance(plans, list) else 0, |
| 100 | } |
| 101 | |
| 102 | |
| 103 | def _usage_windows(status: dict[str, Any]) -> list[dict[str, Any]]: |
| 104 | usage = status.get("usage") if isinstance(status, dict) else {} |
| 105 | if not isinstance(usage, dict) or not usage.get("available"): |
| 106 | return [] |
| 107 | |
| 108 | windows: list[dict[str, Any]] = [] |
| 109 | for key, title in (("primary", "Session"), ("secondary", "Week")): |
| 110 | window = usage.get(key) |
| 111 | if not isinstance(window, dict): |
| 112 | continue |
| 113 | remaining = _remaining_percent(window) |
| 114 | if remaining is None: |
| 115 | continue |
| 116 | windows.append({ |
| 117 | "key": key, |
| 118 | "title": title, |
| 119 | "label": window.get("label") or "", |
| 120 | "remaining_percent": remaining, |
| 121 | "reset_at": window.get("reset_at") or 0, |
| 122 | }) |
| 123 | return windows |
| 124 | |
| 125 | |
| 126 | def _remaining_percent(window: dict[str, Any]) -> float | None: |
| 127 | remaining = window.get("remaining_percent") |
| 128 | if remaining is not None: |
| 129 | try: |
| 130 | return max(0, min(100, float(remaining))) |
| 131 | except (TypeError, ValueError): |
| 132 | pass |
| 133 | |
| 134 | used = window.get("used_percent") |
| 135 | if used is not None: |
| 136 | try: |
| 137 | return max(0, min(100, 100 - float(used))) |
| 138 | except (TypeError, ValueError): |
| 139 | return None |
| 140 | return None |