From 289c71922a093ddfb26da3afb55aaac27684f151 Mon Sep 17 00:00:00 2001 From: "Niemes, Adam" Date: Fri, 4 Sep 2026 14:09:33 -0400 Subject: [PATCH] Validate CLI standard parameters before validation --- core.py | 31 +++++++++++++++++++----------- tests/unit/test_core.py | 42 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 62 insertions(+), 11 deletions(-) create mode 100644 tests/unit/test_core.py diff --git a/core.py b/core.py index ceb203d5b..c47a0b4dd 100644 --- a/core.py +++ b/core.py @@ -61,6 +61,25 @@ def validate_encoding(_ctx, _param, value): ) +def validate_standard(standard: str, custom_standard: bool, logger, ctx) -> str: + if custom_standard: + return standard + + normalized_standard = standard.lower() + supported_standards = StandardTypes.values() + + if normalized_standard not in supported_standards: + supported_list = ", ".join(sorted(supported_standards)) + logger.error( + f"Standard '{standard}' is not a supported standard. " + f"Supported standards: {supported_list}. " + f"Use --custom-standard flag for custom standards." + ) + ctx.exit(2) + + return normalized_standard + + def valid_data_file(data_path: list) -> tuple[list, set]: allowed_formats = [ DataFormatTypes.XPT.value, @@ -653,17 +672,7 @@ def validate( # noqa if resolved.is_file(): define_xml_path = str(resolved) - if not custom_standard: - standard = standard.lower() - supported_standards = StandardTypes.values() - if standard not in supported_standards: - supported_list = ", ".join(sorted(supported_standards)) - logger.error( - f"Standard '{standard}' is not a supported standard. " - f"Supported standards: {supported_list}. " - f"Use --custom-standard flag for custom standards." - ) - ctx.exit(2) + standard = validate_standard(standard, custom_standard, logger, ctx) if raw_report: if not (len(output_format) == 1 and output_format[0] == ReportTypes.JSON.value): diff --git a/tests/unit/test_core.py b/tests/unit/test_core.py new file mode 100644 index 000000000..de6453297 --- /dev/null +++ b/tests/unit/test_core.py @@ -0,0 +1,42 @@ +from unittest.mock import MagicMock + +import pytest + +from core import validate_standard + + +def test_validate_standard_accepts_supported_value(): + logger = MagicMock() + ctx = MagicMock() + + result = validate_standard("SDTMIG", False, logger, ctx) + + assert result == "sdtmig" + logger.error.assert_not_called() + ctx.exit.assert_not_called() + + +def test_validate_standard_rejects_typo(): + logger = MagicMock() + ctx = MagicMock() + ctx.exit.side_effect = SystemExit(2) + + with pytest.raises(SystemExit) as exc_info: + validate_standard("stdtmig", False, logger, ctx) + + assert exc_info.value.code == 2 + logger.error.assert_called_once() + message = logger.error.call_args.args[0] + assert "stdtmig" in message + assert "sdtmig" in message + + +def test_validate_standard_allows_custom_standard(): + logger = MagicMock() + ctx = MagicMock() + + result = validate_standard("cust_standard", True, logger, ctx) + + assert result == "cust_standard" + logger.error.assert_not_called() + ctx.exit.assert_not_called()