Team Ai
Apppublic

TanujInsane/document-classification-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
3likes
client.py104 linesDownload Raw Back to root
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