diff --git a/app/services/burst_location.py b/app/services/burst_location.py index 3925b80..4feffec 100644 --- a/app/services/burst_location.py +++ b/app/services/burst_location.py @@ -186,6 +186,14 @@ def run_burst_location_by_network( normal_pressure_series = normal_pressure_from_payload normal_pressure_samples = 1 observed_source = "scada_burst_payload_normal_timerange" + selected_pressure_ids, burst_pressure_series, normal_pressure_series = ( + _align_observed_series_pair( + ids=selected_pressure_ids, + burst_series=burst_pressure_series, + normal_series=normal_pressure_series, + data_label="压力数据", + ) + ) else: if burst_pressure is None or normal_pressure is None: raise ValueError( @@ -269,6 +277,14 @@ def run_burst_location_by_network( else: normal_flow_series = normal_flow_from_payload normal_flow_samples = 1 + selected_flow_ids, burst_flow_series, normal_flow_series = ( + _align_observed_series_pair( + ids=selected_flow_ids, + burst_series=burst_flow_series, + normal_series=normal_flow_series, + data_label="流量数据", + ) + ) else: if flow_scada_ids is not None: selected_flow_ids = _dedupe_ids(flow_scada_ids) @@ -439,6 +455,23 @@ def _validate_time_window( return start_dt, end_dt +def _align_observed_series_pair( + *, + ids: list[str], + burst_series: pd.Series, + normal_series: pd.Series, + data_label: str, +) -> tuple[list[str], pd.Series, pd.Series]: + common_ids = [ + sensor_id + for sensor_id in _dedupe_ids(ids) + if sensor_id in burst_series.index and sensor_id in normal_series.index + ] + if not common_ids: + raise ValueError(f"{data_label}没有同时具备爆管时段和正常时段有效数据的点位。") + return common_ids, burst_series.loc[common_ids], normal_series.loc[common_ids] + + def _build_observed_series_from_scada( *, network: str, @@ -473,11 +506,13 @@ def _build_observed_series_from_scada( float(item["value"]) for item in records if item.get("value") is not None ] if not numeric_values: - raise ValueError( - f"{_series_display_name(series_name)} 在时间窗内无有效数据: {sensor_id}" - ) + continue values[sensor_id] = float(sum(numeric_values) / len(numeric_values)) sample_counts.append(len(numeric_values)) + if not values: + raise ValueError( + f"{_series_display_name(series_name)} 在时间窗内无有效数据: {', '.join(sensor_ids[:10])}" + ) return pd.Series(values, dtype=float), min(sample_counts) diff --git a/tests/unit/test_burst_location_service.py b/tests/unit/test_burst_location_service.py index a5f00e0..ba08a76 100644 --- a/tests/unit/test_burst_location_service.py +++ b/tests/unit/test_burst_location_service.py @@ -326,6 +326,51 @@ def test_build_observed_series_from_scada_uses_chinese_error_label(monkeypatch): assert "burst_pressure" not in message +def test_build_observed_series_from_scada_skips_missing_sensor_values(monkeypatch): + module = _load_burst_location_module() + + monkeypatch.setattr( + module, + "get_all_scada_info", + lambda network: [ + {"type": "pressure", "associated_element_id": "J1", "api_query_id": "q1"}, + {"type": "pressure", "associated_element_id": "J2", "api_query_id": "q2"}, + {"type": "pressure", "associated_element_id": "J3", "api_query_id": "q3"}, + ], + ) + monkeypatch.setattr( + module.InternalQueries, + "query_scada_by_ids_timerange", + staticmethod( + lambda **kwargs: { + "q1": [ + {"time": kwargs["start_time"], "value": 10.0}, + {"time": kwargs["end_time"], "value": 12.0}, + ], + "q2": [], + "q3": [ + {"time": kwargs["start_time"], "value": None}, + {"time": kwargs["end_time"], "value": 18.0}, + ], + } + ), + ) + + series, sample_count = module._build_observed_series_from_scada( + network="tjwater", + sensor_ids=["J1", "J2", "J3"], + 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", + ) + + assert list(series.index) == ["J1", "J3"] + assert series["J1"] == pytest.approx(11.0) + assert series["J3"] == pytest.approx(18.0) + assert sample_count == 1 + + def test_run_burst_location_monitoring_uses_scada_for_burst_and_normal( monkeypatch, tmp_path ): @@ -467,3 +512,66 @@ def test_run_burst_location_monitoring_reuses_burst_window_for_normal( assert scada_calls[0]["end_time"] == scada_calls[1]["end_time"] assert captured["burst_pressure"]["J1"] == pytest.approx(21.0) assert captured["normal_pressure"]["J1"] == pytest.approx(21.0) + + +def test_run_burst_location_monitoring_aligns_partial_scada_data( + monkeypatch, tmp_path +): + module = _load_burst_location_module() + captured = {} + + monkeypatch.setattr( + module, + "get_all_scada_info", + lambda network: [ + {"type": "pressure", "associated_element_id": "J1", "api_query_id": "q1"}, + {"type": "pressure", "associated_element_id": "J2", "api_query_id": "q2"}, + {"type": "pressure", "associated_element_id": "J3", "api_query_id": "q3"}, + ], + ) + 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): + start_hour = datetime.fromisoformat(kwargs["start_time"]).astimezone( + timezone(timedelta(hours=8)) + ).hour + if start_hour == 8: + return { + "q1": [{"time": kwargs["start_time"], "value": 20.0}], + "q2": [{"time": kwargs["start_time"], "value": 30.0}], + "q3": [], + } + return { + "q1": [{"time": kwargs["start_time"], "value": 10.0}], + "q2": [], + "q3": [{"time": kwargs["start_time"], "value": 12.0}], + } + + 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, tzinfo=timezone(timedelta(hours=8))), + scada_burst_end=datetime(2025, 1, 1, 9, 0, 0, tzinfo=timezone(timedelta(hours=8))), + scada_normal_start=datetime(2025, 1, 1, 7, 0, 0, tzinfo=timezone(timedelta(hours=8))), + scada_normal_end=datetime(2025, 1, 1, 8, 0, 0, tzinfo=timezone(timedelta(hours=8))), + ) + + assert result["pressure_scada_ids"] == ["J1"] + assert captured["pressure_scada_ids"] == ["J1"] + assert list(captured["burst_pressure"].index) == ["J1"] + assert list(captured["normal_pressure"].index) == ["J1"] + assert captured["burst_pressure"]["J1"] == pytest.approx(20.0) + assert captured["normal_pressure"]["J1"] == pytest.approx(10.0)