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.

_head_and_tail.py 2.7 kB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667
  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 FunctionCallMessage, Message, TextMessage
  5. class HeadAndTailChatMemory(ChatMemory[Message]):
  6. """A chat memory that keeps a view of the first n and last m messages,
  7. where n is the head size and m is the tail size. The head and tail sizes
  8. are set at initialization.
  9. Args:
  10. head_size (int): The size of the head.
  11. tail_size (int): The size of the tail.
  12. """
  13. def __init__(self, head_size: int, tail_size: int) -> None:
  14. self._messages: List[Message] = []
  15. self._head_size = head_size
  16. self._tail_size = tail_size
  17. async def add_message(self, message: Message) -> None:
  18. """Add a message to the memory."""
  19. self._messages.append(message)
  20. async def get_messages(self) -> List[Message]:
  21. """Get at most `head_size` recent messages and `tail_size` oldest messages."""
  22. head_messages = self._messages[: self._head_size]
  23. # Handle the last message is a function call message.
  24. if head_messages and isinstance(head_messages[-1], FunctionCallMessage):
  25. # Remove the last message from the head.
  26. head_messages = head_messages[:-1]
  27. tail_messages = self._messages[-self._tail_size :]
  28. # Handle the first message is a function call result message.
  29. if tail_messages and isinstance(tail_messages[0], FunctionExecutionResultMessage):
  30. # Remove the first message from the tail.
  31. tail_messages = tail_messages[1:]
  32. num_skipped = len(self._messages) - self._head_size - self._tail_size
  33. if num_skipped <= 0:
  34. # If there are not enough messages to fill the head and tail,
  35. # return all messages.
  36. return self._messages
  37. placeholder_messages = [TextMessage(content=f"Skipped {num_skipped} messages.", source="System")]
  38. return head_messages + placeholder_messages + tail_messages
  39. async def clear(self) -> None:
  40. """Clear the message memory."""
  41. self._messages = []
  42. def save_state(self) -> Mapping[str, Any]:
  43. return {
  44. "messages": [message for message in self._messages],
  45. "head_size": self._head_size,
  46. "tail_size": self._tail_size,
  47. "placeholder_message": self._placeholder_message,
  48. }
  49. def load_state(self, state: Mapping[str, Any]) -> None:
  50. self._messages = state["messages"]
  51. self._head_size = state["head_size"]
  52. self._tail_size = state["tail_size"]
  53. self._placeholder_message = state["placeholder_message"]