training-monkey/dataoncallenv
0
1"""Baseline inference script for DataOnCallEnv.2"""3 4import os5import json6import time7import httpx8from openai import OpenAI9from environment import DataOnCallEnv10from models import Action11from dotenv import load_dotenv12 13load_dotenv()14 15# Configuration16API_BASE_URL = os.environ.get("API_BASE_URL", "https://router.huggingface.co/v1")17MODEL_NAME = os.environ.get("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct").strip()18 19API_KEY = (20 os.environ.get("HF_TOKEN") or21 os.environ.get("OPENAI_API_KEY")22)23 24if not API_KEY:25 raise EnvironmentError("Set HF_TOKEN or OPENAI_API_KEY before running.")26 27client = OpenAI(api_key=API_KEY, base_url=API_BASE_URL)28 29SUCCESS_SCORE_THRESHOLD = 0.5 # Treat 0.5+ as success for the logger30 31# Validator logger functions32def log_start(task: str, env: str, model: str) -> None:33 print(f"[START] task={task} env={env} model={model}", flush=True)34 35def log_step(step: int, action: str, reward: float, done: bool, error: str) -> None:36 error_val = error if error else "null"37 done_val = str(done).lower()38 # Remove newlines from action string just in case39 action_clean = action.replace("\n", " ").replace("\r", "")40 print(41 f"[STEP] step={step} action={action_clean} reward={reward:.2f} done={done_val} error={error_val}",42 flush=True,43 )44 45def log_end(success: bool, steps: int, score: float, rewards: list) -> None:46 rewards_str = ",".join(f"{r:.2f}" for r in rewards)47 print(f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", flush=True)48 49 50# Tool definitions51TOOLS = [52 {53 "type": "function",54 "function": {55 "name": "run_sql",56 "description": (57 "Execute a SELECT SQL query. "58 "Specify column names — SELECT * is blocked. "59 "Discover tables first with list_tables(). Cost: 2.0"60 ),61 "parameters": {62 "type": "object",63 "properties": {64 "query": {"type": "string", "description": "The SELECT query"},65 "reasoning": {"type": "string", "description": "Why you are running this"},66 },67 "required": ["query", "reasoning"],68 },69 },70 },71 {72 "type": "function",73 "function": {74 "name": "inspect_schema",75 "description": "See column names/types. Discover tables first. Cost: 1.0",76 "parameters": {77 "type": "object",78 "properties": {79 "query": {"type": "string", "description": "Table name"},80 "reasoning": {"type": "string", "description": "Why inspecting"},81 },82 "required": ["query", "reasoning"],83 },84 },85 },86 {87 "type": "function",88 "function": {89 "name": "check_logs",90 "description": "Read the dbt changelog. Cost: 1.0",91 "parameters": {92 "type": "object",93 "properties": {94 "query": {"type": "string", "description": "Pass empty str"},95 "reasoning": {"type": "string"},96 },97 "required": ["query", "reasoning"],98 },99 },100 },101 {102 "type": "function",103 "function": {104 "name": "check_airflow",105 "description": "Read Airflow run history. Cost: 1.0",106 "parameters": {107 "type": "object",108 "properties": {109 "query": {"type": "string", "description": "Pass empty str"},110 "reasoning": {"type": "string"},111 },112 "required": ["query", "reasoning"],113 },114 },115 },116 {117 "type": "function",118 "function": {119 "name": "diff_report",120 "description": "Compare dates. Format: 'YYYY-MM-DD,YYYY-MM-DD'. Cost: 1.5",121 "parameters": {122 "type": "object",123 "properties": {124 "query": {"type": "string", "description": "'date1,date2'"},125 "reasoning": {"type": "string"},126 },127 "required": ["query", "reasoning"],128 },129 },130 },131 {132 "type": "function",133 "function": {134 "name": "list_tables",135 "description": "Discover tables. MUST call this first. Cost: 0.5",136 "parameters": {137 "type": "object",138 "properties": {139 "query": {"type": "string", "description": "Pass empty str"},140 "reasoning": {"type": "string"},141 },142 "required": ["query", "reasoning"],143 },144 },145 },146 {147 "type": "function",148 "function": {149 "name": "submit",150 "description": "End episode. Supply root cause + SQL. Cost: 0.0",151 "parameters": {152 "type": "object",153 "properties": {154 "query": {"type": "string", "description": "ROOT CAUSE: ... CORRECTED SQL: ..."},155 "reasoning": {"type": "string"},156 },157 "required": ["query", "reasoning"],158 },159 },160 },161]162 163SYSTEM_PROMPT = """You are a Data Analyst debugging reports.164RULES:165- Budget: 15 steps, 20.0 cost. Tables HIDDEN until list_tables() used.166- SELECT * is BLOCKED. Specify columns. min 2 steps before submit.167 168STRATEGY:1691. list_tables() and inspect_schema() first.1702. check_logs()/check_airflow() for root causes.1713. run_sql() to confirm hypothesis.1724. submit() immediately when diagnosed.173 174SUBMIT FORMAT:175ROOT CAUSE: [explanation]176CORRECTED SQL: [full SELECT query fixing the bug]"""177 178# Agent loop179 180def run_agent(env: DataOnCallEnv, task_id: int):181 task_name = ["", "weekly_revenue", "monthly_active_users", "revenue_inflation"][task_id]182 183 log_start(task=task_name, env="DataOnCallEnv", model=MODEL_NAME)184 185 obs = env.reset(task_id=task_id)186 187 messages = [188 {"role": "system", "content": SYSTEM_PROMPT},189 {"role": "user", "content": json.dumps(obs.result, indent=2)},190 ]191 192 final_reward = None193 step_rewards = []194 steps = 0195 MAX_STEPS = env.MAX_STEPS196 197 while not env.done and steps < MAX_STEPS:198 199 steps_remaining = MAX_STEPS - steps200 budget_remaining = env.COST_BUDGET - env.cost_spent201 202 if steps_remaining == 5 or budget_remaining <= 5.0:203 messages.append({204 "role": "user",205 "content": f"WARNING: {steps_remaining} steps, budget: {budget_remaining:.1f}. Call submit()NOW if ready."206 })207 elif steps_remaining == 2:208 messages.append({209 "role": "user",210 "content": "FINAL WARNING: 2 steps left. Call submit() next step."211 })212 213 max_retries = 5214 for attempt in range(max_retries):215 try:216 response = client.chat.completions.create(217 model=MODEL_NAME,218 messages=messages,219 tools=TOOLS,220 tool_choice="auto",221 temperature=0,222 )223 break224 except Exception as e:225 err_msg = str(e).lower()226 if "503" in err_msg or "loading" in err_msg:227 time.sleep(15 * (attempt + 1))228 if attempt == max_retries - 1: raise e229 elif "429" in err_msg:230 time.sleep(30)231 if attempt == max_retries - 1: raise e232 else:233 raise e234 235 msg = response.choices[0].message236 237 if msg.tool_calls:238 messages.append(msg)239 for tool_call in msg.tool_calls:240 fn_name = tool_call.function.name241 try:242 fn_args = json.loads(tool_call.function.arguments)243 except Exception:244 fn_args = {}245 246 query = fn_args.get("query", "")247 reasoning = fn_args.get("reasoning", "")248 249 action = Action(tool=fn_name, query=query, reasoning=reasoning)250 obs, reward, done, info = env.step(action)251 steps += 1252 253 # Keep episode running rewards tracked properly254 current_reward_val = reward.score if reward else 0.0255 step_rewards.append(current_reward_val)256 257 if reward: final_reward = reward258 259 # Derive error state for logging260 step_error = None261 if isinstance(obs.result, dict):262 if obs.result.get("error") or not obs.result.get("success", True):263 step_error = str(obs.result.get("error") or obs.result.get("result", "error"))264 step_error = step_error.replace("\n", " ").replace("\r", "")265 266 action_str = f"{fn_name}({query})"267 log_step(steps, action_str, current_reward_val, done, step_error)268 269 messages.append({270 "role": "tool",271 "tool_call_id": tool_call.id,272 "content": json.dumps(obs.result),273 })274 275 if done:276 break277 else:278 steps += 1279 content = msg.content or "No Tool Used"280 log_step(steps, "invalid_output", 0.0, False, "Model did not output a tool call")281 282 step_rewards.append(0.0)283 messages.append({"role": "assistant", "content": content})284 messages.append({285 "role": "user",286 "content": "Use a tool to continue, or call submit() with your answer.",287 })288 289 # Backup grade calculation if ended abruptly290 if final_reward is None:291 from graders import grade292 final_reward = grade(env.task_id, env.conn, env.actions, env.final_answer)293 if len(step_rewards) == 0:294 step_rewards.append(final_reward.score)295 else:296 step_rewards[-1] = final_reward.score297 298 is_success = final_reward.score >= SUCCESS_SCORE_THRESHOLD299 log_end(is_success, steps, final_reward.score, step_rewards)300 301 return {302 "task_id": task_id,303 "score": final_reward.score,304 "steps_taken": steps305 }306 307if __name__ == "__main__":308 env = DataOnCallEnv()309 for task_id in [1, 2, 3]:310 run_agent(env, task_id)