Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion atomic-agents/atomic_agents/context/chat_history.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import json
import uuid
from copy import deepcopy
from enum import Enum
from pathlib import Path
from typing import Dict, List, Optional, Tuple, Type, Union
Expand Down Expand Up @@ -201,7 +202,7 @@ def copy(self) -> "ChatHistory":
ChatHistory: A copy of the chat history.
"""
new_history = ChatHistory(max_messages=self.max_messages)
new_history.load(self.dump())
new_history.history = deepcopy(self.history)
new_history.current_turn_id = self.current_turn_id
return new_history

Expand Down
24 changes: 24 additions & 0 deletions atomic-agents/tests/agents/test_atomic_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,30 @@ def test_initialization(agent, mock_instructor, mock_history, mock_system_prompt
assert "max_tokens" not in agent.model_api_parameters


def test_initial_history_with_local_schema_can_be_reset(mock_instructor):
class LocalInput(BaseIOSchema):
"""Input schema created inside an application factory."""

chat_message: str

history = ChatHistory()
history.add_message("user", LocalInput(chat_message="Initial context"))
agent = AtomicAgent[LocalInput, BasicChatOutputSchema](AgentConfig(client=mock_instructor, history=history))

result = agent.run(LocalInput(chat_message="New question"))
assert result.chat_message == "Test output"
assert agent.history.get_message_count() == 3

agent.reset_history()
assert agent.history.get_message_count() == 1
assert isinstance(agent.history.history[0].content, LocalInput)
assert agent.history.history[0].content.chat_message == "Initial context"

agent.history.history[0].content.chat_message = "Changed"
agent.reset_history()
assert agent.history.history[0].content.chat_message == "Initial context"


# model_api_parameters should have priority over other settings
def test_initialization_temperature_priority(mock_instructor, mock_history, mock_system_prompt_generator):
config = AgentConfig(
Expand Down
23 changes: 23 additions & 0 deletions atomic-agents/tests/context/test_chat_history.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,29 @@ def test_copy(history):
assert copied_history.history[0].content.test_field == history.history[0].content.test_field


def test_copy_with_local_schema_is_independent(history):
class LocalSchema(BaseIOSchema):
"""Message schema defined by an application factory."""

items: List[Dict[str, List[str]]]

history.add_message("user", LocalSchema(items=[{"values": ["original"]}]))
copied_history = history.copy()

assert copied_history.max_messages == history.max_messages
assert copied_history.current_turn_id == history.current_turn_id
assert copied_history.get_history() == history.get_history()
assert isinstance(copied_history.history[0].content, LocalSchema)

copied_history.history[0].content.items[0]["values"].append("changed")
copied_history.history[0].role = "assistant"
copied_history.add_message("user", LocalSchema(items=[]))

assert history.history[0].content.items == [{"values": ["original"]}]
assert history.history[0].role == "user"
assert history.get_message_count() == 1


def test_get_current_turn_id(history):
assert history.get_current_turn_id() is None
history.initialize_turn()
Expand Down
Loading