implemented watcher and rerun logic

This commit is contained in:
2026-06-25 11:31:20 -04:00
parent 21402f7ee3
commit be13849a4f
15 changed files with 871 additions and 97 deletions
+67 -4
View File
@@ -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()]
+158 -1
View File
@@ -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 (1am10am) 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:
+58 -7
View File
@@ -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
+185 -16
View File
@@ -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 /<ipv4>/<share>/... 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))
+1
View File
@@ -3,3 +3,4 @@ uvicorn==0.35.0
pydantic==2.11.7
watchdog==6.0.0
httpx==0.28.1
smbprotocol==1.15.0
+35 -14
View File
@@ -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)")
-2
View File
@@ -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)
+161
View File
@@ -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()