from __future__ import annotations from uuid import uuid4 from openenv.core.env_server.interfaces import Environment from openenv.core.env_server.types import EnvironmentMetadata from bdo_ai_env.models import ACTION_COSTS from bdo_ai_env.scenarios import get_scenario from bdo_ai_env.training import consistency_bonus from bdo_ai_env.world import BDOWorldModel from models import BDOAction, BDOObservation, BDOState, RewardBreakdownModel, VerifierScoresModel class BDOEnvironment(Environment[BDOAction, BDOObservation, BDOState]): """OpenEnv-native BDO.ai environment.""" SUPPORTS_CONCURRENT_SESSIONS: bool = True def __init__(self, scenario: str | None = None, seed: int | None = None): super().__init__() self.default_scenario_name = scenario self.default_seed = seed self._scenario_name = scenario self._seed = seed self.world = BDOWorldModel(seed=seed, scenario=get_scenario(scenario)) self._state = BDOState( episode_id=str(uuid4()), step_count=0, scenario_name=scenario or "randomized", month=self.world.world.month, shock_month=self.world.world.shock_month, seed=seed, terminated=False, last_reward=None, ) def reset( self, seed: int | None = None, episode_id: str | None = None, **kwargs, ) -> BDOObservation: self._reset_rubric() scenario_name = kwargs.get("scenario", self.default_scenario_name) if scenario_name == "randomized": scenario_name = None scenario = get_scenario(scenario_name) resolved_seed = seed if seed is not None else self.default_seed state = self.world.reset(seed=resolved_seed, scenario=scenario) self._scenario_name = scenario_name self._seed = resolved_seed self._state = BDOState( episode_id=episode_id or str(uuid4()), step_count=0, scenario_name=state.info["scenario_name"], month=state.info["month"], shock_month=state.info["shock_month"], seed=state.info["seed"], terminated=state.terminated, last_reward=None, ) return self._build_observation( reward=None, done=False, metadata={ "scenario_name": state.info["scenario_name"], "month": state.info["month"], "shock_month": state.info["shock_month"], "seed": state.info["seed"], "reward_breakdown": RewardBreakdownModel().model_dump(), "verifier_scores": VerifierScoresModel().model_dump(), "training_reward": 0.0, "consistency_bonus": 0.0, "executed_actions": [], "executed_action_objects": [], "truncated_actions": [], "actions_truncated": False, "errors": [], }, ) def step( self, action: BDOAction, timeout_s: float | None = None, **kwargs, ) -> BDOObservation: del timeout_s, kwargs self._state.step_count += 1 syntactic_score = 0.0 info = self._apply_batch(action) if info["actions_truncated"] or info["errors"]: syntactic_score = -1.0 state = self.world.state() outcome_score = ( (0.12 * state.reward.coverage) + (0.18 * state.reward.fraud_prevention) + (0.08 * state.reward.solvency) + (0.34 * state.reward.belief_accuracy) ) verifier_scores = VerifierScoresModel( syntactic_verifier=round(syntactic_score, 4), rationality_verifier=round(state.reward.belief_accuracy, 4), outcome_verifier=round(outcome_score, 4), ) training_reward = round(state.reward.total + info["consistency_bonus"], 4) reward_breakdown = RewardBreakdownModel(**state.reward.to_dict()) self._state = BDOState( episode_id=self._state.episode_id, step_count=self._state.step_count, scenario_name=state.info["scenario_name"], month=state.info["month"], shock_month=state.info["shock_month"], seed=state.info["seed"], terminated=state.terminated, last_reward=state.reward.total, ) return self._build_observation( reward=state.reward.total, done=state.terminated, metadata={ **info, **state.info, "reward_breakdown": reward_breakdown.model_dump(), "verifier_scores": verifier_scores.model_dump(), "training_reward": training_reward, }, ) @property def state(self) -> BDOState: return self._state def get_metadata(self) -> EnvironmentMetadata: return EnvironmentMetadata( name="BDO.ai", description=( "OpenEnv environment for a Block Development Officer managing " "welfare delivery under fraud pressure, infra decay, migration, " "and climate shocks." ), version="0.2.0", author="Tiya + Vaidhav", ) def _build_observation( self, *, reward: float | None, done: bool, metadata: dict, ) -> BDOObservation: state = self.world.state() payload = state.observation.to_dict() payload["reward"] = reward payload["done"] = done payload["metadata"] = metadata payload["info"] = metadata return BDOObservation.model_validate(payload) def _apply_batch(self, action: BDOAction) -> dict: executed: list[str] = [] truncated: list[str] = [] errors: list[str] = [] executed_structured: list[dict[str, object]] = [] if action.predicted_fraud_level is not None: self.world.set_predicted_fraud_level(action.predicted_fraud_level) for command in action.actions: cost = ACTION_COSTS[command.name] if self.world.world.hours_remaining < cost: truncated.append(command.name) continue handler = getattr(self.world, command.name, None) if handler is None: errors.append(f"Unknown action handler: {command.name}") continue try: result = handler(**command.params) executed.append(f"{command.name}: {result}") executed_structured.append( {"name": command.name, "params": dict(command.params)} ) except Exception as exc: errors.append(f"{command.name}: {exc}") month_result = self.world.advance_month() executed.append(f"advance_month: {month_result}") bonus = consistency_bonus(action.thought_process, executed_structured) return { "scenario_name": self._scenario_name or "randomized", "executed_actions": executed, "executed_action_objects": executed_structured, "truncated_actions": truncated, "actions_truncated": bool(truncated), "errors": errors, "thought_process": action.thought_process, "consistency_bonus": bonus, }