diff --git a/.github/workflows/ci-aiohttp.yml b/.github/workflows/ci-aiohttp.yml new file mode 100644 index 0000000..754c94c --- /dev/null +++ b/.github/workflows/ci-aiohttp.yml @@ -0,0 +1,44 @@ +name: aiohttp + +on: + push: + branches: + - '**' + pull_request: + branches: + - '**' + workflow_dispatch: + +permissions: + contents: read + +jobs: + tests: + name: aiohttp ${{ matrix.label }} + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - label: minimum + aiohttp: 'aiohttp==3.13.0' + - label: latest + aiohttp: aiohttp + steps: + - name: Checkout + uses: actions/checkout@v5 + + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: '3.12' + cache: pip + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev,generator]" "${{ matrix.aiohttp }}" + python -m pip check + + - name: Run aiohttp tests + run: python scripts/run_tests.py -q -m aiohttp diff --git a/.github/workflows/ci-httpx.yml b/.github/workflows/ci-httpx.yml new file mode 100644 index 0000000..9b0fc79 --- /dev/null +++ b/.github/workflows/ci-httpx.yml @@ -0,0 +1,44 @@ +name: HTTPX + +on: + push: + branches: + - '**' + pull_request: + branches: + - '**' + workflow_dispatch: + +permissions: + contents: read + +jobs: + tests: + name: HTTPX ${{ matrix.label }} + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - label: minimum + httpx: 'httpx==0.28.0' + - label: latest + httpx: httpx + steps: + - name: Checkout + uses: actions/checkout@v5 + + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: '3.12' + cache: pip + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev,generator]" "${{ matrix.httpx }}" + python -m pip check + + - name: Run HTTPX tests + run: python scripts/run_tests.py -q -m httpx diff --git a/.github/workflows/ci-openapi.yml b/.github/workflows/ci-openapi.yml new file mode 100644 index 0000000..030ccd4 --- /dev/null +++ b/.github/workflows/ci-openapi.yml @@ -0,0 +1,42 @@ +name: OpenAPI Generator + +on: + push: + branches: + - '**' + pull_request: + branches: + - '**' + workflow_dispatch: + +permissions: + contents: read + +jobs: + generated-clients: + name: OpenAPI Generator ${{ matrix['generator-version'] }} + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + generator-version: + - '7.23.0' + - '7.24.0' + steps: + - name: Checkout + uses: actions/checkout@v5 + + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: '3.14' + cache: pip + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev]" "openapi-generator-cli[jdk4py]==${{ matrix['generator-version'] }}" + python -m pip check + + - name: Run generated-client tests + run: python scripts/run_tests.py -q -m generated diff --git a/.github/workflows/ci-python.yml b/.github/workflows/ci-python.yml new file mode 100644 index 0000000..ef3da9e --- /dev/null +++ b/.github/workflows/ci-python.yml @@ -0,0 +1,46 @@ +name: Python + +on: + push: + branches: + - '**' + pull_request: + branches: + - '**' + workflow_dispatch: + +permissions: + contents: read + +jobs: + tests: + name: Python ${{ matrix['python-version'] }} + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - python-version: '3.12' + cryptography: 'cryptography==45.0.1' + - python-version: '3.13' + cryptography: cryptography + - python-version: '3.14' + cryptography: cryptography + steps: + - name: Checkout + uses: actions/checkout@v5 + + - name: Set up Python ${{ matrix['python-version'] }} + uses: actions/setup-python@v6 + with: + python-version: ${{ matrix['python-version'] }} + cache: pip + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev]" "${{ matrix.cryptography }}" + python -m pip check + + - name: Run tests + run: python -m pytest -q tests/unit tests/integration --ignore-glob='**/test_openapi*.py' diff --git a/.github/workflows/ci-urllib3.yml b/.github/workflows/ci-urllib3.yml new file mode 100644 index 0000000..df98eca --- /dev/null +++ b/.github/workflows/ci-urllib3.yml @@ -0,0 +1,44 @@ +name: urllib3 + +on: + push: + branches: + - '**' + pull_request: + branches: + - '**' + workflow_dispatch: + +permissions: + contents: read + +jobs: + tests: + name: urllib3 ${{ matrix.label }} + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - label: minimum + urllib3: 'urllib3==2.1.0' + - label: latest + urllib3: urllib3 + steps: + - name: Checkout + uses: actions/checkout@v5 + + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: '3.12' + cache: pip + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev,generator]" "${{ matrix.urllib3 }}" + python -m pip check + + - name: Run generated urllib3 tests + run: python scripts/run_tests.py -q -m urllib3 diff --git a/.github/workflows/quality.yml b/.github/workflows/quality.yml new file mode 100644 index 0000000..1cefb02 --- /dev/null +++ b/.github/workflows/quality.yml @@ -0,0 +1,47 @@ +name: Quality + +on: + push: + branches: + - '**' + pull_request: + branches: + - '**' + workflow_dispatch: + +permissions: + contents: read + +jobs: + checks: + runs-on: ubuntu-latest + steps: + - name: Checkout + uses: actions/checkout@v5 + + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: '3.12' + cache: pip + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev,generator]" + python -m pip check + + - name: Lint + run: python -m ruff check . + + - name: Type check + run: python -m mypy + + - name: Build distributions + run: python -m build + + - name: Check distributions + run: python -m twine check dist/* + + - name: Check wheel + run: python scripts/check_wheel.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..25cf513 --- /dev/null +++ b/.gitignore @@ -0,0 +1,21 @@ +__pycache__/ +*.py[cod] +.coverage +.env +.env.* +!.env.example +.venv/ +.idea/ +.mypy_cache/ +.pytest_cache/ +.ruff_cache/ +.cache/ +build/ +dist/ +*.egg-info/ +docs/ +.github/agents/ +.github/ai-framework/ +.github/skills/ +.github/copilot-instructions.md +tests/generated/ \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..8674164 --- /dev/null +++ b/LICENSE @@ -0,0 +1,13 @@ +Copyright 2025 - 2026 Mastercard + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..c3135e7 --- /dev/null +++ b/README.md @@ -0,0 +1,310 @@ +# OAuth2 Client Python + + + + Mastercard Developers + + +## Overview + +Authenticate requests to Mastercard APIs with OAuth 2.0, FAPI 2.0 and DPoP. +The library obtains and reuses access tokens, adds the required authentication +headers and works with supported native and OpenAPI-generated HTTP clients. + +For more information, see +[Using OAuth 2.0 to Access Mastercard APIs](https://mstr.cd/43CuHBY). + +## Requirements + +### License + +[![License](https://img.shields.io/badge/License-Apache%202.0-blue.svg)](LICENSE) + +This project is licensed under the Apache License 2.0. + +### Python + +[![Python 3.12+](https://img.shields.io/badge/Python-3.12%2B-3776AB.svg?logo=python&logoColor=white)](.github/workflows/ci-python.yml) + +Python 3.12 or newer is required. + +### Dependencies + +The core package requires only `cryptography`. Optional integrations are +available through extras shown with each supported client below. + +## Usage + +### Installation + +Install the core package: + +```bash +python -m pip install mastercard-oauth2-client +``` + +For HTTPX: + +```bash +python -m pip install "mastercard-oauth2-client[httpx]" +``` + +For aiohttp: + +```bash +python -m pip install "mastercard-oauth2-client[aiohttp]" +``` + +### Configuration + +`OAuth2Config` contains the client credentials, token endpoint, issuer, scopes +and DPoP key configuration. Load your client signing key from PEM, DER or +PKCS#12 and your DPoP key pair from a private JWK: + +```python +from mastercard_oauth2_client import ( + OAuth2Config, + StaticDPoPKeyProvider, + StaticScopeResolver, + load_jwk_key_pair, + load_private_key, +) + +client_key = load_private_key("path/to/client-private-key.pem") +dpop_key_pair = load_jwk_key_pair("path/to/dpop-private-key.json") + +oauth2_config = OAuth2Config( + client_id="your-client-id", + token_endpoint="https://sandbox.api.mastercard.com/oauth/token", + issuer="https://sandbox.api.mastercard.com", + client_key=client_key, + key_id="your-signing-key-id", + scope_resolver=StaticScopeResolver({"your-api:scope"}), + dpop_key_provider=StaticDPoPKeyProvider(dpop_key_pair), +) +``` + +Replace every example value with the credentials and scopes provisioned for +your Mastercard API project. For custom scope resolution, DPoP key management +and token storage, see [Extension Points](#extension-points). + +### Low-Level API + +[`OAuth2Handler`](src/mastercard_oauth2_client/handlers/sync.py) and +[`AsyncOAuth2Handler`](src/mastercard_oauth2_client/handlers/async_.py) run the +OAuth2 and DPoP flow. To support another HTTP client, provide the matching +[`SyncHttpAdapter` or `AsyncHttpAdapter`](src/mastercard_oauth2_client/protocols.py) +and pass it to the handler. For higher-level use, choose a ready-made +integration below. + +### Direct HTTP Client Integration + +For a higher-level experience, use an HTTP client integration to authenticate +requests automatically. Choose the client that best fits your application. +Each integration provides the same OAuth2 and DPoP behavior. + +| HTTP client | Mode | Install | +|---|---|---| +| HTTPX | Sync and async | `mastercard-oauth2-client[httpx]` | +| aiohttp | Async | `mastercard-oauth2-client[aiohttp]` | + +#### HTTPX + +Install `mastercard-oauth2-client[httpx]`. + +`OAuth2Transport` supports `httpx.Client`: + +```python +import httpx + +from mastercard_oauth2_client.integrations.httpx import OAuth2Transport + +with httpx.Client(transport=OAuth2Transport(oauth2_config)) as client: + response = client.get("https://api.mastercard.com/resources") + response.raise_for_status() +``` + +For asynchronous applications, use `AsyncOAuth2Transport` with +`httpx.AsyncClient`: + +```python +import httpx + +from mastercard_oauth2_client.integrations.httpx import AsyncOAuth2Transport + +async with httpx.AsyncClient(transport=AsyncOAuth2Transport(oauth2_config)) as client: + response = await client.get("https://api.mastercard.com/resources") + response.raise_for_status() +``` + +Pass a configured HTTPX transport to the wrapper for custom network settings. + +#### aiohttp + +Install `mastercard-oauth2-client[aiohttp]`, then add `OAuth2Middleware` when +creating the session: + +```python +import aiohttp + +from mastercard_oauth2_client.integrations.aiohttp import OAuth2Middleware + +async with aiohttp.ClientSession( + middlewares=(OAuth2Middleware(oauth2_config),), +) as client: + response = await client.get("https://api.mastercard.com/resources") + response.raise_for_status() + result = await response.json() +``` + +Use normal `ClientSession` options for network settings. + +### OpenAPI Generated Clients + +The library integrates with Python clients produced by OpenAPI Generator. +Attach OAuth2 before the generated client sends its first request. + +| Generated library | Mode | Install | +|---|---|---| +| Default (`urllib3`) | Sync | `mastercard-oauth2-client` | +| `httpx` | Async | `mastercard-oauth2-client[httpx]` | +| `asyncio` | Async | `mastercard-oauth2-client[aiohttp]` | + +#### Python Generator (urllib3) + +The Python generator uses urllib3 when no `--library` option is supplied. + +```python +from generated_client import ApiClient, Configuration, ResourcesApi +from mastercard_oauth2_client import add_oauth2_layer + +api_client = ApiClient(Configuration(host="https://api.mastercard.com")) +add_oauth2_layer(api_client, oauth2_config) + +resources_api = ResourcesApi(api_client) +resources = resources_api.get_resources() +``` + +This adapter is specific to OpenAPI-generated clients. It is not a direct +urllib3 integration. + +#### Python Generator (`httpx`) + +Generate with `--library httpx` and install the `httpx` extra. + +```python +from generated_client import ApiClient, Configuration, ResourcesApi +from mastercard_oauth2_client.integrations.openapi.httpx import add_oauth2_layer + +api_client = ApiClient(Configuration(host="https://api.mastercard.com")) +add_oauth2_layer(api_client, oauth2_config) + +try: + resources_api = ResourcesApi(api_client) + resources = await resources_api.get_resources() +finally: + await api_client.close() +``` + +#### Python Generator (`asyncio`) + +Generate with `--library asyncio` and install the `aiohttp` extra. + +```python +from generated_client import ApiClient, Configuration, ResourcesApi +from mastercard_oauth2_client.integrations.openapi.asyncio import add_oauth2_layer + +api_client = ApiClient(Configuration(host="https://api.mastercard.com")) +add_oauth2_layer(api_client, oauth2_config) + +try: + resources_api = ResourcesApi(api_client) + resources = await resources_api.get_resources() +finally: + await api_client.close() +``` + +## Extension Points + +This library is small by design and can be extended for different application +needs. + +- [`OAuth2Config`](src/mastercard_oauth2_client/config.py) connects credentials, + endpoints and extension points. +- Use `StaticScopeResolver` for fixed scopes, or provide a custom + [`ScopeResolver`](src/mastercard_oauth2_client/protocols.py) to choose scopes + for each request. +- Use `StaticDPoPKeyProvider` for one key pair, or provide a custom + [`DPoPKeyProvider`](src/mastercard_oauth2_client/protocols.py) to supply or + rotate keys. +- Use the default `InMemoryTokenStore`, or provide a custom + [`TokenStore`](src/mastercard_oauth2_client/protocols.py) to control token + storage. +- [Key loading](src/mastercard_oauth2_client/keys.py): loads PEM, DER, PKCS#12 + and private JWK key material. + +## Key Support + +The library supports RSA and EC keys: + +- Load client private keys from PEM, DER or PKCS#12 files. +- Load DPoP key pairs from private JWKs. +- Use RSA keys of at least 2048 bits for PS256 signing. +- Use P-256 EC keys for ES256 signing. + +## Troubleshooting + +### Configuration + +The token endpoint must be an absolute HTTPS URL. Configuration errors raise +`OAuth2ConfigError` with the invalid setting. + +### Logging + +Use Python's standard logging configuration to enable messages from +`mastercard_oauth2_client`. Sensitive tokens, assertions, proofs, nonces, keys +and response bodies are not logged. + +## Test Strategy + +- [`tests/unit`](tests/unit) contains fast tests for protocol rules, caching, + handlers and adapters. +- [`tests/integration`](tests/integration) tests direct and generated clients + against local HTTPS authorization and resource servers. +- [`fake-api.yaml`](tests/resources/openapi/fake-api.yaml) defines the API used + to generate test clients. + +Pytest markers group each integration's unit and integration tests for focused +CI runs. Tests added to an existing marked module inherit its marker. + +## Development + +Install the development tools and pinned default OpenAPI Generator, then run +the quality gates: + +```bash +python -m pip install -e ".[dev,generator]" +python scripts/run_tests.py +python -m ruff check . +python -m mypy +python -m build +``` + +`scripts/run_tests.py` regenerates all OpenAPI test clients before running +pytest. To regenerate without running tests: + +```bash +python scripts/generate_test_clients.py +``` + +Tested Python, HTTP client and generator combinations are defined in the +[CI workflows](.github/workflows). + +### Code Style + +[![Ruff](https://img.shields.io/badge/lint-Ruff-D7FF64.svg?logo=ruff&logoColor=black)](https://docs.astral.sh/ruff/) +[![mypy](https://www.mypy-lang.org/static/mypy_badge.svg)](https://www.mypy-lang.org/) + +[Ruff](https://docs.astral.sh/ruff/) checks code quality. +[Mypy](https://www.mypy-lang.org/) checks the package and test scripts in strict +mode. \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..f43d300 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,65 @@ +[build-system] +requires = ["hatchling>=1.27"] +build-backend = "hatchling.build" + +[project] +name = "mastercard-oauth2-client" +version = "1.0.0" +description = "Mastercard OAuth 2.0 client authentication for Python" +readme = "README.md" +requires-python = ">=3.12" +license = { text = "Apache-2.0" } +dependencies = [ + "cryptography>=45.0.1", +] + +[project.optional-dependencies] +dev = [ + "aiohttp>=3.13,<4", + "aiohttp-retry>=2.9,<3", + "build>=1.2.2", + "mypy>=1.17", + "pydantic>=2.11", + "pytest>=8.4", + "python-dateutil>=2.8.2", + "ruff>=0.12", + "twine>=6.1", + "typing-extensions>=4.14", + "urllib3>=2.1,<3", + "httpx>=0.28,<1", +] +httpx = [ + "httpx>=0.28,<1", +] +aiohttp = [ + "aiohttp>=3.13,<4", +] +generator = [ + "openapi-generator-cli[jdk4py]==7.23.0", +] + +[tool.hatch.build.targets.wheel] +packages = ["src/mastercard_oauth2_client"] + +[tool.pytest.ini_options] +testpaths = ["tests"] +addopts = "--strict-markers --strict-config" +pythonpath = ["src", "tests/generated/openapi", "tests/generated/openapi_httpx", "tests/generated/openapi_asyncio"] +markers = [ + "aiohttp: aiohttp direct and generated-client integration tests", + "generated: OpenAPI-generated client tests", + "httpx: HTTPX direct and generated-client integration tests", + "urllib3: generated urllib3 client tests", +] + +[tool.ruff] +target-version = "py312" +line-length = 120 +extend-exclude = ["tests/generated"] + +[tool.mypy] +python_version = "3.12" +strict = true +files = ["src/mastercard_oauth2_client", "scripts"] +mypy_path = "src" +explicit_package_bases = true \ No newline at end of file diff --git a/scripts/check_wheel.py b/scripts/check_wheel.py new file mode 100644 index 0000000..51ca0e8 --- /dev/null +++ b/scripts/check_wheel.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +import subprocess +import sys +import tempfile +import venv +import zipfile +from collections.abc import Iterable +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] + + +def main() -> None: + wheel = next((ROOT / "dist").glob("*.whl"), None) + if wheel is None: + raise SystemExit("No wheel found in dist. Run 'python -m build' first.") + + _check_contents(wheel) + _check_install(wheel, "core", absent=("httpx", "aiohttp")) + _check_install(wheel, "httpx", imports=("mastercard_oauth2_client.integrations.httpx",)) + _check_install(wheel, "aiohttp", imports=("mastercard_oauth2_client.integrations.aiohttp",)) + + +def _check_contents(wheel: Path) -> None: + with zipfile.ZipFile(wheel) as archive: + names = set(archive.namelist()) + if "mastercard_oauth2_client/py.typed" not in names: + raise SystemExit("Wheel is missing mastercard_oauth2_client/py.typed") + unexpected = sorted(name for name in names if name.startswith(("tests/", "scripts/"))) + if unexpected: + raise SystemExit(f"Wheel contains development files: {unexpected}") + + +def _check_install( + wheel: Path, + extra: str, + *, + imports: Iterable[str] = ("mastercard_oauth2_client",), + absent: Iterable[str] = (), +) -> None: + with tempfile.TemporaryDirectory(prefix=f"oauth2-wheel-{extra}-") as directory: + environment = Path(directory) + venv.EnvBuilder(with_pip=True).create(environment) + python = environment / ("Scripts/python.exe" if sys.platform == "win32" else "bin/python") + requirement = f"{wheel}[{extra}]" if extra != "core" else str(wheel) + subprocess.run([str(python), "-m", "pip", "install", requirement], check=True) + subprocess.run([str(python), "-m", "pip", "check"], check=True) + for module in imports: + subprocess.run([str(python), "-c", f"import {module}"], check=True) + for module in absent: + command = f"import importlib.util; assert importlib.util.find_spec('{module}') is None" + subprocess.run([str(python), "-c", command], check=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_test_clients.py b/scripts/generate_test_clients.py new file mode 100644 index 0000000..433dceb --- /dev/null +++ b/scripts/generate_test_clients.py @@ -0,0 +1,70 @@ +from __future__ import annotations + +import shutil +from pathlib import Path + +import openapi_generator_cli + +ROOT = Path(__file__).resolve().parents[1] +SPECIFICATION = ROOT / "tests" / "resources" / "openapi" / "fake-api.yaml" +GENERATED_ROOT = ROOT / "tests" / "generated" +GENERATED_CLIENTS = ( + (GENERATED_ROOT / "openapi", "generated_test_client", "oauth2-generated-test-client", None), + ( + GENERATED_ROOT / "openapi_httpx", + "generated_httpx_test_client", + "oauth2-generated-httpx-test-client", + "httpx", + ), + ( + GENERATED_ROOT / "openapi_asyncio", + "generated_asyncio_test_client", + "oauth2-generated-asyncio-test-client", + "asyncio", + ), +) + + +def main() -> None: + if not SPECIFICATION.is_file(): + raise SystemExit(f"OpenAPI specification not found: {SPECIFICATION}") + + for output, package_name, project_name, library in GENERATED_CLIENTS: + _generate_client(output, package_name, project_name, library) + + +def _generate_client( + output: Path, + package_name: str, + project_name: str, + library: str | None, +) -> None: + properties = f"packageName={package_name},projectName={project_name},generateSourceCodeOnly=true" + arguments = [ + "generate", + "-i", + str(SPECIFICATION), + "-g", + "python", + "-o", + str(output), + ] + if library is not None: + arguments.extend(("--library", library)) + arguments.extend( + ( + f"--additional-properties={properties}", + "--global-property=apiTests=false,modelTests=false,apiDocs=false,modelDocs=false", + ) + ) + + shutil.rmtree(output, ignore_errors=True) + result = openapi_generator_cli.run(arguments) + if result.returncode != 0: + raise SystemExit(f"OpenAPI generation failed for {package_name}") + (output / f"{package_name}_README.md").unlink(missing_ok=True) + print(f"Generated {package_name}") + + +if __name__ == "__main__": + main() diff --git a/scripts/run_tests.py b/scripts/run_tests.py new file mode 100644 index 0000000..2b470e8 --- /dev/null +++ b/scripts/run_tests.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] + + +def main() -> None: + subprocess.run([sys.executable, str(ROOT / "scripts" / "generate_test_clients.py")], check=True) + result = subprocess.run([sys.executable, "-m", "pytest", *sys.argv[1:]], cwd=ROOT, check=False) + raise SystemExit(result.returncode) + + +if __name__ == "__main__": + main() diff --git a/src/mastercard_oauth2_client/__init__.py b/src/mastercard_oauth2_client/__init__.py new file mode 100644 index 0000000..0d06c71 --- /dev/null +++ b/src/mastercard_oauth2_client/__init__.py @@ -0,0 +1,60 @@ +from .config import OAuth2Config, SecurityProfile +from .dpop import create_client_assertion, create_resource_dpop_proof, create_token_dpop_proof +from .exceptions import OAuth2ConfigError, OAuth2Error +from .handlers import AsyncOAuth2Handler, OAuth2Handler +from .integrations.openapi import add_oauth2_layer +from .keys import ( + load_jwk_key_pair, + load_pkcs12_private_key, + load_private_key, + validate_dpop_key, + validate_key, +) +from .models import AccessToken, AccessTokenFilter, DPoPKey, HttpRequest, HttpResponse, KeyPair, PrivateKey, PublicKey +from .protocols import ( + AsyncHttpAdapter, + DPoPKeyProvider, + ScopeResolver, + StaticDPoPKeyProvider, + StaticScopeResolver, + SyncHttpAdapter, + TokenStore, +) +from .store import InMemoryTokenStore +from .token import build_token_request, parse_token_response + +__all__ = [ + "AccessToken", + "AccessTokenFilter", + "AsyncHttpAdapter", + "AsyncOAuth2Handler", + "DPoPKey", + "DPoPKeyProvider", + "HttpRequest", + "HttpResponse", + "InMemoryTokenStore", + "KeyPair", + "OAuth2Config", + "OAuth2ConfigError", + "OAuth2Error", + "OAuth2Handler", + "PrivateKey", + "PublicKey", + "ScopeResolver", + "SecurityProfile", + "StaticDPoPKeyProvider", + "StaticScopeResolver", + "SyncHttpAdapter", + "TokenStore", + "add_oauth2_layer", + "build_token_request", + "create_client_assertion", + "create_resource_dpop_proof", + "create_token_dpop_proof", + "load_jwk_key_pair", + "load_pkcs12_private_key", + "load_private_key", + "parse_token_response", + "validate_dpop_key", + "validate_key", +] diff --git a/src/mastercard_oauth2_client/_internal/__init__.py b/src/mastercard_oauth2_client/_internal/__init__.py new file mode 100644 index 0000000..38905d5 --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/__init__.py @@ -0,0 +1 @@ +"""Private implementation modules for mastercard_oauth2_client.""" diff --git a/src/mastercard_oauth2_client/_internal/_clock.py b/src/mastercard_oauth2_client/_internal/_clock.py new file mode 100644 index 0000000..2e400ce --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/_clock.py @@ -0,0 +1,10 @@ +from collections.abc import Callable +from datetime import UTC, datetime + +type NowProvider = Callable[[], datetime] + + +def utc_now() -> datetime: + """Return the current timezone-aware UTC time.""" + + return datetime.now(UTC) diff --git a/src/mastercard_oauth2_client/_internal/_encoding.py b/src/mastercard_oauth2_client/_internal/_encoding.py new file mode 100644 index 0000000..c507d7e --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/_encoding.py @@ -0,0 +1,7 @@ +import base64 + + +def base64url_encode(value: bytes) -> str: + """Encode bytes with unpadded Base64url as required by JOSE.""" + + return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii") diff --git a/src/mastercard_oauth2_client/_internal/_scopes.py b/src/mastercard_oauth2_client/_internal/_scopes.py new file mode 100644 index 0000000..8ea8358 --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/_scopes.py @@ -0,0 +1,7 @@ +from collections.abc import Set as AbstractSet + + +def normalize_scopes(scopes: AbstractSet[str]) -> frozenset[str]: + """Return non-empty scopes with surrounding whitespace removed.""" + + return frozenset(scope.strip() for scope in scopes if scope.strip()) diff --git a/src/mastercard_oauth2_client/_internal/http.py b/src/mastercard_oauth2_client/_internal/http.py new file mode 100644 index 0000000..32ac581 --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/http.py @@ -0,0 +1,31 @@ +from collections.abc import Mapping +from urllib.parse import urlsplit + +_HTTP_SUCCESS_MIN = 200 +_HTTP_SUCCESS_MAX_EXCLUSIVE = 300 + + +def is_success_status(status: int) -> bool: + """Return whether an HTTP status is in the successful 2xx range.""" + + return _HTTP_SUCCESS_MIN <= status < _HTTP_SUCCESS_MAX_EXCLUSIVE + + +def is_absolute_https_url(url: str) -> bool: + """Return whether a URL has an HTTPS scheme and network location.""" + + try: + parsed = urlsplit(url) + return parsed.scheme.casefold() == "https" and bool(parsed.netloc) and parsed.hostname is not None + except ValueError: + return False + + +def get_header(headers: Mapping[str, str], name: str) -> str | None: + """Read an HTTP header using case-insensitive field-name matching.""" + + normalized_name = name.casefold() + for header_name, value in headers.items(): + if header_name.casefold() == normalized_name: + return value + return None diff --git a/src/mastercard_oauth2_client/_internal/jose.py b/src/mastercard_oauth2_client/_internal/jose.py new file mode 100644 index 0000000..90d16b0 --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/jose.py @@ -0,0 +1,254 @@ +import base64 +import hashlib +import json +from collections.abc import Mapping +from typing import Literal + +from cryptography.hazmat.primitives import hashes +from cryptography.hazmat.primitives.asymmetric import ec, padding, rsa, utils + +from ..exceptions import OAuth2ConfigError, OAuth2Error +from ..models import KeyPair, PrivateKey, PublicKey +from ._encoding import base64url_encode as _base64url_encode + +type JwsAlgorithm = Literal["PS256", "ES256"] + +_ALGORITHM_PS256: JwsAlgorithm = "PS256" +_ALGORITHM_ES256: JwsAlgorithm = "ES256" +_KEY_TYPE_RSA = "RSA" +_KEY_TYPE_EC = "EC" +_CURVE_P256 = "P-256" + + +def _base64url_uint(value: int, size: int | None = None) -> str: + """Encode an unsigned JWK integer using minimal or fixed-width bytes.""" + + length = size if size is not None else max(1, (value.bit_length() + 7) // 8) + return _base64url_encode(value.to_bytes(length, "big")) + + +def _decode_base64url_uint(value: object) -> int: + """Decode an unsigned Base64url JWK integer.""" + + if not isinstance(value, str) or not value: + raise ValueError("JWK member must be a non-empty string") + encoded = value.encode("ascii") + return int.from_bytes(base64.b64decode(encoded + b"=" * (-len(encoded) % 4), altchars=b"-_", validate=True), "big") + + +def _json_encode(value: Mapping[str, object]) -> bytes: + """Serialize JOSE JSON without insignificant whitespace.""" + + return json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + + +def _require_jwk_members( + value: Mapping[str, str | list[str]], + key_type: str, + members: tuple[str, ...], +) -> tuple[str, ...]: + """Return required string members or raise a clear validation error.""" + + member_values: list[str] = [] + for member in members: + member_value = value.get(member) + if not isinstance(member_value, str) or not member_value: + names = ", ".join(members) + raise OAuth2ConfigError(f"Missing required {key_type} JWK parameters ({names})") + member_values.append(member_value) + return tuple(member_values) + + +def _curve_name(curve: ec.EllipticCurve) -> str: + """Return the RFC 7518 JWK name for the supported P-256 curve.""" + + if isinstance(curve, ec.SECP256R1): + return _CURVE_P256 + raise OAuth2ConfigError(f"Unsupported curve: {curve.name}") + + +def algorithm_for_key(key: object) -> JwsAlgorithm: + """Select the supported JOSE signing algorithm for a native key.""" + + if isinstance(key, rsa.RSAPrivateKey | rsa.RSAPublicKey): + return _ALGORITHM_PS256 + if isinstance(key, ec.EllipticCurvePrivateKey | ec.EllipticCurvePublicKey): + return _ALGORITHM_ES256 + raise OAuth2ConfigError(f"Key algorithm must be RSA or EC, but was: {type(key).__name__}") + + +def _import_rsa_jwk(value: Mapping[str, str | list[str]]) -> KeyPair: + """Construct native RSA keys from RFC 7518 private JWK members.""" + + modulus_value, public_exponent_value, private_exponent_value = _require_jwk_members( + value, + _KEY_TYPE_RSA, + ("n", "e", "d"), + ) + modulus = _decode_base64url_uint(modulus_value) + public_exponent = _decode_base64url_uint(public_exponent_value) + private_exponent = _decode_base64url_uint(private_exponent_value) + + # cryptography requires CRT values, so derive them from the minimal n, e and d members. + p, q = rsa.rsa_recover_prime_factors(modulus, public_exponent, private_exponent) + rsa_public_numbers = rsa.RSAPublicNumbers(e=public_exponent, n=modulus) + private_key = rsa.RSAPrivateNumbers( + p=p, + q=q, + d=private_exponent, + dmp1=rsa.rsa_crt_dmp1(private_exponent, p), + dmq1=rsa.rsa_crt_dmq1(private_exponent, q), + iqmp=rsa.rsa_crt_iqmp(p, q), + public_numbers=rsa_public_numbers, + ).private_key() + return KeyPair(private_key=private_key, public_key=private_key.public_key()) + + +def _import_ec_jwk(value: Mapping[str, str | list[str]]) -> KeyPair: + """Construct native P-256 keys from RFC 7518 private JWK members.""" + + curve_name, x_value, y_value, private_value = _require_jwk_members( + value, + _KEY_TYPE_EC, + ("crv", "x", "y", "d"), + ) + if curve_name != _CURVE_P256: + raise OAuth2ConfigError(f"Unsupported curve: {curve_name}") + + x = _decode_base64url_uint(x_value) + y = _decode_base64url_uint(y_value) + d = _decode_base64url_uint(private_value) + ec_public_numbers = ec.EllipticCurvePublicNumbers(x=x, y=y, curve=ec.SECP256R1()) + private_key = ec.EllipticCurvePrivateNumbers( + private_value=d, + public_numbers=ec_public_numbers, + ).private_key() + return KeyPair(private_key=private_key, public_key=private_key.public_key()) + + +def import_jwk_key_pair(value: Mapping[str, str | list[str]]) -> KeyPair: + """Import an RFC 7517 RSA or EC private JWK as native keys.""" + + try: + key_type = value.get("kty") + if key_type == _KEY_TYPE_RSA: + return _import_rsa_jwk(value) + if key_type == _KEY_TYPE_EC: + return _import_ec_jwk(value) + if key_type is None: + raise OAuth2ConfigError("Missing required JWK parameter: kty") + raise OAuth2ConfigError(f"Unsupported key type: {key_type}") + except OAuth2ConfigError: + raise + except (TypeError, UnicodeError, ValueError) as error: + raise OAuth2ConfigError("Unable to load JWK key pair") from error + + +def public_jwk(key: PublicKey) -> dict[str, str]: + """Export the RFC 7517 public JWK members required by DPoP.""" + + if isinstance(key, rsa.RSAPublicKey): + rsa_numbers = key.public_numbers() + return { + "kty": _KEY_TYPE_RSA, + "n": _base64url_uint(rsa_numbers.n), + "e": _base64url_uint(rsa_numbers.e), + } + if isinstance(key, ec.EllipticCurvePublicKey): + curve_name = _curve_name(key.curve) + ec_numbers = key.public_numbers() + component_size = (key.curve.key_size + 7) // 8 + return { + "kty": _KEY_TYPE_EC, + "crv": curve_name, + "x": _base64url_uint(ec_numbers.x, component_size), + "y": _base64url_uint(ec_numbers.y, component_size), + } + raise OAuth2Error(f"Unsupported public key type: {type(key).__name__}") + + +def jwk_thumbprint(key: PublicKey) -> str: + """Return the RFC 7638 SHA-256 thumbprint of a public key.""" + + canonical_jwk = json.dumps(public_jwk(key), separators=(",", ":"), sort_keys=True).encode("ascii") + return _base64url_encode(hashlib.sha256(canonical_jwk).digest()) + + +def _get_signing_input(header: Mapping[str, object], claims: Mapping[str, object]) -> str: + """Build the RFC 7515 protected-header and payload signing input.""" + + encoded_header = _base64url_encode(_json_encode(header)) + encoded_claims = _base64url_encode(_json_encode(claims)) + return f"{encoded_header}.{encoded_claims}" + + +def _sign_ps256(data: bytes, private_key: PrivateKey) -> bytes: + """Sign data using RSASSA-PSS with SHA-256 and a 32-byte salt.""" + + if not isinstance(private_key, rsa.RSAPrivateKey): + raise OAuth2ConfigError("PS256 requires an RSA private key") + return private_key.sign( + data, + padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=32), + hashes.SHA256(), + ) + + +def _encode_p1363_signature(r: int, s: int, component_size: int) -> bytes: + """Encode ECDSA components as the fixed-width P1363 r || s format used by JOSE.""" + + encoded_r = r.to_bytes(component_size, "big") + encoded_s = s.to_bytes(component_size, "big") + return encoded_r + encoded_s + + +def _sign_es256(data: bytes, private_key: PrivateKey) -> bytes: + """Sign data using ECDSA with P-256 and SHA-256, returning P1363 bytes.""" + + if not isinstance(private_key, ec.EllipticCurvePrivateKey): + raise OAuth2ConfigError("ES256 requires an EC private key") + _curve_name(private_key.curve) + + # cryptography returns ASN.1 DER; JOSE requires the P1363 r || s representation. + der_signature = private_key.sign(data, ec.ECDSA(hashes.SHA256())) + r, s = utils.decode_dss_signature(der_signature) + component_size = (private_key.curve.key_size + 7) // 8 + return _encode_p1363_signature(r, s, component_size) + + +def _sign_bytes(data: bytes, private_key: PrivateKey, algorithm: JwsAlgorithm) -> bytes: + """Sign bytes with one of the supported FAPI 2.0 JWS algorithms.""" + + if algorithm == _ALGORITHM_PS256: + return _sign_ps256(data, private_key) + if algorithm == _ALGORITHM_ES256: + return _sign_es256(data, private_key) + raise OAuth2ConfigError(f"Unsupported algorithm: {algorithm}") + + +def _serialize_jwt(signing_input: str, signature: bytes) -> str: + """Create an RFC 7515 compact JWS from its signing input and signature.""" + + if not signature: + raise OAuth2Error("Signature is required") + return f"{signing_input}.{_base64url_encode(signature)}" + + +def sign_jwt( + header: dict[str, object], + claims: dict[str, object], + private_key: PrivateKey, + algorithm: JwsAlgorithm | None = None, +) -> str: + """Create an RFC 7519 JWT using RFC 7515 compact JWS serialization.""" + + selected_algorithm = algorithm or algorithm_for_key(private_key) + protected_header = {**header, "alg": selected_algorithm} + try: + signing_input = _get_signing_input(protected_header, claims) + signature = _sign_bytes(signing_input.encode("ascii"), private_key, selected_algorithm) + return _serialize_jwt(signing_input, signature) + except OAuth2Error: + raise + except Exception as error: + raise OAuth2Error("Unable to sign JWT") from error diff --git a/src/mastercard_oauth2_client/_internal/orchestration.py b/src/mastercard_oauth2_client/_internal/orchestration.py new file mode 100644 index 0000000..a238e72 --- /dev/null +++ b/src/mastercard_oauth2_client/_internal/orchestration.py @@ -0,0 +1,109 @@ +import json +from collections.abc import Mapping +from threading import RLock +from urllib.parse import urlsplit, urlunsplit + +from .._internal.http import get_header, is_absolute_https_url +from .._internal.jose import jwk_thumbprint +from ..config import OAuth2Config +from ..dpop import create_resource_dpop_proof +from ..exceptions import OAuth2Error +from ..models import AccessToken, AccessTokenFilter, DPoPKey, HttpResponse + +_NONCE_CHALLENGE_STATUSES = frozenset({400, 401}) +_AUTHORIZATION_HEADER = "Authorization" +_DPOP_HEADER = "DPoP" +_DPOP_NONCE_HEADER = "DPoP-Nonce" +_USER_AGENT_HEADER = "User-Agent" +_WWW_AUTHENTICATE_HEADER = "WWW-Authenticate" +_USE_DPOP_NONCE = "use_dpop_nonce" +_OAUTH_CONTROLLED_HEADERS = frozenset({_AUTHORIZATION_HEADER.casefold(), _DPOP_HEADER.casefold()}) + + +class NonceStore: + """Store the client-level DPoP nonce shared by token and resource proofs.""" + + def __init__(self) -> None: + self._nonce: str | None = None + self._lock = RLock() + + def get(self) -> str | None: + with self._lock: + return self._nonce + + def update(self, headers: Mapping[str, str]) -> bool: + nonce = get_header(headers, _DPOP_NONCE_HEADER) + if nonce is None or not nonce.strip(): + return False + with self._lock: + self._nonce = nonce + return True + + +def require_resource_url(url: str) -> None: + if not is_absolute_https_url(url): + raise OAuth2Error(f"FAPI 2.0 requires HTTPS for resource server: {url}") + + +def token_filter(scopes: frozenset[str], dpop_key: DPoPKey) -> AccessTokenFilter: + return AccessTokenFilter(scopes=scopes, jkt=jwk_thumbprint(dpop_key.key_pair.public_key)) + + +def build_resource_headers( + config: OAuth2Config, + existing_headers: Mapping[str, str], + method: str, + url: str, + access_token: AccessToken, + dpop_key: DPoPKey, + nonce: str | None, +) -> dict[str, str]: + headers = { + name: value for name, value in existing_headers.items() if name.casefold() not in _OAUTH_CONTROLLED_HEADERS + } + if get_header(headers, _USER_AGENT_HEADER) is None: + headers[_USER_AGENT_HEADER] = config.user_agent + headers[_AUTHORIZATION_HEADER] = f"DPoP {access_token.token_value}" + headers[_DPOP_HEADER] = create_resource_dpop_proof( + config, + dpop_key.key_id, + method, + url, + access_token.token_value, + nonce, + ) + return headers + + +def has_nonce_challenge_header(status: int, headers: Mapping[str, str]) -> bool: + if not is_nonce_challenge_status(status): + return False + authenticate = get_header(headers, _WWW_AUTHENTICATE_HEADER) + return authenticate is not None and _USE_DPOP_NONCE in authenticate.casefold() + + +def has_nonce_challenge_body(status: int, body: bytes | str | None) -> bool: + return is_nonce_challenge_status(status) and _json_error(body) == _USE_DPOP_NONCE + + +def is_nonce_challenge_status(status: int) -> bool: + return status in _NONCE_CHALLENGE_STATUSES + + +def response_snapshot(status: int, headers: Mapping[str, str], body: bytes | str | None) -> HttpResponse: + return HttpResponse(status=status, headers=headers, body=body) + + +def safe_url(url: str) -> str: + parsed = urlsplit(url) + return urlunsplit((parsed.scheme, parsed.netloc, parsed.path, "", "")) + + +def _json_error(body: bytes | str | None) -> object: + if body is None or not body: + return None + try: + value = json.loads(body) + except (json.JSONDecodeError, UnicodeDecodeError): + return None + return value.get("error") if isinstance(value, dict) else None diff --git a/src/mastercard_oauth2_client/config.py b/src/mastercard_oauth2_client/config.py new file mode 100644 index 0000000..1078b0b --- /dev/null +++ b/src/mastercard_oauth2_client/config.py @@ -0,0 +1,103 @@ +import platform +from dataclasses import dataclass, field +from enum import Enum +from importlib.metadata import PackageNotFoundError, version + +from ._internal.http import is_absolute_https_url +from .exceptions import OAuth2ConfigError +from .keys import validate_dpop_key, validate_private_key +from .models import PrivateKey +from .protocols import DPoPKeyProvider, ScopeResolver, TokenStore +from .store import InMemoryTokenStore + +_DISTRIBUTION_NAME = "mastercard-oauth2-client" +_USER_AGENT_PRODUCT = "Mastercard-OAuth2-Client" +_UNKNOWN_VERSION = "0.0.0-unknown" +_DEFAULT_CLOCK_SKEW_SECONDS = 5 + + +class SecurityProfile(Enum): + """OAuth2 security profiles supported by this library.""" + + FAPI2_PRIVATE_KEY_DPOP = "fapi2-private-key-dpop" + + +def _library_version() -> str: + try: + return version(_DISTRIBUTION_NAME) + except PackageNotFoundError: + return _UNKNOWN_VERSION + + +def default_user_agent() -> str: + """Build the default product identifier sent on OAuth2 HTTP requests.""" + + return ( + f"{_USER_AGENT_PRODUCT}/{_library_version()} " + f"(Python/{platform.python_version()}; {platform.system()} {platform.release()})" + ) + + +@dataclass(frozen=True, slots=True) +class OAuth2Config: + """Immutable inputs and extension points for the OAuth2 client.""" + + client_id: str + token_endpoint: str + issuer: str + client_key: PrivateKey + key_id: str + scope_resolver: ScopeResolver + dpop_key_provider: DPoPKeyProvider + security_profile: SecurityProfile = SecurityProfile.FAPI2_PRIVATE_KEY_DPOP + clock_skew_tolerance: int = _DEFAULT_CLOCK_SKEW_SECONDS + token_store: TokenStore = field(default_factory=InMemoryTokenStore) + user_agent: str = field(default_factory=default_user_agent) + + def __post_init__(self) -> None: + self._validate_required_values() + self._validate_extension_points() + self._validate_settings() + self._validate_signing_keys() + + def _validate_required_values(self) -> None: + client_id = _required_string(self.client_id, "Client ID is required") + object.__setattr__(self, "client_id", client_id) + _required_string(self.token_endpoint, "Token endpoint is required") + if not is_absolute_https_url(self.token_endpoint): + raise OAuth2ConfigError("FAPI 2.0 requires HTTPS token endpoint") + _required_string(self.issuer, "Issuer is required") + if self.client_key is None: + raise OAuth2ConfigError("Client private key is required") + key_id = _required_string(self.key_id, "Key ID (kid) is required") + object.__setattr__(self, "key_id", key_id) + + def _validate_extension_points(self) -> None: + if self.scope_resolver is None: + raise OAuth2ConfigError("Scope resolver is required") + if self.dpop_key_provider is None: + raise OAuth2ConfigError("DPoP key provider is required") + if self.token_store is None: + raise OAuth2ConfigError("Token store is required") + + def _validate_settings(self) -> None: + if self.security_profile is not SecurityProfile.FAPI2_PRIVATE_KEY_DPOP: + raise OAuth2ConfigError("Security profile must be FAPI 2.0 with private_key_jwt and DPoP") + if isinstance(self.clock_skew_tolerance, bool) or not isinstance(self.clock_skew_tolerance, int): + raise OAuth2ConfigError("Clock skew tolerance must be an integer") + if self.clock_skew_tolerance < 0: + raise OAuth2ConfigError("Clock skew tolerance must not be negative") + if not isinstance(self.user_agent, str): + raise OAuth2ConfigError("User agent must be a string") + if not self.user_agent.strip(): + object.__setattr__(self, "user_agent", default_user_agent()) + + def _validate_signing_keys(self) -> None: + validate_private_key(self.client_key, context="Client key") + validate_dpop_key(self.dpop_key_provider.get_current_key()) + + +def _required_string(value: object, message: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise OAuth2ConfigError(message) + return value.strip() diff --git a/src/mastercard_oauth2_client/dpop.py b/src/mastercard_oauth2_client/dpop.py new file mode 100644 index 0000000..d8598b5 --- /dev/null +++ b/src/mastercard_oauth2_client/dpop.py @@ -0,0 +1,149 @@ +import hashlib +import secrets +from collections.abc import Callable +from urllib.parse import urlsplit, urlunsplit + +from ._internal._clock import NowProvider +from ._internal._clock import utc_now as _utc_now +from ._internal._encoding import base64url_encode as _base64url_encode +from ._internal.jose import algorithm_for_key, public_jwk, sign_jwt +from .config import OAuth2Config +from .exceptions import OAuth2Error +from .keys import validate_dpop_key + +type JtiProvider = Callable[[], str] + +_PROOF_LIFETIME_SECONDS = 90 + + +def _random_jti() -> str: + """Generate the 96-bit random jti required by RFC 9449 section 4.2.""" + + return _base64url_encode(secrets.token_bytes(12)) + + +def _canonicalize_htu(url: str) -> str: + """Build the RFC 9449 htu value without query and fragment components.""" + + parsed = urlsplit(url) + if not parsed.scheme or not parsed.hostname: + raise OAuth2Error("Invalid URL for DPoP htu claim") + + hostname = parsed.hostname.lower() + if ":" in hostname: + hostname = f"[{hostname}]" + authority = hostname if parsed.port is None else f"{hostname}:{parsed.port}" + return urlunsplit((parsed.scheme.lower(), authority, parsed.path, "", "")) + + +def _access_token_hash(access_token: str) -> str: + """Compute the RFC 9449 section 4.2 ath claim value.""" + + digest = hashlib.sha256(access_token.encode("utf-8")).digest() + return _base64url_encode(digest) + + +def create_client_assertion(config: OAuth2Config) -> str: + """Create a private_key_jwt client assertion for the token endpoint.""" + + return _create_client_assertion(config, now_provider=_utc_now, jti_provider=_random_jti) + + +def _create_client_assertion( + config: OAuth2Config, + *, + now_provider: NowProvider, + jti_provider: JtiProvider, +) -> str: + issued_at = int(now_provider().timestamp()) + skew = config.clock_skew_tolerance + header: dict[str, object] = {"typ": "JWT", "kid": config.key_id} + claims: dict[str, object] = { + "iss": config.client_id, + "sub": config.client_id, + "aud": config.issuer, + "jti": jti_provider(), + "iat": issued_at, + "nbf": issued_at - skew, + "exp": issued_at + _PROOF_LIFETIME_SECONDS + skew, + } + return sign_jwt(header, claims, config.client_key) + + +def create_token_dpop_proof( + config: OAuth2Config, + dpop_key_id: str, + nonce: str | None = None, +) -> str: + """Create an RFC 9449 DPoP proof for a token endpoint request.""" + + return _create_dpop_proof( + config, + dpop_key_id, + "POST", + config.token_endpoint, + nonce=nonce, + now_provider=_utc_now, + jti_provider=_random_jti, + ) + + +def create_resource_dpop_proof( + config: OAuth2Config, + dpop_key_id: str, + method: str, + url: str, + access_token: str, + nonce: str | None = None, +) -> str: + """Create an RFC 9449 DPoP proof bound to a resource request and token.""" + + return _create_dpop_proof( + config, + dpop_key_id, + method, + url, + access_token=access_token, + nonce=nonce, + now_provider=_utc_now, + jti_provider=_random_jti, + ) + + +def _create_dpop_proof( + config: OAuth2Config, + dpop_key_id: str, + method: str, + url: str, + *, + access_token: str | None = None, + nonce: str | None = None, + now_provider: NowProvider, + jti_provider: JtiProvider, +) -> str: + dpop_key = config.dpop_key_provider.get_key(dpop_key_id) + validate_dpop_key(dpop_key) + issued_at = int(now_provider().timestamp()) + private_key = dpop_key.key_pair.private_key + public_key = dpop_key.key_pair.public_key + algorithm = algorithm_for_key(public_key) + + header: dict[str, object] = { + "typ": "dpop+jwt", + "alg": algorithm, + "kid": dpop_key_id, + "jwk": public_jwk(public_key), + } + claims: dict[str, object] = { + "jti": jti_provider(), + "htm": method.upper(), + "htu": _canonicalize_htu(url), + "iat": issued_at, + "exp": issued_at + _PROOF_LIFETIME_SECONDS + config.clock_skew_tolerance, + } + if access_token is not None: + claims["ath"] = _access_token_hash(access_token) + if nonce is not None: + claims["nonce"] = nonce + + return sign_jwt(header, claims, private_key, algorithm) diff --git a/src/mastercard_oauth2_client/exceptions.py b/src/mastercard_oauth2_client/exceptions.py new file mode 100644 index 0000000..b92832d --- /dev/null +++ b/src/mastercard_oauth2_client/exceptions.py @@ -0,0 +1,6 @@ +class OAuth2Error(Exception): + """Base exception for OAuth2 failures produced by this library.""" + + +class OAuth2ConfigError(OAuth2Error, ValueError): + """Raised when OAuth2 configuration is invalid.""" diff --git a/src/mastercard_oauth2_client/handlers/__init__.py b/src/mastercard_oauth2_client/handlers/__init__.py new file mode 100644 index 0000000..b6584eb --- /dev/null +++ b/src/mastercard_oauth2_client/handlers/__init__.py @@ -0,0 +1,4 @@ +from .async_ import AsyncOAuth2Handler +from .sync import OAuth2Handler + +__all__ = ["AsyncOAuth2Handler", "OAuth2Handler"] diff --git a/src/mastercard_oauth2_client/handlers/async_.py b/src/mastercard_oauth2_client/handlers/async_.py new file mode 100644 index 0000000..89b91f0 --- /dev/null +++ b/src/mastercard_oauth2_client/handlers/async_.py @@ -0,0 +1,178 @@ +import logging +from collections.abc import Mapping + +from .._internal._scopes import normalize_scopes +from .._internal.http import is_success_status +from .._internal.orchestration import ( + NonceStore, + build_resource_headers, + has_nonce_challenge_body, + has_nonce_challenge_header, + is_nonce_challenge_status, + require_resource_url, + response_snapshot, + safe_url, + token_filter, +) +from ..config import OAuth2Config +from ..keys import validate_dpop_key +from ..models import AccessToken, DPoPKey, HttpResponse +from ..protocols import AsyncHttpAdapter +from ..token import build_token_request, parse_token_response + +_LOGGER = logging.getLogger(__name__) + + +class AsyncOAuth2Handler[RequestT, ResponseT]: + """Coordinate asynchronous token acquisition and DPoP resource requests.""" + + def __init__(self, config: OAuth2Config) -> None: + self._config = config + self._nonces = NonceStore() + + async def execute(self, request: RequestT, adapter: AsyncHttpAdapter[RequestT, ResponseT]) -> ResponseT: + """Authenticate and execute one resource request, retrying nonce challenges once.""" + + method = adapter.request_method(request) + url = adapter.request_url(request) + require_resource_url(url) + _LOGGER.info("Authenticating resource request: %s %s", method.upper(), safe_url(url)) + _LOGGER.info("Resolving scopes for resource request") + scopes = normalize_scopes(self._config.scope_resolver.resolve(method, url)) + _LOGGER.debug("Resolved %d scopes", len(scopes)) + + _LOGGER.info("Retrieving DPoP key") + dpop_key = self._config.dpop_key_provider.get_current_key() + validate_dpop_key(dpop_key) + _LOGGER.debug("DPoP key selected") + access_token = await self._get_access_token(request, adapter, scopes, dpop_key) + if not isinstance(access_token, AccessToken): + # Token acquisition failed, so return the native token-endpoint response to the caller. + return access_token + + response = await self._send_resource_request(request, adapter, access_token, dpop_key) + _LOGGER.info("Resource request completed with HTTP %d", adapter.response_status(response)) + return response + + async def _get_access_token( + self, + resource_request: RequestT, + adapter: AsyncHttpAdapter[RequestT, ResponseT], + scopes: frozenset[str], + dpop_key: DPoPKey, + ) -> AccessToken | ResponseT: + """Return a usable token, or the failed native token-endpoint response.""" + + access_token = self._config.token_store.get(token_filter(scopes, dpop_key)) + if access_token is not None: + _LOGGER.info("Access token cache hit") + return access_token + + _LOGGER.info("Access token cache miss; requesting token") + response = await self._send_token_request_with_retry(resource_request, adapter, scopes, dpop_key) + status = adapter.response_status(response) + _LOGGER.info("Token request completed with HTTP %d", status) + if not is_success_status(status): + return response + + try: + snapshot = await self._response_snapshot(adapter, response) + finally: + await adapter.close_response(response) + access_token = parse_token_response(self._config, snapshot, scopes, dpop_key.key_id) + self._config.token_store.put(access_token) + _LOGGER.info("Access token stored") + return access_token + + async def _response_snapshot( + self, + adapter: AsyncHttpAdapter[RequestT, ResponseT], + response: ResponseT, + ) -> HttpResponse: + """Read a native token response into a detached neutral snapshot.""" + + return response_snapshot( + adapter.response_status(response), + adapter.response_headers(response), + await adapter.response_body(response), + ) + + async def _send_token_request_with_retry( + self, + resource_request: RequestT, + adapter: AsyncHttpAdapter[RequestT, ResponseT], + scopes: frozenset[str], + dpop_key: DPoPKey, + ) -> ResponseT: + """Send a token request and retry one nonce challenge with a fresh proof.""" + + request = build_token_request(self._config, scopes, dpop_key.key_id, self._nonces.get()) + response = await adapter.send_token_request(resource_request, request) + self._update_nonce(adapter.response_headers(response)) + if await self._is_nonce_challenge(adapter, response): + await adapter.close_response(response) + _LOGGER.info("DPoP nonce-directed retry for token request") + retry_request = build_token_request(self._config, scopes, dpop_key.key_id, self._nonces.get()) + response = await adapter.send_token_request(resource_request, retry_request) + self._update_nonce(adapter.response_headers(response)) + return response + + async def _send_resource_request( + self, + request: RequestT, + adapter: AsyncHttpAdapter[RequestT, ResponseT], + access_token: AccessToken, + dpop_key: DPoPKey, + ) -> ResponseT: + """Send a resource request and retry one nonce challenge with a fresh proof.""" + + headers = self._resource_headers(request, adapter, access_token, dpop_key) + response = await adapter.send_resource_request(request, headers) + self._update_nonce(adapter.response_headers(response)) + if await self._is_nonce_challenge(adapter, response): + await adapter.close_response(response) + _LOGGER.info("DPoP nonce-directed retry for resource request") + retry_headers = self._resource_headers(request, adapter, access_token, dpop_key) + response = await adapter.send_resource_request(request, retry_headers) + self._update_nonce(adapter.response_headers(response)) + return response + + def _resource_headers( + self, + request: RequestT, + adapter: AsyncHttpAdapter[RequestT, ResponseT], + access_token: AccessToken, + dpop_key: DPoPKey, + ) -> dict[str, str]: + """Build OAuth2 headers without copying or serializing the native request.""" + + return build_resource_headers( + self._config, + adapter.request_headers(request), + adapter.request_method(request), + adapter.request_url(request), + access_token, + dpop_key, + self._nonces.get(), + ) + + def _update_nonce(self, headers: Mapping[str, str]) -> None: + """Store a non-empty DPoP nonce returned in response headers.""" + + if self._nonces.update(headers): + _LOGGER.debug("DPoP nonce received") + + @staticmethod + async def _is_nonce_challenge( + adapter: AsyncHttpAdapter[RequestT, ResponseT], + response: ResponseT, + ) -> bool: + """Detect a nonce challenge without reading successful response bodies.""" + + status = adapter.response_status(response) + headers = adapter.response_headers(response) + if has_nonce_challenge_header(status, headers): + return True + if not is_nonce_challenge_status(status): + return False + return has_nonce_challenge_body(status, await adapter.response_body(response)) diff --git a/src/mastercard_oauth2_client/handlers/sync.py b/src/mastercard_oauth2_client/handlers/sync.py new file mode 100644 index 0000000..78d18f0 --- /dev/null +++ b/src/mastercard_oauth2_client/handlers/sync.py @@ -0,0 +1,175 @@ +import logging +from collections.abc import Mapping + +from .._internal._scopes import normalize_scopes +from .._internal.http import is_success_status +from .._internal.orchestration import ( + NonceStore, + build_resource_headers, + has_nonce_challenge_body, + has_nonce_challenge_header, + is_nonce_challenge_status, + require_resource_url, + response_snapshot, + safe_url, + token_filter, +) +from ..config import OAuth2Config +from ..keys import validate_dpop_key +from ..models import AccessToken, DPoPKey, HttpResponse +from ..protocols import SyncHttpAdapter +from ..token import build_token_request, parse_token_response + +_LOGGER = logging.getLogger(__name__) + + +class OAuth2Handler[RequestT, ResponseT]: + """Coordinate synchronous token acquisition and DPoP resource requests.""" + + def __init__(self, config: OAuth2Config) -> None: + self._config = config + self._nonces = NonceStore() + + def execute(self, request: RequestT, adapter: SyncHttpAdapter[RequestT, ResponseT]) -> ResponseT: + """Authenticate and execute one resource request, retrying nonce challenges once.""" + + method = adapter.request_method(request) + url = adapter.request_url(request) + require_resource_url(url) + _LOGGER.info("Authenticating resource request: %s %s", method.upper(), safe_url(url)) + _LOGGER.info("Resolving scopes for resource request") + scopes = normalize_scopes(self._config.scope_resolver.resolve(method, url)) + _LOGGER.debug("Resolved %d scopes", len(scopes)) + + _LOGGER.info("Retrieving DPoP key") + dpop_key = self._config.dpop_key_provider.get_current_key() + validate_dpop_key(dpop_key) + _LOGGER.debug("DPoP key selected") + access_token = self._get_access_token(request, adapter, scopes, dpop_key) + if not isinstance(access_token, AccessToken): + # Token acquisition failed, so return the native token-endpoint response to the caller. + return access_token + + response = self._send_resource_request(request, adapter, access_token, dpop_key) + _LOGGER.info("Resource request completed with HTTP %d", adapter.response_status(response)) + return response + + def _get_access_token( + self, + resource_request: RequestT, + adapter: SyncHttpAdapter[RequestT, ResponseT], + scopes: frozenset[str], + dpop_key: DPoPKey, + ) -> AccessToken | ResponseT: + """Return a usable token, or the failed native token-endpoint response.""" + + access_token = self._config.token_store.get(token_filter(scopes, dpop_key)) + if access_token is not None: + _LOGGER.info("Access token cache hit") + return access_token + + _LOGGER.info("Access token cache miss; requesting token") + response = self._send_token_request_with_retry(resource_request, adapter, scopes, dpop_key) + status = adapter.response_status(response) + _LOGGER.info("Token request completed with HTTP %d", status) + if not is_success_status(status): + return response + + try: + snapshot = self._response_snapshot(adapter, response) + finally: + adapter.close_response(response) + access_token = parse_token_response(self._config, snapshot, scopes, dpop_key.key_id) + self._config.token_store.put(access_token) + _LOGGER.info("Access token stored") + return access_token + + def _send_token_request_with_retry( + self, + resource_request: RequestT, + adapter: SyncHttpAdapter[RequestT, ResponseT], + scopes: frozenset[str], + dpop_key: DPoPKey, + ) -> ResponseT: + """Send a token request and retry one nonce challenge with a fresh proof.""" + + request = build_token_request(self._config, scopes, dpop_key.key_id, self._nonces.get()) + response = adapter.send_token_request(resource_request, request) + self._update_nonce(adapter.response_headers(response)) + if self._is_nonce_challenge(adapter, response): + adapter.close_response(response) + _LOGGER.info("DPoP nonce-directed retry for token request") + retry_request = build_token_request(self._config, scopes, dpop_key.key_id, self._nonces.get()) + response = adapter.send_token_request(resource_request, retry_request) + self._update_nonce(adapter.response_headers(response)) + return response + + def _send_resource_request( + self, + request: RequestT, + adapter: SyncHttpAdapter[RequestT, ResponseT], + access_token: AccessToken, + dpop_key: DPoPKey, + ) -> ResponseT: + """Send a resource request and retry one nonce challenge with a fresh proof.""" + + headers = self._resource_headers(request, adapter, access_token, dpop_key) + response = adapter.send_resource_request(request, headers) + self._update_nonce(adapter.response_headers(response)) + if self._is_nonce_challenge(adapter, response): + adapter.close_response(response) + _LOGGER.info("DPoP nonce-directed retry for resource request") + retry_headers = self._resource_headers(request, adapter, access_token, dpop_key) + response = adapter.send_resource_request(request, retry_headers) + self._update_nonce(adapter.response_headers(response)) + return response + + def _response_snapshot( + self, + adapter: SyncHttpAdapter[RequestT, ResponseT], + response: ResponseT, + ) -> HttpResponse: + """Read a native token response into a detached neutral snapshot.""" + + return response_snapshot( + adapter.response_status(response), + adapter.response_headers(response), + adapter.response_body(response), + ) + + def _resource_headers( + self, + request: RequestT, + adapter: SyncHttpAdapter[RequestT, ResponseT], + access_token: AccessToken, + dpop_key: DPoPKey, + ) -> dict[str, str]: + """Build OAuth2 headers without copying or serializing the native request.""" + + return build_resource_headers( + self._config, + adapter.request_headers(request), + adapter.request_method(request), + adapter.request_url(request), + access_token, + dpop_key, + self._nonces.get(), + ) + + def _update_nonce(self, headers: Mapping[str, str]) -> None: + """Store a non-empty DPoP nonce returned in response headers.""" + + if self._nonces.update(headers): + _LOGGER.debug("DPoP nonce received") + + @staticmethod + def _is_nonce_challenge(adapter: SyncHttpAdapter[RequestT, ResponseT], response: ResponseT) -> bool: + """Detect a nonce challenge without reading successful response bodies.""" + + status = adapter.response_status(response) + headers = adapter.response_headers(response) + if has_nonce_challenge_header(status, headers): + return True + if not is_nonce_challenge_status(status): + return False + return has_nonce_challenge_body(status, adapter.response_body(response)) diff --git a/src/mastercard_oauth2_client/integrations/__init__.py b/src/mastercard_oauth2_client/integrations/__init__.py new file mode 100644 index 0000000..47b0b6c --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/__init__.py @@ -0,0 +1 @@ +"""Optional integrations for concrete HTTP clients.""" diff --git a/src/mastercard_oauth2_client/integrations/aiohttp.py b/src/mastercard_oauth2_client/integrations/aiohttp.py new file mode 100644 index 0000000..5c1ba4e --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/aiohttp.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +from collections.abc import Mapping + +import aiohttp + +from ..config import OAuth2Config +from ..handlers import AsyncOAuth2Handler +from ..models import HttpRequest + + +class _AiohttpAdapter: + """Map native aiohttp requests and responses into handler operations.""" + + def __init__(self, resource_handler: aiohttp.ClientHandlerType) -> None: + self._resource_handler = resource_handler + + def request_method(self, request: aiohttp.ClientRequest) -> str: + return request.method + + def request_url(self, request: aiohttp.ClientRequest) -> str: + return str(request.url) + + def request_headers(self, request: aiohttp.ClientRequest) -> Mapping[str, str]: + return request.headers + + async def send_token_request( + self, + resource_request: aiohttp.ClientRequest, + request: HttpRequest, + ) -> aiohttp.ClientResponse: + return await resource_request.session.request( + request.method, + request.url, + headers=request.headers, + data=request.body, + ssl=resource_request.ssl, + proxy=resource_request.proxy, + proxy_headers=resource_request.proxy_headers, + raise_for_status=False, + middlewares=(), + ) + + async def send_resource_request( + self, + request: aiohttp.ClientRequest, + headers: Mapping[str, str], + ) -> aiohttp.ClientResponse: + request.headers.clear() + request.headers.update(headers) + return await self._resource_handler(request) + + def response_status(self, response: aiohttp.ClientResponse) -> int: + return response.status + + def response_headers(self, response: aiohttp.ClientResponse) -> Mapping[str, str]: + return response.headers + + async def response_body(self, response: aiohttp.ClientResponse) -> bytes: + return await response.read() + + async def close_response(self, response: aiohttp.ClientResponse) -> None: + await response.read() + response.release() + + +class OAuth2Middleware: + """Authenticate aiohttp client requests with OAuth2 and DPoP.""" + + def __init__(self, config: OAuth2Config) -> None: + self._handler = AsyncOAuth2Handler[aiohttp.ClientRequest, aiohttp.ClientResponse](config) + + async def __call__( + self, + request: aiohttp.ClientRequest, + handler: aiohttp.ClientHandlerType, + ) -> aiohttp.ClientResponse: + """Buffer a replayable request and execute the OAuth2 flow.""" + + if isinstance(request.body, aiohttp.payload.Payload): + body = await request.body.as_bytes() + request.chunked = None + await request.update_body(body) + return await self._handler.execute(request, _AiohttpAdapter(handler)) diff --git a/src/mastercard_oauth2_client/integrations/httpx.py b/src/mastercard_oauth2_client/integrations/httpx.py new file mode 100644 index 0000000..65faf68 --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/httpx.py @@ -0,0 +1,126 @@ +from __future__ import annotations + +from collections.abc import Mapping + +import httpx + +from ..config import OAuth2Config +from ..handlers import AsyncOAuth2Handler, OAuth2Handler +from ..models import HttpRequest + + +class _HttpxAdapterBase: + """Map native HTTPX request and response metadata into handler operations.""" + + def request_method(self, request: httpx.Request) -> str: + return request.method + + def request_url(self, request: httpx.Request) -> str: + return str(request.url) + + def request_headers(self, request: httpx.Request) -> Mapping[str, str]: + return request.headers + + def response_status(self, response: httpx.Response) -> int: + return response.status_code + + def response_headers(self, response: httpx.Response) -> Mapping[str, str]: + return response.headers + + @staticmethod + def _token_request(resource_request: httpx.Request, request: HttpRequest) -> httpx.Request: + return httpx.Request( + method=request.method, + url=request.url, + headers=request.headers, + content=request.body, + extensions=dict(resource_request.extensions), + ) + + @staticmethod + def _authenticate_request(request: httpx.Request, headers: Mapping[str, str]) -> httpx.Request: + # Mutate only OAuth-controlled headers on the buffered native request so + # HTTPX extensions, URL, body stream and caller-visible request identity remain intact. + request.headers.clear() + request.headers.update(headers) + return request + + +class _SyncHttpxAdapter(_HttpxAdapterBase): + """Execute synchronous HTTPX handler operations through a native transport.""" + + def __init__(self, transport: httpx.BaseTransport) -> None: + self._transport = transport + + def send_token_request(self, resource_request: httpx.Request, request: HttpRequest) -> httpx.Response: + return self._transport.handle_request(self._token_request(resource_request, request)) + + def send_resource_request(self, request: httpx.Request, headers: Mapping[str, str]) -> httpx.Response: + return self._transport.handle_request(self._authenticate_request(request, headers)) + + def response_body(self, response: httpx.Response) -> bytes: + return response.read() + + def close_response(self, response: httpx.Response) -> None: + response.read() + response.close() + + +class _AsyncHttpxAdapter(_HttpxAdapterBase): + """Execute asynchronous HTTPX handler operations through a native transport.""" + + def __init__(self, transport: httpx.AsyncBaseTransport) -> None: + self._transport = transport + + async def send_token_request(self, resource_request: httpx.Request, request: HttpRequest) -> httpx.Response: + return await self._transport.handle_async_request(self._token_request(resource_request, request)) + + async def send_resource_request(self, request: httpx.Request, headers: Mapping[str, str]) -> httpx.Response: + return await self._transport.handle_async_request(self._authenticate_request(request, headers)) + + async def response_body(self, response: httpx.Response) -> bytes: + return await response.aread() + + async def close_response(self, response: httpx.Response) -> None: + await response.aread() + await response.aclose() + + +class OAuth2Transport(httpx.BaseTransport): + """Authenticate requests while delegating network I/O to an HTTPX transport.""" + + def __init__(self, config: OAuth2Config, transport: httpx.BaseTransport | None = None) -> None: + self._transport = transport or httpx.HTTPTransport() + self._adapter = _SyncHttpxAdapter(self._transport) + self._handler = OAuth2Handler[httpx.Request, httpx.Response](config) + + def handle_request(self, request: httpx.Request) -> httpx.Response: + """Buffer a replayable request and execute the OAuth2 flow.""" + + request.read() + return self._handler.execute(request, self._adapter) + + def close(self) -> None: + """Close the wrapped HTTPX transport.""" + + self._transport.close() + + +class AsyncOAuth2Transport(httpx.AsyncBaseTransport): + """Authenticate async requests while delegating I/O to an HTTPX transport.""" + + def __init__(self, config: OAuth2Config, transport: httpx.AsyncBaseTransport | None = None) -> None: + self._transport = transport or httpx.AsyncHTTPTransport() + self._adapter = _AsyncHttpxAdapter(self._transport) + self._handler = AsyncOAuth2Handler[httpx.Request, httpx.Response](config) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + """Buffer a replayable request and execute the asynchronous OAuth2 flow.""" + + await request.aread() + return await self._handler.execute(request, self._adapter) + + async def aclose(self) -> None: + """Close the wrapped asynchronous HTTPX transport.""" + + await self._transport.aclose() diff --git a/src/mastercard_oauth2_client/integrations/openapi/__init__.py b/src/mastercard_oauth2_client/integrations/openapi/__init__.py new file mode 100644 index 0000000..f555feb --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/openapi/__init__.py @@ -0,0 +1,5 @@ +"""Integrations for clients generated by OpenAPI Generator.""" + +from .urllib3 import add_oauth2_layer + +__all__ = ["add_oauth2_layer"] diff --git a/src/mastercard_oauth2_client/integrations/openapi/asyncio.py b/src/mastercard_oauth2_client/integrations/openapi/asyncio.py new file mode 100644 index 0000000..c4ac8d9 --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/openapi/asyncio.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any, Protocol, cast + +import aiohttp + +from ...config import OAuth2Config +from ...exceptions import OAuth2Error +from ..aiohttp import OAuth2Middleware + + +class _AsyncioConfiguration(Protocol): + client_session_kwargs: dict[str, Any] | None + + +class _AsyncioRestClient(Protocol): + @property + def configuration(self) -> _AsyncioConfiguration: ... + + @property + def pool_manager(self) -> aiohttp.ClientSession | None: ... + + +class _ApiClient(Protocol): + @property + def rest_client(self) -> _AsyncioRestClient: ... + + +def add_oauth2_layer(api_client: _ApiClient, config: OAuth2Config) -> None: + """Authenticate a generated OpenAPI asyncio client through aiohttp middleware.""" + + rest_client = api_client.rest_client + if rest_client.pool_manager is not None: + raise OAuth2Error("Attach OAuth2 before the generated asyncio client sends its first request") + + session_kwargs = dict(rest_client.configuration.client_session_kwargs or {}) + existing_middlewares = cast(Sequence[aiohttp.ClientMiddlewareType], session_kwargs.get("middlewares") or ()) + if any(isinstance(middleware, OAuth2Middleware) for middleware in existing_middlewares): + raise OAuth2Error("OAuth2 is already attached to the generated asyncio client") + session_kwargs["middlewares"] = (OAuth2Middleware(config), *existing_middlewares) + rest_client.configuration.client_session_kwargs = session_kwargs diff --git a/src/mastercard_oauth2_client/integrations/openapi/httpx.py b/src/mastercard_oauth2_client/integrations/openapi/httpx.py new file mode 100644 index 0000000..498004e --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/openapi/httpx.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import ssl +from collections.abc import Mapping +from typing import Protocol + +import httpx + +from ...config import OAuth2Config +from ...exceptions import OAuth2Error +from ..httpx import AsyncOAuth2Transport + + +class _HttpxRestClient(Protocol): + pool_manager: httpx.AsyncClient | None + + @property + def maxsize(self) -> int: ... + + @property + def ssl_context(self) -> ssl.SSLContext: ... + + @property + def proxy(self) -> str | None: ... + + @property + def proxy_headers(self) -> Mapping[str, str] | None: ... + + +class _ApiClient(Protocol): + @property + def rest_client(self) -> _HttpxRestClient: ... + + +def add_oauth2_layer( + api_client: _ApiClient, + config: OAuth2Config, + transport: httpx.AsyncBaseTransport | None = None, +) -> None: + """Authenticate a generated OpenAPI HTTPX client using native async transport.""" + + rest_client = api_client.rest_client + if rest_client.pool_manager is not None: + raise OAuth2Error("Attach OAuth2 before the generated HTTPX client sends its first request") + + base_transport = transport or _generated_transport(rest_client) + rest_client.pool_manager = httpx.AsyncClient( + transport=AsyncOAuth2Transport(config, base_transport), + ) + + +def _generated_transport(rest_client: _HttpxRestClient) -> httpx.AsyncHTTPTransport: + """Build the transport represented by OpenAPI Generator's HTTPX settings.""" + + proxy = None + if rest_client.proxy: + proxy = httpx.Proxy(url=rest_client.proxy, headers=rest_client.proxy_headers) + return httpx.AsyncHTTPTransport( + verify=rest_client.ssl_context, + proxy=proxy, + limits=httpx.Limits(max_connections=rest_client.maxsize), + trust_env=True, + ) diff --git a/src/mastercard_oauth2_client/integrations/openapi/urllib3.py b/src/mastercard_oauth2_client/integrations/openapi/urllib3.py new file mode 100644 index 0000000..d0a1738 --- /dev/null +++ b/src/mastercard_oauth2_client/integrations/openapi/urllib3.py @@ -0,0 +1,174 @@ +from __future__ import annotations + +import json +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from functools import wraps +from typing import Any, Protocol +from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit + +from ...config import OAuth2Config +from ...exceptions import OAuth2Error +from ...handlers import OAuth2Handler +from ...models import HttpRequest + +_OAUTH2_ATTACHED = "_mastercard_oauth2_attached" + + +class _RawResponse(Protocol): + def drain_conn(self) -> None: ... + + def release_conn(self) -> None: ... + + +class _GeneratedResponse(Protocol): + status: int + + @property + def response(self) -> _RawResponse: ... + + def read(self) -> bytes | str | None: ... + + def getheaders(self) -> Mapping[str, str]: ... + + +type GeneratedRequest = Callable[..., _GeneratedResponse] + + +class _RestClient(Protocol): + def request(self, method: str, url: str, **kwargs: Any) -> _GeneratedResponse: ... + + +class _ApiClient(Protocol): + @property + def rest_client(self) -> _RestClient: ... + + +@dataclass(frozen=True, slots=True) +class _OpenApiRequest: + method: str + url: str + headers: Mapping[str, str] + body: object | None + post_params: object | None + request_timeout: object | None + + +class _OpenApiAdapter: + def __init__(self, request: GeneratedRequest) -> None: + self._request = request + + def request_method(self, request: _OpenApiRequest) -> str: + return request.method + + def request_url(self, request: _OpenApiRequest) -> str: + return request.url + + def request_headers(self, request: _OpenApiRequest) -> Mapping[str, str]: + return request.headers + + def send_token_request(self, resource_request: _OpenApiRequest, request: HttpRequest) -> _GeneratedResponse: + body, post_params = _generated_body(request) + return self._send( + method=request.method, + url=request.url, + headers=request.headers, + body=body, + post_params=post_params, + request_timeout=resource_request.request_timeout, + ) + + def send_resource_request(self, request: _OpenApiRequest, headers: Mapping[str, str]) -> _GeneratedResponse: + return self._send( + method=request.method, + url=request.url, + headers=headers, + body=request.body, + post_params=request.post_params, + request_timeout=request.request_timeout, + ) + + def _send( + self, + *, + method: str, + url: str, + headers: Mapping[str, str], + body: object | None, + post_params: object | None, + request_timeout: object | None, + ) -> _GeneratedResponse: + return self._request( + method, + url, + headers=dict(headers), + body=body, + post_params=post_params, + _request_timeout=request_timeout, + ) + + def response_status(self, response: _GeneratedResponse) -> int: + return response.status + + def response_headers(self, response: _GeneratedResponse) -> Mapping[str, str]: + return _response_headers(response) + + def response_body(self, response: _GeneratedResponse) -> bytes | str | None: + return response.read() + + def close_response(self, response: _GeneratedResponse) -> None: + response.response.drain_conn() + response.response.release_conn() + + +def add_oauth2_layer(api_client: _ApiClient, config: OAuth2Config) -> None: + """Authenticate calls made by a synchronous generated OpenAPI Python client.""" + + original_request = api_client.rest_client.request + if getattr(original_request, _OAUTH2_ATTACHED, False): + raise OAuth2Error("OAuth2 is already attached to the generated urllib3 client") + adapter = _OpenApiAdapter(original_request) + handler = OAuth2Handler[_OpenApiRequest, _GeneratedResponse](config) + + @wraps(original_request) + def authenticated_request(method: str, url: str, **kwargs: Any) -> _GeneratedResponse: + return handler.execute( + _OpenApiRequest( + method=method, + url=_url_with_query(url, kwargs.get("query_params")), + headers=dict(kwargs.get("headers") or {}), + body=kwargs.get("body"), + post_params=kwargs.get("post_params"), + request_timeout=kwargs.get("_request_timeout"), + ), + adapter, + ) + + setattr(authenticated_request, _OAUTH2_ATTACHED, True) + api_client.rest_client.request = authenticated_request # type: ignore[method-assign] + + +def _url_with_query(url: str, query_params: object) -> str: + if not query_params: + return url + parsed = urlsplit(url) + encoded_query = urlencode(query_params, doseq=True) # type: ignore[arg-type] + query = "&".join(part for part in (parsed.query, encoded_query) if part) + return urlunsplit((parsed.scheme, parsed.netloc, parsed.path, query, parsed.fragment)) + + +def _response_headers(response: _GeneratedResponse) -> dict[str, str]: + return {str(name).casefold(): str(value) for name, value in response.getheaders().items()} + + +def _generated_body(request: HttpRequest) -> tuple[object | None, list[tuple[str, str]] | None]: + content_type = _content_type(request.headers) + if content_type == "application/x-www-form-urlencoded" and isinstance(request.body, str): + return None, parse_qsl(request.body, keep_blank_values=True) + if "json" in content_type and isinstance(request.body, str): + return json.loads(request.body), None + return request.body, None + + +def _content_type(headers: Mapping[str, str]) -> str: + return next((value.casefold() for name, value in headers.items() if name.casefold() == "content-type"), "") diff --git a/src/mastercard_oauth2_client/keys.py b/src/mastercard_oauth2_client/keys.py new file mode 100644 index 0000000..ec49590 --- /dev/null +++ b/src/mastercard_oauth2_client/keys.py @@ -0,0 +1,122 @@ +import json +from collections.abc import Mapping +from os import PathLike +from pathlib import Path + +from cryptography.exceptions import UnsupportedAlgorithm +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from cryptography.hazmat.primitives.serialization import pkcs12 + +from ._internal.jose import import_jwk_key_pair +from .exceptions import OAuth2ConfigError +from .models import DPoPKey, KeyPair, PrivateKey + +type KeySource = str | PathLike[str] | bytes +type JwkSource = KeySource | Mapping[str, str | list[str]] + +_MINIMUM_RSA_KEY_SIZE = 2048 +_MINIMUM_EC_KEY_SIZE = 224 + + +def _read_source(source: KeySource) -> bytes: + if isinstance(source, bytes): + return source + return Path(source).read_bytes() + + +def _password_bytes(password: str | None) -> bytes | None: + return password.encode("utf-8") if password is not None else None + + +def validate_key(key: object) -> None: + """Validate an RSA or EC key against the supported FAPI minimums.""" + + if isinstance(key, rsa.RSAPrivateKey | rsa.RSAPublicKey): + if key.key_size < _MINIMUM_RSA_KEY_SIZE: + raise OAuth2ConfigError( + f"RSA keys must have a minimum length of {_MINIMUM_RSA_KEY_SIZE} bits, " + f"but key length was: {key.key_size}" + ) + return + + if isinstance(key, ec.EllipticCurvePrivateKey | ec.EllipticCurvePublicKey): + if key.key_size < _MINIMUM_EC_KEY_SIZE: + raise OAuth2ConfigError( + f"Elliptic curve keys must have a minimum length of {_MINIMUM_EC_KEY_SIZE} bits, " + f"but key length was: {key.key_size}" + ) + return + + algorithm = type(key).__name__ + raise OAuth2ConfigError(f"Key algorithm must be RSA or EC, but was: {algorithm}") + + +def validate_private_key(key: object, *, context: str = "Key") -> PrivateKey: + """Validate that a signing key is a supported private key.""" + + validate_key(key) + if isinstance(key, rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey): + return key + raise OAuth2ConfigError(f"{context} must be an RSA or EC private key") + + +def validate_dpop_key(dpop_key: object) -> None: + """Validate the current DPoP key returned by a provider.""" + + if not isinstance(dpop_key, DPoPKey) or dpop_key.key_pair is None: + raise OAuth2ConfigError("DPoP key provider must return a valid DPoP key") + if not isinstance(dpop_key.key_id, str) or not dpop_key.key_id.strip(): + raise OAuth2ConfigError("DPoP key provider must return a valid DPoP key ID") + validate_key(dpop_key.key_pair.private_key) + validate_key(dpop_key.key_pair.public_key) + + +def load_private_key(source: KeySource, password: str | None = None) -> PrivateKey: + """Load and validate an RSA or EC private key from PEM or DER data.""" + + key_data = _read_source(source) + password_bytes = _password_bytes(password) + + try: + if key_data.lstrip().startswith(b"-----BEGIN"): + loaded_key = serialization.load_pem_private_key(key_data, password=password_bytes) + else: + loaded_key = serialization.load_der_private_key(key_data, password=password_bytes) + except (TypeError, ValueError, UnsupportedAlgorithm) as error: + raise OAuth2ConfigError("Unable to load private key") from error + + return validate_private_key(loaded_key) + + +def load_pkcs12_private_key(source: KeySource, password: str | None = None) -> PrivateKey: + """Load and validate the private key contained in PKCS#12 data.""" + + key_data = _read_source(source) + password_bytes = _password_bytes(password) + + try: + private_key, _, _ = pkcs12.load_key_and_certificates(key_data, password_bytes) + except (TypeError, ValueError, UnsupportedAlgorithm) as error: + raise OAuth2ConfigError("Unable to load PKCS#12 private key") from error + + if private_key is None: + raise OAuth2ConfigError("PKCS#12 data does not contain a private key") + + return validate_private_key(private_key) + + +def load_jwk_key_pair(source: JwkSource) -> KeyPair: + """Parse an RFC 7517 RSA or EC private JWK and construct its key pair.""" + + try: + value = dict(source) if isinstance(source, Mapping) else json.loads(_read_source(source)) + if not isinstance(value, dict): + raise TypeError("JWK must be a JSON object") + return import_jwk_key_pair(value) + except OAuth2ConfigError: + raise + except (json.JSONDecodeError, UnicodeDecodeError) as error: + raise OAuth2ConfigError("Unable to parse JWK JSON") from error + except (TypeError, ValueError) as error: + raise OAuth2ConfigError("Unable to load JWK key pair") from error diff --git a/src/mastercard_oauth2_client/models.py b/src/mastercard_oauth2_client/models.py new file mode 100644 index 0000000..c6fd2b2 --- /dev/null +++ b/src/mastercard_oauth2_client/models.py @@ -0,0 +1,62 @@ +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime + +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +type PrivateKey = rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey +type PublicKey = rsa.RSAPublicKey | ec.EllipticCurvePublicKey + + +@dataclass(frozen=True, slots=True) +class KeyPair: + """An asymmetric private/public key pair used for DPoP signing.""" + + private_key: PrivateKey + public_key: PublicKey + + +@dataclass(frozen=True, slots=True) +class DPoPKey: + """A DPoP signing key pair and its stable key identifier.""" + + key_id: str + key_pair: KeyPair + + +@dataclass(frozen=True, slots=True) +class AccessToken: + """A cached access token and the metadata required for safe reuse.""" + + client_id: str + token_value: str + scopes: frozenset[str] + expires_at: datetime + jkt: str | None = None + + +@dataclass(frozen=True, slots=True) +class AccessTokenFilter: + """Criteria used together to find a reusable access token.""" + + scopes: frozenset[str] + jkt: str | None = None + + +@dataclass(frozen=True, slots=True) +class HttpRequest: + """HTTP-client-neutral token-endpoint request.""" + + method: str + url: str + headers: Mapping[str, str] + body: bytes | str | None = None + + +@dataclass(frozen=True, slots=True) +class HttpResponse: + """HTTP-client-neutral token-endpoint response snapshot.""" + + status: int + headers: Mapping[str, str] + body: bytes | str | None = None diff --git a/src/mastercard_oauth2_client/protocols.py b/src/mastercard_oauth2_client/protocols.py new file mode 100644 index 0000000..65f9de8 --- /dev/null +++ b/src/mastercard_oauth2_client/protocols.py @@ -0,0 +1,116 @@ +from collections.abc import Mapping +from collections.abc import Set as AbstractSet +from typing import Protocol + +from ._internal._scopes import normalize_scopes +from ._internal.jose import jwk_thumbprint +from .models import AccessToken, AccessTokenFilter, DPoPKey, HttpRequest, KeyPair + + +class ScopeResolver(Protocol): + """Determines the OAuth2 scopes required for a resource request.""" + + def resolve(self, method: str, url: str) -> frozenset[str]: ... + + def all_scopes(self) -> frozenset[str]: ... + + +class TokenStore(Protocol): + """Stores reusable access tokens for one OAuth2 client configuration.""" + + def put(self, access_token: AccessToken) -> None: ... + + def get(self, token_filter: AccessTokenFilter) -> AccessToken | None: ... + + +class DPoPKeyProvider(Protocol): + """Supplies the current DPoP key and resolves keys by identifier.""" + + def get_current_key(self) -> DPoPKey: ... + + def get_key(self, key_id: str) -> DPoPKey: ... + + +class SyncHttpAdapter[RequestT, ResponseT](Protocol): + """Bridge OAuth2 orchestration to native resource request and response types. + + Response inspection must preserve any final response returned to the caller. + The handler closes only token responses it consumes and challenged resource + responses it replaces with a retry. + """ + + def request_method(self, request: RequestT) -> str: ... + + def request_url(self, request: RequestT) -> str: ... + + def request_headers(self, request: RequestT) -> Mapping[str, str]: ... + + def send_token_request(self, resource_request: RequestT, request: HttpRequest) -> ResponseT: ... + + def send_resource_request(self, request: RequestT, headers: Mapping[str, str]) -> ResponseT: ... + + def response_status(self, response: ResponseT) -> int: ... + + def response_headers(self, response: ResponseT) -> Mapping[str, str]: ... + + def response_body(self, response: ResponseT) -> bytes | str | None: ... + + def close_response(self, response: ResponseT) -> None: ... + + +class AsyncHttpAdapter[RequestT, ResponseT](Protocol): + """Bridge OAuth2 orchestration to asynchronous native HTTP client types. + + Response inspection must preserve any final response returned to the caller. + The handler closes only token responses it consumes and challenged resource + responses it replaces with a retry. + """ + + def request_method(self, request: RequestT) -> str: ... + + def request_url(self, request: RequestT) -> str: ... + + def request_headers(self, request: RequestT) -> Mapping[str, str]: ... + + async def send_token_request(self, resource_request: RequestT, request: HttpRequest) -> ResponseT: ... + + async def send_resource_request(self, request: RequestT, headers: Mapping[str, str]) -> ResponseT: ... + + def response_status(self, response: ResponseT) -> int: ... + + def response_headers(self, response: ResponseT) -> Mapping[str, str]: ... + + async def response_body(self, response: ResponseT) -> bytes | str | None: ... + + async def close_response(self, response: ResponseT) -> None: ... + + +class StaticScopeResolver: + """Returns one fixed normalized scope set for every resource request.""" + + def __init__(self, scopes: AbstractSet[str]) -> None: + self._scopes = normalize_scopes(scopes) + + def resolve(self, method: str, url: str) -> frozenset[str]: + # Method and URL are intentionally ignored because this resolver is static. + return self._scopes + + def all_scopes(self) -> frozenset[str]: + return self._scopes + + +class StaticDPoPKeyProvider: + """Returns one fixed DPoP key when key rotation is not required.""" + + def __init__(self, key_pair: KeyPair) -> None: + self._dpop_key = DPoPKey( + key_id=jwk_thumbprint(key_pair.public_key), + key_pair=key_pair, + ) + + def get_current_key(self) -> DPoPKey: + return self._dpop_key + + def get_key(self, key_id: str) -> DPoPKey: + # The key ID is intentionally ignored because this provider has one key. + return self._dpop_key diff --git a/src/mastercard_oauth2_client/py.typed b/src/mastercard_oauth2_client/py.typed new file mode 100644 index 0000000..e69de29 diff --git a/src/mastercard_oauth2_client/store.py b/src/mastercard_oauth2_client/store.py new file mode 100644 index 0000000..2b5f376 --- /dev/null +++ b/src/mastercard_oauth2_client/store.py @@ -0,0 +1,63 @@ +from datetime import timedelta +from threading import RLock + +from ._internal._clock import NowProvider +from ._internal._clock import utc_now as _utc_now +from .models import AccessToken, AccessTokenFilter + +type TokenCacheKey = tuple[str | None, tuple[str, ...]] + + +class InMemoryTokenStore: + """Thread-safe token cache for one OAuth2 client configuration. + + A DPoP-bound token can be retrieved by scopes alone or by scopes and its JWK + thumbprint. Use a separate instance for each configured OAuth2 client. + """ + + EARLY_EXPIRATION_THRESHOLD = timedelta(seconds=60) + + def __init__(self, now_provider: NowProvider = _utc_now) -> None: + self._now = now_provider + self._tokens: dict[TokenCacheKey, AccessToken] = {} + # Keep each lookup or update atomic. + self._lock = RLock() + + def put(self, access_token: AccessToken) -> None: + with self._lock: + self._remove_expired_tokens() + + scope_key = self._create_key(jkt=None, scopes=access_token.scopes) + self._tokens[scope_key] = access_token + + if access_token.jkt is not None: + dpop_bound_key = self._create_key(jkt=access_token.jkt, scopes=access_token.scopes) + self._tokens[dpop_bound_key] = access_token + + def get(self, token_filter: AccessTokenFilter) -> AccessToken | None: + key = self._create_key(token_filter.jkt, token_filter.scopes) + with self._lock: + access_token = self._tokens.get(key) + if access_token is None: + return None + + earliest_usable_expiry = self._now() + self.EARLY_EXPIRATION_THRESHOLD + if access_token.expires_at < earliest_usable_expiry: + self._tokens.pop(key, None) + return None + + return access_token + + @staticmethod + def _create_key(jkt: str | None, scopes: frozenset[str]) -> TokenCacheKey: + """Return a stable key independent of the caller's scope ordering.""" + + return jkt, tuple(sorted(scopes)) + + def _remove_expired_tokens(self) -> None: + """Discard expired entries during writes to prevent unbounded growth.""" + + now = self._now() + expired_keys = [key for key, token in self._tokens.items() if token.expires_at < now] + for key in expired_keys: + self._tokens.pop(key, None) diff --git a/src/mastercard_oauth2_client/token.py b/src/mastercard_oauth2_client/token.py new file mode 100644 index 0000000..21a8c01 --- /dev/null +++ b/src/mastercard_oauth2_client/token.py @@ -0,0 +1,160 @@ +import json +import math +from collections.abc import Mapping +from collections.abc import Set as AbstractSet +from datetime import timedelta +from urllib.parse import urlencode + +from ._internal._clock import NowProvider as _NowProvider +from ._internal._clock import utc_now as _utc_now +from ._internal._scopes import normalize_scopes as _normalize_scopes +from ._internal.http import is_success_status as _is_success_status +from ._internal.jose import jwk_thumbprint +from .config import OAuth2Config +from .dpop import create_client_assertion, create_token_dpop_proof +from .exceptions import OAuth2Error +from .keys import validate_dpop_key +from .models import AccessToken, HttpRequest, HttpResponse + +_CLIENT_ASSERTION_TYPE = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" + + +def _build_token_request_body(config: OAuth2Config, scopes: AbstractSet[str]) -> str: + """Build the RFC 6749 client-credentials form body.""" + + form_fields = [ + ("grant_type", "client_credentials"), + ("client_id", config.client_id), + ] + normalized_scopes = _normalize_scopes(scopes) + if normalized_scopes: + form_fields.append(("scope", " ".join(sorted(normalized_scopes)))) + form_fields.extend( + [ + ("client_assertion_type", _CLIENT_ASSERTION_TYPE), + ("client_assertion", create_client_assertion(config)), + ] + ) + return urlencode(form_fields) + + +def _build_token_request_headers( + config: OAuth2Config, + dpop_key_id: str, + nonce: str | None, +) -> dict[str, str]: + """Build token-endpoint headers including a fresh DPoP proof.""" + + return { + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + "DPoP": create_token_dpop_proof(config, dpop_key_id, nonce), + "User-Agent": config.user_agent, + } + + +def build_token_request( + config: OAuth2Config, + scopes: AbstractSet[str], + dpop_key_id: str, + nonce: str | None = None, +) -> HttpRequest: + """Build a neutral RFC 6749 client-credentials request using DPoP.""" + + return HttpRequest( + method="POST", + url=config.token_endpoint, + headers=_build_token_request_headers(config, dpop_key_id, nonce), + body=_build_token_request_body(config, scopes), + ) + + +def _read_token_response_json(body: bytes | str | None) -> Mapping[str, object]: + if body is None or not body.strip(): + raise OAuth2Error("Empty access token response") + try: + value = json.loads(body) + except (json.JSONDecodeError, UnicodeDecodeError) as error: + raise OAuth2Error("Failed to parse JSON access token response") from error + if not isinstance(value, dict): + raise OAuth2Error("Failed to parse JSON access token response") + return value + + +def _read_required_string(value: Mapping[str, object], field: str) -> str: + field_value = value.get(field) + if not isinstance(field_value, str) or not field_value: + raise OAuth2Error(f"Missing value in access token response: {field}") + return field_value + + +def _read_expires_in(value: Mapping[str, object]) -> float: + if "expires_in" not in value: + raise OAuth2Error("Missing value in access token response: expires_in") + expires_in = value["expires_in"] + if isinstance(expires_in, bool) or not isinstance(expires_in, int | float): + raise OAuth2Error( + "FAPI 2.0 requires valid expires_in field in token response for proper token lifetime management" + ) + if not math.isfinite(expires_in) or expires_in <= 0: + raise OAuth2Error( + "FAPI 2.0 requires valid expires_in field in token response for proper token lifetime management" + ) + return float(expires_in) + + +def _read_granted_scopes(value: Mapping[str, object], requested_scopes: AbstractSet[str]) -> frozenset[str]: + """Return granted scopes, treating an omitted RFC 6749 scope as unchanged.""" + + scope_value = value.get("scope") + if scope_value is None: + return _normalize_scopes(requested_scopes) + if not isinstance(scope_value, str): + raise OAuth2Error("Token response scope must be a string") + return frozenset(scope_value.split()) + + +def parse_token_response( + config: OAuth2Config, + response: HttpResponse, + requested_scopes: AbstractSet[str], + dpop_key_id: str, +) -> AccessToken: + """Parse a successful RFC 6749 response into a cache-ready DPoP token.""" + + return _parse_token_response( + config, + response, + requested_scopes, + dpop_key_id, + now_provider=_utc_now, + ) + + +def _parse_token_response( + config: OAuth2Config, + response: HttpResponse, + requested_scopes: AbstractSet[str], + dpop_key_id: str, + *, + now_provider: _NowProvider, +) -> AccessToken: + if not _is_success_status(response.status): + raise OAuth2Error(f"Token request failed with HTTP {response.status}") + + value = _read_token_response_json(response.body) + access_token = _read_required_string(value, "access_token") + token_type = _read_required_string(value, "token_type") + if token_type.casefold() != "dpop": + raise OAuth2Error(f"Expected DPoP token type but received: {token_type}") + expires_in = _read_expires_in(value) + + dpop_key = config.dpop_key_provider.get_key(dpop_key_id) + validate_dpop_key(dpop_key) + return AccessToken( + client_id=config.client_id, + token_value=access_token, + scopes=_read_granted_scopes(value, requested_scopes), + expires_at=now_provider() + timedelta(seconds=expires_in), + jkt=jwk_thumbprint(dpop_key.key_pair.public_key), + ) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py new file mode 100644 index 0000000..ac63e65 --- /dev/null +++ b/tests/integration/conftest.py @@ -0,0 +1,27 @@ +from collections.abc import Iterator + +import pytest + +from tests.integration.support.environment import OAuthTestEnvironment + + +@pytest.fixture(scope="session") +def oauth_servers(tmp_path_factory: pytest.TempPathFactory) -> Iterator[OAuthTestEnvironment]: + """Start one TLS server pair for the session to keep integration tests fast.""" + + environment = OAuthTestEnvironment.start(tmp_path_factory.mktemp("oauth-https")) + try: + yield environment + finally: + environment.close() + + +@pytest.fixture +def oauth_environment(oauth_servers: OAuthTestEnvironment) -> Iterator[OAuthTestEnvironment]: + """Give each test clean request logs, resources, bindings and error modes.""" + + oauth_servers.reset() + try: + yield oauth_servers + finally: + oauth_servers.reset() diff --git a/tests/integration/support/__init__.py b/tests/integration/support/__init__.py new file mode 100644 index 0000000..0480f98 --- /dev/null +++ b/tests/integration/support/__init__.py @@ -0,0 +1 @@ +"""Reusable support for local HTTPS integration tests.""" diff --git a/tests/integration/support/assertions.py b/tests/integration/support/assertions.py new file mode 100644 index 0000000..5ed7bbd --- /dev/null +++ b/tests/integration/support/assertions.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +import base64 +import hashlib +import json +from dataclasses import dataclass +from datetime import UTC, datetime +from typing import Any + +from cryptography.exceptions import InvalidSignature +from cryptography.hazmat.primitives import hashes +from cryptography.hazmat.primitives.asymmetric import ec, padding, rsa, utils + +from mastercard_oauth2_client.models import PublicKey + + +@dataclass(frozen=True, slots=True) +class VerifiedJwt: + header: dict[str, object] + claims: dict[str, object] + + +class JwtValidationError(ValueError): + pass + + +def verify_jwt(compact_jwt: str, public_key: PublicKey | None = None) -> VerifiedJwt: + """Verify compact JOSE independently of the library's signing implementation.""" + + try: + encoded_header, encoded_claims, encoded_signature = compact_jwt.split(".") + header = _json_object(_decode(encoded_header)) + claims = _json_object(_decode(encoded_claims)) + verification_key = public_key or _public_key_from_jwk(header.get("jwk")) + signature = _decode(encoded_signature) + signing_input = f"{encoded_header}.{encoded_claims}".encode("ascii") + _verify_signature(verification_key, header.get("alg"), signing_input, signature) + return VerifiedJwt(header=header, claims=claims) + except JwtValidationError: + raise + except (InvalidSignature, TypeError, UnicodeError, ValueError) as error: + raise JwtValidationError("JWT signature or encoding is invalid") from error + + +def validate_time_claims(claims: dict[str, object], *, now: int | None = None) -> None: + """Reject stale, future or unexpectedly long-lived test assertions and proofs.""" + + current = int(datetime.now(UTC).timestamp()) if now is None else now + issued_at = _integer_claim(claims, "iat") + expires_at = _integer_claim(claims, "exp") + if issued_at > current + 60 or expires_at <= current or expires_at - issued_at > 180: + raise JwtValidationError("JWT time claims are invalid") + if "nbf" in claims and _integer_claim(claims, "nbf") > current + 60: + raise JwtValidationError("JWT nbf claim is invalid") + + +def access_token_hash(access_token: str) -> str: + return _encode(hashlib.sha256(access_token.encode("ascii")).digest()) + + +def jwk_thumbprint(jwk: object) -> str: + if not isinstance(jwk, dict): + raise JwtValidationError("DPoP JWK is missing") + key_type = jwk.get("kty") + members = ("e", "kty", "n") if key_type == "RSA" else ("crv", "kty", "x", "y") + canonical = {name: _string_member(jwk, name) for name in members} + return _encode(hashlib.sha256(json.dumps(canonical, separators=(",", ":"), sort_keys=True).encode("ascii")).digest()) + + +def _verify_signature(public_key: PublicKey, algorithm: object, data: bytes, signature: bytes) -> None: + if algorithm == "PS256" and isinstance(public_key, rsa.RSAPublicKey): + public_key.verify( + signature, + data, + padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=32), + hashes.SHA256(), + ) + return + if algorithm == "ES256" and isinstance(public_key, ec.EllipticCurvePublicKey): + component_size = (public_key.key_size + 7) // 8 + if len(signature) != component_size * 2: + raise JwtValidationError("ES256 signature length is invalid") + r = int.from_bytes(signature[:component_size], "big") + s = int.from_bytes(signature[component_size:], "big") + public_key.verify(utils.encode_dss_signature(r, s), data, ec.ECDSA(hashes.SHA256())) + return + raise JwtValidationError("JWT algorithm does not match its key") + + +def _public_key_from_jwk(value: object) -> PublicKey: + if not isinstance(value, dict): + raise JwtValidationError("DPoP JWK is missing") + if value.get("kty") == "RSA": + return rsa.RSAPublicNumbers( + e=_decode_uint(_string_member(value, "e")), + n=_decode_uint(_string_member(value, "n")), + ).public_key() + if value.get("kty") == "EC" and value.get("crv") == "P-256": + return ec.EllipticCurvePublicNumbers( + x=_decode_uint(_string_member(value, "x")), + y=_decode_uint(_string_member(value, "y")), + curve=ec.SECP256R1(), + ).public_key() + raise JwtValidationError("Unsupported DPoP JWK") + + +def _json_object(value: bytes) -> dict[str, Any]: + decoded = json.loads(value) + if not isinstance(decoded, dict): + raise JwtValidationError("JWT section must be a JSON object") + return decoded + + +def _integer_claim(claims: dict[str, object], name: str) -> int: + value = claims.get(name) + if isinstance(value, bool) or not isinstance(value, int): + raise JwtValidationError(f"JWT {name} claim is invalid") + return value + + +def _string_member(value: dict[object, object], name: str) -> str: + member = value.get(name) + if not isinstance(member, str) or not member: + raise JwtValidationError(f"JWK {name} member is invalid") + return member + + +def _decode_uint(value: str) -> int: + return int.from_bytes(_decode(value), "big") + + +def _decode(value: str) -> bytes: + return base64.b64decode(value + "=" * (-len(value) % 4), altchars=b"-_", validate=True) + + +def _encode(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii") diff --git a/tests/integration/support/authorization_server.py b/tests/integration/support/authorization_server.py new file mode 100644 index 0000000..d6cab3d --- /dev/null +++ b/tests/integration/support/authorization_server.py @@ -0,0 +1,203 @@ +from __future__ import annotations + +import json +import ssl +from collections.abc import Mapping +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from threading import Thread +from typing import Self +from urllib.parse import parse_qs + +from mastercard_oauth2_client._internal.jose import algorithm_for_key + +from .assertions import JwtValidationError, jwk_thumbprint, validate_time_claims, verify_jwt +from .models import AuthorizationState, RecordedRequest +from .tls import TlsFiles + +_ASSERTION_TYPE = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" +TOKEN_PATH = "/oauth/token" + + +class _AuthorizationHttpServer(ThreadingHTTPServer): + daemon_threads = True + state: AuthorizationState + + +class _AuthorizationHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self) -> None: + server = self.server + if not isinstance(server, _AuthorizationHttpServer): + raise TypeError("Unexpected authorization server") + request = _record_request(self) + server.state.requests.append(request) + if server.state.initial_token_request_barrier is not None and len(server.state.requests) <= 2: + server.state.initial_token_request_barrier.wait() + if self.path != TOKEN_PATH: + self._send_json(404, {"error": "not_found"}) + return + try: + self._validate_request(server.state, request) + except JwtValidationError as error: + self._send_json(400, {"error": "invalid_dpop_proof", "error_description": str(error)}) + return + except ValueError as error: + self._send_json(400, {"error": "invalid_client", "error_description": str(error)}) + return + + dpop = verify_jwt(_required_header(request, "dpop")) + provided_nonce = dpop.claims.get("nonce") + if server.state.always_challenge_nonce or provided_nonce != server.state.nonce: + self._send_json( + 400, + {"error": "use_dpop_nonce"}, + {"DPoP-Nonce": server.state.nonce}, + ) + return + + server.state.binding.jkt = jwk_thumbprint(dpop.header.get("jwk")) + self._send_json( + 200, + { + "access_token": server.state.access_token, + "token_type": "DPoP", + "expires_in": 900, + "scope": " ".join(sorted(server.state.scopes)), + }, + {"DPoP-Nonce": server.state.nonce}, + ) + + def log_message(self, format: str, *args: object) -> None: + return + + def _validate_request(self, state: AuthorizationState, request: RecordedRequest) -> None: + _require_single_headers(request, ("accept", "content-type", "dpop", "user-agent")) + if request.header("accept") != "application/json": + raise ValueError("Accept header is invalid") + if request.header("content-type") != "application/x-www-form-urlencoded": + raise ValueError("Content-Type header is invalid") + form = parse_qs(request.body.decode("utf-8"), keep_blank_values=True) + _require_form_value(form, "grant_type", "client_credentials") + _require_form_value(form, "client_id", state.client_id) + _require_form_value(form, "client_assertion_type", _ASSERTION_TYPE) + if set(_single_form_value(form, "scope").split()) != state.scopes: + raise ValueError("Requested scopes are invalid") + if state.reject_client_assertion: + raise ValueError("Client assertion was rejected") + + try: + assertion = verify_jwt(_single_form_value(form, "client_assertion"), state.client_public_key) + except JwtValidationError as error: + raise ValueError("Client assertion signature is invalid") from error + if assertion.header.get("typ") != "JWT" or assertion.header.get("kid") != state.key_id: + raise ValueError("Client assertion header is invalid") + if assertion.header.get("alg") != algorithm_for_key(state.client_public_key): + raise ValueError("Client assertion algorithm is invalid") + for claim in ("iss", "sub"): + if assertion.claims.get(claim) != state.client_id: + raise ValueError(f"Client assertion {claim} is invalid") + if assertion.claims.get("aud") != _request_origin(request): + raise ValueError("Client assertion audience is invalid") + _require_string_claim(assertion.claims, "jti") + validate_time_claims(assertion.claims) + + dpop = verify_jwt(_required_header(request, "dpop")) + if dpop.header.get("typ") != "dpop+jwt": + raise JwtValidationError("DPoP typ is invalid") + dpop_jkt = jwk_thumbprint(dpop.header.get("jwk")) + if dpop.header.get("kid") != dpop_jkt: + raise JwtValidationError("DPoP key ID is invalid") + if dpop.claims.get("htm") != "POST" or dpop.claims.get("htu") != f"{_request_origin(request)}{TOKEN_PATH}": + raise JwtValidationError("DPoP request binding is invalid") + if "ath" in dpop.claims: + raise JwtValidationError("Token DPoP must not contain ath") + _require_string_claim(dpop.claims, "jti") + validate_time_claims(dpop.claims) + + def _send_json(self, status: int, value: object, headers: Mapping[str, str] | None = None) -> None: + body = json.dumps(value, separators=(",", ":")).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + for name, header_value in (headers or {}).items(): + self.send_header(name, header_value) + self.end_headers() + self.wfile.write(body) + + +class FakeAuthorizationServer: + """Local HTTPS token endpoint with configurable deterministic test state.""" + + def __init__(self, state: AuthorizationState, tls_files: TlsFiles) -> None: + self.state = state + self._server = _AuthorizationHttpServer(("127.0.0.1", 0), _AuthorizationHandler) + self._server.state = state + context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(tls_files.certificate, tls_files.private_key) + self._server.socket = context.wrap_socket(self._server.socket, server_side=True) + self._thread = Thread(target=self._server.serve_forever, daemon=True) + + @property + def base_url(self) -> str: + return f"https://127.0.0.1:{self._server.server_port}" + + def __enter__(self) -> Self: + self._thread.start() + return self + + def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None: + self._server.shutdown() + self._server.server_close() + self._thread.join() + + +def _record_request(handler: BaseHTTPRequestHandler) -> RecordedRequest: + content_length = int(handler.headers.get("Content-Length", "0")) + headers = { + name.casefold(): tuple(handler.headers.get_all(name, failobj=[])) + for name in dict(handler.headers) + } + return RecordedRequest( + method=handler.command, + path=handler.path, + headers=headers, + body=handler.rfile.read(content_length), + client_port=handler.client_address[1], + ) + + +def _required_header(request: RecordedRequest, name: str) -> str: + value = request.header(name) + if value is None: + raise ValueError(f"Missing or duplicate header: {name}") + return value + + +def _require_single_headers(request: RecordedRequest, names: tuple[str, ...]) -> None: + for name in names: + _required_header(request, name) + + +def _single_form_value(form: dict[str, list[str]], name: str) -> str: + values = form.get(name, []) + if len(values) != 1 or not values[0]: + raise ValueError(f"Form field is missing or duplicated: {name}") + return values[0] + + +def _require_form_value(form: dict[str, list[str]], name: str, expected: str) -> None: + if _single_form_value(form, name) != expected: + raise ValueError(f"Form field is invalid: {name}") + + +def _require_string_claim(claims: dict[str, object], name: str) -> str: + value = claims.get(name) + if not isinstance(value, str) or not value: + raise JwtValidationError(f"JWT {name} claim is invalid") + return value + + +def _request_origin(request: RecordedRequest) -> str: + host = _required_header(request, "host") + return f"https://{host}" diff --git a/tests/integration/support/environment.py b/tests/integration/support/environment.py new file mode 100644 index 0000000..6bc47c6 --- /dev/null +++ b/tests/integration/support/environment.py @@ -0,0 +1,91 @@ +from __future__ import annotations + +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from typing import Self + +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from mastercard_oauth2_client import ( + InMemoryTokenStore, + KeyPair, + OAuth2Config, + StaticDPoPKeyProvider, + StaticScopeResolver, +) +from mastercard_oauth2_client.models import PrivateKey +from mastercard_oauth2_client.protocols import TokenStore + +from .authorization_server import TOKEN_PATH, FakeAuthorizationServer +from .models import AuthorizationState, DPoPAlgorithm, ResourceState, TokenBinding +from .resource_server import FakeResourceServer +from .tls import TlsFiles, create_tls_files + + +@dataclass(slots=True) +class OAuthTestEnvironment: + """Own the shared HTTPS servers and create isolated OAuth2 configurations.""" + + tls_files: TlsFiles + client_key: PrivateKey + authorization: FakeAuthorizationServer + resource: FakeResourceServer + _stack: ExitStack + + @classmethod + def start(cls, directory: Path) -> Self: + """Start authorization and resource servers sharing one token binding.""" + + tls_files = create_tls_files(directory) + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + binding = TokenBinding() + authorization_state = AuthorizationState( + client_id="integration-client", + key_id="integration-key", + client_public_key=client_key.public_key(), + scopes=frozenset({"resources:read", "resources:write"}), + binding=binding, + ) + resource_state = ResourceState(binding=binding) + stack = ExitStack() + authorization = stack.enter_context(FakeAuthorizationServer(authorization_state, tls_files)) + resource = stack.enter_context(FakeResourceServer(resource_state, tls_files)) + return cls( + tls_files=tls_files, + client_key=client_key, + authorization=authorization, + resource=resource, + _stack=stack, + ) + + def reset(self) -> None: + """Clear all scenario state before and after each integration test.""" + + self.authorization.state.reset() + self.resource.state.reset() + + def config(self, algorithm: DPoPAlgorithm, *, token_store: TokenStore | None = None) -> OAuth2Config: + """Create a fresh client config with the requested DPoP key algorithm.""" + + dpop_private_key: PrivateKey + if algorithm == "ec": + dpop_private_key = ec.generate_private_key(ec.SECP256R1()) + else: + dpop_private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + selected_store = token_store if token_store is not None else InMemoryTokenStore() + return OAuth2Config( + client_id=self.authorization.state.client_id, + token_endpoint=f"{self.authorization.base_url}{TOKEN_PATH}", + issuer=self.authorization.base_url, + client_key=self.client_key, + key_id=self.authorization.state.key_id, + scope_resolver=StaticScopeResolver(self.authorization.state.scopes), + dpop_key_provider=StaticDPoPKeyProvider( + KeyPair(private_key=dpop_private_key, public_key=dpop_private_key.public_key()) + ), + token_store=selected_store, + ) + + def close(self) -> None: + self._stack.close() diff --git a/tests/integration/support/lifecycle.py b/tests/integration/support/lifecycle.py new file mode 100644 index 0000000..90bf0c5 --- /dev/null +++ b/tests/integration/support/lifecycle.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +from typing import Protocol + + +class LifecycleClient(Protocol): + """Common fake-resource operations every synchronous test client exposes.""" + + def create_resource(self, name: str) -> str: ... + + def get_resource(self, resource_id: str) -> str: ... + + def delete_resource(self, resource_id: str) -> None: ... + + +class AsyncLifecycleClient(Protocol): + """Common fake-resource operations every asynchronous test client exposes.""" + + async def create_resource(self, name: str) -> str: ... + + async def get_resource(self, resource_id: str) -> str: ... + + async def delete_resource(self, resource_id: str) -> None: ... + + +def run_resource_lifecycle(client: LifecycleClient, *, name: str = "integration-resource") -> str: + """Create and fetch one resource, always deleting it before returning.""" + + resource_id = client.create_resource(name) + try: + fetched_id = client.get_resource(resource_id) + if fetched_id != resource_id: + raise AssertionError("Fetched resource ID does not match created resource ID") + return resource_id + finally: + client.delete_resource(resource_id) + + +async def run_async_resource_lifecycle( + client: AsyncLifecycleClient, + *, + name: str = "integration-resource", +) -> str: + """Create and fetch one resource asynchronously, always deleting it.""" + + resource_id = await client.create_resource(name) + try: + fetched_id = await client.get_resource(resource_id) + if fetched_id != resource_id: + raise AssertionError("Fetched resource ID does not match created resource ID") + return resource_id + finally: + await client.delete_resource(resource_id) + + +def require_resource_id(value: object) -> str: + """Return a valid resource ID from a decoded JSON object.""" + + if not isinstance(value, dict): + raise TypeError("Response body is not a JSON object") + resource_id = value.get("id") + if not isinstance(resource_id, str) or not resource_id: + raise AssertionError("Resource response has no valid ID") + return resource_id diff --git a/tests/integration/support/lifecycle_clients/__init__.py b/tests/integration/support/lifecycle_clients/__init__.py new file mode 100644 index 0000000..4036b3c --- /dev/null +++ b/tests/integration/support/lifecycle_clients/__init__.py @@ -0,0 +1 @@ +"""Client-specific adapters for shared integration lifecycle scenarios.""" diff --git a/tests/integration/support/lifecycle_clients/aiohttp.py b/tests/integration/support/lifecycle_clients/aiohttp.py new file mode 100644 index 0000000..d98a802 --- /dev/null +++ b/tests/integration/support/lifecycle_clients/aiohttp.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import aiohttp + +from ..lifecycle import require_resource_id + + +@dataclass(slots=True) +class AiohttpLifecycleClient: + """Run the shared lifecycle through a native aiohttp client session.""" + + client: aiohttp.ClientSession + base_url: str + + async def create_resource(self, name: str) -> str: + response = await self.client.post(f"{self.base_url}/api/resources", json={"name": name}) + response.raise_for_status() + return require_resource_id(await response.json()) + + async def get_resource(self, resource_id: str) -> str: + response = await self.client.get( + f"{self.base_url}/api/resources/{resource_id}", + headers={"Accept": "application/json"}, + ) + response.raise_for_status() + return require_resource_id(await response.json()) + + async def delete_resource(self, resource_id: str) -> None: + response = await self.client.delete(f"{self.base_url}/api/resources/{resource_id}") + response.raise_for_status() diff --git a/tests/integration/support/lifecycle_clients/httpx.py b/tests/integration/support/lifecycle_clients/httpx.py new file mode 100644 index 0000000..3c75b19 --- /dev/null +++ b/tests/integration/support/lifecycle_clients/httpx.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import httpx + +from ..lifecycle import require_resource_id + + +@dataclass(slots=True) +class HttpxLifecycleClient: + """Run the shared lifecycle through a native synchronous HTTPX client.""" + + client: httpx.Client + base_url: str + + def create_resource(self, name: str) -> str: + response = self.client.post(f"{self.base_url}/api/resources", json={"name": name}) + response.raise_for_status() + return require_resource_id(response.json()) + + def get_resource(self, resource_id: str) -> str: + response = self.client.get( + f"{self.base_url}/api/resources/{resource_id}", + headers={"Accept": "application/json"}, + ) + response.raise_for_status() + return require_resource_id(response.json()) + + def delete_resource(self, resource_id: str) -> None: + response = self.client.delete(f"{self.base_url}/api/resources/{resource_id}") + response.raise_for_status() + + +@dataclass(slots=True) +class AsyncHttpxLifecycleClient: + """Run the shared lifecycle through a native asynchronous HTTPX client.""" + + client: httpx.AsyncClient + base_url: str + + async def create_resource(self, name: str) -> str: + response = await self.client.post(f"{self.base_url}/api/resources", json={"name": name}) + response.raise_for_status() + return require_resource_id(response.json()) + + async def get_resource(self, resource_id: str) -> str: + response = await self.client.get( + f"{self.base_url}/api/resources/{resource_id}", + headers={"Accept": "application/json"}, + ) + response.raise_for_status() + return require_resource_id(response.json()) + + async def delete_resource(self, resource_id: str) -> None: + response = await self.client.delete(f"{self.base_url}/api/resources/{resource_id}") + response.raise_for_status() diff --git a/tests/integration/support/lifecycle_clients/openapi_asyncio.py b/tests/integration/support/lifecycle_clients/openapi_asyncio.py new file mode 100644 index 0000000..3618374 --- /dev/null +++ b/tests/integration/support/lifecycle_clients/openapi_asyncio.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from generated_asyncio_test_client import ( # type: ignore[import-not-found] + ApiClient, + Configuration, + Resource, + ResourcesApi, +) + +from mastercard_oauth2_client import OAuth2Config +from mastercard_oauth2_client.integrations.openapi.asyncio import add_oauth2_layer + + +@dataclass(slots=True) +class OpenApiAsyncioLifecycleClient: + """Run the shared lifecycle through a generated Python/asyncio client.""" + + api: ResourcesApi + api_client: ApiClient + + @classmethod + def create(cls, base_url: str, ca_file: str, config: OAuth2Config) -> OpenApiAsyncioLifecycleClient: + generated_config = Configuration(host=base_url) + generated_config.ssl_ca_cert = ca_file + generated_config.proxy = None + api_client = ApiClient(generated_config) + add_oauth2_layer(api_client, config) + return cls(ResourcesApi(api_client), api_client) + + async def close(self) -> None: + await self.api_client.close() + + async def create_resource(self, name: str) -> str: + resource = await self.api.create_resource(Resource(name=name)) + if not isinstance(resource.id, str) or not resource.id: + raise AssertionError("Created resource has no valid ID") + return resource.id + + async def get_resource(self, resource_id: str) -> str: + resource = await self.api.get_resource(resource_id) + if not isinstance(resource.id, str) or not resource.id: + raise AssertionError("Fetched resource has no valid ID") + return resource.id + + async def delete_resource(self, resource_id: str) -> None: + await self.api.delete_resource(resource_id) diff --git a/tests/integration/support/lifecycle_clients/openapi_httpx.py b/tests/integration/support/lifecycle_clients/openapi_httpx.py new file mode 100644 index 0000000..03b3207 --- /dev/null +++ b/tests/integration/support/lifecycle_clients/openapi_httpx.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from generated_httpx_test_client import ( # type: ignore[import-not-found] + ApiClient, + Configuration, + Resource, + ResourcesApi, +) + +from mastercard_oauth2_client import OAuth2Config +from mastercard_oauth2_client.integrations.openapi.httpx import add_oauth2_layer + + +@dataclass(slots=True) +class OpenApiHttpxLifecycleClient: + """Run the shared lifecycle through a generated Python/HTTPX client.""" + + api: ResourcesApi + api_client: ApiClient + + @classmethod + def create(cls, base_url: str, ca_file: str, config: OAuth2Config) -> OpenApiHttpxLifecycleClient: + generated_config = Configuration(host=base_url) + generated_config.ssl_ca_cert = ca_file + generated_config.proxy = None + api_client = ApiClient(generated_config) + add_oauth2_layer(api_client, config) + return cls(ResourcesApi(api_client), api_client) + + async def close(self) -> None: + await self.api_client.close() + + async def create_resource(self, name: str) -> str: + resource = await self.api.create_resource(Resource(name=name)) + if not isinstance(resource.id, str) or not resource.id: + raise AssertionError("Created resource has no valid ID") + return resource.id + + async def get_resource(self, resource_id: str) -> str: + resource = await self.api.get_resource(resource_id) + if not isinstance(resource.id, str) or not resource.id: + raise AssertionError("Fetched resource has no valid ID") + return resource.id + + async def delete_resource(self, resource_id: str) -> None: + await self.api.delete_resource(resource_id) diff --git a/tests/integration/support/lifecycle_clients/openapi_urllib3.py b/tests/integration/support/lifecycle_clients/openapi_urllib3.py new file mode 100644 index 0000000..a3a6f73 --- /dev/null +++ b/tests/integration/support/lifecycle_clients/openapi_urllib3.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from generated_test_client import ApiClient, Configuration, Resource, ResourcesApi # type: ignore[import-not-found] + +from mastercard_oauth2_client import OAuth2Config, add_oauth2_layer + + +@dataclass(slots=True) +class OpenApiLifecycleClient: + """Run the shared lifecycle through the generated Python/urllib3 client.""" + + api: ResourcesApi + + @classmethod + def create(cls, base_url: str, ca_file: str, config: OAuth2Config) -> OpenApiLifecycleClient: + generated_config = Configuration(host=base_url) + generated_config.ssl_ca_cert = ca_file + generated_config.proxy = None + api_client = ApiClient(generated_config) + add_oauth2_layer(api_client, config) + return cls(ResourcesApi(api_client)) + + def create_resource(self, name: str) -> str: + resource = self.api.create_resource(Resource(name=name)) + if not isinstance(resource.id, str) or not resource.id: + raise AssertionError("Created resource has no valid ID") + return resource.id + + def get_resource(self, resource_id: str) -> str: + resource = self.api.get_resource(resource_id) + if not isinstance(resource.id, str) or not resource.id: + raise AssertionError("Fetched resource has no valid ID") + return resource.id + + def delete_resource(self, resource_id: str) -> None: + self.api.delete_resource(resource_id) diff --git a/tests/integration/support/models.py b/tests/integration/support/models.py new file mode 100644 index 0000000..74016f5 --- /dev/null +++ b/tests/integration/support/models.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from threading import Barrier +from typing import Literal + +from mastercard_oauth2_client.models import PublicKey + + +@dataclass(frozen=True, slots=True) +class RecordedRequest: + method: str + path: str + headers: dict[str, tuple[str, ...]] + body: bytes + client_port: int + + def header(self, name: str) -> str | None: + values = self.headers.get(name.casefold(), ()) + return values[0] if len(values) == 1 else None + + +@dataclass(slots=True) +class TokenBinding: + jkt: str | None = None + + +@dataclass(slots=True) +class AuthorizationState: + client_id: str + key_id: str + client_public_key: PublicKey + scopes: frozenset[str] + binding: TokenBinding + requests: list[RecordedRequest] = field(default_factory=list) + nonce: str = "authorization-nonce" + access_token: str = "integration-access-token" + reject_client_assertion: bool = False + always_challenge_nonce: bool = False + initial_token_request_barrier: Barrier | None = None + + def reset(self) -> None: + self.requests.clear() + self.binding.jkt = None + self.reject_client_assertion = False + self.always_challenge_nonce = False + self.initial_token_request_barrier = None + + +@dataclass(slots=True) +class ResourceState: + binding: TokenBinding + requests: list[RecordedRequest] = field(default_factory=list) + resources: dict[str, dict[str, object]] = field(default_factory=dict) + nonce: str = "resource-nonce" + access_token: str = "integration-access-token" + next_id: int = 1 + insufficient_scope: bool = False + always_challenge_nonce: bool = False + fail_get: bool = False + + def reset(self) -> None: + self.requests.clear() + self.resources.clear() + self.next_id = 1 + self.insufficient_scope = False + self.always_challenge_nonce = False + self.fail_get = False + + +type DPoPAlgorithm = Literal["ec", "rsa"] diff --git a/tests/integration/support/resource_server.py b/tests/integration/support/resource_server.py new file mode 100644 index 0000000..fc46e82 --- /dev/null +++ b/tests/integration/support/resource_server.py @@ -0,0 +1,250 @@ +from __future__ import annotations + +import json +import ssl +from collections.abc import Mapping +from http import HTTPStatus +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from threading import Thread +from typing import Self +from urllib.parse import urlsplit + +from .assertions import JwtValidationError, access_token_hash, jwk_thumbprint, validate_time_claims, verify_jwt +from .models import RecordedRequest, ResourceState +from .tls import TlsFiles + +RESOURCE_COLLECTION_PATH = "/api/resources" +PAYLOAD_ECHO_PATH = "/api/echo" +_RESOURCE_ITEM_PREFIX = f"{RESOURCE_COLLECTION_PATH}/" + + +class _ResourceHttpServer(ThreadingHTTPServer): + daemon_threads = True + state: ResourceState + + +class _ResourceHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_GET(self) -> None: + self._handle() + + def do_POST(self) -> None: + self._handle() + + def do_DELETE(self) -> None: + self._handle() + + def log_message(self, format: str, *args: object) -> None: + return + + def _handle(self) -> None: + server = self.server + if not isinstance(server, _ResourceHttpServer): + raise TypeError("Unexpected resource server") + request = _record_request(self) + server.state.requests.append(request) + try: + self._validate_request(server.state, request) + except JwtValidationError as error: + self._send_json(401, {"error": "invalid_dpop_proof", "error_description": str(error)}) + return + except ValueError as error: + self._send_json(401, {"error": "invalid_token", "error_description": str(error)}) + return + + if server.state.insufficient_scope: + self._send_json( + 403, + {"error": "insufficient_scope", "error_description": "requested scope is not permitted"}, + {"WWW-Authenticate": 'DPoP error="insufficient_scope"'}, + ) + return + + dpop = verify_jwt(_required_header(request, "dpop")) + if server.state.always_challenge_nonce or dpop.claims.get("nonce") != server.state.nonce: + self._send_json( + 401, + {"error": "use_dpop_nonce"}, + { + "DPoP-Nonce": server.state.nonce, + "WWW-Authenticate": 'DPoP error="use_dpop_nonce"', + }, + ) + return + + self._handle_resource(server.state, request) + + def _validate_request(self, state: ResourceState, request: RecordedRequest) -> None: + _require_single_headers(request, ("authorization", "dpop", "user-agent")) + if request.header("authorization") != f"DPoP {state.access_token}": + raise ValueError("Access token is invalid") + if request.method == "GET" and request.header("accept") != "application/json": + raise ValueError("Accept header is invalid") + path = urlsplit(request.path).path + if ( + path == RESOURCE_COLLECTION_PATH + and request.method == "POST" + and request.header("content-type") != "application/json" + ): + raise ValueError("Content-Type header is invalid") + + dpop = verify_jwt(_required_header(request, "dpop")) + if dpop.header.get("typ") != "dpop+jwt": + raise JwtValidationError("DPoP typ is invalid") + dpop_jkt = jwk_thumbprint(dpop.header.get("jwk")) + if state.binding.jkt is None or dpop_jkt != state.binding.jkt or dpop.header.get("kid") != dpop_jkt: + raise JwtValidationError("DPoP token binding is invalid") + expected_htu = f"{_request_origin(request)}{urlsplit(request.path).path}" + if dpop.claims.get("htm") != request.method or dpop.claims.get("htu") != expected_htu: + raise JwtValidationError("DPoP request binding is invalid") + if dpop.claims.get("ath") != access_token_hash(state.access_token): + raise JwtValidationError("DPoP ath is invalid") + _require_string_claim(dpop.claims, "jti") + validate_time_claims(dpop.claims) + + def _handle_resource(self, state: ResourceState, request: RecordedRequest) -> None: + path = urlsplit(request.path).path + if path == PAYLOAD_ECHO_PATH: + self._handle_echo(state, request) + return + if path == RESOURCE_COLLECTION_PATH: + self._handle_collection(state, request) + return + if path.startswith(_RESOURCE_ITEM_PREFIX): + self._handle_item(state, request, path.removeprefix(_RESOURCE_ITEM_PREFIX)) + return + self._send_json(404, {"error": "not_found"}) + + def _handle_echo(self, state: ResourceState, request: RecordedRequest) -> None: + if request.method == "POST": + self._send_json(200, {"received": len(request.body)}, {"DPoP-Nonce": state.nonce}) + return + self._send_json(404, {"error": "not_found"}) + + def _handle_collection(self, state: ResourceState, request: RecordedRequest) -> None: + if request.method == "POST": + try: + value = _parse_new_resource(request.body) + except (json.JSONDecodeError, UnicodeDecodeError, ValueError): + self._send_json(400, {"error": "invalid_resource"}) + return + resource_id = str(state.next_id) + state.next_id += 1 + resource = {**value, "id": resource_id} + state.resources[resource_id] = resource + self._send_json(201, resource, {"DPoP-Nonce": state.nonce}) + return + self._send_json(404, {"error": "not_found"}) + + def _handle_item(self, state: ResourceState, request: RecordedRequest, resource_id: str) -> None: + if request.method == "GET": + if state.fail_get: + self._send_json(500, {"error": "resource_server_error"}) + return + stored_resource = state.resources.get(resource_id) + if stored_resource is None: + self._send_json(404, {"error": "not_found"}) + else: + self._send_json(200, stored_resource, {"DPoP-Nonce": state.nonce}) + return + if request.method == "DELETE": + if state.resources.pop(resource_id, None) is None: + self._send_json(404, {"error": "not_found"}) + else: + self._send_empty(204, {"DPoP-Nonce": state.nonce}) + return + self._send_json(HTTPStatus.METHOD_NOT_ALLOWED, {"error": "method_not_allowed"}) + + def _send_json(self, status: int, value: object, headers: Mapping[str, str] | None = None) -> None: + body = json.dumps(value, separators=(",", ":")).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + for name, header_value in (headers or {}).items(): + self.send_header(name, header_value) + self.end_headers() + self.wfile.write(body) + + def _send_empty(self, status: int, headers: Mapping[str, str] | None = None) -> None: + self.send_response(status) + self.send_header("Content-Length", "0") + for name, header_value in (headers or {}).items(): + self.send_header(name, header_value) + self.end_headers() + + +class FakeResourceServer: + """Local HTTPS resource API that validates DPoP and stores fake resources.""" + + def __init__(self, state: ResourceState, tls_files: TlsFiles) -> None: + self.state = state + self._server = _ResourceHttpServer(("127.0.0.1", 0), _ResourceHandler) + self._server.state = state + context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(tls_files.certificate, tls_files.private_key) + self._server.socket = context.wrap_socket(self._server.socket, server_side=True) + self._thread = Thread(target=self._server.serve_forever, daemon=True) + + @property + def base_url(self) -> str: + return f"https://127.0.0.1:{self._server.server_port}" + + def __enter__(self) -> Self: + self._thread.start() + return self + + def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None: + self._server.shutdown() + self._server.server_close() + self._thread.join() + + +def _record_request(handler: BaseHTTPRequestHandler) -> RecordedRequest: + content_length = int(handler.headers.get("Content-Length", "0")) + headers = { + name.casefold(): tuple(handler.headers.get_all(name, failobj=[])) + for name in dict(handler.headers) + } + return RecordedRequest( + method=handler.command, + path=handler.path, + headers=headers, + body=handler.rfile.read(content_length), + client_port=handler.client_address[1], + ) + + +def _required_header(request: RecordedRequest, name: str) -> str: + value = request.header(name) + if value is None: + raise ValueError(f"Missing or duplicate header: {name}") + return value + + +def _require_single_headers(request: RecordedRequest, names: tuple[str, ...]) -> None: + for name in names: + _required_header(request, name) + + +def _require_string_claim(claims: dict[str, object], name: str) -> str: + value = claims.get(name) + if not isinstance(value, str) or not value: + raise JwtValidationError(f"JWT {name} claim is invalid") + return value + + +def _request_origin(request: RecordedRequest) -> str: + return f"https://{_required_header(request, 'host')}" + + +def _parse_new_resource(body: bytes) -> dict[str, object]: + """Accept only payloads allowed by the fake OpenAPI Resource schema.""" + + value = json.loads(body) + if not isinstance(value, dict) or set(value) != {"name"}: + raise ValueError("Resource payload must contain only name") + name = value.get("name") + if not isinstance(name, str) or not 1 <= len(name) <= 200: + raise ValueError("Resource name is invalid") + return {"name": name} diff --git a/tests/integration/support/stores.py b/tests/integration/support/stores.py new file mode 100644 index 0000000..67fc05b --- /dev/null +++ b/tests/integration/support/stores.py @@ -0,0 +1,21 @@ +from mastercard_oauth2_client import AccessToken, AccessTokenFilter, InMemoryTokenStore + + +class InvalidatableTokenStore: + """Test store that can force a cache miss without changing wall-clock time.""" + + def __init__(self) -> None: + self._delegate = InMemoryTokenStore() + self._expired = False + + def put(self, access_token: AccessToken) -> None: + self._delegate.put(access_token) + self._expired = False + + def get(self, token_filter: AccessTokenFilter) -> AccessToken | None: + if self._expired: + return None + return self._delegate.get(token_filter) + + def invalidate(self) -> None: + self._expired = True diff --git a/tests/integration/support/tls.py b/tests/integration/support/tls.py new file mode 100644 index 0000000..5785374 --- /dev/null +++ b/tests/integration/support/tls.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import ipaddress +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from pathlib import Path + +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.x509.oid import NameOID + + +@dataclass(frozen=True, slots=True) +class TlsFiles: + certificate: Path + private_key: Path + + +def create_tls_files(directory: Path) -> TlsFiles: + """Create a one-day self-signed certificate trusted only by test clients.""" + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + subject = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "localhost")]) + now = datetime.now(UTC) + certificate = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(subject) + .public_key(private_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - timedelta(minutes=1)) + .not_valid_after(now + timedelta(days=1)) + .add_extension( + x509.SubjectAlternativeName([x509.IPAddress(ipaddress.ip_address("127.0.0.1"))]), + critical=False, + ) + .add_extension(x509.BasicConstraints(ca=True, path_length=None), critical=True) + .sign(private_key, hashes.SHA256()) + ) + + certificate_path = directory / "certificate.pem" + private_key_path = directory / "private-key.pem" + certificate_path.write_bytes(certificate.public_bytes(serialization.Encoding.PEM)) + private_key_path.write_bytes( + private_key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + ) + return TlsFiles(certificate=certificate_path, private_key=private_key_path) diff --git a/tests/integration/test_aiohttp.py b/tests/integration/test_aiohttp.py new file mode 100644 index 0000000..1c56cb1 --- /dev/null +++ b/tests/integration/test_aiohttp.py @@ -0,0 +1,412 @@ +from __future__ import annotations + +import asyncio +import email +import io +import ssl +from collections import Counter +from collections.abc import AsyncIterator +from threading import Barrier +from urllib.parse import parse_qs + +import aiohttp +import aiohttp_retry +import pytest + +from mastercard_oauth2_client import OAuth2Config +from mastercard_oauth2_client.integrations.aiohttp import OAuth2Middleware +from tests.integration.support.environment import OAuthTestEnvironment +from tests.integration.support.lifecycle import run_async_resource_lifecycle +from tests.integration.support.lifecycle_clients.aiohttp import AiohttpLifecycleClient +from tests.integration.support.models import DPoPAlgorithm +from tests.integration.support.resource_server import PAYLOAD_ECHO_PATH +from tests.integration.support.stores import InvalidatableTokenStore + +pytestmark = pytest.mark.aiohttp + + +def _client( + environment: OAuthTestEnvironment, + algorithm: DPoPAlgorithm = "ec", + *, + config: OAuth2Config | None = None, +) -> aiohttp.ClientSession: + ssl_context = ssl.create_default_context(cafile=str(environment.tls_files.certificate)) + return aiohttp.ClientSession( + connector=aiohttp.TCPConnector(ssl=ssl_context), + middlewares=(OAuth2Middleware(config or environment.config(algorithm)),), + ) + + +@pytest.mark.parametrize("dpop_algorithm", ["ec", "rsa"]) +def test_aiohttp_client_completes_stateful_https_lifecycle( + oauth_environment: OAuthTestEnvironment, + dpop_algorithm: DPoPAlgorithm, +) -> None: + async def exercise() -> str: + async with _client(oauth_environment, dpop_algorithm) as client: + lifecycle_client = AiohttpLifecycleClient(client, oauth_environment.resource.base_url) + return await run_async_resource_lifecycle(lifecycle_client) + + resource_id = asyncio.run(exercise()) + + assert resource_id == "1" + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.authorization.state.requests] == ["POST", "POST"] + assert [request.method for request in oauth_environment.resource.state.requests] == [ + "POST", + "POST", + "GET", + "DELETE", + ] + assert len({request.client_port for request in oauth_environment.authorization.state.requests}) == 1 + assert len({request.client_port for request in oauth_environment.resource.state.requests}) == 1 + + +def test_aiohttp_lifecycle_attempts_delete_after_get_failure( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + async with _client(oauth_environment) as client: + lifecycle_client = AiohttpLifecycleClient(client, oauth_environment.resource.base_url) + await run_async_resource_lifecycle(lifecycle_client) + + oauth_environment.resource.state.fail_get = True + + with pytest.raises(aiohttp.ClientResponseError) as error: + asyncio.run(exercise()) + + assert error.value.status == 500 + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.resource.state.requests][-2:] == ["GET", "DELETE"] + + +def test_aiohttp_client_buffers_streaming_body_for_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + body = b'{"name":"streamed-resource"}' + + async def body_stream() -> AsyncIterator[bytes]: + yield body[:10] + yield body[10:] + + async def exercise() -> aiohttp.ClientResponse: + async with _client(oauth_environment) as client: + return await client.post( + f"{oauth_environment.resource.base_url}/api/resources", + headers={"Content-Type": "application/json"}, + data=body_stream(), + ) + + response = asyncio.run(exercise()) + + assert response.status == 201 + assert [request.body for request in oauth_environment.resource.state.requests] == [body, body] + + +def test_aiohttp_client_preserves_form_body_for_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> int: + async with _client(oauth_environment) as client: + response = await client.post( + f"{oauth_environment.resource.base_url}{PAYLOAD_ECHO_PATH}", + data={"name": "form resource", "tag": ["first", "second"]}, + ) + await response.read() + return response.status + + assert asyncio.run(exercise()) == 200 + requests = oauth_environment.resource.state.requests + assert len(requests) == 2 + assert [request.header("content-type") for request in requests] == [ + "application/x-www-form-urlencoded", + "application/x-www-form-urlencoded", + ] + assert [parse_qs(request.body.decode()) for request in requests] == [ + {"name": ["form resource"], "tag": ["first", "second"]}, + {"name": ["form resource"], "tag": ["first", "second"]}, + ] + + +def test_aiohttp_client_preserves_multipart_file_for_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + file_body = b"multipart-file-content" + + async def exercise() -> int: + form = aiohttp.FormData() + form.add_field("description", "multipart resource") + form.add_field("document", io.BytesIO(file_body), filename="resource.txt", content_type="text/plain") + async with _client(oauth_environment) as client: + response = await client.post(f"{oauth_environment.resource.base_url}{PAYLOAD_ECHO_PATH}", data=form) + await response.read() + return response.status + + assert asyncio.run(exercise()) == 200 + requests = oauth_environment.resource.state.requests + assert len(requests) == 2 + assert requests[0].header("content-type") == requests[1].header("content-type") + assert [_multipart_parts(request.header("content-type"), request.body) for request in requests] == [ + {"description": (None, b"multipart resource"), "document": ("resource.txt", file_body)}, + {"description": (None, b"multipart resource"), "document": ("resource.txt", file_body)}, + ] + + +def test_aiohttp_client_preserves_query_json_and_caller_headers( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> tuple[int, dict[str, object]]: + async with _client(oauth_environment) as client: + response = await client.post( + f"{oauth_environment.resource.base_url}/api/resources", + params={"source": "aiohttp"}, + json={"name": "aiohttp-resource"}, + headers={ + "Accept": "application/json", + "User-Agent": "caller-agent", + "Authorization": "Bearer stale", + "DPoP": "stale-proof", + }, + ) + return response.status, await response.json() + + status, response_body = asyncio.run(exercise()) + recorded = oauth_environment.resource.state.requests[-1] + + assert status == 201 + assert response_body == {"id": "1", "name": "aiohttp-resource"} + assert recorded.path == "/api/resources?source=aiohttp" + assert recorded.body == b'{"name": "aiohttp-resource"}' + assert recorded.header("user-agent") == "caller-agent" + assert recorded.header("authorization") == "DPoP integration-access-token" + assert recorded.header("dpop") != "stale-proof" + + +def test_aiohttp_client_reuses_cached_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> tuple[int, int]: + async with _client(oauth_environment) as client: + headers = {"Accept": "application/json"} + first = await client.get(f"{oauth_environment.resource.base_url}/api/resources/missing", headers=headers) + second = await client.get(f"{oauth_environment.resource.base_url}/api/resources/missing", headers=headers) + return first.status, second.status + + assert asyncio.run(exercise()) == (404, 404) + assert len(oauth_environment.authorization.state.requests) == 2 + + +def test_aiohttp_client_reacquires_invalidated_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise(config: OAuth2Config, store: InvalidatableTokenStore) -> None: + async with _client(oauth_environment, config=config) as client: + headers = {"Accept": "application/json"} + await client.get(f"{oauth_environment.resource.base_url}/api/resources/missing", headers=headers) + store.invalidate() + await client.get(f"{oauth_environment.resource.base_url}/api/resources/missing", headers=headers) + + store = InvalidatableTokenStore() + asyncio.run(exercise(oauth_environment.config("ec", token_store=store), store)) + + assert len(oauth_environment.authorization.state.requests) == 4 + + +def test_aiohttp_client_allows_concurrent_cache_misses(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> list[int]: + async with _client(oauth_environment) as client: + responses = await asyncio.gather( + client.get( + f"{oauth_environment.resource.base_url}/api/resources/first", + headers={"Accept": "application/json"}, + ), + client.get( + f"{oauth_environment.resource.base_url}/api/resources/second", + headers={"Accept": "application/json"}, + ), + ) + statuses = [response.status for response in responses] + await asyncio.gather(*(response.read() for response in responses)) + return statuses + + oauth_environment.authorization.state.initial_token_request_barrier = Barrier(2) + + assert asyncio.run(exercise()) == [404, 404] + assert len(oauth_environment.authorization.state.requests) == 4 + resource_paths = Counter(request.path for request in oauth_environment.resource.state.requests) + assert set(resource_paths) == {"/api/resources/first", "/api/resources/second"} + assert all(attempts in {1, 2} for attempts in resource_paths.values()) + + +def test_aiohttp_client_stops_after_one_token_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> int: + async with _client(oauth_environment) as client: + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + return response.status + + oauth_environment.authorization.state.always_challenge_nonce = True + + assert asyncio.run(exercise()) == 400 + assert len(oauth_environment.authorization.state.requests) == 2 + assert oauth_environment.resource.state.requests == [] + + +def test_aiohttp_client_returns_native_token_error(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> tuple[aiohttp.ClientResponse, dict[str, str]]: + async with _client(oauth_environment) as client: + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + return response, await response.json() + + oauth_environment.authorization.state.reject_client_assertion = True + response, body = asyncio.run(exercise()) + + assert isinstance(response, aiohttp.ClientResponse) + assert response.status == 400 + assert body == { + "error": "invalid_client", + "error_description": "Client assertion was rejected", + } + assert oauth_environment.resource.state.requests == [] + + +def test_aiohttp_client_returns_resource_server_error(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> tuple[int, dict[str, str], str | None]: + async with _client(oauth_environment) as client: + response = await client.post( + f"{oauth_environment.resource.base_url}/api/resources", + json={"name": "forbidden"}, + ) + return response.status, await response.json(), response.headers.get("WWW-Authenticate") + + oauth_environment.resource.state.insufficient_scope = True + status, body, challenge = asyncio.run(exercise()) + + assert status == 403 + assert body == { + "error": "insufficient_scope", + "error_description": "requested scope is not permitted", + } + assert challenge == 'DPoP error="insufficient_scope"' + + +def test_aiohttp_client_stops_after_one_resource_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> int: + async with _client(oauth_environment) as client: + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + return response.status + + oauth_environment.resource.state.always_challenge_nonce = True + + assert asyncio.run(exercise()) == 401 + assert len(oauth_environment.resource.state.requests) == 2 + + +def test_aiohttp_final_response_remains_caller_owned(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> bytes: + async with _client(oauth_environment) as client: + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + assert not response.content.at_eof() + return await response.read() + + assert asyncio.run(exercise()) == b'{"error":"not_found"}' + + +def test_aiohttp_middleware_order_exposes_each_resource_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> tuple[list[str], list[str]]: + outer_calls: list[str] = [] + inner_calls: list[str] = [] + + async def outer(request: aiohttp.ClientRequest, handler: aiohttp.ClientHandlerType) -> aiohttp.ClientResponse: + outer_calls.append(str(request.url)) + return await handler(request) + + async def inner(request: aiohttp.ClientRequest, handler: aiohttp.ClientHandlerType) -> aiohttp.ClientResponse: + inner_calls.append(str(request.url)) + return await handler(request) + + ssl_context = ssl.create_default_context(cafile=str(oauth_environment.tls_files.certificate)) + async with aiohttp.ClientSession( + connector=aiohttp.TCPConnector(ssl=ssl_context), + middlewares=(outer, OAuth2Middleware(oauth_environment.config("ec")), inner), + ) as client: + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + await response.read() + return outer_calls, inner_calls + + outer_calls, inner_calls = asyncio.run(exercise()) + + resource_url = f"{oauth_environment.resource.base_url}/api/resources/missing" + assert outer_calls == [resource_url] + assert inner_calls == [resource_url, resource_url] + + +def test_aiohttp_middleware_does_not_own_session(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> tuple[bool, bool]: + client = _client(oauth_environment) + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + await response.read() + open_after_request = not client.closed + await client.close() + return open_after_request, client.closed + + assert asyncio.run(exercise()) == (True, True) + + +def test_aiohttp_retry_client_uses_oauth2_session_middleware( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> int: + async with _client(oauth_environment) as client: + retry_client = aiohttp_retry.RetryClient( + client_session=client, + retry_options=aiohttp_retry.ExponentialRetry(attempts=2), + ) + response = await retry_client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + await response.read() + await retry_client.close() + return response.status + + assert asyncio.run(exercise()) == 404 + assert len(oauth_environment.authorization.state.requests) == 2 + + +def _multipart_parts(content_type: str | None, body: bytes) -> dict[str, tuple[str | None, bytes]]: + if content_type is None: + raise AssertionError("Multipart request must have Content-Type") + message = email.message_from_bytes(f"Content-Type: {content_type}\r\n\r\n".encode() + body) + if not message.is_multipart(): + raise AssertionError("Request body must be multipart") + parts: dict[str, tuple[str | None, bytes]] = {} + for part in message.walk(): + if part.is_multipart(): + continue + name = part.get_param("name", header="content-disposition") + payload = part.get_payload(decode=True) + if not isinstance(name, str) or not isinstance(payload, bytes): + raise TypeError("Multipart part must have a name and byte payload") + parts[name] = (part.get_filename(), payload) + return parts diff --git a/tests/integration/test_httpx.py b/tests/integration/test_httpx.py new file mode 100644 index 0000000..61acbba --- /dev/null +++ b/tests/integration/test_httpx.py @@ -0,0 +1,246 @@ +from __future__ import annotations + +import ssl + +import httpx +import pytest + +from mastercard_oauth2_client import OAuth2Config +from mastercard_oauth2_client.integrations.httpx import OAuth2Transport +from tests.integration.support.environment import OAuthTestEnvironment +from tests.integration.support.lifecycle import run_resource_lifecycle +from tests.integration.support.lifecycle_clients.httpx import HttpxLifecycleClient +from tests.integration.support.models import DPoPAlgorithm +from tests.integration.support.stores import InvalidatableTokenStore + +pytestmark = pytest.mark.httpx + + +def _client( + environment: OAuthTestEnvironment, + algorithm: DPoPAlgorithm = "ec", + *, + config: OAuth2Config | None = None, +) -> httpx.Client: + ssl_context = ssl.create_default_context(cafile=str(environment.tls_files.certificate)) + transport = OAuth2Transport( + config or environment.config(algorithm), + httpx.HTTPTransport(verify=ssl_context), + ) + return httpx.Client(transport=transport) + + +@pytest.mark.parametrize("dpop_algorithm", ["ec", "rsa"]) +def test_httpx_client_completes_stateful_https_lifecycle( + oauth_environment: OAuthTestEnvironment, + dpop_algorithm: DPoPAlgorithm, +) -> None: + with _client(oauth_environment, dpop_algorithm) as client: + resource_id = run_resource_lifecycle(HttpxLifecycleClient(client, oauth_environment.resource.base_url)) + + assert resource_id == "1" + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.authorization.state.requests] == ["POST", "POST"] + assert [request.method for request in oauth_environment.resource.state.requests] == [ + "POST", + "POST", + "GET", + "DELETE", + ] + assert len({request.client_port for request in oauth_environment.authorization.state.requests}) == 1 + assert len({request.client_port for request in oauth_environment.resource.state.requests}) == 1 + + +def test_httpx_lifecycle_attempts_delete_after_get_failure(oauth_environment: OAuthTestEnvironment) -> None: + oauth_environment.resource.state.fail_get = True + + with _client(oauth_environment) as client, pytest.raises(httpx.HTTPStatusError) as error: + run_resource_lifecycle(HttpxLifecycleClient(client, oauth_environment.resource.base_url)) + + assert error.value.response.status_code == 500 + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.resource.state.requests][-2:] == ["GET", "DELETE"] + + +def test_httpx_client_preserves_query_json_and_caller_headers(oauth_environment: OAuthTestEnvironment) -> None: + with _client(oauth_environment) as client: + create_response = client.post( + f"{oauth_environment.resource.base_url}/api/resources", + params={"source": "httpx"}, + json={"name": "httpx-resource"}, + headers={ + "Accept": "application/json", + "User-Agent": "caller-agent", + "Authorization": "Bearer stale", + "DPoP": "stale-proof", + }, + ) + create_response.raise_for_status() + resource_id = create_response.json()["id"] + try: + recorded = oauth_environment.resource.state.requests[-1] + assert recorded.path == "/api/resources?source=httpx" + assert recorded.body == b'{"name":"httpx-resource"}' + assert recorded.header("user-agent") == "caller-agent" + assert recorded.header("authorization") == "DPoP integration-access-token" + assert recorded.header("dpop") != "stale-proof" + assert create_response.request.url.query == b"source=httpx" + assert create_response.request.headers["Authorization"] == "DPoP integration-access-token" + assert create_response.history == [] + finally: + client.delete(f"{oauth_environment.resource.base_url}/api/resources/{resource_id}").raise_for_status() + + +def test_httpx_client_reuses_cached_token(oauth_environment: OAuthTestEnvironment) -> None: + with _client(oauth_environment) as client: + first = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + second = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + assert first.status_code == 404 + assert second.status_code == 404 + assert len(oauth_environment.authorization.state.requests) == 2 + + +def test_httpx_client_reacquires_invalidated_token(oauth_environment: OAuthTestEnvironment) -> None: + store = InvalidatableTokenStore() + config = oauth_environment.config("ec", token_store=store) + with _client(oauth_environment, config=config) as client: + first = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + store.invalidate() + second = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + assert first.status_code == 404 + assert second.status_code == 404 + assert len(oauth_environment.authorization.state.requests) == 4 + + +def test_httpx_client_stops_after_one_token_nonce_retry(oauth_environment: OAuthTestEnvironment) -> None: + oauth_environment.authorization.state.always_challenge_nonce = True + + with _client(oauth_environment) as client: + response = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + assert response.status_code == 400 + assert response.history == [] + assert len(oauth_environment.authorization.state.requests) == 2 + assert oauth_environment.resource.state.requests == [] + + +def test_httpx_client_returns_native_token_error_without_history(oauth_environment: OAuthTestEnvironment) -> None: + oauth_environment.authorization.state.reject_client_assertion = True + + with _client(oauth_environment) as client: + response = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + assert isinstance(response, httpx.Response) + assert response.status_code == 400 + assert response.history == [] + assert response.json() == { + "error": "invalid_client", + "error_description": "Client assertion was rejected", + } + assert oauth_environment.resource.state.requests == [] + + +def test_httpx_client_returns_resource_server_error(oauth_environment: OAuthTestEnvironment) -> None: + oauth_environment.resource.state.insufficient_scope = True + + with _client(oauth_environment) as client: + response = client.post( + f"{oauth_environment.resource.base_url}/api/resources", + json={"name": "forbidden"}, + ) + + assert response.status_code == 403 + assert response.json() == { + "error": "insufficient_scope", + "error_description": "requested scope is not permitted", + } + assert response.headers["WWW-Authenticate"] == 'DPoP error="insufficient_scope"' + + +def test_httpx_client_stops_after_one_resource_nonce_retry(oauth_environment: OAuthTestEnvironment) -> None: + oauth_environment.resource.state.always_challenge_nonce = True + + with _client(oauth_environment) as client: + response = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + assert response.status_code == 401 + assert response.history == [] + assert len(oauth_environment.resource.state.requests) == 2 + + +def test_httpx_streaming_response_remains_caller_owned(oauth_environment: OAuthTestEnvironment) -> None: + with _client(oauth_environment) as client, client.stream( + "GET", + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) as response: + assert response.status_code == 404 + assert not response.is_stream_consumed + assert response.read() == b'{"error":"not_found"}' + + +def test_httpx_client_preserves_resource_event_hooks(oauth_environment: OAuthTestEnvironment) -> None: + request_urls: list[str] = [] + response_statuses: list[int] = [] + ssl_context = ssl.create_default_context(cafile=str(oauth_environment.tls_files.certificate)) + transport = OAuth2Transport( + oauth_environment.config("ec"), + httpx.HTTPTransport(verify=ssl_context), + ) + + with httpx.Client( + transport=transport, + event_hooks={ + "request": [lambda request: request_urls.append(str(request.url))], + "response": [lambda response: response_statuses.append(response.status_code)], + }, + ) as client: + response = client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + assert response.status_code == 404 + assert request_urls == [f"{oauth_environment.resource.base_url}/api/resources/missing"] + assert response_statuses == [404] + + +def test_httpx_transport_closes_wrapped_transport(oauth_environment: OAuthTestEnvironment) -> None: + class TrackingTransport(httpx.BaseTransport): + closed = False + + def handle_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request) + + def close(self) -> None: + self.closed = True + + wrapped = TrackingTransport() + transport = OAuth2Transport(oauth_environment.config("ec"), wrapped) + + transport.close() + + assert wrapped.closed diff --git a/tests/integration/test_httpx_async.py b/tests/integration/test_httpx_async.py new file mode 100644 index 0000000..fd354ed --- /dev/null +++ b/tests/integration/test_httpx_async.py @@ -0,0 +1,299 @@ +from __future__ import annotations + +import asyncio +import ssl + +import httpx +import pytest + +from mastercard_oauth2_client import OAuth2Config +from mastercard_oauth2_client.integrations.httpx import AsyncOAuth2Transport +from tests.integration.support.environment import OAuthTestEnvironment +from tests.integration.support.lifecycle import run_async_resource_lifecycle +from tests.integration.support.lifecycle_clients.httpx import AsyncHttpxLifecycleClient +from tests.integration.support.models import DPoPAlgorithm, RecordedRequest +from tests.integration.support.stores import InvalidatableTokenStore + +pytestmark = pytest.mark.httpx + + +def _client( + environment: OAuthTestEnvironment, + algorithm: DPoPAlgorithm = "ec", + *, + config: OAuth2Config | None = None, +) -> httpx.AsyncClient: + ssl_context = ssl.create_default_context(cafile=str(environment.tls_files.certificate)) + transport = AsyncOAuth2Transport( + config or environment.config(algorithm), + httpx.AsyncHTTPTransport(verify=ssl_context), + ) + return httpx.AsyncClient(transport=transport) + + +@pytest.mark.parametrize("dpop_algorithm", ["ec", "rsa"]) +def test_async_httpx_client_completes_stateful_https_lifecycle( + oauth_environment: OAuthTestEnvironment, + dpop_algorithm: DPoPAlgorithm, +) -> None: + async def exercise() -> str: + async with _client(oauth_environment, dpop_algorithm) as client: + lifecycle_client = AsyncHttpxLifecycleClient(client, oauth_environment.resource.base_url) + return await run_async_resource_lifecycle(lifecycle_client) + + resource_id = asyncio.run(exercise()) + + assert resource_id == "1" + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.authorization.state.requests] == ["POST", "POST"] + assert [request.method for request in oauth_environment.resource.state.requests] == [ + "POST", + "POST", + "GET", + "DELETE", + ] + assert len({request.client_port for request in oauth_environment.authorization.state.requests}) == 1 + assert len({request.client_port for request in oauth_environment.resource.state.requests}) == 1 + + +def test_async_httpx_lifecycle_attempts_delete_after_get_failure( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + async with _client(oauth_environment) as client: + lifecycle_client = AsyncHttpxLifecycleClient(client, oauth_environment.resource.base_url) + await run_async_resource_lifecycle(lifecycle_client) + + oauth_environment.resource.state.fail_get = True + + with pytest.raises(httpx.HTTPStatusError) as error: + asyncio.run(exercise()) + + assert error.value.response.status_code == 500 + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.resource.state.requests][-2:] == ["GET", "DELETE"] + + +def test_async_httpx_client_preserves_query_json_and_caller_headers( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> tuple[httpx.Response, RecordedRequest]: + async with _client(oauth_environment) as client: + create_response = await client.post( + f"{oauth_environment.resource.base_url}/api/resources", + params={"source": "httpx-async"}, + json={"name": "httpx-async-resource"}, + headers={ + "Accept": "application/json", + "User-Agent": "caller-agent", + "Authorization": "Bearer stale", + "DPoP": "stale-proof", + }, + ) + create_response.raise_for_status() + resource_id = create_response.json()["id"] + try: + return create_response, oauth_environment.resource.state.requests[-1] + finally: + response = await client.delete(f"{oauth_environment.resource.base_url}/api/resources/{resource_id}") + response.raise_for_status() + + create_response, recorded = asyncio.run(exercise()) + assert recorded.path == "/api/resources?source=httpx-async" + assert recorded.body == b'{"name":"httpx-async-resource"}' + assert recorded.header("user-agent") == "caller-agent" + assert recorded.header("authorization") == "DPoP integration-access-token" + assert recorded.header("dpop") != "stale-proof" + assert create_response.request.url.query == b"source=httpx-async" + assert create_response.request.headers["Authorization"] == "DPoP integration-access-token" + assert create_response.history == [] + + +def test_async_httpx_client_reuses_cached_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> tuple[httpx.Response, httpx.Response]: + async with _client(oauth_environment) as client: + first = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + second = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + return first, second + + first, second = asyncio.run(exercise()) + + assert first.status_code == 404 + assert second.status_code == 404 + assert len(oauth_environment.authorization.state.requests) == 2 + + +def test_async_httpx_client_reacquires_invalidated_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise(config: OAuth2Config, store: InvalidatableTokenStore) -> tuple[httpx.Response, httpx.Response]: + async with _client(oauth_environment, config=config) as client: + first = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + store.invalidate() + second = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + return first, second + + store = InvalidatableTokenStore() + config = oauth_environment.config("ec", token_store=store) + first, second = asyncio.run(exercise(config, store)) + + assert first.status_code == 404 + assert second.status_code == 404 + assert len(oauth_environment.authorization.state.requests) == 4 + + +def test_async_httpx_client_stops_after_one_token_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> httpx.Response: + async with _client(oauth_environment) as client: + return await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + oauth_environment.authorization.state.always_challenge_nonce = True + response = asyncio.run(exercise()) + + assert response.status_code == 400 + assert response.history == [] + assert len(oauth_environment.authorization.state.requests) == 2 + assert oauth_environment.resource.state.requests == [] + + +def test_async_httpx_client_returns_native_token_error_without_history( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> httpx.Response: + async with _client(oauth_environment) as client: + return await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + oauth_environment.authorization.state.reject_client_assertion = True + response = asyncio.run(exercise()) + + assert isinstance(response, httpx.Response) + assert response.status_code == 400 + assert response.history == [] + assert response.json() == { + "error": "invalid_client", + "error_description": "Client assertion was rejected", + } + assert oauth_environment.resource.state.requests == [] + + +def test_async_httpx_client_returns_resource_server_error(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> httpx.Response: + async with _client(oauth_environment) as client: + return await client.post( + f"{oauth_environment.resource.base_url}/api/resources", + json={"name": "forbidden"}, + ) + + oauth_environment.resource.state.insufficient_scope = True + response = asyncio.run(exercise()) + + assert response.status_code == 403 + assert response.json() == { + "error": "insufficient_scope", + "error_description": "requested scope is not permitted", + } + assert response.headers["WWW-Authenticate"] == 'DPoP error="insufficient_scope"' + + +def test_async_httpx_client_stops_after_one_resource_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> httpx.Response: + async with _client(oauth_environment) as client: + return await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + + oauth_environment.resource.state.always_challenge_nonce = True + response = asyncio.run(exercise()) + + assert response.status_code == 401 + assert response.history == [] + assert len(oauth_environment.resource.state.requests) == 2 + + +def test_async_httpx_streaming_response_remains_caller_owned( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + async with _client(oauth_environment) as client, client.stream( + "GET", + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) as response: + assert response.status_code == 404 + assert not response.is_stream_consumed + assert await response.aread() == b'{"error":"not_found"}' + + asyncio.run(exercise()) + + +def test_async_httpx_client_preserves_resource_event_hooks(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> tuple[httpx.Response, list[str], list[int]]: + request_urls: list[str] = [] + response_statuses: list[int] = [] + + async def record_request(request: httpx.Request) -> None: + request_urls.append(str(request.url)) + + async def record_response(response: httpx.Response) -> None: + response_statuses.append(response.status_code) + + ssl_context = ssl.create_default_context(cafile=str(oauth_environment.tls_files.certificate)) + transport = AsyncOAuth2Transport( + oauth_environment.config("ec"), + httpx.AsyncHTTPTransport(verify=ssl_context), + ) + async with httpx.AsyncClient( + transport=transport, + event_hooks={"request": [record_request], "response": [record_response]}, + ) as client: + response = await client.get( + f"{oauth_environment.resource.base_url}/api/resources/missing", + headers={"Accept": "application/json"}, + ) + return response, request_urls, response_statuses + + response, request_urls, response_statuses = asyncio.run(exercise()) + + assert response.status_code == 404 + assert request_urls == [f"{oauth_environment.resource.base_url}/api/resources/missing"] + assert response_statuses == [404] + + +def test_async_httpx_transport_closes_wrapped_transport(oauth_environment: OAuthTestEnvironment) -> None: + class TrackingTransport(httpx.AsyncBaseTransport): + closed = False + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request) + + async def aclose(self) -> None: + self.closed = True + + async def exercise() -> bool: + wrapped = TrackingTransport() + transport = AsyncOAuth2Transport(oauth_environment.config("ec"), wrapped) + await transport.aclose() + return wrapped.closed + + assert asyncio.run(exercise()) diff --git a/tests/integration/test_openapi.py b/tests/integration/test_openapi.py new file mode 100644 index 0000000..105d89c --- /dev/null +++ b/tests/integration/test_openapi.py @@ -0,0 +1,310 @@ +from __future__ import annotations + +import json +from urllib.parse import parse_qsl, urlencode + +import pytest +import urllib3 +from generated_test_client import Resource # type: ignore[import-not-found] +from generated_test_client.exceptions import ApiException # type: ignore[import-not-found] +from urllib3._collections import HTTPHeaderDict + +from mastercard_oauth2_client import OAuth2Config, build_token_request, create_resource_dpop_proof +from tests.integration.support.environment import OAuthTestEnvironment +from tests.integration.support.lifecycle import run_resource_lifecycle +from tests.integration.support.lifecycle_clients.openapi_urllib3 import OpenApiLifecycleClient +from tests.integration.support.models import DPoPAlgorithm +from tests.integration.support.stores import InvalidatableTokenStore + +pytestmark = [pytest.mark.generated, pytest.mark.urllib3] + + +def _client( + environment: OAuthTestEnvironment, + algorithm: DPoPAlgorithm = "ec", + *, + config: OAuth2Config | None = None, +) -> OpenApiLifecycleClient: + return OpenApiLifecycleClient.create( + environment.resource.base_url, + str(environment.tls_files.certificate), + config or environment.config(algorithm), + ) + + +@pytest.mark.parametrize("dpop_algorithm", ["ec", "rsa"]) +def test_openapi_client_completes_stateful_https_lifecycle( + oauth_environment: OAuthTestEnvironment, + dpop_algorithm: DPoPAlgorithm, +) -> None: + """Prove token/resource nonce handling and POST/GET/DELETE for each DPoP key type.""" + + resource_id = run_resource_lifecycle(_client(oauth_environment, dpop_algorithm)) + + assert resource_id == "1" + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.authorization.state.requests] == ["POST", "POST"] + assert [request.method for request in oauth_environment.resource.state.requests] == [ + "POST", + "POST", + "GET", + "DELETE", + ] + + +def test_openapi_lifecycle_attempts_delete_after_get_failure(oauth_environment: OAuthTestEnvironment) -> None: + """Prevent failed lifecycle assertions from leaving fake resources behind.""" + + oauth_environment.resource.state.fail_get = True + + with pytest.raises(ApiException) as error: + run_resource_lifecycle(_client(oauth_environment)) + + assert error.value.status == 500 + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.resource.state.requests][-2:] == ["GET", "DELETE"] + + +def test_openapi_client_preserves_query_json_and_caller_headers(oauth_environment: OAuthTestEnvironment) -> None: + client = _client(oauth_environment) + client.api.api_client.user_agent = "caller-agent" + resource = client.api.create_resource( + Resource(name="openapi-resource"), + source="openapi", + _headers={ + "Authorization": "Bearer stale", + "DPoP": "stale-proof", + }, + ) + assert resource.id is not None + try: + recorded = oauth_environment.resource.state.requests[-1] + assert recorded.path == "/api/resources?source=openapi" + assert recorded.body == b'{"name": "openapi-resource"}' + assert recorded.header("user-agent") == "caller-agent" + assert recorded.header("authorization") == "DPoP integration-access-token" + assert recorded.header("dpop") != "stale-proof" + finally: + client.delete_resource(resource.id) + + +def test_openapi_without_preload_content_returns_readable_raw_response( + oauth_environment: OAuthTestEnvironment, +) -> None: + """Leave successful resource responses available to generated streaming methods.""" + + client = _client(oauth_environment) + response = client.api.create_resource_without_preload_content(Resource(name="raw-response")) + try: + payload = json.loads(response.read()) + resource_id = payload["id"] + assert response.status == 201 + assert resource_id == "1" + finally: + response.release_conn() + + client.delete_resource(resource_id) + + +def test_openapi_client_reuses_cached_token(oauth_environment: OAuthTestEnvironment) -> None: + client = _client(oauth_environment) + + for _ in range(2): + with pytest.raises(ApiException) as error: + client.get_resource("missing") + assert error.value.status == 404 + + assert len(oauth_environment.authorization.state.requests) == 2 + + +def test_openapi_client_returns_configured_server_errors(oauth_environment: OAuthTestEnvironment) -> None: + """Preserve authorization-server and resource-server failures for callers.""" + + oauth_environment.authorization.state.reject_client_assertion = True + with pytest.raises(ApiException) as assertion_error: + _client(oauth_environment).create_resource("rejected") + assert assertion_error.value.status == 400 + assert json.loads(assertion_error.value.body) == { + "error": "invalid_client", + "error_description": "Client assertion was rejected", + } + + oauth_environment.reset() + oauth_environment.resource.state.insufficient_scope = True + with pytest.raises(ApiException) as scope_error: + _client(oauth_environment).create_resource("forbidden") + assert scope_error.value.status == 403 + assert json.loads(scope_error.value.body) == { + "error": "insufficient_scope", + "error_description": "requested scope is not permitted", + } + assert scope_error.value.headers["WWW-Authenticate"] == 'DPoP error="insufficient_scope"' + + +def test_openapi_client_stops_after_one_nonce_retry(oauth_environment: OAuthTestEnvironment) -> None: + """Fail closed after one authorization-server nonce retry.""" + + oauth_environment.authorization.state.always_challenge_nonce = True + + with pytest.raises(ApiException) as error: + _client(oauth_environment).create_resource("never-created") + + assert error.value.status == 400 + assert len(oauth_environment.authorization.state.requests) == 2 + assert oauth_environment.resource.state.requests == [] + + +def test_openapi_client_stops_after_one_resource_nonce_retry(oauth_environment: OAuthTestEnvironment) -> None: + """Fail closed after one resource-server nonce retry.""" + + oauth_environment.resource.state.always_challenge_nonce = True + + with pytest.raises(ApiException) as error: + _client(oauth_environment).create_resource("never-created") + + assert error.value.status == 401 + assert len(oauth_environment.resource.state.requests) == 2 + + +def test_openapi_client_reacquires_expired_token(oauth_environment: OAuthTestEnvironment) -> None: + """Replace an unusable cached token before the next resource request.""" + + store = InvalidatableTokenStore() + config = oauth_environment.config("ec", token_store=store) + client = _client(oauth_environment, config=config) + resource_id = client.create_resource("expires") + store.invalidate() + try: + assert client.get_resource(resource_id) == resource_id + finally: + client.delete_resource(resource_id) + + # The shared handler nonce contains the resource nonce, so token reacquisition + # receives a fresh authorization-server challenge before succeeding. + assert len(oauth_environment.authorization.state.requests) == 4 + + +def test_authorization_server_rejects_tampered_client_assertion(oauth_environment: OAuthTestEnvironment) -> None: + """Prove the fake token endpoint independently checks assertion signatures.""" + + config = oauth_environment.config("ec") + dpop_key = config.dpop_key_provider.get_current_key() + request = build_token_request(config, config.scope_resolver.all_scopes(), dpop_key.key_id) + form = dict(parse_qsl(str(request.body))) + form["client_assertion"] = _tamper(form["client_assertion"]) + pool = urllib3.PoolManager(ca_certs=str(oauth_environment.tls_files.certificate)) + + response = pool.request(request.method, request.url, headers=request.headers, body=urlencode(form)) + + assert response.status == 400 + assert json.loads(response.data)["error"] == "invalid_client" + + +def test_resource_server_rejects_tampered_dpop_signature(oauth_environment: OAuthTestEnvironment) -> None: + """Prove the fake resource endpoint independently checks DPoP signatures.""" + + client = _client(oauth_environment) + resource_id = client.create_resource("tamper-target") + original = oauth_environment.resource.state.requests[-1] + authorization = original.header("authorization") + dpop = original.header("dpop") + assert authorization is not None + assert dpop is not None + headers = { + "Accept": "application/json", + "Authorization": authorization, + "DPoP": _tamper(dpop), + "User-Agent": "integration-test", + } + pool = urllib3.PoolManager(ca_certs=str(oauth_environment.tls_files.certificate)) + try: + response = pool.request( + "GET", + f"{oauth_environment.resource.base_url}/api/resources/{resource_id}", + headers=headers, + timeout=urllib3.Timeout(connect=2, read=2), + ) + assert response.status == 401 + assert json.loads(response.data)["error"] == "invalid_dpop_proof" + finally: + client.delete_resource(resource_id) + + +def test_resource_server_rejects_incorrect_ath(oauth_environment: OAuthTestEnvironment) -> None: + """Reject a validly signed proof bound to a different access token.""" + + config = oauth_environment.config("ec") + client = _client(oauth_environment, config=config) + resource_id = client.create_resource("wrong-ath") + dpop_key = config.dpop_key_provider.get_current_key() + resource_url = f"{oauth_environment.resource.base_url}/api/resources/{resource_id}" + proof = create_resource_dpop_proof( + config, + dpop_key.key_id, + "GET", + resource_url, + "different-access-token", + oauth_environment.resource.state.nonce, + ) + pool = urllib3.PoolManager(ca_certs=str(oauth_environment.tls_files.certificate)) + try: + response = pool.request( + "GET", + resource_url, + headers={ + "Accept": "application/json", + "Authorization": f"DPoP {oauth_environment.resource.state.access_token}", + "DPoP": proof, + "User-Agent": "integration-test", + }, + timeout=urllib3.Timeout(connect=2, read=2), + ) + assert response.status == 401 + assert json.loads(response.data)["error"] == "invalid_dpop_proof" + finally: + client.delete_resource(resource_id) + + +def test_resource_server_rejects_duplicate_authorization_header(oauth_environment: OAuthTestEnvironment) -> None: + """Reject ambiguous requests containing multiple Authorization values.""" + + config = oauth_environment.config("ec") + client = _client(oauth_environment, config=config) + resource_id = client.create_resource("duplicate-header") + dpop_key = config.dpop_key_provider.get_current_key() + resource_url = f"{oauth_environment.resource.base_url}/api/resources/{resource_id}" + proof = create_resource_dpop_proof( + config, + dpop_key.key_id, + "GET", + resource_url, + oauth_environment.resource.state.access_token, + oauth_environment.resource.state.nonce, + ) + headers = HTTPHeaderDict( + { + "Accept": "application/json", + "DPoP": proof, + "User-Agent": "integration-test", + } + ) + headers.add("Authorization", f"DPoP {oauth_environment.resource.state.access_token}") + headers.add("Authorization", f"DPoP {oauth_environment.resource.state.access_token}") + pool = urllib3.PoolManager(ca_certs=str(oauth_environment.tls_files.certificate)) + try: + response = pool.request( + "GET", + resource_url, + headers=headers, + timeout=urllib3.Timeout(connect=2, read=2), + ) + assert response.status == 401 + assert json.loads(response.data)["error"] == "invalid_token" + finally: + client.delete_resource(resource_id) + + +def _tamper(compact_jwt: str) -> str: + encoded_header, encoded_claims, signature = compact_jwt.split(".") + replacement = "A" if signature[0] != "A" else "B" + return f"{encoded_header}.{encoded_claims}.{replacement}{signature[1:]}" diff --git a/tests/integration/test_openapi_asyncio.py b/tests/integration/test_openapi_asyncio.py new file mode 100644 index 0000000..f74eb13 --- /dev/null +++ b/tests/integration/test_openapi_asyncio.py @@ -0,0 +1,340 @@ +from __future__ import annotations + +import asyncio +import json + +import aiohttp +import aiohttp_retry +import pytest +from generated_asyncio_test_client import ( # type: ignore[import-not-found] + ApiClient as GeneratedApiClient, +) +from generated_asyncio_test_client import ( + Configuration as GeneratedConfiguration, +) +from generated_asyncio_test_client import ( + Resource, + ResourcesApi, +) +from generated_asyncio_test_client.exceptions import ApiException # type: ignore[import-not-found] + +from mastercard_oauth2_client import OAuth2Config +from mastercard_oauth2_client.integrations.openapi.asyncio import add_oauth2_layer +from tests.integration.support.environment import OAuthTestEnvironment +from tests.integration.support.lifecycle import run_async_resource_lifecycle +from tests.integration.support.lifecycle_clients.openapi_asyncio import OpenApiAsyncioLifecycleClient + +pytestmark = [pytest.mark.generated, pytest.mark.aiohttp] +from tests.integration.support.models import DPoPAlgorithm +from tests.integration.support.stores import InvalidatableTokenStore + + +def _client( + environment: OAuthTestEnvironment, + algorithm: DPoPAlgorithm = "ec", + *, + config: OAuth2Config | None = None, +) -> OpenApiAsyncioLifecycleClient: + return OpenApiAsyncioLifecycleClient.create( + environment.resource.base_url, + str(environment.tls_files.certificate), + config or environment.config(algorithm), + ) + + +@pytest.mark.parametrize("dpop_algorithm", ["ec", "rsa"]) +def test_openapi_asyncio_client_completes_stateful_https_lifecycle( + oauth_environment: OAuthTestEnvironment, + dpop_algorithm: DPoPAlgorithm, +) -> None: + async def exercise() -> str: + client = _client(oauth_environment, dpop_algorithm) + try: + return await run_async_resource_lifecycle(client) + finally: + await client.close() + + resource_id = asyncio.run(exercise()) + + assert resource_id == "1" + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.authorization.state.requests] == ["POST", "POST"] + assert [request.method for request in oauth_environment.resource.state.requests] == [ + "POST", + "POST", + "GET", + "DELETE", + ] + assert len({request.client_port for request in oauth_environment.authorization.state.requests}) == 1 + assert len({request.client_port for request in oauth_environment.resource.state.requests}) == 1 + + +def test_openapi_asyncio_lifecycle_attempts_delete_after_get_failure( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + await run_async_resource_lifecycle(client) + finally: + await client.close() + + oauth_environment.resource.state.fail_get = True + + with pytest.raises(ApiException) as error: + asyncio.run(exercise()) + + assert error.value.status == 500 + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.resource.state.requests][-2:] == ["GET", "DELETE"] + + +def test_openapi_asyncio_client_preserves_query_json_and_caller_headers( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> Resource: + client = _client(oauth_environment) + client.api.api_client.user_agent = "caller-agent" + resource_id: str | None = None + try: + resource = await client.api.create_resource( + Resource(name="openapi-asyncio-resource"), + source="openapi-asyncio", + _headers={"Authorization": "Bearer stale", "DPoP": "stale-proof"}, + ) + if not isinstance(resource.id, str): + raise TypeError("Created resource has no ID") + resource_id = resource.id + return resource + finally: + if resource_id is not None: + await client.api.delete_resource(resource_id) + await client.close() + + resource = asyncio.run(exercise()) + recorded = oauth_environment.resource.state.requests[-2] + + assert resource.name == "openapi-asyncio-resource" + assert recorded.path == "/api/resources?source=openapi-asyncio" + assert recorded.body == b'{"name": "openapi-asyncio-resource"}' + assert recorded.header("user-agent") == "caller-agent" + assert recorded.header("authorization") == "DPoP integration-access-token" + assert recorded.header("dpop") != "stale-proof" + + +def test_openapi_asyncio_raw_method_returns_native_readable_response( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> tuple[aiohttp.ClientResponse, str]: + client = _client(oauth_environment) + response = await client.api.create_resource_without_preload_content(Resource(name="raw-response")) + assert isinstance(response, aiohttp.ClientResponse) + assert not response.content.at_eof() + value = await response.json() + resource_id = value["id"] + await client.api.delete_resource(resource_id) + await client.close() + return response, resource_id + + response, resource_id = asyncio.run(exercise()) + + assert response.status == 201 + assert resource_id == "1" + + +def test_openapi_asyncio_client_reuses_cached_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + for _ in range(2): + with pytest.raises(ApiException) as error: + await client.api.get_resource("missing") + assert error.value.status == 404 + finally: + await client.close() + + asyncio.run(exercise()) + + assert len(oauth_environment.authorization.state.requests) == 2 + + +def test_openapi_asyncio_client_reacquires_invalidated_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise(config: OAuth2Config, store: InvalidatableTokenStore) -> None: + client = _client(oauth_environment, config=config) + try: + resource_id = await client.create_resource("expires") + store.invalidate() + assert await client.get_resource(resource_id) == resource_id + await client.delete_resource(resource_id) + finally: + await client.close() + + store = InvalidatableTokenStore() + config = oauth_environment.config("ec", token_store=store) + asyncio.run(exercise(config, store)) + + assert len(oauth_environment.authorization.state.requests) == 4 + + +def test_openapi_asyncio_client_stops_after_one_token_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + await client.create_resource("never-created") + finally: + await client.close() + + oauth_environment.authorization.state.always_challenge_nonce = True + + with pytest.raises(ApiException) as error: + asyncio.run(exercise()) + + assert error.value.status == 400 + assert len(oauth_environment.authorization.state.requests) == 2 + assert oauth_environment.resource.state.requests == [] + + +def test_openapi_asyncio_client_returns_configured_server_errors( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def create(name: str) -> None: + client = _client(oauth_environment) + try: + await client.create_resource(name) + finally: + await client.close() + + oauth_environment.authorization.state.reject_client_assertion = True + with pytest.raises(ApiException) as assertion_error: + asyncio.run(create("rejected")) + assert assertion_error.value.status == 400 + assert json.loads(assertion_error.value.body) == { + "error": "invalid_client", + "error_description": "Client assertion was rejected", + } + + oauth_environment.reset() + oauth_environment.resource.state.insufficient_scope = True + with pytest.raises(ApiException) as scope_error: + asyncio.run(create("forbidden")) + assert scope_error.value.status == 403 + assert json.loads(scope_error.value.body) == { + "error": "insufficient_scope", + "error_description": "requested scope is not permitted", + } + assert scope_error.value.headers["WWW-Authenticate"] == 'DPoP error="insufficient_scope"' + + +def test_openapi_asyncio_client_stops_after_one_resource_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + await client.create_resource("never-created") + finally: + await client.close() + + oauth_environment.resource.state.always_challenge_nonce = True + + with pytest.raises(ApiException) as error: + asyncio.run(exercise()) + + assert error.value.status == 401 + assert len(oauth_environment.resource.state.requests) == 2 + + +def test_openapi_asyncio_preserves_generated_session_settings_and_middleware_order( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> tuple[list[str], list[str], int, int, bool]: + middleware_urls: list[str] = [] + traced_urls: list[str] = [] + + async def record_middleware( + request: aiohttp.ClientRequest, + handler: aiohttp.ClientHandlerType, + ) -> aiohttp.ClientResponse: + middleware_urls.append(str(request.url)) + return await handler(request) + + async def record_trace( + session: aiohttp.ClientSession, + context: object, + params: aiohttp.TraceRequestStartParams, + ) -> None: + traced_urls.append(str(params.url)) + + trace_config = aiohttp.TraceConfig() + trace_config.on_request_start.append(record_trace) + timeout = aiohttp.ClientTimeout(total=17) + generated_config = GeneratedConfiguration( + host=oauth_environment.resource.base_url, + ssl_ca_cert=str(oauth_environment.tls_files.certificate), + trace_configs=[trace_config], + tcp_connector_limit_per_host=3, + connection_pool_maxsize=7, + client_session_kwargs={"timeout": timeout, "middlewares": (record_middleware,)}, + ) + api_client = GeneratedApiClient(generated_config) + add_oauth2_layer(api_client, oauth_environment.config("ec")) + oauth_environment.resource.state.always_challenge_nonce = True + try: + with pytest.raises(ApiException): + await ResourcesApi(api_client).create_resource(Resource(name="never-created")) + pool = api_client.rest_client.pool_manager + if pool is None: + raise AssertionError("Generated asyncio pool was not created") + connector = pool.connector + return middleware_urls, traced_urls, connector.limit, connector.limit_per_host, pool.timeout is timeout + finally: + pool = api_client.rest_client.pool_manager + await api_client.close() + pool_closed = pool is not None and pool.closed + assert pool_closed + + middleware_urls, traced_urls, limit, limit_per_host, timeout_preserved = asyncio.run(exercise()) + + resource_url = f"{oauth_environment.resource.base_url}/api/resources" + assert middleware_urls == [resource_url, resource_url] + assert traced_urls == [ + resource_url, + f"{oauth_environment.authorization.base_url}/oauth/token", + f"{oauth_environment.authorization.base_url}/oauth/token", + ] + assert limit == 7 + assert limit_per_host == 3 + assert timeout_preserved + + +def test_openapi_asyncio_retry_client_reuses_oauth2_session_and_cached_token( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> bool: + generated_config = GeneratedConfiguration( + host=oauth_environment.resource.base_url, + ssl_ca_cert=str(oauth_environment.tls_files.certificate), + retries=aiohttp_retry.ExponentialRetry(attempts=2, statuses={500}), + ) + api_client = GeneratedApiClient(generated_config) + add_oauth2_layer(api_client, oauth_environment.config("ec")) + api = ResourcesApi(api_client) + resource = await api.create_resource(Resource(name="retry-resource")) + if not isinstance(resource.id, str): + raise TypeError("Created resource has no ID") + oauth_environment.resource.state.fail_get = True + try: + with pytest.raises(ApiException) as error: + await api.get_resource(resource.id) + assert error.value.status == 500 + return api_client.rest_client.retry_client is not None + finally: + oauth_environment.resource.state.fail_get = False + await api.delete_resource(resource.id) + await api_client.close() + + assert asyncio.run(exercise()) + assert len(oauth_environment.authorization.state.requests) == 2 + assert [request.method for request in oauth_environment.resource.state.requests].count("GET") == 2 diff --git a/tests/integration/test_openapi_httpx.py b/tests/integration/test_openapi_httpx.py new file mode 100644 index 0000000..9a8a8db --- /dev/null +++ b/tests/integration/test_openapi_httpx.py @@ -0,0 +1,248 @@ +from __future__ import annotations + +import asyncio +import json + +import httpx +import pytest +from generated_httpx_test_client import Resource # type: ignore[import-not-found] +from generated_httpx_test_client.exceptions import ApiException # type: ignore[import-not-found] + +from mastercard_oauth2_client import OAuth2Config +from tests.integration.support.environment import OAuthTestEnvironment +from tests.integration.support.lifecycle import run_async_resource_lifecycle +from tests.integration.support.lifecycle_clients.openapi_httpx import OpenApiHttpxLifecycleClient +from tests.integration.support.models import DPoPAlgorithm +from tests.integration.support.stores import InvalidatableTokenStore + +pytestmark = [pytest.mark.generated, pytest.mark.httpx] + + +def _client( + environment: OAuthTestEnvironment, + algorithm: DPoPAlgorithm = "ec", + *, + config: OAuth2Config | None = None, +) -> OpenApiHttpxLifecycleClient: + return OpenApiHttpxLifecycleClient.create( + environment.resource.base_url, + str(environment.tls_files.certificate), + config or environment.config(algorithm), + ) + + +@pytest.mark.parametrize("dpop_algorithm", ["ec", "rsa"]) +def test_openapi_httpx_client_completes_stateful_https_lifecycle( + oauth_environment: OAuthTestEnvironment, + dpop_algorithm: DPoPAlgorithm, +) -> None: + async def exercise() -> str: + client = _client(oauth_environment, dpop_algorithm) + try: + return await run_async_resource_lifecycle(client) + finally: + await client.close() + + resource_id = asyncio.run(exercise()) + + assert resource_id == "1" + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.authorization.state.requests] == ["POST", "POST"] + assert [request.method for request in oauth_environment.resource.state.requests] == [ + "POST", + "POST", + "GET", + "DELETE", + ] + assert len({request.client_port for request in oauth_environment.authorization.state.requests}) == 1 + assert len({request.client_port for request in oauth_environment.resource.state.requests}) == 1 + + +def test_openapi_httpx_lifecycle_attempts_delete_after_get_failure( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + await run_async_resource_lifecycle(client) + finally: + await client.close() + + oauth_environment.resource.state.fail_get = True + + with pytest.raises(ApiException) as error: + asyncio.run(exercise()) + + assert error.value.status == 500 + assert oauth_environment.resource.state.resources == {} + assert [request.method for request in oauth_environment.resource.state.requests][-2:] == ["GET", "DELETE"] + + +def test_openapi_httpx_client_preserves_query_json_and_caller_headers( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> Resource: + client = _client(oauth_environment) + client.api.api_client.user_agent = "caller-agent" + resource_id: str | None = None + try: + resource = await client.api.create_resource( + Resource(name="openapi-httpx-resource"), + source="openapi-httpx", + _headers={"Authorization": "Bearer stale", "DPoP": "stale-proof"}, + ) + if not isinstance(resource.id, str): + raise TypeError("Created resource has no ID") + resource_id = resource.id + return resource + finally: + if resource_id is not None: + await client.api.delete_resource(resource_id) + await client.close() + + resource = asyncio.run(exercise()) + recorded = oauth_environment.resource.state.requests[-2] + + assert resource.name == "openapi-httpx-resource" + assert recorded.path == "/api/resources?source=openapi-httpx" + assert recorded.body == b'{"name":"openapi-httpx-resource"}' + assert recorded.header("user-agent") == "caller-agent" + assert recorded.header("authorization") == "DPoP integration-access-token" + assert recorded.header("dpop") != "stale-proof" + + +def test_openapi_httpx_raw_method_returns_native_readable_response( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> tuple[httpx.Response, str]: + client = _client(oauth_environment) + response = await client.api.create_resource_without_preload_content(Resource(name="raw-response")) + assert response.is_stream_consumed + resource_id = response.json()["id"] + await client.api.delete_resource(resource_id) + await client.close() + return response, resource_id + + response, resource_id = asyncio.run(exercise()) + + assert isinstance(response, httpx.Response) + assert response.status_code == 201 + assert response.is_stream_consumed + assert resource_id == "1" + assert response.json() == {"id": "1", "name": "raw-response"} + + +def test_openapi_httpx_client_reuses_cached_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + for _ in range(2): + with pytest.raises(ApiException) as error: + await client.api.get_resource("missing") + assert error.value.status == 404 + finally: + await client.close() + + asyncio.run(exercise()) + + assert len(oauth_environment.authorization.state.requests) == 2 + + +def test_openapi_httpx_client_reacquires_invalidated_token(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise(config: OAuth2Config, store: InvalidatableTokenStore) -> None: + client = _client(oauth_environment, config=config) + try: + resource_id = await client.create_resource("expires") + store.invalidate() + assert await client.get_resource(resource_id) == resource_id + await client.delete_resource(resource_id) + finally: + await client.close() + + store = InvalidatableTokenStore() + config = oauth_environment.config("ec", token_store=store) + asyncio.run(exercise(config, store)) + + assert len(oauth_environment.authorization.state.requests) == 4 + + +def test_openapi_httpx_client_stops_after_one_token_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + await client.create_resource("never-created") + finally: + await client.close() + + oauth_environment.authorization.state.always_challenge_nonce = True + + with pytest.raises(ApiException) as error: + asyncio.run(exercise()) + + assert error.value.status == 400 + assert len(oauth_environment.authorization.state.requests) == 2 + assert oauth_environment.resource.state.requests == [] + + +def test_openapi_httpx_client_returns_configured_server_errors( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def create(name: str) -> None: + client = _client(oauth_environment) + try: + await client.create_resource(name) + finally: + await client.close() + + oauth_environment.authorization.state.reject_client_assertion = True + with pytest.raises(ApiException) as assertion_error: + asyncio.run(create("rejected")) + assert assertion_error.value.status == 400 + assert json.loads(assertion_error.value.body) == { + "error": "invalid_client", + "error_description": "Client assertion was rejected", + } + + oauth_environment.reset() + oauth_environment.resource.state.insufficient_scope = True + with pytest.raises(ApiException) as scope_error: + asyncio.run(create("forbidden")) + assert scope_error.value.status == 403 + assert json.loads(scope_error.value.body) == { + "error": "insufficient_scope", + "error_description": "requested scope is not permitted", + } + assert scope_error.value.headers["WWW-Authenticate"] == 'DPoP error="insufficient_scope"' + + +def test_openapi_httpx_client_stops_after_one_resource_nonce_retry( + oauth_environment: OAuthTestEnvironment, +) -> None: + async def exercise() -> None: + client = _client(oauth_environment) + try: + await client.create_resource("never-created") + finally: + await client.close() + + oauth_environment.resource.state.always_challenge_nonce = True + + with pytest.raises(ApiException) as error: + asyncio.run(exercise()) + + assert error.value.status == 401 + assert len(oauth_environment.resource.state.requests) == 2 + + +def test_openapi_httpx_client_close_closes_generated_pool(oauth_environment: OAuthTestEnvironment) -> None: + async def exercise() -> bool: + client = _client(oauth_environment) + pool = client.api_client.rest_client.pool_manager + if pool is None: + raise AssertionError("OAuth2 layer did not install generated HTTPX pool") + await client.close() + return bool(pool.is_closed) + + assert asyncio.run(exercise()) diff --git a/tests/resources/openapi/fake-api.yaml b/tests/resources/openapi/fake-api.yaml new file mode 100644 index 0000000..6544f7e --- /dev/null +++ b/tests/resources/openapi/fake-api.yaml @@ -0,0 +1,69 @@ +openapi: 3.0.3 +info: + title: OAuth2 integration test API + version: 1.0.0 +servers: + - url: https://localhost +paths: + /api/resources: + post: + operationId: createResource + tags: [Resources] + parameters: + - name: source + in: query + required: false + schema: + type: string + requestBody: + required: true + content: + application/json: + schema: + $ref: "#/components/schemas/Resource" + responses: + "201": + description: Resource created + content: + application/json: + schema: + $ref: "#/components/schemas/Resource" + /api/resources/{resource_id}: + parameters: + - name: resource_id + in: path + required: true + schema: + type: string + get: + operationId: getResource + tags: [Resources] + responses: + "200": + description: Resource returned + content: + application/json: + schema: + $ref: "#/components/schemas/Resource" + "404": + description: Resource not found + delete: + operationId: deleteResource + tags: [Resources] + responses: + "204": + description: Resource deleted +components: + schemas: + Resource: + type: object + additionalProperties: false + required: [name] + properties: + id: + type: string + readOnly: true + name: + type: string + minLength: 1 + maxLength: 200 \ No newline at end of file diff --git a/tests/unit/configuration.py b/tests/unit/configuration.py new file mode 100644 index 0000000..3c1934a --- /dev/null +++ b/tests/unit/configuration.py @@ -0,0 +1,17 @@ +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from mastercard_oauth2_client import KeyPair, OAuth2Config, StaticDPoPKeyProvider, StaticScopeResolver + + +def integration_config() -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client-id", + key_id="key-id", + token_endpoint="https://auth.example.com/token", + issuer="https://auth.example.com", + client_key=client_key, + scope_resolver=StaticScopeResolver({"read"}), + dpop_key_provider=StaticDPoPKeyProvider(KeyPair(dpop_key, dpop_key.public_key())), + ) \ No newline at end of file diff --git a/tests/unit/test_aiohttp_integration.py b/tests/unit/test_aiohttp_integration.py new file mode 100644 index 0000000..0bf7800 --- /dev/null +++ b/tests/unit/test_aiohttp_integration.py @@ -0,0 +1,238 @@ +from __future__ import annotations + +import asyncio +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Any, cast + +import aiohttp +import pytest +from multidict import CIMultiDict + +from mastercard_oauth2_client.handlers import AsyncOAuth2Handler +from mastercard_oauth2_client.integrations.aiohttp import OAuth2Middleware, _AiohttpAdapter +from mastercard_oauth2_client.models import HttpRequest +from tests.unit.configuration import integration_config + +pytestmark = pytest.mark.aiohttp + + +@dataclass +class FakeResponse: + status: int = 200 + headers: Mapping[str, str] = field(default_factory=dict) + body: bytes = b"response" + reads: int = 0 + releases: int = 0 + + async def read(self) -> bytes: + self.reads += 1 + return self.body + + def release(self) -> None: + self.releases += 1 + + +class FakeSession: + def __init__(self, response: FakeResponse) -> None: + self.response = response + self.calls: list[tuple[tuple[object, ...], dict[str, object]]] = [] + + async def request(self, *args: object, **kwargs: object) -> FakeResponse: + self.calls.append((args, kwargs)) + return self.response + + +@dataclass +class FakeRequest: + method: str + url: str + headers: CIMultiDict[str] + session: FakeSession + body: object = b"" + ssl: object = True + proxy: object = None + proxy_headers: object = None + chunked: bool | None = None + updated_bodies: list[object] = field(default_factory=list) + + async def update_body(self, body: object) -> None: + self.updated_bodies.append(body) + self.body = body + + +def test_aiohttp_middleware_constructs_async_handler() -> None: + middleware = OAuth2Middleware(integration_config()) + + assert isinstance(middleware._handler, AsyncOAuth2Handler) + + +def test_aiohttp_adapter_maps_metadata_and_token_request_options() -> None: + async def exercise() -> None: + token_response = FakeResponse(status=201) + expected_response = cast(aiohttp.ClientResponse, token_response) + session = FakeSession(token_response) + request = FakeRequest( + method="GET", + url="https://api.example.com/resources?limit=1", + headers=CIMultiDict({"Accept": "application/json"}), + session=session, + ssl="resource-ssl", + proxy="https://proxy.example.com", + proxy_headers={"Proxy-Authorization": "proxy-value"}, + ) + + async def resource_handler(resource_request: aiohttp.ClientRequest) -> aiohttp.ClientResponse: + raise AssertionError(f"Unexpected resource request: {resource_request!r}") + + adapter = _AiohttpAdapter(resource_handler) + native_request = cast(aiohttp.ClientRequest, request) + token_request = HttpRequest( + method="POST", + url="https://auth.example.com/oauth/token", + headers={"Content-Type": "application/x-www-form-urlencoded"}, + body="grant_type=client_credentials", + ) + + response = await adapter.send_token_request(native_request, token_request) + + assert response is expected_response + assert adapter.request_method(native_request) == "GET" + assert adapter.request_url(native_request) == "https://api.example.com/resources?limit=1" + assert adapter.request_headers(native_request) is request.headers + assert session.calls == [ + ( + ("POST", "https://auth.example.com/oauth/token"), + { + "headers": token_request.headers, + "data": "grant_type=client_credentials", + "ssl": "resource-ssl", + "proxy": "https://proxy.example.com", + "proxy_headers": {"Proxy-Authorization": "proxy-value"}, + "raise_for_status": False, + "middlewares": (), + }, + ) + ] + + asyncio.run(exercise()) + + +def test_aiohttp_adapter_replaces_headers_and_uses_native_resource_handler() -> None: + async def exercise() -> None: + response = FakeResponse() + expected_response = cast(aiohttp.ClientResponse, response) + request = FakeRequest( + method="POST", + url="https://api.example.com/resources", + headers=CIMultiDict( + (("Authorization", "Bearer stale"), ("Authorization", "Bearer duplicate"), ("X-Correlation-ID", "old")) + ), + session=FakeSession(response), + ) + seen_requests: list[aiohttp.ClientRequest] = [] + + async def resource_handler(resource_request: aiohttp.ClientRequest) -> aiohttp.ClientResponse: + seen_requests.append(resource_request) + return cast(aiohttp.ClientResponse, response) + + adapter = _AiohttpAdapter(resource_handler) + native_request = cast(aiohttp.ClientRequest, request) + result = await adapter.send_resource_request( + native_request, + {"Authorization": "DPoP access-token", "DPoP": "fresh-proof", "X-Correlation-ID": "new"}, + ) + + assert result is expected_response + assert seen_requests == [native_request] + assert list(request.headers.items()) == [ + ("Authorization", "DPoP access-token"), + ("DPoP", "fresh-proof"), + ("X-Correlation-ID", "new"), + ] + + asyncio.run(exercise()) + + +def test_aiohttp_adapter_reads_and_releases_internal_responses() -> None: + async def exercise() -> None: + response = FakeResponse(status=401, headers={"DPoP-Nonce": "nonce"}, body=b"challenge") + + async def resource_handler(resource_request: aiohttp.ClientRequest) -> aiohttp.ClientResponse: + raise AssertionError(f"Unexpected resource request: {resource_request!r}") + + adapter = _AiohttpAdapter(resource_handler) + native_response = cast(aiohttp.ClientResponse, response) + + assert adapter.response_status(native_response) == 401 + assert adapter.response_headers(native_response) == {"DPoP-Nonce": "nonce"} + assert await adapter.response_body(native_response) == b"challenge" + await adapter.close_response(native_response) + assert response.reads == 2 + assert response.releases == 1 + + asyncio.run(exercise()) + + +def test_aiohttp_middleware_buffers_payload_before_delegating() -> None: + async def exercise() -> None: + response = FakeResponse() + expected_response = cast(aiohttp.ClientResponse, response) + request = FakeRequest( + method="POST", + url="https://api.example.com/resources", + headers=CIMultiDict({"Content-Type": "application/octet-stream"}), + session=FakeSession(response), + body=aiohttp.payload.BytesPayload(b"request-body"), + chunked=True, + ) + seen_adapters: list[object] = [] + + class FakeHandler: + async def execute(self, resource_request: object, adapter: object) -> aiohttp.ClientResponse: + assert resource_request is request + seen_adapters.append(adapter) + return cast(aiohttp.ClientResponse, response) + + async def resource_handler(resource_request: aiohttp.ClientRequest) -> aiohttp.ClientResponse: + raise AssertionError(f"Unexpected resource request: {resource_request!r}") + + middleware = OAuth2Middleware.__new__(OAuth2Middleware) + middleware._handler = cast(Any, FakeHandler()) + result = await middleware(cast(aiohttp.ClientRequest, request), resource_handler) + + assert result is expected_response + assert request.updated_bodies == [b"request-body"] + assert request.chunked is None + assert len(seen_adapters) == 1 + assert isinstance(seen_adapters[0], _AiohttpAdapter) + + asyncio.run(exercise()) + + +def test_aiohttp_middleware_leaves_empty_body_unchanged() -> None: + async def exercise() -> None: + response = FakeResponse() + expected_response = cast(aiohttp.ClientResponse, response) + request = FakeRequest( + method="GET", + url="https://api.example.com/resources", + headers=CIMultiDict(), + session=FakeSession(response), + ) + + class FakeHandler: + async def execute(self, resource_request: object, adapter: object) -> aiohttp.ClientResponse: + return cast(aiohttp.ClientResponse, response) + + async def resource_handler(resource_request: aiohttp.ClientRequest) -> aiohttp.ClientResponse: + raise AssertionError(f"Unexpected resource request: {resource_request!r}") + + middleware = OAuth2Middleware.__new__(OAuth2Middleware) + middleware._handler = cast(Any, FakeHandler()) + result = await middleware(cast(aiohttp.ClientRequest, request), resource_handler) + + assert result is expected_response + assert request.updated_bodies == [] + + asyncio.run(exercise()) diff --git a/tests/unit/test_async_handler.py b/tests/unit/test_async_handler.py new file mode 100644 index 0000000..fc2be8b --- /dev/null +++ b/tests/unit/test_async_handler.py @@ -0,0 +1,417 @@ +import asyncio +import base64 +import json +from collections.abc import Iterable, Mapping +from collections.abc import Set as AbstractSet +from dataclasses import dataclass +from datetime import UTC, datetime +from urllib.parse import parse_qs + +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +import mastercard_oauth2_client.handlers.async_ as async_handler_module +from mastercard_oauth2_client import ( + AccessToken, + AsyncOAuth2Handler, + HttpRequest, + HttpResponse, + InMemoryTokenStore, + KeyPair, + OAuth2Config, + OAuth2Error, + StaticDPoPKeyProvider, + StaticScopeResolver, + parse_token_response, +) +from mastercard_oauth2_client._internal.jose import jwk_thumbprint + + +@dataclass(frozen=True, slots=True) +class NativeRequest: + method: str + url: str + headers: Mapping[str, str] + + +@dataclass(slots=True) +class NativeResponse: + status: int + headers: Mapping[str, str] + body: bytes | str | None = None + body_reads: int = 0 + closed: bool = False + + +class AsyncScriptedAdapter: + def __init__(self, responses: Iterable[NativeResponse]) -> None: + self._responses = iter(responses) + self.requests: list[HttpRequest] = [] + self.token_requests = 0 + + def request_method(self, request: NativeRequest) -> str: + return request.method + + def request_url(self, request: NativeRequest) -> str: + return request.url + + def request_headers(self, request: NativeRequest) -> Mapping[str, str]: + return request.headers + + async def send_token_request(self, resource_request: NativeRequest, request: HttpRequest) -> NativeResponse: + self.token_requests += 1 + self.requests.append(request) + await asyncio.sleep(0) + return next(self._responses) + + async def send_resource_request(self, request: NativeRequest, headers: Mapping[str, str]) -> NativeResponse: + self.requests.append(HttpRequest(method=request.method, url=request.url, headers=dict(headers))) + await asyncio.sleep(0) + return next(self._responses) + + def response_status(self, response: NativeResponse) -> int: + return response.status + + def response_headers(self, response: NativeResponse) -> Mapping[str, str]: + return response.headers + + async def response_body(self, response: NativeResponse) -> bytes | str | None: + response.body_reads += 1 + await asyncio.sleep(0) + return response.body + + async def close_response(self, response: NativeResponse) -> None: + response.closed = True + await asyncio.sleep(0) + + +class FailingAsyncAdapter(AsyncScriptedAdapter): + def __init__(self) -> None: + super().__init__([]) + + async def send_token_request(self, resource_request: NativeRequest, request: HttpRequest) -> NativeResponse: + raise ConnectionError("network unavailable") + + +def _config(token_store: InMemoryTokenStore | None = None) -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_private_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client-id", + token_endpoint="https://auth.example.com/oauth/token", + issuer="https://auth.example.com", + client_key=client_key, + key_id="client-key-id", + scope_resolver=StaticScopeResolver({"payments:read"}), + dpop_key_provider=StaticDPoPKeyProvider( + KeyPair(private_key=dpop_private_key, public_key=dpop_private_key.public_key()) + ), + token_store=token_store or InMemoryTokenStore(), + user_agent="library-agent", + ) + + +def _request(url: str = "https://api.example.com/payments") -> NativeRequest: + return NativeRequest(method="GET", url=url, headers={"Accept": "application/json"}) + + +def _token_response(*, nonce: str | None = None) -> NativeResponse: + headers = {"Content-Type": "application/json"} + if nonce is not None: + headers["DPoP-Nonce"] = nonce + return NativeResponse( + status=200, + headers=headers, + body=json.dumps( + { + "access_token": "access-token", + "token_type": "DPoP", + "expires_in": 900, + "scope": "payments:read", + } + ), + ) + + +def _challenge(nonce: str, *, status: int = 401, header: bool = True) -> NativeResponse: + headers = {"DPoP-Nonce": nonce} + if header: + headers["WWW-Authenticate"] = 'DPoP error="use_dpop_nonce"' + return NativeResponse( + status=status, + headers=headers, + body='{"error":"use_dpop_nonce","secret":"response-secret"}', + ) + + +def _claims(compact_jwt: str) -> dict[str, object]: + payload = compact_jwt.split(".")[1] + return json.loads(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4))) # type: ignore[no-any-return] + + +def test_async_handler_acquires_token_and_preserves_native_response() -> None: + async def exercise() -> None: + config = _config() + token_response = _token_response() + final_response = NativeResponse(status=200, headers={}, body="resource") + adapter = AsyncScriptedAdapter([token_response, final_response]) + + response = await AsyncOAuth2Handler[NativeRequest, NativeResponse](config).execute(_request(), adapter) + + assert response is final_response + assert token_response.closed + assert token_response.body_reads == 1 + assert not final_response.closed + assert final_response.body_reads == 0 + assert adapter.requests[1].headers["Authorization"] == "DPoP access-token" + + asyncio.run(exercise()) + + +def test_async_handler_closes_successful_token_response_before_parsing(monkeypatch: pytest.MonkeyPatch) -> None: + async def exercise() -> None: + config = _config() + token_response = _token_response() + adapter = AsyncScriptedAdapter([token_response, NativeResponse(status=200, headers={})]) + + def parse_after_close( + parse_config: OAuth2Config, + response: HttpResponse, + requested_scopes: AbstractSet[str], + dpop_key_id: str, + ) -> AccessToken: + assert token_response.closed + return parse_token_response(parse_config, response, requested_scopes, dpop_key_id) + + monkeypatch.setitem(async_handler_module.__dict__, "parse_token_response", parse_after_close) + + await AsyncOAuth2Handler[NativeRequest, NativeResponse](config).execute(_request(), adapter) + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("status", [400, 401]) +def test_async_handler_retries_token_nonce_once(status: int) -> None: + async def exercise() -> None: + first = _challenge("token-nonce", status=status, header=False) + adapter = AsyncScriptedAdapter([first, _token_response(), NativeResponse(status=200, headers={})]) + + await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert first.closed + assert first.body_reads == 1 + assert len(adapter.requests) == 3 + assert _claims(adapter.requests[1].headers["DPoP"])["nonce"] == "token-nonce" + assertions = [parse_qs(str(request.body))["client_assertion"][0] for request in adapter.requests[:2]] + assert assertions[0] != assertions[1] + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("status", [400, 401]) +def test_async_handler_retries_resource_nonce_once(status: int) -> None: + async def exercise() -> None: + first = _challenge("resource-nonce", status=status) + final_response = NativeResponse(status=200, headers={}) + adapter = AsyncScriptedAdapter([_token_response(), first, final_response]) + + response = await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert response is final_response + assert first.closed + assert first.body_reads == 0 + assert _claims(adapter.requests[2].headers["DPoP"])["nonce"] == "resource-nonce" + + asyncio.run(exercise()) + + +def test_async_handler_retries_body_only_resource_nonce() -> None: + async def exercise() -> None: + first = _challenge("resource-nonce", header=False) + adapter = AsyncScriptedAdapter([_token_response(), first, NativeResponse(status=200, headers={})]) + + await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert first.body_reads == 1 + assert first.closed + assert _claims(adapter.requests[2].headers["DPoP"])["nonce"] == "resource-nonce" + + asyncio.run(exercise()) + + +def test_async_handler_shares_token_response_nonce_with_resource_proof() -> None: + async def exercise() -> None: + adapter = AsyncScriptedAdapter( + [_token_response(nonce="shared-nonce"), NativeResponse(status=200, headers={})] + ) + + await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert _claims(adapter.requests[1].headers["DPoP"])["nonce"] == "shared-nonce" + + asyncio.run(exercise()) + + +def test_async_handler_stops_after_one_token_nonce_retry() -> None: + async def exercise() -> None: + first = _challenge("first") + second = _challenge("second") + adapter = AsyncScriptedAdapter([first, second]) + + response = await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert response is second + assert first.closed + assert not second.closed + assert adapter.token_requests == 2 + assert len(adapter.requests) == 2 + + asyncio.run(exercise()) + + +def test_async_handler_returns_terminal_token_error_open() -> None: + async def exercise() -> None: + token_error = NativeResponse(status=403, headers={}, body="denied") + adapter = AsyncScriptedAdapter([token_error]) + + response = await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert response is token_error + assert not response.closed + assert response.body_reads == 0 + assert len(adapter.requests) == 1 + + asyncio.run(exercise()) + + +def test_async_handler_stops_after_one_resource_nonce_retry() -> None: + async def exercise() -> None: + first = _challenge("first") + second = _challenge("second") + adapter = AsyncScriptedAdapter([_token_response(), first, second]) + + response = await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute(_request(), adapter) + + assert response is second + assert first.closed + assert not second.closed + assert len(adapter.requests) == 3 + + asyncio.run(exercise()) + + +def test_async_handler_reuses_cached_token() -> None: + async def exercise() -> None: + config = _config() + dpop_key = config.dpop_key_provider.get_current_key() + config.token_store.put( + AccessToken( + client_id=config.client_id, + token_value="cached-token", + scopes=frozenset({"payments:read"}), + expires_at=datetime(2100, 1, 1, tzinfo=UTC), + jkt=jwk_thumbprint(dpop_key.key_pair.public_key), + ) + ) + adapter = AsyncScriptedAdapter([NativeResponse(status=200, headers={})]) + + await AsyncOAuth2Handler[NativeRequest, NativeResponse](config).execute(_request(), adapter) + + assert adapter.token_requests == 0 + assert adapter.requests[0].headers["Authorization"] == "DPoP cached-token" + + asyncio.run(exercise()) + + +def test_async_handler_preserves_caller_headers_and_replaces_oauth_headers() -> None: + async def exercise() -> None: + config = _config() + dpop_key = config.dpop_key_provider.get_current_key() + config.token_store.put( + AccessToken( + client_id=config.client_id, + token_value="cached-token", + scopes=frozenset({"payments:read"}), + expires_at=datetime(2100, 1, 1, tzinfo=UTC), + jkt=jwk_thumbprint(dpop_key.key_pair.public_key), + ) + ) + request = NativeRequest( + method="GET", + url="https://api.example.com/payments", + headers={ + "Accept": "application/json", + "User-Agent": "caller-agent", + "Authorization": "Bearer stale", + "DPoP": "stale-proof", + }, + ) + adapter = AsyncScriptedAdapter([NativeResponse(status=200, headers={})]) + + await AsyncOAuth2Handler[NativeRequest, NativeResponse](config).execute(request, adapter) + + headers = adapter.requests[0].headers + assert headers["Accept"] == "application/json" + assert headers["User-Agent"] == "caller-agent" + assert headers["Authorization"] == "DPoP cached-token" + assert headers["DPoP"] != "stale-proof" + + asyncio.run(exercise()) + + +def test_async_handler_preserves_transport_errors() -> None: + async def exercise() -> None: + with pytest.raises(ConnectionError, match="network unavailable"): + await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute( + _request(), + FailingAsyncAdapter(), + ) + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("resource_url", ["http://api.example.com/payments", "/payments", "not-a-url"]) +def test_async_handler_rejects_non_https_resource_before_transport(resource_url: str) -> None: + async def exercise() -> None: + adapter = AsyncScriptedAdapter([]) + + with pytest.raises(OAuth2Error, match="FAPI 2.0 requires HTTPS for resource server"): + await AsyncOAuth2Handler[NativeRequest, NativeResponse](_config()).execute( + _request(resource_url), + adapter, + ) + + assert adapter.requests == [] + + asyncio.run(exercise()) + + +def test_async_handler_logs_lifecycle_without_sensitive_values(caplog: pytest.LogCaptureFixture) -> None: + async def exercise() -> None: + config = _config() + adapter = AsyncScriptedAdapter( + [_challenge("super-secret-nonce"), _token_response(), NativeResponse(status=200, headers={})] + ) + + with caplog.at_level("DEBUG", logger="mastercard_oauth2_client.handlers.async_"): + await AsyncOAuth2Handler[NativeRequest, NativeResponse](config).execute( + _request("https://api.example.com/payments?account=secret#fragment"), + adapter, + ) + + logs = caplog.text + assert "GET https://api.example.com/payments" in logs + assert "Retrieving DPoP key" in logs + assert "DPoP key selected" in logs + assert "cache miss" in logs + assert "nonce-directed retry" in logs + for sensitive_value in ( + "access-token", + "super-secret-nonce", + "response-secret", + "account=secret", + config.dpop_key_provider.get_current_key().key_id, + ): + assert sensitive_value not in logs + + asyncio.run(exercise()) diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py new file mode 100644 index 0000000..b5e4e2f --- /dev/null +++ b/tests/unit/test_config.py @@ -0,0 +1,255 @@ +import platform +from dataclasses import FrozenInstanceError, replace +from importlib.metadata import PackageNotFoundError +from typing import Any, TypedDict, cast + +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +import mastercard_oauth2_client.config as config_module +from mastercard_oauth2_client import ( + InMemoryTokenStore, + KeyPair, + OAuth2Config, + OAuth2ConfigError, + OAuth2Error, + SecurityProfile, + StaticDPoPKeyProvider, + StaticScopeResolver, +) + + +class ConfigValues(TypedDict): + client_id: str + token_endpoint: str + issuer: str + client_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey + key_id: str + scope_resolver: StaticScopeResolver + dpop_key_provider: StaticDPoPKeyProvider + + +def minimum_config_values() -> ConfigValues: + dpop_private_key = ec.generate_private_key(ec.SECP256R1()) + return { + "client_id": "client-id", + "token_endpoint": "https://sandbox.api.mastercard.com/oauth/token", + "issuer": "https://sandbox.api.mastercard.com", + "client_key": rsa.generate_private_key(public_exponent=65537, key_size=2048), + "key_id": "key-id", + "scope_resolver": StaticScopeResolver({"payments:read"}), + "dpop_key_provider": StaticDPoPKeyProvider( + KeyPair( + private_key=dpop_private_key, + public_key=dpop_private_key.public_key(), + ) + ), + } + + +def test_oauth2_config_accepts_minimum_valid_fapi2_configuration() -> None: + values = minimum_config_values() + + config = OAuth2Config(**values) + + assert config.client_id == "client-id" + assert config.security_profile is SecurityProfile.FAPI2_PRIVATE_KEY_DPOP + assert config.clock_skew_tolerance == 5 + assert isinstance(config.token_store, InMemoryTokenStore) + assert config.user_agent.startswith("Mastercard-OAuth2-Client/") + assert "(Python/" in config.user_agent + assert isinstance(config.client_key, rsa.RSAPrivateKey) + + +def test_oauth2_config_trims_client_and_key_identifiers() -> None: + values = minimum_config_values() + values["client_id"] = " client-id " + values["key_id"] = " key-id " + + config = OAuth2Config(**values) + + assert config.client_id == "client-id" + assert config.key_id == "key-id" + + +def test_default_user_agent_has_exact_runtime_format(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(config_module, "version", lambda distribution: "1.2.3") + monkeypatch.setattr(platform, "python_version", lambda: "3.13.2") + monkeypatch.setattr(platform, "system", lambda: "TestOS") + monkeypatch.setattr(platform, "release", lambda: "9") + + assert config_module.default_user_agent() == "Mastercard-OAuth2-Client/1.2.3 (Python/3.13.2; TestOS 9)" + + +def test_default_user_agent_uses_unknown_version_when_distribution_is_missing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + def missing_distribution(distribution: str) -> str: + raise PackageNotFoundError(distribution) + + monkeypatch.setattr(config_module, "version", missing_distribution) + monkeypatch.setattr(platform, "python_version", lambda: "3.13.2") + monkeypatch.setattr(platform, "system", lambda: "TestOS") + monkeypatch.setattr(platform, "release", lambda: "9") + + assert config_module.default_user_agent() == ( + "Mastercard-OAuth2-Client/0.0.0-unknown (Python/3.13.2; TestOS 9)" + ) + + +def test_oauth2_config_uses_distinct_default_token_stores() -> None: + first = OAuth2Config(**minimum_config_values()) + second = OAuth2Config(**minimum_config_values()) + + assert first.token_store is not second.token_store + + +def test_oauth2_config_accepts_replaceable_defaults() -> None: + custom_store = InMemoryTokenStore() + + config = OAuth2Config( + **minimum_config_values(), + token_store=custom_store, + user_agent="customer-agent/1.0", + ) + + assert config.token_store is custom_store + assert config.user_agent == "customer-agent/1.0" + + +def test_oauth2_config_uses_generated_user_agent_for_blank_override() -> None: + config = OAuth2Config(**minimum_config_values(), user_agent=" ") + + assert config.user_agent.startswith("Mastercard-OAuth2-Client/") + + +def test_oauth2_config_is_immutable() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(FrozenInstanceError): + config.client_id = "different-client" # type: ignore[misc] + + +def test_oauth2_config_rejects_missing_client_id() -> None: + values = minimum_config_values() + values["client_id"] = " " + + with pytest.raises(OAuth2ConfigError, match="Client ID is required"): + OAuth2Config(**values) + + +def test_oauth2_config_rejects_missing_token_endpoint() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Token endpoint is required"): + replace(config, token_endpoint=" ") + + +@pytest.mark.parametrize("token_endpoint", ["http://auth.example.com/token", "/oauth/token", "not-a-url"]) +def test_oauth2_config_requires_absolute_https_token_endpoint(token_endpoint: str) -> None: + values = minimum_config_values() + values["token_endpoint"] = token_endpoint + + with pytest.raises(OAuth2ConfigError, match="FAPI 2.0 requires HTTPS token endpoint"): + OAuth2Config(**values) + + +def test_oauth2_config_rejects_missing_issuer() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Issuer is required"): + replace(config, issuer=" ") + + +def test_oauth2_config_rejects_missing_client_key() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Client private key is required"): + replace(config, client_key=None) # type: ignore[arg-type] + + +def test_oauth2_config_rejects_missing_key_id() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match=r"Key ID \(kid\) is required"): + replace(config, key_id=" ") + + +def test_oauth2_config_rejects_missing_scope_resolver() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Scope resolver is required"): + replace(config, scope_resolver=None) # type: ignore[arg-type] + + +def test_oauth2_config_rejects_missing_dpop_key_provider() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="DPoP key provider is required"): + replace(config, dpop_key_provider=None) # type: ignore[arg-type] + + +def test_oauth2_config_rejects_missing_token_store() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Token store is required"): + replace(config, token_store=None) # type: ignore[arg-type] + + +def test_oauth2_config_rejects_unsupported_security_profile() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Security profile must be FAPI 2.0"): + replace(config, security_profile=object()) # type: ignore[arg-type] + + +def test_oauth2_config_rejects_explicitly_missing_security_profile() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Security profile must be FAPI 2.0"): + replace(config, security_profile=None) # type: ignore[arg-type] + + +def test_oauth2_config_rejects_negative_clock_skew() -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Clock skew tolerance must not be negative"): + replace(config, clock_skew_tolerance=-1) + + +@pytest.mark.parametrize( + ("field_name", "value", "message"), + [ + pytest.param("client_id", None, "Client ID is required", id="client-id-none"), + pytest.param("client_id", 123, "Client ID is required", id="client-id-number"), + pytest.param("token_endpoint", None, "Token endpoint is required", id="token-endpoint-none"), + pytest.param("token_endpoint", 123, "Token endpoint is required", id="token-endpoint-number"), + pytest.param("issuer", None, "Issuer is required", id="issuer-none"), + pytest.param("issuer", 123, "Issuer is required", id="issuer-number"), + pytest.param("key_id", None, r"Key ID \(kid\) is required", id="key-id-none"), + pytest.param("key_id", 123, r"Key ID \(kid\) is required", id="key-id-number"), + pytest.param("user_agent", None, "User agent must be a string", id="user-agent-none"), + pytest.param("user_agent", 123, "User agent must be a string", id="user-agent-number"), + pytest.param("clock_skew_tolerance", None, "Clock skew tolerance must be an integer", id="skew-none"), + pytest.param("clock_skew_tolerance", "5", "Clock skew tolerance must be an integer", id="skew-string"), + pytest.param("clock_skew_tolerance", True, "Clock skew tolerance must be an integer", id="skew-bool"), + ], +) +def test_oauth2_config_rejects_invalid_runtime_types(field_name: str, value: object, message: str) -> None: + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match=message): + cast(Any, replace)(config, **{field_name: value}) + + +def test_oauth2_config_error_is_an_oauth2_error() -> None: + assert issubclass(OAuth2ConfigError, OAuth2Error) + + +def test_oauth2_config_accepts_ec_client_key() -> None: + values = minimum_config_values() + values["client_key"] = ec.generate_private_key(ec.SECP256R1()) + + config = OAuth2Config(**values) + + assert isinstance(config.client_key, ec.EllipticCurvePrivateKey) diff --git a/tests/unit/test_handler.py b/tests/unit/test_handler.py new file mode 100644 index 0000000..8d69468 --- /dev/null +++ b/tests/unit/test_handler.py @@ -0,0 +1,376 @@ +import base64 +import json +from collections.abc import Iterable, Mapping +from collections.abc import Set as AbstractSet +from dataclasses import dataclass +from datetime import UTC, datetime +from threading import Lock +from urllib.parse import parse_qs + +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +import mastercard_oauth2_client.handlers.sync as handler_module +from mastercard_oauth2_client import ( + AccessToken, + AccessTokenFilter, + HttpRequest, + HttpResponse, + InMemoryTokenStore, + KeyPair, + OAuth2Config, + OAuth2Error, + OAuth2Handler, + StaticDPoPKeyProvider, + StaticScopeResolver, + parse_token_response, +) +from mastercard_oauth2_client._internal.jose import jwk_thumbprint + + +def _decode_claims(token: str) -> dict[str, object]: + payload = token.split(".")[1] + return json.loads(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4))) # type: ignore[no-any-return] + + +@dataclass(frozen=True, slots=True) +class NativeRequest: + method: str + url: str + headers: Mapping[str, str] + body: bytes | str | None = None + + +@dataclass(slots=True) +class NativeResponse: + status: int + headers: Mapping[str, str] + body: bytes | str | None = None + body_reads: int = 0 + closed: bool = False + + +class ScriptedAdapter: + def __init__(self, responses: Iterable[NativeResponse]) -> None: + self._responses = iter(responses) + self.requests: list[HttpRequest] = [] + self.closed_responses: list[NativeResponse] = [] + self.token_requests = 0 + self._response_lock = Lock() + + def request_method(self, request: NativeRequest) -> str: + return request.method + + def request_url(self, request: NativeRequest) -> str: + return request.url + + def request_headers(self, request: NativeRequest) -> Mapping[str, str]: + return request.headers + + def send_token_request(self, resource_request: NativeRequest, request: HttpRequest) -> NativeResponse: + self.token_requests += 1 + self.requests.append(request) + with self._response_lock: + return next(self._responses) + + def send_resource_request(self, request: NativeRequest, headers: Mapping[str, str]) -> NativeResponse: + self.requests.append(HttpRequest(method=request.method, url=request.url, headers=dict(headers), body=request.body)) + with self._response_lock: + return next(self._responses) + + def response_status(self, response: NativeResponse) -> int: + return response.status + + def response_headers(self, response: NativeResponse) -> Mapping[str, str]: + return response.headers + + def response_body(self, response: NativeResponse) -> bytes | str | None: + response.body_reads += 1 + return response.body + + def close_response(self, response: NativeResponse) -> None: + response.closed = True + self.closed_responses.append(response) + + +class FailingAdapter(ScriptedAdapter): + def __init__(self) -> None: + super().__init__([]) + + def send_token_request(self, resource_request: NativeRequest, request: HttpRequest) -> NativeResponse: + raise ConnectionError("network unavailable") + + +def _config(token_store: InMemoryTokenStore | None = None) -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_private_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client-id", + token_endpoint="https://auth.example.com/oauth/token?tenant=secret", + issuer="https://auth.example.com", + client_key=client_key, + key_id="client-key-id", + scope_resolver=StaticScopeResolver({"payments:read"}), + dpop_key_provider=StaticDPoPKeyProvider( + KeyPair(private_key=dpop_private_key, public_key=dpop_private_key.public_key()) + ), + token_store=token_store or InMemoryTokenStore(), + user_agent="library-agent", + ) + + +def _resource_request() -> NativeRequest: + return NativeRequest( + method="GET", + url="https://api.example.com/payments?account=secret#fragment", + headers={"Accept": "application/json", "User-Agent": "caller-agent", "Authorization": "Bearer stale"}, + ) + + +def _token_response(*, nonce: str | None = None) -> NativeResponse: + headers = {} if nonce is None else {"DPoP-Nonce": nonce} + return NativeResponse( + status=200, + headers=headers, + body=json.dumps( + { + "access_token": "new-access-token", + "token_type": "DPoP", + "expires_in": 900, + "scope": "payments:read", + } + ), + ) + + +def _nonce_challenge( + nonce: str, + *, + status: int = 401, + body_error: bool = True, + header_error: bool = True, +) -> NativeResponse: + body = json.dumps({"error": "use_dpop_nonce", "secret": "response-secret"}) if body_error else None + headers = {"dpop-nonce": nonce} + if header_error: + headers["WWW-Authenticate"] = 'DPoP error="use_dpop_nonce"' + return NativeResponse( + status=status, + headers=headers, + body=body, + ) + + +def test_handler_reuses_matching_cached_token() -> None: + config = _config() + dpop_key = config.dpop_key_provider.get_current_key() + config.token_store.put( + AccessToken( + client_id=config.client_id, + token_value="cached-token", + scopes=frozenset({"payments:read"}), + expires_at=datetime(2100, 1, 1, tzinfo=UTC), + jkt=jwk_thumbprint(dpop_key.key_pair.public_key), + ) + ) + expected_response = NativeResponse(status=200, headers={}, body="ok") + transport = ScriptedAdapter([expected_response]) + + response = OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert response is expected_response + assert response.status == 200 + assert response.body_reads == 0 + assert len(transport.requests) == 1 + sent = transport.requests[0] + assert sent.url == _resource_request().url + assert sent.headers["Authorization"] == "DPoP cached-token" + assert sent.headers["User-Agent"] == "caller-agent" + assert len(sent.headers["DPoP"].split(".")) == 3 + + +def test_handler_acquires_stores_and_uses_token_on_cache_miss() -> None: + config = _config() + token_response = _token_response() + transport = ScriptedAdapter([token_response, NativeResponse(status=204, headers={})]) + + response = OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert response.status == 204 + assert [request.url for request in transport.requests] == [config.token_endpoint, _resource_request().url] + assert transport.requests[0].method == "POST" + assert transport.requests[1].headers["Authorization"] == "DPoP new-access-token" + assert token_response.closed + dpop_key = config.dpop_key_provider.get_current_key() + cached = config.token_store.get( + AccessTokenFilter( + scopes=frozenset({"payments:read"}), + jkt=jwk_thumbprint(dpop_key.key_pair.public_key), + ) + ) + assert cached is not None + assert cached.token_value == "new-access-token" + + +def test_handler_closes_successful_token_response_before_parsing(monkeypatch: pytest.MonkeyPatch) -> None: + config = _config() + token_response = _token_response() + transport = ScriptedAdapter([token_response, NativeResponse(status=200, headers={})]) + + def parse_after_close( + parse_config: OAuth2Config, + response: HttpResponse, + requested_scopes: AbstractSet[str], + dpop_key_id: str, + ) -> AccessToken: + assert token_response.closed + return parse_token_response(parse_config, response, requested_scopes, dpop_key_id) + + monkeypatch.setitem(handler_module.__dict__, "parse_token_response", parse_after_close) + + OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + +@pytest.mark.parametrize("challenge_status", [400, 401]) +def test_handler_retries_token_request_once_with_server_nonce(challenge_status: int) -> None: + config = _config() + transport = ScriptedAdapter( + [ + _nonce_challenge("token-nonce", status=challenge_status, header_error=False), + _token_response(), + NativeResponse(status=200, headers={}, body="ok"), + ] + ) + + OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert len(transport.requests) == 3 + first_proof = transport.requests[0].headers["DPoP"] + second_proof = transport.requests[1].headers["DPoP"] + assert first_proof != second_proof + assert "nonce" not in _decode_claims(first_proof) + assert _decode_claims(second_proof)["nonce"] == "token-nonce" + assertions = [parse_qs(str(request.body))["client_assertion"][0] for request in transport.requests[:2]] + assert assertions[0] != assertions[1] + + +@pytest.mark.parametrize("challenge_status", [400, 401]) +def test_handler_retries_resource_request_once_with_server_nonce(challenge_status: int) -> None: + config = _config() + first_challenge = _nonce_challenge("resource-nonce", status=challenge_status, body_error=False) + transport = ScriptedAdapter( + [ + _token_response(), + first_challenge, + NativeResponse(status=200, headers={}, body="ok"), + ] + ) + + response = OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert response.status == 200 + assert len(transport.requests) == 3 + first_proof = transport.requests[1].headers["DPoP"] + second_proof = transport.requests[2].headers["DPoP"] + assert first_proof != second_proof + assert "nonce" not in _decode_claims(first_proof) + assert _decode_claims(second_proof)["nonce"] == "resource-nonce" + assert first_challenge.closed + + +def test_handler_retries_body_only_resource_nonce() -> None: + config = _config() + first_challenge = _nonce_challenge("resource-nonce", header_error=False) + transport = ScriptedAdapter( + [_token_response(), first_challenge, NativeResponse(status=200, headers={}, body="ok")] + ) + + OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert first_challenge.body_reads == 1 + assert first_challenge.closed + assert _decode_claims(transport.requests[2].headers["DPoP"])["nonce"] == "resource-nonce" + + +def test_handler_shares_nonce_from_token_response_with_resource_proof() -> None: + config = _config() + transport = ScriptedAdapter( + [_token_response(nonce="shared-nonce"), NativeResponse(status=200, headers={}, body="ok")] + ) + + OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + resource_proof = transport.requests[1].headers["DPoP"] + assert _decode_claims(resource_proof)["nonce"] == "shared-nonce" + + +def test_handler_stops_after_one_resource_nonce_retry() -> None: + config = _config() + transport = ScriptedAdapter([_token_response(), _nonce_challenge("first"), _nonce_challenge("second")]) + + response = OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert response.status == 401 + assert len(transport.requests) == 3 + + +def test_handler_returns_second_token_challenge_without_sending_resource_request() -> None: + config = _config() + first_challenge = _nonce_challenge("first") + second_challenge = _nonce_challenge("second") + transport = ScriptedAdapter([first_challenge, second_challenge]) + + response = OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + assert response is second_challenge + assert response.status == 401 + assert first_challenge.closed + assert not second_challenge.closed + assert len(transport.requests) == 2 + assert all(request.url == config.token_endpoint for request in transport.requests) + + +def test_handler_preserves_transport_errors() -> None: + config = _config() + + with pytest.raises(ConnectionError, match="network unavailable"): + OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), FailingAdapter()) + + +@pytest.mark.parametrize("resource_url", ["http://api.example.com/payments", "/payments", "not-a-url"]) +def test_handler_requires_absolute_https_resource_url(resource_url: str) -> None: + config = _config() + transport = ScriptedAdapter([]) + request = NativeRequest(method="GET", url=resource_url, headers={}) + + with pytest.raises(OAuth2Error, match="FAPI 2.0 requires HTTPS for resource server"): + OAuth2Handler[NativeRequest, NativeResponse](config).execute(request, transport) + + assert transport.requests == [] + + +def test_handler_logs_lifecycle_without_sensitive_values(caplog: pytest.LogCaptureFixture) -> None: + config = _config() + transport = ScriptedAdapter( + [_nonce_challenge("super-secret-nonce"), _token_response(), NativeResponse(status=200, headers={})] + ) + + with caplog.at_level("DEBUG", logger="mastercard_oauth2_client.handlers.sync"): + OAuth2Handler[NativeRequest, NativeResponse](config).execute(_resource_request(), transport) + + logs = caplog.text + assert "GET https://api.example.com/payments" in logs + assert "Retrieving DPoP key" in logs + assert "DPoP key selected" in logs + assert "cache miss" in logs + assert "nonce-directed retry" in logs + for sensitive_value in ( + "new-access-token", + "super-secret-nonce", + "response-secret", + "account=secret", + "tenant=secret", + "Bearer stale", + config.dpop_key_provider.get_current_key().key_id, + ): + assert sensitive_value not in logs diff --git a/tests/unit/test_httpx_async_integration.py b/tests/unit/test_httpx_async_integration.py new file mode 100644 index 0000000..020fd13 --- /dev/null +++ b/tests/unit/test_httpx_async_integration.py @@ -0,0 +1,91 @@ +import asyncio + +import httpx +import pytest + +from mastercard_oauth2_client.integrations.httpx import AsyncOAuth2Transport +from tests.unit.configuration import integration_config + +pytestmark = pytest.mark.httpx + + +def test_async_httpx_transport_buffers_streaming_body_for_nonce_retry() -> None: + async def exercise() -> None: + body = b'{"name":"streamed-resource"}' + challenge_responses: list[httpx.Response] = [] + resource_bodies: list[bytes] = [] + + async def dispatch(request: httpx.Request) -> httpx.Response: + if request.url.host == "auth.example.com": + return httpx.Response( + 200, + json={"access_token": "token", "token_type": "DPoP", "expires_in": 900, "scope": "read"}, + request=request, + ) + resource_bodies.append(await request.aread()) + if len(resource_bodies) == 1: + response = httpx.Response( + 401, + headers={"DPoP-Nonce": "resource-nonce", "WWW-Authenticate": 'DPoP error="use_dpop_nonce"'}, + stream=httpx.ByteStream(b'{"error":"use_dpop_nonce"}'), + request=request, + ) + challenge_responses.append(response) + return response + return httpx.Response(201, json={"id": "1"}, request=request) + + transport = AsyncOAuth2Transport(integration_config(), httpx.MockTransport(dispatch)) + async with httpx.AsyncClient(transport=transport) as client: + response = await client.post( + "https://api.example.com/resources", + headers={"Content-Type": "application/json"}, + content=_body_stream(body), + ) + + assert response.status_code == 201 + assert resource_bodies == [body, body] + assert response.history == [] + assert challenge_responses[0].is_stream_consumed + assert challenge_responses[0].is_closed + + asyncio.run(exercise()) + + +def test_async_httpx_transport_allows_concurrent_cache_misses() -> None: + async def exercise() -> None: + token_requests = 0 + resource_requests = 0 + both_tokens_started = asyncio.Event() + + async def dispatch(request: httpx.Request) -> httpx.Response: + nonlocal resource_requests, token_requests + if request.url.host == "auth.example.com": + token_requests += 1 + if token_requests == 2: + both_tokens_started.set() + await both_tokens_started.wait() + return httpx.Response( + 200, + json={"access_token": "token", "token_type": "DPoP", "expires_in": 900, "scope": "read"}, + request=request, + ) + resource_requests += 1 + return httpx.Response(200, json={"ok": True}, request=request) + + transport = AsyncOAuth2Transport(integration_config(), httpx.MockTransport(dispatch)) + async with httpx.AsyncClient(transport=transport) as client: + responses = await asyncio.gather( + client.get("https://api.example.com/first"), + client.get("https://api.example.com/second"), + ) + + assert [response.status_code for response in responses] == [200, 200] + assert token_requests == 2 + assert resource_requests == 2 + + asyncio.run(exercise()) + + +async def _body_stream(body: bytes): # type: ignore[no-untyped-def] + yield body[:10] + yield body[10:] diff --git a/tests/unit/test_httpx_integration.py b/tests/unit/test_httpx_integration.py new file mode 100644 index 0000000..1fe8eef --- /dev/null +++ b/tests/unit/test_httpx_integration.py @@ -0,0 +1,46 @@ +import httpx +import pytest + +from mastercard_oauth2_client.integrations.httpx import OAuth2Transport +from tests.unit.configuration import integration_config + +pytestmark = pytest.mark.httpx + + +def test_httpx_transport_buffers_streaming_body_for_nonce_retry() -> None: + body = b'{"name":"streamed-resource"}' + challenge_responses: list[httpx.Response] = [] + resource_bodies: list[bytes] = [] + + def dispatch(request: httpx.Request) -> httpx.Response: + if request.url.host == "auth.example.com": + return httpx.Response( + 200, + json={"access_token": "token", "token_type": "DPoP", "expires_in": 900, "scope": "read"}, + request=request, + ) + resource_bodies.append(request.read()) + if len(resource_bodies) == 1: + response = httpx.Response( + 401, + headers={"DPoP-Nonce": "resource-nonce", "WWW-Authenticate": 'DPoP error="use_dpop_nonce"'}, + stream=httpx.ByteStream(b'{"error":"use_dpop_nonce"}'), + request=request, + ) + challenge_responses.append(response) + return response + return httpx.Response(201, json={"id": "1"}, request=request) + + transport = OAuth2Transport(integration_config(), httpx.MockTransport(dispatch)) + with httpx.Client(transport=transport) as client: + response = client.post( + "https://api.example.com/resources", + headers={"Content-Type": "application/json"}, + content=iter([body[:10], body[10:]]), + ) + + assert response.status_code == 201 + assert resource_bodies == [body, body] + assert response.history == [] + assert challenge_responses[0].is_stream_consumed + assert challenge_responses[0].is_closed diff --git a/tests/unit/test_integration_support.py b/tests/unit/test_integration_support.py new file mode 100644 index 0000000..f25e579 --- /dev/null +++ b/tests/unit/test_integration_support.py @@ -0,0 +1,24 @@ +import json + +import pytest + +from tests.integration.support.resource_server import _parse_new_resource + + +def test_fake_resource_server_accepts_schema_valid_record() -> None: + assert _parse_new_resource(b'{"name":"valid-resource"}') == {"name": "valid-resource"} + + +@pytest.mark.parametrize( + "payload", + [ + {}, + {"name": ""}, + {"name": "valid", "id": "client-controlled"}, + {"name": "valid", "unexpected": True}, + {"name": 42}, + ], +) +def test_fake_resource_server_rejects_schema_invalid_record(payload: object) -> None: + with pytest.raises(ValueError): + _parse_new_resource(json.dumps(payload).encode()) diff --git a/tests/unit/test_jose_dpop.py b/tests/unit/test_jose_dpop.py new file mode 100644 index 0000000..de5b624 --- /dev/null +++ b/tests/unit/test_jose_dpop.py @@ -0,0 +1,449 @@ +import base64 +import hashlib +import json +from collections.abc import Mapping +from datetime import UTC, datetime +from pathlib import Path +from typing import cast + +import pytest +from cryptography.hazmat.primitives import hashes +from cryptography.hazmat.primitives.asymmetric import ec, padding, rsa, utils + +from mastercard_oauth2_client import ( + KeyPair, + OAuth2Config, + OAuth2ConfigError, + OAuth2Error, + StaticDPoPKeyProvider, + StaticScopeResolver, + create_client_assertion, + create_resource_dpop_proof, + create_token_dpop_proof, + load_jwk_key_pair, +) +from mastercard_oauth2_client import dpop as dpop_module +from mastercard_oauth2_client._internal import jose as jose_module + +_FIXED_NOW = datetime(2026, 8, 26, 12, 0, tzinfo=UTC) +_FIXED_JTI = "AAECAwQFBgcICQoL" + + +@pytest.fixture +def fixed_proof_values(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(dpop_module, "_utc_now", lambda: _FIXED_NOW) + monkeypatch.setattr(dpop_module, "_random_jti", lambda: _FIXED_JTI) + + +def _base64url_uint(value: int) -> str: + size = max(1, (value.bit_length() + 7) // 8) + return base64.urlsafe_b64encode(value.to_bytes(size, "big")).rstrip(b"=").decode("ascii") + + +def _private_jwk(private_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey) -> dict[str, str]: + if isinstance(private_key, rsa.RSAPrivateKey): + rsa_numbers = private_key.private_numbers() + rsa_public = rsa_numbers.public_numbers + return { + "kty": "RSA", + "n": _base64url_uint(rsa_public.n), + "e": _base64url_uint(rsa_public.e), + "d": _base64url_uint(rsa_numbers.d), + "p": _base64url_uint(rsa_numbers.p), + "q": _base64url_uint(rsa_numbers.q), + "dp": _base64url_uint(rsa_numbers.dmp1), + "dq": _base64url_uint(rsa_numbers.dmq1), + "qi": _base64url_uint(rsa_numbers.iqmp), + } + + ec_numbers = private_key.private_numbers() + ec_public = ec_numbers.public_numbers + return { + "kty": "EC", + "crv": "P-256", + "x": _base64url_uint(ec_public.x), + "y": _base64url_uint(ec_public.y), + "d": _base64url_uint(ec_numbers.private_value), + } + + +def _decode_part(value: str) -> dict[str, object]: + decoded = json.loads(base64.urlsafe_b64decode(value + "=" * (-len(value) % 4))) + return cast(dict[str, object], decoded) + + +def _decode_jwt(token: str) -> tuple[dict[str, object], dict[str, object], bytes, bytes]: + header_part, payload_part, signature_part = token.split(".") + signature = base64.urlsafe_b64decode(signature_part + "=" * (-len(signature_part) % 4)) + return ( + _decode_part(header_part), + _decode_part(payload_part), + f"{header_part}.{payload_part}".encode(), + signature, + ) + + +def _verify_compact_jws( + token: str, + public_key: rsa.RSAPublicKey | ec.EllipticCurvePublicKey, +) -> tuple[dict[str, object], dict[str, object]]: + header, claims, signing_input, signature = _decode_jwt(token) + if isinstance(public_key, rsa.RSAPublicKey): + public_key.verify( + signature, + signing_input, + padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=32), + hashes.SHA256(), + ) + else: + component_size = (public_key.curve.key_size + 7) // 8 + r = int.from_bytes(signature[:component_size], "big") + s = int.from_bytes(signature[component_size:], "big") + public_key.verify(utils.encode_dss_signature(r, s), signing_input, ec.ECDSA(hashes.SHA256())) + return header, claims + + +def _config( + client_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey, + dpop_private_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey | None = None, +) -> OAuth2Config: + dpop_private_key = dpop_private_key or ec.generate_private_key(ec.SECP256R1()) + dpop_pair = KeyPair(private_key=dpop_private_key, public_key=dpop_private_key.public_key()) + return OAuth2Config( + client_id="client-id", + token_endpoint="https://auth.example.com/oauth/token?ignored=yes", + issuer="https://auth.example.com", + client_key=client_key, + key_id="client-key-id", + scope_resolver=StaticScopeResolver({"payments:read"}), + dpop_key_provider=StaticDPoPKeyProvider(dpop_pair), + ) + + +@pytest.mark.parametrize( + "private_key", + [ + pytest.param(rsa.generate_private_key(public_exponent=65537, key_size=2048), id="RSA"), + pytest.param(ec.generate_private_key(ec.SECP256R1()), id="EC"), + ], +) +def test_imports_private_jwk_as_native_key_pair( + private_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey, +) -> None: + loaded = load_jwk_key_pair(_private_jwk(private_key)) + + assert type(loaded.private_key) is type(private_key) + assert loaded.public_key.public_numbers() == private_key.public_key().public_numbers() + + +def test_jwk_loader_accepts_rsa_without_crt_parameters() -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + minimal_jwk = {key: value for key, value in _private_jwk(private_key).items() if key in {"kty", "n", "e", "d"}} + + loaded = load_jwk_key_pair(minimal_jwk) + + assert loaded.private_key.private_numbers() == private_key.private_numbers() + + +@pytest.mark.parametrize( + ("client_key", "algorithm"), + [ + pytest.param(rsa.generate_private_key(public_exponent=65537, key_size=2048), "PS256", id="PS256"), + pytest.param(ec.generate_private_key(ec.SECP256R1()), "ES256", id="ES256"), + ], +) +def test_client_assertion_has_required_claims_and_independent_signature( + client_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey, + algorithm: str, + fixed_proof_values: None, +) -> None: + config = _config(client_key) + + assertion = create_client_assertion(config) + + header, claims = _verify_compact_jws(assertion, client_key.public_key()) + assert header == {"alg": algorithm, "typ": "JWT", "kid": "client-key-id"} + assert claims == { + "iss": "client-id", + "sub": "client-id", + "aud": "https://auth.example.com", + "jti": _FIXED_JTI, + "iat": 1787745600, + "nbf": 1787745595, + "exp": 1787745695, + } + + +def test_token_dpop_proof_contains_public_jwk_and_nonce(fixed_proof_values: None) -> None: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + config = _config(client_key) + dpop_key = config.dpop_key_provider.get_current_key() + + proof = create_token_dpop_proof( + config, + dpop_key.key_id, + nonce="server-nonce", + ) + + header, claims = _verify_compact_jws(proof, dpop_key.key_pair.public_key) + public_jwk = header["jwk"] + assert isinstance(public_jwk, Mapping) + assert "d" not in public_jwk + assert header["alg"] == "ES256" + assert header["typ"] == "dpop+jwt" + assert header["kid"] == dpop_key.key_id + assert claims == { + "jti": _FIXED_JTI, + "htm": "POST", + "htu": "https://auth.example.com/oauth/token", + "iat": 1787745600, + "exp": 1787745695, + "nonce": "server-nonce", + } + + +def test_token_dpop_proof_supports_ps256() -> None: + client_key = ec.generate_private_key(ec.SECP256R1()) + dpop_private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + config = _config(client_key, dpop_private_key) + dpop_key = config.dpop_key_provider.get_current_key() + + proof = create_token_dpop_proof(config, dpop_key.key_id) + + header, _ = _verify_compact_jws(proof, dpop_key.key_pair.public_key) + assert header["alg"] == "PS256" + assert header["jwk"] == { + key: value for key, value in _private_jwk(dpop_private_key).items() if key in {"kty", "n", "e"} + } + + +def test_resource_dpop_proof_binds_method_url_and_access_token(fixed_proof_values: None) -> None: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + config = _config(client_key) + dpop_key = config.dpop_key_provider.get_current_key() + + proof = create_resource_dpop_proof( + config, + dpop_key.key_id, + "get", + "https://api.example.com:8443/payments?limit=10#fragment", + "access-token", + ) + + _, claims = _verify_compact_jws(proof, dpop_key.key_pair.public_key) + assert claims["htm"] == "GET" + assert claims["htu"] == "https://api.example.com:8443/payments" + assert claims["ath"] == base64.urlsafe_b64encode(hashlib.sha256(b"access-token").digest()).rstrip(b"=").decode() + + +def test_static_dpop_provider_accepts_key_pair_and_uses_thumbprint_id() -> None: + private_key = ec.generate_private_key(ec.SECP256R1()) + provider = StaticDPoPKeyProvider(KeyPair(private_key=private_key, public_key=private_key.public_key())) + + current_key = provider.get_current_key() + + public_numbers = private_key.public_key().public_numbers() + canonical = json.dumps( + { + "crv": "P-256", + "kty": "EC", + "x": _base64url_uint(public_numbers.x), + "y": _base64url_uint(public_numbers.y), + }, + separators=(",", ":"), + sort_keys=True, + ).encode() + expected_thumbprint = base64.urlsafe_b64encode(hashlib.sha256(canonical).digest()).rstrip(b"=").decode() + assert current_key.key_id == expected_thumbprint + assert provider.get_key(current_key.key_id) is current_key + + +def test_rsa_jwk_thumbprint_matches_known_vector() -> None: + modulus = ( + "wxnY2XfkJDaA_qIYUHMbT_5RXnE1xK2YdDiwRwuo1JaNa_aZhqqw4u1dg9ztvyCpsbL_VL_FaExSSrK6OSmQJYpisUROux" + "C1ep6Vn7IcuzJmmhUX_vaElWFCEAST5LuFdgnBR8wmChVTh4BHDXmmL0NzJVGXnzwcQN1COP26usmi8-HB5Vr0COYqD8TdVcYyw" + "fsuhbiQY0uFyl8HQuIdiNx4TZBut3nv4Ii33n1HwlESTxgkmTnnOIwEVicug7sep4lh-5mXaGMhObIXzz-SZl2hMwRHpWFr8HH_" + "youIUfbSEgSWmJsvw5PA4XY4awWdUnC-9U7tKGsE36VWfqfDJw" + ) + modulus_bytes = base64.urlsafe_b64decode(modulus + "=" * (-len(modulus) % 4)) + public_key = rsa.RSAPublicNumbers(e=65537, n=int.from_bytes(modulus_bytes, "big")).public_key() + + assert jose_module.jwk_thumbprint(public_key) == "-cSeNq9eyhJsLmX6Nxg_qZ7H0heh0tqnFwkEIlHfRkc" + + +def test_jwk_loader_rejects_public_only_key() -> None: + private_key = ec.generate_private_key(ec.SECP256R1()) + public_only_jwk = {key: value for key, value in _private_jwk(private_key).items() if key != "d"} + + with pytest.raises(OAuth2ConfigError, match=r"Missing required EC JWK parameters \(crv, x, y, d\)"): + load_jwk_key_pair(public_only_jwk) + + +def test_jwk_loader_rejects_non_p256_ec_key() -> None: + private_key = ec.generate_private_key(ec.SECP384R1()) + jwk = _private_jwk(private_key) + jwk["crv"] = "P-384" + + with pytest.raises(OAuth2ConfigError, match="Unsupported curve: P-384"): + load_jwk_key_pair(jwk) + + +@pytest.mark.parametrize( + "private_key", + [ + pytest.param(rsa.generate_private_key(public_exponent=65537, key_size=2048), id="RSA"), + pytest.param(ec.generate_private_key(ec.SECP256R1()), id="EC"), + ], +) +def test_jwk_loader_accepts_json_bytes_and_file_path( + private_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey, + tmp_path: Path, +) -> None: + jwk_bytes = json.dumps(_private_jwk(private_key)).encode() + jwk_path = tmp_path / "dpop-key.json" + jwk_path.write_bytes(jwk_bytes) + + from_bytes = load_jwk_key_pair(jwk_bytes) + from_path = load_jwk_key_pair(jwk_path) + + expected_public_numbers = private_key.public_key().public_numbers() + assert from_bytes.public_key.public_numbers() == expected_public_numbers + assert from_path.public_key.public_numbers() == expected_public_numbers + + +def test_jwk_loader_reports_malformed_json() -> None: + with pytest.raises(OAuth2ConfigError, match="Unable to parse JWK JSON"): + load_jwk_key_pair(b"not-json") + + +@pytest.mark.parametrize( + ("value", "message"), + [ + pytest.param({"n": "AQ", "e": "AQAB", "d": "AQ"}, "Missing required JWK parameter: kty", id="missing-kty"), + pytest.param({"kty": "DSA", "d": "AQ"}, "Unsupported key type: DSA", id="unsupported-kty"), + pytest.param( + {"kty": "RSA", "e": "AQAB"}, + r"Missing required RSA JWK parameters \(n, e, d\)", + id="missing-rsa-members", + ), + pytest.param( + {"kty": "EC", "crv": "P-256"}, + r"Missing required EC JWK parameters \(crv, x, y, d\)", + id="missing-ec-members", + ), + ], +) +def test_jwk_loader_rejects_missing_or_unsupported_members(value: dict[str, str], message: str) -> None: + with pytest.raises(OAuth2ConfigError, match=message): + load_jwk_key_pair(value) + + +def test_public_jwk_rejects_unsupported_key_with_clear_message() -> None: + with pytest.raises(OAuth2Error, match="Unsupported public key type: object"): + jose_module.public_jwk(object()) # type: ignore[arg-type] + + +def test_client_assertion_rejects_non_p256_ec_key() -> None: + config = _config(ec.generate_private_key(ec.SECP384R1())) + + with pytest.raises(OAuth2ConfigError, match="Unsupported curve: secp384r1"): + create_client_assertion(config) + + +def test_signing_rejects_unsupported_algorithm() -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + + with pytest.raises(OAuth2ConfigError, match="Unsupported algorithm: RS256"): + jose_module._sign_bytes(b"data", private_key, "RS256") # type: ignore[arg-type] + + +def test_jwt_serialization_requires_signature() -> None: + with pytest.raises(OAuth2Error, match="Signature is required"): + jose_module._serialize_jwt("header.payload", b"") + + +def test_sign_jwt_wraps_unexpected_signing_failure(monkeypatch: pytest.MonkeyPatch) -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + + def fail_signing(data: bytes, key: object, algorithm: object) -> bytes: + raise RuntimeError("provider failure") + + monkeypatch.setattr(jose_module, "_sign_bytes", fail_signing) + + with pytest.raises(OAuth2Error, match="Unable to sign JWT") as caught: + jose_module.sign_jwt({"typ": "JWT"}, {"sub": "client"}, private_key) + + assert isinstance(caught.value.__cause__, RuntimeError) + + +@pytest.mark.parametrize( + ("private_key", "public_key", "message"), + [ + pytest.param( + rsa.generate_private_key(public_exponent=65537, key_size=2048), + ec.generate_private_key(ec.SECP256R1()).public_key(), + "ES256 requires an EC private key", + id="RSA-private-EC-public", + ), + pytest.param( + ec.generate_private_key(ec.SECP256R1()), + rsa.generate_private_key(public_exponent=65537, key_size=2048).public_key(), + "PS256 requires an RSA private key", + id="EC-private-RSA-public", + ), + ], +) +def test_dpop_proof_rejects_mismatched_key_algorithms_during_signing( + private_key: rsa.RSAPrivateKey | ec.EllipticCurvePrivateKey, + public_key: rsa.RSAPublicKey | ec.EllipticCurvePublicKey, + message: str, +) -> None: + key_pair = KeyPair(private_key=private_key, public_key=public_key) + config = OAuth2Config( + client_id="client-id", + token_endpoint="https://auth.example.com/oauth/token", + issuer="https://auth.example.com", + client_key=rsa.generate_private_key(public_exponent=65537, key_size=2048), + key_id="client-key-id", + scope_resolver=StaticScopeResolver({"payments:read"}), + dpop_key_provider=StaticDPoPKeyProvider(key_pair), + ) + dpop_key = config.dpop_key_provider.get_current_key() + + with pytest.raises(OAuth2ConfigError, match=message): + create_token_dpop_proof(config, dpop_key.key_id) + + +def test_generated_jti_contains_96_random_bits() -> None: + config = _config(rsa.generate_private_key(public_exponent=65537, key_size=2048)) + + assertion = create_client_assertion(config) + _, claims, _, _ = _decode_jwt(assertion) + jti = claims["jti"] + + assert isinstance(jti, str) + assert len(base64.urlsafe_b64decode(jti + "=" * (-len(jti) % 4))) == 12 + + +def test_dpop_proof_rejects_invalid_url() -> None: + config = _config(rsa.generate_private_key(public_exponent=65537, key_size=2048)) + dpop_key = config.dpop_key_provider.get_current_key() + + with pytest.raises(OAuth2Error, match="Invalid URL for DPoP htu claim"): + create_resource_dpop_proof(config, dpop_key.key_id, "GET", "not-a-url", "access-token") + + +def test_dpop_proof_preserves_explicit_default_port(fixed_proof_values: None) -> None: + config = _config(rsa.generate_private_key(public_exponent=65537, key_size=2048)) + dpop_key = config.dpop_key_provider.get_current_key() + + proof = create_resource_dpop_proof( + config, + dpop_key.key_id, + "GET", + "https://api.example.com:443/payments?limit=10#fragment", + "access-token", + ) + + _, claims, _, _ = _decode_jwt(proof) + assert claims["htu"] == "https://api.example.com:443/payments" diff --git a/tests/unit/test_keys.py b/tests/unit/test_keys.py new file mode 100644 index 0000000..f98c951 --- /dev/null +++ b/tests/unit/test_keys.py @@ -0,0 +1,287 @@ +from collections.abc import Callable +from dataclasses import replace +from datetime import UTC, datetime, timedelta +from pathlib import Path + +import pytest +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from cryptography.hazmat.primitives.serialization import pkcs12 +from cryptography.x509.oid import NameOID + +from mastercard_oauth2_client import DPoPKey, KeyPair, OAuth2ConfigError +from mastercard_oauth2_client.keys import ( + load_jwk_key_pair, + load_pkcs12_private_key, + load_private_key, + validate_dpop_key, + validate_key, +) +from tests.unit.test_config import minimum_config_values + + +class FailingDPoPKeyProvider: + def get_current_key(self) -> DPoPKey: + raise ConnectionError("key service unavailable") + + def get_key(self, key_id: str) -> DPoPKey: + raise ConnectionError("key service unavailable") + + +class InvalidIdDPoPKeyProvider: + def __init__(self, key_pair: KeyPair) -> None: + self._key = DPoPKey(key_id=" ", key_pair=key_pair) + + def get_current_key(self) -> DPoPKey: + return self._key + + def get_key(self, key_id: str) -> DPoPKey: + return self._key + + +def test_validate_key_accepts_strong_rsa_and_ec_private_and_public_keys() -> None: + rsa_private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + ec_private_key = ec.generate_private_key(ec.SECP256R1()) + + validate_key(rsa_private_key) + validate_key(rsa_private_key.public_key()) + validate_key(ec_private_key) + validate_key(ec_private_key.public_key()) + + +def test_validate_key_rejects_weak_rsa_with_clear_message() -> None: + weak_key = rsa.generate_private_key(public_exponent=65537, key_size=1024) + + with pytest.raises( + OAuth2ConfigError, + match="RSA keys must have a minimum length of 2048 bits, but key length was: 1024", + ): + validate_key(weak_key) + + +def test_validate_key_rejects_weak_ec_with_clear_message() -> None: + weak_key = ec.generate_private_key(ec.SECP192R1()) + + with pytest.raises( + OAuth2ConfigError, + match="Elliptic curve keys must have a minimum length of 224 bits, but key length was: 192", + ): + validate_key(weak_key) + + +def test_validate_key_rejects_unsupported_algorithm_with_clear_message() -> None: + with pytest.raises(OAuth2ConfigError, match="Key algorithm must be RSA or EC, but was: object"): + validate_key(object()) + + +@pytest.mark.parametrize("encoding", [serialization.Encoding.PEM, serialization.Encoding.DER]) +def test_load_private_key_supports_pkcs8_rsa(encoding: serialization.Encoding) -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + key_data = private_key.private_bytes( + encoding=encoding, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + + loaded = load_private_key(key_data) + + assert isinstance(loaded, rsa.RSAPrivateKey) + assert loaded.private_numbers() == private_key.private_numbers() + + +def test_load_private_key_supports_pkcs1_pem() -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + key_data = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.TraditionalOpenSSL, + encryption_algorithm=serialization.NoEncryption(), + ) + + assert isinstance(load_private_key(key_data), rsa.RSAPrivateKey) + + +def test_load_private_key_supports_encrypted_pem_with_string_password() -> None: + private_key = ec.generate_private_key(ec.SECP256R1()) + key_data = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.BestAvailableEncryption(b"secret"), + ) + + loaded = load_private_key(key_data, password="secret") + + assert isinstance(loaded, ec.EllipticCurvePrivateKey) + + +def test_load_private_key_accepts_a_file_path(tmp_path: Path) -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + key_path = tmp_path / "key.pem" + key_path.write_bytes( + private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + ) + + assert isinstance(load_private_key(key_path), rsa.RSAPrivateKey) + + +def test_load_pkcs12_private_key_supports_string_password() -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + subject = issuer = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "test")]) + certificate = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(issuer) + .public_key(private_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(datetime.now(UTC)) + .not_valid_after(datetime.now(UTC) + timedelta(days=1)) + .sign(private_key, hashes.SHA256()) + ) + key_data = pkcs12.serialize_key_and_certificates( + name=b"key", + key=private_key, + cert=certificate, + cas=None, + encryption_algorithm=serialization.BestAvailableEncryption(b"secret"), + ) + + loaded = load_pkcs12_private_key(key_data, password="secret") + + assert isinstance(loaded, rsa.RSAPrivateKey) + + +def test_load_pkcs12_private_key_accepts_a_file_path(tmp_path: Path) -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + key_data = pkcs12.serialize_key_and_certificates( + name=b"key", + key=private_key, + cert=None, + cas=None, + encryption_algorithm=serialization.BestAvailableEncryption(b"secret"), + ) + key_path = tmp_path / "key.p12" + key_path.write_bytes(key_data) + + loaded = load_pkcs12_private_key(key_path, password="secret") + + assert isinstance(loaded, rsa.RSAPrivateKey) + assert loaded.private_numbers() == private_key.private_numbers() + + +@pytest.mark.parametrize("loader", [load_private_key, load_pkcs12_private_key, load_jwk_key_pair]) +def test_key_loaders_propagate_missing_file_path(loader: Callable[[Path], object], tmp_path: Path) -> None: + missing_path = tmp_path / "missing.key" + + with pytest.raises(FileNotFoundError, match="missing.key"): + loader(missing_path) + + +def test_load_pkcs12_private_key_wraps_wrong_password() -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + key_data = pkcs12.serialize_key_and_certificates( + name=b"key", + key=private_key, + cert=None, + cas=None, + encryption_algorithm=serialization.BestAvailableEncryption(b"correct"), + ) + + with pytest.raises(OAuth2ConfigError, match="Unable to load PKCS#12 private key") as caught: + load_pkcs12_private_key(key_data, password="wrong") + + assert "wrong" not in str(caught.value) + + +def test_load_private_key_wraps_invalid_data_without_echoing_it() -> None: + invalid_data = b"super-secret-invalid-key-data" + + with pytest.raises(OAuth2ConfigError, match="Unable to load private key") as caught: + load_private_key(invalid_data) + + assert invalid_data.decode() not in str(caught.value) + + +def test_config_rejects_weak_client_key() -> None: + values = minimum_config_values() + values["client_key"] = rsa.generate_private_key(public_exponent=65537, key_size=1024) + + with pytest.raises(OAuth2ConfigError, match="RSA keys must have a minimum length of 2048 bits"): + from mastercard_oauth2_client import OAuth2Config + + OAuth2Config(**values) + + +def test_config_rejects_weak_ec_client_key() -> None: + values = minimum_config_values() + values["client_key"] = ec.generate_private_key(ec.SECP192R1()) + + with pytest.raises(OAuth2ConfigError, match="Elliptic curve keys must have a minimum length of 224 bits"): + from mastercard_oauth2_client import OAuth2Config + + OAuth2Config(**values) + + +def test_config_rejects_weak_ec_dpop_key() -> None: + from mastercard_oauth2_client import OAuth2Config + + weak_private_key = ec.generate_private_key(ec.SECP192R1()) + weak_dpop_key = DPoPKey( + key_id="weak-key", + key_pair=KeyPair(private_key=weak_private_key, public_key=weak_private_key.public_key()), + ) + + class WeakKeyProvider: + def get_current_key(self) -> DPoPKey: + return weak_dpop_key + + def get_key(self, key_id: str) -> DPoPKey: + return weak_dpop_key + + values = minimum_config_values() + values["dpop_key_provider"] = WeakKeyProvider() # type: ignore[typeddict-item] + + with pytest.raises(OAuth2ConfigError, match="Elliptic curve keys must have a minimum length of 224 bits"): + OAuth2Config(**values) + + +def test_config_rejects_public_client_key() -> None: + from mastercard_oauth2_client import OAuth2Config + + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(OAuth2ConfigError, match="Client key must be an RSA or EC private key"): + replace(config, client_key=config.client_key.public_key()) # type: ignore[arg-type] + + +def test_config_rejects_invalid_dpop_provider_key_id() -> None: + from mastercard_oauth2_client import OAuth2Config + + config = OAuth2Config(**minimum_config_values()) + provider = InvalidIdDPoPKeyProvider(config.dpop_key_provider.get_current_key().key_pair) + + with pytest.raises(OAuth2ConfigError, match="DPoP key provider must return a valid DPoP key ID"): + replace(config, dpop_key_provider=provider) + + +def test_validate_dpop_key_rejects_missing_key_pair() -> None: + private_key = ec.generate_private_key(ec.SECP256R1()) + dpop_key = DPoPKey( + key_id="dpop-key", + key_pair=KeyPair(private_key=private_key, public_key=private_key.public_key()), + ) + + with pytest.raises(OAuth2ConfigError, match="DPoP key provider must return a valid DPoP key"): + validate_dpop_key(replace(dpop_key, key_pair=None)) # type: ignore[arg-type] + + +def test_config_preserves_dpop_key_provider_errors() -> None: + from mastercard_oauth2_client import OAuth2Config + + config = OAuth2Config(**minimum_config_values()) + + with pytest.raises(ConnectionError, match="key service unavailable"): + replace(config, dpop_key_provider=FailingDPoPKeyProvider()) diff --git a/tests/unit/test_openapi_adapter.py b/tests/unit/test_openapi_adapter.py new file mode 100644 index 0000000..82c3213 --- /dev/null +++ b/tests/unit/test_openapi_adapter.py @@ -0,0 +1,262 @@ +from __future__ import annotations + +import base64 +import json +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from mastercard_oauth2_client import KeyPair, OAuth2Config, OAuth2Error, StaticDPoPKeyProvider, StaticScopeResolver +from mastercard_oauth2_client.integrations.openapi import add_oauth2_layer +from mastercard_oauth2_client.integrations.openapi.urllib3 import _generated_body +from mastercard_oauth2_client.models import HttpRequest + +pytestmark = [pytest.mark.generated, pytest.mark.urllib3] + + +@dataclass +class GeneratedResponse: + status: int + body: bytes + headers: Mapping[str, str] + reason: str = "" + data: bytes | None = None + read_called: bool = False + drain_called: bool = False + release_called: bool = False + + @property + def response(self) -> GeneratedResponse: + return self + + def read(self) -> bytes: + self.read_called = True + self.data = self.body + return self.body + + def drain_conn(self) -> None: + self.drain_called = True + + def release_conn(self) -> None: + self.release_called = True + + def getheaders(self) -> Mapping[str, str]: + return self.headers + + def getheader(self, name: str, default: str | None = None) -> str | None: + return next((value for key, value in self.headers.items() if key.casefold() == name.casefold()), default) + + +class ScriptedRestClient: + def __init__(self, responses: list[GeneratedResponse]) -> None: + self.seen_responses = responses + self.responses = iter(self.seen_responses) + self.requests: list[tuple[str, str, dict[str, Any]]] = [] + + def request(self, method: str, url: str, **kwargs: Any) -> GeneratedResponse: + self.requests.append((method, url, kwargs)) + return next(self.responses) + + +class GeneratedApiClient: + def __init__(self, rest_client: ScriptedRestClient) -> None: + self.rest_client = rest_client + + +def _config() -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_private_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client-id", + token_endpoint="https://auth.example.com/oauth/token", + issuer="https://auth.example.com", + client_key=client_key, + key_id="client-key-id", + scope_resolver=StaticScopeResolver({"payments:read"}), + dpop_key_provider=StaticDPoPKeyProvider( + KeyPair(private_key=dpop_private_key, public_key=dpop_private_key.public_key()) + ), + ) + + +def _token_response() -> GeneratedResponse: + return GeneratedResponse( + status=200, + headers={"Content-Type": "application/json"}, + body=json.dumps( + { + "access_token": "access-token", + "token_type": "DPoP", + "expires_in": 900, + "scope": "payments:read", + } + ).encode(), + ) + + +def _nonce_challenge(nonce: str, *, status: int) -> GeneratedResponse: + return GeneratedResponse( + status=status, + headers={"DPoP-Nonce": nonce, "WWW-Authenticate": 'DPoP error="use_dpop_nonce"'}, + body=b'{"error":"use_dpop_nonce"}', + ) + + +def _claims(compact_jwt: str) -> dict[str, object]: + payload = compact_jwt.split(".")[1] + return json.loads(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4))) # type: ignore[no-any-return] + + +def test_generated_client_call_is_authenticated_without_generated_code_changes() -> None: + rest_client = ScriptedRestClient( + [ + _token_response(), + GeneratedResponse(status=200, headers={"Content-Type": "application/json"}, body=b'{"result":"ok"}'), + GeneratedResponse(status=200, headers={"Content-Type": "application/json"}, body=b'{"result":"cached"}'), + ] + ) + api_client = GeneratedApiClient(rest_client) + add_oauth2_layer(api_client, _config()) + + first = api_client.rest_client.request( + "GET", + "https://api.example.com/payments", + query_params=[("account", "123")], + headers={"Accept": "application/json"}, + _preload_content=True, + _request_timeout=(1, 2), + ) + second = api_client.rest_client.request("GET", "https://api.example.com/payments", headers={}) + + assert first.status == 200 + assert first.read() == b'{"result":"ok"}' + assert first.getheader("content-type") == "application/json" + assert second.read() == b'{"result":"cached"}' + assert [url for _, url, _ in rest_client.requests] == [ + "https://auth.example.com/oauth/token", + "https://api.example.com/payments?account=123", + "https://api.example.com/payments", + ] + token_kwargs = rest_client.requests[0][2] + assert token_kwargs["body"] is None + assert dict(token_kwargs["post_params"])["grant_type"] == "client_credentials" + first_resource_kwargs = rest_client.requests[1][2] + assert first_resource_kwargs["headers"]["Authorization"] == "DPoP access-token" + assert "DPoP" in first_resource_kwargs["headers"] + assert first_resource_kwargs["headers"]["Accept"] == "application/json" + assert first_resource_kwargs["_request_timeout"] == (1, 2) + token_response, first_resource_response, second_resource_response = rest_client.seen_responses + assert token_response.read_called and token_response.drain_called and token_response.release_called + assert first_resource_response.read_called and not first_resource_response.drain_called + assert not first_resource_response.release_called + assert second_resource_response.read_called and not second_resource_response.drain_called + assert not second_resource_response.release_called + + +def test_add_oauth2_layer_rejects_duplicate_attachment() -> None: + api_client = GeneratedApiClient(ScriptedRestClient([])) + config = _config() + add_oauth2_layer(api_client, config) + + with pytest.raises(OAuth2Error, match="already attached"): + add_oauth2_layer(api_client, config) + + +def test_generated_client_retries_token_and_resource_nonce_challenges_once() -> None: + rest_client = ScriptedRestClient( + [ + _nonce_challenge("token-nonce", status=400), + _token_response(), + _nonce_challenge("resource-nonce", status=401), + GeneratedResponse(status=200, headers={}, body=b'{"result":"ok"}'), + ] + ) + api_client = GeneratedApiClient(rest_client) + add_oauth2_layer(api_client, _config()) + + response = api_client.rest_client.request("GET", "https://api.example.com/payments", headers={}) + + assert response.status == 200 + assert [request[1] for request in rest_client.requests] == [ + "https://auth.example.com/oauth/token", + "https://auth.example.com/oauth/token", + "https://api.example.com/payments", + "https://api.example.com/payments", + ] + proofs = [request[2]["headers"]["DPoP"] for request in rest_client.requests] + assert proofs[0] != proofs[1] + assert _claims(proofs[1])["nonce"] == "token-nonce" + assert proofs[2] != proofs[3] + assert _claims(proofs[3])["nonce"] == "resource-nonce" + token_challenge, token_response, resource_challenge, final_response = rest_client.seen_responses + assert token_challenge.drain_called and token_challenge.release_called + assert token_response.read_called and token_response.drain_called and token_response.release_called + assert resource_challenge.drain_called and resource_challenge.release_called + assert not final_response.read_called and not final_response.drain_called and not final_response.release_called + + +def test_generated_client_preserves_json_post_body() -> None: + rest_client = ScriptedRestClient( + [ + _token_response(), + GeneratedResponse(status=201, headers={"Content-Type": "application/json"}, body=b'{"id":"123"}'), + ] + ) + api_client = GeneratedApiClient(rest_client) + add_oauth2_layer(api_client, _config()) + + response = api_client.rest_client.request( + "POST", + "https://api.example.com/payments", + headers={"Content-Type": "application/json", "X-Correlation-ID": "correlation-id"}, + body={"amount": 42, "currency": "USD"}, + ) + + assert response.status == 201 + resource_kwargs = rest_client.requests[1][2] + assert resource_kwargs["body"] == {"amount": 42, "currency": "USD"} + assert resource_kwargs["headers"]["X-Correlation-ID"] == "correlation-id" + assert resource_kwargs["headers"]["Authorization"] == "DPoP access-token" + + +def test_generated_client_preserves_unrecognized_body_without_reencoding() -> None: + body = b"raw-generated-body" + rest_client = ScriptedRestClient( + [ + _token_response(), + GeneratedResponse(status=202, headers={}, body=b""), + ] + ) + api_client = GeneratedApiClient(rest_client) + add_oauth2_layer(api_client, _config()) + + response = api_client.rest_client.request( + "POST", + "https://api.example.com/upload", + headers={"Content-Type": "application/octet-stream"}, + body=body, + ) + + assert response.status == 202 + assert rest_client.requests[1][2]["body"] is body + + +def test_generated_token_body_adapter_supports_json_and_raw_fallbacks() -> None: + json_body = HttpRequest( + method="POST", + url="https://auth.example.com/token", + headers={"Content-Type": "application/json"}, + body='{"audience":"payments"}', + ) + raw_body = HttpRequest( + method="POST", + url="https://auth.example.com/token", + headers={"Content-Type": "application/octet-stream"}, + body=b"raw-token-body", + ) + + assert _generated_body(json_body) == ({"audience": "payments"}, None) + assert _generated_body(raw_body) == (b"raw-token-body", None) diff --git a/tests/unit/test_openapi_asyncio_adapter.py b/tests/unit/test_openapi_asyncio_adapter.py new file mode 100644 index 0000000..b17749b --- /dev/null +++ b/tests/unit/test_openapi_asyncio_adapter.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, field +from typing import Any + +import aiohttp +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from mastercard_oauth2_client import ( + KeyPair, + OAuth2Config, + OAuth2Error, + StaticDPoPKeyProvider, + StaticScopeResolver, +) +from mastercard_oauth2_client.integrations.aiohttp import OAuth2Middleware +from mastercard_oauth2_client.integrations.openapi.asyncio import add_oauth2_layer + +pytestmark = [pytest.mark.generated, pytest.mark.aiohttp] + + +@dataclass +class Configuration: + client_session_kwargs: dict[str, Any] | None = None + + +@dataclass +class RestClient: + configuration: Configuration = field(default_factory=Configuration) + pool_manager: aiohttp.ClientSession | None = None + + +class ApiClient: + def __init__(self, rest_client: RestClient | None = None) -> None: + self._rest_client = rest_client or RestClient() + + @property + def rest_client(self) -> RestClient: + return self._rest_client + + +def _config() -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client-id", + key_id="key-id", + token_endpoint="https://auth.example.com/token", + issuer="https://auth.example.com", + client_key=client_key, + scope_resolver=StaticScopeResolver({"read"}), + dpop_key_provider=StaticDPoPKeyProvider(KeyPair(dpop_key, dpop_key.public_key())), + ) + + +def test_add_oauth2_layer_preserves_session_kwargs_and_existing_middlewares() -> None: + async def existing_middleware( + request: aiohttp.ClientRequest, + handler: aiohttp.ClientHandlerType, + ) -> aiohttp.ClientResponse: + return await handler(request) + + original_kwargs: dict[str, Any] = { + "cookie_jar": "generated-cookie-jar", + "middlewares": (existing_middleware,), + } + api_client = ApiClient(RestClient(Configuration(original_kwargs))) + + add_oauth2_layer(api_client, _config()) + + session_kwargs = api_client.rest_client.configuration.client_session_kwargs + assert session_kwargs is not None + assert session_kwargs is not original_kwargs + assert session_kwargs["cookie_jar"] == "generated-cookie-jar" + middlewares = session_kwargs["middlewares"] + assert isinstance(middlewares, tuple) + assert isinstance(middlewares[0], OAuth2Middleware) + assert middlewares[1:] == (existing_middleware,) + assert original_kwargs["middlewares"] == (existing_middleware,) + + +def test_add_oauth2_layer_rejects_attachment_after_pool_creation() -> None: + async def exercise() -> None: + pool_manager = aiohttp.ClientSession() + api_client = ApiClient(RestClient(pool_manager=pool_manager)) + try: + with pytest.raises(OAuth2Error, match="before the generated asyncio client sends its first request"): + add_oauth2_layer(api_client, _config()) + finally: + await pool_manager.close() + + asyncio.run(exercise()) + + +def test_add_oauth2_layer_rejects_duplicate_attachment_before_pool_creation() -> None: + api_client = ApiClient() + config = _config() + add_oauth2_layer(api_client, config) + + with pytest.raises(OAuth2Error, match="already attached"): + add_oauth2_layer(api_client, config) diff --git a/tests/unit/test_openapi_httpx_adapter.py b/tests/unit/test_openapi_httpx_adapter.py new file mode 100644 index 0000000..838c734 --- /dev/null +++ b/tests/unit/test_openapi_httpx_adapter.py @@ -0,0 +1,165 @@ +from __future__ import annotations + +import asyncio +import ssl +from dataclasses import dataclass, field + +import httpx +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from generated_httpx_test_client import ( # type: ignore[import-not-found] + ApiClient as GeneratedApiClient, +) +from generated_httpx_test_client import ( # type: ignore[import-not-found] + Configuration as GeneratedConfiguration, +) +from generated_httpx_test_client import ( # type: ignore[import-not-found] + ResourcesApi as GeneratedResourcesApi, +) +from generated_httpx_test_client.exceptions import ApiException # type: ignore[import-not-found] + +import mastercard_oauth2_client.integrations.openapi.httpx as openapi_httpx +from mastercard_oauth2_client import ( + KeyPair, + OAuth2Config, + OAuth2Error, + StaticDPoPKeyProvider, + StaticScopeResolver, +) +from mastercard_oauth2_client.integrations.openapi.httpx import add_oauth2_layer + +pytestmark = [pytest.mark.generated, pytest.mark.httpx] + + +@dataclass +class RestClient: + maxsize: int = 10 + ssl_context: ssl.SSLContext = field(default_factory=ssl.create_default_context) + proxy: str | None = None + proxy_headers: dict[str, str] | None = None + pool_manager: httpx.AsyncClient | None = None + + +@dataclass +class ApiClient: + rest_client: RestClient + + +def _config() -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client-id", + key_id="key-id", + token_endpoint="https://auth.example.com/token", + issuer="https://auth.example.com", + client_key=client_key, + scope_resolver=StaticScopeResolver({"read"}), + dpop_key_provider=StaticDPoPKeyProvider(KeyPair(dpop_key, dpop_key.public_key())), + ) + + +def test_add_oauth2_layer_installs_authenticated_async_client() -> None: + async def exercise() -> None: + requests: list[httpx.Request] = [] + + async def dispatch(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.host == "auth.example.com": + return httpx.Response( + 200, + json={"access_token": "token", "token_type": "DPoP", "expires_in": 900, "scope": "read"}, + request=request, + ) + return httpx.Response(200, json={"ok": True}, request=request) + + api_client = ApiClient(RestClient()) + add_oauth2_layer(api_client, _config(), httpx.MockTransport(dispatch)) + assert api_client.rest_client.pool_manager is not None + async with api_client.rest_client.pool_manager as client: + response = await client.get("https://api.example.com/resources") + + assert response.json() == {"ok": True} + assert [request.url.host for request in requests] == ["auth.example.com", "api.example.com"] + assert requests[-1].headers["Authorization"] == "DPoP token" + + asyncio.run(exercise()) + + +def test_add_oauth2_layer_rejects_attachment_after_pool_creation() -> None: + pool_manager = httpx.AsyncClient() + api_client = ApiClient(RestClient(pool_manager=pool_manager)) + try: + with pytest.raises(OAuth2Error, match="before the generated HTTPX client sends its first request"): + add_oauth2_layer(api_client, _config()) + finally: + asyncio.run(pool_manager.aclose()) + + +def test_add_oauth2_layer_preserves_generated_transport_settings(monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict[str, object] = {} + mock_transport = httpx.MockTransport(lambda request: httpx.Response(200, request=request)) + + def create_transport(**kwargs: object) -> httpx.AsyncBaseTransport: + captured.update(kwargs) + return mock_transport + + monkeypatch.setattr(openapi_httpx.httpx, "AsyncHTTPTransport", create_transport) + ssl_context = ssl.create_default_context() + api_client = ApiClient( + RestClient( + maxsize=7, + ssl_context=ssl_context, + proxy="http://proxy.example.com:8080", + proxy_headers={"Proxy-Authorization": "Basic value"}, + ) + ) + + add_oauth2_layer(api_client, _config()) + assert api_client.rest_client.pool_manager is not None + asyncio.run(api_client.rest_client.pool_manager.aclose()) + + assert captured["verify"] is ssl_context + assert captured["trust_env"] is True + limits = captured["limits"] + assert isinstance(limits, httpx.Limits) + assert limits.max_connections == 7 + proxy = captured["proxy"] + assert isinstance(proxy, httpx.Proxy) + assert str(proxy.url) == "http://proxy.example.com:8080" + assert proxy.headers["Proxy-Authorization"] == "Basic value" + + +def test_generated_request_timeout_reaches_resource_and_token_transports() -> None: + async def exercise() -> None: + requests: list[httpx.Request] = [] + + async def dispatch(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.host == "auth.example.com": + return httpx.Response( + 200, + json={"access_token": "token", "token_type": "DPoP", "expires_in": 900, "scope": "read"}, + request=request, + ) + return httpx.Response(404, json={"error": "not_found"}, request=request) + + generated_config = GeneratedConfiguration(host="https://api.example.com") + api_client = GeneratedApiClient(generated_config) + add_oauth2_layer(api_client, _config(), httpx.MockTransport(dispatch)) + try: + with pytest.raises(ApiException): + await GeneratedResourcesApi(api_client).get_resource("missing", _request_timeout=(1.0, 2.0)) + finally: + await api_client.close() + + assert len(requests) == 2 + for request in requests: + assert request.extensions["timeout"] == { + "connect": 1.0, + "read": 2.0, + "write": None, + "pool": None, + } + + asyncio.run(exercise()) diff --git a/tests/unit/test_orchestration.py b/tests/unit/test_orchestration.py new file mode 100644 index 0000000..58b7fec --- /dev/null +++ b/tests/unit/test_orchestration.py @@ -0,0 +1,12 @@ +import pytest + +from mastercard_oauth2_client._internal.orchestration import has_nonce_challenge_body + + +@pytest.mark.parametrize("body", [None, b"", "not-json", "[]"]) +def test_nonce_challenge_body_ignores_absent_or_invalid_error_objects(body: bytes | str | None) -> None: + assert not has_nonce_challenge_body(401, body) + + +def test_nonce_challenge_body_requires_supported_status() -> None: + assert not has_nonce_challenge_body(200, '{"error":"use_dpop_nonce"}') diff --git a/tests/unit/test_protocols.py b/tests/unit/test_protocols.py new file mode 100644 index 0000000..928bca0 --- /dev/null +++ b/tests/unit/test_protocols.py @@ -0,0 +1,312 @@ +import asyncio +from collections.abc import Mapping +from datetime import UTC, datetime, timedelta + +import pytest +from cryptography.hazmat.primitives.asymmetric import ec + +from mastercard_oauth2_client import ( + AccessToken, + AccessTokenFilter, + AsyncHttpAdapter, + DPoPKeyProvider, + HttpRequest, + InMemoryTokenStore, + KeyPair, + ScopeResolver, + StaticDPoPKeyProvider, + StaticScopeResolver, + SyncHttpAdapter, + TokenStore, +) + + +class FixedClock: + def __init__(self, current: datetime) -> None: + self.current = current + + def __call__(self) -> datetime: + return self.current + + +class NativeAdapter: + def request_method(self, request: tuple[str, str]) -> str: + return request[0] + + def request_url(self, request: tuple[str, str]) -> str: + return request[1] + + def request_headers(self, request: tuple[str, str]) -> dict[str, str]: + return {} + + def send_token_request(self, resource_request: tuple[str, str], request: HttpRequest) -> object: + return object() + + def send_resource_request(self, request: tuple[str, str], headers: Mapping[str, str]) -> object: + return object() + + def response_status(self, response: object) -> int: + return 200 + + def response_headers(self, response: object) -> dict[str, str]: + return {} + + def response_body(self, response: object) -> bytes: + return b"" + + def close_response(self, response: object) -> None: + pass + + +class AsyncNativeAdapter: + def request_method(self, request: tuple[str, str]) -> str: + return request[0] + + def request_url(self, request: tuple[str, str]) -> str: + return request[1] + + def request_headers(self, request: tuple[str, str]) -> Mapping[str, str]: + return {} + + async def send_token_request(self, resource_request: tuple[str, str], request: HttpRequest) -> object: + return object() + + async def send_resource_request(self, request: tuple[str, str], headers: Mapping[str, str]) -> object: + return object() + + def response_status(self, response: object) -> int: + return 200 + + def response_headers(self, response: object) -> Mapping[str, str]: + return {} + + async def response_body(self, response: object) -> bytes: + return b"" + + async def close_response(self, response: object) -> None: + pass + + +def test_static_scope_resolver_satisfies_protocol_and_normalizes_scopes() -> None: + resolver: ScopeResolver = StaticScopeResolver({" payments:read ", "payments:write", "payments:read"}) + + assert resolver.resolve("GET", "https://api.example.com/payments") == frozenset( + {"payments:read", "payments:write"} + ) + assert resolver.all_scopes() == frozenset({"payments:read", "payments:write"}) + + +def test_static_scope_resolver_accepts_any_set_like_collection() -> None: + resolver = StaticScopeResolver({"payments:read": None}.keys()) + + assert resolver.all_scopes() == frozenset({"payments:read"}) + + +def test_static_dpop_key_provider_returns_current_and_named_key() -> None: + private_key = ec.generate_private_key(ec.SECP256R1()) + key_pair = KeyPair(private_key=private_key, public_key=private_key.public_key()) + provider: DPoPKeyProvider = StaticDPoPKeyProvider(key_pair) + dpop_key = provider.get_current_key() + + assert provider.get_current_key() is dpop_key + assert dpop_key.key_id + assert provider.get_key(dpop_key.key_id) is dpop_key + assert provider.get_key("unknown") is dpop_key + + +def test_sync_adapter_preserves_native_resource_types() -> None: + adapter: SyncHttpAdapter[tuple[str, str], object] = NativeAdapter() + request = ("GET", "https://api.example.com") + + response = adapter.send_resource_request(request, {"Authorization": "DPoP token"}) + + assert adapter.request_method(request) == "GET" + assert adapter.request_url(request) == "https://api.example.com" + assert type(response) is object + + +def test_async_adapter_preserves_native_resource_types() -> None: + async def exercise() -> None: + adapter: AsyncHttpAdapter[tuple[str, str], object] = AsyncNativeAdapter() + request = ("GET", "https://api.example.com") + + response = await adapter.send_resource_request(request, {"Authorization": "DPoP token"}) + + assert adapter.request_method(request) == "GET" + assert adapter.request_url(request) == "https://api.example.com" + assert type(response) is object + + asyncio.run(exercise()) + + +def test_in_memory_store_satisfies_protocol_and_normalizes_scope_order() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store: TokenStore = InMemoryTokenStore(now_provider=FixedClock(now)) + token = AccessToken( + client_id="client-id", + token_value="access-token", + scopes=frozenset({"payments:write", "payments:read"}), + expires_at=now + timedelta(minutes=5), + jkt="thumbprint", + ) + + store.put(token) + + assert store.get(AccessTokenFilter(scopes=frozenset({"payments:read", "payments:write"}), jkt="thumbprint")) is token + assert store.get(AccessTokenFilter(scopes=token.scopes)) is token + assert store.get(AccessTokenFilter(scopes=token.scopes, jkt="different")) is None + + +def test_in_memory_store_keeps_bound_token_when_unbound_token_replaces_scope_lookup() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + scopes = frozenset({"payments:read"}) + bound_token = AccessToken("client-id", "bound", scopes, now + timedelta(minutes=5), "thumbprint") + unbound_token = AccessToken("client-id", "unbound", scopes, now + timedelta(minutes=5)) + + store.put(bound_token) + store.put(unbound_token) + + assert store.get(AccessTokenFilter(scopes=scopes, jkt="thumbprint")) is bound_token + assert store.get(AccessTokenFilter(scopes=scopes)) is unbound_token + + +def test_in_memory_store_bound_token_replaces_unbound_scope_lookup() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + scopes = frozenset({"payments:read"}) + unbound_token = AccessToken("client-id", "unbound", scopes, now + timedelta(minutes=5)) + bound_token = AccessToken("client-id", "bound", scopes, now + timedelta(minutes=5), "thumbprint") + + store.put(unbound_token) + store.put(bound_token) + + assert store.get(AccessTokenFilter(scopes=scopes)) is bound_token + assert store.get(AccessTokenFilter(scopes=scopes, jkt="thumbprint")) is bound_token + + +def test_in_memory_store_keeps_different_key_bindings_separate() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + scopes = frozenset({"payments:read"}) + first = AccessToken("client-id", "first", scopes, now + timedelta(minutes=5), "first-jkt") + second = AccessToken("client-id", "second", scopes, now + timedelta(minutes=5), "second-jkt") + + store.put(first) + store.put(second) + + assert store.get(AccessTokenFilter(scopes=scopes, jkt="first-jkt")) is first + assert store.get(AccessTokenFilter(scopes=scopes, jkt="second-jkt")) is second + assert store.get(AccessTokenFilter(scopes=scopes)) is second + + +def test_in_memory_store_replaces_bound_token_with_same_binding_and_scopes() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + scopes = frozenset({"payments:read"}) + first = AccessToken("client-id", "first", scopes, now + timedelta(minutes=5), "thumbprint") + second = AccessToken("client-id", "second", scopes, now + timedelta(minutes=5), "thumbprint") + + store.put(first) + store.put(second) + + assert store.get(AccessTokenFilter(scopes=scopes)) is second + assert store.get(AccessTokenFilter(scopes=scopes, jkt="thumbprint")) is second + + +def test_in_memory_store_replaces_unbound_token_with_same_scopes() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + scopes = frozenset({"payments:read"}) + first = AccessToken("client-id", "first", scopes, now + timedelta(minutes=5)) + second = AccessToken("client-id", "second", scopes, now + timedelta(minutes=5)) + + store.put(first) + store.put(second) + + assert store.get(AccessTokenFilter(scopes=scopes)) is second + + +def test_in_memory_store_supports_empty_scopes() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + token = AccessToken("client-id", "token", frozenset(), now + timedelta(minutes=5), "thumbprint") + + store.put(token) + + assert store.get(AccessTokenFilter(scopes=frozenset())) is token + assert store.get(AccessTokenFilter(scopes=frozenset(), jkt="thumbprint")) is token + + +@pytest.mark.parametrize( + "scopes", + [ + pytest.param(frozenset(), id="empty"), + pytest.param(frozenset({"payments:read"}), id="subset"), + pytest.param(frozenset({"payments:read", "payments:write", "payments:admin"}), id="superset"), + pytest.param(frozenset({"payments:admin"}), id="unrelated"), + ], +) +def test_in_memory_store_rejects_scope_mismatches(scopes: frozenset[str]) -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + token_scopes = frozenset({"payments:read", "payments:write"}) + token = AccessToken("client-id", "token", token_scopes, now + timedelta(minutes=5), "thumbprint") + store.put(token) + + assert store.get(AccessTokenFilter(scopes=scopes)) is None + assert store.get(AccessTokenFilter(scopes=scopes, jkt="thumbprint")) is None + + +def test_in_memory_store_rejects_token_inside_early_expiration_threshold() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + token = AccessToken( + client_id="client-id", + token_value="access-token", + scopes=frozenset({"payments:read"}), + expires_at=now + timedelta(seconds=59), + jkt="thumbprint", + ) + store.put(token) + + assert store.get(AccessTokenFilter(scopes=token.scopes, jkt=token.jkt)) is None + + +def test_in_memory_store_accepts_token_at_early_expiration_boundary() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + store = InMemoryTokenStore(now_provider=FixedClock(now)) + token = AccessToken( + client_id="client-id", + token_value="access-token", + scopes=frozenset({"payments:read"}), + expires_at=now + timedelta(seconds=60), + ) + store.put(token) + + assert store.get(AccessTokenFilter(scopes=token.scopes)) is token + + +def test_in_memory_store_removes_expired_entries_during_put() -> None: + now = datetime(2026, 8, 13, tzinfo=UTC) + clock = FixedClock(now) + store = InMemoryTokenStore(now_provider=clock) + expired = AccessToken( + client_id="client-id", + token_value="expired", + scopes=frozenset({"old"}), + expires_at=now + timedelta(seconds=61), + ) + store.put(expired) + clock.current = now + timedelta(minutes=2) + + store.put( + AccessToken( + client_id="client-id", + token_value="current", + scopes=frozenset({"new"}), + expires_at=clock.current + timedelta(minutes=5), + ) + ) + + assert store.get(AccessTokenFilter(scopes=expired.scopes)) is None diff --git a/tests/unit/test_token.py b/tests/unit/test_token.py new file mode 100644 index 0000000..ac8a462 --- /dev/null +++ b/tests/unit/test_token.py @@ -0,0 +1,201 @@ +import json +from datetime import UTC, datetime, timedelta +from urllib.parse import parse_qs + +import pytest +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from mastercard_oauth2_client import ( + HttpResponse, + KeyPair, + OAuth2Config, + OAuth2Error, + StaticDPoPKeyProvider, + StaticScopeResolver, + build_token_request, + parse_token_response, +) +from mastercard_oauth2_client import token as token_module +from mastercard_oauth2_client._internal.jose import jwk_thumbprint + +_FIXED_NOW = datetime(2026, 8, 28, 12, 0, tzinfo=UTC) + + +def _config() -> OAuth2Config: + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + dpop_private_key = ec.generate_private_key(ec.SECP256R1()) + return OAuth2Config( + client_id="client id", + token_endpoint="https://auth.example.com/oauth/token?tenant=test", + issuer="https://auth.example.com", + client_key=client_key, + key_id="client-key-id", + scope_resolver=StaticScopeResolver({"payments:read"}), + dpop_key_provider=StaticDPoPKeyProvider( + KeyPair(private_key=dpop_private_key, public_key=dpop_private_key.public_key()) + ), + user_agent="test-agent", + ) + + +def _response(payload: object, status: int = 200) -> HttpResponse: + return HttpResponse(status=status, headers={"Content-Type": "application/json"}, body=json.dumps(payload)) + + +def test_build_token_request_contains_client_credentials_assertion_and_dpop_proof() -> None: + config = _config() + dpop_key = config.dpop_key_provider.get_current_key() + + request = build_token_request( + config, + {"payments:write", "payments:read"}, + dpop_key.key_id, + nonce="server-nonce", + ) + + assert request.method == "POST" + assert request.url == config.token_endpoint + assert request.headers == { + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + "DPoP": request.headers["DPoP"], + "User-Agent": "test-agent", + } + assert len(request.headers["DPoP"].split(".")) == 3 + assert isinstance(request.body, str) + form = parse_qs(request.body, keep_blank_values=True) + assert form == { + "grant_type": ["client_credentials"], + "client_id": ["client id"], + "scope": ["payments:read payments:write"], + "client_assertion_type": ["urn:ietf:params:oauth:client-assertion-type:jwt-bearer"], + "client_assertion": [form["client_assertion"][0]], + } + assert len(form["client_assertion"][0].split(".")) == 3 + + +def test_build_token_request_omits_empty_scope() -> None: + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + + request = build_token_request(config, set(), dpop_key_id) + + assert isinstance(request.body, str) + assert "scope" not in parse_qs(request.body, keep_blank_values=True) + + +def test_parse_token_response_returns_cache_ready_dpop_token(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(token_module, "_utc_now", lambda: _FIXED_NOW) + config = _config() + dpop_key = config.dpop_key_provider.get_current_key() + + token = parse_token_response( + config, + _response( + { + "access_token": "access-token", + "token_type": "dPoP", + "expires_in": 900, + "scope": "payments:read payments:write", + } + ), + {"payments:read", "payments:write", "payments:admin"}, + dpop_key.key_id, + ) + + assert token.client_id == config.client_id + assert token.token_value == "access-token" + assert token.scopes == frozenset({"payments:read", "payments:write"}) + assert token.expires_at == _FIXED_NOW + timedelta(seconds=900) + assert token.jkt == jwk_thumbprint(dpop_key.key_pair.public_key) + + +def test_parse_token_response_accepts_fractional_expiry(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(token_module, "_utc_now", lambda: _FIXED_NOW) + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + + token = parse_token_response( + config, + _response({"access_token": "access-token", "token_type": "DPoP", "expires_in": 900.5}), + {"payments:read"}, + dpop_key_id, + ) + + assert token.expires_at == _FIXED_NOW + timedelta(seconds=900.5) + + +def test_parse_token_response_uses_requested_scopes_when_scope_is_omitted( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(token_module, "_utc_now", lambda: _FIXED_NOW) + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + + token = parse_token_response( + config, + _response({"access_token": "access-token", "token_type": "DPoP", "expires_in": 900}), + {"payments:read"}, + dpop_key_id, + ) + + assert token.scopes == frozenset({"payments:read"}) + + +@pytest.mark.parametrize( + ("payload", "message"), + [ + ( + {"token_type": "DPoP", "expires_in": 900}, + "Missing value in access token response: access_token", + ), + ( + {"access_token": "token", "expires_in": 900}, + "Missing value in access token response: token_type", + ), + ( + {"access_token": "token", "token_type": "DPoP"}, + "Missing value in access token response: expires_in", + ), + ( + {"access_token": "token", "token_type": "Bearer", "expires_in": 900}, + "Expected DPoP token type but received: Bearer", + ), + ], +) +def test_parse_token_response_requires_dpop_token_fields(payload: object, message: str) -> None: + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + + with pytest.raises(OAuth2Error, match=message): + parse_token_response(config, _response(payload), set(), dpop_key_id) + + +@pytest.mark.parametrize("expires_in", [0, -1, True, "900", None]) +def test_parse_token_response_requires_positive_numeric_expiry(expires_in: object) -> None: + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + payload = {"access_token": "token", "token_type": "DPoP", "expires_in": expires_in} + + with pytest.raises(OAuth2Error, match="FAPI 2.0 requires valid expires_in field"): + parse_token_response(config, _response(payload), set(), dpop_key_id) + + +def test_parse_token_response_rejects_unsuccessful_or_malformed_response() -> None: + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + + with pytest.raises(OAuth2Error, match="Token request failed with HTTP 400"): + parse_token_response(config, _response({"error": "invalid_client"}, status=400), set(), dpop_key_id) + + with pytest.raises(OAuth2Error, match="Failed to parse JSON access token response"): + parse_token_response(config, HttpResponse(status=200, headers={}, body="not-json"), set(), dpop_key_id) + + +@pytest.mark.parametrize("body", [None, "", " ", b""]) +def test_parse_token_response_rejects_empty_body(body: bytes | str | None) -> None: + config = _config() + dpop_key_id = config.dpop_key_provider.get_current_key().key_id + + with pytest.raises(OAuth2Error, match="Empty access token response"): + parse_token_response(config, HttpResponse(status=200, headers={}, body=body), set(), dpop_key_id)