mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Thread storage and serialization (#394)
* Updates for message store support * Added unit tests * Added suspend-resume example * Added example with custom chat message store * Small fix * Addressed PR feedback * Renaming and documentation * More renaming * Addressed more PR feedback * Small fixes in Foundry chat client and examples * Small update * Addressed PR feedback * Increased timeout for Azure tests
This commit is contained in:
committed by
GitHub
Unverified
parent
95cb20ca40
commit
dea736e550
@@ -12,5 +12,6 @@ from ._agents import * # noqa: F403
|
||||
from ._clients import * # noqa: F403
|
||||
from ._logging import * # noqa: F403
|
||||
from ._mcp import * # noqa: F403
|
||||
from ._threads import * # noqa: F403
|
||||
from ._tools import * # noqa: F403
|
||||
from ._types import * # noqa: F403
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
import sys
|
||||
from collections.abc import AsyncIterable, Callable, MutableMapping, Sequence
|
||||
from contextlib import AbstractAsyncContextManager, AsyncExitStack
|
||||
from enum import Enum
|
||||
from itertools import chain
|
||||
from typing import Any, ClassVar, Literal, Protocol, TypeVar, runtime_checkable
|
||||
from uuid import uuid4
|
||||
@@ -13,6 +12,7 @@ from pydantic import BaseModel, Field, PrivateAttr
|
||||
from ._clients import ChatClient
|
||||
from ._mcp import McpTool
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._threads import AgentThread, ChatMessageStore, deserialize_thread_state, thread_on_new_messages
|
||||
from ._tools import AITool
|
||||
from ._types import (
|
||||
AgentRunResponse,
|
||||
@@ -34,47 +34,7 @@ else:
|
||||
|
||||
TThreadType = TypeVar("TThreadType", bound="AgentThread")
|
||||
|
||||
# region AgentThread
|
||||
|
||||
__all__ = [
|
||||
"AIAgent",
|
||||
"AgentBase",
|
||||
"AgentThread",
|
||||
"ChatClientAgent",
|
||||
"ChatClientAgentThread",
|
||||
"ChatClientAgentThreadType",
|
||||
"MessagesRetrievableThread",
|
||||
]
|
||||
|
||||
|
||||
class AgentThread(AFBaseModel):
|
||||
"""Base class for agent threads."""
|
||||
|
||||
id: str | None = None
|
||||
|
||||
async def on_new_messages(
|
||||
self,
|
||||
new_messages: ChatMessage | Sequence[ChatMessage],
|
||||
) -> None:
|
||||
"""Invoked when a new message has been contributed to the chat by any participant."""
|
||||
await self._on_new_messages(new_messages=new_messages)
|
||||
|
||||
async def _on_new_messages(
|
||||
self,
|
||||
new_messages: ChatMessage | Sequence[ChatMessage],
|
||||
) -> None:
|
||||
"""Invoked when a new message has been contributed to the chat by any participant."""
|
||||
pass
|
||||
|
||||
|
||||
# region MessagesRetrievableThread
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class MessagesRetrievableThread(Protocol):
|
||||
def get_messages(self) -> AsyncIterable[ChatMessage]:
|
||||
"""Asynchronously retrieves all messages from thread."""
|
||||
...
|
||||
__all__ = ["AIAgent", "AgentBase", "ChatClientAgent"]
|
||||
|
||||
|
||||
# region Agent Protocol
|
||||
@@ -186,7 +146,7 @@ class AgentBase(AFBaseModel):
|
||||
) -> None:
|
||||
"""Notify the thread of new messages."""
|
||||
if isinstance(new_messages, ChatMessage) or len(new_messages) > 0:
|
||||
await thread.on_new_messages(new_messages)
|
||||
await thread_on_new_messages(thread, new_messages)
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
@@ -196,116 +156,17 @@ class AgentBase(AFBaseModel):
|
||||
"""
|
||||
return self.name or self.id
|
||||
|
||||
def _validate_or_create_thread_type(
|
||||
self,
|
||||
thread: AgentThread | None,
|
||||
construct_thread: Callable[[], TThreadType],
|
||||
expected_type: type[TThreadType],
|
||||
) -> TThreadType:
|
||||
"""Validate or create a AgentThread of the right type.
|
||||
|
||||
Args:
|
||||
thread: The thread to validate or create.
|
||||
construct_thread: A callable that constructs a new thread if `thread` is None.
|
||||
expected_type: The expected type of the thread.
|
||||
|
||||
Returns:
|
||||
The validated or newly created thread of the expected type.
|
||||
|
||||
Raises:
|
||||
AgentExecutionException: If the thread is not of the expected type.
|
||||
"""
|
||||
if thread is None:
|
||||
return construct_thread()
|
||||
|
||||
if not isinstance(thread, expected_type):
|
||||
raise AgentExecutionException(
|
||||
f"{self.__class__.__name__} currently only supports agent threads of type {expected_type.__name__}."
|
||||
)
|
||||
def get_new_thread(self) -> AgentThread:
|
||||
"""Returns AgentThread instance that is compatible with the agent."""
|
||||
return AgentThread()
|
||||
|
||||
async def deserialize_thread(self, serialized_thread: Any, **kwargs: Any) -> AgentThread:
|
||||
"""Deserializes the thread."""
|
||||
thread: AgentThread = self.get_new_thread()
|
||||
await deserialize_thread_state(thread, serialized_thread, **kwargs)
|
||||
return thread
|
||||
|
||||
|
||||
# region ChatClientAgentThread
|
||||
|
||||
|
||||
class ChatClientAgentThreadType(Enum):
|
||||
"""Defines the different supported storage locations for ChatClientAgentThread."""
|
||||
|
||||
IN_MEMORY_MESSAGES = "InMemoryMessages"
|
||||
"""Messages are stored in memory inside the thread object."""
|
||||
|
||||
CONVERSATION_ID = "ConversationId"
|
||||
"""Messages are stored in the service and the thread object just has an id reference to the service storage."""
|
||||
|
||||
|
||||
class ChatClientAgentThread(AgentThread):
|
||||
"""Chat client agent thread.
|
||||
|
||||
This class manages chat threads either locally (in-memory) or via a service based on initialization.
|
||||
"""
|
||||
|
||||
chat_messages: list[ChatMessage] | None = None
|
||||
storage_location: ChatClientAgentThreadType | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
id: str | None = None,
|
||||
messages: Sequence[ChatMessage] | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Initialize the chat client agent thread.
|
||||
|
||||
Args:
|
||||
id: Service thread identifier. If provided, thread is managed by the service and messages are
|
||||
not stored locally. Must not be empty or whitespace.
|
||||
messages: Initial messages for local storage. If provided, thread is managed
|
||||
locally in-memory.
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Raises:
|
||||
ValueError: If both id and messages are provided, or if id is empty/whitespace.
|
||||
|
||||
Notes:
|
||||
- If id is set, _id is assigned and _chat_messages is None (service-managed).
|
||||
- If messages is set, _chat_messages is populated and _id is None (local).
|
||||
- If neither is provided, creates an empty local thread.
|
||||
"""
|
||||
processed_messages: list[ChatMessage] | None = None
|
||||
storage_location: ChatClientAgentThreadType | None = None
|
||||
|
||||
if id and messages:
|
||||
raise ValueError("Cannot specify both id and messages")
|
||||
|
||||
if id:
|
||||
if not id.strip():
|
||||
raise ValueError("ID cannot be empty or whitespace")
|
||||
storage_location = ChatClientAgentThreadType.CONVERSATION_ID
|
||||
elif messages:
|
||||
processed_messages = []
|
||||
processed_messages.extend(messages)
|
||||
storage_location = ChatClientAgentThreadType.IN_MEMORY_MESSAGES
|
||||
|
||||
super().__init__(
|
||||
id=id,
|
||||
chat_messages=processed_messages, # type: ignore[reportCallIssue]
|
||||
storage_location=storage_location, # type: ignore[reportCallIssue]
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def get_messages(self) -> AsyncIterable[ChatMessage]:
|
||||
"""Get all messages in the thread."""
|
||||
for message in self.chat_messages or []:
|
||||
yield message
|
||||
|
||||
async def _on_new_messages(self, new_messages: ChatMessage | Sequence[ChatMessage]) -> None:
|
||||
"""Handle new messages."""
|
||||
if self.storage_location == ChatClientAgentThreadType.IN_MEMORY_MESSAGES:
|
||||
if self.chat_messages is None:
|
||||
self.chat_messages = []
|
||||
self.chat_messages.extend([new_messages] if isinstance(new_messages, ChatMessage) else new_messages)
|
||||
|
||||
|
||||
# region ChatClientAgent
|
||||
|
||||
|
||||
@@ -317,6 +178,7 @@ class ChatClientAgent(AgentBase):
|
||||
chat_client: ChatClient
|
||||
instructions: str | None = None
|
||||
chat_options: ChatOptions
|
||||
chat_message_store_factory: Callable[[], ChatMessageStore] | None = None
|
||||
_local_mcp_tools: list[McpTool] = PrivateAttr(default_factory=list) # type: ignore[reportUnknownVariableType]
|
||||
_async_exit_stack: AsyncExitStack = PrivateAttr(default_factory=AsyncExitStack)
|
||||
|
||||
@@ -350,6 +212,7 @@ class ChatClientAgent(AgentBase):
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
chat_message_store_factory: Callable[[], ChatMessageStore] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Create a ChatClientAgent.
|
||||
@@ -382,6 +245,8 @@ class ChatClientAgent(AgentBase):
|
||||
top_p: the nucleus sampling probability to use.
|
||||
user: the user to associate with the request.
|
||||
additional_properties: additional properties to include in the request.
|
||||
chat_message_store_factory: factory function to create an instance of ChatMessageStore. If not provided,
|
||||
the default in-memory store will be used.
|
||||
kwargs: any additional keyword arguments.
|
||||
Unused, can be used by subclasses of this Agent.
|
||||
"""
|
||||
@@ -394,6 +259,7 @@ class ChatClientAgent(AgentBase):
|
||||
final_tools = [tool for tool in normalized_tools if not isinstance(tool, McpTool)]
|
||||
args: dict[str, Any] = {
|
||||
"chat_client": chat_client,
|
||||
"chat_message_store_factory": chat_message_store_factory,
|
||||
"chat_options": ChatOptions(
|
||||
ai_model_id=model,
|
||||
frequency_penalty=frequency_penalty,
|
||||
@@ -537,7 +403,7 @@ class ChatClientAgent(AgentBase):
|
||||
chat_options=self.chat_options
|
||||
& ChatOptions(
|
||||
ai_model_id=model,
|
||||
conversation_id=thread.id,
|
||||
conversation_id=thread.service_thread_id,
|
||||
frequency_penalty=frequency_penalty,
|
||||
logit_bias=logit_bias,
|
||||
max_tokens=max_tokens,
|
||||
@@ -660,7 +526,7 @@ class ChatClientAgent(AgentBase):
|
||||
messages=thread_messages,
|
||||
chat_options=self.chat_options
|
||||
& ChatOptions(
|
||||
conversation_id=thread.id,
|
||||
conversation_id=thread.service_thread_id,
|
||||
frequency_penalty=frequency_penalty,
|
||||
logit_bias=logit_bias,
|
||||
max_tokens=max_tokens,
|
||||
@@ -705,44 +571,50 @@ class ChatClientAgent(AgentBase):
|
||||
await self._notify_thread_of_new_messages(thread, input_messages)
|
||||
await self._notify_thread_of_new_messages(thread, response.messages)
|
||||
|
||||
def get_new_thread(self) -> ChatClientAgentThread:
|
||||
return ChatClientAgentThread()
|
||||
def get_new_thread(self) -> AgentThread:
|
||||
message_store: ChatMessageStore | None = None
|
||||
|
||||
if self.chat_message_store_factory:
|
||||
message_store = self.chat_message_store_factory()
|
||||
|
||||
return AgentThread() if message_store is None else AgentThread(message_store=message_store)
|
||||
|
||||
def _update_thread_with_type_and_conversation_id(
|
||||
self, chat_client_thread: ChatClientAgentThread, responseConversationId: str | None
|
||||
self, thread: AgentThread, response_conversation_id: str | None
|
||||
) -> None:
|
||||
"""Update thread with storage type and conversation ID.
|
||||
|
||||
Args:
|
||||
chat_client_thread: The thread to update.
|
||||
responseConversationId: The conversation ID from the response, if any.
|
||||
thread: The thread to update.
|
||||
response_conversation_id: The conversation ID from the response, if any.
|
||||
|
||||
Raises:
|
||||
AgentExecutionException: If conversation ID is missing for service-managed thread.
|
||||
"""
|
||||
# Set the thread's storage location, the first time that we use it.
|
||||
if chat_client_thread.storage_location is None:
|
||||
chat_client_thread.storage_location = (
|
||||
ChatClientAgentThreadType.CONVERSATION_ID
|
||||
if responseConversationId is not None
|
||||
else ChatClientAgentThreadType.IN_MEMORY_MESSAGES
|
||||
if response_conversation_id is None and thread.service_thread_id is not None:
|
||||
# We were passed a thread that is service managed, but we got no conversation id back from the chat client,
|
||||
# meaning the service doesn't support service managed threads,
|
||||
# so the thread cannot be used with this service.
|
||||
raise AgentExecutionException(
|
||||
"Service did not return a valid conversation id when using a service managed thread."
|
||||
)
|
||||
|
||||
# If we got a conversation id back from the chat client, it means that the service supports server side thread
|
||||
# storage so we should capture the id and update the thread with the new id.
|
||||
if chat_client_thread.storage_location == ChatClientAgentThreadType.CONVERSATION_ID:
|
||||
if responseConversationId is None:
|
||||
raise AgentExecutionException(
|
||||
"Service did not return a valid conversation id when using a service managed thread."
|
||||
)
|
||||
chat_client_thread.id = responseConversationId
|
||||
if response_conversation_id is not None:
|
||||
# If we got a conversation id back from the chat client, it means that the service
|
||||
# supports server side thread storage so we should update the thread with the new id.
|
||||
thread.service_thread_id = response_conversation_id
|
||||
elif thread.message_store is None and self.chat_message_store_factory is not None:
|
||||
# If the service doesn't use service side thread storage (i.e. we got no id back from invocation), and
|
||||
# the thread has no message_store yet, and we have a custom messages store, we should update the thread
|
||||
# with the custom message_store so that it has somewhere to store the chat history.
|
||||
thread.message_store = self.chat_message_store_factory()
|
||||
|
||||
async def _prepare_thread_and_messages(
|
||||
self,
|
||||
*,
|
||||
thread: AgentThread | None,
|
||||
input_messages: list[ChatMessage] | None = None,
|
||||
) -> tuple[ChatClientAgentThread, list[ChatMessage]]:
|
||||
) -> tuple[AgentThread, list[ChatMessage]]:
|
||||
"""Prepare the messages for agent execution.
|
||||
|
||||
Args:
|
||||
@@ -755,19 +627,14 @@ class ChatClientAgent(AgentBase):
|
||||
Raises:
|
||||
AgentExecutionException: If the thread is not of the expected type.
|
||||
"""
|
||||
validated_thread: ChatClientAgentThread = self._validate_or_create_thread_type( # type: ignore[reportAssignmentType]
|
||||
thread=thread,
|
||||
construct_thread=self.get_new_thread,
|
||||
expected_type=ChatClientAgentThread,
|
||||
)
|
||||
thread = thread or self.get_new_thread()
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
if self.instructions:
|
||||
messages.append(ChatMessage(role=ChatRole.SYSTEM, text=self.instructions))
|
||||
if isinstance(validated_thread, MessagesRetrievableThread):
|
||||
async for message in validated_thread.get_messages():
|
||||
messages.append(message)
|
||||
messages.extend(await thread.list_messages() or [])
|
||||
messages.extend(input_messages or [])
|
||||
return validated_thread, messages
|
||||
return thread, messages
|
||||
|
||||
def _normalize_messages(
|
||||
self,
|
||||
|
||||
@@ -10,6 +10,7 @@ from pydantic import BaseModel
|
||||
|
||||
from ._logging import get_logger
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._threads import ChatMessageStore
|
||||
from ._tools import AIFunction, AITool
|
||||
from ._types import (
|
||||
AIContents,
|
||||
@@ -645,6 +646,7 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
| MutableMapping[str, Any]
|
||||
| list[MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
chat_message_store_factory: Callable[[], ChatMessageStore] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> "ChatClientAgent":
|
||||
"""Create an agent with the given name and instructions.
|
||||
@@ -653,6 +655,8 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
name: The name of the agent.
|
||||
instructions: The instructions for the agent.
|
||||
tools: Optional list of tools to associate with the agent.
|
||||
chat_message_store_factory: Factory function to create an instance of ChatMessageStore. If not provided,
|
||||
the default in-memory store will be used.
|
||||
**kwargs: Additional keyword arguments to pass to the agent.
|
||||
See ChatClientAgent for all the available options.
|
||||
|
||||
@@ -661,7 +665,14 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
"""
|
||||
from ._agents import ChatClientAgent
|
||||
|
||||
return ChatClientAgent(chat_client=self, name=name, instructions=instructions, tools=tools, **kwargs)
|
||||
return ChatClientAgent(
|
||||
chat_client=self,
|
||||
name=name,
|
||||
instructions=instructions,
|
||||
tools=tools,
|
||||
chat_message_store_factory=chat_message_store_factory,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
# region Embedding Client
|
||||
|
||||
@@ -0,0 +1,375 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, Protocol, overload
|
||||
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._types import ChatMessage
|
||||
|
||||
__all__ = ["AgentThread", "ChatMessageList", "ChatMessageStore"]
|
||||
|
||||
|
||||
class ChatMessageStore(Protocol):
|
||||
"""Defines methods for storing and retrieving chat messages associated with a specific thread.
|
||||
|
||||
Implementations of this protocol are responsible for managing the storage of chat messages,
|
||||
including handling large volumes of data by truncating or summarizing messages as necessary.
|
||||
"""
|
||||
|
||||
async def list_messages(self) -> list[ChatMessage]:
|
||||
"""Gets all the messages from the store that should be used for the next agent invocation.
|
||||
|
||||
Messages are returned in ascending chronological order, with the oldest message first.
|
||||
|
||||
If the messages stored in the store become very large, it is up to the store to
|
||||
truncate, summarize or otherwise limit the number of messages returned.
|
||||
|
||||
When using implementations of ChatMessageStore, a new one should be created for each thread
|
||||
since they may contain state that is specific to a thread.
|
||||
"""
|
||||
...
|
||||
|
||||
async def add_messages(self, messages: Sequence[ChatMessage]) -> None:
|
||||
"""Adds messages to the store."""
|
||||
...
|
||||
|
||||
async def deserialize_state(self, serialized_store_state: Any, **kwargs: Any) -> None:
|
||||
"""Deserializes the state into the properties on this store.
|
||||
|
||||
This method, together with serialize_state can be used to save and load messages from a persistent store
|
||||
if this store only has messages in memory.
|
||||
"""
|
||||
...
|
||||
|
||||
async def serialize_state(self, **kwargs: Any) -> Any:
|
||||
"""Serializes the current object's state.
|
||||
|
||||
This method, together with deserialize_state can be used to save and load messages from a persistent store
|
||||
if this store only has messages in memory.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class AgentThread(AFBaseModel):
|
||||
"""Base class for agent threads."""
|
||||
|
||||
_service_thread_id: str | None = None
|
||||
_message_store: ChatMessageStore | None = None
|
||||
|
||||
@overload
|
||||
def __init__(self) -> None:
|
||||
"""Initialize an empty AgentThread with no service thread ID or message store."""
|
||||
...
|
||||
|
||||
@overload
|
||||
def __init__(self, service_thread_id: str) -> None:
|
||||
"""Initialize an AgentThread with a service thread ID.
|
||||
|
||||
Args:
|
||||
service_thread_id: The ID of the thread managed by the agent service.
|
||||
"""
|
||||
...
|
||||
|
||||
@overload
|
||||
def __init__(self, *, message_store: ChatMessageStore) -> None:
|
||||
"""Initialize an AgentThread with a custom message store.
|
||||
|
||||
Args:
|
||||
message_store: The ChatMessageStore implementation for managing chat messages.
|
||||
"""
|
||||
...
|
||||
|
||||
def __init__(self, service_thread_id: str | None = None, *, message_store: ChatMessageStore | None = None) -> None:
|
||||
"""Initialize an AgentThread.
|
||||
|
||||
Args:
|
||||
service_thread_id: Optional ID of the thread managed by the agent service.
|
||||
message_store: Optional ChatMessageStore implementation for managing chat messages.
|
||||
|
||||
Note:
|
||||
Either service_thread_id or message_store may be set, but not both.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.service_thread_id = service_thread_id
|
||||
self.message_store = message_store
|
||||
|
||||
@property
|
||||
def service_thread_id(self) -> str | None:
|
||||
"""Gets the ID of the current thread to support cases where the thread is owned by the agent service."""
|
||||
return self._service_thread_id
|
||||
|
||||
@service_thread_id.setter
|
||||
def service_thread_id(self, service_thread_id: str | None) -> None:
|
||||
"""Sets the ID of the current thread to support cases where the thread is owned by the agent service.
|
||||
|
||||
Note that either service_thread_id or message_store may be set, but not both.
|
||||
"""
|
||||
if not self._service_thread_id and not service_thread_id:
|
||||
return
|
||||
|
||||
if self._message_store is not None:
|
||||
raise ValueError(
|
||||
"Only the service_thread_id or message_store may be set, "
|
||||
"but not both and switching from one to another is not supported."
|
||||
)
|
||||
|
||||
self._service_thread_id = service_thread_id
|
||||
|
||||
@property
|
||||
def message_store(self) -> ChatMessageStore | None:
|
||||
"""Gets the ChatMessageStore used by this thread, when messages should be stored in a custom location."""
|
||||
return self._message_store
|
||||
|
||||
@message_store.setter
|
||||
def message_store(self, message_store: ChatMessageStore | None) -> None:
|
||||
"""Sets the ChatMessageStore used by this thread, when messages should be stored in a custom location.
|
||||
|
||||
Note that either service_thread_id or message_store may be set, but not both.
|
||||
"""
|
||||
if self._message_store is None and message_store is None:
|
||||
return
|
||||
|
||||
if self._service_thread_id:
|
||||
raise ValueError(
|
||||
"Only the service_thread_id or message_store may be set, "
|
||||
"but not both and switching from one to another is not supported."
|
||||
)
|
||||
|
||||
self._message_store = message_store
|
||||
|
||||
async def list_messages(self) -> list[ChatMessage] | None:
|
||||
"""Retrieves any messages stored in ChatMessageStore of the thread, otherwise returns an empty collection."""
|
||||
return await self._message_store.list_messages() if self._message_store is not None else None
|
||||
|
||||
async def serialize(self, **kwargs: Any) -> dict[str, Any]:
|
||||
"""Serializes the current object's state.
|
||||
|
||||
Args:
|
||||
**kwargs: Arguments for serialization.
|
||||
"""
|
||||
chat_message_store_state = None
|
||||
if self._message_store is not None:
|
||||
chat_message_store_state = await self._message_store.serialize_state(**kwargs)
|
||||
|
||||
state = ThreadState(
|
||||
service_thread_id=self._service_thread_id, chat_message_store_state=chat_message_store_state
|
||||
)
|
||||
|
||||
return state.model_dump()
|
||||
|
||||
|
||||
async def thread_on_new_messages(thread: AgentThread, new_messages: ChatMessage | Sequence[ChatMessage]) -> None:
|
||||
"""Invoked when a new message has been contributed to the chat by any participant."""
|
||||
if thread.service_thread_id is not None:
|
||||
# If the thread messages are stored in the service there is nothing to do here,
|
||||
# since invoking the service should already update the thread.
|
||||
return
|
||||
|
||||
if thread.message_store is None:
|
||||
# If there is no conversation id, and no store we can
|
||||
# create a default in memory store.
|
||||
thread.message_store = ChatMessageList()
|
||||
|
||||
# If a store has been provided, we need to add the messages to the store.
|
||||
if isinstance(new_messages, ChatMessage):
|
||||
new_messages = [new_messages]
|
||||
|
||||
await thread.message_store.add_messages(new_messages)
|
||||
|
||||
|
||||
async def deserialize_thread_state(
|
||||
thread: AgentThread,
|
||||
serialized_thread: dict[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Deserializes the state from a dictionary into the thread properties."""
|
||||
state = ThreadState.model_validate(serialized_thread)
|
||||
|
||||
if state.service_thread_id:
|
||||
thread.service_thread_id = state.service_thread_id
|
||||
# Since we have an ID, we should not have a chat message store and we can return here.
|
||||
return
|
||||
|
||||
# If we don't have any ChatMessageStore state return here.
|
||||
if state.chat_message_store_state is None:
|
||||
return
|
||||
|
||||
if thread.message_store is None:
|
||||
# If we don't have a chat message store yet, create an in-memory one.
|
||||
thread.message_store = ChatMessageList()
|
||||
|
||||
await thread.message_store.deserialize_state(state.chat_message_store_state, **kwargs)
|
||||
|
||||
|
||||
class ThreadState(AFBaseModel):
|
||||
"""State model for serializing and deserializing thread information.
|
||||
|
||||
Attributes:
|
||||
service_thread_id: Optional ID of the thread managed by the agent service.
|
||||
chat_message_store_state: Optional serialized state of the chat message store.
|
||||
"""
|
||||
|
||||
service_thread_id: str | None = None
|
||||
chat_message_store_state: Any | None = None
|
||||
|
||||
|
||||
class StoreState(AFBaseModel):
|
||||
"""State model for serializing and deserializing chat message store data.
|
||||
|
||||
Attributes:
|
||||
messages: List of chat messages stored in the message store.
|
||||
"""
|
||||
|
||||
messages: list[ChatMessage]
|
||||
|
||||
|
||||
class ChatMessageList:
|
||||
"""An in-memory implementation of ChatMessageStore that stores messages in a list.
|
||||
|
||||
This implementation provides a simple, list-based storage for chat messages
|
||||
with support for serialization and deserialization. It implements all the
|
||||
required methods of the ChatMessageStore protocol and provides additional
|
||||
list-like operations for direct message manipulation.
|
||||
|
||||
The store maintains messages in memory and provides methods to serialize
|
||||
and deserialize the state for persistence purposes.
|
||||
"""
|
||||
|
||||
def __init__(self, messages: Sequence[ChatMessage] | None = None) -> None:
|
||||
"""Initialize the message store with optional initial messages.
|
||||
|
||||
Args:
|
||||
messages: Optional collection of initial ChatMessage objects to store.
|
||||
"""
|
||||
self._messages: list[ChatMessage] = []
|
||||
if messages:
|
||||
self._messages.extend(messages)
|
||||
|
||||
async def add_messages(self, messages: Sequence[ChatMessage]) -> None:
|
||||
"""Add messages to the store.
|
||||
|
||||
Args:
|
||||
messages: Sequence of ChatMessage objects to add to the store.
|
||||
"""
|
||||
self._messages.extend(messages)
|
||||
|
||||
async def list_messages(self) -> list[ChatMessage]:
|
||||
"""Get all messages from the store in chronological order.
|
||||
|
||||
Returns:
|
||||
List of ChatMessage objects, ordered from oldest to newest.
|
||||
"""
|
||||
return self._messages
|
||||
|
||||
async def deserialize_state(self, serialized_store_state: Any, **kwargs: Any) -> None:
|
||||
"""Deserialize state data into this store instance.
|
||||
|
||||
Args:
|
||||
serialized_store_state: Previously serialized state data containing messages.
|
||||
**kwargs: Additional arguments for deserialization.
|
||||
"""
|
||||
if serialized_store_state:
|
||||
state = StoreState.model_validate(obj=serialized_store_state, **kwargs)
|
||||
if state.messages:
|
||||
self._messages.extend(state.messages)
|
||||
|
||||
async def serialize_state(self, **kwargs: Any) -> Any:
|
||||
"""Serialize the current store state for persistence.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional arguments for serialization.
|
||||
|
||||
Returns:
|
||||
Serialized state data that can be used with deserialize_state.
|
||||
"""
|
||||
state = StoreState(messages=self._messages)
|
||||
return state.model_dump(**kwargs)
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Return the number of messages in the store.
|
||||
|
||||
Returns:
|
||||
The count of messages currently stored.
|
||||
"""
|
||||
return len(self._messages)
|
||||
|
||||
def __getitem__(self, index: int) -> ChatMessage:
|
||||
"""Get a message by index.
|
||||
|
||||
Args:
|
||||
index: The index of the message to retrieve.
|
||||
|
||||
Returns:
|
||||
The ChatMessage at the specified index.
|
||||
"""
|
||||
return self._messages[index]
|
||||
|
||||
def __setitem__(self, index: int, item: ChatMessage) -> None:
|
||||
"""Set a message at the specified index.
|
||||
|
||||
Args:
|
||||
index: The index at which to set the message.
|
||||
item: The ChatMessage to set at the specified index.
|
||||
"""
|
||||
self._messages[index] = item
|
||||
|
||||
def append(self, item: ChatMessage) -> None:
|
||||
"""Append a message to the end of the store.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage to append.
|
||||
"""
|
||||
self._messages.append(item)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Remove all messages from the store."""
|
||||
self._messages.clear()
|
||||
|
||||
def index(self, item: ChatMessage) -> int:
|
||||
"""Return the index of the first occurrence of the specified message.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage to find.
|
||||
|
||||
Returns:
|
||||
The index of the first occurrence of the message.
|
||||
|
||||
Raises:
|
||||
ValueError: If the message is not found in the store.
|
||||
"""
|
||||
return self._messages.index(item)
|
||||
|
||||
def insert(self, index: int, item: ChatMessage) -> None:
|
||||
"""Insert a message at the specified index.
|
||||
|
||||
Args:
|
||||
index: The index at which to insert the message.
|
||||
item: The ChatMessage to insert.
|
||||
"""
|
||||
self._messages.insert(index, item)
|
||||
|
||||
def remove(self, item: ChatMessage) -> None:
|
||||
"""Remove the first occurrence of the specified message from the store.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage to remove.
|
||||
|
||||
Raises:
|
||||
ValueError: If the message is not found in the store.
|
||||
"""
|
||||
self._messages.remove(item)
|
||||
|
||||
def pop(self, index: int = -1) -> ChatMessage:
|
||||
"""Remove and return a message at the specified index.
|
||||
|
||||
Args:
|
||||
index: The index of the message to remove and return. Defaults to -1 (last item).
|
||||
|
||||
Returns:
|
||||
The ChatMessage that was removed.
|
||||
|
||||
Raises:
|
||||
IndexError: If the index is out of range.
|
||||
"""
|
||||
return self._messages.pop(index)
|
||||
@@ -18,8 +18,9 @@ from ._pydantic import AFBaseSettings
|
||||
if TYPE_CHECKING: # pragma: no cover
|
||||
from opentelemetry.util._decorator import _AgnosticContextManager # type: ignore[reportPrivateUsage]
|
||||
|
||||
from ._agents import AgentThread, AIAgent, ChatClientAgent
|
||||
from ._agents import AIAgent, ChatClientAgent
|
||||
from ._clients import ChatClientBase
|
||||
from ._threads import AgentThread
|
||||
from ._tools import AIFunction
|
||||
from ._types import (
|
||||
AgentRunResponse,
|
||||
@@ -650,8 +651,8 @@ def _get_agent_run_span(
|
||||
span.set_attribute(GenAIAttributes.AGENT_NAME.value, agent.name)
|
||||
if agent.description:
|
||||
span.set_attribute(GenAIAttributes.AGENT_DESCRIPTION.value, agent.description)
|
||||
if thread and thread.id:
|
||||
span.set_attribute(GenAIAttributes.CONVERSATION_ID.value, thread.id)
|
||||
if thread and thread.service_thread_id:
|
||||
span.set_attribute(GenAIAttributes.CONVERSATION_ID.value, thread.service_thread_id)
|
||||
if "model" in kwargs:
|
||||
span.set_attribute(GenAIAttributes.MODEL.value, kwargs["model"])
|
||||
if "seed" in kwargs:
|
||||
|
||||
Reference in New Issue
Block a user