from contextlib import contextmanager import pytest from app.native.wndb.commands import executor from app.native.wndb.core import database from app.native.wndb.core.database import ChangeSet from app.native.wndb.model import junctions, pipes, pumps, reservoirs, tanks, valves def test_batch_commits_before_materialized_view_refresh(monkeypatch) -> None: events: list[str] = [] @contextmanager def fake_transaction(_name: str): events.append("transaction-enter") yield object() events.append("transaction-exit") monkeypatch.setattr(executor, "model_mutation_transaction", fake_transaction) monkeypatch.setattr( executor, "expand_command", lambda _name, change_set: change_set, ) monkeypatch.setattr( executor, "_execute_update_command", lambda _name, _change_set: events.append("write") or ChangeSet(), ) monkeypatch.setattr( executor, "refresh_materialized_views_after_commit", lambda _name: events.append("refresh"), ) executor.execute_batch_commands( "project_a", ChangeSet({"operation": "update", "type": "junction", "id": "J1"}), ) assert events == [ "transaction-enter", "write", "transaction-exit", "refresh", ] def test_failed_batch_does_not_refresh_materialized_views(monkeypatch) -> None: events: list[str] = [] @contextmanager def fake_transaction(_name: str): events.append("transaction-enter") try: yield object() finally: events.append("transaction-exit") monkeypatch.setattr(executor, "model_mutation_transaction", fake_transaction) monkeypatch.setattr( executor, "expand_command", lambda _name, change_set: change_set, ) def fail_write(_name, _change_set): events.append("write") raise RuntimeError("write failed") monkeypatch.setattr(executor, "_execute_update_command", fail_write) monkeypatch.setattr( executor, "refresh_materialized_views_after_commit", lambda _name: events.append("refresh"), ) with pytest.raises(RuntimeError, match="write failed"): executor.execute_batch_commands( "project_a", ChangeSet({"operation": "update", "type": "junction", "id": "J1"}), ) assert events == ["transaction-enter", "write", "transaction-exit"] def test_batch_option_update_does_not_refresh_materialized_views(monkeypatch) -> None: events: list[str] = [] @contextmanager def fake_transaction(_name: str): events.append("transaction-enter") yield object() events.append("transaction-exit") monkeypatch.setattr(executor, "model_mutation_transaction", fake_transaction) monkeypatch.setattr(executor, "expand_command", lambda _name, cs: cs) monkeypatch.setattr( executor, "_execute_update_command", lambda _name, _change_set: events.append("write") or ChangeSet(), ) monkeypatch.setattr( executor, "refresh_materialized_views_after_commit", lambda _name: events.append("refresh"), ) executor.execute_batch_commands( "project_a", ChangeSet({"operation": "update", "type": "option", "id": "duration"}), ) assert events == ["transaction-enter", "write", "transaction-exit"] def test_model_mutation_lock_is_acquired_once_per_transaction(monkeypatch) -> None: state = {"held": False} statements: list[tuple[str, tuple[str]]] = [] class FakeCursor: def __enter__(self): return self def __exit__(self, *_args): return None def execute(self, statement: str, params: tuple[str]): statements.append((statement, params)) class FakeConnection: def cursor(self): return FakeCursor() monkeypatch.setattr( database, "is_model_mutation_lock_active", lambda _name: state["held"], raising=False, ) monkeypatch.setattr( database, "mark_model_mutation_lock_active", lambda _name: state.__setitem__("held", True), raising=False, ) monkeypatch.setattr(database, "get_project_database_name", lambda name: name) conn = FakeConnection() database.acquire_model_mutation_lock(conn, "project_a") database.acquire_model_mutation_lock(conn, "project_a") assert len(statements) == 1 def test_locked_command_builds_after_lock_and_refreshes_after_commit( monkeypatch, ) -> None: events: list[str] = [] @contextmanager def fake_model_transaction(_name: str): events.append("lock") yield object() events.append("commit") def build_command() -> database.DatabaseCommand: events.append("read-and-build") return database.DatabaseCommand( "UPDATE network.junctions SET elevation = 1", [{"operation": "update", "type": "junction", "id": "J1"}], ) monkeypatch.setattr(database, "model_mutation_transaction", fake_model_transaction) monkeypatch.setattr( database, "is_project_transaction_active", lambda _name: False, ) monkeypatch.setattr( database, "execute_command", lambda _name, command: events.append("write") or ChangeSet.from_list(command.changes), ) monkeypatch.setattr( database, "refresh_materialized_views_after_commit", lambda _name: events.append("refresh"), ) result = database.execute_locked_command("project_a", build_command) assert events == ["lock", "read-and-build", "write", "commit", "refresh"] assert result.operations[0]["id"] == "J1" def test_pipe_patch_reads_under_lock_and_updates_only_supplied_columns( monkeypatch, ) -> None: events: list[str] = [] current = { "id": "P1", "node1": "J1", "node2": "J2", "length": 10.0, "diameter": 100.0, "roughness": 120.0, "minor_loss": 0.0, "status": "OPEN", } def fake_locked_command(_name: str, builder): events.append("lock") command = builder() assert command is not None captured.append(command) events.append("refresh") return ChangeSet.from_list(command.changes) def fake_get_pipe(_name: str, _id: str): events.append("read") return current.copy() captured: list[database.DatabaseCommand] = [] monkeypatch.setattr( pipes, "execute_locked_command", fake_locked_command, ) monkeypatch.setattr(pipes, "get_pipe", fake_get_pipe) result = pipes.set_pipe( "project_a", ChangeSet({"operation": "update", "type": "pipe", "id": "P1", "length": 20}), ) assert events == ["lock", "read", "refresh"] assert len(captured) == 1 assert "length = 20.0" in captured[0].sql assert "diameter" not in captured[0].sql assert "update network.links" not in captured[0].sql.lower() assert result.operations[0]["length"] == 20.0 @pytest.mark.parametrize( ("module", "setter_name", "getter_name", "builder_name"), [ (junctions, "set_junction", "get_junction", "_set_junction"), (reservoirs, "set_reservoir", "get_reservoir", "_set_reservoir"), (tanks, "set_tank", "get_tank", "_set_tank"), (pumps, "set_pump", "get_pump", "_set_pump"), (valves, "set_valve", "get_valve", "_set_valve"), ], ) def test_element_patch_reads_after_shared_model_lock( monkeypatch, module, setter_name: str, getter_name: str, builder_name: str, ) -> None: events: list[str] = [] current = {"id": "E1"} def fake_getter(_name: str, _id: str): events.append("read") return current def fake_builder(_name: str, _changes: ChangeSet, supplied_current): events.append("build") assert supplied_current is current return database.DatabaseCommand("UPDATE network.nodes SET id = id", []) def fake_locked_command(_name: str, builder): events.append("lock") assert builder() is not None return ChangeSet() monkeypatch.setattr(module, getter_name, fake_getter) monkeypatch.setattr(module, builder_name, fake_builder) monkeypatch.setattr(module, "execute_locked_command", fake_locked_command) getattr(module, setter_name)( "project_a", ChangeSet({"operation": "update", "type": "element", "id": "E1"}), ) assert events == ["lock", "read", "build"]