From 2775033198668658fd4e1dd1d6e92d8f65741bde Mon Sep 17 00:00:00 2001 From: xfy911 Date: Fri, 5 Jun 2026 15:18:01 +0800 Subject: [PATCH] =?UTF-8?q?feat(websocket):=20=E5=AE=9E=E7=8E=B0=20RFC=206?= =?UTF-8?q?455=20WebSocket=20=E6=94=AF=E6=8C=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 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 项单元测试全部通过 --- Makefile | 6 +- server.c | 60 +++++++ tests/integration_test.sh | 15 ++ tests/websocket_test.py | 245 +++++++++++++++++++++++++++ websocket.c | 338 ++++++++++++++++++++++++++++++++++++++ websocket.h | 130 +++++++++++++++ 6 files changed, 791 insertions(+), 3 deletions(-) create mode 100644 tests/websocket_test.py create mode 100644 websocket.c create mode 100644 websocket.h diff --git a/Makefile b/Makefile index 9d5edbb..74bb9cb 100644 --- a/Makefile +++ b/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 diff --git a/server.c b/server.c index 6408838..fecfa03 100644 --- a/server.c +++ b/server.c @@ -24,6 +24,7 @@ #include "tls.h" #include "http2.h" #include "access_log.h" +#include "websocket.h" #include #include #include @@ -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)) { diff --git a/tests/integration_test.sh b/tests/integration_test.sh index a20740f..3663cf9 100755 --- a/tests/integration_test.sh +++ b/tests/integration_test.sh @@ -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" diff --git a/tests/websocket_test.py b/tests/websocket_test.py new file mode 100644 index 0000000..623c702 --- /dev/null +++ b/tests/websocket_test.py @@ -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) diff --git a/websocket.c b/websocket.c new file mode 100644 index 0000000..a9ad845 --- /dev/null +++ b/websocket.c @@ -0,0 +1,338 @@ +/** + * @file websocket.c + * @brief WebSocket 协议实现(RFC 6455) + * + * 支持握手、帧解析、帧编码、连接管理。 + * 当前为简单 echo 服务器,可扩展为消息路由或广播。 + */ + +#include "websocket.h" +#include "log.h" +#include +#include +#include +#include +#include +#include +#include + +/* 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); +} diff --git a/websocket.h b/websocket.h new file mode 100644 index 0000000..4601f86 --- /dev/null +++ b/websocket.h @@ -0,0 +1,130 @@ +#ifndef WEBSOCKET_H +#define WEBSOCKET_H + +#include +#include +#include + +/** + * @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 */