implemented login
This commit is contained in:
@@ -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
@@ -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")
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user