Files
llama_deploy/examples/python_fullstack/workflows/agent_workflow.py
T
2024-09-12 17:05:56 -06:00

83 lines
2.8 KiB
Python

from logging import getLogger
from typing import List
from llama_index.core.llms import ChatMessage
from llama_index.core.memory import ChatMemoryBuffer
from llama_index.core.workflow import Event, Workflow, StartEvent, StopEvent, step
from llama_index.llms.openai import OpenAI
from llama_index.core.tools import FunctionTool
from rag_workflow import RAGWorkflow
logger = getLogger(__name__)
class ChatEvent(Event):
chat_history: List[ChatMessage]
class AgenticWorkflow(Workflow):
llm: OpenAI = OpenAI(model="gpt-4o")
@step
def prepare_chat_history(self, ev: StartEvent) -> ChatEvent:
logger.info(f"Preparing chat history: {ev}")
chat_history_dicts = ev.get("chat_history_dicts", [])
chat_history = [
ChatMessage(**chat_history_dict) for chat_history_dict in chat_history_dicts
]
newest_msg = ev.get("user_input")
if not newest_msg:
raise ValueError("No `user_input` input provided!")
chat_history.append(ChatMessage(role="user", content=newest_msg))
memory = ChatMemoryBuffer.from_defaults(
chat_history=chat_history,
llm=OpenAI(model="gpt-4o-mini"),
)
processed_chat_history = memory.get()
return ChatEvent(chat_history=processed_chat_history)
@step
async def chat(self, ev: ChatEvent, rag_workflow: RAGWorkflow) -> StopEvent:
chat_history = ev.chat_history
async def run_query(query: str) -> str:
"""
Useful for running a natural language query against a general knowledge base containing the paper "Attention is all your need".
If the user asks anything about a paper, use this tool to query it.
The query input should be a senetence or question related to what the user wants.
"""
response = await rag_workflow.run(query=query)
return str(response)
async def return_response(response: str) -> str:
"""Useful for returning a direct response to the user."""
return response
query_tool = FunctionTool.from_defaults(async_fn=run_query)
response_tool = FunctionTool.from_defaults(async_fn=return_response)
# responds using the tool or the LLM directly
response = await self.llm.apredict_and_call(
[query_tool, response_tool],
chat_history=chat_history,
error_on_no_tool_call=False,
)
logger.info(f"Response: {response.response}")
return StopEvent(result=response.response)
def build_agentic_workflow(rag_workflow: RAGWorkflow) -> AgenticWorkflow:
agentic_workflow = AgenticWorkflow(timeout=120.0)
# add the rag workflow as a subworkflow
agentic_workflow.add_workflows(rag_workflow=rag_workflow)
return agentic_workflow