Initial public ModelForge release
This commit is contained in:
@@ -0,0 +1,121 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user