| 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) |
|
|
|
|