from __future__ import annotations import json import pytest from sqlalchemy import create_engine, inspect, text from cloud.schema import ( HEAD_REVISION, SchemaVersionError, current_revision, downgrade_database, require_current_schema, upgrade_database, ) def _database_url(tmp_path) -> str: return f"sqlite:///{(tmp_path / 'cloud.sqlite3').as_posix()}" def test_forward_and_downgrade_migrations_on_empty_database(tmp_path) -> None: database_url = _database_url(tmp_path) upgrade_database(database_url) engine = create_engine(database_url) try: table_names = set(inspect(engine).get_table_names()) assert { "host_registrations", "pooled_devices", "scheduled_tasks", "plugins", "task_attempts", } <= table_names assert current_revision(database_url) == HEAD_REVISION finally: engine.dispose() downgrade_database(database_url) engine = create_engine(database_url) try: inspector = inspect(engine) assert "task_attempts" not in inspector.get_table_names() task_columns = { column["name"] for column in inspector.get_columns("scheduled_tasks") } assert "lease_id" not in task_columns assert current_revision(database_url) is None finally: engine.dispose() def test_legacy_data_survives_upgrade_and_downgrade(tmp_path) -> None: database_url = _database_url(tmp_path) engine = create_engine(database_url) try: with engine.begin() as connection: _create_legacy_schema(connection) connection.execute( text( "insert into host_registrations " "(host_id, address, last_seen_at) values " "('host-a', 'local', '2026-01-01T00:00:00+00:00')" ) ) connection.execute( text( "insert into scheduled_tasks " "(id, goal, workflow_definition_id, constraints_json, status, " "assigned_device_id, assigned_host_id, created_at) values " "('task-a', 'goal', null, :constraints, 'queued', null, null, " "'2026-01-01T00:00:00+00:00')" ), {"constraints": json.dumps({"driver_type": None, "capability_tags": []})}, ) finally: engine.dispose() upgrade_database(database_url) engine = create_engine(database_url) try: with engine.connect() as connection: assert connection.scalar(text("select count(*) from host_registrations")) == 1 assert connection.scalar(text("select count(*) from scheduled_tasks")) == 1 task_columns = { column["name"] for column in inspect(engine).get_columns("scheduled_tasks") } assert {"attempt_count", "lease_id", "result_json"} <= task_columns finally: engine.dispose() downgrade_database(database_url) engine = create_engine(database_url) try: with engine.connect() as connection: assert connection.scalar(text("select count(*) from host_registrations")) == 1 assert connection.scalar(text("select count(*) from scheduled_tasks")) == 1 task_columns = { column["name"] for column in inspect(engine).get_columns("scheduled_tasks") } assert "attempt_count" not in task_columns finally: engine.dispose() def test_schema_readiness_requires_head_revision(tmp_path) -> None: database_url = _database_url(tmp_path) with pytest.raises(SchemaVersionError, match="unversioned"): require_current_schema(database_url) upgrade_database(database_url) require_current_schema(database_url) def _create_legacy_schema(connection) -> None: connection.exec_driver_sql( "create table host_registrations (" "host_id text primary key, address text, last_seen_at text not null)" ) connection.exec_driver_sql( "create table pooled_devices (" "device_id text not null, host_id text not null, driver_type text not null, " "status text not null, capability_tags_json text not null, synced_at text, " "primary key (host_id, device_id))" ) connection.exec_driver_sql( "create table scheduled_tasks (" "id text primary key, goal text, workflow_definition_id text, " "constraints_json text not null, status text not null, " "assigned_device_id text, assigned_host_id text, created_at text not null)" ) connection.exec_driver_sql( "create table plugins (" "name text primary key, version text not null, entry_point_kind text not null, " "target text not null, wired integer not null)" )