thermal-plant-control / env /interface.py
Kaushalraj Puwar
refactor: improve physical simulation stability with RK2 integration and update documentation across the environment and task modules
ba6f178
Raw
History Blame Contribute Delete
3.8 kB
from __future__ import annotations
import threading
from abc import ABC, abstractmethod
from typing import Any, Dict, Tuple
from env.core import ThermalPlantEnv
class OpenEnvInterface(ABC):
"""
Abstract base class defining the contract for an OpenEnv-compliant interface.
This ensures that the API layer can interact with any environment that
adheres to this standard structure.
"""
@abstractmethod
def get_state(self) -> Dict[str, Any]:
"""Returns the full, unrounded internal state of the environment."""
raise NotImplementedError
@abstractmethod
def reset(self, task_id: str, episode_id: int) -> Dict[str, float]:
"""Resets the environment to a new initial state for a given task and episode."""
raise NotImplementedError
@abstractmethod
def step(self, action: Dict[str, float]) -> Tuple[Dict[str, float], float, bool, Dict[str, Any]]:
"""
Executes one time step in the environment.
Args:
action: A dictionary containing the agent's action.
Returns:
A tuple containing:
- observation (Dict[str, float]): The agent's observation of the current environment state.
- reward (float): The amount of reward returned after the previous action.
- done (bool): Whether the episode has ended.
- info (Dict[str, Any]): Contains auxiliary diagnostic information.
"""
raise NotImplementedError
class ConcreteOpenEnvInterface(OpenEnvInterface):
"""
Concrete implementation of the OpenEnvInterface.
This class provides a thread-safe singleton wrapper for the core thermal
plant environment. It decouples the API layer from the simulation
internals, ensuring consistent state management across sequential
evaluation tasks.
"""
_instance: "ConcreteOpenEnvInterface" | None = None
_env: ThermalPlantEnv | None = None
_lock: threading.RLock = threading.RLock()
def __new__(cls, max_steps: int | None = None) -> "ConcreteOpenEnvInterface":
with cls._lock:
if cls._instance is None:
cls._instance = super().__new__(cls)
if max_steps is None:
cls._env = ThermalPlantEnv()
else:
cls._env = ThermalPlantEnv(max_steps=int(max_steps))
elif cls._env is not None and max_steps is not None:
cls._env.max_steps = int(max_steps)
return cls._instance
def __init__(self, max_steps: int | None = None) -> None:
"""Allow non-API callers to configure the shared env without bypassing the interface."""
with self._lock:
if self._env is not None and max_steps is not None:
self._env.max_steps = int(max_steps)
def get_state(self) -> Dict[str, Any]:
"""Returns the full, unrounded internal state from the core environment."""
with self._lock:
if self._env is None:
raise RuntimeError("Environment not initialized.")
return self._env.state()
def reset(self, task_id: str, episode_id: int) -> Dict[str, float]:
"""Resets the core environment and returns the initial observation."""
with self._lock:
if self._env is None:
raise RuntimeError("Environment not initialized.")
return self._env.reset(task_id=task_id, episode_id=episode_id)
def step(self, action: Dict[str, float]) -> Tuple[Dict[str, float], float, bool, Dict[str, Any]]:
"""Performs a step in the core environment."""
with self._lock:
if self._env is None:
raise RuntimeError("Environment not initialized.")
return self._env.step(action)