什么是SPI?

SPI全称是Service Provider Interface,是一种服务发现机制。SPI的本质是将接口的实现类的全限定名配置在文件中,并由服务加载器读取配置文件,完成类的加载。这样可以在运行是,动态为接口替换实现类。通过SPI机制,可以为程序提供拓展功能。

Dubbo SPI如何使用?

我们首先需要知道的是SPI机制不仅仅是Dubbo有的,Java本身也提SPI的实现了的。

  1. 假如我们有这么一个接口,我们希望我们在运行时动态的选择其实现类。
1
2
3
public interface Robot {
void sayHello();
}
  1. 我们准备了两个不同的实现类
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
public class OptimusPrime implements Robot {

@Override
public void sayHello() {
System.out.println("Hello, I am Optimus Prime.");
}
}

public class Bumblebee implements Robot {

@Override
public void sayHello() {
System.out.println("Hello, I am Bumblebee.");
}
}
  1. 它的实现类的全路径我们配置到META-INF/dubbo路径下的配置文件中。
1
2
optimusPrime = org.apache.spi.OptimusPrime
bumblebee = org.apache.spi.Bumblebee
  1. 编写测试代码
1
2
3
4
5
6
7
8
9
10
11
public class DubboSPITest {
@Test
public void sayHello() throws Exception {
ExtensionLoader<Robot> extensionLoader =
ExtensionLoader.getExtensionLoader(Robot.class);
Robot optimusPrime = extensionLoader.getExtension("optimusPrime");
optimusPrime.sayHello();
Robot bumblebee = extensionLoader.getExtension("bumblebee");
bumblebee.sayHello();
}
}

Dubbo SPI源码分析

首先通过getExtensionLoader方法获取与拓展类对应的ExtensionLoader.然后通过getExtension获取拓展类对象。
getExtensionLoader的具体实现如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
public static <T> ExtensionLoader<T> getExtensionLoader(Class<T> type) {
if (type == null)
throw new IllegalArgumentException("Extension type == null");
if(!type.isInterface()) {
throw new IllegalArgumentException("Extension type(" + type + ") is not interface!");
}
if(!withExtensionAnnotation(type)) {
throw new IllegalArgumentException("Extension type(" + type +
") is not extension, because WITHOUT @" + SPI.class.getSimpleName() + " Annotation!");
}

//尝试从缓存中获取ExtensionLoader
ExtensionLoader<T> loader = (ExtensionLoader<T>) EXTENSION_LOADERS.get(type);
if (loader == null) {
//缓存中没有就新建,并加入到缓存中区
EXTENSION_LOADERS.putIfAbsent(type, new ExtensionLoader<T>(type));
loader = (ExtensionLoader<T>) EXTENSION_LOADERS.get(type);
}
return loader;
}

这个方法只做一件事情,就是获取拓展类对应的getExtensionLoader.首先尝试从缓存中获取,如果获取到了直接返回;没有获取到就创建并加入缓存,然后返回新创建的ExtensionLoader.

拿到ExtensionLoader之后,就可以通过T getExtension(String name)方法获取拓展类对象了。这个方法的具体实现如下;

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
public T getExtension(String name) {
if (name == null || name.length() == 0)
throw new IllegalArgumentException("Extension name == null");
if ("true".equals(name)) {
//如果为true,就返回设置的缺省值,否则返回null
return getDefaultExtension();
}

/*这里的Holder就是用来持有拓展类对象的一个pojo(一个保存对象的属性和get,set方法)。
cachedInstances就是一个ConcurrentMap,键为我们传入的name,值为持有拓展类对象的holder对象。
这行代码就是尝试从缓冲中获取拓展类(如果之前已经创建了)

*/
Holder<Object> holder = cachedInstances.get(name);

if (holder == null) {
/*如果缓存中没有,就创建并加入缓存。
注意,这里创建的仅仅是一个用于持有拓展类对象的holder对象,拓展类实际上还是没有创建*/
cachedInstances.putIfAbsent(name, new Holder<Object>());
holder = cachedInstances.get(name);
}
Object instance = holder.get();
//双从锁定机制
if (instance == null) {
synchronized (holder) {
instance = holder.get();
if (instance == null) {
/*如果holder中没有持有拓展类对象,就创建一个拓展类对象,
交由hodler持有*/
instance = createExtension(name);
holder.set(instance);
}
}
}
return (T) instance;
}

