init: 通用库抽离——httpx/logger/wsc/wscclient/jwtx(合并分叉副本); 顶层包布局; module path git.zeroonesoft.cn/golib/zogo

This commit is contained in:
w11
2026-09-19 19:04:01 +08:00
commit f0f5421262
17 changed files with 2336 additions and 0 deletions
+2
View File
@@ -0,0 +1,2 @@
.idea/
*.exe
+28
View File
@@ -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)。
+47
View File
@@ -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
)
+107
View File
@@ -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=
+18
View File
@@ -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}
}
+28
View File
@@ -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)
}
+35
View File
@@ -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,
})
}
+75
View File
@@ -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)
}
+61
View File
@@ -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))
}
+308
View File
@@ -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()
}
+187
View File
@@ -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
}
+103
View File
@@ -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)
}
+357
View File
@@ -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
}
}
+123
View File
@@ -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)
}
+172
View File
@@ -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 回收。
})
}
+98
View File
@@ -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
}
}
+587
View File
@@ -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
}