diff --git a/cmr/queries.py b/cmr/queries.py index 3e4a1f9..2d91505 100644 --- a/cmr/queries.py +++ b/cmr/queries.py @@ -6,7 +6,7 @@ from collections import defaultdict from datetime import date, datetime, timezone from inspect import getmembers, ismethod -from re import search +from re import fullmatch from typing import Iterable, Iterator from typing_extensions import ( @@ -228,7 +228,7 @@ def format(self, output_format: str = "json") -> Self: # check requested format against the valid format regex's for _format in self._valid_formats_regex: - if search(_format, output_format): + if fullmatch(_format, output_format): self._format = output_format return self @@ -1027,7 +1027,7 @@ def __init__(self, mode: str = CMR_OPS): Query.__init__(self, "collections", mode) self.concept_id_chars = {"C"} self._valid_formats_regex.extend([ - "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]_[0-9]" + "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]+_[0-9]+" ]) def archive_center(self, center: str) -> Self: @@ -1242,7 +1242,7 @@ def __init__(self, mode: str = CMR_OPS): Query.__init__(self, "tools", mode) self.concept_id_chars = {"T"} self._valid_formats_regex.extend([ - "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]_[0-9]" + "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]+_[0-9]+" ]) @override @@ -1259,7 +1259,7 @@ def __init__(self, mode: str = CMR_OPS): Query.__init__(self, "services", mode) self.concept_id_chars = {"S"} self._valid_formats_regex.extend([ - "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]_[0-9]" + "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]+_[0-9]+" ]) @override @@ -1273,7 +1273,7 @@ def __init__(self, mode: str = CMR_OPS): Query.__init__(self, "variables", mode) self.concept_id_chars = {"V"} self._valid_formats_regex.extend([ - "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]_[0-9]" + "dif", "dif10", "opendata", "umm_json", "umm_json_v[0-9]+_[0-9]+" ]) def instance_format(self, format: Union[str, Sequence[str]]) -> Self: diff --git a/tests/test_collection.py b/tests/test_collection.py index 2edeb4f..49ea542 100644 --- a/tests/test_collection.py +++ b/tests/test_collection.py @@ -34,7 +34,7 @@ def test_valid_formats(self): formats = [ "json", "xml", "echo10", "iso", "iso19115", "csv", "atom", "kml", "native", "dif", "dif10", - "opendata", "umm_json", "umm_json_v1_1" "umm_json_v1_9"] + "opendata", "umm_json", "umm_json_v1_1", "umm_json_v1_9"] for _format in formats: query.format(_format) diff --git a/tests/test_format_validation.py b/tests/test_format_validation.py new file mode 100644 index 0000000..20edbb9 --- /dev/null +++ b/tests/test_format_validation.py @@ -0,0 +1,18 @@ +import pytest + +from cmr.queries import CollectionQuery, GranuleQuery, ServiceQuery, ToolQuery, VariableQuery + + +@pytest.mark.parametrize("query_type", [CollectionQuery, GranuleQuery, ServiceQuery, ToolQuery, VariableQuery]) +@pytest.mark.parametrize("output_format", ["jsonn", "not-json", "iso19116", "umm_json_v1_9suffix"]) +def test_rejects_partial_format_matches(query_type, output_format): + query = query_type() + with pytest.raises(ValueError, match="Unsupported format"): + query.format(output_format) + assert query._format == "json" + + +@pytest.mark.parametrize("output_format", ["umm_json", "umm_json_v1_9", "umm_json_v1_18"]) +def test_accepts_complete_umm_versions(output_format): + query = CollectionQuery().format(output_format) + assert query._format == output_format diff --git a/tests/test_service.py b/tests/test_service.py index 47e1d34..e664534 100644 --- a/tests/test_service.py +++ b/tests/test_service.py @@ -38,7 +38,7 @@ def test_valid_formats(self): formats = [ "json", "xml", "echo10", "iso", "iso19115", "csv", "atom", "kml", "native", "dif", "dif10", - "opendata", "umm_json", "umm_json_v1_1" "umm_json_v1_9"] + "opendata", "umm_json", "umm_json_v1_1", "umm_json_v1_9"] for _format in formats: query.format(_format) diff --git a/tests/test_tool.py b/tests/test_tool.py index 5505f18..82551d0 100644 --- a/tests/test_tool.py +++ b/tests/test_tool.py @@ -38,7 +38,7 @@ def test_valid_formats(self): formats = [ "json", "xml", "echo10", "iso", "iso19115", "csv", "atom", "kml", "native", "dif", "dif10", - "opendata", "umm_json", "umm_json_v1_1" "umm_json_v1_9"] + "opendata", "umm_json", "umm_json_v1_1", "umm_json_v1_9"] for _format in formats: query.format(_format) diff --git a/tests/test_variable.py b/tests/test_variable.py index cf9fc15..efbc3e8 100644 --- a/tests/test_variable.py +++ b/tests/test_variable.py @@ -38,7 +38,7 @@ def test_valid_formats(self): formats = [ "json", "xml", "echo10", "iso", "iso19115", "csv", "atom", "kml", "native", "dif", "dif10", - "opendata", "umm_json", "umm_json_v1_1" "umm_json_v1_9"] + "opendata", "umm_json", "umm_json_v1_1", "umm_json_v1_9"] for _format in formats: query.format(_format)