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
32 changes: 30 additions & 2 deletions ogc/servers.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
"""

import re
import xml.sax.saxutils
from flask import Flask, request, Response, make_response, send_file
import six
import logging
Expand All @@ -21,6 +22,16 @@
logger = logging.getLogger(__name__)


def _check_query_string(raw_qs: bytes) -> None:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I feel like, if we ever want to extend our input validation, this will serve nicely as the place to do that.

"""Raise WCSException 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.")
try:
raw_qs.decode("utf-8")
except UnicodeDecodeError:
raise WCSException("Request contains invalid UTF-8 encoding.")


def respond_xml(doc, status=200):
# First, validate that XML can be parsed.
from lxml import etree
Expand Down Expand Up @@ -231,6 +242,12 @@ def ogc_render(self, ogc_idx):
return respond_xml("<p>Only GET supported</p>", status=405)

ogc = self.ogcs[ogc_idx]

try:
_check_query_string(request.query_string)
except WCSException as e:
return respond_xml(e.to_xml(), status=400)

if not request.args:
return self.home_func(ogc.endpoint)
try:
Expand All @@ -253,7 +270,7 @@ def ogc_render(self, ogc_idx):
}

if request.base_url:
args["base_url"] = request.base_url + "?"
args["base_url"] = xml.sax.saxutils.escape(request.base_url, {'"': "&quot;"}) + "?"
else:
args["base_url"] = None
ogc_response = None
Expand Down Expand Up @@ -310,6 +327,12 @@ def wrapper(*args, **kwargs) -> Response:
logger.info("OGC server.edr_render")
if request.method != "GET":
return respond_xml("<p>Only GET supported</p>", status=405)

try:
_check_query_string(request.query_string)
except WCSException as e:
return respond_xml(e.to_xml(), status=400)

try:
# We'll filter out any characters from URl parameter values that
# are not in the allowlist.
Expand All @@ -336,7 +359,9 @@ def wrapper(*args, **kwargs) -> Response:
if format_argument is not None:
filtered_args["f"] = format_argument.lower()

filtered_args["base_url"] = request.base_url
filtered_args["base_url"] = (
xml.sax.saxutils.escape(request.base_url, {'"': "&quot;"}) if request.base_url else None
)

# Replace the arguments with the filtered option
request.args = ImmutableMultiDict(filtered_args)
Expand All @@ -347,6 +372,9 @@ def wrapper(*args, **kwargs) -> Response:
if headers:
response.headers = headers
return response
except WCSException as e:
logger.error("OGC: server.edr_render WCSException: %s", str(e), exc_info=True)
return respond_xml(e.to_xml(), status=400)
except Exception as e:
logger.error("OGC: server.edr_render Exception: %s", str(e), exc_info=True)
ee = WCSException()
Expand Down
3 changes: 3 additions & 0 deletions ogc/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,9 @@
# Max WCS/WMS response size
MAX_GRID_COORDS_REQUEST_SIZE = 1024 * 1024

# HTTP input size limits
MAX_QUERY_STRING_BYTES = 8192

# WMS Capabilities limit layers
WMS_LIMIT_LAYERS = False

Expand Down
96 changes: 96 additions & 0 deletions ogc/test/test_input_security.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
"""
Tests verifying that the server rejects oversized query strings.
"""

import pytest
import numpy as np
import podpac

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

# ---------------------------------------------------------------------------
# Shared fixture (mirrors test_servers.py)
# ---------------------------------------------------------------------------


@pytest.fixture
def client():
lat = np.linspace(90, -90, 11)
lon = np.linspace(-180, 180, 21)
data = np.random.default_rng(1).random((11, 21))
coords = podpac.Coordinates([lat, lon], dims=["lat", "lon"])
node = podpac.data.Array(source=data, coordinates=coords)
layer = pogc.Layer(
node=node,
identifier="layer1",
title="Layer 1",
abstract="Layer 1 Data",
group="Layers",
)
ogc_instance = core.OGC(layers=[layer])
app = servers.FlaskServer(__name__, ogcs=[ogc_instance])
app.config["TESTING"] = True
yield app.test_client()


# ---------------------------------------------------------------------------
# 1. _check_query_string unit tests — exercises the real production function
# ---------------------------------------------------------------------------


