Team Ai
Apppublic

RahulSinghPundir/SQL_Wizard

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
workflow_json_schema.py190 linesDownload Raw Back to root
1from langchain.chains import LLMChain
2from langchain.prompts import PromptTemplate, FewShotPromptTemplate
3from langchain_community.utilities.sql_database import SQLDatabase
4from langchain_experimental.sql import SQLDatabaseChain
5from langchain.chains.sql_database.prompt import PROMPT_SUFFIX
6from examples import get_example_selector
7from langgraph.graph import StateGraph, END
8from typing import TypedDict, Optional
9from get_tables import FunctionalTools
10from get_llm import get_llm_sql
11from get_memory import get_memory, get_chat_history
12import json
13from flask import session
14
15memory = get_memory()
16chat_history = None
17llm = get_llm_sql()
18
19file = json.load(open("config.json"))
20DB_USER = file["DB_USER"]
21DB_PASSWORD = file["DB_PASSWORD"]
22DB_HOST = file["DB_HOST"]
23DB_NAME = file["DB_NAME"]
24
25class GraphState(TypedDict):
26    message: Optional[str] = None
27    recommended_tables: Optional[str] = None
28    sql_query: Optional[str] = None
29    query_result: Optional[str] = None
30    answer: Optional[str] = None
31
32
33class States:
34
35    @staticmethod
36    def recommendTables(state):
37        print("workflow_json_schema")
38        user_query = state["message"][0]
39        database_details = FunctionalTools.getDatabaseDetailsJson()
40        table_recommendation_prompt = PromptTemplate(
41            input_variables=["user_query", "database_details", "chat_history"],
42            template="""
43                    You are an assistant that recommends relevant database tables for SQL queries.
44
45                    Return the names of ALL the SQL tables ONLY from the given table descriptions. Do not assume any table that might be relevant based on the user question.
46
47                    The tables are:
48                    {database_details}
49
50                    History:
51                    {chat_history}
52
53                    User input:
54                    {user_query}
55
56                    Answer format:
57                    Provide the table names in a comma-separated format ONLY.
58                """
59        )
60        recommendation_chain = LLMChain(llm=llm, prompt=table_recommendation_prompt, output_key="table_list")
61        recommendations = recommendation_chain.run({
62            "user_query": user_query,
63            "database_details": database_details,
64            "chat_history": chat_history
65        }).strip()
66        recommended_tables = [name.strip() for name in recommendations.split(',')]
67        state["recommended_tables"] = recommended_tables
68        return state
69
70    @staticmethod
71    def generateSqlQuery(state):
72        try:
73            history=chat_history['chat_history']
74        except Exception as e:
75            history=chat_history
76        print(history)
77        user_input = state["message"][0]
78        recommended_tables = state.get("recommended_tables", [])
79        table_details = FunctionalTools.getTableDetailsJson(recommended_tables)
80        table_joins = FunctionalTools.getTablesJoins(recommended_tables)
81
82        sql_database = SQLDatabase.from_uri(
83            f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}/{DB_NAME}", 
84            include_tables=recommended_tables, sample_rows_in_table_info=5
85        )
86        sql_query_prefix = f"""
87            You are a MySQL expert. Given an input question, follow these steps to create a syntactically correct MySQL query and provide the final answer:
88            Create a MySQL Query**: Formulate a correct MySQL query to run based on the input question. Ensure the query is syntactically accurate.
89
90            Guidelines:
91            - If the user does not specify a specific number of examples to obtain, use the LIMIT clause to query for at most 5 results.
92            - Order the results to return the most informative data in the database.
93            - Do not query for all columns from a table. Only query the columns necessary to answer the question.
94            - Wrap each column name in backticks (`) to denote them as delimited identifiers.
95            - In the end of sql query put semicolon (;).
96            - Use the CURDATE() function to get the current date if the question involves "today".
97
98            Tables information:
99            {table_details}
100
101            Table joins:
102            {table_joins}
103
104            If any variable is unknown (e.g., name), refer to the history:
105            {history}
106
107            Use the following format:
108            SELECT COUNT(*) FROM `ee_offices` WHERE `is_exist` = 1;
109            """
110        
111        example_prompt = PromptTemplate(
112            input_variables=["input", "query"],
113            template="Question: {input}\nQuery: {query}",
114        )
115
116        example_selector = get_example_selector()
117
118        query_prompt_template = FewShotPromptTemplate(
119            example_selector=example_selector,
120            example_prompt=example_prompt,
121            prefix=sql_query_prefix,
122            suffix=PROMPT_SUFFIX + "SQL Query.",
123            input_variables=["input", "table_details", "top_k"],
124        )
125
126        sql_chain = SQLDatabaseChain.from_llm(
127            llm, sql_database, prompt=query_prompt_template, verbose=True, return_sql=True, use_query_checker=False
128        )
129        sql_execution_result = sql_chain.invoke({
130            "input": user_input,
131            "query": user_input,  # Include the "query" key
132            "table_details": table_details,
133            "top_k": 5  # Assuming top_k is required in the input
134        })
135        sql_query = sql_execution_result["result"]
136        sql_query = sql_query[sql_query.find("SELECT"):sql_query.find(";") + 1]  # Ensures the semicolon is included
137        state["sql_query"] = sql_query
138        return state
139
140    
141    @staticmethod
142    def executeSqlQuery(state):
143        recommended_tables = state["recommended_tables"]
144        sql_query = state["sql_query"]
145        sql_database = SQLDatabase.from_uri(
146            f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}/{DB_NAME}", 
147            include_tables=recommended_tables, sample_rows_in_table_info=5
148        )
149        query_results = sql_database.run(sql_query)
150        state["query_result"] = query_results
151        return state
152
153    @staticmethod
154    def generateAnswer(state):
155        user_input = state["message"][0]
156        sql_query = state["sql_query"]
157        query_result = state["query_result"] or "No result."
158
159        answer_template = """Rephrase the answer, {query_result}, to the question, {user_input}, based on the SQL query, {sql_query}, in a single sentence."""
160        answer_prompt = PromptTemplate(input_variables=["query_result", "user_input", "sql_query"], template=answer_template)
161        chat_chain = LLMChain(llm=llm, prompt=answer_prompt)
162        answer = chat_chain.run({
163            "user_input": user_input, 
164            "sql_query": sql_query, 
165            "query_result": query_result
166        })
167        state["answer"] = answer
168        return state
169
170
171def get_json_workflow(chat_id):
172    global chat_history
173    chat_history = get_chat_history(chat_id)['chat_history']
174
175
176    workflow = StateGraph(GraphState)
177
178    workflow.add_node("recommend_tables", States.recommendTables)
179    workflow.add_node("generate_sql_query", States.generateSqlQuery)
180    workflow.add_node("execute_sql_query", States.executeSqlQuery)
181    workflow.add_node("generate_answer", States.generateAnswer)
182
183    workflow.add_edge("recommend_tables", "generate_sql_query")
184    workflow.add_edge("generate_sql_query", "execute_sql_query")
185    workflow.add_edge("execute_sql_query", "generate_answer")
186    workflow.add_edge("generate_answer", END)
187
188    workflow.set_entry_point("recommend_tables")
189    return workflow.compile()
190