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)}"
        )