TanujInsane/document-classification-env
3
1"""2Document Classification Environment Client.3 4Connects to the running OpenEnv server (FastAPI + WebSocket) and provides5a typed interface for interacting with the document classification environment.6 7Example (sync):8 from document_classification_env import DocEnv, DocAction9 10 with DocEnv(base_url="https://tanujinsane-document-classification-env.hf.space").sync() as env:11 result = env.reset()12 print("Documents to classify:", result.observation.total_documents)13 14 for i in range(5):15 action = DocAction(action_id=i % 5, difficulty="easy")16 result = env.step(action)17 print(f"reward={result.reward:.2f} done={result.done}")18 if result.done:19 break20 21Example (async):22 import asyncio23 from document_classification_env import DocEnv, DocAction24 25 async def main():26 async with DocEnv(base_url="http://localhost:7860") as env:27 result = await env.reset()28 result = await env.step(DocAction(action_id=0, difficulty="easy"))29 print(result.reward)30 31 asyncio.run(main())32"""33 34from __future__ import annotations35 36from typing import Any, Dict37 38from openenv.core.env_client import EnvClient39from openenv.core.client_types import StepResult40 41from .models import DocAction, DocObservation, DocState42 43 44class DocEnv(EnvClient[DocAction, DocObservation, DocState]):45 """46 Client for the Document Classification Environment.47 48 Maintains a persistent WebSocket connection to the environment server,49 enabling efficient multi-step interactions with lower latency.50 51 Supports three difficulty levels:52 - easy : 5 categories, 100 documents, no SLA53 - medium : 10 categories, 500 documents, SLA 120s54 - hard : 22 categories, 1000 documents, SLA 60s55 56 Actions:57 action_id 0..N-1 → classify document into category N58 action_id N → request metadata (tool, costs -0.05 reward)59 action_id N+1 → escalate to human (safe fallback, 0.0 reward)60 """61 62 def _step_payload(self, action: DocAction) -> Dict[str, Any]:63 """Convert DocAction to JSON payload for the step request."""64 return {65 "action_id": action.action_id,66 "difficulty": action.difficulty,67 }68 69 def _parse_result(self, payload: Dict[str, Any]) -> StepResult[DocObservation]:70 """Parse server WebSocket response into a typed StepResult."""71 obs_data = payload.get("observation", payload)72 73 observation = DocObservation(74 document_id=obs_data.get("document_id", ""),75 content=obs_data.get("content", ""),76 word_count=obs_data.get("word_count", 0),77 has_urgency_markers=obs_data.get("has_urgency_markers", False),78 features=obs_data.get("features", []),79 document_index=obs_data.get("document_index", 0),80 total_documents=obs_data.get("total_documents", 0),81 metadata_response=obs_data.get("metadata_response", "Not requested."),82 sla_remaining=obs_data.get("sla_remaining", 999.0),83 reward=obs_data.get("reward", payload.get("reward", 0.0)),84 done=obs_data.get("done", payload.get("done", False)),85 metadata=obs_data.get("metadata", payload.get("info", {})),86 )87 88 return StepResult(89 observation=observation,90 reward=observation.reward,91 done=observation.done,92 )93 94 def _parse_state(self, payload: Dict[str, Any]) -> DocState:95 """Parse server /state response into a typed DocState."""96 return DocState(97 episode_id=payload.get("episode_id", ""),98 step_count=payload.get("step_count", 0),99 difficulty=payload.get("difficulty", "easy"),100 accuracy=payload.get("accuracy", 0.0),101 total_reward=payload.get("total_reward", 0.0),102 sla_breaches=payload.get("sla_breaches", 0),103 )104 