from __future__ import annotations from enum import Enum from typing import Any, Dict, List, Optional from openenv.core.env_server.types import Action, Observation, State from pydantic import BaseModel, ConfigDict, Field, model_validator class ActionType(str, Enum): CLASSIFY = "classify" SET_PRIORITY = "set_priority" ROUTE = "route" DRAFT_REPLY = "draft_reply" SUBMIT = "submit" class B2BSupportPayload(BaseModel): model_config = ConfigDict(extra="forbid") category: Optional[str] = Field(default=None, description="Ticket category label") priority: Optional[str] = Field(default=None, description="Ticket priority") route_queue: Optional[str] = Field(default=None, description="Destination support queue") sla_minutes: Optional[int] = Field(default=None, ge=1, le=1440) escalate: Optional[bool] = Field(default=None, description="Whether this ticket should be escalated") reply_text: Optional[str] = Field(default=None, description="Customer-facing draft response") class B2BSupportTriageAction(Action): action_type: ActionType = Field(..., description="Triaging operation to perform") ticket_id: Optional[str] = Field(default=None, description="Ticket identifier (required except submit)") payload: B2BSupportPayload = Field(default_factory=B2BSupportPayload) @model_validator(mode="after") def _validate_required_fields(self) -> "B2BSupportTriageAction": if self.action_type != ActionType.SUBMIT and (self.ticket_id is None or not self.ticket_id.strip()): raise ValueError("ticket_id is required for non-submit actions") required_by_action: Dict[ActionType, List[str]] = { ActionType.CLASSIFY: ["category"], ActionType.SET_PRIORITY: ["priority"], ActionType.ROUTE: ["route_queue", "sla_minutes"], ActionType.DRAFT_REPLY: ["reply_text"], ActionType.SUBMIT: [], } missing: List[str] = [] for field_name in required_by_action[self.action_type]: value = getattr(self.payload, field_name) if value is None or (isinstance(value, str) and not value.strip()): missing.append(field_name) if missing: raise ValueError(f"Missing payload fields for {self.action_type.value}: {', '.join(missing)}") return self class RewardBreakdown(BaseModel): model_config = ConfigDict(extra="forbid") correctness_delta: float = 0.0 policy_bonus: float = 0.0 repeat_penalty: float = 0.0 invalid_penalty: float = 0.0 terminal_bonus: float = 0.0 class VisibleTicket(BaseModel): model_config = ConfigDict(extra="forbid") ticket_id: str subject: str body: str customer_tier: str contract_plan: str region: str prior_incidents: int currently_down: bool class B2BSupportTriageObservation(Observation): task_id: str = "easy" step_index: int = 0 max_steps: int = 0 visible_ticket: VisibleTicket current_plan: List[str] = Field(default_factory=list) applied_decisions: Dict[str, Any] = Field(default_factory=dict) last_action_error: Optional[str] = None progress_score: float = 0.0 reward_breakdown: RewardBreakdown = Field(default_factory=RewardBreakdown) class B2BSupportTriageState(State): task_id: str = "" seed: Optional[int] = None max_steps: int = 0 cumulative_reward: float = 0.0 last_action_error: Optional[str] = None applied_decisions: Dict[str, Any] = Field(default_factory=dict) action_history: List[Dict[str, Any]] = Field(default_factory=list) completion_flags: Dict[str, bool] = Field(default_factory=dict) class GraderResult(BaseModel): model_config = ConfigDict(extra="forbid") score: float = Field(..., ge=0.0, le=1.0) criteria: Dict[str, float] = Field(default_factory=dict)