这段代码,逻辑还是比较复杂的,主要是缓存的处理。它的流程是这样的:

  1. 如果name为“true”,就返回默认拓展类,如果没有设置默认拓展类,那么就返回null
  2. 尝试从缓冲中获取拓展类,获取成功就返回
  3. 如果缓冲中,没有就创建拓展类对象,交由一个hodler对象持有后存入缓冲,并返回拓展类对象

那么创建拓展类对象的这个方法就非常的关键了。即这个方法中的这行代码
instance = createExtension(name);

创建拓展类对象的createExtension方法的具体实现如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
private T createExtension(String name) {
/*
getExtensionClasses拿到一个配置项名称到配置类的映射表map
然后通过get方法,拿到配置类
*/
Class<?> clazz = getExtensionClasses().get(name);
if (clazz == null) {
throw findException(name);
}
try {

/*
尝试从缓存中获取类所对应的实例
EXTENSION_INSTANCES是一个ConcurrentMap<Class<?>, Object>
*/
T instance = (T) EXTENSION_INSTANCES.get(clazz);
if (instance == null) {
//如果不存在,就通过反射创建拓展类实例,并加入缓存
EXTENSION_INSTANCES.putIfAbsent(clazz, (T) clazz.newInstance());
instance = (T) EXTENSION_INSTANCES.get(clazz);
}
//向实例中注入依赖
injectExtension(instance);
Set<Class<?>> wrapperClasses = cachedWrapperClasses;
if (wrapperClasses != null && wrapperClasses.size() > 0) {
// 循环创建 Wrapper 实例
for (Class<?> wrapperClass : wrapperClasses) {

// 将当前 instance 作为参数传给 Wrapper 的构造方法,并通过反射创建 Wrapper 实例。
// 然后向 Wrapper 实例中注入依赖,最后将 Wrapper 实例再次赋值给 instance 变量
instance = injectExtension((T) wrapperClass.getConstructor(type).newInstance(instance));
}
}
return instance;
} catch (Throwable t) {
throw new IllegalStateException("Extension instance(name: " + name + ", class: " +
type + ") could not be instantiated: " + t.getMessage(), t);
}
}

这个方法主要完成这么几件事情:

  1. 通过getExtensionClasses获取一个配置项名为键,拓展类类对象为值的map,然后调用get方法获取拓展类
  2. 尝试从缓存中获取拓展类实例,如果缓存中没有,就通过反射进行创建
  3. 向拓展类实例中注入依赖
  4. 将拓展类对象包裹在相应的wrapper对象中。

读完这个方法的代码,我们有几个大大的问号。

  1. 是如何向拓展类实例中注入依赖的?

那下面,我们就来探究这个问题。

想要知道是如何向拓展类实例中注入属性的,我们要从injectExtension(instance);方法入手,这个方法就完成了拓展类的属性注入工作。这也被称为Duboo IOC
它的具体实现如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
private T injectExtension(T instance) {
try {
if (objectFactory != null) {
for (Method method : instance.getClass().getMethods()) {
/*遍历该实例的方法,找出set开头的方法,
且仅有一个参数,且访问级别为public*/
if (method.getName().startsWith("set")
&& method.getParameterTypes().length == 1
&& Modifier.isPublic(method.getModifiers())) {
//获取setter方法的参数类型
Class<?> pt = method.getParameterTypes()[0];
try {
//获取属性名,如setName的属性名为name
String property = method.getName().length() > 3 ? method.getName().substring(3, 4).toLowerCase() + method.getName().substring(4) : "";

//从ObjectFactory中获取依赖对象
Object object = objectFactory.getExtension(pt, property);
if (object != null) {
//通过反射调用setter方法设置依赖
method.invoke(instance, object);
}
} catch (Exception e) {
logger.error("fail to inject via method " + method.getName()
+ " of interface " + type.getName() + ": " + e.getMessage(), e);
}
}
}
}
} catch (Exception e) {
logger.error(e.getMessage(), e);
}
return instance;
}

