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, )