from contextlib import asynccontextmanager from pathlib import Path from typing import Any from datetime import date, datetime import os from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel, Field 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 DB_PATH = APP_ROOT / "scheduler.db" DUT = os.getenv("DUT", "CGW453").strip() REF = os.getenv("REF", "CGW452").strip() class LoadTestsRequest(BaseModel): csv_path: str = Field(..., description="Absolute or backend-relative path to target CSV") 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 = "" daytime_testing_today: bool = False top_priority_tests: list[str] = Field(default_factory=list) lowest_priority_tests: list[str] = Field(default_factory=list) class RemoveActiveTestsRequest(BaseModel): test_ids: list[str] = Field(default_factory=list) class SaveHolidaysRequest(BaseModel): dates: list[str] = Field(default_factory=list) @asynccontextmanager async def lifespan(application: FastAPI): db.init_db(DB_PATH) graph.reset_graph_state() 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) app.add_middleware( CORSMiddleware, allow_origins=["http://localhost:5173", "http://127.0.0.1:5173"], allow_methods=["*"], allow_headers=["*"], ) @app.get("/api/health") def health() -> dict[str, str]: return {"status": "ok"} @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"} @app.get("/api/settings") def get_settings() -> dict[str, Any]: return db.read_settings(DB_PATH) @app.post("/api/tests/load") def load_tests(request: LoadTestsRequest) -> dict[str, Any]: csv_path = Path(request.csv_path) 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, 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: raise HTTPException(status_code=400, detail=str(exc)) from exc count = db.upsert_tests(parsed.tests, DB_PATH) reset_scheduler_state() graph.reset_graph_state() all_dut_tests = db.list_tests_for_device(DUT, DB_PATH) graph.build_and_persist_graph(all_dut_tests, DB_PATH) return { "loaded_tests": count, "warnings": parsed.warnings, "dut_graph_nodes": len(graph.get_graph()), } @app.post("/api/schedule/active/remove") def remove_active_tests(request: RemoveActiveTestsRequest) -> dict[str, Any]: remove_from_active({item.strip() for item in request.test_ids if item.strip()}) return {"status": "ok", "removed": len(request.test_ids)} @app.post("/api/schedule/compile") def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any]: if request.start_date: try: datetime.strptime(request.start_date, "%Y-%m-%d") 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) return { "schedule_version": version, "scheduled_tests": 0, "completion_date": None, } scheduler_tests = [ SchedulerTest( test_id=t.test_id, device=t.device, test_type=t.test_type, rotation=t.rotation, rx_tx=t.rx_tx, has_coe_pair=t.has_coe_pair, coe_pairing=t.coe_pairing or [], config=t.config, throttled=t.throttled, estimated_minutes=t.estimated_minutes, priority=t.priority, raw_payload=t.raw_payload or {}, ) 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( tests=scheduler_tests, start_date=request.start_date, holiday_dates=holiday_dates, top_priority_tests={item.strip() for item in request.top_priority_tests if item.strip()}, lowest_priority_tests={item.strip() for item in request.lowest_priority_tests if item.strip()}, daytime_testing_today=request.daytime_testing_today, ) except RuntimeError as exc: raise HTTPException(status_code=409, detail=str(exc)) from exc version = db.create_schedule_version( [(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), "completion_date": completion_date, } @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()] db.upsert_holidays(dates, DB_PATH) return {"status": "saved", "count": len(dates)} @app.get("/api/holidays") def get_holidays() -> dict[str, Any]: return {"dates": sorted(db.list_holidays(DB_PATH))} @app.get("/api/schedule/week") def get_schedule_week(start: str | None = None) -> dict[str, Any]: week_start = start or date.today().isoformat() try: datetime.strptime(week_start, "%Y-%m-%d") except ValueError as exc: raise HTTPException(status_code=400, detail="start must be YYYY-MM-DD") from exc rows = db.get_schedule_week(week_start, DB_PATH) return { "start_date": week_start, "items": [ { "test_id": row.test_id, "device": row.device, "scheduled_date": row.scheduled_date, "shift_index": row.shift_index, "sequence_in_shift": row.sequence_in_shift, "test_type": row.test_type, "rotation": row.rotation, "config": row.config, "status": row.status, "priority": row.priority, "estimated_minutes": row.estimated_minutes, } for row in rows ], } if __name__ == "__main__": import uvicorn uvicorn.run("app:app", host="0.0.0.0", port=8000, reload=True)