fix: handler.go 心跳/版本号/CheckOrigin/历史上限改为配置驱动

- 心跳间隔和超时从 config.Server 读取
- ServerVersion 从 config.App.Version 读取
- CheckOrigin 通过 config.Server.AllowedOrigins 控制
- GetHistory limit 从 config.Session.MaxHistory 读取
This commit is contained in:
hhs
2026-06-14 11:55:15 +08:00
parent 63f8cc279d
commit f1ce28966c

View File

@@ -10,6 +10,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/gorilla/websocket" "github.com/gorilla/websocket"
"github.com/hhs/camtalk/internal/config"
"github.com/hhs/camtalk/internal/errors" "github.com/hhs/camtalk/internal/errors"
"github.com/hhs/camtalk/internal/logger" "github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models" "github.com/hhs/camtalk/internal/models"
@@ -17,8 +18,23 @@ import (
"github.com/hhs/camtalk/internal/session" "github.com/hhs/camtalk/internal/session"
) )
var upgrader = websocket.Upgrader{ // newUpgrader 根据配置创建 WebSocket upgrader
CheckOrigin: func(r *http.Request) bool { return true }, // 开发阶段允许所有来源 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 客户端连接。 // Client 代表一个 WebSocket 客户端连接。
@@ -75,13 +91,21 @@ func (w *WSClient) SendError(err models.WsError) error {
} }
// ServeWS 处理 WebSocket 升级请求。 // 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) { 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) conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil { if err != nil {
logger.Log.Errorw("websocket upgrade failed", "error", err) 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{ _ = client.SendJSON(models.WsConnected{
Type: "connected", Type: "connected",
SessionID: sessionID, SessionID: sessionID,
ServerVersion: "0.1.0", ServerVersion: version,
}) })
logger.Log.Infow("client connected", "session", sessionID) logger.Log.Infow("client connected", "session", sessionID)
@@ -122,12 +146,12 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
// 启动心跳检查 goroutine // 启动心跳检查 goroutine
done := make(chan struct{}) done := make(chan struct{})
go func() { go func() {
ticker := time.NewTicker(30 * time.Second) ticker := time.NewTicker(heartbeatInterval)
defer ticker.Stop() defer ticker.Stop()
for { for {
select { select {
case <-ticker.C: case <-ticker.C:
if time.Since(lastPong) > 60*time.Second { if time.Since(lastPong) > heartbeatTimeout {
logger.Log.Warnw("heartbeat timeout", "session", sessionID) logger.Log.Warnw("heartbeat timeout", "session", sessionID)
conn.Close() conn.Close()
return 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 // 创建可取消的 context
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())