diff --git a/server/src/main/java/org/opensearch/action/search/StreamTransportSearchAction.java b/server/src/main/java/org/opensearch/action/search/StreamTransportSearchAction.java index 8474121115222..ae47458e4e993 100644 --- a/server/src/main/java/org/opensearch/action/search/StreamTransportSearchAction.java +++ b/server/src/main/java/org/opensearch/action/search/StreamTransportSearchAction.java @@ -30,6 +30,7 @@ import org.opensearch.transport.StreamTransportService; import org.opensearch.transport.Transport; import org.opensearch.transport.client.node.NodeClient; +import org.opensearch.wlm.WorkloadGroupService; import java.util.Map; import java.util.Set; @@ -59,7 +60,8 @@ public StreamTransportSearchAction( SearchRequestOperationsCompositeListenerFactory searchRequestOperationsCompositeListenerFactory, Tracer tracer, TaskResourceTrackingService taskResourceTrackingService, - IndicesService indicesService + IndicesService indicesService, + WorkloadGroupService workloadGroupService ) { super( client, @@ -78,7 +80,8 @@ public StreamTransportSearchAction( searchRequestOperationsCompositeListenerFactory, tracer, taskResourceTrackingService, - indicesService + indicesService, + workloadGroupService ); } diff --git a/server/src/main/java/org/opensearch/action/search/TransportSearchAction.java b/server/src/main/java/org/opensearch/action/search/TransportSearchAction.java index 9efdb2c4cd1d9..1df60e5d3afc2 100644 --- a/server/src/main/java/org/opensearch/action/search/TransportSearchAction.java +++ b/server/src/main/java/org/opensearch/action/search/TransportSearchAction.java @@ -69,6 +69,7 @@ import org.opensearch.core.common.breaker.CircuitBreaker; import org.opensearch.core.common.io.stream.NamedWriteableRegistry; import org.opensearch.core.common.io.stream.Writeable; +import org.opensearch.core.concurrency.OpenSearchRejectedExecutionException; import org.opensearch.core.index.Index; import org.opensearch.core.index.shard.ShardId; import org.opensearch.core.indices.breaker.CircuitBreakerService; @@ -110,6 +111,7 @@ import org.opensearch.transport.client.Client; import org.opensearch.transport.client.OriginSettingClient; import org.opensearch.transport.client.node.NodeClient; +import org.opensearch.wlm.WorkloadGroupService; import org.opensearch.wlm.WorkloadGroupTask; import java.util.ArrayList; @@ -194,6 +196,8 @@ public class TransportSearchAction extends HandledTransportAction) SearchRequest::new); this.client = client; @@ -240,6 +245,7 @@ public TransportSearchAction( clusterService.getClusterSettings(), new ClusterStateFieldDomainProvider() ); + this.workloadGroupService = workloadGroupService; } private Map buildPerIndexAliasFilter( @@ -483,14 +489,24 @@ void executeRequest( originalSearchRequest, taskResourceTrackingService::getTaskResourceUsageFromThreadContext ); - searchRequestContext.getSearchRequestOperationsListener().onRequestStart(searchRequestContext); - // At this point either the QUERY_GROUP_ID header will be present in ThreadContext either via ActionFilter // or HTTP header (HTTP header will be deprecated once ActionFilter is implemented) - if (task instanceof WorkloadGroupTask) { - ((WorkloadGroupTask) task).setWorkloadGroupId(threadPool.getThreadContext()); + if (task instanceof WorkloadGroupTask workloadGroupTask) { + // Coordinator-task admission point. Runs before onRequestStart (keeps the in-flight gauge balanced) and + // before setWorkloadGroupId (so a rejected task is not counted in total_completions on task completion). + try { + workloadGroupService.rejectIfNeeded( + threadPool.getThreadContext().getHeader(WorkloadGroupTask.WORKLOAD_GROUP_ID_HEADER) + ); + } catch (OpenSearchRejectedExecutionException e) { + updatedListener.onFailure(e); + return; + } + workloadGroupTask.setWorkloadGroupId(threadPool.getThreadContext()); } + searchRequestContext.getSearchRequestOperationsListener().onRequestStart(searchRequestContext); + PipelinedRequest searchRequest; ActionListener listener; try { diff --git a/server/src/main/java/org/opensearch/wlm/WorkloadManagementTransportInterceptor.java b/server/src/main/java/org/opensearch/wlm/WorkloadManagementTransportInterceptor.java index bb52440e4db34..749b75c3800d8 100644 --- a/server/src/main/java/org/opensearch/wlm/WorkloadManagementTransportInterceptor.java +++ b/server/src/main/java/org/opensearch/wlm/WorkloadManagementTransportInterceptor.java @@ -56,9 +56,9 @@ public RequestHandler(ThreadPool threadPool, TransportRequestHandler actualHa @Override public void messageReceived(T request, TransportChannel channel, Task task) throws Exception { if (isSearchWorkloadRequest(task)) { + // Reject before setWorkloadGroupId so a rejected task is not tagged and thus not counted as a phantom completion. + workloadGroupService.rejectIfNeeded(threadPool.getThreadContext().getHeader(WorkloadGroupTask.WORKLOAD_GROUP_ID_HEADER)); ((WorkloadGroupTask) task).setWorkloadGroupId(threadPool.getThreadContext()); - final String workloadGroupId = ((WorkloadGroupTask) (task)).getWorkloadGroupId(); - workloadGroupService.rejectIfNeeded(workloadGroupId); } actualHandler.messageReceived(request, channel, task); } diff --git a/server/src/main/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListener.java b/server/src/main/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListener.java index 2ee503c52246b..91f51c41a54b3 100644 --- a/server/src/main/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListener.java +++ b/server/src/main/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListener.java @@ -42,8 +42,7 @@ public WorkloadGroupRequestOperationListener(WorkloadGroupService workloadGroupS */ @Override protected void onRequestStart(SearchRequestContext searchRequestContext) { - final String workloadGroupId = threadPool.getThreadContext().getHeader(WorkloadGroupTask.WORKLOAD_GROUP_ID_HEADER); - workloadGroupService.rejectIfNeeded(workloadGroupId); + // Do not reject here: CompositeListener swallows exceptions thrown in onRequestStart. WorkloadGroup workloadGroup = workloadGroupService.getCurrentWorkloadGroup(); applyWorkloadGroupSearchSettings(workloadGroup, searchRequestContext.getRequest()); } diff --git a/server/src/test/java/org/opensearch/action/search/TransportSearchActionTests.java b/server/src/test/java/org/opensearch/action/search/TransportSearchActionTests.java index e1f515d9f5c19..4281f7f3161be 100644 --- a/server/src/test/java/org/opensearch/action/search/TransportSearchActionTests.java +++ b/server/src/test/java/org/opensearch/action/search/TransportSearchActionTests.java @@ -66,6 +66,7 @@ import org.opensearch.core.common.Strings; import org.opensearch.core.common.io.stream.NamedWriteableRegistry; import org.opensearch.core.common.transport.TransportAddress; +import org.opensearch.core.concurrency.OpenSearchRejectedExecutionException; import org.opensearch.core.index.Index; import org.opensearch.core.index.shard.ShardId; import org.opensearch.core.indices.breaker.CircuitBreakerService; @@ -113,6 +114,8 @@ import org.opensearch.transport.TransportRequestOptions; import org.opensearch.transport.TransportService; import org.opensearch.transport.client.node.NodeClient; +import org.opensearch.wlm.WorkloadGroupService; +import org.opensearch.wlm.WorkloadGroupTask; import java.util.ArrayList; import java.util.Arrays; @@ -125,6 +128,7 @@ import java.util.concurrent.CopyOnWriteArrayList; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; import java.util.function.BiFunction; @@ -135,7 +139,11 @@ import static org.hamcrest.CoreMatchers.containsString; import static org.hamcrest.CoreMatchers.instanceOf; import static org.hamcrest.CoreMatchers.startsWith; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; public class TransportSearchActionTests extends OpenSearchTestCase { @@ -1251,7 +1259,8 @@ public void testResolveIndices() { new SearchRequestOperationsCompositeListenerFactory(), mock(Tracer.class), mock(TaskResourceTrackingService.class), - mock(IndicesService.class) + mock(IndicesService.class), + mock(WorkloadGroupService.class) ); // Actual test cases start here: @@ -1292,4 +1301,82 @@ public void testResolveIndices() { } } } + + public void testCoordinatorSearchTaskRejectedBeforeRequestStart() { + ClusterService clusterService = mock(ClusterService.class); + when(clusterService.getClusterSettings()).thenReturn( + new ClusterSettings(Settings.EMPTY, ClusterSettings.BUILT_IN_CLUSTER_SETTINGS) + ); + + WorkloadGroupService workloadGroupService = mock(WorkloadGroupService.class); + doThrow(new OpenSearchRejectedExecutionException("WorkloadGroup is already contended.")).when(workloadGroupService) + .rejectIfNeeded(anyString()); + + // onRequestStart must not run on rejection, else the in-flight gauge leaks. + AtomicBoolean requestStarted = new AtomicBoolean(false); + SearchRequestOperationsListener requestStartTracker = new SearchRequestOperationsListener() { + @Override + protected void onRequestStart(SearchRequestContext searchRequestContext) { + requestStarted.set(true); + } + }; + + TransportSearchAction action = new TransportSearchAction( + mock(NodeClient.class), + threadPool, + mock(CircuitBreakerService.class), + mock(TransportService.class), + mock(SearchService.class), + mock(SearchTransportService.class), + new SearchPhaseController(new NamedWriteableRegistry(Collections.emptyList()), (searchSourceBuilder) -> null), + clusterService, + mock(ActionFilters.class), + new IndexNameExpressionResolver(new ThreadContext(Settings.EMPTY)), + new NamedWriteableRegistry(Collections.emptyList()), + mock(SearchPipelineService.class), + mock(MetricsRegistry.class), + new SearchRequestOperationsCompositeListenerFactory(requestStartTracker), + NoopTracer.INSTANCE, + mock(TaskResourceTrackingService.class), + mock(IndicesService.class), + workloadGroupService + ); + + SearchTask task = new SearchTask(0, "transport", SearchAction.NAME, () -> "test", null, Collections.emptyMap()); + + AtomicReference failure = new AtomicReference<>(); + try (ThreadContext.StoredContext ignored = threadPool.getThreadContext().stashContext()) { + threadPool.getThreadContext().putHeader(WorkloadGroupTask.WORKLOAD_GROUP_ID_HEADER, "test-workload-group"); + action.executeRequest( + task, + new SearchRequest(), + (TransportSearchAction.SearchAsyncActionProvider) ( + searchTask, + searchRequest, + executor, + shardIterators, + timeProvider, + connectionLookup, + clusterState, + aliasFilter, + concreteIndexBoosts, + indexRoutings, + listener, + preFilter, + tp, + clusters, + searchRequestContext) -> { + return null; + }, + ActionListener.wrap(r -> fail("expected rejection"), failure::set) + ); + } + + assertThat(failure.get(), instanceOf(OpenSearchRejectedExecutionException.class)); + assertFalse("onRequestStart must not run for a rejected request", requestStarted.get()); + // A rejected task must not be tagged, otherwise onTaskCompleted would count it as a phantom completion. + assertFalse("rejected task must not have its workload group id set", task.isWorkloadGroupSet()); + // Admission reads the header directly (before setWorkloadGroupId); pin that value since the stub matches anyString(). + verify(workloadGroupService).rejectIfNeeded(eq("test-workload-group")); + } } diff --git a/server/src/test/java/org/opensearch/snapshots/SnapshotResiliencyTests.java b/server/src/test/java/org/opensearch/snapshots/SnapshotResiliencyTests.java index 95e7ae5384f49..870f8b70f50f2 100644 --- a/server/src/test/java/org/opensearch/snapshots/SnapshotResiliencyTests.java +++ b/server/src/test/java/org/opensearch/snapshots/SnapshotResiliencyTests.java @@ -251,6 +251,7 @@ import org.opensearch.transport.TransportService; import org.opensearch.transport.client.AdminClient; import org.opensearch.transport.client.node.NodeClient; +import org.opensearch.wlm.WorkloadGroupService; import org.junit.After; import org.junit.Before; @@ -2410,7 +2411,8 @@ public void onFailure(final Exception e) { searchRequestOperationsCompositeListenerFactory, NoopTracer.INSTANCE, new TaskResourceTrackingService(settings, clusterSettings, threadPool), - mockIndicesService + mockIndicesService, + mock(WorkloadGroupService.class) ) ); actions.put( diff --git a/server/src/test/java/org/opensearch/wlm/WorkloadManagementTransportRequestHandlerTests.java b/server/src/test/java/org/opensearch/wlm/WorkloadManagementTransportRequestHandlerTests.java index e05aaf941c4e9..26238280bde14 100644 --- a/server/src/test/java/org/opensearch/wlm/WorkloadManagementTransportRequestHandlerTests.java +++ b/server/src/test/java/org/opensearch/wlm/WorkloadManagementTransportRequestHandlerTests.java @@ -21,7 +21,7 @@ import java.util.Collections; -import static org.mockito.Mockito.anyString; +import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.doNothing; import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; @@ -51,17 +51,23 @@ public void tearDown() throws Exception { public void testMessageReceivedForSearchWorkload_nonRejectionCase() throws Exception { ShardSearchRequest request = mock(ShardSearchRequest.class); WorkloadGroupTask spyTask = getSpyTask(); - doNothing().when(workloadGroupService).rejectIfNeeded(anyString()); + doNothing().when(workloadGroupService).rejectIfNeeded(any()); sut.messageReceived(request, mock(TransportChannel.class), spyTask); assertTrue(sut.isSearchWorkloadRequest(spyTask)); + // Admitted task is tagged and forwarded to the wrapped handler. + assertTrue(spyTask.isWorkloadGroupSet()); + assertEquals(1, actualHandler.invokeCount); } public void testMessageReceivedForSearchWorkload_RejectionCase() throws Exception { ShardSearchRequest request = mock(ShardSearchRequest.class); WorkloadGroupTask spyTask = getSpyTask(); - doThrow(OpenSearchRejectedExecutionException.class).when(workloadGroupService).rejectIfNeeded(anyString()); + doThrow(OpenSearchRejectedExecutionException.class).when(workloadGroupService).rejectIfNeeded(any()); assertThrows(OpenSearchRejectedExecutionException.class, () -> sut.messageReceived(request, mock(TransportChannel.class), spyTask)); + // A rejected task must not be tagged (else onTaskCompleted would count it as a phantom completion) and must not be forwarded. + assertFalse(spyTask.isWorkloadGroupSet()); + assertEquals(0, actualHandler.invokeCount); } public void testMessageReceivedForNonSearchWorkload() throws Exception { diff --git a/server/src/test/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListenerTests.java b/server/src/test/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListenerTests.java index 31071d7acf1c3..18d56c11aafee 100644 --- a/server/src/test/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListenerTests.java +++ b/server/src/test/java/org/opensearch/wlm/listeners/WorkloadGroupRequestOperationListenerTests.java @@ -17,7 +17,6 @@ import org.opensearch.common.settings.Settings; import org.opensearch.common.unit.TimeValue; import org.opensearch.common.util.concurrent.ThreadContext; -import org.opensearch.core.concurrency.OpenSearchRejectedExecutionException; import org.opensearch.search.builder.SearchSourceBuilder; import org.opensearch.test.OpenSearchTestCase; import org.opensearch.threadpool.TestThreadPool; @@ -40,8 +39,6 @@ import java.util.List; import java.util.Map; -import static org.mockito.Mockito.doNothing; -import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -82,21 +79,6 @@ public void tearDown() throws Exception { testThreadPool.shutdown(); } - public void testRejectionCase() { - final String testWorkloadGroupId = "asdgasgkajgkw3141_3rt4t"; - testThreadPool.getThreadContext().putHeader(WorkloadGroupTask.WORKLOAD_GROUP_ID_HEADER, testWorkloadGroupId); - doThrow(OpenSearchRejectedExecutionException.class).when(workloadGroupService).rejectIfNeeded(testWorkloadGroupId); - assertThrows(OpenSearchRejectedExecutionException.class, () -> sut.onRequestStart(mockSearchRequestContext)); - } - - public void testNonRejectionCase() { - final String testWorkloadGroupId = "asdgasgkajgkw3141_3rt4t"; - testThreadPool.getThreadContext().putHeader(WorkloadGroupTask.WORKLOAD_GROUP_ID_HEADER, testWorkloadGroupId); - doNothing().when(workloadGroupService).rejectIfNeeded(testWorkloadGroupId); - - sut.onRequestStart(mockSearchRequestContext); - } - public void testValidWorkloadGroupRequestFailure() throws IOException { WorkloadGroupStats expectedStats = new WorkloadGroupStats(