implemented login

This commit is contained in:
2026-08-03 10:50:58 -04:00
parent 2a2eaa204c
commit 54f0de305f
14 changed files with 976 additions and 126 deletions
+1
View File
@@ -4,3 +4,4 @@ pydantic==2.11.7
watchdog==6.0.0
httpx==0.28.1
smbprotocol==1.15.0
PyJWT==2.10.1
+128 -13
View File
@@ -6,10 +6,11 @@ import io
import os
import zipfile
from fastapi import FastAPI, HTTPException, Response
from fastapi import Depends, FastAPI, HTTPException, Response
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
import auth
import db as db
from parser import CsvValidationError
from file_manager import resolve_requested_csv_path, read_file_bytes, path_name
@@ -30,6 +31,26 @@ DB_PATH = Path(os.getenv("DB_PATH", str(APP_ROOT.parent / "data" / "scheduler.db
DUT = os.getenv("DUT", "CGW453").strip()
REF = os.getenv("REF", "CGW452").strip()
def _require_env_value(name: str) -> str:
value = os.getenv(name, "").strip()
if not value:
raise RuntimeError(f"{name} must be set when bootstrapping default users.")
return value
def _bootstrap_default_users() -> None:
if db.count_users(DB_PATH) > 0:
return
admin_username = _require_env_value("DEFAULT_ADMIN_USERNAME")
admin_password = _require_env_value("DEFAULT_ADMIN_PASSWORD")
viewer_username = _require_env_value("DEFAULT_VIEWER_USERNAME")
viewer_password = _require_env_value("DEFAULT_VIEWER_PASSWORD")
db.create_user(admin_username, admin_password, "admin", DB_PATH)
db.create_user(viewer_username, viewer_password, "viewer", DB_PATH)
class LoadTestsRequest(BaseModel):
csv_path: str | None = Field(default=None, description="Absolute or backend-relative path to target CSV")
csv_paths: list[str] = Field(default_factory=list, description="One or more CSV paths to load together")
@@ -40,6 +61,17 @@ class SaveSettingsRequest(BaseModel):
settings: dict[str, Any]
class LoginRequest(BaseModel):
username: str = Field(..., min_length=1)
password: str = Field(..., min_length=1)
class CreateUserRequest(BaseModel):
username: str = Field(..., min_length=1)
password: str = Field(..., min_length=8)
role: str = Field(default="viewer", min_length=1)
def _smb_credentials_from_settings(settings: dict[str, Any]) -> dict[str, str]:
return {
"username": str(settings.get("smbUsername") or "").strip(),
@@ -83,6 +115,15 @@ def _runtime_overrides_from_settings(settings: dict[str, Any]) -> dict[str, dict
return overrides
def _serialize_user_summary(user: db.UserSummary) -> dict[str, str]:
return {
"username": user.username,
"role": user.role,
"created_at": user.created_at,
"updated_at": user.updated_at,
}
@@ -279,7 +320,9 @@ def _build_schedule_windows(week_start: date, all_rows: list[db.ScheduleRow], ho
@asynccontextmanager
async def lifespan(application: FastAPI):
auth.validate_auth_environment()
db.init_db(DB_PATH)
_bootstrap_default_users()
settings = db.read_settings(DB_PATH)
configure_result_watcher(settings)
try:
@@ -308,20 +351,92 @@ def health() -> dict[str, str]:
return {"status": "ok"}
@app.post("/api/auth/login")
def login(request: LoginRequest) -> dict[str, Any]:
user = db.get_user_by_username(request.username, DB_PATH)
if user is None or not db.verify_password(request.password, user.password_hash):
raise HTTPException(status_code=401, detail="Invalid username or password")
token, expires_at = auth.create_access_token(user.username, user.role)
return {
"access_token": token,
"token_type": "bearer",
"username": user.username,
"role": user.role,
"expires_at": expires_at.isoformat(),
}
@app.get("/api/auth/me")
def get_current_user(current_user: auth.AuthUser = Depends(auth.get_current_user)) -> dict[str, str]:
return {
"username": current_user.username,
"role": current_user.role,
}
@app.get("/api/users")
def list_users_endpoint(_admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
users = db.list_users(DB_PATH)
return {
"users": [_serialize_user_summary(user) for user in users],
}
@app.post("/api/users")
def create_user_endpoint(request: CreateUserRequest, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
try:
created_user = db.create_user(request.username, request.password, request.role, DB_PATH)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return {
"user": {
"username": created_user.username,
"role": created_user.role,
},
}
@app.delete("/api/users/{username}")
def delete_user_endpoint(username: str, current_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
normalized_username = (username or "").strip()
if not normalized_username:
raise HTTPException(status_code=400, detail="username is required")
if normalized_username == current_user.username:
raise HTTPException(status_code=400, detail="You cannot delete your own account")
target_user = db.get_user_by_username(normalized_username, DB_PATH)
if target_user is None:
raise HTTPException(status_code=404, detail="User not found")
if target_user.role == "admin" and db.count_admin_users(DB_PATH) <= 1:
raise HTTPException(status_code=409, detail="At least one admin user must remain")
deleted_user = db.delete_user(normalized_username, DB_PATH)
if deleted_user is None:
raise HTTPException(status_code=404, detail="User not found")
return {
"status": "deleted",
"username": deleted_user.username,
}
@app.post("/api/settings")
def save_settings(request: SaveSettingsRequest) -> dict[str, str]:
def save_settings(request: SaveSettingsRequest, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> 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]:
def get_settings(_admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
return db.read_settings(DB_PATH)
@app.post("/api/settings/restart")
def restart_and_clear_data() -> dict[str, Any]:
def restart_and_clear_data(_admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
deleted_counts = db.reset_all_data(DB_PATH)
configure_result_watcher({})
return {
@@ -331,7 +446,7 @@ def restart_and_clear_data() -> dict[str, Any]:
@app.post("/api/tests/load")
def load_tests(request: LoadTestsRequest) -> dict[str, Any]:
def load_tests(request: LoadTestsRequest, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
requested_paths: list[str] = []
if request.csv_path and request.csv_path.strip():
requested_paths.append(request.csv_path.strip())
@@ -362,14 +477,14 @@ def load_tests(request: LoadTestsRequest) -> dict[str, Any]:
}
@app.post("/api/schedule/active/remove")
def remove_active_tests(request: RemoveActiveTestsRequest) -> dict[str, Any]:
def remove_active_tests(request: RemoveActiveTestsRequest, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
# new_scheduler does not keep global mutable active state in app lifecycle.
# Keep endpoint for compatibility with frontend calls.
return {"status": "ok", "removed": len(request.test_ids)}
@app.post("/api/schedule/compile")
def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any]:
def compile_schedule_endpoint(request: CompileScheduleRequest, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> dict[str, Any]:
if request.start_date:
try:
datetime.strptime(request.start_date, "%Y-%m-%d")
@@ -469,7 +584,7 @@ def compile_schedule_endpoint(request: CompileScheduleRequest) -> dict[str, Any]
@app.get("/api/tests/rerun")
def get_rerun_tests() -> dict[str, Any]:
def get_rerun_tests(_current_user: auth.AuthUser = Depends(auth.get_current_user)) -> 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)
@@ -477,24 +592,24 @@ def get_rerun_tests() -> dict[str, Any]:
@app.post("/api/holidays")
def save_holidays(request: SaveHolidaysRequest) -> dict[str, Any]:
def save_holidays(request: SaveHolidaysRequest, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> 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]:
def get_holidays(_current_user: auth.AuthUser = Depends(auth.get_current_user)) -> dict[str, Any]:
return {"dates": sorted(db.list_holidays(DB_PATH))}
@app.get("/api/schedule/versions")
def get_schedule_versions() -> dict[str, Any]:
def get_schedule_versions(_current_user: auth.AuthUser = Depends(auth.get_current_user)) -> dict[str, Any]:
return {"versions": db.get_schedule_versions(DB_PATH)}
@app.get("/api/schedule/week")
def get_schedule_week(start: str | None = None, version: int | None = None) -> dict[str, Any]:
def get_schedule_week(start: str | None = None, version: int | None = None, _current_user: auth.AuthUser = Depends(auth.get_current_user)) -> dict[str, Any]:
week_start = start or date.today().isoformat()
try:
datetime.strptime(week_start, "%Y-%m-%d")
@@ -536,7 +651,7 @@ def get_schedule_week(start: str | None = None, version: int | None = None) -> d
}
@app.get("/api/schedule/export")
def export_window(window_id: str, version: int | None = None) -> Response:
def export_window(window_id: str, version: int | None = None, _admin_user: auth.AuthUser = Depends(auth.require_admin)) -> Response:
resolved_version = db.resolve_schedule_version(version, DB_PATH)
if version is not None and resolved_version is None:
raise HTTPException(status_code=404, detail=f"Schedule version {version} was not found")
+87
View File
@@ -0,0 +1,87 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
import os
import jwt
from fastapi import Depends, HTTPException
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
@dataclass(frozen=True)
class AuthUser:
username: str
role: str
security = HTTPBearer(auto_error=False)
def _get_secret() -> str:
secret = os.getenv("JWT_SECRET", "").strip()
if not secret:
raise RuntimeError("JWT_SECRET must be set.")
return secret
def _get_algorithm() -> str:
return os.getenv("JWT_ALGORITHM", "HS256").strip() or "HS256"
def _get_expires_hours() -> int:
raw = os.getenv("JWT_EXPIRES_HOURS", "8").strip() or "8"
try:
expires_hours = int(raw)
except ValueError as exc:
raise RuntimeError("JWT_EXPIRES_HOURS must be an integer.") from exc
if expires_hours <= 0:
raise RuntimeError("JWT_EXPIRES_HOURS must be greater than zero.")
return expires_hours
def create_access_token(username: str, role: str) -> tuple[str, datetime]:
now = datetime.now(timezone.utc)
expires_at = now + timedelta(hours=_get_expires_hours())
payload = {
"sub": username,
"role": role,
"iat": int(now.timestamp()),
"exp": int(expires_at.timestamp()),
}
token = jwt.encode(payload, _get_secret(), algorithm=_get_algorithm())
return token, expires_at
def decode_access_token(token: str) -> AuthUser:
try:
payload = jwt.decode(token, _get_secret(), algorithms=[_get_algorithm()])
except jwt.ExpiredSignatureError as exc:
raise HTTPException(status_code=401, detail="Token has expired") from exc
except jwt.InvalidTokenError as exc:
raise HTTPException(status_code=401, detail="Invalid token") from exc
username = str(payload.get("sub") or "").strip()
role = str(payload.get("role") or "").strip().lower()
if not username or role not in {"admin", "viewer"}:
raise HTTPException(status_code=401, detail="Invalid token payload")
return AuthUser(username=username, role=role)
def validate_auth_environment() -> None:
_get_secret()
_get_algorithm()
_get_expires_hours()
def get_current_user(credentials: HTTPAuthorizationCredentials | None = Depends(security)) -> AuthUser:
if credentials is None or credentials.scheme.lower() != "bearer":
raise HTTPException(status_code=401, detail="Authorization token is required")
return decode_access_token(credentials.credentials)
def require_admin(user: AuthUser = Depends(get_current_user)) -> AuthUser:
if user.role != "admin":
raise HTTPException(status_code=403, detail="Admin access required")
return user
+177
View File
@@ -2,6 +2,9 @@ import json
import sqlite3
import os
import re
import base64
import hashlib
import hmac
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
@@ -18,6 +21,7 @@ REF = os.getenv("REF", "CGW452").strip()
# Device names for the test database (use hardware device names)
MAX_SCHEDULE_VERSIONS = 50
PBKDF2_ITERATIONS = 390000
@dataclass(frozen=True)
class TestRecord:
@@ -53,6 +57,21 @@ class ScheduleRow:
status: str
estimated_minutes: int
@dataclass(frozen=True)
class UserAccount:
username: str
password_hash: str
role: str
@dataclass(frozen=True)
class UserSummary:
username: str
role: str
created_at: str
updated_at: str
@contextmanager
def get_connection(db_path: str | Path = DB_PATH) -> Iterator[sqlite3.Connection]:
conn = sqlite3.connect(str(db_path))
@@ -122,8 +141,18 @@ def init_db(db_path: str | Path = DB_PATH) -> None:
minutes INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT NOT NULL UNIQUE,
password_hash TEXT NOT NULL,
role TEXT NOT NULL CHECK (role IN ('admin', 'viewer')),
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_tests_status ON tests(status);
CREATE INDEX IF NOT EXISTS idx_schedules_date_shift ON schedules(scheduled_date, shift_index);
CREATE INDEX IF NOT EXISTS idx_users_username ON users(username);
"""
)
@@ -162,6 +191,154 @@ def init_db(db_path: str | Path = DB_PATH) -> None:
"""
)
def count_users(db_path: str | Path = DB_PATH) -> int:
with get_connection(db_path) as conn:
row = conn.execute("SELECT COUNT(*) AS count FROM users").fetchone()
return int(row["count"])
def hash_password(password: str) -> str:
if not password:
raise ValueError("password is required")
salt = os.urandom(16)
digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, PBKDF2_ITERATIONS)
salt_b64 = base64.b64encode(salt).decode("ascii")
digest_b64 = base64.b64encode(digest).decode("ascii")
return f"pbkdf2_sha256${PBKDF2_ITERATIONS}${salt_b64}${digest_b64}"
def verify_password(password: str, password_hash: str) -> bool:
if not password or not password_hash:
return False
parts = password_hash.split("$")
if len(parts) != 4:
return False
algorithm, iterations_raw, salt_b64, expected_digest_b64 = parts
if algorithm != "pbkdf2_sha256":
return False
try:
iterations = int(iterations_raw)
salt = base64.b64decode(salt_b64.encode("ascii"))
expected_digest = base64.b64decode(expected_digest_b64.encode("ascii"))
except (TypeError, ValueError):
return False
computed_digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, iterations)
return hmac.compare_digest(computed_digest, expected_digest)
def get_user_by_username(username: str, db_path: str | Path = DB_PATH) -> UserAccount | None:
normalized_username = (username or "").strip()
if not normalized_username:
return None
with get_connection(db_path) as conn:
row = conn.execute(
"SELECT username, password_hash, role FROM users WHERE username = ?",
(normalized_username,),
).fetchone()
if row is None:
return None
return UserAccount(
username=row["username"],
password_hash=row["password_hash"],
role=row["role"],
)
def create_user(username: str, password: str, role: str, db_path: str | Path = DB_PATH) -> UserAccount:
normalized_username = (username or "").strip()
normalized_role = (role or "").strip().lower()
if not normalized_username:
raise ValueError("username is required")
if not password:
raise ValueError("password is required")
if len(password) < 8:
raise ValueError("password must be at least 8 characters")
if normalized_role not in {"admin", "viewer"}:
raise ValueError("role must be 'admin' or 'viewer'")
password_hash = hash_password(password)
try:
with get_connection(db_path) as conn:
conn.execute(
"""
INSERT INTO users(username, password_hash, role, updated_at)
VALUES (?, ?, ?, CURRENT_TIMESTAMP)
""",
(normalized_username, password_hash, normalized_role),
)
except sqlite3.IntegrityError as exc:
raise ValueError("username already exists") from exc
return UserAccount(
username=normalized_username,
password_hash=password_hash,
role=normalized_role,
)
def list_users(db_path: str | Path = DB_PATH) -> list[UserSummary]:
with get_connection(db_path) as conn:
rows = conn.execute(
"""
SELECT username, role, created_at, updated_at
FROM users
ORDER BY username COLLATE NOCASE ASC
"""
).fetchall()
return [
UserSummary(
username=row["username"],
role=row["role"],
created_at=row["created_at"],
updated_at=row["updated_at"],
)
for row in rows
]
def count_admin_users(db_path: str | Path = DB_PATH) -> int:
with get_connection(db_path) as conn:
row = conn.execute(
"SELECT COUNT(*) AS count FROM users WHERE role = 'admin'"
).fetchone()
return int(row["count"])
def delete_user(username: str, db_path: str | Path = DB_PATH) -> UserAccount | None:
normalized_username = (username or "").strip()
if not normalized_username:
return None
with get_connection(db_path) as conn:
row = conn.execute(
"SELECT username, password_hash, role FROM users WHERE username = ?",
(normalized_username,),
).fetchone()
if row is None:
return None
conn.execute(
"DELETE FROM users WHERE username = ?",
(normalized_username,),
)
return UserAccount(
username=row["username"],
password_hash=row["password_hash"],
role=row["role"],
)
def upsert_tests(records: list[TestRecord], db_path: str | Path = DB_PATH) -> int:
if not records:
return 0