from __future__ import annotations import sys from pathlib import Path import gradio as gr import matplotlib.pyplot as plt import pandas as pd PROJECT_ROOT = Path(__file__).resolve().parents[1] if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) from services.service import BIService service = BIService() def ensure_list(value): if isinstance(value, list): return [str(v) for v in value if str(v).strip()] if isinstance(value, str) and value.strip(): return [value.strip()] return [] def build_schema_info_markdown() -> str: try: schema_catalog = service.clickhouse.get_schema_catalog() except Exception as exc: return f"## Connected Schema\nUnable to load schema info: `{exc}`" grouped: dict[str, list[str]] = {} for column in schema_catalog: table_name = f"{column.database}.{column.table}" grouped.setdefault(table_name, []).append(f"`{column.name}`") preferred_tables = [ "amazon.amazon_reviews", "amazon.product_dim", "amazon.daily_review_metrics", "amazon.product_dim_stage", ] ordered_tables = [table for table in preferred_tables if table in grouped] ordered_tables.extend(sorted(table for table in grouped if table not in preferred_tables)) lines = ["## Connected Schema", "**Database:** `amazon`", ""] for table_name in ordered_tables: lines.append(f"**{table_name}**") lines.append(", ".join(grouped[table_name])) lines.append("") return "\n".join(lines) def validate_visualization(dataframe: pd.DataFrame, viz: dict) -> dict: if dataframe.empty: return {} columns = list(dataframe.columns) chart_type = viz.get("chart_type") x_axis = viz.get("x_axis") y_axis = viz.get("y_axis") if x_axis not in columns: x_axis = columns[0] if y_axis not in columns: y_axis = columns[1] if len(columns) > 1 else None if chart_type not in ["line", "bar", "scatter", "table"]: if x_axis and any(token in x_axis.lower() for token in ["date", "time", "day", "month", "year"]): chart_type = "line" elif x_axis and dataframe[x_axis].dtype == "object": chart_type = "bar" else: chart_type = "scatter" return { "chart_type": chart_type, "x_axis": x_axis, "y_axis": y_axis, "orientation": viz.get("orientation", "vertical"), "top_n": viz.get("top_n", 12), "rotate_x_labels": viz.get("rotate_x_labels", False), "truncate_labels": viz.get("truncate_labels", 28), "title": viz.get("title", "Chart"), } def _truncate_label(value: object, limit: int) -> str: text = str(value) if len(text) <= limit: return text return text[: max(0, limit - 3)] + "..." def render_chart_gradio(dataframe: pd.DataFrame, viz: dict): if dataframe.empty: return None viz = validate_visualization(dataframe, viz) chart_type = viz.get("chart_type") x_axis = viz.get("x_axis") y_axis = viz.get("y_axis") if chart_type == "table" or not x_axis or not y_axis: return None working = dataframe.copy() top_n = viz.get("top_n", 12) truncate_labels = viz.get("truncate_labels", 28) orientation = viz.get("orientation", "vertical") if chart_type == "bar" and y_axis in working.columns and len(working) > top_n: working = working.sort_values(by=y_axis, ascending=False).head(top_n) if working[x_axis].dtype == "object": working[x_axis] = working[x_axis].map(lambda value: _truncate_label(value, truncate_labels)) label_lengths = working[x_axis].astype(str).map(len) if x_axis in working.columns else pd.Series(dtype=int) auto_horizontal = chart_type == "bar" and (len(working) > 10 or (not label_lengths.empty and label_lengths.max() > 18)) horizontal = orientation == "horizontal" or auto_horizontal fig_width = 12 if not horizontal else 14 fig_height = 5 if len(working) <= 12 else 7 fig, ax = plt.subplots(figsize=(fig_width, fig_height)) if chart_type == "line": ax.plot(working[x_axis], working[y_axis], marker="o") elif chart_type == "bar": if horizontal: ax.barh(working[x_axis], working[y_axis], color="#4C78A8") else: ax.bar(working[x_axis], working[y_axis], color="#4C78A8") elif chart_type == "scatter": ax.scatter(working[x_axis], working[y_axis], alpha=0.8, color="#4C78A8") ax.set_title(viz.get("title") or f"{chart_type.title()} Chart") if horizontal and chart_type == "bar": ax.set_ylabel(x_axis) ax.set_xlabel(y_axis) else: ax.set_xlabel(x_axis) ax.set_ylabel(y_axis) if viz.get("rotate_x_labels") and chart_type != "scatter" and not horizontal: plt.setp(ax.get_xticklabels(), rotation=35, ha="right") ax.grid(axis="y", linestyle="--", alpha=0.25) fig.tight_layout() return fig def run_query(question: str): if not question.strip(): return "Please enter a question.", pd.DataFrame(), "", "", None, "", "" try: result = service.ask(question.strip()) sql_payload = result.get("sql", {}) analysis = result.get("analysis", {}) visualization = result.get("visualization", {}) reflection = result.get("reflection", {}) result_payload = sql_payload.get("result", {}) dataframe = pd.DataFrame(result_payload.get("rows", [])) metadata_only = bool(sql_payload.get("metadata_only")) summary = "" if metadata_only else analysis.get("summary", "") sql_query = sql_payload.get("sql", "") insights = "" if metadata_only else "\n".join(ensure_list(analysis.get("insights", []))) followups = "" if metadata_only else "\n".join(ensure_list(analysis.get("follow_ups", []))) risks = "\n".join(ensure_list(reflection.get("risks", []))) figure = None if metadata_only else render_chart_gradio(dataframe, visualization) return summary, dataframe, sql_query, insights, figure, followups, risks except Exception as exc: return f"Error: {exc}", pd.DataFrame(), "", "", None, "", "" with gr.Blocks(title="BI Agent") as demo: gr.Markdown("# Business Intelligence Agent") gr.Markdown("Ask database questions, get SQL, results, checks, and charts when the data shape supports them.") runtime = service.runtime_status() gr.Markdown( f""" **Runtime Status** - Groq: `{runtime.get('has_api_key', False)}` - ClickHouse: `{runtime.get('has_database', False)}` - Model: `{runtime.get('groq_model', 'unknown')}` """ ) with gr.Row(): with gr.Column(scale=4): question_input = gr.Textbox( label="Ask your question", placeholder="Which popular but poorly rated products should we investigate?", lines=3, ) run_btn = gr.Button("Run Analysis") summary_output = gr.Textbox(label="Summary") table_output = gr.Dataframe(label="Result Table") sql_output = gr.Code(label="Generated SQL", language="sql") insights_output = gr.Textbox(label="Insights") visualization_output = gr.Plot(label="Visualization") followups_output = gr.Textbox(label="Follow-ups") risks_output = gr.Textbox(label="Risks") with gr.Column(scale=2, min_width=320): schema_info = gr.Markdown(build_schema_info_markdown()) run_btn.click( run_query, inputs=[question_input], outputs=[ summary_output, table_output, sql_output, insights_output, visualization_output, followups_output, risks_output, ], ) if __name__ == "__main__": demo.launch(share=True)