RahulSinghPundir/SQL_Wizard
0
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 