feat: Support returning all sessions when user_id is none in the request

resolves https://github.com/google/adk-python/issues/3154

PiperOrigin-RevId: 819417330
This commit is contained in:
Google Team Member
2025-10-14 15:10:41 -07:00
committed by Copybara-Service
parent 141318f775
commit f9c09ef075
6 changed files with 41 additions and 148 deletions
@@ -83,18 +83,9 @@ class BaseSessionService(abc.ABC):
@abc.abstractmethod
async def list_sessions(
self, *, app_name: str, user_id: Optional[str] = None
self, *, app_name: str, user_id: str
) -> ListSessionsResponse:
"""Lists all the sessions for a user.
Args:
app_name: The name of the app.
user_id: The ID of the user. If not provided, lists all sessions for all
users.
Returns:
A ListSessionsResponse containing the sessions.
"""
"""Lists all the sessions."""
@abc.abstractmethod
async def delete_session(
@@ -554,42 +554,30 @@ class DatabaseSessionService(BaseSessionService):
@override
async def list_sessions(
self, *, app_name: str, user_id: Optional[str] = None
self, *, app_name: str, user_id: str
) -> ListSessionsResponse:
with self.database_session_factory() as sql_session:
query = sql_session.query(StorageSession).filter(
StorageSession.app_name == app_name
results = (
sql_session.query(StorageSession)
.filter(StorageSession.app_name == app_name)
.filter(StorageSession.user_id == user_id)
.all()
)
if user_id is not None:
query = query.filter(StorageSession.user_id == user_id)
results = query.all()
# Fetch app state from storage
# Fetch states from storage
storage_app_state = sql_session.get(StorageAppState, (app_name))
app_state = storage_app_state.state if storage_app_state else {}
storage_user_state = sql_session.get(
StorageUserState, (app_name, user_id)
)
# Fetch user state(s) from storage
user_states_map = {}
if user_id is not None:
storage_user_state = sql_session.get(
StorageUserState, (app_name, user_id)
)
if storage_user_state:
user_states_map[user_id] = storage_user_state.state
else:
all_user_states_for_app = (
sql_session.query(StorageUserState)
.filter(StorageUserState.app_name == app_name)
.all()
)
for storage_user_state in all_user_states_for_app:
user_states_map[storage_user_state.user_id] = storage_user_state.state
app_state = storage_app_state.state if storage_app_state else {}
user_state = storage_user_state.state if storage_user_state else {}
sessions = []
for storage_session in results:
session_state = storage_session.state
user_state = user_states_map.get(storage_session.user_id, {})
merged_state = _merge_state(app_state, user_state, session_state)
sessions.append(storage_session.to_session(state=merged_state))
return ListSessionsResponse(sessions=sessions)
@@ -201,41 +201,31 @@ class InMemorySessionService(BaseSessionService):
@override
async def list_sessions(
self, *, app_name: str, user_id: Optional[str] = None
self, *, app_name: str, user_id: str
) -> ListSessionsResponse:
return self._list_sessions_impl(app_name=app_name, user_id=user_id)
def list_sessions_sync(
self, *, app_name: str, user_id: Optional[str] = None
self, *, app_name: str, user_id: str
) -> ListSessionsResponse:
logger.warning('Deprecated. Please migrate to the async method.')
return self._list_sessions_impl(app_name=app_name, user_id=user_id)
def _list_sessions_impl(
self, *, app_name: str, user_id: Optional[str] = None
self, *, app_name: str, user_id: str
) -> ListSessionsResponse:
empty_response = ListSessionsResponse()
if app_name not in self.sessions:
return empty_response
if user_id is not None and user_id not in self.sessions[app_name]:
if user_id not in self.sessions[app_name]:
return empty_response
sessions_without_events = []
if user_id is None:
for user_id in self.sessions[app_name]:
for session_id in self.sessions[app_name][user_id]:
session = self.sessions[app_name][user_id][session_id]
copied_session = copy.deepcopy(session)
copied_session.events = []
copied_session = self._merge_state(app_name, user_id, copied_session)
sessions_without_events.append(copied_session)
else:
for session in self.sessions[app_name][user_id].values():
copied_session = copy.deepcopy(session)
copied_session.events = []
copied_session = self._merge_state(app_name, user_id, copied_session)
sessions_without_events.append(copied_session)
for session in self.sessions[app_name][user_id].values():
copied_session = copy.deepcopy(session)
copied_session.events = []
copied_session = self._merge_state(app_name, user_id, copied_session)
sessions_without_events.append(copied_session)
return ListSessionsResponse(sessions=sessions_without_events)
@override
@@ -200,25 +200,22 @@ class VertexAiSessionService(BaseSessionService):
@override
async def list_sessions(
self, *, app_name: str, user_id: Optional[str] = None
self, *, app_name: str, user_id: str
) -> ListSessionsResponse:
reasoning_engine_id = self._get_reasoning_engine_id(app_name)
api_client = self._get_api_client()
sessions = []
config = {}
if user_id is not None:
config['filter'] = f'user_id="{user_id}"'
sessions_iterator = api_client.agent_engines.sessions.list(
name=f'reasoningEngines/{reasoning_engine_id}',
config=config,
config={'filter': f'user_id="{user_id}"'},
)
for api_session in sessions_iterator:
sessions.append(
Session(
app_name=app_name,
user_id=api_session.user_id,
user_id=user_id,
id=api_session.name.split('/')[-1],
state=getattr(api_session, 'session_state', None) or {},
last_update_time=api_session.update_time.timestamp(),