diff --git a/backend/apps/chatbot/agent/__init__.py b/backend/apps/chatbot/agent/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/backend/apps/chatbot/agent/prompts.py b/backend/apps/chatbot/agent/prompts.py new file mode 100644 index 00000000..172a4213 --- /dev/null +++ b/backend/apps/chatbot/agent/prompts.py @@ -0,0 +1,101 @@ +# -*- coding: utf-8 -*- +SQL_AGENT_SYSTEM_PROMPT = """# Persona: Base dos Dados Research Assistant +You are a specialized AI research assistant, an expert in the Base dos Dados (BD) platform and the landscape of Brazilian public data. Your mission is to be a knowledgeable, systematic, and persistent research partner. You don't just execute tools; you guide users through the complexities of Brazil's data ecosystem, explaining the context behind the data and teaching best practices along the way. + +--- + +# Core Directives (Mandatory) + +**1. Embody the Persona**: Act as Base dos Dados Research Asisstant, the expert guide. Be proactive, educational, and systematic in every interaction. +**2. Adhere to the Workflow**: Your reasoning process **MUST** follow this strict cycle for every user request. Do not deviate. + - **Thought**: First, state your goal and reasoning. Formulate a hypothesis, select the appropriate tool, and explain your chosen parameters. This inner monologue is crucial for transparency. + - **Action**: Execute a single tool call based on your thought process (`search_datasets`, `get_dataset_details`, `decode_table_values`, `inspect_column_values`, or `execute_bigquery_sql`). +**3. Search Protocol is Mandatory**: **Always** begin your investigation with the single-keyword search strategy outlined below. This is the most critical rule for success. +**4. Explain Everything**: Never just show data. Summarize your findings, explain the data's source and context, highlight key insights, and suggest logical next steps. + +--- + +# Knowledge Base: Brazilian Data Essentials + +### Key Data Sources +- **Instituto Brasileiro de Geografia e Estatística (IBGE)**: Census, demographics, economic surveys (`censo`, `pnad`, `pof`). +- **Instituto Nacional de Estudos e Pesquisas Educacionais Anísio Teixeira (INEP)**: Education data (`ideb`, `censo escolar`, `enem`). +- **Ministério da Saúde (MS)**: Health data (`pns`, `sinasc`, `sinan`, `sim`). +- **Ministério da Economia (ME)**: Employment & economic data (`rais`, `caged`). +- **Tribunal Superior Eleitoral (TSE)**: Electoral data (`eleicoes`, `filiados`). +- **Banco Central do Brasil (BCB)**: Financial data (`taxa selic`, `cambio`, `ipca`). + +### Common Data Patterns +- **Geographic**: `sigla_uf` (state), `id_municipio` (municipality), `regiao`. +- **Temporal**: `ano` (year), `mes` (month), `data` (date), `semestre`, `trimestre`. +- **Identifiers**: `id_*`, `codigo_*`, `sigla_*`. +- **Coded Values**: Many columns use codes for efficiency (e.g., `id_municipio`). **Always** prioritize `decode_table_values` to understand them. Use `inspect_column_values` as a fallback for exploration. + +--- + +# Search Protocol + +You **MUST** follow this tiered search funnel. Do not skip steps. Justify your keyword choices in your **Thought** process. + +### Tier 1: High-Confidence Single Keywords (Always Try First) +*Start every search with a single, high-probability keyword, tried in this specific order.* +1. **Dataset Name**: If the user's query mentions a known dataset name (`censo`, `rais`, `enem`, `sinasc`), use it directly. +2. **Organization Acronym**: If a government organization is relevant (`ibge`, `inep`, `ms`, `tse`, `bcb`), use its acronym. +3. **Core Theme (Portuguese)**: Use a broad, common theme in Portuguese (`educacao`, `saude`, `economia`, `emprego`, `eleicoes`). + + +**User:** Como foi o desempenho em matemática dos alunos no brasil nos últimos anos? +**Thought:** The user is asking about student performance. The organization `inep` might be a good data source. I will start by searching with the keyword "inep". If that fails, I will try searching for the theme "educacao". +**Action:** `search_datasets(inep)` + + +### Tier 2: Alternative Single Keywords (If Tier 1 Fails) +*If and only if Tier 1 yields no relevant results, document the failure and proceed to these options.* +- **Synonyms**: Try a Portuguese synonym (`ensino` for `educacao`, `trabalho` for `emprego`). +- **Broader Concepts**: Use a more general term (`social`, `demografia`, `infraestrutura`). +- **English Equivalents**: As a last resort for single keywords, try English (`health`, `education`). + +### Tier 3: Multi-Keyword Search (Last Resort) +*Only use 2-3 keywords if all single-keyword searches have failed. This is an exception, not the rule.* +- **Theme + Agency**: `saude ms`, `educacao inep` +- **Dataset + Geography**: `censo municipio`, `rais estado` + +--- + +# BigQuery SQL Protocol + +- **Reference Full IDs**: Always use the full table ID: `project.dataset.table`. +- **Select Specific Columns**: Never use `SELECT *`. Explicitly list the columns you need. +- **Limit for Exploration**: When first inspecting a table, **always** use a `LIMIT` clause. You can query without a `LIMIT` clause later. +- **Filter Early and Often**: Use `WHERE` clauses on partitioned or clustered columns (usually `ano`) to drastically reduce query cost. +- **Default to Most Recent Data**: If the user does not specify a time range, your default behavior **MUST** be to query for the most recent data. Find the latest year or date in the relevant column (e.g., `ano`) and use it to filter the query. You **MUST** state that you queried the most recent data available. +- **Order for Insights**: Use `ORDER BY` to present data logically. +- **No DDL/DML:** NEVER run DDL/DML commands (`CREATE`, `ALTER`, `DROP`, `INSERT`, `UPDATE`, `DELETE`) + +--- + +# User Communication Protocol + +### Response Structure +1. **Summary of Findings**: Start with a clear, concise summary of the answer. +2. **Context**: Explain what the data represents. Mention the source organization (e.g., "Data from IBGE's 2010 Census..."), the time period, and the geographic level. +3. **Data/Results**: Present the data clearly. Use Markdown tables for structured results, bullet points for lists, etc. Display null/empty values as "N/A" for clarity. +4. **Key Insights**: Highlight 1-3 important points or patterns from the results. +5. **Suggested Next Steps**: Propose a relevant follow-up question, a related dataset to explore, or a way to refine the current analysis. + +### Handling Failures +- **Search Fails**: Explain your keyword strategy, state why it failed (e.g., "The search for 'cnes' returned no datasets"), and describe your next attempt based on the Search Protocol. +- **Query Errors**: Analyze the BigQuery error message. Suggest a specific fix (e.g., "The query is too large. I will add a `WHERE` clause to filter by year to reduce the data processed."). +- **Empty Results**: Hypothesize why the result is empty. Check your filters, the data's time range, or if you are filtering on a coded value incorrectly. Suggest a modified query. + +--- + +# Final Reminder + +**Before executing any action, ensure you are compliant:** +- **NEVER** use `SELECT *`. +- **NEVER** query a table without a `LIMIT` clause during initial exploration. +- **NEVER** run Data Definition/Manipulation Language (`CREATE`, `ALTER`, `DROP`, `INSERT`, `UPDATE`, `DELETE`). Your access is strictly read-only. +- **NEVER** start with a multi-keyword search. The Search Protocol is mandatory. +- **NEVER** present raw data without a summary and context first. +- **NEVER** give up after one failed attempt. Show persistence and a systematic problem-solving approach.""" # noqa: E501 diff --git a/backend/apps/chatbot/agent/tools.py b/backend/apps/chatbot/agent/tools.py new file mode 100644 index 00000000..1d421895 --- /dev/null +++ b/backend/apps/chatbot/agent/tools.py @@ -0,0 +1,566 @@ +# -*- coding: utf-8 -*- +import json +import os +from functools import cache +from typing import Any, Literal, Optional, Self + +import httpx +from google.api_core import exceptions as google_api_exceptions +from google.cloud import bigquery as bq +from langchain_core.tools import BaseTool, tool +from pydantic import BaseModel, model_validator + +# 1GB limit for dictionary queries +LIMIT_DICTIONARY_QUERY = 1e9 + +# 5GB limit for inspection queries +LIMIT_INSPECTION_QUERY = 5e9 + +# 20GB limit for other queries +LIMIT_BIGQUERY_QUERY = 20e9 + +SEARCH_URL = "https://backend.basedosdados.org/search/" +GRAPHQL_URL = "https://backend.basedosdados.org/graphql" + +DATASET_DETAILS_QUERY = """ +query GetDatasetOverview($id: ID!) { + allDataset(id: $id, first: 1) { + edges { + node { + id + name + slug + description + organizations { + edges { + node { + name + slug + } + } + } + themes { + edges { + node { + name + } + } + } + tags { + edges { + node { + name + } + } + } + tables { + edges { + node { + id + name + slug + description + cloudTables { + edges { + node { + gcpProjectId + gcpDatasetId + gcpTableId + } + } + } + columns { + edges { + node { + id + name + description + bigqueryType { + name + } + } + } + } + } + } + } + } + } + } +} +""" + + +class Column(BaseModel): + name: str + type: str + description: Optional[str] + + +class Table(BaseModel): + id: str + gcp_id: Optional[str] + name: str + slug: Optional[str] + description: Optional[str] + columns: list[Column] + + +class DatasetOverview(BaseModel): + id: str + name: str + slug: Optional[str] + description: Optional[str] + tags: list[str] + themes: list[str] + organizations: list[str] + + +class Dataset(DatasetOverview): + tables: list[Table] + + +class ToolOutput(BaseModel): + status: Literal["success", "error"] + results: Optional[Any] = None + error_details: Optional[dict[str, Any]] = None + + @model_validator(mode="after") + def check_passwords_match(self) -> Self: + if (self.results is None) ^ (self.error_details is None): + return self + raise ValueError("Only one of 'results' or 'error_details' should be set") + + +@cache +def get_bigquery_client(): + return bq.Client(project=os.environ["QUERY_PROJECT_ID"]) + + +@tool +def search_datasets(query: str) -> str: + """Search for datasets in Base dos Dados using keywords. + + CRITICAL: Use individual KEYWORDS only, not full sentences. The search engine uses Elasticsearch. + + Args: + query (str): 2-3 keywords maximum. Use Portuguese terms, organization acronyms, or dataset acronyms. + Good Examples: "censo", "educacao", "ibge", "inep", "rais", "saude" + Avoid: "Brazilian population data by municipality" + + Returns: + str: JSON array of datasets. If empty/irrelevant results, try different keywords. + + Strategy: Start with broad terms like "censo", "ibge", "inep", "rais", then get specific if needed. + Next step: Use `get_dataset_details()` with returned dataset IDs. + """ # noqa: E501 + try: + with httpx.Client() as client: + response = client.get( + url=SEARCH_URL, + params={"q": query, "page_size": 10}, + timeout=httpx.Timeout(5.0, read=60.0), + ) + response.raise_for_status() + data: dict = response.json() + + datasets = data.get("results", []) + + overviews = [] + + for dataset in datasets: + dataset_overview = DatasetOverview( + id=dataset["id"], + name=dataset["name"], + slug=dataset.get("slug"), + description=dataset.get("description"), + tags=[tag["name"] for tag in dataset.get("tags", [])], + themes=[theme["name"] for theme in dataset.get("themes", [])], + organizations=[org["name"] for org in dataset.get("organizations", [])], + ) + overviews.append(dataset_overview.model_dump()) + + tool_output = ToolOutput(status="success", results=overviews).model_dump(exclude_none=True) + except Exception as e: + tool_output = ToolOutput( + status="error", error_details={"message": f"Error searching datasets:\n{e}"} + ).model_dump(exclude_none=True) + + return json.dumps(tool_output, ensure_ascii=False, indent=2) + + +@tool +def get_dataset_details(dataset_id: str) -> str: + """Get comprehensive details about a specific dataset including all tables and columns. + + Use AFTER `search_datasets()` to understand data structure before writing queries. + + Args: + dataset_id (str): Dataset ID obtained from `search_datasets()`. + This is typically a UUID-like string, not the human-readable name. + + Returns: + str: JSON object with complete dataset information, including: + - Basic metadata (name, description, tags, themes, organizations) + - tables: Array of all tables in the dataset with: + - gcp_id: Full BigQuery table reference (`project.dataset.table`) + - columns: All column names, types, and descriptions + - table descriptions explaining what each table contains + + Next step: Use `execute_bigquery_sql()` to execute queries. + """ # noqa: E501 + try: + with httpx.Client() as client: + response = client.post( + url=GRAPHQL_URL, + json={ + "query": DATASET_DETAILS_QUERY, + "variables": {"id": dataset_id}, + }, + timeout=httpx.Timeout(5.0, read=60.0), + ) + response.raise_for_status() + data: dict[str, dict[str, dict]] = response.json() + + dataset_edges = data.get("data", {}).get("allDataset", {}).get("edges", []) + + if not dataset_edges: + return f"Dataset {dataset_id} not found" + + dataset = dataset_edges[0]["node"] + + dataset_id = dataset["id"] + dataset_name = dataset["name"] + dataset_slug = dataset.get("slug") + dataset_description = dataset.get("description") + + # Tags + dataset_tags = [] + for edge in dataset.get("tags", {}).get("edges", []): + if tag := edge.get("node", {}).get("name"): + dataset_tags.append(tag) + + # Themes + dataset_themes = [] + for edge in dataset.get("themes", {}).get("edges", []): + if theme := edge.get("node", {}).get("name"): + dataset_themes.append(theme) + + # Organizations + dataset_organizations = [] + for edge in dataset.get("organizations", {}).get("edges", []): + if org := edge.get("node", {}).get("name"): + dataset_organizations.append(org) + + # Tables + dataset_tables = [] + for edge in dataset.get("tables", {}).get("edges", []): + table = edge["node"] + + table_id = table["id"] + table_name = table["name"] + table_slug = table.get("slug") + table_description = table.get("description") + + cloud_table_edges = table["cloudTables"]["edges"] + if cloud_table_edges: + cloud_table = cloud_table_edges[0]["node"] + gcp_project_id = cloud_table["gcpProjectId"] + gcp_dataset_id = cloud_table["gcpDatasetId"] + gcp_table_id = cloud_table["gcpTableId"] + table_gcp_id = f"{gcp_project_id}.{gcp_dataset_id}.{gcp_table_id}" + else: + table_gcp_id = None + + table_columns = [] + for edge in table["columns"]["edges"]: + column = edge["node"] + table_columns.append( + Column( + name=column["name"], + type=column["bigqueryType"]["name"], + description=column.get("description"), + ) + ) + + dataset_tables.append( + Table( + id=table_id, + gcp_id=table_gcp_id, + name=table_name, + slug=table_slug, + description=table_description, + columns=table_columns, + ) + ) + + dataset = Dataset( + id=dataset_id, + name=dataset_name, + slug=dataset_slug, + description=dataset_description, + tags=dataset_tags, + themes=dataset_themes, + organizations=dataset_organizations, + tables=dataset_tables, + ).model_dump() + + tool_output = ToolOutput(status="success", results=dataset).model_dump(exclude_none=True) + except Exception as e: + tool_output = ToolOutput( + status="error", error_details={"message": f"Error fetching dataset details:\n{e}"} + ).model_dump(exclude_none=True) + + return json.dumps(tool_output, ensure_ascii=False, indent=2) + + +@tool +def execute_bigquery_sql(sql_query: str) -> str: + """Execute a SQL query against BigQuery tables from the Base dos Dados database. + + Use AFTER identifying the right datasets and understanding tables structure. + It includes a 20GB processing limit for safety. + + Args: + sql_query (str): Standard GoogleSQL query. Must reference + tables using their full `gcp_id` from `get_dataset_details()`. + + Best practices: + - Use fully qualified names: `project.dataset.table` + - Select only needed columns, avoid `SELECT *` + - Add `LIMIT` for exploration + - Filter early with `WHERE` clauses + - Order by relevant columns + - Never use DDL/DML commands + - Use appropriate data types in comparisons + + Returns: + str: Query results as JSON array. Empty results return "[]". + """ # noqa: E501 + forbidden_commands = [ + "CREATE", + "ALTER", + "DROP", + "TRUNCATE", + "INSERT", + "UPDATE", + "DELETE", + "GRANT", + "REVOKE", + ] + + for command in forbidden_commands: + if command in sql_query.upper(): + return ToolOutput( + status="error", + error_details={ + "message": ( + f"Query aborted: Command {command} is forbidden. ", + "Your access is strictly read-only.", + ) + }, + ).model_dump_json(indent=2, exclude_none=True) + + client = get_bigquery_client() + + try: + job_config = bq.QueryJobConfig(dry_run=True, use_query_cache=False) + query_job = client.query(sql_query, job_config=job_config) + + limit_bytes = LIMIT_BIGQUERY_QUERY + total_bytes = query_job.total_bytes_processed + + if total_bytes and total_bytes > limit_bytes: + return ToolOutput( + status="error", + error_details={ + "type": "QueryTooLarge", + "limit_bytes": limit_bytes, + "total_processed_bytes": total_bytes, + "message": ( + "Query aborted: Data processed exceeds the per-query limit. " + "Consider optimizing by adding filters, selecting fewer columns, " + "or using a LIMIT clause before retrying." + ), + }, + ).model_dump_json(indent=2, exclude_none=True) + + rows = client.query(sql_query).result() + + results = [dict(row) for row in rows] + + tool_output = ToolOutput(status="success", results=results).model_dump(exclude_none=True) + except Exception as e: + tool_output = ToolOutput( + status="error", error_details={"message": f"SQL query execution failed:\n{e}"} + ).model_dump(exclude_none=True) + + return json.dumps(tool_output, ensure_ascii=False, default=str) + + +@tool +def decode_table_values(table_gcp_id: str, column_name: Optional[str] = None) -> str: + """Decode coded values from a table. + + Use when column values appear to be codes (e.g., 1,2,3 or A,B,C). + Many datasets use codes for storage efficiency. This tool provides + the authoritative meanings of these codes. + + Args: + table_gcp_id (str): Full BigQuery table reference + column_name (Optional[str], optional): Column with coded values. If `None`, + all columns will be used. Defaults to `None`. + + Returns: + str: JSON array with chave (code) and valor (meaning) mappings. + """ # noqa: E501 + client = get_bigquery_client() + + try: + project_name, dataset_name, table_name = table_gcp_id.split(".") + except ValueError: + return ToolOutput( + status="error", + error_details={ + "message": ( + f"{table_gcp_id} is not a valid table reference. " + "Please, provide a valid table reference in the format `project.dataset.table`" + ) + }, + ).model_dump_json(indent=2, exclude_none=True) + + dataset_id = f"{project_name}.{dataset_name}" + dict_table_id = f"{dataset_id}.dicionario" + + search_query = f""" + SELECT nome_coluna, chave, valor + FROM {dict_table_id} + WHERE id_tabela = '{table_name}' + """ + + if column_name is not None: + search_query += f"AND nome_coluna = '{column_name}'" + + search_query += "ORDER BY nome_coluna, chave" + + try: + job_config = bq.QueryJobConfig(dry_run=True, use_query_cache=False) + query_job = client.query(search_query, job_config=job_config) + + limit_bytes = LIMIT_DICTIONARY_QUERY + total_bytes = query_job.total_bytes_processed + + if total_bytes and total_bytes > limit_bytes: + return ToolOutput( + status="error", + error_details={ + "type": "QueryTooLarge", + "limit_bytes": limit_bytes, + "total_processed_bytes": total_bytes, + "message": ( + "Dictionary table is unexpectedly large. " + "This might not be the right approach. " + "Check if this dataset actually uses " + "encoded values or try a different table." + ), + }, + ).model_dump_json(indent=2, exclude_none=True) + + rows = client.query(search_query).result() + results = [dict(row) for row in rows] + tool_output = ToolOutput(status="success", results=results).model_dump(exclude_none=True) + except google_api_exceptions.NotFound: + return ToolOutput( + status="error", + error_details={ + "type": "TableNotFound", + "message": ( + f"Dictionary table not found for dataset {dataset_id}. " + "This indicates this dataset does not contain a dicionary table. " + "Consider using the `inspect_column_values` tool to inspect column values." + ), + }, + ).model_dump_json(indent=2, exclude_none=True) + except Exception as e: + tool_output = ToolOutput( + status="error", error_details={"message": f"Failed to decode table values:\n{e}"} + ).model_dump(exclude_none=True) + + return json.dumps(tool_output, ensure_ascii=False, default=str) + + +@tool +def inspect_column_values(table_gcp_id: str, column_name: str, limit: int = 100) -> str: + """Show actual distinct values in a column. + + FALLBACK tool to use when `search_dictionary_table()` fails. + Useful for understanding data patterns and planning `WHERE` clauses. + + Args: + table_gcp_id (str): Full BigQuery table reference + column_name (str): Column name from `get_dataset_details()` + limit (int, optional): Max distinct values to return (default: 100) + + Returns: + str: JSON array of distinct values. + """ # noqa: E501 + client = get_bigquery_client() + + sql_query = f"SELECT DISTINCT {column_name} FROM {table_gcp_id} LIMIT {limit}" + + try: + job_config = bq.QueryJobConfig(dry_run=True, use_query_cache=False) + query_job = client.query(sql_query, job_config=job_config) + + limit_bytes = LIMIT_INSPECTION_QUERY + total_bytes = query_job.total_bytes_processed + + if total_bytes and total_bytes > limit_bytes: + return ToolOutput( + status="error", + error_details={ + "status": "error", + "type": "QueryTooLarge", + "limit_bytes": limit_bytes, + "total_processed_bytes": total_bytes, + "message": ( + "Column inspection exceeds the per-query limit for inspection. " + "Try a smaller limit or add WHERE filters to reduce data size." + ), + }, + ).model_dump_json(indent=2, exclude_none=True) + + rows = client.query(sql_query).result() + + results = [row.get(column_name) for row in rows] + + tool_output = ToolOutput(status="success", results=results).model_dump(exclude_none=True) + except Exception as e: + tool_output = ToolOutput( + status="error", error_details={"message": f"Failed to inspect colum values:\n{e}"} + ).model_dump(exclude_none=True) + + return json.dumps(tool_output, ensure_ascii=False, default=str) + + +def get_tools() -> list[BaseTool]: + """Return all available tools for Base dos Dados database interaction. + + This function provides a complete set of tools for discovering, exploring, + and querying Brazilian public datasets through the Base dos Dados platform. + + Returns: + list[BaseTool]: A list of LangChain tool functions in suggested usage order: + - search_datasets: Find datasets using keywords + - get_dataset_details: Get comprehensive dataset information + - execute_bigquery_sql: Execute SQL queries against BigQuery tables + - decode_table_values: Decode coded values using dictionary tables + - inspect_column_values: Inspect actual column values as fallback + """ + return [ + search_datasets, + get_dataset_details, + execute_bigquery_sql, + decode_table_values, + inspect_column_values, + ] diff --git a/backend/apps/chatbot/migrations/0007_rename_steps_messagepair_events.py b/backend/apps/chatbot/migrations/0007_rename_steps_messagepair_events.py new file mode 100644 index 00000000..e5b8c698 --- /dev/null +++ b/backend/apps/chatbot/migrations/0007_rename_steps_messagepair_events.py @@ -0,0 +1,18 @@ +# -*- coding: utf-8 -*- +# Generated by Django 4.2.10 on 2025-08-27 20:15 + +from django.db import migrations + + +class Migration(migrations.Migration): + dependencies = [ + ("chatbot", "0006_messagepair_error_message_messagepair_steps_and_more"), + ] + + operations = [ + migrations.RenameField( + model_name="messagepair", + old_name="steps", + new_name="events", + ), + ] diff --git a/backend/apps/chatbot/models.py b/backend/apps/chatbot/models.py index a2898d40..70735bee 100644 --- a/backend/apps/chatbot/models.py +++ b/backend/apps/chatbot/models.py @@ -26,7 +26,7 @@ class MessagePair(models.Model): generated_queries = models.JSONField(null=True, blank=True) generated_chart = models.JSONField(null=True, blank=True) created_at = models.DateTimeField(auto_now_add=True) - steps = models.JSONField(null=True) + events = models.JSONField(null=True) class Meta: constraints = [ diff --git a/backend/apps/chatbot/tests/test_endpoints.py b/backend/apps/chatbot/tests/test_endpoints.py index 62215606..0d399032 100644 --- a/backend/apps/chatbot/tests/test_endpoints.py +++ b/backend/apps/chatbot/tests/test_endpoints.py @@ -4,33 +4,37 @@ import pytest from django.utils.dateparse import parse_datetime +from langchain_core.messages import AIMessage from rest_framework.test import APIClient from backend.apps.account.models import Account from backend.apps.chatbot import views from backend.apps.chatbot.models import Feedback, MessagePair, Thread -from chatbot.assistants import SQLAssistantMessage -class MockSQLAssistant: +class MockLangSmithFeedbackSender: def __init__(self, *args, **kwargs): ... - def invoke(self, *args, **kwargs): - return SQLAssistantMessage(content="mock response") - - def clear_thread(self, *args, **kwargs): + def send_feedback(self, *args, **kwargs): ... -class MockLangSmithFeedbackSender: +class MockReActAgent: def __init__(self, *args, **kwargs): ... - def send_feedback(self, *args, **kwargs): + def stream(self, *args, **kwargs): + yield "updates", {"agent": AIMessage("mock response")} + + def clear_thread(self, *args, **kwargs): ... +def mock_create_react_agent(): + return MockReActAgent() + + @pytest.fixture def mock_email() -> str: return "mockemail@mockdomain.com" @@ -272,7 +276,7 @@ def test_message_list_view_get_order_invalid(auth_client: APIClient, auth_user: @pytest.mark.django_db def test_message_list_view_post(monkeypatch, auth_client: APIClient, auth_user: Account): - monkeypatch.setattr(views, "SQLAssistant", MockSQLAssistant) + monkeypatch.setattr(views, "create_react_agent", mock_create_react_agent) thread = Thread.objects.create(account=auth_user) diff --git a/backend/apps/chatbot/utils/stream.py b/backend/apps/chatbot/utils/stream.py index e768f1a8..3997ae42 100644 --- a/backend/apps/chatbot/utils/stream.py +++ b/backend/apps/chatbot/utils/stream.py @@ -1,231 +1,89 @@ # -*- coding: utf-8 -*- -import re -from collections.abc import Callable -from typing import TypeAlias +from typing import Any, Literal, Optional -import sqlparse from langchain_core.messages import AIMessage, ToolMessage -from pydantic import BaseModel, Field +from pydantic import UUID4, BaseModel -# SQLAgent nodes -LIST_DATASETS = "list_datasets" -CALL_SELECT_DATASETS = "call_select_datasets" -TABLES_INFO = "tables_info" -QUERY_AGENT = "query_agent" -GET_ANSWER = "get_answer" -SQL_TOOLS = "tools" +class ToolCall(BaseModel): + id: str + name: str + args: dict[str, Any] -class StepContent(BaseModel): - title: str | None = Field(default=None, description="A title to be displayed in the UI") - body: str = Field(description="Content to be displayed in the UI") +class ToolOutput(BaseModel): + status: Literal["error", "success"] + tool_call_id: str + tool_name: str + output: str -class Step(BaseModel): - label: str = Field( - description="Short description of the action taken by the agent at this step" - ) - content: list[StepContent] = Field( - description="Detailed outputs generated by the agent for this step" - ) +EventType = Literal[ + "tool_call", + "tool_output", + "final_answer", + "error", + "complete", +] -Handler: TypeAlias = Callable[[dict], Step | None] +class EventData(BaseModel): + run_id: Optional[UUID4] = None + content: Optional[str] = None + tool_calls: Optional[list[ToolCall]] = None + tool_outputs: Optional[list[ToolOutput]] = None + error_details: Optional[dict[str, Any]] = None -# ============================== Helper Functions ============================== -def _format_sql(sql: str, reindent: bool = True, keyword_case: str = "upper") -> str: - return sqlparse.format(sql=sql, reindent=reindent, keyword_case=keyword_case) +class StreamEvent(BaseModel): + type: EventType + data: EventData -def _format_datasets_info(text: str) -> str: - matches = re.findall( - pattern=r"^# (.*?)\s+### Description:\s+(.*?)### Tables:", - string=text, - flags=re.MULTILINE | re.DOTALL, - ) + def to_sse(self) -> str: + return self.model_dump_json() + "\n\n" - parts = [] - for title, description in matches: - parts.append(f"##### {title.strip()}\n\n{description.strip()}") - - return "\n\n".join(parts) - - -def _format_tables_info(text: str) -> str: - matches = re.findall( - pattern=r"^# (.*?)\s+### Description:\s+(.*?)### Schema:", - string=text, - flags=re.MULTILINE | re.DOTALL, - ) - - parts = [] - - for title, description in matches: - parts.append(f"##### {title}\n\n{description}") - - return "\n\n".join(parts) - - -# ========================== SQLAgent Chunk Handlers ========================== -def _handle_list_datasets(chunk: dict) -> Step: - message: ToolMessage = chunk[LIST_DATASETS]["messages"][0] - - label = "Selecionando conjuntos de dados..." - - content = StepContent( - title="Buscando Conjuntos de Dados:", body=_format_datasets_info(message.content) - ) - - return Step(label=label, content=[content]) - - -def _handle_call_select_datasets(chunk: dict) -> Step: - message: AIMessage = chunk[CALL_SELECT_DATASETS]["messages"][-1] - names: list[str] = [] - - for tool_call in message.tool_calls: - dataset_names: str = tool_call["args"].get("dataset_names") - if dataset_names: - names.extend(dataset_names.split(",")) - - bolded_names = [f"**{name.strip()}**" for name in names] - - label = "Selecionando conjuntos de dados..." - - if not bolded_names: - body = "" - - body = "Conjuntos de dados selecionandos:\n- " + "\n- ".join(bolded_names) - - content = StepContent(title="Selecionando Conjuntos de Dados:", body=body) - - return Step(label=label, content=[content]) - - -def _handle_tables_info(chunk: dict) -> Step: - messages: list[ToolMessage] = chunk[TABLES_INFO]["messages"] - - formatted_tables_info = [] - - for message in messages: - if message.status == "error": - formatted_tables_info.append( - ":red[**Erro ao selecionar tabelas. Tentando novamente...**]" - ) - else: - formatted_tables_info.append(_format_tables_info(message.content)) - - label = "Selecionando tabelas..." - - content = StepContent( - title="Selecionando Tabelas:", body="\n\n---\n\n".join(formatted_tables_info) - ) - - return Step(label=label, content=[content]) - - -def _handle_query_agent(chunk: dict) -> Step | None: - message: AIMessage = chunk[QUERY_AGENT]["messages"][0] - - if not message.tool_calls: - return None - - content_parts = [] - - for tool_call in message.tool_calls: - tool_name = tool_call["name"] - - if tool_name == "sql_query_check": - query = tool_call["args"].get("query") - label = "Verificando consulta..." - if query: - body = f"```sql\n{_format_sql(query)}\n```" - else: - body = "red[**Erro na chamada da ferramenta**]" - part = StepContent(title="Verificando Consulta:", body=body) - content_parts.append(part) - - elif tool_name == "sql_query_exec": - query = tool_call["args"].get("query") - label = "Executando consulta..." - if query: - body = f"```sql\n{_format_sql(query)}\n```" - else: - body = "red[**Erro na chamada da ferramenta**]" - part = StepContent(title="Executando Consulta:", body=body) - content_parts.append(part) - - if len(content_parts) > 1: - label = "Consultando banco de dados..." - - return Step(label=label, content=content_parts) - - -def _handle_sql_tools(chunk: dict) -> Step | None: - updates = chunk[SQL_TOOLS] - - # single tool call - if isinstance(updates, dict): - messages: list[ToolMessage] = updates["messages"] - # multiple parallel tool calls - elif isinstance(updates, list): - messages: list[ToolMessage] = [ - update["messages"][0] for update in updates if "messages" in update - ] - - content_parts = [] - - for message in messages: - if message.status == "success": - continue - - if message.name == "sql_query_check": - label = "Verificando consulta..." - part = StepContent(body=":red[**Erro na verificação da consulta**]") - elif message.name == "sql_query_exec": - label = "Executando consulta..." - part = StepContent(body=":red[**Erro na execução da consulta**]") - - content_parts.append(part) - - if content_parts: - if len(content_parts) > 1: - label = "Consultando banco de dados..." - return Step(label=label, content=content_parts) - - return None - - -def _handle_get_answer(chunk: dict) -> Step: - final_answer = chunk[GET_ANSWER]["final_answer"] - label = "Gerando resposta..." - content = StepContent(title="Resposta Final:", body=final_answer) - return Step(label=label, content=[content]) - - -# ========= Mapping of SQLAgent chunks to their corresponding handlers ========= -SQL_AGENT_HANDLERS: dict[str, Handler] = { - LIST_DATASETS: _handle_list_datasets, - CALL_SELECT_DATASETS: _handle_call_select_datasets, - TABLES_INFO: _handle_tables_info, - QUERY_AGENT: _handle_query_agent, - SQL_TOOLS: _handle_sql_tools, - GET_ANSWER: _handle_get_answer, -} - - -def process_chunk(chunk: dict) -> Step | None: - """Process a single chunk from the stream by dispatching it to the correct handler. +def process_chunk(chunk: dict[str, Any]) -> StreamEvent | None: + """Process a streaming chunk from a react agent workflow into a standardized StreamEvent. Args: - chunk (dict): A chunk from the stream. + chunk (dict[str, Any]): Raw chunk from agent workflow. + Only processes "agent" and "tools" nodes. Returns: - Step | None: The processed results if a handler was found. `None` otherwise. + StreamEvent | None: Structured event or None if the chunk is ignored: + - "tool_call" for agent messages with tool calls + - "tool_output" for tool execution results + - "final_answer" for agent messages without tool calls + - None for ignored chunks """ - for key, handler in SQL_AGENT_HANDLERS.items(): - if key in chunk: - return handler(chunk) + if "agent" in chunk: + message: AIMessage = chunk["agent"]["messages"][0] + + if message.tool_calls: + tool_calls = [ + ToolCall(id=tool_call["id"], name=tool_call["name"], args=tool_call["args"]) + for tool_call in message.tool_calls + ] + event_type = "tool_call" + event_data = EventData(content=message.content, tool_calls=tool_calls) + else: + event_type = "final_answer" + event_data = EventData(content=message.content) + + return StreamEvent(type=event_type, data=event_data) + elif "tools" in chunk: + messages: list[ToolMessage] = chunk["tools"]["messages"] + + tool_outputs = [ + ToolOutput( + status=message.status, + tool_call_id=message.tool_call_id, + tool_name=message.name, + output=message.content, + ) + for message in messages + ] + + return StreamEvent(type="tool_output", data=EventData(tool_outputs=tool_outputs)) return None diff --git a/backend/apps/chatbot/views.py b/backend/apps/chatbot/views.py index 2bf3772f..37364201 100644 --- a/backend/apps/chatbot/views.py +++ b/backend/apps/chatbot/views.py @@ -1,17 +1,21 @@ # -*- coding: utf-8 -*- -import json import os import uuid +from collections.abc import Generator from contextlib import contextmanager from functools import cache from typing import Any, Iterator, Type, TypedDict, TypeVar from django.http import StreamingHttpResponse +from google.api_core import exceptions as google_api_exceptions from graphql_jwt.shortcuts import get_user_by_token from langchain.chat_models import init_chat_model -from langchain_google_vertexai import VertexAIEmbeddings -from langchain_postgres import PGVector +from langchain_core.messages import RemoveMessage, ToolMessage +from langchain_core.messages.utils import count_tokens_approximately, trim_messages from langgraph.checkpoint.postgres import PostgresSaver +from langgraph.graph.graph import CompiledGraph +from langgraph.graph.message import REMOVE_ALL_MESSAGES +from langgraph.prebuilt import create_react_agent from loguru import logger from rest_framework import exceptions, status from rest_framework.parsers import JSONParser @@ -22,7 +26,8 @@ from rest_framework.views import APIView from rest_framework_simplejwt.tokens import RefreshToken -from backend.apps.chatbot.context_provider import PostgresContextProvider +from backend.apps.chatbot.agent.prompts import SQL_AGENT_SYSTEM_PROMPT +from backend.apps.chatbot.agent.tools import get_tools from backend.apps.chatbot.feedback_sender import LangSmithFeedbackSender from backend.apps.chatbot.models import Feedback, MessagePair, Thread from backend.apps.chatbot.serializers import ( @@ -33,16 +38,27 @@ ThreadSerializer, UserMessageSerializer, ) -from backend.apps.chatbot.utils.stream import process_chunk -from chatbot.assistants import SQLAssistant, format_sql_agent_response -from chatbot.formatters import SQLPromptFormatter +from backend.apps.chatbot.utils.stream import EventData, StreamEvent, process_chunk +from chatbot.agents.utils import delete_checkpoints ModelSerializer = TypeVar("ModelSerializer", bound=Serializer) # Model name/URI. Refer to the LangChain docs for valid names/URIs # https://python.langchain.com/api_reference/langchain/chat_models/langchain.chat_models.base.init_chat_model.html -MODEL_URI = os.environ["MODEL_URI"] +MODEL_URI = "google_vertexai:gemini-2.5-flash" + +# Gemini models have a ~1 million tokens context window +CONTEXT_WINDOW = 2**20 + +# Maximum number of tokens allowed at the START of a conversation turn +MAX_TOKENS = CONTEXT_WINDOW // 2 + +# Generic error message for unexpected errors when calling the agent +UNEXPECTED_ERROR_MESSAGE = ( + "Ops, algo deu errado! Ocorreu um erro inesperado. Por favor, tente novamente. " + "Se o problema persistir, avise-nos. Obrigado pela paciência!" +) class ConfigDict(TypedDict): @@ -173,10 +189,16 @@ def delete(self, request: Request, thread_id: uuid.UUID) -> Response: try: thread.deleted = True thread.save() - with _get_sql_assistant() as assistant: - assistant.clear_thread(str(thread_id)) + with _get_sql_agent() as agent: + if agent.checkpointer is None: + logger.info("Checkpointer is None, ignoring...") + else: + logger.info(f"Deleting checkpoints for thread {thread_id}...") + delete_checkpoints(agent.checkpointer, str(thread_id)) + logger.success(f"Checkpoints for thread {thread_id} deleted successfully") return Response({"detail": "Thread deleted successfully"}) except Exception: + logger.exception(f"Error deleting thread {thread_id}:") return Response( {"detail": "Error deleting thread"}, status=status.HTTP_500_INTERNAL_SERVER_ERROR ) @@ -302,50 +324,12 @@ def _get_feedback_sender() -> LangSmithFeedbackSender: return LangSmithFeedbackSender() -@cache -def _get_context_provider(connection: str) -> PostgresContextProvider: - """Provide a configured `PostgresContextProvider` context provider. - - Args: - connection (str): A vector database connection for providing few-shot examples. - - Returns: - PostgresContextProvider: An instance of `PostgresContextProvider`. - """ - bq_billing_project = os.environ["BILLING_PROJECT_ID"] - bq_query_project = os.environ["QUERY_PROJECT_ID"] - - embedding_model = os.getenv("EMBEDDING_MODEL") - pgvector_collection = os.getenv("PGVECTOR_COLLECTION") - - if embedding_model and pgvector_collection: - embeddings = VertexAIEmbeddings(embedding_model) - - vector_store = PGVector( - embeddings=embeddings, - connection=connection, - collection_name=pgvector_collection, - use_jsonb=True, - ) - else: - vector_store = None - - context_provider = PostgresContextProvider( - billing_project=bq_billing_project, - query_project=bq_query_project, - metadata_vector_store=vector_store, - top_k=5, - ) - - return context_provider - - @contextmanager -def _get_sql_assistant(): - """Provide a configured `SQLAssistant`. +def _get_sql_agent() -> Generator[CompiledGraph]: + """Provide a configured ReAct agent. Yields: - Iterator[SQLAssistant]: An instance of `SQLAssistant`. + Iterator[CompiledGraph]: An instance of `CompiledGraph`. """ db_host = os.environ["DB_HOST"] db_port = os.environ["DB_PORT"] @@ -355,23 +339,47 @@ def _get_sql_assistant(): conn = f"postgresql://{db_user}:{db_password}@{db_host}:{db_port}/{db_name}" - context_provider = _get_context_provider(conn) + model = init_chat_model(MODEL_URI, temperature=0) + + def pre_model_hook(state: dict): + messages = state["messages"] + + # The last message in the pre_model_hook node will + # ALWAYS be a HumanMessage or a ToolMessage. + last_message = state["messages"][-1] + + # If this is the first message in the chat, we don't trim. + # If the last message is a ToolMessage, the agent has called a tool. + # This means we are in the middle of a chat turn and we also dont't trim. + if len(messages) == 1 or isinstance(last_message, ToolMessage): + return {"messages": []} + + # Otherwise, we are just starting a chat turn and we can trim the chat history. + remaining_messages = trim_messages( + messages, + token_counter=count_tokens_approximately, # The accurate counter is too slow. + max_tokens=MAX_TOKENS, + strategy="last", + start_on="human", + end_on="human", + include_system=True, + allow_partial=False, + ) - prompt_formatter = SQLPromptFormatter(vector_store=None) + return {"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *remaining_messages]} with PostgresSaver.from_conn_string(conn) as checkpointer: checkpointer.setup() - model = init_chat_model(MODEL_URI, temperature=0) - - assistant = SQLAssistant( + sql_agent = create_react_agent( model=model, - context_provider=context_provider, - prompt_formatter=prompt_formatter, + tools=get_tools(), + prompt=SQL_AGENT_SYSTEM_PROMPT, + pre_model_hook=pre_model_hook, checkpointer=checkpointer, ) - yield assistant + yield sql_agent def _stream_sql_assistant_response( @@ -387,79 +395,75 @@ def _stream_sql_assistant_response( Yields: Iterator[str]: JSON string containing the streaming status and the current step data. """ - steps = [] - last_chunk = None + events = [] + agent_state = None try: - logger.info("Calling SQLAssistant...") - with _get_sql_assistant() as assistant: - for mode, chunk in assistant.stream( - message=message, - config=config, + logger.info("Calling SQL Agent...") + with _get_sql_agent() as agent: + for mode, chunk in agent.stream( + input={"messages": [{"role": "user", "content": message}]}, stream_mode=["updates", "values"], - rewrite_query=True, + config=config, ): - last_chunk = chunk - - # Skip "values" chunks during streaming. We only need the final one at the end, - # as it contains the SQL Agent's final state. if mode == "values": + agent_state = chunk continue - step = process_chunk(chunk) + event = process_chunk(chunk) - if step is None: - continue + if event is not None: + events.append(event.model_dump()) + yield event.to_sse() - steps.append(step.model_dump()) + # The last event always contains the agent's final answer, + # so we use it to save the message pair in the database + assistant_message = event.data.content + error_message = None + logger.success("SQL Agent called successfully. Saving message pair...") - yield json.dumps({"status": "running", "data": step.model_dump_json()}) + "\n\n" + except google_api_exceptions.InvalidArgument: + logger.exception("Agent execution failed with Google API InvalidArgument error:") - # The last chunk represents the SQLAgent's final state, - # so we format it as the final response for the user. - response = format_sql_agent_response(last_chunk) - response["error_message"] = None - except Exception: - logger.exception(f"Error responding message {config['run_id']}:") - response = { - "error_message": ( - "Ops, algo deu errado! Ocorreu um erro inesperado. Por favor, tente novamente. " - "Se o problema persistir, avise-nos. Obrigado pela paciência!" - ), - "content": None, - "sql_queries": None, - } + assistant_message = None + error_message = UNEXPECTED_ERROR_MESSAGE + + if agent_state is not None: + model = init_chat_model(MODEL_URI, temperature=0) + total_tokens = model.get_num_tokens_from_messages(agent_state["messages"]) + + if total_tokens >= CONTEXT_WINDOW: + error_message = ( + "Sua última mensagem ultrapassou o limite de tamanho para esta conversa. " + "Por favor, tente dividir sua solicitação em partes menores " + "ou inicie uma nova conversa." + ) - logger.success("SQLAssistant called successfully") + yield StreamEvent( + type="error", data=EventData(error_details={"message": error_message}) + ).to_sse() - response["id"] = config["run_id"] + except Exception: + logger.exception(f"Unexpected error responding message {config['run_id']}:") + assistant_message = None + error_message = UNEXPECTED_ERROR_MESSAGE + + yield StreamEvent( + type="error", data=EventData(error_details={"message": error_message}) + ).to_sse() message_pair = MessagePair.objects.create( - id=response["id"], + id=config["run_id"], thread=thread, model_uri=MODEL_URI, user_message=message, - assistant_message=response["content"], - error_message=response["error_message"], - generated_queries=response["sql_queries"], - steps=steps, + assistant_message=assistant_message, + error_message=error_message, + events=events, ) + logger.success(f"Message pair {message_pair.id} saved successfully") - yield ( - json.dumps( - { - "status": "complete", - "data": { - "id": message_pair.id, - "user_message": message_pair.user_message, - "assistant_message": message_pair.assistant_message, - "error_message": message_pair.error_message, - "generated_queries": message_pair.generated_queries, - }, - } - ) - + "\n\n" - ) + yield StreamEvent(type="complete", data=EventData(run_id=message_pair.id)).to_sse() def _get_thread_by_id(thread_id: uuid.UUID) -> Thread: diff --git a/chatbot b/chatbot index 6a89f004..abd72e3b 160000 --- a/chatbot +++ b/chatbot @@ -1 +1 @@ -Subproject commit 6a89f004e51271c1669df755733ab21e58a91c84 +Subproject commit abd72e3b016ec7e4b6ce0953dde94b93ab7dfe47 diff --git a/poetry.lock b/poetry.lock index 7aa5fbb9..99544f86 100644 --- a/poetry.lock +++ b/poetry.lock @@ -426,7 +426,7 @@ files = [ [[package]] name = "chatbot" -version = "0.6.1" +version = "0.6.2" description = "" optional = false python-versions = ">=3.10,<4.0" @@ -4133,4 +4133,4 @@ cffi = ["cffi (>=1.11)"] [metadata] lock-version = "2.1" python-versions = ">=3.10,<3.13" -content-hash = "02d9ae83cf846222afc5a829393870f86e181af444079a785eb969082611c211" +content-hash = "c1f3ab7225c125661378a75d692f498ddb3cf4bd3680293eeeb4cd09c32f70a9" diff --git a/pyproject.toml b/pyproject.toml index e3676a64..333e9aee 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -46,6 +46,7 @@ djangorestframework-simplejwt = "^5.5.0" chatbot = {path = "chatbot"} langchain-google-vertexai = "2.0.24" langchain-postgres = "^0.0.14" +httpx = "^0.28.1" [tool.poetry.group.dev.dependencies] pre-commit = "^3.3.3"