누락된 OAuth 콜백 코드 파싱과 계정 metadata 헬퍼 구현
Some checks failed
API contract / OpenAPI type drift (push) Failing after 48s

This commit is contained in:
Yun Chan 2026-09-23 18:42:19 +09:00
parent 7dce818788
commit fefca4743e
2 changed files with 252 additions and 0 deletions

View file

@ -14,6 +14,7 @@ state는 서버가 발급해 15분 보관하며, 교환 시 provider·만료를
from __future__ import annotations from __future__ import annotations
import asyncio
import base64 import base64
import hashlib import hashlib
import json import json
@ -21,6 +22,7 @@ import secrets
import time import time
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any from typing import Any
from urllib.parse import parse_qs, urlsplit
import httpx import httpx
@ -31,6 +33,9 @@ from .provider_credentials import (
) )
OAUTH_PENDING_TTL_SECONDS = 900.0 OAUTH_PENDING_TTL_SECONDS = 900.0
# Code Assist 온보딩(onboardUser)은 장기 실행 작업이라 완료까지 짧게 폴링한다.
_ONBOARD_POLL_ATTEMPTS = 10
_ONBOARD_POLL_INTERVAL_SECONDS = 2.0
# claude: setup-token과 동일한 공개 클라이언트. 콜백 페이지가 code#state를 # claude: setup-token과 동일한 공개 클라이언트. 콜백 페이지가 code#state를
# 화면에 표시하는 수동 흐름이라 사전 등록·로컬 콜백 서버가 필요 없다. # 화면에 표시하는 수동 흐름이라 사전 등록·로컬 콜백 서버가 필요 없다.
@ -78,6 +83,7 @@ _AGY_OAUTH = {
), ),
"user_agent": "antigravity/cli/1.0.0 (aidev_client; os_type=windows; arch=amd64; auth_method=consumer)", "user_agent": "antigravity/cli/1.0.0 (aidev_client; os_type=windows; arch=amd64; auth_method=consumer)",
"code_assist_base": "https://cloudcode-pa.googleapis.com", "code_assist_base": "https://cloudcode-pa.googleapis.com",
"plugin_type": "GEMINI",
} }
OAUTH_PROVIDERS = ("claude", "openrouter", "codex", "agy") OAUTH_PROVIDERS = ("claude", "openrouter", "codex", "agy")
@ -180,6 +186,127 @@ def start_oauth(provider: str, admin_email: str) -> dict[str, Any]:
} }
def _extract_code_from_paste(pasted: str) -> str:
"""붙여넣은 값에서 authorization code를 꺼낸다.
브라우저 주소창의 콜백 URL 전체와 코드만 복사한 값 모두를 받는다.
"""
value = pasted.strip().strip('"').strip("'")
if not value:
raise ProviderCredentialError("인증 코드가 비어 있습니다.")
if "code=" in value or "://" in value or value.startswith("?"):
query = urlsplit(value).query or value.split("?", 1)[-1]
candidates = parse_qs(query).get("code") or []
if not candidates:
raise ProviderCredentialError(
"붙여넣은 값에서 code를 찾지 못했습니다. 브라우저 주소창의 콜백 주소 전체를 붙여넣어 주세요."
)
value = candidates[0].strip()
if not value:
raise ProviderCredentialError("인증 코드가 비어 있습니다.")
return value
def _chatgpt_account_id(id_token: str | None) -> str | None:
"""id_token payload에서 ChatGPT 계정 id를 꺼낸다.
토큰은 방금 TLS로 받은 응답 본문이라 서명 검증 대신 그 신뢰를 그대로 쓴다.
"""
if not id_token:
return None
parts = str(id_token).split(".")
if len(parts) < 2:
return None
payload = parts[1]
try:
decoded = base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4))
claims = json.loads(decoded)
except (ValueError, TypeError):
return None
if not isinstance(claims, dict):
return None
namespace = claims.get("https://api.openai.com/auth")
if not isinstance(namespace, dict):
return None
account_id = namespace.get("chatgpt_account_id")
return str(account_id).strip() if account_id else None
def _code_assist_headers(token: str) -> dict[str, str]:
return {
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
"User-Agent": str(_AGY_OAUTH["user_agent"]),
}
async def _code_assist_post(
client: httpx.AsyncClient, method: str, token: str, payload: dict[str, Any]
) -> dict[str, Any]:
response = await client.post(
f"{_AGY_OAUTH['code_assist_base']}/v1internal:{method}",
json=payload,
headers=_code_assist_headers(token),
)
response.raise_for_status()
body = response.json()
if not isinstance(body, dict):
raise ValueError("Code Assist 응답이 객체가 아닙니다")
return body
def _default_tier_id(loaded: dict[str, Any]) -> str:
"""loadCodeAssist 응답에서 온보딩에 쓸 기본 등급 id를 고른다."""
tiers = loaded.get("allowedTiers")
if isinstance(tiers, list):
for tier in tiers:
if isinstance(tier, dict) and tier.get("isDefault"):
tier_id = str(tier.get("id") or "").strip()
if tier_id:
return tier_id
current = loaded.get("currentTier")
if isinstance(current, dict):
return str(current.get("id") or "").strip()
return ""
async def _antigravity_project_id(token: str) -> str:
"""Code Assist 프로젝트 id를 확보한다. 있으면 재사용하고 없으면 온보딩한다.
Antigravity 생성 요청이 이 project를 요구하므로 확정하지 못하면 연결을 실패시킨다.
"""
metadata = {"pluginType": str(_AGY_OAUTH["plugin_type"])}
try:
async with httpx.AsyncClient(timeout=30) as client:
loaded = await _code_assist_post(client, "loadCodeAssist", token, {"metadata": metadata})
existing = str(loaded.get("cloudaicompanionProject") or "").strip()
if existing:
return existing
tier_id = _default_tier_id(loaded)
if not tier_id:
raise ProviderCredentialError(
"Google 계정에 Code Assist 사용 등급이 없습니다. Antigravity에서 계정 상태를 확인해 주세요."
)
payload = {"tierId": tier_id, "metadata": metadata}
for _attempt in range(_ONBOARD_POLL_ATTEMPTS):
onboarding = await _code_assist_post(client, "onboardUser", token, payload)
if onboarding.get("done"):
created = ""
if isinstance(onboarding.get("response"), dict):
created = str(
(onboarding["response"] or {}).get("cloudaicompanionProject") or ""
).strip()
if created:
return created
break
await asyncio.sleep(_ONBOARD_POLL_INTERVAL_SECONDS)
except ProviderCredentialError:
raise
except (httpx.HTTPError, ValueError) as exc:
raise ProviderCredentialError(f"Agy(Google) Code Assist 프로젝트 확인 실패: {exc}") from exc
raise ProviderCredentialError("Google 계정의 Code Assist 프로젝트를 확정하지 못했습니다.")
def _pop_pending(provider: str, state: str) -> _PendingOAuth: def _pop_pending(provider: str, state: str) -> _PendingOAuth:
_purge_expired() _purge_expired()
attempt = _PENDING.get(state) attempt = _PENDING.get(state)

