training-monkey/dataoncallenv
0
1"""Integration tests for DataOnCallEnv.2"""3 4import sys5import json6 7# Ensure project root is on path8sys.path.insert(0, ".")9 10from environment import DataOnCallEnv11from models import Action, Observation, Reward12 13passed = 014failed = 015 16def test(name, condition, detail=""):17 global passed, failed18 if condition:19 print(f" ✅ {name}")20 passed += 121 else:22 print(f" ❌ {name} — {detail}")23 failed += 124 25# ══════════════════════════════════════════════════════════════════════════════26print("\n" + "="*60)27print("TEST SUITE: DataOnCallEnv v2.0")28print("="*60)29 30# ── 1. Partial Observability ─────────────────────────────────────────────────31print("\n── 1. Partial Observability ──")32 33env = DataOnCallEnv()34obs = env.reset(task_id=1)35 36# No tables should be visible in initial observation37test("Reset hides tables",38 "available_tables" not in str(obs.result) or "tables" not in obs.result,39 f"Tables found in reset: {obs.result}")40 41test("Discovered tables empty at start",42 len(env.discovered_tables) == 0)43 44# Try to inspect schema before discovering tables45action = Action(tool="inspect_schema", query="sales", reasoning="test")46obs, reward, done, info = env.step(action)47test("inspect_schema before discovery → error",48 "error" in str(obs.result).lower() or "not discovered" in str(obs.result).lower(),49 f"Got: {obs.result}")50 51# Try to run SQL before discovering tables52action = Action(tool="run_sql", query="SELECT sale_id FROM sales LIMIT 5", reasoning="test")53obs, reward, done, info = env.step(action)54test("run_sql before discovery → error",55 "error" in str(obs.result).lower(),56 f"Got: {obs.result}")57 58# Now discover tables59action = Action(tool="list_tables", query="", reasoning="discover tables")60obs, reward, done, info = env.step(action)61test("list_tables discovers tables",62 len(env.discovered_tables) > 0,63 f"Discovered: {env.discovered_tables}")64 65# Now inspect_schema should work66action = Action(tool="inspect_schema", query="sales", reasoning="check schema after discovery")67obs, reward, done, info = env.step(action)68test("inspect_schema after discovery → works",69 "columns" in str(obs.result),70 f"Got: {obs.result}")71 72# ── 2. Query Cost System ─────────────────────────────────────────────────────73print("\n── 2. Query Cost System ──")74 75env2 = DataOnCallEnv()76obs = env2.reset(task_id=1)77 78test("Cost starts at 0", obs.cost_spent == 0.0)79test("Budget starts at 20.0", obs.budget_remaining == 20.0)80 81# list_tables costs 0.582action = Action(tool="list_tables", query="", reasoning="test cost")83obs, _, _, info = env2.step(action)84test("list_tables costs 0.5", obs.cost_spent == 0.5, f"Got: {obs.cost_spent}")85test("Budget remaining correct", obs.budget_remaining == 19.5, f"Got: {obs.budget_remaining}")86 87# run_sql costs 2.088action = Action(tool="run_sql", query="SELECT sale_id, amount FROM sales LIMIT 3", reasoning="test cost")89obs, _, _, info = env2.step(action)90test("run_sql costs 2.0", obs.cost_spent == 2.5, f"Got: {obs.cost_spent}")91 92# Cost appears in info93test("Cost in info dict", "cost_spent" in info and "budget_remaining" in info)94 95# ── 3. Real Logs (dbt + Airflow) ─────────────────────────────────────────────96print("\n── 3. Real Logs (dbt + Airflow) ──")97 98env3 = DataOnCallEnv()99env3.reset(task_id=1)100 101# Check dbt logs are expanded102action = Action(tool="check_logs", query="", reasoning="test logs")103obs, _, _, _ = env3.step(action)104log_result = obs.result105test("dbt logs have multiple entries",106 isinstance(log_result, list) and len(log_result) >= 8,107 f"Got {len(log_result) if isinstance(log_result, list) else 'non-list'} entries")108 109# Check logs have realistic fields110if isinstance(log_result, list) and len(log_result) > 0:111 test("dbt logs have duration_s field",112 "duration_s" in log_result[0],113 f"Keys: {log_result[0].keys()}")114 115# Check Airflow logs116action = Action(tool="check_airflow", query="", reasoning="test airflow")117obs, _, _, _ = env3.step(action)118af_result = obs.result119test("Airflow runs available",120 isinstance(af_result, list) and len(af_result) >= 3,121 f"Got {len(af_result) if isinstance(af_result, list) else 'non-list'} entries")122 123if isinstance(af_result, list) and len(af_result) > 0:124 test("Airflow has dag_id field",125 "dag_id" in af_result[0],126 f"Keys: {af_result[0].keys()}")127 128# Test all 3 tasks have logs129for tid in [1, 2, 3]:130 env_t = DataOnCallEnv()131 env_t.reset(task_id=tid)132 obs_l, _, _, _ = env_t.step(Action(tool="check_logs", query="", reasoning="test"))133 obs_a, _, _, _ = env_t.step(Action(tool="check_airflow", query="", reasoning="test"))134 test(f"Task {tid} has dbt logs",135 isinstance(obs_l.result, list) and len(obs_l.result) >= 8)136 test(f"Task {tid} has airflow runs",137 isinstance(obs_a.result, list) and len(obs_a.result) >= 3)138 139# ── 4. Anti-Cheating ─────────────────────────────────────────────────────────140print("\n── 4. Anti-Cheating ──")141 142env4 = DataOnCallEnv()143env4.reset(task_id=1)144 145# Try to submit before minimum steps146action = Action(tool="submit", query="case mismatch", reasoning="test early submit")147obs, reward, done, info = env4.step(action)148test("Early submit blocked",149 "error" in str(obs.result).lower() and not done,150 f"Done={done}, result={obs.result}")151 152# SELECT * should be blocked153action = Action(tool="list_tables", query="", reasoning="discover")154env4.step(action)155action = Action(tool="run_sql", query="SELECT * FROM sales", reasoning="test select star")156obs, _, _, _ = env4.step(action)157test("SELECT * blocked",158 "error" in str(obs.result).lower() and "select *" in str(obs.result).lower(),159 f"Got: {obs.result}")160 161# Specific columns should work162action = Action(tool="run_sql", query="SELECT sale_id, amount FROM sales LIMIT 3", reasoning="test")163obs, _, _, _ = env4.step(action)164test("Specific columns work",165 isinstance(obs.result, list),166 f"Got: {type(obs.result)}")167 168# ── 5. Better Evaluation ─────────────────────────────────────────────────────169print("\n── 5. Better Evaluation (Graders) ──")170 171# Test with a good investigation path172env5 = DataOnCallEnv()173env5.reset(task_id=1)174 175good_actions = [176 Action(tool="list_tables", query="", reasoning="First, discover what tables exist"),177 Action(tool="inspect_schema", query="sales", reasoning="Check sales table structure"),178 Action(tool="inspect_schema", query="currency_rates", reasoning="Check currency rates structure"),179 Action(tool="check_logs", query="", reasoning="Check pipeline logs for errors"),180 Action(tool="run_sql", query="SELECT currency FROM sales LIMIT 5", reasoning="Check currency format in sales"),181 Action(tool="run_sql", query="SELECT currency_code FROM currency_rates", reasoning="Checking if there's a case mismatch between tables"),182 Action(tool="run_sql",183 query="SELECT ROUND(SUM(s.amount * cr.rate_to_usd), 2) as total_revenue_usd FROM sales s JOIN currency_rates cr ON LOWER(s.currency) = cr.currency_code",184 reasoning="Verify fix by joining with LOWER()"),185 Action(tool="submit",186 query="ROOT CAUSE: The currency codes have a case mismatch. The sales table stores currency as uppercase (USD, EUR, GBP) but currency_rates stores them as lowercase (usd, eur, gbp). The JOIN fails silently because of this case-sensitive mismatch, returning NULLs for non-USD currencies. CORRECTED SQL: SELECT ROUND(SUM(s.amount * cr.rate_to_usd), 2) as total_revenue_usd FROM sales s JOIN currency_rates cr ON LOWER(s.currency) = cr.currency_code",187 reasoning="Found the root cause: case mismatch in currency codes"),188]189 190for action in good_actions:191 obs, reward, done, info = env5.step(action)192 193test("Good agent gets reward", reward is not None)194if reward:195 test("Diagnosis scored > 0", reward.breakdown.diagnosis_correct > 0,196 f"Got: {reward.breakdown.diagnosis_correct}")197 test("Fix scored > 0", reward.breakdown.fix_valid > 0,198 f"Got: {reward.breakdown.fix_valid}")199 test("Investigation quality > 0", reward.breakdown.investigation_quality > 0,200 f"Got: {reward.breakdown.investigation_quality}")201 test("Reasoning quality > 0", reward.breakdown.reasoning_quality > 0,202 f"Got: {reward.breakdown.reasoning_quality}")203 test("Total score > 0.5", reward.score > 0.5,204 f"Got: {reward.score}")205 print(f" 📊 Full score: {reward.score:.4f}")206 print(f" Diagnosis: {reward.breakdown.diagnosis_correct}")207 print(f" Fix: {reward.breakdown.fix_valid}")208 print(f" Efficiency:{reward.breakdown.efficiency}")209 print(f" Reasoning: {reward.breakdown.reasoning_quality}")210 print(f" Investig: {reward.breakdown.investigation_quality}")211 print(f" Penalty: -{reward.false_positive_penalty}")212 213# ── 6. Determinism ────────────────────────────────────────────────────────────214print("\n── 6. Determinism ──")215 216# Run same actions twice, should get identical scores217scores = []218for run in range(2):219 env_det = DataOnCallEnv()220 env_det.reset(task_id=1)221 for action in good_actions:222 obs, reward, done, info = env_det.step(action)223 if reward:224 scores.append(reward.score)225 226test("Two runs produce identical scores",227 len(scores) == 2 and scores[0] == scores[1],228 f"Run 1: {scores[0] if scores else 'N/A'}, Run 2: {scores[1] if len(scores) > 1 else 'N/A'}")229 230# ── 7. State endpoint ────────────────────────────────────────────────────────231print("\n── 7. State ──")232 233env7 = DataOnCallEnv()234env7.reset(task_id=2)235state = env7.state()236test("State includes discovered_tables", hasattr(state, "discovered_tables"))237test("State includes cost_spent", hasattr(state, "cost_spent"))238test("State includes budget_remaining", hasattr(state, "budget_remaining"))239 240# After discovering tables, state should reflect it241env7.step(Action(tool="list_tables", query="", reasoning="test"))242state = env7.state()243test("State shows discovered tables after list_tables",244 len(state.discovered_tables) > 0,245 f"Got: {state.discovered_tables}")246 247# ── Summary ───────────────────────────────────────────────────────────────────248print(f"\n{'='*60}")249print(f"RESULTS: {passed} passed, {failed} failed out of {passed + failed} tests")250print('='*60)251 252if failed > 0:253 sys.exit(1)254else:255 print("🎉 All tests passed!")256 sys.exit(0)257 