Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 24 additions & 5 deletions src/env_doctor/server/database.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
"""SQLAlchemy async database engine and session management."""
from pathlib import Path

from sqlalchemy import event
from sqlalchemy import event, text
from sqlalchemy.exc import OperationalError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import DeclarativeBase

Expand All @@ -23,16 +24,34 @@ class Base(DeclarativeBase):


async def init_db():
"""Create all tables and enable WAL mode."""
"""Create all tables, run lightweight migrations, and enable WAL mode."""
_DB_PATH.parent.mkdir(exist_ok=True)

from . import models # noqa: F401 — ensure models are registered

async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
await conn.execute(
__import__("sqlalchemy").text("PRAGMA journal_mode=WAL")
)
await conn.execute(text("PRAGMA journal_mode=WAL"))
await _run_lightweight_migrations(conn)


async def _run_lightweight_migrations(conn):
"""Idempotent ALTER statements for columns added after the initial schema.

SQLite's ``CREATE TABLE`` is no-op when a table exists, so new columns on
existing tables need explicit ``ALTER TABLE``. We run each statement and
swallow ``OperationalError`` (raised when the column/index already exists).
"""
statements = [
"ALTER TABLE machines ADD COLUMN group_name VARCHAR(64)",
"CREATE INDEX IF NOT EXISTS ix_machines_group_name ON machines(group_name)",
]
for stmt in statements:
try:
await conn.execute(text(stmt))
except OperationalError:
# Column or index already exists — expected on every run after the first.
pass


async def get_session():
Expand Down
2 changes: 2 additions & 0 deletions src/env_doctor/server/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ class Machine(Base):
python_version = Column(String, nullable=True)
latest_status = Column(String, nullable=True) # "pass"/"warning"/"fail"
latest_snapshot_id = Column(Integer, nullable=True)
group_name = Column(String(64), nullable=True, index=True)
first_seen = Column(DateTime, default=lambda: datetime.now(timezone.utc))
last_seen = Column(DateTime, default=lambda: datetime.now(timezone.utc))

Expand All @@ -60,6 +61,7 @@ def to_dict(self):
"platform": self.platform,
"python_version": self.python_version,
"latest_status": self.latest_status,
"group_name": self.group_name,
"first_seen": self.first_seen.isoformat() if self.first_seen else None,
"last_seen": self.last_seen.isoformat() if self.last_seen else None,
}
Expand Down
142 changes: 140 additions & 2 deletions src/env_doctor/server/routes.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
"""API route handlers for the dashboard."""
import json
import os
import re
from datetime import datetime, timezone
from typing import Optional

from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from sqlalchemy import select
from pydantic import BaseModel, Field
from sqlalchemy import case, func, select
from sqlalchemy.ext.asyncio import AsyncSession

from .database import get_session
Expand All @@ -29,6 +30,40 @@ def _seconds_since(when: Optional[datetime]) -> Optional[float]:
return (datetime.now(timezone.utc) - when).total_seconds()


# ---------------------------------------------------------------------------
# Group name validation
# ---------------------------------------------------------------------------

# Groups must be human-readable identifiers safe for URLs, file paths, and
# topology labels. Allow alphanumerics, hyphen, underscore, dot, and space.
_GROUP_NAME_PATTERN = re.compile(r"^[\w\-. ]+$")
_UNGROUPED_LABEL = "ungrouped"


def _clean_group_name(raw: Optional[str]) -> Optional[str]:
"""Normalise a group name. Returns None for empty/whitespace input.

Raises HTTPException(400) on invalid characters so callers do not have to
validate separately.
"""
if raw is None:
return None
cleaned = raw.strip()
if not cleaned:
return None
if len(cleaned) > 64:
raise HTTPException(status_code=400, detail="group_name must be 64 characters or fewer")
if not _GROUP_NAME_PATTERN.match(cleaned):
raise HTTPException(
status_code=400,
detail="group_name may only contain letters, numbers, spaces, '-', '_', '.'",
)
if cleaned.lower() == _UNGROUPED_LABEL:
# Reserved synthetic label used by GET /api/groups for NULL machines.
raise HTTPException(status_code=400, detail=f"'{_UNGROUPED_LABEL}' is a reserved name")
return cleaned


# ---------------------------------------------------------------------------
# Pydantic request/response models
# ---------------------------------------------------------------------------
Expand All @@ -40,6 +75,7 @@ class MachineInfo(BaseModel):
platform_release: Optional[str] = None
python_version: Optional[str] = None
reported_at: Optional[str] = None
group_name: Optional[str] = None # Optional self-tag from CLI; dashboard PATCH overrides.


