openenv-etl / models.py
PushkarDesai's picture
Upload folder using huggingface_hub
94355ec verified
Raw
History Blame Contribute Delete
3.69 kB
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
"""
Data models for the ETL Pipeline Fixer Environment.
"""
from typing import Literal, Optional
from openenv.core.env_server.types import Action, Observation
from pydantic import Field, model_validator
# ---------------------------------------------------------------------------
# ── Observation ─────────────────────────────────────────────────────────────
# ---------------------------------------------------------------------------
class EtlObservation(Observation):
console_logs: str = Field(default="", description="Combined stdout / stderr.")
current_script_content: str = Field(default="", description="Full text of script.")
db_state_summary: str = Field(default="", description="Summary of target DB.")
# ---------------------------------------------------------------------------
# ── Action ──────────────────────────────────────────────────────────────────
# ---------------------------------------------------------------------------
class EtlAction(Action):
action_type: Literal["edit_code", "run_pipeline", "query_db"]
filepath: Optional[str] = Field(default=None, min_length=1)
target_line: Optional[int] = Field(default=None, ge=1)
new_code: Optional[str] = Field(default=None)
timeout_override_s: Optional[float] = Field(default=None, ge=0.1)
sql_string: Optional[str] = Field(default=None, min_length=1)
max_rows: int = Field(default=50, ge=1, le=1000)
@model_validator(mode="after")
def _validate_required_fields(self):
if self.action_type == "edit_code":
if self.filepath is None or self.target_line is None or self.new_code is None:
raise ValueError("edit_code requires filepath, target_line, and new_code")
elif self.action_type == "query_db":
if self.sql_string is None:
raise ValueError("query_db requires sql_string")
return self
class EditCodeAction(EtlAction):
action_type: Literal["edit_code"] = Field(default="edit_code")
filepath: str = Field(..., min_length=1)
target_line: int = Field(..., ge=1)
new_code: str = Field(...)
class RunPipelineAction(EtlAction):
action_type: Literal["run_pipeline"] = Field(default="run_pipeline")
timeout_override_s: Optional[float] = Field(default=None, ge=0.1)
class QueryDbAction(EtlAction):
action_type: Literal["query_db"] = Field(default="query_db")
sql_string: str = Field(..., min_length=1)
max_rows: int = Field(default=50, ge=1, le=1000)
# ---------------------------------------------------------------------------
# ── Reward ───────────────────────────────────────────────────────────────────
# ---------------------------------------------------------------------------
class EtlReward:
def __init__(self, value: float = 0.0) -> None:
if not (0.0 <= value <= 1.0):
raise ValueError(f"Reward must be in [0.0, 1.0], got {value!r}.")
self._value = float(value)
@property
def value(self) -> float:
return self._value
def __float__(self) -> float:
return self._value