🐛 fix(chat.py): remove unused import and variable 'flow_data_store' to improve code readability
✨ feat(chat.py): add support for cache_manager to handle chat history and cache langchain_object 🆕 feat(chat/manager.py): add ChatManager class to handle websocket connections and chat history 🆕 feat(chat/manager.py): add ChatHistory class to manage chat history for each client 🆕 feat(chat/manager.py): add methods to handle websocket connections, send messages, and process chat messages 🆕 feat(chat/manager.py): add method to set cache for a client and handle websocket communication
This commit is contained in:
parent
f65efbc43b
commit
1a51e90c43
2 changed files with 227 additions and 11 deletions
|
|
@ -6,19 +6,16 @@ from langflow.api.v1.schemas import BuildStatus, BuiltResponse, InitResponse, St
|
||||||
from langflow.services import service_manager, ServiceType
|
from langflow.services import service_manager, ServiceType
|
||||||
from langflow.graph.graph.base import Graph
|
from langflow.graph.graph.base import Graph
|
||||||
from langflow.utils.logger import logger
|
from langflow.utils.logger import logger
|
||||||
from cachetools import LRUCache
|
|
||||||
|
|
||||||
router = APIRouter(tags=["Chat"])
|
router = APIRouter(tags=["Chat"])
|
||||||
|
|
||||||
flow_data_store: LRUCache = LRUCache(maxsize=10)
|
|
||||||
|
|
||||||
|
|
||||||
@router.websocket("/chat/{client_id}")
|
@router.websocket("/chat/{client_id}")
|
||||||
async def chat(client_id: str, websocket: WebSocket):
|
async def chat(client_id: str, websocket: WebSocket):
|
||||||
"""Websocket endpoint for chat."""
|
"""Websocket endpoint for chat."""
|
||||||
try:
|
try:
|
||||||
chat_manager = service_manager.get(ServiceType.CHAT_MANAGER)
|
chat_manager = service_manager.get(ServiceType.CHAT_MANAGER)
|
||||||
if client_id in chat_manager.in_memory_cache:
|
if client_id in chat_manager.cache_manager:
|
||||||
await chat_manager.handle_websocket(client_id, websocket)
|
await chat_manager.handle_websocket(client_id, websocket)
|
||||||
else:
|
else:
|
||||||
# We accept the connection but close it immediately
|
# We accept the connection but close it immediately
|
||||||
|
|
@ -34,23 +31,23 @@ async def chat(client_id: str, websocket: WebSocket):
|
||||||
@router.post("/build/init/{flow_id}", response_model=InitResponse, status_code=201)
|
@router.post("/build/init/{flow_id}", response_model=InitResponse, status_code=201)
|
||||||
async def init_build(graph_data: dict, flow_id: str):
|
async def init_build(graph_data: dict, flow_id: str):
|
||||||
"""Initialize the build by storing graph data and returning a unique session ID."""
|
"""Initialize the build by storing graph data and returning a unique session ID."""
|
||||||
|
flow_data_store = service_manager.get(ServiceType.CACHE_MANAGER)
|
||||||
try:
|
try:
|
||||||
if flow_id is None:
|
if flow_id is None:
|
||||||
raise ValueError("No ID provided")
|
raise ValueError("No ID provided")
|
||||||
# Check if already building
|
# Check if already building
|
||||||
if (
|
if (
|
||||||
flow_id in flow_data_store
|
flow_id in flow_data_store
|
||||||
and flow_data_store[flow_id]["status"] == BuildStatus.IN_PROGRESS
|
and isinstance(flow_data_store[flow_id], dict)
|
||||||
|
and flow_data_store[flow_id].get("status") == BuildStatus.IN_PROGRESS
|
||||||
):
|
):
|
||||||
return InitResponse(flowId=flow_id)
|
return InitResponse(flowId=flow_id)
|
||||||
|
|
||||||
# Delete from cache if already exists
|
# Delete from cache if already exists
|
||||||
chat_manager = service_manager.get(ServiceType.CHAT_MANAGER)
|
chat_manager = service_manager.get(ServiceType.CHAT_MANAGER)
|
||||||
if flow_id in chat_manager.in_memory_cache:
|
if flow_id in chat_manager.cache_manager:
|
||||||
with chat_manager.in_memory_cache._lock:
|
chat_manager.cache_manager.delete(flow_id)
|
||||||
chat_manager.in_memory_cache.delete(flow_id)
|
logger.debug(f"Deleted flow {flow_id} from cache")
|
||||||
logger.debug(f"Deleted flow {flow_id} from cache")
|
|
||||||
flow_data_store[flow_id] = {
|
flow_data_store[flow_id] = {
|
||||||
"graph_data": graph_data,
|
"graph_data": graph_data,
|
||||||
"status": BuildStatus.STARTED,
|
"status": BuildStatus.STARTED,
|
||||||
|
|
@ -65,6 +62,7 @@ async def init_build(graph_data: dict, flow_id: str):
|
||||||
@router.get("/build/{flow_id}/status", response_model=BuiltResponse)
|
@router.get("/build/{flow_id}/status", response_model=BuiltResponse)
|
||||||
async def build_status(flow_id: str):
|
async def build_status(flow_id: str):
|
||||||
"""Check the flow_id is in the flow_data_store."""
|
"""Check the flow_id is in the flow_data_store."""
|
||||||
|
flow_data_store = service_manager.get(ServiceType.CACHE_MANAGER)
|
||||||
try:
|
try:
|
||||||
built = (
|
built = (
|
||||||
flow_id in flow_data_store
|
flow_id in flow_data_store
|
||||||
|
|
@ -83,6 +81,7 @@ async def build_status(flow_id: str):
|
||||||
@router.get("/build/stream/{flow_id}", response_class=StreamingResponse)
|
@router.get("/build/stream/{flow_id}", response_class=StreamingResponse)
|
||||||
async def stream_build(flow_id: str):
|
async def stream_build(flow_id: str):
|
||||||
"""Stream the build process based on stored flow data."""
|
"""Stream the build process based on stored flow data."""
|
||||||
|
flow_data_store = service_manager.get(ServiceType.CACHE_MANAGER)
|
||||||
|
|
||||||
async def event_stream(flow_id):
|
async def event_stream(flow_id):
|
||||||
final_response = {"end_of_stream": True}
|
final_response = {"end_of_stream": True}
|
||||||
|
|
@ -163,7 +162,7 @@ async def stream_build(flow_id: str):
|
||||||
}
|
}
|
||||||
yield str(StreamData(event="message", data=input_keys_response))
|
yield str(StreamData(event="message", data=input_keys_response))
|
||||||
chat_manager = service_manager.get(ServiceType.CHAT_MANAGER)
|
chat_manager = service_manager.get(ServiceType.CHAT_MANAGER)
|
||||||
chat_manager.set_cache(flow_id, langchain_object)
|
chat_manager.set_cache(f"{flow_id}_chat", langchain_object)
|
||||||
# We need to reset the chat history
|
# We need to reset the chat history
|
||||||
chat_manager.chat_history.empty_history(flow_id)
|
chat_manager.chat_history.empty_history(flow_id)
|
||||||
flow_data_store[flow_id]["status"] = BuildStatus.SUCCESS
|
flow_data_store[flow_id]["status"] = BuildStatus.SUCCESS
|
||||||
|
|
|
||||||
217
src/backend/langflow/chat/manager.py
Normal file
217
src/backend/langflow/chat/manager.py
Normal file
|
|
@ -0,0 +1,217 @@
|
||||||
|
from collections import defaultdict
|
||||||
|
from fastapi import WebSocket, status
|
||||||
|
from langflow.api.v1.schemas import ChatMessage, ChatResponse, FileResponse
|
||||||
|
from langflow.cache import cache_manager
|
||||||
|
from langflow.cache.manager import Subject
|
||||||
|
from langflow.chat.utils import process_graph
|
||||||
|
from langflow.interface.utils import pil_to_base64
|
||||||
|
from langflow.utils.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
from typing import Any, Dict, List
|
||||||
|
|
||||||
|
from langflow.cache.flow import InMemoryCache
|
||||||
|
|
||||||
|
|
||||||
|
class ChatHistory(Subject):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.history: Dict[str, List[ChatMessage]] = defaultdict(list)
|
||||||
|
|
||||||
|
def add_message(self, client_id: str, message: ChatMessage):
|
||||||
|
"""Add a message to the chat history."""
|
||||||
|
|
||||||
|
self.history[client_id].append(message)
|
||||||
|
|
||||||
|
if not isinstance(message, FileResponse):
|
||||||
|
self.notify()
|
||||||
|
|
||||||
|
def get_history(self, client_id: str, filter_messages=True) -> List[ChatMessage]:
|
||||||
|
"""Get the chat history for a client."""
|
||||||
|
if history := self.history.get(client_id, []):
|
||||||
|
if filter_messages:
|
||||||
|
return [msg for msg in history if msg.type not in ["start", "stream"]]
|
||||||
|
return history
|
||||||
|
else:
|
||||||
|
return []
|
||||||
|
|
||||||
|
def empty_history(self, client_id: str):
|
||||||
|
"""Empty the chat history for a client."""
|
||||||
|
self.history[client_id] = []
|
||||||
|
|
||||||
|
|
||||||
|
class ChatManager:
|
||||||
|
def __init__(self):
|
||||||
|
self.active_connections: Dict[str, WebSocket] = {}
|
||||||
|
self.chat_history = ChatHistory()
|
||||||
|
self.cache_manager = cache_manager
|
||||||
|
self.cache_manager.attach(self.update)
|
||||||
|
self.in_memory_cache = InMemoryCache()
|
||||||
|
|
||||||
|
def on_chat_history_update(self):
|
||||||
|
"""Send the last chat message to the client."""
|
||||||
|
client_id = self.cache_manager.current_client_id
|
||||||
|
if client_id in self.active_connections:
|
||||||
|
chat_response = self.chat_history.get_history(
|
||||||
|
client_id, filter_messages=False
|
||||||
|
)[-1]
|
||||||
|
if chat_response.is_bot:
|
||||||
|
# Process FileResponse
|
||||||
|
if isinstance(chat_response, FileResponse):
|
||||||
|
# If data_type is pandas, convert to csv
|
||||||
|
if chat_response.data_type == "pandas":
|
||||||
|
chat_response.data = chat_response.data.to_csv()
|
||||||
|
elif chat_response.data_type == "image":
|
||||||
|
# Base64 encode the image
|
||||||
|
chat_response.data = pil_to_base64(chat_response.data)
|
||||||
|
# get event loop
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
|
||||||
|
coroutine = self.send_json(client_id, chat_response)
|
||||||
|
asyncio.run_coroutine_threadsafe(coroutine, loop)
|
||||||
|
|
||||||
|
def update(self):
|
||||||
|
if self.cache_manager.current_client_id in self.active_connections:
|
||||||
|
self.last_cached_object_dict = self.cache_manager.get_last()
|
||||||
|
# Add a new ChatResponse with the data
|
||||||
|
chat_response = FileResponse(
|
||||||
|
message=None,
|
||||||
|
type="file",
|
||||||
|
data=self.last_cached_object_dict["obj"],
|
||||||
|
data_type=self.last_cached_object_dict["type"],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.chat_history.add_message(
|
||||||
|
self.cache_manager.current_client_id, chat_response
|
||||||
|
)
|
||||||
|
|
||||||
|
async def connect(self, client_id: str, websocket: WebSocket):
|
||||||
|
await websocket.accept()
|
||||||
|
self.active_connections[client_id] = websocket
|
||||||
|
|
||||||
|
def disconnect(self, client_id: str):
|
||||||
|
self.active_connections.pop(client_id, None)
|
||||||
|
|
||||||
|
async def send_message(self, client_id: str, message: str):
|
||||||
|
websocket = self.active_connections[client_id]
|
||||||
|
await websocket.send_text(message)
|
||||||
|
|
||||||
|
async def send_json(self, client_id: str, message: ChatMessage):
|
||||||
|
websocket = self.active_connections[client_id]
|
||||||
|
await websocket.send_json(message.dict())
|
||||||
|
|
||||||
|
async def close_connection(self, client_id: str, code: int, reason: str):
|
||||||
|
if websocket := self.active_connections[client_id]:
|
||||||
|
try:
|
||||||
|
await websocket.close(code=code, reason=reason)
|
||||||
|
self.disconnect(client_id)
|
||||||
|
except RuntimeError as exc:
|
||||||
|
# This is to catch the following error:
|
||||||
|
# Unexpected ASGI message 'websocket.close', after sending 'websocket.close'
|
||||||
|
if "after sending" in str(exc):
|
||||||
|
logger.error(f"Error closing connection: {exc}")
|
||||||
|
|
||||||
|
async def process_message(
|
||||||
|
self, client_id: str, payload: Dict, langchain_object: Any
|
||||||
|
):
|
||||||
|
# Process the graph data and chat message
|
||||||
|
chat_inputs = payload.pop("inputs", {})
|
||||||
|
chat_inputs = ChatMessage(message=chat_inputs)
|
||||||
|
self.chat_history.add_message(client_id, chat_inputs)
|
||||||
|
|
||||||
|
# graph_data = payload
|
||||||
|
start_resp = ChatResponse(message=None, type="start", intermediate_steps="")
|
||||||
|
await self.send_json(client_id, start_resp)
|
||||||
|
|
||||||
|
# is_first_message = len(self.chat_history.get_history(client_id=client_id)) <= 1
|
||||||
|
# Generate result and thought
|
||||||
|
try:
|
||||||
|
logger.debug("Generating result and thought")
|
||||||
|
|
||||||
|
result, intermediate_steps = await process_graph(
|
||||||
|
langchain_object=langchain_object,
|
||||||
|
chat_inputs=chat_inputs,
|
||||||
|
websocket=self.active_connections[client_id],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
# Log stack trace
|
||||||
|
logger.exception(e)
|
||||||
|
self.chat_history.empty_history(client_id)
|
||||||
|
raise e
|
||||||
|
# Send a response back to the frontend, if needed
|
||||||
|
intermediate_steps = intermediate_steps or ""
|
||||||
|
history = self.chat_history.get_history(client_id, filter_messages=False)
|
||||||
|
file_responses = []
|
||||||
|
if history:
|
||||||
|
# Iterate backwards through the history
|
||||||
|
for msg in reversed(history):
|
||||||
|
if isinstance(msg, FileResponse):
|
||||||
|
if msg.data_type == "image":
|
||||||
|
# Base64 encode the image
|
||||||
|
if isinstance(msg.data, str):
|
||||||
|
continue
|
||||||
|
msg.data = pil_to_base64(msg.data)
|
||||||
|
file_responses.append(msg)
|
||||||
|
if msg.type == "start":
|
||||||
|
break
|
||||||
|
|
||||||
|
response = ChatResponse(
|
||||||
|
message=result,
|
||||||
|
intermediate_steps=intermediate_steps.strip(),
|
||||||
|
type="end",
|
||||||
|
files=file_responses,
|
||||||
|
)
|
||||||
|
await self.send_json(client_id, response)
|
||||||
|
self.chat_history.add_message(client_id, response)
|
||||||
|
|
||||||
|
def set_cache(self, client_id: str, langchain_object: Any) -> bool:
|
||||||
|
"""
|
||||||
|
Set the cache for a client.
|
||||||
|
"""
|
||||||
|
client_id = f"{client_id}_chat"
|
||||||
|
self.cache_manager.set(client_id, langchain_object)
|
||||||
|
return client_id in self.cache_manager
|
||||||
|
|
||||||
|
async def handle_websocket(self, client_id: str, websocket: WebSocket):
|
||||||
|
await self.connect(client_id, websocket)
|
||||||
|
|
||||||
|
try:
|
||||||
|
chat_history = self.chat_history.get_history(client_id)
|
||||||
|
# iterate and make BaseModel into dict
|
||||||
|
chat_history = [chat.dict() for chat in chat_history]
|
||||||
|
await websocket.send_json(chat_history)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
json_payload = await websocket.receive_json()
|
||||||
|
try:
|
||||||
|
payload = json.loads(json_payload)
|
||||||
|
except TypeError:
|
||||||
|
payload = json_payload
|
||||||
|
if "clear_history" in payload:
|
||||||
|
self.chat_history.history[client_id] = []
|
||||||
|
continue
|
||||||
|
|
||||||
|
with self.cache_manager.set_client_id(client_id):
|
||||||
|
langchain_object = self.in_memory_cache.get(f"{client_id}_chat")
|
||||||
|
await self.process_message(client_id, payload, langchain_object)
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
# Handle any exceptions that might occur
|
||||||
|
logger.error(f"Error handling websocket: {exc}")
|
||||||
|
await self.close_connection(
|
||||||
|
client_id=client_id,
|
||||||
|
code=status.WS_1011_INTERNAL_ERROR,
|
||||||
|
reason=str(exc)[:120],
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
await self.close_connection(
|
||||||
|
client_id=client_id,
|
||||||
|
code=status.WS_1000_NORMAL_CLOSURE,
|
||||||
|
reason="Client disconnected",
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Error closing connection: {exc}")
|
||||||
|
self.disconnect(client_id)
|
||||||
Loading…
Add table
Add a link
Reference in a new issue