diff --git a/dashboard_server.py b/dashboard_server.py index d0a8056..cb67155 100644 --- a/dashboard_server.py +++ b/dashboard_server.py @@ -44,9 +44,34 @@ in {"1", "true", "yes", "on"} ) DASHBOARD_CORS_ORIGIN = os.getenv("CHELATED_DASHBOARD_CORS_ORIGIN", "").strip() -# Upper bound for the /api/events ?limit= parameter (LIMIT-01): prevents -# unbounded response sizes from huge or garbage input. +# Parsed integer limits are capped here. An omitted /api/events limit returns +# the filtered file. A non-integer /api/events limit is HTTP 400, not an uncapped list. _MAX_API_LIMIT = 5000 + + +def _nonnegative_limit(raw: str) -> int: + """Parse a limit. Zero and negative select nothing. Above the cap is clamped.""" + value = int(raw) + if value <= 0: + return 0 + return min(value, _MAX_API_LIMIT) + + +def _limit_or_default(query_params: Dict[str, List[str]], default: int) -> int: + """Bounded ``limit`` query. A non-integer keeps ``default`` instead of failing the request.""" + values = query_params.get("limit") if query_params else None + if not values: + return default + try: + return _nonnegative_limit(values[0]) + except (TypeError, ValueError): + return default + + +def _tail_paths(paths, limit: int): + if limit <= 0: + return [] + return list(paths)[-limit:] CAMPAIGN_HISTORY_ROOT = "experiment_runs" VALIDATION_HISTORY_ROOT = "experiment_runs" PREFLIGHT_HISTORY_ROOT = "experiment_runs" @@ -1213,10 +1238,7 @@ def load_campaign_history(root: str = CAMPAIGN_HISTORY_ROOT, limit: int = 25) -> ) reports = [] for path in report_paths[: max(0, limit)]: - try: - reports.append(_extract_campaign_record(path, root_path)) - except (OSError, ValueError): - continue + reports.append(_extract_campaign_record(path, root_path)) return { "root": root, @@ -1261,10 +1283,7 @@ def load_validation_history(root: str = VALIDATION_HISTORY_ROOT, limit: int = 10 report_paths = sorted(root_path.rglob("validation_summary.json"), key=lambda item: item.stat().st_mtime, reverse=True) reports = [] for path in report_paths[: max(0, limit)]: - try: - reports.append(_extract_validation_record(path, root_path)) - except (OSError, ValueError): - continue + reports.append(_extract_validation_record(path, root_path)) latest = reports[0] if reports else {} return { "root": root, @@ -1312,18 +1331,12 @@ def load_preflight_history(root: str = PREFLIGHT_HISTORY_ROOT, limit: int = 10) report_paths_with_mtime = [] for path in root_path.rglob("*preflight*.json"): - try: - report_paths_with_mtime.append((path, path.stat().st_mtime)) - except OSError: - continue + report_paths_with_mtime.append((path, path.stat().st_mtime)) report_paths_with_mtime.sort(key=lambda item: item[1], reverse=True) report_paths = [path for path, _mtime in report_paths_with_mtime] reports = [] for path in report_paths[: max(0, limit)]: - try: - reports.append(_extract_preflight_record(path, root_path)) - except (OSError, ValueError): - continue + reports.append(_extract_preflight_record(path, root_path)) latest = reports[0] if reports else {} return { "root": root, @@ -1353,10 +1366,7 @@ def load_evidence_index(path: str = EVIDENCE_INDEX_PATH) -> Dict[str, Any]: }, "artifacts": {}, } - try: - payload = _load_json_object(index_path) - except (OSError, ValueError): - payload = {} + payload = _load_json_object(index_path) summary = payload.get("summary") if not isinstance(summary, dict): summary = {} @@ -1721,10 +1731,7 @@ def do_HEAD(self): super().do_HEAD() def do_OPTIONS(self): - """CORS preflight scoped to the configured origin (P2-01).""" - if not self._is_api_authorized(): - self.send_error_response(405, "Method not allowed") - return + """CORS preflight. Browsers do not send Authorization on this request.""" self.send_response(204) self.send_header("Content-Length", "0") if DASHBOARD_CORS_ORIGIN: @@ -1793,18 +1800,17 @@ def handle_api_events(self, query_params: Dict[str, List[str]]): - limit: Maximum number of events to return - event_type: Filter by event type (e.g., "query", "error") """ + limit = None + if "limit" in query_params: + try: + limit = int(query_params["limit"][0]) + except (ValueError, IndexError): + self.send_error_response(400, "limit must be an integer") + return + limit = min(limit, _MAX_API_LIMIT) try: events = load_events(LOG_FILE_PATH) - # Extract query parameters - limit = None - if "limit" in query_params: - try: - limit = int(query_params["limit"][0]) - limit = min(limit, _MAX_API_LIMIT) - except (ValueError, IndexError): - limit = None - event_type = None if "event_type" in query_params: try: @@ -1975,11 +1981,13 @@ def handle_api_evidence_cleanup_plan(self, query_params: Dict[str, List[str]]): """Handle /api/evidence_cleanup_plan endpoint.""" try: keep_latest = 1 - limit = 25 - if "keep_latest" in query_params: - keep_latest = max(0, int(query_params["keep_latest"][0])) - if "limit" in query_params: - limit = max(0, int(query_params["limit"][0])) + raw_keep = query_params.get("keep_latest") if query_params else None + if raw_keep: + try: + keep_latest = max(0, int(raw_keep[0])) + except (TypeError, ValueError): + keep_latest = 1 + limit = _limit_or_default(query_params, 25) self.send_json_response(load_evidence_cleanup_plan(EVIDENCE_CLEANUP_ROOT, keep_latest=keep_latest, candidate_limit=limit)) except Exception: self.send_error_response(500, "Error reading evidence cleanup plan") @@ -2002,9 +2010,9 @@ def handle_api_model_scope_events(self, query_params): """Handle /api/model_scope/events — lists recent activation event files.""" from model_scope_artifacts import ArtifactStore, load_model_scope_artifact, summarize_model_scope_artifact try: - limit = int(query_params.get("limit", ["20"])[0]) + limit = _limit_or_default(query_params, 20) store = ArtifactStore(base_dir=MODEL_SCOPE_ARTIFACT_ROOT) - paths = store.list_artifacts(pattern="feature_event_*.json")[-limit:] + paths = _tail_paths(store.list_artifacts(pattern="feature_event_*.json"), limit) items = [] for p in reversed(paths): try: @@ -2018,16 +2026,16 @@ def handle_api_model_scope_events(self, query_params): "reason": None if items else "no_artifacts_found", "items": items, }) - except Exception as e: - self.send_error_response(500, f"Error reading model-scope events: {e}") + except Exception: + self.send_error_response(500, "Error reading model-scope events") def handle_api_model_scope_features(self, query_params): """Handle /api/model_scope/features — lists recent sparse feature events.""" from model_scope_artifacts import ArtifactStore, load_model_scope_artifact try: - limit = int(query_params.get("limit", ["20"])[0]) + limit = _limit_or_default(query_params, 20) store = ArtifactStore(base_dir=MODEL_SCOPE_ARTIFACT_ROOT) - paths = store.list_artifacts(pattern="feature_event_*.json")[-limit:] + paths = _tail_paths(store.list_artifacts(pattern="feature_event_*.json"), limit) items = [] for p in reversed(paths): try: @@ -2041,16 +2049,16 @@ def handle_api_model_scope_features(self, query_params): "reason": None if items else "no_feature_events_found", "items": items, }) - except Exception as e: - self.send_error_response(500, f"Error reading model-scope features: {e}") + except Exception: + self.send_error_response(500, "Error reading model-scope features") def handle_api_model_scope_interventions(self, query_params): """Handle /api/model_scope/interventions — lists recent intervention records.""" from model_scope_artifacts import ArtifactStore, load_model_scope_artifact try: - limit = int(query_params.get("limit", ["20"])[0]) + limit = _limit_or_default(query_params, 20) store = ArtifactStore(base_dir=MODEL_SCOPE_ARTIFACT_ROOT) - paths = store.list_artifacts(pattern="intervention_*.json")[-limit:] + paths = _tail_paths(store.list_artifacts(pattern="intervention_*.json"), limit) items = [] for p in reversed(paths): try: @@ -2063,8 +2071,8 @@ def handle_api_model_scope_interventions(self, query_params): "reason": None if items else "no_interventions_found", "items": items, }) - except Exception as e: - self.send_error_response(500, f"Error reading model-scope interventions: {e}") + except Exception: + self.send_error_response(500, "Error reading model-scope interventions") def handle_api_tts_status(self): """Handle /api/tts/status — returns TTS enable state, config, and last result summary.""" diff --git a/test_dashboard_api_limits.py b/test_dashboard_api_limits.py new file mode 100644 index 0000000..24239ae --- /dev/null +++ b/test_dashboard_api_limits.py @@ -0,0 +1,147 @@ +"""Limits, preflight, and corrupt-history contracts.""" + +import json +import os +import sys +import tempfile +import types +import unittest +from io import BytesIO +from unittest.mock import MagicMock, patch + +import dashboard_server + + +def _handler(): + handler = dashboard_server.DashboardHandler.__new__(dashboard_server.DashboardHandler) + handler.wfile = BytesIO() + handler.headers = {} + handler.send_response = MagicMock() + handler.send_header = MagicMock() + handler.end_headers = MagicMock() + handler.send_error_response = MagicMock() + return handler + + +class TestApiLimits(unittest.TestCase): + def test_events_non_integer_limit_is_400(self): + handler = _handler() + handler.handle_api_events({"limit": ["abc"]}) + handler.send_error_response.assert_called_once_with(400, "limit must be an integer") + + def test_model_scope_non_integer_limit_keeps_default_json(self): + handler = _handler() + params = {"limit": ["abc"]} + methods = ( + handler.handle_api_model_scope_events, + handler.handle_api_model_scope_features, + handler.handle_api_model_scope_interventions, + ) + # The handler imports ArtifactStore locally. Stub that module so this + # limit path does not load torch. + artifacts = types.ModuleType("model_scope_artifacts") + + class ArtifactStore: + def __init__(self, base_dir=None): + self.base_dir = base_dir + + def list_artifacts(self, pattern="feature_event_*.json"): + return [] + + artifacts.ArtifactStore = ArtifactStore + artifacts.load_model_scope_artifact = lambda path: {} + artifacts.summarize_model_scope_artifact = lambda artifact: {} + with patch.dict(sys.modules, {"model_scope_artifacts": artifacts}): + with patch.object(ArtifactStore, "list_artifacts", return_value=[]) as listed: + for method in methods: + method(params) + self.assertEqual(listed.call_count, 3) + handler.send_error_response.assert_not_called() + self.assertEqual(handler.send_response.call_count, 3) + handler.send_response.assert_called_with(200) + body = handler.wfile.getvalue().decode("utf-8") + self.assertNotIn("invalid literal", body) + self.assertNotIn("ValueError", body) + self.assertNotIn("Error reading", body) + self.assertEqual(body.count('"status": "not_generated"'), 3) + + def test_cleanup_non_integer_limit_keeps_default_json(self): + handler = _handler() + with patch( + "dashboard_server.load_evidence_cleanup_plan", + return_value={"candidates": [], "dry_run": True}, + ) as load_plan: + handler.handle_api_evidence_cleanup_plan({"limit": ["abc"]}) + handler.send_error_response.assert_not_called() + load_plan.assert_called_once_with( + dashboard_server.EVIDENCE_CLEANUP_ROOT, + keep_latest=1, + candidate_limit=25, + ) + handler.send_response.assert_called_with(200) + body = handler.wfile.getvalue().decode("utf-8") + payload = json.loads(body) + self.assertEqual(payload["candidates"], []) + self.assertNotIn("invalid literal", body) + self.assertNotIn("Error reading", body) + + def test_cleanup_non_integer_keep_latest_keeps_default(self): + handler = _handler() + with patch( + "dashboard_server.load_evidence_cleanup_plan", + return_value={"candidates": [], "dry_run": True}, + ) as load_plan: + handler.handle_api_evidence_cleanup_plan({"keep_latest": ["abc"]}) + handler.send_error_response.assert_not_called() + load_plan.assert_called_once_with( + dashboard_server.EVIDENCE_CLEANUP_ROOT, + keep_latest=1, + candidate_limit=25, + ) + + def test_options_preflight_does_not_require_a_bearer(self): + old_token = dashboard_server.DASHBOARD_TOKEN + old_origin = dashboard_server.DASHBOARD_CORS_ORIGIN + dashboard_server.DASHBOARD_TOKEN = "secret" + dashboard_server.DASHBOARD_CORS_ORIGIN = "https://example.test" + try: + handler = _handler() + handler.headers = {} + handler.do_OPTIONS() + handler.send_response.assert_called_with(204) + handler.send_error_response.assert_not_called() + finally: + dashboard_server.DASHBOARD_TOKEN = old_token + dashboard_server.DASHBOARD_CORS_ORIGIN = old_origin + + def test_zero_model_scope_limit_is_empty(self): + self.assertEqual(dashboard_server._nonnegative_limit("0"), 0) + self.assertEqual(dashboard_server._nonnegative_limit("-3"), 0) + self.assertEqual(dashboard_server._nonnegative_limit("9000"), dashboard_server._MAX_API_LIMIT) + self.assertEqual(dashboard_server._tail_paths([1, 2, 3], 0), []) + self.assertEqual(dashboard_server._limit_or_default({"limit": ["abc"]}, 20), 20) + self.assertEqual(dashboard_server._limit_or_default({"limit": ["abc"]}, 25), 25) + self.assertEqual(dashboard_server._limit_or_default({"limit": ["0"]}, 20), 0) + self.assertEqual(dashboard_server._limit_or_default({"limit": ["-3"]}, 25), 0) + self.assertEqual( + dashboard_server._limit_or_default({"limit": ["9000"]}, 25), + dashboard_server._MAX_API_LIMIT, + ) + + def test_bad_campaign_file_raises(self): + with tempfile.TemporaryDirectory() as tmpdir: + path = os.path.join(tmpdir, "campaign_report.json") + with open(path, "w", encoding="utf-8") as handle: + handle.write("{") + with self.assertRaises(ValueError): + dashboard_server.load_campaign_history(tmpdir, limit=10) + + def test_missing_campaign_root_is_empty(self): + with tempfile.TemporaryDirectory() as tmpdir: + missing = os.path.join(tmpdir, "absent") + result = dashboard_server.load_campaign_history(missing, limit=10) + self.assertEqual(result["summary"]["total_reports"], 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test_dashboard_server.py b/test_dashboard_server.py index e1aa75e..6fb18c3 100644 --- a/test_dashboard_server.py +++ b/test_dashboard_server.py @@ -454,16 +454,14 @@ def test_load_evidence_index_normalizes_summary(self): self.assertFalse(result["summary"]["latest_review_allowed"]) self.assertTrue(result["summary"]["latest_chain_passed"]) - def test_load_evidence_index_returns_empty_for_malformed_json(self): + def test_load_evidence_index_rejects_malformed_json(self): with tempfile.TemporaryDirectory() as tmpdir: path = os.path.join(tmpdir, "evidence_index.json") with open(path, "w", encoding="utf-8") as handle: handle.write("[not-an-object]") - result = dashboard_server.load_evidence_index(path) - - self.assertTrue(result["present"]) - self.assertEqual(result["summary"]["artifact_counts"], {}) + with self.assertRaises(ValueError): + dashboard_server.load_evidence_index(path) class TestEvidenceChainHistory(unittest.TestCase):