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_nodes_id_and_type(name: str) -> dict[str, str]: rows = read_all_typed(name, "SELECT id, node_type FROM network.nodes", ()) return {row["id"]: row["node_type"] for row in rows} def get_major_nodes(name: str, diameter: int) -> list[str]: rows = read_all_typed( name, """ SELECT DISTINCT endpoint FROM network.links AS l JOIN network.pipes AS p ON p.link_id = l.id CROSS JOIN LATERAL (VALUES (l.start_node_id), (l.end_node_id)) AS e(endpoint) WHERE p.diameter > %s """, (diameter,), ) return [row["endpoint"] for row in rows] 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_links_id_and_type(name: str) -> dict[str, str]: rows = read_all_typed(name, "SELECT id, link_type FROM network.links", ()) return {row["id"]: row["link_type"] for row in rows} def get_major_pipes(name: str, diameter: int) -> list[str]: rows = read_all_typed( name, "SELECT link_id FROM network.pipes WHERE diameter > %s ORDER BY link_id", (diameter,), ) return [row["link_id"] for row in rows] 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 link table.""" rows = read_all_typed( name, "SELECT id, start_node_id, end_node_id FROM 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"]