diff --git a/pkg/channels/pico/pico.go b/pkg/channels/pico/pico.go index fda65a7a9..d6e333725 100644 --- a/pkg/channels/pico/pico.go +++ b/pkg/channels/pico/pico.go @@ -20,13 +20,72 @@ import ( "github.com/sipeed/picoclaw/pkg/logger" ) -// picoConn represents a single WebSocket connection. +// picoConn represents a single WebSocket connection with session tracking. type picoConn struct { - id string - conn *websocket.Conn - sessionID string - writeMu sync.Mutex - closed atomic.Bool + id string + conn *websocket.Conn + sessionID string // Primary session ID (from connection establishment) + sessions map[string]struct{} // All session IDs this connection participates in + sessionsMu sync.RWMutex // Protects sessions map + writeMu sync.Mutex + closed atomic.Bool +} + +// addSession adds a session ID to this connection. +func (pc *picoConn) addSession(sessionID string) { + if sessionID == "" || sessionID == pc.sessionID { + return + } + + pc.sessionsMu.Lock() + defer pc.sessionsMu.Unlock() + + pc.sessions[sessionID] = struct{}{} +} + +// removeSession removes a session ID from this connection. +func (pc *picoConn) removeSession(sessionID string) { + if sessionID == "" || sessionID == pc.sessionID { + return + } + + pc.sessionsMu.Lock() + defer pc.sessionsMu.Unlock() + + delete(pc.sessions, sessionID) +} + +// hasSession checks if this connection belongs to a session. +func (pc *picoConn) hasSession(sessionID string) bool { + if sessionID == "" { + return false + } + + pc.sessionsMu.RLock() + defer pc.sessionsMu.RUnlock() + + _, ok := pc.sessions[sessionID] + return ok +} + +// getSessions returns all session IDs this connection participates in. +func (pc *picoConn) getSessions() []string { + pc.sessionsMu.RLock() + defer pc.sessionsMu.RUnlock() + + result := make([]string, 0, len(pc.sessions)) + for sessionID := range pc.sessions { + result = append(result, sessionID) + } + return result +} + +// getSessionCount returns the number of sessions this connection participates in. +func (pc *picoConn) getSessionCount() int { + pc.sessionsMu.RLock() + defer pc.sessionsMu.RUnlock() + + return len(pc.sessions) } // writeJSON sends a JSON message to the connection with write locking. @@ -197,7 +256,7 @@ func (c *PicoChannel) SendPlaceholder(ctx context.Context, chatID string) (strin return msgID, nil } -// broadcastToSession sends a message to all connections with a matching session. +// broadcastToSession sends a message to the connection with a matching session. func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error { // chatID format: "pico:" sessionID := strings.TrimPrefix(chatID, "pico:") @@ -209,21 +268,30 @@ func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error { if !ok { return true } - if key == sessionID { + // Use hasSession to check if connection belongs to this session + if pc.hasSession(sessionID) { if err := pc.writeJSON(msg); err != nil { logger.DebugCF("pico", "Write to connection failed", map[string]any{ "conn_id": pc.id, + "session": sessionID, "error": err.Error(), }) } else { sent = true + return false // Session found and message sent, stop iterating } } return true }) + logger.DebugCF("pico", "Message sent to session", + map[string]any{ + "session": sessionID, + "sent": sent, + }) + if !sent { - return fmt.Errorf("no active connections for session %s: %w", sessionID, channels.ErrSendFailed) + return fmt.Errorf("no active connection for session %s: %w", sessionID, channels.ErrSendFailed) } return nil } @@ -269,14 +337,16 @@ func (c *PicoChannel) handleWebSocket(w http.ResponseWriter, r *http.Request) { id: uuid.New().String(), conn: conn, sessionID: sessionID, + sessions: map[string]struct{}{sessionID: {}}, // Initialize with primary session } c.connections.Store(pc.id, pc) c.connCount.Add(1) logger.InfoCF("pico", "WebSocket client connected", map[string]any{ - "conn_id": pc.id, - "session_id": sessionID, + "conn_id": pc.id, + "session_id": sessionID, + "session_count": pc.getSessionCount(), }) go c.readLoop(pc) @@ -312,11 +382,19 @@ func (c *PicoChannel) authenticate(r *http.Request) bool { func (c *PicoChannel) readLoop(pc *picoConn) { defer func() { pc.close() + + // Get all sessions for logging before cleanup + sessions := pc.getSessions() + c.connections.Delete(pc.id) c.connCount.Add(-1) + + // Sessions will be automatically cleaned up when picoConn is GC'd logger.InfoCF("pico", "WebSocket client disconnected", map[string]any{ - "conn_id": pc.id, - "session_id": pc.sessionID, + "conn_id": pc.id, + "session_id": pc.sessionID, + "sessions": sessions, + "session_count": len(sessions), }) }() @@ -421,11 +499,16 @@ func (c *PicoChannel) handleMessageSend(pc *picoConn, msg PicoMessage) { sessionID := msg.SessionID if sessionID == "" { sessionID = pc.sessionID - } else { - // Check if sessionID already exists - if _, ok := c.connections.Load(sessionID); !ok { - c.connections.Store(sessionID, pc) - } + } else if sessionID != pc.sessionID { + // Add new sessionID to connection's session set + pc.addSession(sessionID) + + logger.DebugCF("pico", "Session added to connection", + map[string]any{ + "conn_id": pc.id, + "session_id": sessionID, + "session_count": pc.getSessionCount(), + }) } chatID := "pico:" + sessionID