1 /*
2 * Copyright 2014 The Netty Project
3 *
4 * The Netty Project licenses this file to you under the Apache License,
5 * version 2.0 (the "License"); you may not use this file except in compliance
6 * with the License. You may obtain a copy of the License at:
7 *
8 * https://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
12 * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
13 * License for the specific language governing permissions and limitations
14 * under the License.
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.ChannelHandlerContext;
24 import io.netty.channel.ChannelPromise;
25 import io.netty.handler.codec.http.DefaultHttpRequest;
26 import io.netty.handler.codec.http.DefaultHttpResponse;
27 import io.netty.handler.codec.http.HttpHeaderNames;
28 import io.netty.handler.codec.http.HttpHeaders;
29 import io.netty.handler.codec.http.HttpRequest;
30 import io.netty.handler.codec.http.HttpResponse;
31 import io.netty.handler.codec.http.HttpResponseStatus;
32 import io.netty.handler.codec.http.LastHttpContent;
33 import io.netty.util.ReferenceCountUtil;
34 import io.netty.util.internal.ObjectUtil;
35
36 import java.util.ArrayDeque;
37 import java.util.ArrayList;
38 import java.util.Arrays;
39 import java.util.Collections;
40 import java.util.Iterator;
41 import java.util.List;
42 import java.util.Queue;
43
44 /**
45 * This handler negotiates and initializes the WebSocket Extensions.
46 *
47 * It negotiates the extensions based on the client desired order,
48 * ensures that the successfully negotiated extensions are consistent between them,
49 * and initializes the channel pipeline with the extension decoder and encoder.
50 *
51 * Find a basic implementation for compression extensions at
52 * <tt>io.netty.handler.codec.http.websocketx.extensions.compression.WebSocketServerCompressionHandler</tt>.
53 */
54 public class WebSocketServerExtensionHandler extends ChannelDuplexHandler {
55 private static final int DEFAULT_MAX_PIPELINE_DEPTH = 128;
56 private final int maxPipelineDepth;
57 private final List<WebSocketServerExtensionHandshaker> extensionHandshakers;
58
59 private final Queue<List<WebSocketServerExtension>> validExtensions = new ArrayDeque<>(4);
60
61 /**
62 * Constructor
63 *
64 * @param extensionHandshakers
65 * The extension handshaker in priority order. A handshaker could be repeated many times
66 * with fallback configuration.
67 */
68 public WebSocketServerExtensionHandler(WebSocketServerExtensionHandshaker... extensionHandshakers) {
69 this(DEFAULT_MAX_PIPELINE_DEPTH, extensionHandshakers);
70 }
71
72 /**
73 * Constructor
74 *
75 * @param maxPipelineDepth
76 * The maximum number of pipelined upgrade requests.
77 * @param extensionHandshakers
78 * The extension handshaker in priority order. A handshaker could be repeated many times
79 * with fallback configuration.
80 */
81 public WebSocketServerExtensionHandler(
82 int maxPipelineDepth, WebSocketServerExtensionHandshaker... extensionHandshakers) {
83 this.maxPipelineDepth = ObjectUtil.checkPositive(maxPipelineDepth, "maxPipelineDepth");
84 this.extensionHandshakers = Arrays.asList(checkNonEmpty(extensionHandshakers, "extensionHandshakers"));
85 }
86
87 @Override
88 public void channelRead(ChannelHandlerContext ctx, Object msg) throws Exception {
89 // JDK type checks vs non-implemented interfaces costs O(N), where
90 // N is the number of interfaces already implemented by the concrete type that's being tested.
91 // The only requirement for this call is to make HttpRequest(s) implementors to call onHttpRequestChannelRead
92 // and super.channelRead the others, but due to the O(n) cost we perform few fast-path for commonly met
93 // singleton and/or concrete types, to save performing such slow type checks.
94 if (msg != LastHttpContent.EMPTY_LAST_CONTENT) {
95 if (msg instanceof DefaultHttpRequest) {
96 // fast-path
97 onHttpRequestChannelRead(ctx, (DefaultHttpRequest) msg);
98 } else if (msg instanceof HttpRequest) {
99 // slow path
100 onHttpRequestChannelRead(ctx, (HttpRequest) msg);
101 } else {
102 super.channelRead(ctx, msg);
103 }
104 } else {
105 super.channelRead(ctx, msg);
106 }
107 }
108
109 /**
110 * This is a method exposed to perform fail-fast checks of user-defined http types.<p>
111 * eg:<br>
112 * If the user has defined a specific {@link HttpRequest} type i.e.{@code CustomHttpRequest} and
113 * {@link #channelRead} can receive {@link LastHttpContent#EMPTY_LAST_CONTENT} {@code msg}
114 * types too, can override it like this:
115 * <pre>
116 * public void channelRead(ChannelHandlerContext ctx, Object msg) throws Exception {
117 * if (msg != LastHttpContent.EMPTY_LAST_CONTENT) {
118 * if (msg instanceof CustomHttpRequest) {
119 * onHttpRequestChannelRead(ctx, (CustomHttpRequest) msg);
120 * } else {
121 * // if it's handling other HttpRequest types it MUST use onHttpRequestChannelRead again
122 * // or have to delegate it to super.channelRead (that can perform redundant checks).
123 * // If msg is not implementing HttpRequest, it can call ctx.fireChannelRead(msg) on it
124 * // ...
125 * super.channelRead(ctx, msg);
126 * }
127 * } else {
128 * // given that msg isn't a HttpRequest type we can just skip calling super.channelRead
129 * ctx.fireChannelRead(msg);
130 * }
131 * }
132 * </pre>
133 * <strong>IMPORTANT:</strong>
134 * It already call {@code super.channelRead(ctx, request)} before returning.
135 */
136 protected void onHttpRequestChannelRead(ChannelHandlerContext ctx, HttpRequest request) throws Exception {
137 if (maxPipelineDepth <= validExtensions.size()) {
138 ReferenceCountUtil.release(request);
139 ctx.close();
140 throw new IllegalStateException("maxPipelineDepth exceeded: " + maxPipelineDepth);
141 }
142
143 List<WebSocketServerExtension> validExtensionsList = null;
144
145 if (WebSocketExtensionUtil.isWebsocketUpgrade(request.headers())) {
146 String extensionsHeader = request.headers().getAsString(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS);
147
148 if (extensionsHeader != null) {
149 List<WebSocketExtensionData> extensions =
150 WebSocketExtensionUtil.extractExtensions(extensionsHeader);
151 int rsv = 0;
152
153 for (WebSocketExtensionData extensionData : extensions) {
154 Iterator<WebSocketServerExtensionHandshaker> extensionHandshakersIterator =
155 extensionHandshakers.iterator();
156 WebSocketServerExtension validExtension = null;
157
158 while (validExtension == null && extensionHandshakersIterator.hasNext()) {
159 WebSocketServerExtensionHandshaker extensionHandshaker =
160 extensionHandshakersIterator.next();
161 validExtension = extensionHandshaker.handshakeExtension(extensionData);
162 }
163
164 if (validExtension != null && ((validExtension.rsv() & rsv) == 0)) {
165 if (validExtensionsList == null) {
166 validExtensionsList = new ArrayList<WebSocketServerExtension>(1);
167 }
168 rsv = rsv | validExtension.rsv();
169 validExtensionsList.add(validExtension);
170 }
171 }
172 }
173 }
174
175 if (validExtensionsList == null) {
176 validExtensionsList = Collections.emptyList();
177 }
178 validExtensions.offer(validExtensionsList);
179
180 super.channelRead(ctx, request);
181 }
182
183 @Override
184 public void write(final ChannelHandlerContext ctx, Object msg, ChannelPromise promise) throws Exception {
185 if (msg != Unpooled.EMPTY_BUFFER && !(msg instanceof ByteBuf)) {
186 if (msg instanceof DefaultHttpResponse) {
187 onHttpResponseWrite(ctx, (DefaultHttpResponse) msg, promise);
188 } else if (msg instanceof HttpResponse) {
189 onHttpResponseWrite(ctx, (HttpResponse) msg, promise);
190 } else {
191 super.write(ctx, msg, promise);
192 }
193 } else {
194 super.write(ctx, msg, promise);
195 }
196 }
197
198 /**
199 * This is a method exposed to perform fail-fast checks of user-defined http types.<p>
200 * eg:<br>
201 * If the user has defined a specific {@link HttpResponse} type i.e.{@code CustomHttpResponse} and
202 * {@link #write} can receive {@link ByteBuf} {@code msg} types too, it can be overridden like this:
203 * <pre>
204 * public void write(final ChannelHandlerContext ctx, Object msg, ChannelPromise promise) throws Exception {
205 * if (msg != Unpooled.EMPTY_BUFFER && !(msg instanceof ByteBuf)) {
206 * if (msg instanceof CustomHttpResponse) {
207 * onHttpResponseWrite(ctx, (CustomHttpResponse) msg, promise);
208 * } else {
209 * // if it's handling other HttpResponse types it MUST use onHttpResponseWrite again
210 * // or have to delegate it to super.write (that can perform redundant checks).
211 * // If msg is not implementing HttpResponse, it can call ctx.write(msg, promise) on it
212 * // ...
213 * super.write(ctx, msg, promise);
214 * }
215 * } else {
216 * // given that msg isn't a HttpResponse type we can just skip calling super.write
217 * ctx.write(msg, promise);
218 * }
219 * }
220 * </pre>
221 * <strong>IMPORTANT:</strong>
222 * It already call {@code super.write(ctx, response, promise)} before returning.
223 */
224 protected void onHttpResponseWrite(ChannelHandlerContext ctx, HttpResponse response, ChannelPromise promise)
225 throws Exception {
226 List<WebSocketServerExtension> validExtensionsList = validExtensions.poll();
227 // checking the status is faster than looking at headers so we do this first
228 if (HttpResponseStatus.SWITCHING_PROTOCOLS.equals(response.status())) {
229 handlePotentialUpgrade(ctx, promise, response, validExtensionsList);
230 }
231 super.write(ctx, response, promise);
232 }
233
234 /**
235 * Returns {@code true} if WebSocket extensions negotiated from the request should be acknowledged for this
236 * response.
237 * <p>
238 * This method can be overridden to make the final extension negotiation decision when the handshake response is
239 * written.
240 */
241 protected boolean isExtensionNegotiationEnabled(ChannelHandlerContext ctx, HttpResponse response) {
242 return true;
243 }
244
245 private void handlePotentialUpgrade(final ChannelHandlerContext ctx,
246 ChannelPromise promise, HttpResponse httpResponse,
247 final List<WebSocketServerExtension> validExtensionsList) {
248 HttpHeaders headers = httpResponse.headers();
249
250 if (WebSocketExtensionUtil.isWebsocketUpgrade(headers)) {
251 if (isExtensionNegotiationEnabled(ctx, httpResponse)
252 && validExtensionsList != null && !validExtensionsList.isEmpty()) {
253 String headerValue = headers.getAsString(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS);
254 List<WebSocketExtensionData> extraExtensions =
255 new ArrayList<>(extensionHandshakers.size());
256 for (WebSocketServerExtension extension : validExtensionsList) {
257 extraExtensions.add(extension.newReponseData());
258 }
259 String newHeaderValue = WebSocketExtensionUtil
260 .computeMergeExtensionsHeaderValue(headerValue, extraExtensions);
261 promise.addListener(future -> {
262 if (future.isSuccess()) {
263 for (WebSocketServerExtension extension : validExtensionsList) {
264 WebSocketExtensionDecoder decoder = extension.newExtensionDecoder();
265 WebSocketExtensionEncoder encoder = extension.newExtensionEncoder();
266 String name = ctx.name();
267 ctx.pipeline()
268 .addAfter(name, decoder.getClass().getName(), decoder)
269 .addAfter(name, encoder.getClass().getName(), encoder);
270 }
271 }
272 });
273
274 headers.set(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS, newHeaderValue);
275 }
276
277 promise.addListener(future -> {
278 if (future.isSuccess()) {
279 ctx.pipeline().remove(WebSocketServerExtensionHandler.this);
280 }
281 });
282 }
283 }
284 }