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
This commit is contained in:
parent
972e34911c
commit
bb4f6df67a
7 changed files with 513 additions and 26 deletions
|
|
@ -123,10 +123,25 @@ def discover_models(profile_id: str | None = None) -> dict[str, str]:
|
||||||
# команда запускалась в ГЛОБАЛЬНОМ окружении, где вход не выполнен, и
|
# команда запускалась в ГЛОБАЛЬНОМ окружении, где вход не выполнен, и
|
||||||
# отвечала «Please sign in to view available models» — при шести рабочих
|
# отвечала «Please sign in to view available models» — при шести рабочих
|
||||||
# OAuth-профилях. Список моделей поэтому был пуст всегда.
|
# 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
|
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(
|
env = build_safe_subprocess_env(
|
||||||
overrides={
|
overrides={
|
||||||
"USERPROFILE": str(profile_dir),
|
"USERPROFILE": str(profile_dir),
|
||||||
|
|
@ -146,6 +161,7 @@ def discover_models(profile_id: str | None = None) -> dict[str, str]:
|
||||||
encoding="utf-8",
|
encoding="utf-8",
|
||||||
errors="replace",
|
errors="replace",
|
||||||
env=env,
|
env=env,
|
||||||
|
stdin=subprocess.DEVNULL,
|
||||||
)
|
)
|
||||||
raw = result.stdout.strip()
|
raw = result.stdout.strip()
|
||||||
if not raw or result.returncode != 0:
|
if not raw or result.returncode != 0:
|
||||||
|
|
|
||||||
|
|
@ -28,6 +28,7 @@ TOKEN_URL = "https://oauth2.googleapis.com/token"
|
||||||
CLIENT_ID = "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com"
|
CLIENT_ID = "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com"
|
||||||
CLIENT_SECRET = "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"
|
CLIENT_SECRET = "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"
|
||||||
SCOPES = [
|
SCOPES = [
|
||||||
|
"openid",
|
||||||
"https://www.googleapis.com/auth/cloud-platform",
|
"https://www.googleapis.com/auth/cloud-platform",
|
||||||
"https://www.googleapis.com/auth/userinfo.email",
|
"https://www.googleapis.com/auth/userinfo.email",
|
||||||
"https://www.googleapis.com/auth/userinfo.profile",
|
"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(
|
def refresh_access_token(
|
||||||
refresh_token: str,
|
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,
|
post_json: Callable[[str, dict[str, str], dict[str, str]], dict[str, Any]] | None = None,
|
||||||
client: tuple[str, str] | None = None,
|
client: tuple[str, str] | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
|
|
@ -85,9 +88,16 @@ def refresh_access_token(
|
||||||
data = post_json(TOKEN_URL, payload, {"Content-Type": "application/x-www-form-urlencoded"})
|
data = post_json(TOKEN_URL, payload, {"Content-Type": "application/x-www-form-urlencoded"})
|
||||||
if not data.get("access_token"):
|
if not data.get("access_token"):
|
||||||
raise ProxyError("OAuth refresh response did not include access_token", status=401, error_type="invalid_request_error")
|
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 {
|
return {
|
||||||
"refresh_token": data.get("refresh_token") or refresh_token,
|
"refresh_token": data.get("refresh_token") or refresh_token,
|
||||||
"access_token": data["access_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")),
|
"expires_at": _expires_at(data.get("expires_in")),
|
||||||
"token_type": data.get("token_type", "Bearer"),
|
"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")
|
refresh = credentials.get("refresh_token") or credentials.get("refresh")
|
||||||
expires = credentials.get("expires_at") or credentials.get("expires")
|
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))):
|
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
|
return credentials
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -177,6 +194,9 @@ def exchange_code_for_tokens(
|
||||||
return {
|
return {
|
||||||
"refresh_token": refresh_token,
|
"refresh_token": refresh_token,
|
||||||
"access_token": data["access_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")),
|
"expires_at": _expires_at(data.get("expires_in")),
|
||||||
"token_type": data.get("token_type", "Bearer"),
|
"token_type": data.get("token_type", "Bearer"),
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -62,12 +62,15 @@ class AntigravityAdapter(BaseProviderAdapter):
|
||||||
|
|
||||||
if profile_auth:
|
if profile_auth:
|
||||||
# Pre-flight check: verify token expiry before calling subprocess to prevent interactive browser login
|
# Pre-flight check: verify token expiry before calling subprocess to prevent interactive browser login
|
||||||
tokens = profile_auth.get("tokens", {})
|
tokens = profile_auth.get("token") or profile_auth.get("tokens", {})
|
||||||
expiry = tokens.get("expiry_date") or profile_auth.get("expiry_date")
|
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:
|
||||||
if expiry > 1e11:
|
if float(expiry) > 1e11:
|
||||||
expiry = expiry / 1000.0
|
expiry = float(expiry) / 1000.0
|
||||||
if time.time() > expiry:
|
if time.time() > float(expiry) and not refresh_tok:
|
||||||
raise AuthExpiredError(
|
raise AuthExpiredError(
|
||||||
"Авторизация истекла, требуется повторный вход.",
|
"Авторизация истекла, требуется повторный вход.",
|
||||||
provider="antigravity",
|
provider="antigravity",
|
||||||
|
|
|
||||||
|
|
@ -208,7 +208,10 @@ class ModelDiscoveryService:
|
||||||
logger.info("Discovered %d models for provider '%s': %s", len(models), provider, models)
|
logger.info("Discovered %d models for provider '%s': %s", len(models), provider, models)
|
||||||
return 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]]:
|
def _probe_provider(self, provider: str) -> Optional[List[str]]:
|
||||||
"""Perform provider-specific model discovery."""
|
"""Perform provider-specific model discovery."""
|
||||||
|
|
@ -217,7 +220,8 @@ class ModelDiscoveryService:
|
||||||
|
|
||||||
if prov == "antigravity":
|
if prov == "antigravity":
|
||||||
from antigravity_provider.agy_subprocess import discover_models
|
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:
|
if res:
|
||||||
return sorted(list(set(res.values())))
|
return sorted(list(set(res.values())))
|
||||||
return None
|
return None
|
||||||
|
|
|
||||||
|
|
@ -15,6 +15,7 @@ import threading
|
||||||
import time
|
import time
|
||||||
import urllib.request
|
import urllib.request
|
||||||
import urllib.error
|
import urllib.error
|
||||||
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
|
@ -113,6 +114,73 @@ class ProfileAuthManager:
|
||||||
"""Official API to get isolated directory for a profile."""
|
"""Official API to get isolated directory for a profile."""
|
||||||
return get_profile_dir(profile_id, provider)
|
return get_profile_dir(profile_id, provider)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def write_agy_oauth_creds(cls, profile_dir: Path, auth_data: dict) -> Path:
|
||||||
|
"""Atomically write <profile_dir>/.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
|
@staticmethod
|
||||||
def read_windows_credential(target_name: str = "gemini:antigravity") -> Optional[dict]:
|
def read_windows_credential(target_name: str = "gemini:antigravity") -> Optional[dict]:
|
||||||
"""Read a credential blob from Windows Credential Manager."""
|
"""Read a credential blob from Windows Credential Manager."""
|
||||||
|
|
@ -140,7 +208,40 @@ class ProfileAuthManager:
|
||||||
if not advapi32 or os.name != "nt":
|
if not advapi32 or os.name != "nt":
|
||||||
return False
|
return False
|
||||||
with _CM_LOCK:
|
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)
|
buf = ctypes.create_string_buffer(blob_bytes)
|
||||||
cred = CREDENTIAL()
|
cred = CREDENTIAL()
|
||||||
cred.Flags = 0
|
cred.Flags = 0
|
||||||
|
|
@ -201,9 +302,24 @@ class ProfileAuthManager:
|
||||||
pdir.mkdir(parents=True, exist_ok=True)
|
pdir.mkdir(parents=True, exist_ok=True)
|
||||||
auth_file = pdir / "auth.json"
|
auth_file = pdir / "auth.json"
|
||||||
existed = auth_file.is_file()
|
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")
|
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)
|
os.replace(temp_file, auth_file)
|
||||||
|
try:
|
||||||
|
os.chmod(auth_file, 0o600)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# For Antigravity, synchronously maintain <profile_dir>/.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 (
|
from antigravity_provider.router.event_bus import (
|
||||||
EVENT_ACCOUNT_ADDED,
|
EVENT_ACCOUNT_ADDED,
|
||||||
|
|
@ -238,7 +354,58 @@ class ProfileAuthManager:
|
||||||
auth_file = get_profile_auth_path(provider, profile_id)
|
auth_file = get_profile_auth_path(provider, profile_id)
|
||||||
if auth_file.is_file():
|
if auth_file.is_file():
|
||||||
try:
|
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:
|
except Exception as e:
|
||||||
logger.warning("Error reading %s: %s", auth_file, e)
|
logger.warning("Error reading %s: %s", auth_file, e)
|
||||||
|
|
||||||
|
|
@ -406,22 +573,38 @@ class ProfileAuthManager:
|
||||||
"error": None,
|
"error": None,
|
||||||
}
|
}
|
||||||
|
|
||||||
if provider == "antigravity":
|
if provider in ("antigravity", "google-antigravity"):
|
||||||
tokens = auth_data.get("tokens", {})
|
tokens = auth_data.get("token") or auth_data.get("tokens", {})
|
||||||
acc_token = tokens.get("access_token") or auth_data.get("access_token")
|
acc_token = tokens.get("access_token") if isinstance(tokens, dict) else (auth_data.get("access_token") or "")
|
||||||
id_token = tokens.get("id_token") or auth_data.get("id_token")
|
id_token = tokens.get("id_token") if isinstance(tokens, dict) else (auth_data.get("id_token") or "")
|
||||||
email = None
|
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
|
acc_id = None
|
||||||
if id_token:
|
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
|
is_expired = False
|
||||||
if expiry:
|
if expiry:
|
||||||
if expiry > 1e11:
|
if float(expiry) > 1e11:
|
||||||
expiry = expiry / 1000.0
|
expiry = float(expiry) / 1000.0
|
||||||
if time.time() > expiry:
|
if time.time() > float(expiry):
|
||||||
is_expired = True
|
is_expired = not bool(refresh_tok)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"authenticated": True,
|
"authenticated": True,
|
||||||
|
|
@ -429,9 +612,10 @@ class ProfileAuthManager:
|
||||||
"profile_id": profile_id,
|
"profile_id": profile_id,
|
||||||
"email_masked": mask_email(email) if email else None,
|
"email_masked": mask_email(email) if email else None,
|
||||||
"account_id_masked": mask_id(acc_id) if acc_id else None,
|
"account_id_masked": mask_id(acc_id) if acc_id else None,
|
||||||
|
"has_refresh_token": bool(refresh_tok),
|
||||||
"is_expired": is_expired,
|
"is_expired": is_expired,
|
||||||
"status": "EXPIRED" if is_expired else "AUTHENTICATED",
|
"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"):
|
elif provider in ("openai-codex", "codex"):
|
||||||
|
|
|
||||||
|
|
@ -238,11 +238,16 @@ class ProfileOAuthSession:
|
||||||
logger.info("OAuth account identity resolved (email_found=%s)", bool(email))
|
logger.info("OAuth account identity resolved (email_found=%s)", bool(email))
|
||||||
|
|
||||||
# Format in standard gemini:antigravity shape
|
# Format in standard gemini:antigravity shape
|
||||||
|
expires_at = tokens.get("expires_at") or (int(time.time()) + 3600)
|
||||||
auth_data = {
|
auth_data = {
|
||||||
"token": {
|
"token": {
|
||||||
"access_token": tokens["access_token"],
|
"access_token": tokens["access_token"],
|
||||||
"refresh_token": tokens["refresh_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 "",
|
"email": email or "",
|
||||||
"auth_method": "oauth",
|
"auth_method": "oauth",
|
||||||
|
|
|
||||||
255
tests/test_agy_oauth_credentials.py
Normal file
255
tests/test_agy_oauth_credentials.py
Normal file
|
|
@ -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
|
||||||
Loading…
Reference in a new issue