Files

122 lines
4.6 KiB
Python

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()