| 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, |
| } |
|
|