🔧 fix(callback.py): add error handling when sending response to websocket to prevent potential errors

🔧 fix(callback.py): change intermediate_steps assignment to use formatted string for better readability and maintainability
✨ feat(callback.py): add observation_prefix parameter to on_tool_end method to allow customization of the observation prefix in the response message
✨ feat(callback.py): add logger to handle potential errors when sending response to websocket
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-07-08 16:49:30 -03:00
commit 2d4642fc36

View file

@ -4,12 +4,14 @@ from langchain.callbacks.base import AsyncCallbackHandler, BaseCallbackHandler
from langflow.api.v1.schemas import ChatResponse from langflow.api.v1.schemas import ChatResponse
from typing import Any, Dict
from typing import Any, Dict, List, Union from typing import Any, Dict, List, Union
from fastapi import WebSocket from fastapi import WebSocket
from langchain.schema import AgentAction, LLMResult, AgentFinish from langchain.schema import AgentAction, LLMResult, AgentFinish
from langflow.utils.logger import logger
# https://github.com/hwchase17/chat-langchain/blob/master/callback.py # https://github.com/hwchase17/chat-langchain/blob/master/callback.py
@ -62,12 +64,23 @@ class AsyncStreamingLLMCallbackHandler(AsyncCallbackHandler):
async def on_tool_end(self, output: str, **kwargs: Any) -> Any: async def on_tool_end(self, output: str, **kwargs: Any) -> Any:
"""Run when tool ends running.""" """Run when tool ends running."""
observation_prefix = kwargs.get("observation_prefix", "Tool output: ")
# Create a formatted message.
intermediate_steps = f"{observation_prefix}{output}"
# Create a ChatResponse instance.
resp = ChatResponse( resp = ChatResponse(
message="", message="",
type="stream", type="stream",
intermediate_steps=f"Tool output: {output}", intermediate_steps=intermediate_steps,
) )
await self.websocket.send_json(resp.dict())
# Try to send the response, handle potential errors.
try:
await self.websocket.send_json(resp.dict())
except Exception as e:
logger.error(e)
async def on_tool_error( async def on_tool_error(
self, error: Union[Exception, KeyboardInterrupt], **kwargs: Any self, error: Union[Exception, KeyboardInterrupt], **kwargs: Any