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 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] 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() yield 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) 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 try: parsed = parse_target_csv(csv_path) 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_graph_once(all_dut_tests) 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 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 ] 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, ) return { "schedule_version": version, "scheduled_tests": len(entries), "completion_date": completion_date, } @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)