akhilsu's picture
Upload 32 files
f392960 verified
Raw
History Blame Contribute Delete
3.88 kB
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)