import base64 import unittest from time import time from typing import Mapping from urllib.parse import parse_qs, urlparse import httpx import jwt from cryptography.hazmat.primitives.asymmetric import rsa from pwo import async_test from kaya.core import HttpContext, KayaApp from kaya.session import InMemorySessionStore, SessionMixin from kaya.oidc import OIDCConfig, OIDCMixin from kaya.oidc._client import OIDCClient class MockOIDCProvider: def __init__(self) -> None: self.private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) self.public_key = self.private_key.public_key() self.kid = 'mock-key' self.issuer = 'https://mock-oidc.local' self.nonce = 'nonce-value' self.discovery = { 'issuer': self.issuer, 'authorization_endpoint': f'{self.issuer}/auth', 'token_endpoint': f'{self.issuer}/token', 'userinfo_endpoint': f'{self.issuer}/userinfo', 'end_session_endpoint': f'{self.issuer}/logout', 'jwks_uri': f'{self.issuer}/jwks', } def _public_key_to_jwk(self) -> dict[str, str]: numbers = self.public_key.public_numbers() e_bytes = numbers.e.to_bytes((numbers.e.bit_length() + 7) // 8, 'big') n_bytes = numbers.n.to_bytes((numbers.n.bit_length() + 7) // 8, 'big') return { 'kty': 'RSA', 'kid': self.kid, 'use': 'sig', 'n': base64.urlsafe_b64encode(n_bytes).rstrip(b'=').decode(), 'e': base64.urlsafe_b64encode(e_bytes).rstrip(b'=').decode(), 'alg': 'RS256', } def issue_id_token(self, nonce: str, audience: str, expired: bool = False, wrong_nonce: bool = False) -> str: now = time() exp = now - 3600 if expired else now + 3600 payload = { 'sub': 'user123', 'iss': self.issuer, 'aud': audience, 'iat': now, 'exp': exp, 'nonce': 'wrong-nonce' if wrong_nonce else nonce, } return jwt.encode( payload, self.private_key, algorithm='RS256', headers={'kid': self.kid}, ) def handle_request(self, request: httpx.Request) -> httpx.Response: url = str(request.url) path = urlparse(url).path # Serve discovery and JWKS for any host so that issuer-mismatch tests can still fetch keys. if path == '/.well-known/openid-configuration': return httpx.Response(200, json=self.discovery) if path == '/jwks': return httpx.Response(200, json={'keys': [self._public_key_to_jwk()]}) if path == '/token' and request.method == 'POST': body = request.content.decode() data = dict(part.split('=') for part in body.split('&')) if body else {} if data.get('grant_type') == 'refresh_token': return httpx.Response(200, json={ 'access_token': 'new-access-token', 'refresh_token': 'new-refresh-token', 'id_token': self.issue_id_token('refreshed-nonce', data.get('client_id', 'client')), 'token_type': 'Bearer', }) return httpx.Response(200, json={ 'access_token': 'mock-access-token', 'refresh_token': 'mock-refresh-token', 'id_token': self.issue_id_token(self.nonce, data.get('client_id', 'client')), 'token_type': 'Bearer', }) if path == '/userinfo': return httpx.Response(200, json={'sub': 'user123', 'email': 'user@example.com'}) if path == '/logout': return httpx.Response(200) return httpx.Response(404) class MockTransport(httpx.AsyncBaseTransport): def __init__(self, provider: MockOIDCProvider) -> None: self._provider = provider async def handle_async_request(self, request: httpx.Request) -> httpx.Response: return self._provider.handle_request(request) class OIDCClientTest(unittest.TestCase): def setUp(self) -> None: self.provider = MockOIDCProvider() self.http_client = httpx.AsyncClient(transport=MockTransport(self.provider)) self.config = OIDCConfig( issuer=self.provider.issuer, client_id='client', client_secret='secret', redirect_uri='http://localhost:8000/auth/callback', http_client=self.http_client, ) self.client = OIDCClient(self.config) @async_test async def test_build_authorization_url(self) -> None: url, state, nonce, verifier = await self.client.build_authorization_url() self.assertIn(self.provider.discovery['authorization_endpoint'], url) parsed = urlparse(url) query = parse_qs(parsed.query) self.assertEqual(['client'], query['client_id']) self.assertEqual(['code'], query['response_type']) self.assertEqual(['S256'], query['code_challenge_method']) self.assertIn('openid', query['scope'][0]) self.assertEqual([state], query['state']) self.assertEqual([nonce], query['nonce']) self.assertTrue(len(verifier) > 0) @async_test async def test_validate_id_token_success(self) -> None: token = self.provider.issue_id_token('nonce-value', 'client') claims = await self.client.validate_id_token(token, 'nonce-value') self.assertEqual('user123', claims['sub']) self.assertEqual('nonce-value', claims['nonce']) @async_test async def test_validate_id_token_wrong_nonce(self) -> None: token = self.provider.issue_id_token('nonce-value', 'client', wrong_nonce=True) with self.assertRaises(ValueError) as ctx: await self.client.validate_id_token(token, 'nonce-value') self.assertIn('nonce', str(ctx.exception)) @async_test async def test_validate_id_token_expired(self) -> None: token = self.provider.issue_id_token('nonce-value', 'client', expired=True) with self.assertRaises(jwt.ExpiredSignatureError): await self.client.validate_id_token(token, 'nonce-value') @async_test async def test_validate_id_token_wrong_issuer(self) -> None: token = self.provider.issue_id_token('nonce-value', 'client') config = OIDCConfig( issuer='https://wrong-issuer.local', client_id='client', redirect_uri='http://localhost:8000/auth/callback', http_client=self.http_client, ) client = OIDCClient(config) with self.assertRaises(jwt.InvalidIssuerError): await client.validate_id_token(token, 'nonce-value') @async_test async def test_validate_id_token_disallowed_algorithm(self) -> None: token = self.provider.issue_id_token('nonce-value', 'client') config = OIDCConfig( issuer=self.provider.issuer, client_id='client', redirect_uri='http://localhost:8000/auth/callback', http_client=self.http_client, allowed_id_token_algorithms=('ES256',), ) client = OIDCClient(config) with self.assertRaises(ValueError) as ctx: await client.validate_id_token(token, 'nonce-value') self.assertIn('RS256', str(ctx.exception)) @async_test async def test_fetch_userinfo(self) -> None: userinfo = await self.client.fetch_userinfo('access-token') self.assertEqual('user123', userinfo['sub']) self.assertEqual('user@example.com', userinfo['email']) @async_test async def test_build_logout_url(self) -> None: url = await self.client.build_logout_url('id-token') self.assertIsNotNone(url) assert url is not None self.assertIn(self.provider.discovery['end_session_endpoint'], url) parsed = urlparse(url) query = parse_qs(parsed.query) self.assertEqual(['id-token'], query['id_token_hint']) self.assertEqual(['/'], query['post_logout_redirect_uri']) class OIDCMixinTest(unittest.TestCase): def _build_app(self, fetch_userinfo: bool = False) -> tuple[KayaApp, OIDCMixin, MockOIDCProvider, InMemorySessionStore]: provider = MockOIDCProvider() http_client = httpx.AsyncClient(transport=MockTransport(provider)) store = InMemorySessionStore() session = SessionMixin(store) config = OIDCConfig( issuer=provider.issuer, client_id='client', client_secret='secret', redirect_uri='http://localhost:8000/auth/callback', http_client=http_client, fetch_userinfo=fetch_userinfo, ) oidc = OIDCMixin(config, session=session) app = KayaApp(mixins=[session, oidc]) return app, oidc, provider, store def _setup_routes(self, app: KayaApp, oidc: OIDCMixin) -> None: @app.GET('/') async def home(ctx: HttpContext) -> None: await ctx.send_str(200, 'home') @app.GET('/profile') @oidc.require_auth async def profile(ctx: HttpContext) -> None: user = oidc.get_user(ctx) if user is None: await ctx.send_empty(401) return await ctx.send_str(200, f'Hello {user.email}') @app.GET('/refresh') @oidc.require_auth async def refresh(ctx: HttpContext) -> None: new_token = await oidc.refresh_access_token(ctx.session) await ctx.send_str(200, new_token or 'no-token') @async_test async def test_login_redirect(self) -> None: app, oidc, provider, store = self._build_app() self._setup_routes(app, oidc) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) self.assertEqual(302, r.status_code) location = r.headers['Location'] self.assertIn(provider.discovery['authorization_endpoint'], location) self.assertIn('Set-Cookie', r.headers) @async_test async def test_callback_success(self) -> None: app, oidc, provider, store = self._build_app(fetch_userinfo=True) self._setup_routes(app, oidc) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) self.assertEqual(302, r.status_code) parsed = urlparse(r.headers['Location']) query = parse_qs(parsed.query) state = query['state'][0] nonce = query['nonce'][0] provider.nonce = nonce r = await client.get('/auth/callback', params={'code': 'mock-code', 'state': state}, follow_redirects=False) self.assertEqual(302, r.status_code) self.assertEqual('/', r.headers['Location']) self.assertIn('Set-Cookie', r.headers) r = await client.get('/profile') self.assertEqual(200, r.status_code) self.assertIn('user@example.com', r.text) @async_test async def test_callback_invalid_state(self) -> None: app, oidc, provider, store = self._build_app() self._setup_routes(app, oidc) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/callback', params={'code': 'mock-code', 'state': 'wrong'}, follow_redirects=False) self.assertEqual(400, r.status_code) @async_test async def test_logout(self) -> None: app, oidc, provider, store = self._build_app() self._setup_routes(app, oidc) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) parsed = urlparse(r.headers['Location']) query = parse_qs(parsed.query) state = query['state'][0] provider.nonce = query['nonce'][0] await client.get('/auth/callback', params={'code': 'mock-code', 'state': state}, follow_redirects=False) session_id = client.cookies['session_id'] self.assertIn(session_id, store._data) r = await client.get('/auth/logout', follow_redirects=False) self.assertEqual(302, r.status_code) self.assertIn(provider.discovery['end_session_endpoint'], r.headers['Location']) self.assertNotIn(session_id, store._data) @async_test async def test_require_auth_redirect(self) -> None: app, oidc, provider, store = self._build_app() self._setup_routes(app, oidc) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/profile', follow_redirects=False) self.assertEqual(302, r.status_code) self.assertEqual('/auth/login', r.headers['Location']) @async_test async def test_refresh_access_token(self) -> None: app, oidc, provider, store = self._build_app() self._setup_routes(app, oidc) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) parsed = urlparse(r.headers['Location']) query = parse_qs(parsed.query) state = query['state'][0] provider.nonce = query['nonce'][0] await client.get('/auth/callback', params={'code': 'mock-code', 'state': state}, follow_redirects=False) r = await client.get('/refresh') self.assertEqual(200, r.status_code) self.assertEqual('new-access-token', r.text) @async_test async def test_dependency_applied_automatically(self) -> None: # Only pass oidc to KayaApp; SessionMixin should be applied via dependencies. provider = MockOIDCProvider() http_client = httpx.AsyncClient(transport=MockTransport(provider)) store = InMemorySessionStore() session = SessionMixin(store) config = OIDCConfig( issuer=provider.issuer, client_id='client', client_secret='secret', redirect_uri='http://localhost:8000/auth/callback', http_client=http_client, ) oidc = OIDCMixin(config, session=session) app = KayaApp(mixins=[oidc]) @app.GET('/') async def home(ctx: HttpContext) -> None: ctx.session['x'] = 1 await ctx.send_str(200, 'ok') transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') self.assertEqual(200, r.status_code) self.assertIn('Set-Cookie', r.headers) if __name__ == '__main__': unittest.main()