vignette/apps/api/app/test_provider_oauth.py
Yun Chan fefca4743e
Some checks failed
API contract / OpenAPI type drift (push) Failing after 48s
누락된 OAuth 콜백 코드 파싱과 계정 metadata 헬퍼 구현
2026-09-23 18:42:19 +09:00

276 lines
12 KiB
Python

"""provider OAuth(PKCE) 서비스 단위 테스트 — 실제 네트워크 호출 없음."""
import base64
import json
import unittest
from types import SimpleNamespace
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
self.status_code = status_code
def raise_for_status(self):
if self.status_code >= 400:
import httpx
raise httpx.HTTPStatusError("err", request=None, response=None)
def json(self):
return self._payload
class _FakeAsyncClient:
def __init__(self, payload):
self._payload = payload
self.captured = {}
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
return None
async def post(self, url, **kwargs):
self.captured = {"url": url, **kwargs}
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()
def tearDown(self):
svc._PENDING.clear()
def test_unsupported_provider_rejected(self):
with self.assertRaises(svc.ProviderCredentialError):
svc.start_oauth("solar", "admin@example.com")
def test_claude_authorize_url_contains_pkce_and_state(self):
result = svc.start_oauth("claude", "admin@example.com")
self.assertIn("https://claude.ai/oauth/authorize", result["authorize_url"])
self.assertIn("code_challenge_method=S256", result["authorize_url"])
self.assertIn(f"&state={result['state']}", result["authorize_url"])
self.assertIn("user%3Ainference", result["authorize_url"])
self.assertIn(result["state"], svc._PENDING)
def test_openrouter_authorize_url_is_headless(self):
result = svc.start_oauth("openrouter", "admin@example.com")
self.assertIn("https://openrouter.ai/auth?", result["authorize_url"])
self.assertIn("code_challenge_method=S256", result["authorize_url"])
self.assertNotIn("callback_url", result["authorize_url"])
def test_agy_authorize_url_uses_absolute_scope_urls(self):
# 구글은 축약형(auth/...) 스코프를 인식하지 못해 invalid_scope로 거부한다.
result = svc.start_oauth("agy", "admin@example.com")
authorize_url = result["authorize_url"]
self.assertIn("https://accounts.google.com/o/oauth2/v2/auth?", authorize_url)
for scope in (
"cloud-platform",
"userinfo.email",
"userinfo.profile",
"cclog",
"experimentsandconfigs",
):
self.assertIn(f"https://www.googleapis.com/auth/{scope}", authorize_url)
self.assertNotIn("scope=auth/", authorize_url)
def test_pending_expires(self):
import time
result = svc.start_oauth("claude", "admin@example.com")
attempt = svc._PENDING[result["state"]]
stale = svc._PendingOAuth(
provider=attempt.provider,
code_verifier=attempt.code_verifier,
created_by=attempt.created_by,
created_at=time.time() - svc.OAUTH_PENDING_TTL_SECONDS - 1,
)
svc._PENDING[result["state"]] = stale
with self.assertRaises(svc.ProviderCredentialError):
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()
def tearDown(self):
svc._PENDING.clear()
async def test_claude_code_state_exchange_and_save(self):
start = svc.start_oauth("claude", "admin@example.com")
fake = _FakeAsyncClient({"access_token": "sk-ant-oat-token-abc", "expires_in": 31536000})
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="claude", token_hint="…oat", auth_kind="oauth_token",
updated_by="admin@example.com", updated_at=1.0,
)
push.return_value = {"synced": True}
result = await svc.finish_oauth(
"claude", f"auth-code-123#{start['state']}", start["state"], "admin@example.com"
)
self.assertEqual(result["stored"].auth_kind, "oauth_token")
sent = fake.captured["json"]
self.assertEqual(sent["grant_type"], "authorization_code")
self.assertEqual(sent["code"], "auth-code-123")
self.assertIn("code_verifier", sent)
self.assertEqual(fake.captured["url"], svc._CLAUDE_OAUTH["token_url"])
async def test_state_mismatch_rejected(self):
start = svc.start_oauth("claude", "admin@example.com")
with self.assertRaises(svc.ProviderCredentialError):
await svc.finish_oauth("claude", "code#other-state", start["state"], "a@b.com")
async def test_openrouter_code_exchange_returns_api_key(self):
start = svc.start_oauth("openrouter", "admin@example.com")
fake = _FakeAsyncClient({"key": "sk-or-v1-oauth-key"})
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="openrouter", token_hint="…key", auth_kind="api_key",
updated_by="admin@example.com", updated_at=1.0,
)
push.return_value = {"synced": True}
result = await svc.finish_oauth(
"openrouter", "or-code-xyz", start["state"], "admin@example.com"
)
self.assertEqual(result["stored"].auth_kind, "api_key")
self.assertEqual(fake.captured["url"], svc._OPENROUTER_OAUTH["exchange_url"])
self.assertEqual(fake.captured["json"]["code"], "or-code-xyz")
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()