View Javadoc
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 }