File size: 3,094 Bytes
2ac71c7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-License-Identifier: BSD-3-Clause

"""Client for the watercolour environment."""

from __future__ import annotations

from typing import Any, Dict

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

from .models import WatercolourAction, WatercolourObservation, WatercolourState


class WatercolourEnv(
    EnvClient[WatercolourAction, WatercolourObservation, WatercolourState]
):
    """Connects to a running watercolour environment server.

    Examples:

    ```python
    with WatercolourEnv(base_url="http://localhost:8000") as env:
        observation = env.reset().observation
        reply = my_model(observation.system_prompt, observation.prompt)
        result = env.step(WatercolourAction(response=reply))
        print(result.reward, result.observation.feedback)
    ```
    """

    def _step_payload(self, action: WatercolourAction) -> Dict[str, Any]:
        """Convert an action into the JSON body of a step request."""
        return {"response": action.response}

    def _parse_result(
        self, payload: Dict[str, Any]
    ) -> StepResult[WatercolourObservation]:
        """Parse a server response into a typed step result.

        The server hoists `reward` and `done` onto the response envelope and
        drops them from the serialised observation, so they are read from the
        envelope first and only then from the observation body.
        """
        data = payload.get("observation", {})
        reward = payload.get("reward", data.get("reward"))
        done = payload.get("done", data.get("done", False))
        observation = WatercolourObservation(
            prompt=data.get("prompt", ""),
            system_prompt=data.get("system_prompt", ""),
            task_id=data.get("task_id", ""),
            subject=data.get("subject", ""),
            feedback=data.get("feedback", ""),
            gate_passed=data.get("gate_passed", False),
            length_score=data.get("length_score", 0.0),
            judge_score=data.get("judge_score", 0.0),
            judged=data.get("judged", False),
            paint_fraction=data.get("paint_fraction", 0.0),
            finished=data.get("finished", False),
            violations=data.get("violations", []),
            js_errors=data.get("js_errors", []),
            breakdown=data.get("breakdown", {}),
            image_png_base64=data.get("image_png_base64"),
            done=bool(done),
            reward=reward,
            metadata=payload.get("metadata", data.get("metadata", {})),
        )
        return StepResult(
            observation=observation,
            reward=observation.reward,
            done=observation.done,
        )

    def _parse_state(self, payload: Dict[str, Any]) -> WatercolourState:
        """Parse a response from the state endpoint into a typed state."""
        return WatercolourState(
            episode_id=payload.get("episode_id"),
            step_count=payload.get("step_count", 0),
            task_id=payload.get("task_id", ""),
            submitted=payload.get("submitted", False),
        )