feat(agent): enhance agent model management and command handling
- Added model switching functionality for agents, allowing dynamic model changes at runtime. - Introduced a new command handler for Telegram and Discord channels to facilitate model switching. - Updated agent registry to manage default agent selection and model resolution. - Enhanced configuration management to support default model settings. - Improved candidate resolution logic for agent instances. - Added tests for default agent retrieval and fallback behavior.
This commit is contained in:
parent
a00ecedeb6
commit
09e3aacf88
16 changed files with 1011 additions and 54 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -47,3 +47,4 @@ dist/
|
||||||
|
|
||||||
# Windows Application Icon/Resource
|
# Windows Application Icon/Resource
|
||||||
*.syso
|
*.syso
|
||||||
|
AGENTS.md
|
||||||
|
|
|
||||||
3
go.mod
3
go.mod
|
|
@ -11,6 +11,7 @@ require (
|
||||||
github.com/gdamore/tcell/v2 v2.13.8
|
github.com/gdamore/tcell/v2 v2.13.8
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
|
github.com/h2non/filetype v1.1.3
|
||||||
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
||||||
github.com/mdp/qrterminal/v3 v3.2.1
|
github.com/mdp/qrterminal/v3 v3.2.1
|
||||||
github.com/modelcontextprotocol/go-sdk v1.3.0
|
github.com/modelcontextprotocol/go-sdk v1.3.0
|
||||||
|
|
@ -37,8 +38,6 @@ require (
|
||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect
|
github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect
|
||||||
github.com/gdamore/encoding v1.0.1 // indirect
|
github.com/gdamore/encoding v1.0.1 // indirect
|
||||||
github.com/gdamore/tcell/v2 v2.13.8 // indirect
|
|
||||||
github.com/h2non/filetype v1.1.3 // indirect
|
|
||||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||||
github.com/lucasb-eyer/go-colorful v1.3.0 // indirect
|
github.com/lucasb-eyer/go-colorful v1.3.0 // indirect
|
||||||
github.com/mattn/go-colorable v0.1.14 // indirect
|
github.com/mattn/go-colorable v0.1.14 // indirect
|
||||||
|
|
|
||||||
|
|
@ -113,7 +113,38 @@ func NewAgentInstance(
|
||||||
summarizeTokenPercent = 75
|
summarizeTokenPercent = 75
|
||||||
}
|
}
|
||||||
|
|
||||||
// Resolve fallback candidates
|
candidates := ResolveCandidatesForModel(cfg, defaults.Provider, model, fallbacks)
|
||||||
|
|
||||||
|
return &AgentInstance{
|
||||||
|
ID: agentID,
|
||||||
|
Name: agentName,
|
||||||
|
Model: model,
|
||||||
|
Fallbacks: fallbacks,
|
||||||
|
Workspace: workspace,
|
||||||
|
MaxIterations: maxIter,
|
||||||
|
MaxTokens: maxTokens,
|
||||||
|
Temperature: temperature,
|
||||||
|
ContextWindow: maxTokens,
|
||||||
|
SummarizeMessageThreshold: summarizeMessageThreshold,
|
||||||
|
SummarizeTokenPercent: summarizeTokenPercent,
|
||||||
|
Provider: provider,
|
||||||
|
Sessions: sessionsManager,
|
||||||
|
ContextBuilder: contextBuilder,
|
||||||
|
Tools: toolsRegistry,
|
||||||
|
Subagents: subagents,
|
||||||
|
SkillsFilter: skillsFilter,
|
||||||
|
Candidates: candidates,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveCandidatesForModel resolves fallback candidates using the same lookup
|
||||||
|
// behavior as AgentInstance construction. It is reused by model hot-switch paths.
|
||||||
|
func ResolveCandidatesForModel(
|
||||||
|
cfg *config.Config,
|
||||||
|
defaultProvider string,
|
||||||
|
model string,
|
||||||
|
fallbacks []string,
|
||||||
|
) []providers.FallbackCandidate {
|
||||||
modelCfg := providers.ModelConfig{
|
modelCfg := providers.ModelConfig{
|
||||||
Primary: model,
|
Primary: model,
|
||||||
Fallbacks: fallbacks,
|
Fallbacks: fallbacks,
|
||||||
|
|
@ -158,28 +189,7 @@ func NewAgentInstance(
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
candidates := providers.ResolveCandidatesWithLookup(modelCfg, defaults.Provider, resolveFromModelList)
|
return providers.ResolveCandidatesWithLookup(modelCfg, defaultProvider, resolveFromModelList)
|
||||||
|
|
||||||
return &AgentInstance{
|
|
||||||
ID: agentID,
|
|
||||||
Name: agentName,
|
|
||||||
Model: model,
|
|
||||||
Fallbacks: fallbacks,
|
|
||||||
Workspace: workspace,
|
|
||||||
MaxIterations: maxIter,
|
|
||||||
MaxTokens: maxTokens,
|
|
||||||
Temperature: temperature,
|
|
||||||
ContextWindow: maxTokens,
|
|
||||||
SummarizeMessageThreshold: summarizeMessageThreshold,
|
|
||||||
SummarizeTokenPercent: summarizeTokenPercent,
|
|
||||||
Provider: provider,
|
|
||||||
Sessions: sessionsManager,
|
|
||||||
ContextBuilder: contextBuilder,
|
|
||||||
Tools: toolsRegistry,
|
|
||||||
Subagents: subagents,
|
|
||||||
SkillsFilter: skillsFilter,
|
|
||||||
Candidates: candidates,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// resolveAgentWorkspace determines the workspace directory for an agent.
|
// resolveAgentWorkspace determines the workspace directory for an agent.
|
||||||
|
|
|
||||||
|
|
@ -37,6 +37,7 @@ type AgentLoop struct {
|
||||||
bus *bus.MessageBus
|
bus *bus.MessageBus
|
||||||
cfg *config.Config
|
cfg *config.Config
|
||||||
registry *AgentRegistry
|
registry *AgentRegistry
|
||||||
|
modelSwitch *ModelSwitchManager
|
||||||
state *state.Manager
|
state *state.Manager
|
||||||
running atomic.Bool
|
running atomic.Bool
|
||||||
summarizing sync.Map
|
summarizing sync.Map
|
||||||
|
|
@ -66,6 +67,7 @@ func NewAgentLoop(
|
||||||
provider providers.LLMProvider,
|
provider providers.LLMProvider,
|
||||||
) *AgentLoop {
|
) *AgentLoop {
|
||||||
registry := NewAgentRegistry(cfg, provider)
|
registry := NewAgentRegistry(cfg, provider)
|
||||||
|
modelSwitch := NewModelSwitchManager(cfg, registry)
|
||||||
|
|
||||||
// Register shared tools to all agents
|
// Register shared tools to all agents
|
||||||
registerSharedTools(cfg, msgBus, registry, provider)
|
registerSharedTools(cfg, msgBus, registry, provider)
|
||||||
|
|
@ -85,6 +87,7 @@ func NewAgentLoop(
|
||||||
bus: msgBus,
|
bus: msgBus,
|
||||||
cfg: cfg,
|
cfg: cfg,
|
||||||
registry: registry,
|
registry: registry,
|
||||||
|
modelSwitch: modelSwitch,
|
||||||
state: stateManager,
|
state: stateManager,
|
||||||
summarizing: sync.Map{},
|
summarizing: sync.Map{},
|
||||||
fallback: fallbackChain,
|
fallback: fallbackChain,
|
||||||
|
|
@ -1459,7 +1462,9 @@ func (al *AgentLoop) handleCommand(ctx context.Context, msg bus.InboundMessage)
|
||||||
return "No default agent configured", true
|
return "No default agent configured", true
|
||||||
}
|
}
|
||||||
oldModel := defaultAgent.Model
|
oldModel := defaultAgent.Model
|
||||||
defaultAgent.Model = value
|
if err := al.modelSwitch.SwitchModel("", value); err != nil {
|
||||||
|
return fmt.Sprintf("Failed to switch model: %v", err), true
|
||||||
|
}
|
||||||
return fmt.Sprintf("Switched model from %s to %s", oldModel, value), true
|
return fmt.Sprintf("Switched model from %s to %s", oldModel, value), true
|
||||||
case "channel":
|
case "channel":
|
||||||
if al.channelManager == nil {
|
if al.channelManager == nil {
|
||||||
|
|
|
||||||
110
pkg/agent/model_switch.go
Normal file
110
pkg/agent/model_switch.go
Normal file
|
|
@ -0,0 +1,110 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ModelSwitchManager handles dynamic model switching at runtime.
|
||||||
|
type ModelSwitchManager struct {
|
||||||
|
config *config.Config
|
||||||
|
registry *AgentRegistry
|
||||||
|
mu sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewModelSwitchManager creates a new model switch manager.
|
||||||
|
func NewModelSwitchManager(cfg *config.Config, registry *AgentRegistry) *ModelSwitchManager {
|
||||||
|
return &ModelSwitchManager{
|
||||||
|
config: cfg,
|
||||||
|
registry: registry,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SwitchModel switches the model for a given session.
|
||||||
|
// If sessionKey is empty, it updates the global default model.
|
||||||
|
func (m *ModelSwitchManager) SwitchModel(sessionKey, modelName string) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
// Validate model exists in config
|
||||||
|
if _, err := m.config.GetModelConfig(modelName); err != nil {
|
||||||
|
return fmt.Errorf("model %q not found. Available models: %v", modelName, m.listAvailableModels())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Session-scoped override is not implemented yet.
|
||||||
|
if sessionKey != "" {
|
||||||
|
return fmt.Errorf("session-scoped model switch is not supported yet")
|
||||||
|
}
|
||||||
|
|
||||||
|
oldModel := m.config.Agents.Defaults.GetModelName()
|
||||||
|
if oldModel == modelName {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := m.config.SetDefaultModel(modelName); err != nil {
|
||||||
|
return fmt.Errorf("failed to set default model: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
newProvider, _, err := providers.CreateProvider(m.config)
|
||||||
|
if err != nil {
|
||||||
|
// Roll back config on provider creation failure.
|
||||||
|
_ = m.config.SetDefaultModel(oldModel)
|
||||||
|
return fmt.Errorf("failed to create provider for model switch: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := m.registry.SwitchModel(m.config, oldModel, modelName, newProvider); err != nil {
|
||||||
|
_ = m.config.SetDefaultModel(oldModel)
|
||||||
|
if cp, ok := newProvider.(providers.StatefulProvider); ok {
|
||||||
|
cp.Close()
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to apply hot model switch: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCurrentModel returns the current model for a session.
|
||||||
|
// If session-specific override exists, returns that. Otherwise returns global default.
|
||||||
|
func (m *ModelSwitchManager) GetCurrentModel(sessionKey string) string {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
|
// For now, always return the global default model
|
||||||
|
// Session-scoped overrides can be added later
|
||||||
|
return m.config.Agents.Defaults.GetModelName()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateModel validates if a model exists in the configuration.
|
||||||
|
func (m *ModelSwitchManager) ValidateModel(modelName string) (*config.ModelConfig, error) {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
|
return m.config.GetModelConfig(modelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// listAvailableModels returns a slice of all available model names.
|
||||||
|
func (m *ModelSwitchManager) listAvailableModels() []string {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
|
models := make([]string, 0, len(m.config.ModelList))
|
||||||
|
for _, mc := range m.config.ModelList {
|
||||||
|
if mc.ModelName != "" {
|
||||||
|
models = append(models, mc.ModelName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return models
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelInfo returns formatted information about the current model.
|
||||||
|
func (m *ModelSwitchManager) GetModelInfo(sessionKey string) (string, string) {
|
||||||
|
currentModel := m.GetCurrentModel(sessionKey)
|
||||||
|
if mc, err := m.config.GetModelConfig(currentModel); err == nil {
|
||||||
|
return currentModel, mc.Model
|
||||||
|
}
|
||||||
|
|
||||||
|
return currentModel, "unknown"
|
||||||
|
}
|
||||||
120
pkg/agent/model_switch_test.go
Normal file
120
pkg/agent/model_switch_test.go
Normal file
|
|
@ -0,0 +1,120 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type statefulMockProvider struct {
|
||||||
|
closed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *statefulMockProvider) Chat(
|
||||||
|
ctx context.Context,
|
||||||
|
messages []providers.Message,
|
||||||
|
tools []providers.ToolDefinition,
|
||||||
|
model string,
|
||||||
|
options map[string]any,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
return &providers.LLMResponse{Content: "ok"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *statefulMockProvider) GetDefaultModel() string { return "mock" }
|
||||||
|
func (m *statefulMockProvider) Close() { m.closed = true }
|
||||||
|
|
||||||
|
func testSwitchConfig() *config.Config {
|
||||||
|
return &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: "/tmp/picoclaw-switch-test",
|
||||||
|
ModelName: "model-a",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
List: []config.AgentConfig{
|
||||||
|
{ID: "main", Default: true},
|
||||||
|
{ID: "worker", Model: &config.AgentModelConfig{Primary: "fixed-model"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{ModelName: "model-a", Model: "codex-cli/codex"},
|
||||||
|
{ModelName: "model-b", Model: "claude-cli/claude"},
|
||||||
|
{ModelName: "fixed-model", Model: "codex-cli/fixed"},
|
||||||
|
{ModelName: "broken-model", Model: "openai/gpt-bad"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelSwitchManager_SwitchModelSuccess(t *testing.T) {
|
||||||
|
cfg := testSwitchConfig()
|
||||||
|
initialProvider := &statefulMockProvider{}
|
||||||
|
registry := NewAgentRegistry(cfg, initialProvider)
|
||||||
|
manager := NewModelSwitchManager(cfg, registry)
|
||||||
|
|
||||||
|
if err := manager.SwitchModel("", "model-b"); err != nil {
|
||||||
|
t.Fatalf("SwitchModel() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := cfg.Agents.Defaults.GetModelName(); got != "model-b" {
|
||||||
|
t.Fatalf("default model = %q, want %q", got, "model-b")
|
||||||
|
}
|
||||||
|
|
||||||
|
mainAgent, ok := registry.GetAgent("main")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("main agent not found")
|
||||||
|
}
|
||||||
|
if mainAgent.Model != "model-b" {
|
||||||
|
t.Fatalf("main agent model = %q, want %q", mainAgent.Model, "model-b")
|
||||||
|
}
|
||||||
|
if len(mainAgent.Candidates) == 0 || mainAgent.Candidates[0].Provider != "claude-cli" {
|
||||||
|
t.Fatalf("main candidates not refreshed: %+v", mainAgent.Candidates)
|
||||||
|
}
|
||||||
|
|
||||||
|
worker, ok := registry.GetAgent("worker")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("worker agent not found")
|
||||||
|
}
|
||||||
|
if worker.Model != "fixed-model" {
|
||||||
|
t.Fatalf("worker model should remain fixed-model, got %q", worker.Model)
|
||||||
|
}
|
||||||
|
if worker.Provider == initialProvider {
|
||||||
|
t.Fatal("worker provider should be hot-swapped to new provider")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !initialProvider.closed {
|
||||||
|
t.Fatal("old stateful provider should be closed after switch")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelSwitchManager_SwitchModelRollbackOnProviderCreateFailure(t *testing.T) {
|
||||||
|
cfg := testSwitchConfig()
|
||||||
|
initialProvider := &statefulMockProvider{}
|
||||||
|
registry := NewAgentRegistry(cfg, initialProvider)
|
||||||
|
manager := NewModelSwitchManager(cfg, registry)
|
||||||
|
|
||||||
|
err := manager.SwitchModel("", "broken-model")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error for broken-model provider creation")
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := cfg.Agents.Defaults.GetModelName(); got != "model-a" {
|
||||||
|
t.Fatalf("default model should roll back to model-a, got %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
mainAgent, ok := registry.GetAgent("main")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("main agent not found")
|
||||||
|
}
|
||||||
|
if mainAgent.Model != "model-a" {
|
||||||
|
t.Fatalf("main model should remain model-a, got %q", mainAgent.Model)
|
||||||
|
}
|
||||||
|
if mainAgent.Provider != initialProvider {
|
||||||
|
t.Fatal("provider should remain unchanged on failure")
|
||||||
|
}
|
||||||
|
if initialProvider.closed {
|
||||||
|
t.Fatal("old provider should not be closed when switch fails")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package agent
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
|
@ -13,6 +15,7 @@ import (
|
||||||
type AgentRegistry struct {
|
type AgentRegistry struct {
|
||||||
agents map[string]*AgentInstance
|
agents map[string]*AgentInstance
|
||||||
resolver *routing.RouteResolver
|
resolver *routing.RouteResolver
|
||||||
|
defaultAgentID string
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -34,6 +37,7 @@ func NewAgentRegistry(
|
||||||
}
|
}
|
||||||
instance := NewAgentInstance(implicitAgent, &cfg.Agents.Defaults, cfg, provider)
|
instance := NewAgentInstance(implicitAgent, &cfg.Agents.Defaults, cfg, provider)
|
||||||
registry.agents["main"] = instance
|
registry.agents["main"] = instance
|
||||||
|
registry.defaultAgentID = "main"
|
||||||
logger.InfoCF("agent", "Created implicit main agent (no agents.list configured)", nil)
|
logger.InfoCF("agent", "Created implicit main agent (no agents.list configured)", nil)
|
||||||
} else {
|
} else {
|
||||||
for i := range agentConfigs {
|
for i := range agentConfigs {
|
||||||
|
|
@ -41,6 +45,9 @@ func NewAgentRegistry(
|
||||||
id := routing.NormalizeAgentID(ac.ID)
|
id := routing.NormalizeAgentID(ac.ID)
|
||||||
instance := NewAgentInstance(ac, &cfg.Agents.Defaults, cfg, provider)
|
instance := NewAgentInstance(ac, &cfg.Agents.Defaults, cfg, provider)
|
||||||
registry.agents[id] = instance
|
registry.agents[id] = instance
|
||||||
|
if ac.Default && registry.defaultAgentID == "" {
|
||||||
|
registry.defaultAgentID = id
|
||||||
|
}
|
||||||
logger.InfoCF("agent", "Registered agent",
|
logger.InfoCF("agent", "Registered agent",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"agent_id": id,
|
"agent_id": id,
|
||||||
|
|
@ -51,6 +58,17 @@ func NewAgentRegistry(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if registry.defaultAgentID == "" {
|
||||||
|
ids := make([]string, 0, len(registry.agents))
|
||||||
|
for id := range registry.agents {
|
||||||
|
ids = append(ids, id)
|
||||||
|
}
|
||||||
|
slices.Sort(ids)
|
||||||
|
if len(ids) > 0 {
|
||||||
|
registry.defaultAgentID = ids[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return registry
|
return registry
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -104,11 +122,61 @@ func (r *AgentRegistry) CanSpawnSubagent(parentAgentID, targetAgentID string) bo
|
||||||
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
defer r.mu.RUnlock()
|
defer r.mu.RUnlock()
|
||||||
|
if r.defaultAgentID != "" {
|
||||||
|
if agent, ok := r.agents[r.defaultAgentID]; ok {
|
||||||
|
return agent
|
||||||
|
}
|
||||||
|
}
|
||||||
if agent, ok := r.agents["main"]; ok {
|
if agent, ok := r.agents["main"]; ok {
|
||||||
return agent
|
return agent
|
||||||
}
|
}
|
||||||
for _, agent := range r.agents {
|
ids := make([]string, 0, len(r.agents))
|
||||||
return agent
|
for id := range r.agents {
|
||||||
|
ids = append(ids, id)
|
||||||
|
}
|
||||||
|
slices.Sort(ids)
|
||||||
|
if len(ids) > 0 {
|
||||||
|
return r.agents[ids[0]]
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SwitchModel hot-swaps provider and effective model at runtime.
|
||||||
|
// It updates all agents that currently use oldModel to newModel, and refreshes
|
||||||
|
// provider + fallback candidates atomically by replacing agent pointers.
|
||||||
|
func (r *AgentRegistry) SwitchModel(
|
||||||
|
cfg *config.Config,
|
||||||
|
oldModel string,
|
||||||
|
newModel string,
|
||||||
|
newProvider providers.LLMProvider,
|
||||||
|
) error {
|
||||||
|
r.mu.Lock()
|
||||||
|
if len(r.agents) == 0 {
|
||||||
|
r.mu.Unlock()
|
||||||
|
return fmt.Errorf("no agents registered")
|
||||||
|
}
|
||||||
|
|
||||||
|
var oldProvider providers.LLMProvider
|
||||||
|
for id, agent := range r.agents {
|
||||||
|
if oldProvider == nil {
|
||||||
|
oldProvider = agent.Provider
|
||||||
|
}
|
||||||
|
|
||||||
|
updated := *agent
|
||||||
|
if updated.Model == oldModel {
|
||||||
|
updated.Model = newModel
|
||||||
|
}
|
||||||
|
updated.Provider = newProvider
|
||||||
|
updated.Candidates = ResolveCandidatesForModel(cfg, cfg.Agents.Defaults.Provider, updated.Model, updated.Fallbacks)
|
||||||
|
r.agents[id] = &updated
|
||||||
|
}
|
||||||
|
r.mu.Unlock()
|
||||||
|
|
||||||
|
if oldProvider != nil && oldProvider != newProvider {
|
||||||
|
if cp, ok := oldProvider.(providers.StatefulProvider); ok {
|
||||||
|
cp.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -104,11 +104,29 @@ func TestAgentRegistry_GetDefaultAgent(t *testing.T) {
|
||||||
})
|
})
|
||||||
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
// GetDefaultAgent first checks for "main", then returns any
|
|
||||||
agent := registry.GetDefaultAgent()
|
agent := registry.GetDefaultAgent()
|
||||||
if agent == nil {
|
if agent == nil {
|
||||||
t.Fatal("expected a default agent")
|
t.Fatal("expected a default agent")
|
||||||
}
|
}
|
||||||
|
if agent.ID != "beta" {
|
||||||
|
t.Fatalf("default agent ID = %q, want %q", agent.ID, "beta")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentRegistry_GetDefaultAgent_FallbackDeterministic(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "zeta"},
|
||||||
|
{ID: "alpha"},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
agent := registry.GetDefaultAgent()
|
||||||
|
if agent == nil {
|
||||||
|
t.Fatal("expected a default agent")
|
||||||
|
}
|
||||||
|
if agent.ID != "alpha" {
|
||||||
|
t.Fatalf("default agent fallback should be deterministic sorted first, got %q", agent.ID)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAgentRegistry_CanSpawnSubagent(t *testing.T) {
|
func TestAgentRegistry_CanSpawnSubagent(t *testing.T) {
|
||||||
|
|
|
||||||
|
|
@ -111,6 +111,12 @@ func NewBaseChannel(
|
||||||
return bc
|
return bc
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Config returns the underlying config object passed when constructing
|
||||||
|
// the BaseChannel. Callers must perform an appropriate type assertion.
|
||||||
|
func (c *BaseChannel) Config() any {
|
||||||
|
return c.config
|
||||||
|
}
|
||||||
|
|
||||||
// MaxMessageLength returns the maximum message length (in runes) for this channel.
|
// MaxMessageLength returns the maximum message length (in runes) for this channel.
|
||||||
// A value of 0 means no limit.
|
// A value of 0 means no limit.
|
||||||
func (c *BaseChannel) MaxMessageLength() int {
|
func (c *BaseChannel) MaxMessageLength() int {
|
||||||
|
|
|
||||||
|
|
@ -29,35 +29,37 @@ const (
|
||||||
type DiscordChannel struct {
|
type DiscordChannel struct {
|
||||||
*channels.BaseChannel
|
*channels.BaseChannel
|
||||||
session *discordgo.Session
|
session *discordgo.Session
|
||||||
config config.DiscordConfig
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
typingMu sync.Mutex
|
typingMu sync.Mutex
|
||||||
typingStop map[string]chan struct{} // chatID → stop signal
|
typingStop map[string]chan struct{} // chatID → stop signal
|
||||||
botUserID string // stored for mention checking
|
botUserID string // stored for mention checking
|
||||||
|
commands DiscordCommander // Discord command handler
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
func NewDiscordChannel(cfg *config.Config, bus *bus.MessageBus) (*DiscordChannel, error) {
|
||||||
session, err := discordgo.New("Bot " + cfg.Token)
|
session, err := discordgo.New("Bot " + cfg.Channels.Discord.Token)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create discord session: %w", err)
|
return nil, fmt.Errorf("failed to create discord session: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := applyDiscordProxy(session, cfg.Proxy); err != nil {
|
if err := applyDiscordProxy(session, cfg.Channels.Discord.Proxy); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
base := channels.NewBaseChannel("discord", cfg, bus, cfg.AllowFrom,
|
base := channels.NewBaseChannel("discord", cfg, bus, cfg.Channels.Discord.AllowFrom,
|
||||||
channels.WithMaxMessageLength(2000),
|
channels.WithMaxMessageLength(2000),
|
||||||
channels.WithGroupTrigger(cfg.GroupTrigger),
|
channels.WithGroupTrigger(cfg.Channels.Discord.GroupTrigger),
|
||||||
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
channels.WithReasoningChannelID(cfg.Channels.Discord.ReasoningChannelID),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
commands := NewDiscordCommands(session, cfg, bus)
|
||||||
|
|
||||||
return &DiscordChannel{
|
return &DiscordChannel{
|
||||||
BaseChannel: base,
|
BaseChannel: base,
|
||||||
session: session,
|
session: session,
|
||||||
config: cfg,
|
|
||||||
ctx: context.Background(),
|
ctx: context.Background(),
|
||||||
typingStop: make(map[string]chan struct{}),
|
typingStop: make(map[string]chan struct{}),
|
||||||
|
commands: commands,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -235,11 +237,20 @@ func (c *DiscordChannel) EditMessage(ctx context.Context, chatID string, message
|
||||||
// It sends a placeholder message that will later be edited to the actual
|
// It sends a placeholder message that will later be edited to the actual
|
||||||
// response via EditMessage (channels.MessageEditor).
|
// response via EditMessage (channels.MessageEditor).
|
||||||
func (c *DiscordChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
|
func (c *DiscordChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
|
||||||
if !c.config.Placeholder.Enabled {
|
// Placeholder configuration is stored in the Discord channel config
|
||||||
|
// under the global config passed to BaseChannel.
|
||||||
|
cfgAny := c.BaseChannel.Config()
|
||||||
|
cfg, ok := cfgAny.(*config.Config)
|
||||||
|
if !ok {
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
text := c.config.Placeholder.Text
|
phCfg := cfg.Channels.Discord.Placeholder
|
||||||
|
if !phCfg.Enabled {
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
text := phCfg.Text
|
||||||
if text == "" {
|
if text == "" {
|
||||||
text = "Thinking... 💭"
|
text = "Thinking... 💭"
|
||||||
}
|
}
|
||||||
|
|
@ -314,6 +325,16 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
|
|
||||||
content := m.Content
|
content := m.Content
|
||||||
|
|
||||||
|
// Check for slash commands first
|
||||||
|
if strings.HasPrefix(strings.TrimSpace(content), "/") {
|
||||||
|
if err := c.handleSlashCommand(c.ctx, s, m); err != nil {
|
||||||
|
logger.DebugCF("discord", "Command error", map[string]any{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// In guild (group) channels, apply unified group trigger filtering
|
// In guild (group) channels, apply unified group trigger filtering
|
||||||
// DMs (GuildID is empty) always get a response
|
// DMs (GuildID is empty) always get a response
|
||||||
if m.GuildID != "" {
|
if m.GuildID != "" {
|
||||||
|
|
@ -519,3 +540,34 @@ func (c *DiscordChannel) stripBotMention(text string) string {
|
||||||
text = strings.ReplaceAll(text, fmt.Sprintf("<@!%s>", c.botUserID), "")
|
text = strings.ReplaceAll(text, fmt.Sprintf("<@!%s>", c.botUserID), "")
|
||||||
return strings.TrimSpace(text)
|
return strings.TrimSpace(text)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// handleSlashCommand processes slash commands and forwards to appropriate handler.
|
||||||
|
func (c *DiscordChannel) handleSlashCommand(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error {
|
||||||
|
content := strings.TrimSpace(m.Content)
|
||||||
|
if !strings.HasPrefix(content, "/") {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
parts := strings.Fields(content)
|
||||||
|
if len(parts) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd := strings.TrimPrefix(parts[0], "/")
|
||||||
|
|
||||||
|
switch cmd {
|
||||||
|
case "help":
|
||||||
|
return c.commands.Help(ctx, s, m)
|
||||||
|
case "start":
|
||||||
|
return c.commands.Help(ctx, s, m) // Start uses same help message
|
||||||
|
case "show":
|
||||||
|
return c.commands.Show(ctx, s, m)
|
||||||
|
case "list":
|
||||||
|
return c.commands.List(ctx, s, m)
|
||||||
|
case "switch":
|
||||||
|
return c.commands.Switch(ctx, s, m)
|
||||||
|
default:
|
||||||
|
// Unknown command - let it fall through to normal message handling
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
268
pkg/channels/discord/discord_commands.go
Normal file
268
pkg/channels/discord/discord_commands.go
Normal file
|
|
@ -0,0 +1,268 @@
|
||||||
|
package discord
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bwmarrin/discordgo"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/identity"
|
||||||
|
)
|
||||||
|
|
||||||
|
type DiscordCommander interface {
|
||||||
|
Help(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error
|
||||||
|
Show(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error
|
||||||
|
List(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error
|
||||||
|
Switch(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type cmd struct {
|
||||||
|
session *discordgo.Session
|
||||||
|
config *config.Config
|
||||||
|
bus *bus.MessageBus
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewDiscordCommands(session *discordgo.Session, cfg *config.Config, bus *bus.MessageBus) DiscordCommander {
|
||||||
|
return &cmd{
|
||||||
|
session: session,
|
||||||
|
config: cfg,
|
||||||
|
bus: bus,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseCommand extracts the command and arguments from message content.
|
||||||
|
func parseCommand(content string) (string, string) {
|
||||||
|
content = strings.TrimSpace(content)
|
||||||
|
if !strings.HasPrefix(content, "/") {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
parts := strings.SplitN(content, " ", 2)
|
||||||
|
cmd := strings.TrimPrefix(parts[0], "/")
|
||||||
|
|
||||||
|
if len(parts) < 2 {
|
||||||
|
return cmd, ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return cmd, strings.TrimSpace(parts[1])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) Help(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error {
|
||||||
|
msg := `**PicoClaw Commands**
|
||||||
|
|
||||||
|
/start - Start the bot
|
||||||
|
/help - Show this help message
|
||||||
|
/show [model|channel] - Show current configuration
|
||||||
|
/list [models|channels] - List available options
|
||||||
|
/switch model <name> - Switch to a different model
|
||||||
|
|
||||||
|
**Examples:**
|
||||||
|
/switch model gpt-4
|
||||||
|
/switch model claude-sonnet-4.6
|
||||||
|
|
||||||
|
Use /list models to see all available models.
|
||||||
|
`
|
||||||
|
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, msg)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) Show(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error {
|
||||||
|
cmd, args := parseCommand(m.Content)
|
||||||
|
if cmd != "show" {
|
||||||
|
return fmt.Errorf("invalid command format")
|
||||||
|
}
|
||||||
|
|
||||||
|
if args == "" {
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, "Usage: /show [model|channel]")
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
var response string
|
||||||
|
switch args {
|
||||||
|
case "model":
|
||||||
|
currentModel := c.config.Agents.Defaults.GetModelName()
|
||||||
|
provider := c.config.Agents.Defaults.Provider
|
||||||
|
response = fmt.Sprintf("**Current Model:** %s\n**Provider:** %s", currentModel, provider)
|
||||||
|
case "channel":
|
||||||
|
response = "**Current Channel:** discord"
|
||||||
|
default:
|
||||||
|
response = fmt.Sprintf("Unknown parameter: %s. Try 'model' or 'channel'.", args)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, response)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) List(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error {
|
||||||
|
cmd, args := parseCommand(m.Content)
|
||||||
|
if cmd != "list" {
|
||||||
|
return fmt.Errorf("invalid command format")
|
||||||
|
}
|
||||||
|
|
||||||
|
if args == "" {
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, "Usage: /list [models|channels]")
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
var response string
|
||||||
|
switch args {
|
||||||
|
case "models":
|
||||||
|
response = c.formatModelsList()
|
||||||
|
case "channels":
|
||||||
|
response = c.formatChannelsList()
|
||||||
|
default:
|
||||||
|
response = fmt.Sprintf("Unknown parameter: %s. Try 'models' or 'channels'.", args)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, response)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) Switch(ctx context.Context, s *discordgo.Session, m *discordgo.MessageCreate) error {
|
||||||
|
cmd, args := parseCommand(m.Content)
|
||||||
|
if cmd != "switch" {
|
||||||
|
return fmt.Errorf("invalid command format")
|
||||||
|
}
|
||||||
|
|
||||||
|
if args == "" {
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, "Usage: /switch model <name>\nUse /list models to see available models.")
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse "model <name>" format
|
||||||
|
parts := strings.SplitN(args, " ", 2)
|
||||||
|
if len(parts) < 2 || parts[0] != "model" {
|
||||||
|
_, err := s.ChannelMessageSend(m.ChannelID, "Usage: /switch model <name>\nUse /list models to see available models.")
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
modelName := strings.TrimSpace(parts[1])
|
||||||
|
|
||||||
|
// Optional: validate model exists to provide immediate feedback
|
||||||
|
if _, err := c.config.GetModelConfig(modelName); err != nil {
|
||||||
|
available := c.formatModelsList()
|
||||||
|
_, sendErr := c.session.ChannelMessageSend(
|
||||||
|
m.ChannelID,
|
||||||
|
fmt.Sprintf("❌ Model not found: %s\n\n**Available models:**\n%s", modelName, available),
|
||||||
|
)
|
||||||
|
return sendErr
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.bus == nil {
|
||||||
|
_, err := c.session.ChannelMessageSend(m.ChannelID, "❌ Internal error: message bus not initialized")
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Forward a normalized switch command to the agent loop via the message bus
|
||||||
|
// so that the agent can apply the change and persist any in-memory state.
|
||||||
|
inbound := bus.InboundMessage{
|
||||||
|
Channel: "discord",
|
||||||
|
SenderID: m.Author.ID,
|
||||||
|
Sender: bus.SenderInfo{
|
||||||
|
Platform: "discord",
|
||||||
|
PlatformID: m.Author.ID,
|
||||||
|
CanonicalID: identity.BuildCanonicalID("discord", m.Author.ID),
|
||||||
|
Username: m.Author.Username,
|
||||||
|
DisplayName: m.Author.Username,
|
||||||
|
},
|
||||||
|
ChatID: m.ChannelID,
|
||||||
|
Content: fmt.Sprintf("/switch model to %s", modelName),
|
||||||
|
MessageID: m.ID,
|
||||||
|
Metadata: map[string]string{
|
||||||
|
"guild_id": m.GuildID,
|
||||||
|
"channel_id": m.ChannelID,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := c.bus.PublishInbound(ctx, inbound); err != nil {
|
||||||
|
_, sendErr := c.session.ChannelMessageSend(m.ChannelID, fmt.Sprintf("❌ Failed to switch model: %v", err))
|
||||||
|
return sendErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// The agent will respond via the normal outbound flow; no immediate reply here.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) formatModelsList() string {
|
||||||
|
if len(c.config.ModelList) == 0 {
|
||||||
|
return "No models configured. Please check your configuration."
|
||||||
|
}
|
||||||
|
|
||||||
|
currentModel := c.config.Agents.Defaults.GetModelName()
|
||||||
|
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("**Available Models:**\n\n")
|
||||||
|
|
||||||
|
for _, mc := range c.config.ModelList {
|
||||||
|
if mc.ModelName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
prefix := " "
|
||||||
|
if mc.ModelName == currentModel {
|
||||||
|
prefix = "✓ "
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := "openai"
|
||||||
|
if strings.Contains(mc.Model, "/") {
|
||||||
|
protocolParts := strings.SplitN(mc.Model, "/", 2)
|
||||||
|
if len(protocolParts) > 0 {
|
||||||
|
provider = protocolParts[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sb.WriteString(fmt.Sprintf("%s**%s** - %s (%s)\n", prefix, mc.ModelName, mc.Model, provider))
|
||||||
|
}
|
||||||
|
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) formatChannelsList() string {
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
if c.config.Channels.LINE.Enabled {
|
||||||
|
enabled = append(enabled, "line")
|
||||||
|
}
|
||||||
|
if c.config.Channels.QQ.Enabled {
|
||||||
|
enabled = append(enabled, "qq")
|
||||||
|
}
|
||||||
|
if c.config.Channels.OneBot.Enabled {
|
||||||
|
enabled = append(enabled, "onebot")
|
||||||
|
}
|
||||||
|
if c.config.Channels.WeCom.Enabled {
|
||||||
|
enabled = append(enabled, "wecom")
|
||||||
|
}
|
||||||
|
if c.config.Channels.WeComApp.Enabled {
|
||||||
|
enabled = append(enabled, "wecom_app")
|
||||||
|
}
|
||||||
|
if c.config.Channels.WeComAIBot.Enabled {
|
||||||
|
enabled = append(enabled, "wecom_aibot")
|
||||||
|
}
|
||||||
|
if c.config.Channels.Pico.Enabled {
|
||||||
|
enabled = append(enabled, "pico")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(enabled) == 0 {
|
||||||
|
return "No channels enabled."
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf("Enabled channels:\n- %s", strings.Join(enabled, "\n- "))
|
||||||
|
}
|
||||||
|
|
@ -8,6 +8,6 @@ import (
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
channels.RegisterFactory("discord", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
channels.RegisterFactory("discord", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||||
return NewDiscordChannel(cfg.Channels.Discord, b)
|
return NewDiscordChannel(cfg, b)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -93,7 +93,7 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
|
||||||
|
|
||||||
return &TelegramChannel{
|
return &TelegramChannel{
|
||||||
BaseChannel: base,
|
BaseChannel: base,
|
||||||
commands: NewTelegramCommands(bot, cfg),
|
commands: NewTelegramCommands(bot, cfg, bus),
|
||||||
bot: bot,
|
bot: bot,
|
||||||
config: cfg,
|
config: cfg,
|
||||||
chatIDs: make(map[string]int64),
|
chatIDs: make(map[string]int64),
|
||||||
|
|
@ -141,6 +141,10 @@ func (c *TelegramChannel) Start(ctx context.Context) error {
|
||||||
return c.commands.List(ctx, message)
|
return c.commands.List(ctx, message)
|
||||||
}, th.CommandEqual("list"))
|
}, th.CommandEqual("list"))
|
||||||
|
|
||||||
|
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
|
||||||
|
return c.commands.Switch(ctx, message)
|
||||||
|
}, th.CommandEqual("switch"))
|
||||||
|
|
||||||
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
|
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
|
||||||
return c.handleMessage(ctx, &message)
|
return c.handleMessage(ctx, &message)
|
||||||
}, th.AnyMessage())
|
}, th.AnyMessage())
|
||||||
|
|
@ -203,6 +207,10 @@ func (c *TelegramChannel) initBotCommands(ctx context.Context) error {
|
||||||
Command: "list",
|
Command: "list",
|
||||||
Description: "List available options",
|
Description: "List available options",
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
Command: "switch",
|
||||||
|
Description: "Switch to a different model",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// Setting commands on each start will hit the rate limit very quickly, that's why we check if an update is needed
|
// Setting commands on each start will hit the rate limit very quickly, that's why we check if an update is needed
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,9 @@ import (
|
||||||
|
|
||||||
"github.com/mymmrac/telego"
|
"github.com/mymmrac/telego"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/identity"
|
||||||
)
|
)
|
||||||
|
|
||||||
type TelegramCommander interface {
|
type TelegramCommander interface {
|
||||||
|
|
@ -15,17 +17,20 @@ type TelegramCommander interface {
|
||||||
Start(ctx context.Context, message telego.Message) error
|
Start(ctx context.Context, message telego.Message) error
|
||||||
Show(ctx context.Context, message telego.Message) error
|
Show(ctx context.Context, message telego.Message) error
|
||||||
List(ctx context.Context, message telego.Message) error
|
List(ctx context.Context, message telego.Message) error
|
||||||
|
Switch(ctx context.Context, message telego.Message) error
|
||||||
}
|
}
|
||||||
|
|
||||||
type cmd struct {
|
type cmd struct {
|
||||||
bot *telego.Bot
|
bot *telego.Bot
|
||||||
config *config.Config
|
config *config.Config
|
||||||
|
bus *bus.MessageBus
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTelegramCommands(bot *telego.Bot, cfg *config.Config) TelegramCommander {
|
func NewTelegramCommands(bot *telego.Bot, cfg *config.Config, bus *bus.MessageBus) TelegramCommander {
|
||||||
return &cmd{
|
return &cmd{
|
||||||
bot: bot,
|
bot: bot,
|
||||||
config: cfg,
|
config: cfg,
|
||||||
|
bus: bus,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -37,11 +42,30 @@ func commandArgs(text string) string {
|
||||||
return strings.TrimSpace(parts[1])
|
return strings.TrimSpace(parts[1])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isTelegramSwitchAllowed(allowFrom config.FlexibleStringSlice, sender bus.SenderInfo) bool {
|
||||||
|
if len(allowFrom) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, allowedEntry := range allowFrom {
|
||||||
|
if identity.MatchAllowed(sender, allowedEntry) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
func (c *cmd) Help(ctx context.Context, message telego.Message) error {
|
func (c *cmd) Help(ctx context.Context, message telego.Message) error {
|
||||||
msg := `/start - Start the bot
|
msg := `/start - Start the bot
|
||||||
/help - Show this help message
|
/help - Show this help message
|
||||||
/show [model|channel] - Show current configuration
|
/show [model|channel] - Show current configuration
|
||||||
/list [models|channels] - List available options
|
/list [models|channels] - List available options
|
||||||
|
/switch model <name> - Switch to a different model
|
||||||
|
|
||||||
|
**Examples:**
|
||||||
|
/switch model gpt-4
|
||||||
|
/switch model claude-sonnet-4.6
|
||||||
|
|
||||||
|
Use /list models to see all available models.
|
||||||
`
|
`
|
||||||
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
ChatID: telego.ChatID{ID: message.Chat.ID},
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
|
@ -80,9 +104,10 @@ func (c *cmd) Show(ctx context.Context, message telego.Message) error {
|
||||||
var response string
|
var response string
|
||||||
switch args {
|
switch args {
|
||||||
case "model":
|
case "model":
|
||||||
|
currentModel := c.config.Agents.Defaults.GetModelName()
|
||||||
|
provider := c.config.Agents.Defaults.Provider
|
||||||
response = fmt.Sprintf("Current Model: %s (Provider: %s)",
|
response = fmt.Sprintf("Current Model: %s (Provider: %s)",
|
||||||
c.config.Agents.Defaults.GetModelName(),
|
currentModel, provider)
|
||||||
c.config.Agents.Defaults.Provider)
|
|
||||||
case "channel":
|
case "channel":
|
||||||
response = "Current Channel: telegram"
|
response = "Current Channel: telegram"
|
||||||
default:
|
default:
|
||||||
|
|
@ -115,12 +140,7 @@ func (c *cmd) List(ctx context.Context, message telego.Message) error {
|
||||||
var response string
|
var response string
|
||||||
switch args {
|
switch args {
|
||||||
case "models":
|
case "models":
|
||||||
provider := c.config.Agents.Defaults.Provider
|
response = c.formatModelsList()
|
||||||
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":
|
case "channels":
|
||||||
var enabled []string
|
var enabled []string
|
||||||
|
|
@ -154,3 +174,146 @@ func (c *cmd) List(ctx context.Context, message telego.Message) error {
|
||||||
})
|
})
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *cmd) Switch(ctx context.Context, message telego.Message) error {
|
||||||
|
if message.From == nil {
|
||||||
|
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: "❌ Cannot determine sender",
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
platformID := fmt.Sprintf("%d", message.From.ID)
|
||||||
|
sender := bus.SenderInfo{
|
||||||
|
Platform: "telegram",
|
||||||
|
PlatformID: platformID,
|
||||||
|
CanonicalID: identity.BuildCanonicalID("telegram", platformID),
|
||||||
|
Username: message.From.Username,
|
||||||
|
DisplayName: message.From.FirstName,
|
||||||
|
}
|
||||||
|
if !isTelegramSwitchAllowed(c.config.Channels.Telegram.AllowFrom, sender) {
|
||||||
|
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: "❌ You are not allowed to use this command.",
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
content := strings.TrimSpace(message.Text)
|
||||||
|
content = strings.TrimPrefix(content, "/switch")
|
||||||
|
content = strings.TrimSpace(content)
|
||||||
|
|
||||||
|
parts := strings.SplitN(content, " ", 2)
|
||||||
|
if len(parts) < 2 || parts[0] != "model" {
|
||||||
|
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: "Usage: /switch model <name>\nUse /list models to see available models.",
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
modelName := strings.TrimSpace(parts[1])
|
||||||
|
if modelName == "" {
|
||||||
|
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: "Usage: /switch model <name>\nUse /list models to see available models.",
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optional: validate model exists to provide immediate feedback
|
||||||
|
if _, err := c.config.GetModelConfig(modelName); err != nil {
|
||||||
|
_, sendErr := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: fmt.Sprintf("❌ Model not found: %s\n\n%s", modelName, c.formatModelsList()),
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return sendErr
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.bus == nil {
|
||||||
|
_, err := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: "❌ Internal error: message bus not initialized",
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
inbound := bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
SenderID: platformID,
|
||||||
|
Sender: sender,
|
||||||
|
ChatID: fmt.Sprintf("%d", message.Chat.ID),
|
||||||
|
Content: fmt.Sprintf("/switch model to %s", modelName),
|
||||||
|
MessageID: fmt.Sprintf("%d", message.MessageID),
|
||||||
|
Metadata: map[string]string{
|
||||||
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := c.bus.PublishInbound(ctx, inbound); err != nil {
|
||||||
|
_, sendErr := c.bot.SendMessage(ctx, &telego.SendMessageParams{
|
||||||
|
ChatID: telego.ChatID{ID: message.Chat.ID},
|
||||||
|
Text: fmt.Sprintf("❌ Failed to switch model: %v", err),
|
||||||
|
ReplyParameters: &telego.ReplyParameters{
|
||||||
|
MessageID: message.MessageID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return sendErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// The agent will respond via the normal outbound flow; no immediate reply here.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *cmd) formatModelsList() string {
|
||||||
|
if len(c.config.ModelList) == 0 {
|
||||||
|
return "No models configured. Please check your configuration."
|
||||||
|
}
|
||||||
|
|
||||||
|
currentModel := c.config.Agents.Defaults.GetModelName()
|
||||||
|
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("*Available Models:*\n\n")
|
||||||
|
|
||||||
|
for _, mc := range c.config.ModelList {
|
||||||
|
if mc.ModelName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
prefix := " "
|
||||||
|
if mc.ModelName == currentModel {
|
||||||
|
prefix = "✓ "
|
||||||
|
}
|
||||||
|
|
||||||
|
providerStr := "openai"
|
||||||
|
if strings.Contains(mc.Model, "/") {
|
||||||
|
protocolParts := strings.SplitN(mc.Model, "/", 2)
|
||||||
|
if len(protocolParts) > 0 {
|
||||||
|
providerStr = protocolParts[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sb.WriteString(fmt.Sprintf("%s%s - %s (%s)\n", prefix, mc.ModelName, mc.Model, providerStr))
|
||||||
|
}
|
||||||
|
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
|
||||||
115
pkg/channels/telegram/telegram_commands_test.go
Normal file
115
pkg/channels/telegram/telegram_commands_test.go
Normal file
|
|
@ -0,0 +1,115 @@
|
||||||
|
package telegram
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/mymmrac/telego"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/identity"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestIsTelegramSwitchAllowed(t *testing.T) {
|
||||||
|
sender := bus.SenderInfo{
|
||||||
|
Platform: "telegram",
|
||||||
|
PlatformID: "123",
|
||||||
|
CanonicalID: identity.BuildCanonicalID("telegram", "123"),
|
||||||
|
Username: "alice",
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
allowFrom config.FlexibleStringSlice
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "empty allowlist allows all",
|
||||||
|
allowFrom: config.FlexibleStringSlice{},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "matches raw platform id",
|
||||||
|
allowFrom: config.FlexibleStringSlice{"123"},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "matches canonical id",
|
||||||
|
allowFrom: config.FlexibleStringSlice{"telegram:123"},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "non matching denied",
|
||||||
|
allowFrom: config.FlexibleStringSlice{"999"},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := isTelegramSwitchAllowed(tt.allowFrom, sender)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Fatalf("isTelegramSwitchAllowed() = %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSwitchPublishesNormalizedCommand(t *testing.T) {
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
t.Cleanup(msgBus.Close)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
ModelName: "model-a",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Channels: config.ChannelsConfig{
|
||||||
|
Telegram: config.TelegramConfig{
|
||||||
|
AllowFrom: config.FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{ModelName: "model-a", Model: "codex-cli/codex"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
commander := &cmd{
|
||||||
|
config: cfg,
|
||||||
|
bus: msgBus,
|
||||||
|
// bot is intentionally nil: success path should not call SendMessage
|
||||||
|
}
|
||||||
|
|
||||||
|
message := telego.Message{
|
||||||
|
MessageID: 7,
|
||||||
|
Text: "/switch model model-a",
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: 42,
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 123,
|
||||||
|
Username: "alice",
|
||||||
|
FirstName: "Alice",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := commander.Switch(context.Background(), message); err != nil {
|
||||||
|
t.Fatalf("Switch() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
got, ok := msgBus.ConsumeInbound(ctx)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected inbound switch command")
|
||||||
|
}
|
||||||
|
if got.Content != "/switch model to model-a" {
|
||||||
|
t.Fatalf("inbound content = %q, want %q", got.Content, "/switch model to model-a")
|
||||||
|
}
|
||||||
|
if got.Channel != "telegram" || got.ChatID != "42" {
|
||||||
|
t.Fatalf("unexpected inbound routing: channel=%q chat=%q", got.Channel, got.ChatID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -832,3 +832,17 @@ func (c *Config) ValidateModelList() error {
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetDefaultModel sets the default model for all agents.
|
||||||
|
// It validates the model exists in ModelList before updating.
|
||||||
|
func (c *Config) SetDefaultModel(modelName string) error {
|
||||||
|
// Validate model exists
|
||||||
|
if _, err := c.GetModelConfig(modelName); err != nil {
|
||||||
|
return fmt.Errorf("model %q not found: %w", modelName, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update the default model name
|
||||||
|
c.Agents.Defaults.ModelName = modelName
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue