joshua400
Initial commit: FairRecovery++ complete multi-agent RL environment
ce75dcf
Raw
History Blame
6.46 kB
"""
FairRecovery β€” World State.
CityState is the mutable simulation world for a single episode.
ZoneState tracks per-zone metrics. Neither imports from server/.
"""
from __future__ import annotations
from copy import deepcopy
from typing import Dict, List, Optional
from .constants import (
RESOURCE_COSTS,
RESOURCE_EFFECTS,
VULNERABILITY_THRESHOLD,
)
class ZoneState:
"""Single disaster zone β€” mutable during an episode."""
__slots__ = ("zone_id", "damage", "service", "vulnerable_ratio")
def __init__(
self,
zone_id: int,
damage: float,
service: float,
vulnerable_ratio: float,
) -> None:
self.zone_id = zone_id
self.damage = float(damage)
self.service = float(service)
self.vulnerable_ratio = float(vulnerable_ratio)
# ── mutation ──────────────────────────────────────────────────────────────
def apply_resource(self, resource: str) -> None:
"""Apply resource effects, clamping to [0, 1]."""
effects = RESOURCE_EFFECTS.get(resource, {})
self.service = float(
min(1.0, max(0.0, self.service + effects.get("service", 0.0)))
)
self.damage = float(
min(1.0, max(0.0, self.damage + effects.get("damage", 0.0)))
)
# ── properties ────────────────────────────────────────────────────────────
@property
def is_vulnerable(self) -> bool:
return self.vulnerable_ratio >= VULNERABILITY_THRESHOLD
@property
def recovery_priority(self) -> float:
"""Higher = needs more help. Used by baseline greedy agent."""
return self.damage * self.vulnerable_ratio
# ── serialisation ─────────────────────────────────────────────────────────
def to_dict(self) -> Dict:
return {
"zone_id": self.zone_id,
"damage": round(self.damage, 3),
"service": round(self.service, 3),
"vulnerable_ratio": round(self.vulnerable_ratio, 3),
}
def __repr__(self) -> str:
return (
f"Zone({self.zone_id}: dmg={self.damage:.2f}, "
f"svc={self.service:.2f}, vuln={self.vulnerable_ratio:.2f})"
)
class CityState:
"""
Full episode world state.
Owns: zones list, budget, day counter, action history, pending allocations.
Responsible for: applying allocations, tracking prev_services for R_exec.
"""
def __init__(self, task_config: Dict) -> None:
self.zones: List[ZoneState] = [
ZoneState(**z) for z in task_config["zones"]
]
self.initial_budget: float = float(task_config.get("initial_budget", 100.0))
self.budget_left: float = self.initial_budget
self.day: int = 0
self.step_stage: str = "analyze"
self.history: List[str] = []
self.pending_allocations: List[Dict] = []
self.violations_total: int = 0
# snapshot services before first step for R_exec computation
self._prev_services: List[float] = [z.service for z in self.zones]
# ── mutation ──────────────────────────────────────────────────────────────
def snapshot_services(self) -> None:
"""Call BEFORE applying allocations to track Ξ”service."""
self._prev_services = [z.service for z in self.zones]
def apply_allocations(self) -> List[str]:
"""
Apply pending_allocations to zones, deducting budget.
Returns list of violation strings for safety reward computation.
"""
violations: List[str] = []
for alloc in self.pending_allocations:
zone_id = alloc.get("zone")
resource = alloc.get("resource")
# Validate zone
if zone_id is None or not (0 <= zone_id < len(self.zones)):
violations.append(f"invalid_zone:{zone_id}")
self.violations_total += 1
continue
# Validate resource
if resource not in RESOURCE_COSTS:
violations.append(f"invalid_resource:{resource}")
self.violations_total += 1
continue
cost = RESOURCE_COSTS[resource]
# Budget check
if self.budget_left < cost:
violations.append(f"budget_exceeded:zone{zone_id}:{resource}")
self.violations_total += 1
continue
# Apply
self.budget_left -= cost
self.zones[zone_id].apply_resource(resource)
self.pending_allocations = []
self.day += 1
return violations
def record(self, msg: str) -> None:
self.history.append(f"Day {self.day}: {msg}")
# ── properties ────────────────────────────────────────────────────────────
@property
def prev_services(self) -> List[float]:
return list(self._prev_services)
@property
def current_services(self) -> List[float]:
return [z.service for z in self.zones]
@property
def vulnerable_zones(self) -> List[ZoneState]:
return [z for z in self.zones if z.is_vulnerable]
@property
def non_vulnerable_zones(self) -> List[ZoneState]:
return [z for z in self.zones if not z.is_vulnerable]
# ── serialisation ─────────────────────────────────────────────────────────
def to_dict(self) -> Dict:
return {
"zones": [z.to_dict() for z in self.zones],
"day": self.day,
"budget_left": round(self.budget_left, 2),
"step_stage": self.step_stage,
"history": self.history[-5:],
}