File size: 2,856 Bytes
0fb7e57
52adb86
0fb7e57
52adb86
 
3843ff9
0fb7e57
 
 
 
 
3843ff9
 
52adb86
0fb7e57
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52adb86
 
 
0fb7e57
 
52adb86
0fb7e57
 
 
 
 
52adb86
0fb7e57
 
 
 
52adb86
0fb7e57
 
 
 
 
 
 
52adb86
 
0fb7e57
 
52adb86
 
0fb7e57
 
 
8c87f5a
0fb7e57
52adb86
0fb7e57
 
 
 
 
52adb86
0fb7e57
 
 
 
 
 
52adb86
0fb7e57
 
 
 
 
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
from fastapi import FastAPI , HTTPException
from src.embedding import create_embeddings
from src.graph import build_agent,  AgentContext
from pydantic import BaseModel , Field
import os 
from dotenv import load_dotenv
import asyncio
from psycopg_pool import AsyncConnectionPool
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from contextlib import asynccontextmanager
import logging

load_dotenv()

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("text2sql")

DB_URI = os.getenv("DATABASE_URI")

@asynccontextmanager
async def lifespan(app: FastAPI):
    async with AsyncConnectionPool(
        conninfo=DB_URI,
        min_size=1,
        max_size=15,
        max_idle=300,          
        max_lifetime=1800,    
        reconnect_timeout=10,
        kwargs={"autocommit": True, "prepare_threshold": 0},
        check=AsyncConnectionPool.check_connection,
        open=False,            
    ) as pool:
        await pool.open(wait=True, timeout=15)
        checkpointer = AsyncPostgresSaver(pool)
        await checkpointer.setup()
        app.state.pool = pool
        app.state.agent = build_agent(checkpointer)
        yield

app = FastAPI(
    title="Text2SQL Agent API",
    description="A production-grade backend powering LangGraph agent.",
    version="1.0.0",
    lifespan=lifespan)

class ChatRequest(BaseModel):
    connection_url : str = Field(...)
    message: str = Field(...)
    user_id: str = Field(...)
    thread_id: str = Field(...)

class ChatResponse(BaseModel):
    status: str
    thread_id: str
    response: str

class UploadRequest(BaseModel) :
    connection_url : str = Field(...)
    user_id : str = Field(...)

@app.post("/upload")
async def upload_url(request : UploadRequest):
    await asyncio.to_thread(create_embeddings , request.connection_url , request.user_id)

    return {
        "status": "success",
        "message": "You can now chat with the agent."
    }

@app.post("/chat",response_model=ChatResponse)
async def chat_endpoint(request: ChatRequest):
    agent = app.state.agent  

    config = {"configurable": {"thread_id": request.thread_id}}

    try:
        result = await agent.ainvoke(
            {"messages": [{"role": "user", "content": request.message}]},
            config=config,
            context=AgentContext(user_id=request.user_id , connection_url=request.connection_url),
        )
    except Exception:
         logger.exception("Agent processing failed")
         raise HTTPException(status_code=500, detail="Agent Processing Error!")
    output_messages = result.get("messages", [])
    if not output_messages:
        raise HTTPException(status_code=500, detail="No messages returned from the agent.")

    return ChatResponse(
        status="success",
        thread_id=request.thread_id,
        response=output_messages[-1].content,
    )