Files
CA/backend/tests/test_utils.py

166 lines
5.7 KiB
Python
Raw Permalink Normal View History

"""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")