Files
TJWaterServerBinary/app/infra/db/timescaledb/internal_queries.py
T
jiang 5966d039de refactor(backend)!: separate algorithm and data layers
Reorganize algorithm packages by business responsibility, move orchestration into services, and keep database access behind pooled repositories.

Harden analysis API validation, remove unsafe legacy simulation endpoints, and add regression and architecture boundary coverage.

BREAKING CHANGE: legacy algorithm module paths and obsolete simulation endpoints are removed.
2026-09-04 17:30:55 +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.domain.time 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}")