ai-data-chatbot / utils /llm_handler.py
Claude
Switch LLM backend from Anthropic to Gemini 2.0 Flash (free tier)
4f62127 unverified
Raw
History Blame Contribute Delete
5.09 kB
"""Gemini API integration for natural language β†’ data analysis code generation."""
import os
import json
import re
import google.generativeai as genai
SYSTEM_PROMPT = """You are an expert data analyst AI assistant. Users ask questions about their dataset and you generate executable Python code to answer those questions with pandas and plotly visualizations.
CRITICAL: Respond with ONLY a valid JSON object β€” no markdown fences, no text before or after the JSON.
Response format:
{
"explanation": "Clear, concise explanation of what you analyzed and the key findings",
"pandas_code": "Python code using pandas. Variable 'df' is the input DataFrame. You MUST create 'result_df' as the output DataFrame.",
"viz_code": "Python plotly code that uses 'result_df' to create a variable called 'fig'. Use null if no chart is needed.",
"viz_type": "bar | line | pie | scatter | heatmap | table_only | null"
}
PANDAS CODE RULES:
- Input: df (pandas DataFrame with the user's data)
- Output: result_df (a pandas DataFrame β€” always required)
- Available imports: pandas as pd, numpy as np
- Use actual column names from the schema provided
- Handle missing values with .fillna() or .dropna() where needed
- For year-over-year comparisons: filter df['year'] == X
- For rankings: use .nlargest() or .sort_values().head()
- Keep result_df focused and readable (avoid too many columns)
VIZ CODE RULES:
- DO NOT include any import statements β€” px, go, pd, np are already available
- Input: result_df (the DataFrame produced by pandas_code)
- Output: fig (a plotly Figure)
- Always set a descriptive title
- Use color parameters for grouped charts
- For grouped bar charts comparing years: use px.bar with barmode='group'
- For time series: use px.line
- Make axis labels human-readable
- Set appropriate figure height (600-700px)
CHART SELECTION GUIDE:
- Comparing values across categories β†’ bar chart
- Year-over-year comparison β†’ grouped bar chart
- Trends over time β†’ line chart
- Parts of a whole β†’ pie chart
- Correlation between two numeric values β†’ scatter chart
- Ranking / top N β†’ horizontal bar chart
EXAMPLE for "top 5 highest paid employees in 2023":
pandas_code: "result_df = df[df['year'] == 2023].nlargest(5, 'salary')[['full_name', 'job_title', 'department', 'salary']].reset_index(drop=True)"
viz_code: "fig = px.bar(result_df, x='full_name', y='salary', color='department', title='Top 5 Highest Paid Employees (2023)', labels={'salary': 'Annual Salary ($)', 'full_name': 'Employee'}, height=500); fig.update_layout(xaxis_tickangle=-30)"
"""
class LLMHandler:
def __init__(self):
self.chat = None
self.conversation_history = []
def _get_model(self):
api_key = os.getenv("GEMINI_API_KEY", "")
if not api_key:
raise ValueError("GEMINI_API_KEY is not set.")
genai.configure(api_key=api_key)
return genai.GenerativeModel("gemini-2.0-flash", system_instruction=SYSTEM_PROMPT)
def reset_conversation(self):
self.conversation_history = []
self.chat = None
def generate_analysis(self, question: str, schema_info: str, sample_rows: str) -> dict:
user_message = (
"Dataset Schema:\n" + schema_info +
"\n\nSample data (first 5 rows):\n" + sample_rows +
"\n\nUser question: " + question +
"\n\nGenerate the pandas + plotly code to answer this question. Remember: respond with ONLY the JSON object."
)
if self.chat is None:
model = self._get_model()
self.chat = model.start_chat(history=[])
response = self.chat.send_message(user_message)
raw_text = response.text
return self._parse_response(raw_text)
def _parse_response(self, raw_text: str) -> dict:
text = raw_text.strip()
text = re.sub(r"^```(?:json)?\s*", "", text)
text = re.sub(r"\s*```$", "", text)
text = text.strip()
try:
result = json.loads(text)
except json.JSONDecodeError:
match = re.search(r"\{[\s\S]*\}", text)
if match:
try:
result = json.loads(match.group())
except json.JSONDecodeError:
result = {
"explanation": "Could not parse LLM response. Please try rephrasing your question.",
"pandas_code": "result_df = df.head(10)",
"viz_code": None,
"viz_type": "table_only",
}
else:
result = {
"explanation": raw_text,
"pandas_code": "result_df = df.head(10)",
"viz_code": None,
"viz_type": "table_only",
}
return {
"explanation": result.get("explanation", ""),
"pandas_code": result.get("pandas_code", "result_df = df.head(10)"),
"viz_code": result.get("viz_code"),
"viz_type": result.get("viz_type", "table_only"),
}