Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -45,9 +46,11 @@ public void testFullIntegration() throws Exception {
System.setProperty("rpc.protocol", protocol);

AtomicReference<Throwable> serverFailure = new AtomicReference<>();
AtomicReference<RpcServer> 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) {
Expand All @@ -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");
}
Expand Down
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -17,6 +17,7 @@
public class ByteBuddyProxyFactory implements ProxyFactory {

private volatile RpcClient rpcClient;
private volatile boolean closed;

public ByteBuddyProxyFactory() {
// 与 JDK Proxy 一致:SPI 扩展加载阶段不初始化注册中心和传输层。
Expand All @@ -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();
Expand All @@ -48,20 +56,16 @@ public <T> T getProxy(Class<T> 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) {
Expand All @@ -71,9 +75,11 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl
}

RpcRequest request = builder.build();
CompletableFuture<Object> future = getRpcClient().sendRequest(request, method.getReturnType());
// 如果业务接口声明的返回类型是异步的,直接返回 Future;否则阻塞等待结果
if (CompletableFuture.class.isAssignableFrom(method.getReturnType())) {
boolean async = RpcReturnTypeResolver.isAsync(method);
Class<?> payloadType = RpcReturnTypeResolver.resolvePayloadType(method);
CompletableFuture<Object> future = getRpcClient().sendRequest(request, payloadType);

if (async) {
return future;
}
return future.get();
Expand All @@ -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();
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
public class JdkProxyFactory implements ProxyFactory {

private volatile RpcClient rpcClient;
private volatile boolean closed;

public JdkProxyFactory() {
// SPI 扩展加载阶段保持轻量,不在构造时初始化注册中心和传输层。
Expand All @@ -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();
Expand All @@ -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) {
Expand All @@ -68,13 +74,32 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl
}

RpcRequest request = builder.build();
CompletableFuture<Object> future = getRpcClient().sendRequest(request, method.getReturnType());
// 如果业务接口声明的返回类型是异步的,直接返回 Future;否则阻塞等待结果
if (CompletableFuture.class.isAssignableFrom(method.getReturnType())) {
boolean async = RpcReturnTypeResolver.isAsync(method);
Class<?> payloadType = RpcReturnTypeResolver.resolvePayloadType(method);
CompletableFuture<Object> 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();
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,11 @@
import com.xiaoyu.rpc.common.extension.SPI;

@SPI
public interface ProxyFactory {
public interface ProxyFactory extends AutoCloseable {
<T> T getProxy(Class<T> clazz);

@Override
default void close() {
// 默认无资源需要释放;具体代理实现可覆盖。
}
}
32 changes: 19 additions & 13 deletions rpc-core/src/main/java/com/xiaoyu/rpc/core/client/RpcClient.java
Original file line number Diff line number Diff line change
Expand Up @@ -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();
}
Expand All @@ -37,23 +37,24 @@ public RpcClient() {
}

public CompletableFuture<Object> 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<Object> future = new CompletableFuture<>();
future.completeExceptionally(new RuntimeException("未发现服务: " + request.getInterfaceName()));
return future;
return CompletableFuture.failedFuture(
new RuntimeException("未发现服务: " + request.getInterfaceName()));
}

// 交给传输层发送,返回异步 Future
CompletableFuture<Object> 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;
Expand All @@ -69,9 +70,14 @@ public CompletableFuture<Object> sendRequest(RpcRequest request, Class<?> return
});

} catch (Exception e) {
CompletableFuture<Object> future = new CompletableFuture<>();
future.completeExceptionally(e);
return future;
return CompletableFuture.failedFuture(e);
}
}

@Override
public void close() {
if (closed.compareAndSet(false, true)) {
transportClient.close();
}
}
}
Original file line number Diff line number Diff line change
@@ -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> T create(Class<T> 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);
}
}
Original file line number Diff line number Diff line change
@@ -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<T>: " + 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];
}
}
Loading