def test_check_query_string_overflow():
"""Raises WCSException 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"):
_check_query_string(oversized)


def test_check_query_string_valid():
"""Does not raise for a normal ASCII query string."""
_check_query_string(b"SERVICE=WCS&REQUEST=GetCapabilities&VERSION=1.0.0")


def test_check_query_string_exactly_at_limit():
"""Does not raise when the query string is exactly at the limit."""
at_limit = b"A" * settings.MAX_QUERY_STRING_BYTES
_check_query_string(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"):
_check_query_string(b"SERVICE=WCS&COVERAGE=\xff\xfe")


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")


# ---------------------------------------------------------------------------
# 2. HTTP integration — overflow rejected via test client
# ---------------------------------------------------------------------------


def test_ogc_render_overflow_returns_400(client):
"""A query string over the limit returns a 400 WCSException response."""
oversized = "A=" + "B" * (settings.MAX_QUERY_STRING_BYTES + 1)
response = client.get("/ogc", environ_overrides={"QUERY_STRING": oversized})
assert response.status_code == 400
assert b"ExceptionReport" in response.data
assert b"maximum allowed length" in response.data


# ---------------------------------------------------------------------------
# 3. Regression — valid WCS request still works end-to-end
# ---------------------------------------------------------------------------


def test_valid_wcs_get_capabilities_unchanged(client):
"""Normal GetCapabilities request is unaffected by the new guards."""
response = client.get("/ogc?SERVICE=WCS&REQUEST=GetCapabilities&VERSION=1.0.0")
assert response.status_code == 200
assert b"WCS_Capabilities" in response.data
20 changes: 11 additions & 9 deletions ogc/wcs_response_1_0_0.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import logging
import datetime
from xml.sax.saxutils import escape

import traitlets as tl
from ogc import GridCoordinates
Expand Down Expand Up @@ -89,11 +90,11 @@ class CoverageDescription(ogc_common.XMLNode):
def coverage_offering(self, coverage):
xml = """ <wcs:CoverageOffering>"""
if coverage.identifier:
xml += " <wcs:name>{coverage.identifier}</wcs:name>\n".format(coverage=coverage)
xml += " <wcs:name>{}</wcs:name>\n".format(escape(coverage.identifier))
if coverage.title:
xml += " <wcs:label>{coverage.title}</wcs:label>\n".format(coverage=coverage)
xml += " <wcs:label>{}</wcs:label>\n".format(escape(coverage.title))
if coverage.abstract:
xml += " <wcs:description>{coverage.abstract}</wcs:description>\n".format(coverage=coverage)
xml += " <wcs:description>{}</wcs:description>\n".format(escape(coverage.abstract))
temporal_domain = ""
if hasattr(coverage.layer, "valid_times"):
if coverage.layer.all_times_valid:
Expand Down Expand Up @@ -226,14 +227,15 @@ class Capabilities(ogc_common.XMLNode):
version = tl.Unicode(default_value=SERVICE_VERSION)

def service(self):
title = escape(self.service_title) if self.service_title else ""
return """\
<wcs:Service>
<wcs:name>{self.service_title}</wcs:name>
<wcs:label>{self.service_title}</wcs:label>
<wcs:name>{title}</wcs:name>
<wcs:label>{title}</wcs:label>
<wcs:fees>UNAVAILABLE</wcs:fees>
<wcs:accessConstraints>{constraints}</wcs:accessConstraints>
</wcs:Service>
""".format(self=self, constraints=settings.CONSTRAINTS)
""".format(title=title, constraints=settings.CONSTRAINTS)

base_url = tl.Unicode(default_value=None, allow_none=True) # e.g., http://hostname:port/path?

Expand Down Expand Up @@ -296,13 +298,13 @@ def contents(self):
for coverage in self.coverages:
xml += " <wcs:CoverageOfferingBrief>\n"
if coverage.abstract:
xml += " <wcs:description>{coverage.abstract}</wcs:description>\n".format(coverage=coverage)
xml += " <wcs:description>{}</wcs:description>\n".format(escape(coverage.abstract))
if coverage.identifier: # required
xml += " <wcs:name>{coverage.identifier}</wcs:name>\n".format(coverage=coverage)
xml += " <wcs:name>{}</wcs:name>\n".format(escape(coverage.identifier))
else:
logger.info("Invalid layer. Missing name.")
if coverage.title: # required
xml += " <wcs:label>{coverage.title}</wcs:label>\n".format(coverage=coverage)
xml += " <wcs:label>{}</wcs:label>\n".format(escape(coverage.title))
else:
logger.info("Invalid layer. Missing label.")
if coverage.wgs84_bounding_box_lower_corner_lat_lon or coverage.wgs84_bounding_box_upper_corner_lat_lon:
Expand Down
15 changes: 10 additions & 5 deletions ogc/wms_response_1_3_0.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import logging
from xml.sax.saxutils import escape

import traitlets as tl
import numpy as np
Expand Down Expand Up @@ -43,17 +44,19 @@ class Capabilities(ogc_common.XMLNode):
version = tl.Unicode(default_value=SERVICE_VERSION)

def service(self):
title = escape(self.service_title) if self.service_title else ""
return """\
<Service>
<Name>WMS</Name>
<Title>{self.service_title}</Title>
<Title>{title}</Title>
<OnlineResource xlink:href="{self.base_url}"/>
<AccessConstraints>{constraints}</AccessConstraints>
<LayerLimit>1</LayerLimit>
<MaxWidth>{maxWidthWMS}</MaxWidth>
<MaxHeight>{maxHeightWMS}</MaxHeight>
</Service>
""".format(
title=title,
self=self,
constraints=settings.CONSTRAINTS,
maxWidthWMS=int(np.sqrt(settings.MAX_GRID_COORDS_REQUEST_SIZE)),
Expand Down Expand Up @@ -121,13 +124,13 @@ def coverage_times_list(self, coverage, default_time):
def coverage_layer(self, coverage):
xml = """ <Layer queryable="0" opaque="0" cascaded="1">\n"""
if coverage.identifier:
xml += f" <Name>{coverage.identifier}</Name>\n"
xml += f" <Name>{escape(coverage.identifier)}</Name>\n"
if coverage.title:
xml += f" <Title>{coverage.title}</Title>\n"
xml += f" <Title>{escape(coverage.title)}</Title>\n"
else:
logger.info("Invalid layer. Missing title.")
if coverage.abstract:
xml += f" <Abstract>{coverage.abstract}</Abstract>\n"
xml += f" <Abstract>{escape(coverage.abstract)}</Abstract>\n"

xml += self._get_CRS_and_BoundingBox()

Expand Down Expand Up @@ -199,7 +202,9 @@ def coverage_layer(self, coverage):

def layers(self):
xml = " <Layer>\n"
xml += " <Title>{}</Title>\n".format(self.service_group_title)
xml += " <Title>{}</Title>\n".format(
escape(self.service_group_title) if self.service_group_title else ""
)
xml += self._get_CRS_and_BoundingBox(depth=2)

# If configured, trim layers list to layers specified in settings
Expand Down
Loading