这个方法的实现还是比较简单的:

  1. 找出实例中的setter方法
  2. 拿到setter方法的参数类型,和字段名
  3. 从objectFactoy中拿到变量
  4. 通过反射调用setter方法,完成设值

整个流程还是非常的清晰的,不过objectFactory是什么呢?
objectFactory 变量的类型为 AdaptiveExtensionFactory,AdaptiveExtensionFactory 内部维护了一个 ExtensionFactory 列表,用于存储其他类型的 ExtensionFactory。Dubbo 目前提供了两种 ExtensionFactory,分别是 SpiExtensionFactory 和 SpringExtensionFactory。前者用于创建自适应的拓展,后者是用于从 Spring 的 IOC 容器中获取所需的拓展。

这两个类的实现是非常简单的。以SpringExtensionFactory为了,它的getExtension就是直接从spring的ApplicationContext中按名字取出bean返回就完事了。这里我们就不再分析。

Dubbo的自适应拓展

通过Dubbo的SPI机制,可以非常方便的加载拓展。如果我们不希望拓展在框架启动时就被加载,而是在拓展被调用时,根据运行时参数来进行加载。
Dubbo是这样做的:
首先 Dubbo 会为拓展接口生成具有代理功能的代码。然后通过 javassist 或 jdk 编译这段代码,得到 Class类。最后再通过反射创建代理类,在代理类中,就可以通过URL对象的参数来确定到底调用哪个实现类。

自适应拓展机制源码分析

整个自适应拓展机制的入口是getAdaptiveExtension()方法。
这个方法的具体实现如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
public T getAdaptiveExtension() {
//从缓存中获取自适应拓展
Object instance = cachedAdaptiveInstance.get();
if (instance == null) {
if(createAdaptiveInstanceError == null) {
synchronized (cachedAdaptiveInstance) {
instance = cachedAdaptiveInstance.get();
if (instance == null) {
try {
//创建自适应拓展
instance = createAdaptiveExtension();
//设置自适应拓展到缓存中
cachedAdaptiveInstance.set(instance);
} catch (Throwable t) {
createAdaptiveInstanceError = t;
throw new IllegalStateException("fail to create adaptive instance: " + t.toString(), t);
}
}
}
}
else {
throw new IllegalStateException("fail to create adaptive instance: " + createAdaptiveInstanceError.toString(), createAdaptiveInstanceError);
}
}

return (T) instance;
}

这个方法的逻辑比较简单,首先检测缓存,缓存未命中就调用createAdaptiveExtension 创建自适应拓展。
我们发现创建自适应拓展的关键在于createAdaptiveExtension 方法,它的具体实现如下:

1
2
3
4
5
6
7
8
9
private T createAdaptiveExtension() {
try {
/*拿到自适应拓展类,使用反射实例化,
然后交由injectExtension进行属性设值*/
return injectExtension((T) getAdaptiveExtensionClass().newInstance());
} catch (Exception e) {
throw new IllegalStateException("Can not create adaptive extenstion " + type + ", cause: " + e.getMessage(), e);
}
}

在这个方法中,获取自适应拓展类的方法尤为重要。
在之前我们也提到,自适应拓展类Dubbo生成的代码编译而来的。那么我们就通过代码来探究到底是如何做的。
getAdaptiveExtensionClass的具体实现如下:

1
2
3
4
5
6
7
8
9
10
private Class<?> getAdaptiveExtensionClass() {
//通过SPI获取所有的拓展类
getExtensionClasses();
if (cachedAdaptiveClass != null) {
//检测缓存,如果缓存不为空,则直接返回缓存
return cachedAdaptiveClass;
}
//创建自适应拓展类
return cachedAdaptiveClass = createAdaptiveExtensionClass();
}

