460 lines
14 KiB
Python
460 lines
14 KiB
Python
from __future__ import annotations
|
|
|
|
import time
|
|
from datetime import timedelta
|
|
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
|
|
import cloud_api.app as app_module
|
|
from cloud_api.app import create_app
|
|
from cloud.control_config import CloudConfigurationError, CloudControlConfig
|
|
from cloud.database import CloudDatabase
|
|
from cloud.pool import PooledDevice
|
|
from cloud.scheduler import ScheduledTask, TaskConstraints
|
|
from core.models import utc_now
|
|
|
|
|
|
def test_create_app_returns_independent_cloud_application() -> None:
|
|
app = create_app()
|
|
|
|
assert app.title == "Device Cloud API"
|
|
assert callable(create_app)
|
|
paths = set(app.openapi()["paths"])
|
|
assert "/v1/tasks" in paths
|
|
assert "/internal/v1/hosts/{host_id}/heartbeat" in paths
|
|
|
|
|
|
def test_cloud_application_owns_database_lifecycle() -> None:
|
|
events: list[str] = []
|
|
|
|
class FakeRepository:
|
|
def health_check(self) -> None:
|
|
events.append("healthy")
|
|
|
|
def list_queued_tasks(self) -> list[object]:
|
|
events.append("scheduled")
|
|
return []
|
|
|
|
def reap_expired_leases(self, *, now, max_attempts: int) -> list[str]:
|
|
events.append("reaped")
|
|
return []
|
|
|
|
class FakeDatabase:
|
|
repository = FakeRepository()
|
|
|
|
def close(self) -> None:
|
|
events.append("closed")
|
|
|
|
fake_database = FakeDatabase()
|
|
app = create_app(
|
|
config=CloudControlConfig(database_url="sqlite:///:memory:"),
|
|
database_factory=lambda _config: fake_database, # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with TestClient(app):
|
|
assert app.state.database is fake_database
|
|
assert (
|
|
app.state.cloud_services.pool.store is app.state.cloud_services.repository
|
|
)
|
|
assert "healthy" in events
|
|
|
|
assert events[-1] == "closed"
|
|
|
|
|
|
def test_repository_proxy_is_available_only_during_lifespan() -> None:
|
|
class FakeRepository:
|
|
def health_check(self) -> None:
|
|
return None
|
|
|
|
def list_queued_tasks(self) -> list[object]:
|
|
return []
|
|
|
|
def reap_expired_leases(self, *, now, max_attempts: int) -> list[str]:
|
|
return []
|
|
|
|
class FakeDatabase:
|
|
repository = FakeRepository()
|
|
|
|
def close(self) -> None:
|
|
return None
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(database_url="sqlite:///:memory:"),
|
|
database_factory=lambda _config: FakeDatabase(), # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with TestClient(app):
|
|
app.state.cloud_services.repository.health_check()
|
|
|
|
with pytest.raises(RuntimeError, match="outside app lifespan"):
|
|
app.state.cloud_services.repository.health_check()
|
|
|
|
|
|
def test_lifespan_runs_scheduler_and_reaper_until_shutdown() -> None:
|
|
events: list[str] = []
|
|
|
|
class FakeRepository:
|
|
def health_check(self) -> None:
|
|
events.append("healthy")
|
|
|
|
def list_queued_tasks(self) -> list[object]:
|
|
events.append("scheduled")
|
|
return []
|
|
|
|
def reap_expired_leases(self, *, now, max_attempts: int) -> list[str]:
|
|
events.append(f"reaped:{max_attempts}")
|
|
return []
|
|
|
|
class FakeDatabase:
|
|
repository = FakeRepository()
|
|
|
|
def close(self) -> None:
|
|
events.append("closed")
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(
|
|
database_url="sqlite:///:memory:",
|
|
scheduler_interval_seconds=0.01,
|
|
lease_reaper_interval_seconds=0.01,
|
|
max_task_attempts=4,
|
|
),
|
|
database_factory=lambda _config: FakeDatabase(), # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with TestClient(app):
|
|
time.sleep(0.04)
|
|
assert events.count("scheduled") >= 2
|
|
assert events.count("reaped:4") >= 2
|
|
assert all(not task.done() for task in app.state.worker_tasks)
|
|
|
|
assert all(task.done() for task in app.state.worker_tasks)
|
|
assert events[-1] == "closed"
|
|
|
|
|
|
def test_lifecycle_workers_log_failures_and_continue(monkeypatch) -> None:
|
|
logged_workers: list[str] = []
|
|
|
|
def record_exception(_message: str, *, extra: dict[str, str]) -> None:
|
|
logged_workers.append(extra["worker"])
|
|
|
|
monkeypatch.setattr(app_module.logger, "exception", record_exception)
|
|
scheduler_calls = 0
|
|
reaper_calls = 0
|
|
|
|
class FakeRepository:
|
|
def health_check(self) -> None:
|
|
return None
|
|
|
|
def list_queued_tasks(self) -> list[object]:
|
|
nonlocal scheduler_calls
|
|
scheduler_calls += 1
|
|
if scheduler_calls == 1:
|
|
raise RuntimeError("scheduler transient failure")
|
|
return []
|
|
|
|
def reap_expired_leases(self, *, now, max_attempts: int) -> list[str]:
|
|
nonlocal reaper_calls
|
|
reaper_calls += 1
|
|
if reaper_calls == 1:
|
|
raise RuntimeError("reaper transient failure")
|
|
return []
|
|
|
|
class FakeDatabase:
|
|
repository = FakeRepository()
|
|
|
|
def close(self) -> None:
|
|
return None
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(
|
|
database_url="sqlite:///:memory:",
|
|
scheduler_interval_seconds=0.01,
|
|
lease_reaper_interval_seconds=0.01,
|
|
),
|
|
database_factory=lambda _config: FakeDatabase(), # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with TestClient(app):
|
|
time.sleep(0.04)
|
|
assert scheduler_calls >= 2
|
|
assert reaper_calls >= 2
|
|
assert all(not task.done() for task in app.state.worker_tasks)
|
|
|
|
assert set(logged_workers) >= {
|
|
"scheduler",
|
|
"lease_reaper",
|
|
}
|
|
|
|
|
|
def test_liveness_stays_up_when_database_readiness_fails() -> None:
|
|
database_healthy = True
|
|
|
|
class FakeRepository:
|
|
def health_check(self) -> None:
|
|
if not database_healthy:
|
|
raise RuntimeError("database unavailable with secret credentials")
|
|
|
|
def list_queued_tasks(self) -> list[object]:
|
|
return []
|
|
|
|
def reap_expired_leases(self, *, now, max_attempts: int) -> list[str]:
|
|
return []
|
|
|
|
class FakeDatabase:
|
|
repository = FakeRepository()
|
|
|
|
def close(self) -> None:
|
|
return None
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(database_url="sqlite:///:memory:"),
|
|
database_factory=lambda _config: FakeDatabase(), # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with TestClient(app) as client:
|
|
assert client.get("/health/live").json() == {"status": "live"}
|
|
ready = client.get("/health/ready")
|
|
assert ready.status_code == 200
|
|
assert all(ready.json()["checks"].values())
|
|
|
|
database_healthy = False
|
|
assert client.get("/health/live").status_code == 200
|
|
not_ready = client.get("/health/ready")
|
|
assert not_ready.status_code == 503
|
|
assert not_ready.json()["checks"]["database"] is False
|
|
assert "secret credentials" not in not_ready.text
|
|
|
|
|
|
def test_readiness_reports_stopped_worker() -> None:
|
|
class FakeRepository:
|
|
def health_check(self) -> None:
|
|
return None
|
|
|
|
def list_queued_tasks(self) -> list[object]:
|
|
return []
|
|
|
|
def reap_expired_leases(self, *, now, max_attempts: int) -> list[str]:
|
|
return []
|
|
|
|
class FakeDatabase:
|
|
repository = FakeRepository()
|
|
|
|
def close(self) -> None:
|
|
return None
|
|
|
|
class StoppedWorker:
|
|
def done(self) -> bool:
|
|
return True
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(database_url="sqlite:///:memory:"),
|
|
database_factory=lambda _config: FakeDatabase(), # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with TestClient(app) as client:
|
|
app.state.worker_tasks = (StoppedWorker(),)
|
|
response = client.get("/health/ready")
|
|
|
|
assert response.status_code == 503
|
|
assert response.json()["checks"]["workers"] is False
|
|
|
|
|
|
def test_request_correlation_id_is_propagated_without_sensitive_headers(
|
|
monkeypatch,
|
|
tmp_path,
|
|
) -> None:
|
|
request_logs: list[dict[str, object]] = []
|
|
|
|
def record_info(_message: str, *, extra: dict[str, object]) -> None:
|
|
request_logs.append(extra)
|
|
|
|
monkeypatch.setattr(app_module.logger, "info", record_info)
|
|
app = create_app(
|
|
config=CloudControlConfig(
|
|
database_url=f"sqlite:///{(tmp_path / 'correlation.sqlite3').as_posix()}"
|
|
)
|
|
)
|
|
|
|
with TestClient(app) as client:
|
|
response = client.get(
|
|
"/health/live",
|
|
headers={
|
|
"X-Correlation-ID": "request-123",
|
|
"Authorization": "Bearer never-log-this",
|
|
},
|
|
)
|
|
|
|
assert response.headers["X-Correlation-ID"] == "request-123"
|
|
assert request_logs == [
|
|
{
|
|
"correlation_id": "request-123",
|
|
"method": "GET",
|
|
"path": "/health/live",
|
|
"status_code": 200,
|
|
}
|
|
]
|
|
assert "never-log-this" not in repr(request_logs)
|
|
|
|
|
|
def test_startup_database_failure_closes_database() -> None:
|
|
closed = False
|
|
|
|
class FailingRepository:
|
|
def health_check(self) -> None:
|
|
raise RuntimeError("database unavailable")
|
|
|
|
class FailingDatabase:
|
|
repository = FailingRepository()
|
|
|
|
def close(self) -> None:
|
|
nonlocal closed
|
|
closed = True
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(database_url="sqlite:///:memory:"),
|
|
database_factory=lambda _config: FailingDatabase(), # type: ignore[arg-type,return-value]
|
|
)
|
|
|
|
with pytest.raises(RuntimeError, match="database unavailable"):
|
|
with TestClient(app):
|
|
pass
|
|
assert closed is True
|
|
|
|
|
|
def test_persisted_queue_is_recovered_after_control_plane_restart(tmp_path) -> None:
|
|
database_url = f"sqlite:///{(tmp_path / 'queue-recovery.sqlite3').as_posix()}"
|
|
seed = CloudDatabase(database_url)
|
|
now = utc_now()
|
|
seed.repository.upsert_host("host-a", address=None, last_seen_at=now)
|
|
seed.repository.replace_host_devices(
|
|
"host-a",
|
|
[
|
|
PooledDevice(
|
|
device_id="device-a",
|
|
host_id="host-a",
|
|
driver_type="wda",
|
|
status="idle",
|
|
synced_at=now,
|
|
)
|
|
],
|
|
)
|
|
seed.repository.enqueue_task(
|
|
ScheduledTask(
|
|
id="persisted-task",
|
|
goal="resume queued work",
|
|
workflow_definition_id=None,
|
|
constraints=TaskConstraints(),
|
|
created_at=now,
|
|
)
|
|
)
|
|
seed.close()
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(
|
|
database_url=database_url,
|
|
scheduler_interval_seconds=0.01,
|
|
lease_reaper_interval_seconds=1,
|
|
)
|
|
)
|
|
with TestClient(app):
|
|
assert _wait_until(
|
|
lambda: (
|
|
app.state.cloud_services.repository.get_task("persisted-task").status
|
|
== "assigned"
|
|
)
|
|
)
|
|
|
|
verification = CloudDatabase(database_url)
|
|
try:
|
|
task = verification.repository.get_task("persisted-task")
|
|
assert task is not None
|
|
assert task.status == "assigned"
|
|
assert task.attempt_count == 1
|
|
finally:
|
|
verification.close()
|
|
|
|
|
|
def test_expired_lease_is_recovered_after_control_plane_restart(tmp_path) -> None:
|
|
database_url = f"sqlite:///{(tmp_path / 'lease-recovery.sqlite3').as_posix()}"
|
|
seed = CloudDatabase(database_url)
|
|
now = utc_now()
|
|
seed.repository.upsert_host("host-a", address=None, last_seen_at=now)
|
|
seed.repository.replace_host_devices(
|
|
"host-a",
|
|
[
|
|
PooledDevice(
|
|
device_id="device-a",
|
|
host_id="host-a",
|
|
driver_type="wda",
|
|
status="idle",
|
|
synced_at=now,
|
|
)
|
|
],
|
|
)
|
|
seed.repository.enqueue_task(
|
|
ScheduledTask(
|
|
id="expired-task",
|
|
goal="recover expired work",
|
|
workflow_definition_id=None,
|
|
constraints=TaskConstraints(),
|
|
created_at=now,
|
|
)
|
|
)
|
|
seed.repository.assign_task(
|
|
task_id="expired-task",
|
|
host_id="host-a",
|
|
device_id="device-a",
|
|
lease_id="expired-lease",
|
|
lease_expires_at=now - timedelta(seconds=1),
|
|
now=now - timedelta(minutes=1),
|
|
)
|
|
seed.close()
|
|
|
|
app = create_app(
|
|
config=CloudControlConfig(
|
|
database_url=database_url,
|
|
scheduler_interval_seconds=1,
|
|
lease_reaper_interval_seconds=0.01,
|
|
max_task_attempts=1,
|
|
)
|
|
)
|
|
with TestClient(app):
|
|
assert _wait_until(
|
|
lambda: (
|
|
app.state.cloud_services.repository.get_task("expired-task").status
|
|
== "failed"
|
|
)
|
|
)
|
|
|
|
verification = CloudDatabase(database_url)
|
|
try:
|
|
task = verification.repository.get_task("expired-task")
|
|
assert task is not None
|
|
assert task.status == "failed"
|
|
assert task.failure_reason == "lease expired after 1 attempts"
|
|
assert verification.repository.list_task_attempts("expired-task")[0].status == (
|
|
"expired"
|
|
)
|
|
finally:
|
|
verification.close()
|
|
|
|
|
|
def _wait_until(predicate, timeout_seconds: float = 1.0) -> bool:
|
|
deadline = time.monotonic() + timeout_seconds
|
|
while time.monotonic() < deadline:
|
|
if predicate():
|
|
return True
|
|
time.sleep(0.01)
|
|
return False
|
|
|
|
|
|
def test_production_app_rejects_missing_credentials() -> None:
|
|
with pytest.raises(CloudConfigurationError, match="credential"):
|
|
create_app(
|
|
config=CloudControlConfig(
|
|
environment="production",
|
|
database_url="postgresql://db/cloud",
|
|
)
|
|
)
|