diff --git a/pkg/channels/interfaces.go b/pkg/channels/interfaces.go index 42f18405c..2bb4da3e0 100644 --- a/pkg/channels/interfaces.go +++ b/pkg/channels/interfaces.go @@ -48,3 +48,10 @@ type PlaceholderRecorder interface { type CommandRegistrarCapable interface { RegisterCommands(ctx context.Context, defs []commands.Definition) error } + +// CommandParserCapable is implemented by channels that expose a command +// dispatch entrypoint backed by shared command definitions/dispatcher. +// It is optional and intended for cross-channel command handling features. +type CommandParserCapable interface { + DispatchCommand(ctx context.Context, req commands.Request) commands.Result +} diff --git a/pkg/channels/interfaces_command_test.go b/pkg/channels/interfaces_command_test.go index de5502644..947f6b5c1 100644 --- a/pkg/channels/interfaces_command_test.go +++ b/pkg/channels/interfaces_command_test.go @@ -11,6 +11,16 @@ type mockRegistrar struct{} func (mockRegistrar) RegisterCommands(context.Context, []commands.Definition) error { return nil } +type mockParser struct{} + +func (mockParser) DispatchCommand(context.Context, commands.Request) commands.Result { + return commands.Result{Matched: false} +} + func TestCommandRegistrarCapable_Compiles(t *testing.T) { var _ CommandRegistrarCapable = mockRegistrar{} } + +func TestCommandParserCapable_Compiles(t *testing.T) { + var _ CommandParserCapable = mockParser{} +} diff --git a/pkg/channels/telegram/command_registration_test.go b/pkg/channels/telegram/command_registration_test.go index 39237ab95..26f891b2e 100644 --- a/pkg/channels/telegram/command_registration_test.go +++ b/pkg/channels/telegram/command_registration_test.go @@ -3,6 +3,7 @@ package telegram import ( "context" "errors" + "sync/atomic" "testing" "time" @@ -28,3 +29,68 @@ func TestStartCommandRegistration_DoesNotBlock(t *testing.T) { t.Fatal("registration did not start asynchronously") } } + +func TestStartCommandRegistration_RetriesUntilSuccessThenStops(t *testing.T) { + ch := &TelegramChannel{} + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + origBackoff := commandRegistrationBackoff + commandRegistrationBackoff = []time.Duration{5 * time.Millisecond} + defer func() { commandRegistrationBackoff = origBackoff }() + + var attempts atomic.Int32 + ch.registerFunc = func(context.Context, []commands.Definition) error { + n := attempts.Add(1) + if n < 3 { + return errors.New("temporary failure") + } + return nil + } + + ch.startCommandRegistration(ctx, []commands.Definition{{Name: "help", Description: "Help"}}) + + deadline := time.Now().Add(250 * time.Millisecond) + for time.Now().Before(deadline) { + if attempts.Load() >= 3 { + break + } + time.Sleep(5 * time.Millisecond) + } + if attempts.Load() < 3 { + t.Fatalf("expected at least 3 attempts, got %d", attempts.Load()) + } + + stable := attempts.Load() + time.Sleep(30 * time.Millisecond) + if attempts.Load() != stable { + t.Fatalf("expected retries to stop after success, got %d -> %d", stable, attempts.Load()) + } +} + +func TestStartCommandRegistration_StopsAfterCancel(t *testing.T) { + ch := &TelegramChannel{} + ctx, cancel := context.WithCancel(context.Background()) + + origBackoff := commandRegistrationBackoff + commandRegistrationBackoff = []time.Duration{5 * time.Millisecond} + defer func() { commandRegistrationBackoff = origBackoff }() + defer cancel() + + var attempts atomic.Int32 + ch.registerFunc = func(context.Context, []commands.Definition) error { + attempts.Add(1) + return errors.New("always fail") + } + + ch.startCommandRegistration(ctx, []commands.Definition{{Name: "help", Description: "Help"}}) + + time.Sleep(20 * time.Millisecond) + cancel() + time.Sleep(20 * time.Millisecond) // allow in-flight attempt to settle + stable := attempts.Load() + time.Sleep(30 * time.Millisecond) + if attempts.Load() != stable { + t.Fatalf("expected retries to quiesce after cancel, got %d -> %d", stable, attempts.Load()) + } +} diff --git a/pkg/channels/telegram/telegram_commands.go b/pkg/channels/telegram/telegram_commands.go index c17961835..cd24fa9e9 100644 --- a/pkg/channels/telegram/telegram_commands.go +++ b/pkg/channels/telegram/telegram_commands.go @@ -40,7 +40,7 @@ func commandArgs(text string) string { func (c *cmd) Help(ctx context.Context, message telego.Message) error { defs := commands.NewRegistry(commands.BuiltinDefinitions(c.config)).ForChannel("telegram") - msg := formatHelpMessage(defs) + msg := commands.FormatHelpMessage(defs) _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ ChatID: telego.ChatID{ID: message.Chat.ID}, Text: msg, @@ -51,26 +51,6 @@ func (c *cmd) Help(ctx context.Context, message telego.Message) error { return err } -func formatHelpMessage(defs []commands.Definition) string { - if len(defs) == 0 { - return "No commands available." - } - - lines := make([]string, 0, len(defs)) - for _, def := range defs { - usage := def.Usage - if usage == "" { - usage = "/" + def.Name - } - desc := def.Description - if desc == "" { - desc = "No description" - } - lines = append(lines, fmt.Sprintf("%s - %s", usage, desc)) - } - return strings.Join(lines, "\n") -} - func (c *cmd) Start(ctx context.Context, message telego.Message) error { _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ ChatID: telego.ChatID{ID: message.Chat.ID}, diff --git a/pkg/channels/telegram/telegram_dispatch.go b/pkg/channels/telegram/telegram_dispatch.go index 5c6014bd8..57ae7951d 100644 --- a/pkg/channels/telegram/telegram_dispatch.go +++ b/pkg/channels/telegram/telegram_dispatch.go @@ -10,56 +10,45 @@ import ( "github.com/sipeed/picoclaw/pkg/logger" ) -func (c *TelegramChannel) dispatchCommand(ctx context.Context, message telego.Message) bool { +func (c *TelegramChannel) DispatchCommand(ctx context.Context, req commands.Request) commands.Result { if c.dispatcher == nil { - return false + return commands.Result{Matched: false} } + return c.dispatcher.Dispatch(ctx, req) +} +func (c *TelegramChannel) dispatchCommand(ctx context.Context, message telego.Message) bool { senderID := "" if message.From != nil { senderID = strconv.FormatInt(message.From.ID, 10) } - res := c.dispatcher.Dispatch(ctx, commands.Request{ + res := c.DispatchCommand(ctx, commands.Request{ Channel: "telegram", ChatID: strconv.FormatInt(message.Chat.ID, 10), SenderID: senderID, Text: message.Text, MessageID: strconv.Itoa(message.MessageID), + Reply: func(text string) error { + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: text, + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err + }, }) if !res.Matched { return false } - switch res.Command { - case "help": - if err := c.commands.Help(ctx, message); err != nil { - logger.ErrorCF("telegram", "Command execution failed", map[string]any{ - "command": "help", - "error": err.Error(), - }) - } - case "start": - if err := c.commands.Start(ctx, message); err != nil { - logger.ErrorCF("telegram", "Command execution failed", map[string]any{ - "command": "start", - "error": err.Error(), - }) - } - case "show": - if err := c.commands.Show(ctx, message); err != nil { - logger.ErrorCF("telegram", "Command execution failed", map[string]any{ - "command": "show", - "error": err.Error(), - }) - } - case "list": - if err := c.commands.List(ctx, message); err != nil { - logger.ErrorCF("telegram", "Command execution failed", map[string]any{ - "command": "list", - "error": err.Error(), - }) - } + if res.Err != nil { + logger.ErrorCF("telegram", "Command execution failed", map[string]any{ + "command": res.Command, + "error": res.Err.Error(), + }) } return true diff --git a/pkg/channels/whatsapp/whatsapp.go b/pkg/channels/whatsapp/whatsapp.go index f0db7b992..c1a95eb5f 100644 --- a/pkg/channels/whatsapp/whatsapp.go +++ b/pkg/channels/whatsapp/whatsapp.go @@ -20,14 +20,14 @@ import ( type WhatsAppChannel struct { *channels.BaseChannel - conn *websocket.Conn - config config.WhatsAppConfig - url string + conn *websocket.Conn + config config.WhatsAppConfig + url string dispatcher commands.Dispatching - ctx context.Context - cancel context.CancelFunc - mu sync.Mutex - connected bool + ctx context.Context + cancel context.CancelFunc + mu sync.Mutex + connected bool } func NewWhatsAppChannel(cfg config.WhatsAppConfig, bus *bus.MessageBus) (*WhatsAppChannel, error) { @@ -262,16 +262,15 @@ func (c *WhatsAppChannel) tryHandleCommand( ctx context.Context, text, chatID, senderID, messageID string, ) bool { - if c.dispatcher == nil { - return false - } - - res := c.dispatcher.Dispatch(ctx, commands.Request{ + res := c.DispatchCommand(ctx, commands.Request{ Channel: "whatsapp", ChatID: chatID, SenderID: senderID, Text: text, MessageID: messageID, + Reply: func(text string) error { + return c.Send(ctx, bus.OutboundMessage{ChatID: chatID, Content: text}) + }, }) if res.Err != nil { logger.WarnCF("whatsapp", "Command execution failed", map[string]any{ @@ -279,6 +278,12 @@ func (c *WhatsAppChannel) tryHandleCommand( "error": res.Err.Error(), }) } - - return res.Handled + return res.Matched +} + +func (c *WhatsAppChannel) DispatchCommand(ctx context.Context, req commands.Request) commands.Result { + if c.dispatcher == nil { + return commands.Result{Matched: false} + } + return c.dispatcher.Dispatch(ctx, req) } diff --git a/pkg/channels/whatsapp/whatsapp_command_test.go b/pkg/channels/whatsapp/whatsapp_command_test.go index 55f7f35da..1de2f2712 100644 --- a/pkg/channels/whatsapp/whatsapp_command_test.go +++ b/pkg/channels/whatsapp/whatsapp_command_test.go @@ -20,3 +20,15 @@ func TestTryHandleCommand_UsesDispatcher(t *testing.T) { t.Fatalf("handled=%v called=%v", handled, called) } } + +func TestTryHandleCommand_MatchedWithoutHandler_DoesNotFallThrough(t *testing.T) { + ch := &WhatsAppChannel{} + ch.dispatcher = commands.DispatchFunc(func(context.Context, commands.Request) commands.Result { + return commands.Result{Matched: true, Handled: false, Command: "unknown"} + }) + + handled := ch.tryHandleCommand(context.Background(), "/unknown", "chat1", "user1", "mid1") + if !handled { + t.Fatal("expected matched command to be treated as handled") + } +} diff --git a/pkg/channels/whatsapp_native/whatsapp_command_test.go b/pkg/channels/whatsapp_native/whatsapp_command_test.go index 32e4c672c..3fe7e31bd 100644 --- a/pkg/channels/whatsapp_native/whatsapp_command_test.go +++ b/pkg/channels/whatsapp_native/whatsapp_command_test.go @@ -22,3 +22,15 @@ func TestTryHandleCommand_UsesDispatcher(t *testing.T) { t.Fatalf("handled=%v called=%v", handled, called) } } + +func TestTryHandleCommand_MatchedWithoutHandler_DoesNotFallThrough(t *testing.T) { + ch := &WhatsAppNativeChannel{} + ch.dispatcher = commands.DispatchFunc(func(context.Context, commands.Request) commands.Result { + return commands.Result{Matched: true, Handled: false, Command: "unknown"} + }) + + handled := ch.tryHandleCommand(context.Background(), "/unknown", "chat1", "user1", "mid1") + if !handled { + t.Fatal("expected matched command to be treated as handled") + } +} diff --git a/pkg/channels/whatsapp_native/whatsapp_native.go b/pkg/channels/whatsapp_native/whatsapp_native.go index a15a72aca..0f093775f 100644 --- a/pkg/channels/whatsapp_native/whatsapp_native.go +++ b/pkg/channels/whatsapp_native/whatsapp_native.go @@ -406,15 +406,15 @@ func (c *WhatsAppNativeChannel) tryHandleCommand( ctx context.Context, text, chatID, senderID, messageID string, ) bool { - if c.dispatcher == nil { - return false - } - res := c.dispatcher.Dispatch(ctx, commands.Request{ + res := c.DispatchCommand(ctx, commands.Request{ Channel: "whatsapp_native", ChatID: chatID, SenderID: senderID, Text: text, MessageID: messageID, + Reply: func(text string) error { + return c.Send(ctx, bus.OutboundMessage{ChatID: chatID, Content: text}) + }, }) if res.Err != nil { logger.WarnCF("whatsapp", "Command execution failed", map[string]any{ @@ -422,7 +422,14 @@ func (c *WhatsAppNativeChannel) tryHandleCommand( "error": res.Err.Error(), }) } - return res.Handled + return res.Matched +} + +func (c *WhatsAppNativeChannel) DispatchCommand(ctx context.Context, req commands.Request) commands.Result { + if c.dispatcher == nil { + return commands.Result{Matched: false} + } + return c.dispatcher.Dispatch(ctx, req) } func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessage) error { diff --git a/pkg/commands/builtin.go b/pkg/commands/builtin.go index 4b97a14a4..f9622bb6c 100644 --- a/pkg/commands/builtin.go +++ b/pkg/commands/builtin.go @@ -1,32 +1,167 @@ package commands -import "github.com/sipeed/picoclaw/pkg/config" +import ( + "context" + "fmt" + "strings" -func BuiltinDefinitions(_ *config.Config) []Definition { + "github.com/sipeed/picoclaw/pkg/config" +) + +func BuiltinDefinitions(cfg *config.Config) []Definition { return []Definition{ { Name: "start", Description: "Start the bot", Usage: "/start", Channels: []string{"telegram", "whatsapp", "whatsapp_native"}, + Handler: replyText("Hello! I am PicoClaw 🦞"), }, { Name: "help", Description: "Show this help message", Usage: "/help", Channels: []string{"telegram", "whatsapp", "whatsapp_native"}, + Handler: func(_ context.Context, req Request) error { + if req.Reply == nil { + return nil + } + defs := NewRegistry(BuiltinDefinitions(cfg)).ForChannel(req.Channel) + return req.Reply(FormatHelpMessage(defs)) + }, }, { Name: "show", Description: "Show current configuration", Usage: "/show [model|channel]", - Channels: []string{"telegram", "whatsapp", "whatsapp_native"}, + Channels: []string{"telegram"}, + Handler: func(_ context.Context, req Request) error { + if req.Reply == nil { + return nil + } + if cfg == nil { + return req.Reply("Command unavailable in current context.") + } + args := commandArgs(req.Text) + if args == "" { + return req.Reply("Usage: /show [model|channel]") + } + + switch args { + case "model": + return req.Reply(fmt.Sprintf( + "Current Model: %s (Provider: %s)", + cfg.Agents.Defaults.GetModelName(), + cfg.Agents.Defaults.Provider, + )) + case "channel": + return req.Reply(fmt.Sprintf("Current Channel: %s", req.Channel)) + default: + return req.Reply(fmt.Sprintf("Unknown parameter: %s. Try 'model' or 'channel'.", args)) + } + }, }, { Name: "list", Description: "List available options", Usage: "/list [models|channels]", - Channels: []string{"telegram", "whatsapp", "whatsapp_native"}, + Channels: []string{"telegram"}, + Handler: func(_ context.Context, req Request) error { + if req.Reply == nil { + return nil + } + if cfg == nil { + return req.Reply("Command unavailable in current context.") + } + args := commandArgs(req.Text) + if args == "" { + return req.Reply("Usage: /list [models|channels]") + } + + switch args { + case "models": + provider := cfg.Agents.Defaults.Provider + if provider == "" { + provider = "configured default" + } + return req.Reply(fmt.Sprintf( + "Configured Model: %s\nProvider: %s\n\nTo change models, update config.json", + cfg.Agents.Defaults.GetModelName(), + provider, + )) + case "channels": + enabled := enabledChannels(cfg) + return req.Reply(fmt.Sprintf("Enabled Channels:\n- %s", strings.Join(enabled, "\n- "))) + default: + return req.Reply(fmt.Sprintf("Unknown parameter: %s. Try 'models' or 'channels'.", args)) + } + }, }, } } + +func FormatHelpMessage(defs []Definition) string { + if len(defs) == 0 { + return "No commands available." + } + + lines := make([]string, 0, len(defs)) + for _, def := range defs { + usage := def.Usage + if usage == "" { + usage = "/" + def.Name + } + desc := def.Description + if desc == "" { + desc = "No description" + } + lines = append(lines, fmt.Sprintf("%s - %s", usage, desc)) + } + return strings.Join(lines, "\n") +} + +func commandArgs(text string) string { + parts := strings.SplitN(text, " ", 2) + if len(parts) < 2 { + return "" + } + return strings.TrimSpace(parts[1]) +} + +func replyText(text string) Handler { + return func(_ context.Context, req Request) error { + if req.Reply == nil { + return nil + } + return req.Reply(text) + } +} + +func enabledChannels(cfg *config.Config) []string { + enabled := make([]string, 0, 8) + if cfg.Channels.Telegram.Enabled { + enabled = append(enabled, "telegram") + } + if cfg.Channels.WhatsApp.Enabled { + enabled = append(enabled, "whatsapp") + } + if cfg.Channels.Feishu.Enabled { + enabled = append(enabled, "feishu") + } + if cfg.Channels.Discord.Enabled { + enabled = append(enabled, "discord") + } + if cfg.Channels.Slack.Enabled { + enabled = append(enabled, "slack") + } + if cfg.Channels.DingTalk.Enabled { + enabled = append(enabled, "dingtalk") + } + if cfg.Channels.LINE.Enabled { + enabled = append(enabled, "line") + } + if cfg.Channels.OneBot.Enabled { + enabled = append(enabled, "onebot") + } + return enabled +} diff --git a/pkg/commands/builtin_test.go b/pkg/commands/builtin_test.go index a84cb65bd..f1a302ebd 100644 --- a/pkg/commands/builtin_test.go +++ b/pkg/commands/builtin_test.go @@ -14,3 +14,17 @@ func TestBuiltinDefinitions_ContainsTelegramDefaults(t *testing.T) { } } } + +func TestBuiltinDefinitions_WhatsAppOnlyHasBasicCommands(t *testing.T) { + defs := NewRegistry(BuiltinDefinitions(nil)).ForChannel("whatsapp") + names := map[string]bool{} + for _, d := range defs { + names[d.Name] = true + } + if !names["start"] || !names["help"] { + t.Fatalf("whatsapp should include start/help, got %+v", names) + } + if names["show"] || names["list"] { + t.Fatalf("whatsapp should not include show/list, got %+v", names) + } +} diff --git a/pkg/commands/dispatcher.go b/pkg/commands/dispatcher.go index c149e0f86..031967018 100644 --- a/pkg/commands/dispatcher.go +++ b/pkg/commands/dispatcher.go @@ -13,6 +13,7 @@ type Request struct { SenderID string Text string MessageID string + Reply func(text string) error } type Result struct { @@ -41,14 +42,13 @@ func NewDispatcher(reg *Registry) *Dispatcher { } func (d *Dispatcher) Dispatch(ctx context.Context, req Request) Result { - token := firstToken(req.Text) - if token == "" { + cmdName, ok := parseCommandName(req.Text) + if !ok { return Result{Matched: false} } - cmdName := strings.TrimPrefix(token, "/") for _, def := range d.reg.ForChannel(req.Channel) { - if def.Name != cmdName { + if def.Name != cmdName && !contains(def.Aliases, cmdName) { continue } if def.Handler == nil { @@ -68,3 +68,29 @@ func firstToken(input string) string { } return parts[0] } + +func parseCommandName(input string) (string, bool) { + token := firstToken(input) + if token == "" || !strings.HasPrefix(token, "/") { + return "", false + } + + name := strings.TrimPrefix(token, "/") + if i := strings.Index(name, "@"); i >= 0 { + name = name[:i] + } + name = strings.TrimSpace(name) + if name == "" { + return "", false + } + return name, true +} + +func contains(items []string, target string) bool { + for _, item := range items { + if item == target { + return true + } + } + return false +} diff --git a/pkg/commands/dispatcher_test.go b/pkg/commands/dispatcher_test.go index 0e1c1c1c0..44626d323 100644 --- a/pkg/commands/dispatcher_test.go +++ b/pkg/commands/dispatcher_test.go @@ -26,3 +26,36 @@ func TestDispatcher_MatchSlashCommand(t *testing.T) { t.Fatalf("dispatch result = %+v, called=%v", res, called) } } + +func TestDispatcher_DoesNotMatchWithoutSlash(t *testing.T) { + d := NewDispatcher(NewRegistry([]Definition{{Name: "help"}})) + + res := d.Dispatch(context.Background(), Request{ + Channel: "telegram", + Text: "help", + }) + if res.Matched { + t.Fatalf("expected unmatched for plain text, got %+v", res) + } +} + +func TestDispatcher_MatchTelegramMentionSyntax(t *testing.T) { + called := false + d := NewDispatcher(NewRegistry([]Definition{ + { + Name: "help", + Handler: func(context.Context, Request) error { + called = true + return nil + }, + }, + })) + + res := d.Dispatch(context.Background(), Request{ + Channel: "telegram", + Text: "/help@my_bot", + }) + if !res.Matched || !res.Handled || !called || res.Err != nil { + t.Fatalf("dispatch result = %+v, called=%v", res, called) + } +}