File size: 3,926 Bytes
b92d20c
 
 
 
 
 
 
2d2c6e4
 
 
 
 
b92d20c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d2c6e4
 
 
b92d20c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d2c6e4
 
 
b92d20c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d2c6e4
 
 
6139c83
b92d20c
 
 
 
 
 
 
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
from __future__ import annotations

import asyncio
from typing import Any, Dict, Optional

import httpx

from codereview_env.models import (
    CodeReviewAction,
    CodeReviewObservation,
    CodeReviewState,
)


class SyncCodeReviewEnv:
    def __init__(self, async_client: "CodeReviewEnv"):
        self._async_client = async_client

    def reset(self, **kwargs: Any) -> CodeReviewObservation:
        return asyncio.run(self._async_client.reset(**kwargs))

    def step(self, action: CodeReviewAction) -> CodeReviewObservation:
        return asyncio.run(self._async_client.step(action))

    def state(self) -> CodeReviewState:
        return asyncio.run(self._async_client.state())

    def __enter__(self) -> "SyncCodeReviewEnv":
        return self

    def __exit__(self, exc_type, exc_val, exc_tb) -> None:
        return None


class CodeReviewEnv:
    def __init__(self, base_url: str = "http://localhost:7860"):
        self.base_url = base_url.rstrip("/")
        self._last_observation: Optional[CodeReviewObservation] = None
        self._session_id: Optional[str] = None

    async def reset(self, **kwargs: Any) -> CodeReviewObservation:
        async with httpx.AsyncClient(base_url=self.base_url) as client:
            response = await client.post("/reset", json=kwargs or {})
            response.raise_for_status()
            payload = response.json()
            self._session_id = payload.get("session_id")
            observation = CodeReviewObservation.model_validate(
                payload.get("observation", payload)
            )
            if "reward" in payload:
                observation.reward = payload["reward"]
            if "done" in payload:
                observation.done = payload["done"]
            self._last_observation = observation
            return observation

    async def step(self, action: CodeReviewAction) -> CodeReviewObservation:
        async with httpx.AsyncClient(base_url=self.base_url) as client:
            body: Dict[str, Any] = {"action": action.model_dump()}
            if self._session_id:
                body["session_id"] = self._session_id
            response = await client.post("/step", json=body)
            response.raise_for_status()
            payload = response.json()
            self._session_id = payload.get("session_id", self._session_id)
            observation = CodeReviewObservation.model_validate(
                payload.get("observation", payload)
            )
            if "reward" in payload:
                observation.reward = payload["reward"]
            if "done" in payload:
                observation.done = payload["done"]
            self._last_observation = observation
            return observation

    async def state(self) -> CodeReviewState:
        if self._session_id:
            async with httpx.AsyncClient(base_url=self.base_url) as client:
                response = await client.get(f"/state/{self._session_id}")
                response.raise_for_status()
                return CodeReviewState.model_validate(response.json())
        if self._last_observation is None:
            return CodeReviewState()
        return CodeReviewState(
            episode_id=self._session_id or self._last_observation.task_id,
            step_count=int(self._last_observation.metadata.get("step_count", 0)),
            task_id=self._last_observation.task_id,
            difficulty=self._last_observation.difficulty,
            title=self._last_observation.title,
            opened_artifact_ids=list(
                self._last_observation.metadata.get("opened_artifact_ids", [])
            ),
            cumulative_reward=0.05,
            score=self._last_observation.score,
            last_action_error=self._last_observation.last_action_error,
            task_metadata={"source": "client-cache"},
        )

    def sync(self) -> SyncCodeReviewEnv:
        return SyncCodeReviewEnv(self)