File size: 3,879 Bytes
f392960
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
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)