diff --git a/EXAMPLES.md b/EXAMPLES.md index 5a10f4a..074481f 100644 --- a/EXAMPLES.md +++ b/EXAMPLES.md @@ -60,6 +60,35 @@ asyncio.run(exchange_on_behalf_of()) In the current implementation, `get_token_on_behalf_of()` forwards the incoming access token as the [RFC 8693](https://datatracker.ietf.org/doc/html/rfc8693#section-2.1) `subject_token` and relies on Auth0 to handle any DPoP-specific behavior for that token. +### Caching the Exchanged Token + +To cache the exchanged token, pass a `token_store` when constructing `ApiClient`. Caching activates automatically once a store is configured. The SDK reads `sub` from the incoming token to build the cache key, so no additional argument is needed on each call. + +```python +from auth0_api_python import ApiClient, ApiClientOptions + +# token_store is your AbstractTokenStore implementation (e.g. Redis-backed). +# See docs/TokenStorage.md for how to build one. +api_client = ApiClient(ApiClientOptions( + domain="your-tenant.auth0.com", + audience="https://mcp-server.example.com", + client_id="", + client_secret="", + token_store=your_token_store, +)) + +claims = await api_client.verify_access_token(access_token=incoming_access_token) + +result = await api_client.get_token_on_behalf_of( + access_token=incoming_access_token, + audience="https://calendar-api.example.com", + scope="calendar:read calendar:write", +) +``` + +See the **[Token Storage Guide](docs/TokenStorage.md)** for a full working example, how to +implement a Redis-backed store, and the built-in encryption helpers. + ## Inspecting Delegation After Token Verification When a downstream API or `MCP` server receives an access token that may have been issued through diff --git a/README.md b/README.md index 932ddd2..6cc400a 100644 --- a/README.md +++ b/README.md @@ -251,6 +251,10 @@ token as the `subject_token` and relies on Auth0 to handle any DPoP-specific beh The OBO result only includes access-token-oriented fields. It does not expose `id_token` or `refresh_token`. +Configuring a `token_store` on `ApiClientOptions` caches the exchanged token, so a repeat call for +the same caller, audience, organization, scopes, and session reuses it instead of exchanging again. +With no `token_store`, every call performs a fresh exchange, matching the existing behavior above. + #### Inspecting Delegation After Token Verification When a downstream API or `MCP` server receives an access token that may have been issued through @@ -407,6 +411,7 @@ For hybrid mode (migration scenarios), resolver patterns, error handling, and ca - **[Multi-Custom Domain Guide](docs/MultipleCustomDomain.md)** - Configuration modes, resolver patterns, migration, error handling - **[Caching Guide](docs/Caching.md)** - Cache tuning, custom adapters (Redis, Memcached) +- **[Token Storage Guide](docs/TokenStorage.md)** - Caching OBO exchanges, custom TokenStore backends, at-rest encryption ### 8. Anonymous Callers diff --git a/docs/TokenStorage.md b/docs/TokenStorage.md new file mode 100644 index 0000000..298d81c --- /dev/null +++ b/docs/TokenStorage.md @@ -0,0 +1,188 @@ +# Token Storage + +The SDK can cache access tokens it mints on the caller's behalf. Currently this covers tokens +returned by `get_token_on_behalf_of()`. This is separate from the `CacheAdapter` described in the +[Caching Guide](Caching.md), which only caches OIDC discovery metadata and JWKS keys, never a live +bearer token. + +## Default Behavior + +Caching is disabled by default. Without a `token_store`, every call to `get_token_on_behalf_of()` +performs a fresh exchange and nothing is stored. + +To enable caching, pass a `token_store` to `ApiClientOptions`. Once a store is configured, the SDK +automatically builds a cache key from the incoming token and no additional argument is needed per +call. + +## On Behalf Of Exchange with Caching + +The following example verifies an incoming token and exchanges for a downstream token. The result +is cached so a second call for the same caller, audience, org, and scopes returns the cached token +without hitting Auth0 again. + +```python +import asyncio +import httpx + +from auth0_api_python import ApiClient, ApiClientOptions + +async def exchange_on_behalf_of_cached(your_token_store): + api_client = ApiClient(ApiClientOptions( + domain="your-tenant.auth0.com", + audience="https://mcp-server.example.com", + client_id="", + client_secret="", + token_store=your_token_store, + )) + + incoming_access_token = "incoming-auth0-access-token" + + # With a store configured, the exchange verifies the incoming token itself, so there is + # no need to call verify_access_token separately here. + result = await api_client.get_token_on_behalf_of( + access_token=incoming_access_token, + audience="https://calendar-api.example.com", + scope="calendar:read calendar:write", + ) + + async with httpx.AsyncClient() as client: + downstream_response = await client.get( + "https://calendar-api.example.com/events", + headers={"Authorization": f"Bearer {result['access_token']}"} + ) + + downstream_response.raise_for_status() + return downstream_response.json() +``` + +The cached entry is scoped to the caller (`sub`), the target `audience`, the `org_id`, the +scopes it was granted, and the session the incoming token belongs to. Different callers or +sessions never share a cached token. + +## Implementing a Token Store + +To supply a store, subclass `AbstractTokenStore` and implement three async methods: `get`, `set`, +and `delete`. The base class requires a `secret` at construction and provides `encrypt` and +`decrypt` helpers that your methods can call to protect tokens at rest. + +### Redis example + +```python +import time +from typing import Optional + +import redis.asyncio as redis + +from auth0_api_python import AbstractTokenStore, TokenSet + + +class RedisTokenStore(AbstractTokenStore): + def __init__(self, redis_client, *, secret: str): + super().__init__(secret=secret) + self.redis = redis_client + + async def get(self, key: str) -> Optional[TokenSet]: + raw = await self.redis.get(key) + if raw is None: + return None + return self.decrypt(key, raw) + + async def set(self, key: str, value: TokenSet) -> None: + encrypted = self.encrypt(key, value) + ttl = max(value["expires_at"] - int(time.time()), 0) + await self.redis.set(key, encrypted, ex=ttl) + + async def delete(self, key: str) -> None: + await self.redis.delete(key) + + +# Usage +redis_client = redis.Redis(host="localhost", port=6379, db=0) + +api_client = ApiClient(ApiClientOptions( + domain="your-tenant.auth0.com", + audience="https://mcp-server.example.com", + client_id="", + client_secret="", + token_store=RedisTokenStore(redis_client, secret=""), +)) +``` + +### Encryption + +`self.encrypt(key, value)` and `self.decrypt(key, data)` are provided by `AbstractTokenStore`. +They use HKDF-SHA256 to derive a per-entry encryption key from `secret` and the cache key, then +wrap the token in a JWE using `alg: dir` and `enc: A256CBC-HS512`. A fresh random `kid` is +generated on every write, so two encryptions of the same value produce different ciphertext. + +`secret` must be kept outside your codebase, for example in an environment variable or a secrets +manager. Rotating it invalidates all existing cached entries, which is safe because the SDK falls +back to a fresh exchange on any cache miss or decryption failure. + +## Matching cached tokens by scope + +By default the SDK only reuses a cached token when a later call asks for exactly the same scopes. +This is the `strict` setting of `scope_matching` on `ApiClientOptions`, and it works with any +store that implements `get`, `set`, and `delete`. + +Set `scope_matching="non_strict"` when a broader token should satisfy a narrower request. If an +earlier exchange was granted `calendar:read calendar:write` and a later call only needs +`calendar:read`, non_strict returns the cached token instead of exchanging again, because the +granted scopes already cover what was asked for. + +To do that without one scope set evicting another, non_strict keeps every distinct token plus an +index of the scopes each one was granted, so any cached token whose scopes cover the request can +be reused. The index is maintained with an atomic add so that several server processes can write +to it at once without losing each other's entries. A plain `AbstractTokenStore` cannot offer that, +so non_strict requires a store that subclasses `IndexedTokenStore` and implements +`add_index_member` and `list_index_members`. Using non_strict with a plain store raises +`ConfigurationError` at construction. + +### Redis IndexedTokenStore example + +This extends the `RedisTokenStore` above and backs the index with a Redis hash, one field per +token. Adding a member is a single `HSET`, which is atomic per field and overwrites any member +stored under the same `token_key`. + +```python +import json +import time + +from auth0_api_python import IndexedTokenStore, TokenIndexMember + + +class RedisIndexedTokenStore(RedisTokenStore, IndexedTokenStore): + async def add_index_member(self, index_key: str, member: TokenIndexMember) -> None: + await self.redis.hset(index_key, member["token_key"], json.dumps(member)) + + async def list_index_members(self, index_key: str) -> list[TokenIndexMember]: + raw = await self.redis.hgetall(index_key) + now = int(time.time()) + live: list[TokenIndexMember] = [] + expired_fields = [] + for field, value in raw.items(): + member = json.loads(value) + if member["expires_at"] > now: + live.append(member) + else: + expired_fields.append(field) + # Drop expired fields so the hash does not grow without bound as tokens age out. + if expired_fields: + await self.redis.hdel(index_key, *expired_fields) + return live + + +# Usage +api_client = ApiClient(ApiClientOptions( + domain="your-tenant.auth0.com", + audience="https://mcp-server.example.com", + client_id="", + client_secret="", + token_store=RedisIndexedTokenStore(redis_client, secret=""), + scope_matching="non_strict", +)) +``` + +An index member holds a hashed `token_key`, the granted scopes, and an expiry, never a bearer +token, so it is not encrypted. The tokens themselves stay encrypted under their own keys as +described above. diff --git a/poetry.lock b/poetry.lock index ee84389..c874bc7 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.3.4 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.4.1 and should not be changed by hand. [[package]] name = "ada-url" @@ -766,58 +766,58 @@ toml = ["tomli ; python_full_version <= \"3.11.0a6\""] [[package]] name = "cryptography" -version = "50.0.0" +version = "50.0.1" description = "cryptography is a package which provides cryptographic recipes and primitives to Python developers." optional = false python-versions = "!=3.9.0,!=3.9.1,>=3.9" groups = ["main"] files = [ - {file = "cryptography-50.0.0-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:031e2d5dd4bb9caa3ca9c82e5a197fd8ae680232cee62603d1a813f3f07e3d03"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:fd9192b7b70c573d7f214eb1ae35e00d359f6f5e4b27c7e21e30de1fc6204645"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:06a32a980526a6ab9a4b9bf8f7385800791e2bb960903cb6b530e4817509a3b7"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:a1b30560f2acc95aa8b2e06e716a13dbfc97314747b80d9707e307f77b40d6b3"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:8d89f3976b10b4ce31118de72329025f70d2c6ead14a8217c5514dd2c6d5a78f"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:b42a28c1844fd9de8f3f7d540e36b66f3a9c83fceac7170ebc7a6a19edd9dcae"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:900131fafd8aead39ac7dd3a7e833be754c17a95cfd91221636949fe4eb0aa8a"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:07949c449a1abcf60d1ee6e88956d89404c7df3c8258f46589e912988e551987"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:f89831ef99dd7dd169ab06d63a831adb9e20a87aac6d380266bbda5823349169"}, - {file = "cryptography-50.0.0-cp311-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:82148ec5bddac30b51a5b3c1945075f896fa022cb93f8e4a01e9f6ee95292c5f"}, - {file = "cryptography-50.0.0-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:1489e263a8048bb8b6a8bac662eb2d402ea5d2b7b4699b72f385f1e2772db105"}, - {file = "cryptography-50.0.0-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:7cec5b856506da6defb290f30c9ee687d5f5e8cb0bd3f6459dde43b0b4fa40ef"}, - {file = "cryptography-50.0.0-cp311-abi3-win_amd64.whl", hash = "sha256:bd1c592e4d5974f0d08d4888e432157adba757c66da0246918e43677fafa2d30"}, - {file = "cryptography-50.0.0-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:49e7d93abdbd2990caced757e5fade25302f719c3c8fb6e6fff2dde98999fc41"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:19736989797678c6af1e55cd49055cdbcb55d8f6b5583ac5335f933aba9101dc"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:80b63928fa35083b33966ce1efb70e5b9607181e49dcd1c22c8c005e319f667f"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_28_aarch64.whl", hash = "sha256:d58c3db7cd6eed54e6c06744db55456b65ebd7492ddeae9c1e93cfca7aa857d3"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_28_ppc64le.whl", hash = "sha256:df2a58a472f332225671c35b0a830208b86d004f82baa8530fa3782c85646533"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_28_x86_64.whl", hash = "sha256:11b74db56cdbe3cdee6e3f6982ecb70334fa10dce99ed58bf7894aaaa3b2a037"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_31_armv7l.whl", hash = "sha256:f59e38625469987d7ef6d495323c55e7db6c212eaf6112267e0d3b565a2e9c9f"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_34_aarch64.whl", hash = "sha256:ecfed7367f965a0328cfbdd70da860f15441f002f613185668c6e6ebf5a0ac11"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_34_ppc64le.whl", hash = "sha256:9aa87839c383bdbab6ef865787a1fb877af8dd03464c4400322726feaaadfc6d"}, - {file = "cryptography-50.0.0-cp314-cp314t-manylinux_2_34_x86_64.whl", hash = "sha256:6ba6a53445bd3cfa809ef3ef5f1589aa6ba08784a1d962bf47d0940e871dab1c"}, - {file = "cryptography-50.0.0-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:3f5735ffe4996d28b809371756219f5354864902a3b9e7c0b9ee87041209fc9c"}, - {file = "cryptography-50.0.0-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:1b4a266766514614f8aa60416e71f2fc6e575d36e7bdc90f644fadb2f4b75b95"}, - {file = "cryptography-50.0.0-cp314-cp314t-win_amd64.whl", hash = "sha256:12b9c6996425c76ea6c457ace4f3073e715b8c545add07cd1a8f3a4f90691269"}, - {file = "cryptography-50.0.0-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:ccdc4a71a4dabae05de219404f9f4abc38e3b58422177ff93d0da05967dafa07"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:910e1d2668e7de9648f2bcee30e180db2a6b15c30f887d7c4c93ddf96e3992e3"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:a91296cb61e8df6f86d0c19cc4068228da256bf59bf86049fbd821084565327f"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:e722f16708d854fe924790e051061f6704a472c3bac347b6fd88033ea8dd0dc5"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:d764dcf130c428ef66786f866dd750f53182bc608813489915e9fc106bb0c82f"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:105110f43a471dbd0060b9c9516cb8a6a79233631a04cc2ba16f28323ac6e025"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:828743d939e9629bc267b8e2d08d8bb67cd4319c771a33d4b18b22dd8fb7440a"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:2a8183b489dc1f7f80f135780fadc1108f14b31b8a40411c7a5b17425f65f28b"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:6e7d61120573a7f2cd94cc095f9e81f6967c61ccdf194285aa143ecec8e0b708"}, - {file = "cryptography-50.0.0-cp39-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:37fdb0d0111f1e2ff07139dfb79f1b49531f8e213c46f1163dd7642979b58c47"}, - {file = "cryptography-50.0.0-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:c87f62a3d3b9888ed0fdde100ec06aa61ca9cd44bad9057d1dff9a516b5f5bb9"}, - {file = "cryptography-50.0.0-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:65c2c3add92b45fd0709db8594536aea39c2a67af0e27ffcf049c498501140b7"}, - {file = "cryptography-50.0.0-cp39-abi3-win_amd64.whl", hash = "sha256:d24fead1d4d076e1bfb006dcec392074a3cd8d7b4fc8a595aa64073b2b7a96ba"}, - {file = "cryptography-50.0.0-pp311-pypy311_pp73-macosx_11_0_arm64.whl", hash = "sha256:5e34edd123674534acd70147f0ca331eaa2c74e6325fb2028c886aa26ba0b68c"}, - {file = "cryptography-50.0.0-pp311-pypy311_pp73-manylinux_2_28_aarch64.whl", hash = "sha256:8eb5e1172eb569ea8a872796576e6a67c276351728b6455d5beb01242b027c6a"}, - {file = "cryptography-50.0.0-pp311-pypy311_pp73-manylinux_2_28_x86_64.whl", hash = "sha256:910d11e1a385c654bf738bf3e6b8e6ed5de0f5610fcae2be9e5b398d8081d20e"}, - {file = "cryptography-50.0.0-pp311-pypy311_pp73-manylinux_2_34_aarch64.whl", hash = "sha256:62598a8a57f815db4c6259a4e97d857dab56697e7de8e8ab02352ab74da1995d"}, - {file = "cryptography-50.0.0-pp311-pypy311_pp73-manylinux_2_34_x86_64.whl", hash = "sha256:07479a1cb08219ab719147e742e76090c9c773321959bb94946fffdd397a6437"}, - {file = "cryptography-50.0.0-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:c99c003e088647b8a5b7c145d6f78c335f6348332b62e142d411c4b63d1460b9"}, - {file = "cryptography-50.0.0.tar.gz", hash = "sha256:eeac2acb5a20ed25e0ad6d1df9891a520b78b404266b6d11778f25d5d691a6c9"}, + {file = "cryptography-50.0.1-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:b8f852c65863251b9e3a1b8c150ce21e59b522dbb6a7d4bc80e680d38388e986"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:53e279950892dc102c6b4e52af03ae5ea92fac572a1ddab78ca73a997f62b69f"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:ff838d62ec1bfce4f9ba7fa16f4a7b554cd8d0c299e6be37502161a660c84eef"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:e74591e283fe6eb956416c929eb58262a719fe0311fd9054c62c3350ed8760d8"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:5fe002589592ed749ce77fe0695fcbd3500dd61d7d6db5858a7544c612fa8e45"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:51593d180cf6d179bde5c5d065bed81386b1f381656ae7d042b7ffc87a9895ad"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:359e62deae718bce96170e223fdcb6357e4fbd3bb7a3a75f4430763532560e49"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:e2ca8fd1b6b4b82a1c4cb02841d0837e3c12336c2e24b520ab8ab3b969733d8f"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:76de83fbd91ac49c0feaaa983d0748fd7a53176afac5fb3bf7478d244f0eb527"}, + {file = "cryptography-50.0.1-cp311-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:51afcfceb15597cf2635068e4ac9a56b2abde622edde17f37d85fd7b5306497a"}, + {file = "cryptography-50.0.1-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:be224a65493ec5b74a158ff22a5522ce4a5ca1e543c647a3a4730d4a09e5f959"}, + {file = "cryptography-50.0.1-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:9ebcdd5519be9b652a46f507817a74591774fc3d6923ac364e4dfa64e36b291b"}, + {file = "cryptography-50.0.1-cp311-abi3-win_amd64.whl", hash = "sha256:aed8db4f6d71c51efb89530e12d9464e7bf2923d46c3205dc794a2a93f8c0648"}, + {file = "cryptography-50.0.1-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:30a125032e5642a21ff816e021152bd4e7e94f03eff3f4b7fca41cd22bc3110f"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:a0b1a59e3a089064a0ec309e9428c8e3ae4e161419d20ac33600767e83fc658a"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:8921d58f426793c5f1b47f0b59575780de9a095214958d0eb37d909593db8367"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_28_aarch64.whl", hash = "sha256:a8f40ea47330e71b594a7e246898f93177c259490c63183dbaf9e571d71ed9a5"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_28_ppc64le.whl", hash = "sha256:a255449073358275b64b67d3f595f268bbef70e72b6edb65e0c70c735bf739c9"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_28_x86_64.whl", hash = "sha256:8df2de9102026855887e4587084f6eabd80ed0f345b8ad8a7ac27ab9bf4723e0"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_31_armv7l.whl", hash = "sha256:ac02b07824d4d1001bd4367599f839c19cb171924c796e52c23508ac14c2c0cc"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_34_aarch64.whl", hash = "sha256:cbf74a81765ee67413503ca6e26dcc4f6f5a519822436cc0a1b97aab6c1b8a17"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_34_ppc64le.whl", hash = "sha256:16c5ecd954b3330ebfb6605eca4fd952da8bef376551d5cc264534e3770a9ee6"}, + {file = "cryptography-50.0.1-cp314-cp314t-manylinux_2_34_x86_64.whl", hash = "sha256:79bf008d1f9af6071c797ad133e39915dfee7614f18f18f4db9072eb715064a3"}, + {file = "cryptography-50.0.1-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:330fbb252391c596f1ae42c5754449dc924e6ad012dca8efe0d703f9f2d12ec6"}, + {file = "cryptography-50.0.1-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:42be3bb70596b3abe4ac097b75be223e8b3ab614a0e5de068e3dcc54d71d6149"}, + {file = "cryptography-50.0.1-cp314-cp314t-win_amd64.whl", hash = "sha256:f74455bb086a85d5e81246412602aaa97ed095e504cd40dd261ef50be42205bf"}, + {file = "cryptography-50.0.1-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:ca83d00d9e69cd5eb63f2e69c3a5a59e0cecae5ae14c6ae0b35830fe3b37bad0"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:05ba322c4da95b262a212c345af888ef2c37c88c0509756ea00a0e6d68850f23"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:e22dfed744bd4002e909464cb23d2f0b05c6f3113a79ef2e9864a53db737c733"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:4c4188f7c0cf655be5c06342b817ed0f9595b69ffa2b12026e5353eed29dea88"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:2ebbfb0f1fed745e91796e3e1080a1440423fdae8ece1b995a1d80883a409054"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:407fe2b6db00939c05c0e945e9914238f2f0a430974839429dafc82b1ee6bee5"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:2b34d76a652ea2b6faf777c35df230c5637842cd904e04f16230c3f9f03e4361"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:01f41478cf33fc605a6a089cd56d28b45c6c0b45a1928b61797f2621a04bac71"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:fc3ed7ebd2a8c96f5b166de0ab9b624996bef3b07bbeb19364dfb78222c22c80"}, + {file = "cryptography-50.0.1-cp39-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:9dde0a357190eb3b1da1bb9ab750e9c85cba82ca5977aa0836cbb94e92611239"}, + {file = "cryptography-50.0.1-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:fd3718b960d0b5dd213cdf03f3bcb7000e69dda0de8b956061947ff6bcff5558"}, + {file = "cryptography-50.0.1-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:2a93d05e34d5f67fba6f891fe85d929999baa7195e853923ea6d7576c9e68c5e"}, + {file = "cryptography-50.0.1-cp39-abi3-win_amd64.whl", hash = "sha256:55d16b1ef3ee0958d893a977b19777887e546c9954ea81b200c3301a864013f2"}, + {file = "cryptography-50.0.1-pp311-pypy311_pp73-macosx_11_0_arm64.whl", hash = "sha256:9cb3cb952cf5a8abd50c782a98a89d71699715e802fe349704b47f2425b42a94"}, + {file = "cryptography-50.0.1-pp311-pypy311_pp73-manylinux_2_28_aarch64.whl", hash = "sha256:5fe939deeb161024a6be98229c953b6591fef1f41214497a78fe793a244c017f"}, + {file = "cryptography-50.0.1-pp311-pypy311_pp73-manylinux_2_28_x86_64.whl", hash = "sha256:fb4b9672d389c738b175c4166e78310f8a70358886aacd9173ee03a85ffdc671"}, + {file = "cryptography-50.0.1-pp311-pypy311_pp73-manylinux_2_34_aarch64.whl", hash = "sha256:d63ae8f6481fec907ac0f588eee8a90aefde112c633131fe540e5711ddbb5a4e"}, + {file = "cryptography-50.0.1-pp311-pypy311_pp73-manylinux_2_34_x86_64.whl", hash = "sha256:804728ce710890870f3aaa344b2e161172d258d768ac139d02cfd9092d0d94e6"}, + {file = "cryptography-50.0.1-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:693c99b49bd37d0d096e4334c10232c77248c415b98d35236094cdf96d57258b"}, + {file = "cryptography-50.0.1.tar.gz", hash = "sha256:5dd9bda1c12b4162f6ff568eeb5e0ff956c28d14406e875cfe8a63a2d414ff20"}, ] [package.dependencies] @@ -980,6 +980,22 @@ cryptography = ">=45.0.1" [package.extras] drafts = ["pycryptodome"] +[[package]] +name = "jwcrypto" +version = "1.6.1" +description = "Implementation of JOSE Web standards" +optional = false +python-versions = ">=3.9" +groups = ["main"] +files = [ + {file = "jwcrypto-1.6.1-py3-none-any.whl", hash = "sha256:77e856818e514cb3d64a1043b6885a7d52f5cc37773f5e944de02103c8d3a54f"}, + {file = "jwcrypto-1.6.1.tar.gz", hash = "sha256:a1e1570da5c2e35dbcd375ec1d2891a24de8eb35ee2c1fe8e91978efb72cfe90"}, +] + +[package.dependencies] +cryptography = ">=49.0.0" +typing_extensions = ">=4.5.0" + [[package]] name = "packaging" version = "26.2" @@ -1314,11 +1330,11 @@ description = "Backported and Experimental Type Hints for Python 3.9+" optional = false python-versions = ">=3.9" groups = ["main", "dev"] -markers = "python_version < \"3.13\"" files = [ {file = "typing_extensions-4.16.0-py3-none-any.whl", hash = "sha256:481caa481374e813c1b176ada14e97f1f67a4539ce9cfeb3f350d78d6370c2e8"}, {file = "typing_extensions-4.16.0.tar.gz", hash = "sha256:dc983d19a509c94dba722ee6abd33940f7c05a89e243c47e907eb4db6f1a43e5"}, ] +markers = {dev = "python_version < \"3.13\""} [[package]] name = "urllib3" @@ -1361,4 +1377,4 @@ zstd = ["backports-zstd (>=1.0.0) ; python_version < \"3.14\""] [metadata] lock-version = "2.1" python-versions = ">=3.9.2" -content-hash = "c623140ac5cc4dd76691c916dd78a2226af163ed8a84f045f6a9a57fe9f387bf" +content-hash = "f78b7a13854c1a3aedd7ae7d4c4d0685dbeb549a5d847abd08bb98561f043b9d" diff --git a/pyproject.toml b/pyproject.toml index 48755ac..75b7fb8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,6 +15,8 @@ python = ">=3.9.2" authlib = "^1.0" # For JWT/OIDC features requests = "^2.31.0" # If you use requests for HTTP calls (e.g., discovery) httpx = "^0.28.1" +cryptography = ">=43.0.1" +jwcrypto = "^1.5.7" ada-url = [ {version = "^1.30.0", python = ">=3.10"}, {version = "^1.27.0", python = ">=3.9,<3.10"} diff --git a/src/auth0_api_python/__init__.py b/src/auth0_api_python/__init__.py index cc77b05..b6a8c6e 100644 --- a/src/auth0_api_python/__init__.py +++ b/src/auth0_api_python/__init__.py @@ -14,6 +14,14 @@ ConfigurationError, DomainsResolverError, GetTokenByExchangeProfileError, + TokenStoreError, +) +from .token_store import ( + AbstractTokenStore, + IndexedTokenStore, + TokenIndexMember, + TokenSet, + VerifiedToken, ) from .types import ( DomainsResolver, @@ -31,8 +39,14 @@ "DomainsResolverContext", "DomainsResolverError", "GetTokenByExchangeProfileError", + "TokenStoreError", "get_current_actor", "get_delegation_chain", "InMemoryCache", + "AbstractTokenStore", + "IndexedTokenStore", + "TokenIndexMember", + "VerifiedToken", "OnBehalfOfTokenResult", + "TokenSet", ] diff --git a/src/auth0_api_python/_internal/__init__.py b/src/auth0_api_python/_internal/__init__.py new file mode 100644 index 0000000..bcd5606 --- /dev/null +++ b/src/auth0_api_python/_internal/__init__.py @@ -0,0 +1 @@ +"""Internal machinery for auth0-api-python. Not part of the public API.""" diff --git a/src/auth0_api_python/_internal/cache_keys.py b/src/auth0_api_python/_internal/cache_keys.py new file mode 100644 index 0000000..85debc9 --- /dev/null +++ b/src/auth0_api_python/_internal/cache_keys.py @@ -0,0 +1,146 @@ +"""Cache-key builders and scope helpers for the tokens the SDK mints.""" + +import hashlib +from typing import Optional + +from ..utils import get_unverified_payload + +_FIELD_SEPARATOR = "\x1f" + + +def _normalized_scopes(scope: Optional[str]) -> str: + """Sort and dedupe scopes so equivalent scope strings produce the same cache key.""" + if not scope: + return "" + return " ".join(sorted(set(scope.split()))) + + +def is_covered_by(requested: Optional[str], granted: Optional[str]) -> bool: + """True when the requested scopes are all covered by the granted scopes.""" + granted_set = set(_normalized_scopes(granted).split()) + requested_set = set(_normalized_scopes(requested).split()) + return requested_set.issubset(granted_set) + + +def session_fingerprint(access_token: str) -> str: + """ + Derive a session-scoped fingerprint for an access token, to keep concurrent + sessions' cached tokens from colliding. Not a trust or authorization signal: + the token is decoded without signature verification. + """ + try: + payload = get_unverified_payload(access_token) + except ValueError: + return hashlib.sha256(access_token.encode()).hexdigest() + + sid = payload.get("sid") + if isinstance(sid, str) and sid: + return sid + + jti = payload.get("jti") + if isinstance(jti, str) and jti: + return jti + + return hashlib.sha256(access_token.encode()).hexdigest() + + +def _obo_identity_fields( + namespace: str, + *, + issuer: str, + incoming_client_id: str, + exchange_tenant: str, + exchange_client_id: str, + audience: str, + org_id: Optional[str], + session_key: str, +) -> list[str]: + """The fields every OBO cache key shares, namespaced by key kind. + + Folds in the verified issuer, the incoming token's client, the exchange tenant, and the + exchange client so a token minted under one of them is never read back under another. + """ + return [ + namespace, + issuer, + incoming_client_id, + exchange_tenant, + exchange_client_id, + audience, + org_id or "", + session_key, + ] + + +def obo_cache_key( + *, + layout: str, + sub: str, + issuer: str, + incoming_client_id: str, + exchange_tenant: str, + exchange_client_id: str, + audience: str, + org_id: Optional[str], + scopes: Optional[str], + session_key: str, +) -> str: + """Token cache key for an On Behalf Of exchange. + + layout is "strict" or "index" so the two scope-matching modes never read each other's tokens. + The normalized scopes are folded in. + """ + fields = _obo_identity_fields( + "obo_token_" + layout, + issuer=issuer, + incoming_client_id=incoming_client_id, + exchange_tenant=exchange_tenant, + exchange_client_id=exchange_client_id, + audience=audience, + org_id=org_id, + session_key=session_key, + ) + fields.append(_normalized_scopes(scopes)) + return sub + ":" + hashlib.sha256(_FIELD_SEPARATOR.join(fields).encode()).hexdigest() + + +def index_cache_key( + *, + sub: str, + issuer: str, + incoming_client_id: str, + exchange_tenant: str, + exchange_client_id: str, + audience: str, + org_id: Optional[str], + session_key: str, +) -> str: + """Cache key for the OBO token index, scoped to caller, issuer, exchange, audience, org, and session. + + The "obo_index" namespace keeps it from colliding with either layout's token cache key. + """ + fields = _obo_identity_fields( + "obo_index", + issuer=issuer, + incoming_client_id=incoming_client_id, + exchange_tenant=exchange_tenant, + exchange_client_id=exchange_client_id, + audience=audience, + org_id=org_id, + session_key=session_key, + ) + return sub + ":" + hashlib.sha256(_FIELD_SEPARATOR.join(fields).encode()).hexdigest() + + +def m2m_cache_key(*, tenant: str, client_id: str, audience: str, scopes: Optional[str]) -> str: + """Cache key for a client-credentials (M2M) exchange, scoped to tenant, client, audience, and scopes.""" + fields = _FIELD_SEPARATOR.join( + ["m2m", tenant, client_id, audience, _normalized_scopes(scopes)] + ) + return hashlib.sha256(fields.encode()).hexdigest() + + +def token_vault_cache_key(sub: str, connection: str) -> str: + """Cache key for a Token Vault exchange, scoped to caller and connection.""" + fields = _FIELD_SEPARATOR.join(["token_vault", sub, connection]) + return hashlib.sha256(fields.encode()).hexdigest() diff --git a/src/auth0_api_python/_internal/obo_cache.py b/src/auth0_api_python/_internal/obo_cache.py new file mode 100644 index 0000000..111f0b7 --- /dev/null +++ b/src/auth0_api_python/_internal/obo_cache.py @@ -0,0 +1,221 @@ +"""Caching for On Behalf Of exchanges: cache-key identity, lookup, and write.""" + +import logging +import time +from typing import NamedTuple, Optional, cast + +from ..errors import TokenStoreError +from ..token_store import ( + AbstractTokenStore, + IndexedTokenStore, + TokenSet, + VerifiedToken, +) +from ..types import OnBehalfOfTokenResult +from .cache_keys import ( + _normalized_scopes, + index_cache_key, + is_covered_by, + obo_cache_key, + session_fingerprint, +) + + +class _OboIdentity(NamedTuple): + """The verified fields that identify one caller's OBO cache entries.""" + + sub: str + issuer: str + incoming_client_id: str + org_id: Optional[str] + session_key: str + + +class OboCache: + """Caches On Behalf Of exchange results in a token store, keyed by verified caller identity. + + scope_matching controls reuse. "strict" reuses a cached token only when its granted scopes + equal the request. "non_strict" reuses any cached token whose granted scopes cover the + request, which needs the per-caller index and so an IndexedTokenStore. + """ + + def __init__( + self, + token_store: AbstractTokenStore, + *, + exchange_tenant: str, + exchange_client_id: str, + scope_matching: str, + ) -> None: + self._store = token_store + self._exchange_tenant = exchange_tenant + self._exchange_client_id = exchange_client_id + self._mode = scope_matching + + def identity(self, verified: VerifiedToken) -> Optional[_OboIdentity]: + """Build the cache identity from verified claims, or None when they lack a usable subject.""" + claims = verified.claims + sub = claims.get("sub") + if not isinstance(sub, str) or not sub: + logging.warning("Verified token has no usable sub claim, skipping cache") + return None + issuer = claims.get("iss") + if not isinstance(issuer, str) or not issuer: + logging.warning("Verified token has no iss claim, skipping cache") + return None + org_id = claims.get("org_id") if isinstance(claims.get("org_id"), str) else None + incoming = claims.get("azp") or claims.get("client_id") + incoming_client_id = incoming if isinstance(incoming, str) else "" + return _OboIdentity( + sub=sub, + issuer=issuer, + incoming_client_id=incoming_client_id, + org_id=org_id, + session_key=session_fingerprint(verified.access_token), + ) + + async def lookup( + self, identity: _OboIdentity, audience: str, scope: Optional[str] + ) -> Optional[OnBehalfOfTokenResult]: + """Return a reusable cached token for this request, or None on a miss.""" + if self._mode == "non_strict": + return await self._index_lookup(identity, audience, scope) + return await self._strict_lookup(identity, audience, scope) + + async def write( + self, + identity: _OboIdentity, + audience: str, + scope: Optional[str], + obo_result: OnBehalfOfTokenResult, + ) -> None: + """Store a freshly exchanged token under its granted scopes.""" + # When Auth0 does not echo the granted scope, the requested scope is the best label available. + granted = _normalized_scopes(obo_result["scope"] if "scope" in obo_result else scope) + entry: TokenSet = { + "access_token": obo_result["access_token"], + "expires_at": obo_result["expires_at"], + "granted_scopes": granted, + } + try: + if self._mode == "non_strict": + token_key = self._token_key(identity, audience, "index", granted) + await self._store.set(token_key, entry) + await cast(IndexedTokenStore, self._store).add_index_member( + self._index_key(identity, audience), + {"token_key": token_key, "granted_scopes": granted, "expires_at": obo_result["expires_at"]}, + ) + else: + await self._store.set(self._token_key(identity, audience, "strict", granted), entry) + except Exception as exc: + store_err = TokenStoreError("Token store write failed", cause=exc) + logging.warning("Token store write failed, token still returned: %s", store_err.cause) + + # ===== Private methods ===== + + def _token_key(self, identity: _OboIdentity, audience: str, layout: str, scopes: Optional[str]) -> str: + return obo_cache_key( + layout=layout, + sub=identity.sub, + issuer=identity.issuer, + incoming_client_id=identity.incoming_client_id, + exchange_tenant=self._exchange_tenant, + exchange_client_id=self._exchange_client_id, + audience=audience, + org_id=identity.org_id, + scopes=scopes, + session_key=identity.session_key, + ) + + def _index_key(self, identity: _OboIdentity, audience: str) -> str: + return index_cache_key( + sub=identity.sub, + issuer=identity.issuer, + incoming_client_id=identity.incoming_client_id, + exchange_tenant=self._exchange_tenant, + exchange_client_id=self._exchange_client_id, + audience=audience, + org_id=identity.org_id, + session_key=identity.session_key, + ) + + async def _store_get(self, key: str) -> Optional[TokenSet]: + try: + return await self._store.get(key) + except Exception as exc: + store_err = TokenStoreError("Token store read failed", cause=exc) + logging.warning("Token store read failed, treating as cache miss: %s", store_err.cause) + return None + + @staticmethod + def _well_formed_member(member: object) -> bool: + """True when an index member has the fields the lookup reads, so a malformed one is a miss.""" + return ( + isinstance(member, dict) + and isinstance(member.get("token_key"), str) + and isinstance(member.get("granted_scopes"), str) + and isinstance(member.get("expires_at"), int) + ) + + @staticmethod + def _usable_entry(cached: Optional[TokenSet], now: int) -> bool: + if cached is None: + return False + if "access_token" not in cached or "expires_at" not in cached: + logging.warning("Token store returned a malformed entry, treating as cache miss") + return False + return cached["expires_at"] > now + + @staticmethod + def _hit(cached: TokenSet) -> OnBehalfOfTokenResult: + result: OnBehalfOfTokenResult = { + "access_token": cached["access_token"], + "expires_in": cached["expires_at"] - int(time.time()), + "expires_at": cached["expires_at"], + } + granted = cached.get("granted_scopes") + if granted: + result["scope"] = granted + return result + + async def _strict_lookup( + self, identity: _OboIdentity, audience: str, scope: Optional[str] + ) -> Optional[OnBehalfOfTokenResult]: + cached = await self._store_get(self._token_key(identity, audience, "strict", scope)) + if not self._usable_entry(cached, int(time.time())): + return None + # Entries are stored under their granted-scope key, so confirm the grant matches exactly. + if _normalized_scopes(cached.get("granted_scopes")) != _normalized_scopes(scope): + return None + return self._hit(cached) + + async def _index_lookup( + self, identity: _OboIdentity, audience: str, scope: Optional[str] + ) -> Optional[OnBehalfOfTokenResult]: + # An unscoped request is covered by every cached token, so never reuse for one. + if not scope: + return None + now = int(time.time()) + indexed = cast(IndexedTokenStore, self._store) + try: + members = await indexed.list_index_members(self._index_key(identity, audience)) + except Exception as exc: + store_err = TokenStoreError("Token store read failed", cause=exc) + logging.warning("Token store read failed, treating as cache miss: %s", store_err.cause) + return None + requested = set(_normalized_scopes(scope).split()) + candidates = [ + m for m in members + if self._well_formed_member(m) and m["expires_at"] > now and is_covered_by(scope, m["granted_scopes"]) + ] + # Prefer an exact grant, then the covering grant with the fewest extra scopes. + candidates.sort(key=lambda m: len(set(_normalized_scopes(m["granted_scopes"]).split()) - requested)) + for member in candidates: + cached = await self._store_get(member["token_key"]) + if not self._usable_entry(cached, now): + continue + # Trust the token value's own grant, not only the index label. + if not is_covered_by(scope, cached.get("granted_scopes")): + continue + return self._hit(cached) + return None diff --git a/src/auth0_api_python/api_client.py b/src/auth0_api_python/api_client.py index c5ba7c0..ea80760 100644 --- a/src/auth0_api_python/api_client.py +++ b/src/auth0_api_python/api_client.py @@ -6,6 +6,7 @@ import httpx from authlib.jose import JsonWebKey, JsonWebToken +from ._internal.obo_cache import OboCache from .cache import InMemoryCache from .config import ApiClientOptions from .errors import ( @@ -23,6 +24,10 @@ OrganizationNotAllowedError, VerifyAccessTokenError, ) +from .token_store import ( + IndexedTokenStore, + VerifiedToken, +) from .types import OnBehalfOfTokenResult from .utils import ( calculate_jwk_thumbprint, @@ -115,6 +120,19 @@ def __init__(self, options: ApiClientOptions): raise ConfigurationError( "organization_id is only valid when organization_policy is 'required'" ) + if options.scope_matching not in ("strict", "non_strict"): + raise ConfigurationError( + "scope_matching must be either 'strict' or 'non_strict'" + ) + if ( + options.scope_matching == "non_strict" + and options.token_store is not None + and not isinstance(options.token_store, IndexedTokenStore) + ): + raise ConfigurationError( + "scope_matching='non_strict' requires a token_store that subclasses IndexedTokenStore. " + "Use scope_matching='strict' or provide an IndexedTokenStore." + ) if options.cache_adapter: self._discovery_cache = options.cache_adapter @@ -123,6 +141,18 @@ def __init__(self, options: ApiClientOptions): self._discovery_cache = InMemoryCache(max_entries=options.cache_max_entries) self._jwks_cache = InMemoryCache(max_entries=options.cache_max_entries) + self._token_store = options.token_store + self._obo_cache = ( + OboCache( + options.token_store, + exchange_tenant=options.domain or "", + exchange_client_id=options.client_id or "", + scope_matching=options.scope_matching, + ) + if options.token_store is not None + else None + ) + self._cache_ttl = options.cache_ttl_seconds self._jwt = JsonWebToken(["RS256"]) @@ -999,6 +1029,8 @@ async def get_token_on_behalf_of( access_token: str, audience: str, scope: Optional[str] = None, + *, + verified: Optional[VerifiedToken] = None, ) -> OnBehalfOfTokenResult: """ Exchange an Auth0 access token for another Auth0 access token targeting a downstream API @@ -1011,6 +1043,9 @@ async def get_token_on_behalf_of( access_token: The Auth0 access token to exchange audience: Target API identifier for the exchanged access token scope: Optional space-separated OAuth 2.0 scopes to request + verified: The already-verified token, supplied by a caller that has verified it (for + example an MCP server). When omitted and a token_store is configured, the + token is verified here before any cache lookup. Returns: Dictionary containing: @@ -1021,14 +1056,38 @@ async def get_token_on_behalf_of( - token_type (str, optional): Token type (typically "Bearer") - issued_token_type (str, optional): RFC 8693 issued token type identifier + Caching is enabled only when a token_store is configured on the client. Without a store + every call performs a fresh exchange and nothing is cached. The scope_matching option + controls whether a cached token is reused only on an exact scope match ("strict") or + whenever its granted scopes cover the request ("non_strict"). + Raises: MissingRequiredArgumentError: If required parameters are missing + VerifyAccessTokenError: If a store is configured and either verified is omitted and the token fails verification, or verified is supplied but does not match the access token being exchanged + MissingOrganizationError: If organization_policy is "required" and the token has no org_id claim + OrganizationNotAllowedError: If the token's org_id is not in the organization_id allowlist GetTokenByExchangeProfileError: If client credentials are not configured or validation fails ApiError: If the token endpoint returns an error """ if not audience: raise MissingRequiredArgumentError("audience") + identity = None + if self._obo_cache is not None: + if verified is None: + claims = await self.verify_access_token(access_token) + verified = VerifiedToken(access_token=access_token, claims=claims) + elif verified.access_token != access_token: + # The cache identity comes from the verified claims, so those claims must belong + # to the token being exchanged or a caller could read back another token's entry. + raise VerifyAccessTokenError("verified token does not match the access token being exchanged") + identity = self._obo_cache.identity(verified) + + if identity is not None: + hit = await self._obo_cache.lookup(identity, audience, scope) + if hit is not None: + return hit + result = await self.get_token_by_exchange_profile( subject_token=access_token, subject_token_type=OBO_ACCESS_TOKEN_TYPE, @@ -1050,6 +1109,9 @@ async def get_token_on_behalf_of( if "issued_token_type" in result: obo_result["issued_token_type"] = result["issued_token_type"] + if identity is not None: + await self._obo_cache.write(identity, audience, scope, obo_result) + return obo_result # ===== Private Methods ===== diff --git a/src/auth0_api_python/config.py b/src/auth0_api_python/config.py index 68ebdad..0c8e4d1 100644 --- a/src/auth0_api_python/config.py +++ b/src/auth0_api_python/config.py @@ -6,6 +6,7 @@ if TYPE_CHECKING: from .cache import CacheAdapter + from .token_store import AbstractTokenStore class ApiClientOptions: @@ -21,6 +22,7 @@ class ApiClientOptions: audience: The expected 'aud' claim in the token. custom_fetch: Optional callable that can replace the default HTTP fetch logic. cache_adapter: Custom cache implementation. If not provided, uses default InMemoryCache. + token_store: Custom token storage for token caching. If not provided, caching is disabled and every exchange makes a fresh network call. cache_ttl_seconds: Time-to-live for cache entries in seconds (default: 600 = 10 minutes). cache_max_entries: Maximum number of cache entries before LRU eviction (default: 100). dpop_enabled: Whether DPoP is enabled (default: True for backward compatibility). @@ -39,6 +41,11 @@ class ApiClientOptions: organization_id: Optional allowlist of org_id claim values (a single value or a list). Only valid when organization_policy is "required" - passing it with "allow" raises ConfigurationError at construction time. + scope_matching: OBO cache scope-matching mode, applied only when token_store is set. + "strict" (default) reuses a cached token only on an exact granted-scope + match. "non_strict" reuses any cached token whose granted scopes cover + the request, and requires a token_store subclassing IndexedTokenStore + (raises ConfigurationError at construction otherwise). """ def __init__( self, @@ -47,6 +54,7 @@ def __init__( domains: Optional[Union[list[str], Callable[[dict], list[str]]]] = None, custom_fetch: Optional[Callable[..., object]] = None, cache_adapter: Optional["CacheAdapter"] = None, + token_store: Optional["AbstractTokenStore"] = None, cache_ttl_seconds: int = 600, cache_max_entries: int = 100, dpop_enabled: bool = True, @@ -58,12 +66,14 @@ def __init__( timeout: float = 10.0, organization_policy: str = "allow", organization_id: Optional[Union[str, list[str]]] = None, + scope_matching: str = "strict", ): self.domain = domain self.domains = domains self.audience = audience self.custom_fetch = custom_fetch self.cache_adapter = cache_adapter + self.token_store = token_store self.cache_ttl_seconds = cache_ttl_seconds self.cache_max_entries = cache_max_entries self.dpop_enabled = dpop_enabled @@ -75,3 +85,4 @@ def __init__( self.timeout = timeout self.organization_policy = organization_policy self.organization_id = organization_id + self.scope_matching = scope_matching diff --git a/src/auth0_api_python/encryption.py b/src/auth0_api_python/encryption.py new file mode 100644 index 0000000..0363889 --- /dev/null +++ b/src/auth0_api_python/encryption.py @@ -0,0 +1,61 @@ +"""JWE encryption helpers for token store implementations.""" + +from __future__ import annotations + +import base64 +import json +import uuid +from typing import Any + +from cryptography.hazmat.primitives import hashes +from cryptography.hazmat.primitives.kdf.hkdf import HKDF +from jwcrypto import jwe, jwk +from jwcrypto.common import base64url_encode + +_ENC = "A256CBC-HS512" +_ALG = "dir" +_KEY_LENGTH = 64 +_INFO = b"Auth0 Generated Encryption" + + +def _derive_key(secret: bytes, salt: bytes) -> bytes: + return HKDF( + algorithm=hashes.SHA256(), + length=_KEY_LENGTH, + salt=salt, + info=_INFO, + ).derive(secret) + + +def _read_jwe_kid(token: str) -> str: + """Extract the kid from a compact JWE protected header.""" + header_b64 = token.split(".")[0] + padding = 4 - len(header_b64) % 4 + if padding != 4: + header_b64 += "=" * padding + header = json.loads(base64.urlsafe_b64decode(header_b64)) + kid = header.get("kid") + if not kid: + raise ValueError('Missing "kid" in JWE header') + return kid + + +def encrypt(payload: dict[str, Any], secret: str, salt: str) -> str: + """Encrypt a dict to a compact JWE string. A fresh random kid is generated per call.""" + kid = str(uuid.uuid4()) + key_bytes = _derive_key(secret.encode(), f"{salt}{kid}".encode()) + key = jwk.JWK(k=base64url_encode(key_bytes), kty="oct") + token = jwe.JWE(json.dumps(payload), protected={"alg": _ALG, "enc": _ENC, "kid": kid}) + token.add_recipient(key) + return token.serialize(compact=True) + + +def decrypt(data: str, secret: str, salt: str) -> dict[str, Any]: + """Decrypt a compact JWE string back to a dict.""" + kid = _read_jwe_kid(data) + key_bytes = _derive_key(secret.encode(), f"{salt}{kid}".encode()) + key = jwk.JWK(k=base64url_encode(key_bytes), kty="oct") + token = jwe.JWE() + token.deserialize(data) + token.decrypt(key) + return json.loads(token.payload.decode()) diff --git a/src/auth0_api_python/errors.py b/src/auth0_api_python/errors.py index 10d4a9f..3b70556 100644 --- a/src/auth0_api_python/errors.py +++ b/src/auth0_api_python/errors.py @@ -174,3 +174,17 @@ def get_status_code(self) -> int: def get_error_code(self) -> str: return "domains_resolver_error" + + +class TokenStoreError(BaseAuthError): + """Raised when the configured TokenStore backend fails.""" + + def __init__(self, message: str, cause: Exception = None) -> None: + super().__init__(message) + self.cause = cause + + def get_status_code(self) -> int: + return 500 + + def get_error_code(self) -> str: + return "token_store_error" diff --git a/src/auth0_api_python/token_store.py b/src/auth0_api_python/token_store.py new file mode 100644 index 0000000..98b8c1b --- /dev/null +++ b/src/auth0_api_python/token_store.py @@ -0,0 +1,79 @@ +"""Token storage for SDK-minted tokens (OBO, M2M, Token Vault), distinct from CacheAdapter.""" + +from abc import ABC, abstractmethod +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any, Optional, TypedDict + +from .encryption import decrypt, encrypt + + +class _TokenSetRequired(TypedDict): + access_token: str + expires_at: int + + +class TokenSet(_TokenSetRequired, total=False): + """A minted access token and its absolute expiration, as stored by a TokenStore.""" + + granted_scopes: str + + +@dataclass(frozen=True) +class VerifiedToken: + """Access token with verified claims from verify_access_token, trusted to build the cache key.""" + + access_token: str + claims: Mapping[str, Any] + + +class AbstractTokenStore(ABC): + """Base class for external token stores with built-in JWE encrypt and decrypt helpers.""" + + def __init__(self, *, secret: str) -> None: + self._secret = secret + + @abstractmethod + async def get(self, key: str) -> Optional[TokenSet]: + """Return the stored token or None if absent.""" + pass + + @abstractmethod + async def set(self, key: str, value: TokenSet) -> None: + """Store a token under key.""" + pass + + @abstractmethod + async def delete(self, key: str) -> None: + """Delete a stored token by key.""" + pass + + def encrypt(self, key: str, value: TokenSet) -> str: + """Encrypt a TokenSet to a JWE string, keyed to this specific cache entry.""" + return encrypt(dict(value), self._secret, key) + + def decrypt(self, key: str, data: str) -> TokenSet: + """Decrypt a JWE string back to a TokenSet.""" + return decrypt(data, self._secret, key) + + +class TokenIndexMember(TypedDict): + """One cached token listed in an index, with the scopes it was granted and when it expires.""" + + token_key: str + granted_scopes: str + expires_at: int + + +class IndexedTokenStore(AbstractTokenStore): + """Token store variant that maintains a scope index with atomic member writes to avoid concurrent-add races.""" + + @abstractmethod + async def add_index_member(self, index_key: str, member: TokenIndexMember) -> None: + """Atomically add member to the index, replacing any member with the same token_key.""" + pass + + @abstractmethod + async def list_index_members(self, index_key: str) -> list[TokenIndexMember]: + """Return the index members, possibly including expired ones, or [] if the index is absent.""" + pass diff --git a/tests/conftest.py b/tests/conftest.py index 4fb119a..2e5d349 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,6 +1,7 @@ """Shared test fixtures and helpers for auth0-api-python tests.""" import base64 +import time import urllib.parse from typing import Optional @@ -9,6 +10,49 @@ from auth0_api_python import ApiClient, ApiClientOptions from auth0_api_python.errors import ApiError +from auth0_api_python.token_store import ( + AbstractTokenStore, + IndexedTokenStore, + TokenIndexMember, + TokenSet, +) + + +class InMemoryTokenStore(AbstractTokenStore): + """Minimal in-memory store for tests. Not for production use.""" + + def __init__(self, *, secret: str = "test-secret") -> None: # noqa: S107 + super().__init__(secret=secret) + self._store: dict[str, TokenSet] = {} + + async def get(self, key: str) -> Optional[TokenSet]: + entry = self._store.get(key) + if entry is None: + return None + if entry["expires_at"] <= int(time.time()): + del self._store[key] + return None + return entry + + async def set(self, key: str, value: TokenSet) -> None: + self._store[key] = value + + async def delete(self, key: str) -> None: + self._store.pop(key, None) + + +class InMemoryIndexedTokenStore(InMemoryTokenStore, IndexedTokenStore): + """In-memory IndexedTokenStore for tests. Not for production use.""" + + def __init__(self, *, secret: str = "test-secret") -> None: # noqa: S107 + super().__init__(secret=secret) + self._index: dict[str, dict[str, TokenIndexMember]] = {} + + async def add_index_member(self, index_key: str, member: TokenIndexMember) -> None: + self._index.setdefault(index_key, {})[member["token_key"]] = member + + async def list_index_members(self, index_key: str) -> list[TokenIndexMember]: + return list(self._index.get(index_key, {}).values()) # ===== Constants ===== @@ -21,6 +65,18 @@ @pytest.fixture def api_client_confidential(): """Fixture for creating a confidential API client with credentials.""" + return ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=InMemoryTokenStore(), + )) + + +@pytest.fixture +def api_client_confidential_no_store(): + """Confidential client without a token_store, so OBO exchanges never cache or verify.""" return ApiClient(ApiClientOptions( domain="auth0.local", audience="my-audience", diff --git a/tests/test_api_client.py b/tests/test_api_client.py index c10c89e..2f741a5 100644 --- a/tests/test_api_client.py +++ b/tests/test_api_client.py @@ -1,5 +1,7 @@ +import asyncio import base64 import json +import logging import time import httpx @@ -10,6 +12,8 @@ DISCOVERY_URL, JWKS_URL, TOKEN_ENDPOINT, + InMemoryIndexedTokenStore, + InMemoryTokenStore, assert_api_error, assert_form_post, assert_no_requests, @@ -20,6 +24,11 @@ from pytest_httpx import HTTPXMock from auth0_api_python import get_current_actor, get_delegation_chain +from auth0_api_python._internal.cache_keys import ( + index_cache_key, + obo_cache_key, + session_fingerprint, +) from auth0_api_python.api_client import MAX_ARRAY_VALUES_PER_KEY, ApiClient from auth0_api_python.config import ApiClientOptions from auth0_api_python.errors import ( @@ -36,6 +45,7 @@ OrganizationNotAllowedError, VerifyAccessTokenError, ) +from auth0_api_python.token_store import VerifiedToken from auth0_api_python.token_utils import ( PRIVATE_EC_JWK, PRIVATE_JWK, @@ -3089,7 +3099,7 @@ async def test_get_token_on_behalf_of_missing_audience(api_client_confidential): @pytest.mark.asyncio -async def test_get_token_on_behalf_of_success(mock_discovery, api_client_confidential, httpx_mock): +async def test_get_token_on_behalf_of_success(mock_discovery, api_client_confidential_no_store, httpx_mock): """Test successful OBO exchange with fixed access-token types.""" httpx_mock.add_response( method="POST", @@ -3103,7 +3113,7 @@ async def test_get_token_on_behalf_of_success(mock_discovery, api_client_confide } ) - result = await api_client_confidential.get_token_on_behalf_of( + result = await api_client_confidential_no_store.get_token_on_behalf_of( access_token="incoming-access-token", audience="https://api.backend.com", scope="read:data write:data" @@ -3133,12 +3143,12 @@ async def test_get_token_on_behalf_of_success(mock_discovery, api_client_confide @pytest.mark.asyncio async def test_get_token_on_behalf_of_preserves_inferred_expires_at( - api_client_confidential, + api_client_confidential_no_store, monkeypatch, ): """Test that OBO reuses the expires_at inferred by the generic exchange path.""" - async def fake_exchange_profile( + async def stub_exchange_profile( *, subject_token, subject_token_type, @@ -3163,12 +3173,12 @@ async def fake_exchange_profile( } monkeypatch.setattr( - api_client_confidential, + api_client_confidential_no_store, "get_token_by_exchange_profile", - fake_exchange_profile, + stub_exchange_profile, ) - result = await api_client_confidential.get_token_on_behalf_of( + result = await api_client_confidential_no_store.get_token_on_behalf_of( access_token="incoming-access-token", audience="https://api.backend.com", ) @@ -3181,7 +3191,7 @@ async def fake_exchange_profile( @pytest.mark.asyncio -async def test_get_token_on_behalf_of_without_scope(mock_discovery, api_client_confidential, httpx_mock): +async def test_get_token_on_behalf_of_without_scope(mock_discovery, api_client_confidential_no_store, httpx_mock): """Test OBO exchange omits scope when not provided.""" httpx_mock.add_response( method="POST", @@ -3193,7 +3203,7 @@ async def test_get_token_on_behalf_of_without_scope(mock_discovery, api_client_c ) ) - result = await api_client_confidential.get_token_on_behalf_of( + result = await api_client_confidential_no_store.get_token_on_behalf_of( access_token="incoming-access-token", audience="https://api.backend.com", ) @@ -3204,7 +3214,7 @@ async def test_get_token_on_behalf_of_without_scope(mock_discovery, api_client_c @pytest.mark.asyncio async def test_get_token_on_behalf_of_does_not_expose_id_or_refresh_token( - mock_discovery, api_client_confidential, httpx_mock + mock_discovery, api_client_confidential_no_store, httpx_mock ): """Test OBO result only exposes access-token-oriented fields.""" httpx_mock.add_response( @@ -3220,7 +3230,7 @@ async def test_get_token_on_behalf_of_does_not_expose_id_or_refresh_token( } ) - result = await api_client_confidential.get_token_on_behalf_of( + result = await api_client_confidential_no_store.get_token_on_behalf_of( access_token="incoming-access-token", audience="https://api.backend.com", ) @@ -3231,7 +3241,7 @@ async def test_get_token_on_behalf_of_does_not_expose_id_or_refresh_token( @pytest.mark.asyncio -async def test_get_token_on_behalf_of_api_error(mock_discovery, api_client_confidential, httpx_mock): +async def test_get_token_on_behalf_of_api_error(mock_discovery, api_client_confidential_no_store, httpx_mock): """Test that OBO reuses the existing exchange error semantics.""" httpx_mock.add_response( method="POST", @@ -3244,7 +3254,7 @@ async def test_get_token_on_behalf_of_api_error(mock_discovery, api_client_confi ) with pytest.raises(ApiError) as err: - await api_client_confidential.get_token_on_behalf_of( + await api_client_confidential_no_store.get_token_on_behalf_of( access_token="incoming-access-token", audience="https://api.backend.com", ) @@ -3263,6 +3273,1009 @@ async def test_get_token_on_behalf_of_empty_access_token(api_client_confidential ) +# ===== Token Storage Tests ===== + +# ----- OBO ----- + +# The confidential fixture's verified identity: issuer, exchange tenant, and exchange client. +CONFIDENTIAL_ISS = "https://auth0.local/" +CONFIDENTIAL_TENANT = "auth0.local" +CONFIDENTIAL_CID = "cid" + + +def make_access_token(sub: str = "auth0|user1", *, org_id: str = None) -> str: + """Build a minimal stub JWT-shaped token carrying the given sub and optional org_id.""" + claims = {"sub": sub} + if org_id is not None: + claims["org_id"] = org_id + payload = base64.urlsafe_b64encode(json.dumps(claims).encode()).rstrip(b"=").decode() + return f"hdr.{payload}.sig" + + +def verified_for(access_token, sub="auth0|user1", *, org_id=None, iss=CONFIDENTIAL_ISS): + """Build the VerifiedToken an MCP server would pass, so the client skips its own verify.""" + claims = {"sub": sub, "iss": iss} + if org_id is not None: + claims["org_id"] = org_id + return VerifiedToken(access_token=access_token, claims=claims) + + +def strict_key(access_token, *, audience, scopes, sub="auth0|user1", org_id=None, iss=CONFIDENTIAL_ISS): + """The strict-layout token key OboCache computes for the confidential fixture's identity.""" + return obo_cache_key( + layout="strict", sub=sub, issuer=iss, incoming_client_id="", + exchange_tenant=CONFIDENTIAL_TENANT, exchange_client_id=CONFIDENTIAL_CID, + audience=audience, org_id=org_id, scopes=scopes, + session_key=session_fingerprint(access_token), + ) + + +def index_key(access_token, *, audience, sub="auth0|user1", org_id=None, iss=CONFIDENTIAL_ISS): + """The index key OboCache computes for the confidential fixture's identity.""" + return index_cache_key( + sub=sub, issuer=iss, incoming_client_id="", + exchange_tenant=CONFIDENTIAL_TENANT, exchange_client_id=CONFIDENTIAL_CID, + audience=audience, org_id=org_id, session_key=session_fingerprint(access_token), + ) + + +def index_token_key(access_token, *, audience, scopes, sub="auth0|user1", org_id=None, iss=CONFIDENTIAL_ISS): + """The index-layout token key OboCache computes for the confidential fixture's identity.""" + return obo_cache_key( + layout="index", sub=sub, issuer=iss, incoming_client_id="", + exchange_tenant=CONFIDENTIAL_TENANT, exchange_client_id=CONFIDENTIAL_CID, + audience=audience, org_id=org_id, scopes=scopes, + session_key=session_fingerprint(access_token), + ) + + +class RaisingGetTokenStore(InMemoryTokenStore): + """Token store whose get() always raises, to verify read failures fall back to a fresh exchange.""" + + async def get(self, key): + raise RuntimeError("store unavailable") + + +class RaisingSetTokenStore(InMemoryTokenStore): + """Token store whose set() always raises, to verify write failures don't fail the caller.""" + + async def set(self, key, value): + raise RuntimeError("store unavailable") + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_no_token_store_always_fresh_exchange( + mock_discovery, httpx_mock +): + """Test that no token_store means every call performs a fresh exchange.""" + api_client = ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + )) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-1"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-2"), + ) + + payload = base64.urlsafe_b64encode(json.dumps({"sub": "auth0|user1"}).encode()).rstrip(b"=").decode() + token = f"hdr.{payload}.sig" + result1 = await api_client.get_token_on_behalf_of(access_token=token, audience="https://api.backend.com") + result2 = await api_client.get_token_on_behalf_of(access_token=token, audience="https://api.backend.com") + + assert result1["access_token"] == "obo-access-token-1" + assert result2["access_token"] == "obo-access-token-2" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 2 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_with_token_store_cache_miss_then_hit( + mock_discovery, api_client_confidential, httpx_mock +): + """Test that a second call for the same token/audience/scope is served from cache.""" + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + result1 = await api_client_confidential.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + scope="read:data", + verified=verified, + ) + result2 = await api_client_confidential.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + scope="read:data", + verified=verified, + ) + + assert result1["access_token"] == "obo-access-token" + assert result2["access_token"] == "obo-access-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 1 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_unverifiable_raw_token_raises( + api_client_confidential, httpx_mock +): + """With a store and no verified context, an unverifiable raw token raises before any exchange.""" + with pytest.raises(VerifyAccessTokenError): + await api_client_confidential.get_token_on_behalf_of( + access_token="not-a-jwt", + audience="https://api.backend.com", + ) + + assert not [r for r in httpx_mock.get_requests() if r.method == "POST"] + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_verified_context_skips_reverification( + mock_discovery, api_client_confidential, httpx_mock, monkeypatch +): + """A supplied verified context is trusted, so the client does not verify the token itself.""" + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token"), + ) + + async def fail_if_called(*args, **kwargs): + raise AssertionError("verify_access_token must not be called when verified is supplied") + + monkeypatch.setattr(api_client_confidential, "verify_access_token", fail_if_called) + + token = make_access_token("auth0|user1") + result = await api_client_confidential.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + verified=verified_for(token, "auth0|user1"), + ) + + assert result["access_token"] == "obo-access-token" + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_verified_claims_without_sub_skips_caching( + mock_discovery, api_client_confidential, httpx_mock, caplog +): + """Verified claims that lack a usable sub skip caching with a warning, so each call re-exchanges.""" + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-1"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-2"), + ) + + caplog.set_level(logging.WARNING) + token = make_access_token("auth0|user1") + no_sub = VerifiedToken(access_token=token, claims={"iss": CONFIDENTIAL_ISS}) + result1 = await api_client_confidential.get_token_on_behalf_of( + access_token=token, audience="https://api.backend.com", verified=no_sub + ) + result2 = await api_client_confidential.get_token_on_behalf_of( + access_token=token, audience="https://api.backend.com", verified=no_sub + ) + + assert result1["access_token"] == "obo-access-token-1" + assert result2["access_token"] == "obo-access-token-2" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 2 + assert any("no usable sub" in record.message for record in caplog.records) + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_verified_token_mismatch_raises( + api_client_confidential, httpx_mock +): + """A verified token whose access_token differs from the one being exchanged is rejected.""" + token = make_access_token("auth0|user1") + other = make_access_token("auth0|user2") + + with pytest.raises(VerifyAccessTokenError): + await api_client_confidential.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + verified=verified_for(other, "auth0|user2"), + ) + + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 0 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_different_sub_does_not_share_cache( + mock_discovery, api_client_confidential, httpx_mock +): + """Test that tokens with different sub values each trigger their own exchange.""" + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-1"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-2"), + ) + + token1 = make_access_token("auth0|user1") + token2 = make_access_token("auth0|user2") + + await api_client_confidential.get_token_on_behalf_of( + access_token=token1, audience="https://api.backend.com", verified=verified_for(token1, "auth0|user1") + ) + await api_client_confidential.get_token_on_behalf_of( + access_token=token2, audience="https://api.backend.com", verified=verified_for(token2, "auth0|user2") + ) + + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 2 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_different_org_id_does_not_share_cache( + mock_discovery, api_client_confidential, httpx_mock +): + """Test that tokens with different org_id values each trigger their own exchange.""" + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-1"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token-2"), + ) + + token1 = make_access_token("auth0|user1", org_id="org_abc") + token2 = make_access_token("auth0|user1", org_id="org_def") + + await api_client_confidential.get_token_on_behalf_of( + access_token=token1, audience="https://api.backend.com", + verified=verified_for(token1, "auth0|user1", org_id="org_abc"), + ) + await api_client_confidential.get_token_on_behalf_of( + access_token=token2, audience="https://api.backend.com", + verified=verified_for(token2, "auth0|user1", org_id="org_def"), + ) + + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 2 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_store_get_failure_falls_back_to_fresh_exchange( + mock_discovery, httpx_mock, caplog +): + """Test that a token store whose get() raises still succeeds via a fresh exchange.""" + api_client = ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=RaisingGetTokenStore(), + )) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token"), + ) + + token = make_access_token("auth0|user1") + caplog.set_level(logging.WARNING) + result = await api_client.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + verified=verified_for(token, "auth0|user1"), + ) + + assert result["access_token"] == "obo-access-token" + assert any("Token store read failed" in record.message for record in caplog.records) + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_store_set_failure_still_returns_token( + mock_discovery, httpx_mock, caplog +): + """Test that a token store whose set() raises still returns the exchanged token.""" + api_client = ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=RaisingSetTokenStore(), + )) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token"), + ) + + token = make_access_token("auth0|user1") + caplog.set_level(logging.WARNING) + result = await api_client.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + verified=verified_for(token, "auth0|user1"), + ) + + assert result["access_token"] == "obo-access-token" + assert any("Token store write failed" in record.message for record in caplog.records) + + +class ExpiredEntryTokenStore(InMemoryTokenStore): + """Token store whose get() always returns an expired TokenSet.""" + + async def get(self, key: str): + return {"access_token": "expired-token", "expires_at": int(time.time()) - 10} + + +class CorruptEntryTokenStore(InMemoryTokenStore): + """Token store whose get() returns a malformed entry missing required keys.""" + + async def get(self, key: str): + return {"not_a_token": "garbage"} + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_cache_hit_has_valid_expires_in( + mock_discovery, api_client_confidential, httpx_mock +): + """Test that a cache hit returns expires_in > 0 and a future expires_at.""" + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="obo-access-token"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + + # First call — populate cache + await api_client_confidential.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + scope="read:data", + verified=verified, + ) + # Second call — cache hit + result = await api_client_confidential.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + scope="read:data", + verified=verified, + ) + + assert result["expires_in"] > 0 + assert "expires_at" in result + assert result["expires_at"] > int(time.time()) + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_expired_store_entry_triggers_fresh_exchange( + mock_discovery, httpx_mock +): + """Test that a store returning an already-expired entry always triggers a fresh exchange.""" + api_client = ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=ExpiredEntryTokenStore(), + )) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token-1"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token-2"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + + await api_client.get_token_on_behalf_of(access_token=token, audience="https://api.backend.com", verified=verified) + await api_client.get_token_on_behalf_of(access_token=token, audience="https://api.backend.com", verified=verified) + + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 2 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_corrupt_store_value_treated_as_cache_miss( + mock_discovery, httpx_mock +): + """Test that a corrupt/partial store value is treated as a cache miss and a fresh exchange happens.""" + api_client = ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=CorruptEntryTokenStore(), + )) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token"), + ) + + token = make_access_token("auth0|user1") + result = await api_client.get_token_on_behalf_of( + access_token=token, + audience="https://api.backend.com", + verified=verified_for(token, "auth0|user1"), + ) + + assert result["access_token"] == "fresh-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 1 + + +@pytest.mark.asyncio +async def test_get_token_on_behalf_of_expired_cache_entry_triggers_fresh_exchange( + mock_discovery, api_client_confidential, httpx_mock +): + """Test that an expired cached entry is not returned and a fresh exchange happens.""" + audience = "https://api.backend.com" + access_token = make_access_token("auth0|user1") + + cache_key = strict_key(access_token, audience=audience, scopes=None) + await api_client_confidential._token_store.set( + cache_key, + {"access_token": "expired-access-token", "expires_at": int(time.time()) - 10}, + ) + + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-obo-access-token"), + ) + + result = await api_client_confidential.get_token_on_behalf_of( + access_token=access_token, + audience=audience, + verified=verified_for(access_token, "auth0|user1"), + ) + + assert result["access_token"] == "fresh-obo-access-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 1 + + +@pytest.mark.asyncio +async def test_strict_downscoped_cached_under_granted_scope(mock_discovery, api_client_confidential, httpx_mock): + """strict downscoping: a grant narrower than the request is reusable only for the granted subset.""" + audience = "https://api.backend.com" + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="narrow-token", scope="read:x")) + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="wide-token", scope="read:x write:x")) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + # Request read:x write:x, Auth0 downscopes to read:x (fresh downscoping). + r1 = await api_client_confidential.get_token_on_behalf_of(access_token=token, audience=audience, scope="read:x write:x", verified=verified) + # The granted subset is a hit (cached downscoping). + r2 = await api_client_confidential.get_token_on_behalf_of(access_token=token, audience=audience, scope="read:x", verified=verified) + # Re-requesting the wider scope is a miss under strict, so a fresh exchange runs. + r3 = await api_client_confidential.get_token_on_behalf_of(access_token=token, audience=audience, scope="read:x write:x", verified=verified) + + assert r1["access_token"] == "narrow-token" + assert r2["access_token"] == "narrow-token" + assert r3["access_token"] == "wide-token" + posts = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(posts) == 2 + + +@pytest.mark.asyncio +async def test_obo_different_issuer_does_not_share_cache(mock_discovery, api_client_confidential, httpx_mock): + """A cached token is not reused when the verified issuer differs, even for the same subject and token.""" + audience = "https://api.backend.com" + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="issuer-a-token")) + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="issuer-b-token")) + + token = make_access_token("auth0|user1") + r1 = await api_client_confidential.get_token_on_behalf_of( + access_token=token, audience=audience, verified=verified_for(token, "auth0|user1", iss="https://issuer-a/")) + r2 = await api_client_confidential.get_token_on_behalf_of( + access_token=token, audience=audience, verified=verified_for(token, "auth0|user1", iss="https://issuer-b/")) + + assert r1["access_token"] == "issuer-a-token" + assert r2["access_token"] == "issuer-b-token" + posts = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(posts) == 2 + + +# ----- OBO non_strict ----- + +def _nss_client(token_store=None): + """Build a confidential client with scope_matching='non_strict', which needs an indexed store.""" + return ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=token_store or InMemoryIndexedTokenStore(), + scope_matching="non_strict", + )) + + +@pytest.mark.asyncio +async def test_nss_exact_scope_hit(mock_discovery, httpx_mock): + """non_strict: exact scope repeat returns cached token with no second exchange.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="nss-token", scope="read:x"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + result1 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + result2 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + + assert result1["access_token"] == "nss-token" + assert result2["access_token"] == "nss-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_superset_granted_covers_subset_request(mock_discovery, httpx_mock): + """non_strict: a request for a subset of the cached granted scopes reuses the token.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="wide-token", scope="read:x write:x"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x write:x", verified=verified) + result = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + + assert result["access_token"] == "wide-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_not_covered_triggers_fresh_exchange(mock_discovery, httpx_mock): + """non_strict: a request whose scopes exceed the cached granted scopes triggers a fresh exchange.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="narrow-token", scope="read:x"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="wide-token", scope="read:x write:x"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + result1 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + result2 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x write:x", verified=verified) + + assert result1["access_token"] == "narrow-token" + assert result2["access_token"] == "wide-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 2 + + +@pytest.mark.asyncio +async def test_nss_downscoped_grant_reused_for_subset(mock_discovery, httpx_mock): + """non_strict downscoping: a grant narrower than the request is cached and reused for its subset.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="downscoped-token", scope="read:x"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + # Request read:x write:x but Auth0 grants only read:x (fresh downscoping). + result1 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x write:x", verified=verified) + # Requesting the granted subset reuses it (cached downscoping), no second exchange. + result2 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + + assert result1["access_token"] == "downscoped-token" + assert result2["access_token"] == "downscoped-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_omitted_scope_never_reuses(mock_discovery, httpx_mock): + """non_strict: a scopeless request never reuses a cached token, since every grant would cover it.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="scoped-token", scope="read:x"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + result = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", verified=verified) + + assert result["access_token"] == "fresh-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 2 + + +@pytest.mark.asyncio +async def test_nss_covering_check_uses_granted_not_requested(mock_discovery, httpx_mock): + """non_strict: the covering check uses a cached member's granted_scopes, not the original request.""" + client = _nss_client() + token = make_access_token("auth0|user1") + audience = "https://api.example.com" + + # Seed a live index member that was granted only read:x. + await client._token_store.add_index_member( + index_key(token, audience=audience), + { + "token_key": index_token_key(token, audience=audience, scopes="read:x"), + "granted_scopes": "read:x", + "expires_at": int(time.time()) + 3600, + }, + ) + await client._token_store.set( + index_token_key(token, audience=audience, scopes="read:x"), + {"access_token": "narrow-cached-token", "expires_at": int(time.time()) + 3600, "granted_scopes": "read:x"}, + ) + + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-wide-token", scope="read:x write:x"), + ) + + # read:x does not cover read:x write:x, so the seeded member is not reused. + result = await client.get_token_on_behalf_of(access_token=token, audience=audience, scope="read:x write:x", verified=verified_for(token, "auth0|user1")) + + assert result["access_token"] == "fresh-wide-token" + token_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(token_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_invalid_scope_matching_raises_configuration_error(): + """ApiClient construction with an invalid scope_matching value raises ConfigurationError.""" + with pytest.raises(ConfigurationError) as err: + ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + scope_matching="invalid_value", + )) + + assert "scope_matching" in str(err.value) + + +def test_nss_requires_indexed_store(): + """non_strict with a plain AbstractTokenStore raises ConfigurationError at construction.""" + with pytest.raises(ConfigurationError) as err: + ApiClient(ApiClientOptions( + domain="auth0.local", + audience="my-audience", + client_id="cid", + client_secret="csecret", + token_store=InMemoryTokenStore(), + scope_matching="non_strict", + )) + assert "IndexedTokenStore" in str(err.value) + + +@pytest.mark.httpx_mock(can_send_already_matched_responses=True) +@pytest.mark.asyncio +async def test_strict_and_non_strict_do_not_share_cache(mock_discovery, httpx_mock): + """A token cached in strict layout is invisible to a non_strict client sharing the same store.""" + store = InMemoryIndexedTokenStore() + strict_client = ApiClient(ApiClientOptions( + domain="auth0.local", audience="my-audience", client_id="cid", client_secret="csecret", + token_store=store, scope_matching="strict", + )) + nss_client = ApiClient(ApiClientOptions( + domain="auth0.local", audience="my-audience", client_id="cid", client_secret="csecret", + token_store=store, scope_matching="non_strict", + )) + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="strict-token", scope="read:x")) + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="nss-token", scope="read:x")) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + r1 = await strict_client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + r2 = await nss_client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + + assert r1["access_token"] == "strict-token" + assert r2["access_token"] == "nss-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 2 + + +# ----- OBO non_strict index ----- + +class RaisingListIndexTokenStore(InMemoryIndexedTokenStore): + """Indexed store whose list_index_members() raises, to verify index read failures fall back to a fresh exchange.""" + + async def list_index_members(self, index_key): + raise RuntimeError("store unavailable") + + +class RaisingGetIndexTokenStore(InMemoryIndexedTokenStore): + """Indexed store whose get() raises, to verify a member's token read failure falls back to a fresh exchange.""" + + async def get(self, key): + raise RuntimeError("store unavailable") + + +class RaisingAddIndexTokenStore(InMemoryIndexedTokenStore): + """Indexed store whose add_index_member() raises, to verify index write failures don't fail the caller.""" + + async def add_index_member(self, index_key, member): + raise RuntimeError("store unavailable") + + +@pytest.mark.asyncio +async def test_nss_index_upscope_reuse(mock_discovery, httpx_mock): + """Index: a wide token minted first is reused for a narrower subsequent request.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="wide-token", scope="read:x write:x"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + result1 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x write:x", verified=verified) + result2 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + + assert result1["access_token"] == "wide-token" + assert result2["access_token"] == "wide-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_index_no_evict(mock_discovery, httpx_mock): + """Index: minting a non-covering token does not evict a previously cached token.""" + client = _nss_client() + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="read-token", scope="read:x"), + ) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="write-token", scope="write:x"), + ) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="write:x", verified=verified) + result3 = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified) + + # read:x entry must still be cached - exchange should have run exactly twice + assert result3["access_token"] == "read-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 2 + + +@pytest.mark.asyncio +async def test_nss_index_expired_member_not_reused(mock_discovery, httpx_mock): + """Index: an expired member is skipped and a fresh exchange runs.""" + client = _nss_client() + token = make_access_token("auth0|user1") + audience = "https://api.example.com" + + # Index member is expired but its token value is still live, so only the index-level expiry check + # keeps it from being reused. Removing that check would surface the live value as a hit. + token_key = index_token_key(token, audience=audience, scopes="read:x") + await client._token_store.add_index_member(index_key(token, audience=audience), { + "token_key": token_key, + "granted_scopes": "read:x", + "expires_at": int(time.time()) - 10, + }) + await client._token_store.set( + token_key, + {"access_token": "stale-cached-token", "expires_at": int(time.time()) + 3600, "granted_scopes": "read:x"}, + ) + + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token", scope="read:x"), + ) + + result = await client.get_token_on_behalf_of(access_token=token, audience=audience, scope="read:x", verified=verified_for(token, "auth0|user1")) + + assert result["access_token"] == "fresh-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_index_list_failure_falls_back_to_exchange(mock_discovery, httpx_mock): + """Index: a failing index read is treated as a miss and a fresh exchange runs.""" + client = _nss_client(token_store=RaisingListIndexTokenStore()) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token", scope="read:x"), + ) + + token = make_access_token("auth0|user1") + result = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified_for(token, "auth0|user1")) + + assert result["access_token"] == "fresh-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_index_member_read_failure_falls_back_to_exchange(mock_discovery, httpx_mock): + """Index: a live member whose token read fails is treated as a miss and a fresh exchange runs.""" + client = _nss_client(token_store=RaisingGetIndexTokenStore()) + token = make_access_token("auth0|user1") + audience = "https://api.example.com" + + await client._token_store.add_index_member(index_key(token, audience=audience), { + "token_key": index_token_key(token, audience=audience, scopes="read:x"), + "granted_scopes": "read:x", + "expires_at": int(time.time()) + 3600, + }) + + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token", scope="read:x"), + ) + + result = await client.get_token_on_behalf_of(access_token=token, audience=audience, scope="read:x", verified=verified_for(token, "auth0|user1")) + + assert result["access_token"] == "fresh-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_index_malformed_member_treated_as_miss(mock_discovery, httpx_mock): + """Index: a member missing required fields is skipped, so a fresh exchange runs without raising.""" + client = _nss_client() + token = make_access_token("auth0|user1") + audience = "https://api.example.com" + + # A member lacking granted_scopes and expires_at must not break the lookup. + await client._token_store.add_index_member( + index_key(token, audience=audience), + {"token_key": index_token_key(token, audience=audience, scopes="read:x")}, + ) + + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token", scope="read:x"), + ) + + result = await client.get_token_on_behalf_of( + access_token=token, audience=audience, scope="read:x", verified=verified_for(token, "auth0|user1") + ) + + assert result["access_token"] == "fresh-token" + post_requests = [r for r in httpx_mock.get_requests() if r.method == "POST"] + assert len(post_requests) == 1 + + +@pytest.mark.asyncio +async def test_nss_index_add_failure_still_returns_token(mock_discovery, httpx_mock): + """Index: a failing index write does not fail the caller, the token is still returned.""" + client = _nss_client(token_store=RaisingAddIndexTokenStore()) + httpx_mock.add_response( + method="POST", + url=TOKEN_ENDPOINT, + json=token_success(access_token="fresh-token", scope="read:x"), + ) + + token = make_access_token("auth0|user1") + result = await client.get_token_on_behalf_of(access_token=token, audience="https://api.example.com", scope="read:x", verified=verified_for(token, "auth0|user1")) + + assert result["access_token"] == "fresh-token" + + +class BlobRewriteIndexedTokenStore(InMemoryIndexedTokenStore): + """Adversarial store that maintains the index with a non-atomic read-modify-write, to show + the atomic per-member add the SDK relies on is what keeps a concurrent addition from being lost.""" + + async def add_index_member(self, index_key, member): + current = dict(self._index.get(index_key, {})) + await asyncio.sleep(0) # yield mid read-modify-write so a concurrent add is lost + current[member["token_key"]] = member + self._index[index_key] = current + + +@pytest.mark.httpx_mock(can_send_already_matched_responses=True) +@pytest.mark.asyncio +async def test_nss_index_concurrent_adds_both_survive(mock_discovery, httpx_mock): + """Index: two concurrent exchanges for different scopes both stay cached, no lost update.""" + client = _nss_client() + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="t")) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + aud = "https://api.example.com" + await asyncio.gather( + client.get_token_on_behalf_of(access_token=token, audience=aud, scope="read:x", verified=verified), + client.get_token_on_behalf_of(access_token=token, audience=aud, scope="write:x", verified=verified), + ) + posts_after_concurrent = len([r for r in httpx_mock.get_requests() if r.method == "POST"]) + + # Both scopes must now be cached, so neither follow-up triggers a fresh exchange. + await client.get_token_on_behalf_of(access_token=token, audience=aud, scope="read:x", verified=verified) + await client.get_token_on_behalf_of(access_token=token, audience=aud, scope="write:x", verified=verified) + + posts_total = len([r for r in httpx_mock.get_requests() if r.method == "POST"]) + assert posts_after_concurrent == 2 + assert posts_total == 2 + + +@pytest.mark.httpx_mock(can_send_already_matched_responses=True) +@pytest.mark.asyncio +async def test_nss_index_non_atomic_store_loses_concurrent_add(mock_discovery, httpx_mock): + """Mutation check for the concurrency guarantee: a non-atomic read-modify-write store loses one + of two concurrent additions, so the dropped scope has to exchange again on the next call.""" + client = _nss_client(token_store=BlobRewriteIndexedTokenStore()) + httpx_mock.add_response(method="POST", url=TOKEN_ENDPOINT, json=token_success(access_token="t")) + + token = make_access_token("auth0|user1") + verified = verified_for(token, "auth0|user1") + aud = "https://api.example.com" + await asyncio.gather( + client.get_token_on_behalf_of(access_token=token, audience=aud, scope="read:x", verified=verified), + client.get_token_on_behalf_of(access_token=token, audience=aud, scope="write:x", verified=verified), + ) + + await client.get_token_on_behalf_of(access_token=token, audience=aud, scope="read:x", verified=verified) + await client.get_token_on_behalf_of(access_token=token, audience=aud, scope="write:x", verified=verified) + + # 2 concurrent exchanges plus 1 re-exchange for the member the non-atomic store dropped. + posts_total = len([r for r in httpx_mock.get_requests() if r.method == "POST"]) + assert posts_total == 3 + + # ===== MCD (Multi-Custom Domain) Tests ===== @pytest.mark.asyncio @@ -3711,8 +4724,8 @@ async def test_mcd_verify_rejects_symmetric_algorithm(): # Encode header and payload (signature doesn't matter for this test) header_b64 = base64.urlsafe_b64encode(json.dumps(header).encode()).decode().rstrip('=') payload_b64 = base64.urlsafe_b64encode(json.dumps(payload).encode()).decode().rstrip('=') - fake_signature = "fake_signature" - hs256_token = f"{header_b64}.{payload_b64}.{fake_signature}" + stub_signature = "stub_signature" + hs256_token = f"{header_b64}.{payload_b64}.{stub_signature}" api_client = ApiClient(ApiClientOptions( domains=["tenant1.auth0.com"], diff --git a/tests/test_token_store.py b/tests/test_token_store.py new file mode 100644 index 0000000..5c2a61f --- /dev/null +++ b/tests/test_token_store.py @@ -0,0 +1,408 @@ +import hashlib +import time +from typing import Optional + +import pytest + +from auth0_api_python._internal.cache_keys import ( + index_cache_key, + is_covered_by, + m2m_cache_key, + obo_cache_key, + session_fingerprint, + token_vault_cache_key, +) +from auth0_api_python.token_store import ( + AbstractTokenStore, + TokenSet, +) +from auth0_api_python.token_utils import generate_token + +# ===== AbstractTokenStore ===== + + +class ConcreteTokenStore(AbstractTokenStore): + """Minimal concrete store for testing the ABC and its helpers.""" + + def __init__(self, *, secret: str = "test-secret") -> None: # noqa: S107 + super().__init__(secret=secret) + self._store: dict[str, str] = {} + + async def get(self, key: str) -> Optional[TokenSet]: + raw = self._store.get(key) + if raw is None: + return None + return self.decrypt(key, raw) + + async def set(self, key: str, value: TokenSet) -> None: + self._store[key] = self.encrypt(key, value) + + async def delete(self, key: str) -> None: + self._store.pop(key, None) + + +@pytest.mark.asyncio +async def test_abstract_token_store_set_then_get_roundtrip(): + """Stored value is returned correctly after an encrypt/decrypt roundtrip.""" + store = ConcreteTokenStore() + value: TokenSet = {"access_token": "at", "expires_at": int(time.time()) + 3600} + + await store.set("key1", value) + + assert await store.get("key1") == value + + +@pytest.mark.asyncio +async def test_abstract_token_store_get_missing_key_returns_none(): + """get() on a missing key returns None.""" + store = ConcreteTokenStore() + + assert await store.get("missing") is None + + +@pytest.mark.asyncio +async def test_abstract_token_store_delete_removes_entry(): + """delete() removes a stored entry.""" + store = ConcreteTokenStore() + value: TokenSet = {"access_token": "at", "expires_at": int(time.time()) + 3600} + + await store.set("key1", value) + await store.delete("key1") + + assert await store.get("key1") is None + + +@pytest.mark.asyncio +async def test_abstract_token_store_delete_missing_key_is_noop(): + """delete() on a missing key does not raise.""" + store = ConcreteTokenStore() + + await store.delete("missing") + + +def test_encrypt_decrypt_roundtrip(): + """encrypt/decrypt helpers round-trip a TokenSet back to the original dict.""" + store = ConcreteTokenStore(secret="some-secret") + value: TokenSet = {"access_token": "tok", "expires_at": 9999999} + + encrypted = store.encrypt("cache-key", value) + decrypted = store.decrypt("cache-key", encrypted) + + assert decrypted == dict(value) + + +def test_encrypt_different_cache_keys_produce_different_ciphertext(): + """Two entries with the same value but different keys produce different ciphertext.""" + store = ConcreteTokenStore(secret="some-secret") + value: TokenSet = {"access_token": "tok", "expires_at": 9999999} + + enc1 = store.encrypt("key-a", value) + enc2 = store.encrypt("key-b", value) + + assert enc1 != enc2 + + +def test_encrypt_different_secrets_produce_different_ciphertext(): + """Two stores with different secrets produce different ciphertext for the same payload.""" + store_a = ConcreteTokenStore(secret="secret-a") + store_b = ConcreteTokenStore(secret="secret-b") + value: TokenSet = {"access_token": "tok", "expires_at": 9999999} + + enc_a = store_a.encrypt("key", value) + enc_b = store_b.encrypt("key", value) + + assert enc_a != enc_b + + +def test_encrypt_same_inputs_produce_different_ciphertext_each_call(): + """Each encrypt call produces unique ciphertext due to the random kid per call.""" + store = ConcreteTokenStore(secret="some-secret") + value: TokenSet = {"access_token": "tok", "expires_at": 9999999} + + enc1 = store.encrypt("key", value) + enc2 = store.encrypt("key", value) + + assert enc1 != enc2 + + +# ===== obo_cache_key ===== + + +def _obo_key(**overrides): + """Build an obo_cache_key with sensible defaults, overriding only the field under test.""" + params = { + "layout": "strict", + "sub": "sub1", + "issuer": "https://tenant.example/", + "incoming_client_id": "incoming1", + "exchange_tenant": "tenant.example.auth0.com", + "exchange_client_id": "cid", + "audience": "aud1", + "org_id": "org1", + "scopes": "read", + "session_key": "sess1", + } + params.update(overrides) + return obo_cache_key(**params) + + +def test_obo_cache_key_same_inputs_same_key(): + """Test that identical inputs produce identical cache keys.""" + assert _obo_key() == _obo_key() + + +def test_obo_cache_key_different_sub(): + """Test that a different sub produces a different cache key.""" + assert _obo_key(sub="sub1") != _obo_key(sub="sub2") + + +def test_obo_cache_key_different_audience(): + """Test that a different audience produces a different cache key.""" + assert _obo_key(audience="aud1") != _obo_key(audience="aud2") + + +def test_obo_cache_key_different_issuer(): + """Test that a different verified issuer produces a different cache key.""" + assert _obo_key(issuer="https://a/") != _obo_key(issuer="https://b/") + + +def test_obo_cache_key_different_exchange_tenant(): + """Test that a different exchange tenant produces a different cache key.""" + assert _obo_key(exchange_tenant="a.auth0.com") != _obo_key(exchange_tenant="b.auth0.com") + + +def test_obo_cache_key_different_exchange_client_id(): + """Test that a different exchange client_id produces a different cache key.""" + assert _obo_key(exchange_client_id="cid1") != _obo_key(exchange_client_id="cid2") + + +def test_obo_cache_key_different_incoming_client_id(): + """Test that two clients sharing a sub and session never share a key.""" + assert _obo_key(incoming_client_id="app1") != _obo_key(incoming_client_id="app2") + + +def test_obo_cache_key_different_layout(): + """Test that the strict and index layouts never share a key for the same fields.""" + assert _obo_key(layout="strict") != _obo_key(layout="index") + + +def test_obo_cache_key_none_org_id_matches_empty_string_org_id(): + """Test that org_id=None and org_id="" both represent an absent org and share a key.""" + assert _obo_key(org_id=None) == _obo_key(org_id="") + + +def test_obo_cache_key_absent_org_differs_from_real_org(): + """Test that an absent org_id (None/"") produces a different key than a real org_id.""" + assert _obo_key(org_id=None) != _obo_key(org_id="org_abc123") + + +def test_obo_cache_key_scope_order_and_duplicates_do_not_matter(): + """Test that scope order and duplicates normalize to the same cache key.""" + base = _obo_key(scopes="a b") + + assert _obo_key(scopes="b a") == base + assert _obo_key(scopes="a a b") == base + + +def test_obo_cache_key_different_scopes(): + """Test that a different scope set produces a different cache key.""" + assert _obo_key(scopes="read") != _obo_key(scopes="read write") + + +def test_obo_cache_key_different_session_key(): + """Test that a different session_key produces a different cache key.""" + assert _obo_key(session_key="sess1") != _obo_key(session_key="sess2") + + +# ===== m2m_cache_key ===== + + +def _m2m_key(**overrides): + """Build an m2m_cache_key with sensible defaults, overriding only the field under test.""" + params = {"tenant": "tenant.auth0.com", "client_id": "cid", "audience": "aud1", "scopes": "read"} + params.update(overrides) + return m2m_cache_key(**params) + + +def test_m2m_cache_key_same_inputs_same_key(): + """Test that identical inputs produce identical cache keys.""" + assert _m2m_key() == _m2m_key() + + +def test_m2m_cache_key_different_inputs_different_key(): + """Test that a different tenant, client, audience, or scope produces a different cache key.""" + base = _m2m_key() + + assert _m2m_key(tenant="other.auth0.com") != base + assert _m2m_key(client_id="cid2") != base + assert _m2m_key(audience="aud2") != base + assert _m2m_key(scopes="write") != base + + +# ===== token_vault_cache_key ===== + + +def test_token_vault_cache_key_same_inputs_same_key(): + """Test that identical inputs produce identical cache keys.""" + assert token_vault_cache_key("sub1", "conn1") == token_vault_cache_key("sub1", "conn1") + + +def test_token_vault_cache_key_different_inputs_different_key(): + """Test that a different sub or connection produces a different cache key.""" + base = token_vault_cache_key("sub1", "conn1") + + assert token_vault_cache_key("sub2", "conn1") != base + assert token_vault_cache_key("sub1", "conn2") != base + + +# ===== session_fingerprint ===== + + +@pytest.mark.asyncio +async def test_session_fingerprint_uses_sid_when_present(): + """Test that a token with a sid claim returns the sid as the fingerprint.""" + token = await generate_token( + domain="auth0.local", + user_id="user1", + claims={"sid": "session-abc"}, + ) + + assert session_fingerprint(token) == "session-abc" + + +@pytest.mark.asyncio +async def test_session_fingerprint_falls_back_to_jti_when_no_sid(): + """Test that a token with no sid but a jti returns the jti as the fingerprint.""" + token = await generate_token( + domain="auth0.local", + user_id="user1", + claims={"jti": "jwt-id-123"}, + ) + + assert session_fingerprint(token) == "jwt-id-123" + + +@pytest.mark.asyncio +async def test_session_fingerprint_falls_back_to_sha256_when_no_sid_or_jti(): + """Test that a token with neither sid nor jti falls back to a sha256 digest of the token.""" + token = await generate_token( + domain="auth0.local", + user_id="user1", + ) + + expected = hashlib.sha256(token.encode()).hexdigest() + + assert session_fingerprint(token) == expected + + +def test_session_fingerprint_falls_back_to_sha256_for_malformed_token(): + """Test that a malformed, non-JWT string falls back to a sha256 digest without raising.""" + malformed = "not-a-real-token" + + expected = hashlib.sha256(malformed.encode()).hexdigest() + + assert session_fingerprint(malformed) == expected + + +# ===== is_covered_by ===== + + +def test_is_covered_by_exact_match(): + """Identical granted and requested sets are covered.""" + assert is_covered_by("read:x write:x", "read:x write:x") is True + + +def test_is_covered_by_superset_granted(): + """Granted superset covers a requested subset.""" + assert is_covered_by("read:x", "read:x write:x") is True + + +def test_is_covered_by_insufficient_granted(): + """Granted does not cover a request that requires more scopes.""" + assert is_covered_by("read:x write:x", "read:x") is False + + +def test_is_covered_by_both_none(): + """None/None is covered (no scopes requested, none granted).""" + assert is_covered_by(None, None) is True + + +def test_is_covered_by_none_granted_nonempty_requested(): + """Empty granted does not cover a non-empty request.""" + assert is_covered_by("read:x", None) is False + + +def test_is_covered_by_order_and_duplicates_do_not_matter(): + """Scope order and duplicates are normalized before comparison.""" + assert is_covered_by("a b", "b a a") is True + + +# ===== index_cache_key ===== + + +def _idx_key(**overrides): + """Build an index_cache_key with sensible defaults, overriding only the field under test.""" + params = { + "sub": "sub1", + "issuer": "https://tenant.example/", + "incoming_client_id": "incoming1", + "exchange_tenant": "tenant.example.auth0.com", + "exchange_client_id": "cid", + "audience": "aud1", + "org_id": "org1", + "session_key": "sess1", + } + params.update(overrides) + return index_cache_key(**params) + + +def test_index_cache_key_same_inputs_same_key(): + """index_cache_key is deterministic.""" + assert _idx_key() == _idx_key() + + +def test_index_cache_key_differs_from_obo_cache_key(): + """index_cache_key never collides with either token layout for the same fields.""" + idx = _idx_key() + assert idx != _obo_key(layout="index", scopes=None) + assert idx != _obo_key(layout="strict", scopes=None) + + +def test_index_cache_key_different_sub(): + """Different sub produces a different index key.""" + assert _idx_key(sub="sub1") != _idx_key(sub="sub2") + + +def test_index_cache_key_different_audience(): + """Different audience produces a different index key.""" + assert _idx_key(audience="aud1") != _idx_key(audience="aud2") + + +def test_index_cache_key_different_issuer(): + """Different verified issuer produces a different index key.""" + assert _idx_key(issuer="https://a/") != _idx_key(issuer="https://b/") + + +def test_index_cache_key_different_exchange_tenant(): + """Different exchange tenant produces a different index key.""" + assert _idx_key(exchange_tenant="a.auth0.com") != _idx_key(exchange_tenant="b.auth0.com") + + +def test_index_cache_key_different_exchange_client_id(): + """Different exchange client_id produces a different index key.""" + assert _idx_key(exchange_client_id="cid1") != _idx_key(exchange_client_id="cid2") + + +def test_index_cache_key_different_incoming_client_id(): + """Two clients sharing a sub and session never share an index key.""" + assert _idx_key(incoming_client_id="app1") != _idx_key(incoming_client_id="app2") + + +def test_index_cache_key_none_org_matches_empty_string(): + """None and empty org_id produce the same index key.""" + assert _idx_key(org_id=None) == _idx_key(org_id="") + + +def test_index_cache_key_different_session(): + """Different session fingerprint produces a different index key.""" + assert _idx_key(session_key="sess1") != _idx_key(session_key="sess2")