CoolFace
Apppublic

training-monkey/dataoncallenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
test_env.py257 linesDownload Raw Back to root
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