fix: add token introspection support in validate_token for improved token validation
This commit is contained in:
+47
-5
@@ -1,3 +1,4 @@
|
|||||||
|
import logging
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
from urllib.parse import urljoin
|
from urllib.parse import urljoin
|
||||||
@@ -9,6 +10,8 @@ from jose import JWTError, jwk, jwt
|
|||||||
from pydantic import BaseModel, Field, HttpUrl
|
from pydantic import BaseModel, Field, HttpUrl
|
||||||
from pydantic_settings import BaseSettings
|
from pydantic_settings import BaseSettings
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
security = HTTPBearer(auto_error=False)
|
security = HTTPBearer(auto_error=False)
|
||||||
|
|
||||||
class Settings(BaseSettings):
|
class Settings(BaseSettings):
|
||||||
@@ -75,15 +78,48 @@ def get_signing_key(token: str) -> Dict[str, Any]:
|
|||||||
|
|
||||||
raise HTTPException(status_code=401, detail="Unable to find matching JWKS key")
|
raise HTTPException(status_code=401, detail="Unable to find matching JWKS key")
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def get_introspection_endpoint() -> Optional[str]:
|
||||||
|
discovery = get_discovery()
|
||||||
|
return discovery.get("introspection_endpoint")
|
||||||
|
|
||||||
|
def introspect_token(token: str) -> Dict[str, Any]:
|
||||||
|
introspection_endpoint = get_introspection_endpoint()
|
||||||
|
if not introspection_endpoint:
|
||||||
|
raise HTTPException(status_code=401, detail="Unable to introspect token")
|
||||||
|
|
||||||
|
data = {"token": token, "token_type_hint": "access_token"}
|
||||||
|
auth = None
|
||||||
|
if settings.oidc_client_secret:
|
||||||
|
auth = (settings.oidc_client_id, settings.oidc_client_secret)
|
||||||
|
else:
|
||||||
|
data["client_id"] = settings.oidc_client_id
|
||||||
|
|
||||||
|
with httpx.Client(timeout=10.0) as client:
|
||||||
|
response = client.post(introspection_endpoint, data=data, headers={"Accept": "application/json"}, auth=auth)
|
||||||
|
|
||||||
|
if response.status_code != 200:
|
||||||
|
raise HTTPException(status_code=401, detail="Unable to introspect token")
|
||||||
|
|
||||||
|
introspection = response.json()
|
||||||
|
if not introspection.get("active"):
|
||||||
|
raise HTTPException(status_code=401, detail="Invalid or expired token")
|
||||||
|
|
||||||
|
discovery = get_discovery()
|
||||||
|
if "iss" not in introspection:
|
||||||
|
introspection["iss"] = discovery.get("issuer")
|
||||||
|
|
||||||
|
return introspection
|
||||||
|
|
||||||
async def validate_token(credentials: HTTPAuthorizationCredentials = Security(security)) -> TokenClaims:
|
async def validate_token(credentials: HTTPAuthorizationCredentials = Security(security)) -> TokenClaims:
|
||||||
if not credentials or credentials.scheme.lower() != "bearer":
|
if not credentials or credentials.scheme.lower() != "bearer":
|
||||||
raise HTTPException(status_code=401, detail="Missing Bearer authorization header")
|
raise HTTPException(status_code=401, detail="Missing Bearer authorization header")
|
||||||
|
|
||||||
token = credentials.credentials
|
token = credentials.credentials
|
||||||
key = get_signing_key(token)
|
|
||||||
public_key = jwk.construct(key)
|
|
||||||
discovery = get_discovery()
|
|
||||||
try:
|
try:
|
||||||
|
key = get_signing_key(token)
|
||||||
|
public_key = jwk.construct(key)
|
||||||
|
discovery = get_discovery()
|
||||||
# Only pass `audience` to the decoder if configured. Some providers
|
# Only pass `audience` to the decoder if configured. Some providers
|
||||||
# (or local development setups) may not include the aud claim in a
|
# (or local development setups) may not include the aud claim in a
|
||||||
# way that matches your API identifier; in that case leave
|
# way that matches your API identifier; in that case leave
|
||||||
@@ -96,8 +132,14 @@ async def validate_token(credentials: HTTPAuthorizationCredentials = Security(se
|
|||||||
jwt_kwargs["audience"] = settings.oidc_audience
|
jwt_kwargs["audience"] = settings.oidc_audience
|
||||||
|
|
||||||
verified = jwt.decode(token, public_key, **jwt_kwargs)
|
verified = jwt.decode(token, public_key, **jwt_kwargs)
|
||||||
except JWTError as exc:
|
except JWTError:
|
||||||
raise HTTPException(status_code=401, detail="Invalid or expired token") from exc
|
logger.exception("JWT validation failed, attempting introspection")
|
||||||
|
verified = introspect_token(token)
|
||||||
|
|
||||||
|
try:
|
||||||
|
return TokenClaims(**verified)
|
||||||
|
except Exception as exc:
|
||||||
|
raise HTTPException(status_code=401, detail="Unable to parse token claims") from exc
|
||||||
|
|
||||||
try:
|
try:
|
||||||
return TokenClaims(**verified)
|
return TokenClaims(**verified)
|
||||||
|
|||||||
Reference in New Issue
Block a user