refactor: extract DI to channels.Manager, resolve upstream merge and add chunking to SendMessageWithID

This commit is contained in:
Dmitrii Balabanov 2026-03-07 23:19:08 +02:00
parent 4768edc67b
commit cafdb51df2
10 changed files with 729 additions and 112 deletions

View file

@ -29,6 +29,7 @@ import (
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/providers"
"github.com/sipeed/picoclaw/pkg/routing"
"github.com/sipeed/picoclaw/pkg/session"
"github.com/sipeed/picoclaw/pkg/skills"
"github.com/sipeed/picoclaw/pkg/state"
"github.com/sipeed/picoclaw/pkg/tools"
@ -48,6 +49,7 @@ type AgentLoop struct {
mediaStore media.MediaStore
transcriber voice.Transcriber
cmdRegistry *commands.Registry
taskManager *session.TaskManager
}
// processOptions configures how a message is processed
@ -80,9 +82,6 @@ func NewAgentLoop(
) *AgentLoop {
registry := NewAgentRegistry(cfg, provider)
// Register shared tools to all agents
registerSharedTools(cfg, msgBus, registry, provider)
// Set up shared fallback chain
cooldown := providers.NewCooldownTracker()
fallbackChain := providers.NewFallbackChain(cooldown)
@ -102,8 +101,12 @@ func NewAgentLoop(
summarizing: sync.Map{},
fallback: fallbackChain,
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
}
@ -113,6 +116,7 @@ func registerSharedTools(
msgBus *bus.MessageBus,
registry *AgentRegistry,
provider providers.LLMProvider,
taskManager *session.TaskManager,
) {
for _, agentID := range registry.ListAgentIDs() {
agent, ok := registry.GetAgent(agentID)
@ -232,6 +236,12 @@ func registerSharedTools(
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) {
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.

View file

@ -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.
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
r.mu.RLock()

View file

@ -50,3 +50,10 @@ type PlaceholderRecorder interface {
type CommandRegistrarCapable interface {
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)
}

View file

@ -833,3 +833,45 @@ func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, conten
channel, _ := m.channels[channelName]
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)
}

View file

@ -164,23 +164,31 @@ func (c *TelegramChannel) Stop(ctx context.Context) 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() {
return channels.ErrNotRunning
return "", channels.ErrNotRunning
}
chatID, err := parseChatID(msg.ChatID)
cid, err := parseChatID(chatID)
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 == "" {
return nil
if content == "" {
return "", nil
}
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
// 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.
queue := []string{msg.Content}
queue := []string{content}
var lastMsgID int
for len(queue) > 0 {
chunk := queue[0]
queue = queue[1:]
@ -200,31 +208,38 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
continue
}
if err := c.sendHTMLChunk(ctx, chatID, htmlContent, chunk); err != nil {
return err
msgID, err := c.sendHTMLChunk(ctx, cid, htmlContent, chunk)
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
// 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.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{
"error": err.Error(),
})
tgMsg.Text = mdFallback
tgMsg.ParseMode = ""
if _, err = c.bot.SendMessage(ctx, tgMsg); err != nil {
return fmt.Errorf("telegram send: %w", channels.ErrTemporary)
msg, err = c.bot.SendMessage(ctx, tgMsg)
if err != nil {
return 0, fmt.Errorf("telegram send: %w", channels.ErrTemporary)
}
}
return nil
return msg.MessageID, nil
}
// StartTyping implements channels.TypingCapable.

View file

@ -1,19 +1,19 @@
package telegram
import (
"context"
"encoding/json"
"errors"
"strings"
"testing"
"context"
"encoding/json"
"errors"
"strings"
"testing"
"github.com/mymmrac/telego"
ta "github.com/mymmrac/telego/telegoapi"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/mymmrac/telego"
ta "github.com/mymmrac/telego/telegoapi"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels"
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels"
)
const testToken = "1234567890:aaaabbbbaaaabbbbaaaabbbbaaaabbbbccc"
@ -42,8 +42,8 @@ func (s *stubConstructor) JSONRequest(parameters any) (*ta.RequestData, error) {
}
func (s *stubConstructor) MultipartRequest(
parameters map[string]string,
files map[string]ta.NamedReader,
parameters map[string]string,
files map[string]ta.NamedReader,
) (*ta.RequestData, error) {
return &ta.RequestData{}, nil
}
@ -62,15 +62,15 @@ func newTestChannel(t *testing.T, caller *stubCaller) *TelegramChannel {
t.Helper()
bot, err := telego.NewBot(testToken,
telego.WithAPICaller(caller),
telego.WithRequestConstructor(&stubConstructor{}),
telego.WithAPICaller(caller),
telego.WithRequestConstructor(&stubConstructor{}),
telego.WithDiscardLogger(),
)
require.NoError(t, err)
base := channels.NewBaseChannel("telegram", nil, nil, nil,
channels.WithMaxMessageLength(4000),
)
channels.WithMaxMessageLength(4000),
)
base.SetRunning(true)
return &TelegramChannel{
@ -80,25 +80,7 @@ func newTestChannel(t *testing.T, caller *stubCaller) *TelegramChannel {
}
}
func TestSend_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)
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) {
func TestSend_Wrapper(t *testing.T) {
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
return successResponse(t), nil
@ -112,14 +94,41 @@ func TestSend_ShortMessage_SingleCall(t *testing.T) {
})
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")
}
func TestSend_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).
func TestSendMessageWithID_LongMessage_SingleCall(t *testing.T) {
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
return successResponse(t), nil
@ -129,21 +138,18 @@ func TestSend_LongMessage_SingleCall(t *testing.T) {
longContent := strings.Repeat("a", 4000)
err := ch.Send(context.Background(), bus.OutboundMessage{
ChatID: "12345",
Content: longContent,
})
msgID, err := ch.SendMessageWithID(context.Background(), "12345", longContent)
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")
}
func TestSend_HTMLFallback_PerChunk(t *testing.T) {
func TestSendMessageWithID_HTMLFallback_PerChunk(t *testing.T) {
callCount := 0
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
callCount++
// Fail on odd calls (HTML attempt), succeed on even calls (plain text fallback)
if callCount%2 == 1 {
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)
err := ch.Send(context.Background(), bus.OutboundMessage{
ChatID: "12345",
Content: "Hello **world**",
})
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello **world**")
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")
}
func TestSend_HTMLFallback_BothFail(t *testing.T) {
func TestSendMessageWithID_HTMLFallback_BothFail(t *testing.T) {
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
return nil, errors.New("send failed")
@ -170,19 +173,15 @@ func TestSend_HTMLFallback_BothFail(t *testing.T) {
}
ch := newTestChannel(t, caller)
err := ch.Send(context.Background(), bus.OutboundMessage{
ChatID: "12345",
Content: "Hello",
})
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello")
assert.Error(t, err)
assert.Empty(t, msgID)
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")
}
func TestSend_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.
func TestSendMessageWithID_LongMessage_HTMLFallback_StopsOnError(t *testing.T) {
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
return nil, errors.New("send failed")
@ -192,17 +191,14 @@ func TestSend_LongMessage_HTMLFallback_StopsOnError(t *testing.T) {
longContent := strings.Repeat("x", 4001)
err := ch.Send(context.Background(), bus.OutboundMessage{
ChatID: "12345",
Content: longContent,
})
msgID, err := ch.SendMessageWithID(context.Background(), "12345", longContent)
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")
}
func TestSend_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
func TestSendMessageWithID_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
return successResponse(t), nil
@ -210,31 +206,17 @@ func TestSend_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
}
ch := newTestChannel(t, caller)
// Create markdown whose length is <= 4000 but whose HTML expansion is much longer.
// "**a** " (6 chars) becomes "<b>a</b> " (9 chars) in HTML, so repeating it many times
// 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")
markdownContent := strings.Repeat("**a** ", 600)
assert.LessOrEqual(t, len([]rune(markdownContent)), 4000)
htmlExpanded := markdownToTelegramHTML(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,
})
msgID, err := ch.SendMessageWithID(context.Background(), "12345", markdownContent)
assert.NoError(t, err)
assert.Greater(
t, len(caller.calls), 1,
"markdown-short but HTML-long message should be split into multiple SendMessage calls",
)
assert.Equal(t, "1", msgID)
assert.Greater(t, len(caller.calls), 1, "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{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
t.Fatal("should not be called")
@ -244,16 +226,14 @@ func TestSend_NotRunning(t *testing.T) {
ch := newTestChannel(t, caller)
ch.SetRunning(false)
err := ch.Send(context.Background(), bus.OutboundMessage{
ChatID: "12345",
Content: "Hello",
})
msgID, err := ch.SendMessageWithID(context.Background(), "12345", "Hello")
assert.ErrorIs(t, err, channels.ErrNotRunning)
assert.Empty(t, msgID)
assert.Empty(t, caller.calls)
}
func TestSend_InvalidChatID(t *testing.T) {
func TestSendMessageWithID_InvalidChatID(t *testing.T) {
caller := &stubCaller{
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
t.Fatal("should not be called")
@ -262,12 +242,10 @@ func TestSend_InvalidChatID(t *testing.T) {
}
ch := newTestChannel(t, caller)
err := ch.Send(context.Background(), bus.OutboundMessage{
ChatID: "not-a-number",
Content: "Hello",
})
msgID, err := ch.SendMessageWithID(context.Background(), "not-a-number", "Hello")
assert.Error(t, err)
assert.Empty(t, msgID)
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
assert.Empty(t, caller.calls)
}

View file

@ -666,6 +666,19 @@ type ToolsConfig struct {
Subagent ToolConfig `json:"subagent" envPrefix:"PICOCLAW_TOOLS_SUBAGENT_"`
WebFetch ToolConfig `json:"web_fetch" envPrefix:"PICOCLAW_TOOLS_WEB_FETCH_"`
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 {

206
pkg/session/tasks.go Normal file
View 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, ":", "_")
}

View file

@ -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
View 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()
}