feat: Support return all sessions when user id is none

PiperOrigin-RevId: 819884236
This commit is contained in:
Shangjie Chen
2025-10-15 13:14:24 -07:00
committed by Copybara-Service
parent b650181384
commit a985cc38ec
6 changed files with 147 additions and 40 deletions
@@ -116,9 +116,70 @@ async def test_create_and_list_sessions(service_type):
app_name=app_name, user_id=user_id
)
sessions = list_sessions_response.sessions
for i in range(len(sessions)):
assert sessions[i].id == session_ids[i]
assert sessions[i].state == {'key': 'value' + session_ids[i]}
assert len(sessions) == len(session_ids)
assert {s.id for s in sessions} == set(session_ids)
for session in sessions:
assert session.state == {'key': 'value' + session.id}
@pytest.mark.asyncio
@pytest.mark.parametrize(
'service_type', [SessionServiceType.IN_MEMORY, SessionServiceType.DATABASE]
)
async def test_list_sessions_all_users(service_type):
session_service = get_session_service(service_type)
app_name = 'my_app'
user_id_1 = 'user1'
user_id_2 = 'user2'
await session_service.create_session(
app_name=app_name,
user_id=user_id_1,
session_id='session1a',
state={'key': 'value1a'},
)
await session_service.create_session(
app_name=app_name,
user_id=user_id_1,
session_id='session1b',
state={'key': 'value1b'},
)
await session_service.create_session(
app_name=app_name,
user_id=user_id_2,
session_id='session2a',
state={'key': 'value2a'},
)
# List sessions for user1
list_sessions_response_1 = await session_service.list_sessions(
app_name=app_name, user_id=user_id_1
)
sessions_1 = list_sessions_response_1.sessions
assert len(sessions_1) == 2
assert {s.id for s in sessions_1} == {'session1a', 'session1b'}
for session in sessions_1:
if session.id == 'session1a':
assert session.state == {'key': 'value1a'}
else:
assert session.state == {'key': 'value1b'}
# List sessions for user2
list_sessions_response_2 = await session_service.list_sessions(
app_name=app_name, user_id=user_id_2
)
sessions_2 = list_sessions_response_2.sessions
assert len(sessions_2) == 1
assert sessions_2[0].id == 'session2a'
assert sessions_2[0].state == {'key': 'value2a'}
# List sessions for all users
list_sessions_response_all = await session_service.list_sessions(
app_name=app_name, user_id=None
)
sessions_all = list_sessions_response_all.sessions
assert len(sessions_all) == 3
assert {s.id for s in sessions_all} == {'session1a', 'session1b', 'session2a'}
@pytest.mark.asyncio
@@ -252,19 +252,22 @@ class MockApiClient:
def _list_sessions(self, name: str, config: dict[str, Any]):
filter_val = config.get('filter', '')
user_id_match = re.search(r'user_id="([^"]+)"', filter_val)
if not user_id_match:
raise ValueError(f'Could not find user_id in filter: {filter_val}')
user_id = user_id_match.group(1)
if user_id == 'user_with_pages':
if user_id_match:
user_id = user_id_match.group(1)
if user_id == 'user_with_pages':
return [
_convert_to_object(MOCK_SESSION_JSON_PAGE1),
_convert_to_object(MOCK_SESSION_JSON_PAGE2),
]
return [
_convert_to_object(MOCK_SESSION_JSON_PAGE1),
_convert_to_object(MOCK_SESSION_JSON_PAGE2),
_convert_to_object(session)
for session in self.session_dict.values()
if session['user_id'] == user_id
]
# No user filter, return all sessions
return [
_convert_to_object(session)
for session in self.session_dict.values()
if session['user_id'] == user_id
_convert_to_object(session) for session in self.session_dict.values()
]
def _delete_session(self, name: str):
@@ -475,6 +478,15 @@ async def test_list_sessions_with_pagination():
assert sessions.sessions[1].id == 'page2'
@pytest.mark.asyncio
@pytest.mark.usefixtures('mock_get_api_client')
async def test_list_sessions_all_users():
session_service = mock_vertex_ai_session_service()
sessions = await session_service.list_sessions(app_name='123', user_id=None)
assert len(sessions.sessions) == 5
assert {s.id for s in sessions.sessions} == {'1', '2', '3', 'page1', 'page2'}
@pytest.mark.asyncio
@pytest.mark.usefixtures('mock_get_api_client')
async def test_create_session():