From d33619d0d576c5c6a2b32d2317913671f2a3d55b Mon Sep 17 00:00:00 2001 From: Huarch Date: Thu, 9 Jul 2026 11:51:40 +0800 Subject: [PATCH] fix(burst-location): correct normal data window --- .../burst_location/burst_location.py | 2 +- app/services/burst_location.py | 47 +++--- tests/unit/test_burst_location_service.py | 153 +++++++++++++++--- 3 files changed, 162 insertions(+), 40 deletions(-) diff --git a/app/algorithms/burst_location/burst_location.py b/app/algorithms/burst_location/burst_location.py index 4f4f971..54a0130 100644 --- a/app/algorithms/burst_location/burst_location.py +++ b/app/algorithms/burst_location/burst_location.py @@ -121,7 +121,7 @@ def run_burst_location( basic_pressure: float = 10.0, n_workers: int = DEFAULT_N_WORKERS, partition_on_full_graph: bool = True, - visualize_partition: bool = True, + visualize_partition: bool = False, visualize_pause_seconds: float = 0.3, final_candidates_csv_path: ( str | None diff --git a/app/services/burst_location.py b/app/services/burst_location.py index 4feffec..a1bcf73 100644 --- a/app/services/burst_location.py +++ b/app/services/burst_location.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from datetime import datetime +from datetime import datetime, timedelta from typing import Any import pandas as pd @@ -124,6 +124,8 @@ def run_burst_location_by_network( if normalized_data_source == "simulation": if not simulation_scheme_name: raise ValueError("模拟方案模式必须提供 simulation_scheme_name。") + normal_start_dt = burst_start_dt + normal_end_dt = burst_end_dt ( burst_pressure_series, burst_pressure_samples, @@ -144,16 +146,21 @@ def run_burst_location_by_network( ) = _build_observed_series_from_simulation( network=network, sensor_ids=selected_pressure_ids, - start_dt=normal_start_dt or burst_start_dt, - end_dt=normal_end_dt or burst_end_dt, + start_dt=normal_start_dt, + end_dt=normal_end_dt, data_type="pressure", series_name="normal_pressure", - simulation_source="scheme", - simulation_scheme_name=simulation_scheme_name, + simulation_source="realtime", + simulation_scheme_name=None, simulation_scheme_type=resolved_simulation_scheme_type, ) - observed_source = "simulation_scheme_timerange" + observed_source = "simulation_scheme_burst_realtime_normal_timerange" else: + if normal_pressure_from_payload is None and ( + normal_start_dt is None or normal_end_dt is None + ): + normal_start_dt = burst_start_dt - timedelta(days=1) + normal_end_dt = burst_end_dt - timedelta(days=1) ( burst_pressure_series, burst_pressure_samples, @@ -172,16 +179,12 @@ def run_burst_location_by_network( ) = _build_observed_series_from_scada( network=network, sensor_ids=selected_pressure_ids, - start_dt=normal_start_dt or burst_start_dt, - end_dt=normal_end_dt or burst_end_dt, + start_dt=normal_start_dt, + end_dt=normal_end_dt, data_type="pressure", series_name="normal_pressure", ) - observed_source = ( - "scada_burst_scada_normal_timerange" - if normal_start_dt is not None and normal_end_dt is not None - else "scada_timerange" - ) + observed_source = "scada_burst_scada_normal_timerange" else: normal_pressure_series = normal_pressure_from_payload normal_pressure_samples = 1 @@ -245,16 +248,21 @@ def run_burst_location_by_network( _build_observed_series_from_simulation( network=network, sensor_ids=selected_flow_ids, - start_dt=normal_start_dt or burst_start_dt, - end_dt=normal_end_dt or burst_end_dt, + start_dt=normal_start_dt, + end_dt=normal_end_dt, data_type="flow", series_name="normal_flow", - simulation_source="scheme", - simulation_scheme_name=simulation_scheme_name, + simulation_source="realtime", + simulation_scheme_name=None, simulation_scheme_type=resolved_simulation_scheme_type, ) ) else: + if normal_flow_from_payload is None and ( + normal_start_dt is None or normal_end_dt is None + ): + normal_start_dt = burst_start_dt - timedelta(days=1) + normal_end_dt = burst_end_dt - timedelta(days=1) burst_flow_series, burst_flow_samples = _build_observed_series_from_scada( network=network, sensor_ids=selected_flow_ids, @@ -268,8 +276,8 @@ def run_burst_location_by_network( _build_observed_series_from_scada( network=network, sensor_ids=selected_flow_ids, - start_dt=normal_start_dt or burst_start_dt, - end_dt=normal_end_dt or burst_end_dt, + start_dt=normal_start_dt, + end_dt=normal_end_dt, data_type="flow", series_name="normal_flow", ) @@ -313,6 +321,7 @@ def run_burst_location_by_network( normal_flow=normal_flow_series, min_dpressure=min_dpressure, basic_pressure=basic_pressure, + visualize_partition=False, ) payload: dict[str, Any] = { diff --git a/tests/unit/test_burst_location_service.py b/tests/unit/test_burst_location_service.py index 3ab7ba2..4ca9af5 100644 --- a/tests/unit/test_burst_location_service.py +++ b/tests/unit/test_burst_location_service.py @@ -190,29 +190,41 @@ def test_run_burst_location_uses_single_timerange_with_burst_source_split(monkey use_scada_flow=True, ) - assert result["observed_source"] == "simulation_scheme_timerange" + assert result["observed_source"] == "simulation_scheme_burst_realtime_normal_timerange" assert result["simulation_scheme"] == { "name": "BurstSchemeA", "type": "burst_analysis", } assert result["pressure_samples"] == {"burst": 4, "normal": 4} assert result["flow_samples"] == {"burst": 4, "normal": 4} + assert captured["visualize_partition"] is False assert list(captured["burst_pressure"].index) == ["J1"] assert captured["burst_pressure"]["J1"] == pytest.approx(15.0) - assert captured["normal_pressure"]["J1"] == pytest.approx(15.0) + assert captured["normal_pressure"]["J1"] == pytest.approx(11.0) assert captured["burst_flow"]["J2"] == pytest.approx(6.0) assert captured["burst_flow"]["P1"] == pytest.approx(8.0) - assert captured["normal_flow"]["J2"] == pytest.approx(6.0) - assert captured["normal_flow"]["P1"] == pytest.approx(8.0) + assert captured["normal_flow"]["J2"] == pytest.approx(4.0) + assert captured["normal_flow"]["P1"] == pytest.approx(5.0) assert all(call["scheme_name"] == "BurstSchemeA" for call in scheme_calls) - assert len(scheme_calls) == 6 + assert len(scheme_calls) == 3 assert any(call["element_type"] == "node" and call["field"] == "pressure" for call in scheme_calls) assert any(call["element_type"] == "link" and call["field"] == "flow" for call in scheme_calls) assert any(call["element_type"] == "node" and call["field"] == "actual_demand" for call in scheme_calls) - assert realtime_calls == [] + assert len(realtime_calls) == 3 + assert any(call["element_type"] == "node" and call["field"] == "pressure" 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 {call["start_time"] for call in scheme_calls + realtime_calls} == { + "2025-01-01T08:00:00+00:00" + } + assert {call["end_time"] for call in scheme_calls + realtime_calls} == { + "2025-01-01T09:00:00+00:00" + } assert result["scada_window"] == { "burst_start": "2025-01-01T08:00:00+00:00", "burst_end": "2025-01-01T09:00:00+00:00", + "normal_start": "2025-01-01T08:00:00+00:00", + "normal_end": "2025-01-01T09:00:00+00:00", } @@ -452,7 +464,7 @@ def test_run_burst_location_monitoring_uses_scada_for_burst_and_normal( } -def test_run_burst_location_monitoring_reuses_burst_window_for_normal( +def test_run_burst_location_monitoring_defaults_normal_window_to_previous_day( monkeypatch, tmp_path ): module = _load_burst_location_module() @@ -476,18 +488,26 @@ def test_run_burst_location_monitoring_reuses_burst_window_for_normal( "run_burst_location", lambda **kwargs: captured.update(kwargs) or {"located_pipe": "Pipe-001"}, ) + + def fake_scada_query(**kwargs): + scada_calls.append(kwargs) + start_time = datetime.fromisoformat(kwargs["start_time"]) + values = ( + [20.0, 22.0] + if start_time.date().isoformat() == "2025-01-01" + else [10.0, 12.0] + ) + return { + "pressure-query": [ + {"time": kwargs["start_time"], "value": values[0]}, + {"time": kwargs["end_time"], "value": values[1]}, + ] + } + monkeypatch.setattr( module.InternalQueries, "query_scada_by_ids_timerange", - staticmethod( - lambda **kwargs: scada_calls.append(kwargs) - or { - "pressure-query": [ - {"time": kwargs["start_time"], "value": 20.0}, - {"time": kwargs["end_time"], "value": 22.0}, - ] - } - ), + staticmethod(fake_scada_query), ) monkeypatch.setattr( module.InternalQueries, @@ -504,12 +524,105 @@ def test_run_burst_location_monitoring_reuses_burst_window_for_normal( scada_burst_end=datetime(2025, 1, 1, 9, 0, 0), ) - assert result["observed_source"] == "scada_timerange" + assert result["observed_source"] == "scada_burst_scada_normal_timerange" assert len(scada_calls) == 2 - assert scada_calls[0]["start_time"] == scada_calls[1]["start_time"] - assert scada_calls[0]["end_time"] == scada_calls[1]["end_time"] + assert datetime.fromisoformat(scada_calls[1]["start_time"]) == ( + datetime.fromisoformat(scada_calls[0]["start_time"]) - timedelta(days=1) + ) + assert datetime.fromisoformat(scada_calls[1]["end_time"]) == ( + datetime.fromisoformat(scada_calls[0]["end_time"]) - timedelta(days=1) + ) assert captured["burst_pressure"]["J1"] == pytest.approx(21.0) - assert captured["normal_pressure"]["J1"] == pytest.approx(21.0) + assert captured["normal_pressure"]["J1"] == pytest.approx(11.0) + assert result["scada_window"] == { + "burst_start": "2025-01-01T08:00:00+00:00", + "burst_end": "2025-01-01T09:00:00+00:00", + "normal_start": "2024-12-31T08:00:00+00:00", + "normal_end": "2024-12-31T09:00:00+00:00", + } + + +def test_run_burst_location_monitoring_flow_uses_previous_day_normal_window( + monkeypatch, tmp_path +): + module = _load_burst_location_module() + captured = {} + scada_calls = [] + + monkeypatch.setattr( + module, + "get_all_scada_info", + lambda network: [ + { + "type": "pressure", + "associated_element_id": "J1", + "api_query_id": "pressure-query", + }, + { + "type": "pipe_flow", + "associated_element_id": "P1", + "api_query_id": "flow-query", + }, + ], + ) + monkeypatch.setattr(module, "_prepare_burst_inp", lambda network: str(tmp_path / "fake.inp")) + monkeypatch.setattr( + module, + "run_burst_location", + lambda **kwargs: captured.update(kwargs) or {"located_pipe": "Pipe-001"}, + ) + + def fake_scada_query(**kwargs): + scada_calls.append(kwargs) + is_burst_day = ( + datetime.fromisoformat(kwargs["start_time"]).date().isoformat() + == "2025-01-01" + ) + if kwargs["device_ids"] == ["pressure-query"]: + values = [20.0, 22.0] if is_burst_day else [10.0, 12.0] + query_id = "pressure-query" + else: + values = [7.0, 9.0] if is_burst_day else [3.0, 5.0] + query_id = "flow-query" + return { + query_id: [ + {"time": kwargs["start_time"], "value": values[0]}, + {"time": kwargs["end_time"], "value": values[1]}, + ] + } + + monkeypatch.setattr( + module.InternalQueries, + "query_scada_by_ids_timerange", + staticmethod(fake_scada_query), + ) + + result = module.run_burst_location_by_network( + network="tjwater", + username="testuser", + data_source="monitoring", + burst_leakage=1.0, + scada_burst_start=datetime(2025, 1, 1, 8, 0, 0), + scada_burst_end=datetime(2025, 1, 1, 9, 0, 0), + use_scada_flow=True, + ) + + assert result["observed_source"] == "scada_burst_scada_normal_timerange" + assert len(scada_calls) == 4 + for burst_call, normal_call in [ + (scada_calls[0], scada_calls[1]), + (scada_calls[2], scada_calls[3]), + ]: + assert datetime.fromisoformat(normal_call["start_time"]) == ( + datetime.fromisoformat(burst_call["start_time"]) - timedelta(days=1) + ) + assert datetime.fromisoformat(normal_call["end_time"]) == ( + datetime.fromisoformat(burst_call["end_time"]) - timedelta(days=1) + ) + assert captured["burst_pressure"]["J1"] == pytest.approx(21.0) + assert captured["normal_pressure"]["J1"] == pytest.approx(11.0) + assert captured["burst_flow"]["P1"] == pytest.approx(8.0) + assert captured["normal_flow"]["P1"] == pytest.approx(4.0) def test_run_burst_location_monitoring_aligns_partial_scada_data(