maikheb/nl2sql
0
1# langchain_sql_pipeline.py
2
3import pickle
4from pathlib import Path
5from dotenv import load_dotenv
6
7from langchain_community.embeddings import HuggingFaceEmbeddings
8from langchain_community.vectorstores import FAISS
9from langchain.prompts import PromptTemplate
10from langchain.chains import LLMChain
11
12from gemini_flash_beta_llm import GeminiFlashBetaLLM
13
14# Load environment variables (for GEMINI_API_KEY, etc.)
15load_dotenv()
16
17# Compute project root (one level above src/)
18ROOT_DIR = Path(__file__).resolve().parent.parent
19VECTORSTORE_DIR = ROOT_DIR / "vectorstore"
20META_PATH = VECTORSTORE_DIR / "schema_meta.pkl"
21
22def load_faiss_retriever():
23 """
24 Load FAISS index with HuggingFace embeddings for column-level retrieval.
25 """
26 if not META_PATH.exists():
27 raise FileNotFoundError(f"FAISS metadata not found at {META_PATH}")
28 with open(META_PATH, "rb") as f:
29 metadata = pickle.load(f)
30
31 texts = [f"{m['table']} - {m['column']}" for m in metadata]
32 embedder = HuggingFaceEmbeddings(model_name="all-MiniLM-L6-v2")
33 db = FAISS.from_texts(texts, embedder)
34 return db.as_retriever(search_type="similarity", k=5), metadata
35
36def generate_sql_with_langchain(user_query: str, schema_dict: dict) -> str:
37 """
38 RAG pipeline: uses FAISS for retrieval of relevant columns,
39 then Gemini-Flash (v1beta) to generate a SQL query.
40 """
41 # 1) Retrieve relevant columns
42 retriever, metadata = load_faiss_retriever()
43 docs = retriever.get_relevant_documents(user_query)
44 semantic_context = "\n".join(d.page_content for d in docs)
45
46 # 2) Format full schema for grounding
47 schema_text = "\n".join(
48 f"Table: {t} — Columns: {', '.join(cols)}"
49 for t, cols in schema_dict.items()
50 )
51
52 # 3) Build prompt template
53 prompt_template = PromptTemplate.from_template("""
54You are a SQL expert. Given the database schema and relevant columns, write a SQL query for the user question.
55
56### DATABASE SCHEMA
57{schema}
58
59### RELEVANT COLUMNS
60{context}
61
62### USER QUESTION
63{question}
64
65Only use valid table and column names. Do not hallucinate.
66
67SQL:
68""")
69
70 # 4) Instantiate the Flash-beta LLM wrapper
71 llm = GeminiFlashBetaLLM()
72
73 # 5) Create and run the chain
74 chain = LLMChain(llm=llm, prompt=prompt_template)
75 return chain.run(schema=schema_text, context=semantic_context, question=user_query)
76
77# ----------------------------------------
78# Optional CLI test
79# ----------------------------------------
80if __name__ == "__main__":
81 example_schema = {
82 "users": ["id", "name", "age"],
83 "orders": ["order_id", "user_id", "amount", "created_at"]
84 }
85 question = "Find all users older than 30 who placed orders over 100"
86 print(generate_sql_with_langchain(question, example_schema))
87 