fix: 修复了飞书,企业微信代码多模态的处理代码,增加单页面WEB
This commit is contained in:
parent
8207c1c7e6
commit
6668c9eed1
117 changed files with 13612 additions and 1674 deletions
|
|
@ -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))
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
|
||||||
25
cmd/picoclaw/internal/onboard/helpers_test.go
Normal file
25
cmd/picoclaw/internal/onboard/helpers_test.go
Normal 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
122
cmd/picoclaw/internal/web/command.go
Normal file
122
cmd/picoclaw/internal/web/command.go
Normal 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()
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
44
docker/Dockerfile.full
Normal 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"]
|
||||||
44
docker/docker-compose.full.yml
Normal file
44
docker/docker-compose.full.yml
Normal 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
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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")
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
122
pkg/agent/loop_media.go
Normal 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
|
||||||
|
}
|
||||||
|
|
@ -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])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
@ -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` 的 channel(Telegram、Discord、Pico)能真正使用占位消息编辑功能。其余 channel 的 `PlaceholderConfig` 为预留字段。
|
||||||
|
|
||||||
|
8. **ReasoningChannelID**:大多数 channel config 都包含 `reasoning_channel_id` 字段,用于将 LLM 的思维链(reasoning/thinking)路由到指定 channel(WhatsApp、Telegram、Feishu、Discord、MaixCam、QQ、DingTalk、Slack、LINE、OneBot、WeCom、WeComApp)。注意:`PicoConfig` 目前不包含该字段。`BaseChannel` 通过 `WithReasoningChannelID` 选项和 `ReasoningChannelID()` 方法暴露此配置。
|
||||||
|
|
@ -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("", 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()
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
91
pkg/channels/discord/discord_test.go
Normal file
91
pkg/channels/discord/discord_test.go
Normal 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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
|
||||||
292
pkg/channels/feishu/common_test.go
Normal file
292
pkg/channels/feishu/common_test.go
Normal 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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
|
|
|
||||||
256
pkg/channels/feishu/feishu_64_test.go
Normal file
256
pkg/channels/feishu/feishu_64_test.go
Normal 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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
1014
pkg/channels/wecom/aibot.go
Normal file
File diff suppressed because it is too large
Load diff
210
pkg/channels/wecom/aibot_test.go
Normal file
210
pkg/channels/wecom/aibot_test.go
Normal 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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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]"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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+"×tamp="+timestamp+"&nonce="+nonce,
|
"/webhook/wecom?msg_signature="+signature+"×tamp="+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+"×tamp="+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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
54
pkg/channels/wecom/dedupe.go
Normal file
54
pkg/channels/wecom/dedupe.go
Normal 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
|
||||||
|
}
|
||||||
83
pkg/channels/wecom/dedupe_test.go
Normal file
83
pkg/channels/wecom/dedupe_test.go
Normal 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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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
532
pkg/mcp/manager.go
Normal 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
298
pkg/mcp/manager_test.go
Normal 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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"`
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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":
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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"`
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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] {
|
||||||
|
|
|
||||||
|
|
@ -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{}{}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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 == "" {
|
||||||
|
|
|
||||||
|
|
@ -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
246
pkg/tools/mcp_tool.go
Normal 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
492
pkg/tools/mcp_tool_test.go
Normal 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
298
pkg/tools/memory.go
Normal 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
244
pkg/tools/memory_search.go
Normal 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
323
pkg/tools/memory_test.go
Normal 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")
|
||||||
|
}
|
||||||
|
|
@ -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),
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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
Loading…
Add table
Reference in a new issue