Team Ai
Apppublic

Aditya9552/chat_with_sql_database_using_llm

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
streamlit_app.py505 linesDownload Raw Back to root
1import os2import sqlite33import streamlit as st4import pandas as pd5from dotenv import load_dotenv6from typing import TypedDict, Annotated, List, Literal, Optional7import operator8 9# --- New Imports for Gemini ---10# Assuming you have installed: pip install langchain-google-genai11from langchain_google_genai import ChatGoogleGenerativeAI12# -----------------------------13 14from langchain_groq import ChatGroq15from langchain_community.utilities import SQLDatabase16from langchain_community.agent_toolkits import SQLDatabaseToolkit17from langchain_core.messages import AIMessage, HumanMessage, ToolMessage, SystemMessage18from langchain_core.messages.tool import ToolCall19from langgraph.graph import StateGraph, END, START20from langgraph.prebuilt import ToolNode21from langchain_core.language_models import BaseChatModel22 23# Load API keys from .env file24load_dotenv()25GROQ_API_KEY = os.getenv("GROQ_API_KEY")26# Assuming GOOGLE_API_KEY is also set in your .env file27GOOGLE_API_KEY = os.getenv("GOOGLE_API_KEY")28 29# --- Streamlit Configuration ---30st.set_page_config(page_title="Chat with SQL Database", page_icon="💬", layout="wide")31 32# Custom CSS for styling the app33st.markdown(34    """35    <style>36    .main-header {37        font-size: 3em;38        font-weight: bold;39        color: #4CAF50;40        text-align: center;41        margin-bottom: 30px;42        text-shadow: 2px 2px 5px rgba(0,0,0,0.2);43    }44    .stChatInput, .stFileUploader {45        border-radius: 15px;46        padding: 10px;47        box-shadow: 2px 2px 10px rgba(0,0,0,0.1);48    }49    .stButton>button {50        background-color: #4CAF50;51        color: white;52        border-radius: 10px;53        padding: 10px 20px;54        font-size: 1.1em;55        box-shadow: 2px 2px 5px rgba(0,0,0,0.2);56        width: 100%;57    }58    .stButton>button:hover {59        background-color: #45a049;60    }61    .stChatMessage {62        border-radius: 10px;63        padding: 1rem;64        margin-bottom: 1rem;65    }66    .stChatMessage.user {67        background-color: #e0f7fa;68    }69    .stChatMessage.assistant {70        background-color: #f1f8e9;71    }72    </style>73    """,74    unsafe_allow_html=True,75)76 77st.markdown('<h1 class="main-header">💬 Chat with Your SQL Database</h1>', unsafe_allow_html=True)78st.write("Upload your `.db` file and ask questions in natural language to get insights from your data!")79 80# --- LangGraph State and Node Definitions ---81 82class MessagesState(TypedDict):83    """Represents the state of the graph, holding a list of messages.84    Messages are appended using operator.add."""85    messages: Annotated[List[AIMessage | HumanMessage | ToolMessage], operator.add]86 87# Global variables for LLM and DB instances, initialized after file upload.88# Type hint BaseChatModel for LLM89llm: Optional[BaseChatModel] = None90db: Optional[SQLDatabase] = None91toolkit = None92tools = []93get_schema_tool = None94run_query_tool = None95list_tables_tool = None96agent = None # LangGraph agent instance97 98# --- LLM Selection and Initialization Logic ---99 100def initialize_llm_and_tools(selected_llm_model: str):101    """Initializes the LLM, SQL toolkit, and necessary tools once the DB is loaded."""102    global llm, toolkit, tools, get_schema_tool, run_query_tool, list_tables_tool103    104    # 1. Initialize LLM based on selection105    try:106        if selected_llm_model == "Llama-4-Maverick (Groq)":107            if not GROQ_API_KEY:108                st.error("GROQ_API_KEY is not set in environment variables.")109                st.stop()110            # Initialize ChatGroq LLM111            llm = ChatGroq(model="meta-llama/llama-4-maverick-17b-128e-instruct", temperature=0)112        113        elif selected_llm_model == "Gemini Flash 2.5 (Google)":114            if not GOOGLE_API_KEY:115                st.error("GOOGLE_API_KEY is not set in environment variables.")116                st.stop()117            # Initialize ChatGoogleGenerativeAI LLM118            # Use the specified model119            llm = ChatGoogleGenerativeAI(model="gemini-2.5-flash", temperature=0, api_key=GOOGLE_API_KEY)120        121        else:122            st.error("Invalid LLM model selection.")123            st.stop()124    except Exception as e:125        st.error(f"Failed to initialize LLM: {e}")126        st.stop()127 128    # 2. Initialize SQLDatabaseToolkit with the loaded database and LLM129    toolkit = SQLDatabaseToolkit(db=db, llm=llm)130    tools = toolkit.get_tools()131 132    # 3. Extract specific tools required for the workflow133    get_schema_tool = next((tool for tool in tools if tool.name == "sql_db_schema"), None)134    run_query_tool = next((tool for tool in tools if tool.name == "sql_db_query"), None)135    list_tables_tool = next((tool for tool in tools if tool.name == "sql_db_list_tables"), None)136 137    # 4. Validate that all required tools are successfully initialized138    if not all([get_schema_tool, run_query_tool, list_tables_tool]):139        st.error("Failed to initialize all necessary SQL tools. Please ensure the database toolkit provides 'sql_db_schema', 'sql_db_query', and 'sql_db_list_tables'.")140        st.stop()141 142# --- Rest of the LangGraph Nodes (Unchanged) ---143 144def list_tables_node(state: MessagesState):145    print("--- Executing list_tables_node ---")146    if not list_tables_tool:147        return {"messages": state["messages"] + [AIMessage(content="Error: list_tables_tool not initialized.")]}148    try:149        tool_result = list_tables_tool.invoke({})150    except Exception as e:151        return {"messages": state["messages"] + [AIMessage(content=f"Tool failed: {e}")]}152 153    tool_call_msg = AIMessage(154        content="",155        tool_calls=[ToolCall(name=list_tables_tool.name, args={}, id="list_tables_tool_call_id")]156    )157    tool_msg = ToolMessage(158        content=str(tool_result),159        tool_call_id="list_tables_tool_call_id",160        name=list_tables_tool.name161    )162    response = AIMessage(content=f"Available tables: {tool_result}")163    return {"messages": state["messages"] + [tool_call_msg, tool_msg, response]}164 165def call_get_schema(state: MessagesState):166    if not get_schema_tool:167        return {"messages": state["messages"] + [AIMessage(content="Error: get_schema_tool not initialized.")]}168    llm_with_tools = llm.bind_tools([get_schema_tool], tool_choice="any")169    try:170        response = llm_with_tools.invoke(state["messages"] + [HumanMessage(content="Please provide the schema for the relevant tables.")])171    except Exception as e:172        return {"messages": state["messages"] + [AIMessage(content=f"Schema tool failed: {e}")]}173    return {"messages": state["messages"] + [response]}174 175# ---- Prompts ----176 177generate_query_system_prompt = f"""You are a highly skilled SQL database agent. Your primary goal is to answer user questions by formulating and executing precise sqlite queries.178 179Follow these strict guidelines:180 1811.  **Understand the User's Intent:** Carefully analyze the user's question to determine the exact data required.1822.  **Generate SQL Query:**183    * Produce a **syntactically correct** sqlite query.184    * **Prioritize relevance:** Only select columns directly relevant to the user's question. **Never use `SELECT *`**.185    * **Order results:** Use `ORDER BY` clause to present the most relevant or interesting results, if applicable.1863.  **Manage Result Quantity (LIMIT Clause):**187    * By default, **limit queries to a maximum of 5 results** (`LIMIT 5`).188    * **Crucial Exception:** If the user explicitly asks for a specific number of results (e.g., "show 10", "get everything") OR if the question implies a need for all results (e.g., contains keywords like "each", "per", "all", "every"), then you **MUST adjust or remove the `LIMIT` clause accordingly**. Always respect the user's explicit quantity request over the default.1894.  **Database Safety:**190    * **Absolutely DO NOT** perform any data manipulation (DML) or schema modification operations. This includes `INSERT`, `UPDATE`, `DELETE`, `DROP`, `ALTER`, etc. Only generate `SELECT` queries.1915.  **Multi-Step Questions (Internal Note for Advanced Agents):**192    * If the user's question requires multiple separate SQL queries to fully answer (e.g., "First find X, then use X to find Y"), generate only the query for the *first logical step*. The system will handle subsequent steps. (This assumes you have a `decide_next_step` node or similar; otherwise, remove this point).193 194**Example Scenario for LIMIT:**195* User: "List all customers" -> `SELECT CustomerId, FirstName, LastName FROM Customer;` (No LIMIT)196* User: "Show 3 products" -> `SELECT ProductId, Name FROM Product LIMIT 3;`197* User: "How many tracks are there per album?" -> `SELECT AlbumId, COUNT(TrackId) FROM Track GROUP BY AlbumId;` (No LIMIT because "per" implies all)198* User: "Who are the customers?" -> `SELECT CustomerId, FirstName, LastName FROM Customer LIMIT 5;` (Default LIMIT)199 200Now, generate the SQL query for the user's request.201"""202 203def generate_query(state: MessagesState):204    if not run_query_tool:205        return {"messages": state["messages"] + [AIMessage(content="Error: run_query_tool not initialized.")]}206    system_msg = SystemMessage(content=generate_query_system_prompt)207    llm_with_tools = llm.bind_tools([run_query_tool])208    try:209        response = llm_with_tools.invoke([system_msg] + state["messages"])210    except Exception as e:211        return {"messages": state["messages"] + [AIMessage(content=f"Query generation failed: {e}")]}212    return {"messages": state["messages"] + [response]}213 214check_query_system_prompt = f"""215You are a SQL expert with a strong attention to detail.216Double check the sqlite query for common mistakes, including:217- Using NOT IN with NULL values218- Using UNION when UNION ALL should have been used219- Using BETWEEN for exclusive ranges220- Data type mismatch in predicates221- Properly quoting identifiers222- Using the correct number of arguments for functions223- Casting to the correct data type224- Using the proper columns for joins225 226If there are any of the above mistakes, rewrite the query. If there are no mistakes,just reproduce the original query.227"""228 229def check_query(state: MessagesState):230    system_msg = SystemMessage(content=check_query_system_prompt)231    last_message = state["messages"][-1]232 233    query = None234    if isinstance(last_message, AIMessage):235        if last_message.tool_calls:236            tool_call = last_message.tool_calls[0]237            # Handle both dicts and objects238            if isinstance(tool_call, dict):239                query = tool_call.get("args", {}).get("query")240            else:241                query = getattr(tool_call, "args", {}).get("query")242        elif "SELECT" in last_message.content.upper():243            query = last_message.content.strip()244 245    if query and llm and run_query_tool:246        user_msg = HumanMessage(content=query, role="user")247        llm_with_tools = llm.bind_tools([run_query_tool], tool_choice="any")248        try:249            response = llm_with_tools.invoke([system_msg, user_msg])250        except Exception as e:251            # Fallback to initial query if check fails252            return {"messages": state["messages"] + [AIMessage(content=f"Query check failed: {e}. Proceeding with original query.")]}253        return {"messages": state["messages"] + [response]}254 255    print("Fallback: regenerating query or no query found.")256    return generate_query(state)257 258 259def interpret_query_result(state: MessagesState):260    last_msg = state["messages"][-1]261    user_q = next((msg.content for msg in reversed(state["messages"]) if isinstance(msg, HumanMessage)), "the user question")262 263    if isinstance(last_msg, ToolMessage) and llm:264        query_result = str(last_msg.content).strip() # .strip() to handle whitespace from empty results265 266        # --- IMPORTANT MODIFICATION HERE ---267        if not query_result or query_result.lower() in ["[]", "no results"]:268            # Explicitly tell the LLM what to do for empty results269            interpretation_prompt = f"""You are a professional data analyst assistant.270 271Your task is to interpret the results of a SQL query and respond in natural language.272 273**User Question:** {user_q}274**SQL Query Results:** [No results returned for the query.]275 276Based on the query results, there are no records matching your request. State clearly that no information was found for the given question. Be brief and helpful.277 278Answer:279"""280        else:281            # Original prompt for actual results282            interpretation_prompt = f"""You are a professional data analyst assistant.283 284Your task is to interpret the results of a SQL query and respond in natural language.285 286- Be brief, clear, and helpful.287- Provide a concise summary, mentioning relevant numbers or trends.288- If the query results contain multiple rows or columns:289  - Format them as a **well-aligned Markdown table**290  - **Pad all cells** to keep columns aligned291  - Ensure that the `|` separators line up correctly, even with long text values292- Always reflect the original user's question context.293- **CRITICAL INSTRUCTION:**294  - **ONLY answer based on the provided "SQL Query Results".**295  - **Do NOT make any assumptions or invent information.**296  - **Do NOT provide hypothetical examples or assume missing data.**297 298**User Question:** {user_q}299**SQL Query Results:** {query_result}300 301Answer:302"""303        try:304            llm_response = llm.invoke([HumanMessage(content=interpretation_prompt)])305        except Exception as e:306            return {"messages": state["messages"] + [AIMessage(content=f"Interpretation failed: {e}")]}307        return {"messages": state["messages"] + [AIMessage(content=llm_response.content)]}308 309    return {"messages": state["messages"] + [AIMessage(content="Unable to interpret query result.")]}310 311 312def should_continue(state: MessagesState) -> Literal[END, "check_query", "interpret_query_result"]:313    last = state["messages"][-1]314 315    # Check for AIMessage with a tool_call to sql_db_query316    if isinstance(last, AIMessage) and last.tool_calls:317        for tc in last.tool_calls:318            tool_name = getattr(tc, "name", tc.get("name") if isinstance(tc, dict) else None)319            if tool_name == "sql_db_query":320                return "check_query"321 322    # Check for raw SQL content in an AIMessage (used by check_query for the tool call)323    if isinstance(last, AIMessage) and "SELECT" in last.content.upper() and "FROM" in last.content.upper():324        print("Detected raw SQL content.")325        # We assume raw SQL output from check_query should go to run_query, 326        # but LangGraph needs an explicit ToolCall to use ToolNode. 327        # Since 'check_query' is the node, we proceed to 'run_query' only if 328        # the previous step was 'check_query' and it returned a ToolCall (handled above).329        # This case here (raw SQL in AIMessage) is primarily a fallback/diagnostic and might be 330        # better handled by ensuring 'check_query' always returns an AIMessage with a ToolCall or a modified message.331        pass # Let the next condition handle ToolMessage332 333    # Check for a ToolMessage (result of run_query)334    if isinstance(last, ToolMessage):335        return "interpret_query_result"336 337    return END338 339def build_langgraph_agent():340    """Compiles the LangGraph agent workflow."""341    global agent342    if not all([get_schema_tool, run_query_tool]):343        st.error("Cannot build agent: Required tools are missing.")344        return345        346    builder = StateGraph(MessagesState)347    builder.add_node("list_tables", list_tables_node)348    builder.add_node("call_get_schema", call_get_schema)349    builder.add_node("get_schema", ToolNode([get_schema_tool]))350    builder.add_node("generate_query", generate_query)351    builder.add_node("check_query", check_query)352    # The ToolNode expects an iterable of tools353    builder.add_node("run_query", ToolNode([run_query_tool])) 354    builder.add_node("interpret_query_result", interpret_query_result)355 356    builder.add_edge(START, "list_tables")357    builder.add_edge("list_tables", "call_get_schema")358    builder.add_edge("call_get_schema", "get_schema")359    builder.add_edge("get_schema", "generate_query")360 361    builder.add_conditional_edges(362        "generate_query", should_continue, {"check_query": "check_query", END: END}363    )364    # The output of check_query is expected to be an AIMessage with a sql_db_query ToolCall365    builder.add_edge("check_query", "run_query") 366    367    # The output of run_query is a ToolMessage (tool result)368    builder.add_conditional_edges(369        "run_query", should_continue, {"interpret_query_result": "interpret_query_result", END: END}370    )371    builder.add_edge("interpret_query_result", END)372 373    agent = builder.compile()374 375 376# --- Streamlit App Logic ---377 378# Sidebar for file upload and controls379with st.sidebar:380    st.subheader("Database Setup")381    uploaded_file = st.file_uploader("Upload your SQLite Database (.db)", type=["db", "sqlite", "sqlite3"])382 383    st.subheader("LLM Configuration")384    # --- NEW: LLM Selection Box ---385    selected_llm_model = st.selectbox(386        "Choose LLM Model:",387        ("Llama-4-Maverick (Groq)", "Gemini Flash 2.5 (Google)"),388        key="selected_llm_model"389    )390    # -----------------------------391 392    # Button to clear the chat history393    if st.button("Clear Chat History"):394        st.session_state.messages = []   # Clear agent's internal history395        st.session_state.display_messages = [] # Clear user-facing history396        st.session_state.db_loaded = False # Re-trigger setup on next run397        st.rerun() # Rerun the app to reflect changes398 399# Handle the initial file upload OR LLM change.400# If the file changes, or the LLM selection changes, we must re-initialize.401current_llm_key = st.session_state.get("selected_llm_model")402 403# Check if a new file is uploaded OR if the LLM model selection has changed since last load404if uploaded_file is not None and (st.session_state.get("uploaded_file_name") != uploaded_file.name or st.session_state.get("last_llm_key") != current_llm_key):405    406    with st.spinner(f"Loading Database and Initializing with {selected_llm_model}..."):407        # Save the uploaded file to a temporary path in /tmp/408        temp_db_filename = f"temp_{uploaded_file.name}"409        temp_db_path = os.path.join("/tmp", temp_db_filename)410        411        with open(temp_db_path, "wb") as f:412            f.write(uploaded_file.getbuffer())413        414        # Store the database path and file name in session state415        st.session_state["db_path"] = f"sqlite:///{temp_db_path}"416        st.session_state["uploaded_file_name"] = uploaded_file.name417        st.session_state["last_llm_key"] = current_llm_key # Store the current selection418        st.session_state["db_loaded"] = True419        420        # Reset chat history for the new database/model421        st.session_state.messages = []422        st.session_state.display_messages = []423        424        st.sidebar.success(f"Database loaded! Model: {selected_llm_model}")425        # Rerun to initialize agent outside the file-upload block426        st.rerun() 427 428 429# If the database is loaded, re-initialize the agent on every script run430if st.session_state.get("db_loaded", False) and st.session_state.get("last_llm_key") == current_llm_key:431    try:432        # Load the SQL database from the stored path433        db = SQLDatabase.from_uri(st.session_state["db_path"])434        # Initialize LLM and tools with the selected model435        initialize_llm_and_tools(selected_llm_model) 436        build_langgraph_agent()      # Build/rebuild the LangGraph agent437        438        # Display available tables in a sidebar expander439        with st.sidebar:440            with st.expander("Tables in Database"):441                try:442                    table_names = db.get_table_names()443                    st.write(", ".join(table_names))444                except Exception as e:445                    st.error(f"Could not list tables: {e}")446 447    except Exception as e:448        st.error(f"Failed to initialize the agent. Please try reloading the file or check API key. Error: {e}")449        st.session_state["db_loaded"] = False # Mark database as not loaded on error450 451 452# Display chat messages453for i, message in enumerate(st.session_state.get("display_messages", [])):454    role = "user" if isinstance(message, HumanMessage) else "assistant"455    with st.chat_message(role):456        if role == "assistant":457            # For assistant messages, use an expander, hidden by default to avoid clutter458            with st.expander("Response to your question", expanded=True): # Changed to expanded=True for better visibility459                st.markdown(message.content)460        else:461            # For user messages, just display the content462            st.markdown(message.content)463 464# Accept user input465if prompt := st.chat_input("Ask your question about the database..."):466    # Check if a database is loaded and the agent is initialized467    if not st.session_state.get("db_loaded", False) or agent is None:468        st.warning("Please upload a database file and ensure the agent is initialized.")469    else:470        # Add user message to user-facing history and display it immediately471        user_message = HumanMessage(content=prompt)472        st.session_state.display_messages.append(user_message)473        with st.chat_message("user"):474            st.markdown(prompt)475 476        # IMPORTANT: For the agent, pass ONLY the *current* user message.477        st.session_state.messages = [user_message] # Reset agent's internal history to only the current prompt478 479        # Invoke the agent and stream the response480        with st.chat_message("assistant"):481            with st.spinner(f"Asking {selected_llm_model} to analyze the database..."):482                try:483                    final_result_state = None484                    # Stream the agent's execution.485                    for s in agent.stream({"messages": st.session_state.messages}, stream_mode="values"):486                        final_result_state = s487 488                    if final_result_state and "messages" in final_result_state:489                        # Update LangGraph's internal trace for the current turn490                        st.session_state.messages = final_result_state["messages"]491 492                        # Get the final AI response (last message in the state) and add it to the display history493                        ai_response = final_result_state["messages"][-1]494                        if isinstance(ai_response, AIMessage):495                            st.markdown(ai_response.content)496                            st.session_state.display_messages.append(ai_response)497                        else:498                            st.error("Received an unexpected response type from the agent.")499                    else:500                        st.error("Agent did not return a valid final state.")501 502                except Exception as e:503                    error_message = f"An error occurred during agent execution: {e}"504                    st.error(error_message)505                    st.session_state.display_messages.append(AIMessage(content=error_message))