diff --git a/rpc-consumer/src/test/java/com/xiaoyu/rpc/consumer/FullIntegrationTest.java b/rpc-consumer/src/test/java/com/xiaoyu/rpc/consumer/FullIntegrationTest.java index 7994f4e..0267aa8 100644 --- a/rpc-consumer/src/test/java/com/xiaoyu/rpc/consumer/FullIntegrationTest.java +++ b/rpc-consumer/src/test/java/com/xiaoyu/rpc/consumer/FullIntegrationTest.java @@ -17,6 +17,7 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.junit.jupiter.api.Assertions.fail; @@ -45,9 +46,11 @@ public void testFullIntegration() throws Exception { System.setProperty("rpc.protocol", protocol); AtomicReference serverFailure = new AtomicReference<>(); + AtomicReference serverRef = new AtomicReference<>(); Thread serverThread = new Thread(() -> { try { RpcServer server = new RpcServer(); + serverRef.set(server); server.register(HelloService.class, new HelloServiceImpl()); server.start(); } catch (Throwable e) { @@ -57,25 +60,32 @@ public void testFullIntegration() throws Exception { serverThread.setDaemon(true); serverThread.start(); - waitForServer(port, serverFailure); - try { - RpcClient rpcClient = new RpcClient(); - Serializer serializer = SerializerCode.getSerializerByCode(RpcConfig.getInstance().getSerializerCode()); + waitForServer(port, serverFailure); + + try (RpcClient rpcClient = new RpcClient()) { + Serializer serializer = SerializerCode.getSerializerByCode(RpcConfig.getInstance().getSerializerCode()); - String result1 = (String) rpcClient - .sendRequest(buildRequest("World1", serializer), String.class) - .get(5, TimeUnit.SECONDS); + String result1 = (String) rpcClient + .sendRequest(buildRequest("World1", serializer), String.class) + .get(5, TimeUnit.SECONDS); - String result2 = (String) rpcClient - .sendRequest(buildRequest("World2", serializer), String.class) - .get(5, TimeUnit.SECONDS); + String result2 = (String) rpcClient + .sendRequest(buildRequest("World2", serializer), String.class) + .get(5, TimeUnit.SECONDS); - assertNotNull(result1, "Result1 should not be null"); - assertNotNull(result2, "Result2 should not be null"); - assertTrue(result1.contains("World1"), "Result1 should contain World1"); - assertTrue(result2.contains("World2"), "Result2 should contain World2"); + assertNotNull(result1, "Result1 should not be null"); + assertNotNull(result2, "Result2 should not be null"); + assertTrue(result1.contains("World1"), "Result1 should contain World1"); + assertTrue(result2.contains("World2"), "Result2 should contain World2"); + } } finally { + RpcServer server = serverRef.get(); + if (server != null) { + server.close(); + } + serverThread.join(5000); + assertFalse(serverThread.isAlive(), "RPC server thread should stop after RpcServer.close()"); System.clearProperty("rpc.server-host"); System.clearProperty("rpc.server-port"); } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ByteBuddyProxyFactory.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ByteBuddyProxyFactory.java index 7a9fcbd..a0f62d0 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ByteBuddyProxyFactory.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ByteBuddyProxyFactory.java @@ -1,9 +1,9 @@ package com.xiaoyu.rpc.core.client; +import com.google.protobuf.ByteString; import com.xiaoyu.rpc.common.serialization.Serializer; import com.xiaoyu.rpc.common.serialization.SerializerCode; import com.xiaoyu.rpc.common.vo.RpcRequest; -import com.google.protobuf.ByteString; import com.xiaoyu.rpc.core.config.RpcConfig; import net.bytebuddy.ByteBuddy; import net.bytebuddy.implementation.InvocationHandlerAdapter; @@ -17,6 +17,7 @@ public class ByteBuddyProxyFactory implements ProxyFactory { private volatile RpcClient rpcClient; + private volatile boolean closed; public ByteBuddyProxyFactory() { // 与 JDK Proxy 一致:SPI 扩展加载阶段不初始化注册中心和传输层。 @@ -27,9 +28,16 @@ public ByteBuddyProxyFactory() { } private RpcClient getRpcClient() { + if (closed) { + throw new IllegalStateException("ByteBuddyProxyFactory 已关闭"); + } + RpcClient client = rpcClient; if (client == null) { synchronized (this) { + if (closed) { + throw new IllegalStateException("ByteBuddyProxyFactory 已关闭"); + } client = rpcClient; if (client == null) { client = new RpcClient(); @@ -48,20 +56,16 @@ public T getProxy(Class clazz) { .intercept(InvocationHandlerAdapter.of(new InvocationHandler() { @Override public Object invoke(Object proxy, Method method, Object[] args) throws Throwable { - // 请求中记录接口名 + 方法名,服务端靠这两项定位目标方法 RpcRequest.Builder builder = RpcRequest.newBuilder() .setInterfaceName(method.getDeclaringClass().getName()) .setMethodName(method.getName()); Class[] parameterTypes = method.getParameterTypes(); - if (parameterTypes != null) { - for (Class paramType : parameterTypes) { - builder.addParamTypes(paramType.getName()); - } + for (Class paramType : parameterTypes) { + builder.addParamTypes(paramType.getName()); } if (args != null) { - // 获取配置的序列化器 Serializer serializer = SerializerCode .getSerializerByCode(RpcConfig.getInstance().getSerializerCode()); for (Object arg : args) { @@ -71,9 +75,11 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl } RpcRequest request = builder.build(); - CompletableFuture future = getRpcClient().sendRequest(request, method.getReturnType()); - // 如果业务接口声明的返回类型是异步的,直接返回 Future;否则阻塞等待结果 - if (CompletableFuture.class.isAssignableFrom(method.getReturnType())) { + boolean async = RpcReturnTypeResolver.isAsync(method); + Class payloadType = RpcReturnTypeResolver.resolvePayloadType(method); + CompletableFuture future = getRpcClient().sendRequest(request, payloadType); + + if (async) { return future; } return future.get(); @@ -83,4 +89,21 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl throw new RuntimeException("ByteBuddy代理创建失败", e); } } + + @Override + public void close() { + RpcClient client; + synchronized (this) { + if (closed) { + return; + } + closed = true; + client = rpcClient; + rpcClient = null; + } + + if (client != null) { + client.close(); + } + } } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/JdkProxyFactory.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/JdkProxyFactory.java index 33d27c9..72365bf 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/JdkProxyFactory.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/JdkProxyFactory.java @@ -15,6 +15,7 @@ public class JdkProxyFactory implements ProxyFactory { private volatile RpcClient rpcClient; + private volatile boolean closed; public JdkProxyFactory() { // SPI 扩展加载阶段保持轻量,不在构造时初始化注册中心和传输层。 @@ -25,9 +26,16 @@ public JdkProxyFactory() { } private RpcClient getRpcClient() { + if (closed) { + throw new IllegalStateException("JdkProxyFactory 已关闭"); + } + RpcClient client = rpcClient; if (client == null) { synchronized (this) { + if (closed) { + throw new IllegalStateException("JdkProxyFactory 已关闭"); + } client = rpcClient; if (client == null) { client = new RpcClient(); @@ -52,10 +60,8 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl .setMethodName(method.getName()); Class[] parameterTypes = method.getParameterTypes(); - if (parameterTypes != null) { - for (Class paramType : parameterTypes) { - builder.addParamTypes(paramType.getName()); - } + for (Class paramType : parameterTypes) { + builder.addParamTypes(paramType.getName()); } if (args != null) { @@ -68,13 +74,32 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl } RpcRequest request = builder.build(); - CompletableFuture future = getRpcClient().sendRequest(request, method.getReturnType()); - // 如果业务接口声明的返回类型是异步的,直接返回 Future;否则阻塞等待结果 - if (CompletableFuture.class.isAssignableFrom(method.getReturnType())) { + boolean async = RpcReturnTypeResolver.isAsync(method); + Class payloadType = RpcReturnTypeResolver.resolvePayloadType(method); + CompletableFuture future = getRpcClient().sendRequest(request, payloadType); + + if (async) { return future; } return future.get(); } }); } + + @Override + public void close() { + RpcClient client; + synchronized (this) { + if (closed) { + return; + } + closed = true; + client = rpcClient; + rpcClient = null; + } + + if (client != null) { + client.close(); + } + } } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ProxyFactory.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ProxyFactory.java index 8bb4a1b..40a7abf 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ProxyFactory.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/ProxyFactory.java @@ -3,6 +3,11 @@ import com.xiaoyu.rpc.common.extension.SPI; @SPI -public interface ProxyFactory { +public interface ProxyFactory extends AutoCloseable { T getProxy(Class clazz); + + @Override + default void close() { + // 默认无资源需要释放;具体代理实现可覆盖。 + } } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClient.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClient.java index acc93fb..686c258 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClient.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClient.java @@ -14,19 +14,19 @@ import java.net.InetSocketAddress; import java.util.Objects; import java.util.concurrent.CompletableFuture; +import java.util.concurrent.atomic.AtomicBoolean; -public class RpcClient { +public class RpcClient implements AutoCloseable { private final TransportClient transportClient; private final ServiceDiscovery serviceDiscovery; + private final AtomicBoolean closed = new AtomicBoolean(false); public RpcClient() { RpcConfig config = RpcConfig.getInstance(); - // 初始化服务发现 this.serviceDiscovery = ExtensionLoader.getExtensionLoader(ServiceDiscovery.class) .getExtension(config.getRegistryType()); - // 初始化传输层客户端 Transport transport = ExtensionLoader.getExtensionLoader(Transport.class).getExtension(config.getTransport()); this.transportClient = transport.createClient(); } @@ -37,23 +37,24 @@ public RpcClient() { } public CompletableFuture sendRequest(RpcRequest request, Class returnType) { + if (closed.get()) { + return CompletableFuture.failedFuture(new IllegalStateException("RpcClient 已关闭")); + } + try { - // 先做一次服务发现(同步查找,通常会命中本地缓存) InetSocketAddress address = serviceDiscovery.lookupService(request.getInterfaceName()); if (address == null) { - CompletableFuture future = new CompletableFuture<>(); - future.completeExceptionally(new RuntimeException("未发现服务: " + request.getInterfaceName())); - return future; + return CompletableFuture.failedFuture( + new RuntimeException("未发现服务: " + request.getInterfaceName())); } - // 交给传输层发送,返回异步 Future CompletableFuture transportFuture = transportClient.sendRequest(request, address); - // 在回调里把响应体反序列化成目标返回类型 return transportFuture.thenApply(result -> { if (!(result instanceof RpcResponse)) { - throw new RuntimeException("Unexpected response type: " + result.getClass()); + String actualType = result == null ? "null" : result.getClass().getName(); + throw new RuntimeException("Unexpected response type: " + actualType); } RpcResponse response = (RpcResponse) result; @@ -69,9 +70,14 @@ public CompletableFuture sendRequest(RpcRequest request, Class return }); } catch (Exception e) { - CompletableFuture future = new CompletableFuture<>(); - future.completeExceptionally(e); - return future; + return CompletableFuture.failedFuture(e); + } + } + + @Override + public void close() { + if (closed.compareAndSet(false, true)) { + transportClient.close(); } } } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClientProxy.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClientProxy.java index 6bf87db..94893cf 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClientProxy.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClientProxy.java @@ -1,14 +1,28 @@ package com.xiaoyu.rpc.core.client; -import com.xiaoyu.rpc.core.config.RpcConfig; import com.xiaoyu.rpc.common.extension.ExtensionLoader; +import com.xiaoyu.rpc.core.config.RpcConfig; public class RpcClientProxy { + private RpcClientProxy() { + } + public static T create(Class clazz) { // 代理类型由配置决定(jdk / bytebuddy 等),这里统一走 SPI 扩展点加载 + return currentProxyFactory().getProxy(clazz); + } + + /** + * 释放当前代理工厂持有的客户端网络资源。 + * 该方法面向应用退出阶段,调用后当前 SPI 代理工厂不再接受新的 RPC 调用。 + */ + public static void shutdown() { + currentProxyFactory().close(); + } + + private static ProxyFactory currentProxyFactory() { String proxyType = RpcConfig.getInstance().getProxyType(); - ProxyFactory proxyFactory = ExtensionLoader.getExtensionLoader(ProxyFactory.class).getExtension(proxyType); - return proxyFactory.getProxy(clazz); + return ExtensionLoader.getExtensionLoader(ProxyFactory.class).getExtension(proxyType); } } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcReturnTypeResolver.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcReturnTypeResolver.java new file mode 100644 index 0000000..5834704 --- /dev/null +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcReturnTypeResolver.java @@ -0,0 +1,40 @@ +package com.xiaoyu.rpc.core.client; + +import java.lang.reflect.Method; +import java.lang.reflect.ParameterizedType; +import java.lang.reflect.Type; +import java.util.concurrent.CompletableFuture; + +/** + * 解析 RPC 代理方法真正需要反序列化的响应类型。 + */ +final class RpcReturnTypeResolver { + + private RpcReturnTypeResolver() { + } + + static boolean isAsync(Method method) { + return CompletableFuture.class.isAssignableFrom(method.getReturnType()); + } + + static Class resolvePayloadType(Method method) { + if (!isAsync(method)) { + return method.getReturnType(); + } + + Type genericReturnType = method.getGenericReturnType(); + if (!(genericReturnType instanceof ParameterizedType)) { + throw new IllegalArgumentException( + "异步 RPC 方法必须声明具体泛型返回值 CompletableFuture: " + method.toGenericString()); + } + + ParameterizedType parameterizedType = (ParameterizedType) genericReturnType; + Type[] typeArguments = parameterizedType.getActualTypeArguments(); + if (typeArguments.length != 1 || !(typeArguments[0] instanceof Class)) { + throw new IllegalArgumentException( + "暂不支持参数化或未解析的异步 RPC 返回类型: " + genericReturnType.getTypeName()); + } + + return (Class) typeArguments[0]; + } +} diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/server/RpcServer.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/server/RpcServer.java index beac1b6..8329af8 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/server/RpcServer.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/server/RpcServer.java @@ -1,50 +1,73 @@ package com.xiaoyu.rpc.core.server; -import com.xiaoyu.rpc.core.config.RpcConfig; import com.xiaoyu.rpc.common.extension.ExtensionLoader; +import com.xiaoyu.rpc.core.config.RpcConfig; import com.xiaoyu.rpc.core.registry.ServiceRegistry; import com.xiaoyu.rpc.core.transport.Transport; import com.xiaoyu.rpc.core.transport.TransportServer; import lombok.extern.slf4j.Slf4j; import java.net.InetSocketAddress; +import java.util.concurrent.atomic.AtomicBoolean; @Slf4j -public class RpcServer { +public class RpcServer implements AutoCloseable { private final String serverHost; private final int serverPort; private final ServiceRegistry serviceRegistry; private final TransportServer transportServer; + private final AtomicBoolean stopped = new AtomicBoolean(false); + private final Thread shutdownHook; public RpcServer() { - RpcConfig config = RpcConfig.getInstance(); - this.serverHost = config.getServerHost(); - this.serverPort = config.getServerPort(); - this.serviceRegistry = ExtensionLoader.getExtensionLoader(ServiceRegistry.class) - .getExtension(config.getRegistryType()); - - // 获取传输层实现 - Transport transport = ExtensionLoader.getExtensionLoader(Transport.class).getExtension(config.getTransport()); - this.transportServer = transport.createServer(this.serverPort); - - // 注册 JVM 关闭挂钩 (优雅下线) - Runtime.getRuntime().addShutdownHook(new Thread(() -> { - log.info("检测到 JVM 关闭信号,正在执行优雅下线..."); - // 先注销服务,阻止新流量进入 - serviceRegistry.clearRegistry(); - // 再关闭网络层,让存量请求有机会处理完成 - transportServer.stop(); - log.info("优雅下线完成。"); - })); + this(RpcConfig.getInstance()); + } + + private RpcServer(RpcConfig config) { + this( + config.getServerHost(), + config.getServerPort(), + ExtensionLoader.getExtensionLoader(ServiceRegistry.class) + .getExtension(config.getRegistryType()), + ExtensionLoader.getExtensionLoader(Transport.class) + .getExtension(config.getTransport()) + .createServer(config.getServerPort()), + true); + } + + RpcServer(String serverHost, int serverPort, ServiceRegistry serviceRegistry, + TransportServer transportServer) { + this(serverHost, serverPort, serviceRegistry, transportServer, false); + } + + private RpcServer(String serverHost, int serverPort, ServiceRegistry serviceRegistry, + TransportServer transportServer, boolean installShutdownHook) { + this.serverHost = serverHost; + this.serverPort = serverPort; + this.serviceRegistry = serviceRegistry; + this.transportServer = transportServer; + + if (installShutdownHook) { + this.shutdownHook = new Thread(() -> { + log.info("检测到 JVM 关闭信号,正在执行优雅下线..."); + stop(); + log.info("优雅下线完成。"); + }, "rpc-server-shutdown"); + Runtime.getRuntime().addShutdownHook(this.shutdownHook); + } else { + this.shutdownHook = null; + } } public void register(Class interfaceClass, T serviceImpl) { + if (stopped.get()) { + throw new IllegalStateException("RpcServer 已关闭"); + } + String serviceName = interfaceClass.getName(); - // 先做本地注册,便于请求分发时快速定位实现类 ServiceRepository.registerService(serviceName, serviceImpl); - // 再注册到注册中心(Nacos / Local) try { serviceRegistry.registerService(serviceName, new InetSocketAddress(serverHost, serverPort)); log.info("Service registered: {}", serviceName); @@ -54,6 +77,50 @@ public void register(Class interfaceClass, T serviceImpl) { } public void start() throws InterruptedException { - transportServer.start(); + if (stopped.get()) { + throw new IllegalStateException("RpcServer 已关闭"); + } + + try { + transportServer.start(); + } finally { + stop(); + } + } + + public void stop() { + if (!stopped.compareAndSet(false, true)) { + return; + } + + try { + serviceRegistry.clearRegistry(); + } catch (Exception e) { + log.warn("清理服务注册信息失败", e); + } + + try { + transportServer.stop(); + } catch (Exception e) { + log.warn("停止 RPC TransportServer 失败", e); + } + + removeShutdownHook(); + } + + @Override + public void close() { + stop(); + } + + private void removeShutdownHook() { + if (shutdownHook == null || Thread.currentThread() == shutdownHook) { + return; + } + try { + Runtime.getRuntime().removeShutdownHook(shutdownHook); + } catch (IllegalStateException | SecurityException ignored) { + // JVM 已进入关闭阶段时无法移除 Hook,忽略即可。 + } } } diff --git a/rpc-core/src/main/java/com/xiaoyu/rpc/core/transport/TransportClient.java b/rpc-core/src/main/java/com/xiaoyu/rpc/core/transport/TransportClient.java index 0fe8fea..edd126e 100644 --- a/rpc-core/src/main/java/com/xiaoyu/rpc/core/transport/TransportClient.java +++ b/rpc-core/src/main/java/com/xiaoyu/rpc/core/transport/TransportClient.java @@ -1,13 +1,14 @@ package com.xiaoyu.rpc.core.transport; import com.xiaoyu.rpc.common.vo.RpcRequest; + import java.net.InetSocketAddress; import java.util.concurrent.CompletableFuture; /** * 传输层客户端接口 */ -public interface TransportClient { +public interface TransportClient extends AutoCloseable { /** * 发送 RPC 请求 @@ -17,4 +18,9 @@ public interface TransportClient { * @return 响应结果 (CompletableFuture) */ CompletableFuture sendRequest(RpcRequest request, InetSocketAddress address); + + @Override + default void close() { + // 无状态传输实现默认无需释放资源。 + } } diff --git a/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/ProxyAsyncReturnTypeTest.java b/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/ProxyAsyncReturnTypeTest.java new file mode 100644 index 0000000..6265db9 --- /dev/null +++ b/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/ProxyAsyncReturnTypeTest.java @@ -0,0 +1,126 @@ +package com.xiaoyu.rpc.core.client; + +import com.google.protobuf.ByteString; +import com.xiaoyu.rpc.common.extension.ExtensionLoader; +import com.xiaoyu.rpc.common.serialization.Serializer; +import com.xiaoyu.rpc.common.vo.RpcResponse; +import com.xiaoyu.rpc.core.config.RpcConfig; +import com.xiaoyu.rpc.core.registry.ServiceDiscovery; +import com.xiaoyu.rpc.core.transport.TransportClient; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.io.Serializable; +import java.lang.reflect.Field; +import java.net.InetSocketAddress; +import java.util.List; +import java.util.Objects; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; + +@DisplayName("异步 RPC 泛型返回值测试") +public class ProxyAsyncReturnTypeTest { + + @BeforeEach + void setUp() throws Exception { + System.setProperty("rpc.registry", "local"); + System.setProperty("rpc.serializer", "java"); + resetRpcConfigSingleton(); + } + + @AfterEach + void tearDown() throws Exception { + System.clearProperty("rpc.registry"); + System.clearProperty("rpc.serializer"); + resetRpcConfigSingleton(); + } + + @Test + @DisplayName("JDK Proxy 按 CompletableFuture 的 T 反序列化") + void testJdkProxyAsyncPayloadType() throws Exception { + RpcClient rpcClient = rpcClientReturning(new TestUser("alice")); + AsyncUserService proxy = new JdkProxyFactory(rpcClient).getProxy(AsyncUserService.class); + + TestUser user = proxy.findUser("alice").get(1, TimeUnit.SECONDS); + + assertEquals(new TestUser("alice"), user); + } + + @Test + @DisplayName("ByteBuddy Proxy 按 CompletableFuture 的 T 反序列化") + void testByteBuddyProxyAsyncPayloadType() throws Exception { + RpcClient rpcClient = rpcClientReturning(new TestUser("bob")); + AsyncUserService proxy = new ByteBuddyProxyFactory(rpcClient).getProxy(AsyncUserService.class); + + TestUser user = proxy.findUser("bob").get(1, TimeUnit.SECONDS); + + assertEquals(new TestUser("bob"), user); + } + + @Test + @DisplayName("嵌套参数化异步返回值明确拒绝而不是错误反序列化") + void testRejectNestedParameterizedAsyncReturnType() { + RpcClient rpcClient = rpcClientReturning(new TestUser("unused")); + NestedAsyncService proxy = new JdkProxyFactory(rpcClient).getProxy(NestedAsyncService.class); + + assertThrows(IllegalArgumentException.class, proxy::findUsers); + } + + private static RpcClient rpcClientReturning(Object value) { + Serializer serializer = ExtensionLoader.getExtensionLoader(Serializer.class).getExtension("java"); + TransportClient transportClient = (request, address) -> { + RpcResponse response = RpcResponse.newBuilder() + .setRequestId(request.getRequestId()) + .setMessage("Success") + .setData(ByteString.copyFrom(serializer.serialize(value))) + .build(); + return CompletableFuture.completedFuture(response); + }; + ServiceDiscovery discovery = serviceName -> new InetSocketAddress("127.0.0.1", 8080); + return new RpcClient(transportClient, discovery); + } + + public interface AsyncUserService { + CompletableFuture findUser(String name); + } + + public interface NestedAsyncService { + CompletableFuture> findUsers(); + } + + public static final class TestUser implements Serializable { + private final String name; + + public TestUser(String name) { + this.name = name; + } + + @Override + public boolean equals(Object obj) { + if (this == obj) { + return true; + } + if (!(obj instanceof TestUser)) { + return false; + } + TestUser other = (TestUser) obj; + return Objects.equals(name, other.name); + } + + @Override + public int hashCode() { + return Objects.hash(name); + } + } + + private static void resetRpcConfigSingleton() throws Exception { + Field field = RpcConfig.class.getDeclaredField("instance"); + field.setAccessible(true); + field.set(null, null); + } +} diff --git a/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/RpcClientTest.java b/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/RpcClientTest.java index a38fb56..5480d03 100644 --- a/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/RpcClientTest.java +++ b/rpc-core/src/test/java/com/xiaoyu/rpc/core/client/RpcClientTest.java @@ -18,6 +18,7 @@ import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutionException; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; import static org.junit.jupiter.api.Assertions.*; @@ -122,6 +123,38 @@ void testVoidReturnType() throws Exception { assertNull(result); } + @Test + @DisplayName("close 幂等释放传输层且关闭后拒绝新请求") + void testCloseIsIdempotentAndRejectsNewRequests() throws Exception { + AtomicInteger closeCount = new AtomicInteger(); + AtomicInteger discoveryCount = new AtomicInteger(); + TransportClient transportClient = new TransportClient() { + @Override + public CompletableFuture sendRequest(RpcRequest request, InetSocketAddress address) { + return CompletableFuture.completedFuture(successfulResponse(new byte[0])); + } + + @Override + public void close() { + closeCount.incrementAndGet(); + } + }; + ServiceDiscovery serviceDiscovery = serviceName -> { + discoveryCount.incrementAndGet(); + return new InetSocketAddress("127.0.0.1", 8080); + }; + RpcClient rpcClient = new RpcClient(transportClient, serviceDiscovery); + + rpcClient.close(); + rpcClient.close(); + + assertEquals(1, closeCount.get()); + CompletableFuture future = rpcClient.sendRequest(minimalRequest(), String.class); + ExecutionException ex = assertThrows(ExecutionException.class, () -> future.get(1, TimeUnit.SECONDS)); + assertInstanceOf(IllegalStateException.class, ex.getCause()); + assertEquals(0, discoveryCount.get(), "关闭后的请求不应继续访问服务发现"); + } + private static RpcResponse successfulResponse(byte[] body) { return RpcResponse.newBuilder() .setRequestId("req-1") diff --git a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/ChannelProvider.java b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/ChannelProvider.java index b17d658..a430141 100644 --- a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/ChannelProvider.java +++ b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/ChannelProvider.java @@ -1,104 +1,185 @@ package com.xiaoyu.rpc.core.client; -import com.xiaoyu.rpc.core.config.RpcConfig; import io.netty.bootstrap.Bootstrap; import io.netty.channel.Channel; +import io.netty.channel.ChannelFuture; import io.netty.channel.ChannelFutureListener; -import lombok.extern.slf4j.Slf4j; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import java.net.InetSocketAddress; +import java.util.ArrayList; import java.util.LinkedHashMap; +import java.util.List; import java.util.Map; +import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; -import java.util.concurrent.CountDownLatch; -import java.util.concurrent.TimeUnit; - -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; +import java.util.concurrent.atomic.AtomicBoolean; -public class ChannelProvider { +/** + * 单个 Netty 客户端实例使用的连接池。 + * 缓存已建立连接,并对同一地址的并发首次建连进行去重。 + */ +public final class ChannelProvider implements AutoCloseable { private static final Logger log = LoggerFactory.getLogger(ChannelProvider.class); - private static final Map channels = new ConcurrentHashMap<>(); - private static final LinkedHashMap lruTracker = new LinkedHashMap<>(16, 0.75f, true); - private static final int MAX_CONNECTIONS = RpcConfig.getInstance().getMaxConnections(); - - public static Channel get(InetSocketAddress inetSocketAddress, Bootstrap bootstrap) { - String key = inetSocketAddress.toString(); - // 先尝试复用已有连接 - if (channels.containsKey(key)) { - Channel channel = channels.get(key); - if (channel != null && channel.isActive()) { - synchronized (lruTracker) { - lruTracker.put(key, System.currentTimeMillis()); - } - return channel; - } else { - removeChannel(key); + @FunctionalInterface + interface ChannelConnector { + ChannelFuture connect(Bootstrap bootstrap, InetSocketAddress address); + } + + private final Map channels = new ConcurrentHashMap<>(); + private final Map> connectingChannels = new ConcurrentHashMap<>(); + private final LinkedHashMap lruTracker = new LinkedHashMap<>(16, 0.75f, true); + private final int maxConnections; + private final ChannelConnector connector; + private final AtomicBoolean closed = new AtomicBoolean(false); + + public ChannelProvider(int maxConnections) { + this(maxConnections, (bootstrap, address) -> bootstrap.connect(address)); + } + + ChannelProvider(int maxConnections, ChannelConnector connector) { + this.maxConnections = Math.max(1, maxConnections); + this.connector = connector; + } + + public CompletableFuture get(InetSocketAddress address, Bootstrap bootstrap) { + if (closed.get()) { + return CompletableFuture.failedFuture(new IllegalStateException("ChannelProvider 已关闭")); + } + + String key = key(address); + Channel cached = channels.get(key); + if (cached != null) { + if (cached.isActive()) { + touch(key); + return CompletableFuture.completedFuture(cached); } + removeChannel(key, cached); } - // 缓存不可用时再新建连接 - Channel channel = connect(bootstrap, inetSocketAddress); + CompletableFuture existing = connectingChannels.get(key); + if (existing != null) { + return existing; + } - // 新连接建立成功后放回缓存 - if (channel != null) { - addChannel(key, channel); + CompletableFuture candidate = new CompletableFuture<>(); + existing = connectingChannels.putIfAbsent(key, candidate); + if (existing != null) { + return existing; } - return channel; + candidate.whenComplete((channel, throwable) -> connectingChannels.remove(key, candidate)); + connect(key, address, bootstrap, candidate); + return candidate; } - private static void addChannel(String key, Channel channel) { + private void connect(String key, InetSocketAddress address, Bootstrap bootstrap, + CompletableFuture result) { + if (closed.get()) { + result.completeExceptionally(new IllegalStateException("ChannelProvider 已关闭")); + return; + } + + try { + connector.connect(bootstrap, address).addListener((ChannelFutureListener) future -> { + if (!future.isSuccess()) { + Throwable cause = future.cause() != null + ? future.cause() + : new RuntimeException("客户端连接失败: " + address); + log.error("客户端连接失败: {}", address, cause); + result.completeExceptionally(cause); + return; + } + + Channel channel = future.channel(); + if (closed.get() || !cacheChannel(key, channel)) { + channel.close(); + result.completeExceptionally(new IllegalStateException("ChannelProvider 已关闭")); + return; + } + + log.info("客户端连接成功: {}", address); + result.complete(channel); + }); + } catch (Exception e) { + result.completeExceptionally(e); + } + } + + private boolean cacheChannel(String key, Channel channel) { synchronized (lruTracker) { - if (channels.size() >= MAX_CONNECTIONS) { + if (closed.get()) { + return false; + } + + while (channels.size() >= maxConnections && !lruTracker.isEmpty()) { evictLRU(); } + channels.put(key, channel); lruTracker.put(key, System.currentTimeMillis()); } + + channel.closeFuture().addListener(ignored -> removeChannel(key, channel)); + return true; } - private static void removeChannel(String key) { + private void touch(String key) { synchronized (lruTracker) { - channels.remove(key); - lruTracker.remove(key); + if (channels.containsKey(key)) { + lruTracker.put(key, System.currentTimeMillis()); + } } } - private static void evictLRU() { - if (lruTracker.isEmpty()) { - return; + private void removeChannel(String key, Channel expected) { + synchronized (lruTracker) { + if (channels.remove(key, expected)) { + lruTracker.remove(key); + } } + } + + private void evictLRU() { String oldestKey = lruTracker.keySet().iterator().next(); Channel oldChannel = channels.remove(oldestKey); lruTracker.remove(oldestKey); - if (oldChannel != null && oldChannel.isActive()) { + if (oldChannel != null) { oldChannel.close(); } log.info("连接池已满,淘汰最久未使用的连接: {}", oldestKey); } - private static Channel connect(Bootstrap bootstrap, InetSocketAddress inetSocketAddress) { - CountDownLatch latch = new CountDownLatch(1); - final Channel[] channelHolder = new Channel[1]; + @Override + public void close() { + if (!closed.compareAndSet(false, true)) { + return; + } - bootstrap.connect(inetSocketAddress).addListener((ChannelFutureListener) future -> { - if (future.isSuccess()) { - log.info("客户端连接成功: " + inetSocketAddress); - channelHolder[0] = future.channel(); - } else { - log.error("客户端连接失败: " + inetSocketAddress); - } - latch.countDown(); - }); + IllegalStateException closeCause = new IllegalStateException("ChannelProvider 已关闭"); + connectingChannels.values().forEach(future -> future.completeExceptionally(closeCause)); + connectingChannels.clear(); - try { - latch.await(5, TimeUnit.SECONDS); - } catch (InterruptedException e) { - Thread.currentThread().interrupt(); + List snapshot; + synchronized (lruTracker) { + snapshot = new ArrayList<>(channels.values()); + channels.clear(); + lruTracker.clear(); } + snapshot.forEach(Channel::close); + } + + int cachedChannelCount() { + return channels.size(); + } + + int connectingChannelCount() { + return connectingChannels.size(); + } - return channelHolder[0]; + private static String key(InetSocketAddress address) { + return address.getHostString() + ':' + address.getPort(); } } diff --git a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandler.java b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandler.java index e97146c..69f61d5 100644 --- a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandler.java +++ b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandler.java @@ -1,14 +1,16 @@ package com.xiaoyu.rpc.core.client; import com.xiaoyu.rpc.common.vo.RpcResponse; +import io.netty.channel.ChannelHandler; import io.netty.channel.ChannelHandlerContext; import io.netty.channel.SimpleChannelInboundHandler; -import java.util.concurrent.CompletableFuture; - import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import io.netty.channel.ChannelHandler; +import java.nio.channels.ClosedChannelException; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.ConcurrentHashMap; /** * 客户端响应处理器。 @@ -18,8 +20,7 @@ public class NettyRpcClientHandler extends SimpleChannelInboundHandler { private static final Logger log = LoggerFactory.getLogger(NettyRpcClientHandler.class); - // 一个连接上可以并发多个请求,靠 requestId 区分各自回调 - private final java.util.Map> pendingRequests = new java.util.concurrent.ConcurrentHashMap<>(); + private final Map> pendingRequests = new ConcurrentHashMap<>(); public void addFuture(String requestId, CompletableFuture future) { pendingRequests.put(requestId, future); @@ -36,6 +37,14 @@ public void failRequest(String requestId, Throwable cause) { } } + public void failAll(Throwable cause) { + pendingRequests.forEach((requestId, future) -> { + if (pendingRequests.remove(requestId, future)) { + future.completeExceptionally(cause); + } + }); + } + @Override protected void channelRead0(ChannelHandlerContext ctx, RpcResponse response) { String requestId = response.getRequestId(); @@ -49,14 +58,20 @@ protected void channelRead0(ChannelHandlerContext ctx, RpcResponse response) { } } + @Override + public void channelInactive(ChannelHandlerContext ctx) throws Exception { + failAll(new ClosedChannelException()); + super.channelInactive(ctx); + } + @Override public void exceptionCaught(ChannelHandlerContext ctx, Throwable cause) { log.error("Client caught exception", cause); - // 连接级异常通常影响当前连接上的全部在途请求,统一失败返回给上层 - for (CompletableFuture future : pendingRequests.values()) { - future.completeExceptionally(cause); - } - pendingRequests.clear(); + failAll(cause); ctx.close(); } + + int pendingRequestCount() { + return pendingRequests.size(); + } } diff --git a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportClient.java b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportClient.java index 272e008..8a46dbf 100644 --- a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportClient.java +++ b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportClient.java @@ -15,6 +15,9 @@ import io.netty.channel.nio.NioEventLoopGroup; import io.netty.channel.socket.SocketChannel; import io.netty.channel.socket.nio.NioSocketChannel; +import io.netty.util.concurrent.DefaultThreadFactory; +import io.netty.util.concurrent.EventExecutor; +import io.netty.util.concurrent.Future; import lombok.extern.slf4j.Slf4j; import java.net.InetSocketAddress; @@ -23,107 +26,144 @@ import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; @Slf4j public class NettyTransportClient implements TransportClient { - private static volatile EventLoopGroup eventLoopGroup; - private static volatile Bootstrap bootstrap; - - private static Bootstrap getBootstrap() { - if (bootstrap == null) { - synchronized (NettyTransportClient.class) { - if (bootstrap == null) { - // Bootstrap 和 EventLoopGroup 进程内复用,避免每次请求都创建线程池 - eventLoopGroup = new NioEventLoopGroup(); - Bootstrap newBootstrap = new Bootstrap(); - newBootstrap.group(eventLoopGroup) - .channel(NioSocketChannel.class) - .handler(new ChannelInitializer() { - @Override - protected void initChannel(SocketChannel ch) { - String protocolName = RpcConfig.getInstance().getProtocol(); - Protocol protocol = ProtocolFactory.getProtocol(protocolName); - protocol.config(ch.pipeline(), false, null); - } - }); - bootstrap = newBootstrap; - } - } - } - return bootstrap; + private final EventLoopGroup eventLoopGroup; + private final Bootstrap bootstrap; + private final ChannelProvider channelProvider; + private final AtomicBoolean closed = new AtomicBoolean(false); + + public NettyTransportClient() { + RpcConfig config = RpcConfig.getInstance(); + this.eventLoopGroup = new NioEventLoopGroup( + 0, + new DefaultThreadFactory("rpc-client-io", true)); + this.channelProvider = new ChannelProvider(config.getMaxConnections()); + this.bootstrap = new Bootstrap(); + this.bootstrap.group(eventLoopGroup) + .channel(NioSocketChannel.class) + .handler(new ChannelInitializer() { + @Override + protected void initChannel(SocketChannel ch) { + String protocolName = RpcConfig.getInstance().getProtocol(); + Protocol protocol = ProtocolFactory.getProtocol(protocolName); + protocol.config(ch.pipeline(), false, null); + } + }); } @Override public CompletableFuture sendRequest(RpcRequest request, InetSocketAddress address) { + if (closed.get()) { + return CompletableFuture.failedFuture(new IllegalStateException("NettyTransportClient 已关闭")); + } + RpcConfig config = RpcConfig.getInstance(); String protocolName = config.getProtocol(); try { - Channel channel = ChannelProvider.get(address, getBootstrap()); - if (channel == null || !channel.isActive()) { - throw new RuntimeException("无法连接到服务器: " + address); - } + return channelProvider.get(address, bootstrap) + .thenCompose(channel -> sendOnChannel(request, channel, config, protocolName)) + .whenComplete((result, throwable) -> { + if (throwable != null) { + log.warn("RPC请求失败: address={}", address, throwable); + } + }); + } catch (Exception e) { + log.error("RPC请求发起失败", e); + return CompletableFuture.failedFuture(e); + } + } - NettyRpcClientHandler handler = channel.pipeline().get(NettyRpcClientHandler.class); - if (handler == null) { - handler = new NettyRpcClientHandler(); - channel.pipeline().addLast(handler); - } - final NettyRpcClientHandler clientHandler = handler; - - String requestId = UUID.randomUUID().toString(); - RpcRequest newRequest = request.toBuilder() - .setRequestId(requestId) - .build(); - - CompletableFuture resultFuture = new CompletableFuture<>(); - // 必须先注册 Future 再发送,避免极端情况下响应先到。 - clientHandler.addFuture(requestId, resultFuture); - - int timeoutMillis = Math.max(1, config.getRequestTimeoutMillis()); - final ScheduledFuture timeoutTask; - try { - timeoutTask = channel.eventLoop().schedule( - () -> clientHandler.failRequest(requestId, - new TimeoutException("RPC请求超时: requestId=" + requestId - + ", timeoutMs=" + timeoutMillis)), - timeoutMillis, - TimeUnit.MILLISECONDS); - } catch (Exception e) { - clientHandler.failRequest(requestId, e); - throw e; - } + private CompletableFuture sendOnChannel(RpcRequest request, Channel channel, + RpcConfig config, String protocolName) { + if (closed.get()) { + return CompletableFuture.failedFuture(new IllegalStateException("NettyTransportClient 已关闭")); + } + if (channel == null || !channel.isActive()) { + return CompletableFuture.failedFuture(new RuntimeException("无法连接到服务器: " + channel)); + } - // 无论正常完成、超时还是异常,都取消定时任务并确保 pendingRequests 被清理。 - resultFuture.whenComplete((result, throwable) -> { - timeoutTask.cancel(false); - clientHandler.removeFuture(requestId); - }); - - Protocol protocol = ProtocolFactory.getProtocol(protocolName); - try { - protocol.sendRequest(channel, newRequest, clientHandler); - } catch (Exception e) { - clientHandler.failRequest(requestId, e); - throw e; + NettyRpcClientHandler handler = channel.pipeline().get(NettyRpcClientHandler.class); + if (handler == null) { + synchronized (channel) { + handler = channel.pipeline().get(NettyRpcClientHandler.class); + if (handler == null) { + handler = new NettyRpcClientHandler(); + channel.pipeline().addLast(handler); + } } + } + final NettyRpcClientHandler clientHandler = handler; - return resultFuture.thenApply(result -> { - if (result instanceof RpcResponse) { - RpcResponse rpcResponse = (RpcResponse) result; - if (!"Success".equals(rpcResponse.getMessage())) { - throw new RuntimeException("服务端报错: " + rpcResponse.getMessage()); - } - return rpcResponse; - } - throw new RuntimeException("服务端返回的不是 RpcResponse 类型"); - }); + String requestId = UUID.randomUUID().toString(); + RpcRequest newRequest = request.toBuilder() + .setRequestId(requestId) + .build(); + + CompletableFuture resultFuture = new CompletableFuture<>(); + clientHandler.addFuture(requestId, resultFuture); + + int timeoutMillis = Math.max(1, config.getRequestTimeoutMillis()); + final ScheduledFuture timeoutTask; + try { + timeoutTask = channel.eventLoop().schedule( + () -> clientHandler.failRequest(requestId, + new TimeoutException("RPC请求超时: requestId=" + requestId + + ", timeoutMs=" + timeoutMillis)), + timeoutMillis, + TimeUnit.MILLISECONDS); } catch (Exception e) { - log.error("RPC请求发起失败", e); - CompletableFuture future = new CompletableFuture<>(); - future.completeExceptionally(e); - return future; + clientHandler.failRequest(requestId, e); + return CompletableFuture.failedFuture(e); + } + + resultFuture.whenComplete((result, throwable) -> { + timeoutTask.cancel(false); + clientHandler.removeFuture(requestId); + }); + + Protocol protocol = ProtocolFactory.getProtocol(protocolName); + try { + protocol.sendRequest(channel, newRequest, clientHandler); + } catch (Exception e) { + clientHandler.failRequest(requestId, e); + } + + return resultFuture.thenApply(result -> { + if (result instanceof RpcResponse) { + RpcResponse rpcResponse = (RpcResponse) result; + if (!"Success".equals(rpcResponse.getMessage())) { + throw new RuntimeException("服务端报错: " + rpcResponse.getMessage()); + } + return rpcResponse; + } + throw new RuntimeException("服务端返回的不是 RpcResponse 类型"); + }); + } + + @Override + public void close() { + if (!closed.compareAndSet(false, true)) { + return; + } + + channelProvider.close(); + Future shutdownFuture = eventLoopGroup.shutdownGracefully(0, 5, TimeUnit.SECONDS); + if (!isInEventLoop()) { + shutdownFuture.syncUninterruptibly(); + } + } + + private boolean isInEventLoop() { + for (EventExecutor executor : eventLoopGroup) { + if (executor.inEventLoop()) { + return true; + } } + return false; } } diff --git a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportServer.java b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportServer.java index d120200..05bf47c 100644 --- a/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportServer.java +++ b/rpc-transport-netty/src/main/java/com/xiaoyu/rpc/core/transport/netty/NettyTransportServer.java @@ -7,23 +7,30 @@ import com.xiaoyu.rpc.core.server.NettyRpcHandler; import com.xiaoyu.rpc.core.transport.TransportServer; import io.netty.bootstrap.ServerBootstrap; +import io.netty.channel.Channel; import io.netty.channel.ChannelInitializer; import io.netty.channel.EventLoopGroup; import io.netty.channel.nio.NioEventLoopGroup; import io.netty.channel.socket.SocketChannel; import io.netty.channel.socket.nio.NioServerSocketChannel; +import io.netty.util.concurrent.EventExecutor; +import io.netty.util.concurrent.Future; import lombok.extern.slf4j.Slf4j; import java.util.concurrent.ArrayBlockingQueue; import java.util.concurrent.ThreadFactory; import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; @Slf4j public class NettyTransportServer implements TransportServer { private final int port; + private final AtomicBoolean started = new AtomicBoolean(false); + private final AtomicBoolean stopped = new AtomicBoolean(false); + private volatile Channel serverChannel; private EventLoopGroup bossGroup; private EventLoopGroup workerGroup; private ThreadPoolExecutor businessExecutor; @@ -34,6 +41,13 @@ public NettyTransportServer(int port) { @Override public void start() throws InterruptedException { + if (stopped.get()) { + throw new IllegalStateException("NettyTransportServer 已关闭"); + } + if (!started.compareAndSet(false, true)) { + throw new IllegalStateException("NettyTransportServer 已启动"); + } + RpcConfig config = RpcConfig.getInstance(); int cpuCores = Runtime.getRuntime().availableProcessors(); int bossThreads = Math.max(1, config.getBossThreads()); @@ -56,7 +70,6 @@ public void start() throws InterruptedException { new NamedThreadFactory("rpc-business-"), new ThreadPoolExecutor.AbortPolicy()); - // 一个服务端实例共享同一个无状态 Handler 和业务线程池,避免按连接创建线程资源。 NettyRpcHandler serverHandler = new NettyRpcHandler(businessExecutor); try { @@ -76,10 +89,11 @@ protected void initChannel(SocketChannel ch) { } }); + serverChannel = bootstrap.bind(port).sync().channel(); log.info("RPC Server (Netty) started on port {}, bossThreads={}, workerThreads={}, businessThreads={}, " + "businessQueueCapacity={}", port, bossThreads, workerThreads, businessThreads, businessQueueCapacity); - bootstrap.bind(port).sync().channel().closeFuture().sync(); + serverChannel.closeFuture().sync(); } finally { stop(); } @@ -87,15 +101,18 @@ protected void initChannel(SocketChannel ch) { @Override public void stop() { - // 先停止接收新连接,并开始关闭 I/O 线程。 - if (bossGroup != null) { - bossGroup.shutdownGracefully(); + if (!stopped.compareAndSet(false, true)) { + return; } - if (workerGroup != null) { - workerGroup.shutdownGracefully(); + + Channel channel = serverChannel; + if (channel != null) { + channel.close().syncUninterruptibly(); } - // 不再接收新任务后,尽量等待已提交的业务请求执行完成。 + shutdownEventLoop(workerGroup); + shutdownEventLoop(bossGroup); + if (businessExecutor != null) { businessExecutor.shutdown(); try { @@ -109,6 +126,24 @@ public void stop() { } } + private static void shutdownEventLoop(EventLoopGroup group) { + if (group == null) { + return; + } + + Future shutdownFuture = group.shutdownGracefully(0, 5, TimeUnit.SECONDS); + boolean calledFromGroup = false; + for (EventExecutor executor : group) { + if (executor.inEventLoop()) { + calledFromGroup = true; + break; + } + } + if (!calledFromGroup) { + shutdownFuture.syncUninterruptibly(); + } + } + private static final class NamedThreadFactory implements ThreadFactory { private final String prefix; private final AtomicInteger sequence = new AtomicInteger(1); diff --git a/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/ChannelProviderTest.java b/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/ChannelProviderTest.java new file mode 100644 index 0000000..adb87a5 --- /dev/null +++ b/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/ChannelProviderTest.java @@ -0,0 +1,80 @@ +package com.xiaoyu.rpc.core.client; + +import io.netty.bootstrap.Bootstrap; +import io.netty.channel.Channel; +import io.netty.channel.DefaultChannelPromise; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.util.concurrent.ImmediateEventExecutor; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.net.InetSocketAddress; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; +import java.util.concurrent.atomic.AtomicInteger; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +@DisplayName("ChannelProvider 并发建连与生命周期测试") +class ChannelProviderTest { + + @Test + @DisplayName("同一地址的并发首次获取共享一个 in-flight 连接") + void testDeduplicateInFlightConnection() { + EmbeddedChannel channel = new EmbeddedChannel(); + DefaultChannelPromise connectPromise = new DefaultChannelPromise(channel, ImmediateEventExecutor.INSTANCE); + AtomicInteger connectAttempts = new AtomicInteger(); + ChannelProvider provider = new ChannelProvider(8, (bootstrap, address) -> { + connectAttempts.incrementAndGet(); + return connectPromise; + }); + Bootstrap bootstrap = new Bootstrap(); + InetSocketAddress address = new InetSocketAddress("127.0.0.1", 18080); + + CompletableFuture first = provider.get(address, bootstrap); + CompletableFuture second = provider.get(address, bootstrap); + CompletableFuture third = provider.get(address, bootstrap); + + assertEquals(1, connectAttempts.get()); + assertEquals(1, provider.connectingChannelCount()); + assertFalse(first.isDone()); + + connectPromise.setSuccess(); + + assertSame(channel, first.join()); + assertSame(channel, second.join()); + assertSame(channel, third.join()); + assertEquals(1, connectAttempts.get()); + assertEquals(1, provider.cachedChannelCount()); + assertEquals(0, provider.connectingChannelCount()); + + provider.close(); + assertFalse(channel.isOpen()); + } + + @Test + @DisplayName("关闭时失败正在建立的连接并拒绝后续获取") + void testCloseFailsConnectingAndRejectsNewGet() { + EmbeddedChannel channel = new EmbeddedChannel(); + DefaultChannelPromise connectPromise = new DefaultChannelPromise(channel, ImmediateEventExecutor.INSTANCE); + ChannelProvider provider = new ChannelProvider(8, (bootstrap, address) -> connectPromise); + Bootstrap bootstrap = new Bootstrap(); + InetSocketAddress address = new InetSocketAddress("127.0.0.1", 18081); + + CompletableFuture connecting = provider.get(address, bootstrap); + provider.close(); + + assertTrue(connecting.isCompletedExceptionally()); + assertThrows(CompletionException.class, connecting::join); + + connectPromise.setSuccess(); + assertFalse(channel.isOpen(), "关闭后才完成的连接必须立即关闭"); + + CompletableFuture afterClose = provider.get(address, bootstrap); + assertThrows(CompletionException.class, afterClose::join); + } +} diff --git a/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandlerTest.java b/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandlerTest.java index 7781c57..b1f1260 100644 --- a/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandlerTest.java +++ b/rpc-transport-netty/src/test/java/com/xiaoyu/rpc/core/client/NettyRpcClientHandlerTest.java @@ -6,6 +6,7 @@ import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; +import java.nio.channels.ClosedChannelException; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutionException; import java.util.concurrent.TimeUnit; @@ -64,12 +65,28 @@ void testExceptionCaughtFailsAllPending() { assertTrue(f1.isCompletedExceptionally()); assertTrue(f2.isCompletedExceptionally()); + assertEquals(0, handler.pendingRequestCount()); assertFalse(channel.isActive(), "Channel should be closed on exception"); assertThrows(ExecutionException.class, f1::get); assertThrows(ExecutionException.class, f2::get); } + @Test + @DisplayName("连接关闭时所有挂起请求应以 ClosedChannelException 失败") + void testChannelInactiveFailsAllPending() throws Exception { + NettyRpcClientHandler handler = new NettyRpcClientHandler(); + EmbeddedChannel channel = new EmbeddedChannel(handler); + CompletableFuture future = new CompletableFuture<>(); + handler.addFuture("pending", future); + + channel.close().syncUninterruptibly(); + + ExecutionException exception = assertThrows(ExecutionException.class, future::get); + assertInstanceOf(ClosedChannelException.class, exception.getCause()); + assertEquals(0, handler.pendingRequestCount()); + } + private static RpcResponse response(String requestId, String message) { return RpcResponse.newBuilder() .setRequestId(requestId)