Files
TJWaterServerBinary/app/native/wndb/core/database.py
T

241 lines
7.1 KiB
Python

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