cocoon/websocket.c
xfy911 2775033198 feat(websocket): 实现 RFC 6455 WebSocket 支持
- 新增 websocket.h / websocket.c:帧解析、编码、握手、连接管理
- server.c: 检测 Upgrade: websocket 请求,执行握手后进入 ws_handle_connection
- 支持文本/二进制帧 echo、ping/pong、close 帧
- 使用 OpenSSL SHA1 计算 Sec-WebSocket-Accept
- 集成测试:新增 2 项 WebSocket 测试(Python 标准库实现)
- 61 项集成测试全部通过,127 项单元测试全部通过
2026-06-05 15:18:01 +08:00

339 lines
10 KiB
C
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/**
* @file websocket.c
* @brief WebSocket 协议实现RFC 6455
*
* 支持握手、帧解析、帧编码、连接管理。
* 当前为简单 echo 服务器,可扩展为消息路由或广播。
*/
#include "websocket.h"
#include "log.h"
#include <stdio.h>
#include <string.h>
#include <stdlib.h>
#include <unistd.h>
#include <errno.h>
#include <openssl/sha.h>
#include <sys/socket.h>
/* WebSocket 魔数字符串GUID */
static const char WS_GUID[] = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
/* 基础 64 编码表 */
static const char BASE64_TABLE[] =
"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
/**
* @brief 简单 Base64 编码(内部使用)
*
* @param in 输入数据
* @param in_len 输入长度
* @param out 输出缓冲区(至少 4/3 * in_len + 4 字节)
* @return 输出字符串长度
*/
static size_t base64_encode(const uint8_t *in, size_t in_len, char *out) {
size_t i, j;
for (i = 0, j = 0; i + 2 < in_len; i += 3, j += 4) {
uint32_t v = ((uint32_t)in[i] << 16) | ((uint32_t)in[i + 1] << 8) | in[i + 2];
out[j] = BASE64_TABLE[(v >> 18) & 0x3F];
out[j + 1] = BASE64_TABLE[(v >> 12) & 0x3F];
out[j + 2] = BASE64_TABLE[(v >> 6) & 0x3F];
out[j + 3] = BASE64_TABLE[v & 0x3F];
}
if (i < in_len) {
uint32_t v = (uint32_t)in[i] << 16;
if (i + 1 < in_len) v |= (uint32_t)in[i + 1] << 8;
out[j] = BASE64_TABLE[(v >> 18) & 0x3F];
out[j + 1] = BASE64_TABLE[(v >> 12) & 0x3F];
out[j + 2] = (i + 1 < in_len) ? BASE64_TABLE[(v >> 6) & 0x3F] : '=';
out[j + 3] = '=';
j += 4;
}
out[j] = '\0';
return j;
}
/**
* @brief 生成 WebSocket 握手响应的 Sec-WebSocket-Accept
*
* 对 key + GUID 计算 SHA1然后 Base64 编码。
*
* @param key 客户端 Sec-WebSocket-Key
* @param accept 输出缓冲区(至少 32 字节)
* @return 0 成功
*/
static int ws_compute_accept(const char *key, char *accept) {
char concat[128];
size_t key_len = strlen(key);
if (key_len + sizeof(WS_GUID) >= sizeof(concat)) return -1;
memcpy(concat, key, key_len);
memcpy(concat + key_len, WS_GUID, sizeof(WS_GUID) - 1);
concat[key_len + sizeof(WS_GUID) - 1] = '\0';
unsigned char digest[SHA_DIGEST_LENGTH];
SHA1((unsigned char *)concat, key_len + sizeof(WS_GUID) - 1, digest);
base64_encode(digest, SHA_DIGEST_LENGTH, accept);
return 0;
}
/**
* @brief 发送完整的 HTTP 响应
*
* @param fd 客户端 socket
* @param buf 数据
* @param len 长度
* @return 0 成功,-1 失败
*/
static int ws_send_all(int fd, const char *buf, size_t len) {
size_t sent = 0;
while (sent < len) {
ssize_t n = send(fd, buf + sent, len - sent, MSG_NOSIGNAL);
if (n > 0) {
sent += (size_t)n;
} else if (n < 0 && (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR)) {
continue;
} else {
return -1;
}
}
return 0;
}
int ws_handshake(int fd, const char *key) {
char accept[32];
if (ws_compute_accept(key, accept) != 0) return -1;
char response[512];
int n = snprintf(response, sizeof(response),
"HTTP/1.1 101 Switching Protocols\r\n"
"Upgrade: websocket\r\n"
"Connection: Upgrade\r\n"
"Sec-WebSocket-Accept: %s\r\n"
"Server: Cocoon/1.0\r\n"
"\r\n",
accept);
if (n < 0 || (size_t)n >= sizeof(response)) return -1;
return ws_send_all(fd, response, (size_t)n);
}
int ws_parse_frame(const uint8_t *data, size_t len, ws_frame_t *frame, size_t *consumed) {
if (len < 2) return -1; /* 至少需要 2 字节 */
uint8_t b0 = data[0];
uint8_t b1 = data[1];
frame->fin = (b0 >> 7) & 1;
frame->opcode = b0 & 0x0F;
frame->masked = (b1 >> 7) & 1;
uint64_t payload_len = b1 & 0x7F;
size_t header_len = 2;
if (payload_len == 126) {
if (len < 4) return -1;
payload_len = ((uint64_t)data[2] << 8) | data[3];
header_len = 4;
} else if (payload_len == 127) {
if (len < 10) return -1;
payload_len = 0;
for (int i = 0; i < 8; i++) {
payload_len = (payload_len << 8) | data[2 + i];
}
header_len = 10;
}
if (frame->masked) {
if (len < header_len + 4) return -1;
memcpy(frame->mask_key, data + header_len, 4);
header_len += 4;
}
if (len < header_len + payload_len) return -1;
frame->payload_len = payload_len;
if (payload_len > 0) {
frame->payload = (uint8_t *)malloc(payload_len + 1);
if (!frame->payload) return -2;
memcpy(frame->payload, data + header_len, payload_len);
if (frame->masked) {
for (uint64_t i = 0; i < payload_len; i++) {
frame->payload[i] ^= frame->mask_key[i % 4];
}
}
frame->payload[payload_len] = '\0';
} else {
frame->payload = NULL;
}
*consumed = header_len + payload_len;
return 0;
}
void ws_frame_free(ws_frame_t *frame) {
if (frame && frame->payload) {
free(frame->payload);
frame->payload = NULL;
}
}
int ws_send_frame(int fd, uint8_t opcode, const uint8_t *payload, size_t len) {
uint8_t header[14];
size_t header_len = 0;
header[0] = 0x80 | (opcode & 0x0F); /* FIN=1, opcode */
if (len <= 125) {
header[1] = (uint8_t)len;
header_len = 2;
} else if (len <= 65535) {
header[1] = 126;
header[2] = (uint8_t)(len >> 8);
header[3] = (uint8_t)(len & 0xFF);
header_len = 4;
} else {
header[1] = 127;
for (int i = 0; i < 8; i++) {
header[2 + i] = (uint8_t)(len >> (56 - i * 8));
}
header_len = 10;
}
/* 服务器发送不掩码 */
if (ws_send_all(fd, (char *)header, header_len) != 0) return -1;
if (len > 0 && ws_send_all(fd, (char *)payload, len) != 0) return -1;
return 0;
}
int ws_send_text(int fd, const char *text) {
return ws_send_frame(fd, WS_OP_TEXT, (const uint8_t *)text, strlen(text));
}
int ws_send_close(int fd, uint16_t code, const char *reason) {
uint8_t payload[128];
size_t len = 0;
if (code != 0) {
payload[0] = (uint8_t)(code >> 8);
payload[1] = (uint8_t)(code & 0xFF);
len = 2;
if (reason) {
size_t rlen = strlen(reason);
if (rlen > sizeof(payload) - 3) rlen = sizeof(payload) - 3;
memcpy(payload + 2, reason, rlen);
len += rlen;
}
}
return ws_send_frame(fd, WS_OP_CLOSE, payload, len);
}
int ws_send_ping(int fd) {
return ws_send_frame(fd, WS_OP_PING, NULL, 0);
}
int ws_send_pong(int fd, const uint8_t *payload, size_t len) {
return ws_send_frame(fd, WS_OP_PONG, payload, len);
}
/**
* @brief 从连接读取数据(简单阻塞读取)
*/
static ssize_t ws_read_data(int fd, uint8_t *buf, size_t max_len) {
ssize_t n = recv(fd, buf, max_len, 0);
return n;
}
void ws_handle_connection(int fd, uint32_t timeout_ms) {
(void)timeout_ms; /* TODO: 超时处理 */
uint8_t buf[8192];
size_t buf_len = 0;
bool closed = false;
log_info("WebSocket 连接建立 fd=%d", fd);
while (!closed) {
ssize_t n = ws_read_data(fd, buf + buf_len, sizeof(buf) - buf_len);
if (n <= 0) {
if (n < 0 && (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR)) {
continue;
}
log_debug("WebSocket fd=%d 读取结束或错误", fd);
break;
}
buf_len += (size_t)n;
/* 解析帧 */
size_t parsed = 0;
while (parsed < buf_len) {
ws_frame_t frame = {0};
size_t consumed = 0;
int ret = ws_parse_frame(buf + parsed, buf_len - parsed, &frame, &consumed);
if (ret == -1) {
break; /* 数据不完整,等待更多 */
}
if (ret == -2) {
log_warn("WebSocket fd=%d 收到畸形帧", fd);
ws_send_close(fd, 1002, "Protocol error");
closed = true;
break;
}
parsed += consumed;
switch (frame.opcode) {
case WS_OP_TEXT:
log_debug("WebSocket fd=%d 收到文本: %s", fd,
frame.payload ? (char *)frame.payload : "(empty)");
if (frame.fin) {
ws_send_text(fd, frame.payload ? (char *)frame.payload : "");
}
break;
case WS_OP_BINARY:
log_debug("WebSocket fd=%d 收到二进制 %llu 字节", fd,
(unsigned long long)frame.payload_len);
/* 简单 echo 回二进制 */
ws_send_frame(fd, WS_OP_BINARY, frame.payload, frame.payload_len);
break;
case WS_OP_CLOSE:
log_info("WebSocket fd=%d 收到关闭帧", fd);
ws_send_close(fd, 1000, NULL);
closed = true;
break;
case WS_OP_PING:
log_debug("WebSocket fd=%d 收到 Ping", fd);
ws_send_pong(fd, frame.payload, frame.payload_len);
break;
case WS_OP_PONG:
log_debug("WebSocket fd=%d 收到 Pong", fd);
break;
case WS_OP_CONT:
log_debug("WebSocket fd=%d 收到继续帧", fd);
break;
default:
log_warn("WebSocket fd=%d 未知操作码 %d", fd, frame.opcode);
break;
}
ws_frame_free(&frame);
}
/* 移动剩余数据到缓冲区开头 */
if (parsed > 0 && parsed < buf_len) {
memmove(buf, buf + parsed, buf_len - parsed);
buf_len -= parsed;
} else if (parsed == buf_len) {
buf_len = 0;
}
}
log_info("WebSocket 连接关闭 fd=%d", fd);
}