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