125 lines
4.5 KiB
Python
125 lines
4.5 KiB
Python
"""provider 자격증명 서비스(암호화·카탈로그) 단위 테스트."""
|
|
|
|
import unittest
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
from pydantic import SecretStr
|
|
|
|
from app.services import provider_credentials as svc
|
|
|
|
|
|
class TokenEncryptionTest(unittest.TestCase):
|
|
def setUp(self):
|
|
patcher = patch.object(
|
|
svc.settings,
|
|
"provider_credential_secret",
|
|
SecretStr("unit-test-master-secret-0123456789"),
|
|
)
|
|
patcher.start()
|
|
self.addCleanup(patcher.stop)
|
|
|
|
def test_roundtrip(self):
|
|
plain = "sk-or-v1-아주-긴-토큰-abc123!@#"
|
|
envelope = svc.encrypt_token(plain)
|
|
self.assertTrue(envelope.startswith("v1."))
|
|
self.assertNotIn(plain, envelope)
|
|
self.assertEqual(svc.decrypt_token(envelope), plain)
|
|
|
|
def test_unique_nonce_produces_unique_ciphertext(self):
|
|
first = svc.encrypt_token("same-token")
|
|
second = svc.encrypt_token("same-token")
|
|
self.assertNotEqual(first, second)
|
|
|
|
def test_tampered_envelope_rejected(self):
|
|
envelope = svc.encrypt_token("secret-token")
|
|
raw = bytearray(envelope.encode("ascii"))
|
|
raw[-1] = raw[-1] # 정상
|
|
corrupted = (envelope[:-1] + ("A" if envelope[-1] != "A" else "B"))
|
|
with self.assertRaises(svc.ProviderCredentialError):
|
|
svc.decrypt_token(corrupted)
|
|
|
|
def test_master_secret_falls_back_to_session_secret(self):
|
|
with patch.object(svc.settings, "provider_credential_secret", SecretStr("")), patch.object(
|
|
svc.settings, "session_secret", "session-secret-long-enough-for-prod-1234"
|
|
):
|
|
self.assertEqual(svc._master_secret(), b"session-secret-long-enough-for-prod-1234")
|
|
|
|
def test_short_master_secret_rejected(self):
|
|
with patch.object(svc.settings, "provider_credential_secret", SecretStr("short")):
|
|
with self.assertRaises(svc.ProviderCredentialError):
|
|
svc._master_secret()
|
|
|
|
|
|
class TokenHintTest(unittest.TestCase):
|
|
def test_long_token_tail_hint(self):
|
|
self.assertEqual(svc.token_hint("sk-1234567890abcdef"), "…cdef")
|
|
|
|
def test_short_token_generic_hint(self):
|
|
self.assertEqual(svc.token_hint("short"), "저장됨")
|
|
|
|
|
|
class CatalogTest(unittest.TestCase):
|
|
def test_catalog_covers_provider_codes(self):
|
|
self.assertEqual(
|
|
tuple(meta.code for meta in svc.PROVIDER_CATALOG),
|
|
svc.PROVIDER_CODES,
|
|
)
|
|
|
|
def test_openrouter_and_agy_present(self):
|
|
codes = set(svc.PROVIDER_CODES)
|
|
self.assertIn("openrouter", codes)
|
|
self.assertIn("agy", codes)
|
|
self.assertIn("claude", codes)
|
|
self.assertIn("codex", codes)
|
|
|
|
|
|
class SaveCredentialQueryTest(unittest.IsolatedAsyncioTestCase):
|
|
async def test_upsert_preserves_omitted_refresh_and_extra(self):
|
|
row = {
|
|
"provider": "openrouter",
|
|
"token_hint": "…next",
|
|
"auth_kind": "oauth_token",
|
|
"updated_by": "admin@example.com",
|
|
"updated_at": None,
|
|
"last_verified_at": None,
|
|
"last_verify_ok": None,
|
|
"last_verify_error": None,
|
|
"refresh_token_encrypted": None,
|
|
"extra": {},
|
|
}
|
|
conn = AsyncMock()
|
|
conn.fetchrow.return_value = row
|
|
|
|
class _Acquire:
|
|
async def __aenter__(self):
|
|
return conn
|
|
|
|
async def __aexit__(self, exc_type, exc, traceback):
|
|
return False
|
|
|
|
class _Pool:
|
|
def acquire(self):
|
|
return _Acquire()
|
|
|
|
with (
|
|
patch.object(svc, "ensure_table", AsyncMock()),
|
|
patch.object(svc, "get_pool", return_value=_Pool()),
|
|
patch.object(svc.settings, "provider_credential_secret", SecretStr("unit-test-master-secret-0123456789")),
|
|
):
|
|
await svc.save_credential(
|
|
provider="openrouter",
|
|
token="new-access-token",
|
|
auth_kind="oauth_token",
|
|
updated_by="admin@example.com",
|
|
)
|
|
|
|
query = conn.fetchrow.await_args.args[0]
|
|
self.assertIn("INSERT INTO app.admin_provider_credential AS existing", query)
|
|
self.assertIn("EXCLUDED.refresh_token_encrypted,\n existing.refresh_token_encrypted", query)
|
|
self.assertIn("WHEN EXCLUDED.extra = '{}'::jsonb THEN existing.extra", query)
|
|
self.assertEqual(conn.fetchrow.await_args.args[6], None)
|
|
self.assertEqual(conn.fetchrow.await_args.args[7], "{}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|