commit f0f542126232aa54744eb9dfd8659752a005ea9a Author: W11 Date: Sat Sep 19 18:22:50 2026 +0800 init: 通用库抽离——httpx/logger/wsc/wscclient/jwtx(合并分叉副本); 顶层包布局; module path git.zeroonesoft.cn/golib/zogo diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..4df69a3 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +.idea/ +*.exe diff --git a/README.md b/README.md new file mode 100644 index 0000000..0a5050b --- /dev/null +++ b/README.md @@ -0,0 +1,28 @@ +# zogo + +zeroone Go 通用库(单仓多包)。被各后端服务以 Go module 依赖引用: +`go get git.zeroonesoft.cn/golib/zogo/xxx` + +## 包清单 + +| 包 | 用途 | 主要依赖 | +|---|---|---| +| `httpx` | Gin 统一出入参:OkJson/ErrorJson(CodeMsg 错误码透传)、Handle/HandleResult、Parse(uri→query→header→json 顺序绑定 + binding 校验收口) | gin, validator | +| `logger` | logrus 日志初始化(文件轮转) | logrus, rotatelogs | +| `wsc` | WebSocket 服务框架(会话/房间/路由/反射绑定) | gin, gorilla/websocket, uuid, logrus | +| `wscclient` | WebSocket 客户端(自动重连) | gorilla/websocket | +| `jwtx` | 轻量 HMAC-SHA256 JWT 签发与校验 | 无 | + +## 约定 + +- 通用性准入:被 ≥2 个服务复制使用过的包才进入本仓库;仅单服务使用的留在服务内。 +- 改动流程:在本仓库修复/演进 → 打 tag(如 v0.1.1)→ 各服务 `go get -u git.zeroonesoft.cn/golib/zogo@v0.1.1`。 +- 本地联动开发:在 `D:\zomaintain\backend\go.work` 中 `use D:/golib/zogo`(go.work 不提交)。 +- go.mod 声明 go 1.24(不高于最低版本消费方),不要随意上调。 + +## 私有模块拉取配置(每台开发机/CI 一次) + +``` +go env -w GOPRIVATE=git.zeroonesoft.cn/* +``` +git 凭据需能访问 git.zeroonesoft.cn(Git Credential Manager 存 PAT)。 diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..6a6f68b --- /dev/null +++ b/go.mod @@ -0,0 +1,47 @@ +module git.zeroonesoft.cn/golib/zogo + +go 1.25.0 + +require ( + github.com/gin-gonic/gin v1.12.0 + github.com/go-playground/validator/v10 v10.30.4 + github.com/google/uuid v1.6.0 + github.com/gorilla/websocket v1.5.3 + github.com/lestrrat-go/file-rotatelogs v2.4.0+incompatible + github.com/mattn/go-colorable v0.1.15 + github.com/sirupsen/logrus v1.10.2 +) + +require ( + github.com/bytedance/gopkg v0.1.3 // indirect + github.com/bytedance/sonic v1.15.0 // indirect + github.com/bytedance/sonic/loader v0.5.0 // indirect + github.com/cloudwego/base64x v0.1.6 // indirect + github.com/gabriel-vasile/mimetype v1.4.15 // indirect + github.com/gin-contrib/sse v1.1.0 // indirect + github.com/go-playground/locales v0.14.1 // indirect + github.com/go-playground/universal-translator v0.18.1 // indirect + github.com/goccy/go-json v0.10.5 // indirect + github.com/goccy/go-yaml v1.19.2 // indirect + github.com/jonboulle/clockwork v0.5.0 // indirect + github.com/json-iterator/go v1.1.12 // indirect + github.com/klauspost/cpuid/v2 v2.3.0 // indirect + github.com/leodido/go-urn v1.5.0 // indirect + github.com/lestrrat-go/strftime v1.2.0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect + github.com/modern-go/reflect2 v1.0.2 // indirect + github.com/pelletier/go-toml/v2 v2.2.4 // indirect + github.com/pkg/errors v0.9.1 // indirect + github.com/quic-go/qpack v0.6.0 // indirect + github.com/quic-go/quic-go v0.59.0 // indirect + github.com/twitchyliquid64/golang-asm v0.15.1 // indirect + github.com/ugorji/go/codec v1.3.1 // indirect + go.mongodb.org/mongo-driver/v2 v2.5.0 // indirect + golang.org/x/arch v0.22.0 // indirect + golang.org/x/crypto v0.55.0 // indirect + golang.org/x/net v0.57.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect + google.golang.org/protobuf v1.36.10 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..0c0ea6d --- /dev/null +++ b/go.sum @@ -0,0 +1,107 @@ +github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M= +github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM= +github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE= +github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k= +github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE= +github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo= +github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M= +github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/gabriel-vasile/mimetype v1.4.15 h1:05iP/CYtZ/w455R/KZM6rZ5ieAdh99UPtd+d3YzLmaI= +github.com/gabriel-vasile/mimetype v1.4.15/go.mod h1:azpTcoLcDZRNgFou5j+APrqQx9HqVPWa6ijYQIIVswQ= +github.com/gin-contrib/sse v1.1.0 h1:n0w2GMuUpWDVp7qSpvze6fAu9iRxJY4Hmj6AmBOU05w= +github.com/gin-contrib/sse v1.1.0/go.mod h1:hxRZ5gVpWMT7Z0B0gSNYqqsSCNIJMjzvm6fqCz9vjwM= +github.com/gin-gonic/gin v1.12.0 h1:b3YAbrZtnf8N//yjKeU2+MQsh2mY5htkZidOM7O0wG8= +github.com/gin-gonic/gin v1.12.0/go.mod h1:VxccKfsSllpKshkBWgVgRniFFAzFb9csfngsqANjnLc= +github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s= +github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= +github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA= +github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY= +github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY= +github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY= +github.com/go-playground/validator/v10 v10.30.4 h1:9Rcod2ZPO6mOEG6b4GqyoHE/H6//Ze0RuhOo1hT1x0w= +github.com/go-playground/validator/v10 v10.30.4/go.mod h1:numpT+RPLE91R9oYWMY/R9zRgJBewr3IXHko4OISPpk= +github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= +github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= +github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= +github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= +github.com/jonboulle/clockwork v0.5.0 h1:Hyh9A8u51kptdkR+cqRpT1EebBwTn1oK9YfGYbdFz6I= +github.com/jonboulle/clockwork v0.5.0/go.mod h1:3mZlmanh0g2NDKO5TWZVJAfofYk64M7XN3SzBPjZF60= +github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= +github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= +github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= +github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= +github.com/leodido/go-urn v1.5.0 h1:pLqT2kq1zpHW/1D18QMjMpdtX7cekxqtJJjg5ANyWw0= +github.com/leodido/go-urn v1.5.0/go.mod h1:9BORnCDhdPBJNDEX+w1bJisa8yOKYi116VeO96s4ifE= +github.com/lestrrat-go/envload v0.0.0-20180220234015-a3eb8ddeffcc h1:RKf14vYWi2ttpEmkA4aQ3j4u9dStX2t4M8UM6qqNsG8= +github.com/lestrrat-go/envload v0.0.0-20180220234015-a3eb8ddeffcc/go.mod h1:kopuH9ugFRkIXf3YoqHKyrJ9YfUFsckUU9S7B+XP+is= +github.com/lestrrat-go/file-rotatelogs v2.4.0+incompatible h1:Y6sqxHMyB1D2YSzWkLibYKgg+SwmyFU9dF2hn6MdTj4= +github.com/lestrrat-go/file-rotatelogs v2.4.0+incompatible/go.mod h1:ZQnN8lSECaebrkQytbHj4xNgtg8CR7RYXnPok8e0EHA= +github.com/lestrrat-go/strftime v1.2.0 h1:8fAUYOeaJKCuLzNvUWBAo8t6I6hkFfodDTndEzJIun0= +github.com/lestrrat-go/strftime v1.2.0/go.mod h1:GtsIA/7ddIGJjEdfadUafEb1sbutvlvpMdPCMglykYo= +github.com/mattn/go-colorable v0.1.15 h1:+u9SLTRGnXv73cEsnsmoZBom+dMU88B2M0aDcWy0/jY= +github.com/mattn/go-colorable v0.1.15/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= +github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= +github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= +github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= +github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= +github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= +github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= +github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= +github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= +github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= +github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw= +github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU= +github.com/sirupsen/logrus v1.10.2 h1:G2SED73/qrAu6YwbdxOD6peLkCBI3z7L+ykJFTXJBBo= +github.com/sirupsen/logrus v1.10.2/go.mod h1:SLEg8TqYulVKKfIGHldVp2K2aYz2DKSVBq4g/H5bR7Q= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= +github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= +github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= +github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= +github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI= +github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= +github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY= +github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= +go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE= +go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= +go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU= +go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= +go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= +golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI= +golang.org/x/arch v0.22.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A= +golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= +golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE= +google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/httpx/errors.go b/httpx/errors.go new file mode 100644 index 0000000..4dd3d9e --- /dev/null +++ b/httpx/errors.go @@ -0,0 +1,18 @@ +package httpx + +import "fmt" + +// CodeMsg 包含 code 与 msg 的错误结构,实现了 error 接口。 +type CodeMsg struct { + Code int + Msg string +} + +func (c *CodeMsg) Error() string { + return fmt.Sprintf("code: %d, msg: %s", c.Code, c.Msg) +} + +// New 创建一个 CodeMsg 错误。 +func New(code int, msg string) error { + return &CodeMsg{Code: code, Msg: msg} +} diff --git a/httpx/handle.go b/httpx/handle.go new file mode 100644 index 0000000..8568b4a --- /dev/null +++ b/httpx/handle.go @@ -0,0 +1,28 @@ +package httpx + +import "github.com/gin-gonic/gin" + +// Handle 泛型 handler 收口: Parse 取参 → 调 logic → 统一出参。 +// 试点新约定: 替代旧模板逐 handler 复制 Parse/OkJson/ErrorJson 的样板。 +func Handle[TReq any, TResp any](c *gin.Context, fn func(req *TReq) (*TResp, error)) { + var req TReq + if err := Parse(c, &req); err != nil { + ErrorJson(c, err) + return + } + resp, err := fn(&req) + if err != nil { + ErrorJson(c, err) + return + } + OkJson(c, resp) +} + +// HandleResult 统一收口:err 非 nil 走 ErrorJson,否则 OkJson(data)。 +func HandleResult(c *gin.Context, data any, err error) { + if err != nil { + ErrorJson(c, err) + return + } + OkJson(c, data) +} diff --git a/httpx/httpx.go b/httpx/httpx.go new file mode 100644 index 0000000..099f3ad --- /dev/null +++ b/httpx/httpx.go @@ -0,0 +1,35 @@ +// Package httpx 提供 Gin 服务的统一出入参收口:统一 JSON 响应格式、 +// 错误码透出(CodeMsg)与请求参数绑定校验(Parse)。 +package httpx + +import ( + "errors" + "net/http" + + "github.com/gin-gonic/gin" +) + +// OkJson 以 200 OK 返回成功响应,统一格式:{code:0, msg:"OK", data:...} +func OkJson(c *gin.Context, data any) { + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "msg": "OK", + "data": data, + }) +} + +// ErrorJson 统一错误出参: CodeMsg 携带的错误码透出(如 401), 其余折叠为 code 1。 +func ErrorJson(c *gin.Context, err error) { + code := 1 + msg := err.Error() + var cm *CodeMsg + if errors.As(err, &cm) { + code = cm.Code + msg = cm.Msg + } + c.JSON(http.StatusOK, gin.H{ + "code": code, + "msg": msg, + "data": nil, + }) +} diff --git a/httpx/parse.go b/httpx/parse.go new file mode 100644 index 0000000..507f726 --- /dev/null +++ b/httpx/parse.go @@ -0,0 +1,75 @@ +package httpx + +import ( + "encoding/json" + + "github.com/gin-gonic/gin" + "github.com/gin-gonic/gin/binding" + "github.com/go-playground/validator/v10" +) + +// bindingValidator 与 gin 同源: 校验 struct 的 binding tag +var bindingValidator = validator.New(validator.WithRequiredStructEnabled()) + +func init() { + bindingValidator.SetTagName("binding") + // 关闭 gin 内置校验 (ShouldBindQuery/Header/JSON 内部各自 validate, 会使 json 字段的 + // required 在 query 阶段误报)。校验统一由 Parse 末尾的 bindingValidator 收口。 + binding.Validator = nil +} + +// Parse 按统一顺序绑定请求参数: uri → query → header → json body。 +// 各来源只做绑定不做校验, 全部绑定完后统一跑 binding 校验。 +// (修复: gin 的 ShouldBindQuery 内部即执行校验, 会使 json 字段的 required 在 query 阶段误报) +func Parse(c *gin.Context, v any) error { + if err := parseUri(c, v); err != nil { + return err + } + if err := parseQuery(c, v); err != nil { + return err + } + if err := parseHeader(c, v); err != nil { + return err + } + if err := parseJSON(c, v); err != nil { + return err + } + return validate(c, v) +} + +func parseUri(c *gin.Context, v any) error { + if _, ok := c.Params.Get("id"); ok { + return c.ShouldBindUri(v) + } + return nil +} + +func parseQuery(c *gin.Context, v any) error { + return c.ShouldBindQuery(v) +} + +func parseHeader(c *gin.Context, v any) error { + return c.ShouldBindHeader(v) +} + +// parseJSON 直接解码 body, 不触发 gin 内部校验 (校验统一在 Parse 末尾) +func parseJSON(c *gin.Context, v any) error { + if c.Request == nil || c.Request.Body == nil || c.Request.ContentLength == 0 { + return nil + } + return json.NewDecoder(c.Request.Body).Decode(v) +} + +// Validator 请求结构体可实现此接口做自定义校验, 在 binding 校验之前调用 +type Validator interface { + Validate() error +} + +func validate(c *gin.Context, v any) error { + if vv, ok := v.(Validator); ok { + if err := vv.Validate(); err != nil { + return err + } + } + return bindingValidator.Struct(v) +} diff --git a/jwtx/jwt.go b/jwtx/jwt.go new file mode 100644 index 0000000..9f02009 --- /dev/null +++ b/jwtx/jwt.go @@ -0,0 +1,61 @@ +// Package jwtx 轻量 JWT(HMAC-SHA256),无外部依赖。 +package jwtx + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "strings" + "time" +) + +// Claims 登录令牌载荷。 +type Claims struct { + Uid int64 `json:"uid"` + Exp int64 `json:"exp"` + Msg string `json:"msg,omitempty"` +} + +var secret = []byte("zo-maintain-jwt-secret-2026") + +// Sign 签发 token(TTL 秒)。 +func Sign(uid int64, ttlSeconds int64) (string, error) { + payload, err := json.Marshal(Claims{Uid: uid, Exp: time.Now().Add(time.Duration(ttlSeconds) * time.Second).Unix()}) + if err != nil { + return "", err + } + body := base64.RawURLEncoding.EncodeToString(payload) + sig := sign(body) + return body + "." + sig, nil +} + +// Parse 校验并解析 token。 +func Parse(token string) (*Claims, error) { + parts := strings.SplitN(token, ".", 2) + if len(parts) != 2 { + return nil, errors.New("令牌格式错误") + } + if sign(parts[0]) != parts[1] { + return nil, errors.New("令牌签名校验失败") + } + payload, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + return nil, errors.New("令牌载荷解析失败") + } + var c Claims + if err = json.Unmarshal(payload, &c); err != nil { + return nil, errors.New("令牌载荷解析失败") + } + if c.Exp < time.Now().Unix() { + return nil, errors.New("令牌已过期") + } + return &c, nil +} + +func sign(body string) string { + mac := hmac.New(sha256.New, secret) + mac.Write([]byte(body)) + return base64.RawURLEncoding.EncodeToString(mac.Sum(nil)) +} diff --git a/logger/logger.go b/logger/logger.go new file mode 100644 index 0000000..13ca692 --- /dev/null +++ b/logger/logger.go @@ -0,0 +1,308 @@ +// Package logger 提供统一的日志初始化功能 +// +// 使用示例: +// +// logger.InitLog(logrus.DebugLevel, true) +// logger.SetProjectPrefix("D:/Projects/MyProject/") +// logrus.Debug("调试信息") +// logrus.Infof("信息: %s", "test") +package logger + +import ( + "archive/zip" + "io" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "sync" + "time" + + "github.com/lestrrat-go/file-rotatelogs" + "github.com/mattn/go-colorable" + "github.com/sirupsen/logrus" +) + +var ( + // 项目根目录前缀,用于简化日志中的调用文件路径 + // 可以通过 SetProjectPrefix 自定义设置 + projectPrefix string + projectPrefixOnce sync.Once +) + +// SetProjectPrefix 设置项目路径前缀,用于简化日志中的文件路径显示 +// 例如:SetProjectPrefix("D:/Projects/ZoSmartCybercafe/service-server-new/") +func SetProjectPrefix(prefix string) { + projectPrefix = prefix +} + +// getProjectPrefix 获取项目路径前缀(懒加载) +func getProjectPrefix() string { + projectPrefixOnce.Do(func() { + // 默认前缀,如果没有通过 SetProjectPrefix 设置,则使用空字符串 + if projectPrefix == "" { + projectPrefix = "" + } + }) + return projectPrefix +} + +// fileLogHook 写入按天轮转的日志文件的 Hook +type fileLogHook struct { + writer io.Writer + formatter logrus.Formatter + levels []logrus.Level // 此 Hook 生效的日志级别 +} + +// NewFileLogHook 创建按天轮转的日志 Hook +// 参数: +// - logDir: 日志目录 +// - filename: 日志文件名模板(如 "app-%Y-%m-%d.log") +// - levels: 此 Hook 生效的日志级别,如果为空则对所有级别生效 +func NewFileLogHook(logDir string, filename string, levels []logrus.Level) (*fileLogHook, error) { + // 创建日志目录 + if err := os.MkdirAll(logDir, 0755); err != nil { + return nil, err + } + + // 创建按天轮转的日志文件 + logPath := filepath.Join(logDir, filename) + rl, err := rotatelogs.New( + logPath, + rotatelogs.WithMaxAge(-1), // 不自动删除旧日志 + rotatelogs.WithRotationTime(24*time.Hour), // 24小时轮转一次 + ) + if err != nil { + return nil, err + } + + // 文件日志使用的 formatter(不带颜色) + formatter := &logrus.TextFormatter{ + TimestampFormat: "2006-01-02 15:04:05", + FullTimestamp: true, + DisableColors: true, // 文件中不使用颜色 + DisableLevelTruncation: true, // 不缩略 INFO -> INF + DisableSorting: true, + CallerPrettyfier: func(f *runtime.Frame) (string, string) { + file := f.File + prefix := getProjectPrefix() + if prefix != "" && strings.HasPrefix(file, prefix) { + file = file[len(prefix):] + } + // 返回 "文件:行号",file= 字段内不要前导空格 + return "", file + ":" + strconv.Itoa(f.Line) + }, + } + + // 如果没有指定级别,则对所有级别生效 + if len(levels) == 0 { + levels = logrus.AllLevels + } + + return &fileLogHook{ + writer: rl, + formatter: formatter, + levels: levels, + }, nil +} + +// Fire 每次日志记录时调用,写入到文件 +func (hook *fileLogHook) Fire(entry *logrus.Entry) error { + line, err := hook.formatter.Format(entry) + if err != nil { + return err + } + _, err = hook.writer.Write(line) + return err +} + +// Levels 返回此 Hook 生效的日志级别 +func (hook *fileLogHook) Levels() []logrus.Level { + return hook.levels +} + +// InitLog 初始化 logrus 日志配置 +// 参数: +// - level: 日志级别,例如 logrus.DebugLevel, logrus.InfoLevel +// - reportCaller: 是否显示调用者信息(文件名和行号) +// - logDir: 日志文件目录,如果为空则不输出到文件 +// +// 注意:初始化后,在其他文件中直接使用 logrus.Debug、logrus.Info 等函数, +// 不要通过 logger.Debug 等方式调用,否则会导致调用位置信息不准确。 +// +// 日志文件: +// - app-YYYY-MM-DD.log: 所有级别的日志 +// - error-YYYY-MM-DD.log: 只记录 Error 及以上级别的日志 +func InitLog(level logrus.Level, reportCaller bool, logDir ...string) { + logrus.SetOutput(colorable.NewColorableStdout()) + logrus.SetLevel(level) + logrus.SetReportCaller(reportCaller) + logrus.SetFormatter(&logrus.TextFormatter{ + TimestampFormat: "2006-01-02 15:04:05", + FullTimestamp: true, + ForceColors: true, // 强制输出颜色,配合 colorable 在 Windows 上显示 + DisableLevelTruncation: true, // 不缩略 INFO -> INF + DisableSorting: true, + CallerPrettyfier: func(f *runtime.Frame) (string, string) { + file := f.File + prefix := getProjectPrefix() + if prefix != "" && strings.HasPrefix(file, prefix) { + file = file[len(prefix):] + } + // 返回 "文件:行号",file= 字段内不要前导空格 + return "", " " + file + ":" + strconv.Itoa(f.Line) + }, + }) + + // 如果指定了日志目录,则添加文件输出 + if len(logDir) > 0 && logDir[0] != "" { + // 创建主日志 Hook(记录所有级别) + mainHook, err := NewFileLogHook(logDir[0], "app-%Y-%m-%d.log", logrus.AllLevels) + if err == nil { + logrus.AddHook(mainHook) + } else { + logrus.Warnf("创建主日志文件失败: %v", err) + } + + // 创建错误日志 Hook(只记录 Error 及以上级别) + errorHook, err := NewFileLogHook(logDir[0], "error-%Y-%m-%d.log", []logrus.Level{ + logrus.ErrorLevel, + logrus.FatalLevel, + logrus.PanicLevel, + }) + if err == nil { + logrus.AddHook(errorHook) + } else { + logrus.Warnf("创建错误日志文件失败: %v", err) + } + + // 启动后台历史日志压缩协程(每日 0:00-0:30 为不压缩窗口) + StartLogCompressor(logDir[0], time.Hour) + } +} + +// 历史日志压缩相关常量 +const ( + // compressSkipStartHour / compressSkipEndMin 定义每日不压缩窗口:0:00 - 0:30 + // 该窗口通常与日志轮转时间重合,期间不执行压缩,避免与正在写入/轮转的日志冲突 + compressSkipStartHour = 0 + compressSkipEndMin = 30 +) + +// inCompressSkipWindow 判断当前时间是否处于每日不压缩窗口(0:00 - 0:30) +func inCompressSkipWindow(now time.Time) bool { + h, m, _ := now.Clock() + return h == compressSkipStartHour && m < compressSkipEndMin +} + +// StartLogCompressor 在后台启动历史日志压缩协程,按 interval 周期扫描并压缩历史日志。 +// 每日 0:00-0:30 为不压缩窗口,期间直接跳过。 +// 参数: +// - logDir: 日志目录 +// - interval: 扫描周期,<=0 时默认 1 小时 +func StartLogCompressor(logDir string, interval time.Duration) { + if logDir == "" { + return + } + if interval <= 0 { + interval = time.Hour + } + go func() { + // 启动时先执行一次,立即清理已有历史日志 + _ = CompressHistoryLogs(logDir) + ticker := time.NewTicker(interval) + defer ticker.Stop() + for range ticker.C { + _ = CompressHistoryLogs(logDir) + } + }() +} + +// CompressHistoryLogs 将历史日志(非当天、且尚未压缩的 .log 文件)逐个压缩为同名的 .zip 文件, +// 压缩成功后删除原始日志文件以完成清理。每日 0:00-0:30 为不压缩窗口,期间直接跳过。 +// 参数 logDir 为日志目录。 +func CompressHistoryLogs(logDir string) error { + // 不压缩窗口内直接跳过 + if inCompressSkipWindow(time.Now()) { + return nil + } + + entries, err := os.ReadDir(logDir) + if err != nil { + return err + } + + // 当天的日期串,用于跳过正在写入的当天日志 + today := time.Now().Format("2006-01-02") + + for _, entry := range entries { + if entry.IsDir() { + continue + } + name := entry.Name() + + // 仅处理 .log 文件 + if !strings.HasSuffix(name, ".log") { + continue + } + // 跳过当天的日志(正在写入,不压缩) + if strings.Contains(name, today) { + continue + } + + logPath := filepath.Join(logDir, name) + // 文件名不变,仅追加 .zip 后缀(如 app-2026-07-16.log -> app-2026-07-16.log.zip) + zipPath := logPath + ".zip" + + // 若已存在同名 zip,说明此前已压缩成功,直接清理原日志避免重复压缩 + if _, err := os.Stat(zipPath); err == nil { + if rmErr := os.Remove(logPath); rmErr != nil { + logrus.Warnf("清理已压缩的历史日志失败 %s: %v", name, rmErr) + } + continue + } + + if err := zipSingleFile(logPath, zipPath); err != nil { + logrus.Warnf("压缩历史日志失败 %s: %v", name, err) + continue + } + + // 压缩成功后清理原始日志文件 + if err := os.Remove(logPath); err != nil { + logrus.Warnf("删除历史日志失败 %s: %v", name, err) + } else { + logrus.Infof("历史日志已压缩并清理: %s -> %s", name, filepath.Base(zipPath)) + } + } + return nil +} + +// zipSingleFile 将单个文件 src 压缩写入 dst(zip 包内仅包含原始文件名,文件名保持不变) +func zipSingleFile(src, dst string) error { + in, err := os.Open(src) + if err != nil { + return err + } + defer in.Close() + + out, err := os.Create(dst) + if err != nil { + return err + } + defer out.Close() + + zw := zip.NewWriter(out) + w, err := zw.Create(filepath.Base(src)) + if err != nil { + _ = zw.Close() + return err + } + + if _, err := io.Copy(w, in); err != nil { + _ = zw.Close() + return err + } + + return zw.Close() +} diff --git a/wsc/context.go b/wsc/context.go new file mode 100644 index 0000000..f64de80 --- /dev/null +++ b/wsc/context.go @@ -0,0 +1,187 @@ +package wsc + +import ( + "context" + "sync" + "time" + + "github.com/gin-gonic/gin" + "github.com/google/uuid" +) + +// Context WebSocket 上下文,实现 context.Context 接口。 +// 在连接建立时从 gin.Context 提取元数据,贯穿整个连接生命周期。 +type Context struct { + // === 从 gin 提取的元数据 === + // 公开字段仅供读取,正常业务流程不应修改。 + SessionID string `json:"sessionId"` + ClientIP string `json:"clientIp"` + UserAgent string `json:"userAgent"` + Path string `json:"path"` + TraceID string `json:"traceId"` + + // UID 由认证中间件注入,业务层可直接读取。 + UID int64 `json:"uid,omitempty"` + + // === 内部字段 === + ctx context.Context + cancel context.CancelFunc + + session *Session + + mu sync.RWMutex + values map[string]any +} + +// newContext 从 gin.Context 创建 ws 上下文。 +func newContext(c *gin.Context, session *Session) *Context { + ctx, cancel := context.WithCancel(context.Background()) + + wsCtx := &Context{ + SessionID: uuid.New().String(), + ClientIP: c.ClientIP(), + UserAgent: c.GetHeader("User-Agent"), + Path: c.FullPath(), + TraceID: c.GetHeader("X-Trace-Id"), + + ctx: ctx, + cancel: cancel, + session: session, + values: make(map[string]any), + } + + // 透传 gin 上下文中的自定义值(如 OnUpgrade 中设置的 userId/userName) + // gin v1.12.0 的 Keys 为 map[any]any,需将 key 断言为 string。 + if c.Keys != nil { + for k, v := range c.Keys { + if key, ok := k.(string); ok { + wsCtx.values[key] = v + } + } + } + + return wsCtx +} + +// ============================================================ +// context.Context 接口实现 +// ============================================================ + +func (c *Context) Deadline() (time.Time, bool) { return c.ctx.Deadline() } +func (c *Context) Done() <-chan struct{} { return c.ctx.Done() } +func (c *Context) Err() error { return c.ctx.Err() } + +func (c *Context) Value(key any) any { + if s, ok := key.(string); ok { + c.mu.RLock() + v, exists := c.values[s] + c.mu.RUnlock() + if exists { + return v // 含 nil + } + } + return c.ctx.Value(key) +} + +// ============================================================ +// 写回方法 +// ============================================================ + +// Write 发送 JSON 消息(原始数据,不带 type 包装)。 +func (c *Context) Write(v any) error { + return c.session.writeJSON(v) +} + +// WriteMessage 发送带 type 的结构化消息。 +func (c *Context) WriteMessage(msgType string, data any) error { + return c.Write(messageOut(msgType, data)) +} + +// WriteError 发送错误消息。 +func (c *Context) WriteError(code int, msg string) error { + return c.WriteMessage(WsActionError, map[string]any{ + "code": code, + "message": msg, + }) +} + +// ============================================================ +// 自定义值存取 +// ============================================================ + +func (c *Context) Set(key string, value any) { + c.mu.Lock() + c.values[key] = value + c.mu.Unlock() +} + +func (c *Context) Get(key string) (any, bool) { + c.mu.RLock() + v, ok := c.values[key] + c.mu.RUnlock() + return v, ok +} + +func (c *Context) GetString(key string) string { + v, _ := c.Get(key) + s, _ := v.(string) + return s +} + +func (c *Context) GetInt64(key string) int64 { + v, _ := c.Get(key) + n, _ := v.(int64) + return n +} + +// ============================================================ +// 房间操作 +// ============================================================ + +func (c *Context) JoinRoom(name string) { + c.session.server.GetOrCreateRoom(name).Join(c) +} + +func (c *Context) LeaveRoom(name string) { + room := c.session.server.GetRoom(name) + if room != nil { + room.Leave(c) + } +} + +func (c *Context) BroadcastToRoom(roomName string, msgType string, data any) error { + room := c.session.server.GetRoom(roomName) + if room == nil { + return nil + } + return room.BroadcastRaw(c.SessionID, messageOut(msgType, data)) +} + +func (c *Context) GetRoom(name string) *Room { + return c.session.server.GetRoom(name) +} + +// ============================================================ +// 生命周期 +// ============================================================ + +func (c *Context) Close() { + c.session.close() +} + +func (c *Context) cancelCtx() { + c.cancel() +} + +// ============================================================ +// 内部辅助 +// ============================================================ + +// messageOut 构造写出的消息 map(单次序列化)。 +func messageOut(msgType string, data any) map[string]any { + m := map[string]any{"action": msgType} + if data != nil { + m["payload"] = data + } + return m +} diff --git a/wsc/room.go b/wsc/room.go new file mode 100644 index 0000000..097546a --- /dev/null +++ b/wsc/room.go @@ -0,0 +1,103 @@ +package wsc + +import ( + "encoding/json" + "sync" + + "github.com/sirupsen/logrus" +) + +// Room 消息房间,用于业务隔离和群组广播。 +type Room struct { + name string + members map[string]*Session + mu sync.RWMutex + onEmpty func(name string) // 房间空时回调 +} + +func newRoom(name string) *Room { + return &Room{ + name: name, + members: make(map[string]*Session), + } +} + +func (r *Room) Name() string { return r.name } + +func (r *Room) Len() int { + r.mu.RLock() + defer r.mu.RUnlock() + return len(r.members) +} + +func (r *Room) Join(ctx *Context) { + r.mu.Lock() + r.members[ctx.SessionID] = ctx.session + r.mu.Unlock() +} + +func (r *Room) Leave(ctx *Context) { + r.mu.Lock() + delete(r.members, ctx.SessionID) + empty := len(r.members) == 0 + r.mu.Unlock() + + if empty && r.onEmpty != nil { + r.onEmpty(r.name) + } +} + +// Broadcast 向房间成员广播 Message(json.Marshal 后发原始字节)。 +func (r *Room) Broadcast(excludeSessionID string, msg Message) error { + m := map[string]any{"action": msg.Action} + if len(msg.Payload) > 0 { + // Payload 已是 json.RawMessage,直接复用 + m["payload"] = msg.Payload + } + raw, err := json.Marshal(m) + if err != nil { + return err + } + return r.broadcastRaw(excludeSessionID, raw) +} + +// BroadcastRaw 向房间成员广播已序列化的消息(性能优化)。 +func (r *Room) BroadcastRaw(excludeSessionID string, data map[string]any) error { + raw, err := json.Marshal(data) + if err != nil { + return err + } + return r.broadcastRaw(excludeSessionID, raw) +} + +// broadcastRaw 向房间成员发送原始字节(锁外已序列化)。 +func (r *Room) broadcastRaw(excludeSessionID string, raw []byte) error { + r.mu.RLock() + defer r.mu.RUnlock() + + for id, s := range r.members { + if id == excludeSessionID { + continue + } + if err := s.sendRaw(raw); err != nil { + logrus.Errorf("[wsc] broadcast send failed: session=%s, err=%v", id[:8], err) + } + } + return nil +} + +// BroadcastAll 向房间所有成员广播。 +func (r *Room) BroadcastAll(msg Message) error { + return r.Broadcast("", msg) +} + +// SendTo 向房间内指定成员发送消息。 +func (r *Room) SendTo(sessionID string, msg Message) error { + r.mu.RLock() + s, ok := r.members[sessionID] + r.mu.RUnlock() + if !ok { + return nil + } + return s.writeJSON(msg) +} diff --git a/wsc/router.go b/wsc/router.go new file mode 100644 index 0000000..862cac6 --- /dev/null +++ b/wsc/router.go @@ -0,0 +1,357 @@ +package wsc + +import ( + "encoding/json" + "reflect" + + "github.com/sirupsen/logrus" +) + +// ============================================================ +// 框架级消息 action 常量 +// 业务消息(ping/pong/register 等)由各模块自行定义,不放在框架层。 +// ============================================================ +const ( + // WsActionError 错误消息(框架统一回包 action) + WsActionError = "error" +) + +// ============================================================ +// MessageHandler — 消息处理器签名 +// ============================================================ + +// MessageHandler 业务消息处理函数。 +// ctx: WS 上下文 data: 已解析的 JSON data 字段。 +// 返回的 any 会自动 JSON 序列化后通过 ctx.Write() 发回。 +type MessageHandler func(ctx *Context, data json.RawMessage) (any, error) + +// ============================================================ +// MiddlewareFunc — 中间件签名 +// ============================================================ + +// MiddlewareFunc 消息级中间件。 +// 返回 error 时中断链路,错误消息自动发送给客户端。 +type MiddlewareFunc func(ctx *Context, msg *Message) error + +// ============================================================ +// Router — 消息路由器 +// ============================================================ + +// Router 消息路由器,按 type 字段分发到不同的 MessageHandler。 +// 支持中间件链(类似 gin)。 +type Router struct { + routes map[string]MessageHandler + middlewares []MiddlewareFunc +} + +// NewRouter 创建路由器。 +func NewRouter() *Router { + return &Router{ + routes: make(map[string]MessageHandler), + } +} + +// On 注册指定 type 的消息处理器。 +// handler 签名:func(ctx *Context, data json.RawMessage) (any, error) +func (r *Router) On(msgType string, handler MessageHandler) { + r.routes[msgType] = handler +} + +// Use 添加消息级中间件。按添加顺序执行。 +func (r *Router) Use(mw ...MiddlewareFunc) { + r.middlewares = append(r.middlewares, mw...) +} + +// dispatch 内部消息分发。先走中间件链,再走路由。 +func (r *Router) dispatch(ctx *Context, msg *Message) { + // 中间件链 + for _, mw := range r.middlewares { + if err := mw(ctx, msg); err != nil { + // 中间件阻断 + if wErr, ok := err.(*Error); ok { + ctx.WriteError(wErr.Code, wErr.Message) + } else { + ctx.WriteError(403, err.Error()) + } + return + } + } + + // 路由分发 + handler, ok := r.routes[msg.Action] + if !ok { + logrus.Warnf("[wsc] unknown message type: %s, session=%s, ip=%s", msg.Action, ctx.SessionID, ctx.ClientIP) + ctx.WriteError(404, "unknown message type: "+msg.Action) + return + } + + resp, err := handler(ctx, msg.Payload) + if err != nil { + if wErr, ok := err.(*Error); ok { + ctx.WriteError(wErr.Code, wErr.Message) + } else { + ctx.WriteError(500, err.Error()) + } + return + } + + // 有返回值才写回 + if resp != nil { + if rwa, ok := resp.(*responseWithAction); ok { + ctx.Write(messageOut(rwa.action, rwa.data)) + } else { + ctx.Write(messageOut(msg.Action+".resp", resp)) + } + } +} + +// ============================================================ +// Bind — 泛型辅助,自动 JSON 反序列化 +// ============================================================ + +// Bind 将带类型的业务函数包装为 MessageHandler。 +// 自动处理 JSON 反序列化和序列化。 +// +// 用法: +// +// router.On("ping", wsc.Bind(func(ctx *wsc.Context, req *PingReq) (*PingResp, error) { +// return &PingResp{Pong: true}, nil +// })) +// +// responseWithAction 携带自定义响应 action 的返回包装,由 BindAs 使用。 +type responseWithAction struct { + action string + data any +} + +// errType error 接口的 reflect.Type,用于处理器签名校验。 +var errType = reflect.TypeOf((*error)(nil)).Elem() + +// isNilValue 判断 reflect.Value 是否为 nil。 +// 仅对可为 nil 的类型执行 IsNil,其余类型(值类型响应)一律视为非 nil, +// 避免 reflect.Value.IsNil 对值类型 panic。 +func isNilValue(v reflect.Value) bool { + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Ptr, reflect.Slice: + return v.IsNil() + default: + return false + } +} + +// bindImpl 自动 JSON 反序列化并调用业务函数。 +// respAction 为空时回包 action 为 msg.Action+".resp"(Bind 行为); +// 非空时回包 action 为指定值(BindAs 行为,用于兼容客户端约定的固定响应类型)。 +func bindImpl(respAction string, fn any) MessageHandler { + fnVal := reflect.ValueOf(fn) + fnType := fnVal.Type() + + // 验证签名:func(ctx *Context, req T) (R, error) + if fnType.Kind() != reflect.Func { + panic("wsc.Bind: argument must be a function") + } + if fnType.NumIn() != 2 { + panic("wsc.Bind: function must have 2 parameters (ctx, req)") + } + if fnType.NumOut() != 2 { + panic("wsc.Bind: function must return 2 values (resp, error)") + } + + reqType := fnType.In(1) // 请求参数类型 + + // 构造可寻址的零值实例(始终为指针 *T)。 + // 不能用 reflect.New(t).Elem().Interface():经 Interface() 拷贝后丢失可寻址性, + // 值类型请求参数时后续 .Addr() 会 panic。 + newReq := func() reflect.Value { + t := reqType + if t.Kind() == reflect.Ptr { + return reflect.New(t.Elem()) + } + return reflect.New(t) + } + + return func(ctx *Context, data json.RawMessage) (any, error) { + req := newReq() + if len(data) > 0 { + if err := json.Unmarshal(data, req.Interface()); err != nil { + return nil, NewError(400, "invalid request data: "+err.Error()) + } + } + + // reqType 为指针类型时传 *T,为值类型时解引用传 T。 + var reqArg reflect.Value + if reqType.Kind() == reflect.Ptr { + reqArg = req + } else { + reqArg = req.Elem() + } + + results := fnVal.Call([]reflect.Value{ + reflect.ValueOf(ctx), + reqArg, + }) + + var resp any + if !isNilValue(results[0]) { + resp = results[0].Interface() + } + + var err error + if !results[1].IsNil() { + err = results[1].Interface().(error) + } + + if err != nil { + return nil, err + } + if resp == nil { + return nil, nil + } + if respAction == "" { + return resp, nil + } + return &responseWithAction{action: respAction, data: resp}, nil + } +} + +// Bind 将带类型的业务函数包装为 MessageHandler,自动 JSON 反序列化。 +// 回包 action 为 msg.Action + ".resp"。 +func Bind(fn any) MessageHandler { + return bindImpl("", fn) +} + +// BindAs 与 Bind 相同,但允许指定响应 action(覆盖默认的 action+".resp")。 +// 适用于客户端约定了固定响应类型(如 "pong"、"register_resp")的场景。 +// +// 用法: +// +// router.On("ping", wsc.BindAs("pong", func(ctx *wsc.Context, req *PingReq) (*PingResp, error) { +// return &PingResp{Pong: true}, nil +// })) +func BindAs(respAction string, fn any) MessageHandler { + return bindImpl(respAction, fn) +} + +// ============================================================ +// BindNoResp — 无返回值的处理器 +// ============================================================ + +// BindNoResp 包装无返回值的业务函数。 +// 签名:func(ctx *Context, req T) error,唯一返回值须为 error。 +func BindNoResp(fn any) MessageHandler { + fnVal := reflect.ValueOf(fn) + fnType := fnVal.Type() + + // 验证签名:func(ctx *Context, req T) error + if fnType.Kind() != reflect.Func { + panic("wsc.BindNoResp: argument must be a function") + } + if fnType.NumIn() != 2 { + panic("wsc.BindNoResp: function must have 2 parameters (ctx, req)") + } + if fnType.NumOut() != 1 || fnType.Out(0) != errType { + panic("wsc.BindNoResp: function must return 1 value (error)") + } + + reqType := fnType.In(1) + + // 构造可寻址的零值实例,见 bindImpl 同名逻辑。 + newReq := func() reflect.Value { + t := reqType + if t.Kind() == reflect.Ptr { + return reflect.New(t.Elem()) + } + return reflect.New(t) + } + + return func(ctx *Context, data json.RawMessage) (any, error) { + req := newReq() + if len(data) > 0 { + if err := json.Unmarshal(data, req.Interface()); err != nil { + return nil, NewError(400, "invalid request data: "+err.Error()) + } + } + + // reqType 为指针类型时传 *T,为值类型时解引用传 T。 + var reqArg reflect.Value + if reqType.Kind() == reflect.Ptr { + reqArg = req + } else { + reqArg = req.Elem() + } + + results := fnVal.Call([]reflect.Value{ + reflect.ValueOf(ctx), + reqArg, + }) + + if !results[0].IsNil() { + return nil, results[0].Interface().(error) + } + return nil, nil + } +} + +// ============================================================ +// Error — 业务错误 +// ============================================================ + +// Error 业务错误,自动序列化发送给客户端。 +type Error struct { + Code int `json:"code"` + Message string `json:"message"` +} + +func (e *Error) Error() string { return e.Message } + +// NewError 创建业务错误。 +func NewError(code int, msg string) *Error { + return &Error{Code: code, Message: msg} +} + +// ============================================================ +// 内置中间件 +// ============================================================ + +// AuthMiddleware 认证中间件工厂。 +// tokenExtractor 从消息中提取 token;validator 验证 token 并返回 uid。 +// 验证通过后 uid 写入 ctx.UID。 +func AuthMiddleware( + tokenExtractor func(msg *Message) string, + validator func(token string) (uid int64, err error), +) MiddlewareFunc { + return func(ctx *Context, msg *Message) error { + token := tokenExtractor(msg) + if token == "" { + return NewError(401, "token required") + } + uid, err := validator(token) + if err != nil { + return NewError(401, "invalid token: "+err.Error()) + } + ctx.UID = uid + ctx.Set("uid", uid) + return nil + } +} + +// RecoveryMiddleware 恢复中间件,捕获 panic。 +func RecoveryMiddleware() MiddlewareFunc { + return func(ctx *Context, msg *Message) (err error) { + defer func() { + if r := recover(); r != nil { + logrus.Errorf("[wsc] panic recovered: %v, session=%s, type=%s", r, ctx.SessionID, msg.Action) + err = NewError(500, "internal server error") + } + }() + return nil + } +} + +// LoggerMiddleware 日志中间件。 +func LoggerMiddleware() MiddlewareFunc { + return func(ctx *Context, msg *Message) error { + logrus.Debugf("[wsc] %s | %s | %s | %s", ctx.ClientIP, ctx.SessionID, ctx.Path, msg.Action) + return nil + } +} diff --git a/wsc/server.go b/wsc/server.go new file mode 100644 index 0000000..041c472 --- /dev/null +++ b/wsc/server.go @@ -0,0 +1,123 @@ +package wsc + +import ( + "net/http" + "sync" + "sync/atomic" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +// ============================================================ +// 默认配置 +// ============================================================ + +var defaultUpgrader = websocket.Upgrader{ + ReadBufferSize: 4096, + WriteBufferSize: 4096, + CheckOrigin: func(r *http.Request) bool { + return true // 生产环境请限制 Origin + }, +} + +// ============================================================ +// Server — 连接管理器 +// ============================================================ + +// Server 管理 WebSocket 连接的全局状态。 +type Server struct { + upgrader websocket.Upgrader + + // 连接生命周期回调 + OnConnect func(ctx *Context) + OnDisconnect func(ctx *Context) + + // OnUpgrade 升级前钩子:可用于鉴权(如校验 URL token),返回 error 即拒绝连接(401)。 + // 钩子中通过 c.Set(k, v) 写入的值会自动透传到 *Context(见 newContext)。 + OnUpgrade func(c *gin.Context) error + + // 房间管理 + mu sync.RWMutex + rooms map[string]*Room + + // 连接统计 + connCount atomic.Int64 +} + +// NewServer 创建 Server。 +func NewServer() *Server { + return &Server{ + upgrader: defaultUpgrader, + rooms: make(map[string]*Room), + } +} + +// SetUpgrader 自定义升级器。 +func (s *Server) SetUpgrader(u websocket.Upgrader) { + s.upgrader = u +} + +// ============================================================ +// 连接统计 +// ============================================================ + +// ActiveCount 返回当前活跃连接数。 +func (s *Server) ActiveCount() int64 { + return s.connCount.Load() +} + +func (s *Server) onSessionOpen() { + s.connCount.Add(1) +} + +func (s *Server) onSessionClosed() { + s.connCount.Add(-1) +} + +// ============================================================ +// 房间管理 +// ============================================================ + +// GetOrCreateRoom 获取或创建房间。 +func (s *Server) GetOrCreateRoom(name string) *Room { + s.mu.Lock() + defer s.mu.Unlock() + + if room, ok := s.rooms[name]; ok { + return room + } + + room := newRoom(name) + room.onEmpty = func(name string) { + s.mu.Lock() + defer s.mu.Unlock() + // 二次判空:防止 leave→onEmpty 之间新成员加入 + if r, exists := s.rooms[name]; exists && r.Len() == 0 { + delete(s.rooms, name) + } + } + s.rooms[name] = room + return room +} + +// GetRoom 获取房间(不存在返回 nil)。 +func (s *Server) GetRoom(name string) *Room { + s.mu.RLock() + defer s.mu.RUnlock() + return s.rooms[name] +} + +// RemoveRoom 移除房间。 +func (s *Server) RemoveRoom(name string) { + s.mu.Lock() + defer s.mu.Unlock() + delete(s.rooms, name) +} + +// RoomCount 返回当前房间数。 +func (s *Server) RoomCount() int { + s.mu.RLock() + defer s.mu.RUnlock() + return len(s.rooms) +} diff --git a/wsc/session.go b/wsc/session.go new file mode 100644 index 0000000..3b4418e --- /dev/null +++ b/wsc/session.go @@ -0,0 +1,172 @@ +package wsc + +import ( + "encoding/json" + "errors" + "io" + "sync" + "time" + + "github.com/gorilla/websocket" + "github.com/sirupsen/logrus" +) + +// ============================================================ +// 消息帧格式 +// ============================================================ + +// Message 统一消息结构(客户端 <-> 服务端)。 +type Message struct { + Action string `json:"action"` + Payload json.RawMessage `json:"payload,omitempty"` +} + +// ============================================================ +// 错误定义 +// ============================================================ + +var ( + ErrConnectionClosed = errors.New("wsc: connection closed") + ErrSendBufferFull = errors.New("wsc: send buffer full") +) + +// ============================================================ +// Session — 单连接会话 +// ============================================================ + +const ( + writeWait = 10 * time.Second + pongWait = 60 * time.Second + pingPeriod = (pongWait * 9) / 10 + maxMessageSize = 65536 +) + +// Session 管理单个 WebSocket 连接。 +type Session struct { + wsCtx *Context + conn *websocket.Conn + server *Server + + send chan []byte + closeCh chan struct{} + once sync.Once +} + +// writeJSON 序列化并发送 JSON 到写通道(单次 marshal)。 +func (s *Session) writeJSON(v any) error { + data, err := json.Marshal(v) + if err != nil { + return err + } + return s.sendRaw(data) +} + +// sendRaw 将已序列化的字节写入发送通道。 +func (s *Session) sendRaw(data []byte) error { + select { + case s.send <- data: + return nil + case <-s.closeCh: + return ErrConnectionClosed + default: + logrus.Warnf("[wsc] send buffer full, dropping message for session=%s", s.wsCtx.SessionID[:8]) + return ErrSendBufferFull + } +} + +// writePump 写协程,从 send 通道取数据写入 WebSocket。 +func (s *Session) writePump() { + ticker := time.NewTicker(pingPeriod) + defer func() { + ticker.Stop() + s.close() + }() + + for { + select { + case message := <-s.send: + s.conn.SetWriteDeadline(time.Now().Add(writeWait)) + if err := s.conn.WriteMessage(websocket.TextMessage, message); err != nil { + return + } + + case <-ticker.C: + s.conn.SetWriteDeadline(time.Now().Add(writeWait)) + if err := s.conn.WriteMessage(websocket.PingMessage, nil); err != nil { + return + } + + case <-s.closeCh: + return + } + } +} + +// readPump 读协程,读取消息并路由。 +func (s *Session) readPump(router *Router) { + defer func() { + if r := recover(); r != nil { + logrus.Errorf("[wsc] readPump panic recovered: %v, session=%s", r, s.wsCtx.SessionID[:8]) + } + s.close() + }() + + s.conn.SetReadLimit(maxMessageSize) + s.conn.SetReadDeadline(time.Now().Add(pongWait)) + s.conn.SetPongHandler(func(string) error { + s.conn.SetReadDeadline(time.Now().Add(pongWait)) + return nil + }) + + for { + _, data, err := s.conn.ReadMessage() + if err != nil { + if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) && + !errors.Is(err, io.EOF) { + logrus.Errorf("[wsc] read error: %v", err) + } + return + } + + var msg Message + if err := json.Unmarshal(data, &msg); err != nil { + s.wsCtx.WriteError(400, "invalid message format") + continue + } + + router.dispatch(s.wsCtx, &msg) + } +} + +// close 安全关闭连接。 +func (s *Session) close() { + s.once.Do(func() { + close(s.closeCh) + s.wsCtx.cancelCtx() + + // 离开所有房间 + // 先快照房间列表再释放读锁,避免 leave 触发 onEmpty → RemoveRoom 抢写锁死锁 + s.server.mu.RLock() + rooms := make([]*Room, 0, len(s.server.rooms)) + for _, room := range s.server.rooms { + rooms = append(rooms, room) + } + s.server.mu.RUnlock() + + for _, room := range rooms { + room.Leave(s.wsCtx) + } + + // 触发断开回调 + if s.server.OnDisconnect != nil { + s.server.OnDisconnect(s.wsCtx) + } + + s.server.onSessionClosed() + + s.conn.Close() + // s.send 不关闭:readPump 的 dispatch 可能在 close 生效后仍并发回包, + // 向已关闭 channel 发送会 panic。writePump 退出由 closeCh 驱动, + // 残留消息随 Session 对象一起被 GC 回收。 + }) +} diff --git a/wsc/ws.go b/wsc/ws.go new file mode 100644 index 0000000..7e0fca5 --- /dev/null +++ b/wsc/ws.go @@ -0,0 +1,98 @@ +package wsc + +import ( + "net/http" + + "github.com/gin-gonic/gin" +) + +// HandleWS 一站式 WebSocket 注册。 +// +// 用法: +// +// wsc.HandleWS(r, "/ws/chat", svcCtx, func(router *wsc.Router) { +// router.On("ping", wsc.Bind(func(ctx *wsc.Context, req *PingReq) (*PingResp, error) { +// return &PingResp{Pong: true}, nil +// })) +// }) +func HandleWS(r *gin.RouterGroup, path string, server *Server, register func(router *Router)) { + // 预处理:创建路由表(只执行一次) + router := NewRouter() + if register != nil { + register(router) + } + + r.GET(path, func(c *gin.Context) { + // 升级前钩子:鉴权失败直接 401 拒绝连接 + if server.OnUpgrade != nil { + if err := server.OnUpgrade(c); err != nil { + c.AbortWithStatusJSON(401, gin.H{"error": err.Error()}) + return + } + } + + conn, err := server.upgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + return + } + + // 创建会话 + session := &Session{ + conn: conn, + server: server, + send: make(chan []byte, 256), + closeCh: make(chan struct{}), + } + + // 创建上下文(从 gin 提取元数据) + wsCtx := newContext(c, session) + session.wsCtx = wsCtx + wsCtx.session = session + + // 连接回调 + server.onSessionOpen() + if server.OnConnect != nil { + server.OnConnect(wsCtx) + } + + // 启动读写协程 + go session.writePump() + go session.readPump(router) + }) +} + +// HandleWSDefault 内部自动创建默认 Server 的注册,等价于 HandleWS(r, path, NewServer(), register)。 +func HandleWSDefault(r *gin.RouterGroup, path string, register func(router *Router)) { + HandleWS(r, path, NewServer(), register) +} + +// ============================================================ +// QuickServer — 快捷方式,省去手动创建 Server +// ============================================================ + +// QuickServer 创建带常用配置的 Server。 +func QuickServer(onConnect, onDisconnect func(ctx *Context)) *Server { + return &Server{ + upgrader: defaultUpgrader, + OnConnect: onConnect, + OnDisconnect: onDisconnect, + rooms: make(map[string]*Room), + } +} + +// ============================================================ +// CORS 辅助 +// ============================================================ + +// SetCheckOrigin 设置允许的 Origin。 +func (s *Server) SetCheckOrigin(origins ...string) { + s.upgrader.CheckOrigin = func(r *http.Request) bool { + origin := r.Header.Get("Origin") + for _, o := range origins { + if o == "*" || o == origin { + return true + } + } + return false + } +} diff --git a/wscclient/client.go b/wscclient/client.go new file mode 100644 index 0000000..7f758ef --- /dev/null +++ b/wscclient/client.go @@ -0,0 +1,587 @@ +// Package wscclient 提供 wsc 服务端的客户端封装。 +// 对称设计:cli.On == router.On,cli.Send == ctx.WriteMessage。 +package wscclient + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "reflect" + "sync" + "sync/atomic" + "time" + + "github.com/gorilla/websocket" + "github.com/sirupsen/logrus" +) + +// ============================================================ +// 配置 +// ============================================================ + +// Config 客户端配置。 +type Config struct { + URL string // WebSocket 地址(必填) + Header http.Header // 自定义请求头(如 token) + DialTimeout time.Duration // 连接超时,默认 10s + AutoReconnect bool // 自动重连 + MaxRetry int // 最大重试次数(0 = 无限),默认 0 + MinReconDelay time.Duration // 最小重连延迟,默认 1s + MaxReconDelay time.Duration // 最大重连延迟,默认 30s + PingInterval time.Duration // 心跳间隔(0 = 不发送),默认 25s + PongTimeout time.Duration // 心跳超时,默认 60s + WriteTimeout time.Duration // 写超时,默认 10s + SendBufferSize int // 发送缓冲区大小,默认 256 +} + +// DefaultConfig 返回推荐默认配置。 +func DefaultConfig(url string) Config { + return Config{ + URL: url, + DialTimeout: 10 * time.Second, + AutoReconnect: true, + MaxRetry: 0, + MinReconDelay: 1 * time.Second, + MaxReconDelay: 30 * time.Second, + PingInterval: 25 * time.Second, + PongTimeout: 60 * time.Second, + WriteTimeout: 10 * time.Second, + SendBufferSize: 256, + } +} + +// ============================================================ +// Handler — 客户端消息处理器 +// ============================================================ + +// Handler 客户端消息处理函数。payload 为服务端发来的 JSON,已去掉 action 包装。 +type Handler func(payload json.RawMessage) + +// ============================================================ +// Client +// ============================================================ + +// Client WebSocket 客户端,并发安全。 +// 所有公开方法可在任意 goroutine 调用。 +type Client struct { + url string + urlMu sync.RWMutex + cfg Config + + // 消息路由 + mu sync.RWMutex + routes map[string]Handler + + // 连接 + conn *websocket.Conn + connMu sync.Mutex + + writeCh chan []byte + closeCh chan struct{} + doneCh chan struct{} + + // 重连 + reconnecting atomic.Bool + reconMu sync.Mutex // 保护 reconDelay / reconCount / reconStableTimer + reconDelay time.Duration + reconCount int + reconStop chan struct{} + reconDone chan struct{} + reconStableTimer *time.Timer // 连接稳定后重置退避的定时器 + + // 回调 + OnConnected func() + OnDisconnected func(err error) + OnReconnecting func(attempt int, delay time.Duration) + + // 生命周期控制 + closed atomic.Bool + closeOnce sync.Once + wg sync.WaitGroup +} + +// New 创建客户端(尚未连接)。 +func New(cfg Config) *Client { + if cfg.DialTimeout == 0 { + cfg.DialTimeout = 10 * time.Second + } + if cfg.MinReconDelay == 0 { + cfg.MinReconDelay = time.Second + } + if cfg.MaxReconDelay == 0 { + cfg.MaxReconDelay = 30 * time.Second + } + if cfg.PingInterval == 0 { + cfg.PingInterval = 25 * time.Second + } + if cfg.PongTimeout == 0 { + cfg.PongTimeout = 60 * time.Second + } + if cfg.WriteTimeout == 0 { + cfg.WriteTimeout = 10 * time.Second + } + if cfg.SendBufferSize == 0 { + cfg.SendBufferSize = 256 + } + + return &Client{ + url: cfg.URL, + cfg: cfg, + routes: make(map[string]Handler), + writeCh: make(chan []byte, cfg.SendBufferSize), + closeCh: make(chan struct{}), + doneCh: make(chan struct{}), + } +} + +// ============================================================ +// 连接控制 +// ============================================================ + +// Connect 连接 WebSocket 服务端。AutoReconnect=true 时阻塞直到连接成功或达到最大重试。 +func (c *Client) Connect() error { + return c.connect() +} + +// Close 关闭连接,停止重连。 +func (c *Client) Close() { + c.closeOnce.Do(func() { + c.closed.Store(true) + close(c.closeCh) + + // 停止重连 + select { + case c.reconStop <- struct{}{}: + default: + } + + c.cancelBackoffReset() + + c.connMu.Lock() + if c.conn != nil { + c.conn.Close() + } + c.connMu.Unlock() + + c.wg.Wait() + close(c.doneCh) + }) +} + +// Connected 返回当前是否已连接。 +func (c *Client) Connected() bool { + c.connMu.Lock() + defer c.connMu.Unlock() + return c.conn != nil +} + +// Done 返回一个通道,Client 完全关闭后关闭。 +func (c *Client) Done() <-chan struct{} { + return c.doneCh +} + +// SetURL 动态更新连接地址(含 token)。并发安全,下一次 dial/重连时生效。 +// 注意:已建立的连接不会立即断开,仍使用旧地址直到下一次重连。 +func (c *Client) SetURL(url string) { + c.urlMu.Lock() + c.url = url + c.urlMu.Unlock() +} + +// getURL 并发安全地读取当前连接地址。 +func (c *Client) getURL() string { + c.urlMu.RLock() + defer c.urlMu.RUnlock() + return c.url +} + +// ============================================================ +// 消息路由(注册后再 Connect) +// ============================================================ + +// On 注册 action 对应的消息处理器。 +func (c *Client) On(action string, handler Handler) { + c.mu.Lock() + c.routes[action] = handler + c.mu.Unlock() +} + +// Send 发送结构化消息。并发安全。 +func (c *Client) Send(action string, payload any) error { + m := map[string]any{"action": action} + if payload != nil { + m["payload"] = payload + } + data, err := json.Marshal(m) + if err != nil { + return err + } + return c.SendRaw(data) +} + +// SendRaw 发送已序列化的字节。并发安全。 +func (c *Client) SendRaw(data []byte) error { + if c.closed.Load() { + return errors.New("wscclient: client closed") + } + select { + case c.writeCh <- data: + return nil + case <-c.closeCh: + return errors.New("wscclient: client closed") + default: + logrus.Warnf("[wscclient] send buffer full, dropping message") + return errors.New("wscclient: send buffer full") + } +} + +// ============================================================ +// Bind — 自动 JSON 反序列化 +// ============================================================ + +// Bind 将带类型的函数包装为 Handler,自动 JSON 反序列化 payload。 +// +// cli.On("ping.resp", wscclient.Bind(func(resp *PingResp) { +// log.Println(resp.Message) +// })) +func Bind(fn any) Handler { + fnVal := reflect.ValueOf(fn) + fnType := fnVal.Type() + + if fnType.Kind() != reflect.Func || fnType.NumIn() != 1 { + panic("wscclient.Bind: function must have 1 parameter") + } + + reqType := fnType.In(0) + if reqType.Kind() != reflect.Ptr { + panic("wscclient.Bind: parameter must be a pointer") + } + + return func(raw json.RawMessage) { + req := reflect.New(reqType.Elem()).Interface() + if len(raw) > 0 { + if err := json.Unmarshal(raw, req); err != nil { + logrus.Errorf("[wscclient] Bind unmarshal error: %v, raw=%s", err, string(raw)) + return + } + } + fnVal.Call([]reflect.Value{reflect.ValueOf(req)}) + } +} + +// ============================================================ +// 内部:连接与重连 +// ============================================================ + +func (c *Client) connect() error { + for { + err := c.dial() + if err == nil { + return nil + } + + if !c.cfg.AutoReconnect { + return err + } + + if c.closed.Load() { + return errors.New("wscclient: closed") + } + + c.reconCount++ + if c.cfg.MaxRetry > 0 && c.reconCount > c.cfg.MaxRetry { + return err + } + + delay := c.nextDelay() + logrus.Warnf("[wscclient] connect failed (attempt %d): %v, retry in %v", c.reconCount, err, delay) + + if c.OnReconnecting != nil { + c.OnReconnecting(c.reconCount, delay) + } + + select { + case <-time.After(delay): + case <-c.closeCh: + return errors.New("wscclient: closed during reconnect") + } + } +} + +func (c *Client) dial() error { + dialer := websocket.Dialer{ + HandshakeTimeout: c.cfg.DialTimeout, + } + if c.cfg.Header != nil { + dialer.Proxy = http.ProxyFromEnvironment // 无操作,只是占位 + } + + ctx, cancel := context.WithTimeout(context.Background(), c.cfg.DialTimeout) + defer cancel() + + conn, _, err := dialer.DialContext(ctx, c.getURL(), c.cfg.Header) + if err != nil { + return err + } + + c.connMu.Lock() + c.conn = conn + c.connMu.Unlock() + + c.scheduleBackoffReset() + + // 重启内部通道(每次重连重新创建) + if c.reconStop != nil { + select { + case c.reconStop <- struct{}{}: + default: + } + } + c.reconStop = make(chan struct{}, 1) + c.reconDone = make(chan struct{}) + + // 启动读写协程 + c.wg.Add(2) + go c.readPump() + go c.writePump() + + if c.OnConnected != nil { + c.OnConnected() + } + + return nil +} + +// ============================================================ +// 内部:读写协程 +// ============================================================ + +func (c *Client) readPump() { + defer c.wg.Done() + defer func() { + if r := recover(); r != nil { + logrus.Errorf("[wscclient] readPump panic: %v", r) + } + c.onDisconnect() + }() + + c.connMu.Lock() + conn := c.conn + c.connMu.Unlock() + if conn == nil { + return + } + conn.SetReadLimit(65536) + conn.SetReadDeadline(time.Now().Add(c.cfg.PongTimeout)) + conn.SetPongHandler(func(string) error { + conn.SetReadDeadline(time.Now().Add(c.cfg.PongTimeout)) + return nil + }) + + for { + _, data, err := conn.ReadMessage() + if err != nil { + if !c.closed.Load() { + if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) && + !errors.Is(err, context.DeadlineExceeded) { + logrus.Errorf("[wscclient] read error: %v", err) + } + } + return + } + + var msg struct { + Action string `json:"action"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &msg); err != nil { + continue + } + + c.mu.RLock() + handler, ok := c.routes[msg.Action] + c.mu.RUnlock() + + if ok { + handler(msg.Payload) + } else { + logrus.Warnf("[wscclient] unhandled message action=%q, payload=%s", msg.Action, string(msg.Payload)) + } + } +} + +func (c *Client) writePump() { + defer c.wg.Done() + + c.connMu.Lock() + conn := c.conn + c.connMu.Unlock() + if conn == nil { + return + } + defer conn.Close() + + pingTicker := time.NewTicker(c.cfg.PingInterval) + defer pingTicker.Stop() + + for { + select { + case data, ok := <-c.writeCh: + if !ok { + return + } + conn.SetWriteDeadline(time.Now().Add(c.cfg.WriteTimeout)) + if err := conn.WriteMessage(websocket.TextMessage, data); err != nil { + logrus.Errorf("[wscclient] write error: %v", err) + return + } + + case <-pingTicker.C: + conn.SetWriteDeadline(time.Now().Add(c.cfg.WriteTimeout)) + if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { + return + } + + case <-c.reconStop: + return + + case <-c.closeCh: + return + } + } +} + +// ============================================================ +// 内部:断连处理 +// ============================================================ + +func (c *Client) onDisconnect() { + c.connMu.Lock() + c.conn = nil + c.connMu.Unlock() + + if c.OnDisconnected != nil { + c.OnDisconnected(errors.New("connection lost")) + } + + if c.closed.Load() || !c.cfg.AutoReconnect { + close(c.writeCh) + return + } + + // 连接已断开,取消“稳定后重置退避”的定时器,使退避继续累积 + c.cancelBackoffReset() + + // 启动重连 goroutine + if c.reconnecting.CompareAndSwap(false, true) { + go c.reconnectLoop() + } +} + +func (c *Client) reconnectLoop() { + defer c.reconnecting.Store(false) + + for { + if c.closed.Load() { + return + } + + conn, _, err := (&websocket.Dialer{ + HandshakeTimeout: c.cfg.DialTimeout, + }).DialContext(context.Background(), c.getURL(), c.cfg.Header) + + if err == nil { + c.scheduleBackoffReset() + + c.connMu.Lock() + c.conn = conn + c.connMu.Unlock() + + // 新 writeCh(旧的可能还有残留,丢弃) + c.writeCh = make(chan []byte, c.cfg.SendBufferSize) + + c.reconStop = make(chan struct{}, 1) + + c.wg.Add(2) + go c.readPump() + go c.writePump() + + if c.OnConnected != nil { + c.OnConnected() + } + return + } + + c.reconCount++ + if c.cfg.MaxRetry > 0 && c.reconCount > c.cfg.MaxRetry { + logrus.Errorf("[wscclient] reconnect max retry exceeded (%d)", c.cfg.MaxRetry) + c.Close() + return + } + + delay := c.nextDelay() + logrus.Warnf("[wscclient] reconnect attempt %d failed: %v, retry in %v", c.reconCount, err, delay) + + if c.OnReconnecting != nil { + c.OnReconnecting(c.reconCount, delay) + } + + select { + case <-time.After(delay): + case <-c.reconStop: + return + case <-c.closeCh: + return + } + } +} + +// ============================================================ +// 内部:辅助 +// ============================================================ + +// defaultReconStableGrace 连接持续稳定达到该时长后,重连退避才重置为初始值。 +// 避免在连接频繁抖动(握手成功但随即断开)时退避被反复清零,导致“永远 1 秒重连一次”。 +const defaultReconStableGrace = 10 * time.Second + +// scheduleBackoffReset 在连接稳定持续 grace 后,将退避延迟与重试计数重置为初始值。 +// 若在此期间连接再次断开(cancelBackoffReset),则退避继续累积,不会被清零。 +func (c *Client) scheduleBackoffReset() { + c.reconMu.Lock() + defer c.reconMu.Unlock() + if c.reconStableTimer != nil { + c.reconStableTimer.Stop() + } + c.reconStableTimer = time.AfterFunc(defaultReconStableGrace, func() { + c.reconMu.Lock() + c.reconDelay = 0 + c.reconCount = 0 + c.reconStableTimer = nil + c.reconMu.Unlock() + }) +} + +// cancelBackoffReset 取消待执行的退避重置(连接再次断开时调用)。 +func (c *Client) cancelBackoffReset() { + c.reconMu.Lock() + defer c.reconMu.Unlock() + if c.reconStableTimer != nil { + c.reconStableTimer.Stop() + c.reconStableTimer = nil + } +} + +func (c *Client) nextDelay() time.Duration { + c.reconMu.Lock() + defer c.reconMu.Unlock() + if c.reconDelay == 0 { + c.reconDelay = c.cfg.MinReconDelay + } + delay := c.reconDelay + c.reconDelay *= 2 + if c.reconDelay > c.cfg.MaxReconDelay { + c.reconDelay = c.cfg.MaxReconDelay + } + return delay +}