Untitled

 avatar
unknown
plain_text
9 months ago
2.5 kB
11
Indexable
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from typing import Dict, List
import datetime

from rag_logic import get_rag_chain
from langchain.memory import ConversationSummaryBufferMemory
from langchain_google_genai import ChatGoogleGenerativeAI

app = FastAPI()

sessions: Dict[str, List[Dict[str, str]]] = {}
memories: Dict[str, ConversationSummaryBufferMemory] = {}

gemini_llm = ChatGoogleGenerativeAI(model="gemini-2.5-flash")

class ChatRequest(BaseModel):
    session_id: str
    input: str

class ChatResponse(BaseModel):
    answer: str
    chat_history: List[Dict[str, str]]

def append_history(session_id, role, content):
    ts = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
    sessions.setdefault(session_id, []).append({
        "role": role,
        "content": content,
        "timestamp": ts
    })

@app.post("/", response_model=ChatResponse)
def chat(req: ChatRequest):
    session_id = req.session_id
    user_input = req.input.strip()

    if not user_input:
        raise HTTPException(status_code=400, detail="Empty input")

    # Initialize or reuse session memory
    if session_id not in memories:
        memories[session_id] = ConversationSummaryBufferMemory(
            llm=gemini_llm,
            memory_key="chat_history",
            return_messages=True,
            max_token_limit=500
        )

    memory = memories[session_id]
    rag_chain = get_rag_chain()

    # Get past chat history from memory
    memory_variables = memory.load_memory_variables({})
    print(memory_variables)

    append_history(session_id, "human", user_input)

    try:
        # Inject summarized chat history into RAG chain
        response = rag_chain.invoke({
            "input": user_input,
            "chat_history": memory_variables.get("chat_history")
        })
        
        answer = response["answer"]

        # Update the memory with the new interaction
        memory.save_context({"input": user_input}, {"answer": answer})

    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))

    append_history(session_id, "ai", answer)

    return {"answer": answer, "chat_history": sessions[session_id]}

@app.post("/clear/{session_id}")
def clear_chat(session_id: str):
    sessions.pop(session_id, None)
    memories.pop(session_id, None)
    return {"status": "cleared"}
Editor is loading...
Leave a Comment