Spaces:
Sleeping
Sleeping
File size: 3,269 Bytes
cf796c5 f9e8eed cf796c5 8d73fdc f9e8eed cf796c5 f9e8eed cf796c5 f9e8eed cf796c5 f9e8eed cf796c5 f9e8eed d2c5868 f9e8eed d2c5868 f9e8eed d2c5868 f9e8eed d2c5868 f9e8eed d2c5868 f9e8eed d2c5868 f9e8eed cf796c5 f9e8eed cf796c5 f9e8eed | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 | import sqlite3
import os
from src.state import AgentState
DB_PATH = os.getenv("DB_PATH", "data/olist.db")
MAX_RETRIES = 3
def execute_sql(state: AgentState) -> dict:
"""Execute SQL query with retry logic"""
# Check if safety checker already flagged an error
sql_result = state.get("sql_result")
if sql_result and isinstance(sql_result, dict) and sql_result.get("error"):
return {} # safety checker already flagged this
query = state.get("sql_query")
if not query:
return {"sql_result": {"error": "no_query_generated"}}
retry = state.get("retry_count", 0)
if retry >= MAX_RETRIES:
return {
"sql_result": {"error": "max_retries_exceeded"},
"error_log": state.get("error_log", []) + ["[sql_executor] Max retries hit, giving up"]
}
try:
conn = sqlite3.connect(DB_PATH)
conn.row_factory = sqlite3.Row
cursor = conn.execute(query)
rows = [dict(row) for row in cursor.fetchall()]
conn.close()
if not rows:
return {"sql_result": {"rows": [], "empty": True}, "retry_count": 0}
else:
result = {"sql_result": {"rows": rows, "count": len(rows)}, "retry_count": 0}
# Persist entities from query results for next turn
row = rows[0]
session_context_updates = {}
if row.get("order_id"):
result["last_order_id"] = row["order_id"]
session_context_updates["order_id"] = row["order_id"]
print(f"[sql_executor] Persisted order_id: {row['order_id']}")
if row.get("seller_id"):
result["last_seller_id"] = row["seller_id"]
session_context_updates["seller_id"] = row["seller_id"]
print(f"[sql_executor] Persisted seller_id: {row['seller_id']}")
# Handle comma-separated seller_ids (from aggregated queries)
if row.get("seller_ids"):
first_seller = row["seller_ids"].split(",")[0].strip()
result["last_seller_id"] = first_seller
session_context_updates["seller_id"] = first_seller
print(f"[sql_executor] Persisted first seller_id: {first_seller}")
if row.get("product_id"):
result["last_product_id"] = row["product_id"]
session_context_updates["product_id"] = row["product_id"]
print(f"[sql_executor] Persisted product_id: {row['product_id']}")
if session_context_updates:
result["session_context"] = {**state.get("session_context", {}), **session_context_updates}
return result
except sqlite3.Error as e:
return {
"retry_count": retry + 1,
"sql_result": {"error": str(e), "failed_query": query},
"error_log": state.get("error_log", []) + [f"[sql_executor] SQL error (attempt {retry+1}): {str(e)}"]
}
except Exception as e:
return {
"sql_result": {"error": f"unexpected_error: {str(e)}"},
"error_log": state.get("error_log", []) + [f"[sql_executor] Unexpected error: {str(e)}"]
}
|