from collections.abc import Iterator, Mapping, Sequence from contextlib import contextmanager from typing import Any, Callable from psycopg import sql from psycopg.rows import Row, dict_row from app.infra.db.project_routing import get_project_database_name from .connection import ( is_model_mutation_lock_active, is_project_transaction_active, mark_model_mutation_lock_active, project_connection, project_transaction, ) API_ADD = "add" API_UPDATE = "update" API_DELETE = "delete" g_add_prefix = {"operation": API_ADD} g_update_prefix = {"operation": API_UPDATE} g_delete_prefix = {"operation": API_DELETE} class ChangeSet: def __init__(self, ps: dict[str, Any] | None = None): self.operations: list[dict[str, Any]] = [] if ps is not None: self.append(ps) @staticmethod def from_list(ps: list[dict[str, Any]]): change_set = ChangeSet() for item in ps: change_set.append(item) return change_set def add(self, ps: dict[str, Any]): self.operations.append(g_add_prefix | ps) return self def update(self, ps: dict[str, Any]): self.operations.append(g_update_prefix | ps) return self def delete(self, ps: dict[str, Any]): self.operations.append(g_delete_prefix | ps) return self def append(self, ps: dict[str, Any]): self.operations.append(ps) return self def merge(self, change_set): self.operations.extend(change_set.operations) return self def dump(self): for operation in self.operations: print(operation) def compress(self): return self class DatabaseCommand: def __init__(self, statement: str, changes: list[dict[str, Any]]) -> None: self.sql = statement self.changes = changes class MaterializedViewRefreshAfterCommitError(RuntimeError): """Report a failed view refresh without implying that the write rolled back.""" changes_committed = True def __init__(self, project: str) -> None: self.project = project super().__init__( f"Project {project!r} changes were committed, but materialized view " "refresh failed" ) QueryParams = Sequence[Any] | Mapping[str, Any] def sql_literal(value: Any) -> str: """Render one PostgreSQL literal for legacy WNDB SQL batch builders. WNDB still assembles multi-statement model changes before executing them as one transaction. Every interpolated value must pass through this helper; identifiers remain static strings owned by the backend. """ return sql.Literal(value).as_string() def _execute(cur, query: str, params: QueryParams | None = None): return cur.execute(query, params) if params is not None else cur.execute(query) def acquire_model_mutation_lock(conn, name: str) -> None: """Serialize model replacement and ordinary WNDB mutations per database.""" if is_model_mutation_lock_active(name): return physical_name = get_project_database_name(name) with conn.cursor() as cur: cur.execute( "select pg_advisory_xact_lock(hashtextextended(%s, 0))", (f"tjwater:wndb:model:{physical_name}",), ) mark_model_mutation_lock_active(name) @contextmanager def model_mutation_transaction(name: str) -> Iterator[Any]: """Open a project transaction and acquire its model lock before reading.""" with project_transaction(name) as conn: acquire_model_mutation_lock(conn, name) yield conn def read(name: str, query: str, params: QueryParams | None = None) -> Row: with project_connection(name) as conn, conn.cursor(row_factory=dict_row) as cur: _execute(cur, query, params) row = cur.fetchone() if row is None: raise LookupError(query) return row def read_all( name: str, query: str, params: QueryParams | None = None ) -> list[Row]: with project_connection(name) as conn, conn.cursor(row_factory=dict_row) as cur: _execute(cur, query, params) return cur.fetchall() def try_read( name: str, query: str, params: QueryParams | None = None ) -> Row | None: with project_connection(name) as conn, conn.cursor(row_factory=dict_row) as cur: _execute(cur, query, params) return cur.fetchone() def write(name: str, query: str, params: QueryParams | None = None) -> None: connection_context = ( project_connection(name) if is_project_transaction_active(name) else model_mutation_transaction(name) ) with connection_context as conn: acquire_model_mutation_lock(conn, name) with conn.cursor() as cur: _execute(cur, query, params) def refresh_materialized_views(name: str, *, concurrently: bool = True) -> None: """Refresh the GIS query layer after committed model or asset changes.""" with project_connection(name) as conn, conn.cursor() as cur: cur.execute("CALL gis.refresh_all_materialized_views(%s)", (concurrently,)) def refresh_materialized_views_after_commit(name: str) -> None: try: refresh_materialized_views(name) except Exception as exc: raise MaterializedViewRefreshAfterCommitError(name) from exc _MATERIALIZED_VIEW_SOURCES = ( "network.nodes", "network.junctions", "network.reservoirs", "network.tanks", "network.links", "network.pipes", "network.pumps", "network.valves", "network.demands", "gis.node_geometries", "gis.link_vertices", "asset.scada_devices", ) def _affects_materialized_views(command: DatabaseCommand) -> bool: statement = command.sql.lower() return any(source in statement for source in _MATERIALIZED_VIEW_SOURCES) _MATERIALIZED_VIEW_ELEMENT_TYPES = frozenset( { "junction", "reservoir", "tank", "pipe", "pump", "valve", "demand", "vertex", } ) def changes_affect_materialized_views(change_set: ChangeSet) -> bool: """Return whether a dispatched WNDB batch changes a published GIS source.""" return any( operation.get("type") in _MATERIALIZED_VIEW_ELEMENT_TYPES for operation in change_set.operations ) def execute_command(name: str, command: DatabaseCommand) -> ChangeSet: """Apply a model mutation without the removed database undo/redo journal.""" write(name, command.sql) if _affects_materialized_views(command) and not is_project_transaction_active(name): refresh_materialized_views_after_commit(name) return ChangeSet.from_list(command.changes) def execute_locked_command( name: str, builder: Callable[[], DatabaseCommand | None], ) -> ChangeSet: """Build a read-modify-write command only after acquiring the model lock.""" nested_transaction = is_project_transaction_active(name) command: DatabaseCommand | None = None with model_mutation_transaction(name): command = builder() if command is None: return ChangeSet() result = execute_command(name, command) if not nested_transaction and _affects_materialized_views(command): refresh_materialized_views_after_commit(name) return result