forked from strands-agents/harness-sdk
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_session.py
More file actions
149 lines (126 loc) · 6.41 KB
/
Copy pathtest_session.py
File metadata and controls
149 lines (126 loc) · 6.41 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
"""Integration tests for session management."""
import tempfile
from uuid import uuid4
import boto3
import pytest
from botocore.client import ClientError
from strands import Agent
from strands.agent.conversation_manager.sliding_window_conversation_manager import SlidingWindowConversationManager
from strands.session.file_session_manager import FileSessionManager
from strands.session.s3_session_manager import S3SessionManager
# yellow_img imported from conftest
@pytest.fixture
def temp_dir():
"""Create a temporary directory for testing."""
with tempfile.TemporaryDirectory() as temp_dir:
yield temp_dir
@pytest.fixture
def bucket_name():
bucket_name = f"test-strands-session-bucket-{boto3.client('sts').get_caller_identity()['Account']}"
s3_client = boto3.resource("s3", region_name="us-west-2")
try:
s3_client.create_bucket(Bucket=bucket_name, CreateBucketConfiguration={"LocationConstraint": "us-west-2"})
except ClientError as e:
if "BucketAlreadyOwnedByYou" not in str(e):
raise e
yield bucket_name
def test_agent_with_file_session(temp_dir):
# Set up the session manager and add an agent
test_session_id = str(uuid4())
# Create a session
session_manager = FileSessionManager(session_id=test_session_id, storage_dir=temp_dir)
try:
agent = Agent(session_manager=session_manager)
agent("Hello!")
assert len(session_manager.list_messages(test_session_id, agent.agent_id)) == 2
# After agent is persisted and run, restore the agent and run it again
session_manager_2 = FileSessionManager(session_id=test_session_id, storage_dir=temp_dir)
agent_2 = Agent(session_manager=session_manager_2)
assert len(agent_2.messages) == 2
agent_2("Hello!")
assert len(agent_2.messages) == 4
assert len(session_manager_2.list_messages(test_session_id, agent_2.agent_id)) == 4
finally:
# Delete the session
session_manager.delete_session(test_session_id)
assert session_manager.read_session(test_session_id) is None
def test_agent_with_file_session_and_conversation_manager(temp_dir):
# Set up the session manager and add an agent
test_session_id = str(uuid4())
# Create a session
session_manager = FileSessionManager(session_id=test_session_id, storage_dir=temp_dir)
try:
agent = Agent(
session_manager=session_manager, conversation_manager=SlidingWindowConversationManager(window_size=1)
)
agent("Hello!")
assert len(session_manager.list_messages(test_session_id, agent.agent_id)) == 2
# Conversation Manager reduced messages
assert len(agent.messages) == 1
# After agent is persisted and run, restore the agent and run it again
session_manager_2 = FileSessionManager(session_id=test_session_id, storage_dir=temp_dir)
agent_2 = Agent(
session_manager=session_manager_2, conversation_manager=SlidingWindowConversationManager(window_size=1)
)
assert len(agent_2.messages) == 1
assert agent_2.conversation_manager.removed_message_count == 1
agent_2("Hello!")
assert len(agent_2.messages) == 1
assert len(session_manager_2.list_messages(test_session_id, agent_2.agent_id)) == 4
finally:
# Delete the session
session_manager.delete_session(test_session_id)
assert session_manager.read_session(test_session_id) is None
def test_agent_with_file_session_with_image(temp_dir, yellow_img):
test_session_id = str(uuid4())
# Create a session
session_manager = FileSessionManager(session_id=test_session_id, storage_dir=temp_dir)
try:
agent = Agent(session_manager=session_manager)
agent([{"image": {"format": "png", "source": {"bytes": yellow_img}}}])
assert len(session_manager.list_messages(test_session_id, agent.agent_id)) == 2
# After agent is persisted and run, restore the agent and run it again
session_manager_2 = FileSessionManager(session_id=test_session_id, storage_dir=temp_dir)
agent_2 = Agent(session_manager=session_manager_2)
assert len(agent_2.messages) == 2
agent_2("Hello!")
assert len(agent_2.messages) == 4
assert len(session_manager_2.list_messages(test_session_id, agent_2.agent_id)) == 4
finally:
# Delete the session
session_manager.delete_session(test_session_id)
assert session_manager.read_session(test_session_id) is None
def test_agent_with_s3_session(bucket_name):
test_session_id = str(uuid4())
session_manager = S3SessionManager(session_id=test_session_id, bucket=bucket_name, region_name="us-west-2")
try:
agent = Agent(session_manager=session_manager)
agent("Hello!")
assert len(session_manager.list_messages(test_session_id, agent.agent_id)) == 2
# After agent is persisted and run, restore the agent and run it again
session_manager_2 = S3SessionManager(session_id=test_session_id, bucket=bucket_name, region_name="us-west-2")
agent_2 = Agent(session_manager=session_manager_2)
assert len(agent_2.messages) == 2
agent_2("Hello!")
assert len(agent_2.messages) == 4
assert len(session_manager_2.list_messages(test_session_id, agent_2.agent_id)) == 4
finally:
session_manager.delete_session(test_session_id)
assert session_manager.read_session(test_session_id) is None
def test_agent_with_s3_session_with_image(yellow_img, bucket_name):
test_session_id = str(uuid4())
session_manager = S3SessionManager(session_id=test_session_id, bucket=bucket_name, region_name="us-west-2")
try:
agent = Agent(session_manager=session_manager)
agent([{"image": {"format": "png", "source": {"bytes": yellow_img}}}])
assert len(session_manager.list_messages(test_session_id, agent.agent_id)) == 2
# After agent is persisted and run, restore the agent and run it again
session_manager_2 = S3SessionManager(session_id=test_session_id, bucket=bucket_name, region_name="us-west-2")
agent_2 = Agent(session_manager=session_manager_2)
assert len(agent_2.messages) == 2
agent_2("Hello!")
assert len(agent_2.messages) == 4
assert len(session_manager_2.list_messages(test_session_id, agent_2.agent_id)) == 4
finally:
session_manager.delete_session(test_session_id)
assert session_manager.read_session(test_session_id) is None