Replace wrapper-based SessionMiddleware/OIDCApp with KayaMixin subclasses applied via KayaApp(mixins=[...]). Mixins hook into handle_request and handle_websocket via before/after hooks, so both ASGI and RSGI keep working. Mixin dependencies are applied automatically and deduplicated.
361 lines
15 KiB
Python
361 lines
15 KiB
Python
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()
|