100 lines
3.0 KiB
Python
100 lines
3.0 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import logging
|
||
|
|
|
||
|
|
import jwt
|
||
|
|
from jwt import ExpiredSignatureError, InvalidAudienceError, InvalidIssuerError, InvalidTokenError
|
||
|
|
|
||
|
|
from .jwks import JWKSNetworkError, collect_all_issuers, get_jwks_public_key, resolve_jwks_uri_for_iss
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
JWT_LEEWAY = 300
|
||
|
|
|
||
|
|
|
||
|
|
class MSTeamsAuthError(Exception):
|
||
|
|
pass
|
||
|
|
|
||
|
|
|
||
|
|
class MSTeamsTokenExpiredError(MSTeamsAuthError):
|
||
|
|
pass
|
||
|
|
|
||
|
|
|
||
|
|
class MSTeamsInvalidTokenError(MSTeamsAuthError):
|
||
|
|
pass
|
||
|
|
|
||
|
|
|
||
|
|
async def verify_botframework_jwt(token: str, app_id: str, tenant_id: str) -> dict:
|
||
|
|
try:
|
||
|
|
unverified_header = jwt.get_unverified_header(token)
|
||
|
|
unverified_payload = jwt.decode(token, options={"verify_signature": False})
|
||
|
|
except Exception as e:
|
||
|
|
raise MSTeamsInvalidTokenError(f"Invalid JWT format: {e}") from e
|
||
|
|
|
||
|
|
kid = unverified_header.get("kid")
|
||
|
|
iss = unverified_payload.get("iss", "")
|
||
|
|
|
||
|
|
jwks_uri = resolve_jwks_uri_for_iss(iss)
|
||
|
|
if not jwks_uri:
|
||
|
|
raise MSTeamsInvalidTokenError(f"Unknown issuer: {iss}")
|
||
|
|
|
||
|
|
try:
|
||
|
|
public_key_pem = await get_jwks_public_key(jwks_uri, kid)
|
||
|
|
except JWKSNetworkError as e:
|
||
|
|
raise MSTeamsAuthError(f"JWKS network error: {e}") from e
|
||
|
|
except Exception as e:
|
||
|
|
raise MSTeamsInvalidTokenError(f"JWKS key not found for kid '{kid}': {e}") from e
|
||
|
|
|
||
|
|
audiences = [
|
||
|
|
app_id,
|
||
|
|
f"api://{app_id}",
|
||
|
|
"https://api.botframework.com",
|
||
|
|
]
|
||
|
|
|
||
|
|
valid_issuers = list(collect_all_issuers(tenant_id))
|
||
|
|
|
||
|
|
try:
|
||
|
|
payload = jwt.decode(
|
||
|
|
token,
|
||
|
|
public_key_pem,
|
||
|
|
algorithms=["RS256", "RS384", "RS512"],
|
||
|
|
audience=audiences,
|
||
|
|
issuer=valid_issuers,
|
||
|
|
options={"require": ["exp", "nbf", "aud", "iss"]},
|
||
|
|
leeway=JWT_LEEWAY,
|
||
|
|
)
|
||
|
|
except ExpiredSignatureError as e:
|
||
|
|
logger.warning("JWT token expired")
|
||
|
|
raise MSTeamsTokenExpiredError("Token expired") from e
|
||
|
|
except InvalidAudienceError as e:
|
||
|
|
logger.warning("JWT invalid audience")
|
||
|
|
raise MSTeamsInvalidTokenError("Invalid audience") from e
|
||
|
|
except InvalidIssuerError as e:
|
||
|
|
logger.warning("JWT invalid issuer")
|
||
|
|
raise MSTeamsInvalidTokenError("Invalid issuer") from e
|
||
|
|
except InvalidTokenError as e:
|
||
|
|
logger.warning("JWT validation failed: %s", e)
|
||
|
|
raise MSTeamsInvalidTokenError(f"Token validation failed: {e}") from e
|
||
|
|
|
||
|
|
aud_value = payload.get("aud", [])
|
||
|
|
if isinstance(aud_value, str):
|
||
|
|
aud_value = [aud_value]
|
||
|
|
if "https://api.botframework.com" in aud_value:
|
||
|
|
appid = payload.get("appid") or payload.get("azp")
|
||
|
|
if appid and str(appid) != app_id:
|
||
|
|
raise MSTeamsInvalidTokenError("appid claim mismatch")
|
||
|
|
|
||
|
|
return payload
|
||
|
|
|
||
|
|
|
||
|
|
def extract_bearer_token(authorization: str | None) -> str | None:
|
||
|
|
if not authorization:
|
||
|
|
return None
|
||
|
|
|
||
|
|
parts = authorization.split(" ")
|
||
|
|
if len(parts) == 2 and parts[0].lower() == "bearer":
|
||
|
|
token = parts[1].strip()
|
||
|
|
if token:
|
||
|
|
return token
|
||
|
|
return None
|