122 lines
4.6 KiB
Python
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()
|