diff --git a/.gitmodules b/.gitmodules deleted file mode 100644 index 417097d..0000000 --- a/.gitmodules +++ /dev/null @@ -1,3 +0,0 @@ -[submodule "coco"] - path = coco - url = https://github.com/xfy911/coco.git diff --git a/coco b/coco deleted file mode 160000 index 8b42e98..0000000 --- a/coco +++ /dev/null @@ -1 +0,0 @@ -Subproject commit 8b42e984efd55f5c18fdf2d249d1aba0ca4e70eb diff --git a/coco/include/coco.h b/coco/include/coco.h new file mode 100644 index 0000000..706365e --- /dev/null +++ b/coco/include/coco.h @@ -0,0 +1,53 @@ +#ifndef COCO_H +#define COCO_H + +#include +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +#define COCO_OK 0 +#define COCO_ERROR -1 +#define COCO_ERROR_NOMEM -2 +#define COCO_ERROR_CANCELLED -3 +#define COCO_ERROR_WOULD_BLOCK -4 +#define COCO_ERROR_TIMEOUT -5 + +typedef struct coco_sched coco_sched_t; +typedef struct coco_coro coco_coro_t; +typedef struct coco_timer coco_timer_t; + +typedef void (*coco_func_t)(void *arg); +typedef void (*coco_timer_handler_t)(void *arg); + +typedef struct { + int flags; + int stack_size; + int prio; + const char *name; +} coco_go_opts_t; + +int coco_global_sched_start(int num_workers); +void coco_global_sched_wait(void); +void coco_global_sched_stop(void); + +coco_coro_t *coco_go(coco_func_t func, void *arg); +coco_coro_t *coco_go_with_opts(coco_func_t func, void *arg, const coco_go_opts_t *opts); +coco_coro_t *coco_self(void); +int coco_cancel(coco_coro_t *coro); +coco_sched_t *coco_sched_get_current(void); + +ssize_t coco_read(int fd, void *buf, size_t len); +ssize_t coco_write(int fd, const void *buf, size_t len); + +coco_timer_t *coco_timer(uint32_t ms, coco_timer_handler_t handler, void *arg); +void coco_timer_cancel(coco_timer_t *timer); + +#ifdef __cplusplus +} +#endif + +#endif /* COCO_H */ diff --git a/grpc.c b/grpc.c new file mode 100644 index 0000000..5e1d374 --- /dev/null +++ b/grpc.c @@ -0,0 +1,502 @@ +/** + * @file grpc.c - gRPC over HTTP/2 支持实现 + * + * 基于现有 HTTP/2 传输层实现 gRPC 协议支持。 + * 包含消息帧编解码、请求解析、trailers 格式化、状态码转换。 + * + * @author Cocoon Team + */ + +#include "grpc.h" +#include +#include +#include + +/* ===== 内部辅助函数 ===== */ + +/** + * grpc_parse_path - 从 HTTP 路径解析 gRPC service/method + * + * 路径格式: "/package.Service/Method" 或 "/Service/Method" + * 解析结果: + * - service_name = "package.Service" 或 "Service" + * - method_name = "Method" + * + * @param path HTTP :path 伪头值 + * @param svc_buf 服务名输出缓冲区 + * @param svc_size 服务名缓冲区大小 + * @param method_buf 方法名输出缓冲区 + * @param method_size 方法名缓冲区大小 + * @return 0 成功,-1 失败 + */ +static int grpc_parse_path(const char *path, + char *svc_buf, size_t svc_size, + char *method_buf, size_t method_size) { + if (!path || path[0] != '/') { + return -1; + } + + /* 跳过开头的 '/' */ + const char *p = path + 1; + + /* 查找第二个 '/'(分隔 service 和 method) */ + const char *slash = strrchr(p, '/'); + if (!slash || slash == p) { + return -1; /* 格式错误:没有 method 部分 */ + } + + /* service_name = path[1..slash-1] */ + size_t svc_len = (size_t)(slash - p); + if (svc_len >= svc_size) { + svc_len = svc_size - 1; + } + memcpy(svc_buf, p, svc_len); + svc_buf[svc_len] = '\0'; + + /* method_name = slash[1..end] */ + size_t method_len = strlen(slash + 1); + if (method_len >= method_size) { + method_len = method_size - 1; + } + memcpy(method_buf, slash + 1, method_len); + method_buf[method_len] = '\0'; + + return 0; +} + +/** + * extract_content_type - 从 HTTP 请求中提取 Content-Type 头值 + * + * @param req HTTP 请求 + * @return Content-Type 头值指针,未找到返回 NULL + */ +static const char *extract_content_type(const http_request_t *req) { + if (!req) return NULL; + + /* 优先使用已解析的 content_type 字段 */ + if (req->content_type[0] != '\0') { + return req->content_type; + } + + /* 否则遍历 headers 数组查找 */ + for (int i = 0; i < req->num_headers; i++) { + if (strcasecmp(req->headers[i].name, "content-type") == 0) { + return req->headers[i].value; + } + } + return NULL; +} + +/* ===== gRPC 请求检测 ===== */ + +/** + * grpc_detect - 检测请求是否为 gRPC 请求 + * + * 检查 Content-Type 是否以 "application/grpc" 开头。 + * 覆盖: + * - application/grpc + * - application/grpc+proto + * - application/grpc+json + * - application/grpc-web (gRPC-Web) + * - application/grpc-web+proto + */ +bool grpc_detect(const http_request_t *req) { + if (!req) return false; + + const char *ct = extract_content_type(req); + if (!ct) return false; + + /* 检查 "application/grpc" 前缀(大小写不敏感) */ + if (strncasecmp(ct, "application/grpc", 16) == 0) { + char c = ct[16]; + /* 标准 gRPC: application/grpc, application/grpc+proto, ... */ + if (c == '\0' || c == '+' || c == ';' || c == ' ') { + return true; + } + /* gRPC-Web: application/grpc-web, application/grpc-web+proto */ + if (c == '-' && strncasecmp(ct + 17, "web", 3) == 0) { + char d = ct[20]; + if (d == '\0' || d == '+' || d == ';' || d == ' ') { + return true; + } + } + } + + return false; +} + +/** + * grpc_is_grpc_web - 检测请求是否为 gRPC-Web + * + * gRPC-Web 使用 application/grpc-web 前缀。 + */ +bool grpc_is_grpc_web(const http_request_t *req) { + if (!req) return false; + + const char *ct = extract_content_type(req); + if (!ct) return false; + + if (strncasecmp(ct, "application/grpc-web", 20) == 0) { + char c = ct[20]; + if (c == '\0' || c == '+' || c == ';' || c == ' ') { + return true; + } + } + + return false; +} + +/* ===== gRPC 请求解析 ===== */ + +/** + * grpc_parse_request - 从 HTTP/2 请求解析 gRPC 上下文 + * + * 解析路径获取 service/method,并解析消息体中的 gRPC 消息帧。 + * 默认标记为非流式(Unary),流模式由调用方根据 protobuf 定义设置。 + */ +int grpc_parse_request(const http_request_t *req, grpc_request_t *grpc_req) { + if (!req || !grpc_req) return -1; + + /* 清空输出结构 */ + memset(grpc_req, 0, sizeof(grpc_request_t)); + + /* 检测 gRPC-Web */ + grpc_req->is_grpc_web = grpc_is_grpc_web(req); + + /* 解析 service/method */ + if (grpc_parse_path(req->path, + grpc_req->service_name, sizeof(grpc_req->service_name), + grpc_req->method_name, sizeof(grpc_req->method_name)) != 0) { + /* 路径格式错误,设置默认值 */ + grpc_req->status = GRPC_INTERNAL; + snprintf(grpc_req->status_message, sizeof(grpc_req->status_message), + "Invalid gRPC path format: %.200s", req->path); + return -1; + } + + /* 默认流模式:Unary(非流式) */ + grpc_req->client_streaming = false; + grpc_req->server_streaming = false; + grpc_req->is_streaming = false; + + /* 解析消息体中的 gRPC 消息帧 */ + if (req->body && req->body_len > 0) { + int decoded = grpc_decode_message((const uint8_t *)req->body, + req->body_len, &grpc_req->message); + if (decoded < 0) { + grpc_req->status = GRPC_INTERNAL; + snprintf(grpc_req->status_message, sizeof(grpc_req->status_message), + "Failed to decode gRPC message frame"); + return -1; + } + } + + /* 提取元数据(非 Content-Type 的请求头) */ + grpc_req->metadata_count = 0; + for (int i = 0; i < req->num_headers && grpc_req->metadata_count < 16; i++) { + const char *name = req->headers[i].name; + /* 跳过 HTTP/2 伪头和标准传输头 */ + if (name[0] == ':') continue; + if (strcasecmp(name, "content-type") == 0) continue; + if (strcasecmp(name, "content-length") == 0) continue; + if (strcasecmp(name, "te") == 0) continue; + if (strcasecmp(name, "host") == 0) continue; + + snprintf(grpc_req->metadata[grpc_req->metadata_count][0], 256, "%s", name); + snprintf(grpc_req->metadata[grpc_req->metadata_count][1], 256, "%s", + req->headers[i].value); + grpc_req->metadata_count++; + } + + grpc_req->status = GRPC_OK; + return 0; +} + +/* ===== gRPC 消息帧编解码 ===== */ + +/** + * grpc_decode_message - 解码 gRPC 长度前缀消息帧 + * + * 帧格式: [compressed:1][length:4(BE)][payload:N] + * payload 动态分配,调用者需通过 grpc_message_free() 释放。 + */ +int grpc_decode_message(const uint8_t *buf, size_t len, grpc_message_t *msg) { + if (!buf || !msg || len == 0) return -1; + + /* 至少需要 5 字节前缀 */ + if (len < 5) return -1; + + /* 解析压缩标志 */ + msg->compressed = buf[0]; + + /* 解析长度(big-endian uint32) */ + msg->length = ((uint32_t)buf[1] << 24) | + ((uint32_t)buf[2] << 16) | + ((uint32_t)buf[3] << 8) | + (uint32_t)buf[4]; + + /* 检查数据完整性 */ + if (len < 5 + msg->length) return -1; + + /* 分配 payload 内存 */ + if (msg->length > 0) { + msg->payload = (uint8_t *)malloc(msg->length); + if (!msg->payload) return -1; + memcpy(msg->payload, buf + 5, msg->length); + } else { + msg->payload = NULL; + } + + return (int)(5 + msg->length); +} + +/** + * grpc_encode_message - 编码 gRPC 长度前缀消息帧 + * + * 帧格式: [compressed:1][length:4(BE)][payload:N] + */ +int grpc_encode_message(const grpc_message_t *msg, uint8_t *buf, size_t buf_size) { + if (!msg || !buf) return -1; + + /* 检查缓冲区是否足够(5 字节前缀 + payload) */ + if (buf_size < 5 + msg->length) return -1; + + /* 写入压缩标志 */ + buf[0] = msg->compressed; + + /* 写入长度(big-endian uint32) */ + buf[1] = (uint8_t)((msg->length >> 24) & 0xFF); + buf[2] = (uint8_t)((msg->length >> 16) & 0xFF); + buf[3] = (uint8_t)((msg->length >> 8) & 0xFF); + buf[4] = (uint8_t)(msg->length & 0xFF); + + /* 复制 payload */ + if (msg->length > 0 && msg->payload) { + memcpy(buf + 5, msg->payload, msg->length); + } + + return (int)(5 + msg->length); +} + +/** + * grpc_message_free - 释放 gRPC 消息中动态分配的资源 + */ +void grpc_message_free(grpc_message_t *msg) { + if (!msg) return; + + if (msg->payload) { + free(msg->payload); + msg->payload = NULL; + } + msg->length = 0; + msg->compressed = 0; +} + +/* ===== gRPC Trailers ===== */ + +/** + * grpc_format_response_trailers - 格式化 gRPC trailers + * + * gRPC 使用 HTTP/2 trailing headers 传递最终状态。 + * 格式: grpc-status: N\r\ngrpc-message: text\r\n + * + * gRPC-Web 不使用 trailers,而是在响应头中发送 grpc-status。 + */ +int grpc_format_response_trailers(const grpc_request_t *grpc_req, char *buf, size_t buf_size) { + if (!grpc_req || !buf || buf_size == 0) return -1; + + int n = snprintf(buf, buf_size, + "grpc-status: %d\r\n" + "grpc-message: %s\r\n", + (int)grpc_req->status, + grpc_req->status_message[0] ? grpc_req->status_message : ""); + + if ((size_t)n >= buf_size) return -1; + return n; +} + +/* ===== 状态码转换 ===== */ + +/** + * grpc_status_to_http - gRPC 状态码转 HTTP 状态码 + * + * 映射关系参照 gRPC 规范: + * https://github.com/grpc/grpc/blob/master/doc/http-grpc-status-mapping.md + */ +int grpc_status_to_http(grpc_status_t status) { + switch (status) { + case GRPC_OK: return 200; + case GRPC_CANCELLED: return 499; /* Client Closed Request */ + case GRPC_UNKNOWN: return 500; /* Internal Server Error */ + case GRPC_INVALID_ARGUMENT: return 400; /* Bad Request */ + case GRPC_DEADLINE_EXCEEDED: return 504; /* Gateway Timeout */ + case GRPC_NOT_FOUND: return 404; /* Not Found */ + case GRPC_ALREADY_EXISTS: return 409; /* Conflict */ + case GRPC_PERMISSION_DENIED: return 403; /* Forbidden */ + case GRPC_RESOURCE_EXHAUSTED: return 429; /* Too Many Requests */ + case GRPC_FAILED_PRECONDITION: return 400; /* Bad Request */ + case GRPC_ABORTED: return 409; /* Conflict */ + case GRPC_OUT_OF_RANGE: return 400; /* Bad Request */ + case GRPC_UNIMPLEMENTED: return 501; /* Not Implemented */ + case GRPC_INTERNAL: return 500; /* Internal Server Error */ + case GRPC_UNAVAILABLE: return 503; /* Service Unavailable */ + case GRPC_DATA_LOSS: return 500; /* Internal Server Error */ + case GRPC_UNAUTHENTICATED: return 401; /* Unauthorized */ + default: return 500; + } +} + +/** + * grpc_http_to_status - HTTP 状态码转 gRPC 状态码 + * + * HTTP → gRPC 反向映射。 + */ +grpc_status_t grpc_http_to_status(int http_status) { + switch (http_status) { + case 200: return GRPC_OK; + case 400: return GRPC_INVALID_ARGUMENT; + case 401: return GRPC_UNAUTHENTICATED; + case 403: return GRPC_PERMISSION_DENIED; + case 404: return GRPC_NOT_FOUND; + case 409: return GRPC_ABORTED; + case 412: return GRPC_FAILED_PRECONDITION; + case 429: return GRPC_RESOURCE_EXHAUSTED; + case 499: return GRPC_CANCELLED; + case 500: return GRPC_INTERNAL; + case 501: return GRPC_UNIMPLEMENTED; + case 503: return GRPC_UNAVAILABLE; + case 504: return GRPC_DEADLINE_EXCEEDED; + default: + if (http_status >= 200 && http_status < 300) return GRPC_OK; + if (http_status >= 400 && http_status < 500) return GRPC_INVALID_ARGUMENT; + return GRPC_INTERNAL; + } +} + +/** + * grpc_status_to_string - gRPC 状态码转可读字符串 + * + * @return 状态码名称,未知值返回 "UNKNOWN" + */ +const char *grpc_status_to_string(grpc_status_t status) { + switch (status) { + case GRPC_OK: return "OK"; + case GRPC_CANCELLED: return "CANCELLED"; + case GRPC_UNKNOWN: return "UNKNOWN"; + case GRPC_INVALID_ARGUMENT: return "INVALID_ARGUMENT"; + case GRPC_DEADLINE_EXCEEDED: return "DEADLINE_EXCEEDED"; + case GRPC_NOT_FOUND: return "NOT_FOUND"; + case GRPC_ALREADY_EXISTS: return "ALREADY_EXISTS"; + case GRPC_PERMISSION_DENIED: return "PERMISSION_DENIED"; + case GRPC_RESOURCE_EXHAUSTED: return "RESOURCE_EXHAUSTED"; + case GRPC_FAILED_PRECONDITION: return "FAILED_PRECONDITION"; + case GRPC_ABORTED: return "ABORTED"; + case GRPC_OUT_OF_RANGE: return "OUT_OF_RANGE"; + case GRPC_UNIMPLEMENTED: return "UNIMPLEMENTED"; + case GRPC_INTERNAL: return "INTERNAL"; + case GRPC_UNAVAILABLE: return "UNAVAILABLE"; + case GRPC_DATA_LOSS: return "DATA_LOSS"; + case GRPC_UNAUTHENTICATED: return "UNAUTHENTICATED"; + default: return "UNKNOWN"; + } +} + +/* ===== gRPC 响应发送 ===== */ + +/** + * grpc_send_unary_response - 发送 Unary RPC 响应 + * + * 将响应 payload 编码为 gRPC 消息帧格式,通过 socket 发送。 + * 编码格式: [flag:1][length:4(BE)][payload] + */ +int grpc_send_unary_response(cocoon_socket_t fd, grpc_request_t *grpc_req, + const uint8_t *resp_payload, size_t resp_len) { + if (fd == COCOON_INVALID_SOCKET || !grpc_req) return -1; + + /* 构建响应消息帧 */ + grpc_message_t resp_msg = { + .compressed = 0, /* 默认不压缩 */ + .length = (uint32_t)resp_len, + .payload = (uint8_t *)(uintptr_t)resp_payload /* const 转换,编码时不修改 */ + }; + + /* 分配编码缓冲区 */ + size_t buf_size = 5 + resp_len; + uint8_t *buf = (uint8_t *)malloc(buf_size); + if (!buf) return -1; + + int encoded = grpc_encode_message(&resp_msg, buf, buf_size); + if (encoded < 0) { + free(buf); + return -1; + } + + /* 通过 socket 发送 */ + ssize_t sent = cocoon_socket_send(fd, (const char *)buf, (size_t)encoded); + free(buf); + + if (sent < 0) return -1; + return (int)sent; +} + +/** + * grpc_send_trailers - 发送 gRPC trailers + * + * 格式化并发送 gRPC trailing headers。 + * 在当前简化实现中,将 trailers 格式化为文本缓冲区返回。 + * 实际 HTTP/2 trailers 发送需由调用方通过 nghttp2 完成。 + */ +int grpc_send_trailers(cocoon_socket_t fd, uint32_t stream_id, + grpc_status_t status, const char *message) { + if (fd == COCOON_INVALID_SOCKET) return -1; + + /* 构建临时 grpc_request_t 用于格式化 */ + grpc_request_t tmp_req = {0}; + tmp_req.status = status; + if (message) { + snprintf(tmp_req.status_message, sizeof(tmp_req.status_message), "%s", message); + } + + char trailers_buf[512]; + int n = grpc_format_response_trailers(&tmp_req, trailers_buf, sizeof(trailers_buf)); + if (n < 0) return -1; + + /* 发送 trailers 文本(实际应由 HTTP/2 层包装为 HEADERS 帧) */ + ssize_t sent = cocoon_socket_send(fd, trailers_buf, (size_t)n); + (void)stream_id; /* 在完整 HTTP/2 实现中用于 nghttp2_submit_trailer */ + + if (sent < 0) return -1; + return (int)sent; +} + +/** + * grpc_request_free - 释放 gRPC 请求资源 + * + * 释放 grpc_request_t 中所有动态分配的内存。 + * 包括 message.payload 和关联资源。 + */ +void grpc_request_free(grpc_request_t *grpc_req) { + if (!grpc_req) return; + + /* 释放消息 payload */ + grpc_message_free(&grpc_req->message); + + /* 清空其他字段 */ + grpc_req->metadata_count = 0; +} + +/** + * grpc_error_response - 发送 gRPC 错误响应 + * + * 直接发送 grpc-status trailers,不发送消息体。 + * 适用于请求解析失败、方法未找到、未实现等场景。 + */ +void grpc_error_response(cocoon_socket_t fd, uint32_t stream_id, + grpc_status_t status, const char *message) { + if (fd == COCOON_INVALID_SOCKET) return; + + const char *status_msg = message ? message : grpc_status_to_string(status); + + /* 发送错误 trailers */ + grpc_send_trailers(fd, stream_id, status, status_msg); +} diff --git a/grpc.h b/grpc.h new file mode 100644 index 0000000..08200ac --- /dev/null +++ b/grpc.h @@ -0,0 +1,267 @@ +/** + * @file grpc.h - gRPC over HTTP/2 支持模块 + * + * 基于现有 HTTP/2 传输层,实现 gRPC 协议支持。 + * 包含 gRPC 状态码、消息帧编解码、四种 RPC 模式、trailers 格式化。 + * + * gRPC message frame format: + * [1 byte: compressed flag] [4 bytes: length (big-endian)] [N bytes: payload] + * + * @author Cocoon Team + */ + +#ifndef COCOON_GRPC_H +#define COCOON_GRPC_H + +#include "http.h" +#include "platform.h" +#include +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +/* ===== gRPC 状态码 (17 个标准状态码) ===== */ +/** + * grpc_status_t - gRPC 标准状态码 + * + * 参照 https://grpc.io/docs/guides/status-codes/ + */ +typedef enum { + GRPC_OK = 0, /**< 成功 */ + GRPC_CANCELLED = 1, /**< 操作已取消 */ + GRPC_UNKNOWN = 2, /**< 未知错误 */ + GRPC_INVALID_ARGUMENT = 3, /**< 无效参数 */ + GRPC_DEADLINE_EXCEEDED = 4, /**< 超时 */ + GRPC_NOT_FOUND = 5, /**< 未找到 */ + GRPC_ALREADY_EXISTS = 6, /**< 已存在 */ + GRPC_PERMISSION_DENIED = 7, /**< 权限拒绝 */ + GRPC_RESOURCE_EXHAUSTED = 8, /**< 资源耗尽 */ + GRPC_FAILED_PRECONDITION = 9, /**< 前置条件失败 */ + GRPC_ABORTED = 10, /**< 操作已中止 */ + GRPC_OUT_OF_RANGE = 11, /**< 超出范围 */ + GRPC_UNIMPLEMENTED = 12, /**< 未实现 */ + GRPC_INTERNAL = 13, /**< 内部错误 */ + GRPC_UNAVAILABLE = 14, /**< 服务不可用 */ + GRPC_DATA_LOSS = 15, /**< 数据丢失 */ + GRPC_UNAUTHENTICATED = 16, /**< 未认证 */ + GRPC_STATUS_MAX /**< 状态码数量上限 */ +} grpc_status_t; + +/* ===== gRPC 消息帧 ===== */ +/** + * grpc_message_t - gRPC 长度前缀消息帧 + * + * gRPC 消息格式: + * [1 byte flag] [4 bytes length (big-endian)] [length bytes payload] + * flag: 0x00 = 未压缩, 0x01 = 压缩 (protobuf 数据) + * + * payload 为 opaque 字节流,不包含 protobuf 字段级解析。 + */ +typedef struct { + uint8_t compressed; /**< 压缩标志:0x00=未压缩, 0x01=压缩 */ + uint32_t length; /**< payload 长度 */ + uint8_t *payload; /**< 消息数据(动态分配) */ +} grpc_message_t; + +/* ===== gRPC 请求上下文 ===== */ +/** + * grpc_request_t - gRPC 请求上下文 + * + * 从 HTTP/2 请求解析得到的 gRPC 调用信息。 + * 包含 service/method 名、消息、元数据、流模式标记。 + */ +typedef struct { + char service_name[256]; /**< 服务名(从路径 /Service/Method 解析) */ + char method_name[256]; /**< 方法名 */ + grpc_message_t message; /**< 请求消息帧 */ + grpc_status_t status; /**< 响应状态 */ + char status_message[256]; /**< 状态消息文本 */ + + /* 元数据 (key-value pairs) */ + char metadata[16][2][256]; /**< 元数据键值对数组 */ + size_t metadata_count; /**< 元数据数量 */ + + /* RPC 流模式标记 */ + bool is_streaming; /**< 是否为流式 RPC(任一方向) */ + bool client_streaming; /**< 客户端流式(Client Streaming / Bidi) */ + bool server_streaming; /**< 服务端流式(Server Streaming / Bidi) */ + + /* gRPC-Web 兼容 */ + bool is_grpc_web; /**< 是否为 gRPC-Web 请求 */ +} grpc_request_t; + +/* ===== gRPC 流管理 ===== */ +/** + * grpc_stream_t - gRPC 流状态管理 + * + * 用于跟踪流式 RPC 的流生命周期。 + */ +typedef struct { + uint32_t stream_id; /**< HTTP/2 流 ID */ + cocoon_socket_t fd; /**< 客户端 socket */ + grpc_request_t *req; /**< 关联的请求上下文 */ + bool half_closed; /**< 客户端已半关闭(发送 END_STREAM) */ + bool closed; /**< 流已完全关闭 */ +} grpc_stream_t; + +/* ===== API ===== */ + +/** + * grpc_detect - 检测请求是否为 gRPC 请求 + * + * 检查 Content-Type 是否以 "application/grpc" 开头。 + * 支持 application/grpc、application/grpc+proto、application/grpc+json。 + * 也支持 application/grpc-web 变体(设置 is_grpc_web 标记)。 + * + * @param req HTTP 请求 + * @return true 是 gRPC/gRPC-Web 请求 + */ +bool grpc_detect(const http_request_t *req); + +/** + * grpc_is_grpc_web - 检测请求是否为 gRPC-Web + * + * @param req HTTP 请求 + * @return true 是 gRPC-Web 请求 + */ +bool grpc_is_grpc_web(const http_request_t *req); + +/** + * grpc_parse_request - 从 HTTP/2 请求解析 gRPC 上下文 + * + * 解析 :path 伪头获取 service/method,提取消息体。 + * 路径格式: "/package.Service/Method" 或 "/Service/Method" + * + * @param req HTTP 请求(已解析) + * @param grpc_req 输出 gRPC 请求上下文(调用者已 zero-init) + * @return 0 成功,-1 失败 + */ +int grpc_parse_request(const http_request_t *req, grpc_request_t *grpc_req); + +/** + * grpc_decode_message - 解码 gRPC 长度前缀消息帧 + * + * 从缓冲区解码一个 gRPC 消息帧。 + * 不解析 protobuf 内容,仅提取 opaque payload。 + * + * @param buf 输入缓冲区 + * @param len 缓冲区长度 + * @param msg 输出消息结构体(payload 动态分配) + * @return 成功解码的字节数(含 5 字节前缀),< 0 表示错误或数据不完整 + */ +int grpc_decode_message(const uint8_t *buf, size_t len, grpc_message_t *msg); + +/** + * grpc_encode_message - 编码 gRPC 长度前缀消息帧 + * + * @param msg 消息结构体 + * @param buf 输出缓冲区 + * @param buf_size 缓冲区大小 + * @return 编码后的字节数(含 5 字节前缀),< 0 表示缓冲区不足 + */ +int grpc_encode_message(const grpc_message_t *msg, uint8_t *buf, size_t buf_size); + +/** + * grpc_message_free - 释放 gRPC 消息中动态分配的资源 + * + * @param msg 消息结构体 + */ +void grpc_message_free(grpc_message_t *msg); + +/** + * grpc_format_response_trailers - 格式化 gRPC trailers(状态) + * + * gRPC 使用 HTTP/2 trailers 传递最终状态。 + * 格式: grpc-status: N\r\ngrpc-message: text\r\n + * + * @param grpc_req gRPC 请求(含 status 和 status_message) + * @param buf 输出缓冲区 + * @param buf_size 缓冲区大小 + * @return 写入的字节数,< 0 表示缓冲区不足 + */ +int grpc_format_response_trailers(const grpc_request_t *grpc_req, char *buf, size_t buf_size); + +/** + * grpc_status_to_http - gRPC 状态码转 HTTP 状态码 + * + * @param status gRPC 状态码 + * @return 对应的 HTTP 状态码 + */ +int grpc_status_to_http(grpc_status_t status); + +/** + * grpc_http_to_status - HTTP 状态码转 gRPC 状态码 + * + * @param http_status HTTP 状态码 + * @return 对应的 gRPC 状态码 + */ +grpc_status_t grpc_http_to_status(int http_status); + +/** + * grpc_status_to_string - gRPC 状态码转可读字符串 + * + * @param status gRPC 状态码 + * @return 状态码名称字符串 + */ +const char *grpc_status_to_string(grpc_status_t status); + +/** + * grpc_send_unary_response - 发送 Unary RPC 响应 + * + * 发送 gRPC 响应消息(DATA 帧)+ trailers。 + * 编码格式: [flag:1][length:4][payload] + * + * @param fd 客户端 socket + * @param grpc_req 请求上下文 + * @param resp_payload 响应 protobuf payload + * @param resp_len payload 长度 + * @return 成功发送的字节数,< 0 错误 + */ +int grpc_send_unary_response(cocoon_socket_t fd, grpc_request_t *grpc_req, + const uint8_t *resp_payload, size_t resp_len); + +/** + * grpc_send_trailers - 发送 gRPC trailers 结束流 + * + * 发送 HTTP/2 trailing headers 形式的 gRPC 状态。 + * + * @param fd 客户端 socket + * @param stream_id HTTP/2 流 ID + * @param status gRPC 状态码 + * @param message 状态消息(可为 NULL) + * @return 0 成功,< 0 错误 + */ +int grpc_send_trailers(cocoon_socket_t fd, uint32_t stream_id, + grpc_status_t status, const char *message); + +/** + * grpc_request_free - 释放 gRPC 请求资源 + * + * 释放 grpc_request_t 中所有动态分配的内存(包括 message.payload)。 + * + * @param grpc_req gRPC 请求上下文 + */ +void grpc_request_free(grpc_request_t *grpc_req); + +/** + * grpc_error_response - 发送 gRPC 错误响应 + * + * 直接发送带有 grpc-status 的错误 trailers,不发送消息体。 + * 用于请求解析失败、方法未找到等场景。 + * + * @param fd 客户端 socket + * @param stream_id HTTP/2 流 ID + * @param status gRPC 错误状态码 + * @param message 错误消息(可为 NULL) + */ +void grpc_error_response(cocoon_socket_t fd, uint32_t stream_id, + grpc_status_t status, const char *message); + +#ifdef __cplusplus +} +#endif + +#endif /* COCOON_GRPC_H */ diff --git a/http3.c b/http3.c new file mode 100644 index 0000000..40db253 --- /dev/null +++ b/http3.c @@ -0,0 +1,1671 @@ +/** + * @file http3.c + * @brief HTTP/3 (QUIC) 支持模块实现 + * + * 包含: + * - QUIC Variable-Length Integer 编解码 + * - HTTP/3 帧处理(HEADERS / DATA / SETTINGS / GOAWAY) + * - QPACK 静态表头部压缩(RFC 9204 Appendix A) + * - QUIC 流管理(创建、数据读写、FIN 控制) + * - QUIC 连接管理(ID 分配、超时、关闭) + * - HTTP/3 会话管理 + * - TLS 1.3 集成接口(条件编译) + * + * @author Cocoon Team + */ + +#include "http3.h" +#include +#include +#include +#include +#include + +/* ===== TLS 条件编译 ===== */ +#if defined(OPENSSL_HAS_QUIC) && defined(OPENSSL_VERSION_NUMBER) + #include +#endif + +/* ===== 全局连接链表 ===== */ +static quic_connection_t *g_connections = NULL; +static size_t g_connection_count = 0; + +/* ===== QPACK 静态表(RFC 9204 Appendix A) ===== */ +static const qpack_static_entry_t s_qpack_static_table[QPACK_STATIC_TABLE_SIZE] = { + /* 0 */ {":authority", ""}, + /* 1 */ {":path", "/"}, + /* 2 */ {":method", "GET"}, + /* 3 */ {":method", "POST"}, + /* 4 */ {":scheme", "https"}, + /* 5 */ {":scheme", "http"}, + /* 6 */ {":status", "200"}, + /* 7 */ {":status", "204"}, + /* 8 */ {":status", "206"}, + /* 9 */ {":status", "304"}, + /* 10 */ {":status", "400"}, + /* 11 */ {":status", "404"}, + /* 12 */ {":status", "500"}, + /* 13 */ {"accept-charset", ""}, + /* 14 */ {"accept-encoding", "gzip, deflate, br"}, + /* 15 */ {"accept-language", ""}, + /* 16 */ {"accept-ranges", ""}, + /* 17 */ {"accept", ""}, + /* 18 */ {"access-control-allow-origin", ""}, + /* 19 */ {"age", ""}, + /* 20 */ {"allow", ""}, + /* 21 */ {"authorization", ""}, + /* 22 */ {"cache-control", ""}, + /* 23 */ {"content-disposition", ""}, + /* 24 */ {"content-encoding", ""}, + /* 25 */ {"content-language", ""}, + /* 26 */ {"content-length", ""}, + /* 27 */ {"content-location", ""}, + /* 28 */ {"content-range", ""}, + /* 29 */ {"content-type", ""}, + /* 30 */ {"cookie", ""}, + /* 31 */ {"date", ""}, + /* 32 */ {"etag", ""}, + /* 33 */ {"expect", ""}, + /* 34 */ {"expires", ""}, + /* 35 */ {"from", ""}, + /* 36 */ {"host", ""}, + /* 37 */ {"if-match", ""}, + /* 38 */ {"if-modified-since", ""}, + /* 39 */ {"if-none-match", ""}, + /* 40 */ {"if-range", ""}, + /* 41 */ {"if-unmodified-since", ""}, + /* 42 */ {"last-modified", ""}, + /* 43 */ {"link", ""}, + /* 44 */ {"location", ""}, + /* 45 */ {"max-forwards", ""}, + /* 46 */ {"proxy-authenticate", ""}, + /* 47 */ {"proxy-authorization", ""}, + /* 48 */ {"range", ""}, + /* 49 */ {"referer", ""}, + /* 50 */ {"refresh", ""}, + /* 51 */ {"retry-after", ""}, + /* 52 */ {"server", ""}, + /* 53 */ {"set-cookie", ""}, + /* 54 */ {"strict-transport-security", ""}, + /* 55 */ {"transfer-encoding", ""}, + /* 56 */ {"user-agent", ""}, + /* 57 */ {"vary", ""}, + /* 58 */ {"via", ""}, + /* 59 */ {"www-authenticate", ""}, + /* 60 */ {":status", "100"}, + /* 61 */ {":status", "103"}, + /* 62 */ {"accept-encoding", ""}, + /* 63 */ {":path", ""}, + /* 64 */ {"content-type", "application/json"}, + /* 65 */ {"content-type", "text/html"}, + /* 66 */ {"content-type", "text/plain"}, + /* 67 */ {":authority", ""}, + /* 68 */ {":method", "CONNECT"}, + /* 69 */ {":method", "DELETE"}, + /* 70 */ {":method", "HEAD"}, + /* 71 */ {":method", "OPTIONS"}, + /* 72 */ {":method", "PUT"}, + /* 73 */ {":scheme", "https"}, + /* 74 */ {":scheme", "http"}, + /* 75 */ {":status", "100"}, + /* 76 */ {":status", "101"}, + /* 77 */ {":status", "103"}, + /* 78 */ {":status", "201"}, + /* 79 */ {":status", "301"}, + /* 80 */ {":status", "302"}, + /* 81 */ {":status", "304"}, + /* 82 */ {":status", "400"}, + /* 83 */ {":status", "401"}, + /* 84 */ {":status", "403"}, + /* 85 */ {":status", "404"}, + /* 86 */ {":status", "405"}, + /* 87 */ {":status", "406"}, + /* 88 */ {":status", "407"}, + /* 89 */ {":status", "408"}, + /* 90 */ {":status", "409"}, + /* 91 */ {":status", "410"}, + /* 92 */ {":status", "411"}, + /* 93 */ {":status", "412"}, + /* 94 */ {":status", "413"}, + /* 95 */ {":status", "414"}, + /* 96 */ {":status", "415"}, + /* 97 */ {":status", "416"}, + /* 98 */ {":status", "421"}, +}; + +/* ===== Variable-Length Integer 编解码 ===== */ + +/** + * @brief 编码 QUIC variable-length integer + */ +size_t http3_encode_varint(uint64_t value, uint8_t *buf) { + if (!buf) return 0; + + if (value <= 63ULL) { + buf[0] = (uint8_t)value; + return 1; + } else if (value <= 16383ULL) { + buf[0] = (uint8_t)(((value >> 8) & 0x3F) | 0x40); + buf[1] = (uint8_t)(value & 0xFF); + return 2; + } else if (value <= 1073741823ULL) { + buf[0] = (uint8_t)(((value >> 24) & 0x3F) | 0x80); + buf[1] = (uint8_t)((value >> 16) & 0xFF); + buf[2] = (uint8_t)((value >> 8) & 0xFF); + buf[3] = (uint8_t)(value & 0xFF); + return 4; + } else { + buf[0] = (uint8_t)(((value >> 56) & 0x3F) | 0xC0); + buf[1] = (uint8_t)((value >> 48) & 0xFF); + buf[2] = (uint8_t)((value >> 40) & 0xFF); + buf[3] = (uint8_t)((value >> 32) & 0xFF); + buf[4] = (uint8_t)((value >> 24) & 0xFF); + buf[5] = (uint8_t)((value >> 16) & 0xFF); + buf[6] = (uint8_t)((value >> 8) & 0xFF); + buf[7] = (uint8_t)(value & 0xFF); + return 8; + } +} + +/** + * @brief 解码 QUIC variable-length integer + */ +int http3_decode_varint(const uint8_t *buf, size_t len, uint64_t *value) { + if (!buf || !value || len == 0) return -1; + + uint8_t prefix = (buf[0] & 0xC0) >> 6; + size_t need = 1; + uint64_t result = 0; + + switch (prefix) { + case 0: + need = 1; + result = buf[0] & 0x3F; + break; + case 1: + need = 2; + if (len < need) return -1; + result = (((uint64_t)(buf[0] & 0x3F)) << 8) | buf[1]; + break; + case 2: + need = 4; + if (len < need) return -1; + result = (((uint64_t)(buf[0] & 0x3F)) << 24) | + (((uint64_t)buf[1]) << 16) | + (((uint64_t)buf[2]) << 8) | + (uint64_t)buf[3]; + break; + case 3: + need = 8; + if (len < need) return -1; + result = (((uint64_t)(buf[0] & 0x3F)) << 56) | + (((uint64_t)buf[1]) << 48) | + (((uint64_t)buf[2]) << 40) | + (((uint64_t)buf[3]) << 32) | + (((uint64_t)buf[4]) << 24) | + (((uint64_t)buf[5]) << 16) | + (((uint64_t)buf[6]) << 8) | + (uint64_t)buf[7]; + break; + } + + *value = result; + return (int)need; +} + +/** + * @brief 计算帧头编码后的大小 + */ +size_t http3_frame_header_size(uint64_t frame_type, uint64_t length) { + uint8_t dummy[16]; + return http3_encode_varint(frame_type, dummy) + + http3_encode_varint(length, dummy + 8); +} + +/* ===== HTTP/3 帧处理 ===== */ + +/** + * @brief 编码 HTTP/3 帧头 + */ +size_t http3_encode_frame_header(uint64_t frame_type, uint64_t length, + uint8_t *buf) { + if (!buf) return 0; + size_t n = 0; + n += http3_encode_varint(frame_type, buf); + n += http3_encode_varint(length, buf + n); + return n; +} + +/** + * @brief 解码 HTTP/3 帧头 + */ +int http3_decode_frame_header(const uint8_t *buf, size_t len, + uint64_t *frame_type, uint64_t *length) { + if (!buf || !frame_type || !length || len == 0) return -1; + + uint64_t ft = 0, flen = 0; + int n1 = http3_decode_varint(buf, len, &ft); + if (n1 < 0) return -1; + + int n2 = http3_decode_varint(buf + n1, len - (size_t)n1, &flen); + if (n2 < 0) return -1; + + *frame_type = ft; + *length = flen; + return n1 + n2; +} + +/** + * @brief 编码完整 HTTP/3 帧 + */ +int http3_encode_frame(uint64_t frame_type, + const uint8_t *payload, size_t payload_len, + uint8_t *buf, size_t buf_cap) { + if (!buf) return -1; + size_t header_len = http3_frame_header_size(frame_type, payload_len); + if (header_len + payload_len > buf_cap) return -1; + + size_t n = 0; + n += http3_encode_varint(frame_type, buf); + n += http3_encode_varint(payload_len, buf + n); + + if (payload && payload_len > 0) { + memcpy(buf + n, payload, payload_len); + n += payload_len; + } + return (int)n; +} + +/** + * @brief 从 QUIC 流接收缓冲区解析完整帧 + */ +int http3_parse_frame(quic_stream_t *stream, + uint64_t *frame_type, + const uint8_t **payload, + size_t *payload_len) { + if (!stream || !frame_type || !payload || !payload_len) return -1; + if (stream->recv_buf_len == 0) return 1; /* 数据不足 */ + + uint64_t ft = 0, flen = 0; + int header_len = http3_decode_frame_header(stream->recv_buf, + stream->recv_buf_len, + &ft, &flen); + if (header_len < 0) return 1; /* 数据不足 */ + + if ((size_t)header_len + flen > stream->recv_buf_len) return 1; /* 数据不足 */ + + *frame_type = ft; + *payload = stream->recv_buf + header_len; + *payload_len = (size_t)flen; + + /* 消费已解析的帧 */ + size_t total = (size_t)header_len + (size_t)flen; + if (total < stream->recv_buf_len) { + memmove(stream->recv_buf, stream->recv_buf + total, + stream->recv_buf_len - total); + } + stream->recv_buf_len -= total; + + return 0; +} + +/* ===== QPACK 编码/解码 ===== */ + +/** + * @brief 在静态表中查找完全匹配的条目 + * + * @param name 字段名 + * @param value 字段值 + * @return 索引(>= 0),未找到返回 -1 + */ +static int qpack_find_static_index(const char *name, const char *value) { + if (!name) return -1; + + for (int i = 0; i < QPACK_STATIC_TABLE_SIZE; i++) { + if (strcmp(s_qpack_static_table[i].name, name) == 0) { + if (!value || strcmp(s_qpack_static_table[i].value, value) == 0) { + return i; + } + } + } + return -1; +} + +/** + * @brief 在静态表中查找名称匹配的条目 + * + * @param name 字段名 + * @return 索引(>= 0),未找到返回 -1 + */ +static int qpack_find_static_name_index(const char *name) { + if (!name) return -1; + + for (int i = 0; i < QPACK_STATIC_TABLE_SIZE; i++) { + if (strcmp(s_qpack_static_table[i].name, name) == 0) { + return i; + } + } + return -1; +} + +/** + * @brief 编码字面量字符串(带长度前缀) + * + * @param str 字符串 + * @param out 输出缓冲区 + * @param out_cap 缓冲区容量 + * @return 编码后字节数,< 0 错误 + */ +static int qpack_encode_literal(const char *str, uint8_t *out, size_t out_cap) { + if (!str || !out) return -1; + size_t len = strlen(str); + uint8_t len_buf[8]; + size_t len_bytes = http3_encode_varint(len, len_buf); + if (len_bytes + len > out_cap) return -1; + + memcpy(out, len_buf, len_bytes); + memcpy(out + len_bytes, str, len); + return (int)(len_bytes + len); +} + +/** + * @brief 解码字面量字符串(带长度前缀) + * + * @param in 输入缓冲区 + * @param in_len 输入长度 + * @param out 输出字符串缓冲区 + * @param out_cap 输出缓冲区容量 + * @param consumed 消耗字节数 + * @return 解码后字符串长度,< 0 错误 + */ +static int qpack_decode_literal(const uint8_t *in, size_t in_len, + char *out, size_t out_cap, + size_t *consumed) { + if (!in || !out || !consumed || in_len == 0) return -1; + + uint64_t len = 0; + int n = http3_decode_varint(in, in_len, &len); + if (n < 0) return -1; + if ((size_t)n + len > in_len) return -1; + + size_t copy_len = (len < (uint64_t)out_cap) ? (size_t)len : out_cap - 1; + memcpy(out, in + n, copy_len); + out[copy_len] = '\0'; + *consumed = (size_t)n + (size_t)len; + return (int)copy_len; +} + +/** + * @brief QPACK 编码头部字段 + * + * 编码策略: + * 1. 静态表完全匹配 → 索引编码(0b11XXXXXX...) + * 2. 静态表名称匹配 → 字面量编码,引用名称索引 + * 3. 完全不匹配 → 字面量编码,原始名称+值 + */ +int qpack_encode_header(const char *name, const char *value, + uint8_t *out, size_t out_cap) { + if (!name || !value || !out || out_cap == 0) return -1; + + /* 策略1:静态表完全匹配 */ + int idx = qpack_find_static_index(name, value); + if (idx >= 0) { + /* Indexed Field Line: 0b1TXXXXXX ... (T=1 for static table) */ + /* QPACK 静态表索引是 0-based */ + uint64_t encoded_idx = (uint64_t)idx; + if (encoded_idx < 63) { + if (out_cap < 1) return -1; + out[0] = (uint8_t)(0xC0 | (encoded_idx & 0x3F)); + return 1; + } else { + /* 多字节编码 */ + if (out_cap < 2) return -1; + out[0] = (uint8_t)(0xC0 | 0x3F); /* 0b111111 = multi-byte marker */ + out[1] = (uint8_t)(encoded_idx - 63); + return 2; + } + } + + /* 策略2:静态表名称匹配,字面量值 */ + int name_idx = qpack_find_static_name_index(name); + if (name_idx >= 0) { + /* Literal Field Line with Name Reference: 0b0101XXXX ... */ + /* QPACK 静态表索引是 0-based */ + uint64_t encoded_idx = (uint64_t)name_idx; + size_t n = 0; + if (encoded_idx < 15) { + if (out_cap < 2) return -1; + out[n++] = (uint8_t)(0x50 | (encoded_idx & 0x0F)); + } else { + /* 多字节编码 */ + if (out_cap < 3) return -1; + out[n++] = (uint8_t)(0x5F); /* 0b1111 = multi-byte marker */ + out[n++] = (uint8_t)(encoded_idx - 15); + } + int val_n = qpack_encode_literal(value, out + n, out_cap - n); + if (val_n < 0) return -1; + return (int)(n + (size_t)val_n); + } + + /* 策略3:完全字面量编码 */ + /* Literal Field Line with Literal Name: 0b001XXXXX [name] [value] */ + if (out_cap < 1) return -1; + size_t n = 0; + out[n++] = 0x20; /* 0b00100000 */ + int name_n = qpack_encode_literal(name, out + n, out_cap - n); + if (name_n < 0) return -1; + n += (size_t)name_n; + int val_n = qpack_encode_literal(value, out + n, out_cap - n); + if (val_n < 0) return -1; + return (int)(n + (size_t)val_n); +} + +/** + * @brief QPACK 解码 prefix 整数(RFC 9204 Section 4.1.1) + * + * QPACK 使用前缀编码:前 N 位表示值的一部分。 + * 如果值超出前缀位数,则后续字节使用 base-128 延续格式。 + * + * @param in 输入缓冲区 + * @param in_len 输入长度 + * @param prefix_bits 前缀位数(6 或 4) + * @param out_value 输出值 + * @return 消耗字节数,< 0 错误 + */ +static int qpack_decode_prefix_int(const uint8_t *in, size_t in_len, + int prefix_bits, uint64_t *out_value) { + if (!in || !out_value || in_len == 0) return -1; + + uint8_t prefix_mask = (uint8_t)((1U << prefix_bits) - 1); + uint64_t value = (uint64_t)(in[0] & prefix_mask); + + if (value < (uint64_t)prefix_mask) { + *out_value = value; + return 1; + } + + /* Multi-byte: following bytes use base-128 continuation (RFC 9204) */ + size_t pos = 1; + uint64_t m = 0; + while (pos < in_len) { + uint8_t byte = in[pos]; + pos++; + value += ((uint64_t)(byte & 0x7F)) << m; + m += 7; + if ((byte & 0x80) == 0) break; + if (pos > 8) return -1; /* Too many bytes */ + } + + *out_value = value; + return (int)pos; +} + +/** + * @brief QPACK 解码头部字段 + */ +int qpack_decode_header(const uint8_t *in, size_t in_len, + qpack_decoded_t *out, size_t *consumed) { + if (!in || !out || !consumed || in_len == 0) return -1; + + memset(out, 0, sizeof(*out)); + size_t pos = 0; + + uint8_t first = in[pos]; + + /* 判断编码类型 */ + if ((first & 0x80) != 0) { + /* Indexed Field Line: 1TXXXXXX ... */ + /* T bit: 1=static table, 0=dynamic table */ + bool is_static = (first & 0x40) != 0; + if (!is_static) { + /* Dynamic table not implemented */ + return -1; + } + + uint64_t idx = 0; + int n = qpack_decode_prefix_int(in + pos, in_len - pos, 6, &idx); + if (n < 0) return -1; + pos += (size_t)n; + + /* Static table index is 1-based */ + if (idx >= (uint64_t)QPACK_STATIC_TABLE_SIZE) return -1; + const qpack_static_entry_t *entry = &s_qpack_static_table[(size_t)idx]; + size_t nlen = strlen(entry->name); + size_t vlen = strlen(entry->value); + memcpy(out->name, entry->name, nlen + 1); + memcpy(out->value, entry->value, vlen + 1); + out->valid = true; + *consumed = pos; + return 0; + } else if ((first & 0x40) != 0) { + /* Literal Field Line with Name Reference: 01NTXXXX ... */ + /* N bit = 0x20 (post-base), T bit = 0x10 (1=static, 0=dynamic) */ + bool is_static_name = (first & 0x10) != 0; + if (!is_static_name) { + /* Dynamic table not implemented */ + return -1; + } + bool value_is_huffman = (first & 0x08) != 0; + (void)value_is_huffman; /* Huffman decoding not implemented */ + + uint64_t idx = 0; + int n = qpack_decode_prefix_int(in + pos, in_len - pos, 4, &idx); + if (n < 0) return -1; + pos += (size_t)n; + + if (idx >= (uint64_t)QPACK_STATIC_TABLE_SIZE) return -1; + const qpack_static_entry_t *entry = &s_qpack_static_table[(size_t)idx]; + memcpy(out->name, entry->name, strlen(entry->name) + 1); + + /* Decode value literal */ + size_t val_consumed = 0; + int val_len = qpack_decode_literal(in + pos, in_len - pos, + out->value, sizeof(out->value), + &val_consumed); + if (val_len < 0) return -1; + pos += val_consumed; + + out->valid = true; + *consumed = pos; + return 0; + } else if ((first & 0x20) != 0) { + /* Literal Field Line with Literal Name: 001XXXXX [H|name_len] [name] [value] */ + bool name_is_huffman = (first & 0x10) != 0; + (void)name_is_huffman; + pos++; + + /* Decode name */ + size_t name_consumed = 0; + int name_len = qpack_decode_literal(in + pos, in_len - pos, + out->name, sizeof(out->name), + &name_consumed); + if (name_len < 0) return -1; + pos += name_consumed; + + /* Decode value */ + size_t val_consumed = 0; + int val_len = qpack_decode_literal(in + pos, in_len - pos, + out->value, sizeof(out->value), + &val_consumed); + if (val_len < 0) return -1; + pos += val_consumed; + + out->valid = true; + *consumed = pos; + return 0; + } else { + /* 其他类型(动态表更新等)未实现 */ + return -1; + } +} + +/** + * @brief QPACK 编码完整请求头 + */ +int qpack_encode_request_headers(const http_request_t *req, + uint8_t *out, size_t out_cap) { + if (!req || !out || out_cap == 0) return -1; + + size_t total = 0; + + /* Required pseudo-headers for request: + * :method, :scheme, :authority, :path + */ + int n; + + /* :method */ + const char *method_str = http_method_str(req->method); + n = qpack_encode_header(":method", method_str, out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + + /* :scheme (always https for HTTP/3) */ + n = qpack_encode_header(":scheme", "https", out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + + /* :authority (from Host header) */ + const char *authority = ""; + for (int i = 0; i < req->num_headers; i++) { + if (strcasecmp(req->headers[i].name, "host") == 0) { + authority = req->headers[i].value; + break; + } + } + n = qpack_encode_header(":authority", authority, out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + + /* :path */ + n = qpack_encode_header(":path", req->path, out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + + /* Regular headers */ + for (int i = 0; i < req->num_headers; i++) { + n = qpack_encode_header(req->headers[i].name, req->headers[i].value, + out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + } + + return (int)total; +} + +/** + * @brief QPACK 解码完整请求头 + */ +int qpack_decode_request_headers(const uint8_t *in, size_t in_len, + http_request_t *req) { + if (!in || !req) return -1; + + memset(req, 0, sizeof(*req)); + req->method = HTTP_UNKNOWN; + req->content_length = -1; + + size_t pos = 0; + while (pos < in_len) { + qpack_decoded_t decoded; + size_t consumed = 0; + if (qpack_decode_header(in + pos, in_len - pos, &decoded, &consumed) != 0) { + return -1; + } + pos += consumed; + + if (decoded.valid) { + if (strcmp(decoded.name, ":method") == 0) { + if (strcmp(decoded.value, "GET") == 0) req->method = HTTP_GET; + else if (strcmp(decoded.value, "HEAD") == 0) req->method = HTTP_HEAD; + else if (strcmp(decoded.value, "POST") == 0) req->method = HTTP_POST; + else if (strcmp(decoded.value, "PUT") == 0) req->method = HTTP_PUT; + else if (strcmp(decoded.value, "DELETE") == 0) req->method = HTTP_DELETE; + else if (strcmp(decoded.value, "OPTIONS") == 0) req->method = HTTP_OPTIONS; + else req->method = HTTP_UNKNOWN; + } else if (strcmp(decoded.name, ":path") == 0) { + size_t plen = strlen(decoded.value); + if (plen > sizeof(req->path) - 1) plen = sizeof(req->path) - 1; + memcpy(req->path, decoded.value, plen); + req->path[plen] = '\0'; + } else if (strcmp(decoded.name, ":authority") == 0) { + if (req->num_headers < HTTP_MAX_HEADERS) { + memcpy(req->headers[req->num_headers].name, "Host", 5); + size_t vcopy = strlen(decoded.value); + if (vcopy > sizeof(req->headers[0].value) - 1) + vcopy = sizeof(req->headers[0].value) - 1; + memcpy(req->headers[req->num_headers].value, decoded.value, vcopy); + req->headers[req->num_headers].value[vcopy] = '\0'; + req->num_headers++; + } + } else if (decoded.name[0] != ':') { + /* Regular header (skip pseudo-headers starting with ":") */ + if (req->num_headers < HTTP_MAX_HEADERS) { + size_t ncopy = strlen(decoded.name); + if (ncopy > sizeof(req->headers[0].name) - 1) + ncopy = sizeof(req->headers[0].name) - 1; + memcpy(req->headers[req->num_headers].name, decoded.name, ncopy); + req->headers[req->num_headers].name[ncopy] = '\0'; + size_t vcopy = strlen(decoded.value); + if (vcopy > sizeof(req->headers[0].value) - 1) + vcopy = sizeof(req->headers[0].value) - 1; + memcpy(req->headers[req->num_headers].value, decoded.value, vcopy); + req->headers[req->num_headers].value[vcopy] = '\0'; + if (strcmp(decoded.name, "content-length") == 0) { + req->content_length = atoll(decoded.value); + } + req->num_headers++; + } + } + } + } + + return 0; +} + +/** + * @brief QPACK 编码响应头 + */ +int qpack_encode_response_headers(const http_response_t *resp, + uint8_t *out, size_t out_cap) { + if (!resp || !out || out_cap == 0) return -1; + + size_t total = 0; + int n; + char status_str[16]; + + /* :status pseudo-header */ + snprintf(status_str, sizeof(status_str), "%d", resp->status_code); + n = qpack_encode_header(":status", status_str, out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + + /* Content-Type */ + if (resp->content_type && resp->content_type[0]) { + n = qpack_encode_header("content-type", resp->content_type, + out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + } + + /* Content-Length */ + if (resp->content_length >= 0) { + char clen_str[32]; + snprintf(clen_str, sizeof(clen_str), "%ld", (long)resp->content_length); + n = qpack_encode_header("content-length", clen_str, + out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + } + + /* Server */ + n = qpack_encode_header("server", "Cocoon/1.0 (HTTP/3)", + out + total, out_cap - total); + if (n < 0) return -1; + total += (size_t)n; + + return (int)total; +} + +/** + * @brief QPACK 解码响应头 + */ +int qpack_decode_response_headers(const uint8_t *in, size_t in_len, + http_response_t *resp) { + if (!in || !resp) return -1; + + memset(resp, 0, sizeof(*resp)); + resp->status_code = 0; + resp->content_length = -1; + + size_t pos = 0; + while (pos < in_len) { + qpack_decoded_t decoded; + size_t consumed = 0; + if (qpack_decode_header(in + pos, in_len - pos, &decoded, &consumed) != 0) { + return -1; + } + pos += consumed; + + if (decoded.valid) { + if (strcmp(decoded.name, ":status") == 0) { + resp->status_code = atoi(decoded.value); + } else if (strcmp(decoded.name, "content-type") == 0) { + resp->content_type = strdup(decoded.value); + } else if (strcmp(decoded.name, "content-length") == 0) { + resp->content_length = atoll(decoded.value); + } + } + } + + return 0; +} + + +/* ===== QUIC 连接管理 ===== */ + +/** + * @brief 获取当前时间戳(毫秒) + */ +uint64_t quic_current_time_ms(void) { + struct timespec ts; + if (clock_gettime(CLOCK_MONOTONIC, &ts) != 0) { + return (uint64_t)(time(NULL) * 1000ULL); + } + return (uint64_t)(ts.tv_sec * 1000ULL + ts.tv_nsec / 1000000ULL); +} + +/** + * @brief 生成 64-bit 随机连接 ID + */ +uint64_t quic_generate_conn_id(void) { + /* 简单的伪随机生成器,使用当前时间作为种子 */ + static unsigned int seed = 0; + if (seed == 0) { + seed = (unsigned int)time(NULL); + } + uint64_t id = 0; + id |= ((uint64_t)rand_r(&seed) << 48) & 0xFFFF000000000000ULL; + id |= ((uint64_t)rand_r(&seed) << 32) & 0x0000FFFF00000000ULL; + id |= ((uint64_t)rand_r(&seed) << 16) & 0x00000000FFFF0000ULL; + id |= ((uint64_t)rand_r(&seed)) & 0x000000000000FFFFULL; + return id; +} + +/** + * @brief 创建 QUIC 连接 + */ +quic_connection_t *quic_connection_create(uint64_t conn_id, + cocoon_socket_t udp_fd, const struct sockaddr_storage *peer_addr) { + + quic_connection_t *conn = (quic_connection_t *)calloc(1, sizeof(quic_connection_t)); + if (!conn) return NULL; + + conn->conn_id = conn_id; + conn->udp_fd = udp_fd; + conn->handshake_complete = false; + conn->tls_conn = NULL; + conn->max_streams_bidi = QUIC_MAX_STREAMS_PER_CONN; + conn->next_stream_id = 0; /* 客户端 bidirectional 从 0 开始 */ + conn->streams = NULL; + conn->streams_tail = NULL; + conn->bytes_received = 0; + conn->bytes_sent = 0; + conn->idle_timeout_ms = QUIC_DEFAULT_IDLE_TIMEOUT; + conn->last_activity = quic_current_time_ms(); + conn->closed = false; + conn->closing = false; + conn->next = NULL; + + if (peer_addr) { + memcpy(&conn->peer_addr, peer_addr, sizeof(*peer_addr)); + conn->peer_addr_len = sizeof(*peer_addr); + } + + /* 插入全局链表头部 */ + conn->next = g_connections; + g_connections = conn; + g_connection_count++; + + return conn; +} + +/** + * @brief 销毁 QUIC 连接 + */ +void quic_connection_destroy(quic_connection_t *conn) { + if (!conn) return; + + /* 释放所有流 */ + quic_stream_t *stream = conn->streams; + while (stream) { + quic_stream_t *next = stream->next; + if (stream->recv_buf) { + free(stream->recv_buf); + } + free(stream); + stream = next; + } + conn->streams = NULL; + conn->streams_tail = NULL; + + /* 释放 TLS 连接 */ + if (conn->tls_conn) { +#if defined(OPENSSL_HAS_QUIC) && defined(OPENSSL_VERSION_NUMBER) + /* 使用 OpenSSL QUIC-TLS 接口 */ + SSL_free((SSL *)conn->tls_conn); +#else + free(conn->tls_conn); +#endif + conn->tls_conn = NULL; + } + + /* 从全局链表中移除 */ + quic_connection_t **pp = &g_connections; + while (*pp) { + if (*pp == conn) { + *pp = conn->next; + g_connection_count--; + break; + } + pp = &(*pp)->next; + } + + free(conn); +} + +/** + * @brief 在全局链表中查找 QUIC 连接 + */ +quic_connection_t *quic_find_connection(uint64_t conn_id) { + quic_connection_t *conn = g_connections; + while (conn) { + if (conn->conn_id == conn_id && !conn->closed) { + return conn; + } + conn = conn->next; + } + return NULL; +} + +/** + * @brief 移除并清理超时连接 + */ +void quic_cleanup_timeout_connections(uint64_t timeout_ms) { + uint64_t now = quic_current_time_ms(); + quic_connection_t *conn = g_connections; + quic_connection_t *prev = NULL; + + while (conn) { + quic_connection_t *next = conn->next; + if (!conn->closed && (now - conn->last_activity) > timeout_ms) { + /* 超时,移除连接 */ + if (prev) { + prev->next = next; + } else { + g_connections = next; + } + g_connection_count--; + + /* 清理连接 */ + quic_stream_t *stream = conn->streams; + while (stream) { + quic_stream_t *snext = stream->next; + if (stream->recv_buf) free(stream->recv_buf); + free(stream); + stream = snext; + } + if (conn->tls_conn) { +#if defined(OPENSSL_HAS_QUIC) && defined(OPENSSL_VERSION_NUMBER) + SSL_free((SSL *)conn->tls_conn); +#else + free(conn->tls_conn); +#endif + } + free(conn); + } else { + prev = conn; + } + conn = next; + } +} + +/** + * @brief 获取活跃 QUIC 连接数 + */ +size_t quic_get_connection_count(void) { + return g_connection_count; +} + +/** + * @brief 发送 QUIC 数据报 + */ +int quic_send_datagram(quic_connection_t *conn, const uint8_t *data, size_t len) { + if (!conn || !data || len == 0) return -1; + if (conn->udp_fd == COCOON_INVALID_SOCKET) return -1; + + ssize_t sent = sendto(conn->udp_fd, (const char *)data, len, 0, + (const struct sockaddr *)&conn->peer_addr, + conn->peer_addr_len); + if (sent < 0) return -1; + + conn->bytes_sent += (uint64_t)sent; + conn->last_activity = quic_current_time_ms(); + return 0; +} + +/* ===== QUIC 流管理 ===== */ + +/** + * @brief 获取或创建 QUIC 流 + */ +quic_stream_t *quic_stream_get_or_create(quic_connection_t *conn, uint64_t stream_id) { + if (!conn) return NULL; + + /* 查找现有流 */ + quic_stream_t *stream = conn->streams; + while (stream) { + if (stream->stream_id == stream_id) { + return stream; + } + stream = stream->next; + } + + /* 创建新流 */ + stream = (quic_stream_t *)calloc(1, sizeof(quic_stream_t)); + if (!stream) return NULL; + + stream->stream_id = stream_id; + stream->offset = 0; + stream->recv_offset = 0; + stream->peer_fin = false; + stream->local_fin = false; + stream->reset = false; + stream->recv_buf = NULL; + stream->recv_buf_len = 0; + stream->recv_buf_cap = 0; + stream->conn = conn; + stream->next = NULL; + + /* 插入链表尾部 */ + if (conn->streams_tail) { + conn->streams_tail->next = stream; + } else { + conn->streams = stream; + } + conn->streams_tail = stream; + + return stream; +} + +/** + * @brief 销毁 QUIC 流 + */ +void quic_stream_destroy(quic_connection_t *conn, quic_stream_t *stream) { + if (!conn || !stream) return; + + /* 从链表中移除 */ + quic_stream_t **pp = &conn->streams; + while (*pp) { + if (*pp == stream) { + *pp = stream->next; + if (conn->streams_tail == stream) { + conn->streams_tail = (*pp == NULL) ? NULL : conn->streams; + /* Fix tail pointer */ + if (conn->streams_tail == NULL && conn->streams != NULL) { + quic_stream_t *s = conn->streams; + while (s->next) s = s->next; + conn->streams_tail = s; + } + } + break; + } + pp = &(*pp)->next; + } + + if (stream->recv_buf) { + free(stream->recv_buf); + } + free(stream); +} + +/** + * @brief 查找 QUIC 流 + */ +quic_stream_t *quic_stream_find(quic_connection_t *conn, uint64_t stream_id) { + if (!conn) return NULL; + + quic_stream_t *stream = conn->streams; + while (stream) { + if (stream->stream_id == stream_id) { + return stream; + } + stream = stream->next; + } + return NULL; +} + +/** + * @brief 向 QUIC 流写入数据(追加到流的接收缓冲区) + */ +int quic_stream_write(quic_stream_t *stream, const uint8_t *data, size_t len) { + if (!stream || !data) return -1; + + /* 确保缓冲区容量足够 */ + size_t need = stream->recv_buf_len + len; + if (need > stream->recv_buf_cap) { + size_t new_cap = stream->recv_buf_cap * 2; + if (new_cap == 0) new_cap = 4096; + while (new_cap < need) new_cap *= 2; + + uint8_t *new_buf = (uint8_t *)realloc(stream->recv_buf, new_cap); + if (!new_buf) return -1; + + stream->recv_buf = new_buf; + stream->recv_buf_cap = new_cap; + } + + memcpy(stream->recv_buf + stream->recv_buf_len, data, len); + stream->recv_buf_len += len; + stream->recv_offset += len; + + if (stream->conn) { + stream->conn->last_activity = quic_current_time_ms(); + } + + return 0; +} + +/** + * @brief 从 QUIC 流读取数据 + */ +ssize_t quic_stream_read(quic_stream_t *stream, uint8_t *buf, size_t len) { + if (!stream || !buf) return -1; + + size_t to_read = stream->recv_buf_len < len ? stream->recv_buf_len : len; + if (to_read == 0) return 0; + + memcpy(buf, stream->recv_buf, to_read); + + /* 消费已读取的数据 */ + if (to_read < stream->recv_buf_len) { + memmove(stream->recv_buf, stream->recv_buf + to_read, + stream->recv_buf_len - to_read); + } + stream->recv_buf_len -= to_read; + + return (ssize_t)to_read; +} + +/** + * @brief 设置 QUIC 流 FIN 标志 + */ +void quic_stream_set_fin(quic_stream_t *stream) { + if (!stream) return; + stream->local_fin = true; +} + +/* ===== HTTP/3 会话管理 ===== */ + +/** + * @brief 全局 HTTP/3 初始化 + */ +bool http3_init(void) { + /* 重置全局状态 */ + g_connections = NULL; + g_connection_count = 0; + + /* 初始化随机种子 */ + srand((unsigned int)time(NULL)); + + return true; +} + +/** + * @brief 全局 HTTP/3 清理 + */ +void http3_cleanup(void) { + /* 销毁所有连接 */ + while (g_connections) { + quic_connection_t *conn = g_connections; + g_connections = conn->next; + + quic_stream_t *stream = conn->streams; + while (stream) { + quic_stream_t *next = stream->next; + if (stream->recv_buf) free(stream->recv_buf); + free(stream); + stream = next; + } + if (conn->tls_conn) { +#if defined(OPENSSL_HAS_QUIC) && defined(OPENSSL_VERSION_NUMBER) + SSL_free((SSL *)conn->tls_conn); +#else + free(conn->tls_conn); +#endif + } + free(conn); + } + g_connection_count = 0; +} + +/** + * @brief 创建 HTTP/3 会话 + */ +http3_session_t *http3_session_create(quic_connection_t *conn) { + if (!conn) return NULL; + + http3_session_t *session = (http3_session_t *)calloc(1, sizeof(http3_session_t)); + if (!session) return NULL; + + session->conn = conn; + session->max_field_section_size = HTTP3_DEFAULT_MAX_FIELD_SECTION_SIZE; + session->qpack_encoder_max_capacity = 0; /* 不使用动态表 */ + session->qpack_decoder_max_capacity = 0; + session->goaway_stream_id = UINT64_MAX; + session->settings_received = false; + session->settings_sent = false; + + /* 发送服务端 SETTINGS */ + http3_send_settings(session); + + return session; +} + +/** + * @brief 销毁 HTTP/3 会话 + */ +void http3_session_destroy(http3_session_t *session) { + if (!session) return; + + /* 释放所有 HTTP/3 流 */ + for (int i = 0; i < QUIC_MAX_STREAMS_PER_CONN; i++) { + if (session->h3_streams[i]) { + /* 不销毁底层 QUIC 流,只释放 HTTP/3 包装 */ + free(session->h3_streams[i]); + session->h3_streams[i] = NULL; + } + } + + free(session); +} + +/** + * @brief 发送 SETTINGS 帧 + */ +int http3_send_settings(http3_session_t *session) { + if (!session || !session->conn) return -1; + + uint8_t payload[64]; + size_t pos = 0; + + /* 编码 SETTINGS 参数(键值对 varint) */ + pos += http3_encode_varint(HTTP3_SETTING_MAX_FIELD_SECTION_SIZE, + payload + pos); + pos += http3_encode_varint(session->max_field_section_size, + payload + pos); + pos += http3_encode_varint(HTTP3_SETTING_QPACK_MAX_TABLE_CAPACITY, + payload + pos); + pos += http3_encode_varint(0, payload + pos); /* 0 capacity = no dynamic table */ + + /* 编码完整帧 */ + uint8_t frame[128]; + int frame_len = http3_encode_frame(HTTP3_FRAME_SETTINGS, payload, pos, + frame, sizeof(frame)); + if (frame_len < 0) return -1; + + /* 通过流 ID 2(服务器控制流)发送 */ + /* 简化:直接发送数据报 */ + if (quic_send_datagram(session->conn, frame, (size_t)frame_len) != 0) { + return -1; + } + + session->settings_sent = true; + return 0; +} + +/** + * @brief 处理 HTTP/3 控制流 + */ +int http3_handle_control_stream(http3_session_t *session, quic_stream_t *stream, + const uint8_t *data, size_t len) { + if (!session || !stream || !data || len == 0) return -1; + + /* 将数据追加到流的接收缓冲区 */ + if (quic_stream_write(stream, data, len) != 0) { + return -1; + } + + /* 尝试解析帧 */ + uint64_t frame_type = 0; + const uint8_t *payload = NULL; + size_t payload_len = 0; + + while (http3_parse_frame(stream, &frame_type, &payload, &payload_len) == 0) { + switch (frame_type) { + case HTTP3_FRAME_SETTINGS: { + /* 解析 SETTINGS 帧 */ + size_t pos = 0; + while (pos + 2 <= payload_len) { + uint64_t setting_id = 0, setting_value = 0; + int n1 = http3_decode_varint(payload + pos, payload_len - pos, + &setting_id); + if (n1 < 0) break; + pos += (size_t)n1; + + int n2 = http3_decode_varint(payload + pos, payload_len - pos, + &setting_value); + if (n2 < 0) break; + pos += (size_t)n2; + + if (setting_id == HTTP3_SETTING_MAX_FIELD_SECTION_SIZE) { + session->max_field_section_size = setting_value; + } + } + session->settings_received = true; + break; + } + case HTTP3_FRAME_GOAWAY: { + /* 解析 GOAWAY 帧 */ + uint64_t stream_id = 0; + if (payload_len >= 1) { + http3_decode_varint(payload, payload_len, &stream_id); + session->goaway_stream_id = stream_id; + session->conn->closing = true; + } + break; + } + default: + /* 忽略未知帧 */ + break; + } + } + + return 0; +} + +/* ===== HTTP/3 请求/响应处理 ===== */ + +/** + * @brief 查找或创建 HTTP/3 流 + */ +static http3_stream_t *h3_stream_get_or_create(http3_session_t *session, + uint64_t stream_id) { + if (!session) return NULL; + + /* 查找流在数组中的索引 */ + size_t idx = (size_t)(stream_id / 4); + if (idx >= QUIC_MAX_STREAMS_PER_CONN) return NULL; + + if (session->h3_streams[idx]) { + return session->h3_streams[idx]; + } + + /* 创建新的 HTTP/3 流 */ + http3_stream_t *h3s = (http3_stream_t *)calloc(1, sizeof(http3_stream_t)); + if (!h3s) return NULL; + + /* 获取或创建底层 QUIC 流 */ + h3s->qstream = quic_stream_get_or_create(session->conn, stream_id); + if (!h3s->qstream) { + free(h3s); + return NULL; + } + + h3s->headers_received = false; + h3s->data_received = false; + h3s->headers_sent = false; + h3s->trailers_sent = false; + h3s->error_code = HTTP3_NO_ERROR; + h3s->request_complete = false; + h3s->response_complete = false; + + session->h3_streams[idx] = h3s; + return h3s; +} + +/** + * @brief 从 HTTP/3 流读取请求 + */ +int64_t http3_read_request(http3_session_t *session, http_request_t *req) { + if (!session || !req || !session->conn) return -1; + + /* 遍历所有 HTTP/3 流,寻找有完整请求的 */ + for (int i = 0; i < QUIC_MAX_STREAMS_PER_CONN; i++) { + http3_stream_t *h3s = session->h3_streams[i]; + if (!h3s || !h3s->qstream) continue; + + quic_stream_t *qstream = h3s->qstream; + if (qstream->recv_buf_len == 0) continue; + + /* 尝试解析帧 */ + uint64_t frame_type = 0; + const uint8_t *payload = NULL; + size_t payload_len = 0; + + /* 保存缓冲区状态,以便回滚 */ + size_t saved_len = qstream->recv_buf_len; + uint8_t *saved_buf = NULL; + if (saved_len > 0) { + saved_buf = (uint8_t *)malloc(saved_len); + if (saved_buf) memcpy(saved_buf, qstream->recv_buf, saved_len); + } + + int parse_result = http3_parse_frame(qstream, &frame_type, &payload, &payload_len); + + if (parse_result != 0 || frame_type != HTTP3_FRAME_HEADERS) { + /* 恢复缓冲区 */ + if (saved_buf) { + if (qstream->recv_buf_cap < saved_len) { + uint8_t *new_buf = (uint8_t *)realloc(qstream->recv_buf, saved_len); + if (new_buf) { + qstream->recv_buf = new_buf; + qstream->recv_buf_cap = saved_len; + } + } + if (qstream->recv_buf) { + memcpy(qstream->recv_buf, saved_buf, saved_len); + qstream->recv_buf_len = saved_len; + } + free(saved_buf); + } + continue; + } + + free(saved_buf); + + /* 解码 HEADERS */ + if (qpack_decode_request_headers(payload, payload_len, req) != 0) { + continue; + } + + h3s->headers_received = true; + + /* 检查是否有 DATA 帧跟随 */ + if (qstream->recv_buf_len > 0) { + size_t saved_len2 = qstream->recv_buf_len; + uint8_t *saved_buf2 = (uint8_t *)malloc(saved_len2); + if (saved_buf2) memcpy(saved_buf2, qstream->recv_buf, saved_len2); + + uint64_t data_frame_type = 0; + const uint8_t *data_payload = NULL; + size_t data_payload_len = 0; + + if (http3_parse_frame(qstream, &data_frame_type, + &data_payload, &data_payload_len) == 0 && + data_frame_type == HTTP3_FRAME_DATA) { + /* 有 DATA 帧 */ + if (data_payload_len > 0) { + req->body = (char *)malloc(data_payload_len + 1); + if (req->body) { + memcpy(req->body, data_payload, data_payload_len); + req->body[data_payload_len] = '\0'; + req->body_len = data_payload_len; + } + } + h3s->data_received = true; + } else { + /* 恢复缓冲区 */ + if (saved_buf2) { + if (qstream->recv_buf_cap < saved_len2) { + uint8_t *nb = realloc(qstream->recv_buf, saved_len2); + if (nb) { + qstream->recv_buf = nb; + qstream->recv_buf_cap = saved_len2; + } + } + if (qstream->recv_buf) { + memcpy(qstream->recv_buf, saved_buf2, saved_len2); + qstream->recv_buf_len = saved_len2; + } + } + } + free(saved_buf2); + } + + /* 如果 HEADERS 帧已收到,检查 FIN */ + if (qstream->peer_fin || !h3s->data_received) { + h3s->request_complete = true; + } + + return (int64_t)qstream->stream_id; + } + + return -1; /* 没有完整请求 */ +} + +/** + * @brief 发送 HTTP/3 响应 + */ +int http3_send_response(http3_session_t *session, uint64_t stream_id, + const http_response_t *resp, + const uint8_t *body, size_t body_len) { + if (!session || !session->conn) return -1; + + http3_stream_t *h3s = h3_stream_get_or_create(session, stream_id); + if (!h3s) return -1; + + /* 编码响应头 */ + uint8_t headers_buf[4096]; + int headers_len = qpack_encode_response_headers(resp, headers_buf, + sizeof(headers_buf)); + if (headers_len < 0) return -1; + + /* 编码 HEADERS 帧 */ + uint8_t headers_frame[8192]; + int headers_frame_len = http3_encode_frame(HTTP3_FRAME_HEADERS, + headers_buf, (size_t)headers_len, + headers_frame, sizeof(headers_frame)); + if (headers_frame_len < 0) return -1; + + /* 发送 HEADERS 帧 */ + if (quic_send_datagram(session->conn, headers_frame, + (size_t)headers_frame_len) != 0) { + return -1; + } + + h3s->headers_sent = true; + + /* 发送 DATA 帧(如果有 body) */ + if (body && body_len > 0) { + uint8_t data_frame_buf[16384]; + int data_frame_len = http3_encode_frame(HTTP3_FRAME_DATA, + body, body_len, + data_frame_buf, + sizeof(data_frame_buf)); + if (data_frame_len < 0) return -1; + + if (quic_send_datagram(session->conn, data_frame_buf, + (size_t)data_frame_len) != 0) { + return -1; + } + } + + h3s->response_complete = true; + return 0; +} + +/** + * @brief 发送 HTTP/3 错误响应 + */ +void http3_send_error(http3_session_t *session, uint64_t stream_id, + int status_code, const char *message) { + if (!session) return; + + http_response_t resp = {0}; + resp.status_code = status_code; + resp.content_type = "text/html; charset=utf-8"; + + /* 构造简单错误页面 */ + char body[1024]; + int body_len = snprintf(body, sizeof(body), + "%d Error" + "

