第一次提交
This commit is contained in:
@@ -0,0 +1,189 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/gogf/gf/v2/database/gdb"
|
||||
"github.com/gogf/gf/v2/encoding/gyaml"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/os/gcfg"
|
||||
"github.com/gogf/gf/v2/os/gctx"
|
||||
"github.com/gogf/gf/v2/os/gfile"
|
||||
"github.com/gogf/gf/v2/os/glog"
|
||||
"github.com/gogf/gf/v2/util/gconv"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Server ServiceConfig `yaml:"server"`
|
||||
Database *DatabaseConfig `yaml:"database"`
|
||||
SkipUrl string `yaml:"skipUrl"`
|
||||
OpenAPITitle string `yaml:"openAPITitle"`
|
||||
OpenAPIDescription string `yaml:"openAPIDescription"`
|
||||
OpenAPIUrl string `yaml:"openAPIUrl"`
|
||||
OpenAPIName string `yaml:"openAPIName"`
|
||||
DoMain []string `yaml:"doMain"`
|
||||
OpenAPIVersion string `yaml:"openAPIVersion"`
|
||||
Logger LoggerConfig `yaml:"logger"`
|
||||
Dns string `yaml:"dns"`
|
||||
}
|
||||
|
||||
type ServiceConfig struct {
|
||||
Default ServiceDefault `yaml:"default"`
|
||||
}
|
||||
|
||||
type ServiceDefault struct {
|
||||
Address string `yaml:"address"`
|
||||
LogPath string `yaml:"logPath"`
|
||||
LogStdout bool `yaml:"logStdout"`
|
||||
ErrorStack bool `yaml:"errorStack"`
|
||||
ErrorLogEnabled bool `yaml:"errorLogEnabled"`
|
||||
ErrorLogPattern string `yaml:"errorLogPattern"`
|
||||
AccessLogEnable bool `yaml:"accessLogEnable"`
|
||||
AccessLogPattern string `yaml:"accessLogPattern"`
|
||||
FileServerEnabled bool `yaml:"fileServerEnabled"`
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Default DatabaseDefault `yaml:"default"`
|
||||
}
|
||||
|
||||
type DatabaseDefault struct {
|
||||
Host string `yaml:"host" json:"host"`
|
||||
Link string `yaml:"link" dc:"数据库连接字符串" json:"link"`
|
||||
Port string `yaml:"port" json:"port"`
|
||||
User string `yaml:"user" json:"user"`
|
||||
Pass string `yaml:"pass" json:"pass"`
|
||||
Name string `yaml:"name" json:"name"`
|
||||
Type string `yaml:"type" json:"type"`
|
||||
Timezone string `yaml:"timezone" json:"timezone"`
|
||||
Debug bool `yaml:"debug" json:"debug"`
|
||||
Charset string `yaml:"charset" json:"charset"`
|
||||
CreatedAt string `yaml:"createdAt" json:"createdAt"`
|
||||
UpdatedAt string `yaml:"updatedAt" json:"updatedAt"`
|
||||
}
|
||||
|
||||
type LoggerConfig struct {
|
||||
Path string `yaml:"path" json:"path"`
|
||||
File string `yaml:"file" json:"file"`
|
||||
Level string `yaml:"level" json:"level"`
|
||||
TimeFormat string `yaml:"timeFormat" json:"timeFormat"`
|
||||
CtxKeys []string `yaml:"ctxKeys" json:"ctxKeys"`
|
||||
Header bool `yaml:"header" json:"header"`
|
||||
Stdout bool `yaml:"stdout" json:"stdout"`
|
||||
RotateSize string `yaml:"rotateSize" json:"rotateSize"`
|
||||
RotateBackupLimit int `yaml:"rotateBackupLimit" json:"rotateBackupLimit"`
|
||||
StdoutColorDisabled bool `yaml:"stdoutColorDisabled" json:"stdoutColorDisabled"`
|
||||
WriterColorEnable bool `yaml:"writerColorEnable" json:"writerColorEnable"`
|
||||
}
|
||||
|
||||
var DefaultConfig = Config{
|
||||
Server: ServiceConfig{
|
||||
Default: ServiceDefault{
|
||||
Address: ":8080",
|
||||
LogPath: "./log/",
|
||||
LogStdout: true,
|
||||
ErrorStack: true,
|
||||
ErrorLogEnabled: true,
|
||||
ErrorLogPattern: "error-{Ymd}.log",
|
||||
AccessLogEnable: false,
|
||||
FileServerEnabled: true,
|
||||
},
|
||||
},
|
||||
OpenAPITitle: "",
|
||||
OpenAPIDescription: "Api列表 包含各端接口信息 字段注释 枚举说明",
|
||||
OpenAPIUrl: "https://panel.magicany.cc:8888/btpanel",
|
||||
OpenAPIName: "",
|
||||
DoMain: []string{"localhost", "127.0.0.1"},
|
||||
OpenAPIVersion: "v1.0",
|
||||
Dns: "root:123456@tcp(127.0.0.1:3306)/test",
|
||||
Logger: LoggerConfig{
|
||||
Path: "./log/",
|
||||
File: "access-{Ymd}.log",
|
||||
Level: "all",
|
||||
TimeFormat: "2006-01-02 15:04:05",
|
||||
CtxKeys: []string{},
|
||||
Header: true,
|
||||
Stdout: true,
|
||||
RotateSize: "1M",
|
||||
RotateBackupLimit: 10,
|
||||
},
|
||||
}
|
||||
|
||||
func DefaultConfigInit() {
|
||||
database := &DatabaseConfig{Default: DatabaseDefault{
|
||||
Host: "127.0.0.1",
|
||||
Port: "3306",
|
||||
User: "root",
|
||||
Pass: "123456",
|
||||
Name: "database",
|
||||
Type: "mysql",
|
||||
Timezone: "Local",
|
||||
Debug: true,
|
||||
Charset: "utf8",
|
||||
CreatedAt: "create_time",
|
||||
UpdatedAt: "update_time",
|
||||
}}
|
||||
DefaultConfig.Database = database
|
||||
yaml, err := gyaml.Encode(DefaultConfig)
|
||||
if err != nil {
|
||||
g.Log().Error(gctx.New(), "转换yaml失败", err)
|
||||
}
|
||||
|
||||
if !gfile.IsDir(uploadPath) {
|
||||
_ = gfile.Mkdir(uploadPath)
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "template"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "scripts"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "html"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "resource", "css"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "resource", "image"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "resource", "js"))
|
||||
}
|
||||
g.Log().Info(gctx.New(), "正在检查配置文件", gfile.IsFile(ConfigPath))
|
||||
if !gfile.IsFile(ConfigPath) {
|
||||
g.Log().Info(gctx.New(), "正在创建配置文件", ConfigPath)
|
||||
_, _ = gfile.Create(ConfigPath)
|
||||
g.Log().Info(gctx.New(), "正在写入配置文件", ConfigPath)
|
||||
_ = gfile.PutContents(ConfigPath, gconv.String(yaml))
|
||||
g.Log().Info(gctx.New(), "配置文件创建成功!")
|
||||
} else {
|
||||
gcfg.Instance().GetAdapter().(*gcfg.AdapterFile).SetFileName(ConfigPath)
|
||||
}
|
||||
}
|
||||
|
||||
// DefaultSqliteConfigInit 创建默认的sqlite数据库配置 不会再生成配置文件
|
||||
// @param path sqlite数据库路径
|
||||
// @param autoTime 自动时间字段[]string{"create_time","update_time"}
|
||||
// @param debug 数据库调试模式
|
||||
// @param prefix 表前缀可空
|
||||
func DefaultSqliteConfigInit(path string, autoTime []string, debug bool, prefix ...string) {
|
||||
glog.Info(gctx.New(), "正在检查文件夹", gfile.IsFile(uploadPath))
|
||||
if !gfile.IsDir(uploadPath) {
|
||||
_ = gfile.Mkdir(uploadPath)
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "template"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "scripts"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "html"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "resource", "css"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "resource", "image"))
|
||||
_ = gfile.Mkdir(gfile.Join(gfile.Pwd(), "resource", "public", "resource", "js"))
|
||||
}
|
||||
g.Log().Info(gctx.New(), "正在设置数据库配置")
|
||||
node := gdb.ConfigNode{
|
||||
Link: fmt.Sprintf("sqlite::@file(%s)", path),
|
||||
Timezone: "Local",
|
||||
Charset: "utf8",
|
||||
CreatedAt: autoTime[0],
|
||||
UpdatedAt: autoTime[1],
|
||||
Debug: debug,
|
||||
}
|
||||
if len(prefix) > 0 {
|
||||
node.Prefix = prefix[0]
|
||||
}
|
||||
err := gdb.SetConfig(gdb.Config{
|
||||
"default": gdb.ConfigGroup{
|
||||
node,
|
||||
}})
|
||||
if err != nil {
|
||||
g.Log().Error(gctx.New(), "设置数据库配置失败", err)
|
||||
}
|
||||
g.Log().Info(gctx.New(), "设置数据库配置成功")
|
||||
}
|
||||
+401
@@ -0,0 +1,401 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/os/glog"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
"github.com/gogf/gf/v2/net/goai"
|
||||
"github.com/gogf/gf/v2/os/gcfg"
|
||||
"github.com/gogf/gf/v2/os/gctx"
|
||||
"github.com/gogf/gf/v2/os/gfile"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
"github.com/gogf/gf/v2/text/gstr"
|
||||
"github.com/gogf/gf/v2/util/gconv"
|
||||
)
|
||||
|
||||
type Json struct {
|
||||
Code int `json:"code" d:"1"`
|
||||
Data any `json:"data"`
|
||||
Msg string `json:"msg" d:"操作成功"`
|
||||
}
|
||||
|
||||
type ApiRes struct {
|
||||
ctx context.Context
|
||||
json *Json
|
||||
}
|
||||
|
||||
func Success(ctx context.Context) *ApiRes {
|
||||
json := Json{
|
||||
Code: 1,
|
||||
}
|
||||
|
||||
var a = ApiRes{
|
||||
ctx: ctx,
|
||||
json: &json,
|
||||
}
|
||||
return &a
|
||||
}
|
||||
|
||||
func Error(ctx context.Context) *ApiRes {
|
||||
json := Json{
|
||||
Code: 0,
|
||||
}
|
||||
|
||||
var a = ApiRes{
|
||||
ctx: ctx,
|
||||
json: &json,
|
||||
}
|
||||
return &a
|
||||
}
|
||||
|
||||
func (a *ApiRes) SetCode(code int) *ApiRes {
|
||||
a.json.Code = code
|
||||
return a
|
||||
}
|
||||
|
||||
func (a *ApiRes) SetData(data interface{}) *ApiRes {
|
||||
a.json.Data = data
|
||||
return a
|
||||
}
|
||||
|
||||
func (a *ApiRes) SetMsg(msg string) *ApiRes {
|
||||
a.json.Msg = msg
|
||||
return a
|
||||
}
|
||||
|
||||
func (a *ApiRes) End() {
|
||||
from := g.RequestFromCtx(a.ctx)
|
||||
from.Header.Set("Access-Control-Expose-Headers", "Set-Cookie")
|
||||
from.Response.Status = 200
|
||||
from.Response.WriteJson(a.json)
|
||||
return
|
||||
}
|
||||
|
||||
func (a *ApiRes) FileDownload(path, name string) {
|
||||
from := g.RequestFromCtx(a.ctx)
|
||||
from.Response.ServeFileDownload(path)
|
||||
return
|
||||
}
|
||||
|
||||
func (a *ApiRes) FileSelect(path string) {
|
||||
from := g.RequestFromCtx(a.ctx)
|
||||
from.Response.ServeFile(path)
|
||||
return
|
||||
}
|
||||
|
||||
// LoginJson 返回登录json数据
|
||||
/*
|
||||
* @param ctx 上下文
|
||||
* @param msg 返回信息
|
||||
* @param data 返回数据
|
||||
*/
|
||||
func LoginJson(r *ghttp.Request, msg string, data ...interface{}) {
|
||||
var info interface{}
|
||||
if len(data) > 0 {
|
||||
info = data[0]
|
||||
} else {
|
||||
info = nil
|
||||
}
|
||||
r.Response.WriteJsonExit(Json{
|
||||
Code: 1,
|
||||
Data: info,
|
||||
Msg: msg,
|
||||
})
|
||||
}
|
||||
|
||||
// ResponseJson 返回json数据
|
||||
/*
|
||||
* @param ctx 上下文
|
||||
* @param data 返回数据
|
||||
*/
|
||||
func ResponseJson(ctx context.Context, data interface{}) {
|
||||
g.RequestFromCtx(ctx).Response.WriteJson(data)
|
||||
return
|
||||
}
|
||||
|
||||
type PageSize struct {
|
||||
CurrentPage int `json:"currentPage"`
|
||||
Data interface{} `json:"data"`
|
||||
LastPage int `json:"lastPage"`
|
||||
PerPage int `json:"per_page"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
|
||||
// SetPage 设置分页
|
||||
/*
|
||||
* @param page 当前页
|
||||
* @param limit 每页显示条数
|
||||
* @param total 总条数
|
||||
* @param data 返回数据
|
||||
* @return PageSize
|
||||
*/
|
||||
func SetPage(page, limit, total int, data interface{}) *PageSize {
|
||||
var size = new(PageSize)
|
||||
if page == 1 {
|
||||
size.LastPage = 1
|
||||
} else {
|
||||
size.LastPage = page - 1
|
||||
}
|
||||
size.PerPage = limit
|
||||
size.Total = total
|
||||
size.CurrentPage = page
|
||||
size.Data = data
|
||||
return size
|
||||
}
|
||||
|
||||
// MiddlewareError 异常处理中间件
|
||||
func MiddlewareError(r *ghttp.Request) {
|
||||
r.Middleware.Next()
|
||||
var (
|
||||
err = r.GetError()
|
||||
msg string
|
||||
res = r.GetHandlerResponse()
|
||||
status = r.Response.Status
|
||||
)
|
||||
json := new(Json)
|
||||
json.Code = 1
|
||||
json.Data = res
|
||||
json.Msg = "操作成功"
|
||||
if err != nil {
|
||||
bo := gstr.Contains(err.Error(), ": ")
|
||||
if bo {
|
||||
msg = gstr.SubStrFromEx(err.Error(), ": ")
|
||||
} else {
|
||||
msg = err.Error()
|
||||
}
|
||||
r.Response.ClearBuffer()
|
||||
json.Code = 0
|
||||
json.Msg = msg
|
||||
r.Response.Status = http.StatusInternalServerError
|
||||
}
|
||||
if r.Response.BufferLength() > 0 {
|
||||
return
|
||||
}
|
||||
if status == 401 {
|
||||
json.Code = 0
|
||||
json.Msg = "请登录后操作"
|
||||
}
|
||||
r.Response.WriteJson(json)
|
||||
}
|
||||
|
||||
// AuthBase 鉴权中间件,只有前端或者后端登录成功之后才能通过
|
||||
func AuthBase(r *ghttp.Request, name string) {
|
||||
info, err := r.Session.Get(name, nil)
|
||||
if err != nil {
|
||||
panic(err.Error())
|
||||
}
|
||||
if !info.IsEmpty() {
|
||||
r.Middleware.Next()
|
||||
} else {
|
||||
NoLogin(r)
|
||||
}
|
||||
}
|
||||
|
||||
// AuthAdmin 鉴权中间件,只有后端登录成功之后才能通过
|
||||
func AuthAdmin(r *ghttp.Request) {
|
||||
AuthBase(r, "admin")
|
||||
}
|
||||
|
||||
// AuthIndex 鉴权中间件,只有前端登录成功之后才能通过
|
||||
func AuthIndex(r *ghttp.Request) {
|
||||
AuthBase(r, "user")
|
||||
}
|
||||
|
||||
// NoLogin 未登录返回
|
||||
func NoLogin(r *ghttp.Request) {
|
||||
r.Response.Status = 401
|
||||
r.Response.WriteJsonExit(Json{
|
||||
Code: 401,
|
||||
Data: nil,
|
||||
Msg: "请登录后操作",
|
||||
})
|
||||
}
|
||||
|
||||
// CreateFileDir 创建文件目录
|
||||
func CreateFileDir() error {
|
||||
path := gfile.Pwd() + "/resource"
|
||||
if !gfile.IsDir(path) {
|
||||
if err := gfile.Mkdir(path); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := gfile.Mkdir(path + "/public/upload"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func AuthLoginSession(ctx context.Context, sessionKey string) {
|
||||
ti, err := g.RequestFromCtx(ctx).Session.Get(sessionKey+"LoginTime", "")
|
||||
if err != nil {
|
||||
panic(err.Error())
|
||||
}
|
||||
if !ti.IsEmpty() {
|
||||
now := gtime.Now().Timestamp()
|
||||
if now-gconv.Int64(ti) <= 300 {
|
||||
number, err := g.RequestFromCtx(ctx).Session.Get(sessionKey+"LoginNum", 0)
|
||||
if err != nil {
|
||||
panic(err.Error())
|
||||
}
|
||||
if !number.IsEmpty() {
|
||||
count := gconv.Int(number)
|
||||
if count == 3 {
|
||||
panic("请等待5分钟后再次尝试或修改后尝试登录")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func LoginCountSession(ctx context.Context, sessionKey string) {
|
||||
ti, err := g.RequestFromCtx(ctx).Session.Get(sessionKey+"LoginTime", "")
|
||||
if err != nil {
|
||||
panic(err.Error())
|
||||
}
|
||||
if ti.IsEmpty() {
|
||||
_ = g.RequestFromCtx(ctx).Session.Set(sessionKey+"LoginTime", gtime.Now().Timestamp())
|
||||
}
|
||||
now := gtime.Now().Timestamp()
|
||||
if now-gconv.Int64(ti) <= 300 {
|
||||
number, err := g.RequestFromCtx(ctx).Session.Get(sessionKey+"LoginNum", 0)
|
||||
if err != nil {
|
||||
panic(err.Error())
|
||||
}
|
||||
if number.IsEmpty() {
|
||||
_ = g.RequestFromCtx(ctx).Session.Set(sessionKey+"LoginNum", 1)
|
||||
} else {
|
||||
count := gconv.Int(number)
|
||||
if count == 3 {
|
||||
panic("尝试登录已超过限制,请等待5分钟后再次尝试或修改后尝试登录")
|
||||
}
|
||||
_ = g.RequestFromCtx(ctx).Session.Set(sessionKey+"LoginNum", count+1)
|
||||
}
|
||||
} else {
|
||||
_ = g.RequestFromCtx(ctx).Session.Set(sessionKey+"LoginTime", gtime.Now().Timestamp())
|
||||
_ = g.RequestFromCtx(ctx).Session.Set(sessionKey+"LoginNum", 1)
|
||||
}
|
||||
}
|
||||
|
||||
func enhanceOpenAPIDoc(s *ghttp.Server) {
|
||||
openapi := s.GetOpenApi()
|
||||
openapi.Config.CommonResponse = ghttp.DefaultHandlerResponse{}
|
||||
openapi.Config.CommonResponseDataField = `Data`
|
||||
|
||||
// API description.
|
||||
openapi.Info = goai.Info{
|
||||
Title: "Api列表",
|
||||
Description: "Api列表 包含各端接口信息 字段注释 枚举说明",
|
||||
Contact: &goai.Contact{
|
||||
Name: "Api列表",
|
||||
URL: "https://panel.magicany.cc:8888/btpanel",
|
||||
},
|
||||
License: &goai.License{
|
||||
Name: "马国栋",
|
||||
URL: "https://panel.magicany.cc:8888/btpanel",
|
||||
},
|
||||
Version: "Api列表",
|
||||
}
|
||||
}
|
||||
|
||||
var ConfigPath = filepath.Join(gfile.Pwd(), "manifest", "config", "config.yaml")
|
||||
var uploadPath = filepath.Join(gfile.Pwd(), "resource")
|
||||
|
||||
// Start 启动服务
|
||||
/*
|
||||
* @param agent string 浏览器标识
|
||||
* @param maxSessionTime time.Duration session最大时间
|
||||
* @param isApi bool 是否开启api
|
||||
* @param maxBody ...int64 最大上传文件大小 默认200M
|
||||
* @return *ghttp.Server 服务实例
|
||||
*/
|
||||
func Start(agent string, maxSessionTime time.Duration, isApi bool, maxBody ...int64) *ghttp.Server {
|
||||
// var s *ghttp.Server
|
||||
s := g.Server()
|
||||
s.SetDumpRouterMap(false)
|
||||
s.AddStaticPath(fmt.Sprintf("%vstatic", gfile.Separator), uploadPath)
|
||||
err := s.SetLogPath(gfile.Join(gfile.Pwd(), "resource", "log"))
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
}
|
||||
s.SetLogLevel("all")
|
||||
s.SetLogStdout(false)
|
||||
if len(maxBody) > 0 {
|
||||
s.SetClientMaxBodySize(maxBody[0])
|
||||
} else {
|
||||
s.SetClientMaxBodySize(200 * 1024 * 1024)
|
||||
}
|
||||
s.SetFormParsingMemory(50 * 1024 * 1024)
|
||||
if isApi {
|
||||
s.SetOpenApiPath("/api.json")
|
||||
s.SetSwaggerPath("/swagger")
|
||||
}
|
||||
s.SetMaxHeaderBytes(1024 * 20)
|
||||
s.SetErrorStack(true)
|
||||
s.SetSessionIdName("zrSession")
|
||||
s.SetAccessLogEnabled(true)
|
||||
s.SetSessionMaxAge(maxSessionTime)
|
||||
err = s.SetConfigWithMap(g.Map{
|
||||
"sessionPath": gfile.Join(gfile.Pwd(), "resource", "session"),
|
||||
"serverAgent": agent,
|
||||
})
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
}
|
||||
s.Use(MiddlewareError)
|
||||
enhanceOpenAPIDoc(s)
|
||||
return s
|
||||
}
|
||||
|
||||
// SetConfigAndRun 设置配置并运行服务
|
||||
// @param s *ghttp.Server 服务实例
|
||||
// @param address string 监听地址
|
||||
func SetConfigAndRun(s *ghttp.Server, address string) {
|
||||
g.Log().Info(gctx.New(), "正在设置日志配置")
|
||||
g.Log().File("{Y-m-d}.log")
|
||||
g.Log().Path(gfile.Join(gfile.Pwd(), "log"))
|
||||
g.Log().Level(glog.LEVEL_ALL)
|
||||
g.Log().SetWriterColorEnable(false)
|
||||
g.Log().SetTimeFormat("2006-01-02 15:04:05")
|
||||
g.Log().Stdout(false)
|
||||
cfg := g.Log().GetConfig()
|
||||
cfg.RotateBackupLimit = 10
|
||||
cfg.RotateSize = 1024 * 1024 * 2
|
||||
err := g.Log().SetConfig(cfg)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("设置日志配置失败: %+v", err))
|
||||
}
|
||||
s.SetAccessLogEnabled(false)
|
||||
s.SetErrorLogEnabled(true)
|
||||
sLog := s.Logger()
|
||||
sLog.Level(glog.LEVEL_ERRO)
|
||||
err = sLog.SetPath(gfile.Join(gfile.Pwd(), "resource", "log"))
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("添加服务日志路径失败: %+v", err))
|
||||
}
|
||||
sLog.SetLevelPrefix(glog.LEVEL_ERRO, "error")
|
||||
s.SetLogger(sLog)
|
||||
g.Log().Info(gctx.New(), "设置日志配置完成")
|
||||
g.Log().Info(gctx.New(), "正在设置服务监听")
|
||||
s.SetAddr(address)
|
||||
s.SetFileServerEnabled(true)
|
||||
s.SetCookieDomain(fmt.Sprintf("http://%s", address))
|
||||
g.Log().Info(gctx.New(), "设置服务监听完成,执行自动服务")
|
||||
s.Run()
|
||||
}
|
||||
|
||||
func CORSMiddleware(r *ghttp.Request) {
|
||||
corsOptions := r.Response.DefaultCORSOptions()
|
||||
cfg, _ := gcfg.Instance().Get(r.Context(), "doMain", nil)
|
||||
if !cfg.IsNil() {
|
||||
corsOptions.AllowDomain = cfg.Strings()
|
||||
}
|
||||
r.Response.CORS(corsOptions)
|
||||
r.Middleware.Next()
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/util/gconv"
|
||||
)
|
||||
|
||||
var manager = NewWs()
|
||||
|
||||
func NewWs() *Manager {
|
||||
// 1. 自定义配置(可选,也可使用默认配置)
|
||||
customConfig := &Config{
|
||||
AllowAllOrigins: true,
|
||||
HeartbeatInterval: 20 * time.Second, // 20秒发一次心跳
|
||||
HeartbeatTimeout: 40 * time.Second, // 40秒超时
|
||||
}
|
||||
|
||||
// 2. 创建管理器
|
||||
m := NewManager(customConfig)
|
||||
|
||||
// 3. 覆盖业务回调(核心:自定义消息处理逻辑)
|
||||
// 连接建立回调
|
||||
m.OnConnect = func(connID string) {
|
||||
log.Printf("业务回调:连接[%s]上线,当前在线数:%d", connID, m.GetOnlineCount())
|
||||
// 欢迎消息
|
||||
_ = m.SendToConn(connID, []byte("欢迎连接WebSocket服务!"))
|
||||
}
|
||||
|
||||
// 收到消息回调
|
||||
m.OnMessage = func(connID string, msgType int, data any) {
|
||||
log.Printf("业务回调:收到连接[%s]消息:%s", connID, gconv.String(data))
|
||||
// 示例:echo回复
|
||||
reply := []byte("服务端回复:" + gconv.String(data))
|
||||
_ = m.SendToConn(connID, reply)
|
||||
|
||||
// 示例:广播消息给所有连接
|
||||
_ = m.Broadcast([]byte("广播:" + connID + "说:" + gconv.String(data)))
|
||||
}
|
||||
|
||||
// 连接断开回调
|
||||
m.OnDisconnect = func(connID string, err error) {
|
||||
log.Printf("业务回调:连接[%s]下线,原因:%v,当前在线数:%d", connID, err, m.GetOnlineCount())
|
||||
}
|
||||
return m
|
||||
}
|
||||
func Upgrade(w http.ResponseWriter, r *http.Request, connID string) {
|
||||
_, err := manager.Upgrade(w, r, connID)
|
||||
if err != nil {
|
||||
log.Printf("升级连接失败:%v", err)
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}
|
||||
func main() {
|
||||
// 4. 注册WebSocket路由
|
||||
http.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
|
||||
// 自定义连接ID(示例:使用请求参数中的user_id)
|
||||
connID := r.URL.Query().Get("user_id")
|
||||
if connID == "" {
|
||||
http.Error(w, "user_id不能为空", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
// 升级连接
|
||||
Upgrade(w, r, connID)
|
||||
})
|
||||
|
||||
// 5. 启动服务
|
||||
log.Println("WebSocket服务启动:http://localhost:8080/ws")
|
||||
log.Fatal(http.ListenAndServe(":8080", nil))
|
||||
}
|
||||
@@ -0,0 +1,488 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/encoding/gjson"
|
||||
"github.com/gogf/gf/v2/os/gctx"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
"github.com/gogf/gf/v2/os/gtimer"
|
||||
"github.com/gogf/gf/v2/text/gstr"
|
||||
"github.com/gogf/gf/v2/util/gconv"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// 常量定义:默认配置
|
||||
const (
|
||||
// 默认读写缓冲区大小(字节)
|
||||
DefaultReadBufferSize = 1024
|
||||
DefaultWriteBufferSize = 1024
|
||||
// 默认心跳间隔(秒):每30秒发送一次心跳
|
||||
DefaultHeartbeatInterval = 30 * time.Second
|
||||
// 默认心跳超时(秒):60秒未收到客户端心跳响应则关闭连接
|
||||
DefaultHeartbeatTimeout = 60 * time.Second
|
||||
// 默认读写超时(秒)
|
||||
DefaultReadTimeout = 60 * time.Second
|
||||
DefaultWriteTimeout = 10 * time.Second
|
||||
// 消息类型
|
||||
MessageTypeText = websocket.TextMessage
|
||||
MessageTypeBinary = websocket.BinaryMessage
|
||||
// 心跳最大重试次数
|
||||
HeartbeatMaxRetry = 3
|
||||
)
|
||||
|
||||
// Config WebSocket服务端配置
|
||||
type Config struct {
|
||||
// 读写缓冲区大小
|
||||
ReadBufferSize int
|
||||
WriteBufferSize int
|
||||
// 跨域配置:是否允许所有跨域(生产环境建议指定Origin)
|
||||
AllowAllOrigins bool
|
||||
// 允许的跨域Origin列表(AllowAllOrigins=false时生效)
|
||||
AllowedOrigins []string
|
||||
// 心跳配置
|
||||
HeartbeatInterval time.Duration // 心跳发送间隔
|
||||
HeartbeatTimeout time.Duration // 心跳超时时间
|
||||
// 读写超时
|
||||
ReadTimeout time.Duration
|
||||
WriteTimeout time.Duration
|
||||
MsgType int // 发送消息的默认类型
|
||||
HeartbeatValue string // 心跳消息的标识字段值(如"heartbeat"、"pong")
|
||||
HeartbeatKey string // 心跳消息的标识字段名(如"type")
|
||||
}
|
||||
|
||||
// 默认配置
|
||||
func DefaultConfig() *Config {
|
||||
return &Config{
|
||||
ReadBufferSize: DefaultReadBufferSize,
|
||||
WriteBufferSize: DefaultWriteBufferSize,
|
||||
AllowAllOrigins: true,
|
||||
AllowedOrigins: []string{},
|
||||
HeartbeatInterval: DefaultHeartbeatInterval,
|
||||
HeartbeatTimeout: DefaultHeartbeatTimeout,
|
||||
ReadTimeout: DefaultReadTimeout,
|
||||
WriteTimeout: DefaultWriteTimeout,
|
||||
MsgType: MessageTypeText,
|
||||
HeartbeatValue: "heartbeat",
|
||||
HeartbeatKey: "type", // 心跳消息的标识字段名,默认"type"
|
||||
}
|
||||
}
|
||||
|
||||
// Connection WebSocket连接结构体
|
||||
type Connection struct {
|
||||
conn *websocket.Conn // 底层连接
|
||||
connID string // 唯一连接ID
|
||||
manager *Manager // 所属管理器
|
||||
createTime time.Time // 连接创建时间
|
||||
heartbeatChan time.Time // 心跳通道(用于检测客户端响应)
|
||||
heartbeatTime *gtimer.Entry
|
||||
ctx context.Context // 上下文
|
||||
cancel context.CancelFunc // 上下文取消函数
|
||||
writeMutex sync.Mutex // 写消息互斥锁(防止并发写)
|
||||
heartbeatRetry int // 心跳发送重试次数
|
||||
}
|
||||
|
||||
// Manager WebSocket连接管理器
|
||||
type Manager struct {
|
||||
config *Config // 配置
|
||||
upgrader *websocket.Upgrader // HTTP升级器
|
||||
connections map[string]*Connection // 所有在线连接(connID -> Connection)
|
||||
mutex sync.RWMutex // 读写锁(保护connections)
|
||||
// 业务回调:收到消息时触发(用户自定义处理逻辑)
|
||||
OnMessage func(connID string, msgType int, data any)
|
||||
// 业务回调:连接建立时触发
|
||||
OnConnect func(connID string)
|
||||
// 业务回调:连接关闭时触发
|
||||
OnDisconnect func(connID string, err error)
|
||||
}
|
||||
|
||||
// Merge 合并配置,用传入的配置覆盖非零值部分
|
||||
func (c *Config) Merge(other *Config) *Config {
|
||||
result := *c // 复制当前配置
|
||||
|
||||
if other == nil {
|
||||
return &result
|
||||
}
|
||||
|
||||
if other.ReadBufferSize > 0 {
|
||||
result.ReadBufferSize = other.ReadBufferSize
|
||||
}
|
||||
if other.WriteBufferSize > 0 {
|
||||
result.WriteBufferSize = other.WriteBufferSize
|
||||
}
|
||||
if other.HeartbeatInterval > 0 {
|
||||
result.HeartbeatInterval = other.HeartbeatInterval
|
||||
}
|
||||
if other.HeartbeatTimeout > 0 {
|
||||
result.HeartbeatTimeout = other.HeartbeatTimeout
|
||||
}
|
||||
if other.ReadTimeout > 0 {
|
||||
result.ReadTimeout = other.ReadTimeout
|
||||
}
|
||||
if other.WriteTimeout > 0 {
|
||||
result.WriteTimeout = other.WriteTimeout
|
||||
}
|
||||
if other.AllowAllOrigins {
|
||||
result.AllowAllOrigins = other.AllowAllOrigins
|
||||
}
|
||||
if other.HeartbeatValue != "" {
|
||||
result.HeartbeatValue = other.HeartbeatValue
|
||||
}
|
||||
if other.HeartbeatKey != "" {
|
||||
result.HeartbeatKey = other.HeartbeatKey
|
||||
}
|
||||
if len(other.AllowedOrigins) > 0 {
|
||||
result.AllowedOrigins = other.AllowedOrigins
|
||||
}
|
||||
if other.MsgType != 0 {
|
||||
result.MsgType = other.MsgType
|
||||
}
|
||||
|
||||
return &result
|
||||
}
|
||||
|
||||
// NewManager 创建连接管理器
|
||||
func NewManager(config *Config) *Manager {
|
||||
defaultConfig := DefaultConfig()
|
||||
finalConfig := defaultConfig.Merge(config)
|
||||
// 初始化升级器
|
||||
upgrader := &websocket.Upgrader{
|
||||
ReadBufferSize: config.ReadBufferSize,
|
||||
WriteBufferSize: config.WriteBufferSize,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
// 跨域检查
|
||||
if config.AllowAllOrigins {
|
||||
return true
|
||||
}
|
||||
origin := r.Header.Get("Origin")
|
||||
for _, allowed := range finalConfig.AllowedOrigins {
|
||||
if origin == allowed {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
},
|
||||
}
|
||||
|
||||
return &Manager{
|
||||
config: finalConfig,
|
||||
upgrader: upgrader,
|
||||
connections: make(map[string]*Connection),
|
||||
mutex: sync.RWMutex{},
|
||||
// 默认回调(用户可覆盖)
|
||||
OnMessage: func(connID string, msgType int, data any) {
|
||||
log.Printf("[默认回调] 收到连接[%s]消息:%s", connID, gconv.String(data))
|
||||
},
|
||||
OnConnect: func(connID string) {
|
||||
log.Printf("[默认回调] 连接[%s]已建立", connID)
|
||||
},
|
||||
OnDisconnect: func(connID string, err error) {
|
||||
log.Printf("[默认回调] 连接[%s]已关闭:%v", connID, err)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Upgrade HTTP升级为WebSocket连接
|
||||
// connID:自定义连接唯一ID(如用户ID、设备ID)
|
||||
func (m *Manager) Upgrade(w http.ResponseWriter, r *http.Request, connID string) (*Connection, error) {
|
||||
if connID == "" {
|
||||
return nil, errors.New("连接ID不能为空")
|
||||
}
|
||||
|
||||
// 检查连接ID是否已存在
|
||||
m.mutex.RLock()
|
||||
_, exists := m.connections[connID]
|
||||
m.mutex.RUnlock()
|
||||
if exists {
|
||||
return nil, fmt.Errorf("连接ID[%s]已存在", connID)
|
||||
}
|
||||
|
||||
// 升级HTTP连接
|
||||
conn, err := m.upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("升级WebSocket失败:%w", err)
|
||||
}
|
||||
|
||||
// 创建上下文(用于优雅关闭)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
// 创建连接实例
|
||||
wsConn := &Connection{
|
||||
conn: conn,
|
||||
connID: connID,
|
||||
manager: m,
|
||||
createTime: time.Now(),
|
||||
heartbeatChan: time.Now(), // 缓冲1,防止阻塞
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
writeMutex: sync.Mutex{},
|
||||
heartbeatRetry: 0,
|
||||
}
|
||||
wsConn.heartbeatTime = gtimer.AddSingleton(gctx.New(), m.config.HeartbeatTimeout, func(ctx context.Context) {
|
||||
log.Printf("[心跳检测] 连接[%s]已关闭:心跳超时", wsConn.connID)
|
||||
wsConn.heartbeatTime.Close()
|
||||
wsConn.heartbeatTime.Stop()
|
||||
wsConn.heartbeatTime = nil
|
||||
wsConn.ctx.Done()
|
||||
wsConn.Close(fmt.Errorf("心跳超时"))
|
||||
})
|
||||
// 添加到管理器
|
||||
m.mutex.Lock()
|
||||
m.connections[connID] = wsConn
|
||||
m.mutex.Unlock()
|
||||
|
||||
// 触发连接建立回调
|
||||
m.OnConnect(connID)
|
||||
|
||||
// 启动读消息协程
|
||||
go wsConn.ReadPump()
|
||||
// 启动写消息协程(处理异步发送)
|
||||
go wsConn.WritePump()
|
||||
// 启动心跳检测协程
|
||||
go wsConn.Heartbeat()
|
||||
|
||||
return wsConn, nil
|
||||
}
|
||||
|
||||
// ReadPump 读取客户端消息(持续运行)
|
||||
func (c *Connection) ReadPump() {
|
||||
defer func() {
|
||||
// 发生panic时关闭连接
|
||||
if err := recover(); err != nil {
|
||||
log.Printf("连接[%s]读消息协程panic:%v", c.connID, err)
|
||||
}
|
||||
// 关闭连接并清理
|
||||
c.Close(fmt.Errorf("读消息协程退出"))
|
||||
}()
|
||||
|
||||
// 循环读取消息
|
||||
for {
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return // 上下文已取消,退出
|
||||
default:
|
||||
// 设置读超时(每次读取前重置,防止长时间无消息超时)
|
||||
c.conn.SetReadDeadline(time.Now().Add(c.manager.config.ReadTimeout))
|
||||
// 读取客户端消息
|
||||
msgType, data, err := c.conn.ReadMessage()
|
||||
if err != nil {
|
||||
// 区分正常关闭和异常错误
|
||||
var closeErr *websocket.CloseError
|
||||
if errors.As(err, &closeErr) {
|
||||
c.Close(fmt.Errorf("客户端主动关闭:%s(代码:%d)", closeErr.Text, closeErr.Code))
|
||||
} else {
|
||||
c.Close(fmt.Errorf("读取消息失败:%w", err))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// 尝试解析JSON格式的心跳消息(精准判断,替代包含判断)
|
||||
isHeartbeat := false
|
||||
// 先尝试解析为JSON对象
|
||||
var msgMap map[string]interface{}
|
||||
if err := gjson.DecodeTo(data, &msgMap); err == nil {
|
||||
// 获取心跳标识字段的值
|
||||
heartbeatValue := gconv.String(msgMap[c.manager.config.HeartbeatKey])
|
||||
if heartbeatValue == c.manager.config.HeartbeatValue {
|
||||
isHeartbeat = true
|
||||
}
|
||||
} else {
|
||||
// 非JSON格式,降级为包含判断(兼容纯文本心跳)
|
||||
str := gconv.String(data)
|
||||
if gstr.Contains(str, c.manager.config.HeartbeatValue) {
|
||||
isHeartbeat = true
|
||||
}
|
||||
}
|
||||
if isHeartbeat {
|
||||
log.Printf("[心跳] 收到连接[%s]心跳消息:%s", c.connID, string(data))
|
||||
// 心跳消息:重置重试次数 + 发送心跳信号 + 重置读超时
|
||||
js, err := gjson.Encode(&Msg[any]{c.manager.config.HeartbeatValue, nil, gtime.Timestamp()})
|
||||
if err != nil {
|
||||
log.Printf("[心跳] 客户端[%s]json编码失败", c.connID)
|
||||
continue
|
||||
}
|
||||
err = c.Send(js)
|
||||
if err != nil {
|
||||
log.Printf("[心跳] 客户端[%s]发送心跳消息失败", c.connID)
|
||||
continue
|
||||
}
|
||||
c.heartbeatTime.Reset()
|
||||
continue // 跳过业务回调
|
||||
}
|
||||
|
||||
// 非心跳消息:触发业务回调
|
||||
c.manager.OnMessage(c.connID, msgType, data)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type Msg[T any] struct {
|
||||
Type string `json:"type"`
|
||||
Data T `json:"data"`
|
||||
Timestamp int64 `json:"timestamp"`
|
||||
}
|
||||
|
||||
// WritePump 处理异步写消息(持续运行)
|
||||
// 扩展为监听写队列,防止消息丢失
|
||||
func (c *Connection) WritePump() {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
log.Printf("连接[%s]写消息协程panic:%v", c.connID, err)
|
||||
}
|
||||
}()
|
||||
|
||||
// 暂时保持简化,实际可扩展为带缓冲的写队列
|
||||
<-c.ctx.Done()
|
||||
}
|
||||
|
||||
// Heartbeat 心跳检测(持续运行)
|
||||
func (c *Connection) Heartbeat() {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
log.Printf("连接[%s]心跳协程panic:%v", c.connID, err)
|
||||
}
|
||||
}()
|
||||
c.heartbeatTime.Start()
|
||||
}
|
||||
|
||||
// Send 发送消息到客户端(线程安全)
|
||||
func (c *Connection) Send(data []byte) error {
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return errors.New("连接已关闭,无法发送消息")
|
||||
default:
|
||||
// 加锁防止并发写
|
||||
c.writeMutex.Lock()
|
||||
defer c.writeMutex.Unlock()
|
||||
|
||||
// 设置写超时
|
||||
c.conn.SetWriteDeadline(time.Now().Add(c.manager.config.WriteTimeout))
|
||||
|
||||
// 发送消息(使用连接的默认类型,支持动态调整)
|
||||
err := c.conn.WriteMessage(c.manager.config.MsgType, data)
|
||||
if err != nil {
|
||||
return fmt.Errorf("发送消息失败:%w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// Close 关闭连接(优雅清理)
|
||||
func (c *Connection) Close(err error) {
|
||||
// 防止重复关闭
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
// 取消上下文(终止所有协程)
|
||||
c.cancel()
|
||||
|
||||
// 关闭底层连接(友好关闭)
|
||||
_ = c.conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, err.Error()))
|
||||
_ = c.conn.Close()
|
||||
|
||||
// 从管理器移除
|
||||
c.manager.mutex.Lock()
|
||||
delete(c.manager.connections, c.connID)
|
||||
c.manager.mutex.Unlock()
|
||||
|
||||
// 触发断开回调
|
||||
c.manager.OnDisconnect(c.connID, err)
|
||||
|
||||
log.Printf("连接[%s]已关闭,当前在线数:%d,原因:%v", c.connID, c.manager.GetOnlineCount(), err)
|
||||
}
|
||||
|
||||
// GetOnlineCount 获取在线连接数
|
||||
func (m *Manager) GetOnlineCount() int {
|
||||
m.mutex.RLock()
|
||||
defer m.mutex.RUnlock()
|
||||
return len(m.connections)
|
||||
}
|
||||
|
||||
// Broadcast 广播消息到所有在线连接
|
||||
func (m *Manager) Broadcast(data []byte) error {
|
||||
m.mutex.RLock()
|
||||
defer m.mutex.RUnlock()
|
||||
|
||||
if len(m.connections) == 0 {
|
||||
return errors.New("无在线连接")
|
||||
}
|
||||
|
||||
// 并发发送(非阻塞)
|
||||
var wg sync.WaitGroup
|
||||
var errMsg string
|
||||
|
||||
for _, conn := range m.connections {
|
||||
wg.Add(1)
|
||||
go func(c *Connection) {
|
||||
defer wg.Done()
|
||||
if err := c.Send(data); err != nil {
|
||||
errMsg += fmt.Sprintf("连接[%s]广播失败:%v;", c.connID, err)
|
||||
}
|
||||
}(conn)
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
|
||||
if errMsg != "" {
|
||||
return errors.New(errMsg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SendToConn 定向发送消息到指定连接
|
||||
func (m *Manager) SendToConn(connID string, data []byte) error {
|
||||
m.mutex.RLock()
|
||||
conn, exists := m.connections[connID]
|
||||
m.mutex.RUnlock()
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("连接[%s]不存在", connID)
|
||||
}
|
||||
|
||||
return conn.Send(data)
|
||||
}
|
||||
|
||||
func (m *Manager) GetAllConn() map[string]*Connection {
|
||||
m.mutex.RLock()
|
||||
defer m.mutex.RUnlock()
|
||||
// 返回副本,防止外部修改
|
||||
connCopy := make(map[string]*Connection, len(m.connections))
|
||||
for k, v := range m.connections {
|
||||
connCopy[k] = v
|
||||
}
|
||||
return connCopy
|
||||
}
|
||||
|
||||
func (m *Manager) GetConn(connID string) *Connection {
|
||||
m.mutex.RLock()
|
||||
defer m.mutex.RUnlock()
|
||||
return m.connections[connID]
|
||||
}
|
||||
|
||||
// CloseAll 关闭所有连接
|
||||
func (m *Manager) CloseAll() {
|
||||
m.mutex.RLock()
|
||||
connIDs := make([]string, 0, len(m.connections))
|
||||
for connID := range m.connections {
|
||||
connIDs = append(connIDs, connID)
|
||||
}
|
||||
m.mutex.RUnlock()
|
||||
|
||||
for _, connID := range connIDs {
|
||||
m.mutex.RLock()
|
||||
conn := m.connections[connID]
|
||||
m.mutex.RUnlock()
|
||||
if conn != nil {
|
||||
conn.Close(errors.New("服务端主动关闭所有连接"))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user