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:
xfy911 2026-06-10 11:36:45 +08:00
parent 525cd8954c
commit bf9787e6e9
15 changed files with 8885 additions and 4 deletions

3
.gitmodules vendored
View File

@ -1,3 +0,0 @@
[submodule "coco"]
path = coco
url = https://github.com/xfy911/coco.git

1
coco

@ -1 +0,0 @@
Subproject commit 8b42e984efd55f5c18fdf2d249d1aba0ca4e70eb

53
coco/include/coco.h Normal file
View 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
View 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
View 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/grpcapplication/grpc+protoapplication/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 */

1671
http3.c Normal file

File diff suppressed because it is too large Load Diff

625
http3.h Normal file
View File

@ -0,0 +1,625 @@
/**
* @file http3.h
* @brief HTTP/3 (QUIC)
*
* HTTP/3 over QUIC
* - QUIC UDP socket ID
* - HTTP/3 HEADERSDATASETTINGSGOAWAY
* - 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
*
* SETTINGSGOAWAY
*
* @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
View 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
View 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
View File

@ -0,0 +1,662 @@
/**
* @file middleware_ext.c -
*
* Phase 4
* - JWT HS256, Base64Url, exp
* - Security Headers
* - Request ID32 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
View 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.signatureBase64Url
*
* @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

File diff suppressed because it is too large Load Diff

989
tests/unit/test_http3.c Normal file
View 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 编码631字节最大值 */
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 编码163832字节最大值 */
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 编码10737418234字节最大值 */
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();
}

File diff suppressed because it is too large Load Diff

View 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();
}