Add kaya-oidc package for OpenID Connect authentication
This commit is contained in:
@@ -36,6 +36,10 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
.venv/bin/python -m mypy -p kaya.session
|
.venv/bin/python -m mypy -p kaya.session
|
||||||
.venv/bin/python -m unittest discover -s packages/kaya-session/tests
|
.venv/bin/python -m unittest discover -s packages/kaya-session/tests
|
||||||
|
- name: Check kaya-oidc
|
||||||
|
run: |
|
||||||
|
.venv/bin/python -m mypy -p kaya.oidc
|
||||||
|
.venv/bin/python -m unittest discover -s packages/kaya-oidc/tests
|
||||||
- name: Publish kaya-core artifacts
|
- name: Publish kaya-core artifacts
|
||||||
env:
|
env:
|
||||||
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
||||||
@@ -60,3 +64,11 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
.venv/bin/pyproject-build packages/kaya-session
|
.venv/bin/pyproject-build packages/kaya-session
|
||||||
.venv/bin/twine upload --repository gitea packages/kaya-session/dist/*.whl packages/kaya-session/dist/*.tar.gz
|
.venv/bin/twine upload --repository gitea packages/kaya-session/dist/*.whl packages/kaya-session/dist/*.tar.gz
|
||||||
|
- name: Publish kaya-oidc artifacts
|
||||||
|
env:
|
||||||
|
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
||||||
|
TWINE_USERNAME: ${{ vars.PUBLISHER_USERNAME }}
|
||||||
|
TWINE_PASSWORD: ${{ secrets.PUBLISHER_TOKEN }}
|
||||||
|
run: |
|
||||||
|
.venv/bin/pyproject-build packages/kaya-oidc
|
||||||
|
.venv/bin/twine upload --repository gitea packages/kaya-oidc/dist/*.whl packages/kaya-oidc/dist/*.tar.gz
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ This repository is a monorepo for the Kaya framework. The code is split into ind
|
|||||||
- **kaya-core** — core routing, HTTP/WS abstractions, and ASGI adapter (`packages/kaya-core/`)
|
- **kaya-core** — core routing, HTTP/WS abstractions, and ASGI adapter (`packages/kaya-core/`)
|
||||||
- **kaya-rsgi** — RSGI/Granian integration (`packages/kaya-rsgi/`)
|
- **kaya-rsgi** — RSGI/Granian integration (`packages/kaya-rsgi/`)
|
||||||
- **kaya-session** — server-side HTTP session management (`packages/kaya-session/`)
|
- **kaya-session** — server-side HTTP session management (`packages/kaya-session/`)
|
||||||
|
- **kaya-oidc** — OpenID Connect authentication (`packages/kaya-oidc/`)
|
||||||
|
|
||||||
Additional `kaya-*` packages can be added as new directories under `packages/`.
|
Additional `kaya-*` packages can be added as new directories under `packages/`.
|
||||||
|
|
||||||
@@ -23,7 +24,7 @@ pip install --index-url https://gitea.woggioni.net/api/packages/woggioni/pypi/si
|
|||||||
Install the packages in development mode:
|
Install the packages in development mode:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -e packages/kaya-core -e packages/kaya-rsgi -e packages/kaya-session
|
pip install -e packages/kaya-core -e packages/kaya-rsgi -e packages/kaya-session -e packages/kaya-oidc
|
||||||
```
|
```
|
||||||
|
|
||||||
Run the example:
|
Run the example:
|
||||||
@@ -38,6 +39,7 @@ python example/hello.py
|
|||||||
python -m unittest discover -s packages/kaya-core/tests
|
python -m unittest discover -s packages/kaya-core/tests
|
||||||
python -m unittest discover -s packages/kaya-rsgi/tests
|
python -m unittest discover -s packages/kaya-rsgi/tests
|
||||||
python -m unittest discover -s packages/kaya-session/tests
|
python -m unittest discover -s packages/kaya-session/tests
|
||||||
|
python -m unittest discover -s packages/kaya-oidc/tests
|
||||||
```
|
```
|
||||||
|
|
||||||
## Static analysis
|
## Static analysis
|
||||||
@@ -46,6 +48,7 @@ python -m unittest discover -s packages/kaya-session/tests
|
|||||||
mypy -p kaya.core
|
mypy -p kaya.core
|
||||||
mypy -p kaya.rsgi
|
mypy -p kaya.rsgi
|
||||||
mypy -p kaya.session
|
mypy -p kaya.session
|
||||||
|
mypy -p kaya.oidc
|
||||||
```
|
```
|
||||||
|
|
||||||
## Building packages
|
## Building packages
|
||||||
@@ -54,4 +57,5 @@ mypy -p kaya.session
|
|||||||
python -m build packages/kaya-core
|
python -m build packages/kaya-core
|
||||||
python -m build packages/kaya-rsgi
|
python -m build packages/kaya-rsgi
|
||||||
python -m build packages/kaya-session
|
python -m build packages/kaya-session
|
||||||
|
python -m build packages/kaya-oidc
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -0,0 +1,34 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
from kaya.core import HttpContext, KayaApp
|
||||||
|
from kaya.oidc import OIDCConfig, OIDCApp
|
||||||
|
from kaya.session import InMemorySessionStore, SessionMiddleware
|
||||||
|
|
||||||
|
app = KayaApp()
|
||||||
|
session_app = SessionMiddleware(app, InMemorySessionStore())
|
||||||
|
|
||||||
|
oidc = OIDCApp(
|
||||||
|
session_app,
|
||||||
|
OIDCConfig(
|
||||||
|
issuer=os.environ.get('OIDC_ISSUER', 'https://accounts.google.com'),
|
||||||
|
client_id=os.environ.get('OIDC_CLIENT_ID', 'replace-me'),
|
||||||
|
client_secret=os.environ.get('OIDC_CLIENT_SECRET'),
|
||||||
|
redirect_uri=os.environ.get('OIDC_REDIRECT_URI', 'http://localhost:8000/auth/callback'),
|
||||||
|
fetch_userinfo=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@oidc.GET('/')
|
||||||
|
async def home(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(200, 'public home')
|
||||||
|
|
||||||
|
|
||||||
|
@oidc.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.name or user.email or user.sub}')
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
# kaya-oidc
|
||||||
|
|
||||||
|
OpenID Connect authentication for the Kaya web framework.
|
||||||
|
|
||||||
|
Built on top of `kaya-session` and implements the OIDC **Authorization Code
|
||||||
|
Flow with PKCE**.
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
```python
|
||||||
|
import os
|
||||||
|
from kaya.core import HttpContext, KayaApp
|
||||||
|
from kaya.session import SessionMiddleware, InMemorySessionStore
|
||||||
|
from kaya.oidc import OIDCConfig, OIDCApp
|
||||||
|
|
||||||
|
app = KayaApp()
|
||||||
|
session_app = SessionMiddleware(app, InMemorySessionStore())
|
||||||
|
|
||||||
|
oidc = OIDCApp(
|
||||||
|
session_app,
|
||||||
|
OIDCConfig(
|
||||||
|
issuer=os.environ['OIDC_ISSUER'],
|
||||||
|
client_id=os.environ['OIDC_CLIENT_ID'],
|
||||||
|
client_secret=os.environ.get('OIDC_CLIENT_SECRET'),
|
||||||
|
redirect_uri='http://localhost:8000/auth/callback',
|
||||||
|
fetch_userinfo=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@oidc.GET('/')
|
||||||
|
async def home(ctx: HttpContext):
|
||||||
|
await ctx.send_str(200, 'public home')
|
||||||
|
|
||||||
|
@oidc.GET('/profile')
|
||||||
|
@oidc.require_auth
|
||||||
|
async def profile(ctx: HttpContext):
|
||||||
|
user = oidc.get_user(ctx)
|
||||||
|
await ctx.send_str(200, f'Hello {user.email or user.sub}')
|
||||||
|
```
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- Generic OIDC discovery
|
||||||
|
- Authorization Code Flow with PKCE (S256)
|
||||||
|
- ID token signature validation with JWKS
|
||||||
|
- `state` and `nonce` protection
|
||||||
|
- Session fixation defense via `regenerate_id()` after login
|
||||||
|
- Optional userinfo endpoint fetch
|
||||||
|
- Refresh token support
|
||||||
|
- RP-initiated logout (when provider advertises `end_session_endpoint`)
|
||||||
|
|
||||||
|
## Security notes
|
||||||
|
|
||||||
|
- The `none` signing algorithm is rejected by default.
|
||||||
|
- Only algorithms listed in `OIDCConfig.allowed_id_token_algorithms` are accepted.
|
||||||
|
- Always use HTTPS for `redirect_uri` in production.
|
||||||
|
|
||||||
|
## Supported flows
|
||||||
|
|
||||||
|
Only the Authorization Code Flow with PKCE is supported. Implicit and Hybrid
|
||||||
|
flows are intentionally not implemented.
|
||||||
|
|
||||||
|
## Supported algorithms
|
||||||
|
|
||||||
|
ID token signature verification supports:
|
||||||
|
`RS256`, `RS384`, `RS512`, `PS256`, `PS384`, `PS512`,
|
||||||
|
`ES256`, `ES384`, `ES512`, and `EdDSA`.
|
||||||
|
|
||||||
|
HMAC algorithms (`HS*`) are disabled by default and can be enabled by adding
|
||||||
|
them to `allowed_id_token_algorithms` if your provider uses them.
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
[build-system]
|
||||||
|
requires = ["setuptools>=61.0", "setuptools-scm>=8"]
|
||||||
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[project]
|
||||||
|
name = "kaya-oidc"
|
||||||
|
dynamic = ["version"]
|
||||||
|
authors = [
|
||||||
|
{ name="Walter Oggioni", email="oggioni.walter@gmail.com" },
|
||||||
|
]
|
||||||
|
description = "OpenID Connect authentication for the Kaya lightweight ASGI web framework"
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.10"
|
||||||
|
license = "MIT"
|
||||||
|
classifiers = [
|
||||||
|
'Development Status :: 3 - Alpha',
|
||||||
|
'Topic :: Utilities',
|
||||||
|
'Intended Audience :: System Administrators',
|
||||||
|
'Intended Audience :: Developers',
|
||||||
|
'Environment :: Console',
|
||||||
|
'Programming Language :: Python :: 3',
|
||||||
|
]
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
"kaya-core",
|
||||||
|
"kaya-session",
|
||||||
|
"httpx",
|
||||||
|
"PyJWT[crypto]",
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = [
|
||||||
|
"build", "mypy", "ipdb", "twine", "httpx", "httpx-ws"
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.urls]
|
||||||
|
"Homepage" = "https://github.com/woggioni/kaya"
|
||||||
|
"Bug Tracker" = "https://github.com/woggioni/kaya/issues"
|
||||||
|
|
||||||
|
[tool.setuptools.packages.find]
|
||||||
|
where = ["src"]
|
||||||
|
namespaces = true
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
python_version = "3.12"
|
||||||
|
disallow_untyped_defs = true
|
||||||
|
show_error_codes = true
|
||||||
|
no_implicit_optional = true
|
||||||
|
warn_return_any = true
|
||||||
|
warn_unused_ignores = true
|
||||||
|
exclude = ["scripts", "docs", "test"]
|
||||||
|
strict = true
|
||||||
|
|
||||||
|
[tool.setuptools_scm]
|
||||||
|
root = "../.."
|
||||||
|
version_file = "src/kaya/oidc/_version.py"
|
||||||
|
|
||||||
|
[tool.setuptools_scm.tag]
|
||||||
|
prefix = "release/"
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
from ._app import OIDCApp, OIDCUser
|
||||||
|
from ._client import OIDCClient
|
||||||
|
from ._config import OIDCConfig
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'OIDCApp',
|
||||||
|
'OIDCClient',
|
||||||
|
'OIDCConfig',
|
||||||
|
'OIDCUser',
|
||||||
|
]
|
||||||
@@ -0,0 +1,234 @@
|
|||||||
|
from typing import Any, Awaitable, Callable, Mapping, MutableMapping, Optional, Sequence, cast
|
||||||
|
from urllib.parse import parse_qs, urlencode
|
||||||
|
|
||||||
|
from kaya.core import HttpContext, HttpMethod, KayaApp
|
||||||
|
from kaya.session import Session, SessionMiddleware
|
||||||
|
|
||||||
|
from ._client import OIDCClient
|
||||||
|
from ._config import OIDCConfig
|
||||||
|
|
||||||
|
|
||||||
|
type HttpHandler = Callable[..., Awaitable[None]]
|
||||||
|
type WebSocketHandler = Callable[..., Awaitable[None]]
|
||||||
|
type RouteDecorator = Callable[[HttpHandler], HttpHandler]
|
||||||
|
type WebSocketDecorator = Callable[[WebSocketHandler], WebSocketHandler]
|
||||||
|
type ASGIApp = Callable[
|
||||||
|
[MutableMapping[str, Any], Callable[[], Awaitable[Any]], Callable[[MutableMapping[str, Any]], Awaitable[None]]],
|
||||||
|
Awaitable[None],
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class OIDCUser(Mapping[str, Any]):
|
||||||
|
"""Read-only view of the OIDC userinfo / ID token claims."""
|
||||||
|
|
||||||
|
def __init__(self, data: Mapping[str, Any]) -> None:
|
||||||
|
self._data = dict(data)
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> Any:
|
||||||
|
return self._data[key]
|
||||||
|
|
||||||
|
def __iter__(self) -> Any:
|
||||||
|
return iter(self._data)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self._data)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sub(self) -> str:
|
||||||
|
return self._data['sub']
|
||||||
|
|
||||||
|
@property
|
||||||
|
def email(self) -> Optional[str]:
|
||||||
|
return self._data.get('email')
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> Optional[str]:
|
||||||
|
return self._data.get('name')
|
||||||
|
|
||||||
|
@property
|
||||||
|
def picture(self) -> Optional[str]:
|
||||||
|
return self._data.get('picture')
|
||||||
|
|
||||||
|
|
||||||
|
class OIDCApp:
|
||||||
|
"""ASGI app wrapper that adds OIDC authentication routes to a Kaya app.
|
||||||
|
|
||||||
|
The wrapped app must be a ``SessionMiddleware`` instance so that OIDC state,
|
||||||
|
nonce, and user data can be stored in ``ctx.session``.
|
||||||
|
|
||||||
|
Built-in routes:
|
||||||
|
|
||||||
|
- ``login_path`` (default ``/auth/login``): redirects to the OIDC provider.
|
||||||
|
- ``callback_path`` (default ``/auth/callback``): handles the provider callback.
|
||||||
|
- ``logout_path`` (default ``/auth/logout``): logs the user out.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, app: SessionMiddleware, config: OIDCConfig) -> None:
|
||||||
|
self._app = app
|
||||||
|
self._config = config
|
||||||
|
self._client = OIDCClient(config)
|
||||||
|
self._register_routes()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _session(ctx: HttpContext) -> Session:
|
||||||
|
session = ctx.session
|
||||||
|
assert isinstance(session, Session)
|
||||||
|
return session
|
||||||
|
|
||||||
|
def _register_routes(self) -> None:
|
||||||
|
@self._app.GET(self._config.login_path)
|
||||||
|
async def login(ctx: HttpContext) -> None:
|
||||||
|
auth_url, state, nonce, code_verifier = await self._client.build_authorization_url()
|
||||||
|
session = self._session(ctx)
|
||||||
|
session['oidc_state'] = state
|
||||||
|
session['oidc_nonce'] = nonce
|
||||||
|
session['oidc_code_verifier'] = code_verifier
|
||||||
|
await ctx.send_empty(302, {'Location': auth_url})
|
||||||
|
|
||||||
|
@self._app.GET(self._config.callback_path)
|
||||||
|
async def callback(ctx: HttpContext) -> None:
|
||||||
|
query = parse_qs(ctx.query_string)
|
||||||
|
code = self._first_value(query.get('code'))
|
||||||
|
state = self._first_value(query.get('state'))
|
||||||
|
error = self._first_value(query.get('error'))
|
||||||
|
error_description = self._first_value(query.get('error_description'))
|
||||||
|
|
||||||
|
if error is not None:
|
||||||
|
message = f'OIDC error: {error}'
|
||||||
|
if error_description is not None:
|
||||||
|
message += f': {error_description}'
|
||||||
|
await ctx.send_str(400, message)
|
||||||
|
return
|
||||||
|
|
||||||
|
if code is None or state is None:
|
||||||
|
await ctx.send_str(400, 'Missing code or state')
|
||||||
|
return
|
||||||
|
|
||||||
|
session = self._session(ctx)
|
||||||
|
expected_state = session.get('oidc_state')
|
||||||
|
if state != expected_state:
|
||||||
|
await ctx.send_str(400, 'Invalid state')
|
||||||
|
return
|
||||||
|
|
||||||
|
nonce = session.get('oidc_nonce', '')
|
||||||
|
code_verifier = session.get('oidc_code_verifier', '')
|
||||||
|
|
||||||
|
try:
|
||||||
|
token_response = await self._client.fetch_token(code, code_verifier)
|
||||||
|
id_token = token_response.get('id_token')
|
||||||
|
if not isinstance(id_token, str):
|
||||||
|
await ctx.send_str(400, 'No id_token in token response')
|
||||||
|
return
|
||||||
|
|
||||||
|
claims = await self._client.validate_id_token(id_token, nonce)
|
||||||
|
|
||||||
|
user: dict[str, Any] = dict(claims)
|
||||||
|
access_token = token_response.get('access_token')
|
||||||
|
if self._config.fetch_userinfo and isinstance(access_token, str):
|
||||||
|
userinfo = await self._client.fetch_userinfo(access_token)
|
||||||
|
user.update(userinfo)
|
||||||
|
|
||||||
|
session.pop('oidc_state', None)
|
||||||
|
session.pop('oidc_nonce', None)
|
||||||
|
session.pop('oidc_code_verifier', None)
|
||||||
|
|
||||||
|
session['oidc_user'] = user
|
||||||
|
session['oidc_id_token'] = id_token
|
||||||
|
if isinstance(access_token, str):
|
||||||
|
session['oidc_access_token'] = access_token
|
||||||
|
if 'refresh_token' in token_response:
|
||||||
|
session['oidc_refresh_token'] = token_response['refresh_token']
|
||||||
|
|
||||||
|
session.regenerate_id()
|
||||||
|
await ctx.send_empty(302, {'Location': self._config.post_login_redirect})
|
||||||
|
except ValueError as exc:
|
||||||
|
await ctx.send_str(400, f'Authentication failed: {exc}')
|
||||||
|
|
||||||
|
@self._app.GET(self._config.logout_path)
|
||||||
|
async def logout(ctx: HttpContext) -> None:
|
||||||
|
session = self._session(ctx)
|
||||||
|
id_token = session.get('oidc_id_token')
|
||||||
|
session.invalidate()
|
||||||
|
logout_url = await self._client.build_logout_url(id_token if isinstance(id_token, str) else None)
|
||||||
|
location = logout_url if logout_url is not None else self._config.post_logout_redirect
|
||||||
|
await ctx.send_empty(302, {'Location': location})
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _first_value(values: Optional[Sequence[str]]) -> Optional[str]:
|
||||||
|
if values and len(values) > 0:
|
||||||
|
return values[0]
|
||||||
|
return None
|
||||||
|
|
||||||
|
def route(
|
||||||
|
self,
|
||||||
|
paths: str | Sequence[str],
|
||||||
|
methods: Optional[HttpMethod | Sequence[HttpMethod]] = None,
|
||||||
|
recursive: bool = False,
|
||||||
|
) -> RouteDecorator:
|
||||||
|
return self._app.route(paths, methods, recursive)
|
||||||
|
|
||||||
|
def GET(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.GET(path, recursive)
|
||||||
|
|
||||||
|
def POST(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.POST(path, recursive)
|
||||||
|
|
||||||
|
def PUT(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.PUT(path, recursive)
|
||||||
|
|
||||||
|
def DELETE(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.DELETE(path, recursive)
|
||||||
|
|
||||||
|
def OPTIONS(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.OPTIONS(path, recursive)
|
||||||
|
|
||||||
|
def HEAD(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.HEAD(path, recursive)
|
||||||
|
|
||||||
|
def PATCH(self, path: str, recursive: bool = False) -> RouteDecorator:
|
||||||
|
return self._app.PATCH(path, recursive)
|
||||||
|
|
||||||
|
def websocket(self, path: str, recursive: bool = False) -> WebSocketDecorator:
|
||||||
|
return self._app.websocket(path, recursive)
|
||||||
|
|
||||||
|
async def __call__(
|
||||||
|
self,
|
||||||
|
scope: MutableMapping[str, Any],
|
||||||
|
receive: Callable[[], Awaitable[Any]],
|
||||||
|
send: Callable[[MutableMapping[str, Any]], Awaitable[None]],
|
||||||
|
) -> None:
|
||||||
|
await cast(ASGIApp, self._app)(scope, receive, send)
|
||||||
|
|
||||||
|
def is_authenticated(self, ctx: HttpContext) -> bool:
|
||||||
|
return 'oidc_user' in self._session(ctx)
|
||||||
|
|
||||||
|
def get_user(self, ctx: HttpContext) -> Optional[OIDCUser]:
|
||||||
|
user = self._session(ctx).get('oidc_user')
|
||||||
|
if user is None or not isinstance(user, Mapping):
|
||||||
|
return None
|
||||||
|
return OIDCUser(user)
|
||||||
|
|
||||||
|
def require_auth(self, handler: HttpHandler) -> HttpHandler:
|
||||||
|
async def wrapper(ctx: HttpContext, *args: Any, **kwargs: Any) -> None:
|
||||||
|
if not self.is_authenticated(ctx):
|
||||||
|
await ctx.send_empty(302, {'Location': self._config.login_path})
|
||||||
|
return
|
||||||
|
await handler(ctx, *args, **kwargs)
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
async def refresh_access_token(self, session: Session) -> Optional[str]:
|
||||||
|
refresh_token = session.get('oidc_refresh_token')
|
||||||
|
if not isinstance(refresh_token, str):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
token_response = await self._client.refresh_token(refresh_token)
|
||||||
|
access_token = token_response.get('access_token')
|
||||||
|
if isinstance(access_token, str):
|
||||||
|
session['oidc_access_token'] = access_token
|
||||||
|
if 'id_token' in token_response:
|
||||||
|
session['oidc_id_token'] = token_response['id_token']
|
||||||
|
if 'refresh_token' in token_response:
|
||||||
|
session['oidc_refresh_token'] = token_response['refresh_token']
|
||||||
|
return access_token if isinstance(access_token, str) else None
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
@@ -0,0 +1,182 @@
|
|||||||
|
from typing import Any, Mapping, Optional, Sequence, cast
|
||||||
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import jwt
|
||||||
|
from jwt import PyJWK
|
||||||
|
|
||||||
|
from ._config import OIDCConfig
|
||||||
|
from ._utils import generate_nonce, generate_pkce, generate_state
|
||||||
|
|
||||||
|
|
||||||
|
class OIDCClient:
|
||||||
|
"""Low-level OIDC client implementing discovery, PKCE, and token validation."""
|
||||||
|
|
||||||
|
def __init__(self, config: OIDCConfig) -> None:
|
||||||
|
self._config = config
|
||||||
|
self._metadata: Optional[Mapping[str, Any]] = None
|
||||||
|
self._jwks: Optional[Mapping[str, Any]] = None
|
||||||
|
|
||||||
|
async def _ensure_metadata(self) -> Mapping[str, Any]:
|
||||||
|
if self._metadata is None:
|
||||||
|
url = f"{self._config.issuer.rstrip('/')}/.well-known/openid-configuration"
|
||||||
|
resp = await self._request('GET', url)
|
||||||
|
resp.raise_for_status()
|
||||||
|
self._metadata = cast(Mapping[str, Any], resp.json())
|
||||||
|
return self._metadata
|
||||||
|
|
||||||
|
async def _ensure_jwks(self) -> Mapping[str, Any]:
|
||||||
|
if self._jwks is None:
|
||||||
|
metadata = await self._ensure_metadata()
|
||||||
|
jwks_uri = metadata.get('jwks_uri')
|
||||||
|
if not isinstance(jwks_uri, str):
|
||||||
|
raise RuntimeError('OIDC provider does not advertise a jwks_uri')
|
||||||
|
resp = await self._request('GET', jwks_uri)
|
||||||
|
resp.raise_for_status()
|
||||||
|
self._jwks = cast(Mapping[str, Any], resp.json())
|
||||||
|
return self._jwks
|
||||||
|
|
||||||
|
async def _request(self, method: str, url: str, **kwargs: Any) -> httpx.Response:
|
||||||
|
if self._config.http_client is not None:
|
||||||
|
return await self._config.http_client.request(method, url, **kwargs)
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
return await client.request(method, url, **kwargs)
|
||||||
|
|
||||||
|
async def build_authorization_url(self) -> tuple[str, str, str, str]:
|
||||||
|
"""Return (authorization_url, state, nonce, code_verifier)."""
|
||||||
|
metadata = await self._ensure_metadata()
|
||||||
|
authorization_endpoint = metadata.get('authorization_endpoint')
|
||||||
|
if not isinstance(authorization_endpoint, str):
|
||||||
|
raise RuntimeError('OIDC provider does not advertise an authorization_endpoint')
|
||||||
|
|
||||||
|
state = generate_state()
|
||||||
|
nonce = generate_nonce()
|
||||||
|
code_verifier, code_challenge = generate_pkce()
|
||||||
|
|
||||||
|
params = {
|
||||||
|
'client_id': self._config.client_id,
|
||||||
|
'response_type': 'code',
|
||||||
|
'scope': ' '.join(self._config.scopes),
|
||||||
|
'redirect_uri': self._config.redirect_uri,
|
||||||
|
'state': state,
|
||||||
|
'nonce': nonce,
|
||||||
|
'code_challenge': code_challenge,
|
||||||
|
'code_challenge_method': 'S256',
|
||||||
|
}
|
||||||
|
url = authorization_endpoint + '?' + urlencode(params)
|
||||||
|
return url, state, nonce, code_verifier
|
||||||
|
|
||||||
|
async def fetch_token(self, code: str, code_verifier: str) -> Mapping[str, Any]:
|
||||||
|
metadata = await self._ensure_metadata()
|
||||||
|
token_endpoint = metadata.get('token_endpoint')
|
||||||
|
if not isinstance(token_endpoint, str):
|
||||||
|
raise RuntimeError('OIDC provider does not advertise a token_endpoint')
|
||||||
|
|
||||||
|
data = {
|
||||||
|
'grant_type': 'authorization_code',
|
||||||
|
'code': code,
|
||||||
|
'redirect_uri': self._config.redirect_uri,
|
||||||
|
'client_id': self._config.client_id,
|
||||||
|
'code_verifier': code_verifier,
|
||||||
|
}
|
||||||
|
auth: Optional[tuple[str, str]] = None
|
||||||
|
if self._config.client_secret is not None:
|
||||||
|
auth = (self._config.client_id, self._config.client_secret)
|
||||||
|
|
||||||
|
resp = await self._request('POST', token_endpoint, data=data, auth=auth)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return cast(Mapping[str, Any], resp.json())
|
||||||
|
|
||||||
|
async def refresh_token(self, refresh_token: str) -> Mapping[str, Any]:
|
||||||
|
metadata = await self._ensure_metadata()
|
||||||
|
token_endpoint = metadata.get('token_endpoint')
|
||||||
|
if not isinstance(token_endpoint, str):
|
||||||
|
raise RuntimeError('OIDC provider does not advertise a token_endpoint')
|
||||||
|
|
||||||
|
data = {
|
||||||
|
'grant_type': 'refresh_token',
|
||||||
|
'refresh_token': refresh_token,
|
||||||
|
'client_id': self._config.client_id,
|
||||||
|
}
|
||||||
|
auth: Optional[tuple[str, str]] = None
|
||||||
|
if self._config.client_secret is not None:
|
||||||
|
auth = (self._config.client_id, self._config.client_secret)
|
||||||
|
|
||||||
|
resp = await self._request('POST', token_endpoint, data=data, auth=auth)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return cast(Mapping[str, Any], resp.json())
|
||||||
|
|
||||||
|
async def validate_id_token(self, id_token: str, nonce: str) -> Mapping[str, Any]:
|
||||||
|
header = jwt.get_unverified_header(id_token)
|
||||||
|
alg = header.get('alg')
|
||||||
|
if not isinstance(alg, str):
|
||||||
|
raise ValueError('ID token header is missing alg')
|
||||||
|
|
||||||
|
if alg == 'none':
|
||||||
|
if not self._config.allow_unsigned_id_tokens:
|
||||||
|
raise ValueError('Unsigned ID tokens are not allowed')
|
||||||
|
payload = jwt.decode(
|
||||||
|
id_token,
|
||||||
|
options={'verify_signature': False},
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
if alg not in self._config.allowed_id_token_algorithms:
|
||||||
|
raise ValueError(f'ID token algorithm {alg} is not allowed')
|
||||||
|
jwks = await self._ensure_jwks()
|
||||||
|
signing_key = self._get_signing_key(jwks, header.get('kid'), alg)
|
||||||
|
payload = jwt.decode(
|
||||||
|
id_token,
|
||||||
|
signing_key.key,
|
||||||
|
algorithms=[alg],
|
||||||
|
issuer=self._config.issuer,
|
||||||
|
audience=self._config.client_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
raise ValueError('ID token payload is not a JSON object')
|
||||||
|
if payload.get('nonce') != nonce:
|
||||||
|
raise ValueError('ID token nonce mismatch')
|
||||||
|
return payload
|
||||||
|
|
||||||
|
def _get_signing_key(self, jwks: Mapping[str, Any], kid: Optional[str], alg: str) -> PyJWK:
|
||||||
|
keys = jwks.get('keys')
|
||||||
|
if not isinstance(keys, Sequence):
|
||||||
|
raise ValueError('JWKS contains no keys')
|
||||||
|
|
||||||
|
if kid is not None:
|
||||||
|
for key in keys:
|
||||||
|
if isinstance(key, Mapping) and key.get('kid') == kid:
|
||||||
|
return PyJWK(cast(dict[str, Any], dict(key)), algorithm=alg)
|
||||||
|
raise ValueError(f'No signing key found for kid {kid}')
|
||||||
|
|
||||||
|
if len(keys) == 1 and isinstance(keys[0], Mapping):
|
||||||
|
return PyJWK(cast(dict[str, Any], dict(keys[0])), algorithm=alg)
|
||||||
|
|
||||||
|
raise ValueError('ID token has no kid and JWKS contains multiple keys')
|
||||||
|
|
||||||
|
async def fetch_userinfo(self, access_token: str) -> Mapping[str, Any]:
|
||||||
|
metadata = await self._ensure_metadata()
|
||||||
|
userinfo_endpoint = metadata.get('userinfo_endpoint')
|
||||||
|
if not isinstance(userinfo_endpoint, str):
|
||||||
|
raise RuntimeError('OIDC provider does not advertise a userinfo_endpoint')
|
||||||
|
|
||||||
|
resp = await self._request(
|
||||||
|
'GET',
|
||||||
|
userinfo_endpoint,
|
||||||
|
headers={'Authorization': f'Bearer {access_token}'},
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return cast(Mapping[str, Any], resp.json())
|
||||||
|
|
||||||
|
async def build_logout_url(self, id_token: Optional[str] = None) -> Optional[str]:
|
||||||
|
metadata = await self._ensure_metadata()
|
||||||
|
end_session_endpoint = metadata.get('end_session_endpoint')
|
||||||
|
if not isinstance(end_session_endpoint, str):
|
||||||
|
return None
|
||||||
|
|
||||||
|
params: dict[str, str] = {
|
||||||
|
'post_logout_redirect_uri': self._config.post_logout_redirect,
|
||||||
|
}
|
||||||
|
if id_token is not None:
|
||||||
|
params['id_token_hint'] = id_token
|
||||||
|
return end_session_endpoint + '?' + urlencode(params)
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Optional, Sequence
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OIDCConfig:
|
||||||
|
"""Configuration for an OpenID Connect provider.
|
||||||
|
|
||||||
|
``client_secret`` is optional; omit it for public clients.
|
||||||
|
``fetch_userinfo`` controls whether the userinfo endpoint is queried after
|
||||||
|
token exchange.
|
||||||
|
``allowed_id_token_algorithms`` lists the JWS algorithms the client will
|
||||||
|
accept when validating the ID token. The ``none`` algorithm is rejected
|
||||||
|
unless ``allow_unsigned_id_tokens`` is set to ``True``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
issuer: str
|
||||||
|
client_id: str
|
||||||
|
redirect_uri: str
|
||||||
|
client_secret: Optional[str] = None
|
||||||
|
scopes: Sequence[str] = ('openid', 'email', 'profile')
|
||||||
|
login_path: str = '/auth/login'
|
||||||
|
callback_path: str = '/auth/callback'
|
||||||
|
logout_path: str = '/auth/logout'
|
||||||
|
post_login_redirect: str = '/'
|
||||||
|
post_logout_redirect: str = '/'
|
||||||
|
fetch_userinfo: bool = False
|
||||||
|
allow_unsigned_id_tokens: bool = False
|
||||||
|
allowed_id_token_algorithms: Sequence[str] = field(default_factory=lambda: (
|
||||||
|
'RS256', 'RS384', 'RS512',
|
||||||
|
'PS256', 'PS384', 'PS512',
|
||||||
|
'ES256', 'ES384', 'ES512',
|
||||||
|
'EdDSA',
|
||||||
|
))
|
||||||
|
http_client: Optional[httpx.AsyncClient] = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
scopes_list = list(self.scopes)
|
||||||
|
if 'openid' not in scopes_list:
|
||||||
|
self.scopes = ('openid', *scopes_list)
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
import base64
|
||||||
|
import hashlib
|
||||||
|
import secrets
|
||||||
|
|
||||||
|
|
||||||
|
def generate_state() -> str:
|
||||||
|
"""Generate a random CSRF state parameter."""
|
||||||
|
return secrets.token_urlsafe(32)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_nonce() -> str:
|
||||||
|
"""Generate a random nonce for replay protection."""
|
||||||
|
return secrets.token_urlsafe(32)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_pkce() -> tuple[str, str]:
|
||||||
|
"""Return a PKCE (code_verifier, code_challenge) pair using S256."""
|
||||||
|
verifier = secrets.token_urlsafe(64)
|
||||||
|
challenge = base64.urlsafe_b64encode(
|
||||||
|
hashlib.sha256(verifier.encode()).digest()
|
||||||
|
).rstrip(b'=').decode()
|
||||||
|
return verifier, challenge
|
||||||
@@ -0,0 +1,334 @@
|
|||||||
|
import base64
|
||||||
|
import json
|
||||||
|
import unittest
|
||||||
|
from time import time
|
||||||
|
from typing import Any, Mapping
|
||||||
|
from urllib.parse import parse_qs, urlparse
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import jwt
|
||||||
|
from cryptography.hazmat.primitives import serialization
|
||||||
|
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||||
|
from pwo import async_test
|
||||||
|
from kaya.core import HttpContext, KayaApp
|
||||||
|
from kaya.session import InMemorySessionStore, SessionMiddleware
|
||||||
|
|
||||||
|
from kaya.oidc import OIDCApp, OIDCConfig
|
||||||
|
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 OIDCAppTest(unittest.TestCase):
|
||||||
|
|
||||||
|
def _build_app(self, fetch_userinfo: bool = False) -> tuple[OIDCApp, MockOIDCProvider, InMemorySessionStore]:
|
||||||
|
provider = MockOIDCProvider()
|
||||||
|
http_client = httpx.AsyncClient(transport=MockTransport(provider))
|
||||||
|
store = InMemorySessionStore()
|
||||||
|
app = KayaApp()
|
||||||
|
session_app = SessionMiddleware(app, 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_app = OIDCApp(session_app, config)
|
||||||
|
return oidc_app, provider, store
|
||||||
|
|
||||||
|
def _setup_routes(self, oidc_app: OIDCApp) -> None:
|
||||||
|
@oidc_app.GET('/')
|
||||||
|
async def home(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(200, 'home')
|
||||||
|
|
||||||
|
@oidc_app.GET('/profile')
|
||||||
|
@oidc_app.require_auth
|
||||||
|
async def profile(ctx: HttpContext) -> None:
|
||||||
|
user = oidc_app.get_user(ctx)
|
||||||
|
if user is None:
|
||||||
|
await ctx.send_empty(401)
|
||||||
|
return
|
||||||
|
await ctx.send_str(200, f'Hello {user.email}')
|
||||||
|
|
||||||
|
@oidc_app.GET('/refresh')
|
||||||
|
@oidc_app.require_auth
|
||||||
|
async def refresh(ctx: HttpContext) -> None:
|
||||||
|
new_token = await oidc_app.refresh_access_token(ctx.session)
|
||||||
|
await ctx.send_str(200, new_token or 'no-token')
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_login_redirect(self) -> None:
|
||||||
|
oidc_app, provider, store = self._build_app()
|
||||||
|
self._setup_routes(oidc_app)
|
||||||
|
transport = httpx.ASGITransport(app=oidc_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:
|
||||||
|
oidc_app, provider, store = self._build_app(fetch_userinfo=True)
|
||||||
|
self._setup_routes(oidc_app)
|
||||||
|
transport = httpx.ASGITransport(app=oidc_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:
|
||||||
|
oidc_app, provider, store = self._build_app()
|
||||||
|
self._setup_routes(oidc_app)
|
||||||
|
transport = httpx.ASGITransport(app=oidc_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:
|
||||||
|
oidc_app, provider, store = self._build_app()
|
||||||
|
self._setup_routes(oidc_app)
|
||||||
|
transport = httpx.ASGITransport(app=oidc_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:
|
||||||
|
oidc_app, provider, store = self._build_app()
|
||||||
|
self._setup_routes(oidc_app)
|
||||||
|
transport = httpx.ASGITransport(app=oidc_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:
|
||||||
|
oidc_app, provider, store = self._build_app()
|
||||||
|
self._setup_routes(oidc_app)
|
||||||
|
transport = httpx.ASGITransport(app=oidc_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)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
unittest.main()
|
||||||
@@ -7,8 +7,8 @@ from ._session import Session
|
|||||||
from ._store import SessionStore
|
from ._store import SessionStore
|
||||||
|
|
||||||
|
|
||||||
type HttpHandler = Callable[[Any, Any], Awaitable[None]]
|
type HttpHandler = Callable[..., Awaitable[None]]
|
||||||
type WebSocketHandler = Callable[[Any, Any], Awaitable[None]]
|
type WebSocketHandler = Callable[..., Awaitable[None]]
|
||||||
type RouteDecorator = Callable[[HttpHandler], HttpHandler]
|
type RouteDecorator = Callable[[HttpHandler], HttpHandler]
|
||||||
type WebSocketDecorator = Callable[[WebSocketHandler], WebSocketHandler]
|
type WebSocketDecorator = Callable[[WebSocketHandler], WebSocketHandler]
|
||||||
type ASGIApp = Callable[
|
type ASGIApp = Callable[
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
kaya-core @ file:./packages/kaya-core
|
kaya-core @ file:./packages/kaya-core
|
||||||
kaya-rsgi @ file:./packages/kaya-rsgi
|
kaya-rsgi @ file:./packages/kaya-rsgi
|
||||||
kaya-session @ file:./packages/kaya-session
|
kaya-session @ file:./packages/kaya-session
|
||||||
|
kaya-oidc @ file:./packages/kaya-oidc
|
||||||
build
|
build
|
||||||
mypy
|
mypy
|
||||||
ipdb
|
ipdb
|
||||||
|
|||||||
@@ -82,11 +82,16 @@ jeepney==0.9.0
|
|||||||
file:./packages/kaya-core
|
file:./packages/kaya-core
|
||||||
# via
|
# via
|
||||||
# -r requirements-dev.in
|
# -r requirements-dev.in
|
||||||
|
# kaya-oidc
|
||||||
# kaya-rsgi
|
# kaya-rsgi
|
||||||
# kaya-session
|
# kaya-session
|
||||||
file:./packages/kaya-rsgi
|
file:./packages/kaya-rsgi
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
file:./packages/kaya-session
|
file:./packages/kaya-session
|
||||||
|
# via
|
||||||
|
# -r requirements-dev.in
|
||||||
|
# kaya-oidc
|
||||||
|
file:./packages/kaya-oidc
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
keyring==25.7.0
|
keyring==25.7.0
|
||||||
# via twine
|
# via twine
|
||||||
|
|||||||
Reference in New Issue
Block a user