fix: handler.go 心跳/版本号/CheckOrigin/历史上限改为配置驱动
- 心跳间隔和超时从 config.Server 读取 - ServerVersion 从 config.App.Version 读取 - CheckOrigin 通过 config.Server.AllowedOrigins 控制 - GetHistory limit 从 config.Session.MaxHistory 读取
This commit is contained in:
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/errors"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
@@ -17,8 +18,23 @@ import (
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
)
|
||||
|
||||
var upgrader = websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool { return true }, // 开发阶段允许所有来源
|
||||
// newUpgrader 根据配置创建 WebSocket upgrader。
|
||||
func newUpgrader(cfg *config.Config) websocket.Upgrader {
|
||||
allowedOrigins := cfg.Server.AllowedOrigins
|
||||
return websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
if len(allowedOrigins) == 0 {
|
||||
return true // 未配置则允许所有来源(开发模式)
|
||||
}
|
||||
origin := r.Header.Get("Origin")
|
||||
for _, o := range allowedOrigins {
|
||||
if o == origin || o == "*" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Client 代表一个 WebSocket 客户端连接。
|
||||
@@ -75,13 +91,21 @@ func (w *WSClient) SendError(err models.WsError) error {
|
||||
}
|
||||
|
||||
// ServeWS 处理 WebSocket 升级请求。
|
||||
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator) gin.HandlerFunc {
|
||||
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config) gin.HandlerFunc {
|
||||
upgrader := newUpgrader(cfg)
|
||||
heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second
|
||||
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
|
||||
version := cfg.App.Version
|
||||
|
||||
maxHistory := cfg.Session.MaxHistory
|
||||
|
||||
return func(c *gin.Context) {
|
||||
serveWS(c, sessionMgr, orch)
|
||||
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, maxHistory)
|
||||
}
|
||||
}
|
||||
|
||||
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator) {
|
||||
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
|
||||
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
logger.Log.Errorw("websocket upgrade failed", "error", err)
|
||||
@@ -108,7 +132,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
_ = client.SendJSON(models.WsConnected{
|
||||
Type: "connected",
|
||||
SessionID: sessionID,
|
||||
ServerVersion: "0.1.0",
|
||||
ServerVersion: version,
|
||||
})
|
||||
logger.Log.Infow("client connected", "session", sessionID)
|
||||
|
||||
@@ -122,12 +146,12 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
// 启动心跳检查 goroutine
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
ticker := time.NewTicker(heartbeatInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if time.Since(lastPong) > 60*time.Second {
|
||||
if time.Since(lastPong) > heartbeatTimeout {
|
||||
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
|
||||
conn.Close()
|
||||
return
|
||||
@@ -181,7 +205,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
}
|
||||
|
||||
// 获取对话历史
|
||||
history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, 20)
|
||||
history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, maxHistory)
|
||||
|
||||
// 创建可取消的 context
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
Reference in New Issue
Block a user