refactor(agent): delegate generic command execution to commands executor
This commit is contained in:
parent
33ba65fc91
commit
32638082cf
2 changed files with 204 additions and 140 deletions
|
|
@ -12,7 +12,6 @@ import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
@ -21,6 +20,7 @@ import (
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/commands"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/constants"
|
"github.com/sipeed/picoclaw/pkg/constants"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
|
@ -1266,144 +1266,55 @@ func (al *AgentLoop) handleCommand(
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
runtime := newAgentCommandRuntime(msg, route, agent, al.cfg)
|
||||||
|
executor := commands.NewExecutor(commands.NewRegistry(commands.BuiltinDefinitionsWithRuntime(al.cfg, runtime)))
|
||||||
|
|
||||||
|
var commandReply string
|
||||||
|
result := executor.Execute(ctx, commands.Request{
|
||||||
|
Channel: msg.Channel,
|
||||||
|
ChatID: msg.ChatID,
|
||||||
|
SenderID: msg.SenderID,
|
||||||
|
Text: msg.Content,
|
||||||
|
Reply: func(text string) error {
|
||||||
|
commandReply = text
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
switch result.Outcome {
|
||||||
|
case commands.OutcomeHandled:
|
||||||
|
if result.Err != nil {
|
||||||
|
return mapCommandError(result), true
|
||||||
|
}
|
||||||
|
if commandReply != "" {
|
||||||
|
return commandReply, true
|
||||||
|
}
|
||||||
|
if result.Reply != "" {
|
||||||
|
return result.Reply, true
|
||||||
|
}
|
||||||
|
return "", true
|
||||||
|
case commands.OutcomeRejected:
|
||||||
|
if result.Reply != "" {
|
||||||
|
return result.Reply, true
|
||||||
|
}
|
||||||
|
if result.Command != "" {
|
||||||
|
return fmt.Sprintf("Command /%s is not supported on %s.", result.Command, msg.Channel), true
|
||||||
|
}
|
||||||
|
return "Command is not supported on this channel.", true
|
||||||
|
case commands.OutcomePassthrough:
|
||||||
parts := strings.Fields(content)
|
parts := strings.Fields(content)
|
||||||
if len(parts) == 0 {
|
if len(parts) == 0 {
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd := parts[0]
|
cmd := parts[0]
|
||||||
if at := strings.Index(cmd, "@"); at > 0 {
|
if at := strings.Index(cmd, "@"); at > 0 {
|
||||||
cmd = cmd[:at]
|
cmd = cmd[:at]
|
||||||
}
|
}
|
||||||
|
if cmd != "/switch" {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
args := parts[1:]
|
args := parts[1:]
|
||||||
|
|
||||||
switch cmd {
|
|
||||||
case "/new", "/reset":
|
|
||||||
scopeKey := resolveScopeKey(route, msg.SessionKey)
|
|
||||||
newSessionKey, err := agent.Sessions.StartNew(scopeKey)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf("Failed to start new session: %v", err), true
|
|
||||||
}
|
|
||||||
|
|
||||||
backlogLimit := config.DefaultSessionBacklogLimit
|
|
||||||
if al.cfg != nil {
|
|
||||||
backlogLimit = al.cfg.Session.EffectiveBacklogLimit()
|
|
||||||
}
|
|
||||||
|
|
||||||
pruned, err := agent.Sessions.Prune(scopeKey, backlogLimit)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf(
|
|
||||||
"Started new session (%s), but pruning old sessions failed: %v",
|
|
||||||
newSessionKey,
|
|
||||||
err,
|
|
||||||
), true
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(pruned) == 0 {
|
|
||||||
return fmt.Sprintf("Started new session: %s", newSessionKey), true
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("Started new session: %s (pruned %d old session(s))", newSessionKey, len(pruned)), true
|
|
||||||
|
|
||||||
case "/session":
|
|
||||||
if len(args) < 1 {
|
|
||||||
return "Usage: /session [list|resume <index>]", true
|
|
||||||
}
|
|
||||||
|
|
||||||
scopeKey := resolveScopeKey(route, msg.SessionKey)
|
|
||||||
switch args[0] {
|
|
||||||
case "list":
|
|
||||||
list, err := agent.Sessions.List(scopeKey)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf("Failed to list sessions: %v", err), true
|
|
||||||
}
|
|
||||||
if len(list) == 0 {
|
|
||||||
return "No sessions found for current chat.", true
|
|
||||||
}
|
|
||||||
|
|
||||||
lines := make([]string, 0, len(list)+1)
|
|
||||||
lines = append(lines, "Sessions for current chat:")
|
|
||||||
for _, item := range list {
|
|
||||||
activeMarker := " "
|
|
||||||
if item.Active {
|
|
||||||
activeMarker = "*"
|
|
||||||
}
|
|
||||||
updated := "-"
|
|
||||||
if !item.UpdatedAt.IsZero() {
|
|
||||||
updated = item.UpdatedAt.Format("2006-01-02 15:04")
|
|
||||||
}
|
|
||||||
lines = append(lines, fmt.Sprintf(
|
|
||||||
"%d. [%s] %s (%d msgs, updated %s)",
|
|
||||||
item.Ordinal,
|
|
||||||
activeMarker,
|
|
||||||
item.SessionKey,
|
|
||||||
item.MessageCnt,
|
|
||||||
updated,
|
|
||||||
))
|
|
||||||
}
|
|
||||||
return strings.Join(lines, "\n"), true
|
|
||||||
|
|
||||||
case "resume":
|
|
||||||
if len(args) != 2 {
|
|
||||||
return "Usage: /session resume <index>", true
|
|
||||||
}
|
|
||||||
index, err := strconv.Atoi(args[1])
|
|
||||||
if err != nil || index < 1 {
|
|
||||||
return "Usage: /session resume <index>", true
|
|
||||||
}
|
|
||||||
sessionKey, err := agent.Sessions.Resume(scopeKey, index)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf("Failed to resume session %d: %v", index, err), true
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("Resumed session %d: %s", index, sessionKey), true
|
|
||||||
|
|
||||||
default:
|
|
||||||
return "Usage: /session [list|resume <index>]", true
|
|
||||||
}
|
|
||||||
|
|
||||||
case "/show":
|
|
||||||
if len(args) < 1 {
|
|
||||||
return "Usage: /show [model|channel|agents]", true
|
|
||||||
}
|
|
||||||
switch args[0] {
|
|
||||||
case "model":
|
|
||||||
defaultAgent := al.registry.GetDefaultAgent()
|
|
||||||
if defaultAgent == nil {
|
|
||||||
return "No default agent configured", true
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("Current model: %s", defaultAgent.Model), true
|
|
||||||
case "channel":
|
|
||||||
return fmt.Sprintf("Current channel: %s", msg.Channel), true
|
|
||||||
case "agents":
|
|
||||||
agentIDs := al.registry.ListAgentIDs()
|
|
||||||
return fmt.Sprintf("Registered agents: %s", strings.Join(agentIDs, ", ")), true
|
|
||||||
default:
|
|
||||||
return fmt.Sprintf("Unknown show target: %s", args[0]), true
|
|
||||||
}
|
|
||||||
|
|
||||||
case "/list":
|
|
||||||
if len(args) < 1 {
|
|
||||||
return "Usage: /list [models|channels|agents]", true
|
|
||||||
}
|
|
||||||
switch args[0] {
|
|
||||||
case "models":
|
|
||||||
return "Available models: configured in config.json per agent", true
|
|
||||||
case "channels":
|
|
||||||
if al.channelManager == nil {
|
|
||||||
return "Channel manager not initialized", true
|
|
||||||
}
|
|
||||||
channels := al.channelManager.GetEnabledChannels()
|
|
||||||
if len(channels) == 0 {
|
|
||||||
return "No channels enabled", true
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("Enabled channels: %s", strings.Join(channels, ", ")), true
|
|
||||||
case "agents":
|
|
||||||
agentIDs := al.registry.ListAgentIDs()
|
|
||||||
return fmt.Sprintf("Registered agents: %s", strings.Join(agentIDs, ", ")), true
|
|
||||||
default:
|
|
||||||
return fmt.Sprintf("Unknown list target: %s", args[0]), true
|
|
||||||
}
|
|
||||||
|
|
||||||
case "/switch":
|
|
||||||
if len(args) < 3 || args[1] != "to" {
|
if len(args) < 3 || args[1] != "to" {
|
||||||
return "Usage: /switch [model|channel] to <name>", true
|
return "Usage: /switch [model|channel] to <name>", true
|
||||||
}
|
}
|
||||||
|
|
@ -1435,6 +1346,50 @@ func (al *AgentLoop) handleCommand(
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type agentCommandRuntime struct {
|
||||||
|
channel string
|
||||||
|
scope string
|
||||||
|
sess commands.SessionOps
|
||||||
|
cfg *config.Config
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAgentCommandRuntime(
|
||||||
|
msg bus.InboundMessage,
|
||||||
|
route routing.ResolvedRoute,
|
||||||
|
agent *AgentInstance,
|
||||||
|
cfg *config.Config,
|
||||||
|
) commands.Runtime {
|
||||||
|
return agentCommandRuntime{
|
||||||
|
channel: msg.Channel,
|
||||||
|
scope: resolveScopeKey(route, msg.SessionKey),
|
||||||
|
sess: agent.Sessions,
|
||||||
|
cfg: cfg,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r agentCommandRuntime) Channel() string {
|
||||||
|
return r.channel
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r agentCommandRuntime) ScopeKey() string {
|
||||||
|
return r.scope
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r agentCommandRuntime) SessionOps() commands.SessionOps {
|
||||||
|
return r.sess
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r agentCommandRuntime) Config() *config.Config {
|
||||||
|
return r.cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func mapCommandError(result commands.ExecuteResult) string {
|
||||||
|
if result.Command == "" {
|
||||||
|
return fmt.Sprintf("Failed to execute command: %v", result.Err)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("Failed to execute /%s: %v", result.Command, result.Err)
|
||||||
|
}
|
||||||
|
|
||||||
// extractPeer extracts the routing peer from the inbound message's structured Peer field.
|
// extractPeer extracts the routing peer from the inbound message's structured Peer field.
|
||||||
func extractPeer(msg bus.InboundMessage) *routing.RoutePeer {
|
func extractPeer(msg bus.InboundMessage) *routing.RoutePeer {
|
||||||
if msg.Peer.Kind == "" {
|
if msg.Peer.Kind == "" {
|
||||||
|
|
|
||||||
|
|
@ -367,6 +367,29 @@ func (m *simpleMockProvider) GetDefaultModel() string {
|
||||||
return "mock-model"
|
return "mock-model"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type countingMockProvider struct {
|
||||||
|
response string
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *countingMockProvider) Chat(
|
||||||
|
ctx context.Context,
|
||||||
|
messages []providers.Message,
|
||||||
|
tools []providers.ToolDefinition,
|
||||||
|
model string,
|
||||||
|
opts map[string]any,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
m.calls++
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: m.response,
|
||||||
|
ToolCalls: []providers.ToolCall{},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *countingMockProvider) GetDefaultModel() string {
|
||||||
|
return "counting-mock-model"
|
||||||
|
}
|
||||||
|
|
||||||
// mockCustomTool is a simple mock tool for registration testing
|
// mockCustomTool is a simple mock tool for registration testing
|
||||||
type mockCustomTool struct{}
|
type mockCustomTool struct{}
|
||||||
|
|
||||||
|
|
@ -606,6 +629,92 @@ func TestHandleCommand_NewAndSessionCommands(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProcessMessage_CommandOutcomes(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Session: config.SessionConfig{
|
||||||
|
DMScope: "per-channel-peer",
|
||||||
|
BacklogLimit: 20,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &countingMockProvider{response: "LLM reply"}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
helper := testHelper{al: al}
|
||||||
|
|
||||||
|
baseMsg := bus.InboundMessage{
|
||||||
|
Channel: "whatsapp",
|
||||||
|
SenderID: "user1",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Peer: bus.Peer{
|
||||||
|
Kind: "direct",
|
||||||
|
ID: "user1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
showResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{
|
||||||
|
Channel: baseMsg.Channel,
|
||||||
|
SenderID: baseMsg.SenderID,
|
||||||
|
ChatID: baseMsg.ChatID,
|
||||||
|
Content: "/show channel",
|
||||||
|
Peer: baseMsg.Peer,
|
||||||
|
})
|
||||||
|
if showResp != "Command /show is not supported on whatsapp." {
|
||||||
|
t.Fatalf("unexpected /show reply: %q", showResp)
|
||||||
|
}
|
||||||
|
if provider.calls != 0 {
|
||||||
|
t.Fatalf("LLM should not be called for rejected command, calls=%d", provider.calls)
|
||||||
|
}
|
||||||
|
|
||||||
|
fooResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{
|
||||||
|
Channel: baseMsg.Channel,
|
||||||
|
SenderID: baseMsg.SenderID,
|
||||||
|
ChatID: baseMsg.ChatID,
|
||||||
|
Content: "/foo",
|
||||||
|
Peer: baseMsg.Peer,
|
||||||
|
})
|
||||||
|
if fooResp != "LLM reply" {
|
||||||
|
t.Fatalf("unexpected /foo reply: %q", fooResp)
|
||||||
|
}
|
||||||
|
if provider.calls != 1 {
|
||||||
|
t.Fatalf("LLM should be called exactly once after /foo passthrough, calls=%d", provider.calls)
|
||||||
|
}
|
||||||
|
|
||||||
|
route := al.registry.ResolveRoute(routing.RouteInput{
|
||||||
|
Channel: baseMsg.Channel,
|
||||||
|
Peer: extractPeer(baseMsg),
|
||||||
|
})
|
||||||
|
scopeKey := route.SessionKey
|
||||||
|
|
||||||
|
newResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{
|
||||||
|
Channel: baseMsg.Channel,
|
||||||
|
SenderID: baseMsg.SenderID,
|
||||||
|
ChatID: baseMsg.ChatID,
|
||||||
|
Content: "/new",
|
||||||
|
Peer: baseMsg.Peer,
|
||||||
|
})
|
||||||
|
if !strings.Contains(newResp, "Started new session: "+scopeKey+"#2") {
|
||||||
|
t.Fatalf("unexpected /new reply: %q", newResp)
|
||||||
|
}
|
||||||
|
if provider.calls != 1 {
|
||||||
|
t.Fatalf("LLM should not be called for handled /new command, calls=%d", provider.calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestToolResult_SilentToolDoesNotSendUserMessage verifies silent tools don't trigger outbound
|
// TestToolResult_SilentToolDoesNotSendUserMessage verifies silent tools don't trigger outbound
|
||||||
func TestToolResult_SilentToolDoesNotSendUserMessage(t *testing.T) {
|
func TestToolResult_SilentToolDoesNotSendUserMessage(t *testing.T) {
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue