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}")