feat: Phase 7.2 — WS 连接 JWT 认证

- 从 ?token=xxx 查询参数提取 access_token
- 校验失败返回 401(missing token / invalid token)
- 校验成功后将 userID 用于创建会话
- 更新现有测试:setupTestServer 自动生成有效 token
This commit is contained in:
hhs
2026-06-14 17:47:05 +08:00
parent 2aa3c98ab6
commit 905b56640e
2 changed files with 24 additions and 5 deletions

View File

@@ -107,6 +107,21 @@ func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *co
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, tokenMgr *auth.TokenManager) { upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int, tokenMgr *auth.TokenManager) {
// --- JWT 认证upgrade 前完成,失败直接返回 HTTP 错误) ---
token := c.Query("token")
if token == "" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing token"})
return
}
claims, err := tokenMgr.ValidateAccess(token)
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
return
}
userID := claims.UserID
username := claims.Username
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)
@@ -115,7 +130,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
defer conn.Close() defer conn.Close()
// 创建会话 // 创建会话
sessionID, err := sessionMgr.Create(context.Background(), "", models.DefaultConfig()) sessionID, err := sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
if err != nil { if err != nil {
logger.Log.Errorw("create session failed", "error", err) logger.Log.Errorw("create session failed", "error", err)
return return
@@ -135,7 +150,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
SessionID: sessionID, SessionID: sessionID,
ServerVersion: version, ServerVersion: version,
}) })
logger.Log.Infow("client connected", "session", sessionID) logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username)
// 心跳检测 // 心跳检测
lastPong := time.Now() lastPong := time.Now()

View File

@@ -133,25 +133,29 @@ func (m *MockOrchestrator) ProcessQuery(
// --- 测试辅助函数 --- // --- 测试辅助函数 ---
// setupTestServer 创建测试用 Gin 服务器和 WebSocket URL。 // setupTestServer 创建测试用 Gin 服务器和 WebSocket URL。
// 返回的 wsURL 已包含有效 token可直接连接。
func setupTestServer(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, string) { func setupTestServer(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, string) {
t.Helper() t.Helper()
sessionMgr := session.NewMemoryManager(5*time.Minute, 20) sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
t.Cleanup(func() { sessionMgr.Stop() }) t.Cleanup(func() { sessionMgr.Stop() })
tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
r := gin.New() r := gin.New()
cfg := &config.Config{ cfg := &config.Config{
App: config.AppConfig{Version: "test"}, App: config.AppConfig{Version: "test"},
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60}, Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
Session: config.SessionConfig{MaxHistory: 20}, Session: config.SessionConfig{MaxHistory: 20},
} }
tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr)) r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr))
srv := httptest.NewServer(r) srv := httptest.NewServer(r)
// 构造 WebSocket URL // 生成有效 token 并构造 WebSocket URL
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws" token, _, err := tokenMgr.GeneratePair("test-user", "testuser")
require.NoError(t, err)
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
return srv, wsURL return srv, wsURL
} }