refactor: extract DI to channels.Manager, resolve upstream merge and add chunking to SendMessageWithID
This commit is contained in:
parent
4768edc67b
commit
cafdb51df2
10 changed files with 729 additions and 112 deletions
|
|
@ -29,6 +29,7 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/media"
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/routing"
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/session"
|
||||||
"github.com/sipeed/picoclaw/pkg/skills"
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
"github.com/sipeed/picoclaw/pkg/state"
|
"github.com/sipeed/picoclaw/pkg/state"
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
|
|
@ -48,6 +49,7 @@ type AgentLoop struct {
|
||||||
mediaStore media.MediaStore
|
mediaStore media.MediaStore
|
||||||
transcriber voice.Transcriber
|
transcriber voice.Transcriber
|
||||||
cmdRegistry *commands.Registry
|
cmdRegistry *commands.Registry
|
||||||
|
taskManager *session.TaskManager
|
||||||
}
|
}
|
||||||
|
|
||||||
// processOptions configures how a message is processed
|
// processOptions configures how a message is processed
|
||||||
|
|
@ -80,9 +82,6 @@ func NewAgentLoop(
|
||||||
) *AgentLoop {
|
) *AgentLoop {
|
||||||
registry := NewAgentRegistry(cfg, provider)
|
registry := NewAgentRegistry(cfg, provider)
|
||||||
|
|
||||||
// Register shared tools to all agents
|
|
||||||
registerSharedTools(cfg, msgBus, registry, provider)
|
|
||||||
|
|
||||||
// Set up shared fallback chain
|
// Set up shared fallback chain
|
||||||
cooldown := providers.NewCooldownTracker()
|
cooldown := providers.NewCooldownTracker()
|
||||||
fallbackChain := providers.NewFallbackChain(cooldown)
|
fallbackChain := providers.NewFallbackChain(cooldown)
|
||||||
|
|
@ -102,8 +101,12 @@ func NewAgentLoop(
|
||||||
summarizing: sync.Map{},
|
summarizing: sync.Map{},
|
||||||
fallback: fallbackChain,
|
fallback: fallbackChain,
|
||||||
cmdRegistry: commands.NewRegistry(commands.BuiltinDefinitions()),
|
cmdRegistry: commands.NewRegistry(commands.BuiltinDefinitions()),
|
||||||
|
taskManager: session.NewTaskManager(filepath.Join(cfg.WorkspacePath(), "tasks")),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Register shared tools to all agents (TaskTool needs task manager from AgentLoop)
|
||||||
|
registerSharedTools(cfg, msgBus, registry, provider, al.taskManager)
|
||||||
|
|
||||||
return al
|
return al
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -113,6 +116,7 @@ func registerSharedTools(
|
||||||
msgBus *bus.MessageBus,
|
msgBus *bus.MessageBus,
|
||||||
registry *AgentRegistry,
|
registry *AgentRegistry,
|
||||||
provider providers.LLMProvider,
|
provider providers.LLMProvider,
|
||||||
|
taskManager *session.TaskManager,
|
||||||
) {
|
) {
|
||||||
for _, agentID := range registry.ListAgentIDs() {
|
for _, agentID := range registry.ListAgentIDs() {
|
||||||
agent, ok := registry.GetAgent(agentID)
|
agent, ok := registry.GetAgent(agentID)
|
||||||
|
|
@ -232,6 +236,12 @@ func registerSharedTools(
|
||||||
logger.WarnCF("agent", "spawn tool requires subagent to be enabled", nil)
|
logger.WarnCF("agent", "spawn tool requires subagent to be enabled", nil)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Task planning tool
|
||||||
|
if cfg.Tools.IsToolEnabled("tasktool") {
|
||||||
|
taskTool := tools.NewTaskTool(taskManager, cfg.Tools.TaskTool.Icons)
|
||||||
|
agent.Tools.Register(taskTool)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -390,6 +400,26 @@ func (al *AgentLoop) RegisterTool(tool tools.Tool) {
|
||||||
|
|
||||||
func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
||||||
al.channelManager = cm
|
al.channelManager = cm
|
||||||
|
al.bindAdvancedMessageManagers(cm)
|
||||||
|
}
|
||||||
|
|
||||||
|
// bindAdvancedMessageManagers wires up channel callbacks to any tools that
|
||||||
|
// require asynchronous, advanced message management (e.g., TaskTool)
|
||||||
|
func (al *AgentLoop) bindAdvancedMessageManagers(cm *channels.Manager) {
|
||||||
|
al.registry.ForEachToolInstance(func(t tools.Tool) {
|
||||||
|
if advancedManager, ok := t.(tools.AdvancedMessageManager); ok {
|
||||||
|
advancedManager.SetCallbacks(
|
||||||
|
// sendPlaceholder
|
||||||
|
func(channelName, chatID, content string) (string, error) {
|
||||||
|
return cm.SendMessageWithID(context.Background(), channelName, chatID, content)
|
||||||
|
},
|
||||||
|
// editMessage
|
||||||
|
func(channelName, chatID, messageID, content string) error {
|
||||||
|
return cm.EditMessage(context.Background(), channelName, chatID, messageID, content)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetMediaStore injects a MediaStore for media lifecycle management.
|
// SetMediaStore injects a MediaStore for media lifecycle management.
|
||||||
|
|
|
||||||
|
|
@ -114,6 +114,20 @@ func (r *AgentRegistry) ForEachTool(name string, fn func(tools.Tool)) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ForEachToolInstance calls fn for every tool registered across all agents.
|
||||||
|
// This is useful for propagating dependencies matching a specific interface.
|
||||||
|
func (r *AgentRegistry) ForEachToolInstance(fn func(tools.Tool)) {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
for _, agent := range r.agents {
|
||||||
|
for _, name := range agent.Tools.List() {
|
||||||
|
if t, ok := agent.Tools.Get(name); ok {
|
||||||
|
fn(t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// GetDefaultAgent returns the default agent instance.
|
// GetDefaultAgent returns the default agent instance.
|
||||||
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
|
|
|
||||||
|
|
@ -50,3 +50,10 @@ type PlaceholderRecorder interface {
|
||||||
type CommandRegistrarCapable interface {
|
type CommandRegistrarCapable interface {
|
||||||
RegisterCommands(ctx context.Context, defs []commands.Definition) error
|
RegisterCommands(ctx context.Context, defs []commands.Definition) error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SyncSender — channels that can bypass the async bus to send a message synchronously.
|
||||||
|
// This is typically used by internal tools (like TaskTool) that must immediately
|
||||||
|
// receive the generated message ID in order to edit it later.
|
||||||
|
type SyncSender interface {
|
||||||
|
SendMessageWithID(ctx context.Context, chatID, content string) (string, error)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -833,3 +833,45 @@ func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, conten
|
||||||
channel, _ := m.channels[channelName]
|
channel, _ := m.channels[channelName]
|
||||||
return channel.Send(ctx, msg)
|
return channel.Send(ctx, msg)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SendMessageWithID sends a message synchronously via the channel's native API if supported,
|
||||||
|
// returning the platform-specific message ID. If the channel does not support SyncSender,
|
||||||
|
// it falls back to the async bus and returns an error.
|
||||||
|
func (m *Manager) SendMessageWithID(ctx context.Context, channelName, chatID, content string) (string, error) {
|
||||||
|
ch, ok := m.GetChannel(channelName)
|
||||||
|
if !ok {
|
||||||
|
return "", fmt.Errorf("channel %s not found", channelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
if syncSender, ok := ch.(SyncSender); ok {
|
||||||
|
msgID, err := syncSender.SendMessageWithID(ctx, chatID, content)
|
||||||
|
if err == nil && msgID != "" {
|
||||||
|
return msgID, nil
|
||||||
|
}
|
||||||
|
logger.ErrorCF("manager", "SendMessageWithID failed", map[string]any{"error": err, "msgID": msgID})
|
||||||
|
} else {
|
||||||
|
logger.WarnCF("manager", "channel does not implement SyncSender", map[string]any{"channel": channelName})
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.WarnCF("manager", "falling back to bus publish", nil)
|
||||||
|
m.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||||
|
Channel: channelName,
|
||||||
|
ChatID: chatID,
|
||||||
|
Content: content,
|
||||||
|
})
|
||||||
|
|
||||||
|
return "", fmt.Errorf("channel does not support returning message ID")
|
||||||
|
}
|
||||||
|
|
||||||
|
// EditMessage synchronously edits an existing message if the channel supports MessageEditor.
|
||||||
|
func (m *Manager) EditMessage(ctx context.Context, channelName, chatID, messageID, content string) error {
|
||||||
|
ch, ok := m.GetChannel(channelName)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("channel %s not found", channelName)
|
||||||
|
}
|
||||||
|
editor, ok := ch.(MessageEditor)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("channel %s does not support message editing", channelName)
|
||||||
|
}
|
||||||
|
return editor.EditMessage(ctx, chatID, messageID, content)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -164,23 +164,31 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
_, err := c.SendMessageWithID(ctx, msg.ChatID, msg.Content)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendMessageWithID implements an optional interface for AgentLoop to send a message synchronously and get the MessageID.
|
||||||
|
func (c *TelegramChannel) SendMessageWithID(ctx context.Context, chatID string, content string) (string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return "", channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, err := parseChatID(msg.ChatID)
|
cid, err := parseChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return "", fmt.Errorf("invalid chat ID %s: %w", chatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
if msg.Content == "" {
|
if content == "" {
|
||||||
return nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
||||||
// so msg.Content is guaranteed to be within that limit. We still need to
|
// so msg.Content is guaranteed to be within that limit. We still need to
|
||||||
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
||||||
queue := []string{msg.Content}
|
queue := []string{content}
|
||||||
|
var lastMsgID int
|
||||||
|
|
||||||
for len(queue) > 0 {
|
for len(queue) > 0 {
|
||||||
chunk := queue[0]
|
chunk := queue[0]
|
||||||
queue = queue[1:]
|
queue = queue[1:]
|
||||||
|
|
@ -200,31 +208,38 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendHTMLChunk(ctx, chatID, htmlContent, chunk); err != nil {
|
msgID, err := c.sendHTMLChunk(ctx, cid, htmlContent, chunk)
|
||||||
return err
|
if err != nil {
|
||||||
|
return "", err
|
||||||
}
|
}
|
||||||
|
lastMsgID = msgID
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
if lastMsgID == 0 {
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%d", lastMsgID), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendHTMLChunk sends a single HTML message, falling back to the original
|
// sendHTMLChunk sends a single HTML message, falling back to the original
|
||||||
// markdown as plain text on parse failure so users never see raw HTML tags.
|
// markdown as plain text on parse failure so users never see raw HTML tags.
|
||||||
func (c *TelegramChannel) sendHTMLChunk(ctx context.Context, chatID int64, htmlContent, mdFallback string) error {
|
func (c *TelegramChannel) sendHTMLChunk(ctx context.Context, chatID int64, htmlContent, mdFallback string) (int, error) {
|
||||||
tgMsg := tu.Message(tu.ID(chatID), htmlContent)
|
tgMsg := tu.Message(tu.ID(chatID), htmlContent)
|
||||||
tgMsg.ParseMode = telego.ModeHTML
|
tgMsg.ParseMode = telego.ModeHTML
|
||||||
|
|
||||||
if _, err := c.bot.SendMessage(ctx, tgMsg); err != nil {
|
msg, err := c.bot.SendMessage(ctx, tgMsg)
|
||||||
|
if err != nil {
|
||||||
logger.ErrorCF("telegram", "HTML parse failed, falling back to plain text", map[string]any{
|
logger.ErrorCF("telegram", "HTML parse failed, falling back to plain text", map[string]any{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
tgMsg.Text = mdFallback
|
tgMsg.Text = mdFallback
|
||||||
tgMsg.ParseMode = ""
|
tgMsg.ParseMode = ""
|
||||||
if _, err = c.bot.SendMessage(ctx, tgMsg); err != nil {
|
msg, err = c.bot.SendMessage(ctx, tgMsg)
|
||||||
return fmt.Errorf("telegram send: %w", channels.ErrTemporary)
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("telegram send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return msg.MessageID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartTyping implements channels.TypingCapable.
|
// StartTyping implements channels.TypingCapable.
|
||||||
|
|
|
||||||
|
|
@ -1,19 +1,19 @@
|
||||||
package telegram
|
package telegram
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/mymmrac/telego"
|
"github.com/mymmrac/telego"
|
||||||
ta "github.com/mymmrac/telego/telegoapi"
|
ta "github.com/mymmrac/telego/telegoapi"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
)
|
)
|
||||||
|
|
||||||
const testToken = "1234567890:aaaabbbbaaaabbbbaaaabbbbaaaabbbbccc"
|
const testToken = "1234567890:aaaabbbbaaaabbbbaaaabbbbaaaabbbbccc"
|
||||||
|
|
@ -42,8 +42,8 @@ func (s *stubConstructor) JSONRequest(parameters any) (*ta.RequestData, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *stubConstructor) MultipartRequest(
|
func (s *stubConstructor) MultipartRequest(
|
||||||
parameters map[string]string,
|
parameters map[string]string,
|
||||||
files map[string]ta.NamedReader,
|
files map[string]ta.NamedReader,
|
||||||
) (*ta.RequestData, error) {
|
) (*ta.RequestData, error) {
|
||||||
return &ta.RequestData{}, nil
|
return &ta.RequestData{}, nil
|
||||||
}
|
}
|
||||||
|
|
@ -62,15 +62,15 @@ func newTestChannel(t *testing.T, caller *stubCaller) *TelegramChannel {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
bot, err := telego.NewBot(testToken,
|
bot, err := telego.NewBot(testToken,
|
||||||
telego.WithAPICaller(caller),
|
telego.WithAPICaller(caller),
|
||||||
telego.WithRequestConstructor(&stubConstructor{}),
|
telego.WithRequestConstructor(&stubConstructor{}),
|
||||||
telego.WithDiscardLogger(),
|
telego.WithDiscardLogger(),
|
||||||
)
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
base := channels.NewBaseChannel("telegram", nil, nil, nil,
|
base := channels.NewBaseChannel("telegram", nil, nil, nil,
|
||||||
channels.WithMaxMessageLength(4000),
|
channels.WithMaxMessageLength(4000),
|
||||||
)
|
)
|
||||||
base.SetRunning(true)
|
base.SetRunning(true)
|
||||||
|
|
||||||
return &TelegramChannel{
|
return &TelegramChannel{
|
||||||
|
|
@ -80,25 +80,7 @@ func newTestChannel(t *testing.T, caller *stubCaller) *TelegramChannel {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_EmptyContent(t *testing.T) {
|
func TestSend_Wrapper(t *testing.T) {
|
||||||
caller := &stubCaller{
|
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
|
||||||
t.Fatal("SendMessage should not be called for empty content")
|
|
||||||
return nil, nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
ch := newTestChannel(t, caller)
|
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
|
||||||
ChatID: "12345",
|
|
||||||
Content: "",
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.NoError(t, err)
|
|
||||||
assert.Empty(t, caller.calls, "no API calls should be made for empty content")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSend_ShortMessage_SingleCall(t *testing.T) {
|
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
return successResponse(t), nil
|
return successResponse(t), nil
|
||||||
|
|
@ -112,14 +94,41 @@ func TestSend_ShortMessage_SingleCall(t *testing.T) {
|
||||||
})
|
})
|
||||||
|
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, caller.calls, 1, "wrapper should call inner function")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessageWithID_EmptyContent(t *testing.T) {
|
||||||
|
caller := &stubCaller{
|
||||||
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
|
t.Fatal("SendMessage should not be called for empty content")
|
||||||
|
return nil, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "")
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Empty(t, msgID)
|
||||||
|
assert.Empty(t, caller.calls, "no API calls should be made for empty content")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessageWithID_ShortMessage_SingleCall(t *testing.T) {
|
||||||
|
caller := &stubCaller{
|
||||||
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
|
return successResponse(t), nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello, world!")
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, "1", msgID)
|
||||||
assert.Len(t, caller.calls, 1, "short message should result in exactly one SendMessage call")
|
assert.Len(t, caller.calls, 1, "short message should result in exactly one SendMessage call")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_LongMessage_SingleCall(t *testing.T) {
|
func TestSendMessageWithID_LongMessage_SingleCall(t *testing.T) {
|
||||||
// With WithMaxMessageLength(4000), the Manager pre-splits messages before
|
|
||||||
// they reach Send(). A message at exactly 4000 chars should go through
|
|
||||||
// as a single SendMessage call (no re-split needed since HTML expansion
|
|
||||||
// won't exceed 4096 for plain text).
|
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
return successResponse(t), nil
|
return successResponse(t), nil
|
||||||
|
|
@ -129,21 +138,18 @@ func TestSend_LongMessage_SingleCall(t *testing.T) {
|
||||||
|
|
||||||
longContent := strings.Repeat("a", 4000)
|
longContent := strings.Repeat("a", 4000)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", longContent)
|
||||||
ChatID: "12345",
|
|
||||||
Content: longContent,
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, "1", msgID)
|
||||||
assert.Len(t, caller.calls, 1, "pre-split message within limit should result in one SendMessage call")
|
assert.Len(t, caller.calls, 1, "pre-split message within limit should result in one SendMessage call")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_HTMLFallback_PerChunk(t *testing.T) {
|
func TestSendMessageWithID_HTMLFallback_PerChunk(t *testing.T) {
|
||||||
callCount := 0
|
callCount := 0
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
callCount++
|
callCount++
|
||||||
// Fail on odd calls (HTML attempt), succeed on even calls (plain text fallback)
|
|
||||||
if callCount%2 == 1 {
|
if callCount%2 == 1 {
|
||||||
return nil, errors.New("Bad Request: can't parse entities")
|
return nil, errors.New("Bad Request: can't parse entities")
|
||||||
}
|
}
|
||||||
|
|
@ -152,17 +158,14 @@ func TestSend_HTMLFallback_PerChunk(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello **world**")
|
||||||
ChatID: "12345",
|
|
||||||
Content: "Hello **world**",
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
// One short message → 1 HTML attempt (fail) + 1 plain text fallback (success) = 2 calls
|
assert.Equal(t, "1", msgID)
|
||||||
assert.Equal(t, 2, len(caller.calls), "should have HTML attempt + plain text fallback")
|
assert.Equal(t, 2, len(caller.calls), "should have HTML attempt + plain text fallback")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_HTMLFallback_BothFail(t *testing.T) {
|
func TestSendMessageWithID_HTMLFallback_BothFail(t *testing.T) {
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
return nil, errors.New("send failed")
|
return nil, errors.New("send failed")
|
||||||
|
|
@ -170,19 +173,15 @@ func TestSend_HTMLFallback_BothFail(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello")
|
||||||
ChatID: "12345",
|
|
||||||
Content: "Hello",
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
assert.Error(t, err)
|
||||||
|
assert.Empty(t, msgID)
|
||||||
assert.True(t, errors.Is(err, channels.ErrTemporary), "error should wrap ErrTemporary")
|
assert.True(t, errors.Is(err, channels.ErrTemporary), "error should wrap ErrTemporary")
|
||||||
assert.Equal(t, 2, len(caller.calls), "should have HTML attempt + plain text attempt")
|
assert.Equal(t, 2, len(caller.calls), "should have HTML attempt + plain text attempt")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_LongMessage_HTMLFallback_StopsOnError(t *testing.T) {
|
func TestSendMessageWithID_LongMessage_HTMLFallback_StopsOnError(t *testing.T) {
|
||||||
// With a long message that gets split into 2 chunks, if both HTML and
|
|
||||||
// plain text fail on the first chunk, Send should return early.
|
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
return nil, errors.New("send failed")
|
return nil, errors.New("send failed")
|
||||||
|
|
@ -192,17 +191,14 @@ func TestSend_LongMessage_HTMLFallback_StopsOnError(t *testing.T) {
|
||||||
|
|
||||||
longContent := strings.Repeat("x", 4001)
|
longContent := strings.Repeat("x", 4001)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", longContent)
|
||||||
ChatID: "12345",
|
|
||||||
Content: longContent,
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
assert.Error(t, err)
|
||||||
// Should fail on the first chunk (2 calls: HTML + fallback), never reaching the second chunk.
|
assert.Empty(t, msgID)
|
||||||
assert.Equal(t, 2, len(caller.calls), "should stop after first chunk fails both HTML and plain text")
|
assert.Equal(t, 2, len(caller.calls), "should stop after first chunk fails both HTML and plain text")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
|
func TestSendMessageWithID_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
return successResponse(t), nil
|
return successResponse(t), nil
|
||||||
|
|
@ -210,31 +206,17 @@ func TestSend_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
// Create markdown whose length is <= 4000 but whose HTML expansion is much longer.
|
markdownContent := strings.Repeat("**a** ", 600)
|
||||||
// "**a** " (6 chars) becomes "<b>a</b> " (9 chars) in HTML, so repeating it many times
|
assert.LessOrEqual(t, len([]rune(markdownContent)), 4000)
|
||||||
// yields HTML that exceeds Telegram's limit while markdown stays within it.
|
|
||||||
markdownContent := strings.Repeat("**a** ", 600) // 3600 chars markdown, HTML ~5400+ chars
|
|
||||||
assert.LessOrEqual(t, len([]rune(markdownContent)), 4000, "markdown content must not exceed chunk size")
|
|
||||||
|
|
||||||
htmlExpanded := markdownToTelegramHTML(markdownContent)
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", markdownContent)
|
||||||
assert.Greater(
|
|
||||||
t, len([]rune(htmlExpanded)), 4096,
|
|
||||||
"HTML expansion must exceed Telegram limit for this test to be meaningful",
|
|
||||||
)
|
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
|
||||||
ChatID: "12345",
|
|
||||||
Content: markdownContent,
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.Greater(
|
assert.Equal(t, "1", msgID)
|
||||||
t, len(caller.calls), 1,
|
assert.Greater(t, len(caller.calls), 1, "markdown-short but HTML-long message should be split into multiple SendMessage calls")
|
||||||
"markdown-short but HTML-long message should be split into multiple SendMessage calls",
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_NotRunning(t *testing.T) {
|
func TestSendMessageWithID_NotRunning(t *testing.T) {
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
t.Fatal("should not be called")
|
t.Fatal("should not be called")
|
||||||
|
|
@ -244,16 +226,14 @@ func TestSend_NotRunning(t *testing.T) {
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
ch.SetRunning(false)
|
ch.SetRunning(false)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello")
|
||||||
ChatID: "12345",
|
|
||||||
Content: "Hello",
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.ErrorIs(t, err, channels.ErrNotRunning)
|
assert.ErrorIs(t, err, channels.ErrNotRunning)
|
||||||
|
assert.Empty(t, msgID)
|
||||||
assert.Empty(t, caller.calls)
|
assert.Empty(t, caller.calls)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSend_InvalidChatID(t *testing.T) {
|
func TestSendMessageWithID_InvalidChatID(t *testing.T) {
|
||||||
caller := &stubCaller{
|
caller := &stubCaller{
|
||||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
t.Fatal("should not be called")
|
t.Fatal("should not be called")
|
||||||
|
|
@ -262,12 +242,10 @@ func TestSend_InvalidChatID(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
msgID, err := ch.SendMessageWithID(context.Background(), "not-a-number", "Hello")
|
||||||
ChatID: "not-a-number",
|
|
||||||
Content: "Hello",
|
|
||||||
})
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
assert.Error(t, err)
|
||||||
|
assert.Empty(t, msgID)
|
||||||
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
|
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
|
||||||
assert.Empty(t, caller.calls)
|
assert.Empty(t, caller.calls)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -666,6 +666,19 @@ type ToolsConfig struct {
|
||||||
Subagent ToolConfig `json:"subagent" envPrefix:"PICOCLAW_TOOLS_SUBAGENT_"`
|
Subagent ToolConfig `json:"subagent" envPrefix:"PICOCLAW_TOOLS_SUBAGENT_"`
|
||||||
WebFetch ToolConfig `json:"web_fetch" envPrefix:"PICOCLAW_TOOLS_WEB_FETCH_"`
|
WebFetch ToolConfig `json:"web_fetch" envPrefix:"PICOCLAW_TOOLS_WEB_FETCH_"`
|
||||||
WriteFile ToolConfig `json:"write_file" envPrefix:"PICOCLAW_TOOLS_WRITE_FILE_"`
|
WriteFile ToolConfig `json:"write_file" envPrefix:"PICOCLAW_TOOLS_WRITE_FILE_"`
|
||||||
|
TaskTool TaskToolConfig `json:"tasktool"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type TaskToolConfig struct {
|
||||||
|
ToolConfig `envPrefix:"PICOCLAW_TOOLS_TASK_TOOL_"`
|
||||||
|
Icons TaskToolIconsConfig `json:"icons"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type TaskToolIconsConfig struct {
|
||||||
|
Pending string `json:"pending" env:"PICOCLAW_TOOLS_TASK_TOOL_ICONS_PENDING" default:"⚪"`
|
||||||
|
InProgress string `json:"in_progress" env:"PICOCLAW_TOOLS_TASK_TOOL_ICONS_IN_PROGRESS" default:"🟡"`
|
||||||
|
Completed string `json:"completed" env:"PICOCLAW_TOOLS_TASK_TOOL_ICONS_COMPLETED" default:"🟢"`
|
||||||
|
Failed string `json:"failed" env:"PICOCLAW_TOOLS_TASK_TOOL_ICONS_FAILED" default:"🔴"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type SearchCacheConfig struct {
|
type SearchCacheConfig struct {
|
||||||
|
|
|
||||||
206
pkg/session/tasks.go
Normal file
206
pkg/session/tasks.go
Normal file
|
|
@ -0,0 +1,206 @@
|
||||||
|
package session
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/fileutil"
|
||||||
|
)
|
||||||
|
|
||||||
|
type TaskStatus string
|
||||||
|
|
||||||
|
const (
|
||||||
|
TaskStatusPending TaskStatus = "pending"
|
||||||
|
TaskStatusInProgress TaskStatus = "in_progress"
|
||||||
|
TaskStatusCompleted TaskStatus = "completed"
|
||||||
|
TaskStatusFailed TaskStatus = "failed"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Task struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Status TaskStatus `json:"status"`
|
||||||
|
Result string `json:"result,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type SessionTasks struct {
|
||||||
|
SessionKey string `json:"session_key"`
|
||||||
|
MessageID string `json:"message_id,omitempty"` // ID of the message to edit with progress
|
||||||
|
Tasks []Task `json:"tasks"`
|
||||||
|
Updated time.Time `json:"updated"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type TaskManager struct {
|
||||||
|
storage string
|
||||||
|
tasks map[string]*SessionTasks
|
||||||
|
mu sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTaskManager(storage string) *TaskManager {
|
||||||
|
tm := &TaskManager{
|
||||||
|
storage: storage,
|
||||||
|
tasks: make(map[string]*SessionTasks),
|
||||||
|
}
|
||||||
|
if storage != "" {
|
||||||
|
if err := tm.loadAll(); err != nil {
|
||||||
|
// just log
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return tm
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tm *TaskManager) GetOrCreate(sessionKey string) *SessionTasks {
|
||||||
|
tm.mu.Lock()
|
||||||
|
defer tm.mu.Unlock()
|
||||||
|
|
||||||
|
tasks, ok := tm.tasks[sessionKey]
|
||||||
|
if ok {
|
||||||
|
return tasks
|
||||||
|
}
|
||||||
|
|
||||||
|
tasks = &SessionTasks{
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
Tasks: []Task{},
|
||||||
|
Updated: time.Now(),
|
||||||
|
}
|
||||||
|
tm.tasks[sessionKey] = tasks
|
||||||
|
return tasks
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tm *TaskManager) CreatePlan(sessionKey string, tasks []Task) *SessionTasks {
|
||||||
|
tm.mu.Lock()
|
||||||
|
defer tm.mu.Unlock()
|
||||||
|
|
||||||
|
st := &SessionTasks{
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
Tasks: make([]Task, len(tasks)),
|
||||||
|
Updated: time.Now(),
|
||||||
|
}
|
||||||
|
copy(st.Tasks, tasks)
|
||||||
|
tm.tasks[sessionKey] = st
|
||||||
|
return st
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tm *TaskManager) UpdateTask(sessionKey, taskID string, status TaskStatus, result string) (*SessionTasks, error) {
|
||||||
|
tm.mu.Lock()
|
||||||
|
defer tm.mu.Unlock()
|
||||||
|
|
||||||
|
st, ok := tm.tasks[sessionKey]
|
||||||
|
if !ok || len(st.Tasks) == 0 {
|
||||||
|
return nil, fmt.Errorf("no active plan for session %s", sessionKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
found := false
|
||||||
|
for i := range st.Tasks {
|
||||||
|
if st.Tasks[i].ID == taskID {
|
||||||
|
st.Tasks[i].Status = status
|
||||||
|
if result != "" {
|
||||||
|
st.Tasks[i].Result = result
|
||||||
|
}
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !found {
|
||||||
|
return nil, fmt.Errorf("task %s not found in plan", taskID)
|
||||||
|
}
|
||||||
|
|
||||||
|
st.Updated = time.Now()
|
||||||
|
// Attempt to save immediately but do not block return on error.
|
||||||
|
go func() {
|
||||||
|
// Just saving this session, we need a separate lock and method for fine-grained
|
||||||
|
_ = tm.Save(sessionKey)
|
||||||
|
}()
|
||||||
|
|
||||||
|
return st, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tm *TaskManager) SetMessageID(sessionKey, messageID string) {
|
||||||
|
tm.mu.Lock()
|
||||||
|
defer tm.mu.Unlock()
|
||||||
|
|
||||||
|
if st, ok := tm.tasks[sessionKey]; ok {
|
||||||
|
st.MessageID = messageID
|
||||||
|
st.Updated = time.Now()
|
||||||
|
go func() { _ = tm.Save(sessionKey) }()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tm *TaskManager) Save(key string) error {
|
||||||
|
if tm.storage == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
filename := sanitizeFilenameTasks(key) + "_tasks.json"
|
||||||
|
|
||||||
|
if filename == "." || !filepath.IsLocal(filename) || strings.ContainsAny(filename, "/\\") {
|
||||||
|
return os.ErrInvalid
|
||||||
|
}
|
||||||
|
|
||||||
|
tm.mu.RLock()
|
||||||
|
stored, ok := tm.tasks[key]
|
||||||
|
if !ok {
|
||||||
|
tm.mu.RUnlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Make a safe copy to marshal
|
||||||
|
snapshot := SessionTasks{
|
||||||
|
SessionKey: stored.SessionKey,
|
||||||
|
MessageID: stored.MessageID,
|
||||||
|
Updated: stored.Updated,
|
||||||
|
Tasks: make([]Task, len(stored.Tasks)),
|
||||||
|
}
|
||||||
|
copy(snapshot.Tasks, stored.Tasks)
|
||||||
|
tm.mu.RUnlock()
|
||||||
|
|
||||||
|
data, err := json.MarshalIndent(snapshot, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionPath := filepath.Join(tm.storage, filename)
|
||||||
|
return fileutil.WriteFileAtomic(sessionPath, data, 0o644)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tm *TaskManager) loadAll() error {
|
||||||
|
files, err := os.ReadDir(tm.storage)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
tm.mu.Lock()
|
||||||
|
defer tm.mu.Unlock()
|
||||||
|
|
||||||
|
for _, file := range files {
|
||||||
|
if file.IsDir() || !strings.HasSuffix(file.Name(), "_tasks.json") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
path := filepath.Join(tm.storage, file.Name())
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var st SessionTasks
|
||||||
|
if err := json.Unmarshal(data, &st); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
tm.tasks[st.SessionKey] = &st
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure sanitizeFilename is accessible by importing from another file or duplicating.
|
||||||
|
// For now, copying it since it's an unexported utility in manager.go.
|
||||||
|
func sanitizeFilenameTasks(key string) string {
|
||||||
|
return strings.ReplaceAll(key, ":", "_")
|
||||||
|
}
|
||||||
|
|
@ -91,3 +91,13 @@ func ToolToSchema(tool Tool) map[string]any {
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AdvancedMessageManager represents tools that require direct, synchronous
|
||||||
|
// interaction with messaging channels (e.g., sending placeholders and editing messages).
|
||||||
|
type AdvancedMessageManager interface {
|
||||||
|
Tool
|
||||||
|
SetCallbacks(
|
||||||
|
sendPlaceholder func(channel, chatID, content string) (string, error),
|
||||||
|
editMessage func(channel, chatID, messageID, content string) error,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
|
||||||
302
pkg/tools/tasktool.go
Normal file
302
pkg/tools/tasktool.go
Normal file
|
|
@ -0,0 +1,302 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
type TaskTool struct {
|
||||||
|
taskManager *session.TaskManager
|
||||||
|
sendPlaceholder func(channel, chatID, content string) (string, error)
|
||||||
|
editMessage func(channel, chatID, messageID, content string) error
|
||||||
|
icons config.TaskToolIconsConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTaskTool(taskManager *session.TaskManager, icons config.TaskToolIconsConfig) *TaskTool {
|
||||||
|
if icons.Pending == "" {
|
||||||
|
icons.Pending = "🔘"
|
||||||
|
}
|
||||||
|
if icons.InProgress == "" {
|
||||||
|
icons.InProgress = "🟡"
|
||||||
|
}
|
||||||
|
if icons.Completed == "" {
|
||||||
|
icons.Completed = "🟢"
|
||||||
|
}
|
||||||
|
if icons.Failed == "" {
|
||||||
|
icons.Failed = "🔴"
|
||||||
|
}
|
||||||
|
|
||||||
|
return &TaskTool{
|
||||||
|
taskManager: taskManager,
|
||||||
|
icons: icons,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) Name() string {
|
||||||
|
return "tasktool"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) Description() string {
|
||||||
|
return "Manage planning mode tasks. Use action='create_plan' to start a new plan with a list of tasks. Use action='update_task' to update the status of an existing task and return the current plan state.\n\n" +
|
||||||
|
"CRITICAL INSTRUCTIONS:\n" +
|
||||||
|
"- If you determine a user request is complex and requires planning, use 'create_plan' to define a checklist. Wait for the user to accept the plan. Once accepted, execute the plan and update the status of each step using 'update_task'.\n" +
|
||||||
|
"- Use ONLY `tasktool` for storing and updating tasks. Do NOT save tasks into files.\n" +
|
||||||
|
"- Do not duplicate the plan text into the chat. `tasktool` already sends the plan automatically. Only write messages when completing a task or if there are questions/problems."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"action": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Action to perform: 'create_plan', 'update_task', 'list_plan', or 'resend_plan'",
|
||||||
|
"enum": []string{"create_plan", "update_task", "list_plan", "resend_plan"},
|
||||||
|
},
|
||||||
|
"tasks": map[string]any{
|
||||||
|
"type": "array",
|
||||||
|
"items": map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"id": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Unique identifier for the task (e.g. 'task_1')",
|
||||||
|
},
|
||||||
|
"description": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Description of the task to be completed",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"id", "description"},
|
||||||
|
},
|
||||||
|
"description": "List of tasks (only used for 'create_plan')",
|
||||||
|
},
|
||||||
|
"task_id": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "ID of the task to update (only used for 'update_task')",
|
||||||
|
},
|
||||||
|
"status": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "New status for the task (only used for 'update_task')",
|
||||||
|
"enum": []string{"pending", "in_progress", "completed", "failed"},
|
||||||
|
},
|
||||||
|
"result": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Optional brief result or note about the task update (only used for 'update_task')",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"action"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) SetCallbacks(
|
||||||
|
sendPlaceholder func(channel, chatID, content string) (string, error),
|
||||||
|
editMessage func(channel, chatID, messageID, content string) error,
|
||||||
|
) {
|
||||||
|
t.sendPlaceholder = sendPlaceholder
|
||||||
|
t.editMessage = editMessage
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
if t.taskManager == nil {
|
||||||
|
return &ToolResult{ForLLM: "tasktool: task manager not configured", IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
channel := ToolChannel(ctx)
|
||||||
|
chatID := ToolChatID(ctx)
|
||||||
|
// We use the same combination for task state as session manager might.
|
||||||
|
// But note: AgentLoop uses scopes out of routes. We'll use channel:chatID as implicit for now
|
||||||
|
// To be perfectly aligned with SessionKey, we'd need to extract SessionKey from context.
|
||||||
|
// We'll add SessionKey to context later if needed, or just use channel:chatID for tasks since planning is chat-specific.
|
||||||
|
sessionKey := fmt.Sprintf("%s:%s", channel, chatID)
|
||||||
|
|
||||||
|
action, ok := args["action"].(string)
|
||||||
|
if !ok {
|
||||||
|
return &ToolResult{ForLLM: "tasktool: action is required", IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch action {
|
||||||
|
case "create_plan":
|
||||||
|
return t.handleCreatePlan(sessionKey, channel, chatID, args)
|
||||||
|
case "update_task":
|
||||||
|
return t.handleUpdateTask(sessionKey, channel, chatID, args)
|
||||||
|
case "list_plan":
|
||||||
|
return t.handleListPlan(sessionKey)
|
||||||
|
case "resend_plan":
|
||||||
|
return t.handleResendPlan(sessionKey, channel, chatID)
|
||||||
|
default:
|
||||||
|
return &ToolResult{ForLLM: fmt.Sprintf("tasktool: unknown action '%s'", action), IsError: true}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) handleCreatePlan(sessionKey, channel, chatID string, args map[string]any) *ToolResult {
|
||||||
|
tasksRaw, ok := args["tasks"].([]interface{})
|
||||||
|
if !ok || len(tasksRaw) == 0 {
|
||||||
|
return &ToolResult{ForLLM: "tasktool: tasks array is required and cannot be empty for 'create_plan'", IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
var parsedTasks []session.Task
|
||||||
|
for i, raw := range tasksRaw {
|
||||||
|
taskMap, ok := raw.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
return &ToolResult{ForLLM: fmt.Sprintf("tasktool: invalid task at index %d", i), IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
id, ok := taskMap["id"].(string)
|
||||||
|
if !ok || id == "" {
|
||||||
|
return &ToolResult{ForLLM: fmt.Sprintf("tasktool: missing id for task at index %d", i), IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
desc, ok := taskMap["description"].(string)
|
||||||
|
if !ok || desc == "" {
|
||||||
|
return &ToolResult{ForLLM: fmt.Sprintf("tasktool: missing description for task at index %d", i), IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
parsedTasks = append(parsedTasks, session.Task{
|
||||||
|
ID: id,
|
||||||
|
Description: desc,
|
||||||
|
Status: session.TaskStatusPending,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
st := t.taskManager.CreatePlan(sessionKey, parsedTasks)
|
||||||
|
|
||||||
|
content := t.formatPlanMessage(st.Tasks)
|
||||||
|
|
||||||
|
// Send message through callback if available
|
||||||
|
if t.sendPlaceholder != nil {
|
||||||
|
msgID, err := t.sendPlaceholder(channel, chatID, content)
|
||||||
|
if err == nil && msgID != "" {
|
||||||
|
t.taskManager.SetMessageID(sessionKey, msgID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
tasksJSON, _ := json.Marshal(parsedTasks)
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf("Plan created with %d tasks.\nTasks: %s", len(parsedTasks), string(tasksJSON)),
|
||||||
|
Silent: true, // We already sent the message via callback
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) handleListPlan(sessionKey string) *ToolResult {
|
||||||
|
st := t.taskManager.GetOrCreate(sessionKey)
|
||||||
|
if len(st.Tasks) == 0 {
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: "No active plan found for this session.",
|
||||||
|
Silent: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
content := t.formatPlanMessage(st.Tasks)
|
||||||
|
tasksJSON, _ := json.Marshal(st.Tasks)
|
||||||
|
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf("Current plan state:\n%s\n\nRaw JSON:\n%s", content, string(tasksJSON)),
|
||||||
|
Silent: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) handleResendPlan(sessionKey, channel, chatID string) *ToolResult {
|
||||||
|
st := t.taskManager.GetOrCreate(sessionKey)
|
||||||
|
if len(st.Tasks) == 0 {
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: "No active plan found for this session to resend.",
|
||||||
|
IsError: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
content := t.formatPlanMessage(st.Tasks)
|
||||||
|
|
||||||
|
if t.sendPlaceholder != nil {
|
||||||
|
msgID, err := t.sendPlaceholder(channel, chatID, content)
|
||||||
|
if err == nil && msgID != "" {
|
||||||
|
t.taskManager.SetMessageID(sessionKey, msgID)
|
||||||
|
} else {
|
||||||
|
return &ToolResult{ForLLM: fmt.Sprintf("Failed to resend message: %v", err), IsError: true}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
return &ToolResult{ForLLM: "tasktool: channel sending callbacks are not configured", IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
tasksJSON, _ := json.Marshal(st.Tasks)
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf("Plan successfully resent as a new message.\nTasks: %s", string(tasksJSON)),
|
||||||
|
Silent: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) handleUpdateTask(sessionKey, channel, chatID string, args map[string]any) *ToolResult {
|
||||||
|
taskID, _ := args["task_id"].(string)
|
||||||
|
if taskID == "" {
|
||||||
|
return &ToolResult{ForLLM: "tasktool: task_id is required for 'update_task'", IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
statusStr, _ := args["status"].(string)
|
||||||
|
if statusStr == "" {
|
||||||
|
return &ToolResult{ForLLM: "tasktool: status is required for 'update_task'", IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
result, _ := args["result"].(string)
|
||||||
|
|
||||||
|
st, err := t.taskManager.UpdateTask(sessionKey, taskID, session.TaskStatus(statusStr), result)
|
||||||
|
if err != nil {
|
||||||
|
return &ToolResult{ForLLM: fmt.Sprintf("tasktool: %v", err), IsError: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
content := t.formatPlanMessage(st.Tasks)
|
||||||
|
|
||||||
|
// Edit message through callback if available
|
||||||
|
if t.editMessage != nil && st.MessageID != "" {
|
||||||
|
_ = t.editMessage(channel, chatID, st.MessageID, content)
|
||||||
|
} else if t.sendPlaceholder != nil && st.MessageID == "" {
|
||||||
|
// Fallback: send new progress message if we didn't have one
|
||||||
|
msgID, err := t.sendPlaceholder(channel, chatID, content)
|
||||||
|
if err == nil && msgID != "" {
|
||||||
|
t.taskManager.SetMessageID(sessionKey, msgID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
tasksJSON, _ := json.Marshal(st.Tasks)
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf("Task '%s' updated to '%s'. Current plan:\n%s", taskID, statusStr, string(tasksJSON)),
|
||||||
|
Silent: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TaskTool) formatPlanMessage(tasks []session.Task) string {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("📋 **Execution Plan**:\n\n")
|
||||||
|
|
||||||
|
for _, task := range tasks {
|
||||||
|
var icon string
|
||||||
|
switch task.Status {
|
||||||
|
case session.TaskStatusPending:
|
||||||
|
icon = t.icons.Pending
|
||||||
|
case session.TaskStatusInProgress:
|
||||||
|
icon = t.icons.InProgress
|
||||||
|
case session.TaskStatusCompleted:
|
||||||
|
icon = t.icons.Completed
|
||||||
|
case session.TaskStatusFailed:
|
||||||
|
icon = t.icons.Failed
|
||||||
|
default:
|
||||||
|
icon = t.icons.Pending
|
||||||
|
}
|
||||||
|
|
||||||
|
// Primitive markdown-to-html regex parser across multiple lines.
|
||||||
|
// For the description and result, we replace lone underscores to prevent similar italic bugs.
|
||||||
|
safeDesc := strings.ReplaceAll(task.Description, "_", " ")
|
||||||
|
|
||||||
|
sb.WriteString(fmt.Sprintf("%s %s\n", icon, safeDesc))
|
||||||
|
if task.Result != "" {
|
||||||
|
safeResult := strings.ReplaceAll(task.Result, "_", " ")
|
||||||
|
sb.WriteString(fmt.Sprintf(" **Result**: %s\n", safeResult))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue