fix(burst-location): normalize simulation ids

This commit is contained in:
2026-07-08 17:17:41 +08:00
parent c7947a7481
commit 18943314f8
3 changed files with 121 additions and 22 deletions
+17 -9
View File
@@ -229,7 +229,14 @@ class InternalQueries:
scheme_type: str | None = None, scheme_type: str | None = None,
scheme_name: str | None = None, scheme_name: str | None = None,
) -> dict[str, list[dict]]: ) -> dict[str, list[dict]]:
if not element_ids: 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 {} return {}
start_dt = parse_utc_time(start_time, field_name="start_time") start_dt = parse_utc_time(start_time, field_name="start_time")
@@ -253,9 +260,9 @@ class InternalQueries:
with conn.cursor(row_factory=dict_row) as cur: with conn.cursor(row_factory=dict_row) as cur:
if schema_name == "scheme": if schema_name == "scheme":
query = sql.SQL( query = sql.SQL(
"SELECT id, time, {} FROM {}.{} " "SELECT btrim(id::text) AS id, time, {} FROM {}.{} "
"WHERE scheme_type = %s AND scheme_name = %s " "WHERE scheme_type = %s AND scheme_name = %s "
"AND time >= %s AND time <= %s AND id = ANY(%s)" "AND time >= %s AND time <= %s AND btrim(id::text) = ANY(%s)"
).format( ).format(
sql.Identifier(field), sql.Identifier(field),
sql.Identifier(schema_name), sql.Identifier(schema_name),
@@ -268,25 +275,26 @@ class InternalQueries:
scheme_name, scheme_name,
start_dt, start_dt,
end_dt, end_dt,
element_ids, normalized_element_ids,
), ),
) )
else: else:
query = sql.SQL( query = sql.SQL(
"SELECT id, time, {} FROM {}.{} " "SELECT btrim(id::text) AS id, time, {} FROM {}.{} "
"WHERE time >= %s AND time <= %s AND id = ANY(%s)" "WHERE time >= %s AND time <= %s AND btrim(id::text) = ANY(%s)"
).format( ).format(
sql.Identifier(field), sql.Identifier(field),
sql.Identifier(schema_name), sql.Identifier(schema_name),
sql.Identifier(table_name), sql.Identifier(table_name),
) )
cur.execute(query, (start_dt, end_dt, element_ids)) cur.execute(query, (start_dt, end_dt, normalized_element_ids))
rows = cur.fetchall() rows = cur.fetchall()
result: dict[str, list[dict]] = { result: dict[str, list[dict]] = {
element_id: [] for element_id in element_ids element_id: [] for element_id in normalized_element_ids
} }
for row in rows: for row in rows:
result.setdefault(row["id"], []).append( element_id = str(row["id"]).strip()
result.setdefault(element_id, []).append(
{"time": row["time"].isoformat(), "value": row[field]} {"time": row["time"].isoformat(), "value": row[field]}
) )
for element_id in result: for element_id in result:
+35 -10
View File
@@ -40,7 +40,7 @@ def _normalize_series(data: SeriesInput, field_name: str) -> pd.Series:
else: else:
raise ValueError(f"Unsupported data format for {field_name}.") raise ValueError(f"Unsupported data format for {field_name}.")
series.index = series.index.map(str) series.index = series.index.map(_normalize_identifier)
return pd.to_numeric(series, errors="raise") return pd.to_numeric(series, errors="raise")
@@ -385,6 +385,7 @@ def _build_observed_series_from_scada(
data_type: str, data_type: str,
series_name: str, series_name: str,
) -> tuple[pd.Series, int]: ) -> tuple[pd.Series, int]:
sensor_ids = _dedupe_ids(sensor_ids)
scada_mapping = _build_scada_mapping(network=network, data_type=data_type) scada_mapping = _build_scada_mapping(network=network, data_type=data_type)
missing_ids = [ missing_ids = [
sensor_id for sensor_id in sensor_ids if sensor_id not in scada_mapping sensor_id for sensor_id in sensor_ids if sensor_id not in scada_mapping
@@ -400,6 +401,7 @@ def _build_observed_series_from_scada(
start_time=start_dt.isoformat(), start_time=start_dt.isoformat(),
end_time=end_dt.isoformat(), end_time=end_dt.isoformat(),
) )
scada_data = _normalize_timeseries_by_id(scada_data)
values: dict[str, float] = {} values: dict[str, float] = {}
sample_counts: list[int] = [] sample_counts: list[int] = []
for sensor_id, query_id in zip(sensor_ids, query_ids): for sensor_id, query_id in zip(sensor_ids, query_ids):
@@ -427,6 +429,7 @@ def _build_observed_series_from_simulation(
simulation_scheme_name: str | None, simulation_scheme_name: str | None,
simulation_scheme_type: str, simulation_scheme_type: str,
) -> tuple[pd.Series, int]: ) -> tuple[pd.Series, int]:
sensor_ids = _dedupe_ids(sensor_ids)
sensor_metadata = _build_sensor_metadata(network=network, data_type=data_type) sensor_metadata = _build_sensor_metadata(network=network, data_type=data_type)
missing_ids = [ missing_ids = [
sensor_id for sensor_id in sensor_ids if sensor_id not in sensor_metadata sensor_id for sensor_id in sensor_ids if sensor_id not in sensor_metadata
@@ -446,6 +449,7 @@ def _build_observed_series_from_simulation(
simulation_scheme_name=simulation_scheme_name, simulation_scheme_name=simulation_scheme_name,
simulation_scheme_type=simulation_scheme_type, simulation_scheme_type=simulation_scheme_type,
) )
simulation_data = _normalize_timeseries_by_id(simulation_data)
values: dict[str, float] = {} values: dict[str, float] = {}
sample_counts: list[int] = [] sample_counts: list[int] = []
for sensor_id in sensor_ids: for sensor_id in sensor_ids:
@@ -476,6 +480,7 @@ def _query_simulation_data_by_sensor_ids(
if simulation_source not in {"scheme", "realtime"}: if simulation_source not in {"scheme", "realtime"}:
raise ValueError(f"Unsupported simulation_source: {simulation_source}") raise ValueError(f"Unsupported simulation_source: {simulation_source}")
sensor_ids = _dedupe_ids(sensor_ids)
result: dict[str, list[dict[str, Any]]] = { result: dict[str, list[dict[str, Any]]] = {
sensor_id: [] for sensor_id in sensor_ids sensor_id: [] for sensor_id in sensor_ids
} }
@@ -556,6 +561,7 @@ def _query_simulation_values(
simulation_scheme_name: str | None, simulation_scheme_name: str | None,
simulation_scheme_type: str, simulation_scheme_type: str,
) -> dict[str, list[dict[str, Any]]]: ) -> dict[str, list[dict[str, Any]]]:
element_ids = _dedupe_ids(element_ids)
if not element_ids: if not element_ids:
return {} return {}
if simulation_source == "scheme": if simulation_source == "scheme":
@@ -595,14 +601,9 @@ def _build_sensor_metadata(network: str, data_type: str) -> dict[str, dict[str,
continue continue
else: else:
raise ValueError(f"Unsupported data_type: {data_type}") raise ValueError(f"Unsupported data_type: {data_type}")
element_id = item.get("associated_element_id") element_id = _normalize_identifier(item.get("associated_element_id"))
query_id = item.get("api_query_id") query_id = _normalize_identifier(item.get("api_query_id"))
if ( if element_id and query_id:
isinstance(element_id, str)
and element_id
and isinstance(query_id, str)
and query_id
):
metadata[element_id] = {"query_id": query_id, "scada_type": scada_type} metadata[element_id] = {"query_id": query_id, "scada_type": scada_type}
return metadata return metadata
@@ -638,7 +639,31 @@ def _get_sensor_nodes(network: str, data_type: str) -> list[str]:
def _dedupe_ids(ids: list[str] | None) -> list[str]: def _dedupe_ids(ids: list[str] | None) -> list[str]:
if ids is None: if ids is None:
return [] return []
return list(dict.fromkeys([str(item) for item in ids if item])) return list(
dict.fromkeys(
normalized
for normalized in (_normalize_identifier(item) for item in ids)
if normalized
)
)
def _normalize_identifier(value: Any) -> str:
if value is None:
return ""
return str(value).strip()
def _normalize_timeseries_by_id(
data: dict[Any, list[dict[str, Any]]] | None,
) -> dict[str, list[dict[str, Any]]]:
normalized_data: dict[str, list[dict[str, Any]]] = {}
for raw_id, records in (data or {}).items():
normalized_id = _normalize_identifier(raw_id)
if not normalized_id:
continue
normalized_data.setdefault(normalized_id, []).extend(records or [])
return normalized_data
def _to_datetime(value: datetime | str) -> datetime: def _to_datetime(value: datetime | str) -> datetime:
+69 -3
View File
@@ -1,7 +1,7 @@
import importlib.util import importlib.util
import sys import sys
import types import types
from datetime import datetime, timedelta from datetime import datetime, timedelta, timezone
from pathlib import Path from pathlib import Path
import pytest import pytest
@@ -30,6 +30,24 @@ def _load_burst_location_module():
]: ]:
ensure_package(package_name) ensure_package(package_name)
def parse_utc_time(value, field_name="datetime"):
parsed = value if isinstance(value, datetime) else datetime.fromisoformat(value)
if parsed.tzinfo is None:
return parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
time_api_module = types.ModuleType("app.services.time_api")
time_api_module.parse_utc_time = parse_utc_time
time_api_module.extract_date = (
lambda value, field_name="date": (
value.date()
if isinstance(value, datetime)
else datetime.fromisoformat(value).date()
)
)
time_api_module.utc_now = lambda: datetime.now(timezone.utc)
sys.modules["app.services.time_api"] = time_api_module
algorithms_module = types.ModuleType("app.algorithms.burst_location") algorithms_module = types.ModuleType("app.algorithms.burst_location")
algorithms_module.run_burst_location = lambda **kwargs: {} algorithms_module.run_burst_location = lambda **kwargs: {}
sys.modules["app.algorithms.burst_location"] = algorithms_module sys.modules["app.algorithms.burst_location"] = algorithms_module
@@ -198,8 +216,8 @@ def test_run_burst_location_uses_single_timerange_with_burst_source_split(monkey
assert any(call["element_type"] == "link" and call["field"] == "flow" for call in realtime_calls) assert any(call["element_type"] == "link" and call["field"] == "flow" for call in realtime_calls)
assert any(call["element_type"] == "node" and call["field"] == "actual_demand" for call in realtime_calls) assert any(call["element_type"] == "node" and call["field"] == "actual_demand" for call in realtime_calls)
assert result["scada_window"] == { assert result["scada_window"] == {
"burst_start": "2025-01-01T08:00:00", "burst_start": "2025-01-01T08:00:00+00:00",
"burst_end": "2025-01-01T09:00:00", "burst_end": "2025-01-01T09:00:00+00:00",
} }
@@ -230,6 +248,54 @@ def test_run_burst_location_requires_simulation_scheme_name(monkeypatch, tmp_pat
) )
def test_build_observed_series_from_simulation_normalizes_result_ids(monkeypatch):
module = _load_burst_location_module()
query_calls = []
monkeypatch.setattr(
module,
"get_all_scada_info",
lambda network: [
{
"type": "pressure",
"associated_element_id": " 100026 ",
"api_query_id": " pressure-query ",
}
],
)
def fake_scheme_query(**kwargs):
query_calls.append(kwargs)
return {
100026: [
{"time": kwargs["start_time"], "value": 10.0},
{"time": kwargs["end_time"], "value": 14.0},
]
}
monkeypatch.setattr(
module.InternalQueries,
"query_scheme_simulation_by_ids_timerange",
staticmethod(fake_scheme_query),
)
series, sample_count = module._build_observed_series_from_simulation(
network="tjwater",
sensor_ids=["100026"],
start_dt=datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc),
end_dt=datetime(2025, 1, 1, 1, 0, 0, tzinfo=timezone.utc),
data_type="pressure",
series_name="burst_pressure",
simulation_source="scheme",
simulation_scheme_name="BurstSchemeA",
simulation_scheme_type="burst_analysis",
)
assert query_calls[0]["element_ids"] == ["100026"]
assert sample_count == 2
assert series["100026"] == pytest.approx(12.0)
def test_run_burst_location_monitoring_uses_scada_for_burst_and_realtime_for_normal( def test_run_burst_location_monitoring_uses_scada_for_burst_and_realtime_for_normal(
monkeypatch, tmp_path monkeypatch, tmp_path
): ):