diff --git a/example.env b/example.env index bcee7a6..918fade 100644 --- a/example.env +++ b/example.env @@ -1,5 +1,4 @@ -# Airbyte -GITHUB_TOKEN= +GITHUB_TOKEN= # Database credentials DB_HOST=db @@ -12,5 +11,5 @@ DB_PASSWORD=postgres DB_URL=postgresql://${DB_USER}:${DB_PASSWORD}@${DB_HOST}:${DB_PORT}/${DB_NAME} # Google cloud gemini -GEMINI_API_KEY= +GEMINI_API_KEY= GEMINI_MODEL_NAME=gemini-2.0-flash \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 2bd7127..c355f8d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,28 +1,283 @@ +aiohappyeyeballs==2.6.1 +aiohttp==3.13.2 +aiosignal==1.4.0 airbyte==0.31.3 -FastAPI==0.115.9 -uvicorn==0.34.3 -psycopg2-binary==2.9.10 -python-dotenv==1.1.0 -vanna==0.7.9 -chromadb -google-genai +airbyte-api==0.52.2 +airbyte-cdk==6.61.6 +airbyte_protocol_models_dataclasses==0.17.1 +airbyte_protocol_models_pdv2==0.13.1 +annotated-types==0.7.0 +anyascii==0.3.3 +anyio==4.11.0 +asn1crypto==1.5.1 +async-timeout==4.0.3 +attributes-doc==0.4.0 +attrs==25.4.0 +Authlib==1.6.5 +backoff==2.2.1 +backports.tarfile==1.2.0 +bcrypt==5.0.0 +beartype==0.22.6 +blinker==1.9.0 +boltons==25.0.0 +boto3==1.41.5 +botocore==1.41.5 +bracex==2.6 +build==1.3.0 +cachetools==6.2.2 +cattrs==25.3.0 +certifi==2025.10.5 +cffi==1.17.1 +charset-normalizer==3.4.4 +choreographer==1.2.1 +chromadb==1.3.5 +click==8.3.1 +coloredlogs==15.0.1 +contourpy==1.3.2 +coverage==7.12.0 +cryptography==44.0.3 +cycler==0.12.1 +cyclopts==4.3.0 +dataclasses-json==0.6.7 +dateparser==1.2.2 +Deprecated==1.3.1 +diskcache==5.6.3 +distro==1.9.0 +dnspython==2.8.0 +docstring_parser==0.17.0 +docutils==0.22.3 +dpath==2.2.0 +duckdb==1.4.2 +duckdb_engine==0.13.6 +dunamai==1.25.0 +durationpy==0.10 +email-validator==2.3.0 +exceptiongroup==1.3.0 +fastapi==0.115.9 +fastmcp==2.13.0.2 +filelock==3.20.0 +flasgger==0.9.7.1 +Flask==3.1.2 +flask-sock==0.7.0 +flatbuffers==25.9.23 +fonttools==4.60.1 +frozenlist==1.8.0 +fsspec==2025.10.0 +genson==1.3.0 +google-ai-generativelanguage==0.4.0 +google-api-core==2.28.1 +google-auth==2.43.0 google-cloud-aiplatform==1.96.0 +google-cloud-bigquery==3.30.0 +google-cloud-bigquery-storage==2.34.0 +google-cloud-core==2.5.0 +google-cloud-resource-manager==1.15.0 +google-cloud-secret-manager==2.25.0 +google-cloud-storage==2.19.0 +google-crc32c==1.7.1 +google-genai==1.52.0 +google-generativeai==0.4.1 +google-resumable-media==2.8.0 +googleapis-common-protos==1.72.0 +greenlet==3.2.4 +grpc-google-iam-v1==0.14.3 +grpcio==1.76.0 +grpcio-status==1.62.3 +h11==0.16.0 +hf-xet==1.2.0 +httpcore==1.0.9 +httptools==0.7.1 +httpx==0.28.1 +httpx-sse==0.4.3 +huggingface-hub==0.36.0 +humanfriendly==10.0 +idna==3.11 +importlib_metadata==8.4.0 +importlib_resources==6.5.2 +iniconfig==2.3.0 +isodate==0.6.1 +itsdangerous==2.2.0 +jaraco.classes==3.4.0 +jaraco.context==6.0.1 +jaraco.functools==4.3.0 +jeepney==0.9.0 +Jinja2==3.1.6 +jmespath==1.0.1 +joblib==1.5.2 +jsonpatch==1.33 +jsonpath-python==1.1.4 +jsonpointer==3.0.0 +jsonref==0.2 +jsonschema==4.25.1 +jsonschema-path==0.3.4 +jsonschema-specifications==2025.9.1 +kaleido==1.2.0 +keyring==25.7.0 +kiwisolver==1.4.9 +kubernetes==34.1.0 +langchain==0.1.16 +langchain-community==0.0.32 +langchain-core==0.1.42 +langchain-google-genai==0.0.11 +langchain-text-splitters==0.0.2 +langsmith==0.1.147 +logistro==2.0.1 +markdown-it-py==4.0.0 +MarkupSafe==3.0.3 +marshmallow==3.26.1 +matplotlib==3.9.2 +mcp==1.22.0 +mdurl==0.1.2 +mistune==3.1.4 +mmh3==5.2.0 +more-itertools==10.8.0 +mpmath==1.3.0 +multidict==6.7.0 +mypy_extensions==1.1.0 +narwhals==2.12.0 +networkx==3.4.2 +nltk==3.9.1 +numpy==1.26.4 +nvidia-cublas-cu12==12.6.4.1 +nvidia-cuda-cupti-cu12==12.6.80 +nvidia-cuda-nvrtc-cu12==12.6.77 +nvidia-cuda-runtime-cu12==12.6.77 +nvidia-cudnn-cu12==9.5.1.17 +nvidia-cufft-cu12==11.3.0.4 +nvidia-cufile-cu12==1.11.1.6 +nvidia-curand-cu12==10.3.7.77 +nvidia-cusolver-cu12==11.7.1.2 +nvidia-cusparse-cu12==12.5.4.2 +nvidia-cusparselt-cu12==0.6.3 +nvidia-nccl-cu12==2.26.2 +nvidia-nvjitlink-cu12==12.6.85 +nvidia-nvtx-cu12==12.6.77 +oauthlib==3.3.1 +ollama==0.6.0 onnxruntime==1.22.0 -langchain-google-genai<1.0.0 -langchain - +openapi-pydantic==0.5.1 +opentelemetry-api==1.27.0 +opentelemetry-exporter-otlp-proto-common==1.27.0 +opentelemetry-exporter-otlp-proto-grpc==1.27.0 +opentelemetry-proto==1.27.0 +opentelemetry-sdk==1.27.0 +opentelemetry-semantic-conventions==0.48b0 +orjson==3.11.4 +overrides==7.7.0 +packaging==23.2 +pandas==2.2.3 +pathable==0.4.4 +pathvalidate==3.3.1 +pillow==12.0.0 +platformdirs==4.5.0 +plotly==6.5.0 +pluggy==1.6.0 +posthog==5.4.0 +propcache==0.4.1 +proto-plus==1.26.1 +protobuf==4.25.8 +psutil==6.1.0 +psycopg==3.2.13 +psycopg-binary==3.2.13 +psycopg-pool==3.2.8 +psycopg2-binary==2.9.10 +py-key-value-aio==0.2.8 +py-key-value-shared==0.2.8 +pyarrow==21.0.0 +pyasn1==0.6.1 +pyasn1_modules==0.4.2 +pybase64==1.4.2 +pycparser==2.23 +pydantic==2.12.3 +pydantic-settings==2.12.0 +pydantic_core==2.41.4 +Pygments==2.19.2 +PyJWT==2.10.1 +pyOpenSSL==25.1.0 +pyparsing==3.2.5 +pyperclip==1.11.0 +PyPika==0.48.9 +pyproject_hooks==1.2.0 +pyrate-limiter==3.1.1 pytest==7.4.0 pytest-asyncio==0.21.1 -pytest-mock==3.12.0 pytest-cov==4.1.0 -httpx +pytest-mock==3.12.0 +pytest-timeout==2.4.0 +python-dateutil==2.9.0.post0 +python-dotenv==1.1.0 +python-multipart==0.0.20 +python-ulid==3.1.0 +pytz==2024.2 +PyYAML==6.0.3 +RapidFuzz==3.14.3 +referencing==0.36.2 +regex==2025.11.3 +requests==2.32.5 +requests-cache==1.2.1 requests-mock==1.11.0 - -torch>=2.2.0,<2.8.0 -transformers>=4.35.0,<5.0.0 -tokenizers>=0.15.0,<1.0.0 -scikit-learn>=1.5.2,<2.0.0 -numpy>=1.21.0,<1.27.0 - -matplotlib==3.9.2 -pandas==2.2.3 \ No newline at end of file +requests-oauthlib==2.0.0 +requests-toolbelt==1.0.0 +rich==13.9.4 +rich-click==1.9.4 +rich-rst==1.3.2 +rpds-py==0.29.0 +rsa==4.9.1 +s3transfer==0.15.0 +safetensors==0.7.0 +scikit-learn==1.7.2 +scipy==1.15.3 +SecretStorage==3.5.0 +serpyco-rs==1.17.1 +shapely==2.1.2 +shellingham==1.5.4 +simple-websocket==1.1.0 +simplejson==3.20.2 +six==1.17.0 +sniffio==1.3.1 +snowflake-connector-python==3.18.0 +snowflake-sqlalchemy==1.7.7 +sortedcontainers==2.4.0 +SQLAlchemy==2.0.44 +sqlalchemy-bigquery==1.12.0 +sqlparse==0.5.3 +sse-starlette==3.0.3 +starlette==0.45.3 +structlog==24.4.0 +sympy==1.14.0 +tabulate==0.9.0 +tenacity==8.5.0 +threadpoolctl==3.6.0 +tokenizers==0.22.1 +tomli==2.3.0 +tomlkit==0.13.3 +torch==2.7.1 +tqdm==4.67.1 +transformers==4.57.3 +triton==3.3.1 +typer==0.20.0 +typing-inspect==0.9.0 +typing-inspection==0.4.2 +typing_extensions==4.15.0 +tzdata==2025.2 +tzlocal==5.3.1 +Unidecode==1.4.0 +url-normalize==2.2.1 +urllib3==2.3.0 +uuid==1.30 +uuid7==0.1.0 +uv==0.8.24 +uvicorn==0.34.3 +uvloop==0.22.1 +vanna==0.7.9 +watchfiles==1.1.1 +wcmatch==10.0 +websocket-client==1.9.0 +websockets==15.0.1 +Werkzeug==3.1.3 +whenever==0.6.17 +wrapt==2.0.1 +wsproto==1.3.2 +xmltodict==0.14.2 +yarl==1.22.0 +zipp==3.23.0 diff --git a/src/api/controller/AskController.py b/src/api/controller/AskController.py index 523d814..adc8848 100644 --- a/src/api/controller/AskController.py +++ b/src/api/controller/AskController.py @@ -1,31 +1,37 @@ from src.assets.pattern.singleton import SingletonMeta from src.api.models import Question, Response -# Novas imports para gráficos -import io -import os -import uuid -import matplotlib.pyplot as plt -import pandas as pd - from langchain_google_genai import ChatGoogleGenerativeAI -from langchain.schema import SystemMessage, HumanMessage +from langchain.schema import SystemMessage, HumanMessage, AIMessage +from langchain.memory import ConversationBufferWindowMemory from google import genai from src.api.database.MyVanna import MyVanna +import json +from typing import Dict, Optional +import pandas as pd +import matplotlib.pyplot as plt +import uuid +import os +import hashlib +import re from src.assets.aux.env import env + # Gemini env vars GEMINI_API_KEY = env["GEMINI_API_KEY"] GEMINI_MODEL_NAME = env["GEMINI_MODEL_NAME"] class AskController(metaclass=SingletonMeta): - STATIC_DIR = "src/api/static/graficos/" + # Diretório estático para salvar gráficos + STATIC_DIR = os.path.join(os.getcwd(), "static", "graficos") + def __init__(self): self.client = genai.Client(api_key=GEMINI_API_KEY) - self.gen = ChatGoogleGenerativeAI( + # LLM principal para geração de respostas + self.llm = ChatGoogleGenerativeAI( model=GEMINI_MODEL_NAME, google_api_key=GEMINI_API_KEY, temperature=0, @@ -35,134 +41,374 @@ def __init__(self): convert_system_message_to_human=True ) + # Memória conversacional (mantém últimas 5 interações) + self.memory = ConversationBufferWindowMemory( + k=5, + memory_key="chat_history", + return_messages=True, + input_key="question", + output_key="answer" + ) + + # Instância do vanna self.vn = MyVanna(config={ 'print_prompt': False, 'print_sql': False, 'api_key': GEMINI_API_KEY, 'model_name': GEMINI_MODEL_NAME }) - self.vn.prepare() - # Garantir que a pasta para gráficos exista - if not os.path.exists(self.STATIC_DIR): - os.makedirs(self.STATIC_DIR) + # Cache simples para queries SQL (em memória) + self.sql_cache: Dict[str, str] = {} + self.result_cache: Dict[str, any] = {} + + def _generate_chart_if_requested(self, resultado, wants_chart: bool): + """ + Gera um gráfico a partir do resultado se wants_chart for True. + Salva o gráfico em arquivo e retorna o link, ou mensagem de erro se não houver dados. + """ + if not wants_chart: + return None + # Garante que o diretório existe + os.makedirs(self.STATIC_DIR, exist_ok=True) + df = pd.DataFrame(resultado) + if df.empty: + return {"output": "Não há dados suficientes para gerar um gráfico."} + + plt.figure(figsize=(8, 5)) + if df.shape[1] >= 2: + x = df.columns[0] + y = df.columns[1] + plt.bar(df[x], df[y]) + plt.xlabel(x) + plt.ylabel(y) + plt.title("Gráfico gerado a partir dos dados") + else: + plt.plot(df[df.columns[0]]) + plt.title("Gráfico gerado a partir dos dados") + + filename = f"{uuid.uuid4()}.png" + filepath = os.path.join(self.STATIC_DIR, filename) + plt.tight_layout() + plt.savefig(filepath) + plt.close() - def ask(self, question: Question): + link = f"http://localhost:8000/static/graficos/{filename}" + return {"output": f"Gráfico gerado: [Clique aqui para visualizar]({link})", "grafico_url": link} + + def _detect_chart_request(self, question: str) -> bool: + """ + Detecta se a pergunta do usuário sugere a geração de um gráfico. + """ + chart_keywords = [ + "gráfico", "grafico", "plot", "visualização", "visualizacao", + "chart", "plotar", "desenhar gráfico", "desenhar grafico", "mostrar gráfico", + "mostrar grafico", "visualize", "visualizar", "figure", "figura" + ] + question_lower = question.lower() + return any(kw in question_lower for kw in chart_keywords) + + def _preprocess_question(self, question: str) -> str: + """ + Pré-processa a pergunta usando LLM com contexto da memória + Versão simplificada sem LLMChain para evitar problemas de parsing + """ try: - # Detectar se usuário quer gráfico - gerar_grafico = any( - kw in question.question.lower() - for kw in ["gráfico", "grafico", "chart", "visualizar", "plot"] + # Recupera histórico da memória + memory_vars = self.memory.load_memory_variables({}) + chat_history = memory_vars.get("chat_history", []) + + # Monta contexto do histórico + history_text = "" + if chat_history: + history_text = "Histórico da conversa:\n" + for msg in chat_history[-5:]: # Últimas 5 mensagens + if isinstance(msg, HumanMessage): + history_text += f"Usuário: {msg.content}\n" + elif isinstance(msg, AIMessage): + history_text += f"Assistente: {msg.content}\n" + history_text += "\n" + + # Monta o prompt + system_prompt = """Você é um assistente especializado em processar perguntas sobre dados do GitHub. + + Sua função é: + 1. Analisar o histórico da conversa para entender o contexto + 2. Resolver referências contextuais (ex: "e no mês passado?", "mostre mais detalhes", "e o outro repositório?") + 3. Normalizar expressões temporais: + - "3 meses" → "90 dias" + - "1 ano e 2 meses" → "425 dias" + - Meses separados = 30 dias cada + 4. Normalizar terminologia: + - "mudança" → "commit" + - "alteração" → "commit" + 5. Expandir a pergunta com contexto necessário do histórico + 6. Desenvolver análises robustas dos dados extraídos para que insight valiosos sejam extraídos + 7. + + REGRAS CRÍTICAS: + - Se a pergunta fizer referência a algo anterior ("e aquele", "o outro", "também"), inclua o contexto explícito + - Se não houver referência contextual, retorne a pergunta apenas normalizada + """ + + # Monta mensagens + messages = [ + SystemMessage(content=system_prompt) + ] + + # Adiciona histórico se existir + if history_text: + messages.append(HumanMessage(content=history_text)) + + # Adiciona pergunta atual + messages.append(HumanMessage(content=f"Pergunta a processar: {question}")) + + # Chama LLM + response = self.llm.invoke(messages) + + # Extrai conteúdo da resposta + if hasattr(response, 'content'): + processed = response.content.strip() + else: + processed = str(response).strip() + + # Remove qualquer explicação extra (pega só a primeira linha) + processed = processed.split('\n')[0].strip() + + return processed + + except Exception as e: + print(f"[Warning] Erro no preprocessing: {e}. Usando pergunta original.") + return question + + def _get_cache_key(self, text: str) -> str: + """Gera chave de cache baseada no hash da pergunta normalizada""" + normalized = text.lower().strip() + return hashlib.md5(normalized.encode()).hexdigest() + + def _validate_sql(self, sql: str) -> tuple[bool, Optional[str]]: + """Valida SQL gerado para segurança""" + sql_upper = sql.upper().strip() + + # Whitelist: apenas SELECT permitido + if not sql_upper.startswith("SELECT"): + return False, "Apenas queries SELECT são permitidas" + + # Blacklist com word boundaries (evita falsos positivos) + dangerous_patterns = [ + r'\bDELETE\b', + r'\bDROP\b', + r'\bTRUNCATE\b', + r'\bINSERT\b', + r'\bUPDATE\b', + r'\bALTER\b', + r'\bCREATE\s+TABLE\b', + r'\bCREATE\s+INDEX\b', + r'\bCREATE\s+DATABASE\b', + r'\bGRANT\b', + r'\bREVOKE\b', + r'\bEXEC\b', + r'\bEXECUTE\b', + r';\s*\w+', # SQL injection + ] + + for pattern in dangerous_patterns: + if re.search(pattern, sql_upper): + keyword = pattern.replace(r'\b', '').replace(r'\s+', ' ') + return False, f"Operação '{keyword}' não é permitida" + + # Limite de complexidade (número de JOINs) + join_count = sql_upper.count("JOIN") + if join_count > 10: + return False, "Query muito complexa (máximo 10 JOINs)" + + return True, None + + def _format_response_with_context(self, question: str, sql: str, result: any) -> str: + """Formata resposta final usando LLM com contexto conversacional""" + + # Recupera histórico da memória + memory_vars = self.memory.load_memory_variables({}) + chat_history = memory_vars.get("chat_history", []) + + # Monta contexto do histórico + history_context = "" + if chat_history: + history_context = "\n\nContexto da conversa anterior:\n" + for msg in chat_history[-3:]: # Últimas 3 mensagens + if isinstance(msg, HumanMessage): + history_context += f"Usuário: {msg.content}\n" + elif isinstance(msg, AIMessage): + history_context += f"Assistente: {msg.content}\n" + + prompt = f""" + Você é um assistente especializado em análise de dados do GitHub. + + {history_context} + + Pergunta atual: "{question}" + + SQL gerado e executado: + ```sql + {sql} + ``` + + Resultado da consulta: {result} + + Com base no contexto da conversa e nos resultados, gere uma resposta: + 1. Clara e direta + 2. Em linguagem natural + 3. Destacando insights relevantes + 4. Relacionando com perguntas anteriores se aplicável + 5. Formato estruturado se houver múltiplos dados + + Responda de forma conversacional e útil. + """ + + try: + response = self.client.models.generate_content( + model=GEMINI_MODEL_NAME, + contents=prompt, + config={ + "response_mime_type": "application/json", + "response_schema": list[Response], + } ) + return response.parsed[0].texto + except Exception as e: + print(f"[Error] Erro ao formatar resposta: {e}") + # Fallback para resposta simples + return f"Consulta executada com sucesso. Resultado: {result}" - # ----------------------------- - # Pergunta padrão (sem gráfico) - # ----------------------------- - if not gerar_grafico: - mensagem = [ - SystemMessage(content="""Você é um assistente em um sistema que de chat AI, - seu trabalho é receber uma mensagem de um humano e tratar ela para ser processada - por outra IA que possui falhas. Dentro dos tratamentos necessários estão: - -Trocar datas que são dadas em formato que não sejam dias, por exemplo 3 mêses, - e transformar em dias, 90 dias. Outro exemplo, 1 ano e 2 meses, trocar por 425 dias - (serão considerados que os mêses separados terão 30 dias) - -Quando usada a expressão "mudança" referente ao repositório, você trocara por "commit", - por exemplo, "qual foi a ultima mudança feita no repositório?" será trocado por - "qual foi o ultimo commit feito no repositório?") - NÃO explique, NÃO confirme, NÃO dê exemplos. Apenas RETORNE a mensagem tratada. - Mensagem a ser processada:"""), - HumanMessage(content=question.question) - ] - - ai_mensagem = self.gen.invoke(mensagem) - sql_gerado = self.vn.generate_sql(ai_mensagem.content) - - if "SELECT" not in sql_gerado.upper(): - return {"output": "Não consegui entender sua pergunta bem o suficiente para gerar uma resposta SQL válida."} + def ask(self, question: Question, session_id: Optional[str] = None) -> dict: + """ + Processa pergunta com contexto conversacional + + Args: + question: Objeto Question com a pergunta do usuário + session_id: ID da sessão para memória multi-usuário (futuro) + """ + + try: + original_question = question.question + print(f"[Original] {original_question}") - resultado = self.vn.run_sql(sql_gerado) + # Detectar intenção de gráfico + wants_chart = self._detect_chart_request(original_question) - if not resultado: - return {"output": "A consulta foi feita, mas não há dados correspondentes no banco."} - - prompt = f""" - Você é um assistente que responde perguntas sobre dados extraídos do GitHub. - Pergunta do usuário: "{question.question}" - Resultado da consulta SQL: {resultado} - Gere uma resposta clara e útil para o usuário, explicando o que o resultado significa. - """ - response = self.client.models.generate_content( - model=GEMINI_MODEL_NAME, - contents=prompt, - config={ - "response_mime_type": "application/json", - "response_schema": list[Response], - } - ) - texto = response.parsed[0].texto - return {"output": texto} + # Etapa 1: Pré-processar com contexto + processed_question = self._preprocess_question(original_question) + print(f"[Preprocessed] {processed_question}") - # ----------------------------- - # Pergunta com gráfico - # ----------------------------- + # Etapa 2: Verificar cache de SQL + cache_key = self._get_cache_key(processed_question) + + if cache_key in self.sql_cache: + print(f"[Cache Hit] SQL encontrado no cache") + sql_gerado = self.sql_cache[cache_key] else: - mensagem = [ - SystemMessage(content="""Você é um assistente em um sistema que de chat AI, - seu trabalho é receber uma mensagem de um humano e tratar ela para ser processada - por outra IA que possui falhas. Dentro dos tratamentos necessários estão: - -Trocar datas que são dadas em formato que não sejam dias, por exemplo 3 mêses, - e transformar em dias, 90 dias. Outro exemplo, 1 ano e 2 meses, trocar por 425 dias - (serão considerados que os mêses separados terão 30 dias) - -Quando usada a expressão "mudança" referente ao repositório, você trocara por "commit", - por exemplo, "qual foi a ultima mudança feita no repositório?" será trocado por - "qual foi o ultimo commit feito no repositório?") - retire as palavras faça gráfico da mensagem, - NÃO explique, NÃO confirme, NÃO dê exemplos. Apenas RETORNE a mensagem tratada. - Mensagem a ser processada:"""), - HumanMessage(content=question.question) - ] - - ai_mensagem = self.gen.invoke(mensagem) - sql_gerado = self.vn.generate_sql(ai_mensagem.content) - - if "SELECT" not in sql_gerado.upper(): - return {"output": "Não consegui entender sua pergunta bem o suficiente para gerar uma resposta SQL válida."} + # Gerar SQL com Vanna + sql_gerado = self.vn.generate_sql(processed_question) - resultado = self.vn.run_sql(sql_gerado) + # Validar SQL + is_valid, error_msg = self._validate_sql(sql_gerado) + if not is_valid: + return { + "output": f"Query inválida: {error_msg}", + "error": True + } + + # Armazenar no cache + self.sql_cache[cache_key] = sql_gerado + print(f"[Cache Miss] SQL gerado e armazenado") + + print(f"[SQL] {sql_gerado}") + + # Etapa 3: Verificar cache de resultados + result_cache_key = hashlib.md5(sql_gerado.encode()).hexdigest() + if result_cache_key in self.result_cache: + print(f"[Cache Hit] Resultado encontrado no cache") + resultado = self.result_cache[result_cache_key] + else: + # Executar SQL + resultado = self.vn.run_sql(sql_gerado) + if not resultado: - return {"output": "A consulta foi feita, mas não há dados correspondentes no banco."} - - # ----------------------------- - # Geração do gráfico - # ----------------------------- - df = pd.DataFrame(resultado) - if df.empty: - return {"output": "Não há dados suficientes para gerar um gráfico."} - - plt.figure(figsize=(8,5)) - if df.shape[1] >= 2: - x = df.columns[0] - y = df.columns[1] - plt.bar(df[x], df[y]) - plt.xlabel(x) - plt.ylabel(y) - plt.title("Gráfico gerado a partir dos dados") - else: - plt.plot(df[df.columns[0]]) - plt.title("Gráfico gerado a partir dos dados") - - # Salvar gráfico como arquivo - filename = f"{uuid.uuid4()}.png" - filepath = os.path.join(self.STATIC_DIR, filename) - plt.tight_layout() - plt.savefig(filepath) - plt.close() - - # Montar link clicável - link = f"http://localhost:8000/static/graficos/{filename}" - return {"output": f"Gráfico gerado: [Clique aqui para visualizar]({link})", "grafico_url": link} + # Salvar na memória mesmo sem resultado + self.memory.save_context( + inputs={"question": original_question}, + outputs={"answer": "Não há dados correspondentes no banco."} + ) + return { + "output": "A consulta foi executada, mas não há dados correspondentes.", + "sql": sql_gerado + } + + # Armazenar resultado no cache + self.result_cache[result_cache_key] = resultado + print(f"[Cache Miss] Resultado obtido e armazenado") + + # Etapa 4: Gerar gráfico se solicitado + chart_result = self._generate_chart_if_requested(resultado, wants_chart) + + # Etapa 5: Formatar resposta com contexto + resposta_formatada = self._format_response_with_context( + question=original_question, + sql=sql_gerado, + result=resultado + ) + + # Etapa 6: Salvar na memória + self.memory.save_context( + inputs={"question": original_question}, + outputs={"answer": resposta_formatada} + ) + + response = { + "output": resposta_formatada, + "sql": sql_gerado, + "cached": result_cache_key in self.result_cache, + "wants_chart": wants_chart + } + if chart_result: + response.update(chart_result) + return response except Exception as e: - return {"output": f"Ocorreu um erro ao processar sua pergunta: {str(e)}"} \ No newline at end of file + import traceback + error_msg = f"Erro ao processar pergunta: {str(e)}" + print(f"[Error] {error_msg}") + print(f"[Error] Traceback: {traceback.format_exc()}") + + # Salvar erro na memória + try: + self.memory.save_context( + inputs={"question": question.question}, + outputs={"answer": error_msg} + ) + except: + pass + + return { + "output": error_msg, + "error": True, + "wants_chart": False + } + + def clear_memory(self): + """Limpa o histórico da conversa""" + self.memory.clear() + print("[Memory] Histórico limpo") + + def get_conversation_history(self) -> list: + """Retorna o histórico da conversa""" + memory_vars = self.memory.load_memory_variables({}) + return memory_vars.get("chat_history", []) + + def clear_cache(self): + """Limpa os caches de SQL e resultados""" + self.sql_cache.clear() + self.result_cache.clear() + print("[Cache] Caches limpos") diff --git a/src/api/database/MyVanna.py b/src/api/database/MyVanna.py index a21ab42..18d01ef 100644 --- a/src/api/database/MyVanna.py +++ b/src/api/database/MyVanna.py @@ -164,224 +164,50 @@ def run_sql(self, sql): def prepare(self): + """ + Prepara o Vanna com treinamento inicial CORRIGIDO + + Args: + force_retrain: Se True, força o retreinamento mesmo se já existir cache + """ + + # Conectar ao banco self.connect_to_postgres( - host = DB_HOST, - port = DB_PORT, - dbname = DB_NAME, - user = DB_USER, - password = DB_PASSWORD + host=DB_HOST, + port=DB_PORT, + dbname=DB_NAME, + user=DB_USER, + password=DB_PASSWORD ) + + print("[Vanna] → Iniciando treinamento (isso consome API quota)...") + + # ========================================================================= + # ETAPA 1: Treinar com DDL (estrutura real do banco) + # ========================================================================= + print("[Vanna] 1/2 Treinando DDL...") self.train(ddl=self.get_schema()) - self.train(documentation=""" -Table: user_info - - id: Bigint primary key with default value from sequence - - login: Required username field (character varying) - - html_url: Required profile URL field (text) - -Table: milestone - - id: Bigint primary key with default value from sequence - - repository_id: Associated repository ID (integer, required) - - title: Milestone title (text, required) - - description: Milestone description (text, optional) - - number: Milestone number (integer, required) - - state: Milestone state (character varying, required) - - created_at: Creation timestamp with time zone - - updated_at: Update timestamp with time zone - - creator: Creator user ID (bigint, required) - -Table: repository - - id: Integer primary key with default value from sequence - - name: Repository name (character varying, required) - -Table: branch - - id: Bigint primary key with default value from sequence - - name: Branch name (character varying, required) - - repository_id: Associated repository ID (integer, required) - -Table: issue - - id: Bigint primary key with default value from sequence - - title: Issue title (text, required) - - body: Issue body/description (text, optional) - - number: Issue number (integer, required) - - html_url: Issue URL (text, optional) - - created_at: Creation timestamp with time zone - - updated_at: Update timestamp with time zone - - created_by: Creator user ID (bigint, required) - - repository_id: Associated repository ID (bigint, required) + print("[Vanna] 2/2 Treinando SQL Examples e documentação...") + self.train(sql=open("src/api/database/sql_examples.sql").read()) - milestone_id: Associated milestone ID (bigint, optional) - -Table: pull_requests - - id: Bigint primary key with default value from sequence - - created_by: Creator user ID (bigint, required) - - repository_id: Associated repository ID (integer, required) - - number: Pull request number (integer, required) - - state: Pull request state (character varying, required) - - title: Pull request title (text, required) - - body: Pull request body/description (text, optional) - - html_url: Pull request URL (text, required) - - created_at: Creation timestamp with time zone - - updated_at: Update timestamp with time zone - - milestone_id: Associated milestone ID (bigint, optional) - -Table: commits - - id: Bigint primary key with default value from sequence - - user_id: Author user ID (bigint, required) - - branch_id: Associated branch ID (integer, optional) - - pull_request_id: Associated pull request ID (bigint, optional) - - created_at: Creation timestamp with time zone - - message: Commit message (text, required) - - sha: Commit SHA hash (character varying, required) - - html_url: Commit URL (text, optional) - -Table: parents_commits - - id: Integer primary key with default value from sequence - - parent_sha: Parent commit SHA hash (character varying, required) - - commit_id: Child commit ID (integer, required) - -Table: issue_assignees - - issue_id: Issue ID (bigint, required, part of primary key) - - user_id: Assigned user ID (bigint, required, part of primary key) + self.train(documentation=""" + O banco contempla atividades de GitHub: -Table: pull_request_assignees + - user_info: usuários + - repository: repositórios + - branch: branches dos repositórios + - issue: issues criadas + - pull_requests: PRs + - commits: commits de usuários, podendo referenciar PRs + - issue_assignees / pull_request_assignees: responsáveis + - milestone: grupo de issues e PRs - pull_request_id: Pull request ID (bigint, required, part of primary key) + Consultas esperadas: ranking, agregações, contagem de atividades, obtenção de repositórios mais ativos e etc. - user_id: Assigned user ID (bigint, required, part of primary key) - """) + """) - self.train(sql=""" - -- 1. Repositórios com mais issues abertas - SELECT - r.name AS repositorio, - COUNT(*) AS total_issues_abertas, - MAX(i.created_at) AS data_ultima_issue - FROM - issue i - JOIN - repository r ON i.repository_id = r.id - WHERE - i.state = 'open' - GROUP BY - r.name - ORDER BY - total_issues_abertas DESC - LIMIT 10; - """) - - self.train(sql=""" - -- 2. Top 5 usuários com mais commits registrados - SELECT - u.login, - COUNT(*) AS total_commits - FROM - commits c - JOIN - user_info u ON c.user_id = u.id - GROUP BY - u.login - ORDER BY - total_commits DESC - LIMIT 5; - """) - - self.train(sql=""" - -- 3. Total de pull requests abertos por repositório - SELECT - r.name AS repositorio, - COUNT(*) AS total_pr_abertos - FROM - pull_requests pr - JOIN - repository r ON pr.repository_id = r.id - WHERE - pr.state = 'open' - GROUP BY - r.name - ORDER BY - total_pr_abertos DESC; - """) - - self.train(sql=""" - -- 4. Número de issues por milestone - SELECT - m.title AS milestone, - COUNT(*) AS total_issues - FROM - issue i - JOIN - milestone m ON i.milestone_id = m.id - GROUP BY - m.title - ORDER BY - total_issues DESC; - """) - - self.train(sql=""" - -- 5. Commits feitos por branch - SELECT - b.name AS branch, - COUNT(*) AS total_commits - FROM - commits c - JOIN - branch b ON c.branch_id = b.id - GROUP BY - b.name - ORDER BY - total_commits DESC; - """) + \ No newline at end of file diff --git a/src/api/database/sql_examples.sql b/src/api/database/sql_examples.sql new file mode 100644 index 0000000..4a320f7 --- /dev/null +++ b/src/api/database/sql_examples.sql @@ -0,0 +1,35 @@ +-- Exemplo 1: Contar issues por repositório +SELECT r.name, COUNT(*) AS total_issues +FROM issue i +JOIN repository r ON r.id = i.repository_id +GROUP BY r.name; + +-- Exemplo 2: Contar commits por usuário +SELECT u.login, COUNT(*) AS total_commits +FROM commits c +JOIN user_info u ON u.id = c.user_id +GROUP BY u.login; + +-- Exemplo 3: Repositório mais movimentado (issues + PRs + commits) +WITH activity AS ( + SELECT + r.id, + COUNT(DISTINCT i.id) + + COUNT(DISTINCT pr.id) + + COUNT(DISTINCT c.id) AS score + FROM repository r + LEFT JOIN issue i ON i.repository_id = r.id + LEFT JOIN pull_requests pr ON pr.repository_id = r.id + LEFT JOIN commits c ON c.pull_request_id = pr.id + GROUP BY r.id +) +SELECT id FROM activity ORDER BY score DESC LIMIT 1; + +-- Exemplo 4: Tasks (issues) por usuário em um repositório específico +SELECT + u.login, + COUNT(*) AS total_tasks +FROM issue i +JOIN user_info u ON u.id = i.created_by +WHERE i.repository_id = 1 +GROUP BY u.login;