from fastapi.testclient import TestClient from hardware_fakes import FakeAcceleratorCollector, FakeHostCollector, accelerator from pydantic import SecretStr from sqlalchemy import create_engine from sqlalchemy.orm import Session from sqlalchemy.pool import StaticPool from modelforge_api.api.routes.hardware import get_hardware_service from modelforge_api.main import app from modelforge_api.persistence.models import Base from modelforge_api.services.hardware_inventory import HardwareInventoryService, HardwareRefreshBusy from modelforge_api.settings import Settings, get_settings def _test_operator_credential() -> str: return "hardware-test-operator" def _operator_headers() -> dict[str, str]: return {"X-ModelForge-Admin-Token": _test_operator_credential()} def _configure_operator_auth() -> None: settings = Settings( _env_file=None, operator_api_key=SecretStr(_test_operator_credential()), ) app.dependency_overrides[get_settings] = lambda: settings def test_hardware_refresh_and_resource_endpoints_serialize_real_and_unknown_values() -> None: engine = create_engine( "sqlite+pysqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) Base.metadata.create_all(engine) session = Session(engine) service = HardwareInventoryService( session, FakeHostCollector(), FakeAcceleratorCollector([accelerator()]) ) _configure_operator_auth() app.dependency_overrides[get_hardware_service] = lambda: service try: with TestClient(app) as client: refreshed = client.post("/api/v1/hardware/refresh", headers=_operator_headers()) assert refreshed.status_code == 200 payload = refreshed.json() assert payload["overview"]["accelerator_count"] == 1 assert ( payload["nodes"][0]["accelerators"][0]["mig_mode_current"]["availability"] == "unsupported" ) node_id = payload["nodes"][0]["id"] accelerator_id = payload["nodes"][0]["accelerators"][0]["id"] assert client.get("/api/v1/hardware", headers=_operator_headers()).status_code == 200 assert ( client.get("/api/v1/hardware/nodes", headers=_operator_headers()).json()[0]["id"] == node_id ) assert ( client.get( f"/api/v1/hardware/nodes/{node_id}", headers=_operator_headers() ).status_code == 200 ) assert ( client.get("/api/v1/hardware/accelerators", headers=_operator_headers()).json()[0][ "id" ] == accelerator_id ) assert ( client.get( f"/api/v1/hardware/accelerators/{accelerator_id}", headers=_operator_headers(), ).status_code == 200 ) missing = client.get( "/api/v1/hardware/nodes/00000000-0000-0000-0000-000000000000", headers=_operator_headers(), ) assert missing.status_code == 404 assert missing.json()["error"]["code"] == "http_404" finally: app.dependency_overrides.clear() session.close() def test_concurrent_refresh_returns_normalized_conflict() -> None: class BusyService: def refresh(self): raise HardwareRefreshBusy("hardware refresh already in progress") _configure_operator_auth() app.dependency_overrides[get_hardware_service] = lambda: BusyService() try: with TestClient(app) as client: response = client.post("/api/v1/hardware/refresh", headers=_operator_headers()) assert response.status_code == 409 assert response.json()["error"]["code"] == "http_409" finally: app.dependency_overrides.clear() def test_refresh_failure_returns_normalized_service_unavailable() -> None: class FailingService: def refresh(self): raise RuntimeError("probe failed") _configure_operator_auth() app.dependency_overrides[get_hardware_service] = lambda: FailingService() try: with TestClient(app) as client: response = client.post("/api/v1/hardware/refresh", headers=_operator_headers()) assert response.status_code == 503 assert response.json()["error"]["code"] == "http_503" assert response.json()["error"]["message"] == "hardware inventory failed: RuntimeError" finally: app.dependency_overrides.clear()