File size: 3,431 Bytes
a69c08b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

"""Emailtriage Environment Client."""

from typing import Dict

from openenv.core import EnvClient
from openenv.core.client_types import StepResult

from .models import EmailtriageAction, EmailtriageObservation, EmailtriageState


class EmailtriageEnv(
    EnvClient[EmailtriageAction, EmailtriageObservation, EmailtriageState]
):
    """
    Client for the Emailtriage Environment.

    This client maintains a persistent WebSocket
    connection to the environment server,
    enabling efficient multi-step interactions with lower latency.

    Example:
        >>> with EmailtriageEnv(base_url="http://localhost:8000") as client:
        ...     result = client.reset(options={"task_id": "easy"})
        ...     print(result.observation.inbox_remaining)
        ...
        ...     result = client.step(EmailtriageAction(
        ...         action_type="archive", target_email_id=101))
        ...     print(result.reward)

    Example with Docker:
        >>> client = EmailtriageEnv.from_docker_image("emailtriage-env:latest")
        >>> try:
        ...     result = client.reset(options={"task_id": "medium"})
        ...     result = client.step(EmailtriageAction(
        ...         action_type="read", target_email_id=101))
        ... finally:
        ...     client.close()
    """

    def _step_payload(self, action: EmailtriageAction) -> Dict:
        """Convert EmailtriageAction to JSON payload for step message."""
        return {
            "action_type": action.action_type,
            "target_email_id": action.target_email_id,
            "draft_content": action.draft_content,
            "proposed_slot": action.proposed_slot,
        }

    def _parse_result(
        self, payload: Dict
    ) -> StepResult[EmailtriageObservation]:
        """Parse server response into StepResult[EmailtriageObservation]."""
        obs_data = payload.get("observation", {})
        observation = EmailtriageObservation(
            inbox_preview=obs_data.get("inbox_preview", []),
            returned_emails=obs_data.get("returned_emails", []),
            calendar_slots=obs_data.get("calendar_slots", []),
            last_action_result=obs_data.get("last_action_result", ""),
            conversation_history=obs_data.get("conversation_history", []),
            inbox_remaining=obs_data.get("inbox_remaining", 0),
            done=payload.get("done", False),
            reward=payload.get("reward"),
            metadata=obs_data.get("metadata", {}),
        )

        return StepResult(
            observation=observation,
            reward=payload.get("reward"),
            done=payload.get("done", False),
        )

    def _parse_state(self, payload: Dict) -> EmailtriageState:
        """Parse server response into EmailtriageState object."""
        return EmailtriageState(
            episode_id=payload.get("episode_id"),
            step_count=payload.get("step_count", 0),
            task_id=payload.get("task_id", "hard"),
            inbox=payload.get("inbox", []),
            calendar_slots=payload.get("calendar_slots", []),
            queried_calendar=payload.get("queried_calendar", False),
            processed_email_ids=payload.get("processed_email_ids", []),
        )