1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16 package io.netty.handler.codec.http.websocketx.extensions;
17
18 import static io.netty.util.internal.ObjectUtil.checkNonEmpty;
19
20 import io.netty.buffer.ByteBuf;
21 import io.netty.buffer.Unpooled;
22 import io.netty.channel.ChannelDuplexHandler;
23 import io.netty.channel.ChannelFuture;
24 import io.netty.channel.ChannelFutureListener;
25 import io.netty.channel.ChannelHandlerContext;
26 import io.netty.channel.ChannelPromise;
27 import io.netty.handler.codec.http.DefaultHttpRequest;
28 import io.netty.handler.codec.http.DefaultHttpResponse;
29 import io.netty.handler.codec.http.HttpHeaderNames;
30 import io.netty.handler.codec.http.HttpHeaders;
31 import io.netty.handler.codec.http.HttpRequest;
32 import io.netty.handler.codec.http.HttpResponse;
33 import io.netty.handler.codec.http.HttpResponseStatus;
34 import io.netty.handler.codec.http.LastHttpContent;
35 import io.netty.util.ReferenceCountUtil;
36 import io.netty.util.internal.ObjectUtil;
37
38 import java.util.ArrayDeque;
39 import java.util.ArrayList;
40 import java.util.Arrays;
41 import java.util.Collections;
42 import java.util.Iterator;
43 import java.util.List;
44 import java.util.Queue;
45
46
47
48
49
50
51
52
53
54
55
56 public class WebSocketServerExtensionHandler extends ChannelDuplexHandler {
57 private static final int DEFAULT_MAX_PIPELINE_DEPTH = 128;
58 private final int maxPipelineDepth;
59 private final List<WebSocketServerExtensionHandshaker> extensionHandshakers;
60
61 private final Queue<List<WebSocketServerExtension>> validExtensions =
62 new ArrayDeque<List<WebSocketServerExtension>>(4);
63
64
65
66
67
68
69
70
71 public WebSocketServerExtensionHandler(WebSocketServerExtensionHandshaker... extensionHandshakers) {
72 this(DEFAULT_MAX_PIPELINE_DEPTH, extensionHandshakers);
73 }
74
75
76
77
78
79
80
81
82
83
84 public WebSocketServerExtensionHandler(
85 int maxPipelineDepth, WebSocketServerExtensionHandshaker... extensionHandshakers) {
86 this.maxPipelineDepth = ObjectUtil.checkPositive(maxPipelineDepth, "maxPipelineDepth");
87 this.extensionHandshakers = Arrays.asList(checkNonEmpty(extensionHandshakers, "extensionHandshakers"));
88 }
89
90 @Override
91 public void channelRead(ChannelHandlerContext ctx, Object msg) throws Exception {
92
93
94
95
96
97 if (msg != LastHttpContent.EMPTY_LAST_CONTENT) {
98 if (msg instanceof DefaultHttpRequest) {
99
100 onHttpRequestChannelRead(ctx, (DefaultHttpRequest) msg);
101 } else if (msg instanceof HttpRequest) {
102
103 onHttpRequestChannelRead(ctx, (HttpRequest) msg);
104 } else {
105 super.channelRead(ctx, msg);
106 }
107 } else {
108 super.channelRead(ctx, msg);
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 protected void onHttpRequestChannelRead(ChannelHandlerContext ctx, HttpRequest request) throws Exception {
140 if (maxPipelineDepth <= validExtensions.size()) {
141 ReferenceCountUtil.release(request);
142 ctx.close();
143 throw new IllegalStateException("maxPipelineDepth exceeded: " + maxPipelineDepth);
144 }
145
146 List<WebSocketServerExtension> validExtensionsList = null;
147
148 if (WebSocketExtensionUtil.isWebsocketUpgrade(request.headers())) {
149 String extensionsHeader = request.headers().getAsString(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS);
150
151 if (extensionsHeader != null) {
152 List<WebSocketExtensionData> extensions =
153 WebSocketExtensionUtil.extractExtensions(extensionsHeader);
154 int rsv = 0;
155
156 for (WebSocketExtensionData extensionData : extensions) {
157 Iterator<WebSocketServerExtensionHandshaker> extensionHandshakersIterator =
158 extensionHandshakers.iterator();
159 WebSocketServerExtension validExtension = null;
160
161 while (validExtension == null && extensionHandshakersIterator.hasNext()) {
162 WebSocketServerExtensionHandshaker extensionHandshaker =
163 extensionHandshakersIterator.next();
164 validExtension = extensionHandshaker.handshakeExtension(extensionData);
165 }
166
167 if (validExtension != null && ((validExtension.rsv() & rsv) == 0)) {
168 if (validExtensionsList == null) {
169 validExtensionsList = new ArrayList<WebSocketServerExtension>(1);
170 }
171 rsv = rsv | validExtension.rsv();
172 validExtensionsList.add(validExtension);
173 }
174 }
175 }
176 }
177
178 if (validExtensionsList == null) {
179 validExtensionsList = Collections.emptyList();
180 }
181 validExtensions.offer(validExtensionsList);
182
183 super.channelRead(ctx, request);
184 }
185
186 @Override
187 public void write(final ChannelHandlerContext ctx, Object msg, ChannelPromise promise) throws Exception {
188 if (msg != Unpooled.EMPTY_BUFFER && !(msg instanceof ByteBuf)) {
189 if (msg instanceof DefaultHttpResponse) {
190 onHttpResponseWrite(ctx, (DefaultHttpResponse) msg, promise);
191 } else if (msg instanceof HttpResponse) {
192 onHttpResponseWrite(ctx, (HttpResponse) msg, promise);
193 } else {
194 super.write(ctx, msg, promise);
195 }
196 } else {
197 super.write(ctx, msg, promise);
198 }
199 }
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227 protected void onHttpResponseWrite(ChannelHandlerContext ctx, HttpResponse response, ChannelPromise promise)
228 throws Exception {
229 List<WebSocketServerExtension> validExtensionsList = validExtensions.poll();
230
231 if (HttpResponseStatus.SWITCHING_PROTOCOLS.equals(response.status())) {
232 handlePotentialUpgrade(ctx, promise, response, validExtensionsList);
233 }
234 super.write(ctx, response, promise);
235 }
236
237 private void handlePotentialUpgrade(final ChannelHandlerContext ctx,
238 ChannelPromise promise, HttpResponse httpResponse,
239 final List<WebSocketServerExtension> validExtensionsList) {
240 HttpHeaders headers = httpResponse.headers();
241
242 if (WebSocketExtensionUtil.isWebsocketUpgrade(headers)) {
243 if (validExtensionsList != null && !validExtensionsList.isEmpty()) {
244 String headerValue = headers.getAsString(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS);
245 List<WebSocketExtensionData> extraExtensions =
246 new ArrayList<WebSocketExtensionData>(extensionHandshakers.size());
247 for (WebSocketServerExtension extension : validExtensionsList) {
248 extraExtensions.add(extension.newReponseData());
249 }
250 String newHeaderValue = WebSocketExtensionUtil
251 .computeMergeExtensionsHeaderValue(headerValue, extraExtensions);
252 promise.addListener(new ChannelFutureListener() {
253 @Override
254 public void operationComplete(ChannelFuture future) {
255 if (future.isSuccess()) {
256 for (WebSocketServerExtension extension : validExtensionsList) {
257 WebSocketExtensionDecoder decoder = extension.newExtensionDecoder();
258 WebSocketExtensionEncoder encoder = extension.newExtensionEncoder();
259 String name = ctx.name();
260 ctx.pipeline()
261 .addAfter(name, decoder.getClass().getName(), decoder)
262 .addAfter(name, encoder.getClass().getName(), encoder);
263 }
264 }
265 }
266 });
267
268 if (newHeaderValue != null) {
269 headers.set(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS, newHeaderValue);
270 }
271 }
272
273 promise.addListener(new ChannelFutureListener() {
274 @Override
275 public void operationComplete(ChannelFuture future) {
276 if (future.isSuccess()) {
277 ctx.pipeline().remove(WebSocketServerExtensionHandler.this);
278 }
279 }
280 });
281 }
282 }
283 }