diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index e1f46e1b0..96509760d 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -143,6 +143,7 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers cmdRegistry.Register(&command.ShowCommand{}) cmdRegistry.Register(&command.ListCommand{}) cmdRegistry.Register(&command.SwitchCommand{}) + cmdRegistry.Register(&command.SessionCommand{}) // Register new session command cmdRegistry.Register(&command.StartCommand{}) cmdRegistry.Register(&command.HelpCommand{Registry: cmdRegistry}) @@ -162,6 +163,16 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers } } +// GetSessionManager exposes the session manager. +func (al *AgentLoop) GetSessionManager() interface{} { + return al.sessions +} + +// GetStateManager exposes the state manager. +func (al *AgentLoop) GetStateManager() interface{} { + return al.state +} + func (al *AgentLoop) Run(ctx context.Context) error { al.running.Store(true) @@ -175,6 +186,19 @@ func (al *AgentLoop) Run(ctx context.Context) error { continue } + // Resolve effective session key based on user's active session + // Default key is channel:chatID + // If user has active session "xyz", key becomes channel:chatID:xyz + baseKey := fmt.Sprintf("%s:%s", msg.Channel, msg.ChatID) + activeSession := al.state.GetUserSession(baseKey) + + effectiveSessionKey := msg.SessionKey + if activeSession != "" && activeSession != "default" { + effectiveSessionKey = fmt.Sprintf("%s:%s", baseKey, activeSession) + // Update msg.SessionKey so downstream logic uses the correct context + msg.SessionKey = effectiveSessionKey + } + response, err := al.processMessage(ctx, msg) if err != nil { response = fmt.Sprintf("Error processing message: %v", err) @@ -246,12 +270,21 @@ func (al *AgentLoop) ProcessDirect(ctx context.Context, content, sessionKey stri } func (al *AgentLoop) ProcessDirectWithChannel(ctx context.Context, content, sessionKey, channel, chatID string) (string, error) { + // Resolve effective session key based on user's active session + baseKey := fmt.Sprintf("%s:%s", channel, chatID) + activeSession := al.state.GetUserSession(baseKey) + + effectiveSessionKey := sessionKey + if activeSession != "" && activeSession != "default" { + effectiveSessionKey = fmt.Sprintf("%s:%s", baseKey, activeSession) + } + msg := bus.InboundMessage{ Channel: channel, SenderID: "cron", ChatID: chatID, Content: content, - SessionKey: sessionKey, + SessionKey: effectiveSessionKey, } return al.processMessage(ctx, msg) diff --git a/pkg/command/session.go b/pkg/command/session.go new file mode 100644 index 000000000..fb69d3717 --- /dev/null +++ b/pkg/command/session.go @@ -0,0 +1,160 @@ +package command + +import ( + "context" + "fmt" + "sort" + "strings" + + "github.com/sipeed/picoclaw/pkg/bus" +) + +type SessionCommand struct{} + +func (c *SessionCommand) Name() string { + return "/session" +} + +func (c *SessionCommand) Description() string { + return "Manage conversation sessions (new, switch, ls, remove)" +} + +func (c *SessionCommand) Execute(ctx context.Context, agent AgentState, args []string, msg bus.InboundMessage) (string, error) { + if len(args) < 1 { + return "Usage: /session [name]", nil + } + + subcmd := args[0] + // Use ChatID from message for unique session base key channel:chatID + baseKey := fmt.Sprintf("%s:%s", msg.Channel, msg.ChatID) + + // Define interfaces for accessing SessionManager and StateManager + type SessionManager interface { + ListSessions(prefix string) []string + DeleteSession(key string) error + } + + type StateManager interface { + SetUserSession(userID, sessionName string) error + GetUserSession(userID string) string + } + + type ManagersProvider interface { + GetSessionManager() interface{} + GetStateManager() interface{} + } + + provider, ok := agent.(ManagersProvider) + if !ok { + return "Error: Agent does not support session management", nil + } + + smRaw := provider.GetSessionManager() + stRaw := provider.GetStateManager() + + sm, okSm := smRaw.(SessionManager) + st, okSt := stRaw.(StateManager) + + if !okSm || !okSt { + return "Error: Failed to access internal managers", nil + } + + switch subcmd { + case "ls", "list": + return c.handleList(sm, st, baseKey) + case "new", "switch": + if len(args) < 2 { + return fmt.Sprintf("Usage: /session %s ", subcmd), nil + } + name := args[1] + return c.handleSwitch(st, baseKey, name) + case "remove", "rm", "delete": + if len(args) < 2 { + return "Usage: /session remove ", nil + } + name := args[1] + return c.handleRemove(sm, st, baseKey, name) + default: + return fmt.Sprintf("Unknown subcommand: %s", subcmd), nil + } +} + +func (c *SessionCommand) handleList(sm interface { + ListSessions(prefix string) []string +}, st interface { + GetUserSession(userID string) string +}, baseKey string) (string, error) { + keys := sm.ListSessions(baseKey) + currentSessionName := st.GetUserSession(baseKey) + if currentSessionName == "" { + currentSessionName = "default" + } + + sessionMap := make(map[string]bool) + sessionMap["default"] = true // Default always implicitly exists + + for _, key := range keys { + if key == baseKey { + sessionMap["default"] = true + } else if strings.HasPrefix(key, baseKey+":") { + name := strings.TrimPrefix(key, baseKey+":") + sessionMap[name] = true + } + } + + var names []string + for name := range sessionMap { + names = append(names, name) + } + sort.Strings(names) + + var sb strings.Builder + sb.WriteString("Available sessions:\n") + for _, name := range names { + marker := " " + if name == currentSessionName { + marker = "*" + } + sb.WriteString(fmt.Sprintf("%s %s\n", marker, name)) + } + + return sb.String(), nil +} + +func (c *SessionCommand) handleSwitch(st interface { + SetUserSession(userID, sessionName string) error +}, baseKey, name string) (string, error) { + target := name + if name == "default" { + target = "" // empty string denotes default session in state + } + + if err := st.SetUserSession(baseKey, target); err != nil { + return "", fmt.Errorf("failed to switch session: %w", err) + } + + return fmt.Sprintf("Switched to session: %s", name), nil +} + +func (c *SessionCommand) handleRemove(sm interface { + DeleteSession(key string) error +}, st interface { + GetUserSession(userID string) string +}, baseKey, name string) (string, error) { + if name == "default" { + return "Error: Cannot remove default session", nil + } + + current := st.GetUserSession(baseKey) + if current == name { + return "Error: Cannot remove active session. Switch to another session first.", nil + } + + // Session key format: baseKey:name + sessionKey := fmt.Sprintf("%s:%s", baseKey, name) + if err := sm.DeleteSession(sessionKey); err != nil { + return "", fmt.Errorf("failed to remove session: %w", err) + } + + return fmt.Sprintf("Session '%s' removed.", name), nil +} diff --git a/pkg/session/manager.go b/pkg/session/manager.go index 12bf33df0..0ba397b4a 100644 --- a/pkg/session/manager.go +++ b/pkg/session/manager.go @@ -280,3 +280,41 @@ func (sm *SessionManager) SetHistory(key string, history []providers.Message) { session.Updated = time.Now() } } + +// ListSessions returns a list of sessions that match the given prefix. +// The prefix is typically the channel:chatID part. +func (sm *SessionManager) ListSessions(prefix string) []string { + sm.mu.RLock() + defer sm.mu.RUnlock() + + var sessions []string + for key := range sm.sessions { + if strings.HasPrefix(key, prefix) { + sessions = append(sessions, key) + } + } + return sessions +} + +// DeleteSession removes a session from memory and disk. +func (sm *SessionManager) DeleteSession(key string) error { + sm.mu.Lock() + defer sm.mu.Unlock() + + // Remove from memory + delete(sm.sessions, key) + + if sm.storage == "" { + return nil + } + + // Remove from disk + filename := sanitizeFilename(key) + ".json" + sessionPath := filepath.Join(sm.storage, filename) + + if err := os.Remove(sessionPath); err != nil && !os.IsNotExist(err) { + return err + } + + return nil +} diff --git a/pkg/state/state.go b/pkg/state/state.go index 0bb9cd497..1d26dd413 100644 --- a/pkg/state/state.go +++ b/pkg/state/state.go @@ -21,6 +21,9 @@ type State struct { // Timestamp is the last time this state was updated Timestamp time.Time `json:"timestamp"` + + // UserSessions maps user IDs to their active session names + UserSessions map[string]string `json:"user_sessions,omitempty"` } // Manager manages persistent state with atomic saves. @@ -153,7 +156,7 @@ func (sm *Manager) saveAtomic() error { return nil } -// load loads the state from disk. +// Load loads the state from disk. func (sm *Manager) load() error { data, err := os.ReadFile(sm.stateFile) if err != nil { @@ -170,3 +173,34 @@ func (sm *Manager) load() error { return nil } + +// SetUserSession atomically updates the active session for a user. +func (sm *Manager) SetUserSession(userID, sessionName string) error { + sm.mu.Lock() + defer sm.mu.Unlock() + + if sm.state.UserSessions == nil { + sm.state.UserSessions = make(map[string]string) + } + + sm.state.UserSessions[userID] = sessionName + sm.state.Timestamp = time.Now() + + // Atomic save using temp file + rename + if err := sm.saveAtomic(); err != nil { + return fmt.Errorf("failed to save state atomically: %w", err) + } + + return nil +} + +// GetUserSession returns the active session for a user. +func (sm *Manager) GetUserSession(userID string) string { + sm.mu.RLock() + defer sm.mu.RUnlock() + + if sm.state.UserSessions == nil { + return "" + } + return sm.state.UserSessions[userID] +}