text2sql_backend / src /main.py
LightRT's picture
Fixed Langgraph Workflow
0fb7e57
Raw History Blame Contribute Delete
2.86 kB
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,
)