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