dprajjwal/supportops-env
0
1"""2SupportOps-Env: Client3HTTP client for interacting with a running SupportOps-Env server.4"""5from __future__ import annotations6from typing import Any, Dict, List, Optional, Tuple7 8import httpx9 10from models import SupportAction, SupportObservation, SupportState11 12 13class StepResult:14 """Result from a step or reset call."""15 def __init__(self, observation: SupportObservation, reward: Optional[float], done: bool):16 self.observation = observation17 self.reward = reward18 self.done = done19 20 def __repr__(self) -> str:21 return f"StepResult(reward={self.reward}, done={self.done})"22 23 24class SupportOpsEnv:25 """26 Synchronous HTTP client for SupportOps-Env.27 28 Usage:29 with SupportOpsEnv(base_url="http://localhost:8000") as env:30 result = env.reset(task_name="ticket_classification")31 print(result.observation.ticket)32 33 action = SupportAction(34 action_type="classify",35 payload={"category": "Bug"}36 )37 result = env.step(action)38 print(result.reward)39 """40 41 def __init__(self, base_url: str = "http://localhost:8000", timeout: float = 30.0):42 self.base_url = base_url.rstrip("/")43 self.timeout = timeout44 self._client: Optional[httpx.Client] = None45 46 def __enter__(self) -> "SupportOpsEnv":47 self._client = httpx.Client(base_url=self.base_url, timeout=self.timeout)48 return self49 50 def __exit__(self, *args: Any) -> None:51 if self._client:52 self._client.close()53 self._client = None54 55 def _get_client(self) -> httpx.Client:56 if self._client is None:57 self._client = httpx.Client(base_url=self.base_url, timeout=self.timeout)58 return self._client59 60 def reset(61 self,62 task_name: Optional[str] = None,63 seed: Optional[int] = None,64 episode_id: Optional[str] = None,65 ) -> StepResult:66 """Reset the environment and start a new episode."""67 payload: Dict[str, Any] = {}68 if task_name:69 payload["task_name"] = task_name70 if seed is not None:71 payload["seed"] = seed72 if episode_id:73 payload["episode_id"] = episode_id74 75 resp = self._get_client().post("/reset", json=payload)76 resp.raise_for_status()77 data = resp.json()78 return self._parse_result(data)79 80 def step(self, action: SupportAction) -> StepResult:81 """Execute an action and return the result."""82 resp = self._get_client().post(83 "/step",84 json={"action": action.model_dump()},85 )86 resp.raise_for_status()87 return self._parse_result(resp.json())88 89 def state(self) -> SupportState:90 """Get current episode state."""91 resp = self._get_client().get("/state")92 resp.raise_for_status()93 return SupportState(**resp.json())94 95 def health(self) -> Dict[str, str]:96 """Check server health."""97 resp = self._get_client().get("/health")98 resp.raise_for_status()99 return resp.json()100 101 def schema(self) -> Dict[str, Any]:102 """Get action/observation/state schemas."""103 resp = self._get_client().get("/schema")104 resp.raise_for_status()105 return resp.json()106 107 def list_tasks(self) -> List[Dict[str, Any]]:108 """List available tasks."""109 resp = self._get_client().get("/tasks")110 resp.raise_for_status()111 return resp.json().get("tasks", [])112 113 def close(self) -> None:114 if self._client:115 self._client.close()116 self._client = None117 118 @staticmethod119 def _parse_result(data: Dict[str, Any]) -> StepResult:120 obs_data = data.get("observation", {})121 obs = SupportObservation(**obs_data)122 return StepResult(123 observation=obs,124 reward=data.get("reward"),125 done=data.get("done", False),126 )127 