From be13849a4f624794891bda66fea526eb92db360e Mon Sep 17 00:00:00 2001 From: Mia Wu Date: Thu, 25 Jun 2026 11:31:20 -0400 Subject: [PATCH] implemented watcher and rerun logic --- backend/app.py | 71 +++++++- backend/db.py | 159 ++++++++++++++++- backend/graph.py | 65 ++++++- backend/parser.py | 201 ++++++++++++++++++++-- backend/requirements.txt | 1 + backend/scanner.py | 49 ++++-- backend/scheduler.py | 2 - backend/watcher.py | 161 +++++++++++++++++ frontend/src/App.jsx | 108 ++++++++++-- frontend/src/api.js | 1 + frontend/src/components/Calendar.jsx | 3 +- frontend/src/components/FailedBanner.jsx | 18 +- frontend/src/components/RightPanel.jsx | 14 ++ frontend/src/components/SettingsModal.jsx | 112 +++++++++--- frontend/src/components/TestCard.jsx | 3 +- 15 files changed, 871 insertions(+), 97 deletions(-) diff --git a/backend/app.py b/backend/app.py index 31eafed..5e36462 100644 --- a/backend/app.py +++ b/backend/app.py @@ -12,6 +12,7 @@ import db import graph from parser import CsvValidationError, parse_target_csv from scheduler import SchedulerTest, compile_schedule, remove_from_active, reset_scheduler_state +from watcher import configure_result_watcher, stop_result_watcher APP_ROOT = Path(__file__).resolve().parent @@ -28,6 +29,41 @@ class SaveSettingsRequest(BaseModel): settings: dict[str, Any] +def _smb_credentials_from_settings(settings: dict[str, Any]) -> dict[str, str]: + return { + "username": str(settings.get("smbUsername") or "").strip(), + "password": str(settings.get("smbPassword") or ""), + "domain": str(settings.get("smbDomain") or "").strip(), + } + + +def _runtime_overrides_from_settings(settings: dict[str, Any]) -> dict[str, int]: + def _parse_positive_int(value: Any) -> int | None: + if value is None: + return None + if isinstance(value, str): + value = value.strip() + if not value: + return None + try: + parsed = int(value) + except (TypeError, ValueError): + return None + return parsed if parsed > 0 else None + + overrides: dict[str, int] = {} + for test_type, key in ( + ("P2P", "p2pRuntimeMinutes"), + ("COE", "coeRuntimeMinutes"), + ("P3P", "p3pRuntimeMinutes"), + ): + minutes = _parse_positive_int(settings.get(key)) + if minutes is not None: + overrides[test_type] = minutes + + return overrides + + class CompileScheduleRequest(BaseModel): start_date: str | None = Field(default=None, description="YYYY-MM-DD") rule: str = "" @@ -48,7 +84,12 @@ class SaveHolidaysRequest(BaseModel): async def lifespan(application: FastAPI): db.init_db(DB_PATH) graph.reset_graph_state() - yield + settings = db.read_settings(DB_PATH) + configure_result_watcher(settings) + try: + yield + finally: + stop_result_watcher() app = FastAPI(title="Scheduler API", version="0.1.0", lifespan=lifespan) @@ -69,6 +110,7 @@ def health() -> dict[str, str]: @app.post("/api/settings") def save_settings(request: SaveSettingsRequest) -> dict[str, str]: db.save_settings(request.settings, DB_PATH) + configure_result_watcher(db.read_settings(DB_PATH)) return {"status": "saved"} @@ -83,8 +125,16 @@ def load_tests(request: LoadTestsRequest) -> dict[str, Any]: if not csv_path.is_absolute(): csv_path = APP_ROOT / csv_path + settings = db.read_settings(DB_PATH) + smb_credentials = _smb_credentials_from_settings(settings) + runtime_overrides = _runtime_overrides_from_settings(settings) + try: - parsed = parse_target_csv(csv_path) + parsed = parse_target_csv( + csv_path, + smb_credentials=smb_credentials, + runtime_overrides=runtime_overrides, + ) except FileNotFoundError as exc: raise HTTPException(status_code=404, detail=str(exc)) from exc except CsvValidationError as exc: @@ -94,7 +144,7 @@ def load_tests(request: LoadTestsRequest) -> dict[str, Any]: reset_scheduler_state() graph.reset_graph_state() all_dut_tests = db.list_tests_for_device(DUT, DB_PATH) - graph.build_graph_once(all_dut_tests) + graph.build_and_persist_graph(all_dut_tests, DB_PATH) return { "loaded_tests": count, @@ -117,6 +167,7 @@ def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any] except ValueError as exc: raise HTTPException(status_code=400, detail="start_date must be YYYY-MM-DD") from exc + reset_scheduler_state() stored_tests = db.list_schedulable_tests(DB_PATH, rule=request.rule) if not stored_tests: version = db.create_schedule_version([], DB_PATH) @@ -144,6 +195,10 @@ def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any] for t in stored_tests ] + # DB-backed graph retrieval ensures compile works after restart without manual save/load. + all_dut_tests = db.list_tests_for_device(DUT, DB_PATH) + graph.get_graph(DB_PATH, all_dut_tests) + holiday_dates = db.list_holidays(DB_PATH) try: entries, completion_date = compile_schedule( @@ -161,7 +216,7 @@ def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any] [(e.test_id, e.device, e.scheduled_date, e.shift_index, e.sequence_in_shift) for e in entries], DB_PATH, ) - + print(f"Schedule version {version} created with {len(entries)} entries, completion date: {completion_date}") return { "schedule_version": version, "scheduled_tests": len(entries), @@ -169,6 +224,14 @@ def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any] } +@app.get("/api/tests/rerun") +def get_rerun_tests() -> dict[str, Any]: + db.mark_overdue_as_rerun(DB_PATH) + tests = db.get_rerun_tests(DB_PATH) + total_minutes = sum(t["estimated_minutes"] for t in tests) + return {"tests": tests, "total_estimated_minutes": total_minutes} + + @app.post("/api/holidays") def save_holidays(request: SaveHolidaysRequest) -> dict[str, Any]: dates = [d.strip() for d in request.dates if d.strip()] diff --git a/backend/db.py b/backend/db.py index 302fcb4..9b86daf 100644 --- a/backend/db.py +++ b/backend/db.py @@ -87,7 +87,7 @@ def init_db(db_path: str | Path = DB_PATH) -> None: config_json TEXT, throttled INTEGER NOT NULL DEFAULT 0, estimated_minutes INTEGER NOT NULL, - status TEXT NOT NULL CHECK (status IN ('pending', 'completed', 'failed', 'invalid')), + status TEXT NOT NULL CHECK (status IN ('pending', 'completed', 'rerun')), excluded INTEGER NOT NULL DEFAULT 0, raw_payload TEXT, created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP, @@ -123,6 +123,13 @@ def init_db(db_path: str | Path = DB_PATH) -> None: minutes INTEGER NOT NULL ); + CREATE TABLE IF NOT EXISTS graph_cache ( + name TEXT PRIMARY KEY, + payload_json TEXT NOT NULL, + test_count INTEGER NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP + ); + CREATE TABLE IF NOT EXISTS rerun_events ( id INTEGER PRIMARY KEY AUTOINCREMENT, detected_date TEXT NOT NULL, @@ -246,6 +253,52 @@ def read_settings(db_path: str | Path = DB_PATH) -> dict[str, Any]: return {row["key"]: json.loads(row["value_json"]) for row in rows} +def save_graph_cache( + name: str, + payload: dict[str, list[str]], + test_count: int, + db_path: str | Path = DB_PATH, +) -> None: + with get_connection(db_path) as conn: + conn.execute( + """ + INSERT INTO graph_cache(name, payload_json, test_count, updated_at) + VALUES (?, ?, ?, CURRENT_TIMESTAMP) + ON CONFLICT(name) DO UPDATE SET + payload_json = excluded.payload_json, + test_count = excluded.test_count, + updated_at = CURRENT_TIMESTAMP + """, + (name, json.dumps(payload), int(test_count)), + ) + + +def load_graph_cache(name: str, db_path: str | Path = DB_PATH) -> dict[str, Any] | None: + with get_connection(db_path) as conn: + row = conn.execute( + """ + SELECT payload_json, test_count, updated_at + FROM graph_cache + WHERE name = ? + """, + (name,), + ).fetchone() + + if row is None: + return None + + return { + "payload": json.loads(row["payload_json"] or "{}"), + "test_count": int(row["test_count"]), + "updated_at": row["updated_at"], + } + + +def delete_graph_cache(name: str, db_path: str | Path = DB_PATH) -> None: + with get_connection(db_path) as conn: + conn.execute("DELETE FROM graph_cache WHERE name = ?", (name,)) + + def _parse_rule_tokens(rule: str | None) -> list[str]: if not rule: return [] @@ -543,6 +596,110 @@ def mark_tests_completed(test_ids_with_device: list[tuple[str, str]], db_path: s ) +def mark_overdue_as_rerun(db_path: str | Path = DB_PATH) -> int: + """Mark tests from the last overnight window that are still pending as rerun. + + 'Last overnight window' = all pending tests scheduled before today (any shift) + plus today's shift 1 (1am–10am) if it has already ended (current hour >= 10). + + Returns the number of tests newly marked as rerun. + """ + from datetime import date as _date, datetime as _datetime + now = _datetime.now() + today = now.date().isoformat() + shift1_ended = now.hour >= 10 # shift 1 ends ~10am + + with get_connection(db_path) as conn: + version_row = conn.execute("SELECT MAX(schedule_version) AS latest FROM schedules").fetchone() + latest = version_row["latest"] + if latest is None: + return 0 + + if shift1_ended: + rows = conn.execute( + """ + SELECT DISTINCT s.test_id, s.device + FROM schedules s + JOIN tests t ON t.test_id = s.test_id AND t.device = s.device + WHERE s.schedule_version = ? + AND t.status = 'pending' + AND ( + s.scheduled_date < ? + OR (s.scheduled_date = ? AND s.shift_index = 1) + ) + """, + (latest, today, today), + ).fetchall() + else: + rows = conn.execute( + """ + SELECT DISTINCT s.test_id, s.device + FROM schedules s + JOIN tests t ON t.test_id = s.test_id AND t.device = s.device + WHERE s.schedule_version = ? + AND s.scheduled_date < ? + AND t.status = 'pending' + """, + (latest, today), + ).fetchall() + + if not rows: + return 0 + + conn.executemany( + """ + UPDATE tests + SET status = 'rerun', updated_at = CURRENT_TIMESTAMP + WHERE test_id = ? AND device = ? AND status = 'pending' + """, + [(r["test_id"], r["device"]) for r in rows], + ) + return len(rows) + + +def get_rerun_tests(db_path: str | Path = DB_PATH) -> list[dict]: + """Return tests scheduled before shift 2 today that are not yet completed. + + This checks the latest schedule version and returns tests from: + - Yesterday's shift 3 (5pm-1am) + - Today's shift 1 (1am-10am) + + Once tests are rescheduled to shift 2 or later today, they no longer appear. + """ + from datetime import date as _date + today = _date.today().isoformat() + with get_connection(db_path) as conn: + version_row = conn.execute("SELECT MAX(schedule_version) AS latest FROM schedules").fetchone() + latest = version_row["latest"] + if latest is None: + return [] + + rows = conn.execute( + """ + SELECT DISTINCT t.test_id, t.device, t.test_type, t.estimated_minutes + FROM tests t + JOIN schedules s ON s.test_id = t.test_id AND s.device = t.device + WHERE s.schedule_version = ? + AND t.status != 'completed' + AND ( + (s.scheduled_date = date(?, '-1 day') AND s.shift_index = 3) + OR (s.scheduled_date = ? AND s.shift_index = 1) + ) + ORDER BY t.test_id, t.device + """, + (latest, today, today), + ).fetchall() + return [ + { + "test_id": r["test_id"], + "device": r["device"], + "test_type": r["test_type"], + "estimated_minutes": r["estimated_minutes"], + } + for r in rows + ] + + def upsert_holidays(dates: list[str], db_path: str | Path = DB_PATH) -> None: """Replace all holidays with the provided list of YYYY-MM-DD date strings.""" with get_connection(db_path) as conn: diff --git a/backend/graph.py b/backend/graph.py index d1a5cbf..364429a 100644 --- a/backend/graph.py +++ b/backend/graph.py @@ -1,12 +1,17 @@ from __future__ import annotations import json -from typing import Any +import os from pathlib import Path +from typing import Any # Immutable-in-practice graph state for the loaded DUT test set. _TESTS_BY_ID: dict[str, Any] = {} _GRAPH: dict[str, set[str]] = {} +DUT = os.getenv("DUT", "CGW453").strip() +APP_ROOT = Path(__file__).resolve().parent +DB_PATH = APP_ROOT / "scheduler.db" +GRAPH_CACHE_NAME = "dut_compatibility" def reset_graph_state() -> None: @@ -44,7 +49,58 @@ def build_graph_once(tests: list[Any]) -> dict[str, set[str]]: return _GRAPH -def get_graph() -> dict[str, set[str]]: +def serialize_graph(graph: dict[str, set[str]]) -> dict[str, list[str]]: + return {str(test_id): sorted(str(neighbor) for neighbor in neighbors) for test_id, neighbors in graph.items()} + + +def deserialize_graph(data: dict[str, list[str]]) -> dict[str, set[str]]: + deserialized: dict[str, set[str]] = {} + for test_id, neighbors in (data or {}).items(): + deserialized[str(test_id)] = {str(neighbor) for neighbor in (neighbors or [])} + return deserialized + + +def _persist_graph(db_path: str | Path = DB_PATH) -> None: + import db as db_module + db_module.save_graph_cache( + name=GRAPH_CACHE_NAME, + payload=serialize_graph(_GRAPH), + test_count=len(_GRAPH), + db_path=db_path, + ) + + +def _load_graph_from_db(db_path: str | Path = DB_PATH) -> bool: + global _GRAPH + import db as db_module + cached = db_module.load_graph_cache(GRAPH_CACHE_NAME, db_path) + if cached is None: + return False + _GRAPH = deserialize_graph(cached.get("payload") or {}) + return True + + +def build_and_persist_graph(tests: list[Any], db_path: str | Path = DB_PATH) -> dict[str, set[str]]: + graph = build_graph_once(tests) + _persist_graph(db_path) + return graph + + +def get_graph( + db_path: str | Path = DB_PATH, + dut_tests: list[Any] | None = None, +) -> dict[str, set[str]]: + if _GRAPH: + return _GRAPH + + if _load_graph_from_db(db_path): + return _GRAPH + + if dut_tests is None: + import db as db_module + dut_tests = db_module.list_tests_for_device(DUT, db_path) + + build_and_persist_graph(dut_tests or [], db_path) return _GRAPH @@ -60,10 +116,6 @@ def _build_graph(tests: dict[str, Any]) -> dict[str, set[str]]: if _compatible(a, b): graph[a_id].add(b_id) graph[b_id].add(a_id) - # Print out graph for debugging - out = Path('output1.txt') - with out.open('w', encoding='utf-8') as f: - f.write(json.dumps({k: list(v) for k, v in graph.items()}, indent=2)) return graph @@ -93,7 +145,6 @@ def _same_testpoint_to_station( stations_a = {s for s, t in station_map_a.items() if t == testpoint} stations_b = {s for s, t in station_map_b.items() if t == testpoint} if stations_a != stations_b: - print(f"Tests {test_a.test_id} and {test_b.test_id} have conflicting stations for testpoint. {station_map_a} vs {station_map_b}") return False return True diff --git a/backend/parser.py b/backend/parser.py index 0ff6a0b..26e75a6 100644 --- a/backend/parser.py +++ b/backend/parser.py @@ -1,7 +1,14 @@ import csv import os +import re from dataclasses import dataclass from pathlib import Path +from typing import Any + +try: + import smbclient # type: ignore[import-not-found] +except ModuleNotFoundError: + smbclient = None from db import DEVICE_DUT, DEVICE_REF, TestRecord @@ -81,6 +88,126 @@ class CsvValidationError(ValueError): pass +_SMB_SESSIONS: set[str] = set() + + +def _normalize_input_path(path_value: str | Path) -> str: + path = str(path_value).strip() + if not path: + return path + + # Accept //server/share style and normalize to UNC for smbclient. + if path.startswith("//"): + return "\\\\" + path.lstrip("/").replace("/", "\\") + + # Accept ///... and normalize to UNC for Linux-hosted inputs. + if re.match(r"^/\d{1,3}(?:\.\d{1,3}){3}/[^/]+", path): + return "\\\\" + path.lstrip("/").replace("/", "\\") + + return path + + +def _is_unc_path(path: str) -> bool: + return path.startswith("\\\\") + + +def _extract_unc_server(path: str) -> str | None: + if not _is_unc_path(path): + return None + rest = path[2:] + return rest.split("\\", 1)[0] if rest else None + + +def _normalize_smb_credentials(smb_credentials: dict[str, Any] | None) -> tuple[str, str, str]: + username = "" + password = "" + domain = "" + + if smb_credentials: + username = str(smb_credentials.get("username", "")).strip() + password = str(smb_credentials.get("password", "")) + domain = str(smb_credentials.get("domain", "")).strip() + + if not username: + username = os.getenv("SMB_USERNAME", "").strip() + if not password: + password = os.getenv("SMB_PASSWORD", "") + if not domain: + domain = os.getenv("SMB_DOMAIN", "").strip() + + if username and domain and "\\" not in username and "@" not in username: + username = f"{domain}\\{username}" + + return username, password, domain + + +def _register_smb_session_if_needed(path: str, smb_credentials: dict[str, Any] | None) -> None: + if not _is_unc_path(path): + return + + if smbclient is None: + raise ModuleNotFoundError("smbclient is required to read CSV files from UNC paths") + + server = _extract_unc_server(path) + if not server or server in _SMB_SESSIONS: + return + + username, password, _ = _normalize_smb_credentials(smb_credentials) + if username: + smbclient.register_session(server, username=username, password=password) + else: + smbclient.register_session(server) + + _SMB_SESSIONS.add(server) + + +def _path_exists(path: str, smb_credentials: dict[str, Any] | None) -> bool: + if _is_unc_path(path): + _register_smb_session_if_needed(path, smb_credentials) + try: + smbclient.stat(path) + return True + except OSError: + return False + + return Path(path).exists() + + +def _is_dir(path: str, smb_credentials: dict[str, Any] | None) -> bool: + if _is_unc_path(path): + _register_smb_session_if_needed(path, smb_credentials) + try: + entry_iter = smbclient.scandir(path) + for _ in entry_iter: + break + return True + except OSError: + return False + + return Path(path).is_dir() + + +def _path_name(path: str) -> str: + trimmed = path.rstrip("\\/") + if not trimmed: + return path + parts = re.split(r"[\\/]", trimmed) + return parts[-1] if parts else trimmed + + +def _csv_paths_from_dir(path: str, smb_credentials: dict[str, Any] | None) -> list[str]: + if _is_unc_path(path): + _register_smb_session_if_needed(path, smb_credentials) + entries = [] + for entry in smbclient.scandir(path): + name = getattr(entry, "name", "") + if name and name.lower().endswith(".csv") and entry.is_file(): + entries.append(path.rstrip("\\/") + "\\" + name) + return sorted(entries) + + return [str(p) for p in sorted(Path(path).glob("*.csv"))] + + def _columns_missing(normalized_fieldnames: set[str], required_columns: list[str]) -> list[str]: return [ column @@ -104,20 +231,25 @@ def _detect_csv_format(normalized_fieldnames: set[str]) -> str: ) -def parse_target_csv(csv_path: str | Path | list[str | Path] | tuple[str | Path, ...]) -> ParseResult: - paths = _resolve_csv_paths(csv_path) +def parse_target_csv( + csv_path: str | Path | list[str | Path] | tuple[str | Path, ...], + smb_credentials: dict[str, Any] | None = None, + runtime_overrides: dict[str, Any] | None = None, +) -> ParseResult: + paths = _resolve_csv_paths(csv_path, smb_credentials=smb_credentials) + runtime_defaults = _resolve_runtime_defaults(runtime_overrides) all_tests: list[TestRecord] = [] all_warnings: list[str] = [] seen_test_keys: set[tuple[str, str]] = set() for path in paths: - parsed = _parse_single_csv(path) + parsed = _parse_single_csv(path, smb_credentials=smb_credentials, runtime_defaults=runtime_defaults) all_warnings.extend(parsed.warnings) for record in parsed.tests: key = (record.test_id, record.device) if key in seen_test_keys: all_warnings.append( - f"File {path.name}: duplicate test/device '{record.test_id}/{record.device}', record skipped." + f"File {_path_name(path)}: duplicate test/device '{record.test_id}/{record.device}', record skipped." ) continue seen_test_keys.add(key) @@ -126,17 +258,43 @@ def parse_target_csv(csv_path: str | Path | list[str | Path] | tuple[str | Path, return ParseResult(tests=all_tests, warnings=all_warnings) +def _resolve_runtime_defaults(runtime_overrides: dict[str, Any] | None) -> dict[str, int]: + defaults = dict(RUNTIME_DEFAULTS) + if not runtime_overrides: + return defaults -def _resolve_csv_paths(csv_path: str | Path | list[str | Path] | tuple[str | Path, ...]) -> list[Path]: + for test_type in ("P2P", "COE", "P3P"): + raw_value = runtime_overrides.get(test_type) + if raw_value is None: + continue + if isinstance(raw_value, str): + raw_value = raw_value.strip() + if not raw_value: + continue + try: + minutes = int(raw_value) + except (TypeError, ValueError): + continue + if minutes > 0: + defaults[test_type] = minutes + + return defaults + + + +def _resolve_csv_paths( + csv_path: str | Path | list[str | Path] | tuple[str | Path, ...], + smb_credentials: dict[str, Any] | None, +) -> list[str]: if isinstance(csv_path, (str, Path)): - paths = [Path(csv_path)] + paths = [_normalize_input_path(csv_path)] else: - paths = [Path(item) for item in csv_path] + paths = [_normalize_input_path(item) for item in csv_path] - resolved_paths: list[Path] = [] + resolved_paths: list[str] = [] for path in paths: - if path.is_dir(): - resolved_paths.extend(sorted(path.glob("*.csv"))) + if _is_dir(path, smb_credentials=smb_credentials): + resolved_paths.extend(_csv_paths_from_dir(path, smb_credentials=smb_credentials)) else: resolved_paths.append(path) @@ -144,15 +302,26 @@ def _resolve_csv_paths(csv_path: str | Path | list[str | Path] | tuple[str | Pat raise FileNotFoundError("No CSV files found to parse.") for path in resolved_paths: - if not path.exists(): + if not _path_exists(path, smb_credentials=smb_credentials): raise FileNotFoundError(f"CSV file not found: {path}") return resolved_paths -def _parse_single_csv(path: Path) -> ParseResult: - with path.open("r", encoding="utf-8-sig", newline="") as handle: +def _open_csv_handle(path: str, smb_credentials: dict[str, Any] | None): + if _is_unc_path(path): + _register_smb_session_if_needed(path, smb_credentials) + return smbclient.open_file(path, mode="r", encoding="utf-8-sig", newline="") + return Path(path).open("r", encoding="utf-8-sig", newline="") + + +def _parse_single_csv( + path: str, + smb_credentials: dict[str, Any] | None, + runtime_defaults: dict[str, int], +) -> ParseResult: + with _open_csv_handle(path, smb_credentials=smb_credentials) as handle: reader = csv.DictReader(handle) if not reader.fieldnames: raise CsvValidationError("CSV is missing a header row.") @@ -169,11 +338,11 @@ def _parse_single_csv(path: Path) -> ParseResult: for row_num, row in enumerate(reader, start=2): test_id = (row.get("TC ID") or "").strip() if not test_id: - warnings.append(f"File {path.name}, row {row_num}: missing TC ID, row skipped.") + warnings.append(f"File {_path_name(path)}, row {row_num}: missing TC ID, row skipped.") continue if test_id in seen_test_ids: - warnings.append(f"File {path.name}, row {row_num}: duplicate TC ID '{test_id}', row skipped.") + warnings.append(f"File {_path_name(path)}, row {row_num}: duplicate TC ID '{test_id}', row skipped.") continue seen_test_ids.add(test_id) @@ -184,7 +353,7 @@ def _parse_single_csv(path: Path) -> ParseResult: has_coe_pair = _normalize_yes_no(_row_get(row, "COE Pair")) if csv_format == "p2p_coe" else False config = _build_config(row, csv_format) signature = _victim_band_signature(row) if csv_format == "p2p_coe" else None - estimated_minutes = RUNTIME_DEFAULTS.get(test_type) + estimated_minutes = runtime_defaults.get(test_type, RUNTIME_DEFAULTS[test_type]) victim_band_source = "Victim Band" if csv_format == "p2p_coe" else "Band" victim_band = _normalize_victim_band(_row_get(row, victim_band_source)) diff --git a/backend/requirements.txt b/backend/requirements.txt index 0969745..9670c67 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -3,3 +3,4 @@ uvicorn==0.35.0 pydantic==2.11.7 watchdog==6.0.0 httpx==0.28.1 +smbprotocol==1.15.0 diff --git a/backend/scanner.py b/backend/scanner.py index 4e49d29..329b339 100644 --- a/backend/scanner.py +++ b/backend/scanner.py @@ -3,11 +3,11 @@ import re import threading try: - import smbclient + import smbclient # type: ignore[import-not-found] except ModuleNotFoundError: smbclient = None -from db import DEVICE_DUT, DEVICE_REF, mark_tests_completed +from db import DEVICE_DUT, DEVICE_REF, mark_tests_completed, mark_overdue_as_rerun _SMB_SESSIONS = set() _SCAN_STATE_LOCK = threading.Lock() @@ -15,6 +15,29 @@ _ACTIVE_SCAN_COUNT = 0 _RESULT_TEST_ID_PATTERN = re.compile(r"(?:COE|P2P|P3P)(?:RX|TX)?[A-Z]{2}\d{3}", re.IGNORECASE) +def _normalize_smb_credentials(smb_credentials=None): + username = "" + password = "" + domain = "" + + if smb_credentials: + username = str(smb_credentials.get("username", "")).strip() + password = str(smb_credentials.get("password", "")) + domain = str(smb_credentials.get("domain", "")).strip() + + if not username: + username = os.getenv("SMB_USERNAME", "").strip() + if not password: + password = os.getenv("SMB_PASSWORD", "") + if not domain: + domain = os.getenv("SMB_DOMAIN", "").strip() + + if username and domain and "\\" not in username and "@" not in username: + username = f"{domain}\\{username}" + + return username, password + + def _extract_test_id_from_result_dir_name(dir_name): if not dir_name: return None @@ -171,7 +194,7 @@ def _extract_unc_server(path): return rest.split("\\", 1)[0] if rest else None -def _register_smb_session_if_needed(path): +def _register_smb_session_if_needed(path, smb_credentials=None): if not _is_unc_path(path): return @@ -182,12 +205,7 @@ def _register_smb_session_if_needed(path): if not server or server in _SMB_SESSIONS: return - username = os.getenv("SMB_USERNAME", "").strip() - password = os.getenv("SMB_PASSWORD", "") - domain = os.getenv("SMB_DOMAIN", "").strip() - - if username and domain and "\\" not in username and "@" not in username: - username = f"{domain}\\{username}" + username, password = _normalize_smb_credentials(smb_credentials) if username: smbclient.register_session(server, username=username, password=password) @@ -197,17 +215,17 @@ def _register_smb_session_if_needed(path): _SMB_SESSIONS.add(server) -def _iter_dir_entries(path): +def _iter_dir_entries(path, smb_credentials=None): path = _normalize_input_path(path) if _is_unc_path(path): if smbclient is None: raise ModuleNotFoundError("smbclient is required to scan UNC result paths") - _register_smb_session_if_needed(path) + _register_smb_session_if_needed(path, smb_credentials=smb_credentials) return list(smbclient.scandir(path)) return list(os.scandir(path)) -def scan_results(results_dir_dut, results_dir_ref): +def scan_results(results_dir_dut, results_dir_ref, smb_credentials=None): results_dir_dut = _normalize_input_path(results_dir_dut) results_dir_ref = _normalize_input_path(results_dir_ref) @@ -217,7 +235,7 @@ def scan_results(results_dir_dut, results_dir_ref): try: dut_entries = [ entry.name - for entry in _iter_dir_entries(results_dir_dut) + for entry in _iter_dir_entries(results_dir_dut, smb_credentials=smb_credentials) if entry.is_dir() ] except OSError as exc: @@ -229,7 +247,7 @@ def scan_results(results_dir_dut, results_dir_ref): try: ref_entries = [ entry.name - for entry in _iter_dir_entries(results_dir_ref) + for entry in _iter_dir_entries(results_dir_ref, smb_credentials=smb_credentials) if entry.is_dir() ] except OSError as exc: @@ -259,4 +277,7 @@ def scan_results(results_dir_dut, results_dir_ref): print(f"[scanner] skipped {len(unmatched_entries)} result dir(s) with no recognizable test id") mark_tests_completed(completed_batch) + newly_rerun = mark_overdue_as_rerun() + if newly_rerun: + print(f"[scanner] {newly_rerun} test(s) marked as rerun-required (scheduled but not completed)") diff --git a/backend/scheduler.py b/backend/scheduler.py index 8948612..46603d6 100644 --- a/backend/scheduler.py +++ b/backend/scheduler.py @@ -125,8 +125,6 @@ def compile_schedule( """ if not tests: return [], None - if not graph.is_graph_built(): - raise RuntimeError("Graph not initialized. Load tests first to build DUT compatibility graph.") initialize_scheduler_state(tests) set_user_priorities(top_priority_tests, lowest_priority_tests) diff --git a/backend/watcher.py b/backend/watcher.py index e69de29..a3ab930 100644 --- a/backend/watcher.py +++ b/backend/watcher.py @@ -0,0 +1,161 @@ +from __future__ import annotations + +import logging +import os +import threading +import time +from dataclasses import dataclass +from typing import Any + +from watchdog.events import FileSystemEvent, FileSystemEventHandler, FileSystemMovedEvent +from watchdog.observers.polling import PollingObserver + +from scanner import resolve_runtime_path, scan_results + + +LOGGER = logging.getLogger("scheduler.watcher") + + +@dataclass(frozen=True) +class WatchConfig: + dut_dir: str + ref_dir: str + smb_username: str + smb_password: str + smb_domain: str + + +class _ResultDirEventHandler(FileSystemEventHandler): + def __init__(self, on_new_directory): + super().__init__() + self._on_new_directory = on_new_directory + + def on_created(self, event: FileSystemEvent) -> None: + if event.is_directory: + self._on_new_directory(event.src_path) + + def on_moved(self, event: FileSystemMovedEvent) -> None: + if event.is_directory: + self._on_new_directory(event.dest_path) + + +class ResultDirectoryWatcher: + """Watch DUT/REF result directories and trigger scans when new result folders appear.""" + + def __init__(self, poll_interval_seconds: float = 2.0, min_scan_interval_seconds: float = 1.0): + self._poll_interval_seconds = poll_interval_seconds + self._min_scan_interval_seconds = min_scan_interval_seconds + self._lock = threading.RLock() + self._scan_lock = threading.Lock() + self._observer: PollingObserver | None = None + self._current_config: WatchConfig | None = None + self._last_scan_monotonic = 0.0 + + def configure_from_settings(self, settings: dict[str, Any]) -> None: + dut_dir = _normalize_watch_path(settings.get("dutResultDir")) + ref_dir = _normalize_watch_path(settings.get("refResultDir")) + new_config = WatchConfig( + dut_dir=dut_dir, + ref_dir=ref_dir, + smb_username=str(settings.get("smbUsername") or "").strip(), + smb_password=str(settings.get("smbPassword") or ""), + smb_domain=str(settings.get("smbDomain") or "").strip(), + ) + + with self._lock: + if new_config == self._current_config and self._observer is not None: + return + + self._stop_locked() + self._current_config = new_config + + if not dut_dir or not ref_dir: + LOGGER.info("result watcher disabled: DUT/REF result directories are not both set") + return + + handler = _ResultDirEventHandler(self._on_new_directory) + observer = PollingObserver(timeout=self._poll_interval_seconds) + + scheduled_count = 0 + for result_dir, label in ((dut_dir, "DUT"), (ref_dir, "REF")): + if not os.path.isdir(result_dir): + LOGGER.warning("%s result directory does not exist yet, skipping watch: %s", label, result_dir) + continue + observer.schedule(handler, result_dir, recursive=False) + scheduled_count += 1 + + if scheduled_count == 0: + LOGGER.warning("result watcher not started: no valid result directories to watch") + return + + observer.start() + self._observer = observer + LOGGER.info("result watcher started for DUT=%s REF=%s", dut_dir, ref_dir) + + self.scan_now(reason="startup") + + def _on_new_directory(self, directory_path: str) -> None: + LOGGER.info("new result directory detected: %s", directory_path) + self.scan_now(reason="directory-created") + + def scan_now(self, reason: str = "manual") -> None: + with self._lock: + config = self._current_config + if config is None or not config.dut_dir or not config.ref_dir: + return + + now = time.monotonic() + if now - self._last_scan_monotonic < self._min_scan_interval_seconds: + return + self._last_scan_monotonic = now + + if not self._scan_lock.acquire(blocking=False): + return + + def _run_scan() -> None: + credentials = { + "username": config.smb_username, + "password": config.smb_password, + "domain": config.smb_domain, + } + try: + scan_results(config.dut_dir, config.ref_dir, smb_credentials=credentials) + LOGGER.info("result scan finished (%s)", reason) + except Exception: + LOGGER.exception("result scan failed (%s)", reason) + finally: + self._scan_lock.release() + + threading.Thread(target=_run_scan, daemon=True).start() + + def stop(self) -> None: + with self._lock: + self._stop_locked() + + def _stop_locked(self) -> None: + observer = self._observer + self._observer = None + if observer is not None: + observer.stop() + observer.join(timeout=5) + LOGGER.info("result watcher stopped") + + +_WATCHER = ResultDirectoryWatcher() + + +def configure_result_watcher(settings: dict[str, Any]) -> None: + _WATCHER.configure_from_settings(settings) + + +def stop_result_watcher() -> None: + _WATCHER.stop() + + +def _normalize_watch_path(path_value: Any) -> str: + if path_value is None: + return "" + cleaned = str(path_value).strip() + if not cleaned: + return "" + return str(resolve_runtime_path(cleaned)).strip() diff --git a/frontend/src/App.jsx b/frontend/src/App.jsx index 07a145d..46a3a61 100644 --- a/frontend/src/App.jsx +++ b/frontend/src/App.jsx @@ -17,6 +17,17 @@ function getMondayOfWeek(date) { return d } +function parseLocalDate(value) { + if (!value) return null + + const [year, month, day] = value.split('-').map(Number) + if (!year || !month || !day) return null + + const parsed = new Date(year, month - 1, day) + parsed.setHours(0, 0, 0, 0) + return Number.isNaN(parsed.getTime()) ? null : parsed +} + function toKey(d) { return d.toISOString().slice(0, 10) } function addDays(date, n) { @@ -85,9 +96,22 @@ const DEFAULT_SETTINGS = { p3pCsvPath: '', refResultDir: '', dutResultDir: '', + smbUsername: '', + smbPassword: '', + smbDomain: '', + p2pRuntimeMinutes: '', + coeRuntimeMinutes: '', + p3pRuntimeMinutes: '', testExclusion: '', holidays: '', - startDateOverride: '', +} + +function sanitizeSettings(saved = {}) { + const { startDateOverride: _ignored, ...rest } = saved + return { + ...DEFAULT_SETTINGS, + ...rest, + } } // --------------------------------------------------------------------------- @@ -98,6 +122,7 @@ export default function App() { const [daytimeEnabled, setDaytimeEnabled] = useState(false) const [topPriority, setTopPriority] = useState('') const [lowestPriority, setLowestPriority] = useState('') + const [startDateOverride, setStartDateOverride] = useState('') const [failedTests, setFailedTests] = useState([]) const [scheduleData, setScheduleData] = useState({}) const [completionDate, setCompletionDate] = useState(null) @@ -116,6 +141,15 @@ export default function App() { } }, []) + const fetchRerunTests = useCallback(async () => { + try { + const data = await api.getRerunTests() + setFailedTests(data.tests ?? []) + } catch (e) { + console.error('Failed to fetch rerun tests:', e) + } + }, []) + // On mount: load settings + holidays + schedule for today's week useEffect(() => { async function init() { @@ -124,9 +158,8 @@ export default function App() { api.getSettings(), api.getHolidays(), ]) - setSettings(prev => ({ - ...prev, - ...saved, + setSettings(() => ({ + ...sanitizeSettings(saved), holidays: (holidayData.dates ?? []).join(', '), })) } catch (e) { @@ -142,26 +175,43 @@ export default function App() { fetchSchedule(weekStart) }, [weekStart, fetchSchedule]) + // Keep calendar statuses fresh when backend watcher marks tests from new result folders. + useEffect(() => { + const timer = setInterval(() => { + fetchSchedule(weekStart) + fetchRerunTests() + }, 5000) + + return () => clearInterval(timer) + }, [weekStart, fetchSchedule, fetchRerunTests]) + + // Load rerun tests on mount + useEffect(() => { + fetchRerunTests() + }, [fetchRerunTests]) + // ------------------------------------------------------------------------- async function handleSaveSettings(newSettings) { setLoading(true) setError(null) try { - await api.saveSettings(newSettings) + const sanitizedSettings = sanitizeSettings(newSettings) - const holidayDates = (newSettings.holidays ?? '') + await api.saveSettings(sanitizedSettings) + + const holidayDates = (sanitizedSettings.holidays ?? '') .split(',').map(s => s.trim()).filter(Boolean) await api.saveHolidays(holidayDates) - if (newSettings.p2pCoeCsvPath?.trim()) { - await api.loadCsv(newSettings.p2pCoeCsvPath.trim()) + if (sanitizedSettings.p2pCoeCsvPath?.trim()) { + await api.loadCsv(sanitizedSettings.p2pCoeCsvPath.trim()) } - if (newSettings.p3pCsvPath?.trim()) { - await api.loadCsv(newSettings.p3pCsvPath.trim()) + if (sanitizedSettings.p3pCsvPath?.trim()) { + await api.loadCsv(sanitizedSettings.p3pCsvPath.trim()) } - setSettings(newSettings) + setSettings(sanitizedSettings) } catch (e) { setError(e.message) } finally { @@ -173,15 +223,30 @@ export default function App() { setLoading(true) setError(null) try { + const effectiveStartDate = startDateOverride.trim() const result = await api.compileSchedule({ - start_date: settings.startDateOverride?.trim() || null, + start_date: effectiveStartDate || null, daytime_testing_today: daytimeEnabled, top_priority_tests: topPriority.split(',').map(s => s.trim()).filter(Boolean), lowest_priority_tests: lowestPriority.split(',').map(s => s.trim()).filter(Boolean), rule: settings.testExclusion ?? '', }) setCompletionDate(result.completion_date ?? null) - await fetchSchedule(weekStart) + + if (effectiveStartDate) { + const overrideDate = parseLocalDate(effectiveStartDate) + if (overrideDate) { + const overrideWeekStart = getMondayOfWeek(overrideDate) + setWeekStart(overrideWeekStart) + await fetchSchedule(overrideWeekStart) + } else { + await fetchSchedule(weekStart) + } + } else { + await fetchSchedule(weekStart) + } + + setStartDateOverride('') } catch (e) { setError(e.message) } finally { @@ -190,8 +255,16 @@ export default function App() { } function handleRerunDecision(rerunDuringDay) { - // TODO: POST /api/failed-tests/rerun-decision once backend endpoint exists - console.log('Rerun during day:', rerunDuringDay) + const rerunIds = failedTests.map(t => t.test_id).join(', ') + setTopPriority(prev => { + const existing = prev.split(',').map(s => s.trim()).filter(Boolean) + const incoming = failedTests.map(t => t.test_id) + const merged = [...new Set([...existing, ...incoming])] + return merged.join(', ') + }) + if (rerunDuringDay) { + setDaytimeEnabled(true) + } setFailedTests([]) } @@ -209,8 +282,7 @@ export default function App() { )} @@ -226,6 +298,8 @@ export default function App() { request('POST', '/schedule/compile', opts), getScheduleWeek: (start) => request('GET', `/schedule/week?start=${start}`), + getRerunTests: () => request('GET', '/tests/rerun'), } // Transform the flat items array from GET /api/schedule/week into the diff --git a/frontend/src/components/Calendar.jsx b/frontend/src/components/Calendar.jsx index e86d14e..174fb91 100644 --- a/frontend/src/components/Calendar.jsx +++ b/frontend/src/components/Calendar.jsx @@ -121,8 +121,7 @@ export default function Calendar({ scheduleData = {}, daytimeDateKey = null, wee
Pending Completed - Failed - Invalid + Rerun Required
) diff --git a/frontend/src/components/FailedBanner.jsx b/frontend/src/components/FailedBanner.jsx index b80fc2b..9ea1b8a 100644 --- a/frontend/src/components/FailedBanner.jsx +++ b/frontend/src/components/FailedBanner.jsx @@ -1,9 +1,11 @@ -export default function FailedBanner({ failedTests, estimatedMinutes, onDecision }) { - if (!failedTests || failedTests.length === 0) return null +export default function FailedBanner({ rerunTests, onDecision }) { + if (!rerunTests || rerunTests.length === 0) return null - const hours = Math.floor(estimatedMinutes / 60) - const mins = estimatedMinutes % 60 + const totalMinutes = rerunTests.reduce((sum, t) => sum + (t.estimated_minutes ?? 0), 0) + const hours = Math.floor(totalMinutes / 60) + const mins = totalMinutes % 60 const timeStr = hours > 0 ? `${hours}h ${mins}m` : `${mins}m` + const testIds = rerunTests.map(t => t.test_id).join(', ') return (
@@ -23,13 +25,13 @@ export default function FailedBanner({ failedTests, estimatedMinutes, onDecision

- {failedTests.length} test{failedTests.length !== 1 ? 's' : ''} - {' '}failed last night and require a rerun —{' '} - {failedTests.join(', ')} + {rerunTests.length} test{rerunTests.length !== 1 ? 's' : ''} + {' '}require a rerun — not completed in last overnight window:{' '} + {testIds}

Estimated rerun time: {timeStr} - {' '}— Rerun these tests during the day today? + {' '}— Schedule rerun during daytime today?

diff --git a/frontend/src/components/RightPanel.jsx b/frontend/src/components/RightPanel.jsx index a3392f0..4e8c979 100644 --- a/frontend/src/components/RightPanel.jsx +++ b/frontend/src/components/RightPanel.jsx @@ -14,6 +14,8 @@ function formatCompletionDate(value) { export default function RightPanel({ completionDate, + startDateOverride, + onStartDateOverrideChange, daytimeEnabled, onDaytimeEnabledChange, topPriority, @@ -75,6 +77,18 @@ export default function RightPanel({ />
+
+ + onStartDateOverrideChange(e.target.value)} + className="w-full bg-gray-900 border border-gray-600 rounded-md px-2.5 py-1.5 text-xs text-gray-200 focus:outline-none focus:border-blue-500 font-mono" + /> +
+ {/* Remake schedule */}