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
1 change: 1 addition & 0 deletions newsfragments/3589.change.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Validate decoded Model Target request values before processing them.
36 changes: 24 additions & 12 deletions src/mock_vws/_model_target_web_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from urllib.parse import parse_qs

from beartype import beartype
from beartype.door import TypeHint

from mock_vws._mock_common import RequestData, json_dump
from mock_vws._services_validators.exceptions import (
Expand All @@ -25,6 +26,10 @@
OAuth2ClientCredential,
)

type _JSONValue = (
bool | int | float | str | list[_JSONValue] | dict[str, _JSONValue] | None
)

_ResponseType = tuple[int, dict[str, str], str | bytes]


Expand Down Expand Up @@ -664,10 +669,10 @@ def _client_credential_not_found(*, client_id: str) -> _ResponseType:
@beartype
def _string_list(value: object) -> list[str] | None:
"""Return a string list when ``value`` contains only strings."""
if not isinstance(value, list):
if not _is_json_array(value):
return None
strings: list[str] = []
for item in value: # pyright: ignore[reportUnknownVariableType]
for item in value:
if not isinstance(item, str):
return None
strings.append(item)
Expand Down Expand Up @@ -805,13 +810,21 @@ def delete_oauth2_client_credential(


@beartype
def _is_json_object(*, value: object) -> bool:
def _is_json_object(value: object, /) -> TypeGuard[dict[str, _JSONValue]]:
"""Return whether a decoded JSON value is an object."""
return isinstance(value, dict)
return TypeHint(hint=dict[str, _JSONValue]).is_bearable(obj=value)


@beartype
def _is_json_array(value: object, /) -> TypeGuard[list[_JSONValue]]:
"""Return whether a decoded JSON value is an array."""
return TypeHint(hint=list[_JSONValue]).is_bearable(obj=value)


@beartype
def _load_request_json(request: RequestData) -> dict[str, Any] | _ResponseType: # pyrefly: ignore [explicit-any]
def _load_request_json(
request: RequestData,
) -> dict[str, _JSONValue] | _ResponseType:
"""Load a Model Target dataset creation request body."""
content_type_header = _get_header(request=request, name="Content-Type")
content_type = (
Expand All @@ -826,7 +839,7 @@ def _load_request_json(request: RequestData) -> dict[str, Any] | _ResponseType:
details=None,
)
try:
request_json: dict[str, Any] = json.loads( # pyrefly: ignore [explicit-any]
request_json: object = json.loads(
s=request.body.decode(encoding="utf-8"),
)
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
Expand All @@ -837,7 +850,7 @@ def _load_request_json(request: RequestData) -> dict[str, Any] | _ResponseType:
target=None,
details=None,
)
if not _is_json_object(value=request_json):
if not _is_json_object(request_json):
# The required top-level fields are read from the request body, so a
# body which is valid JSON but not a JSON object is reported as
# having every required field missing.
Expand Down Expand Up @@ -1142,7 +1155,7 @@ def _configuration_states(
"error.expected.validjson"
),
}
if not _is_json_object(value=configuration):
if not _is_json_object(configuration):
return None, {
"code": "VALIDATION_ERROR",
"message": (
Expand All @@ -1151,16 +1164,15 @@ def _configuration_states(
),
}
configuration_states_value: object = configuration.get("states")
if not _is_json_object(value=configuration_states_value):
if not _is_json_object(configuration_states_value):
return None, {
"code": "VALIDATION_ERROR",
"message": (
f"/models({model_index})/stateBasedConfigurationJsonString/"
"states: error.expected.jsobject"
),
}
configuration_states: dict[str, Any] = configuration["states"] # pyrefly: ignore [explicit-any]
state_names = frozenset(configuration_states)
state_names = frozenset(configuration_states_value)
return state_names, None


Expand Down Expand Up @@ -1412,7 +1424,7 @@ def create_model_target_dataset(
is_state_based = isinstance(models_value, list) and any(
isinstance(model, dict)
and "stateBasedConfigurationJsonString" in model
for model in models_value # pyright: ignore[reportUnknownVariableType]
for model in models_value
)
if is_state_based:
state_scope_error = _require_state_based_scope(
Expand Down