File size: 3,968 Bytes
a77725d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""
models.py — Core data models for AI Customer Support Simulator
"""

from dataclasses import dataclass, field
from typing import List, Optional, Dict, Any
from enum import Enum


class StepName(str, Enum):
    EMPATHY       = "empathy"
    COLLECT_INFO  = "collect_info"
    INVESTIGATE   = "investigate"
    RESOLUTION    = "resolution"


class DifficultyLevel(str, Enum):
    EASY   = "easy"
    MEDIUM = "medium"
    HARD   = "hard"


class EpisodeStatus(str, Enum):
    RUNNING = "running"
    SUCCESS = "success"
    FAIL    = "fail"


@dataclass
class StepResult:
    step: StepName
    agent_response: str
    detected_action: str          # what grader detected
    expected_action: str          # what was required
    correct: bool
    base_score: float             # 0.0 – 1.0
    step_bonus: float             # 0.0 – 0.2
    penalty: float                # 0.0+ (subtracted)
    penalty_reasons: List[str]
    reward: float                 # base_score + step_bonus - penalty
    fail_triggered: bool = False
    fail_reason: str = ""


@dataclass
class Episode:
    task_id: str
    difficulty: DifficultyLevel
    steps: List[StepResult] = field(default_factory=list)
    total_reward: float = 0.0
    wrong_step_count: int = 0
    status: EpisodeStatus = EpisodeStatus.RUNNING
    fail_reason: str = ""

    # Thresholds
    MAX_WRONG_STEPS: int = 3
    MIN_TOTAL_REWARD: float = 1.5   # out of 4.0 max

    def add_step(self, result: StepResult) -> None:
        self.steps.append(result)
        self.total_reward += result.reward
        if not result.correct:
            self.wrong_step_count += 1
        self._check_fail_conditions(result)

    def _check_fail_conditions(self, result: StepResult) -> None:
        if self.status != EpisodeStatus.RUNNING:
            return

        # Too many wrong steps
        if self.wrong_step_count >= self.MAX_WRONG_STEPS:
            self.status    = EpisodeStatus.FAIL
            self.fail_reason = f"Too many wrong steps ({self.wrong_step_count}/{self.MAX_WRONG_STEPS})"
            return

        # Step-level fail
        if result.fail_triggered:
            self.status    = EpisodeStatus.FAIL
            self.fail_reason = result.fail_reason
            return

        # All 4 steps done — check total score
        if len(self.steps) == 4:
            if self.total_reward < self.MIN_TOTAL_REWARD:
                self.status    = EpisodeStatus.FAIL
                self.fail_reason = (
                    f"Total reward {self.total_reward:.2f} below threshold "
                    f"{self.MIN_TOTAL_REWARD}"
                )
            else:
                self.status = EpisodeStatus.SUCCESS

    def summary(self) -> Dict[str, Any]:
        return {
            "task_id":         self.task_id,
            "difficulty":      self.difficulty.value,
            "status":          self.status.value,
            "total_reward":    round(self.total_reward, 3),
            "wrong_steps":     self.wrong_step_count,
            "fail_reason":     self.fail_reason,
            "steps": [
                {
                    "step":            s.step.value,
                    "correct":         s.correct,
                    "base_score":      round(s.base_score, 3),
                    "step_bonus":      round(s.step_bonus, 3),
                    "penalty":         round(s.penalty, 3),
                    "reward":          round(s.reward, 3),
                    "penalty_reasons": s.penalty_reasons,
                }
                for s in self.steps
            ],
        }


@dataclass
class Task:
    task_id: str
    difficulty: DifficultyLevel
    customer_message: str
    scenario_context: str
    required_steps: List[StepName]
    step_keywords: Dict[StepName, List[str]]   # keywords that prove step done
    escalation_risk: bool = False               # hard tasks only
    customer_emotion: str = "neutral"          # neutral | frustrated | angry