Spaces:
Sleeping
Sleeping
File size: 5,781 Bytes
13fe504 | 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 | """OpenEnv-core adapter for the DataForge RL environment."""
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from pydantic import BaseModel, Field
from dataforge.env.environment import DataForgeEnv
if TYPE_CHECKING:
class OpenEnvAction(BaseModel):
"""Typed stand-in for openenv-core's action model."""
metadata: dict[str, Any] = Field(default_factory=dict)
class OpenEnvObservation(BaseModel):
"""Typed stand-in for openenv-core's observation model."""
done: bool = False
reward: float | None = None
metadata: dict[str, Any] = Field(default_factory=dict)
class OpenEnvEnvironment:
"""Typed stand-in for openenv-core's environment base."""
def __init__(self) -> None: ...
def create_app(*args: Any, **kwargs: Any) -> Any:
"""Typed stand-in for openenv-core's FastAPI app factory."""
...
else:
try:
from openenv.core.env_server import Action as OpenEnvAction
from openenv.core.env_server import Environment as OpenEnvEnvironment
from openenv.core.env_server import Observation as OpenEnvObservation
from openenv.core.env_server import create_app
except ImportError as exc: # pragma: no cover - exercised only without openenv extra
raise RuntimeError(
"The OpenEnv adapter requires the openenv extra: pip install 'dataforge[openenv]'."
) from exc
class DataForgeOpenEnvAction(OpenEnvAction):
"""OpenEnv action wrapper for DataForge's typed action payloads."""
action_type: str = Field(min_length=1)
row_indices: list[int] | None = None
column_names: list[str] | None = None
query: str | None = None
sql: str | None = None
test_type: str | None = None
test: str | None = None
column: str | None = None
threshold: float | None = None
pattern: str | None = None
regex: str | None = None
expect_match: bool | None = None
claim: str | None = None
affected_rows: list[int] | None = None
affected_columns: list[str] | None = None
root_cause_type: str | None = None
error_indices: list[int] | None = None
row: int | None = None
issue_type: str | None = None
new_value: str | None = None
proposed_value: str | None = None
justification: str | None = None
fix_type: str | None = None
def as_dataforge_payload(self) -> dict[str, Any]:
"""Return the action payload expected by ``DataForgeEnv.step``."""
payload: dict[str, Any] = self.model_dump(exclude_none=True)
payload.pop("metadata", None)
return payload
class DataForgeOpenEnvObservation(OpenEnvObservation):
"""OpenEnv observation model mirroring DataForge's native observation."""
visible_rows: list[dict[str, Any]] | None = None
detector_hints: list[str] | None = None
scratchpad_summary: str = ""
step_budget_remaining: int = 0
tool_usage_history: list[dict[str, Any]] = Field(default_factory=list)
latest_result: dict[str, Any] | None = None
cumulative_reward: float = 0.0
def _to_openenv_observation(payload: dict[str, Any]) -> DataForgeOpenEnvObservation:
"""Convert a native DataForge observation dictionary into OpenEnv shape."""
return DataForgeOpenEnvObservation(
visible_rows=payload.get("visible_rows"),
detector_hints=payload.get("detector_hints"),
scratchpad_summary=str(payload.get("scratchpad_summary", "")),
step_budget_remaining=int(payload.get("step_budget_remaining", 0)),
tool_usage_history=list(payload.get("tool_usage_history") or []),
latest_result=payload.get("latest_result"),
done=bool(payload.get("done", False)),
reward=payload.get("reward"),
cumulative_reward=float(payload.get("cumulative_reward", 0.0)),
metadata=dict(payload.get("metadata") or {}),
)
class DataForgeOpenEnv(OpenEnvEnvironment):
"""OpenEnv-native environment wrapper."""
SUPPORTS_CONCURRENT_SESSIONS = True
def __init__(self) -> None:
super().__init__()
self._env = DataForgeEnv()
self._last_observation: DataForgeOpenEnvObservation | None = None
def reset(
self,
seed: int | None = None,
episode_id: str | None = None,
**kwargs: Any,
) -> DataForgeOpenEnvObservation:
"""Reset the wrapped DataForge environment."""
del episode_id, kwargs
result = self._env.reset(seed=seed)
observation = _to_openenv_observation(result.observation.model_dump(mode="json"))
self._last_observation = observation
return observation
def step(
self,
action: DataForgeOpenEnvAction,
timeout_s: float | None = None,
**kwargs: Any,
) -> DataForgeOpenEnvObservation:
"""Step the wrapped DataForge environment."""
del timeout_s, kwargs
result = self._env.step(action.as_dataforge_payload())
observation = _to_openenv_observation(result.observation.model_dump(mode="json"))
self._last_observation = observation
return observation
def state(self) -> DataForgeOpenEnvObservation:
"""Return the latest observation or reset lazily."""
if self._last_observation is None:
return self.reset()
return self._last_observation
def close(self) -> None:
"""Close the wrapped environment."""
self._env.close()
app = create_app(
DataForgeOpenEnv,
DataForgeOpenEnvAction,
DataForgeOpenEnvObservation,
env_name="dataforge-env",
max_concurrent_envs=64,
)
|