Files
CamTalk/backend/internal/ws/handler.go
hhs da87ec776b feat: main.go 接入配置 + graceful shutdown + logger 迁移
- main.go: 接入 config.Load() 替换硬编码 :8080
- main.go: 添加 signal.NotifyContext + http.Server.Shutdown(10s drain)
- main.go: 初始化 Zap 日志,prod 环境切换 Gin release 模式
- ws/handler.go: 所有 log.Printf 替换为 logger.Log 结构化日志
- 修复 .gitignore 排除规则(/server 仅匹配根目录二进制)
- 新增 zap、viper 依赖
2026-06-13 15:18:03 +08:00

150 lines
3.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package ws
import (
"encoding/json"
"net/http"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true }, // 开发阶段允许所有来源
}
// Client 代表一个 WebSocket 客户端连接。
type Client struct {
conn *websocket.Conn
sessionID string
mu sync.Mutex
}
func (c *Client) sendJSON(v any) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.conn.WriteJSON(v)
}
// ServeWS 处理 WebSocket 升级请求。
func ServeWS(c *gin.Context) {
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
logger.Log.Errorw("websocket upgrade failed", "error", err)
return
}
defer conn.Close()
sessionID := uuid.New().String()
client := &Client{conn: conn, sessionID: sessionID}
// 发送 connected 消息
_ = client.sendJSON(models.WsConnected{
Type: "connected",
SessionID: sessionID,
ServerVersion: "0.1.0",
})
logger.Log.Infow("client connected", "session", sessionID)
// 心跳检测
lastPong := time.Now()
conn.SetPongHandler(func(string) error {
lastPong = time.Now()
return nil
})
// 启动心跳检查 goroutine
done := make(chan struct{})
go func() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if time.Since(lastPong) > 60*time.Second {
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
conn.Close()
return
}
case <-done:
return
}
}
}()
// 消息读取循环
for {
_, message, err := conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
logger.Log.Warnw("ws read error", "error", err)
}
break
}
// 解析消息类型
var envelope struct {
Type string `json:"type"`
}
if err := json.Unmarshal(message, &envelope); err != nil {
_ = client.sendJSON(models.WsError{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "invalid JSON",
})
continue
}
switch envelope.Type {
case "ping":
_ = client.sendJSON(models.WsPong{Type: "pong"})
case "query":
var msg models.WsQuery
if err := json.Unmarshal(message, &msg); err != nil {
_ = client.sendJSON(models.WsError{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "invalid query message",
RequestID: msg.RequestID,
})
continue
}
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
// TODO: 调用 AI 编排流程STT → LLM → TTS
case "config":
var msg models.WsConfig
if err := json.Unmarshal(message, &msg); err != nil {
_ = client.sendJSON(models.WsError{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "invalid config message",
})
continue
}
logger.Log.Infow("config update", "session", sessionID)
// TODO: 更新会话配置
case "interrupt":
logger.Log.Infow("interrupt received", "session", sessionID)
// TODO: 中断当前 AI 响应
default:
_ = client.sendJSON(models.WsError{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "unknown message type: " + envelope.Type,
})
}
}
close(done)
logger.Log.Infow("client disconnected", "session", sessionID)
}