Add default in-memory history provider for workflow agents (#3918)

This commit is contained in:
Eduard van Valkenburg
2026-02-13 14:55:39 +01:00
committed by GitHub
Unverified
parent e9b3a5bbc7
commit f3ea872156
2 changed files with 48 additions and 2 deletions
@@ -12,7 +12,13 @@ from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast, overload
from .._agents import BaseAgent
from .._sessions import AgentSession, BaseContextProvider, BaseHistoryProvider, SessionContext
from .._sessions import (
AgentSession,
BaseContextProvider,
BaseHistoryProvider,
InMemoryHistoryProvider,
SessionContext,
)
from .._types import (
AgentResponse,
AgentResponseUpdate,
@@ -112,7 +118,17 @@ class WorkflowAgent(BaseAgent):
if not any(is_type_compatible(list[Message], input_type) for input_type in start_executor.input_types):
raise ValueError("Workflow's start executor cannot handle list[Message]")
super().__init__(id=id, name=name, description=description, context_providers=context_providers, **kwargs)
resolved_context_providers = list(context_providers) if context_providers is not None else []
if not resolved_context_providers:
resolved_context_providers.append(InMemoryHistoryProvider("memory"))
super().__init__(
id=id,
name=name,
description=description,
context_providers=resolved_context_providers,
**kwargs,
)
self._workflow: Workflow = workflow
self._pending_requests: dict[str, WorkflowEvent[Any]] = {}
@@ -14,6 +14,7 @@ from agent_framework import (
AgentSession,
Content,
Executor,
InMemoryHistoryProvider,
Message,
ResponseStream,
SupportsAgentRun,
@@ -562,6 +563,35 @@ class TestWorkflowAgent:
assert len(capturing_executor.received_messages) == 1
assert capturing_executor.received_messages[0].text == "Just a new message"
async def test_workflow_as_agent_adds_default_history_provider(self) -> None:
"""Test that workflow.as_agent() defaults to in-memory history when no providers are configured."""
capturing_executor = ConversationHistoryCapturingExecutor(id="default_history_provider_test")
workflow = WorkflowBuilder(start_executor=capturing_executor).build()
agent = workflow.as_agent(name="Default History Provider Agent")
session = AgentSession()
await agent.run("first message", session=session)
await agent.run("second message", session=session)
assert any(isinstance(provider, InMemoryHistoryProvider) for provider in agent.context_providers)
texts = [message.text for message in capturing_executor.received_messages]
assert "first message" in texts
assert "second message" in texts
async def test_workflow_agent_keeps_explicit_context_providers(self) -> None:
"""Test that WorkflowAgent does not append defaults when context providers are explicitly provided."""
workflow = WorkflowBuilder(
start_executor=ConversationHistoryCapturingExecutor(id="explicit_provider_test")
).build()
explicit_provider = InMemoryHistoryProvider("custom-memory")
agent = WorkflowAgent(
workflow=workflow,
name="Explicit Provider Agent",
context_providers=[explicit_provider],
)
assert agent.context_providers == [explicit_provider]
async def test_checkpoint_storage_passed_to_workflow(self) -> None:
"""Test that checkpoint_storage parameter is passed through to the workflow."""
from agent_framework import InMemoryCheckpointStorage