refactor: Move SanitizeFilename to fileutil, propagate context.Context in task tool callbacks, and improve task tool and channel manager error handling.

This commit is contained in:
Dmitrii Balabanov 2026-03-07 23:49:46 +02:00
parent e938241c0d
commit 896fb08f93
11 changed files with 91 additions and 76 deletions

View file

@ -410,16 +410,16 @@ func (al *AgentLoop) bindAdvancedMessageManagers(cm *channels.Manager) {
if advancedManager, ok := t.(tools.AdvancedMessageManager); ok {
advancedManager.SetCallbacks(
// sendPlaceholder
func(channelName, chatID, content string) (string, error) {
return cm.SendMessageWithID(context.Background(), bus.OutboundMessage{
func(ctx context.Context, channelName, chatID, content string) (string, error) {
return cm.SendMessageWithID(ctx, bus.OutboundMessage{
Channel: channelName,
ChatID: chatID,
Content: content,
})
},
// editMessage
func(channelName, chatID, messageID, content string) error {
return cm.EditMessage(context.Background(), channelName, chatID, messageID, content)
func(ctx context.Context, channelName, chatID, messageID, content string) error {
return cm.EditMessage(ctx, channelName, chatID, messageID, content)
},
)
}

View file

@ -22,6 +22,7 @@ import (
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/fileutil"
"github.com/sipeed/picoclaw/pkg/identity"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/media"
@ -612,7 +613,7 @@ func (c *FeishuChannel) downloadResource(
return ""
}
ext := filepath.Ext(filename)
localPath := filepath.Join(mediaDir, utils.SanitizeFilename(messageID+"-"+fileKey+ext))
localPath := filepath.Join(mediaDir, fileutil.SanitizeFilename(messageID+"-"+fileKey+ext))
out, err := os.Create(localPath)
if err != nil {

View file

@ -849,9 +849,13 @@ func (m *Manager) SendMessageWithID(ctx context.Context, msg bus.OutboundMessage
return msgID, nil
}
logger.ErrorCF("manager", "SendMessageWithID failed", map[string]any{"error": err, "msgID": msgID})
logger.WarnCF("manager", "channel does not implement SyncSender", map[string]any{"channel": msg.Channel})
if err == nil {
err = fmt.Errorf("sync sender returned empty message ID")
}
return "", err
}
logger.WarnCF("manager", "channel does not implement SyncSender", map[string]any{"channel": msg.Channel})
logger.WarnCF("manager", "falling back to bus publish", nil)
m.bus.PublishOutbound(ctx, msg)

View file

@ -675,7 +675,7 @@ type TaskToolConfig struct {
}
type TaskToolIconsConfig struct {
Pending string `json:"pending" env:"PICOCLAW_TOOLS_TASK_TOOL_ICONS_PENDING" default:""`
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:"🔴"`

View file

@ -11,9 +11,27 @@ import (
"fmt"
"os"
"path/filepath"
"strings"
"time"
)
// SanitizeFilename removes potentially dangerous characters from a filename
// and returns a safe version for local filesystem storage.
func SanitizeFilename(filename string) string {
// First, replace common directory separators and colons with underscores
safe := strings.ReplaceAll(filename, "/", "_")
safe = strings.ReplaceAll(safe, "\\", "_")
safe = strings.ReplaceAll(safe, ":", "_")
// Then get the base filename to ensure no path components remain
base := filepath.Base(safe)
// Finally, remove any exact ".." sequences that might have slipped through
base = strings.ReplaceAll(base, "..", "")
return base
}
// WriteFileAtomic atomically writes data to a file using a temp file + rename pattern.
//
// This guarantees that the target file is either:

View file

@ -8,6 +8,7 @@ import (
"sync"
"time"
"github.com/sipeed/picoclaw/pkg/fileutil"
"github.com/sipeed/picoclaw/pkg/providers"
)
@ -145,21 +146,12 @@ func (sm *SessionManager) TruncateHistory(key string, keepLast int) {
session.Updated = time.Now()
}
// sanitizeFilename converts a session key into a cross-platform safe filename.
// Session keys use "channel:chatID" (e.g. "telegram:123456") but ':' is the
// volume separator on Windows, so filepath.Base would misinterpret the key.
// We replace it with '_'. The original key is preserved inside the JSON file,
// so loadSessions still maps back to the right in-memory key.
func sanitizeFilename(key string) string {
return strings.ReplaceAll(key, ":", "_")
}
func (sm *SessionManager) Save(key string) error {
if sm.storage == "" {
return nil
}
filename := sanitizeFilename(key)
filename := fileutil.SanitizeFilename(key)
// filepath.IsLocal rejects empty names, "..", absolute paths, and
// OS-reserved device names (NUL, COM1 … on Windows).

View file

@ -4,6 +4,8 @@ import (
"os"
"path/filepath"
"testing"
"github.com/sipeed/picoclaw/pkg/fileutil"
)
func TestSanitizeFilename(t *testing.T) {
@ -21,9 +23,9 @@ func TestSanitizeFilename(t *testing.T) {
for _, tt := range tests {
t.Run(tt.input, func(t *testing.T) {
got := sanitizeFilename(tt.input)
got := fileutil.SanitizeFilename(tt.input)
if got != tt.expected {
t.Errorf("sanitizeFilename(%q) = %q, want %q", tt.input, got, tt.expected)
t.Errorf("SanitizeFilename(%q) = %q, want %q", tt.input, got, tt.expected)
}
})
}

View file

@ -10,6 +10,7 @@ import (
"time"
"github.com/sipeed/picoclaw/pkg/fileutil"
"github.com/sipeed/picoclaw/pkg/logger"
)
type TaskStatus string
@ -48,12 +49,25 @@ func NewTaskManager(storage string) *TaskManager {
}
if storage != "" {
if err := tm.loadAll(); err != nil {
// just log
logger.ErrorCF("tasks", "Failed to load session tasks on startup", map[string]any{
"error": err.Error(),
})
}
}
return tm
}
func (tm *TaskManager) Get(sessionKey string) *SessionTasks {
tm.mu.RLock()
defer tm.mu.RUnlock()
tasks, ok := tm.tasks[sessionKey]
if ok {
return tasks
}
return nil
}
func (tm *TaskManager) GetOrCreate(sessionKey string) *SessionTasks {
tm.mu.Lock()
defer tm.mu.Unlock()
@ -137,7 +151,10 @@ func (tm *TaskManager) Save(key string) error {
return nil
}
filename := sanitizeFilenameTasks(key) + "_tasks.json"
if err := os.MkdirAll(tm.storage, 0o755); err != nil {
return err
}
filename := fileutil.SanitizeFilename(key) + "_tasks.json"
if filename == "." || !filepath.IsLocal(filename) || strings.ContainsAny(filename, "/\\") {
return os.ErrInvalid
@ -198,9 +215,3 @@ func (tm *TaskManager) loadAll() error {
}
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

@ -97,7 +97,7 @@ func ToolToSchema(tool Tool) map[string]any {
type AdvancedMessageManager interface {
Tool
SetCallbacks(
sendPlaceholder func(channel, chatID, content string) (string, error),
editMessage func(channel, chatID, messageID, content string) error,
sendPlaceholder func(ctx context.Context, channel, chatID, content string) (string, error),
editMessage func(ctx context.Context, channel, chatID, messageID, content string) error,
)
}

View file

@ -7,29 +7,18 @@ import (
"strings"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/logger"
"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
sendPlaceholder func(ctx context.Context, channel, chatID, content string) (string, error)
editMessage func(ctx context.Context, 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,
@ -95,8 +84,8 @@ func (t *TaskTool) Parameters() map[string]any {
}
func (t *TaskTool) SetCallbacks(
sendPlaceholder func(channel, chatID, content string) (string, error),
editMessage func(channel, chatID, messageID, content string) error,
sendPlaceholder func(ctx context.Context, channel, chatID, content string) (string, error),
editMessage func(ctx context.Context, channel, chatID, messageID, content string) error,
) {
t.sendPlaceholder = sendPlaceholder
t.editMessage = editMessage
@ -122,19 +111,19 @@ func (t *TaskTool) Execute(ctx context.Context, args map[string]any) *ToolResult
switch action {
case "create_plan":
return t.handleCreatePlan(sessionKey, channel, chatID, args)
return t.handleCreatePlan(ctx, sessionKey, channel, chatID, args)
case "update_task":
return t.handleUpdateTask(sessionKey, channel, chatID, args)
return t.handleUpdateTask(ctx, sessionKey, channel, chatID, args)
case "list_plan":
return t.handleListPlan(sessionKey)
case "resend_plan":
return t.handleResendPlan(sessionKey, channel, chatID)
return t.handleResendPlan(ctx, 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 {
func (t *TaskTool) handleCreatePlan(ctx context.Context, 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}
@ -170,7 +159,7 @@ func (t *TaskTool) handleCreatePlan(sessionKey, channel, chatID string, args map
// Send message through callback if available
if t.sendPlaceholder != nil {
msgID, err := t.sendPlaceholder(channel, chatID, content)
msgID, err := t.sendPlaceholder(ctx, channel, chatID, content)
if err == nil && msgID != "" {
t.taskManager.SetMessageID(sessionKey, msgID)
}
@ -184,8 +173,8 @@ func (t *TaskTool) handleCreatePlan(sessionKey, channel, chatID string, args map
}
func (t *TaskTool) handleListPlan(sessionKey string) *ToolResult {
st := t.taskManager.GetOrCreate(sessionKey)
if len(st.Tasks) == 0 {
st := t.taskManager.Get(sessionKey)
if st == nil || len(st.Tasks) == 0 {
return &ToolResult{
ForLLM: "No active plan found for this session.",
Silent: true,
@ -201,7 +190,7 @@ func (t *TaskTool) handleListPlan(sessionKey string) *ToolResult {
}
}
func (t *TaskTool) handleResendPlan(sessionKey, channel, chatID string) *ToolResult {
func (t *TaskTool) handleResendPlan(ctx context.Context, sessionKey, channel, chatID string) *ToolResult {
st := t.taskManager.GetOrCreate(sessionKey)
if len(st.Tasks) == 0 {
return &ToolResult{
@ -213,9 +202,13 @@ func (t *TaskTool) handleResendPlan(sessionKey, channel, chatID string) *ToolRes
content := t.formatPlanMessage(st.Tasks)
if t.sendPlaceholder != nil {
msgID, err := t.sendPlaceholder(channel, chatID, content)
if err == nil && msgID != "" {
msgID, err := t.sendPlaceholder(ctx, channel, chatID, content)
if err == nil {
if msgID != "" {
t.taskManager.SetMessageID(sessionKey, msgID)
}
// If err == nil but msgID == "", the channel delivered the message
// (or is async) but doesn't support returning IDs. We consider this a success.
} else {
return &ToolResult{ForLLM: fmt.Sprintf("Failed to resend message: %v", err), IsError: true}
}
@ -230,7 +223,7 @@ func (t *TaskTool) handleResendPlan(sessionKey, channel, chatID string) *ToolRes
}
}
func (t *TaskTool) handleUpdateTask(sessionKey, channel, chatID string, args map[string]any) *ToolResult {
func (t *TaskTool) handleUpdateTask(ctx context.Context, 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}
@ -252,10 +245,17 @@ func (t *TaskTool) handleUpdateTask(sessionKey, channel, chatID string, args map
// Edit message through callback if available
if t.editMessage != nil && st.MessageID != "" {
_ = t.editMessage(channel, chatID, st.MessageID, content)
if err := t.editMessage(ctx, channel, chatID, st.MessageID, content); err != nil {
logger.WarnCF("tasktool", "Failed to edit task message", map[string]any{
"channel": channel,
"chat_id": chatID,
"message_id": st.MessageID,
"error": err.Error(),
})
}
} 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)
msgID, err := t.sendPlaceholder(ctx, channel, chatID, content)
if err == nil && msgID != "" {
t.taskManager.SetMessageID(sessionKey, msgID)
}
@ -289,11 +289,11 @@ func (t *TaskTool) formatPlanMessage(tasks []session.Task) string {
// 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, "_", " ")
safeDesc := strings.ReplaceAll(task.Description, "_", "\\_")
sb.WriteString(fmt.Sprintf("%s %s\n", icon, safeDesc))
if task.Result != "" {
safeResult := strings.ReplaceAll(task.Result, "_", " ")
safeResult := strings.ReplaceAll(task.Result, "_", "\\_")
sb.WriteString(fmt.Sprintf(" **Result**: %s\n", safeResult))
}
}

View file

@ -11,6 +11,7 @@ import (
"github.com/google/uuid"
"github.com/sipeed/picoclaw/pkg/fileutil"
"github.com/sipeed/picoclaw/pkg/logger"
)
@ -34,20 +35,6 @@ func IsAudioFile(filename, contentType string) bool {
return false
}
// SanitizeFilename removes potentially dangerous characters from a filename
// and returns a safe version for local filesystem storage.
func SanitizeFilename(filename string) string {
// Get the base filename without path
base := filepath.Base(filename)
// Remove any directory traversal attempts
base = strings.ReplaceAll(base, "..", "")
base = strings.ReplaceAll(base, "/", "_")
base = strings.ReplaceAll(base, "\\", "_")
return base
}
// DownloadOptions holds optional parameters for downloading files
type DownloadOptions struct {
Timeout time.Duration
@ -76,7 +63,7 @@ func DownloadFile(urlStr, filename string, opts DownloadOptions) string {
}
// Generate unique filename with UUID prefix to prevent conflicts
safeName := SanitizeFilename(filename)
safeName := fileutil.SanitizeFilename(filename)
localPath := filepath.Join(mediaDir, uuid.New().String()[:8]+"_"+safeName)
// Create HTTP request