Spaces:
Sleeping
Sleeping
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)
|