Files
TJWaterServerBinary/app/infra/db/timescaledb/internal_queries.py
T
jiang 9b095c7439 refactor(db)!: clean up business SQL access
- make realtime replacement and analysis result writes transactional\n- consolidate SCADA repositories and remove process-global project state\n- validate SCADA batches and use indexed GIS-backed business queries\n\nBREAKING CHANGE: remove the public analysis result writer and the pipeline-health network_name query parameter.
2026-08-28 11:37:36 +08:00

311 lines
12 KiB
Python

from typing import List
from fastapi.logger import logger
from datetime import datetime, timedelta
from psycopg import sql
from psycopg.rows import dict_row
import time
from app.infra.db.project_routing import get_project_timescale_pgconn_string
from app.infra.db.timescaledb.sync_pool import timescale_connection
from app.infra.db.timescaledb.repositories.analysis import AnalysisResultsRepository
from app.infra.db.timescaledb.repositories.realtime import RealtimeRepository
from app.infra.db.timescaledb.repositories.scada import ScadaRepository
from app.services.time_api import parse_utc_time
class InternalStorage:
@staticmethod
def store_realtime_simulation(
node_result_list: List[dict],
link_result_list: List[dict],
result_start_time: str,
db_name: str = None,
max_retries: int = 3,
):
"""存储实时模拟结果"""
for attempt in range(max_retries):
try:
with timescale_connection(db_name) as conn:
RealtimeRepository.store_realtime_simulation_result_sync(
conn, node_result_list, link_result_list, result_start_time
)
break # 成功
except Exception as e:
logger.error(f"存储尝试 {attempt + 1} 失败: {e}")
if attempt < max_retries - 1:
time.sleep(1) # 重试前等待
else:
raise # 达到最大重试次数后抛出异常
@staticmethod
def store_analysis_simulation(
run_id,
node_result_list: List[dict],
link_result_list: List[dict],
result_start_time: str,
num_periods: int = 1,
result_timestep_seconds: int | None = None,
db_name: str = None,
max_retries: int = 3,
):
"""Store immutable simulation results for one analysis run."""
for attempt in range(max_retries):
try:
with timescale_connection(db_name) as conn:
node_rows, link_rows = AnalysisResultsRepository.prepare_simulation_rows(
node_result_list, link_result_list, result_start_time,
num_periods, result_timestep_seconds or 3600,
)
AnalysisResultsRepository.store_results_sync(
conn, run_id, node_rows, link_rows
)
break # 成功
except Exception as e:
logger.error(f"存储尝试 {attempt + 1} 失败: {e}")
if attempt < max_retries - 1:
time.sleep(1) # 重试前等待
else:
raise # 达到最大重试次数后抛出异常
class InternalQueries:
@staticmethod
def query_scada_by_ids_time(
device_ids: List[str],
query_time: str,
db_name: str = None,
max_retries: int = 3,
) -> dict:
"""查询指定时间点的 SCADA 数据"""
target_time = parse_utc_time(query_time, field_name="query_time")
start_time = target_time - timedelta(seconds=1)
end_time = target_time + timedelta(seconds=1)
for attempt in range(max_retries):
try:
with timescale_connection(db_name) as conn:
rows = ScadaRepository.get_scada_by_ids_time_range_sync(
conn, device_ids, start_time, end_time
)
# Rows are ordered by device/time; retain the first sample
# for each requested device in one pass.
result = {device_id: None for device_id in device_ids}
seen: set[str] = set()
for row in rows:
device_id = str(row["device_id"])
if device_id in result and device_id not in seen:
result[device_id] = row["monitored_value"]
seen.add(device_id)
return result
except Exception as e:
logger.error(f"查询尝试 {attempt + 1} 失败: {e}")
if attempt < max_retries - 1:
time.sleep(1)
else:
raise
@staticmethod
def query_scada_by_ids_timerange(
device_ids: List[str],
start_time: str | datetime,
end_time: str | datetime,
db_name: str = None,
max_retries: int = 3,
) -> dict[str, list[dict]]:
"""查询指定时间窗的 SCADA 数据,返回 {device_id: [{time, value}, ...]}。"""
start_dt = parse_utc_time(start_time, field_name="start_time")
end_dt = parse_utc_time(end_time, field_name="end_time")
for attempt in range(max_retries):
try:
with timescale_connection(db_name) as conn:
rows = ScadaRepository.get_scada_by_ids_time_range_sync(
conn, device_ids, start_dt, end_dt
)
result: dict[str, list[dict]] = {
device_id: [] for device_id in device_ids
}
for row in rows:
device_id = row["device_id"]
value = row.get("cleaned_value")
if value is None:
value = row.get("monitored_value")
result.setdefault(device_id, []).append(
{"time": row["time"].isoformat(), "value": value}
)
return result
except Exception as e:
logger.error(f"查询尝试 {attempt + 1} 失败: {e}")
if attempt < max_retries - 1:
time.sleep(1)
else:
raise
@staticmethod
def query_latest_scada_time(
device_ids: List[str],
before_time: str | datetime | None = None,
db_name: str = None,
max_retries: int = 3,
) -> datetime | None:
"""Return the latest SCADA timestamp for the selected devices."""
before_dt = (
parse_utc_time(before_time, field_name="before_time")
if before_time is not None
else None
)
for attempt in range(max_retries):
try:
with timescale_connection(db_name) as conn:
return ScadaRepository.get_latest_scada_time_sync(
conn,
device_ids,
before_dt,
)
except Exception:
if attempt < max_retries - 1:
time.sleep(1)
else:
raise
@staticmethod
def query_realtime_simulation_by_ids_timerange(
element_ids: List[str],
start_time: str | datetime,
end_time: str | datetime,
element_type: str,
field: str,
db_name: str = None,
max_retries: int = 3,
) -> dict[str, list[dict]]:
"""查询实时模拟结果,返回 {id: [{time, value}, ...]}。"""
return InternalQueries._query_simulation_by_ids_timerange(
schema_name="realtime",
element_ids=element_ids,
start_time=start_time,
end_time=end_time,
element_type=element_type,
field=field,
db_name=db_name,
max_retries=max_retries,
)
@staticmethod
def query_analysis_simulation_by_ids_timerange(
element_ids: List[str],
start_time: str | datetime,
end_time: str | datetime,
element_type: str,
field: str,
run_id,
db_name: str = None,
max_retries: int = 3,
) -> dict[str, list[dict]]:
"""Query one analysis run, returning {id: [{time, value}, ...]}."""
return InternalQueries._query_simulation_by_ids_timerange(
schema_name="analysis",
element_ids=element_ids,
start_time=start_time,
end_time=end_time,
element_type=element_type,
field=field,
db_name=db_name,
max_retries=max_retries,
run_id=run_id,
)
@staticmethod
def _query_simulation_by_ids_timerange(
*,
schema_name: str,
element_ids: List[str],
start_time: str | datetime,
end_time: str | datetime,
element_type: str,
field: str,
db_name: str = None,
max_retries: int = 3,
run_id=None,
) -> dict[str, list[dict]]:
normalized_element_ids = list(
dict.fromkeys(
normalized
for normalized in (str(element_id).strip() for element_id in element_ids)
if normalized
)
)
if not normalized_element_ids:
return {}
start_dt = parse_utc_time(start_time, field_name="start_time")
end_dt = parse_utc_time(end_time, field_name="end_time")
table_name, id_column, valid_fields = InternalQueries._resolve_simulation_table(element_type)
if field not in valid_fields:
raise ValueError(f"Invalid field for {element_type}: {field}")
if schema_name not in {"realtime", "analysis"}:
raise ValueError(f"Unsupported schema_name: {schema_name}")
if schema_name == "analysis" and run_id is None:
raise ValueError("analysis query requires run_id")
for attempt in range(max_retries):
try:
with timescale_connection(db_name) as conn:
with conn.cursor(row_factory=dict_row) as cur:
if schema_name == "analysis":
query = sql.SQL(
"SELECT btrim({}::text) AS id, time, {} FROM {}.{} "
"WHERE run_id = %s AND time >= %s AND time <= %s "
"AND btrim({}::text) = ANY(%s) ORDER BY id, time"
).format(
sql.Identifier(id_column),
sql.Identifier(field),
sql.Identifier(schema_name),
sql.Identifier(table_name),
sql.Identifier(id_column),
)
cur.execute(
query,
(run_id, start_dt, end_dt, normalized_element_ids),
)
else:
query = sql.SQL(
"SELECT btrim({}::text) AS id, time, {} FROM {}.{} "
"WHERE time >= %s AND time <= %s "
"AND btrim({}::text) = ANY(%s) ORDER BY id, time"
).format(
sql.Identifier(id_column),
sql.Identifier(field),
sql.Identifier(schema_name),
sql.Identifier(table_name),
sql.Identifier(id_column),
)
cur.execute(query, (start_dt, end_dt, normalized_element_ids))
rows = cur.fetchall()
result: dict[str, list[dict]] = {
element_id: [] for element_id in normalized_element_ids
}
for row in rows:
element_id = str(row["id"]).strip()
result.setdefault(element_id, []).append(
{"time": row["time"].isoformat(), "value": row[field]}
)
for element_id in result:
result[element_id].sort(key=lambda item: item["time"])
return result
except Exception as e:
logger.error(f"查询尝试 {attempt + 1} 失败: {e}")
if attempt < max_retries - 1:
time.sleep(1)
else:
raise
@staticmethod
def _resolve_simulation_table(element_type: str) -> tuple[str, str, set[str]]:
normalized_type = element_type.lower()
if normalized_type == "node":
return "node_results", "node_id", set(RealtimeRepository.NODE_FIELDS)
if normalized_type == "link":
return "link_results", "link_id", set(RealtimeRepository.LINK_FIELDS)
raise ValueError(f"Unsupported element_type: {element_type}")