Spaces:
Sleeping
Sleeping
File size: 5,715 Bytes
007fbdd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 | from __future__ import annotations
from dataclasses import asdict, dataclass, field, is_dataclass
from typing import Any, Generic, Optional, TypeVar
try: # pragma: no cover - used when openenv-core is installed
from openenv.core import EnvClient, StepResult
from openenv.core.env_server import Action, Environment, Observation, State, create_fastapi_app
except (ImportError, ModuleNotFoundError): # pragma: no cover - local fallback
import httpx
from fastapi import Body, FastAPI
from pydantic import BaseModel
@dataclass
class Action:
"""Fallback Action base class."""
@dataclass
class Observation:
"""Fallback Observation base class."""
@dataclass
class State:
episode_id: str = ""
step_count: int = 0
class Environment:
"""Fallback Environment base class."""
pass
ObsT_co = TypeVar("ObsT_co")
@dataclass
class StepResult(Generic[ObsT_co]):
observation: Any
reward: float = 0.0
done: bool = False
info: dict[str, Any] = field(default_factory=dict)
ActT = TypeVar("ActT", bound=Action)
ObsT = TypeVar("ObsT", bound=Observation)
StateT = TypeVar("StateT", bound=State)
class EnvClient(Generic[ActT, ObsT, StateT]):
"""
Minimal fallback client compatible with the subset of the OpenEnv API
used by this package.
"""
def __init__(self, base_url: str):
self.base_url = base_url.rstrip("/")
self._client: Optional[httpx.Client] = None
def _ensure_client(self) -> httpx.Client:
if self._client is None:
self._client = httpx.Client(base_url=self.base_url, timeout=30.0)
return self._client
def close(self) -> None:
if self._client is not None:
self._client.close()
self._client = None
def __enter__(self):
self._ensure_client()
return self
def __exit__(self, exc_type, exc, tb):
self.close()
return False
def sync(self):
return self
def reset(self, seed: int | None = None, case_id: str | None = None):
client = self._ensure_client()
payload = {"seed": seed, "case_id": case_id}
response = client.post("/reset", json=payload)
response.raise_for_status()
return self._parse_result(response.json())
def step(self, action: ActT):
client = self._ensure_client()
response = client.post("/step", json=self._step_payload(action))
response.raise_for_status()
return self._parse_result(response.json())
def state(self) -> StateT:
client = self._ensure_client()
response = client.get("/state")
response.raise_for_status()
return self._parse_state(response.json())
@classmethod
def from_docker_image(
cls,
image_name: str,
base_url: str = "http://localhost:8000",
):
del image_name
return cls(base_url=base_url)
def _step_payload(self, action: ActT) -> dict[str, Any]:
raise NotImplementedError
def _parse_result(self, payload: dict[str, Any]) -> StepResult[ObsT]:
raise NotImplementedError
def _parse_state(self, payload: dict[str, Any]) -> StateT:
raise NotImplementedError
class _ResetRequest(BaseModel):
seed: int | None = None
case_id: str | None = None
def _serialize(value: Any) -> Any:
if is_dataclass(value):
return asdict(value)
return value
def create_fastapi_app(env: Any, action_cls: Any, observation_cls: Any) -> FastAPI:
del observation_cls
app = FastAPI(title="LedgerShield OpenEnv", version="0.3.0")
@app.get("/")
def root() -> dict[str, Any]:
return {
"status": "ok",
"service": "LedgerShield OpenEnv",
}
@app.get("/health")
def health() -> dict[str, Any]:
return {"status": "ok"}
@app.post("/reset")
def reset(request: _ResetRequest | None = Body(default=None)) -> dict[str, Any]:
seed = request.seed if request is not None else None
case_id = request.case_id if request is not None else None
obs = env.reset(seed=seed, case_id=case_id)
if hasattr(env, "result_payload"):
return env.result_payload(obs)
return {
"observation": _serialize(obs),
"reward": 0.0,
"done": False,
"info": {},
}
@app.post("/step")
def step(request: dict[str, Any]) -> dict[str, Any]:
action = action_cls(**request)
obs = env.step(action)
if hasattr(env, "result_payload"):
return env.result_payload(obs)
return {
"observation": _serialize(obs),
"reward": 0.0,
"done": False,
"info": {},
}
@app.get("/state")
def state() -> dict[str, Any]:
if hasattr(env, "public_state"):
return _serialize(env.public_state())
current_state = env.state
if is_dataclass(current_state):
return asdict(current_state)
return current_state
return app
__all__ = [
"Action",
"Observation",
"State",
"Environment",
"EnvClient",
"StepResult",
"create_fastapi_app",
]
|