diff --git a/cmd/picoclaw/internal/gateway/helpers.go b/cmd/picoclaw/internal/gateway/helpers.go index fed3d5ffb..2d84d9e0a 100644 --- a/cmd/picoclaw/internal/gateway/helpers.go +++ b/cmd/picoclaw/internal/gateway/helpers.go @@ -25,6 +25,7 @@ import ( _ "github.com/sipeed/picoclaw/pkg/channels/qq" _ "github.com/sipeed/picoclaw/pkg/channels/slack" _ "github.com/sipeed/picoclaw/pkg/channels/telegram" + tgchannel "github.com/sipeed/picoclaw/pkg/channels/telegram" _ "github.com/sipeed/picoclaw/pkg/channels/wecom" _ "github.com/sipeed/picoclaw/pkg/channels/whatsapp" _ "github.com/sipeed/picoclaw/pkg/channels/whatsapp_native" @@ -143,6 +144,12 @@ func gatewayCmd(debug bool) error { logger.InfoCF("voice", "Transcription enabled (agent-level)", map[string]any{"provider": transcriber.Name()}) } + if telegramCh, ok := channelManager.GetChannel("telegram"); ok { + if tc, ok := telegramCh.(*tgchannel.TelegramChannel); ok { + tc.SetRegistry(agentLoop.GetRegistry()) + } + } + enabledChannels := channelManager.GetEnabledChannels() if len(enabledChannels) > 0 { fmt.Printf("✓ Channels enabled: %s\n", enabledChannels) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 235d42fcc..76c7847c1 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -546,6 +546,10 @@ func inferMediaType(filename, contentType string) string { return "file" } +func (al *AgentLoop) GetRegistry() *AgentRegistry { + return al.registry +} + // RecordLastChannel records the last active channel for this workspace. // This uses the atomic state save mechanism to prevent data loss on crash. func (al *AgentLoop) RecordLastChannel(channel string) error { diff --git a/pkg/agent/registry.go b/pkg/agent/registry.go index 58b7ce440..b3ce46da6 100644 --- a/pkg/agent/registry.go +++ b/pkg/agent/registry.go @@ -1,6 +1,8 @@ package agent import ( + "fmt" + "strings" "sync" "github.com/sipeed/picoclaw/pkg/config" @@ -13,6 +15,7 @@ import ( // AgentRegistry manages multiple agent instances and routes messages to them. type AgentRegistry struct { agents map[string]*AgentInstance + cfg *config.Config resolver *routing.RouteResolver mu sync.RWMutex } @@ -24,6 +27,7 @@ func NewAgentRegistry( ) *AgentRegistry { registry := &AgentRegistry{ agents: make(map[string]*AgentInstance), + cfg: cfg, resolver: routing.NewRouteResolver(cfg), } @@ -130,6 +134,68 @@ func (r *AgentRegistry) Close() { func (r *AgentRegistry) GetDefaultAgent() *AgentInstance { r.mu.RLock() defer r.mu.RUnlock() + return r.defaultAgentLocked() +} + +// GetDefaultAgentModel returns the active model name of the default agent. +func (r *AgentRegistry) GetDefaultAgentModel() string { + r.mu.RLock() + defer r.mu.RUnlock() + agent := r.defaultAgentLocked() + if agent == nil { + return "" + } + return agent.Model +} + +// SwitchDefaultAgentModel switches the default agent to a named model from config.model_list. +// It returns old and new runtime model IDs. +func (r *AgentRegistry) SwitchDefaultAgentModel(modelName string) (string, string, error) { + modelName = strings.TrimSpace(modelName) + if modelName == "" { + return "", "", fmt.Errorf("model name is required") + } + + if r.cfg == nil { + return "", "", fmt.Errorf("registry config not available") + } + + modelCfg, err := r.cfg.GetModelConfig(modelName) + if err != nil { + return "", "", err + } + + resolved := *modelCfg + if resolved.Workspace == "" { + resolved.Workspace = r.cfg.WorkspacePath() + } + + provider, modelID, err := providers.CreateProviderFromConfig(&resolved) + if err != nil { + return "", "", err + } + + protocol, _ := providers.ExtractProtocol(resolved.Model) + + r.mu.Lock() + defer r.mu.Unlock() + agent := r.defaultAgentLocked() + if agent == nil { + return "", "", fmt.Errorf("no default agent configured") + } + + oldModel := agent.Model + agent.Provider = provider + agent.Model = modelID + agent.Candidates = []providers.FallbackCandidate{{ + Provider: protocol, + Model: modelID, + }} + + return oldModel, modelID, nil +} + +func (r *AgentRegistry) defaultAgentLocked() *AgentInstance { if agent, ok := r.agents["main"]; ok { return agent } diff --git a/pkg/agent/registry_test.go b/pkg/agent/registry_test.go index 518bb441f..417f2452a 100644 --- a/pkg/agent/registry_test.go +++ b/pkg/agent/registry_test.go @@ -203,3 +203,68 @@ func TestAgentInstance_FallbackExplicitEmpty(t *testing.T) { t.Errorf("expected 0 fallbacks (explicit empty), got %d: %v", len(agent.Fallbacks), agent.Fallbacks) } } + +func TestAgentRegistry_SwitchDefaultAgentModel(t *testing.T) { + cfg := testCfg(nil) + cfg.Agents.Defaults.ModelName = "test-openai-mini" + cfg.ModelList = []config.ModelConfig{ + { + ModelName: "test-openai-mini", + Model: "openai/fake-openai-mini", + APIKey: "test-openai-key", + }, + { + ModelName: "test-qwen-plus", + Model: "qwen/fake-qwen-plus", + APIKey: "test-qwen-key", + }, + } + + registry := NewAgentRegistry(cfg, &mockRegistryProvider{}) + + oldModel, newModel, err := registry.SwitchDefaultAgentModel("test-qwen-plus") + if err != nil { + t.Fatalf("SwitchDefaultAgentModel() error = %v", err) + } + if oldModel != "test-openai-mini" { + t.Errorf("oldModel = %q, want %q", oldModel, "test-openai-mini") + } + if newModel != "fake-qwen-plus" { + t.Errorf("newModel = %q, want %q", newModel, "fake-qwen-plus") + } + + agent := registry.GetDefaultAgent() + if agent == nil { + t.Fatal("expected default agent") + } + if agent.Model != "fake-qwen-plus" { + t.Errorf("agent.Model = %q, want %q", agent.Model, "fake-qwen-plus") + } + + if len(agent.Candidates) != 1 { + t.Fatalf("len(agent.Candidates) = %d, want 1", len(agent.Candidates)) + } + if agent.Candidates[0].Provider != "qwen" || agent.Candidates[0].Model != "fake-qwen-plus" { + t.Errorf( + "candidate = %s/%s, want qwen/fake-qwen-plus", + agent.Candidates[0].Provider, + agent.Candidates[0].Model, + ) + } +} + +func TestAgentRegistry_SwitchDefaultAgentModel_NotFound(t *testing.T) { + cfg := testCfg(nil) + cfg.ModelList = []config.ModelConfig{ + { + ModelName: "test-openai-mini", + Model: "openai/fake-openai-mini", + APIKey: "test-openai-key", + }, + } + + registry := NewAgentRegistry(cfg, &mockRegistryProvider{}) + if _, _, err := registry.SwitchDefaultAgentModel("missing-model"); err == nil { + t.Fatal("expected error for missing model") + } +} diff --git a/pkg/channels/telegram/telegram.go b/pkg/channels/telegram/telegram.go index b04beeb6e..2462df1a9 100644 --- a/pkg/channels/telegram/telegram.go +++ b/pkg/channels/telegram/telegram.go @@ -40,12 +40,13 @@ var ( type TelegramChannel struct { *channels.BaseChannel - bot *telego.Bot - bh *th.BotHandler - config *config.Config - chatIDs map[string]int64 - ctx context.Context - cancel context.CancelFunc + bot *telego.Bot + bh *th.BotHandler + commands TelegramCommander + config *config.Config + chatIDs map[string]int64 + ctx context.Context + cancel context.CancelFunc registerFunc func(context.Context, []commands.Definition) error commandRegCancel context.CancelFunc @@ -95,12 +96,17 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann return &TelegramChannel{ BaseChannel: base, + commands: NewTelegramCommands(bot, cfg, nil), bot: bot, config: cfg, chatIDs: make(map[string]int64), }, nil } +func (c *TelegramChannel) SetRegistry(switcher AgentModelSwitcher) { + c.commands = NewTelegramCommands(c.bot, c.config, switcher) +} + func (c *TelegramChannel) Start(ctx context.Context) error { logger.InfoC("telegram", "Starting Telegram bot (polling mode)...") @@ -121,6 +127,25 @@ func (c *TelegramChannel) Start(ctx context.Context) error { } c.bh = bh + bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { + return c.commands.Start(ctx, message) + }, th.CommandEqual("start")) + bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { + return c.commands.Help(ctx, message) + }, th.CommandEqual("help")) + + bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { + return c.commands.Show(ctx, message) + }, th.CommandEqual("show")) + + bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { + return c.commands.List(ctx, message) + }, th.CommandEqual("list")) + + bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { + return c.commands.Model(ctx, message) + }, th.CommandEqual("model")) + bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { return c.handleMessage(ctx, &message) }, th.AnyMessage()) @@ -130,7 +155,11 @@ func (c *TelegramChannel) Start(ctx context.Context) error { "username": c.bot.Username(), }) - c.startCommandRegistration(c.ctx, commands.BuiltinDefinitions()) + commandDefs := append(commands.BuiltinDefinitions(), commands.Definition{ + Name: "model", + Description: "Show or switch the active model", + }) + c.startCommandRegistration(c.ctx, commandDefs) go func() { if err = bh.Start(); err != nil { diff --git a/pkg/channels/telegram/telegram_commands.go b/pkg/channels/telegram/telegram_commands.go new file mode 100644 index 000000000..e603b0117 --- /dev/null +++ b/pkg/channels/telegram/telegram_commands.go @@ -0,0 +1,234 @@ +package telegram + +import ( + "context" + "fmt" + "strings" + + "github.com/mymmrac/telego" + + "github.com/sipeed/picoclaw/pkg/config" + "github.com/sipeed/picoclaw/pkg/providers" +) + +// AgentModelSwitcher is the minimal interface needed to get/set the active model. +type AgentModelSwitcher interface { + GetDefaultAgentModel() string + SwitchDefaultAgentModel(modelName string) (string, string, error) +} + +type TelegramCommander interface { + Help(ctx context.Context, message telego.Message) error + Start(ctx context.Context, message telego.Message) error + Show(ctx context.Context, message telego.Message) error + List(ctx context.Context, message telego.Message) error + Model(ctx context.Context, message telego.Message) error +} + +type cmd struct { + bot *telego.Bot + config *config.Config + switcher AgentModelSwitcher +} + +func NewTelegramCommands(bot *telego.Bot, cfg *config.Config, switcher AgentModelSwitcher) TelegramCommander { + return &cmd{ + bot: bot, + config: cfg, + switcher: switcher, + } +} + +func commandArgs(text string) string { + parts := strings.SplitN(text, " ", 2) + if len(parts) < 2 { + return "" + } + return strings.TrimSpace(parts[1]) +} + +func (c *cmd) Help(ctx context.Context, message telego.Message) error { + msg := `/start - Start the bot +/help - Show this help message +/show [model|channel] - Show current configuration +/list [models|channels] - List available options +/model - Show or switch the active model + /model - show current model + /model list - list available models + /model - switch to named model + ` + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: msg, + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err +} + +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}, + Text: "Hello! I am PicoClaw 🦞", + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err +} + +func (c *cmd) Show(ctx context.Context, message telego.Message) error { + args := commandArgs(message.Text) + if args == "" { + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: "Usage: /show [model|channel]", + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err + } + + var response string + switch args { + case "model": + if c.switcher == nil { + response = fmt.Sprintf("Current model: %s", c.config.Agents.Defaults.GetModelName()) + } else { + response = fmt.Sprintf("Current model: %s", c.switcher.GetDefaultAgentModel()) + } + case "channel": + response = "Current Channel: telegram" + default: + response = fmt.Sprintf("Unknown parameter: %s. Try 'model' or 'channel'.", args) + } + + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: response, + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err +} + +func (c *cmd) List(ctx context.Context, message telego.Message) error { + args := commandArgs(message.Text) + if args == "" { + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: "Usage: /list [models|channels]", + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err + } + + var response string + switch args { + case "models": + provider := c.config.Agents.Defaults.Provider + if provider == "" { + provider = "configured default" + } + response = fmt.Sprintf("Configured Model: %s\nProvider: %s\n\nTo change models, update config.json", + c.config.Agents.Defaults.GetModelName(), provider) + + case "channels": + var enabled []string + if c.config.Channels.Telegram.Enabled { + enabled = append(enabled, "telegram") + } + if c.config.Channels.WhatsApp.Enabled { + enabled = append(enabled, "whatsapp") + } + if c.config.Channels.Feishu.Enabled { + enabled = append(enabled, "feishu") + } + if c.config.Channels.Discord.Enabled { + enabled = append(enabled, "discord") + } + if c.config.Channels.Slack.Enabled { + enabled = append(enabled, "slack") + } + response = fmt.Sprintf("Enabled Channels:\n- %s", strings.Join(enabled, "\n- ")) + + default: + response = fmt.Sprintf("Unknown parameter: %s. Try 'models' or 'channels'.", args) + } + + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: response, + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err +} + +func (c *cmd) Model(ctx context.Context, message telego.Message) error { + args := commandArgs(message.Text) + + var response string + switch { + case args == "": + // Show current model + if c.switcher == nil { + response = "No agent configured." + } else { + response = fmt.Sprintf("Current model: %s", c.switcher.GetDefaultAgentModel()) + } + + case args == "list": + // List models from config + if len(c.config.ModelList) == 0 { + response = "No models configured in model_list." + } else { + currentModel := "" + if c.switcher != nil { + currentModel = c.switcher.GetDefaultAgentModel() + } + lines := make([]string, 0, len(c.config.ModelList)) + for _, m := range c.config.ModelList { + _, modelID := providers.ExtractProtocol(m.Model) + line := "• " + m.ModelName + if modelID == currentModel { + line += " (active)" + } + line += " -> " + m.Model + lines = append(lines, line) + } + response = "Available models:\n" + strings.Join(lines, "\n") + } + + default: + // Switch to the named model + modelName := args + if c.switcher == nil { + response = "No agent configured." + } else { + oldModel, newModel, err := c.switcher.SwitchDefaultAgentModel(modelName) + if err != nil { + response = fmt.Sprintf("Failed to switch model: %v", err) + } else if oldModel == newModel { + response = fmt.Sprintf("Model already active: %s", newModel) + } else { + response = fmt.Sprintf("Switched model: %s -> %s", oldModel, newModel) + } + } + } + + _, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{ + ChatID: telego.ChatID{ID: message.Chat.ID}, + Text: response, + ReplyParameters: &telego.ReplyParameters{ + MessageID: message.MessageID, + }, + }) + return err +}