""" Spreadsheet per-step reward transform (Layer 3). Used when: --reward-mode openenv --gym spreadsheet Scoring by tool: read_range / read_cell: Successful read → +0.02 (neutral, slightly positive) write_cell / write_range: Write to target region → +0.05 Write preceded by recent read (≤4 steps) → +0.05 (on top of base) Write outside target region → -0.10 Repeated write to same cell (≥3) → -0.05 inspect_formula: Successful → +0.05 validate_partial: Called successfully → +0.05 Shows improvement over last validate → +0.10 submit_workbook: All checks pass (pass_rate == 1.0) → +0.50 Partial pass (pass_rate > 0.5) → +0.20 Mostly failing (pass_rate < 0.3) → -0.10 No prior validate_partial call → -0.05 (unsupported submission) list_sheets / list_scenarios / get_session_info / load_scenario / get_edit_history / reset_scenario / list_named_targets: Successful → 0.0 (neutral) Error on any non-neutral tool → -0.05 """ from __future__ import annotations import json from typing import Any from openenv.core.env_server.mcp_types import CallToolObservation from openenv.core.env_server.types import Observation from .base import StepRewardTransform WRITE_TOOLS = frozenset({"write_cell", "write_range"}) READ_TOOLS = frozenset({"read_range", "read_cell"}) NEUTRAL_TOOLS = frozenset({ "list_sheets", "list_scenarios", "get_session_info", "load_scenario", "get_edit_history", "reset_scenario", "list_named_targets", "list_tools", }) def _extract_result(observation) -> Any: result = getattr(observation, "result", None) if hasattr(result, "data"): return result.data if isinstance(result, dict) and "data" in result: return result["data"] if isinstance(result, str): try: return json.loads(result) except (json.JSONDecodeError, TypeError): return result return result class SpreadsheetStepTransform(StepRewardTransform): """Per-step reward for Spreadsheet gym (Layer 3, trajectory-aware).""" def __init__(self, scenario: dict | None = None): super().__init__() self._scenario = scenario or {} self._recent_tools: list[str] = [] self._write_counts: dict[str, int] = {} self._last_validate_passed: int = 0 self._has_validated: bool = False def set_scenario(self, scenario: Any) -> None: """Set scenario context (called by runner at start of each scenario).""" if hasattr(scenario, "id"): self._scenario = {"id": scenario.id} elif isinstance(scenario, dict): self._scenario = scenario self._recent_tools = [] self._write_counts = {} self._last_validate_passed = 0 self._has_validated = False def _compute_reward(self, observation: Observation) -> float: if not isinstance(observation, CallToolObservation): return 0.0 tool_name = getattr(observation, "tool_name", "") or "" result = _extract_result(observation) if not isinstance(result, dict): result = {} has_error = ( observation.error is not None or (isinstance(result, dict) and result.get("error")) ) if tool_name in NEUTRAL_TOOLS: self._recent_tools.append(tool_name) return 0.0 if has_error: self._recent_tools.append(tool_name) return -0.05 reward = self._score_tool(tool_name, result) self._recent_tools.append(tool_name) return reward def _score_tool(self, tool_name: str, result: dict) -> float: if tool_name in READ_TOOLS: return 0.02 if tool_name == "inspect_formula": return 0.05 if tool_name in WRITE_TOOLS: return self._score_write(tool_name, result) if tool_name == "validate_partial": return self._score_validate(result) if tool_name == "submit_workbook": return self._score_submit(result) return 0.0 def _score_write(self, tool_name: str, result: dict) -> float: outside_target = result.get("outside_target", False) if outside_target: return -0.10 reward = 0.05 lookback = self._recent_tools[-4:] if any(t in READ_TOOLS for t in lookback): reward += 0.05 cell_key = f"{result.get('sheet', '')}:{result.get('cell', result.get('start_cell', ''))}" self._write_counts[cell_key] = self._write_counts.get(cell_key, 0) + 1 if self._write_counts[cell_key] >= 3: reward -= 0.05 return reward def _score_validate(self, result: dict) -> float: self._has_validated = True new_passed = result.get("passed", 0) if new_passed > self._last_validate_passed: self._last_validate_passed = new_passed return 0.10 self._last_validate_passed = new_passed return 0.05 def _score_submit(self, result: dict) -> float: pass_rate = result.get("pass_rate", 0) reward = 0.0 if pass_rate == 1.0: reward = 0.50 elif pass_rate > 0.5: reward = 0.20 elif pass_rate < 0.3: reward = -0.10 if not self._has_validated: reward -= 0.05 return reward def transform(trajectory: list, scenario: dict) -> list: """ Apply per-step rewards to trajectory (used by run_eval transform_factory). Returns trajectory with each step's reward populated. """ t = SpreadsheetStepTransform(scenario=scenario) for step in trajectory: if hasattr(step, "observation"): obs = step.observation step.reward = t._compute_reward(obs) return trajectory