这个方法首先会尝试从缓存中取自适应拓展类,如果缓存命中则直接返回,如果缓存未命中就调用createAdaptiveExtensionClass方法创建。
它的实现如下:

1
2
3
4
5
6
7
8
private Class<?> createAdaptiveExtensionClass() {
//组装代码
String code = createAdaptiveExtensionClassCode();
ClassLoader classLoader = findClassLoader();
com.alibaba.dubbo.common.compiler.Compiler compiler = ExtensionLoader.getExtensionLoader(com.alibaba.dubbo.common.compiler.Compiler.class).getAdaptiveExtension();
//编译代码,得到class返回
return compiler.compile(code, classLoader);
}

组装代码的代码非常的复杂:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
private String createAdaptiveExtensionClassCode() {
StringBuilder codeBuidler = new StringBuilder();
Method[] methods = type.getMethods();
boolean hasAdaptiveAnnotation = false;
for(Method m : methods) {
if(m.isAnnotationPresent(Adaptive.class)) {
hasAdaptiveAnnotation = true;
break;
}
}
// 完全没有Adaptive方法,则不需要生成Adaptive类
if(! hasAdaptiveAnnotation)
throw new IllegalStateException("No adaptive method on extension " + type.getName() + ", refuse to create the adaptive class!");

codeBuidler.append("package " + type.getPackage().getName() + ";");
codeBuidler.append("\nimport " + ExtensionLoader.class.getName() + ";");
codeBuidler.append("\npublic class " + type.getSimpleName() + "$Adpative" + " implements " + type.getCanonicalName() + " {");

for (Method method : methods) {
Class<?> rt = method.getReturnType();
Class<?>[] pts = method.getParameterTypes();
Class<?>[] ets = method.getExceptionTypes();

Adaptive adaptiveAnnotation = method.getAnnotation(Adaptive.class);
StringBuilder code = new StringBuilder(512);
if (adaptiveAnnotation == null) {
code.append("throw new UnsupportedOperationException(\"method ")
.append(method.toString()).append(" of interface ")
.append(type.getName()).append(" is not adaptive method!\");");
} else {
int urlTypeIndex = -1;
for (int i = 0; i < pts.length; ++i) {
if (pts[i].equals(URL.class)) {
urlTypeIndex = i;
break;
}
}
// 有类型为URL的参数
if (urlTypeIndex != -1) {
// Null Point check
String s = String.format("\nif (arg%d == null) throw new IllegalArgumentException(\"url == null\");",
urlTypeIndex);
code.append(s);

s = String.format("\n%s url = arg%d;", URL.class.getName(), urlTypeIndex);
code.append(s);
}
// 参数没有URL类型
else {
String attribMethod = null;

// 找到参数的URL属性
LBL_PTS:
for (int i = 0; i < pts.length; ++i) {
Method[] ms = pts[i].getMethods();
for (Method m : ms) {
String name = m.getName();
if ((name.startsWith("get") || name.length() > 3)
&& Modifier.isPublic(m.getModifiers())
&& !Modifier.isStatic(m.getModifiers())
&& m.getParameterTypes().length == 0
&& m.getReturnType() == URL.class) {
urlTypeIndex = i;
attribMethod = name;
break LBL_PTS;
}
}
}
if(attribMethod == null) {
throw new IllegalStateException("fail to create adative class for interface " + type.getName()
+ ": not found url parameter or url attribute in parameters of method " + method.getName());
}

// Null point check
String s = String.format("\nif (arg%d == null) throw new IllegalArgumentException(\"%s argument == null\");",
urlTypeIndex, pts[urlTypeIndex].getName());
code.append(s);
s = String.format("\nif (arg%d.%s() == null) throw new IllegalArgumentException(\"%s argument %s() == null\");",
urlTypeIndex, attribMethod, pts[urlTypeIndex].getName(), attribMethod);
code.append(s);

s = String.format("%s url = arg%d.%s();",URL.class.getName(), urlTypeIndex, attribMethod);
code.append(s);
}

String[] value = adaptiveAnnotation.value();
// 没有设置Key,则使用“扩展点接口名的点分隔 作为Key
if(value.length == 0) {
char[] charArray = type.getSimpleName().toCharArray();
StringBuilder sb = new StringBuilder(128);
for (int i = 0; i < charArray.length; i++) {
if(Character.isUpperCase(charArray[i])) {
if(i != 0) {
sb.append(".");
}
sb.append(Character.toLowerCase(charArray[i]));
}
else {
sb.append(charArray[i]);
}
}
value = new String[] {sb.toString()};
}

boolean hasInvocation = false;
for (int i = 0; i < pts.length; ++i) {
if (pts[i].getName().equals("com.alibaba.dubbo.rpc.Invocation")) {
// Null Point check
String s = String.format("\nif (arg%d == null) throw new IllegalArgumentException(\"invocation == null\");", i);
code.append(s);
s = String.format("\nString methodName = arg%d.getMethodName();", i);
code.append(s);
hasInvocation = true;
break;
}
}

String defaultExtName = cachedDefaultName;
String getNameCode = null;
for (int i = value.length - 1; i >= 0; --i) {
if(i == value.length - 1) {
if(null != defaultExtName) {
if(!"protocol".equals(value[i]))
if (hasInvocation)
getNameCode = String.format("url.getMethodParameter(methodName, \"%s\", \"%s\")", value[i], defaultExtName);
else
getNameCode = String.format("url.getParameter(\"%s\", \"%s\")", value[i], defaultExtName);
else
getNameCode = String.format("( url.getProtocol() == null ? \"%s\" : url.getProtocol() )", defaultExtName);
}
else {
if(!"protocol".equals(value[i]))
if (hasInvocation)
getNameCode = String.format("url.getMethodParameter(methodName, \"%s\", \"%s\")", value[i], defaultExtName);
else
getNameCode = String.format("url.getParameter(\"%s\")", value[i]);
else
getNameCode = "url.getProtocol()";
}
}
else {
if(!"protocol".equals(value[i]))
if (hasInvocation)
getNameCode = String.format("url.getMethodParameter(methodName, \"%s\", \"%s\")", value[i], defaultExtName);
else
getNameCode = String.format("url.getParameter(\"%s\", %s)", value[i], getNameCode);
else
getNameCode = String.format("url.getProtocol() == null ? (%s) : url.getProtocol()", getNameCode);
}
}
code.append("\nString extName = ").append(getNameCode).append(";");
// check extName == null?
String s = String.format("\nif(extName == null) " +
"throw new IllegalStateException(\"Fail to get extension(%s) name from url(\" + url.toString() + \") use keys(%s)\");",
type.getName(), Arrays.toString(value));
code.append(s);

s = String.format("\n%s extension = (%<s)%s.getExtensionLoader(%s.class).getExtension(extName);",
type.getName(), ExtensionLoader.class.getSimpleName(), type.getName());
code.append(s);

// return statement
if (!rt.equals(void.class)) {
code.append("\nreturn ");
}

s = String.format("extension.%s(", method.getName());
code.append(s);
for (int i = 0; i < pts.length; i++) {
if (i != 0)
code.append(", ");
code.append("arg").append(i);
}
code.append(");");
}

codeBuidler.append("\npublic " + rt.getCanonicalName() + " " + method.getName() + "(");
for (int i = 0; i < pts.length; i ++) {
if (i > 0) {
codeBuidler.append(", ");
}
codeBuidler.append(pts[i].getCanonicalName());
codeBuidler.append(" ");
codeBuidler.append("arg" + i);
}
codeBuidler.append(")");
if (ets.length > 0) {
codeBuidler.append(" throws ");
for (int i = 0; i < ets.length; i ++) {
if (i > 0) {
codeBuidler.append(", ");
}
codeBuidler.append(pts[i].getCanonicalName());
}
}
codeBuidler.append(" {");
codeBuidler.append(code.toString());
codeBuidler.append("\n}");
}
codeBuidler.append("\n}");
if (logger.isDebugEnabled()) {
logger.debug(codeBuidler.toString());
}
return codeBuidler.toString();
}