init: 通用库抽离——httpx/logger/wsc/wscclient/jwtx(合并分叉副本); 顶层包布局; module path git.zeroonesoft.cn/golib/zogo
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
.idea/
|
||||
*.exe
|
||||
@@ -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)。
|
||||
@@ -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
|
||||
)
|
||||
@@ -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=
|
||||
@@ -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}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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 回收。
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user