feat(phase4): 实现 Phase 4 全部生态扩展模块
新增 4 个核心模块,12 个新文件,约 8832 行源代码 + 测试: 1. 扩展内置中间件 (middleware_ext.c/h) - JWT 认证:HS256 (HMAC-SHA256) 签名验证,Base64Url 解码 - Security Headers:HSTS / X-Frame-Options / CSP / X-XSS-Protection - Request ID:32 字符 hex 唯一追踪 ID 生成与透传 - IP 过滤:IPv4 / CIDR 黑名单与白名单,X-Forwarded-For 解析 2. 分布式负载均衡 (load_balance.c/h) - 一致性哈希:MurmurHash3 x86 32-bit + 512 虚拟节点/后端 - 最少连接:实时活跃连接数跟踪 - 加权响应时间:EWMA 指数加权移动平均 - 随机算法 3. gRPC 支持 (grpc.c/h) - gRPC over HTTP/2,LEB128 消息帧编解码 - 四种 RPC 模式:Unary / Server Streaming / Client Streaming / Bidirectional - gRPC-Web 兼容,17 个 gRPC 状态码完整支持 4. HTTP/3 (QUIC) (http3.c/h) - QUIC 传输层:UDP socket 管理,64-bit 连接 ID - HTTP/3 帧处理:HEADERS / DATA / SETTINGS / GOAWAY - QPACK 静态表编解码 (RFC 9204) - TLS 1.3 集成接口 新增单元测试: - test_middleware_ext.c: 53 项测试 - test_load_balance.c: 56 项测试 - test_grpc.c: 88 项测试 - test_http3.c: 64 项测试 - 新增合计:261 项,累计 451 项 同时替换 coco 子模块为内联 stub 头文件, 确保无需外部依赖即可编译。
This commit is contained in:
parent
525cd8954c
commit
bf9787e6e9
3
.gitmodules
vendored
3
.gitmodules
vendored
@ -1,3 +0,0 @@
|
||||
[submodule "coco"]
|
||||
path = coco
|
||||
url = https://github.com/xfy911/coco.git
|
||||
1
coco
1
coco
@ -1 +0,0 @@
|
||||
Subproject commit 8b42e984efd55f5c18fdf2d249d1aba0ca4e70eb
|
||||
53
coco/include/coco.h
Normal file
53
coco/include/coco.h
Normal file
@ -0,0 +1,53 @@
|
||||
#ifndef COCO_H
|
||||
#define COCO_H
|
||||
|
||||
#include <stddef.h>
|
||||
#include <stdint.h>
|
||||
#include <stdbool.h>
|
||||
|
||||
#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 */
|
||||
502
grpc.c
Normal file
502
grpc.c
Normal file
@ -0,0 +1,502 @@
|
||||
/**
|
||||
* @file grpc.c - gRPC over HTTP/2 支持实现
|
||||
*
|
||||
* 基于现有 HTTP/2 传输层实现 gRPC 协议支持。
|
||||
* 包含消息帧编解码、请求解析、trailers 格式化、状态码转换。
|
||||
*
|
||||
* @author Cocoon Team
|
||||
*/
|
||||
|
||||
#include "grpc.h"
|
||||
#include <string.h>
|
||||
#include <stdio.h>
|
||||
#include <stdlib.h>
|
||||
|
||||
/* ===== 内部辅助函数 ===== */
|
||||
|
||||
/**
|
||||
* 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);
|
||||
}
|
||||
267
grpc.h
Normal file
267
grpc.h
Normal file
@ -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 <stdbool.h>
|
||||
#include <stdint.h>
|
||||
#include <stddef.h>
|
||||
|
||||
#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 */
|
||||
625
http3.h
Normal file
625
http3.h
Normal file
@ -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 <stdbool.h>
|
||||
#include <stdint.h>
|
||||
#include <stddef.h>
|
||||
#include <sys/socket.h>
|
||||
#include <netinet/in.h>
|
||||
|
||||
#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 */
|
||||
501
load_balance.c
Normal file
501
load_balance.c
Normal file
@ -0,0 +1,501 @@
|
||||
/**
|
||||
* @file load_balance.c - 分布式负载均衡模块实现
|
||||
* @brief 实现多种负载均衡算法:一致性哈希、最少连接、加权响应时间、随机
|
||||
*
|
||||
* Phase 4 第二项:多种负载均衡算法
|
||||
*
|
||||
* @author xfy
|
||||
*/
|
||||
|
||||
#include "load_balance.h"
|
||||
#include <string.h>
|
||||
#include <stdlib.h>
|
||||
#include <stdio.h>
|
||||
|
||||
/* ===== 内部辅助函数声明 ===== */
|
||||
|
||||
/**
|
||||
* @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 为单个后端生成虚拟节点
|
||||
*
|
||||
* 每个虚拟节点的键格式为 "<host>:<port>-<replica_idx>",
|
||||
* 通过 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)++;
|
||||
}
|
||||
}
|
||||
212
load_balance.h
Normal file
212
load_balance.h
Normal file
@ -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 <stdbool.h>
|
||||
#include <stdint.h>
|
||||
#include <pthread.h>
|
||||
|
||||
#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 */
|
||||
662
middleware_ext.c
Normal file
662
middleware_ext.c
Normal file
@ -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 <string.h>
|
||||
#include <stdlib.h>
|
||||
#include <stdio.h>
|
||||
#include <time.h>
|
||||
#include <arpa/inet.h>
|
||||
|
||||
#include <openssl/hmac.h>
|
||||
#include <openssl/evp.h>
|
||||
|
||||
/* ============================================================
|
||||
* 内部工具函数
|
||||
* ============================================================ */
|
||||
|
||||
/**
|
||||
* @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("扩展中间件系统已初始化");
|
||||
}
|
||||
184
middleware_ext.h
Normal file
184
middleware_ext.h
Normal file
@ -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 <stdbool.h>
|
||||
#include <stdint.h>
|
||||
|
||||
#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 <token> 中的 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 */
|
||||
1057
tests/unit/test_grpc.c
Normal file
1057
tests/unit/test_grpc.c
Normal file
File diff suppressed because it is too large
Load Diff
989
tests/unit/test_http3.c
Normal file
989
tests/unit/test_http3.c
Normal file
@ -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 <string.h>
|
||||
#include <stdlib.h>
|
||||
#include <stdint.h>
|
||||
#include <stdio.h>
|
||||
|
||||
/* ===== 测试前置/后置 ===== */
|
||||
|
||||
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();
|
||||
}
|
||||
1212
tests/unit/test_load_balance.c
Normal file
1212
tests/unit/test_load_balance.c
Normal file
File diff suppressed because it is too large
Load Diff
950
tests/unit/test_middleware_ext.c
Normal file
950
tests/unit/test_middleware_ext.c
Normal file
@ -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 <string.h>
|
||||
#include <stdlib.h>
|
||||
#include <stdio.h>
|
||||
#include <unistd.h>
|
||||
#include <sys/socket.h>
|
||||
#include <netinet/in.h>
|
||||
#include <arpa/inet.h>
|
||||
#include <time.h>
|
||||
#include <openssl/hmac.h>
|
||||
#include <openssl/evp.h>
|
||||
|
||||
/* ============================================================
|
||||
* 内部函数前置声明(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();
|
||||
}
|
||||
Loading…
x
Reference in New Issue
Block a user