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
60 changes: 59 additions & 1 deletion ogc/edr/edr_api.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,13 @@
import json
import logging
import pyproj
import numpy as np
import pygeoapi.api
import pygeoapi.api.environmental_data_retrieval as pygeoedr
from functools import wraps
from http import HTTPStatus
from datetime import datetime, timezone
from typing import Tuple, List, Dict, Any, Union
from typing import Tuple, List, Dict, Any, Union, Callable

from traitlets import TraitError
from ogc import podpac as pogc
Expand All @@ -17,6 +19,8 @@
from .edr_provider import EdrProvider
from .. import settings

logger = logging.getLogger(__file__)


class EdrAPI:
"""Used to modify the default responses before returning data to the user."""
Expand All @@ -30,7 +34,56 @@ class EdrAPI:
)
SCHEMA_CLASS = "https://schemas.opengis.net/ogcapi/edr/1.1/openapi"

@staticmethod
def raise_edr_exception(func: Callable[..., Tuple[dict, int, str]]) -> Callable[..., Tuple[dict, int, str]]:
"""Decorator that raises an EDRException for non-success HTTP status codes.

The decorated function must return a tuple containing (headers, status_code, contents).
An internal error is raised if the return does not match the expected tuple.
This decorator is useful to ensure output sanitization occurs before returning responses to client.

Parameters
----------
func : Callable[..., Tuple[dict, int, str]]
Function returning a tuple containing (headers, status_code, contents).

Returns
-------
Callable[..., Tuple[dict, int, str]]
A wrapped function that validates the response status code.

Raises
------
EDRException
If the returned response has a non-success status code.
EDRException
If the function did not return the expected tuple.
"""

@wraps(func)
def wrapper(*args: Any, **kwargs: Any) -> Any:
output = func(*args, **kwargs)

# Handle (headers, status_code, contents)
if isinstance(output, tuple) and len(output) == 3:
_, status_code, content = output

if not (200 <= status_code < 300):
logger.warning(f"Unsuccessful Response ({status_code}): {content}")
status = HTTPStatus(status_code)
raise EDRException(
status_code=status.value, exception_code=status.phrase, exception_text=status.description
)
else:
logger.warning("Unhandled return type.")
raise EDRException()

return output

return wrapper

@jsonldify
@raise_edr_exception
@staticmethod
def landing_page(api: API, request: APIRequest) -> Tuple[dict, int, str]:
"""Provide the API landing page.
Expand All @@ -49,6 +102,7 @@ def landing_page(api: API, request: APIRequest) -> Tuple[dict, int, str]:
"""
return pygeoapi.api.landing_page(api, request)

@raise_edr_exception
@staticmethod
def openapi_(api: API, request: APIRequest) -> Tuple[dict, int, str]:
"""Provide the OpenAPI documentation.
Expand Down Expand Up @@ -87,6 +141,7 @@ def openapi_(api: API, request: APIRequest) -> Tuple[dict, int, str]:
else:
return headers, HTTPStatus.OK, api.openapi

@raise_edr_exception
@staticmethod
def conformance(api: API, request: APIRequest) -> Tuple[dict, int, str]:
"""Provide the conformance definition.
Expand Down Expand Up @@ -117,6 +172,7 @@ def conformance(api: API, request: APIRequest) -> Tuple[dict, int, str]:
return headers, HTTPStatus.OK, to_json(conformance, api.pretty_print)

