joelgilbert/NL2SQL
0
1"""2Embedding generation utilities for vector search.3 4Note: Upstash Vector provides built-in embedding generation,5so this module provides helper functions for text preparation.6"""7 8import logging9import re10from typing import Dict, List11 12logger = logging.getLogger(__name__)13 14 15class EmbeddingHelper:16 """Helper functions for preparing text for embedding."""17 18 @staticmethod19 def normalize_text(text: str) -> str:20 """21 Normalize text before embedding generation.22 23 Args:24 text: Raw text to normalize25 26 Returns:27 Normalized text28 """29 # Convert to lowercase30 text = text.lower()31 32 # Remove extra whitespace33 text = re.sub(r'\s+', ' ', text)34 35 # Strip leading/trailing whitespace36 text = text.strip()37 38 return text39 40 @staticmethod41 def embed_schema(table_structure: Dict) -> str:42 """43 Format table schema for embedding.44 45 Args:46 table_structure: Table structure dictionary from schema manager47 48 Returns:49 Formatted text suitable for embedding50 """51 table_name = table_structure.get('table_name', '')52 columns = table_structure.get('columns', [])53 54 # Create descriptive text55 parts = [f"Table {table_name}"]56 57 # Add column information58 if columns:59 column_desc = []60 for col in columns:61 col_name = col.get('column_name', '')62 col_type = col.get('data_type', '')63 column_desc.append(f"{col_name} {col_type}")64 65 parts.append(f"with columns: {', '.join(column_desc)}")66 67 text = " ".join(parts)68 return EmbeddingHelper.normalize_text(text)69 70 @staticmethod71 def embed_question(question: str) -> str:72 """73 Normalize and prepare user question for embedding.74 75 Args:76 question: User's natural language question77 78 Returns:79 Normalized question80 """81 # Basic normalization82 question = EmbeddingHelper.normalize_text(question)83 84 # Remove question marks and punctuation for better matching85 question = re.sub(r'[?!.,;:]', '', question)86 87 return question.strip()88 89 @staticmethod90 def format_example_queries(similar_queries: List[Dict]) -> str:91 """92 Format similar queries into a examples string for prompt.93 94 Args:95 similar_queries: List of similar query dictionaries96 97 Returns:98 Formatted examples text99 """100 if not similar_queries:101 return "No similar examples found."102 103 examples = []104 for i, query in enumerate(similar_queries, 1):105 question = query.get('question', 'N/A')106 sql = query.get('sql', 'N/A')107 examples.append(f"Example {i}:\nQuestion: {question}\nSQL: {sql}\n")108 109 return "\n".join(examples)110 111 @staticmethod112 def expand_abbreviations(text: str) -> str:113 """114 Expand common abbreviations in queries.115 116 Args:117 text: Text with potential abbreviations118 119 Returns:120 Text with expanded abbreviations121 """122 abbreviations = {123 r'\bqty\b': 'quantity',124 r'\bamt\b': 'amount',125 r'\btot\b': 'total',126 r'\bavg\b': 'average',127 r'\bmax\b': 'maximum',128 r'\bmin\b': 'minimum',129 r'\bnum\b': 'number',130 r'\bpct\b': 'percent',131 r'\bid\b': 'identifier',132 }133 134 for abbrev, full in abbreviations.items():135 text = re.sub(abbrev, full, text, flags=re.IGNORECASE)136 137 return text138 139 140# Global embedding helper instance141embedding_helper = EmbeddingHelper()142 