21#include <freerdp/log.h>
24#define TAG FREERDP_TAG("core.gateway.websocket")
26#define RESPONSE_SIZE_LIMIT (64ULL * 1024ULL * 1024ULL)
28struct s_websocket_context
35 BYTE fragmentOriginalOpcode;
36 BYTE lengthAndMaskPosition;
37 WEBSOCKET_STATE state;
42static BOOL Stream_Reset(
wStream* s)
44 Stream_ResetPosition(s);
45 return Stream_SetLength(s, Stream_Capacity(s));
49static int websocket_write_all(BIO* bio,
const BYTE* data,
size_t length);
51BOOL websocket_context_mask_and_send(BIO* bio,
wStream* sPacket,
wStream* sDataPacket,
54 const size_t len = Stream_Length(sDataPacket);
55 Stream_ResetPosition(sDataPacket);
57 if (!Stream_EnsureRemainingCapacity(sPacket, len))
62 for (; streamPos + 4 <= len; streamPos += 4)
64 const uint32_t data = Stream_Get_UINT32(sDataPacket);
65 Stream_Write_UINT32(sPacket, data ^ maskingKey);
69 for (; streamPos < len; streamPos++)
72 BYTE* partialMask = ((BYTE*)&maskingKey) + (streamPos % 4);
73 Stream_Read_UINT8(sDataPacket, data);
74 Stream_Write_UINT8(sPacket, data ^ *partialMask);
77 Stream_SealLength(sPacket);
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);
84 return !((status < 0) || ((
size_t)status != size));
87wStream* websocket_context_packet_new(
size_t len, WEBSOCKET_OPCODE opcode, UINT32* pMaskingKey)
89 WINPR_ASSERT(pMaskingKey);
96 else if (len < 0x10000)
101 UINT32 maskingKey = 0;
102 if (winpr_RAND(&maskingKey,
sizeof(maskingKey)) < 0)
105 wStream* sWS = Stream_New(
nullptr, fullLen);
109 Stream_Write_UINT8(sWS, (UINT8)(WEBSOCKET_FIN_BIT | opcode));
111 Stream_Write_UINT8(sWS, (UINT8)len | WEBSOCKET_MASK_BIT);
112 else if (len < 0x10000)
114 Stream_Write_UINT8(sWS, 126 | WEBSOCKET_MASK_BIT);
115 Stream_Write_UINT16_BE(sWS, (UINT16)len);
119 Stream_Write_UINT8(sWS, 127 | WEBSOCKET_MASK_BIT);
120 Stream_Write_UINT32_BE(sWS, 0);
121 Stream_Write_UINT32_BE(sWS, (UINT32)len);
123 Stream_Write_UINT32(sWS, maskingKey);
124 *pMaskingKey = maskingKey;
128BOOL websocket_context_write_wstream(websocket_context* context, BIO* bio,
wStream* sPacket,
129 WEBSOCKET_OPCODE opcode)
131 WINPR_ASSERT(context);
133 if (context->closeSent)
136 if (opcode == WebsocketCloseOpcode)
137 context->closeSent = TRUE;
140 WINPR_ASSERT(sPacket);
142 const size_t len = Stream_Length(sPacket);
143 uint32_t maskingKey = 0;
144 wStream* sWS = websocket_context_packet_new(len, opcode, &maskingKey);
148 return websocket_context_mask_and_send(bio, sWS, sPacket, maskingKey);
151int websocket_write_all(BIO* bio,
const BYTE* data,
size_t length)
157 if (length > INT32_MAX)
160 while (offset < length)
163 const size_t diff = length - offset;
164 int status = BIO_write(bio, &data[offset], (
int)diff);
167 offset += (size_t)status;
170 if (!BIO_should_retry(bio))
173 if (BIO_write_blocked(bio))
175 const long rstatus = BIO_wait_write(bio, 100);
179 else if (BIO_read_blocked(bio))
189int websocket_context_write(websocket_context* context, BIO* bio,
const BYTE* buf,
int isize,
190 WEBSOCKET_OPCODE opcode)
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)))
202 if (!websocket_context_write_wstream(context, bio, s, opcode))
208static int websocket_read_data(BIO* bio, BYTE* pBuffer,
size_t size,
209 websocket_context* encodingContext)
214 WINPR_ASSERT(pBuffer);
215 WINPR_ASSERT(encodingContext);
217 if (encodingContext->payloadLength == 0)
219 encodingContext->state = WebsocketStateOpcodeAndFin;
224 (encodingContext->payloadLength < size ? encodingContext->payloadLength : size);
225 if (rlen > INT32_MAX)
229 status = BIO_read(bio, pBuffer, (
int)rlen);
234 if (!BIO_should_retry(bio))
240 if ((
size_t)status > encodingContext->payloadLength)
243 encodingContext->payloadLength -= (size_t)status;
245 if (encodingContext->payloadLength == 0)
246 encodingContext->state = WebsocketStateOpcodeAndFin;
252static int websocket_read_wstream(BIO* bio, websocket_context* encodingContext)
255 WINPR_ASSERT(encodingContext);
257 wStream* s = encodingContext->responseStreamBuffer;
260 if (encodingContext->payloadLength == 0)
262 encodingContext->state = WebsocketStateOpcodeAndFin;
266 if (!Stream_EnsureRemainingCapacity(s, encodingContext->payloadLength))
269 "wStream::capacity [%" PRIuz
"] != encodingContext::paylaodLangth [%" PRIuz
"]",
270 Stream_GetRemainingCapacity(s), encodingContext->payloadLength);
274 const int status = websocket_read_data(bio, Stream_Pointer(s), Stream_GetRemainingCapacity(s),
279 if (!Stream_SafeSeek(s, (
size_t)status))
286static BOOL websocket_reply_close(BIO* bio, websocket_context* context,
wStream* s)
290 return websocket_context_write_wstream(context, bio, s, WebsocketCloseOpcode);
294static BOOL websocket_reply_pong(BIO* bio, websocket_context* context,
wStream* s)
299 if (Stream_GetPosition(s) != 0)
300 return websocket_context_write_wstream(context, bio, s, WebsocketPongOpcode);
302 return websocket_reply_close(bio, context, s);
306static int websocket_handle_payload(BIO* bio, BYTE* pBuffer,
size_t size,
307 websocket_context* encodingContext)
312 WINPR_ASSERT(pBuffer);
313 WINPR_ASSERT(encodingContext);
315 const BYTE effectiveOpcode = ((encodingContext->opcode & 0xf) == WebsocketContinuationOpcode
316 ? encodingContext->fragmentOriginalOpcode & 0xf
317 : encodingContext->opcode & 0xf);
319 switch (effectiveOpcode)
321 case WebsocketBinaryOpcode:
323 status = websocket_read_data(bio, pBuffer, size, encodingContext);
329 case WebsocketPingOpcode:
331 status = websocket_read_wstream(bio, encodingContext);
335 if (encodingContext->payloadLength == 0)
337 if (!websocket_reply_pong(bio, encodingContext,
338 encodingContext->responseStreamBuffer))
340 if (!Stream_Reset(encodingContext->responseStreamBuffer))
345 case WebsocketPongOpcode:
347 status = websocket_read_wstream(bio, encodingContext);
351 if (!Stream_Reset(encodingContext->responseStreamBuffer))
355 case WebsocketCloseOpcode:
357 status = websocket_read_wstream(bio, encodingContext);
361 if (encodingContext->payloadLength == 0)
363 if (!websocket_reply_close(bio, encodingContext,
364 encodingContext->responseStreamBuffer))
366 encodingContext->closeSent = TRUE;
367 if (!Stream_Reset(encodingContext->responseStreamBuffer))
373 WLog_WARN(TAG,
"Unimplemented websocket opcode %" PRIx8
". Dropping", effectiveOpcode);
375 status = websocket_read_wstream(bio, encodingContext);
378 if (!Stream_Reset(encodingContext->responseStreamBuffer))
387int websocket_context_read(websocket_context* encodingContext, BIO* bio, BYTE* pBuffer,
size_t size)
390 size_t effectiveDataLen = 0;
393 WINPR_ASSERT(pBuffer);
394 WINPR_ASSERT(encodingContext);
398 switch (encodingContext->state)
400 case WebsocketStateOpcodeAndFin:
402 BYTE buffer[1] = WINPR_C_ARRAY_INIT;
405 status = BIO_read(bio, (
char*)buffer,
sizeof(buffer));
407 return (effectiveDataLen > 0 ? WINPR_ASSERTING_INT_CAST(
int, effectiveDataLen)
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;
417 case WebsocketStateLengthAndMasking:
419 BYTE buffer[1] = WINPR_C_ARRAY_INIT;
422 status = BIO_read(bio, (
char*)buffer,
sizeof(buffer));
424 return (effectiveDataLen > 0 ? WINPR_ASSERTING_INT_CAST(
int, effectiveDataLen)
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;
433 encodingContext->payloadLength = len;
434 encodingContext->state = (encodingContext->masking ? WebSocketStateMaskingKey
435 : WebSocketStatePayload);
438 encodingContext->state = WebsocketStateShortLength;
440 encodingContext->state = WebsocketStateLongLength;
443 case WebsocketStateShortLength:
444 case WebsocketStateLongLength:
446 BYTE buffer[1] = WINPR_C_ARRAY_INIT;
447 const BYTE lenLength =
448 (encodingContext->state == WebsocketStateShortLength ? 2 : 8);
449 while (encodingContext->lengthAndMaskPosition < lenLength)
452 status = BIO_read(bio, (
char*)buffer,
sizeof(buffer));
454 return (effectiveDataLen > 0
455 ? WINPR_ASSERTING_INT_CAST(
int, effectiveDataLen)
457 if (status > UINT8_MAX)
459 encodingContext->payloadLength =
460 (encodingContext->payloadLength) << 8 | buffer[0];
461 encodingContext->lengthAndMaskPosition +=
462 WINPR_ASSERTING_INT_CAST(BYTE, status);
464 if (encodingContext->payloadLength > RESPONSE_SIZE_LIMIT)
466 WLog_ERR(TAG,
"received excessive payload size %" PRIuz
", aborting",
467 encodingContext->payloadLength);
470 encodingContext->state =
471 (encodingContext->masking ? WebSocketStateMaskingKey : WebSocketStatePayload);
474 case WebSocketStateMaskingKey:
477 TAG,
"Websocket Server sends data with masking key. This is against RFC 6455.");
480 case WebSocketStatePayload:
482 status = websocket_handle_payload(bio, pBuffer, size, encodingContext);
484 return (effectiveDataLen > 0 ? WINPR_ASSERTING_INT_CAST(
int, effectiveDataLen)
487 effectiveDataLen += WINPR_ASSERTING_INT_CAST(
size_t, status);
489 if (WINPR_ASSERTING_INT_CAST(
size_t, status) >= size)
490 return WINPR_ASSERTING_INT_CAST(
int, effectiveDataLen);
492 size -= WINPR_ASSERTING_INT_CAST(
size_t, status);
502websocket_context* websocket_context_new(
void)
504 websocket_context* context = calloc(1,
sizeof(websocket_context));
508 context->responseStreamBuffer = Stream_New(
nullptr, 1024);
509 if (!context->responseStreamBuffer)
512 if (!websocket_context_reset(context))
517 websocket_context_free(context);
521void websocket_context_free(websocket_context* context)
526 Stream_Free(context->responseStreamBuffer, TRUE);
530BOOL websocket_context_reset(websocket_context* context)
532 WINPR_ASSERT(context);
534 context->state = WebsocketStateOpcodeAndFin;
535 return Stream_Reset(context->responseStreamBuffer);