class ReportPayload(BaseModel):
Expand Down Expand Up @@ -111,6 +147,7 @@ async def receive_report(
first_seen=now,
last_seen=now,
latest_status=payload.status,
group_name=_clean_group_name(payload.machine.group_name),
)
session.add(machine)
else:
Expand All @@ -119,6 +156,11 @@ async def receive_report(
machine.python_version = payload.machine.python_version
machine.last_seen = now
machine.latest_status = payload.status
# Only honour CLI self-tag when no group was set via the dashboard PATCH.
# Dashboard-assigned groups are the source of truth and must not be
# silently overwritten by every check-in.
if machine.group_name is None and payload.machine.group_name:
machine.group_name = _clean_group_name(payload.machine.group_name)

# Create snapshot
fields = _extract_fields(payload)
Expand Down Expand Up @@ -228,6 +270,102 @@ async def get_machine(
return result


# ---------------------------------------------------------------------------
# PATCH /api/machines/{id} (mutate machine metadata — currently group_name)
# ---------------------------------------------------------------------------

class MachineUpdate(BaseModel):
# Use Field(...) sentinel so we can distinguish "set to null" from "field omitted".
group_name: Optional[str] = Field(default=None, max_length=64)


@router.patch("/machines/{machine_id}")
async def update_machine(
machine_id: str,
body: MachineUpdate,
session: AsyncSession = Depends(get_session),
):
"""Update mutable machine fields. Currently supports group_name.

Pass ``{"group_name": "training-east"}`` to assign, or ``{"group_name": ""}``
/ ``{"group_name": null}`` to ungroup.
"""
machine = await session.get(Machine, machine_id)
if not machine:
raise HTTPException(status_code=404, detail="Machine not found")

# _clean_group_name raises 400 on invalid characters.
machine.group_name = _clean_group_name(body.group_name)
await session.commit()
await session.refresh(machine)

result = machine.to_dict()
elapsed = _seconds_since(machine.last_seen)
result["last_seen_seconds"] = elapsed
result["stale"] = elapsed is not None and elapsed > _STALE_AFTER_SECONDS
if machine.latest_snapshot_id:
snap = await session.get(Snapshot, machine.latest_snapshot_id)
if snap:
# Match GET /api/machines/{id} shape so callers (MachineDetail) can
# safely setMachine(patchResponse) without losing diagnostics.
result["latest_report"] = json.loads(snap.report_json)
result["gpu_name"] = snap.gpu_name
result["driver_version"] = snap.driver_version
result["cuda_version"] = snap.cuda_version
result["torch_version"] = snap.torch_version
return result


# ---------------------------------------------------------------------------
# GET /api/groups
# ---------------------------------------------------------------------------

@router.get("/groups")
async def list_groups(session: AsyncSession = Depends(get_session)):
"""Return all distinct group names with member counts and status breakdown.

NULL group_name values are aggregated under the synthetic "ungrouped" entry
and listed last; named groups are sorted alphabetically.
"""
pass_sum = func.sum(case((Machine.latest_status == "pass", 1), else_=0))
warn_sum = func.sum(case((Machine.latest_status == "warning", 1), else_=0))
fail_sum = func.sum(case((Machine.latest_status == "fail", 1), else_=0))

query = (
select(
Machine.group_name,
func.count(Machine.id),
pass_sum,
warn_sum,
fail_sum,
)
.group_by(Machine.group_name)
)
result = await session.execute(query)

named: list[dict] = []
ungrouped: Optional[dict] = None
for group_name, count, n_pass, n_warn, n_fail in result.all():
entry = {
"name": group_name if group_name else _UNGROUPED_LABEL,
"machine_count": int(count or 0),
"status_breakdown": {
"pass": int(n_pass or 0),
"warning": int(n_warn or 0),
"fail": int(n_fail or 0),
},
}
if group_name:
named.append(entry)
else:
ungrouped = entry

named.sort(key=lambda g: g["name"].lower())
if ungrouped is not None:
named.append(ungrouped)
return named


# ---------------------------------------------------------------------------
# GET /api/machines/{id}/history
# ---------------------------------------------------------------------------
Expand Down
27 changes: 27 additions & 0 deletions web/src/api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ import type {
CommandActivityRow,
CommandRecord,
MachineDetail,
MachineGroup,
MachineListItem,
SnapshotSummary,
} from "./types";
Expand Down Expand Up @@ -104,6 +105,32 @@ export function getCommands(machineId: string): Promise<CommandRecord[]> {
return fetchJson(`${BASE}/machines/${machineId}/commands`);
}

export function getGroups(): Promise<MachineGroup[]> {
return fetchJson(`${BASE}/groups`);
}

export async function updateMachineGroup(
id: string,
group_name: string | null
): Promise<MachineDetail> {
const res = await apiFetch(`${BASE}/machines/${id}`, {
method: "PATCH",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({ group_name }),
});
if (!res.ok) {
let detail = `HTTP ${res.status}`;
try {
const body = await res.json();
if (body && typeof body.detail === "string") detail = body.detail;
} catch {
/* ignore */
}
throw new Error(detail);
}
return res.json();
}

export function getCommandActivity(
filters: CommandActivityFilters = {}
): Promise<CommandActivityRow[]> {
Expand Down
Loading
Loading