from psycopg import Connection, sql from .connection import project_connection, project_transaction from .database import acquire_model_mutation_lock _MODEL_SCHEMAS = ("network", "gis") _ALLOWED_EXTERNAL_REFERENCES = { ("analysis", "results", "network", "nodes"), ("analysis", "results", "network", "links"), ("asset", "scada_devices", "network", "nodes"), ("asset", "scada_devices", "network", "links"), } def _model_tables(conn: Connection) -> list[tuple[str, str]]: with conn.cursor() as cur: cur.execute( """ select n.nspname as schema_name, c.relname as table_name from pg_class c join pg_namespace n on n.oid = c.relnamespace where n.nspname = any(%s) and c.relkind in ('r', 'p') order by n.nspname, c.relname """, (list(_MODEL_SCHEMAS),), ) return [(row["schema_name"], row["table_name"]) for row in cur.fetchall()] def _copy_order( conn: Connection, tables: list[tuple[str, str]] ) -> list[tuple[str, str]]: table_set = set(tables) dependencies: dict[tuple[str, str], set[tuple[str, str]]] = { table: set() for table in tables } with conn.cursor() as cur: cur.execute( """ select source_ns.nspname as source_schema, source.relname as source_table, target_ns.nspname as target_schema, target.relname as target_table from pg_constraint constraint_row join pg_class source on source.oid = constraint_row.conrelid join pg_namespace source_ns on source_ns.oid = source.relnamespace join pg_class target on target.oid = constraint_row.confrelid join pg_namespace target_ns on target_ns.oid = target.relnamespace where constraint_row.contype = 'f' and source_ns.nspname = any(%s) and target_ns.nspname = any(%s) """, (list(_MODEL_SCHEMAS), list(_MODEL_SCHEMAS)), ) for row in cur.fetchall(): source = (row["source_schema"], row["source_table"]) target = (row["target_schema"], row["target_table"]) if source in table_set and target in table_set and source != target: dependencies[source].add(target) ordered: list[tuple[str, str]] = [] remaining = set(tables) while remaining: ready = sorted( table for table in remaining if not (dependencies[table] & remaining) ) if not ready: cycle = ", ".join(f"{schema}.{table}" for schema, table in sorted(remaining)) raise RuntimeError(f"Model table foreign-key cycle detected: {cycle}") ordered.extend(ready) remaining.difference_update(ready) return ordered def _table_columns( conn: Connection, schema_name: str, table_name: str ) -> list[str]: with conn.cursor() as cur: cur.execute( """ select column_name from information_schema.columns where table_schema = %s and table_name = %s and is_generated = 'NEVER' order by ordinal_position """, (schema_name, table_name), ) return [row["column_name"] for row in cur.fetchall()] def _geometry_srids( conn: Connection, schema_name: str, table_name: str ) -> dict[str, int]: with conn.cursor() as cur: cur.execute( """ select f_geometry_column as column_name, srid from geometry_columns where f_table_schema = %s and f_table_name = %s order by f_geometry_column """, (schema_name, table_name), ) return {row["column_name"]: int(row["srid"]) for row in cur.fetchall()} def _copy_out_statement( schema_name: str, table_name: str, columns: list[str], *, source_geometry_srids: dict[str, int], target_geometry_srids: dict[str, int], ) -> sql.Composed: relation = sql.Identifier(schema_name, table_name) column_list = sql.SQL(", ").join(map(sql.Identifier, columns)) srid_changes = { column_name: target_geometry_srids[column_name] for column_name, source_srid in source_geometry_srids.items() if target_geometry_srids[column_name] != source_srid } if not srid_changes: return sql.SQL("copy {} ({}) to stdout").format(relation, column_list) select_list = sql.SQL(", ").join( sql.SQL("st_setsrid({}, {}) as {}").format( sql.Identifier(column_name), sql.Literal(srid_changes[column_name]), sql.Identifier(column_name), ) if column_name in srid_changes else sql.Identifier(column_name) for column_name in columns ) return sql.SQL("copy (select {} from {}) to stdout").format( select_list, relation, ) def _external_references(conn: Connection) -> set[tuple[str, str, str, str]]: with conn.cursor() as cur: cur.execute( """ select source_ns.nspname as source_schema, source.relname as source_table, target_ns.nspname as target_schema, target.relname as target_table from pg_constraint constraint_row join pg_class source on source.oid = constraint_row.conrelid join pg_namespace source_ns on source_ns.oid = source.relnamespace join pg_class target on target.oid = constraint_row.confrelid join pg_namespace target_ns on target_ns.oid = target.relnamespace where constraint_row.contype = 'f' and target_ns.nspname = any(%s) and source_ns.nspname <> all(%s) """, (list(_MODEL_SCHEMAS), list(_MODEL_SCHEMAS)), ) return { ( row["source_schema"], row["source_table"], row["target_schema"], row["target_table"], ) for row in cur.fetchall() } def _copy_table( source_conn: Connection, target_conn: Connection, schema_name: str, table_name: str, columns: list[str], ) -> None: relation = sql.Identifier(schema_name, table_name) column_list = sql.SQL(", ").join(map(sql.Identifier, columns)) source_geometry_srids = _geometry_srids(source_conn, schema_name, table_name) target_geometry_srids = _geometry_srids(target_conn, schema_name, table_name) if source_geometry_srids.keys() != target_geometry_srids.keys(): raise RuntimeError( "Source/target geometry columns differ for " f"{schema_name}.{table_name}" ) copy_out = _copy_out_statement( schema_name, table_name, columns, source_geometry_srids=source_geometry_srids, target_geometry_srids=target_geometry_srids, ) copy_in = sql.SQL("copy {} ({}) from stdin").format(relation, column_list) with source_conn.cursor().copy(copy_out) as source_copy: with target_conn.cursor().copy(copy_in) as target_copy: for chunk in source_copy: target_copy.write(chunk) def copy_network_tables(source_project: str, target_project: str) -> None: """Copy only the immutable simulation input tables into a project template.""" with project_connection(source_project) as source_conn, source_conn.transaction(): with source_conn.cursor() as cur: cur.execute("set transaction isolation level repeatable read, read only") source_tables = [ table for table in _model_tables(source_conn) if table[0] == "network" ] source_columns = { table: _table_columns(source_conn, *table) for table in source_tables } copy_order = _copy_order(source_conn, source_tables) with project_transaction(target_project) as target_conn: target_tables = [ table for table in _model_tables(target_conn) if table[0] == "network" ] if set(target_tables) != set(source_tables): raise RuntimeError("Source/target network schemas differ") for table, columns in source_columns.items(): if _table_columns(target_conn, *table) != columns: raise RuntimeError( f"Source/target columns differ for {table[0]}.{table[1]}" ) with target_conn.cursor() as cur: for schema_name, table_name in reversed(copy_order): cur.execute( sql.SQL("delete from {}").format( sql.Identifier(schema_name, table_name) ) ) for schema_name, table_name in copy_order: _copy_table( source_conn, target_conn, schema_name, table_name, source_columns[(schema_name, table_name)], ) def replace_project_model( target_project: str, source_project: str, *, copy_source_scada: bool = False, ) -> None: """Atomically replace WNDB/GIS model tables from a validated staging DB. Business and analysis tables remain in the target database. Historical analysis rows keep references that still exist and become element-neutral when an element disappeared. SCADA devices are retained only when their bound node or link still exists in the replacement model. ``copy_source_scada`` is used only for temporary project clones. Normal INP replacement keeps the target project's existing device mappings and drops mappings whose model element no longer exists. """ with project_connection(source_project) as source_conn, source_conn.transaction(): with source_conn.cursor() as cur: cur.execute("set transaction isolation level repeatable read, read only") source_tables = _model_tables(source_conn) source_columns = { table: _table_columns(source_conn, *table) for table in source_tables } source_scada_columns = ( _table_columns(source_conn, "asset", "scada_devices") if copy_source_scada else [] ) copy_order = _copy_order(source_conn, source_tables) with project_transaction(target_project) as target_conn: acquire_model_mutation_lock(target_conn, target_project) target_tables = _model_tables(target_conn) if set(target_tables) != set(source_tables): missing = sorted(set(source_tables) - set(target_tables)) extra = sorted(set(target_tables) - set(source_tables)) raise RuntimeError( f"Staging/target model schema mismatch; missing={missing}, extra={extra}" ) for table, columns in source_columns.items(): if _table_columns(target_conn, *table) != columns: raise RuntimeError( f"Staging/target columns differ for {table[0]}.{table[1]}" ) if copy_source_scada and _table_columns( target_conn, "asset", "scada_devices" ) != source_scada_columns: raise RuntimeError( "Source/target columns differ for asset.scada_devices" ) unexpected = _external_references(target_conn) - _ALLOWED_EXTERNAL_REFERENCES if unexpected: formatted = ", ".join( f"{source_schema}.{source_table}->{target_schema}.{target_table}" for source_schema, source_table, target_schema, target_table in sorted( unexpected ) ) raise RuntimeError( f"Model replacement has unsupported external references: {formatted}" ) with target_conn.cursor() as cur: if not copy_source_scada: cur.execute( "create temporary table model_scada_snapshot on commit drop " "as table asset.scada_devices" ) cur.execute("delete from asset.scada_devices") cur.execute( "create temporary table model_result_refs on commit drop as " "select result_id, node_id, link_id from analysis.results" ) cur.execute("update analysis.results set node_id = null, link_id = null") for schema_name, table_name in reversed(copy_order): cur.execute( sql.SQL("delete from {}").format( sql.Identifier(schema_name, table_name) ) ) for schema_name, table_name in copy_order: _copy_table( source_conn, target_conn, schema_name, table_name, source_columns[(schema_name, table_name)], ) if copy_source_scada: _copy_table( source_conn, target_conn, "asset", "scada_devices", source_scada_columns, ) with target_conn.cursor() as cur: if not copy_source_scada: cur.execute( """ insert into asset.scada_devices select snapshot.* from model_scada_snapshot snapshot where ( snapshot.node_id is not null and exists ( select 1 from network.nodes node where node.id = snapshot.node_id ) ) or ( snapshot.link_id is not null and exists ( select 1 from network.links link where link.id = snapshot.link_id ) ) """ ) cur.execute( """ update analysis.results result set node_id = case when exists ( select 1 from network.nodes node where node.id = refs.node_id ) then refs.node_id else null end, link_id = case when exists ( select 1 from network.links link where link.id = refs.link_id ) then refs.link_id else null end from model_result_refs refs where refs.result_id = result.result_id """ )