249 lines
6.4 KiB
Python
249 lines
6.4 KiB
Python
from typing import Any
|
|
|
|
from psycopg import sql
|
|
from psycopg.rows import Row, dict_row
|
|
|
|
from ..core.connection import project_connection
|
|
from ..core.database import read
|
|
|
|
_NODE = "network.nodes"
|
|
_LINK = "network.links"
|
|
_CURVE = "network.curves"
|
|
_PATTERN = "network.patterns"
|
|
_REGION = "gis.regions"
|
|
|
|
JUNCTION = "junction"
|
|
RESERVOIR = "reservoir"
|
|
TANK = "tank"
|
|
PIPE = "pipe"
|
|
PUMP = "pump"
|
|
VALVE = "valve"
|
|
PATTERN = "pattern"
|
|
CURVE = "curve"
|
|
REGION = "region"
|
|
|
|
ELEMENT_TYPES: dict[str, int] = {
|
|
RESERVOIR: 0,
|
|
TANK: 1,
|
|
JUNCTION: 2,
|
|
PIPE: 3,
|
|
PUMP: 4,
|
|
VALVE: 5,
|
|
}
|
|
|
|
|
|
def _table_identifier(table: str):
|
|
return sql.Identifier(*table.split("."))
|
|
|
|
|
|
def _get_from(name: str, element_id: str, table: str) -> Row | None:
|
|
query = sql.SQL("SELECT * FROM {} WHERE id = %s").format(
|
|
_table_identifier(table)
|
|
)
|
|
with project_connection(name) as conn, conn.cursor(row_factory=dict_row) as cur:
|
|
cur.execute(query, (element_id,))
|
|
return cur.fetchone()
|
|
|
|
|
|
def is_node(name: str, element_id: str) -> bool:
|
|
return _get_from(name, element_id, _NODE) is not None
|
|
|
|
|
|
def is_junction(name: str, element_id: str) -> bool:
|
|
row = _get_from(name, element_id, _NODE)
|
|
return row is not None and row["node_type"] == JUNCTION
|
|
|
|
|
|
def is_reservoir(name: str, element_id: str) -> bool:
|
|
row = _get_from(name, element_id, _NODE)
|
|
return row is not None and row["node_type"] == RESERVOIR
|
|
|
|
|
|
def is_tank(name: str, element_id: str) -> bool:
|
|
row = _get_from(name, element_id, _NODE)
|
|
return row is not None and row["node_type"] == TANK
|
|
|
|
|
|
def is_link(name: str, element_id: str) -> bool:
|
|
return _get_from(name, element_id, _LINK) is not None
|
|
|
|
|
|
def is_pipe(name: str, element_id: str) -> bool:
|
|
row = _get_from(name, element_id, _LINK)
|
|
return row is not None and row["link_type"] == PIPE
|
|
|
|
|
|
def is_pump(name: str, element_id: str) -> bool:
|
|
row = _get_from(name, element_id, _LINK)
|
|
return row is not None and row["link_type"] == PUMP
|
|
|
|
|
|
def is_valve(name: str, element_id: str) -> bool:
|
|
row = _get_from(name, element_id, _LINK)
|
|
return row is not None and row["link_type"] == VALVE
|
|
|
|
|
|
def get_node_type(name: str, node_id: str) -> str:
|
|
row = _get_from(name, node_id, _NODE)
|
|
if row is None:
|
|
raise LookupError(node_id)
|
|
return row["node_type"]
|
|
|
|
|
|
def get_link_type(name: str, link_id: str) -> str:
|
|
row = _get_from(name, link_id, _LINK)
|
|
if row is None:
|
|
raise LookupError(link_id)
|
|
return row["link_type"]
|
|
|
|
|
|
def get_element_type(name: str, element_id: str) -> str | None:
|
|
if is_node(name, element_id):
|
|
return get_node_type(name, element_id)
|
|
if is_link(name, element_id):
|
|
return get_link_type(name, element_id)
|
|
return None
|
|
|
|
|
|
def get_element_type_value(name: str, element_id: str) -> int:
|
|
element_type = get_element_type(name, element_id)
|
|
if element_type is None:
|
|
raise LookupError(element_id)
|
|
return ELEMENT_TYPES[element_type]
|
|
|
|
|
|
def is_curve(name: str, element_id: str) -> bool:
|
|
return _get_from(name, element_id, _CURVE) is not None
|
|
|
|
|
|
def is_pattern(name: str, element_id: str) -> bool:
|
|
return _get_from(name, element_id, _PATTERN) is not None
|
|
|
|
|
|
def is_region(name: str, element_id: str) -> bool:
|
|
return _get_from(name, element_id, _REGION) is not None
|
|
|
|
|
|
def _get_all(name: str, table: str) -> list[str]:
|
|
query = sql.SQL("SELECT id FROM {} ORDER BY id").format(_table_identifier(table))
|
|
with project_connection(name) as conn, conn.cursor(row_factory=dict_row) as cur:
|
|
cur.execute(query)
|
|
return [row["id"] for row in cur]
|
|
|
|
|
|
def _get_nodes_by_type(name: str, node_type: str) -> list[str]:
|
|
rows = read_all_typed(
|
|
name,
|
|
"SELECT id FROM network.nodes WHERE node_type = %s ORDER BY id",
|
|
(node_type,),
|
|
)
|
|
return [row["id"] for row in rows]
|
|
|
|
|
|
def _get_links_by_type(name: str, link_type: str) -> list[str]:
|
|
rows = read_all_typed(
|
|
name,
|
|
"SELECT id FROM network.links WHERE link_type = %s ORDER BY id",
|
|
(link_type,),
|
|
)
|
|
return [row["id"] for row in rows]
|
|
|
|
|
|
def read_all_typed(name: str, query: str, params: tuple[Any, ...]) -> list[Row]:
|
|
with project_connection(name) as conn, conn.cursor(row_factory=dict_row) as cur:
|
|
cur.execute(query, params)
|
|
return cur.fetchall()
|
|
|
|
|
|
def get_nodes(name: str) -> list[str]:
|
|
return _get_all(name, _NODE)
|
|
|
|
|
|
def get_junctions(name: str) -> list[str]:
|
|
return _get_nodes_by_type(name, JUNCTION)
|
|
|
|
|
|
def get_reservoirs(name: str) -> list[str]:
|
|
return _get_nodes_by_type(name, RESERVOIR)
|
|
|
|
|
|
def get_tanks(name: str) -> list[str]:
|
|
return _get_nodes_by_type(name, TANK)
|
|
|
|
|
|
def get_links(name: str) -> list[str]:
|
|
return _get_all(name, _LINK)
|
|
|
|
|
|
def get_pipes(name: str) -> list[str]:
|
|
return _get_links_by_type(name, PIPE)
|
|
|
|
|
|
def get_pumps(name: str) -> list[str]:
|
|
return _get_links_by_type(name, PUMP)
|
|
|
|
|
|
def get_valves(name: str) -> list[str]:
|
|
return _get_links_by_type(name, VALVE)
|
|
|
|
|
|
def get_curves(name: str) -> list[str]:
|
|
return _get_all(name, _CURVE)
|
|
|
|
|
|
def get_patterns(name: str) -> list[str]:
|
|
return _get_all(name, _PATTERN)
|
|
|
|
|
|
def get_regions(name: str) -> list[str]:
|
|
return _get_all(name, _REGION)
|
|
|
|
|
|
def get_node_links(name: str, node_id: str) -> list[str]:
|
|
rows = read_all_typed(
|
|
name,
|
|
"""
|
|
SELECT id FROM network.links
|
|
WHERE start_node_id = %s OR end_node_id = %s
|
|
ORDER BY id
|
|
""",
|
|
(node_id, node_id),
|
|
)
|
|
return [row["id"] for row in rows]
|
|
|
|
|
|
def get_all_node_links(name: str) -> dict[str, list[str]]:
|
|
"""Build the node adjacency map with one scan of the unified GIS view."""
|
|
rows = read_all_typed(
|
|
name,
|
|
"SELECT id, start_node_id, end_node_id FROM gis.network_links ORDER BY id",
|
|
(),
|
|
)
|
|
result: dict[str, list[str]] = {}
|
|
for row in rows:
|
|
link_id = str(row["id"])
|
|
result.setdefault(str(row["start_node_id"]), []).append(link_id)
|
|
result.setdefault(str(row["end_node_id"]), []).append(link_id)
|
|
return result
|
|
|
|
|
|
def get_link_nodes(name: str, link_id: str) -> list[str]:
|
|
row = read(
|
|
name,
|
|
"""
|
|
SELECT start_node_id, end_node_id
|
|
FROM network.links WHERE id = %s
|
|
""",
|
|
(link_id,),
|
|
)
|
|
return [str(row["start_node_id"]), str(row["end_node_id"])]
|
|
|
|
|
|
def get_region_type(name: str, region_id: str) -> str:
|
|
row = read(
|
|
name,
|
|
"SELECT region_type FROM gis.regions WHERE id = %s",
|
|
(region_id,),
|
|
)
|
|
return row["region_type"]
|