diff --git a/invokeai/app/services/boards/boards_default.py b/invokeai/app/services/boards/boards_default.py index 750a4504929..8088e2b0e47 100644 --- a/invokeai/app/services/boards/boards_default.py +++ b/invokeai/app/services/boards/boards_default.py @@ -1,6 +1,6 @@ from typing import Optional -from invokeai.app.services.board_records.board_records_common import BoardChanges, BoardRecordOrderBy +from invokeai.app.services.board_records.board_records_common import BoardChanges, BoardRecord, BoardRecordOrderBy from invokeai.app.services.boards.boards_base import BoardServiceABC from invokeai.app.services.boards.boards_common import BoardDTO, board_record_to_dto from invokeai.app.services.invoker import Invoker @@ -100,31 +100,7 @@ def get_many( board_records = self.__invoker.services.board_records.get_many( user_id, is_admin, order_by, direction, offset, limit, include_archived ) - summaries = self.__invoker.services.gallery.get_board_media_summaries( - [record.board_id for record in board_records.items] - ) - board_dtos = [] - for r in board_records.items: - summary = summaries[r.board_id] - - # For admin users, include owner username - owner_username = None - if is_admin: - owner = self.__invoker.services.users.get(r.user_id) - if owner: - owner_username = owner.display_name or owner.email - - board_dtos.append( - board_record_to_dto( - r, - summary.cover_image_name, - summary.image_count, - summary.asset_count, - owner_username, - cover_video_name=summary.cover_video_name, - video_count=summary.video_count, - ) - ) + board_dtos = self._to_dtos(board_records.items, is_admin) return OffsetPaginatedResults[BoardDTO](items=board_dtos, offset=offset, limit=limit, total=len(board_dtos)) @@ -139,19 +115,30 @@ def get_all( board_records = self.__invoker.services.board_records.get_all( user_id, is_admin, order_by, direction, include_archived ) + return self._to_dtos(board_records, is_admin) + + def _to_dtos(self, board_records: list[BoardRecord], is_admin: bool) -> list[BoardDTO]: + """Builds board DTOs for a listing with a fixed number of queries. + + Both the media summaries and (for admins) the owner display names are fetched for + the whole page at once. The owner lookup used to run one `users.get` per board, so + an admin listing 50 boards issued 50 extra queries for what is usually a handful of + distinct owners. + """ summaries = self.__invoker.services.gallery.get_board_media_summaries( [record.board_id for record in board_records] ) - board_dtos = [] + owners = ( + self.__invoker.services.users.get_many([record.user_id for record in board_records]) if is_admin else {} + ) + + board_dtos: list[BoardDTO] = [] for r in board_records: summary = summaries[r.board_id] # For admin users, include owner username - owner_username = None - if is_admin: - owner = self.__invoker.services.users.get(r.user_id) - if owner: - owner_username = owner.display_name or owner.email + owner = owners.get(r.user_id) + owner_username = (owner.display_name or owner.email) if owner else None board_dtos.append( board_record_to_dto( diff --git a/invokeai/app/services/users/users_base.py b/invokeai/app/services/users/users_base.py index dd789b561ee..65c58ee2315 100644 --- a/invokeai/app/services/users/users_base.py +++ b/invokeai/app/services/users/users_base.py @@ -1,6 +1,7 @@ """Abstract base class for user service.""" from abc import ABC, abstractmethod +from collections.abc import Sequence from invokeai.app.services.users.users_common import UserCreateRequest, UserDTO, UserUpdateRequest @@ -37,6 +38,18 @@ def get(self, user_id: str) -> UserDTO | None: """ pass + @abstractmethod + def get_many(self, user_ids: Sequence[str]) -> dict[str, UserDTO]: + """Get several users at once. + + Args: + user_ids: The user IDs to look up. Duplicates are collapsed. + + Returns: + A dict keyed by user_id; ids with no matching user are absent. + """ + pass + @abstractmethod def get_by_email(self, email: str) -> UserDTO | None: """Get user by email. diff --git a/invokeai/app/services/users/users_default.py b/invokeai/app/services/users/users_default.py index 6e472882124..3ddc1a03274 100644 --- a/invokeai/app/services/users/users_default.py +++ b/invokeai/app/services/users/users_default.py @@ -1,6 +1,7 @@ """Default SQLite implementation of user service.""" import sqlite3 +from collections.abc import Sequence from datetime import datetime, timezone from uuid import uuid4 @@ -9,6 +10,10 @@ from invokeai.app.services.users.users_base import UserServiceBase from invokeai.app.services.users.users_common import UserCreateRequest, UserDTO, UserUpdateRequest +# Bound-parameter chunk size for `get_many`. SQLite's SQLITE_MAX_VARIABLE_NUMBER is 32766 +# on modern builds but 999 on older ones; stay under the smaller limit. +USER_LOOKUP_CHUNK_SIZE = 900 + class UserService(UserServiceBase): """SQLite-based user service.""" @@ -82,6 +87,41 @@ def get(self, user_id: str) -> UserDTO | None: last_login_at=datetime.fromisoformat(row[7]) if row[7] else None, ) + def get_many(self, user_ids: Sequence[str]) -> dict[str, UserDTO]: + """Get users by ID, keyed by user_id. Unknown ids are absent from the result.""" + unique_ids = list(dict.fromkeys(user_ids)) + if not unique_ids: + return {} + + users: dict[str, UserDTO] = {} + # Chunked so a caller with many distinct owners can't exceed SQLite's bound-parameter + # limit (999 on builds predating 3.32). + for start in range(0, len(unique_ids), USER_LOOKUP_CHUNK_SIZE): + chunk = unique_ids[start : start + USER_LOOKUP_CHUNK_SIZE] + placeholders = ",".join("?" for _ in chunk) + with self._db.transaction() as cursor: + cursor.execute( + f""" + SELECT user_id, email, display_name, is_admin, is_active, created_at, updated_at, last_login_at + FROM users + WHERE user_id IN ({placeholders}) + """, + chunk, + ) + rows = cursor.fetchall() + for row in rows: + users[row[0]] = UserDTO( + user_id=row[0], + email=row[1], + display_name=row[2], + is_admin=bool(row[3]), + is_active=bool(row[4]), + created_at=datetime.fromisoformat(row[5]), + updated_at=datetime.fromisoformat(row[6]), + last_login_at=datetime.fromisoformat(row[7]) if row[7] else None, + ) + return users + def get_by_email(self, email: str) -> UserDTO | None: """Get user by email.""" with self._db.transaction() as cursor: diff --git a/tests/app/services/boards/test_boards_default.py b/tests/app/services/boards/test_boards_default.py index 96cf63bb5e9..6c13099a1fe 100644 --- a/tests/app/services/boards/test_boards_default.py +++ b/tests/app/services/boards/test_boards_default.py @@ -55,3 +55,70 @@ def test_board_listing_fetches_media_summaries_once(mock_invoker: Invoker) -> No mock_invoker.services.board_image_records.get_image_count_for_board.assert_not_called() mock_invoker.services.board_image_records.get_asset_count_for_board.assert_not_called() mock_invoker.services.board_video_records.get_video_count_for_board.assert_not_called() + + +def test_admin_board_listing_batches_owner_lookup(mock_invoker: Invoker) -> None: + """Owner display names are fetched for the whole page in one call. + + An admin listing used to issue one `users.get` per board — 50 boards meant 50 extra + queries for what is usually a handful of distinct owners. + """ + owners = ["alice", "bob", "alice", "carol"] * 5 + board_ids = [ + mock_invoker.services.board_records.save(f"Board {index}", owner).board_id for index, owner in enumerate(owners) + ] + summaries = { + board_id: SimpleNamespace( + cover_image_name=None, + cover_video_name=None, + image_count=0, + video_count=0, + asset_count=0, + ) + for board_id in board_ids + } + mock_invoker.services.gallery.get_board_media_summaries = MagicMock(return_value=summaries) # type: ignore[attr-defined] + mock_invoker.services.users.get = MagicMock() # type: ignore[method-assign] + mock_invoker.services.users.get_many = MagicMock( # type: ignore[method-assign] + return_value={ + name: SimpleNamespace(display_name=name.title(), email=f"{name}@example.com") + for name in ("alice", "bob", "carol") + } + ) + + result = mock_invoker.services.boards.get_all( + user_id="admin", + is_admin=True, + order_by=BoardRecordOrderBy.Name, + direction=SQLiteDirection.Ascending, + ) + + assert len(result) == len(board_ids) + assert {dto.owner_username for dto in result} == {"Alice", "Bob", "Carol"} + mock_invoker.services.users.get.assert_not_called() # type: ignore[attr-defined] + mock_invoker.services.users.get_many.assert_called_once() # type: ignore[attr-defined] + + +def test_non_admin_board_listing_skips_owner_lookup(mock_invoker: Invoker) -> None: + """Non-admin listings don't show owner names, so they must not query for them at all.""" + board_id = mock_invoker.services.board_records.save("Board", "user").board_id + mock_invoker.services.gallery.get_board_media_summaries = MagicMock( # type: ignore[attr-defined] + return_value={ + board_id: SimpleNamespace( + cover_image_name=None, cover_video_name=None, image_count=0, video_count=0, asset_count=0 + ) + } + ) + mock_invoker.services.users.get = MagicMock() # type: ignore[method-assign] + mock_invoker.services.users.get_many = MagicMock() # type: ignore[method-assign] + + result = mock_invoker.services.boards.get_all( + user_id="user", + is_admin=False, + order_by=BoardRecordOrderBy.Name, + direction=SQLiteDirection.Ascending, + ) + + assert [dto.owner_username for dto in result] == [None] + mock_invoker.services.users.get.assert_not_called() # type: ignore[attr-defined] + mock_invoker.services.users.get_many.assert_not_called() # type: ignore[attr-defined] diff --git a/tests/app/services/users/test_user_service.py b/tests/app/services/users/test_user_service.py index d5d04964005..1cdca16b00c 100644 --- a/tests/app/services/users/test_user_service.py +++ b/tests/app/services/users/test_user_service.py @@ -6,7 +6,7 @@ from invokeai.app.services.shared.sqlite.sqlite_database import SqliteDatabase from invokeai.app.services.users.users_common import UserCreateRequest, UserUpdateRequest -from invokeai.app.services.users.users_default import UserService +from invokeai.app.services.users.users_default import USER_LOOKUP_CHUNK_SIZE, UserService @pytest.fixture @@ -270,3 +270,40 @@ def test_list_users(user_service: UserService): limited_users = user_service.list_users(limit=2) assert len(limited_users) == 2 + + +def test_get_many_returns_users_keyed_by_id(user_service: UserService): + """Batch lookup: dedups input, keys by user_id, and omits unknown ids.""" + created = [ + user_service.create( + UserCreateRequest( + email=f"batch{index}@example.com", + display_name=f"Batch User {index}", + password="TestPassword123", + ) + ) + for index in range(3) + ] + + requested = [created[0].user_id, created[1].user_id, created[0].user_id, "does-not-exist"] + users = user_service.get_many(requested) + + assert set(users) == {created[0].user_id, created[1].user_id} + assert users[created[0].user_id].email == "batch0@example.com" + assert users[created[1].user_id].display_name == "Batch User 1" + + +def test_get_many_with_no_ids_returns_empty(user_service: UserService): + assert user_service.get_many([]) == {} + + +def test_get_many_chunks_beyond_sqlite_parameter_limit(user_service: UserService): + """More ids than SQLite's bound-parameter limit must not raise.""" + user = user_service.create( + UserCreateRequest(email="chunked@example.com", display_name="Chunked", password="TestPassword123") + ) + ids = [f"missing-{index}" for index in range(USER_LOOKUP_CHUNK_SIZE * 2 + 5)] + [user.user_id] + + users = user_service.get_many(ids) + + assert set(users) == {user.user_id}