Files
CamTalk/backend/internal/ws/handler.go
hhs 9f876df816 feat: 初始化 Go 后端项目骨架
- 使用 Gin 框架替代原始 net/http,方便调试
- 实现 WebSocket 端点 (/ws),支持 ping/query/config/interrupt 消息
- 实现健康检查 REST 接口 (/api/health)
- 定义协议数据结构 (models)
- 心跳检测(60 秒超时断开)
2026-06-12 17:34:43 +08:00

150 lines
3.3 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"
"log"
"net/http"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"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 {
log.Printf("websocket upgrade failed: %v", 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",
})
log.Printf("client connected: session=%s", 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 {
log.Printf("heartbeat timeout: session=%s", sessionID)
conn.Close()
return
}
case <-done:
return
}
}
}()
// 消息读取循环
for {
_, message, err := conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
log.Printf("ws read error: %v", 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
}
log.Printf("query received: session=%s request=%s", sessionID, 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
}
log.Printf("config update: session=%s", sessionID)
// TODO: 更新会话配置
case "interrupt":
log.Printf("interrupt received: session=%s", sessionID)
// TODO: 中断当前 AI 响应
default:
_ = client.sendJSON(models.WsError{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "unknown message type: " + envelope.Type,
})
}
}
close(done)
log.Printf("client disconnected: session=%s", sessionID)
}