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
4 changes: 4 additions & 0 deletions ogc/edr/edr_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
55 changes: 47 additions & 8 deletions ogc/ogc_common.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import logging
import string

import json
import lxml
import lxml.etree
import numpy as np
Expand All @@ -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):
Expand Down Expand Up @@ -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 "<OutputFormat>%s</OutputFormat>" % 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):
Expand Down Expand Up @@ -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="",
):
"""
Expand Down Expand Up @@ -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="",
):
"""
Expand All @@ -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,
}
)
45 changes: 29 additions & 16 deletions ogc/servers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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, {'"': "&quot;"}) if request.base_url else None
Expand All @@ -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

Expand Down
9 changes: 4 additions & 5 deletions ogc/test/test_input_security.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -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)


Expand All @@ -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")


Expand Down
Loading
Loading