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/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())
|
||||||
|
|||||||
Reference in New Issue
Block a user