Repository navigation
feat: add caching for Token Vault connection token exchanges #139
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. Weβll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: feat/m2m-client-credentials
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,68 @@ | ||
| """Caching for Token Vault (federated connection) exchanges: lookup and write.""" | ||
|
|
||
| import logging | ||
| import time | ||
| from typing import Optional | ||
|
|
||
| from ..errors import TokenStoreError | ||
| from ..token_store import AbstractTokenStore, TokenSet | ||
| from .cache_keys import _normalized_scopes, token_vault_cache_key | ||
|
|
||
|
|
||
| class TokenVaultCache: | ||
| """Caches Token Vault exchange results in a token store, keyed by caller and connection. | ||
|
|
||
| ApiClient builds one TokenVaultCache when a token_store is configured and delegates | ||
| every cache decision to it. | ||
| """ | ||
|
|
||
| def __init__(self, token_store: AbstractTokenStore) -> None: | ||
| self._store = token_store | ||
|
|
||
| def cache_key(self, tenant: str, client_id: str, sub: Optional[str], connection: str) -> Optional[str]: | ||
| """Return the cache key, or None when sub is absent (cache skipped with a warning).""" | ||
| if not sub: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This guards only with Can we use |
||
| logging.warning("Access token has no usable sub claim, skipping Token Vault cache") | ||
| return None | ||
| return token_vault_cache_key(sub=sub, connection=connection, tenant=tenant, client_id=client_id) | ||
|
|
||
| async def lookup(self, key: str) -> Optional[dict]: | ||
| """Return a live cached token or None. Miss cases: absent, malformed, expired.""" | ||
| cached: Optional[TokenSet] = None | ||
| try: | ||
| cached = 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 | ||
|
|
||
| if cached is None: | ||
| return None | ||
| if "access_token" not in cached or "expires_at" not in cached: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This malformed entry check only tests for key presence. A non int Can we also require |
||
| logging.warning("Token store returned a malformed entry, treating as cache miss") | ||
| return None | ||
| if cached["expires_at"] <= int(time.time()): | ||
| return None | ||
|
|
||
| result = { | ||
| "access_token": cached["access_token"], | ||
| "expires_in": cached["expires_at"] - int(time.time()), | ||
| "expires_at": cached["expires_at"], | ||
| } | ||
| granted = cached.get("granted_scopes") | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. A fresh call always sets Should we set |
||
| if granted: | ||
| result["scope"] = granted | ||
| return result | ||
|
|
||
| async def write(self, key: str, result: dict) -> None: | ||
| """Store a freshly exchanged connection token. Failures are logged and swallowed.""" | ||
| entry: TokenSet = { | ||
| "access_token": result["access_token"], | ||
| "expires_at": result["expires_at"], | ||
| "granted_scopes": _normalized_scopes(result.get("scope")), | ||
| } | ||
| try: | ||
| await self._store.set(key, 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) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -9,6 +9,7 @@ | |
|
|
||
| from ._internal.cache_keys import m2m_cache_key | ||
| from ._internal.obo_cache import OboCache | ||
| from ._internal.token_vault_cache import TokenVaultCache | ||
| from .cache import InMemoryCache | ||
| from .config import ApiClientOptions | ||
| from .errors import ( | ||
|
|
@@ -169,6 +170,11 @@ def __init__(self, options: ApiClientOptions): | |
| if options.token_store is not None | ||
| else None | ||
| ) | ||
| self._token_vault_cache = ( | ||
| TokenVaultCache(options.token_store) | ||
| if options.token_store is not None | ||
| else None | ||
| ) | ||
|
|
||
| self._cache_ttl = options.cache_ttl_seconds | ||
|
|
||
|
|
@@ -725,18 +731,27 @@ async def verify_dpop_proof( | |
|
|
||
| return claims | ||
|
|
||
| async def get_access_token_for_connection(self, options: dict[str, Any]) -> dict[str, Any]: | ||
| async def get_access_token_for_connection( | ||
| self, | ||
| options: dict[str, Any], | ||
| *, | ||
| verified: Optional[VerifiedToken] = None, | ||
| ) -> dict[str, Any]: | ||
| """ | ||
| Retrieves a token for a connection. | ||
|
|
||
| Args: | ||
| options: Options for retrieving an access token for a connection. | ||
| Must include 'connection' and 'access_token' keys. | ||
| May optionally include 'login_hint'. | ||
| 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. | ||
|
|
||
| Raises: | ||
| GetAccessTokenForConnectionError: If there was an issue requesting the access token. | ||
| ApiError: If the token exchange endpoint returns an error. | ||
| 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. | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Minor, with a store configured and Can we add both, like the OBO doc-string does? |
||
|
|
||
| Returns: | ||
| Dictionary containing the token response with access_token, expires_in, and scope. | ||
|
|
@@ -745,8 +760,8 @@ async def get_access_token_for_connection(self, options: dict[str, Any]) -> dict | |
| SUBJECT_TYPE_ACCESS_TOKEN = "urn:ietf:params:oauth:token-type:access_token" # noqa S105 | ||
| REQUESTED_TOKEN_TYPE_FEDERATED_CONNECTION_ACCESS_TOKEN = "http://auth0.com/oauth/token-type/federated-connection-access-token" # noqa S105 | ||
| GRANT_TYPE_FEDERATED_CONNECTION_ACCESS_TOKEN = "urn:auth0:params:oauth:grant-type:token-exchange:federated-connection-access-token" # noqa S105 | ||
| connection = options.get("connection") | ||
| access_token = options.get("access_token") | ||
| connection = options.get("connection", "") | ||
| access_token = options.get("access_token", "") | ||
|
|
||
| if not connection: | ||
| raise MissingRequiredArgumentError("connection") | ||
|
|
@@ -759,6 +774,28 @@ async def get_access_token_for_connection(self, options: dict[str, Any]) -> dict | |
| if not client_id or not client_secret: | ||
| raise GetAccessTokenForConnectionError("You must configure the SDK with a client_id and client_secret to use get_access_token_for_connection.") | ||
|
|
||
| cache_key = None | ||
| if self._token_vault_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: | ||
| # Claims must belong to the token being exchanged or a caller could read another token's entry. | ||
| raise VerifyAccessTokenError( | ||
| "verified token does not match the access token being exchanged" | ||
| ) | ||
| cache_key = self._token_vault_cache.cache_key( | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The cache key is For a user with more than one linked account on the same connection, a first call with It is scoped to the same user, not a cross user leak. Can we fold a normalized |
||
| tenant=self.options.domain or "", | ||
| client_id=self.options.client_id or "", | ||
| sub=verified.claims.get("sub"), | ||
| connection=connection, | ||
| ) | ||
|
|
||
| if cache_key is not None: | ||
| hit = await self._token_vault_cache.lookup(cache_key) | ||
| if hit is not None: | ||
| return hit | ||
|
|
||
| metadata = await self._discover() | ||
|
|
||
| token_endpoint = metadata.get("token_endpoint") | ||
|
|
@@ -817,12 +854,18 @@ async def get_access_token_for_connection(self, options: dict[str, Any]) -> dict | |
| except (TypeError, ValueError): | ||
| raise ApiError("invalid_response", "expires_in is not an integer.", 502) | ||
|
|
||
| return { | ||
| result = { | ||
| "access_token": access_token, | ||
| "expires_in": expires_in, | ||
| "expires_at": int(time.time()) + expires_in, | ||
| "scope": token_endpoint_response.get("scope", "") | ||
| } | ||
|
|
||
| if cache_key is not None: | ||
| await self._token_vault_cache.write(cache_key, result) | ||
|
|
||
| return result | ||
|
|
||
| except httpx.TimeoutException as exc: | ||
| raise ApiError( | ||
| "timeout_error", | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
This example does
verified = await api_client.verify_access_token(...)and then passesverified=verified, butverify_access_tokenreturns a claimsdictwhile theverifiedkwarg expects aVerifiedToken.With a store configured the method reads
verified.access_token, so this raisesAttributeError: 'dict' object has no attribute 'access_token'. Anyone copy pasting the headline example will hit a runtime crash.Can we construct
verified=VerifiedToken(access_token=incoming_access_token, claims=claims)and reword the prose (README line 125 has the same wording)?