Merge remote-tracking branch 'github' into feat-agent-message
This commit is contained in:
commit
15da43a1c4
28 changed files with 2261 additions and 163 deletions
14
README.md
14
README.md
|
|
@ -721,6 +721,20 @@ PicoClaw stores data in your configured workspace (default: `~/.picoclaw/workspa
|
||||||
└── USER.md # User preferences
|
└── USER.md # User preferences
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Skill Sources
|
||||||
|
|
||||||
|
By default, skills are loaded from:
|
||||||
|
|
||||||
|
1. `~/.picoclaw/workspace/skills` (workspace)
|
||||||
|
2. `~/.picoclaw/skills` (global)
|
||||||
|
3. `<current-working-directory>/skills` (builtin)
|
||||||
|
|
||||||
|
For advanced/test setups, you can override the builtin skills root with:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export PICOCLAW_BUILTIN_SKILLS=/path/to/skills
|
||||||
|
```
|
||||||
|
|
||||||
### 🔒 Security Sandbox
|
### 🔒 Security Sandbox
|
||||||
|
|
||||||
PicoClaw runs in a sandboxed environment by default. The agent can only access files and execute commands within the configured workspace.
|
PicoClaw runs in a sandboxed environment by default. The agent can only access files and execute commands within the configured workspace.
|
||||||
|
|
|
||||||
14
README.zh.md
14
README.zh.md
|
|
@ -362,6 +362,20 @@ PicoClaw 将数据存储在您配置的工作区中(默认:`~/.picoclaw/work
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### 技能来源 (Skill Sources)
|
||||||
|
|
||||||
|
默认情况下,技能会按以下顺序加载:
|
||||||
|
|
||||||
|
1. `~/.picoclaw/workspace/skills`(工作区)
|
||||||
|
2. `~/.picoclaw/skills`(全局)
|
||||||
|
3. `<current-working-directory>/skills`(内置)
|
||||||
|
|
||||||
|
在高级/测试场景下,可通过以下环境变量覆盖内置技能目录:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export PICOCLAW_BUILTIN_SKILLS=/path/to/skills
|
||||||
|
```
|
||||||
|
|
||||||
### 心跳 / 周期性任务 (Heartbeat)
|
### 心跳 / 周期性任务 (Heartbeat)
|
||||||
|
|
||||||
PicoClaw 可以自动执行周期性任务。在工作区创建 `HEARTBEAT.md` 文件:
|
PicoClaw 可以自动执行周期性任务。在工作区创建 `HEARTBEAT.md` 文件:
|
||||||
|
|
|
||||||
|
|
@ -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,6 +59,7 @@
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"token": "YOUR_DISCORD_BOT_TOKEN",
|
"token": "YOUR_DISCORD_BOT_TOKEN",
|
||||||
|
"proxy": "",
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"group_trigger": {
|
"group_trigger": {
|
||||||
"mention_only": false
|
"mention_only": false
|
||||||
|
|
@ -246,6 +248,14 @@
|
||||||
"mcp": {
|
"mcp": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"servers": {
|
"servers": {
|
||||||
|
"context7": {
|
||||||
|
"enabled": false,
|
||||||
|
"type": "http",
|
||||||
|
"url": "https://mcp.context7.com/mcp",
|
||||||
|
"headers": {
|
||||||
|
"CONTEXT7_API_KEY": "ctx7sk-xx"
|
||||||
|
}
|
||||||
|
},
|
||||||
"filesystem": {
|
"filesystem": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"command": "npx",
|
"command": "npx",
|
||||||
|
|
|
||||||
2
go.mod
2
go.mod
|
|
@ -9,6 +9,7 @@ require (
|
||||||
github.com/caarlos0/env/v11 v11.3.1
|
github.com/caarlos0/env/v11 v11.3.1
|
||||||
github.com/chzyer/readline v1.5.1
|
github.com/chzyer/readline v1.5.1
|
||||||
github.com/gdamore/tcell/v2 v2.13.8
|
github.com/gdamore/tcell/v2 v2.13.8
|
||||||
|
github.com/gdamore/tcell/v2 v2.13.8
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
||||||
|
|
@ -37,6 +38,7 @@ require (
|
||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect
|
github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect
|
||||||
github.com/gdamore/encoding v1.0.1 // indirect
|
github.com/gdamore/encoding v1.0.1 // indirect
|
||||||
|
github.com/h2non/filetype v1.1.3 // indirect
|
||||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||||
github.com/lucasb-eyer/go-colorful v1.3.0 // indirect
|
github.com/lucasb-eyer/go-colorful v1.3.0 // indirect
|
||||||
github.com/mattn/go-colorable v0.1.14 // indirect
|
github.com/mattn/go-colorable v0.1.14 // indirect
|
||||||
|
|
|
||||||
2
go.sum
2
go.sum
|
|
@ -98,6 +98,8 @@ github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aN
|
||||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||||
github.com/grbit/go-json v0.11.0 h1:bAbyMdYrYl/OjYsSqLH99N2DyQ291mHy726Mx+sYrnc=
|
github.com/grbit/go-json v0.11.0 h1:bAbyMdYrYl/OjYsSqLH99N2DyQ291mHy726Mx+sYrnc=
|
||||||
github.com/grbit/go-json v0.11.0/go.mod h1:IYpHsdybQ386+6g3VE6AXQ3uTGa5mquBme5/ZWmtzek=
|
github.com/grbit/go-json v0.11.0/go.mod h1:IYpHsdybQ386+6g3VE6AXQ3uTGa5mquBme5/ZWmtzek=
|
||||||
|
github.com/h2non/filetype v1.1.3 h1:FKkx9QbD7HR/zjK1Ia5XiBsq9zdLi5Kf3zGyFTAFkGg=
|
||||||
|
github.com/h2non/filetype v1.1.3/go.mod h1:319b3zT68BvV+WRj7cwy856M2ehB3HqNOt6sy1HndBY=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
||||||
github.com/hpcloud/tail v1.0.0/go.mod h1:ab1qPbhIpdTxEkNHXyeSf5vhxWSCs/tWer42PpOxQnU=
|
github.com/hpcloud/tail v1.0.0/go.mod h1:ab1qPbhIpdTxEkNHXyeSf5vhxWSCs/tWer42PpOxQnU=
|
||||||
|
|
|
||||||
|
|
@ -34,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 {
|
||||||
|
|
@ -47,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{
|
||||||
|
|
@ -167,6 +175,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{
|
||||||
|
|
@ -186,14 +195,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"),
|
||||||
|
|
@ -204,23 +213,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 {
|
||||||
|
|
@ -231,17 +256,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:
|
||||||
|
|
@ -253,7 +282,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
|
||||||
|
|
@ -273,21 +302,17 @@ func (cb *ContextBuilder) sourceFilesChangedLocked() bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- Skills directory (handled separately from sourcePaths) ---
|
// --- Skill roots (workspace/global/builtin) ---
|
||||||
//
|
//
|
||||||
// 1. Creation/deletion: tracked via existedAtCache, same as bootstrap files.
|
// For each root:
|
||||||
skillsDir := filepath.Join(cb.workspace, "skills")
|
// 1. Creation/deletion and root directory mtime changes are tracked by fileChangedSince.
|
||||||
if cb.fileChangedSince(skillsDir) {
|
// 2. Nested file create/delete/mtime changes are tracked by the skill file snapshot.
|
||||||
return true
|
for _, root := range cb.skillRoots() {
|
||||||
|
if cb.fileChangedSince(root) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
if skillFilesChangedSince(cb.skillRoots(), cb.skillFilesAtCache) {
|
||||||
// 2. Structural changes (add/remove entries inside the dir) are reflected
|
|
||||||
// in the directory's own mtime, which fileChangedSince already checks.
|
|
||||||
//
|
|
||||||
// 3. Content-only edits to files inside skills/ do NOT update the parent
|
|
||||||
// directory mtime on most filesystems, so we recursively walk to check
|
|
||||||
// individual file mtimes at any nesting depth.
|
|
||||||
if skillFilesModifiedSince(skillsDir, cb.cachedAt) {
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -328,28 +353,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 {
|
||||||
|
|
@ -485,10 +546,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.
|
||||||
|
|
|
||||||
|
|
@ -51,14 +51,15 @@ 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."
|
||||||
|
|
@ -196,6 +197,17 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
// Initialize MCP servers for all agents
|
// Initialize MCP servers for all agents
|
||||||
if al.cfg.Tools.MCP.Enabled {
|
if al.cfg.Tools.MCP.Enabled {
|
||||||
mcpManager := mcp.NewManager()
|
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()
|
defaultAgent := al.registry.GetDefaultAgent()
|
||||||
var workspacePath string
|
var workspacePath string
|
||||||
if defaultAgent != nil && defaultAgent.Workspace != "" {
|
if defaultAgent != nil && defaultAgent.Workspace != "" {
|
||||||
|
|
@ -210,16 +222,6 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
// Ensure MCP connections are cleaned up on exit, only if initialization succeeded
|
|
||||||
defer func() {
|
|
||||||
if err := mcpManager.Close(); err != nil {
|
|
||||||
logger.ErrorCF("agent", "Failed to close MCP manager",
|
|
||||||
map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Register MCP tools for all agents
|
// Register MCP tools for all agents
|
||||||
servers := mcpManager.GetServers()
|
servers := mcpManager.GetServers()
|
||||||
uniqueTools := 0
|
uniqueTools := 0
|
||||||
|
|
@ -696,6 +698,7 @@ func (al *AgentLoop) processMessageWithTask(ctx context.Context, task *Task, msg
|
||||||
Channel: msg.Channel,
|
Channel: msg.Channel,
|
||||||
ChatID: msg.ChatID,
|
ChatID: msg.ChatID,
|
||||||
UserMessage: msg.Content,
|
UserMessage: msg.Content,
|
||||||
|
Media: msg.Media,
|
||||||
EnableSummary: true,
|
EnableSummary: true,
|
||||||
SendResponse: false,
|
SendResponse: false,
|
||||||
})
|
})
|
||||||
|
|
@ -831,11 +834,15 @@ func (al *AgentLoop) runAgentLoop(
|
||||||
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)
|
||||||
|
|
||||||
|
|
|
||||||
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
|
||||||
|
}
|
||||||
|
|
@ -6,12 +6,14 @@ import (
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"slices"
|
"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"
|
||||||
)
|
)
|
||||||
|
|
@ -808,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])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -181,14 +181,19 @@ type AgentDefaults struct {
|
||||||
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
||||||
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
||||||
|
|
||||||
// Steering architecture (nanobot-inspired, opt-in)
|
|
||||||
|
|
||||||
// Legacy: Phase 2 concurrent task management (to be deprecated)
|
|
||||||
MaxConcurrentTasks int `json:"max_concurrent_tasks,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_CONCURRENT_TASKS"` // Maximum concurrent tasks (0=unlimited)
|
MaxConcurrentTasks int `json:"max_concurrent_tasks,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_CONCURRENT_TASKS"` // Maximum concurrent tasks (0=unlimited)
|
||||||
EnableSteeringLoop bool `json:"enable_steering_loop,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_ENABLE_STEERING_LOOP"` // Enable steering loop for interrupt monitoring
|
|
||||||
SteeringLoopIntervalMs int `json:"steering_loop_interval_ms,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_STEERING_LOOP_INTERVAL_MS"` // Steering loop check interval (ms)
|
|
||||||
TaskCleanupIntervalMins int `json:"task_cleanup_interval_mins,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TASK_CLEANUP_INTERVAL_MINS"` // Task cleanup interval (minutes)
|
TaskCleanupIntervalMins int `json:"task_cleanup_interval_mins,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TASK_CLEANUP_INTERVAL_MINS"` // Task cleanup interval (minutes)
|
||||||
TaskRetentionHours int `json:"task_retention_hours,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TASK_RETENTION_HOURS"` // Task retention time (hours)
|
TaskRetentionHours int `json:"task_retention_hours,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TASK_RETENTION_HOURS"` // Task retention time (hours)
|
||||||
|
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.
|
||||||
|
|
@ -246,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"`
|
||||||
|
|
@ -262,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"`
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import (
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||||
|
|
||||||
|
|
@ -108,7 +109,7 @@ type ServerConnection struct {
|
||||||
type Manager struct {
|
type Manager struct {
|
||||||
servers map[string]*ServerConnection
|
servers map[string]*ServerConnection
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
closed bool
|
closed atomic.Bool // changed from bool to atomic.Bool to avoid TOCTOU race
|
||||||
wg sync.WaitGroup // tracks in-flight CallTool calls
|
wg sync.WaitGroup // tracks in-flight CallTool calls
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -440,14 +441,20 @@ func (m *Manager) CallTool(
|
||||||
serverName, toolName string,
|
serverName, toolName string,
|
||||||
arguments map[string]any,
|
arguments map[string]any,
|
||||||
) (*mcp.CallToolResult, error) {
|
) (*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()
|
m.mu.RLock()
|
||||||
if m.closed {
|
// Double-check after acquiring lock to prevent TOCTOU race
|
||||||
|
if m.closed.Load() {
|
||||||
m.mu.RUnlock()
|
m.mu.RUnlock()
|
||||||
return nil, fmt.Errorf("manager is closed")
|
return nil, fmt.Errorf("manager is closed")
|
||||||
}
|
}
|
||||||
conn, ok := m.servers[serverName]
|
conn, ok := m.servers[serverName]
|
||||||
if ok {
|
if ok {
|
||||||
m.wg.Add(1)
|
m.wg.Add(1) // Add to WaitGroup while holding the lock
|
||||||
}
|
}
|
||||||
m.mu.RUnlock()
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
|
@ -471,15 +478,14 @@ func (m *Manager) CallTool(
|
||||||
|
|
||||||
// Close closes all server connections
|
// Close closes all server connections
|
||||||
func (m *Manager) Close() error {
|
func (m *Manager) Close() error {
|
||||||
m.mu.Lock()
|
// Use Swap to atomically set closed=true and get the previous value
|
||||||
if m.closed {
|
// This prevents TOCTOU race with CallTool's closed check
|
||||||
m.mu.Unlock()
|
if m.closed.Swap(true) {
|
||||||
return nil
|
return nil // already closed
|
||||||
}
|
}
|
||||||
m.closed = true
|
|
||||||
m.mu.Unlock()
|
|
||||||
|
|
||||||
// Wait for all in-flight CallTool calls to finish before closing sessions
|
// 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.wg.Wait()
|
||||||
|
|
||||||
m.mu.Lock()
|
m.mu.Lock()
|
||||||
|
|
|
||||||
|
|
@ -268,7 +268,7 @@ func TestGetAllTools_FiltersEmptyTools(t *testing.T) {
|
||||||
func TestCallTool_ErrorsForClosedOrMissingServer(t *testing.T) {
|
func TestCallTool_ErrorsForClosedOrMissingServer(t *testing.T) {
|
||||||
t.Run("manager closed", func(t *testing.T) {
|
t.Run("manager closed", func(t *testing.T) {
|
||||||
mgr := NewManager()
|
mgr := NewManager()
|
||||||
mgr.closed = true
|
mgr.closed.Store(true)
|
||||||
|
|
||||||
_, err := mgr.CallTool(context.Background(), "s1", "tool", nil)
|
_, err := mgr.CallTool(context.Background(), "s1", "tool", nil)
|
||||||
if err == nil || !strings.Contains(err.Error(), "manager is closed") {
|
if err == nil || !strings.Contains(err.Error(), "manager is closed") {
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
@ -296,19 +296,55 @@ type openaiMessage struct {
|
||||||
ToolCallID string `json:"tool_call_id,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{
|
||||||
ReasoningContent: m.ReasoningContent,
|
Role: m.Role,
|
||||||
ToolCalls: m.ToolCalls,
|
Content: m.Content,
|
||||||
ToolCallID: m.ToolCallID,
|
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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
@ -416,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"`
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -109,6 +109,10 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
return "", fmt.Errorf("failed to read response: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("brave api error (status %d): %s", resp.StatusCode, string(body))
|
||||||
|
}
|
||||||
|
|
||||||
var searchResp struct {
|
var searchResp struct {
|
||||||
Web struct {
|
Web struct {
|
||||||
Results []struct {
|
Results []struct {
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,7 @@ package utils
|
||||||
import (
|
import (
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
@ -52,11 +53,12 @@ type DownloadOptions struct {
|
||||||
Timeout time.Duration
|
Timeout time.Duration
|
||||||
ExtraHeaders map[string]string
|
ExtraHeaders map[string]string
|
||||||
LoggerPrefix string
|
LoggerPrefix string
|
||||||
|
ProxyURL string
|
||||||
}
|
}
|
||||||
|
|
||||||
// DownloadFile downloads a file from URL to a local temp directory.
|
// DownloadFile downloads a file from URL to a local temp directory.
|
||||||
// Returns the local file path or empty string on error.
|
// Returns the local file path or empty string on error.
|
||||||
func DownloadFile(url, filename string, opts DownloadOptions) string {
|
func DownloadFile(urlStr, filename string, opts DownloadOptions) string {
|
||||||
// Set defaults
|
// Set defaults
|
||||||
if opts.Timeout == 0 {
|
if opts.Timeout == 0 {
|
||||||
opts.Timeout = 60 * time.Second
|
opts.Timeout = 60 * time.Second
|
||||||
|
|
@ -78,7 +80,7 @@ func DownloadFile(url, filename string, opts DownloadOptions) string {
|
||||||
localPath := filepath.Join(mediaDir, uuid.New().String()[:8]+"_"+safeName)
|
localPath := filepath.Join(mediaDir, uuid.New().String()[:8]+"_"+safeName)
|
||||||
|
|
||||||
// Create HTTP request
|
// Create HTTP request
|
||||||
req, err := http.NewRequest("GET", url, nil)
|
req, err := http.NewRequest("GET", urlStr, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF(opts.LoggerPrefix, "Failed to create download request", map[string]any{
|
logger.ErrorCF(opts.LoggerPrefix, "Failed to create download request", map[string]any{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
|
|
@ -92,11 +94,24 @@ func DownloadFile(url, filename string, opts DownloadOptions) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
client := &http.Client{Timeout: opts.Timeout}
|
client := &http.Client{Timeout: opts.Timeout}
|
||||||
|
if opts.ProxyURL != "" {
|
||||||
|
proxyURL, parseErr := url.Parse(opts.ProxyURL)
|
||||||
|
if parseErr != nil {
|
||||||
|
logger.ErrorCF(opts.LoggerPrefix, "Invalid proxy URL for download", map[string]any{
|
||||||
|
"error": parseErr.Error(),
|
||||||
|
"proxy": opts.ProxyURL,
|
||||||
|
})
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
client.Transport = &http.Transport{
|
||||||
|
Proxy: http.ProxyURL(proxyURL),
|
||||||
|
}
|
||||||
|
}
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF(opts.LoggerPrefix, "Failed to download file", map[string]any{
|
logger.ErrorCF(opts.LoggerPrefix, "Failed to download file", map[string]any{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
"url": url,
|
"url": urlStr,
|
||||||
})
|
})
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
@ -105,7 +120,7 @@ func DownloadFile(url, filename string, opts DownloadOptions) string {
|
||||||
if resp.StatusCode != http.StatusOK {
|
if resp.StatusCode != http.StatusOK {
|
||||||
logger.ErrorCF(opts.LoggerPrefix, "File download returned non-200 status", map[string]any{
|
logger.ErrorCF(opts.LoggerPrefix, "File download returned non-200 status", map[string]any{
|
||||||
"status": resp.StatusCode,
|
"status": resp.StatusCode,
|
||||||
"url": url,
|
"url": urlStr,
|
||||||
})
|
})
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue