Spaces:
Sleeping
Sleeping
File size: 5,735 Bytes
4db8795 99bbd9b 4db8795 99bbd9b 4db8795 99bbd9b 4db8795 99bbd9b 4db8795 7c8d332 5c20f4d 37d4305 5c20f4d 99bbd9b 4db8795 37d4305 99bbd9b 4db8795 99bbd9b 4db8795 37d4305 4db8795 99bbd9b 7c8d332 5c20f4d | 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 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 | # # from fastapi import APIRouter, HTTPException
# # from app.models import SQLQueryRequest, SQLQueryResponse
# # from app.services.sql_agent import execute_query
# # router = APIRouter()
# # @router.post("/query", response_model=SQLQueryResponse)
# # async def query_database(request: SQLQueryRequest):
# # try:
# # result = execute_query(request.query)
# # return SQLQueryResponse(result=result)
# # except ValueError as e:
# # raise HTTPException(status_code=400, detail=str(e))
# # except Exception as e:
# # raise HTTPException(status_code=500, detail=str(e))
# # app/api/v1/endpoints/sql_query.py
# from fastapi import APIRouter, HTTPException
# from pydantic import BaseModel
# from app.services.sql_agent_instance import sql_agent
# router = APIRouter()
# class SQLQueryRequest(BaseModel):
# query: str
# class SQLQueryResponse(BaseModel):
# result: str
# @router.post("/query", response_model=SQLQueryResponse)
# async def query_database(request: SQLQueryRequest):
# try:
# result = sql_agent.execute_query(request.query)
# return SQLQueryResponse(result=result)
# except ValueError as e:
# raise HTTPException(status_code=400, detail=str(e))
# except Exception as e:
# raise HTTPException(status_code=500, detail=str(e))
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel
from app.services.sql_agent_instance import sql_agent
from typing import Optional
import uuid
from langchain_core.messages import AIMessage
from sqlalchemy import text
from app.api.v1.auth import get_db
router = APIRouter()
class SQLQueryRequest(BaseModel):
query: str
thread_id: Optional[str] = None
username: Optional[str] = "developer"
class SQLQueryResponse(BaseModel):
result: str
thread_id: str ## client can use this to continue the conversation
@router.post("/query", response_model=SQLQueryResponse)
async def query_database(request: SQLQueryRequest):
try:
## generate if not provided thread id
thread_id = request.thread_id or str(uuid.uuid4())
## add debug
print(f"Thread ID: {thread_id}, Query: {request.query}")
result = sql_agent.execute_query(request.query, config={"configurable": {"thread_id": thread_id}})
print(f"Result: {result}")
# Extract the SQL query from the LangGraph state messages
sql_query = None
try:
state = sql_agent.app.get_state({"configurable": {"thread_id": thread_id}})
messages = state.values.get("messages", [])
for msg in reversed(messages):
if hasattr(msg, "tool_calls") and msg.tool_calls:
for tc in msg.tool_calls:
# Check if execute_query tool was called
if tc.get("name") == "execute_query":
# The argument name might be query
sql_query = tc.get("args", {}).get("query")
break
if sql_query:
break
except Exception as e:
print(f"Error extracting SQL query from state: {e}")
# Save to query history if found
if sql_query:
try:
username = request.username or "developer"
conn = get_db()
cur = conn.cursor()
cur.execute(
"INSERT INTO query_history (username, natural_query, generated_sql) VALUES (?, ?, ?)",
(username, request.query, sql_query)
)
conn.commit()
conn.close()
print("Logged successfully executed query to history.")
except Exception as e:
print(f"Error saving query to history: {e}")
return SQLQueryResponse(result=result, thread_id=thread_id)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@router.post("/query/explain")
async def explain_query(request: SQLQueryRequest):
if not sql_agent.db:
raise HTTPException(
status_code=400,
detail="Database connection is not established. Please set up the connection first."
)
try:
query = request.query.strip().rstrip(";")
dialect = sql_agent.db._engine.dialect.name
# Dialect specific syntax
if dialect == "postgresql":
explain_sql = f"EXPLAIN (FORMAT JSON) {query}"
elif dialect == "mysql":
explain_sql = f"EXPLAIN FORMAT=JSON {query}"
else:
explain_sql = f"EXPLAIN QUERY PLAN {query}"
with sql_agent.db._engine.connect() as connection:
result = connection.execute(text(explain_sql))
rows = result.fetchall()
# Format results into a serializable list of dicts/lists
plan_output = []
for row in rows:
if dialect in ["postgresql", "mysql"]:
# Usually returns a single column containing JSON
plan_output.append(row[0])
else:
# SQLite returns (id, parent, notused, detail)
plan_output.append(dict(row._mapping))
return {
"dialect": dialect,
"explain_query": explain_sql,
"plan": plan_output
}
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"An error occurred while explaining the query: {str(e)}"
)
|