From 55e9b4ba1c71edb9e5a78686ed0456b55aad7d43 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Thu, 21 Dec 2023 14:52:12 -0300 Subject: [PATCH] Add AIMessage support and update Result model --- src/backend/langflow/processing/process.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/src/backend/langflow/processing/process.py b/src/backend/langflow/processing/process.py index 1fb4b322e..659141446 100644 --- a/src/backend/langflow/processing/process.py +++ b/src/backend/langflow/processing/process.py @@ -5,12 +5,14 @@ from langchain.agents import AgentExecutor from langchain.chains.base import Chain from langchain.schema import AgentAction, Document from langchain.vectorstores.base import VectorStore +from langchain_core.messages import AIMessage from langchain_core.runnables.base import Runnable +from loguru import logger +from pydantic import BaseModel + from langflow.interface.custom.custom_component import CustomComponent from langflow.interface.run import build_sorted_vertices, get_memory_key, update_memory_keys from langflow.services.deps import get_session_service -from loguru import logger -from pydantic import BaseModel def fix_memory_inputs(langchain_object): @@ -125,6 +127,11 @@ async def process_runnable(runnable: Runnable, inputs: Union[dict, List[dict]]): result = await runnable.ainvoke(inputs) else: raise ValueError(f"Runnable {runnable} does not support inputs of type {type(inputs)}") + # Check if the result is a list of AIMessages + if isinstance(result, list) and all(isinstance(r, AIMessage) for r in result): + result = [r.content for r in result] + elif isinstance(result, AIMessage): + result = result.content return result @@ -181,7 +188,7 @@ async def generate_result(built_object: Union[Chain, VectorStore, Runnable], inp class Result(BaseModel): - result: Any + result: Union[dict, List[dict], str, List[str], AIMessage] session_id: str