diff --git a/src/mock_vws/_model_target_web_api.py b/src/mock_vws/_model_target_web_api.py index cab79217a..b1d3f9a5b 100644 --- a/src/mock_vws/_model_target_web_api.py +++ b/src/mock_vws/_model_target_web_api.py @@ -6,8 +6,9 @@ import secrets import uuid import zipfile +from collections.abc import Mapping from http import HTTPStatus -from typing import Any, Protocol, runtime_checkable +from typing import Any, Protocol, TypeGuard, runtime_checkable from urllib.parse import parse_qs from beartype import beartype @@ -88,6 +89,13 @@ def remove_oauth2_client_credential(self, client_id: str) -> None: ) _CLIENT_CREDENTIALS_SCOPE = "oauth2.clientcredentials.all" _MAX_CLIENT_CREDENTIALS = 100 + + +def _is_object_mapping(value: object, /) -> TypeGuard[Mapping[object, object]]: + """Return whether a value is a mapping with unchecked entries.""" + return isinstance(value, Mapping) + + # A stable mock value standing in for the user-id segment that real # Vuforia embeds in some Model Target error targets such as # ``userId:7635391``. The numeric portion is per-account in real Vuforia; @@ -290,8 +298,10 @@ def _jwt_header_error(*, bearer_token: str) -> str | None: @beartype -def _jwt_payload_error(*, bearer_token: str) -> str | None: - """Return the Vuforia error for an invalid JSON Web Token payload.""" +def _jwt_payload(*, bearer_token: str) -> tuple[Mapping[object, object], bool]: + """Decode a JSON Web Token payload and report whether it is an + object. + """ encoded_payload = bearer_token.split(sep=".")[1] try: padding = "=" * (-len(encoded_payload) % 4) @@ -300,11 +310,20 @@ def _jwt_payload_error(*, bearer_token: str) -> str | None: altchars=b"-_", validate=True, ) - payload = json.loads(s=decoded_payload) + payload: object = json.loads(s=decoded_payload) except ValueError: payload = None - if not isinstance(payload, dict): + if not _is_object_mapping(payload): + return dict[object, object](), False + return payload, True + + +@beartype +def _jwt_payload_error(*, bearer_token: str) -> str | None: + """Return the Vuforia error for an invalid JSON Web Token payload.""" + _, is_object = _jwt_payload(bearer_token=bearer_token) + if not is_object: return "Payload of JWS object is not a valid JSON object" return None @@ -334,19 +353,12 @@ def _jwt_signature_error(*, bearer_token: str) -> str | None: @beartype def _jwt_scopes(*, bearer_token: str) -> frozenset[str]: """Return scopes from a valid mock JSON Web Token.""" - encoded_payload = bearer_token.split(sep=".")[1] - padding = "=" * (-len(encoded_payload) % 4) - payload = json.loads( - s=base64.b64decode( - s=encoded_payload + padding, - altchars=b"-_", - validate=True, - ), - ) - scope = payload.get("scope", "") # pyrefly: ignore [unknown-variable-type] + empty_scopes = frozenset[str]() + payload, _ = _jwt_payload(bearer_token=bearer_token) + scope: object = payload.get("scope", "") if not isinstance(scope, str): - return frozenset() # ty: ignore[unsound-return-statement] - return frozenset(scope.split()) # ty: ignore[unsound-return-statement] + return empty_scopes + return frozenset(scope.split()) @beartype diff --git a/tests/mock_vws/test_model_target_web_api.py b/tests/mock_vws/test_model_target_web_api.py index 206f72152..57bade5da 100644 --- a/tests/mock_vws/test_model_target_web_api.py +++ b/tests/mock_vws/test_model_target_web_api.py @@ -2445,8 +2445,20 @@ def test_client_credential_authentication_errors() -> None: ) @staticmethod - def test_non_string_token_scope() -> None: - """A token with a non-string scope has no usable scopes.""" + @pytest.mark.parametrize( + argnames=("payload", "status_code"), + argvalues=[ + (b'{"scope":[]}', HTTPStatus.FORBIDDEN), + (b"[]", HTTPStatus.UNAUTHORIZED), + ], + ) + def test_invalid_token_scope( + payload: bytes, + status_code: HTTPStatus, + ) -> None: + """A non-string scope or non-object payload has no usable + scopes. + """ encoded_header = ( base64.urlsafe_b64encode( s=b'{"alg":"mock"}', @@ -2456,7 +2468,7 @@ def test_non_string_token_scope() -> None: ) encoded_payload = ( base64.urlsafe_b64encode( - s=b'{"scope":[]}', + s=payload, ) .decode(encoding="ascii") .rstrip("=") @@ -2470,7 +2482,7 @@ def test_non_string_token_scope() -> None: ) assert_model_target_status( response=response, - status_codes=HTTPStatus.FORBIDDEN, + status_codes=status_code, ) @staticmethod