prazy1208/text2sql
0
1"""2Metadata-to-FAISS pipeline for PostgreSQL schemas (healthcare, retail, finance).3Extracts table/column metadata and comments, builds embeddings with sentence-transformers,4and stores them in per-schema FAISS indexes and metadata JSON.5 6Outputs per schema:7- Table-level: faiss_indexes/{schema}.index + metadata_store/{schema}_metadata.json8- Column-level: faiss_indexes/{schema}_columns.index + metadata_store/{schema}_columns_metadata.json9 (flat list order matches FAISS row indices 0..n-1).10 11Run: python build_vector_store.py12"""13 14import json15import logging16import os17from pathlib import Path18 19from dotenv import load_dotenv20from sqlalchemy import create_engine, text21 22# ---------------------------------------------------------------------------23# Logging24# ---------------------------------------------------------------------------25logging.basicConfig(26 level=logging.INFO,27 format="%(asctime)s [%(levelname)s] %(message)s",28 datefmt="%Y-%m-%d %H:%M:%S",29)30logger = logging.getLogger(__name__)31 32load_dotenv()33 34# Output directories35FAISS_INDEX_DIR = Path("faiss_indexes")36METADATA_STORE_DIR = Path("metadata_store")37SCHEMAS = ["healthcare_schema", "retail_schema", "finance_schema"]38 39 40def get_engine():41 """Build SQLAlchemy engine from .env (DATABASE_URL or DB_* variables)."""42 database_url = os.getenv("DATABASE_URL")43 if database_url:44 return create_engine(database_url)45 host = os.getenv("DB_HOST", "localhost")46 port = os.getenv("DB_PORT", "5432")47 user = os.getenv("DB_USER", "postgres")48 password = os.getenv("DB_PASSWORD", "")49 dbname = os.getenv("DB_NAME", "postgres")50 url = f"postgresql://{user}:{password}@{host}:{port}/{dbname}"51 return create_engine(url)52 53 54def extract_metadata(engine, schema_name: str) -> list[dict]:55 """56 Extract schema/table/column metadata and COMMENT ON TABLE / COMMENT ON COLUMN.57 Column types come from pg_catalog.format_type (Postgres-native, includes length/precision).58 Returns a list of dicts, one per table, with keys: schema_name, table_name,59 table_description, columns (list of {name, description, data_type}).60 """61 # Tables with comments (pg_catalog)62 tables_sql = text("""63 SELECT64 c.relname AS table_name,65 COALESCE(pg_catalog.obj_description(c.oid, 'pg_class'), '') AS table_comment66 FROM pg_catalog.pg_class c67 JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace68 WHERE n.nspname = :schema_name69 AND c.relkind = 'r'70 AND c.relname <> 'table_relationships'71 AND c.relname NOT LIKE '%\\_business\\_rules' ESCAPE '\\'72 ORDER BY c.relname73 """)74 # Columns: comments + data type (format_type matches psql / CREATE TABLE spelling)75 columns_sql = text("""76 SELECT77 c.relname AS table_name,78 a.attname AS column_name,79 pg_catalog.format_type(a.atttypid, a.atttypmod) AS data_type,80 COALESCE(pg_catalog.col_description(c.oid, a.attnum), '') AS column_comment81 FROM pg_catalog.pg_attribute a82 JOIN pg_catalog.pg_class c ON c.oid = a.attrelid83 JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace84 WHERE n.nspname = :schema_name85 AND c.relkind = 'r'86 AND c.relname <> 'table_relationships'87 AND c.relname NOT LIKE '%\\_business\\_rules' ESCAPE '\\'88 AND a.attnum > 089 AND NOT a.attisdropped90 ORDER BY c.relname, a.attnum91 """)92 93 with engine.connect() as conn:94 tables = conn.execute(tables_sql, {"schema_name": schema_name}).fetchall()95 columns = conn.execute(columns_sql, {"schema_name": schema_name}).fetchall()96 97 # Group columns by table98 cols_by_table: dict[str, list[dict]] = {}99 for row in columns:100 tbl = row.table_name101 if tbl not in cols_by_table:102 cols_by_table[tbl] = []103 dtype = (row.data_type or "").strip() if row.data_type is not None else ""104 cols_by_table[tbl].append({105 "name": row.column_name,106 "description": (row.column_comment or "").strip(),107 "data_type": dtype,108 })109 110 # Build one record per table111 result = []112 for row in tables:113 table_name = row.table_name114 table_comment = (row.table_comment or "").strip()115 result.append({116 "schema_name": schema_name,117 "table_name": table_name,118 "table_description": table_comment,119 "columns": cols_by_table.get(table_name, []),120 })121 122 logger.info("Extracted metadata for %d tables in %s", len(result), schema_name)123 return result124 125 126def nested_metadata_to_flat_columns(metadata_list: list[dict]) -> list[dict]:127 """128 Flatten nested table metadata to one dict per column for column-level FAISS.129 130 Stable order (documented for rebuild alignment with vectors):131 tables sorted by table_name; within each table, columns sorted by name.132 Each row: schema_name, table_name, column_name, data_type, description,133 table_description, fqn (schema.table.column).134 """135 flat: list[dict] = []136 for m in sorted(metadata_list, key=lambda x: x["table_name"]):137 schema = m["schema_name"]138 table = m["table_name"]139 table_desc = (m.get("table_description") or "").strip()140 for col in sorted(m.get("columns") or [], key=lambda c: c["name"]):141 name = col["name"]142 desc = (col.get("description") or "").strip()143 dtype = (col.get("data_type") or "").strip()144 fqn = f"{schema}.{table}.{name}"145 flat.append({146 "schema_name": schema,147 "table_name": table,148 "column_name": name,149 "data_type": dtype,150 "description": desc,151 "table_description": table_desc,152 "fqn": fqn,153 })154 return flat155 156 157def column_flat_rows_to_texts(flat_rows: list[dict]) -> list[str]:158 """One embedding string per column (metadata only; no row data)."""159 texts = []160 for r in flat_rows:161 tdesc = r["table_description"] or "(no table description)"162 cdesc = r["description"] or "(no description)"163 dtype = r["data_type"] or "(unknown type)"164 text = "\n".join(165 [166 f"Schema: {r['schema_name']}",167 f"Table: {r['table_name']}",168 f"Table description: {tdesc}",169 f"Column: {r['column_name']}",170 f"Type: {dtype}",171 f"Column description: {cdesc}",172 ]173 )174 texts.append(text)175 return texts176 177 178def metadata_to_texts(metadata_list: list[dict]) -> list[str]:179 """180 Turn per-table metadata into structured text chunks suitable for embedding.181 One string per table.182 """183 texts = []184 for m in metadata_list:185 lines = [186 f"Schema: {m['schema_name']}",187 f"Table: {m['table_name']}",188 f"Description: {m['table_description'] or '(no description)'}",189 "Columns:",190 ]191 for col in m["columns"]:192 desc = col["description"] or "(no description)"193 dtype = (col.get("data_type") or "").strip() or "(unknown type)"194 lines.append(f" - {col['name']} ({dtype}): {desc}")195 texts.append("\n".join(lines))196 return texts197 198 199def get_embedding_model(model_name: str = "all-MiniLM-L6-v2"):200 """Load the sentence-transformers model once. Reuse the returned model for all schemas."""201 from sentence_transformers import SentenceTransformer202 203 logger.info("Loading embedding model: %s", model_name)204 return SentenceTransformer(model_name)205 206 207def build_embeddings(texts: list[str], model):208 """209 Generate embeddings for a list of texts using a pre-loaded sentence-transformers model.210 Returns a numpy array of shape (len(texts), embedding_dim).211 """212 logger.info("Encoding %d text(s)", len(texts))213 embeddings = model.encode(texts, show_progress_bar=len(texts) > 10)214 return embeddings215 216 217def build_faiss_index(schema_name: str, embeddings, metadata_list: list[dict]):218 """219 Create a FAISS IndexFlatL2 for the given embeddings, save the index to220 faiss_indexes/{schema}.index and metadata to metadata_store/{schema}_metadata.json.221 """222 import faiss223 import numpy as np224 225 FAISS_INDEX_DIR.mkdir(parents=True, exist_ok=True)226 METADATA_STORE_DIR.mkdir(parents=True, exist_ok=True)227 228 embeddings = np.asarray(embeddings, dtype=np.float32)229 if embeddings.ndim == 1:230 embeddings = embeddings.reshape(1, -1)231 d = embeddings.shape[1]232 index = faiss.IndexFlatL2(d)233 index.add(embeddings)234 n = index.ntotal235 logger.info("Built FAISS index for %s: %d vectors, dim=%d", schema_name, n, d)236 237 index_path = FAISS_INDEX_DIR / f"{schema_name}.index"238 faiss.write_index(index, str(index_path))239 logger.info("Saved FAISS index to %s", index_path)240 241 metadata_path = METADATA_STORE_DIR / f"{schema_name}_metadata.json"242 with open(metadata_path, "w", encoding="utf-8") as f:243 json.dump(metadata_list, f, indent=2, ensure_ascii=False)244 logger.info("Saved metadata to %s", metadata_path)245 246 247def build_column_faiss_index(schema_name: str, embeddings, flat_column_metadata: list[dict]):248 """249 Column-level FAISS: vector i corresponds to flat_column_metadata[i].250 Writes faiss_indexes/{schema}_columns.index and251 metadata_store/{schema}_columns_metadata.json.252 """253 import faiss254 import numpy as np255 256 if not flat_column_metadata:257 logger.warning("No columns to index for %s; skipping column FAISS.", schema_name)258 return259 260 FAISS_INDEX_DIR.mkdir(parents=True, exist_ok=True)261 METADATA_STORE_DIR.mkdir(parents=True, exist_ok=True)262 263 embeddings = np.asarray(embeddings, dtype=np.float32)264 if embeddings.ndim == 1:265 embeddings = embeddings.reshape(1, -1)266 n_vec = embeddings.shape[0]267 if n_vec != len(flat_column_metadata):268 raise ValueError(269 f"Column embedding count ({n_vec}) != flat metadata rows ({len(flat_column_metadata)}) "270 f"for {schema_name}"271 )272 273 d = embeddings.shape[1]274 index = faiss.IndexFlatL2(d)275 index.add(embeddings)276 index_path = FAISS_INDEX_DIR / f"{schema_name}_columns.index"277 faiss.write_index(index, str(index_path))278 logger.info(279 "Built column FAISS index for %s: %d vectors, dim=%d -> %s",280 schema_name,281 index.ntotal,282 d,283 index_path,284 )285 286 metadata_path = METADATA_STORE_DIR / f"{schema_name}_columns_metadata.json"287 with open(metadata_path, "w", encoding="utf-8") as f:288 json.dump(flat_column_metadata, f, indent=2, ensure_ascii=False)289 logger.info("Saved column flat metadata to %s", metadata_path)290 291 292def run_pipeline_for_schema(engine, schema_name: str, model):293 """Extract metadata, build embeddings, and write FAISS index + metadata for one schema."""294 logger.info("Processing schema: %s", schema_name)295 metadata_list = extract_metadata(engine, schema_name)296 if not metadata_list:297 logger.warning("No tables found in %s, skipping.", schema_name)298 return299 texts = metadata_to_texts(metadata_list)300 embeddings = build_embeddings(texts, model)301 build_faiss_index(schema_name, embeddings, metadata_list)302 303 flat_columns = nested_metadata_to_flat_columns(metadata_list)304 if flat_columns:305 col_texts = column_flat_rows_to_texts(flat_columns)306 col_embeddings = build_embeddings(col_texts, model)307 build_column_faiss_index(schema_name, col_embeddings, flat_columns)308 else:309 logger.warning("No column rows after flatten for %s; skipping column index.", schema_name)310 311 312def main():313 logger.info("Starting metadata-to-FAISS pipeline.")314 engine = get_engine()315 model = get_embedding_model()316 for schema_name in SCHEMAS:317 run_pipeline_for_schema(engine, schema_name, model)318 logger.info("Pipeline finished.")319 320 321if __name__ == "__main__":322 main()323 