Spaces:
Running
Running
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,
) |