monsoon-rl / test_weather_forecast_env.py
DHDRL's picture
Update test_weather_forecast_env.py
33e85dd verified
Raw
History Blame
31.8 kB
"""
tests/test_weather_forecast_env.py
====================================
Integration tests for WeatherForecastEnv and NaNSafetyWrapper.
These tests lock the contracts between the env and the zone_observation
schema.
Run with: pytest tests/ -v
"""
from __future__ import annotations
import sys
import os
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
# ---------------------------------------------------------------------------
# Path setup — allow running from repo root or tests/ directory
# ---------------------------------------------------------------------------
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
from zone_observation import (
AlertLevel,
EpisodeContext,
ForecastConfig,
RiskScore,
ZoneObs,
make_synthetic_episode_context,
make_synthetic_forecast_result,
make_synthetic_zone_obs,
)
# ---------------------------------------------------------------------------
# crop_risk_scorer mock
#
# The env calls compute_risk_score(obs, forecast, config) and reads:
# risk_score.alert_level (AlertLevel enum)
# risk_score.alert_level.severity() (int)
# risk_score.supply_shortfall_prob (float [0,1])
# risk_score.flood_risk (float [0,1])
# risk_score.drought_risk (float [0,1])
#
# We inject a configurable factory so individual tests can control
# what the scorer returns without re-importing the module each time.
# ---------------------------------------------------------------------------
def _make_mock_risk_score(
alert_level: AlertLevel = AlertLevel.ADVISORY,
supply_shortfall_prob: float = 0.4,
flood_risk: float = 0.3,
drought_risk: float = 0.1,
) -> RiskScore:
return RiskScore(
zone_id="mock_zone",
scored_at=datetime.now(tz=timezone.utc),
alert_level=alert_level,
supply_shortfall_prob=supply_shortfall_prob,
flood_risk=flood_risk,
drought_risk=drought_risk,
)
@pytest.fixture(autouse=True)
def mock_crop_risk_scorer():
mock_module = MagicMock()
mock_module.compute_risk_score.return_value = _make_mock_risk_score()
with patch.dict("sys.modules", {"crop_risk_scorer": mock_module}):
yield mock_module
@pytest.fixture
def make_env(mock_crop_risk_scorer):
def _factory(config: ForecastConfig = None, nan_wrapper: bool = True):
from weather_forecast_env import make_weather_env
return make_weather_env(config, use_nan_wrapper=nan_wrapper)
return _factory
# ---------------------------------------------------------------------------
# Observation space contract
# ---------------------------------------------------------------------------
class TestObservationSpaceContract:
def test_reset_obs_keys_are_complete(self, make_env):
env = make_env()
obs, _ = env.reset(seed=0)
expected = {
"zone_belief", "forecast_precip", "forecast_uncertainty",
"action_mask", "prior_belief", "basin_context",
}
assert set(obs.keys()) == expected
def test_reset_obs_shapes_match_space(self, make_env):
cfg = ForecastConfig(n_zones=3, horizon_days=14)
env = make_env(cfg)
obs, _ = env.reset(seed=0)
assert obs["zone_belief"].shape == (3,)
assert obs["forecast_precip"].shape == (3, 14)
assert obs["forecast_uncertainty"].shape == (3,)
assert obs["action_mask"].shape == (4,) # n_zones + 1
assert obs["prior_belief"].shape == (1,)
def test_step_obs_shapes_consistent_with_reset(self, make_env):
cfg = ForecastConfig(n_zones=2, horizon_days=10)
env = make_env(cfg)
reset_obs, _ = env.reset(seed=1)
step_obs, _, _, _, _ = env.step(0)
for key in reset_obs:
assert reset_obs[key].shape == step_obs[key].shape, (
f"Shape mismatch on '{key}' between reset and step"
)
def test_obs_values_are_all_finite_after_reset(self, make_env):
env = make_env()
obs, _ = env.reset(seed=5)
for key, arr in obs.items():
arr_f = arr.astype(float)
assert np.all(np.isfinite(arr_f)), f"Non-finite value in obs['{key}'] after reset"
def test_obs_values_are_all_finite_after_step(self, make_env):
env = make_env()
env.reset(seed=5)
obs, _, _, _, _ = env.step(0)
for key, arr in obs.items():
arr_f = arr.astype(float)
assert np.all(np.isfinite(arr_f)), f"Non-finite value in obs['{key}'] after step"
def test_zone_belief_in_unit_interval(self, make_env):
env = make_env(ForecastConfig(n_zones=4))
obs, _ = env.reset(seed=2)
assert np.all(obs["zone_belief"] >= 0.0)
assert np.all(obs["zone_belief"] <= 1.0)
def test_forecast_precip_non_negative(self, make_env):
env = make_env()
obs, _ = env.reset(seed=3)
assert np.all(obs["forecast_precip"] >= 0.0)
def test_terminate_action_always_in_mask(self, make_env):
cfg = ForecastConfig(n_zones=3)
env = make_env(cfg)
obs, _ = env.reset(seed=0)
assert obs["action_mask"][env.terminate_action], (
"terminate_action must never be masked"
)
# ---------------------------------------------------------------------------
# Belief map dynamics
# ---------------------------------------------------------------------------
class TestBeliefMapDynamics:
def test_belief_decreases_after_inspection(self, make_env):
env = make_env(ForecastConfig(n_zones=1))
obs_before, _ = env.reset(seed=0)
belief_before = obs_before["zone_belief"][0]
obs_after, _, _, _, _ = env.step(0)
belief_after = obs_after["zone_belief"][0]
assert belief_after <= belief_before, (
f"Belief increased after inspection: {belief_before}{belief_after}"
)
def test_belief_never_goes_below_floor(self, make_env):
cfg = ForecastConfig(n_zones=1, max_steps=100)
env = make_env(cfg)
env.reset(seed=0)
for _ in range(50):
obs, _, terminated, _, _ = env.step(0)
assert obs["zone_belief"][0] >= cfg.belief_floor - 1e-6, (
f"Belief fell below floor: {obs['zone_belief'][0]} < {cfg.belief_floor}"
)
if terminated:
break
def test_belief_initialised_from_composite_risk_on_reset(self, make_env):
cfg = ForecastConfig(prior_belief=0.12, n_zones=2)
env = make_env(cfg)
obs, _ = env.reset(seed=0)
assert obs["zone_belief"].shape == (2,)
assert np.all(obs["zone_belief"] >= cfg.prior_belief * 0.5 - 1e-6)
assert np.all(obs["zone_belief"] <= 1.0)
def test_belief_reset_between_episodes(self, make_env):
cfg = ForecastConfig(n_zones=1, max_steps=5)
env = make_env(cfg)
obs_seed1_first, _ = env.reset(seed=1)
for _ in range(3):
env.step(0)
obs_seed1_again, _ = env.reset(seed=1)
np.testing.assert_allclose(
obs_seed1_first["zone_belief"],
obs_seed1_again["zone_belief"],
atol=1e-5,
err_msg="Same seed must produce same initial belief (no leakage)",
)
# ---------------------------------------------------------------------------
# Reward contract
# ---------------------------------------------------------------------------
class TestRewardContract:
def test_inspection_reward_is_finite(self, make_env):
env = make_env()
env.reset(seed=0)
_, reward, _, _, _ = env.step(0)
assert np.isfinite(reward)
def test_termination_reward_is_finite(self, make_env):
env = make_env()
env.reset(seed=0)
_, reward, _, _, _ = env.step(env.terminate_action)
assert np.isfinite(reward)
def test_reward_is_clipped_within_wrapper_bounds(self, make_env):
env = make_env()
env.reset(seed=0)
for _ in range(3):
_, reward, terminated, _, _ = env.step(0)
assert -31_000.0 <= reward <= 15_500.0
if terminated:
break
def test_termination_reward_positive_when_high_belief(
self, make_env, mock_crop_risk_scorer
):
cfg = ForecastConfig(alert_value=100.0, false_alert_penalty=20.0, miss_penalty=50.0)
mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score(
supply_shortfall_prob=0.9,
alert_level=AlertLevel.CRITICAL,
)
env = make_env(cfg)
env.reset(seed=0)
env.unwrapped._belief_map[:] = 0.85
_, reward, terminated, _, info = env.step(env.terminate_action)
assert terminated
assert reward > 0, (
f"Expected positive termination reward for high belief, got {reward}"
)
assert info["believed_p"] >= 0.8
def test_termination_reward_negative_when_low_risk(
self, make_env, mock_crop_risk_scorer
):
cfg = ForecastConfig(alert_value=100.0, false_alert_penalty=20.0)
mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score(
supply_shortfall_prob=0.02,
alert_level=AlertLevel.NONE,
)
env = make_env(cfg)
env.reset(seed=0)
_, reward, terminated, _, _ = env.step(env.terminate_action)
assert terminated
assert reward < 0, (
f"Expected negative termination reward for low-risk zone, got {reward}"
)
def test_invalid_action_is_penalised(self, make_env):
cfg = ForecastConfig(n_zones=1)
env = make_env(cfg)
env.reset(seed=0)
env2 = make_env(ForecastConfig(n_zones=4))
ctx = make_synthetic_episode_context("z", seed=1)
env2.reset(options={"context": ctx})
_, reward, _, _, info = env2.step(2) # padding zone
assert reward < 0 or info.get("padding_action"), (
"Inspecting a padding zone should penalise or flag padding_action"
)
# ---------------------------------------------------------------------------
# Termination conditions
# ---------------------------------------------------------------------------
class TestTerminationConditions:
def test_terminate_action_ends_episode(self, make_env):
env = make_env()
env.reset(seed=0)
_, _, terminated, truncated, info = env.step(env.terminate_action)
assert terminated
assert not truncated
assert info.get("early_termination")
def test_max_steps_exhaustion_ends_episode(self, make_env):
env = make_env(ForecastConfig(n_zones=1, max_steps=3))
env.reset(seed=0)
for step in range(4):
_, _, terminated, truncated, info = env.step(0)
if terminated or truncated:
assert step == 2, f"Expected termination at step 2, got {step}"
assert info.get("budget_exhausted")
break
else:
pytest.fail("Episode did not terminate after max_steps")
def test_step_before_reset_raises(self, make_env):
from weather_forecast_env import WeatherForecastEnv
env = WeatherForecastEnv(ForecastConfig(n_zones=1))
with pytest.raises(RuntimeError, match="reset()"):
env.step(0)
def test_budget_saved_decreases_with_steps(self, make_env):
env = make_env(ForecastConfig(n_zones=1, max_steps=10))
env.reset(seed=0)
env.step(0) # 1 step
_, _, _, _, info = env.step(env.terminate_action)
assert info["budget_saved"] == 8 # max_steps - steps_taken(2)
def test_miss_penalty_applied_when_ground_truth_present(
self, make_env, mock_crop_risk_scorer
):
from zone_observation import make_synthetic_episode_context, RiskScore, AlertLevel
mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score(
alert_level=AlertLevel.ADVISORY,
supply_shortfall_prob=0.6,
)
ctx = make_synthetic_episode_context("gt_zone", seed=7)
gt = _make_mock_risk_score(supply_shortfall_prob=0.8, alert_level=AlertLevel.WARNING)
object.__setattr__(gt, "zone_id", ctx.obs.zone_id)
from zone_observation import EpisodeContext
ctx = EpisodeContext(
obs=ctx.obs, forecast=ctx.forecast, config=ctx.config,
ground_truth=gt, zone_ids=ctx.zone_ids, data_source=ctx.data_source,
)
env = make_env(ForecastConfig(n_zones=1, max_steps=1, miss_penalty=40.0, alert_value=100.0))
env.reset(options={"context": ctx})
_, reward, terminated, _, info = env.step(env.terminate_action)
assert terminated
assert reward < 0.5, f"Expected miss penalty to pull reward down, got {reward}"
# ---------------------------------------------------------------------------
# EpisodeContext injection
# ---------------------------------------------------------------------------
class TestContextInjection:
def test_injected_context_produces_valid_obs(self, make_env):
ctx = make_synthetic_episode_context("inject_zone", seed=42)
env = make_env(ForecastConfig(n_zones=1))
obs, _ = env.reset(options={"context": ctx})
assert obs["zone_belief"].shape == (1,)
assert np.all(np.isfinite(obs["forecast_precip"]))
def test_wrong_context_type_raises_type_error(self, make_env):
env = make_env()
env.reset(seed=0)
with pytest.raises(TypeError):
env.reset(options={"context": "not_an_episode_context"})
def test_two_different_contexts_produce_different_forecasts(self, make_env):
cfg = ForecastConfig(n_zones=1)
env = make_env(cfg)
ctx_a = make_synthetic_episode_context("zone_A", flood=True, seed=1)
ctx_b = make_synthetic_episode_context("zone_B", drought=True, seed=2)
obs_a, _ = env.reset(options={"context": ctx_a})
obs_b, _ = env.reset(options={"context": ctx_b})
assert not np.array_equal(obs_a["forecast_precip"], obs_b["forecast_precip"]), (
"Different injected contexts must produce different forecast arrays"
)
def test_synthetic_fallback_works_without_context(self, make_env):
env = make_env()
obs, info = env.reset(seed=7)
assert obs is not None
assert isinstance(info, dict)
def test_context_forecast_horizon_matches_env_config(self, make_env):
cfg = ForecastConfig(n_zones=1, horizon_days=14)
obs_14d = make_synthetic_zone_obs("z", seed=1)
fc_14d = make_synthetic_forecast_result("z", horizon_days=14, valid_time=obs_14d.valid_time)
ctx = EpisodeContext(obs=obs_14d, forecast=fc_14d, config=cfg)
env = make_env(cfg)
obs, _ = env.reset(options={"context": ctx})
assert obs["forecast_precip"].shape == (1, 14)
# ---------------------------------------------------------------------------
# Soft reset — no cross-episode leakage
# ---------------------------------------------------------------------------
class TestSoftReset:
def test_two_seeded_resets_differ(self, make_env):
env = make_env()
obs0, _ = env.reset(seed=0)
obs1, _ = env.reset(seed=1)
assert not np.array_equal(
obs0["forecast_precip"], obs1["forecast_precip"]
), "Two resets with different seeds produced identical forecast_precip"
def test_belief_array_cleared_between_episodes(self, make_env):
cfg = ForecastConfig(n_zones=2, prior_belief=0.3, max_steps=20)
env = make_env(cfg)
obs_clean, _ = env.reset(seed=5)
belief_clean = obs_clean["zone_belief"].copy()
env.reset(seed=0)
for _ in range(10):
_, _, terminated, _, _ = env.step(0)
if terminated:
break
obs_after, _ = env.reset(seed=5)
np.testing.assert_allclose(
obs_after["zone_belief"], belief_clean, atol=1e-5,
err_msg="Belief leaked from previous episode into fresh reset",
)
def test_forecast_array_cleared_between_episodes(self, make_env):
cfg = ForecastConfig(n_zones=1, horizon_days=10)
env = make_env(cfg)
ctx_flood = make_synthetic_episode_context("z", flood=True, seed=1)
ctx_drought = make_synthetic_episode_context("z", drought=True, seed=2)
obs_flood, _ = env.reset(options={"context": ctx_flood})
obs_drought, _ = env.reset(options={"context": ctx_drought})
assert not np.array_equal(
obs_flood["forecast_precip"], obs_drought["forecast_precip"]
), "forecast_precip not updated on second reset — stale data from first episode"
def test_steps_taken_resets_to_zero(self, make_env):
env = make_env(ForecastConfig(n_zones=1, max_steps=10))
env.reset(seed=0)
for _ in range(3):
env.step(0)
env.reset(seed=1)
done = False
steps = 0
while not done and steps < 15:
_, _, terminated, truncated, _ = env.step(0)
done = terminated or truncated
steps += 1
assert steps <= 11, "steps_taken not reset — episode ended prematurely"
# ---------------------------------------------------------------------------
# NaNSafetyWrapper
# ---------------------------------------------------------------------------
class TestNaNSafetyWrapper:
def test_nan_in_belief_map_is_cleaned(self, make_env):
from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv
base = WeatherForecastEnv(ForecastConfig(n_zones=1))
wrapped = NaNSafetyWrapper(base)
wrapped.reset(seed=0)
base._belief_map[0] = float("nan")
obs, _, _, _, _ = wrapped.step(wrapped.terminate_action)
assert np.all(np.isfinite(obs["zone_belief"])), "NaN in belief_map not sanitised"
def test_inf_in_forecast_is_cleaned(self, make_env):
from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv
base = WeatherForecastEnv(ForecastConfig(n_zones=1, horizon_days=5))
wrapped = NaNSafetyWrapper(base)
wrapped.reset(seed=0)
base._forecast_arr[0, 2] = float("inf")
obs, _, _, _, _ = wrapped.step(0)
assert np.all(np.isfinite(obs["forecast_precip"])), "Inf in forecast not sanitised"
def test_non_finite_reward_replaced(self, mock_crop_risk_scorer):
from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv
base = WeatherForecastEnv(ForecastConfig(n_zones=1))
wrapped = NaNSafetyWrapper(base, reward_default=-999.0)
wrapped.reset(seed=0)
original_step = base.step
def _bad_step(action):
obs, _, term, trunc, info = original_step(action)
return obs, float("inf"), term, trunc, info
base.step = _bad_step
_, reward, _, _, _ = wrapped.step(0)
assert reward == -999.0 or np.isfinite(reward), (
"Non-finite reward from env was not replaced by NaNSafetyWrapper"
)
def test_terminate_action_proxied_through_wrapper(self, mock_crop_risk_scorer):
from weather_forecast_env import make_weather_env
env = make_weather_env(use_nan_wrapper=True)
env.reset(seed=0)
assert hasattr(env, "terminate_action")
assert isinstance(env.terminate_action, int)
def test_nan_limit_raises_runtime_error(self, mock_crop_risk_scorer):
from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv
base = WeatherForecastEnv(ForecastConfig(n_zones=1))
wrapped = NaNSafetyWrapper(base, nan_limit=0)
wrapped.reset(seed=0)
base._belief_map[0] = float("nan")
with pytest.raises(RuntimeError, match="invalid values"):
wrapped.step(0)
# ---------------------------------------------------------------------------
# End-to-end episode rollout
# ---------------------------------------------------------------------------
class TestEpisodeRollout:
def test_single_episode_terminates(self, make_env):
env = make_env(ForecastConfig(n_zones=1, max_steps=20))
env.reset(seed=0)
done = False
steps = 0
while not done:
action = env.action_space.sample()
_, _, terminated, truncated, _ = env.step(action)
done = terminated or truncated
steps += 1
assert steps <= 25, "Episode ran beyond max_steps + buffer"
def test_10_episodes_all_terminate_cleanly(self, make_env):
env = make_env(ForecastConfig(n_zones=1, max_steps=10))
for ep in range(10):
env.reset(seed=ep)
done = False
steps = 0
while not done:
_, _, terminated, truncated, info = env.step(0)
done = terminated or truncated
steps += 1
assert done, f"Episode {ep} never terminated"
def test_cumulative_reward_in_info(self, make_env):
env = make_env(ForecastConfig(n_zones=1, max_steps=5))
env.reset(seed=0)
for _ in range(3):
_, _, terminated, _, info = env.step(0)
assert "cumulative_reward" in info
assert np.isfinite(info["cumulative_reward"])
if terminated:
break
def test_obs_space_contains_all_returned_obs(self, make_env):
from weather_forecast_env import WeatherForecastEnv
cfg = ForecastConfig(n_zones=2, horizon_days=7, max_steps=5)
base = WeatherForecastEnv(cfg)
obs, _ = base.reset(seed=0)
assert base.observation_space.contains(obs), (
"Reset obs not contained in observation_space"
)
obs2, _, _, _, _ = base.step(0)
assert base.observation_space.contains(obs2), (
"Step obs not contained in observation_space"
)
# ---------------------------------------------------------------------------
# Visited-zone masking
# ---------------------------------------------------------------------------
class TestVisitedZoneMasking:
def test_visited_zone_is_masked_after_inspection(self, make_env):
cfg = ForecastConfig(n_zones=3, max_steps=10)
env = make_env(cfg)
obs, _ = env.reset(seed=0)
assert obs["action_mask"][0]
obs, _, _, _, _ = env.step(0)
assert not obs["action_mask"][0], "Zone 0 should be masked after inspection"
def test_terminate_remains_valid_after_all_zones_visited(self, make_env):
cfg = ForecastConfig(n_zones=2, max_steps=10)
env = make_env(cfg)
env.reset(seed=0)
env.step(0)
obs, _, terminated, _, _ = env.step(1)
if not terminated:
assert obs["action_mask"][env.terminate_action], (
"Terminate must remain valid even after all zones visited"
)
def test_revisit_penalty_applied(self, make_env):
cfg = ForecastConfig(n_zones=2, max_steps=10)
env = make_env(cfg)
env.reset(seed=0)
_, r_first, _, _, _ = env.step(0)
_, r_revisit, _, _, info = env.step(0)
assert info.get("revisit_penalty") or r_revisit < 0, (
"Re-inspecting a visited zone must incur a penalty"
)
def test_beliefs_differ_across_zones_on_reset(self, make_env):
cfg = ForecastConfig(n_zones=4)
env = make_env(cfg)
found_diverse = False
for seed in range(10):
obs, _ = env.reset(seed=seed)
beliefs = obs["zone_belief"]
if not np.all(np.isclose(beliefs, beliefs[0], atol=1e-4)):
found_diverse = True
break
assert found_diverse, (
"No seed in range(10) produced diverse zone beliefs — "
"per-zone seeding not working"
)
def test_beliefs_in_unit_interval_after_composite_risk_init(self, make_env):
cfg = ForecastConfig(n_zones=4)
env = make_env(cfg)
for seed in range(5):
obs, _ = env.reset(seed=seed)
b = obs["zone_belief"]
assert np.all(b >= 0.0) and np.all(b <= 1.0), (
f"Belief out of [0,1] at seed={seed}: {b}"
)
def test_same_seed_produces_same_initial_beliefs(self, make_env):
cfg = ForecastConfig(n_zones=3)
env = make_env(cfg)
obs_a, _ = env.reset(seed=7)
obs_b, _ = env.reset(seed=7)
np.testing.assert_allclose(
obs_a["zone_belief"], obs_b["zone_belief"], atol=1e-5,
err_msg="Same seed must produce same initial beliefs",
)
# ---------------------------------------------------------------------------
# Reward scale
# ---------------------------------------------------------------------------
class TestRewardScale:
def test_inspection_and_termination_rewards_comparable(
self, make_env, mock_crop_risk_scorer
):
mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score(
supply_shortfall_prob=0.5,
alert_level=AlertLevel.ADVISORY,
)
env = make_env(ForecastConfig(n_zones=2, max_steps=10))
env.reset(seed=0)
_, step_reward, _, _, _ = env.step(0)
env.reset(seed=0)
_, term_reward, _, _, _ = env.step(env.terminate_action)
ratio = abs(term_reward) / max(abs(step_reward), 1e-6)
assert ratio < 100.0, (
f"Reward scale mismatch: step={step_reward:.3f}, "
f"terminal={term_reward:.3f}, ratio={ratio:.1f}"
)
def test_termination_reward_scaled_by_alert_value(
self, make_env, mock_crop_risk_scorer
):
mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score(
supply_shortfall_prob=0.9,
alert_level=AlertLevel.CRITICAL,
)
cfg = ForecastConfig(alert_value=100.0, false_alert_penalty=20.0)
env = make_env(cfg)
env.reset(seed=0)
_, reward, _, _, _ = env.step(env.terminate_action)
assert abs(reward) < 20.0, (
f"Terminal reward {reward:.3f} appears unscaled (expected < 20)"
)
# ---------------------------------------------------------------------------
# Multi-zone properties
# ---------------------------------------------------------------------------
class TestMultiZoneProperties:
def test_multi_zone_env_has_correct_action_space(self, make_env):
cfg = ForecastConfig(n_zones=4)
env = make_env(cfg)
assert env.action_space.n == 5 # n_zones + terminate
def test_n_zones_info_reflects_context(self, make_env):
cfg = ForecastConfig(n_zones=3)
env = make_env(cfg)
_, info = env.reset(seed=0)
assert info["n_zones"] == 3
def test_visited_info_in_step_output(self, make_env):
cfg = ForecastConfig(n_zones=2, max_steps=5)
env = make_env(cfg)
env.reset(seed=0)
_, _, _, _, info = env.step(0)
assert "visited" in info
assert info["visited"].shape == (2,)
# ---------------------------------------------------------------------------
# Per-zone forecast diversity, uncertainty decay
# ---------------------------------------------------------------------------
class TestPerZoneForecastDiversity:
def test_forecast_rows_differ_across_zones(self, make_env):
cfg = ForecastConfig(n_zones=4, horizon_days=14)
env = make_env(cfg)
obs, _ = env.reset(seed=0)
precip = obs["forecast_precip"]
all_same = all(np.allclose(precip[0], precip[i]) for i in range(1, 4))
assert not all_same, (
"All zones have identical forecast_precip — per-zone forecasts not wired"
)
def test_forecast_diversity_across_seeds(self, make_env):
cfg = ForecastConfig(n_zones=3, horizon_days=10)
env = make_env(cfg)
found_diverse = False
for seed in range(5):
obs, _ = env.reset(seed=seed)
precip = obs["forecast_precip"]
if not np.allclose(precip[0], precip[1], atol=0.1):
found_diverse = True
break
assert found_diverse, "No seed produced diverse per-zone forecasts"
def test_uncertainty_differs_across_zones(self, make_env):
cfg = ForecastConfig(n_zones=4)
env = make_env(cfg)
found_diverse = False
for seed in range(10):
obs, _ = env.reset(seed=seed)
unc = obs["forecast_uncertainty"]
if not np.all(np.isclose(unc, unc[0], atol=1e-4)):
found_diverse = True
break
assert found_diverse, "Forecast uncertainty is identical across all zones"
class TestUncertaintyDecay:
def test_uncertainty_decreases_after_inspection(self, make_env):
cfg = ForecastConfig(n_zones=2, max_steps=10)
env = make_env(cfg)
obs_before, _ = env.reset(seed=0)
unc_before = obs_before["forecast_uncertainty"][0]
obs_after, _, _, _, _ = env.step(0)
unc_after = obs_after["forecast_uncertainty"][0]
assert unc_after < unc_before + 1e-6, (
f"Uncertainty did not decrease after inspection: "
f"{unc_before:.4f}{unc_after:.4f}"
)
def test_uninspected_zone_uncertainty_unchanged(self, make_env):
cfg = ForecastConfig(n_zones=3, max_steps=10)
env = make_env(cfg)
obs_before, _ = env.reset(seed=0)
unc_zone1_before = obs_before["forecast_uncertainty"][1]
obs_after, _, _, _, _ = env.step(0)
unc_zone1_after = obs_after["forecast_uncertainty"][1]
assert abs(unc_zone1_after - unc_zone1_before) < 1e-6, (
"Uninspected zone 1 uncertainty changed when zone 0 was inspected"
)
def test_uncertainty_bounded_after_multiple_inspections(self, make_env):
cfg = ForecastConfig(n_zones=1, max_steps=20)
env = make_env(cfg)
env.reset(seed=0)
for _ in range(15):
obs, _, terminated, _, _ = env.step(0)
unc = obs["forecast_uncertainty"][0]
assert 0.0 <= unc <= 1.0, f"Uncertainty out of [0,1]: {unc}"
if terminated:
break
def test_uncertainty_reset_between_episodes(self, make_env):
cfg = ForecastConfig(n_zones=1, max_steps=10)
env = make_env(cfg)
obs_init, _ = env.reset(seed=5)
unc_init = obs_init["forecast_uncertainty"][0]
for _ in range(5):
_, _, terminated, _, _ = env.step(0)
if terminated:
break
obs_fresh, _ = env.reset(seed=5)
unc_fresh = obs_fresh["forecast_uncertainty"][0]
assert abs(unc_fresh - unc_init) < 1e-5, (
f"Uncertainty not restored on reset: init={unc_init:.4f}, fresh={unc_fresh:.4f}"
)
def test_lower_uncertainty_improves_termination_reward(
self, make_env, mock_crop_risk_scorer
):
mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score(
supply_shortfall_prob=0.5,
alert_level=AlertLevel.ADVISORY,
)
cfg = ForecastConfig(n_zones=2, max_steps=10)
env = make_env(cfg)
env.reset(seed=0)
_, reward_early, _, _, _ = env.step(env.terminate_action)
env.reset(seed=0)
env.step(0)
env.step(1)
_, reward_late, _, _, _ = env.step(env.terminate_action)
assert reward_late >= reward_early - 1e-6, (
f"Terminating after inspection should reward at least as well: "
f"early={reward_early:.4f}, late={reward_late:.4f}"
)