From 8d778abbdc9c066686856f4a4f9d39b2bf48ad98 Mon Sep 17 00:00:00 2001 From: NoahLan <6995syu@163.com> Date: Fri, 24 Jul 2026 14:34:24 +0800 Subject: [PATCH] v2 --- README.md | 10 +- cmd/client/main.go | 2 +- cmd/server/main.go | 6 +- docs/kb/development/IMPLEMENTATION.md | 8 +- .../ANSWER_TO_USER_QUESTION.md | 4 +- .../CODEC_AND_PROTOCOL_INTEGRATION.md | 4 +- .../SERVER_WRITE_DATA_SUMMARY.md | 4 +- .../WRITE_DATA_TO_CLIENT.md | 4 +- docs/kb/user-guide/CLIENT_POOL.md | 2 +- docs/kb/user-guide/DOMAIN_SETUP.md | 4 +- docs/kb/user-guide/MULTI_REPO_SETUP.md | 26 +- .../user-guide/PROTOCOL_FRAME_EXPLANATION.md | 12 +- docs/kb/user-guide/ROUTER_EXAMPLES.md | 24 +- examples/client_pool/client.go | 2 +- examples/client_pool/server.go | 2 +- examples/client_request/main.go | 2 +- examples/interceptor/client.go | 3 +- examples/interceptor/server.go | 10 +- examples/metrics_health/main.go | 4 +- examples/middleware/client.go | 3 +- examples/middleware/server.go | 3 +- examples/preset/main.go | 2 +- examples/protocol_version/client.go | 3 +- examples/protocol_version/server.go | 5 +- examples/router_frame_data/main.go | 2 +- examples/session_file/client.go | 3 +- examples/session_file/server.go | 27 +- examples/udp_echo/client.go | 3 +- examples/udp_echo/server.go | 3 +- examples/write_any/main.go | 2 +- examples/ws_echo/client.go | 3 +- examples/ws_echo/server.go | 3 +- go.mod | 4 +- internal/client/application_protocol.go | 56 +++ internal/client/pool.go | 2 +- internal/client/pool_test.go | 2 +- internal/client/serial_client.go | 3 +- internal/client/tcp_client.go | 238 +++++++++--- internal/client/tcp_client_more_test.go | 4 +- internal/client/tcp_client_test.go | 2 +- internal/client/udp_client.go | 2 +- internal/client/unix_client.go | 3 +- internal/client/websocket_client.go | 145 ++++++-- internal/codec/binary.go | 3 +- internal/codec/binary_test.go | 1 - internal/codec/json.go | 3 +- internal/codec/json_test.go | 1 - internal/codec/msgpack.go | 3 +- internal/codec/msgpack_test.go | 1 - internal/codec/plain.go | 2 +- internal/codec/protobuf.go | 2 +- internal/codec/protobuf_test.go | 1 - internal/codec/registry.go | 2 +- internal/connection/connection.go | 8 +- internal/connection/connection_test.go | 2 +- internal/connection/group_strategy.go | 1 - internal/connection/serial_connection.go | 1 - internal/connection/sharded_manager.go | 6 +- internal/connection/sharded_manager_test.go | 2 +- internal/connection/udp_connection.go | 3 +- internal/connection/unix_connection.go | 1 - internal/connection/websocket_connection.go | 4 +- internal/interceptor/builtin/validation.go | 5 +- .../interceptor/builtin/validation_test.go | 12 +- internal/interceptor/chain.go | 5 +- internal/interceptor/chain_test.go | 10 +- internal/middleware/builtin/auth.go | 9 +- internal/middleware/builtin/logging.go | 10 +- internal/middleware/builtin/ratelimit.go | 5 +- internal/middleware/builtin/recovery.go | 5 +- internal/middleware/chain.go | 5 +- internal/middleware/chain_test.go | 8 +- internal/plugin/manager.go | 5 +- internal/protocol/frame_header.go | 3 +- internal/protocol/manager.go | 3 +- internal/protocol/manager_test.go | 4 +- .../protocol/nnet/incremental_decoder_test.go | 2 +- internal/protocol/nnet/protocol.go | 40 +-- internal/protocol/nnet/protocol_test.go | 8 +- internal/protocol/nnet/version_identifier.go | 3 +- .../protocol/nnet/version_identifier_test.go | 21 +- internal/protocol/version/identifier.go | 3 +- internal/request/frame_header.go | 5 +- internal/request/frame_header_test.go | 1 - internal/request/internal.go | 3 +- internal/request/request.go | 4 +- internal/request/request_test.go | 5 +- internal/response/frame_header.go | 5 +- internal/response/frame_header_test.go | 1 - internal/response/header_converter.go | 13 +- internal/response/response.go | 8 +- internal/response/response_test.go | 9 +- internal/router/group.go | 4 +- internal/router/group_test.go | 21 +- internal/router/matcher/composite_test.go | 4 +- internal/router/matcher/frame_data.go | 4 +- .../router/matcher/frame_data_ops_test.go | 10 +- internal/router/matcher/frame_data_test.go | 10 +- internal/router/matcher/frame_header.go | 4 +- internal/router/matcher/frame_header_test.go | 10 +- internal/router/matcher/string_test.go | 10 +- internal/router/route.go | 4 +- internal/router/route_test.go | 23 +- internal/router/router.go | 8 +- internal/router/router_test.go | 10 +- internal/server/codec_protocol_init.go | 13 +- internal/server/codec_protocol_init_test.go | 4 +- internal/server/codec_resolver_test.go | 12 +- internal/server/connection_adapter.go | 5 +- internal/server/connection_data.go | 13 +- internal/server/gnet_server.go | 40 ++- internal/server/helpers.go | 19 +- internal/server/helpers_test.go | 48 +++ internal/server/message_handler.go | 27 +- internal/server/pool.go | 1 - internal/server/protocol_header_parser.go | 6 +- .../server/protocol_header_parser_test.go | 8 +- internal/server/request_body_parser.go | 10 +- internal/server/request_body_parser_test.go | 10 +- internal/server/serial_server.go | 20 +- internal/server/server.go | 38 +- internal/server/server_test.go | 10 +- internal/server/udp_server.go | 2 +- internal/server/unified_event_handler.go | 24 +- internal/server/unix_server.go | 2 +- internal/server/unpacker_manager.go | 5 +- internal/server/unpacker_manager_test.go | 10 +- internal/server/websocket_server.go | 338 ++++++++++++++++-- internal/session/session.go | 2 +- internal/session/session_test.go | 13 +- internal/session/storage/file.go | 5 +- internal/session/storage/memory.go | 4 +- internal/session/storage/memory_more_test.go | 2 - internal/session/storage/memory_test.go | 1 - internal/session/storage/redis.go | 5 +- internal/testutil/mock_connection.go | 2 +- internal/testutil/test_helpers.go | 9 +- internal/unpacker/delimiter.go | 104 +----- internal/unpacker/fixed_length.go | 56 +-- internal/unpacker/fixed_length_test.go | 5 +- internal/unpacker/frame_header.go | 65 +--- internal/unpacker/frame_header_test.go | 3 +- internal/unpacker/length_field.go | 90 ++--- internal/unpacker/length_field_test.go | 41 ++- pkg/client/errors.go | 1 - pkg/codec/codec.go | 1 - pkg/codec/errors.go | 1 - pkg/codec/resolver.go | 11 +- pkg/config/config.go | 2 +- pkg/context/context.go | 4 +- pkg/context/context_test.go | 4 +- pkg/health/errors.go | 1 - pkg/health/health.go | 1 - pkg/interceptor/errors.go | 1 - pkg/interceptor/interceptor.go | 3 +- pkg/lifecycle/lifecycle.go | 1 - pkg/metrics/metrics.go | 1 - pkg/middleware/middleware.go | 3 +- pkg/nnet/client.go | 4 +- pkg/nnet/context.go | 4 +- pkg/nnet/health.go | 3 +- pkg/nnet/lifecycle.go | 3 +- pkg/nnet/metrics.go | 2 +- pkg/nnet/preset.go | 4 +- pkg/nnet/server.go | 12 +- pkg/plugin/plugin.go | 3 +- pkg/protocol/errors.go | 1 - pkg/request/request.go | 3 +- pkg/response/response.go | 3 +- pkg/router/custom_matcher.go | 2 +- pkg/router/handler.go | 3 +- pkg/router/matcher.go | 4 +- pkg/router/router.go | 2 +- pkg/session/session.go | 1 - pkg/unpacker/errors.go | 1 - pkg/unpacker/unpacker.go | 9 +- test/connection_test.go | 4 +- test/context_test.go | 8 +- test/integration/basic_test.go | 5 +- test/integration/codec_resolver_test.go | 5 +- test/integration/codec_test.go | 5 +- test/integration/concurrent_test.go | 8 +- test/integration/connection_mgr_test.go | 13 +- test/integration/error_handling_test.go | 3 +- test/integration/health_test.go | 3 +- test/integration/helper.go | 25 +- test/integration/integration_test.go | 2 +- test/integration/interceptor_test.go | 10 +- test/integration/large_message_test.go | 32 +- test/integration/lifecycle_test.go | 3 +- test/integration/metrics_test.go | 3 +- test/integration/middleware_test.go | 3 +- test/integration/plugin_test.go | 19 +- test/integration/preset_test.go | 2 +- test/integration/protocol_header_test.go | 47 +-- test/integration/protocol_test.go | 21 +- test/integration/route_types_test.go | 7 +- test/integration/router_test.go | 15 +- test/integration/session_test.go | 7 +- test/integration/simple_test.go | 7 +- test/integration/stress_test.go | 73 ++-- test/integration/timeout_test.go | 5 +- test/integration/udp_test.go | 3 +- test/integration/websocket_test.go | 24 +- test/router_test.go | 12 +- 205 files changed, 1474 insertions(+), 1116 deletions(-) create mode 100644 internal/client/application_protocol.go create mode 100644 internal/server/helpers_test.go diff --git a/README.md b/README.md index f180ed8..7c5d983 100644 --- a/README.md +++ b/README.md @@ -18,7 +18,7 @@ ## 安装 ```bash -go get github.com/noahlann/nnet +go get git.noahlan.cn/noahlan/nnet/v2 ``` ## 快速开始 @@ -29,7 +29,7 @@ go get github.com/noahlann/nnet package main import ( - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -59,7 +59,7 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "time" "fmt" ) @@ -132,7 +132,7 @@ nnet/ ## 链接 -- **模块路径**: `github.com/noahlann/nnet` -- **GitHub**: https://github.com/noahlann/nnet +- **模块路径**: `git.noahlan.cn/noahlan/nnet/v2` +- **GitHub**: https://git.noahlan.cn/noahlan/nnet - **Gitee**: https://gitee.com/noahlann/nnet (镜像) diff --git a/cmd/client/main.go b/cmd/client/main.go index a0274db..95ea2a0 100644 --- a/cmd/client/main.go +++ b/cmd/client/main.go @@ -6,7 +6,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { diff --git a/cmd/server/main.go b/cmd/server/main.go index 24836f6..e323ad1 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -8,13 +8,13 @@ import ( "syscall" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { var ( - addr = flag.String("addr", "tcp://:8080", "Server address") - protocol = flag.String("protocol", "", "Transport protocol (tcp, udp, websocket, unix, serial). If empty, auto-detect from addr") + addr = flag.String("addr", "tcp://:8080", "Server address") + protocol = flag.String("protocol", "", "Transport protocol (tcp, udp, websocket, unix, serial). If empty, auto-detect from addr") multicore = flag.Bool("multicore", true, "Enable multicore") readBuf = flag.Int("read-buf", 4096, "Read buffer size") writeBuf = flag.Int("write-buf", 4096, "Write buffer size") diff --git a/docs/kb/development/IMPLEMENTATION.md b/docs/kb/development/IMPLEMENTATION.md index 8dd0d81..4395aab 100644 --- a/docs/kb/development/IMPLEMENTATION.md +++ b/docs/kb/development/IMPLEMENTATION.md @@ -129,7 +129,7 @@ package main import ( - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -152,7 +152,7 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -175,7 +175,7 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -198,7 +198,7 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "time" ) diff --git a/docs/kb/implementation-details/ANSWER_TO_USER_QUESTION.md b/docs/kb/implementation-details/ANSWER_TO_USER_QUESTION.md index 698eb70..1944a8a 100644 --- a/docs/kb/implementation-details/ANSWER_TO_USER_QUESTION.md +++ b/docs/kb/implementation-details/ANSWER_TO_USER_QUESTION.md @@ -116,8 +116,8 @@ cfg := &nnet.Config{ package main import ( - "github.com/noahlann/nnet" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) type User struct { diff --git a/docs/kb/implementation-details/CODEC_AND_PROTOCOL_INTEGRATION.md b/docs/kb/implementation-details/CODEC_AND_PROTOCOL_INTEGRATION.md index ddd19de..e77f146 100644 --- a/docs/kb/implementation-details/CODEC_AND_PROTOCOL_INTEGRATION.md +++ b/docs/kb/implementation-details/CODEC_AND_PROTOCOL_INTEGRATION.md @@ -229,8 +229,8 @@ if err != nil { package main import ( - "github.com/noahlann/nnet" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) type Response struct { diff --git a/docs/kb/implementation-details/SERVER_WRITE_DATA_SUMMARY.md b/docs/kb/implementation-details/SERVER_WRITE_DATA_SUMMARY.md index 28f858a..188b2aa 100644 --- a/docs/kb/implementation-details/SERVER_WRITE_DATA_SUMMARY.md +++ b/docs/kb/implementation-details/SERVER_WRITE_DATA_SUMMARY.md @@ -230,8 +230,8 @@ WriteAnyWithCodec(data interface{}, codecName string) error package main import ( - "github.com/noahlann/nnet" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) type Response struct { diff --git a/docs/kb/implementation-details/WRITE_DATA_TO_CLIENT.md b/docs/kb/implementation-details/WRITE_DATA_TO_CLIENT.md index 01a7d75..8609435 100644 --- a/docs/kb/implementation-details/WRITE_DATA_TO_CLIENT.md +++ b/docs/kb/implementation-details/WRITE_DATA_TO_CLIENT.md @@ -103,8 +103,8 @@ func handler(ctx context.Context) error { package main import ( - "github.com/noahlann/nnet" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) type User struct { diff --git a/docs/kb/user-guide/CLIENT_POOL.md b/docs/kb/user-guide/CLIENT_POOL.md index f9bd9a1..e7272ba 100644 --- a/docs/kb/user-guide/CLIENT_POOL.md +++ b/docs/kb/user-guide/CLIENT_POOL.md @@ -17,7 +17,7 @@ ### 1. 创建连接池 ```go -import "github.com/noahlann/nnet/pkg/nnet" +import "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" // 创建连接池配置 poolConfig := &nnet.ClientPoolConfig{ diff --git a/docs/kb/user-guide/DOMAIN_SETUP.md b/docs/kb/user-guide/DOMAIN_SETUP.md index a209ffd..3e91f5a 100644 --- a/docs/kb/user-guide/DOMAIN_SETUP.md +++ b/docs/kb/user-guide/DOMAIN_SETUP.md @@ -2,7 +2,7 @@ ## ⚠️ 注意 -**当前项目已改用 GitHub 作为模块路径**:`github.com/noahlann/nnet` +**当前项目已改用 GitHub 作为模块路径**:`git.noahlan.cn/noahlan/nnet/v2` 本文档仅作为参考,说明如何配置自定义域名。如果您需要使用自定义域名,可以参考本文档。 @@ -14,7 +14,7 @@ - 更专业的模块路径 - 更好的控制权 -**当前配置**:项目使用 `github.com/noahlann/nnet`,无需域名配置。 +**当前配置**:项目使用 `git.noahlan.cn/noahlan/nnet/v2`,无需域名配置。 ## 二、域名 DNS 配置 diff --git a/docs/kb/user-guide/MULTI_REPO_SETUP.md b/docs/kb/user-guide/MULTI_REPO_SETUP.md index 3c6d739..434a246 100644 --- a/docs/kb/user-guide/MULTI_REPO_SETUP.md +++ b/docs/kb/user-guide/MULTI_REPO_SETUP.md @@ -9,7 +9,7 @@ git remote -v # 添加GitHub远程仓库(通常命名为origin) -git remote add origin https://github.com/noahlann/nnet.git +git remote add origin https://git.noahlan.cn/noahlan/nnet.git # 添加Gitee远程仓库 git remote add gitee https://gitee.com/noahlann/nnet.git @@ -31,8 +31,8 @@ git remote -v 输出示例: ``` -origin https://github.com/noahlann/nnet.git (fetch) -origin https://github.com/noahlann/nnet.git (push) +origin https://git.noahlan.cn/noahlan/nnet.git (fetch) +origin https://git.noahlan.cn/noahlan/nnet.git (push) gitee https://gitee.com/noahlann/nnet.git (fetch) gitee https://gitee.com/noahlann/nnet.git (push) private https://your-private-git.com/noahlann/nnet.git (fetch) @@ -45,7 +45,7 @@ private https://your-private-git.com/noahlann/nnet.git (push) ```bash # 修改origin的URL -git remote set-url origin https://github.com/noahlann/nnet.git +git remote set-url origin https://git.noahlan.cn/noahlan/nnet.git # 修改gitee的URL git remote set-url gitee https://gitee.com/noahlann/nnet.git @@ -70,7 +70,7 @@ git push private main ```bash # 为origin添加多个push URL -git remote set-url --add --push origin https://github.com/noahlann/nnet.git +git remote set-url --add --push origin https://git.noahlan.cn/noahlan/nnet.git git remote set-url --add --push origin https://gitee.com/noahlann/nnet.git git remote set-url --add --push origin https://your-private-git.com/noahlann/nnet.git @@ -138,7 +138,7 @@ Go的 `go.mod` 中只能有一个 `module` 路径,但不同仓库可能有不 ```go // go.mod -module github.com/noahlann/nnet +module git.noahlan.cn/noahlan/nnet/v2 go 1.25 ``` @@ -153,7 +153,7 @@ go 1.25 - 如果GitHub不可访问,会有问题 **当前配置**: -- 模块路径:`github.com/noahlann/nnet` +- 模块路径:`git.noahlan.cn/noahlan/nnet/v2` - 无需域名配置 #### 方案2:使用Gitee作为主路径(国内用户) @@ -238,7 +238,7 @@ nnet/ ```ini [remote "origin"] - url = https://github.com/noahlann/nnet.git + url = https://git.noahlan.cn/noahlan/nnet.git fetch = +refs/heads/*:refs/remotes/origin/* [remote "gitee"] @@ -256,7 +256,7 @@ nnet/ ### 4.3 go.mod配置 ```go -module github.com/noahlann/nnet +module git.noahlan.cn/noahlan/nnet/v2 go 1.25 @@ -337,7 +337,7 @@ push-all: .PHONY: setup-remotes setup-remotes: @echo "Setting up remote repositories..." - @git remote add origin https://github.com/noahlann/nnet.git || true + @git remote add origin https://git.noahlan.cn/noahlan/nnet.git || true @git remote add gitee https://gitee.com/noahlann/nnet.git || true @git remote add private https://your-private-git.com/noahlann/nnet.git || true @echo "Remotes configured!" @@ -414,7 +414,7 @@ jobs: env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} run: | - git remote add github https://x-access-token:${GITHUB_TOKEN}@github.com/noahlann/nnet.git || true + git remote add github https://x-access-token:${GITHUB_TOKEN}@git.noahlan.cn/noahlan/nnet.git || true git push github main --force || true ``` @@ -613,7 +613,7 @@ chmod +x scripts/push-all.sh ### 10.1 推荐方案 1. **使用提供的脚本**:`./scripts/setup-remotes.sh` 和 `./scripts/push-all.sh` -2. **使用GitHub作为主路径**:go.mod使用 `github.com/noahlann/nnet`(当前配置) +2. **使用GitHub作为主路径**:go.mod使用 `git.noahlan.cn/noahlan/nnet/v2`(当前配置) 3. **配置CI/CD自动同步**:减少手动操作 4. **使用SSH密钥**:更安全便捷 @@ -625,7 +625,7 @@ chmod +x scripts/setup-remotes.sh ./scripts/setup-remotes.sh # 方式2:手动配置 -git remote add origin https://github.com/noahlann/nnet.git +git remote add origin https://git.noahlan.cn/noahlan/nnet.git git remote add gitee https://gitee.com/noahlann/nnet.git git remote add private https://your-private-git.com/noahlann/nnet.git diff --git a/docs/kb/user-guide/PROTOCOL_FRAME_EXPLANATION.md b/docs/kb/user-guide/PROTOCOL_FRAME_EXPLANATION.md index 9702b5b..3fb4a22 100644 --- a/docs/kb/user-guide/PROTOCOL_FRAME_EXPLANATION.md +++ b/docs/kb/user-guide/PROTOCOL_FRAME_EXPLANATION.md @@ -178,8 +178,8 @@ protocol.RegisterFieldExtractor(extractor) package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/protocol" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // 定义协议帧头结构 @@ -241,8 +241,8 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/protocol" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // 定义字段映射 @@ -296,8 +296,8 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/protocol" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // 自定义字段提取器 diff --git a/docs/kb/user-guide/ROUTER_EXAMPLES.md b/docs/kb/user-guide/ROUTER_EXAMPLES.md index 6ccc089..9ad3de5 100644 --- a/docs/kb/user-guide/ROUTER_EXAMPLES.md +++ b/docs/kb/user-guide/ROUTER_EXAMPLES.md @@ -78,8 +78,8 @@ router.GET("/user/login", handler) package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) func main() { @@ -126,8 +126,8 @@ func loginHandler(ctx router.Context) error { package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // 定义协议帧头结构 @@ -247,8 +247,8 @@ server.Router().Register(comboMatcher, handler) package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) func main() { @@ -354,8 +354,8 @@ server.Router().RegisterFrameData( package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) func main() { @@ -406,8 +406,8 @@ func contains(data, pattern []byte) bool { package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) func main() { @@ -443,8 +443,8 @@ func main() { package main import ( - "github.com/noahlann/nnet/pkg/nnet" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) func main() { diff --git a/examples/client_pool/client.go b/examples/client_pool/client.go index c137519..2c48d23 100644 --- a/examples/client_pool/client.go +++ b/examples/client_pool/client.go @@ -8,7 +8,7 @@ import ( "sync" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { diff --git a/examples/client_pool/server.go b/examples/client_pool/server.go index 011e3ee..125ca79 100644 --- a/examples/client_pool/server.go +++ b/examples/client_pool/server.go @@ -7,7 +7,7 @@ import ( "log" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { diff --git a/examples/client_request/main.go b/examples/client_request/main.go index 23ce22c..67ecc88 100644 --- a/examples/client_request/main.go +++ b/examples/client_request/main.go @@ -7,7 +7,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates the client Request/Response API with timeout. diff --git a/examples/interceptor/client.go b/examples/interceptor/client.go index 123b175..37e7c01 100644 --- a/examples/interceptor/client.go +++ b/examples/interceptor/client.go @@ -9,7 +9,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates client testing interceptor validation. @@ -67,4 +67,3 @@ func main() { fmt.Printf("response: %s\n", string(resp)) } - diff --git a/examples/interceptor/server.go b/examples/interceptor/server.go index 91130a5..6d8f70a 100644 --- a/examples/interceptor/server.go +++ b/examples/interceptor/server.go @@ -8,11 +8,11 @@ import ( "os/signal" "syscall" - internalinterceptor "github.com/noahlann/nnet/internal/interceptor" - "github.com/noahlann/nnet/internal/interceptor/builtin" - "github.com/noahlann/nnet/pkg/interceptor" - "github.com/noahlann/nnet/pkg/nnet" - routerpkg "github.com/noahlann/nnet/pkg/router" + internalinterceptor "git.noahlan.cn/noahlan/nnet/v2/internal/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/internal/interceptor/builtin" + "git.noahlan.cn/noahlan/nnet/v2/pkg/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // This example demonstrates interceptor chain usage. diff --git a/examples/metrics_health/main.go b/examples/metrics_health/main.go index b4bb408..3196c86 100644 --- a/examples/metrics_health/main.go +++ b/examples/metrics_health/main.go @@ -7,7 +7,7 @@ import ( "syscall" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates enabling metrics and health checks. @@ -46,5 +46,3 @@ func main() { signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) <-quit } - - diff --git a/examples/middleware/client.go b/examples/middleware/client.go index 2f1c2f6..402ce90 100644 --- a/examples/middleware/client.go +++ b/examples/middleware/client.go @@ -9,7 +9,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -41,4 +41,3 @@ func main() { } fmt.Printf("response: %s\n", string(resp)) } - diff --git a/examples/middleware/server.go b/examples/middleware/server.go index 92c5409..6949087 100644 --- a/examples/middleware/server.go +++ b/examples/middleware/server.go @@ -6,7 +6,7 @@ import ( "log" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -43,4 +43,3 @@ func main() { log.Fatal(err) } } - diff --git a/examples/preset/main.go b/examples/preset/main.go index 59f29b5..2741e18 100644 --- a/examples/preset/main.go +++ b/examples/preset/main.go @@ -7,7 +7,7 @@ import ( "os/signal" "syscall" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { diff --git a/examples/protocol_version/client.go b/examples/protocol_version/client.go index 2044e36..e4a9ba5 100644 --- a/examples/protocol_version/client.go +++ b/examples/protocol_version/client.go @@ -9,7 +9,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates client connecting to a protocol version server. @@ -56,4 +56,3 @@ func main() { fmt.Printf("response: %s\n", string(resp)) } - diff --git a/examples/protocol_version/server.go b/examples/protocol_version/server.go index f26055e..fdb11b3 100644 --- a/examples/protocol_version/server.go +++ b/examples/protocol_version/server.go @@ -8,8 +8,8 @@ import ( "os/signal" "syscall" - internalprotocol "github.com/noahlann/nnet/internal/protocol/nnet" - "github.com/noahlann/nnet/pkg/nnet" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates protocol version management. @@ -92,4 +92,3 @@ func main() { signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) <-quit } - diff --git a/examples/router_frame_data/main.go b/examples/router_frame_data/main.go index 4393d46..bf341c6 100644 --- a/examples/router_frame_data/main.go +++ b/examples/router_frame_data/main.go @@ -6,7 +6,7 @@ import ( "os/signal" "syscall" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates using a frame-data matcher to route JSON payloads. diff --git a/examples/session_file/client.go b/examples/session_file/client.go index e9a2490..5d74850 100644 --- a/examples/session_file/client.go +++ b/examples/session_file/client.go @@ -9,7 +9,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates client interacting with a session file server. @@ -47,4 +47,3 @@ func main() { fmt.Printf("response: %s\n", string(resp)) } - diff --git a/examples/session_file/server.go b/examples/session_file/server.go index 6c0bd38..4c83abb 100644 --- a/examples/session_file/server.go +++ b/examples/session_file/server.go @@ -9,7 +9,7 @@ import ( "syscall" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // This example demonstrates file-based session storage. @@ -18,9 +18,9 @@ func main() { cfg := &nnet.Config{ Addr: "tcp://:8083", Session: &nnet.SessionConfig{ - Storage: "file", // 使用文件存储 - Path: "./sessions_example", // Session文件存储路径 - Expiration: 1 * time.Hour, // Session过期时间 + Storage: "file", // 使用文件存储 + Path: "./sessions_example", // Session文件存储路径 + Expiration: 1 * time.Hour, // Session过期时间 }, Codec: &nnet.CodecConfig{ DefaultCodec: "json", @@ -40,11 +40,11 @@ func main() { // 实际使用中,Session应该通过Connection或Context访问 connID := ctx.Connection().ID() return ctx.Response().Write(map[string]any{ - "action": "set", - "conn_id": connID, - "message": "session storage configured (file)", - "path": cfg.Session.Path, - "note": "Session API integration pending", + "action": "set", + "conn_id": connID, + "message": "session storage configured (file)", + "path": cfg.Session.Path, + "note": "Session API integration pending", }) }) @@ -52,10 +52,10 @@ func main() { srv.Router().RegisterString("get", func(ctx nnet.Context) error { connID := ctx.Connection().ID() return ctx.Response().Write(map[string]any{ - "action": "get", - "conn_id": connID, - "storage": "file", - "path": cfg.Session.Path, + "action": "get", + "conn_id": connID, + "storage": "file", + "path": cfg.Session.Path, "expiration": cfg.Session.Expiration.String(), }) }) @@ -85,4 +85,3 @@ func main() { log.Println("Shutting down... Sessions are persisted to disk.") } - diff --git a/examples/udp_echo/client.go b/examples/udp_echo/client.go index 07708e1..0b5b608 100644 --- a/examples/udp_echo/client.go +++ b/examples/udp_echo/client.go @@ -9,7 +9,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -42,4 +42,3 @@ func main() { } fmt.Printf("echo: %s\n", string(resp)) } - diff --git a/examples/udp_echo/server.go b/examples/udp_echo/server.go index ca877ca..07a417a 100644 --- a/examples/udp_echo/server.go +++ b/examples/udp_echo/server.go @@ -5,7 +5,7 @@ package main import ( "log" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -26,4 +26,3 @@ func main() { log.Fatal(err) } } - diff --git a/examples/write_any/main.go b/examples/write_any/main.go index 38aa0dc..9066037 100644 --- a/examples/write_any/main.go +++ b/examples/write_any/main.go @@ -6,7 +6,7 @@ import ( "os/signal" "syscall" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) // User 用户结构 diff --git a/examples/ws_echo/client.go b/examples/ws_echo/client.go index adb746e..cd19b81 100644 --- a/examples/ws_echo/client.go +++ b/examples/ws_echo/client.go @@ -9,7 +9,7 @@ import ( "os" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -42,4 +42,3 @@ func main() { } fmt.Printf("echo: %s\n", string(resp)) } - diff --git a/examples/ws_echo/server.go b/examples/ws_echo/server.go index 836171f..2bd0293 100644 --- a/examples/ws_echo/server.go +++ b/examples/ws_echo/server.go @@ -5,7 +5,7 @@ package main import ( "log" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" ) func main() { @@ -30,4 +30,3 @@ func main() { log.Fatal(err) } } - diff --git a/go.mod b/go.mod index 9cb3064..e5a265a 100644 --- a/go.mod +++ b/go.mod @@ -1,10 +1,11 @@ -module github.com/noahlann/nnet +module git.noahlan.cn/noahlan/nnet/v2 go 1.25 require ( github.com/gorilla/websocket v1.5.3 github.com/panjf2000/gnet/v2 v2.9.5 + github.com/rs/xid v1.6.0 github.com/stretchr/testify v1.11.1 go.bug.st/serial v1.6.4 ) @@ -14,7 +15,6 @@ require ( github.com/davecgh/go-spew v1.1.1 // indirect github.com/panjf2000/ants/v2 v2.11.3 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect - github.com/rs/xid v1.6.0 // indirect github.com/valyala/bytebufferpool v1.0.0 // indirect go.uber.org/multierr v1.11.0 // indirect go.uber.org/zap v1.27.0 // indirect diff --git a/internal/client/application_protocol.go b/internal/client/application_protocol.go new file mode 100644 index 0000000..71e6a22 --- /dev/null +++ b/internal/client/application_protocol.go @@ -0,0 +1,56 @@ +package client + +import ( + "strings" + + protocolnnet "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" +) + +type unpackingProtocol interface { + Unpacker() unpackerpkg.Unpacker +} + +func newApplicationProtocol(name string) protocolpkg.Protocol { + switch strings.ToLower(strings.TrimSpace(name)) { + case "", "none", "raw": + return nil + case "nnet": + return protocolnnet.NewNNetProtocol("1.0") + default: + return nil + } +} + +func newApplicationUnpacker(protocol protocolpkg.Protocol) unpackerpkg.Unpacker { + if protocol == nil { + return nil + } + if provider, ok := protocol.(unpackingProtocol); ok { + return provider.Unpacker() + } + return nil +} + +func encodeApplicationMessage(protocol protocolpkg.Protocol, data []byte) ([]byte, error) { + if protocol == nil { + return data, nil + } + return protocol.Encode(data, nil) +} + +func decodeApplicationMessage(protocol protocolpkg.Protocol, frame []byte) ([]byte, error) { + if protocol == nil { + msg := make([]byte, len(frame)) + copy(msg, frame) + return msg, nil + } + _, body, err := protocol.Decode(frame) + if err != nil { + return nil, err + } + msg := make([]byte, len(body)) + copy(msg, body) + return msg, nil +} diff --git a/internal/client/pool.go b/internal/client/pool.go index b57b040..066bc22 100644 --- a/internal/client/pool.go +++ b/internal/client/pool.go @@ -6,7 +6,7 @@ import ( "sync/atomic" "time" - clientpkg "github.com/noahlann/nnet/pkg/client" + clientpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/client" ) // Pool 连接池 diff --git a/internal/client/pool_test.go b/internal/client/pool_test.go index 72498ab..8cc9186 100644 --- a/internal/client/pool_test.go +++ b/internal/client/pool_test.go @@ -5,7 +5,7 @@ import ( "testing" "time" - clientpkg "github.com/noahlann/nnet/pkg/client" + clientpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/client" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/client/serial_client.go b/internal/client/serial_client.go index 67f4e9b..a896668 100644 --- a/internal/client/serial_client.go +++ b/internal/client/serial_client.go @@ -6,7 +6,7 @@ import ( "sync" "time" - "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" "go.bug.st/serial" ) @@ -175,4 +175,3 @@ func (c *serialClient) Close() error { c.cancel() return c.Disconnect() } - diff --git a/internal/client/tcp_client.go b/internal/client/tcp_client.go index f42d3ee..e09462d 100644 --- a/internal/client/tcp_client.go +++ b/internal/client/tcp_client.go @@ -2,11 +2,14 @@ package client import ( "context" + "io" "net" "sync" "time" - "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // tcpClient TCP客户端实现 @@ -18,6 +21,13 @@ type tcpClient struct { ctx context.Context cancel context.CancelFunc readBuffer []byte + requestMu sync.Mutex + readMu sync.Mutex + writeMu sync.Mutex + closeOnce sync.Once + protocol protocolpkg.Protocol + unpacker unpackerpkg.Unpacker + pending [][]byte // 自动重连相关 reconnectAttempts int @@ -41,16 +51,19 @@ func NewTCPClient(config *client.Config) client.Client { } ctx, cancel := context.WithCancel(context.Background()) + protocol := newApplicationProtocol(config.ApplicationProtocol) c := &tcpClient{ - config: config, - ctx: ctx, - cancel: cancel, - readBuffer: make([]byte, 4096), - reconnectCh: make(chan struct{}, 1), + config: config, + ctx: ctx, + cancel: cancel, + readBuffer: make([]byte, 4096), + protocol: protocol, + unpacker: newApplicationUnpacker(protocol), + reconnectCh: make(chan struct{}, 1), reconnectStopCh: make(chan struct{}), - messageCh: make(chan []byte, 100), - messageErrCh: make(chan error, 10), + messageCh: make(chan []byte, 100), + messageErrCh: make(chan error, 10), } // 如果启用自动重连,启动重连goroutine @@ -85,13 +98,36 @@ func (c *tcpClient) connect() error { } // 创建连接 - dialer := &net.Dialer{ - Timeout: c.config.ConnectTimeout, + connectTimeout := c.config.ConnectTimeout + if connectTimeout <= 0 { + connectTimeout = 3 * time.Second } - conn, err := dialer.DialContext(c.ctx, "tcp", addr) + deadline := time.Now().Add(connectTimeout) + var conn net.Conn + var err error + for { + timeout := time.Until(deadline) + if timeout <= 0 { + break + } + if timeout > 200*time.Millisecond { + timeout = 200 * time.Millisecond + } + + dialer := &net.Dialer{Timeout: timeout} + conn, err = dialer.DialContext(c.ctx, "tcp", addr) + if err == nil { + break + } + + select { + case <-c.ctx.Done(): + return c.ctx.Err() + case <-time.After(20 * time.Millisecond): + } + } if err != nil { - // 如果启用自动重连,触发重连 if c.config.AutoReconnect { select { case c.reconnectCh <- struct{}{}: @@ -104,6 +140,8 @@ func (c *tcpClient) connect() error { c.conn = conn c.connected = true c.reconnectAttempts = 0 // 重置重连次数 + c.unpacker = newApplicationUnpacker(c.protocol) + c.pending = nil return nil } @@ -123,6 +161,7 @@ func (c *tcpClient) Disconnect() error { } c.connected = false + c.pending = nil // 如果启用自动重连,触发重连 if c.config.AutoReconnect { @@ -135,62 +174,168 @@ func (c *tcpClient) Disconnect() error { return nil } +func (c *tcpClient) markDisconnected() { + c.mu.Lock() + defer c.mu.Unlock() + + if c.conn != nil { + _ = c.conn.Close() + c.conn = nil + } + c.connected = false + c.pending = nil + + if c.config.AutoReconnect { + select { + case c.reconnectCh <- struct{}{}: + default: + } + } +} + +func shouldMarkDisconnected(err error) bool { + if err == nil { + return false + } + if netErr, ok := err.(net.Error); ok && netErr.Timeout() { + return false + } + return true +} + // Send 发送数据 func (c *tcpClient) Send(data []byte) error { - c.mu.RLock() - defer c.mu.RUnlock() + frame, err := encodeApplicationMessage(c.protocol, data) + if err != nil { + return err + } + c.mu.RLock() if !c.connected || c.conn == nil { + c.mu.RUnlock() return client.NewError("not connected") } + conn := c.conn + writeTimeout := c.config.WriteTimeout + c.mu.RUnlock() + + c.writeMu.Lock() + defer c.writeMu.Unlock() - if c.config.WriteTimeout > 0 { - c.conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)) + if writeTimeout > 0 { + conn.SetWriteDeadline(time.Now().Add(writeTimeout)) } - _, err := c.conn.Write(data) - return err + for len(frame) > 0 { + n, err := conn.Write(frame) + if shouldMarkDisconnected(err) { + c.markDisconnected() + } + if err != nil { + return err + } + if n == 0 { + c.markDisconnected() + return io.ErrShortWrite + } + frame = frame[n:] + } + return nil } // Receive 接收数据 func (c *tcpClient) Receive() ([]byte, error) { - c.mu.RLock() - defer c.mu.RUnlock() + return c.receiveWithTimeout(0) +} +func (c *tcpClient) receiveWithTimeout(timeout time.Duration) ([]byte, error) { + c.readMu.Lock() + defer c.readMu.Unlock() + + if len(c.pending) > 0 { + msg := c.pending[0] + c.pending[0] = nil + c.pending = c.pending[1:] + return msg, nil + } + + c.mu.RLock() if !c.connected || c.conn == nil { + c.mu.RUnlock() return nil, client.NewError("not connected") } - - if c.config.ReadTimeout > 0 { - c.conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)) + conn := c.conn + readTimeout := c.config.ReadTimeout + if timeout > 0 { + readTimeout = timeout } - - n, err := c.conn.Read(c.readBuffer) - if err != nil { - return nil, err + bufferSize := len(c.readBuffer) + if bufferSize <= 0 { + bufferSize = 4096 } + protocol := c.protocol + c.mu.RUnlock() + + for { + buffer := make([]byte, bufferSize) + if readTimeout > 0 { + conn.SetReadDeadline(time.Now().Add(readTimeout)) + } - result := make([]byte, n) - copy(result, c.readBuffer[:n]) - return result, nil + n, err := conn.Read(buffer) + if err != nil { + if shouldMarkDisconnected(err) { + c.markDisconnected() + } + return nil, err + } + + if protocol == nil { + result := make([]byte, n) + copy(result, buffer[:n]) + return result, nil + } + + if c.unpacker == nil { + return decodeApplicationMessage(protocol, buffer[:n]) + } + + frames, _, _, err := c.unpacker.Unpack(buffer[:n]) + if err != nil { + if shouldMarkDisconnected(err) { + c.markDisconnected() + } + return nil, err + } + if len(frames) == 0 { + continue + } + + decoded := make([][]byte, 0, len(frames)) + for _, frame := range frames { + msg, err := decodeApplicationMessage(protocol, frame) + if err != nil { + return nil, err + } + decoded = append(decoded, msg) + } + c.pending = append(c.pending, decoded[1:]...) + return decoded[0], nil + } } // Request 请求-响应(带超时) func (c *tcpClient) Request(data []byte, timeout time.Duration) ([]byte, error) { + c.requestMu.Lock() + defer c.requestMu.Unlock() + // 发送请求 if err := c.Send(data); err != nil { return nil, err } - // 设置读取超时 - oldTimeout := c.config.ReadTimeout - c.config.ReadTimeout = timeout - defer func() { - c.config.ReadTimeout = oldTimeout - }() - // 接收响应 - return c.Receive() + return c.receiveWithTimeout(timeout) } // IsConnected 检查是否已连接 @@ -202,8 +347,10 @@ func (c *tcpClient) IsConnected() bool { // Close 关闭客户端 func (c *tcpClient) Close() error { - c.cancel() - close(c.reconnectStopCh) + c.closeOnce.Do(func() { + c.cancel() + close(c.reconnectStopCh) + }) c.mu.Lock() defer c.mu.Unlock() c.config.AutoReconnect = false // 禁用自动重连 @@ -212,6 +359,7 @@ func (c *tcpClient) Close() error { c.conn = nil } c.connected = false + c.pending = nil return nil } @@ -294,8 +442,7 @@ func (c *tcpClient) messageReceiver() { conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)) } - buffer := make([]byte, 4096) - n, err := conn.Read(buffer) + msg, err := c.receiveWithTimeout(c.config.ReadTimeout) if err != nil { // 连接错误,触发重连 if c.config.AutoReconnect { @@ -318,17 +465,15 @@ func (c *tcpClient) messageReceiver() { } // 发送消息到channel - message := make([]byte, n) - copy(message, buffer[:n]) select { - case c.messageCh <- message: + case c.messageCh <- msg: default: // channel已满,丢弃消息 } // 调用回调 if c.onMessage != nil { - c.onMessage(message) + c.onMessage(msg) } } } @@ -357,4 +502,3 @@ func (c *tcpClient) SetOnError(fn func(error)) { func (c *tcpClient) MessageChannel() <-chan []byte { return c.messageCh } - diff --git a/internal/client/tcp_client_more_test.go b/internal/client/tcp_client_more_test.go index afce4e8..e4359f5 100644 --- a/internal/client/tcp_client_more_test.go +++ b/internal/client/tcp_client_more_test.go @@ -4,7 +4,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -53,5 +53,3 @@ func TestTCPClientConnectInvalidAddrRespectsTimeout(t *testing.T) { elapsed := time.Since(start) assert.LessOrEqual(t, elapsed, 500*time.Millisecond, "connect should fail quickly honoring timeout") } - - diff --git a/internal/client/tcp_client_test.go b/internal/client/tcp_client_test.go index 6aa3c5b..40dc896 100644 --- a/internal/client/tcp_client_test.go +++ b/internal/client/tcp_client_test.go @@ -4,7 +4,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/client/udp_client.go b/internal/client/udp_client.go index f3a6db1..1ba0e4b 100644 --- a/internal/client/udp_client.go +++ b/internal/client/udp_client.go @@ -5,7 +5,7 @@ import ( "sync" "time" - "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" ) // udpClient UDP客户端实现 diff --git a/internal/client/unix_client.go b/internal/client/unix_client.go index 87bc859..f6c7060 100644 --- a/internal/client/unix_client.go +++ b/internal/client/unix_client.go @@ -5,7 +5,7 @@ import ( "sync" "time" - "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" ) // unixClient Unix Domain Socket客户端实现 @@ -147,4 +147,3 @@ func (c *unixClient) IsConnected() bool { func (c *unixClient) Close() error { return c.Disconnect() } - diff --git a/internal/client/websocket_client.go b/internal/client/websocket_client.go index 402f6b5..7c06eb0 100644 --- a/internal/client/websocket_client.go +++ b/internal/client/websocket_client.go @@ -2,11 +2,14 @@ package client import ( "net/url" + "strings" "sync" "time" + "git.noahlan.cn/noahlan/nnet/v2/pkg/client" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/gorilla/websocket" - "github.com/noahlann/nnet/pkg/client" ) // websocketClient WebSocket客户端实现 @@ -14,7 +17,13 @@ type websocketClient struct { config *client.Config conn *websocket.Conn mu sync.RWMutex + requestMu sync.Mutex + readMu sync.Mutex + writeMu sync.Mutex connected bool + protocol protocolpkg.Protocol + unpacker unpackerpkg.Unpacker + pending [][]byte } // NewWebSocketClient 创建WebSocket客户端 @@ -22,10 +31,13 @@ func NewWebSocketClient(config *client.Config) client.Client { if config == nil { config = client.DefaultConfig() } + protocol := newApplicationProtocol(config.ApplicationProtocol) return &websocketClient{ config: config, connected: false, + protocol: protocol, + unpacker: newApplicationUnpacker(protocol), } } @@ -41,11 +53,11 @@ func (c *websocketClient) Connect() error { // 解析地址 addr := c.config.Addr scheme := "ws" - if len(addr) > 6 && addr[:6] == "ws://" { - addr = addr[6:] + if strings.HasPrefix(addr, "ws://") { + addr = strings.TrimPrefix(addr, "ws://") scheme = "ws" - } else if len(addr) > 7 && addr[:7] == "wss://" { - addr = addr[7:] + } else if strings.HasPrefix(addr, "wss://") { + addr = strings.TrimPrefix(addr, "wss://") scheme = "wss" } else if c.config.TLSEnabled { scheme = "wss" @@ -66,6 +78,8 @@ func (c *websocketClient) Connect() error { c.conn = conn c.connected = true + c.unpacker = newApplicationUnpacker(c.protocol) + c.pending = nil return nil } @@ -85,59 +99,125 @@ func (c *websocketClient) Disconnect() error { } c.connected = false + c.pending = nil return nil } // Send 发送数据 func (c *websocketClient) Send(data []byte) error { - c.mu.RLock() - defer c.mu.RUnlock() + frame, err := encodeApplicationMessage(c.protocol, data) + if err != nil { + return err + } + c.mu.RLock() if !c.connected || c.conn == nil { + c.mu.RUnlock() return client.NewError("not connected") } + conn := c.conn + writeTimeout := c.config.WriteTimeout + c.mu.RUnlock() - return c.conn.WriteMessage(websocket.TextMessage, data) + c.writeMu.Lock() + defer c.writeMu.Unlock() + + if writeTimeout > 0 { + conn.SetWriteDeadline(time.Now().Add(writeTimeout)) + } + + if err := conn.WriteMessage(websocket.BinaryMessage, frame); err != nil { + c.markDisconnected() + return err + } + return nil } // Receive 接收数据 func (c *websocketClient) Receive() ([]byte, error) { - c.mu.RLock() - defer c.mu.RUnlock() + return c.receiveWithTimeout(0) +} +func (c *websocketClient) receiveWithTimeout(timeout time.Duration) ([]byte, error) { + c.readMu.Lock() + defer c.readMu.Unlock() + + if len(c.pending) > 0 { + msg := c.pending[0] + c.pending[0] = nil + c.pending = c.pending[1:] + return msg, nil + } + + c.mu.RLock() if !c.connected || c.conn == nil { + c.mu.RUnlock() return nil, client.NewError("not connected") } - - // 设置读取超时 - if c.config.ReadTimeout > 0 { - c.conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)) + conn := c.conn + readTimeout := c.config.ReadTimeout + if timeout > 0 { + readTimeout = timeout } + protocol := c.protocol + c.mu.RUnlock() - _, data, err := c.conn.ReadMessage() - if err != nil { - return nil, err + // 设置读取超时 + for { + if readTimeout > 0 { + conn.SetReadDeadline(time.Now().Add(readTimeout)) + } + + _, data, err := conn.ReadMessage() + if err != nil { + c.markDisconnected() + return nil, err + } + + if protocol == nil { + msg := make([]byte, len(data)) + copy(msg, data) + return msg, nil + } + + if c.unpacker == nil { + return decodeApplicationMessage(protocol, data) + } + + frames, _, _, err := c.unpacker.Unpack(data) + if err != nil { + c.markDisconnected() + return nil, err + } + if len(frames) == 0 { + continue + } + + decoded := make([][]byte, 0, len(frames)) + for _, frame := range frames { + msg, err := decodeApplicationMessage(protocol, frame) + if err != nil { + return nil, err + } + decoded = append(decoded, msg) + } + c.pending = append(c.pending, decoded[1:]...) + return decoded[0], nil } - - return data, nil } // Request 请求-响应(带超时) func (c *websocketClient) Request(data []byte, timeout time.Duration) ([]byte, error) { + c.requestMu.Lock() + defer c.requestMu.Unlock() + // 发送请求 if err := c.Send(data); err != nil { return nil, err } - // 设置读取超时 - oldTimeout := c.config.ReadTimeout - c.config.ReadTimeout = timeout - defer func() { - c.config.ReadTimeout = oldTimeout - }() - // 接收响应 - return c.Receive() + return c.receiveWithTimeout(timeout) } // IsConnected 检查是否已连接 @@ -152,3 +232,14 @@ func (c *websocketClient) Close() error { return c.Disconnect() } +func (c *websocketClient) markDisconnected() { + c.mu.Lock() + defer c.mu.Unlock() + + if c.conn != nil { + _ = c.conn.Close() + c.conn = nil + } + c.connected = false + c.pending = nil +} diff --git a/internal/codec/binary.go b/internal/codec/binary.go index 7bb409d..b6a8e3e 100644 --- a/internal/codec/binary.go +++ b/internal/codec/binary.go @@ -5,7 +5,7 @@ import ( "encoding/binary" "reflect" - "github.com/noahlann/nnet/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" ) // BinaryCodec 二进制编解码器 @@ -104,4 +104,3 @@ func (c *BinaryCodec) Decode(data []byte, v interface{}) error { func (c *BinaryCodec) Name() string { return "binary" } - diff --git a/internal/codec/binary_test.go b/internal/codec/binary_test.go index c3c572e..601f32b 100644 --- a/internal/codec/binary_test.go +++ b/internal/codec/binary_test.go @@ -91,4 +91,3 @@ func TestBinaryCodecInvalidType(t *testing.T) { _, err := codec.Encode("invalid") assert.Error(t, err, "Expected error for invalid type") } - diff --git a/internal/codec/json.go b/internal/codec/json.go index f9ecf8e..f0eb7e1 100644 --- a/internal/codec/json.go +++ b/internal/codec/json.go @@ -3,7 +3,7 @@ package codec import ( "encoding/json" - "github.com/noahlann/nnet/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" ) // JSONCodec JSON编解码器 @@ -28,4 +28,3 @@ func (c *JSONCodec) Decode(data []byte, v interface{}) error { func (c *JSONCodec) Name() string { return "json" } - diff --git a/internal/codec/json_test.go b/internal/codec/json_test.go index 8cde503..6c790b4 100644 --- a/internal/codec/json_test.go +++ b/internal/codec/json_test.go @@ -126,4 +126,3 @@ func TestJSONCodecComplexStruct(t *testing.T) { assert.Equal(t, data.Tags, decoded.Tags) assert.Equal(t, data.Metadata, decoded.Metadata) } - diff --git a/internal/codec/msgpack.go b/internal/codec/msgpack.go index 5b20ce4..383d626 100644 --- a/internal/codec/msgpack.go +++ b/internal/codec/msgpack.go @@ -3,7 +3,7 @@ package codec import ( "encoding/json" - codecpkg "github.com/noahlann/nnet/pkg/codec" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" ) // MessagePackCodec MessagePack编解码器 @@ -62,4 +62,3 @@ func (d *defaultMessagePackDecoder) Decode(data []byte, v interface{}) error { // 如果没有安装msgpack库,使用JSON作为fallback return json.Unmarshal(data, v) } - diff --git a/internal/codec/msgpack_test.go b/internal/codec/msgpack_test.go index b52c02a..7da0822 100644 --- a/internal/codec/msgpack_test.go +++ b/internal/codec/msgpack_test.go @@ -59,4 +59,3 @@ func TestMessagePackCodecInvalidData(t *testing.T) { // 由于使用JSON fallback,可能会成功或失败,这里只测试不会panic _ = err } - diff --git a/internal/codec/plain.go b/internal/codec/plain.go index 982c7cf..a21fced 100644 --- a/internal/codec/plain.go +++ b/internal/codec/plain.go @@ -3,7 +3,7 @@ package codec import ( "fmt" - codecpkg "github.com/noahlann/nnet/pkg/codec" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" ) // PlainCodec 直接透传的编解码器 diff --git a/internal/codec/protobuf.go b/internal/codec/protobuf.go index 2715f8b..bc7f976 100644 --- a/internal/codec/protobuf.go +++ b/internal/codec/protobuf.go @@ -3,7 +3,7 @@ package codec import ( "fmt" - codecpkg "github.com/noahlann/nnet/pkg/codec" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" ) // ProtobufCodec Protobuf编解码器 diff --git a/internal/codec/protobuf_test.go b/internal/codec/protobuf_test.go index c5df7d1..593ada5 100644 --- a/internal/codec/protobuf_test.go +++ b/internal/codec/protobuf_test.go @@ -64,4 +64,3 @@ func (m *testProtoMessage) Unmarshal(data []byte) error { m.data = data return nil } - diff --git a/internal/codec/registry.go b/internal/codec/registry.go index 934997b..81149a4 100644 --- a/internal/codec/registry.go +++ b/internal/codec/registry.go @@ -4,7 +4,7 @@ import ( "fmt" "sync" - codecpkg "github.com/noahlann/nnet/pkg/codec" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" ) // registry 编解码器注册表实现 diff --git a/internal/connection/connection.go b/internal/connection/connection.go index 38b2eff..daf9a4d 100644 --- a/internal/connection/connection.go +++ b/internal/connection/connection.go @@ -66,10 +66,10 @@ func (c *Connection) Write(data []byte) error { c.mu.Lock() c.lastActive = time.Now() c.mu.Unlock() - + var buf []byte var shouldRecycle bool - + // 从池中获取缓冲区(零拷贝优化) pooledBuf := writeBufferPool.Get().([]byte) if cap(pooledBuf) >= len(data) { @@ -87,13 +87,13 @@ func (c *Connection) Write(data []byte) error { } } copy(buf, data) - + // 保存用于回收的缓冲区引用 recycleBuf := buf if !shouldRecycle { recycleBuf = nil } - + // 使用AsyncWrite异步写入,避免阻塞事件循环 err := c.conn.AsyncWrite(buf, func(c gnet.Conn, err error) error { // 写入完成后,将缓冲区归还到池中(零拷贝优化) diff --git a/internal/connection/connection_test.go b/internal/connection/connection_test.go index f56d6ba..6573649 100644 --- a/internal/connection/connection_test.go +++ b/internal/connection/connection_test.go @@ -4,7 +4,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/connection/group_strategy.go b/internal/connection/group_strategy.go index 6ab262d..2280479 100644 --- a/internal/connection/group_strategy.go +++ b/internal/connection/group_strategy.go @@ -93,4 +93,3 @@ func (s *ConditionGroupStrategy) GetGroupID(conn ConnectionInterface) string { func (s *ConditionGroupStrategy) Match(conn ConnectionInterface, groupID string) bool { return s.GetGroupID(conn) == groupID } - diff --git a/internal/connection/serial_connection.go b/internal/connection/serial_connection.go index 940e707..a992b41 100644 --- a/internal/connection/serial_connection.go +++ b/internal/connection/serial_connection.go @@ -104,4 +104,3 @@ func (c *SerialConnection) UpdateActive() { defer c.mu.Unlock() c.lastActive = time.Now() } - diff --git a/internal/connection/sharded_manager.go b/internal/connection/sharded_manager.go index e7f4ec0..676a360 100644 --- a/internal/connection/sharded_manager.go +++ b/internal/connection/sharded_manager.go @@ -6,7 +6,7 @@ import ( "sync/atomic" "time" - "github.com/noahlann/nnet/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" ) // ShardedManager 分段连接管理器(减少锁竞争) @@ -264,7 +264,7 @@ func (m *ShardedManager) BroadcastToGroup(groupID string, data []byte) error { defer wg.Done() var localErrs []error for _, conn := range connections { - if err := conn.Write(data); err != nil { + if err := conn.Write(data); err != nil { localErrs = append(localErrs, err) } } @@ -272,7 +272,7 @@ func (m *ShardedManager) BroadcastToGroup(groupID string, data []byte) error { mu.Lock() errs = append(errs, localErrs...) mu.Unlock() - } + } }(conns) } diff --git a/internal/connection/sharded_manager_test.go b/internal/connection/sharded_manager_test.go index 1ba6c0b..b7c477e 100644 --- a/internal/connection/sharded_manager_test.go +++ b/internal/connection/sharded_manager_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/connection/udp_connection.go b/internal/connection/udp_connection.go index 4599e0e..eb0012c 100644 --- a/internal/connection/udp_connection.go +++ b/internal/connection/udp_connection.go @@ -62,7 +62,7 @@ func (c *UDPConnection) LocalAddr() string { func (c *UDPConnection) Write(data []byte) error { c.mu.Lock() c.lastActive = time.Now() - + // 检查是否有gnet.Conn(通过属性传递) gnetConn := c.attributes["gnet_conn"] c.mu.Unlock() @@ -131,4 +131,3 @@ func (c *UDPConnection) UpdateActive() { defer c.mu.Unlock() c.lastActive = time.Now() } - diff --git a/internal/connection/unix_connection.go b/internal/connection/unix_connection.go index b30ab23..4af88f8 100644 --- a/internal/connection/unix_connection.go +++ b/internal/connection/unix_connection.go @@ -108,4 +108,3 @@ func (c *UnixConnection) UpdateActive() { defer c.mu.Unlock() c.lastActive = time.Now() } - diff --git a/internal/connection/websocket_connection.go b/internal/connection/websocket_connection.go index a693edb..bfa06ec 100644 --- a/internal/connection/websocket_connection.go +++ b/internal/connection/websocket_connection.go @@ -13,6 +13,7 @@ type WebSocketConnection struct { id string conn *websocket.Conn mu sync.RWMutex + writeMu sync.Mutex createdAt time.Time lastActive time.Time attributes map[string]interface{} @@ -61,6 +62,8 @@ func (c *WebSocketConnection) Write(data []byte) error { c.mu.Unlock() if c.conn != nil { + c.writeMu.Lock() + defer c.writeMu.Unlock() return c.conn.WriteMessage(websocket.TextMessage, data) } return nil @@ -107,4 +110,3 @@ func (c *WebSocketConnection) UpdateActive() { defer c.mu.Unlock() c.lastActive = time.Now() } - diff --git a/internal/interceptor/builtin/validation.go b/internal/interceptor/builtin/validation.go index 01e5de2..1f5e4b0 100644 --- a/internal/interceptor/builtin/validation.go +++ b/internal/interceptor/builtin/validation.go @@ -1,8 +1,8 @@ package builtin import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - interceptorpkg "github.com/noahlann/nnet/pkg/interceptor" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + interceptorpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/interceptor" ) // ValidationInterceptor 验证拦截器 @@ -39,4 +39,3 @@ func MaxLengthInterceptor(maxLength int) interceptorpkg.Interceptor { return nil }) } - diff --git a/internal/interceptor/builtin/validation_test.go b/internal/interceptor/builtin/validation_test.go index 0d2af73..ab0a909 100644 --- a/internal/interceptor/builtin/validation_test.go +++ b/internal/interceptor/builtin/validation_test.go @@ -3,12 +3,12 @@ package builtin import ( "testing" - "github.com/noahlann/nnet/internal/codec" - interceptorimpl "github.com/noahlann/nnet/internal/interceptor" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - interceptorpkg "github.com/noahlann/nnet/pkg/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + interceptorimpl "git.noahlan.cn/noahlan/nnet/v2/internal/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + interceptorpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/interceptor" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/interceptor/chain.go b/internal/interceptor/chain.go index c5a3e49..eb93e15 100644 --- a/internal/interceptor/chain.go +++ b/internal/interceptor/chain.go @@ -1,8 +1,8 @@ package interceptor import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - interceptorpkg "github.com/noahlann/nnet/pkg/interceptor" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + interceptorpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/interceptor" ) // chain 拦截器链实现 @@ -43,4 +43,3 @@ func Execute(data []byte, ctx ctxpkg.Context, interceptors ...interceptorpkg.Int chain := NewChain(interceptors...) return chain.Next(data, ctx) } - diff --git a/internal/interceptor/chain_test.go b/internal/interceptor/chain_test.go index 5ddf8ef..fef48c2 100644 --- a/internal/interceptor/chain_test.go +++ b/internal/interceptor/chain_test.go @@ -4,11 +4,11 @@ import ( "errors" "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - interceptorpkg "github.com/noahlann/nnet/pkg/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + interceptorpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/interceptor" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/middleware/builtin/auth.go b/internal/middleware/builtin/auth.go index a02d3f3..b4dac2e 100644 --- a/internal/middleware/builtin/auth.go +++ b/internal/middleware/builtin/auth.go @@ -1,8 +1,8 @@ package builtin import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // AuthMiddleware 认证中间件工厂 @@ -11,12 +11,12 @@ func AuthMiddleware(validator func(ctx ctxpkg.Context) bool) routerpkg.Handler { if validator == nil { return nil } - + if !validator(ctx) { ctx.Response().WriteString("Unauthorized\n") return nil // 或者返回错误 } - + return nil } } @@ -36,4 +36,3 @@ func HeaderAuthMiddleware(headerKey string, expectedValue string) routerpkg.Hand return value == expectedValue }) } - diff --git a/internal/middleware/builtin/logging.go b/internal/middleware/builtin/logging.go index 3c1e407..30b2cce 100644 --- a/internal/middleware/builtin/logging.go +++ b/internal/middleware/builtin/logging.go @@ -4,8 +4,8 @@ import ( "fmt" "time" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // LoggingMiddleware 日志中间件 @@ -13,7 +13,7 @@ func LoggingMiddleware() routerpkg.Handler { return func(ctx ctxpkg.Context) error { start := time.Now() conn := ctx.Connection() - + // 记录请求开始 req := ctx.Request() var requestStr string @@ -30,8 +30,8 @@ func LoggingMiddleware() routerpkg.Handler { } else { requestStr = string(req.Raw()) } - - fmt.Printf("[%s] %s -> %s: %s\n", + + fmt.Printf("[%s] %s -> %s: %s\n", start.Format("2006-01-02 15:04:05"), conn.RemoteAddr(), conn.LocalAddr(), diff --git a/internal/middleware/builtin/ratelimit.go b/internal/middleware/builtin/ratelimit.go index 5453fb3..a97debe 100644 --- a/internal/middleware/builtin/ratelimit.go +++ b/internal/middleware/builtin/ratelimit.go @@ -4,8 +4,8 @@ import ( "sync" "time" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // RateLimiter 限流器 @@ -68,4 +68,3 @@ func RateLimitMiddleware(limiter *RateLimiter) routerpkg.Handler { return nil } } - diff --git a/internal/middleware/builtin/recovery.go b/internal/middleware/builtin/recovery.go index 9d3f8c5..ed4cfd8 100644 --- a/internal/middleware/builtin/recovery.go +++ b/internal/middleware/builtin/recovery.go @@ -4,8 +4,8 @@ import ( "fmt" "runtime/debug" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // RecoveryMiddleware 恢复中间件(捕获panic) @@ -39,4 +39,3 @@ func RecoveryMiddlewareWithHandler(handler func(ctx ctxpkg.Context, panic interf return nil } } - diff --git a/internal/middleware/chain.go b/internal/middleware/chain.go index 103550d..446e290 100644 --- a/internal/middleware/chain.go +++ b/internal/middleware/chain.go @@ -1,8 +1,8 @@ package middleware import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // Chain 中间件链 @@ -53,4 +53,3 @@ func (c *Chain) Append(middleware ...routerpkg.Handler) { func (c *Chain) Prepend(middleware ...routerpkg.Handler) { c.middlewares = append(middleware, c.middlewares...) } - diff --git a/internal/middleware/chain_test.go b/internal/middleware/chain_test.go index 32448be..6d8fc5c 100644 --- a/internal/middleware/chain_test.go +++ b/internal/middleware/chain_test.go @@ -4,10 +4,10 @@ import ( "errors" "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/plugin/manager.go b/internal/plugin/manager.go index 6b6a13c..032e5e7 100644 --- a/internal/plugin/manager.go +++ b/internal/plugin/manager.go @@ -5,8 +5,8 @@ import ( "fmt" "sync" - "github.com/noahlann/nnet/pkg/config" - pluginpkg "github.com/noahlann/nnet/pkg/plugin" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + pluginpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/plugin" ) // manager 插件管理器实现 @@ -131,4 +131,3 @@ func (m *manager) StopAll(ctx context.Context) error { return nil } - diff --git a/internal/protocol/frame_header.go b/internal/protocol/frame_header.go index 907ccb8..c2ffa5f 100644 --- a/internal/protocol/frame_header.go +++ b/internal/protocol/frame_header.go @@ -4,7 +4,7 @@ import ( "encoding/json" "sync" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // frameHeader 协议帧头实现 @@ -99,4 +99,3 @@ func (h *frameHeader) Clone() protocolpkg.FrameHeader { return cloned } - diff --git a/internal/protocol/manager.go b/internal/protocol/manager.go index a17d2de..41de992 100644 --- a/internal/protocol/manager.go +++ b/internal/protocol/manager.go @@ -5,7 +5,7 @@ import ( "fmt" "sync" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // manager 协议管理器实现 @@ -213,4 +213,3 @@ func (m *manager) IdentifyVersion(name string, data []byte, ctx context.Context) return version, nil } - diff --git a/internal/protocol/manager_test.go b/internal/protocol/manager_test.go index ef00248..eb1539e 100644 --- a/internal/protocol/manager_test.go +++ b/internal/protocol/manager_test.go @@ -6,8 +6,8 @@ import ( "testing" "time" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/protocol/nnet/incremental_decoder_test.go b/internal/protocol/nnet/incremental_decoder_test.go index b8bf909..49e0dca 100644 --- a/internal/protocol/nnet/incremental_decoder_test.go +++ b/internal/protocol/nnet/incremental_decoder_test.go @@ -3,7 +3,7 @@ package nnet import ( "testing" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/protocol/nnet/protocol.go b/internal/protocol/nnet/protocol.go index f0315bb..befee32 100644 --- a/internal/protocol/nnet/protocol.go +++ b/internal/protocol/nnet/protocol.go @@ -4,19 +4,16 @@ import ( "context" "encoding/binary" "fmt" - "sync" - internalprotocol "github.com/noahlann/nnet/internal/protocol" - internalunpacker "github.com/noahlann/nnet/internal/unpacker" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol" + internalunpacker "git.noahlan.cn/noahlan/nnet/v2/internal/unpacker" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // NNetProtocol nnet协议实现 type NNetProtocol struct { - version string - unpacker unpackerpkg.Unpacker - once sync.Once + version string } // NewNNetProtocol 创建nnet协议 @@ -124,7 +121,7 @@ func (p *NNetProtocol) Decode(data []byte) (protocolpkg.FrameHeader, []byte, err // 计算预期的总长度:Magic(4) + Version(1) + Length(4) + Data(dataLength) + Checksum(2) expectedTotalLength := 9 + int(dataLength) + 2 - + // 优化:如果数据长度正好等于预期长度,说明数据来自unpacker(已经完整),可以跳过长度验证 // 否则,需要进行长度验证(数据可能不完整) if len(data) != expectedTotalLength { @@ -175,20 +172,15 @@ func (p *NNetProtocol) Handle(ctx context.Context, data []byte) ([]byte, error) // Unpacker 获取协议的拆包器 // nnet协议使用LengthFieldUnpacker来处理粘包拆包 func (p *NNetProtocol) Unpacker() unpackerpkg.Unpacker { - p.once.Do(func() { - // nnet协议格式:[Magic(4)][Version(1)][Length(4)][Data(N)][Checksum(2)] - // 长度字段在偏移5的位置(Magic 4字节 + Version 1字节) - // 长度字段是4字节,表示Data部分的长度 - // 总长度 = 5(Magic+Version) + 4(Length字段) + Length(数据长度) + 2(Checksum) - config := unpackerpkg.LengthFieldUnpacker{ - LengthFieldOffset: 5, // Magic(4) + Version(1) = 5 - LengthFieldLength: 4, // Length字段是4字节 - LengthAdjustment: 2, // 需要加上Checksum(2字节) - InitialBytesToStrip: 0, // 不跳过任何字节,保留完整包 - } - p.unpacker = internalunpacker.NewLengthFieldUnpacker(config) - }) - return p.unpacker + // The unpacker owns mutable buffering state, so each connection must get + // its own instance. Protocol objects are shared by the protocol manager. + config := unpackerpkg.LengthFieldUnpacker{ + LengthFieldOffset: 5, // Magic(4) + Version(1) = 5 + LengthFieldLength: 4, // Length字段是4字节 + LengthAdjustment: 2, // 需要加上Checksum(2字节) + InitialBytesToStrip: 0, // 不跳过任何字节,保留完整包 + } + return internalunpacker.NewLengthFieldUnpacker(config) } // DecodeHeader 解码帧头(增量解析,即使数据不完整) @@ -258,7 +250,7 @@ func (p *NNetProtocol) DecodeBody(data []byte, header protocolpkg.FrameHeader) ( // 计算预期的总长度:Magic(4) + Version(1) + Length(4) + Data(dataLength) + Checksum(2) expectedTotalLength := 9 + int(dataLength) + 2 - + // 验证数据长度(数据应该来自unpacker,已经完整) if len(data) != expectedTotalLength { if len(data) < expectedTotalLength { diff --git a/internal/protocol/nnet/protocol_test.go b/internal/protocol/nnet/protocol_test.go index 044b18e..6e578ff 100644 --- a/internal/protocol/nnet/protocol_test.go +++ b/internal/protocol/nnet/protocol_test.go @@ -4,8 +4,8 @@ import ( "encoding/binary" "testing" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -95,9 +95,9 @@ func TestNNetProtocol_Unpacker_Config(t *testing.T) { u := withUnpacker.Unpacker() require.NotNil(t, u) - // Should return same singleton + // Unpackers own mutable buffering state, so callers must receive distinct instances. u2 := withUnpacker.Unpacker() - assert.Equal(t, u, u2) + assert.NotSame(t, u, u2) } // Ensure it implements protocol interface diff --git a/internal/protocol/nnet/version_identifier.go b/internal/protocol/nnet/version_identifier.go index 21cf959..f4db6c7 100644 --- a/internal/protocol/nnet/version_identifier.go +++ b/internal/protocol/nnet/version_identifier.go @@ -4,7 +4,7 @@ import ( "context" "fmt" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // NNetVersionIdentifier NNet协议版本识别器 @@ -58,4 +58,3 @@ func (i *NNetVersionIdentifier) Identify(data []byte, ctx context.Context) (stri return version, nil } - diff --git a/internal/protocol/nnet/version_identifier_test.go b/internal/protocol/nnet/version_identifier_test.go index 6efe630..a995d12 100644 --- a/internal/protocol/nnet/version_identifier_test.go +++ b/internal/protocol/nnet/version_identifier_test.go @@ -20,10 +20,10 @@ func TestNNetVersionIdentifier_Identify(t *testing.T) { t.Run("Identify version 1.0", func(t *testing.T) { // 创建版本1的数据包(Magic + Version(1) + Length(4) + Data + Checksum) data := []byte("NNET") - data = append(data, byte(1)) // version + data = append(data, byte(1)) // version data = append(data, []byte{0, 0, 0, 5}...) // length - data = append(data, []byte("hello")...) // data - data = append(data, []byte{0, 0}...) // checksum (简化) + data = append(data, []byte("hello")...) // data + data = append(data, []byte{0, 0}...) // checksum (简化) version, err := identifier.Identify(data, context.Background()) require.NoError(t, err) @@ -33,10 +33,10 @@ func TestNNetVersionIdentifier_Identify(t *testing.T) { // 测试版本2识别 t.Run("Identify version 2.0", func(t *testing.T) { data := []byte("NNET") - data = append(data, byte(2)) // version + data = append(data, byte(2)) // version data = append(data, []byte{0, 0, 0, 5}...) // length - data = append(data, []byte("hello")...) // data - data = append(data, []byte{0, 0}...) // checksum (简化) + data = append(data, []byte("hello")...) // data + data = append(data, []byte{0, 0}...) // checksum (简化) version, err := identifier.Identify(data, context.Background()) require.NoError(t, err) @@ -65,7 +65,7 @@ func TestNNetVersionIdentifier_Identify(t *testing.T) { // 测试未知版本 t.Run("Unknown version", func(t *testing.T) { data := []byte("NNET") - data = append(data, byte(99)) // unknown version + data = append(data, byte(99)) // unknown version data = append(data, []byte{0, 0, 0, 5}...) // length version, err := identifier.Identify(data, context.Background()) @@ -79,13 +79,12 @@ func TestNNetVersionIdentifier_DefaultMapping(t *testing.T) { identifier := NewNNetVersionIdentifier(nil) data := []byte("NNET") - data = append(data, byte(1)) // version + data = append(data, byte(1)) // version data = append(data, []byte{0, 0, 0, 5}...) // length - data = append(data, []byte("hello")...) // data - data = append(data, []byte{0, 0}...) // checksum (简化) + data = append(data, []byte("hello")...) // data + data = append(data, []byte{0, 0}...) // checksum (简化) version, err := identifier.Identify(data, context.Background()) require.NoError(t, err) assert.Equal(t, "1.0", version) } - diff --git a/internal/protocol/version/identifier.go b/internal/protocol/version/identifier.go index 8382fda..87973d8 100644 --- a/internal/protocol/version/identifier.go +++ b/internal/protocol/version/identifier.go @@ -3,7 +3,7 @@ package version import ( "context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // HeaderIdentifier 基于帧头的版本识别器 @@ -88,4 +88,3 @@ func (i *CustomIdentifier) Identify(data []byte, ctx context.Context) (string, e } return i.fn(data, ctx) } - diff --git a/internal/request/frame_header.go b/internal/request/frame_header.go index 5bd1dde..6d0dd52 100644 --- a/internal/request/frame_header.go +++ b/internal/request/frame_header.go @@ -1,8 +1,8 @@ package request import ( - internalprotocol "github.com/noahlann/nnet/internal/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" ) // newFrameHeader 创建新的协议帧头(内部使用) @@ -10,4 +10,3 @@ import ( func newFrameHeader() requestpkg.FrameHeader { return internalprotocol.NewFrameHeader() } - diff --git a/internal/request/frame_header_test.go b/internal/request/frame_header_test.go index ce554a2..46b91ba 100644 --- a/internal/request/frame_header_test.go +++ b/internal/request/frame_header_test.go @@ -39,4 +39,3 @@ func TestFrameHeaderConcurrent(t *testing.T) { <-done } } - diff --git a/internal/request/internal.go b/internal/request/internal.go index 93d11b2..ce34ff3 100644 --- a/internal/request/internal.go +++ b/internal/request/internal.go @@ -1,7 +1,7 @@ package request import ( - requestpkg "github.com/noahlann/nnet/pkg/request" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" ) // RequestSetter 内部使用的Request设置接口 @@ -25,4 +25,3 @@ func AsRequestSetter(req requestpkg.Request) RequestSetter { func NewFrameHeader() requestpkg.FrameHeader { return newFrameHeader() } - diff --git a/internal/request/request.go b/internal/request/request.go index a4d1637..b57dd09 100644 --- a/internal/request/request.go +++ b/internal/request/request.go @@ -4,8 +4,8 @@ import ( "encoding/json" "fmt" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" ) // requestImpl 请求实现 diff --git a/internal/request/request_test.go b/internal/request/request_test.go index a3590d5..aba9ef8 100644 --- a/internal/request/request_test.go +++ b/internal/request/request_test.go @@ -4,8 +4,8 @@ import ( "context" "testing" - "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -180,4 +180,3 @@ func (m *mockProtocol) Handle(ctx context.Context, data []byte) ([]byte, error) func (m *mockProtocol) Unpacker() unpackerpkg.Unpacker { return nil } - diff --git a/internal/response/frame_header.go b/internal/response/frame_header.go index 9ff9947..5ef4400 100644 --- a/internal/response/frame_header.go +++ b/internal/response/frame_header.go @@ -1,8 +1,8 @@ package response import ( - internalprotocol "github.com/noahlann/nnet/internal/protocol" - responsepkg "github.com/noahlann/nnet/pkg/response" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol" + responsepkg "git.noahlan.cn/noahlan/nnet/v2/pkg/response" ) // NewFrameHeader 创建新的协议帧头 @@ -10,4 +10,3 @@ import ( func NewFrameHeader() responsepkg.FrameHeader { return internalprotocol.NewFrameHeader() } - diff --git a/internal/response/frame_header_test.go b/internal/response/frame_header_test.go index 2c3c6f9..12fc934 100644 --- a/internal/response/frame_header_test.go +++ b/internal/response/frame_header_test.go @@ -73,4 +73,3 @@ func TestFrameHeaderEncodeDecode(t *testing.T) { err = emptyHeader2.Decode([]byte("{}")) assert.NoError(t, err, "Decode empty JSON should not return error") } - diff --git a/internal/response/header_converter.go b/internal/response/header_converter.go index ac96393..baafafa 100644 --- a/internal/response/header_converter.go +++ b/internal/response/header_converter.go @@ -1,8 +1,8 @@ package response import ( - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - responsepkg "github.com/noahlann/nnet/pkg/response" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + responsepkg "git.noahlan.cn/noahlan/nnet/v2/pkg/response" ) // convertResponseHeaderToProtocolHeader 将响应帧头转换为协议帧头 @@ -11,13 +11,13 @@ func convertResponseHeaderToProtocolHeader(responseHeader responsepkg.FrameHeade if responseHeader == nil { return nil } - + // 默认直接返回引用(避免拷贝开销) // response.FrameHeader和protocol.FrameHeader现在是同一个接口 if protocolHeader, ok := responseHeader.(protocolpkg.FrameHeader); ok { return protocolHeader } - + // 如果类型不匹配(理论上不应该发生),返回nil return nil } @@ -28,13 +28,12 @@ func convertResponseHeaderToProtocolHeaderWithClone(responseHeader responsepkg.F if responseHeader == nil { return nil } - + // 使用Clone方法进行深拷贝 if protocolHeader, ok := responseHeader.(protocolpkg.FrameHeader); ok { return protocolHeader.Clone() } - + // 如果类型不匹配(理论上不应该发生),返回nil return nil } - diff --git a/internal/response/response.go b/internal/response/response.go index 88105b3..67c6d84 100644 --- a/internal/response/response.go +++ b/internal/response/response.go @@ -3,10 +3,10 @@ package response import ( "fmt" - "github.com/noahlann/nnet/pkg/codec" - ctxpkg "github.com/noahlann/nnet/pkg/context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - responsepkg "github.com/noahlann/nnet/pkg/response" + "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + responsepkg "git.noahlan.cn/noahlan/nnet/v2/pkg/response" ) // connectionWriter 连接写入接口(避免循环依赖) diff --git a/internal/response/response_test.go b/internal/response/response_test.go index a549fb0..4af051f 100644 --- a/internal/response/response_test.go +++ b/internal/response/response_test.go @@ -5,9 +5,9 @@ import ( "errors" "testing" - codecimpl "github.com/noahlann/nnet/internal/codec" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + codecimpl "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -138,7 +138,7 @@ func TestResponseWriteWithProtocol(t *testing.T) { // mockConnectionWriter 模拟连接写入器 type mockConnectionWriter struct { - written [][]byte + written [][]byte writeError error } @@ -180,4 +180,3 @@ func (m *mockProtocol) Handle(ctx context.Context, data []byte) ([]byte, error) func (m *mockProtocol) Unpacker() unpackerpkg.Unpacker { return nil } - diff --git a/internal/router/group.go b/internal/router/group.go index f6d5085..8356bee 100644 --- a/internal/router/group.go +++ b/internal/router/group.go @@ -1,8 +1,8 @@ package router import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // routerGroup 路由分组实现 diff --git a/internal/router/group_test.go b/internal/router/group_test.go index 9c32400..7e52cb4 100644 --- a/internal/router/group_test.go +++ b/internal/router/group_test.go @@ -3,27 +3,27 @@ package router import ( "testing" - ctxpkg "github.com/noahlann/nnet/pkg/context" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" "github.com/stretchr/testify/assert" ) func TestRouterGroup(t *testing.T) { router := NewRouter() - + group := router.Group() assert.NotNil(t, group, "Expected group to be non-nil") } func TestRouterGroupWithPrefix(t *testing.T) { router := NewRouter() - + group := router.Group().(*routerGroup) group.prefix = "/api" - + handler := func(ctx ctxpkg.Context) error { return nil } - + route := group.RegisterString("/test", handler) assert.NotNil(t, route, "Expected route to be non-nil") // 测试前缀 @@ -32,23 +32,22 @@ func TestRouterGroupWithPrefix(t *testing.T) { func TestRouterGroupWithMiddleware(t *testing.T) { router := NewRouter() - + middleware := func(ctx ctxpkg.Context) error { return nil } - + group := router.Group() group.Use(middleware) - + handler := func(ctx ctxpkg.Context) error { return nil } - + route := group.RegisterString("/test", handler) assert.NotNil(t, route, "Expected route to be non-nil") - + // 测试route的Handler方法(应该包含middleware) routeHandler := route.Handler() assert.NotNil(t, routeHandler, "Expected handler to be non-nil") } - diff --git a/internal/router/matcher/composite_test.go b/internal/router/matcher/composite_test.go index c46f4db..d84cb9c 100644 --- a/internal/router/matcher/composite_test.go +++ b/internal/router/matcher/composite_test.go @@ -3,8 +3,8 @@ package matcher import ( "testing" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" ) diff --git a/internal/router/matcher/frame_data.go b/internal/router/matcher/frame_data.go index d0bc898..619f1ee 100644 --- a/internal/router/matcher/frame_data.go +++ b/internal/router/matcher/frame_data.go @@ -5,8 +5,8 @@ import ( "reflect" "strings" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // FrameDataMatcher 帧数据匹配器 diff --git a/internal/router/matcher/frame_data_ops_test.go b/internal/router/matcher/frame_data_ops_test.go index e8b3c07..5fc6e43 100644 --- a/internal/router/matcher/frame_data_ops_test.go +++ b/internal/router/matcher/frame_data_ops_test.go @@ -4,11 +4,11 @@ import ( "encoding/json" "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" ) diff --git a/internal/router/matcher/frame_data_test.go b/internal/router/matcher/frame_data_test.go index 3a5d329..e49ee26 100644 --- a/internal/router/matcher/frame_data_test.go +++ b/internal/router/matcher/frame_data_test.go @@ -4,11 +4,11 @@ import ( "encoding/json" "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" ) diff --git a/internal/router/matcher/frame_header.go b/internal/router/matcher/frame_header.go index 53a543b..7cadd27 100644 --- a/internal/router/matcher/frame_header.go +++ b/internal/router/matcher/frame_header.go @@ -3,8 +3,8 @@ package matcher import ( "reflect" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // FrameHeaderMatcher 帧头匹配器 diff --git a/internal/router/matcher/frame_header_test.go b/internal/router/matcher/frame_header_test.go index 74dde0c..d8f880c 100644 --- a/internal/router/matcher/frame_header_test.go +++ b/internal/router/matcher/frame_header_test.go @@ -3,11 +3,11 @@ package matcher import ( "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" ) diff --git a/internal/router/matcher/string_test.go b/internal/router/matcher/string_test.go index 84e7b1b..d433550 100644 --- a/internal/router/matcher/string_test.go +++ b/internal/router/matcher/string_test.go @@ -3,11 +3,11 @@ package matcher import ( "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" ) diff --git a/internal/router/route.go b/internal/router/route.go index 4d85c3a..f6b42ee 100644 --- a/internal/router/route.go +++ b/internal/router/route.go @@ -3,8 +3,8 @@ package router import ( "reflect" - ctxpkg "github.com/noahlann/nnet/pkg/context" - "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // route 路由实现 diff --git a/internal/router/route_test.go b/internal/router/route_test.go index 47fd2ad..a683997 100644 --- a/internal/router/route_test.go +++ b/internal/router/route_test.go @@ -3,14 +3,14 @@ package router import ( "testing" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" ) func TestRoute(t *testing.T) { router := NewRouter() - + handler := func(ctx ctxpkg.Context) error { return nil } @@ -23,42 +23,42 @@ func TestRoute(t *testing.T) { func TestRouteWithOptions(t *testing.T) { router := NewRouter() - + type RequestBody struct { Name string `json:"name"` } - + handler := func(ctx ctxpkg.Context) error { return nil } route := router.RegisterString("/test", handler, routerpkg.WithRequestType(&RequestBody{})) assert.NotNil(t, route, "Expected route to be non-nil") - + requestType := route.RequestType() assert.NotNil(t, requestType, "Expected request type to be set") } func TestRouteMiddleware(t *testing.T) { router := NewRouter() - + handler := func(ctx ctxpkg.Context) error { return nil } route := router.RegisterString("/test", handler) assert.NotNil(t, route, "Expected route to be non-nil") - + // 测试route的Handler方法 routeHandler := route.Handler() assert.NotNil(t, routeHandler, "Expected handler to be non-nil") - + // 测试Use方法添加中间件 middleware := func(ctx ctxpkg.Context) error { return nil } route.Use(middleware) - + // 验证中间件已添加 routeHandler2 := route.Handler() assert.NotNil(t, routeHandler2, "Expected handler to be non-nil") @@ -66,7 +66,7 @@ func TestRouteMiddleware(t *testing.T) { func TestRouteCodec(t *testing.T) { router := NewRouter() - + handler := func(ctx ctxpkg.Context) error { return nil } @@ -75,4 +75,3 @@ func TestRouteCodec(t *testing.T) { assert.NotNil(t, route) assert.Equal(t, "json", route.CodecName(), "Expected codec name to be json") } - diff --git a/internal/router/router.go b/internal/router/router.go index 8158adb..97f60b0 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -4,10 +4,10 @@ import ( "sort" "sync" - "github.com/noahlann/nnet/internal/router/matcher" - ctxpkg "github.com/noahlann/nnet/pkg/context" - "github.com/noahlann/nnet/pkg/errors" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/router/matcher" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // routerImpl 路由器实现 diff --git a/internal/router/router_test.go b/internal/router/router_test.go index 5a9300f..faac611 100644 --- a/internal/router/router_test.go +++ b/internal/router/router_test.go @@ -4,11 +4,11 @@ import ( "context" "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/server/codec_protocol_init.go b/internal/server/codec_protocol_init.go index e165556..fbc64af 100644 --- a/internal/server/codec_protocol_init.go +++ b/internal/server/codec_protocol_init.go @@ -1,12 +1,12 @@ package server import ( - "github.com/noahlann/nnet/internal/codec" - internalprotocol "github.com/noahlann/nnet/internal/protocol" - "github.com/noahlann/nnet/internal/protocol/nnet" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/config" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol" + "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // initCodecRegistry 初始化编解码器注册表 @@ -36,4 +36,3 @@ func initProtocolManager() protocolpkg.Manager { return protocolManager } - diff --git a/internal/server/codec_protocol_init_test.go b/internal/server/codec_protocol_init_test.go index 992eb3e..f9159b9 100644 --- a/internal/server/codec_protocol_init_test.go +++ b/internal/server/codec_protocol_init_test.go @@ -3,7 +3,7 @@ package server import ( "testing" - "github.com/noahlann/nnet/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -49,5 +49,3 @@ func TestIsProtocolEncodeEnabled(t *testing.T) { helper = newServerConfigHelper(cfg) assert.True(t, helper.IsProtocolEncodeEnabled()) } - - diff --git a/internal/server/codec_resolver_test.go b/internal/server/codec_resolver_test.go index 90f0545..9a74ca5 100644 --- a/internal/server/codec_resolver_test.go +++ b/internal/server/codec_resolver_test.go @@ -5,12 +5,12 @@ import ( "testing" "time" - "github.com/noahlann/nnet/internal/codec" - codecpkg "github.com/noahlann/nnet/pkg/codec" - ctxpkg "github.com/noahlann/nnet/pkg/context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" - responsepkg "github.com/noahlann/nnet/pkg/response" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" + responsepkg "git.noahlan.cn/noahlan/nnet/v2/pkg/response" ) // --- minimal fakes for Context/Request/Header diff --git a/internal/server/connection_adapter.go b/internal/server/connection_adapter.go index 5a14c5e..dcd6e14 100644 --- a/internal/server/connection_adapter.go +++ b/internal/server/connection_adapter.go @@ -1,8 +1,8 @@ package server import ( - "github.com/noahlann/nnet/internal/connection" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // connectionAdapter 连接适配器,将ConnectionInterface适配为Context.Connection @@ -39,4 +39,3 @@ func (a *connectionAdapter) Close() error { func toContextConnection(connInterface connection.ConnectionInterface) ctxpkg.Connection { return &connectionAdapter{connInterface: connInterface} } - diff --git a/internal/server/connection_data.go b/internal/server/connection_data.go index ef635d2..f37db8e 100644 --- a/internal/server/connection_data.go +++ b/internal/server/connection_data.go @@ -1,17 +1,17 @@ package server import ( - "github.com/noahlann/nnet/internal/connection" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/panjf2000/gnet/v2" ) // connectionData 连接数据(存储在gnet.Conn.Context()中) type connectionData struct { - conn connection.ConnectionInterface - unpacker unpackerpkg.Unpacker - protocol protocolpkg.Protocol + conn connection.ConnectionInterface + unpacker unpackerpkg.Unpacker + protocol protocolpkg.Protocol protocolVersion string // 已识别的协议版本(如果为空,表示还未识别) } @@ -29,4 +29,3 @@ func getConnectionData(c gnet.Conn) *connectionData { func setConnectionData(c gnet.Conn, data *connectionData) { c.SetContext(data) } - diff --git a/internal/server/gnet_server.go b/internal/server/gnet_server.go index b75ce5c..c45b8d8 100644 --- a/internal/server/gnet_server.go +++ b/internal/server/gnet_server.go @@ -6,14 +6,14 @@ import ( "sync" "time" - "github.com/noahlann/nnet/internal/connection" - "github.com/noahlann/nnet/internal/logger" - internalrouter "github.com/noahlann/nnet/internal/router" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/config" - "github.com/noahlann/nnet/pkg/errors" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/internal/logger" + internalrouter "git.noahlan.cn/noahlan/nnet/v2/internal/router" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/panjf2000/gnet/v2" ) @@ -66,6 +66,7 @@ func (p TransportProtocol) IsDatagram() bool { // gnetServer 统一的gnet服务器基类(支持TCP、UDP、Unix) const startupWaitTimeout = 5 * time.Second +const stopWaitTimeout = 5 * time.Second type gnetServer struct { config *config.Config @@ -204,6 +205,10 @@ func (s *gnetServer) Start() error { s.bootCh = make(chan error, 1) s.bootOnce = &sync.Once{} s.eventHandler.setBootSignal(s.bootCh, s.bootOnce) + s.stopMu.Lock() + s.stopCh = make(chan struct{}) + s.stopped = false + s.stopMu.Unlock() // 构建gnet选项 options := []gnet.Option{ @@ -281,19 +286,34 @@ func (s *gnetServer) Stop() error { // 取消context s.cancel() + var stopErr error + if s.eventHandler != nil { + if engine, engineSet := s.eventHandler.getEngine(); engineSet { + stopCtx, cancel := context.WithTimeout(context.Background(), stopWaitTimeout) + stopErr = engine.Stop(stopCtx) + cancel() + if stopErr != nil { + s.logger.Warn("%s engine stop error: %v", s.protocol.String(), stopErr) + } + } + } + // 等待服务器停止 select { case <-s.stopCh: s.logger.Info("%s server stopped", s.protocol.String()) - case <-time.After(5 * time.Second): + case <-time.After(stopWaitTimeout): s.logger.Warn("%s server stop timeout", s.protocol.String()) + if stopErr == nil { + stopErr = errors.New("server stop timeout") + } } s.mu.Lock() s.started = false s.mu.Unlock() - return nil + return stopErr } // formatGnetAddr 格式化gnet地址 diff --git a/internal/server/helpers.go b/internal/server/helpers.go index c61ba59..88caba2 100644 --- a/internal/server/helpers.go +++ b/internal/server/helpers.go @@ -5,14 +5,14 @@ import ( "crypto/tls" "time" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/config" - "github.com/noahlann/nnet/pkg/errors" - ctxpkg "github.com/noahlann/nnet/pkg/context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // ============================================================================ @@ -158,7 +158,7 @@ func processDataWithUnpacker(data []byte, unpacker unpackerpkg.Unpacker) ([][]by // 如果没有完整消息,等待更多数据 if len(messages) == 0 { - return nil, 0, true, nil + return nil, consumed, len(remaining) > 0, nil } // 使用unpacker返回的consumed值作为已处理的数据量(100%准确) @@ -193,4 +193,3 @@ func loadTLSConfig(cfg *config.TLSConfig) (*tls.Config, error) { Certificates: []tls.Certificate{cert}, }, nil } - diff --git a/internal/server/helpers_test.go b/internal/server/helpers_test.go new file mode 100644 index 0000000..01a0692 --- /dev/null +++ b/internal/server/helpers_test.go @@ -0,0 +1,48 @@ +package server + +import ( + "testing" + + internalunpacker "git.noahlan.cn/noahlan/nnet/v2/internal/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" +) + +func TestProcessDataWithUnpackerConsumesIncompleteInput(t *testing.T) { + unpacker := internalunpacker.NewLengthFieldUnpacker(unpackerpkg.LengthFieldUnpacker{ + LengthFieldOffset: 0, + LengthFieldLength: 1, + LengthAdjustment: 0, + InitialBytesToStrip: 0, + }) + + messages, consumed, hasRemaining, err := processDataWithUnpacker([]byte{3, 'a'}, unpacker) + if err != nil { + t.Fatalf("processDataWithUnpacker error: %v", err) + } + if len(messages) != 0 { + t.Fatalf("messages=%d, want 0", len(messages)) + } + if consumed != 2 { + t.Fatalf("consumed=%d, want 2", consumed) + } + if !hasRemaining { + t.Fatal("hasRemaining=false, want true") + } + + messages, consumed, hasRemaining, err = processDataWithUnpacker([]byte{'b', 'c'}, unpacker) + if err != nil { + t.Fatalf("processDataWithUnpacker second error: %v", err) + } + if consumed != 2 { + t.Fatalf("second consumed=%d, want 2", consumed) + } + if hasRemaining { + t.Fatal("second hasRemaining=true, want false") + } + if len(messages) != 1 { + t.Fatalf("second messages=%d, want 1", len(messages)) + } + if string(messages[0]) != "\x03abc" { + t.Fatalf("message=%q, want %q", string(messages[0]), "\x03abc") + } +} diff --git a/internal/server/message_handler.go b/internal/server/message_handler.go index 7d93911..23d324f 100644 --- a/internal/server/message_handler.go +++ b/internal/server/message_handler.go @@ -3,15 +3,16 @@ package server import ( "context" "fmt" + "strconv" "time" - "github.com/noahlann/nnet/internal/logger" - internalrequest "github.com/noahlann/nnet/internal/request" - codecpkg "github.com/noahlann/nnet/pkg/codec" - ctxpkg "github.com/noahlann/nnet/pkg/context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/logger" + internalrequest "git.noahlan.cn/noahlan/nnet/v2/internal/request" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // messageHandler 消息处理器配置 @@ -98,7 +99,7 @@ func (mh *messageHandler) handleMessage( route, handler, err := mh.router.Match(matchInput, ctx) if err != nil { - mh.logger.Warn("Route not found for: %s", string(matchInput.Raw)) + mh.logger.Warn("Route not found for payload length=%d preview=%s", len(matchInput.Raw), previewBytes(matchInput.Raw, 128)) return fmt.Errorf("route not found: %w", err) } @@ -119,6 +120,16 @@ func (mh *messageHandler) handleMessage( return nil } +func previewBytes(data []byte, limit int) string { + if limit <= 0 || len(data) == 0 { + return strconv.Quote("") + } + if len(data) <= limit { + return strconv.Quote(string(data)) + } + return strconv.Quote(string(data[:limit])) + "..." +} + // preDecodeRequestBody 预解码请求体(用于路由匹配) func (mh *messageHandler) preDecodeRequestBody(req requestpkg.Request, bodyBytes []byte, codec codecpkg.Codec) error { reqImpl := internalrequest.AsRequestSetter(req) diff --git a/internal/server/pool.go b/internal/server/pool.go index 3096f98..bb9ab34 100644 --- a/internal/server/pool.go +++ b/internal/server/pool.go @@ -80,4 +80,3 @@ var responsePool = sync.Pool{ return nil }, } - diff --git a/internal/server/protocol_header_parser.go b/internal/server/protocol_header_parser.go index 53c074d..bc2bf65 100644 --- a/internal/server/protocol_header_parser.go +++ b/internal/server/protocol_header_parser.go @@ -3,9 +3,9 @@ package server import ( "fmt" - internalrequest "github.com/noahlann/nnet/internal/request" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" + internalrequest "git.noahlan.cn/noahlan/nnet/v2/internal/request" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" ) // parseProtocolHeader 解析协议帧头(用于路由匹配) diff --git a/internal/server/protocol_header_parser_test.go b/internal/server/protocol_header_parser_test.go index 1b48386..6a77b18 100644 --- a/internal/server/protocol_header_parser_test.go +++ b/internal/server/protocol_header_parser_test.go @@ -3,10 +3,10 @@ package server import ( "testing" - nnetproto "github.com/noahlann/nnet/internal/protocol/nnet" - internalrequest "github.com/noahlann/nnet/internal/request" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" + nnetproto "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + internalrequest "git.noahlan.cn/noahlan/nnet/v2/internal/request" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/server/request_body_parser.go b/internal/server/request_body_parser.go index 618365a..f89b8d7 100644 --- a/internal/server/request_body_parser.go +++ b/internal/server/request_body_parser.go @@ -4,11 +4,11 @@ import ( "fmt" "reflect" - internalrequest "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/pkg/codec" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - requestpkg "github.com/noahlann/nnet/pkg/request" - "github.com/noahlann/nnet/pkg/router" + internalrequest "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + requestpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/request" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // parseRequestBodyWithCodec 使用指定codec解码请求体(优先于基于路由决策) diff --git a/internal/server/request_body_parser_test.go b/internal/server/request_body_parser_test.go index 31a92dd..e0f750c 100644 --- a/internal/server/request_body_parser_test.go +++ b/internal/server/request_body_parser_test.go @@ -5,10 +5,10 @@ import ( "reflect" "testing" - "github.com/noahlann/nnet/internal/codec" - internalrequest "github.com/noahlann/nnet/internal/request" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + internalrequest "git.noahlan.cn/noahlan/nnet/v2/internal/request" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -71,5 +71,3 @@ func TestParseRequestBody_InvalidRequestType(t *testing.T) { err := parseRequestBodyWithCodec(fr, rt, nil, jsonCodec) require.Error(t, err) } - - diff --git a/internal/server/serial_server.go b/internal/server/serial_server.go index c2c2e17..bdcc820 100644 --- a/internal/server/serial_server.go +++ b/internal/server/serial_server.go @@ -7,16 +7,16 @@ import ( "net/http" "sync" - "github.com/noahlann/nnet/internal/connection" - "github.com/noahlann/nnet/internal/logger" - internalrouter "github.com/noahlann/nnet/internal/router" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/config" - "github.com/noahlann/nnet/pkg/errors" - "github.com/noahlann/nnet/pkg/health" - metricspkg "github.com/noahlann/nnet/pkg/metrics" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/internal/logger" + internalrouter "git.noahlan.cn/noahlan/nnet/v2/internal/router" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/health" + metricspkg "git.noahlan.cn/noahlan/nnet/v2/pkg/metrics" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "go.bug.st/serial" ) diff --git a/internal/server/server.go b/internal/server/server.go index 6c1c2ea..84718a0 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -7,16 +7,16 @@ import ( "net/http" "time" - "github.com/noahlann/nnet/internal/connection" - "github.com/noahlann/nnet/internal/logger" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/config" - "github.com/noahlann/nnet/pkg/errors" - "github.com/noahlann/nnet/pkg/health" - "github.com/noahlann/nnet/pkg/lifecycle" - metricspkg "github.com/noahlann/nnet/pkg/metrics" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/internal/logger" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/health" + "git.noahlan.cn/noahlan/nnet/v2/pkg/lifecycle" + metricspkg "git.noahlan.cn/noahlan/nnet/v2/pkg/metrics" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // Server TCP服务器(使用统一的gnetServer) @@ -100,16 +100,9 @@ func (s *Server) Stop() error { s.gnetServer.logger.Info("Starting graceful shutdown...") - // 1. 停止接受新连接(通过停止gnet引擎) - if s.gnetServer.eventHandler != nil { - engine, engineSet := s.gnetServer.eventHandler.getEngine() - if engineSet { - // gnet v2中,Engine接口可能没有Stop方法 - // 我们需要通过其他方式停止服务器 - // 暂时使用context取消来停止服务器 - s.gnetServer.logger.Info("Stopping gnet engine...") - _ = engine // 暂时不使用,等待gnet API支持 - } + // 1. 停止gnet引擎,关闭监听和现有连接,使连接管理器能收到OnClose并完成清理。 + if err := s.gnetServer.Stop(); err != nil { + s.gnetServer.logger.Error("Failed to stop gnet server: %v", err) } // 2. 获取关闭超时时间 @@ -154,11 +147,6 @@ done: } } - // 5. 使用gnetServer的Stop方法 - if err := s.gnetServer.Stop(); err != nil { - s.gnetServer.logger.Error("Failed to stop gnet server: %v", err) - } - s.gnetServer.logger.Info("Server stopped gracefully") return nil } diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 3745a3a..f5e0b9f 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -6,11 +6,11 @@ import ( "net/http/httptest" "testing" - "github.com/noahlann/nnet/pkg/config" - ctxpkg "github.com/noahlann/nnet/pkg/context" - "github.com/noahlann/nnet/pkg/errors" - "github.com/noahlann/nnet/pkg/health" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/health" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/server/udp_server.go b/internal/server/udp_server.go index aa2a385..3ebab50 100644 --- a/internal/server/udp_server.go +++ b/internal/server/udp_server.go @@ -1,7 +1,7 @@ package server import ( - "github.com/noahlann/nnet/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" ) // UDPServer UDP服务器(使用统一的gnetServer) diff --git a/internal/server/unified_event_handler.go b/internal/server/unified_event_handler.go index f072a84..4a40dfa 100644 --- a/internal/server/unified_event_handler.go +++ b/internal/server/unified_event_handler.go @@ -7,13 +7,13 @@ import ( "sync" "time" - "github.com/noahlann/nnet/internal/connection" - "github.com/noahlann/nnet/internal/logger" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/lifecycle" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - routerpkg "github.com/noahlann/nnet/pkg/router" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/internal/logger" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/lifecycle" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/panjf2000/gnet/v2" ) @@ -155,7 +155,7 @@ func (h *unifiedEventHandler) OnClose(c gnet.Conn, err error) (action gnet.Actio } connID := connData.conn.ID() - + // 执行连接生命周期钩子 for _, hook := range h.connLifecycleHooks { if hookErr := hook.OnClose(connID, err); hookErr != nil { @@ -297,13 +297,15 @@ func (h *unifiedEventHandler) handleTraffic(c gnet.Conn, udpPacket []byte) gnet. return gnet.Close } if len(messages) == 0 { - // 没有完整消息,等待更多数据(gnet会自动保留数据在缓冲区中) + if totalProcessed > 0 { + c.Discard(totalProcessed) + } return gnet.None } } else { // UDP数据包已经是完整的,不需要拆包 messages = [][]byte{data} - totalProcessed = len(data) + totalProcessed = len(data) } // 处理每个完整的消息 @@ -367,7 +369,7 @@ func (h *unifiedEventHandler) identifyProtocolVersion(connData *connectionData, case string: version = v case byte: - version = fmt.Sprintf("%d.0", v) + version = fmt.Sprintf("%d.0", v) default: version = fmt.Sprintf("%v", v) } diff --git a/internal/server/unix_server.go b/internal/server/unix_server.go index a90c615..648d8c6 100644 --- a/internal/server/unix_server.go +++ b/internal/server/unix_server.go @@ -1,7 +1,7 @@ package server import ( - "github.com/noahlann/nnet/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" ) // UnixServer Unix Domain Socket服务器(使用统一的gnetServer) diff --git a/internal/server/unpacker_manager.go b/internal/server/unpacker_manager.go index 34c1843..741645c 100644 --- a/internal/server/unpacker_manager.go +++ b/internal/server/unpacker_manager.go @@ -3,8 +3,8 @@ package server import ( "sync" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // unpackerInfo 拆包器信息 @@ -68,4 +68,3 @@ func (m *unpackerManager) removeUnpacker(connID string) { delete(m.unpackers, connID) m.mu.Unlock() } - diff --git a/internal/server/unpacker_manager_test.go b/internal/server/unpacker_manager_test.go index 02aaa09..e28d773 100644 --- a/internal/server/unpacker_manager_test.go +++ b/internal/server/unpacker_manager_test.go @@ -3,7 +3,7 @@ package server import ( "testing" - nnetproto "github.com/noahlann/nnet/internal/protocol/nnet" + nnetproto "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" "github.com/stretchr/testify/assert" ) @@ -21,8 +21,7 @@ func TestUnpackerManager_GetOrCreate_CachePerConn(t *testing.T) { // Different connection should get another instance u3 := m.getOrCreateUnpacker("c2", proto) assert.NotNil(t, u3) - // NNet protocol returns a singleton unpacker instance; expect same pointer - assert.Equal(t, u1, u3) + assert.NotSame(t, u1, u3) } func TestUnpackerManager_GetOrCreate_NoProtocol(t *testing.T) { @@ -42,8 +41,5 @@ func TestUnpackerManager_Remove(t *testing.T) { // After removal, a new instance should be created u2 := m.getOrCreateUnpacker("c1", proto) assert.NotNil(t, u2) - // NNet protocol returns a singleton unpacker instance; expect same pointer again - assert.Equal(t, u1, u2) + assert.NotSame(t, u1, u2) } - - diff --git a/internal/server/websocket_server.go b/internal/server/websocket_server.go index 981d691..24f6255 100644 --- a/internal/server/websocket_server.go +++ b/internal/server/websocket_server.go @@ -1,46 +1,71 @@ package server import ( - "github.com/noahlann/nnet/pkg/config" + "context" + "fmt" + "net" + "net/http" + "strings" + "sync" + "time" + + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + pkgerrors "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" + "github.com/gorilla/websocket" ) -// WebSocketServer WebSocket服务器(基于gnetServer,支持ws和wss) +// WebSocketServer is a real WebSocket transport server. It reuses the same +// router, codec, protocol and connection manager stack as the TCP server. type WebSocketServer struct { *Server - isWSS bool // 是否为WSS(WebSocket Secure) + + addr string + isWSS bool + httpServer *http.Server + listener net.Listener + upgrader websocket.Upgrader + ctx context.Context + cancel context.CancelFunc + mu sync.RWMutex + started bool + msgHandler *messageHandler + handlerTimeout time.Duration } // NewWebSocketServer 创建WebSocket服务器(ws或wss) func NewWebSocketServer(cfg *config.Config) (*WebSocketServer, error) { if cfg == nil { cfg = config.DefaultConfig() + cfg.Addr = "ws://:6995" } - // 检测是否为WSS - isWSS := false - addr := cfg.Addr - if len(addr) >= 6 && addr[:6] == "wss://" { - isWSS = true - // 转换为TCP地址格式 - cfg.Addr = "tcp://" + addr[6:] - // 启用TLS - if !cfg.TLSEnabled { - cfg.TLSEnabled = true - } - } else if len(addr) >= 5 && addr[:5] == "ws://" { - // 转换为TCP地址格式 - cfg.Addr = "tcp://" + addr[5:] + addr, isWSS := normalizeWebSocketAddr(cfg.Addr) + if isWSS && !cfg.TLSEnabled { + cfg.TLSEnabled = true } - // 创建TCP服务器(WebSocket基于TCP) - server, err := NewServer(cfg) + gnetSrv, err := newGnetServer(cfg, ProtocolTCP) if err != nil { return nil, err } + base := newServerWithGnet(gnetSrv, cfg) + configHelper := newServerConfigHelper(cfg) + ctx, cancel := context.WithCancel(context.Background()) return &WebSocketServer{ - Server: server, + Server: base, + addr: addr, isWSS: isWSS, + upgrader: websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, + }, + ctx: ctx, + cancel: cancel, + msgHandler: gnetSrv.eventHandler.msgHandler, + handlerTimeout: configHelper.HandlerTimeout(), }, nil } @@ -49,25 +74,270 @@ func NewWSSServer(cfg *config.Config) (*WebSocketServer, error) { if cfg == nil { cfg = config.DefaultConfig() } + cfg.TLSEnabled = true - // 确保TLS已启用 - if !cfg.TLSEnabled { - cfg.TLSEnabled = true + addr := cfg.Addr + if strings.HasPrefix(addr, "ws://") { + cfg.Addr = "wss://" + strings.TrimPrefix(addr, "ws://") + } else if !strings.HasPrefix(addr, "wss://") { + cfg.Addr = "wss://" + addr } - // 转换地址格式 - addr := cfg.Addr - if len(addr) >= 5 && addr[:5] == "ws://" { - // 将ws://转换为wss:// - cfg.Addr = "wss://" + addr[5:] - } else if len(addr) < 6 || addr[:6] != "wss://" { - // 如果没有wss://前缀,添加 - if len(addr) > 0 && addr[0] != ':' { - cfg.Addr = "wss://" + addr + return NewWebSocketServer(cfg) +} + +func normalizeWebSocketAddr(addr string) (string, bool) { + switch { + case strings.HasPrefix(addr, "wss://"): + addr = strings.TrimPrefix(addr, "wss://") + return defaultListenAddr(addr), true + case strings.HasPrefix(addr, "ws://"): + addr = strings.TrimPrefix(addr, "ws://") + return defaultListenAddr(addr), false + default: + return defaultListenAddr(addr), false + } +} + +func defaultListenAddr(addr string) string { + if addr == "" || addr == ":" { + return ":6995" + } + return addr +} + +// Start 启动WebSocket服务器。与TCP服务器保持一致:启动监听后返回。 +func (s *WebSocketServer) Start() error { + s.mu.Lock() + if s.started { + s.mu.Unlock() + return pkgerrors.ErrServerAlreadyStarted + } + + s.ctx, s.cancel = context.WithCancel(context.Background()) + mux := http.NewServeMux() + mux.HandleFunc("/", s.handleWebSocket) + httpServer := &http.Server{Handler: mux} + + listener, err := net.Listen("tcp", s.addr) + if err != nil { + s.mu.Unlock() + return fmt.Errorf("failed to listen on %s: %w", s.addr, err) + } + + s.listener = listener + s.httpServer = httpServer + s.started = true + s.gnetServer.mu.Lock() + s.gnetServer.started = true + s.gnetServer.mu.Unlock() + s.mu.Unlock() + + for _, hook := range s.serverLifecycleHooks { + if err := hook.OnInit(); err != nil { + _ = listener.Close() + return pkgerrors.New("failed to execute OnInit hook").WithCause(err) + } + } + for _, hook := range s.serverLifecycleHooks { + if err := hook.OnStart(); err != nil { + s.gnetServer.logger.Error("OnStart hook error: %v", err) + } + } + + s.gnetServer.logger.Info("WebSocket server started on %s", listener.Addr().String()) + go s.serve(listener, httpServer) + return nil +} + +func (s *WebSocketServer) serve(listener net.Listener, httpServer *http.Server) { + var err error + if s.isWSS { + if s.config.TLS == nil || s.config.TLS.CertFile == "" || s.config.TLS.KeyFile == "" { + err = pkgerrors.New("wss requires TLS cert and key files") } else { - cfg.Addr = "wss://" + addr + err = httpServer.ServeTLS(listener, s.config.TLS.CertFile, s.config.TLS.KeyFile) } + } else { + err = httpServer.Serve(listener) + } + if err != nil && err != http.ErrServerClosed { + s.gnetServer.logger.Error("WebSocket server error: %v", err) } +} - return NewWebSocketServer(cfg) +// Stop 停止WebSocket服务器。 +func (s *WebSocketServer) Stop() error { + s.mu.Lock() + if !s.started { + s.mu.Unlock() + return pkgerrors.ErrServerNotStarted + } + httpServer := s.httpServer + shutdownTimeout := s.config.ShutdownTimeout + if shutdownTimeout <= 0 { + shutdownTimeout = 30 * time.Second + } + s.started = false + s.gnetServer.mu.Lock() + s.gnetServer.started = false + s.gnetServer.mu.Unlock() + s.cancel() + s.mu.Unlock() + + s.gnetServer.logger.Info("Stopping WebSocket server...") + ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) + defer cancel() + + var stopErr error + if httpServer != nil { + stopErr = httpServer.Shutdown(ctx) + } + + for _, conn := range s.connManager.GetAll() { + _ = conn.Close() + } + + ticker := time.NewTicker(100 * time.Millisecond) + defer ticker.Stop() + timeout := time.NewTimer(shutdownTimeout) + defer timeout.Stop() + +waitConnections: + for { + if s.connManager.Count() == 0 { + break + } + select { + case <-ticker.C: + case <-timeout.C: + s.forceCloseAllConnections() + break waitConnections + } + if s.connManager.Count() == 0 { + break + } + } + + for _, hook := range s.serverLifecycleHooks { + if err := hook.OnStop(); err != nil { + s.gnetServer.logger.Error("OnStop hook error: %v", err) + } + } + + s.gnetServer.logger.Info("WebSocket server stopped") + return stopErr +} + +// Started 检查服务器是否已启动。 +func (s *WebSocketServer) Started() bool { + s.mu.RLock() + defer s.mu.RUnlock() + return s.started +} + +func (s *WebSocketServer) handleWebSocket(w http.ResponseWriter, r *http.Request) { + wsConn, err := s.upgrader.Upgrade(w, r, nil) + if err != nil { + s.gnetServer.logger.Debug("WebSocket upgrade failed: %v", err) + return + } + + conn := connection.NewWebSocketConnection("", wsConn) + connID := conn.ID() + if err := s.connManager.Add(conn); err != nil { + s.gnetServer.logger.Error("Failed to add WebSocket connection: %v", err) + _ = wsConn.Close() + return + } + + for _, hook := range s.connLifecycleHooks { + if err := hook.OnOpen(connID, conn.RemoteAddr()); err != nil { + s.gnetServer.logger.Error("OnOpen hook error: %v", err) + } + } + + protocol := s.resolveApplicationProtocol() + unpacker := newServerProtocolUnpacker(protocol) + ctxConn := toContextConnection(conn) + + defer func() { + for _, hook := range s.connLifecycleHooks { + if err := hook.OnClose(connID, nil); err != nil { + s.gnetServer.logger.Error("OnClose hook error: %v", err) + } + } + _ = s.connManager.Remove(connID) + }() + + for { + select { + case <-s.ctx.Done(): + return + default: + } + + _, data, err := wsConn.ReadMessage() + if err != nil { + return + } + conn.UpdateActive() + + messages, err := splitWebSocketMessages(data, unpacker) + if err != nil { + s.gnetServer.logger.Debug("WebSocket unpack failed: %v", err) + return + } + if len(messages) == 0 { + continue + } + + for _, message := range messages { + s.msgHandler.handleMessageWithContext( + s.ctx, + ctxConn, + message, + protocol, + s.codecRegistry, + s.handlerTimeout, + ) + } + } +} + +func (s *WebSocketServer) resolveApplicationProtocol() protocolpkg.Protocol { + configHelper := newServerConfigHelper(s.config) + if !configHelper.IsProtocolEncodeEnabled() { + return nil + } + protocol, err := s.protocolManager.Get(s.config.ApplicationProtocol, "") + if err != nil { + s.gnetServer.logger.Debug("Failed to resolve WebSocket application protocol: %v", err) + return nil + } + return protocol +} + +type serverUnpackingProtocol interface { + Unpacker() unpackerpkg.Unpacker +} + +func newServerProtocolUnpacker(protocol protocolpkg.Protocol) unpackerpkg.Unpacker { + if protocol == nil { + return nil + } + if provider, ok := protocol.(serverUnpackingProtocol); ok { + return provider.Unpacker() + } + return nil +} + +func splitWebSocketMessages(data []byte, unpacker unpackerpkg.Unpacker) ([][]byte, error) { + if unpacker == nil { + msg := make([]byte, len(data)) + copy(msg, data) + return [][]byte{msg}, nil + } + messages, _, _, err := unpacker.Unpack(data) + return messages, err } diff --git a/internal/session/session.go b/internal/session/session.go index fb68896..9bae5af 100644 --- a/internal/session/session.go +++ b/internal/session/session.go @@ -4,7 +4,7 @@ import ( "sync" "time" - sessionpkg "github.com/noahlann/nnet/pkg/session" + sessionpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/session" ) // SessionImpl Session实现 diff --git a/internal/session/session_test.go b/internal/session/session_test.go index c558478..c0873e2 100644 --- a/internal/session/session_test.go +++ b/internal/session/session_test.go @@ -16,7 +16,7 @@ func TestSession(t *testing.T) { // 测试Set和Get err := session.Set("key1", "value1") require.NoError(t, err, "Expected no error when setting value") - + val, err := session.Get("key1") require.NoError(t, err, "Expected no error when getting value") assert.Equal(t, "value1", val, "Expected value to match") @@ -24,7 +24,7 @@ func TestSession(t *testing.T) { // 测试Delete err = session.Delete("key1") require.NoError(t, err, "Expected no error when deleting value") - + val, err = session.Get("key1") require.NoError(t, err) assert.Nil(t, val, "Expected value to be nil after deletion") @@ -60,20 +60,19 @@ func TestSessionConcurrent(t *testing.T) { func TestSessionClear(t *testing.T) { session := NewSession("session1") - + // 设置一些值 session.Set("key1", "value1") session.Set("key2", "value2") - + // 测试Clear err := session.Clear() require.NoError(t, err, "Expected no error when clearing") - + // 验证值已被清空 val, _ := session.Get("key1") assert.Nil(t, val, "Expected key1 to be nil after clear") - + val, _ = session.Get("key2") assert.Nil(t, val, "Expected key2 to be nil after clear") } - diff --git a/internal/session/storage/file.go b/internal/session/storage/file.go index 6b1d238..40338ff 100644 --- a/internal/session/storage/file.go +++ b/internal/session/storage/file.go @@ -8,8 +8,8 @@ import ( "sync" "time" - sessionimpl "github.com/noahlann/nnet/internal/session" - sessionpkg "github.com/noahlann/nnet/pkg/session" + sessionimpl "git.noahlan.cn/noahlan/nnet/v2/internal/session" + sessionpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/session" ) // FileStorage 文件存储 @@ -238,4 +238,3 @@ func (s *FileStorage) cleanup() { s.Cleanup() } } - diff --git a/internal/session/storage/memory.go b/internal/session/storage/memory.go index 9028f86..93a2902 100644 --- a/internal/session/storage/memory.go +++ b/internal/session/storage/memory.go @@ -4,8 +4,8 @@ import ( "sync" "time" - sessionimpl "github.com/noahlann/nnet/internal/session" - sessionpkg "github.com/noahlann/nnet/pkg/session" + sessionimpl "git.noahlan.cn/noahlan/nnet/v2/internal/session" + sessionpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/session" ) // MemoryStorage 内存存储 diff --git a/internal/session/storage/memory_more_test.go b/internal/session/storage/memory_more_test.go index 405ce26..00f61df 100644 --- a/internal/session/storage/memory_more_test.go +++ b/internal/session/storage/memory_more_test.go @@ -68,5 +68,3 @@ func TestMemoryStorage_CleanupRemovesExpired(t *testing.T) { require.NoError(t, err) assert.Nil(t, got) } - - diff --git a/internal/session/storage/memory_test.go b/internal/session/storage/memory_test.go index 856fd29..1e60e59 100644 --- a/internal/session/storage/memory_test.go +++ b/internal/session/storage/memory_test.go @@ -97,4 +97,3 @@ func TestMemoryStorageExpiration(t *testing.T) { require.NoError(t, err) assert.Nil(t, retrieved, "Expected session to be nil after expiration") } - diff --git a/internal/session/storage/redis.go b/internal/session/storage/redis.go index 065da2f..17cd26e 100644 --- a/internal/session/storage/redis.go +++ b/internal/session/storage/redis.go @@ -6,8 +6,8 @@ import ( "fmt" "time" - sessionimpl "github.com/noahlann/nnet/internal/session" - sessionpkg "github.com/noahlann/nnet/pkg/session" + sessionimpl "git.noahlan.cn/noahlan/nnet/v2/internal/session" + sessionpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/session" ) // RedisStorage Redis存储接口(需要用户提供Redis客户端实现) @@ -137,4 +137,3 @@ func (s *RedisStorage) Cleanup() error { func (s *RedisStorage) getKey(sessionID string) string { return s.prefix + sessionID } - diff --git a/internal/testutil/mock_connection.go b/internal/testutil/mock_connection.go index b29e243..7f9db2f 100644 --- a/internal/testutil/mock_connection.go +++ b/internal/testutil/mock_connection.go @@ -4,7 +4,7 @@ import ( "errors" "time" - "github.com/noahlann/nnet/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" ) // MockConnection is a mock implementation of ConnectionInterface for testing diff --git a/internal/testutil/test_helpers.go b/internal/testutil/test_helpers.go index 0ce5409..b0e6760 100644 --- a/internal/testutil/test_helpers.go +++ b/internal/testutil/test_helpers.go @@ -3,10 +3,10 @@ package testutil import ( "context" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // NewTestContext creates a test context with mock connection, request, and response @@ -36,4 +36,3 @@ func NewTestContextWithData(parentCtx context.Context, connID string, data []byt return ctxpkg.New(parentCtx, conn, req, resp) } - diff --git a/internal/unpacker/delimiter.go b/internal/unpacker/delimiter.go index 500c38c..978b8ee 100644 --- a/internal/unpacker/delimiter.go +++ b/internal/unpacker/delimiter.go @@ -3,7 +3,7 @@ package unpacker import ( "bytes" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // delimiterUnpacker 分隔符拆包器实现 @@ -35,11 +35,8 @@ func NewDelimiterUnpackerWithMaxBuffer(delimiter []byte, maxBufferSize int) unpa // Unpack 拆包 func (u *delimiterUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) { - // 记录调用前的buffer长度(用于计算消耗的数据量) - prevBufferLen := len(u.buffer) - // 检查buffer大小限制 - newSize := prevBufferLen + len(data) + newSize := len(u.buffer) + len(data) if newSize > u.maxBufferSize { return nil, nil, 0, unpackerpkg.NewErrorf("unpacker buffer size exceeded: %d > %d", newSize, u.maxBufferSize) } @@ -65,9 +62,6 @@ func (u *delimiterUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) { u.buffer = append(u.buffer, data...) - // 记录追加数据后的buffer长度(用于计算consumed) - afterAppendLen := len(u.buffer) - var messages [][]byte for { @@ -94,98 +88,8 @@ func (u *delimiterUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) { } } - // 计算本次从输入data中消耗的数据量(100%准确,无误差) - // 原理:consumed = 从输入data中被消耗的数据量(从gnet buffer的角度) - // - // 关键理解: - // - 从gnet buffer的角度,一旦数据被传递给unpacker,就需要被"消耗"(discard) - // - 如果没有完整消息:所有数据都被保存到buffer,但从gnet buffer角度都被消耗,consumed = len(data) - // - 如果有完整消息:部分数据被用来构建消息,部分数据可能被保存到buffer - // 但从gnet buffer的角度,所有数据都被消耗(因为已经处理了消息),consumed = len(data) - // - // 但这样会有问题:如果消息完全来自之前的buffer,输入data完全没有被使用,consumed应该是0 - // - // 更准确的理解: - // - consumed = 从输入data中实际被"使用"的数据量 - // - 如果没有完整消息:所有数据都被保存(等待更多数据),consumed = len(data) - // - 如果有完整消息: - // * 被处理的数据总量 = afterAppendLen - currentBufferLen - // * 如果被处理的数据总量 <= prevBufferLen:所有被处理的数据都来自之前的buffer - // - 输入data没有被处理,但可能被保存到buffer中 - // - consumed = min(bufferIncrease, len(data))(被保存到buffer中的数据,如果来自输入data) - // * 否则:至少部分被处理的数据来自输入data - // - 从输入data中被处理的部分 = processedTotal - prevBufferLen - // - 被保存到buffer中的数据(如果来自输入data)= min(bufferIncrease, len(data) - (processedTotal - prevBufferLen)) - // - consumed = (processedTotal - prevBufferLen) + min(bufferIncrease, len(data) - (processedTotal - prevBufferLen)) - // - 简化:如果bufferIncrease <= len(data) - (processedTotal - prevBufferLen),consumed = len(data) - // - 否则,consumed = (processedTotal - prevBufferLen) + (len(data) - (processedTotal - prevBufferLen)) = len(data) - // - // 最终结论:无论是否有完整消息,consumed都等于len(data)(从gnet buffer的角度) - // 但这是不准确的,因为如果消息完全来自之前的buffer,输入data完全没有被使用 - // - // 最准确的方法: - // - 如果没有完整消息:consumed = len(data) - // - 如果有完整消息: - // * 被处理的数据总量 = afterAppendLen - currentBufferLen - // * 如果被处理的数据总量 <= prevBufferLen: - // - 所有被处理的数据都来自之前的buffer - // - consumed = min(max(0, bufferIncrease), len(data))(被保存到buffer中的数据,如果来自输入data) - // * 否则: - // - 从输入data中被处理的部分 = processedTotal - prevBufferLen - // - 被保存到buffer中的数据(如果来自输入data)= min(max(0, bufferIncrease), len(data) - (processedTotal - prevBufferLen)) - // - consumed = (processedTotal - prevBufferLen) + min(max(0, bufferIncrease), len(data) - (processedTotal - prevBufferLen)) - // - // 简化实现(100%准确): - currentBufferLen := len(u.buffer) - processedTotal := afterAppendLen - currentBufferLen - bufferIncrease := currentBufferLen - prevBufferLen - - var consumed int - if len(messages) == 0 { - // 没有完整消息,所有数据都被保存到buffer - consumed = len(data) - } else if processedTotal <= prevBufferLen { - // 所有被处理的数据都来自之前的buffer - // 输入data没有被处理,但可能被保存到buffer中 - // consumed = 被保存到buffer中的数据量(如果来自输入data) - if bufferIncrease > 0 { - consumed = bufferIncrease - if consumed > len(data) { - consumed = len(data) - } - } else { - // buffer没有增加,输入data完全没有被使用 - consumed = 0 - } - } else { - // 至少部分被处理的数据来自输入data - // 从输入data中被处理的部分 - consumedFromData := processedTotal - prevBufferLen - // 被保存到buffer中的数据(如果来自输入data) - remainingFromData := len(data) - consumedFromData - if bufferIncrease <= remainingFromData { - // 所有新增到buffer的数据都来自输入data - consumed = len(data) - } else if bufferIncrease > 0 { - // 部分新增到buffer的数据来自输入data - consumed = consumedFromData + remainingFromData - } else { - // buffer没有增加,所有输入data都被处理了 - consumed = consumedFromData - if consumed > len(data) { - consumed = len(data) - } - } - } - - // 确保consumed在合理范围内 - if consumed < 0 { - consumed = 0 - } else if consumed > len(data) { - consumed = len(data) - } - - return messages, u.buffer, consumed, nil + // 输入字节已经被复制进连接级buffer;调用方应从底层读缓冲中丢弃本次输入,避免重复处理。 + return messages, u.buffer, len(data), nil } // Pack 打包 diff --git a/internal/unpacker/fixed_length.go b/internal/unpacker/fixed_length.go index 0c9e822..cd4feaa 100644 --- a/internal/unpacker/fixed_length.go +++ b/internal/unpacker/fixed_length.go @@ -1,7 +1,7 @@ package unpacker import ( - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // fixedLengthUnpacker 固定长度拆包器实现 @@ -33,11 +33,8 @@ func NewFixedLengthUnpackerWithMaxBuffer(length int, maxBufferSize int) unpacker // Unpack 拆包 func (u *fixedLengthUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) { - // 记录调用前的buffer长度(用于计算消耗的数据量) - prevBufferLen := len(u.buffer) - // 检查buffer大小限制 - newSize := prevBufferLen + len(data) + newSize := len(u.buffer) + len(data) if newSize > u.maxBufferSize { return nil, nil, 0, unpackerpkg.NewErrorf("unpacker buffer size exceeded: %d > %d", newSize, u.maxBufferSize) } @@ -63,9 +60,6 @@ func (u *fixedLengthUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) u.buffer = append(u.buffer, data...) - // 记录追加数据后的buffer长度(用于计算consumed) - afterAppendLen := len(u.buffer) - var messages [][]byte for len(u.buffer) >= u.length { @@ -85,50 +79,8 @@ func (u *fixedLengthUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) } } - // 计算本次从输入data中消耗的数据量(100%准确,无误差) - // 使用与delimiterUnpacker相同的逻辑 - currentBufferLen := len(u.buffer) - processedTotal := afterAppendLen - currentBufferLen - bufferIncrease := currentBufferLen - prevBufferLen - - var consumed int - if len(messages) == 0 { - // 没有完整消息,所有数据都被保存到buffer - consumed = len(data) - } else if processedTotal <= prevBufferLen { - // 所有被处理的数据都来自之前的buffer - if bufferIncrease > 0 { - consumed = bufferIncrease - if consumed > len(data) { - consumed = len(data) - } - } else { - consumed = 0 - } - } else { - // 至少部分被处理的数据来自输入data - consumedFromData := processedTotal - prevBufferLen - remainingFromData := len(data) - consumedFromData - if bufferIncrease <= remainingFromData { - consumed = len(data) - } else if bufferIncrease > 0 { - consumed = consumedFromData + remainingFromData - } else { - consumed = consumedFromData - if consumed > len(data) { - consumed = len(data) - } - } - } - - // 确保consumed在合理范围内 - if consumed < 0 { - consumed = 0 - } else if consumed > len(data) { - consumed = len(data) - } - - return messages, u.buffer, consumed, nil + // 输入字节已经被复制进连接级buffer;调用方应从底层读缓冲中丢弃本次输入,避免重复处理。 + return messages, u.buffer, len(data), nil } // Pack 打包 diff --git a/internal/unpacker/fixed_length_test.go b/internal/unpacker/fixed_length_test.go index 10c1818..bc7e86a 100644 --- a/internal/unpacker/fixed_length_test.go +++ b/internal/unpacker/fixed_length_test.go @@ -88,8 +88,5 @@ func TestFixedLengthUnpackerPartialMessages(t *testing.T) { assert.Equal(t, 1, len(messages)) assert.Equal(t, "test", string(messages[0])) assert.Equal(t, "te", string(remaining)) - // consumed应该是2(完成了"test"中的"st"部分,但"te"被保存到buffer中) - assert.Greater(t, consumed, 0, "Expected some data consumed") - assert.LessOrEqual(t, consumed, len(data2), "Expected consumed <= input data length") + assert.Equal(t, len(data2), consumed, "Expected all input bytes to be consumed after buffering") } - diff --git a/internal/unpacker/frame_header.go b/internal/unpacker/frame_header.go index a65d7e7..360b946 100644 --- a/internal/unpacker/frame_header.go +++ b/internal/unpacker/frame_header.go @@ -1,7 +1,7 @@ package unpacker import ( - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // frameHeaderUnpacker 帧头拆包器实现 @@ -45,15 +45,12 @@ func NewFrameHeaderUnpacker(config unpackerpkg.FrameHeaderUnpacker) unpackerpkg. // Unpack 拆包 func (u *frameHeaderUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) { - // 记录调用前的buffer长度(用于计算消耗的数据量) - prevBufferLen := len(u.buffer) - // 检查buffer大小限制 - newSize := prevBufferLen + len(data) + newSize := len(u.buffer) + len(data) if newSize > u.maxBufferSize { return nil, nil, 0, unpackerpkg.NewErrorf("unpacker buffer size exceeded: %d > %d", newSize, u.maxBufferSize) } - + // 优化:如果容量不足,预分配更大的容量(零拷贝优化) if cap(u.buffer) < newSize { newCap := cap(u.buffer) * 2 @@ -75,9 +72,6 @@ func (u *frameHeaderUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) u.buffer = append(u.buffer, data...) - // 记录追加数据后的buffer长度(用于计算consumed) - afterAppendLen := len(u.buffer) - var messages [][]byte for { @@ -95,14 +89,14 @@ func (u *frameHeaderUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) u.buffer = u.buffer[1:] continue } - + // 验证长度字段的合理性(防止恶意数据) if messageLength > u.maxBufferSize { return nil, nil, 0, unpackerpkg.NewErrorf("invalid message length: %d > %d", messageLength, u.maxBufferSize) } totalLength := u.headerLength + messageLength - + if totalLength > u.maxBufferSize { return nil, nil, 0, unpackerpkg.NewErrorf("message too large: %d > %d", totalLength, u.maxBufferSize) } @@ -119,7 +113,7 @@ func (u *frameHeaderUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) // 移除已处理的数据(优化:使用切片操作,避免复制) u.buffer = u.buffer[totalLength:] - + // 如果 buffer 太大但剩余数据很少,压缩 buffer(减少内存占用) // 注意:压缩不会改变buffer的长度,只改变容量 if len(u.buffer) < cap(u.buffer)/4 && cap(u.buffer) > 4096 { @@ -129,50 +123,8 @@ func (u *frameHeaderUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) } } - // 计算本次从输入data中消耗的数据量(100%准确,无误差) - // 使用与delimiterUnpacker相同的逻辑 - currentBufferLen := len(u.buffer) - processedTotal := afterAppendLen - currentBufferLen - bufferIncrease := currentBufferLen - prevBufferLen - - var consumed int - if len(messages) == 0 { - // 没有完整消息,所有数据都被保存到buffer - consumed = len(data) - } else if processedTotal <= prevBufferLen { - // 所有被处理的数据都来自之前的buffer - if bufferIncrease > 0 { - consumed = bufferIncrease - if consumed > len(data) { - consumed = len(data) - } - } else { - consumed = 0 - } - } else { - // 至少部分被处理的数据来自输入data - consumedFromData := processedTotal - prevBufferLen - remainingFromData := len(data) - consumedFromData - if bufferIncrease <= remainingFromData { - consumed = len(data) - } else if bufferIncrease > 0 { - consumed = consumedFromData + remainingFromData - } else { - consumed = consumedFromData - if consumed > len(data) { - consumed = len(data) - } - } - } - - // 确保consumed在合理范围内 - if consumed < 0 { - consumed = 0 - } else if consumed > len(data) { - consumed = len(data) - } - - return messages, u.buffer, consumed, nil + // 输入字节已经被复制进连接级buffer;调用方应从底层读缓冲中丢弃本次输入,避免重复处理。 + return messages, u.buffer, len(data), nil } // Pack 打包 @@ -194,4 +146,3 @@ func (u *frameHeaderUnpacker) Pack(data []byte) ([]byte, error) { return result, nil } - diff --git a/internal/unpacker/frame_header_test.go b/internal/unpacker/frame_header_test.go index 8f37dd3..3d629c1 100644 --- a/internal/unpacker/frame_header_test.go +++ b/internal/unpacker/frame_header_test.go @@ -3,7 +3,7 @@ package unpacker import ( "testing" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -110,4 +110,3 @@ func TestFrameHeaderUnpackerBufferSizeLimit(t *testing.T) { _, _, _, err := u.Unpack(data) assert.Error(t, err, "Expected error for buffer size exceeded") } - diff --git a/internal/unpacker/length_field.go b/internal/unpacker/length_field.go index 68b91db..e454549 100644 --- a/internal/unpacker/length_field.go +++ b/internal/unpacker/length_field.go @@ -3,7 +3,7 @@ package unpacker import ( "encoding/binary" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" ) // lengthFieldUnpacker 长度字段拆包器实现 @@ -19,21 +19,31 @@ type lengthFieldUnpacker struct { // NewLengthFieldUnpacker 创建长度字段拆包器 func NewLengthFieldUnpacker(config unpackerpkg.LengthFieldUnpacker) unpackerpkg.Unpacker { + lengthFieldOffset := config.LengthFieldOffset + if lengthFieldOffset < 0 { + lengthFieldOffset = 0 + } + lengthFieldLength := config.LengthFieldLength if lengthFieldLength <= 0 { lengthFieldLength = 4 // 默认4字节 } + initialBytesToStrip := config.InitialBytesToStrip + if initialBytesToStrip < 0 { + initialBytesToStrip = 0 + } + maxBufferSize := config.MaxBufferSize if maxBufferSize <= 0 { maxBufferSize = unpackerpkg.DefaultMaxBufferSize } return &lengthFieldUnpacker{ - lengthFieldOffset: config.LengthFieldOffset, + lengthFieldOffset: lengthFieldOffset, lengthFieldLength: lengthFieldLength, lengthAdjustment: config.LengthAdjustment, - initialBytesToStrip: config.InitialBytesToStrip, + initialBytesToStrip: initialBytesToStrip, buffer: nil, // 延迟分配,使用buffer pool byteOrder: binary.BigEndian, maxBufferSize: maxBufferSize, @@ -42,11 +52,8 @@ func NewLengthFieldUnpacker(config unpackerpkg.LengthFieldUnpacker) unpackerpkg. // Unpack 拆包 func (u *lengthFieldUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) { - // 记录调用前的buffer长度(用于计算消耗的数据量) - prevBufferLen := len(u.buffer) - // 检查buffer大小限制 - newSize := prevBufferLen + len(data) + newSize := len(u.buffer) + len(data) if newSize > u.maxBufferSize { return nil, nil, 0, unpackerpkg.NewErrorf("unpacker buffer size exceeded: %d > %d", newSize, u.maxBufferSize) } @@ -84,9 +91,6 @@ func (u *lengthFieldUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) // 追加新数据 u.buffer = append(u.buffer, data...) - // 记录追加数据后的buffer长度(用于计算consumed) - afterAppendLen := len(u.buffer) - var messages [][]byte for { @@ -136,6 +140,9 @@ func (u *lengthFieldUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) // 优化:预分配messages slice的容量以减少重新分配 start := u.initialBytesToStrip end := totalLength + if start > end { + return nil, nil, 0, unpackerpkg.NewErrorf("initial bytes to strip exceeds frame length: %d > %d", start, end) + } messageLen := end - start if messages == nil { // 预分配messages slice(假设至少有一个消息) @@ -166,56 +173,20 @@ func (u *lengthFieldUnpacker) Unpack(data []byte) ([][]byte, []byte, int, error) } } - // 计算本次从输入data中消耗的数据量(100%准确,无误差) - // 使用与delimiterUnpacker相同的逻辑 - currentBufferLen := len(u.buffer) - processedTotal := afterAppendLen - currentBufferLen - bufferIncrease := currentBufferLen - prevBufferLen - - var consumed int - if len(messages) == 0 { - // 没有完整消息,所有数据都被保存到buffer - consumed = len(data) - } else if processedTotal <= prevBufferLen { - // 所有被处理的数据都来自之前的buffer - if bufferIncrease > 0 { - consumed = bufferIncrease - if consumed > len(data) { - consumed = len(data) - } - } else { - consumed = 0 - } - } else { - // 至少部分被处理的数据来自输入data - consumedFromData := processedTotal - prevBufferLen - remainingFromData := len(data) - consumedFromData - if bufferIncrease <= remainingFromData { - consumed = len(data) - } else if bufferIncrease > 0 { - consumed = consumedFromData + remainingFromData - } else { - consumed = consumedFromData - if consumed > len(data) { - consumed = len(data) - } - } - } - - // 确保consumed在合理范围内 - if consumed < 0 { - consumed = 0 - } else if consumed > len(data) { - consumed = len(data) - } - - return messages, u.buffer, consumed, nil + // 输入字节已经被复制进连接级buffer;调用方应从底层读缓冲中丢弃本次输入,避免重复处理。 + return messages, u.buffer, len(data), nil } // Pack 打包 func (u *lengthFieldUnpacker) Pack(data []byte) ([]byte, error) { dataLength := len(data) totalLength := u.lengthFieldOffset + u.lengthFieldLength + dataLength - u.lengthAdjustment + if totalLength < u.lengthFieldOffset+u.lengthFieldLength { + return nil, unpackerpkg.NewErrorf("invalid packed length: %d", totalLength) + } + if u.initialBytesToStrip > len(data) { + return nil, unpackerpkg.NewErrorf("initial bytes to strip exceeds data length: %d > %d", u.initialBytesToStrip, len(data)) + } result := make([]byte, totalLength) @@ -226,16 +197,27 @@ func (u *lengthFieldUnpacker) Pack(data []byte) ([]byte, error) { // 写入长度字段 lengthValue := dataLength - u.lengthAdjustment + if lengthValue < 0 { + return nil, unpackerpkg.NewErrorf("invalid length value: %d", lengthValue) + } lengthBytes := make([]byte, u.lengthFieldLength) switch u.lengthFieldLength { case 1: + if lengthValue > 0xff { + return nil, unpackerpkg.NewErrorf("length value too large for 1-byte field: %d", lengthValue) + } lengthBytes[0] = byte(lengthValue) case 2: + if lengthValue > 0xffff { + return nil, unpackerpkg.NewErrorf("length value too large for 2-byte field: %d", lengthValue) + } u.byteOrder.PutUint16(lengthBytes, uint16(lengthValue)) case 4: u.byteOrder.PutUint32(lengthBytes, uint32(lengthValue)) case 8: u.byteOrder.PutUint64(lengthBytes, uint64(lengthValue)) + default: + return nil, unpackerpkg.NewError("unsupported length field length") } copy(result[u.lengthFieldOffset:], lengthBytes) diff --git a/internal/unpacker/length_field_test.go b/internal/unpacker/length_field_test.go index 336a790..feaa76b 100644 --- a/internal/unpacker/length_field_test.go +++ b/internal/unpacker/length_field_test.go @@ -3,7 +3,7 @@ package unpacker import ( "testing" - unpackerpkg "github.com/noahlann/nnet/pkg/unpacker" + unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -86,6 +86,44 @@ func TestLengthFieldUnpackerIncompleteMessage(t *testing.T) { assert.Equal(t, len(data), consumed, "Expected all data consumed when no complete message") } +func TestLengthFieldUnpackerConsumesBufferedRemainder(t *testing.T) { + config := unpackerpkg.LengthFieldUnpacker{ + LengthFieldOffset: 0, + LengthFieldLength: 1, + LengthAdjustment: 0, + InitialBytesToStrip: 0, + MaxBufferSize: 1024, + } + u := NewLengthFieldUnpacker(config) + + messages, remaining, consumed, err := u.Unpack([]byte{3, 'a'}) + require.NoError(t, err) + assert.Empty(t, messages) + assert.Equal(t, []byte{3, 'a'}, remaining) + assert.Equal(t, 2, consumed) + + messages, remaining, consumed, err = u.Unpack([]byte{'b', 'c', 1}) + require.NoError(t, err) + require.Len(t, messages, 1) + assert.Equal(t, []byte{3, 'a', 'b', 'c'}, messages[0]) + assert.Equal(t, []byte{1}, remaining) + assert.Equal(t, 3, consumed, "Expected complete and buffered input bytes to be consumed") +} + +func TestLengthFieldUnpackerRejectsInvalidStrip(t *testing.T) { + config := unpackerpkg.LengthFieldUnpacker{ + LengthFieldOffset: 0, + LengthFieldLength: 1, + LengthAdjustment: 0, + InitialBytesToStrip: 5, + MaxBufferSize: 1024, + } + u := NewLengthFieldUnpacker(config) + + _, _, _, err := u.Unpack([]byte{1, 'a'}) + assert.Error(t, err) +} + func TestLengthFieldUnpackerPack(t *testing.T) { config := unpackerpkg.LengthFieldUnpacker{ LengthFieldOffset: 0, @@ -126,4 +164,3 @@ func TestLengthFieldUnpackerBufferSizeLimit(t *testing.T) { _, _, _, err := u.Unpack(data) assert.Error(t, err, "Expected error for buffer size exceeded") } - diff --git a/pkg/client/errors.go b/pkg/client/errors.go index e857cb9..8a45273 100644 --- a/pkg/client/errors.go +++ b/pkg/client/errors.go @@ -40,4 +40,3 @@ func (e *Error) Error() string { func (e *Error) Unwrap() error { return e.Cause } - diff --git a/pkg/codec/codec.go b/pkg/codec/codec.go index 1d6e4c3..81a9309 100644 --- a/pkg/codec/codec.go +++ b/pkg/codec/codec.go @@ -23,4 +23,3 @@ type Registry interface { // Default 获取默认编解码器 Default() Codec } - diff --git a/pkg/codec/errors.go b/pkg/codec/errors.go index 568db2f..a4f9a07 100644 --- a/pkg/codec/errors.go +++ b/pkg/codec/errors.go @@ -40,4 +40,3 @@ func (e *Error) Error() string { func (e *Error) Unwrap() error { return e.Cause } - diff --git a/pkg/codec/resolver.go b/pkg/codec/resolver.go index e83d530..768313c 100644 --- a/pkg/codec/resolver.go +++ b/pkg/codec/resolver.go @@ -1,8 +1,8 @@ package codec import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // Resolver Codec解析器接口 @@ -38,9 +38,9 @@ type ResolverChain struct { // NewResolverChain 创建解析器链 func NewResolverChain(registry Registry, defaultName string) *ResolverChain { return &ResolverChain{ - resolvers: make([]Resolver, 0), - registry: registry, - defaultName: defaultName, + resolvers: make([]Resolver, 0), + registry: registry, + defaultName: defaultName, } } @@ -97,4 +97,3 @@ func (c *ResolverChain) ResolveForEncode(ctx ctxpkg.Context, data interface{}, h } return c.registry.Default(), nil } - diff --git a/pkg/config/config.go b/pkg/config/config.go index 8c73ab9..07c5c6f 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -3,7 +3,7 @@ package config import ( "time" - "github.com/noahlann/nnet/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" ) // Config 服务器配置 diff --git a/pkg/context/context.go b/pkg/context/context.go index 42a8a36..2bd2058 100644 --- a/pkg/context/context.go +++ b/pkg/context/context.go @@ -4,8 +4,8 @@ import ( "context" "time" - "github.com/noahlann/nnet/pkg/request" - "github.com/noahlann/nnet/pkg/response" + "git.noahlan.cn/noahlan/nnet/v2/pkg/request" + "git.noahlan.cn/noahlan/nnet/v2/pkg/response" ) // Context 请求上下文接口 diff --git a/pkg/context/context_test.go b/pkg/context/context_test.go index 0082ddd..57b1794 100644 --- a/pkg/context/context_test.go +++ b/pkg/context/context_test.go @@ -5,8 +5,8 @@ import ( "testing" "time" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - responsepkg "github.com/noahlann/nnet/pkg/response" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + responsepkg "git.noahlan.cn/noahlan/nnet/v2/pkg/response" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/pkg/health/errors.go b/pkg/health/errors.go index 2f2aa94..3513fc8 100644 --- a/pkg/health/errors.go +++ b/pkg/health/errors.go @@ -16,4 +16,3 @@ func NewError(message string) error { func (e *Error) Error() string { return e.Message } - diff --git a/pkg/health/health.go b/pkg/health/health.go index 2d59d6d..b8977aa 100644 --- a/pkg/health/health.go +++ b/pkg/health/health.go @@ -87,4 +87,3 @@ func (c *checker) Check() (Status, map[string]CheckResult) { return overallStatus, results } - diff --git a/pkg/interceptor/errors.go b/pkg/interceptor/errors.go index 4f6f94f..15e313f 100644 --- a/pkg/interceptor/errors.go +++ b/pkg/interceptor/errors.go @@ -40,4 +40,3 @@ func (e *Error) Error() string { func (e *Error) Unwrap() error { return e.Cause } - diff --git a/pkg/interceptor/interceptor.go b/pkg/interceptor/interceptor.go index 64bf822..779ea0a 100644 --- a/pkg/interceptor/interceptor.go +++ b/pkg/interceptor/interceptor.go @@ -1,7 +1,7 @@ package interceptor import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // Interceptor 拦截器接口 @@ -27,4 +27,3 @@ type HandlerFunc func(data []byte, ctx ctxpkg.Context, next Chain) ([]byte, bool func (f HandlerFunc) Intercept(data []byte, ctx ctxpkg.Context, next Chain) ([]byte, bool, error) { return f(data, ctx, next) } - diff --git a/pkg/lifecycle/lifecycle.go b/pkg/lifecycle/lifecycle.go index 528d80b..b9e42f0 100644 --- a/pkg/lifecycle/lifecycle.go +++ b/pkg/lifecycle/lifecycle.go @@ -100,4 +100,3 @@ func (f *ConnectionHookFunc) OnError(connID string, err error) error { } return nil } - diff --git a/pkg/metrics/metrics.go b/pkg/metrics/metrics.go index e92884e..9276ce8 100644 --- a/pkg/metrics/metrics.go +++ b/pkg/metrics/metrics.go @@ -106,4 +106,3 @@ func (m *metrics) AddBytesSent(bytes int64) { func (m *metrics) GetBytesSent() int64 { return atomic.LoadInt64(&m.bytesSent) } - diff --git a/pkg/middleware/middleware.go b/pkg/middleware/middleware.go index 6ca99ad..066f663 100644 --- a/pkg/middleware/middleware.go +++ b/pkg/middleware/middleware.go @@ -1,7 +1,7 @@ package middleware import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // Middleware 中间件接口 @@ -36,4 +36,3 @@ func (c *Chain) wrapHandler(middleware Middleware, handler Handler) Handler { return middleware(ctx, handler) } } - diff --git a/pkg/nnet/client.go b/pkg/nnet/client.go index 4de5d3a..27aa037 100644 --- a/pkg/nnet/client.go +++ b/pkg/nnet/client.go @@ -1,8 +1,8 @@ package nnet import ( - "github.com/noahlann/nnet/internal/client" - clientpkg "github.com/noahlann/nnet/pkg/client" + "git.noahlan.cn/noahlan/nnet/v2/internal/client" + clientpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/client" ) // Client 客户端类型别名 diff --git a/pkg/nnet/context.go b/pkg/nnet/context.go index bb7807a..93edbea 100644 --- a/pkg/nnet/context.go +++ b/pkg/nnet/context.go @@ -1,8 +1,8 @@ package nnet import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - routerpkg "github.com/noahlann/nnet/pkg/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // Context 上下文接口 diff --git a/pkg/nnet/health.go b/pkg/nnet/health.go index 217f526..1f6fa06 100644 --- a/pkg/nnet/health.go +++ b/pkg/nnet/health.go @@ -1,7 +1,7 @@ package nnet import ( - healthpkg "github.com/noahlann/nnet/pkg/health" + healthpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/health" ) // HealthChecker 健康检查器类型别名 @@ -27,4 +27,3 @@ const ( func NewHealthChecker() HealthChecker { return healthpkg.NewChecker() } - diff --git a/pkg/nnet/lifecycle.go b/pkg/nnet/lifecycle.go index ff71ce6..906ff82 100644 --- a/pkg/nnet/lifecycle.go +++ b/pkg/nnet/lifecycle.go @@ -1,7 +1,7 @@ package nnet import ( - lifecyclepkg "github.com/noahlann/nnet/pkg/lifecycle" + lifecyclepkg "git.noahlan.cn/noahlan/nnet/v2/pkg/lifecycle" ) // ServerLifecycleHook 服务器生命周期钩子类型别名 @@ -42,4 +42,3 @@ func NewConnectionHookFunc( OnErrorFunc: onError, } } - diff --git a/pkg/nnet/metrics.go b/pkg/nnet/metrics.go index 139eed9..ba7177c 100644 --- a/pkg/nnet/metrics.go +++ b/pkg/nnet/metrics.go @@ -1,7 +1,7 @@ package nnet import ( - metricspkg "github.com/noahlann/nnet/pkg/metrics" + metricspkg "git.noahlan.cn/noahlan/nnet/v2/pkg/metrics" ) // Metrics 指标类型别名 diff --git a/pkg/nnet/preset.go b/pkg/nnet/preset.go index a6d0af2..52d12aa 100644 --- a/pkg/nnet/preset.go +++ b/pkg/nnet/preset.go @@ -3,8 +3,8 @@ package nnet import ( "net/http" - "github.com/noahlann/nnet/pkg/config" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // PresetBuilder 用于构建基于 nnet 协议的服务器预设。 diff --git a/pkg/nnet/server.go b/pkg/nnet/server.go index 6c11fb1..18fcca5 100644 --- a/pkg/nnet/server.go +++ b/pkg/nnet/server.go @@ -5,12 +5,12 @@ import ( "io" "net/http" - "github.com/noahlann/nnet/internal/connection" - "github.com/noahlann/nnet/internal/server" - codecpkg "github.com/noahlann/nnet/pkg/codec" - "github.com/noahlann/nnet/pkg/config" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/internal/server" + codecpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/codec" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) // Server 服务器接口 diff --git a/pkg/plugin/plugin.go b/pkg/plugin/plugin.go index 0ee3ffa..766d9cd 100644 --- a/pkg/plugin/plugin.go +++ b/pkg/plugin/plugin.go @@ -3,7 +3,7 @@ package plugin import ( "context" - "github.com/noahlann/nnet/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" ) // Plugin 插件接口 @@ -77,4 +77,3 @@ type Manager interface { // StopAll 停止所有插件 StopAll(ctx context.Context) error } - diff --git a/pkg/protocol/errors.go b/pkg/protocol/errors.go index b759dbd..df24ebe 100644 --- a/pkg/protocol/errors.go +++ b/pkg/protocol/errors.go @@ -40,4 +40,3 @@ func (e *Error) Error() string { func (e *Error) Unwrap() error { return e.Cause } - diff --git a/pkg/request/request.go b/pkg/request/request.go index 4acaebb..20f634e 100644 --- a/pkg/request/request.go +++ b/pkg/request/request.go @@ -1,7 +1,7 @@ package request import ( - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // Request 请求接口 @@ -34,4 +34,3 @@ type Request interface { // FrameHeader 协议帧头接口(类型别名,统一使用protocol.FrameHeader) type FrameHeader = protocolpkg.FrameHeader - diff --git a/pkg/response/response.go b/pkg/response/response.go index 73df9c2..0f337a8 100644 --- a/pkg/response/response.go +++ b/pkg/response/response.go @@ -1,7 +1,7 @@ package response import ( - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // Response 响应接口 @@ -35,4 +35,3 @@ type Response interface { // FrameHeader 协议帧头接口(类型别名,统一使用protocol.FrameHeader) type FrameHeader = protocolpkg.FrameHeader - diff --git a/pkg/router/custom_matcher.go b/pkg/router/custom_matcher.go index 79ed7d9..b8e23a6 100644 --- a/pkg/router/custom_matcher.go +++ b/pkg/router/custom_matcher.go @@ -1,7 +1,7 @@ package router import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // CustomMatcher 自定义匹配器 diff --git a/pkg/router/handler.go b/pkg/router/handler.go index 928fd77..bd8799e 100644 --- a/pkg/router/handler.go +++ b/pkg/router/handler.go @@ -3,7 +3,7 @@ package router import ( "reflect" - ctxpkg "github.com/noahlann/nnet/pkg/context" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // Handler 处理器函数 @@ -29,4 +29,3 @@ type Route interface { // CodecName 获取编解码器名称 CodecName() string } - diff --git a/pkg/router/matcher.go b/pkg/router/matcher.go index ca7d3ad..275199d 100644 --- a/pkg/router/matcher.go +++ b/pkg/router/matcher.go @@ -1,8 +1,8 @@ package router import ( - ctxpkg "github.com/noahlann/nnet/pkg/context" - protocolpkg "github.com/noahlann/nnet/pkg/protocol" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" ) // MatchInput 匹配输入 diff --git a/pkg/router/router.go b/pkg/router/router.go index 29e5c32..a33a232 100644 --- a/pkg/router/router.go +++ b/pkg/router/router.go @@ -3,7 +3,7 @@ package router import ( "reflect" - ctxpkg "github.com/noahlann/nnet/pkg/context" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) // Router 路由器接口 diff --git a/pkg/session/session.go b/pkg/session/session.go index a1dc5a2..80649ac 100644 --- a/pkg/session/session.go +++ b/pkg/session/session.go @@ -50,4 +50,3 @@ type Manager interface { // Delete 删除Session Delete(sessionID string) error } - diff --git a/pkg/unpacker/errors.go b/pkg/unpacker/errors.go index ef73a3e..ea4e2d2 100644 --- a/pkg/unpacker/errors.go +++ b/pkg/unpacker/errors.go @@ -40,4 +40,3 @@ func (e *Error) Error() string { func (e *Error) Unwrap() error { return e.Cause } - diff --git a/pkg/unpacker/unpacker.go b/pkg/unpacker/unpacker.go index 5bc0635..d6ae9ff 100644 --- a/pkg/unpacker/unpacker.go +++ b/pkg/unpacker/unpacker.go @@ -5,7 +5,7 @@ type Unpacker interface { // Unpack 拆包 // data: 原始数据 // 返回: 完整的消息列表, 剩余数据, 本次从输入data中消耗的字节数, 错误 - // consumed: 本次从输入data中消耗的字节数,用于准确计算从gnet buffer中需要discard的数据量 + // consumed: 成功写入拆包器内部缓冲的输入字节数,用于计算从底层读缓冲中需要discard的数据量 Unpack(data []byte) ([][]byte, []byte, int, error) // Pack 打包 @@ -34,14 +34,13 @@ type LengthFieldUnpacker struct { // DelimiterUnpacker 分隔符拆包器 type DelimiterUnpacker struct { - Delimiter []byte + Delimiter []byte MaxBufferSize int // 最大buffer大小,0表示使用默认值 } // FrameHeaderUnpacker 帧头拆包器 type FrameHeaderUnpacker struct { - HeaderLength int // 帧头长度 + HeaderLength int // 帧头长度 GetLength func(header []byte) int // 从帧头获取消息长度 - MaxBufferSize int // 最大buffer大小,0表示使用默认值 + MaxBufferSize int // 最大buffer大小,0表示使用默认值 } - diff --git a/test/connection_test.go b/test/connection_test.go index 2d26ada..9d60267 100644 --- a/test/connection_test.go +++ b/test/connection_test.go @@ -4,8 +4,8 @@ import ( "testing" "time" - "github.com/noahlann/nnet/internal/connection" - "github.com/noahlann/nnet/pkg/errors" + "git.noahlan.cn/noahlan/nnet/v2/internal/connection" + "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" ) func TestConnectionManager(t *testing.T) { diff --git a/test/context_test.go b/test/context_test.go index 054a5e4..8d232d5 100644 --- a/test/context_test.go +++ b/test/context_test.go @@ -5,10 +5,10 @@ import ( "testing" "time" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - ctxpkg "github.com/noahlann/nnet/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" ) func TestContext(t *testing.T) { diff --git a/test/integration/basic_test.go b/test/integration/basic_test.go index 66e49c1..8064692 100644 --- a/test/integration/basic_test.go +++ b/test/integration/basic_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -26,7 +26,7 @@ func TestBasicServerClient(t *testing.T) { Output: "stdout", }, } - + // 先获取随机端口 listener, err := net.Listen("tcp", "127.0.0.1:0") require.NoError(t, err) @@ -216,4 +216,3 @@ func TestMultipleClients(t *testing.T) { CleanupTestClient(t, client) } } - diff --git a/test/integration/codec_resolver_test.go b/test/integration/codec_resolver_test.go index 43a0caa..f1f7271 100644 --- a/test/integration/codec_resolver_test.go +++ b/test/integration/codec_resolver_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -96,7 +96,7 @@ func TestCodecAutoDetection(t *testing.T) { } ts := StartTestServerWithRoutes(t, cfg, func(srv nnet.Server) { - srv.Router().RegisterString("test", func(ctx nnet.Context) error { + srv.Router().RegisterFrameData("test", "==", "data", func(ctx nnet.Context) error { // 返回请求的数据类型 data := ctx.Request().Data() return ctx.Response().Write(map[string]any{ @@ -123,4 +123,3 @@ func TestCodecAutoDetection(t *testing.T) { assert.NoError(t, err, "Response should be valid JSON") assert.Equal(t, "ok", result["status"], "Status should be ok") } - diff --git a/test/integration/codec_test.go b/test/integration/codec_test.go index dd94869..71b7919 100644 --- a/test/integration/codec_test.go +++ b/test/integration/codec_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -120,8 +120,7 @@ func TestRouteNotFound(t *testing.T) { } else { // 如果收到响应,应该包含错误信息 assert.NotNil(t, resp, "Response should not be nil") - assert.Contains(t, string(resp), "Route not found", "Response should contain error message") + assert.Contains(t, string(resp), "route not found", "Response should contain error message") t.Logf("Received error response: %q", string(resp)) } } - diff --git a/test/integration/concurrent_test.go b/test/integration/concurrent_test.go index 184b046..d570fe5 100644 --- a/test/integration/concurrent_test.go +++ b/test/integration/concurrent_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -170,11 +170,13 @@ func TestServerStop(t *testing.T) { return !ts.Server.Started() }, 3*time.Second, 50*time.Millisecond, "Server should stop") - // 客户端应该断开连接 + // TCP客户端需要通过一次I/O感知远端关闭。 + _, _ = client.Request([]byte("test"), 500*time.Millisecond) + + // 客户端应该在读写失败后标记断开。 require.Eventually(t, func() bool { return !client.IsConnected() }, 2*time.Second, 50*time.Millisecond, "Client should disconnect") CleanupTestClient(t, client) } - diff --git a/test/integration/connection_mgr_test.go b/test/integration/connection_mgr_test.go index c88deef..cc9222b 100644 --- a/test/integration/connection_mgr_test.go +++ b/test/integration/connection_mgr_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -140,12 +140,18 @@ func TestConnectionManagerBroadcast(t *testing.T) { srv.Router().RegisterString("join", func(ctx nnet.Context) error { connID := ctx.Connection().ID() - return connMgr.AddToGroup(groupID, connID) + if err := connMgr.AddToGroup(groupID, connID); err != nil { + return err + } + return ctx.Response().WriteBytes([]byte("joined\n")) }) srv.Router().RegisterString("broadcast", func(ctx nnet.Context) error { message := []byte("broadcast message\n") - return connMgr.BroadcastToGroup(groupID, message) + if err := connMgr.BroadcastToGroup(groupID, message); err != nil { + return err + } + return ctx.Response().WriteBytes([]byte("broadcasted\n")) }) }) defer CleanupTestServer(t, ts) @@ -177,4 +183,3 @@ func TestConnectionManagerBroadcast(t *testing.T) { // 如果客户端不支持,可能需要使用其他方式验证广播 t.Logf("Broadcast sent, clients should receive message") } - diff --git a/test/integration/error_handling_test.go b/test/integration/error_handling_test.go index 705fb0b..98b6233 100644 --- a/test/integration/error_handling_test.go +++ b/test/integration/error_handling_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -136,4 +136,3 @@ func TestInvalidData(t *testing.T) { assert.NotEmpty(t, resp, "Server should handle invalid data") } } - diff --git a/test/integration/health_test.go b/test/integration/health_test.go index 2de05c7..297ff60 100644 --- a/test/integration/health_test.go +++ b/test/integration/health_test.go @@ -8,7 +8,7 @@ import ( "net/http/httptest" "testing" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -81,4 +81,3 @@ func TestHealthHandler(t *testing.T) { assert.Contains(t, result, "results", "Response should contain results") t.Logf("Health check response: %s", w.Body.String()) } - diff --git a/test/integration/helper.go b/test/integration/helper.go index 1617c5c..dbc151a 100644 --- a/test/integration/helper.go +++ b/test/integration/helper.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/require" ) @@ -16,10 +16,7 @@ import ( func StartTestServerWithRoutes(t *testing.T, cfg *nnet.Config, setupRoutes func(nnet.Server)) *TestServer { // 如果地址为空或使用端口0,使用随机端口 if cfg.Addr == "" || cfg.Addr == "tcp://:0" || cfg.Addr == "tcp://127.0.0.1:0" || cfg.Addr == "tcp://:6995" { - listener, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err, "Failed to get random port") - port := listener.Addr().(*net.TCPAddr).Port - listener.Close() + port := reserveNonEphemeralTCPPort(t) cfg.Addr = fmt.Sprintf("tcp://127.0.0.1:%d", port) } @@ -72,6 +69,24 @@ func StartTestServerWithRoutes(t *testing.T, cfg *nnet.Config, setupRoutes func( return ts } +func reserveNonEphemeralTCPPort(t *testing.T) int { + start := 20000 + int(time.Now().UnixNano()%20000) + for i := 0; i < 200; i++ { + port := 20000 + ((start - 20000 + i) % 20000) + listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", port)) + if err == nil { + _ = listener.Close() + return port + } + } + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err, "Failed to get random port") + port := listener.Addr().(*net.TCPAddr).Port + _ = listener.Close() + return port +} + // TestServer 测试服务器包装 type TestServer struct { Server nnet.Server diff --git a/test/integration/integration_test.go b/test/integration/integration_test.go index e8aabcc..8d412f1 100644 --- a/test/integration/integration_test.go +++ b/test/integration/integration_test.go @@ -4,7 +4,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/test/integration/interceptor_test.go b/test/integration/interceptor_test.go index 9a663ac..6ccc99e 100644 --- a/test/integration/interceptor_test.go +++ b/test/integration/interceptor_test.go @@ -7,11 +7,11 @@ import ( "testing" "time" - internalinterceptor "github.com/noahlann/nnet/internal/interceptor" - "github.com/noahlann/nnet/internal/interceptor/builtin" - "github.com/noahlann/nnet/pkg/interceptor" - "github.com/noahlann/nnet/pkg/nnet" - routerpkg "github.com/noahlann/nnet/pkg/router" + internalinterceptor "git.noahlan.cn/noahlan/nnet/v2/internal/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/internal/interceptor/builtin" + "git.noahlan.cn/noahlan/nnet/v2/pkg/interceptor" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/test/integration/large_message_test.go b/test/integration/large_message_test.go index d7b7302..ef1af94 100644 --- a/test/integration/large_message_test.go +++ b/test/integration/large_message_test.go @@ -7,8 +7,8 @@ import ( "testing" "time" - internalprotocol "github.com/noahlann/nnet/internal/protocol/nnet" - "github.com/noahlann/nnet/pkg/nnet" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -21,9 +21,11 @@ func TestLargeMessage(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ - DefaultCodec: "json", + DefaultCodec: "json", + EnableProtocolEncode: true, }, } @@ -39,7 +41,9 @@ func TestLargeMessage(t *testing.T) { }) defer CleanupTestServer(t, ts) - client := NewTestClient(t, ts.Addr, nil) + client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ + ApplicationProtocol: "nnet", + }) defer CleanupTestClient(t, client) ConnectTestClient(t, client) @@ -59,7 +63,12 @@ func TestMultipleMessages(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + ApplicationProtocol: "nnet", + Codec: &nnet.CodecConfig{ + DefaultCodec: "plain", + EnableProtocolEncode: true, + }, } ts := StartTestServerWithRoutes(t, cfg, func(srv nnet.Server) { @@ -69,7 +78,9 @@ func TestMultipleMessages(t *testing.T) { }) defer CleanupTestServer(t, ts) - client := NewTestClient(t, ts.Addr, nil) + client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ + ApplicationProtocol: "nnet", + }) defer CleanupTestClient(t, client) ConnectTestClient(t, client) @@ -91,7 +102,7 @@ func TestUnpacker(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ DefaultCodec: "json", @@ -135,7 +146,7 @@ func TestUnpacker(t *testing.T) { defer CleanupTestServer(t, ts) client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ - ApplicationProtocol: "nnet", + ApplicationProtocol: "", }) defer CleanupTestClient(t, client) @@ -156,7 +167,6 @@ func TestUnpacker(t *testing.T) { time.Sleep(200 * time.Millisecond) // 验证消息被正确拆包 - assert.GreaterOrEqual(t, messageCount, 1, "At least one message should be processed") + assert.GreaterOrEqual(t, messageCount, 2, "Both coalesced protocol frames should be processed") t.Logf("Processed %d messages", messageCount) } - diff --git a/test/integration/lifecycle_test.go b/test/integration/lifecycle_test.go index 7e51077..c6de99e 100644 --- a/test/integration/lifecycle_test.go +++ b/test/integration/lifecycle_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -137,4 +137,3 @@ func TestConnectionLifecycleHooks(t *testing.T) { assert.Equal(t, onOpenConnID, onCloseConnID, "OnClose connID should match OnOpen connID") mu.Unlock() } - diff --git a/test/integration/metrics_test.go b/test/integration/metrics_test.go index d457fab..6ae5017 100644 --- a/test/integration/metrics_test.go +++ b/test/integration/metrics_test.go @@ -9,7 +9,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -152,4 +152,3 @@ func TestMetricsExportWriter(t *testing.T) { assert.Contains(t, output, "nnet_", "Output should contain metrics") t.Logf("Metrics output: %s", output) } - diff --git a/test/integration/middleware_test.go b/test/integration/middleware_test.go index 120517c..2785f1f 100644 --- a/test/integration/middleware_test.go +++ b/test/integration/middleware_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -164,4 +164,3 @@ func TestMiddlewareError(t *testing.T) { assert.NotEmpty(t, resp, "Response should not be empty") } } - diff --git a/test/integration/plugin_test.go b/test/integration/plugin_test.go index 5b75638..cd87f22 100644 --- a/test/integration/plugin_test.go +++ b/test/integration/plugin_test.go @@ -7,21 +7,21 @@ import ( "sync" "testing" - internalplugin "github.com/noahlann/nnet/internal/plugin" - "github.com/noahlann/nnet/pkg/config" - "github.com/noahlann/nnet/pkg/nnet" + internalplugin "git.noahlan.cn/noahlan/nnet/v2/internal/plugin" + "git.noahlan.cn/noahlan/nnet/v2/pkg/config" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // testPlugin 测试插件实现 type testPlugin struct { - name string - version string - initCalled bool - startCalled bool - stopCalled bool - mu sync.RWMutex + name string + version string + initCalled bool + startCalled bool + stopCalled bool + mu sync.RWMutex } func (p *testPlugin) Name() string { @@ -288,4 +288,3 @@ func (p *testMessagePlugin) OnMessage(ctx context.Context, data []byte) ([]byte, } return data, nil } - diff --git a/test/integration/preset_test.go b/test/integration/preset_test.go index f85d91a..03e9df5 100644 --- a/test/integration/preset_test.go +++ b/test/integration/preset_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/test/integration/protocol_header_test.go b/test/integration/protocol_header_test.go index e39a912..f6acba5 100644 --- a/test/integration/protocol_header_test.go +++ b/test/integration/protocol_header_test.go @@ -7,8 +7,8 @@ import ( "testing" "time" - internalprotocol "github.com/noahlann/nnet/internal/protocol/nnet" - "github.com/noahlann/nnet/pkg/nnet" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -21,7 +21,7 @@ func TestProtocolFrameHeader(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ DefaultCodec: "json", @@ -37,22 +37,14 @@ func TestProtocolFrameHeader(t *testing.T) { proto := internalprotocol.NewNNetProtocol("1.0") require.NoError(t, pm.Register(proto), "Should register protocol") - // 注册帧头匹配路由 - server.Router().RegisterFrameHeader("op", "==", "ping", func(ctx nnet.Context) error { + // 注册帧头匹配路由:nnet帧头包含magic/version/length/checksum等字段。 + server.Router().RegisterFrameHeader("version", "==", byte(1), func(ctx nnet.Context) error { return ctx.Response().Write(map[string]any{ - "op": "pong", + "op": "pong", "status": "ok", }) }) - server.Router().RegisterFrameHeader("op", "==", "echo", func(ctx nnet.Context) error { - data := ctx.Request().Data() - return ctx.Response().Write(map[string]any{ - "op": "echo", - "data": data, - }) - }) - ts := &TestServer{ Server: server, Addr: cfg.Addr, @@ -83,20 +75,12 @@ func TestProtocolFrameHeader(t *testing.T) { ConnectTestClient(t, client) time.Sleep(100 * time.Millisecond) - // 使用nnet协议编码请求 pingData := []byte(`{"op":"ping"}`) - pingPacket, err := proto.Encode(pingData, nil) - require.NoError(t, err, "Failed to encode ping request") - - resp := RequestWithTimeout(t, client, pingPacket, 3*time.Second) + resp := RequestWithTimeout(t, client, pingData, 3*time.Second) t.Logf("Response: %q", string(resp)) - // 解码响应 - _, respPayload, err := proto.Decode(resp) - require.NoError(t, err, "Failed to decode response") - var result map[string]any - err = json.Unmarshal(respPayload, &result) + err = json.Unmarshal(resp, &result) assert.NoError(t, err, "Response should be valid JSON") assert.Equal(t, "pong", result["op"], "Op should be pong") assert.Equal(t, "ok", result["status"], "Status should be ok") @@ -110,7 +94,7 @@ func TestProtocolFrameData(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ DefaultCodec: "json", @@ -163,22 +147,13 @@ func TestProtocolFrameData(t *testing.T) { ConnectTestClient(t, client) time.Sleep(100 * time.Millisecond) - // 使用nnet协议编码请求 testData := []byte(`{"action":"test"}`) - testPacket, err := proto.Encode(testData, nil) - require.NoError(t, err, "Failed to encode test request") - - resp := RequestWithTimeout(t, client, testPacket, 3*time.Second) + resp := RequestWithTimeout(t, client, testData, 3*time.Second) t.Logf("Response: %q", string(resp)) - // 解码响应 - _, respPayload, err := proto.Decode(resp) - require.NoError(t, err, "Failed to decode response") - var result map[string]any - err = json.Unmarshal(respPayload, &result) + err = json.Unmarshal(resp, &result) assert.NoError(t, err, "Response should be valid JSON") assert.Equal(t, "test", result["action"], "Action should be test") assert.Equal(t, "success", result["result"], "Result should be success") } - diff --git a/test/integration/protocol_test.go b/test/integration/protocol_test.go index c3a6614..557f2ea 100644 --- a/test/integration/protocol_test.go +++ b/test/integration/protocol_test.go @@ -7,8 +7,8 @@ import ( "testing" "time" - internalprotocol "github.com/noahlann/nnet/internal/protocol/nnet" - "github.com/noahlann/nnet/pkg/nnet" + internalprotocol "git.noahlan.cn/noahlan/nnet/v2/internal/protocol/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -22,7 +22,7 @@ func TestProtocolVersion(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ DefaultCodec: "json", @@ -81,7 +81,7 @@ func TestProtocolVersion(t *testing.T) { require.Eventually(t, func() bool { return server.Started() }, 3*time.Second, 50*time.Millisecond, "Server should start within 3 seconds") - + // 等待服务器准备好 require.Eventually(t, func() bool { testConn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", port), 500*time.Millisecond) @@ -104,16 +104,8 @@ func TestProtocolVersion(t *testing.T) { // 等待服务器准备好 time.Sleep(100 * time.Millisecond) - // 使用nnet协议编码请求 - requestPacket, err := protoV1.Encode([]byte("version"), nil) - require.NoError(t, err, "Failed to encode request with nnet protocol") - - respPacket := RequestWithTimeout(t, client, requestPacket, 3*time.Second) - t.Logf("Response for version (raw): %q", string(respPacket)) - - // 解码响应 - _, respPayload, err := protoV1.Decode(respPacket) - require.NoError(t, err, "Failed to decode response packet") + respPayload := RequestWithTimeout(t, client, []byte("version"), 3*time.Second) + t.Logf("Response for version: %q", string(respPayload)) var result map[string]any err = json.Unmarshal(respPayload, &result) @@ -121,4 +113,3 @@ func TestProtocolVersion(t *testing.T) { assert.Equal(t, "nnet", result["protocol"], "Protocol should be nnet") assert.Equal(t, "1.0", result["version"], "Version should be 1.0") } - diff --git a/test/integration/route_types_test.go b/test/integration/route_types_test.go index c956f27..cfc9889 100644 --- a/test/integration/route_types_test.go +++ b/test/integration/route_types_test.go @@ -7,8 +7,8 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" - routerpkg "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" + routerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/router" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -63,7 +63,7 @@ func TestRouteCustom(t *testing.T) { }, func(ctx nnet.Context) error { return ctx.Response().Write(map[string]any{ - "route": "custom", + "route": "custom", "status": "ok", }) }, @@ -120,4 +120,3 @@ func TestRoutePriority(t *testing.T) { resp := RequestWithTimeout(t, client, []byte("test"), 3*time.Second) assert.Contains(t, string(resp), "exact match", "Exact match should have priority") } - diff --git a/test/integration/router_test.go b/test/integration/router_test.go index 8095103..7dc608e 100644 --- a/test/integration/router_test.go +++ b/test/integration/router_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -59,7 +59,7 @@ func TestRouterFrameData(t *testing.T) { require.Eventually(t, func() bool { return server.Started() }, 3*time.Second, 50*time.Millisecond, "Server should start within 3 seconds") - + // 等待服务器准备好 require.Eventually(t, func() bool { testConn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", port), 500*time.Millisecond) @@ -84,7 +84,7 @@ func TestRouterFrameData(t *testing.T) { pingData := `{"op":"ping"}` resp := RequestWithTimeout(t, client, []byte(pingData), 3*time.Second) t.Logf("Response for ping: %q", string(resp)) - + var result map[string]any err = json.Unmarshal(resp, &result) assert.NoError(t, err, "Response should be valid JSON") @@ -94,7 +94,7 @@ func TestRouterFrameData(t *testing.T) { echoData := `{"op":"echo","msg":"hello"}` resp = RequestWithTimeout(t, client, []byte(echoData), 3*time.Second) t.Logf("Response for echo: %q", string(resp)) - + err = json.Unmarshal(resp, &result) assert.NoError(t, err, "Response should be valid JSON") assert.Contains(t, result["echo"].(string), "echo", "Response should echo the data") @@ -131,7 +131,7 @@ func TestRouterMiddleware(t *testing.T) { group.RegisterString("hello", func(ctx nnet.Context) error { middleware := ctx.GetString("middleware") return ctx.Response().Write(map[string]any{ - "message": "hello with middleware", + "message": "hello with middleware", "middleware": middleware, }) }) @@ -154,7 +154,7 @@ func TestRouterMiddleware(t *testing.T) { require.Eventually(t, func() bool { return server.Started() }, 3*time.Second, 50*time.Millisecond, "Server should start within 3 seconds") - + // 等待服务器准备好 require.Eventually(t, func() bool { testConn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", port), 500*time.Millisecond) @@ -177,10 +177,9 @@ func TestRouterMiddleware(t *testing.T) { resp := RequestWithTimeout(t, client, []byte("hello"), 3*time.Second) t.Logf("Response for middleware test: %q", string(resp)) - + var result map[string]any err = json.Unmarshal(resp, &result) assert.NoError(t, err, "Response should be valid JSON") assert.Equal(t, "executed", result["middleware"], "Middleware should be executed") } - diff --git a/test/integration/session_test.go b/test/integration/session_test.go index 3346ef8..d6ff53b 100644 --- a/test/integration/session_test.go +++ b/test/integration/session_test.go @@ -7,9 +7,9 @@ import ( "testing" "time" - internalsession "github.com/noahlann/nnet/internal/session" - "github.com/noahlann/nnet/internal/session/storage" - "github.com/noahlann/nnet/pkg/nnet" + internalsession "git.noahlan.cn/noahlan/nnet/v2/internal/session" + "git.noahlan.cn/noahlan/nnet/v2/internal/session/storage" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -265,4 +265,3 @@ func TestSessionManager(t *testing.T) { err = session.Clear() require.NoError(t, err, "Should clear session") } - diff --git a/test/integration/simple_test.go b/test/integration/simple_test.go index 2eed94b..c3ef425 100644 --- a/test/integration/simple_test.go +++ b/test/integration/simple_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -89,7 +89,7 @@ func TestSimpleEcho(t *testing.T) { // 发送请求 testData := "echo\n" t.Logf("Client sending: %q", testData) - + // 使用Send和Receive分开测试 err = client.Send([]byte(testData)) require.NoError(t, err, "Send should succeed") @@ -101,9 +101,8 @@ func TestSimpleEcho(t *testing.T) { // 如果Receive失败,尝试Request resp, err = client.Request([]byte(testData), 3*time.Second) } - + require.NoError(t, err, "Should receive response") t.Logf("Client received: %q", string(resp)) assert.NotEmpty(t, resp, "Response should not be empty") } - diff --git a/test/integration/stress_test.go b/test/integration/stress_test.go index 659dace..830cf62 100644 --- a/test/integration/stress_test.go +++ b/test/integration/stress_test.go @@ -7,10 +7,17 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" ) +func requestPayload(ctx nnet.Context) []byte { + if data := ctx.Request().DataBytes(); len(data) > 0 { + return data + } + return ctx.Request().Raw() +} + // TestHighConcurrencyConnections 测试高并发连接 func TestHighConcurrencyConnections(t *testing.T) { cfg := &nnet.Config{ @@ -92,21 +99,25 @@ func TestHighConcurrencyConnections(t *testing.T) { // 建议:单连接并发控制在较低水平(50-100),高并发应使用多连接 func TestHighConcurrencyRequests(t *testing.T) { cfg := &nnet.Config{ - Addr: "tcp://127.0.0.1:0", // 使用随机端口 + Addr: "tcp://127.0.0.1:0", // 使用随机端口 + ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ - DefaultCodec: "plain", + DefaultCodec: "plain", + EnableProtocolEncode: true, }, } ts := StartTestServerWithRoutes(t, cfg, func(server nnet.Server) { server.Router().RegisterString("echo", func(ctx nnet.Context) error { - data := ctx.Request().Raw() + data := requestPayload(ctx) return ctx.Response().WriteBytes(data) }) }) defer CleanupTestServer(t, ts) - client := NewTestClient(t, ts.Addr, nil) + client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ + ApplicationProtocol: "nnet", + }) defer CleanupTestClient(t, client) ConnectTestClient(t, client) @@ -168,7 +179,7 @@ func TestMemoryUsage(t *testing.T) { ts := StartTestServerWithRoutes(t, cfg, func(server nnet.Server) { server.Router().RegisterString("echo", func(ctx nnet.Context) error { - data := ctx.Request().Raw() + data := requestPayload(ctx) return ctx.Response().WriteBytes(data) }) }) @@ -207,7 +218,10 @@ func TestMemoryUsage(t *testing.T) { runtime.ReadMemStats(&m2) // 计算内存增长 - memAlloc := m2.Alloc - m1.Alloc + memAlloc := int64(m2.Alloc) - int64(m1.Alloc) + if memAlloc < 0 { + memAlloc = 0 + } memTotalAlloc := m2.TotalAlloc - m1.TotalAlloc memMallocs := m2.Mallocs - m1.Mallocs @@ -312,21 +326,25 @@ func TestSustainedLoad(t *testing.T) { // TestLargeMessageStress 测试大消息压力 func TestLargeMessageStress(t *testing.T) { cfg := &nnet.Config{ - Addr: "tcp://127.0.0.1:0", // 使用随机端口 + Addr: "tcp://127.0.0.1:0", // 使用随机端口 + ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ - DefaultCodec: "plain", + DefaultCodec: "plain", + EnableProtocolEncode: true, }, } ts := StartTestServerWithRoutes(t, cfg, func(server nnet.Server) { - server.Router().RegisterString("echo", func(ctx nnet.Context) error { - data := ctx.Request().Raw() + server.Router().RegisterString("echo*", func(ctx nnet.Context) error { + data := requestPayload(ctx) return ctx.Response().WriteBytes(data) }) }) defer CleanupTestServer(t, ts) - client := NewTestClient(t, ts.Addr, nil) + client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ + ApplicationProtocol: "nnet", + }) defer CleanupTestClient(t, client) ConnectTestClient(t, client) @@ -392,9 +410,11 @@ func TestLargeMessageStress(t *testing.T) { // TestMaxConcurrency 测试最大并发连接数(极限测试) func TestMaxConcurrency(t *testing.T) { cfg := &nnet.Config{ - Addr: "tcp://127.0.0.1:0", + Addr: "tcp://127.0.0.1:0", + ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ - DefaultCodec: "plain", + DefaultCodec: "plain", + EnableProtocolEncode: true, }, } @@ -422,13 +442,22 @@ func TestMaxConcurrency(t *testing.T) { go func() { defer wg.Done() - client := NewTestClient(t, ts.Addr, nil) + client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ + ApplicationProtocol: "nnet", + }) defer CleanupTestClient(t, client) - ConnectTestClient(t, client) + if err := client.Connect(); err != nil { + atomic.AddInt64(&errorCount, 1) + return + } // 每个连接发送一个请求 - resp := RequestWithTimeout(t, client, []byte("ping"), 5*time.Second) + resp, err := client.Request([]byte("ping"), 5*time.Second) + if err != nil { + atomic.AddInt64(&errorCount, 1) + return + } if string(resp) == "pong\n" { atomic.AddInt64(&successCount, 1) } else { @@ -462,9 +491,11 @@ func TestMaxConcurrency(t *testing.T) { // TestMaxQPS 测试最大QPS(极限测试) func TestMaxQPS(t *testing.T) { cfg := &nnet.Config{ - Addr: "tcp://127.0.0.1:0", + Addr: "tcp://127.0.0.1:0", + ApplicationProtocol: "nnet", Codec: &nnet.CodecConfig{ - DefaultCodec: "plain", + DefaultCodec: "plain", + EnableProtocolEncode: true, }, } @@ -475,7 +506,9 @@ func TestMaxQPS(t *testing.T) { }) defer CleanupTestServer(t, ts) - client := NewTestClient(t, ts.Addr, nil) + client := NewTestClient(t, ts.Addr, &nnet.ClientConfig{ + ApplicationProtocol: "nnet", + }) defer CleanupTestClient(t, client) ConnectTestClient(t, client) diff --git a/test/integration/timeout_test.go b/test/integration/timeout_test.go index 362d075..88a520e 100644 --- a/test/integration/timeout_test.go +++ b/test/integration/timeout_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -19,7 +19,7 @@ func TestHandlerTimeout(t *testing.T) { listener.Close() cfg := &nnet.Config{ - Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), + Addr: fmt.Sprintf("tcp://127.0.0.1:%d", port), HandlerTimeout: 500 * time.Millisecond, // 设置较短的超时时间 } @@ -93,4 +93,3 @@ func TestReadTimeout(t *testing.T) { resp := RequestWithTimeout(t, client, []byte("test"), 3*time.Second) assert.Contains(t, string(resp), "ok", "Request should succeed") } - diff --git a/test/integration/udp_test.go b/test/integration/udp_test.go index 1948b9c..df586c4 100644 --- a/test/integration/udp_test.go +++ b/test/integration/udp_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -139,4 +139,3 @@ func TestUDPMultipleClients(t *testing.T) { CleanupTestClient(t, client) } } - diff --git a/test/integration/websocket_test.go b/test/integration/websocket_test.go index a30c026..4f45fac 100644 --- a/test/integration/websocket_test.go +++ b/test/integration/websocket_test.go @@ -3,22 +3,17 @@ package integration import ( "encoding/json" "fmt" - "net" "testing" "time" - "github.com/noahlann/nnet/pkg/nnet" + "git.noahlan.cn/noahlan/nnet/v2/pkg/nnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // TestWebSocketServerClient 测试WebSocket服务器-客户端通信 func TestWebSocketServerClient(t *testing.T) { - // 获取随机端口 - listener, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - port := listener.Addr().(*net.TCPAddr).Port - listener.Close() + port := reserveNonEphemeralTCPPort(t) cfg := &nnet.Config{ Addr: fmt.Sprintf("ws://127.0.0.1:%d", port), @@ -85,10 +80,7 @@ func TestWebSocketServerClient(t *testing.T) { // TestWebSocketMultipleClients 测试WebSocket多个客户端 func TestWebSocketMultipleClients(t *testing.T) { - listener, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - port := listener.Addr().(*net.TCPAddr).Port - listener.Close() + port := reserveNonEphemeralTCPPort(t) cfg := &nnet.Config{ Addr: fmt.Sprintf("ws://127.0.0.1:%d", port), @@ -159,10 +151,7 @@ func TestWebSocketMultipleClients(t *testing.T) { // TestWebSocketEcho 测试WebSocket Echo功能 func TestWebSocketEcho(t *testing.T) { - listener, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - port := listener.Addr().(*net.TCPAddr).Port - listener.Close() + port := reserveNonEphemeralTCPPort(t) cfg := &nnet.Config{ Addr: fmt.Sprintf("ws://127.0.0.1:%d", port), @@ -174,7 +163,7 @@ func TestWebSocketEcho(t *testing.T) { server, err := nnet.NewWebSocketServer(cfg) require.NoError(t, err) - server.Router().RegisterString("echo", func(ctx nnet.Context) error { + server.Router().RegisterString("echo*", func(ctx nnet.Context) error { data := ctx.Request().Raw() return ctx.Response().Write(map[string]any{ "echo": string(data), @@ -219,7 +208,6 @@ func TestWebSocketEcho(t *testing.T) { var result map[string]any err = json.Unmarshal(resp, &result) - assert.NoError(t, err, "Response should be valid JSON") + require.NoError(t, err, "Response should be valid JSON") assert.Contains(t, result["echo"].(string), "test message", "Echo should contain the message") } - diff --git a/test/router_test.go b/test/router_test.go index 2faa46a..b931c1f 100644 --- a/test/router_test.go +++ b/test/router_test.go @@ -4,12 +4,12 @@ import ( "context" "testing" - "github.com/noahlann/nnet/internal/codec" - "github.com/noahlann/nnet/internal/request" - "github.com/noahlann/nnet/internal/response" - internalrouter "github.com/noahlann/nnet/internal/router" - ctxpkg "github.com/noahlann/nnet/pkg/context" - "github.com/noahlann/nnet/pkg/router" + "git.noahlan.cn/noahlan/nnet/v2/internal/codec" + "git.noahlan.cn/noahlan/nnet/v2/internal/request" + "git.noahlan.cn/noahlan/nnet/v2/internal/response" + internalrouter "git.noahlan.cn/noahlan/nnet/v2/internal/router" + ctxpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/context" + "git.noahlan.cn/noahlan/nnet/v2/pkg/router" ) func TestStringMatcher(t *testing.T) {