diff --git a/.bcr/MODULE.bazel b/.bcr/MODULE.bazel index c09a051..39a0c92 100644 --- a/.bcr/MODULE.bazel +++ b/.bcr/MODULE.bazel @@ -2,3 +2,26 @@ module( name = "khttpd", version = "{{VERSION}}", ) + +bazel_dep(name = "platforms", version = "1.1.0") +bazel_dep(name = "bazel_skylib", version = "1.9.0") +bazel_dep(name = "rules_cc", version = "0.2.20") +bazel_dep(name = "rules_shell", version = "0.8.0") +bazel_dep(name = "rules_perl", version = "1.1.1") +bazel_dep(name = "fmt", version = "12.1.0") +bazel_dep(name = "googletest", version = "1.17.0.bcr.2") +bazel_dep(name = "sqlite3", version = "3.53.2") +bazel_dep(name = "openssl", version = "4.0.1.bcr.0") +bazel_dep(name = "boringssl", version = "0.20260616.0") +bazel_dep(name = "boost", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.asio", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.json", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.mysql", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.beast", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.filesystem", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.url", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.uuid", version = "1.90.0.bcr.1") +bazel_dep(name = "spdlog", version = "1.17.0") + +cc_configure = use_extension("@rules_cc//cc:extensions.bzl", "cc_configure_extension") +use_repo(cc_configure, "local_config_cc") diff --git a/.github/workflows/bazel.yml b/.github/workflows/bazel.yml index ac175c2..1454659 100644 --- a/.github/workflows/bazel.yml +++ b/.github/workflows/bazel.yml @@ -61,7 +61,7 @@ jobs: - name: Bazel Test run: | - bazel test framework/... --test_output=errors --test_verbose_timeout_warnings --verbose_failures + bazel test //framework/... --test_output=errors --test_verbose_timeout_warnings --verbose_failures example: needs: build strategy: @@ -108,7 +108,10 @@ jobs: release: needs: [build, example] runs-on: ubuntu-latest - if: github.event_name == 'push' && github.ref == 'refs/heads/main' + if: github.repository == 'ClangTools/khttpd' && github.event_name == 'push' && github.ref == 'refs/heads/main' + outputs: + tag_name: ${{ steps.version.outputs.tag_name }} + release_created: ${{ steps.check_tag.outputs.exists == 'false' }} permissions: contents: write steps: @@ -122,6 +125,7 @@ jobs: run: | VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' MODULE.bazel | head -1) echo "version=${VERSION}" >> "$GITHUB_OUTPUT" + echo "tag_name=v${VERSION}" >> "$GITHUB_OUTPUT" echo "Found version: ${VERSION}" - name: Check if tag exists @@ -153,3 +157,12 @@ jobs: name: Release v${{ steps.version.outputs.version }} generate_release_notes: true draft: false + + publish_to_bcr: + name: Publish release to BCR + needs: release + if: github.repository == 'ClangTools/khttpd' && needs.release.outputs.release_created == 'true' + uses: ./.github/workflows/publish_to_bcr.yml + with: + tag_name: ${{ needs.release.outputs.tag_name }} + secrets: inherit diff --git a/.github/workflows/publish_to_bcr.yml b/.github/workflows/publish_to_bcr.yml index 5c00e00..315bd02 100644 --- a/.github/workflows/publish_to_bcr.yml +++ b/.github/workflows/publish_to_bcr.yml @@ -3,12 +3,22 @@ name: Publish to BCR on: release: types: [published] + workflow_call: + inputs: + tag_name: + description: Release tag to publish to the Bazel Central Registry + required: true + type: string + secrets: + BCR_PUBLISH_TOKEN: + required: true jobs: publish: + if: github.repository == 'ClangTools/khttpd' uses: kekxv/bcr/.github/workflows/publish_to_bcr.yml@publish-to-bcr with: - tag_name: ${{ github.event.release.tag_name }} + tag_name: ${{ inputs.tag_name || github.event.release.tag_name }} module_name: "khttpd" secrets: publish_token: ${{ secrets.BCR_PUBLISH_TOKEN }} diff --git a/MODULE.bazel b/MODULE.bazel index 86939f6..cad0c01 100644 --- a/MODULE.bazel +++ b/MODULE.bazel @@ -1,6 +1,6 @@ module( name = "khttpd", - version = "0.2.0", + version = "0.3.0", ) bazel_dep(name = "platforms", version = "1.1.0") @@ -11,7 +11,7 @@ bazel_dep(name = "rules_perl", version = "1.1.1") bazel_dep(name = "fmt", version = "12.1.0") bazel_dep(name = "googletest", version = "1.17.0.bcr.2") bazel_dep(name = "sqlite3", version = "3.53.2") -bazel_dep(name = "openssl", version = "3.5.5.bcr.4") +bazel_dep(name = "openssl", version = "4.0.1.bcr.0") bazel_dep(name = "boringssl", version = "0.20260616.0") bazel_dep(name = "boost", version = "1.90.0.bcr.1") bazel_dep(name = "boost.asio", version = "1.90.0.bcr.1") diff --git a/README.md b/README.md index fa2a495..5b741fb 100644 --- a/README.md +++ b/README.md @@ -9,7 +9,7 @@ and [Boost.Asio](https://www.boost.org/doc/libs/release/libs/asio/), managed wit ## Features - **HTTP Server** — Multi-threaded, async I/O server powered by Boost.Asio strand-based concurrency -- **WebSocket Support** — Full WebSocket lifecycle management (onopen / onmessage / onclose / onerror) +- **WebSocket Support** — Dynamic routes, handshake metadata, typed text/binary/control frames, and full lifecycle management - **Routing** — Express-style route registration with path parameters (`/users/:id`), query params, and method specificity sorting - **Controller Pattern** — CRTP-based `BaseController` with `KHTTPD_ROUTE` / `KHTTPD_WSROUTE` macros for clean route @@ -20,6 +20,7 @@ and [Boost.Asio](https://www.boost.org/doc/libs/release/libs/asio/), managed wit - **Interceptors** — Pre-request / post-response middleware pipeline - **Exception Handling** — Type-safe exception dispatcher with per-type handlers - **Chunked Streaming** — Server-sent chunked transfer encoding via `HttpContext::chunked()` +- **Bidirectional HTTP Streaming** — Header-first request routing, fixed-buffer upload/download, proxy backpressure, and Range forwarding - **Cookie Support** — Read / write cookies with configurable `CookieOptions` (path, domain, SameSite, etc.) - **Form & Multipart** — `application/x-www-form-urlencoded` and `multipart/form-data` parsing (file uploads) - **JSON** — Native `boost::json` integration with `get_json()`, `set_body_json()`, `set_body_from()` @@ -32,12 +33,12 @@ and [Boost.Asio](https://www.boost.org/doc/libs/release/libs/asio/), managed wit | Component | Version | |---------------------|----------------| -| Boost | 1.89.0 | -| Boost.Beast | 1.89.0 | -| Boost.Asio | 1.89.0 | -| fmt | 12.0.0 | -| OpenSSL / BoringSSL | 3.3.1 / latest | -| SQLite3 | 3.50.4 | +| Boost | 1.90.0 | +| Boost.Beast | 1.90.0 | +| Boost.Asio | 1.90.0 | +| fmt | 12.1.0 | +| OpenSSL / BoringSSL | 4.0.1 / 0.20260616.0 | +| SQLite3 | 3.53.2 | | Build System | Bazel (bzlmod) | ## Quick Start @@ -49,21 +50,23 @@ In your project's `MODULE.bazel`: ```python http_archive = use_repo_rule("@bazel_tools//tools/build_defs/repo:http.bzl", "http_archive") -bazel_dep(name="platforms", version="1.0.0") -bazel_dep(name="rules_cc", version="0.2.13") -bazel_dep(name="fmt", version="12.0.0") -bazel_dep(name="boost", version="1.89.0.bcr.2") -bazel_dep(name="boost.asio", version="1.89.0.bcr.2") -bazel_dep(name="boost.beast", version="1.89.0.bcr.2") -bazel_dep(name="boost.json", version="1.89.0.bcr.2") -bazel_dep(name="boost.filesystem", version="1.89.0.bcr.2") -bazel_dep(name="boost.url", version="1.89.0.bcr.2") -bazel_dep(name="boringssl", version="0.20251110.0") +bazel_dep(name="platforms", version="1.1.0") +bazel_dep(name="rules_cc", version="0.2.20") +bazel_dep(name="fmt", version="12.1.0") +bazel_dep(name="boost", version="1.90.0.bcr.1") +bazel_dep(name="boost.asio", version="1.90.0.bcr.1") +bazel_dep(name="boost.beast", version="1.90.0.bcr.1") +bazel_dep(name="boost.json", version="1.90.0.bcr.1") +bazel_dep(name="boost.filesystem", version="1.90.0.bcr.1") +bazel_dep(name="boost.url", version="1.90.0.bcr.1") +bazel_dep(name="boost.uuid", version="1.90.0.bcr.1") +bazel_dep(name="boringssl", version="0.20260616.0") +bazel_dep(name="spdlog", version="1.17.0") http_archive( name="khttpd", - strip_prefix="khttpd-0.1.0", - url="https://github.com/ClangTools/khttpd/archive/refs/tags/v0.1.0.tar.gz", + strip_prefix="khttpd-0.3.0", + url="https://github.com/ClangTools/khttpd/archive/refs/tags/v0.3.0.tar.gz", ) ``` @@ -87,6 +90,9 @@ int main() { auto& router = server->get_http_router(); + // Buffered JSON/form routes default to 16 MiB. Configure the limit in bytes. + server->set_max_buffered_request_body_size(32ULL * 1024 * 1024); + // Simple route router.get("/hello", [](khttpd::framework::HttpContext& ctx) { std::string name = ctx.get_query_param("name").value_or("World"); @@ -128,6 +134,8 @@ framework/ ├── io_context_pool.hpp # Asio io_context thread pool ├── context/ │ ├── http_context.hpp/cpp # Request/response abstraction (params, body, cookies, streaming) +│ ├── http_request_stream.hpp # Fixed-buffer inbound request body +│ ├── http_response_stream.hpp # Fixed-buffer outbound response body │ └── websocket_context.hpp/cpp # WebSocket session context (send, attributes) ├── router/ │ ├── http_router.hpp/cpp # Route matching, interceptors, exception dispatch @@ -136,6 +144,8 @@ framework/ │ └── http_controller.hpp # CRTP BaseController + KHTTPD_ROUTE / KHTTPD_WSROUTE macros ├── client/ │ ├── http_client.hpp/cpp # Sync/async HTTP client with SSL +│ ├── http_client_stream.hpp/cpp # Fixed-buffer HTTP streaming client +│ ├── http_proxy_session.hpp/cpp # Bidirectional streaming proxy pump │ └── websocket_client.hpp/cpp # WebSocket client ├── interceptor/ │ └── interceptor.hpp # Pre/Post middleware interface @@ -184,9 +194,17 @@ framework/ ```cpp auto& ws = server->get_websocket_router(); -ws.add_handler("/ws", - [](WebsocketContext& ctx) { /* onopen */ ctx.send("Welcome!"); }, - [](WebsocketContext& ctx) { /* onmessage */ ctx.send("Echo: " + ctx.message, ctx.is_text); }, +ws.add_handler("/gateway/:target", + [](WebsocketContext& ctx) { + auto target = ctx.get_path_param("target"); // may contain multiple path segments + auto token = ctx.get_header("Authorization"); + auto trace = ctx.get_query_param("trace"); + ctx.send("Welcome!"); + }, + [](WebsocketContext& ctx) { + // frame.type preserves text/binary; payload may contain arbitrary bytes. + ctx.send(ctx.frame); + }, [](WebsocketContext& ctx) { /* onclose */ }, [](WebsocketContext& ctx) { /* onerror */ } ); @@ -212,6 +230,49 @@ class MyController : public khttpd::framework::BaseController { MyController::create()->register_routes(server->get_http_router()); ``` +### Streaming HTTP routes and proxying + +Large request and response bodies can bypass `string_body` buffering by using a +stream route. Reads and writes are serialized through fixed-size buffers, so +backpressure is propagated between the downstream and upstream connections. + +```cpp +router.stream("/gateway/upload", boost::beast::http::verb::post, + [](HttpContext& ctx, + std::shared_ptr request, + std::shared_ptr response, + HttpStreamComplete complete) + { + client::HttpClientStream::RequestHead head{ + ctx.method(), "/upload", ctx.get_request().version()}; + for (const auto& field : ctx.get_request()) + head.insert(field.name_string(), field.value()); + + auto proxy = std::make_shared(request, response); + proxy->start("http://upstream.internal/upload", std::move(head)); + }); +``` + +Normal routes remain buffered for JSON/form compatibility and reject request +bodies above the configurable limit (16 MiB by default; use +`Server::set_max_buffered_request_body_size(bytes)`). Stream routes are not +subject to this buffered-body limit and have no size-dependent allocation; +the proxy buffer defaults to 64 KiB and can be configured in its constructor. + +`HttpClientStream` supports both `http://` and `https://` without changing its +fixed-buffer behavior. Its default TLS context verifies the system trust store; +an application can inject an `ssl::context` into `HttpClientStream` or +`HttpProxySession` for private CAs and test certificates. + +### Tests + +Run the complete framework suite, including buffered-body boundaries, streaming +edge cases, proxy cancellation, WebSocket dynamic routing, and frame fidelity: + +```bash +bazel test //framework/... --test_output=errors +``` + ### Interceptors ```cpp diff --git a/doc/advanced.md b/doc/advanced.md index 5670965..ec6fd5f 100644 --- a/doc/advanced.md +++ b/doc/advanced.md @@ -44,6 +44,24 @@ Request → Interceptor1.handle_request → Interceptor2.handle_request → Hand - **前置拦截器**:按注册**正序**执行 - **后置拦截器**:按注册**逆序**执行(洋葱模型) - 任一前置返回 `Stop` → 跳过剩余前置和 handler → 执行全部后置 +- WebSocket Upgrade 在握手前也执行这条链,因此可复用 HTTP 鉴权 + +远程鉴权可覆盖异步入口,完成回调可以从任意线程调用,但必须恰好调用一次: + +```cpp +void async_handle_request(HttpContext& ctx, RequestCompletion complete) override { + auth_client.check(ctx.get_header("Authorization"), + [&ctx, complete = std::move(complete)](bool allowed) mutable { + if (!allowed) { + ctx.set_status(http::status::unauthorized); + ctx.set_body("Unauthorized"); + } + complete(allowed ? InterceptorResult::Continue : InterceptorResult::Stop); + }); +} +``` + +做可信 `X-Forwarded-For` 解析时,应先用 `ctx.peer_endpoint()` 判断直连 peer 是否属于受信代理网段;IP 限流的默认 key 应使用 `ctx.peer_address()`,不能直接信任请求头。 ### 上下文数据传递 @@ -122,9 +140,10 @@ router.set_unknown_exception_handler([](HttpContext& ctx) { auto& ws_router = server->get_websocket_router(); ws_router.add_handler( - "/chat", + "/chat/:room", // on_open [](WebsocketContext& ctx) { + auto room = ctx.get_path_param("room").value_or("lobby"); ctx.send("Welcome to the chat!"); }, // on_message @@ -142,6 +161,28 @@ ws_router.add_handler( ); ``` +WebSocket 路由和 HTTP 路由一样支持动态参数。最后一个参数可以包含 `/`,因此 `/gateway/:target` 能匹配 `/gateway/orders/ws/v1`。静态路由始终优先于动态路由。 + +### 握手信息与帧类型 + +```cpp +ws_router.add_handler("/gateway/:target", + [](WebsocketContext& ctx) { + const auto& request = ctx.handshake(); // target/path/headers/query/subprotocols + auto authorization = ctx.get_header("Authorization"); + auto cookies = ctx.get_headers("Cookie"); // 保留重复字段 + auto trace = ctx.get_query_param("trace"); + }, + [](WebsocketContext& ctx) { + if (ctx.frame.type == WebsocketFrameType::binary) { + // payload 是原始字节,可包含 \0;不会转成文本帧。 + ctx.send(ctx.frame); + } + }); +``` + +`WebsocketFrame` 还可表示 ping、pong 和 close;close 帧包含 `close_code` 与 `close_reason`。 + ### 广播消息 ```cpp @@ -188,6 +229,8 @@ ChatController::create()->register_routes(ws_router); ## 分块流式响应 +`HttpContext::chunked()` 适合服务端逐块生成响应,但请求体仍属于普通缓冲模型。需要流式上传、下载或代理时,应使用下一节的 `router.stream()`。 + ```cpp router.get("/stream/:count", [](HttpContext& ctx) { int count = std::stoi(ctx.get_path_param("count").value_or("10")); @@ -209,6 +252,42 @@ router.get("/stream/:count", [](HttpContext& ctx) { --- +## 双向 HTTP 流与大文件代理 + +普通 JSON、form 和 multipart handler 需要完整请求体,默认最大 16 MiB。可以全局调整: + +```cpp +server->set_max_buffered_request_body_size(32ULL * 1024 * 1024); +``` + +带 `Content-Length` 的超限请求会在读取 body、发送 `100 Continue` 之前返回 413;chunked 请求在累计越界时返回 413。大文件不要简单调高该上限,应注册流式路由: + +```cpp +router.stream("/gateway/:target", http::verb::post, + [](HttpContext& ctx, + std::shared_ptr request, + std::shared_ptr response, + HttpStreamComplete) + { + client::HttpClientStream::RequestHead head{ + ctx.method(), "/", ctx.get_request().version()}; + for (const auto& field : ctx.get_request()) + head.insert(field.name_string(), field.value()); + + auto proxy = std::make_shared( + request, response, 64 * 1024); + proxy->start("http://upstream.internal/upload", std::move(head)); + }); +``` + +请求和响应各自只保持固定缓冲区,前一次写完成后才读取下一块。该模型支持 Content-Length、chunked、206 Range 响应和 hop-by-hop header 过滤。任一侧错误或取消时会联动取消其他方向。 + +若上游已提前拒绝请求,可调用 `request->cancel_read()` 或 `response->cancel_request_body()`。这只终止入站请求体读取,响应仍可正常写回;为避免未消费字节污染下一条请求,该连接不会再 keep-alive。 + +`HttpClientStream` 和 `HttpProxySession` 同时支持 `http://` 与 `https://`,两种传输都保持固定缓冲模型。默认 TLS context 使用系统信任库并校验证书;私有 CA 可通过接受 `ssl::context&` 的构造函数注入。 + +--- + ## Cron 定时任务 ### Lambda 任务 diff --git a/doc/api-reference.md b/doc/api-reference.md index f8e39f7..67db174 100644 --- a/doc/api-reference.md +++ b/doc/api-reference.md @@ -21,6 +21,8 @@ Server(const tcp::endpoint& endpoint, std::string web_root, int num_threads = 1) | `get_http_router()` | `HttpRouter&` | 获取 HTTP 路由器引用,用于注册路由 | | `get_websocket_router()` | `WebsocketRouter&` | 获取 WebSocket 路由器引用 | | `add_interceptor(interceptor)` | `void` | 添加全局请求/响应拦截器 | +| `set_max_buffered_request_body_size(bytes)` | `void` | 设置普通缓冲路由的请求体上限;默认 16 MiB,仅影响之后建立的连接 | +| `get_max_buffered_request_body_size()` | `std::uint64_t` | 获取当前普通缓冲路由请求体上限 | | `run()` | `void` | 启动服务器(阻塞调用,直到收到 SIGINT/SIGTERM) | | `stop()` | `void` | 停止服务器,关闭 acceptor 和线程池 | @@ -58,6 +60,8 @@ server->run(); | `get_path_param(key)` | `std::optional` | 路径参数,如 `/users/:id` 中的 `id` | | `get_header(name)` | `std::optional` | 请求头(支持 `http::field` 枚举和字符串) | | `get_headers(name)` | `std::optional>` | 同名请求头列表 | +| `peer_endpoint()` | `const std::optional&` | TCP 真实对端;不受 `X-Forwarded-For` 伪造影响 | +| `peer_address()` | `std::optional` | TCP 真实对端地址,适合可信代理判断和 IP 限流 | ### Cookie 操作 @@ -142,9 +146,22 @@ struct MultipartFile { | `put(path, handler)` | 注册 PUT 路由 | | `del(path, handler)` | 注册 DELETE 路由 | | `options(path, handler)` | 注册 OPTIONS 路由 | +| `stream(path, method, handler)` | 注册在读取完整 body 前分发的流式路由 | +| `async_route(path, method, handler)` | 注册异步路由;handler 完成时调用一次 `complete()` | `handler` 签名:`void(HttpContext&)` +流式 `handler` 签名: + +```cpp +void(HttpContext&, + std::shared_ptr, + std::shared_ptr, + HttpStreamComplete) +``` + +普通 handler 仍使用 `string_body`,超过 Server 配置上限会返回 413。流式路由使用固定缓冲区读取,不受该缓冲上限约束。 + ### 路由语法 | 语法 | 示例 | 匹配 | @@ -190,7 +207,7 @@ using WebsocketErrorHandler = std::function; | 方法 | 说明 | |------|------| -| `add_handler(path, on_open, on_message, on_close, on_error)` | 注册 WebSocket 路径的所有生命周期处理器。`path` 为精确匹配(不支持动态参数) | +| `add_handler(path, on_open, on_message, on_close, on_error)` | 注册 WebSocket 生命周期处理器;支持 `/gateway/:target` 动态参数,最后一个参数可匹配多层路径 | | `dispatch_open(path, ctx)` | 分发 open 事件 | | `dispatch_message(path, ctx)` | 分发 message 事件 | | `dispatch_close(path, ctx)` | 分发 close 事件 | @@ -209,6 +226,7 @@ using WebsocketErrorHandler = std::function; | `is_text` | `bool` | 消息是否为文本(仅 message 事件有效) | | `error_code` | `beast::error_code` | 错误码(仅 error/close 事件有效) | | `path` | `std::string` | 连接路径 | +| `frame` | `WebsocketFrame` | 当前帧的类型、原始 payload、关闭码和关闭原因 | | `session_weak_ptr` | `weak_ptr` | 会话的弱引用 | ### 方法 @@ -216,6 +234,11 @@ using WebsocketErrorHandler = std::function; | 方法 | 说明 | |------|------| | `send(msg, is_text)` | 发送消息给客户端 | +| `send(frame)` | 按 `WebsocketFrameType` 发送 text、binary、ping、pong 或 close 帧 | +| `handshake()` | 获取原始 target、路径、重复 headers、query 参数和客户端请求的 subprotocol 列表 | +| `get_header(name)` / `get_headers(name)` | 大小写不敏感地读取握手 header;复数版本保留重复字段 | +| `get_query_param(key)` | 读取握手查询参数 | +| `get_path_param(key)` | 读取 WebSocket 动态路由参数 | | `set_attribute(key, value)` | 存储扩展数据 | | `get_attribute_as(key)` | 获取并类型转换扩展数据 | @@ -266,9 +289,11 @@ enum class InterceptorResult { Continue, Stop }; | 方法 | 默认返回 | 调用时机 | |------|----------|----------| | `handle_request(ctx)` | `Continue` | 路由处理前,按添加顺序执行 | +| `async_handle_request(ctx, complete)` | 调用同步 `handle_request` | 异步前置检查;完成时调用一次 `complete(result)` | | `handle_response(ctx)` | 空 | 响应生成后,按添加**逆序**执行 | 返回 `Stop` 时中断后续拦截器和路由处理器,直接执行后置拦截器。 +HTTP 与 WebSocket Upgrade 都会执行同一条前置拦截器链。 --- @@ -357,6 +382,56 @@ public: --- +## HTTP 流式 API + +### HttpRequestStream + +| 方法 | 说明 | +|------|------| +| `async_read_some(buffer, callback)` | 将下一段请求体读入调用方缓冲区;回调参数为 `(ec, bytes, done)` | +| `cancel_read()` | 只取消请求体读取;响应通道仍可发送,连接随后以非 keep-alive 结束 | +| `cancel()` | 取消读取并关闭对应连接 | + +同一个方向必须等待前一次回调完成后再发起下一次读取。 + +### HttpResponseStream + +| 方法 | 说明 | +|------|------| +| `async_start(head, callback)` | 发送响应头;未指定 Content-Length/Chunked 时自动使用 chunked | +| `async_write_some(buffer, callback)` | 发送一段响应体 | +| `async_finish(callback)` | 完成响应体并写入终止块(如需要) | +| `cancel_request_body()` | 只取消配对的请求体读取,保留响应通道 | +| `cancel()` | 取消下游响应 | + +### HttpClientStream + +| 方法 | 说明 | +|------|------| +| `async_start(url, head, callback)` | 连接上游并发送请求头 | +| `async_write_some(buffer, callback)` | 发送一段请求体 | +| `async_finish_request(callback)` | 完成请求体 | +| `async_read_response_head(callback)` | 读取上游响应头 | +| `async_read_some(buffer, callback)` | 固定缓冲区读取响应体 | +| `cancel()` | 取消解析、连接和未完成 I/O | + +流式客户端支持 `http://` 和 `https://`,TLS 传输仍使用同一套固定缓冲 serializer/parser。默认 context 校验系统信任库;也可通过 `HttpClientStream(ssl_context)` 或 `HttpClientStream(ioc, ssl_context)` 注入私有 CA 配置。 +客户端会跳过连续的 100/103 等 informational response,向调用方交付最终响应;HEAD 按响应头完成,不等待 `Content-Length` 指示的正文。 + +### HttpProxySession + +`HttpProxySession` 把入站请求流、上游 `HttpClientStream` 和下游响应流串联起来。请求和响应方向都遵循“读一块、写一块、写完再读下一块”,从而形成自然背压。 + +```cpp +auto proxy = std::make_shared( + request_stream, response_stream, 64 * 1024); +proxy->start(upstream_url, std::move(request_head), complete_callback); +``` + +代理会过滤 hop-by-hop headers,透明保留 Range、Content-Range、Accept-Ranges、ETag 等端到端字段,并在任一侧失败时取消其余方向。 + +--- + ## HttpClient ### 构造函数 @@ -419,6 +494,10 @@ API_CALL(http::verb::get, "/users/:id", get_user, | `send(message)` | 发送消息(线程安全) | | `close()` | 关闭连接 | | `set_header(key, value)` | 设置握手头 | +| `set_subprotocols(protocols)` | 设置 `Sec-WebSocket-Protocol` 请求列表 | +| `negotiated_subprotocol()` | 返回握手响应中服务端最终选择的子协议 | | `set_on_message(handler)` | 设置消息回调 | +| `set_on_frame(handler)` | 设置保留 text/binary/control 类型的帧回调 | | `set_on_error(handler)` | 设置错误回调 | | `set_on_close(handler)` | 设置关闭回调 | +| `send(frame)` | 发送带明确类型和关闭信息的 `WebsocketFrame` | diff --git a/doc/architecture.md b/doc/architecture.md index 87742b3..e1d7fdb 100644 --- a/doc/architecture.md +++ b/doc/architecture.md @@ -52,11 +52,17 @@ Client Request │ ▼ ┌─────────────┐ -│ HttpSession │ async_read 读取 HTTP 请求 +│ HttpSession │ async_read_header 先读取请求行和 headers └──────┬──────┘ │ ▼ ┌──────────────────┐ +│ Header Route │ 流式路由 → buffer_body 固定缓冲读取 +│ Decision │ 普通路由 → 在可配置上限内累计 body +└──────┬───────────┘ + │ + ▼ +┌──────────────────┐ │ Pre-Interceptor │ 链式执行,可中断 │ Chain │ 任一返回 Stop → 跳到 Post-Interceptor └──────┬───────────┘ @@ -122,21 +128,33 @@ HTTP Request (Upgrade header) - 使用 **正则表达式** 解析动态路径参数 - 按**特异性排序**:字面段多的路由优先,同数量下动态段少的优先 -- WebSocket 路由器仅支持精确路径匹配(不支持动态参数) +- HTTP 与 WebSocket 路由均支持 `:param` 动态参数和特异性排序 +- WebSocket 最后一个动态参数可捕获剩余多层路径,例如 `/gateway/:target` 匹配 `/gateway/orders/ws/v1` +- WebSocket 路由表使用共享锁读取、独占锁更新;热更新只影响后续分发 ### 3. 响应发送 - 普通响应:`beast::async_write` 异步发送 - Chunked 流式响应:通过 `HttpContext::chunked()` 启用,内部将同步写入转换为异步写链 +- 双向 HTTP 流:`request_parser` 与 `response` 使用调用方固定缓冲区;每次 write 完成后才继续 read,形成背压 - 静态文件:使用 Beast 的 `file_body` 零拷贝发送 ### 4. 内存管理 - `HttpSession` / `WebsocketSession` 由 `shared_ptr` 管理 -- `HttpContext` / `WebsocketContext` 为临时栈对象,生命周期仅限于单个请求/事件 +- 普通 JSON/form 路由保留完整 body,但由 `Server::set_max_buffered_request_body_size()` 控制上限(默认 16 MiB) +- 流式路由、客户端和代理使用固定大小缓冲区,内存不随单个 body 大小线性增长 +- 流对象持有对应会话,异步操作完成或取消前不会悬空 - 拦截器和异常处理器存储为 `shared_ptr` 在路由器中 -### 5. 静态文件安全 +### 5. WebSocket 帧与握手元数据 + +- `WebsocketFrame` 统一表示 text、binary、ping、pong 和 close 帧 +- binary payload 使用 `std::string` 作为字节容器,但不会转换为文本,类型由 `frame.type` 明确保留 +- `WebsocketHandshakeRequest` 保留原始 target、path、重复 headers、query 参数和 subprotocol 请求列表 +- 写操作通过每连接串行队列执行,避免并发 `async_write` + +### 6. 静态文件安全 - 使用 `boost::filesystem::canonical()` 规范化路径,防止 `../` 目录遍历 - 规范化后校验路径仍在 `web_root` 内 @@ -160,6 +178,8 @@ framework/ ├── io_context_pool.hpp # IO 线程池(单例) ├── context/ │ ├── http_context.hpp/cpp # HTTP 请求/响应上下文 +│ ├── http_request_stream.hpp # 入站请求流接口 +│ ├── http_response_stream.hpp # 下游响应流接口 │ └── websocket_context.hpp/cpp # WebSocket 上下文 ├── router/ │ ├── http_router.hpp/cpp # HTTP 路由匹配与分发 @@ -168,6 +188,8 @@ framework/ │ └── http_controller.hpp # CRTP Controller + 路由宏 ├── client/ │ ├── http_client.hpp/cpp # HTTP 客户端(同步/异步) +│ ├── http_client_stream.hpp/cpp # 固定缓冲 HTTP 流客户端 +│ ├── http_proxy_session.hpp/cpp # 双向代理泵与背压 │ ├── websocket_client.hpp/cpp # WebSocket 客户端 │ └── macros.hpp # API_CALL 宏 ├── interceptor/ diff --git a/doc/http-client.md b/doc/http-client.md index cee296d..586c244 100644 --- a/doc/http-client.md +++ b/doc/http-client.md @@ -171,6 +171,32 @@ KHTTPD_API_CLIENT_END() --- +## 流式 HTTP 客户端 + +`HttpClient` 返回 `response`,适合普通 API。大请求或响应应使用 `HttpClientStream`,由调用方提供固定大小缓冲区: + +```cpp +#include "framework/client/http_client_stream.hpp" + +auto stream = std::make_shared(ioc); +HttpClientStream::RequestHead head{http::verb::post, "/upload", 11}; +head.chunked(true); + +stream->async_start("http://storage.internal/upload", std::move(head), + [stream](beast::error_code ec) { + // async_write_some(...) -> async_finish_request(...) + // -> async_read_response_head(...) -> async_read_some(...) + }); +``` + +同一方向必须串行调用:等待当前 read/write 回调后再提交下一块。这既限制 in-flight 数据,也让 TCP 自然提供背压。调用 `cancel()` 会取消解析、连接和未完成 I/O。 + +流式客户端同时支持 `http://` 和 `https://`,TLS 不会回退到全量缓存。默认构造函数使用系统信任库并校验证书;私有 CA 或测试环境可以使用 `HttpClientStream(ioc, ssl_context)` 注入自定义 context。 + +缓冲和流式客户端都会连续消费上游 `100 Continue`、`103 Early Hints` 等 1xx 响应,并只返回最终响应。HEAD 请求在最终响应头后即完成,即使响应包含非零 `Content-Length` 也不会等待正文。 + +--- + ## WebSocket 客户端 ### 基本使用 @@ -187,6 +213,16 @@ ws->set_on_message([](const std::string& msg) { fmt::print("Received: {}\n", msg); }); +// 需要保留 text/binary/control 类型时使用帧回调。 +ws->set_on_frame([](const WebsocketFrame& frame) { + if (frame.type == WebsocketFrameType::binary) { + fmt::print("binary bytes: {}\n", frame.payload.size()); + } +}); + +// 请求子协议;连接后可读取服务端最终选择的协议。 +ws->set_subprotocols({"chat.v1", "chat.v2"}); + ws->set_on_error([](beast::error_code ec) { if (ec != boost::asio::error::operation_aborted) { fmt::print(stderr, "WS Error: {}\n", ec.message()); @@ -198,8 +234,9 @@ ws->set_on_close([]() { }); // 连接 -ws->connect("wss://echo.websocket.org", [](beast::error_code ec) { +ws->connect("wss://echo.websocket.org", [ws](beast::error_code ec) { if (!ec) { + fmt::print("protocol: {}\n", ws->negotiated_subprotocol()); fmt::print("Connected!\n"); } }); @@ -215,6 +252,9 @@ ws->send("Hello, server!"); ws->send("Message 1"); ws->send("Message 2"); ws->send("Message 3"); + +// 发送二进制帧,payload 可以包含 NUL 字节。 +ws->send({WebsocketFrameType::binary, std::string("\x00\x01", 2)}); ``` ### 完整示例:Echo 客户端 diff --git a/doc/index.md b/doc/index.md index 6885638..2dd011f 100644 --- a/doc/index.md +++ b/doc/index.md @@ -7,8 +7,8 @@ | [快速开始指南](quick-start.md) | 10 分钟搭建第一个 khttpd 服务 | | [API 参考文档](api-reference.md) | 完整 API 方法签名和参数说明 | | [架构指南](architecture.md) | 框架设计、请求流程、线程模型、扩展点 | -| [高级功能](advanced.md) | 拦截器、异常处理、WebSocket、Cron、DI 容器、Cookie | -| [HTTP 与 WebSocket 客户端](http-client.md) | 内置客户端 API、API_CALL 宏、WebSocket 客户端 | +| [高级功能](advanced.md) | 拦截器、异常处理、动态 WebSocket、双向 HTTP 流、Cron、DI、Cookie | +| [HTTP 与 WebSocket 客户端](http-client.md) | 缓冲/流式 HTTP 客户端、API_CALL 宏、WebSocket 帧客户端 | ## 按主题查找 @@ -32,6 +32,8 @@ - 文件上传 → [API 参考](api-reference.md#表单与文件上传) - 设置响应 → [API 参考](api-reference.md#响应设置) - 分块流式响应 → [高级功能](advanced.md#分块流式响应) +- 大文件双向流式代理 → [高级功能](advanced.md#双向-http-流与大文件代理) +- 普通请求体 413 上限 → [API 参考](api-reference.md#server) - Cookie 操作 → [高级功能](advanced.md#cookie-操作) ### 中间件 @@ -50,6 +52,7 @@ ### 客户端 - HTTP 客户端 → [HTTP 客户端](http-client.md#http-客户端) +- 流式 HTTP 客户端 → [HTTP 客户端](http-client.md#流式-http-客户端) - Oat++ 风格 API 定义 → [HTTP 客户端](http-client.md#oat-风格-api-定义) - 多 Host 权重分发 → [HTTP 客户端](http-client.md#多-host-权重分发) - API_CALL 宏 → [HTTP 客户端](http-client.md#api_call-宏自动生成客户端方法) diff --git a/doc/quick-start.md b/doc/quick-start.md index 101ae9f..0150662 100644 --- a/doc/quick-start.md +++ b/doc/quick-start.md @@ -35,23 +35,24 @@ mkdir my-khttpd-app && cd my-khttpd-app ```python module(name = "my-khttpd-app", version = "0.1.0") -bazel_dep(name = "platforms", version = "1.0.0") -bazel_dep(name = "rules_cc", version = "0.2.13") -bazel_dep(name = "fmt", version = "12.0.0") -bazel_dep(name = "boost", version = "1.89.0.bcr.2") -bazel_dep(name = "boost.asio", version = "1.89.0.bcr.2") -bazel_dep(name = "boost.beast", version = "1.89.0.bcr.2") -bazel_dep(name = "boost.json", version = "1.89.0.bcr.2") -bazel_dep(name = "boost.filesystem", version = "1.89.0.bcr.2") -bazel_dep(name = "boost.url", version = "1.89.0.bcr.2") -bazel_dep(name = "boost.uuid", version = "1.89.0.bcr.2") -bazel_dep(name = "boringssl", version = "0.20251110.0") +bazel_dep(name = "platforms", version = "1.1.0") +bazel_dep(name = "rules_cc", version = "0.2.20") +bazel_dep(name = "fmt", version = "12.1.0") +bazel_dep(name = "boost", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.asio", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.beast", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.json", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.filesystem", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.url", version = "1.90.0.bcr.1") +bazel_dep(name = "boost.uuid", version = "1.90.0.bcr.1") +bazel_dep(name = "boringssl", version = "0.20260616.0") +bazel_dep(name = "spdlog", version = "1.17.0") http_archive = use_repo_rule("@bazel_tools//tools/build_defs/repo:http.bzl", "http_archive") http_archive( name = "khttpd", - strip_prefix = "khttpd-0.1.0", - url = "https://github.com/ClangTools/khttpd/archive/refs/tags/v0.1.0.tar.gz", + strip_prefix = "khttpd-0.3.0", + url = "https://github.com/ClangTools/khttpd/archive/refs/tags/v0.3.0.tar.gz", ) ``` @@ -88,6 +89,9 @@ int main() auto server = std::make_shared( tcp::endpoint{address, port}, "web_root", threads); + // 普通 JSON/form 路由的 body 默认最多 16 MiB,可按服务需要调整。 + server->set_max_buffered_request_body_size(32ULL * 1024 * 1024); + auto& router = server->get_http_router(); // 简单路由 @@ -142,6 +146,12 @@ curl -X POST -H "Content-Type: application/json" \ # {"msg":"hi"} ``` +框架开发与回归测试: + +```bash +bazel test //framework/... --test_output=errors +``` + ## 下一步 - [API 文档](api-reference.md) — 完整 API 参考 diff --git a/example/MODULE.bazel b/example/MODULE.bazel index 1027494..30cc3eb 100644 --- a/example/MODULE.bazel +++ b/example/MODULE.bazel @@ -7,7 +7,7 @@ bazel_dep(name = "boost", version = "1.89.0.bcr.2") bazel_dep(name = "boost.asio", version = "1.89.0.bcr.2") bazel_dep(name = "boost.mysql", version = "1.89.0.bcr.2") bazel_dep(name = "spdlog", version = "1.17.0") -bazel_dep(name = "khttpd", version = "0.2.0") +bazel_dep(name = "khttpd", version = "0.3.0") local_path_override( module_name = "khttpd", path = "..", diff --git a/framework/client/http_client.cpp b/framework/client/http_client.cpp index e8cf60f..e8b5440 100644 --- a/framework/client/http_client.cpp +++ b/framework/client/http_client.cpp @@ -3,6 +3,7 @@ #include #include #include +#include #include "io_context_pool.hpp" namespace khttpd::framework::client @@ -27,7 +28,7 @@ namespace khttpd::framework::client protected: HttpClient::ResponseCallback callback_; http::request req_; - http::response res_; + std::optional> response_parser_; beast::flat_buffer buffer_; std::chrono::seconds timeout_; std::atomic completed_{false}; @@ -120,7 +121,15 @@ namespace khttpd::framework::client boost::ignore_unused(bytes_transferred); if (ec) return on_fail(ec, "write"); - http::async_read(stream_, buffer_, res_, + read_response(); + } + + void read_response() + { + response_parser_.emplace(); + response_parser_->body_limit((std::numeric_limits::max)()); + response_parser_->skip(req_.method() == http::verb::head); + http::async_read(stream_, buffer_, *response_parser_, beast::bind_front_handler(&HttpSession::on_read, get_shared())); } @@ -129,9 +138,13 @@ namespace khttpd::framework::client boost::ignore_unused(bytes_transferred); if (ec) return on_fail(ec, "read"); + auto response = response_parser_->release(); + if (response.result_int() >= 100 && response.result_int() < 200 && + response.result() != http::status::switching_protocols) + return read_response(); beast::error_code ignored; stream_.socket().shutdown(tcp::socket::shutdown_both, ignored); - complete({}, std::move(res_)); + complete({}, std::move(response)); } }; @@ -210,7 +223,15 @@ namespace khttpd::framework::client { boost::ignore_unused(bytes_transferred); if (ec) return on_fail(ec, "write"); - http::async_read(stream_, buffer_, res_, + read_response(); + } + + void read_response() + { + response_parser_.emplace(); + response_parser_->body_limit((std::numeric_limits::max)()); + response_parser_->skip(req_.method() == http::verb::head); + http::async_read(stream_, buffer_, *response_parser_, beast::bind_front_handler(&HttpsSession::on_read, get_shared())); } @@ -219,14 +240,20 @@ namespace khttpd::framework::client boost::ignore_unused(bytes_transferred); if (ec) { - if ((ec == ssl::error::stream_truncated || ec == net::error::eof) && res_.result_int() != 0) + if ((ec == ssl::error::stream_truncated || ec == net::error::eof) && + response_parser_ && response_parser_->is_done()) { - complete({}, std::move(res_)); + complete({}, response_parser_->release()); return; } return on_fail(ec, "read"); } + auto response = response_parser_->release(); + if (response.result_int() >= 100 && response.result_int() < 200 && + response.result() != http::status::switching_protocols) + return read_response(); + final_response_ = std::move(response); stream_.async_shutdown(beast::bind_front_handler(&HttpsSession::on_shutdown, get_shared())); } @@ -234,8 +261,10 @@ namespace khttpd::framework::client { if (ec == net::error::eof || ec == ssl::error::stream_truncated) ec = {}; - complete(ec, std::move(res_)); + complete(ec, std::move(final_response_)); } + + http::response final_response_; }; // 1. 傻瓜式:全局 IO + 默认 SSL diff --git a/framework/client/http_client_stream.cpp b/framework/client/http_client_stream.cpp new file mode 100644 index 0000000..bd78db6 --- /dev/null +++ b/framework/client/http_client_stream.cpp @@ -0,0 +1,299 @@ +#include "http_client_stream.hpp" + +#include +#include +#include +#include +#include +#include "io_context_pool.hpp" + +namespace khttpd::framework::client +{ + namespace beast = boost::beast; + namespace http = beast::http; + namespace net = boost::asio; + namespace ssl = net::ssl; + using tcp = net::ip::tcp; + + namespace + { + std::shared_ptr default_ssl_context() + { + auto context = std::make_shared(ssl::context::tls_client); + context->set_default_verify_paths(); + context->set_verify_mode(ssl::verify_peer); + return context; + } + + std::shared_ptr borrowed_ssl_context(ssl::context& context) + { + return {&context, [](ssl::context*) {}}; + } + } + + struct HttpClientStream::Impl : std::enable_shared_from_this + { + using TlsStream = beast::ssl_stream; + + net::any_io_executor executor; + tcp::resolver resolver; + std::shared_ptr ssl_context; + std::unique_ptr plain_stream; + std::unique_ptr tls_stream; + beast::flat_buffer read_buffer; + http::request request; + std::optional> request_serializer; + std::optional> response_parser; + http::verb request_method = http::verb::unknown; + std::string host; + std::string port; + bool use_tls = false; + bool started = false; + bool request_finished = false; + + Impl(net::io_context& ioc, std::shared_ptr context) + : executor(net::make_strand(ioc)), resolver(executor), ssl_context(std::move(context)) + { + } + + void start(const std::string& url, RequestHead head, Callback callback) + { + net::post(executor, [self = shared_from_this(), url, head = std::move(head), + callback = std::move(callback)]() mutable + { self->start_on_executor(url, std::move(head), std::move(callback)); }); + } + + void start_on_executor(const std::string& url, RequestHead head, Callback callback) + { + const auto parsed = boost::urls::parse_uri(url); + if (!parsed || (parsed->scheme() != "http" && parsed->scheme() != "https")) + return callback(make_error_code(boost::system::errc::operation_not_supported)); + + use_tls = parsed->scheme() == "https"; + host = parsed->host(); + port = parsed->port().empty() ? (use_tls ? "443" : "80") : std::string(parsed->port()); + std::string target(parsed->encoded_target()); + if (target.empty()) target = "/"; + else if (target.front() == '?') target.insert(target.begin(), '/'); + request.method(head.method()); + request_method = head.method(); + request.target(target); + request.version(head.version()); + request.keep_alive(head.keep_alive()); + for (const auto& field : head) request.insert(field.name_string(), field.value()); + const bool default_port = (use_tls && port == "443") || (!use_tls && port == "80"); + request.set(http::field::host, default_port ? host : host + ":" + port); + request.body().more = true; + request_serializer.emplace(request); + + if (use_tls) + { + tls_stream = std::make_unique(executor, *ssl_context); + if (!SSL_set_tlsext_host_name(tls_stream->native_handle(), host.c_str())) + return callback(beast::error_code(static_cast(::ERR_get_error()), net::error::get_ssl_category())); + tls_stream->set_verify_callback(ssl::host_name_verification(host)); + } + else + { + plain_stream = std::make_unique(executor); + } + + resolver.async_resolve(host, port, [self = shared_from_this(), callback = std::move(callback)] + (beast::error_code ec, tcp::resolver::results_type results) mutable + { + if (ec) return callback(ec); + self->async_connect(std::move(results), std::move(callback)); + }); + } + + void async_connect(tcp::resolver::results_type results, Callback callback) + { + if (use_tls) + { + beast::get_lowest_layer(*tls_stream).async_connect(results, + [self = shared_from_this(), callback = std::move(callback)] + (beast::error_code ec, const tcp::endpoint&) mutable + { + if (ec) return callback(ec); + self->tls_stream->async_handshake(ssl::stream_base::client, + [self, callback = std::move(callback)](beast::error_code handshake_ec) mutable + { + if (handshake_ec) return callback(handshake_ec); + self->async_write_header(*self->tls_stream, std::move(callback)); + }); + }); + } + else + { + plain_stream->async_connect(results, + [self = shared_from_this(), callback = std::move(callback)] + (beast::error_code ec, const tcp::endpoint&) mutable + { + if (ec) return callback(ec); + self->async_write_header(*self->plain_stream, std::move(callback)); + }); + } + } + + template + void async_write_header(Stream& stream, Callback callback) + { + http::async_write_header(stream, *request_serializer, + [self = shared_from_this(), callback = std::move(callback)] + (beast::error_code ec, std::size_t) mutable + { + self->started = !ec; + callback(ec); + }); + } + + void write(net::const_buffer source, Callback callback) + { + net::post(executor, [self = shared_from_this(), source, callback = std::move(callback)]() mutable + { self->write_on_executor(source, std::move(callback)); }); + } + + void write_on_executor(net::const_buffer source, Callback callback) + { + if (!started || request_finished) return callback(net::error::operation_aborted); + request.body().data = const_cast(source.data()); + request.body().size = source.size(); + request.body().more = true; + if (use_tls) async_write_body(*tls_stream, std::move(callback)); + else async_write_body(*plain_stream, std::move(callback)); + } + + template + void async_write_body(Stream& stream, Callback callback) + { + http::async_write(stream, *request_serializer, + [callback = std::move(callback)](beast::error_code ec, std::size_t) mutable + { if (ec == http::error::need_buffer) ec = {}; callback(ec); }); + } + + void finish(Callback callback) + { + net::post(executor, [self = shared_from_this(), callback = std::move(callback)]() mutable + { self->finish_on_executor(std::move(callback)); }); + } + + void finish_on_executor(Callback callback) + { + if (!started || request_finished) return callback(net::error::operation_aborted); + request_finished = true; + if (request_serializer->is_done()) return callback({}); + request.body().data = nullptr; + request.body().size = 0; + request.body().more = false; + if (use_tls) async_write_body(*tls_stream, std::move(callback)); + else async_write_body(*plain_stream, std::move(callback)); + } + + void read_head(ResponseHeadCallback callback) + { + net::post(executor, [self = shared_from_this(), callback = std::move(callback)]() mutable + { self->read_head_on_executor(std::move(callback)); }); + } + + void read_head_on_executor(ResponseHeadCallback callback) + { + if (!started) return callback(net::error::operation_aborted, {}); + if (use_tls) async_read_head(*tls_stream, std::move(callback)); + else async_read_head(*plain_stream, std::move(callback)); + } + + template + void async_read_head(Stream& stream, ResponseHeadCallback callback) + { + response_parser.emplace(); + response_parser->body_limit((std::numeric_limits::max)()); + response_parser->skip(request_method == http::verb::head); + http::async_read_header(stream, read_buffer, *response_parser, + [self = shared_from_this(), callback = std::move(callback)](beast::error_code ec, std::size_t) mutable + { + ResponseHead head; + if (!ec) + { + const auto& source = self->response_parser->get(); + head.result(source.result()); + head.version(source.version()); + head.keep_alive(source.keep_alive()); + for (const auto& field : source) head.insert(field.name_string(), field.value()); + if (head.result_int() >= 100 && head.result_int() < 200 && + head.result() != http::status::switching_protocols) + { + if (self->use_tls) return self->async_read_head(*self->tls_stream, std::move(callback)); + return self->async_read_head(*self->plain_stream, std::move(callback)); + } + } + callback(ec, std::move(head)); + }); + } + + void read(net::mutable_buffer target, ReadCallback callback) + { + net::post(executor, [self = shared_from_this(), target, callback = std::move(callback)]() mutable + { self->read_on_executor(target, std::move(callback)); }); + } + + void read_on_executor(net::mutable_buffer target, ReadCallback callback) + { + if (!started) return callback(net::error::operation_aborted, 0, true); + if (!response_parser) return callback(net::error::operation_aborted, 0, true); + if (response_parser->is_done()) return callback({}, 0, true); + auto& body = response_parser->get().body(); + body.data = target.data(); + body.size = target.size(); + if (use_tls) async_read_body(*tls_stream, target.size(), std::move(callback)); + else async_read_body(*plain_stream, target.size(), std::move(callback)); + } + + template + void async_read_body(Stream& stream, std::size_t capacity, ReadCallback callback) + { + http::async_read_some(stream, read_buffer, *response_parser, + [self = shared_from_this(), capacity, callback = std::move(callback)] + (beast::error_code ec, std::size_t) mutable + { + if (ec == http::error::need_buffer) ec = {}; + const auto produced = capacity - self->response_parser->get().body().size; + callback(ec, produced, self->response_parser->is_done()); + }); + } + + void cancel() + { + beast::error_code ignored; + resolver.cancel(); + if (tls_stream) + { + beast::get_lowest_layer(*tls_stream).cancel(); + beast::get_lowest_layer(*tls_stream).socket().close(ignored); + } + if (plain_stream) + { + plain_stream->cancel(); + plain_stream->socket().close(ignored); + } + } + }; + + HttpClientStream::HttpClientStream() + : impl_(std::make_shared(IoContextPool::instance().get_io_context(), default_ssl_context())) {} + HttpClientStream::HttpClientStream(ssl::context& context) + : impl_(std::make_shared(IoContextPool::instance().get_io_context(), borrowed_ssl_context(context))) {} + HttpClientStream::HttpClientStream(net::io_context& ioc) + : impl_(std::make_shared(ioc, default_ssl_context())) {} + HttpClientStream::HttpClientStream(net::io_context& ioc, ssl::context& context) + : impl_(std::make_shared(ioc, borrowed_ssl_context(context))) {} + HttpClientStream::~HttpClientStream() { if (impl_) impl_->cancel(); } + void HttpClientStream::async_start(const std::string& url, RequestHead head, Callback cb) + { impl_->start(url, std::move(head), std::move(cb)); } + void HttpClientStream::async_write_some(net::const_buffer buffer, Callback cb) + { impl_->write(buffer, std::move(cb)); } + void HttpClientStream::async_finish_request(Callback cb) { impl_->finish(std::move(cb)); } + void HttpClientStream::async_read_response_head(ResponseHeadCallback cb) { impl_->read_head(std::move(cb)); } + void HttpClientStream::async_read_some(net::mutable_buffer buffer, ReadCallback cb) + { impl_->read(buffer, std::move(cb)); } + void HttpClientStream::cancel() { impl_->cancel(); } +} diff --git a/framework/client/http_client_stream.hpp b/framework/client/http_client_stream.hpp new file mode 100644 index 0000000..7a099a4 --- /dev/null +++ b/framework/client/http_client_stream.hpp @@ -0,0 +1,40 @@ +#ifndef KHTTPD_FRAMEWORK_CLIENT_HTTP_CLIENT_STREAM_HPP +#define KHTTPD_FRAMEWORK_CLIENT_HTTP_CLIENT_STREAM_HPP + +#include +#include +#include +#include +#include +#include + +namespace khttpd::framework::client +{ + class HttpClientStream : public std::enable_shared_from_this + { + public: + using RequestHead = boost::beast::http::request; + using ResponseHead = boost::beast::http::response; + using Callback = std::function; + using ResponseHeadCallback = std::function; + using ReadCallback = std::function; + + HttpClientStream(); + explicit HttpClientStream(boost::asio::ssl::context& ssl_context); + explicit HttpClientStream(boost::asio::io_context& ioc); + HttpClientStream(boost::asio::io_context& ioc, boost::asio::ssl::context& ssl_context); + ~HttpClientStream(); + void async_start(const std::string& url, RequestHead head, Callback callback); + void async_write_some(boost::asio::const_buffer buffer, Callback callback); + void async_finish_request(Callback callback); + void async_read_response_head(ResponseHeadCallback callback); + void async_read_some(boost::asio::mutable_buffer buffer, ReadCallback callback); + void cancel(); + + private: + struct Impl; + std::shared_ptr impl_; + }; +} + +#endif diff --git a/framework/client/http_proxy_session.cpp b/framework/client/http_proxy_session.cpp new file mode 100644 index 0000000..ff03e3f --- /dev/null +++ b/framework/client/http_proxy_session.cpp @@ -0,0 +1,152 @@ +#include "http_proxy_session.hpp" +#include +#include "io_context_pool.hpp" + +namespace khttpd::framework::client +{ + namespace net = boost::asio; + namespace beast = boost::beast; + + namespace + { + template + void filter_hop_by_hop(Message& message) + { + const bool chunked = message.chunked(); + if (const auto connection = message[beast::http::field::connection]; !connection.empty()) + { + std::string tokens(connection); + std::size_t begin = 0; + while (begin < tokens.size()) + { + const auto end = tokens.find(',', begin); + auto token = tokens.substr(begin, end == std::string::npos ? end : end - begin); + const auto first = token.find_first_not_of(" \t"); + if (first != std::string::npos) + { + token = token.substr(first, token.find_last_not_of(" \t") - first + 1); + message.erase(token); + } + if (end == std::string::npos) break; + begin = end + 1; + } + } + message.erase(beast::http::field::connection); + message.erase(beast::http::field::keep_alive); + message.erase(beast::http::field::proxy_authenticate); + message.erase(beast::http::field::proxy_authorization); + message.erase(beast::http::field::te); + message.erase(beast::http::field::trailer); + message.erase(beast::http::field::transfer_encoding); + message.erase(beast::http::field::upgrade); + message.erase("Proxy-Connection"); + if (chunked) message.chunked(true); + } + } + + HttpProxySession::HttpProxySession(std::shared_ptr inbound, + std::shared_ptr downstream, std::size_t buffer_size) + : HttpProxySession(IoContextPool::instance().get_io_context(), std::move(inbound), + std::move(downstream), buffer_size) {} + + HttpProxySession::HttpProxySession(std::shared_ptr inbound, + std::shared_ptr downstream, + net::ssl::context& ssl_context, std::size_t buffer_size) + : HttpProxySession(IoContextPool::instance().get_io_context(), std::move(inbound), + std::move(downstream), ssl_context, buffer_size) {} + + HttpProxySession::HttpProxySession(net::io_context& ioc, std::shared_ptr inbound, + std::shared_ptr downstream, std::size_t buffer_size) + : inbound_(std::move(inbound)), downstream_(std::move(downstream)), + upstream_(std::make_shared(ioc)), + request_buffer_(buffer_size), response_buffer_(buffer_size) {} + + HttpProxySession::HttpProxySession(net::io_context& ioc, std::shared_ptr inbound, + std::shared_ptr downstream, + net::ssl::context& ssl_context, std::size_t buffer_size) + : inbound_(std::move(inbound)), downstream_(std::move(downstream)), + upstream_(std::make_shared(ioc, ssl_context)), + request_buffer_(buffer_size), response_buffer_(buffer_size) {} + + void HttpProxySession::start(const std::string& url, HttpClientStream::RequestHead head, CompleteCallback callback) + { + complete_ = std::move(callback); + head.erase(boost::beast::http::field::expect); + filter_hop_by_hop(head); + upstream_->async_start(url, std::move(head), [self = shared_from_this()](beast::error_code ec) + { if (ec) self->finish(ec); else self->pump_request(); }); + } + + void HttpProxySession::pump_request() + { + inbound_->async_read_some(net::buffer(request_buffer_), [self = shared_from_this()] + (beast::error_code ec, std::size_t n, bool done) + { + if (ec) return self->finish(ec); + if (n) + { + self->upstream_->async_write_some(net::buffer(self->request_buffer_.data(), n), + [self, done](beast::error_code write_ec) + { if (write_ec) self->finish(write_ec); else if (done) self->finish_request(); else self->pump_request(); }); + } + else if (done) self->finish_request(); + else self->pump_request(); + }); + } + + void HttpProxySession::finish_request() + { + upstream_->async_finish_request([self = shared_from_this()](beast::error_code ec) + { if (ec) self->finish(ec); else self->read_response_head(); }); + } + + void HttpProxySession::read_response_head() + { + upstream_->async_read_response_head([self = shared_from_this()] + (beast::error_code ec, HttpClientStream::ResponseHead head) + { + if (ec) return self->finish(ec); + filter_hop_by_hop(head); + self->downstream_->async_start(std::move(head), [self](beast::error_code start_ec) + { if (start_ec) self->finish(start_ec); else self->pump_response(); }); + }); + } + + void HttpProxySession::pump_response() + { + upstream_->async_read_some(net::buffer(response_buffer_), [self = shared_from_this()] + (beast::error_code ec, std::size_t n, bool done) + { + if (ec) return self->finish(ec); + if (n) + { + self->downstream_->async_write_some(net::buffer(self->response_buffer_.data(), n), + [self, done](beast::error_code write_ec) + { + if (write_ec) return self->finish(write_ec); + if (!done) return self->pump_response(); + self->downstream_->async_finish([self](beast::error_code finish_ec) { self->finish(finish_ec); }); + }); + } + else if (done) self->downstream_->async_finish([self](beast::error_code finish_ec) { self->finish(finish_ec); }); + else self->pump_response(); + }); + } + + void HttpProxySession::cancel() + { + if (inbound_) inbound_->cancel(); + if (upstream_) upstream_->cancel(); + if (downstream_) downstream_->cancel(); + } + + void HttpProxySession::finish(beast::error_code ec) + { + if (completed_) return; + completed_ = true; + if (ec) { spdlog::error("HttpProxySession failed: {}", ec.message()); cancel(); } + inbound_.reset(); + downstream_.reset(); + if (complete_) complete_(ec); + } +} diff --git a/framework/client/http_proxy_session.hpp b/framework/client/http_proxy_session.hpp new file mode 100644 index 0000000..c2eddd5 --- /dev/null +++ b/framework/client/http_proxy_session.hpp @@ -0,0 +1,51 @@ +#ifndef KHTTPD_FRAMEWORK_CLIENT_HTTP_PROXY_SESSION_HPP +#define KHTTPD_FRAMEWORK_CLIENT_HTTP_PROXY_SESSION_HPP + +#include "http_client_stream.hpp" +#include "context/http_request_stream.hpp" +#include "context/http_response_stream.hpp" +#include + +namespace khttpd::framework::client +{ + class HttpProxySession : public std::enable_shared_from_this + { + public: + using CompleteCallback = std::function; + HttpProxySession(std::shared_ptr inbound, + std::shared_ptr downstream, + std::size_t buffer_size = 64 * 1024); + HttpProxySession(std::shared_ptr inbound, + std::shared_ptr downstream, + boost::asio::ssl::context& ssl_context, + std::size_t buffer_size = 64 * 1024); + HttpProxySession(boost::asio::io_context& ioc, + std::shared_ptr inbound, + std::shared_ptr downstream, + std::size_t buffer_size = 64 * 1024); + HttpProxySession(boost::asio::io_context& ioc, + std::shared_ptr inbound, + std::shared_ptr downstream, + boost::asio::ssl::context& ssl_context, + std::size_t buffer_size = 64 * 1024); + void start(const std::string& upstream_url, HttpClientStream::RequestHead head, + CompleteCallback callback = {}); + void cancel(); + + private: + std::shared_ptr inbound_; + std::shared_ptr downstream_; + std::shared_ptr upstream_; + std::vector request_buffer_; + std::vector response_buffer_; + CompleteCallback complete_; + bool completed_ = false; + void pump_request(); + void finish_request(); + void read_response_head(); + void pump_response(); + void finish(boost::system::error_code ec); + }; +} + +#endif diff --git a/framework/client/websocket_client.cpp b/framework/client/websocket_client.cpp index 702f54f..f47293b 100644 --- a/framework/client/websocket_client.cpp +++ b/framework/client/websocket_client.cpp @@ -14,8 +14,10 @@ namespace khttpd::framework::client bool alive = true; bool close_notified = false; MessageHandler on_message; + FrameHandler on_frame; ErrorHandler on_error; CloseHandler on_close; + std::string negotiated_subprotocol; }; // ========================================== @@ -26,7 +28,7 @@ namespace khttpd::framework::client std::weak_ptr state_; std::string host_; beast::flat_buffer buffer_; - std::deque write_queue_; // 写队列 + std::deque write_queue_; bool is_writing_ = false; bool closing_ = false; bool close_started_ = false; @@ -42,28 +44,40 @@ namespace khttpd::framework::client const std::map& headers, WebsocketClient::ConnectCallback cb) = 0; virtual void close() = 0; + void set_negotiated_subprotocol(const beast::string_view value) + { + if (auto state = state_.lock()) + { + std::lock_guard lock{state->mutex}; + if (state->alive) state->negotiated_subprotocol.assign(value.data(), value.size()); + } + } + // 核心发送逻辑:入队 - void queue_write(std::string message) + void queue_write(WebsocketFrame frame) { net::post(get_executor(), beast::bind_front_handler( - &WebsocketSessionImpl::on_queue_write, shared_from_this(), std::move(message))); + &WebsocketSessionImpl::on_queue_write, shared_from_this(), std::move(frame))); } protected: virtual net::any_io_executor get_executor() = 0; virtual void do_write_from_queue() = 0; - void notify_message(const std::string& message) + void notify_frame(const WebsocketFrame& frame) { auto state = state_.lock(); if (!state) return; - WebsocketClient::MessageHandler handler; + WebsocketClient::FrameHandler frame_handler; + WebsocketClient::MessageHandler message_handler; { std::lock_guard lock{state->mutex}; if (!state->alive) return; - handler = state->on_message; + frame_handler = state->on_frame; + message_handler = state->on_message; } - if (handler) handler(message); + if (frame_handler) frame_handler(frame); + if (frame.type == WebsocketFrameType::text && message_handler) message_handler(frame.payload); } void notify_error(beast::error_code ec) @@ -107,10 +121,10 @@ namespace khttpd::framework::client if (callback) callback(ec); } - void on_queue_write(std::string message) + void on_queue_write(WebsocketFrame frame) { if (closing_) return; - write_queue_.push_back(std::move(message)); + write_queue_.push_back(std::move(frame)); if (!is_writing_) { is_writing_ = true; @@ -119,7 +133,7 @@ namespace khttpd::framework::client } // 通用的读循环处理 - void process_read_result(beast::error_code ec, std::size_t bytes) + void process_read_result(beast::error_code ec, std::size_t bytes, bool text, const websocket::close_reason& reason) { boost::ignore_unused(bytes); if (ec) @@ -133,6 +147,8 @@ namespace khttpd::framework::client ec == boost::asio::error::bad_descriptor || ec == boost::asio::error::operation_aborted) { + if (ec == websocket::error::closed) + notify_frame({WebsocketFrameType::close, {}, static_cast(reason.code), reason.reason.c_str()}); notify_close(); } else @@ -142,7 +158,8 @@ namespace khttpd::framework::client return; } - notify_message(beast::buffers_to_string(buffer_.data())); + notify_frame({text ? WebsocketFrameType::text : WebsocketFrameType::binary, + beast::buffers_to_string(buffer_.data())}); buffer_.consume(buffer_.size()); } @@ -182,6 +199,7 @@ namespace khttpd::framework::client class PlainWebsocketSession : public WebsocketSessionImpl { websocket::stream ws_; + websocket::response_type handshake_response_; tcp::resolver resolver_; WebsocketClient::ConnectCallback connect_cb_; @@ -233,9 +251,14 @@ namespace khttpd::framework::client protected: void do_write_from_queue() override { - ws_.async_write(net::buffer(write_queue_.front()), - beast::bind_front_handler(&PlainWebsocketSession::on_write, - std::static_pointer_cast(shared_from_this()))); + const auto& frame = write_queue_.front(); + if (frame.type == WebsocketFrameType::ping) + return ws_.async_ping(websocket::ping_data(frame.payload), [self = std::static_pointer_cast(shared_from_this())](beast::error_code ec) { self->process_write_result(ec); }); + if (frame.type == WebsocketFrameType::pong) + return ws_.async_pong(websocket::ping_data(frame.payload), [self = std::static_pointer_cast(shared_from_this())](beast::error_code ec) { self->process_write_result(ec); }); + ws_.text(frame.type == WebsocketFrameType::text); + ws_.async_write(net::buffer(frame.payload), beast::bind_front_handler(&PlainWebsocketSession::on_write, + std::static_pointer_cast(shared_from_this()))); } private: @@ -265,7 +288,7 @@ namespace khttpd::framework::client for (const auto& h : headers) req.set(h.first, h.second); })); - ws_.async_handshake(host_, target, + ws_.async_handshake(handshake_response_, host_, target, beast::bind_front_handler(&PlainWebsocketSession::on_handshake, std::static_pointer_cast( shared_from_this()))); @@ -275,6 +298,7 @@ namespace khttpd::framework::client { if (closing_) return; if (ec) return fail(ec); + set_negotiated_subprotocol(handshake_response_[beast::http::field::sec_websocket_protocol]); notify_connect(connect_cb_, ec); do_read(); } @@ -288,7 +312,7 @@ namespace khttpd::framework::client void on_read(beast::error_code ec, std::size_t bytes) { - process_read_result(ec, bytes); + process_read_result(ec, bytes, ws_.got_text(), ws_.reason()); if (!ec) do_read(); } @@ -309,6 +333,7 @@ namespace khttpd::framework::client class SslWebsocketSession : public WebsocketSessionImpl { websocket::stream> ws_; + websocket::response_type handshake_response_; tcp::resolver resolver_; WebsocketClient::ConnectCallback connect_cb_; @@ -365,9 +390,14 @@ namespace khttpd::framework::client protected: void do_write_from_queue() override { - ws_.async_write(net::buffer(write_queue_.front()), - beast::bind_front_handler(&SslWebsocketSession::on_write, - std::static_pointer_cast(shared_from_this()))); + const auto& frame = write_queue_.front(); + if (frame.type == WebsocketFrameType::ping) + return ws_.async_ping(websocket::ping_data(frame.payload), [self = std::static_pointer_cast(shared_from_this())](beast::error_code ec) { self->process_write_result(ec); }); + if (frame.type == WebsocketFrameType::pong) + return ws_.async_pong(websocket::ping_data(frame.payload), [self = std::static_pointer_cast(shared_from_this())](beast::error_code ec) { self->process_write_result(ec); }); + ws_.text(frame.type == WebsocketFrameType::text); + ws_.async_write(net::buffer(frame.payload), beast::bind_front_handler(&SslWebsocketSession::on_write, + std::static_pointer_cast(shared_from_this()))); } private: @@ -405,7 +435,7 @@ namespace khttpd::framework::client for (const auto& h : headers) req.set(h.first, h.second); })); - ws_.async_handshake(host_, target, + ws_.async_handshake(handshake_response_, host_, target, beast::bind_front_handler(&SslWebsocketSession::on_handshake, std::static_pointer_cast(shared_from_this()))); } @@ -414,6 +444,7 @@ namespace khttpd::framework::client { if (closing_) return; if (ec) return fail(ec); + set_negotiated_subprotocol(handshake_response_[beast::http::field::sec_websocket_protocol]); notify_connect(connect_cb_, ec); do_read(); } @@ -427,7 +458,7 @@ namespace khttpd::framework::client void on_read(beast::error_code ec, std::size_t bytes) { - process_read_result(ec, bytes); + process_read_result(ec, bytes, ws_.got_text(), ws_.reason()); if (!ec) do_read(); } @@ -489,6 +520,25 @@ namespace khttpd::framework::client headers_[key] = value; } + void WebsocketClient::set_subprotocols(const std::vector& subprotocols) + { + std::string value; + for (const auto& protocol : subprotocols) + { + if (protocol.empty()) continue; + if (!value.empty()) value += ", "; + value += protocol; + } + if (value.empty()) headers_.erase("Sec-WebSocket-Protocol"); + else headers_["Sec-WebSocket-Protocol"] = std::move(value); + } + + std::string WebsocketClient::negotiated_subprotocol() const + { + std::lock_guard lock{state_->mutex}; + return state_->negotiated_subprotocol; + } + void WebsocketClient::connect(const std::string& url, ConnectCallback callback) { auto url_result = boost::urls::parse_uri(url); @@ -501,8 +551,15 @@ namespace khttpd::framework::client std::string host = u.host(); std::string scheme = u.scheme(); std::string port = u.port(); - std::string target = u.encoded_path().data(); + std::string target(u.encoded_target()); if (target.empty()) target = "/"; + else if (target.front() == '?') target.insert(target.begin(), '/'); + + if (scheme != "ws" && scheme != "wss") + { + if (callback) callback(make_error_code(boost::system::errc::operation_not_supported)); + return; + } if (port.empty()) port = (scheme == "wss") ? "443" : "80"; @@ -516,6 +573,7 @@ namespace khttpd::framework::client { std::lock_guard lock{state_->mutex}; state_->close_notified = false; + state_->negotiated_subprotocol.clear(); } auto s = std::make_shared(ioc_, *ssl_ctx_ptr_, state_); session_ = s; @@ -526,6 +584,7 @@ namespace khttpd::framework::client { std::lock_guard lock{state_->mutex}; state_->close_notified = false; + state_->negotiated_subprotocol.clear(); } auto s = std::make_shared(ioc_, state_); session_ = s; @@ -534,10 +593,15 @@ namespace khttpd::framework::client } void WebsocketClient::send(const std::string& message) + { + send({WebsocketFrameType::text, message}); + } + + void WebsocketClient::send(WebsocketFrame frame) { if (session_) { - session_->queue_write(message); + session_->queue_write(std::move(frame)); } } @@ -556,6 +620,12 @@ namespace khttpd::framework::client state_->on_message = std::move(handler); } + void WebsocketClient::set_on_frame(FrameHandler handler) + { + std::lock_guard lock{state_->mutex}; + state_->on_frame = std::move(handler); + } + void WebsocketClient::set_on_error(ErrorHandler handler) { std::lock_guard lock{state_->mutex}; diff --git a/framework/client/websocket_client.hpp b/framework/client/websocket_client.hpp index 3d00489..77c807c 100644 --- a/framework/client/websocket_client.hpp +++ b/framework/client/websocket_client.hpp @@ -14,6 +14,7 @@ #include #include #include +#include "context/websocket_context.hpp" namespace khttpd::framework::client { @@ -33,6 +34,7 @@ namespace khttpd::framework::client using ConnectCallback = std::function; using MessageHandler = std::function; + using FrameHandler = std::function; using ErrorHandler = std::function; using CloseHandler = std::function; @@ -47,13 +49,17 @@ namespace khttpd::framework::client // 发送消息 (线程安全,支持并发调用) void send(const std::string& message); + void send(WebsocketFrame frame); // 关闭连接 void close(); // 配置 void set_header(const std::string& key, const std::string& value); + void set_subprotocols(const std::vector& subprotocols); + std::string negotiated_subprotocol() const; void set_on_message(MessageHandler handler); + void set_on_frame(FrameHandler handler); void set_on_error(ErrorHandler handler); void set_on_close(CloseHandler handler); diff --git a/framework/context/http_context.cpp b/framework/context/http_context.cpp index 733f744..e812a7c 100644 --- a/framework/context/http_context.cpp +++ b/framework/context/http_context.cpp @@ -55,8 +55,9 @@ namespace khttpd::framework } - HttpContext::HttpContext(Request& req, Response& res) - : req_(req), res_(res) + HttpContext::HttpContext(Request& req, Response& res, + std::optional peer_endpoint) + : req_(req), res_(res), peer_endpoint_(std::move(peer_endpoint)) { res_.version(req_.version()); res_.keep_alive(req_.keep_alive()); diff --git a/framework/context/http_context.hpp b/framework/context/http_context.hpp index a15515f..eda0d1a 100644 --- a/framework/context/http_context.hpp +++ b/framework/context/http_context.hpp @@ -6,6 +6,7 @@ #include #include #include +#include #include #include #include @@ -40,7 +41,8 @@ namespace khttpd::framework using WriteHandler = std::function; using HttpStreamHandler = std::function; - HttpContext(Request& req, Response& res); + HttpContext(Request& req, Response& res, + std::optional peer_endpoint = std::nullopt); ~HttpContext(); const std::string& path() const; @@ -53,6 +55,12 @@ namespace khttpd::framework std::optional> get_headers(boost::beast::string_view name) const; std::optional> get_headers(boost::beast::http::field name) const; + // The transport peer, captured from the accepted socket. Unlike forwarded + // headers this value cannot be supplied by the HTTP client. + const std::optional& peer_endpoint() const { return peer_endpoint_; } + std::optional peer_address() const + { return peer_endpoint_ ? std::optional{peer_endpoint_->address()} : std::nullopt; } + // Cookie support std::optional get_cookie(const std::string& key) const; std::vector get_cookies(const std::string& key) const; @@ -131,6 +139,7 @@ namespace khttpd::framework private: Request& req_; Response& res_; + std::optional peer_endpoint_; mutable std::map query_params_; mutable std::string cached_path_; mutable boost::urls::url_view parsed_url_; diff --git a/framework/context/http_request_stream.hpp b/framework/context/http_request_stream.hpp new file mode 100644 index 0000000..330884a --- /dev/null +++ b/framework/context/http_request_stream.hpp @@ -0,0 +1,26 @@ +#ifndef KHTTPD_FRAMEWORK_CONTEXT_HTTP_REQUEST_STREAM_HPP +#define KHTTPD_FRAMEWORK_CONTEXT_HTTP_REQUEST_STREAM_HPP + +#include +#include +#include + +namespace khttpd::framework +{ + // A request body is consumed serially: callers must wait for a read callback + // before issuing the next read. This is the backpressure boundary used by + // streaming proxy handlers. + class HttpRequestStream + { + public: + using ReadCallback = std::function; + virtual ~HttpRequestStream() = default; + + virtual void async_read_some(boost::asio::mutable_buffer buffer, ReadCallback callback) = 0; + // Cancels only request-body consumption. The response side remains usable. + virtual void cancel_read() { cancel(); } + virtual void cancel() = 0; + }; +} + +#endif diff --git a/framework/context/http_response_stream.hpp b/framework/context/http_response_stream.hpp new file mode 100644 index 0000000..7406709 --- /dev/null +++ b/framework/context/http_response_stream.hpp @@ -0,0 +1,25 @@ +#ifndef KHTTPD_FRAMEWORK_CONTEXT_HTTP_RESPONSE_STREAM_HPP +#define KHTTPD_FRAMEWORK_CONTEXT_HTTP_RESPONSE_STREAM_HPP + +#include +#include +#include + +namespace khttpd::framework +{ + class HttpResponseStream + { + public: + using ResponseHead = boost::beast::http::response; + using Callback = std::function; + virtual ~HttpResponseStream() = default; + virtual void async_start(ResponseHead head, Callback callback) = 0; + virtual void async_write_some(boost::asio::const_buffer buffer, Callback callback) = 0; + virtual void async_finish(Callback callback) = 0; + // Stops an inbound request-body read while preserving this response stream. + virtual void cancel_request_body() { cancel(); } + virtual void cancel() = 0; + }; +} + +#endif diff --git a/framework/context/websocket_context.cpp b/framework/context/websocket_context.cpp index 2a0384d..a5430f7 100644 --- a/framework/context/websocket_context.cpp +++ b/framework/context/websocket_context.cpp @@ -4,6 +4,8 @@ #include #include +#include +#include namespace khttpd::framework { @@ -15,6 +17,8 @@ namespace khttpd::framework { id = session_shared_ptr->id; } + frame.type = text ? WebsocketFrameType::text : WebsocketFrameType::binary; + frame.payload = message; } WebsocketContext::WebsocketContext(std::weak_ptr session, std::string path_str, @@ -39,4 +43,54 @@ namespace khttpd::framework spdlog::error("Attempted to send WS message to expired session (path: {}).", path); } } + + void WebsocketContext::send(WebsocketFrame outbound_frame) + { + if (auto session = session_weak_ptr.lock()) session->send_frame(std::move(outbound_frame)); + else spdlog::error("Attempted to send WS frame to expired session (path: {}).", path); + } + + const WebsocketHandshakeRequest& WebsocketContext::handshake() const + { + static const WebsocketHandshakeRequest empty; + if (const auto session = session_weak_ptr.lock()) return session->handshake(); + return empty; + } + + std::optional WebsocketContext::get_header(const std::string& name) const + { + const auto values = get_headers(name); + if (values.empty()) return std::nullopt; + return values.front(); + } + + std::vector WebsocketContext::get_headers(const std::string& name) const + { + std::vector values; + const auto equal_ci = [](const std::string& a, const std::string& b) + { + return a.size() == b.size() && std::equal(a.begin(), a.end(), b.begin(), + [](unsigned char x, unsigned char y) { return std::tolower(x) == std::tolower(y); }); + }; + for (const auto& [key, value] : handshake().headers) + if (equal_ci(key, name)) values.push_back(value); + return values; + } + + std::optional WebsocketContext::get_query_param(const std::string& key) const + { + const auto it = handshake().query_params.find(key); + return it == handshake().query_params.end() ? std::nullopt : std::optional(it->second); + } + + std::optional WebsocketContext::get_path_param(const std::string& key) const + { + const auto it = path_params_.find(key); + return it == path_params_.end() ? std::nullopt : std::optional(it->second); + } + + void WebsocketContext::set_path_params(std::map params) + { + path_params_ = std::move(params); + } } diff --git a/framework/context/websocket_context.hpp b/framework/context/websocket_context.hpp index e560b0b..a755f92 100644 --- a/framework/context/websocket_context.hpp +++ b/framework/context/websocket_context.hpp @@ -7,10 +7,35 @@ #include #include #include +#include namespace khttpd::framework { class WebsocketSession; + + enum class WebsocketFrameType { text, binary, ping, pong, close }; + + struct WebsocketFrame + { + WebsocketFrameType type = WebsocketFrameType::text; + std::string payload; + uint16_t close_code = 1000; + std::string close_reason; + }; + + // HeaderList deliberately keeps duplicate fields (notably Cookie and + // Sec-WebSocket-Extensions) in their original handshake order. + using WebsocketHeaderList = std::vector>; + + struct WebsocketHandshakeRequest + { + std::string target; + std::string path; + WebsocketHeaderList headers; + std::map query_params; + std::vector subprotocols; + }; + class WebsocketContext { public: @@ -20,6 +45,7 @@ namespace khttpd::framework bool is_text; boost::beast::error_code error_code; std::string path; + WebsocketFrame frame; std::map extended_data; @@ -29,6 +55,14 @@ namespace khttpd::framework boost::beast::error_code ec = {}); void send(const std::string& msg, bool is_text = true); + void send(WebsocketFrame frame); + + const WebsocketHandshakeRequest& handshake() const; + std::optional get_header(const std::string& name) const; + std::vector get_headers(const std::string& name) const; + std::optional get_query_param(const std::string& key) const; + std::optional get_path_param(const std::string& key) const; + void set_path_params(std::map params); void set_attribute(const std::string& key, std::any value) { extended_data[key] = std::move(value); @@ -54,6 +88,9 @@ namespace khttpd::framework } return std::nullopt; } + + private: + std::map path_params_; }; } #endif // KHTTPD_FRAMEWORK_WEBSOCKET_CONTEXT_HPP diff --git a/framework/interceptor/interceptor.hpp b/framework/interceptor/interceptor.hpp index 3e4cbf8..7d6e124 100644 --- a/framework/interceptor/interceptor.hpp +++ b/framework/interceptor/interceptor.hpp @@ -2,6 +2,7 @@ #define KHTTPD_FRAMEWORK_INTERCEPTOR_INTERCEPTOR_HPP_ #include "context/http_context.hpp" +#include namespace khttpd::framework { @@ -14,6 +15,7 @@ namespace khttpd::framework class Interceptor { public: + using RequestCompletion = std::function; virtual ~Interceptor() = default; /** @@ -26,6 +28,14 @@ namespace khttpd::framework return InterceptorResult::Continue; } + // Override this for remote authentication or other asynchronous policy + // checks. Invoke complete exactly once; it may be invoked from any thread. + // The default preserves existing synchronous interceptors. + virtual void async_handle_request(HttpContext& ctx, RequestCompletion complete) + { + complete(handle_request(ctx)); + } + /** * @brief Post-response handler. * @param ctx The HTTP context. diff --git a/framework/router/http_router.cpp b/framework/router/http_router.cpp index 14e0723..2cf5b91 100644 --- a/framework/router/http_router.cpp +++ b/framework/router/http_router.cpp @@ -151,11 +151,142 @@ namespace khttpd::framework add_route(path, boost::beast::http::verb::options, std::move(handler)); } + void HttpRouter::stream(const std::string& path_pattern, const boost::beast::http::verb method, + HttpStreamHandler handler) + { + for (auto& entry : routes_) + { + if (entry.original_path == path_pattern) + { + entry.stream_handlers[method] = std::move(handler); + return; + } + } + RouteEntry entry; + entry.original_path = path_pattern; + auto [regex, params, literal_count, dynamic_count] = parse_path_pattern(path_pattern); + entry.path_regex = std::move(regex); + entry.param_names = std::move(params); + entry.literal_segments_count = literal_count; + entry.dynamic_segments_count = dynamic_count; + entry.stream_handlers[method] = std::move(handler); + routes_.push_back(std::move(entry)); + std::sort(routes_.begin(), routes_.end(), RouteEntry::compare_specificity); + } + + bool HttpRouter::is_stream_route(const std::string& path, const boost::beast::http::verb method) const + { + for (const auto& entry : routes_) + { + if (!std::regex_match(path, entry.path_regex)) continue; + return entry.stream_handlers.find(method) != entry.stream_handlers.end(); + } + return false; + } + + bool HttpRouter::dispatch_stream(HttpContext& ctx, std::shared_ptr stream, + std::shared_ptr response_stream, + HttpStreamComplete complete) const + { + for (const auto& entry : routes_) + { + std::smatch matches; + if (!std::regex_match(ctx.path(), matches, entry.path_regex)) continue; + const auto handler = entry.stream_handlers.find(ctx.method()); + if (handler == entry.stream_handlers.end()) return false; + std::map params; + for (size_t i = 0; i < entry.param_names.size() && i + 1 < matches.size(); ++i) + params[entry.param_names[i]] = matches[i + 1].str(); + ctx.set_path_params(std::move(params)); + handler->second(ctx, std::move(stream), std::move(response_stream), std::move(complete)); + return true; + } + return false; + } + void HttpRouter::add_interceptor(std::shared_ptr interceptor) { interceptors_.push_back(std::move(interceptor)); } + void HttpRouter::async_run_pre_interceptors(HttpContext& ctx, InterceptorCompletion complete) const + { + struct State : std::enable_shared_from_this + { + const HttpRouter* router; + const std::vector>* interceptors; + HttpContext* ctx; + InterceptorCompletion complete; + std::size_t index = 0; + + void advance(InterceptorResult result) + { + if (result == InterceptorResult::Stop || index == interceptors->size()) + return complete(result); + auto interceptor = (*interceptors)[index++]; + try + { + interceptor->async_handle_request(*ctx, [self = shared_from_this()](InterceptorResult next_result) + { self->advance(next_result); }); + } + catch (...) + { + router->handle_exception(std::current_exception(), *ctx); + complete(InterceptorResult::Stop); + } + } + }; + auto state = std::make_shared(); + state->router = this; + state->interceptors = &interceptors_; + state->ctx = &ctx; + state->complete = std::move(complete); + state->advance(InterceptorResult::Continue); + } + + void HttpRouter::async_route(const std::string& path, boost::beast::http::verb method, + HttpAsyncHandler handler) + { + auto [path_regex, param_names, literal_count, dynamic_count] = parse_path_pattern(path); + for (auto& entry : routes_) + { + if (entry.original_path == path) + { + entry.async_handlers[method] = std::move(handler); + return; + } + } + RouteEntry entry; + entry.original_path = path; + entry.path_regex = std::move(path_regex); + entry.param_names = std::move(param_names); + entry.literal_segments_count = literal_count; + entry.dynamic_segments_count = dynamic_count; + entry.async_handlers[method] = std::move(handler); + routes_.push_back(std::move(entry)); + std::sort(routes_.begin(), routes_.end(), RouteEntry::compare_specificity); + } + + bool HttpRouter::dispatch_async(HttpContext& ctx, HttpAsyncComplete complete) const + { + for (const auto& entry : routes_) + { + std::smatch matches; + if (!std::regex_match(ctx.path(), matches, entry.path_regex)) continue; + auto handler = entry.async_handlers.find(ctx.method()); + if (handler == entry.async_handlers.end() && ctx.method() == boost::beast::http::verb::head) + handler = entry.async_handlers.find(boost::beast::http::verb::get); + if (handler == entry.async_handlers.end()) return false; + std::map params; + for (size_t i = 0; i < entry.param_names.size() && i + 1 < matches.size(); ++i) + params[entry.param_names[i]] = matches[i + 1].str(); + ctx.set_path_params(std::move(params)); + handler->second(ctx, std::move(complete)); + return true; + } + return false; + } + InterceptorResult HttpRouter::run_pre_interceptors(HttpContext& ctx) const { for (const auto& interceptor : interceptors_) diff --git a/framework/router/http_router.hpp b/framework/router/http_router.hpp index 2630b67..183b9ad 100644 --- a/framework/router/http_router.hpp +++ b/framework/router/http_router.hpp @@ -3,6 +3,8 @@ #define KHTTPD_FRAMEWORK_ROUTER_HTTP_ROUT #include "context/http_context.hpp" +#include "context/http_request_stream.hpp" +#include "context/http_response_stream.hpp" #include "interceptor/interceptor.hpp" #include "exception/exception_handler.hpp" #include @@ -15,6 +17,11 @@ namespace khttpd::framework { using HttpHandler = std::function; + using HttpAsyncComplete = std::function; + using HttpAsyncHandler = std::function; + using HttpStreamComplete = std::function; + using HttpStreamHandler = std::function, + std::shared_ptr, HttpStreamComplete)>; using UnknownExceptionHandler = std::function; // 路由条目结构 @@ -24,6 +31,8 @@ namespace khttpd::framework std::regex path_regex; std::vector param_names; std::map handlers; + std::map async_handlers; + std::map stream_handlers; int literal_segments_count = 0; int dynamic_segments_count = 0; @@ -51,9 +60,20 @@ namespace khttpd::framework void put(const std::string& path, HttpHandler handler); void del(const std::string& path, HttpHandler handler); void options(const std::string& path, HttpHandler handler); + // Async handlers must invoke complete exactly once, from any thread. + void async_route(const std::string& path, boost::beast::http::verb method, HttpAsyncHandler handler); + void stream(const std::string& path, boost::beast::http::verb method, HttpStreamHandler handler); + + // Used by HttpSession after it has parsed only the request header. + bool is_stream_route(const std::string& path, boost::beast::http::verb method) const; + bool dispatch_stream(HttpContext& ctx, std::shared_ptr stream, + std::shared_ptr response_stream, + HttpStreamComplete complete) const; void add_interceptor(std::shared_ptr interceptor); InterceptorResult run_pre_interceptors(HttpContext& ctx) const; + using InterceptorCompletion = std::function; + void async_run_pre_interceptors(HttpContext& ctx, InterceptorCompletion complete) const; void run_post_interceptors(HttpContext& ctx) const; // Exception handling @@ -63,6 +83,7 @@ namespace khttpd::framework void handle_unknown_exception(HttpContext& ctx) const; bool dispatch(HttpContext& ctx, const std::function& static_file_fun = nullptr) const; + bool dispatch_async(HttpContext& ctx, HttpAsyncComplete complete) const; private: std::vector routes_; diff --git a/framework/router/websocket_router.cpp b/framework/router/websocket_router.cpp index d4605b8..eacce20 100644 --- a/framework/router/websocket_router.cpp +++ b/framework/router/websocket_router.cpp @@ -1,5 +1,6 @@ -// framework/router/websocket_router.cpp #include "websocket_router.hpp" + +#include #include #include "websocket/websocket_session.hpp" @@ -17,64 +18,73 @@ namespace khttpd::framework WebsocketRouter::WebsocketRouter() = default; - void WebsocketRouter::add_handler(const std::string& path, - WebsocketOpenHandler on_open, - WebsocketMessageHandler on_message, - WebsocketCloseHandler on_close, + void WebsocketRouter::add_handler(const std::string& path, WebsocketOpenHandler on_open, + WebsocketMessageHandler on_message, WebsocketCloseHandler on_close, WebsocketErrorHandler on_error) { - WebsocketRouteEntry entry; - entry.on_open = std::move(on_open); - entry.on_message = std::move(on_message); - entry.on_close = std::move(on_close); - entry.on_error = std::move(on_error); - handlers_[path] = entry; + WebsocketRouteEntry entry{std::move(on_open), std::move(on_message), std::move(on_close), std::move(on_error)}; + std::unique_lock lock(handlers_mutex_); + for (auto& route : handlers_) + if (route.original_path == path) { route.handlers = std::move(entry); return; } + auto [regex, params, literal_count, dynamic_count] = parse_path_pattern(path); + handlers_.push_back({path, std::move(regex), std::move(params), literal_count, dynamic_count, std::move(entry)}); + std::sort(handlers_.begin(), handlers_.end(), WebsocketRoute::compare_specificity); spdlog::debug("Registered WebSocket handlers for path: {}", path); } void WebsocketRouter::dispatch_open(const std::string& path, WebsocketContext& ctx) - { - const auto it = handlers_.find(path); - if (it == handlers_.end() || !it->second.on_open) - { - spdlog::warn("No on_open handler found for WS path: {}", path); - return; - } - it->second.on_open(ctx); - } - + { dispatch(path, ctx, [&](const WebsocketRouteEntry& h) { if (h.on_open) h.on_open(ctx); }); } void WebsocketRouter::dispatch_message(const std::string& path, WebsocketContext& ctx) - { - const auto it = handlers_.find(path); - if (it == handlers_.end() || !it->second.on_message) - { - spdlog::warn("No on_message handler found for WS path: {}", path); - // Default behavior: if no handler, just close connection? Echo? - // For now, nothing happens. - return; - } - it->second.on_message(ctx); - } - + { dispatch(path, ctx, [&](const WebsocketRouteEntry& h) { if (h.on_message) h.on_message(ctx); }); } void WebsocketRouter::dispatch_close(const std::string& path, WebsocketContext& ctx) + { dispatch(path, ctx, [&](const WebsocketRouteEntry& h) { if (h.on_close) h.on_close(ctx); }); } + void WebsocketRouter::dispatch_error(const std::string& path, WebsocketContext& ctx) + { dispatch(path, ctx, [&](const WebsocketRouteEntry& h) { if (h.on_error) h.on_error(ctx); }); } + + std::tuple, int, int> + WebsocketRouter::parse_path_pattern(const std::string& pattern) { - const auto it = handlers_.find(path); - if (it == handlers_.end() || !it->second.on_close) + std::string expression = "^"; + std::vector names; + std::regex parameter(":([a-zA-Z_][a-zA-Z0-9_]*)"); + std::regex escaped(R"([\\\.\+\*\?\|\(\)\[\]\{\}\^\$])"); + std::sregex_iterator end, it(pattern.begin(), pattern.end(), parameter); + const int total = std::distance(it, end); + auto cursor = pattern.begin(); + int literals = 0, dynamics = 0; + for (int index = 0; it != end; ++it, ++index) { - spdlog::warn("No on_close handler found for WS path: {}", path); - return; + const std::string literal(cursor, it->prefix().second); + literals += static_cast(std::count(literal.begin(), literal.end(), '/')); + expression += std::regex_replace(literal, escaped, "\\$&"); + names.push_back((*it)[1]); ++dynamics; + expression += index + 1 == total ? "(.*)" : "([^/]+)"; + cursor = it->suffix().first; } - it->second.on_close(ctx); + const std::string tail(cursor, pattern.end()); + literals += static_cast(std::count(tail.begin(), tail.end(), '/')); + expression += std::regex_replace(tail, escaped, "\\$&"); + return {std::regex(expression + "$"), names, literals, dynamics}; } - void WebsocketRouter::dispatch_error(const std::string& path, WebsocketContext& ctx) + void WebsocketRouter::dispatch(const std::string& path, WebsocketContext& ctx, + const std::function& invoke) { - const auto it = handlers_.find(path); - if (it == handlers_.end() || !it->second.on_error) { - spdlog::warn("No on_error handler found for WS path: {} (error: {})", path, ctx.error_code.message()); - return; + std::shared_lock lock(handlers_mutex_); + for (const auto& route : handlers_) + { + std::smatch match; + if (!std::regex_match(path, match, route.path_regex)) continue; + std::map params; + for (size_t i = 0; i < route.param_names.size(); ++i) params[route.param_names[i]] = match[i + 1].str(); + ctx.set_path_params(std::move(params)); + const auto handler = route.handlers; + lock.unlock(); + invoke(handler); + return; + } } - it->second.on_error(ctx); + spdlog::warn("No WebSocket handler found for path: {}", path); } } diff --git a/framework/router/websocket_router.hpp b/framework/router/websocket_router.hpp index 133105b..cf4e440 100644 --- a/framework/router/websocket_router.hpp +++ b/framework/router/websocket_router.hpp @@ -5,6 +5,10 @@ #include #include #include +#include +#include +#include +#include #include "context/websocket_context.hpp" namespace khttpd::framework @@ -29,6 +33,22 @@ namespace khttpd::framework WebsocketErrorHandler on_error; }; + struct WebsocketRoute + { + std::string original_path; + std::regex path_regex; + std::vector param_names; + int literal_segments_count = 0; + int dynamic_segments_count = 0; + WebsocketRouteEntry handlers; + static bool compare_specificity(const WebsocketRoute& a, const WebsocketRoute& b) + { + return a.literal_segments_count == b.literal_segments_count + ? a.dynamic_segments_count < b.dynamic_segments_count + : a.literal_segments_count > b.literal_segments_count; + } + }; + class WebsocketRouter { public: @@ -46,7 +66,11 @@ namespace khttpd::framework void dispatch_error(const std::string& path, WebsocketContext& ctx); private: - std::map handlers_; // 路径 -> 处理器集合 + static std::tuple, int, int> parse_path_pattern(const std::string& pattern); + void dispatch(const std::string& path, WebsocketContext& ctx, + const std::function& invoke); + std::vector handlers_; + mutable std::shared_mutex handlers_mutex_; }; } #endif // KHTTPD_FRAMEWORK_ROUTER_WEBSOCKET_ROUTER_HPP diff --git a/framework/server.cpp b/framework/server.cpp index a895b2a..62db1b3 100644 --- a/framework/server.cpp +++ b/framework/server.cpp @@ -95,6 +95,16 @@ namespace khttpd::framework return acceptor_.local_endpoint(); } + void Server::set_max_buffered_request_body_size(std::uint64_t bytes) + { + max_buffered_request_body_size_.store(bytes, std::memory_order_relaxed); + } + + std::uint64_t Server::get_max_buffered_request_body_size() const + { + return max_buffered_request_body_size_.load(std::memory_order_relaxed); + } + void Server::run() { spdlog::info("Server listening on {}:{}", acceptor_.local_endpoint().address().to_string(), @@ -140,7 +150,8 @@ namespace khttpd::framework } else { - std::make_shared(std::move(socket), http_router_, websocket_router_, web_root_, canonical_web_root_)->run(); + std::make_shared(std::move(socket), http_router_, websocket_router_, web_root_, + canonical_web_root_, get_max_buffered_request_body_size())->run(); } if (acceptor_.is_open()) diff --git a/framework/server.hpp b/framework/server.hpp index e40825c..70b3f33 100644 --- a/framework/server.hpp +++ b/framework/server.hpp @@ -8,6 +8,8 @@ #include #include #include +#include +#include #include "router/http_router.hpp" #include "router/websocket_router.hpp" @@ -20,6 +22,8 @@ namespace khttpd::framework class Server : public std::enable_shared_from_this { public: + static constexpr std::uint64_t default_max_buffered_request_body_size = 16ULL * 1024 * 1024; + Server(const tcp::endpoint& endpoint, std::string web_root, int num_threads = 1); HttpRouter& get_http_router(); @@ -32,6 +36,11 @@ namespace khttpd::framework tcp::endpoint local_endpoint() const; + // Limit for routes that buffer the complete request body in memory. + // Stream routes are not subject to this limit. + void set_max_buffered_request_body_size(std::uint64_t bytes); + std::uint64_t get_max_buffered_request_body_size() const; + void run(); void stop(); @@ -45,6 +54,7 @@ namespace khttpd::framework HttpRouter http_router_; WebsocketRouter websocket_router_; + std::atomic max_buffered_request_body_size_{default_max_buffered_request_body_size}; void do_accept(); void on_accept(boost::beast::error_code ec, tcp::socket socket); diff --git a/framework/session/http_session.cpp b/framework/session/http_session.cpp index 4db8600..697bba8 100644 --- a/framework/session/http_session.cpp +++ b/framework/session/http_session.cpp @@ -6,6 +6,7 @@ #include #include #include +#include using namespace khttpd::framework; @@ -26,15 +27,52 @@ struct HttpSession::ChunkWriteState } }; +class HttpSession::RequestStreamImpl final : public HttpRequestStream +{ + std::shared_ptr session_; +public: + explicit RequestStreamImpl(const std::shared_ptr& session) : session_(session) {} + void async_read_some(net::mutable_buffer target, ReadCallback callback) override + { + if (session_) session_->async_read_stream_body(target, std::move(callback)); + else callback(net::error::operation_aborted, 0, true); + } + void cancel() override + { + if (session_) session_->cancel_session(); + } + void cancel_read() override { if (session_) session_->cancel_stream_body(); } +}; + +class HttpSession::ResponseStreamImpl final : public HttpResponseStream +{ + std::shared_ptr session_; +public: + explicit ResponseStreamImpl(const std::shared_ptr& session) : session_(session) {} + void async_start(ResponseHead head, Callback cb) override + { if (session_) session_->start_stream_response(std::move(head), std::move(cb)); else cb(net::error::operation_aborted); } + void async_write_some(net::const_buffer b, Callback cb) override + { if (session_) session_->write_stream_response(b, std::move(cb)); else cb(net::error::operation_aborted); } + void async_finish(Callback cb) override + { if (session_) session_->finish_stream_response(std::move(cb)); else cb(net::error::operation_aborted); } + void cancel_request_body() override { if (session_) session_->cancel_stream_body(); } + void cancel() override { if (session_) session_->cancel_session(); } +}; + HttpSession::HttpSession(tcp::socket&& socket, HttpRouter& router, WebsocketRouter& ws_router, const std::string& web_root, - const boost::filesystem::path& canonical_web_root) + const boost::filesystem::path& canonical_web_root, + std::uint64_t max_buffered_request_body_size) : stream_(std::move(socket)), router_(router), websocket_router_(ws_router), web_root_path_(web_root), - canonical_web_root_path_(canonical_web_root) + canonical_web_root_path_(canonical_web_root), + max_buffered_request_body_size_(max_buffered_request_body_size) { + beast::error_code peer_ec; + peer_endpoint_ = stream_.socket().remote_endpoint(peer_ec); + if (peer_ec) peer_endpoint_.reset(); if (canonical_web_root_path_.empty()) { disable_web_root_ = true; @@ -50,8 +88,257 @@ void HttpSession::run() void HttpSession::do_read() { req_ = {}; - http::async_read(stream_, buffer_, req_, - beast::bind_front_handler(&HttpSession::on_read, shared_from_this())); + res_ = {}; + ctx.reset(); + buffered_body_.clear(); + stream_completed_ = false; + request_body_cancelled_ = false; + request_parser_.emplace(); + request_parser_->header_limit(64 * 1024); + request_parser_->body_limit((std::numeric_limits::max)()); + http::async_read_header(stream_, buffer_, *request_parser_, + beast::bind_front_handler(&HttpSession::on_read_header, shared_from_this())); +} + +void HttpSession::copy_request_head() +{ + const auto& source = request_parser_->get(); + req_.method(source.method()); + req_.target(source.target()); + req_.version(source.version()); + req_.keep_alive(source.keep_alive()); + for (const auto& field : source) req_.insert(field.name_string(), field.value()); +} + +void HttpSession::on_read_header(const beast::error_code& ec, std::size_t bytes_transferred) +{ + boost::ignore_unused(bytes_transferred); + if (ec == http::error::end_of_stream) return do_close(); + if (ec) { spdlog::error("HttpSession header read error: {}", ec.message()); return do_close(); } + copy_request_head(); + if (beast::websocket::is_upgrade(req_)) + { + ctx = std::make_shared(req_, res_, peer_endpoint_); + return router_.async_run_pre_interceptors(*ctx, [self = shared_from_this()](InterceptorResult result) + { + net::post(self->stream_.get_executor(), [self, result] + { + if (result == InterceptorResult::Continue) self->handle_websocket_upgrade(); + else + { + self->router_.run_post_interceptors(*self->ctx); + self->send_context_response(); + } + }); + }); + } + + std::string path(req_.target()); + if (const auto query = path.find('?'); query != std::string::npos) path.resize(query); + const bool stream_route = router_.is_stream_route(path, req_.method()); + if (!stream_route) + { + const auto content_length = request_parser_->content_length(); + if (content_length && *content_length > max_buffered_request_body_size_) + return send_payload_too_large(); + } + const auto expect = req_[http::field::expect]; + if (boost::beast::iequals(expect, "100-continue")) + { + auto response = std::make_shared>(http::status::continue_, req_.version()); + http::async_write(stream_, *response, [self = shared_from_this(), response, stream_route] + (beast::error_code write_ec, std::size_t) + { + if (write_ec) return self->do_close(); + if (stream_route) self->handle_stream_request(); else self->read_buffered_body(); + }); + return; + } + if (stream_route) return handle_stream_request(); + read_buffered_body(); +} + +void HttpSession::read_buffered_body() +{ + if (request_parser_->is_done()) + { + req_.body() = std::move(buffered_body_); + return handle_request(); + } + auto& body = request_parser_->get().body(); + body.data = buffered_body_chunk_.data(); + body.size = buffered_body_chunk_.size(); + http::async_read_some(stream_, buffer_, *request_parser_, + beast::bind_front_handler(&HttpSession::on_read_buffered_body, shared_from_this())); +} + +void HttpSession::on_read_buffered_body(beast::error_code ec, std::size_t bytes_transferred) +{ + boost::ignore_unused(bytes_transferred); + if (ec == http::error::need_buffer) ec = {}; + if (ec) { spdlog::error("HttpSession body read error: {}", ec.message()); return do_close(); } + const auto produced = buffered_body_chunk_.size() - request_parser_->get().body().size; + if (buffered_body_.size() > max_buffered_request_body_size_ || + produced > max_buffered_request_body_size_ - buffered_body_.size()) + return send_payload_too_large(); + buffered_body_.append(buffered_body_chunk_.data(), produced); + read_buffered_body(); +} + +void HttpSession::send_payload_too_large() +{ + http::response response{http::status::payload_too_large, req_.version()}; + response.keep_alive(false); + response.body() = "request body exceeds buffered route limit"; + response.prepare_payload(); + send_response(std::move(response)); +} + +void HttpSession::handle_stream_request() +{ + ctx = std::make_shared(req_, res_, peer_endpoint_); + auto stream = std::make_shared(shared_from_this()); + auto response_stream = std::make_shared(shared_from_this()); + std::weak_ptr weak = shared_from_this(); + const auto complete = [weak] + { + if (auto self = weak.lock()) net::post(self->stream_.get_executor(), [self] + { + if (self->stream_completed_) return; + self->stream_completed_ = true; + self->router_.run_post_interceptors(*self->ctx); + self->send_context_response(); + }); + }; + try + { + router_.async_run_pre_interceptors(*ctx, + [self = shared_from_this(), stream = std::move(stream), response_stream = std::move(response_stream), complete] + (InterceptorResult result) mutable + { + net::post(self->stream_.get_executor(), + [self, result, stream = std::move(stream), response_stream = std::move(response_stream), complete]() mutable + { + if (result == InterceptorResult::Stop) return complete(); + try + { + if (!self->router_.dispatch_stream(*self->ctx, std::move(stream), std::move(response_stream), complete)) + { + self->ctx->set_status(http::status::not_found); + self->ctx->set_body("stream route not found"); + complete(); + } + } + catch (...) { self->router_.handle_exception(std::current_exception(), *self->ctx); complete(); } + }); + }); + } + catch (...) { router_.handle_exception(std::current_exception(), *ctx); complete(); } +} + +void HttpSession::async_read_stream_body(net::mutable_buffer target, HttpRequestStream::ReadCallback callback) +{ + auto self = shared_from_this(); + net::post(stream_.get_executor(), [self, target, callback = std::move(callback)]() mutable + { + if (self->request_body_cancelled_) return callback(net::error::operation_aborted, 0, true); + if (!self->request_parser_ || self->request_parser_->is_done()) return callback({}, 0, true); + auto& body = self->request_parser_->get().body(); + body.data = target.data(); + body.size = target.size(); + http::async_read_some(self->stream_, self->buffer_, *self->request_parser_, + [self, capacity = target.size(), callback = std::move(callback)](beast::error_code ec, std::size_t) mutable + { + if (ec == http::error::need_buffer) ec = {}; + const auto produced = capacity - self->request_parser_->get().body().size; + callback(ec, produced, self->request_parser_->is_done()); + }); + }); +} + +void HttpSession::cancel_stream_body() +{ + auto self = shared_from_this(); + net::post(stream_.get_executor(), [self] + { + self->request_body_cancelled_ = true; + // Cancel an outstanding body read without closing the socket. Subsequent + // response writes remain valid; the connection is made non-persistent + // because unread request bytes cannot be parsed as a next request. + self->res_.keep_alive(false); + self->stream_.cancel(); + }); +} + +void HttpSession::cancel_session() +{ + auto self = shared_from_this(); + net::post(stream_.get_executor(), [self] + { + beast::error_code ignored; + self->stream_.cancel(); + self->stream_.socket().shutdown(tcp::socket::shutdown_both, ignored); + }); +} + +void HttpSession::start_stream_response(HttpResponseStream::ResponseHead head, HttpResponseStream::Callback callback) +{ + auto self = shared_from_this(); + net::post(stream_.get_executor(), [self, head = std::move(head), callback = std::move(callback)]() mutable + { + self->streaming_response_ = {}; + self->streaming_response_.result(head.result()); self->streaming_response_.version(head.version()); + self->streaming_response_.keep_alive(head.keep_alive()); + if (self->request_body_cancelled_) self->streaming_response_.keep_alive(false); + for (const auto& field : head) self->streaming_response_.insert(field.name_string(), field.value()); + if (!self->streaming_response_.has_content_length() && !self->streaming_response_.chunked()) + self->streaming_response_.chunked(true); + self->streaming_response_.body().more = true; + self->streaming_response_serializer_.emplace(self->streaming_response_); + http::async_write_header(self->stream_, *self->streaming_response_serializer_, + [callback = std::move(callback)](beast::error_code ec, std::size_t) mutable + { callback(ec); }); + }); +} + +void HttpSession::write_stream_response(net::const_buffer source, HttpResponseStream::Callback callback) +{ + auto self = shared_from_this(); + net::post(stream_.get_executor(), [self, source, callback = std::move(callback)]() mutable + { + if (!self->streaming_response_serializer_) return callback(net::error::operation_aborted); + auto& body = self->streaming_response_.body(); + body.data = const_cast(source.data()); body.size = source.size(); body.more = true; + http::async_write(self->stream_, *self->streaming_response_serializer_, + [callback = std::move(callback)](beast::error_code ec, std::size_t) mutable + { if (ec == http::error::need_buffer) ec = {}; callback(ec); }); + }); +} + +void HttpSession::finish_stream_response(HttpResponseStream::Callback callback) +{ + auto self = shared_from_this(); + net::post(stream_.get_executor(), [self, callback = std::move(callback)]() mutable + { + if (!self->streaming_response_serializer_) return callback(net::error::operation_aborted); + if (self->streaming_response_serializer_->is_done()) + { + callback({}); + const bool keep_alive = self->streaming_response_.keep_alive(); + self->streaming_response_serializer_.reset(); + if (keep_alive) self->do_read(); else self->do_close(); + return; + } + auto& body = self->streaming_response_.body(); body.data = nullptr; body.size = 0; body.more = false; + http::async_write(self->stream_, *self->streaming_response_serializer_, + [self, callback = std::move(callback)](beast::error_code ec, std::size_t) mutable + { + callback(ec); + const bool keep_alive = self->streaming_response_.keep_alive(); + self->streaming_response_serializer_.reset(); + if (!ec && keep_alive) self->do_read(); else self->do_close(); + }); + }); } void HttpSession::on_read(const beast::error_code& ec, std::size_t bytes_transferred) @@ -82,59 +369,70 @@ void HttpSession::handle_request() { res_ = {}; - ctx = std::make_shared(req_, res_); + ctx = std::make_shared(req_, res_, peer_endpoint_); + + try + { + router_.async_run_pre_interceptors(*ctx, [self = shared_from_this()](InterceptorResult result) + { net::post(self->stream_.get_executor(), [self, result] { self->dispatch_request_after_interceptors(result); }); }); + } + catch (...) + { + router_.handle_exception(std::current_exception(), *ctx); + send_response(std::move(res_)); + } +} +void HttpSession::dispatch_request_after_interceptors(InterceptorResult result) +{ try { - // 1. Run Pre-interceptors - if (router_.run_pre_interceptors(*ctx) == InterceptorResult::Stop) + if (result == InterceptorResult::Stop) { router_.run_post_interceptors(*ctx); - - if (res_.chunked()) - { - send_chunked_response(); - } - else - { - send_response(std::move(res_)); - } - return; + return send_context_response(); } + // Asynchronous routes take precedence when registered for this method. + if (router_.dispatch_async(*ctx, [self = shared_from_this()] + { net::post(self->stream_.get_executor(), [self] + { + try { self->router_.run_post_interceptors(*self->ctx); self->send_context_response(); } + catch (...) { self->router_.handle_exception(std::current_exception(), *self->ctx); self->send_context_response(); } + }); })) return; + bool static_file_served = false; - // 2. Dispatch to routes or static files router_.dispatch(*ctx, [this, &static_file_served] { if (req_.method() == http::verb::get || req_.method() == http::verb::head) - { static_file_served = do_serve_static_file(); - } return static_file_served; }); - - if (static_file_served) - { - return; - } - - // 3. Run Post-interceptors + if (static_file_served) return; router_.run_post_interceptors(*ctx); - - if (res_.chunked()) - { - send_chunked_response(); - } - else - { - send_response(std::move(res_)); - } + send_context_response(); } catch (...) { router_.handle_exception(std::current_exception(), *ctx); - send_response(std::move(res_)); + send_context_response(); + } +} + +void HttpSession::send_context_response() +{ + if (request_body_cancelled_) res_.keep_alive(false); + if (req_.method() == http::verb::head) + { + http::response head{res_.result(), res_.version()}; + head.keep_alive(res_.keep_alive()); + for (const auto& field : res_) head.insert(field.name_string(), field.value()); + head.erase(http::field::transfer_encoding); + if (!head.has_content_length()) head.content_length(res_.body().size()); + return send_response(std::move(head)); } + if (res_.chunked()) send_chunked_response(); + else send_response(std::move(res_)); } // Extract path from request target (query-stripped) @@ -452,6 +750,7 @@ void HttpSession::on_write(bool keep_alive, beast::error_code ec, std::size_t by void HttpSession::do_close() { + spdlog::debug("HttpSession closing connection"); beast::error_code ec; stream_.socket().shutdown(tcp::socket::shutdown_send, ec); if (ec) diff --git a/framework/session/http_session.hpp b/framework/session/http_session.hpp index 365d5d2..1abc502 100644 --- a/framework/session/http_session.hpp +++ b/framework/session/http_session.hpp @@ -7,6 +7,8 @@ #include #include #include +#include +#include #include "router/http_router.hpp" #include "websocket/websocket_session.hpp" @@ -25,28 +27,42 @@ namespace khttpd::framework class HttpSession : public std::enable_shared_from_this { public: + static constexpr std::uint64_t default_max_buffered_request_body_size = 16ULL * 1024 * 1024; + HttpSession(tcp::socket&& socket, HttpRouter& router, WebsocketRouter& ws_router, const std::string& web_root, - const boost::filesystem::path& canonical_web_root); + const boost::filesystem::path& canonical_web_root, + std::uint64_t max_buffered_request_body_size = default_max_buffered_request_body_size); // 启动会话 void run(); private: struct ChunkWriteState; + class RequestStreamImpl; + class ResponseStreamImpl; bool disable_web_root_ = false; beast::tcp_stream stream_; beast::flat_buffer buffer_; http::request req_; + std::optional> request_parser_; + std::array buffered_body_chunk_{}; + std::string buffered_body_; http::response res_; HttpRouter& router_; WebsocketRouter& websocket_router_; const boost::filesystem::path web_root_path_; const boost::filesystem::path canonical_web_root_path_; + const std::uint64_t max_buffered_request_body_size_; std::shared_ptr ws_session_; std::optional> sr_; + http::response streaming_response_; + std::optional> streaming_response_serializer_; std::shared_ptr ctx = nullptr; + bool stream_completed_ = false; + std::optional peer_endpoint_; + bool request_body_cancelled_ = false; // Chunked streaming support std::shared_ptr> chunk_queue_; @@ -55,9 +71,23 @@ namespace khttpd::framework std::shared_ptr chunk_error_; void do_read(); + void on_read_header(const beast::error_code& ec, std::size_t bytes_transferred); + void read_buffered_body(); + void on_read_buffered_body(beast::error_code ec, std::size_t bytes_transferred); + void send_payload_too_large(); + void copy_request_head(); + void handle_stream_request(); + void async_read_stream_body(net::mutable_buffer target, HttpRequestStream::ReadCallback callback); + void cancel_stream_body(); + void cancel_session(); + void start_stream_response(HttpResponseStream::ResponseHead head, HttpResponseStream::Callback callback); + void write_stream_response(net::const_buffer source, HttpResponseStream::Callback callback); + void finish_stream_response(HttpResponseStream::Callback callback); void on_read(const beast::error_code& ec, std::size_t bytes_transferred); void handle_request(); + void dispatch_request_after_interceptors(InterceptorResult result); + void send_context_response(); // 新增:尝试处理静态文件请求 bool do_serve_static_file(); diff --git a/framework/tests/BUILD.bazel b/framework/tests/BUILD.bazel index b3b234a..b0a4a1f 100644 --- a/framework/tests/BUILD.bazel +++ b/framework/tests/BUILD.bazel @@ -136,6 +136,48 @@ cc_test( ], ) +cc_test( + name = "buffered_body_limit_test", + srcs = ["buffered_body_limit_test.cpp", "http_session_test_harness.hpp"], + copts = ["-std=c++17", "-Wall", "-pedantic"], + deps = ["//framework", "@googletest//:gtest", "@googletest//:gtest_main"], +) + +cc_test( + name = "http_stream_edge_test", + srcs = ["http_stream_edge_test.cpp", "http_session_test_harness.hpp"], + copts = ["-std=c++17", "-Wall", "-pedantic"], + deps = ["//framework", "@googletest//:gtest", "@googletest//:gtest_main"], +) + +cc_test( + name = "http_client_stream_edge_test", + srcs = ["http_client_stream_edge_test.cpp"], + copts = ["-std=c++17", "-Wall", "-pedantic"], + deps = ["//framework", "@googletest//:gtest", "@googletest//:gtest_main"], +) + +cc_test( + name = "http_proxy_session_edge_test", + srcs = ["http_proxy_session_edge_test.cpp"], + copts = ["-std=c++17", "-Wall", "-pedantic"], + deps = ["//framework", "@googletest//:gtest", "@googletest//:gtest_main"], +) + +cc_test( + name = "websocket_router_dynamic_test", + srcs = ["websocket_router_dynamic_test.cpp"], + copts = ["-std=c++17", "-Wall", "-pedantic"], + deps = ["//framework", "@googletest//:gtest", "@googletest//:gtest_main"], +) + +cc_test( + name = "websocket_client_handshake_test", + srcs = ["websocket_client_handshake_test.cpp"], + copts = ["-std=c++17", "-Wall", "-pedantic"], + deps = ["//framework", "@googletest//:gtest", "@googletest//:gtest_main"], +) + cc_test( name = "server_stability_test", srcs = ["server_stability_test.cpp"], diff --git a/framework/tests/buffered_body_limit_test.cpp b/framework/tests/buffered_body_limit_test.cpp new file mode 100644 index 0000000..8f08e1c --- /dev/null +++ b/framework/tests/buffered_body_limit_test.cpp @@ -0,0 +1,70 @@ +#include "http_session_test_harness.hpp" +#include "framework/server.hpp" + +#include + +namespace fw = khttpd::framework; +namespace test = khttpd::framework::tests; +namespace http = boost::beast::http; + +namespace +{ + http::response send_body(std::uint64_t limit, std::string body, + bool chunked, bool& handled) + { + test::TempWebRoot web_root; + fw::HttpRouter router; + fw::WebsocketRouter websocket_router; + router.post("/body", [&handled](fw::HttpContext& ctx) + { + handled = true; + ctx.set_body(std::to_string(ctx.get_request().body().size())); + }); + http::request request{http::verb::post, "/body", 11}; + request.body() = std::move(body); + if (chunked) request.chunked(true); else request.prepare_payload(); + request.keep_alive(false); + return test::round_trip(router, websocket_router, web_root.path, std::move(request), limit); + } +} + +TEST(BufferedBodyLimitTest, DefaultLimitRemainsSixteenMiB) +{ + EXPECT_EQ(fw::HttpSession::default_max_buffered_request_body_size, 16ULL * 1024 * 1024); + EXPECT_EQ(fw::Server::default_max_buffered_request_body_size, 16ULL * 1024 * 1024); +} + +TEST(BufferedBodyLimitTest, ZeroLimitAcceptsEmptyBody) +{ + bool handled = false; + const auto response = send_body(0, {}, false, handled); + EXPECT_EQ(response.result(), http::status::ok); + EXPECT_TRUE(handled); + EXPECT_EQ(response.body(), "0"); +} + +TEST(BufferedBodyLimitTest, ZeroLimitRejectsOneByte) +{ + bool handled = false; + const auto response = send_body(0, "x", false, handled); + EXPECT_EQ(response.result(), http::status::payload_too_large); + EXPECT_FALSE(handled); + EXPECT_FALSE(response.keep_alive()); +} + +TEST(BufferedBodyLimitTest, ContentLengthAtLimitIsAccepted) +{ + bool handled = false; + const auto response = send_body(17, std::string(17, 'x'), false, handled); + EXPECT_EQ(response.result(), http::status::ok); + EXPECT_TRUE(handled); + EXPECT_EQ(response.body(), "17"); +} + +TEST(BufferedBodyLimitTest, ChunkedBodyIsRejectedAfterCumulativeLimit) +{ + bool handled = false; + const auto response = send_body(31, std::string(32, 'x'), true, handled); + EXPECT_EQ(response.result(), http::status::payload_too_large); + EXPECT_FALSE(handled); +} diff --git a/framework/tests/client_test.cpp b/framework/tests/client_test.cpp index c13cd1d..53d1060 100644 --- a/framework/tests/client_test.cpp +++ b/framework/tests/client_test.cpp @@ -1,4 +1,5 @@ #include "framework/client/http_client.hpp" +#include "framework/client/http_client_stream.hpp" #include "framework/client/websocket_client.hpp" #include "framework/client/api_macros.hpp" #include "framework/client/host_pool.hpp" @@ -134,9 +135,11 @@ class LocalHttpEchoServer { beast::flat_buffer buffer; boost::system::error_code ec; - http::request req; - http::read(socket, buffer, req, ec); + http::request_parser parser; + parser.body_limit(8 * 1024 * 1024); + http::read(socket, buffer, parser, ec); if (ec) return; + auto req = parser.release(); http::response res{http::status::ok, req.version()}; res.set(http::field::server, "khttpd-local-echo"); @@ -163,6 +166,11 @@ class LocalHttpEchoServer { res.body() = "{\"data\":" + req.body() + "}"; } + else if (target.rfind("/stream", 0) == 0) + { + res.set(http::field::content_type, "application/octet-stream"); + res.body() = req.body(); + } else { res.body() = "{\"target\":\"" + target + "\"}"; @@ -608,6 +616,58 @@ TEST(HttpClientLocalTest, SyncRequestTimeoutClosesStalledConnection) // WebSocket 测试 // ========================================== +TEST(HttpClientStreamTest, StreamsRequestAndResponseWithBoundedBuffers) +{ + LocalHttpEchoServer server; + net::io_context ioc; + auto client = std::make_shared(ioc); + auto payload = std::make_shared(2 * 1024 * 1024 + 31, 's'); + auto offset = std::make_shared(0); + auto received = std::make_shared(0); + auto write_next = std::make_shared>(); + auto read_next = std::make_shared>(); + auto read_buffer = std::make_shared>(); + bool completed = false; + + HttpClientStream::RequestHead head{http::verb::post, "/stream", 11}; + head.content_length(payload->size()); + *read_next = [client, received, read_buffer, read_next, &completed] + { + client->async_read_some(net::buffer(*read_buffer), + [client, received, read_buffer, read_next, &completed](beast::error_code ec, std::size_t n, bool done) + { + ASSERT_FALSE(ec) << ec.message(); *received += n; + if (!done) return (*read_next)(); + completed = true; + }); + }; + *write_next = [client, payload, offset, write_next, read_next] + { + if (*offset == payload->size()) + { + return client->async_finish_request([client, read_next](beast::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + client->async_read_response_head([read_next](beast::error_code head_ec, HttpClientStream::ResponseHead head) + { + ASSERT_FALSE(head_ec) << head_ec.message(); EXPECT_EQ(head.result(), http::status::ok); (*read_next)(); + }); + }); + } + const auto size = std::min(32 * 1024, payload->size() - *offset); + const auto chunk = net::buffer(payload->data() + *offset, size); + client->async_write_some(chunk, [offset, size, write_next](beast::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); *offset += size; (*write_next)(); + }); + }; + client->async_start(server.base_url() + "/stream", std::move(head), + [write_next](beast::error_code ec) { ASSERT_FALSE(ec) << ec.message(); (*write_next)(); }); + ioc.run(); + EXPECT_TRUE(completed); + EXPECT_EQ(*received, payload->size()); +} + class WebsocketTest : public ::testing::Test { protected: @@ -706,6 +766,36 @@ TEST_F(WebsocketTest, WssEchoAndWriteQueue) EXPECT_TRUE(closed_gracefully) << "on_close should be triggered"; } +TEST_F(WebsocketTest, FrameHandlerPreservesBinaryTypeAndPayload) +{ + LocalWebSocketEchoServer server; + bool connected = false; + bool received = false; + boost::asio::steady_timer timer(ioc, std::chrono::seconds(5)); + + ws_client->set_on_frame([&](const khttpd::framework::WebsocketFrame& frame) + { + if (frame.type != khttpd::framework::WebsocketFrameType::binary) return; + received = frame.payload == std::string("\x00\x01\xff", 3); + ws_client->close(); + timer.cancel(); + }); + ws_client->connect(server.url(), [&](boost::beast::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + connected = true; + ws_client->send({khttpd::framework::WebsocketFrameType::binary, std::string("\x00\x01\xff", 3)}); + }); + timer.async_wait([&](boost::system::error_code ec) + { + if (!ec) { ws_client->close(); ADD_FAILURE() << "Frame echo timed out"; } + }); + ioc.run(); + + EXPECT_TRUE(connected); + EXPECT_TRUE(received); +} + TEST_F(WebsocketTest, ConnectFailure) { // 测试连接不可达端口 diff --git a/framework/tests/http_client_stream_edge_test.cpp b/framework/tests/http_client_stream_edge_test.cpp new file mode 100644 index 0000000..6ad06f5 --- /dev/null +++ b/framework/tests/http_client_stream_edge_test.cpp @@ -0,0 +1,273 @@ +#include "framework/client/http_client_stream.hpp" +#include "framework/client/http_client.hpp" + +#include +#include +#include +#include +#include + +namespace client = khttpd::framework::client; +namespace http = boost::beast::http; +namespace net = boost::asio; +namespace ssl = net::ssl; +using tcp = net::ip::tcp; + +namespace +{ + constexpr char test_certificate[] = R"PEM(-----BEGIN CERTIFICATE----- +MIIDCTCCAfGgAwIBAgIUThDTQ+lb7kmBiXurHn9Wy6vf+54wDQYJKoZIhvcNAQEL +BQAwFDESMBAGA1UEAwwJbG9jYWxob3N0MB4XDTI2MDgwOTE0NDcwN1oXDTI2MDgx +MDE0NDcwN1owFDESMBAGA1UEAwwJbG9jYWxob3N0MIIBIjANBgkqhkiG9w0BAQEF +AAOCAQ8AMIIBCgKCAQEAkZayJOWhyvs9PBYbxg9pZns1OoVkr5Gs7DmDuzZSpsKW +0MrWOMCyvQVOxyOwsjJ3V4x3K4aOZMhLzBe0mewiQ+2LDTD0bqfg2mHf8x1xkwAc +WM25qiNo6yNDiFzCrxFgMRck4ejlr+YGSvKclrVoHny2Rmr8Zsvb/4SDRMIL/Qej +eOhWecl+XiE99+YQUJWbPenXHWfgGhcyvdQmiSQ3FJYv73bpsbKQ6Dw5eaTe7hlN +6M2qW39I6p7Y2I0RnzEPJ0USGXo2DLi5pzWzQDcXupC9ciH8nq2MpFvj8gZcfoEg +5ml/RAFVmLWiGyV18JIFFQAmFUTN58iw7vxKzbtvwQIDAQABo1MwUTAdBgNVHQ4E +FgQUGbjhcOG4N6JZLiLCWhxbYBM7LbQwHwYDVR0jBBgwFoAUGbjhcOG4N6JZLiLC +WhxbYBM7LbQwDwYDVR0TAQH/BAUwAwEB/zANBgkqhkiG9w0BAQsFAAOCAQEAK37S +esIrxwdHPpKbIokLsf2/1QCVMtxNqTgFJ7GUNeg8Psu69uXbPCS3+1gRUP0385cW +goX/Ok0NVymKn+qIdZU84t+ClEFFVgrZ6onknSdRKdvEqTUsFJ5UAf/Tbq0cpOmz +bhYuFAbwwZnaLhRdsv08A8m0ql29lTMzCByVD5crvADEuXR3wqjdbbBI0zxcjK+J +pj819xaKeS37VEQ9jjwq5+w6rlaIIqy1flwsSBDp+DWB2A9nhIkwE3qfmxlOnTC9 +alJ1edM4GDkgOW+q/iIrLSLq7UFl+JgQYD+jHXF3gDgjqK/riDjWYxO3TabwjcBz +d+xV3kUBYw3hrnwi7A== +-----END CERTIFICATE-----)PEM"; + + constexpr char test_private_key[] = R"PEM(-----BEGIN PRIVATE KEY----- +MIIEvAIBADANBgkqhkiG9w0BAQEFAASCBKYwggSiAgEAAoIBAQCRlrIk5aHK+z08 +FhvGD2lmezU6hWSvkazsOYO7NlKmwpbQytY4wLK9BU7HI7CyMndXjHcrho5kyEvM +F7SZ7CJD7YsNMPRup+DaYd/zHXGTABxYzbmqI2jrI0OIXMKvEWAxFyTh6OWv5gZK +8pyWtWgefLZGavxmy9v/hINEwgv9B6N46FZ5yX5eIT335hBQlZs96dcdZ+AaFzK9 +1CaJJDcUli/vdumxspDoPDl5pN7uGU3ozapbf0jqntjYjRGfMQ8nRRIZejYMuLmn +NbNANxe6kL1yIfyerYykW+PyBlx+gSDmaX9EAVWYtaIbJXXwkgUVACYVRM3nyLDu +/ErNu2/BAgMBAAECggEAAP4QvVEmavKO/o2dB1rcClONL5awssSws9SJihlq81GQ +wyAa2TyxCzpRyOg8oF5ZM2rU9iI+7r9xytSficwTCLkCEWczx1xUG1D+/JKHD2w5 +BT7zxM3kfXPaVj/hoN1itTr16KdUh4AvK0wflqRqbwjFGlJI4a+CkqmV1n5nJASq +bb1XgM6aPJwek8MuYg8c5HY7Ul3vJyTWJKWobdG+1w0fvnWA501W7eXfBexTaCZQ +VPwpSGC4tz00XbS17defuUD8zlDJS3jETZAq5sGZMsVNoLXskaW5h3ex5aJRYNak +JC/93l/5znFLtIX2igDmis7BfO9qf9hJeeed+b/qUQKBgQDHBGRFr74poaniQIez +AXpddagd+rZ657v2GJJn+HyoyMTWt7wqDNlZgOmXmp31z4WOvs5Anb1WlaoUOM+R +m1rzuS3ool2AkkqxHo8v95xJjUKQ0iO0NsHSGhPwq/8ftEwT7mWqoz54s81wEO3J +1YRt5r1lQyrCNPDDnsGD6w4beQKBgQC7RhbDxmb9g6z4786YMXRf+lNTFIHSmaaR +mzMWrlobF0Na46OeqbkjNighWCpkekdCEqaVaYUm7gAGbiEhJgtcZTsVzwTdcTnf +qOZs+qxPCjTgQti7iyMRYwsky5BTw7/gzPYwY4lQHpS8PkQ+PzA7hTCtOnKdq82S +YO/sELmciQKBgGNs0T9zRhh8WGfc/y4xrdUlI4EesK2EOgX/Tp08qeKUsqnmjs2f +L7KkUY7YwtN8AmhG8LmdVGr+SELkAubmazDZsZLIEthZvZDxCG3ZUS35sWiyYv30 +YS46sv2In+NR6rQGZKoz9dDNWvQCsRklX4ycOsBtJt5xHltMY7co5hpZAoGAU3fn +yZZybOf1fnaT5C2WqviNjugC/PTS0u8TlDZdntl9gdMYKC2JgPIwbLw5GNOPUxmw ++cMwP6uwgy0uwvGL+sB71zqP9oryuoczPLt1dT0dWB8zLlPTa3pzixDX4R3MNcvk +pqiWmQkoTcaK8BuFyeGRUoRMdY4PcACYruS9ddECgYBUAM/XcdBamOkHavC1Khut +lL7kS/l+M8K6Q2ghxEPLQzPD+gQvlUT+zjpoP9jHK3gHq3Im4bayFvmaqX+t0Piy +CQmqJob1eraZfRUHk7QnsOO0TF+s6UWf1bO3rYZBLdr15sWcAkC7e0GhfGEhSFzW +C6fgselIUr+XYmAQoT5Tqw== +-----END PRIVATE KEY-----)PEM"; +} + +TEST(HttpClientStreamEdgeTest, StreamsRequestAndResponseOverTls) +{ + net::io_context server_ioc; + ssl::context server_context(ssl::context::tls_server); + server_context.use_certificate_chain(net::buffer(test_certificate)); + server_context.use_private_key(net::buffer(test_private_key), ssl::context::file_format::pem); + tcp::acceptor acceptor(server_ioc, {net::ip::address_v4::loopback(), 0}); + const auto port = acceptor.local_endpoint().port(); + const std::string payload(4097, 't'); + std::thread server([&] + { + tcp::socket socket(server_ioc); + acceptor.accept(socket); + boost::beast::ssl_stream tls(std::move(socket), server_context); + tls.handshake(ssl::stream_base::server); + boost::beast::flat_buffer buffer; + http::request_parser parser; + parser.body_limit(8192); + http::read(tls, buffer, parser); + auto request = parser.release(); + http::response response{http::status::ok, 11}; + response.body() = std::move(request.body()); + response.prepare_payload(); + response.keep_alive(false); + boost::system::error_code ignored; + http::write(tls, response, ignored); + }); + + net::io_context ioc; + ssl::context client_context(ssl::context::tls_client); + client_context.set_verify_mode(ssl::verify_none); + auto stream = std::make_shared(ioc, client_context); + auto write_offset = std::make_shared(0); + auto received = std::make_shared(); + auto write_next = std::make_shared>(); + auto read_next = std::make_shared>(); + auto read_buffer = std::make_shared>(); + bool completed = false; + client::HttpClientStream::RequestHead head{http::verb::post, "/ignored", 11}; + head.content_length(payload.size()); + + *read_next = [stream, received, read_buffer, read_next, &completed] + { + stream->async_read_some(net::buffer(*read_buffer), + [stream, received, read_buffer, read_next, &completed] + (boost::system::error_code ec, std::size_t size, bool done) + { + ASSERT_FALSE(ec) << ec.message(); + received->append(read_buffer->data(), size); + if (!done) return (*read_next)(); + completed = true; + }); + }; + *write_next = [stream, &payload, write_offset, write_next, read_next] + { + if (*write_offset == payload.size()) + return stream->async_finish_request([stream, read_next](boost::system::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + stream->async_read_response_head([read_next] + (boost::system::error_code head_ec, client::HttpClientStream::ResponseHead head) + { + ASSERT_FALSE(head_ec) << head_ec.message(); + EXPECT_EQ(head.result(), http::status::ok); + (*read_next)(); + }); + }); + const auto size = std::min(7, payload.size() - *write_offset); + stream->async_write_some(net::buffer(payload.data() + *write_offset, size), + [write_offset, size, write_next](boost::system::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + *write_offset += size; + (*write_next)(); + }); + }; + stream->async_start("https://127.0.0.1:" + std::to_string(port) + "/upload?part=1", + std::move(head), [write_next](boost::system::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + (*write_next)(); + }); + ioc.run(); + server.join(); + EXPECT_TRUE(completed); + EXPECT_EQ(*received, payload); +} + +TEST(HttpClientStreamEdgeTest, RejectsMalformedUrl) +{ + net::io_context ioc; + auto stream = std::make_shared(ioc); + client::HttpClientStream::RequestHead head{http::verb::get, "/", 11}; + boost::system::error_code result; + stream->async_start("not a url", std::move(head), + [&](boost::system::error_code ec) { result = ec; }); + ioc.run(); + EXPECT_EQ(result, make_error_code(boost::system::errc::operation_not_supported)); +} + +TEST(HttpClientStreamEdgeTest, WriteBeforeStartIsAborted) +{ + net::io_context ioc; + auto stream = std::make_shared(ioc); + const std::array data{'x'}; + boost::system::error_code result; + stream->async_write_some(net::buffer(data), [&](boost::system::error_code ec) { result = ec; }); + ioc.run(); + EXPECT_EQ(result, net::error::operation_aborted); +} + +TEST(HttpClientStreamEdgeTest, FinishBeforeStartIsAborted) +{ + net::io_context ioc; + auto stream = std::make_shared(ioc); + boost::system::error_code result; + stream->async_finish_request([&](boost::system::error_code ec) { result = ec; }); + ioc.run(); + EXPECT_EQ(result, net::error::operation_aborted); +} + +TEST(HttpClientStreamEdgeTest, HeadSkipsBodyAfterConsecutiveInformationalResponses) +{ + net::io_context server_ioc; + tcp::acceptor acceptor(server_ioc, {net::ip::address_v4::loopback(), 0}); + const auto port = acceptor.local_endpoint().port(); + std::thread server([&] + { + tcp::socket socket(server_ioc); + acceptor.accept(socket); + boost::beast::flat_buffer request_buffer; + http::request request; + http::read(socket, request_buffer, request); + const std::string wire = + "HTTP/1.1 100 Continue\r\n\r\n" + "HTTP/1.1 103 Early Hints\r\nLink: ; rel=preload\r\n\r\n" + "HTTP/1.1 200 OK\r\nContent-Length: 1234\r\nConnection: close\r\n\r\n"; + net::write(socket, net::buffer(wire)); + }); + + net::io_context ioc; + auto stream = std::make_shared(ioc); + client::HttpClientStream::RequestHead request{http::verb::head, "/", 11}; + bool done = false; + stream->async_start("http://127.0.0.1:" + std::to_string(port) + "/", std::move(request), + [stream, &done](boost::system::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + stream->async_finish_request([stream, &done](boost::system::error_code finish_ec) + { + ASSERT_FALSE(finish_ec) << finish_ec.message(); + stream->async_read_response_head([stream, &done] + (boost::system::error_code head_ec, client::HttpClientStream::ResponseHead head) + { + ASSERT_FALSE(head_ec) << head_ec.message(); + EXPECT_EQ(head.result(), http::status::ok); + EXPECT_EQ(head[http::field::content_length], "1234"); + std::array body{}; + stream->async_read_some(net::buffer(body), [&done] + (boost::system::error_code read_ec, std::size_t size, bool body_done) + { + EXPECT_FALSE(read_ec); + EXPECT_EQ(size, 0u); + EXPECT_TRUE(body_done); + done = true; + }); + }); + }); + }); + ioc.run(); + server.join(); + EXPECT_TRUE(done); +} + +TEST(HttpClientStreamEdgeTest, BufferedClientAlsoHandlesHeadAndInformationalResponses) +{ + net::io_context server_ioc; + tcp::acceptor acceptor(server_ioc, {net::ip::address_v4::loopback(), 0}); + const auto port = acceptor.local_endpoint().port(); + std::thread server([&] + { + tcp::socket socket(server_ioc); + acceptor.accept(socket); + boost::beast::flat_buffer request_buffer; + http::request request; + http::read(socket, request_buffer, request); + const std::string wire = "HTTP/1.1 103 Early Hints\r\n\r\n" + "HTTP/1.1 200 OK\r\nContent-Length: 99\r\nConnection: close\r\n\r\n"; + net::write(socket, net::buffer(wire)); + }); + net::io_context ioc; + auto http_client = std::make_shared(ioc); + bool done = false; + http_client->request(http::verb::head, "http://127.0.0.1:" + std::to_string(port) + "/", + {}, {}, {}, [&](boost::system::error_code ec, http::response response) + { + EXPECT_FALSE(ec) << ec.message(); + EXPECT_EQ(response.result(), http::status::ok); + EXPECT_EQ(response[http::field::content_length], "99"); + EXPECT_TRUE(response.body().empty()); + done = true; + }); + ioc.run(); + server.join(); + EXPECT_TRUE(done); +} diff --git a/framework/tests/http_proxy_session_edge_test.cpp b/framework/tests/http_proxy_session_edge_test.cpp new file mode 100644 index 0000000..153a1c5 --- /dev/null +++ b/framework/tests/http_proxy_session_edge_test.cpp @@ -0,0 +1,77 @@ +#include "framework/client/http_proxy_session.hpp" + +#include +#include + +namespace fw = khttpd::framework; +namespace client = khttpd::framework::client; +namespace http = boost::beast::http; +namespace net = boost::asio; + +namespace +{ + class ObservableRequestStream final : public fw::HttpRequestStream + { + public: + bool read = false; + bool cancelled = false; + + void async_read_some(net::mutable_buffer, ReadCallback callback) override + { + read = true; + callback({}, 0, true); + } + void cancel() override { cancelled = true; } + }; + + class ObservableResponseStream final : public fw::HttpResponseStream + { + public: + bool started = false; + bool cancelled = false; + + void async_start(ResponseHead, Callback callback) override + { + started = true; + callback({}); + } + void async_write_some(net::const_buffer, Callback callback) override { callback({}); } + void async_finish(Callback callback) override { callback({}); } + void cancel() override { cancelled = true; } + }; +} + +TEST(HttpProxySessionEdgeTest, UnsupportedUpstreamSchemeCancelsBothDownstreamSides) +{ + net::io_context ioc; + auto inbound = std::make_shared(); + auto downstream = std::make_shared(); + auto proxy = std::make_shared(ioc, inbound, downstream, 17); + client::HttpClientStream::RequestHead head{http::verb::post, "/", 11}; + boost::system::error_code result; + bool completed = false; + + proxy->start("ftp://unsupported.test/upload", std::move(head), + [&](boost::system::error_code ec) { result = ec; completed = true; }); + ioc.run(); + + EXPECT_TRUE(completed); + EXPECT_EQ(result, make_error_code(boost::system::errc::operation_not_supported)); + EXPECT_TRUE(inbound->cancelled); + EXPECT_TRUE(downstream->cancelled); + EXPECT_FALSE(inbound->read); + EXPECT_FALSE(downstream->started); +} + +TEST(HttpProxySessionEdgeTest, ExplicitCancelCancelsInboundAndDownstream) +{ + net::io_context ioc; + auto inbound = std::make_shared(); + auto downstream = std::make_shared(); + auto proxy = std::make_shared(ioc, inbound, downstream); + + proxy->cancel(); + + EXPECT_TRUE(inbound->cancelled); + EXPECT_TRUE(downstream->cancelled); +} diff --git a/framework/tests/http_session_test_harness.hpp b/framework/tests/http_session_test_harness.hpp new file mode 100644 index 0000000..acdb96f --- /dev/null +++ b/framework/tests/http_session_test_harness.hpp @@ -0,0 +1,78 @@ +#ifndef KHTTPD_FRAMEWORK_TESTS_HTTP_SESSION_TEST_HARNESS_HPP +#define KHTTPD_FRAMEWORK_TESTS_HTTP_SESSION_TEST_HARNESS_HPP + +#include "framework/router/http_router.hpp" +#include "framework/router/websocket_router.hpp" +#include "framework/session/http_session.hpp" + +#include +#include +#include +#include +#include + +namespace khttpd::framework::tests +{ + namespace beast = boost::beast; + namespace http = beast::http; + namespace net = boost::asio; + namespace fs = boost::filesystem; + using tcp = net::ip::tcp; + + struct TempWebRoot + { + fs::path path = fs::temp_directory_path() / fs::unique_path("khttpd-http-test-%%%%-%%%%"); + + TempWebRoot() { fs::create_directories(path); } + ~TempWebRoot() + { + boost::system::error_code ec; + fs::remove_all(path, ec); + } + }; + + template + http::response round_trip(HttpRouter& router, WebsocketRouter& websocket_router, + const fs::path& web_root, + http::request request, + std::uint64_t max_buffered_body = + HttpSession::default_max_buffered_request_body_size) + { + net::io_context server_ioc; + auto guard = net::make_work_guard(server_ioc); + tcp::acceptor acceptor(server_ioc, {net::ip::make_address("127.0.0.1"), 0}); + const auto endpoint = acceptor.local_endpoint(); + const auto canonical_web_root = fs::canonical(web_root); + + acceptor.async_accept([&](beast::error_code ec, tcp::socket socket) + { + if (ec) return; + std::make_shared(std::move(socket), router, websocket_router, + web_root.string(), canonical_web_root, + max_buffered_body)->run(); + }); + std::thread server_thread([&] { server_ioc.run(); }); + + net::io_context client_ioc; + tcp::socket client(client_ioc); + client.connect(endpoint); + beast::error_code ec; + http::write(client, request, ec); + if (ec) throw boost::system::system_error(ec); + + beast::flat_buffer buffer; + http::response_parser parser; + http::read(client, buffer, parser, ec); + if (ec) throw boost::system::system_error(ec); + auto response = parser.release(); + + beast::error_code ignored; + client.close(ignored); + guard.reset(); + server_ioc.stop(); + server_thread.join(); + return response; + } +} + +#endif diff --git a/framework/tests/http_stream_edge_test.cpp b/framework/tests/http_stream_edge_test.cpp new file mode 100644 index 0000000..3b375f0 --- /dev/null +++ b/framework/tests/http_stream_edge_test.cpp @@ -0,0 +1,117 @@ +#include "http_session_test_harness.hpp" + +#include +#include + +namespace fw = khttpd::framework; +namespace test = khttpd::framework::tests; +namespace http = boost::beast::http; +namespace net = boost::asio; + +namespace +{ + void install_counting_stream(fw::HttpRouter& router, std::size_t read_buffer_size = 5) + { + router.stream("/stream", http::verb::post, + [read_buffer_size](fw::HttpContext& ctx, std::shared_ptr request, + std::shared_ptr, fw::HttpStreamComplete complete) + { + struct State + { + std::vector buffer; + std::size_t total = 0; + fw::HttpContext* ctx; + std::shared_ptr request; + fw::HttpStreamComplete complete; + std::function read; + }; + auto state = std::make_shared(); + state->buffer.resize(read_buffer_size); + state->ctx = &ctx; + state->request = std::move(request); + state->complete = std::move(complete); + state->read = [state] + { + state->request->async_read_some(net::buffer(state->buffer), + [state](boost::system::error_code ec, std::size_t size, bool done) + { + ASSERT_FALSE(ec) << ec.message(); + state->total += size; + if (!done) return state->read(); + state->ctx->set_body(std::to_string(state->total)); + state->complete(); + }); + }; + state->read(); + }); + } + + http::response send_stream(std::string body, bool chunked, + std::uint64_t buffered_limit = 0) + { + test::TempWebRoot web_root; + fw::HttpRouter router; + fw::WebsocketRouter websocket_router; + install_counting_stream(router); + http::request request{http::verb::post, "/stream", 11}; + request.body() = std::move(body); + if (chunked) request.chunked(true); else request.prepare_payload(); + request.keep_alive(false); + return test::round_trip(router, websocket_router, web_root.path, std::move(request), buffered_limit); + } +} + +TEST(HttpStreamEdgeTest, EmptyContentLengthBodyCompletesWithoutReadData) +{ + const auto response = send_stream({}, false); + EXPECT_EQ(response.result(), http::status::ok); + EXPECT_EQ(response.body(), "0"); +} + +TEST(HttpStreamEdgeTest, StreamRouteBypassesZeroBufferedBodyLimit) +{ + const auto response = send_stream(std::string(4097, 's'), false, 0); + EXPECT_EQ(response.result(), http::status::ok); + EXPECT_EQ(response.body(), "4097"); +} + +TEST(HttpStreamEdgeTest, ChunkedBodyCrossingManySmallReadsPreservesByteCount) +{ + const auto response = send_stream(std::string(103, 'c'), true, 1); + EXPECT_EQ(response.result(), http::status::ok); + EXPECT_EQ(response.body(), "103"); +} + +TEST(HttpStreamEdgeTest, CancellingRequestBodyKeepsResponseChannelUsable) +{ + test::TempWebRoot web_root; + fw::HttpRouter router; + fw::WebsocketRouter websocket_router; + router.stream("/early", http::verb::post, + [](fw::HttpContext&, std::shared_ptr, + std::shared_ptr response, fw::HttpStreamComplete) + { + response->cancel_request_body(); + fw::HttpResponseStream::ResponseHead head{http::status::forbidden, 11}; + head.content_length(6); + head.keep_alive(false); + response->async_start(std::move(head), [response](boost::system::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + static constexpr char denied[] = "denied"; + response->async_write_some(net::buffer(denied, 6), [response](boost::system::error_code write_ec) + { + ASSERT_FALSE(write_ec) << write_ec.message(); + response->async_finish([](boost::system::error_code finish_ec) + { ASSERT_FALSE(finish_ec) << finish_ec.message(); }); + }); + }); + }); + http::request request{http::verb::post, "/early", 11}; + request.body().assign(4096, 'x'); + request.prepare_payload(); + request.keep_alive(false); + const auto response = test::round_trip(router, websocket_router, web_root.path, std::move(request), 0); + EXPECT_EQ(response.result(), http::status::forbidden); + EXPECT_EQ(response.body(), "denied"); +} diff --git a/framework/tests/router_test.cpp b/framework/tests/router_test.cpp index 0ed7d58..0276a45 100644 --- a/framework/tests/router_test.cpp +++ b/framework/tests/router_test.cpp @@ -839,6 +839,30 @@ TEST(WebsocketRouterTest, HandlerOverwrite) ASSERT_EQ(state.path_received, "v2"); // Should use the second handler } +TEST(WebsocketRouterTest, DynamicRouteCapturesTrailingPathAndPrefersStaticRoute) +{ + khttpd_fw::WebsocketRouter router; + auto session = std::make_shared(); + std::shared_ptr base = session; + std::string selected; + std::string target; + + router.add_handler("/gateway/:service", [&](khttpd_fw::WebsocketContext& ctx) { + selected = "dynamic"; + target = ctx.get_path_param("service").value_or(""); + }); + router.add_handler("/gateway/health", [&](khttpd_fw::WebsocketContext&) { selected = "static"; }); + + khttpd_fw::WebsocketContext nested(base, "/gateway/orders/ws/v1"); + router.dispatch_open("/gateway/orders/ws/v1", nested); + ASSERT_EQ(selected, "dynamic"); + ASSERT_EQ(target, "orders/ws/v1"); + + khttpd_fw::WebsocketContext health(base, "/gateway/health"); + router.dispatch_open("/gateway/health", health); + ASSERT_EQ(selected, "static"); +} + TEST(WebsocketRouterTest, SpecificHandlerRegistered) { khttpd_fw::WebsocketRouter router; diff --git a/framework/tests/session_test.cpp b/framework/tests/session_test.cpp index c855fd6..c09f1c5 100644 --- a/framework/tests/session_test.cpp +++ b/framework/tests/session_test.cpp @@ -2,6 +2,7 @@ #include "framework/router/http_router.hpp" #include "framework/router/websocket_router.hpp" #include "framework/context/http_context.hpp" +#include "framework/client/http_proxy_session.hpp" #include #include @@ -50,9 +51,13 @@ namespace khttpd_fw::WebsocketRouter& websocket_router, const fs::path& web_root, http::request req, - bool skip_body = false) + bool skip_body = false, + std::function completion = {}, + std::uint64_t max_buffered_request_body_size = + khttpd_fw::HttpSession::default_max_buffered_request_body_size) { net::io_context ioc; + auto work_guard = net::make_work_guard(ioc); tcp::acceptor acceptor(ioc, {net::ip::make_address("127.0.0.1"), 0}); const auto endpoint = acceptor.local_endpoint(); const auto canonical_web_root = fs::canonical(web_root); @@ -61,7 +66,8 @@ namespace { ASSERT_FALSE(ec) << ec.message(); std::make_shared( - std::move(socket), router, websocket_router, web_root.string(), canonical_web_root)->run(); + std::move(socket), router, websocket_router, web_root.string(), canonical_web_root, + max_buffered_request_body_size)->run(); }); std::thread server_thread([&] @@ -72,17 +78,24 @@ namespace net::io_context client_ioc; tcp::socket client(client_ioc); client.connect(endpoint); - http::write(client, req); + beast::error_code io_ec; + http::write(client, req, io_ec); + EXPECT_FALSE(io_ec) << io_ec.message(); beast::flat_buffer buffer; http::response_parser parser; parser.skip(skip_body); - http::read(client, buffer, parser); + http::read(client, buffer, parser, io_ec); + EXPECT_FALSE(io_ec) << io_ec.message(); auto res = parser.release(); + for (int i = 0; completion && !completion() && i < 1000; ++i) + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + beast::error_code ignored; client.shutdown(tcp::socket::shutdown_both, ignored); client.close(ignored); + work_guard.reset(); ioc.stop(); server_thread.join(); @@ -90,6 +103,68 @@ namespace } } +TEST(HttpSessionTest, ConfigurableBufferedBodyLimitAcceptsBodyAtLimit) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + constexpr std::uint64_t limit = 1024; + router.post("/buffered", [](khttpd_fw::HttpContext& ctx) + { + ctx.set_body(std::to_string(ctx.get_request().body().size())); + }); + + http::request req{http::verb::post, "/buffered", 11}; + req.body().assign(limit, 'x'); + req.prepare_payload(); + req.keep_alive(false); + + auto res = round_trip(router, websocket_router, tree.web, std::move(req), + false, {}, limit); + EXPECT_EQ(res.result(), http::status::ok); + EXPECT_EQ(res.body(), std::to_string(limit)); +} + +TEST(HttpSessionTest, ConfigurableBufferedBodyLimitRejectsContentLengthBeforeHandler) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + constexpr std::uint64_t limit = 1024; + bool handled = false; + router.post("/buffered", [&handled](khttpd_fw::HttpContext&) { handled = true; }); + + http::request req{http::verb::post, "/buffered", 11}; + req.body().assign(limit + 1, 'x'); + req.prepare_payload(); + req.keep_alive(false); + + auto res = round_trip(router, websocket_router, tree.web, std::move(req), + false, {}, limit); + EXPECT_EQ(res.result(), http::status::payload_too_large); + EXPECT_FALSE(handled); +} + +TEST(HttpSessionTest, ConfigurableBufferedBodyLimitRejectsChunkedBodyWhileReading) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + constexpr std::uint64_t limit = 1024; + bool handled = false; + router.post("/buffered", [&handled](khttpd_fw::HttpContext&) { handled = true; }); + + http::request req{http::verb::post, "/buffered", 11}; + req.body().assign(limit + 1, 'x'); + req.chunked(true); + req.keep_alive(false); + + auto res = round_trip(router, websocket_router, tree.web, std::move(req), + false, {}, limit); + EXPECT_EQ(res.result(), http::status::payload_too_large); + EXPECT_FALSE(handled); +} + TEST(HttpSessionTest, StaticFileRejectsSiblingPrefixTraversal) { TempStaticTree tree; @@ -124,6 +199,69 @@ TEST(HttpSessionTest, StaticHeadReturnsHeadersWithoutBody) EXPECT_EQ(res[http::field::content_length], "5"); } +TEST(HttpSessionTest, DynamicHeadUsesGetMetadataWithoutSendingBody) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + router.get("/resource", [](khttpd_fw::HttpContext& ctx) { ctx.set_body("generated-body"); }); + http::request req{http::verb::head, "/resource", 11}; + req.keep_alive(false); + auto res = round_trip(router, websocket_router, tree.web, std::move(req), true); + EXPECT_EQ(res.result(), http::status::ok); + EXPECT_EQ(res[http::field::content_length], "14"); +} + +TEST(HttpSessionTest, AsyncInterceptorSeesTransportPeerAndCanDenyRequest) +{ + struct RemoteAuth final : khttpd_fw::Interceptor + { + std::optional seen_peer; + void async_handle_request(khttpd_fw::HttpContext& ctx, RequestCompletion complete) override + { + seen_peer = ctx.peer_endpoint(); + ctx.set_status(http::status::unauthorized); + ctx.set_body("remote auth denied"); + std::thread([complete = std::move(complete)] { complete(khttpd_fw::InterceptorResult::Stop); }).detach(); + } + }; + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + auto auth = std::make_shared(); + router.add_interceptor(auth); + router.get("/private", [](khttpd_fw::HttpContext& ctx) { ctx.set_body("should not run"); }); + http::request req{http::verb::get, "/private", 11}; + req.keep_alive(false); + auto res = round_trip(router, websocket_router, tree.web, std::move(req)); + EXPECT_EQ(res.result(), http::status::unauthorized); + EXPECT_EQ(res.body(), "remote auth denied"); + ASSERT_TRUE(auth->seen_peer); + EXPECT_TRUE(auth->seen_peer->address().is_loopback()); + EXPECT_NE(auth->seen_peer->port(), 0); +} + +TEST(HttpSessionTest, AsyncRouteCompletesResponseFromAnotherThread) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + router.async_route("/remote", http::verb::get, + [](khttpd_fw::HttpContext& ctx, khttpd_fw::HttpAsyncComplete complete) + { + std::thread([&ctx, complete = std::move(complete)]() mutable + { + ctx.set_body("async result"); + complete(); + }).detach(); + }); + http::request req{http::verb::get, "/remote", 11}; + req.keep_alive(false); + auto res = round_trip(router, websocket_router, tree.web, std::move(req)); + EXPECT_EQ(res.result(), http::status::ok); + EXPECT_EQ(res.body(), "async result"); +} + TEST(HttpSessionTest, ChunkedResponseCompletesWithSingleIoThread) { TempStaticTree tree; @@ -153,6 +291,150 @@ TEST(HttpSessionTest, ChunkedResponseCompletesWithSingleIoThread) EXPECT_EQ(res.body(), "onetwo"); } +TEST(HttpSessionTest, StreamsContentLengthRequestBodyInFixedBuffers) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + // Deliberately exceeds the 16 MiB buffered-route limit. A stream route must + // consume it successfully without accumulating the body in HttpSession. + constexpr std::size_t body_size = 20 * 1024 * 1024 + 17; + + router.stream("/upload", http::verb::post, + [](khttpd_fw::HttpContext& ctx, std::shared_ptr stream, + std::shared_ptr, + khttpd_fw::HttpStreamComplete complete) + { + struct State { + std::array buffer{}; + std::size_t total = 0; + std::shared_ptr stream; + khttpd_fw::HttpContext* ctx; + khttpd_fw::HttpStreamComplete complete; + std::function next; + }; + auto state = std::make_shared(); + state->stream = std::move(stream); state->ctx = &ctx; state->complete = std::move(complete); + state->next = [state] + { + state->stream->async_read_some(net::buffer(state->buffer), [state](beast::error_code ec, std::size_t n, bool done) + { + ASSERT_FALSE(ec) << ec.message(); + state->total += n; + if (!done) return state->next(); + state->ctx->set_body(std::to_string(state->total)); + state->complete(); + }); + }; + state->next(); + }); + + http::request req{http::verb::post, "/upload", 11}; + req.body().assign(body_size, 'x'); + req.prepare_payload(); + req.keep_alive(false); + auto res = round_trip(router, websocket_router, tree.web, std::move(req)); + EXPECT_EQ(res.result(), http::status::ok); + EXPECT_EQ(res.body(), std::to_string(body_size)); +} + +TEST(HttpSessionTest, StreamsChunkedRequestBody) +{ + TempStaticTree tree; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + router.stream("/chunked", http::verb::post, + [](khttpd_fw::HttpContext& ctx, std::shared_ptr stream, + std::shared_ptr, + khttpd_fw::HttpStreamComplete complete) + { + auto buffer = std::make_shared>(); + auto total = std::make_shared(0); + auto next = std::make_shared>(); + *next = [&ctx, stream, complete, buffer, total, next] + { + stream->async_read_some(net::buffer(*buffer), [&ctx, stream, complete, buffer, total, next] + (beast::error_code ec, std::size_t n, bool done) + { + ASSERT_FALSE(ec) << ec.message(); *total += n; + if (!done) return (*next)(); + ctx.set_body(std::to_string(*total)); complete(); + }); + }; + (*next)(); + }); + http::request req{http::verb::post, "/chunked", 11}; + const std::string chunked_body = "chunked-body-with-several-pieces"; + req.body() = chunked_body; + req.chunked(true); + req.keep_alive(false); + auto res = round_trip(router, websocket_router, tree.web, std::move(req)); + EXPECT_EQ(res.result(), http::status::ok); + EXPECT_EQ(res.body(), std::to_string(chunked_body.size())); +} + +TEST(HttpSessionTest, ProxySessionStreamsUploadDownloadAndRangeHeaders) +{ + TempStaticTree tree; + net::io_context upstream_ioc; + tcp::acceptor upstream_acceptor(upstream_ioc, {net::ip::make_address("127.0.0.1"), 0}); + const auto upstream_endpoint = upstream_acceptor.local_endpoint(); + upstream_acceptor.async_accept([&](beast::error_code accept_ec, tcp::socket socket) + { + if (accept_ec) return; + beast::flat_buffer buffer; + http::request_parser parser; parser.body_limit(8 * 1024 * 1024); + beast::error_code ec; http::read(socket, buffer, parser, ec); + if (ec) return; + auto request = parser.release(); + http::response response{http::status::partial_content, 11}; + response.set(http::field::content_range, "bytes 0-9/100"); + response.set(http::field::accept_ranges, "bytes"); + response.body() = std::move(request.body()); response.keep_alive(false); response.prepare_payload(); + http::write(socket, response, ec); + }); + std::thread upstream_thread([&] { upstream_ioc.run(); }); + + net::io_context proxy_ioc; + auto proxy_guard = net::make_work_guard(proxy_ioc); + std::thread proxy_thread([&] { proxy_ioc.run(); }); + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + std::shared_ptr active_proxy; + std::atomic proxy_handler_called{false}; + std::atomic proxy_error{0}; + std::atomic proxy_completed{false}; + router.stream("/proxy", http::verb::post, + [&](khttpd_fw::HttpContext& ctx, std::shared_ptr request_stream, + std::shared_ptr response_stream, khttpd_fw::HttpStreamComplete) + { + proxy_handler_called = true; + khttpd::framework::client::HttpClientStream::RequestHead head{ctx.method(), "/", ctx.get_request().version()}; + for (const auto& field : ctx.get_request()) head.insert(field.name_string(), field.value()); + active_proxy = std::make_shared( + proxy_ioc, std::move(request_stream), std::move(response_stream), 32 * 1024); + const auto url = "http://127.0.0.1:" + std::to_string(upstream_endpoint.port()) + "/echo"; + active_proxy->start(url, std::move(head), [&](beast::error_code ec) + { proxy_error = ec.value(); proxy_completed = true; }); + }); + + const std::string payload(2 * 1024 * 1024 + 9, 'p'); + http::request req{http::verb::post, "/proxy", 11}; + req.set(http::field::range, "bytes=0-9"); req.body() = payload; req.prepare_payload(); req.keep_alive(false); + auto res = round_trip(router, websocket_router, tree.web, std::move(req), false, + [&] { return proxy_completed.load(); }); + + proxy_guard.reset(); proxy_ioc.stop(); proxy_thread.join(); + beast::error_code ignored; upstream_acceptor.close(ignored); upstream_ioc.stop(); upstream_thread.join(); + EXPECT_TRUE(proxy_handler_called); + EXPECT_TRUE(proxy_completed); + EXPECT_EQ(proxy_error, 0); + EXPECT_EQ(res.result(), http::status::partial_content); + EXPECT_EQ(res[http::field::content_range], "bytes 0-9/100"); + EXPECT_EQ(res[http::field::accept_ranges], "bytes"); + EXPECT_EQ(res.body(), payload); +} + TEST(HttpSessionTest, WebSocketDrainsQueuedMessages) { TempStaticTree tree; @@ -218,3 +500,114 @@ TEST(HttpSessionTest, WebSocketDrainsQueuedMessages) EXPECT_EQ(first, "first"); EXPECT_EQ(second, "second"); } + +TEST(HttpSessionTest, WebSocketUpgradeRunsHttpAuthenticationFirst) +{ + struct DenyUpgrade final : khttpd_fw::Interceptor + { + khttpd_fw::InterceptorResult handle_request(khttpd_fw::HttpContext& ctx) override + { + ctx.set_status(http::status::unauthorized); + ctx.set_body("websocket auth required"); + return khttpd_fw::InterceptorResult::Stop; + } + }; + TempStaticTree tree; + net::io_context server_ioc; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + router.add_interceptor(std::make_shared()); + bool opened = false; + websocket_router.add_handler("/private-ws", [&](khttpd_fw::WebsocketContext&) { opened = true; }); + tcp::acceptor acceptor(server_ioc, {net::ip::address_v4::loopback(), 0}); + const auto endpoint = acceptor.local_endpoint(); + acceptor.async_accept([&](beast::error_code ec, tcp::socket socket) + { + ASSERT_FALSE(ec); + std::make_shared(std::move(socket), router, websocket_router, + tree.web.string(), fs::canonical(tree.web))->run(); + }); + std::thread server_thread([&] { server_ioc.run(); }); + net::io_context client_ioc; + websocket::stream client(client_ioc); + client.next_layer().connect(endpoint); + websocket::response_type response; + beast::error_code ec; + client.handshake(response, "127.0.0.1", "/private-ws", ec); + EXPECT_TRUE(ec); + EXPECT_EQ(response.result(), http::status::unauthorized); + EXPECT_FALSE(opened); + beast::error_code ignored; + client.next_layer().close(ignored); + server_ioc.stop(); + server_thread.join(); +} + +TEST(HttpSessionTest, WebSocketHandshakePreservesRoutingAndRequestMetadata) +{ + TempStaticTree tree; + net::io_context server_ioc; + khttpd_fw::HttpRouter router; + khttpd_fw::WebsocketRouter websocket_router; + std::mutex captured_mutex; + std::string target_param; + khttpd_fw::WebsocketHandshakeRequest captured; + std::atomic opened{false}; + std::atomic closed{false}; + + websocket_router.add_handler("/gateway/:target", [&](khttpd_fw::WebsocketContext& ctx) + { + std::lock_guard lock(captured_mutex); + target_param = ctx.get_path_param("target").value_or(""); + captured = ctx.handshake(); + opened = true; + }, nullptr, [&](khttpd_fw::WebsocketContext&) { closed = true; }); + + tcp::acceptor acceptor(server_ioc, {net::ip::make_address("127.0.0.1"), 0}); + const auto endpoint = acceptor.local_endpoint(); + const auto canonical_web_root = fs::canonical(tree.web); + acceptor.async_accept([&](beast::error_code ec, tcp::socket socket) + { + ASSERT_FALSE(ec) << ec.message(); + std::make_shared( + std::move(socket), router, websocket_router, tree.web.string(), canonical_web_root)->run(); + }); + std::thread server_thread([&] { server_ioc.run(); }); + + net::io_context client_ioc; + websocket::stream client(client_ioc); + client.set_option(websocket::stream_base::decorator([](websocket::request_type& request) + { + request.set(http::field::authorization, "Bearer test-token"); + request.insert(http::field::cookie, "first=1"); + request.insert(http::field::cookie, "second=2"); + request.set(http::field::origin, "https://gateway-client.example"); + request.set(http::field::sec_websocket_protocol, "chat, telemetry"); + })); + client.next_layer().connect(endpoint); + client.handshake("127.0.0.1", "/gateway/orders/ws?tenant=acme&trace=abc"); + + for (int i = 0; i < 100 && !opened; ++i) + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + beast::error_code ignored; + client.close(websocket::close_code::normal, ignored); + for (int i = 0; i < 100 && !closed; ++i) + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + server_ioc.stop(); + server_thread.join(); + + ASSERT_TRUE(opened); + ASSERT_TRUE(closed); + std::lock_guard lock(captured_mutex); + EXPECT_EQ(target_param, "orders/ws"); + EXPECT_EQ(captured.target, "/gateway/orders/ws?tenant=acme&trace=abc"); + EXPECT_EQ(captured.path, "/gateway/orders/ws"); + EXPECT_EQ(captured.query_params.at("tenant"), "acme"); + EXPECT_EQ(captured.query_params.at("trace"), "abc"); + EXPECT_EQ(captured.subprotocols, (std::vector{"chat", "telemetry"})); + + std::vector cookies; + for (const auto& [name, value] : captured.headers) + if (name == "Cookie") cookies.push_back(value); + EXPECT_EQ(cookies, (std::vector{"first=1", "second=2"})); +} diff --git a/framework/tests/websocket_client_handshake_test.cpp b/framework/tests/websocket_client_handshake_test.cpp new file mode 100644 index 0000000..277be7f --- /dev/null +++ b/framework/tests/websocket_client_handshake_test.cpp @@ -0,0 +1,64 @@ +#include "framework/client/websocket_client.hpp" + +#include +#include +#include +#include + +namespace fw = khttpd::framework; +namespace client = khttpd::framework::client; +namespace beast = boost::beast; +namespace http = beast::http; +namespace websocket = beast::websocket; +namespace net = boost::asio; +using tcp = net::ip::tcp; + +TEST(WebsocketClientHandshakeTest, PreservesEncodedTargetAndNegotiatesSubprotocol) +{ + net::io_context server_ioc; + tcp::acceptor acceptor(server_ioc, {net::ip::address_v4::loopback(), 0}); + const auto port = acceptor.local_endpoint().port(); + std::string received_target; + std::string requested_protocols; + std::thread server([&] + { + tcp::socket socket(server_ioc); + acceptor.accept(socket); + beast::flat_buffer request_buffer; + http::request request; + http::read(socket, request_buffer, request); + received_target = std::string(request.target()); + requested_protocols = std::string(request[http::field::sec_websocket_protocol]); + + websocket::stream ws(std::move(socket)); + ws.set_option(websocket::stream_base::decorator([](websocket::response_type& response) + { + response.set(http::field::sec_websocket_protocol, "chat.v2"); + })); + ws.accept(request); + beast::error_code ignored; + ws.next_layer().shutdown(tcp::socket::shutdown_both, ignored); + ws.next_layer().close(ignored); + }); + + net::io_context client_ioc; + auto ws = std::make_shared(client_ioc); + ws->set_subprotocols({"chat.v1", "chat.v2"}); + bool connected = false; + std::string selected; + ws->connect("ws://127.0.0.1:" + std::to_string(port) + + "/socket%2Fv1?tenant=acme&trace=a%2Fb", + [&](beast::error_code ec) + { + ASSERT_FALSE(ec) << ec.message(); + connected = true; + selected = ws->negotiated_subprotocol(); + }); + client_ioc.run(); + server.join(); + + EXPECT_TRUE(connected); + EXPECT_EQ(received_target, "/socket%2Fv1?tenant=acme&trace=a%2Fb"); + EXPECT_EQ(requested_protocols, "chat.v1, chat.v2"); + EXPECT_EQ(selected, "chat.v2"); +} diff --git a/framework/tests/websocket_router_dynamic_test.cpp b/framework/tests/websocket_router_dynamic_test.cpp new file mode 100644 index 0000000..cb34c51 --- /dev/null +++ b/framework/tests/websocket_router_dynamic_test.cpp @@ -0,0 +1,91 @@ +#include "framework/router/websocket_router.hpp" + +#include +#include +#include + +namespace fw = khttpd::framework; + +namespace +{ + fw::WebsocketContext context_for(std::string path) + { + return fw::WebsocketContext(std::weak_ptr{}, std::move(path)); + } +} + +TEST(WebsocketRouterDynamicTest, MultipleParametersKeepFinalParameterGreedy) +{ + fw::WebsocketRouter router; + std::string tenant; + std::string target; + router.add_handler("/gateway/:tenant/:target", [&](fw::WebsocketContext& ctx) + { + tenant = ctx.get_path_param("tenant").value_or(""); + target = ctx.get_path_param("target").value_or(""); + }); + auto ctx = context_for("/gateway/acme/orders/ws/v2"); + router.dispatch_open(ctx.path, ctx); + EXPECT_EQ(tenant, "acme"); + EXPECT_EQ(target, "orders/ws/v2"); +} + +TEST(WebsocketRouterDynamicTest, LiteralRegexCharactersAreMatchedLiterally) +{ + fw::WebsocketRouter router; + bool called = false; + router.add_handler("/socket.v1/:target", [&](fw::WebsocketContext&) { called = true; }); + auto wrong = context_for("/socketXv1/orders"); + router.dispatch_open(wrong.path, wrong); + EXPECT_FALSE(called); + auto exact = context_for("/socket.v1/orders"); + router.dispatch_open(exact.path, exact); + EXPECT_TRUE(called); +} + +TEST(WebsocketRouterDynamicTest, StaticRouteWinsRegardlessOfRegistrationOrder) +{ + fw::WebsocketRouter router; + std::string selected; + router.add_handler("/gateway/:target", [&](fw::WebsocketContext&) { selected = "dynamic"; }); + router.add_handler("/gateway/health", [&](fw::WebsocketContext&) { selected = "static"; }); + auto ctx = context_for("/gateway/health"); + router.dispatch_open(ctx.path, ctx); + EXPECT_EQ(selected, "static"); +} + +TEST(WebsocketRouterDynamicTest, HandlerCanReplaceItselfDuringDispatch) +{ + fw::WebsocketRouter router; + int version = 0; + router.add_handler("/hot", [&](fw::WebsocketContext&) + { + version = 1; + router.add_handler("/hot", [&](fw::WebsocketContext&) { version = 2; }); + }); + auto first = context_for("/hot"); + router.dispatch_open(first.path, first); + EXPECT_EQ(version, 1); + auto second = context_for("/hot"); + router.dispatch_open(second.path, second); + EXPECT_EQ(version, 2); +} + +TEST(WebsocketRouterDynamicTest, ConcurrentDispatchAndHotUpdateRemainSafe) +{ + fw::WebsocketRouter router; + std::atomic calls{0}; + router.add_handler("/gateway/:target", [&](fw::WebsocketContext&) { ++calls; }); + std::thread dispatcher([&] + { + for (int i = 0; i < 500; ++i) + { + auto ctx = context_for("/gateway/service/ws"); + router.dispatch_open(ctx.path, ctx); + } + }); + for (int i = 0; i < 100; ++i) + router.add_handler("/gateway/:target", [&](fw::WebsocketContext&) { ++calls; }); + dispatcher.join(); + EXPECT_EQ(calls, 500); +} diff --git a/framework/websocket/websocket_session.cpp b/framework/websocket/websocket_session.cpp index 77f1bb5..510aaac 100644 --- a/framework/websocket/websocket_session.cpp +++ b/framework/websocket/websocket_session.cpp @@ -39,6 +39,30 @@ namespace khttpd::framework } spdlog::debug("WebSocket handshake successful for path: {}", initial_path_); + std::weak_ptr weak = shared_from_this(); + ws_.control_callback([weak](ws::frame_type type, beast::string_view payload) + { + const auto self = weak.lock(); + if (!self) return; + WebsocketFrameType frame_type; + switch (type) + { + case ws::frame_type::ping: frame_type = WebsocketFrameType::ping; break; + case ws::frame_type::pong: frame_type = WebsocketFrameType::pong; break; + case ws::frame_type::close: frame_type = WebsocketFrameType::close; break; + } + WebsocketContext ctx(self, self->initial_path_); + ctx.frame.type = frame_type; + ctx.frame.payload = std::string(payload); + if (type == ws::frame_type::close) + { + const auto reason = self->ws_.reason(); + ctx.frame.close_code = static_cast(reason.code); + ctx.frame.close_reason = reason.reason.c_str(); + } + self->websocket_router_.dispatch_message(self->initial_path_, ctx); + }); + WebsocketContext open_ctx(shared_from_this(), initial_path_); { std::unique_lock lock{m_sessions_mutex}; @@ -86,16 +110,20 @@ namespace khttpd::framework } void WebsocketSession::send_message(const std::string& msg, bool is_text_msg) + { + send_frame({is_text_msg ? WebsocketFrameType::text : WebsocketFrameType::binary, msg}); + } + + void WebsocketSession::send_frame(WebsocketFrame frame) { auto self = shared_from_this(); - auto ss = std::make_shared(msg); - net::post(ws_.get_executor(), [self, ss, is_text_msg]() + net::post(ws_.get_executor(), [self, frame = std::move(frame)]() mutable { if (self->closed_) { return; } - self->write_queue_.emplace(ss, is_text_msg); + self->write_queue_.push(std::move(frame)); if (!self->writing_) { self->writing_ = true; @@ -132,12 +160,27 @@ namespace khttpd::framework return; } - auto item = std::move(write_queue_.front()); + auto frame = std::move(write_queue_.front()); write_queue_.pop(); - auto& ss = item.first; - auto is_text_msg = item.second; + if (frame.type == WebsocketFrameType::close) + { + closed_ = true; + ws_.async_close({static_cast(frame.close_code), frame.close_reason}, + [self = shared_from_this()](beast::error_code ec) { self->on_write(ec, 0); }); + return; + } + if (frame.type == WebsocketFrameType::ping || frame.type == WebsocketFrameType::pong) + { + const ws::ping_data payload(frame.payload); + if (frame.type == WebsocketFrameType::ping) + ws_.async_ping(payload, [self = shared_from_this()](beast::error_code ec) { self->on_write(ec, 0); }); + else + ws_.async_pong(payload, [self = shared_from_this()](beast::error_code ec) { self->on_write(ec, 0); }); + return; + } + auto ss = std::make_shared(std::move(frame.payload)); - ws_.text(is_text_msg); + ws_.text(frame.type == WebsocketFrameType::text); if (ss->length() < auto_fragment_threshold_) { @@ -254,6 +297,10 @@ namespace khttpd::framework else { WebsocketContext close_ctx(shared_from_this(), initial_path_, ec); + const auto reason = ws_.reason(); + close_ctx.frame.type = WebsocketFrameType::close; + close_ctx.frame.close_code = static_cast(reason.code); + close_ctx.frame.close_reason = reason.reason.c_str(); websocket_router_.dispatch_close(initial_path_, close_ctx); } diff --git a/framework/websocket/websocket_session.hpp b/framework/websocket/websocket_session.hpp index 91f979c..c6f9bcb 100644 --- a/framework/websocket/websocket_session.hpp +++ b/framework/websocket/websocket_session.hpp @@ -28,6 +28,9 @@ namespace khttpd::framework void run_handshake(http::request> req); virtual void send_message(const std::string& msg, bool is_text); + virtual void send_frame(WebsocketFrame frame); + + const WebsocketHandshakeRequest& handshake() const { return handshake_; } static bool send_message(const std::string& id, const std::string& msg, bool is_text); static size_t send_message(const std::vector& ids, const std::string& msg, bool is_text); @@ -40,6 +43,7 @@ namespace khttpd::framework beast::flat_buffer buffer_; WebsocketRouter& websocket_router_; std::string initial_path_; + WebsocketHandshakeRequest handshake_; static std::mutex m_gen_mutex; static boost::uuids::random_generator gen; static std::mutex m_sessions_mutex; @@ -52,7 +56,7 @@ namespace khttpd::framework static constexpr size_t const auto_fragment_threshold_ = fragment_size_ * 2; // Write queue to serialize concurrent async_write calls - std::queue, bool>> write_queue_; + std::queue write_queue_; bool writing_ = false; bool closed_ = false; bool close_pending_ = false; @@ -69,6 +73,43 @@ namespace khttpd::framework template void WebsocketSession::run_handshake(http::request> req) { + handshake_.target = std::string(req.target()); + const auto query_pos = handshake_.target.find('?'); + handshake_.path = handshake_.target.substr(0, query_pos); + initial_path_ = handshake_.path.empty() ? "/" : handshake_.path; + for (const auto& field : req) + { + const std::string name(field.name_string()); + const std::string value(field.value()); + handshake_.headers.emplace_back(name, value); + if (field.name() == http::field::sec_websocket_protocol) + { + size_t begin = 0; + while (begin < value.size()) + { + const size_t end = value.find(',', begin); + const auto token = value.substr(begin, end == std::string::npos ? end : end - begin); + const auto first = token.find_first_not_of(" \t"); + if (first != std::string::npos) handshake_.subprotocols.push_back(token.substr(first, token.find_last_not_of(" \t") - first + 1)); + if (end == std::string::npos) break; + begin = end + 1; + } + } + } + if (query_pos != std::string::npos) + { + std::string query = handshake_.target.substr(query_pos + 1); + size_t begin = 0; + while (begin <= query.size()) + { + const size_t end = query.find('&', begin); + const std::string pair = query.substr(begin, end == std::string::npos ? end : end - begin); + const size_t equal = pair.find('='); + handshake_.query_params.emplace(pair.substr(0, equal), equal == std::string::npos ? "" : pair.substr(equal + 1)); + if (end == std::string::npos) break; + begin = end + 1; + } + } ws_.async_accept(req, beast::bind_front_handler(&WebsocketSession::on_handshake, shared_from_this())); }