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"
)
// 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
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.
func (pc *picoConn) writeJSON(v any) error {
if pc.closed.Load() {
@ -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>"
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,6 +337,7 @@ 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)
@ -277,6 +346,7 @@ func (c *PicoChannel) handleWebSocket(w http.ResponseWriter, r *http.Request) {
logger.InfoCF("pico", "WebSocket client connected", map[string]any{
"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,
"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