diff --git a/model/jpa/src/main/java/org/keycloak/models/jpa/GroupAdapter.java b/model/jpa/src/main/java/org/keycloak/models/jpa/GroupAdapter.java index 7d0eb58d4d10..320948cd075f 100755 --- a/model/jpa/src/main/java/org/keycloak/models/jpa/GroupAdapter.java +++ b/model/jpa/src/main/java/org/keycloak/models/jpa/GroupAdapter.java @@ -185,10 +185,10 @@ public Stream getSubGroupsStream() { @Override public Stream getSubGroupsStream(String search, Boolean exact, Integer firstResult, Integer maxResults) { CriteriaBuilder builder = em.getCriteriaBuilder(); - CriteriaQuery queryBuilder = builder.createQuery(String.class); + CriteriaQuery queryBuilder = builder.createQuery(GroupEntity.class); Root root = queryBuilder.from(GroupEntity.class); - queryBuilder.select(root.get("id")); + queryBuilder.select(root); List predicates = new ArrayList<>(); @@ -216,6 +216,9 @@ public Stream getSubGroupsStream(String search, Boolean exact, Integ queryBuilder.orderBy(builder.asc(root.get("name"))); return closing(paginateQuery(em.createQuery(queryBuilder), firstResult, maxResults).getResultStream() + // The selected entities are managed by the persistence context, so the by-id lookup is served + // from it (or the group cache) without issuing one additional query per group. + .map(GroupEntity::getId) .map(realm::getGroupById) // In concurrent tests, the group might be deleted in another thread, therefore, skip those null values. .filter(Objects::nonNull) diff --git a/model/jpa/src/main/java/org/keycloak/models/jpa/JpaRealmProvider.java b/model/jpa/src/main/java/org/keycloak/models/jpa/JpaRealmProvider.java index bec0ae2253fa..0a8d65700911 100644 --- a/model/jpa/src/main/java/org/keycloak/models/jpa/JpaRealmProvider.java +++ b/model/jpa/src/main/java/org/keycloak/models/jpa/JpaRealmProvider.java @@ -781,10 +781,10 @@ public Stream getGroupsByRoleStream(RealmModel realm, RoleModel role @Override public Stream getTopLevelGroupsStream(RealmModel realm, String search, Boolean exact, Integer firstResult, Integer maxResults) { CriteriaBuilder builder = em.getCriteriaBuilder(); - CriteriaQuery queryBuilder = builder.createQuery(String.class); + CriteriaQuery queryBuilder = builder.createQuery(GroupEntity.class); Root root = queryBuilder.from(GroupEntity.class); - queryBuilder.select(root.get("id")); + queryBuilder.select(root); List predicates = new ArrayList<>(); @@ -804,6 +804,9 @@ public Stream getTopLevelGroupsStream(RealmModel realm, String searc queryBuilder.orderBy(builder.asc(root.get("name"))); return closing(paginateQuery(em.createQuery(queryBuilder), firstResult, maxResults).getResultStream() + // The selected entities are managed by the persistence context, so the by-id lookup is served + // from it (or the group cache) without issuing one additional query per group. + .map(GroupEntity::getId) .map(realm::getGroupById) // In concurrent tests, the group might be deleted in another thread, therefore, skip those null values. .filter(Objects::nonNull) @@ -1422,10 +1425,10 @@ public Map getClientScopes(RealmModel realm, ClientMod @Override public Stream searchForGroupByNameStream(RealmModel realm, String search, Boolean exact, Integer first, Integer max) { CriteriaBuilder builder = em.getCriteriaBuilder(); - CriteriaQuery queryBuilder = builder.createQuery(String.class); + CriteriaQuery queryBuilder = builder.createQuery(GroupEntity.class); Root root = queryBuilder.from(GroupEntity.class); - queryBuilder.select(root.get("id")); + queryBuilder.select(root); List predicates = new ArrayList<>(); @@ -1445,6 +1448,9 @@ public Stream searchForGroupByNameStream(RealmModel realm, String se queryBuilder.orderBy(builder.asc(root.get("name"))); return closing(paginateQuery(em.createQuery(queryBuilder), first, max).getResultStream() + // The selected entities are managed by the persistence context, so the by-id lookup is served + // from it (or the group cache) without issuing one additional query per group. + .map(GroupEntity::getId) .map(id -> session.groups().getGroupById(realm, id)) .filter(Objects::nonNull) .distinct()); diff --git a/tests/base/src/test/java/org/keycloak/tests/model/GroupListingQueryCountTest.java b/tests/base/src/test/java/org/keycloak/tests/model/GroupListingQueryCountTest.java new file mode 100644 index 000000000000..ddefdc409e9e --- /dev/null +++ b/tests/base/src/test/java/org/keycloak/tests/model/GroupListingQueryCountTest.java @@ -0,0 +1,113 @@ +package org.keycloak.tests.model; + +import java.util.List; +import java.util.function.Supplier; +import java.util.stream.Collectors; + +import org.keycloak.connections.jpa.JpaConnectionProvider; +import org.keycloak.models.GroupModel; +import org.keycloak.models.KeycloakSession; +import org.keycloak.models.RealmModel; +import org.keycloak.models.cache.CacheRealmProvider; +import org.keycloak.testframework.annotations.InjectRealm; +import org.keycloak.testframework.annotations.KeycloakIntegrationTest; +import org.keycloak.testframework.realm.ManagedRealm; +import org.keycloak.testframework.realm.RealmBuilder; +import org.keycloak.testframework.realm.RealmConfig; +import org.keycloak.testframework.remote.annotations.TestOnServer; + +import org.hibernate.SessionFactory; +import org.hibernate.stat.Statistics; + +import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.is; + +@KeycloakIntegrationTest +public class GroupListingQueryCountTest { + + private static final int GROUP_COUNT = 25; + private static final int PAGE_SIZE = 20; + + @InjectRealm(config = GroupListingRealmConfig.class) + ManagedRealm realm; + + @TestOnServer + public void topLevelGroupsPageIsHydratedByASingleQuery(KeycloakSession session) { + RealmModel realm = session.getContext().getRealm(); + long count = countColdPageQueries(session, () -> + session.groups().getTopLevelGroupsStream(realm, 0, PAGE_SIZE).collect(Collectors.toList())); + // The page of groups is selected as entities in one query; hydrating the models must not + // issue one additional lookup per group. + assertThat(count, is(1L)); + } + + @TestOnServer + public void searchByNamePageIsHydratedByASingleQuery(KeycloakSession session) { + RealmModel realm = session.getContext().getRealm(); + long count = countColdPageQueries(session, () -> + session.groups().searchForGroupByNameStream(realm, "grp", false, GROUP_COUNT - PAGE_SIZE, PAGE_SIZE) + .collect(Collectors.toList())); + assertThat(count, is(1L)); + } + + @TestOnServer + public void subGroupsPageIsHydratedByASingleQuery(KeycloakSession session) { + RealmModel realm = session.getContext().getRealm(); + GroupModel parent = session.groups().searchForGroupByNameStream(realm, "sub-parent", true, 0, 1) + .findFirst().orElseGet(() -> { + GroupModel created = session.groups().createGroup(realm, "sub-parent"); + for (int i = 0; i < GROUP_COUNT; i++) { + session.groups().createGroup(realm, String.format("sub-%02d", i), created); + } + return created; + }); + long count = countColdPageQueries(session, () -> + parent.getSubGroupsStream(null, false, 0, PAGE_SIZE).collect(Collectors.toList())); + assertThat(count, is(1L)); + } + + /** + * Measures the statements needed to load a page of groups when the group cache cannot serve them, + * as after an eviction or on another cluster node. The warmup run collects the ids of the page, + * registering an invalidation for each id forces the by-id resolution to the store, and clearing + * the persistence context makes sure the warmup run cannot satisfy those lookups either. + * ({@code CacheRealmProvider.clear()} is a cluster notification that does not take effect within + * the running session, so the miss path is forced per group instead.) + */ + private static long countColdPageQueries(KeycloakSession session, Supplier> loadPage) { + List warmup = loadPage.get(); + assertThat(warmup.size(), is(PAGE_SIZE)); + + CacheRealmProvider cache = session.getProvider(CacheRealmProvider.class); + warmup.forEach(group -> cache.registerGroupInvalidation(group.getId())); + session.getProvider(JpaConnectionProvider.class).getEntityManager().clear(); + + return countQueries(session, () -> assertThat(loadPage.get().size(), is(PAGE_SIZE))); + } + + private static long countQueries(KeycloakSession session, Runnable walk) { + Statistics stats = session.getProvider(JpaConnectionProvider.class) + .getEntityManager() + .getEntityManagerFactory() + .unwrap(SessionFactory.class) + .getStatistics(); + boolean wasEnabled = stats.isStatisticsEnabled(); + stats.setStatisticsEnabled(true); + stats.clear(); + walk.run(); + long count = stats.getPrepareStatementCount(); + stats.setStatisticsEnabled(wasEnabled); + return count; + } + + private static class GroupListingRealmConfig implements RealmConfig { + @Override + public RealmBuilder configure(RealmBuilder r) { + r.name("group-listing-query-count-realm"); + for (int i = 0; i < GROUP_COUNT; i++) { + r.groups(String.format("grp-%02d", i)); + } + return r; + } + } +}