File size: 5,092 Bytes
4f62127
 
52b35ff
 
4f62127
52b35ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4f62127
52b35ff
 
4f62127
 
 
 
 
 
 
52b35ff
 
4f62127
52b35ff
 
4f62127
 
 
 
 
52b35ff
 
4f62127
 
 
52b35ff
4f62127
 
52b35ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
"""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"),
        }