TanujInsane/document-classification-env
3
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 