Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
gatekeeper.py219 linesDownload Raw Back to agents
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