- 新增 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 项单元测试全部通过
339 lines
10 KiB
C
339 lines
10 KiB
C
/**
|
||
* @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);
|
||
}
|