joelgilbert/NL2SQL
0
1"""2Agent 1: Gatekeeper - Intent classification and question validation.3"""4 5import logging6import json7import re8from typing import Dict, Tuple9from datetime import datetime, timedelta10from groq import Groq11from tenacity import retry, stop_after_attempt, wait_exponential12 13from config.settings import settings14from config.prompts import PromptTemplates15 16logger = logging.getLogger(__name__)17 18 19class GatekeeperAgent:20 """Agent responsible for intent classification and question validation."""21 22 def __init__(self):23 """Initialize gatekeeper agent with Groq API client."""24 self.client = Groq(api_key=settings.api.groq_api_key)25 self.model = "llama-3.1-8b-instant"26 27 @retry(stop=stop_after_attempt(3), wait=wait_exponential(multiplier=1, min=2, max=10))28 def classify_intent(self, message: str) -> Dict[str, any]:29 """30 Classify user intent using Llama-3.1-8B via Groq.31 32 Args:33 message: User's input message34 35 Returns:36 Dictionary with:37 - intent: str (greeting, data_query, vague_question, off_topic)38 - confidence: float (0.0-1.0)39 - needs_clarification: bool40 - response: str (response message to user)41 """42 try:43 completion = self.client.chat.completions.create(44 model=self.model,45 messages=[46 {"role": "system", "content": PromptTemplates.GATEKEEPER_SYSTEM},47 {"role": "user", "content": message}48 ],49 temperature=0.3,50 max_tokens=500,51 response_format={"type": "json_object"}52 )53 54 response_text = completion.choices[0].message.content55 result = json.loads(response_text)56 57 logger.info(f"Intent classified as: {result.get('intent')} with confidence {result.get('confidence')}")58 return result59 60 except json.JSONDecodeError as e:61 logger.error(f"Failed to parse JSON response: {e}")62 # Fallback to rule-based classification63 return self._fallback_classification(message)64 65 except Exception as e:66 logger.error(f"Groq API call failed: {e}")67 # Fallback to rule-based classification68 return self._fallback_classification(message)69 70 def _fallback_classification(self, message: str) -> Dict[str, any]:71 """72 Rule-based fallback classification when API fails.73 74 Args:75 message: User's input message76 77 Returns:78 Classification dictionary79 """80 message_lower = message.lower().strip()81 82 # Check for greetings83 greetings = ['hello', 'hi', 'hey', 'good morning', 'good afternoon', 'good evening']84 if any(word in message_lower for word in greetings):85 return {86 "intent": "greeting",87 "confidence": 0.8,88 "needs_clarification": False,89 "response": "Hello! I'm your SQL assistant. I can help you query your database using natural language. What would you like to know about your data?"90 }91 92 # Check for data-related keywords93 data_keywords = ['show', 'get', 'find', 'list', 'count', 'total', 'sum', 'average', 94 'select', 'query', 'data', 'table', 'records', 'customers', 'orders', 95 'sales', 'revenue', 'users']96 97 if any(keyword in message_lower for keyword in data_keywords):98 # Check if question is specific enough99 if len(message.split()) < 4:100 return {101 "intent": "vague_question",102 "confidence": 0.7,103 "needs_clarification": True,104 "response": "Could you be more specific? Please provide more details about what data you're looking for."105 }106 107 return {108 "intent": "data_query",109 "confidence": 0.7,110 "needs_clarification": False,111 "response": "I'll help you with that query."112 }113 114 # Default to off-topic115 return {116 "intent": "off_topic",117 "confidence": 0.6,118 "needs_clarification": False,119 "response": "I can only help with database queries. Please ask a question about your data."120 }121 122 def validate_question(self, question: str) -> Tuple[bool, str]:123 """124 Validate if question is specific enough for SQL generation.125 126 Args:127 question: User's question128 129 Returns:130 Tuple of (is_valid, clarification_message)131 """132 question = question.strip()133 134 # Check minimum length135 if len(question) < 10:136 return False, "Your question seems too short. Could you provide more details?"137 138 # Check if it's just a greeting139 greetings = ['hello', 'hi', 'hey']140 if question.lower() in greetings:141 return False, "Please ask a specific question about your data."142 143 # Check for question words or data keywords144 data_indicators = ['show', 'get', 'find', 'list', 'count', 'how many', 'what', 145 'which', 'total', 'sum', 'average', 'display']146 147 has_indicator = any(ind in question.lower() for ind in data_indicators)148 149 if not has_indicator:150 return False, "Please rephrase your question to be more specific. For example: 'Show me total sales by region' or 'How many active customers do we have?'"151 152 return True, ""153 154 def preprocess_question(self, question: str) -> str:155 """156 Clean and normalize user question.157 158 Args:159 question: Raw user question160 161 Returns:162 Preprocessed question163 """164 # Strip whitespace165 question = question.strip()166 167 # Expand abbreviations168 abbreviations = {169 r'\bqty\b': 'quantity',170 r'\bamt\b': 'amount',171 r'\btot\b': 'total',172 r'\bavg\b': 'average',173 r'\bpct\b': 'percent',174 }175 176 for abbrev, full in abbreviations.items():177 question = re.sub(abbrev, full, question, flags=re.IGNORECASE)178 179 # Resolve relative time references180 question = self._resolve_time_references(question)181 182 return question183 184 def _resolve_time_references(self, question: str) -> str:185 """186 Resolve relative time references to specific dates.187 188 Args:189 question: Question with potential time references190 191 Returns:192 Question with resolved dates193 """194 now = datetime.now()195 196 # Define time mappings197 replacements = {198 'today': now.strftime('%Y-%m-%d'),199 'yesterday': (now - timedelta(days=1)).strftime('%Y-%m-%d'),200 'last week': f"between '{(now - timedelta(days=7)).strftime('%Y-%m-%d')}' and '{now.strftime('%Y-%m-%d')}'",201 'last month': f"in {(now - timedelta(days=30)).strftime('%Y-%m')}",202 'this month': f"in {now.strftime('%Y-%m')}",203 'this year': f"in {now.strftime('%Y')}",204 }205 206 question_lower = question.lower()207 208 for ref, replacement in replacements.items():209 if ref in question_lower:210 # Replace while preserving original case if possible211 pattern = re.compile(re.escape(ref), re.IGNORECASE)212 question = pattern.sub(replacement, question)213 214 return question215 216 217# Global gatekeeper agent instance218gatekeeper = GatekeeperAgent()219 