From e341d19abd94797c6a384e4611fee966641edc0d Mon Sep 17 00:00:00 2001 From: Ayobami Adegoke Date: Mon, 13 Jul 2026 14:55:55 +0100 Subject: [PATCH] fix(salary): validate salary data shape before lookup (#141) --- salary_lookup.py | 48 ++++++++++++++++++++++- tests/test_salary_lookup.py | 76 +++++++++++++++++++++++++++++++++++++ 2 files changed, 122 insertions(+), 2 deletions(-) diff --git a/salary_lookup.py b/salary_lookup.py index 65a7b88..0295e29 100644 --- a/salary_lookup.py +++ b/salary_lookup.py @@ -43,6 +43,46 @@ STRIP_PATTERNS = [ ] +def fail_data_error(message): + """Exit with a user-facing salary data setup error.""" + print(f"Error: invalid salary_data.json: {message}", file=sys.stderr) + print("", file=sys.stderr) + print("See tools/README_SALARY_TOOL.md for the expected format.", file=sys.stderr) + sys.exit(1) + + +def validate_data(data): + """Validate the salary data shape before lookups use it.""" + if not isinstance(data, dict): + fail_data_error("top-level JSON value must be an object") + + metadata = data.get("metadata", {}) + if metadata is not None and not isinstance(metadata, dict): + fail_data_error("'metadata' must be an object when provided") + + companies = data.get("companies") + if not isinstance(companies, list): + fail_data_error("'companies' must be a list") + + for index, entry in enumerate(companies, start=1): + if not isinstance(entry, dict): + fail_data_error(f"companies[{index}] must be an object") + + company = entry.get("company") + if not isinstance(company, str) or not company.strip(): + fail_data_error(f"companies[{index}].company must be a non-empty string") + + city = entry.get("city") + if city is not None and not isinstance(city, str): + fail_data_error(f"companies[{index}].city must be a string when provided") + + categories = entry.get("categories", {}) + if categories is not None and not isinstance(categories, dict): + fail_data_error(f"companies[{index}].categories must be an object when provided") + + return data + + def load_data(): if not DATA_FILE.exists(): print("Error: salary_data.json not found.", file=sys.stderr) @@ -53,8 +93,12 @@ def load_data(): print("If you don't have salary data, the salary lookup", file=sys.stderr) print("step will be skipped during /apply.", file=sys.stderr) sys.exit(1) - with open(DATA_FILE, "r", encoding="utf-8") as f: - return json.load(f) + try: + with open(DATA_FILE, "r", encoding="utf-8") as f: + data = json.load(f) + except json.JSONDecodeError as exc: + fail_data_error(f"invalid JSON at line {exc.lineno}, column {exc.colno}: {exc.msg}") + return validate_data(data) def normalize(s): diff --git a/tests/test_salary_lookup.py b/tests/test_salary_lookup.py index 10f0dcd..dd5e46f 100644 --- a/tests/test_salary_lookup.py +++ b/tests/test_salary_lookup.py @@ -1,7 +1,12 @@ """Tests for salary_lookup.py — format_entry, match_score, and search_company.""" +import io +import tempfile import unittest +from contextlib import redirect_stderr +from pathlib import Path +import salary_lookup from salary_lookup import ( format_entry, normalize, @@ -9,6 +14,7 @@ from salary_lookup import ( extract_core_words, match_score, search_company, + validate_data, ) @@ -165,6 +171,76 @@ class SearchCompanyTests(unittest.TestCase): self.assertEqual(results, []) +class ValidateDataTests(unittest.TestCase): + def assert_invalid_data(self, data, expected_message): + stderr = io.StringIO() + with self.assertRaises(SystemExit) as raised: + with redirect_stderr(stderr): + validate_data(data) + + self.assertEqual(raised.exception.code, 1) + self.assertIn("Error: invalid salary_data.json", stderr.getvalue()) + self.assertIn(expected_message, stderr.getvalue()) + self.assertIn("tools/README_SALARY_TOOL.md", stderr.getvalue()) + + def test_valid_minimal_data_is_returned(self): + data = {"metadata": {}, "companies": [{"company": "Example Corp"}]} + + self.assertIs(validate_data(data), data) + + def test_top_level_value_must_be_object(self): + self.assert_invalid_data([], "top-level JSON value must be an object") + + def test_companies_must_be_list(self): + self.assert_invalid_data({"companies": {"company": "Example Corp"}}, "'companies' must be a list") + + def test_metadata_must_be_object_when_provided(self): + self.assert_invalid_data( + {"metadata": [], "companies": [{"company": "Example Corp"}]}, + "'metadata' must be an object when provided", + ) + + def test_load_data_reports_json_parse_errors_without_traceback(self): + with tempfile.TemporaryDirectory() as tmpdir: + data_file = Path(tmpdir) / "salary_data.json" + data_file.write_text('{"companies": [', encoding="utf-8") + + original_data_file = salary_lookup.DATA_FILE + salary_lookup.DATA_FILE = data_file + try: + stderr = io.StringIO() + with self.assertRaises(SystemExit) as raised: + with redirect_stderr(stderr): + salary_lookup.load_data() + finally: + salary_lookup.DATA_FILE = original_data_file + + self.assertEqual(raised.exception.code, 1) + self.assertIn("invalid JSON at line", stderr.getvalue()) + self.assertIn("tools/README_SALARY_TOOL.md", stderr.getvalue()) + + def test_company_entry_must_be_object(self): + self.assert_invalid_data({"companies": ["Example Corp"]}, "companies[1] must be an object") + + def test_company_name_is_required(self): + self.assert_invalid_data({"companies": [{"city": "Aarhus"}]}, "companies[1].company must be a non-empty string") + + def test_company_name_must_not_be_blank(self): + self.assert_invalid_data({"companies": [{"company": " "}]}, "companies[1].company must be a non-empty string") + + def test_city_must_be_string_when_provided(self): + self.assert_invalid_data( + {"companies": [{"company": "Example Corp", "city": 123}]}, + "companies[1].city must be a string when provided", + ) + + def test_categories_must_be_object_when_provided(self): + self.assert_invalid_data( + {"companies": [{"company": "Example Corp", "categories": []}]}, + "companies[1].categories must be an object when provided", + ) + + class UtilityTests(unittest.TestCase): def test_normalize_strips_suffix_and_noise(self): self.assertEqual(normalize("Novo Nordisk A/S"), "novonordisk")