diff --git a/docs/logs.md b/docs/logs.md index f7b9ca43..839997a2 100644 --- a/docs/logs.md +++ b/docs/logs.md @@ -26,6 +26,11 @@ for event in page.data: `search()` returns an immutable `LogSearchResponse` with `data`, `limit`, `has_more`, and `next_cursor`. +Both log methods reject malformed response envelopes and row collections with +`TypeError`. The response must be an object whose `data` field is a list of +objects; an empty object is not an empty page. Page limits and activity totals +must be integers, and `has_more` must be a boolean. + ## Continue a search ```python diff --git a/src/volcano_sdk/_log_response.py b/src/volcano_sdk/_log_response.py new file mode 100644 index 00000000..d374b6ef --- /dev/null +++ b/src/volcano_sdk/_log_response.py @@ -0,0 +1,84 @@ +"""Validate log response containers before generated model normalization.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import TYPE_CHECKING, cast + +if TYPE_CHECKING: + from .models import JSONValue + +INVALID_LOG_RESPONSE = "Expected a complete log response" + + +def response_values(payload: object) -> Mapping[str, object]: + """Require a JSON object for the response envelope. + + Raises: + TypeError: The response envelope is not an object. + + Returns: + The response fields without normalizing their values. + + """ + if not isinstance(payload, Mapping): + raise TypeError(INVALID_LOG_RESPONSE) + return cast("Mapping[str, object]", payload) + + +def response_data(values: Mapping[str, object]) -> tuple[Mapping[str, JSONValue], ...]: + """Require a list of objects without coercing empty mappings into lists. + + Raises: + TypeError: The rows are not a list of objects. + + Returns: + The validated rows in their original order. + + """ + raw_data = values.get("data") + if not isinstance(raw_data, list): + raise TypeError(INVALID_LOG_RESPONSE) + data = cast("list[object]", raw_data) + if any(not isinstance(item, Mapping) for item in data): + raise TypeError(INVALID_LOG_RESPONSE) + return tuple(cast("Mapping[str, JSONValue]", item) for item in data) + + +def search_metadata(values: Mapping[str, object]) -> tuple[int, bool, str | None]: + """Require search pagination fields without coercing their values. + + Raises: + TypeError: Required metadata is missing or has an invalid type. + + Returns: + The page limit, continuation flag, and optional cursor. + + """ + limit = values.get("limit") + has_more = values.get("has_more") + next_cursor = values.get("next_cursor") + if ( + not isinstance(limit, int) + or isinstance(limit, bool) + or not isinstance(has_more, bool) + or (next_cursor is not None and not isinstance(next_cursor, str)) + ): + raise TypeError(INVALID_LOG_RESPONSE) + return limit, has_more, next_cursor + + +def activity_total(values: Mapping[str, object]) -> int: + """Require the activity total without treating booleans as integers. + + Raises: + TypeError: The total is missing or is not an integer. + + Returns: + The total reported by the server. + + """ + total = values.get("total") + if not isinstance(total, int) or isinstance(total, bool): + raise TypeError(INVALID_LOG_RESPONSE) + return total diff --git a/src/volcano_sdk/_transport.py b/src/volcano_sdk/_transport.py index 14c09eaa..bec37471 100644 --- a/src/volcano_sdk/_transport.py +++ b/src/volcano_sdk/_transport.py @@ -232,6 +232,12 @@ UploadStorageObjectFilesBody, ) from ._generated.types import UNSET, File +from ._log_response import ( + activity_total, + response_data, + response_values, + search_metadata, +) from .errors import ( AuthenticationError, ConflictError, @@ -1708,6 +1714,17 @@ def invoke_function_url( ) return self._raw_response(response) + @staticmethod + def _validate_log_response( + response: _RawHTTPResponse, + metadata: Callable[[Mapping[str, object]], object], + ) -> None: + if response.status_code == HTTP_OK: + payload = GeneratedTransport._raw_response(response).payload + values = response_values(payload) + _ = response_data(values) + _ = metadata(values) + def search_project_logs( self, *, @@ -1724,6 +1741,7 @@ def search_project_logs( raw_response = client.get_httpx_client().request(**request_kwargs) if raw_response.status_code == HTTP_UNAUTHORIZED: return self._raw_response(raw_response) + self._validate_log_response(raw_response, search_metadata) response = build_log_search_response( client=client, response=raw_response, @@ -1746,6 +1764,7 @@ def get_project_log_activity( raw_response = client.get_httpx_client().request(**request_kwargs) if raw_response.status_code == HTTP_UNAUTHORIZED: return self._raw_response(raw_response) + self._validate_log_response(raw_response, activity_total) response = build_log_activity_response( client=client, response=raw_response, diff --git a/src/volcano_sdk/logs.py b/src/volcano_sdk/logs.py index 5303a5aa..326ebce6 100644 --- a/src/volcano_sdk/logs.py +++ b/src/volcano_sdk/logs.py @@ -5,6 +5,12 @@ from collections.abc import Mapping from typing import TYPE_CHECKING, Protocol, cast +from ._log_response import ( + activity_total, + response_data, + response_values, + search_metadata, +) from ._transport import TransportResponse, invoke, response_payload from .models import JSONValue, LogActivityResponse, LogSearchResponse, _freeze_json @@ -13,7 +19,6 @@ _INVALID_PROJECT_ID = "project_id must be a non-empty string" _INVALID_LOG_REQUEST = "Log request must be a mapping" -_INVALID_LOG_RESPONSE = "Expected a complete log response" class LogsTransport(Protocol): @@ -110,36 +115,11 @@ def _log_request( return project_id, cast("Mapping[str, JSONValue]", snapshot) -def _response_values(payload: object) -> Mapping[str, object]: - if not isinstance(payload, Mapping): - raise TypeError(_INVALID_LOG_RESPONSE) - return cast("Mapping[str, object]", payload) - - -def _response_data(values: Mapping[str, object]) -> tuple[Mapping[str, JSONValue], ...]: - raw_data = values.get("data") - if not isinstance(raw_data, list): - raise TypeError(_INVALID_LOG_RESPONSE) - data = cast("list[object]", raw_data) - if any(not isinstance(item, Mapping) for item in data): - raise TypeError(_INVALID_LOG_RESPONSE) - return tuple(cast("Mapping[str, JSONValue]", item) for item in data) - - def _search_response(payload: object) -> LogSearchResponse: - values = _response_values(payload) - limit = values.get("limit") - has_more = values.get("has_more") - next_cursor = values.get("next_cursor") - if ( - not isinstance(limit, int) - or isinstance(limit, bool) - or not isinstance(has_more, bool) - or (next_cursor is not None and not isinstance(next_cursor, str)) - ): - raise TypeError(_INVALID_LOG_RESPONSE) + values = response_values(payload) + limit, has_more, next_cursor = search_metadata(values) return LogSearchResponse( - data=_response_data(values), + data=response_data(values), limit=limit, has_more=has_more, next_cursor=next_cursor, @@ -147,8 +127,5 @@ def _search_response(payload: object) -> LogSearchResponse: def _activity_response(payload: object) -> LogActivityResponse: - values = _response_values(payload) - total = values.get("total") - if not isinstance(total, int) or isinstance(total, bool): - raise TypeError(_INVALID_LOG_RESPONSE) - return LogActivityResponse(data=_response_data(values), total=total) + values = response_values(payload) + return LogActivityResponse(data=response_data(values), total=activity_total(values)) diff --git a/tests/unit/test_log_response_validation.py b/tests/unit/test_log_response_validation.py new file mode 100644 index 00000000..9f11f7a7 --- /dev/null +++ b/tests/unit/test_log_response_validation.py @@ -0,0 +1,114 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import httpx +import pytest +from test_logs import FakeLogsTransport, FakeResponse, logs_client +from test_logs_refresh import make_client + +if TYPE_CHECKING: + from volcano_sdk import VolcanoClient + +PROJECT_ID = "00000000-0000-4000-8000-000000000001" +REQUEST = {"resource": {"type": "function"}} + + +def response_client(payload: object, *, native_transport: bool) -> VolcanoClient: + def handle(_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=payload) + + if native_transport: + return make_client(handle) + transport = FakeLogsTransport() + transport.search_response = FakeResponse(200, payload, {}) + transport.activity_response = FakeResponse(200, payload, {}) + return logs_client(transport) + + +@pytest.mark.parametrize("native_transport", [False, True]) +@pytest.mark.parametrize("operation", ["search", "activity"]) +@pytest.mark.parametrize("payload", [None, [], "invalid", 1, True]) +def test_logs_reject_non_object_responses( + operation: str, payload: object, *, native_transport: bool +) -> None: + logs = response_client(payload, native_transport=native_transport).logs + read = logs.search if operation == "search" else logs.activity + + with pytest.raises(TypeError, match="Expected a complete log response"): + read(PROJECT_ID, REQUEST) + + +@pytest.mark.parametrize("native_transport", [False, True]) +@pytest.mark.parametrize("operation", ["search", "activity"]) +@pytest.mark.parametrize("data", [None, {}, "invalid", [None], [1], [[], {}]]) +def test_logs_reject_invalid_response_rows( + operation: str, data: object, *, native_transport: bool +) -> None: + payload = {"data": data, "limit": 25, "has_more": False, "total": 0} + logs = response_client(payload, native_transport=native_transport).logs + read = logs.search if operation == "search" else logs.activity + + with pytest.raises(TypeError, match="Expected a complete log response"): + read(PROJECT_ID, REQUEST) + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("limit", None), + ("limit", "25"), + ("limit", 25.0), + ("limit", True), + ("has_more", None), + ("has_more", "false"), + ("has_more", 0), + ("next_cursor", 1), + ("next_cursor", []), + ], +) +@pytest.mark.parametrize("native_transport", [False, True]) +def test_logs_search_rejects_invalid_page_metadata( + field: str, value: object, *, native_transport: bool +) -> None: + payload: dict[str, object] = {"data": [], "limit": 25, "has_more": False} + payload[field] = value + client = response_client(payload, native_transport=native_transport) + with pytest.raises(TypeError, match="Expected a complete log response"): + client.logs.search(PROJECT_ID, REQUEST) + + +@pytest.mark.parametrize("native_transport", [False, True]) +@pytest.mark.parametrize("total", [None, "0", 0.0, True]) +def test_logs_activity_rejects_invalid_totals( + total: object, *, native_transport: bool +) -> None: + client = response_client( + {"data": [], "total": total}, native_transport=native_transport + ) + with pytest.raises(TypeError, match="Expected a complete log response"): + client.logs.activity(PROJECT_ID, REQUEST) + + +@pytest.mark.parametrize("native_transport", [False, True]) +@pytest.mark.parametrize("field", ["data", "limit", "has_more"]) +def test_logs_search_rejects_missing_required_fields( + field: str, *, native_transport: bool +) -> None: + payload: dict[str, object] = {"data": [], "limit": 25, "has_more": False} + del payload[field] + client = response_client(payload, native_transport=native_transport) + with pytest.raises(TypeError, match="Expected a complete log response"): + client.logs.search(PROJECT_ID, REQUEST) + + +@pytest.mark.parametrize("native_transport", [False, True]) +@pytest.mark.parametrize("field", ["data", "total"]) +def test_logs_activity_rejects_missing_required_fields( + field: str, *, native_transport: bool +) -> None: + payload: dict[str, object] = {"data": [], "total": 0} + del payload[field] + client = response_client(payload, native_transport=native_transport) + with pytest.raises(TypeError, match="Expected a complete log response"): + client.logs.activity(PROJECT_ID, REQUEST)