Team Ai
Apppublic

maikheb/nl2sql

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
langchain_sql_pipeline.py87 linesDownload Raw Back to src
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