Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 29 additions & 17 deletions src/mock_vws/_model_target_web_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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)
Expand All @@ -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

Expand Down Expand Up @@ -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
Expand Down
20 changes: 16 additions & 4 deletions tests/mock_vws/test_model_target_web_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"}',
Expand All @@ -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("=")
Expand All @@ -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
Expand Down
Loading