%d %s

%s

", + status_code, status_code, + status_code == 404 ? "Not Found" : + status_code == 500 ? "Internal Server Error" : + status_code == 400 ? "Bad Request" : + status_code == 405 ? "Method Not Allowed" : "Error", + message ? message : ""); + + resp.content_length = body_len; + + http3_send_response(session, stream_id, &resp, + (const uint8_t *)body, (size_t)body_len); +} + +/** + * @brief 关闭 QUIC 连接 + */ +void http3_close_connection(quic_connection_t *conn, http3_error_t error) { + if (!conn || conn->closed) return; + + (void)error; /* 错误码可用于发送 CONNECTION_CLOSE 帧 */ + + conn->closed = true; + conn->closing = true; + + /* 发送 GOAWAY 帧 */ + if (conn->udp_fd != COCOON_INVALID_SOCKET) { + uint8_t goaway_payload[16]; + size_t pos = 0; + pos += http3_encode_varint(UINT64_MAX, goaway_payload + pos); + + uint8_t frame[32]; + int frame_len = http3_encode_frame(HTTP3_FRAME_GOAWAY, + goaway_payload, pos, + frame, sizeof(frame)); + if (frame_len > 0) { + quic_send_datagram(conn, frame, (size_t)frame_len); + } + } +} + +/* ===== UDP 服务器循环 ===== */ + +/** + * @brief 处理 UDP 数据报 + * + * 简化版处理:不实现完整 QUIC 数据报解析, + * 而是将数据直接分发给对应连接的流。 + */ +void http3_process_datagram(cocoon_socket_t udp_fd, + const uint8_t *buf, size_t len, + const struct sockaddr_storage *peer_addr) { + if (!buf || len == 0 || !peer_addr) return; + + /* 简化处理:假设数据报前 8 字节是连接 ID */ + uint64_t conn_id = 0; + if (len >= 8) { + conn_id = ((uint64_t)buf[0] << 56) | + ((uint64_t)buf[1] << 48) | + ((uint64_t)buf[2] << 40) | + ((uint64_t)buf[3] << 32) | + ((uint64_t)buf[4] << 24) | + ((uint64_t)buf[5] << 16) | + ((uint64_t)buf[6] << 8) | + (uint64_t)buf[7]; + } + + /* 查找或创建连接 */ + quic_connection_t *conn = quic_find_connection(conn_id); + if (!conn) { + conn = quic_connection_create(conn_id, udp_fd, peer_addr); + if (!conn) return; + } + + /* 更新活动时间 */ + conn->last_activity = quic_current_time_ms(); + conn->bytes_received += len; + + /* 简化:假设剩余数据包含 [stream_id:8][data...] + * 在实际完整实现中,这里需要解析 QUIC 数据包头、包号、帧等 + */ + if (len > 16) { + uint64_t stream_id = ((uint64_t)buf[8] << 56) | + ((uint64_t)buf[9] << 48) | + ((uint64_t)buf[10] << 40) | + ((uint64_t)buf[11] << 32) | + ((uint64_t)buf[12] << 24) | + ((uint64_t)buf[13] << 16) | + ((uint64_t)buf[14] << 8) | + (uint64_t)buf[15]; + + quic_stream_t *stream = quic_stream_get_or_create(conn, stream_id); + if (stream) { + quic_stream_write(stream, buf + 16, len - 16); + } + } +} + +/** + * @brief 接受新的 QUIC 连接 + * + * 从 UDP socket 接收数据报并处理。 + */ +quic_connection_t *http3_accept(cocoon_socket_t udp_fd) { + if (udp_fd == COCOON_INVALID_SOCKET) return NULL; + + uint8_t buf[QUIC_MAX_DATAGRAM_SIZE]; + struct sockaddr_storage peer_addr; + socklen_t peer_addr_len = sizeof(peer_addr); + + ssize_t n = recvfrom(udp_fd, (char *)buf, sizeof(buf), 0, + (struct sockaddr *)&peer_addr, &peer_addr_len); + if (n <= 0) return NULL; + + http3_process_datagram(udp_fd, buf, (size_t)n, &peer_addr); + + /* 返回新创建或找到的连接 */ + if (n >= 8) { + uint64_t conn_id = ((uint64_t)buf[0] << 56) | + ((uint64_t)buf[1] << 48) | + ((uint64_t)buf[2] << 40) | + ((uint64_t)buf[3] << 32) | + ((uint64_t)buf[4] << 24) | + ((uint64_t)buf[5] << 16) | + ((uint64_t)buf[6] << 8) | + (uint64_t)buf[7]; + return quic_find_connection(conn_id); + } + + return NULL; +} diff --git a/http3.h b/http3.h new file mode 100644 index 0000000..d51e153 --- /dev/null +++ b/http3.h @@ -0,0 +1,625 @@ +/** + * @file http3.h + * @brief HTTP/3 (QUIC) 支持模块头文件 + * + * 实现 HTTP/3 over QUIC 协议支持,包括: + * - QUIC 基础传输层(简化版):UDP socket、连接 ID、流管理 + * - HTTP/3 帧处理:HEADERS、DATA、SETTINGS、GOAWAY 等 + * - QPACK 头部压缩:静态表编码/解码(RFC 9204 Appendix A) + * - Variable-length integer 编解码 + * - TLS 1.3 集成接口(条件编译,OpenSSL 3.2+) + * + * @author Cocoon Team + */ + +#ifndef COCOON_HTTP3_H +#define COCOON_HTTP3_H + +#include "http.h" +#include "platform.h" +#include +#include +#include +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +/* ===== HTTP/3 帧类型 ===== */ +#define HTTP3_FRAME_DATA 0x00 +#define HTTP3_FRAME_HEADERS 0x01 +#define HTTP3_FRAME_CANCEL_PUSH 0x03 +#define HTTP3_FRAME_SETTINGS 0x04 +#define HTTP3_FRAME_PUSH_PROMISE 0x05 +#define HTTP3_FRAME_GOAWAY 0x06 +#define HTTP3_FRAME_MAX_PUSH_ID 0x07 +#define HTTP3_FRAME_PRIORITY_UPDATE_REQ 0xF0700 +#define HTTP3_FRAME_PRIORITY_UPDATE_PUSH 0xF0701 + +/* ===== HTTP/3 设置参数 ===== */ +#define HTTP3_SETTING_MAX_FIELD_SECTION_SIZE 0x06 +#define HTTP3_DEFAULT_MAX_FIELD_SECTION_SIZE (16384) /**< 16KB */ +#define HTTP3_SETTING_QPACK_MAX_TABLE_CAPACITY 0x01 +#define HTTP3_SETTING_BLOCKED_STREAMS 0x07 + +/* ===== QUIC 常量 ===== */ +#define QUIC_MAX_DATAGRAM_SIZE 1200 +#define QUIC_MAX_STREAMS_PER_CONN 100 +#define QUIC_DEFAULT_IDLE_TIMEOUT 30000 /**< 30 秒(毫秒) */ +#define QUIC_MAX_CONN_ID_LEN 8 +#define HTTP3_MAX_HEADER_ENTRIES 64 +#define QPACK_STATIC_TABLE_SIZE 99 + +/* ===== HTTP/3 错误码 ===== */ +typedef enum { + HTTP3_NO_ERROR = 0x0100, + HTTP3_GENERAL_PROTOCOL_ERROR = 0x0101, + HTTP3_INTERNAL_ERROR = 0x0102, + HTTP3_STREAM_CREATION_ERROR = 0x0103, + HTTP3_CLOSED_CRITICAL_STREAM = 0x0104, + HTTP3_FRAME_UNEXPECTED = 0x0105, + HTTP3_FRAME_ERROR = 0x0106, + HTTP3_EXCESSIVE_LOAD = 0x0107, + HTTP3_ID_ERROR = 0x0108, + HTTP3_SETTINGS_ERROR = 0x0109, + HTTP3_MISSING_SETTINGS = 0x010A, + HTTP3_REQUEST_REJECTED = 0x010B, + HTTP3_REQUEST_CANCELLED = 0x010C, + HTTP3_REQUEST_INCOMPLETE = 0x010D, + HTTP3_EARLY_RESPONSE = 0x010E, + HTTP3_CONNECT_ERROR = 0x010F, + HTTP3_VERSION_FALLBACK = 0x0110 +} http3_error_t; + +/* ===== 前向声明 ===== */ +typedef struct quic_stream quic_stream_t; +typedef struct quic_connection quic_connection_t; + +/** + * @brief QPACK 静态表条目 + * + * RFC 9204 Appendix A 定义的静态表。 + * 包含常见 HTTP 字段名和值的预定义索引。 + */ +typedef struct { + const char *name; /**< 字段名 */ + const char *value; /**< 字段值(可为空字符串) */ +} qpack_static_entry_t; + +/** + * @brief QUIC 流结构 + * + * 每个 QUIC 流对应一个此结构体。 + * 流 ID 的低 2 位表示类型: + * - 0x00: 客户端发起的双向流 + * - 0x01: 服务器发起的双向流 + * - 0x02: 客户端发起的单向流 + * - 0x03: 服务器发起的单向流 + */ +struct quic_stream { + uint64_t stream_id; /**< 流 ID */ + uint64_t offset; /**< 当前发送偏移 */ + uint64_t recv_offset; /**< 当前接收偏移 */ + bool peer_fin; /**< 对端已发送 FIN */ + bool local_fin; /**< 本端已发送 FIN */ + bool reset; /**< 流是否被重置 */ + uint8_t *recv_buf; /**< 接收缓冲区(动态分配) */ + size_t recv_buf_len; /**< 接收缓冲区已用长度 */ + size_t recv_buf_cap; /**< 接收缓冲区容量 */ + quic_connection_t *conn; /**< 所属连接 */ + quic_stream_t *next; /**< 链表下一个节点 */ +}; + +/** + * @brief QUIC 连接结构 + * + * 简化版 QUIC 连接管理。 + * 不实现完整拥塞控制,依赖内核 UDP 传输。 + */ +struct quic_connection { + uint64_t conn_id; /**< 64-bit 连接 ID */ + cocoon_socket_t udp_fd; /**< 底层 UDP socket */ + struct sockaddr_storage peer_addr; /**< 对端地址 */ + socklen_t peer_addr_len; /**< 对端地址长度 */ + bool handshake_complete; /**< TLS 1.3 握手完成标志 */ + void *tls_conn; /**< TLS 连接指针( opaque,避免暴露 tls_conn_t ) */ + uint32_t max_streams_bidi; /**< 最大双向流数 */ + uint32_t next_stream_id; /**< 下一个客户端发起的双向流 ID */ + quic_stream_t *streams; /**< 活跃流链表头 */ + quic_stream_t *streams_tail; /**< 活跃流链表尾 */ + uint64_t bytes_received; /**< 接收字节数 */ + uint64_t bytes_sent; /**< 发送字节数 */ + uint64_t idle_timeout_ms; /**< 空闲超时(毫秒) */ + uint64_t last_activity; /**< 最后活动时间戳(毫秒) */ + bool closed; /**< 连接已关闭 */ + bool closing; /**< 连接正在关闭 */ + quic_connection_t *next; /**< 全局连接链表下一个节点 */ +}; + +/** + * @brief HTTP/3 流结构 + * + * 封装 QUIC 流,添加 HTTP/3 语义状态。 + */ +typedef struct { + quic_stream_t *qstream; /**< 底层 QUIC 流 */ + bool headers_received; /**< 已接收 HEADERS 帧 */ + bool data_received; /**< 已接收 DATA 帧 */ + bool headers_sent; /**< 已发送 HEADERS 帧 */ + bool trailers_sent; /**< 已发送 trailers */ + http3_error_t error_code; /**< 流错误码 */ + bool request_complete; /**< 请求接收完整 */ + bool response_complete; /**< 响应发送完整 */ +} http3_stream_t; + +/** + * @brief HTTP/3 会话结构 + * + * 管理 HTTP/3 连接状态,包含 SETTINGS 和 QPACK 上下文。 + */ +typedef struct { + quic_connection_t *conn; /**< 底层 QUIC 连接 */ + uint64_t max_field_section_size; /**< SETTINGS: 最大字段段大小 */ + uint64_t qpack_encoder_max_capacity; /**< QPACK 编码器最大容量 */ + uint64_t qpack_decoder_max_capacity; /**< QPACK 解码器最大容量 */ + http3_stream_t *h3_streams[QUIC_MAX_STREAMS_PER_CONN]; /**< HTTP/3 流数组 */ + uint64_t goaway_stream_id; /**< GOAWAY 流 ID */ + bool settings_received; /**< 已收到客户端 SETTINGS */ + bool settings_sent; /**< 已发送服务端 SETTINGS */ +} http3_session_t; + +/** + * @brief QPACK 编码头部字段结果 + */ +typedef struct { + uint8_t *data; /**< 编码后的数据(动态分配) */ + size_t len; /**< 编码后长度 */ +} qpack_encoded_t; + +/** + * @brief QPACK 解码头部字段结果 + */ +typedef struct { + char name[HTTP_HEADER_NAME_MAX]; /**< 字段名 */ + char value[HTTP_HEADER_VALUE_MAX]; /**< 字段值 */ + bool valid; /**< 解码是否成功 */ +} qpack_decoded_t; + +/* ===== 全局初始化 / 清理 ===== */ + +/** + * @brief 全局 HTTP/3 初始化 + * + * 初始化 QPACK 静态表等全局资源。 + * 应在服务器启动前调用一次。 + * + * @return true 成功,false 失败 + */ +bool http3_init(void); + +/** + * @brief 全局 HTTP/3 清理 + * + * 释放所有全局资源。 + */ +void http3_cleanup(void); + +/* ===== 会话管理 ===== */ + +/** + * @brief 创建 HTTP/3 会话 + * + * 为给定的 QUIC 连接创建 HTTP/3 会话, + * 发送初始 SETTINGS 帧。 + * + * @param conn QUIC 连接 + * @return HTTP/3 会话,失败返回 NULL + */ +http3_session_t *http3_session_create(quic_connection_t *conn); + +/** + * @brief 销毁 HTTP/3 会话 + * + * 释放会话及其管理的所有 HTTP/3 流。 + * + * @param session HTTP/3 会话 + */ +void http3_session_destroy(http3_session_t *session); + +/* ===== QUIC 连接管理 ===== */ + +/** + * @brief 接受新的 QUIC 连接 + * + * 从 UDP socket 接收数据报,查找或创建 QUIC 连接。 + * + * @param udp_fd UDP socket 描述符 + * @return 新连接或已有连接,NULL 表示无新连接 + */ +quic_connection_t *http3_accept(cocoon_socket_t udp_fd); + +/** + * @brief 创建 QUIC 连接(内部使用/测试用) + * + * @param conn_id 连接 ID + * @param udp_fd UDP socket + * @param peer_addr 对端地址 + * @return 新连接,失败返回 NULL + */ +quic_connection_t *quic_connection_create(uint64_t conn_id, + cocoon_socket_t udp_fd, const struct sockaddr_storage *peer_addr); + +/** + * @brief 销毁 QUIC 连接 + * + * @param conn QUIC 连接 + */ +void quic_connection_destroy(quic_connection_t *conn); + +/** + * @brief 获取或创建 QUIC 流 + * + * @param conn QUIC 连接 + * @param stream_id 流 ID + * @return 流指针,失败返回 NULL + */ +quic_stream_t *quic_stream_get_or_create(quic_connection_t *conn, uint64_t stream_id); + +/** + * @brief 销毁 QUIC 流 + * + * @param conn 所属连接 + * @param stream 要销毁的流 + */ +void quic_stream_destroy(quic_connection_t *conn, quic_stream_t *stream); + +/** + * @brief 查找 QUIC 流 + * + * @param conn QUIC 连接 + * @param stream_id 流 ID + * @return 流指针,未找到返回 NULL + */ +quic_stream_t *quic_stream_find(quic_connection_t *conn, uint64_t stream_id); + +/** + * @brief 向 QUIC 流写入数据 + * + * @param stream QUIC 流 + * @param data 数据 + * @param len 数据长度 + * @return 0 成功,< 0 错误 + */ +int quic_stream_write(quic_stream_t *stream, const uint8_t *data, size_t len); + +/** + * @brief 从 QUIC 流读取数据 + * + * @param stream QUIC 流 + * @param buf 输出缓冲区 + * @param len 最大读取长度 + * @return 读取字节数,< 0 错误 + */ +ssize_t quic_stream_read(quic_stream_t *stream, uint8_t *buf, size_t len); + +/** + * @brief 设置 QUIC 流 FIN 标志 + * + * @param stream QUIC 流 + */ +void quic_stream_set_fin(quic_stream_t *stream); + +/* ===== HTTP/3 请求/响应处理 ===== */ + +/** + * @brief 从 HTTP/3 流读取请求 + * + * 解析 HEADERS 帧(QPACK 解码)和 DATA 帧,构建 http_request_t。 + * + * @param session HTTP/3 会话 + * @param req 输出 HTTP 请求 + * @return 成功返回流 ID(>= 0),< 0 表示无可用请求 + */ +int64_t http3_read_request(http3_session_t *session, http_request_t *req); + +/** + * @brief 发送 HTTP/3 响应 + * + * 发送 HEADERS 帧(QPACK 编码)和 DATA 帧。 + * + * @param session HTTP/3 会话 + * @param stream_id 流 ID + * @param resp HTTP 响应 + * @param body 响应体数据 + * @param body_len 响应体长度 + * @return 0 成功,< 0 错误 + */ +int http3_send_response(http3_session_t *session, uint64_t stream_id, + const http_response_t *resp, + const uint8_t *body, size_t body_len); + +/** + * @brief 发送 HTTP/3 错误响应 + * + * @param session HTTP/3 会话 + * @param stream_id 流 ID + * @param status_code HTTP 状态码 + * @param message 错误消息 + */ +void http3_send_error(http3_session_t *session, uint64_t stream_id, + int status_code, const char *message); + +/** + * @brief 关闭 QUIC 连接 + * + * @param conn QUIC 连接 + * @param error HTTP/3 错误码 + */ +void http3_close_connection(quic_connection_t *conn, http3_error_t error); + +/** + * @brief 发送 SETTINGS 帧 + * + * @param session HTTP/3 会话 + * @return 0 成功,< 0 错误 + */ +int http3_send_settings(http3_session_t *session); + +/** + * @brief 处理 UDP 数据报 + * + * 接收并处理 UDP 数据报,分发给对应的 QUIC 连接。 + * + * @param udp_fd UDP socket + * @param buf 数据报内容 + * @param len 数据报长度 + * @param peer_addr 对端地址 + */ +void http3_process_datagram(cocoon_socket_t udp_fd, + const uint8_t *buf, size_t len, + const struct sockaddr_storage *peer_addr); + +/** + * @brief 处理 HTTP/3 控制流 + * + * 处理 SETTINGS、GOAWAY 等控制帧。 + * + * @param session HTTP/3 会话 + * @param stream QUIC 控制流 + * @param data 数据 + * @param len 数据长度 + * @return 0 成功,< 0 错误 + */ +int http3_handle_control_stream(http3_session_t *session, quic_stream_t *stream, + const uint8_t *data, size_t len); + +/* ===== Variable-Length Integer 编解码 ===== */ + +/** + * @brief 编码 QUIC variable-length integer + * + * 2-bit prefix 指示长度: + * - 00 = 1 byte (0..63) + * - 01 = 2 bytes (0..16383) + * - 10 = 4 bytes (0..1073741823) + * - 11 = 8 bytes (0..4611686018427387903) + * + * @param value 要编码的值 + * @param buf 输出缓冲区(至少 8 字节) + * @return 编码后的字节数(1, 2, 4, or 8) + */ +size_t http3_encode_varint(uint64_t value, uint8_t *buf); + +/** + * @brief 解码 QUIC variable-length integer + * + * @param buf 输入缓冲区 + * @param len 缓冲区长度 + * @param value 输出值 + * @return 解码后的字节数,< 0 表示错误 + */ +int http3_decode_varint(const uint8_t *buf, size_t len, uint64_t *value); + +/** + * @brief 计算帧头编码后的大小 + * + * @param frame_type 帧类型 + * @param length 帧负载长度 + * @return 帧头总字节数 + */ +size_t http3_frame_header_size(uint64_t frame_type, uint64_t length); + +/* ===== QPACK 编码/解码 ===== */ + +/** + * @brief QPACK 编码头部字段 + * + * 使用静态表索引或字面量编码单个头部字段。 + * + * @param name 字段名 + * @param value 字段值 + * @param out 输出缓冲区 + * @param out_cap 输出缓冲区容量 + * @return 编码后字节数,< 0 表示错误 + */ +int qpack_encode_header(const char *name, const char *value, + uint8_t *out, size_t out_cap); + +/** + * @brief QPACK 解码头部字段 + * + * @param in 输入缓冲区 + * @param in_len 输入长度 + * @param out 输出解码结果 + * @param consumed 消耗字节数 + * @return 0 成功,< 0 错误 + */ +int qpack_decode_header(const uint8_t *in, size_t in_len, + qpack_decoded_t *out, size_t *consumed); + +/** + * @brief QPACK 编码完整请求头 + * + * 编码所有请求伪头和常规头。 + * + * @param req HTTP 请求 + * @param out 输出缓冲区 + * @param out_cap 输出缓冲区容量 + * @return 编码后字节数,< 0 错误 + */ +int qpack_encode_request_headers(const http_request_t *req, + uint8_t *out, size_t out_cap); + +/** + * @brief QPACK 解码完整请求头 + * + * @param in 输入缓冲区 + * @param in_len 输入长度 + * @param req 输出 HTTP 请求 + * @return 0 成功,< 0 错误 + */ +int qpack_decode_request_headers(const uint8_t *in, size_t in_len, + http_request_t *req); + +/** + * @brief QPACK 编码响应头 + * + * @param resp HTTP 响应 + * @param out 输出缓冲区 + * @param out_cap 输出缓冲区容量 + * @return 编码后字节数,< 0 错误 + */ +int qpack_encode_response_headers(const http_response_t *resp, + uint8_t *out, size_t out_cap); + +/** + * @brief QPACK 解码响应头 + * + * @param in 输入缓冲区 + * @param in_len 输入长度 + * @param resp 输出 HTTP 响应 + * @return 0 成功,< 0 错误 + */ +int qpack_decode_response_headers(const uint8_t *in, size_t in_len, + http_response_t *resp); + +/* ===== 帧处理辅助函数 ===== */ + +/** + * @brief 编码 HTTP/3 帧头 + * + * 格式:[Type: varint][Length: varint] + * + * @param frame_type 帧类型 + * @param length 负载长度 + * @param buf 输出缓冲区(至少 16 字节) + * @return 编码后帧头字节数 + */ +size_t http3_encode_frame_header(uint64_t frame_type, uint64_t length, + uint8_t *buf); + +/** + * @brief 解码 HTTP/3 帧头 + * + * @param buf 输入缓冲区 + * @param len 缓冲区长度 + * @param frame_type 输出帧类型 + * @param length 输出负载长度 + * @return 解码后帧头字节数,< 0 错误 + */ +int http3_decode_frame_header(const uint8_t *buf, size_t len, + uint64_t *frame_type, uint64_t *length); + +/** + * @brief 编码完整 HTTP/3 帧 + * + * @param frame_type 帧类型 + * @param payload 负载数据 + * @param payload_len 负载长度 + * @param buf 输出缓冲区 + * @param buf_cap 缓冲区容量 + * @return 编码后总字节数,< 0 错误 + */ +int http3_encode_frame(uint64_t frame_type, + const uint8_t *payload, size_t payload_len, + uint8_t *buf, size_t buf_cap); + +/** + * @brief 解析 HTTP/3 帧(从流缓冲区) + * + * 从 QUIC 流接收缓冲区解析完整帧。 + * + * @param stream QUIC 流 + * @param frame_type 输出帧类型 + * @param payload 输出负载指针(指向流缓冲区内部) + * @param payload_len 输出负载长度 + * @return 0 成功,1 数据不足,< 0 错误 + */ +int http3_parse_frame(quic_stream_t *stream, + uint64_t *frame_type, + const uint8_t **payload, + size_t *payload_len); + +/* ===== 全局连接管理 ===== */ + +/** + * @brief 查找现有 QUIC 连接 + * + * 根据连接 ID 在全局链表中查找。 + * + * @param conn_id 连接 ID + * @return 连接指针,未找到返回 NULL + */ +quic_connection_t *quic_find_connection(uint64_t conn_id); + +/** + * @brief 移除并清理超时连接 + * + * 检查全局连接链表,关闭并释放超时的连接。 + * + * @param timeout_ms 超时毫秒数 + */ +void quic_cleanup_timeout_connections(uint64_t timeout_ms); + +/** + * @brief 生成 64-bit 随机连接 ID + * + * @return 随机连接 ID + */ +uint64_t quic_generate_conn_id(void); + +/** + * @brief 获取当前时间戳(毫秒) + * + * @return 毫秒时间戳 + */ +uint64_t quic_current_time_ms(void); + +/** + * @brief 发送 QUIC 数据报 + * + * 通过 UDP socket 发送数据报给对端。 + * + * @param conn QUIC 连接 + * @param data 数据 + * @param len 数据长度 + * @return 0 成功,< 0 错误 + */ +int quic_send_datagram(quic_connection_t *conn, const uint8_t *data, size_t len); + +/* ===== 连接计数 ===== */ + +/** + * @brief 获取活跃 QUIC 连接数 + * + * @return 活跃连接数 + */ +size_t quic_get_connection_count(void); + +#ifdef __cplusplus +} +#endif + +#endif /* COCOON_HTTP3_H */ diff --git a/load_balance.c b/load_balance.c new file mode 100644 index 0000000..aa41813 --- /dev/null +++ b/load_balance.c @@ -0,0 +1,501 @@ +/** + * @file load_balance.c - 分布式负载均衡模块实现 + * @brief 实现多种负载均衡算法:一致性哈希、最少连接、加权响应时间、随机 + * + * Phase 4 第二项:多种负载均衡算法 + * + * @author xfy + */ + +#include "load_balance.h" +#include +#include +#include + +/* ===== 内部辅助函数声明 ===== */ + +/** + * @brief 获取健康的后端数量 + */ +static size_t count_healthy_backends(cocoon_proxy_rule_t *rule); + +/** + * @brief 最少连接算法实现 + */ +static int select_least_connections(cocoon_load_balancer_t *lb, + cocoon_proxy_rule_t *rule); + +/** + * @brief 加权响应时间 (EWMA) 算法实现 + */ +static int select_weighted_response(cocoon_load_balancer_t *lb, + cocoon_proxy_rule_t *rule); + +/** + * @brief 一致性哈希算法实现 + */ +static int select_consistent_hash(cocoon_load_balancer_t *lb, + cocoon_proxy_rule_t *rule, + const char *hash_key); + +/** + * @brief 随机算法实现 + */ +static int select_random(cocoon_proxy_rule_t *rule); + +/** + * @brief 虚拟节点比较函数(用于 qsort) + */ +static int node_compare(const void *a, const void *b); + +/** + * @brief 为单个后端生成虚拟节点 + */ +static void generate_virtual_nodes(cocoon_hash_ring_node_t *nodes, + size_t *count, + size_t backend_idx, + const char *host, + uint16_t port, + size_t replicas); + +/* ===== 公共 API 实现 ===== */ + +void lb_init(cocoon_load_balancer_t *lb, cocoon_lb_algorithm_t algo) { + if (!lb) return; + + memset(lb, 0, sizeof(*lb)); + lb->algorithm = algo; + lb->alpha = 80; /**< 默认 EWMA 平滑因子 0.8(以百分制表示) */ + pthread_mutex_init(&lb->mutex, NULL); + + /* 初始化哈希环 */ + lb->hash_ring.node_count = 0; + lb->hash_ring.initialized = false; + + /* 初始化所有后端统计 */ + for (size_t i = 0; i < COCOON_MAX_PROXY_BACKENDS; i++) { + lb->stats[i].active_connections = 0; + lb->stats[i].total_requests = 0; + lb->stats[i].total_failures = 0; + lb->stats[i].total_response_time_us = 0; + lb->stats[i].last_response_time_us = 0; + lb->stats[i].ewma_response_time_us = 0; + } +} + +void lb_destroy(cocoon_load_balancer_t *lb) { + if (!lb) return; + pthread_mutex_destroy(&lb->mutex); + memset(lb, 0, sizeof(*lb)); +} + +int lb_select_backend(cocoon_load_balancer_t *lb, cocoon_proxy_rule_t *rule, + const char *hash_key) { + if (!lb || !rule || rule->backend_count == 0) { + return -1; + } + + switch (lb->algorithm) { + case COCOON_LB_ROUND_ROBIN: + /** + * 加权轮询使用 proxy.c 中已有的平滑加权轮询算法。 + * 返回 -1 表示调用者应使用 select_backend_sww()。 + */ + return -1; + + case COCOON_LB_LEAST_CONNECTIONS: + return select_least_connections(lb, rule); + + case COCOON_LB_WEIGHTED_RESPONSE: + return select_weighted_response(lb, rule); + + case COCOON_LB_CONSISTENT_HASH: + return select_consistent_hash(lb, rule, hash_key); + + case COCOON_LB_RANDOM: + return select_random(rule); + + default: + return -1; + } +} + +void lb_update_stats_request_start(cocoon_load_balancer_t *lb, size_t backend_idx) { + if (!lb || backend_idx >= COCOON_MAX_PROXY_BACKENDS) return; + + pthread_mutex_lock(&lb->mutex); + lb->stats[backend_idx].active_connections++; + lb->stats[backend_idx].total_requests++; + pthread_mutex_unlock(&lb->mutex); +} + +void lb_update_stats_request_end(cocoon_load_balancer_t *lb, size_t backend_idx, + bool success, uint64_t response_time_us) { + if (!lb || backend_idx >= COCOON_MAX_PROXY_BACKENDS) return; + + pthread_mutex_lock(&lb->mutex); + + cocoon_backend_stats_t *s = &lb->stats[backend_idx]; + + /* 减少活跃连接数 */ + if (s->active_connections > 0) { + s->active_connections--; + } + + s->total_response_time_us += response_time_us; + s->last_response_time_us = response_time_us; + + if (!success) { + s->total_failures++; + } + + /** + * EWMA 更新公式: + * new_ewma = (alpha * current_rtt + (100 - alpha) * old_ewma) / 100 + * + * alpha = 80 表示新值占 80%,旧值占 20%,对近期响应更敏感。 + * 首次更新(old_ewma == 0)时直接使用当前值。 + */ + uint32_t a = lb->alpha; + if (s->ewma_response_time_us == 0) { + s->ewma_response_time_us = response_time_us; + } else { + s->ewma_response_time_us = (a * response_time_us + + (100 - a) * s->ewma_response_time_us) / 100; + } + + pthread_mutex_unlock(&lb->mutex); +} + +void lb_build_hash_ring(cocoon_hash_ring_t *ring, cocoon_proxy_rule_t *rule) { + if (!ring || !rule) return; + + ring->node_count = 0; + ring->initialized = false; + + if (rule->backend_count == 0) return; + + /** + * 为每个后端生成 COCOON_HASH_RING_SIZE 个虚拟节点。 + * 虚拟节点数量 = 后端数 × 每个后端的副本数 + */ + for (size_t i = 0; i < rule->backend_count; i++) { + cocoon_proxy_backend_t *be = &rule->backends[i]; + + /* 跳过重和不健康的后端 */ + if (!be->healthy) continue; + + generate_virtual_nodes(ring->nodes, &ring->node_count, + i, be->target_host, be->target_port, + COCOON_HASH_RING_SIZE); + } + + if (ring->node_count == 0) return; + + /** + * 按哈希值升序排序,二分查找依赖有序数组。 + */ + qsort(ring->nodes, ring->node_count, + sizeof(cocoon_hash_ring_node_t), node_compare); + + ring->initialized = true; +} + +uint32_t lb_hash_key(const char *key, size_t len) { + if (!key || len == 0) return 0; + + /** + * MurmurHash3 x86 32-bit 实现 + * + * 参考原始 MurmurHash3 算法参数: + * c1 = 0xcc9e2d51 (第一次混合常数) + * c2 = 0x1b873593 (第二次混合常数) + * r1 = 15 (第一次右移位数) + * r2 = 13 (第二次右移位数) + * m = 5 (累乘常数) + * n = 0xe6546b64 (累加常数) + * + * 此实现为纯软件、零外部依赖版本。 + */ + const uint32_t c1 = 0xcc9e2d51; + const uint32_t c2 = 0x1b873593; + const uint32_t r1 = 15; + const uint32_t r2 = 13; + const uint32_t m = 5; + const uint32_t n = 0xe6546b64; + + uint32_t hash = 0; /**< seed = 0,保证相同输入总有相同输出 */ + const size_t nblocks = len / 4; + const uint32_t *blocks = (const uint32_t *)(const void *)key; + + /* === 主体:4-byte 块处理 === */ + for (size_t i = 0; i < nblocks; i++) { + uint32_t k = blocks[i]; + + /* 第一次混合 */ + k *= c1; + k = (k << r1) | (k >> (32 - r1)); /**< ROTL32(k, r1) */ + k *= c2; + + /* 哈希混合 */ + hash ^= k; + hash = ((hash << r2) | (hash >> (32 - r2))); /**< ROTL32(hash, r2) */ + hash = hash * m + n; + } + + /* === 尾部:不足 4 字节的部分 === */ + const uint8_t *tail = (const uint8_t *)(key + nblocks * 4); + uint32_t k1 = 0; + + switch (len & 3) { /**< len % 4 */ + case 3: + k1 ^= (uint32_t)tail[2] << 16; + /* fall through */ + case 2: + k1 ^= (uint32_t)tail[1] << 8; + /* fall through */ + case 1: + k1 ^= (uint32_t)tail[0]; + k1 *= c1; + k1 = (k1 << r1) | (k1 >> (32 - r1)); + k1 *= c2; + hash ^= k1; + break; + default: + break; + } + + /* === 最终化(finalization) === */ + hash ^= (uint32_t)len; + hash ^= hash >> 16; + hash *= 0x85ebca6b; + hash ^= hash >> 13; + hash *= 0xc2b2ae35; + hash ^= hash >> 16; + + return hash; +} + +size_t lb_pick_from_ring(cocoon_hash_ring_t *ring, uint32_t hash) { + if (!ring || ring->node_count == 0) return 0; + + /** + * 二分查找:定位第一个 node_hash >= hash 的虚拟节点 + * + * 如果 hash 大于环中所有节点的哈希值,环绕到第一个节点 + *(取模运算自然处理此情况)。 + */ + size_t lo = 0; + size_t hi = ring->node_count; + + while (lo < hi) { + size_t mid = (lo + hi) / 2; + if (ring->nodes[mid].node_hash < hash) { + lo = mid + 1; + } else { + hi = mid; + } + } + + return ring->nodes[lo % ring->node_count].backend_index; +} + +const char *lb_algorithm_name(cocoon_lb_algorithm_t algo) { + switch (algo) { + case COCOON_LB_ROUND_ROBIN: return "round_robin"; + case COCOON_LB_LEAST_CONNECTIONS: return "least_connections"; + case COCOON_LB_WEIGHTED_RESPONSE: return "weighted_response"; + case COCOON_LB_CONSISTENT_HASH: return "consistent_hash"; + case COCOON_LB_RANDOM: return "random"; + default: return "unknown"; + } +} + +/* ===== 内部辅助函数实现 ===== */ + +/** + * @brief 统计健康后端数量 + */ +static size_t count_healthy_backends(cocoon_proxy_rule_t *rule) { + size_t count = 0; + for (size_t i = 0; i < rule->backend_count; i++) { + if (rule->backends[i].healthy) { + count++; + } + } + return count; +} + +/** + * @brief 最少连接算法 + * + * 遍历所有健康后端,选择 active_connections 最小者。 + * 如果多个后端连接数相同,选择第一个(保持确定性)。 + */ +static int select_least_connections(cocoon_load_balancer_t *lb, + cocoon_proxy_rule_t *rule) { + int best_idx = -1; + uint32_t min_conns = UINT32_MAX; + + for (size_t i = 0; i < rule->backend_count; i++) { + if (!rule->backends[i].healthy) continue; + + uint32_t conns = lb->stats[i].active_connections; + if (conns < min_conns) { + min_conns = conns; + best_idx = (int)i; + } + } + + return best_idx; +} + +/** + * @brief 加权响应时间算法 (EWMA) + * + * 选择 ewma_response_time_us / weight 比值最小的健康后端。 + * 后端的 weight 越高,可以承受更高的 EWMA 仍被选中。 + * 如果后端没有 EWMA 记录(首次请求),视为 0 优先选择。 + */ +static int select_weighted_response(cocoon_load_balancer_t *lb, + cocoon_proxy_rule_t *rule) { + int best_idx = -1; + double best_score = 1e308; /**< 初始化为极大值 */ + + for (size_t i = 0; i < rule->backend_count; i++) { + if (!rule->backends[i].healthy) continue; + + uint32_t weight = rule->backends[i].weight; + if (weight == 0) weight = 1; /**< 防止除零 */ + + uint64_t ewma = lb->stats[i].ewma_response_time_us; + + /** + * 评分 = EWMA / weight + * weight 越大,相同 EWMA 下评分越低(越优先)。 + */ + double score = (double)ewma / (double)weight; + + if (score < best_score) { + best_score = score; + best_idx = (int)i; + } + } + + return best_idx; +} + +/** + * @brief 一致性哈希算法 + * + * 如果提供了 hash_key,使用 MurmurHash3 哈希后从环中选取; + * 如果 hash_key 为 NULL 或哈希环未初始化,回退到随机选择。 + */ +static int select_consistent_hash(cocoon_load_balancer_t *lb, + cocoon_proxy_rule_t *rule, + const char *hash_key) { + /** + * 如果哈希环未初始化,尝试构建。 + */ + if (!lb->hash_ring.initialized) { + lb_build_hash_ring(&lb->hash_ring, rule); + } + + if (!lb->hash_ring.initialized || lb->hash_ring.node_count == 0) { + return select_random(rule); + } + + /** + * 如果没有提供 hash_key,回退到随机选择。 + * 这样保证在无 key 场景下仍有可用后端。 + */ + if (!hash_key || hash_key[0] == '\0') { + return select_random(rule); + } + + uint32_t h = lb_hash_key(hash_key, strlen(hash_key)); + size_t idx = lb_pick_from_ring(&lb->hash_ring, h); + + /** + * 确保选中的后端是健康的。如果该后端不健康, + * 回退到随机选择(简单处理,生产环境可改为查找下一个节点)。 + */ + if (idx < rule->backend_count && rule->backends[idx].healthy) { + return (int)idx; + } + + return select_random(rule); +} + +/** + * @brief 随机算法 + * + * 在健康后端中均匀随机选择一个。 + * 使用标准库的 rand(),调用者应事先 srand()。 + */ +static int select_random(cocoon_proxy_rule_t *rule) { + size_t healthy_count = count_healthy_backends(rule); + if (healthy_count == 0) return -1; + + /** + * 生成 [0, healthy_count) 范围内的随机数, + * 然后映射到第 n 个健康后端。 + */ + size_t pick = (size_t)rand() % healthy_count; + + size_t count = 0; + for (size_t i = 0; i < rule->backend_count; i++) { + if (rule->backends[i].healthy) { + if (count == pick) { + return (int)i; + } + count++; + } + } + + return -1; /**< 不应到达此处 */ +} + +/** + * @brief 虚拟节点比较函数(qsort 回调) + * + * 按 node_hash 升序排列。 + */ +static int node_compare(const void *a, const void *b) { + const cocoon_hash_ring_node_t *na = (const cocoon_hash_ring_node_t *)a; + const cocoon_hash_ring_node_t *nb = (const cocoon_hash_ring_node_t *)b; + + if (na->node_hash < nb->node_hash) return -1; + if (na->node_hash > nb->node_hash) return 1; + return 0; +} + +/** + * @brief 为单个后端生成虚拟节点 + * + * 每个虚拟节点的键格式为 ":-", + * 通过 MurmurHash3 生成唯一哈希值。 + */ +static void generate_virtual_nodes(cocoon_hash_ring_node_t *nodes, + size_t *count, + size_t backend_idx, + const char *host, + uint16_t port, + size_t replicas) { + for (size_t i = 0; i < replicas; i++) { + if (*count >= COCOON_HASH_RING_MAX_NODES) break; + + /** + * 虚拟节点键:"hostname:port-replica_index" + * 格式保证不同后端、不同副本产生不同的哈希值。 + */ + char key[512]; + int n = snprintf(key, sizeof(key), "%s:%u-%zu", + host, (unsigned int)port, i); + if (n < 0 || (size_t)n >= sizeof(key)) continue; + + nodes[*count].node_hash = lb_hash_key(key, (size_t)n); + nodes[*count].backend_index = backend_idx; + (*count)++; + } +} diff --git a/load_balance.h b/load_balance.h new file mode 100644 index 0000000..96f076c --- /dev/null +++ b/load_balance.h @@ -0,0 +1,212 @@ +/** + * @file load_balance.h - 分布式负载均衡模块头文件 + * @brief 支持多种负载均衡算法:一致性哈希、最少连接、加权响应时间、随机 + * + * Phase 4 第二项:多种负载均衡算法 + * 设计为独立模块,可集成到 proxy.c 替换现有轮询逻辑。 + * + * @author xfy + */ + +#ifndef COCOON_LOAD_BALANCE_H +#define COCOON_LOAD_BALANCE_H + +#include "proxy.h" +#include "healthcheck.h" +#include +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +/* ===== 负载均衡算法枚举 ===== */ +/** + * @enum cocoon_lb_algorithm_t + * @brief 支持的负载均衡算法 + */ +typedef enum { + COCOON_LB_ROUND_ROBIN, /**< 加权轮询(已有平滑加权轮询在 proxy.c 中) */ + COCOON_LB_LEAST_CONNECTIONS, /**< 最少连接 */ + COCOON_LB_WEIGHTED_RESPONSE, /**< 加权响应时间 (EWMA) */ + COCOON_LB_CONSISTENT_HASH, /**< 一致性哈希 */ + COCOON_LB_RANDOM /**< 随机 */ +} cocoon_lb_algorithm_t; + +/* ===== 后端状态跟踪 ===== */ +/** + * @struct cocoon_backend_stats_t + * @brief 后端实时统计信息 + * + * 记录每个后端的连接数、请求数、响应时间等关键指标, + * 为负载均衡算法提供决策依据。 + */ +typedef struct { + uint32_t active_connections; /**< 当前活跃连接数 */ + uint32_t total_requests; /**< 总请求数 */ + uint32_t total_failures; /**< 总失败数 */ + uint64_t total_response_time_us; /**< 总响应时间(微秒) */ + uint64_t last_response_time_us; /**< 上次响应时间 */ + uint64_t ewma_response_time_us; /**< 指数加权移动平均响应时间 */ +} cocoon_backend_stats_t; + +/* ===== 一致性哈希环 ===== */ +/** + * @def COCOON_HASH_RING_SIZE + * @brief 每个后端生成的虚拟节点数量 + */ +#define COCOON_HASH_RING_SIZE 512 + +/** + * @def COCOON_HASH_RING_MAX_NODES + * @brief 哈希环最大节点总数 + */ +#define COCOON_HASH_RING_MAX_NODES (COCOON_HASH_RING_SIZE * COCOON_MAX_PROXY_BACKENDS) + +/** + * @struct cocoon_hash_ring_node_t + * @brief 哈希环上的单个虚拟节点 + */ +typedef struct { + uint32_t node_hash; /**< 虚拟节点哈希值 */ + size_t backend_index; /**< 指向的后端索引 */ +} cocoon_hash_ring_node_t; + +/** + * @struct cocoon_hash_ring_t + * @brief 一致性哈希环 + * + * 虚拟节点按哈希值升序排列,支持二分查找。 + */ +typedef struct { + cocoon_hash_ring_node_t nodes[COCOON_HASH_RING_MAX_NODES]; /**< 排序后的虚拟节点数组 */ + size_t node_count; /**< 实际节点数量 */ + bool initialized; /**< 是否已初始化 */ +} cocoon_hash_ring_t; + +/* ===== 负载均衡器 ===== */ +/** + * @struct cocoon_load_balancer_t + * @brief 负载均衡器主结构体 + * + * 包含算法选择、后端统计、哈希环和线程安全锁。 + */ +typedef struct { + cocoon_lb_algorithm_t algorithm; /**< 当前使用的算法 */ + cocoon_backend_stats_t stats[COCOON_MAX_PROXY_BACKENDS]; /**< 各后端统计 */ + cocoon_hash_ring_t hash_ring; /**< 一致性哈希环 */ + uint32_t alpha; /**< EWMA 平滑因子(默认 80,即 0.8) */ + pthread_mutex_t mutex; /**< 统计更新锁(线程安全) */ +} cocoon_load_balancer_t; + +/* ===== API ===== */ + +/** + * @brief 初始化负载均衡器 + * + * 设置算法类型,初始化统计数组和互斥锁。 + * 如果算法为 CONSISTENT_HASH,不会自动构建哈希环, + * 需要显式调用 lb_build_hash_ring()。 + * + * @param lb 负载均衡器指针 + * @param algo 负载均衡算法 + */ +void lb_init(cocoon_load_balancer_t *lb, cocoon_lb_algorithm_t algo); + +/** + * @brief 销毁负载均衡器 + * + * 销毁互斥锁,清空所有状态。 + * + * @param lb 负载均衡器指针 + */ +void lb_destroy(cocoon_load_balancer_t *lb); + +/** + * @brief 根据算法选择后端索引 + * + * 根据配置的负载均衡算法,从规则的健康后端中选择一个。 + * 对于 CONSISTENT_HASH 算法,需要提供 hash_key; + * 如果 hash_key 为 NULL,则回退到随机选择。 + * + * @param lb 负载均衡器 + * @param rule 代理规则(包含后端列表) + * @param hash_key 一致性哈希的键(如请求路径或客户端 IP),可为 NULL + * @return 选择的后端索引,无可用后端返回 -1 + */ +int lb_select_backend(cocoon_load_balancer_t *lb, cocoon_proxy_rule_t *rule, + const char *hash_key); + +/** + * @brief 请求开始时更新统计 + * + * 增加指定后端的活跃连接数和总请求数。 + * 线程安全:内部加锁。 + * + * @param lb 负载均衡器 + * @param backend_idx 后端索引 + */ +void lb_update_stats_request_start(cocoon_load_balancer_t *lb, size_t backend_idx); + +/** + * @brief 请求结束时更新统计 + * + * 减少活跃连接数,更新 EWMA 响应时间,记录成功/失败。 + * 线程安全:内部加锁。 + * + * @param lb 负载均衡器 + * @param backend_idx 后端索引 + * @param success 请求是否成功 + * @param response_time_us 响应时间(微秒) + */ +void lb_update_stats_request_end(cocoon_load_balancer_t *lb, size_t backend_idx, + bool success, uint64_t response_time_us); + +/** + * @brief 构建一致性哈希环 + * + * 为规则中的所有健康后端生成虚拟节点并排序。 + * 应在后端配置变更或健康状态变化时重新构建。 + * + * @param ring 哈希环 + * @param rule 代理规则 + */ +void lb_build_hash_ring(cocoon_hash_ring_t *ring, cocoon_proxy_rule_t *rule); + +/** + * @brief MurmurHash3 x86 32-bit 哈希函数 + * + * 纯软件实现,不依赖外部库。 + * 用于一致性哈希的键哈希。 + * + * @param key 输入键 + * @param len 键长度 + * @return 32-bit 哈希值 + */ +uint32_t lb_hash_key(const char *key, size_t len); + +/** + * @brief 从哈希环中选取后端 + * + * 使用二分查找定位最近的虚拟节点,支持哈希环环绕。 + * + * @param ring 已构建的哈希环 + * @param hash 要查找的哈希值 + * @return 选中的后端索引 + */ +size_t lb_pick_from_ring(cocoon_hash_ring_t *ring, uint32_t hash); + +/** + * @brief 获取算法名称字符串 + * + * @param algo 负载均衡算法 + * @return 算法名称(静态字符串) + */ +const char *lb_algorithm_name(cocoon_lb_algorithm_t algo); + +#ifdef __cplusplus +} +#endif + +#endif /* COCOON_LOAD_BALANCE_H */ diff --git a/middleware_ext.c b/middleware_ext.c new file mode 100644 index 0000000..ed86199 --- /dev/null +++ b/middleware_ext.c @@ -0,0 +1,662 @@ +/** + * @file middleware_ext.c - 扩展内置中间件实现 + * + * Phase 4 第一项:扩展中间件系统 + * - JWT 认证(HS256, Base64Url, exp 验证) + * - Security Headers(全局配置) + * - Request ID(32 字符 hex UUID) + * - IP 过滤(IPv4 CIDR 黑白名单) + * + * @author xfy + */ + +#include "middleware_ext.h" +#include "log.h" +#include +#include +#include +#include +#include + +#include +#include + +/* ============================================================ + * 内部工具函数 + * ============================================================ */ + +/** + * @brief Base64Url 解码表 + * + * 标准 Base64: A-Z a-z 0-9 + / = + * Base64Url: A-Z a-z 0-9 - _ (无填充 =) + */ +static const char base64url_chars[] = + "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_"; + +/** + * @brief Base64Url 解码 + * + * 将 Base64Url 编码字符串解码为二进制数据。 + * Base64Url 变体使用 '-' 代替 '+','_' 代替 '/',省略 '=' 填充。 + * + * @param in 输入 Base64Url 字符串 + * @param out 输出缓冲区 + * @param out_size 输出缓冲区大小 + * @return 解码后的字节数,< 0 表示错误 + */ +/* 测试可见的内部函数(非 static,通过前置声明在测试中访问) */ +int base64url_decode(const char *in, unsigned char *out, int out_size) { + if (!in || !out || out_size <= 0) return -1; + + int val = 0, valb = -8; + int out_len = 0; + for (const char *p = in; *p; p++) { + const char *pos = strchr(base64url_chars, *p); + if (!pos) { + /* Base64Url 不应有填充字符,但兼容处理 */ + if (*p == '=') break; + /* 非法字符 */ + return -1; + } + int c = (int)(pos - base64url_chars); + val = (val << 6) + c; + valb += 6; + if (valb >= 0) { + if (out_len < out_size) { + out[out_len++] = (unsigned char)((val >> valb) & 0xFF); + } + valb -= 8; + } + } + return out_len; +} + +/** + * @brief Base64Url 编码 + * + * 将二进制数据编码为 Base64Url 字符串(无填充)。 + * + * @param in 输入数据 + * @param in_len 输入长度 + * @param out 输出缓冲区 + * @param out_size 输出缓冲区大小 + * @return 编码后的字符串长度,< 0 表示缓冲区不足 + */ +int base64url_encode(const unsigned char *in, int in_len, char *out, int out_size) { + if (!in || !out || out_size <= 0 || in_len < 0) return -1; + + int i = 0, j = 0; + unsigned char a, b, c; + int val; + + while (i < in_len) { + a = i < in_len ? in[i] : 0; + b = i + 1 < in_len ? in[i + 1] : 0; + c = i + 2 < in_len ? in[i + 2] : 0; + + val = (a << 16) | (b << 8) | c; + + if (j >= out_size - 1) return -1; + out[j++] = base64url_chars[(val >> 18) & 0x3F]; + if (j >= out_size - 1) return -1; + out[j++] = base64url_chars[(val >> 12) & 0x3F]; + if (i + 1 < in_len) { + if (j >= out_size - 1) return -1; + out[j++] = base64url_chars[(val >> 6) & 0x3F]; + } + if (i + 2 < in_len) { + if (j >= out_size - 1) return -1; + out[j++] = base64url_chars[val & 0x3F]; + } + i += 3; + } + out[j] = '\0'; + return j; +} + +/** + * @brief 从请求头中查找指定名称的值 + * + * @param req HTTP 请求 + * @param name 头名称(不区分大小写) + * @return 头值指针,未找到返回 NULL + */ +const char *find_header(const http_request_t *req, const char *name) { + for (int i = 0; i < req->num_headers; i++) { + if (strcasecmp(req->headers[i].name, name) == 0) { + return req->headers[i].value; + } + } + return NULL; +} + +/** + * @brief 发送 JSON 格式的错误响应 + * + * @param fd 客户端 socket + * @param status HTTP 状态码 + * @param body 响应体 + * @param keep_alive 是否保持连接 + */ +void send_json_error(cocoon_socket_t fd, int status, const char *body, bool keep_alive) { + char response[1024]; + int n = snprintf(response, sizeof(response), + "HTTP/1.1 %d %s\r\n" + "Content-Type: application/json\r\n" + "Content-Length: %zu\r\n" + "Connection: %s\r\n" + "Server: Cocoon/1.0\r\n" + "\r\n" + "%s", + status, + status == 401 ? "Unauthorized" : (status == 403 ? "Forbidden" : "Error"), + strlen(body), + keep_alive ? "keep-alive" : "close", + body); + cocoon_socket_send(fd, response, (size_t)n); +} + +/* ============================================================ + * JWT 中间件 + * ============================================================ */ + +/** + * @brief 从 JWT payload JSON 中解析 exp 字段 + * + * 简单字符串搜索,不引入完整 JSON 解析器。 + * 查找 "exp" 后面的数字值。 + * + * @param payload 解码后的 payload 字符串 + * @return exp 时间戳,未找到返回 0 + */ +time_t jwt_parse_exp(const char *payload) { + if (!payload) return 0; + + const char *p = payload; + while (*p) { + /* 寻找 "exp" 字段 */ + const char *exp_key = strstr(p, "\"exp\""); + if (!exp_key) break; + + /* 查找冒号后的值 */ + const char *val = exp_key + 5; + while (*val && (*val == ':' || *val == ' ' || *val == '\t')) val++; + + if (*val) { + char *endptr = NULL; + long exp_val = strtol(val, &endptr, 10); + if (endptr != val) { + return (time_t)exp_val; + } + } + p = exp_key + 5; + } + return 0; +} + +/** + * @brief 验证 JWT token 的签名 + * + * 使用 HMAC-SHA256 重新计算签名并与 token 中的签名比较。 + * + * @param header_payload header.payload 部分(Base64Url 编码) + * @param signature_b64 signature 部分(Base64Url 编码) + * @param secret 密钥 + * @param secret_len 密钥长度 + * @return true 签名验证通过,false 失败 + */ +bool jwt_verify_signature(const char *header_payload, + const char *signature_b64, + const char *secret, + size_t secret_len) { + unsigned char computed_sig[EVP_MAX_MD_SIZE]; + unsigned int computed_len = 0; + + /* 使用 HMAC-SHA256 计算签名 */ + if (!HMAC(EVP_sha256(), + secret, (int)secret_len, + (const unsigned char *)header_payload, strlen(header_payload), + computed_sig, &computed_len)) { + return false; + } + + /* 解码 token 中的签名 */ + unsigned char token_sig[EVP_MAX_MD_SIZE]; + int token_sig_len = base64url_decode(signature_b64, token_sig, sizeof(token_sig)); + if (token_sig_len < 0) return false; + + /* 比较签名(常量时间比较防止时序攻击) */ + if ((size_t)token_sig_len != computed_len) return false; + + unsigned char diff = 0; + for (size_t i = 0; i < computed_len; i++) { + diff |= (computed_sig[i] ^ token_sig[i]); + } + return diff == 0; +} + +int cocoon_middleware_jwt(http_request_t *req, cocoon_socket_t fd, void *user_data) { + const cocoon_jwt_config_t *cfg = (const cocoon_jwt_config_t *)user_data; + if (!cfg || cfg->secret[0] == '\0') return 0; /* 未配置,跳过 */ + + /* 跳过 OPTIONS 预检请求 */ + if (cfg->skip_preflight && req->method == HTTP_OPTIONS) { + return 0; + } + + /* 从请求头获取 Authorization */ + const char *auth_header = find_header(req, cfg->header_name[0] ? cfg->header_name : "Authorization"); + if (!auth_header) { + log_warn("JWT 认证失败: 缺少 Authorization 头"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + + /* 检查前缀(默认 "Bearer ") */ + const char *prefix = cfg->prefix[0] ? cfg->prefix : "Bearer "; + size_t prefix_len = strlen(prefix); + if (strncasecmp(auth_header, prefix, prefix_len) != 0) { + log_warn("JWT 认证失败: 前缀不匹配"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + + /* 提取 token */ + const char *token = auth_header + prefix_len; + if (strlen(token) == 0 || strlen(token) >= COCOON_JWT_TOKEN_MAX) { + log_warn("JWT 认证失败: Token 长度异常"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + + /* 复制 token 到可修改缓冲区 */ + char token_buf[COCOON_JWT_TOKEN_MAX]; + strncpy(token_buf, token, sizeof(token_buf) - 1); + token_buf[sizeof(token_buf) - 1] = '\0'; + + /* 解析三段式:header.payload.signature */ + char *first_dot = strchr(token_buf, '.'); + if (!first_dot) { + log_warn("JWT 认证失败: Token 格式错误(缺少第一段分隔符)"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + *first_dot = '\0'; + + char *second_dot = strchr(first_dot + 1, '.'); + if (!second_dot) { + log_warn("JWT 认证失败: Token 格式错误(缺少第二段分隔符)"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + *second_dot = '\0'; + + const char *header_b64 = token_buf; + const char *payload_b64 = first_dot + 1; + const char *signature_b64 = second_dot + 1; + + /* 签名验证 */ + /* header_payload = "header_b64.payload_b64" */ + char header_payload[COCOON_JWT_TOKEN_MAX * 2]; + int n = snprintf(header_payload, sizeof(header_payload), "%s.%s", + header_b64, payload_b64); + if (n < 0 || (size_t)n >= sizeof(header_payload)) { + log_warn("JWT 认证失败: header.payload 组合过长"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + + if (!jwt_verify_signature(header_payload, signature_b64, + cfg->secret, strlen(cfg->secret))) { + log_warn("JWT 认证失败: 签名验证未通过"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + + /* 解码 payload 验证 exp */ + unsigned char payload_decoded[COCOON_JWT_TOKEN_MAX]; + int payload_len = base64url_decode(payload_b64, payload_decoded, sizeof(payload_decoded) - 1); + if (payload_len < 0) { + log_warn("JWT 认证失败: payload 解码失败"); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + payload_decoded[payload_len] = '\0'; + + /* 验证 exp 声明 */ + time_t exp = jwt_parse_exp((const char *)payload_decoded); + if (exp > 0) { + time_t now = time(NULL); + if (now > exp) { + log_warn("JWT 认证失败: Token 已过期 (exp=%ld, now=%ld)", (long)exp, (long)now); + send_json_error(fd, 401, "{\"error\": \"Unauthorized\"}", req->keep_alive); + return 1; + } + } + + log_debug("JWT 认证成功"); + return 0; /* 验证通过,继续处理 */ +} + +/* ============================================================ + * Security Headers 中间件 + * ============================================================ */ + +/** 全局 Security Headers 配置 */ +static cocoon_security_headers_config_t g_security_headers_cfg = {0}; +static bool g_security_headers_initialized = false; + +int cocoon_middleware_security_headers(http_request_t *req, cocoon_socket_t fd, void *user_data) { + (void)req; + (void)fd; + + const cocoon_security_headers_config_t *cfg = + (const cocoon_security_headers_config_t *)user_data; + if (!cfg) return 0; + + /* 复制配置到全局变量 */ + g_security_headers_cfg = *cfg; + g_security_headers_initialized = true; + + return 0; /* 始终继续处理 */ +} + +const cocoon_security_headers_config_t *cocoon_middleware_security_headers_get(void) { + if (!g_security_headers_initialized) return NULL; + return &g_security_headers_cfg; +} + +/* ============================================================ + * Request ID 中间件 + * ============================================================ */ + +/** 请求 ID 生成计数器,增加随机性 */ +static uint32_t g_request_id_counter = 0; + +/** + * @brief 生成 32 字符 hex 请求 ID + * + * 使用 4 个 rand() 输出组合成 16 字节(32 hex 字符)的 ID。 + * 格式:%08x%08x%08x%08x + * + * @param buf 输出缓冲区(至少 33 字节) + * @param buf_size 缓冲区大小 + */ +static void generate_request_id(char *buf, size_t buf_size) { + if (buf_size < COCOON_REQUEST_ID_LEN + 1) return; + + static int seeded = 0; + if (!seeded) { + srand((unsigned int)time(NULL) ^ (unsigned int)g_request_id_counter); + seeded = 1; + } + + g_request_id_counter++; + + uint32_t r1 = (uint32_t)rand() ^ g_request_id_counter; + uint32_t r2 = (uint32_t)rand() ^ (g_request_id_counter << 7); + uint32_t r3 = (uint32_t)rand() ^ (g_request_id_counter << 13); + uint32_t r4 = (uint32_t)rand() ^ (g_request_id_counter << 19); + + snprintf(buf, buf_size, "%08x%08x%08x%08x", + (unsigned int)r1, (unsigned int)r2, + (unsigned int)r3, (unsigned int)r4); +} + +int cocoon_middleware_request_id(http_request_t *req, cocoon_socket_t fd, void *user_data) { + (void)fd; + + const cocoon_request_id_config_t *cfg = + (const cocoon_request_id_config_t *)user_data; + if (!cfg) return 0; + + const char *header_name = cfg->header_name[0] ? cfg->header_name : "X-Request-ID"; + char request_id[COCOON_REQUEST_ID_LEN + 1]; + + /* 如果 trust_incoming 为 true 且请求头已有 ID,则复用 */ + if (cfg->trust_incoming) { + const char *incoming = find_header(req, header_name); + if (incoming && strlen(incoming) == COCOON_REQUEST_ID_LEN) { + /* 验证是否全为 hex 字符 */ + bool valid = true; + for (size_t i = 0; i < COCOON_REQUEST_ID_LEN; i++) { + char c = incoming[i]; + if (!((c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F'))) { + valid = false; + break; + } + } + if (valid) { + strncpy(request_id, incoming, COCOON_REQUEST_ID_LEN); + request_id[COCOON_REQUEST_ID_LEN] = '\0'; + log_debug("Request ID 复用客户端传入: %s", request_id); + /* TODO: 传递给响应头的机制(通过 server.c 集成) */ + return 0; + } + } + } + + /* 生成新的请求 ID */ + generate_request_id(request_id, sizeof(request_id)); + log_debug("Request ID 新生成: %s", request_id); + + /* TODO: 将 request_id 添加到响应头需要 server.c 集成 */ + (void)header_name; + + return 0; /* 继续处理 */ +} + +/* ============================================================ + * IP 过滤中间件 + * ============================================================ */ + +/** + * @brief 解析 IPv4 地址字符串为 32 位整数 + * + * @param ip_str IP 地址字符串 + * @param addr 输出 32 位网络字节序地址 + * @return true 解析成功 + */ +bool parse_ipv4(const char *ip_str, uint32_t *addr) { + struct in_addr sin_addr; + if (inet_pton(AF_INET, ip_str, &sin_addr) != 1) return false; + *addr = ntohl(sin_addr.s_addr); + return true; +} + +/** + * @brief 解析 CIDR 字符串为地址和掩码 + * + * 支持格式: + * - "192.168.1.1" -> addr, mask=32 + * - "192.168.1.0/24" -> addr, mask=24 + * + * @param cidr_str CIDR 字符串 + * @param addr 输出网络地址(主机字节序) + * @param mask 输出掩码位数(0-32) + * @return true 解析成功 + */ +bool parse_cidr(const char *cidr_str, uint32_t *addr, int *mask) { + if (!cidr_str || !addr || !mask) return false; + + char buf[64]; + strncpy(buf, cidr_str, sizeof(buf) - 1); + buf[sizeof(buf) - 1] = '\0'; + + char *slash = strchr(buf, '/'); + if (slash) { + *slash = '\0'; + *mask = atoi(slash + 1); + if (*mask < 0 || *mask > 32) return false; + } else { + *mask = 32; /* 精确匹配 */ + } + + if (!parse_ipv4(buf, addr)) return false; + + /* 将地址按掩码对齐到网络地址 */ + if (*mask < 32) { + *addr &= ~(uint32_t)((1U << (32 - *mask)) - 1); + } + return true; +} + +/** + * @brief 检查 IP 是否匹配 CIDR 条目 + * + * @param ip_addr 要检查的 IP(主机字节序) + * @param cidr_str CIDR 字符串 + * @return true 匹配 + */ +bool ip_match_cidr(uint32_t ip_addr, const char *cidr_str) { + uint32_t net_addr; + int mask; + if (!parse_cidr(cidr_str, &net_addr, &mask)) return false; + + if (mask == 32) { + return ip_addr == net_addr; + } + + uint32_t mask_bits = (mask == 0) ? 0 : ~(uint32_t)((1U << (32 - mask)) - 1); + return (ip_addr & mask_bits) == (net_addr & mask_bits); +} + +/** + * @brief 从 socket 获取客户端 IP(网络字节序转主机字节序) + * + * @param fd 客户端 socket + * @param addr 输出 IP 地址(主机字节序) + * @return true 成功 + */ +bool get_client_ip(cocoon_socket_t fd, uint32_t *addr) { + struct sockaddr_storage ss; + socklen_t len = sizeof(ss); + if (getpeername(fd, (struct sockaddr *)&ss, &len) != 0) return false; + + if (ss.ss_family == AF_INET) { + struct sockaddr_in *sin = (struct sockaddr_in *)&ss; + *addr = ntohl(sin->sin_addr.s_addr); + return true; + } + /* IPv6 不匹配任何 IPv4 规则 */ + return false; +} + +/** + * @brief 从 X-Forwarded-For 头解析第一个 IP + * + * X-Forwarded-For 格式:client, proxy1, proxy2, ... + * 取第一个(最左边的)即客户端真实 IP。 + * + * @param header_value X-Forwarded-For 头值 + * @param addr 输出 IP 地址(主机字节序) + * @return true 成功解析到有效 IP + */ +bool parse_x_forwarded_for(const char *header_value, uint32_t *addr) { + if (!header_value || !addr) return false; + + /* 复制并取第一个 IP(逗号前的部分) */ + char buf[64]; + strncpy(buf, header_value, sizeof(buf) - 1); + buf[sizeof(buf) - 1] = '\0'; + + char *comma = strchr(buf, ','); + if (comma) *comma = '\0'; + + /* 去除前后空格 */ + char *start = buf; + while (*start == ' ') start++; + char *end = start + strlen(start) - 1; + while (end > start && *end == ' ') *end-- = '\0'; + + return parse_ipv4(start, addr); +} + +int cocoon_middleware_ip_filter(http_request_t *req, cocoon_socket_t fd, void *user_data) { + const cocoon_ip_filter_config_t *cfg = (const cocoon_ip_filter_config_t *)user_data; + if (!cfg || cfg->count == 0) return 0; /* 未配置,跳过 */ + + uint32_t client_ip = 0; + bool have_ip = false; + + /* 优先使用 X-Forwarded-For 获取真实 IP */ + const char *xff = find_header(req, "X-Forwarded-For"); + if (xff) { + have_ip = parse_x_forwarded_for(xff, &client_ip); + if (have_ip) { + log_debug("IP 过滤: 使用 X-Forwarded-For IP %s", xff); + } + } + + /* 没有 X-Forwarded-For 或解析失败,使用直接连接 IP */ + if (!have_ip) { + have_ip = get_client_ip(fd, &client_ip); + } + + if (!have_ip) { + /* 无法获取 IP,白名单模式下拒绝,黑名单模式下允许 */ + if (cfg->mode == COCOON_IP_FILTER_ALLOW) { + log_warn("IP 过滤: 无法获取客户端 IP,白名单模式拒绝"); + send_json_error(fd, 403, + cfg->deny_message[0] ? cfg->deny_message : "{\"error\": \"Forbidden\"}", + req->keep_alive); + return 1; + } + return 0; /* 黑名单模式下允许 */ + } + + /* 检查是否匹配列表 */ + bool matched = false; + for (size_t i = 0; i < cfg->count; i++) { + if (ip_match_cidr(client_ip, cfg->entries[i])) { + matched = true; + break; + } + } + + /* 黑名单模式:匹配则拒绝 */ + if (cfg->mode == COCOON_IP_FILTER_DENY) { + if (matched) { + log_warn("IP 过滤: 客户端 IP 匹配黑名单规则,拒绝访问"); + send_json_error(fd, 403, + cfg->deny_message[0] ? cfg->deny_message : "{\"error\": \"Forbidden\"}", + req->keep_alive); + return 1; + } + return 0; /* 未匹配,允许 */ + } + + /* 白名单模式:不匹配则拒绝 */ + if (cfg->mode == COCOON_IP_FILTER_ALLOW) { + if (!matched) { + log_warn("IP 过滤: 客户端 IP 不匹配任何白名单规则,拒绝访问"); + send_json_error(fd, 403, + cfg->deny_message[0] ? cfg->deny_message : "{\"error\": \"Forbidden\"}", + req->keep_alive); + return 1; + } + return 0; /* 匹配,允许 */ + } + + return 0; +} + +/* ============================================================ + * 一键初始化 + * ============================================================ */ + +void cocoon_middleware_init_extended(void *server_config) { + (void)server_config; + + /* + * 此函数由 server.c 在启动时调用,根据配置注册扩展中间件。 + * 由于 cocoon_middleware_register 接受用户数据指针, + * 调用者需确保配置对象的生命周期覆盖整个服务运行期。 + * + * 当前实现为占位,实际注册由 server.c / config.c 完成, + * 它们在解析配置后调用 cocoon_middleware_register() 注册 + * 各个扩展中间件并传入对应的配置结构体。 + */ + log_info("扩展中间件系统已初始化"); +} diff --git a/middleware_ext.h b/middleware_ext.h new file mode 100644 index 0000000..4d52960 --- /dev/null +++ b/middleware_ext.h @@ -0,0 +1,184 @@ +/** + * @file middleware_ext.h - 扩展内置中间件接口 + * + * Phase 4 第一项:扩展中间件系统,新增 4 个内置中间件 + * - JWT 认证中间件 (middleware_jwt) + * - Security Headers 中间件 (middleware_security_headers) + * - Request ID 中间件 (middleware_request_id) + * - IP 过滤中间件 (middleware_ip_filter) + * + * @author xfy + */ + +#ifndef COCOON_MIDDLEWARE_EXT_H +#define COCOON_MIDDLEWARE_EXT_H + +#include "middleware.h" +#include "platform.h" +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +/* ===== JWT 认证中间件 ===== */ + +/** JWT 密钥最大长度 */ +#define COCOON_JWT_SECRET_MAX 256 + +/** JWT Token 最大长度 */ +#define COCOON_JWT_TOKEN_MAX 2048 + +/** + * @brief JWT 中间件配置 + * + * 用于 cocoon_middleware_jwt() 的用户数据。 + */ +typedef struct { + char secret[COCOON_JWT_SECRET_MAX]; /**< JWT 签名密钥 */ + char header_name[64]; /**< 默认 "Authorization" */ + char prefix[16]; /**< 默认 "Bearer " */ + uint32_t max_age; /**< Token 最大有效期(秒,默认 3600) */ + bool skip_preflight; /**< 是否跳过 OPTIONS 预检请求 */ +} cocoon_jwt_config_t; + +/** + * @brief JWT 认证中间件 + * + * 验证 Authorization: Bearer 中的 JWT token。 + * 使用 HS256 (HMAC-SHA256) 签名验证。 + * Token 格式:header.payload.signature(Base64Url 编码) + * + * @param req HTTP 请求 + * @param fd 客户端 socket + * @param user_data 指向 cocoon_jwt_config_t 的指针 + * @return 0 验证通过继续处理,1 验证失败(已发送 401) + */ +int cocoon_middleware_jwt(http_request_t *req, cocoon_socket_t fd, void *user_data); + +/* ===== Security Headers 中间件 ===== */ + +/** + * @brief Security Headers 全局配置 + * + * 由中间件设置,server.c 发送响应前读取。 + */ +typedef struct { + bool hsts_enabled; /**< Strict-Transport-Security */ + uint32_t hsts_max_age; /**< max-age(秒,默认 31536000 = 1年) */ + bool hsts_include_subdomains;/**< includeSubDomains */ + bool frame_options_enabled; /**< X-Frame-Options */ + char frame_options[32]; /**< DENY 或 SAMEORIGIN(默认 DENY) */ + bool xss_protection_enabled; /**< X-XSS-Protection */ + bool csp_enabled; /**< Content-Security-Policy */ + char csp_policy[512]; /**< CSP 策略字符串 */ + bool content_type_options; /**< X-Content-Type-Options: nosniff */ + bool referrer_policy_enabled;/**< Referrer-Policy */ + char referrer_policy[32]; /**< strict-origin-when-cross-origin */ +} cocoon_security_headers_config_t; + +/** + * @brief Security Headers 中间件 + * + * 设置全局安全头配置。中间件本身返回 0(继续处理), + * 由 server.c 在发送响应时调用 cocoon_middleware_security_headers_get() 获取配置。 + * + * @param req HTTP 请求 + * @param fd 客户端 socket(未使用) + * @param user_data 指向 cocoon_security_headers_config_t 的指针 + * @return 始终返回 0(继续处理) + */ +int cocoon_middleware_security_headers(http_request_t *req, cocoon_socket_t fd, void *user_data); + +/** + * @brief 获取当前 Security Headers 配置 + * + * 供 server.c 在发送响应前调用,将安全头追加到响应中。 + * + * @return 指向当前全局配置的指针,未设置时返回 NULL + */ +const cocoon_security_headers_config_t *cocoon_middleware_security_headers_get(void); + +/* ===== Request ID 中间件 ===== */ + +/** Request ID 长度(32 字符 hex) */ +#define COCOON_REQUEST_ID_LEN 32 + +/** + * @brief Request ID 中间件配置 + */ +typedef struct { + char header_name[32]; /**< 默认 "X-Request-ID" */ + bool trust_incoming; /**< 是否信任客户端传入的 ID(默认 true) */ +} cocoon_request_id_config_t; + +/** + * @brief Request ID 中间件 + * + * 生成 32 字符 hex 唯一追踪 ID。 + * 如果 trust_incoming=true 且请求头已有 ID,则复用。 + * 将 ID 添加到响应头 X-Request-ID。 + * + * @param req HTTP 请求 + * @param fd 客户端 socket + * @param user_data 指向 cocoon_request_id_config_t 的指针 + * @return 始终返回 0(继续处理) + */ +int cocoon_middleware_request_id(http_request_t *req, cocoon_socket_t fd, void *user_data); + +/* ===== IP 过滤中间件 ===== */ + +/** IP 过滤列表最大条目数 */ +#define COCOON_IP_FILTER_MAX 64 + +/** + * @brief IP 过滤模式 + */ +typedef enum { + COCOON_IP_FILTER_DENY, /**< 黑名单模式:匹配则拒绝 */ + COCOON_IP_FILTER_ALLOW /**< 白名单模式:不匹配则拒绝 */ +} cocoon_ip_filter_mode_t; + +/** + * @brief IP 过滤中间件配置 + * + * 支持 IPv4 地址和 CIDR(如 "192.168.1.0/24")。 + */ +typedef struct { + char entries[COCOON_IP_FILTER_MAX][64]; /**< IP 或 CIDR 字符串 */ + size_t count; /**< 实际条目数 */ + cocoon_ip_filter_mode_t mode; /**< 黑名单或白名单模式 */ + char deny_message[256]; /**< 拒绝时的响应消息 */ +} cocoon_ip_filter_config_t; + +/** + * @brief IP 过滤中间件 + * + * 黑名单模式:匹配则返回 403 + * 白名单模式:不匹配则返回 403 + * 支持解析 X-Forwarded-For 头获取真实 IP。 + * + * @param req HTTP 请求 + * @param fd 客户端 socket + * @param user_data 指向 cocoon_ip_filter_config_t 的指针 + * @return 0 允许继续处理,1 已拒绝(已发送 403) + */ +int cocoon_middleware_ip_filter(http_request_t *req, cocoon_socket_t fd, void *user_data); + +/* ===== 一键初始化 ===== */ + +/** + * @brief 一键注册所有扩展中间件 + * + * 根据配置注册 Phase 4 新增的 4 个内置中间件。 + * + * @param server_config 服务器配置(可包含中间件配置) + */ +void cocoon_middleware_init_extended(void *server_config); + +#ifdef __cplusplus +} +#endif + +#endif /* COCOON_MIDDLEWARE_EXT_H */ diff --git a/tests/unit/test_grpc.c b/tests/unit/test_grpc.c new file mode 100644 index 0000000..0815d01 --- /dev/null +++ b/tests/unit/test_grpc.c @@ -0,0 +1,1057 @@ +#include "unity.h" +#include "grpc.h" +#include +#include + +/* ===== 辅助函数 ===== */ + +/** + * fill_http_request - 填充 HTTP 请求结构体(用于测试) + */ +static void fill_http_request(http_request_t *req, + const char *path, + const char *content_type, + const uint8_t *body, + size_t body_len) { + memset(req, 0, sizeof(http_request_t)); + if (path) { + snprintf(req->path, sizeof(req->path), "%s", path); + } + if (content_type) { + snprintf(req->content_type, sizeof(req->content_type), "%s", content_type); + /* 同时加入 headers 数组 */ + snprintf(req->headers[0].name, sizeof(req->headers[0].name), "content-type"); + snprintf(req->headers[0].value, sizeof(req->headers[0].value), "%s", content_type); + req->num_headers = 1; + } + if (body && body_len > 0) { + req->body = (char *)malloc(body_len); + memcpy(req->body, body, body_len); + req->body_len = body_len; + } +} + +/** + * make_grpc_frame - 构造 gRPC 消息帧(用于测试输入) + * + * @param compressed 压缩标志 + * @param payload payload 数据 + * @param payload_len payload 长度 + * @param out_len 输出帧长度 + * @return 帧缓冲区(需由调用者释放) + */ +static uint8_t *make_grpc_frame(uint8_t compressed, const uint8_t *payload, + uint32_t payload_len, size_t *out_len) { + *out_len = 5 + payload_len; + uint8_t *frame = (uint8_t *)malloc(*out_len); + frame[0] = compressed; + frame[1] = (uint8_t)((payload_len >> 24) & 0xFF); + frame[2] = (uint8_t)((payload_len >> 16) & 0xFF); + frame[3] = (uint8_t)((payload_len >> 8) & 0xFF); + frame[4] = (uint8_t)(payload_len & 0xFF); + if (payload_len > 0) { + memcpy(frame + 5, payload, payload_len); + } + return frame; +} + +/* ===== grpc_detect ===== */ + +void test_detect_grpc_basic(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc"); + TEST_ASSERT_TRUE(grpc_detect(&req)); +} + +void test_detect_grpc_proto(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc+proto"); + TEST_ASSERT_TRUE(grpc_detect(&req)); +} + +void test_detect_grpc_json(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc+json"); + TEST_ASSERT_TRUE(grpc_detect(&req)); +} + +void test_detect_grpc_with_charset(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), + "application/grpc; charset=utf-8"); + TEST_ASSERT_TRUE(grpc_detect(&req)); +} + +void test_detect_grpc_web(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc-web"); + /* grpc_detect 应该也返回 true(因为 grpc-web 以 grpc 开头) */ + TEST_ASSERT_TRUE(grpc_detect(&req)); + /* grpc_is_grpc_web 应该返回 true */ + TEST_ASSERT_TRUE(grpc_is_grpc_web(&req)); +} + +void test_detect_grpc_web_proto(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), + "application/grpc-web+proto"); + TEST_ASSERT_TRUE(grpc_detect(&req)); + TEST_ASSERT_TRUE(grpc_is_grpc_web(&req)); +} + +void test_detect_not_grpc(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/json"); + TEST_ASSERT_FALSE(grpc_detect(&req)); +} + +void test_detect_text_plain(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "text/plain"); + TEST_ASSERT_FALSE(grpc_detect(&req)); +} + +void test_detect_grpc_uppercase(void) { + /* Content-Type 检测应该是大小写不敏感的 */ + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "APPLICATION/GRPC"); + TEST_ASSERT_TRUE(grpc_detect(&req)); +} + +void test_detect_null_request(void) { + TEST_ASSERT_FALSE(grpc_detect(NULL)); +} + +void test_detect_empty_content_type(void) { + http_request_t req = {0}; + /* content_type 为空字符串 */ + TEST_ASSERT_FALSE(grpc_detect(&req)); +} + +void test_detect_grpc_similar_prefix(void) { + /* "application/grpc-like" 不应该匹配 */ + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc-like"); + TEST_ASSERT_FALSE(grpc_detect(&req)); +} + +/* ===== grpc_is_grpc_web ===== */ + +void test_is_grpc_web_direct(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc-web"); + TEST_ASSERT_TRUE(grpc_is_grpc_web(&req)); +} + +void test_is_grpc_web_not_grpc(void) { + http_request_t req = {0}; + snprintf(req.content_type, sizeof(req.content_type), "application/grpc"); + TEST_ASSERT_FALSE(grpc_is_grpc_web(&req)); +} + +void test_is_grpc_web_null(void) { + TEST_ASSERT_FALSE(grpc_is_grpc_web(NULL)); +} + +/* ===== grpc_decode_message ===== */ + +void test_decode_simple_message(void) { + const uint8_t payload[] = {0x0A, 0x04, 0x74, 0x65, 0x73, 0x74}; /* protobuf-like */ + size_t frame_len; + uint8_t *frame = make_grpc_frame(0x00, payload, sizeof(payload), &frame_len); + + grpc_message_t msg = {0}; + int ret = grpc_decode_message(frame, frame_len, &msg); + + TEST_ASSERT_EQUAL_INT(5 + (int)sizeof(payload), ret); + TEST_ASSERT_EQUAL_UINT8(0x00, msg.compressed); + TEST_ASSERT_EQUAL_UINT32(sizeof(payload), msg.length); + TEST_ASSERT_NOT_NULL(msg.payload); + TEST_ASSERT_EQUAL_UINT8_ARRAY(payload, msg.payload, sizeof(payload)); + + grpc_message_free(&msg); + free(frame); +} + +void test_decode_empty_payload(void) { + size_t frame_len; + uint8_t *frame = make_grpc_frame(0x00, NULL, 0, &frame_len); + + grpc_message_t msg = {0}; + int ret = grpc_decode_message(frame, frame_len, &msg); + + TEST_ASSERT_EQUAL_INT(5, ret); + TEST_ASSERT_EQUAL_UINT8(0x00, msg.compressed); + TEST_ASSERT_EQUAL_UINT32(0, msg.length); + TEST_ASSERT_NULL(msg.payload); + + grpc_message_free(&msg); + free(frame); +} + +void test_decode_compressed_flag(void) { + const uint8_t payload[] = {0x01, 0x02, 0x03}; + size_t frame_len; + uint8_t *frame = make_grpc_frame(0x01, payload, sizeof(payload), &frame_len); + + grpc_message_t msg = {0}; + int ret = grpc_decode_message(frame, frame_len, &msg); + + TEST_ASSERT_EQUAL_INT(5 + (int)sizeof(payload), ret); + TEST_ASSERT_EQUAL_UINT8(0x01, msg.compressed); + TEST_ASSERT_EQUAL_UINT32(sizeof(payload), msg.length); + + grpc_message_free(&msg); + free(frame); +} + +void test_decode_large_payload(void) { + /* 测试大 payload(1KB) */ + uint8_t *large_payload = (uint8_t *)malloc(1024); + for (int i = 0; i < 1024; i++) large_payload[i] = (uint8_t)(i & 0xFF); + + size_t frame_len; + uint8_t *frame = make_grpc_frame(0x00, large_payload, 1024, &frame_len); + + grpc_message_t msg = {0}; + int ret = grpc_decode_message(frame, frame_len, &msg); + + TEST_ASSERT_EQUAL_INT(5 + 1024, ret); + TEST_ASSERT_EQUAL_UINT32(1024, msg.length); + TEST_ASSERT_NOT_NULL(msg.payload); + TEST_ASSERT_EQUAL_UINT8_ARRAY(large_payload, msg.payload, 1024); + + grpc_message_free(&msg); + free(frame); + free(large_payload); +} + +void test_decode_incomplete_header(void) { + /* 只有 3 字节,不足 5 字节头部 */ + const uint8_t buf[] = {0x00, 0x00, 0x00}; + grpc_message_t msg = {0}; + int ret = grpc_decode_message(buf, sizeof(buf), &msg); + + TEST_ASSERT_EQUAL_INT(-1, ret); + /* msg 不应该被修改(payload 为 NULL) */ + TEST_ASSERT_NULL(msg.payload); +} + +void test_decode_incomplete_payload(void) { + /* 头部完整但 payload 不完整 */ + const uint8_t buf[] = {0x00, 0x00, 0x00, 0x00, 0x10}; /* length=16, 但无 payload */ + grpc_message_t msg = {0}; + int ret = grpc_decode_message(buf, sizeof(buf), &msg); + + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_decode_null_buffer(void) { + grpc_message_t msg = {0}; + int ret = grpc_decode_message(NULL, 10, &msg); + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_decode_zero_length(void) { + const uint8_t buf[] = {0x00, 0x00, 0x00, 0x00, 0x00}; + grpc_message_t msg = {0}; + int ret = grpc_decode_message(buf, 0, &msg); + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_decode_null_message(void) { + const uint8_t buf[] = {0x00, 0x00, 0x00, 0x00, 0x00}; + int ret = grpc_decode_message(buf, sizeof(buf), NULL); + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +/* ===== grpc_encode_message ===== */ + +void test_encode_simple_message(void) { + const uint8_t payload[] = {0x0A, 0x04, 0x74, 0x65, 0x73, 0x74}; + grpc_message_t msg = { + .compressed = 0x00, + .length = sizeof(payload), + .payload = (uint8_t *)payload + }; + + uint8_t buf[64]; + int ret = grpc_encode_message(&msg, buf, sizeof(buf)); + + TEST_ASSERT_EQUAL_INT(5 + (int)sizeof(payload), ret); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[0]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[1]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[2]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[3]); + TEST_ASSERT_EQUAL_UINT8((int)sizeof(payload), buf[4]); + TEST_ASSERT_EQUAL_UINT8_ARRAY(payload, buf + 5, sizeof(payload)); +} + +void test_encode_empty_payload(void) { + grpc_message_t msg = { + .compressed = 0x00, + .length = 0, + .payload = NULL + }; + + uint8_t buf[16]; + int ret = grpc_encode_message(&msg, buf, sizeof(buf)); + + TEST_ASSERT_EQUAL_INT(5, ret); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[0]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[1]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[2]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[3]); + TEST_ASSERT_EQUAL_UINT8(0x00, buf[4]); +} + +void test_encode_buffer_too_small(void) { + const uint8_t payload[] = {0x01, 0x02, 0x03, 0x04, 0x05}; + grpc_message_t msg = { + .compressed = 0x00, + .length = sizeof(payload), + .payload = (uint8_t *)payload + }; + + uint8_t buf[8]; /* 太小,需要 10 字节 */ + int ret = grpc_encode_message(&msg, buf, sizeof(buf)); + + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_encode_null_params(void) { + uint8_t buf[16]; + int ret = grpc_encode_message(NULL, buf, sizeof(buf)); + TEST_ASSERT_EQUAL_INT(-1, ret); + + uint8_t payload[] = {0x01}; + grpc_message_t msg = {0}; + msg.payload = payload; + msg.length = 1; + ret = grpc_encode_message(&msg, NULL, 16); + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_encode_exact_buffer_size(void) { + /* 缓冲区大小恰好等于所需大小 */ + const uint8_t payload[] = {0xAB, 0xCD}; + grpc_message_t msg = { + .compressed = 0x01, + .length = sizeof(payload), + .payload = (uint8_t *)payload + }; + + uint8_t buf[7]; /* 5 + 2 = 7 */ + int ret = grpc_encode_message(&msg, buf, sizeof(buf)); + + TEST_ASSERT_EQUAL_INT(7, ret); + TEST_ASSERT_EQUAL_UINT8(0x01, buf[0]); + TEST_ASSERT_EQUAL_UINT8_ARRAY(payload, buf + 5, sizeof(payload)); +} + +/* ===== grpc_parse_request ===== */ + +void test_parse_request_basic(void) { + http_request_t req = {0}; + fill_http_request(&req, "/myPackage.MyService/MyMethod", + "application/grpc", NULL, 0); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(0, ret); + TEST_ASSERT_EQUAL_STRING("myPackage.MyService", grpc_req.service_name); + TEST_ASSERT_EQUAL_STRING("MyMethod", grpc_req.method_name); + TEST_ASSERT_FALSE(grpc_req.is_grpc_web); + TEST_ASSERT_FALSE(grpc_req.client_streaming); + TEST_ASSERT_FALSE(grpc_req.server_streaming); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +void test_parse_request_no_package(void) { + http_request_t req = {0}; + fill_http_request(&req, "/MyService/MyMethod", + "application/grpc", NULL, 0); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(0, ret); + TEST_ASSERT_EQUAL_STRING("MyService", grpc_req.service_name); + TEST_ASSERT_EQUAL_STRING("MyMethod", grpc_req.method_name); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +void test_parse_request_with_body(void) { + const uint8_t payload[] = {0x08, 0x01, 0x12, 0x04, 0x74, 0x65, 0x73, 0x74}; + size_t frame_len; + uint8_t *frame = make_grpc_frame(0x00, payload, sizeof(payload), &frame_len); + + http_request_t req = {0}; + fill_http_request(&req, "/TestService/Echo", + "application/grpc", frame, frame_len); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(0, ret); + TEST_ASSERT_EQUAL_STRING("TestService", grpc_req.service_name); + TEST_ASSERT_EQUAL_STRING("Echo", grpc_req.method_name); + TEST_ASSERT_EQUAL_UINT8(0x00, grpc_req.message.compressed); + TEST_ASSERT_EQUAL_UINT32(sizeof(payload), grpc_req.message.length); + TEST_ASSERT_NOT_NULL(grpc_req.message.payload); + TEST_ASSERT_EQUAL_UINT8_ARRAY(payload, grpc_req.message.payload, sizeof(payload)); + + grpc_request_free(&grpc_req); + http_request_free(&req); + free(frame); +} + +void test_parse_request_grpc_web(void) { + http_request_t req = {0}; + fill_http_request(&req, "/WebService/WebMethod", + "application/grpc-web", NULL, 0); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(0, ret); + TEST_ASSERT_TRUE(grpc_req.is_grpc_web); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +void test_parse_request_invalid_path_no_slash(void) { + http_request_t req = {0}; + fill_http_request(&req, "invalid", + "application/grpc", NULL, 0); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(-1, ret); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +void test_parse_request_empty_path(void) { + http_request_t req = {0}; + fill_http_request(&req, "", + "application/grpc", NULL, 0); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(-1, ret); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +void test_parse_request_null_params(void) { + int ret = grpc_parse_request(NULL, NULL); + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_parse_request_metadata(void) { + http_request_t req = {0}; + fill_http_request(&req, "/Svc/Method", + "application/grpc", NULL, 0); + /* 添加额外元数据头 */ + snprintf(req.headers[1].name, sizeof(req.headers[1].name), "x-request-id"); + snprintf(req.headers[1].value, sizeof(req.headers[1].value), "abc-123"); + snprintf(req.headers[2].name, sizeof(req.headers[2].name), "authorization"); + snprintf(req.headers[2].value, sizeof(req.headers[2].value), "Bearer token"); + req.num_headers = 3; + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(0, ret); + TEST_ASSERT_GREATER_THAN(0, (int)grpc_req.metadata_count); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +/* ===== grpc_format_response_trailers ===== */ + +void test_format_trailers_ok(void) { + grpc_request_t grpc_req = {0}; + grpc_req.status = GRPC_OK; + snprintf(grpc_req.status_message, sizeof(grpc_req.status_message), "OK"); + + char buf[256]; + int ret = grpc_format_response_trailers(&grpc_req, buf, sizeof(buf)); + + TEST_ASSERT_GREATER_THAN(0, ret); + TEST_ASSERT_TRUE(strstr(buf, "grpc-status: 0") != NULL); + TEST_ASSERT_TRUE(strstr(buf, "grpc-message: OK") != NULL); +} + +void test_format_trailers_error(void) { + grpc_request_t grpc_req = {0}; + grpc_req.status = GRPC_NOT_FOUND; + snprintf(grpc_req.status_message, sizeof(grpc_req.status_message), + "Method not found"); + + char buf[256]; + int ret = grpc_format_response_trailers(&grpc_req, buf, sizeof(buf)); + + TEST_ASSERT_GREATER_THAN(0, ret); + TEST_ASSERT_TRUE(strstr(buf, "grpc-status: 5") != NULL); + TEST_ASSERT_TRUE(strstr(buf, "grpc-message: Method not found") != NULL); +} + +void test_format_trailers_empty_message(void) { + grpc_request_t grpc_req = {0}; + grpc_req.status = GRPC_OK; + /* status_message 保持空 */ + + char buf[256]; + int ret = grpc_format_response_trailers(&grpc_req, buf, sizeof(buf)); + + TEST_ASSERT_GREATER_THAN(0, ret); + TEST_ASSERT_TRUE(strstr(buf, "grpc-status: 0") != NULL); +} + +void test_format_trailers_buffer_too_small(void) { + grpc_request_t grpc_req = {0}; + grpc_req.status = GRPC_OK; + snprintf(grpc_req.status_message, sizeof(grpc_req.status_message), + "A very long message that will definitely not fit in a tiny buffer"); + + char buf[10]; + int ret = grpc_format_response_trailers(&grpc_req, buf, sizeof(buf)); + + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +void test_format_trailers_null_params(void) { + char buf[256]; + int ret = grpc_format_response_trailers(NULL, buf, sizeof(buf)); + TEST_ASSERT_EQUAL_INT(-1, ret); + + grpc_request_t grpc_req = {0}; + ret = grpc_format_response_trailers(&grpc_req, NULL, sizeof(buf)); + TEST_ASSERT_EQUAL_INT(-1, ret); +} + +/* ===== grpc_status_to_http ===== */ + +void test_status_to_http_ok(void) { + TEST_ASSERT_EQUAL_INT(200, grpc_status_to_http(GRPC_OK)); +} + +void test_status_to_http_cancelled(void) { + TEST_ASSERT_EQUAL_INT(499, grpc_status_to_http(GRPC_CANCELLED)); +} + +void test_status_to_http_unknown(void) { + TEST_ASSERT_EQUAL_INT(500, grpc_status_to_http(GRPC_UNKNOWN)); +} + +void test_status_to_http_invalid_argument(void) { + TEST_ASSERT_EQUAL_INT(400, grpc_status_to_http(GRPC_INVALID_ARGUMENT)); +} + +void test_status_to_http_deadline_exceeded(void) { + TEST_ASSERT_EQUAL_INT(504, grpc_status_to_http(GRPC_DEADLINE_EXCEEDED)); +} + +void test_status_to_http_not_found(void) { + TEST_ASSERT_EQUAL_INT(404, grpc_status_to_http(GRPC_NOT_FOUND)); +} + +void test_status_to_http_already_exists(void) { + TEST_ASSERT_EQUAL_INT(409, grpc_status_to_http(GRPC_ALREADY_EXISTS)); +} + +void test_status_to_http_permission_denied(void) { + TEST_ASSERT_EQUAL_INT(403, grpc_status_to_http(GRPC_PERMISSION_DENIED)); +} + +void test_status_to_http_resource_exhausted(void) { + TEST_ASSERT_EQUAL_INT(429, grpc_status_to_http(GRPC_RESOURCE_EXHAUSTED)); +} + +void test_status_to_http_unimplemented(void) { + TEST_ASSERT_EQUAL_INT(501, grpc_status_to_http(GRPC_UNIMPLEMENTED)); +} + +void test_status_to_http_unavailable(void) { + TEST_ASSERT_EQUAL_INT(503, grpc_status_to_http(GRPC_UNAVAILABLE)); +} + +void test_status_to_http_unauthenticated(void) { + TEST_ASSERT_EQUAL_INT(401, grpc_status_to_http(GRPC_UNAUTHENTICATED)); +} + +void test_status_to_http_internal(void) { + TEST_ASSERT_EQUAL_INT(500, grpc_status_to_http(GRPC_INTERNAL)); +} + +void test_status_to_http_data_loss(void) { + TEST_ASSERT_EQUAL_INT(500, grpc_status_to_http(GRPC_DATA_LOSS)); +} + +void test_status_to_http_all_codes(void) { + /* 验证所有 17 个状态码都能映射(不崩溃) */ + for (int i = 0; i < GRPC_STATUS_MAX; i++) { + int http = grpc_status_to_http((grpc_status_t)i); + TEST_ASSERT_GREATER_THAN(0, http); + TEST_ASSERT_LESS_THAN(600, http); + } +} + +/* ===== grpc_http_to_status ===== */ + +void test_http_to_status_200(void) { + TEST_ASSERT_EQUAL_INT(GRPC_OK, grpc_http_to_status(200)); +} + +void test_http_to_status_400(void) { + TEST_ASSERT_EQUAL_INT(GRPC_INVALID_ARGUMENT, grpc_http_to_status(400)); +} + +void test_http_to_status_401(void) { + TEST_ASSERT_EQUAL_INT(GRPC_UNAUTHENTICATED, grpc_http_to_status(401)); +} + +void test_http_to_status_403(void) { + TEST_ASSERT_EQUAL_INT(GRPC_PERMISSION_DENIED, grpc_http_to_status(403)); +} + +void test_http_to_status_404(void) { + TEST_ASSERT_EQUAL_INT(GRPC_NOT_FOUND, grpc_http_to_status(404)); +} + +void test_http_to_status_409(void) { + TEST_ASSERT_EQUAL_INT(GRPC_ABORTED, grpc_http_to_status(409)); +} + +void test_http_to_status_429(void) { + TEST_ASSERT_EQUAL_INT(GRPC_RESOURCE_EXHAUSTED, grpc_http_to_status(429)); +} + +void test_http_to_status_500(void) { + TEST_ASSERT_EQUAL_INT(GRPC_INTERNAL, grpc_http_to_status(500)); +} + +void test_http_to_status_501(void) { + TEST_ASSERT_EQUAL_INT(GRPC_UNIMPLEMENTED, grpc_http_to_status(501)); +} + +void test_http_to_status_503(void) { + TEST_ASSERT_EQUAL_INT(GRPC_UNAVAILABLE, grpc_http_to_status(503)); +} + +void test_http_to_status_504(void) { + TEST_ASSERT_EQUAL_INT(GRPC_DEADLINE_EXCEEDED, grpc_http_to_status(504)); +} + +void test_http_to_status_2xx_success(void) { + /* 201-299 都应映射到 OK */ + TEST_ASSERT_EQUAL_INT(GRPC_OK, grpc_http_to_status(201)); + TEST_ASSERT_EQUAL_INT(GRPC_OK, grpc_http_to_status(204)); + TEST_ASSERT_EQUAL_INT(GRPC_OK, grpc_http_to_status(299)); +} + +void test_http_to_status_4xx_default(void) { + /* 未明确映射的 4xx 应返回 INVALID_ARGUMENT */ + TEST_ASSERT_EQUAL_INT(GRPC_INVALID_ARGUMENT, grpc_http_to_status(418)); +} + +/* ===== grpc_status_to_string ===== */ + +void test_status_to_string_all(void) { + TEST_ASSERT_EQUAL_STRING("OK", grpc_status_to_string(GRPC_OK)); + TEST_ASSERT_EQUAL_STRING("CANCELLED", grpc_status_to_string(GRPC_CANCELLED)); + TEST_ASSERT_EQUAL_STRING("UNKNOWN", grpc_status_to_string(GRPC_UNKNOWN)); + TEST_ASSERT_EQUAL_STRING("INVALID_ARGUMENT", + grpc_status_to_string(GRPC_INVALID_ARGUMENT)); + TEST_ASSERT_EQUAL_STRING("DEADLINE_EXCEEDED", + grpc_status_to_string(GRPC_DEADLINE_EXCEEDED)); + TEST_ASSERT_EQUAL_STRING("NOT_FOUND", grpc_status_to_string(GRPC_NOT_FOUND)); + TEST_ASSERT_EQUAL_STRING("ALREADY_EXISTS", + grpc_status_to_string(GRPC_ALREADY_EXISTS)); + TEST_ASSERT_EQUAL_STRING("PERMISSION_DENIED", + grpc_status_to_string(GRPC_PERMISSION_DENIED)); + TEST_ASSERT_EQUAL_STRING("RESOURCE_EXHAUSTED", + grpc_status_to_string(GRPC_RESOURCE_EXHAUSTED)); + TEST_ASSERT_EQUAL_STRING("FAILED_PRECONDITION", + grpc_status_to_string(GRPC_FAILED_PRECONDITION)); + TEST_ASSERT_EQUAL_STRING("ABORTED", grpc_status_to_string(GRPC_ABORTED)); + TEST_ASSERT_EQUAL_STRING("OUT_OF_RANGE", + grpc_status_to_string(GRPC_OUT_OF_RANGE)); + TEST_ASSERT_EQUAL_STRING("UNIMPLEMENTED", + grpc_status_to_string(GRPC_UNIMPLEMENTED)); + TEST_ASSERT_EQUAL_STRING("INTERNAL", + grpc_status_to_string(GRPC_INTERNAL)); + TEST_ASSERT_EQUAL_STRING("UNAVAILABLE", + grpc_status_to_string(GRPC_UNAVAILABLE)); + TEST_ASSERT_EQUAL_STRING("DATA_LOSS", + grpc_status_to_string(GRPC_DATA_LOSS)); + TEST_ASSERT_EQUAL_STRING("UNAUTHENTICATED", + grpc_status_to_string(GRPC_UNAUTHENTICATED)); +} + +void test_status_to_string_unknown(void) { + /* 非法状态码应返回 "UNKNOWN" */ + TEST_ASSERT_EQUAL_STRING("UNKNOWN", grpc_status_to_string(99)); + TEST_ASSERT_EQUAL_STRING("UNKNOWN", grpc_status_to_string(GRPC_STATUS_MAX)); +} + +/* ===== grpc_message_free ===== */ + +void test_message_free_null(void) { + /* 不应崩溃 */ + grpc_message_free(NULL); +} + +void test_message_free_null_payload(void) { + grpc_message_t msg = {0}; + msg.payload = NULL; + msg.length = 0; + /* 不应崩溃 */ + grpc_message_free(&msg); +} + +void test_message_free_with_payload(void) { + grpc_message_t msg = {0}; + msg.payload = (uint8_t *)malloc(16); + msg.length = 16; + msg.compressed = 1; + + grpc_message_free(&msg); + + TEST_ASSERT_NULL(msg.payload); + TEST_ASSERT_EQUAL_UINT32(0, msg.length); + TEST_ASSERT_EQUAL_UINT8(0, msg.compressed); +} + +/* ===== grpc_request_free ===== */ + +void test_request_free_null(void) { + /* 不应崩溃 */ + grpc_request_free(NULL); +} + +void test_request_free_empty(void) { + grpc_request_t grpc_req = {0}; + /* 不应崩溃 */ + grpc_request_free(&grpc_req); +} + +void test_request_free_with_message(void) { + grpc_request_t grpc_req = {0}; + grpc_req.message.payload = (uint8_t *)malloc(32); + grpc_req.message.length = 32; + grpc_req.message.compressed = 0; + + grpc_request_free(&grpc_req); + + TEST_ASSERT_NULL(grpc_req.message.payload); + TEST_ASSERT_EQUAL_UINT32(0, grpc_req.message.length); +} + +/* ===== grpc_error_response ===== */ + +void test_error_response_invalid_socket(void) { + /* 无效 socket 不应崩溃 */ + grpc_error_response(COCOON_INVALID_SOCKET, 1, GRPC_INTERNAL, "error"); +} + +/* ===== 双向转换一致性 ===== */ + +void test_status_http_roundtrip(void) { + /* 常见状态码的双向转换一致性 */ + grpc_status_t statuses[] = { + GRPC_OK, GRPC_UNKNOWN, GRPC_INVALID_ARGUMENT, GRPC_NOT_FOUND, + GRPC_PERMISSION_DENIED, GRPC_UNIMPLEMENTED, GRPC_UNAVAILABLE, + GRPC_INTERNAL, GRPC_UNAUTHENTICATED + }; + int http_codes[] = {200, 500, 400, 404, 403, 501, 503, 500, 401}; + + for (size_t i = 0; i < sizeof(statuses)/sizeof(statuses[0]); i++) { + int http = grpc_status_to_http(statuses[i]); + TEST_ASSERT_EQUAL_INT(http_codes[i], http); + + grpc_status_t back = grpc_http_to_status(http); + /* 由于多对一映射,反向不一定完全一致,但应合理 */ + TEST_ASSERT_GREATER_OR_EQUAL(GRPC_OK, (int)back); + TEST_ASSERT_LESS_THAN(GRPC_STATUS_MAX, (int)back); + } +} + +/* ===== RPC 模式标记 ===== */ + +void test_rpc_mode_unary(void) { + grpc_request_t grpc_req = {0}; + grpc_req.client_streaming = false; + grpc_req.server_streaming = false; + grpc_req.is_streaming = grpc_req.client_streaming || grpc_req.server_streaming; + TEST_ASSERT_FALSE(grpc_req.is_streaming); +} + +void test_rpc_mode_server_streaming(void) { + grpc_request_t grpc_req = {0}; + grpc_req.client_streaming = false; + grpc_req.server_streaming = true; + grpc_req.is_streaming = grpc_req.client_streaming || grpc_req.server_streaming; + TEST_ASSERT_TRUE(grpc_req.is_streaming); +} + +void test_rpc_mode_client_streaming(void) { + grpc_request_t grpc_req = {0}; + grpc_req.client_streaming = true; + grpc_req.server_streaming = false; + grpc_req.is_streaming = grpc_req.client_streaming || grpc_req.server_streaming; + TEST_ASSERT_TRUE(grpc_req.is_streaming); +} + +void test_rpc_mode_bidirectional(void) { + grpc_request_t grpc_req = {0}; + grpc_req.client_streaming = true; + grpc_req.server_streaming = true; + grpc_req.is_streaming = grpc_req.client_streaming || grpc_req.server_streaming; + TEST_ASSERT_TRUE(grpc_req.is_streaming); +} + +/* ===== 边界条件:超长路径 ===== */ + +void test_parse_long_path(void) { + /* 构造接近 256 字节限制的路径 */ + char long_svc[260]; + char path[300]; + memset(long_svc, 'A', 250); + long_svc[250] = '\0'; + snprintf(path, sizeof(path), "/%s/MyMethod", long_svc); + + http_request_t req = {0}; + fill_http_request(&req, path, "application/grpc", NULL, 0); + + grpc_request_t grpc_req = {0}; + int ret = grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_EQUAL_INT(0, ret); + /* service_name 被截断到 255 字符 */ + TEST_ASSERT_GREATER_THAN(0, strlen(grpc_req.service_name)); + TEST_ASSERT_EQUAL_STRING("MyMethod", grpc_req.method_name); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +/* ===== 编码/解码对称性 ===== */ + +void test_encode_decode_symmetry(void) { + /* 编码后再解码应得到原始数据 */ + const uint8_t payload[] = { + 0x12, 0x34, 0x56, 0x78, 0x9A, 0xBC, 0xDE, 0xF0 + }; + grpc_message_t orig = { + .compressed = 0x01, + .length = sizeof(payload), + .payload = (uint8_t *)payload + }; + + uint8_t buf[64]; + int encoded = grpc_encode_message(&orig, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, encoded); + + grpc_message_t decoded = {0}; + int decoded_len = grpc_decode_message(buf, (size_t)encoded, &decoded); + + TEST_ASSERT_EQUAL_INT(encoded, decoded_len); + TEST_ASSERT_EQUAL_UINT8(orig.compressed, decoded.compressed); + TEST_ASSERT_EQUAL_UINT32(orig.length, decoded.length); + TEST_ASSERT_EQUAL_UINT8_ARRAY(payload, decoded.payload, sizeof(payload)); + + grpc_message_free(&decoded); +} + +void test_encode_decode_empty(void) { + /* 空 payload 编码/解码对称性 */ + grpc_message_t orig = { + .compressed = 0x00, + .length = 0, + .payload = NULL + }; + + uint8_t buf[16]; + int encoded = grpc_encode_message(&orig, buf, sizeof(buf)); + TEST_ASSERT_EQUAL_INT(5, encoded); + + grpc_message_t decoded = {0}; + int decoded_len = grpc_decode_message(buf, (size_t)encoded, &decoded); + + TEST_ASSERT_EQUAL_INT(5, decoded_len); + TEST_ASSERT_EQUAL_UINT8(0x00, decoded.compressed); + TEST_ASSERT_EQUAL_UINT32(0, decoded.length); + TEST_ASSERT_NULL(decoded.payload); + + grpc_message_free(&decoded); +} + +/* ===== gRPC-Web 特殊处理 ===== */ + +void test_grpc_web_flag_set(void) { + http_request_t req = {0}; + fill_http_request(&req, "/WebService/Method", + "application/grpc-web", NULL, 0); + + grpc_request_t grpc_req = {0}; + grpc_parse_request(&req, &grpc_req); + + TEST_ASSERT_TRUE(grpc_req.is_grpc_web); + TEST_ASSERT_FALSE(grpc_req.client_streaming); + TEST_ASSERT_FALSE(grpc_req.server_streaming); + + grpc_request_free(&grpc_req); + http_request_free(&req); +} + +/* ===== 主函数 ===== */ + +void setUp(void) { + /* 每个测试前执行 */ +} + +void tearDown(void) { + /* 每个测试后执行 */ +} + +int main(void) { + UNITY_BEGIN(); + + /* grpc_detect */ + RUN_TEST(test_detect_grpc_basic); + RUN_TEST(test_detect_grpc_proto); + RUN_TEST(test_detect_grpc_json); + RUN_TEST(test_detect_grpc_with_charset); + RUN_TEST(test_detect_grpc_web); + RUN_TEST(test_detect_grpc_web_proto); + RUN_TEST(test_detect_not_grpc); + RUN_TEST(test_detect_text_plain); + RUN_TEST(test_detect_grpc_uppercase); + RUN_TEST(test_detect_null_request); + RUN_TEST(test_detect_empty_content_type); + RUN_TEST(test_detect_grpc_similar_prefix); + + /* grpc_is_grpc_web */ + RUN_TEST(test_is_grpc_web_direct); + RUN_TEST(test_is_grpc_web_not_grpc); + RUN_TEST(test_is_grpc_web_null); + + /* grpc_decode_message */ + RUN_TEST(test_decode_simple_message); + RUN_TEST(test_decode_empty_payload); + RUN_TEST(test_decode_compressed_flag); + RUN_TEST(test_decode_large_payload); + RUN_TEST(test_decode_incomplete_header); + RUN_TEST(test_decode_incomplete_payload); + RUN_TEST(test_decode_null_buffer); + RUN_TEST(test_decode_zero_length); + RUN_TEST(test_decode_null_message); + + /* grpc_encode_message */ + RUN_TEST(test_encode_simple_message); + RUN_TEST(test_encode_empty_payload); + RUN_TEST(test_encode_buffer_too_small); + RUN_TEST(test_encode_null_params); + RUN_TEST(test_encode_exact_buffer_size); + + /* grpc_parse_request */ + RUN_TEST(test_parse_request_basic); + RUN_TEST(test_parse_request_no_package); + RUN_TEST(test_parse_request_with_body); + RUN_TEST(test_parse_request_grpc_web); + RUN_TEST(test_parse_request_invalid_path_no_slash); + RUN_TEST(test_parse_request_empty_path); + RUN_TEST(test_parse_request_null_params); + RUN_TEST(test_parse_request_metadata); + + /* grpc_format_response_trailers */ + RUN_TEST(test_format_trailers_ok); + RUN_TEST(test_format_trailers_error); + RUN_TEST(test_format_trailers_empty_message); + RUN_TEST(test_format_trailers_buffer_too_small); + RUN_TEST(test_format_trailers_null_params); + + /* grpc_status_to_http */ + RUN_TEST(test_status_to_http_ok); + RUN_TEST(test_status_to_http_cancelled); + RUN_TEST(test_status_to_http_unknown); + RUN_TEST(test_status_to_http_invalid_argument); + RUN_TEST(test_status_to_http_deadline_exceeded); + RUN_TEST(test_status_to_http_not_found); + RUN_TEST(test_status_to_http_already_exists); + RUN_TEST(test_status_to_http_permission_denied); + RUN_TEST(test_status_to_http_resource_exhausted); + RUN_TEST(test_status_to_http_unimplemented); + RUN_TEST(test_status_to_http_unavailable); + RUN_TEST(test_status_to_http_unauthenticated); + RUN_TEST(test_status_to_http_internal); + RUN_TEST(test_status_to_http_data_loss); + RUN_TEST(test_status_to_http_all_codes); + + /* grpc_http_to_status */ + RUN_TEST(test_http_to_status_200); + RUN_TEST(test_http_to_status_400); + RUN_TEST(test_http_to_status_401); + RUN_TEST(test_http_to_status_403); + RUN_TEST(test_http_to_status_404); + RUN_TEST(test_http_to_status_409); + RUN_TEST(test_http_to_status_429); + RUN_TEST(test_http_to_status_500); + RUN_TEST(test_http_to_status_501); + RUN_TEST(test_http_to_status_503); + RUN_TEST(test_http_to_status_504); + RUN_TEST(test_http_to_status_2xx_success); + RUN_TEST(test_http_to_status_4xx_default); + + /* grpc_status_to_string */ + RUN_TEST(test_status_to_string_all); + RUN_TEST(test_status_to_string_unknown); + + /* grpc_message_free */ + RUN_TEST(test_message_free_null); + RUN_TEST(test_message_free_null_payload); + RUN_TEST(test_message_free_with_payload); + + /* grpc_request_free */ + RUN_TEST(test_request_free_null); + RUN_TEST(test_request_free_empty); + RUN_TEST(test_request_free_with_message); + + /* grpc_error_response */ + RUN_TEST(test_error_response_invalid_socket); + + /* 双向转换一致性 */ + RUN_TEST(test_status_http_roundtrip); + + /* RPC 模式标记 */ + RUN_TEST(test_rpc_mode_unary); + RUN_TEST(test_rpc_mode_server_streaming); + RUN_TEST(test_rpc_mode_client_streaming); + RUN_TEST(test_rpc_mode_bidirectional); + + /* 边界条件 */ + RUN_TEST(test_parse_long_path); + + /* 编码/解码对称性 */ + RUN_TEST(test_encode_decode_symmetry); + RUN_TEST(test_encode_decode_empty); + + /* gRPC-Web */ + RUN_TEST(test_grpc_web_flag_set); + + return UNITY_END(); +} diff --git a/tests/unit/test_http3.c b/tests/unit/test_http3.c new file mode 100644 index 0000000..55cb9d1 --- /dev/null +++ b/tests/unit/test_http3.c @@ -0,0 +1,989 @@ +/** + * @file test_http3.c + * @brief HTTP/3 (QUIC) 模块单元测试 + * + * 使用 Unity 测试框架,覆盖以下功能: + * - Variable-length integer 编解码 + * - HTTP/3 帧头编解码 + * - QPACK 静态表编码/解码 + * - QUIC 流管理 + * - QUIC 连接管理 + * - HTTP/3 会话管理 + * - SETTINGS / GOAWAY 帧处理 + * + * @author Cocoon Team + */ + +#include "unity.h" +#include "http3.h" +#include +#include +#include +#include + +/* ===== 测试前置/后置 ===== */ + +void setUp(void) { + http3_init(); +} + +void tearDown(void) { + http3_cleanup(); +} + +/* ===== Variable-Length Integer 测试 ===== */ + +/** @test varint 编码:0(最小1字节值) */ +void test_varint_encode_zero(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(0, buf); + TEST_ASSERT_EQUAL(1, n); + TEST_ASSERT_EQUAL(0, buf[0]); +} + +/** @test varint 编码:63(1字节最大值) */ +void test_varint_encode_63(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(63, buf); + TEST_ASSERT_EQUAL(1, n); + TEST_ASSERT_EQUAL(63, buf[0]); +} + +/** @test varint 编码:64(需要2字节) */ +void test_varint_encode_64(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(64, buf); + TEST_ASSERT_EQUAL(2, n); + TEST_ASSERT_EQUAL(0x40, buf[0]); + TEST_ASSERT_EQUAL(0x40, buf[1]); +} + +/** @test varint 编码:16383(2字节最大值) */ +void test_varint_encode_16383(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(16383, buf); + TEST_ASSERT_EQUAL(2, n); + TEST_ASSERT_EQUAL(0x7F, buf[0]); + TEST_ASSERT_EQUAL(0xFF, buf[1]); +} + +/** @test varint 编码:16384(需要4字节) */ +void test_varint_encode_16384(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(16384, buf); + TEST_ASSERT_EQUAL(4, n); + TEST_ASSERT_EQUAL(0x80, buf[0]); + TEST_ASSERT_EQUAL(0x00, buf[1]); + TEST_ASSERT_EQUAL(0x40, buf[2]); + TEST_ASSERT_EQUAL(0x00, buf[3]); +} + +/** @test varint 编码:1073741823(4字节最大值) */ +void test_varint_encode_4byte_max(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(1073741823ULL, buf); + TEST_ASSERT_EQUAL(4, n); + TEST_ASSERT_EQUAL(0xBF, buf[0]); + TEST_ASSERT_EQUAL(0xFF, buf[1]); + TEST_ASSERT_EQUAL(0xFF, buf[2]); + TEST_ASSERT_EQUAL(0xFF, buf[3]); +} + +/** @test varint 编码:1073741824(需要8字节) */ +void test_varint_encode_8byte(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(1073741824ULL, buf); + TEST_ASSERT_EQUAL(8, n); + TEST_ASSERT_EQUAL(0xC0, buf[0]); +} + +/** @test varint 编码:最大值 */ +void test_varint_encode_max(void) { + uint8_t buf[8]; + size_t n = http3_encode_varint(4611686018427387903ULL, buf); + TEST_ASSERT_EQUAL(8, n); + TEST_ASSERT_EQUAL(0xFF, buf[0]); + TEST_ASSERT_EQUAL(0xFF, buf[1]); + TEST_ASSERT_EQUAL(0xFF, buf[2]); + TEST_ASSERT_EQUAL(0xFF, buf[3]); + TEST_ASSERT_EQUAL(0xFF, buf[4]); + TEST_ASSERT_EQUAL(0xFF, buf[5]); + TEST_ASSERT_EQUAL(0xFF, buf[6]); + TEST_ASSERT_EQUAL(0xFF, buf[7]); +} + +/** @test varint 编码:NULL 缓冲区 */ +void test_varint_encode_null_buf(void) { + size_t n = http3_encode_varint(42, NULL); + TEST_ASSERT_EQUAL(0, n); +} + +/** @test varint 解码:0 */ +void test_varint_decode_zero(void) { + uint8_t buf[1] = {0x00}; + uint64_t value = 0; + int n = http3_decode_varint(buf, 1, &value); + TEST_ASSERT_EQUAL(1, n); + TEST_ASSERT_EQUAL(0ULL, value); +} + +/** @test varint 解码:63 */ +void test_varint_decode_63(void) { + uint8_t buf[1] = {0x3F}; + uint64_t value = 0; + int n = http3_decode_varint(buf, 1, &value); + TEST_ASSERT_EQUAL(1, n); + TEST_ASSERT_EQUAL(63ULL, value); +} + +/** @test varint 解码:64 */ +void test_varint_decode_64(void) { + uint8_t buf[2] = {0x40, 0x40}; + uint64_t value = 0; + int n = http3_decode_varint(buf, 2, &value); + TEST_ASSERT_EQUAL(2, n); + TEST_ASSERT_EQUAL(64ULL, value); +} + +/** @test varint 解码:16383 */ +void test_varint_decode_16383(void) { + uint8_t buf[2] = {0x7F, 0xFF}; + uint64_t value = 0; + int n = http3_decode_varint(buf, 2, &value); + TEST_ASSERT_EQUAL(2, n); + TEST_ASSERT_EQUAL(16383ULL, value); +} + +/** @test varint 解码:4字节值 */ +void test_varint_decode_4byte(void) { + uint8_t buf[4] = {0x80, 0x00, 0x00, 0x01}; + uint64_t value = 0; + int n = http3_decode_varint(buf, 4, &value); + TEST_ASSERT_EQUAL(4, n); + TEST_ASSERT_EQUAL(1ULL, value); +} + +/** @test varint 解码:数据不足 */ +void test_varint_decode_insufficient_data(void) { + uint8_t buf[1] = {0x40}; /* 需要2字节,但只给1字节 */ + uint64_t value = 0; + int n = http3_decode_varint(buf, 1, &value); + TEST_ASSERT_EQUAL(-1, n); +} + +/** @test varint 解码:NULL 输入 */ +void test_varint_decode_null(void) { + uint64_t value = 0; + int n = http3_decode_varint(NULL, 0, &value); + TEST_ASSERT_EQUAL(-1, n); +} + +/** @test varint 编解码往返测试 */ +void test_varint_roundtrip(void) { + uint64_t test_values[] = {0, 1, 63, 64, 100, 16383, 16384, 100000, + 1073741823ULL, 1073741824ULL, + 4294967295ULL, 4611686018427387903ULL}; + int num_values = sizeof(test_values) / sizeof(test_values[0]); + + for (int i = 0; i < num_values; i++) { + uint8_t buf[8]; + size_t enc_n = http3_encode_varint(test_values[i], buf); + TEST_ASSERT_MESSAGE(enc_n > 0, "编码失败"); + + uint64_t decoded = 0; + int dec_n = http3_decode_varint(buf, enc_n, &decoded); + TEST_ASSERT_EQUAL_MESSAGE(enc_n, (size_t)dec_n, "编解码长度不匹配"); + TEST_ASSERT_EQUAL_MESSAGE(test_values[i], decoded, "编解码值不匹配"); + } +} + +/** @test 帧头大小计算 */ +void test_frame_header_size(void) { + TEST_ASSERT_EQUAL(2, http3_frame_header_size(0, 0)); + TEST_ASSERT_EQUAL(2, http3_frame_header_size(1, 10)); + TEST_ASSERT_EQUAL(4, http3_frame_header_size(100, 1000)); +} + +/* ===== HTTP/3 帧处理测试 ===== */ + +/** @test 帧头编码/解码 */ +void test_frame_header_encode_decode(void) { + uint8_t buf[16]; + uint64_t ft = HTTP3_FRAME_HEADERS; + uint64_t flen = 42; + + size_t n = http3_encode_frame_header(ft, flen, buf); + TEST_ASSERT_GREATER_THAN(0, n); + + uint64_t decoded_ft = 0, decoded_len = 0; + int dn = http3_decode_frame_header(buf, n, &decoded_ft, &decoded_len); + TEST_ASSERT_EQUAL(n, (size_t)dn); + TEST_ASSERT_EQUAL(ft, decoded_ft); + TEST_ASSERT_EQUAL(flen, decoded_len); +} + +/** @test 完整帧编码 */ +void test_full_frame_encode(void) { + uint8_t payload[] = "Hello HTTP/3"; + uint8_t buf[64]; + + int n = http3_encode_frame(HTTP3_FRAME_DATA, payload, sizeof(payload) - 1, + buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + + /* 验证帧头 */ + uint64_t ft = 0, flen = 0; + int hd = http3_decode_frame_header(buf, (size_t)n, &ft, &flen); + TEST_ASSERT_GREATER_THAN(0, hd); + TEST_ASSERT_EQUAL(HTTP3_FRAME_DATA, ft); + TEST_ASSERT_EQUAL(sizeof(payload) - 1, flen); +} + +/** @test 帧头解码:数据不足 */ +void test_frame_header_decode_insufficient(void) { + uint8_t buf[1] = {0x80}; /* 需要更多数据 */ + uint64_t ft = 0, flen = 0; + int n = http3_decode_frame_header(buf, 1, &ft, &flen); + TEST_ASSERT_EQUAL(-1, n); +} + +/** @test 帧头解码:NULL 输入 */ +void test_frame_header_decode_null(void) { + uint64_t ft = 0, flen = 0; + int n = http3_decode_frame_header(NULL, 0, &ft, &flen); + TEST_ASSERT_EQUAL(-1, n); +} + +/** @test 编码 DATA 帧 */ +void test_encode_data_frame(void) { + uint8_t payload[] = {0x01, 0x02, 0x03, 0x04}; + uint8_t buf[32]; + int n = http3_encode_frame(HTTP3_FRAME_DATA, payload, 4, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + + uint64_t ft = 0, flen = 0; + int hd = http3_decode_frame_header(buf, (size_t)n, &ft, &flen); + TEST_ASSERT_GREATER_THAN(0, hd); + TEST_ASSERT_EQUAL(HTTP3_FRAME_DATA, ft); + TEST_ASSERT_EQUAL(4, flen); +} + +/** @test 编码 HEADERS 帧 */ +void test_encode_headers_frame(void) { + uint8_t payload[] = {0xC0, 0x01, 0xD5}; /* 一些 QPACK 编码数据 */ + uint8_t buf[32]; + int n = http3_encode_frame(HTTP3_FRAME_HEADERS, payload, 3, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + + uint64_t ft = 0, flen = 0; + int hd = http3_decode_frame_header(buf, (size_t)n, &ft, &flen); + TEST_ASSERT_GREATER_THAN(0, hd); + TEST_ASSERT_EQUAL(HTTP3_FRAME_HEADERS, ft); + TEST_ASSERT_EQUAL(3, flen); +} + +/** @test 编码 SETTINGS 帧 */ +void test_encode_settings_frame(void) { + uint8_t payload[16]; + size_t pos = 0; + pos += http3_encode_varint(HTTP3_SETTING_MAX_FIELD_SECTION_SIZE, payload + pos); + pos += http3_encode_varint(16384, payload + pos); + + uint8_t buf[32]; + int n = http3_encode_frame(HTTP3_FRAME_SETTINGS, payload, pos, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + + uint64_t ft = 0, flen = 0; + int hd = http3_decode_frame_header(buf, (size_t)n, &ft, &flen); + TEST_ASSERT_GREATER_THAN(0, hd); + TEST_ASSERT_EQUAL(HTTP3_FRAME_SETTINGS, ft); +} + +/** @test 编码 GOAWAY 帧 */ +void test_encode_goaway_frame(void) { + uint8_t payload[8]; + size_t pos = http3_encode_varint(100, payload); + + uint8_t buf[32]; + int n = http3_encode_frame(HTTP3_FRAME_GOAWAY, payload, pos, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + + uint64_t ft = 0, flen = 0; + int hd = http3_decode_frame_header(buf, (size_t)n, &ft, &flen); + TEST_ASSERT_GREATER_THAN(0, hd); + TEST_ASSERT_EQUAL(HTTP3_FRAME_GOAWAY, ft); +} + +/* ===== QPACK 编码/解码测试 ===== */ + +/** @test QPACK 编码:静态表完全匹配 :method=GET */ +void test_qpack_encode_static_match_method_get(void) { + uint8_t buf[16]; + int n = qpack_encode_header(":method", "GET", buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + /* Should use indexed encoding (0b11xxxxxx) with static index 2 */ + /* buf[0] should be 0xC0 | 0x02 = 0xC2 = 194 */ + TEST_ASSERT_EQUAL(0xC0, buf[0] & 0xC0); /* 0b11 prefix */ +} + +/** @test QPACK 编码:静态表完全匹配 :path=/ */ +void test_qpack_encode_static_match_path(void) { + uint8_t buf[16]; + int n = qpack_encode_header(":path", "/", buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + TEST_ASSERT_EQUAL(0xC0, buf[0] & 0xC0); /* 0b11 prefix */ +} + +/** @test QPACK 编码:静态表完全匹配 :scheme=https */ +void test_qpack_encode_static_match_scheme(void) { + uint8_t buf[16]; + int n = qpack_encode_header(":scheme", "https", buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + TEST_ASSERT_EQUAL(0xC0, buf[0] & 0xC0); /* 0b11 prefix */ +} + +/** @test QPACK 编码:字面量编码未知字段 */ +void test_qpack_encode_literal_field(void) { + uint8_t buf[64]; + int n = qpack_encode_header("x-custom-header", "custom-value", buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + /* Should use literal encoding (0b001xxxxx) */ + TEST_ASSERT_EQUAL(0x20, buf[0]); +} + +/** @test QPACK 编码:名称匹配静态表,值用字面量 */ +void test_qpack_encode_name_ref_value_literal(void) { + uint8_t buf[64]; + int n = qpack_encode_header("accept", "text/html", buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); + /* accept=text/html is NOT a full static table match (accept="" at idx 17), + * so it uses literal with name reference (0b0101xxxx) */ + TEST_ASSERT_EQUAL(0x40, buf[0] & 0xE0); /* 0b010 prefix */ +} + +/** @test QPACK 解码:索引编码的头部 */ +void test_qpack_decode_indexed_header(void) { + /* Encode :method=GET - static index 2 */ + uint8_t encoded[16]; + int enc_n = qpack_encode_header(":method", "GET", encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, enc_n); + + /* Now decode it */ + qpack_decoded_t decoded; + size_t consumed = 0; + int rc = qpack_decode_header(encoded, (size_t)enc_n, &decoded, &consumed); + TEST_ASSERT_EQUAL(0, rc); + TEST_ASSERT_TRUE(decoded.valid); + TEST_ASSERT_EQUAL_STRING(":method", decoded.name); + TEST_ASSERT_EQUAL_STRING("GET", decoded.value); +} + +/** @test QPACK 解码:字面量编码的头部 */ +void test_qpack_decode_literal_header(void) { + /* Encode a literal field */ + uint8_t encoded[64]; + int enc_n = qpack_encode_header("x-test", "test-value", encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, enc_n); + + /* Decode it */ + qpack_decoded_t decoded; + size_t consumed = 0; + int rc = qpack_decode_header(encoded, (size_t)enc_n, &decoded, &consumed); + TEST_ASSERT_EQUAL(0, rc); + TEST_ASSERT_TRUE(decoded.valid); + TEST_ASSERT_EQUAL_STRING("x-test", decoded.name); + TEST_ASSERT_EQUAL_STRING("test-value", decoded.value); +} + +/** @test QPACK 编解码往返 */ +void test_qpack_roundtrip(void) { + const char *test_names[] = { + ":method", ":path", ":scheme", ":authority", + "content-type", "accept", "user-agent" + }; + const char *test_values[] = { + "GET", "/api/v1/test", "https", "example.com", + "application/json", "*/*", "TestAgent/1.0" + }; + int num = sizeof(test_names) / sizeof(test_names[0]); + + for (int i = 0; i < num; i++) { + uint8_t encoded[128]; + int enc_n = qpack_encode_header(test_names[i], test_values[i], + encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN_MESSAGE(0, enc_n, "编码失败"); + + qpack_decoded_t decoded; + size_t consumed = 0; + int rc = qpack_decode_header(encoded, (size_t)enc_n, &decoded, &consumed); + TEST_ASSERT_EQUAL_MESSAGE(0, rc, "解码失败"); + TEST_ASSERT_TRUE_MESSAGE(decoded.valid, "解码结果无效"); + TEST_ASSERT_EQUAL_STRING_MESSAGE(test_names[i], decoded.name, "名称不匹配"); + TEST_ASSERT_EQUAL_STRING_MESSAGE(test_values[i], decoded.value, "值不匹配"); + } +} + +/** @test QPACK 编码完整请求头 */ +void test_qpack_encode_request_headers(void) { + http_request_t req; + memset(&req, 0, sizeof(req)); + req.method = HTTP_GET; + strncpy(req.path, "/index.html", sizeof(req.path) - 1); + strncpy(req.headers[0].name, "host", sizeof(req.headers[0].name) - 1); + strncpy(req.headers[0].value, "example.com", sizeof(req.headers[0].value) - 1); + req.num_headers = 1; + + uint8_t buf[1024]; + int n = qpack_encode_request_headers(&req, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); +} + +/** @test QPACK 解码完整请求头 */ +void test_qpack_decode_request_headers(void) { + /* Build an encoded request */ + http_request_t req; + memset(&req, 0, sizeof(req)); + req.method = HTTP_GET; + strncpy(req.path, "/test", sizeof(req.path) - 1); + strncpy(req.headers[0].name, "host", sizeof(req.headers[0].name) - 1); + strncpy(req.headers[0].value, "example.com", sizeof(req.headers[0].value) - 1); + req.num_headers = 1; + + uint8_t encoded[1024]; + int enc_n = qpack_encode_request_headers(&req, encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, enc_n); + + /* Decode it */ + http_request_t decoded_req; + int rc = qpack_decode_request_headers(encoded, (size_t)enc_n, &decoded_req); + TEST_ASSERT_EQUAL(0, rc); + TEST_ASSERT_EQUAL(HTTP_GET, decoded_req.method); + TEST_ASSERT_EQUAL_STRING("/test", decoded_req.path); + http_request_free(&decoded_req); +} + +/** @test QPACK 编码响应头 */ +void test_qpack_encode_response_headers(void) { + http_response_t resp; + memset(&resp, 0, sizeof(resp)); + resp.status_code = 200; + resp.content_type = "text/html"; + resp.content_length = 42; + + uint8_t buf[1024]; + int n = qpack_encode_response_headers(&resp, buf, sizeof(buf)); + TEST_ASSERT_GREATER_THAN(0, n); +} + +/** @test QPACK 解码响应头 */ +void test_qpack_decode_response_headers(void) { + /* Build an encoded response */ + http_response_t resp; + memset(&resp, 0, sizeof(resp)); + resp.status_code = 404; + resp.content_type = "text/plain"; + resp.content_length = 0; + + uint8_t encoded[1024]; + int enc_n = qpack_encode_response_headers(&resp, encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, enc_n); + + /* Decode it */ + http_response_t decoded_resp; + int rc = qpack_decode_response_headers(encoded, (size_t)enc_n, &decoded_resp); + TEST_ASSERT_EQUAL(0, rc); + TEST_ASSERT_EQUAL(404, decoded_resp.status_code); + + if (decoded_resp.content_type) { + free((void *)decoded_resp.content_type); + } +} + +/** @test QPACK 解码:无效输入 */ +void test_qpack_decode_invalid(void) { + qpack_decoded_t decoded; + size_t consumed = 0; + int rc = qpack_decode_header(NULL, 0, &decoded, &consumed); + TEST_ASSERT_EQUAL(-1, rc); +} + +/* ===== QUIC 流管理测试 ===== */ + +/** @test QUIC 流创建和查找 */ +void test_quic_stream_create_and_find(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(12345, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_stream_t *s = quic_stream_get_or_create(conn, 0); + TEST_ASSERT_NOT_NULL(s); + TEST_ASSERT_EQUAL(0, s->stream_id); + + quic_stream_t *found = quic_stream_find(conn, 0); + TEST_ASSERT_NOT_NULL(found); + TEST_ASSERT_EQUAL_PTR(s, found); + + quic_connection_destroy(conn); +} + +/** @test QUIC 流数据写入和读取 */ +void test_quic_stream_write_read(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(12345, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_stream_t *s = quic_stream_get_or_create(conn, 0); + TEST_ASSERT_NOT_NULL(s); + + uint8_t data[] = "Hello, QUIC Stream!"; + int rc = quic_stream_write(s, data, sizeof(data)); + TEST_ASSERT_EQUAL(0, rc); + TEST_ASSERT_EQUAL(sizeof(data), s->recv_buf_len); + + uint8_t read_buf[64]; + ssize_t n = quic_stream_read(s, read_buf, sizeof(read_buf)); + TEST_ASSERT_EQUAL(sizeof(data), n); + TEST_ASSERT_EQUAL_MEMORY(data, read_buf, sizeof(data)); + TEST_ASSERT_EQUAL(0, s->recv_buf_len); + + quic_connection_destroy(conn); +} + +/** @test QUIC 流 FIN 标志 */ +void test_quic_stream_fin(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(12345, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_stream_t *s = quic_stream_get_or_create(conn, 0); + TEST_ASSERT_NOT_NULL(s); + TEST_ASSERT_FALSE(s->local_fin); + + quic_stream_set_fin(s); + TEST_ASSERT_TRUE(s->local_fin); + + quic_connection_destroy(conn); +} + +/** @test QUIC 流销毁 */ +void test_quic_stream_destroy(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(12345, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_stream_t *s = quic_stream_get_or_create(conn, 4); + TEST_ASSERT_NOT_NULL(s); + + quic_stream_destroy(conn, s); + quic_stream_t *found = quic_stream_find(conn, 4); + TEST_ASSERT_NULL(found); + + quic_connection_destroy(conn); +} + +/** @test QUIC 流:多个流 */ +void test_quic_multiple_streams(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(12345, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_stream_t *s0 = quic_stream_get_or_create(conn, 0); + quic_stream_t *s4 = quic_stream_get_or_create(conn, 4); + quic_stream_t *s8 = quic_stream_get_or_create(conn, 8); + + TEST_ASSERT_NOT_NULL(s0); + TEST_ASSERT_NOT_NULL(s4); + TEST_ASSERT_NOT_NULL(s8); + + TEST_ASSERT(s0 != s4); + TEST_ASSERT(s4 != s8); + + quic_connection_destroy(conn); +} + +/** @test QUIC 流:写入大数据 */ +void test_quic_stream_large_write(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(12345, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_stream_t *s = quic_stream_get_or_create(conn, 0); + TEST_ASSERT_NOT_NULL(s); + + uint8_t *large_data = (uint8_t *)malloc(10000); + TEST_ASSERT_NOT_NULL(large_data); + memset(large_data, 0xAB, 10000); + + int rc = quic_stream_write(s, large_data, 10000); + TEST_ASSERT_EQUAL(0, rc); + TEST_ASSERT_EQUAL(10000, s->recv_buf_len); + + uint8_t *read_buf = (uint8_t *)malloc(10000); + TEST_ASSERT_NOT_NULL(read_buf); + ssize_t n = quic_stream_read(s, read_buf, 10000); + TEST_ASSERT_EQUAL(10000, n); + TEST_ASSERT_EQUAL_UINT8_ARRAY_MESSAGE(large_data, read_buf, 10000, "大数据读写不一致"); + + free(large_data); + free(read_buf); + quic_connection_destroy(conn); +} + +/* ===== QUIC 连接管理测试 ===== */ + +/** @test QUIC 连接创建和销毁 */ +void test_quic_connection_create_destroy(void) { + struct sockaddr_storage addr = {0}; + ((struct sockaddr_in *)&addr)->sin_family = AF_INET; + + quic_connection_t *conn = quic_connection_create(99999, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + TEST_ASSERT_EQUAL(99999, conn->conn_id); + TEST_ASSERT_FALSE(conn->handshake_complete); + TEST_ASSERT_FALSE(conn->closed); + + quic_connection_destroy(conn); +} + +/** @test QUIC 连接查找 */ +void test_quic_find_connection(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(11111, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + quic_connection_t *found = quic_find_connection(11111); + TEST_ASSERT_NOT_NULL(found); + TEST_ASSERT_EQUAL_PTR(conn, found); + + quic_connection_t *not_found = quic_find_connection(99999); + TEST_ASSERT_NULL(not_found); + + quic_connection_destroy(conn); +} + +/** @test QUIC 连接计数 */ +void test_quic_connection_count(void) { + size_t before = quic_get_connection_count(); + + struct sockaddr_storage addr = {0}; + quic_connection_t *c1 = quic_connection_create(100, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(c1); + TEST_ASSERT_EQUAL(before + 1, quic_get_connection_count()); + + quic_connection_t *c2 = quic_connection_create(200, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(c2); + TEST_ASSERT_EQUAL(before + 2, quic_get_connection_count()); + + quic_connection_destroy(c1); + TEST_ASSERT_EQUAL(before + 1, quic_get_connection_count()); + + quic_connection_destroy(c2); + TEST_ASSERT_EQUAL(before, quic_get_connection_count()); +} + +/** @test QUIC 连接 ID 生成 */ +void test_quic_generate_conn_id(void) { + uint64_t id1 = quic_generate_conn_id(); + uint64_t id2 = quic_generate_conn_id(); + + TEST_ASSERT_NOT_EQUAL(id1, id2); + TEST_ASSERT_NOT_EQUAL(0, id1); + TEST_ASSERT_NOT_EQUAL(0, id2); +} + +/** @test QUIC 超时连接清理 */ +void test_quic_cleanup_timeout_connections(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(300, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + size_t count_before = quic_get_connection_count(); + TEST_ASSERT_GREATER_THAN(0, count_before); + + /* 设置最后活动时间为很久以前 */ + conn->last_activity = 0; + + /* 清理超时连接(使用 1ms 超时) */ + quic_cleanup_timeout_connections(1); + + /* 连接应该被清理 */ + quic_connection_t *found = quic_find_connection(300); + TEST_ASSERT_NULL(found); +} + +/* ===== HTTP/3 会话管理测试 ===== */ + +/** @test HTTP/3 会话创建和销毁 */ +void test_http3_session_create_destroy(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(400, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + http3_session_t *session = http3_session_create(conn); + TEST_ASSERT_NOT_NULL(session); + TEST_ASSERT_EQUAL(conn, session->conn); + TEST_ASSERT_EQUAL(HTTP3_DEFAULT_MAX_FIELD_SECTION_SIZE, + session->max_field_section_size); + + http3_session_destroy(session); + quic_connection_destroy(conn); +} + +/** @test HTTP/3 多次会话创建 */ +void test_http3_session_multiple(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn1 = quic_connection_create(500, + COCOON_INVALID_SOCKET, &addr); + quic_connection_t *conn2 = quic_connection_create(600, + COCOON_INVALID_SOCKET, &addr); + + http3_session_t *s1 = http3_session_create(conn1); + http3_session_t *s2 = http3_session_create(conn2); + + TEST_ASSERT_NOT_NULL(s1); + TEST_ASSERT_NOT_NULL(s2); + TEST_ASSERT(s1 != s2); + + http3_session_destroy(s1); + http3_session_destroy(s2); + quic_connection_destroy(conn1); + quic_connection_destroy(conn2); +} + +/** @test HTTP/3 会话创建:NULL 连接 */ +void test_http3_session_create_null(void) { + http3_session_t *session = http3_session_create(NULL); + TEST_ASSERT_NULL(session); +} + +/* ===== 错误码测试 ===== */ + +/** @test HTTP/3 错误码值 */ +void test_http3_error_codes(void) { + TEST_ASSERT_EQUAL(0x0100, HTTP3_NO_ERROR); + TEST_ASSERT_EQUAL(0x0101, HTTP3_GENERAL_PROTOCOL_ERROR); + TEST_ASSERT_EQUAL(0x0102, HTTP3_INTERNAL_ERROR); + TEST_ASSERT_EQUAL(0x0103, HTTP3_STREAM_CREATION_ERROR); + TEST_ASSERT_EQUAL(0x0110, HTTP3_VERSION_FALLBACK); +} + +/* ===== 时间戳测试 ===== */ + +/** @test 当前时间戳 */ +void test_quic_current_time_ms(void) { + uint64_t t1 = quic_current_time_ms(); + TEST_ASSERT_GREATER_THAN(0, t1); + + /* 应该是单调递增的(至少不减) */ + uint64_t t2 = quic_current_time_ms(); + TEST_ASSERT_GREATER_OR_EQUAL(t1, t2); +} + +/* ===== 边界条件测试 ===== */ + +/** @test varint 编解码边界值 */ +void test_varint_boundary_values(void) { + uint64_t boundaries[] = {0, 63, 64, 16383, 16384, 1073741823ULL, + 1073741824ULL, 4611686018427387903ULL}; + int num = sizeof(boundaries) / sizeof(boundaries[0]); + + for (int i = 0; i < num; i++) { + uint8_t buf[8]; + size_t enc_n = http3_encode_varint(boundaries[i], buf); + TEST_ASSERT_GREATER_THAN(0, enc_n); + + uint64_t decoded = 0; + int dec_n = http3_decode_varint(buf, enc_n, &decoded); + TEST_ASSERT_EQUAL(enc_n, (size_t)dec_n); + TEST_ASSERT_EQUAL(boundaries[i], decoded); + } +} + +/** @test QUIC 流:NULL 参数处理 */ +void test_quic_stream_null_params(void) { + int rc = quic_stream_write(NULL, (uint8_t *)"test", 4); + TEST_ASSERT_EQUAL(-1, rc); + + ssize_t n = quic_stream_read(NULL, NULL, 0); + TEST_ASSERT_EQUAL(-1, n); + + quic_stream_set_fin(NULL); /* 不应崩溃 */ +} + +/** @test QUIC 连接:NULL 参数处理 */ +void test_quic_connection_null(void) { + quic_connection_destroy(NULL); /* 不应崩溃 */ + + quic_stream_t *s = quic_stream_get_or_create(NULL, 0); + TEST_ASSERT_NULL(s); + + quic_stream_t *f = quic_stream_find(NULL, 0); + TEST_ASSERT_NULL(f); +} + +/** @test http3_close_connection: NULL */ +void test_http3_close_connection_null(void) { + http3_close_connection(NULL, HTTP3_NO_ERROR); /* 不应崩溃 */ +} + +/** @test http3_send_error: NULL session */ +void test_http3_send_error_null(void) { + http3_send_error(NULL, 0, 404, "Not Found"); /* 不应崩溃 */ +} + +/** @test http3_send_settings: NULL */ +void test_http3_send_settings_null(void) { + int rc = http3_send_settings(NULL); + TEST_ASSERT_EQUAL(-1, rc); +} + +/** @test http3_process_datagram: NULL params */ +void test_http3_process_datagram_null(void) { + uint8_t data[] = "test"; + struct sockaddr_storage addr = {0}; + http3_process_datagram(COCOON_INVALID_SOCKET, NULL, 0, &addr); + http3_process_datagram(COCOON_INVALID_SOCKET, data, 4, NULL); +} + +/** @test quic_send_datagram: invalid socket */ +void test_quic_send_datagram_invalid(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(700, + COCOON_INVALID_SOCKET, &addr); + TEST_ASSERT_NOT_NULL(conn); + + uint8_t data[] = "test"; + int rc = quic_send_datagram(conn, data, 4); + TEST_ASSERT_EQUAL(-1, rc); + + quic_connection_destroy(conn); +} + +/** @test 帧解码:数据不足 */ +void test_parse_frame_insufficient_data(void) { + struct sockaddr_storage addr = {0}; + quic_connection_t *conn = quic_connection_create(800, + COCOON_INVALID_SOCKET, &addr); + quic_stream_t *stream = quic_stream_get_or_create(conn, 0); + + /* 写入不完整的帧数据 */ + uint8_t partial[] = {0x00}; /* 只有类型,没有长度 */ + quic_stream_write(stream, partial, 1); + + uint64_t ft = 0; + const uint8_t *payload = NULL; + size_t payload_len = 0; + int rc = http3_parse_frame(stream, &ft, &payload, &payload_len); + TEST_ASSERT_EQUAL(1, rc); /* 数据不足 */ + + quic_connection_destroy(conn); +} + +/* ===== main ===== */ + +int main(void) { + UNITY_BEGIN(); + + /* Variable-length integer encode */ + RUN_TEST(test_varint_encode_zero); + RUN_TEST(test_varint_encode_63); + RUN_TEST(test_varint_encode_64); + RUN_TEST(test_varint_encode_16383); + RUN_TEST(test_varint_encode_16384); + RUN_TEST(test_varint_encode_4byte_max); + RUN_TEST(test_varint_encode_8byte); + RUN_TEST(test_varint_encode_max); + RUN_TEST(test_varint_encode_null_buf); + + /* Variable-length integer decode */ + RUN_TEST(test_varint_decode_zero); + RUN_TEST(test_varint_decode_63); + RUN_TEST(test_varint_decode_64); + RUN_TEST(test_varint_decode_16383); + RUN_TEST(test_varint_decode_4byte); + RUN_TEST(test_varint_decode_insufficient_data); + RUN_TEST(test_varint_decode_null); + RUN_TEST(test_varint_roundtrip); + RUN_TEST(test_varint_boundary_values); + + /* Frame header */ + RUN_TEST(test_frame_header_size); + RUN_TEST(test_frame_header_encode_decode); + RUN_TEST(test_frame_header_decode_insufficient); + RUN_TEST(test_frame_header_decode_null); + + /* Full frame encoding */ + RUN_TEST(test_full_frame_encode); + RUN_TEST(test_encode_data_frame); + RUN_TEST(test_encode_headers_frame); + RUN_TEST(test_encode_settings_frame); + RUN_TEST(test_encode_goaway_frame); + + /* QPACK encode */ + RUN_TEST(test_qpack_encode_static_match_method_get); + RUN_TEST(test_qpack_encode_static_match_path); + RUN_TEST(test_qpack_encode_static_match_scheme); + RUN_TEST(test_qpack_encode_literal_field); + RUN_TEST(test_qpack_encode_name_ref_value_literal); + + /* QPACK decode */ + RUN_TEST(test_qpack_decode_indexed_header); + RUN_TEST(test_qpack_decode_literal_header); + RUN_TEST(test_qpack_decode_invalid); + + /* QPACK roundtrip */ + RUN_TEST(test_qpack_roundtrip); + RUN_TEST(test_qpack_encode_request_headers); + RUN_TEST(test_qpack_decode_request_headers); + RUN_TEST(test_qpack_encode_response_headers); + RUN_TEST(test_qpack_decode_response_headers); + + /* QUIC stream management */ + RUN_TEST(test_quic_stream_create_and_find); + RUN_TEST(test_quic_stream_write_read); + RUN_TEST(test_quic_stream_fin); + RUN_TEST(test_quic_stream_destroy); + RUN_TEST(test_quic_multiple_streams); + RUN_TEST(test_quic_stream_large_write); + RUN_TEST(test_quic_stream_null_params); + + /* QUIC connection management */ + RUN_TEST(test_quic_connection_create_destroy); + RUN_TEST(test_quic_find_connection); + RUN_TEST(test_quic_connection_count); + RUN_TEST(test_quic_generate_conn_id); + RUN_TEST(test_quic_cleanup_timeout_connections); + RUN_TEST(test_quic_connection_null); + + /* HTTP/3 session management */ + RUN_TEST(test_http3_session_create_destroy); + RUN_TEST(test_http3_session_multiple); + RUN_TEST(test_http3_session_create_null); + + /* Error codes */ + RUN_TEST(test_http3_error_codes); + + /* Time utilities */ + RUN_TEST(test_quic_current_time_ms); + + /* Boundary / NULL tests */ + RUN_TEST(test_http3_close_connection_null); + RUN_TEST(test_http3_send_error_null); + RUN_TEST(test_http3_send_settings_null); + RUN_TEST(test_http3_process_datagram_null); + RUN_TEST(test_quic_send_datagram_invalid); + RUN_TEST(test_parse_frame_insufficient_data); + + return UNITY_END(); +} diff --git a/tests/unit/test_load_balance.c b/tests/unit/test_load_balance.c new file mode 100644 index 0000000..7e8dbc8 --- /dev/null +++ b/tests/unit/test_load_balance.c @@ -0,0 +1,1212 @@ +/** + * @file test_load_balance.c - 负载均衡模块单元测试 + * @brief 使用 Unity 框架测试所有负载均衡算法 + * + * 测试覆盖: + * - MurmurHash3 哈希函数已知值验证 + * - 一致性哈希环构建与查找 + * - 一致性哈希 key→后端映射稳定性 + * - 最少连接选择正确性 + * - EWMA 计算验证 + * - 加权响应时间权重影响 + * - 随机算法分布 + * - 边界条件和错误处理 + * + * @author xfy + */ + +#include "unity.h" +#include "load_balance.h" +#include +#include +#include + +/* ===== 测试辅助函数 ===== */ + +/** + * @brief 创建简单的代理规则(n 个后端)用于测试 + */ +static void make_test_rule(cocoon_proxy_rule_t *rule, size_t n) { + memset(rule, 0, sizeof(*rule)); + rule->backend_count = n; + for (size_t i = 0; i < n; i++) { + snprintf(rule->backends[i].target_host, + sizeof(rule->backends[i].target_host), + "backend%zu", i); + rule->backends[i].target_port = (uint16_t)(8000 + i); + rule->backends[i].healthy = true; + rule->backends[i].weight = 1; + rule->backends[i].current_weight = 0; + } +} + +/** + * @brief 创建带权重的代理规则 + */ +static void make_weighted_rule(cocoon_proxy_rule_t *rule, size_t n, + const uint32_t *weights) { + make_test_rule(rule, n); + for (size_t i = 0; i < n; i++) { + rule->backends[i].weight = weights[i]; + } +} + +void setUp(void) { + /* 每次测试前重置随机种子,保证可重复 */ + srand(42); +} + +void tearDown(void) { } + +/* ===== 测试组 1: MurmurHash3 哈希函数 ===== */ + +/** + * @test test_hash_known_values + * @brief MurmurHash3 已知输入输出验证 + * + * 使用参考实现验证相同输入产生相同的哈希值。 + */ +void test_hash_known_values(void) { + /* 空字符串:seed=0 时 hash=0 */ + TEST_ASSERT_EQUAL_UINT32(0, lb_hash_key("", 0)); + + /* 常见字符串的已知哈希值(seed=0 的 MurmurHash3 x86 32-bit) */ + TEST_ASSERT_EQUAL_UINT32(0x248bfa47U, lb_hash_key("hello", 5)); + TEST_ASSERT_EQUAL_UINT32(0x149bbb7fU, lb_hash_key("hello, world", 12)); + TEST_ASSERT_EQUAL_UINT32(0x38d82c45U, lb_hash_key("19Jan2038at03:14:07UTC", 22)); + + /* 较长的字符串 */ + TEST_ASSERT_EQUAL_UINT32(0x2e4ff723U, + lb_hash_key("The quick brown fox jumps over the lazy dog", 43)); +} + +/** + * @test test_hash_empty_string_zero_length + * @brief 空字符串返回 0 + */ +void test_hash_empty_string_zero_length(void) { + TEST_ASSERT_EQUAL_UINT32(0, lb_hash_key("", 0)); +} + +/** + * @test test_hash_null_returns_zero + * @brief NULL 键返回 0 + */ +void test_hash_null_returns_zero(void) { + TEST_ASSERT_EQUAL_UINT32(0, lb_hash_key(NULL, 0)); +} + +/** + * @test test_hash_consistency + * @brief 相同输入必须产生相同输出(确定性) + */ +void test_hash_consistency(void) { + const char *key = "/api/users/12345"; + size_t len = strlen(key); + uint32_t h1 = lb_hash_key(key, len); + uint32_t h2 = lb_hash_key(key, len); + uint32_t h3 = lb_hash_key(key, len); + + TEST_ASSERT_EQUAL_UINT32(h1, h2); + TEST_ASSERT_EQUAL_UINT32(h2, h3); +} + +/** + * @test test_hash_different_inputs + * @brief 不同输入产生不同哈希值(碰撞概率极低) + */ +void test_hash_different_inputs(void) { + uint32_t h1 = lb_hash_key("/api/a", 6); + uint32_t h2 = lb_hash_key("/api/b", 6); + uint32_t h3 = lb_hash_key("/api/c", 6); + + TEST_ASSERT_UINT32_WITHIN(0xFFFFFFFFU, 0, h1 ^ h2); /* 非零差异 */ + TEST_ASSERT(h1 != h2); + TEST_ASSERT(h2 != h3); + TEST_ASSERT(h1 != h3); +} + +/** + * @test test_hash_binary_data + * @brief 二进制数据也能正确处理 + */ +void test_hash_binary_data(void) { + const uint8_t data[] = {0x00, 0x01, 0x02, 0x03, 0xFF, 0xFE, 0xFD, 0xFC}; + uint32_t h = lb_hash_key((const char *)data, sizeof(data)); + TEST_ASSERT_NOT_EQUAL(0, h); +} + +/* ===== 测试组 2: lb_init 和 lb_destroy ===== */ + +/** + * @test test_init_sets_algorithm + * @brief 初始化设置正确的算法 + */ +void test_init_sets_algorithm(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + TEST_ASSERT_EQUAL(COCOON_LB_LEAST_CONNECTIONS, lb.algorithm); + lb_destroy(&lb); +} + +/** + * @test test_init_default_alpha + * @brief 初始化设置默认 alpha 值 + */ +void test_init_default_alpha(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + TEST_ASSERT_EQUAL_UINT32(80, lb.alpha); + lb_destroy(&lb); +} + +/** + * @test test_init_clears_stats + * @brief 初始化清零所有统计 + */ +void test_init_clears_stats(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_RANDOM); + for (size_t i = 0; i < COCOON_MAX_PROXY_BACKENDS; i++) { + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[i].active_connections); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[i].total_requests); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[i].total_failures); + TEST_ASSERT_EQUAL_UINT64(0, lb.stats[i].ewma_response_time_us); + } + lb_destroy(&lb); +} + +/** + * @test test_init_clears_hash_ring + * @brief 初始化清零哈希环 + */ +void test_init_clears_hash_ring(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_CONSISTENT_HASH); + TEST_ASSERT_EQUAL_size_t(0, lb.hash_ring.node_count); + TEST_ASSERT_FALSE(lb.hash_ring.initialized); + lb_destroy(&lb); +} + +/** + * @test test_destroy_cleans_up + * @brief 销毁后结构体应被清零 + */ +void test_destroy_cleans_up(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + lb_destroy(&lb); + /* lb.algorithm 应被设为 0(ROUND_ROBIN 的值) */ + TEST_ASSERT_EQUAL(COCOON_LB_ROUND_ROBIN, lb.algorithm); +} + +/* ===== 测试组 3: lb_algorithm_name ===== */ + +/** + * @test test_algorithm_name_all + * @brief 所有算法都有名称 + */ +void test_algorithm_name_all(void) { + TEST_ASSERT_EQUAL_STRING("round_robin", + lb_algorithm_name(COCOON_LB_ROUND_ROBIN)); + TEST_ASSERT_EQUAL_STRING("least_connections", + lb_algorithm_name(COCOON_LB_LEAST_CONNECTIONS)); + TEST_ASSERT_EQUAL_STRING("weighted_response", + lb_algorithm_name(COCOON_LB_WEIGHTED_RESPONSE)); + TEST_ASSERT_EQUAL_STRING("consistent_hash", + lb_algorithm_name(COCOON_LB_CONSISTENT_HASH)); + TEST_ASSERT_EQUAL_STRING("random", + lb_algorithm_name(COCOON_LB_RANDOM)); +} + +/** + * @test test_algorithm_name_unknown + * @brief 未知算法返回 "unknown" + */ +void test_algorithm_name_unknown(void) { + TEST_ASSERT_EQUAL_STRING("unknown", + lb_algorithm_name((cocoon_lb_algorithm_t)999)); +} + +/* ===== 测试组 4: 一致性哈希环 ===== */ + +/** + * @test test_build_hash_ring_creates_nodes + * @brief 构建哈希环创建正确数量的虚拟节点 + */ +void test_build_hash_ring_creates_nodes(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + TEST_ASSERT_TRUE(ring.initialized); + TEST_ASSERT_EQUAL_size_t(3 * COCOON_HASH_RING_SIZE, ring.node_count); +} + +/** + * @test test_build_hash_ring_empty_rule + * @brief 空规则构建空哈希环 + */ +void test_build_hash_ring_empty_rule(void) { + cocoon_proxy_rule_t rule; + memset(&rule, 0, sizeof(rule)); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + TEST_ASSERT_FALSE(ring.initialized); + TEST_ASSERT_EQUAL_size_t(0, ring.node_count); +} + +/** + * @test test_build_hash_ring_skips_unhealthy + * @brief 构建哈希环跳过不健康后端 + */ +void test_build_hash_ring_skips_unhealthy(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 4); + rule.backends[1].healthy = false; /**< backend1 不健康 */ + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + TEST_ASSERT_TRUE(ring.initialized); + /**< 3 个健康后端 × 512 虚拟节点 */ + TEST_ASSERT_EQUAL_size_t(3 * COCOON_HASH_RING_SIZE, ring.node_count); +} + +/** + * @test test_build_hash_ring_all_unhealthy + * @brief 所有后端不健康时构建空环 + */ +void test_build_hash_ring_all_unhealthy(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + for (size_t i = 0; i < 3; i++) { + rule.backends[i].healthy = false; + } + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + TEST_ASSERT_FALSE(ring.initialized); + TEST_ASSERT_EQUAL_size_t(0, ring.node_count); +} + +/** + * @test test_hash_ring_nodes_sorted + * @brief 哈希环虚拟节点按哈希值排序 + */ +void test_hash_ring_nodes_sorted(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + for (size_t i = 1; i < ring.node_count; i++) { + TEST_ASSERT_MESSAGE( + ring.nodes[i - 1].node_hash <= ring.nodes[i].node_hash, + "哈希环虚拟节点未按升序排列" + ); + } +} + +/** + * @test test_hash_ring_backend_coverage + * @brief 所有健康后端都有虚拟节点 + */ +void test_hash_ring_backend_coverage(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + bool has_backend[3] = {false, false, false}; + for (size_t i = 0; i < ring.node_count; i++) { + if (ring.nodes[i].backend_index < 3) { + has_backend[ring.nodes[i].backend_index] = true; + } + } + TEST_ASSERT_TRUE(has_backend[0]); + TEST_ASSERT_TRUE(has_backend[1]); + TEST_ASSERT_TRUE(has_backend[2]); +} + +/** + * @test test_pick_from_ring_basic + * @brief 从哈希环选取返回有效后端索引 + */ +void test_pick_from_ring_basic(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + uint32_t h = lb_hash_key("/api/users", 10); + size_t idx = lb_pick_from_ring(&ring, h); + TEST_ASSERT(idx < 3); +} + +/** + * @test test_pick_from_ring_wraparound + * @brief 哈希值大于最大节点时正确环绕 + */ +void test_pick_from_ring_wraparound(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 2); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + /**< 使用最大可能的哈希值,强制环绕 */ + size_t idx = lb_pick_from_ring(&ring, 0xFFFFFFFFU); + TEST_ASSERT(idx < 2); +} + +/** + * @test test_pick_from_ring_zero_hash + * @brief 哈希值为 0 时正确返回第一个节点 + */ +void test_pick_from_ring_zero_hash(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 2); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + size_t idx = lb_pick_from_ring(&ring, 0); + TEST_ASSERT(idx < 2); +} + +/** + * @test test_pick_from_ring_deterministic + * @brief 相同哈希值总是选取相同后端 + */ +void test_pick_from_ring_deterministic(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + uint32_t h = lb_hash_key("/static/file.css", 16); + size_t idx1 = lb_pick_from_ring(&ring, h); + size_t idx2 = lb_pick_from_ring(&ring, h); + size_t idx3 = lb_pick_from_ring(&ring, h); + + TEST_ASSERT_EQUAL_size_t(idx1, idx2); + TEST_ASSERT_EQUAL_size_t(idx2, idx3); +} + +/** + * @test test_pick_from_ring_empty + * @brief 空环返回 0(安全回退) + */ +void test_pick_from_ring_empty(void) { + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + + size_t idx = lb_pick_from_ring(&ring, 12345); + TEST_ASSERT_EQUAL_size_t(0, idx); +} + +/* ===== 测试组 5: 一致性哈希映射稳定性 ===== */ + +/** + * @test test_consistent_hash_stability + * @brief 相同 key 总是映射到相同后端 + */ +void test_consistent_hash_stability(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_CONSISTENT_HASH); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 4); + + lb_build_hash_ring(&lb.hash_ring, &rule); + + /**< 多次选择相同 key */ + int idx1 = lb_select_backend(&lb, &rule, "/api/resource/42"); + int idx2 = lb_select_backend(&lb, &rule, "/api/resource/42"); + int idx3 = lb_select_backend(&lb, &rule, "/api/resource/42"); + + TEST_ASSERT_EQUAL_INT(idx1, idx2); + TEST_ASSERT_EQUAL_INT(idx2, idx3); + TEST_ASSERT(idx1 >= 0 && idx1 < 4); + + lb_destroy(&lb); +} + +/** + * @test test_consistent_hash_different_keys + * @brief 不同 key 可以映射到不同后端 + */ +void test_consistent_hash_different_keys(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_CONSISTENT_HASH); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 4); + + lb_build_hash_ring(&lb.hash_ring, &rule); + + /**< 使用多个不同的 key */ + int idx1 = lb_select_backend(&lb, &rule, "/api/users"); + int idx2 = lb_select_backend(&lb, &rule, "/api/products"); + int idx3 = lb_select_backend(&lb, &rule, "/api/orders"); + + TEST_ASSERT(idx1 >= 0 && idx1 < 4); + TEST_ASSERT(idx2 >= 0 && idx2 < 4); + TEST_ASSERT(idx3 >= 0 && idx3 < 4); + + lb_destroy(&lb); +} + +/** + * @test test_consistent_hash_null_key_fallback + * @brief NULL key 回退到随机选择 + */ +void test_consistent_hash_null_key_fallback(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_CONSISTENT_HASH); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + lb_build_hash_ring(&lb.hash_ring, &rule); + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT(idx >= 0 && idx < 3); + + lb_destroy(&lb); +} + +/** + * @test test_consistent_hash_distribution + * @brief 多个 key 的分布检查(不应全部映射到同一后端) + */ +void test_consistent_hash_distribution(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_CONSISTENT_HASH); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 4); + + lb_build_hash_ring(&lb.hash_ring, &rule); + + /**< 使用大量不同的 key 测试分布 */ + int count[4] = {0, 0, 0, 0}; + char key[64]; + for (int i = 0; i < 1000; i++) { + snprintf(key, sizeof(key), "/api/item/%d", i); + int idx = lb_select_backend(&lb, &rule, key); + TEST_ASSERT(idx >= 0 && idx < 4); + count[idx]++; + } + + /**< 每个后端都应被分配到一些 key(分布应相对均匀) */ + for (int i = 0; i < 4; i++) { + TEST_ASSERT_MESSAGE(count[i] > 50, + "一致性哈希分布不均匀,某个后端分配过少"); + } + + lb_destroy(&lb); +} + +/* ===== 测试组 6: 最少连接算法 ===== */ + +/** + * @test test_least_connections_selects_min + * @brief 最少连接选择连接数最少的后端 + */ +void test_least_connections_selects_min(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + /**< 设置不同的活跃连接数 */ + lb.stats[0].active_connections = 5; + lb.stats[1].active_connections = 2; /**< 最少 */ + lb.stats[2].active_connections = 8; + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_least_connections_all_zero + * @brief 所有连接数为 0 时选择第一个健康后端 + */ +void test_least_connections_all_zero(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(0, idx); + + lb_destroy(&lb); +} + +/** + * @test test_least_connections_skips_unhealthy + * @brief 最少连接跳过不健康后端 + */ +void test_least_connections_skips_unhealthy(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + rule.backends[0].healthy = false; + + /**< backend0 连接数最少,但不健康 */ + lb.stats[0].active_connections = 0; + lb.stats[1].active_connections = 3; + lb.stats[2].active_connections = 5; + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_least_connections_no_healthy + * @brief 没有健康后端返回 -1 + */ +void test_least_connections_no_healthy(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + for (size_t i = 0; i < 3; i++) { + rule.backends[i].healthy = false; + } + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(-1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_least_connections_tie_break + * @brief 连接数相同时选择索引小的(确定性) + */ +void test_least_connections_tie_break(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + /**< 所有后端连接数相同 */ + lb.stats[0].active_connections = 3; + lb.stats[1].active_connections = 3; + lb.stats[2].active_connections = 3; + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(0, idx); /**< 第一个达到最小值的后端 */ + + lb_destroy(&lb); +} + +/* ===== 测试组 7: 统计更新 ===== */ + +/** + * @test test_stats_request_start + * @brief 请求开始正确更新连接数和请求数 + */ +void test_stats_request_start(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + lb_update_stats_request_start(&lb, 0); + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[0].active_connections); + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[0].total_requests); + + lb_update_stats_request_start(&lb, 0); + TEST_ASSERT_EQUAL_UINT32(2, lb.stats[0].active_connections); + TEST_ASSERT_EQUAL_UINT32(2, lb.stats[0].total_requests); + + lb_destroy(&lb); +} + +/** + * @test test_stats_request_end_success + * @brief 请求成功正确更新统计 + */ +void test_stats_request_end_success(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + lb_update_stats_request_start(&lb, 0); + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[0].active_connections); + + lb_update_stats_request_end(&lb, 0, true, 1000); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[0].active_connections); + TEST_ASSERT_EQUAL_UINT64(1000, lb.stats[0].last_response_time_us); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[0].total_failures); + + lb_destroy(&lb); +} + +/** + * @test test_stats_request_end_failure + * @brief 请求失败正确记录失败数 + */ +void test_stats_request_end_failure(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + lb_update_stats_request_start(&lb, 0); + lb_update_stats_request_end(&lb, 0, false, 500); + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[0].total_failures); + + lb_update_stats_request_start(&lb, 0); + lb_update_stats_request_end(&lb, 0, false, 600); + TEST_ASSERT_EQUAL_UINT32(2, lb.stats[0].total_failures); + + lb_destroy(&lb); +} + +/** + * @test test_stats_end_without_start + * @brief 未调用 start 直接调用 end 不会下溢 + */ +void test_stats_end_without_start(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + /**< active_connections 已经是 0 */ + lb_update_stats_request_end(&lb, 0, true, 1000); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[0].active_connections); + + lb_destroy(&lb); +} + +/** + * @test test_stats_multiple_backends + * @brief 多后端统计互不干扰 + */ +void test_stats_multiple_backends(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + lb_update_stats_request_start(&lb, 0); + lb_update_stats_request_start(&lb, 1); + lb_update_stats_request_start(&lb, 1); + + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[0].active_connections); + TEST_ASSERT_EQUAL_UINT32(2, lb.stats[1].active_connections); + + lb_destroy(&lb); +} + +/* ===== 测试组 8: EWMA 计算 ===== */ + +/** + * @test test_ewma_first_update + * @brief 首次 EWMA 更新直接使用当前值 + */ +void test_ewma_first_update(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + lb_update_stats_request_end(&lb, 0, true, 1000); + /**< 首次更新:ewma 直接设为当前值 */ + TEST_ASSERT_EQUAL_UINT64(1000, lb.stats[0].ewma_response_time_us); + + lb_destroy(&lb); +} + +/** + * @test test_ewma_subsequent_updates + * @brief 后续 EWMA 更新应用平滑公式 + * + * alpha = 80,即 new = 0.8 * current + 0.2 * old + * old = 1000, current = 2000 + * new = (80 * 2000 + 20 * 1000) / 100 = (160000 + 20000) / 100 = 1800 + */ +void test_ewma_subsequent_updates(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + lb_update_stats_request_end(&lb, 0, true, 1000); /**< 首次:ewma=1000 */ + lb_update_stats_request_end(&lb, 0, true, 2000); /**< 第二次:ewma=1800 */ + + TEST_ASSERT_EQUAL_UINT64(1800, lb.stats[0].ewma_response_time_us); + + lb_destroy(&lb); +} + +/** + * @test test_ewma_converges + * @brief EWMA 向稳定值收敛 + * + * 多次更新相同值后,EWMA 应接近该值。 + */ +void test_ewma_converges(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + lb_update_stats_request_end(&lb, 0, true, 1000); /**< 首次 */ + for (int i = 0; i < 20; i++) { + lb_update_stats_request_end(&lb, 0, true, 1000); + } + + /**< 经过 20 次相同值更新,EWMA 应非常接近 1000 */ + TEST_ASSERT_UINT64_WITHIN(10, 1000, + lb.stats[0].ewma_response_time_us); + + lb_destroy(&lb); +} + +/** + * @test test_ewma_alpha_effect + * @brief alpha 值影响平滑程度 + * + * alpha 较大时(接近 100),对新值更敏感。 + */ +void test_ewma_alpha_effect(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + lb.alpha = 90; /**< 更高 alpha,更敏感 */ + + lb_update_stats_request_end(&lb, 0, true, 1000); + lb_update_stats_request_end(&lb, 0, true, 2000); + + /**< alpha=90: new = (90*2000 + 10*1000)/100 = 1900 */ + TEST_ASSERT_EQUAL_UINT64(1900, lb.stats[0].ewma_response_time_us); + + lb.alpha = 50; /**< 更低 alpha,更平滑 */ + lb_update_stats_request_end(&lb, 0, true, 2000); + + /**< alpha=50: new = (50*2000 + 50*1900)/100 = 1950 */ + TEST_ASSERT_EQUAL_UINT64(1950, lb.stats[0].ewma_response_time_us); + + lb_destroy(&lb); +} + +/* ===== 测试组 9: 加权响应时间选择 ===== */ + +/** + * @test test_weighted_response_selects_lowest_ratio + * @brief 选择 EWMA/weight 比值最小的后端 + */ +void test_weighted_response_selects_lowest_ratio(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + cocoon_proxy_rule_t rule; + uint32_t weights[] = {1, 1, 1}; + make_weighted_rule(&rule, 3, weights); + + /**< backend1 的 EWMA 最低 */ + lb.stats[0].ewma_response_time_us = 5000; + lb.stats[1].ewma_response_time_us = 1000; /**< 最优 */ + lb.stats[2].ewma_response_time_us = 3000; + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_weighted_response_weight_influence + * @brief 权重影响选择(高权重可抵消高 EWMA) + * + * backend0: ewma=5000, weight=5 => score=1000 + * backend1: ewma=3000, weight=1 => score=3000 + * backend2: ewma=4000, weight=2 => score=2000 + * + * backend0 应被选中(最低 score) + */ +void test_weighted_response_weight_influence(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + cocoon_proxy_rule_t rule; + uint32_t weights[] = {5, 1, 2}; + make_weighted_rule(&rule, 3, weights); + + lb.stats[0].ewma_response_time_us = 5000; /**< weight=5, score=1000 */ + lb.stats[1].ewma_response_time_us = 3000; /**< weight=1, score=3000 */ + lb.stats[2].ewma_response_time_us = 4000; /**< weight=2, score=2000 */ + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(0, idx); + + lb_destroy(&lb); +} + +/** + * @test test_weighted_response_zero_weight + * @brief weight=0 时按 weight=1 处理(防止除零) + */ +void test_weighted_response_zero_weight(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + cocoon_proxy_rule_t rule; + uint32_t weights[] = {0, 1}; /**< backend0 weight=0 */ + make_weighted_rule(&rule, 2, weights); + + lb.stats[0].ewma_response_time_us = 1000; + lb.stats[1].ewma_response_time_us = 2000; + + /**< backend0: 1000/1=1000, backend1: 2000/1=2000 */ + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(0, idx); + + lb_destroy(&lb); +} + +/* ===== 测试组 10: 随机算法 ===== */ + +/** + * @test test_random_returns_valid_index + * @brief 随机选择返回有效后端索引 + */ +void test_random_returns_valid_index(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_RANDOM); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 5); + + for (int i = 0; i < 50; i++) { + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT(idx >= 0 && idx < 5); + } + + lb_destroy(&lb); +} + +/** + * @test test_random_distribution + * @brief 随机选择分布大致均匀 + */ +void test_random_distribution(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_RANDOM); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 4); + + int count[4] = {0, 0, 0, 0}; + for (int i = 0; i < 4000; i++) { + int idx = lb_select_backend(&lb, &rule, NULL); + count[idx]++; + } + + /**< 每个后端应被选中约 1000 次,允许 ±200 偏差 */ + for (int i = 0; i < 4; i++) { + TEST_ASSERT_INT_WITHIN(300, 1000, count[i]); + } + + lb_destroy(&lb); +} + +/** + * @test test_random_no_healthy + * @brief 没有健康后端时返回 -1 + */ +void test_random_no_healthy(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_RANDOM); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + for (size_t i = 0; i < 3; i++) { + rule.backends[i].healthy = false; + } + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(-1, idx); + + lb_destroy(&lb); +} + +/* ===== 测试组 11: 边界条件和错误处理 ===== */ + +/** + * @test test_select_null_lb + * @brief NULL 负载均衡器返回 -1 + */ +void test_select_null_lb(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 2); + + int idx = lb_select_backend(NULL, &rule, NULL); + TEST_ASSERT_EQUAL_INT(-1, idx); +} + +/** + * @test test_select_null_rule + * @brief NULL 规则返回 -1 + */ +void test_select_null_rule(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_RANDOM); + + int idx = lb_select_backend(&lb, NULL, NULL); + TEST_ASSERT_EQUAL_INT(-1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_select_empty_rule + * @brief 空规则(无后端)返回 -1 + */ +void test_select_empty_rule(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_RANDOM); + + cocoon_proxy_rule_t rule; + memset(&rule, 0, sizeof(rule)); + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(-1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_update_stats_null_lb + * @brief NULL 负载均衡器更新统计不崩溃 + */ +void test_update_stats_null_lb(void) { + /**< 应安全返回,不崩溃 */ + lb_update_stats_request_start(NULL, 0); + lb_update_stats_request_end(NULL, 0, true, 100); + TEST_PASS_MESSAGE("NULL lb stats update handled safely"); +} + +/** + * @test test_round_robin_returns_neg_one + * @brief 轮询算法返回 -1(由调用者使用 proxy.c 的 SWW) + */ +void test_round_robin_returns_neg_one(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_ROUND_ROBIN); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(-1, idx); + + lb_destroy(&lb); +} + +/* ===== 测试组 12: 综合集成测试 ===== */ + +/** + * @test test_full_request_lifecycle_least_connections + * @brief 最少连接完整请求生命周期 + */ +void test_full_request_lifecycle_least_connections(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_LEAST_CONNECTIONS); + + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + /**< 初始选择:backend0(连接数都是 0,选第一个) */ + int idx1 = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(0, idx1); + + /**< 请求开始 */ + lb_update_stats_request_start(&lb, (size_t)idx1); + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[0].active_connections); + + /**< 第二次选择:backend1(backend0 已有 1 个连接) */ + int idx2 = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(1, idx2); + lb_update_stats_request_start(&lb, (size_t)idx2); + + /**< 第三次选择:backend2 */ + int idx3 = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(2, idx3); + lb_update_stats_request_start(&lb, (size_t)idx3); + + /**< 请求结束 */ + lb_update_stats_request_end(&lb, (size_t)idx1, true, 500); + lb_update_stats_request_end(&lb, (size_t)idx2, true, 600); + lb_update_stats_request_end(&lb, (size_t)idx3, false, 100); + + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[0].active_connections); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[1].active_connections); + TEST_ASSERT_EQUAL_UINT32(0, lb.stats[2].active_connections); + TEST_ASSERT_EQUAL_UINT32(1, lb.stats[2].total_failures); + + lb_destroy(&lb); +} + +/** + * @test test_full_request_lifecycle_weighted_response + * @brief 加权响应时间完整生命周期 + */ +void test_full_request_lifecycle_weighted_response(void) { + cocoon_load_balancer_t lb; + lb_init(&lb, COCOON_LB_WEIGHTED_RESPONSE); + + cocoon_proxy_rule_t rule; + uint32_t weights[] = {2, 1}; + make_weighted_rule(&rule, 2, weights); + + /**< 初始无 EWMA,两个后端都是 0,选第一个健康后端 */ + int idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(0, idx); + + /**< 模拟请求并记录不同响应时间 */ + lb_update_stats_request_start(&lb, 0); + lb_update_stats_request_end(&lb, 0, true, 2000); + + lb_update_stats_request_start(&lb, 1); + lb_update_stats_request_end(&lb, 1, true, 500); + + /**< backend0: ewma=2000, weight=2 => score=1000 */ + /**< backend1: ewma=500, weight=1 => score=500 */ + /**< 应选择 backend1 */ + idx = lb_select_backend(&lb, &rule, NULL); + TEST_ASSERT_EQUAL_INT(1, idx); + + lb_destroy(&lb); +} + +/** + * @test test_virtual_node_distribution_uniform + * @brief 虚拟节点在环上均匀分布(哈希值范围覆盖) + */ +void test_virtual_node_distribution_uniform(void) { + cocoon_proxy_rule_t rule; + make_test_rule(&rule, 3); + + cocoon_hash_ring_t ring; + memset(&ring, 0, sizeof(ring)); + lb_build_hash_ring(&ring, &rule); + + /**< 检查最小和最大哈希值存在合理差距 */ + uint32_t min_hash = ring.nodes[0].node_hash; + uint32_t max_hash = ring.nodes[ring.node_count - 1].node_hash; + + /**< 哈希值应分布在较大范围内 */ + TEST_ASSERT(min_hash < max_hash); + TEST_ASSERT_UINT32_WITHIN(0xFFFFFFFFU / 4, 0x80000000U, + min_hash + (max_hash - min_hash) / 2); +} + +/* ===== 主函数 ===== */ + +int main(void) { + UNITY_BEGIN(); + + /* 测试组 1: MurmurHash3 */ + RUN_TEST(test_hash_known_values); + RUN_TEST(test_hash_empty_string_zero_length); + RUN_TEST(test_hash_null_returns_zero); + RUN_TEST(test_hash_consistency); + RUN_TEST(test_hash_different_inputs); + RUN_TEST(test_hash_binary_data); + + /* 测试组 2: lb_init / lb_destroy */ + RUN_TEST(test_init_sets_algorithm); + RUN_TEST(test_init_default_alpha); + RUN_TEST(test_init_clears_stats); + RUN_TEST(test_init_clears_hash_ring); + RUN_TEST(test_destroy_cleans_up); + + /* 测试组 3: lb_algorithm_name */ + RUN_TEST(test_algorithm_name_all); + RUN_TEST(test_algorithm_name_unknown); + + /* 测试组 4: 一致性哈希环 */ + RUN_TEST(test_build_hash_ring_creates_nodes); + RUN_TEST(test_build_hash_ring_empty_rule); + RUN_TEST(test_build_hash_ring_skips_unhealthy); + RUN_TEST(test_build_hash_ring_all_unhealthy); + RUN_TEST(test_hash_ring_nodes_sorted); + RUN_TEST(test_hash_ring_backend_coverage); + RUN_TEST(test_pick_from_ring_basic); + RUN_TEST(test_pick_from_ring_wraparound); + RUN_TEST(test_pick_from_ring_zero_hash); + RUN_TEST(test_pick_from_ring_deterministic); + RUN_TEST(test_pick_from_ring_empty); + + /* 测试组 5: 一致性哈希映射稳定性 */ + RUN_TEST(test_consistent_hash_stability); + RUN_TEST(test_consistent_hash_different_keys); + RUN_TEST(test_consistent_hash_null_key_fallback); + RUN_TEST(test_consistent_hash_distribution); + + /* 测试组 6: 最少连接算法 */ + RUN_TEST(test_least_connections_selects_min); + RUN_TEST(test_least_connections_all_zero); + RUN_TEST(test_least_connections_skips_unhealthy); + RUN_TEST(test_least_connections_no_healthy); + RUN_TEST(test_least_connections_tie_break); + + /* 测试组 7: 统计更新 */ + RUN_TEST(test_stats_request_start); + RUN_TEST(test_stats_request_end_success); + RUN_TEST(test_stats_request_end_failure); + RUN_TEST(test_stats_end_without_start); + RUN_TEST(test_stats_multiple_backends); + + /* 测试组 8: EWMA 计算 */ + RUN_TEST(test_ewma_first_update); + RUN_TEST(test_ewma_subsequent_updates); + RUN_TEST(test_ewma_converges); + RUN_TEST(test_ewma_alpha_effect); + + /* 测试组 9: 加权响应时间选择 */ + RUN_TEST(test_weighted_response_selects_lowest_ratio); + RUN_TEST(test_weighted_response_weight_influence); + RUN_TEST(test_weighted_response_zero_weight); + + /* 测试组 10: 随机算法 */ + RUN_TEST(test_random_returns_valid_index); + RUN_TEST(test_random_distribution); + RUN_TEST(test_random_no_healthy); + + /* 测试组 11: 边界条件 */ + RUN_TEST(test_select_null_lb); + RUN_TEST(test_select_null_rule); + RUN_TEST(test_select_empty_rule); + RUN_TEST(test_update_stats_null_lb); + RUN_TEST(test_round_robin_returns_neg_one); + + /* 测试组 12: 综合集成 */ + RUN_TEST(test_full_request_lifecycle_least_connections); + RUN_TEST(test_full_request_lifecycle_weighted_response); + RUN_TEST(test_virtual_node_distribution_uniform); + + return UNITY_END(); +} diff --git a/tests/unit/test_middleware_ext.c b/tests/unit/test_middleware_ext.c new file mode 100644 index 0000000..ffbc09c --- /dev/null +++ b/tests/unit/test_middleware_ext.c @@ -0,0 +1,950 @@ +/** + * @file test_middleware_ext.c - 扩展中间件单元测试 + * + * 测试覆盖: + * - Base64Url 编解码 + * - JWT 认证(成功/失败/过期/格式错误) + * - Security Headers 配置 + * - Request ID 生成与复用 + * - IP 过滤(白名单/黑名单/CIDR/X-Forwarded-For) + * + * 使用 Unity 测试框架。 + */ + +#include "unity.h" +#include "middleware_ext.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +/* ============================================================ + * 内部函数前置声明(middleware_ext.c 中的非 static 函数) + * ============================================================ */ + +extern int base64url_decode(const char *in, unsigned char *out, int out_size); +extern int base64url_encode(const unsigned char *in, int in_len, char *out, int out_size); +extern const char *find_header(const http_request_t *req, const char *name); +extern time_t jwt_parse_exp(const char *payload); +extern bool jwt_verify_signature(const char *header_payload, + const char *signature_b64, + const char *secret, + size_t secret_len); +extern void send_json_error(cocoon_socket_t fd, int status, const char *body, bool keep_alive); +extern bool parse_ipv4(const char *ip_str, uint32_t *addr); +extern bool parse_cidr(const char *cidr_str, uint32_t *addr, int *mask); +extern bool ip_match_cidr(uint32_t ip_addr, const char *cidr_str); +extern bool parse_x_forwarded_for(const char *header_value, uint32_t *addr); + +/* ============================================================ + * 测试辅助函数 + * ============================================================ */ + +/** + * @brief 创建一对已连接的 socket 用于测试 + */ +static int create_socket_pair(int fds[2]) { + return socketpair(AF_UNIX, SOCK_STREAM, 0, fds); +} + +/** + * @brief 读取 socket 中可用数据到缓冲区 + * + * 使用非阻塞读取,避免在 socket 未关闭时无限等待 EOF。 + */ +static ssize_t read_all(int fd, char *buf, size_t buf_size) { + /* 设置 500ms 接收超时 */ + struct timeval tv = {0, 500000}; + setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &tv, sizeof(tv)); + + ssize_t total = 0; + ssize_t n; + while (total < (ssize_t)buf_size - 1 && + (n = read(fd, buf + total, buf_size - 1 - total)) > 0) { + total += n; + } + buf[total] = '\0'; + return total; +} + +/** + * @brief 生成测试用的 JWT token + * + * 使用 HS256 签名。 + * + * @param payload JWT payload JSON + * @param secret 签名密钥 + * @param token_out 输出 token 缓冲区 + * @param out_size 输出缓冲区大小 + * @return 0 成功,-1 失败 + */ +static int generate_jwt_token(const char *payload, const char *secret, + char *token_out, size_t out_size) { + /* header: {"alg":"HS256","typ":"JWT"} */ + const char *header = "{\"alg\":\"HS256\",\"typ\":\"JWT\"}"; + + char header_b64[512]; + char payload_b64[1024]; + + int header_len = base64url_encode((const unsigned char *)header, + (int)strlen(header), header_b64, sizeof(header_b64)); + int payload_len = base64url_encode((const unsigned char *)payload, + (int)strlen(payload), payload_b64, sizeof(payload_b64)); + if (header_len < 0 || payload_len < 0) return -1; + + /* 计算签名 */ + char to_sign[1536]; + snprintf(to_sign, sizeof(to_sign), "%s.%s", header_b64, payload_b64); + + unsigned char sig[EVP_MAX_MD_SIZE]; + unsigned int sig_len = 0; + if (!HMAC(EVP_sha256(), secret, (int)strlen(secret), + (const unsigned char *)to_sign, strlen(to_sign), + sig, &sig_len)) { + return -1; + } + + char sig_b64[EVP_MAX_MD_SIZE * 2]; + int sig_b64_len = base64url_encode(sig, (int)sig_len, sig_b64, sizeof(sig_b64)); + if (sig_b64_len < 0) return -1; + + int n = snprintf(token_out, out_size, "%s.%s.%s", + header_b64, payload_b64, sig_b64); + if (n < 0 || (size_t)n >= out_size) return -1; + return 0; +} + +/** + * @brief 创建带有指定 header 的 HTTP 请求 + */ +static void make_request_with_auth(http_request_t *req, const char *auth_header) { + memset(req, 0, sizeof(*req)); + req->method = HTTP_GET; + strcpy(req->path, "/api/test"); + strcpy(req->version, "HTTP/1.1"); + req->keep_alive = true; + + if (auth_header) { + strcpy(req->headers[0].name, "Authorization"); + strcpy(req->headers[0].value, auth_header); + req->num_headers = 1; + } +} + +/* ============================================================ + * setUp / tearDown + * ============================================================ */ + +void setUp(void) { + /* 每个测试前执行 */ +} + +void tearDown(void) { + /* 每个测试后执行 */ +} + +/* ============================================================ + * Base64Url 编解码测试 + * ============================================================ */ + +void test_base64url_decode_basic(void) { + /* 编码 "hello" -> aGVsbG8 */ + const char *encoded = "aGVsbG8"; + unsigned char decoded[16]; + int len = base64url_decode(encoded, decoded, sizeof(decoded)); + + TEST_ASSERT_EQUAL(5, len); + TEST_ASSERT_EQUAL_MEMORY("hello", decoded, 5); +} + +void test_base64url_decode_with_special_chars(void) { + /* Base64Url 使用 - 和 _ 代替 + 和 / */ + /* "test+data/ok" 在 Base64Url 中是 "dGVzdCtkYXRhL29r" */ + const char *encoded = "dGVzdCtkYXRhL29r"; + unsigned char decoded[32]; + int len = base64url_decode(encoded, decoded, sizeof(decoded)); + + TEST_ASSERT_EQUAL(12, len); + TEST_ASSERT_EQUAL_MEMORY("test+data/ok", decoded, 12); +} + +void test_base64url_decode_empty(void) { + unsigned char decoded[8]; + int len = base64url_decode("", decoded, sizeof(decoded)); + TEST_ASSERT_EQUAL(0, len); +} + +void test_base64url_decode_null_params(void) { + unsigned char decoded[8]; + TEST_ASSERT_EQUAL(-1, base64url_decode(NULL, decoded, sizeof(decoded))); + TEST_ASSERT_EQUAL(-1, base64url_decode("abc", NULL, 8)); +} + +void test_base64url_decode_binary(void) { + /* 测试二进制数据 */ + unsigned char binary[3] = {0xFF, 0x00, 0xAB}; + char encoded[16]; + int enc_len = base64url_encode(binary, 3, encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, enc_len); + + unsigned char decoded[8]; + int dec_len = base64url_decode(encoded, decoded, sizeof(decoded)); + TEST_ASSERT_EQUAL(3, dec_len); + TEST_ASSERT_EQUAL_UINT8(0xFF, decoded[0]); + TEST_ASSERT_EQUAL_UINT8(0x00, decoded[1]); + TEST_ASSERT_EQUAL_UINT8(0xAB, decoded[2]); +} + +void test_base64url_encode_decode_roundtrip(void) { + const char *orig = "The quick brown fox jumps over the lazy dog."; + char encoded[256]; + unsigned char decoded[256]; + + int enc_len = base64url_encode((const unsigned char *)orig, + (int)strlen(orig), encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, enc_len); + + int dec_len = base64url_decode(encoded, decoded, sizeof(decoded)); + TEST_ASSERT_EQUAL((int)strlen(orig), dec_len); + TEST_ASSERT_EQUAL_MEMORY(orig, decoded, strlen(orig)); +} + +void test_base64url_encode_no_padding(void) { + /* Base64Url 不应有填充 */ + unsigned char data[1] = {'a'}; + char encoded[16]; + int len = base64url_encode(data, 1, encoded, sizeof(encoded)); + TEST_ASSERT_GREATER_THAN(0, len); + /* 不应包含 = */ + TEST_ASSERT_NULL(strchr(encoded, '=')); +} + +/* ============================================================ + * find_header 测试 + * ============================================================ */ + +void test_find_header_exists(void) { + http_request_t req = {0}; + strcpy(req.headers[0].name, "Authorization"); + strcpy(req.headers[0].value, "Bearer token123"); + req.num_headers = 1; + + const char *val = find_header(&req, "Authorization"); + TEST_ASSERT_NOT_NULL(val); + TEST_ASSERT_EQUAL_STRING("Bearer token123", val); +} + +void test_find_header_case_insensitive(void) { + http_request_t req = {0}; + strcpy(req.headers[0].name, "X-Custom-Header"); + strcpy(req.headers[0].value, "custom-value"); + req.num_headers = 1; + + const char *val = find_header(&req, "x-custom-header"); + TEST_ASSERT_NOT_NULL(val); + TEST_ASSERT_EQUAL_STRING("custom-value", val); +} + +void test_find_header_not_found(void) { + http_request_t req = {0}; + req.num_headers = 0; + + const char *val = find_header(&req, "Authorization"); + TEST_ASSERT_NULL(val); +} + +void test_find_header_multiple_headers(void) { + http_request_t req = {0}; + strcpy(req.headers[0].name, "Host"); + strcpy(req.headers[0].value, "localhost"); + strcpy(req.headers[1].name, "Authorization"); + strcpy(req.headers[1].value, "Bearer abc"); + strcpy(req.headers[2].name, "Content-Type"); + strcpy(req.headers[2].value, "application/json"); + req.num_headers = 3; + + TEST_ASSERT_EQUAL_STRING("localhost", find_header(&req, "Host")); + TEST_ASSERT_EQUAL_STRING("Bearer abc", find_header(&req, "Authorization")); + TEST_ASSERT_EQUAL_STRING("application/json", find_header(&req, "Content-Type")); + TEST_ASSERT_NULL(find_header(&req, "X-Not-Found")); +} + +/* ============================================================ + * jwt_parse_exp 测试 + * ============================================================ */ + +void test_jwt_parse_exp_present(void) { + const char *payload = "{\"sub\":\"user1\",\"exp\":1893456000,\"iat\":1609459200}"; + time_t exp = jwt_parse_exp(payload); + TEST_ASSERT_EQUAL((time_t)1893456000, exp); +} + +void test_jwt_parse_exp_not_present(void) { + const char *payload = "{\"sub\":\"user1\",\"iat\":1609459200}"; + time_t exp = jwt_parse_exp(payload); + TEST_ASSERT_EQUAL(0, exp); +} + +void test_jwt_parse_exp_first_field(void) { + const char *payload = "{\"exp\":2000000000,\"sub\":\"user1\"}"; + time_t exp = jwt_parse_exp(payload); + TEST_ASSERT_EQUAL((time_t)2000000000, exp); +} + +void test_jwt_parse_exp_null(void) { + TEST_ASSERT_EQUAL(0, jwt_parse_exp(NULL)); +} + +/* ============================================================ + * JWT 签名验证测试 + * ============================================================ */ + +void test_jwt_verify_signature_valid(void) { + const char *secret = "my-secret-key"; + const char *payload = "{\"sub\":\"user1\",\"exp\":1893456000}"; + char token[2048]; + + TEST_ASSERT_EQUAL(0, generate_jwt_token(payload, secret, token, sizeof(token))); + + /* 提取 header.payload 部分 */ + char *sig_dot = strrchr(token, '.'); + TEST_ASSERT_NOT_NULL(sig_dot); + *sig_dot = '\0'; + + const char *signature_b64 = sig_dot + 1; + + bool valid = jwt_verify_signature(token, signature_b64, secret, strlen(secret)); + TEST_ASSERT_TRUE(valid); +} + +void test_jwt_verify_signature_invalid_secret(void) { + const char *secret = "my-secret-key"; + const char *payload = "{\"sub\":\"user1\"}"; + char token[2048]; + + TEST_ASSERT_EQUAL(0, generate_jwt_token(payload, secret, token, sizeof(token))); + + char *sig_dot = strrchr(token, '.'); + TEST_ASSERT_NOT_NULL(sig_dot); + *sig_dot = '\0'; + + const char *signature_b64 = sig_dot + 1; + + /* 使用错误的密钥验证 */ + bool valid = jwt_verify_signature(token, signature_b64, "wrong-secret", strlen("wrong-secret")); + TEST_ASSERT_FALSE(valid); +} + +void test_jwt_verify_signature_tampered_payload(void) { + const char *secret = "my-secret-key"; + const char *payload = "{\"sub\":\"user1\"}"; + char token[2048]; + + TEST_ASSERT_EQUAL(0, generate_jwt_token(payload, secret, token, sizeof(token))); + + /* 篡改 payload:在 header_b64 和 payload_b64 之间添加字符 */ + char tampered[2048]; + strncpy(tampered, token, sizeof(tampered) - 1); + tampered[sizeof(tampered) - 1] = '\0'; + + char *sig_dot = strrchr(tampered, '.'); + TEST_ASSERT_NOT_NULL(sig_dot); + *sig_dot = '\0'; + + const char *signature_b64 = sig_dot + 1; + + /* 篡改 header.payload 部分 */ + strcat(tampered, ".extra"); + + bool valid = jwt_verify_signature(tampered, signature_b64, secret, strlen(secret)); + TEST_ASSERT_FALSE(valid); +} + +/* ============================================================ + * JWT 中间件完整流程测试 + * ============================================================ */ + +void test_jwt_middleware_success(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .header_name = "Authorization", + .prefix = "Bearer ", + .skip_preflight = false, + }; + + /* 生成有效的 token */ + char token[2048]; + const char *payload = "{\"sub\":\"user1\",\"exp\":4102444800}"; /* 2099 年 */ + TEST_ASSERT_EQUAL(0, generate_jwt_token(payload, "test-secret", token, sizeof(token))); + + char auth_header[2300]; + snprintf(auth_header, sizeof(auth_header), "Bearer %s", token); + + http_request_t req; + make_request_with_auth(&req, auth_header); + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 验证通过 */ + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_missing_header(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .header_name = "Authorization", + .prefix = "Bearer ", + }; + + http_request_t req; + make_request_with_auth(&req, NULL); /* 无 Authorization 头 */ + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(1, ret); /* 短路 */ + + /* 读取响应 */ + char response[1024]; + ssize_t n = read_all(fds[1], response, sizeof(response)); + TEST_ASSERT_GREATER_THAN(0, n); + TEST_ASSERT_NOT_NULL(strstr(response, "401")); + TEST_ASSERT_NOT_NULL(strstr(response, "Unauthorized")); + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_wrong_prefix(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .prefix = "Bearer ", + }; + + http_request_t req; + make_request_with_auth(&req, "Basic dXNlcjpwYXNz"); /* Basic 而非 Bearer */ + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(1, ret); /* 短路 */ + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_invalid_signature(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .prefix = "Bearer ", + }; + + /* 使用错误签名的 token */ + http_request_t req; + make_request_with_auth(&req, "Bearer eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiJ1c2VyMSJ9.invalidsignature"); + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(1, ret); /* 短路 */ + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_expired_token(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .prefix = "Bearer ", + }; + + /* 生成已过期 token */ + char token[2048]; + const char *payload = "{\"sub\":\"user1\",\"exp\":1000000000}"; /* 2001 年,已过期 */ + TEST_ASSERT_EQUAL(0, generate_jwt_token(payload, "test-secret", token, sizeof(token))); + + char auth_header[2300]; + snprintf(auth_header, sizeof(auth_header), "Bearer %s", token); + + http_request_t req; + make_request_with_auth(&req, auth_header); + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(1, ret); /* 过期,短路 */ + + /* 读取响应确认 401 */ + char response[1024]; + ssize_t n = read_all(fds[1], response, sizeof(response)); + TEST_ASSERT_GREATER_THAN(0, n); + TEST_ASSERT_NOT_NULL(strstr(response, "401")); + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_skip_options(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .skip_preflight = true, + }; + + /* OPTIONS 请求应跳过验证 */ + http_request_t req = {0}; + req.method = HTTP_OPTIONS; + strcpy(req.path, "/api/test"); + req.keep_alive = true; + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 跳过 */ + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_no_config(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + /* 空密钥,应跳过 */ + cocoon_jwt_config_t cfg = {0}; + + http_request_t req; + make_request_with_auth(&req, NULL); + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 跳过 */ + + close(fds[0]); + close(fds[1]); +} + +void test_jwt_middleware_malformed_token(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + cocoon_jwt_config_t cfg = { + .secret = "test-secret", + .prefix = "Bearer ", + }; + + /* 缺少分隔符的 token */ + http_request_t req; + make_request_with_auth(&req, "Bearer malformedtoken"); + + int ret = cocoon_middleware_jwt(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(1, ret); /* 格式错误,短路 */ + + close(fds[0]); + close(fds[1]); +} + +/* ============================================================ + * Security Headers 中间件测试 + * ============================================================ */ + +void test_security_headers_set_config(void) { + cocoon_security_headers_config_t cfg = { + .hsts_enabled = true, + .hsts_max_age = 31536000, + .hsts_include_subdomains = true, + .frame_options_enabled = true, + .frame_options = "DENY", + .xss_protection_enabled = true, + .csp_enabled = true, + .csp_policy = "default-src 'self'", + .content_type_options = true, + .referrer_policy_enabled = true, + .referrer_policy = "strict-origin-when-cross-origin", + }; + + http_request_t req = {0}; + int ret = cocoon_middleware_security_headers(&req, -1, &cfg); + TEST_ASSERT_EQUAL(0, ret); + + const cocoon_security_headers_config_t *got = + cocoon_middleware_security_headers_get(); + TEST_ASSERT_NOT_NULL(got); + TEST_ASSERT_TRUE(got->hsts_enabled); + TEST_ASSERT_EQUAL(31536000, got->hsts_max_age); + TEST_ASSERT_TRUE(got->hsts_include_subdomains); + TEST_ASSERT_TRUE(got->frame_options_enabled); + TEST_ASSERT_EQUAL_STRING("DENY", got->frame_options); + TEST_ASSERT_TRUE(got->xss_protection_enabled); + TEST_ASSERT_TRUE(got->csp_enabled); + TEST_ASSERT_EQUAL_STRING("default-src 'self'", got->csp_policy); + TEST_ASSERT_TRUE(got->content_type_options); + TEST_ASSERT_TRUE(got->referrer_policy_enabled); + TEST_ASSERT_EQUAL_STRING("strict-origin-when-cross-origin", got->referrer_policy); +} + +void test_security_headers_null_config(void) { + http_request_t req = {0}; + int ret = cocoon_middleware_security_headers(&req, -1, NULL); + TEST_ASSERT_EQUAL(0, ret); + /* 配置应保持不变 */ +} + +void test_security_headers_get_not_initialized(void) { + /* 注意:如果之前测试已初始化,此处可能不为 NULL */ + /* 测试中不做断言,仅确认不崩溃 */ + (void)cocoon_middleware_security_headers_get(); +} + +void test_security_headers_sameorigin(void) { + cocoon_security_headers_config_t cfg = { + .frame_options_enabled = true, + .frame_options = "SAMEORIGIN", + }; + + http_request_t req = {0}; + cocoon_middleware_security_headers(&req, -1, &cfg); + + const cocoon_security_headers_config_t *got = + cocoon_middleware_security_headers_get(); + TEST_ASSERT_NOT_NULL(got); + TEST_ASSERT_EQUAL_STRING("SAMEORIGIN", got->frame_options); +} + +/* ============================================================ + * Request ID 中间件测试 + * ============================================================ */ + +void test_request_id_generates_32_chars(void) { + cocoon_request_id_config_t cfg = { + .header_name = "X-Request-ID", + .trust_incoming = true, + }; + + http_request_t req = {0}; + req.method = HTTP_GET; + + int ret = cocoon_middleware_request_id(&req, -1, &cfg); + TEST_ASSERT_EQUAL(0, ret); + /* 中间件内部生成了 32 字符 ID, + 但当前版本未暴露给外部,仅确认不崩溃 */ +} + +void test_request_id_null_config(void) { + http_request_t req = {0}; + int ret = cocoon_middleware_request_id(&req, -1, NULL); + TEST_ASSERT_EQUAL(0, ret); +} + +void test_request_id_trust_incoming_valid(void) { + cocoon_request_id_config_t cfg = { + .header_name = "X-Request-ID", + .trust_incoming = true, + }; + + http_request_t req = {0}; + strcpy(req.headers[0].name, "X-Request-ID"); + strcpy(req.headers[0].value, "aabbccdd11223344556677889900aabb"); /* 32 字符 hex */ + req.num_headers = 1; + + int ret = cocoon_middleware_request_id(&req, -1, &cfg); + TEST_ASSERT_EQUAL(0, ret); +} + +void test_request_id_trust_incoming_invalid_hex(void) { + cocoon_request_id_config_t cfg = { + .header_name = "X-Request-ID", + .trust_incoming = true, + }; + + /* 非 hex 字符 */ + http_request_t req = {0}; + strcpy(req.headers[0].name, "X-Request-ID"); + strcpy(req.headers[0].value, "zzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzz"); + req.num_headers = 1; + + int ret = cocoon_middleware_request_id(&req, -1, &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 忽略无效值,生成新 ID */ +} + +void test_request_id_not_trust_incoming(void) { + cocoon_request_id_config_t cfg = { + .header_name = "X-Request-ID", + .trust_incoming = false, + }; + + http_request_t req = {0}; + strcpy(req.headers[0].name, "X-Request-ID"); + strcpy(req.headers[0].value, "aabbccdd11223344556677889900aabb"); + req.num_headers = 1; + + int ret = cocoon_middleware_request_id(&req, -1, &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 不信任传入,生成新 ID */ +} + +/* ============================================================ + * IP 过滤工具函数测试 + * ============================================================ */ + +void test_parse_ipv4_valid(void) { + uint32_t addr; + TEST_ASSERT_TRUE(parse_ipv4("192.168.1.1", &addr)); + TEST_ASSERT_EQUAL((uint32_t)0xC0A80101, addr); /* 192.168.1.1 */ + + TEST_ASSERT_TRUE(parse_ipv4("0.0.0.0", &addr)); + TEST_ASSERT_EQUAL(0U, addr); + + TEST_ASSERT_TRUE(parse_ipv4("255.255.255.255", &addr)); + TEST_ASSERT_EQUAL(0xFFFFFFFFU, addr); +} + +void test_parse_ipv4_invalid(void) { + uint32_t addr; + TEST_ASSERT_FALSE(parse_ipv4("invalid", &addr)); + TEST_ASSERT_FALSE(parse_ipv4("", &addr)); + TEST_ASSERT_FALSE(parse_ipv4("256.1.1.1", &addr)); /* 超出范围 */ + TEST_ASSERT_FALSE(parse_ipv4("1.1.1", &addr)); /* 不足 4 段 */ +} + +void test_parse_cidr_exact(void) { + uint32_t addr; + int mask; + TEST_ASSERT_TRUE(parse_cidr("192.168.1.1", &addr, &mask)); + TEST_ASSERT_EQUAL(32, mask); + TEST_ASSERT_EQUAL((uint32_t)0xC0A80101, addr); +} + +void test_parse_cidr_with_mask(void) { + uint32_t addr; + int mask; + TEST_ASSERT_TRUE(parse_cidr("192.168.1.0/24", &addr, &mask)); + TEST_ASSERT_EQUAL(24, mask); + /* 网络地址应为 192.168.1.0 */ + TEST_ASSERT_EQUAL((uint32_t)0xC0A80100, addr); +} + +void test_parse_cidr_16(void) { + uint32_t addr; + int mask; + TEST_ASSERT_TRUE(parse_cidr("10.0.0.0/16", &addr, &mask)); + TEST_ASSERT_EQUAL(16, mask); + TEST_ASSERT_EQUAL((uint32_t)0x0A000000, addr); +} + +void test_parse_cidr_8(void) { + uint32_t addr; + int mask; + TEST_ASSERT_TRUE(parse_cidr("172.0.0.0/8", &addr, &mask)); + TEST_ASSERT_EQUAL(8, mask); + TEST_ASSERT_EQUAL((uint32_t)0xAC000000, addr); +} + +void test_parse_cidr_invalid(void) { + uint32_t addr; + int mask; + TEST_ASSERT_FALSE(parse_cidr("invalid", &addr, &mask)); + TEST_ASSERT_FALSE(parse_cidr("192.168.1.0/33", &addr, &mask)); /* mask > 32 */ + TEST_ASSERT_FALSE(parse_cidr("", &addr, &mask)); +} + +void test_ip_match_cidr_exact(void) { + uint32_t ip; + parse_ipv4("192.168.1.1", &ip); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "192.168.1.1")); + TEST_ASSERT_FALSE(ip_match_cidr(ip, "192.168.1.2")); +} + +void test_ip_match_cidr_24(void) { + uint32_t ip; + parse_ipv4("192.168.1.100", &ip); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "192.168.1.0/24")); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "192.168.1.0/16")); + TEST_ASSERT_FALSE(ip_match_cidr(ip, "10.0.0.0/24")); +} + +void test_ip_match_cidr_16(void) { + uint32_t ip; + parse_ipv4("10.0.50.100", &ip); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "10.0.0.0/16")); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "10.0.0.0/8")); + TEST_ASSERT_FALSE(ip_match_cidr(ip, "10.1.0.0/16")); +} + +void test_ip_match_cidr_edge_cases(void) { + uint32_t ip; + parse_ipv4("0.0.0.0", &ip); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "0.0.0.0/0")); + + parse_ipv4("255.255.255.255", &ip); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "255.255.255.255")); + TEST_ASSERT_TRUE(ip_match_cidr(ip, "0.0.0.0/0")); +} + +void test_parse_x_forwarded_for_single(void) { + uint32_t addr; + TEST_ASSERT_TRUE(parse_x_forwarded_for("192.168.1.100", &addr)); + TEST_ASSERT_EQUAL((uint32_t)0xC0A80164, addr); +} + +void test_parse_x_forwarded_for_chain(void) { + uint32_t addr; + TEST_ASSERT_TRUE(parse_x_forwarded_for("192.168.1.100, 10.0.0.1, 172.16.0.1", &addr)); + TEST_ASSERT_EQUAL((uint32_t)0xC0A80164, addr); /* 取第一个 */ +} + +void test_parse_x_forwarded_for_with_spaces(void) { + uint32_t addr; + TEST_ASSERT_TRUE(parse_x_forwarded_for(" 192.168.1.50 ", &addr)); + TEST_ASSERT_EQUAL((uint32_t)0xC0A80132, addr); +} + +void test_parse_x_forwarded_for_invalid(void) { + uint32_t addr; + TEST_ASSERT_FALSE(parse_x_forwarded_for("invalid-ip", &addr)); + TEST_ASSERT_FALSE(parse_x_forwarded_for("", &addr)); + TEST_ASSERT_FALSE(parse_x_forwarded_for(NULL, &addr)); +} + +/* ============================================================ + * IP 过滤中间件完整流程测试 + * ============================================================ */ + +void test_ip_filter_no_config(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + /* 空配置应跳过 */ + cocoon_ip_filter_config_t cfg = {0}; + + http_request_t req = {0}; + int ret = cocoon_middleware_ip_filter(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 跳过 */ + + close(fds[0]); + close(fds[1]); +} + +void test_ip_filter_blacklist_allow(void) { + int fds[2]; + TEST_ASSERT_EQUAL(0, create_socket_pair(fds)); + + /* 黑名单模式,不包含 127.0.0.1,应允许 */ + cocoon_ip_filter_config_t cfg = { + .count = 1, + .mode = COCOON_IP_FILTER_DENY, + }; + strcpy(cfg.entries[0], "192.168.1.0/24"); + + http_request_t req = {0}; + int ret = cocoon_middleware_ip_filter(&req, fds[0], &cfg); + TEST_ASSERT_EQUAL(0, ret); /* 允许(127.0.0.1 不在黑名单中) */ + + close(fds[0]); + close(fds[1]); +} + +/* ============================================================ + * 一键初始化测试 + * ============================================================ */ + +void test_middleware_init_extended(void) { + /* 确认不崩溃 */ + cocoon_middleware_init_extended(NULL); + cocoon_middleware_init_extended((void *)0x1234); /* 无效指针,但不应崩溃 */ +} + +/* ============================================================ + * 主函数 + * ============================================================ */ + +int main(void) { + UNITY_BEGIN(); + + /* Base64Url 编解码 (7) */ + RUN_TEST(test_base64url_decode_basic); + RUN_TEST(test_base64url_decode_with_special_chars); + RUN_TEST(test_base64url_decode_empty); + RUN_TEST(test_base64url_decode_null_params); + RUN_TEST(test_base64url_decode_binary); + RUN_TEST(test_base64url_encode_decode_roundtrip); + RUN_TEST(test_base64url_encode_no_padding); + + /* find_header (4) */ + RUN_TEST(test_find_header_exists); + RUN_TEST(test_find_header_case_insensitive); + RUN_TEST(test_find_header_not_found); + RUN_TEST(test_find_header_multiple_headers); + + /* jwt_parse_exp (4) */ + RUN_TEST(test_jwt_parse_exp_present); + RUN_TEST(test_jwt_parse_exp_not_present); + RUN_TEST(test_jwt_parse_exp_first_field); + RUN_TEST(test_jwt_parse_exp_null); + + /* JWT 签名验证 (3) */ + RUN_TEST(test_jwt_verify_signature_valid); + RUN_TEST(test_jwt_verify_signature_invalid_secret); + RUN_TEST(test_jwt_verify_signature_tampered_payload); + + /* JWT 中间件完整流程 (8) */ + RUN_TEST(test_jwt_middleware_success); + RUN_TEST(test_jwt_middleware_missing_header); + RUN_TEST(test_jwt_middleware_wrong_prefix); + RUN_TEST(test_jwt_middleware_invalid_signature); + RUN_TEST(test_jwt_middleware_expired_token); + RUN_TEST(test_jwt_middleware_skip_options); + RUN_TEST(test_jwt_middleware_no_config); + RUN_TEST(test_jwt_middleware_malformed_token); + + /* Security Headers (4) */ + RUN_TEST(test_security_headers_set_config); + RUN_TEST(test_security_headers_null_config); + RUN_TEST(test_security_headers_get_not_initialized); + RUN_TEST(test_security_headers_sameorigin); + + /* Request ID (5) */ + RUN_TEST(test_request_id_generates_32_chars); + RUN_TEST(test_request_id_null_config); + RUN_TEST(test_request_id_trust_incoming_valid); + RUN_TEST(test_request_id_trust_incoming_invalid_hex); + RUN_TEST(test_request_id_not_trust_incoming); + + /* IP 工具函数 (15) */ + RUN_TEST(test_parse_ipv4_valid); + RUN_TEST(test_parse_ipv4_invalid); + RUN_TEST(test_parse_cidr_exact); + RUN_TEST(test_parse_cidr_with_mask); + RUN_TEST(test_parse_cidr_16); + RUN_TEST(test_parse_cidr_8); + RUN_TEST(test_parse_cidr_invalid); + RUN_TEST(test_ip_match_cidr_exact); + RUN_TEST(test_ip_match_cidr_24); + RUN_TEST(test_ip_match_cidr_16); + RUN_TEST(test_ip_match_cidr_edge_cases); + RUN_TEST(test_parse_x_forwarded_for_single); + RUN_TEST(test_parse_x_forwarded_for_chain); + RUN_TEST(test_parse_x_forwarded_for_with_spaces); + RUN_TEST(test_parse_x_forwarded_for_invalid); + + /* IP 过滤中间件 (2) */ + RUN_TEST(test_ip_filter_no_config); + RUN_TEST(test_ip_filter_blacklist_allow); + + /* 一键初始化 (1) */ + RUN_TEST(test_middleware_init_extended); + + return UNITY_END(); +}