@jsonldify
@raise_edr_exception
@staticmethod
def describe_collections(api: API, request: APIRequest, dataset: str | None = None) -> Tuple[dict, int, str]:
"""Provide the collection/collections metadata.
Expand Down Expand Up @@ -178,6 +234,7 @@ def describe_collections(api: API, request: APIRequest, dataset: str | None = No

return headers, status, to_json(collection_description, api.pretty_print)

@raise_edr_exception
@staticmethod
def get_collection_edr_instances(
api: API, request: APIRequest, dataset: str, instance_id: str | None = None
Expand Down Expand Up @@ -236,6 +293,7 @@ def get_collection_edr_instances(

return headers, status, to_json(instance_description, api.pretty_print)

@raise_edr_exception
@staticmethod
def get_collection_edr_query(
api: API,
Expand Down
32 changes: 18 additions & 14 deletions ogc/edr/test/test_edr_routes.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import os
import json
import pytest
import numpy as np
import tempfile
from pygeoapi.api import APIRequest
Expand All @@ -9,6 +10,7 @@
from werkzeug.wrappers import Request
from werkzeug.datastructures import ImmutableMultiDict
from ogc import podpac as pogc
from ogc.ogc_common import EDRException
from ogc.edr.edr_routes import EdrRoutes


Expand Down Expand Up @@ -335,14 +337,15 @@ def test_edr_routes_collection_query_invalid_type(layers: List[pogc.Layer], sing
request = mock_request(single_layer_cube_args)
edr_routes = EdrRoutes(layers=layers)

_, status, _ = edr_routes.collection_query(
request,
collection_id=collection_id,
instance_id=instance_id,
query_type="corridor",
)
with pytest.raises(EDRException) as exception_info:
edr_routes.collection_query(
request,
collection_id=collection_id,
instance_id=instance_id,
query_type="corridor",
)

assert status == HTTPStatus.BAD_REQUEST
assert exception_info.value.status_code == 400


def test_edr_routes_collection_query_invalid_bbox(layers: List[pogc.Layer], single_layer_cube_args: Dict[str, Any]):
Expand All @@ -360,14 +363,15 @@ def test_edr_routes_collection_query_invalid_bbox(layers: List[pogc.Layer], sing
request = mock_request(single_layer_cube_args)
edr_routes = EdrRoutes(layers=layers)

_, status, _ = edr_routes.collection_query(
request,
collection_id=layers[0].group,
instance_id=next(iter(layers[0].time_instances())),
query_type="cube",
)
with pytest.raises(EDRException) as exception_info:
edr_routes.collection_query(
request,
collection_id=layers[0].group,
instance_id=next(iter(layers[0].time_instances())),
query_type="cube",
)

assert status == HTTPStatus.BAD_REQUEST
assert exception_info.value.status_code == 400


def test_edr_routes_collection_query_missing_parameter(
Expand Down
19 changes: 18 additions & 1 deletion ogc/servers.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,26 @@


def _check_query_string(raw_qs: bytes) -> None:
"""Raise ValueError if the raw query string exceeds the maximum allowed length or contains invalid UTF-8."""
"""Checks the query string for malicious and invalid content.

Parameters
----------
raw_qs : bytes
The raw query string bytes.

Raises
------
ValueError
Raised if the query string exceeds a maximum allowed size.
ValueError
Raised if the query string contains null bytes.
ValueError
Raised if the query string contains invalid UTF-8 character.
"""
if len(raw_qs) > settings.MAX_QUERY_STRING_BYTES:
raise ValueError("Request query string exceeds maximum allowed length.")
if b"%00" in raw_qs:
raise ValueError("Request contains null bytes.")
try:
raw_qs.decode("utf-8")
except UnicodeDecodeError:
Expand Down
6 changes: 6 additions & 0 deletions ogc/test/test_input_security.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,12 @@ def test_check_query_string_invalid_utf8():
_check_query_string(b"SERVICE=WCS&COVERAGE=\xff\xfe")


def test_check_query_string_null_byte_injection():
"""Raises ValueError when the query string contains null bytes."""
with pytest.raises(ValueError, match="null bytes"):
_check_query_string(b"SERVICE=WCS&COVERAGE=%00")


def test_check_query_string_valid_percent_encoded():
"""Does not raise for percent-encoded non-ASCII (all bytes are ASCII in the raw QS)."""
_check_query_string(b"SERVICE=WCS&COVERAGE=%C3%A9")
Expand Down
Loading