VAKYA/Report_generation
0
1"""WebSocket client for daily_report_env."""2 3from typing import Any, Dict4 5try:6 from openenv.core.client_types import StepResult7 from openenv.core.env_client import EnvClient8 9 from .models import (10 DailyReportAction,11 DailyReportObservation,12 DailyReportState,13 ReportReward,14 )15except ImportError:16 from models import ( # type: ignore17 DailyReportAction,18 DailyReportObservation,19 DailyReportState,20 ReportReward,21 )22 23 from openenv.core.client_types import StepResult # type: ignore24 from openenv.core.env_client import EnvClient # type: ignore25 26 27class DailyReportEnv(EnvClient[DailyReportAction, DailyReportObservation, DailyReportState]):28 """Async client; use `.sync()` for synchronous inference scripts."""29 30 def _step_payload(self, action: DailyReportAction) -> Dict[str, Any]:31 return action.model_dump(exclude_none=True)32 33 def _parse_result(self, payload: Dict[str, Any]) -> StepResult[DailyReportObservation]:34 obs_data = payload.get("observation", {}) or {}35 rd = obs_data.get("reward_detail")36 reward_detail = ReportReward.model_validate(rd) if isinstance(rd, dict) else None37 38 observation = DailyReportObservation(39 task=obs_data.get("task", "daily_header"),40 instructions=obs_data.get("instructions", ""),41 static_data=obs_data.get("static_data") or {},42 header_fields=obs_data.get("header_fields") or {},43 summary_metrics=obs_data.get("summary_metrics") or {},44 kpi_rows=obs_data.get("kpi_rows") or [],45 pdf_generated=bool(obs_data.get("pdf_generated", False)),46 submitted=bool(obs_data.get("submitted", False)),47 graded_score=float(obs_data.get("graded_score", 0.0)),48 max_steps=int(obs_data.get("max_steps", 30)),49 steps_remaining=int(obs_data.get("steps_remaining", 30)),50 feedback=obs_data.get("feedback", ""),51 last_action_error=obs_data.get("last_action_error"),52 reward_detail=reward_detail,53 done=payload.get("done", False),54 reward=payload.get("reward"),55 metadata=dict(obs_data.get("metadata") or {}),56 )57 return StepResult(58 observation=observation,59 reward=payload.get("reward"),60 done=payload.get("done", False),61 )62 63 def _parse_state(self, payload: Dict[str, Any]) -> DailyReportState:64 return DailyReportState.model_validate(payload)65 