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 项单元测试全部通过
This commit is contained in:
parent
0ebf8b997d
commit
2775033198
6
Makefile
6
Makefile
@ -20,7 +20,7 @@ PREFIX ?= /usr/local
|
||||
BINDIR = $(PREFIX)/bin
|
||||
|
||||
# 源文件
|
||||
SRCS = main.c server.c http.c static.c log.c config.c multipart.c tls.c http2.c access_log.c
|
||||
SRCS = main.c server.c http.c static.c log.c config.c multipart.c tls.c http2.c access_log.c websocket.c
|
||||
OBJS = $(SRCS:.c=.o)
|
||||
TARGET = cocoon
|
||||
|
||||
@ -78,8 +78,8 @@ unit-test: $(UNIT_TEST_BINS)
|
||||
fi
|
||||
|
||||
# 单元测试编译规则
|
||||
$(UNIT_TEST_DIR)/test_server: $(UNIT_TEST_DIR)/test_server.c server.c http.c static.c log.c config.c multipart.c tls.c http2.c access_log.c $(UNITY_SRC)
|
||||
$(CC) $(CFLAGS) -I. -I$(UNIT_TEST_DIR)/../unity -o $@ $(UNIT_TEST_DIR)/test_server.c server.c http.c static.c log.c config.c multipart.c tls.c http2.c access_log.c $(UNITY_SRC) $(LDFLAGS)
|
||||
$(UNIT_TEST_DIR)/test_server: $(UNIT_TEST_DIR)/test_server.c server.c http.c static.c log.c config.c multipart.c tls.c http2.c access_log.c websocket.c $(UNITY_SRC)
|
||||
$(CC) $(CFLAGS) -I. -I$(UNIT_TEST_DIR)/../unity -o $@ $(UNIT_TEST_DIR)/test_server.c server.c http.c static.c log.c config.c multipart.c tls.c http2.c access_log.c websocket.c $(UNITY_SRC) $(LDFLAGS)
|
||||
|
||||
$(UNIT_TEST_DIR)/test_multipart: $(UNIT_TEST_DIR)/test_multipart.c multipart.c $(UNITY_SRC)
|
||||
$(CC) $(CFLAGS) -I. -I$(UNIT_TEST_DIR)/../unity -o $@ $(UNIT_TEST_DIR)/test_multipart.c multipart.c $(UNITY_SRC) -lm
|
||||
|
||||
60
server.c
60
server.c
@ -24,6 +24,7 @@
|
||||
#include "tls.h"
|
||||
#include "http2.h"
|
||||
#include "access_log.h"
|
||||
#include "websocket.h"
|
||||
#include <stdio.h>
|
||||
#include <stdlib.h>
|
||||
#include <string.h>
|
||||
@ -633,6 +634,38 @@ static bool is_h2c_upgrade_request(const http_request_t *req) {
|
||||
return has_upgrade && has_connection_upgrade;
|
||||
}
|
||||
|
||||
/**
|
||||
* is_websocket_upgrade_request - 检查是否是 WebSocket Upgrade 请求
|
||||
*
|
||||
* 检查请求头是否包含 Upgrade: websocket、Connection: Upgrade 和 Sec-WebSocket-Key。
|
||||
*
|
||||
* @param req HTTP 请求
|
||||
* @return true 是 WebSocket 升级请求
|
||||
*/
|
||||
static bool is_websocket_upgrade_request(const http_request_t *req) {
|
||||
if (strcmp(req->version, "HTTP/1.1") != 0) return false;
|
||||
if (req->method != HTTP_GET) return false;
|
||||
|
||||
bool has_upgrade = false;
|
||||
bool has_connection_upgrade = false;
|
||||
const char *key = NULL;
|
||||
|
||||
for (int i = 0; i < req->num_headers; i++) {
|
||||
if (strcasecmp(req->headers[i].name, "upgrade") == 0) {
|
||||
if (strcasestr(req->headers[i].value, "websocket") != NULL) {
|
||||
has_upgrade = true;
|
||||
}
|
||||
} else if (strcasecmp(req->headers[i].name, "connection") == 0) {
|
||||
if (strcasestr(req->headers[i].value, "upgrade") != NULL) {
|
||||
has_connection_upgrade = true;
|
||||
}
|
||||
} else if (strcasecmp(req->headers[i].name, "sec-websocket-key") == 0) {
|
||||
key = req->headers[i].value;
|
||||
}
|
||||
}
|
||||
return has_upgrade && has_connection_upgrade && key != NULL;
|
||||
}
|
||||
|
||||
/**
|
||||
* send_h2c_upgrade_response - 发送 101 Switching Protocols
|
||||
*
|
||||
@ -750,6 +783,31 @@ static void client_handler(void *arg) {
|
||||
http_request_free(&req);
|
||||
break;
|
||||
}
|
||||
|
||||
/* 检查 WebSocket Upgrade 请求 */
|
||||
if (parsed > 0 && is_websocket_upgrade_request(&req)) {
|
||||
const char *key = NULL;
|
||||
for (int i = 0; i < req.num_headers; i++) {
|
||||
if (strcasecmp(req.headers[i].name, "sec-websocket-key") == 0) {
|
||||
key = req.headers[i].value;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (key) {
|
||||
conn_cancel_timer(conn);
|
||||
if (ws_handshake(conn->fd, key) == 0) {
|
||||
http_request_free(&req);
|
||||
/* 消费已解析的请求数据 */
|
||||
if ((size_t)parsed < conn->buf_len) {
|
||||
memmove(conn->buf, conn->buf + parsed, conn->buf_len - (size_t)parsed);
|
||||
}
|
||||
conn->buf_len -= (size_t)parsed;
|
||||
ws_handle_connection(conn->fd, conn->timeout_ms);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
http_request_free(&req);
|
||||
|
||||
/* 尝试处理请求 */
|
||||
@ -775,6 +833,7 @@ static void client_handler(void *arg) {
|
||||
* @param arg 服务器上下文指针
|
||||
*/
|
||||
static void accept_loop(void *arg) {
|
||||
log_debug("=== accept_loop 启动 ===");
|
||||
server_context_t *ctx = (server_context_t *)arg;
|
||||
if (!ctx) return;
|
||||
|
||||
@ -805,6 +864,7 @@ static void accept_loop(void *arg) {
|
||||
|
||||
if (ctx->config.threaded) {
|
||||
/* 多线程模式:非阻塞 accept + poll 阻塞等待 */
|
||||
log_debug("accept_loop: 尝试 accept()...");
|
||||
client_fd = accept(ctx->listen_fd,
|
||||
(struct sockaddr *)&client_addr, &addr_len);
|
||||
if (client_fd < 0 && (errno == EAGAIN || errno == EWOULDBLOCK)) {
|
||||
|
||||
@ -720,6 +720,21 @@ else
|
||||
fail
|
||||
fi
|
||||
|
||||
# WebSocket 测试
|
||||
echo ""
|
||||
echo "=== WebSocket 测试 ==="
|
||||
|
||||
if python3 "$ROOT/../websocket_test.py" > "$TMPDIR/ws_test.log" 2>&1; then
|
||||
echo " ✓ WebSocket 握手 + echo — 通过"
|
||||
pass
|
||||
pass
|
||||
else
|
||||
echo " ✗ WebSocket 测试失败"
|
||||
cat "$TMPDIR/ws_test.log"
|
||||
fail
|
||||
fail
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== 结果汇总 ==="
|
||||
echo "通过: $PASS"
|
||||
|
||||
245
tests/websocket_test.py
Normal file
245
tests/websocket_test.py
Normal file
@ -0,0 +1,245 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
WebSocket 集成测试 — 使用 Python 标准库实现
|
||||
|
||||
测试内容:
|
||||
1. HTTP 升级握手(101 Switching Protocols)
|
||||
2. 发送文本帧,接收 echo 回显
|
||||
3. 发送关闭帧,接收关闭响应
|
||||
"""
|
||||
|
||||
import socket
|
||||
import base64
|
||||
import hashlib
|
||||
import struct
|
||||
import sys
|
||||
|
||||
HOST = "localhost"
|
||||
PORT = 9999
|
||||
GUID = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
|
||||
|
||||
|
||||
def build_handshake():
|
||||
"""构建 WebSocket 握手请求"""
|
||||
key = base64.b64encode(b"\x00" * 16).decode()
|
||||
req = (
|
||||
f"GET /ws HTTP/1.1\r\n"
|
||||
f"Host: {HOST}:{PORT}\r\n"
|
||||
f"Upgrade: websocket\r\n"
|
||||
f"Connection: Upgrade\r\n"
|
||||
f"Sec-WebSocket-Key: {key}\r\n"
|
||||
f"Sec-WebSocket-Version: 13\r\n"
|
||||
f"\r\n"
|
||||
)
|
||||
return req, key
|
||||
|
||||
|
||||
def parse_response(data):
|
||||
"""解析 HTTP 101 响应"""
|
||||
lines = data.split(b"\r\n")
|
||||
status = lines[0].decode()
|
||||
headers = {}
|
||||
for line in lines[1:]:
|
||||
if line == b"":
|
||||
break
|
||||
if b":" in line:
|
||||
k, v = line.split(b":", 1)
|
||||
headers[k.decode().strip().lower()] = v.decode().strip()
|
||||
return status, headers
|
||||
|
||||
|
||||
def compute_accept(key):
|
||||
"""计算 Sec-WebSocket-Accept"""
|
||||
concat = key.encode() + GUID
|
||||
digest = hashlib.sha1(concat).digest()
|
||||
return base64.b64encode(digest).decode()
|
||||
|
||||
|
||||
def build_frame(opcode, payload, masked=True):
|
||||
"""构建 WebSocket 帧(客户端发送,带掩码)"""
|
||||
length = len(payload)
|
||||
header = bytearray()
|
||||
header.append(0x80 | opcode) # FIN=1, opcode
|
||||
|
||||
if length <= 125:
|
||||
header.append((0x80 if masked else 0x00) | length)
|
||||
elif length <= 65535:
|
||||
header.append((0x80 if masked else 0x00) | 126)
|
||||
header.extend(struct.pack(">H", length))
|
||||
else:
|
||||
header.append((0x80 if masked else 0x00) | 127)
|
||||
header.extend(struct.pack(">Q", length))
|
||||
|
||||
if masked:
|
||||
mask = b"\x12\x34\x56\x78"
|
||||
masked_payload = bytearray()
|
||||
for i, b in enumerate(payload):
|
||||
masked_payload.append(b ^ mask[i % 4])
|
||||
return bytes(header) + mask + bytes(masked_payload)
|
||||
|
||||
return bytes(header) + payload
|
||||
|
||||
|
||||
def parse_frame(data):
|
||||
"""解析服务器发来的 WebSocket 帧(无掩码)"""
|
||||
if len(data) < 2:
|
||||
return None, 0
|
||||
|
||||
b0 = data[0]
|
||||
b1 = data[1]
|
||||
fin = (b0 >> 7) & 1
|
||||
opcode = b0 & 0x0F
|
||||
payload_len = b1 & 0x7F
|
||||
offset = 2
|
||||
|
||||
if payload_len == 126:
|
||||
if len(data) < 4:
|
||||
return None, 0
|
||||
payload_len = struct.unpack(">H", data[2:4])[0]
|
||||
offset = 4
|
||||
elif payload_len == 127:
|
||||
if len(data) < 10:
|
||||
return None, 0
|
||||
payload_len = struct.unpack(">Q", data[2:10])[0]
|
||||
offset = 10
|
||||
|
||||
# 服务器不掩码
|
||||
if len(data) < offset + payload_len:
|
||||
return None, 0
|
||||
|
||||
payload = data[offset:offset + payload_len]
|
||||
return {"opcode": opcode, "fin": fin, "payload": payload}, offset + payload_len
|
||||
|
||||
|
||||
def recv_all(sock, n):
|
||||
"""接收至少 n 字节数据"""
|
||||
data = b""
|
||||
while len(data) < n:
|
||||
chunk = sock.recv(n - len(data))
|
||||
if not chunk:
|
||||
break
|
||||
data += chunk
|
||||
return data
|
||||
|
||||
|
||||
def recv_frame(sock):
|
||||
"""接收一个完整帧(自动读取足够的数据)"""
|
||||
data = b""
|
||||
while True:
|
||||
frame, consumed = parse_frame(data)
|
||||
if frame is not None:
|
||||
return frame, consumed
|
||||
chunk = sock.recv(1024)
|
||||
if not chunk:
|
||||
return None, 0
|
||||
data += chunk
|
||||
|
||||
|
||||
def test_handshake():
|
||||
"""测试 WebSocket 握手"""
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.connect((HOST, PORT))
|
||||
|
||||
req, key = build_handshake()
|
||||
sock.send(req.encode())
|
||||
|
||||
data = b""
|
||||
while b"\r\n\r\n" not in data:
|
||||
chunk = sock.recv(1024)
|
||||
if not chunk:
|
||||
break
|
||||
data += chunk
|
||||
|
||||
status, headers = parse_response(data)
|
||||
if "101" not in status:
|
||||
print(f"FAIL: 期望 101,实际: {status}")
|
||||
sock.close()
|
||||
return False
|
||||
|
||||
accept = compute_accept(key)
|
||||
if headers.get("sec-websocket-accept") != accept:
|
||||
print(f"FAIL: Sec-WebSocket-Accept 不匹配")
|
||||
print(f" 期望: {accept}")
|
||||
print(f" 实际: {headers.get('sec-websocket-accept')}")
|
||||
sock.close()
|
||||
return False
|
||||
|
||||
print("PASS: WebSocket 握手 101 + Sec-WebSocket-Accept 正确")
|
||||
sock.close()
|
||||
return True
|
||||
|
||||
|
||||
def test_echo():
|
||||
"""测试文本帧 echo"""
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.settimeout(5.0)
|
||||
sock.connect((HOST, PORT))
|
||||
|
||||
req, key = build_handshake()
|
||||
sock.send(req.encode())
|
||||
|
||||
data = b""
|
||||
while b"\r\n\r\n" not in data:
|
||||
chunk = sock.recv(1024)
|
||||
if not chunk:
|
||||
break
|
||||
data += chunk
|
||||
|
||||
# 发送文本帧
|
||||
msg = b"Hello, Cocoon!"
|
||||
frame = build_frame(0x01, msg)
|
||||
sock.send(frame)
|
||||
|
||||
# 接收 echo 回显
|
||||
frame, _ = recv_frame(sock)
|
||||
if frame is None:
|
||||
print("FAIL: 未收到响应帧")
|
||||
sock.close()
|
||||
return False
|
||||
|
||||
if frame["opcode"] != 0x01:
|
||||
print(f"FAIL: 期望文本帧(1),实际操作码: {frame['opcode']}")
|
||||
sock.close()
|
||||
return False
|
||||
|
||||
if frame["payload"] != msg:
|
||||
print(f"FAIL: 回显内容不匹配")
|
||||
print(f" 期望: {msg}")
|
||||
print(f" 实际: {frame['payload']}")
|
||||
sock.close()
|
||||
return False
|
||||
|
||||
print("PASS: 文本帧 echo 正确")
|
||||
|
||||
# 发送关闭帧
|
||||
close_frame = build_frame(0x08, b"\x03\xe8") # 1000
|
||||
sock.send(close_frame)
|
||||
|
||||
# 接收关闭响应
|
||||
frame, _ = recv_frame(sock)
|
||||
if frame is None or frame["opcode"] != 0x08:
|
||||
print("FAIL: 未收到关闭帧响应")
|
||||
sock.close()
|
||||
return False
|
||||
|
||||
print("PASS: 关闭帧响应正确")
|
||||
sock.close()
|
||||
return True
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
passed = 0
|
||||
failed = 0
|
||||
|
||||
if test_handshake():
|
||||
passed += 1
|
||||
else:
|
||||
failed += 1
|
||||
|
||||
if test_echo():
|
||||
passed += 1
|
||||
else:
|
||||
failed += 1
|
||||
|
||||
print(f"\n通过: {passed}, 失败: {failed}")
|
||||
sys.exit(0 if failed == 0 else 1)
|
||||
338
websocket.c
Normal file
338
websocket.c
Normal file
@ -0,0 +1,338 @@
|
||||
/**
|
||||
* @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);
|
||||
}
|
||||
130
websocket.h
Normal file
130
websocket.h
Normal file
@ -0,0 +1,130 @@
|
||||
#ifndef WEBSOCKET_H
|
||||
#define WEBSOCKET_H
|
||||
|
||||
#include <stdint.h>
|
||||
#include <stddef.h>
|
||||
#include <stdbool.h>
|
||||
|
||||
/**
|
||||
* @file websocket.h
|
||||
* @brief WebSocket 协议实现
|
||||
*
|
||||
* 支持 RFC 6455 WebSocket 握手、帧解析与编码。
|
||||
* 服务器端实现:不发送掩码(mask=0),接收客户端掩码帧。
|
||||
*/
|
||||
|
||||
/**
|
||||
* @brief WebSocket 操作码
|
||||
*/
|
||||
typedef enum {
|
||||
WS_OP_CONT = 0x0, /**< 继续帧 */
|
||||
WS_OP_TEXT = 0x1, /**< 文本帧 */
|
||||
WS_OP_BINARY = 0x2, /**< 二进制帧 */
|
||||
WS_OP_CLOSE = 0x8, /**< 关闭帧 */
|
||||
WS_OP_PING = 0x9, /**< Ping 帧 */
|
||||
WS_OP_PONG = 0xA /**< Pong 帧 */
|
||||
} ws_opcode_t;
|
||||
|
||||
/**
|
||||
* @brief WebSocket 帧结构
|
||||
*/
|
||||
typedef struct {
|
||||
uint8_t opcode; /**< 操作码 */
|
||||
bool fin; /**< 是否为最后一帧 */
|
||||
bool masked; /**< 是否掩码 */
|
||||
uint64_t payload_len;/**< 负载长度 */
|
||||
uint8_t mask_key[4]; /**< 掩码密钥(仅客户端发送时有效) */
|
||||
uint8_t *payload; /**< 负载数据(已解掩码) */
|
||||
} ws_frame_t;
|
||||
|
||||
/**
|
||||
* @brief 执行 WebSocket 握手响应
|
||||
*
|
||||
* 根据 RFC 6455,对 Sec-WebSocket-Key 计算 SHA1 + Base64 响应。
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @param key 客户端发来的 Sec-WebSocket-Key
|
||||
* @return 0 成功,-1 失败
|
||||
*/
|
||||
int ws_handshake(int fd, const char *key);
|
||||
|
||||
/**
|
||||
* @brief 解析单个 WebSocket 帧
|
||||
*
|
||||
* 从数据流中解析一个完整帧。如果数据不完整,返回 -1 且不修改 frame。
|
||||
*
|
||||
* @param data 输入数据
|
||||
* @param len 数据长度
|
||||
* @param frame 输出帧结构(调用者需初始化)
|
||||
* @param consumed 输出:消耗的字节数
|
||||
* @return 0 成功,-1 数据不完整,-2 格式错误
|
||||
*/
|
||||
int ws_parse_frame(const uint8_t *data, size_t len, ws_frame_t *frame, size_t *consumed);
|
||||
|
||||
/**
|
||||
* @brief 释放帧占用的负载内存
|
||||
*
|
||||
* @param frame 帧指针
|
||||
*/
|
||||
void ws_frame_free(ws_frame_t *frame);
|
||||
|
||||
/**
|
||||
* @brief 发送 WebSocket 帧
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @param opcode 操作码
|
||||
* @param payload 负载数据
|
||||
* @param len 负载长度
|
||||
* @return 0 成功,-1 失败
|
||||
*/
|
||||
int ws_send_frame(int fd, uint8_t opcode, const uint8_t *payload, size_t len);
|
||||
|
||||
/**
|
||||
* @brief 发送文本帧
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @param text 文本内容(UTF-8)
|
||||
* @return 0 成功,-1 失败
|
||||
*/
|
||||
int ws_send_text(int fd, const char *text);
|
||||
|
||||
/**
|
||||
* @brief 发送关闭帧
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @param code 关闭码(如 1000)
|
||||
* @param reason 关闭原因(可为 NULL)
|
||||
* @return 0 成功,-1 失败
|
||||
*/
|
||||
int ws_send_close(int fd, uint16_t code, const char *reason);
|
||||
|
||||
/**
|
||||
* @brief 发送 Ping 帧
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @return 0 成功,-1 失败
|
||||
*/
|
||||
int ws_send_ping(int fd);
|
||||
|
||||
/**
|
||||
* @brief 发送 Pong 帧
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @param payload Ping 的负载(可为 NULL)
|
||||
* @param len 负载长度
|
||||
* @return 0 成功,-1 失败
|
||||
*/
|
||||
int ws_send_pong(int fd, const uint8_t *payload, size_t len);
|
||||
|
||||
/**
|
||||
* @brief 处理 WebSocket 连接(主循环)
|
||||
*
|
||||
* 进入 WebSocket 帧循环,处理文本/二进制/ping/pong/close。
|
||||
* 当前实现为简单 echo 服务器。
|
||||
*
|
||||
* @param fd 客户端 socket
|
||||
* @param timeout_ms 超时毫秒(0 表示默认)
|
||||
*/
|
||||
void ws_handle_connection(int fd, uint32_t timeout_ms);
|
||||
|
||||
#endif /* WEBSOCKET_H */
|
||||
Loading…
x
Reference in New Issue
Block a user