You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

_buffered.py 1.6 kB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647
  1. from typing import Any, List, Mapping
  2. from agnext.components.memory import ChatMemory
  3. from agnext.components.models import FunctionExecutionResultMessage
  4. from ..types import Message
  5. class BufferedChatMemory(ChatMemory[Message]):
  6. """A buffered chat memory that keeps a view of the last n messages,
  7. where n is the buffer size. The buffer size is set at initialization.
  8. Args:
  9. buffer_size (int): The size of the buffer.
  10. """
  11. def __init__(self, buffer_size: int) -> None:
  12. self._messages: List[Message] = []
  13. self._buffer_size = buffer_size
  14. async def add_message(self, message: Message) -> None:
  15. """Add a message to the memory."""
  16. self._messages.append(message)
  17. async def get_messages(self) -> List[Message]:
  18. """Get at most `buffer_size` recent messages."""
  19. messages = self._messages[-self._buffer_size :]
  20. # Handle the first message is a function call result message.
  21. if messages and isinstance(messages[0], FunctionExecutionResultMessage):
  22. # Remove the first message from the list.
  23. messages = messages[1:]
  24. return messages
  25. async def clear(self) -> None:
  26. """Clear the message memory."""
  27. self._messages = []
  28. def save_state(self) -> Mapping[str, Any]:
  29. return {
  30. "messages": [message for message in self._messages],
  31. "buffer_size": self._buffer_size,
  32. }
  33. def load_state(self, state: Mapping[str, Any]) -> None:
  34. self._messages = state["messages"]
  35. self._buffer_size = state["buffer_size"]