166 lines
5.7 KiB
Python
166 lines
5.7 KiB
Python
|
|
"""US-009: Utility function and data processing tests."""
|
||
|
|
import pytest
|
||
|
|
import math
|
||
|
|
from pathlib import Path
|
||
|
|
from datetime import datetime
|
||
|
|
|
||
|
|
|
||
|
|
class TestRiskValueToLevel:
|
||
|
|
def test_high(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(0.9) == "high"
|
||
|
|
assert risk_value_to_level(0.8) == "high"
|
||
|
|
assert risk_value_to_level(1.0) == "high"
|
||
|
|
|
||
|
|
def test_medium_high(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(0.7) == "medium_high"
|
||
|
|
assert risk_value_to_level(0.6) == "medium_high"
|
||
|
|
|
||
|
|
def test_medium(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(0.5) == "medium"
|
||
|
|
assert risk_value_to_level(0.4) == "medium"
|
||
|
|
|
||
|
|
def test_medium_low(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(0.3) == "medium_low"
|
||
|
|
assert risk_value_to_level(0.2) == "medium_low"
|
||
|
|
|
||
|
|
def test_low(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(0.1) == "low"
|
||
|
|
assert risk_value_to_level(0.0) == "low"
|
||
|
|
|
||
|
|
def test_negative_value(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(-0.1) == "low"
|
||
|
|
|
||
|
|
def test_boundaries(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
assert risk_value_to_level(0.8) == "high"
|
||
|
|
assert risk_value_to_level(0.6) == "medium_high"
|
||
|
|
assert risk_value_to_level(0.4) == "medium"
|
||
|
|
assert risk_value_to_level(0.2) == "medium_low"
|
||
|
|
|
||
|
|
|
||
|
|
class TestCalculateTrend:
|
||
|
|
def test_upward_trend(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
assert calculate_trend([0.1, 0.3, 0.5, 0.7, 0.9]) == "up"
|
||
|
|
|
||
|
|
def test_downward_trend(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
assert calculate_trend([0.9, 0.7, 0.5, 0.3, 0.1]) == "down"
|
||
|
|
|
||
|
|
def test_stable(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
assert calculate_trend([0.5, 0.51, 0.49, 0.5, 0.5]) == "stable"
|
||
|
|
|
||
|
|
def test_single_value(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
assert calculate_trend([0.5]) == "stable"
|
||
|
|
|
||
|
|
def test_empty_list(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
assert calculate_trend([]) == "stable"
|
||
|
|
|
||
|
|
def test_all_zeros(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
assert calculate_trend([0.0, 0.0, 0.0]) == "stable"
|
||
|
|
|
||
|
|
|
||
|
|
class TestValidateDateFormat:
|
||
|
|
def test_valid_dates(self):
|
||
|
|
from utils.date_helpers import validate_date_format
|
||
|
|
assert validate_date_format("20231201")
|
||
|
|
assert validate_date_format("20240101")
|
||
|
|
assert validate_date_format("20221215")
|
||
|
|
|
||
|
|
def test_invalid_dates(self):
|
||
|
|
from utils.date_helpers import validate_date_format
|
||
|
|
assert not validate_date_format("2023-12-01")
|
||
|
|
assert not validate_date_format("2023121")
|
||
|
|
assert not validate_date_format("202312011")
|
||
|
|
assert not validate_date_format("abc")
|
||
|
|
assert not validate_date_format("")
|
||
|
|
|
||
|
|
def test_edge_cases(self):
|
||
|
|
from utils.date_helpers import validate_date_format
|
||
|
|
assert not validate_date_format("2023-1-1")
|
||
|
|
assert not validate_date_format("2023/12/01")
|
||
|
|
|
||
|
|
|
||
|
|
class TestGetLatestDate:
|
||
|
|
def test_returns_valid_format(self):
|
||
|
|
from utils.date_helpers import get_latest_date
|
||
|
|
result = get_latest_date()
|
||
|
|
assert len(result) == 8
|
||
|
|
assert result.isdigit()
|
||
|
|
int(result)
|
||
|
|
|
||
|
|
def test_consistent_result(self):
|
||
|
|
from utils.date_helpers import get_latest_date
|
||
|
|
d1 = get_latest_date()
|
||
|
|
d2 = get_latest_date()
|
||
|
|
assert d1 == d2
|
||
|
|
|
||
|
|
|
||
|
|
class TestGridIdConversion:
|
||
|
|
def test_roundtrip(self):
|
||
|
|
from routers.alerts import lat_lon_to_grid_id, grid_id_to_center
|
||
|
|
lat, lon = 30.5, 114.3
|
||
|
|
grid_id = lat_lon_to_grid_id(lat, lon)
|
||
|
|
rlat, rlon = grid_id_to_center(grid_id)
|
||
|
|
assert abs(lat - rlat) < 0.001
|
||
|
|
assert abs(lon - rlon) < 0.001
|
||
|
|
|
||
|
|
def test_multiple_locations(self):
|
||
|
|
from routers.alerts import lat_lon_to_grid_id, grid_id_to_center
|
||
|
|
test_points = [
|
||
|
|
(30.59276, 114.30524), # Wuhan center area
|
||
|
|
(30.0, 114.0),
|
||
|
|
(31.0, 115.0),
|
||
|
|
]
|
||
|
|
for lat, lon in test_points:
|
||
|
|
grid_id = lat_lon_to_grid_id(lat, lon)
|
||
|
|
rlat, rlon = grid_id_to_center(grid_id)
|
||
|
|
assert abs(lat - rlat) < 0.001, f"lat mismatch: {lat} vs {rlat}"
|
||
|
|
assert abs(lon - rlon) < 0.001, f"lon mismatch: {lon} vs {rlon}"
|
||
|
|
|
||
|
|
|
||
|
|
class TestParseGeoJSON:
|
||
|
|
def test_parse_valid_file(self):
|
||
|
|
from utils.geojson import parse_geojson_file
|
||
|
|
from config import DATA_DIR
|
||
|
|
filepath = DATA_DIR / "risk_20231201.geojson"
|
||
|
|
grids = parse_geojson_file(filepath)
|
||
|
|
assert len(grids) > 0
|
||
|
|
g = grids[0]
|
||
|
|
assert "grid_id" in g
|
||
|
|
assert "latitude" in g
|
||
|
|
assert "longitude" in g
|
||
|
|
assert "risk_value" in g
|
||
|
|
assert "risk_level" in g
|
||
|
|
assert isinstance(g["risk_value"], (int, float))
|
||
|
|
assert g["risk_level"] in ("high", "medium_high", "medium", "medium_low", "low")
|
||
|
|
|
||
|
|
def test_parse_nonexistent_file(self):
|
||
|
|
from utils.geojson import parse_geojson_file
|
||
|
|
grids = parse_geojson_file(Path("/nonexistent/file.geojson"))
|
||
|
|
assert grids == []
|
||
|
|
|
||
|
|
|
||
|
|
class TestNoNaNNorInf:
|
||
|
|
"""Verify no NaN or Inf propagation in calculations."""
|
||
|
|
|
||
|
|
def test_risk_value_to_level_no_nan(self):
|
||
|
|
from utils.risk import risk_value_to_level
|
||
|
|
result = risk_value_to_level(float('nan'))
|
||
|
|
assert result in ("high", "medium_high", "medium", "medium_low", "low")
|
||
|
|
|
||
|
|
def test_trend_no_nan(self):
|
||
|
|
from utils.risk import calculate_trend
|
||
|
|
result = calculate_trend([0.5, float('nan')])
|
||
|
|
assert result in ("up", "down", "stable")
|