View file

@ -1,5 +1,7 @@
"""provider OAuth(PKCE) 서비스 단위 테스트 — 실제 네트워크 호출 없음.""" """provider OAuth(PKCE) 서비스 단위 테스트 — 실제 네트워크 호출 없음."""
import base64
import json
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
@ -7,6 +9,11 @@ from unittest.mock import patch
from app.services import provider_oauth as svc from app.services import provider_oauth as svc
def _b64url(claims: dict) -> str:
raw = json.dumps(claims).encode("utf-8")
return base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=")
class _FakeResponse: class _FakeResponse:
def __init__(self, payload, status_code=200): def __init__(self, payload, status_code=200):
self._payload = payload self._payload = payload
@ -38,6 +45,25 @@ class _FakeAsyncClient:
return _FakeResponse(self._payload) return _FakeResponse(self._payload)
class _SequencedAsyncClient:
"""호출 순서대로 준비한 응답을 돌려주고 호출을 기록하는 가짜 클라이언트."""
def __init__(self, payloads):
self._payloads = list(payloads)
self.calls = []
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
return None
async def post(self, url, **kwargs):
self.calls.append({"url": url, **kwargs})
payload = self._payloads.pop(0) if self._payloads else {}
return _FakeResponse(payload)
class StartOAuthTest(unittest.TestCase): class StartOAuthTest(unittest.TestCase):
def setUp(self): def setUp(self):
svc._PENDING.clear() svc._PENDING.clear()
@ -94,6 +120,52 @@ class StartOAuthTest(unittest.TestCase):
svc._pop_pending("claude", result["state"]) svc._pop_pending("claude", result["state"])
class PasteAndMetadataTest(unittest.TestCase):
"""콜백 붙여넣기 파싱과 제공자 metadata 추출 헬퍼."""
def test_extract_code_from_callback_url(self):
url = (
"http://localhost:1455/auth/callback?state=s1&iss=https://accounts.google.com"
"&code=4/0AXqlp5_x-3&scope=email%20profile"
)
self.assertEqual(svc._extract_code_from_paste(url), "4/0AXqlp5_x-3")
def test_extract_code_accepts_bare_and_quoted_values(self):
self.assertEqual(svc._extract_code_from_paste(" 4/0AX "), "4/0AX")
self.assertEqual(svc._extract_code_from_paste('"4/0AX"'), "4/0AX")
def test_extract_code_percent_decodes(self):
self.assertEqual(svc._extract_code_from_paste("code=4%2F0A%20x&state=s"), "4/0A x")
def test_extract_code_rejects_useless_paste(self):
for value in ("", " ", "http://localhost:1455/auth/callback?state=only"):
with self.assertRaises(svc.ProviderCredentialError):
svc._extract_code_from_paste(value)
def test_chatgpt_account_id_reads_namespace_claim(self):
claims = {"https://api.openai.com/auth": {"chatgpt_account_id": "acct-1"}}
token = "header." + _b64url(claims) + ".signature"
self.assertEqual(svc._chatgpt_account_id(token), "acct-1")
def test_chatgpt_account_id_returns_none_without_claim(self):
self.assertIsNone(svc._chatgpt_account_id(None))
self.assertIsNone(svc._chatgpt_account_id("not-a-jwt"))
self.assertIsNone(svc._chatgpt_account_id("h." + _b64url({"sub": "x"}) + ".s"))
def test_default_tier_prefers_is_default_flag(self):
loaded = {
"allowedTiers": [{"id": "free-tier"}, {"id": "paid-tier", "isDefault": True}],
"currentTier": {"id": "current-tier"},
}
self.assertEqual(svc._default_tier_id(loaded), "paid-tier")
def test_default_tier_falls_back_to_current_tier(self):
self.assertEqual(
svc._default_tier_id({"currentTier": {"id": "current-tier"}}), "current-tier"
)
self.assertEqual(svc._default_tier_id({}), "")
class FinishOAuthTest(unittest.IsolatedAsyncioTestCase): class FinishOAuthTest(unittest.IsolatedAsyncioTestCase):
def setUp(self): def setUp(self):
svc._PENDING.clear() svc._PENDING.clear()
@ -147,5 +219,58 @@ class FinishOAuthTest(unittest.IsolatedAsyncioTestCase):
self.assertEqual(fake.captured["json"]["code_challenge_method"], "S256") self.assertEqual(fake.captured["json"]["code_challenge_method"], "S256")
async def test_antigravity_project_id_reuses_existing_project(self):
fake = _SequencedAsyncClient([{"cloudaicompanionProject": "proj-existing"}])
with patch.object(svc.httpx, "AsyncClient", return_value=fake):
self.assertEqual(await svc._antigravity_project_id("ya29-token"), "proj-existing")
self.assertEqual(len(fake.calls), 1)
self.assertEqual(
fake.calls[0]["url"], "https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist"
)
async def test_antigravity_project_id_requires_a_tier(self):
fake = _SequencedAsyncClient([{}])
with patch.object(svc.httpx, "AsyncClient", return_value=fake):
with self.assertRaises(svc.ProviderCredentialError):
await svc._antigravity_project_id("ya29-token")
async def test_agy_callback_url_paste_onboards_code_assist_project(self):
start = svc.start_oauth("agy", "admin@example.com")
fake = _SequencedAsyncClient(
[
{"access_token": "ya29-token", "refresh_token": "1//refresh"},
{"allowedTiers": [{"id": "legacy-tier", "isDefault": True}]},
{"done": True, "response": {"cloudaicompanionProject": "proj-123"}},
]
)
with patch.object(svc.httpx, "AsyncClient", return_value=fake), patch.object(
svc, "save_credential"
) as save, patch.object(svc, "push_credentials_to_gateway") as push:
save.return_value = SimpleNamespace(
provider="agy", token_hint="…ken", auth_kind="oauth_token",
updated_by="admin@example.com", updated_at=1.0,
)
push.return_value = {"synced": True}
callback = (
f"http://localhost:1455/auth/callback?state={start['state']}"
"&code=4/0AXqlp5_x-3&scope=email%20profile"
)
result = await svc.finish_oauth("agy", callback, start["state"], "admin@example.com")
self.assertEqual(result["stored"].auth_kind, "oauth_token")
self.assertEqual(fake.calls[0]["url"], svc._AGY_OAUTH["token_url"])
self.assertEqual(fake.calls[0]["data"]["code"], "4/0AXqlp5_x-3")
self.assertEqual(
fake.calls[1]["url"], "https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist"
)
self.assertEqual(
fake.calls[2]["url"], "https://cloudcode-pa.googleapis.com/v1internal:onboardUser"
)
self.assertEqual(fake.calls[2]["json"]["tierId"], "legacy-tier")
self.assertEqual(save.call_args.kwargs["token"], "ya29-token")
self.assertEqual(save.call_args.kwargs["refresh_token"], "1//refresh")
self.assertEqual(save.call_args.kwargs["extra"], {"project_id": "proj-123"})
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()