| | | 1 | | // Licensed to the .NET Foundation under one or more agreements. |
| | | 2 | | // The .NET Foundation licenses this file to you under the MIT license. |
| | | 3 | | |
| | | 4 | | using System.Buffers; |
| | | 5 | | using System.Diagnostics; |
| | | 6 | | using static System.IO.Compression.ZLibNative; |
| | | 7 | | |
| | | 8 | | namespace System.Net.WebSockets.Compression |
| | | 9 | | { |
| | | 10 | | /// <summary> |
| | | 11 | | /// Provides a wrapper around the ZLib decompression API. |
| | | 12 | | /// </summary> |
| | | 13 | | internal sealed class WebSocketInflater : IDisposable |
| | | 14 | | { |
| | | 15 | | internal const int FlushMarkerLength = 4; |
| | 195 | 16 | | internal static ReadOnlySpan<byte> FlushMarker => [0x00, 0x00, 0xFF, 0xFF]; |
| | | 17 | | |
| | | 18 | | private readonly int _windowBits; |
| | | 19 | | private ZLibStreamHandle? _stream; |
| | | 20 | | private readonly bool _persisted; |
| | | 21 | | |
| | | 22 | | /// <summary> |
| | | 23 | | /// There is no way of knowing, when decoding data, if the underlying inflater |
| | | 24 | | /// has flushed all outstanding data to consumer other than to provide a buffer |
| | | 25 | | /// and see whether any bytes are written. There are cases when the consumers |
| | | 26 | | /// provide a buffer exactly the size of the uncompressed data and in this case |
| | | 27 | | /// to avoid requiring another read we will use this field. |
| | | 28 | | /// </summary> |
| | | 29 | | private byte? _remainingByte; |
| | | 30 | | |
| | | 31 | | /// <summary> |
| | | 32 | | /// The last added bytes to the inflater were part of the final |
| | | 33 | | /// payload for the message being sent. |
| | | 34 | | /// </summary> |
| | | 35 | | private bool _endOfMessage; |
| | | 36 | | |
| | | 37 | | private byte[]? _buffer; |
| | | 38 | | |
| | | 39 | | /// <summary> |
| | | 40 | | /// The position for the next unconsumed byte in the inflate buffer. |
| | | 41 | | /// </summary> |
| | | 42 | | private int _position; |
| | | 43 | | |
| | | 44 | | /// <summary> |
| | | 45 | | /// How many unconsumed bytes are left in the inflate buffer. |
| | | 46 | | /// </summary> |
| | | 47 | | private int _available; |
| | | 48 | | |
| | 2025 | 49 | | internal WebSocketInflater(int windowBits, bool persisted) |
| | 2025 | 50 | | { |
| | 2025 | 51 | | _windowBits = -windowBits; // Negative for raw deflate |
| | 2025 | 52 | | _persisted = persisted; |
| | 2025 | 53 | | } |
| | | 54 | | |
| | 38 | 55 | | public Memory<byte> Memory => _buffer.AsMemory(_position + _available); |
| | | 56 | | |
| | 336 | 57 | | public Span<byte> Span => _buffer.AsSpan(_position + _available); |
| | | 58 | | |
| | | 59 | | public void Dispose() |
| | 2025 | 60 | | { |
| | 2025 | 61 | | _stream?.Dispose(); |
| | 2025 | 62 | | ReleaseBuffer(); |
| | 2025 | 63 | | } |
| | | 64 | | |
| | | 65 | | /// <summary> |
| | | 66 | | /// Initializes the inflater by allocating a buffer so the websocket can receive directly onto it. |
| | | 67 | | /// </summary> |
| | | 68 | | /// <param name="payloadLength">the length of the message payload</param> |
| | | 69 | | /// <param name="userBufferLength">the length of the buffer where the payload will be inflated</param> |
| | | 70 | | public void Prepare(long payloadLength, int userBufferLength) |
| | 165 | 71 | | { |
| | 165 | 72 | | if (_buffer is not null) |
| | 14 | 73 | | { |
| | 14 | 74 | | Debug.Assert(_available > 0); |
| | | 75 | | |
| | 14 | 76 | | _buffer.AsSpan(_position, _available).CopyTo(_buffer); |
| | 14 | 77 | | _position = 0; |
| | 14 | 78 | | } |
| | | 79 | | else |
| | 151 | 80 | | { |
| | | 81 | | // Rent a buffer as close to the size of the user buffer as possible. |
| | | 82 | | // If the payload is smaller than the user buffer, rent only as much as we need. |
| | 151 | 83 | | _buffer = ArrayPool<byte>.Shared.Rent((int)Math.Min(userBufferLength, payloadLength)); |
| | 151 | 84 | | } |
| | 165 | 85 | | } |
| | | 86 | | |
| | | 87 | | public void AddBytes(int totalBytesReceived, bool endOfMessage) |
| | 808 | 88 | | { |
| | 808 | 89 | | Debug.Assert(totalBytesReceived == 0 || _buffer is not null, "Prepare must be called."); |
| | | 90 | | |
| | 808 | 91 | | _available += totalBytesReceived; |
| | 808 | 92 | | _endOfMessage = endOfMessage; |
| | | 93 | | |
| | 808 | 94 | | if (endOfMessage) |
| | 158 | 95 | | { |
| | 158 | 96 | | if (_buffer is null) |
| | 113 | 97 | | { |
| | 113 | 98 | | Debug.Assert(_available == 0); |
| | | 99 | | |
| | 113 | 100 | | _buffer = ArrayPool<byte>.Shared.Rent(FlushMarkerLength); |
| | 113 | 101 | | _available = FlushMarkerLength; |
| | 113 | 102 | | FlushMarker.CopyTo(_buffer); |
| | 113 | 103 | | } |
| | | 104 | | else |
| | 45 | 105 | | { |
| | 45 | 106 | | if (_buffer.Length < _available + FlushMarkerLength) |
| | 1 | 107 | | { |
| | 1 | 108 | | byte[] newBuffer = ArrayPool<byte>.Shared.Rent(_available + FlushMarkerLength); |
| | 1 | 109 | | _buffer.AsSpan(0, _available).CopyTo(newBuffer); |
| | | 110 | | |
| | 1 | 111 | | byte[] toReturn = _buffer; |
| | 1 | 112 | | _buffer = newBuffer; |
| | | 113 | | |
| | 1 | 114 | | ArrayPool<byte>.Shared.Return(toReturn); |
| | 1 | 115 | | } |
| | | 116 | | |
| | 45 | 117 | | FlushMarker.CopyTo(_buffer.AsSpan(_available)); |
| | 45 | 118 | | _available += FlushMarkerLength; |
| | 45 | 119 | | } |
| | 158 | 120 | | } |
| | 808 | 121 | | } |
| | | 122 | | |
| | | 123 | | /// <summary> |
| | | 124 | | /// Inflates the last receive payload into the provided buffer. |
| | | 125 | | /// </summary> |
| | | 126 | | public unsafe bool Inflate(Span<byte> output, out int written) |
| | 682 | 127 | | { |
| | 682 | 128 | | _stream ??= CreateInflater(); |
| | | 129 | | |
| | 682 | 130 | | bool streamEnded = false; |
| | | 131 | | |
| | 682 | 132 | | if (_available > 0 && output.Length > 0) |
| | 152 | 133 | | { |
| | | 134 | | int consumed; |
| | | 135 | | |
| | 152 | 136 | | fixed (byte* bufferPtr = _buffer) |
| | 152 | 137 | | { |
| | 152 | 138 | | _stream.NextIn = (IntPtr)(bufferPtr + _position); |
| | 152 | 139 | | _stream.AvailIn = (uint)_available; |
| | | 140 | | |
| | 152 | 141 | | written = Inflate(_stream, output, FlushCode.NoFlush, out streamEnded); |
| | 127 | 142 | | consumed = _available - (int)_stream.AvailIn; |
| | 127 | 143 | | } |
| | | 144 | | |
| | 127 | 145 | | _position += consumed; |
| | 127 | 146 | | _available -= consumed; |
| | 127 | 147 | | } |
| | | 148 | | else |
| | 530 | 149 | | { |
| | 530 | 150 | | written = 0; |
| | 530 | 151 | | } |
| | | 152 | | |
| | 657 | 153 | | if (_available == 0) |
| | 653 | 154 | | { |
| | 653 | 155 | | ReleaseBuffer(); |
| | 653 | 156 | | return _endOfMessage ? Finish(output, ref written) : true; |
| | | 157 | | } |
| | | 158 | | |
| | 4 | 159 | | if (streamEnded && _available > 0) |
| | 4 | 160 | | { |
| | | 161 | | // zlib reached the end of the DEFLATE stream (a BFINAL=1 final block) while compressed |
| | | 162 | | // bytes still remain that it will never consume. permessage-deflate messages are not |
| | | 163 | | // expected to contain a final block; continuing would make no forward progress (the |
| | | 164 | | // inflater would report empty results forever and hang the caller's receive loop), so |
| | | 165 | | // reject the message. |
| | 4 | 166 | | throw new WebSocketException(SR.net_WebSockets_DataAfterBFinal); |
| | | 167 | | } |
| | | 168 | | |
| | 0 | 169 | | return false; |
| | 653 | 170 | | } |
| | | 171 | | |
| | | 172 | | /// <summary> |
| | | 173 | | /// Finishes the decoding by flushing any outstanding data to the output. |
| | | 174 | | /// </summary> |
| | | 175 | | /// <returns>true if the flush completed, false to indicate that there is more outstanding data.</returns> |
| | | 176 | | private bool Finish(Span<byte> output, ref int written) |
| | 17 | 177 | | { |
| | 17 | 178 | | Debug.Assert(_stream is not null && _stream.AvailIn == 0); |
| | 17 | 179 | | Debug.Assert(_available == 0); |
| | | 180 | | |
| | 17 | 181 | | if (_remainingByte is not null) |
| | 0 | 182 | | { |
| | 0 | 183 | | if (output.Length == written) |
| | 0 | 184 | | { |
| | 0 | 185 | | return false; |
| | | 186 | | } |
| | 0 | 187 | | output[written] = _remainingByte.GetValueOrDefault(); |
| | 0 | 188 | | _remainingByte = null; |
| | 0 | 189 | | written += 1; |
| | 0 | 190 | | } |
| | | 191 | | |
| | | 192 | | // If we have more space in the output, try to inflate |
| | 17 | 193 | | if (output.Length > written) |
| | 17 | 194 | | { |
| | 17 | 195 | | written += Inflate(_stream, output[written..], FlushCode.SyncFlush, out _); |
| | 17 | 196 | | } |
| | | 197 | | |
| | | 198 | | // After inflate, if we have more space in the output then it means that we |
| | | 199 | | // have finished. Otherwise we need to manually check for more data. |
| | 17 | 200 | | if (written < output.Length || IsFinished(_stream, out _remainingByte)) |
| | 17 | 201 | | { |
| | 17 | 202 | | if (!_persisted) |
| | 0 | 203 | | { |
| | 0 | 204 | | _stream.Dispose(); |
| | 0 | 205 | | _stream = null; |
| | 0 | 206 | | } |
| | 17 | 207 | | return true; |
| | | 208 | | } |
| | | 209 | | |
| | 0 | 210 | | return false; |
| | 17 | 211 | | } |
| | | 212 | | |
| | | 213 | | private void ReleaseBuffer() |
| | 2678 | 214 | | { |
| | 2678 | 215 | | if (_buffer is byte[] toReturn) |
| | 264 | 216 | | { |
| | 264 | 217 | | _buffer = null; |
| | 264 | 218 | | _available = 0; |
| | 264 | 219 | | _position = 0; |
| | | 220 | | |
| | 264 | 221 | | ArrayPool<byte>.Shared.Return(toReturn); |
| | 264 | 222 | | } |
| | 2678 | 223 | | } |
| | | 224 | | |
| | | 225 | | private static bool IsFinished(ZLibStreamHandle stream, out byte? remainingByte) |
| | 0 | 226 | | { |
| | | 227 | | // There is no other way to make sure that we've consumed all data |
| | | 228 | | // but to try to inflate again with at least one byte of output buffer. |
| | 0 | 229 | | byte b = 0; |
| | 0 | 230 | | if (Inflate(stream, new Span<byte>(ref b), FlushCode.SyncFlush, out _) == 0) |
| | 0 | 231 | | { |
| | 0 | 232 | | remainingByte = null; |
| | 0 | 233 | | return true; |
| | | 234 | | } |
| | | 235 | | |
| | 0 | 236 | | remainingByte = b; |
| | 0 | 237 | | return false; |
| | 0 | 238 | | } |
| | | 239 | | |
| | | 240 | | private static unsafe int Inflate(ZLibStreamHandle stream, Span<byte> destination, FlushCode flushCode, out bool |
| | 169 | 241 | | { |
| | 169 | 242 | | Debug.Assert(destination.Length > 0); |
| | | 243 | | ErrorCode errorCode; |
| | | 244 | | |
| | 169 | 245 | | fixed (byte* bufPtr = destination) |
| | 169 | 246 | | { |
| | 169 | 247 | | stream.NextOut = (IntPtr)bufPtr; |
| | 169 | 248 | | stream.AvailOut = (uint)destination.Length; |
| | | 249 | | |
| | 169 | 250 | | errorCode = stream.Inflate(flushCode); |
| | | 251 | | |
| | 169 | 252 | | if (errorCode is ErrorCode.Ok or ErrorCode.StreamEnd or ErrorCode.BufError) |
| | 144 | 253 | | { |
| | 144 | 254 | | streamEnded = errorCode == ErrorCode.StreamEnd; |
| | 144 | 255 | | return destination.Length - (int)stream.AvailOut; |
| | | 256 | | } |
| | 25 | 257 | | } |
| | | 258 | | |
| | 25 | 259 | | string message = errorCode switch |
| | 25 | 260 | | { |
| | 0 | 261 | | ErrorCode.MemError => SR.ZLibErrorNotEnoughMemory, |
| | 25 | 262 | | ErrorCode.DataError => SR.ZLibUnsupportedCompression, |
| | 0 | 263 | | ErrorCode.StreamError => SR.ZLibErrorInconsistentStream, |
| | 0 | 264 | | _ => SR.Format(SR.ZLibErrorUnexpected, (int)errorCode) |
| | 25 | 265 | | }; |
| | 25 | 266 | | throw new WebSocketException(message); |
| | 144 | 267 | | } |
| | | 268 | | |
| | | 269 | | private ZLibStreamHandle CreateInflater() |
| | 241 | 270 | | { |
| | | 271 | | try |
| | 241 | 272 | | { |
| | 241 | 273 | | return ZLibStreamHandle.CreateForInflate(_windowBits); |
| | | 274 | | } |
| | 0 | 275 | | catch (Exception ex) |
| | 0 | 276 | | { |
| | 0 | 277 | | throw new WebSocketException(ex.Message, ex.InnerException); |
| | | 278 | | } |
| | 241 | 279 | | } |
| | | 280 | | } |
| | | 281 | | } |
| | | 282 | | |