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