File size: 3,776 Bytes
cf796c5
b6ae869
cf796c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f9e8eed
cf796c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8d73fdc
 
 
 
 
 
 
 
cf796c5
 
b6ae869
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
from src.utils.prompt_templates import SQL_SYSTEM
from src.utils.llm_factory import get_llm, _invoke_with_backoff
from src.state import AgentState
from langchain_core.messages import HumanMessage

# Dynamic schema injection: only give LLM the tables it needs
SUB_INTENT_SCHEMA_MAP = {
    "order_tracking": ["v_order_full"],
    "price_query": ["v_order_full"],
    "product_recommendation": ["products", "category_translation", "order_items", "order_reviews"],
    "seller_review": ["v_seller_summary", "order_reviews"],
    "delivery_date": ["v_order_full"],
    "default": ["v_order_full", "products", "v_seller_summary", "category_translation"],
}

SCHEMA_DEFINITIONS = {
    "v_order_full": """
        v_order_full(order_id TEXT, customer_id TEXT, order_status TEXT,
        order_purchase_timestamp TEXT, order_approved_at TEXT,
        order_delivered_carrier_date TEXT, order_delivered_customer_date TEXT,
        order_estimated_delivery_date TEXT, total_amount REAL,
        item_count INT, seller_ids TEXT, payment_methods TEXT,
        is_late INT, days_overdue REAL)""",
    "v_seller_summary": """
        v_seller_summary(seller_id TEXT, seller_city TEXT, seller_state TEXT,
        total_orders INT, avg_rating REAL, negative_reviews INT, positive_reviews INT)""",
    "products": """
        products(product_id TEXT, product_category_name TEXT, product_name_length REAL,
        product_description_length REAL, product_photos_qty REAL, product_weight_g REAL,
        product_length_cm REAL, product_height_cm REAL, product_width_cm REAL)""",
    "category_translation": """
        category_translation(product_category_name TEXT, category_name_english TEXT)""",
    "order_items": """
        order_items(order_id TEXT, order_item_id INT, product_id TEXT, seller_id TEXT,
        shipping_limit_date TEXT, price REAL, freight_value REAL)""",
    "order_reviews": """
        order_reviews(review_id TEXT, order_id TEXT, review_score INT,
        review_comment_title TEXT, review_comment_message TEXT,
        review_creation_date TEXT, review_answer_timestamp TEXT)""",
}

def generate_sql(state: AgentState) -> dict:
    """Generate SQL query with context-aware schema injection"""
    try:
        sub = state.get("sub_intents", ["default"])
        relevant_tables = []
        for s in sub:
            relevant_tables.extend(SUB_INTENT_SCHEMA_MAP.get(s, SUB_INTENT_SCHEMA_MAP["default"]))

        schema_str = "\n".join(SCHEMA_DEFINITIONS.get(t, "") for t in set(relevant_tables))
        prompt = SQL_SYSTEM.format(
            schema=schema_str,
            question=state.get("user_input", ""),
            order_id=state.get("last_order_id", "NULL"),
            seller_id=state.get("last_seller_id", "NULL"),
            category=state.get("last_category", "NULL"),
        )
        
        # If this is a retry, inject the previous error to help LLM fix the query
        retry_count = state.get("retry_count", 0)
        if retry_count > 0:
            last_error = state.get("sql_result", {}).get("error", "")
            failed_query = state.get("sql_result", {}).get("failed_query", "")
            if last_error and failed_query:
                prompt += f"\n\nPREVIOUS ATTEMPT FAILED (attempt {retry_count}):\nQuery: {failed_query}\nError: {last_error}\n\nPlease fix the query to resolve this error."

        llm = get_llm(temperature=0.0)
        response = _invoke_with_backoff(llm, [HumanMessage(content=prompt)], provider="groq")
        sql = response.content.strip().replace("```sql", "").replace("```", "").strip()
        return {"sql_query": sql}
    except Exception as e:
        return {
            "sql_query": None,
            "error_log": state.get("error_log", []) + [f"[sql_generator] Error: {str(e)}"]
        }