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)}"]
        }