Team Ai
Apppublic

TanujInsane/document-classification-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
3likes
test_environment.py95 linesDownload Raw Back to root
1import sys2import numpy as np3from environment import DocumentClassificationEnv4from grading import AgentGrader5from models import Observation, Action, Reward, EnvironmentState6from agent import TicketAgent7 8def test_environment_creation():9    print("Testing environment creation...")10    for difficulty in ["easy", "medium", "hard"]:11        try:12            env = DocumentClassificationEnv(task_difficulty=difficulty)13            obs, info = env.reset()14            print(f"  ✓ {difficulty} environment created successfully")15            assert "features" in obs, f"Missing features in {difficulty} observation"16            assert "content" in obs, f"Missing content in {difficulty} observation"17        except Exception as e:18            print(f"  ✗ Failed to create {difficulty} environment: {e}")19            return False20    return True21 22def test_step_function():23    print("\nTesting step function...")24    env = DocumentClassificationEnv(task_difficulty="easy")25    obs, _ = env.reset()26 27    for i in range(5):28        action = i % env.num_categories29        obs, reward, done, truncated, info = env.step(action)30        assert isinstance(reward, float), f"Reward should be float, got {type(reward)}"31        assert "is_correct" in info, "Missing is_correct in info"32        print(f"  ✓ Step {i+1}: reward={reward:.3f}, accuracy={info.get('episode_accuracy', 0):.3f}")33        if done: break34    return True35 36def test_ticket_agent_loop():37    print("\nTesting TicketAgent interaction loop...")38    agent = TicketAgent(model_dir=".")39    env = DocumentClassificationEnv(task_difficulty="easy")40    41    state = agent.run_episode(env)42    print(f"  ✓ Agent finished episode. Accuracy: {state['accuracy']:.2%}, Total Reward: {state['total_reward']:.2f}")43    assert state['total_documents'] > 044    assert 'accuracy' in state45    return True46 47def test_pydantic_models():48    print("\nTesting Pydantic typed models...")49    env = DocumentClassificationEnv(task_difficulty="easy")50    obs_raw, _ = env.reset()51 52    obs_typed = Observation.from_gym_obs(obs_raw)53    assert obs_typed.document_id != ""54    print(f"  ✓ Observation model: doc_id={obs_typed.document_id}")55 56    state_raw = env.state()57    state_typed = EnvironmentState.from_state_dict(state_raw)58    assert state_typed.difficulty == "easy"59    print(f"  ✓ EnvironmentState model: difficulty={state_typed.difficulty}, accuracy={state_typed.accuracy:.3f}")60    return True61 62def run_all_tests():63    print("=" * 60)64    print("Document Classification Environment - Fixed Test Suite")65    print("=" * 60)66 67    tests = [68        ("Environment Creation", test_environment_creation),69        ("Step Function", test_step_function),70        ("TicketAgent Loop", test_ticket_agent_loop),71        ("Pydantic Models", test_pydantic_models),72    ]73 74    results = []75    for name, func in tests:76        try:77            res = func()78            results.append((name, res))79        except Exception as e:80            print(f"  ✗ {name} failed: {e}")81            import traceback82            traceback.print_exc()83            results.append((name, False))84 85    print("\nSUMMARY")86    passed = sum(1 for _, r in results if r)87    for name, r in results:88        print(f"  [{'PASS' if r else 'FAIL'}] {name}")89    print(f"\nTotal: {passed}/{len(tests)} passed")90    return passed == len(tests)91 92if __name__ == "__main__":93    success = run_all_tests()94    sys.exit(0 if success else 1)95