Akhilllll/data_quality_agent
0
1"""2Data models for the Data Quality Agent Environment.3 4An ML-focused environment where an AI agent must inspect, diagnose,5and remediate data quality issues in realistic datasets - the kind of6work every ML engineer does before training a model.7"""8 9from typing import Dict, List, Optional10from openenv.core.env_server.types import Action, Observation11from pydantic import Field12 13 14class DataQualityAgentAction(Action):15 """Action for the Data Quality Agent environment - analyst commands."""16 17 message: str = Field(18 ...,19 description=(20 "Data quality command(s), semicolon-separated. "21 "Examples: DESCRIBE, INSPECT age, CHECK_MISSING, "22 "FIX_MISSING age median, VALIDATE, SUBMIT_REPORT"23 ),24 )25 26 27class DataQualityAgentObservation(Observation):28 """Observation from the Data Quality Agent environment."""29 30 # Main output31 dashboard: str = Field(default="", description="Human-readable data quality report")32 step_number: int = Field(default=0, description="Current step number")33 total_steps: int = Field(default=40, description="Max steps in episode")34 35 # Dataset overview36 dataset_name: str = Field(default="", description="Name of the dataset being audited")37 num_rows: int = Field(default=0, description="Number of rows in dataset")38 num_columns: int = Field(default=0, description="Number of columns in dataset")39 40 # Issue tracking41 issues_detected: int = Field(default=0, description="Number of distinct issues detected so far")42 issues_fixed: int = Field(default=0, description="Number of issues remediated so far")43 total_issues: int = Field(default=0, description="Total ground-truth issues (hidden)")44 45 # Quality metrics46 quality_score: float = Field(default=0.0, description="Current overall data quality 0-1")47 completeness: float = Field(default=0.0, description="Data completeness score 0-1")48 consistency: float = Field(default=0.0, description="Data consistency score 0-1")49 accuracy: float = Field(default=0.0, description="Data accuracy score 0-1")50 51 # Phase tracking52 phase: str = Field(default="inspection", description="Current phase: inspection / remediation / validation")53 54 task_name: str = Field(default="", description="Current scenario name")55 task_description: str = Field(default="", description="Scenario description")56 