Message history
This commit is contained in:
parent
9f1d82daa8
commit
5a1cb253e5
1 changed files with 36 additions and 1 deletions
|
|
@ -6,6 +6,13 @@ from agent import ChatDeps, ChatSessionState, create_chat_agent
|
||||||
from anyio import create_memory_object_stream, create_task_group
|
from anyio import create_memory_object_stream, create_task_group
|
||||||
from anyio.streams.memory import MemoryObjectSendStream
|
from anyio.streams.memory import MemoryObjectSendStream
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
from pydantic_ai.messages import (
|
||||||
|
ModelMessage,
|
||||||
|
ModelRequest,
|
||||||
|
ModelResponse,
|
||||||
|
TextPart,
|
||||||
|
UserPromptPart,
|
||||||
|
)
|
||||||
from starlette.applications import Starlette
|
from starlette.applications import Starlette
|
||||||
from starlette.middleware import Middleware
|
from starlette.middleware import Middleware
|
||||||
from starlette.middleware.cors import CORSMiddleware
|
from starlette.middleware.cors import CORSMiddleware
|
||||||
|
|
@ -19,6 +26,30 @@ from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.graph.agui.emitter import AGUIEmitter
|
from haiku.rag.graph.agui.emitter import AGUIEmitter
|
||||||
from haiku.rag.graph.agui.server import RunAgentInput, format_sse_event
|
from haiku.rag.graph.agui.server import RunAgentInput, format_sse_event
|
||||||
|
|
||||||
|
|
||||||
|
def convert_messages_to_history(
|
||||||
|
messages: list[dict[str, str]],
|
||||||
|
) -> list[ModelMessage]:
|
||||||
|
"""Convert AG-UI/CopilotKit messages to pydantic-ai message history.
|
||||||
|
|
||||||
|
Skips the last message since it will be passed as user_prompt to agent.run().
|
||||||
|
"""
|
||||||
|
history: list[ModelMessage] = []
|
||||||
|
|
||||||
|
# Skip the last message - it will be the current user prompt
|
||||||
|
for msg in messages[:-1]:
|
||||||
|
role = msg.get("role", "")
|
||||||
|
content = msg.get("content", "")
|
||||||
|
|
||||||
|
if role == "user":
|
||||||
|
history.append(ModelRequest(parts=[UserPromptPart(content=content)]))
|
||||||
|
elif role == "assistant":
|
||||||
|
history.append(ModelResponse(parts=[TextPart(content=content)]))
|
||||||
|
# Skip other roles (system, tool, etc.) for now
|
||||||
|
|
||||||
|
return history
|
||||||
|
|
||||||
|
|
||||||
load_dotenv()
|
load_dotenv()
|
||||||
|
|
||||||
logging.basicConfig(
|
logging.basicConfig(
|
||||||
|
|
@ -65,8 +96,10 @@ async def stream_chat(request: Request) -> StreamingResponse:
|
||||||
input_data = RunAgentInput(**body)
|
input_data = RunAgentInput(**body)
|
||||||
|
|
||||||
user_message = ""
|
user_message = ""
|
||||||
|
message_history: list[ModelMessage] = []
|
||||||
if input_data.messages:
|
if input_data.messages:
|
||||||
user_message = input_data.messages[-1].get("content", "")
|
user_message = input_data.messages[-1].get("content", "")
|
||||||
|
message_history = convert_messages_to_history(input_data.messages)
|
||||||
|
|
||||||
send_stream, receive_stream = create_memory_object_stream[str]()
|
send_stream, receive_stream = create_memory_object_stream[str]()
|
||||||
|
|
||||||
|
|
@ -113,7 +146,9 @@ async def stream_chat(request: Request) -> StreamingResponse:
|
||||||
async with create_task_group() as tg:
|
async with create_task_group() as tg:
|
||||||
tg.start_soon(forward_events)
|
tg.start_soon(forward_events)
|
||||||
|
|
||||||
result = await chat_agent.run(user_message, deps=deps)
|
result = await chat_agent.run(
|
||||||
|
user_message, deps=deps, message_history=message_history
|
||||||
|
)
|
||||||
emitter.log(result.output)
|
emitter.log(result.output)
|
||||||
emitter.finish_run(result.output)
|
emitter.finish_run(result.output)
|
||||||
await emitter.close()
|
await emitter.close()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue