diff --git a/apps/api/app/services/provider_oauth.py b/apps/api/app/services/provider_oauth.py index 9ceae97..fd2f389 100644 --- a/apps/api/app/services/provider_oauth.py +++ b/apps/api/app/services/provider_oauth.py @@ -14,6 +14,7 @@ state는 서버가 발급해 15분 보관하며, 교환 시 provider·만료를 from __future__ import annotations +import asyncio import base64 import hashlib import json @@ -21,6 +22,7 @@ import secrets import time from dataclasses import dataclass from typing import Any +from urllib.parse import parse_qs, urlsplit import httpx @@ -31,6 +33,9 @@ from .provider_credentials import ( ) OAUTH_PENDING_TTL_SECONDS = 900.0 +# Code Assist 온보딩(onboardUser)은 장기 실행 작업이라 완료까지 짧게 폴링한다. +_ONBOARD_POLL_ATTEMPTS = 10 +_ONBOARD_POLL_INTERVAL_SECONDS = 2.0 # 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)", "code_assist_base": "https://cloudcode-pa.googleapis.com", + "plugin_type": "GEMINI", } 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: _purge_expired() attempt = _PENDING.get(state) diff --git a/apps/api/app/test_provider_oauth.py b/apps/api/app/test_provider_oauth.py index c6ee527..b58f78d 100644 --- a/apps/api/app/test_provider_oauth.py +++ b/apps/api/app/test_provider_oauth.py @@ -1,5 +1,7 @@ """provider OAuth(PKCE) 서비스 단위 테스트 — 실제 네트워크 호출 없음.""" +import base64 +import json import unittest from types import SimpleNamespace from unittest.mock import patch @@ -7,6 +9,11 @@ from unittest.mock import patch 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: def __init__(self, payload, status_code=200): self._payload = payload @@ -38,6 +45,25 @@ class _FakeAsyncClient: 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): def setUp(self): svc._PENDING.clear() @@ -94,6 +120,52 @@ class StartOAuthTest(unittest.TestCase): 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): def setUp(self): svc._PENDING.clear() @@ -147,5 +219,58 @@ class FinishOAuthTest(unittest.IsolatedAsyncioTestCase): 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__": unittest.main()