From bb4f6df67a7da0a204c03958a0e14aa7c6026619 Mon Sep 17 00:00:00 2001 From: Hermes Team Date: Sun, 23 Aug 2026 23:09:58 +0700 Subject: [PATCH] fix(oauth): preserve complete credentials and sync oauth_creds.json for Antigravity profiles (A20) - Add openid to OAuth scopes for id_token issuance - Preserve id_token, scope, and token_type on token exchange and refresh - Atomically write and sync .gemini/oauth_creds.json in profile directories - Auto-resolve active profile environment in discover_models - Cache discovered models in models_cache.json with graceful timeout handling - Add unit test coverage for full OAuth lifecycle and model discovery caching --- src/antigravity_provider/agy_subprocess.py | 20 +- src/antigravity_provider/oauth.py | 22 +- .../router/adapters/antigravity_adapter.py | 13 +- .../router/model_discovery_service.py | 8 +- .../router/profile_manager.py | 214 +++++++++++++-- .../router/profile_oauth.py | 7 +- tests/test_agy_oauth_credentials.py | 255 ++++++++++++++++++ 7 files changed, 513 insertions(+), 26 deletions(-) create mode 100644 tests/test_agy_oauth_credentials.py diff --git a/src/antigravity_provider/agy_subprocess.py b/src/antigravity_provider/agy_subprocess.py index 8d12dcb..20d9698 100644 --- a/src/antigravity_provider/agy_subprocess.py +++ b/src/antigravity_provider/agy_subprocess.py @@ -123,10 +123,25 @@ def discover_models(profile_id: str | None = None) -> dict[str, str]: # команда запускалась в ГЛОБАЛЬНОМ окружении, где вход не выполнен, и # отвечала «Please sign in to view available models» — при шести рабочих # OAuth-профилях. Список моделей поэтому был пуст всегда. - if profile_id: + target_profile_id = profile_id + if not target_profile_id: + try: + from antigravity_provider.router.profile_manager import ProfileAuthManager + main_p = ProfileAuthManager.get_main_profile("antigravity") + if main_p and ProfileAuthManager.load_profile_auth("antigravity", main_p): + target_profile_id = main_p + else: + for candidate in ["ag-orch-primary", "ag-w1", "ag-w2", "ag-w3", "ag-w4", "ag-w5"]: + if ProfileAuthManager.load_profile_auth("antigravity", candidate): + target_profile_id = candidate + break + except Exception: + target_profile_id = None + + if target_profile_id: from antigravity_provider.router.adapters.antigravity_adapter import get_profile_env_dir - profile_dir = get_profile_env_dir(profile_id) + profile_dir = get_profile_env_dir(target_profile_id) env = build_safe_subprocess_env( overrides={ "USERPROFILE": str(profile_dir), @@ -146,6 +161,7 @@ def discover_models(profile_id: str | None = None) -> dict[str, str]: encoding="utf-8", errors="replace", env=env, + stdin=subprocess.DEVNULL, ) raw = result.stdout.strip() if not raw or result.returncode != 0: diff --git a/src/antigravity_provider/oauth.py b/src/antigravity_provider/oauth.py index 591ede1..fd5e726 100644 --- a/src/antigravity_provider/oauth.py +++ b/src/antigravity_provider/oauth.py @@ -28,6 +28,7 @@ TOKEN_URL = "https://oauth2.googleapis.com/token" CLIENT_ID = "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com" CLIENT_SECRET = "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" SCOPES = [ + "openid", "https://www.googleapis.com/auth/cloud-platform", "https://www.googleapis.com/auth/userinfo.email", "https://www.googleapis.com/auth/userinfo.profile", @@ -71,6 +72,8 @@ def _get_json(url: str, headers: dict[str, str]) -> dict[str, Any]: def refresh_access_token( refresh_token: str, *, + existing_id_token: str | None = None, + existing_scope: str | None = None, post_json: Callable[[str, dict[str, str], dict[str, str]], dict[str, Any]] | None = None, client: tuple[str, str] | None = None, ) -> dict[str, Any]: @@ -85,9 +88,16 @@ def refresh_access_token( data = post_json(TOKEN_URL, payload, {"Content-Type": "application/x-www-form-urlencoded"}) if not data.get("access_token"): raise ProxyError("OAuth refresh response did not include access_token", status=401, error_type="invalid_request_error") + + id_token = data.get("id_token") or existing_id_token or "" + scope = data.get("scope") or existing_scope or "" + return { "refresh_token": data.get("refresh_token") or refresh_token, "access_token": data["access_token"], + "id_token": id_token, + "scope": scope, + "expires_in": data.get("expires_in"), "expires_at": _expires_at(data.get("expires_in")), "token_type": data.get("token_type", "Bearer"), } @@ -98,7 +108,14 @@ def refresh_if_needed(credentials: dict[str, Any], *, skew_seconds: int = 60) -> refresh = credentials.get("refresh_token") or credentials.get("refresh") expires = credentials.get("expires_at") or credentials.get("expires") if refresh and (not access or (isinstance(expires, (int, float)) and time.time() + skew_seconds >= float(expires))): - credentials = {**credentials, **refresh_access_token(str(refresh))} + existing_id = credentials.get("id_token") + existing_scope = credentials.get("scope") + refreshed = refresh_access_token( + str(refresh), + existing_id_token=str(existing_id) if existing_id else None, + existing_scope=str(existing_scope) if existing_scope else None, + ) + credentials = {**credentials, **refreshed} return credentials @@ -177,6 +194,9 @@ def exchange_code_for_tokens( return { "refresh_token": refresh_token, "access_token": data["access_token"], + "id_token": data.get("id_token", ""), + "scope": data.get("scope", ""), + "expires_in": data.get("expires_in"), "expires_at": _expires_at(data.get("expires_in")), "token_type": data.get("token_type", "Bearer"), } diff --git a/src/antigravity_provider/router/adapters/antigravity_adapter.py b/src/antigravity_provider/router/adapters/antigravity_adapter.py index f25c77f..61f0fea 100644 --- a/src/antigravity_provider/router/adapters/antigravity_adapter.py +++ b/src/antigravity_provider/router/adapters/antigravity_adapter.py @@ -62,12 +62,15 @@ class AntigravityAdapter(BaseProviderAdapter): if profile_auth: # Pre-flight check: verify token expiry before calling subprocess to prevent interactive browser login - tokens = profile_auth.get("tokens", {}) - expiry = tokens.get("expiry_date") or profile_auth.get("expiry_date") + tokens = profile_auth.get("token") or profile_auth.get("tokens", {}) + refresh_tok = tokens.get("refresh_token") if isinstance(tokens, dict) else profile_auth.get("refresh_token") + expiry = tokens.get("expiry_date") if isinstance(tokens, dict) else profile_auth.get("expiry_date") + if not expiry and isinstance(tokens, dict): + expiry = tokens.get("expires_at") if expiry: - if expiry > 1e11: - expiry = expiry / 1000.0 - if time.time() > expiry: + if float(expiry) > 1e11: + expiry = float(expiry) / 1000.0 + if time.time() > float(expiry) and not refresh_tok: raise AuthExpiredError( "Авторизация истекла, требуется повторный вход.", provider="antigravity", diff --git a/src/antigravity_provider/router/model_discovery_service.py b/src/antigravity_provider/router/model_discovery_service.py index f8684fd..addb95d 100644 --- a/src/antigravity_provider/router/model_discovery_service.py +++ b/src/antigravity_provider/router/model_discovery_service.py @@ -208,7 +208,10 @@ class ModelDiscoveryService: logger.info("Discovered %d models for provider '%s': %s", len(models), provider, models) return models - return None + # If probe returned None, retain existing cache if any + with self._cache_lock: + entry = self._cache.get(provider.lower()) + return list(entry["models"]) if entry and "models" in entry else None def _probe_provider(self, provider: str) -> Optional[List[str]]: """Perform provider-specific model discovery.""" @@ -217,7 +220,8 @@ class ModelDiscoveryService: if prov == "antigravity": from antigravity_provider.agy_subprocess import discover_models - res = discover_models() + main_p = ProfileAuthManager.get_main_profile("antigravity") + res = discover_models(profile_id=main_p) if res: return sorted(list(set(res.values()))) return None diff --git a/src/antigravity_provider/router/profile_manager.py b/src/antigravity_provider/router/profile_manager.py index 00e5a4b..4691433 100644 --- a/src/antigravity_provider/router/profile_manager.py +++ b/src/antigravity_provider/router/profile_manager.py @@ -15,6 +15,7 @@ import threading import time import urllib.request import urllib.error +from datetime import datetime from pathlib import Path from typing import Any, Dict, List, Optional, Tuple @@ -113,6 +114,73 @@ class ProfileAuthManager: """Official API to get isolated directory for a profile.""" return get_profile_dir(profile_id, provider) + @classmethod + def write_agy_oauth_creds(cls, profile_dir: Path, auth_data: dict) -> Path: + """Atomically write /.gemini/oauth_creds.json in exact agy CLI format.""" + token_info = auth_data.get("token") or auth_data.get("tokens") or auth_data + if not isinstance(token_info, dict): + token_info = {} + + access_token = token_info.get("access_token") or auth_data.get("access_token") or "" + refresh_token = token_info.get("refresh_token") or auth_data.get("refresh_token") or "" + scope = token_info.get("scope") or auth_data.get("scope") or "" + token_type = token_info.get("token_type") or auth_data.get("token_type") or "Bearer" + id_token = token_info.get("id_token") or auth_data.get("id_token") or "" + + # Expiry date in milliseconds (int) + expiry_date = token_info.get("expiry_date") or auth_data.get("expiry_date") + if not expiry_date: + expires_at = token_info.get("expires_at") or auth_data.get("expires_at") + if expires_at: + try: + expiry_date = int(float(expires_at) * 1000) + except Exception: + expiry_date = int((time.time() + 3600) * 1000) + else: + expiry_str = token_info.get("expiry") or auth_data.get("expiry") + if expiry_str: + try: + dt = datetime.fromisoformat(str(expiry_str).replace("Z", "+00:00")) + expiry_date = int(dt.timestamp() * 1000) + except Exception: + expiry_date = int((time.time() + 3600) * 1000) + else: + expiry_date = int((time.time() + 3600) * 1000) + elif float(expiry_date) < 1e11: # in seconds + expiry_date = int(float(expiry_date) * 1000) + else: + expiry_date = int(expiry_date) + + creds_dict = { + "access_token": str(access_token), + "refresh_token": str(refresh_token), + "scope": str(scope), + "token_type": str(token_type), + "id_token": str(id_token), + "expiry_date": expiry_date, + } + + gemini_dir = profile_dir / ".gemini" + gemini_dir.mkdir(parents=True, exist_ok=True) + try: + os.chmod(gemini_dir, 0o700) + except OSError: + pass + + target_file = gemini_dir / "oauth_creds.json" + temp_file = gemini_dir / f"oauth_creds.json.tmp-{threading.get_ident()}-{time.time_ns()}" + temp_file.write_text(json.dumps(creds_dict, indent=2), encoding="utf-8") + try: + os.chmod(temp_file, 0o600) + except OSError: + pass + os.replace(temp_file, target_file) + try: + os.chmod(target_file, 0o600) + except OSError: + pass + return target_file + @staticmethod def read_windows_credential(target_name: str = "gemini:antigravity") -> Optional[dict]: """Read a credential blob from Windows Credential Manager.""" @@ -140,7 +208,40 @@ class ProfileAuthManager: if not advapi32 or os.name != "nt": return False with _CM_LOCK: - blob_bytes = json.dumps(auth_data).encode("utf-8") + payload_data = auth_data + if target_name == "gemini:antigravity" and isinstance(auth_data, dict): + # Ensure 6-field agy-compatible schema in Credential Manager + token_info = auth_data.get("token") or auth_data.get("tokens") or auth_data + if isinstance(token_info, dict) and ("access_token" in token_info or "access_token" in auth_data): + acc = token_info.get("access_token") or auth_data.get("access_token") or "" + ref = token_info.get("refresh_token") or auth_data.get("refresh_token") or "" + sc = token_info.get("scope") or auth_data.get("scope") or "" + tt = token_info.get("token_type") or auth_data.get("token_type") or "Bearer" + idt = token_info.get("id_token") or auth_data.get("id_token") or "" + exp = token_info.get("expiry_date") or auth_data.get("expiry_date") + if not exp: + exp_at = token_info.get("expires_at") or auth_data.get("expires_at") + if exp_at: + try: + exp = int(float(exp_at) * 1000) + except Exception: + exp = int((time.time() + 3600) * 1000) + else: + exp = int((time.time() + 3600) * 1000) + elif float(exp) < 1e11: + exp = int(float(exp) * 1000) + else: + exp = int(exp) + payload_data = { + "access_token": str(acc), + "refresh_token": str(ref), + "scope": str(sc), + "token_type": str(tt), + "id_token": str(idt), + "expiry_date": exp, + } + + blob_bytes = json.dumps(payload_data).encode("utf-8") buf = ctypes.create_string_buffer(blob_bytes) cred = CREDENTIAL() cred.Flags = 0 @@ -201,9 +302,24 @@ class ProfileAuthManager: pdir.mkdir(parents=True, exist_ok=True) auth_file = pdir / "auth.json" existed = auth_file.is_file() - temp_file = pdir / f"auth.json.tmp-{threading.get_ident()}" + temp_file = pdir / f"auth.json.tmp-{threading.get_ident()}-{time.time_ns()}" temp_file.write_text(json.dumps(auth_data, indent=2), encoding="utf-8") + try: + os.chmod(temp_file, 0o600) + except OSError: + pass os.replace(temp_file, auth_file) + try: + os.chmod(auth_file, 0o600) + except OSError: + pass + + # For Antigravity, synchronously maintain /.gemini/oauth_creds.json + if provider in ("antigravity", "google-antigravity"): + try: + cls.write_agy_oauth_creds(pdir, auth_data) + except Exception as e: + logger.warning("Failed to write .gemini/oauth_creds.json for profile=%s: %s", profile_id, e) from antigravity_provider.router.event_bus import ( EVENT_ACCOUNT_ADDED, @@ -238,7 +354,58 @@ class ProfileAuthManager: auth_file = get_profile_auth_path(provider, profile_id) if auth_file.is_file(): try: - return json.loads(auth_file.read_text(encoding="utf-8")) + data = json.loads(auth_file.read_text(encoding="utf-8")) + if provider in ("antigravity", "google-antigravity") and isinstance(data, dict): + pdir = get_profile_dir(profile_id, provider) + gemini_creds = pdir / ".gemini" / "oauth_creds.json" + if not gemini_creds.is_file(): + try: + cls.write_agy_oauth_creds(pdir, data) + except Exception as e: + logger.debug("Failed to create missing oauth_creds.json: %s", e) + + # Auto-refresh expired or expiring access tokens if refresh_token is present + tokens = data.get("token") or data.get("tokens") + if isinstance(tokens, dict): + refresh_tok = tokens.get("refresh_token") + acc_tok = tokens.get("access_token") + exp_at = tokens.get("expires_at") + if not exp_at: + exp_str = tokens.get("expiry") + if exp_str: + try: + dt = datetime.fromisoformat(str(exp_str).replace("Z", "+00:00")) + exp_at = dt.timestamp() + except Exception: + exp_at = None + + now = time.time() + if refresh_tok and (not acc_tok or (exp_at and now + 60 >= float(exp_at))): + try: + from antigravity_provider.oauth import refresh_access_token + + existing_id = tokens.get("id_token") + existing_scope = tokens.get("scope") + refreshed = refresh_access_token( + str(refresh_tok), + existing_id_token=str(existing_id) if existing_id else None, + existing_scope=str(existing_scope) if existing_scope else None, + ) + tokens.update({ + "access_token": refreshed["access_token"], + "refresh_token": refreshed.get("refresh_token") or refresh_tok, + "id_token": refreshed.get("id_token") or existing_id or "", + "scope": refreshed.get("scope") or existing_scope or "", + "token_type": refreshed.get("token_type", "Bearer"), + "expires_at": refreshed.get("expires_at"), + "expiry": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(refreshed["expires_at"])), + }) + data["token"] = tokens + cls.save_profile_auth(provider, profile_id, data) + except Exception as re_err: + logger.warning("Silent token refresh failed for profile=%s: %s", profile_id, re_err) + + return data except Exception as e: logger.warning("Error reading %s: %s", auth_file, e) @@ -406,22 +573,38 @@ class ProfileAuthManager: "error": None, } - if provider == "antigravity": - tokens = auth_data.get("tokens", {}) - acc_token = tokens.get("access_token") or auth_data.get("access_token") - id_token = tokens.get("id_token") or auth_data.get("id_token") - email = None + if provider in ("antigravity", "google-antigravity"): + tokens = auth_data.get("token") or auth_data.get("tokens", {}) + acc_token = tokens.get("access_token") if isinstance(tokens, dict) else (auth_data.get("access_token") or "") + id_token = tokens.get("id_token") if isinstance(tokens, dict) else (auth_data.get("id_token") or "") + refresh_tok = tokens.get("refresh_token") if isinstance(tokens, dict) else (auth_data.get("refresh_token") or "") + email = auth_data.get("email") acc_id = None if id_token: - email, acc_id = cls.extract_jwt_identity(id_token) + email_from_jwt, acc_id = cls.extract_jwt_identity(id_token) + email = email or email_from_jwt + if not email and acc_token: + email_from_jwt, acc_id = cls.extract_jwt_identity(acc_token) + email = email or email_from_jwt + + expiry = tokens.get("expiry_date") if isinstance(tokens, dict) else auth_data.get("expiry_date") + if not expiry and isinstance(tokens, dict): + expiry = tokens.get("expires_at") + if not expiry: + expiry_str = tokens.get("expiry") if isinstance(tokens, dict) else auth_data.get("expiry") + if expiry_str: + try: + dt = datetime.fromisoformat(str(expiry_str).replace("Z", "+00:00")) + expiry = dt.timestamp() + except Exception: + expiry = None - expiry = tokens.get("expiry_date") or auth_data.get("expiry_date") is_expired = False if expiry: - if expiry > 1e11: - expiry = expiry / 1000.0 - if time.time() > expiry: - is_expired = True + if float(expiry) > 1e11: + expiry = float(expiry) / 1000.0 + if time.time() > float(expiry): + is_expired = not bool(refresh_tok) return { "authenticated": True, @@ -429,9 +612,10 @@ class ProfileAuthManager: "profile_id": profile_id, "email_masked": mask_email(email) if email else None, "account_id_masked": mask_id(acc_id) if acc_id else None, + "has_refresh_token": bool(refresh_tok), "is_expired": is_expired, "status": "EXPIRED" if is_expired else "AUTHENTICATED", - "error": "Token expired" if is_expired else None, + "error": "Token expired without refresh token" if is_expired else None, } elif provider in ("openai-codex", "codex"): diff --git a/src/antigravity_provider/router/profile_oauth.py b/src/antigravity_provider/router/profile_oauth.py index b227dc5..9f375e6 100644 --- a/src/antigravity_provider/router/profile_oauth.py +++ b/src/antigravity_provider/router/profile_oauth.py @@ -238,11 +238,16 @@ class ProfileOAuthSession: logger.info("OAuth account identity resolved (email_found=%s)", bool(email)) # Format in standard gemini:antigravity shape + expires_at = tokens.get("expires_at") or (int(time.time()) + 3600) auth_data = { "token": { "access_token": tokens["access_token"], "refresh_token": tokens["refresh_token"], - "expiry": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(tokens["expires_at"])), + "id_token": tokens.get("id_token", ""), + "scope": tokens.get("scope", ""), + "token_type": tokens.get("token_type", "Bearer"), + "expires_at": expires_at, + "expiry": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(expires_at)), }, "email": email or "", "auth_method": "oauth", diff --git a/tests/test_agy_oauth_credentials.py b/tests/test_agy_oauth_credentials.py new file mode 100644 index 0000000..2c5d3bb --- /dev/null +++ b/tests/test_agy_oauth_credentials.py @@ -0,0 +1,255 @@ +"""Tests for Antigravity OAuth full credentials preservation (P0-1), +writing .gemini/oauth_creds.json (P0-2), and Model Discovery caching (P1-4). +""" +from __future__ import annotations + +import json +import time +from pathlib import Path +from unittest.mock import patch, MagicMock +import pytest + +from antigravity_provider.oauth import ( + SCOPES, + build_auth_url, + exchange_code_for_tokens, + refresh_access_token, + refresh_if_needed, +) +from antigravity_provider.router.profile_manager import ProfileAuthManager, get_profile_dir +from antigravity_provider.router.profile_oauth import ProfileOAuthSession +from antigravity_provider.router.model_discovery_service import ModelDiscoveryService + + +# ── TEST P0-1: Full Credential Preservation in OAuth ── + + +@pytest.mark.unit +def test_oauth_scopes_include_openid(): + """P0-1: SCOPES must include 'openid' so Google issues an OpenID Connect id_token.""" + assert "openid" in SCOPES + url, verifier = build_auth_url() + assert "openid" in url + assert "code_challenge=" in url + + +@pytest.mark.unit +def test_exchange_code_returns_all_required_fields(): + """P0-1: exchange_code_for_tokens must return id_token, scope, token_type, expires_in, expires_at.""" + mock_resp = { + "access_token": "ya29.mock_access_123", + "refresh_token": "1//mock_refresh_456", + "id_token": "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.mock_jwt_payload.signature", + "scope": "openid https://www.googleapis.com/auth/userinfo.email", + "token_type": "Bearer", + "expires_in": 3600, + } + + with patch("antigravity_provider.oauth._post_form_json", return_value=mock_resp): + tokens = exchange_code_for_tokens("test_code", code_verifier="test_verifier") + + assert tokens["access_token"] == "ya29.mock_access_123" + assert tokens["refresh_token"] == "1//mock_refresh_456" + assert tokens["id_token"] == mock_resp["id_token"] + assert tokens["scope"] == mock_resp["scope"] + assert tokens["token_type"] == "Bearer" + assert tokens["expires_in"] == 3600 + assert isinstance(tokens["expires_at"], int) + + +@pytest.mark.unit +def test_refresh_access_token_preserves_id_token_and_scope(): + """P0-1: refresh_access_token preserves id_token and scope from response or existing fallback.""" + # Case 1: Google returns updated id_token and scope + mock_resp_full = { + "access_token": "ya29.new_access_token", + "id_token": "eyJhbGciOiJSUzI1NiJ9.new_jwt.sig", + "scope": "openid email profile", + "token_type": "Bearer", + "expires_in": 3600, + } + with patch("antigravity_provider.oauth._post_form_json", return_value=mock_resp_full): + res1 = refresh_access_token("1//mock_refresh") + assert res1["access_token"] == "ya29.new_access_token" + assert res1["id_token"] == "eyJhbGciOiJSUzI1NiJ9.new_jwt.sig" + assert res1["scope"] == "openid email profile" + + # Case 2: Google returns only access_token and expires_in (common on refresh) + mock_resp_minimal = { + "access_token": "ya29.refreshed_access", + "expires_in": 3600, + "token_type": "Bearer", + } + with patch("antigravity_provider.oauth._post_form_json", return_value=mock_resp_minimal): + res2 = refresh_access_token( + "1//mock_refresh", + existing_id_token="eyJhbGciOiJSUzI1NiJ9.preserved_jwt.sig", + existing_scope="openid https://www.googleapis.com/auth/userinfo.email", + ) + assert res2["access_token"] == "ya29.refreshed_access" + assert res2["id_token"] == "eyJhbGciOiJSUzI1NiJ9.preserved_jwt.sig" + assert res2["scope"] == "openid https://www.googleapis.com/auth/userinfo.email" + + +@pytest.mark.unit +def test_refresh_if_needed_passes_existing_id_and_scope(): + """P0-1: refresh_if_needed supplies existing id_token and scope to refresh.""" + creds = { + "refresh_token": "1//test_refresh", + "access_token": "ya29.old_token", + "id_token": "eyJhbGci.existing_jwt.sig", + "scope": "openid email", + "expires_at": time.time() - 100, # Expired + } + mock_resp = { + "access_token": "ya29.refreshed_token", + "expires_in": 3600, + } + with patch("antigravity_provider.oauth._post_form_json", return_value=mock_resp): + updated = refresh_if_needed(creds) + + assert updated["access_token"] == "ya29.refreshed_token" + assert updated["id_token"] == "eyJhbGci.existing_jwt.sig" + assert updated["scope"] == "openid email" + + +# ── TEST P0-2: Writing .gemini/oauth_creds.json in Profile Directory ── + + +@pytest.mark.unit +def test_save_profile_auth_creates_oauth_creds_json_with_6_fields(tmp_path, monkeypatch): + """P0-2: Saving an Antigravity profile creates .gemini/oauth_creds.json with exact 6 fields.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + + profile_id = "ag-test-profile" + auth_data = { + "token": { + "access_token": "ya29.a0AfH6SM...", + "refresh_token": "1//0gK9...", + "id_token": "eyJhbGciOiJSUzI1NiIsImtpZCI...", + "scope": "openid https://www.googleapis.com/auth/userinfo.email", + "token_type": "Bearer", + "expires_at": 1786634144.0, + "expiry": "2026-08-24T12:00:00Z", + }, + "email": "test.user@gmail.com", + "auth_method": "oauth", + } + + saved_path = ProfileAuthManager.save_profile_auth("antigravity", profile_id, auth_data) + assert saved_path.is_file() + + pdir = get_profile_dir(profile_id, "antigravity") + oauth_creds_path = pdir / ".gemini" / "oauth_creds.json" + assert oauth_creds_path.is_file() + + creds_content = json.loads(oauth_creds_path.read_text(encoding="utf-8")) + + # Exact 6 fields + expected_keys = {"access_token", "refresh_token", "scope", "token_type", "id_token", "expiry_date"} + assert set(creds_content.keys()) == expected_keys + assert creds_content["access_token"] == "ya29.a0AfH6SM..." + assert creds_content["refresh_token"] == "1//0gK9..." + assert creds_content["id_token"] == "eyJhbGciOiJSUzI1NiIsImtpZCI..." + assert creds_content["scope"] == "openid https://www.googleapis.com/auth/userinfo.email" + assert creds_content["token_type"] == "Bearer" + + # expiry_date must be integer milliseconds > 1e12 + assert isinstance(creds_content["expiry_date"], int) + assert creds_content["expiry_date"] > 1000000000000 + assert creds_content["expiry_date"] == int(1786634144.0 * 1000) + + +@pytest.mark.unit +def test_oauth_session_callback_persists_full_credentials(tmp_path, monkeypatch): + """P0-2: Full OAuth session flow populates id_token, scope, and creates oauth_creds.json.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + + session = ProfileOAuthSession("ag-callback-test") + session.verifier = "test_verifier" + session.state = "test_state" + + mock_tokens = { + "access_token": "ya29.callback_access", + "refresh_token": "1//callback_refresh", + "id_token": "eyJhbGciOiJSUzI1NiJ9.test_id_token_jwt.sig", + "scope": "openid email profile", + "token_type": "Bearer", + "expires_at": int(time.time()) + 3600, + } + + with patch("antigravity_provider.router.profile_oauth.exchange_code_for_tokens", return_value=mock_tokens), \ + patch("antigravity_provider.router.profile_oauth.fetch_user_email", return_value="callback_user@gmail.com"): + + ok, msg = session.handle_callback("test_code", "test_state", source="test") + assert ok is True + + pdir = get_profile_dir("ag-callback-test", "antigravity") + auth_json = json.loads((pdir / "auth.json").read_text(encoding="utf-8")) + assert auth_json["token"]["id_token"] == mock_tokens["id_token"] + assert auth_json["token"]["scope"] == "openid email profile" + + oauth_creds = json.loads((pdir / ".gemini" / "oauth_creds.json").read_text(encoding="utf-8")) + assert oauth_creds["id_token"] == mock_tokens["id_token"] + assert oauth_creds["scope"] == "openid email profile" + assert isinstance(oauth_creds["expiry_date"], int) + assert oauth_creds["expiry_date"] > 1e12 + + +# ── TEST P1-4: Model Discovery Caching & Resilience ── + + +@pytest.mark.unit +def test_model_discovery_cache_persistence(tmp_path): + """P1-4: Models are persisted to models_cache.json and survive service re-creation.""" + cache_file = tmp_path / "models_cache.json" + service1 = ModelDiscoveryService(cache_path=cache_file) + + # Initially empty + assert service1.get_models("antigravity") is None + meta = service1.get_models_with_metadata("antigravity") + assert meta["models"] is None + assert meta["has_cache"] is False + + # Simulate discovery + mock_models = ["gemini-3.7-flash", "gemini-2.5-pro"] + with service1._cache_lock: + service1._cache["antigravity"] = { + "models": mock_models, + "discovered_at": time.time(), + } + service1._save_cache_to_disk() + + assert cache_file.is_file() + + # Re-instantiate service reading the same file + service2 = ModelDiscoveryService(cache_path=cache_file) + cached = service2.get_models("antigravity") + assert cached == mock_models + meta2 = service2.get_models_with_metadata("antigravity") + assert meta2["has_cache"] is True + assert meta2["is_stale"] is False + + +@pytest.mark.unit +def test_model_discovery_failure_preserves_existing_cache(tmp_path): + """P1-4: Background discovery failure/timeout does NOT wipe existing cache.""" + cache_file = tmp_path / "models_cache.json" + service = ModelDiscoveryService(cache_path=cache_file) + + initial_models = ["gemini-3.7-flash"] + with service._cache_lock: + service._cache["antigravity"] = { + "models": initial_models, + "discovered_at": time.time() - 7200, # Stale + } + service._save_cache_to_disk() + + # Probe fails (returns None) + with patch.object(service, "_probe_provider", return_value=None): + res = service.discover_models_sync("antigravity", timeout=1.0) + # Retains existing cache + assert res == initial_models + + # Cache must still contain initial_models + assert service.get_models("antigravity") == initial_models