support multiple sessions per connection

This commit is contained in:
jdhxyy 2026-03-05 17:41:34 +08:00
parent 37d1713247
commit 1d36b4b0f0

View file

@ -20,15 +20,74 @@ import (
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
) )
// picoConn represents a single WebSocket connection. // picoConn represents a single WebSocket connection with session tracking.
type picoConn struct { type picoConn struct {
id string id string
conn *websocket.Conn conn *websocket.Conn
sessionID string 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 writeMu sync.Mutex
closed atomic.Bool 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. // writeJSON sends a JSON message to the connection with write locking.
func (pc *picoConn) writeJSON(v any) error { func (pc *picoConn) writeJSON(v any) error {
if pc.closed.Load() { if pc.closed.Load() {
@ -197,7 +256,7 @@ func (c *PicoChannel) SendPlaceholder(ctx context.Context, chatID string) (strin
return msgID, nil 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 { func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error {
// chatID format: "pico:<sessionID>" // chatID format: "pico:<sessionID>"
sessionID := strings.TrimPrefix(chatID, "pico:") sessionID := strings.TrimPrefix(chatID, "pico:")
@ -209,21 +268,30 @@ func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error {
if !ok { if !ok {
return true 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 { if err := pc.writeJSON(msg); err != nil {
logger.DebugCF("pico", "Write to connection failed", map[string]any{ logger.DebugCF("pico", "Write to connection failed", map[string]any{
"conn_id": pc.id, "conn_id": pc.id,
"session": sessionID,
"error": err.Error(), "error": err.Error(),
}) })
} else { } else {
sent = true sent = true
return false // Session found and message sent, stop iterating
} }
} }
return true return true
}) })
logger.DebugCF("pico", "Message sent to session",
map[string]any{
"session": sessionID,
"sent": sent,
})
if !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 return nil
} }
@ -269,6 +337,7 @@ func (c *PicoChannel) handleWebSocket(w http.ResponseWriter, r *http.Request) {
id: uuid.New().String(), id: uuid.New().String(),
conn: conn, conn: conn,
sessionID: sessionID, sessionID: sessionID,
sessions: map[string]struct{}{sessionID: {}}, // Initialize with primary session
} }
c.connections.Store(pc.id, pc) c.connections.Store(pc.id, pc)
@ -277,6 +346,7 @@ func (c *PicoChannel) handleWebSocket(w http.ResponseWriter, r *http.Request) {
logger.InfoCF("pico", "WebSocket client connected", map[string]any{ logger.InfoCF("pico", "WebSocket client connected", map[string]any{
"conn_id": pc.id, "conn_id": pc.id,
"session_id": sessionID, "session_id": sessionID,
"session_count": pc.getSessionCount(),
}) })
go c.readLoop(pc) go c.readLoop(pc)
@ -312,11 +382,19 @@ func (c *PicoChannel) authenticate(r *http.Request) bool {
func (c *PicoChannel) readLoop(pc *picoConn) { func (c *PicoChannel) readLoop(pc *picoConn) {
defer func() { defer func() {
pc.close() pc.close()
// Get all sessions for logging before cleanup
sessions := pc.getSessions()
c.connections.Delete(pc.id) c.connections.Delete(pc.id)
c.connCount.Add(-1) c.connCount.Add(-1)
// Sessions will be automatically cleaned up when picoConn is GC'd
logger.InfoCF("pico", "WebSocket client disconnected", map[string]any{ logger.InfoCF("pico", "WebSocket client disconnected", map[string]any{
"conn_id": pc.id, "conn_id": pc.id,
"session_id": pc.sessionID, "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 sessionID := msg.SessionID
if sessionID == "" { if sessionID == "" {
sessionID = pc.sessionID sessionID = pc.sessionID
} else { } else if sessionID != pc.sessionID {
// Check if sessionID already exists // Add new sessionID to connection's session set
if _, ok := c.connections.Load(sessionID); !ok { pc.addSession(sessionID)
c.connections.Store(sessionID, pc)
} logger.DebugCF("pico", "Session added to connection",
map[string]any{
"conn_id": pc.id,
"session_id": sessionID,
"session_count": pc.getSessionCount(),
})
} }
chatID := "pico:" + sessionID chatID := "pico:" + sessionID