fix: 修复了飞书,企业微信代码多模态的处理代码,增加单页面WEB

This commit is contained in:
opcache 2026-03-12 11:53:05 +08:00
parent 8207c1c7e6
commit 6668c9eed1
117 changed files with 13612 additions and 1674 deletions

View file

@ -10,8 +10,8 @@ import (
picoclawconfig "github.com/sipeed/picoclaw/pkg/config" picoclawconfig "github.com/sipeed/picoclaw/pkg/config"
) )
func (s *appState) channelMenu() tview.Primitive { func (s *appState) buildChannelMenuItems() []MenuItem {
items := []MenuItem{ return []MenuItem{
{Label: "Back", Description: "Return to main menu", Action: func() { s.pop() }}, {Label: "Back", Description: "Return to main menu", Action: func() { s.pop() }},
channelItem( channelItem(
"Telegram", "Telegram",
@ -86,8 +86,10 @@ func (s *appState) channelMenu() tview.Primitive {
func() { s.push("channel-wecomapp", s.wecomAppForm()) }, func() { s.push("channel-wecomapp", s.wecomAppForm()) },
), ),
} }
}
menu := NewMenu("Channels", items) func (s *appState) channelMenu() tview.Primitive {
menu := NewMenu("Channels", s.buildChannelMenuItems())
menu.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey { menu.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
if event.Key() == tcell.KeyEsc { if event.Key() == tcell.KeyEsc {
s.pop() s.pop()
@ -103,199 +105,72 @@ func (s *appState) channelMenu() tview.Primitive {
} }
func refreshChannelMenuFromState(menu *Menu, s *appState) { func refreshChannelMenuFromState(menu *Menu, s *appState) {
items := []MenuItem{ menu.applyItems(s.buildChannelMenuItems())
{Label: "Back", Description: "Return to main menu", Action: func() { s.pop() }},
channelItem(
"Telegram",
"Telegram bot settings",
s.config.Channels.Telegram.Enabled,
func() { s.push("channel-telegram", s.telegramForm()) },
),
channelItem(
"Discord",
"Discord bot settings",
s.config.Channels.Discord.Enabled,
func() { s.push("channel-discord", s.discordForm()) },
),
channelItem(
"QQ",
"QQ bot settings",
s.config.Channels.QQ.Enabled,
func() { s.push("channel-qq", s.qqForm()) },
),
channelItem(
"MaixCam",
"MaixCam gateway",
s.config.Channels.MaixCam.Enabled,
func() { s.push("channel-maixcam", s.maixcamForm()) },
),
channelItem(
"WhatsApp",
"WhatsApp bridge",
s.config.Channels.WhatsApp.Enabled,
func() { s.push("channel-whatsapp", s.whatsappForm()) },
),
channelItem(
"Feishu",
"Feishu bot settings",
s.config.Channels.Feishu.Enabled,
func() { s.push("channel-feishu", s.feishuForm()) },
),
channelItem(
"DingTalk",
"DingTalk bot settings",
s.config.Channels.DingTalk.Enabled,
func() { s.push("channel-dingtalk", s.dingtalkForm()) },
),
channelItem(
"Slack",
"Slack bot settings",
s.config.Channels.Slack.Enabled,
func() { s.push("channel-slack", s.slackForm()) },
),
channelItem(
"LINE",
"LINE bot settings",
s.config.Channels.LINE.Enabled,
func() { s.push("channel-line", s.lineForm()) },
),
channelItem(
"OneBot",
"OneBot settings",
s.config.Channels.OneBot.Enabled,
func() { s.push("channel-onebot", s.onebotForm()) },
),
channelItem(
"WeCom",
"WeCom bot settings",
s.config.Channels.WeCom.Enabled,
func() { s.push("channel-wecom", s.wecomForm()) },
),
channelItem(
"WeCom App",
"WeCom App settings",
s.config.Channels.WeComApp.Enabled,
func() { s.push("channel-wecomapp", s.wecomAppForm()) },
),
}
menu.applyItems(items)
} }
func (s *appState) telegramForm() tview.Primitive { func (s *appState) telegramForm() tview.Primitive {
cfg := &s.config.Channels.Telegram cfg := &s.config.Channels.Telegram
form := baseChannelForm("Telegram", cfg.Enabled, func(v bool) { form := baseChannelForm("Telegram", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Token", cfg.Token, 128, nil, func(text string) { form.AddInputField("Token", cfg.Token, 128, nil, func(text string) {
cfg.Token = strings.TrimSpace(text) cfg.Token = strings.TrimSpace(text)
}) })
form.AddInputField("Proxy", cfg.Proxy, 128, nil, func(text string) { form.AddInputField("Proxy", cfg.Proxy, 128, nil, func(text string) {
cfg.Proxy = strings.TrimSpace(text) cfg.Proxy = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) discordForm() tview.Primitive { func (s *appState) discordForm() tview.Primitive {
cfg := &s.config.Channels.Discord cfg := &s.config.Channels.Discord
form := baseChannelForm("Discord", cfg.Enabled, func(v bool) { form := baseChannelForm("Discord", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Token", cfg.Token, 128, nil, func(text string) { form.AddInputField("Token", cfg.Token, 128, nil, func(text string) {
cfg.Token = strings.TrimSpace(text) cfg.Token = strings.TrimSpace(text)
}) })
form.AddCheckbox("Mention Only", cfg.MentionOnly, func(checked bool) { form.AddCheckbox("Mention Only", cfg.MentionOnly, func(checked bool) {
cfg.MentionOnly = checked cfg.MentionOnly = checked
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) qqForm() tview.Primitive { func (s *appState) qqForm() tview.Primitive {
cfg := &s.config.Channels.QQ cfg := &s.config.Channels.QQ
form := baseChannelForm("QQ", cfg.Enabled, func(v bool) { form := baseChannelForm("QQ", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("App ID", cfg.AppID, 64, nil, func(text string) { form.AddInputField("App ID", cfg.AppID, 64, nil, func(text string) {
cfg.AppID = strings.TrimSpace(text) cfg.AppID = strings.TrimSpace(text)
}) })
form.AddInputField("App Secret", cfg.AppSecret, 128, nil, func(text string) { form.AddInputField("App Secret", cfg.AppSecret, 128, nil, func(text string) {
cfg.AppSecret = strings.TrimSpace(text) cfg.AppSecret = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) maixcamForm() tview.Primitive { func (s *appState) maixcamForm() tview.Primitive {
cfg := &s.config.Channels.MaixCam cfg := &s.config.Channels.MaixCam
form := baseChannelForm("MaixCam", cfg.Enabled, func(v bool) { form := baseChannelForm("MaixCam", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Host", cfg.Host, 64, nil, func(text string) { form.AddInputField("Host", cfg.Host, 64, nil, func(text string) {
cfg.Host = strings.TrimSpace(text) cfg.Host = strings.TrimSpace(text)
}) })
addIntField(form, "Port", cfg.Port, func(value int) { cfg.Port = value }) addIntField(form, "Port", cfg.Port, func(value int) { cfg.Port = value })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) whatsappForm() tview.Primitive { func (s *appState) whatsappForm() tview.Primitive {
cfg := &s.config.Channels.WhatsApp cfg := &s.config.Channels.WhatsApp
form := baseChannelForm("WhatsApp", cfg.Enabled, func(v bool) { form := baseChannelForm("WhatsApp", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Bridge URL", cfg.BridgeURL, 128, nil, func(text string) { form.AddInputField("Bridge URL", cfg.BridgeURL, 128, nil, func(text string) {
cfg.BridgeURL = strings.TrimSpace(text) cfg.BridgeURL = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) feishuForm() tview.Primitive { func (s *appState) feishuForm() tview.Primitive {
cfg := &s.config.Channels.Feishu cfg := &s.config.Channels.Feishu
form := baseChannelForm("Feishu", cfg.Enabled, func(v bool) { form := baseChannelForm("Feishu", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("App ID", cfg.AppID, 64, nil, func(text string) { form.AddInputField("App ID", cfg.AppID, 64, nil, func(text string) {
cfg.AppID = strings.TrimSpace(text) cfg.AppID = strings.TrimSpace(text)
}) })
@ -308,66 +183,39 @@ func (s *appState) feishuForm() tview.Primitive {
form.AddInputField("Verification Token", cfg.VerificationToken, 128, nil, func(text string) { form.AddInputField("Verification Token", cfg.VerificationToken, 128, nil, func(text string) {
cfg.VerificationToken = strings.TrimSpace(text) cfg.VerificationToken = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) dingtalkForm() tview.Primitive { func (s *appState) dingtalkForm() tview.Primitive {
cfg := &s.config.Channels.DingTalk cfg := &s.config.Channels.DingTalk
form := baseChannelForm("DingTalk", cfg.Enabled, func(v bool) { form := baseChannelForm("DingTalk", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Client ID", cfg.ClientID, 64, nil, func(text string) { form.AddInputField("Client ID", cfg.ClientID, 64, nil, func(text string) {
cfg.ClientID = strings.TrimSpace(text) cfg.ClientID = strings.TrimSpace(text)
}) })
form.AddInputField("Client Secret", cfg.ClientSecret, 128, nil, func(text string) { form.AddInputField("Client Secret", cfg.ClientSecret, 128, nil, func(text string) {
cfg.ClientSecret = strings.TrimSpace(text) cfg.ClientSecret = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) slackForm() tview.Primitive { func (s *appState) slackForm() tview.Primitive {
cfg := &s.config.Channels.Slack cfg := &s.config.Channels.Slack
form := baseChannelForm("Slack", cfg.Enabled, func(v bool) { form := baseChannelForm("Slack", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Bot Token", cfg.BotToken, 128, nil, func(text string) { form.AddInputField("Bot Token", cfg.BotToken, 128, nil, func(text string) {
cfg.BotToken = strings.TrimSpace(text) cfg.BotToken = strings.TrimSpace(text)
}) })
form.AddInputField("App Token", cfg.AppToken, 128, nil, func(text string) { form.AddInputField("App Token", cfg.AppToken, 128, nil, func(text string) {
cfg.AppToken = strings.TrimSpace(text) cfg.AppToken = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) lineForm() tview.Primitive { func (s *appState) lineForm() tview.Primitive {
cfg := &s.config.Channels.LINE cfg := &s.config.Channels.LINE
form := baseChannelForm("LINE", cfg.Enabled, func(v bool) { form := baseChannelForm("LINE", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Channel Secret", cfg.ChannelSecret, 128, nil, func(text string) { form.AddInputField("Channel Secret", cfg.ChannelSecret, 128, nil, func(text string) {
cfg.ChannelSecret = strings.TrimSpace(text) cfg.ChannelSecret = strings.TrimSpace(text)
}) })
@ -381,22 +229,13 @@ func (s *appState) lineForm() tview.Primitive {
form.AddInputField("Webhook Path", cfg.WebhookPath, 64, nil, func(text string) { form.AddInputField("Webhook Path", cfg.WebhookPath, 64, nil, func(text string) {
cfg.WebhookPath = strings.TrimSpace(text) cfg.WebhookPath = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) onebotForm() tview.Primitive { func (s *appState) onebotForm() tview.Primitive {
cfg := &s.config.Channels.OneBot cfg := &s.config.Channels.OneBot
form := baseChannelForm("OneBot", cfg.Enabled, func(v bool) { form := baseChannelForm("OneBot", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("WS URL", cfg.WSUrl, 128, nil, func(text string) { form.AddInputField("WS URL", cfg.WSUrl, 128, nil, func(text string) {
cfg.WSUrl = strings.TrimSpace(text) cfg.WSUrl = strings.TrimSpace(text)
}) })
@ -418,22 +257,13 @@ func (s *appState) onebotForm() tview.Primitive {
cfg.GroupTriggerPrefix = splitCSV(text) cfg.GroupTriggerPrefix = splitCSV(text)
}, },
) )
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) wecomForm() tview.Primitive { func (s *appState) wecomForm() tview.Primitive {
cfg := &s.config.Channels.WeCom cfg := &s.config.Channels.WeCom
form := baseChannelForm("WeCom", cfg.Enabled, func(v bool) { form := baseChannelForm("WeCom", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Token", cfg.Token, 128, nil, func(text string) { form.AddInputField("Token", cfg.Token, 128, nil, func(text string) {
cfg.Token = strings.TrimSpace(text) cfg.Token = strings.TrimSpace(text)
}) })
@ -450,9 +280,7 @@ func (s *appState) wecomForm() tview.Primitive {
form.AddInputField("Webhook Path", cfg.WebhookPath, 64, nil, func(text string) { form.AddInputField("Webhook Path", cfg.WebhookPath, 64, nil, func(text string) {
cfg.WebhookPath = strings.TrimSpace(text) cfg.WebhookPath = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
addIntField( addIntField(
form, form,
"Reply Timeout", "Reply Timeout",
@ -464,14 +292,7 @@ func (s *appState) wecomForm() tview.Primitive {
func (s *appState) wecomAppForm() tview.Primitive { func (s *appState) wecomAppForm() tview.Primitive {
cfg := &s.config.Channels.WeComApp cfg := &s.config.Channels.WeComApp
form := baseChannelForm("WeCom App", cfg.Enabled, func(v bool) { form := baseChannelForm("WeCom App", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
cfg.Enabled = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
})
form.AddInputField("Corp ID", cfg.CorpID, 64, nil, func(text string) { form.AddInputField("Corp ID", cfg.CorpID, 64, nil, func(text string) {
cfg.CorpID = strings.TrimSpace(text) cfg.CorpID = strings.TrimSpace(text)
}) })
@ -492,9 +313,7 @@ func (s *appState) wecomAppForm() tview.Primitive {
form.AddInputField("Webhook Path", cfg.WebhookPath, 64, nil, func(text string) { form.AddInputField("Webhook Path", cfg.WebhookPath, 64, nil, func(text string) {
cfg.WebhookPath = strings.TrimSpace(text) cfg.WebhookPath = strings.TrimSpace(text)
}) })
form.AddInputField("Allow From", strings.Join(cfg.AllowFrom, ","), 128, nil, func(text string) { addAllowFromField(form, &cfg.AllowFrom)
cfg.AllowFrom = splitCSV(text)
})
addIntField( addIntField(
form, form,
"Reply Timeout", "Reply Timeout",
@ -504,6 +323,23 @@ func (s *appState) wecomAppForm() tview.Primitive {
return wrapWithBack(form, s) return wrapWithBack(form, s)
} }
func (s *appState) makeChannelOnEnabled(enabledPtr *bool) func(bool) {
return func(v bool) {
*enabledPtr = v
s.dirty = true
refreshMainMenuIfPresent(s)
if menu, ok := s.menus["channel"]; ok {
refreshChannelMenuFromState(menu, s)
}
}
}
func addAllowFromField(form *tview.Form, allowFrom *picoclawconfig.FlexibleStringSlice) {
form.AddInputField("Allow From", strings.Join(*allowFrom, ","), 128, nil, func(text string) {
*allowFrom = splitCSV(text)
})
}
func baseChannelForm(title string, enabled bool, onEnabled func(bool)) *tview.Form { func baseChannelForm(title string, enabled bool, onEnabled func(bool)) *tview.Form {
form := tview.NewForm() form := tview.NewForm()
form.SetBorder(true).SetTitle(fmt.Sprintf("Channel: %s", title)) form.SetBorder(true).SetTitle(fmt.Sprintf("Channel: %s", title))

View file

@ -5,6 +5,19 @@ import (
"github.com/rivo/tview" "github.com/rivo/tview"
) )
const (
colorBlue = "[#3e5db9]"
colorRed = "[#d54646]"
banner = "\r\n[::b]" +
colorBlue + "██████╗ ██╗ ██████╗ ██████╗ " + colorRed + " ██████╗██╗ █████╗ ██╗ ██╗\n" +
colorBlue + "██╔══██╗██║██╔════╝██╔═══██╗" + colorRed + "██╔════╝██║ ██╔══██╗██║ ██║\n" +
colorBlue + "██████╔╝██║██║ ██║ ██║" + colorRed + "██║ ██║ ███████║██║ █╗ ██║\n" +
colorBlue + "██╔═══╝ ██║██║ ██║ ██║" + colorRed + "██║ ██║ ██╔══██║██║███╗██║\n" +
colorBlue + "██║ ██║╚██████╗╚██████╔╝" + colorRed + "╚██████╗███████╗██║ ██║╚███╔███╔╝\n" +
colorBlue + "╚═╝ ╚═╝ ╚═════╝ ╚═════╝ " + colorRed + " ╚═════╝╚══════╝╚═╝ ╚═╝ ╚══╝╚══╝\n " +
"[:]"
)
func applyStyles() { func applyStyles() {
tview.Styles.PrimitiveBackgroundColor = tcell.NewRGBColor(12, 13, 22) tview.Styles.PrimitiveBackgroundColor = tcell.NewRGBColor(12, 13, 22)
tview.Styles.ContrastBackgroundColor = tcell.NewRGBColor(34, 19, 53) tview.Styles.ContrastBackgroundColor = tcell.NewRGBColor(34, 19, 53)
@ -24,14 +37,7 @@ func bannerView() *tview.TextView {
text.SetDynamicColors(true) text.SetDynamicColors(true)
text.SetTextAlign(tview.AlignCenter) text.SetTextAlign(tview.AlignCenter)
text.SetBackgroundColor(tview.Styles.PrimitiveBackgroundColor) text.SetBackgroundColor(tview.Styles.PrimitiveBackgroundColor)
text.SetText( text.SetText(banner)
"[::b][#84aaff]██████╗ ██╗ ██████╗ ██████╗ ██████╗██╗ █████╗ ██╗ ██╗\n" +
"[#84aaff]██╔══██╗██║██╔════╝██╔═══██╗██╔════╝██║ ██╔══██╗██║ ██║\n" +
"[#84aaff]██████╔╝██║██║ ██║ ██║██║ ██║ ███████║██║ █╗ ██║\n" +
"[#84aaff]██╔═══╝ ██║██║ ██║ ██║██║ ██║ ██╔══██║██║███╗██║\n" +
"[#84aaff]██║ ██║╚██████╗╚██████╔╝╚██████╗███████╗██║ ██║╚███╔███╔╝\n" +
"[#84aaff]╚═╝ ╚═╝ ╚═════╝ ╚═════╝ ╚═════╝╚══════╝╚═╝ ╚═╝ ╚══╝╚══╝",
)
text.SetBorder(false) text.SetBorder(false)
return text return text
} }

View file

@ -19,6 +19,9 @@ var (
) )
func GetConfigPath() string { func GetConfigPath() string {
if configPath := os.Getenv("PICOCLAW_CONFIG"); configPath != "" {
return configPath
}
home, _ := os.UserHomeDir() home, _ := os.UserHomeDir()
return filepath.Join(home, ".picoclaw", "config.json") return filepath.Join(home, ".picoclaw", "config.json")
} }

View file

@ -95,3 +95,13 @@ func TestGetConfigPath_Windows(t *testing.T) {
func TestGetVersion(t *testing.T) { func TestGetVersion(t *testing.T) {
assert.Equal(t, "dev", GetVersion()) assert.Equal(t, "dev", GetVersion())
} }
func TestGetConfigPath_WithEnv(t *testing.T) {
t.Setenv("PICOCLAW_CONFIG", "/tmp/custom/config.json")
t.Setenv("HOME", "/tmp/home") // Also set home to ensure env is preferred
got := GetConfigPath()
want := "/tmp/custom/config.json"
assert.Equal(t, want, got)
}

View file

@ -0,0 +1,25 @@
package onboard
import (
"os"
"path/filepath"
"testing"
)
func TestCopyEmbeddedToTargetUsesAgentsMarkdown(t *testing.T) {
targetDir := t.TempDir()
if err := copyEmbeddedToTarget(targetDir); err != nil {
t.Fatalf("copyEmbeddedToTarget() error = %v", err)
}
agentsPath := filepath.Join(targetDir, "AGENTS.md")
if _, err := os.Stat(agentsPath); err != nil {
t.Fatalf("expected %s to exist: %v", agentsPath, err)
}
legacyPath := filepath.Join(targetDir, "AGENT.md")
if _, err := os.Stat(legacyPath); !os.IsNotExist(err) {
t.Fatalf("expected legacy file %s to be absent, got err=%v", legacyPath, err)
}
}

View file

@ -71,7 +71,7 @@ func NewSkillsCommand() *cobra.Command {
newInstallBuiltinCommand(workspaceFn), newInstallBuiltinCommand(workspaceFn),
newListBuiltinCommand(), newListBuiltinCommand(),
newRemoveCommand(installerFn), newRemoveCommand(installerFn),
newSearchCommand(installerFn), newSearchCommand(),
newShowCommand(loaderFn), newShowCommand(loaderFn),
) )

View file

@ -15,6 +15,8 @@ import (
"github.com/sipeed/picoclaw/pkg/utils" "github.com/sipeed/picoclaw/pkg/utils"
) )
const skillsSearchMaxResults = 20
func skillsListCmd(loader *skills.SkillsLoader) { func skillsListCmd(loader *skills.SkillsLoader) {
allSkills := loader.ListSkills() allSkills := loader.ListSkills()
@ -215,34 +217,43 @@ func skillsListBuiltinCmd() {
} }
} }
func skillsSearchCmd(installer *skills.SkillInstaller) { func skillsSearchCmd(query string) {
fmt.Println("Searching for available skills...") fmt.Println("Searching for available skills...")
cfg, err := internal.LoadConfig()
if err != nil {
fmt.Printf("✗ Failed to load config: %v\n", err)
return
}
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
})
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel() defer cancel()
availableSkills, err := installer.ListAvailableSkills(ctx) results, err := registryMgr.SearchAll(ctx, query, skillsSearchMaxResults)
if err != nil { if err != nil {
fmt.Printf("✗ Failed to fetch skills list: %v\n", err) fmt.Printf("✗ Failed to fetch skills list: %v\n", err)
return return
} }
if len(availableSkills) == 0 { if len(results) == 0 {
fmt.Println("No skills available.") fmt.Println("No skills available.")
return return
} }
fmt.Printf("\nAvailable Skills (%d):\n", len(availableSkills)) fmt.Printf("\nAvailable Skills (%d):\n", len(results))
fmt.Println("--------------------") fmt.Println("--------------------")
for _, skill := range availableSkills { for _, result := range results {
fmt.Printf(" 📦 %s\n", skill.Name) fmt.Printf(" 📦 %s\n", result.DisplayName)
fmt.Printf(" %s\n", skill.Description) fmt.Printf(" %s\n", result.Summary)
fmt.Printf(" Repo: %s\n", skill.Repository) fmt.Printf(" Slug: %s\n", result.Slug)
if skill.Author != "" { fmt.Printf(" Registry: %s\n", result.RegistryName)
fmt.Printf(" Author: %s\n", skill.Author) if result.Version != "" {
} fmt.Printf(" Version: %s\n", result.Version)
if len(skill.Tags) > 0 {
fmt.Printf(" Tags: %v\n", skill.Tags)
} }
fmt.Println() fmt.Println()
} }

View file

@ -2,20 +2,19 @@ package skills
import ( import (
"github.com/spf13/cobra" "github.com/spf13/cobra"
"github.com/sipeed/picoclaw/pkg/skills"
) )
func newSearchCommand(installerFn func() (*skills.SkillInstaller, error)) *cobra.Command { func newSearchCommand() *cobra.Command {
cmd := &cobra.Command{ cmd := &cobra.Command{
Use: "search", Use: "search [query]",
Short: "Search available skills", Short: "Search available skills",
RunE: func(_ *cobra.Command, _ []string) error { Args: cobra.MaximumNArgs(1),
installer, err := installerFn() RunE: func(_ *cobra.Command, args []string) error {
if err != nil { query := ""
return err if len(args) == 1 {
query = args[0]
} }
skillsSearchCmd(installer) skillsSearchCmd(query)
return nil return nil
}, },
} }

View file

@ -8,11 +8,11 @@ import (
) )
func TestNewSearchSubcommand(t *testing.T) { func TestNewSearchSubcommand(t *testing.T) {
cmd := newSearchCommand(nil) cmd := newSearchCommand()
require.NotNil(t, cmd) require.NotNil(t, cmd)
assert.Equal(t, "search", cmd.Use) assert.Equal(t, "search [query]", cmd.Use)
assert.Equal(t, "Search available skills", cmd.Short) assert.Equal(t, "Search available skills", cmd.Short)
assert.Nil(t, cmd.Run) assert.Nil(t, cmd.Run)

View file

@ -0,0 +1,122 @@
package web
import (
"context"
"fmt"
"github.com/spf13/cobra"
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
"github.com/sipeed/picoclaw/pkg/agent"
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/providers"
pkgweb "github.com/sipeed/picoclaw/pkg/web"
)
// NewWebCommand creates the cobra command for the web management UI.
func NewWebCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "web",
Short: "启动 Web 管理界面",
Long: `启动 PicoClaw Web 配置管理界面
通过浏览器访问管理界面来配置 AI 模型消息通道和工具选项
配置保存到 config.json 需要重启相关服务才能生效
首次使用前请在 config.json 中设置 web.username web.password
{
"web": {
"host": "0.0.0.0",
"port": 18799,
"username": "admin",
"password": "your-secure-password"
}
}`,
RunE: func(cmd *cobra.Command, args []string) error {
return runWeb()
},
}
return cmd
}
func runWeb() error {
cfg, err := internal.LoadConfig()
if err != nil {
return fmt.Errorf("加载配置失败: %w", err)
}
configPath := internal.GetConfigPath()
// Warn if password not set
if cfg.Web.Password == "" {
return fmt.Errorf(
"Web 管理界面密码未配置。\n\n"+
"请在 %s 中设置:\n\n"+
" \"web\": {\n"+
" \"host\": \"0.0.0.0\",\n"+
" \"port\": 18799,\n"+
" \"username\": \"admin\",\n"+
" \"password\": \"your-secure-password\"\n"+
" }\n",
configPath,
)
}
// Use default username if empty
if cfg.Web.Username == "" {
cfg.Web.Username = "admin"
}
// Create provider for agent loop
provider, modelID, err := providers.CreateProvider(cfg)
if err != nil {
return fmt.Errorf("error creating provider: %w", err)
}
// Use the resolved model ID from provider creation
if modelID != "" {
cfg.Agents.Defaults.ModelName = modelID
}
// Create message bus and agent loop
msgBus := bus.NewMessageBus()
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
// Print agent startup info
fmt.Println("\n📦 Agent Status:")
startupInfo := agentLoop.GetStartupInfo()
toolsInfo := startupInfo["tools"].(map[string]any)
skillsInfo := startupInfo["skills"].(map[string]any)
fmt.Printf(" • Tools: %d loaded\n", toolsInfo["count"])
fmt.Printf(" • Skills: %d/%d available\n",
skillsInfo["available"],
skillsInfo["total"])
logger.InfoCF("agent", "Agent initialized",
map[string]any{
"tools_count": toolsInfo["count"],
"skills_total": skillsInfo["total"],
"skills_available": skillsInfo["available"],
})
// Start agent loop in background
ctx, cancel := context.WithCancel(context.Background())
go func() {
if err := agentLoop.Run(ctx); err != nil {
logger.ErrorCF("web", "Agent loop error", map[string]any{
"error": err.Error(),
})
}
}()
defer cancel()
// Create web server
server := pkgweb.NewServer(cfg, configPath)
// Inject agent loop and message bus into web server
server.SetAgentLoop(agentLoop, msgBus)
return server.Start()
}

View file

@ -22,6 +22,7 @@ import (
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/skills" "github.com/sipeed/picoclaw/cmd/picoclaw/internal/skills"
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/status" "github.com/sipeed/picoclaw/cmd/picoclaw/internal/status"
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/version" "github.com/sipeed/picoclaw/cmd/picoclaw/internal/version"
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/web"
) )
func NewPicoclawCommand() *cobra.Command { func NewPicoclawCommand() *cobra.Command {
@ -43,12 +44,27 @@ func NewPicoclawCommand() *cobra.Command {
migrate.NewMigrateCommand(), migrate.NewMigrateCommand(),
skills.NewSkillsCommand(), skills.NewSkillsCommand(),
version.NewVersionCommand(), version.NewVersionCommand(),
web.NewWebCommand(),
) )
return cmd return cmd
} }
const (
colorBlue = "\033[1;38;2;62;93;185m"
colorRed = "\033[1;38;2;213;70;70m"
banner = "\r\n" +
colorBlue + "██████╗ ██╗ ██████╗ ██████╗ " + colorRed + " ██████╗██╗ █████╗ ██╗ ██╗\n" +
colorBlue + "██╔══██╗██║██╔════╝██╔═══██╗" + colorRed + "██╔════╝██║ ██╔══██╗██║ ██║\n" +
colorBlue + "██████╔╝██║██║ ██║ ██║" + colorRed + "██║ ██║ ███████║██║ █╗ ██║\n" +
colorBlue + "██╔═══╝ ██║██║ ██║ ██║" + colorRed + "██║ ██║ ██╔══██║██║███╗██║\n" +
colorBlue + "██║ ██║╚██████╗╚██████╔╝" + colorRed + "╚██████╗███████╗██║ ██║╚███╔███╔╝\n" +
colorBlue + "╚═╝ ╚═╝ ╚═════╝ ╚═════╝ " + colorRed + " ╚═════╝╚══════╝╚═╝ ╚═╝ ╚══╝╚══╝\n " +
"\033[0m\r\n"
)
func main() { func main() {
fmt.Printf("%s", banner)
cmd := NewPicoclawCommand() cmd := NewPicoclawCommand()
if err := cmd.Execute(); err != nil { if err := cmd.Execute(); err != nil {
os.Exit(1) os.Exit(1)

View file

@ -49,6 +49,7 @@
"telegram": { "telegram": {
"enabled": false, "enabled": false,
"token": "YOUR_TELEGRAM_BOT_TOKEN", "token": "YOUR_TELEGRAM_BOT_TOKEN",
"base_url": "",
"proxy": "", "proxy": "",
"allow_from": [ "allow_from": [
"YOUR_USER_ID" "YOUR_USER_ID"
@ -58,8 +59,11 @@
"discord": { "discord": {
"enabled": false, "enabled": false,
"token": "YOUR_DISCORD_BOT_TOKEN", "token": "YOUR_DISCORD_BOT_TOKEN",
"proxy": "",
"allow_from": [], "allow_from": [],
"mention_only": false, "group_trigger": {
"mention_only": false
},
"reasoning_channel_id": "" "reasoning_channel_id": ""
}, },
"qq": { "qq": {
@ -111,8 +115,6 @@
"enabled": false, "enabled": false,
"channel_secret": "YOUR_LINE_CHANNEL_SECRET", "channel_secret": "YOUR_LINE_CHANNEL_SECRET",
"channel_access_token": "YOUR_LINE_CHANNEL_ACCESS_TOKEN", "channel_access_token": "YOUR_LINE_CHANNEL_ACCESS_TOKEN",
"webhook_host": "0.0.0.0",
"webhook_port": 18791,
"webhook_path": "/webhook/line", "webhook_path": "/webhook/line",
"allow_from": [], "allow_from": [],
"reasoning_channel_id": "" "reasoning_channel_id": ""
@ -127,32 +129,38 @@
"reasoning_channel_id": "" "reasoning_channel_id": ""
}, },
"wecom": { "wecom": {
"_comment": "WeCom Bot (智能机器人) - Easier setup, supports group chats", "_comment": "WeCom Bot - Easier setup, supports group chats",
"enabled": false, "enabled": false,
"token": "YOUR_TOKEN", "token": "YOUR_TOKEN",
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY", "encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
"webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=YOUR_KEY", "webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=YOUR_KEY",
"webhook_host": "0.0.0.0",
"webhook_port": 18793,
"webhook_path": "/webhook/wecom", "webhook_path": "/webhook/wecom",
"allow_from": [], "allow_from": [],
"reply_timeout": 5, "reply_timeout": 5,
"reasoning_channel_id": "" "reasoning_channel_id": ""
}, },
"wecom_app": { "wecom_app": {
"_comment": "WeCom App (自建应用) - More features, proactive messaging, private chat only. See docs/wecom-app-configuration.md", "_comment": "WeCom App (自建应用) - More features, proactive messaging, private chat only.",
"enabled": false, "enabled": false,
"corp_id": "YOUR_CORP_ID", "corp_id": "YOUR_CORP_ID",
"corp_secret": "YOUR_CORP_SECRET", "corp_secret": "YOUR_CORP_SECRET",
"agent_id": 1000002, "agent_id": 1000002,
"token": "YOUR_TOKEN", "token": "YOUR_TOKEN",
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY", "encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
"webhook_host": "0.0.0.0",
"webhook_port": 18792,
"webhook_path": "/webhook/wecom-app", "webhook_path": "/webhook/wecom-app",
"allow_from": [], "allow_from": [],
"reply_timeout": 5, "reply_timeout": 5,
"reasoning_channel_id": "" "reasoning_channel_id": ""
},
"wecom_aibot": {
"_comment": "WeCom AI Bot (智能机器人) - Official WeCom AI Bot integration, supports proactive messaging and private chats.",
"enabled": false,
"token": "YOUR_TOKEN",
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
"webhook_path": "/webhook/wecom-aibot",
"max_steps": 10,
"welcome_message": "Hello! I'm your AI assistant. How can I help you today?",
"reasoning_channel_id": ""
} }
}, },
"providers": { "providers": {
@ -237,6 +245,71 @@
"cron": { "cron": {
"exec_timeout_minutes": 5 "exec_timeout_minutes": 5
}, },
"mcp": {
"enabled": false,
"servers": {
"context7": {
"enabled": false,
"type": "http",
"url": "https://mcp.context7.com/mcp",
"headers": {
"CONTEXT7_API_KEY": "ctx7sk-xx"
}
},
"filesystem": {
"enabled": false,
"command": "npx",
"args": [
"-y",
"@modelcontextprotocol/server-filesystem",
"/tmp"
]
},
"github": {
"enabled": false,
"command": "npx",
"args": [
"-y",
"@modelcontextprotocol/server-github"
],
"env": {
"GITHUB_PERSONAL_ACCESS_TOKEN": "YOUR_GITHUB_TOKEN"
}
},
"brave-search": {
"enabled": false,
"command": "npx",
"args": [
"-y",
"@modelcontextprotocol/server-brave-search"
],
"env": {
"BRAVE_API_KEY": "YOUR_BRAVE_API_KEY"
}
},
"postgres": {
"enabled": false,
"command": "npx",
"args": [
"-y",
"@modelcontextprotocol/server-postgres",
"postgresql://user:password@localhost/dbname"
]
},
"slack": {
"enabled": false,
"command": "npx",
"args": [
"-y",
"@modelcontextprotocol/server-slack"
],
"env": {
"SLACK_BOT_TOKEN": "YOUR_SLACK_BOT_TOKEN",
"SLACK_TEAM_ID": "YOUR_SLACK_TEAM_ID"
}
}
}
},
"exec": { "exec": {
"enable_deny_patterns": false, "enable_deny_patterns": false,
"custom_deny_patterns": [] "custom_deny_patterns": []
@ -265,4 +338,4 @@
"host": "127.0.0.1", "host": "127.0.0.1",
"port": 18790 "port": 18790
} }
} }

44
docker/Dockerfile.full Normal file
View file

@ -0,0 +1,44 @@
# ============================================================
# Stage 1: Build the picoclaw binary
# ============================================================
FROM golang:1.26.0-alpine AS builder
RUN apk add --no-cache git make
WORKDIR /src
# Cache dependencies
COPY go.mod go.sum ./
RUN go mod download
# Copy source and build
COPY . .
RUN make build
# ============================================================
# Stage 2: Node.js-based runtime with full MCP support
# ============================================================
FROM node:24-alpine3.23
# Install runtime dependencies
RUN apk add --no-cache \
ca-certificates \
curl \
git \
python3 \
py3-pip
# Install uv and symlink to system path
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
ln -s /root/.local/bin/uv /usr/local/bin/uv && \
ln -s /root/.local/bin/uvx /usr/local/bin/uvx && \
uv --version
# Copy binary
COPY --from=builder /src/build/picoclaw /usr/local/bin/picoclaw
# Create picoclaw home directory
RUN /usr/local/bin/picoclaw onboard
ENTRYPOINT ["picoclaw"]
CMD ["gateway"]

View file

@ -0,0 +1,44 @@
services:
# ─────────────────────────────────────────────
# PicoClaw Agent (one-shot query) - Full MCP Support
# docker compose -f docker/docker-compose.full.yml run --rm picoclaw-agent -m "Hello"
# ─────────────────────────────────────────────
picoclaw-agent:
build:
context: ..
dockerfile: docker/Dockerfile.full
container_name: picoclaw-agent-full
profiles:
- agent
volumes:
- ../config/config.json:/root/.picoclaw/config.json:ro
- picoclaw-workspace:/root/.picoclaw/workspace
- picoclaw-npm-cache:/root/.npm # npm cache for faster MCP server installs
entrypoint: ["picoclaw", "agent"]
stdin_open: true
tty: true
# ─────────────────────────────────────────────
# PicoClaw Gateway (Long-running Bot) - Full MCP Support
# docker compose -f docker/docker-compose.full.yml --profile gateway up
# ─────────────────────────────────────────────
picoclaw-gateway:
build:
context: ..
dockerfile: docker/Dockerfile.full
container_name: picoclaw-gateway-full
restart: unless-stopped
profiles:
- gateway
volumes:
# Configuration file
- ../config/config.json:/root/.picoclaw/config.json:ro
# Persistent workspace (sessions, memory, logs)
- picoclaw-workspace:/root/.picoclaw/workspace
# NPM cache for faster MCP server installs
- picoclaw-npm-cache:/root/.npm
command: ["gateway"]
volumes:
picoclaw-workspace:
picoclaw-npm-cache: # Cache npm packages to speed up MCP server installations

View file

@ -7,6 +7,7 @@ import (
"os" "os"
"path/filepath" "path/filepath"
"runtime" "runtime"
"slices"
"strings" "strings"
"sync" "sync"
"time" "time"
@ -33,6 +34,11 @@ type ContextBuilder struct {
// created (didn't exist at cache time, now exist) or deleted (existed at // created (didn't exist at cache time, now exist) or deleted (existed at
// cache time, now gone) — both of which should trigger a cache rebuild. // cache time, now gone) — both of which should trigger a cache rebuild.
existedAtCache map[string]bool existedAtCache map[string]bool
// skillFilesAtCache snapshots the skill tree file set and mtimes at cache
// build time. This catches nested file creations/deletions/mtime changes
// that may not update the top-level skill root directory mtime.
skillFilesAtCache map[string]time.Time
} }
func getGlobalConfigDir() string { func getGlobalConfigDir() string {
@ -46,8 +52,11 @@ func getGlobalConfigDir() string {
func NewContextBuilder(workspace string) *ContextBuilder { func NewContextBuilder(workspace string) *ContextBuilder {
// builtin skills: skills directory in current project // builtin skills: skills directory in current project
// Use the skills/ directory under the current working directory // Use the skills/ directory under the current working directory
wd, _ := os.Getwd() builtinSkillsDir := strings.TrimSpace(os.Getenv("PICOCLAW_BUILTIN_SKILLS"))
builtinSkillsDir := filepath.Join(wd, "skills") if builtinSkillsDir == "" {
wd, _ := os.Getwd()
builtinSkillsDir = filepath.Join(wd, "skills")
}
globalSkillsDir := filepath.Join(getGlobalConfigDir(), "skills") globalSkillsDir := filepath.Join(getGlobalConfigDir(), "skills")
return &ContextBuilder{ return &ContextBuilder{
@ -57,12 +66,17 @@ func NewContextBuilder(workspace string) *ContextBuilder {
} }
} }
// Memory returns the memory store for accessing long-term and daily memories.
func (cb *ContextBuilder) Memory() *MemoryStore {
return cb.memory
}
func (cb *ContextBuilder) getIdentity() string { func (cb *ContextBuilder) getIdentity() string {
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace)) workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
return fmt.Sprintf(`# picoclaw 🦞 return fmt.Sprintf(`# zhuiai claw 🦞
You are picoclaw, a helpful AI assistant. You are zhuiai claw, a helpful AI assistant.
## Workspace ## Workspace
Your workspace is at: %s Your workspace is at: %s
@ -76,7 +90,45 @@ Your workspace is at: %s
2. **Be helpful and accurate** - When using tools, briefly explain what you're doing. 2. **Be helpful and accurate** - When using tools, briefly explain what you're doing.
3. **Memory** - When interacting with me if something seems memorable, update %s/memory/MEMORY.md 3. **Memory Management** - You have dedicated memory tools:
- **update_memory**: Save important information immediately
- **search_memory**: Find existing memories by keywords
### When to use update_memory:
User mentions personal info (name, job, location)
User expresses preferences (language, timezone, habits)
Important events or deadlines mentioned
Task completions worth remembering
### Examples:
- User: "我下周要去北京出差" Call update_memory with category="important_note"
- User: "我对海鲜过敏" Call update_memory with category="preference"
- User: "今天完成了项目部署" Call update_memory with memory_type="daily_note"
4. **Task Management** - Use **manage_tasks** tool for schedules and reminders:
- **add**: Add a new task with due date/time and repeat pattern
- **list**: List tasks (scope: all, today, upcoming, completed)
- **complete**: Mark a task as completed
- **delete**: Delete a task
- **get_today**: Quick way to get today's tasks
### When to use manage_tasks:
User mentions schedules, plans, or to-dos
User asks "今天有什么任务" or "我的工作计划是什么"
User wants to track recurring tasks
### Task Parameters:
- **due_date**: YYYY-MM-DD format (e.g., "2026-02-03")
- **due_weekday**: 1=Monday, 2=Tuesday, ..., 7=Sunday
- **repeat**: none, daily, weekly, monthly, weekdays
### Examples:
- User: "帮我记一下周1去吴晓客户那里周2去小明家吃饭"
Call manage_tasks with action="add", title="去吴晓客户那里", due_weekday=1, repeat="weekly"
- User: "今天我要做什么?"
Call manage_tasks with action="get_today"
- User: "完成了上面的任务"
Call manage_tasks with action="complete", task_id="[from previous list]"
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.`, 4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.`,
workspacePath, workspacePath, workspacePath, workspacePath, workspacePath) workspacePath, workspacePath, workspacePath, workspacePath, workspacePath)
@ -147,6 +199,7 @@ func (cb *ContextBuilder) BuildSystemPromptWithCache() string {
cb.cachedSystemPrompt = prompt cb.cachedSystemPrompt = prompt
cb.cachedAt = baseline.maxMtime cb.cachedAt = baseline.maxMtime
cb.existedAtCache = baseline.existed cb.existedAtCache = baseline.existed
cb.skillFilesAtCache = baseline.skillFiles
logger.DebugCF("agent", "System prompt cached", logger.DebugCF("agent", "System prompt cached",
map[string]any{ map[string]any{
@ -166,14 +219,14 @@ func (cb *ContextBuilder) InvalidateCache() {
cb.cachedSystemPrompt = "" cb.cachedSystemPrompt = ""
cb.cachedAt = time.Time{} cb.cachedAt = time.Time{}
cb.existedAtCache = nil cb.existedAtCache = nil
cb.skillFilesAtCache = nil
logger.DebugCF("agent", "System prompt cache invalidated", nil) logger.DebugCF("agent", "System prompt cache invalidated", nil)
} }
// sourcePaths returns the workspace source file paths tracked for cache // sourcePaths returns non-skill workspace source files tracked for cache
// invalidation (bootstrap files + memory). The skills directory is handled // invalidation (bootstrap files + memory). Skill roots are handled separately
// separately in sourceFilesChangedLocked because it requires both directory- // because they require both directory-level and recursive file-level checks.
// level and recursive file-level mtime checks.
func (cb *ContextBuilder) sourcePaths() []string { func (cb *ContextBuilder) sourcePaths() []string {
return []string{ return []string{
filepath.Join(cb.workspace, "AGENTS.md"), filepath.Join(cb.workspace, "AGENTS.md"),
@ -184,23 +237,39 @@ func (cb *ContextBuilder) sourcePaths() []string {
} }
} }
// skillRoots returns all skill root directories that can affect
// BuildSkillsSummary output (workspace/global/builtin).
func (cb *ContextBuilder) skillRoots() []string {
if cb.skillsLoader == nil {
return []string{filepath.Join(cb.workspace, "skills")}
}
roots := cb.skillsLoader.SkillRoots()
if len(roots) == 0 {
return []string{filepath.Join(cb.workspace, "skills")}
}
return roots
}
// cacheBaseline holds the file existence snapshot and the latest observed // cacheBaseline holds the file existence snapshot and the latest observed
// mtime across all tracked paths. Used as the cache reference point. // mtime across all tracked paths. Used as the cache reference point.
type cacheBaseline struct { type cacheBaseline struct {
existed map[string]bool existed map[string]bool
maxMtime time.Time skillFiles map[string]time.Time
maxMtime time.Time
} }
// buildCacheBaseline records which tracked paths currently exist and computes // buildCacheBaseline records which tracked paths currently exist and computes
// the latest mtime across all tracked files + skills directory contents. // the latest mtime across all tracked files + skills directory contents.
// Called under write lock when the cache is built. // Called under write lock when the cache is built.
func (cb *ContextBuilder) buildCacheBaseline() cacheBaseline { func (cb *ContextBuilder) buildCacheBaseline() cacheBaseline {
skillsDir := filepath.Join(cb.workspace, "skills") skillRoots := cb.skillRoots()
// All paths whose existence we track: source files + skills dir. // All paths whose existence we track: source files + all skill roots.
allPaths := append(cb.sourcePaths(), skillsDir) allPaths := append(cb.sourcePaths(), skillRoots...)
existed := make(map[string]bool, len(allPaths)) existed := make(map[string]bool, len(allPaths))
skillFiles := make(map[string]time.Time)
var maxMtime time.Time var maxMtime time.Time
for _, p := range allPaths { for _, p := range allPaths {
@ -211,17 +280,21 @@ func (cb *ContextBuilder) buildCacheBaseline() cacheBaseline {
} }
} }
// Walk skills files to capture their mtimes too. // Walk all skill roots recursively to snapshot skill files and mtimes.
// Use os.Stat (not d.Info) to match the stat method used in // Use os.Stat (not d.Info) for consistency with sourceFilesChanged checks.
// fileChangedSince / skillFilesModifiedSince for consistency. for _, root := range skillRoots {
_ = filepath.WalkDir(skillsDir, func(path string, d fs.DirEntry, walkErr error) error { _ = filepath.WalkDir(root, func(path string, d fs.DirEntry, walkErr error) error {
if walkErr == nil && !d.IsDir() { if walkErr == nil && !d.IsDir() {
if info, err := os.Stat(path); err == nil && info.ModTime().After(maxMtime) { if info, err := os.Stat(path); err == nil {
maxMtime = info.ModTime() skillFiles[path] = info.ModTime()
if info.ModTime().After(maxMtime) {
maxMtime = info.ModTime()
}
}
} }
} return nil
return nil })
}) }
// If no tracked files exist yet (empty workspace), maxMtime is zero. // If no tracked files exist yet (empty workspace), maxMtime is zero.
// Use a very old non-zero time so that: // Use a very old non-zero time so that:
@ -233,7 +306,7 @@ func (cb *ContextBuilder) buildCacheBaseline() cacheBaseline {
maxMtime = time.Unix(1, 0) maxMtime = time.Unix(1, 0)
} }
return cacheBaseline{existed: existed, maxMtime: maxMtime} return cacheBaseline{existed: existed, skillFiles: skillFiles, maxMtime: maxMtime}
} }
// sourceFilesChangedLocked checks whether any workspace source file has been // sourceFilesChangedLocked checks whether any workspace source file has been
@ -249,27 +322,21 @@ func (cb *ContextBuilder) sourceFilesChangedLocked() bool {
} }
// Check tracked source files (bootstrap + memory). // Check tracked source files (bootstrap + memory).
for _, p := range cb.sourcePaths() { if slices.ContainsFunc(cb.sourcePaths(), cb.fileChangedSince) {
if cb.fileChangedSince(p) {
return true
}
}
// --- Skills directory (handled separately from sourcePaths) ---
//
// 1. Creation/deletion: tracked via existedAtCache, same as bootstrap files.
skillsDir := filepath.Join(cb.workspace, "skills")
if cb.fileChangedSince(skillsDir) {
return true return true
} }
// 2. Structural changes (add/remove entries inside the dir) are reflected // --- Skill roots (workspace/global/builtin) ---
// in the directory's own mtime, which fileChangedSince already checks.
// //
// 3. Content-only edits to files inside skills/ do NOT update the parent // For each root:
// directory mtime on most filesystems, so we recursively walk to check // 1. Creation/deletion and root directory mtime changes are tracked by fileChangedSince.
// individual file mtimes at any nesting depth. // 2. Nested file create/delete/mtime changes are tracked by the skill file snapshot.
if skillFilesModifiedSince(skillsDir, cb.cachedAt) { for _, root := range cb.skillRoots() {
if cb.fileChangedSince(root) {
return true
}
}
if skillFilesChangedSince(cb.skillRoots(), cb.skillFilesAtCache) {
return true return true
} }
@ -310,28 +377,64 @@ func (cb *ContextBuilder) fileChangedSince(path string) bool {
// if the callback returned nil when its err parameter is non-nil. // if the callback returned nil when its err parameter is non-nil.
var errWalkStop = errors.New("walk stop") var errWalkStop = errors.New("walk stop")
// skillFilesModifiedSince recursively walks the skills directory and checks // skillFilesChangedSince compares the current recursive skill file tree
// whether any file was modified after t. This catches content-only edits at // against the cache-time snapshot. Any create/delete/mtime drift invalidates
// any nesting depth (e.g. skills/name/docs/extra.md) that don't update // the cache.
// parent directory mtimes. func skillFilesChangedSince(skillRoots []string, filesAtCache map[string]time.Time) bool {
func skillFilesModifiedSince(skillsDir string, t time.Time) bool { // Defensive: if the snapshot was never initialized, force rebuild.
changed := false if filesAtCache == nil {
err := filepath.WalkDir(skillsDir, func(path string, d fs.DirEntry, walkErr error) error { return true
if walkErr == nil && !d.IsDir() {
if info, statErr := os.Stat(path); statErr == nil && info.ModTime().After(t) {
changed = true
return errWalkStop // stop walking
}
}
return nil
})
// errWalkStop is expected (early exit on first changed file).
// os.IsNotExist means the skills dir doesn't exist yet — not an error.
// Any other error is unexpected and worth logging.
if err != nil && !errors.Is(err, errWalkStop) && !os.IsNotExist(err) {
logger.DebugCF("agent", "skills walk error", map[string]any{"error": err.Error()})
} }
return changed
// Check cached files still exist and keep the same mtime.
for path, cachedMtime := range filesAtCache {
info, err := os.Stat(path)
if err != nil {
// A previously tracked file disappeared (or became inaccessible):
// either way, cached skill summary may now be stale.
return true
}
if !info.ModTime().Equal(cachedMtime) {
return true
}
}
// Check no new files appeared under any skill root.
changed := false
for _, root := range skillRoots {
if strings.TrimSpace(root) == "" {
continue
}
err := filepath.WalkDir(root, func(path string, d fs.DirEntry, walkErr error) error {
if walkErr != nil {
// Treat unexpected walk errors as changed to avoid stale cache.
if !os.IsNotExist(walkErr) {
changed = true
return errWalkStop
}
return nil
}
if d.IsDir() {
return nil
}
if _, ok := filesAtCache[path]; !ok {
changed = true
return errWalkStop
}
return nil
})
if changed {
return true
}
if err != nil && !errors.Is(err, errWalkStop) && !os.IsNotExist(err) {
logger.DebugCF("agent", "skills walk error", map[string]any{"error": err.Error()})
return true
}
}
return false
} }
func (cb *ContextBuilder) LoadBootstrapFiles() string { func (cb *ContextBuilder) LoadBootstrapFiles() string {
@ -467,10 +570,14 @@ func (cb *ContextBuilder) BuildMessages(
// Add current user message // Add current user message
if strings.TrimSpace(currentMessage) != "" { if strings.TrimSpace(currentMessage) != "" {
messages = append(messages, providers.Message{ msg := providers.Message{
Role: "user", Role: "user",
Content: currentMessage, Content: currentMessage,
}) }
if len(media) > 0 {
msg.Media = media
}
messages = append(messages, msg)
} }
return messages return messages

View file

@ -383,6 +383,162 @@ Updated content.`
} }
} }
// TestGlobalSkillFileContentChange verifies that modifying a global skill
// (~/.picoclaw/skills) invalidates the cached system prompt.
func TestGlobalSkillFileContentChange(t *testing.T) {
tmpHome := t.TempDir()
t.Setenv("HOME", tmpHome)
tmpDir := setupWorkspace(t, nil)
defer os.RemoveAll(tmpDir)
globalSkillPath := filepath.Join(tmpHome, ".picoclaw", "skills", "global-skill", "SKILL.md")
if err := os.MkdirAll(filepath.Dir(globalSkillPath), 0o755); err != nil {
t.Fatal(err)
}
v1 := `---
name: global-skill
description: global-v1
---
# Global Skill v1`
if err := os.WriteFile(globalSkillPath, []byte(v1), 0o644); err != nil {
t.Fatal(err)
}
cb := NewContextBuilder(tmpDir)
sp1 := cb.BuildSystemPromptWithCache()
if !strings.Contains(sp1, "global-v1") {
t.Fatal("expected initial prompt to contain global skill description")
}
v2 := `---
name: global-skill
description: global-v2
---
# Global Skill v2`
if err := os.WriteFile(globalSkillPath, []byte(v2), 0o644); err != nil {
t.Fatal(err)
}
future := time.Now().Add(2 * time.Second)
if err := os.Chtimes(globalSkillPath, future, future); err != nil {
t.Fatalf("failed to update mtime for %s: %v", globalSkillPath, err)
}
cb.systemPromptMutex.RLock()
changed := cb.sourceFilesChangedLocked()
cb.systemPromptMutex.RUnlock()
if !changed {
t.Fatal("sourceFilesChangedLocked() should detect global skill file content change")
}
sp2 := cb.BuildSystemPromptWithCache()
if !strings.Contains(sp2, "global-v2") {
t.Error("rebuilt prompt should contain updated global skill description")
}
if sp1 == sp2 {
t.Error("cache should be invalidated when global skill file content changes")
}
}
// TestBuiltinSkillFileContentChange verifies that modifying a builtin skill
// invalidates the cached system prompt.
func TestBuiltinSkillFileContentChange(t *testing.T) {
tmpHome := t.TempDir()
t.Setenv("HOME", tmpHome)
tmpDir := setupWorkspace(t, nil)
defer os.RemoveAll(tmpDir)
builtinRoot := t.TempDir()
t.Setenv("PICOCLAW_BUILTIN_SKILLS", builtinRoot)
builtinSkillPath := filepath.Join(builtinRoot, "builtin-skill", "SKILL.md")
if err := os.MkdirAll(filepath.Dir(builtinSkillPath), 0o755); err != nil {
t.Fatal(err)
}
v1 := `---
name: builtin-skill
description: builtin-v1
---
# Builtin Skill v1`
if err := os.WriteFile(builtinSkillPath, []byte(v1), 0o644); err != nil {
t.Fatal(err)
}
cb := NewContextBuilder(tmpDir)
sp1 := cb.BuildSystemPromptWithCache()
if !strings.Contains(sp1, "builtin-v1") {
t.Fatal("expected initial prompt to contain builtin skill description")
}
v2 := `---
name: builtin-skill
description: builtin-v2
---
# Builtin Skill v2`
if err := os.WriteFile(builtinSkillPath, []byte(v2), 0o644); err != nil {
t.Fatal(err)
}
future := time.Now().Add(2 * time.Second)
if err := os.Chtimes(builtinSkillPath, future, future); err != nil {
t.Fatalf("failed to update mtime for %s: %v", builtinSkillPath, err)
}
cb.systemPromptMutex.RLock()
changed := cb.sourceFilesChangedLocked()
cb.systemPromptMutex.RUnlock()
if !changed {
t.Fatal("sourceFilesChangedLocked() should detect builtin skill file content change")
}
sp2 := cb.BuildSystemPromptWithCache()
if !strings.Contains(sp2, "builtin-v2") {
t.Error("rebuilt prompt should contain updated builtin skill description")
}
if sp1 == sp2 {
t.Error("cache should be invalidated when builtin skill file content changes")
}
}
// TestSkillFileDeletionInvalidatesCache verifies that deleting a nested skill
// file invalidates the cached system prompt.
func TestSkillFileDeletionInvalidatesCache(t *testing.T) {
tmpDir := setupWorkspace(t, map[string]string{
"skills/delete-me/SKILL.md": `---
name: delete-me
description: delete-me-v1
---
# Delete Me`,
})
defer os.RemoveAll(tmpDir)
cb := NewContextBuilder(tmpDir)
sp1 := cb.BuildSystemPromptWithCache()
if !strings.Contains(sp1, "delete-me-v1") {
t.Fatal("expected initial prompt to contain skill description")
}
skillPath := filepath.Join(tmpDir, "skills", "delete-me", "SKILL.md")
if err := os.Remove(skillPath); err != nil {
t.Fatal(err)
}
cb.systemPromptMutex.RLock()
changed := cb.sourceFilesChangedLocked()
cb.systemPromptMutex.RUnlock()
if !changed {
t.Fatal("sourceFilesChangedLocked() should detect deleted skill file")
}
sp2 := cb.BuildSystemPromptWithCache()
if strings.Contains(sp2, "delete-me-v1") {
t.Error("rebuilt prompt should not contain deleted skill description")
}
if sp1 == sp2 {
t.Error("cache should be invalidated when skill file is deleted")
}
}
// TestConcurrentBuildSystemPromptWithCache verifies that multiple goroutines // TestConcurrentBuildSystemPromptWithCache verifies that multiple goroutines
// can safely call BuildSystemPromptWithCache concurrently without producing // can safely call BuildSystemPromptWithCache concurrently without producing
// empty results, panics, or data races. // empty results, panics, or data races.
@ -404,11 +560,11 @@ func TestConcurrentBuildSystemPromptWithCache(t *testing.T) {
var wg sync.WaitGroup var wg sync.WaitGroup
errs := make(chan string, goroutines*iterations) errs := make(chan string, goroutines*iterations)
for g := 0; g < goroutines; g++ { for g := range goroutines {
wg.Add(1) wg.Add(1)
go func(id int) { go func(id int) {
defer wg.Done() defer wg.Done()
for i := 0; i < iterations; i++ { for i := range iterations {
result := cb.BuildSystemPromptWithCache() result := cb.BuildSystemPromptWithCache()
if result == "" { if result == "" {
errs <- "empty prompt returned" errs <- "empty prompt returned"

View file

@ -1,9 +1,11 @@
package agent package agent
import ( import (
"fmt"
"log" "log"
"os" "os"
"path/filepath" "path/filepath"
"regexp"
"strings" "strings"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
@ -16,22 +18,24 @@ import (
// AgentInstance represents a fully configured agent with its own workspace, // AgentInstance represents a fully configured agent with its own workspace,
// session manager, context builder, and tool registry. // session manager, context builder, and tool registry.
type AgentInstance struct { type AgentInstance struct {
ID string ID string
Name string Name string
Model string Model string
Fallbacks []string Fallbacks []string
Workspace string Workspace string
MaxIterations int MaxIterations int
MaxTokens int MaxTokens int
Temperature float64 Temperature float64
ContextWindow int ContextWindow int
Provider providers.LLMProvider SummarizeMessageThreshold int
Sessions *session.SessionManager SummarizeTokenPercent int
ContextBuilder *ContextBuilder Provider providers.LLMProvider
Tools *tools.ToolRegistry Sessions *session.SessionManager
Subagents *config.SubagentsConfig ContextBuilder *ContextBuilder
SkillsFilter []string Tools *tools.ToolRegistry
Candidates []providers.FallbackCandidate Subagents *config.SubagentsConfig
SkillsFilter []string
Candidates []providers.FallbackCandidate
} }
// NewAgentInstance creates an agent instance from config. // NewAgentInstance creates an agent instance from config.
@ -48,18 +52,33 @@ func NewAgentInstance(
fallbacks := resolveAgentFallbacks(agentCfg, defaults) fallbacks := resolveAgentFallbacks(agentCfg, defaults)
restrict := defaults.RestrictToWorkspace restrict := defaults.RestrictToWorkspace
readRestrict := restrict && !defaults.AllowReadOutsideWorkspace
// Compile path whitelist patterns from config.
allowReadPaths := compilePatterns(cfg.Tools.AllowReadPaths)
allowWritePaths := compilePatterns(cfg.Tools.AllowWritePaths)
toolsRegistry := tools.NewToolRegistry() toolsRegistry := tools.NewToolRegistry()
toolsRegistry.Register(tools.NewReadFileTool(workspace, restrict)) toolsRegistry.Register(tools.NewReadFileTool(workspace, readRestrict, allowReadPaths))
toolsRegistry.Register(tools.NewWriteFileTool(workspace, restrict)) toolsRegistry.Register(tools.NewWriteFileTool(workspace, restrict, allowWritePaths))
toolsRegistry.Register(tools.NewListDirTool(workspace, restrict)) toolsRegistry.Register(tools.NewListDirTool(workspace, readRestrict, allowReadPaths))
execTool, err := tools.NewExecToolWithConfig(workspace, restrict, cfg) execTool, err := tools.NewExecToolWithConfig(workspace, restrict, cfg)
if err != nil { if err != nil {
log.Fatalf("Critical error: unable to initialize exec tool: %v", err) log.Fatalf("Critical error: unable to initialize exec tool: %v", err)
} }
toolsRegistry.Register(execTool) toolsRegistry.Register(execTool)
toolsRegistry.Register(tools.NewEditFileTool(workspace, restrict)) toolsRegistry.Register(tools.NewEditFileTool(workspace, restrict, allowWritePaths))
toolsRegistry.Register(tools.NewAppendFileTool(workspace, restrict)) toolsRegistry.Register(tools.NewAppendFileTool(workspace, restrict, allowWritePaths))
// Memory tools - create MemoryStore directly in tools package
memoryStore := tools.NewMemoryStore(workspace)
toolsRegistry.Register(tools.NewUpdateMemoryTool(memoryStore))
toolsRegistry.Register(tools.NewSearchMemoryTool(memoryStore))
// Task management tools
taskStore := tools.NewTaskStore(workspace)
toolsRegistry.Register(tools.NewTaskTool(taskStore))
sessionsDir := filepath.Join(workspace, "sessions") sessionsDir := filepath.Join(workspace, "sessions")
sessionsManager := session.NewSessionManager(sessionsDir) sessionsManager := session.NewSessionManager(sessionsDir)
@ -93,6 +112,16 @@ func NewAgentInstance(
temperature = *defaults.Temperature temperature = *defaults.Temperature
} }
summarizeMessageThreshold := defaults.SummarizeMessageThreshold
if summarizeMessageThreshold == 0 {
summarizeMessageThreshold = 20
}
summarizeTokenPercent := defaults.SummarizeTokenPercent
if summarizeTokenPercent == 0 {
summarizeTokenPercent = 75
}
// Resolve fallback candidates // Resolve fallback candidates
modelCfg := providers.ModelConfig{ modelCfg := providers.ModelConfig{
Primary: model, Primary: model,
@ -141,22 +170,24 @@ func NewAgentInstance(
candidates := providers.ResolveCandidatesWithLookup(modelCfg, defaults.Provider, resolveFromModelList) candidates := providers.ResolveCandidatesWithLookup(modelCfg, defaults.Provider, resolveFromModelList)
return &AgentInstance{ return &AgentInstance{
ID: agentID, ID: agentID,
Name: agentName, Name: agentName,
Model: model, Model: model,
Fallbacks: fallbacks, Fallbacks: fallbacks,
Workspace: workspace, Workspace: workspace,
MaxIterations: maxIter, MaxIterations: maxIter,
MaxTokens: maxTokens, MaxTokens: maxTokens,
Temperature: temperature, Temperature: temperature,
ContextWindow: maxTokens, ContextWindow: maxTokens,
Provider: provider, SummarizeMessageThreshold: summarizeMessageThreshold,
Sessions: sessionsManager, SummarizeTokenPercent: summarizeTokenPercent,
ContextBuilder: contextBuilder, Provider: provider,
Tools: toolsRegistry, Sessions: sessionsManager,
Subagents: subagents, ContextBuilder: contextBuilder,
SkillsFilter: skillsFilter, Tools: toolsRegistry,
Candidates: candidates, Subagents: subagents,
SkillsFilter: skillsFilter,
Candidates: candidates,
} }
} }
@ -189,6 +220,19 @@ func resolveAgentFallbacks(agentCfg *config.AgentConfig, defaults *config.AgentD
return defaults.ModelFallbacks return defaults.ModelFallbacks
} }
func compilePatterns(patterns []string) []*regexp.Regexp {
compiled := make([]*regexp.Regexp, 0, len(patterns))
for _, p := range patterns {
re, err := regexp.Compile(p)
if err != nil {
fmt.Printf("Warning: invalid path pattern %q: %v\n", p, err)
continue
}
compiled = append(compiled, re)
}
return compiled
}
func expandHome(path string) string { func expandHome(path string) string {
if path == "" { if path == "" {
return path return path

View file

@ -95,75 +95,68 @@ func TestNewAgentInstance_DefaultsTemperatureWhenUnset(t *testing.T) {
} }
func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) { func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*") tests := []struct {
if err != nil { name string
t.Fatalf("Failed to create temp dir: %v", err) aliasName string
} modelName string
defer os.RemoveAll(tmpDir) apiBase string
wantProvider string
cfg := &config.Config{ wantModel string
Agents: config.AgentsConfig{ }{
Defaults: config.AgentDefaults{ {
Workspace: tmpDir, name: "alias with provider prefix",
Model: "step-3.5-flash", aliasName: "step-3.5-flash",
}, modelName: "openrouter/stepfun/step-3.5-flash:free",
apiBase: "https://openrouter.ai/api/v1",
wantProvider: "openrouter",
wantModel: "stepfun/step-3.5-flash:free",
}, },
ModelList: []config.ModelConfig{ {
{ name: "alias without provider prefix",
ModelName: "step-3.5-flash", aliasName: "glm-5",
Model: "openrouter/stepfun/step-3.5-flash:free", modelName: "glm-5",
APIBase: "https://openrouter.ai/api/v1", apiBase: "https://api.z.ai/api/coding/paas/v4",
}, wantProvider: "openai",
wantModel: "glm-5",
}, },
} }
provider := &mockProvider{} for _, tt := range tests {
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) t.Run(tt.name, func(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
if len(agent.Candidates) != 1 { cfg := &config.Config{
t.Fatalf("len(Candidates) = %d, want 1", len(agent.Candidates)) Agents: config.AgentsConfig{
} Defaults: config.AgentDefaults{
if agent.Candidates[0].Provider != "openrouter" { Workspace: tmpDir,
t.Fatalf("candidate provider = %q, want %q", agent.Candidates[0].Provider, "openrouter") Model: tt.aliasName,
} },
if agent.Candidates[0].Model != "stepfun/step-3.5-flash:free" { },
t.Fatalf("candidate model = %q, want %q", agent.Candidates[0].Model, "stepfun/step-3.5-flash:free") ModelList: []config.ModelConfig{
} {
} ModelName: tt.aliasName,
Model: tt.modelName,
func TestNewAgentInstance_ResolveCandidatesFromModelListAliasWithoutProtocol(t *testing.T) { APIBase: tt.apiBase,
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*") },
if err != nil { },
t.Fatalf("Failed to create temp dir: %v", err) }
}
defer os.RemoveAll(tmpDir) provider := &mockProvider{}
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
cfg := &config.Config{
Agents: config.AgentsConfig{ if len(agent.Candidates) != 1 {
Defaults: config.AgentDefaults{ t.Fatalf("len(Candidates) = %d, want 1", len(agent.Candidates))
Workspace: tmpDir, }
Model: "glm-5", if agent.Candidates[0].Provider != tt.wantProvider {
}, t.Fatalf("candidate provider = %q, want %q", agent.Candidates[0].Provider, tt.wantProvider)
}, }
ModelList: []config.ModelConfig{ if agent.Candidates[0].Model != tt.wantModel {
{ t.Fatalf("candidate model = %q, want %q", agent.Candidates[0].Model, tt.wantModel)
ModelName: "glm-5", }
Model: "glm-5", })
APIBase: "https://api.z.ai/api/coding/paas/v4",
},
},
}
provider := &mockProvider{}
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
if len(agent.Candidates) != 1 {
t.Fatalf("len(Candidates) = %d, want 1", len(agent.Candidates))
}
if agent.Candidates[0].Provider != "openai" {
t.Fatalf("candidate provider = %q, want %q", agent.Candidates[0].Provider, "openai")
}
if agent.Candidates[0].Model != "glm-5" {
t.Fatalf("candidate model = %q, want %q", agent.Candidates[0].Model, "glm-5")
} }
} }

View file

@ -23,6 +23,7 @@ import (
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/constants" "github.com/sipeed/picoclaw/pkg/constants"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/mcp"
"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"
@ -46,19 +47,24 @@ type AgentLoop struct {
// processOptions configures how a message is processed // processOptions configures how a message is processed
type processOptions struct { type processOptions struct {
SessionKey string // Session identifier for history/context SessionKey string // Session identifier for history/context
Channel string // Target channel for tool execution Channel string // Target channel for tool execution
ChatID string // Target chat ID for tool execution ChatID string // Target chat ID for tool execution
UserMessage string // User message content (may include prefix) UserMessage string // User message content (may include prefix)
DefaultResponse string // Response when LLM returns empty Media []string // media:// refs from inbound message
EnableSummary bool // Whether to trigger summarization DefaultResponse string // Response when LLM returns empty
SendResponse bool // Whether to send response via bus EnableSummary bool // Whether to trigger summarization
NoHistory bool // If true, don't load session history (for heartbeat) SendResponse bool // Whether to send response via bus
NoHistory bool // If true, don't load session history (for heartbeat)
} }
const defaultResponse = "I've completed processing but have no response to give. Increase `max_tool_iterations` in config.json." const defaultResponse = "I've completed processing but have no response to give. Increase `max_tool_iterations` in config.json."
func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers.LLMProvider) *AgentLoop { func NewAgentLoop(
cfg *config.Config,
msgBus *bus.MessageBus,
provider providers.LLMProvider,
) *AgentLoop {
registry := NewAgentRegistry(cfg, provider) registry := NewAgentRegistry(cfg, provider)
// Register shared tools to all agents // Register shared tools to all agents
@ -99,7 +105,7 @@ func registerSharedTools(
} }
// Web tools // Web tools
if searchTool := tools.NewWebSearchTool(tools.WebSearchToolOptions{ searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
BraveAPIKey: cfg.Tools.Web.Brave.APIKey, BraveAPIKey: cfg.Tools.Web.Brave.APIKey,
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults, BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
BraveEnabled: cfg.Tools.Web.Brave.Enabled, BraveEnabled: cfg.Tools.Web.Brave.Enabled,
@ -112,11 +118,24 @@ func registerSharedTools(
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey, PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults, PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled, PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
GLMSearchAPIKey: cfg.Tools.Web.GLMSearch.APIKey,
GLMSearchBaseURL: cfg.Tools.Web.GLMSearch.BaseURL,
GLMSearchEngine: cfg.Tools.Web.GLMSearch.SearchEngine,
GLMSearchMaxResults: cfg.Tools.Web.GLMSearch.MaxResults,
GLMSearchEnabled: cfg.Tools.Web.GLMSearch.Enabled,
Proxy: cfg.Tools.Web.Proxy, Proxy: cfg.Tools.Web.Proxy,
}); searchTool != nil { })
if err != nil {
logger.ErrorCF("agent", "Failed to create web search tool", map[string]any{"error": err.Error()})
} else if searchTool != nil {
agent.Tools.Register(searchTool) agent.Tools.Register(searchTool)
} }
agent.Tools.Register(tools.NewWebFetchToolWithProxy(50000, cfg.Tools.Web.Proxy)) fetchTool, err := tools.NewWebFetchToolWithProxy(50000, cfg.Tools.Web.Proxy, cfg.Tools.Web.FetchLimitBytes)
if err != nil {
logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
} else {
agent.Tools.Register(fetchTool)
}
// Hardware tools (I2C, SPI) - Linux only, returns error on other platforms // Hardware tools (I2C, SPI) - Linux only, returns error on other platforms
agent.Tools.Register(tools.NewI2CTool()) agent.Tools.Register(tools.NewI2CTool())
@ -162,6 +181,72 @@ func registerSharedTools(
func (al *AgentLoop) Run(ctx context.Context) error { func (al *AgentLoop) Run(ctx context.Context) error {
al.running.Store(true) al.running.Store(true)
// Initialize MCP servers for all agents
if al.cfg.Tools.MCP.Enabled {
mcpManager := mcp.NewManager()
// Ensure MCP connections are cleaned up on exit, regardless of initialization success
// This fixes resource leak when LoadFromMCPConfig partially succeeds then fails
defer func() {
if err := mcpManager.Close(); err != nil {
logger.ErrorCF("agent", "Failed to close MCP manager",
map[string]any{
"error": err.Error(),
})
}
}()
defaultAgent := al.registry.GetDefaultAgent()
var workspacePath string
if defaultAgent != nil && defaultAgent.Workspace != "" {
workspacePath = defaultAgent.Workspace
} else {
workspacePath = al.cfg.WorkspacePath()
}
if err := mcpManager.LoadFromMCPConfig(ctx, al.cfg.Tools.MCP, workspacePath); err != nil {
logger.WarnCF("agent", "Failed to load MCP servers, MCP tools will not be available",
map[string]any{
"error": err.Error(),
})
} else {
// Register MCP tools for all agents
servers := mcpManager.GetServers()
uniqueTools := 0
totalRegistrations := 0
agentIDs := al.registry.ListAgentIDs()
agentCount := len(agentIDs)
for serverName, conn := range servers {
uniqueTools += len(conn.Tools)
for _, tool := range conn.Tools {
for _, agentID := range agentIDs {
agent, ok := al.registry.GetAgent(agentID)
if !ok {
continue
}
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
agent.Tools.Register(mcpTool)
totalRegistrations++
logger.DebugCF("agent", "Registered MCP tool",
map[string]any{
"agent_id": agentID,
"server": serverName,
"tool": tool.Name,
"name": mcpTool.Name(),
})
}
}
}
logger.InfoCF("agent", "MCP tools registered successfully",
map[string]any{
"server_count": len(servers),
"unique_tools": uniqueTools,
"total_registrations": totalRegistrations,
"agent_count": agentCount,
})
}
}
for al.running.Load() { for al.running.Load() {
select { select {
case <-ctx.Done(): case <-ctx.Done():
@ -302,7 +387,10 @@ func (al *AgentLoop) RecordLastChatID(chatID string) error {
return al.state.SetLastChatID(chatID) return al.state.SetLastChatID(chatID)
} }
func (al *AgentLoop) ProcessDirect(ctx context.Context, content, sessionKey string) (string, error) { func (al *AgentLoop) ProcessDirect(
ctx context.Context,
content, sessionKey string,
) (string, error) {
return al.ProcessDirectWithChannel(ctx, content, sessionKey, "cli", "direct") return al.ProcessDirectWithChannel(ctx, content, sessionKey, "cli", "direct")
} }
@ -323,7 +411,10 @@ func (al *AgentLoop) ProcessDirectWithChannel(
// ProcessHeartbeat processes a heartbeat request without session history. // ProcessHeartbeat processes a heartbeat request without session history.
// Each heartbeat is independent and doesn't accumulate context. // Each heartbeat is independent and doesn't accumulate context.
func (al *AgentLoop) ProcessHeartbeat(ctx context.Context, content, channel, chatID string) (string, error) { func (al *AgentLoop) ProcessHeartbeat(
ctx context.Context,
content, channel, chatID string,
) (string, error) {
agent := al.registry.GetDefaultAgent() agent := al.registry.GetDefaultAgent()
if agent == nil { if agent == nil {
return "", fmt.Errorf("no default agent for heartbeat") return "", fmt.Errorf("no default agent for heartbeat")
@ -348,13 +439,16 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
} else { } else {
logContent = utils.Truncate(msg.Content, 80) logContent = utils.Truncate(msg.Content, 80)
} }
logger.InfoCF("agent", fmt.Sprintf("Processing message from %s:%s: %s", msg.Channel, msg.SenderID, logContent), logger.InfoCF(
"agent",
fmt.Sprintf("Processing message from %s:%s: %s", msg.Channel, msg.SenderID, logContent),
map[string]any{ map[string]any{
"channel": msg.Channel, "channel": msg.Channel,
"chat_id": msg.ChatID, "chat_id": msg.ChatID,
"sender_id": msg.SenderID, "sender_id": msg.SenderID,
"session_key": msg.SessionKey, "session_key": msg.SessionKey,
}) },
)
// Route system messages to processSystemMessage // Route system messages to processSystemMessage
if msg.Channel == "system" { if msg.Channel == "system" {
@ -409,15 +503,22 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
Channel: msg.Channel, Channel: msg.Channel,
ChatID: msg.ChatID, ChatID: msg.ChatID,
UserMessage: msg.Content, UserMessage: msg.Content,
Media: msg.Media,
DefaultResponse: defaultResponse, DefaultResponse: defaultResponse,
EnableSummary: true, EnableSummary: true,
SendResponse: false, SendResponse: false,
}) })
} }
func (al *AgentLoop) processSystemMessage(ctx context.Context, msg bus.InboundMessage) (string, error) { func (al *AgentLoop) processSystemMessage(
ctx context.Context,
msg bus.InboundMessage,
) (string, error) {
if msg.Channel != "system" { if msg.Channel != "system" {
return "", fmt.Errorf("processSystemMessage called with non-system message channel: %s", msg.Channel) return "", fmt.Errorf(
"processSystemMessage called with non-system message channel: %s",
msg.Channel,
)
} }
logger.InfoCF("agent", "Processing system message", logger.InfoCF("agent", "Processing system message",
@ -475,14 +576,22 @@ func (al *AgentLoop) processSystemMessage(ctx context.Context, msg bus.InboundMe
} }
// runAgentLoop is the core message processing logic. // runAgentLoop is the core message processing logic.
func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opts processOptions) (string, error) { func (al *AgentLoop) runAgentLoop(
ctx context.Context,
agent *AgentInstance,
opts processOptions,
) (string, error) {
// 0. Record last channel for heartbeat notifications (skip internal channels) // 0. Record last channel for heartbeat notifications (skip internal channels)
if opts.Channel != "" && opts.ChatID != "" { if opts.Channel != "" && opts.ChatID != "" {
// Don't record internal channels (cli, system, subagent) // Don't record internal channels (cli, system, subagent)
if !constants.IsInternalChannel(opts.Channel) { if !constants.IsInternalChannel(opts.Channel) {
channelKey := fmt.Sprintf("%s:%s", opts.Channel, opts.ChatID) channelKey := fmt.Sprintf("%s:%s", opts.Channel, opts.ChatID)
if err := al.RecordLastChannel(channelKey); err != nil { if err := al.RecordLastChannel(channelKey); err != nil {
logger.WarnCF("agent", "Failed to record last channel", map[string]any{"error": err.Error()}) logger.WarnCF(
"agent",
"Failed to record last channel",
map[string]any{"error": err.Error()},
)
} }
} }
} }
@ -501,11 +610,15 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt
history, history,
summary, summary,
opts.UserMessage, opts.UserMessage,
nil, opts.Media,
opts.Channel, opts.Channel,
opts.ChatID, opts.ChatID,
) )
// Resolve media:// refs to base64 data URLs (streaming)
maxMediaSize := al.cfg.Agents.Defaults.GetMaxMediaSize()
messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize)
// 3. Save user message to session // 3. Save user message to session
agent.Sessions.AddMessage(opts.SessionKey, "user", opts.UserMessage) agent.Sessions.AddMessage(opts.SessionKey, "user", opts.UserMessage)
@ -564,7 +677,10 @@ func (al *AgentLoop) targetReasoningChannelID(channelName string) (chatID string
return "" return ""
} }
func (al *AgentLoop) handleReasoning(ctx context.Context, reasoningContent, channelName, channelID string) { func (al *AgentLoop) handleReasoning(
ctx context.Context,
reasoningContent, channelName, channelID string,
) {
if reasoningContent == "" || channelName == "" || channelID == "" { if reasoningContent == "" || channelName == "" || channelID == "" {
return return
} }
@ -657,22 +773,33 @@ func (al *AgentLoop) runLLMIteration(
callLLM := func() (*providers.LLMResponse, error) { callLLM := func() (*providers.LLMResponse, error) {
if len(agent.Candidates) > 1 && al.fallback != nil { if len(agent.Candidates) > 1 && al.fallback != nil {
fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates, fbResult, fbErr := al.fallback.Execute(
ctx,
agent.Candidates,
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) { func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
return agent.Provider.Chat(ctx, messages, providerToolDefs, model, map[string]any{ return agent.Provider.Chat(
"max_tokens": agent.MaxTokens, ctx,
"temperature": agent.Temperature, messages,
"prompt_cache_key": agent.ID, providerToolDefs,
}) model,
map[string]any{
"max_tokens": agent.MaxTokens,
"temperature": agent.Temperature,
"prompt_cache_key": agent.ID,
},
)
}, },
) )
if fbErr != nil { if fbErr != nil {
return nil, fbErr return nil, fbErr
} }
if fbResult.Provider != "" && len(fbResult.Attempts) > 0 { if fbResult.Provider != "" && len(fbResult.Attempts) > 0 {
logger.InfoCF("agent", fmt.Sprintf("Fallback: succeeded with %s/%s after %d attempts", logger.InfoCF(
fbResult.Provider, fbResult.Model, len(fbResult.Attempts)+1), "agent",
map[string]any{"agent_id": agent.ID, "iteration": iteration}) fmt.Sprintf("Fallback: succeeded with %s/%s after %d attempts",
fbResult.Provider, fbResult.Model, len(fbResult.Attempts)+1),
map[string]any{"agent_id": agent.ID, "iteration": iteration},
)
} }
return fbResult.Response, nil return fbResult.Response, nil
} }
@ -723,10 +850,14 @@ func (al *AgentLoop) runLLMIteration(
} }
if isContextError && retry < maxRetries { if isContextError && retry < maxRetries {
logger.WarnCF("agent", "Context window error detected, attempting compression", map[string]any{ logger.WarnCF(
"error": err.Error(), "agent",
"retry": retry, "Context window error detected, attempting compression",
}) map[string]any{
"error": err.Error(),
"retry": retry,
},
)
if retry == 0 && !constants.IsInternalChannel(opts.Channel) { if retry == 0 && !constants.IsInternalChannel(opts.Channel) {
al.bus.PublishOutbound(ctx, bus.OutboundMessage{ al.bus.PublishOutbound(ctx, bus.OutboundMessage{
@ -758,7 +889,12 @@ func (al *AgentLoop) runLLMIteration(
return "", iteration, fmt.Errorf("LLM call failed after retries: %w", err) return "", iteration, fmt.Errorf("LLM call failed after retries: %w", err)
} }
go al.handleReasoning(ctx, response.Reasoning, opts.Channel, al.targetReasoningChannelID(opts.Channel)) go al.handleReasoning(
ctx,
response.Reasoning,
opts.Channel,
al.targetReasoningChannelID(opts.Channel),
)
logger.DebugCF("agent", "LLM response", logger.DebugCF("agent", "LLM response",
map[string]any{ map[string]any{
@ -833,62 +969,76 @@ func (al *AgentLoop) runLLMIteration(
// Save assistant message with tool calls to session // Save assistant message with tool calls to session
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg) agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
// Execute tool calls // Execute tool calls in parallel
for _, tc := range normalizedToolCalls { type indexedAgentResult struct {
argsJSON, _ := json.Marshal(tc.Arguments) result *tools.ToolResult
argsPreview := utils.Truncate(string(argsJSON), 200) tc providers.ToolCall
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview), }
map[string]any{
"agent_id": agent.ID,
"tool": tc.Name,
"iteration": iteration,
})
// Create async callback for tools that implement AsyncTool agentResults := make([]indexedAgentResult, len(normalizedToolCalls))
// NOTE: Following openclaw's design, async tools do NOT send results directly to users. var wg sync.WaitGroup
// Instead, they notify the agent via PublishInbound, and the agent decides
// whether to forward the result to the user (in processSystemMessage). for i, tc := range normalizedToolCalls {
asyncCallback := func(callbackCtx context.Context, result *tools.ToolResult) { agentResults[i].tc = tc
// Log the async completion but don't send directly to user
// The agent will handle user notification via processSystemMessage wg.Add(1)
if !result.Silent && result.ForUser != "" { go func(idx int, tc providers.ToolCall) {
logger.InfoCF("agent", "Async tool completed, agent will handle notification", defer wg.Done()
map[string]any{
"tool": tc.Name, argsJSON, _ := json.Marshal(tc.Arguments)
"content_len": len(result.ForUser), argsPreview := utils.Truncate(string(argsJSON), 200)
}) logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
map[string]any{
"agent_id": agent.ID,
"tool": tc.Name,
"iteration": iteration,
})
// Create async callback for tools that implement AsyncTool
asyncCallback := func(callbackCtx context.Context, result *tools.ToolResult) {
if !result.Silent && result.ForUser != "" {
logger.InfoCF("agent", "Async tool completed, agent will handle notification",
map[string]any{
"tool": tc.Name,
"content_len": len(result.ForUser),
})
}
} }
}
toolResult := agent.Tools.ExecuteWithContext( toolResult := agent.Tools.ExecuteWithContext(
ctx, ctx,
tc.Name, tc.Name,
tc.Arguments, tc.Arguments,
opts.Channel, opts.Channel,
opts.ChatID, opts.ChatID,
asyncCallback, asyncCallback,
) )
agentResults[idx].result = toolResult
}(i, tc)
}
wg.Wait()
// Process results in original order (send to user, save to session)
for _, r := range agentResults {
// Send ForUser content to user immediately if not Silent // Send ForUser content to user immediately if not Silent
if !toolResult.Silent && toolResult.ForUser != "" && opts.SendResponse { if !r.result.Silent && r.result.ForUser != "" && opts.SendResponse {
al.bus.PublishOutbound(ctx, bus.OutboundMessage{ al.bus.PublishOutbound(ctx, bus.OutboundMessage{
Channel: opts.Channel, Channel: opts.Channel,
ChatID: opts.ChatID, ChatID: opts.ChatID,
Content: toolResult.ForUser, Content: r.result.ForUser,
}) })
logger.DebugCF("agent", "Sent tool result to user", logger.DebugCF("agent", "Sent tool result to user",
map[string]any{ map[string]any{
"tool": tc.Name, "tool": r.tc.Name,
"content_len": len(toolResult.ForUser), "content_len": len(r.result.ForUser),
}) })
} }
// If tool returned media refs, publish them as outbound media // If tool returned media refs, publish them as outbound media
if len(toolResult.Media) > 0 && opts.SendResponse { if len(r.result.Media) > 0 && opts.SendResponse {
parts := make([]bus.MediaPart, 0, len(toolResult.Media)) parts := make([]bus.MediaPart, 0, len(r.result.Media))
for _, ref := range toolResult.Media { for _, ref := range r.result.Media {
part := bus.MediaPart{Ref: ref} part := bus.MediaPart{Ref: ref}
// Populate metadata from MediaStore when available
if al.mediaStore != nil { if al.mediaStore != nil {
if _, meta, err := al.mediaStore.ResolveWithMeta(ref); err == nil { if _, meta, err := al.mediaStore.ResolveWithMeta(ref); err == nil {
part.Filename = meta.Filename part.Filename = meta.Filename
@ -906,15 +1056,15 @@ func (al *AgentLoop) runLLMIteration(
} }
// Determine content for LLM based on tool result // Determine content for LLM based on tool result
contentForLLM := toolResult.ForLLM contentForLLM := r.result.ForLLM
if contentForLLM == "" && toolResult.Err != nil { if contentForLLM == "" && r.result.Err != nil {
contentForLLM = toolResult.Err.Error() contentForLLM = r.result.Err.Error()
} }
toolResultMsg := providers.Message{ toolResultMsg := providers.Message{
Role: "tool", Role: "tool",
Content: contentForLLM, Content: contentForLLM,
ToolCallID: tc.ID, ToolCallID: r.tc.ID,
} }
messages = append(messages, toolResultMsg) messages = append(messages, toolResultMsg)
@ -950,9 +1100,9 @@ func (al *AgentLoop) updateToolContexts(agent *AgentInstance, channel, chatID st
func (al *AgentLoop) maybeSummarize(agent *AgentInstance, sessionKey, channel, chatID string) { func (al *AgentLoop) maybeSummarize(agent *AgentInstance, sessionKey, channel, chatID string) {
newHistory := agent.Sessions.GetHistory(sessionKey) newHistory := agent.Sessions.GetHistory(sessionKey)
tokenEstimate := al.estimateTokens(newHistory) tokenEstimate := al.estimateTokens(newHistory)
threshold := agent.ContextWindow * 75 / 100 threshold := agent.ContextWindow * agent.SummarizeTokenPercent / 100
if len(newHistory) > 20 || tokenEstimate > threshold { if len(newHistory) > agent.SummarizeMessageThreshold || tokenEstimate > threshold {
summarizeKey := agent.ID + ":" + sessionKey summarizeKey := agent.ID + ":" + sessionKey
if _, loading := al.summarizing.LoadOrStore(summarizeKey, true); !loading { if _, loading := al.summarizing.LoadOrStore(summarizeKey, true); !loading {
go func() { go func() {
@ -1060,7 +1210,11 @@ func formatMessagesForLog(messages []providers.Message) string {
for _, tc := range msg.ToolCalls { for _, tc := range msg.ToolCalls {
fmt.Fprintf(&sb, " - ID: %s, Type: %s, Name: %s\n", tc.ID, tc.Type, tc.Name) fmt.Fprintf(&sb, " - ID: %s, Type: %s, Name: %s\n", tc.ID, tc.Type, tc.Name)
if tc.Function != nil { if tc.Function != nil {
fmt.Fprintf(&sb, " Arguments: %s\n", utils.Truncate(tc.Function.Arguments, 200)) fmt.Fprintf(
&sb,
" Arguments: %s\n",
utils.Truncate(tc.Function.Arguments, 200),
)
} }
} }
} }
@ -1089,7 +1243,11 @@ func formatToolsForLog(toolDefs []providers.ToolDefinition) string {
fmt.Fprintf(&sb, " [%d] Type: %s, Name: %s\n", i, tool.Type, tool.Function.Name) fmt.Fprintf(&sb, " [%d] Type: %s, Name: %s\n", i, tool.Type, tool.Function.Name)
fmt.Fprintf(&sb, " Description: %s\n", tool.Function.Description) fmt.Fprintf(&sb, " Description: %s\n", tool.Function.Description)
if len(tool.Function.Parameters) > 0 { if len(tool.Function.Parameters) > 0 {
fmt.Fprintf(&sb, " Parameters: %s\n", utils.Truncate(fmt.Sprintf("%v", tool.Function.Parameters), 200)) fmt.Fprintf(
&sb,
" Parameters: %s\n",
utils.Truncate(fmt.Sprintf("%v", tool.Function.Parameters), 200),
)
} }
} }
sb.WriteString("]") sb.WriteString("]")
@ -1186,7 +1344,9 @@ func (al *AgentLoop) summarizeBatch(
existingSummary string, existingSummary string,
) (string, error) { ) (string, error) {
var sb strings.Builder var sb strings.Builder
sb.WriteString("Provide a concise summary of this conversation segment, preserving core context and key points.\n") sb.WriteString(
"Provide a concise summary of this conversation segment, preserving core context and key points.\n",
)
if existingSummary != "" { if existingSummary != "" {
sb.WriteString("Existing context: ") sb.WriteString("Existing context: ")
sb.WriteString(existingSummary) sb.WriteString(existingSummary)

122
pkg/agent/loop_media.go Normal file
View file

@ -0,0 +1,122 @@
// PicoClaw - Ultra-lightweight personal AI agent
// Inspired by and based on nanobot: https://github.com/HKUDS/nanobot
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package agent
import (
"bytes"
"encoding/base64"
"io"
"os"
"strings"
"github.com/h2non/filetype"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/providers"
)
// resolveMediaRefs replaces media:// refs in message Media fields with base64 data URLs.
// Uses streaming base64 encoding (file handle → encoder → buffer) to avoid holding
// both raw bytes and encoded string in memory simultaneously.
// Returns a new slice; original messages are not mutated.
func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxSize int) []providers.Message {
if store == nil {
return messages
}
result := make([]providers.Message, len(messages))
copy(result, messages)
for i, m := range result {
if len(m.Media) == 0 {
continue
}
resolved := make([]string, 0, len(m.Media))
for _, ref := range m.Media {
if !strings.HasPrefix(ref, "media://") {
resolved = append(resolved, ref)
continue
}
localPath, meta, err := store.ResolveWithMeta(ref)
if err != nil {
logger.WarnCF("agent", "Failed to resolve media ref", map[string]any{
"ref": ref,
"error": err.Error(),
})
continue
}
info, err := os.Stat(localPath)
if err != nil {
logger.WarnCF("agent", "Failed to stat media file", map[string]any{
"path": localPath,
"error": err.Error(),
})
continue
}
if info.Size() > int64(maxSize) {
logger.WarnCF("agent", "Media file too large, skipping", map[string]any{
"path": localPath,
"size": info.Size(),
"max_size": maxSize,
})
continue
}
// Determine MIME type: prefer metadata, fallback to magic-bytes detection
mime := meta.ContentType
if mime == "" {
kind, ftErr := filetype.MatchFile(localPath)
if ftErr != nil || kind == filetype.Unknown {
logger.WarnCF("agent", "Unknown media type, skipping", map[string]any{
"path": localPath,
})
continue
}
mime = kind.MIME.Value
}
// Streaming base64: open file → base64 encoder → buffer
// Peak memory: ~1.33x file size (buffer only, no raw bytes copy)
f, err := os.Open(localPath)
if err != nil {
logger.WarnCF("agent", "Failed to open media file", map[string]any{
"path": localPath,
"error": err.Error(),
})
continue
}
prefix := "data:" + mime + ";base64,"
encodedLen := base64.StdEncoding.EncodedLen(int(info.Size()))
var buf bytes.Buffer
buf.Grow(len(prefix) + encodedLen)
buf.WriteString(prefix)
encoder := base64.NewEncoder(base64.StdEncoding, &buf)
if _, err := io.Copy(encoder, f); err != nil {
f.Close()
logger.WarnCF("agent", "Failed to encode media file", map[string]any{
"path": localPath,
"error": err.Error(),
})
continue
}
encoder.Close()
f.Close()
resolved = append(resolved, buf.String())
}
result[i].Media = resolved
}
return result
}

View file

@ -5,12 +5,15 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"slices"
"strings"
"testing" "testing"
"time" "time"
"github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels" "github.com/sipeed/picoclaw/pkg/channels"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/providers"
"github.com/sipeed/picoclaw/pkg/tools" "github.com/sipeed/picoclaw/pkg/tools"
) )
@ -26,16 +29,15 @@ func (f *fakeChannel) IsAllowed(string) bool {
func (f *fakeChannel) IsAllowedSender(sender bus.SenderInfo) bool { return true } func (f *fakeChannel) IsAllowedSender(sender bus.SenderInfo) bool { return true }
func (f *fakeChannel) ReasoningChannelID() string { return f.id } func (f *fakeChannel) ReasoningChannelID() string { return f.id }
func TestRecordLastChannel(t *testing.T) { func newTestAgentLoop(
// Create temp workspace t *testing.T,
) (al *AgentLoop, cfg *config.Config, msgBus *bus.MessageBus, provider *mockProvider, cleanup func()) {
t.Helper()
tmpDir, err := os.MkdirTemp("", "agent-test-*") tmpDir, err := os.MkdirTemp("", "agent-test-*")
if err != nil { if err != nil {
t.Fatalf("Failed to create temp dir: %v", err) t.Fatalf("Failed to create temp dir: %v", err)
} }
defer os.RemoveAll(tmpDir) cfg = &config.Config{
// Create test config
cfg := &config.Config{
Agents: config.AgentsConfig{ Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{ Defaults: config.AgentDefaults{
Workspace: tmpDir, Workspace: tmpDir,
@ -45,74 +47,43 @@ func TestRecordLastChannel(t *testing.T) {
}, },
}, },
} }
msgBus = bus.NewMessageBus()
provider = &mockProvider{}
al = NewAgentLoop(cfg, msgBus, provider)
return al, cfg, msgBus, provider, func() { os.RemoveAll(tmpDir) }
}
// Create agent loop func TestRecordLastChannel(t *testing.T) {
msgBus := bus.NewMessageBus() al, cfg, msgBus, provider, cleanup := newTestAgentLoop(t)
provider := &mockProvider{} defer cleanup()
al := NewAgentLoop(cfg, msgBus, provider)
// Test RecordLastChannel
testChannel := "test-channel" testChannel := "test-channel"
err = al.RecordLastChannel(testChannel) if err := al.RecordLastChannel(testChannel); err != nil {
if err != nil {
t.Fatalf("RecordLastChannel failed: %v", err) t.Fatalf("RecordLastChannel failed: %v", err)
} }
if got := al.state.GetLastChannel(); got != testChannel {
// Verify channel was saved t.Errorf("Expected channel '%s', got '%s'", testChannel, got)
lastChannel := al.state.GetLastChannel()
if lastChannel != testChannel {
t.Errorf("Expected channel '%s', got '%s'", testChannel, lastChannel)
} }
// Verify persistence by creating a new agent loop
al2 := NewAgentLoop(cfg, msgBus, provider) al2 := NewAgentLoop(cfg, msgBus, provider)
if al2.state.GetLastChannel() != testChannel { if got := al2.state.GetLastChannel(); got != testChannel {
t.Errorf("Expected persistent channel '%s', got '%s'", testChannel, al2.state.GetLastChannel()) t.Errorf("Expected persistent channel '%s', got '%s'", testChannel, got)
} }
} }
func TestRecordLastChatID(t *testing.T) { func TestRecordLastChatID(t *testing.T) {
// Create temp workspace al, cfg, msgBus, provider, cleanup := newTestAgentLoop(t)
tmpDir, err := os.MkdirTemp("", "agent-test-*") defer cleanup()
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
// Create test config
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
MaxTokens: 4096,
MaxToolIterations: 10,
},
},
}
// Create agent loop
msgBus := bus.NewMessageBus()
provider := &mockProvider{}
al := NewAgentLoop(cfg, msgBus, provider)
// Test RecordLastChatID
testChatID := "test-chat-id-123" testChatID := "test-chat-id-123"
err = al.RecordLastChatID(testChatID) if err := al.RecordLastChatID(testChatID); err != nil {
if err != nil {
t.Fatalf("RecordLastChatID failed: %v", err) t.Fatalf("RecordLastChatID failed: %v", err)
} }
if got := al.state.GetLastChatID(); got != testChatID {
// Verify chat ID was saved t.Errorf("Expected chat ID '%s', got '%s'", testChatID, got)
lastChatID := al.state.GetLastChatID()
if lastChatID != testChatID {
t.Errorf("Expected chat ID '%s', got '%s'", testChatID, lastChatID)
} }
// Verify persistence by creating a new agent loop
al2 := NewAgentLoop(cfg, msgBus, provider) al2 := NewAgentLoop(cfg, msgBus, provider)
if al2.state.GetLastChatID() != testChatID { if got := al2.state.GetLastChatID(); got != testChatID {
t.Errorf("Expected persistent chat ID '%s', got '%s'", testChatID, al2.state.GetLastChatID()) t.Errorf("Expected persistent chat ID '%s', got '%s'", testChatID, got)
} }
} }
@ -187,13 +158,7 @@ func TestToolRegistry_ToolRegistration(t *testing.T) {
toolsList := toolsInfo["names"].([]string) toolsList := toolsInfo["names"].([]string)
// Check that our custom tool name is in the list // Check that our custom tool name is in the list
found := false found := slices.Contains(toolsList, "mock_custom")
for _, name := range toolsList {
if name == "mock_custom" {
found = true
break
}
}
if !found { if !found {
t.Error("Expected custom tool to be registered") t.Error("Expected custom tool to be registered")
} }
@ -262,13 +227,7 @@ func TestToolRegistry_GetDefinitions(t *testing.T) {
toolsList := toolsInfo["names"].([]string) toolsList := toolsInfo["names"].([]string)
// Check that our custom tool name is in the list // Check that our custom tool name is in the list
found := false found := slices.Contains(toolsList, "mock_custom")
for _, name := range toolsList {
if name == "mock_custom" {
found = true
break
}
}
if !found { if !found {
t.Error("Expected custom tool to be registered") t.Error("Expected custom tool to be registered")
} }
@ -851,3 +810,142 @@ func TestHandleReasoning(t *testing.T) {
} }
}) })
} }
func TestResolveMediaRefs_ResolvesToBase64(t *testing.T) {
store := media.NewFileMediaStore()
dir := t.TempDir()
// Create a minimal valid PNG (8-byte header is enough for filetype detection)
pngPath := filepath.Join(dir, "test.png")
// PNG magic: 0x89 P N G \r \n 0x1A \n + minimal IHDR
pngHeader := []byte{
0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, // PNG signature
0x00, 0x00, 0x00, 0x0D, // IHDR length
0x49, 0x48, 0x44, 0x52, // "IHDR"
0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02, // 1x1 RGB
0x00, 0x00, 0x00, // no interlace
0x90, 0x77, 0x53, 0xDE, // CRC
}
if err := os.WriteFile(pngPath, pngHeader, 0o644); err != nil {
t.Fatal(err)
}
ref, err := store.Store(pngPath, media.MediaMeta{}, "test")
if err != nil {
t.Fatal(err)
}
messages := []providers.Message{
{Role: "user", Content: "describe this", Media: []string{ref}},
}
result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize)
if len(result[0].Media) != 1 {
t.Fatalf("expected 1 resolved media, got %d", len(result[0].Media))
}
if !strings.HasPrefix(result[0].Media[0], "data:image/png;base64,") {
t.Fatalf("expected data:image/png;base64, prefix, got %q", result[0].Media[0][:40])
}
}
func TestResolveMediaRefs_SkipsOversizedFile(t *testing.T) {
store := media.NewFileMediaStore()
dir := t.TempDir()
bigPath := filepath.Join(dir, "big.png")
// Write PNG header + padding to exceed limit
data := make([]byte, 1024+1) // 1KB + 1 byte
copy(data, []byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A})
if err := os.WriteFile(bigPath, data, 0o644); err != nil {
t.Fatal(err)
}
ref, _ := store.Store(bigPath, media.MediaMeta{}, "test")
messages := []providers.Message{
{Role: "user", Content: "hi", Media: []string{ref}},
}
// Use a tiny limit (1KB) so the file is oversized
result := resolveMediaRefs(messages, store, 1024)
if len(result[0].Media) != 0 {
t.Fatalf("expected 0 media (oversized), got %d", len(result[0].Media))
}
}
func TestResolveMediaRefs_SkipsUnknownType(t *testing.T) {
store := media.NewFileMediaStore()
dir := t.TempDir()
txtPath := filepath.Join(dir, "readme.txt")
if err := os.WriteFile(txtPath, []byte("hello world"), 0o644); err != nil {
t.Fatal(err)
}
ref, _ := store.Store(txtPath, media.MediaMeta{}, "test")
messages := []providers.Message{
{Role: "user", Content: "hi", Media: []string{ref}},
}
result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize)
if len(result[0].Media) != 0 {
t.Fatalf("expected 0 media (unknown type), got %d", len(result[0].Media))
}
}
func TestResolveMediaRefs_PassesThroughNonMediaRefs(t *testing.T) {
messages := []providers.Message{
{Role: "user", Content: "hi", Media: []string{"https://example.com/img.png"}},
}
result := resolveMediaRefs(messages, nil, config.DefaultMaxMediaSize)
if len(result[0].Media) != 1 || result[0].Media[0] != "https://example.com/img.png" {
t.Fatalf("expected passthrough of non-media:// URL, got %v", result[0].Media)
}
}
func TestResolveMediaRefs_DoesNotMutateOriginal(t *testing.T) {
store := media.NewFileMediaStore()
dir := t.TempDir()
pngPath := filepath.Join(dir, "test.png")
pngHeader := []byte{
0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A,
0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52,
0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02,
0x00, 0x00, 0x00, 0x90, 0x77, 0x53, 0xDE,
}
os.WriteFile(pngPath, pngHeader, 0o644)
ref, _ := store.Store(pngPath, media.MediaMeta{}, "test")
original := []providers.Message{
{Role: "user", Content: "hi", Media: []string{ref}},
}
originalRef := original[0].Media[0]
resolveMediaRefs(original, store, config.DefaultMaxMediaSize)
if original[0].Media[0] != originalRef {
t.Fatal("resolveMediaRefs mutated original message slice")
}
}
func TestResolveMediaRefs_UsesMetaContentType(t *testing.T) {
store := media.NewFileMediaStore()
dir := t.TempDir()
// File with JPEG content but stored with explicit content type
jpegPath := filepath.Join(dir, "photo")
jpegHeader := []byte{0xFF, 0xD8, 0xFF, 0xE0} // JPEG magic bytes
os.WriteFile(jpegPath, jpegHeader, 0o644)
ref, _ := store.Store(jpegPath, media.MediaMeta{ContentType: "image/jpeg"}, "test")
messages := []providers.Message{
{Role: "user", Content: "hi", Media: []string{ref}},
}
result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize)
if len(result[0].Media) != 1 {
t.Fatalf("expected 1 media, got %d", len(result[0].Media))
}
if !strings.HasPrefix(result[0].Media[0], "data:image/jpeg;base64,") {
t.Fatalf("expected jpeg prefix, got %q", result[0].Media[0][:30])
}
}

View file

@ -41,6 +41,11 @@ func NewMemoryStore(workspace string) *MemoryStore {
} }
} }
// GetWorkspace returns the workspace path.
func (ms *MemoryStore) GetWorkspace() string {
return ms.workspace
}
// getTodayFile returns the path to today's daily note file (memory/YYYYMM/YYYYMMDD.md). // getTodayFile returns the path to today's daily note file (memory/YYYYMM/YYYYMMDD.md).
func (ms *MemoryStore) getTodayFile() string { func (ms *MemoryStore) getTodayFile() string {
today := time.Now().Format("20060102") // YYYYMMDD today := time.Now().Format("20060102") // YYYYMMDD
@ -111,7 +116,7 @@ func (ms *MemoryStore) GetRecentDailyNotes(days int) string {
var sb strings.Builder var sb strings.Builder
first := true first := true
for i := 0; i < days; i++ { for i := range days {
date := time.Now().AddDate(0, 0, -i) date := time.Now().AddDate(0, 0, -i)
dateStr := date.Format("20060102") // YYYYMMDD dateStr := date.Format("20060102") // YYYYMMDD
monthDir := dateStr[:6] // YYYYMM monthDir := dateStr[:6] // YYYYMM

View file

@ -67,7 +67,7 @@ func TestPublishInbound_ContextCancel(t *testing.T) {
// Fill the buffer // Fill the buffer
ctx := context.Background() ctx := context.Background()
for i := 0; i < defaultBusBufferSize; i++ { for i := range defaultBusBufferSize {
if err := mb.PublishInbound(ctx, InboundMessage{Content: "fill"}); err != nil { if err := mb.PublishInbound(ctx, InboundMessage{Content: "fill"}); err != nil {
t.Fatalf("fill failed at %d: %v", i, err) t.Fatalf("fill failed at %d: %v", i, err)
} }
@ -154,7 +154,7 @@ func TestConcurrentPublishClose(t *testing.T) {
wg.Add(numGoroutines + 1) wg.Add(numGoroutines + 1)
// Spawn many goroutines trying to publish // Spawn many goroutines trying to publish
for i := 0; i < numGoroutines; i++ { for range numGoroutines {
go func() { go func() {
defer wg.Done() defer wg.Done()
// Use a short timeout context so we don't block forever after close // Use a short timeout context so we don't block forever after close
@ -194,7 +194,7 @@ func TestPublishInbound_FullBuffer(t *testing.T) {
ctx := context.Background() ctx := context.Background()
// Fill the buffer // Fill the buffer
for i := 0; i < defaultBusBufferSize; i++ { for i := range defaultBusBufferSize {
if err := mb.PublishInbound(ctx, InboundMessage{Content: "fill"}); err != nil { if err := mb.PublishInbound(ctx, InboundMessage{Content: "fill"}); err != nil {
t.Fatalf("fill failed at %d: %v", i, err) t.Fatalf("fill failed at %d: %v", i, err)
} }

View file

@ -1,7 +1,5 @@
# PicoClaw Channel System Refactor: Complete Development Guide # PicoClaw Channel System: Complete Development Guide
> **Branch**: `refactor/channel-system`
> **Status**: Active development (~40 commits)
> **Scope**: `pkg/channels/`, `pkg/bus/`, `pkg/media/`, `pkg/identity/`, `cmd/picoclaw/internal/gateway/` > **Scope**: `pkg/channels/`, `pkg/bus/`, `pkg/media/`, `pkg/identity/`, `cmd/picoclaw/internal/gateway/`
--- ---
@ -46,6 +44,8 @@ pkg/channels/
pkg/channels/ pkg/channels/
├── base.go # BaseChannel shared abstraction layer ├── base.go # BaseChannel shared abstraction layer
├── interfaces.go # Optional capability interfaces (TypingCapable, MessageEditor, ReactionCapable, PlaceholderCapable, PlaceholderRecorder) ├── interfaces.go # Optional capability interfaces (TypingCapable, MessageEditor, ReactionCapable, PlaceholderCapable, PlaceholderRecorder)
├── README.md # English documentation
├── README.zh.md # Chinese documentation
├── media.go # MediaSender optional interface ├── media.go # MediaSender optional interface
├── webhook.go # WebhookHandler, HealthChecker optional interfaces ├── webhook.go # WebhookHandler, HealthChecker optional interfaces
├── errors.go # Sentinel errors (ErrNotRunning, ErrRateLimit, ErrTemporary, ErrSendFailed) ├── errors.go # Sentinel errors (ErrNotRunning, ErrRateLimit, ErrTemporary, ErrSendFailed)
@ -60,7 +60,7 @@ pkg/channels/
├── discord/ ├── discord/
│ ├── init.go │ ├── init.go
│ └── discord.go │ └── discord.go
├── slack/ line/ onebot/ dingtalk/ feishu/ wecom/ qq/ whatsapp/ maixcam/ pico/ ├── slack/ line/ onebot/ dingtalk/ feishu/ wecom/ qq/ whatsapp/ whatsapp_native/ maixcam/ pico/
│ └── ... │ └── ...
pkg/bus/ pkg/bus/
@ -111,7 +111,7 @@ pkg/identity/
|-----------|-------------| |-----------|-------------|
| **Sub-package Isolation** | Each channel is a standalone Go sub-package, depending on `BaseChannel` and interfaces from the `channels` parent package | | **Sub-package Isolation** | Each channel is a standalone Go sub-package, depending on `BaseChannel` and interfaces from the `channels` parent package |
| **Factory Registration** | Sub-packages self-register via `init()`, Manager looks up factories by name, eliminating import coupling | | **Factory Registration** | Sub-packages self-register via `init()`, Manager looks up factories by name, eliminating import coupling |
| **Capability Discovery** | Optional capabilities are declared via interfaces (`MediaSender`, `TypingCapable`, `ReactionCapable`, `PlaceholderCapable`, `MessageEditor`, `WebhookHandler`), discovered by Manager via runtime type assertions | | **Capability Discovery** | Optional capabilities are declared via interfaces (`MediaSender`, `TypingCapable`, `ReactionCapable`, `PlaceholderCapable`, `MessageEditor`, `WebhookHandler`, `HealthChecker`), discovered by Manager via runtime type assertions |
| **Structured Messages** | Peer, MessageID, and SenderInfo promoted from Metadata to first-class fields on InboundMessage | | **Structured Messages** | Peer, MessageID, and SenderInfo promoted from Metadata to first-class fields on InboundMessage |
| **Error Classification** | Channels return sentinel errors (`ErrRateLimit`, `ErrTemporary`, etc.), Manager uses these to determine retry strategy | | **Error Classification** | Channels return sentinel errors (`ErrRateLimit`, `ErrTemporary`, etc.), Manager uses these to determine retry strategy |
| **Centralized Orchestration** | Rate limiting, message splitting, retries, and Typing/Reaction/Placeholder management are all handled by Manager and BaseChannel; channels only need to implement Send | | **Centralized Orchestration** | Rate limiting, message splitting, retries, and Typing/Reaction/Placeholder management are all handled by Manager and BaseChannel; channels only need to implement Send |
@ -145,6 +145,7 @@ After refactoring, these files have been removed and code moved to corresponding
| _(did not exist)_ | `pkg/channels/interfaces.go` | New optional capability interfaces | | _(did not exist)_ | `pkg/channels/interfaces.go` | New optional capability interfaces |
| _(did not exist)_ | `pkg/channels/media.go` | New MediaSender interface | | _(did not exist)_ | `pkg/channels/media.go` | New MediaSender interface |
| _(did not exist)_ | `pkg/channels/webhook.go` | New WebhookHandler/HealthChecker | | _(did not exist)_ | `pkg/channels/webhook.go` | New WebhookHandler/HealthChecker |
| _(did not exist)_ | `pkg/channels/whatsapp_native/` | New WhatsApp native mode (whatsmeow) |
| _(did not exist)_ | `pkg/channels/split.go` | New message splitting (migrated from utils) | | _(did not exist)_ | `pkg/channels/split.go` | New message splitting (migrated from utils) |
| _(did not exist)_ | `pkg/bus/types.go` | New structured message types | | _(did not exist)_ | `pkg/bus/types.go` | New structured message types |
| _(did not exist)_ | `pkg/media/store.go` | New media file lifecycle management | | _(did not exist)_ | `pkg/media/store.go` | New media file lifecycle management |
@ -220,6 +221,7 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
cfg.Channels.Telegram.AllowFrom, // Allow list cfg.Channels.Telegram.AllowFrom, // Allow list
channels.WithMaxMessageLength(4096), // Platform message length limit channels.WithMaxMessageLength(4096), // Platform message length limit
channels.WithGroupTrigger(cfg.Channels.Telegram.GroupTrigger), // Group trigger config channels.WithGroupTrigger(cfg.Channels.Telegram.GroupTrigger), // Group trigger config
channels.WithReasoningChannelID(cfg.Channels.Telegram.ReasoningChannelID), // Reasoning chain routing
) )
return &TelegramChannel{ return &TelegramChannel{
BaseChannel: base, BaseChannel: base,
@ -466,6 +468,7 @@ func NewMatrixChannel(cfg *config.Config, msgBus *bus.MessageBus) (*MatrixChanne
matrixCfg.AllowFrom, // Allow list matrixCfg.AllowFrom, // Allow list
channels.WithMaxMessageLength(65536), // Matrix message length limit channels.WithMaxMessageLength(65536), // Matrix message length limit
channels.WithGroupTrigger(matrixCfg.GroupTrigger), channels.WithGroupTrigger(matrixCfg.GroupTrigger),
channels.WithReasoningChannelID(matrixCfg.ReasoningChannelID), // Reasoning chain routing (optional)
) )
return &MatrixChannel{ return &MatrixChannel{
@ -666,6 +669,32 @@ func (c *MatrixChannel) EditMessage(ctx context.Context, chatID, messageID, cont
} }
``` ```
#### PlaceholderCapable — Placeholder Messages
```go
// If the platform supports sending placeholder messages (e.g. "Thinking... 💭"),
// and the channel also implements MessageEditor, then Manager's preSend will
// automatically edit the placeholder into the final response on outbound.
// SendPlaceholder checks PlaceholderConfig.Enabled internally;
// returning ("", nil) means skip.
func (c *MatrixChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
cfg := c.config.Channels.Matrix.Placeholder
if !cfg.Enabled {
return "", nil
}
text := cfg.Text
if text == "" {
text = "Thinking... 💭"
}
// Call Matrix API to send placeholder message
msg, err := c.sendText(ctx, chatID, text)
if err != nil {
return "", err
}
return msg.ID, nil
}
```
#### WebhookHandler — HTTP Webhook Reception #### WebhookHandler — HTTP Webhook Reception
```go ```go
@ -746,15 +775,17 @@ When the Agent finishes processing a message, Manager's `preSend` automatically:
```go ```go
type ChannelsConfig struct { type ChannelsConfig struct {
// ... existing channels // ... existing channels
Matrix MatrixChannelConfig `yaml:"matrix" json:"matrix"` Matrix MatrixChannelConfig `json:"matrix"`
} }
type MatrixChannelConfig struct { type MatrixChannelConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"` Enabled bool `json:"enabled"`
HomeServer string `yaml:"home_server" json:"home_server"` HomeServer string `json:"home_server"`
Token string `yaml:"token" json:"token"` Token string `json:"token"`
AllowFrom []string `yaml:"allow_from" json:"allow_from"` AllowFrom []string `json:"allow_from"`
GroupTrigger GroupTriggerConfig `yaml:"group_trigger" json:"group_trigger"` GroupTrigger GroupTriggerConfig `json:"group_trigger"`
Placeholder PlaceholderConfig `json:"placeholder"`
ReasoningChannelID string `json:"reasoning_channel_id"`
} }
``` ```
@ -767,6 +798,15 @@ if m.config.Channels.Matrix.Enabled && m.config.Channels.Matrix.Token != "" {
} }
``` ```
> **Note**: If your channel has multiple modes (like WhatsApp Bridge vs Native), branch in initChannels based on config:
> ```go
> if cfg.UseNative {
> m.initChannel("whatsapp_native", "WhatsApp Native")
> } else {
> m.initChannel("whatsapp", "WhatsApp")
> }
> ```
#### Add blank import in Gateway #### Add blank import in Gateway
```go ```go
@ -882,19 +922,21 @@ BaseChannel is the shared abstraction layer for all channels, providing the foll
| `IsRunning() bool` | Atomically read running state | | `IsRunning() bool` | Atomically read running state |
| `SetRunning(bool)` | Atomically set running state | | `SetRunning(bool)` | Atomically set running state |
| `MaxMessageLength() int` | Message length limit (rune count), 0 = unlimited | | `MaxMessageLength() int` | Message length limit (rune count), 0 = unlimited |
| `ReasoningChannelID() string` | Reasoning chain routing target channel ID (empty = no routing) |
| `IsAllowed(senderID string) bool` | Legacy allow-list check (supports `"id\|username"` and `"@username"` formats) | | `IsAllowed(senderID string) bool` | Legacy allow-list check (supports `"id\|username"` and `"@username"` formats) |
| `IsAllowedSender(sender SenderInfo) bool` | New allow-list check (delegates to `identity.MatchAllowed`) | | `IsAllowedSender(sender SenderInfo) bool` | New allow-list check (delegates to `identity.MatchAllowed`) |
| `ShouldRespondInGroup(isMentioned, content) (bool, string)` | Unified group chat trigger filtering logic | | `ShouldRespondInGroup(isMentioned, content) (bool, string)` | Unified group chat trigger filtering logic |
| `HandleMessage(...)` | Unified inbound message handling: permission check → build MediaScope → auto-trigger Typing/Reaction → publish to Bus | | `HandleMessage(...)` | Unified inbound message handling: permission check → build MediaScope → auto-trigger Typing/Reaction/Placeholder → publish to Bus |
| `SetMediaStore(s) / GetMediaStore()` | MediaStore injected by Manager | | `SetMediaStore(s) / GetMediaStore()` | MediaStore injected by Manager |
| `SetPlaceholderRecorder(r) / GetPlaceholderRecorder()` | PlaceholderRecorder injected by Manager | | `SetPlaceholderRecorder(r) / GetPlaceholderRecorder()` | PlaceholderRecorder injected by Manager |
| `SetOwner(ch)` | Concrete channel reference injected by Manager (used for Typing/Reaction type assertions in HandleMessage) | | `SetOwner(ch)` | Concrete channel reference injected by Manager (used for Typing/Reaction/Placeholder type assertions in HandleMessage) |
**Functional Options**: **Functional Options**:
```go ```go
channels.WithMaxMessageLength(4096) // Set platform message length limit channels.WithMaxMessageLength(4096) // Set platform message length limit
channels.WithGroupTrigger(groupTriggerCfg) // Set group trigger configuration channels.WithGroupTrigger(groupTriggerCfg) // Set group trigger configuration
channels.WithReasoningChannelID(id) // Set reasoning chain routing target channel
``` ```
### 4.4 Factory Registry ### 4.4 Factory Registry
@ -998,7 +1040,7 @@ StartAll:
- runMediaWorker (per-channel outbound media) - runMediaWorker (per-channel outbound media)
- dispatchOutbound (route from bus to worker queues) - dispatchOutbound (route from bus to worker queues)
- dispatchOutboundMedia (route from bus to media worker queues) - dispatchOutboundMedia (route from bus to media worker queues)
- runTTLJanitor (every 10s clean up expired typing/placeholder) - runTTLJanitor (every 10s clean up expired typing/reaction/placeholder)
4. Start shared HTTP server (if configured) 4. Start shared HTTP server (if configured)
StopAll: StopAll:
@ -1206,18 +1248,20 @@ make test # Full test suite
| Sub-package | Registered Name | Optional Interfaces | | Sub-package | Registered Name | Optional Interfaces |
|-------------|----------------|-------------------| |-------------|----------------|-------------------|
| `pkg/channels/telegram/` | `"telegram"` | MessageEditor, MediaSender, TypingCapable, PlaceholderCapable | | `pkg/channels/telegram/` | `"telegram"` | TypingCapable, PlaceholderCapable, MessageEditor, MediaSender |
| `pkg/channels/discord/` | `"discord"` | MessageEditor, TypingCapable, PlaceholderCapable | | `pkg/channels/discord/` | `"discord"` | TypingCapable, PlaceholderCapable, MessageEditor, MediaSender |
| `pkg/channels/slack/` | `"slack"` | ReactionCapable | | `pkg/channels/slack/` | `"slack"` | ReactionCapable, MediaSender |
| `pkg/channels/line/` | `"line"` | WebhookHandler, HealthChecker, TypingCapable | | `pkg/channels/line/` | `"line"` | TypingCapable, MediaSender, WebhookHandler |
| `pkg/channels/onebot/` | `"onebot"` | ReactionCapable | | `pkg/channels/onebot/` | `"onebot"` | ReactionCapable, MediaSender |
| `pkg/channels/dingtalk/` | `"dingtalk"` | WebhookHandler | | `pkg/channels/dingtalk/` | `"dingtalk"` | — |
| `pkg/channels/feishu/` | `"feishu"` | WebhookHandler (architecture-specific build tags) | | `pkg/channels/feishu/` | `"feishu"` | — (architecture-specific build tags: `feishu_32.go` / `feishu_64.go`) |
| `pkg/channels/wecom/` | `"wecom"` + `"wecom_app"` | WebhookHandler | | `pkg/channels/wecom/` | `"wecom"` | WebhookHandler, HealthChecker |
| `pkg/channels/wecom/` | `"wecom_app"` | MediaSender, WebhookHandler, HealthChecker |
| `pkg/channels/qq/` | `"qq"` | — | | `pkg/channels/qq/` | `"qq"` | — |
| `pkg/channels/whatsapp/` | `"whatsapp"` | — | | `pkg/channels/whatsapp/` | `"whatsapp"` | — (Bridge mode) |
| `pkg/channels/whatsapp_native/` | `"whatsapp_native"` | — (Native whatsmeow mode) |
| `pkg/channels/maixcam/` | `"maixcam"` | — | | `pkg/channels/maixcam/` | `"maixcam"` | — |
| `pkg/channels/pico/` | `"pico"` | WebhookHandler (Pico Protocol), TypingCapable, PlaceholderCapable | | `pkg/channels/pico/` | `"pico"` | TypingCapable, PlaceholderCapable, MessageEditor, WebhookHandler |
### A.3 Interface Quick Reference ### A.3 Interface Quick Reference
@ -1231,6 +1275,7 @@ type Channel interface {
IsRunning() bool IsRunning() bool
IsAllowed(senderID string) bool IsAllowed(senderID string) bool
IsAllowedSender(sender bus.SenderInfo) bool IsAllowedSender(sender bus.SenderInfo) bool
ReasoningChannelID() string
} }
// ===== Optional ===== // ===== Optional =====
@ -1324,8 +1369,16 @@ agentLoop.Stop() // Stop Agent
1. **Media cleanup temporarily disabled**: The `ReleaseAll` call in the Agent loop is commented out (`refactor(loop): disable media cleanup to prevent premature file deletion`) because session boundaries are not yet clearly defined. TTL cleanup remains active. 1. **Media cleanup temporarily disabled**: The `ReleaseAll` call in the Agent loop is commented out (`refactor(loop): disable media cleanup to prevent premature file deletion`) because session boundaries are not yet clearly defined. TTL cleanup remains active.
2. **Feishu architecture-specific compilation**: The Feishu channel uses build tags to distinguish 32-bit and 64-bit architectures (`feishu_32.go` / `feishu_64.go`). 2. **Feishu architecture-specific compilation**: The Feishu channel uses build tags to distinguish 32-bit and 64-bit architectures (`feishu_32.go` / `feishu_64.go`). Feishu uses the SDK's WebSocket mode (not HTTP webhook), so it does not implement `WebhookHandler`.
3. **WeCom has two factories**: `"wecom"` (Bot mode) and `"wecom_app"` (App mode) are registered separately. 3. **WeCom has two factories**: `"wecom"` (Bot mode, webhook only) and `"wecom_app"` (App mode, supports MediaSender) are registered separately. Both implement `WebhookHandler` and `HealthChecker`.
4. **Pico Protocol**: `pkg/channels/pico/` implements a custom PicoClaw native protocol channel that receives messages via webhook. 4. **Pico Protocol**: `pkg/channels/pico/` implements a custom PicoClaw native protocol channel that receives messages via WebSocket webhook (`/pico/ws`).
5. **WhatsApp has two modes**: `"whatsapp"` (Bridge mode, communicates via external bridge URL) and `"whatsapp_native"` (native whatsmeow mode, connects directly to WhatsApp). Manager selects which to initialize based on `WhatsAppConfig.UseNative`.
6. **DingTalk uses Stream mode**: DingTalk uses the SDK's Stream/WebSocket mode (not HTTP webhook), so it does not implement `WebhookHandler`.
7. **PlaceholderConfig vs implementation**: `PlaceholderConfig` appears in 6 channel configs (Telegram, Discord, Slack, LINE, OneBot, Pico), but only channels that implement both `PlaceholderCapable` + `MessageEditor` (Telegram, Discord, Pico) can actually use placeholder message editing. The rest are reserved fields.
8. **ReasoningChannelID**: Most channel configs include a `reasoning_channel_id` field to route LLM reasoning/thinking output to a designated channel (WhatsApp, Telegram, Feishu, Discord, MaixCam, QQ, DingTalk, Slack, LINE, OneBot, WeCom, WeComApp). Note: `PicoConfig` does not currently expose this field. `BaseChannel` exposes this via the `WithReasoningChannelID` option and `ReasoningChannelID()` method.

View file

@ -1,7 +1,5 @@
# PicoClaw Channel System 重构:完整开发指南 # PicoClaw Channel System完整开发指南
> **分支**: `refactor/channel-system`
> **状态**: 活跃开发中(约 40 commits
> **影响范围**: `pkg/channels/`, `pkg/bus/`, `pkg/media/`, `pkg/identity/`, `cmd/picoclaw/internal/gateway/` > **影响范围**: `pkg/channels/`, `pkg/bus/`, `pkg/media/`, `pkg/identity/`, `cmd/picoclaw/internal/gateway/`
--- ---
@ -46,6 +44,8 @@ pkg/channels/
pkg/channels/ pkg/channels/
├── base.go # BaseChannel 共享抽象层 ├── base.go # BaseChannel 共享抽象层
├── interfaces.go # 可选能力接口TypingCapable, MessageEditor, ReactionCapable, PlaceholderCapable, PlaceholderRecorder ├── interfaces.go # 可选能力接口TypingCapable, MessageEditor, ReactionCapable, PlaceholderCapable, PlaceholderRecorder
├── README.md # 英文文档
├── README.zh.md # 中文文档
├── media.go # MediaSender 可选接口 ├── media.go # MediaSender 可选接口
├── webhook.go # WebhookHandler, HealthChecker 可选接口 ├── webhook.go # WebhookHandler, HealthChecker 可选接口
├── errors.go # 错误哨兵值ErrNotRunning, ErrRateLimit, ErrTemporary, ErrSendFailed ├── errors.go # 错误哨兵值ErrNotRunning, ErrRateLimit, ErrTemporary, ErrSendFailed
@ -60,7 +60,7 @@ pkg/channels/
├── discord/ ├── discord/
│ ├── init.go │ ├── init.go
│ └── discord.go │ └── discord.go
├── slack/ line/ onebot/ dingtalk/ feishu/ wecom/ qq/ whatsapp/ maixcam/ pico/ ├── slack/ line/ onebot/ dingtalk/ feishu/ wecom/ qq/ whatsapp/ whatsapp_native/ maixcam/ pico/
│ └── ... │ └── ...
pkg/bus/ pkg/bus/
@ -111,7 +111,7 @@ pkg/identity/
|------|------| |------|------|
| **子包隔离** | 每个 channel 一个独立 Go 子包,依赖 `channels` 父包提供的 `BaseChannel` 和接口 | | **子包隔离** | 每个 channel 一个独立 Go 子包,依赖 `channels` 父包提供的 `BaseChannel` 和接口 |
| **工厂注册** | 各子包通过 `init()` 自注册Manager 通过名字查找工厂,消除 import 耦合 | | **工厂注册** | 各子包通过 `init()` 自注册Manager 通过名字查找工厂,消除 import 耦合 |
| **能力发现** | 可选能力通过接口(`MediaSender`, `TypingCapable`, `ReactionCapable`, `PlaceholderCapable`, `MessageEditor`, `WebhookHandler`声明Manager 运行时类型断言发现 | | **能力发现** | 可选能力通过接口(`MediaSender`, `TypingCapable`, `ReactionCapable`, `PlaceholderCapable`, `MessageEditor`, `WebhookHandler`, `HealthChecker`声明Manager 运行时类型断言发现 |
| **结构化消息** | Peer、MessageID、SenderInfo 从 Metadata 提升为 InboundMessage 的一等字段 | | **结构化消息** | Peer、MessageID、SenderInfo 从 Metadata 提升为 InboundMessage 的一等字段 |
| **错误分类** | Channel 返回哨兵错误(`ErrRateLimit`, `ErrTemporary`Manager 据此决定重试策略 | | **错误分类** | Channel 返回哨兵错误(`ErrRateLimit`, `ErrTemporary`Manager 据此决定重试策略 |
| **集中编排** | 速率限制、消息分割、重试、Typing/Reaction/Placeholder 全部由 Manager 和 BaseChannel 统一处理Channel 只负责 Send | | **集中编排** | 速率限制、消息分割、重试、Typing/Reaction/Placeholder 全部由 Manager 和 BaseChannel 统一处理Channel 只负责 Send |
@ -145,6 +145,7 @@ pkg/identity/
| _(不存在)_ | `pkg/channels/interfaces.go` | 新增可选能力接口 | | _(不存在)_ | `pkg/channels/interfaces.go` | 新增可选能力接口 |
| _(不存在)_ | `pkg/channels/media.go` | 新增 MediaSender 接口 | | _(不存在)_ | `pkg/channels/media.go` | 新增 MediaSender 接口 |
| _(不存在)_ | `pkg/channels/webhook.go` | 新增 WebhookHandler/HealthChecker | | _(不存在)_ | `pkg/channels/webhook.go` | 新增 WebhookHandler/HealthChecker |
| _(不存在)_ | `pkg/channels/whatsapp_native/` | 新增 WhatsApp 原生模式whatsmeow |
| _(不存在)_ | `pkg/channels/split.go` | 新增消息分割(从 utils 迁入) | | _(不存在)_ | `pkg/channels/split.go` | 新增消息分割(从 utils 迁入) |
| _(不存在)_ | `pkg/bus/types.go` | 新增结构化消息类型 | | _(不存在)_ | `pkg/bus/types.go` | 新增结构化消息类型 |
| _(不存在)_ | `pkg/media/store.go` | 新增媒体文件生命周期管理 | | _(不存在)_ | `pkg/media/store.go` | 新增媒体文件生命周期管理 |
@ -220,6 +221,7 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
cfg.Channels.Telegram.AllowFrom, // 允许列表 cfg.Channels.Telegram.AllowFrom, // 允许列表
channels.WithMaxMessageLength(4096), // 平台消息长度上限 channels.WithMaxMessageLength(4096), // 平台消息长度上限
channels.WithGroupTrigger(cfg.Channels.Telegram.GroupTrigger), // 群聊触发配置 channels.WithGroupTrigger(cfg.Channels.Telegram.GroupTrigger), // 群聊触发配置
channels.WithReasoningChannelID(cfg.Channels.Telegram.ReasoningChannelID), // 思维链路由
) )
return &TelegramChannel{ return &TelegramChannel{
BaseChannel: base, BaseChannel: base,
@ -466,6 +468,7 @@ func NewMatrixChannel(cfg *config.Config, msgBus *bus.MessageBus) (*MatrixChanne
matrixCfg.AllowFrom, // 允许列表 matrixCfg.AllowFrom, // 允许列表
channels.WithMaxMessageLength(65536), // Matrix 消息长度限制 channels.WithMaxMessageLength(65536), // Matrix 消息长度限制
channels.WithGroupTrigger(matrixCfg.GroupTrigger), channels.WithGroupTrigger(matrixCfg.GroupTrigger),
channels.WithReasoningChannelID(matrixCfg.ReasoningChannelID), // 思维链路由(可选)
) )
return &MatrixChannel{ return &MatrixChannel{
@ -666,6 +669,31 @@ func (c *MatrixChannel) EditMessage(ctx context.Context, chatID, messageID, cont
} }
``` ```
#### PlaceholderCapable — 占位消息
```go
// 如果平台支持发送占位消息(如 "Thinking... 💭"),并且实现了 MessageEditor
// 则 Manager 的 preSend 会在出站时自动将占位消息编辑为最终回复。
// SendPlaceholder 内部根据 PlaceholderConfig.Enabled 决定是否发送;
// 返回 ("", nil) 表示跳过。
func (c *MatrixChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
cfg := c.config.Channels.Matrix.Placeholder
if !cfg.Enabled {
return "", nil
}
text := cfg.Text
if text == "" {
text = "Thinking... 💭"
}
// 调用 Matrix API 发送占位消息
msg, err := c.sendText(ctx, chatID, text)
if err != nil {
return "", err
}
return msg.ID, nil
}
```
#### WebhookHandler — HTTP Webhook 接收 #### WebhookHandler — HTTP Webhook 接收
```go ```go
@ -746,15 +774,17 @@ if c.owner != nil && c.placeholderRecorder != nil {
```go ```go
type ChannelsConfig struct { type ChannelsConfig struct {
// ... 现有 channels // ... 现有 channels
Matrix MatrixChannelConfig `yaml:"matrix" json:"matrix"` Matrix MatrixChannelConfig `json:"matrix"`
} }
type MatrixChannelConfig struct { type MatrixChannelConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"` Enabled bool `json:"enabled"`
HomeServer string `yaml:"home_server" json:"home_server"` HomeServer string `json:"home_server"`
Token string `yaml:"token" json:"token"` Token string `json:"token"`
AllowFrom []string `yaml:"allow_from" json:"allow_from"` AllowFrom []string `json:"allow_from"`
GroupTrigger GroupTriggerConfig `yaml:"group_trigger" json:"group_trigger"` GroupTrigger GroupTriggerConfig `json:"group_trigger"`
Placeholder PlaceholderConfig `json:"placeholder"`
ReasoningChannelID string `json:"reasoning_channel_id"`
} }
``` ```
@ -767,6 +797,15 @@ if m.config.Channels.Matrix.Enabled && m.config.Channels.Matrix.Token != "" {
} }
``` ```
> **注意**:如果你的 channel 有多种模式(如 WhatsApp Bridge vs Native需要在 initChannels 中根据配置分支:
> ```go
> if cfg.UseNative {
> m.initChannel("whatsapp_native", "WhatsApp Native")
> } else {
> m.initChannel("whatsapp", "WhatsApp")
> }
> ```
#### 在 Gateway 中添加 blank import #### 在 Gateway 中添加 blank import
```go ```go
@ -882,19 +921,21 @@ BaseChannel 是所有 channel 的共享抽象层,提供以下能力:
| `IsRunning() bool` | 原子读取运行状态 | | `IsRunning() bool` | 原子读取运行状态 |
| `SetRunning(bool)` | 原子设置运行状态 | | `SetRunning(bool)` | 原子设置运行状态 |
| `MaxMessageLength() int` | 消息长度限制rune 计数0 = 无限制 | | `MaxMessageLength() int` | 消息长度限制rune 计数0 = 无限制 |
| `ReasoningChannelID() string` | 思维链路由目标 channel ID空 = 不路由) |
| `IsAllowed(senderID string) bool` | 旧格式允许列表检查(支持 `"id\|username"``"@username"` 格式) | | `IsAllowed(senderID string) bool` | 旧格式允许列表检查(支持 `"id\|username"``"@username"` 格式) |
| `IsAllowedSender(sender SenderInfo) bool` | 新格式允许列表检查(委托给 `identity.MatchAllowed` | | `IsAllowedSender(sender SenderInfo) bool` | 新格式允许列表检查(委托给 `identity.MatchAllowed` |
| `ShouldRespondInGroup(isMentioned, content) (bool, string)` | 统一群聊触发过滤逻辑 | | `ShouldRespondInGroup(isMentioned, content) (bool, string)` | 统一群聊触发过滤逻辑 |
| `HandleMessage(...)` | 统一入站消息处理:权限检查 → 构建 MediaScope → 自动触发 Typing/Reaction → 发布到 Bus | | `HandleMessage(...)` | 统一入站消息处理:权限检查 → 构建 MediaScope → 自动触发 Typing/Reaction/Placeholder → 发布到 Bus |
| `SetMediaStore(s) / GetMediaStore()` | Manager 注入的媒体存储 | | `SetMediaStore(s) / GetMediaStore()` | Manager 注入的媒体存储 |
| `SetPlaceholderRecorder(r) / GetPlaceholderRecorder()` | Manager 注入的占位符记录器 | | `SetPlaceholderRecorder(r) / GetPlaceholderRecorder()` | Manager 注入的占位符记录器 |
| `SetOwner(ch) ` | Manager 注入的具体 channel 引用(用于 HandleMessage 内部的 Typing/Reaction 类型断言) | | `SetOwner(ch) ` | Manager 注入的具体 channel 引用(用于 HandleMessage 内部的 Typing/Reaction/Placeholder 类型断言) |
**功能选项** **功能选项**
```go ```go
channels.WithMaxMessageLength(4096) // 设置平台消息长度限制 channels.WithMaxMessageLength(4096) // 设置平台消息长度限制
channels.WithGroupTrigger(groupTriggerCfg) // 设置群聊触发配置 channels.WithGroupTrigger(groupTriggerCfg) // 设置群聊触发配置
channels.WithReasoningChannelID(id) // 设置思维链路由目标 channel
``` ```
### 4.4 工厂注册表 ### 4.4 工厂注册表
@ -998,7 +1039,7 @@ StartAll:
- runMediaWorker (per-channel 出站媒体) - runMediaWorker (per-channel 出站媒体)
- dispatchOutbound (从 bus 路由到 worker 队列) - dispatchOutbound (从 bus 路由到 worker 队列)
- dispatchOutboundMedia (从 bus 路由到 media worker 队列) - dispatchOutboundMedia (从 bus 路由到 media worker 队列)
- runTTLJanitor (每 10s 清理过期 typing/placeholder) - runTTLJanitor (每 10s 清理过期 typing/reaction/placeholder)
4. 启动共享 HTTP 服务器(如已配置) 4. 启动共享 HTTP 服务器(如已配置)
StopAll: StopAll:
@ -1206,18 +1247,20 @@ make test # 全量测试
| 子包 | 注册名 | 可选接口 | | 子包 | 注册名 | 可选接口 |
|------|--------|----------| |------|--------|----------|
| `pkg/channels/telegram/` | `"telegram"` | MessageEditor, MediaSender, TypingCapable, PlaceholderCapable | | `pkg/channels/telegram/` | `"telegram"` | TypingCapable, PlaceholderCapable, MessageEditor, MediaSender |
| `pkg/channels/discord/` | `"discord"` | MessageEditor, TypingCapable, PlaceholderCapable | | `pkg/channels/discord/` | `"discord"` | TypingCapable, PlaceholderCapable, MessageEditor, MediaSender |
| `pkg/channels/slack/` | `"slack"` | ReactionCapable | | `pkg/channels/slack/` | `"slack"` | ReactionCapable, MediaSender |
| `pkg/channels/line/` | `"line"` | WebhookHandler, HealthChecker, TypingCapable | | `pkg/channels/line/` | `"line"` | TypingCapable, MediaSender, WebhookHandler |
| `pkg/channels/onebot/` | `"onebot"` | ReactionCapable | | `pkg/channels/onebot/` | `"onebot"` | ReactionCapable, MediaSender |
| `pkg/channels/dingtalk/` | `"dingtalk"` | WebhookHandler | | `pkg/channels/dingtalk/` | `"dingtalk"` | — |
| `pkg/channels/feishu/` | `"feishu"` | WebhookHandler (架构特定 build tags) | | `pkg/channels/feishu/` | `"feishu"` | — (架构特定 build tags: `feishu_32.go` / `feishu_64.go`) |
| `pkg/channels/wecom/` | `"wecom"` + `"wecom_app"` | WebhookHandler | | `pkg/channels/wecom/` | `"wecom"` | WebhookHandler, HealthChecker |
| `pkg/channels/wecom/` | `"wecom_app"` | MediaSender, WebhookHandler, HealthChecker |
| `pkg/channels/qq/` | `"qq"` | — | | `pkg/channels/qq/` | `"qq"` | — |
| `pkg/channels/whatsapp/` | `"whatsapp"` | — | | `pkg/channels/whatsapp/` | `"whatsapp"` | — (Bridge 模式) |
| `pkg/channels/whatsapp_native/` | `"whatsapp_native"` | — (原生 whatsmeow 模式) |
| `pkg/channels/maixcam/` | `"maixcam"` | — | | `pkg/channels/maixcam/` | `"maixcam"` | — |
| `pkg/channels/pico/` | `"pico"` | WebhookHandler (Pico Protocol), TypingCapable, PlaceholderCapable | | `pkg/channels/pico/` | `"pico"` | TypingCapable, PlaceholderCapable, MessageEditor, WebhookHandler |
### A.3 接口速查表 ### A.3 接口速查表
@ -1231,6 +1274,7 @@ type Channel interface {
IsRunning() bool IsRunning() bool
IsAllowed(senderID string) bool IsAllowed(senderID string) bool
IsAllowedSender(sender bus.SenderInfo) bool IsAllowedSender(sender bus.SenderInfo) bool
ReasoningChannelID() string
} }
// ===== 可选实现 ===== // ===== 可选实现 =====
@ -1324,8 +1368,16 @@ agentLoop.Stop() // 停止 Agent
1. **媒体清理暂时禁用**Agent loop 中的 `ReleaseAll` 调用被注释掉了(`refactor(loop): disable media cleanup to prevent premature file deletion`因为会话边界尚未明确定义。TTL 清理仍然有效。 1. **媒体清理暂时禁用**Agent loop 中的 `ReleaseAll` 调用被注释掉了(`refactor(loop): disable media cleanup to prevent premature file deletion`因为会话边界尚未明确定义。TTL 清理仍然有效。
2. **Feishu 架构特定编译**Feishu channel 使用 build tags 区分 32 位和 64 位架构(`feishu_32.go` / `feishu_64.go`)。 2. **Feishu 架构特定编译**Feishu channel 使用 build tags 区分 32 位和 64 位架构(`feishu_32.go` / `feishu_64.go`)。Feishu 使用 SDK 的 WebSocket 模式(非 HTTP webhook因此不实现 `WebhookHandler`
3. **WeCom 有两个工厂**`"wecom"`Bot 模式)和 `"wecom_app"`(应用模式)分别注册 3. **WeCom 有两个工厂**`"wecom"`Bot 模式,纯 webhook`"wecom_app"`(应用模式,支持 MediaSender分别注册。两者都实现了 `WebhookHandler``HealthChecker`
4. **Pico Protocol**`pkg/channels/pico/` 实现了一个自定义的 PicoClaw 原生协议 channel通过 webhook 接收消息。 4. **Pico Protocol**`pkg/channels/pico/` 实现了一个自定义的 PicoClaw 原生协议 channel通过 WebSocket webhook (`/pico/ws`) 接收消息。
5. **WhatsApp 有两种模式**`"whatsapp"`Bridge 模式,通过外部 bridge URL 通信)和 `"whatsapp_native"`(原生 whatsmeow 模式,直接连接 WhatsApp。Manager 根据 `WhatsAppConfig.UseNative` 决定初始化哪个。
6. **DingTalk 使用 Stream 模式**DingTalk 使用 SDK 的 Stream/WebSocket 模式(非 HTTP webhook因此不实现 `WebhookHandler`
7. **PlaceholderConfig 的配置与实现**`PlaceholderConfig` 出现在 6 个 channel config 中Telegram、Discord、Slack、LINE、OneBot、Pico但只有实现了 `PlaceholderCapable` + `MessageEditor` 的 channelTelegram、Discord、Pico能真正使用占位消息编辑功能。其余 channel 的 `PlaceholderConfig` 为预留字段。
8. **ReasoningChannelID**:大多数 channel config 都包含 `reasoning_channel_id` 字段,用于将 LLM 的思维链reasoning/thinking路由到指定 channelWhatsApp、Telegram、Feishu、Discord、MaixCam、QQ、DingTalk、Slack、LINE、OneBot、WeCom、WeComApp。注意`PicoConfig` 目前不包含该字段。`BaseChannel` 通过 `WithReasoningChannelID` 选项和 `ReasoningChannelID()` 方法暴露此配置。

View file

@ -4,9 +4,17 @@
package dingtalk package dingtalk
import ( import (
"bytes"
"context" "context"
"encoding/json"
"fmt" "fmt"
"io"
"mime/multipart"
"net/http"
"os"
"path/filepath"
"sync" "sync"
"time"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/chatbot" "github.com/open-dingtalk/dingtalk-stream-sdk-go/chatbot"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/client" "github.com/open-dingtalk/dingtalk-stream-sdk-go/client"
@ -17,6 +25,7 @@ import (
"github.com/sipeed/picoclaw/pkg/identity" "github.com/sipeed/picoclaw/pkg/identity"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/utils" "github.com/sipeed/picoclaw/pkg/utils"
"github.com/sipeed/picoclaw/pkg/voice"
) )
// DingTalkChannel implements the Channel interface for DingTalk (钉钉) // DingTalkChannel implements the Channel interface for DingTalk (钉钉)
@ -29,6 +38,7 @@ type DingTalkChannel struct {
streamClient *client.StreamClient streamClient *client.StreamClient
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
transcriber *voice.GroqTranscriber
// Map to store session webhooks for each chat // Map to store session webhooks for each chat
sessionWebhooks sync.Map // chatID -> sessionWebhook sessionWebhooks sync.Map // chatID -> sessionWebhook
} }
@ -53,6 +63,11 @@ func NewDingTalkChannel(cfg config.DingTalkConfig, messageBus *bus.MessageBus) (
}, nil }, nil
} }
// SetTranscriber sets the voice transcriber for voice message processing
func (c *DingTalkChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
c.transcriber = transcriber
}
// Start initializes the DingTalk channel with Stream Mode // Start initializes the DingTalk channel with Stream Mode
func (c *DingTalkChannel) Start(ctx context.Context) error { func (c *DingTalkChannel) Start(ctx context.Context) error {
logger.InfoC("dingtalk", "Starting DingTalk channel (Stream Mode)...") logger.InfoC("dingtalk", "Starting DingTalk channel (Stream Mode)...")
@ -131,14 +146,59 @@ func (c *DingTalkChannel) onChatBotMessageReceived(
ctx context.Context, ctx context.Context,
data *chatbot.BotCallbackDataModel, data *chatbot.BotCallbackDataModel,
) ([]byte, error) { ) ([]byte, error) {
// Extract message content from Text field // Extract message content
content := data.Text.Content content := data.Text.Content
var localFiles []string
// If content is empty, try to extract from Content interface{}
if content == "" { if content == "" {
// Try to extract from Content interface{} if Text is empty
if contentMap, ok := data.Content.(map[string]any); ok { if contentMap, ok := data.Content.(map[string]any); ok {
if textContent, ok := contentMap["content"].(string); ok { // Check for different message types in the content map
if textContent, ok := contentMap["content"].(string); ok && textContent != "" {
content = textContent content = textContent
} }
// Handle image
if imgURL, ok := contentMap["imageUrl"].(string); ok && imgURL != "" {
imagePath := c.downloadImage(imgURL)
if imagePath != "" {
localFiles = append(localFiles, imagePath)
if content != "" {
content += "\n"
}
content += "[image]"
}
}
// Handle file
if fileURL, ok := contentMap["fileUrl"].(string); ok && fileURL != "" {
filePath := c.downloadFile(fileURL)
if filePath != "" {
localFiles = append(localFiles, filePath)
if content != "" {
content += "\n"
}
content += "[file]"
}
}
// Handle voice
if voiceURL, ok := contentMap["mediaUrl"].(string); ok && voiceURL != "" {
voicePath := c.downloadVoice(voiceURL)
if voicePath != "" {
localFiles = append(localFiles, voicePath)
// Try to transcribe
if c.transcriber != nil && c.transcriber.IsAvailable() {
transcribeCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
result, err := c.transcriber.Transcribe(transcribeCtx, voicePath)
if err != nil {
content += "[voice transcription failed]"
} else {
content += fmt.Sprintf("[voice: %s]", result.Text)
}
} else {
content += "[voice]"
}
}
}
} }
} }
@ -182,6 +242,7 @@ func (c *DingTalkChannel) onChatBotMessageReceived(
"sender_nick": senderNick, "sender_nick": senderNick,
"sender_id": senderID, "sender_id": senderID,
"preview": utils.Truncate(content, 50), "preview": utils.Truncate(content, 50),
"files": len(localFiles),
}) })
// Build sender info // Build sender info
@ -197,13 +258,286 @@ func (c *DingTalkChannel) onChatBotMessageReceived(
} }
// Handle the message through the base channel // Handle the message through the base channel
c.HandleMessage(ctx, peer, "", senderID, chatID, content, nil, metadata, sender) c.HandleMessage(ctx, peer, "", senderID, chatID, content, localFiles, metadata, sender)
// Return nil to indicate we've handled the message asynchronously // Return nil to indicate we've handled the message asynchronously
// The response will be sent through the message bus // The response will be sent through the message bus
return nil, nil return nil, nil
} }
// downloadImage downloads an image from DingTalk
func (c *DingTalkChannel) downloadImage(url string) string {
if url == "" {
return ""
}
// Get access token first
token, err := c.getAccessToken()
if err != nil {
logger.ErrorCF("dingtalk", "Failed to get access token", map[string]any{
"error": err.Error(),
})
return ""
}
// Construct download URL
downloadURL := fmt.Sprintf("https://oapi.dingtalk.com/media/download?access_token=%s&media_id=%s", token, url)
filename := fmt.Sprintf("dingtalk_image_%d.jpg", time.Now().Unix())
return utils.DownloadFile(downloadURL, filename, utils.DownloadOptions{
LoggerPrefix: "dingtalk",
})
}
// downloadFile downloads a file from DingTalk
func (c *DingTalkChannel) downloadFile(url string) string {
if url == "" {
return ""
}
// Get access token first
token, err := c.getAccessToken()
if err != nil {
logger.ErrorCF("dingtalk", "Failed to get access token", map[string]any{
"error": err.Error(),
})
return ""
}
// Construct download URL
downloadURL := fmt.Sprintf("https://oapi.dingtalk.com/media/download?access_token=%s&media_id=%s", token, url)
filename := fmt.Sprintf("dingtalk_file_%d", time.Now().Unix())
return utils.DownloadFile(downloadURL, filename, utils.DownloadOptions{
LoggerPrefix: "dingtalk",
})
}
// downloadVoice downloads a voice file from DingTalk
func (c *DingTalkChannel) downloadVoice(url string) string {
if url == "" {
return ""
}
// Get access token first
token, err := c.getAccessToken()
if err != nil {
logger.ErrorCF("dingtalk", "Failed to get access token", map[string]any{
"error": err.Error(),
})
return ""
}
// Construct download URL
downloadURL := fmt.Sprintf("https://oapi.dingtalk.com/media/download?access_token=%s&media_id=%s", token, url)
filename := fmt.Sprintf("dingtalk_voice_%d.amr", time.Now().Unix())
return utils.DownloadFile(downloadURL, filename, utils.DownloadOptions{
LoggerPrefix: "dingtalk",
})
}
// getAccessToken gets the DingTalk access token
func (c *DingTalkChannel) getAccessToken() (string, error) {
url := fmt.Sprintf("https://oapi.dingtalk.com/gettoken?appkey=%s&appsecret=%s", c.clientID, c.clientSecret)
resp, err := http.Get(url)
if err != nil {
return "", err
}
defer resp.Body.Close()
var result struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
AccessToken string `json:"access_token"`
}
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
return "", err
}
if result.ErrCode != 0 {
return "", fmt.Errorf("dingtalk API error: %s", result.ErrMsg)
}
return result.AccessToken, nil
}
// SendMedia implements channels.MediaSender for sending media messages
func (c *DingTalkChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
if !c.IsRunning() {
return channels.ErrNotRunning
}
store := c.GetMediaStore()
if store == nil {
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
}
// Get session webhook from storage
sessionWebhookRaw, ok := c.sessionWebhooks.Load(msg.ChatID)
if !ok {
return fmt.Errorf("no session_webhook found for chat %s, cannot send media", msg.ChatID)
}
sessionWebhook, ok := sessionWebhookRaw.(string)
if !ok {
return fmt.Errorf("invalid session_webhook type for chat %s", msg.ChatID)
}
for _, part := range msg.Parts {
localPath, err := store.Resolve(part.Ref)
if err != nil {
logger.ErrorCF("dingtalk", "Failed to resolve media ref", map[string]any{
"ref": part.Ref,
"error": err.Error(),
})
continue
}
// Upload media and send
mediaID, err := c.uploadMedia(ctx, localPath, part.Type)
if err != nil {
logger.ErrorCF("dingtalk", "Failed to upload media", map[string]any{
"type": part.Type,
"error": err.Error(),
})
continue
}
// Send media info via session webhook using direct HTTP call
caption := part.Caption
if caption == "" {
caption = part.Filename
}
var content string
switch part.Type {
case "image":
// Send image URL as markdown
content = fmt.Sprintf("![image](%s)", mediaID)
if caption != "" {
content = caption + "\n" + content
}
default:
content = fmt.Sprintf("[%s: %s](%s)", part.Type, part.Filename, mediaID)
if caption != "" {
content = caption + "\n" + content
}
}
if err := c.sendTextToWebhook(ctx, sessionWebhook, content); err != nil {
logger.ErrorCF("dingtalk", "Failed to send media message", map[string]any{
"error": err.Error(),
})
}
}
return nil
}
// sendTextToWebhook sends text to a session webhook
func (c *DingTalkChannel) sendTextToWebhook(ctx context.Context, webhook, content string) error {
payload := map[string]any{
"msgtype": "text",
"text": map[string]string{
"content": content,
},
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, webhook, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("webhook status: %d", resp.StatusCode)
}
return nil
}
// uploadMedia uploads media to DingTalk and returns the media ID
func (c *DingTalkChannel) uploadMedia(ctx context.Context, filePath, mediaType string) (string, error) {
token, err := c.getAccessToken()
if err != nil {
return "", fmt.Errorf("get access token: %w", err)
}
// Determine media type for DingTalk API
dingtalkMediaType := "file"
switch mediaType {
case "image":
dingtalkMediaType = "image"
case "voice":
dingtalkMediaType = "voice"
case "video":
dingtalkMediaType = "video"
}
// Upload to DingTalk
uploadURL := fmt.Sprintf("https://oapi.dingtalk.com/media/upload?access_token=%s&type=%s", token, dingtalkMediaType)
// Read file
file, err := os.Open(filePath)
if err != nil {
return "", fmt.Errorf("open file: %w", err)
}
defer file.Close()
// Create multipart form request
var body bytes.Buffer
writer := multipart.NewWriter(&body)
part, err := writer.CreateFormFile("media", filepath.Base(filePath))
if err != nil {
return "", fmt.Errorf("create form file: %w", err)
}
if _, err := io.Copy(part, file); err != nil {
return "", fmt.Errorf("copy file: %w", err)
}
writer.Close()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, uploadURL, &body)
if err != nil {
return "", fmt.Errorf("create request: %w", err)
}
req.Header.Set("Content-Type", writer.FormDataContentType())
resp, err := http.DefaultClient.Do(req)
if err != nil {
return "", fmt.Errorf("upload request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("upload status: %d", resp.StatusCode)
}
var uploadResp struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
MediaID string `json:"media_id"`
Type string `json:"type"`
}
if err := json.NewDecoder(resp.Body).Decode(&uploadResp); err != nil {
return "", fmt.Errorf("parse response: %w", err)
}
if uploadResp.ErrCode != 0 {
return "", fmt.Errorf("upload error: %s", uploadResp.ErrMsg)
}
return uploadResp.MediaID, nil
}
// SendDirectReply sends a direct reply using the session webhook // SendDirectReply sends a direct reply using the session webhook
func (c *DingTalkChannel) SendDirectReply(ctx context.Context, sessionWebhook, content string) error { func (c *DingTalkChannel) SendDirectReply(ctx context.Context, sessionWebhook, content string) error {
replier := chatbot.NewChatbotReplier() replier := chatbot.NewChatbotReplier()

View file

@ -3,12 +3,15 @@ package discord
import ( import (
"context" "context"
"fmt" "fmt"
"net/http"
"net/url"
"os" "os"
"strings" "strings"
"sync" "sync"
"time" "time"
"github.com/bwmarrin/discordgo" "github.com/bwmarrin/discordgo"
"github.com/gorilla/websocket"
"github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels" "github.com/sipeed/picoclaw/pkg/channels"
@ -40,6 +43,9 @@ func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordC
return nil, fmt.Errorf("failed to create discord session: %w", err) return nil, fmt.Errorf("failed to create discord session: %w", err)
} }
if err := applyDiscordProxy(session, cfg.Proxy); err != nil {
return nil, err
}
base := channels.NewBaseChannel("discord", cfg, bus, cfg.AllowFrom, base := channels.NewBaseChannel("discord", cfg, bus, cfg.AllowFrom,
channels.WithMaxMessageLength(2000), channels.WithMaxMessageLength(2000),
channels.WithGroupTrigger(cfg.GroupTrigger), channels.WithGroupTrigger(cfg.GroupTrigger),
@ -465,9 +471,43 @@ func (c *DiscordChannel) StartTyping(ctx context.Context, chatID string) (func()
func (c *DiscordChannel) downloadAttachment(url, filename string) string { func (c *DiscordChannel) downloadAttachment(url, filename string) string {
return utils.DownloadFile(url, filename, utils.DownloadOptions{ return utils.DownloadFile(url, filename, utils.DownloadOptions{
LoggerPrefix: "discord", LoggerPrefix: "discord",
ProxyURL: c.config.Proxy,
}) })
} }
func applyDiscordProxy(session *discordgo.Session, proxyAddr string) error {
var proxyFunc func(*http.Request) (*url.URL, error)
if proxyAddr != "" {
proxyURL, err := url.Parse(proxyAddr)
if err != nil {
return fmt.Errorf("invalid discord proxy URL %q: %w", proxyAddr, err)
}
proxyFunc = http.ProxyURL(proxyURL)
} else if os.Getenv("HTTP_PROXY") != "" || os.Getenv("HTTPS_PROXY") != "" {
proxyFunc = http.ProxyFromEnvironment
}
if proxyFunc == nil {
return nil
}
transport := &http.Transport{Proxy: proxyFunc}
session.Client = &http.Client{
Timeout: sendTimeout,
Transport: transport,
}
if session.Dialer != nil {
dialerCopy := *session.Dialer
dialerCopy.Proxy = proxyFunc
session.Dialer = &dialerCopy
} else {
session.Dialer = &websocket.Dialer{Proxy: proxyFunc}
}
return nil
}
// stripBotMention removes the bot mention from the message content. // stripBotMention removes the bot mention from the message content.
// Discord mentions have the format <@USER_ID> or <@!USER_ID> (with nickname). // Discord mentions have the format <@USER_ID> or <@!USER_ID> (with nickname).
func (c *DiscordChannel) stripBotMention(text string) string { func (c *DiscordChannel) stripBotMention(text string) string {

View file

@ -0,0 +1,91 @@
package discord
import (
"net/http"
"net/url"
"testing"
"github.com/bwmarrin/discordgo"
)
func TestApplyDiscordProxy_CustomProxy(t *testing.T) {
session, err := discordgo.New("Bot test-token")
if err != nil {
t.Fatalf("discordgo.New() error: %v", err)
}
if err = applyDiscordProxy(session, "http://127.0.0.1:7890"); err != nil {
t.Fatalf("applyDiscordProxy() error: %v", err)
}
req, err := http.NewRequest("GET", "https://discord.com/api/v10/gateway", nil)
if err != nil {
t.Fatalf("http.NewRequest() error: %v", err)
}
restProxy := session.Client.Transport.(*http.Transport).Proxy
restProxyURL, err := restProxy(req)
if err != nil {
t.Fatalf("rest proxy func error: %v", err)
}
if got, want := restProxyURL.String(), "http://127.0.0.1:7890"; got != want {
t.Fatalf("REST proxy = %q, want %q", got, want)
}
wsProxyURL, err := session.Dialer.Proxy(req)
if err != nil {
t.Fatalf("ws proxy func error: %v", err)
}
if got, want := wsProxyURL.String(), "http://127.0.0.1:7890"; got != want {
t.Fatalf("WS proxy = %q, want %q", got, want)
}
}
func TestApplyDiscordProxy_FromEnvironment(t *testing.T) {
t.Setenv("HTTP_PROXY", "http://127.0.0.1:8888")
t.Setenv("http_proxy", "http://127.0.0.1:8888")
t.Setenv("HTTPS_PROXY", "http://127.0.0.1:8888")
t.Setenv("https_proxy", "http://127.0.0.1:8888")
t.Setenv("ALL_PROXY", "")
t.Setenv("all_proxy", "")
t.Setenv("NO_PROXY", "")
t.Setenv("no_proxy", "")
session, err := discordgo.New("Bot test-token")
if err != nil {
t.Fatalf("discordgo.New() error: %v", err)
}
if err = applyDiscordProxy(session, ""); err != nil {
t.Fatalf("applyDiscordProxy() error: %v", err)
}
req, err := http.NewRequest("GET", "https://discord.com/api/v10/gateway", nil)
if err != nil {
t.Fatalf("http.NewRequest() error: %v", err)
}
gotURL, err := session.Dialer.Proxy(req)
if err != nil {
t.Fatalf("ws proxy func error: %v", err)
}
wantURL, err := url.Parse("http://127.0.0.1:8888")
if err != nil {
t.Fatalf("url.Parse() error: %v", err)
}
if gotURL.String() != wantURL.String() {
t.Fatalf("WS proxy = %q, want %q", gotURL.String(), wantURL.String())
}
}
func TestApplyDiscordProxy_InvalidProxyURL(t *testing.T) {
session, err := discordgo.New("Bot test-token")
if err != nil {
t.Fatalf("discordgo.New() error: %v", err)
}
if err = applyDiscordProxy(session, "://bad-proxy"); err == nil {
t.Fatal("applyDiscordProxy() expected error for invalid proxy URL, got nil")
}
}

View file

@ -1,5 +1,16 @@
package feishu package feishu
import (
"encoding/json"
"regexp"
"strings"
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
)
// mentionPlaceholderRegex matches @_user_N placeholders inserted by Feishu for mentions.
var mentionPlaceholderRegex = regexp.MustCompile(`@_user_\d+`)
// stringValue safely dereferences a *string pointer. // stringValue safely dereferences a *string pointer.
func stringValue(v *string) string { func stringValue(v *string) string {
if v == nil { if v == nil {
@ -7,3 +18,69 @@ func stringValue(v *string) string {
} }
return *v return *v
} }
// buildMarkdownCard builds a Feishu Interactive Card JSON 2.0 string with markdown content.
// JSON 2.0 cards support full CommonMark standard markdown syntax.
func buildMarkdownCard(content string) (string, error) {
card := map[string]any{
"schema": "2.0",
"body": map[string]any{
"elements": []map[string]any{
{
"tag": "markdown",
"content": content,
},
},
},
}
data, err := json.Marshal(card)
if err != nil {
return "", err
}
return string(data), nil
}
// extractJSONStringField unmarshals content as JSON and returns the value of the given string field.
// Returns "" if the content is invalid JSON or the field is missing/empty.
func extractJSONStringField(content, field string) string {
var m map[string]json.RawMessage
if err := json.Unmarshal([]byte(content), &m); err != nil {
return ""
}
raw, ok := m[field]
if !ok {
return ""
}
var s string
if err := json.Unmarshal(raw, &s); err != nil {
return ""
}
return s
}
// extractImageKey extracts the image_key from a Feishu image message content JSON.
// Format: {"image_key": "img_xxx"}
func extractImageKey(content string) string { return extractJSONStringField(content, "image_key") }
// extractFileKey extracts the file_key from a Feishu file/audio message content JSON.
// Format: {"file_key": "file_xxx", "file_name": "...", ...}
func extractFileKey(content string) string { return extractJSONStringField(content, "file_key") }
// extractFileName extracts the file_name from a Feishu file message content JSON.
func extractFileName(content string) string { return extractJSONStringField(content, "file_name") }
// stripMentionPlaceholders removes @_user_N placeholders from the text content.
// These are inserted by Feishu when users @mention someone in a message.
func stripMentionPlaceholders(content string, mentions []*larkim.MentionEvent) string {
if len(mentions) == 0 {
return content
}
for _, m := range mentions {
if m.Key != nil && *m.Key != "" {
content = strings.ReplaceAll(content, *m.Key, "")
}
}
// Also clean up any remaining @_user_N patterns
content = mentionPlaceholderRegex.ReplaceAllString(content, "")
return strings.TrimSpace(content)
}

View file

@ -0,0 +1,292 @@
package feishu
import (
"encoding/json"
"testing"
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
)
func TestExtractJSONStringField(t *testing.T) {
tests := []struct {
name string
content string
field string
want string
}{
{
name: "valid field",
content: `{"image_key": "img_v2_xxx"}`,
field: "image_key",
want: "img_v2_xxx",
},
{
name: "missing field",
content: `{"image_key": "img_v2_xxx"}`,
field: "file_key",
want: "",
},
{
name: "invalid JSON",
content: `not json at all`,
field: "image_key",
want: "",
},
{
name: "empty content",
content: "",
field: "image_key",
want: "",
},
{
name: "non-string field value",
content: `{"count": 42}`,
field: "count",
want: "",
},
{
name: "empty string value",
content: `{"image_key": ""}`,
field: "image_key",
want: "",
},
{
name: "multiple fields",
content: `{"file_key": "file_xxx", "file_name": "test.pdf"}`,
field: "file_name",
want: "test.pdf",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := extractJSONStringField(tt.content, tt.field)
if got != tt.want {
t.Errorf("extractJSONStringField(%q, %q) = %q, want %q", tt.content, tt.field, got, tt.want)
}
})
}
}
func TestExtractImageKey(t *testing.T) {
tests := []struct {
name string
content string
want string
}{
{
name: "normal",
content: `{"image_key": "img_v2_abc123"}`,
want: "img_v2_abc123",
},
{
name: "missing key",
content: `{"file_key": "file_xxx"}`,
want: "",
},
{
name: "malformed JSON",
content: `{broken`,
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := extractImageKey(tt.content)
if got != tt.want {
t.Errorf("extractImageKey(%q) = %q, want %q", tt.content, got, tt.want)
}
})
}
}
func TestExtractFileKey(t *testing.T) {
tests := []struct {
name string
content string
want string
}{
{
name: "normal",
content: `{"file_key": "file_v2_abc123", "file_name": "test.doc"}`,
want: "file_v2_abc123",
},
{
name: "missing key",
content: `{"image_key": "img_xxx"}`,
want: "",
},
{
name: "malformed JSON",
content: `not json`,
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := extractFileKey(tt.content)
if got != tt.want {
t.Errorf("extractFileKey(%q) = %q, want %q", tt.content, got, tt.want)
}
})
}
}
func TestExtractFileName(t *testing.T) {
tests := []struct {
name string
content string
want string
}{
{
name: "normal",
content: `{"file_key": "file_xxx", "file_name": "report.pdf"}`,
want: "report.pdf",
},
{
name: "missing name",
content: `{"file_key": "file_xxx"}`,
want: "",
},
{
name: "malformed JSON",
content: `{bad`,
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := extractFileName(tt.content)
if got != tt.want {
t.Errorf("extractFileName(%q) = %q, want %q", tt.content, got, tt.want)
}
})
}
}
func TestBuildMarkdownCard(t *testing.T) {
tests := []struct {
name string
content string
}{
{
name: "normal content",
content: "Hello **world**",
},
{
name: "empty content",
content: "",
},
{
name: "special characters",
content: `Code: "foo" & <bar> 'baz'`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result, err := buildMarkdownCard(tt.content)
if err != nil {
t.Fatalf("buildMarkdownCard(%q) unexpected error: %v", tt.content, err)
}
// Verify valid JSON
var parsed map[string]any
if err := json.Unmarshal([]byte(result), &parsed); err != nil {
t.Fatalf("buildMarkdownCard(%q) produced invalid JSON: %v", tt.content, err)
}
// Verify schema
if parsed["schema"] != "2.0" {
t.Errorf("schema = %v, want %q", parsed["schema"], "2.0")
}
// Verify body.elements[0].content == input
body, ok := parsed["body"].(map[string]any)
if !ok {
t.Fatal("missing body in card JSON")
}
elements, ok := body["elements"].([]any)
if !ok || len(elements) == 0 {
t.Fatal("missing or empty elements in card JSON")
}
elem, ok := elements[0].(map[string]any)
if !ok {
t.Fatal("first element is not an object")
}
if elem["tag"] != "markdown" {
t.Errorf("tag = %v, want %q", elem["tag"], "markdown")
}
if elem["content"] != tt.content {
t.Errorf("content = %v, want %q", elem["content"], tt.content)
}
})
}
}
func TestStripMentionPlaceholders(t *testing.T) {
strPtr := func(s string) *string { return &s }
tests := []struct {
name string
content string
mentions []*larkim.MentionEvent
want string
}{
{
name: "no mentions",
content: "Hello world",
mentions: nil,
want: "Hello world",
},
{
name: "single mention",
content: "@_user_1 hello",
mentions: []*larkim.MentionEvent{
{Key: strPtr("@_user_1")},
},
want: "hello",
},
{
name: "multiple mentions",
content: "@_user_1 @_user_2 hey",
mentions: []*larkim.MentionEvent{
{Key: strPtr("@_user_1")},
{Key: strPtr("@_user_2")},
},
want: "hey",
},
{
name: "empty content",
content: "",
mentions: []*larkim.MentionEvent{{Key: strPtr("@_user_1")}},
want: "",
},
{
name: "empty mentions slice",
content: "@_user_1 test",
mentions: []*larkim.MentionEvent{},
want: "@_user_1 test",
},
{
name: "mention with nil key",
content: "@_user_1 test",
mentions: []*larkim.MentionEvent{
{Key: nil},
},
want: "test",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := stripMentionPlaceholders(tt.content, tt.mentions)
if got != tt.want {
t.Errorf("stripMentionPlaceholders(%q, ...) = %q, want %q", tt.content, got, tt.want)
}
})
}
}

View file

@ -16,6 +16,8 @@ type FeishuChannel struct {
*channels.BaseChannel *channels.BaseChannel
} }
var errUnsupported = errors.New("feishu channel is not supported on 32-bit architectures")
// NewFeishuChannel returns an error on 32-bit architectures where the Feishu SDK is not supported // NewFeishuChannel returns an error on 32-bit architectures where the Feishu SDK is not supported
func NewFeishuChannel(cfg config.FeishuConfig, bus *bus.MessageBus) (*FeishuChannel, error) { func NewFeishuChannel(cfg config.FeishuConfig, bus *bus.MessageBus) (*FeishuChannel, error) {
return nil, errors.New( return nil, errors.New(
@ -25,15 +27,35 @@ func NewFeishuChannel(cfg config.FeishuConfig, bus *bus.MessageBus) (*FeishuChan
// Start is a stub method to satisfy the Channel interface // Start is a stub method to satisfy the Channel interface
func (c *FeishuChannel) Start(ctx context.Context) error { func (c *FeishuChannel) Start(ctx context.Context) error {
return nil return errUnsupported
} }
// Stop is a stub method to satisfy the Channel interface // Stop is a stub method to satisfy the Channel interface
func (c *FeishuChannel) Stop(ctx context.Context) error { func (c *FeishuChannel) Stop(ctx context.Context) error {
return nil return errUnsupported
} }
// Send is a stub method to satisfy the Channel interface // Send is a stub method to satisfy the Channel interface
func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error { func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
return errors.New("feishu channel is not supported on 32-bit architectures") return errUnsupported
}
// EditMessage is a stub method to satisfy MessageEditor
func (c *FeishuChannel) EditMessage(ctx context.Context, chatID, messageID, content string) error {
return errUnsupported
}
// SendPlaceholder is a stub method to satisfy PlaceholderCapable
func (c *FeishuChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
return "", errUnsupported
}
// ReactToMessage is a stub method to satisfy ReactionCapable
func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID string) (func(), error) {
return func() {}, errUnsupported
}
// SendMedia is a stub method to satisfy MediaSender
func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
return errUnsupported
} }

View file

@ -6,10 +6,15 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"io"
"net/http"
"os"
"path/filepath"
"sync" "sync"
"time" "sync/atomic"
lark "github.com/larksuite/oapi-sdk-go/v3" lark "github.com/larksuite/oapi-sdk-go/v3"
larkcore "github.com/larksuite/oapi-sdk-go/v3/core"
larkdispatcher "github.com/larksuite/oapi-sdk-go/v3/event/dispatcher" larkdispatcher "github.com/larksuite/oapi-sdk-go/v3/event/dispatcher"
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1" larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
larkws "github.com/larksuite/oapi-sdk-go/v3/ws" larkws "github.com/larksuite/oapi-sdk-go/v3/ws"
@ -19,6 +24,7 @@ import (
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/identity" "github.com/sipeed/picoclaw/pkg/identity"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/utils" "github.com/sipeed/picoclaw/pkg/utils"
) )
@ -28,6 +34,8 @@ type FeishuChannel struct {
client *lark.Client client *lark.Client
wsClient *larkws.Client wsClient *larkws.Client
botOpenID atomic.Value // stores string; populated lazily for @mention detection
mu sync.Mutex mu sync.Mutex
cancel context.CancelFunc cancel context.CancelFunc
} }
@ -38,11 +46,13 @@ func NewFeishuChannel(cfg config.FeishuConfig, bus *bus.MessageBus) (*FeishuChan
channels.WithReasoningChannelID(cfg.ReasoningChannelID), channels.WithReasoningChannelID(cfg.ReasoningChannelID),
) )
return &FeishuChannel{ ch := &FeishuChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
client: lark.NewClient(cfg.AppID, cfg.AppSecret), client: lark.NewClient(cfg.AppID, cfg.AppSecret),
}, nil }
ch.SetOwner(ch)
return ch, nil
} }
func (c *FeishuChannel) Start(ctx context.Context) error { func (c *FeishuChannel) Start(ctx context.Context) error {
@ -50,6 +60,13 @@ func (c *FeishuChannel) Start(ctx context.Context) error {
return fmt.Errorf("feishu app_id or app_secret is empty") return fmt.Errorf("feishu app_id or app_secret is empty")
} }
// Fetch bot open_id via API for reliable @mention detection.
if err := c.fetchBotOpenID(ctx); err != nil {
logger.ErrorCF("feishu", "Failed to fetch bot open_id, @mention detection may not work", map[string]any{
"error": err.Error(),
})
}
dispatcher := larkdispatcher.NewEventDispatcher(c.config.VerificationToken, c.config.EncryptKey). dispatcher := larkdispatcher.NewEventDispatcher(c.config.VerificationToken, c.config.EncryptKey).
OnP2MessageReceiveV1(c.handleMessageReceive) OnP2MessageReceiveV1(c.handleMessageReceive)
@ -93,46 +110,213 @@ func (c *FeishuChannel) Stop(ctx context.Context) error {
return nil return nil
} }
// Send sends a message using Interactive Card format for markdown rendering.
func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error { func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
if !c.IsRunning() { if !c.IsRunning() {
return channels.ErrNotRunning return channels.ErrNotRunning
} }
if msg.ChatID == "" { if msg.ChatID == "" {
return fmt.Errorf("chat ID is empty") return fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
} }
payload, err := json.Marshal(map[string]string{"text": msg.Content}) // Build interactive card with markdown content
cardContent, err := buildMarkdownCard(msg.Content)
if err != nil { if err != nil {
return fmt.Errorf("failed to marshal feishu content: %w", err) return fmt.Errorf("feishu send: card build failed: %w", err)
}
return c.sendCard(ctx, msg.ChatID, cardContent)
}
// EditMessage implements channels.MessageEditor.
// Uses Message.Patch to update an interactive card message.
func (c *FeishuChannel) EditMessage(ctx context.Context, chatID, messageID, content string) error {
cardContent, err := buildMarkdownCard(content)
if err != nil {
return fmt.Errorf("feishu edit: card build failed: %w", err)
}
req := larkim.NewPatchMessageReqBuilder().
MessageId(messageID).
Body(larkim.NewPatchMessageReqBodyBuilder().Content(cardContent).Build()).
Build()
resp, err := c.client.Im.V1.Message.Patch(ctx, req)
if err != nil {
return fmt.Errorf("feishu edit: %w", err)
}
if !resp.Success() {
return fmt.Errorf("feishu edit api error (code=%d msg=%s)", resp.Code, resp.Msg)
}
return nil
}
// SendPlaceholder implements channels.PlaceholderCapable.
// Sends an interactive card with placeholder text and returns its message ID.
func (c *FeishuChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
if !c.config.Placeholder.Enabled {
logger.DebugCF("feishu", "Placeholder disabled, skipping", map[string]any{
"chat_id": chatID,
})
return "", nil
}
text := c.config.Placeholder.Text
if text == "" {
text = "Thinking..."
}
cardContent, err := buildMarkdownCard(text)
if err != nil {
return "", fmt.Errorf("feishu placeholder: card build failed: %w", err)
} }
req := larkim.NewCreateMessageReqBuilder(). req := larkim.NewCreateMessageReqBuilder().
ReceiveIdType(larkim.ReceiveIdTypeChatId). ReceiveIdType(larkim.ReceiveIdTypeChatId).
Body(larkim.NewCreateMessageReqBodyBuilder(). Body(larkim.NewCreateMessageReqBodyBuilder().
ReceiveId(msg.ChatID). ReceiveId(chatID).
MsgType(larkim.MsgTypeText). MsgType(larkim.MsgTypeInteractive).
Content(string(payload)). Content(cardContent).
Uuid(fmt.Sprintf("picoclaw-%d", time.Now().UnixNano())).
Build()). Build()).
Build() Build()
resp, err := c.client.Im.V1.Message.Create(ctx, req) resp, err := c.client.Im.V1.Message.Create(ctx, req)
if err != nil { if err != nil {
return fmt.Errorf("feishu send: %w", channels.ErrTemporary) return "", fmt.Errorf("feishu placeholder send: %w", err)
} }
if !resp.Success() { if !resp.Success() {
return fmt.Errorf("feishu api error (code=%d msg=%s): %w", resp.Code, resp.Msg, channels.ErrTemporary) return "", fmt.Errorf("feishu placeholder api error (code=%d msg=%s)", resp.Code, resp.Msg)
} }
logger.DebugCF("feishu", "Feishu message sent", map[string]any{ if resp.Data != nil && resp.Data.MessageId != nil {
"chat_id": msg.ChatID, return *resp.Data.MessageId, nil
}) }
return "", nil
}
// ReactToMessage implements channels.ReactionCapable.
// Adds an "Pin" reaction and returns an undo function to remove it.
func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID string) (func(), error) {
req := larkim.NewCreateMessageReactionReqBuilder().
MessageId(messageID).
Body(larkim.NewCreateMessageReactionReqBodyBuilder().
ReactionType(larkim.NewEmojiBuilder().EmojiType("Pin").Build()).
Build()).
Build()
resp, err := c.client.Im.V1.MessageReaction.Create(ctx, req)
if err != nil {
logger.ErrorCF("feishu", "Failed to add reaction", map[string]any{
"message_id": messageID,
"error": err.Error(),
})
return func() {}, fmt.Errorf("feishu react: %w", err)
}
if !resp.Success() {
logger.ErrorCF("feishu", "Reaction API error", map[string]any{
"message_id": messageID,
"code": resp.Code,
"msg": resp.Msg,
})
return func() {}, fmt.Errorf("feishu react api error (code=%d msg=%s)", resp.Code, resp.Msg)
}
var reactionID string
if resp.Data != nil && resp.Data.ReactionId != nil {
reactionID = *resp.Data.ReactionId
}
if reactionID == "" {
return func() {}, nil
}
var undone atomic.Bool
undo := func() {
if !undone.CompareAndSwap(false, true) {
return
}
delReq := larkim.NewDeleteMessageReactionReqBuilder().
MessageId(messageID).
ReactionId(reactionID).
Build()
_, _ = c.client.Im.V1.MessageReaction.Delete(context.Background(), delReq)
}
return undo, nil
}
// SendMedia implements channels.MediaSender.
// Uploads images/files via Feishu API then sends as messages.
func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
if !c.IsRunning() {
return channels.ErrNotRunning
}
if msg.ChatID == "" {
return fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
}
store := c.GetMediaStore()
if store == nil {
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
}
for _, part := range msg.Parts {
if err := c.sendMediaPart(ctx, msg.ChatID, part, store); err != nil {
return err
}
}
return nil return nil
} }
// sendMediaPart resolves and sends a single media part.
func (c *FeishuChannel) sendMediaPart(
ctx context.Context,
chatID string,
part bus.MediaPart,
store media.MediaStore,
) error {
localPath, err := store.Resolve(part.Ref)
if err != nil {
logger.ErrorCF("feishu", "Failed to resolve media ref", map[string]any{
"ref": part.Ref,
"error": err.Error(),
})
return nil // skip this part
}
file, err := os.Open(localPath)
if err != nil {
logger.ErrorCF("feishu", "Failed to open media file", map[string]any{
"path": localPath,
"error": err.Error(),
})
return nil // skip this part
}
defer file.Close()
switch part.Type {
case "image":
err = c.sendImage(ctx, chatID, file)
default:
filename := part.Filename
if filename == "" {
filename = "file"
}
err = c.sendFile(ctx, chatID, file, filename, part.Type)
}
if err != nil {
logger.ErrorCF("feishu", "Failed to send media", map[string]any{
"type": part.Type,
"error": err.Error(),
})
return fmt.Errorf("feishu send media: %w", channels.ErrTemporary)
}
return nil
}
// --- Inbound message handling ---
func (c *FeishuChannel) handleMessageReceive(ctx context.Context, event *larkim.P2MessageReceiveV1) error { func (c *FeishuChannel) handleMessageReceive(ctx context.Context, event *larkim.P2MessageReceiveV1) error {
if event == nil || event.Event == nil || event.Event.Message == nil { if event == nil || event.Event == nil || event.Event.Message == nil {
return nil return nil
@ -151,34 +335,68 @@ func (c *FeishuChannel) handleMessageReceive(ctx context.Context, event *larkim.
senderID = "unknown" senderID = "unknown"
} }
content := extractFeishuMessageContent(message) messageType := stringValue(message.MessageType)
messageID := stringValue(message.MessageId)
rawContent := stringValue(message.Content)
// Check allowlist early to avoid downloading media for rejected senders.
// BaseChannel.HandleMessage will check again, but this avoids wasted network I/O.
senderInfo := bus.SenderInfo{
Platform: "feishu",
PlatformID: senderID,
CanonicalID: identity.BuildCanonicalID("feishu", senderID),
}
if !c.IsAllowedSender(senderInfo) {
return nil
}
// Extract content based on message type
content := extractContent(messageType, rawContent)
// Handle media messages (download and store)
var mediaRefs []string
if store := c.GetMediaStore(); store != nil && messageID != "" {
mediaRefs = c.downloadInboundMedia(ctx, chatID, messageID, messageType, rawContent, store)
}
// Append media tags to content (like Telegram does)
content = appendMediaTags(content, messageType, mediaRefs)
if content == "" { if content == "" {
content = "[empty message]" content = "[empty message]"
} }
metadata := map[string]string{} metadata := map[string]string{}
messageID := "" if messageID != "" {
if mid := stringValue(message.MessageId); mid != "" { metadata["message_id"] = messageID
messageID = mid
} }
if messageType := stringValue(message.MessageType); messageType != "" { if messageType != "" {
metadata["message_type"] = messageType metadata["message_type"] = messageType
} }
if chatType := stringValue(message.ChatType); chatType != "" { chatType := stringValue(message.ChatType)
if chatType != "" {
metadata["chat_type"] = chatType metadata["chat_type"] = chatType
} }
if sender != nil && sender.TenantKey != nil { if sender != nil && sender.TenantKey != nil {
metadata["tenant_key"] = *sender.TenantKey metadata["tenant_key"] = *sender.TenantKey
} }
chatType := stringValue(message.ChatType)
var peer bus.Peer var peer bus.Peer
if chatType == "p2p" { if chatType == "p2p" {
peer = bus.Peer{Kind: "direct", ID: senderID} peer = bus.Peer{Kind: "direct", ID: senderID}
} else { } else {
peer = bus.Peer{Kind: "group", ID: chatID} peer = bus.Peer{Kind: "group", ID: chatID}
// Check if bot was mentioned
isMentioned := c.isBotMentioned(message)
// Strip mention placeholders from content before group trigger check
if len(message.Mentions) > 0 {
content = stripMentionPlaceholders(content, message.Mentions)
}
// In group chats, apply unified group trigger filtering // In group chats, apply unified group trigger filtering
respond, cleaned := c.ShouldRespondInGroup(false, content) respond, cleaned := c.ShouldRespondInGroup(isMentioned, content)
if !respond { if !respond {
return nil return nil
} }
@ -186,22 +404,398 @@ func (c *FeishuChannel) handleMessageReceive(ctx context.Context, event *larkim.
} }
logger.InfoCF("feishu", "Feishu message received", map[string]any{ logger.InfoCF("feishu", "Feishu message received", map[string]any{
"sender_id": senderID, "sender_id": senderID,
"chat_id": chatID, "chat_id": chatID,
"preview": utils.Truncate(content, 80), "message_id": messageID,
"preview": utils.Truncate(content, 80),
}) })
senderInfo := bus.SenderInfo{ c.HandleMessage(ctx, peer, messageID, senderID, chatID, content, mediaRefs, metadata, senderInfo)
Platform: "feishu", return nil
PlatformID: senderID, }
CanonicalID: identity.BuildCanonicalID("feishu", senderID),
// --- Internal helpers ---
// fetchBotOpenID calls the Feishu bot info API to retrieve and store the bot's open_id.
func (c *FeishuChannel) fetchBotOpenID(ctx context.Context) error {
resp, err := c.client.Do(ctx, &larkcore.ApiReq{
HttpMethod: http.MethodGet,
ApiPath: "/open-apis/bot/v3/info",
SupportedAccessTokenTypes: []larkcore.AccessTokenType{larkcore.AccessTokenTypeTenant},
})
if err != nil {
return fmt.Errorf("bot info request: %w", err)
} }
if !c.IsAllowedSender(senderInfo) { var result struct {
return nil Code int `json:"code"`
Bot struct {
OpenID string `json:"open_id"`
} `json:"bot"`
}
if err := json.Unmarshal(resp.RawBody, &result); err != nil {
return fmt.Errorf("bot info parse: %w", err)
}
if result.Code != 0 {
return fmt.Errorf("bot info api error (code=%d)", result.Code)
}
if result.Bot.OpenID == "" {
return fmt.Errorf("bot info: empty open_id")
} }
c.HandleMessage(ctx, peer, messageID, senderID, chatID, content, nil, metadata, senderInfo) c.botOpenID.Store(result.Bot.OpenID)
logger.InfoCF("feishu", "Fetched bot open_id from API", map[string]any{
"open_id": result.Bot.OpenID,
})
return nil
}
// isBotMentioned checks if the bot was @mentioned in the message.
func (c *FeishuChannel) isBotMentioned(message *larkim.EventMessage) bool {
if message.Mentions == nil {
return false
}
knownID, _ := c.botOpenID.Load().(string)
if knownID == "" {
logger.DebugCF("feishu", "Bot open_id unknown, cannot detect @mention", nil)
return false
}
for _, m := range message.Mentions {
if m.Id == nil {
continue
}
if m.Id.OpenId != nil && *m.Id.OpenId == knownID {
return true
}
}
return false
}
// extractContent extracts text content from different message types.
func extractContent(messageType, rawContent string) string {
if rawContent == "" {
return ""
}
switch messageType {
case larkim.MsgTypeText:
var textPayload struct {
Text string `json:"text"`
}
if err := json.Unmarshal([]byte(rawContent), &textPayload); err == nil {
return textPayload.Text
}
return rawContent
case larkim.MsgTypePost:
// Pass raw JSON to LLM — structured rich text is more informative than flattened plain text
return rawContent
case larkim.MsgTypeImage:
// Image messages don't have text content
return ""
case larkim.MsgTypeFile, larkim.MsgTypeAudio, larkim.MsgTypeMedia:
// File/audio/video messages may have a filename
name := extractFileName(rawContent)
if name != "" {
return name
}
return ""
default:
return rawContent
}
}
// downloadInboundMedia downloads media from inbound messages and stores in MediaStore.
func (c *FeishuChannel) downloadInboundMedia(
ctx context.Context,
chatID, messageID, messageType, rawContent string,
store media.MediaStore,
) []string {
var refs []string
scope := channels.BuildMediaScope("feishu", chatID, messageID)
switch messageType {
case larkim.MsgTypeImage:
imageKey := extractImageKey(rawContent)
if imageKey == "" {
return nil
}
ref := c.downloadResource(ctx, messageID, imageKey, "image", ".jpg", store, scope)
if ref != "" {
refs = append(refs, ref)
}
case larkim.MsgTypeFile, larkim.MsgTypeAudio, larkim.MsgTypeMedia:
fileKey := extractFileKey(rawContent)
if fileKey == "" {
return nil
}
// Derive a fallback extension from the message type.
var ext string
switch messageType {
case larkim.MsgTypeAudio:
ext = ".ogg"
case larkim.MsgTypeMedia:
ext = ".mp4"
default:
ext = "" // generic file — rely on resp.FileName
}
ref := c.downloadResource(ctx, messageID, fileKey, "file", ext, store, scope)
if ref != "" {
refs = append(refs, ref)
}
}
return refs
}
// downloadResource downloads a message resource (image/file) from Feishu,
// writes it to the project media directory, and stores the reference in MediaStore.
// fallbackExt (e.g. ".jpg") is appended when the resolved filename has no extension.
func (c *FeishuChannel) downloadResource(
ctx context.Context,
messageID, fileKey, resourceType, fallbackExt string,
store media.MediaStore,
scope string,
) string {
req := larkim.NewGetMessageResourceReqBuilder().
MessageId(messageID).
FileKey(fileKey).
Type(resourceType).
Build()
resp, err := c.client.Im.V1.MessageResource.Get(ctx, req)
if err != nil {
logger.ErrorCF("feishu", "Failed to download resource", map[string]any{
"message_id": messageID,
"file_key": fileKey,
"error": err.Error(),
})
return ""
}
if !resp.Success() {
logger.ErrorCF("feishu", "Resource download api error", map[string]any{
"code": resp.Code,
"msg": resp.Msg,
})
return ""
}
if resp.File == nil {
return ""
}
// Safely close the underlying reader if it implements io.Closer (e.g. HTTP response body).
if closer, ok := resp.File.(io.Closer); ok {
defer closer.Close()
}
filename := resp.FileName
if filename == "" {
filename = fileKey
}
// If filename still has no extension, append the fallback (like Telegram's ext parameter).
if filepath.Ext(filename) == "" && fallbackExt != "" {
filename += fallbackExt
}
// Write to the shared picoclaw_media directory using a unique name to avoid collisions.
mediaDir := filepath.Join(os.TempDir(), "picoclaw_media")
if mkdirErr := os.MkdirAll(mediaDir, 0o700); mkdirErr != nil {
logger.ErrorCF("feishu", "Failed to create media directory", map[string]any{
"error": mkdirErr.Error(),
})
return ""
}
ext := filepath.Ext(filename)
localPath := filepath.Join(mediaDir, utils.SanitizeFilename(messageID+"-"+fileKey+ext))
out, err := os.Create(localPath)
if err != nil {
logger.ErrorCF("feishu", "Failed to create local file for resource", map[string]any{
"error": err.Error(),
})
return ""
}
if _, copyErr := io.Copy(out, resp.File); copyErr != nil {
out.Close()
os.Remove(localPath)
logger.ErrorCF("feishu", "Failed to write resource to file", map[string]any{
"error": copyErr.Error(),
})
return ""
}
out.Close()
ref, err := store.Store(localPath, media.MediaMeta{
Filename: filename,
Source: "feishu",
}, scope)
if err != nil {
logger.ErrorCF("feishu", "Failed to store downloaded resource", map[string]any{
"file_key": fileKey,
"error": err.Error(),
})
os.Remove(localPath)
return ""
}
return ref
}
// appendMediaTags appends media type tags to content (like Telegram's "[image: photo]").
func appendMediaTags(content, messageType string, mediaRefs []string) string {
if len(mediaRefs) == 0 {
return content
}
var tag string
switch messageType {
case larkim.MsgTypeImage:
tag = "[image: photo]"
case larkim.MsgTypeAudio:
tag = "[audio]"
case larkim.MsgTypeMedia:
tag = "[video]"
case larkim.MsgTypeFile:
tag = "[file]"
default:
tag = "[attachment]"
}
if content == "" {
return tag
}
return content + " " + tag
}
// sendCard sends an interactive card message to a chat.
func (c *FeishuChannel) sendCard(ctx context.Context, chatID, cardContent string) error {
req := larkim.NewCreateMessageReqBuilder().
ReceiveIdType(larkim.ReceiveIdTypeChatId).
Body(larkim.NewCreateMessageReqBodyBuilder().
ReceiveId(chatID).
MsgType(larkim.MsgTypeInteractive).
Content(cardContent).
Build()).
Build()
resp, err := c.client.Im.V1.Message.Create(ctx, req)
if err != nil {
return fmt.Errorf("feishu send card: %w", channels.ErrTemporary)
}
if !resp.Success() {
return fmt.Errorf("feishu api error (code=%d msg=%s): %w", resp.Code, resp.Msg, channels.ErrTemporary)
}
logger.DebugCF("feishu", "Feishu card message sent", map[string]any{
"chat_id": chatID,
})
return nil
}
// sendImage uploads an image and sends it as a message.
func (c *FeishuChannel) sendImage(ctx context.Context, chatID string, file *os.File) error {
// Upload image to get image_key
uploadReq := larkim.NewCreateImageReqBuilder().
Body(larkim.NewCreateImageReqBodyBuilder().
ImageType("message").
Image(file).
Build()).
Build()
uploadResp, err := c.client.Im.V1.Image.Create(ctx, uploadReq)
if err != nil {
return fmt.Errorf("feishu image upload: %w", err)
}
if !uploadResp.Success() {
return fmt.Errorf("feishu image upload api error (code=%d msg=%s)", uploadResp.Code, uploadResp.Msg)
}
if uploadResp.Data == nil || uploadResp.Data.ImageKey == nil {
return fmt.Errorf("feishu image upload: no image_key returned")
}
imageKey := *uploadResp.Data.ImageKey
// Send image message
content, _ := json.Marshal(map[string]string{"image_key": imageKey})
req := larkim.NewCreateMessageReqBuilder().
ReceiveIdType(larkim.ReceiveIdTypeChatId).
Body(larkim.NewCreateMessageReqBodyBuilder().
ReceiveId(chatID).
MsgType(larkim.MsgTypeImage).
Content(string(content)).
Build()).
Build()
resp, err := c.client.Im.V1.Message.Create(ctx, req)
if err != nil {
return fmt.Errorf("feishu image send: %w", err)
}
if !resp.Success() {
return fmt.Errorf("feishu image send api error (code=%d msg=%s)", resp.Code, resp.Msg)
}
return nil
}
// sendFile uploads a file and sends it as a message.
func (c *FeishuChannel) sendFile(ctx context.Context, chatID string, file *os.File, filename, fileType string) error {
// Map part type to Feishu file type
feishuFileType := "stream"
switch fileType {
case "audio":
feishuFileType = "opus"
case "video":
feishuFileType = "mp4"
}
// Upload file to get file_key
uploadReq := larkim.NewCreateFileReqBuilder().
Body(larkim.NewCreateFileReqBodyBuilder().
FileType(feishuFileType).
FileName(filename).
File(file).
Build()).
Build()
uploadResp, err := c.client.Im.V1.File.Create(ctx, uploadReq)
if err != nil {
return fmt.Errorf("feishu file upload: %w", err)
}
if !uploadResp.Success() {
return fmt.Errorf("feishu file upload api error (code=%d msg=%s)", uploadResp.Code, uploadResp.Msg)
}
if uploadResp.Data == nil || uploadResp.Data.FileKey == nil {
return fmt.Errorf("feishu file upload: no file_key returned")
}
fileKey := *uploadResp.Data.FileKey
// Send file message
content, _ := json.Marshal(map[string]string{"file_key": fileKey})
req := larkim.NewCreateMessageReqBuilder().
ReceiveIdType(larkim.ReceiveIdTypeChatId).
Body(larkim.NewCreateMessageReqBodyBuilder().
ReceiveId(chatID).
MsgType(larkim.MsgTypeFile).
Content(string(content)).
Build()).
Build()
resp, err := c.client.Im.V1.Message.Create(ctx, req)
if err != nil {
return fmt.Errorf("feishu file send: %w", err)
}
if !resp.Success() {
return fmt.Errorf("feishu file send api error (code=%d msg=%s)", resp.Code, resp.Msg)
}
return nil return nil
} }
@ -222,20 +816,3 @@ func extractFeishuSenderID(sender *larkim.EventSender) string {
return "" return ""
} }
func extractFeishuMessageContent(message *larkim.EventMessage) string {
if message == nil || message.Content == nil || *message.Content == "" {
return ""
}
if message.MessageType != nil && *message.MessageType == larkim.MsgTypeText {
var textPayload struct {
Text string `json:"text"`
}
if err := json.Unmarshal([]byte(*message.Content), &textPayload); err == nil {
return textPayload.Text
}
}
return *message.Content
}

View file

@ -0,0 +1,256 @@
//go:build amd64 || arm64 || riscv64 || mips64 || ppc64
package feishu
import (
"testing"
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
)
func TestExtractContent(t *testing.T) {
tests := []struct {
name string
messageType string
rawContent string
want string
}{
{
name: "text message",
messageType: "text",
rawContent: `{"text": "hello world"}`,
want: "hello world",
},
{
name: "text message invalid JSON",
messageType: "text",
rawContent: `not json`,
want: "not json",
},
{
name: "post message returns raw JSON",
messageType: "post",
rawContent: `{"title": "test post"}`,
want: `{"title": "test post"}`,
},
{
name: "image message returns empty",
messageType: "image",
rawContent: `{"image_key": "img_xxx"}`,
want: "",
},
{
name: "file message with filename",
messageType: "file",
rawContent: `{"file_key": "file_xxx", "file_name": "report.pdf"}`,
want: "report.pdf",
},
{
name: "file message without filename",
messageType: "file",
rawContent: `{"file_key": "file_xxx"}`,
want: "",
},
{
name: "audio message with filename",
messageType: "audio",
rawContent: `{"file_key": "file_xxx", "file_name": "recording.ogg"}`,
want: "recording.ogg",
},
{
name: "media message with filename",
messageType: "media",
rawContent: `{"file_key": "file_xxx", "file_name": "video.mp4"}`,
want: "video.mp4",
},
{
name: "unknown message type returns raw",
messageType: "sticker",
rawContent: `{"sticker_id": "sticker_xxx"}`,
want: `{"sticker_id": "sticker_xxx"}`,
},
{
name: "empty raw content",
messageType: "text",
rawContent: "",
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := extractContent(tt.messageType, tt.rawContent)
if got != tt.want {
t.Errorf("extractContent(%q, %q) = %q, want %q", tt.messageType, tt.rawContent, got, tt.want)
}
})
}
}
func TestAppendMediaTags(t *testing.T) {
tests := []struct {
name string
content string
messageType string
mediaRefs []string
want string
}{
{
name: "no refs returns content unchanged",
content: "hello",
messageType: "image",
mediaRefs: nil,
want: "hello",
},
{
name: "empty refs returns content unchanged",
content: "hello",
messageType: "image",
mediaRefs: []string{},
want: "hello",
},
{
name: "image with content",
content: "check this",
messageType: "image",
mediaRefs: []string{"ref1"},
want: "check this [image: photo]",
},
{
name: "image empty content",
content: "",
messageType: "image",
mediaRefs: []string{"ref1"},
want: "[image: photo]",
},
{
name: "audio",
content: "listen",
messageType: "audio",
mediaRefs: []string{"ref1"},
want: "listen [audio]",
},
{
name: "media/video",
content: "watch",
messageType: "media",
mediaRefs: []string{"ref1"},
want: "watch [video]",
},
{
name: "file",
content: "report.pdf",
messageType: "file",
mediaRefs: []string{"ref1"},
want: "report.pdf [file]",
},
{
name: "unknown type",
content: "something",
messageType: "sticker",
mediaRefs: []string{"ref1"},
want: "something [attachment]",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := appendMediaTags(tt.content, tt.messageType, tt.mediaRefs)
if got != tt.want {
t.Errorf(
"appendMediaTags(%q, %q, %v) = %q, want %q",
tt.content,
tt.messageType,
tt.mediaRefs,
got,
tt.want,
)
}
})
}
}
func TestExtractFeishuSenderID(t *testing.T) {
strPtr := func(s string) *string { return &s }
tests := []struct {
name string
sender *larkim.EventSender
want string
}{
{
name: "nil sender",
sender: nil,
want: "",
},
{
name: "nil sender ID",
sender: &larkim.EventSender{SenderId: nil},
want: "",
},
{
name: "userId preferred",
sender: &larkim.EventSender{
SenderId: &larkim.UserId{
UserId: strPtr("u_abc123"),
OpenId: strPtr("ou_def456"),
UnionId: strPtr("on_ghi789"),
},
},
want: "u_abc123",
},
{
name: "openId fallback",
sender: &larkim.EventSender{
SenderId: &larkim.UserId{
UserId: strPtr(""),
OpenId: strPtr("ou_def456"),
UnionId: strPtr("on_ghi789"),
},
},
want: "ou_def456",
},
{
name: "unionId fallback",
sender: &larkim.EventSender{
SenderId: &larkim.UserId{
UserId: strPtr(""),
OpenId: strPtr(""),
UnionId: strPtr("on_ghi789"),
},
},
want: "on_ghi789",
},
{
name: "all empty strings",
sender: &larkim.EventSender{
SenderId: &larkim.UserId{
UserId: strPtr(""),
OpenId: strPtr(""),
UnionId: strPtr(""),
},
},
want: "",
},
{
name: "nil userId pointer falls through",
sender: &larkim.EventSender{
SenderId: &larkim.UserId{
UserId: nil,
OpenId: strPtr("ou_def456"),
UnionId: nil,
},
},
want: "ou_def456",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := extractFeishuSenderID(tt.sender)
if got != tt.want {
t.Errorf("extractFeishuSenderID() = %q, want %q", got, tt.want)
}
})
}
}

View file

@ -10,6 +10,7 @@ import (
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"os"
"strings" "strings"
"sync" "sync"
"time" "time"
@ -45,11 +46,13 @@ type replyTokenEntry struct {
type LINEChannel struct { type LINEChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.LINEConfig config config.LINEConfig
botUserID string // Bot's user ID infoClient *http.Client // for bot info lookups (short timeout)
botBasicID string // Bot's basic ID (e.g. @216ru...) apiClient *http.Client // for messaging API calls
botDisplayName string // Bot's display name for text-based mention detection botUserID string // Bot's user ID
replyTokens sync.Map // chatID -> replyTokenEntry botBasicID string // Bot's basic ID (e.g. @216ru...)
quoteTokens sync.Map // chatID -> quoteToken (string) botDisplayName string // Bot's display name for text-based mention detection
replyTokens sync.Map // chatID -> replyTokenEntry
quoteTokens sync.Map // chatID -> quoteToken (string)
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
} }
@ -69,6 +72,8 @@ func NewLINEChannel(cfg config.LINEConfig, messageBus *bus.MessageBus) (*LINECha
return &LINEChannel{ return &LINEChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
infoClient: &http.Client{Timeout: 10 * time.Second},
apiClient: &http.Client{Timeout: 30 * time.Second},
}, nil }, nil
} }
@ -104,8 +109,7 @@ func (c *LINEChannel) fetchBotInfo() error {
} }
req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken) req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken)
client := &http.Client{Timeout: 10 * time.Second} resp, err := c.infoClient.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return err return err
} }
@ -325,7 +329,11 @@ func (c *LINEChannel) processEvent(event lineEvent) {
content = "[video]" content = "[video]"
} }
case "file": case "file":
content = "[file]" localPath := c.downloadContent(msg.ID, "file")
if localPath != "" {
mediaPaths = append(mediaPaths, storeMedia(localPath, "file"))
content = "[file]"
}
case "sticker": case "sticker":
content = "[sticker]" content = "[sticker]"
default: default:
@ -514,9 +522,7 @@ func (c *LINEChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
} }
// SendMedia implements the channels.MediaSender interface. // SendMedia implements the channels.MediaSender interface.
// LINE requires media to be accessible via public URL; since we only have local files, // Uploads media to LINE and sends as media messages.
// we fall back to sending a text message with the filename/caption.
// For full support, an external file hosting service would be needed.
func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error { func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
if !c.IsRunning() { if !c.IsRunning() {
return channels.ErrNotRunning return channels.ErrNotRunning
@ -527,15 +533,67 @@ func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessag
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed) return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
} }
// LINE Messaging API requires publicly accessible URLs for media messages.
// Since we only have local file paths, send caption text as fallback.
for _, part := range msg.Parts { for _, part := range msg.Parts {
caption := part.Caption // Resolve local file path
if caption == "" { localPath, err := store.Resolve(part.Ref)
caption = fmt.Sprintf("[%s: %s]", part.Type, part.Filename) if err != nil {
logger.ErrorCF("line", "Failed to resolve media ref", map[string]any{
"ref": part.Ref,
"error": err.Error(),
})
continue
} }
if err := c.sendPush(ctx, msg.ChatID, caption, ""); err != nil { // Upload media and send as media message
var mediaID string
switch part.Type {
case "image":
id, err := c.uploadMedia(ctx, localPath, "image", part.Filename)
if err != nil {
logger.ErrorCF("line", "Failed to upload image", map[string]any{
"error": err.Error(),
})
continue
}
mediaID = id
case "video":
id, err := c.uploadMedia(ctx, localPath, "video", part.Filename)
if err != nil {
logger.ErrorCF("line", "Failed to upload video", map[string]any{
"error": err.Error(),
})
continue
}
mediaID = id
case "audio":
id, err := c.uploadMedia(ctx, localPath, "audio", part.Filename)
if err != nil {
logger.ErrorCF("line", "Failed to upload audio", map[string]any{
"error": err.Error(),
})
continue
}
mediaID = id
default:
// For unknown types, treat as file
id, err := c.uploadMedia(ctx, localPath, "file", part.Filename)
if err != nil {
logger.ErrorCF("line", "Failed to upload file", map[string]any{
"error": err.Error(),
})
continue
}
mediaID = id
}
// Build caption
caption := part.Caption
if caption == "" && part.Filename != "" {
caption = part.Filename
}
// Send media message
if err := c.sendMediaMessage(ctx, msg.ChatID, part.Type, mediaID, caption); err != nil {
return err return err
} }
} }
@ -543,6 +601,112 @@ func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessag
return nil return nil
} }
// uploadMedia uploads media to LINE and returns the media ID.
func (c *LINEChannel) uploadMedia(ctx context.Context, filePath, mediaType, filename string) (string, error) {
// First, get upload URL from LINE
uploadURL := lineDataAPIBase + "/bot/message/upload"
req, err := http.NewRequestWithContext(ctx, http.MethodGet, uploadURL, nil)
if err != nil {
return "", fmt.Errorf("create request: %w", err)
}
req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken)
resp, err := c.apiClient.Do(req)
if err != nil {
return "", fmt.Errorf("get upload URL: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("get upload URL status: %d", resp.StatusCode)
}
var uploadResp struct {
UploadURL string `json:"uploadUrl"`
ID string `json:"messageId"`
}
if err := json.NewDecoder(resp.Body).Decode(&uploadResp); err != nil {
return "", fmt.Errorf("parse upload response: %w", err)
}
if uploadResp.UploadURL == "" || uploadResp.ID == "" {
return "", fmt.Errorf("empty upload URL or message ID")
}
// Upload the file to the provided URL
file, err := os.Open(filePath)
if err != nil {
return "", fmt.Errorf("open file: %w", err)
}
defer file.Close()
// Determine content type
contentType := "application/octet-stream"
switch mediaType {
case "image":
contentType = "image/jpeg"
case "video":
contentType = "video/mp4"
case "audio":
contentType = "audio/mp4"
}
uploadReq, err := http.NewRequestWithContext(ctx, http.MethodPost, uploadResp.UploadURL, file)
if err != nil {
return "", fmt.Errorf("create upload request: %w", err)
}
uploadReq.Header.Set("Content-Type", contentType)
uploadResp2, err := http.DefaultClient.Do(uploadReq)
if err != nil {
return "", fmt.Errorf("upload media: %w", err)
}
defer uploadResp2.Body.Close()
if uploadResp2.StatusCode != http.StatusOK {
return "", fmt.Errorf("upload media status: %d", uploadResp2.StatusCode)
}
return uploadResp.ID, nil
}
// sendMediaMessage sends a media message (image/video/audio/file) to LINE.
func (c *LINEChannel) sendMediaMessage(ctx context.Context, chatID, mediaType, mediaID, caption string) error {
var msgType string
switch mediaType {
case "image":
msgType = "image"
case "video":
msgType = "video"
case "audio":
msgType = "audio"
case "file":
msgType = "file"
default:
msgType = "file"
}
content := map[string]string{
"type": msgType,
"id": mediaID,
}
if caption != "" {
content["originalContentUrl"] = caption // LINE uses this field for caption in media messages
}
payload := map[string]any{
"to": chatID,
"messages": []map[string]string{{
"type": msgType,
"id": mediaID,
"originalContentUrl": caption,
}},
}
return c.callAPI(ctx, linePushEndpoint, payload)
}
// buildTextMessage creates a text message object, optionally with quoteToken. // buildTextMessage creates a text message object, optionally with quoteToken.
func buildTextMessage(content, quoteToken string) map[string]string { func buildTextMessage(content, quoteToken string) map[string]string {
msg := map[string]string{ msg := map[string]string{
@ -644,8 +808,7 @@ func (c *LINEChannel) callAPI(ctx context.Context, endpoint string, payload any)
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken) req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken)
client := &http.Client{Timeout: 30 * time.Second} resp, err := c.apiClient.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }

View file

@ -255,6 +255,10 @@ func (m *Manager) initChannels() error {
m.initChannel("wecom", "WeCom") m.initChannel("wecom", "WeCom")
} }
if m.config.Channels.WeComAIBot.Enabled && m.config.Channels.WeComAIBot.Token != "" {
m.initChannel("wecom_aibot", "WeCom AI Bot")
}
if m.config.Channels.WeComApp.Enabled && m.config.Channels.WeComApp.CorpID != "" { if m.config.Channels.WeComApp.Enabled && m.config.Channels.WeComApp.CorpID != "" {
m.initChannel("wecom_app", "WeCom App") m.initChannel("wecom_app", "WeCom App")
} }
@ -539,86 +543,88 @@ func (m *Manager) sendWithRetry(ctx context.Context, name string, w *channelWork
}) })
} }
func (m *Manager) dispatchOutbound(ctx context.Context) { func dispatchLoop[M any](
logger.InfoC("channels", "Outbound dispatcher started") ctx context.Context,
m *Manager,
subscribe func(context.Context) (M, bool),
getChannel func(M) string,
enqueue func(context.Context, *channelWorker, M) bool,
startMsg, stopMsg, unknownMsg, noWorkerMsg string,
) {
logger.InfoC("channels", startMsg)
for { for {
msg, ok := m.bus.SubscribeOutbound(ctx) msg, ok := subscribe(ctx)
if !ok { if !ok {
logger.InfoC("channels", "Outbound dispatcher stopped") logger.InfoC("channels", stopMsg)
return return
} }
channel := getChannel(msg)
// Silently skip internal channels // Silently skip internal channels
if constants.IsInternalChannel(msg.Channel) { if constants.IsInternalChannel(channel) {
continue continue
} }
m.mu.RLock() m.mu.RLock()
_, exists := m.channels[msg.Channel] _, exists := m.channels[channel]
w, wExists := m.workers[msg.Channel] w, wExists := m.workers[channel]
m.mu.RUnlock() m.mu.RUnlock()
if !exists { if !exists {
logger.WarnCF("channels", "Unknown channel for outbound message", map[string]any{ logger.WarnCF("channels", unknownMsg, map[string]any{"channel": channel})
"channel": msg.Channel,
})
continue continue
} }
if wExists && w != nil { if wExists && w != nil {
select { if !enqueue(ctx, w, msg) {
case w.queue <- msg:
case <-ctx.Done():
return return
} }
} else if exists { } else if exists {
logger.WarnCF("channels", "Channel has no active worker, skipping message", map[string]any{ logger.WarnCF("channels", noWorkerMsg, map[string]any{"channel": channel})
"channel": msg.Channel,
})
} }
} }
} }
func (m *Manager) dispatchOutbound(ctx context.Context) {
dispatchLoop(
ctx, m,
m.bus.SubscribeOutbound,
func(msg bus.OutboundMessage) string { return msg.Channel },
func(ctx context.Context, w *channelWorker, msg bus.OutboundMessage) bool {
select {
case w.queue <- msg:
return true
case <-ctx.Done():
return false
}
},
"Outbound dispatcher started",
"Outbound dispatcher stopped",
"Unknown channel for outbound message",
"Channel has no active worker, skipping message",
)
}
func (m *Manager) dispatchOutboundMedia(ctx context.Context) { func (m *Manager) dispatchOutboundMedia(ctx context.Context) {
logger.InfoC("channels", "Outbound media dispatcher started") dispatchLoop(
ctx, m,
for { m.bus.SubscribeOutboundMedia,
msg, ok := m.bus.SubscribeOutboundMedia(ctx) func(msg bus.OutboundMediaMessage) string { return msg.Channel },
if !ok { func(ctx context.Context, w *channelWorker, msg bus.OutboundMediaMessage) bool {
logger.InfoC("channels", "Outbound media dispatcher stopped")
return
}
// Silently skip internal channels
if constants.IsInternalChannel(msg.Channel) {
continue
}
m.mu.RLock()
_, exists := m.channels[msg.Channel]
w, wExists := m.workers[msg.Channel]
m.mu.RUnlock()
if !exists {
logger.WarnCF("channels", "Unknown channel for outbound media message", map[string]any{
"channel": msg.Channel,
})
continue
}
if wExists && w != nil {
select { select {
case w.mediaQueue <- msg: case w.mediaQueue <- msg:
return true
case <-ctx.Done(): case <-ctx.Done():
return return false
} }
} else if exists { },
logger.WarnCF("channels", "Channel has no active worker, skipping media message", map[string]any{ "Outbound media dispatcher started",
"channel": msg.Channel, "Outbound media dispatcher stopped",
}) "Unknown channel for outbound media message",
} "Channel has no active worker, skipping media message",
} )
} }
// runMediaWorker processes outbound media messages for a single channel. // runMediaWorker processes outbound media messages for a single channel.

View file

@ -274,13 +274,12 @@ func TestWorkerRateLimiter(t *testing.T) {
limiter: rate.NewLimiter(2, 1), limiter: rate.NewLimiter(2, 1),
} }
ctx, cancel := context.WithCancel(context.Background()) ctx := t.Context()
defer cancel()
go m.runWorker(ctx, "test", w) go m.runWorker(ctx, "test", w)
// Enqueue 4 messages // Enqueue 4 messages
for i := 0; i < 4; i++ { for i := range 4 {
w.queue <- bus.OutboundMessage{Channel: "test", ChatID: "1", Content: fmt.Sprintf("msg%d", i)} w.queue <- bus.OutboundMessage{Channel: "test", ChatID: "1", Content: fmt.Sprintf("msg%d", i)}
} }
@ -352,8 +351,7 @@ func TestRunWorker_MessageSplitting(t *testing.T) {
limiter: rate.NewLimiter(rate.Inf, 1), limiter: rate.NewLimiter(rate.Inf, 1),
} }
ctx, cancel := context.WithCancel(context.Background()) ctx := t.Context()
defer cancel()
go m.runWorker(ctx, "test", w) go m.runWorker(ctx, "test", w)
@ -576,7 +574,7 @@ func TestRecordPlaceholder_ConcurrentSafe(t *testing.T) {
m := newTestManager() m := newTestManager()
var wg sync.WaitGroup var wg sync.WaitGroup
for i := 0; i < 100; i++ { for i := range 100 {
wg.Add(1) wg.Add(1)
go func(i int) { go func(i int) {
defer wg.Done() defer wg.Done()
@ -591,7 +589,7 @@ func TestRecordTypingStop_ConcurrentSafe(t *testing.T) {
m := newTestManager() m := newTestManager()
var wg sync.WaitGroup var wg sync.WaitGroup
for i := 0; i < 100; i++ { for i := range 100 {
wg.Add(1) wg.Add(1)
go func(i int) { go func(i int) {
defer wg.Done() defer wg.Done()
@ -834,7 +832,7 @@ func TestLazyWorkerCreation(t *testing.T) {
func TestBuildMediaScope_FastIDUniqueness(t *testing.T) { func TestBuildMediaScope_FastIDUniqueness(t *testing.T) {
seen := make(map[string]bool) seen := make(map[string]bool)
for i := 0; i < 1000; i++ { for range 1000 {
scope := BuildMediaScope("test", "chat1", "") scope := BuildMediaScope("test", "chat1", "")
if seen[scope] { if seen[scope] {
t.Fatalf("duplicate scope generated: %s", scope) t.Fatalf("duplicate scope generated: %s", scope)

View file

@ -337,10 +337,7 @@ func (c *OneBotChannel) sendAPIRequest(action string, params any, timeout time.D
} }
func (c *OneBotChannel) reconnectLoop() { func (c *OneBotChannel) reconnectLoop() {
interval := time.Duration(c.config.ReconnectInterval) * time.Second interval := max(time.Duration(c.config.ReconnectInterval)*time.Second, 5*time.Second)
if interval < 5*time.Second {
interval = 5 * time.Second
}
for { for {
select { select {

View file

@ -292,8 +292,8 @@ func (c *PicoChannel) authenticate(r *http.Request) bool {
// Check Authorization header // Check Authorization header
auth := r.Header.Get("Authorization") auth := r.Header.Get("Authorization")
if strings.HasPrefix(auth, "Bearer ") { if after, ok := strings.CutPrefix(auth, "Bearer "); ok {
if strings.TrimPrefix(auth, "Bearer ") == token { if after == token {
return true return true
} }
} }

View file

@ -23,10 +23,7 @@ func SplitMessage(content string, maxLen int) []string {
var messages []string var messages []string
// Dynamic buffer: 10% of maxLen, but at least 50 chars if possible // Dynamic buffer: 10% of maxLen, but at least 50 chars if possible
codeBlockBuffer := maxLen / 10 codeBlockBuffer := max(maxLen/10, 50)
if codeBlockBuffer < 50 {
codeBlockBuffer = 50
}
if codeBlockBuffer > maxLen/2 { if codeBlockBuffer > maxLen/2 {
codeBlockBuffer = maxLen / 2 codeBlockBuffer = maxLen / 2
} }
@ -40,10 +37,7 @@ func SplitMessage(content string, maxLen int) []string {
} }
// Effective split point: maxLen minus buffer, to leave room for code blocks // Effective split point: maxLen minus buffer, to leave room for code blocks
effectiveLimit := maxLen - codeBlockBuffer effectiveLimit := max(maxLen-codeBlockBuffer, maxLen/2)
if effectiveLimit < maxLen/2 {
effectiveLimit = maxLen / 2
}
end := start + effectiveLimit end := start + effectiveLimit
@ -85,10 +79,9 @@ func SplitMessage(content string, maxLen int) []string {
// If we have a reasonable amount of content after the header, split inside // If we have a reasonable amount of content after the header, split inside
if msgEnd > headerEndIdx+20 { if msgEnd > headerEndIdx+20 {
// Find a better split point closer to maxLen // Find a better split point closer to maxLen
innerLimit := start + maxLen - 5 // Leave room for "\n```" innerLimit := min(
if innerLimit > totalLen { // Leave room for "\n```"
innerLimit = totalLen start+maxLen-5, totalLen)
}
betterEnd := findLastNewlineInRange(runes, start, innerLimit, 200) betterEnd := findLastNewlineInRange(runes, start, innerLimit, 200)
if betterEnd > headerEndIdx { if betterEnd > headerEndIdx {
msgEnd = betterEnd msgEnd = betterEnd
@ -117,10 +110,7 @@ func SplitMessage(content string, maxLen int) []string {
if unclosedIdx-start > 20 { if unclosedIdx-start > 20 {
msgEnd = unclosedIdx msgEnd = unclosedIdx
} else { } else {
splitAt := start + maxLen - 5 splitAt := min(start+maxLen-5, totalLen)
if splitAt > totalLen {
splitAt = totalLen
}
chunk := strings.TrimRight(string(runes[start:splitAt]), " \t\n\r") + "\n```" chunk := strings.TrimRight(string(runes[start:splitAt]), " \t\n\r") + "\n```"
messages = append(messages, chunk) messages = append(messages, chunk)
remaining := strings.TrimSpace(header + "\n" + string(runes[splitAt:totalLen])) remaining := strings.TrimSpace(header + "\n" + string(runes[splitAt:totalLen]))
@ -196,10 +186,7 @@ func findNewlineFrom(runes []rune, from int) int {
// findLastNewlineInRange finds the last newline within the last searchWindow runes // findLastNewlineInRange finds the last newline within the last searchWindow runes
// of the range runes[start:end]. Returns the absolute index or start-1 (indicating not found). // of the range runes[start:end]. Returns the absolute index or start-1 (indicating not found).
func findLastNewlineInRange(runes []rune, start, end, searchWindow int) int { func findLastNewlineInRange(runes []rune, start, end, searchWindow int) int {
searchStart := end - searchWindow searchStart := max(end-searchWindow, start)
if searchStart < start {
searchStart = start
}
for i := end - 1; i >= searchStart; i-- { for i := end - 1; i >= searchStart; i-- {
if runes[i] == '\n' { if runes[i] == '\n' {
return i return i
@ -211,10 +198,7 @@ func findLastNewlineInRange(runes []rune, start, end, searchWindow int) int {
// findLastSpaceInRange finds the last space/tab within the last searchWindow runes // findLastSpaceInRange finds the last space/tab within the last searchWindow runes
// of the range runes[start:end]. Returns the absolute index or start-1 (indicating not found). // of the range runes[start:end]. Returns the absolute index or start-1 (indicating not found).
func findLastSpaceInRange(runes []rune, start, end, searchWindow int) int { func findLastSpaceInRange(runes []rune, start, end, searchWindow int) int {
searchStart := end - searchWindow searchStart := max(end-searchWindow, start)
if searchStart < start {
searchStart = start
}
for i := end - 1; i >= searchStart; i-- { for i := end - 1; i >= searchStart; i-- {
if runes[i] == ' ' || runes[i] == '\t' { if runes[i] == ' ' || runes[i] == '\t' {
return i return i

View file

@ -7,12 +7,12 @@ import (
"net/url" "net/url"
"os" "os"
"regexp" "regexp"
"slices"
"strconv" "strconv"
"strings" "strings"
"time" "time"
"github.com/mymmrac/telego" "github.com/mymmrac/telego"
"github.com/mymmrac/telego/telegohandler"
th "github.com/mymmrac/telego/telegohandler" th "github.com/mymmrac/telego/telegohandler"
tu "github.com/mymmrac/telego/telegoutil" tu "github.com/mymmrac/telego/telegoutil"
@ -41,7 +41,7 @@ var (
type TelegramChannel struct { type TelegramChannel struct {
*channels.BaseChannel *channels.BaseChannel
bot *telego.Bot bot *telego.Bot
bh *telegohandler.BotHandler bh *th.BotHandler
commands TelegramCommander commands TelegramCommander
config *config.Config config *config.Config
chatIDs map[string]int64 chatIDs map[string]int64
@ -72,6 +72,10 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
})) }))
} }
if baseURL := strings.TrimRight(strings.TrimSpace(telegramCfg.BaseURL), "/"); baseURL != "" {
opts = append(opts, telego.WithAPIServer(baseURL))
}
bot, err := telego.NewBot(telegramCfg.Token, opts...) bot, err := telego.NewBot(telegramCfg.Token, opts...)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to create telegram bot: %w", err) return nil, fmt.Errorf("failed to create telegram bot: %w", err)
@ -101,6 +105,12 @@ func (c *TelegramChannel) Start(ctx context.Context) error {
c.ctx, c.cancel = context.WithCancel(ctx) c.ctx, c.cancel = context.WithCancel(ctx)
if err := c.initBotCommands(c.ctx); err != nil {
logger.WarnCF("telegram", "Failed to initialize bot commands", map[string]any{
"error": err.Error(),
})
}
updates, err := c.bot.UpdatesViaLongPolling(c.ctx, &telego.GetUpdatesParams{ updates, err := c.bot.UpdatesViaLongPolling(c.ctx, &telego.GetUpdatesParams{
Timeout: 30, Timeout: 30,
}) })
@ -109,20 +119,19 @@ func (c *TelegramChannel) Start(ctx context.Context) error {
return fmt.Errorf("failed to start long polling: %w", err) return fmt.Errorf("failed to start long polling: %w", err)
} }
bh, err := telegohandler.NewBotHandler(c.bot, updates) bh, err := th.NewBotHandler(c.bot, updates)
if err != nil { if err != nil {
c.cancel() c.cancel()
return fmt.Errorf("failed to create bot handler: %w", err) return fmt.Errorf("failed to create bot handler: %w", err)
} }
c.bh = bh c.bh = bh
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
c.commands.Help(ctx, message)
return nil
}, th.CommandEqual("help"))
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
return c.commands.Start(ctx, message) return c.commands.Start(ctx, message)
}, th.CommandEqual("start")) }, th.CommandEqual("start"))
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
return c.commands.Help(ctx, message)
}, th.CommandEqual("help"))
bh.HandleMessage(func(ctx *th.Context, message telego.Message) error { bh.HandleMessage(func(ctx *th.Context, message telego.Message) error {
return c.commands.Show(ctx, message) return c.commands.Show(ctx, message)
@ -141,7 +150,13 @@ func (c *TelegramChannel) Start(ctx context.Context) error {
"username": c.bot.Username(), "username": c.bot.Username(),
}) })
go bh.Start() go func() {
if err = bh.Start(); err != nil {
logger.ErrorCF("telegram", "Bot handler failed", map[string]any{
"error": err.Error(),
})
}
}()
return nil return nil
} }
@ -152,7 +167,7 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
// Stop the bot handler // Stop the bot handler
if c.bh != nil { if c.bh != nil {
c.bh.Stop() _ = c.bh.StopWithContext(ctx)
} }
// Cancel our context (stops long polling) // Cancel our context (stops long polling)
@ -163,6 +178,51 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
return nil return nil
} }
func (c *TelegramChannel) initBotCommands(ctx context.Context) error {
currentCommands, err := c.bot.GetMyCommands(ctx, &telego.GetMyCommandsParams{
Scope: tu.ScopeDefault(),
})
if err != nil {
return fmt.Errorf("get commands: %w", err)
}
commands := []telego.BotCommand{
{
Command: "start",
Description: "Start the bot",
},
{
Command: "help",
Description: "Show a help message",
},
{
Command: "show",
Description: "Show current configuration",
},
{
Command: "list",
Description: "List available options",
},
}
// Setting commands on each start will hit the rate limit very quickly, that's why we check if an update is needed
if !slices.Equal(currentCommands, commands) {
logger.InfoC("telegram", "Updating bot commands")
err = c.bot.SetMyCommands(ctx, &telego.SetMyCommandsParams{
Commands: commands,
Scope: tu.ScopeDefault(),
})
if err != nil {
return fmt.Errorf("set commands: %w", err)
}
} else {
logger.DebugC("telegram", "Bot commands are up to date")
}
return nil
}
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error { func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
if !c.IsRunning() { if !c.IsRunning() {
return channels.ErrNotRunning return channels.ErrNotRunning

1014
pkg/channels/wecom/aibot.go Normal file

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,210 @@
package wecom
import (
"context"
"testing"
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/config"
)
func TestNewWeComAIBotChannel(t *testing.T) {
t.Run("success with valid config", func(t *testing.T) {
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
EncodingAESKey: "testkey1234567890123456789012345678901234567",
WebhookPath: "/webhook/test",
}
messageBus := bus.NewMessageBus()
ch, err := NewWeComAIBotChannel(cfg, messageBus)
if err != nil {
t.Fatalf("Expected no error, got %v", err)
}
if ch == nil {
t.Fatal("Expected channel to be created")
}
if ch.Name() != "wecom_aibot" {
t.Errorf("Expected name 'wecom_aibot', got '%s'", ch.Name())
}
})
t.Run("error with missing token", func(t *testing.T) {
cfg := config.WeComAIBotConfig{
Enabled: true,
EncodingAESKey: "testkey1234567890123456789012345678901234567",
}
messageBus := bus.NewMessageBus()
_, err := NewWeComAIBotChannel(cfg, messageBus)
if err == nil {
t.Fatal("Expected error for missing token, got nil")
}
})
t.Run("error with missing encoding key", func(t *testing.T) {
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
}
messageBus := bus.NewMessageBus()
_, err := NewWeComAIBotChannel(cfg, messageBus)
if err == nil {
t.Fatal("Expected error for missing encoding key, got nil")
}
})
}
func TestWeComAIBotChannelStartStop(t *testing.T) {
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
EncodingAESKey: "testkey1234567890123456789012345678901234567",
}
messageBus := bus.NewMessageBus()
ch, err := NewWeComAIBotChannel(cfg, messageBus)
if err != nil {
t.Fatalf("Failed to create channel: %v", err)
}
ctx := context.Background()
// Test Start
if err := ch.Start(ctx); err != nil {
t.Fatalf("Failed to start channel: %v", err)
}
if !ch.IsRunning() {
t.Error("Expected channel to be running")
}
// Test Stop
if err := ch.Stop(ctx); err != nil {
t.Fatalf("Failed to stop channel: %v", err)
}
if ch.IsRunning() {
t.Error("Expected channel to be stopped")
}
}
func TestWeComAIBotChannelWebhookPath(t *testing.T) {
t.Run("default path", func(t *testing.T) {
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
EncodingAESKey: "testkey1234567890123456789012345678901234567",
}
messageBus := bus.NewMessageBus()
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
expectedPath := "/webhook/wecom-aibot"
if ch.WebhookPath() != expectedPath {
t.Errorf("Expected webhook path '%s', got '%s'", expectedPath, ch.WebhookPath())
}
})
t.Run("custom path", func(t *testing.T) {
customPath := "/custom/webhook"
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
EncodingAESKey: "testkey1234567890123456789012345678901234567",
WebhookPath: customPath,
}
messageBus := bus.NewMessageBus()
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
if ch.WebhookPath() != customPath {
t.Errorf("Expected webhook path '%s', got '%s'", customPath, ch.WebhookPath())
}
})
}
func TestGenerateStreamID(t *testing.T) {
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
EncodingAESKey: "testkey1234567890123456789012345678901234567",
}
messageBus := bus.NewMessageBus()
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
// Generate multiple IDs and check they are unique
ids := make(map[string]bool)
for i := 0; i < 100; i++ {
id := ch.generateStreamID()
if len(id) != 10 {
t.Errorf("Expected stream ID length 10, got %d", len(id))
}
if ids[id] {
t.Errorf("Duplicate stream ID generated: %s", id)
}
ids[id] = true
}
}
func TestEncryptDecrypt(t *testing.T) {
// Use a valid 43-character base64 key (企业微信标准格式)
cfg := config.WeComAIBotConfig{
Enabled: true,
Token: "test_token",
EncodingAESKey: "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG", // 43 characters
}
messageBus := bus.NewMessageBus()
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
plaintext := "Hello, World!"
receiveid := ""
// Encrypt
encrypted, err := ch.encryptMessage(plaintext, receiveid)
if err != nil {
t.Fatalf("Failed to encrypt message: %v", err)
}
if encrypted == "" {
t.Fatal("Encrypted message is empty")
}
// Decrypt
decrypted, err := decryptMessageWithVerify(encrypted, cfg.EncodingAESKey, receiveid)
if err != nil {
t.Fatalf("Failed to decrypt message: %v", err)
}
if decrypted != plaintext {
t.Errorf("Expected decrypted message '%s', got '%s'", plaintext, decrypted)
}
}
func TestGenerateSignature(t *testing.T) {
token := "test_token"
timestamp := "1234567890"
nonce := "test_nonce"
encrypt := "encrypted_msg"
signature := computeSignature(token, timestamp, nonce, encrypt)
if signature == "" {
t.Error("Generated signature is empty")
}
// Verify signature using verifySignature function
if !verifySignature(token, signature, timestamp, nonce, encrypt) {
t.Error("Generated signature does not verify correctly")
}
}

View file

@ -21,6 +21,7 @@ import (
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/identity" "github.com/sipeed/picoclaw/pkg/identity"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/utils" "github.com/sipeed/picoclaw/pkg/utils"
) )
@ -32,13 +33,13 @@ const (
type WeComAppChannel struct { type WeComAppChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.WeComAppConfig config config.WeComAppConfig
client *http.Client
accessToken string accessToken string
tokenExpiry time.Time tokenExpiry time.Time
tokenMu sync.RWMutex tokenMu sync.RWMutex
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
processedMsgs map[string]bool // Message deduplication: msg_id -> processed processedMsgs *MessageDeduplicator
msgMu sync.RWMutex
} }
// WeComXMLMessage represents the XML message structure from WeCom // WeComXMLMessage represents the XML message structure from WeCom
@ -129,10 +130,21 @@ func NewWeComAppChannel(cfg config.WeComAppConfig, messageBus *bus.MessageBus) (
channels.WithReasoningChannelID(cfg.ReasoningChannelID), channels.WithReasoningChannelID(cfg.ReasoningChannelID),
) )
// Client timeout must be >= the configured ReplyTimeout so the
// per-request context deadline is always the effective limit.
clientTimeout := 30 * time.Second
if d := time.Duration(cfg.ReplyTimeout) * time.Second; d > clientTimeout {
clientTimeout = d
}
ctx, cancel := context.WithCancel(context.Background())
return &WeComAppChannel{ return &WeComAppChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
processedMsgs: make(map[string]bool), client: &http.Client{Timeout: clientTimeout},
ctx: ctx,
cancel: cancel,
processedMsgs: NewMessageDeduplicator(wecomMaxProcessedMessages),
}, nil }, nil
} }
@ -145,6 +157,10 @@ func (c *WeComAppChannel) Name() string {
func (c *WeComAppChannel) Start(ctx context.Context) error { func (c *WeComAppChannel) Start(ctx context.Context) error {
logger.InfoC("wecom_app", "Starting WeCom App channel...") logger.InfoC("wecom_app", "Starting WeCom App channel...")
// Cancel the context created in the constructor to avoid a resource leak.
if c.cancel != nil {
c.cancel()
}
c.ctx, c.cancel = context.WithCancel(ctx) c.ctx, c.cancel = context.WithCancel(ctx)
// Get initial access token // Get initial access token
@ -249,10 +265,16 @@ func (c *WeComAppChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
} }
// Send media message using the media_id // Send media message using the media_id
if mediaType == "image" { switch mediaType {
case "image":
err = c.sendImageMessage(ctx, accessToken, msg.ChatID, mediaID) err = c.sendImageMessage(ctx, accessToken, msg.ChatID, mediaID)
} else { case "video":
// For non-image types, send as text fallback with caption err = c.sendVideoMessage(ctx, accessToken, msg.ChatID, mediaID)
case "voice":
err = c.sendVoiceMessage(ctx, accessToken, msg.ChatID, mediaID)
case "file":
err = c.sendFileMessage(ctx, accessToken, msg.ChatID, mediaID, part.Filename)
default:
caption := part.Caption caption := part.Caption
if caption == "" { if caption == "" {
caption = fmt.Sprintf("[%s: %s]", part.Type, part.Filename) caption = fmt.Sprintf("[%s: %s]", part.Type, part.Filename)
@ -299,8 +321,7 @@ func (c *WeComAppChannel) uploadMedia(ctx context.Context, accessToken, mediaTyp
} }
req.Header.Set("Content-Type", writer.FormDataContentType()) req.Header.Set("Content-Type", writer.FormDataContentType())
client := &http.Client{Timeout: 30 * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return "", channels.ClassifyNetError(err) return "", channels.ClassifyNetError(err)
} }
@ -327,18 +348,11 @@ func (c *WeComAppChannel) uploadMedia(ctx context.Context, accessToken, mediaTyp
return result.MediaID, nil return result.MediaID, nil
} }
// sendImageMessage sends an image message using a media_id. // sendWeComMessage marshals payload and POSTs it to the WeCom message API.
func (c *WeComAppChannel) sendImageMessage(ctx context.Context, accessToken, userID, mediaID string) error { func (c *WeComAppChannel) sendWeComMessage(ctx context.Context, accessToken string, payload any) error {
apiURL := fmt.Sprintf("%s/cgi-bin/message/send?access_token=%s", wecomAPIBase, accessToken) apiURL := fmt.Sprintf("%s/cgi-bin/message/send?access_token=%s", wecomAPIBase, accessToken)
msg := WeComImageMessage{ jsonData, err := json.Marshal(payload)
ToUser: userID,
MsgType: "image",
AgentID: c.config.AgentID,
}
msg.Image.MediaID = mediaID
jsonData, err := json.Marshal(msg)
if err != nil { if err != nil {
return fmt.Errorf("failed to marshal message: %w", err) return fmt.Errorf("failed to marshal message: %w", err)
} }
@ -357,8 +371,7 @@ func (c *WeComAppChannel) sendImageMessage(ctx context.Context, accessToken, use
} }
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(timeout) * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }
@ -386,6 +399,82 @@ func (c *WeComAppChannel) sendImageMessage(ctx context.Context, accessToken, use
return nil return nil
} }
// sendImageMessage sends an image message using a media_id.
func (c *WeComAppChannel) sendImageMessage(ctx context.Context, accessToken, userID, mediaID string) error {
msg := WeComImageMessage{
ToUser: userID,
MsgType: "image",
AgentID: c.config.AgentID,
}
msg.Image.MediaID = mediaID
return c.sendWeComMessage(ctx, accessToken, msg)
}
// WeComVideoMessage represents video message for sending
type WeComVideoMessage struct {
ToUser string `json:"touser"`
MsgType string `json:"msgtype"`
AgentID int64 `json:"agentid"`
Video struct {
MediaID string `json:"media_id"`
Title string `json:"title"`
Desc string `json:"description"`
} `json:"video"`
}
// sendVideoMessage sends a video message using a media_id.
func (c *WeComAppChannel) sendVideoMessage(ctx context.Context, accessToken, userID, mediaID string) error {
msg := WeComVideoMessage{
ToUser: userID,
MsgType: "video",
AgentID: c.config.AgentID,
}
msg.Video.MediaID = mediaID
return c.sendWeComMessage(ctx, accessToken, msg)
}
// WeComVoiceMessage represents voice message for sending
type WeComVoiceMessage struct {
ToUser string `json:"touser"`
MsgType string `json:"msgtype"`
AgentID int64 `json:"agentid"`
Voice struct {
MediaID string `json:"media_id"`
} `json:"voice"`
}
// sendVoiceMessage sends a voice message using a media_id.
func (c *WeComAppChannel) sendVoiceMessage(ctx context.Context, accessToken, userID, mediaID string) error {
msg := WeComVoiceMessage{
ToUser: userID,
MsgType: "voice",
AgentID: c.config.AgentID,
}
msg.Voice.MediaID = mediaID
return c.sendWeComMessage(ctx, accessToken, msg)
}
// WeComFileMessage represents file message for sending
type WeComFileMessage struct {
ToUser string `json:"touser"`
MsgType string `json:"msgtype"`
AgentID int64 `json:"agentid"`
File struct {
MediaID string `json:"media_id"`
} `json:"file"`
}
// sendFileMessage sends a file message using a media_id.
func (c *WeComAppChannel) sendFileMessage(ctx context.Context, accessToken, userID, mediaID, filename string) error {
msg := WeComFileMessage{
ToUser: userID,
MsgType: "file",
AgentID: c.config.AgentID,
}
msg.File.MediaID = mediaID
return c.sendWeComMessage(ctx, accessToken, msg)
}
// WebhookPath returns the path for registering on the shared HTTP server. // WebhookPath returns the path for registering on the shared HTTP server.
func (c *WeComAppChannel) WebhookPath() string { func (c *WeComAppChannel) WebhookPath() string {
if c.config.WebhookPath != "" { if c.config.WebhookPath != "" {
@ -567,8 +656,9 @@ func (c *WeComAppChannel) handleMessageCallback(ctx context.Context, w http.Resp
return return
} }
// Process the message with context // Process the message with the channel's long-lived context (not the HTTP
go c.processMessage(ctx, msg) // request context, which is canceled as soon as we return the response).
go c.processMessage(c.ctx, msg)
// Return success response immediately // Return success response immediately
// WeCom App requires response within configured timeout (default 5 seconds) // WeCom App requires response within configured timeout (default 5 seconds)
@ -577,8 +667,8 @@ func (c *WeComAppChannel) handleMessageCallback(ctx context.Context, w http.Resp
// processMessage processes the received message // processMessage processes the received message
func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessage) { func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessage) {
// Skip non-text messages for now (can be extended) // Handle different message types
if msg.MsgType != "text" && msg.MsgType != "image" && msg.MsgType != "voice" { if msg.MsgType != "text" && msg.MsgType != "image" && msg.MsgType != "voice" && msg.MsgType != "video" && msg.MsgType != "file" {
logger.DebugCF("wecom_app", "Skipping non-supported message type", map[string]any{ logger.DebugCF("wecom_app", "Skipping non-supported message type", map[string]any{
"msg_type": msg.MsgType, "msg_type": msg.MsgType,
}) })
@ -588,23 +678,12 @@ func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessag
// Message deduplication: Use msg_id to prevent duplicate processing // Message deduplication: Use msg_id to prevent duplicate processing
// As per WeCom documentation, use msg_id for deduplication // As per WeCom documentation, use msg_id for deduplication
msgID := fmt.Sprintf("%d", msg.MsgId) msgID := fmt.Sprintf("%d", msg.MsgId)
c.msgMu.Lock() if !c.processedMsgs.MarkMessageProcessed(msgID) {
if c.processedMsgs[msgID] {
c.msgMu.Unlock()
logger.DebugCF("wecom_app", "Skipping duplicate message", map[string]any{ logger.DebugCF("wecom_app", "Skipping duplicate message", map[string]any{
"msg_id": msgID, "msg_id": msgID,
}) })
return return
} }
c.processedMsgs[msgID] = true
c.msgMu.Unlock()
// Clean up old messages periodically (keep last 1000)
if len(c.processedMsgs) > 1000 {
c.msgMu.Lock()
c.processedMsgs = make(map[string]bool)
c.msgMu.Unlock()
}
senderID := msg.FromUserName senderID := msg.FromUserName
chatID := senderID // WeCom App uses user ID as chat ID for direct messages chatID := senderID // WeCom App uses user ID as chat ID for direct messages
@ -625,10 +704,28 @@ func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessag
content := msg.Content content := msg.Content
// Handle media messages (download and store)
var mediaRefs []string
store := c.GetMediaStore()
if store != nil && msg.MediaId != "" {
mediaRefs = c.downloadInboundMedia(ctx, msg.MsgType, msg.MediaId, messageID, store)
}
// Append media tags to content
if len(mediaRefs) > 0 {
mediaTag := c.getMediaTag(msg.MsgType)
if content != "" {
content += "\n" + mediaTag
} else {
content = mediaTag
}
}
logger.DebugCF("wecom_app", "Received message", map[string]any{ logger.DebugCF("wecom_app", "Received message", map[string]any{
"sender_id": senderID, "sender_id": senderID,
"msg_type": msg.MsgType, "msg_type": msg.MsgType,
"preview": utils.Truncate(content, 50), "preview": utils.Truncate(content, 50),
"mediaRefs": len(mediaRefs),
}) })
// Build sender info // Build sender info
@ -639,7 +736,7 @@ func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessag
} }
// Handle the message through the base channel // Handle the message through the base channel
c.HandleMessage(ctx, peer, messageID, senderID, chatID, content, nil, metadata, appSender) c.HandleMessage(ctx, peer, messageID, senderID, chatID, content, mediaRefs, metadata, appSender)
} }
// tokenRefreshLoop periodically refreshes the access token // tokenRefreshLoop periodically refreshes the access token
@ -707,64 +804,15 @@ func (c *WeComAppChannel) getAccessToken() string {
return c.accessToken return c.accessToken
} }
// sendTextMessage sends a text message to a user // sendTextMessage sends a text message to a user.
func (c *WeComAppChannel) sendTextMessage(ctx context.Context, accessToken, userID, content string) error { func (c *WeComAppChannel) sendTextMessage(ctx context.Context, accessToken, userID, content string) error {
apiURL := fmt.Sprintf("%s/cgi-bin/message/send?access_token=%s", wecomAPIBase, accessToken)
msg := WeComTextMessage{ msg := WeComTextMessage{
ToUser: userID, ToUser: userID,
MsgType: "text", MsgType: "text",
AgentID: c.config.AgentID, AgentID: c.config.AgentID,
} }
msg.Text.Content = content msg.Text.Content = content
return c.sendWeComMessage(ctx, accessToken, msg)
jsonData, err := json.Marshal(msg)
if err != nil {
return fmt.Errorf("failed to marshal message: %w", err)
}
// Use configurable timeout (default 5 seconds)
timeout := c.config.ReplyTimeout
if timeout <= 0 {
timeout = 5
}
reqCtx, cancel := context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
defer cancel()
req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
if err != nil {
return fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(timeout) * time.Second}
resp, err := client.Do(req)
if err != nil {
return channels.ClassifyNetError(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
return channels.ClassifySendError(resp.StatusCode, fmt.Errorf("wecom_app API error: %s", string(body)))
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return fmt.Errorf("failed to read response: %w", err)
}
var sendResp WeComSendMessageResponse
if err := json.Unmarshal(body, &sendResp); err != nil {
return fmt.Errorf("failed to parse response: %w", err)
}
if sendResp.ErrCode != 0 {
return fmt.Errorf("API error: %s (code: %d)", sendResp.ErrMsg, sendResp.ErrCode)
}
return nil
} }
// handleHealth handles health check requests // handleHealth handles health check requests
@ -778,3 +826,187 @@ func (c *WeComAppChannel) handleHealth(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json") w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(status) json.NewEncoder(w).Encode(status)
} }
// downloadInboundMedia downloads media from inbound messages and stores in MediaStore.
func (c *WeComAppChannel) downloadInboundMedia(
ctx context.Context,
msgType, mediaID, messageID string,
store media.MediaStore,
) []string {
var refs []string
scope := channels.BuildMediaScope("wecom_app", messageID, mediaID)
// Determine file extension based on message type
var ext string
switch msgType {
case "image":
ext = ".jpg"
case "voice":
ext = ".amr"
case "video":
ext = ".mp4"
case "file":
ext = ""
default:
ext = ""
}
// Download the media file from WeCom server
localPath := c.downloadMedia(ctx, mediaID, ext)
if localPath == "" {
logger.ErrorCF("wecom_app", "Failed to download media", map[string]any{
"media_id": mediaID,
"msg_type": msgType,
})
return refs
}
// Determine filename
filename := mediaID + ext
if msgType == "file" {
filename = "file"
}
// Store in media store
ref, err := store.Store(localPath, media.MediaMeta{
Filename: filename,
Source: "wecom_app",
}, scope)
if err != nil {
logger.ErrorCF("wecom_app", "Failed to store media", map[string]any{
"error": err.Error(),
"path": localPath,
})
return refs
}
refs = append(refs, ref)
return refs
}
// downloadMedia downloads media from WeCom API and saves to local file.
func (c *WeComAppChannel) downloadMedia(ctx context.Context, mediaID, ext string) string {
accessToken := c.getAccessToken()
if accessToken == "" {
logger.ErrorCF("wecom_app", "No access token available for media download", nil)
return ""
}
// WeCom API: GET /cgi-bin/media/get?access_token=ACCESS_TOKEN&media_id=MEDIA_ID
apiURL := fmt.Sprintf("%s/cgi-bin/media/get?access_token=%s&media_id=%s",
wecomAPIBase, url.QueryEscape(accessToken), url.QueryEscape(mediaID))
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
if err != nil {
logger.ErrorCF("wecom_app", "Failed to create media download request", map[string]any{
"error": err.Error(),
})
return ""
}
resp, err := c.client.Do(req)
if err != nil {
logger.ErrorCF("wecom_app", "Failed to download media", map[string]any{
"error": err.Error(),
"media_id": mediaID,
})
return ""
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
logger.ErrorCF("wecom_app", "Media download failed with status", map[string]any{
"status": resp.StatusCode,
"media_id": mediaID,
})
return ""
}
// Check content-type to determine file type
contentType := resp.Header.Get("Content-Type")
// Determine file extension from content-type if not provided
if ext == "" {
switch {
case strings.Contains(contentType, "image/jpeg"):
ext = ".jpg"
case strings.Contains(contentType, "image/png"):
ext = ".png"
case strings.Contains(contentType, "image/gif"):
ext = ".gif"
case strings.Contains(contentType, "audio/amr"):
ext = ".amr"
case strings.Contains(contentType, "audio/mp3") || strings.Contains(contentType, "audio/mpeg"):
ext = ".mp3"
case strings.Contains(contentType, "video/mp4"):
ext = ".mp4"
default:
ext = ".bin"
}
}
// Generate temp file path
tempDir := os.TempDir()
mediaDir := filepath.Join(tempDir, "picoclaw_media", "wecom")
os.MkdirAll(mediaDir, 0o755)
localPath := filepath.Join(mediaDir, mediaID+ext)
// Write to file
body, err := io.ReadAll(resp.Body)
if err != nil {
logger.ErrorCF("wecom_app", "Failed to read media body", map[string]any{
"error": err.Error(),
})
return ""
}
// Check if response is error JSON
if len(body) > 0 && body[0] == '{' {
var errResp struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
}
if json.Unmarshal(body, &errResp) == nil && errResp.ErrCode != 0 {
logger.ErrorCF("wecom_app", "Media download API error", map[string]any{
"errcode": errResp.ErrCode,
"errmsg": errResp.ErrMsg,
"media_id": mediaID,
})
return ""
}
}
err = os.WriteFile(localPath, body, 0o644)
if err != nil {
logger.ErrorCF("wecom_app", "Failed to write media file", map[string]any{
"error": err.Error(),
"path": localPath,
})
return ""
}
logger.DebugCF("wecom_app", "Media downloaded successfully", map[string]any{
"media_id": mediaID,
"path": localPath,
"size": len(body),
})
return localPath
}
// getMediaTag returns a media tag string for the message type.
func (c *WeComAppChannel) getMediaTag(msgType string) string {
switch msgType {
case "image":
return "[image]"
case "voice":
return "[voice message]"
case "video":
return "[video]"
case "file":
return "[file]"
default:
return "[media]"
}
}

View file

@ -43,7 +43,7 @@ func encryptTestMessageApp(message, aesKey string) (string, error) {
// Prepare message: random(16) + msg_len(4) + msg + corp_id // Prepare message: random(16) + msg_len(4) + msg + corp_id
random := make([]byte, 0, 16) random := make([]byte, 0, 16)
for i := 0; i < 16; i++ { for i := range 16 {
random = append(random, byte(i+1)) random = append(random, byte(i+1))
} }
@ -323,60 +323,6 @@ func TestWeComAppDecryptMessage(t *testing.T) {
}) })
} }
func TestWeComAppPKCS7Unpad(t *testing.T) {
tests := []struct {
name string
input []byte
expected []byte
}{
{
name: "empty input",
input: []byte{},
expected: []byte{},
},
{
name: "valid padding 3 bytes",
input: append([]byte("hello"), bytes.Repeat([]byte{3}, 3)...),
expected: []byte("hello"),
},
{
name: "valid padding 16 bytes (full block)",
input: append([]byte("123456789012345"), bytes.Repeat([]byte{16}, 16)...),
expected: []byte("123456789012345"),
},
{
name: "invalid padding larger than data",
input: []byte{20},
expected: nil, // should return error
},
{
name: "invalid padding zero",
input: append([]byte("test"), byte(0)),
expected: nil, // should return error
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result, err := pkcs7Unpad(tt.input)
if tt.expected == nil {
// This case should return an error
if err == nil {
t.Errorf("pkcs7Unpad() expected error for invalid padding, got result: %v", result)
}
return
}
if err != nil {
t.Errorf("pkcs7Unpad() unexpected error: %v", err)
return
}
if !bytes.Equal(result, tt.expected) {
t.Errorf("pkcs7Unpad() = %v, want %v", result, tt.expected)
}
})
}
}
func TestWeComAppHandleVerification(t *testing.T) { func TestWeComAppHandleVerification(t *testing.T) {
msgBus := bus.NewMessageBus() msgBus := bus.NewMessageBus()
aesKey := generateTestAESKeyApp() aesKey := generateTestAESKeyApp()

View file

@ -9,7 +9,6 @@ import (
"io" "io"
"net/http" "net/http"
"strings" "strings"
"sync"
"time" "time"
"github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/bus"
@ -25,10 +24,10 @@ import (
type WeComBotChannel struct { type WeComBotChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.WeComConfig config config.WeComConfig
client *http.Client
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
processedMsgs map[string]bool // Message deduplication: msg_id -> processed processedMsgs *MessageDeduplicator
msgMu sync.RWMutex
} }
// WeComBotMessage represents the JSON message structure from WeCom Bot (AIBOT) // WeComBotMessage represents the JSON message structure from WeCom Bot (AIBOT)
@ -93,10 +92,21 @@ func NewWeComBotChannel(cfg config.WeComConfig, messageBus *bus.MessageBus) (*We
channels.WithReasoningChannelID(cfg.ReasoningChannelID), channels.WithReasoningChannelID(cfg.ReasoningChannelID),
) )
// Client timeout must be >= the configured ReplyTimeout so the
// per-request context deadline is always the effective limit.
clientTimeout := 30 * time.Second
if d := time.Duration(cfg.ReplyTimeout) * time.Second; d > clientTimeout {
clientTimeout = d
}
ctx, cancel := context.WithCancel(context.Background())
return &WeComBotChannel{ return &WeComBotChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
processedMsgs: make(map[string]bool), client: &http.Client{Timeout: clientTimeout},
ctx: ctx,
cancel: cancel,
processedMsgs: NewMessageDeduplicator(wecomMaxProcessedMessages),
}, nil }, nil
} }
@ -109,6 +119,10 @@ func (c *WeComBotChannel) Name() string {
func (c *WeComBotChannel) Start(ctx context.Context) error { func (c *WeComBotChannel) Start(ctx context.Context) error {
logger.InfoC("wecom", "Starting WeCom Bot channel...") logger.InfoC("wecom", "Starting WeCom Bot channel...")
// Cancel the context created in the constructor to avoid a resource leak.
if c.cancel != nil {
c.cancel()
}
c.ctx, c.cancel = context.WithCancel(ctx) c.ctx, c.cancel = context.WithCancel(ctx)
c.SetRunning(true) c.SetRunning(true)
@ -292,8 +306,9 @@ func (c *WeComBotChannel) handleMessageCallback(ctx context.Context, w http.Resp
return return
} }
// Process the message asynchronously with context // Process the message with the channel's long-lived context (not the HTTP
go c.processMessage(ctx, msg) // request context, which is canceled as soon as we return the response).
go c.processMessage(c.ctx, msg)
// Return success response immediately // Return success response immediately
// WeCom Bot requires response within configured timeout (default 5 seconds) // WeCom Bot requires response within configured timeout (default 5 seconds)
@ -313,23 +328,12 @@ func (c *WeComBotChannel) processMessage(ctx context.Context, msg WeComBotMessag
// Message deduplication: Use msg_id to prevent duplicate processing // Message deduplication: Use msg_id to prevent duplicate processing
msgID := msg.MsgID msgID := msg.MsgID
c.msgMu.Lock() if !c.processedMsgs.MarkMessageProcessed(msgID) {
if c.processedMsgs[msgID] {
c.msgMu.Unlock()
logger.DebugCF("wecom", "Skipping duplicate message", map[string]any{ logger.DebugCF("wecom", "Skipping duplicate message", map[string]any{
"msg_id": msgID, "msg_id": msgID,
}) })
return return
} }
c.processedMsgs[msgID] = true
c.msgMu.Unlock()
// Clean up old messages periodically (keep last 1000)
if len(c.processedMsgs) > 1000 {
c.msgMu.Lock()
c.processedMsgs = make(map[string]bool)
c.msgMu.Unlock()
}
senderID := msg.From.UserID senderID := msg.From.UserID
@ -442,8 +446,7 @@ func (c *WeComBotChannel) sendWebhookReply(ctx context.Context, userID, content
} }
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(timeout) * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }

View file

@ -42,7 +42,7 @@ func encryptTestMessage(message, aesKey string) (string, error) {
// Prepare message: random(16) + msg_len(4) + msg + receiveid // Prepare message: random(16) + msg_len(4) + msg + receiveid
random := make([]byte, 0, 16) random := make([]byte, 0, 16)
for i := 0; i < 16; i++ { for i := range 16 {
random = append(random, byte(i)) random = append(random, byte(i))
} }
@ -412,22 +412,9 @@ func TestWeComBotHandleMessageCallback(t *testing.T) {
} }
ch, _ := NewWeComBotChannel(cfg, msgBus) ch, _ := NewWeComBotChannel(cfg, msgBus)
t.Run("valid direct message callback", func(t *testing.T) { runBotMessageCallback := func(t *testing.T, jsonMsg string) *httptest.ResponseRecorder {
// Create JSON message for direct chat (single) t.Helper()
jsonMsg := `{
"msgid": "test_msg_id_123",
"aibotid": "test_aibot_id",
"chattype": "single",
"from": {"userid": "user123"},
"response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
"msgtype": "text",
"text": {"content": "Hello World"}
}`
// Encrypt message
encrypted, _ := encryptTestMessage(jsonMsg, aesKey) encrypted, _ := encryptTestMessage(jsonMsg, aesKey)
// Create encrypted XML wrapper
encryptedWrapper := struct { encryptedWrapper := struct {
XMLName xml.Name `xml:"xml"` XMLName xml.Name `xml:"xml"`
Encrypt string `xml:"Encrypt"` Encrypt string `xml:"Encrypt"`
@ -435,20 +422,29 @@ func TestWeComBotHandleMessageCallback(t *testing.T) {
Encrypt: encrypted, Encrypt: encrypted,
} }
wrapperData, _ := xml.Marshal(encryptedWrapper) wrapperData, _ := xml.Marshal(encryptedWrapper)
timestamp := "1234567890" timestamp := "1234567890"
nonce := "test_nonce" nonce := "test_nonce"
signature := generateSignature("test_token", timestamp, nonce, encrypted) signature := generateSignature("test_token", timestamp, nonce, encrypted)
req := httptest.NewRequest( req := httptest.NewRequest(
http.MethodPost, http.MethodPost,
"/webhook/wecom?msg_signature="+signature+"&timestamp="+timestamp+"&nonce="+nonce, "/webhook/wecom?msg_signature="+signature+"&timestamp="+timestamp+"&nonce="+nonce,
bytes.NewReader(wrapperData), bytes.NewReader(wrapperData),
) )
w := httptest.NewRecorder() w := httptest.NewRecorder()
ch.handleMessageCallback(context.Background(), w, req) ch.handleMessageCallback(context.Background(), w, req)
return w
}
t.Run("valid direct message callback", func(t *testing.T) {
w := runBotMessageCallback(t, `{
"msgid": "test_msg_id_123",
"aibotid": "test_aibot_id",
"chattype": "single",
"from": {"userid": "user123"},
"response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
"msgtype": "text",
"text": {"content": "Hello World"}
}`)
if w.Code != http.StatusOK { if w.Code != http.StatusOK {
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK) t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
} }
@ -458,8 +454,7 @@ func TestWeComBotHandleMessageCallback(t *testing.T) {
}) })
t.Run("valid group message callback", func(t *testing.T) { t.Run("valid group message callback", func(t *testing.T) {
// Create JSON message for group chat w := runBotMessageCallback(t, `{
jsonMsg := `{
"msgid": "test_msg_id_456", "msgid": "test_msg_id_456",
"aibotid": "test_aibot_id", "aibotid": "test_aibot_id",
"chatid": "group_chat_id_123", "chatid": "group_chat_id_123",
@ -468,33 +463,7 @@ func TestWeComBotHandleMessageCallback(t *testing.T) {
"response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test", "response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
"msgtype": "text", "msgtype": "text",
"text": {"content": "Hello Group"} "text": {"content": "Hello Group"}
}` }`)
// Encrypt message
encrypted, _ := encryptTestMessage(jsonMsg, aesKey)
// Create encrypted XML wrapper
encryptedWrapper := struct {
XMLName xml.Name `xml:"xml"`
Encrypt string `xml:"Encrypt"`
}{
Encrypt: encrypted,
}
wrapperData, _ := xml.Marshal(encryptedWrapper)
timestamp := "1234567890"
nonce := "test_nonce"
signature := generateSignature("test_token", timestamp, nonce, encrypted)
req := httptest.NewRequest(
http.MethodPost,
"/webhook/wecom?msg_signature="+signature+"&timestamp="+timestamp+"&nonce="+nonce,
bytes.NewReader(wrapperData),
)
w := httptest.NewRecorder()
ch.handleMessageCallback(context.Background(), w, req)
if w.Code != http.StatusOK { if w.Code != http.StatusOK {
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK) t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
} }

View file

@ -1,12 +1,15 @@
package wecom package wecom
import ( import (
"bytes"
"crypto/aes" "crypto/aes"
"crypto/cipher" "crypto/cipher"
"crypto/rand"
"crypto/sha1" "crypto/sha1"
"encoding/base64" "encoding/base64"
"encoding/binary" "encoding/binary"
"fmt" "fmt"
"math/big"
"sort" "sort"
"strings" "strings"
) )
@ -14,25 +17,23 @@ import (
// blockSize is the PKCS7 block size used by WeCom (32) // blockSize is the PKCS7 block size used by WeCom (32)
const blockSize = 32 const blockSize = 32
// computeSignature computes the WeCom message signature from the given parameters.
// It sorts [token, timestamp, nonce, encrypt], concatenates them and returns the SHA1 hex digest.
func computeSignature(token, timestamp, nonce, encrypt string) string {
params := []string{token, timestamp, nonce, encrypt}
sort.Strings(params)
str := strings.Join(params, "")
hash := sha1.Sum([]byte(str))
return fmt.Sprintf("%x", hash)
}
// verifySignature verifies the message signature for WeCom // verifySignature verifies the message signature for WeCom
// This is a common function used by both WeCom Bot and WeCom App // This is a common function used by both WeCom Bot and WeCom App
func verifySignature(token, msgSignature, timestamp, nonce, msgEncrypt string) bool { func verifySignature(token, msgSignature, timestamp, nonce, msgEncrypt string) bool {
if token == "" { if token == "" {
return true // Skip verification if token is not set return true // Skip verification if token is not set
} }
return computeSignature(token, timestamp, nonce, msgEncrypt) == msgSignature
// Sort parameters
params := []string{token, timestamp, nonce, msgEncrypt}
sort.Strings(params)
// Concatenate
str := strings.Join(params, "")
// SHA1 hash
hash := sha1.Sum([]byte(str))
expectedSignature := fmt.Sprintf("%x", hash)
return expectedSignature == msgSignature
} }
// decryptMessage decrypts the encrypted message using AES // decryptMessage decrypts the encrypted message using AES
@ -53,64 +54,128 @@ func decryptMessageWithVerify(encryptedMsg, encodingAESKey, receiveid string) (s
return string(decoded), nil return string(decoded), nil
} }
// Decode AES key (base64) aesKey, err := decodeWeComAESKey(encodingAESKey)
aesKey, err := base64.StdEncoding.DecodeString(encodingAESKey + "=")
if err != nil { if err != nil {
return "", fmt.Errorf("failed to decode AES key: %w", err) return "", err
} }
// Decode encrypted message
cipherText, err := base64.StdEncoding.DecodeString(encryptedMsg) cipherText, err := base64.StdEncoding.DecodeString(encryptedMsg)
if err != nil { if err != nil {
return "", fmt.Errorf("failed to decode message: %w", err) return "", fmt.Errorf("failed to decode message: %w", err)
} }
// AES decrypt plainText, err := decryptAESCBC(aesKey, cipherText)
if err != nil {
return "", err
}
return unpackWeComFrame(plainText, receiveid)
}
// decodeWeComAESKey base64-decodes the 43-character EncodingAESKey (trailing "=" is
// appended automatically) and validates that the result is exactly 32 bytes.
// It is the single place that handles this repeated pattern in both encrypt and decrypt paths.
func decodeWeComAESKey(encodingAESKey string) ([]byte, error) {
aesKey, err := base64.StdEncoding.DecodeString(encodingAESKey + "=")
if err != nil {
return nil, fmt.Errorf("failed to decode AES key: %w", err)
}
if len(aesKey) != 32 {
return nil, fmt.Errorf("invalid AES key length: %d", len(aesKey))
}
return aesKey, nil
}
// encryptAESCBC encrypts plaintext using AES-CBC with the given key, mirroring
// decryptAESCBC. IV = aesKey[:aes.BlockSize]. The caller must PKCS7-pad the
// plaintext to a multiple of aes.BlockSize before calling.
func encryptAESCBC(aesKey, plaintext []byte) ([]byte, error) {
block, err := aes.NewCipher(aesKey) block, err := aes.NewCipher(aesKey)
if err != nil { if err != nil {
return "", fmt.Errorf("failed to create cipher: %w", err) return nil, fmt.Errorf("failed to create cipher: %w", err)
} }
if len(cipherText) < aes.BlockSize {
return "", fmt.Errorf("ciphertext too short")
}
// IV is the first 16 bytes of AESKey
iv := aesKey[:aes.BlockSize] iv := aesKey[:aes.BlockSize]
mode := cipher.NewCBCDecrypter(block, iv) ciphertext := make([]byte, len(plaintext))
plainText := make([]byte, len(cipherText)) cipher.NewCBCEncrypter(block, iv).CryptBlocks(ciphertext, plaintext)
mode.CryptBlocks(plainText, cipherText) return ciphertext, nil
}
// Remove PKCS7 padding // packWeComFrame builds the WeCom wire format:
plainText, err = pkcs7Unpad(plainText) //
if err != nil { // random(16 ASCII digits) + msg_len(4, big-endian) + msg + receiveid
return "", fmt.Errorf("failed to unpad: %w", err) func packWeComFrame(msg, receiveid string) ([]byte, error) {
randomBytes := make([]byte, 16)
for i := range 16 {
n, err := rand.Int(rand.Reader, big.NewInt(10))
if err != nil {
return nil, fmt.Errorf("failed to generate random: %w", err)
}
randomBytes[i] = byte('0' + n.Int64())
} }
msgBytes := []byte(msg)
msgLenBytes := make([]byte, 4)
binary.BigEndian.PutUint32(msgLenBytes, uint32(len(msgBytes)))
var buf bytes.Buffer
buf.Write(randomBytes)
buf.Write(msgLenBytes)
buf.Write(msgBytes)
buf.WriteString(receiveid)
return buf.Bytes(), nil
}
// Parse message structure // unpackWeComFrame parses the WeCom wire format produced by packWeComFrame.
// Format: random(16) + msg_len(4) + msg + receiveid // If receiveid is non-empty it verifies the frame's trailing receiveid field.
if len(plainText) < 20 { func unpackWeComFrame(data []byte, receiveid string) (string, error) {
return "", fmt.Errorf("decrypted message too short") if len(data) < 20 {
return "", fmt.Errorf("decrypted frame too short: %d bytes", len(data))
} }
msgLen := binary.BigEndian.Uint32(data[16:20])
msgLen := binary.BigEndian.Uint32(plainText[16:20]) if int(msgLen) > len(data)-20 {
if int(msgLen) > len(plainText)-20 { return "", fmt.Errorf("invalid message length: %d", msgLen)
return "", fmt.Errorf("invalid message length")
} }
msg := data[20 : 20+msgLen]
msg := plainText[20 : 20+msgLen] if receiveid != "" && len(data) > 20+int(msgLen) {
actualReceiveID := string(data[20+msgLen:])
// Verify receiveid if provided
if receiveid != "" && len(plainText) > 20+int(msgLen) {
actualReceiveID := string(plainText[20+msgLen:])
if actualReceiveID != receiveid { if actualReceiveID != receiveid {
return "", fmt.Errorf("receiveid mismatch: expected %s, got %s", receiveid, actualReceiveID) return "", fmt.Errorf("receiveid mismatch: expected %s, got %s", receiveid, actualReceiveID)
} }
} }
return string(msg), nil return string(msg), nil
} }
// decryptAESCBC decrypts ciphertext using AES-CBC with the given key.
// IV = aesKey[:aes.BlockSize]. PKCS7 padding is stripped from the returned plaintext.
func decryptAESCBC(aesKey, ciphertext []byte) ([]byte, error) {
if len(ciphertext) == 0 {
return nil, fmt.Errorf("ciphertext is empty")
}
if len(ciphertext)%aes.BlockSize != 0 {
return nil, fmt.Errorf("ciphertext length %d is not a multiple of block size", len(ciphertext))
}
block, err := aes.NewCipher(aesKey)
if err != nil {
return nil, fmt.Errorf("failed to create cipher: %w", err)
}
iv := aesKey[:aes.BlockSize]
plaintext := make([]byte, len(ciphertext))
cipher.NewCBCDecrypter(block, iv).CryptBlocks(plaintext, ciphertext)
plaintext, err = pkcs7Unpad(plaintext)
if err != nil {
return nil, fmt.Errorf("failed to unpad: %w", err)
}
return plaintext, nil
}
// pkcs7Pad adds PKCS7 padding
func pkcs7Pad(data []byte, blockSize int) []byte {
padding := blockSize - (len(data) % blockSize)
if padding == 0 {
padding = blockSize
}
padText := bytes.Repeat([]byte{byte(padding)}, padding)
return append(data, padText...)
}
// pkcs7Unpad removes PKCS7 padding with validation // pkcs7Unpad removes PKCS7 padding with validation
func pkcs7Unpad(data []byte) ([]byte, error) { func pkcs7Unpad(data []byte) ([]byte, error) {
if len(data) == 0 { if len(data) == 0 {
@ -125,7 +190,7 @@ func pkcs7Unpad(data []byte) ([]byte, error) {
return nil, fmt.Errorf("padding size larger than data") return nil, fmt.Errorf("padding size larger than data")
} }
// Verify all padding bytes // Verify all padding bytes
for i := 0; i < padding; i++ { for i := range padding {
if data[len(data)-1-i] != byte(padding) { if data[len(data)-1-i] != byte(padding) {
return nil, fmt.Errorf("invalid padding byte at position %d", i) return nil, fmt.Errorf("invalid padding byte at position %d", i)
} }

View file

@ -0,0 +1,54 @@
package wecom
import "sync"
const wecomMaxProcessedMessages = 1000
// MessageDeduplicator provides thread-safe message deduplication using a circular queue (ring buffer)
// combined with a hash map. This ensures fast O(1) lookups while naturally evicting the oldest
// messages without causing "amnesia cliffs" when the limit is reached.
type MessageDeduplicator struct {
mu sync.Mutex
msgs map[string]bool
ring []string
idx int
max int
}
// NewMessageDeduplicator creates a new deduplicator with the specified capacity.
func NewMessageDeduplicator(maxEntries int) *MessageDeduplicator {
if maxEntries <= 0 {
maxEntries = wecomMaxProcessedMessages
}
return &MessageDeduplicator{
msgs: make(map[string]bool, maxEntries),
ring: make([]string, maxEntries),
max: maxEntries,
}
}
// MarkMessageProcessed marks msgID as processed and returns false for duplicates.
func (d *MessageDeduplicator) MarkMessageProcessed(msgID string) bool {
d.mu.Lock()
defer d.mu.Unlock()
// 1. Check for duplicate
if d.msgs[msgID] {
return false
}
// 2. Evict the oldest message at our current ring position (if any)
oldestID := d.ring[d.idx]
if oldestID != "" {
delete(d.msgs, oldestID)
}
// 3. Store the new message
d.msgs[msgID] = true
d.ring[d.idx] = msgID
// 4. Advance the circle queue index
d.idx = (d.idx + 1) % d.max
return true
}

View file

@ -0,0 +1,83 @@
package wecom
import (
"sync"
"testing"
)
func TestMessageDeduplicator_DuplicateDetection(t *testing.T) {
d := NewMessageDeduplicator(wecomMaxProcessedMessages)
if ok := d.MarkMessageProcessed("msg-1"); !ok {
t.Fatalf("first message should be accepted")
}
if ok := d.MarkMessageProcessed("msg-1"); ok {
t.Fatalf("duplicate message should be rejected")
}
}
func TestMessageDeduplicator_ConcurrentSameMessage(t *testing.T) {
d := NewMessageDeduplicator(wecomMaxProcessedMessages)
const goroutines = 64
var wg sync.WaitGroup
wg.Add(goroutines)
results := make(chan bool, goroutines)
for i := 0; i < goroutines; i++ {
go func() {
defer wg.Done()
results <- d.MarkMessageProcessed("msg-concurrent")
}()
}
wg.Wait()
close(results)
successes := 0
for ok := range results {
if ok {
successes++
}
}
if successes != 1 {
t.Fatalf("expected exactly 1 successful mark, got %d", successes)
}
}
func TestMessageDeduplicator_CircularQueueEviction(t *testing.T) {
// Create a deduplicator with a very small capacity to test eviction easily.
capacity := 3
d := NewMessageDeduplicator(capacity)
// Fill the queue.
d.MarkMessageProcessed("msg-1")
d.MarkMessageProcessed("msg-2")
d.MarkMessageProcessed("msg-3")
// At this point, the queue is full. msg-1 is the oldest.
if len(d.msgs) != 3 {
t.Fatalf("expected map size to be 3, got %d", len(d.msgs))
}
// This should evict msg-1 and add msg-4.
if ok := d.MarkMessageProcessed("msg-4"); !ok {
t.Fatalf("msg-4 should be accepted")
}
if len(d.msgs) != 3 {
t.Fatalf("expected map size to remain at max capacity (3), got %d", len(d.msgs))
}
// msg-1 should now be forgotten (evicted).
if ok := d.MarkMessageProcessed("msg-1"); !ok {
t.Fatalf("msg-1 should be accepted again because it was evicted")
}
// msg-2 should have been evicted when we added msg-1 back.
if ok := d.MarkMessageProcessed("msg-2"); !ok {
t.Fatalf("msg-2 should be accepted again because it was evicted")
}
}

View file

@ -13,4 +13,7 @@ func init() {
channels.RegisterFactory("wecom_app", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) { channels.RegisterFactory("wecom_app", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
return NewWeComAppChannel(cfg.Channels.WeComApp, b) return NewWeComAppChannel(cfg.Channels.WeComApp, b)
}) })
channels.RegisterFactory("wecom_aibot", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
return NewWeComAIBotChannel(cfg.Channels.WeComAIBot, b)
})
} }

View file

@ -58,6 +58,8 @@ type Config struct {
Tools ToolsConfig `json:"tools"` Tools ToolsConfig `json:"tools"`
Heartbeat HeartbeatConfig `json:"heartbeat"` Heartbeat HeartbeatConfig `json:"heartbeat"`
Devices DevicesConfig `json:"devices"` Devices DevicesConfig `json:"devices"`
Web WebConfig `json:"web"`
Voice VoiceConfig `json:"voice"`
} }
// MarshalJSON implements custom JSON marshaling for Config // MarshalJSON implements custom JSON marshaling for Config
@ -168,17 +170,30 @@ type SessionConfig struct {
} }
type AgentDefaults struct { type AgentDefaults struct {
Workspace string `json:"workspace" env:"PICOCLAW_AGENTS_DEFAULTS_WORKSPACE"` Workspace string `json:"workspace" env:"PICOCLAW_AGENTS_DEFAULTS_WORKSPACE"`
RestrictToWorkspace bool `json:"restrict_to_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_RESTRICT_TO_WORKSPACE"` RestrictToWorkspace bool `json:"restrict_to_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_RESTRICT_TO_WORKSPACE"`
Provider string `json:"provider" env:"PICOCLAW_AGENTS_DEFAULTS_PROVIDER"` AllowReadOutsideWorkspace bool `json:"allow_read_outside_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_ALLOW_READ_OUTSIDE_WORKSPACE"`
ModelName string `json:"model_name,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MODEL_NAME"` Provider string `json:"provider" env:"PICOCLAW_AGENTS_DEFAULTS_PROVIDER"`
Model string `json:"model" env:"PICOCLAW_AGENTS_DEFAULTS_MODEL"` // Deprecated: use model_name instead ModelName string `json:"model_name,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MODEL_NAME"`
ModelFallbacks []string `json:"model_fallbacks,omitempty"` Model string `json:"model" env:"PICOCLAW_AGENTS_DEFAULTS_MODEL"` // Deprecated: use model_name instead
ImageModel string `json:"image_model,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_IMAGE_MODEL"` ModelFallbacks []string `json:"model_fallbacks,omitempty"`
ImageModelFallbacks []string `json:"image_model_fallbacks,omitempty"` ImageModel string `json:"image_model,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_IMAGE_MODEL"`
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"` ImageModelFallbacks []string `json:"image_model_fallbacks,omitempty"`
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"` MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"` Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
SummarizeMessageThreshold int `json:"summarize_message_threshold" env:"PICOCLAW_AGENTS_DEFAULTS_SUMMARIZE_MESSAGE_THRESHOLD"`
SummarizeTokenPercent int `json:"summarize_token_percent" env:"PICOCLAW_AGENTS_DEFAULTS_SUMMARIZE_TOKEN_PERCENT"`
MaxMediaSize int `json:"max_media_size,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_MEDIA_SIZE"`
}
const DefaultMaxMediaSize = 20 * 1024 * 1024 // 20 MB
func (d *AgentDefaults) GetMaxMediaSize() int {
if d.MaxMediaSize > 0 {
return d.MaxMediaSize
}
return DefaultMaxMediaSize
} }
// GetModelName returns the effective model name for the agent defaults. // GetModelName returns the effective model name for the agent defaults.
@ -191,19 +206,20 @@ func (d *AgentDefaults) GetModelName() string {
} }
type ChannelsConfig struct { type ChannelsConfig struct {
WhatsApp WhatsAppConfig `json:"whatsapp"` WhatsApp WhatsAppConfig `json:"whatsapp"`
Telegram TelegramConfig `json:"telegram"` Telegram TelegramConfig `json:"telegram"`
Feishu FeishuConfig `json:"feishu"` Feishu FeishuConfig `json:"feishu"`
Discord DiscordConfig `json:"discord"` Discord DiscordConfig `json:"discord"`
MaixCam MaixCamConfig `json:"maixcam"` MaixCam MaixCamConfig `json:"maixcam"`
QQ QQConfig `json:"qq"` QQ QQConfig `json:"qq"`
DingTalk DingTalkConfig `json:"dingtalk"` DingTalk DingTalkConfig `json:"dingtalk"`
Slack SlackConfig `json:"slack"` Slack SlackConfig `json:"slack"`
LINE LINEConfig `json:"line"` LINE LINEConfig `json:"line"`
OneBot OneBotConfig `json:"onebot"` OneBot OneBotConfig `json:"onebot"`
WeCom WeComConfig `json:"wecom"` WeCom WeComConfig `json:"wecom"`
WeComApp WeComAppConfig `json:"wecom_app"` WeComApp WeComAppConfig `json:"wecom_app"`
Pico PicoConfig `json:"pico"` WeComAIBot WeComAIBotConfig `json:"wecom_aibot"`
Pico PicoConfig `json:"pico"`
} }
// GroupTriggerConfig controls when the bot responds in group chats. // GroupTriggerConfig controls when the bot responds in group chats.
@ -235,6 +251,7 @@ type WhatsAppConfig struct {
type TelegramConfig struct { type TelegramConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_TELEGRAM_ENABLED"` Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_TELEGRAM_ENABLED"`
Token string `json:"token" env:"PICOCLAW_CHANNELS_TELEGRAM_TOKEN"` Token string `json:"token" env:"PICOCLAW_CHANNELS_TELEGRAM_TOKEN"`
BaseURL string `json:"base_url" env:"PICOCLAW_CHANNELS_TELEGRAM_BASE_URL"`
Proxy string `json:"proxy" env:"PICOCLAW_CHANNELS_TELEGRAM_PROXY"` Proxy string `json:"proxy" env:"PICOCLAW_CHANNELS_TELEGRAM_PROXY"`
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_TELEGRAM_ALLOW_FROM"` AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_TELEGRAM_ALLOW_FROM"`
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"` GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
@ -251,12 +268,14 @@ type FeishuConfig struct {
VerificationToken string `json:"verification_token" env:"PICOCLAW_CHANNELS_FEISHU_VERIFICATION_TOKEN"` VerificationToken string `json:"verification_token" env:"PICOCLAW_CHANNELS_FEISHU_VERIFICATION_TOKEN"`
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_FEISHU_ALLOW_FROM"` AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_FEISHU_ALLOW_FROM"`
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"` GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_FEISHU_REASONING_CHANNEL_ID"` ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_FEISHU_REASONING_CHANNEL_ID"`
} }
type DiscordConfig struct { type DiscordConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_DISCORD_ENABLED"` Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_DISCORD_ENABLED"`
Token string `json:"token" env:"PICOCLAW_CHANNELS_DISCORD_TOKEN"` Token string `json:"token" env:"PICOCLAW_CHANNELS_DISCORD_TOKEN"`
Proxy string `json:"proxy" env:"PICOCLAW_CHANNELS_DISCORD_PROXY"`
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_DISCORD_ALLOW_FROM"` AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_DISCORD_ALLOW_FROM"`
MentionOnly bool `json:"mention_only" env:"PICOCLAW_CHANNELS_DISCORD_MENTION_ONLY"` MentionOnly bool `json:"mention_only" env:"PICOCLAW_CHANNELS_DISCORD_MENTION_ONLY"`
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"` GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
@ -359,6 +378,18 @@ type WeComAppConfig struct {
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_APP_REASONING_CHANNEL_ID"` ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_APP_REASONING_CHANNEL_ID"`
} }
type WeComAIBotConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ENABLED"`
Token string `json:"token" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_TOKEN"`
EncodingAESKey string `json:"encoding_aes_key" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ENCODING_AES_KEY"`
WebhookPath string `json:"webhook_path" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_WEBHOOK_PATH"`
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ALLOW_FROM"`
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_REPLY_TIMEOUT"`
MaxSteps int `json:"max_steps" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_MAX_STEPS"` // Maximum streaming steps
WelcomeMessage string `json:"welcome_message" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_WELCOME_MESSAGE"` // Sent on enter_chat event; empty = no welcome
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_REASONING_CHANNEL_ID"`
}
type PicoConfig struct { type PicoConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_PICO_ENABLED"` Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_PICO_ENABLED"`
Token string `json:"token" env:"PICOCLAW_CHANNELS_PICO_TOKEN"` Token string `json:"token" env:"PICOCLAW_CHANNELS_PICO_TOKEN"`
@ -385,6 +416,7 @@ type DevicesConfig struct {
type ProvidersConfig struct { type ProvidersConfig struct {
Anthropic ProviderConfig `json:"anthropic"` Anthropic ProviderConfig `json:"anthropic"`
OpenAI OpenAIProviderConfig `json:"openai"` OpenAI OpenAIProviderConfig `json:"openai"`
LiteLLM ProviderConfig `json:"litellm"`
OpenRouter ProviderConfig `json:"openrouter"` OpenRouter ProviderConfig `json:"openrouter"`
Groq ProviderConfig `json:"groq"` Groq ProviderConfig `json:"groq"`
Zhipu ProviderConfig `json:"zhipu"` Zhipu ProviderConfig `json:"zhipu"`
@ -408,6 +440,7 @@ type ProvidersConfig struct {
func (p ProvidersConfig) IsEmpty() bool { func (p ProvidersConfig) IsEmpty() bool {
return p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" && return p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" &&
p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" && p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" &&
p.LiteLLM.APIKey == "" && p.LiteLLM.APIBase == "" &&
p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" && p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" &&
p.Groq.APIKey == "" && p.Groq.APIBase == "" && p.Groq.APIKey == "" && p.Groq.APIBase == "" &&
p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" && p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" &&
@ -516,14 +549,26 @@ type PerplexityConfig struct {
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
} }
type GLMSearchConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_GLM_ENABLED"`
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_GLM_API_KEY"`
BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_GLM_BASE_URL"`
// SearchEngine specifies the search backend: "search_std" (default),
// "search_pro", "search_pro_sogou", or "search_pro_quark".
SearchEngine string `json:"search_engine" env:"PICOCLAW_TOOLS_WEB_GLM_SEARCH_ENGINE"`
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_GLM_MAX_RESULTS"`
}
type WebToolsConfig struct { type WebToolsConfig struct {
Brave BraveConfig `json:"brave"` Brave BraveConfig `json:"brave"`
Tavily TavilyConfig `json:"tavily"` Tavily TavilyConfig `json:"tavily"`
DuckDuckGo DuckDuckGoConfig `json:"duckduckgo"` DuckDuckGo DuckDuckGoConfig `json:"duckduckgo"`
Perplexity PerplexityConfig `json:"perplexity"` Perplexity PerplexityConfig `json:"perplexity"`
GLMSearch GLMSearchConfig `json:"glm_search"`
// Proxy is an optional proxy URL for web tools (http/https/socks5/socks5h). // Proxy is an optional proxy URL for web tools (http/https/socks5/socks5h).
// For authenticated proxies, prefer HTTP_PROXY/HTTPS_PROXY env vars instead of embedding credentials in config. // For authenticated proxies, prefer HTTP_PROXY/HTTPS_PROXY env vars instead of embedding credentials in config.
Proxy string `json:"proxy,omitempty" env:"PICOCLAW_TOOLS_WEB_PROXY"` Proxy string `json:"proxy,omitempty" env:"PICOCLAW_TOOLS_WEB_PROXY"`
FetchLimitBytes int64 `json:"fetch_limit_bytes,omitempty" env:"PICOCLAW_TOOLS_WEB_FETCH_LIMIT_BYTES"`
} }
type CronToolsConfig struct { type CronToolsConfig struct {
@ -531,8 +576,9 @@ type CronToolsConfig struct {
} }
type ExecConfig struct { type ExecConfig struct {
EnableDenyPatterns bool `json:"enable_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS"` EnableDenyPatterns bool `json:"enable_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS"`
CustomDenyPatterns []string `json:"custom_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS"` CustomDenyPatterns []string `json:"custom_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS"`
CustomAllowPatterns []string `json:"custom_allow_patterns" env:"PICOCLAW_TOOLS_EXEC_CUSTOM_ALLOW_PATTERNS"`
} }
type MediaCleanupConfig struct { type MediaCleanupConfig struct {
@ -542,11 +588,14 @@ type MediaCleanupConfig struct {
} }
type ToolsConfig struct { type ToolsConfig struct {
Web WebToolsConfig `json:"web"` AllowReadPaths []string `json:"allow_read_paths" env:"PICOCLAW_TOOLS_ALLOW_READ_PATHS"`
Cron CronToolsConfig `json:"cron"` AllowWritePaths []string `json:"allow_write_paths" env:"PICOCLAW_TOOLS_ALLOW_WRITE_PATHS"`
Exec ExecConfig `json:"exec"` Web WebToolsConfig `json:"web"`
Skills SkillsToolsConfig `json:"skills"` Cron CronToolsConfig `json:"cron"`
MediaCleanup MediaCleanupConfig `json:"media_cleanup"` Exec ExecConfig `json:"exec"`
Skills SkillsToolsConfig `json:"skills"`
MediaCleanup MediaCleanupConfig `json:"media_cleanup"`
MCP MCPConfig `json:"mcp"`
} }
type SkillsToolsConfig struct { type SkillsToolsConfig struct {
@ -576,6 +625,123 @@ type ClawHubRegistryConfig struct {
MaxResponseSize int `json:"max_response_size" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_MAX_RESPONSE_SIZE"` MaxResponseSize int `json:"max_response_size" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_MAX_RESPONSE_SIZE"`
} }
// MCPServerConfig defines configuration for a single MCP server
type MCPServerConfig struct {
// Enabled indicates whether this MCP server is active
Enabled bool `json:"enabled"`
// Command is the executable to run (e.g., "npx", "python", "/path/to/server")
Command string `json:"command"`
// Args are the arguments to pass to the command
Args []string `json:"args,omitempty"`
// Env are environment variables to set for the server process (stdio only)
Env map[string]string `json:"env,omitempty"`
// EnvFile is the path to a file containing environment variables (stdio only)
EnvFile string `json:"env_file,omitempty"`
// Type is "stdio", "sse", or "http" (default: stdio if command is set, sse if url is set)
Type string `json:"type,omitempty"`
// URL is used for SSE/HTTP transport
URL string `json:"url,omitempty"`
// Headers are HTTP headers to send with requests (sse/http only)
Headers map[string]string `json:"headers,omitempty"`
}
// MCPConfig defines configuration for all MCP servers
type MCPConfig struct {
// Enabled globally enables/disables MCP integration
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_MCP_ENABLED"`
// Servers is a map of server name to server configuration
Servers map[string]MCPServerConfig `json:"servers,omitempty"`
}
// WebChatConfig configures web chat modes
type WebChatConfig struct {
HTTP bool `json:"http" env:"PICOCLAW_WEB_CHAT_HTTP"` // HTTP polling mode (default: true)
WebSocket bool `json:"websocket" env:"PICOCLAW_WEB_CHAT_WEBSOCKET"` // WebSocket real-time mode (default: true)
Timeout int `json:"timeout" env:"PICOCLAW_WEB_CHAT_TIMEOUT"` // Response timeout in seconds (default: 10)
}
// UnmarshalJSON handles JSON unmarshaling for WebChatConfig, accepting both boolean and numeric values
func (w *WebChatConfig) UnmarshalJSON(data []byte) error {
type Alias WebChatConfig
aux := &struct {
HTTP interface{} `json:"http"`
WebSocket interface{} `json:"websocket"`
Timeout int `json:"timeout"`
}{
Timeout: 10, // default
}
if err := json.Unmarshal(data, aux); err != nil {
return err
}
// Convert HTTP to boolean
w.HTTP = toBool(aux.HTTP)
// Convert WebSocket to boolean
w.WebSocket = toBool(aux.WebSocket)
// Set timeout
w.Timeout = aux.Timeout
if w.Timeout <= 0 {
w.Timeout = 10 // fallback to default
}
return nil
}
// Helper function to convert various types to boolean
func toBool(v interface{}) bool {
if v == nil {
return false
}
switch val := v.(type) {
case bool:
return val
case float64:
return val != 0
case string:
return val == "true" || val == "1" || val == "yes"
default:
return false
}
}
// WebConfig configures the web management UI
type WebConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_WEB_ENABLED"`
Host string `json:"host" env:"PICOCLAW_WEB_HOST"`
Port int `json:"port" env:"PICOCLAW_WEB_PORT"`
Username string `json:"username" env:"PICOCLAW_WEB_USERNAME"`
Password string `json:"password" env:"PICOCLAW_WEB_PASSWORD"`
Chat WebChatConfig `json:"chat"`
}
// VoiceConfig configures voice features
type VoiceConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_VOICE_ENABLED"`
TTS TTSConfig `json:"tts"`
STT STTConfig `json:"stt"`
}
// TTSConfig configures text-to-speech
type TTSConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_VOICE_TTS_ENABLED"`
Provider string `json:"provider" env:"PICOCLAW_VOICE_TTS_PROVIDER"` // aliyun/openai
APIKey string `json:"api_key" env:"PICOCLAW_VOICE_TTS_API_KEY"`
Voice string `json:"voice" env:"PICOCLAW_VOICE_TTS_VOICE"`
Speed int `json:"speed" env:"PICOCLAW_VOICE_TTS_SPEED"`
Volume int `json:"volume" env:"PICOCLAW_VOICE_TTS_VOLUME"`
Pitch int `json:"pitch" env:"PICOCLAW_VOICE_TTS_PITCH"`
Format string `json:"format" env:"PICOCLAW_VOICE_TTS_FORMAT"`
SampleRate int `json:"sample_rate" env:"PICOCLAW_VOICE_TTS_SAMPLE_RATE"`
}
// STTConfig configures speech-to-text
type STTConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_VOICE_STT_ENABLED"`
Provider string `json:"provider" env:"PICOCLAW_VOICE_STT_PROVIDER"` // groq/aliyun
APIKey string `json:"api_key" env:"PICOCLAW_VOICE_STT_API_KEY"`
}
func LoadConfig(path string) (*Config, error) { func LoadConfig(path string) (*Config, error) {
cfg := DefaultConfig() cfg := DefaultConfig()
@ -632,7 +798,8 @@ func (c *Config) migrateChannelConfigs() {
} }
// OneBot: group_trigger_prefix -> group_trigger.prefixes // OneBot: group_trigger_prefix -> group_trigger.prefixes
if len(c.Channels.OneBot.GroupTriggerPrefix) > 0 && len(c.Channels.OneBot.GroupTrigger.Prefixes) == 0 { if len(c.Channels.OneBot.GroupTriggerPrefix) > 0 &&
len(c.Channels.OneBot.GroupTrigger.Prefixes) == 0 {
c.Channels.OneBot.GroupTrigger.Prefixes = c.Channels.OneBot.GroupTriggerPrefix c.Channels.OneBot.GroupTrigger.Prefixes = c.Channels.OneBot.GroupTriggerPrefix
} }
} }
@ -742,25 +909,7 @@ func (c *Config) findMatches(modelName string) []ModelConfig {
// HasProvidersConfig checks if any provider in the old providers config has configuration. // HasProvidersConfig checks if any provider in the old providers config has configuration.
func (c *Config) HasProvidersConfig() bool { func (c *Config) HasProvidersConfig() bool {
v := c.Providers return !c.Providers.IsEmpty()
return v.Anthropic.APIKey != "" || v.Anthropic.APIBase != "" ||
v.OpenAI.APIKey != "" || v.OpenAI.APIBase != "" ||
v.OpenRouter.APIKey != "" || v.OpenRouter.APIBase != "" ||
v.Groq.APIKey != "" || v.Groq.APIBase != "" ||
v.Zhipu.APIKey != "" || v.Zhipu.APIBase != "" ||
v.VLLM.APIKey != "" || v.VLLM.APIBase != "" ||
v.Gemini.APIKey != "" || v.Gemini.APIBase != "" ||
v.Nvidia.APIKey != "" || v.Nvidia.APIBase != "" ||
v.Ollama.APIKey != "" || v.Ollama.APIBase != "" ||
v.Moonshot.APIKey != "" || v.Moonshot.APIBase != "" ||
v.ShengSuanYun.APIKey != "" || v.ShengSuanYun.APIBase != "" ||
v.DeepSeek.APIKey != "" || v.DeepSeek.APIBase != "" ||
v.Cerebras.APIKey != "" || v.Cerebras.APIBase != "" ||
v.VolcEngine.APIKey != "" || v.VolcEngine.APIBase != "" ||
v.GitHubCopilot.APIKey != "" || v.GitHubCopilot.APIBase != "" ||
v.Antigravity.APIKey != "" || v.Antigravity.APIBase != "" ||
v.Qwen.APIKey != "" || v.Qwen.APIBase != "" ||
v.Mistral.APIKey != "" || v.Mistral.APIBase != ""
} }
// ValidateModelList validates all ModelConfig entries in the model_list. // ValidateModelList validates all ModelConfig entries in the model_list.

View file

@ -435,6 +435,18 @@ func TestLoadConfig_WebToolsProxy(t *testing.T) {
} }
// TestDefaultConfig_DMScope verifies the default dm_scope value // TestDefaultConfig_DMScope verifies the default dm_scope value
// TestDefaultConfig_SummarizationThresholds verifies summarization defaults
func TestDefaultConfig_SummarizationThresholds(t *testing.T) {
cfg := DefaultConfig()
if cfg.Agents.Defaults.SummarizeMessageThreshold != 20 {
t.Errorf("SummarizeMessageThreshold = %d, want 20", cfg.Agents.Defaults.SummarizeMessageThreshold)
}
if cfg.Agents.Defaults.SummarizeTokenPercent != 75 {
t.Errorf("SummarizeTokenPercent = %d, want 75", cfg.Agents.Defaults.SummarizeTokenPercent)
}
}
func TestDefaultConfig_DMScope(t *testing.T) { func TestDefaultConfig_DMScope(t *testing.T) {
cfg := DefaultConfig() cfg := DefaultConfig()
@ -442,3 +454,28 @@ func TestDefaultConfig_DMScope(t *testing.T) {
t.Errorf("Session.DMScope = %q, want 'per-channel-peer'", cfg.Session.DMScope) t.Errorf("Session.DMScope = %q, want 'per-channel-peer'", cfg.Session.DMScope)
} }
} }
func TestDefaultConfig_WorkspacePath_Default(t *testing.T) {
// Unset to ensure we test the default
t.Setenv("PICOCLAW_HOME", "")
// Set a known home for consistent test results
t.Setenv("HOME", "/tmp/home")
cfg := DefaultConfig()
want := filepath.Join("/tmp/home", ".picoclaw", "workspace")
if cfg.Agents.Defaults.Workspace != want {
t.Errorf("Default workspace path = %q, want %q", cfg.Agents.Defaults.Workspace, want)
}
}
func TestDefaultConfig_WorkspacePath_WithPicoclawHome(t *testing.T) {
t.Setenv("PICOCLAW_HOME", "/custom/picoclaw/home")
cfg := DefaultConfig()
want := "/custom/picoclaw/home/workspace"
if cfg.Agents.Defaults.Workspace != want {
t.Errorf("Workspace path with PICOCLAW_HOME = %q, want %q", cfg.Agents.Defaults.Workspace, want)
}
}

View file

@ -5,18 +5,36 @@
package config package config
import (
"os"
"path/filepath"
)
// DefaultConfig returns the default configuration for PicoClaw. // DefaultConfig returns the default configuration for PicoClaw.
func DefaultConfig() *Config { func DefaultConfig() *Config {
// Determine the base path for the workspace.
// Priority: $PICOCLAW_HOME > ~/.picoclaw
var homePath string
if picoclawHome := os.Getenv("PICOCLAW_HOME"); picoclawHome != "" {
homePath = picoclawHome
} else {
userHome, _ := os.UserHomeDir()
homePath = filepath.Join(userHome, ".picoclaw")
}
workspacePath := filepath.Join(homePath, "workspace")
return &Config{ return &Config{
Agents: AgentsConfig{ Agents: AgentsConfig{
Defaults: AgentDefaults{ Defaults: AgentDefaults{
Workspace: "~/.picoclaw/workspace", Workspace: workspacePath,
RestrictToWorkspace: true, RestrictToWorkspace: true,
Provider: "", Provider: "",
Model: "", Model: "",
MaxTokens: 32768, MaxTokens: 32768,
Temperature: nil, // nil means use provider default Temperature: nil, // nil means use provider default
MaxToolIterations: 50, MaxToolIterations: 50,
SummarizeMessageThreshold: 20,
SummarizeTokenPercent: 75,
}, },
}, },
Bindings: []AgentBinding{}, Bindings: []AgentBinding{},
@ -121,6 +139,16 @@ func DefaultConfig() *Config {
AllowFrom: FlexibleStringSlice{}, AllowFrom: FlexibleStringSlice{},
ReplyTimeout: 5, ReplyTimeout: 5,
}, },
WeComAIBot: WeComAIBotConfig{
Enabled: false,
Token: "",
EncodingAESKey: "",
WebhookPath: "/webhook/wecom-aibot",
AllowFrom: FlexibleStringSlice{},
ReplyTimeout: 5,
MaxSteps: 10,
WelcomeMessage: "Hello! I'm your AI assistant. How can I help you today?",
},
Pico: PicoConfig{ Pico: PicoConfig{
Enabled: false, Enabled: false,
Token: "", Token: "",
@ -299,7 +327,8 @@ func DefaultConfig() *Config {
Interval: 5, Interval: 5,
}, },
Web: WebToolsConfig{ Web: WebToolsConfig{
Proxy: "", Proxy: "",
FetchLimitBytes: 10 * 1024 * 1024, // 10MB by default
Brave: BraveConfig{ Brave: BraveConfig{
Enabled: false, Enabled: false,
APIKey: "", APIKey: "",
@ -314,6 +343,13 @@ func DefaultConfig() *Config {
APIKey: "", APIKey: "",
MaxResults: 5, MaxResults: 5,
}, },
GLMSearch: GLMSearchConfig{
Enabled: false,
APIKey: "",
BaseURL: "https://open.bigmodel.cn/api/paas/v4/web_search",
SearchEngine: "search_std",
MaxResults: 5,
},
}, },
Cron: CronToolsConfig{ Cron: CronToolsConfig{
ExecTimeoutMinutes: 5, ExecTimeoutMinutes: 5,
@ -334,6 +370,10 @@ func DefaultConfig() *Config {
TTLSeconds: 300, TTLSeconds: 300,
}, },
}, },
MCP: MCPConfig{
Enabled: false,
Servers: map[string]MCPServerConfig{},
},
}, },
Heartbeat: HeartbeatConfig{ Heartbeat: HeartbeatConfig{
Enabled: true, Enabled: true,

View file

@ -88,6 +88,23 @@ func ConvertProvidersToModelList(cfg *Config) []ModelConfig {
}, true }, true
}, },
}, },
{
providerNames: []string{"litellm"},
protocol: "litellm",
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
if p.LiteLLM.APIKey == "" && p.LiteLLM.APIBase == "" {
return ModelConfig{}, false
}
return ModelConfig{
ModelName: "litellm",
Model: "litellm/auto",
APIKey: p.LiteLLM.APIKey,
APIBase: p.LiteLLM.APIBase,
Proxy: p.LiteLLM.Proxy,
RequestTimeout: p.LiteLLM.RequestTimeout,
}, true
},
},
{ {
providerNames: []string{"openrouter"}, providerNames: []string{"openrouter"},
protocol: "openrouter", protocol: "openrouter",

View file

@ -63,6 +63,33 @@ func TestConvertProvidersToModelList_Anthropic(t *testing.T) {
} }
} }
func TestConvertProvidersToModelList_LiteLLM(t *testing.T) {
cfg := &Config{
Providers: ProvidersConfig{
LiteLLM: ProviderConfig{
APIKey: "litellm-key",
APIBase: "http://localhost:4000/v1",
},
},
}
result := ConvertProvidersToModelList(cfg)
if len(result) != 1 {
t.Fatalf("len(result) = %d, want 1", len(result))
}
if result[0].ModelName != "litellm" {
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "litellm")
}
if result[0].Model != "litellm/auto" {
t.Errorf("Model = %q, want %q", result[0].Model, "litellm/auto")
}
if result[0].APIBase != "http://localhost:4000/v1" {
t.Errorf("APIBase = %q, want %q", result[0].APIBase, "http://localhost:4000/v1")
}
}
func TestConvertProvidersToModelList_Multiple(t *testing.T) { func TestConvertProvidersToModelList_Multiple(t *testing.T) {
cfg := &Config{ cfg := &Config{
Providers: ProvidersConfig{ Providers: ProvidersConfig{
@ -115,6 +142,7 @@ func TestConvertProvidersToModelList_AllProviders(t *testing.T) {
cfg := &Config{ cfg := &Config{
Providers: ProvidersConfig{ Providers: ProvidersConfig{
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "key1"}}, OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "key1"}},
LiteLLM: ProviderConfig{APIKey: "key-litellm", APIBase: "http://localhost:4000/v1"},
Anthropic: ProviderConfig{APIKey: "key2"}, Anthropic: ProviderConfig{APIKey: "key2"},
OpenRouter: ProviderConfig{APIKey: "key3"}, OpenRouter: ProviderConfig{APIKey: "key3"},
Groq: ProviderConfig{APIKey: "key4"}, Groq: ProviderConfig{APIKey: "key4"},
@ -137,9 +165,9 @@ func TestConvertProvidersToModelList_AllProviders(t *testing.T) {
result := ConvertProvidersToModelList(cfg) result := ConvertProvidersToModelList(cfg)
// All 18 providers should be converted // All 19 providers should be converted
if len(result) != 18 { if len(result) != 19 {
t.Errorf("len(result) = %d, want 18", len(result)) t.Errorf("len(result) = %d, want 19", len(result))
} }
} }

View file

@ -64,7 +64,7 @@ func TestGetModelConfig_RoundRobin(t *testing.T) {
// Test round-robin distribution // Test round-robin distribution
results := make(map[string]int) results := make(map[string]int)
for i := 0; i < 30; i++ { for range 30 {
result, err := cfg.GetModelConfig("lb-model") result, err := cfg.GetModelConfig("lb-model")
if err != nil { if err != nil {
t.Fatalf("GetModelConfig() error = %v", err) t.Fatalf("GetModelConfig() error = %v", err)
@ -94,17 +94,15 @@ func TestGetModelConfig_Concurrent(t *testing.T) {
var wg sync.WaitGroup var wg sync.WaitGroup
errors := make(chan error, goroutines*iterations) errors := make(chan error, goroutines*iterations)
for i := 0; i < goroutines; i++ { for range goroutines {
wg.Add(1) wg.Go(func() {
go func() { for range iterations {
defer wg.Done()
for j := 0; j < iterations; j++ {
_, err := cfg.GetModelConfig("concurrent-model") _, err := cfg.GetModelConfig("concurrent-model")
if err != nil { if err != nil {
errors <- err errors <- err
} }
} }
}() })
} }
wg.Wait() wg.Wait()

View file

@ -4,6 +4,7 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"maps"
"net/http" "net/http"
"sync" "sync"
"time" "time"
@ -122,9 +123,7 @@ func (s *Server) readyHandler(w http.ResponseWriter, r *http.Request) {
s.mu.RLock() s.mu.RLock()
ready := s.ready ready := s.ready
checks := make(map[string]Check) checks := make(map[string]Check)
for k, v := range s.checks { maps.Copy(checks, s.checks)
checks[k] = v
}
s.mu.RUnlock() s.mu.RUnlock()
if !ready { if !ready {

View file

@ -47,79 +47,63 @@ func TestExecuteHeartbeat_Async(t *testing.T) {
} }
} }
func TestExecuteHeartbeat_Error(t *testing.T) { func TestExecuteHeartbeat_ResultLogging(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "heartbeat-test-*") tests := []struct {
if err != nil { name string
t.Fatalf("Failed to create temp dir: %v", err) result *tools.ToolResult
} wantLog string
defer os.RemoveAll(tmpDir) }{
{
hs := NewHeartbeatService(tmpDir, 30, true) name: "error result",
hs.stopChan = make(chan struct{}) // Enable for testing result: &tools.ToolResult{
ForLLM: "Heartbeat failed: connection error",
hs.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult { ForUser: "",
return &tools.ToolResult{ Silent: false,
ForLLM: "Heartbeat failed: connection error", IsError: true,
ForUser: "", Async: false,
Silent: false, },
IsError: true, wantLog: "error message",
Async: false, },
} {
}) name: "silent result",
result: &tools.ToolResult{
// Create HEARTBEAT.md ForLLM: "Heartbeat completed successfully",
os.WriteFile(filepath.Join(tmpDir, "HEARTBEAT.md"), []byte("Test task"), 0o644) ForUser: "",
Silent: true,
hs.executeHeartbeat() IsError: false,
Async: false,
// Check log file for error message },
logFile := filepath.Join(tmpDir, "heartbeat.log") wantLog: "completion message",
data, err := os.ReadFile(logFile) },
if err != nil {
t.Fatalf("Failed to read log file: %v", err)
} }
logContent := string(data) for _, tt := range tests {
if logContent == "" { t.Run(tt.name, func(t *testing.T) {
t.Error("Expected log file to contain error message") tmpDir, err := os.MkdirTemp("", "heartbeat-test-*")
} if err != nil {
} t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
func TestExecuteHeartbeat_Silent(t *testing.T) { hs := NewHeartbeatService(tmpDir, 30, true)
tmpDir, err := os.MkdirTemp("", "heartbeat-test-*") hs.stopChan = make(chan struct{}) // Enable for testing
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
hs := NewHeartbeatService(tmpDir, 30, true) hs.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult {
hs.stopChan = make(chan struct{}) // Enable for testing return tt.result
})
hs.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult { os.WriteFile(filepath.Join(tmpDir, "HEARTBEAT.md"), []byte("Test task"), 0o644)
return &tools.ToolResult{ hs.executeHeartbeat()
ForLLM: "Heartbeat completed successfully",
ForUser: "",
Silent: true,
IsError: false,
Async: false,
}
})
// Create HEARTBEAT.md logFile := filepath.Join(tmpDir, "heartbeat.log")
os.WriteFile(filepath.Join(tmpDir, "HEARTBEAT.md"), []byte("Test task"), 0o644) data, err := os.ReadFile(logFile)
if err != nil {
hs.executeHeartbeat() t.Fatalf("Failed to read log file: %v", err)
}
// Check log file for completion message if string(data) == "" {
logFile := filepath.Join(tmpDir, "heartbeat.log") t.Errorf("Expected log file to contain %s", tt.wantLog)
data, err := os.ReadFile(logFile) }
if err != nil { })
t.Fatalf("Failed to read log file: %v", err)
}
logContent := string(data)
if logContent == "" {
t.Error("Expected log file to contain completion message")
} }
} }

532
pkg/mcp/manager.go Normal file
View file

@ -0,0 +1,532 @@
package mcp
import (
"bufio"
"context"
"errors"
"fmt"
"net/http"
"os"
"os/exec"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/logger"
)
// headerTransport is an http.RoundTripper that adds custom headers to requests
type headerTransport struct {
base http.RoundTripper
headers map[string]string
}
func (t *headerTransport) RoundTrip(req *http.Request) (*http.Response, error) {
// Clone the request to avoid modifying the original
req = req.Clone(req.Context())
// Add custom headers
for key, value := range t.headers {
req.Header.Set(key, value)
}
// Use the base transport
base := t.base
if base == nil {
base = http.DefaultTransport
}
return base.RoundTrip(req)
}
// loadEnvFile loads environment variables from a file in .env format
// Each line should be in the format: KEY=value
// Lines starting with # are comments
// Empty lines are ignored
func loadEnvFile(path string) (map[string]string, error) {
file, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("failed to open env file: %w", err)
}
defer file.Close()
envVars := make(map[string]string)
scanner := bufio.NewScanner(file)
lineNum := 0
for scanner.Scan() {
lineNum++
line := strings.TrimSpace(scanner.Text())
// Skip empty lines and comments
if line == "" || strings.HasPrefix(line, "#") {
continue
}
// Parse KEY=value
parts := strings.SplitN(line, "=", 2)
if len(parts) != 2 {
return nil, fmt.Errorf("invalid format at line %d: %s", lineNum, line)
}
key := strings.TrimSpace(parts[0])
value := strings.TrimSpace(parts[1])
if key == "" {
return nil, fmt.Errorf("invalid format at line %d: empty key", lineNum)
}
// Remove surrounding quotes if present
if len(value) >= 2 {
if (value[0] == '"' && value[len(value)-1] == '"') ||
(value[0] == '\'' && value[len(value)-1] == '\'') {
value = value[1 : len(value)-1]
}
}
envVars[key] = value
}
if err := scanner.Err(); err != nil {
return nil, fmt.Errorf("error reading env file: %w", err)
}
return envVars, nil
}
// ServerConnection represents a connection to an MCP server
type ServerConnection struct {
Name string
Client *mcp.Client
Session *mcp.ClientSession
Tools []*mcp.Tool
}
// Manager manages multiple MCP server connections
type Manager struct {
servers map[string]*ServerConnection
mu sync.RWMutex
closed atomic.Bool // changed from bool to atomic.Bool to avoid TOCTOU race
wg sync.WaitGroup // tracks in-flight CallTool calls
}
// NewManager creates a new MCP manager
func NewManager() *Manager {
return &Manager{
servers: make(map[string]*ServerConnection),
}
}
// LoadFromConfig loads MCP servers from configuration
func (m *Manager) LoadFromConfig(ctx context.Context, cfg *config.Config) error {
return m.LoadFromMCPConfig(ctx, cfg.Tools.MCP, cfg.WorkspacePath())
}
// LoadFromMCPConfig loads MCP servers from MCP configuration and workspace path.
// This is the minimal dependency version that doesn't require the full Config object.
func (m *Manager) LoadFromMCPConfig(
ctx context.Context,
mcpCfg config.MCPConfig,
workspacePath string,
) error {
if !mcpCfg.Enabled {
logger.InfoCF("mcp", "MCP integration is disabled", nil)
return nil
}
if len(mcpCfg.Servers) == 0 {
logger.InfoCF("mcp", "No MCP servers configured", nil)
return nil
}
logger.InfoCF("mcp", "Initializing MCP servers",
map[string]any{
"count": len(mcpCfg.Servers),
})
var wg sync.WaitGroup
errs := make(chan error, len(mcpCfg.Servers))
enabledCount := 0
for name, serverCfg := range mcpCfg.Servers {
if !serverCfg.Enabled {
logger.DebugCF("mcp", "Skipping disabled server",
map[string]any{
"server": name,
})
continue
}
enabledCount++
wg.Add(1)
go func(name string, serverCfg config.MCPServerConfig, workspace string) {
defer wg.Done()
// Resolve relative envFile paths relative to workspace
if serverCfg.EnvFile != "" && !filepath.IsAbs(serverCfg.EnvFile) {
if workspace == "" {
err := fmt.Errorf(
"workspace path is empty while resolving relative envFile %q for server %s",
serverCfg.EnvFile,
name,
)
logger.ErrorCF("mcp", "Invalid MCP server configuration",
map[string]any{
"server": name,
"env_file": serverCfg.EnvFile,
"error": err.Error(),
})
errs <- err
return
}
serverCfg.EnvFile = filepath.Join(workspace, serverCfg.EnvFile)
}
if err := m.ConnectServer(ctx, name, serverCfg); err != nil {
logger.ErrorCF("mcp", "Failed to connect to MCP server",
map[string]any{
"server": name,
"error": err.Error(),
})
errs <- fmt.Errorf("failed to connect to server %s: %w", name, err)
}
}(name, serverCfg, workspacePath)
}
wg.Wait()
close(errs)
// Collect errors
var allErrors []error
for err := range errs {
allErrors = append(allErrors, err)
}
connectedCount := len(m.GetServers())
// If all enabled servers failed to connect, return aggregated error
if enabledCount > 0 && connectedCount == 0 {
logger.ErrorCF("mcp", "All MCP servers failed to connect",
map[string]any{
"failed": len(allErrors),
"total": enabledCount,
})
return errors.Join(allErrors...)
}
if len(allErrors) > 0 {
logger.WarnCF("mcp", "Some MCP servers failed to connect",
map[string]any{
"failed": len(allErrors),
"connected": connectedCount,
"total": enabledCount,
})
// Don't fail completely if some servers successfully connected
}
logger.InfoCF("mcp", "MCP server initialization complete",
map[string]any{
"connected": connectedCount,
"total": enabledCount,
})
return nil
}
// ConnectServer connects to a single MCP server
func (m *Manager) ConnectServer(
ctx context.Context,
name string,
cfg config.MCPServerConfig,
) error {
logger.InfoCF("mcp", "Connecting to MCP server",
map[string]any{
"server": name,
"command": cfg.Command,
"args_count": len(cfg.Args),
})
// Create client
client := mcp.NewClient(&mcp.Implementation{
Name: "picoclaw",
Version: "1.0.0",
}, nil)
// Create transport based on configuration
// Auto-detect transport type if not explicitly specified
var transport mcp.Transport
transportType := cfg.Type
// Auto-detect: if URL is provided, use SSE; if command is provided, use stdio
if transportType == "" {
if cfg.URL != "" {
transportType = "sse"
} else if cfg.Command != "" {
transportType = "stdio"
} else {
return fmt.Errorf("either URL or command must be provided")
}
}
switch transportType {
case "sse", "http":
if cfg.URL == "" {
return fmt.Errorf("URL is required for SSE/HTTP transport")
}
logger.DebugCF("mcp", "Using SSE/HTTP transport",
map[string]any{
"server": name,
"url": cfg.URL,
})
sseTransport := &mcp.StreamableClientTransport{
Endpoint: cfg.URL,
}
// Add custom headers if provided
if len(cfg.Headers) > 0 {
// Create a custom HTTP client with header-injecting transport
sseTransport.HTTPClient = &http.Client{
Transport: &headerTransport{
base: http.DefaultTransport,
headers: cfg.Headers,
},
}
logger.DebugCF("mcp", "Added custom HTTP headers",
map[string]any{
"server": name,
"header_count": len(cfg.Headers),
})
}
transport = sseTransport
case "stdio":
if cfg.Command == "" {
return fmt.Errorf("command is required for stdio transport")
}
logger.DebugCF("mcp", "Using stdio transport",
map[string]any{
"server": name,
"command": cfg.Command,
})
// Create command with context
cmd := exec.CommandContext(ctx, cfg.Command, cfg.Args...)
// Build environment variables with proper override semantics
// Use a map to ensure config variables override file variables
envMap := make(map[string]string)
// Start with parent process environment
for _, e := range cmd.Environ() {
if idx := strings.Index(e, "="); idx > 0 {
envMap[e[:idx]] = e[idx+1:]
}
}
// Load environment variables from file if specified
if cfg.EnvFile != "" {
envVars, err := loadEnvFile(cfg.EnvFile)
if err != nil {
return fmt.Errorf("failed to load env file %s: %w", cfg.EnvFile, err)
}
for k, v := range envVars {
envMap[k] = v
}
logger.DebugCF("mcp", "Loaded environment variables from file",
map[string]any{
"server": name,
"envFile": cfg.EnvFile,
"var_count": len(envVars),
})
}
// Environment variables from config override those from file
for k, v := range cfg.Env {
envMap[k] = v
}
// Convert map to slice
env := make([]string, 0, len(envMap))
for k, v := range envMap {
env = append(env, fmt.Sprintf("%s=%s", k, v))
}
cmd.Env = env
transport = &mcp.CommandTransport{Command: cmd}
default:
return fmt.Errorf(
"unsupported transport type: %s (supported: stdio, sse, http)",
transportType,
)
}
// Connect to server
session, err := client.Connect(ctx, transport, nil)
if err != nil {
return fmt.Errorf("failed to connect: %w", err)
}
// Get server info
initResult := session.InitializeResult()
logger.InfoCF("mcp", "Connected to MCP server",
map[string]any{
"server": name,
"serverName": initResult.ServerInfo.Name,
"serverVersion": initResult.ServerInfo.Version,
"protocol": initResult.ProtocolVersion,
})
// List available tools if supported
var tools []*mcp.Tool
if initResult.Capabilities.Tools != nil {
for tool, err := range session.Tools(ctx, nil) {
if err != nil {
logger.WarnCF("mcp", "Error listing tool",
map[string]any{
"server": name,
"error": err.Error(),
})
continue
}
tools = append(tools, tool)
}
logger.InfoCF("mcp", "Listed tools from MCP server",
map[string]any{
"server": name,
"toolCount": len(tools),
})
}
// Store connection
m.mu.Lock()
m.servers[name] = &ServerConnection{
Name: name,
Client: client,
Session: session,
Tools: tools,
}
m.mu.Unlock()
return nil
}
// GetServers returns all connected servers
func (m *Manager) GetServers() map[string]*ServerConnection {
m.mu.RLock()
defer m.mu.RUnlock()
result := make(map[string]*ServerConnection, len(m.servers))
for k, v := range m.servers {
result[k] = v
}
return result
}
// GetServer returns a specific server connection
func (m *Manager) GetServer(name string) (*ServerConnection, bool) {
m.mu.RLock()
defer m.mu.RUnlock()
conn, ok := m.servers[name]
return conn, ok
}
// CallTool calls a tool on a specific server
func (m *Manager) CallTool(
ctx context.Context,
serverName, toolName string,
arguments map[string]any,
) (*mcp.CallToolResult, error) {
// Check if closed before acquiring lock (fast path)
if m.closed.Load() {
return nil, fmt.Errorf("manager is closed")
}
m.mu.RLock()
// Double-check after acquiring lock to prevent TOCTOU race
if m.closed.Load() {
m.mu.RUnlock()
return nil, fmt.Errorf("manager is closed")
}
conn, ok := m.servers[serverName]
if ok {
m.wg.Add(1) // Add to WaitGroup while holding the lock
}
m.mu.RUnlock()
if !ok {
return nil, fmt.Errorf("server %s not found", serverName)
}
defer m.wg.Done()
params := &mcp.CallToolParams{
Name: toolName,
Arguments: arguments,
}
result, err := conn.Session.CallTool(ctx, params)
if err != nil {
return nil, fmt.Errorf("failed to call tool: %w", err)
}
return result, nil
}
// Close closes all server connections
func (m *Manager) Close() error {
// Use Swap to atomically set closed=true and get the previous value
// This prevents TOCTOU race with CallTool's closed check
if m.closed.Swap(true) {
return nil // already closed
}
// Wait for all in-flight CallTool calls to finish before closing sessions
// After closed=true is set, no new CallTool can start (they check closed first)
m.wg.Wait()
m.mu.Lock()
defer m.mu.Unlock()
logger.InfoCF("mcp", "Closing all MCP server connections",
map[string]any{
"count": len(m.servers),
})
var errs []error
for name, conn := range m.servers {
if err := conn.Session.Close(); err != nil {
logger.ErrorCF("mcp", "Failed to close server connection",
map[string]any{
"server": name,
"error": err.Error(),
})
errs = append(errs, fmt.Errorf("server %s: %w", name, err))
}
}
m.servers = make(map[string]*ServerConnection)
if len(errs) > 0 {
return fmt.Errorf("failed to close %d server(s): %w", len(errs), errors.Join(errs...))
}
return nil
}
// GetAllTools returns all tools from all connected servers
func (m *Manager) GetAllTools() map[string][]*mcp.Tool {
m.mu.RLock()
defer m.mu.RUnlock()
result := make(map[string][]*mcp.Tool)
for name, conn := range m.servers {
if len(conn.Tools) > 0 {
result[name] = conn.Tools
}
}
return result
}

298
pkg/mcp/manager_test.go Normal file
View file

@ -0,0 +1,298 @@
package mcp
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/sipeed/picoclaw/pkg/config"
)
func TestLoadEnvFile(t *testing.T) {
tests := []struct {
name string
content string
expected map[string]string
expectErr bool
}{
{
name: "basic env file",
content: `API_KEY=secret123
DATABASE_URL=postgres://localhost/db
PORT=8080`,
expected: map[string]string{
"API_KEY": "secret123",
"DATABASE_URL": "postgres://localhost/db",
"PORT": "8080",
},
expectErr: false,
},
{
name: "with comments and empty lines",
content: `# This is a comment
API_KEY=secret123
# Another comment
DATABASE_URL=postgres://localhost/db
PORT=8080`,
expected: map[string]string{
"API_KEY": "secret123",
"DATABASE_URL": "postgres://localhost/db",
"PORT": "8080",
},
expectErr: false,
},
{
name: "with quoted values",
content: `API_KEY="secret with spaces"
NAME='single quoted'
PLAIN=no-quotes`,
expected: map[string]string{
"API_KEY": "secret with spaces",
"NAME": "single quoted",
"PLAIN": "no-quotes",
},
expectErr: false,
},
{
name: "with spaces around equals",
content: `API_KEY = secret123
DATABASE_URL= postgres://localhost/db
PORT =8080`,
expected: map[string]string{
"API_KEY": "secret123",
"DATABASE_URL": "postgres://localhost/db",
"PORT": "8080",
},
expectErr: false,
},
{
name: "invalid format - no equals",
content: `INVALID_LINE`,
expectErr: true,
},
{
name: "empty file",
content: ``,
expected: map[string]string{},
expectErr: false,
},
{
name: "only comments",
content: `# Comment 1
# Comment 2`,
expected: map[string]string{},
expectErr: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tmpDir := t.TempDir()
envFile := filepath.Join(tmpDir, ".env")
if err := os.WriteFile(envFile, []byte(tt.content), 0o644); err != nil {
t.Fatalf("Failed to create test file: %v", err)
}
result, err := loadEnvFile(envFile)
if tt.expectErr {
if err == nil {
t.Errorf("Expected error but got none")
}
return
}
if err != nil {
t.Errorf("Unexpected error: %v", err)
return
}
if len(result) != len(tt.expected) {
t.Errorf("Expected %d variables, got %d", len(tt.expected), len(result))
}
for key, expectedValue := range tt.expected {
if actualValue, ok := result[key]; !ok {
t.Errorf("Expected key %s not found", key)
} else if actualValue != expectedValue {
t.Errorf("For key %s: expected %q, got %q", key, expectedValue, actualValue)
}
}
})
}
}
func TestLoadEnvFileNotFound(t *testing.T) {
_, err := loadEnvFile("/nonexistent/file.env")
if err == nil {
t.Error("Expected error for nonexistent file")
}
}
func TestEnvFilePriority(t *testing.T) {
// Create a temporary .env file
tmpDir := t.TempDir()
envFile := filepath.Join(tmpDir, ".env")
envContent := `API_KEY=from_file
DATABASE_URL=from_file
SHARED_VAR=from_file`
if err := os.WriteFile(envFile, []byte(envContent), 0o644); err != nil {
t.Fatalf("Failed to create .env file: %v", err)
}
// Load envFile
envVars, err := loadEnvFile(envFile)
if err != nil {
t.Fatalf("Failed to load env file: %v", err)
}
// Verify envFile variables
if envVars["API_KEY"] != "from_file" {
t.Errorf("Expected API_KEY=from_file, got %s", envVars["API_KEY"])
}
// Simulate config.Env overriding envFile
configEnv := map[string]string{
"SHARED_VAR": "from_config",
"NEW_VAR": "from_config",
}
// Merge: envFile first, then config overrides
merged := make(map[string]string)
for k, v := range envVars {
merged[k] = v
}
for k, v := range configEnv {
merged[k] = v
}
// Verify priority: config.Env should override envFile
if merged["SHARED_VAR"] != "from_config" {
t.Errorf(
"Expected SHARED_VAR=from_config (config should override file), got %s",
merged["SHARED_VAR"],
)
}
if merged["API_KEY"] != "from_file" {
t.Errorf("Expected API_KEY=from_file, got %s", merged["API_KEY"])
}
if merged["NEW_VAR"] != "from_config" {
t.Errorf("Expected NEW_VAR=from_config, got %s", merged["NEW_VAR"])
}
}
func TestLoadFromMCPConfig_EmptyWorkspaceWithRelativeEnvFile(t *testing.T) {
mgr := NewManager()
mcpCfg := config.MCPConfig{
Enabled: true,
Servers: map[string]config.MCPServerConfig{
"test-server": {
Enabled: true,
Command: "echo",
Args: []string{"ok"},
EnvFile: ".env",
},
},
}
err := mgr.LoadFromMCPConfig(context.Background(), mcpCfg, "")
if err == nil {
t.Fatal("expected error for relative env_file with empty workspace path, got nil")
}
if !strings.Contains(err.Error(), "workspace path is empty") {
t.Fatalf("expected workspace path validation error, got: %v", err)
}
}
func TestNewManager_InitialState(t *testing.T) {
mgr := NewManager()
if mgr == nil {
t.Fatal("expected manager instance, got nil")
}
if len(mgr.GetServers()) != 0 {
t.Fatalf("expected no servers on new manager, got %d", len(mgr.GetServers()))
}
}
func TestLoadFromMCPConfig_DisabledOrEmptyServers(t *testing.T) {
mgr := NewManager()
err := mgr.LoadFromMCPConfig(context.Background(), config.MCPConfig{Enabled: false}, "/tmp")
if err != nil {
t.Fatalf("expected nil error when MCP disabled, got: %v", err)
}
err = mgr.LoadFromMCPConfig(context.Background(), config.MCPConfig{Enabled: true}, "/tmp")
if err != nil {
t.Fatalf("expected nil error when no servers configured, got: %v", err)
}
}
func TestGetServers_ReturnsCopy(t *testing.T) {
mgr := NewManager()
mgr.servers["s1"] = &ServerConnection{Name: "s1"}
servers := mgr.GetServers()
delete(servers, "s1")
if _, ok := mgr.GetServer("s1"); !ok {
t.Fatal("expected internal manager state to remain unchanged")
}
}
func TestGetAllTools_FiltersEmptyTools(t *testing.T) {
mgr := NewManager()
mgr.servers["empty"] = &ServerConnection{Name: "empty", Tools: nil}
mgr.servers["with-tools"] = &ServerConnection{Name: "with-tools", Tools: []*sdkmcp.Tool{{}}}
all := mgr.GetAllTools()
if _, ok := all["empty"]; ok {
t.Fatal("expected server without tools to be excluded")
}
if _, ok := all["with-tools"]; !ok {
t.Fatal("expected server with tools to be included")
}
}
func TestCallTool_ErrorsForClosedOrMissingServer(t *testing.T) {
t.Run("manager closed", func(t *testing.T) {
mgr := NewManager()
mgr.closed.Store(true)
_, err := mgr.CallTool(context.Background(), "s1", "tool", nil)
if err == nil || !strings.Contains(err.Error(), "manager is closed") {
t.Fatalf("expected manager closed error, got: %v", err)
}
})
t.Run("server missing", func(t *testing.T) {
mgr := NewManager()
_, err := mgr.CallTool(context.Background(), "missing", "tool", nil)
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("expected server not found error, got: %v", err)
}
})
}
func TestClose_IdempotentOnEmptyManager(t *testing.T) {
mgr := NewManager()
if err := mgr.Close(); err != nil {
t.Fatalf("first close should succeed, got: %v", err)
}
if err := mgr.Close(); err != nil {
t.Fatalf("second close should be idempotent, got: %v", err)
}
}

View file

@ -49,7 +49,7 @@ func TestReleaseAll(t *testing.T) {
paths := make([]string, 3) paths := make([]string, 3)
refs := make([]string, 3) refs := make([]string, 3)
for i := 0; i < 3; i++ { for i := range 3 {
paths[i] = createTempFile(t, dir, strings.Repeat("a", i+1)+".jpg") paths[i] = createTempFile(t, dir, strings.Repeat("a", i+1)+".jpg")
var err error var err error
refs[i], err = store.Store(paths[i], MediaMeta{Source: "test"}, "scope1") refs[i], err = store.Store(paths[i], MediaMeta{Source: "test"}, "scope1")
@ -228,12 +228,12 @@ func TestConcurrentSafety(t *testing.T) {
var wg sync.WaitGroup var wg sync.WaitGroup
wg.Add(goroutines) wg.Add(goroutines)
for g := 0; g < goroutines; g++ { for g := range goroutines {
go func(gIdx int) { go func(gIdx int) {
defer wg.Done() defer wg.Done()
scope := strings.Repeat("s", gIdx+1) scope := strings.Repeat("s", gIdx+1)
for i := 0; i < filesPerGoroutine; i++ { for i := range filesPerGoroutine {
path := createTempFile(t, dir, strings.Repeat("f", gIdx*filesPerGoroutine+i+1)+".tmp") path := createTempFile(t, dir, strings.Repeat("f", gIdx*filesPerGoroutine+i+1)+".tmp")
ref, err := store.Store(path, MediaMeta{Source: "test"}, scope) ref, err := store.Store(path, MediaMeta{Source: "test"}, scope)
if err != nil { if err != nil {
@ -448,11 +448,11 @@ func TestConcurrentCleanupSafety(t *testing.T) {
wg.Add(workers * 4) wg.Add(workers * 4)
// Store workers // Store workers
for w := 0; w < workers; w++ { for w := range workers {
go func(wIdx int) { go func(wIdx int) {
defer wg.Done() defer wg.Done()
scope := fmt.Sprintf("scope-%d", wIdx) scope := fmt.Sprintf("scope-%d", wIdx)
for i := 0; i < ops; i++ { for i := range ops {
p := createTempFile(t, dir, fmt.Sprintf("w%d-f%d.tmp", wIdx, i)) p := createTempFile(t, dir, fmt.Sprintf("w%d-f%d.tmp", wIdx, i))
store.Store(p, MediaMeta{Source: "test"}, scope) store.Store(p, MediaMeta{Source: "test"}, scope)
} }
@ -460,30 +460,30 @@ func TestConcurrentCleanupSafety(t *testing.T) {
} }
// Resolve workers // Resolve workers
for w := 0; w < workers; w++ { for range workers {
go func() { go func() {
defer wg.Done() defer wg.Done()
for i := 0; i < ops; i++ { for range ops {
store.Resolve("media://nonexistent") store.Resolve("media://nonexistent")
} }
}() }()
} }
// ReleaseAll workers // ReleaseAll workers
for w := 0; w < workers; w++ { for w := range workers {
go func(wIdx int) { go func(wIdx int) {
defer wg.Done() defer wg.Done()
for i := 0; i < ops; i++ { for range ops {
store.ReleaseAll(fmt.Sprintf("scope-%d", wIdx)) store.ReleaseAll(fmt.Sprintf("scope-%d", wIdx))
} }
}(w) }(w)
} }
// CleanExpired workers // CleanExpired workers
for w := 0; w < workers; w++ { for range workers {
go func() { go func() {
defer wg.Done() defer wg.Done()
for i := 0; i < ops; i++ { for range ops {
store.CleanExpired() store.CleanExpired()
} }
}() }()

View file

@ -118,64 +118,55 @@ func TestPlanWorkspaceMigration(t *testing.T) {
assert.GreaterOrEqual(t, len(actions), 1) assert.GreaterOrEqual(t, len(actions), 1)
} }
func TestPlanWorkspaceMigrationWithExistingDestination(t *testing.T) { func TestPlanWorkspaceMigrationExistingFile(t *testing.T) {
tmpDir := t.TempDir() tests := []struct {
srcWorkspace := filepath.Join(tmpDir, "src", "workspace") name string
dstWorkspace := filepath.Join(tmpDir, "dst", "workspace") force bool
wantActionType ActionType
}{
{
name: "backup when not forced",
force: false,
wantActionType: ActionBackup,
},
{
name: "copy when forced",
force: true,
wantActionType: ActionCopy,
},
}
err := os.MkdirAll(srcWorkspace, 0o755) for _, tt := range tests {
require.NoError(t, err) t.Run(tt.name, func(t *testing.T) {
tmpDir := t.TempDir()
srcWorkspace := filepath.Join(tmpDir, "src", "workspace")
dstWorkspace := filepath.Join(tmpDir, "dst", "workspace")
err = os.MkdirAll(dstWorkspace, 0o755) err := os.MkdirAll(srcWorkspace, 0o755)
require.NoError(t, err) require.NoError(t, err)
err = os.WriteFile(filepath.Join(srcWorkspace, "file1.txt"), []byte("source"), 0o644) err = os.MkdirAll(dstWorkspace, 0o755)
require.NoError(t, err) require.NoError(t, err)
err = os.WriteFile(filepath.Join(dstWorkspace, "file1.txt"), []byte("existing"), 0o644) err = os.WriteFile(filepath.Join(srcWorkspace, "file1.txt"), []byte("source"), 0o644)
require.NoError(t, err) require.NoError(t, err)
actions, err := PlanWorkspaceMigration( err = os.WriteFile(filepath.Join(dstWorkspace, "file1.txt"), []byte("existing"), 0o644)
srcWorkspace, require.NoError(t, err)
dstWorkspace,
[]string{"file1.txt"},
[]string{},
false,
)
require.NoError(t, err)
require.GreaterOrEqual(t, len(actions), 1) actions, err := PlanWorkspaceMigration(
assert.Equal(t, ActionBackup, actions[0].Type) srcWorkspace,
} dstWorkspace,
[]string{"file1.txt"},
[]string{},
tt.force,
)
require.NoError(t, err)
func TestPlanWorkspaceMigrationForce(t *testing.T) { require.GreaterOrEqual(t, len(actions), 1)
tmpDir := t.TempDir() assert.Equal(t, tt.wantActionType, actions[0].Type)
srcWorkspace := filepath.Join(tmpDir, "src", "workspace") })
dstWorkspace := filepath.Join(tmpDir, "dst", "workspace") }
err := os.MkdirAll(srcWorkspace, 0o755)
require.NoError(t, err)
err = os.MkdirAll(dstWorkspace, 0o755)
require.NoError(t, err)
err = os.WriteFile(filepath.Join(srcWorkspace, "file1.txt"), []byte("source"), 0o644)
require.NoError(t, err)
err = os.WriteFile(filepath.Join(dstWorkspace, "file1.txt"), []byte("existing"), 0o644)
require.NoError(t, err)
actions, err := PlanWorkspaceMigration(
srcWorkspace,
dstWorkspace,
[]string{"file1.txt"},
[]string{},
true,
)
require.NoError(t, err)
require.GreaterOrEqual(t, len(actions), 1)
assert.Equal(t, ActionCopy, actions[0].Type)
} }
func TestPlanWorkspaceMigrationNonExistentSource(t *testing.T) { func TestPlanWorkspaceMigrationNonExistentSource(t *testing.T) {

View file

@ -212,14 +212,14 @@ func translateTools(tools []ToolDefinition) []anthropic.ToolUnionParam {
} }
func parseResponse(resp *anthropic.Message) *LLMResponse { func parseResponse(resp *anthropic.Message) *LLMResponse {
var content string var content strings.Builder
var toolCalls []ToolCall var toolCalls []ToolCall
for _, block := range resp.Content { for _, block := range resp.Content {
switch block.Type { switch block.Type {
case "text": case "text":
tb := block.AsText() tb := block.AsText()
content += tb.Text content.WriteString(tb.Text)
case "tool_use": case "tool_use":
tu := block.AsToolUse() tu := block.AsToolUse()
var args map[string]any var args map[string]any
@ -246,7 +246,7 @@ func parseResponse(resp *anthropic.Message) *LLMResponse {
} }
return &LLMResponse{ return &LLMResponse{
Content: content, Content: content.String(),
ToolCalls: toolCalls, ToolCalls: toolCalls,
FinishReason: finishReason, FinishReason: finishReason,
Usage: &UsageInfo{ Usage: &UsageInfo{
@ -264,8 +264,8 @@ func normalizeBaseURL(apiBase string) string {
} }
base = strings.TrimRight(base, "/") base = strings.TrimRight(base, "/")
if strings.HasSuffix(base, "/v1") { if before, ok := strings.CutSuffix(base, "/v1"); ok {
base = strings.TrimSuffix(base, "/v1") base = before
} }
if base == "" { if base == "" {
return defaultBaseURL return defaultBaseURL

View file

@ -100,44 +100,12 @@ func (p *ClaudeCliProvider) buildSystemPrompt(messages []Message, tools []ToolDe
} }
if len(tools) > 0 { if len(tools) > 0 {
parts = append(parts, p.buildToolsPrompt(tools)) parts = append(parts, buildCLIToolsPrompt(tools))
} }
return strings.Join(parts, "\n\n") return strings.Join(parts, "\n\n")
} }
// buildToolsPrompt creates the tool definitions section for the system prompt.
func (p *ClaudeCliProvider) buildToolsPrompt(tools []ToolDefinition) string {
var sb strings.Builder
sb.WriteString("## Available Tools\n\n")
sb.WriteString("When you need to use a tool, respond with ONLY a JSON object:\n\n")
sb.WriteString("```json\n")
sb.WriteString(
`{"tool_calls":[{"id":"call_xxx","type":"function","function":{"name":"tool_name","arguments":"{...}"}}]}`,
)
sb.WriteString("\n```\n\n")
sb.WriteString("CRITICAL: The 'arguments' field MUST be a JSON-encoded STRING.\n\n")
sb.WriteString("### Tool Definitions:\n\n")
for _, tool := range tools {
if tool.Type != "function" {
continue
}
sb.WriteString(fmt.Sprintf("#### %s\n", tool.Function.Name))
if tool.Function.Description != "" {
sb.WriteString(fmt.Sprintf("Description: %s\n", tool.Function.Description))
}
if len(tool.Function.Parameters) > 0 {
paramsJSON, _ := json.Marshal(tool.Function.Parameters)
sb.WriteString(fmt.Sprintf("Parameters:\n```json\n%s\n```\n", string(paramsJSON)))
}
sb.WriteString("\n")
}
return sb.String()
}
// parseClaudeCliResponse parses the JSON output from the claude CLI. // parseClaudeCliResponse parses the JSON output from the claude CLI.
func (p *ClaudeCliProvider) parseClaudeCliResponse(output string) (*LLMResponse, error) { func (p *ClaudeCliProvider) parseClaudeCliResponse(output string) (*LLMResponse, error) {
var resp claudeCliJSONResponse var resp claudeCliJSONResponse

View file

@ -660,12 +660,11 @@ func TestBuildSystemPrompt_ToolsOnlyNoSystem(t *testing.T) {
// --- buildToolsPrompt tests --- // --- buildToolsPrompt tests ---
func TestBuildToolsPrompt_SkipsNonFunction(t *testing.T) { func TestBuildToolsPrompt_SkipsNonFunction(t *testing.T) {
p := NewClaudeCliProvider("/workspace")
tools := []ToolDefinition{ tools := []ToolDefinition{
{Type: "other", Function: ToolFunctionDefinition{Name: "skip_me"}}, {Type: "other", Function: ToolFunctionDefinition{Name: "skip_me"}},
{Type: "function", Function: ToolFunctionDefinition{Name: "include_me", Description: "Included"}}, {Type: "function", Function: ToolFunctionDefinition{Name: "include_me", Description: "Included"}},
} }
got := p.buildToolsPrompt(tools) got := buildCLIToolsPrompt(tools)
if strings.Contains(got, "skip_me") { if strings.Contains(got, "skip_me") {
t.Error("buildToolsPrompt() should skip non-function tools") t.Error("buildToolsPrompt() should skip non-function tools")
} }
@ -675,11 +674,10 @@ func TestBuildToolsPrompt_SkipsNonFunction(t *testing.T) {
} }
func TestBuildToolsPrompt_NoDescription(t *testing.T) { func TestBuildToolsPrompt_NoDescription(t *testing.T) {
p := NewClaudeCliProvider("/workspace")
tools := []ToolDefinition{ tools := []ToolDefinition{
{Type: "function", Function: ToolFunctionDefinition{Name: "bare_tool"}}, {Type: "function", Function: ToolFunctionDefinition{Name: "bare_tool"}},
} }
got := p.buildToolsPrompt(tools) got := buildCLIToolsPrompt(tools)
if !strings.Contains(got, "bare_tool") { if !strings.Contains(got, "bare_tool") {
t.Error("should include tool name") t.Error("should include tool name")
} }
@ -689,14 +687,13 @@ func TestBuildToolsPrompt_NoDescription(t *testing.T) {
} }
func TestBuildToolsPrompt_NoParameters(t *testing.T) { func TestBuildToolsPrompt_NoParameters(t *testing.T) {
p := NewClaudeCliProvider("/workspace")
tools := []ToolDefinition{ tools := []ToolDefinition{
{Type: "function", Function: ToolFunctionDefinition{ {Type: "function", Function: ToolFunctionDefinition{
Name: "no_params_tool", Name: "no_params_tool",
Description: "A tool with no parameters", Description: "A tool with no parameters",
}}, }},
} }
got := p.buildToolsPrompt(tools) got := buildCLIToolsPrompt(tools)
if strings.Contains(got, "Parameters:") { if strings.Contains(got, "Parameters:") {
t.Error("should not include Parameters: section when nil") t.Error("should not include Parameters: section when nil")
} }

View file

@ -115,7 +115,7 @@ func (p *CodexCliProvider) buildPrompt(messages []Message, tools []ToolDefinitio
} }
if len(tools) > 0 { if len(tools) > 0 {
sb.WriteString(p.buildToolsPrompt(tools)) sb.WriteString(buildCLIToolsPrompt(tools))
sb.WriteString("\n\n") sb.WriteString("\n\n")
} }
@ -128,38 +128,6 @@ func (p *CodexCliProvider) buildPrompt(messages []Message, tools []ToolDefinitio
return sb.String() return sb.String()
} }
// buildToolsPrompt creates a tool definitions section for the prompt.
func (p *CodexCliProvider) buildToolsPrompt(tools []ToolDefinition) string {
var sb strings.Builder
sb.WriteString("## Available Tools\n\n")
sb.WriteString("When you need to use a tool, respond with ONLY a JSON object:\n\n")
sb.WriteString("```json\n")
sb.WriteString(
`{"tool_calls":[{"id":"call_xxx","type":"function","function":{"name":"tool_name","arguments":"{...}"}}]}`,
)
sb.WriteString("\n```\n\n")
sb.WriteString("CRITICAL: The 'arguments' field MUST be a JSON-encoded STRING.\n\n")
sb.WriteString("### Tool Definitions:\n\n")
for _, tool := range tools {
if tool.Type != "function" {
continue
}
sb.WriteString(fmt.Sprintf("#### %s\n", tool.Function.Name))
if tool.Function.Description != "" {
sb.WriteString(fmt.Sprintf("Description: %s\n", tool.Function.Description))
}
if len(tool.Function.Parameters) > 0 {
paramsJSON, _ := json.Marshal(tool.Function.Parameters)
sb.WriteString(fmt.Sprintf("Parameters:\n```json\n%s\n```\n", string(paramsJSON)))
}
sb.WriteString("\n")
}
return sb.String()
}
// codexEvent represents a single JSONL event from `codex exec --json`. // codexEvent represents a single JSONL event from `codex exec --json`.
type codexEvent struct { type codexEvent struct {
Type string `json:"type"` Type string `json:"type"`

View file

@ -163,8 +163,8 @@ func resolveCodexModel(model string) (string, string) {
return codexDefaultModel, "empty model" return codexDefaultModel, "empty model"
} }
if strings.HasPrefix(m, "openai/") { if after, ok := strings.CutPrefix(m, "openai/"); ok {
m = strings.TrimPrefix(m, "openai/") m = after
} else if strings.Contains(m, "/") { } else if strings.Contains(m, "/") {
return codexDefaultModel, "non-openai model namespace" return codexDefaultModel, "non-openai model namespace"
} }

View file

@ -138,7 +138,7 @@ func TestCooldown_FailureWindowReset(t *testing.T) {
ct, current := newTestTracker(now) ct, current := newTestTracker(now)
// 4 errors → 1h cooldown // 4 errors → 1h cooldown
for i := 0; i < 4; i++ { for range 4 {
ct.MarkFailure("openai", FailoverRateLimit) ct.MarkFailure("openai", FailoverRateLimit)
*current = current.Add(2 * time.Second) // small advance between errors *current = current.Add(2 * time.Second) // small advance between errors
} }
@ -230,7 +230,7 @@ func TestCooldown_ConcurrentAccess(t *testing.T) {
ct := NewCooldownTracker() ct := NewCooldownTracker()
var wg sync.WaitGroup var wg sync.WaitGroup
for i := 0; i < 100; i++ { for range 100 {
wg.Add(3) wg.Add(3)
go func() { go func() {
defer wg.Done() defer wg.Done()

View file

@ -6,6 +6,13 @@ import (
"strings" "strings"
) )
// Common patterns in Go HTTP error messages
var httpStatusPatterns = []*regexp.Regexp{
regexp.MustCompile(`status[:\s]+(\d{3})`),
regexp.MustCompile(`http[/\s]+\d*\.?\d*\s+(\d{3})`),
regexp.MustCompile(`\b([3-5]\d{2})\b`),
}
// errorPattern defines a single pattern (string or regex) for error classification. // errorPattern defines a single pattern (string or regex) for error classification.
type errorPattern struct { type errorPattern struct {
substring string substring string
@ -198,20 +205,13 @@ func classifyByMessage(msg string) FailoverReason {
} }
// extractHTTPStatus extracts an HTTP status code from an error message. // extractHTTPStatus extracts an HTTP status code from an error message.
// Looks for patterns like "status: 429", "status 429", "HTTP 429", or standalone "429". // Looks for patterns like "status: 429", "status 429", "http/1.1 429", "http 429", or standalone "429".
func extractHTTPStatus(msg string) int { func extractHTTPStatus(msg string) int {
// Common patterns in Go HTTP error messages for _, p := range httpStatusPatterns {
patterns := []*regexp.Regexp{
regexp.MustCompile(`status[:\s]+(\d{3})`),
regexp.MustCompile(`HTTP[/\s]+\d*\.?\d*\s+(\d{3})`),
}
for _, p := range patterns {
if m := p.FindStringSubmatch(msg); len(m) > 1 { if m := p.FindStringSubmatch(msg); len(m) > 1 {
return parseDigits(m[1]) return parseDigits(m[1])
} }
} }
return 0 return 0
} }

View file

@ -305,7 +305,8 @@ func TestExtractHTTPStatus(t *testing.T) {
}{ }{
{"status: 429 rate limited", 429}, {"status: 429 rate limited", 429},
{"status 401 unauthorized", 401}, {"status 401 unauthorized", 401},
{"HTTP/1.1 502 Bad Gateway", 502}, {"http/1.1 502 bad gateway", 502},
{"error 429", 429},
{"no status code here", 0}, {"no status code here", 0},
{"random number 12345", 0}, {"random number 12345", 0},
} }

View file

@ -102,6 +102,15 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
sel.apiBase = "https://openrouter.ai/api/v1" sel.apiBase = "https://openrouter.ai/api/v1"
} }
} }
case "litellm":
if cfg.Providers.LiteLLM.APIKey != "" || cfg.Providers.LiteLLM.APIBase != "" {
sel.apiKey = cfg.Providers.LiteLLM.APIKey
sel.apiBase = cfg.Providers.LiteLLM.APIBase
sel.proxy = cfg.Providers.LiteLLM.Proxy
if sel.apiBase == "" {
sel.apiBase = "http://localhost:4000/v1"
}
}
case "zhipu", "glm": case "zhipu", "glm":
if cfg.Providers.Zhipu.APIKey != "" { if cfg.Providers.Zhipu.APIKey != "" {
sel.apiKey = cfg.Providers.Zhipu.APIKey sel.apiKey = cfg.Providers.Zhipu.APIKey

View file

@ -53,7 +53,7 @@ func ExtractProtocol(model string) (protocol, modelID string) {
// CreateProviderFromConfig creates a provider based on the ModelConfig. // CreateProviderFromConfig creates a provider based on the ModelConfig.
// It uses the protocol prefix in the Model field to determine which provider to create. // It uses the protocol prefix in the Model field to determine which provider to create.
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot // Supported protocols: openai, litellm, anthropic, antigravity, claude-cli, codex-cli, github-copilot
// Returns the provider, the model ID (without protocol prefix), and any error. // Returns the provider, the model ID (without protocol prefix), and any error.
func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) { func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) {
if cfg == nil { if cfg == nil {
@ -92,7 +92,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout, cfg.RequestTimeout,
), modelID, nil ), modelID, nil
case "openrouter", "groq", "zhipu", "gemini", "nvidia", case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras", "ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
"volcengine", "vllm", "qwen", "mistral": "volcengine", "vllm", "qwen", "mistral":
// All other OpenAI-compatible HTTP providers // All other OpenAI-compatible HTTP providers
@ -180,6 +180,8 @@ func getDefaultAPIBase(protocol string) string {
return "https://api.openai.com/v1" return "https://api.openai.com/v1"
case "openrouter": case "openrouter":
return "https://openrouter.ai/api/v1" return "https://openrouter.ai/api/v1"
case "litellm":
return "http://localhost:4000/v1"
case "groq": case "groq":
return "https://api.groq.com/openai/v1" return "https://api.groq.com/openai/v1"
case "zhipu": case "zhipu":

View file

@ -135,6 +135,32 @@ func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
} }
} }
func TestGetDefaultAPIBase_LiteLLM(t *testing.T) {
if got := getDefaultAPIBase("litellm"); got != "http://localhost:4000/v1" {
t.Fatalf("getDefaultAPIBase(%q) = %q, want %q", "litellm", got, "http://localhost:4000/v1")
}
}
func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
cfg := &config.ModelConfig{
ModelName: "test-litellm",
Model: "litellm/my-proxy-alias",
APIKey: "test-key",
APIBase: "http://localhost:4000/v1",
}
provider, modelID, err := CreateProviderFromConfig(cfg)
if err != nil {
t.Fatalf("CreateProviderFromConfig() error = %v", err)
}
if provider == nil {
t.Fatal("CreateProviderFromConfig() returned nil provider")
}
if modelID != "my-proxy-alias" {
t.Errorf("modelID = %q, want %q", modelID, "my-proxy-alias")
}
}
func TestCreateProviderFromConfig_Anthropic(t *testing.T) { func TestCreateProviderFromConfig_Anthropic(t *testing.T) {
cfg := &config.ModelConfig{ cfg := &config.ModelConfig{
ModelName: "test-anthropic", ModelName: "test-anthropic",

View file

@ -17,6 +17,27 @@ func TestResolveProviderSelection(t *testing.T) {
wantProxy string wantProxy string
wantErrSubstr string wantErrSubstr string
}{ }{
{
name: "explicit litellm provider uses configured base",
setup: func(cfg *config.Config) {
cfg.Agents.Defaults.Provider = "litellm"
cfg.Providers.LiteLLM.APIKey = "litellm-key"
cfg.Providers.LiteLLM.APIBase = "http://localhost:4000/v1"
cfg.Providers.LiteLLM.Proxy = "http://127.0.0.1:7890"
},
wantType: providerTypeHTTPCompat,
wantAPIBase: "http://localhost:4000/v1",
wantProxy: "http://127.0.0.1:7890",
},
{
name: "explicit litellm provider defaults base when only key is configured",
setup: func(cfg *config.Config) {
cfg.Agents.Defaults.Provider = "litellm"
cfg.Providers.LiteLLM.APIKey = "litellm-key"
},
wantType: providerTypeHTTPCompat,
wantAPIBase: "http://localhost:4000/v1",
},
{ {
name: "explicit claude-cli provider routes to cli provider type", name: "explicit claude-cli provider routes to cli provider type",
setup: func(cfg *config.Config) { setup: func(cfg *config.Config) {

View file

@ -26,8 +26,9 @@ func NewGitHubCopilotProvider(uri string, connectMode string, model string) (*Gi
switch connectMode { switch connectMode {
case "stdio": case "stdio":
// TODO: // TODO: Implement stdio mode for GitHub Copilot provider
return nil, fmt.Errorf("stdio mode not implemented") // See https://github.com/github/copilot-sdk/blob/main/docs/getting-started.md for details
return nil, fmt.Errorf("stdio mode not implemented for GitHub Copilot provider; please use 'grpc' mode instead")
case "grpc": case "grpc":
client := copilot.NewClient(&copilot.ClientOptions{ client := copilot.NewClient(&copilot.ClientOptions{
CLIUrl: uri, CLIUrl: uri,
@ -100,9 +101,12 @@ func (p *GitHubCopilotProvider) Chat(
return nil, fmt.Errorf("provider closed") return nil, fmt.Errorf("provider closed")
} }
resp, _ := session.SendAndWait(ctx, copilot.MessageOptions{ resp, err := session.SendAndWait(ctx, copilot.MessageOptions{
Prompt: string(fullcontent), Prompt: string(fullcontent),
}) })
if err != nil {
return nil, fmt.Errorf("failed to send message to copilot: %w", err)
}
if resp == nil { if resp == nil {
return nil, fmt.Errorf("empty response from copilot") return nil, fmt.Errorf("empty response from copilot")

View file

@ -116,7 +116,7 @@ func (p *Provider) Chat(
requestBody := map[string]any{ requestBody := map[string]any{
"model": model, "model": model,
"messages": stripSystemParts(messages), "messages": serializeMessages(messages),
} }
if len(tools) > 0 { if len(tools) > 0 {
@ -289,31 +289,69 @@ func parseResponse(body []byte) (*LLMResponse, error) {
// It mirrors protocoltypes.Message but omits SystemParts, which is an // It mirrors protocoltypes.Message but omits SystemParts, which is an
// internal field that would be unknown to third-party endpoints. // internal field that would be unknown to third-party endpoints.
type openaiMessage struct { type openaiMessage struct {
Role string `json:"role"` Role string `json:"role"`
Content string `json:"content"` Content string `json:"content"`
ToolCalls []ToolCall `json:"tool_calls,omitempty"` ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"` ToolCalls []ToolCall `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
} }
// stripSystemParts converts []Message to []openaiMessage, dropping the // serializeMessages converts internal Message structs to the OpenAI wire format.
// SystemParts field so it doesn't leak into the JSON payload sent to // - Strips SystemParts (unknown to third-party endpoints)
// OpenAI-compatible APIs (some strict endpoints reject unknown fields). // - Converts messages with Media to multipart content format (text + image_url parts)
func stripSystemParts(messages []Message) []openaiMessage { // - Preserves ToolCallID, ToolCalls, and ReasoningContent for all messages
out := make([]openaiMessage, len(messages)) func serializeMessages(messages []Message) []any {
for i, m := range messages { out := make([]any, 0, len(messages))
out[i] = openaiMessage{ for _, m := range messages {
Role: m.Role, if len(m.Media) == 0 {
Content: m.Content, out = append(out, openaiMessage{
ToolCalls: m.ToolCalls, Role: m.Role,
ToolCallID: m.ToolCallID, Content: m.Content,
ReasoningContent: m.ReasoningContent,
ToolCalls: m.ToolCalls,
ToolCallID: m.ToolCallID,
})
continue
} }
// Multipart content format for messages with media
parts := make([]map[string]any, 0, 1+len(m.Media))
if m.Content != "" {
parts = append(parts, map[string]any{
"type": "text",
"text": m.Content,
})
}
for _, mediaURL := range m.Media {
parts = append(parts, map[string]any{
"type": "image_url",
"image_url": map[string]any{
"url": mediaURL,
},
})
}
msg := map[string]any{
"role": m.Role,
"content": parts,
}
if m.ToolCallID != "" {
msg["tool_call_id"] = m.ToolCallID
}
if len(m.ToolCalls) > 0 {
msg["tool_calls"] = m.ToolCalls
}
if m.ReasoningContent != "" {
msg["reasoning_content"] = m.ReasoningContent
}
out = append(out, msg)
} }
return out return out
} }
func normalizeModel(model, apiBase string) string { func normalizeModel(model, apiBase string) string {
idx := strings.Index(model, "/") before, after, ok := strings.Cut(model, "/")
if idx == -1 { if !ok {
return model return model
} }
@ -321,10 +359,10 @@ func normalizeModel(model, apiBase string) string {
return model return model
} }
prefix := strings.ToLower(model[:idx]) prefix := strings.ToLower(before)
switch prefix { switch prefix {
case "moonshot", "nvidia", "groq", "ollama", "deepseek", "google", "openrouter", "zhipu", "mistral": case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google", "openrouter", "zhipu", "mistral":
return model[idx+1:] return after
default: default:
return model return model
} }

View file

@ -5,8 +5,11 @@ import (
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/url" "net/url"
"strings"
"testing" "testing"
"time" "time"
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
) )
func TestProviderChat_UsesMaxCompletionTokensForGLM(t *testing.T) { func TestProviderChat_UsesMaxCompletionTokensForGLM(t *testing.T) {
@ -146,6 +149,56 @@ func TestProviderChat_ParsesReasoningContent(t *testing.T) {
} }
} }
func TestProviderChat_PreservesReasoningContentInHistory(t *testing.T) {
var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
resp := map[string]any{
"choices": []map[string]any{
{
"message": map[string]any{"content": "ok"},
"finish_reason": "stop",
},
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
}))
defer server.Close()
p := NewProvider("key", server.URL, "")
// Simulate a multi-turn conversation where the assistant's previous
// reply included reasoning_content (e.g. from kimi-k2.5).
messages := []Message{
{Role: "user", Content: "What is 1+1?"},
{Role: "assistant", Content: "2", ReasoningContent: "Let me think... 1+1=2"},
{Role: "user", Content: "What about 2+2?"},
}
_, err := p.Chat(t.Context(), messages, nil, "kimi-k2.5", nil)
if err != nil {
t.Fatalf("Chat() error = %v", err)
}
// Verify reasoning_content is preserved in the serialized request.
reqMessages, ok := requestBody["messages"].([]any)
if !ok {
t.Fatalf("messages is not []any: %T", requestBody["messages"])
}
assistantMsg, ok := reqMessages[1].(map[string]any)
if !ok {
t.Fatalf("assistant message is not map[string]any: %T", reqMessages[1])
}
if assistantMsg["reasoning_content"] != "Let me think... 1+1=2" {
t.Errorf("reasoning_content not preserved in request, got %v", assistantMsg["reasoning_content"])
}
}
func TestProviderChat_HTTPError(t *testing.T) { func TestProviderChat_HTTPError(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "bad request", http.StatusBadRequest) http.Error(w, "bad request", http.StatusBadRequest)
@ -206,6 +259,11 @@ func TestProviderChat_StripsGroqAndOllamaPrefixes(t *testing.T) {
input string input string
wantModel string wantModel string
}{ }{
{
name: "strips litellm prefix and preserves proxy model name",
input: "litellm/my-proxy-alias",
wantModel: "my-proxy-alias",
},
{ {
name: "strips groq prefix and keeps nested model", name: "strips groq prefix and keeps nested model",
input: "groq/openai/gpt-oss-120b", input: "groq/openai/gpt-oss-120b",
@ -361,3 +419,97 @@ func TestProvider_FunctionalOptionRequestTimeoutNonPositive(t *testing.T) {
t.Fatalf("http timeout = %v, want %v", p.httpClient.Timeout, defaultRequestTimeout) t.Fatalf("http timeout = %v, want %v", p.httpClient.Timeout, defaultRequestTimeout)
} }
} }
func TestSerializeMessages_PlainText(t *testing.T) {
messages := []protocoltypes.Message{
{Role: "user", Content: "hello"},
{Role: "assistant", Content: "hi", ReasoningContent: "thinking..."},
}
result := serializeMessages(messages)
data, err := json.Marshal(result)
if err != nil {
t.Fatal(err)
}
var msgs []map[string]any
json.Unmarshal(data, &msgs)
if msgs[0]["content"] != "hello" {
t.Fatalf("expected plain string content, got %v", msgs[0]["content"])
}
if msgs[1]["reasoning_content"] != "thinking..." {
t.Fatalf("reasoning_content not preserved, got %v", msgs[1]["reasoning_content"])
}
}
func TestSerializeMessages_WithMedia(t *testing.T) {
messages := []protocoltypes.Message{
{Role: "user", Content: "describe this", Media: []string{"data:image/png;base64,abc123"}},
}
result := serializeMessages(messages)
data, _ := json.Marshal(result)
var msgs []map[string]any
json.Unmarshal(data, &msgs)
content, ok := msgs[0]["content"].([]any)
if !ok {
t.Fatalf("expected array content for media message, got %T", msgs[0]["content"])
}
if len(content) != 2 {
t.Fatalf("expected 2 content parts, got %d", len(content))
}
textPart := content[0].(map[string]any)
if textPart["type"] != "text" || textPart["text"] != "describe this" {
t.Fatalf("text part mismatch: %v", textPart)
}
imgPart := content[1].(map[string]any)
if imgPart["type"] != "image_url" {
t.Fatalf("expected image_url type, got %v", imgPart["type"])
}
imgURL := imgPart["image_url"].(map[string]any)
if imgURL["url"] != "data:image/png;base64,abc123" {
t.Fatalf("image url mismatch: %v", imgURL["url"])
}
}
func TestSerializeMessages_MediaWithToolCallID(t *testing.T) {
messages := []protocoltypes.Message{
{Role: "tool", Content: "image result", Media: []string{"data:image/png;base64,xyz"}, ToolCallID: "call_1"},
}
result := serializeMessages(messages)
data, _ := json.Marshal(result)
var msgs []map[string]any
json.Unmarshal(data, &msgs)
if msgs[0]["tool_call_id"] != "call_1" {
t.Fatalf("tool_call_id not preserved with media, got %v", msgs[0]["tool_call_id"])
}
// Content should be multipart array
if _, ok := msgs[0]["content"].([]any); !ok {
t.Fatalf("expected array content, got %T", msgs[0]["content"])
}
}
func TestSerializeMessages_StripsSystemParts(t *testing.T) {
messages := []protocoltypes.Message{
{
Role: "system",
Content: "you are helpful",
SystemParts: []protocoltypes.ContentBlock{
{Type: "text", Text: "you are helpful"},
},
},
}
result := serializeMessages(messages)
data, _ := json.Marshal(result)
raw := string(data)
if strings.Contains(raw, "system_parts") {
t.Fatal("system_parts should not appear in serialized output")
}
}

View file

@ -65,6 +65,7 @@ type ContentBlock struct {
type Message struct { type Message struct {
Role string `json:"role"` Role string `json:"role"`
Content string `json:"content"` Content string `json:"content"`
Media []string `json:"media,omitempty"`
ReasoningContent string `json:"reasoning_content,omitempty"` ReasoningContent string `json:"reasoning_content,omitempty"`
SystemParts []ContentBlock `json:"system_parts,omitempty"` // structured system blocks for cache-aware adapters SystemParts []ContentBlock `json:"system_parts,omitempty"` // structured system blocks for cache-aware adapters
ToolCalls []ToolCall `json:"tool_calls,omitempty"` ToolCalls []ToolCall `json:"tool_calls,omitempty"`

View file

@ -5,7 +5,43 @@
package providers package providers
import "encoding/json" import (
"encoding/json"
"fmt"
"strings"
)
// buildCLIToolsPrompt creates the tool definitions section for a CLI provider system prompt.
func buildCLIToolsPrompt(tools []ToolDefinition) string {
var sb strings.Builder
sb.WriteString("## Available Tools\n\n")
sb.WriteString("When you need to use a tool, respond with ONLY a JSON object:\n\n")
sb.WriteString("```json\n")
sb.WriteString(
`{"tool_calls":[{"id":"call_xxx","type":"function","function":{"name":"tool_name","arguments":"{...}"}}]}`,
)
sb.WriteString("\n```\n\n")
sb.WriteString("CRITICAL: The 'arguments' field MUST be a JSON-encoded STRING.\n\n")
sb.WriteString("### Tool Definitions:\n\n")
for _, tool := range tools {
if tool.Type != "function" {
continue
}
sb.WriteString(fmt.Sprintf("#### %s\n", tool.Function.Name))
if tool.Function.Description != "" {
sb.WriteString(fmt.Sprintf("Description: %s\n", tool.Function.Description))
}
if len(tool.Function.Parameters) > 0 {
paramsJSON, _ := json.Marshal(tool.Function.Parameters)
sb.WriteString(fmt.Sprintf("Parameters:\n```json\n%s\n```\n", string(paramsJSON)))
}
sb.WriteString("\n")
}
return sb.String()
}
// NormalizeToolCall normalizes a ToolCall to ensure all fields are properly populated. // NormalizeToolCall normalizes a ToolCall to ensure all fields are properly populated.
// It handles cases where Name/Arguments might be in different locations (top-level vs Function) // It handles cases where Name/Arguments might be in different locations (top-level vs Function)

View file

@ -1,6 +1,9 @@
package routing package routing
import "testing" import (
"strings"
"testing"
)
func TestNormalizeAgentID_Empty(t *testing.T) { func TestNormalizeAgentID_Empty(t *testing.T) {
if got := NormalizeAgentID(""); got != DefaultAgentID { if got := NormalizeAgentID(""); got != DefaultAgentID {
@ -57,11 +60,11 @@ func TestNormalizeAgentID_AllInvalid(t *testing.T) {
} }
func TestNormalizeAgentID_TruncatesAt64(t *testing.T) { func TestNormalizeAgentID_TruncatesAt64(t *testing.T) {
long := "" var long strings.Builder
for i := 0; i < 100; i++ { for range 100 {
long += "a" long.WriteString("a")
} }
got := NormalizeAgentID(long) got := NormalizeAgentID(long.String())
if len(got) > MaxAgentIDLength { if len(got) > MaxAgentIDLength {
t.Errorf("length = %d, want <= %d", len(got), MaxAgentIDLength) t.Errorf("length = %d, want <= %d", len(got), MaxAgentIDLength)
} }

View file

@ -2,7 +2,6 @@ package skills
import ( import (
"context" "context"
"encoding/json"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
@ -18,14 +17,6 @@ type SkillInstaller struct {
workspace string workspace string
} }
type AvailableSkill struct {
Name string `json:"name"`
Repository string `json:"repository"`
Description string `json:"description"`
Author string `json:"author"`
Tags []string `json:"tags"`
}
func NewSkillInstaller(workspace string) *SkillInstaller { func NewSkillInstaller(workspace string) *SkillInstaller {
return &SkillInstaller{ return &SkillInstaller{
workspace: workspace, workspace: workspace,
@ -89,35 +80,3 @@ func (si *SkillInstaller) Uninstall(skillName string) error {
return nil return nil
} }
func (si *SkillInstaller) ListAvailableSkills(ctx context.Context) ([]AvailableSkill, error) {
url := "https://raw.githubusercontent.com/sipeed/picoclaw-skills/main/skills.json"
client := &http.Client{Timeout: 15 * time.Second}
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err)
}
resp, err := utils.DoRequestWithRetry(client, req)
if err != nil {
return nil, fmt.Errorf("failed to fetch skills list: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
return nil, fmt.Errorf("failed to fetch skills list: HTTP %d", resp.StatusCode)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read response: %w", err)
}
var skills []AvailableSkill
if err := json.Unmarshal(body, &skills); err != nil {
return nil, fmt.Errorf("failed to parse skills list: %w", err)
}
return skills, nil
}

View file

@ -64,6 +64,29 @@ type SkillsLoader struct {
builtinSkills string // builtin skills builtinSkills string // builtin skills
} }
// SkillRoots returns all unique skill root directories used by this loader.
// The order follows resolution priority: workspace > global > builtin.
func (sl *SkillsLoader) SkillRoots() []string {
roots := []string{sl.workspaceSkills, sl.globalSkills, sl.builtinSkills}
seen := make(map[string]struct{}, len(roots))
out := make([]string, 0, len(roots))
for _, root := range roots {
trimmed := strings.TrimSpace(root)
if trimmed == "" {
continue
}
clean := filepath.Clean(trimmed)
if _, ok := seen[clean]; ok {
continue
}
seen[clean] = struct{}{}
out = append(out, clean)
}
return out
}
func NewSkillsLoader(workspace string, globalSkills string, builtinSkills string) *SkillsLoader { func NewSkillsLoader(workspace string, globalSkills string, builtinSkills string) *SkillsLoader {
return &SkillsLoader{ return &SkillsLoader{
workspace: workspace, workspace: workspace,
@ -240,7 +263,7 @@ func (sl *SkillsLoader) parseSimpleYAML(content string) map[string]string {
normalized := strings.ReplaceAll(content, "\r\n", "\n") normalized := strings.ReplaceAll(content, "\r\n", "\n")
normalized = strings.ReplaceAll(normalized, "\r", "\n") normalized = strings.ReplaceAll(normalized, "\r", "\n")
for _, line := range strings.Split(normalized, "\n") { for line := range strings.SplitSeq(normalized, "\n") {
line = strings.TrimSpace(line) line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") { if line == "" || strings.HasPrefix(line, "#") {
continue continue

View file

@ -326,3 +326,19 @@ func TestStripFrontmatter(t *testing.T) {
}) })
} }
} }
func TestSkillRootsTrimsWhitespaceAndDedups(t *testing.T) {
tmp := t.TempDir()
workspace := filepath.Join(tmp, "workspace")
global := filepath.Join(tmp, "global")
builtin := filepath.Join(tmp, "builtin")
sl := NewSkillsLoader(workspace, " "+global+" ", "\t"+builtin+"\n")
roots := sl.SkillRoots()
assert.Equal(t, []string{
filepath.Join(workspace, "skills"),
global,
builtin,
}, roots)
}

View file

@ -1,7 +1,7 @@
package skills package skills
import ( import (
"sort" "slices"
"strings" "strings"
"sync" "sync"
"time" "time"
@ -183,7 +183,7 @@ func buildTrigrams(s string) []uint32 {
} }
// Sort and Deduplication // Sort and Deduplication
sort.Slice(trigrams, func(i, j int) bool { return trigrams[i] < trigrams[j] }) slices.Sort(trigrams)
n := 1 n := 1
for i := 1; i < len(trigrams); i++ { for i := 1; i < len(trigrams); i++ {
if trigrams[i] != trigrams[i-1] { if trigrams[i] != trigrams[i-1] {

View file

@ -153,7 +153,7 @@ func TestSearchCacheConcurrency(t *testing.T) {
// Concurrent writes // Concurrent writes
go func() { go func() {
for i := 0; i < 100; i++ { for i := range 100 {
cache.Put("query-write-"+string(rune('a'+i%26)), []SearchResult{{Slug: "x"}}) cache.Put("query-write-"+string(rune('a'+i%26)), []SearchResult{{Slug: "x"}})
} }
done <- struct{}{} done <- struct{}{}
@ -161,7 +161,7 @@ func TestSearchCacheConcurrency(t *testing.T) {
// Concurrent reads // Concurrent reads
go func() { go func() {
for i := 0; i < 100; i++ { for range 100 {
cache.Get("query-write-a") cache.Get("query-write-a")
} }
done <- struct{}{} done <- struct{}{}

View file

@ -40,7 +40,9 @@ func NewManager(workspace string) *Manager {
oldStateFile := filepath.Join(workspace, "state.json") oldStateFile := filepath.Join(workspace, "state.json")
// Create state directory if it doesn't exist // Create state directory if it doesn't exist
os.MkdirAll(stateDir, 0o755) if err := os.MkdirAll(stateDir, 0o755); err != nil {
log.Fatalf("[FATAL] state: failed to create state directory: %v", err)
}
sm := &Manager{ sm := &Manager{
workspace: workspace, workspace: workspace,
@ -54,13 +56,17 @@ func NewManager(workspace string) *Manager {
if data, err := os.ReadFile(oldStateFile); err == nil { if data, err := os.ReadFile(oldStateFile); err == nil {
if err := json.Unmarshal(data, sm.state); err == nil { if err := json.Unmarshal(data, sm.state); err == nil {
// Migrate to new location // Migrate to new location
sm.saveAtomic() if err := sm.saveAtomic(); err != nil {
log.Printf("[WARN] state: failed to save state: %v", err)
}
log.Printf("[INFO] state: migrated state from %s to %s", oldStateFile, stateFile) log.Printf("[INFO] state: migrated state from %s to %s", oldStateFile, stateFile)
} }
} }
} else { } else {
// Load from new location // Load from new location
sm.load() if err := sm.load(); err != nil {
log.Printf("[WARN] state: failed to load state: %v", err)
}
} }
return sm return sm

View file

@ -2,8 +2,10 @@ package state
import ( import (
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"os" "os"
"os/exec"
"path/filepath" "path/filepath"
"testing" "testing"
) )
@ -135,7 +137,7 @@ func TestConcurrentAccess(t *testing.T) {
// Test concurrent writes // Test concurrent writes
done := make(chan bool, 10) done := make(chan bool, 10)
for i := 0; i < 10; i++ { for i := range 10 {
go func(idx int) { go func(idx int) {
channel := fmt.Sprintf("channel-%d", idx) channel := fmt.Sprintf("channel-%d", idx)
sm.SetLastChannel(channel) sm.SetLastChannel(channel)
@ -144,7 +146,7 @@ func TestConcurrentAccess(t *testing.T) {
} }
// Wait for all goroutines to complete // Wait for all goroutines to complete
for i := 0; i < 10; i++ { for range 10 {
<-done <-done
} }
@ -214,3 +216,39 @@ func TestNewManager_EmptyWorkspace(t *testing.T) {
t.Error("Expected zero timestamp for new state") t.Error("Expected zero timestamp for new state")
} }
} }
func TestNewManager_MkdirFailureCrashes(t *testing.T) {
// Since log.Fatalf calls os.Exit(1), we cannot test it normally
// Otherwise, the test suite would stop altogether.
// We use the standard pattern of Go: rerun this test in a subprocess.
if os.Getenv("BE_CRASHER") == "1" {
tmpDir := os.Getenv("CRASH_DIR")
statePath := filepath.Join(tmpDir, "state")
if err := os.WriteFile(statePath, []byte("I'm a file, not a folder"), 0o644); err != nil {
fmt.Printf("setup failed: %v", err)
os.Exit(0)
}
NewManager(tmpDir)
os.Exit(0)
}
tmpDir, err := os.MkdirTemp("", "state-crash-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cmd := exec.Command(os.Args[0], "-test.run=TestNewManager_MkdirFailureCrashes")
cmd.Env = append(os.Environ(), "BE_CRASHER=1", "CRASH_DIR="+tmpDir)
err = cmd.Run()
var e *exec.ExitError
if errors.As(err, &e) && !e.Success() {
return
}
t.Fatalf("The process ended without error, a crash was expected via os.Exit(1). Err: %v", err)
}

View file

@ -3,6 +3,7 @@ package tools
import ( import (
"context" "context"
"fmt" "fmt"
"strings"
"sync" "sync"
"time" "time"
@ -222,7 +223,8 @@ func (t *CronTool) listJobs() *ToolResult {
return SilentResult("No scheduled jobs") return SilentResult("No scheduled jobs")
} }
result := "Scheduled jobs:\n" var result strings.Builder
result.WriteString("Scheduled jobs:\n")
for _, j := range jobs { for _, j := range jobs {
var scheduleInfo string var scheduleInfo string
if j.Schedule.Kind == "every" && j.Schedule.EveryMS != nil { if j.Schedule.Kind == "every" && j.Schedule.EveryMS != nil {
@ -234,10 +236,10 @@ func (t *CronTool) listJobs() *ToolResult {
} else { } else {
scheduleInfo = "unknown" scheduleInfo = "unknown"
} }
result += fmt.Sprintf("- %s (id: %s, %s)\n", j.Name, j.ID, scheduleInfo) result.WriteString(fmt.Sprintf("- %s (id: %s, %s)\n", j.Name, j.ID, scheduleInfo))
} }
return SilentResult(result) return SilentResult(result.String())
} }
func (t *CronTool) removeJob(args map[string]any) *ToolResult { func (t *CronTool) removeJob(args map[string]any) *ToolResult {

View file

@ -5,6 +5,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"io/fs" "io/fs"
"regexp"
"strings" "strings"
) )
@ -15,14 +16,12 @@ type EditFileTool struct {
} }
// NewEditFileTool creates a new EditFileTool with optional directory restriction. // NewEditFileTool creates a new EditFileTool with optional directory restriction.
func NewEditFileTool(workspace string, restrict bool) *EditFileTool { func NewEditFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *EditFileTool {
var fs fileSystem var patterns []*regexp.Regexp
if restrict { if len(allowPaths) > 0 {
fs = &sandboxFs{workspace: workspace} patterns = allowPaths[0]
} else {
fs = &hostFs{}
} }
return &EditFileTool{fs: fs} return &EditFileTool{fs: buildFs(workspace, restrict, patterns)}
} }
func (t *EditFileTool) Name() string { func (t *EditFileTool) Name() string {
@ -80,14 +79,12 @@ type AppendFileTool struct {
fs fileSystem fs fileSystem
} }
func NewAppendFileTool(workspace string, restrict bool) *AppendFileTool { func NewAppendFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *AppendFileTool {
var fs fileSystem var patterns []*regexp.Regexp
if restrict { if len(allowPaths) > 0 {
fs = &sandboxFs{workspace: workspace} patterns = allowPaths[0]
} else {
fs = &hostFs{}
} }
return &AppendFileTool{fs: fs} return &AppendFileTool{fs: buildFs(workspace, restrict, patterns)}
} }
func (t *AppendFileTool) Name() string { func (t *AppendFileTool) Name() string {

View file

@ -6,6 +6,7 @@ import (
"io/fs" "io/fs"
"os" "os"
"path/filepath" "path/filepath"
"regexp"
"strings" "strings"
"time" "time"
@ -87,14 +88,12 @@ type ReadFileTool struct {
fs fileSystem fs fileSystem
} }
func NewReadFileTool(workspace string, restrict bool) *ReadFileTool { func NewReadFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *ReadFileTool {
var fs fileSystem var patterns []*regexp.Regexp
if restrict { if len(allowPaths) > 0 {
fs = &sandboxFs{workspace: workspace} patterns = allowPaths[0]
} else {
fs = &hostFs{}
} }
return &ReadFileTool{fs: fs} return &ReadFileTool{fs: buildFs(workspace, restrict, patterns)}
} }
func (t *ReadFileTool) Name() string { func (t *ReadFileTool) Name() string {
@ -135,14 +134,12 @@ type WriteFileTool struct {
fs fileSystem fs fileSystem
} }
func NewWriteFileTool(workspace string, restrict bool) *WriteFileTool { func NewWriteFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *WriteFileTool {
var fs fileSystem var patterns []*regexp.Regexp
if restrict { if len(allowPaths) > 0 {
fs = &sandboxFs{workspace: workspace} patterns = allowPaths[0]
} else {
fs = &hostFs{}
} }
return &WriteFileTool{fs: fs} return &WriteFileTool{fs: buildFs(workspace, restrict, patterns)}
} }
func (t *WriteFileTool) Name() string { func (t *WriteFileTool) Name() string {
@ -192,14 +189,12 @@ type ListDirTool struct {
fs fileSystem fs fileSystem
} }
func NewListDirTool(workspace string, restrict bool) *ListDirTool { func NewListDirTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *ListDirTool {
var fs fileSystem var patterns []*regexp.Regexp
if restrict { if len(allowPaths) > 0 {
fs = &sandboxFs{workspace: workspace} patterns = allowPaths[0]
} else {
fs = &hostFs{}
} }
return &ListDirTool{fs: fs} return &ListDirTool{fs: buildFs(workspace, restrict, patterns)}
} }
func (t *ListDirTool) Name() string { func (t *ListDirTool) Name() string {
@ -394,6 +389,57 @@ func (r *sandboxFs) ReadDir(path string) ([]os.DirEntry, error) {
return entries, err return entries, err
} }
// whitelistFs wraps a sandboxFs and allows access to specific paths outside
// the workspace when they match any of the provided patterns.
type whitelistFs struct {
sandbox *sandboxFs
host hostFs
patterns []*regexp.Regexp
}
func (w *whitelistFs) matches(path string) bool {
for _, p := range w.patterns {
if p.MatchString(path) {
return true
}
}
return false
}
func (w *whitelistFs) ReadFile(path string) ([]byte, error) {
if w.matches(path) {
return w.host.ReadFile(path)
}
return w.sandbox.ReadFile(path)
}
func (w *whitelistFs) WriteFile(path string, data []byte) error {
if w.matches(path) {
return w.host.WriteFile(path, data)
}
return w.sandbox.WriteFile(path, data)
}
func (w *whitelistFs) ReadDir(path string) ([]os.DirEntry, error) {
if w.matches(path) {
return w.host.ReadDir(path)
}
return w.sandbox.ReadDir(path)
}
// buildFs returns the appropriate fileSystem implementation based on restriction
// settings and optional path whitelist patterns.
func buildFs(workspace string, restrict bool, patterns []*regexp.Regexp) fileSystem {
if !restrict {
return &hostFs{}
}
sandbox := &sandboxFs{workspace: workspace}
if len(patterns) > 0 {
return &whitelistFs{sandbox: sandbox, patterns: patterns}
}
return sandbox
}
// Helper to get a safe relative path for os.Root usage // Helper to get a safe relative path for os.Root usage
func getSafeRelPath(workspace, path string) (string, error) { func getSafeRelPath(workspace, path string) (string, error) {
if workspace == "" { if workspace == "" {

View file

@ -5,6 +5,7 @@ import (
"io" "io"
"os" "os"
"path/filepath" "path/filepath"
"regexp"
"strings" "strings"
"testing" "testing"
@ -486,3 +487,36 @@ func TestRootRW_Write(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
assert.Equal(t, newData, content) assert.Equal(t, newData, content)
} }
// TestWhitelistFs_AllowsMatchingPaths verifies that whitelistFs allows access to
// paths matching the whitelist patterns while blocking non-matching paths.
func TestWhitelistFs_AllowsMatchingPaths(t *testing.T) {
workspace := t.TempDir()
outsideDir := t.TempDir()
outsideFile := filepath.Join(outsideDir, "allowed.txt")
os.WriteFile(outsideFile, []byte("outside content"), 0o644)
// Pattern allows access to the outsideDir.
patterns := []*regexp.Regexp{regexp.MustCompile(`^` + regexp.QuoteMeta(outsideDir))}
tool := NewReadFileTool(workspace, true, patterns)
// Read from whitelisted path should succeed.
result := tool.Execute(context.Background(), map[string]any{"path": outsideFile})
if result.IsError {
t.Errorf("expected whitelisted path to be readable, got: %s", result.ForLLM)
}
if !strings.Contains(result.ForLLM, "outside content") {
t.Errorf("expected file content, got: %s", result.ForLLM)
}
// Read from non-whitelisted path outside workspace should fail.
otherDir := t.TempDir()
otherFile := filepath.Join(otherDir, "blocked.txt")
os.WriteFile(otherFile, []byte("blocked"), 0o644)
result = tool.Execute(context.Background(), map[string]any{"path": otherFile})
if !result.IsError {
t.Errorf("expected non-whitelisted path to be blocked, got: %s", result.ForLLM)
}
}

246
pkg/tools/mcp_tool.go Normal file
View file

@ -0,0 +1,246 @@
package tools
import (
"context"
"encoding/json"
"fmt"
"hash/fnv"
"strings"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// MCPManager defines the interface for MCP manager operations
// This allows for easier testing with mock implementations
type MCPManager interface {
CallTool(
ctx context.Context,
serverName, toolName string,
arguments map[string]any,
) (*mcp.CallToolResult, error)
}
// MCPTool wraps an MCP tool to implement the Tool interface
type MCPTool struct {
manager MCPManager
serverName string
tool *mcp.Tool
}
// NewMCPTool creates a new MCP tool wrapper
func NewMCPTool(manager MCPManager, serverName string, tool *mcp.Tool) *MCPTool {
return &MCPTool{
manager: manager,
serverName: serverName,
tool: tool,
}
}
// sanitizeIdentifierComponent normalizes a string so it can be safely used
// as part of a tool/function identifier for downstream providers.
// It:
// - lowercases the string
// - replaces any character not in [a-z0-9_-] with '_'
// - collapses multiple consecutive '_' into a single '_'
// - trims leading/trailing '_'
// - falls back to "unnamed" if the result is empty
// - truncates overly long components to a reasonable length
func sanitizeIdentifierComponent(s string) string {
const maxLen = 64
s = strings.ToLower(s)
var b strings.Builder
b.Grow(len(s))
prevUnderscore := false
for _, r := range s {
isAllowed := (r >= 'a' && r <= 'z') ||
(r >= '0' && r <= '9') ||
r == '_' || r == '-'
if !isAllowed {
// Normalize any disallowed character to '_'
if !prevUnderscore {
b.WriteRune('_')
prevUnderscore = true
}
continue
}
if r == '_' {
if prevUnderscore {
continue
}
prevUnderscore = true
} else {
prevUnderscore = false
}
b.WriteRune(r)
}
result := strings.Trim(b.String(), "_")
if result == "" {
result = "unnamed"
}
if len(result) > maxLen {
result = result[:maxLen]
}
return result
}
// Name returns the tool name, prefixed with the server name.
// The total length is capped at 64 characters (OpenAI-compatible API limit).
// A short hash of the original (unsanitized) server and tool names is appended
// whenever sanitization is lossy or the name is truncated, ensuring that two
// names which differ only in disallowed characters remain distinct after sanitization.
func (t *MCPTool) Name() string {
// Prefix with server name to avoid conflicts, and sanitize components
sanitizedServer := sanitizeIdentifierComponent(t.serverName)
sanitizedTool := sanitizeIdentifierComponent(t.tool.Name)
full := fmt.Sprintf("mcp_%s_%s", sanitizedServer, sanitizedTool)
// Check if sanitization was lossless (only lowercasing, no char replacement/truncation)
lossless := strings.ToLower(t.serverName) == sanitizedServer &&
strings.ToLower(t.tool.Name) == sanitizedTool
const maxTotal = 64
if lossless && len(full) <= maxTotal {
return full
}
// Sanitization was lossy or name too long: append hash of the ORIGINAL names
// (not the sanitized names) so different originals always yield different hashes.
h := fnv.New32a()
_, _ = h.Write([]byte(t.serverName + "\x00" + t.tool.Name))
suffix := fmt.Sprintf("%08x", h.Sum32()) // 8 chars
base := full
if len(base) > maxTotal-9 {
base = strings.TrimRight(full[:maxTotal-9], "_")
}
return base + "_" + suffix
}
// Description returns the tool description
func (t *MCPTool) Description() string {
desc := t.tool.Description
if desc == "" {
desc = fmt.Sprintf("MCP tool from %s server", t.serverName)
}
// Add server info to description
return fmt.Sprintf("[MCP:%s] %s", t.serverName, desc)
}
// Parameters returns the tool parameters schema
func (t *MCPTool) Parameters() map[string]any {
// The InputSchema is already a JSON Schema object
schema := t.tool.InputSchema
// Handle nil schema
if schema == nil {
return map[string]any{
"type": "object",
"properties": map[string]any{},
"required": []string{},
}
}
// Try direct conversion first (fast path)
if schemaMap, ok := schema.(map[string]any); ok {
return schemaMap
}
// Handle json.RawMessage and []byte - unmarshal directly
var jsonData []byte
if rawMsg, ok := schema.(json.RawMessage); ok {
jsonData = rawMsg
} else if bytes, ok := schema.([]byte); ok {
jsonData = bytes
}
if jsonData != nil {
var result map[string]any
if err := json.Unmarshal(jsonData, &result); err == nil {
return result
}
// Fallback on error
return map[string]any{
"type": "object",
"properties": map[string]any{},
"required": []string{},
}
}
// For other types (structs, etc.), convert via JSON marshal/unmarshal
var err error
jsonData, err = json.Marshal(schema)
if err != nil {
// Fallback to empty schema if marshaling fails
return map[string]any{
"type": "object",
"properties": map[string]any{},
"required": []string{},
}
}
var result map[string]any
if err := json.Unmarshal(jsonData, &result); err != nil {
// Fallback to empty schema if unmarshaling fails
return map[string]any{
"type": "object",
"properties": map[string]any{},
"required": []string{},
}
}
return result
}
// Execute executes the MCP tool
func (t *MCPTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
result, err := t.manager.CallTool(ctx, t.serverName, t.tool.Name, args)
if err != nil {
return ErrorResult(fmt.Sprintf("MCP tool execution failed: %v", err)).WithError(err)
}
if result == nil {
nilErr := fmt.Errorf("MCP tool returned nil result without error")
return ErrorResult("MCP tool execution failed: nil result").WithError(nilErr)
}
// Handle error result from server
if result.IsError {
errMsg := extractContentText(result.Content)
return ErrorResult(fmt.Sprintf("MCP tool returned error: %s", errMsg)).
WithError(fmt.Errorf("MCP tool error: %s", errMsg))
}
// Extract text content from result
output := extractContentText(result.Content)
return &ToolResult{
ForLLM: output,
IsError: false,
}
}
// extractContentText extracts text from MCP content array
func extractContentText(content []mcp.Content) string {
var parts []string
for _, c := range content {
switch v := c.(type) {
case *mcp.TextContent:
parts = append(parts, v.Text)
case *mcp.ImageContent:
// For images, just indicate that an image was returned
parts = append(parts, fmt.Sprintf("[Image: %s]", v.MIMEType))
default:
// For other content types, use string representation
parts = append(parts, fmt.Sprintf("[Content: %T]", v))
}
}
return strings.Join(parts, "\n")
}

492
pkg/tools/mcp_tool_test.go Normal file
View file

@ -0,0 +1,492 @@
package tools
import (
"context"
"fmt"
"strings"
"testing"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// MockMCPManager is a mock implementation of MCPManager interface for testing
type MockMCPManager struct {
callToolFunc func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error)
}
func (m *MockMCPManager) CallTool(
ctx context.Context,
serverName, toolName string,
arguments map[string]any,
) (*mcp.CallToolResult, error) {
if m.callToolFunc != nil {
return m.callToolFunc(ctx, serverName, toolName, arguments)
}
return &mcp.CallToolResult{
Content: []mcp.Content{
&mcp.TextContent{Text: "mock result"},
},
IsError: false,
}, nil
}
// TestNewMCPTool verifies MCP tool creation
func TestNewMCPTool(t *testing.T) {
manager := &MockMCPManager{}
tool := &mcp.Tool{
Name: "test_tool",
Description: "A test tool",
InputSchema: map[string]any{
"type": "object",
"properties": map[string]any{
"input": map[string]any{
"type": "string",
"description": "Test input",
},
},
},
}
mcpTool := NewMCPTool(manager, "test_server", tool)
if mcpTool == nil {
t.Fatal("NewMCPTool should not return nil")
}
// Verify tool properties we can access
if mcpTool.Name() != "mcp_test_server_test_tool" {
t.Errorf("Expected tool name with prefix, got '%s'", mcpTool.Name())
}
}
// TestMCPTool_Name verifies tool name with server prefix
func TestMCPTool_Name(t *testing.T) {
tests := []struct {
name string
serverName string
toolName string
expected string
}{
{
name: "simple name",
serverName: "github",
toolName: "create_issue",
expected: "mcp_github_create_issue",
},
{
name: "filesystem server",
serverName: "filesystem",
toolName: "read_file",
expected: "mcp_filesystem_read_file",
},
{
name: "remote server",
serverName: "remote-api",
toolName: "fetch_data",
expected: "mcp_remote-api_fetch_data",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
manager := &MockMCPManager{}
tool := &mcp.Tool{Name: tt.toolName}
mcpTool := NewMCPTool(manager, tt.serverName, tool)
result := mcpTool.Name()
if result != tt.expected {
t.Errorf("Expected name '%s', got '%s'", tt.expected, result)
}
})
}
}
// TestMCPTool_Description verifies tool description generation
func TestMCPTool_Description(t *testing.T) {
tests := []struct {
name string
serverName string
toolDescription string
expectContains []string
}{
{
name: "with description",
serverName: "github",
toolDescription: "Create a GitHub issue",
expectContains: []string{"[MCP:github]", "Create a GitHub issue"},
},
{
name: "empty description",
serverName: "filesystem",
toolDescription: "",
expectContains: []string{"[MCP:filesystem]", "MCP tool from filesystem server"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
manager := &MockMCPManager{}
tool := &mcp.Tool{
Name: "test_tool",
Description: tt.toolDescription,
}
mcpTool := NewMCPTool(manager, tt.serverName, tool)
result := mcpTool.Description()
for _, expected := range tt.expectContains {
if !strings.Contains(result, expected) {
t.Errorf("Description should contain '%s', got: %s", expected, result)
}
}
})
}
}
// TestMCPTool_Parameters verifies parameter schema conversion
func TestMCPTool_Parameters(t *testing.T) {
tests := []struct {
name string
inputSchema any
expectType string
checkProperty string
expectProperty bool
}{
{
name: "map schema",
inputSchema: map[string]any{
"type": "object",
"properties": map[string]any{
"query": map[string]any{
"type": "string",
"description": "Search query",
},
},
"required": []string{"query"},
},
expectType: "object",
checkProperty: "query",
expectProperty: true,
},
{
name: "nil schema",
inputSchema: nil,
expectType: "object",
expectProperty: false,
},
{
name: "json.RawMessage schema",
inputSchema: []byte(`{
"type": "object",
"properties": {
"repo": {
"type": "string",
"description": "Repository name"
},
"stars": {
"type": "integer",
"description": "Minimum stars"
}
},
"required": ["repo"]
}`),
expectType: "object",
checkProperty: "repo",
expectProperty: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
manager := &MockMCPManager{}
tool := &mcp.Tool{
Name: "test_tool",
InputSchema: tt.inputSchema,
}
mcpTool := NewMCPTool(manager, "test_server", tool)
params := mcpTool.Parameters()
if params == nil {
t.Fatal("Parameters should not be nil")
}
if params["type"] != tt.expectType {
t.Errorf("Expected type '%s', got '%v'", tt.expectType, params["type"])
}
// Check if property exists when expected
if tt.checkProperty != "" {
properties, ok := params["properties"].(map[string]any)
if !ok && tt.expectProperty {
t.Errorf("Expected properties to be a map")
return
}
if ok {
_, hasProperty := properties[tt.checkProperty]
if hasProperty != tt.expectProperty {
t.Errorf("Expected property '%s' existence: %v, got: %v",
tt.checkProperty, tt.expectProperty, hasProperty)
}
}
}
})
}
}
// TestMCPTool_Execute_Success tests successful tool execution
func TestMCPTool_Execute_Success(t *testing.T) {
manager := &MockMCPManager{
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
// Verify correct parameters passed
if serverName != "github" {
t.Errorf("Expected serverName 'github', got '%s'", serverName)
}
if toolName != "search_repos" {
t.Errorf("Expected toolName 'search_repos', got '%s'", toolName)
}
return &mcp.CallToolResult{
Content: []mcp.Content{
&mcp.TextContent{Text: "Found 3 repositories"},
},
IsError: false,
}, nil
},
}
tool := &mcp.Tool{
Name: "search_repos",
Description: "Search GitHub repositories",
}
mcpTool := NewMCPTool(manager, "github", tool)
ctx := context.Background()
args := map[string]any{
"query": "golang mcp",
}
result := mcpTool.Execute(ctx, args)
if result == nil {
t.Fatal("Result should not be nil")
}
if result.IsError {
t.Errorf("Expected no error, got error: %s", result.ForLLM)
}
if result.ForLLM != "Found 3 repositories" {
t.Errorf("Expected 'Found 3 repositories', got '%s'", result.ForLLM)
}
}
// TestMCPTool_Execute_ManagerError tests execution when manager returns error
func TestMCPTool_Execute_ManagerError(t *testing.T) {
manager := &MockMCPManager{
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
return nil, fmt.Errorf("connection failed")
},
}
tool := &mcp.Tool{Name: "test_tool"}
mcpTool := NewMCPTool(manager, "test_server", tool)
ctx := context.Background()
result := mcpTool.Execute(ctx, map[string]any{})
if result == nil {
t.Fatal("Result should not be nil")
}
if !result.IsError {
t.Error("Expected IsError to be true")
}
if !strings.Contains(result.ForLLM, "MCP tool execution failed") {
t.Errorf("Error message should mention execution failure, got: %s", result.ForLLM)
}
if !strings.Contains(result.ForLLM, "connection failed") {
t.Errorf("Error message should include original error, got: %s", result.ForLLM)
}
}
// TestMCPTool_Execute_ServerError tests execution when server returns error
func TestMCPTool_Execute_ServerError(t *testing.T) {
manager := &MockMCPManager{
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
return &mcp.CallToolResult{
Content: []mcp.Content{
&mcp.TextContent{Text: "Invalid API key"},
},
IsError: true,
}, nil
},
}
tool := &mcp.Tool{Name: "test_tool"}
mcpTool := NewMCPTool(manager, "test_server", tool)
ctx := context.Background()
result := mcpTool.Execute(ctx, map[string]any{})
if result == nil {
t.Fatal("Result should not be nil")
}
if !result.IsError {
t.Error("Expected IsError to be true")
}
if !strings.Contains(result.ForLLM, "MCP tool returned error") {
t.Errorf("Error message should mention server error, got: %s", result.ForLLM)
}
if !strings.Contains(result.ForLLM, "Invalid API key") {
t.Errorf("Error message should include server message, got: %s", result.ForLLM)
}
}
// TestMCPTool_Execute_MultipleContent tests execution with multiple content items
func TestMCPTool_Execute_MultipleContent(t *testing.T) {
manager := &MockMCPManager{
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
return &mcp.CallToolResult{
Content: []mcp.Content{
&mcp.TextContent{Text: "First line"},
&mcp.TextContent{Text: "Second line"},
&mcp.TextContent{Text: "Third line"},
},
IsError: false,
}, nil
},
}
tool := &mcp.Tool{Name: "multi_output"}
mcpTool := NewMCPTool(manager, "test_server", tool)
ctx := context.Background()
result := mcpTool.Execute(ctx, map[string]any{})
if result.IsError {
t.Errorf("Expected no error, got: %s", result.ForLLM)
}
expected := "First line\nSecond line\nThird line"
if result.ForLLM != expected {
t.Errorf("Expected '%s', got '%s'", expected, result.ForLLM)
}
}
// TestExtractContentText_TextContent tests text content extraction
func TestExtractContentText_TextContent(t *testing.T) {
content := []mcp.Content{
&mcp.TextContent{Text: "Hello World"},
&mcp.TextContent{Text: "Second message"},
}
result := extractContentText(content)
expected := "Hello World\nSecond message"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
}
// TestExtractContentText_ImageContent tests image content extraction
func TestExtractContentText_ImageContent(t *testing.T) {
content := []mcp.Content{
&mcp.ImageContent{
Data: []byte("base64data"),
MIMEType: "image/png",
},
}
result := extractContentText(content)
if !strings.Contains(result, "[Image:") {
t.Errorf("Expected image indicator, got: %s", result)
}
if !strings.Contains(result, "image/png") {
t.Errorf("Expected MIME type in output, got: %s", result)
}
}
// TestExtractContentText_MixedContent tests mixed content types
func TestExtractContentText_MixedContent(t *testing.T) {
content := []mcp.Content{
&mcp.TextContent{Text: "Description"},
&mcp.ImageContent{
Data: []byte("data"),
MIMEType: "image/jpeg",
},
&mcp.TextContent{Text: "More text"},
}
result := extractContentText(content)
if !strings.Contains(result, "Description") {
t.Errorf("Should contain text content, got: %s", result)
}
if !strings.Contains(result, "[Image:") {
t.Errorf("Should contain image indicator, got: %s", result)
}
if !strings.Contains(result, "More text") {
t.Errorf("Should contain second text, got: %s", result)
}
}
// TestExtractContentText_EmptyContent tests empty content array
func TestExtractContentText_EmptyContent(t *testing.T) {
content := []mcp.Content{}
result := extractContentText(content)
if result != "" {
t.Errorf("Expected empty string for empty content, got: %s", result)
}
}
// TestMCPTool_InterfaceCompliance verifies MCPTool implements Tool interface
func TestMCPTool_InterfaceCompliance(t *testing.T) {
manager := &MockMCPManager{}
tool := &mcp.Tool{Name: "test"}
mcpTool := NewMCPTool(manager, "test_server", tool)
// Verify it implements Tool interface
var _ Tool = mcpTool
}
// TestMCPTool_Parameters_MapSchema tests schema that's already a map
func TestMCPTool_Parameters_MapSchema(t *testing.T) {
manager := &MockMCPManager{}
schema := map[string]any{
"type": "object",
"properties": map[string]any{
"name": map[string]any{
"type": "string",
"description": "The name parameter",
},
},
"required": []string{"name"},
}
tool := &mcp.Tool{
Name: "test_tool",
InputSchema: schema,
}
mcpTool := NewMCPTool(manager, "test_server", tool)
params := mcpTool.Parameters()
// Should return the schema as-is when it's already a map
if params["type"] != "object" {
t.Errorf("Expected type 'object', got '%v'", params["type"])
}
props, ok := params["properties"].(map[string]any)
if !ok {
t.Error("Properties should be a map")
}
nameParam, ok := props["name"].(map[string]any)
if !ok {
t.Error("Name parameter should exist")
}
if nameParam["type"] != "string" {
t.Errorf("Name type should be 'string', got '%v'", nameParam["type"])
}
}

298
pkg/tools/memory.go Normal file
View file

@ -0,0 +1,298 @@
// PicoClaw - Ultra-lightweight personal AI agent
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package tools
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"time"
)
// MemoryStore manages persistent memory for the agent.
// - Long-term memory: memory/MEMORY.md
// - Daily notes: memory/YYYYMM/YYYYMMDD.md
type MemoryStore struct {
workspace string
memoryDir string
memoryFile string
}
// NewMemoryStore creates a new MemoryStore with the given workspace path.
func NewMemoryStore(workspace string) *MemoryStore {
memoryDir := filepath.Join(workspace, "memory")
os.MkdirAll(memoryDir, 0755)
return &MemoryStore{
workspace: workspace,
memoryDir: memoryDir,
memoryFile: filepath.Join(memoryDir, "MEMORY.md"),
}
}
// ReadLongTerm reads the long-term memory file.
func (m *MemoryStore) ReadLongTerm() string {
data, err := os.ReadFile(m.memoryFile)
if err != nil {
return ""
}
return string(data)
}
// WriteLongTerm writes the long-term memory file.
func (m *MemoryStore) WriteLongTerm(content string) error {
return os.WriteFile(m.memoryFile, []byte(content), 0644)
}
// GetTodayFilePath returns the path for today's daily note file.
func (m *MemoryStore) GetTodayFilePath() string {
now := time.Now()
yearMonth := now.Format("200601")
dayFile := now.Format("20060102.md")
return filepath.Join(m.memoryDir, yearMonth, dayFile)
}
// AppendToday appends content to today's daily note.
func (m *MemoryStore) AppendToday(content string) error {
now := time.Now()
yearMonth := now.Format("200601")
monthDir := filepath.Join(m.memoryDir, yearMonth)
os.MkdirAll(monthDir, 0755)
filePath := m.GetTodayFilePath()
var existingContent string
if data, err := os.ReadFile(filePath); err == nil {
existingContent = string(data)
}
// Add timestamp header if file is new
if existingContent == "" {
existingContent = "# " + now.Format("2006-01-02") + "\n\n"
}
updated := existingContent + content + "\n"
return os.WriteFile(filePath, []byte(updated), 0644)
}
// ReadToday reads today's daily note.
func (m *MemoryStore) ReadToday() string {
filePath := m.GetTodayFilePath()
data, err := os.ReadFile(filePath)
if err != nil {
return ""
}
return string(data)
}
// GetWorkspace returns the workspace path.
func (m *MemoryStore) GetWorkspace() string {
return m.workspace
}
// UpdateMemoryTool updates long-term memory or daily notes.
type UpdateMemoryTool struct {
memory *MemoryStore
}
// NewUpdateMemoryTool creates a new UpdateMemoryTool.
func NewUpdateMemoryTool(memory *MemoryStore) *UpdateMemoryTool {
return &UpdateMemoryTool{
memory: memory,
}
}
// Name returns the tool name.
func (t *UpdateMemoryTool) Name() string {
return "update_memory"
}
// Description returns the tool description.
func (t *UpdateMemoryTool) Description() string {
return `Save important information to memory. Use this when:
- User shares personal info (name, job, location, relationships)
- User expresses preferences (language, timezone, habits, likes/dislikes)
- Important events, deadlines, or plans are mentioned
- Task completions worth remembering
Memory types:
- long_term: Permanent storage (user info, preferences, important notes)
- daily_note: Temporary note for today's activities`
}
// Parameters returns the tool parameters.
func (t *UpdateMemoryTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{
"memory_type": map[string]any{
"type": "string",
"enum": []string{"long_term", "daily_note"},
"description": "Type of memory: 'long_term' for permanent storage, 'daily_note' for today's temporary note",
},
"content": map[string]any{
"type": "string",
"description": "The content to remember (concise and clear)",
},
"category": map[string]any{
"type": "string",
"enum": []string{"user_info", "preference", "important_note", "task", "configuration"},
"description": "Category for long_term memory (ignored for daily_note)",
},
},
"required": []string{"memory_type", "content"},
}
}
// Execute executes the tool.
func (t *UpdateMemoryTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
memoryType, ok := args["memory_type"].(string)
if !ok {
return ErrorResult("memory_type is required and must be 'long_term' or 'daily_note'").
WithError(fmt.Errorf("invalid memory_type"))
}
content, ok := args["content"].(string)
if !ok {
return ErrorResult("content is required and must be a string").
WithError(fmt.Errorf("invalid content"))
}
// Validate content length
if strings.TrimSpace(content) == "" {
return ErrorResult("content cannot be empty").
WithError(fmt.Errorf("empty content"))
}
switch memoryType {
case "long_term":
return t.updateLongTerm(content, args)
case "daily_note":
return t.updateDailyNote(content)
default:
return ErrorResult("memory_type must be 'long_term' or 'daily_note'").
WithError(fmt.Errorf("invalid memory_type: %s", memoryType))
}
}
// updateLongTerm updates the long-term memory file (MEMORY.md).
func (t *UpdateMemoryTool) updateLongTerm(content string, args map[string]any) *ToolResult {
category, _ := args["category"].(string)
if category == "" {
category = "important_note" // default category
}
// Read current memory
currentMemory := t.memory.ReadLongTerm()
// Append to the appropriate section
updatedMemory := t.appendToSection(currentMemory, category, content)
// Write back atomically
err := t.memory.WriteLongTerm(updatedMemory)
if err != nil {
return ErrorResult("Failed to write long-term memory: " + err.Error()).
WithError(err)
}
return &ToolResult{
ForLLM: fmt.Sprintf("Successfully saved to long-term memory under '%s' category.", category),
ForUser: fmt.Sprintf("✅ 已保存到长期记忆 (%s)", category),
Silent: false,
IsError: false,
Async: false,
}
}
// updateDailyNote appends content to today's daily note.
func (t *UpdateMemoryTool) updateDailyNote(content string) *ToolResult {
// Format with timestamp
timestamp := time.Now().Format("15:04")
formattedContent := fmt.Sprintf("- [%s] %s", timestamp, content)
err := t.memory.AppendToday(formattedContent)
if err != nil {
return ErrorResult("Failed to write daily note: " + err.Error()).
WithError(err)
}
return &ToolResult{
ForLLM: "Successfully added to today's daily note.",
ForUser: "✅ 已添加到今日笔记",
Silent: false,
IsError: false,
Async: false,
}
}
// appendToSection appends content to a specific section in MEMORY.md.
func (t *UpdateMemoryTool) appendToSection(currentMemory, category, content string) string {
// Define section headers
sectionHeaders := map[string]string{
"user_info": "## User Information",
"preference": "## Preferences",
"important_note": "## Important Notes",
"task": "## Tasks & Activities",
"configuration": "## Configuration",
}
header, exists := sectionHeaders[category]
if !exists {
header = "## Important Notes"
}
// Check if section exists
if strings.Contains(currentMemory, header) {
// Append to existing section
return t.appendAfterHeader(currentMemory, header, content)
}
// Create section if it doesn't exist
return t.createNewSection(currentMemory, header, content)
}
// appendAfterHeader appends content after a section header.
func (t *UpdateMemoryTool) appendAfterHeader(memory, header, content string) string {
lines := strings.Split(memory, "\n")
var result []string
sectionFound := false
for i, line := range lines {
result = append(result, line)
// Find the header
if strings.TrimSpace(line) == header {
sectionFound = true
// Skip empty lines after header
j := i + 1
for j < len(lines) && strings.TrimSpace(lines[j]) == "" {
result = append(result, lines[j])
j++
}
// Add content as bullet point
result = append(result, fmt.Sprintf("- %s", content))
result = append(result, "") // Add spacing
}
}
// Fallback if header not found (shouldn't happen)
if !sectionFound {
return memory + "\n\n" + header + "\n\n- " + content + "\n"
}
return strings.Join(result, "\n")
}
// createNewSection creates a new section in MEMORY.md.
func (t *UpdateMemoryTool) createNewSection(currentMemory, header, content string) string {
if currentMemory == "" {
// Create new file with header
return "# Long-term Memory\n\n" + header + "\n\n- " + content + "\n"
}
// Append to end of file
return strings.TrimRight(currentMemory, "\n") + "\n\n" + header + "\n\n- " + content + "\n"
}

244
pkg/tools/memory_search.go Normal file
View file

@ -0,0 +1,244 @@
// PicoClaw - Ultra-lightweight personal AI agent
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package tools
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"time"
)
// SearchMemoryTool searches through memory files.
type SearchMemoryTool struct {
memory *MemoryStore
}
// NewSearchMemoryTool creates a new SearchMemoryTool.
func NewSearchMemoryTool(memory *MemoryStore) *SearchMemoryTool {
return &SearchMemoryTool{
memory: memory,
}
}
// Name returns the tool name.
func (t *SearchMemoryTool) Name() string {
return "search_memory"
}
// Description returns the tool description.
func (t *SearchMemoryTool) Description() string {
return `Search through long-term memory and daily notes by keywords.
Use this to find previously stored information.
Search scope:
- long_term: Search only MEMORY.md
- daily_notes: Search daily note files (YYYYMMDD.md)
- all: Search both (default)
Returns matching excerpts with context.`
}
// Parameters returns the tool parameters.
func (t *SearchMemoryTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{
"query": map[string]any{
"type": "string",
"description": "Search query (keywords or phrases)",
},
"memory_type": map[string]any{
"type": "string",
"enum": []string{"all", "long_term", "daily_notes"},
"description": "Scope of search: 'all' (default), 'long_term', or 'daily_notes'",
},
"days": map[string]any{
"type": "number",
"description": "Number of recent days to search for daily_notes (default: 7)",
},
},
"required": []string{"query"},
}
}
// Execute executes the tool.
func (t *SearchMemoryTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
query, ok := args["query"].(string)
if !ok || strings.TrimSpace(query) == "" {
return ErrorResult("query is required and cannot be empty").
WithError(fmt.Errorf("invalid query"))
}
memoryType, _ := args["memory_type"].(string)
if memoryType == "" {
memoryType = "all"
}
daysFloat, _ := args["days"].(float64)
days := int(daysFloat)
if days <= 0 {
days = 7 // default to 7 days
}
var results []string
// Search long-term memory
if memoryType == "all" || memoryType == "long_term" {
lmResults := t.searchLongTerm(query)
if lmResults != "" {
results = append(results, lmResults)
}
}
// Search daily notes
if memoryType == "all" || memoryType == "daily_notes" {
dnResults := t.searchDailyNotes(query, days)
if dnResults != "" {
results = append(results, dnResults)
}
}
if len(results) == 0 {
return &ToolResult{
ForLLM: "No relevant memories found for the query.",
ForUser: "🔍 未找到相关记忆",
Silent: false,
IsError: false,
Async: false,
}
}
// Format results
formattedResults := strings.Join(results, "\n\n---\n\n")
return &ToolResult{
ForLLM: fmt.Sprintf("Found %d result(s):\n\n%s", len(results), formattedResults),
ForUser: fmt.Sprintf("🔍 找到 %d 条相关记忆", len(results)),
Silent: false,
IsError: false,
Async: false,
}
}
// searchLongTerm searches in MEMORY.md.
func (t *SearchMemoryTool) searchLongTerm(query string) string {
content := t.memory.ReadLongTerm()
if content == "" {
return ""
}
// Case-insensitive search
queryLower := strings.ToLower(query)
contentLower := strings.ToLower(content)
if !strings.Contains(contentLower, queryLower) {
return ""
}
// Extract relevant paragraphs
excerpts := t.extractRelevantExcerpts(content, queryLower)
if len(excerpts) == 0 {
return ""
}
return "**长期记忆**\n\n" + strings.Join(excerpts, "\n\n")
}
// searchDailyNotes searches in daily note files.
func (t *SearchMemoryTool) searchDailyNotes(query string, days int) string {
var results []string
queryLower := strings.ToLower(query)
for i := 0; i < days; i++ {
date := time.Now().AddDate(0, 0, -i)
dateStr := date.Format("20060102") // YYYYMMDD
monthDir := dateStr[:6] // YYYYMM
// Get workspace from memory store
workspace := t.getWorkspacePath()
filePath := filepath.Join(workspace, "memory", monthDir, dateStr+".md")
content, err := os.ReadFile(filePath)
if err != nil {
continue // File doesn't exist, skip
}
contentStr := string(content)
contentLower := strings.ToLower(contentStr)
if strings.Contains(contentLower, queryLower) {
excerpts := t.extractRelevantExcerpts(contentStr, queryLower)
if len(excerpts) > 0 {
header := fmt.Sprintf("**%s 的笔记**", date.Format("2006-01-02"))
results = append(results, header+"\n\n"+strings.Join(excerpts, "\n\n"))
}
}
}
if len(results) == 0 {
return ""
}
return "**日常笔记**\n\n" + strings.Join(results, "\n\n")
}
// extractRelevantExcerpts extracts paragraphs containing the query.
func (t *SearchMemoryTool) extractRelevantExcerpts(content, queryLower string) []string {
var excerpts []string
// Split into sections by headers (## or #)
sections := strings.Split(content, "\n#")
for _, section := range sections {
sectionLower := strings.ToLower(section)
// Check if section contains query
if strings.Contains(sectionLower, queryLower) {
// Clean up section
cleanSection := strings.TrimSpace(section)
if !strings.HasPrefix(cleanSection, "#") {
cleanSection = "#" + cleanSection
}
// Limit excerpt length
if len(cleanSection) > 500 {
cleanSection = cleanSection[:500] + "..."
}
excerpts = append(excerpts, cleanSection)
}
}
// If no section matches, try paragraph-level search
if len(excerpts) == 0 {
paragraphs := strings.Split(content, "\n\n")
for _, para := range paragraphs {
if strings.Contains(strings.ToLower(para), queryLower) {
cleanPara := strings.TrimSpace(para)
if len(cleanPara) > 300 {
cleanPara = cleanPara[:300] + "..."
}
excerpts = append(excerpts, cleanPara)
// Limit to 3 paragraphs
if len(excerpts) >= 3 {
break
}
}
}
}
return excerpts
}
// getWorkspacePath extracts workspace path from MemoryStore.
func (t *SearchMemoryTool) getWorkspacePath() string {
return t.memory.GetWorkspace()
}

323
pkg/tools/memory_test.go Normal file
View file

@ -0,0 +1,323 @@
// PicoClaw - Ultra-lightweight personal AI agent
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package tools
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/sipeed/picoclaw/pkg/agent"
"github.com/stretchr/testify/assert"
)
func setupTestMemory(t *testing.T) (*agent.MemoryStore, func()) {
// Create temp directory
tmpDir := t.TempDir()
memory := agent.NewMemoryStore(tmpDir)
cleanup := func() {
os.RemoveAll(tmpDir)
}
return memory, cleanup
}
func TestUpdateMemoryTool_LongTerm(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
tool := NewUpdateMemoryTool(memory)
ctx := context.Background()
// Test adding user info
args := map[string]any{
"memory_type": "long_term",
"category": "user_info",
"content": "用户名叫小明,是一名 Go 语言开发者",
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
assert.Contains(t, result.ForUser, "长期记忆")
// Verify content was written
content := memory.ReadLongTerm()
assert.Contains(t, content, "## User Information")
assert.Contains(t, content, "用户名叫小明,是一名 Go 语言开发者")
}
func TestUpdateMemoryTool_DailyNote(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
tool := NewUpdateMemoryTool(memory)
ctx := context.Background()
args := map[string]any{
"memory_type": "daily_note",
"content": "完成了代码审查",
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
assert.Contains(t, result.ForUser, "今日笔记")
// Verify daily note was created
today := time.Now().Format("20060102")
monthDir := today[:6]
todayFile := filepath.Join(memory.GetWorkspace(), "memory", monthDir, today+".md")
content, err := os.ReadFile(todayFile)
assert.NoError(t, err)
assert.Contains(t, string(content), "完成了代码审查")
}
func TestUpdateMemoryTool_EmptyContent(t *testing.T) {
memory, _ := setupTestMemory(t)
tool := NewUpdateMemoryTool(memory)
ctx := context.Background()
args := map[string]any{
"memory_type": "long_term",
"content": " ", // empty after trim
}
result := tool.Execute(ctx, args)
assert.True(t, result.IsError)
assert.Contains(t, result.ForLLM, "empty")
}
func TestUpdateMemoryTool_InvalidType(t *testing.T) {
memory, _ := setupTestMemory(t)
tool := NewUpdateMemoryTool(memory)
ctx := context.Background()
args := map[string]any{
"memory_type": "invalid_type",
"content": "test content",
}
result := tool.Execute(ctx, args)
assert.True(t, result.IsError)
assert.Contains(t, result.ForLLM, "invalid memory_type")
}
func TestUpdateMemoryTool_AppendToExistingSection(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
// Create initial memory
initialContent := `# Long-term Memory
## User Information
- 初始信息`
err := memory.WriteLongTerm(initialContent)
assert.NoError(t, err)
tool := NewUpdateMemoryTool(memory)
ctx := context.Background()
args := map[string]any{
"memory_type": "long_term",
"category": "user_info",
"content": "新增的用户信息",
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
// Verify both items exist
content := memory.ReadLongTerm()
assert.Contains(t, content, "初始信息")
assert.Contains(t, content, "新增的用户信息")
}
func TestUpdateMemoryTool_CreateNewSection(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
// Create minimal initial memory
initialContent := "# Long-term Memory\n\n## Preferences\n\n- 偏好 1"
err := memory.WriteLongTerm(initialContent)
assert.NoError(t, err)
tool := NewUpdateMemoryTool(memory)
ctx := context.Background()
// Add to non-existing section
args := map[string]any{
"memory_type": "long_term",
"category": "task", // This section doesn't exist yet
"content": "新任务",
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
content := memory.ReadLongTerm()
assert.Contains(t, content, "## Tasks & Activities")
assert.Contains(t, content, "新任务")
}
func TestSearchMemoryTool_LongTerm(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
// Setup test data
testContent := `# Long-term Memory
## User Information
- 用户名叫小明
- 是一名 Go 语言开发者
## Preferences
- 喜欢喝拿铁咖啡
- 偏好中文交流`
err := memory.WriteLongTerm(testContent)
assert.NoError(t, err)
tool := NewSearchMemoryTool(memory)
ctx := context.Background()
// Test search that should find results
args := map[string]any{
"query": "小明",
"memory_type": "long_term",
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "小明")
assert.Contains(t, result.ForUser, "找到")
}
func TestSearchMemoryTool_NoResults(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
// Setup test data
testContent := `# Long-term Memory
## User Information
- 用户名叫小明`
err := memory.WriteLongTerm(testContent)
assert.NoError(t, err)
tool := NewSearchMemoryTool(memory)
ctx := context.Background()
// Test search that should NOT find results
args := map[string]any{
"query": "不存在的关键词",
"memory_type": "long_term",
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "No relevant memories")
assert.Contains(t, result.ForUser, "未找到")
}
func TestSearchMemoryTool_DailyNotes(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
// Create today's daily note
today := time.Now().Format("20060102")
monthDir := today[:6]
todayFile := filepath.Join(memory.GetWorkspace(), "memory", monthDir, today+".md")
err := os.MkdirAll(filepath.Dir(todayFile), 0o755)
assert.NoError(t, err)
noteContent := `# ` + today[:4] + "-" + today[4:6] + "-" + today[6:] + `
## Conversations
- 讨论了项目架构
- 帮助调试代码`
err = os.WriteFile(todayFile, []byte(noteContent), 0o600)
assert.NoError(t, err)
tool := NewSearchMemoryTool(memory)
ctx := context.Background()
args := map[string]any{
"query": "项目架构",
"memory_type": "daily_notes",
"days": float64(7),
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "项目架构")
}
func TestSearchMemoryTool_EmptyQuery(t *testing.T) {
memory, _ := setupTestMemory(t)
tool := NewSearchMemoryTool(memory)
ctx := context.Background()
args := map[string]any{
"query": "",
}
result := tool.Execute(ctx, args)
assert.True(t, result.IsError)
assert.Contains(t, result.ForLLM, "query is required")
}
func TestSearchMemoryTool_CaseInsensitive(t *testing.T) {
memory, cleanup := setupTestMemory(t)
defer cleanup()
// Setup test data with mixed case
testContent := `# Long-term Memory
## User Information
- 用户喜欢 GO 编程`
err := memory.WriteLongTerm(testContent)
assert.NoError(t, err)
tool := NewSearchMemoryTool(memory)
ctx := context.Background()
// Search with different case
args := map[string]any{
"query": "go", // lowercase
}
result := tool.Execute(ctx, args)
assert.False(t, result.IsError)
// Should find "GO" even though query is lowercase
assert.Contains(t, result.ForLLM, "GO")
}

View file

@ -3,6 +3,7 @@ package tools
import ( import (
"context" "context"
"fmt" "fmt"
"sync/atomic"
) )
type SendCallback func(channel, chatID, content string) error type SendCallback func(channel, chatID, content string) error
@ -11,7 +12,7 @@ type MessageTool struct {
sendCallback SendCallback sendCallback SendCallback
defaultChannel string defaultChannel string
defaultChatID string defaultChatID string
sentInRound bool // Tracks whether a message was sent in the current processing round sentInRound atomic.Bool // Tracks whether a message was sent in the current processing round
} }
func NewMessageTool() *MessageTool { func NewMessageTool() *MessageTool {
@ -50,12 +51,12 @@ func (t *MessageTool) Parameters() map[string]any {
func (t *MessageTool) SetContext(channel, chatID string) { func (t *MessageTool) SetContext(channel, chatID string) {
t.defaultChannel = channel t.defaultChannel = channel
t.defaultChatID = chatID t.defaultChatID = chatID
t.sentInRound = false // Reset send tracking for new processing round t.sentInRound.Store(false) // Reset send tracking for new processing round
} }
// HasSentInRound returns true if the message tool sent a message during the current round. // HasSentInRound returns true if the message tool sent a message during the current round.
func (t *MessageTool) HasSentInRound() bool { func (t *MessageTool) HasSentInRound() bool {
return t.sentInRound return t.sentInRound.Load()
} }
func (t *MessageTool) SetSendCallback(callback SendCallback) { func (t *MessageTool) SetSendCallback(callback SendCallback) {
@ -94,7 +95,7 @@ func (t *MessageTool) Execute(ctx context.Context, args map[string]any) *ToolRes
} }
} }
t.sentInRound = true t.sentInRound.Store(true)
// Silent: user already received the message directly // Silent: user already received the message directly
return &ToolResult{ return &ToolResult{
ForLLM: fmt.Sprintf("Message sent to %s:%s", channel, chatID), ForLLM: fmt.Sprintf("Message sent to %s:%s", channel, chatID),

View file

@ -25,7 +25,12 @@ func NewToolRegistry() *ToolRegistry {
func (r *ToolRegistry) Register(tool Tool) { func (r *ToolRegistry) Register(tool Tool) {
r.mu.Lock() r.mu.Lock()
defer r.mu.Unlock() defer r.mu.Unlock()
r.tools[tool.Name()] = tool name := tool.Name()
if _, exists := r.tools[name]; exists {
logger.WarnCF("tools", "Tool registration overwrites existing tool",
map[string]any{"name": name})
}
r.tools[name] = tool
} }
func (r *ToolRegistry) Get(name string) (Tool, bool) { func (r *ToolRegistry) Get(name string) (Tool, bool) {

View file

@ -329,7 +329,7 @@ func TestToolRegistry_ConcurrentAccess(t *testing.T) {
r := NewToolRegistry() r := NewToolRegistry()
var wg sync.WaitGroup var wg sync.WaitGroup
for i := 0; i < 50; i++ { for i := range 50 {
wg.Add(1) wg.Add(1)
go func(n int) { go func(n int) {
defer wg.Done() defer wg.Done()

View file

@ -21,53 +21,77 @@ type ExecTool struct {
timeout time.Duration timeout time.Duration
denyPatterns []*regexp.Regexp denyPatterns []*regexp.Regexp
allowPatterns []*regexp.Regexp allowPatterns []*regexp.Regexp
customAllowPatterns []*regexp.Regexp
restrictToWorkspace bool restrictToWorkspace bool
} }
var defaultDenyPatterns = []*regexp.Regexp{ var (
regexp.MustCompile(`\brm\s+-[rf]{1,2}\b`), defaultDenyPatterns = []*regexp.Regexp{
regexp.MustCompile(`\bdel\s+/[fq]\b`), regexp.MustCompile(`\brm\s+-[rf]{1,2}\b`),
regexp.MustCompile(`\brmdir\s+/s\b`), regexp.MustCompile(`\bdel\s+/[fq]\b`),
regexp.MustCompile(`\b(format|mkfs|diskpart)\b\s`), // Match disk wiping commands (must be followed by space/args) regexp.MustCompile(`\brmdir\s+/s\b`),
regexp.MustCompile(`\bdd\s+if=`), // Match disk wiping commands (must be followed by space/args)
regexp.MustCompile(`>\s*/dev/sd[a-z]\b`), // Block writes to disk devices (but allow /dev/null) regexp.MustCompile(
regexp.MustCompile(`\b(shutdown|reboot|poweroff)\b`), `\b(format|mkfs|diskpart)\b\s`,
regexp.MustCompile(`:\(\)\s*\{.*\};\s*:`), ),
regexp.MustCompile(`\$\([^)]+\)`), regexp.MustCompile(`\bdd\s+if=`),
regexp.MustCompile(`\$\{[^}]+\}`), // Block writes to block devices (all common naming schemes).
regexp.MustCompile("`[^`]+`"), regexp.MustCompile(
regexp.MustCompile(`\|\s*sh\b`), `>\s*/dev/(sd[a-z]|hd[a-z]|vd[a-z]|xvd[a-z]|nvme\d|mmcblk\d|loop\d|dm-\d|md\d|sr\d|nbd\d)`,
regexp.MustCompile(`\|\s*bash\b`), ),
regexp.MustCompile(`;\s*rm\s+-[rf]`), regexp.MustCompile(`\b(shutdown|reboot|poweroff)\b`),
regexp.MustCompile(`&&\s*rm\s+-[rf]`), regexp.MustCompile(`:\(\)\s*\{.*\};\s*:`),
regexp.MustCompile(`\|\|\s*rm\s+-[rf]`), regexp.MustCompile(`\$\([^)]+\)`),
regexp.MustCompile(`>\s*/dev/null\s*>&?\s*\d?`), regexp.MustCompile(`\$\{[^}]+\}`),
regexp.MustCompile(`<<\s*EOF`), regexp.MustCompile("`[^`]+`"),
regexp.MustCompile(`\$\(\s*cat\s+`), regexp.MustCompile(`\|\s*sh\b`),
regexp.MustCompile(`\$\(\s*curl\s+`), regexp.MustCompile(`\|\s*bash\b`),
regexp.MustCompile(`\$\(\s*wget\s+`), regexp.MustCompile(`;\s*rm\s+-[rf]`),
regexp.MustCompile(`\$\(\s*which\s+`), regexp.MustCompile(`&&\s*rm\s+-[rf]`),
regexp.MustCompile(`\bsudo\b`), regexp.MustCompile(`\|\|\s*rm\s+-[rf]`),
regexp.MustCompile(`\bchmod\s+[0-7]{3,4}\b`), regexp.MustCompile(`<<\s*EOF`),
regexp.MustCompile(`\bchown\b`), regexp.MustCompile(`\$\(\s*cat\s+`),
regexp.MustCompile(`\bpkill\b`), regexp.MustCompile(`\$\(\s*curl\s+`),
regexp.MustCompile(`\bkillall\b`), regexp.MustCompile(`\$\(\s*wget\s+`),
regexp.MustCompile(`\bkill\s+-[9]\b`), regexp.MustCompile(`\$\(\s*which\s+`),
regexp.MustCompile(`\bcurl\b.*\|\s*(sh|bash)`), regexp.MustCompile(`\bsudo\b`),
regexp.MustCompile(`\bwget\b.*\|\s*(sh|bash)`), regexp.MustCompile(`\bchmod\s+[0-7]{3,4}\b`),
regexp.MustCompile(`\bnpm\s+install\s+-g\b`), regexp.MustCompile(`\bchown\b`),
regexp.MustCompile(`\bpip\s+install\s+--user\b`), regexp.MustCompile(`\bpkill\b`),
regexp.MustCompile(`\bapt\s+(install|remove|purge)\b`), regexp.MustCompile(`\bkillall\b`),
regexp.MustCompile(`\byum\s+(install|remove)\b`), regexp.MustCompile(`\bkill\s+-[9]\b`),
regexp.MustCompile(`\bdnf\s+(install|remove)\b`), regexp.MustCompile(`\bcurl\b.*\|\s*(sh|bash)`),
regexp.MustCompile(`\bdocker\s+run\b`), regexp.MustCompile(`\bwget\b.*\|\s*(sh|bash)`),
regexp.MustCompile(`\bdocker\s+exec\b`), regexp.MustCompile(`\bnpm\s+install\s+-g\b`),
regexp.MustCompile(`\bgit\s+push\b`), regexp.MustCompile(`\bpip\s+install\s+--user\b`),
regexp.MustCompile(`\bgit\s+force\b`), regexp.MustCompile(`\bapt\s+(install|remove|purge)\b`),
regexp.MustCompile(`\bssh\b.*@`), regexp.MustCompile(`\byum\s+(install|remove)\b`),
regexp.MustCompile(`\beval\b`), regexp.MustCompile(`\bdnf\s+(install|remove)\b`),
regexp.MustCompile(`\bsource\s+.*\.sh\b`), regexp.MustCompile(`\bdocker\s+run\b`),
} regexp.MustCompile(`\bdocker\s+exec\b`),
regexp.MustCompile(`\bgit\s+push\b`),
regexp.MustCompile(`\bgit\s+force\b`),
regexp.MustCompile(`\bssh\b.*@`),
regexp.MustCompile(`\beval\b`),
regexp.MustCompile(`\bsource\s+.*\.sh\b`),
}
// absolutePathPattern matches absolute file paths in commands (Unix and Windows).
absolutePathPattern = regexp.MustCompile(`[A-Za-z]:\\[^\\\"']+|/[^\s\"']+`)
// safePaths are kernel pseudo-devices that are always safe to reference in
// commands, regardless of workspace restriction. They contain no user data
// and cannot cause destructive writes.
safePaths = map[string]bool{
"/dev/null": true,
"/dev/zero": true,
"/dev/random": true,
"/dev/urandom": true,
"/dev/stdin": true,
"/dev/stdout": true,
"/dev/stderr": true,
}
)
func NewExecTool(workingDir string, restrict bool) (*ExecTool, error) { func NewExecTool(workingDir string, restrict bool) (*ExecTool, error) {
return NewExecToolWithConfig(workingDir, restrict, nil) return NewExecToolWithConfig(workingDir, restrict, nil)
@ -75,6 +99,7 @@ func NewExecTool(workingDir string, restrict bool) (*ExecTool, error) {
func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Config) (*ExecTool, error) { func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Config) (*ExecTool, error) {
denyPatterns := make([]*regexp.Regexp, 0) denyPatterns := make([]*regexp.Regexp, 0)
customAllowPatterns := make([]*regexp.Regexp, 0)
if config != nil { if config != nil {
execConfig := config.Tools.Exec execConfig := config.Tools.Exec
@ -95,6 +120,13 @@ func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Conf
// If deny patterns are disabled, we won't add any patterns, allowing all commands. // If deny patterns are disabled, we won't add any patterns, allowing all commands.
fmt.Println("Warning: deny patterns are disabled. All commands will be allowed.") fmt.Println("Warning: deny patterns are disabled. All commands will be allowed.")
} }
for _, pattern := range execConfig.CustomAllowPatterns {
re, err := regexp.Compile(pattern)
if err != nil {
return nil, fmt.Errorf("invalid custom allow pattern %q: %w", pattern, err)
}
customAllowPatterns = append(customAllowPatterns, re)
}
} else { } else {
denyPatterns = append(denyPatterns, defaultDenyPatterns...) denyPatterns = append(denyPatterns, defaultDenyPatterns...)
} }
@ -104,6 +136,7 @@ func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Conf
timeout: 60 * time.Second, timeout: 60 * time.Second,
denyPatterns: denyPatterns, denyPatterns: denyPatterns,
allowPatterns: nil, allowPatterns: nil,
customAllowPatterns: customAllowPatterns,
restrictToWorkspace: restrict, restrictToWorkspace: restrict,
}, nil }, nil
} }
@ -258,9 +291,20 @@ func (t *ExecTool) guardCommand(command, cwd string) string {
cmd := strings.TrimSpace(command) cmd := strings.TrimSpace(command)
lower := strings.ToLower(cmd) lower := strings.ToLower(cmd)
for _, pattern := range t.denyPatterns { // Custom allow patterns exempt a command from deny checks.
explicitlyAllowed := false
for _, pattern := range t.customAllowPatterns {
if pattern.MatchString(lower) { if pattern.MatchString(lower) {
return "Command blocked by safety guard (dangerous pattern detected)" explicitlyAllowed = true
break
}
}
if !explicitlyAllowed {
for _, pattern := range t.denyPatterns {
if pattern.MatchString(lower) {
return "Command blocked by safety guard (dangerous pattern detected)"
}
} }
} }
@ -287,8 +331,7 @@ func (t *ExecTool) guardCommand(command, cwd string) string {
return "" return ""
} }
pathPattern := regexp.MustCompile(`[A-Za-z]:\\[^\\\"']+|/[^\s\"']+`) matches := absolutePathPattern.FindAllString(cmd, -1)
matches := pathPattern.FindAllString(cmd, -1)
for _, raw := range matches { for _, raw := range matches {
p, err := filepath.Abs(raw) p, err := filepath.Abs(raw)
@ -296,6 +339,10 @@ func (t *ExecTool) guardCommand(command, cwd string) string {
continue continue
} }
if safePaths[p] {
continue
}
rel, err := filepath.Rel(cwdPath, p) rel, err := filepath.Rel(cwdPath, p)
if err != nil { if err != nil {
continue continue

Some files were not shown because too many files have changed in this diff Show more