diff --git a/repository/src/main/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategy.java b/repository/src/main/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategy.java index 7375b0a60c..1541eb34ec 100644 --- a/repository/src/main/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategy.java +++ b/repository/src/main/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategy.java @@ -125,6 +125,11 @@ public class ScrollSearchStrategy extends SearchExecutionStrategy LOGGER.trace("Scroll response JSON: {}", scrollResponse.toJsonString()); validateResponse(scrollResponse); + if (scrollResponse.hits().hits() == null || scrollResponse.hits().hits().isEmpty()) + { + break; + } + scrollId = scrollResponse.scrollId(); resultList.addAll(skipHits(scrollResponse.hits().hits(), skipCount, limit - resultList.size())); } diff --git a/repository/src/test/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategyTest.java b/repository/src/test/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategyTest.java index 141675ab70..2413393713 100644 --- a/repository/src/test/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategyTest.java +++ b/repository/src/test/java/org/alfresco/repo/search/impl/elasticsearch/query/ScrollSearchStrategyTest.java @@ -27,6 +27,7 @@ package org.alfresco.repo.search.impl.elasticsearch.query; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyInt; import static org.mockito.ArgumentMatchers.anyList; @@ -35,11 +36,13 @@ import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; import java.util.ArrayList; import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.Before; import org.junit.Test; @@ -50,11 +53,13 @@ import org.opensearch.client.opensearch.OpenSearchClient; import org.opensearch.client.opensearch._types.ShardStatistics; import org.opensearch.client.opensearch._types.Time; import org.opensearch.client.opensearch._types.query_dsl.Query; +import org.opensearch.client.opensearch.core.ClearScrollRequest; import org.opensearch.client.opensearch.core.ScrollRequest; import org.opensearch.client.opensearch.core.ScrollResponse; import org.opensearch.client.opensearch.core.SearchRequest; import org.opensearch.client.opensearch.core.SearchResponse; import org.opensearch.client.opensearch.core.search.Hit; +import org.opensearch.client.opensearch.core.search.HitsMetadata; import org.opensearch.client.opensearch.core.search.TotalHits; import org.opensearch.client.opensearch.core.search.TotalHitsRelation; @@ -66,6 +71,8 @@ import org.alfresco.service.cmr.search.SearchParameters; public class ScrollSearchStrategyTest { + private static final int SCROLL_CALL_SAFETY_LIMIT = 10; + @Mock private SearchRequestBuilderService requestBuilderService; @Mock @@ -269,4 +276,175 @@ public class ScrollSearchStrategyTest resultSet.close(); } } + + @Test + public void search_scrollExhausted_stopsScrollingAndReturnsPartialPage() throws Exception + { + when(searchParameters.getStores()).thenReturn(new ArrayList<>()); + when(searchParameters.getSkipCount()).thenReturn(3); + when(searchParameters.getLimit()).thenReturn(4); + + when(requestBuilderService.buildSearchRequest(any(), any(), anyInt(), any(Time.class), anyString())) + .thenReturn(new SearchRequestWrapper.Builder().searchRequest(searchRequest).build()); + + when(client.search(any(SearchRequest.class), eq(Object.class))) + .thenReturn(searchResponse(hits("1", "2", "3"), 5L, "scroll-0")); + + ScrollResponse lastBatch = scrollResponse(hits("4", "5"), 5L, "scroll-1"); + ScrollResponse exhausted = scrollResponse(List.of(), 5L, "scroll-1"); + AtomicInteger scrollCalls = new AtomicInteger(); + when(client.scroll(any(ScrollRequest.class), eq(Object.class))).thenAnswer(invocation -> { + int call = scrollCalls.incrementAndGet(); + if (call > SCROLL_CALL_SAFETY_LIMIT) + { + throw new AssertionError("Scroll loop did not terminate: " + call + " scroll requests issued"); + } + return call == 1 ? lastBatch : exhausted; + }); + + ArgumentCaptor collected = ArgumentCaptor.forClass(List.class); + ElasticsearchResultSet rsMock = mock(ElasticsearchResultSet.class); + when(resultSetBuilder.build(eq(searchParameters), collected.capture(), anyLong(), anyLong())).thenReturn(rsMock); + + ResultSet resultSet = strategy.executeSearch(searchParameters, queryWithPermissions); + try + { + assertNotNull(resultSet); + assertEquals(2, scrollCalls.get()); + verify(client, times(2)).scroll(any(ScrollRequest.class), eq(Object.class)); + + List page = collected.getValue(); + assertEquals(2, page.size()); + assertEquals("4", ((Hit) page.get(0)).id()); + assertEquals("5", ((Hit) page.get(1)).id()); + verify(resultSetBuilder).build(eq(searchParameters), anyList(), eq(5L), anyLong()); + verify(client).clearScroll(any(ClearScrollRequest.class)); + } + finally + { + resultSet.close(); + } + } + + @Test + public void search_skipCountBeyondTotalHits_stopsScrollingAndReturnsEmptyPage() throws Exception + { + when(searchParameters.getStores()).thenReturn(new ArrayList<>()); + when(searchParameters.getSkipCount()).thenReturn(6); + when(searchParameters.getLimit()).thenReturn(4); + + when(requestBuilderService.buildSearchRequest(any(), any(), anyInt(), any(Time.class), anyString())) + .thenReturn(new SearchRequestWrapper.Builder().searchRequest(searchRequest).build()); + + when(client.search(any(SearchRequest.class), eq(Object.class))) + .thenReturn(searchResponse(hits("1", "2", "3"), 5L, "scroll-0")); + + ScrollResponse secondBatch = scrollResponse(hits("4", "5"), 5L, "scroll-1"); + ScrollResponse exhausted = scrollResponse(List.of(), 5L, "scroll-1"); + AtomicInteger scrollCalls = new AtomicInteger(); + when(client.scroll(any(ScrollRequest.class), eq(Object.class))).thenAnswer(invocation -> { + int call = scrollCalls.incrementAndGet(); + if (call > SCROLL_CALL_SAFETY_LIMIT) + { + throw new AssertionError("Scroll loop did not terminate: " + call + " scroll requests issued"); + } + return call == 1 ? secondBatch : exhausted; + }); + + ArgumentCaptor collected = ArgumentCaptor.forClass(List.class); + ElasticsearchResultSet rsMock = mock(ElasticsearchResultSet.class); + when(resultSetBuilder.build(eq(searchParameters), collected.capture(), anyLong(), anyLong())).thenReturn(rsMock); + + ResultSet resultSet = strategy.executeSearch(searchParameters, queryWithPermissions); + try + { + assertNotNull(resultSet); + assertEquals(2, scrollCalls.get()); + assertTrue(collected.getValue().isEmpty()); + verify(resultSetBuilder).build(eq(searchParameters), anyList(), eq(5L), anyLong()); + } + finally + { + resultSet.close(); + } + } + + @Test + @SuppressWarnings("unchecked") + public void search_scrollReturnsNullHits_stopsScrolling() throws Exception + { + when(searchParameters.getStores()).thenReturn(new ArrayList<>()); + when(searchParameters.getSkipCount()).thenReturn(0); + when(searchParameters.getLimit()).thenReturn(4); + + when(requestBuilderService.buildSearchRequest(any(), any(), anyInt(), any(Time.class), anyString())) + .thenReturn(new SearchRequestWrapper.Builder().searchRequest(searchRequest).build()); + + when(client.search(any(SearchRequest.class), eq(Object.class))) + .thenReturn(searchResponse(hits("1", "2"), 2L, "scroll-0")); + + HitsMetadata nullHits = mock(HitsMetadata.class); + when(nullHits.hits()).thenReturn(null); + when(nullHits.total()).thenReturn(new TotalHits.Builder().value(2).relation(TotalHitsRelation.Eq).build()); + ScrollResponse nullHitsResponse = mock(ScrollResponse.class); + when(nullHitsResponse.hits()).thenReturn(nullHits); + when(nullHitsResponse.scrollId()).thenReturn("scroll-1"); + + AtomicInteger scrollCalls = new AtomicInteger(); + when(client.scroll(any(ScrollRequest.class), eq(Object.class))).thenAnswer(invocation -> { + if (scrollCalls.incrementAndGet() > SCROLL_CALL_SAFETY_LIMIT) + { + throw new AssertionError("Scroll loop did not terminate: " + scrollCalls.get() + " scroll requests issued"); + } + return nullHitsResponse; + }); + + ArgumentCaptor collected = ArgumentCaptor.forClass(List.class); + ElasticsearchResultSet rsMock = mock(ElasticsearchResultSet.class); + when(resultSetBuilder.build(eq(searchParameters), collected.capture(), anyLong(), anyLong())).thenReturn(rsMock); + + ResultSet resultSet = strategy.executeSearch(searchParameters, queryWithPermissions); + try + { + assertNotNull(resultSet); + assertEquals(1, scrollCalls.get()); + assertEquals(2, collected.getValue().size()); + } + finally + { + resultSet.close(); + } + } + + private static List> hits(String... ids) + { + List> hits = new ArrayList<>(ids.length); + for (String id : ids) + { + hits.add(new Hit.Builder<>().id(id).build()); + } + return hits; + } + + private static SearchResponse searchResponse(List> hits, long totalHits, String scrollId) + { + return new SearchResponse.Builder() + .took(1).timedOut(false) + .shards(new ShardStatistics.Builder().total(1).successful(1).skipped(0).failed(0).build()) + .hits(hb -> hb.total(new TotalHits.Builder().value(totalHits).relation(TotalHitsRelation.Eq).build()) + .hits(hits)) + .scrollId(scrollId) + .build(); + } + + private static ScrollResponse scrollResponse(List> hits, long totalHits, String scrollId) + { + return new ScrollResponse.Builder() + .took(1L).timedOut(false) + .shards(new ShardStatistics.Builder().total(1).successful(1).skipped(0).failed(0).build()) + .hits(hb -> hb.total(new TotalHits.Builder().value(totalHits).relation(TotalHitsRelation.Eq).build()) + .hits(hits)) + .scrollId(scrollId) + .build(); + } }