FreeRDP
Loading...
Searching...
No Matches
websocket.c
1
20#include "websocket.h"
21#include <freerdp/log.h>
22#include "../tcp.h"
23
24#define TAG FREERDP_TAG("core.gateway.websocket")
25
26#define RESPONSE_SIZE_LIMIT (64ULL * 1024ULL * 1024ULL)
27
28struct s_websocket_context
29{
30 size_t payloadLength;
31 uint32_t maskingKey;
32 BOOL masking;
33 BOOL closeSent;
34 BYTE opcode;
35 BYTE fragmentOriginalOpcode;
36 BYTE lengthAndMaskPosition;
37 WEBSOCKET_STATE state;
38 wStream* responseStreamBuffer;
39};
40
41WINPR_ATTR_NODISCARD
42static BOOL Stream_Reset(wStream* s)
43{
44 Stream_ResetPosition(s);
45 return Stream_SetLength(s, Stream_Capacity(s));
46}
47
48WINPR_ATTR_NODISCARD
49static int websocket_write_all(BIO* bio, const BYTE* data, size_t length);
50
51BOOL websocket_context_mask_and_send(BIO* bio, wStream* sPacket, wStream* sDataPacket,
52 UINT32 maskingKey)
53{
54 const size_t len = Stream_Length(sDataPacket);
55 Stream_ResetPosition(sDataPacket);
56
57 if (!Stream_EnsureRemainingCapacity(sPacket, len))
58 return FALSE;
59
60 /* mask as much as possible with 32bit access */
61 size_t streamPos = 0;
62 for (; streamPos + 4 <= len; streamPos += 4)
63 {
64 const uint32_t data = Stream_Get_UINT32(sDataPacket);
65 Stream_Write_UINT32(sPacket, data ^ maskingKey);
66 }
67
68 /* mask the rest byte by byte */
69 for (; streamPos < len; streamPos++)
70 {
71 BYTE data = 0;
72 BYTE* partialMask = ((BYTE*)&maskingKey) + (streamPos % 4);
73 Stream_Read_UINT8(sDataPacket, data);
74 Stream_Write_UINT8(sPacket, data ^ *partialMask);
75 }
76
77 Stream_SealLength(sPacket);
78
79 ERR_clear_error();
80 const size_t size = Stream_Length(sPacket);
81 const int status = websocket_write_all(bio, Stream_Buffer(sPacket), size);
82 Stream_Free(sPacket, TRUE);
83
84 return !((status < 0) || ((size_t)status != size));
85}
86
87wStream* websocket_context_packet_new(size_t len, WEBSOCKET_OPCODE opcode, UINT32* pMaskingKey)
88{
89 WINPR_ASSERT(pMaskingKey);
90 if (len > INT_MAX)
91 return nullptr;
92
93 size_t fullLen = 0;
94 if (len < 126)
95 fullLen = len + 6; /* 2 byte "mini header" + 4 byte masking key */
96 else if (len < 0x10000)
97 fullLen = len + 8; /* 2 byte "mini header" + 2 byte length + 4 byte masking key */
98 else
99 fullLen = len + 14; /* 2 byte "mini header" + 8 byte length + 4 byte masking key */
100
101 UINT32 maskingKey = 0;
102 if (winpr_RAND(&maskingKey, sizeof(maskingKey)) < 0)
103 return nullptr;
104
105 wStream* sWS = Stream_New(nullptr, fullLen);
106 if (!sWS)
107 return nullptr;
108
109 Stream_Write_UINT8(sWS, (UINT8)(WEBSOCKET_FIN_BIT | opcode));
110 if (len < 126)
111 Stream_Write_UINT8(sWS, (UINT8)len | WEBSOCKET_MASK_BIT);
112 else if (len < 0x10000)
113 {
114 Stream_Write_UINT8(sWS, 126 | WEBSOCKET_MASK_BIT);
115 Stream_Write_UINT16_BE(sWS, (UINT16)len);
116 }
117 else
118 {
119 Stream_Write_UINT8(sWS, 127 | WEBSOCKET_MASK_BIT);
120 Stream_Write_UINT32_BE(sWS, 0); /* payload is limited to INT_MAX */
121 Stream_Write_UINT32_BE(sWS, (UINT32)len);
122 }
123 Stream_Write_UINT32(sWS, maskingKey);
124 *pMaskingKey = maskingKey;
125 return sWS;
126}
127
128BOOL websocket_context_write_wstream(websocket_context* context, BIO* bio, wStream* sPacket,
129 WEBSOCKET_OPCODE opcode)
130{
131 WINPR_ASSERT(context);
132
133 if (context->closeSent)
134 return FALSE;
135
136 if (opcode == WebsocketCloseOpcode)
137 context->closeSent = TRUE;
138
139 WINPR_ASSERT(bio);
140 WINPR_ASSERT(sPacket);
141
142 const size_t len = Stream_Length(sPacket);
143 uint32_t maskingKey = 0;
144 wStream* sWS = websocket_context_packet_new(len, opcode, &maskingKey);
145 if (!sWS)
146 return FALSE;
147
148 return websocket_context_mask_and_send(bio, sWS, sPacket, maskingKey);
149}
150
151int websocket_write_all(BIO* bio, const BYTE* data, size_t length)
152{
153 WINPR_ASSERT(bio);
154 WINPR_ASSERT(data);
155 size_t offset = 0;
156
157 if (length > INT32_MAX)
158 return -1;
159
160 while (offset < length)
161 {
162 ERR_clear_error();
163 const size_t diff = length - offset;
164 int status = BIO_write(bio, &data[offset], (int)diff);
165
166 if (status > 0)
167 offset += (size_t)status;
168 else
169 {
170 if (!BIO_should_retry(bio))
171 return -1;
172
173 if (BIO_write_blocked(bio))
174 {
175 const long rstatus = BIO_wait_write(bio, 100);
176 if (rstatus < 0)
177 return -1;
178 }
179 else if (BIO_read_blocked(bio))
180 return -2; /* Abort write, there is data that must be read */
181 else
182 USleep(100);
183 }
184 }
185
186 return (int)length;
187}
188
189int websocket_context_write(websocket_context* context, BIO* bio, const BYTE* buf, int isize,
190 WEBSOCKET_OPCODE opcode)
191{
192 WINPR_ASSERT(bio);
193 WINPR_ASSERT(buf);
194
195 if (isize < 0)
196 return -1;
197
198 wStream sbuffer = WINPR_C_ARRAY_INIT;
199 wStream* s = Stream_StaticConstInit(&sbuffer, buf, (size_t)isize);
200 if (!Stream_SetLength(s, Stream_Capacity(s)))
201 return -3;
202 if (!websocket_context_write_wstream(context, bio, s, opcode))
203 return -2;
204 return isize;
205}
206
207WINPR_ATTR_NODISCARD
208static int websocket_read_data(BIO* bio, BYTE* pBuffer, size_t size,
209 websocket_context* encodingContext)
210{
211 int status = 0;
212
213 WINPR_ASSERT(bio);
214 WINPR_ASSERT(pBuffer);
215 WINPR_ASSERT(encodingContext);
216
217 if (encodingContext->payloadLength == 0)
218 {
219 encodingContext->state = WebsocketStateOpcodeAndFin;
220 return 0;
221 }
222
223 const size_t rlen =
224 (encodingContext->payloadLength < size ? encodingContext->payloadLength : size);
225 if (rlen > INT32_MAX)
226 return -1;
227
228 ERR_clear_error();
229 status = BIO_read(bio, pBuffer, (int)rlen);
230 if (status <= 0)
231 {
232 if (status == 0)
233 {
234 if (!BIO_should_retry(bio))
235 return -1;
236 }
237 return status;
238 }
239
240 if ((size_t)status > encodingContext->payloadLength)
241 return status;
242
243 encodingContext->payloadLength -= (size_t)status;
244
245 if (encodingContext->payloadLength == 0)
246 encodingContext->state = WebsocketStateOpcodeAndFin;
247
248 return status;
249}
250
251WINPR_ATTR_NODISCARD
252static int websocket_read_wstream(BIO* bio, websocket_context* encodingContext)
253{
254 WINPR_ASSERT(bio);
255 WINPR_ASSERT(encodingContext);
256
257 wStream* s = encodingContext->responseStreamBuffer;
258 WINPR_ASSERT(s);
259
260 if (encodingContext->payloadLength == 0)
261 {
262 encodingContext->state = WebsocketStateOpcodeAndFin;
263 return 0;
264 }
265
266 if (!Stream_EnsureRemainingCapacity(s, encodingContext->payloadLength))
267 {
268 WLog_WARN(TAG,
269 "wStream::capacity [%" PRIuz "] != encodingContext::paylaodLangth [%" PRIuz "]",
270 Stream_GetRemainingCapacity(s), encodingContext->payloadLength);
271 return -1;
272 }
273
274 const int status = websocket_read_data(bio, Stream_Pointer(s), Stream_GetRemainingCapacity(s),
275 encodingContext);
276 if (status < 0)
277 return status;
278
279 if (!Stream_SafeSeek(s, (size_t)status))
280 return -1;
281
282 return status;
283}
284
285WINPR_ATTR_NODISCARD
286static BOOL websocket_reply_close(BIO* bio, websocket_context* context, wStream* s)
287{
288 WINPR_ASSERT(bio);
289
290 return websocket_context_write_wstream(context, bio, s, WebsocketCloseOpcode);
291}
292
293WINPR_ATTR_NODISCARD
294static BOOL websocket_reply_pong(BIO* bio, websocket_context* context, wStream* s)
295{
296 WINPR_ASSERT(bio);
297 WINPR_ASSERT(s);
298
299 if (Stream_GetPosition(s) != 0)
300 return websocket_context_write_wstream(context, bio, s, WebsocketPongOpcode);
301
302 return websocket_reply_close(bio, context, s);
303}
304
305WINPR_ATTR_NODISCARD
306static int websocket_handle_payload(BIO* bio, BYTE* pBuffer, size_t size,
307 websocket_context* encodingContext)
308{
309 int status = 0;
310
311 WINPR_ASSERT(bio);
312 WINPR_ASSERT(pBuffer);
313 WINPR_ASSERT(encodingContext);
314
315 const BYTE effectiveOpcode = ((encodingContext->opcode & 0xf) == WebsocketContinuationOpcode
316 ? encodingContext->fragmentOriginalOpcode & 0xf
317 : encodingContext->opcode & 0xf);
318
319 switch (effectiveOpcode)
320 {
321 case WebsocketBinaryOpcode:
322 {
323 status = websocket_read_data(bio, pBuffer, size, encodingContext);
324 if (status < 0)
325 return status;
326
327 return status;
328 }
329 case WebsocketPingOpcode:
330 {
331 status = websocket_read_wstream(bio, encodingContext);
332 if (status < 0)
333 return status;
334
335 if (encodingContext->payloadLength == 0)
336 {
337 if (!websocket_reply_pong(bio, encodingContext,
338 encodingContext->responseStreamBuffer))
339 return -1;
340 if (!Stream_Reset(encodingContext->responseStreamBuffer))
341 return -1;
342 }
343 }
344 break;
345 case WebsocketPongOpcode:
346 {
347 status = websocket_read_wstream(bio, encodingContext);
348 if (status < 0)
349 return status;
350 /* We don“t care about pong response data, discard. */
351 if (!Stream_Reset(encodingContext->responseStreamBuffer))
352 return -1;
353 }
354 break;
355 case WebsocketCloseOpcode:
356 {
357 status = websocket_read_wstream(bio, encodingContext);
358 if (status < 0)
359 return status;
360
361 if (encodingContext->payloadLength == 0)
362 {
363 if (!websocket_reply_close(bio, encodingContext,
364 encodingContext->responseStreamBuffer))
365 return -1;
366 encodingContext->closeSent = TRUE;
367 if (!Stream_Reset(encodingContext->responseStreamBuffer))
368 return -1;
369 }
370 }
371 break;
372 default:
373 WLog_WARN(TAG, "Unimplemented websocket opcode %" PRIx8 ". Dropping", effectiveOpcode);
374
375 status = websocket_read_wstream(bio, encodingContext);
376 if (status < 0)
377 return status;
378 if (!Stream_Reset(encodingContext->responseStreamBuffer))
379 return -1;
380 break;
381 }
382 /* return how many bytes have been written to pBuffer.
383 * Only WebsocketBinaryOpcode writes into it and it returns directly */
384 return 0;
385}
386
387int websocket_context_read(websocket_context* encodingContext, BIO* bio, BYTE* pBuffer, size_t size)
388{
389 int status = 0;
390 size_t effectiveDataLen = 0;
391
392 WINPR_ASSERT(bio);
393 WINPR_ASSERT(pBuffer);
394 WINPR_ASSERT(encodingContext);
395
396 while (TRUE)
397 {
398 switch (encodingContext->state)
399 {
400 case WebsocketStateOpcodeAndFin:
401 {
402 BYTE buffer[1] = WINPR_C_ARRAY_INIT;
403
404 ERR_clear_error();
405 status = BIO_read(bio, (char*)buffer, sizeof(buffer));
406 if (status <= 0)
407 return (effectiveDataLen > 0 ? WINPR_ASSERTING_INT_CAST(int, effectiveDataLen)
408 : status);
409
410 encodingContext->opcode = buffer[0];
411 if (((encodingContext->opcode & 0xf) != WebsocketContinuationOpcode) &&
412 (encodingContext->opcode & 0xf) < 0x08)
413 encodingContext->fragmentOriginalOpcode = encodingContext->opcode;
414 encodingContext->state = WebsocketStateLengthAndMasking;
415 }
416 break;
417 case WebsocketStateLengthAndMasking:
418 {
419 BYTE buffer[1] = WINPR_C_ARRAY_INIT;
420
421 ERR_clear_error();
422 status = BIO_read(bio, (char*)buffer, sizeof(buffer));
423 if (status <= 0)
424 return (effectiveDataLen > 0 ? WINPR_ASSERTING_INT_CAST(int, effectiveDataLen)
425 : status);
426
427 encodingContext->masking = ((buffer[0] & WEBSOCKET_MASK_BIT) == WEBSOCKET_MASK_BIT);
428 encodingContext->lengthAndMaskPosition = 0;
429 encodingContext->payloadLength = 0;
430 const BYTE len = buffer[0] & 0x7f;
431 if (len < 126)
432 {
433 encodingContext->payloadLength = len;
434 encodingContext->state = (encodingContext->masking ? WebSocketStateMaskingKey
435 : WebSocketStatePayload);
436 }
437 else if (len == 126)
438 encodingContext->state = WebsocketStateShortLength;
439 else
440 encodingContext->state = WebsocketStateLongLength;
441 }
442 break;
443 case WebsocketStateShortLength:
444 case WebsocketStateLongLength:
445 {
446 BYTE buffer[1] = WINPR_C_ARRAY_INIT;
447 const BYTE lenLength =
448 (encodingContext->state == WebsocketStateShortLength ? 2 : 8);
449 while (encodingContext->lengthAndMaskPosition < lenLength)
450 {
451 ERR_clear_error();
452 status = BIO_read(bio, (char*)buffer, sizeof(buffer));
453 if (status <= 0)
454 return (effectiveDataLen > 0
455 ? WINPR_ASSERTING_INT_CAST(int, effectiveDataLen)
456 : status);
457 if (status > UINT8_MAX)
458 return -1;
459 encodingContext->payloadLength =
460 (encodingContext->payloadLength) << 8 | buffer[0];
461 encodingContext->lengthAndMaskPosition +=
462 WINPR_ASSERTING_INT_CAST(BYTE, status);
463 }
464 if (encodingContext->payloadLength > RESPONSE_SIZE_LIMIT)
465 {
466 WLog_ERR(TAG, "received excessive payload size %" PRIuz ", aborting",
467 encodingContext->payloadLength);
468 return -1;
469 }
470 encodingContext->state =
471 (encodingContext->masking ? WebSocketStateMaskingKey : WebSocketStatePayload);
472 }
473 break;
474 case WebSocketStateMaskingKey:
475 {
476 WLog_WARN(
477 TAG, "Websocket Server sends data with masking key. This is against RFC 6455.");
478 return -1;
479 }
480 case WebSocketStatePayload:
481 {
482 status = websocket_handle_payload(bio, pBuffer, size, encodingContext);
483 if (status < 0)
484 return (effectiveDataLen > 0 ? WINPR_ASSERTING_INT_CAST(int, effectiveDataLen)
485 : status);
486
487 effectiveDataLen += WINPR_ASSERTING_INT_CAST(size_t, status);
488
489 if (WINPR_ASSERTING_INT_CAST(size_t, status) >= size)
490 return WINPR_ASSERTING_INT_CAST(int, effectiveDataLen);
491 pBuffer += status;
492 size -= WINPR_ASSERTING_INT_CAST(size_t, status);
493 }
494 break;
495 default:
496 break;
497 }
498 }
499 /* should be unreachable */
500}
501
502websocket_context* websocket_context_new(void)
503{
504 websocket_context* context = calloc(1, sizeof(websocket_context));
505 if (!context)
506 goto fail;
507
508 context->responseStreamBuffer = Stream_New(nullptr, 1024);
509 if (!context->responseStreamBuffer)
510 goto fail;
511
512 if (!websocket_context_reset(context))
513 goto fail;
514
515 return context;
516fail:
517 websocket_context_free(context);
518 return nullptr;
519}
520
521void websocket_context_free(websocket_context* context)
522{
523 if (!context)
524 return;
525
526 Stream_Free(context->responseStreamBuffer, TRUE);
527 free(context);
528}
529
530BOOL websocket_context_reset(websocket_context* context)
531{
532 WINPR_ASSERT(context);
533
534 context->state = WebsocketStateOpcodeAndFin;
535 return Stream_Reset(context->responseStreamBuffer);
536}