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
51 changes: 19 additions & 32 deletions invokeai/app/services/boards/boards_default.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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))

Expand All @@ -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(
Expand Down
13 changes: 13 additions & 0 deletions invokeai/app/services/users/users_base.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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.
Expand Down
40 changes: 40 additions & 0 deletions invokeai/app/services/users/users_default.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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."""
Expand Down Expand Up @@ -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:
Expand Down
67 changes: 67 additions & 0 deletions tests/app/services/boards/test_boards_default.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
39 changes: 38 additions & 1 deletion tests/app/services/users/test_user_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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}
Loading