Files
InterferenceETL/tests/test_pipeline.py
T

275 lines
10 KiB
Python

from __future__ import annotations
import csv
from datetime import datetime
import io
import json
from pathlib import Path
import tempfile
import types
import unittest
from unittest.mock import patch
import zipfile
from openpyxl import Workbook
import main
COMPLETE_WINDOW = "2026073110001100"
INCOMPLETE_WINDOW = "2026073111001200"
class PipelineTest(unittest.TestCase):
def test_latest_complete_window_is_converted_and_merged(self) -> None:
with tempfile.TemporaryDirectory() as temp:
root = Path(temp) / "source"
output = Path(temp) / "output"
for source_type in main.EXPECTED_TYPES:
create_archive(root, source_type, COMPLETE_WINDOW)
create_archive(root, main.EXPECTED_TYPES[0], INCOMPLETE_WINDOW)
store = RecordingStore()
result = main.process(main.LocalSource(root), output, lookback_days=3, store=store)
self.assertEqual(result.name, COMPLETE_WINDOW)
converted = sorted((result / "converted").glob("*.csv"))
self.assertEqual(len(converted), 7)
with (result / f"interference_summary_{COMPLETE_WINDOW}.csv").open(encoding="utf-8-sig", newline="") as file:
rows = list(csv.DictReader(file))
self.assertEqual(len(rows), 7)
self.assertEqual({row["source_type"] for row in rows}, set(main.EXPECTED_TYPES))
rows_by_type = {row["source_type"]: row for row in rows}
self.assertEqual(rows_by_type["5G干扰监控"]["cgi"], "460-00-200-1")
self.assertEqual(rows_by_type["700M干扰监控"]["cgi"], "460-00-200-1")
for source_type in set(main.EXPECTED_TYPES) - {"5G干扰监控", "700M干扰监控"}:
self.assertEqual(rows_by_type[source_type]["cgi"], "460-00-100-1")
self.assertTrue(all(row["interference_dbm"] == "-100.5" for row in rows))
manifest = json.loads((result / "manifest.json").read_text(encoding="utf-8"))
self.assertEqual(manifest["source_file_count"], 7)
self.assertEqual(manifest["summary_rows"], 7)
self.assertEqual(len(manifest["warnings"]), 1)
self.assertIn(INCOMPLETE_WINDOW, manifest["warnings"][0])
self.assertEqual(manifest["database"]["metric_time"], "2026-07-31 10:00:00")
self.assertEqual(len(store.rows), 7)
def test_requested_incomplete_window_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as temp:
root = Path(temp) / "source"
create_archive(root, main.EXPECTED_TYPES[0], INCOMPLETE_WINDOW)
with self.assertRaisesRegex(main.ProcessingError, "Requested window is incomplete"):
main.process(main.LocalSource(root), Path(temp) / "output", 3, INCOMPLETE_WINDOW)
def test_schema_change_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as temp:
root = Path(temp) / "source"
for source_type in main.EXPECTED_TYPES:
create_archive(root, source_type, COMPLETE_WINDOW, bad_header=source_type == main.EXPECTED_TYPES[0])
with self.assertRaisesRegex(main.ProcessingError, "Unexpected Sheet0 header"):
main.process(main.LocalSource(root), Path(temp) / "output", 3)
def test_cross_midnight_window(self) -> None:
self.assertEqual(
main.window_bounds("2026073023000000"),
("2026-07-30 23:00:00", "2026-07-31 00:00:00"),
)
def test_mysql_store_replaces_same_hour_and_deletes_other_hours(self) -> None:
connection = FakeConnection(datetime(2026, 7, 31, 10), same_hour_rows=7, old_rows=20)
store = main.MySQLSummaryStore(connect_factory=lambda: connection)
result = store.replace_latest([database_row()])
self.assertEqual(result["metric_time"], "2026-07-31 10:00:00")
self.assertEqual(result["inserted_rows"], 1)
self.assertEqual(result["refreshed_rows"], 7)
self.assertEqual(result["old_rows_deleted"], 20)
self.assertEqual(connection.cursor_instance.inserted[0][0], datetime(2026, 7, 31, 10))
self.assertTrue(connection.committed)
self.assertFalse(connection.rolled_back)
self.assertTrue(connection.closed)
def test_mysql_store_refuses_to_replace_a_newer_hour(self) -> None:
connection = FakeConnection(datetime(2026, 7, 31, 11))
store = main.MySQLSummaryStore(connect_factory=lambda: connection)
with self.assertRaisesRegex(main.ProcessingError, "already contains newer metric time"):
store.replace_latest([database_row()])
self.assertFalse(connection.committed)
self.assertTrue(connection.rolled_back)
self.assertEqual(connection.cursor_instance.inserted, [])
self.assertTrue(connection.closed)
def test_mysql_store_creates_and_connects_to_own_database(self) -> None:
bootstrap = BootstrapConnection()
target = object()
calls: list[dict[str, object]] = []
def connect(**kwargs: object) -> object:
calls.append(kwargs)
return bootstrap if len(calls) == 1 else target
with patch.dict("sys.modules", {"pymysql": types.SimpleNamespace(connect=connect)}):
connection = main.MySQLSummaryStore()._connect()
self.assertIs(connection, target)
self.assertNotIn("database", calls[0])
self.assertEqual(calls[1]["database"], "interference_etl")
self.assertIn("CREATE DATABASE IF NOT EXISTS `interference_etl`", bootstrap.cursor_instance.executed[0])
self.assertTrue(bootstrap.cursor_instance.closed)
self.assertTrue(bootstrap.closed)
class RecordingStore:
def __init__(self) -> None:
self.rows: list[dict[str, str]] = []
def replace_latest(self, rows: list[dict[str, str]]) -> dict[str, object]:
self.rows = list(rows)
return {
"enabled": True,
"table": main.DATABASE_TABLE,
"metric_time": rows[0]["hour_start"],
"inserted_rows": len(rows),
"refreshed_rows": 0,
"old_rows_deleted": 0,
}
class FakeCursor:
def __init__(self, latest_time: datetime | None, same_hour_rows: int, old_rows: int) -> None:
self.latest_time = latest_time
self.same_hour_rows = same_hour_rows
self.old_rows = old_rows
self.rowcount = 0
self.inserted: list[tuple[object, ...]] = []
self.closed = False
def execute(self, query: str, params: tuple[object, ...] | None = None) -> None:
normalized = " ".join(query.split())
if normalized.startswith("DELETE") and "metric_time = %s" in normalized:
self.rowcount = self.same_hour_rows
elif normalized.startswith("DELETE") and "metric_time <> %s" in normalized:
self.rowcount = self.old_rows
def executemany(self, query: str, values: list[tuple[object, ...]]) -> None:
del query
self.inserted = list(values)
self.rowcount = len(values)
def fetchone(self) -> tuple[datetime | None]:
return (self.latest_time,)
def close(self) -> None:
self.closed = True
class FakeConnection:
def __init__(self, latest_time: datetime | None, same_hour_rows: int = 0, old_rows: int = 0) -> None:
self.cursor_instance = FakeCursor(latest_time, same_hour_rows, old_rows)
self.committed = False
self.rolled_back = False
self.closed = False
def cursor(self) -> FakeCursor:
return self.cursor_instance
def commit(self) -> None:
self.committed = True
def rollback(self) -> None:
self.rolled_back = True
def close(self) -> None:
self.closed = True
class BootstrapCursor:
def __init__(self) -> None:
self.executed: list[str] = []
self.closed = False
def execute(self, query: str) -> None:
self.executed.append(" ".join(query.split()))
def close(self) -> None:
self.closed = True
class BootstrapConnection:
def __init__(self) -> None:
self.cursor_instance = BootstrapCursor()
self.closed = False
def cursor(self) -> BootstrapCursor:
return self.cursor_instance
def close(self) -> None:
self.closed = True
def database_row() -> dict[str, str]:
return {
"hour_start": "2026-07-31 10:00:00",
"hour_end": "2026-07-31 11:00:00",
"source_type": "5G干扰监控",
"cgi": "460-00-200-1",
"cell_name": "测试小区",
"interference_dbm": "-100.5",
"source_path": "/source.zip",
}
def create_archive(root: Path, source_type: str, window: str, bad_header: bool = False) -> None:
date_dir = root / f"{window[:4]}-{window[4:6]}-{window[6:8]}"
date_dir.mkdir(parents=True, exist_ok=True)
filename = f"{source_type}_LWP_每小时_过滤110_{window}"
workbook = Workbook()
sheet = workbook.active
sheet.title = "Sheet0"
header = list(main.EXPECTED_HEADERS[source_type])
if bad_header:
header[-1] = "unexpected"
sheet.append(header)
sheet.append(mock_row(source_type, window))
metadata = workbook.create_sheet("指标(计数器)")
metadata.append(["指标或计数器", "指标或计数器描述", "指标公式", "指标或计数器状态"])
content = io.BytesIO()
workbook.save(content)
workbook.close()
with zipfile.ZipFile(date_dir / f"{filename}.zip", "w", zipfile.ZIP_DEFLATED) as archive:
archive.writestr(f"{filename}.xlsx", content.getvalue())
def mock_row(source_type: str, window: str) -> list[object]:
header = main.EXPECTED_HEADERS[source_type]
values: dict[str, object] = {column: "mock" for column in header}
values.update(
{
"开始时间": datetime.strptime(window[:12], "%Y%m%d%H%M"),
"粒度": "1 小时",
"eNodeBId": 100,
"eNodeBID": 100,
"gNBId": 200,
"gNBplmn": "460-00",
"cellId": 1,
"小区ID": 1,
"masterOperatorId": "unused-source-value",
"E-UTRAN FDD小区名称": f"{source_type}-小区",
"E-UTRAN TDD小区名称": f"{source_type}-小区",
"CU小区配置名称": f"{source_type}-小区",
"小区名称": f"{source_type}-小区",
"载波平均噪声干扰(dBm)": -100.5,
"小区上行平均干扰电平(dBm)": -100.5,
}
)
return [values[column] for column in header]
if __name__ == "__main__":
unittest.main()