File size: 5,316 Bytes
182889e
a77725d
182889e
fa3f31a
4e39ea2
182889e
 
 
 
 
 
a77725d
4e39ea2
a77725d
 
 
 
 
 
 
 
182889e
 
 
 
 
 
 
a77725d
 
 
 
 
 
 
fa3f31a
 
 
 
 
 
 
 
 
a77725d
 
182889e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a77725d
4e39ea2
 
a77725d
 
 
 
 
182889e
 
 
 
 
 
 
 
 
 
 
a77725d
 
 
 
 
 
182889e
 
 
 
a77725d
182889e
 
 
 
 
 
 
 
 
 
 
 
4e39ea2
a77725d
 
 
 
 
 
182889e
 
a77725d
 
 
 
 
 
 
 
182889e
 
 
 
 
 
 
 
 
 
 
a77725d
 
 
 
 
 
 
 
4e39ea2
a77725d
 
182889e
 
4e39ea2
 
182889e
ca25a0a
 
 
 
 
 
 
 
 
 
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
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.openapi.docs import get_swagger_ui_html
from pydantic import BaseModel, model_validator
from typing import Optional

from environment import CustomerSupportEnv, STEP_ORDER
from graders.base_grader import BaseGrader, HardTaskGrader
from tasks.easy_task import EASY_TASK
from tasks.medium_task import MEDIUM_TASK
from tasks.hard_task import HARD_TASK

app = FastAPI(title="Customer Support AI Environment")

app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_methods=["*"],
    allow_headers=["*"],
)

_env: Optional[CustomerSupportEnv] = None

TASK_MAP = {
    "easy":   EASY_TASK,
    "medium": MEDIUM_TASK,
    "hard":   HARD_TASK,
}


class ResetRequest(BaseModel):
    task: str = "easy"


class StepRequest(BaseModel):
    action: Optional[str] = None    # accepts {"action": "..."}
    response: Optional[str] = None  # accepts {"response": "..."}

    @model_validator(mode="after")
    def resolve_action(self):
        self.action = self.action or self.response
        if not self.action:
            raise ValueError("Provide either 'action' or 'response' field with the agent reply.")
        return self


def _build_observation():
    ep = _env.episode
    task = _env.task
    step_index = _env._step_index
    current_step_name = (
        STEP_ORDER[step_index].value if step_index < len(STEP_ORDER) else "done"
    )
    return {
        "task_id": task.task_id,
        "difficulty": task.difficulty.value,
        "customer_emotion": task.customer_emotion,
        "customer_message": task.customer_message,
        "scenario_context": task.scenario_context,
        "current_step": current_step_name,
        "step_number": step_index + 1,
        "total_steps": len(STEP_ORDER),
        "episode_status": ep.status.value,
        "total_reward": round(ep.total_reward, 3),
        "wrong_step_count": ep.wrong_step_count,
        "steps_completed": [
            {
                "step": s.step.value,
                "correct": s.correct,
                "detected_action": s.detected_action,
                "reward": round(s.reward, 3),
                "penalty": round(s.penalty, 3),
                "penalty_reasons": s.penalty_reasons,
            }
            for s in ep.steps
        ],
    }


@app.get("/")
def health_check():
    return {"status": "ok", "message": "Customer Support AI Environment"}


@app.post("/reset")
def reset_env(body: ResetRequest = None):
    global _env
    task_key = (body.task if body else "easy").lower()
    if task_key not in TASK_MAP:
        raise HTTPException(
            status_code=400,
            detail=f"Unknown task '{task_key}'. Choose from: {list(TASK_MAP.keys())}"
        )
    task = TASK_MAP[task_key]
    grader = HardTaskGrader() if task_key == "hard" else BaseGrader()
    _env = CustomerSupportEnv(task=task, grader=grader)
    _env.reset()
    return {"observation": _build_observation()}


@app.post("/step")
def step_env(body: StepRequest):
    global _env
    if _env is None:
        raise HTTPException(status_code=400, detail="Environment not initialized. Call /reset first.")
    if _env.episode.status.value != "running":
        raise HTTPException(status_code=400, detail=f"Episode already ended: {_env.episode.status.value}")
    result, done = _env.step(body.action)
    return {
        "observation": _build_observation(),
        "reward": round(result.reward, 3),
        "done": done,
        "info": {
            "step": result.step.value,
            "correct": result.correct,
            "detected_action": result.detected_action,
            "expected_action": result.expected_action,
            "penalty": round(result.penalty, 3),
            "penalty_reasons": result.penalty_reasons,
            "fail_triggered": result.fail_triggered,
            "fail_reason": result.fail_reason,
        },
    }


@app.get("/state")
def get_state():
    if _env is None:
        raise HTTPException(status_code=400, detail="Environment not initialized. Call /reset first.")
    return _env.summary()


@app.get("/observation_space")
def observation_space():
    return {
        "type": "Dict",
        "fields": {
            "task_id": "str",
            "difficulty": "str (easy|medium|hard)",
            "customer_emotion": "str",
            "customer_message": "str",
            "scenario_context": "str",
            "current_step": "str (empathy|collect_info|investigate|resolution|done)",
            "step_number": "int",
            "total_steps": "int",
            "episode_status": "str (running|success|fail)",
            "total_reward": "float",
            "wrong_step_count": "int",
            "steps_completed": "List[Dict]",
        }
    }


@app.get("/action_space")
def action_space():
    return {
        "type": "Text",
        "description": "Agent free-text reply to customer",
        "min_words": 6,
    }


@app.get("/docs", include_in_schema=False)
def custom_docs():
    return get_swagger_ui_html(openapi_url="/openapi.json", title="API Docs")


def main():
    """Entry point for 'server' script defined in pyproject.toml."""
    import uvicorn
    uvicorn.run("api:app", host="0.0.0.0", port=7860)


if __name__ == "__main__":
    main()