Files
TJWaterServerBinary/app/native/wndb/core/model_replace.py
T
jiang 90b02057bc
Generic Container CI/CD / test-build-publish (push) Successful in 1m13s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 1m13s
feat(projects): automate project infrastructure provisioning
2026-09-11 10:57:51 +08:00

389 lines
15 KiB
Python

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
"""
)