diff --git a/ogc/edr/edr_api.py b/ogc/edr/edr_api.py index 4146f32..2e840b1 100644 --- a/ogc/edr/edr_api.py +++ b/ogc/edr/edr_api.py @@ -9,6 +9,7 @@ from traitlets import TraitError from ogc import podpac as pogc +from ogc.ogc_common import EDRException from pygeoapi.plugin import load_plugin from pygeoapi.util import filter_dict_by_key_value, to_json, get_provider_by_type from pygeoapi.api import API, APIRequest @@ -64,6 +65,9 @@ def openapi_(api: API, request: APIRequest) -> Tuple[dict, int, str]: Tuple[dict, int, str] Headers, HTTP Status, and Content returned as a tuple. """ + if request._args.get("ui") is not None and request._args.get("ui") != "redoc": + raise EDRException(status_code=400, exception_code="InvalidQuery", exception_text="") + html_path = "openapi/redoc.html" if request._args.get("ui") == "redoc" else "openapi/swagger.html" headers = request.get_response_headers(**api.api_headers) diff --git a/ogc/ogc_common.py b/ogc/ogc_common.py index db091ad..8e17dbb 100755 --- a/ogc/ogc_common.py +++ b/ogc/ogc_common.py @@ -1,6 +1,6 @@ import logging import string - +import json import lxml import lxml.etree import numpy as np @@ -16,6 +16,8 @@ "urn:ogc:def:crs:EPSG::26910", # used in an example, "urn:ogc:def:crs:EPSG::4326", # used in an example, ) +INTERNAL_APPLICATION_ERROR = "Internal application error." +NO_APPLICABLE_CODE = "NoApplicableCode" class EscapeFormatter(string.Formatter): @@ -69,12 +71,16 @@ class OutputFormat(XMLNode): value = tl.Unicode(default_value=None, allow_none=True) + # Allowed values of None mean that all values are allowed, an empty list means no values allowed + allowed_values = tl.List(tl.Unicode(), default_value=None, allow_none=True) + def validate(self): assert bool(self.value) is True, "error validating output format" - # Can check here for specific allowed formats if desired, prob. not necessary. - def to_xml(self): - return "%s" % self.value + if self.allowed_values is not None: + assert self.value is not None and self.value.lower() in [ + allowed_value.lower() for allowed_value in self.allowed_values + ], "error validating output format, value not in allowed values" class BoundingBox(XMLNode): @@ -121,8 +127,8 @@ def to_xml(self): class WCSException(Exception): def __init__( self, - exception_text="Internal application error.", - exception_code="NoApplicableCode", + exception_text=INTERNAL_APPLICATION_ERROR, + exception_code=NO_APPLICABLE_CODE, locator="", ): """ @@ -161,8 +167,8 @@ class WMTSException(WCSException): def __init__( self, - exception_text="Internal application error.", - exception_code="NoApplicableCode", + exception_text=INTERNAL_APPLICATION_ERROR, + exception_code=NO_APPLICABLE_CODE, locator="", ): """ @@ -171,3 +177,36 @@ def __init__( 'OperationNotSupported', 'TileOutOfRange' """ super().__init__(exception_text, exception_code, locator) + + +class EDRException(Exception): + def __init__( + self, + status_code=500, + exception_code=NO_APPLICABLE_CODE, + exception_text=INTERNAL_APPLICATION_ERROR, + ): + """ + exception_code: 'NoApplicableCode', 'NotFound', 'InvalidParameterValue', 'InvalidQuery' + """ + super().__init__(status_code, exception_text, exception_code) + + self.status_code = status_code + self.exception_code = exception_code + self.exception_text = exception_text + + def to_json(self) -> str: + """Return JSON string for the exception. + + Returns + ------- + str + The exception in JSON string format. + """ + return json.dumps( + { + "code": self.status_code, + "type": self.exception_code, + "description": self.exception_text, + } + ) diff --git a/ogc/servers.py b/ogc/servers.py index 51b2a1b..071d230 100755 --- a/ogc/servers.py +++ b/ogc/servers.py @@ -14,22 +14,23 @@ from typing import Callable from werkzeug.datastructures import ImmutableMultiDict -from ogc.ogc_common import WCSException +from ogc.ogc_common import WCSException, EDRException from pygeoapi.api import APIRequest from pygeoapi.util import get_api_rules from . import settings logger = logging.getLogger(__name__) +INVALID_ARGUMENTS = "Invalid arguments" def _check_query_string(raw_qs: bytes) -> None: - """Raise WCSException if the raw query string exceeds the maximum allowed length or contains invalid UTF-8.""" + """Raise ValueError if the raw query string exceeds the maximum allowed length or contains invalid UTF-8.""" if len(raw_qs) > settings.MAX_QUERY_STRING_BYTES: - raise WCSException("Request query string exceeds maximum allowed length.") + raise ValueError("Request query string exceeds maximum allowed length.") try: raw_qs.decode("utf-8") except UnicodeDecodeError: - raise WCSException("Request contains invalid UTF-8 encoding.") + raise ValueError("Request contains invalid UTF-8 encoding.") def respond_xml(doc, status=200): @@ -259,8 +260,9 @@ def ogc_render(self, ogc_idx): try: _check_query_string(request.query_string) - except WCSException as e: - return respond_xml(e.to_xml(), status=400) + except ValueError as e: + ee = WCSException(exception_code="InvalidParameterValue", exception_text=str(e)) + return respond_xml(ee.to_xml(), status=400) if not request.args: return self.home_func(ogc.endpoint) @@ -346,8 +348,9 @@ def wrapper(*args, **kwargs) -> Response: try: _check_query_string(request.query_string) - except WCSException as e: - return respond_xml(e.to_xml(), status=400) + except ValueError as e: + ee = EDRException(status_code=400, exception_code="InvalidQuery", exception_text=str(e)) + return Response(ee.to_json(), status=ee.status_code) try: # We'll filter out any characters from URl parameter values that @@ -369,11 +372,21 @@ def wrapper(*args, **kwargs) -> Response: # Replace format with its lowercase version to match pygeoapi expectations query_type = kwargs.get("query_type") default_format = settings.JSON + query_formats = [settings.HTML, settings.JSON] + if query_type is not None: default_format = settings.EDR_QUERY_DEFAULTS.get(query_type, default_format) - format_argument = filtered_args.get("f", default_format) - if format_argument is not None: - filtered_args["f"] = format_argument.lower() + query_formats = settings.EDR_QUERY_FORMATS.get(query_type, []) + + format_argument = filtered_args.get("f", default_format).lower() + filtered_args["f"] = format_argument + + if format_argument not in [item.lower() for item in query_formats]: + raise EDRException( + status_code=400, + exception_code="InvalidQuery", + exception_text=INVALID_ARGUMENTS, + ) filtered_args["base_url"] = ( xml.sax.saxutils.escape(request.base_url, {'"': """}) if request.base_url else None @@ -388,13 +401,13 @@ def wrapper(*args, **kwargs) -> Response: if headers: response.headers = headers return response - except WCSException as e: - logger.exception("OGC: server.edr_render WCSException: %s", str(e)) - return respond_xml(e.to_xml(), status=400) + except EDRException as e: + logger.exception("OGC: server.edr_render EDRException: %s", str(e)) + return Response(e.to_json(), status=e.status_code) except Exception as e: # noqa: B902 logger.exception("OGC: server.edr_render Exception: %s", str(e)) - ee = WCSException() - return respond_xml(ee.to_xml(), status=500) + ee = EDRException() + return Response(ee.to_json(), status=ee.status_code) return wrapper diff --git a/ogc/test/test_input_security.py b/ogc/test/test_input_security.py index 3093681..ac05783 100644 --- a/ogc/test/test_input_security.py +++ b/ogc/test/test_input_security.py @@ -8,7 +8,6 @@ from ogc import servers, core, settings from ogc import podpac as pogc -from ogc.ogc_common import WCSException from ogc.servers import _check_query_string # --------------------------------------------------------------------------- @@ -42,9 +41,9 @@ def client(): def test_check_query_string_overflow(): - """Raises WCSException when byte length exceeds MAX_QUERY_STRING_BYTES.""" + """Raises ValueError when byte length exceeds MAX_QUERY_STRING_BYTES.""" oversized = b"A" * (settings.MAX_QUERY_STRING_BYTES + 1) - with pytest.raises(WCSException, match="maximum allowed length"): + with pytest.raises(ValueError, match="maximum allowed length"): _check_query_string(oversized) @@ -60,8 +59,8 @@ def test_check_query_string_exactly_at_limit(): def test_check_query_string_invalid_utf8(): - """Raises WCSException when the query string contains raw non-UTF-8 bytes.""" - with pytest.raises(WCSException, match="invalid UTF-8 encoding"): + """Raises ValueError when the query string contains raw non-UTF-8 bytes.""" + with pytest.raises(ValueError, match="invalid UTF-8 encoding"): _check_query_string(b"SERVICE=WCS&COVERAGE=\xff\xfe") diff --git a/ogc/test/test_input_validation.py b/ogc/test/test_input_validation.py new file mode 100644 index 0000000..0874d3e --- /dev/null +++ b/ogc/test/test_input_validation.py @@ -0,0 +1,614 @@ +import pytest +import importlib +import podpac +import datetime +import numpy as np +from flask.testing import FlaskClient +from itertools import chain +from collections.abc import Iterator +from unittest.mock import patch +from ogc import core +from ogc import servers +from ogc import settings +from ogc import podpac as pogc +from ogc.settings import EDR_TIME_INSTANCE_DIMENSION + +lat = np.linspace(90, -90, 11) +lon = np.linspace(-180, 180, 21) +time = np.array(["2025-10-24T12:00:00"], dtype="datetime64") +instance = np.array(["2025-10-24T00:00:00"], dtype="datetime64") +data_static = np.random.default_rng(1).random((11, 21)) +coords_static = podpac.Coordinates([lat, lon], dims=["lat", "lon"]) +data_with_time = np.random.default_rng(1).random((11, 21, 1)) +coords_with_time = podpac.Coordinates([lat, lon, time], dims=["lat", "lon", "time"]) +data_with_instance = np.random.default_rng(1).random((11, 21, 1, 1)) +coords_with_instance = podpac.Coordinates( + [lat, lon, time, instance], dims=["lat", "lon", "time", EDR_TIME_INSTANCE_DIMENSION] +) + +# Define a layer which does not include temporal coordinates +node_static = podpac.data.Array(source=data_static, coordinates=coords_static) +layer_static = pogc.Layer( + node=node_static, + identifier="layerStatic", + title="Layer Static", + abstract="Layer Static", + group="Layers", +) + +# Define a layer which includes time coordinates +node_time = podpac.data.Array(source=data_with_time, coordinates=coords_with_time) +layer_time = pogc.Layer( + node=node_time, + identifier="layerTime", + title="Layer Time", + abstract="Layer Time", + group="Layers", + valid_times=[dt.astype(datetime.datetime) for dt in time], +) + +# Define a layer which includes both time coordinates and instance coordinates +node_instance = podpac.data.Array(source=data_with_instance, coordinates=coords_with_instance) +layer_instance = pogc.Layer( + node=node_instance, + identifier="layerInstance", + title="Layer Instance", + abstract="Layer Instance", + group="Layers", + valid_times=[dt.astype(datetime.datetime) for dt in time], +) + + +@pytest.fixture +def enable_all_formats_in_env(): + """Test client for FlaskServer with all formats enabled.""" + with patch.dict("os.environ", {"OGC_SUPPORTED_FORMATS": "wms,wcs,wmts,edr"}): + importlib.reload(settings) + yield + importlib.reload(settings) + + +@pytest.fixture +def client(): + """ + Create a test client for the Flask server. + + Yields + ------ + client : FlaskClient + A test client for the Flask server. + """ + # Create an OGC instance with the test layers + ogc = core.OGC(layers=[layer_static, layer_time, layer_instance]) + + # Create a FlaskServer instance + app = servers.FlaskServer(__name__, ogcs=[ogc]) + app.config.update({"TESTING": True}) + yield app.test_client() + + +def make_valid_ogc_wms_get_capabilities_args() -> dict: + """Valid argument dictionary for WMS get capabilities. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMS", + "REQUEST": "GetCapabilities", + } + + +def make_valid_ogc_wms_get_feature_info_args() -> dict: + """Valid argument dictionary for WMS get feature info. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMS", + "REQUEST": "GetFeatureInfo", + "VERSION": "1.3.0", + } + + +def make_valid_ogc_wms_get_legend_graphic_args(layer: str) -> dict: + """Valid argument dictionary for WMS get legend graphic. + + Parameters + ---------- + layer : str + Identifier for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMS", + "REQUEST": "GetLegendGraphic", + "VERSION": "1.3.0", + "LAYER": layer, + } + + +def make_valid_ogc_wms_get_map_args(layer: str, time: str) -> dict: + """Valid argument dictionary for WMS get map. + + Parameters + ---------- + layer : str + Identifier for the layer. + time : str + Available time for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMS", + "REQUEST": "GetMap", + "VERSION": "1.3.0", + "LAYERS": [layer], + "CRS": "EPSG:4326", + "BBOX": "-180,-90,180,90", + "FORMAT": "image/png", + "TIME": time, + "HEIGHT": 512, + "WIDTH": 512, + } + + +def make_valid_ogc_wcs_describe_coverage_args(layer: str) -> dict: + """Valid argument dictionary for WCS describe coverage. + + Parameters + ---------- + layer : str + Identifier for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WCS", + "REQUEST": "DescribeCoverage", + "VERSION": "1.0.0", + "COVERAGE": layer, + } + + +def make_valid_ogc_wcs_get_capabilities_args() -> dict: + """Valid argument dictionary for WCS get capabilities. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WCS", + "REQUEST": "GetCapabilities", + } + + +def make_valid_ogc_wcs_get_coverage_args(layer: str, time: str) -> dict: + """Valid argument dictionary for WCS get coverage. + + Parameters + ---------- + layer : str + Identifier for the layer. + time : str + Available time for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WCS", + "REQUEST": "GetCoverage", + "VERSION": "1.0.0", + "COVERAGE": layer, + "REQUEST_CRS": "EPSG:4326", + "CRS": "EPSG:4326", + "BBOX": "-180,-90,180,90", + "FORMAT": "geotiff", + "TIME": time, + "HEIGHT": 512, + "WIDTH": 512, + } + + +def make_valid_ogc_wmts_get_capabilities_args() -> dict: + """Valid argument dictionary for WMTS get capabilities. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMTS", + "REQUEST": "GetCapabilities", + } + + +def make_valid_ogc_wmts_get_feature_info_args() -> dict: + """Valid argument dictionary for WMTS get feature info. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMTS", + "REQUEST": "GetFeatureInfo", + "VERSION": "1.0.0", + } + + +def make_valid_ogc_wmts_get_tile_args(layer: str, time: str) -> dict: + """Valid argument dictionary for WMTS get tile. + + Parameters + ---------- + layer : str + Identifier for the layer. + time : str + Available time for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "SERVICE": "WMTS", + "REQUEST": "GetTile", + "VERSION": "1.0.0", + "LAYER": layer, + "TILEROW": "0", + "TILECOL": "0", + "TILEMATRIX": "0", + "TILEMATRIXSET": "WebMercatorQuad", + "FORMAT": "image/png", + "TIME": time, + } + + +def make_valid_ogc_edr_format_args() -> dict: + """Valid argument dictionary for EDR requests using only format. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "json", + } + + +def make_valid_ogc_edr_api_args() -> dict: + """Valid argument dictionary for EDR api requests. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "json", + "ui": "redoc", + } + + +def make_valid_ogc_edr_static_cube_args(layer: str) -> dict: + """Valid argument dictionary for EDR cube query without time or instances. + + Parameters + ---------- + layer : str + Identifier for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "CoverageJSON", + "bbox": "-180,-90,180,90", + "crs": "http://www.opengis.net/def/crs/OGC/1.3/CRS84", + "parameter-name": layer, + "resolution-x": 512, + "resolution-y": 512, + } + + +def make_valid_ogc_edr_static_area_args(layer: str) -> dict: + """Valid argument dictionary for EDR area query without time or instances. + + Parameters + ---------- + layer : str + Identifier for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "CoverageJSON", + "coords": "POLYGON((-180 90, -180 -90, 180 -90, 180 90, -180 90))", + "crs": "http://www.opengis.net/def/crs/OGC/1.3/CRS84", + "parameter-name": layer, + "resolution-x": 512, + "resolution-y": 512, + } + + +def make_valid_ogc_edr_static_position_args(layer: str) -> dict: + """Valid argument dictionary for EDR position query without time or instances. + + Parameters + ---------- + layer : str + Identifier for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "CoverageJSON", + "coords": "POINT(40 50)", + "crs": "http://www.opengis.net/def/crs/OGC/1.3/CRS84", + "parameter-name": layer, + } + + +def make_valid_ogc_edr_instance_cube_args(layer: str, time: str) -> dict: + """Valid argument dictionary for EDR cube query with time and instances. + + Parameters + ---------- + layer : str + Identifier for the layer. + time : str + Available time for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "CoverageJSON", + "bbox": "-180,-90,180,90", + "crs": "http://www.opengis.net/def/crs/OGC/1.3/CRS84", + "datetime": time, + "parameter-name": layer, + "resolution-x": 512, + "resolution-y": 512, + } + + +def make_valid_ogc_edr_instance_area_args(layer: str, time: str) -> dict: + """Valid argument dictionary for EDR area query with time and instances. + + Parameters + ---------- + layer : str + Identifier for the layer. + time : str + Available time for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "CoverageJSON", + "coords": "POLYGON((-180 90, -180 -90, 180 -90, 180 90, -180 90))", + "crs": "http://www.opengis.net/def/crs/OGC/1.3/CRS84", + "datetime": time, + "parameter-name": layer, + "resolution-x": 512, + "resolution-y": 512, + } + + +def make_valid_ogc_edr_instance_position_args(layer: str, time: str) -> dict: + """Valid argument dictionary for EDR position query with time and instances. + + Parameters + ---------- + layer : str + Identifier for the layer. + time : str + Available time for the layer. + + Returns + ------- + dict + The argument dictionary. + """ + return { + "f": "CoverageJSON", + "coords": "POINT(40 50)", + "crs": "http://www.opengis.net/def/crs/OGC/1.3/CRS84", + "datetime": time, + "parameter-name": layer, + } + + +class BaseValidation: + @staticmethod + def generate_cases(url: str, params: dict) -> Iterator[tuple[str, dict, bool]]: + """Generate cases for each argument group. + The first case uses all valid arguments and returns true. + The remaining cases change a single argument to "invalid" and return false. + + Parameters + ---------- + url: str + The URL to request from. + params : dict + The valid argument group. + + Yields + ------ + Iterator[tuple[dict, bool]] + An iterator containing the URL, updated arguments, and a boolean whether it is valid or not. + """ + yield url, params, True + + for key in params: + invalid = params.copy() + invalid[key] = "invalid" + yield url, invalid, False + + @staticmethod + def input_validation( + client: FlaskClient, + url: str, + params: dict, + should_pass: bool, + ): + """Check that the application validates input properly for the provided query arguments. + + Parameters + ---------- + client: FlaskClient + The client used to make requests. + url : str + The URL to request from. + params: dict + The arguments for the request. + should_pass: bool + Whether the test should pass or fail. + """ + response = client.get(url, query_string=params) + + if should_pass: + assert response.status_code == 200 + else: + assert response.status_code == 400 + + +class TestWcsValidation: + @pytest.mark.parametrize( + "url, params, should_pass", + chain( + BaseValidation.generate_cases("/ogc?", make_valid_ogc_wcs_describe_coverage_args(layer_static.identifier)), + BaseValidation.generate_cases("/ogc?", make_valid_ogc_wcs_get_capabilities_args()), + BaseValidation.generate_cases( + "/ogc?", make_valid_ogc_wcs_get_coverage_args(layer_time.identifier, str(time[0])) + ), + ), + ) + def test_input_validation( + self, enable_all_formats_in_env, client: FlaskClient, url: str, params: dict, should_pass: bool + ): + BaseValidation.input_validation(client, url, params, should_pass) + + +class TestWmsValidation: + # Ignore the following until implemented + # BaseValidation.generate_cases("/ogc?", make_valid_ogc_wms_get_feature_info_args()) + @pytest.mark.parametrize( + "url, params, should_pass", + chain( + BaseValidation.generate_cases("/ogc?", make_valid_ogc_wms_get_capabilities_args()), + BaseValidation.generate_cases("/ogc?", make_valid_ogc_wms_get_legend_graphic_args(layer_static.identifier)), + BaseValidation.generate_cases( + "/ogc?", make_valid_ogc_wms_get_map_args(layer_time.identifier, str(time[0])) + ), + ), + ) + def test_input_validation( + self, enable_all_formats_in_env, client: FlaskClient, url: str, params: dict, should_pass: bool + ): + BaseValidation.input_validation(client, url, params, should_pass) + + +class TestWmtsValidation: + # Ignore the following until implemented + # BaseValidation.generate_cases("/ogc?", make_valid_ogc_wmts_get_feature_info_args()) + @pytest.mark.parametrize( + "url, params, should_pass", + chain( + BaseValidation.generate_cases("/ogc?", make_valid_ogc_wmts_get_capabilities_args()), + BaseValidation.generate_cases( + "/ogc?", make_valid_ogc_wmts_get_tile_args(layer_time.identifier, str(time[0])) + ), + ), + ) + def test_input_validation( + self, enable_all_formats_in_env, client: FlaskClient, url: str, params: dict, should_pass: bool + ): + BaseValidation.input_validation(client, url, params, should_pass) + + +class TestEdrValidation: + @pytest.mark.parametrize( + "url, params, should_pass", + chain( + BaseValidation.generate_cases("/ogc/edr?", make_valid_ogc_edr_format_args()), + BaseValidation.generate_cases("/ogc/edr/api?", make_valid_ogc_edr_api_args()), + BaseValidation.generate_cases("/ogc/edr/openapi?", make_valid_ogc_edr_api_args()), + BaseValidation.generate_cases("/ogc/edr/conformance?", make_valid_ogc_edr_format_args()), + BaseValidation.generate_cases("/ogc/edr/collections?", make_valid_ogc_edr_format_args()), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_static.group}?", make_valid_ogc_edr_format_args() + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_static.group}/instances?", make_valid_ogc_edr_format_args() + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_static.group}/cube?", + make_valid_ogc_edr_static_cube_args(layer_static.identifier), + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_static.group}/area?", + make_valid_ogc_edr_static_area_args(layer_static.identifier), + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_static.group}/position?", + make_valid_ogc_edr_static_position_args(layer_static.identifier), + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_instance.group}/instances/{instance[0]}/cube?", + make_valid_ogc_edr_instance_cube_args(layer_instance.identifier, str(time[0])), + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_instance.group}/instances/{instance[0]}/area?", + make_valid_ogc_edr_instance_area_args(layer_instance.identifier, str(time[0])), + ), + BaseValidation.generate_cases( + f"/ogc/edr/collections/{layer_instance.group}/instances/{instance[0]}/position?", + make_valid_ogc_edr_instance_position_args(layer_instance.identifier, str(time[0])), + ), + ), + ) + def test_input_validation( + self, enable_all_formats_in_env, client: FlaskClient, url: str, params: dict, should_pass: bool + ): + BaseValidation.input_validation(client, url, params, should_pass) diff --git a/ogc/test/test_servers.py b/ogc/test/test_servers.py index 2d65c33..a1df634 100644 --- a/ogc/test/test_servers.py +++ b/ogc/test/test_servers.py @@ -4,7 +4,7 @@ from ogc import core from ogc import podpac as pogc from ogc import settings -from ogc.ogc_common import WCSException +from ogc.ogc_common import EDRException from pygeoapi.api import APIRequest from unittest.mock import patch @@ -243,19 +243,19 @@ def capturing(request, locales): assert captured_args.get("f") == "json" -def test_edr_render_wcs_exception_returns_400(enable_edr_in_env, client): +def test_edr_render_edr_exception_returns_400(enable_edr_in_env, client): """WCSException raised by a handler is returned as a 400 XML response.""" app = client.application - def raises_wcs_exception(api_request, *args, **kwargs): - raise WCSException("test error") + def raises_edr_exception(api_request, *args, **kwargs): + raise EDRException(status_code=400, exception_code="InvalidQuery", exception_text="") - wrapper = app.edr_render(raises_wcs_exception) + wrapper = app.edr_render(raises_edr_exception) app.add_url_rule("/test_edr_wcs", endpoint="test_edr_wcs", view_func=wrapper, methods=["GET"]) response = client.get("/test_edr_wcs") assert response.status_code == 400 - assert "ExceptionReport" in response.get_data(as_text=True) + assert "InvalidQuery" in response.get_data(as_text=True) def test_edr_render_exception_returns_500(enable_edr_in_env, client): @@ -270,4 +270,4 @@ def raises_runtime_error(api_request, *args, **kwargs): response = client.get("/test_edr_exc") assert response.status_code == 500 - assert "ExceptionReport" in response.get_data(as_text=True) + assert "NoApplicableCode" in response.get_data(as_text=True) diff --git a/ogc/wcs_request_1_0_0.py b/ogc/wcs_request_1_0_0.py index 0f7fcc9..49b4a23 100755 --- a/ogc/wcs_request_1_0_0.py +++ b/ogc/wcs_request_1_0_0.py @@ -206,7 +206,7 @@ def _load_from_kv(self, args): upper_corner=(float(bbox[3]), float(bbox[2])), ) - self.output_format = ogc_common.OutputFormat(value=args["format"]) + self.output_format = ogc_common.OutputFormat(value=args["format"], allowed_values=["geotiff"]) if "time" in args: # TIME : time1, time2,... # or diff --git a/ogc/wmts/wmts_routes.py b/ogc/wmts/wmts_routes.py index 94ea5cc..35efbb0 100644 --- a/ogc/wmts/wmts_routes.py +++ b/ogc/wmts/wmts_routes.py @@ -134,7 +134,7 @@ def get_capabilities(self, args: Dict[str, Any]) -> str: try: get_capabilities.load_from_kv(args) get_capabilities.validate() - except AssertionError: + except Exception: # noqa: B902 logger.exception(LOAD_FAILURE) raise WMTSException(exception_text=INVALID_ARGUMENTS) @@ -175,7 +175,7 @@ def get_tile(self, args: Dict[str, Any]) -> Dict[str, Any]: try: get_tile.load_from_kv(args) get_tile.validate() - except AssertionError: + except Exception: # noqa: B902 logger.exception(LOAD_FAILURE) raise WMTSException(exception_text=INVALID_ARGUMENTS)