diff --git a/trpc-proto/trpc-proto-http/src/main/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutor.java b/trpc-proto/trpc-proto-http/src/main/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutor.java index 73eeb11f9..2ab588049 100644 --- a/trpc-proto/trpc-proto-http/src/main/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutor.java +++ b/trpc-proto/trpc-proto-http/src/main/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutor.java @@ -38,11 +38,15 @@ import com.tencent.trpc.proto.http.common.RpcServerContextWithHttp; import com.tencent.trpc.proto.http.common.TrpcServletRequestWrapper; import com.tencent.trpc.proto.http.common.TrpcServletResponseWrapper; +import java.io.IOException; import java.lang.reflect.Type; import java.net.InetSocketAddress; import java.nio.charset.StandardCharsets; import java.util.Enumeration; import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.CompletionStage; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; @@ -65,15 +69,15 @@ public abstract class AbstractHttpExecutor { protected void execute(HttpServletRequest request, HttpServletResponse response, RpcMethodInfoAndInvoker methodInfoAndInvoker) { - + AtomicBoolean responded = new AtomicBoolean(false); try { DefRequest rpcRequest = buildDefRequest(request, response, methodInfoAndInvoker); - CountDownLatch countDownLatch = new CountDownLatch(1); + CompletableFuture completionFuture = new CompletableFuture<>(); // use a thread pool for asynchronous processing - invokeRpcRequest(methodInfoAndInvoker.getInvoker(), rpcRequest, countDownLatch); + invokeRpcRequest(methodInfoAndInvoker.getInvoker(), rpcRequest, completionFuture, responded); // If the request carries a timeout, use this timeout to wait for the request to be processed. // If not carried, use the default timeout. @@ -81,18 +85,25 @@ protected void execute(HttpServletRequest request, HttpServletResponse response, if (requestTimeout <= 0) { requestTimeout = methodInfoAndInvoker.getInvoker().getConfig().getRequestTimeout(); } - if (requestTimeout > 0 && !countDownLatch.await(requestTimeout, TimeUnit.MILLISECONDS)) { - throw TRpcException.newFrameException(ErrorCode.TRPC_SERVER_TIMEOUT_ERR, - "wait http request execute timeout"); + if (requestTimeout > 0) { + try { + completionFuture.get(requestTimeout, TimeUnit.MILLISECONDS); + } catch (TimeoutException ex) { + if (responded.compareAndSet(false, true)) { + doErrorReply(request, response, + TRpcException.newFrameException(ErrorCode.TRPC_SERVER_TIMEOUT_ERR, + "wait http request execute timeout")); + } + } } else { - countDownLatch.await(); + completionFuture.get(); } - } catch (Exception ex) { logger.error("dispatch request [{}] error", request, ex); - doErrorReply(request, response, ex); + if (responded.compareAndSet(false, true)) { + doErrorReply(request, response, ex); + } } - } /** @@ -107,55 +118,83 @@ protected void execute(HttpServletRequest request, HttpServletResponse response, /** * Request processing * - * @param countDownLatch latch used to wait for the request processing + * @param invoker the invoker + * @param rpcRequest the rpc request + * @param completionFuture the completion future + * @param responded the responded flag */ - private void invokeRpcRequest(ProviderInvoker invoker, DefRequest rpcRequest, CountDownLatch countDownLatch) { + private void invokeRpcRequest(ProviderInvoker invoker, DefRequest rpcRequest, + CompletableFuture completionFuture, + AtomicBoolean responded) { WorkerPool workerPool = invoker.getConfig().getWorkerPoolObj(); if (null == workerPool) { logger.error("dispatch rpcRequest [{}] error, workerPool is empty", rpcRequest); - throw TRpcException.newFrameException(ErrorCode.TRPC_SERVER_NOSERVICE_ERR, - "not found service, workerPool is empty"); + completionFuture.completeExceptionally(TRpcException.newFrameException(ErrorCode.TRPC_SERVER_NOSERVICE_ERR, + "not found service, workerPool is empty")); + return; } workerPool.execute(() -> { - - // Get the original http response - HttpServletResponse response = getOriginalResponse(rpcRequest); - - // Invoke the routing implementation method to handle the request. - CompletionStage future = invoker.invoke(rpcRequest); - future.whenComplete((result, t) -> { - try { - // Throw the call exception, which will be handled uniformly by the exception handling program. - if (t != null) { - throw t; - } - - // Throw a business logic exception, which will be handled uniformly - // by the exception handling program. - Throwable ex = result.getException(); - if (ex != null) { - throw ex; + try { + // Get the original http response + HttpServletResponse response = getOriginalResponse(rpcRequest); + // Invoke the routing implementation method to handle the request. + CompletionStage rpcFuture = invoker.invoke(rpcRequest); + + rpcFuture.whenComplete((result, throwable) -> { + try { + if (responded.get()) { + return; + } + + // Throw the call exception, which will be handled uniformly by the exception handling program. + if (throwable != null) { + throw throwable; + } + + // Throw a business logic exception, which will be handled uniformly + // by the exception handling program. + if (result.getException() != null) { + throw result.getException(); + } + + // normal response + if (responded.compareAndSet(false, true)) { + response.setStatus(HttpStatus.SC_OK); + httpCodec.writeHttpResponse(response, result); + response.flushBuffer(); + } + + completionFuture.complete(null); + } catch (Throwable t) { + handleError(t, rpcRequest, response, responded, completionFuture); } + }); - // normal response - response.setStatus(HttpStatus.SC_OK); - httpCodec.writeHttpResponse(response, result); - response.flushBuffer(); - } catch (Throwable e) { - HttpServletRequest request = getOriginalRequest(rpcRequest); - logger.warn("reply message error, channel: [{}], msg:[{}]", request.getRemoteAddr(), request, e); - httpErrorReply(request, response, - ErrorResponse.create(request, HttpStatus.SC_SERVICE_UNAVAILABLE, e)); - } finally { - countDownLatch.countDown(); - } - }); + } catch (Exception e) { + handleError(e, rpcRequest, getOriginalResponse(rpcRequest), responded, completionFuture); + } }); } + /** + * Handle error + */ + private void handleError(Throwable t, DefRequest rpcRequest, HttpServletResponse response, + AtomicBoolean responded, CompletableFuture completionFuture) { + try { + if (responded.compareAndSet(false, true)) { + HttpServletRequest request = getOriginalRequest(rpcRequest); + logger.warn("reply message error, channel: [{}], msg:[{}]", request.getRemoteAddr(), request, t); + httpErrorReply(request, response, ErrorResponse.create(request, HttpStatus.SC_SERVICE_UNAVAILABLE, t)); + } + } finally { + completionFuture.completeExceptionally(t); + } + } + /** * Build the context request. * @@ -480,4 +519,4 @@ private String getString(String[] callInfos, int length, int cursor) { return callInfos.length < length ? StringUtils.EMPTY : callInfos[cursor]; } -} +} \ No newline at end of file diff --git a/trpc-proto/trpc-proto-http/src/test/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutorTest.java b/trpc-proto/trpc-proto-http/src/test/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutorTest.java index 2abbe5643..acce946be 100644 --- a/trpc-proto/trpc-proto-http/src/test/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutorTest.java +++ b/trpc-proto/trpc-proto-http/src/test/java/com/tencent/trpc/proto/http/server/AbstractHttpExecutorTest.java @@ -9,18 +9,37 @@ * A copy of the Apache 2.0 License can be found in the LICENSE file. */ + package com.tencent.trpc.proto.http.server; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; +import static org.mockito.Matchers.any; +import static org.mockito.Mockito.verify; +import static org.powermock.api.mockito.PowerMockito.doAnswer; +import static org.powermock.api.mockito.PowerMockito.doCallRealMethod; import static org.powermock.api.mockito.PowerMockito.doReturn; import static org.powermock.api.mockito.PowerMockito.mock; import static org.powermock.api.mockito.PowerMockito.when; +import com.tencent.trpc.core.common.config.ProviderConfig; +import com.tencent.trpc.core.rpc.ProviderInvoker; import com.tencent.trpc.core.rpc.RpcInvocation; +import com.tencent.trpc.core.rpc.Response; import com.tencent.trpc.core.rpc.common.RpcMethodInfo; +import com.tencent.trpc.core.rpc.common.RpcMethodInfoAndInvoker; +import com.tencent.trpc.core.rpc.def.DefRequest; +import com.tencent.trpc.core.rpc.def.DefResponse; +import com.tencent.trpc.core.worker.spi.WorkerPool; +import com.tencent.trpc.core.worker.spi.WorkerPool.Task; +import com.tencent.trpc.proto.http.common.HttpCodec; import com.tencent.trpc.proto.http.common.HttpConstants; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.atomic.AtomicBoolean; import javax.servlet.http.HttpServletRequest; +import javax.servlet.http.HttpServletResponse; +import org.apache.http.HttpStatus; import org.junit.Test; import org.junit.runner.RunWith; import org.powermock.core.classloader.annotations.PowerMockIgnore; @@ -34,18 +53,210 @@ @PrepareForTest(AbstractHttpExecutor.class) public class AbstractHttpExecutorTest { + private static final String TEST_SERVICE = "trpc.demo.server"; + private static final String TEST_METHOD = "hello"; + private static final String TEST_IP = "127.0.0.1"; + private static final int TEST_PORT = 8080; + + private HttpServletRequest mockRequest() { + HttpServletRequest request = mock(HttpServletRequest.class); + when(request.getAttribute(HttpConstants.REQUEST_ATTRIBUTE_TRPC_SERVICE)).thenReturn(TEST_SERVICE); + when(request.getAttribute(HttpConstants.REQUEST_ATTRIBUTE_TRPC_METHOD)).thenReturn(TEST_METHOD); + when(request.getRemoteAddr()).thenReturn(TEST_IP); + when(request.getRemotePort()).thenReturn(TEST_PORT); + return request; + } + + private WorkerPool mockSyncWorkerPool() { + WorkerPool workerPool = mock(WorkerPool.class); + doAnswer(invocation -> { + Object arg = invocation.getArguments()[0]; + if (arg instanceof Runnable) { + ((Runnable) arg).run(); + } else if (arg instanceof Task) { + ((Task) arg).run(); + } + return null; + }).when(workerPool).execute(any()); + return workerPool; + } + + private ProviderConfig mockProviderConfig(int timeout) { + ProviderConfig config = mock(ProviderConfig.class); + when(config.getRequestTimeout()).thenReturn(timeout); + WorkerPool workerPool = mockSyncWorkerPool(); + when(config.getWorkerPoolObj()).thenReturn(workerPool); + return config; + } + + private DefRequest mockDefRequest(HttpServletRequest request, HttpServletResponse response) { + DefRequest defRequest = new DefRequest(); + defRequest.getAttachments().put(HttpConstants.TRPC_ATTACH_SERVLET_RESPONSE, response); + defRequest.getAttachments().put(HttpConstants.TRPC_ATTACH_SERVLET_REQUEST, request); + return defRequest; + } + + private AbstractHttpExecutor mockExecutorWithCodec() { + AbstractHttpExecutor executor = mock(AbstractHttpExecutor.class); + HttpCodec httpCodec = mock(HttpCodec.class); + Whitebox.setInternalState(executor, "httpCodec", httpCodec); + return executor; + } @Test - public void buildRpcInvocation_shouldSuccess() throws Exception { + public void testBuildRpcInvocation() throws Exception { HttpServletRequest request = mock(HttpServletRequest.class); + doReturn(TEST_SERVICE).when(request).getAttribute(HttpConstants.REQUEST_ATTRIBUTE_TRPC_SERVICE); + doReturn(TEST_METHOD).when(request).getAttribute(HttpConstants.REQUEST_ATTRIBUTE_TRPC_METHOD); + + RpcMethodInfo methodInfo = mock(RpcMethodInfo.class); + AbstractHttpExecutor executor = mock(AbstractHttpExecutor.class); + doReturn(null).when(executor, "parseRpcParams", request, methodInfo); + when(executor, "buildRpcInvocation", request, methodInfo).thenCallRealMethod(); + RpcInvocation invocation = Whitebox.invokeMethod(executor, "buildRpcInvocation", request, methodInfo); + + assertEquals("/trpc.demo.server/hello", invocation.getFunc()); + } + + @Test + public void testExecuteSuccess() throws Exception { + ProviderInvoker invoker = mock(ProviderInvoker.class); + DefResponse successResponse = new DefResponse(); + successResponse.setValue("success"); + CompletableFuture successFuture = CompletableFuture.completedFuture(successResponse); + when(invoker.invoke(any())).thenReturn(successFuture); + ProviderConfig config = mockProviderConfig(0); + when(invoker.getConfig()).thenReturn(config); + RpcMethodInfo methodInfo = mock(RpcMethodInfo.class); + RpcMethodInfoAndInvoker methodInfoAndInvoker = mock(RpcMethodInfoAndInvoker.class); + when(methodInfoAndInvoker.getMethodInfo()).thenReturn(methodInfo); + doReturn(invoker).when(methodInfoAndInvoker, "getInvoker"); + HttpServletRequest request = mockRequest(); + HttpServletResponse response = mock(HttpServletResponse.class); + DefRequest defRequest = mockDefRequest(request, response); + AbstractHttpExecutor executor = mockExecutorWithCodec(); + doReturn(defRequest).when(executor, "buildDefRequest", any(), any(), any()); + doReturn(response).when(executor, "getOriginalResponse", any()); + doCallRealMethod().when(executor, "execute", request, response, methodInfoAndInvoker); + doCallRealMethod().when(executor, "invokeRpcRequest", any(), any(), any(), any()); + Whitebox.invokeMethod(executor, "execute", request, response, methodInfoAndInvoker); + verify(response).setStatus(HttpStatus.SC_OK); + HttpCodec httpCodec = Whitebox.getInternalState(executor, "httpCodec"); + verify(httpCodec).writeHttpResponse(response, successResponse); + verify(response).flushBuffer(); + } + + @Test + public void testExecuteTimeout() throws Exception { + ProviderInvoker invoker = mock(ProviderInvoker.class); + CompletableFuture neverCompleteFuture = new CompletableFuture<>(); + when(invoker.invoke(any())).thenReturn(neverCompleteFuture); + ProviderConfig config = mockProviderConfig(100); + when(invoker.getConfig()).thenReturn(config); RpcMethodInfo methodInfo = mock(RpcMethodInfo.class); - AbstractHttpExecutor abstractHttpExecutor = mock(AbstractHttpExecutor.class); - doReturn(null).when(abstractHttpExecutor, "parseRpcParams", request, methodInfo); - doReturn("trpc.demo.server").when(request).getAttribute(HttpConstants.REQUEST_ATTRIBUTE_TRPC_SERVICE); - doReturn("hello").when(request).getAttribute(HttpConstants.REQUEST_ATTRIBUTE_TRPC_METHOD); - when(abstractHttpExecutor, "buildRpcInvocation", request, methodInfo).thenCallRealMethod(); - RpcInvocation rpcInvocation = Whitebox.invokeMethod(abstractHttpExecutor, "buildRpcInvocation", request, - methodInfo); - assertEquals(rpcInvocation.getFunc(), "/trpc.demo.server/hello"); + RpcMethodInfoAndInvoker methodInfoAndInvoker = mock(RpcMethodInfoAndInvoker.class); + when(methodInfoAndInvoker.getMethodInfo()).thenReturn(methodInfo); + doReturn(invoker).when(methodInfoAndInvoker, "getInvoker"); + HttpServletRequest request = mockRequest(); + HttpServletResponse response = mock(HttpServletResponse.class); + DefRequest defRequest = mockDefRequest(request, response); + AbstractHttpExecutor executor = mockExecutorWithCodec(); + doReturn(null).when(executor, "parseRpcParams", any(), any()); + doReturn(defRequest).when(executor, "buildDefRequest", any(), any(), any()); + when(executor, "execute", request, response, methodInfoAndInvoker).thenCallRealMethod(); + doCallRealMethod().when(executor, "doErrorReply", any(), any(), any()); + doCallRealMethod().when(executor, "httpErrorReply", any(), any(), any()); + when(executor, "invokeRpcRequest", any(), any(), any(), any()).thenCallRealMethod(); + Whitebox.invokeMethod(executor, "execute", request, response, methodInfoAndInvoker); + verify(response).setStatus(HttpStatus.SC_REQUEST_TIMEOUT); + } + + @Test + public void testHandleError() throws Exception { + HttpServletRequest request = mock(HttpServletRequest.class); + when(request.getRemoteAddr()).thenReturn(TEST_IP); + when(request.getMethod()).thenReturn("POST"); + when(request.getRequestURI()).thenReturn("/api/test"); + when(request.getQueryString()).thenReturn("param=value"); + HttpServletResponse response = mock(HttpServletResponse.class); + DefRequest defRequest = mockDefRequest(request, response); + AbstractHttpExecutor executor = mockExecutorWithCodec(); + doCallRealMethod().when(executor, "handleError", any(Throwable.class), any(DefRequest.class), + any(HttpServletResponse.class), any(AtomicBoolean.class), any(CompletableFuture.class)); + doCallRealMethod().when(executor, "httpErrorReply", any(), any(), any()); + doReturn(request).when(executor, "getOriginalRequest", any()); + AtomicBoolean responded = new AtomicBoolean(false); + CompletableFuture completionFuture = new CompletableFuture<>(); + Throwable testException = new RuntimeException("Test error"); + Whitebox.invokeMethod(executor, "handleError", testException, defRequest, response, + responded, completionFuture); + assertTrue(responded.get()); + assertTrue(completionFuture.isCompletedExceptionally()); + verify(response).setStatus(HttpStatus.SC_SERVICE_UNAVAILABLE); + HttpCodec httpCodec = Whitebox.getInternalState(executor, "httpCodec"); + verify(httpCodec).writeHttpResponse(any(HttpServletResponse.class), any()); + } + + @Test + public void testInvokeRpcWithException() throws Exception { + + + ProviderConfig config = mockProviderConfig(0); + ProviderInvoker invoker = mock(ProviderInvoker.class); + when(invoker.getConfig()).thenReturn(config); + CompletableFuture failedFuture = new CompletableFuture<>(); + failedFuture.completeExceptionally(new RuntimeException("boom")); + when(invoker.invoke(any())).thenReturn(failedFuture); + + AbstractHttpExecutor executor = mockExecutorWithCodec(); + HttpServletRequest request = mock(HttpServletRequest.class); + HttpServletResponse response = mock(HttpServletResponse.class); + doReturn(response).when(executor, "getOriginalResponse", any()); + doReturn(request).when(executor, "getOriginalRequest", any()); + doCallRealMethod().when(executor, "invokeRpcRequest", any(), any(), any(), any()); + doCallRealMethod().when(executor, "httpErrorReply", any(), any(), any()); + doCallRealMethod().when(executor, "handleError", any(Throwable.class), any(DefRequest.class), + any(HttpServletResponse.class), any(AtomicBoolean.class), any(CompletableFuture.class)); + + AtomicBoolean responded = new AtomicBoolean(false); + CompletableFuture completionFuture = new CompletableFuture<>(); + + DefRequest defRequest = mockDefRequest(request, response); + Whitebox.invokeMethod(executor, "invokeRpcRequest", invoker, defRequest, completionFuture, responded); + assertTrue(responded.get()); + assertTrue(completionFuture.isCompletedExceptionally()); + verify(response).setStatus(HttpStatus.SC_SERVICE_UNAVAILABLE); + HttpCodec httpCodec = Whitebox.getInternalState(executor, "httpCodec"); + verify(httpCodec).writeHttpResponse(any(HttpServletResponse.class), any()); + } + + @Test + public void testInvokeRpcThrowsDirectly() throws Exception { + HttpServletRequest request = mock(HttpServletRequest.class); + HttpServletResponse response = mock(HttpServletResponse.class); + DefRequest defRequest = mockDefRequest(request, response); + + ProviderConfig config = mockProviderConfig(0); + ProviderInvoker invoker = mock(ProviderInvoker.class); + when(invoker.getConfig()).thenReturn(config); + when(invoker.invoke(any())).thenThrow(new RuntimeException("boom-direct")); + + AbstractHttpExecutor executor = mockExecutorWithCodec(); + doReturn(response).when(executor, "getOriginalResponse", any()); + doReturn(request).when(executor, "getOriginalRequest", any()); + doCallRealMethod().when(executor, "invokeRpcRequest", any(), any(), any(), any()); + doCallRealMethod().when(executor, "httpErrorReply", any(), any(), any()); + doCallRealMethod().when(executor, "handleError", any(Throwable.class), any(DefRequest.class), + any(HttpServletResponse.class), any(AtomicBoolean.class), any(CompletableFuture.class)); + + AtomicBoolean responded = new AtomicBoolean(false); + CompletableFuture completionFuture = new CompletableFuture<>(); + Whitebox.invokeMethod(executor, "invokeRpcRequest", invoker, defRequest, completionFuture, responded); + + assertTrue(responded.get()); + assertTrue(completionFuture.isCompletedExceptionally()); + verify(response).setStatus(HttpStatus.SC_SERVICE_UNAVAILABLE); + HttpCodec httpCodec = Whitebox.getInternalState(executor, "httpCodec"); + verify(httpCodec).writeHttpResponse(any(HttpServletResponse.class), any()); } } \ No newline at end of file