Message history

This commit is contained in:
Yiorgis Gozadinos 2026-01-07 14:14:33 +02:00
parent 9f1d82daa8
commit 5a1cb253e5
No known key found for this signature in database

View file

@ -6,6 +6,13 @@ from agent import ChatDeps, ChatSessionState, create_chat_agent
from anyio import create_memory_object_stream, create_task_group
from anyio.streams.memory import MemoryObjectSendStream
from dotenv import load_dotenv
from pydantic_ai.messages import (
ModelMessage,
ModelRequest,
ModelResponse,
TextPart,
UserPromptPart,
)
from starlette.applications import Starlette
from starlette.middleware import Middleware
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.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()
logging.basicConfig(
@ -65,8 +96,10 @@ async def stream_chat(request: Request) -> StreamingResponse:
input_data = RunAgentInput(**body)
user_message = ""
message_history: list[ModelMessage] = []
if input_data.messages:
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]()
@ -113,7 +146,9 @@ async def stream_chat(request: Request) -> StreamingResponse:
async with create_task_group() as tg:
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.finish_run(result.output)
await emitter.close()