diff --git a/app/native/wndb/connection.py b/app/native/wndb/connection.py index b42b481..7dd5de2 100644 --- a/app/native/wndb/connection.py +++ b/app/native/wndb/connection.py @@ -1,3 +1,66 @@ +from collections.abc import Iterator +from contextlib import contextmanager +from threading import RLock + import psycopg as pg -g_conn_dict : dict[str, pg.Connection] = {} \ No newline at end of file +from app.core.config import get_pgconn_string + +g_conn_dict: dict[str, pg.Connection] = {} +_registry_lock = RLock() +_project_locks: dict[str, RLock] = {} + + +def _is_closed(connection: pg.Connection) -> bool: + return bool(getattr(connection, "closed", False)) + + +def _close_connection(connection: pg.Connection) -> None: + if not _is_closed(connection): + connection.close() + + +def _get_project_lock(name: str) -> RLock: + with _registry_lock: + lock = _project_locks.get(name) + if lock is None: + lock = RLock() + _project_locks[name] = lock + return lock + + +def open_connection(name: str) -> pg.Connection: + with _get_project_lock(name): + connection = g_conn_dict.get(name) + if connection is None or _is_closed(connection): + if connection is not None: + _close_connection(connection) + connection = pg.connect( + conninfo=get_pgconn_string(db_name=name), autocommit=True + ) + g_conn_dict[name] = connection + return connection + + +def is_connection_open(name: str) -> bool: + with _get_project_lock(name): + connection = g_conn_dict.get(name) + if connection is None: + return False + if _is_closed(connection): + del g_conn_dict[name] + return False + return True + + +def close_connection(name: str) -> None: + with _get_project_lock(name): + connection = g_conn_dict.pop(name, None) + if connection is not None: + _close_connection(connection) + + +@contextmanager +def project_connection(name: str) -> Iterator[pg.Connection]: + with _get_project_lock(name): + yield open_connection(name) diff --git a/app/native/wndb/database.py b/app/native/wndb/database.py index 248b4a3..6d2893b 100644 --- a/app/native/wndb/database.py +++ b/app/native/wndb/database.py @@ -1,6 +1,6 @@ from typing import Any from psycopg.rows import dict_row, Row -from .connection import g_conn_dict as conn +from .connection import project_connection API_ADD = 'add' API_UPDATE = 'update' @@ -83,29 +83,33 @@ class DbChangeSet: def read(name: str, sql: str) -> Row: - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(sql) - row = cur.fetchone() - if row == None: - raise Exception(sql) - return row + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(sql) + row = cur.fetchone() + if row == None: + raise Exception(sql) + return row def read_all(name: str, sql: str) -> list[Row]: - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(sql) - return cur.fetchall() + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(sql) + return cur.fetchall() def try_read(name: str, sql: str) -> Row | None: - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(sql) - return cur.fetchone() + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(sql) + return cur.fetchone() def write(name: str, sql: str) -> None: - with conn[name].cursor() as cur: - cur.execute(sql) + with project_connection(name) as conn: + with conn.cursor() as cur: + cur.execute(sql) def get_current_operation(name: str) -> int: diff --git a/app/native/wndb/project.py b/app/native/wndb/project.py index 25b9328..5403d0f 100644 --- a/app/native/wndb/project.py +++ b/app/native/wndb/project.py @@ -2,7 +2,11 @@ import os import psycopg as pg from psycopg import sql from psycopg.rows import dict_row -from .connection import g_conn_dict as conn +from .connection import ( + close_connection, + is_connection_open, + open_connection, +) from app.core.config import get_pgconn_string, get_pg_config, get_pg_password # no undo/redo @@ -31,9 +35,7 @@ def have_project(name: str) -> bool: def copy_project(source: str, new: str) -> None: - if source in conn: - conn[source].close() - del conn[source] + close_connection(source) with pg.connect( conninfo=get_pgconn_string(db_name="postgres"), autocommit=True @@ -176,17 +178,12 @@ def clean_project(excluded: list[str] = []) -> None: def open_project(name: str) -> None: - if name not in conn: - conn[name] = pg.connect( - conninfo=get_pgconn_string(db_name=name), autocommit=True - ) + open_connection(name) def is_project_open(name: str) -> bool: - return name in conn + return is_connection_open(name) def close_project(name: str) -> None: - if name in conn: - conn[name].close() - del conn[name] + close_connection(name) diff --git a/app/native/wndb/s0_base.py b/app/native/wndb/s0_base.py index 65882ab..2986af2 100644 --- a/app/native/wndb/s0_base.py +++ b/app/native/wndb/s0_base.py @@ -1,5 +1,5 @@ from psycopg.rows import dict_row, Row -from .connection import g_conn_dict as conn +from .connection import project_connection from .database import read from typing import Any @@ -47,9 +47,10 @@ ELEMENT_TYPES : dict[str, int] = { } def _get_from(name: str, id: str, base_type: str) -> Row | None: - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select * from {base_type} where id = '{id}'") - return cur.fetchone() + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select * from {base_type} where id = '{id}'") + return cur.fetchone() def is_node(name: str, id: str) -> bool: @@ -125,10 +126,11 @@ def is_region(name: str, id: str) -> bool: def _get_all(name: str, base_type: str) -> list[str]: ids : list[str] = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id from {base_type} order by id") - for record in cur: - ids.append(record['id']) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id from {base_type} order by id") + for record in cur: + ids.append(record['id']) return ids @@ -138,29 +140,32 @@ def get_nodes(name: str) -> list[str]: # DingZQ def _get_nodes_by_type(name: str, type: str) -> list[str]: ids : list[str] = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id from {_NODE} where type = '{type}' order by id") - for record in cur: - ids.append(record['id']) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id from {_NODE} where type = '{type}' order by id") + for record in cur: + ids.append(record['id']) return ids # DingZQ def get_nodes_id_and_type(name: str) -> dict[str, str]: nodes_id_and_type: dict[str, str] = {} - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id, type from {_NODE} order by id") - for record in cur: - nodes_id_and_type[record['id']] = record['type'] + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id, type from {_NODE} order by id") + for record in cur: + nodes_id_and_type[record['id']] = record['type'] return nodes_id_and_type # DingZQ 2024-12-31 def get_major_nodes(name: str, diameter: int) -> list[str]: major_nodes_set = set() - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select node1, node2 from pipes where diameter > {diameter}") - for record in cur: - major_nodes_set.add(record['node1']) - major_nodes_set.add(record['node2']) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select node1, node2 from pipes where diameter > {diameter}") + for record in cur: + major_nodes_set.add(record['node1']) + major_nodes_set.add(record['node2']) return list(major_nodes_set) @@ -183,29 +188,32 @@ def get_links(name: str) -> list[str]: # DingZQ def _get_links_by_type(name: str, type: str) -> list[str]: ids : list[str] = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id from {_LINK} where type = '{type}' order by id") - for record in cur: - ids.append(record['id']) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id from {_LINK} where type = '{type}' order by id") + for record in cur: + ids.append(record['id']) return ids # DingZQ def get_links_id_and_type(name: str) -> dict[str, str]: links_id_and_type: dict[str, str] = {} - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id, type from {_LINK} order by id") - for record in cur: - links_id_and_type[record['id']] = record['type'] + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id, type from {_LINK} order by id") + for record in cur: + links_id_and_type[record['id']] = record['type'] return links_id_and_type # DingZQ 2024-12-31 # 获取直径大于800的管道 def get_major_pipes(name: str, diameter: int) -> list[str]: major_pipe_ids: list[str] = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id from pipes where diameter > {diameter} order by id") - for record in cur: - major_pipe_ids.append(record['id']) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id from pipes where diameter > {diameter} order by id") + for record in cur: + major_pipe_ids.append(record['id']) return major_pipe_ids # DingZQ @@ -232,15 +240,16 @@ def get_regions(name: str) -> list[str]: return _get_all(name, _REGION) def get_node_links(name: str, id: str) -> list[str]: - with conn[name].cursor(row_factory=dict_row) as cur: - links: list[str] = [] - for p in cur.execute(f"select id from pipes where node1 = '{id}' or node2 = '{id}'").fetchall(): - links.append(p['id']) - for p in cur.execute(f"select id from pumps where node1 = '{id}' or node2 = '{id}'").fetchall(): - links.append(p['id']) - for p in cur.execute(f"select id from valves where node1 = '{id}' or node2 = '{id}'").fetchall(): - links.append(p['id']) - return links + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + links: list[str] = [] + for p in cur.execute(f"select id from pipes where node1 = '{id}' or node2 = '{id}'").fetchall(): + links.append(p['id']) + for p in cur.execute(f"select id from pumps where node1 = '{id}' or node2 = '{id}'").fetchall(): + links.append(p['id']) + for p in cur.execute(f"select id from valves where node1 = '{id}' or node2 = '{id}'").fetchall(): + links.append(p['id']) + return links def get_link_nodes(name: str, id: str) -> list[str]: @@ -259,4 +268,3 @@ def get_region_type(name: str, id: str)->str: return type - diff --git a/app/native/wndb/s24_coordinates.py b/app/native/wndb/s24_coordinates.py index d038a96..df0fcc6 100644 --- a/app/native/wndb/s24_coordinates.py +++ b/app/native/wndb/s24_coordinates.py @@ -1,5 +1,7 @@ from .database import * +from .connection import project_connection from .s0_base import get_link_nodes +from psycopg.rows import dict_row def sql_update_coord(node: str, x: float, y: float) -> str: coord = f"st_geomfromtext('point({x} {y})')" @@ -49,10 +51,11 @@ def get_links_in_extent(name: str, x1: float, y1: float, x2: float, y2: float) - node_ids = set([s.split(':')[0] for s in get_nodes_in_extent(name, x1, y1, x2, y2)]) all_link_ids = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select id from pipes") - for record in cur: - all_link_ids.append(record['id']) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select id from pipes") + for record in cur: + all_link_ids.append(record['id']) links = [] for link_id in all_link_ids: diff --git a/app/native/wndb/s41_pipe_risk_probability.py b/app/native/wndb/s41_pipe_risk_probability.py index 33f0fe9..b441305 100644 --- a/app/native/wndb/s41_pipe_risk_probability.py +++ b/app/native/wndb/s41_pipe_risk_probability.py @@ -1,5 +1,7 @@ from .database import * +from .connection import project_connection from .s0_base import * +from psycopg.rows import dict_row import json def get_pipe_risk_probability_now(name: str, pipe_id: str) -> dict[str, Any]: @@ -28,29 +30,31 @@ def get_pipe_risk_probability(name: str, pipe_id: str) -> dict[str, Any]: def get_network_pipe_risk_probability_now(name: str) -> list[dict[str, Any]]: pipe_risk_probability_list = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select * from pipe_risk_probability") - for record in cur: - #pipe_risk_probability_list.append(record) - t = {} - t['pipeid'] = record['pipeid'] - t['pipeage'] = record['pipeage'] - t['risk_probability_now'] = record['risk_probability_now'] - pipe_risk_probability_list.append(t) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select * from pipe_risk_probability") + for record in cur: + #pipe_risk_probability_list.append(record) + t = {} + t['pipeid'] = record['pipeid'] + t['pipeage'] = record['pipeage'] + t['risk_probability_now'] = record['risk_probability_now'] + pipe_risk_probability_list.append(t) return pipe_risk_probability_list def get_pipes_risk_probability(name: str, pipe_ids: list[str]) -> list[dict[str, Any]]: pipe_risk_probability_list = [] - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select * from pipe_risk_probability") - for record in cur: - if record['pipeid'] in pipe_ids: - t = {} - t['pipeid'] = record['pipeid'] - t['x'] = record['x'] - t['y'] = record['y'] - pipe_risk_probability_list.append(t) + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select * from pipe_risk_probability") + for record in cur: + if record['pipeid'] in pipe_ids: + t = {} + t['pipeid'] = record['pipeid'] + t['x'] = record['x'] + t['y'] = record['y'] + pipe_risk_probability_list.append(t) return pipe_risk_probability_list @@ -67,21 +71,22 @@ def get_pipe_risk_probability_geometries(name: str) -> dict[str, Any]: # key_endnode = '下游节点' key_geometry = 'geometry' - with conn[name].cursor(row_factory=dict_row) as cur: - cur.execute(f"select *, ST_AsGeoJSON(geometry) AS {key_geometry} from gis_pipe") + with project_connection(name) as conn: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute(f"select *, ST_AsGeoJSON(geometry) AS {key_geometry} from gis_pipe") - for record in cur: - id = record[key_pipeId] - geom = json.loads(record[key_geometry]) + for record in cur: + id = record[key_pipeId] + geom = json.loads(record[key_geometry]) - pipe_risk_probability_geometries[id] = { - 'points': geom['coordinates'] - } + pipe_risk_probability_geometries[id] = { + 'points': geom['coordinates'] + } - for col in record: - if col != key_geometry: - pipe_risk_probability_geometries[id][col] = record[col] + for col in record: + if col != key_geometry: + pipe_risk_probability_geometries[id][col] = record[col] # print(len(pipe_risk_probability_geometries)) - return pipe_risk_probability_geometries \ No newline at end of file + return pipe_risk_probability_geometries diff --git a/tests/unit/test_wndb_connection.py b/tests/unit/test_wndb_connection.py new file mode 100644 index 0000000..ea68d90 --- /dev/null +++ b/tests/unit/test_wndb_connection.py @@ -0,0 +1,78 @@ +import pytest + +from app.native.wndb import connection +from app.native.wndb import database +from app.native.wndb import project + + +class _FakeCursor: + def __init__(self, rows): + self.rows = rows + self.executed = [] + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + def execute(self, sql): + self.executed.append(sql) + + def fetchall(self): + return self.rows + + +class _FakeConnection: + def __init__(self, rows=None, *, closed=False): + self.rows = list(rows or []) + self.closed = closed + self.close_calls = 0 + + def cursor(self, row_factory=None): + if self.closed: + raise RuntimeError("the connection is closed") + return _FakeCursor(self.rows) + + def close(self): + self.close_calls += 1 + self.closed = True + + +@pytest.fixture(autouse=True) +def clear_native_connections(): + connection.g_conn_dict.clear() + connection._project_locks.clear() + yield + connection.g_conn_dict.clear() + connection._project_locks.clear() + + +def test_is_project_open_drops_closed_cached_connection(): + connection.g_conn_dict["fengyang"] = _FakeConnection(closed=True) + + assert project.is_project_open("fengyang") is False + assert "fengyang" not in connection.g_conn_dict + + +def test_read_all_reopens_closed_cached_connection(monkeypatch): + stale = _FakeConnection(closed=True) + fresh = _FakeConnection(rows=[{"key": "DURATION", "value": "01:00:00"}]) + connection.g_conn_dict["fengyang"] = stale + + opened = [] + + def fake_connect(*, conninfo, autocommit): + opened.append((conninfo, autocommit)) + return fresh + + monkeypatch.setattr(connection.pg, "connect", fake_connect) + monkeypatch.setattr( + connection, "get_pgconn_string", lambda db_name: f"dbname={db_name}" + ) + + rows = database.read_all("fengyang", "select * from times") + + assert rows == [{"key": "DURATION", "value": "01:00:00"}] + assert opened == [("dbname=fengyang", True)] + assert connection.g_conn_dict["fengyang"] is fresh