codereview-env / client.py
Anurag137's picture
fix: comprehensive Phase 2 score range audit - eliminate all 0.0/1.0 leaks
6139c83
Raw
History Blame Contribute Delete
3.3 kB
from __future__ import annotations
import asyncio
from typing import Any, 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
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()
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:
response = await client.post("/step", json={"action": action.model_dump()})
response.raise_for_status()
payload = response.json()
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._last_observation is None:
return CodeReviewState()
return CodeReviewState(
episode_id=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)