feat: Phase 7.2 — WS 连接 JWT 认证
- 从 ?token=xxx 查询参数提取 access_token - 校验失败返回 401(missing token / invalid token) - 校验成功后将 userID 用于创建会话 - 更新现有测试:setupTestServer 自动生成有效 token
This commit is contained in:
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user