Merge main into d4
This commit is contained in:
commit
0fd7a401a2
126 changed files with 7924 additions and 1024 deletions
|
|
@ -18,10 +18,10 @@ builds:
|
||||||
- stdjson
|
- stdjson
|
||||||
ldflags:
|
ldflags:
|
||||||
- -s -w
|
- -s -w
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.version={{ .Version }}
|
- -X github.com/sipeed/picoclaw/pkg/config.Version={{ .Version }}
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.gitCommit={{ .ShortCommit }}
|
- -X github.com/sipeed/picoclaw/pkg/config.GitCommit={{ .ShortCommit }}
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.buildTime={{ .Date }}
|
- -X github.com/sipeed/picoclaw/pkg/config.BuildTime={{ .Date }}
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.goVersion={{ .Env.GOVERSION }}
|
- -X github.com/sipeed/picoclaw/pkg/config.GoVersion={{ .Env.GOVERSION }}
|
||||||
goos:
|
goos:
|
||||||
- linux
|
- linux
|
||||||
- windows
|
- windows
|
||||||
|
|
@ -125,6 +125,23 @@ dockers_v2:
|
||||||
- linux/arm64
|
- linux/arm64
|
||||||
- linux/riscv64
|
- linux/riscv64
|
||||||
|
|
||||||
|
- id: picoclaw-launcher
|
||||||
|
dockerfile: docker/Dockerfile.goreleaser.launcher
|
||||||
|
ids:
|
||||||
|
- picoclaw
|
||||||
|
- picoclaw-launcher
|
||||||
|
- picoclaw-launcher-tui
|
||||||
|
images:
|
||||||
|
- "ghcr.io/{{ .Env.GITHUB_REPOSITORY_OWNER }}/picoclaw"
|
||||||
|
- '{{ if not (isEnvSet "NIGHTLY_BUILD") }}docker.io/{{ .Env.DOCKERHUB_IMAGE_NAME }}{{ end }}'
|
||||||
|
tags:
|
||||||
|
- "{{ .Tag }}-launcher"
|
||||||
|
- '{{ if isEnvSet "NIGHTLY_BUILD" }}nightly-launcher{{ else }}launcher{{ end }}'
|
||||||
|
platforms:
|
||||||
|
- linux/amd64
|
||||||
|
- linux/arm64
|
||||||
|
- linux/riscv64
|
||||||
|
|
||||||
notarize:
|
notarize:
|
||||||
macos:
|
macos:
|
||||||
- enabled: '{{ isEnvSet "MACOS_SIGN_P12" }}'
|
- enabled: '{{ isEnvSet "MACOS_SIGN_P12" }}'
|
||||||
|
|
|
||||||
4
Makefile
4
Makefile
|
|
@ -11,8 +11,8 @@ VERSION?=$(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
|
||||||
GIT_COMMIT=$(shell git rev-parse --short=8 HEAD 2>/dev/null || echo "dev")
|
GIT_COMMIT=$(shell git rev-parse --short=8 HEAD 2>/dev/null || echo "dev")
|
||||||
BUILD_TIME=$(shell date +%FT%T%z)
|
BUILD_TIME=$(shell date +%FT%T%z)
|
||||||
GO_VERSION=$(shell $(GO) version | awk '{print $$3}')
|
GO_VERSION=$(shell $(GO) version | awk '{print $$3}')
|
||||||
INTERNAL=github.com/sipeed/picoclaw/cmd/picoclaw/internal
|
CONFIG_PKG=github.com/sipeed/picoclaw/pkg/config
|
||||||
LDFLAGS=-ldflags "-X $(INTERNAL).version=$(VERSION) -X $(INTERNAL).gitCommit=$(GIT_COMMIT) -X $(INTERNAL).buildTime=$(BUILD_TIME) -X $(INTERNAL).goVersion=$(GO_VERSION) -s -w"
|
LDFLAGS=-ldflags "-X $(CONFIG_PKG).Version=$(VERSION) -X $(CONFIG_PKG).GitCommit=$(GIT_COMMIT) -X $(CONFIG_PKG).BuildTime=$(BUILD_TIME) -X $(CONFIG_PKG).GoVersion=$(GO_VERSION) -s -w"
|
||||||
|
|
||||||
# Go variables
|
# Go variables
|
||||||
GO?=CGO_ENABLED=0 go
|
GO?=CGO_ENABLED=0 go
|
||||||
|
|
|
||||||
|
|
@ -649,7 +649,6 @@ PicoClaw stocke les données dans votre workspace configuré (par défaut : `~/.
|
||||||
├── HEARTBEAT.md # Invites de tâches périodiques (vérifiées toutes les 30 min)
|
├── HEARTBEAT.md # Invites de tâches périodiques (vérifiées toutes les 30 min)
|
||||||
├── IDENTITY.md # Identité de l'Agent
|
├── IDENTITY.md # Identité de l'Agent
|
||||||
├── SOUL.md # Âme de l'Agent
|
├── SOUL.md # Âme de l'Agent
|
||||||
├── TOOLS.md # Description des outils
|
|
||||||
└── USER.md # Préférences utilisateur
|
└── USER.md # Préférences utilisateur
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -980,6 +979,7 @@ Cette conception permet également le **support multi-agent** avec une sélectio
|
||||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Obtenir Clé](https://cerebras.ai) |
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Obtenir Clé](https://cerebras.ai) |
|
||||||
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Obtenir Clé](https://console.volcengine.com) |
|
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Obtenir Clé](https://console.volcengine.com) |
|
||||||
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [Obtenir une clé](https://longcat.chat/platform) |
|
||||||
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth uniquement |
|
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth uniquement |
|
||||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -610,7 +610,6 @@ PicoClaw は設定されたワークスペース(デフォルト: `~/.picoclaw
|
||||||
├── HEARTBEAT.md # 定期タスクプロンプト(30分ごとに確認)
|
├── HEARTBEAT.md # 定期タスクプロンプト(30分ごとに確認)
|
||||||
├── IDENTITY.md # エージェントのアイデンティティ
|
├── IDENTITY.md # エージェントのアイデンティティ
|
||||||
├── SOUL.md # エージェントのソウル
|
├── SOUL.md # エージェントのソウル
|
||||||
├── TOOLS.md # ツールの説明
|
|
||||||
└── USER.md # ユーザー設定
|
└── USER.md # ユーザー設定
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -921,6 +920,7 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [キーを取得](https://cerebras.ai) |
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [キーを取得](https://cerebras.ai) |
|
||||||
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [キーを取得](https://console.volcengine.com) |
|
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [キーを取得](https://console.volcengine.com) |
|
||||||
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [キーを取得](https://longcat.chat/platform) |
|
||||||
| **Antigravity** | `antigravity/` | Google Cloud | カスタム | OAuthのみ |
|
| **Antigravity** | `antigravity/` | Google Cloud | カスタム | OAuthのみ |
|
||||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
|
|
||||||
16
README.md
16
README.md
|
|
@ -194,6 +194,19 @@ docker compose -f docker/docker-compose.yml logs -f picoclaw-gateway
|
||||||
docker compose -f docker/docker-compose.yml --profile gateway down
|
docker compose -f docker/docker-compose.yml --profile gateway down
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Launcher Mode (Web Console)
|
||||||
|
|
||||||
|
The `launcher` image includes all three binaries (`picoclaw`, `picoclaw-launcher`, `picoclaw-launcher-tui`) and starts the web console by default, which provides a browser-based UI for configuration and chat.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker compose -f docker/docker-compose.yml --profile launcher up -d
|
||||||
|
```
|
||||||
|
|
||||||
|
Open http://localhost:18800 in your browser. The launcher manages the gateway process automatically.
|
||||||
|
|
||||||
|
> [!WARNING]
|
||||||
|
> The web console does not yet support authentication. Avoid exposing it to the public internet.
|
||||||
|
|
||||||
### Agent Mode (One-shot)
|
### Agent Mode (One-shot)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
@ -774,7 +787,6 @@ PicoClaw stores data in your configured workspace (default: `~/.picoclaw/workspa
|
||||||
├── HEARTBEAT.md # Periodic task prompts (checked every 30 min)
|
├── HEARTBEAT.md # Periodic task prompts (checked every 30 min)
|
||||||
├── IDENTITY.md # Agent identity
|
├── IDENTITY.md # Agent identity
|
||||||
├── SOUL.md # Agent soul
|
├── SOUL.md # Agent soul
|
||||||
├── TOOLS.md # Tool descriptions
|
|
||||||
└── USER.md # User preferences
|
└── USER.md # User preferences
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -1021,6 +1033,7 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://console.volcengine.com) |
|
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://console.volcengine.com) |
|
||||||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
| **Vivgrid** | `vivgrid/` | `https://api.vivgrid.com/v1` | OpenAI | [Get Key](https://vivgrid.com) |
|
| **Vivgrid** | `vivgrid/` | `https://api.vivgrid.com/v1` | OpenAI | [Get Key](https://vivgrid.com) |
|
||||||
|
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [Get Key](https://longcat.chat/platform) |
|
||||||
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth only |
|
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth only |
|
||||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
|
@ -1491,3 +1504,4 @@ This happens when another instance of the bot is running. Make sure only one `pi
|
||||||
| **SearXNG** | Unlimited (self-hosted) | Privacy-focused metasearch (70+ engines) |
|
| **SearXNG** | Unlimited (self-hosted) | Privacy-focused metasearch (70+ engines) |
|
||||||
| **Groq** | Free tier available | Fast inference (Llama, Mixtral) |
|
| **Groq** | Free tier available | Fast inference (Llama, Mixtral) |
|
||||||
| **Cerebras** | Free tier available | Fast inference (Llama, Qwen, etc.) |
|
| **Cerebras** | Free tier available | Fast inference (Llama, Qwen, etc.) |
|
||||||
|
| **LongCat** | Up to 5M tokens/day | Fast inference (free tier) |
|
||||||
|
|
|
||||||
|
|
@ -645,7 +645,6 @@ O PicoClaw armazena dados no workspace configurado (padrão: `~/.picoclaw/worksp
|
||||||
├── HEARTBEAT.md # Prompts de tarefas periodicas (verificado a cada 30 min)
|
├── HEARTBEAT.md # Prompts de tarefas periodicas (verificado a cada 30 min)
|
||||||
├── IDENTITY.md # Identidade do Agente
|
├── IDENTITY.md # Identidade do Agente
|
||||||
├── SOUL.md # Alma do Agente
|
├── SOUL.md # Alma do Agente
|
||||||
├── TOOLS.md # Descrição das ferramentas
|
|
||||||
└── USER.md # Preferencias do usuario
|
└── USER.md # Preferencias do usuario
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -976,6 +975,7 @@ Este design também possibilita o **suporte multi-agent** com seleção flexíve
|
||||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Obter Chave](https://cerebras.ai) |
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Obter Chave](https://cerebras.ai) |
|
||||||
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Obter Chave](https://console.volcengine.com) |
|
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Obter Chave](https://console.volcengine.com) |
|
||||||
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [Obter Chave](https://longcat.chat/platform) |
|
||||||
| **Antigravity** | `antigravity/` | Google Cloud | Custom | Apenas OAuth |
|
| **Antigravity** | `antigravity/` | Google Cloud | Custom | Apenas OAuth |
|
||||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -617,7 +617,6 @@ PicoClaw lưu trữ dữ liệu trong workspace đã cấu hình (mặc định:
|
||||||
├── HEARTBEAT.md # Prompt tác vụ định kỳ (kiểm tra mỗi 30 phút)
|
├── HEARTBEAT.md # Prompt tác vụ định kỳ (kiểm tra mỗi 30 phút)
|
||||||
├── IDENTITY.md # Danh tính Agent
|
├── IDENTITY.md # Danh tính Agent
|
||||||
├── SOUL.md # Tâm hồn/Tính cách Agent
|
├── SOUL.md # Tâm hồn/Tính cách Agent
|
||||||
├── TOOLS.md # Mô tả công cụ
|
|
||||||
└── USER.md # Tùy chọn người dùng
|
└── USER.md # Tùy chọn người dùng
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -945,6 +944,7 @@ Thiết kế này cũng cho phép **hỗ trợ đa tác nhân** với lựa ch
|
||||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Lấy Khóa](https://cerebras.ai) |
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Lấy Khóa](https://cerebras.ai) |
|
||||||
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Lấy Khóa](https://console.volcengine.com) |
|
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Lấy Khóa](https://console.volcengine.com) |
|
||||||
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [Lấy Key](https://longcat.chat/platform) |
|
||||||
| **Antigravity** | `antigravity/` | Google Cloud | Tùy chỉnh | Chỉ OAuth |
|
| **Antigravity** | `antigravity/` | Google Cloud | Tùy chỉnh | Chỉ OAuth |
|
||||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -365,7 +365,6 @@ PicoClaw 将数据存储在您配置的工作区中(默认:`~/.picoclaw/work
|
||||||
├── HEARTBEAT.md # 周期性任务提示词 (每 30 分钟检查一次)
|
├── HEARTBEAT.md # 周期性任务提示词 (每 30 分钟检查一次)
|
||||||
├── IDENTITY.md # Agent 身份设定
|
├── IDENTITY.md # Agent 身份设定
|
||||||
├── SOUL.md # Agent 灵魂/性格
|
├── SOUL.md # Agent 灵魂/性格
|
||||||
├── TOOLS.md # 工具描述
|
|
||||||
└── USER.md # 用户偏好
|
└── USER.md # 用户偏好
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
@ -517,6 +516,7 @@ Agent 读取 HEARTBEAT.md
|
||||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [获取密钥](https://cerebras.ai) |
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [获取密钥](https://cerebras.ai) |
|
||||||
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [获取密钥](https://console.volcengine.com) |
|
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [获取密钥](https://console.volcengine.com) |
|
||||||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [获取密钥](https://longcat.chat/platform) |
|
||||||
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
|
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
|
||||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
|
@ -879,3 +879,4 @@ Discord: [https://discord.gg/V4sAZ9XWpN](https://discord.gg/V4sAZ9XWpN)
|
||||||
| **Brave Search** | 2000 次查询/月 | 网络搜索功能 |
|
| **Brave Search** | 2000 次查询/月 | 网络搜索功能 |
|
||||||
| **Tavily** | 1000 次查询/月 | AI Agent 搜索优化 |
|
| **Tavily** | 1000 次查询/月 | AI Agent 搜索优化 |
|
||||||
| **Groq** | 提供免费层级 | 极速推理 (Llama, Mixtral) |
|
| **Groq** | 提供免费层级 | 极速推理 (Llama, Mixtral) |
|
||||||
|
| **LongCat** | 最多 5M tokens/天 | 推理速度快 (免费额度) |
|
||||||
|
|
|
||||||
Binary file not shown.
|
Before Width: | Height: | Size: 348 KiB After Width: | Height: | Size: 345 KiB |
|
|
@ -1,23 +1,42 @@
|
||||||
package gateway
|
package gateway
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewGatewayCommand() *cobra.Command {
|
func NewGatewayCommand() *cobra.Command {
|
||||||
var debug bool
|
var debug bool
|
||||||
|
var noTruncate bool
|
||||||
|
|
||||||
cmd := &cobra.Command{
|
cmd := &cobra.Command{
|
||||||
Use: "gateway",
|
Use: "gateway",
|
||||||
Aliases: []string{"g"},
|
Aliases: []string{"g"},
|
||||||
Short: "Start picoclaw gateway",
|
Short: "Start picoclaw gateway",
|
||||||
Args: cobra.NoArgs,
|
Args: cobra.NoArgs,
|
||||||
|
PreRunE: func(_ *cobra.Command, _ []string) error {
|
||||||
|
if noTruncate && !debug {
|
||||||
|
return fmt.Errorf("the --no-truncate option can only be used in conjunction with --debug (-d)")
|
||||||
|
}
|
||||||
|
|
||||||
|
if noTruncate {
|
||||||
|
utils.SetDisableTruncation(true)
|
||||||
|
logger.Info("String truncation is globally disabled via 'no-truncate' flag")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
RunE: func(_ *cobra.Command, _ []string) error {
|
RunE: func(_ *cobra.Command, _ []string) error {
|
||||||
return gatewayCmd(debug)
|
return gatewayCmd(debug)
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd.Flags().BoolVarP(&debug, "debug", "d", false, "Enable debug logging")
|
cmd.Flags().BoolVarP(&debug, "debug", "d", false, "Enable debug logging")
|
||||||
|
cmd.Flags().BoolVarP(&noTruncate, "no-truncate", "T", false, "Disable string truncation in debug logs")
|
||||||
|
|
||||||
return cmd
|
return cmd
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,23 +1,14 @@
|
||||||
package internal
|
package internal
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
const Logo = "🦞"
|
const Logo = "🦞"
|
||||||
|
|
||||||
var (
|
|
||||||
version = "dev"
|
|
||||||
gitCommit string
|
|
||||||
buildTime string
|
|
||||||
goVersion string
|
|
||||||
)
|
|
||||||
|
|
||||||
// GetPicoclawHome returns the picoclaw home directory.
|
// GetPicoclawHome returns the picoclaw home directory.
|
||||||
// Priority: $PICOCLAW_HOME > ~/.picoclaw
|
// Priority: $PICOCLAW_HOME > ~/.picoclaw
|
||||||
func GetPicoclawHome() string {
|
func GetPicoclawHome() string {
|
||||||
|
|
@ -40,25 +31,19 @@ func LoadConfig() (*config.Config, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// FormatVersion returns the version string with optional git commit
|
// FormatVersion returns the version string with optional git commit
|
||||||
|
// Deprecated: Use pkg/config.FormatVersion instead
|
||||||
func FormatVersion() string {
|
func FormatVersion() string {
|
||||||
v := version
|
return config.FormatVersion()
|
||||||
if gitCommit != "" {
|
|
||||||
v += fmt.Sprintf(" (git: %s)", gitCommit)
|
|
||||||
}
|
|
||||||
return v
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// FormatBuildInfo returns build time and go version info
|
// FormatBuildInfo returns build time and go version info
|
||||||
|
// Deprecated: Use pkg/config.FormatBuildInfo instead
|
||||||
func FormatBuildInfo() (string, string) {
|
func FormatBuildInfo() (string, string) {
|
||||||
build := buildTime
|
return config.FormatBuildInfo()
|
||||||
goVer := goVersion
|
|
||||||
if goVer == "" {
|
|
||||||
goVer = runtime.Version()
|
|
||||||
}
|
|
||||||
return build, goVer
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetVersion returns the version string
|
// GetVersion returns the version string
|
||||||
|
// Deprecated: Use pkg/config.GetVersion instead
|
||||||
func GetVersion() string {
|
func GetVersion() string {
|
||||||
return version
|
return config.GetVersion()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -40,65 +40,6 @@ func TestGetConfigPath_WithPICOCLAW_CONFIG(t *testing.T) {
|
||||||
assert.Equal(t, want, got)
|
assert.Equal(t, want, got)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFormatVersion_NoGitCommit(t *testing.T) {
|
|
||||||
oldVersion, oldGit := version, gitCommit
|
|
||||||
t.Cleanup(func() { version, gitCommit = oldVersion, oldGit })
|
|
||||||
|
|
||||||
version = "1.2.3"
|
|
||||||
gitCommit = ""
|
|
||||||
|
|
||||||
assert.Equal(t, "1.2.3", FormatVersion())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatVersion_WithGitCommit(t *testing.T) {
|
|
||||||
oldVersion, oldGit := version, gitCommit
|
|
||||||
t.Cleanup(func() { version, gitCommit = oldVersion, oldGit })
|
|
||||||
|
|
||||||
version = "1.2.3"
|
|
||||||
gitCommit = "abc123"
|
|
||||||
|
|
||||||
assert.Equal(t, "1.2.3 (git: abc123)", FormatVersion())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatBuildInfo_UsesBuildTimeAndGoVersion_WhenSet(t *testing.T) {
|
|
||||||
oldBuildTime, oldGoVersion := buildTime, goVersion
|
|
||||||
t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion })
|
|
||||||
|
|
||||||
buildTime = "2026-02-20T00:00:00Z"
|
|
||||||
goVersion = "go1.23.0"
|
|
||||||
|
|
||||||
build, goVer := FormatBuildInfo()
|
|
||||||
|
|
||||||
assert.Equal(t, buildTime, build)
|
|
||||||
assert.Equal(t, goVersion, goVer)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatBuildInfo_EmptyBuildTime_ReturnsEmptyBuild(t *testing.T) {
|
|
||||||
oldBuildTime, oldGoVersion := buildTime, goVersion
|
|
||||||
t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion })
|
|
||||||
|
|
||||||
buildTime = ""
|
|
||||||
goVersion = "go1.23.0"
|
|
||||||
|
|
||||||
build, goVer := FormatBuildInfo()
|
|
||||||
|
|
||||||
assert.Empty(t, build)
|
|
||||||
assert.Equal(t, goVersion, goVer)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatBuildInfo_EmptyGoVersion_FallsBackToRuntimeVersion(t *testing.T) {
|
|
||||||
oldBuildTime, oldGoVersion := buildTime, goVersion
|
|
||||||
t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion })
|
|
||||||
|
|
||||||
buildTime = "x"
|
|
||||||
goVersion = ""
|
|
||||||
|
|
||||||
build, goVer := FormatBuildInfo()
|
|
||||||
|
|
||||||
assert.Equal(t, "x", build)
|
|
||||||
assert.Equal(t, runtime.Version(), goVer)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetConfigPath_Windows(t *testing.T) {
|
func TestGetConfigPath_Windows(t *testing.T) {
|
||||||
if runtime.GOOS != "windows" {
|
if runtime.GOOS != "windows" {
|
||||||
t.Skip("windows-specific HOME behavior varies; run on windows")
|
t.Skip("windows-specific HOME behavior varies; run on windows")
|
||||||
|
|
@ -112,17 +53,3 @@ func TestGetConfigPath_Windows(t *testing.T) {
|
||||||
|
|
||||||
require.True(t, strings.EqualFold(got, want), "GetConfigPath() = %q, want %q", got, want)
|
require.True(t, strings.EqualFold(got, want), "GetConfigPath() = %q, want %q", got, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGetVersion(t *testing.T) {
|
|
||||||
assert.Equal(t, "dev", GetVersion())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetConfigPath_WithEnv(t *testing.T) {
|
|
||||||
t.Setenv("PICOCLAW_CONFIG", "/tmp/custom/config.json")
|
|
||||||
t.Setenv("HOME", "/tmp/home") // Also set home to ensure env is preferred
|
|
||||||
|
|
||||||
got := GetConfigPath()
|
|
||||||
want := "/tmp/custom/config.json"
|
|
||||||
|
|
||||||
assert.Equal(t, want, got)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func statusCmd() {
|
func statusCmd() {
|
||||||
|
|
@ -18,8 +19,8 @@ func statusCmd() {
|
||||||
configPath := internal.GetConfigPath()
|
configPath := internal.GetConfigPath()
|
||||||
|
|
||||||
fmt.Printf("%s picoclaw Status\n", internal.Logo)
|
fmt.Printf("%s picoclaw Status\n", internal.Logo)
|
||||||
fmt.Printf("Version: %s\n", internal.FormatVersion())
|
fmt.Printf("Version: %s\n", config.FormatVersion())
|
||||||
build, _ := internal.FormatBuildInfo()
|
build, _ := config.FormatBuildInfo()
|
||||||
if build != "" {
|
if build != "" {
|
||||||
fmt.Printf("Build: %s\n", build)
|
fmt.Printf("Build: %s\n", build)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewVersionCommand() *cobra.Command {
|
func NewVersionCommand() *cobra.Command {
|
||||||
|
|
@ -22,8 +23,8 @@ func NewVersionCommand() *cobra.Command {
|
||||||
}
|
}
|
||||||
|
|
||||||
func printVersion() {
|
func printVersion() {
|
||||||
fmt.Printf("%s picoclaw %s\n", internal.Logo, internal.FormatVersion())
|
fmt.Printf("%s picoclaw %s\n", internal.Logo, config.FormatVersion())
|
||||||
build, goVer := internal.FormatBuildInfo()
|
build, goVer := config.FormatBuildInfo()
|
||||||
if build != "" {
|
if build != "" {
|
||||||
fmt.Printf(" Build: %s\n", build)
|
fmt.Printf(" Build: %s\n", build)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -22,10 +22,11 @@ import (
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/skills"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/skills"
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/status"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/status"
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/version"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/version"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewPicoclawCommand() *cobra.Command {
|
func NewPicoclawCommand() *cobra.Command {
|
||||||
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, internal.GetVersion())
|
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, config.GetVersion())
|
||||||
|
|
||||||
cmd := &cobra.Command{
|
cmd := &cobra.Command{
|
||||||
Use: "picoclaw",
|
Use: "picoclaw",
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ import (
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewPicoclawCommand(t *testing.T) {
|
func TestNewPicoclawCommand(t *testing.T) {
|
||||||
|
|
@ -16,7 +17,7 @@ func TestNewPicoclawCommand(t *testing.T) {
|
||||||
|
|
||||||
require.NotNil(t, cmd)
|
require.NotNil(t, cmd)
|
||||||
|
|
||||||
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, internal.GetVersion())
|
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, config.GetVersion())
|
||||||
|
|
||||||
assert.Equal(t, "picoclaw", cmd.Use)
|
assert.Equal(t, "picoclaw", cmd.Use)
|
||||||
assert.Equal(t, short, cmd.Short)
|
assert.Equal(t, short, cmd.Short)
|
||||||
|
|
|
||||||
|
|
@ -35,6 +35,11 @@
|
||||||
"model": "deepseek/deepseek-chat",
|
"model": "deepseek/deepseek-chat",
|
||||||
"api_key": "sk-your-deepseek-key"
|
"api_key": "sk-your-deepseek-key"
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"model_name": "longcat",
|
||||||
|
"model": "longcat/LongCat-Flash-Thinking",
|
||||||
|
"api_key": "your-longcat-api-key"
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"model_name": "loadbalanced-gpt4",
|
"model_name": "loadbalanced-gpt4",
|
||||||
"model": "openai/gpt-5.2",
|
"model": "openai/gpt-5.2",
|
||||||
|
|
@ -274,6 +279,10 @@
|
||||||
"avian": {
|
"avian": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": "https://api.avian.io/v1"
|
"api_base": "https://api.avian.io/v1"
|
||||||
|
},
|
||||||
|
"longcat": {
|
||||||
|
"api_key": "",
|
||||||
|
"api_base": "https://api.longcat.chat/openai"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"tools": {
|
"tools": {
|
||||||
|
|
@ -284,6 +293,9 @@
|
||||||
"brave": {
|
"brave": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "YOUR_BRAVE_API_KEY",
|
"api_key": "YOUR_BRAVE_API_KEY",
|
||||||
|
"api_keys": [
|
||||||
|
"YOUR_BRAVE_API_KEY"
|
||||||
|
],
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
},
|
},
|
||||||
"tavily": {
|
"tavily": {
|
||||||
|
|
@ -298,7 +310,10 @@
|
||||||
},
|
},
|
||||||
"perplexity": {
|
"perplexity": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "",
|
"api_key": "pplx-xxx",
|
||||||
|
"api_keys": [
|
||||||
|
"pplx-xxx"
|
||||||
|
],
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
},
|
},
|
||||||
"searxng": {
|
"searxng": {
|
||||||
|
|
@ -471,6 +486,9 @@
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"monitor_usb": true
|
"monitor_usb": true
|
||||||
},
|
},
|
||||||
|
"voice": {
|
||||||
|
"echo_transcription": false
|
||||||
|
},
|
||||||
"gateway": {
|
"gateway": {
|
||||||
"host": "127.0.0.1",
|
"host": "127.0.0.1",
|
||||||
"port": 18790
|
"port": 18790
|
||||||
|
|
|
||||||
12
docker/Dockerfile.goreleaser.launcher
Normal file
12
docker/Dockerfile.goreleaser.launcher
Normal file
|
|
@ -0,0 +1,12 @@
|
||||||
|
FROM alpine:3.21
|
||||||
|
|
||||||
|
ARG TARGETPLATFORM
|
||||||
|
|
||||||
|
RUN apk add --no-cache ca-certificates tzdata
|
||||||
|
|
||||||
|
COPY $TARGETPLATFORM/picoclaw /usr/local/bin/picoclaw
|
||||||
|
COPY $TARGETPLATFORM/picoclaw-launcher /usr/local/bin/picoclaw-launcher
|
||||||
|
COPY $TARGETPLATFORM/picoclaw-launcher-tui /usr/local/bin/picoclaw-launcher-tui
|
||||||
|
|
||||||
|
ENTRYPOINT ["picoclaw-launcher"]
|
||||||
|
CMD ["-public", "-no-browser"]
|
||||||
|
|
@ -19,7 +19,7 @@ services:
|
||||||
|
|
||||||
# ─────────────────────────────────────────────
|
# ─────────────────────────────────────────────
|
||||||
# PicoClaw Gateway (Long-running Bot)
|
# PicoClaw Gateway (Long-running Bot)
|
||||||
# docker compose -f docker/docker-compose.yml up picoclaw-gateway
|
# docker compose -f docker/docker-compose.yml --profile gateway up
|
||||||
# ─────────────────────────────────────────────
|
# ─────────────────────────────────────────────
|
||||||
picoclaw-gateway:
|
picoclaw-gateway:
|
||||||
image: docker.io/sipeed/picoclaw:latest
|
image: docker.io/sipeed/picoclaw:latest
|
||||||
|
|
@ -32,3 +32,21 @@ services:
|
||||||
# - "host.docker.internal:host-gateway"
|
# - "host.docker.internal:host-gateway"
|
||||||
volumes:
|
volumes:
|
||||||
- ./data:/root/.picoclaw
|
- ./data:/root/.picoclaw
|
||||||
|
|
||||||
|
# ─────────────────────────────────────────────
|
||||||
|
# PicoClaw Launcher (Web Console + Gateway)
|
||||||
|
# docker compose -f docker/docker-compose.yml --profile launcher up
|
||||||
|
# ─────────────────────────────────────────────
|
||||||
|
picoclaw-launcher:
|
||||||
|
image: docker.io/sipeed/picoclaw:launcher
|
||||||
|
container_name: picoclaw-launcher
|
||||||
|
restart: on-failure
|
||||||
|
profiles:
|
||||||
|
- launcher
|
||||||
|
environment:
|
||||||
|
- PICOCLAW_GATEWAY_HOST=0.0.0.0
|
||||||
|
ports:
|
||||||
|
- "127.0.0.1:18800:18800"
|
||||||
|
- "127.0.0.1:18790:18790"
|
||||||
|
volumes:
|
||||||
|
- ./data:/root/.picoclaw
|
||||||
|
|
|
||||||
|
|
@ -22,7 +22,8 @@ Add this to `config.json`:
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"text": "Thinking..."
|
"text": "Thinking..."
|
||||||
},
|
},
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": "",
|
||||||
|
"message_format": "richtext"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -42,10 +43,12 @@ Add this to `config.json`:
|
||||||
| group_trigger | object | No | Group trigger strategy (`mention_only` / `prefixes`) |
|
| group_trigger | object | No | Group trigger strategy (`mention_only` / `prefixes`) |
|
||||||
| placeholder | object | No | Placeholder message config |
|
| placeholder | object | No | Placeholder message config |
|
||||||
| reasoning_channel_id | string | No | Target channel for reasoning output |
|
| reasoning_channel_id | string | No | Target channel for reasoning output |
|
||||||
|
| message_format | string | No | Output format: `"richtext"` (default) renders markdown as HTML; `"plain"` sends plain text only |
|
||||||
|
|
||||||
## 3. Currently Supported
|
## 3. Currently Supported
|
||||||
|
|
||||||
- Text message send/receive
|
- Text message send/receive with markdown rendering (bold, italic, headers, code blocks, etc.)
|
||||||
|
- Configurable message format (`richtext` / `plain`)
|
||||||
- Incoming image/audio/video/file download (MediaStore first, local path fallback)
|
- Incoming image/audio/video/file download (MediaStore first, local path fallback)
|
||||||
- Incoming audio normalization into existing transcription flow (`[audio: ...]`)
|
- Incoming audio normalization into existing transcription flow (`[audio: ...]`)
|
||||||
- Outgoing image/audio/video/file upload and send
|
- Outgoing image/audio/video/file upload and send
|
||||||
|
|
|
||||||
33
docs/debug.md
Normal file
33
docs/debug.md
Normal file
|
|
@ -0,0 +1,33 @@
|
||||||
|
# Debugging PicoClaw
|
||||||
|
|
||||||
|
PicoClaw performs multiple complex interactions under the hood for every single request it receives—from routing messages and evaluating complexity, to executing tools and adapting to model failures. Being able to see exactly what is happening is crucial, not just for troubleshooting potential issues, but also for truly understanding how the agent operates.
|
||||||
|
## Starting PicoClaw in Debug Mode
|
||||||
|
|
||||||
|
To get detailed information about what the agent is doing (LLM requests, tool calls, message routing), you can start the PicoClaw gateway with the debug flag:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw gateway --debug
|
||||||
|
# or
|
||||||
|
picoclaw gateway -d
|
||||||
|
```
|
||||||
|
|
||||||
|
In this mode, the system will format the logs extensively and display previews of system prompts and tool execution results.
|
||||||
|
|
||||||
|
## Disabling Log Truncation (Full Logs)
|
||||||
|
|
||||||
|
By default, PicoClaw truncates very long strings (such as the *System Prompt* or large JSON output results) in the debug logs to keep the console readable.
|
||||||
|
|
||||||
|
If you need to inspect the complete output of a command or the exact payload sent to the LLM model, you can use the `--no-truncate` flag.
|
||||||
|
|
||||||
|
**Note:** This flag *only* works when combined with the `--debug` mode.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw gateway --debug --no-truncate
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
When this flag is active, the global truncation function is disabled. This is extremely useful for:
|
||||||
|
|
||||||
|
* Verifying the exact syntax of the messages sent to the provider.
|
||||||
|
* Reading the complete output of tools like `exec`, `web_fetch`, or `read_file`.
|
||||||
|
* Debugging the session history saved in memory.
|
||||||
3
go.mod
3
go.mod
|
|
@ -11,6 +11,7 @@ require (
|
||||||
github.com/ergochat/irc-go v0.5.0
|
github.com/ergochat/irc-go v0.5.0
|
||||||
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/gomarkdown/markdown v0.0.0-20260217112301-37c66b85d6ab
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
github.com/h2non/filetype v1.1.3
|
github.com/h2non/filetype v1.1.3
|
||||||
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
||||||
|
|
@ -20,6 +21,7 @@ require (
|
||||||
github.com/open-dingtalk/dingtalk-stream-sdk-go v0.9.1
|
github.com/open-dingtalk/dingtalk-stream-sdk-go v0.9.1
|
||||||
github.com/openai/openai-go/v3 v3.22.0
|
github.com/openai/openai-go/v3 v3.22.0
|
||||||
github.com/rivo/tview v0.42.0
|
github.com/rivo/tview v0.42.0
|
||||||
|
github.com/rs/zerolog v1.34.0
|
||||||
github.com/slack-go/slack v0.17.3
|
github.com/slack-go/slack v0.17.3
|
||||||
github.com/spf13/cobra v1.10.2
|
github.com/spf13/cobra v1.10.2
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
|
|
@ -49,7 +51,6 @@ require (
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||||
github.com/rivo/uniseg v0.4.7 // indirect
|
github.com/rivo/uniseg v0.4.7 // indirect
|
||||||
github.com/rs/zerolog v1.34.0 // indirect
|
|
||||||
github.com/segmentio/asm v1.1.3 // indirect
|
github.com/segmentio/asm v1.1.3 // indirect
|
||||||
github.com/segmentio/encoding v0.5.3 // indirect
|
github.com/segmentio/encoding v0.5.3 // indirect
|
||||||
github.com/spf13/pflag v1.0.10 // indirect
|
github.com/spf13/pflag v1.0.10 // indirect
|
||||||
|
|
|
||||||
2
go.sum
2
go.sum
|
|
@ -79,6 +79,8 @@ github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvq
|
||||||
github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
|
github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
|
||||||
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
|
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
|
||||||
github.com/golang/protobuf v1.5.2/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
|
github.com/golang/protobuf v1.5.2/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
|
||||||
|
github.com/gomarkdown/markdown v0.0.0-20260217112301-37c66b85d6ab h1:VYNivV7P8IRHUam2swVUNkhIdp0LRRFKe4hXNnoZKTc=
|
||||||
|
github.com/gomarkdown/markdown v0.0.0-20260217112301-37c66b85d6ab/go.mod h1:JDGcbDT52eL4fju3sZ4TeHGsQwhG9nbDV21aMyhwPoA=
|
||||||
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
||||||
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
||||||
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||||
|
|
|
||||||
|
|
@ -12,9 +12,11 @@ import (
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/skills"
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ContextBuilder struct {
|
type ContextBuilder struct {
|
||||||
|
|
@ -80,8 +82,10 @@ func NewContextBuilder(workspace string) *ContextBuilder {
|
||||||
func (cb *ContextBuilder) getIdentity() string {
|
func (cb *ContextBuilder) getIdentity() string {
|
||||||
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
||||||
toolDiscovery := cb.getDiscoveryRule()
|
toolDiscovery := cb.getDiscoveryRule()
|
||||||
|
version := config.FormatVersion()
|
||||||
|
|
||||||
return fmt.Sprintf(`# picoclaw 🦞
|
return fmt.Sprintf(
|
||||||
|
`# picoclaw 🦞 (%s)
|
||||||
|
|
||||||
You are picoclaw, a helpful AI assistant.
|
You are picoclaw, a helpful AI assistant.
|
||||||
|
|
||||||
|
|
@ -102,7 +106,7 @@ Your workspace is at: %s
|
||||||
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.
|
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.
|
||||||
|
|
||||||
%s`,
|
%s`,
|
||||||
workspacePath, workspacePath, workspacePath, workspacePath, workspacePath, toolDiscovery)
|
version, workspacePath, workspacePath, workspacePath, workspacePath, workspacePath, toolDiscovery)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) getDiscoveryRule() string {
|
func (cb *ContextBuilder) getDiscoveryRule() string {
|
||||||
|
|
@ -535,10 +539,7 @@ func (cb *ContextBuilder) BuildMessages(
|
||||||
})
|
})
|
||||||
|
|
||||||
// Log preview of system prompt (avoid logging huge content)
|
// Log preview of system prompt (avoid logging huge content)
|
||||||
preview := fullSystemPrompt
|
preview := utils.Truncate(fullSystemPrompt, 500)
|
||||||
if len(preview) > 500 {
|
|
||||||
preview = preview[:500] + "... (truncated)"
|
|
||||||
}
|
|
||||||
logger.DebugCF("agent", "System prompt preview",
|
logger.DebugCF("agent", "System prompt preview",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"preview": preview,
|
"preview": preview,
|
||||||
|
|
|
||||||
|
|
@ -25,7 +25,6 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/constants"
|
"github.com/sipeed/picoclaw/pkg/constants"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/mcp"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/media"
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/routing"
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
|
|
@ -48,6 +47,7 @@ type AgentLoop struct {
|
||||||
mediaStore media.MediaStore
|
mediaStore media.MediaStore
|
||||||
transcriber voice.Transcriber
|
transcriber voice.Transcriber
|
||||||
cmdRegistry *commands.Registry
|
cmdRegistry *commands.Registry
|
||||||
|
mcp mcpRuntime
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -121,19 +121,21 @@ func registerSharedTools(
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Web tools
|
|
||||||
if cfg.Tools.IsToolEnabled("web") {
|
if cfg.Tools.IsToolEnabled("web") {
|
||||||
searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
|
searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
|
||||||
BraveAPIKey: cfg.Tools.Web.Brave.APIKey,
|
BraveAPIKeys: config.MergeAPIKeys(cfg.Tools.Web.Brave.APIKey, cfg.Tools.Web.Brave.APIKeys),
|
||||||
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
||||||
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
||||||
TavilyAPIKey: cfg.Tools.Web.Tavily.APIKey,
|
TavilyAPIKeys: config.MergeAPIKeys(cfg.Tools.Web.Tavily.APIKey, cfg.Tools.Web.Tavily.APIKeys),
|
||||||
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
||||||
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
||||||
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
||||||
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
||||||
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
||||||
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
|
PerplexityAPIKeys: config.MergeAPIKeys(
|
||||||
|
cfg.Tools.Web.Perplexity.APIKey,
|
||||||
|
cfg.Tools.Web.Perplexity.APIKeys,
|
||||||
|
),
|
||||||
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
||||||
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
||||||
SearXNGBaseURL: cfg.Tools.Web.SearXNG.BaseURL,
|
SearXNGBaseURL: cfg.Tools.Web.SearXNG.BaseURL,
|
||||||
|
|
@ -239,123 +241,8 @@ func registerSharedTools(
|
||||||
func (al *AgentLoop) Run(ctx context.Context) error {
|
func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
al.running.Store(true)
|
al.running.Store(true)
|
||||||
|
|
||||||
// Initialize MCP servers for all agents
|
if err := al.ensureMCPInitialized(ctx); err != nil {
|
||||||
cfg := al.GetConfig()
|
return err
|
||||||
if cfg.Tools.IsToolEnabled("mcp") {
|
|
||||||
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.GetRegistry().GetDefaultAgent()
|
|
||||||
var workspacePath string
|
|
||||||
if defaultAgent != nil && defaultAgent.Workspace != "" {
|
|
||||||
workspacePath = defaultAgent.Workspace
|
|
||||||
} else {
|
|
||||||
workspacePath = cfg.WorkspacePath()
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := mcpManager.LoadFromMCPConfig(ctx, cfg.Tools.MCP, workspacePath); err != nil {
|
|
||||||
logger.WarnCF("agent", "Failed to load MCP servers, MCP tools will not be available",
|
|
||||||
map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
} else {
|
|
||||||
// Register MCP tools for all agents
|
|
||||||
servers := mcpManager.GetServers()
|
|
||||||
uniqueTools := 0
|
|
||||||
totalRegistrations := 0
|
|
||||||
registry := al.GetRegistry()
|
|
||||||
agentIDs := registry.ListAgentIDs()
|
|
||||||
agentCount := len(agentIDs)
|
|
||||||
|
|
||||||
for serverName, conn := range servers {
|
|
||||||
uniqueTools += len(conn.Tools)
|
|
||||||
for _, tool := range conn.Tools {
|
|
||||||
for _, agentID := range agentIDs {
|
|
||||||
agent, ok := registry.GetAgent(agentID)
|
|
||||||
if !ok {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
|
|
||||||
|
|
||||||
cfg := al.GetConfig()
|
|
||||||
if cfg.Tools.MCP.Discovery.Enabled {
|
|
||||||
agent.Tools.RegisterHidden(mcpTool)
|
|
||||||
} else {
|
|
||||||
agent.Tools.Register(mcpTool)
|
|
||||||
}
|
|
||||||
|
|
||||||
totalRegistrations++
|
|
||||||
logger.DebugCF("agent", "Registered MCP tool",
|
|
||||||
map[string]any{
|
|
||||||
"agent_id": agentID,
|
|
||||||
"server": serverName,
|
|
||||||
"tool": tool.Name,
|
|
||||||
"name": mcpTool.Name(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
logger.InfoCF("agent", "MCP tools registered successfully",
|
|
||||||
map[string]any{
|
|
||||||
"server_count": len(servers),
|
|
||||||
"unique_tools": uniqueTools,
|
|
||||||
"total_registrations": totalRegistrations,
|
|
||||||
"agent_count": agentCount,
|
|
||||||
})
|
|
||||||
|
|
||||||
// Initializes Discovery Tools only if enabled by configuration
|
|
||||||
cfg := al.GetConfig()
|
|
||||||
if cfg.Tools.MCP.Enabled && cfg.Tools.MCP.Discovery.Enabled {
|
|
||||||
useBM25 := cfg.Tools.MCP.Discovery.UseBM25
|
|
||||||
useRegex := cfg.Tools.MCP.Discovery.UseRegex
|
|
||||||
|
|
||||||
// Fail fast: If discovery is enabled but no search method is turned on
|
|
||||||
if !useBM25 && !useRegex {
|
|
||||||
return fmt.Errorf(
|
|
||||||
"tool discovery is enabled but neither 'use_bm25' nor 'use_regex' is set to true in the configuration",
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
ttl := cfg.Tools.MCP.Discovery.TTL
|
|
||||||
if ttl <= 0 {
|
|
||||||
ttl = 5 // Default value
|
|
||||||
}
|
|
||||||
|
|
||||||
maxSearchResults := cfg.Tools.MCP.Discovery.MaxSearchResults
|
|
||||||
if maxSearchResults <= 0 {
|
|
||||||
maxSearchResults = 5 // Default value
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.InfoCF("agent", "Initializing tool discovery", map[string]any{
|
|
||||||
"bm25": useBM25, "regex": useRegex, "ttl": ttl, "max_results": maxSearchResults,
|
|
||||||
})
|
|
||||||
|
|
||||||
registry := al.GetRegistry()
|
|
||||||
for _, agentID := range agentIDs {
|
|
||||||
agent, ok := registry.GetAgent(agentID)
|
|
||||||
if !ok {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if useRegex {
|
|
||||||
agent.Tools.Register(tools.NewRegexSearchTool(agent.Tools, ttl, maxSearchResults))
|
|
||||||
}
|
|
||||||
if useBM25 {
|
|
||||||
agent.Tools.Register(tools.NewBM25SearchTool(agent.Tools, ttl, maxSearchResults))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for al.running.Load() {
|
for al.running.Load() {
|
||||||
|
|
@ -435,6 +322,17 @@ func (al *AgentLoop) Stop() {
|
||||||
|
|
||||||
// Close releases resources held by agent session stores. Call after Stop.
|
// Close releases resources held by agent session stores. Call after Stop.
|
||||||
func (al *AgentLoop) Close() {
|
func (al *AgentLoop) Close() {
|
||||||
|
mcpManager := al.mcp.takeManager()
|
||||||
|
|
||||||
|
if mcpManager != nil {
|
||||||
|
if err := mcpManager.Close(); err != nil {
|
||||||
|
logger.ErrorCF("agent", "Failed to close MCP manager",
|
||||||
|
map[string]any{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
al.GetRegistry().Close()
|
al.GetRegistry().Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -583,9 +481,10 @@ var audioAnnotationRe = regexp.MustCompile(`\[(voice|audio)(?::[^\]]*)?\]`)
|
||||||
|
|
||||||
// transcribeAudioInMessage resolves audio media refs, transcribes them, and
|
// transcribeAudioInMessage resolves audio media refs, transcribes them, and
|
||||||
// replaces audio annotations in msg.Content with the transcribed text.
|
// replaces audio annotations in msg.Content with the transcribed text.
|
||||||
func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.InboundMessage) bus.InboundMessage {
|
// Returns the (possibly modified) message and true if audio was transcribed.
|
||||||
|
func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.InboundMessage) (bus.InboundMessage, bool) {
|
||||||
if al.transcriber == nil || al.mediaStore == nil || len(msg.Media) == 0 {
|
if al.transcriber == nil || al.mediaStore == nil || len(msg.Media) == 0 {
|
||||||
return msg
|
return msg, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Transcribe each audio media ref in order.
|
// Transcribe each audio media ref in order.
|
||||||
|
|
@ -609,9 +508,11 @@ func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.Inbou
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(transcriptions) == 0 {
|
if len(transcriptions) == 0 {
|
||||||
return msg
|
return msg, false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
al.sendTranscriptionFeedback(ctx, msg.Channel, msg.ChatID, msg.MessageID, transcriptions)
|
||||||
|
|
||||||
// Replace audio annotations sequentially with transcriptions.
|
// Replace audio annotations sequentially with transcriptions.
|
||||||
idx := 0
|
idx := 0
|
||||||
newContent := audioAnnotationRe.ReplaceAllStringFunc(msg.Content, func(match string) string {
|
newContent := audioAnnotationRe.ReplaceAllStringFunc(msg.Content, func(match string) string {
|
||||||
|
|
@ -629,7 +530,48 @@ func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.Inbou
|
||||||
}
|
}
|
||||||
|
|
||||||
msg.Content = newContent
|
msg.Content = newContent
|
||||||
return msg
|
return msg, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendTranscriptionFeedback sends feedback to the user with the result of
|
||||||
|
// audio transcription if the option is enabled. It uses Manager.SendMessage
|
||||||
|
// which executes synchronously (rate limiting, splitting, retry) so that
|
||||||
|
// ordering with the subsequent placeholder is guaranteed.
|
||||||
|
func (al *AgentLoop) sendTranscriptionFeedback(
|
||||||
|
ctx context.Context,
|
||||||
|
channel, chatID, messageID string,
|
||||||
|
validTexts []string,
|
||||||
|
) {
|
||||||
|
if !al.cfg.Voice.EchoTranscription {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if al.channelManager == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var nonEmpty []string
|
||||||
|
for _, t := range validTexts {
|
||||||
|
if t != "" {
|
||||||
|
nonEmpty = append(nonEmpty, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var feedbackMsg string
|
||||||
|
if len(nonEmpty) > 0 {
|
||||||
|
feedbackMsg = "Transcript: " + strings.Join(nonEmpty, "\n")
|
||||||
|
} else {
|
||||||
|
feedbackMsg = "No voice detected in the audio"
|
||||||
|
}
|
||||||
|
|
||||||
|
err := al.channelManager.SendMessage(ctx, bus.OutboundMessage{
|
||||||
|
Channel: channel,
|
||||||
|
ChatID: chatID,
|
||||||
|
Content: feedbackMsg,
|
||||||
|
ReplyToMessageID: messageID,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("voice", "Failed to send transcription feedback", map[string]any{"error": err.Error()})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// inferMediaType determines the media type ("image", "audio", "video", "file")
|
// inferMediaType determines the media type ("image", "audio", "video", "file")
|
||||||
|
|
@ -691,6 +633,10 @@ func (al *AgentLoop) ProcessDirectWithChannel(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
content, sessionKey, channel, chatID string,
|
content, sessionKey, channel, chatID string,
|
||||||
) (string, error) {
|
) (string, error) {
|
||||||
|
if err := al.ensureMCPInitialized(ctx); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
msg := bus.InboundMessage{
|
msg := bus.InboundMessage{
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
SenderID: "cron",
|
SenderID: "cron",
|
||||||
|
|
@ -743,7 +689,14 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
msg = al.transcribeAudioInMessage(ctx, msg)
|
var hadAudio bool
|
||||||
|
msg, hadAudio = al.transcribeAudioInMessage(ctx, msg)
|
||||||
|
|
||||||
|
// For audio messages the placeholder was deferred by the channel.
|
||||||
|
// Now that transcription (and optional feedback) is done, send it.
|
||||||
|
if hadAudio && al.channelManager != nil {
|
||||||
|
al.channelManager.SendPlaceholder(ctx, msg.Channel, msg.ChatID)
|
||||||
|
}
|
||||||
|
|
||||||
// Route system messages to processSystemMessage
|
// Route system messages to processSystemMessage
|
||||||
if msg.Channel == "system" {
|
if msg.Channel == "system" {
|
||||||
|
|
|
||||||
184
pkg/agent/loop_mcp.go
Normal file
184
pkg/agent/loop_mcp.go
Normal file
|
|
@ -0,0 +1,184 @@
|
||||||
|
// 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 (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/mcp"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mcpRuntime struct {
|
||||||
|
initOnce sync.Once
|
||||||
|
mu sync.Mutex
|
||||||
|
manager *mcp.Manager
|
||||||
|
initErr error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *mcpRuntime) setManager(manager *mcp.Manager) {
|
||||||
|
r.mu.Lock()
|
||||||
|
r.manager = manager
|
||||||
|
r.initErr = nil
|
||||||
|
r.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *mcpRuntime) setInitErr(err error) {
|
||||||
|
r.mu.Lock()
|
||||||
|
r.initErr = err
|
||||||
|
r.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *mcpRuntime) getInitErr() error {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return r.initErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *mcpRuntime) takeManager() *mcp.Manager {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
manager := r.manager
|
||||||
|
r.manager = nil
|
||||||
|
return manager
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *mcpRuntime) hasManager() bool {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return r.manager != nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureMCPInitialized loads MCP servers/tools once so both Run() and direct
|
||||||
|
// agent mode share the same initialization path.
|
||||||
|
func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
|
||||||
|
if !al.cfg.Tools.IsToolEnabled("mcp") {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
al.mcp.initOnce.Do(func() {
|
||||||
|
mcpManager := mcp.NewManager()
|
||||||
|
|
||||||
|
defaultAgent := al.registry.GetDefaultAgent()
|
||||||
|
workspacePath := al.cfg.WorkspacePath()
|
||||||
|
if defaultAgent != nil && defaultAgent.Workspace != "" {
|
||||||
|
workspacePath = defaultAgent.Workspace
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := mcpManager.LoadFromMCPConfig(ctx, al.cfg.Tools.MCP, workspacePath); err != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to load MCP servers, MCP tools will not be available",
|
||||||
|
map[string]any{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
if closeErr := mcpManager.Close(); closeErr != nil {
|
||||||
|
logger.ErrorCF("agent", "Failed to close MCP manager",
|
||||||
|
map[string]any{
|
||||||
|
"error": closeErr.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register MCP tools for all agents
|
||||||
|
servers := mcpManager.GetServers()
|
||||||
|
uniqueTools := 0
|
||||||
|
totalRegistrations := 0
|
||||||
|
agentIDs := al.registry.ListAgentIDs()
|
||||||
|
agentCount := len(agentIDs)
|
||||||
|
|
||||||
|
for serverName, conn := range servers {
|
||||||
|
uniqueTools += len(conn.Tools)
|
||||||
|
for _, tool := range conn.Tools {
|
||||||
|
for _, agentID := range agentIDs {
|
||||||
|
agent, ok := al.registry.GetAgent(agentID)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
|
||||||
|
|
||||||
|
if al.cfg.Tools.MCP.Discovery.Enabled {
|
||||||
|
agent.Tools.RegisterHidden(mcpTool)
|
||||||
|
} else {
|
||||||
|
agent.Tools.Register(mcpTool)
|
||||||
|
}
|
||||||
|
|
||||||
|
totalRegistrations++
|
||||||
|
logger.DebugCF("agent", "Registered MCP tool",
|
||||||
|
map[string]any{
|
||||||
|
"agent_id": agentID,
|
||||||
|
"server": serverName,
|
||||||
|
"tool": tool.Name,
|
||||||
|
"name": mcpTool.Name(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
logger.InfoCF("agent", "MCP tools registered successfully",
|
||||||
|
map[string]any{
|
||||||
|
"server_count": len(servers),
|
||||||
|
"unique_tools": uniqueTools,
|
||||||
|
"total_registrations": totalRegistrations,
|
||||||
|
"agent_count": agentCount,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Initializes Discovery Tools only if enabled by configuration
|
||||||
|
if al.cfg.Tools.MCP.Enabled && al.cfg.Tools.MCP.Discovery.Enabled {
|
||||||
|
useBM25 := al.cfg.Tools.MCP.Discovery.UseBM25
|
||||||
|
useRegex := al.cfg.Tools.MCP.Discovery.UseRegex
|
||||||
|
|
||||||
|
// Fail fast: If discovery is enabled but no search method is turned on
|
||||||
|
if !useBM25 && !useRegex {
|
||||||
|
al.mcp.setInitErr(fmt.Errorf(
|
||||||
|
"tool discovery is enabled but neither 'use_bm25' nor 'use_regex' is set to true in the configuration",
|
||||||
|
))
|
||||||
|
if closeErr := mcpManager.Close(); closeErr != nil {
|
||||||
|
logger.ErrorCF("agent", "Failed to close MCP manager",
|
||||||
|
map[string]any{
|
||||||
|
"error": closeErr.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ttl := al.cfg.Tools.MCP.Discovery.TTL
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = 5 // Default value
|
||||||
|
}
|
||||||
|
|
||||||
|
maxSearchResults := al.cfg.Tools.MCP.Discovery.MaxSearchResults
|
||||||
|
if maxSearchResults <= 0 {
|
||||||
|
maxSearchResults = 5 // Default value
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.InfoCF("agent", "Initializing tool discovery", map[string]any{
|
||||||
|
"bm25": useBM25, "regex": useRegex, "ttl": ttl, "max_results": maxSearchResults,
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, agentID := range agentIDs {
|
||||||
|
agent, ok := al.registry.GetAgent(agentID)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if useRegex {
|
||||||
|
agent.Tools.Register(tools.NewRegexSearchTool(agent.Tools, ttl, maxSearchResults))
|
||||||
|
}
|
||||||
|
if useBM25 {
|
||||||
|
agent.Tools.Register(tools.NewBM25SearchTool(agent.Tools, ttl, maxSearchResults))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
al.mcp.setManager(mcpManager)
|
||||||
|
})
|
||||||
|
|
||||||
|
return al.mcp.getInitErr()
|
||||||
|
}
|
||||||
|
|
@ -770,6 +770,56 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProcessDirectWithChannel_InitializesMCPInAgentMode(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Tools: config.ToolsConfig{
|
||||||
|
MCP: config.MCPConfig{
|
||||||
|
ToolConfig: config.ToolConfig{
|
||||||
|
Enabled: true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &mockProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
defer al.Close()
|
||||||
|
|
||||||
|
if al.mcp.hasManager() {
|
||||||
|
t.Fatal("expected MCP manager to be nil before first direct processing")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = al.ProcessDirectWithChannel(
|
||||||
|
context.Background(),
|
||||||
|
"hello",
|
||||||
|
"session-1",
|
||||||
|
"cli",
|
||||||
|
"direct",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessDirectWithChannel failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !al.mcp.hasManager() {
|
||||||
|
t.Fatal("expected MCP manager to be initialized in direct agent mode")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestTargetReasoningChannelID_AllChannels(t *testing.T) {
|
func TestTargetReasoningChannelID_AllChannels(t *testing.T) {
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,7 @@ type OutboundMessage struct {
|
||||||
Channel string `json:"channel"`
|
Channel string `json:"channel"`
|
||||||
ChatID string `json:"chat_id"`
|
ChatID string `json:"chat_id"`
|
||||||
Content string `json:"content"`
|
Content string `json:"content"`
|
||||||
|
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// MediaPart describes a single media attachment to send.
|
// MediaPart describes a single media attachment to send.
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
@ -32,6 +33,9 @@ func init() {
|
||||||
uniqueIDPrefix = hex.EncodeToString(b[:])
|
uniqueIDPrefix = hex.EncodeToString(b[:])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// audioAnnotationRe matches audio/voice annotations injected by channels (e.g. [voice], [audio: file.ogg]).
|
||||||
|
var audioAnnotationRe = regexp.MustCompile(`\[(voice|audio)(?::[^\]]*)?\]`)
|
||||||
|
|
||||||
// uniqueID generates a process-unique ID using a random prefix and an atomic counter.
|
// uniqueID generates a process-unique ID using a random prefix and an atomic counter.
|
||||||
// This ID is intended for internal correlation (e.g. media scope keys) and is NOT
|
// This ID is intended for internal correlation (e.g. media scope keys) and is NOT
|
||||||
// cryptographically secure — it must not be used in contexts where unpredictability matters.
|
// cryptographically secure — it must not be used in contexts where unpredictability matters.
|
||||||
|
|
@ -284,13 +288,18 @@ func (c *BaseChannel) HandleMessage(
|
||||||
c.placeholderRecorder.RecordReactionUndo(c.name, chatID, undo)
|
c.placeholderRecorder.RecordReactionUndo(c.name, chatID, undo)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Placeholder — independent pipeline
|
// Placeholder — independent pipeline.
|
||||||
|
// Skip when the message contains audio: the agent will send the
|
||||||
|
// placeholder after transcription completes, so the user sees
|
||||||
|
// "Thinking…" only once the voice has been processed.
|
||||||
|
if !audioAnnotationRe.MatchString(content) {
|
||||||
if pc, ok := c.owner.(PlaceholderCapable); ok {
|
if pc, ok := c.owner.(PlaceholderCapable); ok {
|
||||||
if phID, err := pc.SendPlaceholder(ctx, chatID); err == nil && phID != "" {
|
if phID, err := pc.SendPlaceholder(ctx, chatID); err == nil && phID != "" {
|
||||||
c.placeholderRecorder.RecordPlaceholder(c.name, chatID, phID)
|
c.placeholderRecorder.RecordPlaceholder(c.name, chatID, phID)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if err := c.bus.PublishInbound(ctx, msg); err != nil {
|
if err := c.bus.PublishInbound(ctx, msg); err != nil {
|
||||||
logger.ErrorCF("channels", "Failed to publish inbound message", map[string]any{
|
logger.ErrorCF("channels", "Failed to publish inbound message", map[string]any{
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ import (
|
||||||
|
|
||||||
"github.com/open-dingtalk/dingtalk-stream-sdk-go/chatbot"
|
"github.com/open-dingtalk/dingtalk-stream-sdk-go/chatbot"
|
||||||
"github.com/open-dingtalk/dingtalk-stream-sdk-go/client"
|
"github.com/open-dingtalk/dingtalk-stream-sdk-go/client"
|
||||||
|
dinglog "github.com/open-dingtalk/dingtalk-stream-sdk-go/logger"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
|
@ -39,6 +40,9 @@ func NewDingTalkChannel(cfg config.DingTalkConfig, messageBus *bus.MessageBus) (
|
||||||
return nil, fmt.Errorf("dingtalk client_id and client_secret are required")
|
return nil, fmt.Errorf("dingtalk client_id and client_secret are required")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Set the logger for the Stream SDK
|
||||||
|
dinglog.SetLogger(logger.NewLogger("dingtalk"))
|
||||||
|
|
||||||
base := channels.NewBaseChannel("dingtalk", cfg, messageBus, cfg.AllowFrom,
|
base := channels.NewBaseChannel("dingtalk", cfg, messageBus, cfg.AllowFrom,
|
||||||
channels.WithMaxMessageLength(20000),
|
channels.WithMaxMessageLength(20000),
|
||||||
channels.WithGroupTrigger(cfg.GroupTrigger),
|
channels.WithGroupTrigger(cfg.GroupTrigger),
|
||||||
|
|
|
||||||
|
|
@ -45,6 +45,14 @@ type DiscordChannel struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
||||||
|
discordgo.Logger = logger.NewLogger("discord").
|
||||||
|
WithLevels(map[int]logger.LogLevel{
|
||||||
|
discordgo.LogError: logger.ERROR,
|
||||||
|
discordgo.LogWarning: logger.WARN,
|
||||||
|
discordgo.LogInformational: logger.INFO,
|
||||||
|
discordgo.LogDebug: logger.DEBUG,
|
||||||
|
}).Log
|
||||||
|
|
||||||
session, err := discordgo.New("Bot " + cfg.Token)
|
session, err := discordgo.New("Bot " + cfg.Token)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create discord session: %w", err)
|
return nil, fmt.Errorf("failed to create discord session: %w", err)
|
||||||
|
|
@ -134,7 +142,7 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return c.sendChunk(ctx, channelID, msg.Content)
|
return c.sendChunk(ctx, channelID, msg.Content, msg.ReplyToMessageID)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
|
|
@ -259,14 +267,29 @@ func (c *DiscordChannel) SendPlaceholder(ctx context.Context, chatID string) (st
|
||||||
return msg.ID, nil
|
return msg.ID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content string) error {
|
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content, replyToID string) error {
|
||||||
// Use the passed ctx for timeout control
|
// Use the passed ctx for timeout control
|
||||||
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
done := make(chan error, 1)
|
done := make(chan error, 1)
|
||||||
go func() {
|
go func() {
|
||||||
_, err := c.session.ChannelMessageSend(channelID, content)
|
var err error
|
||||||
|
|
||||||
|
// If we have an ID, we send the message as "Reply"
|
||||||
|
if replyToID != "" {
|
||||||
|
_, err = c.session.ChannelMessageSendComplex(channelID, &discordgo.MessageSend{
|
||||||
|
Content: content,
|
||||||
|
Reference: &discordgo.MessageReference{
|
||||||
|
MessageID: replyToID,
|
||||||
|
ChannelID: channelID,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
// Otherwise, we send a normal message
|
||||||
|
_, err = c.session.ChannelMessageSend(channelID, content)
|
||||||
|
}
|
||||||
|
|
||||||
done <- err
|
done <- err
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -102,6 +102,27 @@ func (m *Manager) RecordPlaceholder(channel, chatID, placeholderID string) {
|
||||||
m.placeholders.Store(key, placeholderEntry{id: placeholderID, createdAt: time.Now()})
|
m.placeholders.Store(key, placeholderEntry{id: placeholderID, createdAt: time.Now()})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SendPlaceholder sends a "Thinking…" placeholder for the given channel/chatID
|
||||||
|
// and records it for later editing. Returns true if a placeholder was sent.
|
||||||
|
func (m *Manager) SendPlaceholder(ctx context.Context, channel, chatID string) bool {
|
||||||
|
m.mu.RLock()
|
||||||
|
ch, ok := m.channels[channel]
|
||||||
|
m.mu.RUnlock()
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
pc, ok := ch.(PlaceholderCapable)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
phID, err := pc.SendPlaceholder(ctx, chatID)
|
||||||
|
if err != nil || phID == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
m.RecordPlaceholder(channel, chatID, phID)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// RecordTypingStop registers a typing stop function for later invocation.
|
// RecordTypingStop registers a typing stop function for later invocation.
|
||||||
// Implements PlaceholderRecorder.
|
// Implements PlaceholderRecorder.
|
||||||
func (m *Manager) RecordTypingStop(channel, chatID string, stop func()) {
|
func (m *Manager) RecordTypingStop(channel, chatID string, stop func()) {
|
||||||
|
|
@ -813,6 +834,39 @@ func (m *Manager) UnregisterChannel(name string) {
|
||||||
delete(m.channels, name)
|
delete(m.channels, name)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SendMessage sends an outbound message synchronously through the channel
|
||||||
|
// worker's rate limiter and retry logic. It blocks until the message is
|
||||||
|
// delivered (or all retries are exhausted), which preserves ordering when
|
||||||
|
// a subsequent operation depends on the message having been sent.
|
||||||
|
func (m *Manager) SendMessage(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
m.mu.RLock()
|
||||||
|
_, exists := m.channels[msg.Channel]
|
||||||
|
w, wExists := m.workers[msg.Channel]
|
||||||
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
if !exists {
|
||||||
|
return fmt.Errorf("channel %s not found", msg.Channel)
|
||||||
|
}
|
||||||
|
if !wExists || w == nil {
|
||||||
|
return fmt.Errorf("channel %s has no active worker", msg.Channel)
|
||||||
|
}
|
||||||
|
|
||||||
|
maxLen := 0
|
||||||
|
if mlp, ok := w.ch.(MessageLengthProvider); ok {
|
||||||
|
maxLen = mlp.MaxMessageLength()
|
||||||
|
}
|
||||||
|
if maxLen > 0 && len([]rune(msg.Content)) > maxLen {
|
||||||
|
for _, chunk := range SplitMessage(msg.Content, maxLen) {
|
||||||
|
chunkMsg := msg
|
||||||
|
chunkMsg.Content = chunk
|
||||||
|
m.sendWithRetry(ctx, msg.Channel, w, chunkMsg)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
m.sendWithRetry(ctx, msg.Channel, w, msg)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, content string) error {
|
func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, content string) error {
|
||||||
m.mu.RLock()
|
m.mu.RLock()
|
||||||
_, exists := m.channels[channelName]
|
_, exists := m.channels[channelName]
|
||||||
|
|
|
||||||
|
|
@ -18,15 +18,31 @@ import (
|
||||||
type mockChannel struct {
|
type mockChannel struct {
|
||||||
BaseChannel
|
BaseChannel
|
||||||
sendFn func(ctx context.Context, msg bus.OutboundMessage) error
|
sendFn func(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
|
sentMessages []bus.OutboundMessage
|
||||||
|
placeholdersSent int
|
||||||
|
editedMessages int
|
||||||
|
lastPlaceholderID string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *mockChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (m *mockChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
m.sentMessages = append(m.sentMessages, msg)
|
||||||
return m.sendFn(ctx, msg)
|
return m.sendFn(ctx, msg)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *mockChannel) Start(ctx context.Context) error { return nil }
|
func (m *mockChannel) Start(ctx context.Context) error { return nil }
|
||||||
func (m *mockChannel) Stop(ctx context.Context) error { return nil }
|
func (m *mockChannel) Stop(ctx context.Context) error { return nil }
|
||||||
|
|
||||||
|
func (m *mockChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
|
||||||
|
m.placeholdersSent++
|
||||||
|
m.lastPlaceholderID = "mock-ph-123"
|
||||||
|
return m.lastPlaceholderID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockChannel) EditMessage(ctx context.Context, chatID, messageID, content string) error {
|
||||||
|
m.editedMessages++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// newTestManager creates a minimal Manager suitable for unit tests.
|
// newTestManager creates a minimal Manager suitable for unit tests.
|
||||||
func newTestManager() *Manager {
|
func newTestManager() *Manager {
|
||||||
return &Manager{
|
return &Manager{
|
||||||
|
|
@ -860,3 +876,286 @@ func TestBuildMediaScope_WithMessageID(t *testing.T) {
|
||||||
t.Fatalf("expected %s, got %s", expected, scope)
|
t.Fatalf("expected %s, got %s", expected, scope)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestManager_PlaceholderConsumedByResponse(t *testing.T) {
|
||||||
|
mgr := &Manager{
|
||||||
|
channels: make(map[string]Channel),
|
||||||
|
workers: make(map[string]*channelWorker),
|
||||||
|
placeholders: sync.Map{},
|
||||||
|
}
|
||||||
|
|
||||||
|
mockCh := &mockChannel{
|
||||||
|
sendFn: func(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
worker := newChannelWorker("mock", mockCh)
|
||||||
|
mgr.channels["mock"] = mockCh
|
||||||
|
mgr.workers["mock"] = worker
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
key := "mock:chat-1"
|
||||||
|
|
||||||
|
// Simulate a placeholder recorded by base.go HandleMessage
|
||||||
|
mgr.RecordPlaceholder("mock", "chat-1", "ph-123")
|
||||||
|
|
||||||
|
if _, ok := mgr.placeholders.Load(key); !ok {
|
||||||
|
t.Fatal("expected placeholder to be recorded")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Transcription feedback arrives first — it should consume the placeholder
|
||||||
|
// and be delivered via EditMessage, not Send.
|
||||||
|
msgTranscript := bus.OutboundMessage{
|
||||||
|
Channel: "mock",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Content: "Transcript: hello",
|
||||||
|
}
|
||||||
|
mgr.sendWithRetry(ctx, "mock", worker, msgTranscript)
|
||||||
|
|
||||||
|
if mockCh.editedMessages != 1 {
|
||||||
|
t.Errorf("expected 1 edited message (placeholder consumed by transcript), got %d", mockCh.editedMessages)
|
||||||
|
}
|
||||||
|
if len(mockCh.sentMessages) != 0 {
|
||||||
|
t.Errorf("expected 0 normal messages (transcript used edit), got %d", len(mockCh.sentMessages))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Placeholder should be gone now
|
||||||
|
if _, ok := mgr.placeholders.Load(key); ok {
|
||||||
|
t.Error("expected placeholder to be removed after being consumed")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Final LLM response arrives — no placeholder left, so it goes through Send
|
||||||
|
msgFinal := bus.OutboundMessage{
|
||||||
|
Channel: "mock",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Content: "Final Answer",
|
||||||
|
}
|
||||||
|
mgr.sendWithRetry(ctx, "mock", worker, msgFinal)
|
||||||
|
|
||||||
|
if len(mockCh.sentMessages) != 1 {
|
||||||
|
t.Errorf("expected 1 normal message sent, got %d", len(mockCh.sentMessages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessage_Synchronous(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
|
||||||
|
var received []bus.OutboundMessage
|
||||||
|
ch := &mockChannel{
|
||||||
|
sendFn: func(_ context.Context, msg bus.OutboundMessage) error {
|
||||||
|
received = append(received, msg)
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
msg := bus.OutboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "123",
|
||||||
|
Content: "hello world",
|
||||||
|
ReplyToMessageID: "msg-456",
|
||||||
|
}
|
||||||
|
|
||||||
|
err := m.SendMessage(context.Background(), msg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendMessage is synchronous — message should already be delivered
|
||||||
|
if len(received) != 1 {
|
||||||
|
t.Fatalf("expected 1 message sent, got %d", len(received))
|
||||||
|
}
|
||||||
|
if received[0].ReplyToMessageID != "msg-456" {
|
||||||
|
t.Fatalf("expected ReplyToMessageID msg-456, got %s", received[0].ReplyToMessageID)
|
||||||
|
}
|
||||||
|
if received[0].Content != "hello world" {
|
||||||
|
t.Fatalf("expected content 'hello world', got %s", received[0].Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessage_UnknownChannel(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
|
||||||
|
msg := bus.OutboundMessage{
|
||||||
|
Channel: "nonexistent",
|
||||||
|
ChatID: "123",
|
||||||
|
Content: "hello",
|
||||||
|
}
|
||||||
|
|
||||||
|
err := m.SendMessage(context.Background(), msg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error for unknown channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessage_NoWorker(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
|
||||||
|
ch := &mockChannel{
|
||||||
|
sendFn: func(_ context.Context, _ bus.OutboundMessage) error { return nil },
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
// No worker registered
|
||||||
|
|
||||||
|
msg := bus.OutboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "123",
|
||||||
|
Content: "hello",
|
||||||
|
}
|
||||||
|
|
||||||
|
err := m.SendMessage(context.Background(), msg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error when no worker exists")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessage_WithRetry(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
|
||||||
|
var callCount int
|
||||||
|
ch := &mockChannel{
|
||||||
|
sendFn: func(_ context.Context, _ bus.OutboundMessage) error {
|
||||||
|
callCount++
|
||||||
|
if callCount == 1 {
|
||||||
|
return fmt.Errorf("transient: %w", ErrTemporary)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
msg := bus.OutboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "123",
|
||||||
|
Content: "retry me",
|
||||||
|
}
|
||||||
|
|
||||||
|
err := m.SendMessage(context.Background(), msg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if callCount != 2 {
|
||||||
|
t.Fatalf("expected 2 Send calls (1 failure + 1 success), got %d", callCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessage_WithSplitting(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
|
||||||
|
var received []string
|
||||||
|
ch := &mockChannelWithLength{
|
||||||
|
mockChannel: mockChannel{
|
||||||
|
sendFn: func(_ context.Context, msg bus.OutboundMessage) error {
|
||||||
|
received = append(received, msg.Content)
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
maxLen: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
msg := bus.OutboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "123",
|
||||||
|
Content: "hello world",
|
||||||
|
}
|
||||||
|
|
||||||
|
err := m.SendMessage(context.Background(), msg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(received) < 2 {
|
||||||
|
t.Fatalf("expected message to be split into at least 2 chunks, got %d", len(received))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMessage_PreservesOrdering(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
|
||||||
|
var order []string
|
||||||
|
ch := &mockChannel{
|
||||||
|
sendFn: func(_ context.Context, msg bus.OutboundMessage) error {
|
||||||
|
order = append(order, msg.Content)
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
// Send two messages sequentially — they must arrive in order
|
||||||
|
_ = m.SendMessage(context.Background(), bus.OutboundMessage{
|
||||||
|
Channel: "test", ChatID: "1", Content: "first",
|
||||||
|
})
|
||||||
|
_ = m.SendMessage(context.Background(), bus.OutboundMessage{
|
||||||
|
Channel: "test", ChatID: "1", Content: "second",
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(order) != 2 {
|
||||||
|
t.Fatalf("expected 2 messages, got %d", len(order))
|
||||||
|
}
|
||||||
|
if order[0] != "first" || order[1] != "second" {
|
||||||
|
t.Fatalf("expected [first, second], got %v", order)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManager_SendPlaceholder(t *testing.T) {
|
||||||
|
mgr := &Manager{
|
||||||
|
channels: make(map[string]Channel),
|
||||||
|
workers: make(map[string]*channelWorker),
|
||||||
|
placeholders: sync.Map{},
|
||||||
|
}
|
||||||
|
|
||||||
|
mockCh := &mockChannel{
|
||||||
|
sendFn: func(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
mgr.channels["mock"] = mockCh
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// SendPlaceholder should send a placeholder and record it
|
||||||
|
ok := mgr.SendPlaceholder(ctx, "mock", "chat-1")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected SendPlaceholder to succeed")
|
||||||
|
}
|
||||||
|
if mockCh.placeholdersSent != 1 {
|
||||||
|
t.Errorf("expected 1 placeholder sent, got %d", mockCh.placeholdersSent)
|
||||||
|
}
|
||||||
|
|
||||||
|
key := "mock:chat-1"
|
||||||
|
if _, loaded := mgr.placeholders.Load(key); !loaded {
|
||||||
|
t.Error("expected placeholder to be recorded in manager")
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendPlaceholder on unknown channel should return false
|
||||||
|
ok = mgr.SendPlaceholder(ctx, "unknown", "chat-1")
|
||||||
|
if ok {
|
||||||
|
t.Error("expected SendPlaceholder to fail for unknown channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,9 @@ import (
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gomarkdown/markdown"
|
||||||
|
mdhtml "github.com/gomarkdown/markdown/html"
|
||||||
|
"github.com/gomarkdown/markdown/parser"
|
||||||
"maunium.net/go/mautrix"
|
"maunium.net/go/mautrix"
|
||||||
"maunium.net/go/mautrix/event"
|
"maunium.net/go/mautrix/event"
|
||||||
"maunium.net/go/mautrix/id"
|
"maunium.net/go/mautrix/id"
|
||||||
|
|
@ -268,6 +271,12 @@ func (c *MatrixChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func markdownToHTML(md string) string {
|
||||||
|
p := parser.NewWithExtensions(parser.CommonExtensions | parser.AutoHeadingIDs)
|
||||||
|
renderer := mdhtml.NewRenderer(mdhtml.RendererOptions{Flags: mdhtml.CommonFlags})
|
||||||
|
return strings.TrimSpace(string(markdown.ToHTML([]byte(md), p, renderer)))
|
||||||
|
}
|
||||||
|
|
||||||
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
|
|
@ -283,16 +292,22 @@ func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, &event.MessageEventContent{
|
_, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, c.messageContent(content))
|
||||||
MsgType: event.MsgText,
|
|
||||||
Body: content,
|
|
||||||
})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("matrix send: %w", channels.ErrTemporary)
|
return fmt.Errorf("matrix send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *MatrixChannel) messageContent(text string) *event.MessageEventContent {
|
||||||
|
mc := &event.MessageEventContent{MsgType: event.MsgText, Body: text}
|
||||||
|
if c.config.MessageFormat != "plain" {
|
||||||
|
mc.Format = event.FormatHTML
|
||||||
|
mc.FormattedBody = markdownToHTML(text)
|
||||||
|
}
|
||||||
|
return mc
|
||||||
|
}
|
||||||
|
|
||||||
// SendMedia implements channels.MediaSender.
|
// SendMedia implements channels.MediaSender.
|
||||||
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
|
|
@ -482,10 +497,7 @@ func (c *MatrixChannel) EditMessage(ctx context.Context, chatID string, messageI
|
||||||
return fmt.Errorf("matrix message ID is empty")
|
return fmt.Errorf("matrix message ID is empty")
|
||||||
}
|
}
|
||||||
|
|
||||||
editContent := &event.MessageEventContent{
|
editContent := c.messageContent(content)
|
||||||
MsgType: event.MsgText,
|
|
||||||
Body: content,
|
|
||||||
}
|
|
||||||
editContent.SetEdit(id.EventID(messageID))
|
editContent.SetEdit(id.EventID(messageID))
|
||||||
|
|
||||||
_, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, editContent)
|
_, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, editContent)
|
||||||
|
|
|
||||||
|
|
@ -4,12 +4,15 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"maunium.net/go/mautrix"
|
"maunium.net/go/mautrix"
|
||||||
"maunium.net/go/mautrix/event"
|
"maunium.net/go/mautrix/event"
|
||||||
"maunium.net/go/mautrix/id"
|
"maunium.net/go/mautrix/id"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestMatrixLocalpartMentionRegexp(t *testing.T) {
|
func TestMatrixLocalpartMentionRegexp(t *testing.T) {
|
||||||
|
|
@ -289,3 +292,50 @@ func TestMatrixOutboundContent(t *testing.T) {
|
||||||
t.Fatalf("unexpected fallback body: %q", noCaption.Body)
|
t.Fatalf("unexpected fallback body: %q", noCaption.Body)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMarkdownToHTML(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
contains string
|
||||||
|
}{
|
||||||
|
{"bold", "**hello**", "<strong>hello</strong>"},
|
||||||
|
{"italic", "_world_", "<em>world</em>"},
|
||||||
|
{"header", "### Title", "<h3"},
|
||||||
|
{"code block", "```\nfoo()\n```", "<code>"},
|
||||||
|
{"inline code", "`x`", "<code>x</code>"},
|
||||||
|
{"plain text", "just text", "just text"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := markdownToHTML(tt.input)
|
||||||
|
if !strings.Contains(got, tt.contains) {
|
||||||
|
t.Fatalf("markdownToHTML(%q) = %q, want it to contain %q", tt.input, got, tt.contains)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMessageContent(t *testing.T) {
|
||||||
|
richtext := &MatrixChannel{config: config.MatrixConfig{MessageFormat: "richtext"}}
|
||||||
|
plain := &MatrixChannel{config: config.MatrixConfig{MessageFormat: "plain"}}
|
||||||
|
defaultt := &MatrixChannel{config: config.MatrixConfig{}}
|
||||||
|
|
||||||
|
for _, c := range []*MatrixChannel{richtext, defaultt} {
|
||||||
|
mc := c.messageContent("**hi**")
|
||||||
|
if mc.Format != event.FormatHTML {
|
||||||
|
t.Errorf("format %q: expected FormatHTML, got %q", c.config.MessageFormat, mc.Format)
|
||||||
|
}
|
||||||
|
if !strings.Contains(mc.FormattedBody, "<strong>hi</strong>") {
|
||||||
|
t.Errorf("format %q: FormattedBody %q missing <strong>", c.config.MessageFormat, mc.FormattedBody)
|
||||||
|
}
|
||||||
|
if mc.Body != "**hi**" {
|
||||||
|
t.Errorf("format %q: Body should remain plain, got %q", c.config.MessageFormat, mc.Body)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mc := plain.messageContent("**hi**")
|
||||||
|
if mc.Format != "" || mc.FormattedBody != "" {
|
||||||
|
t.Errorf("plain: expected no formatting, got format=%q formattedBody=%q", mc.Format, mc.FormattedBody)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -78,6 +78,7 @@ func (c *QQChannel) Start(ctx context.Context) error {
|
||||||
return fmt.Errorf("QQ app_id and app_secret not configured")
|
return fmt.Errorf("QQ app_id and app_secret not configured")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
botgo.SetLogger(logger.NewLogger("botgo"))
|
||||||
logger.InfoC("qq", "Starting QQ bot (WebSocket mode)")
|
logger.InfoC("qq", "Starting QQ bot (WebSocket mode)")
|
||||||
|
|
||||||
// Reinitialize shutdown signal for clean restart.
|
// Reinitialize shutdown signal for clean restart.
|
||||||
|
|
|
||||||
|
|
@ -122,7 +122,11 @@ func (c *SlackChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
slack.MsgOptionText(msg.Content, false),
|
slack.MsgOptionText(msg.Content, false),
|
||||||
}
|
}
|
||||||
|
|
||||||
if threadTS != "" {
|
if msg.ReplyToMessageID != "" && threadTS == "" {
|
||||||
|
// Answer to the message by creating a Thread under it
|
||||||
|
opts = append(opts, slack.MsgOptionTS(msg.ReplyToMessageID))
|
||||||
|
} else if threadTS != "" {
|
||||||
|
// If we are already in a thread, continue in the thread
|
||||||
opts = append(opts, slack.MsgOptionTS(threadTS))
|
opts = append(opts, slack.MsgOptionTS(threadTS))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -77,6 +77,7 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
|
||||||
if baseURL := strings.TrimRight(strings.TrimSpace(telegramCfg.BaseURL), "/"); baseURL != "" {
|
if baseURL := strings.TrimRight(strings.TrimSpace(telegramCfg.BaseURL), "/"); baseURL != "" {
|
||||||
opts = append(opts, telego.WithAPIServer(baseURL))
|
opts = append(opts, telego.WithAPIServer(baseURL))
|
||||||
}
|
}
|
||||||
|
opts = append(opts, telego.WithLogger(logger.NewLogger("telego")))
|
||||||
|
|
||||||
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -168,7 +169,7 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, err := parseChatID(msg.ChatID)
|
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
@ -180,6 +181,7 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
||||||
// so msg.Content is guaranteed to be within that limit. We still need to
|
// so msg.Content is guaranteed to be within that limit. We still need to
|
||||||
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
||||||
|
replyToID := msg.ReplyToMessageID
|
||||||
queue := []string{msg.Content}
|
queue := []string{msg.Content}
|
||||||
for len(queue) > 0 {
|
for len(queue) > 0 {
|
||||||
chunk := queue[0]
|
chunk := queue[0]
|
||||||
|
|
@ -200,9 +202,11 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendHTMLChunk(ctx, chatID, htmlContent, chunk); err != nil {
|
if err := c.sendHTMLChunk(ctx, chatID, threadID, htmlContent, chunk, replyToID); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
// Only the first chunk should be a reply; subsequent chunks are normal messages.
|
||||||
|
replyToID = ""
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|
@ -210,9 +214,20 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
|
|
||||||
// sendHTMLChunk sends a single HTML message, falling back to the original
|
// sendHTMLChunk sends a single HTML message, falling back to the original
|
||||||
// markdown as plain text on parse failure so users never see raw HTML tags.
|
// markdown as plain text on parse failure so users never see raw HTML tags.
|
||||||
func (c *TelegramChannel) sendHTMLChunk(ctx context.Context, chatID int64, htmlContent, mdFallback string) error {
|
func (c *TelegramChannel) sendHTMLChunk(
|
||||||
|
ctx context.Context, chatID int64, threadID int, htmlContent, mdFallback string, replyToID string,
|
||||||
|
) error {
|
||||||
tgMsg := tu.Message(tu.ID(chatID), htmlContent)
|
tgMsg := tu.Message(tu.ID(chatID), htmlContent)
|
||||||
tgMsg.ParseMode = telego.ModeHTML
|
tgMsg.ParseMode = telego.ModeHTML
|
||||||
|
tgMsg.MessageThreadID = threadID
|
||||||
|
|
||||||
|
if replyToID != "" {
|
||||||
|
if mid, parseErr := strconv.Atoi(replyToID); parseErr == nil {
|
||||||
|
tgMsg.ReplyParameters = &telego.ReplyParameters{
|
||||||
|
MessageID: mid,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if _, err := c.bot.SendMessage(ctx, tgMsg); err != nil {
|
if _, err := c.bot.SendMessage(ctx, tgMsg); err != nil {
|
||||||
logger.ErrorCF("telegram", "HTML parse failed, falling back to plain text", map[string]any{
|
logger.ErrorCF("telegram", "HTML parse failed, falling back to plain text", map[string]any{
|
||||||
|
|
@ -232,13 +247,16 @@ func (c *TelegramChannel) sendHTMLChunk(ctx context.Context, chatID int64, htmlC
|
||||||
// (Telegram's typing indicator expires after ~5s) in a background goroutine.
|
// (Telegram's typing indicator expires after ~5s) in a background goroutine.
|
||||||
// The returned stop function is idempotent and cancels the goroutine.
|
// The returned stop function is idempotent and cancels the goroutine.
|
||||||
func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
||||||
cid, err := parseChatID(chatID)
|
cid, threadID, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return func() {}, err
|
return func() {}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
action := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
||||||
|
action.MessageThreadID = threadID
|
||||||
|
|
||||||
// Send the first typing action immediately
|
// Send the first typing action immediately
|
||||||
_ = c.bot.SendChatAction(ctx, tu.ChatAction(tu.ID(cid), telego.ChatActionTyping))
|
_ = c.bot.SendChatAction(ctx, action)
|
||||||
|
|
||||||
typingCtx, cancel := context.WithCancel(ctx)
|
typingCtx, cancel := context.WithCancel(ctx)
|
||||||
go func() {
|
go func() {
|
||||||
|
|
@ -249,7 +267,9 @@ func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(
|
||||||
case <-typingCtx.Done():
|
case <-typingCtx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
_ = c.bot.SendChatAction(typingCtx, tu.ChatAction(tu.ID(cid), telego.ChatActionTyping))
|
a := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
||||||
|
a.MessageThreadID = threadID
|
||||||
|
_ = c.bot.SendChatAction(typingCtx, a)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
@ -259,7 +279,7 @@ func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(
|
||||||
|
|
||||||
// EditMessage implements channels.MessageEditor.
|
// EditMessage implements channels.MessageEditor.
|
||||||
func (c *TelegramChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
func (c *TelegramChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
||||||
cid, err := parseChatID(chatID)
|
cid, _, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
@ -288,12 +308,14 @@ func (c *TelegramChannel) SendPlaceholder(ctx context.Context, chatID string) (s
|
||||||
text = "Thinking... 💭"
|
text = "Thinking... 💭"
|
||||||
}
|
}
|
||||||
|
|
||||||
cid, err := parseChatID(chatID)
|
cid, threadID, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
pMsg, err := c.bot.SendMessage(ctx, tu.Message(tu.ID(cid), text))
|
phMsg := tu.Message(tu.ID(cid), text)
|
||||||
|
phMsg.MessageThreadID = threadID
|
||||||
|
pMsg, err := c.bot.SendMessage(ctx, phMsg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
@ -307,7 +329,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, err := parseChatID(msg.ChatID)
|
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
@ -340,6 +362,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
case "image":
|
case "image":
|
||||||
params := &telego.SendPhotoParams{
|
params := &telego.SendPhotoParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
|
MessageThreadID: threadID,
|
||||||
Photo: telego.InputFile{File: file},
|
Photo: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
|
|
@ -347,6 +370,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
case "audio":
|
case "audio":
|
||||||
params := &telego.SendAudioParams{
|
params := &telego.SendAudioParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
|
MessageThreadID: threadID,
|
||||||
Audio: telego.InputFile{File: file},
|
Audio: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
|
|
@ -354,6 +378,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
case "video":
|
case "video":
|
||||||
params := &telego.SendVideoParams{
|
params := &telego.SendVideoParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
|
MessageThreadID: threadID,
|
||||||
Video: telego.InputFile{File: file},
|
Video: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
|
|
@ -361,6 +386,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
default: // "file" or unknown types
|
default: // "file" or unknown types
|
||||||
params := &telego.SendDocumentParams{
|
params := &telego.SendDocumentParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
|
MessageThreadID: threadID,
|
||||||
Document: telego.InputFile{File: file},
|
Document: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
|
|
@ -506,19 +532,28 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
content = cleaned
|
content = cleaned
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// For forum topics, embed the thread ID as "chatID/threadID" so replies
|
||||||
|
// route to the correct topic and each topic gets its own session.
|
||||||
|
// Only forum groups (IsForum) are handled; regular group reply threads
|
||||||
|
// must share one session per group.
|
||||||
|
compositeChatID := fmt.Sprintf("%d", chatID)
|
||||||
|
threadID := message.MessageThreadID
|
||||||
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
compositeChatID = fmt.Sprintf("%d/%d", chatID, threadID)
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("telegram", "Received message", map[string]any{
|
logger.DebugCF("telegram", "Received message", map[string]any{
|
||||||
"sender_id": sender.CanonicalID,
|
"sender_id": sender.CanonicalID,
|
||||||
"chat_id": fmt.Sprintf("%d", chatID),
|
"chat_id": compositeChatID,
|
||||||
|
"thread_id": threadID,
|
||||||
"preview": utils.Truncate(content, 50),
|
"preview": utils.Truncate(content, 50),
|
||||||
})
|
})
|
||||||
|
|
||||||
// Placeholder is now auto-triggered by BaseChannel.HandleMessage via PlaceholderCapable
|
|
||||||
|
|
||||||
peerKind := "direct"
|
peerKind := "direct"
|
||||||
peerID := fmt.Sprintf("%d", user.ID)
|
peerID := fmt.Sprintf("%d", user.ID)
|
||||||
if message.Chat.Type != "private" {
|
if message.Chat.Type != "private" {
|
||||||
peerKind = "group"
|
peerKind = "group"
|
||||||
peerID = fmt.Sprintf("%d", chatID)
|
peerID = compositeChatID
|
||||||
}
|
}
|
||||||
|
|
||||||
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
||||||
|
|
@ -531,11 +566,17 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Set parent_peer metadata for per-topic agent binding.
|
||||||
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
metadata["parent_peer_kind"] = "topic"
|
||||||
|
metadata["parent_peer_id"] = fmt.Sprintf("%d", threadID)
|
||||||
|
}
|
||||||
|
|
||||||
c.HandleMessage(c.ctx,
|
c.HandleMessage(c.ctx,
|
||||||
peer,
|
peer,
|
||||||
messageID,
|
messageID,
|
||||||
platformID,
|
platformID,
|
||||||
fmt.Sprintf("%d", chatID),
|
compositeChatID,
|
||||||
content,
|
content,
|
||||||
mediaPaths,
|
mediaPaths,
|
||||||
metadata,
|
metadata,
|
||||||
|
|
@ -583,10 +624,23 @@ func (c *TelegramChannel) downloadFile(ctx context.Context, fileID, ext string)
|
||||||
return c.downloadFileWithInfo(file, ext)
|
return c.downloadFileWithInfo(file, ext)
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseChatID(chatIDStr string) (int64, error) {
|
// parseTelegramChatID splits "chatID/threadID" into its components.
|
||||||
var id int64
|
// Returns threadID=0 when no "/" is present (non-forum messages).
|
||||||
_, err := fmt.Sscanf(chatIDStr, "%d", &id)
|
func parseTelegramChatID(chatID string) (int64, int, error) {
|
||||||
return id, err
|
idx := strings.Index(chatID, "/")
|
||||||
|
if idx == -1 {
|
||||||
|
cid, err := strconv.ParseInt(chatID, 10, 64)
|
||||||
|
return cid, 0, err
|
||||||
|
}
|
||||||
|
cid, err := strconv.ParseInt(chatID[:idx], 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, err
|
||||||
|
}
|
||||||
|
tid, err := strconv.Atoi(chatID[idx+1:])
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, fmt.Errorf("invalid thread ID in chat ID %q: %w", chatID, err)
|
||||||
|
}
|
||||||
|
return cid, tid, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func markdownToTelegramHTML(text string) string {
|
func markdownToTelegramHTML(text string) string {
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/mymmrac/telego"
|
"github.com/mymmrac/telego"
|
||||||
ta "github.com/mymmrac/telego/telegoapi"
|
ta "github.com/mymmrac/telego/telegoapi"
|
||||||
|
|
@ -271,3 +272,191 @@ func TestSend_InvalidChatID(t *testing.T) {
|
||||||
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
|
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
|
||||||
assert.Empty(t, caller.calls)
|
assert.Empty(t, caller.calls)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_Plain(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("12345")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(12345), cid)
|
||||||
|
assert.Equal(t, 0, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_NegativeGroup(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("-1001234567890")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(-1001234567890), cid)
|
||||||
|
assert.Equal(t, 0, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_WithThreadID(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("-1001234567890/42")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(-1001234567890), cid)
|
||||||
|
assert.Equal(t, 42, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_GeneralTopic(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("-100123/1")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(-100123), cid)
|
||||||
|
assert.Equal(t, 1, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_Invalid(t *testing.T) {
|
||||||
|
_, _, err := parseTelegramChatID("not-a-number")
|
||||||
|
assert.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_InvalidThreadID(t *testing.T) {
|
||||||
|
_, _, err := parseTelegramChatID("-100123/not-a-thread")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "invalid thread ID")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSend_WithForumThreadID(t *testing.T) {
|
||||||
|
caller := &stubCaller{
|
||||||
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
|
return successResponse(t), nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
|
err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
|
ChatID: "-1001234567890/42",
|
||||||
|
Content: "Hello from topic",
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, caller.calls, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ForumTopic_SetsMetadata(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "hello from topic",
|
||||||
|
MessageID: 10,
|
||||||
|
MessageThreadID: 42,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: -1001234567890,
|
||||||
|
Type: "supergroup",
|
||||||
|
IsForum: true,
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 7,
|
||||||
|
FirstName: "Alice",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
inbound, ok := messageBus.ConsumeInbound(ctx)
|
||||||
|
require.True(t, ok, "expected inbound message")
|
||||||
|
|
||||||
|
// Composite chatID should include thread ID
|
||||||
|
assert.Equal(t, "-1001234567890/42", inbound.ChatID)
|
||||||
|
|
||||||
|
// Peer ID should include thread ID for session key isolation
|
||||||
|
assert.Equal(t, "group", inbound.Peer.Kind)
|
||||||
|
assert.Equal(t, "-1001234567890/42", inbound.Peer.ID)
|
||||||
|
|
||||||
|
// Parent peer metadata should be set for agent binding
|
||||||
|
assert.Equal(t, "topic", inbound.Metadata["parent_peer_kind"])
|
||||||
|
assert.Equal(t, "42", inbound.Metadata["parent_peer_id"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_NoForum_NoThreadMetadata(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "regular group message",
|
||||||
|
MessageID: 11,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: -100999,
|
||||||
|
Type: "group",
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 8,
|
||||||
|
FirstName: "Bob",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
inbound, ok := messageBus.ConsumeInbound(ctx)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
// Plain chatID without thread suffix
|
||||||
|
assert.Equal(t, "-100999", inbound.ChatID)
|
||||||
|
|
||||||
|
// Peer ID should be raw chat ID (no thread suffix)
|
||||||
|
assert.Equal(t, "group", inbound.Peer.Kind)
|
||||||
|
assert.Equal(t, "-100999", inbound.Peer.ID)
|
||||||
|
|
||||||
|
// No parent peer metadata
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_kind"])
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ReplyThread_NonForum_NoIsolation(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// In regular groups, reply threads set MessageThreadID to the original
|
||||||
|
// message ID. This should NOT trigger per-thread session isolation.
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "reply in thread",
|
||||||
|
MessageID: 20,
|
||||||
|
MessageThreadID: 15,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: -100999,
|
||||||
|
Type: "supergroup",
|
||||||
|
IsForum: false,
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 9,
|
||||||
|
FirstName: "Carol",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
inbound, ok := messageBus.ConsumeInbound(ctx)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
// chatID should NOT include thread suffix for non-forum groups
|
||||||
|
assert.Equal(t, "-100999", inbound.ChatID)
|
||||||
|
|
||||||
|
// Peer ID should be raw chat ID (shared session for whole group)
|
||||||
|
assert.Equal(t, "group", inbound.Peer.Kind)
|
||||||
|
assert.Equal(t, "-100999", inbound.Peer.ID)
|
||||||
|
|
||||||
|
// No parent peer metadata
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_kind"])
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -209,7 +209,7 @@ func TestWeComAppVerifySignature(t *testing.T) {
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("empty token skips verification", func(t *testing.T) {
|
t.Run("empty token rejects verification (fail-closed)", func(t *testing.T) {
|
||||||
cfgEmpty := config.WeComAppConfig{
|
cfgEmpty := config.WeComAppConfig{
|
||||||
CorpID: "test_corp_id",
|
CorpID: "test_corp_id",
|
||||||
CorpSecret: "test_secret",
|
CorpSecret: "test_secret",
|
||||||
|
|
@ -218,8 +218,8 @@ func TestWeComAppVerifySignature(t *testing.T) {
|
||||||
}
|
}
|
||||||
chEmpty, _ := NewWeComAppChannel(cfgEmpty, msgBus)
|
chEmpty, _ := NewWeComAppChannel(cfgEmpty, msgBus)
|
||||||
|
|
||||||
if !verifySignature(chEmpty.config.Token, "any_sig", "any_ts", "any_nonce", "any_msg") {
|
if verifySignature(chEmpty.config.Token, "any_sig", "any_ts", "any_nonce", "any_msg") {
|
||||||
t.Error("empty token should skip verification and return true")
|
t.Error("empty token should reject verification (fail-closed)")
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -189,8 +189,7 @@ func TestWeComBotVerifySignature(t *testing.T) {
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("empty token skips verification", func(t *testing.T) {
|
t.Run("empty token rejects verification (fail-closed)", func(t *testing.T) {
|
||||||
// Create a channel manually with empty token to test the behavior
|
|
||||||
cfgEmpty := config.WeComConfig{
|
cfgEmpty := config.WeComConfig{
|
||||||
Token: "",
|
Token: "",
|
||||||
WebhookURL: "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
WebhookURL: "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
||||||
|
|
@ -199,8 +198,8 @@ func TestWeComBotVerifySignature(t *testing.T) {
|
||||||
config: cfgEmpty,
|
config: cfgEmpty,
|
||||||
}
|
}
|
||||||
|
|
||||||
if !verifySignature(chEmpty.config.Token, "any_sig", "any_ts", "any_nonce", "any_msg") {
|
if verifySignature(chEmpty.config.Token, "any_sig", "any_ts", "any_nonce", "any_msg") {
|
||||||
t.Error("empty token should skip verification and return true")
|
t.Error("empty token should reject verification (fail-closed)")
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -31,7 +31,7 @@ func computeSignature(token, timestamp, nonce, encrypt string) string {
|
||||||
// This is a common function used by both WeCom Bot and WeCom App
|
// This is a common function used by both WeCom Bot and WeCom App
|
||||||
func verifySignature(token, msgSignature, timestamp, nonce, msgEncrypt string) bool {
|
func verifySignature(token, msgSignature, timestamp, nonce, msgEncrypt string) bool {
|
||||||
if token == "" {
|
if token == "" {
|
||||||
return true // Skip verification if token is not set
|
return false
|
||||||
}
|
}
|
||||||
return computeSignature(token, timestamp, nonce, msgEncrypt) == msgSignature
|
return computeSignature(token, timestamp, nonce, msgEncrypt) == msgSignature
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/caarlos0/env/v11"
|
"github.com/caarlos0/env/v11"
|
||||||
|
|
@ -16,6 +17,8 @@ var rrCounter atomic.Uint64
|
||||||
|
|
||||||
// FlexibleStringSlice is a []string that also accepts JSON numbers,
|
// FlexibleStringSlice is a []string that also accepts JSON numbers,
|
||||||
// so allow_from can contain both "123" and 123.
|
// so allow_from can contain both "123" and 123.
|
||||||
|
// It also supports parsing comma-separated strings from environment variables,
|
||||||
|
// including both English (,) and Chinese (,) commas.
|
||||||
type FlexibleStringSlice []string
|
type FlexibleStringSlice []string
|
||||||
|
|
||||||
func (f *FlexibleStringSlice) UnmarshalJSON(data []byte) error {
|
func (f *FlexibleStringSlice) UnmarshalJSON(data []byte) error {
|
||||||
|
|
@ -47,6 +50,30 @@ func (f *FlexibleStringSlice) UnmarshalJSON(data []byte) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// UnmarshalText implements encoding.TextUnmarshaler to support env variable parsing.
|
||||||
|
// It handles comma-separated values with both English (,) and Chinese (,) commas.
|
||||||
|
func (f *FlexibleStringSlice) UnmarshalText(text []byte) error {
|
||||||
|
if len(text) == 0 {
|
||||||
|
*f = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
s := string(text)
|
||||||
|
// Replace Chinese comma with English comma, then split
|
||||||
|
s = strings.ReplaceAll(s, ",", ",")
|
||||||
|
parts := strings.Split(s, ",")
|
||||||
|
|
||||||
|
result := make([]string, 0, len(parts))
|
||||||
|
for _, part := range parts {
|
||||||
|
part = strings.TrimSpace(part)
|
||||||
|
if part != "" {
|
||||||
|
result = append(result, part)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
*f = result
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Agents AgentsConfig `json:"agents"`
|
Agents AgentsConfig `json:"agents"`
|
||||||
Bindings []AgentBinding `json:"bindings,omitempty"`
|
Bindings []AgentBinding `json:"bindings,omitempty"`
|
||||||
|
|
@ -58,6 +85,17 @@ type Config struct {
|
||||||
Tools ToolsConfig `json:"tools"`
|
Tools ToolsConfig `json:"tools"`
|
||||||
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
||||||
Devices DevicesConfig `json:"devices"`
|
Devices DevicesConfig `json:"devices"`
|
||||||
|
Voice VoiceConfig `json:"voice"`
|
||||||
|
// BuildInfo contains build-time version information
|
||||||
|
BuildInfo BuildInfo `json:"build_info,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildInfo contains build-time version information
|
||||||
|
type BuildInfo struct {
|
||||||
|
Version string `json:"version"`
|
||||||
|
GitCommit string `json:"git_commit"`
|
||||||
|
BuildTime string `json:"build_time"`
|
||||||
|
GoVersion string `json:"go_version"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// MarshalJSON implements custom JSON marshaling for Config
|
// MarshalJSON implements custom JSON marshaling for Config
|
||||||
|
|
@ -344,6 +382,7 @@ type MatrixConfig struct {
|
||||||
AccessToken string `json:"access_token" env:"PICOCLAW_CHANNELS_MATRIX_ACCESS_TOKEN"`
|
AccessToken string `json:"access_token" env:"PICOCLAW_CHANNELS_MATRIX_ACCESS_TOKEN"`
|
||||||
DeviceID string `json:"device_id,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_DEVICE_ID"`
|
DeviceID string `json:"device_id,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_DEVICE_ID"`
|
||||||
JoinOnInvite bool `json:"join_on_invite" env:"PICOCLAW_CHANNELS_MATRIX_JOIN_ON_INVITE"`
|
JoinOnInvite bool `json:"join_on_invite" env:"PICOCLAW_CHANNELS_MATRIX_JOIN_ON_INVITE"`
|
||||||
|
MessageFormat string `json:"message_format,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_MESSAGE_FORMAT"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_MATRIX_ALLOW_FROM"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_MATRIX_ALLOW_FROM"`
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
||||||
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
|
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
|
||||||
|
|
@ -461,6 +500,10 @@ type DevicesConfig struct {
|
||||||
MonitorUSB bool `json:"monitor_usb" env:"PICOCLAW_DEVICES_MONITOR_USB"`
|
MonitorUSB bool `json:"monitor_usb" env:"PICOCLAW_DEVICES_MONITOR_USB"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type VoiceConfig struct {
|
||||||
|
EchoTranscription bool `json:"echo_transcription" env:"PICOCLAW_VOICE_ECHO_TRANSCRIPTION"`
|
||||||
|
}
|
||||||
|
|
||||||
type ProvidersConfig struct {
|
type ProvidersConfig struct {
|
||||||
Anthropic ProviderConfig `json:"anthropic"`
|
Anthropic ProviderConfig `json:"anthropic"`
|
||||||
OpenAI OpenAIProviderConfig `json:"openai"`
|
OpenAI OpenAIProviderConfig `json:"openai"`
|
||||||
|
|
@ -484,6 +527,7 @@ type ProvidersConfig struct {
|
||||||
Mistral ProviderConfig `json:"mistral"`
|
Mistral ProviderConfig `json:"mistral"`
|
||||||
Avian ProviderConfig `json:"avian"`
|
Avian ProviderConfig `json:"avian"`
|
||||||
Minimax ProviderConfig `json:"minimax"`
|
Minimax ProviderConfig `json:"minimax"`
|
||||||
|
LongCat ProviderConfig `json:"longcat"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsEmpty checks if all provider configs are empty (no API keys or API bases set)
|
// IsEmpty checks if all provider configs are empty (no API keys or API bases set)
|
||||||
|
|
@ -510,7 +554,8 @@ func (p ProvidersConfig) IsEmpty() bool {
|
||||||
p.Qwen.APIKey == "" && p.Qwen.APIBase == "" &&
|
p.Qwen.APIKey == "" && p.Qwen.APIBase == "" &&
|
||||||
p.Mistral.APIKey == "" && p.Mistral.APIBase == "" &&
|
p.Mistral.APIKey == "" && p.Mistral.APIBase == "" &&
|
||||||
p.Avian.APIKey == "" && p.Avian.APIBase == "" &&
|
p.Avian.APIKey == "" && p.Avian.APIBase == "" &&
|
||||||
p.Minimax.APIKey == "" && p.Minimax.APIBase == ""
|
p.Minimax.APIKey == "" && p.Minimax.APIBase == "" &&
|
||||||
|
p.LongCat.APIKey == "" && p.LongCat.APIBase == ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// MarshalJSON implements custom JSON marshaling for ProvidersConfig
|
// MarshalJSON implements custom JSON marshaling for ProvidersConfig
|
||||||
|
|
@ -595,12 +640,14 @@ type ToolConfig struct {
|
||||||
type BraveConfig struct {
|
type BraveConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
||||||
|
APIKeys []string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEYS"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"`
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type TavilyConfig struct {
|
type TavilyConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_TAVILY_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_TAVILY_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEY"`
|
||||||
|
APIKeys []string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEYS"`
|
||||||
BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_TAVILY_BASE_URL"`
|
BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_TAVILY_BASE_URL"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"`
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
@ -613,6 +660,7 @@ type DuckDuckGoConfig struct {
|
||||||
type PerplexityConfig struct {
|
type PerplexityConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEY"`
|
||||||
|
APIKeys []string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEYS"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -654,6 +702,7 @@ type CronToolsConfig struct {
|
||||||
type ExecConfig struct {
|
type ExecConfig struct {
|
||||||
ToolConfig ` envPrefix:"PICOCLAW_TOOLS_EXEC_"`
|
ToolConfig ` envPrefix:"PICOCLAW_TOOLS_EXEC_"`
|
||||||
EnableDenyPatterns bool ` env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS" json:"enable_deny_patterns"`
|
EnableDenyPatterns bool ` env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS" json:"enable_deny_patterns"`
|
||||||
|
AllowRemote bool ` env:"PICOCLAW_TOOLS_EXEC_ALLOW_REMOTE" json:"allow_remote"`
|
||||||
CustomDenyPatterns []string ` env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS" json:"custom_deny_patterns"`
|
CustomDenyPatterns []string ` env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS" json:"custom_deny_patterns"`
|
||||||
CustomAllowPatterns []string ` env:"PICOCLAW_TOOLS_EXEC_CUSTOM_ALLOW_PATTERNS" json:"custom_allow_patterns"`
|
CustomAllowPatterns []string ` env:"PICOCLAW_TOOLS_EXEC_CUSTOM_ALLOW_PATTERNS" json:"custom_allow_patterns"`
|
||||||
TimeoutSeconds int ` env:"PICOCLAW_TOOLS_EXEC_TIMEOUT_SECONDS" json:"timeout_seconds"` // 0 means use default (60s)
|
TimeoutSeconds int ` env:"PICOCLAW_TOOLS_EXEC_TIMEOUT_SECONDS" json:"timeout_seconds"` // 0 means use default (60s)
|
||||||
|
|
@ -933,6 +982,29 @@ func (c *Config) ValidateModelList() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func MergeAPIKeys(apiKey string, apiKeys []string) []string {
|
||||||
|
seen := make(map[string]struct{})
|
||||||
|
var all []string
|
||||||
|
|
||||||
|
if k := strings.TrimSpace(apiKey); k != "" {
|
||||||
|
if _, exists := seen[k]; !exists {
|
||||||
|
seen[k] = struct{}{}
|
||||||
|
all = append(all, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, k := range apiKeys {
|
||||||
|
if trimmed := strings.TrimSpace(k); trimmed != "" {
|
||||||
|
if _, exists := seen[trimmed]; !exists {
|
||||||
|
seen[trimmed] = struct{}{}
|
||||||
|
all = append(all, trimmed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return all
|
||||||
|
}
|
||||||
|
|
||||||
func (t *ToolsConfig) IsToolEnabled(name string) bool {
|
func (t *ToolsConfig) IsToolEnabled(name string) bool {
|
||||||
switch name {
|
switch name {
|
||||||
case "web":
|
case "web":
|
||||||
|
|
|
||||||
|
|
@ -296,7 +296,7 @@ func TestDefaultConfig_WebTools(t *testing.T) {
|
||||||
if cfg.Tools.Web.Brave.MaxResults != 5 {
|
if cfg.Tools.Web.Brave.MaxResults != 5 {
|
||||||
t.Error("Expected Brave MaxResults 5, got ", cfg.Tools.Web.Brave.MaxResults)
|
t.Error("Expected Brave MaxResults 5, got ", cfg.Tools.Web.Brave.MaxResults)
|
||||||
}
|
}
|
||||||
if cfg.Tools.Web.Brave.APIKey != "" {
|
if len(cfg.Tools.Web.Brave.APIKeys) != 0 {
|
||||||
t.Error("Brave API key should be empty by default")
|
t.Error("Brave API key should be empty by default")
|
||||||
}
|
}
|
||||||
if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 {
|
if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 {
|
||||||
|
|
@ -384,6 +384,13 @@ func TestDefaultConfig_OpenAIWebSearchEnabled(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDefaultConfig_ExecAllowRemoteEnabled(t *testing.T) {
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
if !cfg.Tools.Exec.AllowRemote {
|
||||||
|
t.Fatal("DefaultConfig().Tools.Exec.AllowRemote should be true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLoadConfig_OpenAIWebSearchDefaultsTrueWhenUnset(t *testing.T) {
|
func TestLoadConfig_OpenAIWebSearchDefaultsTrueWhenUnset(t *testing.T) {
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
configPath := filepath.Join(dir, "config.json")
|
configPath := filepath.Join(dir, "config.json")
|
||||||
|
|
@ -400,6 +407,22 @@ func TestLoadConfig_OpenAIWebSearchDefaultsTrueWhenUnset(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoadConfig_ExecAllowRemoteDefaultsTrueWhenUnset(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
configPath := filepath.Join(dir, "config.json")
|
||||||
|
if err := os.WriteFile(configPath, []byte(`{"tools":{"exec":{"enable_deny_patterns":true}}}`), 0o600); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
if !cfg.Tools.Exec.AllowRemote {
|
||||||
|
t.Fatal("tools.exec.allow_remote should remain true when unset in config file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLoadConfig_OpenAIWebSearchCanBeDisabled(t *testing.T) {
|
func TestLoadConfig_OpenAIWebSearchCanBeDisabled(t *testing.T) {
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
configPath := filepath.Join(dir, "config.json")
|
configPath := filepath.Join(dir, "config.json")
|
||||||
|
|
@ -482,3 +505,119 @@ func TestDefaultConfig_WorkspacePath_WithPicoclawHome(t *testing.T) {
|
||||||
t.Errorf("Workspace path with PICOCLAW_HOME = %q, want %q", cfg.Agents.Defaults.Workspace, want)
|
t.Errorf("Workspace path with PICOCLAW_HOME = %q, want %q", cfg.Agents.Defaults.Workspace, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestFlexibleStringSlice_UnmarshalText tests UnmarshalText with various comma separators
|
||||||
|
func TestFlexibleStringSlice_UnmarshalText(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
expected []string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "English commas only",
|
||||||
|
input: "123,456,789",
|
||||||
|
expected: []string{"123", "456", "789"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Chinese commas only",
|
||||||
|
input: "123,456,789",
|
||||||
|
expected: []string{"123", "456", "789"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Mixed English and Chinese commas",
|
||||||
|
input: "123,456,789",
|
||||||
|
expected: []string{"123", "456", "789"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Single value",
|
||||||
|
input: "123",
|
||||||
|
expected: []string{"123"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Values with whitespace",
|
||||||
|
input: " 123 , 456 , 789 ",
|
||||||
|
expected: []string{"123", "456", "789"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Empty string",
|
||||||
|
input: "",
|
||||||
|
expected: nil,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Only commas - English",
|
||||||
|
input: ",,",
|
||||||
|
expected: []string{},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Only commas - Chinese",
|
||||||
|
input: ",,",
|
||||||
|
expected: []string{},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Mixed commas with empty parts",
|
||||||
|
input: "123,,456,,789",
|
||||||
|
expected: []string{"123", "456", "789"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Complex mixed values",
|
||||||
|
input: "user1@example.com,user2@test.com, admin@domain.org",
|
||||||
|
expected: []string{"user1@example.com", "user2@test.com", "admin@domain.org"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var f FlexibleStringSlice
|
||||||
|
err := f.UnmarshalText([]byte(tt.input))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("UnmarshalText(%q) error = %v", tt.input, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if tt.expected == nil {
|
||||||
|
if f != nil {
|
||||||
|
t.Errorf("UnmarshalText(%q) = %v, want nil", tt.input, f)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(f) != len(tt.expected) {
|
||||||
|
t.Errorf("UnmarshalText(%q) length = %d, want %d", tt.input, len(f), len(tt.expected))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, v := range tt.expected {
|
||||||
|
if f[i] != v {
|
||||||
|
t.Errorf("UnmarshalText(%q)[%d] = %q, want %q", tt.input, i, f[i], v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFlexibleStringSlice_UnmarshalText_EmptySliceConsistency tests nil vs empty slice behavior
|
||||||
|
func TestFlexibleStringSlice_UnmarshalText_EmptySliceConsistency(t *testing.T) {
|
||||||
|
t.Run("Empty string returns nil", func(t *testing.T) {
|
||||||
|
var f FlexibleStringSlice
|
||||||
|
err := f.UnmarshalText([]byte(""))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("UnmarshalText error = %v", err)
|
||||||
|
}
|
||||||
|
if f != nil {
|
||||||
|
t.Errorf("Empty string should return nil, got %v", f)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Commas only returns empty slice", func(t *testing.T) {
|
||||||
|
var f FlexibleStringSlice
|
||||||
|
err := f.UnmarshalText([]byte(",,,"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("UnmarshalText error = %v", err)
|
||||||
|
}
|
||||||
|
if f == nil {
|
||||||
|
t.Error("Commas only should return empty slice, not nil")
|
||||||
|
}
|
||||||
|
if len(f) != 0 {
|
||||||
|
t.Errorf("Expected empty slice, got %v", f)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -355,6 +355,14 @@ func DefaultConfig() *Config {
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
},
|
},
|
||||||
|
|
||||||
|
// LongCat - https://longcat.chat/platform
|
||||||
|
{
|
||||||
|
ModelName: "LongCat-Flash-Thinking",
|
||||||
|
Model: "longcat/LongCat-Flash-Thinking",
|
||||||
|
APIBase: "https://api.longcat.chat/openai",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
// VLLM (local) - http://localhost:8000
|
// VLLM (local) - http://localhost:8000
|
||||||
{
|
{
|
||||||
ModelName: "local-model",
|
ModelName: "local-model",
|
||||||
|
|
@ -384,6 +392,13 @@ func DefaultConfig() *Config {
|
||||||
Brave: BraveConfig{
|
Brave: BraveConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
|
APIKeys: nil,
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
Tavily: TavilyConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
APIKeys: nil,
|
||||||
MaxResults: 5,
|
MaxResults: 5,
|
||||||
},
|
},
|
||||||
DuckDuckGo: DuckDuckGoConfig{
|
DuckDuckGo: DuckDuckGoConfig{
|
||||||
|
|
@ -393,6 +408,7 @@ func DefaultConfig() *Config {
|
||||||
Perplexity: PerplexityConfig{
|
Perplexity: PerplexityConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
|
APIKeys: nil,
|
||||||
MaxResults: 5,
|
MaxResults: 5,
|
||||||
},
|
},
|
||||||
SearXNG: SearXNGConfig{
|
SearXNG: SearXNGConfig{
|
||||||
|
|
@ -419,6 +435,7 @@ func DefaultConfig() *Config {
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
},
|
},
|
||||||
EnableDenyPatterns: true,
|
EnableDenyPatterns: true,
|
||||||
|
AllowRemote: true,
|
||||||
TimeoutSeconds: 60,
|
TimeoutSeconds: 60,
|
||||||
},
|
},
|
||||||
Skills: SkillsToolsConfig{
|
Skills: SkillsToolsConfig{
|
||||||
|
|
@ -502,5 +519,14 @@ func DefaultConfig() *Config {
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
MonitorUSB: true,
|
MonitorUSB: true,
|
||||||
},
|
},
|
||||||
|
Voice: VoiceConfig{
|
||||||
|
EchoTranscription: false,
|
||||||
|
},
|
||||||
|
BuildInfo: BuildInfo{
|
||||||
|
Version: Version,
|
||||||
|
GitCommit: GitCommit,
|
||||||
|
BuildTime: BuildTime,
|
||||||
|
GoVersion: GoVersion,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -407,6 +407,23 @@ func ConvertProvidersToModelList(cfg *Config) []ModelConfig {
|
||||||
}, true
|
}, true
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"longcat"},
|
||||||
|
protocol: "longcat",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.LongCat.APIKey == "" && p.LongCat.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "longcat",
|
||||||
|
Model: "longcat/LongCat-Flash-Thinking",
|
||||||
|
APIKey: p.LongCat.APIKey,
|
||||||
|
APIBase: p.LongCat.APIBase,
|
||||||
|
Proxy: p.LongCat.Proxy,
|
||||||
|
RequestTimeout: p.LongCat.RequestTimeout,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// Process each provider migration
|
// Process each provider migration
|
||||||
|
|
|
||||||
|
|
@ -162,14 +162,15 @@ func TestConvertProvidersToModelList_AllProviders(t *testing.T) {
|
||||||
Qwen: ProviderConfig{APIKey: "key17"},
|
Qwen: ProviderConfig{APIKey: "key17"},
|
||||||
Mistral: ProviderConfig{APIKey: "key18"},
|
Mistral: ProviderConfig{APIKey: "key18"},
|
||||||
Avian: ProviderConfig{APIKey: "key19"},
|
Avian: ProviderConfig{APIKey: "key19"},
|
||||||
|
LongCat: ProviderConfig{APIKey: "key-longcat"},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
result := ConvertProvidersToModelList(cfg)
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
// All 21 providers should be converted
|
// All 22 providers should be converted
|
||||||
if len(result) != 21 {
|
if len(result) != 22 {
|
||||||
t.Errorf("len(result) = %d, want 21", len(result))
|
t.Errorf("len(result) = %d, want 22", len(result))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
44
pkg/config/version.go
Normal file
44
pkg/config/version.go
Normal file
|
|
@ -0,0 +1,44 @@
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"runtime"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Build-time variables injected via ldflags during build process.
|
||||||
|
// These are set by the Makefile or .goreleaser.yaml using the -X flag:
|
||||||
|
//
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.Version=<version>
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.GitCommit=<commit>
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.BuildTime=<timestamp>
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.GoVersion=<go-version>
|
||||||
|
var (
|
||||||
|
Version = "dev" // Default value when not built with ldflags
|
||||||
|
GitCommit string // Git commit SHA (short)
|
||||||
|
BuildTime string // Build timestamp in RFC3339 format
|
||||||
|
GoVersion string // Go version used for building
|
||||||
|
)
|
||||||
|
|
||||||
|
// FormatVersion returns the version string with optional git commit
|
||||||
|
func FormatVersion() string {
|
||||||
|
v := Version
|
||||||
|
if GitCommit != "" {
|
||||||
|
v += fmt.Sprintf(" (git: %s)", GitCommit)
|
||||||
|
}
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatBuildInfo returns build time and go version info
|
||||||
|
func FormatBuildInfo() (string, string) {
|
||||||
|
build := BuildTime
|
||||||
|
goVer := GoVersion
|
||||||
|
if goVer == "" {
|
||||||
|
goVer = runtime.Version()
|
||||||
|
}
|
||||||
|
return build, goVer
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetVersion returns the version string
|
||||||
|
func GetVersion() string {
|
||||||
|
return Version
|
||||||
|
}
|
||||||
92
pkg/config/version_test.go
Normal file
92
pkg/config/version_test.go
Normal file
|
|
@ -0,0 +1,92 @@
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestFormatVersion_NoGitCommit(t *testing.T) {
|
||||||
|
oldVersion, oldGit := Version, GitCommit
|
||||||
|
t.Cleanup(func() { Version, GitCommit = oldVersion, oldGit })
|
||||||
|
|
||||||
|
Version = "1.2.3"
|
||||||
|
GitCommit = ""
|
||||||
|
|
||||||
|
assert.Equal(t, "1.2.3", FormatVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatVersion_WithGitCommit(t *testing.T) {
|
||||||
|
oldVersion, oldGit := Version, GitCommit
|
||||||
|
t.Cleanup(func() { Version, GitCommit = oldVersion, oldGit })
|
||||||
|
|
||||||
|
Version = "1.2.3"
|
||||||
|
GitCommit = "abc123"
|
||||||
|
|
||||||
|
assert.Equal(t, "1.2.3 (git: abc123)", FormatVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatBuildInfo_UsesBuildTimeAndGoVersion_WhenSet(t *testing.T) {
|
||||||
|
oldBuildTime, oldGoVersion := BuildTime, GoVersion
|
||||||
|
t.Cleanup(func() { BuildTime, GoVersion = oldBuildTime, oldGoVersion })
|
||||||
|
|
||||||
|
BuildTime = "2026-02-20T00:00:00Z"
|
||||||
|
GoVersion = "go1.23.0"
|
||||||
|
|
||||||
|
build, goVer := FormatBuildInfo()
|
||||||
|
|
||||||
|
assert.Equal(t, BuildTime, build)
|
||||||
|
assert.Equal(t, GoVersion, goVer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatBuildInfo_EmptyBuildTime_ReturnsEmptyBuild(t *testing.T) {
|
||||||
|
oldBuildTime, oldGoVersion := BuildTime, GoVersion
|
||||||
|
t.Cleanup(func() { BuildTime, GoVersion = oldBuildTime, oldGoVersion })
|
||||||
|
|
||||||
|
BuildTime = ""
|
||||||
|
GoVersion = "go1.23.0"
|
||||||
|
|
||||||
|
build, goVer := FormatBuildInfo()
|
||||||
|
|
||||||
|
assert.Empty(t, build)
|
||||||
|
assert.Equal(t, GoVersion, goVer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatBuildInfo_EmptyGoVersion_FallsBackToRuntimeVersion(t *testing.T) {
|
||||||
|
oldBuildTime, oldGoVersion := BuildTime, GoVersion
|
||||||
|
t.Cleanup(func() { BuildTime, GoVersion = oldBuildTime, oldGoVersion })
|
||||||
|
|
||||||
|
BuildTime = "x"
|
||||||
|
GoVersion = ""
|
||||||
|
|
||||||
|
build, goVer := FormatBuildInfo()
|
||||||
|
|
||||||
|
assert.Equal(t, "x", build)
|
||||||
|
assert.Equal(t, runtime.Version(), goVer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetVersion(t *testing.T) {
|
||||||
|
oldVersion := Version
|
||||||
|
t.Cleanup(func() { Version = oldVersion })
|
||||||
|
|
||||||
|
Version = "dev"
|
||||||
|
assert.Equal(t, "dev", GetVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetVersion_Custom(t *testing.T) {
|
||||||
|
oldVersion := Version
|
||||||
|
t.Cleanup(func() { Version = oldVersion })
|
||||||
|
|
||||||
|
Version = "v1.0.0"
|
||||||
|
assert.Equal(t, "v1.0.0", GetVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestVersion_DefaultIsDev(t *testing.T) {
|
||||||
|
// Reset to default values
|
||||||
|
oldVersion := Version
|
||||||
|
Version = "dev"
|
||||||
|
t.Cleanup(func() { Version = oldVersion })
|
||||||
|
|
||||||
|
assert.Equal(t, "dev", Version)
|
||||||
|
}
|
||||||
|
|
@ -1,24 +1,24 @@
|
||||||
package logger
|
package logger
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
|
||||||
"os"
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"runtime"
|
"runtime"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
|
||||||
|
"github.com/rs/zerolog"
|
||||||
)
|
)
|
||||||
|
|
||||||
type LogLevel int
|
type LogLevel = zerolog.Level
|
||||||
|
|
||||||
const (
|
const (
|
||||||
DEBUG LogLevel = iota
|
DEBUG = zerolog.DebugLevel
|
||||||
INFO
|
INFO = zerolog.InfoLevel
|
||||||
WARN
|
WARN = zerolog.WarnLevel
|
||||||
ERROR
|
ERROR = zerolog.ErrorLevel
|
||||||
FATAL
|
FATAL = zerolog.FatalLevel
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|
@ -31,27 +31,24 @@ var (
|
||||||
}
|
}
|
||||||
|
|
||||||
currentLevel = INFO
|
currentLevel = INFO
|
||||||
logger *Logger
|
logger zerolog.Logger
|
||||||
|
fileLogger zerolog.Logger
|
||||||
|
logFile *os.File
|
||||||
once sync.Once
|
once sync.Once
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
)
|
)
|
||||||
|
|
||||||
type Logger struct {
|
|
||||||
file *os.File
|
|
||||||
}
|
|
||||||
|
|
||||||
type LogEntry struct {
|
|
||||||
Level string `json:"level"`
|
|
||||||
Timestamp string `json:"timestamp"`
|
|
||||||
Component string `json:"component,omitempty"`
|
|
||||||
Message string `json:"message"`
|
|
||||||
Fields map[string]any `json:"fields,omitempty"`
|
|
||||||
Caller string `json:"caller,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
once.Do(func() {
|
once.Do(func() {
|
||||||
logger = &Logger{}
|
zerolog.SetGlobalLevel(zerolog.InfoLevel)
|
||||||
|
|
||||||
|
consoleWriter := zerolog.ConsoleWriter{
|
||||||
|
Out: os.Stdout,
|
||||||
|
TimeFormat: "15:04:05", // TODO: make it configurable???
|
||||||
|
}
|
||||||
|
|
||||||
|
logger = zerolog.New(consoleWriter).With().Timestamp().Logger()
|
||||||
|
fileLogger = zerolog.Logger{}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -59,6 +56,7 @@ func SetLevel(level LogLevel) {
|
||||||
mu.Lock()
|
mu.Lock()
|
||||||
defer mu.Unlock()
|
defer mu.Unlock()
|
||||||
currentLevel = level
|
currentLevel = level
|
||||||
|
zerolog.SetGlobalLevel(level)
|
||||||
}
|
}
|
||||||
|
|
||||||
func GetLevel() LogLevel {
|
func GetLevel() LogLevel {
|
||||||
|
|
@ -71,17 +69,22 @@ func EnableFileLogging(filePath string) error {
|
||||||
mu.Lock()
|
mu.Lock()
|
||||||
defer mu.Unlock()
|
defer mu.Unlock()
|
||||||
|
|
||||||
file, err := os.OpenFile(filePath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
|
if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil {
|
||||||
|
return fmt.Errorf("failed to create log directory: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
newFile, err := os.OpenFile(filePath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to open log file: %w", err)
|
return fmt.Errorf("failed to open log file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if logger.file != nil {
|
// Close old file if exists
|
||||||
logger.file.Close()
|
if logFile != nil {
|
||||||
|
logFile.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.file = file
|
logFile = newFile
|
||||||
log.Println("File logging enabled:", filePath)
|
fileLogger = zerolog.New(logFile).With().Timestamp().Caller().Logger()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -89,10 +92,57 @@ func DisableFileLogging() {
|
||||||
mu.Lock()
|
mu.Lock()
|
||||||
defer mu.Unlock()
|
defer mu.Unlock()
|
||||||
|
|
||||||
if logger.file != nil {
|
if logFile != nil {
|
||||||
logger.file.Close()
|
logFile.Close()
|
||||||
logger.file = nil
|
logFile = nil
|
||||||
log.Println("File logging disabled")
|
}
|
||||||
|
fileLogger = zerolog.Logger{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func getCallerInfo() (string, int, string) {
|
||||||
|
for i := 2; i < 15; i++ {
|
||||||
|
pc, file, line, ok := runtime.Caller(i)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fn := runtime.FuncForPC(pc)
|
||||||
|
if fn == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// bypass common loggers
|
||||||
|
if strings.HasSuffix(file, "/logger.go") ||
|
||||||
|
strings.HasSuffix(file, "/log.go") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
funcName := fn.Name()
|
||||||
|
if strings.HasPrefix(funcName, "runtime.") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
return filepath.Base(file), line, filepath.Base(funcName)
|
||||||
|
}
|
||||||
|
|
||||||
|
return "???", 0, "???"
|
||||||
|
}
|
||||||
|
|
||||||
|
//nolint:zerologlint
|
||||||
|
func getEvent(logger zerolog.Logger, level LogLevel) *zerolog.Event {
|
||||||
|
switch level {
|
||||||
|
case zerolog.DebugLevel:
|
||||||
|
return logger.Debug()
|
||||||
|
case zerolog.InfoLevel:
|
||||||
|
return logger.Info()
|
||||||
|
case zerolog.WarnLevel:
|
||||||
|
return logger.Warn()
|
||||||
|
case zerolog.ErrorLevel:
|
||||||
|
return logger.Error()
|
||||||
|
case zerolog.FatalLevel:
|
||||||
|
return logger.Fatal()
|
||||||
|
default:
|
||||||
|
return logger.Info()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -101,65 +151,41 @@ func logMessage(level LogLevel, component string, message string, fields map[str
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
entry := LogEntry{
|
callerFile, callerLine, callerFunc := getCallerInfo()
|
||||||
Level: logLevelNames[level],
|
|
||||||
Timestamp: time.Now().UTC().Format(time.RFC3339),
|
|
||||||
Component: component,
|
|
||||||
Message: message,
|
|
||||||
Fields: fields,
|
|
||||||
}
|
|
||||||
|
|
||||||
if pc, file, line, ok := runtime.Caller(2); ok {
|
event := getEvent(logger, level)
|
||||||
fn := runtime.FuncForPC(pc)
|
|
||||||
if fn != nil {
|
|
||||||
entry.Caller = fmt.Sprintf("%s:%d (%s)", file, line, fn.Name())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if logger.file != nil {
|
// Build combined field with component and caller
|
||||||
jsonData, err := json.Marshal(entry)
|
if component != "" {
|
||||||
if err == nil {
|
event.Str("caller", fmt.Sprintf("%-6s %s:%d (%s)", component, callerFile, callerLine, callerFunc))
|
||||||
logger.file.Write(append(jsonData, '\n'))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var fieldStr string
|
|
||||||
if len(fields) > 0 {
|
|
||||||
fieldStr = " " + formatFields(fields)
|
|
||||||
} else {
|
} else {
|
||||||
fieldStr = ""
|
event.Str("caller", fmt.Sprintf("<none> %s:%d (%s)", callerFile, callerLine, callerFunc))
|
||||||
}
|
}
|
||||||
|
|
||||||
logLine := fmt.Sprintf("[%s] [%s]%s %s%s",
|
for k, v := range fields {
|
||||||
entry.Timestamp,
|
event.Interface(k, v)
|
||||||
logLevelNames[level],
|
}
|
||||||
formatComponent(component),
|
|
||||||
message,
|
|
||||||
fieldStr,
|
|
||||||
)
|
|
||||||
|
|
||||||
log.Println(logLine)
|
event.Msg(message)
|
||||||
|
|
||||||
|
// Also log to file if enabled
|
||||||
|
if fileLogger.GetLevel() != zerolog.NoLevel {
|
||||||
|
fileEvent := getEvent(fileLogger, level)
|
||||||
|
|
||||||
|
if component != "" {
|
||||||
|
fileEvent.Str("component", component)
|
||||||
|
}
|
||||||
|
for k, v := range fields {
|
||||||
|
fileEvent.Interface(k, v)
|
||||||
|
}
|
||||||
|
fileEvent.Msg(message)
|
||||||
|
}
|
||||||
|
|
||||||
if level == FATAL {
|
if level == FATAL {
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func formatComponent(component string) string {
|
|
||||||
if component == "" {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return fmt.Sprintf(" %s:", component)
|
|
||||||
}
|
|
||||||
|
|
||||||
func formatFields(fields map[string]any) string {
|
|
||||||
parts := make([]string, 0, len(fields))
|
|
||||||
for k, v := range fields {
|
|
||||||
parts = append(parts, fmt.Sprintf("%s=%v", k, v))
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("{%s}", strings.Join(parts, ", "))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Debug(message string) {
|
func Debug(message string) {
|
||||||
logMessage(DEBUG, "", message, nil)
|
logMessage(DEBUG, "", message, nil)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
95
pkg/logger/logger_3rd_party.go
Normal file
95
pkg/logger/logger_3rd_party.go
Normal file
|
|
@ -0,0 +1,95 @@
|
||||||
|
// this file is for compatible with 3rd party loggers, should not be called in PicoClaw project
|
||||||
|
|
||||||
|
package logger
|
||||||
|
|
||||||
|
import "fmt"
|
||||||
|
|
||||||
|
// Logger implements common Logger interface
|
||||||
|
type Logger struct {
|
||||||
|
component string
|
||||||
|
levels map[int]LogLevel
|
||||||
|
}
|
||||||
|
|
||||||
|
// Debug logs debug messages
|
||||||
|
func (b *Logger) Debug(v ...any) {
|
||||||
|
logMessage(DEBUG, b.component, fmt.Sprint(v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Info logs info messages
|
||||||
|
func (b *Logger) Info(v ...any) {
|
||||||
|
logMessage(INFO, b.component, fmt.Sprint(v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warn logs warning messages
|
||||||
|
func (b *Logger) Warn(v ...any) {
|
||||||
|
logMessage(WARN, b.component, fmt.Sprint(v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Error logs error messages
|
||||||
|
func (b *Logger) Error(v ...any) {
|
||||||
|
logMessage(ERROR, b.component, fmt.Sprint(v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Debugf logs formatted debug messages
|
||||||
|
func (b *Logger) Debugf(format string, v ...any) {
|
||||||
|
logMessage(DEBUG, b.component, fmt.Sprintf(format, v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Infof logs formatted info messages
|
||||||
|
func (b *Logger) Infof(format string, v ...any) {
|
||||||
|
logMessage(INFO, b.component, fmt.Sprintf(format, v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warnf logs formatted warning messages
|
||||||
|
func (b *Logger) Warnf(format string, v ...any) {
|
||||||
|
logMessage(WARN, b.component, fmt.Sprintf(format, v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warningf logs formatted warning messages
|
||||||
|
func (b *Logger) Warningf(format string, v ...any) {
|
||||||
|
logMessage(WARN, b.component, fmt.Sprintf(format, v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Errorf logs formatted error messages
|
||||||
|
func (b *Logger) Errorf(format string, v ...any) {
|
||||||
|
logMessage(ERROR, b.component, fmt.Sprintf(format, v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fatalf logs formatted fatal messages and exits
|
||||||
|
func (b *Logger) Fatalf(format string, v ...any) {
|
||||||
|
logMessage(FATAL, b.component, fmt.Sprintf(format, v...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Log logs a message at a given level with caller information
|
||||||
|
// the func name must be this because 3rd party loggers expect this
|
||||||
|
// msgL: message level (DEBUG, INFO, WARN, ERROR, FATAL)
|
||||||
|
// caller: unused parameter reserved for compatibility
|
||||||
|
// format: format string
|
||||||
|
// a: format arguments
|
||||||
|
//
|
||||||
|
//nolint:goprintffuncname
|
||||||
|
func (b *Logger) Log(msgL, caller int, format string, a ...any) {
|
||||||
|
level := LogLevel(msgL)
|
||||||
|
if b.levels != nil {
|
||||||
|
if lvl, ok := b.levels[msgL]; ok {
|
||||||
|
level = lvl
|
||||||
|
}
|
||||||
|
}
|
||||||
|
logMessage(level, b.component, fmt.Sprintf(format, a...), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sync flushes log buffer (no-op for this implementation)
|
||||||
|
func (b *Logger) Sync() error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithLevels sets log levels mapping for this logger
|
||||||
|
func (b *Logger) WithLevels(levels map[int]LogLevel) *Logger {
|
||||||
|
b.levels = levels
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewLogger creates a new logger instance with optional component name
|
||||||
|
func NewLogger(component string) *Logger {
|
||||||
|
return &Logger{component: component}
|
||||||
|
}
|
||||||
|
|
@ -86,14 +86,14 @@ func (s *JSONLStore) metaPath(key string) string {
|
||||||
|
|
||||||
// sanitizeKey converts a session key to a safe filename component.
|
// sanitizeKey converts a session key to a safe filename component.
|
||||||
// Mirrors pkg/session.sanitizeFilename so that migration paths match.
|
// Mirrors pkg/session.sanitizeFilename so that migration paths match.
|
||||||
//
|
// Replaces ':' with '_' (session key separator) and '/' and '\' with '_'
|
||||||
// Note: this is a lossy mapping — "telegram:123" and "telegram_123"
|
// so composite IDs (e.g. Telegram forum "chatID/threadID", Slack "channel/thread_ts")
|
||||||
// both produce the same filename. This is an intentional tradeoff:
|
// do not create subdirectories or break on Windows.
|
||||||
// keys with colons (e.g. from channels) are by far the common case,
|
|
||||||
// and a bidirectional encoding (like URL-encoding) would complicate
|
|
||||||
// file listings and debugging.
|
|
||||||
func sanitizeKey(key string) string {
|
func sanitizeKey(key string) string {
|
||||||
return strings.ReplaceAll(key, ":", "_")
|
s := strings.ReplaceAll(key, ":", "_")
|
||||||
|
s = strings.ReplaceAll(s, "/", "_")
|
||||||
|
s = strings.ReplaceAll(s, "\\", "_")
|
||||||
|
return s
|
||||||
}
|
}
|
||||||
|
|
||||||
// readMeta loads the metadata file for a session.
|
// readMeta loads the metadata file for a session.
|
||||||
|
|
|
||||||
|
|
@ -48,6 +48,12 @@ func MigrateFromJSON(
|
||||||
if !strings.HasSuffix(name, ".json") {
|
if !strings.HasSuffix(name, ".json") {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
// Skip JSONL metadata files. They are part of the new storage format,
|
||||||
|
// not legacy session snapshots, and re-importing them would overwrite
|
||||||
|
// the paired .jsonl history with an empty message list.
|
||||||
|
if strings.HasSuffix(name, ".meta.json") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
// Skip already-migrated files.
|
// Skip already-migrated files.
|
||||||
if strings.HasSuffix(name, ".migrated") {
|
if strings.HasSuffix(name, ".migrated") {
|
||||||
continue
|
continue
|
||||||
|
|
|
||||||
|
|
@ -382,3 +382,55 @@ func TestMigrateFromJSON_NonexistentDir(t *testing.T) {
|
||||||
t.Errorf("expected 0, got %d", count)
|
t.Errorf("expected 0, got %d", count)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMigrateFromJSON_SkipsMetaJSONFiles(t *testing.T) {
|
||||||
|
sessionsDir := t.TempDir()
|
||||||
|
store, err := NewJSONLStore(sessionsDir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewJSONLStore: %v", err)
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
if addErr := store.AddMessage(ctx, "agent:main:pico:direct:pico:test", "user", "keep me"); addErr != nil {
|
||||||
|
t.Fatalf("AddMessage: %v", addErr)
|
||||||
|
}
|
||||||
|
if summaryErr := store.SetSummary(ctx, "agent:main:pico:direct:pico:test", "keep summary"); summaryErr != nil {
|
||||||
|
t.Fatalf("SetSummary: %v", summaryErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
metaPath := filepath.Join(sessionsDir, "agent_main_pico_direct_pico_test.meta.json")
|
||||||
|
if _, statErr := os.Stat(metaPath); statErr != nil {
|
||||||
|
t.Fatalf("meta file missing before migration: %v", statErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
count, err := MigrateFromJSON(ctx, sessionsDir, store)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("MigrateFromJSON: %v", err)
|
||||||
|
}
|
||||||
|
if count != 0 {
|
||||||
|
t.Fatalf("expected 0 migrated, got %d", count)
|
||||||
|
}
|
||||||
|
|
||||||
|
history, err := store.GetHistory(ctx, "agent:main:pico:direct:pico:test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetHistory: %v", err)
|
||||||
|
}
|
||||||
|
if len(history) != 1 || history[0].Content != "keep me" {
|
||||||
|
t.Fatalf("history = %+v, want preserved single message", history)
|
||||||
|
}
|
||||||
|
|
||||||
|
summary, err := store.GetSummary(ctx, "agent:main:pico:direct:pico:test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetSummary: %v", err)
|
||||||
|
}
|
||||||
|
if summary != "keep summary" {
|
||||||
|
t.Fatalf("summary = %q, want %q", summary, "keep summary")
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, statErr := os.Stat(metaPath); statErr != nil {
|
||||||
|
t.Fatalf("meta file should remain in place: %v", statErr)
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(metaPath + ".migrated"); !os.IsNotExist(statErr) {
|
||||||
|
t.Fatalf("meta file should not be renamed, stat err = %v", statErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,6 @@ var migrateableFiles = []string{
|
||||||
"AGENTS.md",
|
"AGENTS.md",
|
||||||
"SOUL.md",
|
"SOUL.md",
|
||||||
"USER.md",
|
"USER.md",
|
||||||
"TOOLS.md",
|
|
||||||
"HEARTBEAT.md",
|
"HEARTBEAT.md",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -735,12 +735,14 @@ type WebToolsConfig struct {
|
||||||
type BraveConfig struct {
|
type BraveConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
APIKey string `json:"api_key"`
|
APIKey string `json:"api_key"`
|
||||||
|
APIKeys []string `json:"api_keys"`
|
||||||
MaxResults int `json:"max_results"`
|
MaxResults int `json:"max_results"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type TavilyConfig struct {
|
type TavilyConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
APIKey string `json:"api_key"`
|
APIKey string `json:"api_key"`
|
||||||
|
APIKeys []string `json:"api_keys"`
|
||||||
BaseURL string `json:"base_url"`
|
BaseURL string `json:"base_url"`
|
||||||
MaxResults int `json:"max_results"`
|
MaxResults int `json:"max_results"`
|
||||||
}
|
}
|
||||||
|
|
@ -753,6 +755,7 @@ type DuckDuckGoConfig struct {
|
||||||
type PerplexityConfig struct {
|
type PerplexityConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
APIKey string `json:"api_key"`
|
APIKey string `json:"api_key"`
|
||||||
|
APIKeys []string `json:"api_keys"`
|
||||||
MaxResults int `json:"max_results"`
|
MaxResults int `json:"max_results"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1082,6 +1085,7 @@ func (c ToolsConfig) ToStandardTools() config.ToolsConfig {
|
||||||
Brave: config.BraveConfig{
|
Brave: config.BraveConfig{
|
||||||
Enabled: c.Web.Brave.Enabled,
|
Enabled: c.Web.Brave.Enabled,
|
||||||
APIKey: c.Web.Brave.APIKey,
|
APIKey: c.Web.Brave.APIKey,
|
||||||
|
APIKeys: c.Web.Brave.APIKeys,
|
||||||
MaxResults: c.Web.Brave.MaxResults,
|
MaxResults: c.Web.Brave.MaxResults,
|
||||||
},
|
},
|
||||||
Tavily: config.TavilyConfig{
|
Tavily: config.TavilyConfig{
|
||||||
|
|
@ -1107,6 +1111,7 @@ func (c ToolsConfig) ToStandardTools() config.ToolsConfig {
|
||||||
Exec: config.ExecConfig{
|
Exec: config.ExecConfig{
|
||||||
EnableDenyPatterns: c.Exec.EnableDenyPatterns,
|
EnableDenyPatterns: c.Exec.EnableDenyPatterns,
|
||||||
CustomDenyPatterns: c.Exec.CustomDenyPatterns,
|
CustomDenyPatterns: c.Exec.CustomDenyPatterns,
|
||||||
|
AllowRemote: config.DefaultConfig().Tools.Exec.AllowRemote,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -290,6 +290,20 @@ func TestConvertToPicoClaw(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToStandardConfig_ExecAllowRemoteDefaultsTrue(t *testing.T) {
|
||||||
|
cfg := (&PicoClawConfig{
|
||||||
|
Tools: ToolsConfig{
|
||||||
|
Exec: ExecConfig{
|
||||||
|
EnableDenyPatterns: true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}).ToStandardConfig()
|
||||||
|
|
||||||
|
if !cfg.Tools.Exec.AllowRemote {
|
||||||
|
t.Fatal("ToStandardConfig() should preserve the default tools.exec.allow_remote=true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestConvertToPicoClawWithQQAndDingTalk(t *testing.T) {
|
func TestConvertToPicoClawWithQQAndDingTalk(t *testing.T) {
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
configPath := filepath.Join(tmpDir, "openclaw.json")
|
configPath := filepath.Join(tmpDir, "openclaw.json")
|
||||||
|
|
|
||||||
|
|
@ -40,6 +40,10 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
||||||
lowerModel := strings.ToLower(model)
|
lowerModel := strings.ToLower(model)
|
||||||
|
|
||||||
|
if providerName == "" && model == "" {
|
||||||
|
return providerSelection{}, fmt.Errorf("no model configured: agents.defaults.model is empty")
|
||||||
|
}
|
||||||
|
|
||||||
sel := providerSelection{
|
sel := providerSelection{
|
||||||
providerType: providerTypeHTTPCompat,
|
providerType: providerTypeHTTPCompat,
|
||||||
model: model,
|
model: model,
|
||||||
|
|
@ -217,6 +221,15 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
sel.apiBase = "https://api.minimaxi.com/v1"
|
sel.apiBase = "https://api.minimaxi.com/v1"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
case "longcat":
|
||||||
|
if cfg.Providers.LongCat.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.LongCat.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.LongCat.APIBase
|
||||||
|
sel.proxy = cfg.Providers.LongCat.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.longcat.chat/openai"
|
||||||
|
}
|
||||||
|
}
|
||||||
case "github_copilot", "copilot":
|
case "github_copilot", "copilot":
|
||||||
sel.providerType = providerTypeGitHubCopilot
|
sel.providerType = providerTypeGitHubCopilot
|
||||||
if cfg.Providers.GitHubCopilot.APIBase != "" {
|
if cfg.Providers.GitHubCopilot.APIBase != "" {
|
||||||
|
|
@ -348,6 +361,13 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
if sel.apiBase == "" {
|
if sel.apiBase == "" {
|
||||||
sel.apiBase = "https://api.avian.io/v1"
|
sel.apiBase = "https://api.avian.io/v1"
|
||||||
}
|
}
|
||||||
|
case (strings.Contains(lowerModel, "longcat") || strings.HasPrefix(model, "longcat/")) && cfg.Providers.LongCat.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.LongCat.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.LongCat.APIBase
|
||||||
|
sel.proxy = cfg.Providers.LongCat.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.longcat.chat/openai"
|
||||||
|
}
|
||||||
case cfg.Providers.VLLM.APIBase != "":
|
case cfg.Providers.VLLM.APIBase != "":
|
||||||
sel.apiKey = cfg.Providers.VLLM.APIKey
|
sel.apiKey = cfg.Providers.VLLM.APIKey
|
||||||
sel.apiBase = cfg.Providers.VLLM.APIBase
|
sel.apiBase = cfg.Providers.VLLM.APIBase
|
||||||
|
|
|
||||||
|
|
@ -95,7 +95,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
"vivgrid", "volcengine", "vllm", "qwen", "mistral", "avian",
|
"vivgrid", "volcengine", "vllm", "qwen", "mistral", "avian",
|
||||||
"minimax":
|
"minimax", "longcat":
|
||||||
// All other OpenAI-compatible HTTP providers
|
// All other OpenAI-compatible HTTP providers
|
||||||
if cfg.APIKey == "" && cfg.APIBase == "" {
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
||||||
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
|
|
@ -215,6 +215,8 @@ func getDefaultAPIBase(protocol string) string {
|
||||||
return "https://api.avian.io/v1"
|
return "https://api.avian.io/v1"
|
||||||
case "minimax":
|
case "minimax":
|
||||||
return "https://api.minimaxi.com/v1"
|
return "https://api.minimaxi.com/v1"
|
||||||
|
case "longcat":
|
||||||
|
return "https://api.longcat.chat/openai"
|
||||||
default:
|
default:
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -113,6 +113,7 @@ func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
|
||||||
{"vllm", "vllm"},
|
{"vllm", "vllm"},
|
||||||
{"deepseek", "deepseek"},
|
{"deepseek", "deepseek"},
|
||||||
{"ollama", "ollama"},
|
{"ollama", "ollama"},
|
||||||
|
{"longcat", "longcat"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
|
|
@ -162,6 +163,29 @@ func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_LongCat(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-longcat",
|
||||||
|
Model: "longcat/LongCat-Flash-Thinking",
|
||||||
|
APIKey: "test-key",
|
||||||
|
APIBase: "https://api.longcat.chat/openai",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "LongCat-Flash-Thinking" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "LongCat-Flash-Thinking")
|
||||||
|
}
|
||||||
|
if _, ok := provider.(*HTTPProvider); !ok {
|
||||||
|
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_Anthropic(t *testing.T) {
|
func TestCreateProviderFromConfig_Anthropic(t *testing.T) {
|
||||||
cfg := &config.ModelConfig{
|
cfg := &config.ModelConfig{
|
||||||
ModelName: "test-anthropic",
|
ModelName: "test-anthropic",
|
||||||
|
|
|
||||||
|
|
@ -178,6 +178,26 @@ func TestResolveProviderSelection(t *testing.T) {
|
||||||
wantAPIBase: "https://api.moonshot.cn/v1",
|
wantAPIBase: "https://api.moonshot.cn/v1",
|
||||||
wantProxy: "http://127.0.0.1:7890",
|
wantProxy: "http://127.0.0.1:7890",
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
name: "explicit longcat provider uses defaults",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "longcat"
|
||||||
|
cfg.Providers.LongCat.APIKey = "longcat-key"
|
||||||
|
cfg.Providers.LongCat.Proxy = "http://127.0.0.1:7890"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://api.longcat.chat/openai",
|
||||||
|
wantProxy: "http://127.0.0.1:7890",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "longcat model fallback uses longcat base default",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "longcat/LongCat-Flash-Thinking"
|
||||||
|
cfg.Providers.LongCat.APIKey = "longcat-key"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://api.longcat.chat/openai",
|
||||||
|
},
|
||||||
{
|
{
|
||||||
name: "missing keys returns model config error",
|
name: "missing keys returns model config error",
|
||||||
setup: func(cfg *config.Config) {
|
setup: func(cfg *config.Config) {
|
||||||
|
|
|
||||||
|
|
@ -156,9 +156,10 @@ func (p *Provider) Chat(
|
||||||
// The key is typically the agent ID — stable per agent, shared across requests.
|
// The key is typically the agent ID — stable per agent, shared across requests.
|
||||||
// See: https://platform.openai.com/docs/guides/prompt-caching
|
// See: https://platform.openai.com/docs/guides/prompt-caching
|
||||||
// Prompt caching is only supported by OpenAI-native endpoints.
|
// Prompt caching is only supported by OpenAI-native endpoints.
|
||||||
// Gemini and other providers reject unknown fields, so skip for non-OpenAI APIs.
|
// Non-OpenAI providers (Mistral, Gemini, DeepSeek, etc.) reject unknown
|
||||||
|
// fields with 422 errors, so only include it for OpenAI APIs.
|
||||||
if cacheKey, ok := options["prompt_cache_key"].(string); ok && cacheKey != "" {
|
if cacheKey, ok := options["prompt_cache_key"].(string); ok && cacheKey != "" {
|
||||||
if !strings.Contains(p.apiBase, "generativelanguage.googleapis.com") {
|
if supportsPromptCacheKey(p.apiBase) {
|
||||||
requestBody["prompt_cache_key"] = cacheKey
|
requestBody["prompt_cache_key"] = cacheKey
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -284,7 +285,7 @@ func parseResponse(body io.Reader) (*LLMResponse, error) {
|
||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
Function *struct {
|
Function *struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Arguments string `json:"arguments"`
|
Arguments json.RawMessage `json:"arguments"`
|
||||||
} `json:"function"`
|
} `json:"function"`
|
||||||
ExtraContent *struct {
|
ExtraContent *struct {
|
||||||
Google *struct {
|
Google *struct {
|
||||||
|
|
@ -323,12 +324,7 @@ func parseResponse(body io.Reader) (*LLMResponse, error) {
|
||||||
|
|
||||||
if tc.Function != nil {
|
if tc.Function != nil {
|
||||||
name = tc.Function.Name
|
name = tc.Function.Name
|
||||||
if tc.Function.Arguments != "" {
|
arguments = decodeToolCallArguments(tc.Function.Arguments, name)
|
||||||
if err := json.Unmarshal([]byte(tc.Function.Arguments), &arguments); err != nil {
|
|
||||||
log.Printf("openai_compat: failed to decode tool call arguments for %q: %v", name, err)
|
|
||||||
arguments["raw"] = tc.Function.Arguments
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build ToolCall with ExtraContent for Gemini 3 thought_signature persistence
|
// Build ToolCall with ExtraContent for Gemini 3 thought_signature persistence
|
||||||
|
|
@ -361,6 +357,39 @@ func parseResponse(body io.Reader) (*LLMResponse, error) {
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func decodeToolCallArguments(raw json.RawMessage, name string) map[string]any {
|
||||||
|
arguments := make(map[string]any)
|
||||||
|
raw = bytes.TrimSpace(raw)
|
||||||
|
if len(raw) == 0 || bytes.Equal(raw, []byte("null")) {
|
||||||
|
return arguments
|
||||||
|
}
|
||||||
|
|
||||||
|
var decoded any
|
||||||
|
if err := json.Unmarshal(raw, &decoded); err != nil {
|
||||||
|
log.Printf("openai_compat: failed to decode tool call arguments payload for %q: %v", name, err)
|
||||||
|
arguments["raw"] = string(raw)
|
||||||
|
return arguments
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v := decoded.(type) {
|
||||||
|
case string:
|
||||||
|
if strings.TrimSpace(v) == "" {
|
||||||
|
return arguments
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(v), &arguments); err != nil {
|
||||||
|
log.Printf("openai_compat: failed to decode tool call arguments for %q: %v", name, err)
|
||||||
|
arguments["raw"] = v
|
||||||
|
}
|
||||||
|
return arguments
|
||||||
|
case map[string]any:
|
||||||
|
return v
|
||||||
|
default:
|
||||||
|
log.Printf("openai_compat: unsupported tool call arguments type for %q: %T", name, decoded)
|
||||||
|
arguments["raw"] = string(raw)
|
||||||
|
return arguments
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// openaiMessage is the wire-format message for OpenAI-compatible APIs.
|
// openaiMessage is the wire-format message for OpenAI-compatible APIs.
|
||||||
// It mirrors protocoltypes.Message but omits SystemParts, which is an
|
// It mirrors protocoltypes.Message but omits SystemParts, which is an
|
||||||
// internal field that would be unknown to third-party endpoints.
|
// internal field that would be unknown to third-party endpoints.
|
||||||
|
|
@ -476,3 +505,16 @@ func asFloat(v any) (float64, bool) {
|
||||||
return 0, false
|
return 0, false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// supportsPromptCacheKey reports whether the given API base is known to
|
||||||
|
// support the prompt_cache_key request field. Currently only OpenAI's own
|
||||||
|
// API and Azure OpenAI support this. All other OpenAI-compatible providers
|
||||||
|
// (Mistral, Gemini, DeepSeek, Groq, etc.) reject unknown fields with 422 errors.
|
||||||
|
func supportsPromptCacheKey(apiBase string) bool {
|
||||||
|
u, err := url.Parse(apiBase)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
host := u.Hostname()
|
||||||
|
return host == "api.openai.com" || strings.HasSuffix(host, ".openai.azure.com")
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -108,6 +108,55 @@ func TestProviderChat_ParsesToolCalls(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_ParsesToolCallsWithObjectArguments(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
resp := map[string]any{
|
||||||
|
"choices": []map[string]any{
|
||||||
|
{
|
||||||
|
"message": map[string]any{
|
||||||
|
"content": "",
|
||||||
|
"tool_calls": []map[string]any{
|
||||||
|
{
|
||||||
|
"id": "call_1",
|
||||||
|
"type": "function",
|
||||||
|
"function": map[string]any{
|
||||||
|
"name": "get_weather",
|
||||||
|
"arguments": map[string]any{
|
||||||
|
"city": "SF",
|
||||||
|
"metric": true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"finish_reason": "tool_calls",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "gpt-4o", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(out.ToolCalls) != 1 {
|
||||||
|
t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
|
||||||
|
}
|
||||||
|
if out.ToolCalls[0].Name != "get_weather" {
|
||||||
|
t.Fatalf("ToolCalls[0].Name = %q, want %q", out.ToolCalls[0].Name, "get_weather")
|
||||||
|
}
|
||||||
|
if out.ToolCalls[0].Arguments["city"] != "SF" {
|
||||||
|
t.Fatalf("ToolCalls[0].Arguments[city] = %v, want SF", out.ToolCalls[0].Arguments["city"])
|
||||||
|
}
|
||||||
|
if out.ToolCalls[0].Arguments["metric"] != true {
|
||||||
|
t.Fatalf("ToolCalls[0].Arguments[metric] = %v, want true", out.ToolCalls[0].Arguments["metric"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestProviderChat_ParsesReasoningContent(t *testing.T) {
|
func TestProviderChat_ParsesReasoningContent(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
resp := map[string]any{
|
resp := map[string]any{
|
||||||
|
|
@ -669,6 +718,111 @@ func TestSerializeMessages_MediaWithToolCallID(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// chatWithCacheKey sets up a test server, sends a Chat request with prompt_cache_key,
|
||||||
|
// and returns the decoded request body for assertion.
|
||||||
|
func chatWithCacheKey(t *testing.T, apiBase string) map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
var requestBody map[string]any
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resp := map[string]any{
|
||||||
|
"choices": []map[string]any{
|
||||||
|
{
|
||||||
|
"message": map[string]any{"content": "ok"},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
p.apiBase = apiBase
|
||||||
|
p.httpClient = &http.Client{
|
||||||
|
Transport: roundTripperFunc(func(r *http.Request) (*http.Response, error) {
|
||||||
|
r.URL, _ = url.Parse(server.URL + r.URL.Path)
|
||||||
|
return http.DefaultTransport.RoundTrip(r)
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := p.Chat(
|
||||||
|
t.Context(),
|
||||||
|
[]Message{{Role: "user", Content: "hi"}},
|
||||||
|
nil,
|
||||||
|
"test-model",
|
||||||
|
map[string]any{"prompt_cache_key": "agent-main"},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
return requestBody
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_PromptCacheKeySentToOpenAI(t *testing.T) {
|
||||||
|
body := chatWithCacheKey(t, "https://api.openai.com/v1")
|
||||||
|
if body["prompt_cache_key"] != "agent-main" {
|
||||||
|
t.Fatalf("prompt_cache_key = %v, want %q", body["prompt_cache_key"], "agent-main")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_PromptCacheKeyOmittedForNonOpenAI(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
apiBase string
|
||||||
|
}{
|
||||||
|
{"mistral", "https://api.mistral.ai/v1"},
|
||||||
|
{"gemini", "https://generativelanguage.googleapis.com/v1beta"},
|
||||||
|
{"deepseek", "https://api.deepseek.com/v1"},
|
||||||
|
{"groq", "https://api.groq.com/openai/v1"},
|
||||||
|
{"minimax", "https://api.minimaxi.com/v1"},
|
||||||
|
{"ollama_local", "http://localhost:11434/v1"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
body := chatWithCacheKey(t, tt.apiBase)
|
||||||
|
if _, exists := body["prompt_cache_key"]; exists {
|
||||||
|
t.Fatalf("prompt_cache_key should NOT be sent to %s, but was included in request", tt.name)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSupportsPromptCacheKey(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
apiBase string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"https://api.openai.com/v1", true},
|
||||||
|
{"https://api.openai.com/v1/", true},
|
||||||
|
{"https://myresource.openai.azure.com/openai/deployments/gpt-4", true},
|
||||||
|
{"https://eastus.openai.azure.com/v1", true},
|
||||||
|
{"https://api.mistral.ai/v1", false},
|
||||||
|
{"https://generativelanguage.googleapis.com/v1beta", false},
|
||||||
|
{"https://api.deepseek.com/v1", false},
|
||||||
|
{"https://api.groq.com/openai/v1", false},
|
||||||
|
{"http://localhost:11434/v1", false},
|
||||||
|
{"https://openrouter.ai/api/v1", false},
|
||||||
|
// Edge cases: proxy URLs with openai.com in path should NOT match
|
||||||
|
{"https://my-proxy.com/api.openai.com/v1", false},
|
||||||
|
{"https://proxy.example.com/openai.azure.com/v1", false},
|
||||||
|
// Malformed or empty
|
||||||
|
{"", false},
|
||||||
|
{"not-a-url", false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := supportsPromptCacheKey(tt.apiBase); got != tt.want {
|
||||||
|
t.Errorf("supportsPromptCacheKey(%q) = %v, want %v", tt.apiBase, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestSerializeMessages_StripsSystemParts(t *testing.T) {
|
func TestSerializeMessages_StripsSystemParts(t *testing.T) {
|
||||||
messages := []protocoltypes.Message{
|
messages := []protocoltypes.Message{
|
||||||
{
|
{
|
||||||
|
|
|
||||||
|
|
@ -32,7 +32,7 @@ func NewSessionManager(storage string) *SessionManager {
|
||||||
}
|
}
|
||||||
|
|
||||||
if storage != "" {
|
if storage != "" {
|
||||||
os.MkdirAll(storage, 0o755)
|
os.MkdirAll(storage, 0o700)
|
||||||
sm.loadSessions()
|
sm.loadSessions()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -146,12 +146,15 @@ func (sm *SessionManager) TruncateHistory(key string, keepLast int) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// sanitizeFilename converts a session key into a cross-platform safe filename.
|
// sanitizeFilename converts a session key into a cross-platform safe filename.
|
||||||
// Session keys use "channel:chatID" (e.g. "telegram:123456") but ':' is the
|
// Replaces ':' with '_' (session key separator) and '/' and '\' with '_' so
|
||||||
// volume separator on Windows, so filepath.Base would misinterpret the key.
|
// composite IDs (e.g. Telegram forum "chatID/threadID") do not create
|
||||||
// We replace it with '_'. The original key is preserved inside the JSON file,
|
// subdirectories or break on Windows. The original key is preserved inside
|
||||||
// so loadSessions still maps back to the right in-memory key.
|
// the JSON file, so loadSessions still maps back to the right in-memory key.
|
||||||
func sanitizeFilename(key string) string {
|
func sanitizeFilename(key string) string {
|
||||||
return strings.ReplaceAll(key, ":", "_")
|
s := strings.ReplaceAll(key, ":", "_")
|
||||||
|
s = strings.ReplaceAll(s, "/", "_")
|
||||||
|
s = strings.ReplaceAll(s, "\\", "_")
|
||||||
|
return s
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sm *SessionManager) Save(key string) error {
|
func (sm *SessionManager) Save(key string) error {
|
||||||
|
|
@ -162,10 +165,9 @@ func (sm *SessionManager) Save(key string) error {
|
||||||
filename := sanitizeFilename(key)
|
filename := sanitizeFilename(key)
|
||||||
|
|
||||||
// filepath.IsLocal rejects empty names, "..", absolute paths, and
|
// filepath.IsLocal rejects empty names, "..", absolute paths, and
|
||||||
// OS-reserved device names (NUL, COM1 … on Windows).
|
// OS-reserved device names (NUL, COM1 … on Windows). sanitizeFilename
|
||||||
// The extra checks reject "." and any directory separators so that
|
// already replaced '/' and '\' with '_', so no subdirs are created.
|
||||||
// the session file is always written directly inside sm.storage.
|
if filename == "." || !filepath.IsLocal(filename) {
|
||||||
if filename == "." || !filepath.IsLocal(filename) || strings.ContainsAny(filename, `/\`) {
|
|
||||||
return os.ErrInvalid
|
return os.ErrInvalid
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -214,7 +216,7 @@ func (sm *SessionManager) Save(key string) error {
|
||||||
_ = tmpFile.Close()
|
_ = tmpFile.Close()
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := tmpFile.Chmod(0o644); err != nil {
|
if err := tmpFile.Chmod(0o600); err != nil {
|
||||||
_ = tmpFile.Close()
|
_ = tmpFile.Close()
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -17,6 +17,7 @@ func TestSanitizeFilename(t *testing.T) {
|
||||||
{"slack:C01234", "slack_C01234"},
|
{"slack:C01234", "slack_C01234"},
|
||||||
{"no-colons-here", "no-colons-here"},
|
{"no-colons-here", "no-colons-here"},
|
||||||
{"multiple:colons:here", "multiple_colons_here"},
|
{"multiple:colons:here", "multiple_colons_here"},
|
||||||
|
{"agent:main:telegram:group:-1003822706455/12", "agent_main_telegram_group_-1003822706455_12"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
|
|
@ -64,11 +65,21 @@ func TestSave_RejectsPathTraversal(t *testing.T) {
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
sm := NewSessionManager(tmpDir)
|
sm := NewSessionManager(tmpDir)
|
||||||
|
|
||||||
badKeys := []string{"", ".", "..", "foo/bar", "foo\\bar"}
|
// Invalid names that must still be rejected.
|
||||||
|
badKeys := []string{"", ".", ".."}
|
||||||
for _, key := range badKeys {
|
for _, key := range badKeys {
|
||||||
sm.GetOrCreate(key)
|
sm.GetOrCreate(key)
|
||||||
if err := sm.Save(key); err == nil {
|
if err := sm.Save(key); err == nil {
|
||||||
t.Errorf("Save(%q) should have failed but didn't", key)
|
t.Errorf("Save(%q) should have failed but didn't", key)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Keys containing path separators are sanitized (no subdirs created).
|
||||||
|
sm.GetOrCreate("foo/bar")
|
||||||
|
if err := sm.Save("foo/bar"); err != nil {
|
||||||
|
t.Fatalf("Save(\"foo/bar\") after sanitize should succeed: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(filepath.Join(tmpDir, "foo_bar.json")); os.IsNotExist(err) {
|
||||||
|
t.Errorf("expected foo_bar.json in storage (sanitized from foo/bar)")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -10,14 +10,15 @@ import (
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gomarkdown/markdown"
|
||||||
|
"github.com/gomarkdown/markdown/ast"
|
||||||
|
"github.com/gomarkdown/markdown/parser"
|
||||||
|
"gopkg.in/yaml.v3"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var namePattern = regexp.MustCompile(`^[a-zA-Z0-9]+(-[a-zA-Z0-9]+)*$`)
|
||||||
namePattern = regexp.MustCompile(`^[a-zA-Z0-9]+(-[a-zA-Z0-9]+)*$`)
|
|
||||||
reFrontmatter = regexp.MustCompile(`(?s)^---(?:\r\n|\n|\r)(.*?)(?:\r\n|\n|\r)---`)
|
|
||||||
reStripFrontmatter = regexp.MustCompile(`(?s)^---(?:\r\n|\n|\r)(.*?)(?:\r\n|\n|\r)---(?:\r\n|\n|\r)*`)
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
MaxNameLength = 64
|
MaxNameLength = 64
|
||||||
|
|
@ -226,11 +227,20 @@ func (sl *SkillsLoader) getSkillMetadata(skillPath string) *SkillMetadata {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
frontmatter := sl.extractFrontmatter(string(content))
|
frontmatter, bodyContent := splitFrontmatter(string(content))
|
||||||
if frontmatter == "" {
|
dirName := filepath.Base(filepath.Dir(skillPath))
|
||||||
return &SkillMetadata{
|
title, bodyDescription := extractMarkdownMetadata(bodyContent)
|
||||||
Name: filepath.Base(filepath.Dir(skillPath)),
|
|
||||||
|
metadata := &SkillMetadata{
|
||||||
|
Name: dirName,
|
||||||
|
Description: bodyDescription,
|
||||||
}
|
}
|
||||||
|
if title != "" && namePattern.MatchString(title) && len(title) <= MaxNameLength {
|
||||||
|
metadata.Name = title
|
||||||
|
}
|
||||||
|
|
||||||
|
if frontmatter == "" {
|
||||||
|
return metadata
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try JSON first (for backward compatibility)
|
// Try JSON first (for backward compatibility)
|
||||||
|
|
@ -239,60 +249,133 @@ func (sl *SkillsLoader) getSkillMetadata(skillPath string) *SkillMetadata {
|
||||||
Description string `json:"description"`
|
Description string `json:"description"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal([]byte(frontmatter), &jsonMeta); err == nil {
|
if err := json.Unmarshal([]byte(frontmatter), &jsonMeta); err == nil {
|
||||||
return &SkillMetadata{
|
if jsonMeta.Name != "" {
|
||||||
Name: jsonMeta.Name,
|
metadata.Name = jsonMeta.Name
|
||||||
Description: jsonMeta.Description,
|
|
||||||
}
|
}
|
||||||
|
if jsonMeta.Description != "" {
|
||||||
|
metadata.Description = jsonMeta.Description
|
||||||
|
}
|
||||||
|
return metadata
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fall back to simple YAML parsing
|
// Fall back to simple YAML parsing
|
||||||
yamlMeta := sl.parseSimpleYAML(frontmatter)
|
yamlMeta := sl.parseSimpleYAML(frontmatter)
|
||||||
return &SkillMetadata{
|
if name := yamlMeta["name"]; name != "" {
|
||||||
Name: yamlMeta["name"],
|
metadata.Name = name
|
||||||
Description: yamlMeta["description"],
|
|
||||||
}
|
}
|
||||||
|
if description := yamlMeta["description"]; description != "" {
|
||||||
|
metadata.Description = description
|
||||||
|
}
|
||||||
|
return metadata
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseSimpleYAML parses simple key: value YAML format
|
func extractMarkdownMetadata(content string) (title, description string) {
|
||||||
// Example: name: github\n description: "..."
|
p := parser.NewWithExtensions(parser.CommonExtensions)
|
||||||
// Normalizes line endings to handle \n (Unix), \r\n (Windows), and \r (classic Mac)
|
doc := markdown.Parse([]byte(content), p)
|
||||||
|
if doc == nil {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
ast.WalkFunc(doc, func(node ast.Node, entering bool) ast.WalkStatus {
|
||||||
|
if !entering {
|
||||||
|
return ast.GoToNext
|
||||||
|
}
|
||||||
|
|
||||||
|
switch n := node.(type) {
|
||||||
|
case *ast.Heading:
|
||||||
|
if title == "" && n.Level == 1 {
|
||||||
|
title = nodeText(n)
|
||||||
|
if title != "" && description != "" {
|
||||||
|
return ast.Terminate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case *ast.Paragraph:
|
||||||
|
if description == "" {
|
||||||
|
description = nodeText(n)
|
||||||
|
if title != "" && description != "" {
|
||||||
|
return ast.Terminate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ast.GoToNext
|
||||||
|
})
|
||||||
|
|
||||||
|
return title, description
|
||||||
|
}
|
||||||
|
|
||||||
|
func nodeText(n ast.Node) string {
|
||||||
|
var b strings.Builder
|
||||||
|
ast.WalkFunc(n, func(node ast.Node, entering bool) ast.WalkStatus {
|
||||||
|
if !entering {
|
||||||
|
return ast.GoToNext
|
||||||
|
}
|
||||||
|
|
||||||
|
switch t := node.(type) {
|
||||||
|
case *ast.Text:
|
||||||
|
b.Write(t.Literal)
|
||||||
|
case *ast.Code:
|
||||||
|
b.Write(t.Literal)
|
||||||
|
case *ast.Softbreak, *ast.Hardbreak, *ast.NonBlockingSpace:
|
||||||
|
b.WriteByte(' ')
|
||||||
|
}
|
||||||
|
return ast.GoToNext
|
||||||
|
})
|
||||||
|
return strings.Join(strings.Fields(b.String()), " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseSimpleYAML parses YAML frontmatter and extracts known metadata fields.
|
||||||
func (sl *SkillsLoader) parseSimpleYAML(content string) map[string]string {
|
func (sl *SkillsLoader) parseSimpleYAML(content string) map[string]string {
|
||||||
result := make(map[string]string)
|
result := make(map[string]string)
|
||||||
|
|
||||||
// Normalize line endings: convert \r\n and \r to \n
|
var meta struct {
|
||||||
normalized := strings.ReplaceAll(content, "\r\n", "\n")
|
Name string `yaml:"name"`
|
||||||
normalized = strings.ReplaceAll(normalized, "\r", "\n")
|
Description string `yaml:"description"`
|
||||||
|
|
||||||
for line := range strings.SplitSeq(normalized, "\n") {
|
|
||||||
line = strings.TrimSpace(line)
|
|
||||||
if line == "" || strings.HasPrefix(line, "#") {
|
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
|
if err := yaml.Unmarshal([]byte(content), &meta); err != nil {
|
||||||
parts := strings.SplitN(line, ":", 2)
|
return result
|
||||||
if len(parts) == 2 {
|
|
||||||
key := strings.TrimSpace(parts[0])
|
|
||||||
value := strings.TrimSpace(parts[1])
|
|
||||||
// Remove quotes if present
|
|
||||||
value = strings.Trim(value, "\"'")
|
|
||||||
result[key] = value
|
|
||||||
}
|
}
|
||||||
|
if meta.Name != "" {
|
||||||
|
result["name"] = meta.Name
|
||||||
|
}
|
||||||
|
if meta.Description != "" {
|
||||||
|
result["description"] = meta.Description
|
||||||
}
|
}
|
||||||
|
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sl *SkillsLoader) extractFrontmatter(content string) string {
|
func (sl *SkillsLoader) extractFrontmatter(content string) string {
|
||||||
// Support \n (Unix), \r\n (Windows), and \r (classic Mac) line endings for frontmatter blocks
|
frontmatter, _ := splitFrontmatter(content)
|
||||||
match := reFrontmatter.FindStringSubmatch(content)
|
return frontmatter
|
||||||
if len(match) > 1 {
|
|
||||||
return match[1]
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sl *SkillsLoader) stripFrontmatter(content string) string {
|
func (sl *SkillsLoader) stripFrontmatter(content string) string {
|
||||||
return reStripFrontmatter.ReplaceAllString(content, "")
|
_, body := splitFrontmatter(content)
|
||||||
|
return body
|
||||||
|
}
|
||||||
|
|
||||||
|
func splitFrontmatter(content string) (frontmatter, body string) {
|
||||||
|
normalized := string(parser.NormalizeNewlines([]byte(content)))
|
||||||
|
lines := strings.Split(normalized, "\n")
|
||||||
|
if len(lines) == 0 || lines[0] != "---" {
|
||||||
|
return "", content
|
||||||
|
}
|
||||||
|
|
||||||
|
end := -1
|
||||||
|
for i := 1; i < len(lines); i++ {
|
||||||
|
if lines[i] == "---" {
|
||||||
|
end = i
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if end == -1 {
|
||||||
|
return "", content
|
||||||
|
}
|
||||||
|
|
||||||
|
frontmatter = strings.Join(lines[1:end], "\n")
|
||||||
|
body = strings.Join(lines[end+1:], "\n")
|
||||||
|
body = strings.TrimLeft(body, "\n")
|
||||||
|
return frontmatter, body
|
||||||
}
|
}
|
||||||
|
|
||||||
func escapeXML(s string) string {
|
func escapeXML(s string) string {
|
||||||
|
|
|
||||||
|
|
@ -342,3 +342,78 @@ func TestSkillRootsTrimsWhitespaceAndDedups(t *testing.T) {
|
||||||
builtin,
|
builtin,
|
||||||
}, roots)
|
}, roots)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestGetSkillMetadata_UsesMarkdownParagraphWhenNoFrontmatter(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
skillDir := filepath.Join(tmp, "workspace", "skills", "plain-skill")
|
||||||
|
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
||||||
|
|
||||||
|
content := "# Plain Skill\n\nThis is parsed from markdown paragraph.\n"
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0o644))
|
||||||
|
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
meta := sl.getSkillMetadata(filepath.Join(skillDir, "SKILL.md"))
|
||||||
|
require.NotNil(t, meta)
|
||||||
|
assert.Equal(t, "plain-skill", meta.Name)
|
||||||
|
assert.Equal(t, "This is parsed from markdown paragraph.", meta.Description)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetSkillMetadata_FrontmatterOverridesMarkdown(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
skillDir := filepath.Join(tmp, "workspace", "skills", "plain-skill")
|
||||||
|
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
||||||
|
|
||||||
|
content := "---\nname: frontmatter-skill\ndescription: frontmatter description\n---\n\n# Plain Skill\n\nBody description.\n"
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0o644))
|
||||||
|
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
meta := sl.getSkillMetadata(filepath.Join(skillDir, "SKILL.md"))
|
||||||
|
require.NotNil(t, meta)
|
||||||
|
assert.Equal(t, "frontmatter-skill", meta.Name)
|
||||||
|
assert.Equal(t, "frontmatter description", meta.Description)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetSkillMetadata_YAMLMultilineDescription(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
skillDir := filepath.Join(tmp, "workspace", "skills", "plain-skill")
|
||||||
|
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
||||||
|
|
||||||
|
content := "---\nname: frontmatter-skill\ndescription: |\n line 1: with colon\n line 2\n---\n\n# Plain Skill\n\nBody description.\n"
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0o644))
|
||||||
|
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
meta := sl.getSkillMetadata(filepath.Join(skillDir, "SKILL.md"))
|
||||||
|
require.NotNil(t, meta)
|
||||||
|
assert.Equal(t, "frontmatter-skill", meta.Name)
|
||||||
|
assert.Equal(t, "line 1: with colon\nline 2", meta.Description)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetSkillMetadata_InvalidHeadingNameFallsBackToDirName(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
skillDir := filepath.Join(tmp, "workspace", "skills", "valid-name")
|
||||||
|
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
||||||
|
|
||||||
|
content := "# Invalid Heading Name\n\nBody description.\n"
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0o644))
|
||||||
|
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
meta := sl.getSkillMetadata(filepath.Join(skillDir, "SKILL.md"))
|
||||||
|
require.NotNil(t, meta)
|
||||||
|
assert.Equal(t, "valid-name", meta.Name)
|
||||||
|
assert.Equal(t, "Body description.", meta.Description)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetSkillMetadata_IgnoresHTMLCommentBlocks(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
skillDir := filepath.Join(tmp, "workspace", "skills", "biomed-skill")
|
||||||
|
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
||||||
|
|
||||||
|
content := "<!--\n# COPYRIGHT NOTICE\n# This file is part of the \"Universal Biomedical Skills\" project.\n# Copyright (c) 2026 MD BABU MIA, PhD <md.babu.mia@mssm.edu>\n# All Rights Reserved.\n#\n# This code is proprietary and confidential.\n# Unauthorized copying of this file, via any medium is strictly prohibited.\n#\n# Provenance: Authenticated by MD BABU MIA\n\n-->\n\n# Biomed Skill\n\nSummarize biomedical papers.\n"
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0o644))
|
||||||
|
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
meta := sl.getSkillMetadata(filepath.Join(skillDir, "SKILL.md"))
|
||||||
|
require.NotNil(t, meta)
|
||||||
|
assert.Equal(t, "biomed-skill", meta.Name)
|
||||||
|
assert.Equal(t, "Summarize biomedical papers.", meta.Description)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -40,8 +40,8 @@ func NewManager(workspace string) *Manager {
|
||||||
oldStateFile := filepath.Join(workspace, "state.json")
|
oldStateFile := filepath.Join(workspace, "state.json")
|
||||||
|
|
||||||
// Create state directory if it doesn't exist
|
// Create state directory if it doesn't exist
|
||||||
if err := os.MkdirAll(stateDir, 0o755); err != nil {
|
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||||
log.Fatalf("[FATAL] state: failed to create state directory: %v", err)
|
log.Printf("[WARN] state: failed to create state directory %s: %v", stateDir, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
sm := &Manager{
|
sm := &Manager{
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,6 @@ package state
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
|
|
@ -217,10 +216,7 @@ func TestNewManager_EmptyWorkspace(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNewManager_MkdirFailureCrashes(t *testing.T) {
|
func TestNewManager_MkdirFailureDoesNotCrash(t *testing.T) {
|
||||||
// Since log.Fatalf calls os.Exit(1), we cannot test it normally
|
|
||||||
// Otherwise, the test suite would stop altogether.
|
|
||||||
// We use the standard pattern of Go: rerun this test in a subprocess.
|
|
||||||
if os.Getenv("BE_CRASHER") == "1" {
|
if os.Getenv("BE_CRASHER") == "1" {
|
||||||
tmpDir := os.Getenv("CRASH_DIR")
|
tmpDir := os.Getenv("CRASH_DIR")
|
||||||
|
|
||||||
|
|
@ -240,15 +236,11 @@ func TestNewManager_MkdirFailureCrashes(t *testing.T) {
|
||||||
}
|
}
|
||||||
defer os.RemoveAll(tmpDir)
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
cmd := exec.Command(os.Args[0], "-test.run=TestNewManager_MkdirFailureCrashes")
|
cmd := exec.Command(os.Args[0], "-test.run=TestNewManager_MkdirFailureDoesNotCrash")
|
||||||
cmd.Env = append(os.Environ(), "BE_CRASHER=1", "CRASH_DIR="+tmpDir)
|
cmd.Env = append(os.Environ(), "BE_CRASHER=1", "CRASH_DIR="+tmpDir)
|
||||||
|
|
||||||
err = cmd.Run()
|
err = cmd.Run()
|
||||||
|
if err != nil {
|
||||||
var e *exec.ExitError
|
t.Fatalf("NewManager should not crash when state dir creation fails, got: %v", err)
|
||||||
if errors.As(err, &e) && !e.Success() {
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Fatalf("The process ended without error, a crash was expected via os.Exit(1). Err: %v", err)
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,7 @@ import (
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/constants"
|
||||||
"github.com/sipeed/picoclaw/pkg/cron"
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
@ -73,6 +74,10 @@ func (t *CronTool) Parameters() map[string]any {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Optional: Shell command to execute directly (e.g., 'df -h'). If set, the agent will run this command and report output instead of just showing the message. 'deliver' will be forced to false for commands.",
|
"description": "Optional: Shell command to execute directly (e.g., 'df -h'). If set, the agent will run this command and report output instead of just showing the message. 'deliver' will be forced to false for commands.",
|
||||||
},
|
},
|
||||||
|
"command_confirm": map[string]any{
|
||||||
|
"type": "boolean",
|
||||||
|
"description": "Required when using command=true. Must be true to explicitly confirm scheduling a shell command.",
|
||||||
|
},
|
||||||
"at_seconds": map[string]any{
|
"at_seconds": map[string]any{
|
||||||
"type": "integer",
|
"type": "integer",
|
||||||
"description": "One-time reminder: seconds from now when to trigger (e.g., 600 for 10 minutes later). Use this for one-time reminders like 'remind me in 10 minutes'.",
|
"description": "One-time reminder: seconds from now when to trigger (e.g., 600 for 10 minutes later). Use this for one-time reminders like 'remind me in 10 minutes'.",
|
||||||
|
|
@ -175,12 +180,17 @@ func (t *CronTool) addJob(ctx context.Context, args map[string]any) *ToolResult
|
||||||
deliver = d
|
deliver = d
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GHSA-pv8c-p6jf-3fpp: command scheduling requires internal channel + explicit confirm.
|
||||||
|
// Non-command reminders (plain messages) remain open to all channels.
|
||||||
command, _ := args["command"].(string)
|
command, _ := args["command"].(string)
|
||||||
|
commandConfirm, _ := args["command_confirm"].(bool)
|
||||||
if command != "" {
|
if command != "" {
|
||||||
// Commands must be processed by agent/exec tool, so deliver must be false (or handled specifically)
|
if !constants.IsInternalChannel(channel) {
|
||||||
// Actually, let's keep deliver=false to let the system know it's not a simple chat message
|
return ErrorResult("scheduling command execution is restricted to internal channels")
|
||||||
// But for our new logic in ExecuteJob, we can handle it regardless of deliver flag if Payload.Command is set.
|
}
|
||||||
// However, logically, it's not "delivered" to chat directly as is.
|
if !commandConfirm {
|
||||||
|
return ErrorResult("command_confirm=true is required to schedule command execution")
|
||||||
|
}
|
||||||
deliver = false
|
deliver = false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -282,6 +292,8 @@ func (t *CronTool) ExecuteJob(ctx context.Context, job *cron.CronJob) string {
|
||||||
if job.Payload.Command != "" {
|
if job.Payload.Command != "" {
|
||||||
args := map[string]any{
|
args := map[string]any{
|
||||||
"command": job.Payload.Command,
|
"command": job.Payload.Command,
|
||||||
|
"__channel": channel,
|
||||||
|
"__chat_id": chatID,
|
||||||
}
|
}
|
||||||
|
|
||||||
result := t.execTool.Execute(ctx, args)
|
result := t.execTool.Execute(ctx, args)
|
||||||
|
|
|
||||||
116
pkg/tools/cron_test.go
Normal file
116
pkg/tools/cron_test.go
Normal file
|
|
@ -0,0 +1,116 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestCronTool(t *testing.T) *CronTool {
|
||||||
|
t.Helper()
|
||||||
|
storePath := filepath.Join(t.TempDir(), "cron.json")
|
||||||
|
cronService := cron.NewCronService(storePath, nil)
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
tool, err := NewCronTool(cronService, nil, msgBus, t.TempDir(), true, 0, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewCronTool() error: %v", err)
|
||||||
|
}
|
||||||
|
return tool
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCronTool_CommandBlockedFromRemoteChannel verifies command scheduling is restricted to internal channels
|
||||||
|
func TestCronTool_CommandBlockedFromRemoteChannel(t *testing.T) {
|
||||||
|
tool := newTestCronTool(t)
|
||||||
|
ctx := WithToolContext(context.Background(), "telegram", "chat-1")
|
||||||
|
result := tool.Execute(ctx, map[string]any{
|
||||||
|
"action": "add",
|
||||||
|
"message": "check disk",
|
||||||
|
"command": "df -h",
|
||||||
|
"command_confirm": true,
|
||||||
|
"at_seconds": float64(60),
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Fatal("expected command scheduling to be blocked from remote channel")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "restricted to internal channels") {
|
||||||
|
t.Errorf("expected 'restricted to internal channels', got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCronTool_CommandRequiresConfirm verifies command_confirm=true is required
|
||||||
|
func TestCronTool_CommandRequiresConfirm(t *testing.T) {
|
||||||
|
tool := newTestCronTool(t)
|
||||||
|
ctx := WithToolContext(context.Background(), "cli", "direct")
|
||||||
|
result := tool.Execute(ctx, map[string]any{
|
||||||
|
"action": "add",
|
||||||
|
"message": "check disk",
|
||||||
|
"command": "df -h",
|
||||||
|
"at_seconds": float64(60),
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Fatal("expected error when command_confirm is missing")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "command_confirm=true") {
|
||||||
|
t.Errorf("expected 'command_confirm=true' message, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCronTool_CommandAllowedFromInternalChannel verifies command scheduling works from internal channels
|
||||||
|
func TestCronTool_CommandAllowedFromInternalChannel(t *testing.T) {
|
||||||
|
tool := newTestCronTool(t)
|
||||||
|
ctx := WithToolContext(context.Background(), "cli", "direct")
|
||||||
|
result := tool.Execute(ctx, map[string]any{
|
||||||
|
"action": "add",
|
||||||
|
"message": "check disk",
|
||||||
|
"command": "df -h",
|
||||||
|
"command_confirm": true,
|
||||||
|
"at_seconds": float64(60),
|
||||||
|
})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("expected command scheduling to succeed from internal channel, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "Cron job added") {
|
||||||
|
t.Errorf("expected 'Cron job added', got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCronTool_AddJobRequiresSessionContext verifies fail-closed when channel/chatID missing
|
||||||
|
func TestCronTool_AddJobRequiresSessionContext(t *testing.T) {
|
||||||
|
tool := newTestCronTool(t)
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"action": "add",
|
||||||
|
"message": "reminder",
|
||||||
|
"at_seconds": float64(60),
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Fatal("expected error when session context is missing")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "no session context") {
|
||||||
|
t.Errorf("expected 'no session context' message, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCronTool_NonCommandJobAllowedFromRemoteChannel verifies regular reminders work from any channel
|
||||||
|
func TestCronTool_NonCommandJobAllowedFromRemoteChannel(t *testing.T) {
|
||||||
|
tool := newTestCronTool(t)
|
||||||
|
ctx := WithToolContext(context.Background(), "telegram", "chat-1")
|
||||||
|
result := tool.Execute(ctx, map[string]any{
|
||||||
|
"action": "add",
|
||||||
|
"message": "time to stretch",
|
||||||
|
"at_seconds": float64(600),
|
||||||
|
})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("expected non-command reminder to succeed from remote channel, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -14,6 +14,7 @@ import (
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/constants"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ExecTool struct {
|
type ExecTool struct {
|
||||||
|
|
@ -23,6 +24,7 @@ type ExecTool struct {
|
||||||
allowPatterns []*regexp.Regexp
|
allowPatterns []*regexp.Regexp
|
||||||
customAllowPatterns []*regexp.Regexp
|
customAllowPatterns []*regexp.Regexp
|
||||||
restrictToWorkspace bool
|
restrictToWorkspace bool
|
||||||
|
allowRemote bool
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|
@ -100,10 +102,12 @@ func NewExecTool(workingDir string, restrict bool) (*ExecTool, error) {
|
||||||
func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Config) (*ExecTool, error) {
|
func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Config) (*ExecTool, error) {
|
||||||
denyPatterns := make([]*regexp.Regexp, 0)
|
denyPatterns := make([]*regexp.Regexp, 0)
|
||||||
customAllowPatterns := make([]*regexp.Regexp, 0)
|
customAllowPatterns := make([]*regexp.Regexp, 0)
|
||||||
|
allowRemote := true
|
||||||
|
|
||||||
if config != nil {
|
if config != nil {
|
||||||
execConfig := config.Tools.Exec
|
execConfig := config.Tools.Exec
|
||||||
enableDenyPatterns := execConfig.EnableDenyPatterns
|
enableDenyPatterns := execConfig.EnableDenyPatterns
|
||||||
|
allowRemote = execConfig.AllowRemote
|
||||||
if enableDenyPatterns {
|
if enableDenyPatterns {
|
||||||
denyPatterns = append(denyPatterns, defaultDenyPatterns...)
|
denyPatterns = append(denyPatterns, defaultDenyPatterns...)
|
||||||
if len(execConfig.CustomDenyPatterns) > 0 {
|
if len(execConfig.CustomDenyPatterns) > 0 {
|
||||||
|
|
@ -143,6 +147,7 @@ func NewExecToolWithConfig(workingDir string, restrict bool, config *config.Conf
|
||||||
allowPatterns: nil,
|
allowPatterns: nil,
|
||||||
customAllowPatterns: customAllowPatterns,
|
customAllowPatterns: customAllowPatterns,
|
||||||
restrictToWorkspace: restrict,
|
restrictToWorkspace: restrict,
|
||||||
|
allowRemote: allowRemote,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -177,6 +182,19 @@ func (t *ExecTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
||||||
return ErrorResult("command is required")
|
return ErrorResult("command is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GHSA-pv8c-p6jf-3fpp: block exec from remote channels (e.g. Telegram webhooks)
|
||||||
|
// unless explicitly opted-in via config. Fail-closed: empty channel = blocked.
|
||||||
|
if !t.allowRemote {
|
||||||
|
channel := ToolChannel(ctx)
|
||||||
|
if channel == "" {
|
||||||
|
channel, _ = args["__channel"].(string)
|
||||||
|
}
|
||||||
|
channel = strings.TrimSpace(channel)
|
||||||
|
if channel == "" || !constants.IsInternalChannel(channel) {
|
||||||
|
return ErrorResult("exec is restricted to internal channels")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
cwd := t.workingDir
|
cwd := t.workingDir
|
||||||
if wd, ok := args["working_dir"].(string); ok && wd != "" {
|
if wd, ok := args["working_dir"].(string); ok && wd != "" {
|
||||||
if t.restrictToWorkspace && t.workingDir != "" {
|
if t.restrictToWorkspace && t.workingDir != "" {
|
||||||
|
|
@ -201,6 +219,25 @@ func (t *ExecTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
||||||
return ErrorResult(guardError)
|
return ErrorResult(guardError)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Re-resolve symlinks immediately before execution to shrink the TOCTOU window
|
||||||
|
// between validation and cmd.Dir assignment.
|
||||||
|
if t.restrictToWorkspace && t.workingDir != "" && cwd != t.workingDir {
|
||||||
|
resolved, err := filepath.EvalSymlinks(cwd)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("Command blocked by safety guard (path resolution failed: %v)", err))
|
||||||
|
}
|
||||||
|
absWorkspace, _ := filepath.Abs(t.workingDir)
|
||||||
|
wsResolved, _ := filepath.EvalSymlinks(absWorkspace)
|
||||||
|
if wsResolved == "" {
|
||||||
|
wsResolved = absWorkspace
|
||||||
|
}
|
||||||
|
rel, err := filepath.Rel(wsResolved, resolved)
|
||||||
|
if err != nil || !filepath.IsLocal(rel) {
|
||||||
|
return ErrorResult("Command blocked by safety guard (working directory escaped workspace)")
|
||||||
|
}
|
||||||
|
cwd = resolved
|
||||||
|
}
|
||||||
|
|
||||||
// timeout == 0 means no timeout
|
// timeout == 0 means no timeout
|
||||||
var cmdCtx context.Context
|
var cmdCtx context.Context
|
||||||
var cancel context.CancelFunc
|
var cancel context.CancelFunc
|
||||||
|
|
|
||||||
|
|
@ -301,6 +301,85 @@ func TestShellTool_WorkingDir_SymlinkEscape(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestShellTool_RemoteChannelBlockedByDefault verifies exec is blocked for remote channels
|
||||||
|
func TestShellTool_RemoteChannelBlockedByDefault(t *testing.T) {
|
||||||
|
cfg := &config.Config{}
|
||||||
|
cfg.Tools.Exec.EnableDenyPatterns = true
|
||||||
|
cfg.Tools.Exec.AllowRemote = false
|
||||||
|
|
||||||
|
tool, err := NewExecToolWithConfig("", false, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewExecToolWithConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
ctx := WithToolContext(context.Background(), "telegram", "chat-1")
|
||||||
|
result := tool.Execute(ctx, map[string]any{"command": "echo hi"})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Fatal("expected remote-channel exec to be blocked")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "restricted to internal channels") {
|
||||||
|
t.Errorf("expected 'restricted to internal channels' message, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestShellTool_InternalChannelAllowed verifies exec is allowed for internal channels
|
||||||
|
func TestShellTool_InternalChannelAllowed(t *testing.T) {
|
||||||
|
cfg := &config.Config{}
|
||||||
|
cfg.Tools.Exec.EnableDenyPatterns = true
|
||||||
|
cfg.Tools.Exec.AllowRemote = false
|
||||||
|
|
||||||
|
tool, err := NewExecToolWithConfig("", false, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewExecToolWithConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
ctx := WithToolContext(context.Background(), "cli", "direct")
|
||||||
|
result := tool.Execute(ctx, map[string]any{"command": "echo hi"})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("expected internal channel exec to succeed, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "hi") {
|
||||||
|
t.Errorf("expected output to contain 'hi', got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestShellTool_EmptyChannelBlockedWhenNotAllowRemote verifies fail-closed when no channel context
|
||||||
|
func TestShellTool_EmptyChannelBlockedWhenNotAllowRemote(t *testing.T) {
|
||||||
|
cfg := &config.Config{}
|
||||||
|
cfg.Tools.Exec.EnableDenyPatterns = true
|
||||||
|
cfg.Tools.Exec.AllowRemote = false
|
||||||
|
|
||||||
|
tool, err := NewExecToolWithConfig("", false, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewExecToolWithConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"command": "echo hi",
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Fatal("expected exec with empty channel to be blocked when allowRemote=false")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestShellTool_AllowRemoteBypassesChannelCheck verifies allowRemote=true permits any channel
|
||||||
|
func TestShellTool_AllowRemoteBypassesChannelCheck(t *testing.T) {
|
||||||
|
cfg := &config.Config{}
|
||||||
|
cfg.Tools.Exec.EnableDenyPatterns = true
|
||||||
|
cfg.Tools.Exec.AllowRemote = true
|
||||||
|
|
||||||
|
tool, err := NewExecToolWithConfig("", false, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewExecToolWithConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
ctx := WithToolContext(context.Background(), "telegram", "chat-1")
|
||||||
|
result := tool.Execute(ctx, map[string]any{"command": "echo hi"})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("expected allowRemote=true to permit remote channel, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestShellTool_RestrictToWorkspace verifies workspace restriction
|
// TestShellTool_RestrictToWorkspace verifies workspace restriction
|
||||||
func TestShellTool_RestrictToWorkspace(t *testing.T) {
|
func TestShellTool_RestrictToWorkspace(t *testing.T) {
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
|
||||||
312
pkg/tools/web.go
312
pkg/tools/web.go
|
|
@ -7,10 +7,12 @@ import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -76,12 +78,50 @@ func createHTTPClient(proxyURL string, timeout time.Duration) (*http.Client, err
|
||||||
return client, nil
|
return client, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type APIKeyPool struct {
|
||||||
|
keys []string
|
||||||
|
current uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewAPIKeyPool(keys []string) *APIKeyPool {
|
||||||
|
return &APIKeyPool{
|
||||||
|
keys: keys,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type APIKeyIterator struct {
|
||||||
|
pool *APIKeyPool
|
||||||
|
startIdx uint32
|
||||||
|
attempt uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *APIKeyPool) NewIterator() *APIKeyIterator {
|
||||||
|
if len(p.keys) == 0 {
|
||||||
|
return &APIKeyIterator{pool: p}
|
||||||
|
}
|
||||||
|
idx := atomic.AddUint32(&p.current, 1) - 1
|
||||||
|
return &APIKeyIterator{
|
||||||
|
pool: p,
|
||||||
|
startIdx: idx,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (it *APIKeyIterator) Next() (string, bool) {
|
||||||
|
length := uint32(len(it.pool.keys))
|
||||||
|
if length == 0 || it.attempt >= length {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
key := it.pool.keys[(it.startIdx+it.attempt)%length]
|
||||||
|
it.attempt++
|
||||||
|
return key, true
|
||||||
|
}
|
||||||
|
|
||||||
type SearchProvider interface {
|
type SearchProvider interface {
|
||||||
Search(ctx context.Context, query string, count int) (string, error)
|
Search(ctx context.Context, query string, count int) (string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type BraveSearchProvider struct {
|
type BraveSearchProvider struct {
|
||||||
apiKey string
|
keyPool *APIKeyPool
|
||||||
proxy string
|
proxy string
|
||||||
client *http.Client
|
client *http.Client
|
||||||
}
|
}
|
||||||
|
|
@ -90,27 +130,46 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in
|
||||||
searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d",
|
searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d",
|
||||||
url.QueryEscape(query), count)
|
url.QueryEscape(query), count)
|
||||||
|
|
||||||
|
var lastErr error
|
||||||
|
iter := p.keyPool.NewIterator()
|
||||||
|
|
||||||
|
for {
|
||||||
|
apiKey, ok := iter.Next()
|
||||||
|
if !ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil)
|
req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to create request: %w", err)
|
return "", fmt.Errorf("failed to create request: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
req.Header.Set("Accept", "application/json")
|
req.Header.Set("Accept", "application/json")
|
||||||
req.Header.Set("X-Subscription-Token", p.apiKey)
|
req.Header.Set("X-Subscription-Token", apiKey)
|
||||||
|
|
||||||
resp, err := p.client.Do(req)
|
resp, err := p.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
lastErr = fmt.Errorf("request failed: %w", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
lastErr = fmt.Errorf("failed to read response: %w", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
if resp.StatusCode != http.StatusOK {
|
||||||
return "", fmt.Errorf("brave api error (status %d): %s", resp.StatusCode, string(body))
|
lastErr = fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||||
|
if resp.StatusCode == http.StatusTooManyRequests ||
|
||||||
|
resp.StatusCode == http.StatusUnauthorized ||
|
||||||
|
resp.StatusCode == http.StatusForbidden ||
|
||||||
|
resp.StatusCode >= 500 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return "", lastErr
|
||||||
}
|
}
|
||||||
|
|
||||||
var searchResp struct {
|
var searchResp struct {
|
||||||
|
|
@ -125,7 +184,6 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in
|
||||||
|
|
||||||
if err := json.Unmarshal(body, &searchResp); err != nil {
|
if err := json.Unmarshal(body, &searchResp); err != nil {
|
||||||
// Log error body for debugging
|
// Log error body for debugging
|
||||||
fmt.Printf("Brave API Error Body: %s\n", string(body))
|
|
||||||
return "", fmt.Errorf("failed to parse response: %w", err)
|
return "", fmt.Errorf("failed to parse response: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -147,10 +205,13 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in
|
||||||
}
|
}
|
||||||
|
|
||||||
return strings.Join(lines, "\n"), nil
|
return strings.Join(lines, "\n"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", fmt.Errorf("all api keys failed, last error: %w", lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
type TavilySearchProvider struct {
|
type TavilySearchProvider struct {
|
||||||
apiKey string
|
keyPool *APIKeyPool
|
||||||
baseURL string
|
baseURL string
|
||||||
proxy string
|
proxy string
|
||||||
client *http.Client
|
client *http.Client
|
||||||
|
|
@ -162,8 +223,17 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
|
||||||
searchURL = "https://api.tavily.com/search"
|
searchURL = "https://api.tavily.com/search"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var lastErr error
|
||||||
|
iter := p.keyPool.NewIterator()
|
||||||
|
|
||||||
|
for {
|
||||||
|
apiKey, ok := iter.Next()
|
||||||
|
if !ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
payload := map[string]any{
|
payload := map[string]any{
|
||||||
"api_key": p.apiKey,
|
"api_key": apiKey,
|
||||||
"query": query,
|
"query": query,
|
||||||
"search_depth": "advanced",
|
"search_depth": "advanced",
|
||||||
"include_answer": false,
|
"include_answer": false,
|
||||||
|
|
@ -187,17 +257,27 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
|
||||||
|
|
||||||
resp, err := p.client.Do(req)
|
resp, err := p.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
lastErr = fmt.Errorf("request failed: %w", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
lastErr = fmt.Errorf("failed to read response: %w", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
if resp.StatusCode != http.StatusOK {
|
||||||
return "", fmt.Errorf("tavily api error (status %d): %s", resp.StatusCode, string(body))
|
lastErr = fmt.Errorf("tavily api error (status %d): %s", resp.StatusCode, string(body))
|
||||||
|
if resp.StatusCode == http.StatusTooManyRequests ||
|
||||||
|
resp.StatusCode == http.StatusUnauthorized ||
|
||||||
|
resp.StatusCode == http.StatusForbidden ||
|
||||||
|
resp.StatusCode >= 500 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return "", lastErr
|
||||||
}
|
}
|
||||||
|
|
||||||
var searchResp struct {
|
var searchResp struct {
|
||||||
|
|
@ -230,6 +310,9 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
|
||||||
}
|
}
|
||||||
|
|
||||||
return strings.Join(lines, "\n"), nil
|
return strings.Join(lines, "\n"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", fmt.Errorf("all api keys failed, last error: %w", lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
type DuckDuckGoSearchProvider struct {
|
type DuckDuckGoSearchProvider struct {
|
||||||
|
|
@ -324,7 +407,7 @@ func stripTags(content string) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
type PerplexitySearchProvider struct {
|
type PerplexitySearchProvider struct {
|
||||||
apiKey string
|
keyPool *APIKeyPool
|
||||||
proxy string
|
proxy string
|
||||||
client *http.Client
|
client *http.Client
|
||||||
}
|
}
|
||||||
|
|
@ -332,6 +415,15 @@ type PerplexitySearchProvider struct {
|
||||||
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
||||||
searchURL := "https://api.perplexity.ai/chat/completions"
|
searchURL := "https://api.perplexity.ai/chat/completions"
|
||||||
|
|
||||||
|
var lastErr error
|
||||||
|
iter := p.keyPool.NewIterator()
|
||||||
|
|
||||||
|
for {
|
||||||
|
apiKey, ok := iter.Next()
|
||||||
|
if !ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
payload := map[string]any{
|
payload := map[string]any{
|
||||||
"model": "sonar",
|
"model": "sonar",
|
||||||
"messages": []map[string]string{
|
"messages": []map[string]string{
|
||||||
|
|
@ -358,22 +450,32 @@ func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, cou
|
||||||
}
|
}
|
||||||
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||||
req.Header.Set("User-Agent", userAgent)
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
resp, err := p.client.Do(req)
|
resp, err := p.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
lastErr = fmt.Errorf("request failed: %w", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
lastErr = fmt.Errorf("failed to read response: %w", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
if resp.StatusCode != http.StatusOK {
|
||||||
return "", fmt.Errorf("Perplexity API error: %s", string(body))
|
lastErr = fmt.Errorf("Perplexity API error: %s", string(body))
|
||||||
|
if resp.StatusCode == http.StatusTooManyRequests ||
|
||||||
|
resp.StatusCode == http.StatusUnauthorized ||
|
||||||
|
resp.StatusCode == http.StatusForbidden ||
|
||||||
|
resp.StatusCode >= 500 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return "", lastErr
|
||||||
}
|
}
|
||||||
|
|
||||||
var searchResp struct {
|
var searchResp struct {
|
||||||
|
|
@ -393,6 +495,9 @@ func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, cou
|
||||||
}
|
}
|
||||||
|
|
||||||
return fmt.Sprintf("Results for: %s (via Perplexity)\n%s", query, searchResp.Choices[0].Message.Content), nil
|
return fmt.Sprintf("Results for: %s (via Perplexity)\n%s", query, searchResp.Choices[0].Message.Content), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", fmt.Errorf("all api keys failed, last error: %w", lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
type SearXNGSearchProvider struct {
|
type SearXNGSearchProvider struct {
|
||||||
|
|
@ -545,16 +650,16 @@ type WebSearchTool struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type WebSearchToolOptions struct {
|
type WebSearchToolOptions struct {
|
||||||
BraveAPIKey string
|
BraveAPIKeys []string
|
||||||
BraveMaxResults int
|
BraveMaxResults int
|
||||||
BraveEnabled bool
|
BraveEnabled bool
|
||||||
TavilyAPIKey string
|
TavilyAPIKeys []string
|
||||||
TavilyBaseURL string
|
TavilyBaseURL string
|
||||||
TavilyMaxResults int
|
TavilyMaxResults int
|
||||||
TavilyEnabled bool
|
TavilyEnabled bool
|
||||||
DuckDuckGoMaxResults int
|
DuckDuckGoMaxResults int
|
||||||
DuckDuckGoEnabled bool
|
DuckDuckGoEnabled bool
|
||||||
PerplexityAPIKey string
|
PerplexityAPIKeys []string
|
||||||
PerplexityMaxResults int
|
PerplexityMaxResults int
|
||||||
PerplexityEnabled bool
|
PerplexityEnabled bool
|
||||||
SearXNGBaseURL string
|
SearXNGBaseURL string
|
||||||
|
|
@ -571,23 +676,26 @@ type WebSearchToolOptions struct {
|
||||||
func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
|
func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
|
||||||
var provider SearchProvider
|
var provider SearchProvider
|
||||||
maxResults := 5
|
maxResults := 5
|
||||||
|
|
||||||
// Priority: Perplexity > Brave > SearXNG > Tavily > DuckDuckGo > GLM Search
|
// Priority: Perplexity > Brave > SearXNG > Tavily > DuckDuckGo > GLM Search
|
||||||
if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" {
|
if opts.PerplexityEnabled && len(opts.PerplexityAPIKeys) > 0 {
|
||||||
client, err := createHTTPClient(opts.Proxy, perplexityTimeout)
|
client, err := createHTTPClient(opts.Proxy, perplexityTimeout)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for Perplexity: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for Perplexity: %w", err)
|
||||||
}
|
}
|
||||||
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey, proxy: opts.Proxy, client: client}
|
provider = &PerplexitySearchProvider{
|
||||||
|
keyPool: NewAPIKeyPool(opts.PerplexityAPIKeys),
|
||||||
|
proxy: opts.Proxy,
|
||||||
|
client: client,
|
||||||
|
}
|
||||||
if opts.PerplexityMaxResults > 0 {
|
if opts.PerplexityMaxResults > 0 {
|
||||||
maxResults = opts.PerplexityMaxResults
|
maxResults = opts.PerplexityMaxResults
|
||||||
}
|
}
|
||||||
} else if opts.BraveEnabled && opts.BraveAPIKey != "" {
|
} else if opts.BraveEnabled && len(opts.BraveAPIKeys) > 0 {
|
||||||
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for Brave: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for Brave: %w", err)
|
||||||
}
|
}
|
||||||
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey, proxy: opts.Proxy, client: client}
|
provider = &BraveSearchProvider{keyPool: NewAPIKeyPool(opts.BraveAPIKeys), proxy: opts.Proxy, client: client}
|
||||||
if opts.BraveMaxResults > 0 {
|
if opts.BraveMaxResults > 0 {
|
||||||
maxResults = opts.BraveMaxResults
|
maxResults = opts.BraveMaxResults
|
||||||
}
|
}
|
||||||
|
|
@ -596,13 +704,13 @@ func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
|
||||||
if opts.SearXNGMaxResults > 0 {
|
if opts.SearXNGMaxResults > 0 {
|
||||||
maxResults = opts.SearXNGMaxResults
|
maxResults = opts.SearXNGMaxResults
|
||||||
}
|
}
|
||||||
} else if opts.TavilyEnabled && opts.TavilyAPIKey != "" {
|
} else if opts.TavilyEnabled && len(opts.TavilyAPIKeys) > 0 {
|
||||||
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for Tavily: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for Tavily: %w", err)
|
||||||
}
|
}
|
||||||
provider = &TavilySearchProvider{
|
provider = &TavilySearchProvider{
|
||||||
apiKey: opts.TavilyAPIKey,
|
keyPool: NewAPIKeyPool(opts.TavilyAPIKeys),
|
||||||
baseURL: opts.TavilyBaseURL,
|
baseURL: opts.TavilyBaseURL,
|
||||||
proxy: opts.Proxy,
|
proxy: opts.Proxy,
|
||||||
client: client,
|
client: client,
|
||||||
|
|
@ -711,6 +819,10 @@ func NewWebFetchTool(maxChars int, fetchLimitBytes int64) (*WebFetchTool, error)
|
||||||
return NewWebFetchToolWithProxy(maxChars, "", fetchLimitBytes)
|
return NewWebFetchToolWithProxy(maxChars, "", fetchLimitBytes)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// allowPrivateWebFetchHosts controls whether loopback/private hosts are allowed.
|
||||||
|
// This is false in normal runtime to reduce SSRF exposure, and tests can override it temporarily.
|
||||||
|
var allowPrivateWebFetchHosts atomic.Bool
|
||||||
|
|
||||||
func NewWebFetchToolWithProxy(maxChars int, proxy string, fetchLimitBytes int64) (*WebFetchTool, error) {
|
func NewWebFetchToolWithProxy(maxChars int, proxy string, fetchLimitBytes int64) (*WebFetchTool, error) {
|
||||||
if maxChars <= 0 {
|
if maxChars <= 0 {
|
||||||
maxChars = defaultMaxChars
|
maxChars = defaultMaxChars
|
||||||
|
|
@ -719,10 +831,20 @@ func NewWebFetchToolWithProxy(maxChars int, proxy string, fetchLimitBytes int64)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for web fetch: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for web fetch: %w", err)
|
||||||
}
|
}
|
||||||
|
if transport, ok := client.Transport.(*http.Transport); ok {
|
||||||
|
dialer := &net.Dialer{
|
||||||
|
Timeout: 15 * time.Second,
|
||||||
|
KeepAlive: 30 * time.Second,
|
||||||
|
}
|
||||||
|
transport.DialContext = newSafeDialContext(dialer)
|
||||||
|
}
|
||||||
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
|
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
|
||||||
if len(via) >= maxRedirects {
|
if len(via) >= maxRedirects {
|
||||||
return fmt.Errorf("stopped after %d redirects", maxRedirects)
|
return fmt.Errorf("stopped after %d redirects", maxRedirects)
|
||||||
}
|
}
|
||||||
|
if isObviousPrivateHost(req.URL.Hostname()) {
|
||||||
|
return fmt.Errorf("redirect target is private or local network host")
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if fetchLimitBytes <= 0 {
|
if fetchLimitBytes <= 0 {
|
||||||
|
|
@ -781,6 +903,13 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
return ErrorResult("missing domain in URL")
|
return ErrorResult("missing domain in URL")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Lightweight pre-flight: block obvious localhost/literal-IP without DNS resolution.
|
||||||
|
// The real SSRF guard is newSafeDialContext at connect time.
|
||||||
|
hostname := parsedURL.Hostname()
|
||||||
|
if isObviousPrivateHost(hostname) {
|
||||||
|
return ErrorResult("fetching private or local network hosts is not allowed")
|
||||||
|
}
|
||||||
|
|
||||||
maxChars := t.maxChars
|
maxChars := t.maxChars
|
||||||
if mc, ok := args["maxChars"].(float64); ok {
|
if mc, ok := args["maxChars"].(float64); ok {
|
||||||
if int(mc) > 100 {
|
if int(mc) > 100 {
|
||||||
|
|
@ -794,7 +923,6 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
}
|
}
|
||||||
|
|
||||||
req.Header.Set("User-Agent", userAgent)
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
resp, err := t.client.Do(req)
|
resp, err := t.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(fmt.Sprintf("request failed: %v", err))
|
return ErrorResult(fmt.Sprintf("request failed: %v", err))
|
||||||
|
|
@ -885,3 +1013,127 @@ func (t *WebFetchTool) extractText(htmlContent string) string {
|
||||||
|
|
||||||
return strings.Join(cleanLines, "\n")
|
return strings.Join(cleanLines, "\n")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// newSafeDialContext re-resolves DNS at connect time to mitigate DNS rebinding (TOCTOU)
|
||||||
|
// where a hostname resolves to a public IP during pre-flight but a private IP at connect time.
|
||||||
|
func newSafeDialContext(dialer *net.Dialer) func(context.Context, string, string) (net.Conn, error) {
|
||||||
|
return func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||||
|
if allowPrivateWebFetchHosts.Load() {
|
||||||
|
return dialer.DialContext(ctx, network, address)
|
||||||
|
}
|
||||||
|
|
||||||
|
host, port, err := net.SplitHostPort(address)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid target address %q: %w", address, err)
|
||||||
|
}
|
||||||
|
if host == "" {
|
||||||
|
return nil, fmt.Errorf("empty target host")
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip := net.ParseIP(host); ip != nil {
|
||||||
|
if isPrivateOrRestrictedIP(ip) {
|
||||||
|
return nil, fmt.Errorf("blocked private or local target: %s", host)
|
||||||
|
}
|
||||||
|
return dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port))
|
||||||
|
}
|
||||||
|
|
||||||
|
ipAddrs, err := net.DefaultResolver.LookupIPAddr(ctx, host)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to resolve %s: %w", host, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
attempted := 0
|
||||||
|
var lastErr error
|
||||||
|
for _, ipAddr := range ipAddrs {
|
||||||
|
if isPrivateOrRestrictedIP(ipAddr.IP) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
attempted++
|
||||||
|
conn, err := dialer.DialContext(ctx, network, net.JoinHostPort(ipAddr.IP.String(), port))
|
||||||
|
if err == nil {
|
||||||
|
return conn, nil
|
||||||
|
}
|
||||||
|
lastErr = err
|
||||||
|
}
|
||||||
|
|
||||||
|
if attempted == 0 {
|
||||||
|
return nil, fmt.Errorf("all resolved addresses for %s are private or restricted", host)
|
||||||
|
}
|
||||||
|
if lastErr != nil {
|
||||||
|
return nil, fmt.Errorf("failed connecting to public addresses for %s: %w", host, lastErr)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed connecting to public addresses for %s", host)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isObviousPrivateHost performs a lightweight, no-DNS check for obviously private hosts.
|
||||||
|
// It catches localhost, literal private IPs, and empty hosts. It does NOT resolve DNS —
|
||||||
|
// the real SSRF guard is newSafeDialContext which checks IPs at connect time.
|
||||||
|
func isObviousPrivateHost(host string) bool {
|
||||||
|
if allowPrivateWebFetchHosts.Load() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
h := strings.ToLower(strings.TrimSpace(host))
|
||||||
|
h = strings.TrimSuffix(h, ".")
|
||||||
|
if h == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if h == "localhost" || strings.HasSuffix(h, ".localhost") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip := net.ParseIP(h); ip != nil {
|
||||||
|
return isPrivateOrRestrictedIP(ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// isPrivateOrRestrictedIP returns true for IPs that should never be reached via web_fetch:
|
||||||
|
// RFC 1918, loopback, link-local (incl. cloud metadata 169.254.x.x), carrier-grade NAT,
|
||||||
|
// IPv6 unique-local (fc00::/7), 6to4 (2002::/16), and Teredo (2001:0000::/32).
|
||||||
|
func isPrivateOrRestrictedIP(ip net.IP) bool {
|
||||||
|
if ip == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() ||
|
||||||
|
ip.IsMulticast() || ip.IsUnspecified() {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip4 := ip.To4(); ip4 != nil {
|
||||||
|
// IPv4 private, loopback, link-local, and carrier-grade NAT ranges.
|
||||||
|
if ip4[0] == 10 ||
|
||||||
|
ip4[0] == 127 ||
|
||||||
|
ip4[0] == 0 ||
|
||||||
|
(ip4[0] == 172 && ip4[1] >= 16 && ip4[1] <= 31) ||
|
||||||
|
(ip4[0] == 192 && ip4[1] == 168) ||
|
||||||
|
(ip4[0] == 169 && ip4[1] == 254) ||
|
||||||
|
(ip4[0] == 100 && ip4[1] >= 64 && ip4[1] <= 127) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(ip) == net.IPv6len {
|
||||||
|
// IPv6 unique local addresses (fc00::/7)
|
||||||
|
if (ip[0] & 0xfe) == 0xfc {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// 6to4 addresses (2002::/16): check the embedded IPv4 at bytes [2:6].
|
||||||
|
if ip[0] == 0x20 && ip[1] == 0x02 {
|
||||||
|
embedded := net.IPv4(ip[2], ip[3], ip[4], ip[5])
|
||||||
|
return isPrivateOrRestrictedIP(embedded)
|
||||||
|
}
|
||||||
|
// Teredo (2001:0000::/32): client IPv4 is at bytes [12:16], XOR-inverted.
|
||||||
|
if ip[0] == 0x20 && ip[1] == 0x01 && ip[2] == 0x00 && ip[3] == 0x00 {
|
||||||
|
client := net.IPv4(ip[12]^0xff, ip[13]^0xff, ip[14]^0xff, ip[15]^0xff)
|
||||||
|
return isPrivateOrRestrictedIP(client)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
@ -18,6 +19,8 @@ const testFetchLimit = int64(10 * 1024 * 1024)
|
||||||
|
|
||||||
// TestWebTool_WebFetch_Success verifies successful URL fetching
|
// TestWebTool_WebFetch_Success verifies successful URL fetching
|
||||||
func TestWebTool_WebFetch_Success(t *testing.T) {
|
func TestWebTool_WebFetch_Success(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.Header().Set("Content-Type", "text/html")
|
w.Header().Set("Content-Type", "text/html")
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
|
|
@ -55,6 +58,8 @@ func TestWebTool_WebFetch_Success(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebFetch_JSON verifies JSON content handling
|
// TestWebTool_WebFetch_JSON verifies JSON content handling
|
||||||
func TestWebTool_WebFetch_JSON(t *testing.T) {
|
func TestWebTool_WebFetch_JSON(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
testData := map[string]string{"key": "value", "number": "123"}
|
testData := map[string]string{"key": "value", "number": "123"}
|
||||||
expectedJSON, _ := json.MarshalIndent(testData, "", " ")
|
expectedJSON, _ := json.MarshalIndent(testData, "", " ")
|
||||||
|
|
||||||
|
|
@ -163,6 +168,8 @@ func TestWebTool_WebFetch_MissingURL(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebFetch_Truncation verifies content truncation
|
// TestWebTool_WebFetch_Truncation verifies content truncation
|
||||||
func TestWebTool_WebFetch_Truncation(t *testing.T) {
|
func TestWebTool_WebFetch_Truncation(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
longContent := strings.Repeat("x", 20000)
|
longContent := strings.Repeat("x", 20000)
|
||||||
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|
@ -205,6 +212,8 @@ func TestWebTool_WebFetch_Truncation(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestWebFetchTool_PayloadTooLarge(t *testing.T) {
|
func TestWebFetchTool_PayloadTooLarge(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
// Create a mock HTTP server
|
// Create a mock HTTP server
|
||||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.Header().Set("Content-Type", "text/html")
|
w.Header().Set("Content-Type", "text/html")
|
||||||
|
|
@ -249,7 +258,7 @@ func TestWebFetchTool_PayloadTooLarge(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
|
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
|
||||||
func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""})
|
tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKeys: nil})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Unexpected error: %v", err)
|
t.Fatalf("Unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -269,7 +278,11 @@ func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
|
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
|
||||||
func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5})
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
BraveEnabled: true,
|
||||||
|
BraveAPIKeys: []string{"test-key"},
|
||||||
|
BraveMaxResults: 5,
|
||||||
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Unexpected error: %v", err)
|
t.Fatalf("Unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -286,6 +299,8 @@ func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebFetch_HTMLExtraction verifies HTML text extraction
|
// TestWebTool_WebFetch_HTMLExtraction verifies HTML text extraction
|
||||||
func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.Header().Set("Content-Type", "text/html")
|
w.Header().Set("Content-Type", "text/html")
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
|
|
@ -400,6 +415,205 @@ func TestWebFetchTool_extractText(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func withPrivateWebFetchHostsAllowed(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
previous := allowPrivateWebFetchHosts.Load()
|
||||||
|
allowPrivateWebFetchHosts.Store(true)
|
||||||
|
t.Cleanup(func() {
|
||||||
|
allowPrivateWebFetchHosts.Store(previous)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebTool_WebFetch_PrivateHostBlocked(t *testing.T) {
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": "http://127.0.0.1:0",
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Errorf("expected error for private host URL, got success")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "private or local network") &&
|
||||||
|
!strings.Contains(result.ForUser, "private or local network") {
|
||||||
|
t.Errorf("expected private host block message, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebTool_WebFetch_PrivateHostAllowedForTests(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "text/plain")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
w.Write([]byte("ok"))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": server.URL,
|
||||||
|
})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Errorf("expected success when private host access is allowed in tests, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebFetch_BlocksIPv4MappedIPv6Loopback verifies ::ffff:127.0.0.1 is blocked
|
||||||
|
func TestWebFetch_BlocksIPv4MappedIPv6Loopback(t *testing.T) {
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": "http://[::ffff:127.0.0.1]:0",
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("expected error for IPv4-mapped IPv6 loopback URL, got success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebFetch_BlocksMetadataIP verifies 169.254.169.254 is blocked
|
||||||
|
func TestWebFetch_BlocksMetadataIP(t *testing.T) {
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": "http://169.254.169.254/latest/meta-data",
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("expected error for cloud metadata IP, got success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebFetch_BlocksIPv6UniqueLocal verifies fc00::/7 addresses are blocked
|
||||||
|
func TestWebFetch_BlocksIPv6UniqueLocal(t *testing.T) {
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": "http://[fd00::1]:0",
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("expected error for IPv6 unique local address, got success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebFetch_Blocks6to4WithPrivateEmbed verifies 6to4 with private embedded IPv4 is blocked
|
||||||
|
func TestWebFetch_Blocks6to4WithPrivateEmbed(t *testing.T) {
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
// 2002:7f00:0001::1 embeds 127.0.0.1
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": "http://[2002:7f00:0001::1]:0",
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("expected error for 6to4 with private embedded IPv4, got success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebFetch_Allows6to4WithPublicEmbed verifies 6to4 with public embedded IPv4 is NOT blocked
|
||||||
|
func TestWebFetch_Allows6to4WithPublicEmbed(t *testing.T) {
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
// 2002:0801:0101::1 embeds 8.1.1.1 (public) — pre-flight should pass,
|
||||||
|
// connection will fail (no listener) but that's after the SSRF check.
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": "http://[2002:0801:0101::1]:0",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Should NOT be blocked by SSRF check — error should be connection failure, not "private"
|
||||||
|
if result.IsError && strings.Contains(result.ForLLM, "private") {
|
||||||
|
t.Error("6to4 with public embedded IPv4 should not be blocked as private")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebFetch_RedirectToPrivateBlocked verifies redirects to private IPs are blocked
|
||||||
|
func TestWebFetch_RedirectToPrivateBlocked(t *testing.T) {
|
||||||
|
withPrivateWebFetchHostsAllowed(t)
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Redirect to a private IP
|
||||||
|
http.Redirect(w, r, "http://10.0.0.1/secret", http.StatusFound)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
// Temporarily disable private host allowance for the redirect check
|
||||||
|
allowPrivateWebFetchHosts.Store(false)
|
||||||
|
defer allowPrivateWebFetchHosts.Store(true)
|
||||||
|
|
||||||
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create web fetch tool: %v", err)
|
||||||
|
}
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"url": server.URL,
|
||||||
|
})
|
||||||
|
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("expected error when redirecting to private IP, got success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIsPrivateOrRestrictedIP_Table tests IP classification logic
|
||||||
|
func TestIsPrivateOrRestrictedIP_Table(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
blocked bool
|
||||||
|
desc string
|
||||||
|
}{
|
||||||
|
{"127.0.0.1", true, "IPv4 loopback"},
|
||||||
|
{"10.0.0.1", true, "IPv4 private class A"},
|
||||||
|
{"172.16.0.1", true, "IPv4 private class B"},
|
||||||
|
{"192.168.1.1", true, "IPv4 private class C"},
|
||||||
|
{"169.254.169.254", true, "link-local / cloud metadata"},
|
||||||
|
{"100.64.0.1", true, "carrier-grade NAT"},
|
||||||
|
{"0.0.0.0", true, "unspecified"},
|
||||||
|
{"8.8.8.8", false, "public DNS"},
|
||||||
|
{"1.1.1.1", false, "public DNS"},
|
||||||
|
{"::1", true, "IPv6 loopback"},
|
||||||
|
{"::ffff:127.0.0.1", true, "IPv4-mapped IPv6 loopback"},
|
||||||
|
{"::ffff:10.0.0.1", true, "IPv4-mapped IPv6 private"},
|
||||||
|
{"fc00::1", true, "IPv6 unique local"},
|
||||||
|
{"fd00::1", true, "IPv6 unique local"},
|
||||||
|
{"2002:7f00:0001::1", true, "6to4 with embedded 127.x (private)"},
|
||||||
|
{"2002:0a00:0001::1", true, "6to4 with embedded 10.0.0.1 (private)"},
|
||||||
|
{"2002:0801:0101::1", false, "6to4 with embedded 8.1.1.1 (public)"},
|
||||||
|
{"2001:0000:4136:e378:8000:63bf:f5ff:fffe", true, "Teredo with client 10.0.0.1 (private)"},
|
||||||
|
{"2001:0000:4136:e378:8000:63bf:f7f6:fefe", false, "Teredo with client 8.9.1.1 (public)"},
|
||||||
|
{"2607:f8b0:4004:800::200e", false, "public IPv6 (Google)"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.desc, func(t *testing.T) {
|
||||||
|
ip := net.ParseIP(tt.ip)
|
||||||
|
if ip == nil {
|
||||||
|
t.Fatalf("failed to parse IP: %s", tt.ip)
|
||||||
|
}
|
||||||
|
got := isPrivateOrRestrictedIP(ip)
|
||||||
|
if got != tt.blocked {
|
||||||
|
t.Errorf("isPrivateOrRestrictedIP(%s) = %v, want %v", tt.ip, got, tt.blocked)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestWebTool_WebFetch_MissingDomain verifies error handling for URL without domain
|
// TestWebTool_WebFetch_MissingDomain verifies error handling for URL without domain
|
||||||
func TestWebTool_WebFetch_MissingDomain(t *testing.T) {
|
func TestWebTool_WebFetch_MissingDomain(t *testing.T) {
|
||||||
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
tool, err := NewWebFetchTool(50000, testFetchLimit)
|
||||||
|
|
@ -553,7 +767,7 @@ func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
|
||||||
t.Run("perplexity", func(t *testing.T) {
|
t.Run("perplexity", func(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
PerplexityEnabled: true,
|
PerplexityEnabled: true,
|
||||||
PerplexityAPIKey: "k",
|
PerplexityAPIKeys: []string{"k"},
|
||||||
PerplexityMaxResults: 3,
|
PerplexityMaxResults: 3,
|
||||||
Proxy: "http://127.0.0.1:7890",
|
Proxy: "http://127.0.0.1:7890",
|
||||||
})
|
})
|
||||||
|
|
@ -572,7 +786,7 @@ func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
|
||||||
t.Run("brave", func(t *testing.T) {
|
t.Run("brave", func(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
BraveEnabled: true,
|
BraveEnabled: true,
|
||||||
BraveAPIKey: "k",
|
BraveAPIKeys: []string{"k"},
|
||||||
BraveMaxResults: 3,
|
BraveMaxResults: 3,
|
||||||
Proxy: "http://127.0.0.1:7890",
|
Proxy: "http://127.0.0.1:7890",
|
||||||
})
|
})
|
||||||
|
|
@ -650,7 +864,7 @@ func TestWebTool_TavilySearch_Success(t *testing.T) {
|
||||||
|
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
TavilyEnabled: true,
|
TavilyEnabled: true,
|
||||||
TavilyAPIKey: "test-key",
|
TavilyAPIKeys: []string{"test-key"},
|
||||||
TavilyBaseURL: server.URL,
|
TavilyBaseURL: server.URL,
|
||||||
TavilyMaxResults: 5,
|
TavilyMaxResults: 5,
|
||||||
})
|
})
|
||||||
|
|
@ -682,6 +896,121 @@ func TestWebTool_TavilySearch_Success(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAPIKeyPool(t *testing.T) {
|
||||||
|
pool := NewAPIKeyPool([]string{"key1", "key2", "key3"})
|
||||||
|
if len(pool.keys) != 3 {
|
||||||
|
t.Fatalf("expected 3 keys, got %d", len(pool.keys))
|
||||||
|
}
|
||||||
|
if pool.keys[0] != "key1" || pool.keys[1] != "key2" || pool.keys[2] != "key3" {
|
||||||
|
t.Fatalf("unexpected keys: %v", pool.keys)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test Iterator: each iterator should cover all keys exactly once
|
||||||
|
iter := pool.NewIterator()
|
||||||
|
expected := []string{"key1", "key2", "key3"}
|
||||||
|
for i, want := range expected {
|
||||||
|
k, ok := iter.Next()
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("iter.Next() returned false at step %d", i)
|
||||||
|
}
|
||||||
|
if k != want {
|
||||||
|
t.Errorf("step %d: expected %s, got %s", i, want, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Should be exhausted
|
||||||
|
if _, ok := iter.Next(); ok {
|
||||||
|
t.Errorf("expected iterator exhausted after all keys")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second iterator starts at next position (load balancing)
|
||||||
|
iter2 := pool.NewIterator()
|
||||||
|
k, ok := iter2.Next()
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("iter2.Next() returned false")
|
||||||
|
}
|
||||||
|
if k != "key2" {
|
||||||
|
t.Errorf("expected key2 (round-robin), got %s", k)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Empty pool
|
||||||
|
emptyPool := NewAPIKeyPool([]string{})
|
||||||
|
emptyIter := emptyPool.NewIterator()
|
||||||
|
if _, ok := emptyIter.Next(); ok {
|
||||||
|
t.Errorf("expected false for empty pool")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single key pool
|
||||||
|
singlePool := NewAPIKeyPool([]string{"single"})
|
||||||
|
singleIter := singlePool.NewIterator()
|
||||||
|
if k, ok := singleIter.Next(); !ok || k != "single" {
|
||||||
|
t.Errorf("expected single, got %s (ok=%v)", k, ok)
|
||||||
|
}
|
||||||
|
if _, ok := singleIter.Next(); ok {
|
||||||
|
t.Errorf("expected exhausted after single key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebTool_TavilySearch_Failover(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var payload map[string]any
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
|
||||||
|
t.Fatalf("failed to decode payload: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
apiKey := payload["api_key"].(string)
|
||||||
|
|
||||||
|
if apiKey == "key1" {
|
||||||
|
w.WriteHeader(http.StatusTooManyRequests)
|
||||||
|
w.Write([]byte("Rate limited"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if apiKey == "key2" {
|
||||||
|
// Success
|
||||||
|
response := map[string]any{
|
||||||
|
"results": []map[string]any{
|
||||||
|
{
|
||||||
|
"title": "Success Result",
|
||||||
|
"url": "https://example.com/success",
|
||||||
|
"content": "Success content",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
json.NewEncoder(w).Encode(response)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
TavilyEnabled: true,
|
||||||
|
TavilyAPIKeys: []string{"key1", "key2"},
|
||||||
|
TavilyBaseURL: server.URL,
|
||||||
|
TavilyMaxResults: 5,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewWebSearchTool() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]any{
|
||||||
|
"query": "test query",
|
||||||
|
}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Errorf("Expected success, got Error: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForUser, "Success Result") {
|
||||||
|
t.Errorf("Expected failover to second key and success result, got: %s", result.ForUser)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestWebTool_GLMSearch_Success(t *testing.T) {
|
func TestWebTool_GLMSearch_Success(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.Method != "POST" {
|
if r.Method != "POST" {
|
||||||
|
|
|
||||||
|
|
@ -2,9 +2,18 @@ package utils
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"unicode"
|
"unicode"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Global variable to disable truncation
|
||||||
|
var disableTruncation atomic.Bool
|
||||||
|
|
||||||
|
// SetDisableTruncation globally enables or disables string truncation
|
||||||
|
func SetDisableTruncation(enabled bool) {
|
||||||
|
disableTruncation.Store(enabled)
|
||||||
|
}
|
||||||
|
|
||||||
// SanitizeMessageContent removes Unicode control characters, format characters (RTL overrides,
|
// SanitizeMessageContent removes Unicode control characters, format characters (RTL overrides,
|
||||||
// zero-width characters), and other non-graphic characters that could confuse an LLM
|
// zero-width characters), and other non-graphic characters that could confuse an LLM
|
||||||
// or cause display issues in the agent UI.
|
// or cause display issues in the agent UI.
|
||||||
|
|
@ -30,6 +39,10 @@ func SanitizeMessageContent(input string) string {
|
||||||
// Handles multi-byte Unicode characters properly.
|
// Handles multi-byte Unicode characters properly.
|
||||||
// If the string is truncated, "..." is appended to indicate truncation.
|
// If the string is truncated, "..." is appended to indicate truncation.
|
||||||
func Truncate(s string, maxLen int) string {
|
func Truncate(s string, maxLen int) string {
|
||||||
|
// If the no-truncate flag is active, it returns the full string
|
||||||
|
if disableTruncation.Load() {
|
||||||
|
return s
|
||||||
|
}
|
||||||
if maxLen <= 0 {
|
if maxLen <= 0 {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,6 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
@ -17,36 +16,11 @@ func (h *Handler) registerConfigRoutes(mux *http.ServeMux) {
|
||||||
mux.HandleFunc("PATCH /api/config", h.handlePatchConfig)
|
mux.HandleFunc("PATCH /api/config", h.handlePatchConfig)
|
||||||
}
|
}
|
||||||
|
|
||||||
// loadFilteredConfig loads the configuration and filters out default placeholder credentials
|
|
||||||
// (like API limits/keys) if the configuration file has not been created yet by the user.
|
|
||||||
func (h *Handler) loadFilteredConfig() (*config.Config, error) {
|
|
||||||
cfg, err := config.LoadConfig(h.configPath)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
configExists := false
|
|
||||||
if h.configPath != "" {
|
|
||||||
if _, err := os.Stat(h.configPath); err == nil {
|
|
||||||
configExists = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !configExists {
|
|
||||||
for i := range cfg.ModelList {
|
|
||||||
cfg.ModelList[i].APIKey = ""
|
|
||||||
cfg.ModelList[i].AuthMethod = ""
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return cfg, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleGetConfig returns the complete system configuration.
|
// handleGetConfig returns the complete system configuration.
|
||||||
//
|
//
|
||||||
// GET /api/config
|
// GET /api/config
|
||||||
func (h *Handler) handleGetConfig(w http.ResponseWriter, r *http.Request) {
|
func (h *Handler) handleGetConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
cfg, err := h.loadFilteredConfig()
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
return
|
return
|
||||||
|
|
@ -74,6 +48,9 @@ func (h *Handler) handleUpdateConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if execAllowRemoteOmitted(body) {
|
||||||
|
cfg.Tools.Exec.AllowRemote = config.DefaultConfig().Tools.Exec.AllowRemote
|
||||||
|
}
|
||||||
|
|
||||||
if errs := validateConfig(&cfg); len(errs) > 0 {
|
if errs := validateConfig(&cfg); len(errs) > 0 {
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
|
@ -94,6 +71,20 @@ func (h *Handler) handleUpdateConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func execAllowRemoteOmitted(body []byte) bool {
|
||||||
|
var raw struct {
|
||||||
|
Tools *struct {
|
||||||
|
Exec *struct {
|
||||||
|
AllowRemote *bool `json:"allow_remote"`
|
||||||
|
} `json:"exec"`
|
||||||
|
} `json:"tools"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &raw); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return raw.Tools == nil || raw.Tools.Exec == nil || raw.Tools.Exec.AllowRemote == nil
|
||||||
|
}
|
||||||
|
|
||||||
// handlePatchConfig partially updates the system configuration using JSON Merge Patch (RFC 7396).
|
// handlePatchConfig partially updates the system configuration using JSON Merge Patch (RFC 7396).
|
||||||
// Only the fields present in the request body will be updated; all other fields remain unchanged.
|
// Only the fields present in the request body will be updated; all other fields remain unchanged.
|
||||||
//
|
//
|
||||||
|
|
|
||||||
88
web/backend/api/config_test.go
Normal file
88
web/backend/api/config_test.go
Normal file
|
|
@ -0,0 +1,88 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestHandleUpdateConfig_PreservesExecAllowRemoteDefaultWhenOmitted(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPut, "/api/config", bytes.NewBufferString(`{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"workspace": "~/.picoclaw/workspace"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "custom-default",
|
||||||
|
"model": "openai/gpt-4o",
|
||||||
|
"api_key": "sk-default"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if !cfg.Tools.Exec.AllowRemote {
|
||||||
|
t.Fatal("tools.exec.allow_remote should remain true when omitted from PUT /api/config")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleUpdateConfig_DoesNotInheritDefaultModelFields(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPut, "/api/config", bytes.NewBufferString(`{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"workspace": "~/.picoclaw/workspace"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "custom-default",
|
||||||
|
"model": "openai/gpt-4o",
|
||||||
|
"api_key": "sk-default"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if got := cfg.ModelList[0].APIBase; got != "" {
|
||||||
|
t.Fatalf("model_list[0].api_base = %q, want empty string", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -10,7 +10,6 @@ import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
"runtime"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
@ -19,6 +18,7 @@ import (
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
// gateway holds the state for the managed gateway process.
|
// gateway holds the state for the managed gateway process.
|
||||||
|
|
@ -36,6 +36,7 @@ var gateway = struct {
|
||||||
func (h *Handler) registerGatewayRoutes(mux *http.ServeMux) {
|
func (h *Handler) registerGatewayRoutes(mux *http.ServeMux) {
|
||||||
mux.HandleFunc("GET /api/gateway/status", h.handleGatewayStatus)
|
mux.HandleFunc("GET /api/gateway/status", h.handleGatewayStatus)
|
||||||
mux.HandleFunc("GET /api/gateway/events", h.handleGatewayEvents)
|
mux.HandleFunc("GET /api/gateway/events", h.handleGatewayEvents)
|
||||||
|
mux.HandleFunc("POST /api/gateway/logs/clear", h.handleGatewayClearLogs)
|
||||||
mux.HandleFunc("POST /api/gateway/start", h.handleGatewayStart)
|
mux.HandleFunc("POST /api/gateway/start", h.handleGatewayStart)
|
||||||
mux.HandleFunc("POST /api/gateway/stop", h.handleGatewayStop)
|
mux.HandleFunc("POST /api/gateway/stop", h.handleGatewayStop)
|
||||||
mux.HandleFunc("POST /api/gateway/restart", h.handleGatewayRestart)
|
mux.HandleFunc("POST /api/gateway/restart", h.handleGatewayRestart)
|
||||||
|
|
@ -89,11 +90,12 @@ func (h *Handler) gatewayStartReady() (bool, string, error) {
|
||||||
return false, fmt.Sprintf("default model %q is invalid", modelName), nil
|
return false, fmt.Sprintf("default model %q is invalid", modelName), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
hasCredential := strings.TrimSpace(modelCfg.APIKey) != "" ||
|
if !hasModelConfiguration(*modelCfg) {
|
||||||
strings.TrimSpace(modelCfg.AuthMethod) != ""
|
|
||||||
if !hasCredential {
|
|
||||||
return false, fmt.Sprintf("default model %q has no credentials configured", modelName), nil
|
return false, fmt.Sprintf("default model %q has no credentials configured", modelName), nil
|
||||||
}
|
}
|
||||||
|
if requiresRuntimeProbe(*modelCfg) && !probeLocalModelAvailability(*modelCfg) {
|
||||||
|
return false, fmt.Sprintf("default model %q is not reachable", modelName), nil
|
||||||
|
}
|
||||||
|
|
||||||
return true, "", nil
|
return true, "", nil
|
||||||
}
|
}
|
||||||
|
|
@ -131,9 +133,19 @@ func isCmdProcessAliveLocked(cmd *exec.Cmd) bool {
|
||||||
|
|
||||||
func (h *Handler) startGatewayLocked() (int, error) {
|
func (h *Handler) startGatewayLocked() (int, error) {
|
||||||
// Locate the picoclaw executable
|
// Locate the picoclaw executable
|
||||||
execPath := findPicoclawBinary()
|
execPath := utils.FindPicoclawBinary()
|
||||||
|
|
||||||
cmd := exec.Command(execPath, "gateway")
|
cmd := exec.Command(execPath, "gateway")
|
||||||
|
cmd.Env = os.Environ()
|
||||||
|
// Forward the launcher's config path via the environment variable that
|
||||||
|
// GetConfigPath() already reads, so the gateway sub-process uses the same
|
||||||
|
// config file without requiring a --config flag on the gateway subcommand.
|
||||||
|
if h.configPath != "" {
|
||||||
|
cmd.Env = append(cmd.Env, "PICOCLAW_CONFIG="+h.configPath)
|
||||||
|
}
|
||||||
|
if host := h.gatewayHostOverride(); host != "" {
|
||||||
|
cmd.Env = append(cmd.Env, "PICOCLAW_GATEWAY_HOST="+host)
|
||||||
|
}
|
||||||
|
|
||||||
stdoutPipe, err := cmd.StdoutPipe()
|
stdoutPipe, err := cmd.StdoutPipe()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -201,10 +213,7 @@ func (h *Handler) startGatewayLocked() (int, error) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
healthHost := "127.0.0.1"
|
healthHost := gatewayProbeHost(h.effectiveGatewayBindHost(cfg))
|
||||||
if cfg.Gateway.Host != "" && cfg.Gateway.Host != "0.0.0.0" {
|
|
||||||
healthHost = cfg.Gateway.Host
|
|
||||||
}
|
|
||||||
healthPort := cfg.Gateway.Port
|
healthPort := cfg.Gateway.Port
|
||||||
if healthPort == 0 {
|
if healthPort == 0 {
|
||||||
healthPort = 18790
|
healthPort = 18790
|
||||||
|
|
@ -347,6 +356,20 @@ func (h *Handler) handleGatewayRestart(w http.ResponseWriter, r *http.Request) {
|
||||||
h.handleGatewayStart(w, r)
|
h.handleGatewayStart(w, r)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// handleGatewayClearLogs clears the in-memory gateway log buffer.
|
||||||
|
//
|
||||||
|
// POST /api/gateway/logs/clear
|
||||||
|
func (h *Handler) handleGatewayClearLogs(w http.ResponseWriter, r *http.Request) {
|
||||||
|
gateway.logs.Clear()
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "cleared",
|
||||||
|
"log_total": 0,
|
||||||
|
"log_run_id": gateway.logs.RunID(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// handleGatewayStatus returns the gateway run status, health info, and logs.
|
// handleGatewayStatus returns the gateway run status, health info, and logs.
|
||||||
//
|
//
|
||||||
// GET /api/gateway/status
|
// GET /api/gateway/status
|
||||||
|
|
@ -369,9 +392,7 @@ func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) {
|
||||||
host := "127.0.0.1"
|
host := "127.0.0.1"
|
||||||
port := 18790
|
port := 18790
|
||||||
if err == nil && cfg != nil {
|
if err == nil && cfg != nil {
|
||||||
if cfg.Gateway.Host != "" && cfg.Gateway.Host != "0.0.0.0" {
|
host = gatewayProbeHost(h.effectiveGatewayBindHost(cfg))
|
||||||
host = cfg.Gateway.Host
|
|
||||||
}
|
|
||||||
if cfg.Gateway.Port != 0 {
|
if cfg.Gateway.Port != 0 {
|
||||||
port = cfg.Gateway.Port
|
port = cfg.Gateway.Port
|
||||||
}
|
}
|
||||||
|
|
@ -529,22 +550,6 @@ func (h *Handler) currentGatewayStatus() string {
|
||||||
return string(encoded)
|
return string(encoded)
|
||||||
}
|
}
|
||||||
|
|
||||||
// findPicoclawBinary locates the picoclaw executable.
|
|
||||||
// Tries the same directory as the current executable first, then falls back to $PATH.
|
|
||||||
func findPicoclawBinary() string {
|
|
||||||
if exe, err := os.Executable(); err == nil {
|
|
||||||
dir := filepath.Dir(exe)
|
|
||||||
candidate := filepath.Join(dir, "picoclaw")
|
|
||||||
if runtime.GOOS == "windows" {
|
|
||||||
candidate += ".exe"
|
|
||||||
}
|
|
||||||
if info, err := os.Stat(candidate); err == nil && !info.IsDir() {
|
|
||||||
return candidate
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return "picoclaw"
|
|
||||||
}
|
|
||||||
|
|
||||||
// scanPipe reads lines from r and appends them to buf. Returns when r reaches EOF.
|
// scanPipe reads lines from r and appends them to buf. Returns when r reaches EOF.
|
||||||
func scanPipe(r io.Reader, buf *LogBuffer) {
|
func scanPipe(r io.Reader, buf *LogBuffer) {
|
||||||
scanner := bufio.NewScanner(r)
|
scanner := bufio.NewScanner(r)
|
||||||
|
|
|
||||||
66
web/backend/api/gateway_host.go
Normal file
66
web/backend/api/gateway_host.go
Normal file
|
|
@ -0,0 +1,66 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (h *Handler) effectiveLauncherPublic() bool {
|
||||||
|
if h.serverPublicExplicit {
|
||||||
|
return h.serverPublic
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := h.loadLauncherConfig()
|
||||||
|
if err == nil {
|
||||||
|
return cfg.Public
|
||||||
|
}
|
||||||
|
|
||||||
|
return h.serverPublic
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) gatewayHostOverride() string {
|
||||||
|
if h.effectiveLauncherPublic() {
|
||||||
|
return "0.0.0.0"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) effectiveGatewayBindHost(cfg *config.Config) string {
|
||||||
|
if override := h.gatewayHostOverride(); override != "" {
|
||||||
|
return override
|
||||||
|
}
|
||||||
|
if cfg == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(cfg.Gateway.Host)
|
||||||
|
}
|
||||||
|
|
||||||
|
func gatewayProbeHost(bindHost string) string {
|
||||||
|
if bindHost == "" || bindHost == "0.0.0.0" {
|
||||||
|
return "127.0.0.1"
|
||||||
|
}
|
||||||
|
return bindHost
|
||||||
|
}
|
||||||
|
|
||||||
|
func requestHostName(r *http.Request) string {
|
||||||
|
reqHost, _, err := net.SplitHostPort(r.Host)
|
||||||
|
if err == nil {
|
||||||
|
return reqHost
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(r.Host) != "" {
|
||||||
|
return r.Host
|
||||||
|
}
|
||||||
|
return "127.0.0.1"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) buildWsURL(r *http.Request, cfg *config.Config) string {
|
||||||
|
host := h.effectiveGatewayBindHost(cfg)
|
||||||
|
if host == "" || host == "0.0.0.0" {
|
||||||
|
host = requestHostName(r)
|
||||||
|
}
|
||||||
|
return "ws://" + net.JoinHostPort(host, strconv.Itoa(cfg.Gateway.Port)) + "/pico/ws"
|
||||||
|
}
|
||||||
59
web/backend/api/gateway_host_test.go
Normal file
59
web/backend/api/gateway_host_test.go
Normal file
|
|
@ -0,0 +1,59 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGatewayHostOverrideUsesExplicitRuntimePublic(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
launcherPath := launcherconfig.PathForAppConfig(configPath)
|
||||||
|
if err := launcherconfig.Save(launcherPath, launcherconfig.Config{
|
||||||
|
Port: 18800,
|
||||||
|
Public: false,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("launcherconfig.Save() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
h.SetServerOptions(18800, true, true, nil)
|
||||||
|
|
||||||
|
if got := h.gatewayHostOverride(); got != "0.0.0.0" {
|
||||||
|
t.Fatalf("gatewayHostOverride() = %q, want %q", got, "0.0.0.0")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildWsURLUsesRequestHostWhenLauncherPublicSaved(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
launcherPath := launcherconfig.PathForAppConfig(configPath)
|
||||||
|
if err := launcherconfig.Save(launcherPath, launcherconfig.Config{
|
||||||
|
Port: 18800,
|
||||||
|
Public: true,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("launcherconfig.Save() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
h.SetServerOptions(18800, false, false, nil)
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Gateway.Host = "127.0.0.1"
|
||||||
|
cfg.Gateway.Port = 18790
|
||||||
|
|
||||||
|
req := httptest.NewRequest("GET", "http://launcher.local/api/pico/token", nil)
|
||||||
|
req.Host = "192.168.1.9:18800"
|
||||||
|
|
||||||
|
if got := h.buildWsURL(req, cfg); got != "ws://192.168.1.9:18790/pico/ws" {
|
||||||
|
t.Fatalf("buildWsURL() = %q, want %q", got, "ws://192.168.1.9:18790/pico/ws")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayProbeHostUsesLoopbackForWildcardBind(t *testing.T) {
|
||||||
|
if got := gatewayProbeHost("0.0.0.0"); got != "127.0.0.1" {
|
||||||
|
t.Fatalf("gatewayProbeHost() = %q, want %q", got, "127.0.0.1")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -4,11 +4,15 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestGatewayStartReady_NoDefaultModel(t *testing.T) {
|
func TestGatewayStartReady_NoDefaultModel(t *testing.T) {
|
||||||
|
|
@ -31,7 +35,8 @@ func TestGatewayStartReady_InvalidDefaultModel(t *testing.T) {
|
||||||
configPath := filepath.Join(t.TempDir(), "config.json")
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Model = "missing-model"
|
cfg.Agents.Defaults.Model = "missing-model"
|
||||||
if err := config.SaveConfig(configPath, cfg); err != nil {
|
err := config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("SaveConfig() error = %v", err)
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -53,7 +58,8 @@ func TestGatewayStartReady_ValidDefaultModel(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
|
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
|
||||||
cfg.ModelList[0].APIKey = "test-key"
|
cfg.ModelList[0].APIKey = "test-key"
|
||||||
if err := config.SaveConfig(configPath, cfg); err != nil {
|
err := config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("SaveConfig() error = %v", err)
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -73,7 +79,8 @@ func TestGatewayStartReady_DefaultModelWithoutCredential(t *testing.T) {
|
||||||
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
|
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
|
||||||
cfg.ModelList[0].APIKey = ""
|
cfg.ModelList[0].APIKey = ""
|
||||||
cfg.ModelList[0].AuthMethod = ""
|
cfg.ModelList[0].AuthMethod = ""
|
||||||
if err := config.SaveConfig(configPath, cfg); err != nil {
|
err := config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("SaveConfig() error = %v", err)
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -90,6 +97,195 @@ func TestGatewayStartReady_DefaultModelWithoutCredential(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_LocalModelWithoutAPIKey(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "local-vllm",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "http://localhost:8000/v1",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "local-vllm"
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = true, want false without a running local service")
|
||||||
|
}
|
||||||
|
if !strings.Contains(reason, "not reachable") {
|
||||||
|
t.Fatalf("gatewayStartReady() reason = %q, want contains %q", reason, "not reachable")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_LocalModelWithRunningService(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
return apiBase == "http://127.0.0.1:8000/v1" && modelID == "custom-model"
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "local-vllm",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "http://127.0.0.1:8000/v1",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "local-vllm"
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = false, want true with a running local service (reason=%q)", reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_RemoteVLLMWithAPIKeyDoesNotProbe(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
t.Fatalf("unexpected OpenAI-compatible probe for %q (%q)", apiBase, modelID)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "remote-vllm",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "https://models.example.com/v1",
|
||||||
|
APIKey: "remote-key",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "remote-vllm"
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = false, want true for remote vllm with api key (reason=%q)", reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_LocalOllamaUsesDefaultProbeBase(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
probeOllamaModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
return apiBase == "http://localhost:11434/v1" && modelID == "llama3"
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "local-ollama",
|
||||||
|
Model: "ollama/llama3",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "local-ollama"
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = false, want true with default Ollama probe base (reason=%q)", reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_OAuthModelRequiresStoredCredential(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "openai-oauth",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "openai-oauth"
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = true, want false without stored credential")
|
||||||
|
}
|
||||||
|
if !strings.Contains(reason, "no credentials configured") {
|
||||||
|
t.Fatalf("gatewayStartReady() reason = %q, want contains %q", reason, "no credentials configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
err = auth.SetCredential(oauthProviderOpenAI, &auth.AuthCredential{
|
||||||
|
AccessToken: "openai-token",
|
||||||
|
Provider: oauthProviderOpenAI,
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SetCredential() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ready, reason, err = h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = false, want true with stored credential (reason=%q)", reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestGatewayStatusIncludesStartConditionWhenNotReady(t *testing.T) {
|
func TestGatewayStatusIncludesStartConditionWhenNotReady(t *testing.T) {
|
||||||
configPath := filepath.Join(t.TempDir(), "config.json")
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
h := NewHandler(configPath)
|
h := NewHandler(configPath)
|
||||||
|
|
@ -120,3 +316,95 @@ func TestGatewayStatusIncludesStartConditionWhenNotReady(t *testing.T) {
|
||||||
t.Fatalf("gateway_start_reason missing or not string: %#v", body["gateway_start_reason"])
|
t.Fatalf("gateway_start_reason missing or not string: %#v", body["gateway_start_reason"])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestGatewayClearLogsResetsBufferedHistory(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
gateway.logs.Clear()
|
||||||
|
gateway.logs.Append("first line")
|
||||||
|
gateway.logs.Append("second line")
|
||||||
|
previousRunID := gateway.logs.RunID()
|
||||||
|
|
||||||
|
clearRec := httptest.NewRecorder()
|
||||||
|
clearReq := httptest.NewRequest(http.MethodPost, "/api/gateway/logs/clear", nil)
|
||||||
|
mux.ServeHTTP(clearRec, clearReq)
|
||||||
|
|
||||||
|
if clearRec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("clear status = %d, want %d", clearRec.Code, http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
var clearBody map[string]any
|
||||||
|
if err := json.Unmarshal(clearRec.Body.Bytes(), &clearBody); err != nil {
|
||||||
|
t.Fatalf("unmarshal clear response: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := clearBody["status"]; got != "cleared" {
|
||||||
|
t.Fatalf("clear status body = %#v, want %q", got, "cleared")
|
||||||
|
}
|
||||||
|
|
||||||
|
clearRunID, ok := clearBody["log_run_id"].(float64)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("log_run_id missing or not number: %#v", clearBody["log_run_id"])
|
||||||
|
}
|
||||||
|
if int(clearRunID) <= previousRunID {
|
||||||
|
t.Fatalf("log_run_id = %d, want > %d", int(clearRunID), previousRunID)
|
||||||
|
}
|
||||||
|
|
||||||
|
statusRec := httptest.NewRecorder()
|
||||||
|
statusReq := httptest.NewRequest(
|
||||||
|
http.MethodGet,
|
||||||
|
"/api/gateway/status?log_offset=0&log_run_id="+strconv.Itoa(previousRunID),
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
mux.ServeHTTP(statusRec, statusReq)
|
||||||
|
|
||||||
|
if statusRec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status code = %d, want %d", statusRec.Code, http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
var statusBody map[string]any
|
||||||
|
if err := json.Unmarshal(statusRec.Body.Bytes(), &statusBody); err != nil {
|
||||||
|
t.Fatalf("unmarshal status response: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
logs, ok := statusBody["logs"].([]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("logs missing or not array: %#v", statusBody["logs"])
|
||||||
|
}
|
||||||
|
if len(logs) != 0 {
|
||||||
|
t.Fatalf("logs len = %d, want 0", len(logs))
|
||||||
|
}
|
||||||
|
if got := statusBody["log_total"]; got != float64(0) {
|
||||||
|
t.Fatalf("log_total = %#v, want 0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindPicoclawBinary_EnvOverride(t *testing.T) {
|
||||||
|
// Create a temporary file to act as the mock binary
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
mockBinary := filepath.Join(tmpDir, "picoclaw-mock")
|
||||||
|
if err := os.WriteFile(mockBinary, []byte("mock"), 0o755); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Setenv("PICOCLAW_BINARY", mockBinary)
|
||||||
|
|
||||||
|
got := utils.FindPicoclawBinary()
|
||||||
|
if got != mockBinary {
|
||||||
|
t.Errorf("FindPicoclawBinary() = %q, want %q", got, mockBinary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindPicoclawBinary_EnvOverride_InvalidPath(t *testing.T) {
|
||||||
|
// When PICOCLAW_BINARY points to a non-existent path, fall through to next strategy
|
||||||
|
t.Setenv("PICOCLAW_BINARY", "/nonexistent/picoclaw-binary")
|
||||||
|
|
||||||
|
got := utils.FindPicoclawBinary()
|
||||||
|
// Should not return the invalid path; falls back to "picoclaw" or another found path
|
||||||
|
if got == "/nonexistent/picoclaw-binary" {
|
||||||
|
t.Errorf("FindPicoclawBinary() returned invalid env path %q, expected fallback", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -14,7 +14,7 @@ import (
|
||||||
func TestGetLauncherConfigUsesRuntimeFallback(t *testing.T) {
|
func TestGetLauncherConfigUsesRuntimeFallback(t *testing.T) {
|
||||||
configPath := filepath.Join(t.TempDir(), "config.json")
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
h := NewHandler(configPath)
|
h := NewHandler(configPath)
|
||||||
h.SetServerOptions(19999, true, []string{"192.168.1.0/24"})
|
h.SetServerOptions(19999, true, false, []string{"192.168.1.0/24"})
|
||||||
|
|
||||||
mux := http.NewServeMux()
|
mux := http.NewServeMux()
|
||||||
h.RegisterRoutes(mux)
|
h.RegisterRoutes(mux)
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@ import "sync"
|
||||||
|
|
||||||
// LogBuffer is a thread-safe ring buffer that stores the most recent N log lines.
|
// LogBuffer is a thread-safe ring buffer that stores the most recent N log lines.
|
||||||
// It supports incremental reads via LinesSince and tracks a runID that increments
|
// It supports incremental reads via LinesSince and tracks a runID that increments
|
||||||
// on each Reset (used to detect gateway restarts).
|
// whenever the buffer is reset or cleared so clients can detect log history resets.
|
||||||
type LogBuffer struct {
|
type LogBuffer struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
lines []string
|
lines []string
|
||||||
|
|
@ -45,6 +45,12 @@ func (b *LogBuffer) Reset() {
|
||||||
b.runID++
|
b.runID++
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Clear removes all buffered lines and increments the runID so clients treat
|
||||||
|
// subsequent reads as a new log stream.
|
||||||
|
func (b *LogBuffer) Clear() {
|
||||||
|
b.Reset()
|
||||||
|
}
|
||||||
|
|
||||||
// LinesSince returns lines appended after the given offset, the current total count, and the runID.
|
// LinesSince returns lines appended after the given offset, the current total count, and the runID.
|
||||||
// If offset >= total, no lines are returned. If offset is too old (evicted), all buffered lines are returned.
|
// If offset >= total, no lines are returned. If offset is too old (evicted), all buffered lines are returned.
|
||||||
func (b *LogBuffer) LinesSince(offset int) (lines []string, total int, runID int) {
|
func (b *LogBuffer) LinesSince(offset int) (lines []string, total int, runID int) {
|
||||||
|
|
|
||||||
324
web/backend/api/model_status.go
Normal file
324
web/backend/api/model_status.go
Normal file
|
|
@ -0,0 +1,324 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
const modelProbeTimeout = 800 * time.Millisecond
|
||||||
|
|
||||||
|
var (
|
||||||
|
probeTCPServiceFunc = probeTCPService
|
||||||
|
probeOllamaModelFunc = probeOllamaModel
|
||||||
|
probeOpenAICompatibleModelFunc = probeOpenAICompatibleModel
|
||||||
|
)
|
||||||
|
|
||||||
|
func hasModelConfiguration(m config.ModelConfig) bool {
|
||||||
|
authMethod := strings.ToLower(strings.TrimSpace(m.AuthMethod))
|
||||||
|
apiKey := strings.TrimSpace(m.APIKey)
|
||||||
|
|
||||||
|
if authMethod == "oauth" || authMethod == "token" {
|
||||||
|
if provider, ok := oauthProviderForModel(m.Model); ok {
|
||||||
|
cred, err := oauthGetCredential(provider)
|
||||||
|
if err != nil || cred == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(cred.AccessToken) != "" || strings.TrimSpace(cred.RefreshToken) != ""
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if requiresRuntimeProbe(m) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return apiKey != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// isModelConfigured reports whether a model is currently available to use.
|
||||||
|
// Local models must be reachable; remote/API-key models only need saved config.
|
||||||
|
func isModelConfigured(m config.ModelConfig) bool {
|
||||||
|
if !hasModelConfiguration(m) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if requiresRuntimeProbe(m) {
|
||||||
|
return probeLocalModelAvailability(m)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func requiresRuntimeProbe(m config.ModelConfig) bool {
|
||||||
|
authMethod := strings.ToLower(strings.TrimSpace(m.AuthMethod))
|
||||||
|
if authMethod == "local" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
switch modelProtocol(m.Model) {
|
||||||
|
case "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", "copilot":
|
||||||
|
return true
|
||||||
|
case "ollama", "vllm":
|
||||||
|
apiBase := strings.TrimSpace(m.APIBase)
|
||||||
|
return apiBase == "" || hasLocalAPIBase(apiBase)
|
||||||
|
}
|
||||||
|
|
||||||
|
if hasLocalAPIBase(m.APIBase) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func probeLocalModelAvailability(m config.ModelConfig) bool {
|
||||||
|
apiBase := modelProbeAPIBase(m)
|
||||||
|
protocol, modelID := splitModel(m.Model)
|
||||||
|
switch protocol {
|
||||||
|
case "ollama":
|
||||||
|
return probeOllamaModelFunc(apiBase, modelID)
|
||||||
|
case "vllm":
|
||||||
|
return probeOpenAICompatibleModelFunc(apiBase, modelID)
|
||||||
|
case "github-copilot", "copilot":
|
||||||
|
return probeTCPServiceFunc(apiBase)
|
||||||
|
case "claude-cli", "claudecli", "codex-cli", "codexcli":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
if hasLocalAPIBase(apiBase) {
|
||||||
|
return probeOpenAICompatibleModelFunc(apiBase, modelID)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func modelProbeAPIBase(m config.ModelConfig) string {
|
||||||
|
if apiBase := strings.TrimSpace(m.APIBase); apiBase != "" {
|
||||||
|
return normalizeModelProbeAPIBase(apiBase)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch modelProtocol(m.Model) {
|
||||||
|
case "ollama":
|
||||||
|
return "http://localhost:11434/v1"
|
||||||
|
case "vllm":
|
||||||
|
return "http://localhost:8000/v1"
|
||||||
|
case "github-copilot", "copilot":
|
||||||
|
return "localhost:4321"
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeModelProbeAPIBase(raw string) string {
|
||||||
|
u, err := parseAPIBase(raw)
|
||||||
|
if err != nil {
|
||||||
|
return strings.TrimSpace(raw)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch strings.ToLower(u.Hostname()) {
|
||||||
|
case "0.0.0.0":
|
||||||
|
u.Host = net.JoinHostPort("127.0.0.1", u.Port())
|
||||||
|
case "::":
|
||||||
|
u.Host = net.JoinHostPort("::1", u.Port())
|
||||||
|
default:
|
||||||
|
return strings.TrimSpace(raw)
|
||||||
|
}
|
||||||
|
|
||||||
|
if u.Port() == "" {
|
||||||
|
u.Host = u.Hostname()
|
||||||
|
}
|
||||||
|
|
||||||
|
return u.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func oauthProviderForModel(model string) (string, bool) {
|
||||||
|
switch modelProtocol(model) {
|
||||||
|
case "openai":
|
||||||
|
return oauthProviderOpenAI, true
|
||||||
|
case "anthropic":
|
||||||
|
return oauthProviderAnthropic, true
|
||||||
|
case "antigravity", "google-antigravity":
|
||||||
|
return oauthProviderGoogleAntigravity, true
|
||||||
|
default:
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func modelProtocol(model string) string {
|
||||||
|
protocol, _ := splitModel(model)
|
||||||
|
return protocol
|
||||||
|
}
|
||||||
|
|
||||||
|
func splitModel(model string) (protocol, modelID string) {
|
||||||
|
model = strings.ToLower(strings.TrimSpace(model))
|
||||||
|
protocol, _, found := strings.Cut(model, "/")
|
||||||
|
if !found {
|
||||||
|
return "openai", model
|
||||||
|
}
|
||||||
|
return protocol, strings.TrimSpace(model[strings.Index(model, "/")+1:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasLocalAPIBase(raw string) bool {
|
||||||
|
raw = strings.TrimSpace(raw)
|
||||||
|
if raw == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
u, err := url.Parse(raw)
|
||||||
|
if err != nil || u.Hostname() == "" {
|
||||||
|
u, err = url.Parse("//" + raw)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch strings.ToLower(u.Hostname()) {
|
||||||
|
case "localhost", "127.0.0.1", "::1", "0.0.0.0":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func probeTCPService(raw string) bool {
|
||||||
|
hostPort, err := hostPortFromAPIBase(raw)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := net.DialTimeout("tcp", hostPort, modelProbeTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = conn.Close()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func probeOllamaModel(apiBase, modelID string) bool {
|
||||||
|
root, err := apiRootFromAPIBase(apiBase)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Models []struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Model string `json:"model"`
|
||||||
|
} `json:"models"`
|
||||||
|
}
|
||||||
|
if err := getJSON(root+"/api/tags", &resp); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, model := range resp.Models {
|
||||||
|
if ollamaModelMatches(model.Name, modelID) || ollamaModelMatches(model.Model, modelID) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func probeOpenAICompatibleModel(apiBase, modelID string) bool {
|
||||||
|
if strings.TrimSpace(apiBase) == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Data []struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
} `json:"data"`
|
||||||
|
}
|
||||||
|
if err := getJSON(strings.TrimRight(strings.TrimSpace(apiBase), "/")+"/models", &resp); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, model := range resp.Data {
|
||||||
|
if strings.EqualFold(strings.TrimSpace(model.ID), modelID) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func getJSON(rawURL string, out any) error {
|
||||||
|
req, err := http.NewRequest(http.MethodGet, rawURL, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: modelProbeTimeout}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return fmt.Errorf("unexpected status %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
return json.NewDecoder(resp.Body).Decode(out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func apiRootFromAPIBase(raw string) (string, error) {
|
||||||
|
u, err := parseAPIBase(raw)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return (&url.URL{Scheme: u.Scheme, Host: u.Host}).String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func hostPortFromAPIBase(raw string) (string, error) {
|
||||||
|
u, err := parseAPIBase(raw)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
if port := u.Port(); port != "" {
|
||||||
|
return u.Host, nil
|
||||||
|
}
|
||||||
|
switch strings.ToLower(u.Scheme) {
|
||||||
|
case "https":
|
||||||
|
return net.JoinHostPort(u.Hostname(), "443"), nil
|
||||||
|
default:
|
||||||
|
return net.JoinHostPort(u.Hostname(), "80"), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseAPIBase(raw string) (*url.URL, error) {
|
||||||
|
raw = strings.TrimSpace(raw)
|
||||||
|
if raw == "" {
|
||||||
|
return nil, fmt.Errorf("empty api base")
|
||||||
|
}
|
||||||
|
|
||||||
|
u, err := url.Parse(raw)
|
||||||
|
if err == nil && u.Hostname() != "" {
|
||||||
|
return u, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
u, err = url.Parse("//" + raw)
|
||||||
|
if err != nil || u.Hostname() == "" {
|
||||||
|
return nil, fmt.Errorf("invalid api base %q", raw)
|
||||||
|
}
|
||||||
|
if u.Scheme == "" {
|
||||||
|
u.Scheme = "http"
|
||||||
|
}
|
||||||
|
return u, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func ollamaModelMatches(candidate, want string) bool {
|
||||||
|
candidate = strings.TrimSpace(candidate)
|
||||||
|
want = strings.TrimSpace(want)
|
||||||
|
if candidate == "" || want == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if strings.EqualFold(candidate, want) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
base, _, _ := strings.Cut(candidate, ":")
|
||||||
|
return strings.EqualFold(base, want)
|
||||||
|
}
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
@ -45,13 +46,24 @@ type modelResponse struct {
|
||||||
//
|
//
|
||||||
// GET /api/models
|
// GET /api/models
|
||||||
func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
||||||
cfg, err := h.loadFilteredConfig()
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
defaultModel := cfg.Agents.Defaults.GetModelName()
|
defaultModel := cfg.Agents.Defaults.GetModelName()
|
||||||
|
configured := make([]bool, len(cfg.ModelList))
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(len(cfg.ModelList))
|
||||||
|
for i, m := range cfg.ModelList {
|
||||||
|
go func(i int, m config.ModelConfig) {
|
||||||
|
defer wg.Done()
|
||||||
|
configured[i] = isModelConfigured(m)
|
||||||
|
}(i, m)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
models := make([]modelResponse, 0, len(cfg.ModelList))
|
models := make([]modelResponse, 0, len(cfg.ModelList))
|
||||||
for i, m := range cfg.ModelList {
|
for i, m := range cfg.ModelList {
|
||||||
|
|
@ -69,7 +81,7 @@ func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
||||||
MaxTokensField: m.MaxTokensField,
|
MaxTokensField: m.MaxTokensField,
|
||||||
RequestTimeout: m.RequestTimeout,
|
RequestTimeout: m.RequestTimeout,
|
||||||
ThinkingLevel: m.ThinkingLevel,
|
ThinkingLevel: m.ThinkingLevel,
|
||||||
Configured: m.APIKey != "" || m.AuthMethod != "",
|
Configured: configured[i],
|
||||||
IsDefault: m.ModelName == defaultModel,
|
IsDefault: m.ModelName == defaultModel,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
313
web/backend/api/models_test.go
Normal file
313
web/backend/api/models_test.go
Normal file
|
|
@ -0,0 +1,313 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func resetModelProbeHooks(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
origTCPProbe := probeTCPServiceFunc
|
||||||
|
origOllamaProbe := probeOllamaModelFunc
|
||||||
|
origOpenAIProbe := probeOpenAICompatibleModelFunc
|
||||||
|
t.Cleanup(func() {
|
||||||
|
probeTCPServiceFunc = origTCPProbe
|
||||||
|
probeOllamaModelFunc = origOllamaProbe
|
||||||
|
probeOpenAICompatibleModelFunc = origOpenAIProbe
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleListModels_ConfiguredStatusUsesRuntimeProbesForLocalModels(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
var mu sync.Mutex
|
||||||
|
var openAIProbes []string
|
||||||
|
var ollamaProbes []string
|
||||||
|
var tcpProbes []string
|
||||||
|
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
mu.Lock()
|
||||||
|
openAIProbes = append(openAIProbes, apiBase+"|"+modelID)
|
||||||
|
mu.Unlock()
|
||||||
|
return apiBase == "http://127.0.0.1:8000/v1" && modelID == "custom-model"
|
||||||
|
}
|
||||||
|
probeOllamaModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
mu.Lock()
|
||||||
|
ollamaProbes = append(ollamaProbes, apiBase+"|"+modelID)
|
||||||
|
mu.Unlock()
|
||||||
|
return apiBase == "http://localhost:11434/v1" && modelID == "llama3"
|
||||||
|
}
|
||||||
|
probeTCPServiceFunc = func(apiBase string) bool {
|
||||||
|
mu.Lock()
|
||||||
|
tcpProbes = append(tcpProbes, apiBase)
|
||||||
|
mu.Unlock()
|
||||||
|
return apiBase == "http://127.0.0.1:4321"
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "openai-oauth",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "vllm-local",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "http://127.0.0.1:8000/v1",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "ollama-default",
|
||||||
|
Model: "ollama/llama3",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "vllm-remote",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "https://models.example.com/v1",
|
||||||
|
APIKey: "remote-key",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "copilot-gpt-5.2",
|
||||||
|
Model: "github-copilot/gpt-5.2",
|
||||||
|
APIBase: "http://127.0.0.1:4321",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.ModelName = "openai-oauth"
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/models", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Models []modelResponse `json:"models"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := make(map[string]bool, len(resp.Models))
|
||||||
|
for _, model := range resp.Models {
|
||||||
|
got[model.ModelName] = model.Configured
|
||||||
|
}
|
||||||
|
|
||||||
|
if got["openai-oauth"] {
|
||||||
|
t.Fatalf("openai oauth model configured = true, want false without stored credential")
|
||||||
|
}
|
||||||
|
if !got["vllm-local"] {
|
||||||
|
t.Fatalf("vllm local model configured = false, want true when local probe succeeds")
|
||||||
|
}
|
||||||
|
if !got["ollama-default"] {
|
||||||
|
t.Fatalf("ollama default model configured = false, want true when default local probe succeeds")
|
||||||
|
}
|
||||||
|
if !got["vllm-remote"] {
|
||||||
|
t.Fatalf("remote vllm model configured = false, want true with api_key")
|
||||||
|
}
|
||||||
|
if !got["copilot-gpt-5.2"] {
|
||||||
|
t.Fatalf("copilot model configured = false, want true when local bridge probe succeeds")
|
||||||
|
}
|
||||||
|
if len(openAIProbes) != 1 || openAIProbes[0] != "http://127.0.0.1:8000/v1|custom-model" {
|
||||||
|
t.Fatalf("openAI probes = %#v, want only local vllm probe", openAIProbes)
|
||||||
|
}
|
||||||
|
if len(ollamaProbes) != 1 || ollamaProbes[0] != "http://localhost:11434/v1|llama3" {
|
||||||
|
t.Fatalf("ollama probes = %#v, want default local probe", ollamaProbes)
|
||||||
|
}
|
||||||
|
if len(tcpProbes) != 1 || tcpProbes[0] != "http://127.0.0.1:4321" {
|
||||||
|
t.Fatalf("tcp probes = %#v, want only local copilot probe", tcpProbes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleListModels_ConfiguredStatusForOAuthModelWithCredential(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "claude-oauth",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "claude-oauth"
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential(oauthProviderAnthropic, &auth.AuthCredential{
|
||||||
|
AccessToken: "anthropic-token",
|
||||||
|
Provider: oauthProviderAnthropic,
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("SetCredential() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/models", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Models []modelResponse `json:"models"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(resp.Models) != 1 {
|
||||||
|
t.Fatalf("len(models) = %d, want 1", len(resp.Models))
|
||||||
|
}
|
||||||
|
if !resp.Models[0].Configured {
|
||||||
|
t.Fatalf("oauth model configured = false, want true with stored credential")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleListModels_ProbesLocalModelsConcurrently(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
started := make(chan string, 2)
|
||||||
|
release := make(chan struct{})
|
||||||
|
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
started <- apiBase + "|" + modelID
|
||||||
|
<-release
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "local-vllm-a",
|
||||||
|
Model: "vllm/custom-a",
|
||||||
|
APIBase: "http://127.0.0.1:8000/v1",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "local-vllm-b",
|
||||||
|
Model: "vllm/custom-b",
|
||||||
|
APIBase: "http://127.0.0.1:8001/v1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
recCh := make(chan *httptest.ResponseRecorder, 1)
|
||||||
|
go func() {
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/models", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
recCh <- rec
|
||||||
|
}()
|
||||||
|
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
select {
|
||||||
|
case <-started:
|
||||||
|
case <-time.After(200 * time.Millisecond):
|
||||||
|
t.Fatal("expected both local probes to start before the first one completed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
close(release)
|
||||||
|
|
||||||
|
rec := <-recCh
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleListModels_NormalizesWildcardLocalAPIBaseForProbe(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
resetModelProbeHooks(t)
|
||||||
|
|
||||||
|
var gotProbe string
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID string) bool {
|
||||||
|
gotProbe = apiBase + "|" + modelID
|
||||||
|
return apiBase == "http://127.0.0.1:8000/v1" && modelID == "custom-model"
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "vllm-local",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "http://0.0.0.0:8000/v1",
|
||||||
|
}}
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/models", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Models []modelResponse `json:"models"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(resp.Models) != 1 {
|
||||||
|
t.Fatalf("len(models) = %d, want 1", len(resp.Models))
|
||||||
|
}
|
||||||
|
if !resp.Models[0].Configured {
|
||||||
|
t.Fatal("wildcard-bound local model configured = false, want true after probe host normalization")
|
||||||
|
}
|
||||||
|
if gotProbe != "http://127.0.0.1:8000/v1|custom-model" {
|
||||||
|
t.Fatalf("probe api base = %q, want %q", gotProbe, "http://127.0.0.1:8000/v1|custom-model")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -5,9 +5,7 @@ import (
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
|
@ -30,7 +28,7 @@ func (h *Handler) handleGetPicoToken(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
wsURL := buildWsURL(r, cfg)
|
wsURL := h.buildWsURL(r, cfg)
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
|
@ -58,7 +56,7 @@ func (h *Handler) handleRegenPicoToken(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
wsURL := fmt.Sprintf("ws://%s/pico/ws", net.JoinHostPort(cfg.Gateway.Host, strconv.Itoa(cfg.Gateway.Port)))
|
wsURL := h.buildWsURL(r, cfg)
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
|
@ -123,7 +121,7 @@ func (h *Handler) handlePicoSetup(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
wsURL := buildWsURL(r, cfg)
|
wsURL := h.buildWsURL(r, cfg)
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
|
@ -134,22 +132,6 @@ func (h *Handler) handlePicoSetup(w http.ResponseWriter, r *http.Request) {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildWsURL creates a WebSocket URL for the Pico Channel.
|
|
||||||
// When the gateway host is "0.0.0.0" or empty, it uses the hostname from the
|
|
||||||
// incoming HTTP request so the browser gets a connectable address.
|
|
||||||
func buildWsURL(r *http.Request, cfg *config.Config) string {
|
|
||||||
host := cfg.Gateway.Host
|
|
||||||
if host == "" || host == "0.0.0.0" {
|
|
||||||
// Use the hostname the browser used to reach this backend
|
|
||||||
reqHost, _, err := net.SplitHostPort(r.Host)
|
|
||||||
if err != nil {
|
|
||||||
reqHost = r.Host // r.Host might not have a port
|
|
||||||
}
|
|
||||||
host = reqHost
|
|
||||||
}
|
|
||||||
return "ws://" + net.JoinHostPort(host, strconv.Itoa(cfg.Gateway.Port)) + "/pico/ws"
|
|
||||||
}
|
|
||||||
|
|
||||||
// generateSecureToken creates a random 32-character hex string.
|
// generateSecureToken creates a random 32-character hex string.
|
||||||
func generateSecureToken() string {
|
func generateSecureToken() string {
|
||||||
b := make([]byte, 16)
|
b := make([]byte, 16)
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ type Handler struct {
|
||||||
configPath string
|
configPath string
|
||||||
serverPort int
|
serverPort int
|
||||||
serverPublic bool
|
serverPublic bool
|
||||||
|
serverPublicExplicit bool
|
||||||
serverCIDRs []string
|
serverCIDRs []string
|
||||||
oauthMu sync.Mutex
|
oauthMu sync.Mutex
|
||||||
oauthFlows map[string]*oauthFlow
|
oauthFlows map[string]*oauthFlow
|
||||||
|
|
@ -29,9 +30,10 @@ func NewHandler(configPath string) *Handler {
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetServerOptions stores current backend listen options for fallback behavior.
|
// SetServerOptions stores current backend listen options for fallback behavior.
|
||||||
func (h *Handler) SetServerOptions(port int, public bool, allowedCIDRs []string) {
|
func (h *Handler) SetServerOptions(port int, public bool, publicExplicit bool, allowedCIDRs []string) {
|
||||||
h.serverPort = port
|
h.serverPort = port
|
||||||
h.serverPublic = public
|
h.serverPublic = public
|
||||||
|
h.serverPublicExplicit = publicExplicit
|
||||||
h.serverCIDRs = append([]string(nil), allowedCIDRs...)
|
h.serverCIDRs = append([]string(nil), allowedCIDRs...)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -58,6 +60,10 @@ func (h *Handler) RegisterRoutes(mux *http.ServeMux) {
|
||||||
// Channel catalog (for frontend navigation/config pages)
|
// Channel catalog (for frontend navigation/config pages)
|
||||||
h.registerChannelRoutes(mux)
|
h.registerChannelRoutes(mux)
|
||||||
|
|
||||||
|
// Skills and tools support/actions
|
||||||
|
h.registerSkillRoutes(mux)
|
||||||
|
h.registerToolRoutes(mux)
|
||||||
|
|
||||||
// OS startup / launch-at-login
|
// OS startup / launch-at-login
|
||||||
h.registerStartupRoutes(mux)
|
h.registerStartupRoutes(mux)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,9 @@
|
||||||
package api
|
package api
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
|
@ -33,12 +35,22 @@ type sessionFile struct {
|
||||||
// sessionListItem is a lightweight summary returned by GET /api/sessions.
|
// sessionListItem is a lightweight summary returned by GET /api/sessions.
|
||||||
type sessionListItem struct {
|
type sessionListItem struct {
|
||||||
ID string `json:"id"`
|
ID string `json:"id"`
|
||||||
|
Title string `json:"title"`
|
||||||
Preview string `json:"preview"`
|
Preview string `json:"preview"`
|
||||||
MessageCount int `json:"message_count"`
|
MessageCount int `json:"message_count"`
|
||||||
Created string `json:"created"`
|
Created string `json:"created"`
|
||||||
Updated string `json:"updated"`
|
Updated string `json:"updated"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type sessionMetaFile struct {
|
||||||
|
Key string `json:"key"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
Skip int `json:"skip"`
|
||||||
|
Count int `json:"count"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
|
}
|
||||||
|
|
||||||
// picoSessionPrefix is the key prefix used by the gateway's routing for Pico
|
// picoSessionPrefix is the key prefix used by the gateway's routing for Pico
|
||||||
// channel sessions. The full key format is:
|
// channel sessions. The full key format is:
|
||||||
//
|
//
|
||||||
|
|
@ -47,7 +59,12 @@ type sessionListItem struct {
|
||||||
// The sanitized filename replaces ':' with '_', so on disk it becomes:
|
// The sanitized filename replaces ':' with '_', so on disk it becomes:
|
||||||
//
|
//
|
||||||
// agent_main_pico_direct_pico_<session-uuid>.json
|
// agent_main_pico_direct_pico_<session-uuid>.json
|
||||||
const picoSessionPrefix = "agent:main:pico:direct:pico:"
|
const (
|
||||||
|
picoSessionPrefix = "agent:main:pico:direct:pico:"
|
||||||
|
sanitizedPicoSessionPrefix = "agent_main_pico_direct_pico_"
|
||||||
|
maxSessionJSONLLineSize = 10 * 1024 * 1024 // 10 MB
|
||||||
|
maxSessionTitleRunes = 60
|
||||||
|
)
|
||||||
|
|
||||||
// extractPicoSessionID extracts the session UUID from a full session key.
|
// extractPicoSessionID extracts the session UUID from a full session key.
|
||||||
// Returns the UUID and true if the key matches the Pico session pattern.
|
// Returns the UUID and true if the key matches the Pico session pattern.
|
||||||
|
|
@ -58,6 +75,178 @@ func extractPicoSessionID(key string) (string, bool) {
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func extractPicoSessionIDFromSanitizedKey(key string) (string, bool) {
|
||||||
|
if strings.HasPrefix(key, sanitizedPicoSessionPrefix) {
|
||||||
|
return strings.TrimPrefix(key, sanitizedPicoSessionPrefix), true
|
||||||
|
}
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeSessionKey(key string) string {
|
||||||
|
return strings.ReplaceAll(key, ":", "_")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) readLegacySession(dir, sessionID string) (sessionFile, error) {
|
||||||
|
path := filepath.Join(dir, sanitizeSessionKey(picoSessionPrefix+sessionID)+".json")
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return sessionFile{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var sess sessionFile
|
||||||
|
if err := json.Unmarshal(data, &sess); err != nil {
|
||||||
|
return sessionFile{}, err
|
||||||
|
}
|
||||||
|
return sess, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) readSessionMeta(path, sessionKey string) (sessionMetaFile, error) {
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return sessionMetaFile{Key: sessionKey}, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return sessionMetaFile{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var meta sessionMetaFile
|
||||||
|
if err := json.Unmarshal(data, &meta); err != nil {
|
||||||
|
return sessionMetaFile{}, err
|
||||||
|
}
|
||||||
|
if meta.Key == "" {
|
||||||
|
meta.Key = sessionKey
|
||||||
|
}
|
||||||
|
return meta, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) readSessionMessages(path string, skip int) ([]providers.Message, error) {
|
||||||
|
f, err := os.Open(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer f.Close()
|
||||||
|
|
||||||
|
msgs := make([]providers.Message, 0)
|
||||||
|
scanner := bufio.NewScanner(f)
|
||||||
|
scanner.Buffer(make([]byte, 0, 64*1024), maxSessionJSONLLineSize)
|
||||||
|
|
||||||
|
seen := 0
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Bytes()
|
||||||
|
if len(line) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
seen++
|
||||||
|
if seen <= skip {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var msg providers.Message
|
||||||
|
if err := json.Unmarshal(line, &msg); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
msgs = append(msgs, msg)
|
||||||
|
}
|
||||||
|
if err := scanner.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return msgs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) readJSONLSession(dir, sessionID string) (sessionFile, error) {
|
||||||
|
sessionKey := picoSessionPrefix + sessionID
|
||||||
|
base := filepath.Join(dir, sanitizeSessionKey(sessionKey))
|
||||||
|
jsonlPath := base + ".jsonl"
|
||||||
|
metaPath := base + ".meta.json"
|
||||||
|
|
||||||
|
meta, err := h.readSessionMeta(metaPath, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return sessionFile{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
messages, err := h.readSessionMessages(jsonlPath, meta.Skip)
|
||||||
|
if err != nil {
|
||||||
|
return sessionFile{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
updated := meta.UpdatedAt
|
||||||
|
created := meta.CreatedAt
|
||||||
|
if created.IsZero() || updated.IsZero() {
|
||||||
|
if info, statErr := os.Stat(jsonlPath); statErr == nil {
|
||||||
|
if created.IsZero() {
|
||||||
|
created = info.ModTime()
|
||||||
|
}
|
||||||
|
if updated.IsZero() {
|
||||||
|
updated = info.ModTime()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sessionFile{
|
||||||
|
Key: meta.Key,
|
||||||
|
Messages: messages,
|
||||||
|
Summary: meta.Summary,
|
||||||
|
Created: created,
|
||||||
|
Updated: updated,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildSessionListItem(sessionID string, sess sessionFile) sessionListItem {
|
||||||
|
preview := ""
|
||||||
|
for _, msg := range sess.Messages {
|
||||||
|
if msg.Role == "user" && strings.TrimSpace(msg.Content) != "" {
|
||||||
|
preview = msg.Content
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
title := strings.TrimSpace(sess.Summary)
|
||||||
|
if title == "" {
|
||||||
|
title = preview
|
||||||
|
}
|
||||||
|
|
||||||
|
title = truncateRunes(title, maxSessionTitleRunes)
|
||||||
|
preview = truncateRunes(preview, maxSessionTitleRunes)
|
||||||
|
|
||||||
|
if preview == "" {
|
||||||
|
preview = "(empty)"
|
||||||
|
}
|
||||||
|
if title == "" {
|
||||||
|
title = preview
|
||||||
|
}
|
||||||
|
|
||||||
|
validMessageCount := 0
|
||||||
|
for _, msg := range sess.Messages {
|
||||||
|
if (msg.Role == "user" || msg.Role == "assistant") && strings.TrimSpace(msg.Content) != "" {
|
||||||
|
validMessageCount++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sessionListItem{
|
||||||
|
ID: sessionID,
|
||||||
|
Title: title,
|
||||||
|
Preview: preview,
|
||||||
|
MessageCount: validMessageCount,
|
||||||
|
Created: sess.Created.Format(time.RFC3339),
|
||||||
|
Updated: sess.Updated.Format(time.RFC3339),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isEmptySession(sess sessionFile) bool {
|
||||||
|
return len(sess.Messages) == 0 && strings.TrimSpace(sess.Summary) == ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func truncateRunes(s string, maxLen int) string {
|
||||||
|
if maxLen <= 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
runes := []rune(strings.TrimSpace(s))
|
||||||
|
if len(runes) <= maxLen {
|
||||||
|
return string(runes)
|
||||||
|
}
|
||||||
|
return string(runes[:maxLen]) + "..."
|
||||||
|
}
|
||||||
|
|
||||||
// sessionsDir resolves the path to the gateway's session storage directory.
|
// sessionsDir resolves the path to the gateway's session storage directory.
|
||||||
// It reads the workspace from config, falling back to ~/.picoclaw/workspace.
|
// It reads the workspace from config, falling back to ~/.picoclaw/workspace.
|
||||||
func (h *Handler) sessionsDir() (string, error) {
|
func (h *Handler) sessionsDir() (string, error) {
|
||||||
|
|
@ -104,58 +293,76 @@ func (h *Handler) handleListSessions(w http.ResponseWriter, r *http.Request) {
|
||||||
}
|
}
|
||||||
|
|
||||||
items := []sessionListItem{}
|
items := []sessionListItem{}
|
||||||
|
seen := make(map[string]struct{})
|
||||||
|
|
||||||
for _, entry := range entries {
|
for _, entry := range entries {
|
||||||
if entry.IsDir() || filepath.Ext(entry.Name()) != ".json" {
|
if entry.IsDir() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
data, err := os.ReadFile(filepath.Join(dir, entry.Name()))
|
name := entry.Name()
|
||||||
if err != nil {
|
var (
|
||||||
continue
|
sessionID string
|
||||||
}
|
sess sessionFile
|
||||||
|
loadErr error
|
||||||
|
ok bool
|
||||||
|
)
|
||||||
|
|
||||||
var sess sessionFile
|
switch {
|
||||||
if err := json.Unmarshal(data, &sess); err != nil {
|
case strings.HasSuffix(name, ".jsonl"):
|
||||||
continue
|
sessionID, ok = extractPicoSessionIDFromSanitizedKey(strings.TrimSuffix(name, ".jsonl"))
|
||||||
}
|
|
||||||
|
|
||||||
// Only include Pico channel sessions
|
|
||||||
sessionID, ok := extractPicoSessionID(sess.Key)
|
|
||||||
if !ok {
|
if !ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
sess, loadErr = h.readJSONLSession(dir, sessionID)
|
||||||
// Build a preview from the first user message
|
if loadErr == nil && isEmptySession(sess) {
|
||||||
preview := ""
|
continue
|
||||||
for _, msg := range sess.Messages {
|
}
|
||||||
if msg.Role == "user" && strings.TrimSpace(msg.Content) != "" {
|
case strings.HasSuffix(name, ".meta.json"):
|
||||||
preview = msg.Content
|
continue
|
||||||
break
|
case filepath.Ext(name) == ".json":
|
||||||
|
base := strings.TrimSuffix(name, ".json")
|
||||||
|
if _, statErr := os.Stat(filepath.Join(dir, base+".jsonl")); statErr == nil {
|
||||||
|
if jsonlSessionID, found := extractPicoSessionIDFromSanitizedKey(base); found {
|
||||||
|
if jsonlSess, jsonlErr := h.readJSONLSession(
|
||||||
|
dir,
|
||||||
|
jsonlSessionID,
|
||||||
|
); jsonlErr == nil &&
|
||||||
|
!isEmptySession(jsonlSess) {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if len([]rune(preview)) > 60 {
|
|
||||||
preview = string([]rune(preview)[:60]) + "..."
|
|
||||||
}
|
}
|
||||||
if preview == "" {
|
data, err := os.ReadFile(filepath.Join(dir, name))
|
||||||
preview = "(empty)"
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(data, &sess); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if isEmptySession(sess) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
sessionID, ok = extractPicoSessionID(sess.Key)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, exists := seen[sessionID]; exists {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Only count non-empty user and assistant messages
|
if loadErr != nil {
|
||||||
validMessageCount := 0
|
continue
|
||||||
for _, msg := range sess.Messages {
|
|
||||||
if (msg.Role == "user" || msg.Role == "assistant") && strings.TrimSpace(msg.Content) != "" {
|
|
||||||
validMessageCount++
|
|
||||||
}
|
}
|
||||||
|
if _, exists := seen[sessionID]; exists {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
items = append(items, sessionListItem{
|
seen[sessionID] = struct{}{}
|
||||||
ID: sessionID,
|
items = append(items, buildSessionListItem(sessionID, sess))
|
||||||
Preview: preview,
|
|
||||||
MessageCount: validMessageCount,
|
|
||||||
Created: sess.Created.Format(time.RFC3339),
|
|
||||||
Updated: sess.Updated.Format(time.RFC3339),
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sort by updated descending (most recent first)
|
// Sort by updated descending (most recent first)
|
||||||
|
|
@ -209,20 +416,25 @@ func (h *Handler) handleGetSession(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// The sanitized filename replaces ':' with '_':
|
sess, err := h.readJSONLSession(dir, sessionID)
|
||||||
// agent:main:pico:direct:pico:<uuid> -> agent_main_pico_direct_pico_<uuid>.json
|
if err == nil && isEmptySession(sess) {
|
||||||
filename := strings.ReplaceAll(picoSessionPrefix+sessionID, ":", "_") + ".json"
|
err = os.ErrNotExist
|
||||||
|
}
|
||||||
data, err := os.ReadFile(filepath.Join(dir, filename))
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if errors.Is(err, os.ErrNotExist) {
|
||||||
|
sess, err = h.readLegacySession(dir, sessionID)
|
||||||
|
if err == nil && isEmptySession(sess) {
|
||||||
|
err = os.ErrNotExist
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, os.ErrNotExist) {
|
||||||
http.Error(w, "session not found", http.StatusNotFound)
|
http.Error(w, "session not found", http.StatusNotFound)
|
||||||
|
} else {
|
||||||
|
http.Error(w, "failed to parse session", http.StatusInternalServerError)
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var sess sessionFile
|
|
||||||
if err := json.Unmarshal(data, &sess); err != nil {
|
|
||||||
http.Error(w, "failed to parse session", http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Convert to a simpler format for the frontend
|
// Convert to a simpler format for the frontend
|
||||||
|
|
@ -268,17 +480,25 @@ func (h *Handler) handleDeleteSession(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// The sanitized filename replaces ':' with '_':
|
base := filepath.Join(dir, sanitizeSessionKey(picoSessionPrefix+sessionID))
|
||||||
// agent:main:pico:direct:pico:<uuid> -> agent_main_pico_direct_pico_<uuid>.json
|
jsonlPath := base + ".jsonl"
|
||||||
filename := strings.ReplaceAll(picoSessionPrefix+sessionID, ":", "_") + ".json"
|
metaPath := base + ".meta.json"
|
||||||
filePath := filepath.Join(dir, filename)
|
legacyPath := base + ".json"
|
||||||
|
|
||||||
if err := os.Remove(filePath); err != nil {
|
removed := false
|
||||||
|
for _, path := range []string{jsonlPath, metaPath, legacyPath} {
|
||||||
|
if err := os.Remove(path); err != nil {
|
||||||
if os.IsNotExist(err) {
|
if os.IsNotExist(err) {
|
||||||
http.Error(w, "session not found", http.StatusNotFound)
|
continue
|
||||||
} else {
|
|
||||||
http.Error(w, "failed to delete session", http.StatusInternalServerError)
|
|
||||||
}
|
}
|
||||||
|
http.Error(w, "failed to delete session", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
removed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if !removed {
|
||||||
|
http.Error(w, "session not found", http.StatusNotFound)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
322
web/backend/api/session_test.go
Normal file
322
web/backend/api/session_test.go
Normal file
|
|
@ -0,0 +1,322 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/memory"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
func sessionsTestDir(t *testing.T, configPath string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
dir := filepath.Join(cfg.Agents.Defaults.Workspace, "sessions")
|
||||||
|
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll() error = %v", err)
|
||||||
|
}
|
||||||
|
return dir
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleListSessions_JSONLStorage(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
dir := sessionsTestDir(t, configPath)
|
||||||
|
store, err := memory.NewJSONLStore(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewJSONLStore() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionKey := picoSessionPrefix + "history-jsonl"
|
||||||
|
if err := store.AddFullMessage(nil, sessionKey, providers.Message{
|
||||||
|
Role: "user",
|
||||||
|
Content: "Explain why the history API is empty after migration.",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("AddFullMessage(user) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := store.AddFullMessage(nil, sessionKey, providers.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: "Because the API still reads only legacy JSON session files.",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("AddFullMessage(assistant) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := store.AddFullMessage(nil, sessionKey, providers.Message{
|
||||||
|
Role: "tool",
|
||||||
|
Content: "ignored",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("AddFullMessage(tool) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := store.SetSummary(nil, sessionKey, "JSONL-backed session"); err != nil {
|
||||||
|
t.Fatalf("SetSummary() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/sessions", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var items []sessionListItem
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &items); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(items) != 1 {
|
||||||
|
t.Fatalf("len(items) = %d, want 1", len(items))
|
||||||
|
}
|
||||||
|
if items[0].ID != "history-jsonl" {
|
||||||
|
t.Fatalf("items[0].ID = %q, want %q", items[0].ID, "history-jsonl")
|
||||||
|
}
|
||||||
|
if items[0].MessageCount != 2 {
|
||||||
|
t.Fatalf("items[0].MessageCount = %d, want 2", items[0].MessageCount)
|
||||||
|
}
|
||||||
|
if items[0].Title != "JSONL-backed session" {
|
||||||
|
t.Fatalf("items[0].Title = %q, want %q", items[0].Title, "JSONL-backed session")
|
||||||
|
}
|
||||||
|
if items[0].Preview != "Explain why the history API is empty after migration." {
|
||||||
|
t.Fatalf("items[0].Preview = %q", items[0].Preview)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleListSessions_TitleUsesTrimmedSummary(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
dir := sessionsTestDir(t, configPath)
|
||||||
|
store, err := memory.NewJSONLStore(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewJSONLStore() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionKey := picoSessionPrefix + "summary-title"
|
||||||
|
if err := store.AddFullMessage(nil, sessionKey, providers.Message{
|
||||||
|
Role: "user",
|
||||||
|
Content: "fallback preview",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("AddFullMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := store.SetSummary(
|
||||||
|
nil,
|
||||||
|
sessionKey,
|
||||||
|
" This summary is intentionally longer than sixty characters so it must be truncated in the history menu. ",
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("SetSummary() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/sessions", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var items []sessionListItem
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &items); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(items) != 1 {
|
||||||
|
t.Fatalf("len(items) = %d, want 1", len(items))
|
||||||
|
}
|
||||||
|
expectedTitle := truncateRunes(
|
||||||
|
"This summary is intentionally longer than sixty characters so it must be truncated in the history menu.",
|
||||||
|
maxSessionTitleRunes,
|
||||||
|
)
|
||||||
|
if items[0].Title != expectedTitle {
|
||||||
|
t.Fatalf("items[0].Title = %q", items[0].Title)
|
||||||
|
}
|
||||||
|
if items[0].Preview != "fallback preview" {
|
||||||
|
t.Fatalf("items[0].Preview = %q, want %q", items[0].Preview, "fallback preview")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleGetSession_JSONLStorage(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
dir := sessionsTestDir(t, configPath)
|
||||||
|
store, err := memory.NewJSONLStore(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewJSONLStore() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionKey := picoSessionPrefix + "detail-jsonl"
|
||||||
|
for _, msg := range []providers.Message{
|
||||||
|
{Role: "user", Content: "first"},
|
||||||
|
{Role: "assistant", Content: "second"},
|
||||||
|
{Role: "tool", Content: "ignored"},
|
||||||
|
} {
|
||||||
|
if err := store.AddFullMessage(nil, sessionKey, msg); err != nil {
|
||||||
|
t.Fatalf("AddFullMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := store.SetSummary(nil, sessionKey, "detail summary"); err != nil {
|
||||||
|
t.Fatalf("SetSummary() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/sessions/detail-jsonl", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
Messages []struct {
|
||||||
|
Role string `json:"role"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"messages"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if resp.ID != "detail-jsonl" {
|
||||||
|
t.Fatalf("resp.ID = %q, want %q", resp.ID, "detail-jsonl")
|
||||||
|
}
|
||||||
|
if resp.Summary != "detail summary" {
|
||||||
|
t.Fatalf("resp.Summary = %q, want %q", resp.Summary, "detail summary")
|
||||||
|
}
|
||||||
|
if len(resp.Messages) != 2 {
|
||||||
|
t.Fatalf("len(resp.Messages) = %d, want 2", len(resp.Messages))
|
||||||
|
}
|
||||||
|
if resp.Messages[0].Role != "user" || resp.Messages[0].Content != "first" {
|
||||||
|
t.Fatalf("first message = %#v, want user/first", resp.Messages[0])
|
||||||
|
}
|
||||||
|
if resp.Messages[1].Role != "assistant" || resp.Messages[1].Content != "second" {
|
||||||
|
t.Fatalf("second message = %#v, want assistant/second", resp.Messages[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleDeleteSession_JSONLStorage(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
dir := sessionsTestDir(t, configPath)
|
||||||
|
store, err := memory.NewJSONLStore(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewJSONLStore() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionKey := picoSessionPrefix + "delete-jsonl"
|
||||||
|
if err := store.AddFullMessage(nil, sessionKey, providers.Message{
|
||||||
|
Role: "user",
|
||||||
|
Content: "delete me",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("AddFullMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := store.SetSummary(nil, sessionKey, "delete summary"); err != nil {
|
||||||
|
t.Fatalf("SetSummary() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodDelete, "/api/sessions/delete-jsonl", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusNoContent {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusNoContent, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
base := filepath.Join(dir, sanitizeSessionKey(sessionKey))
|
||||||
|
for _, path := range []string{base + ".jsonl", base + ".meta.json"} {
|
||||||
|
if _, err := os.Stat(path); !os.IsNotExist(err) {
|
||||||
|
t.Fatalf("expected %s to be removed, stat err = %v", path, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleGetSession_LegacyJSONFallback(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
dir := sessionsTestDir(t, configPath)
|
||||||
|
manager := session.NewSessionManager(dir)
|
||||||
|
sessionKey := picoSessionPrefix + "legacy-json"
|
||||||
|
manager.AddMessage(sessionKey, "user", "legacy user")
|
||||||
|
manager.AddMessage(sessionKey, "assistant", "legacy assistant")
|
||||||
|
if err := manager.Save(sessionKey); err != nil {
|
||||||
|
t.Fatalf("Save() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/sessions/legacy-json", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleSessions_FiltersEmptyJSONLFiles(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
dir := sessionsTestDir(t, configPath)
|
||||||
|
base := filepath.Join(dir, sanitizeSessionKey(picoSessionPrefix+"empty-jsonl"))
|
||||||
|
if err := os.WriteFile(base+".jsonl", []byte{}, 0o644); err != nil {
|
||||||
|
t.Fatalf("WriteFile(jsonl) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
listRec := httptest.NewRecorder()
|
||||||
|
listReq := httptest.NewRequest(http.MethodGet, "/api/sessions", nil)
|
||||||
|
mux.ServeHTTP(listRec, listReq)
|
||||||
|
|
||||||
|
if listRec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("list status = %d, want %d, body=%s", listRec.Code, http.StatusOK, listRec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var items []sessionListItem
|
||||||
|
if err := json.Unmarshal(listRec.Body.Bytes(), &items); err != nil {
|
||||||
|
t.Fatalf("Unmarshal(list) error = %v", err)
|
||||||
|
}
|
||||||
|
if len(items) != 0 {
|
||||||
|
t.Fatalf("len(items) = %d, want 0", len(items))
|
||||||
|
}
|
||||||
|
|
||||||
|
detailRec := httptest.NewRecorder()
|
||||||
|
detailReq := httptest.NewRequest(http.MethodGet, "/api/sessions/empty-jsonl", nil)
|
||||||
|
mux.ServeHTTP(detailRec, detailReq)
|
||||||
|
|
||||||
|
if detailRec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("detail status = %d, want %d, body=%s", detailRec.Code, http.StatusNotFound, detailRec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
331
web/backend/api/skills.go
Normal file
331
web/backend/api/skills.go
Normal file
|
|
@ -0,0 +1,331 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
)
|
||||||
|
|
||||||
|
type skillSupportResponse struct {
|
||||||
|
Skills []skills.SkillInfo `json:"skills"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type skillDetailResponse struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
Source string `json:"source"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
skillNameSanitizer = regexp.MustCompile(`[^a-z0-9-]+`)
|
||||||
|
importedSkillFrontmatter = regexp.MustCompile(`(?s)^---(?:\r\n|\n|\r)(.*?)(?:\r\n|\n|\r)---(?:\r\n|\n|\r)*`)
|
||||||
|
skillFrontmatterStripper = regexp.MustCompile(`(?s)^---(?:\r\n|\n|\r)(.*?)(?:\r\n|\n|\r)---(?:\r\n|\n|\r)*`)
|
||||||
|
)
|
||||||
|
|
||||||
|
func (h *Handler) registerSkillRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/skills", h.handleListSkills)
|
||||||
|
mux.HandleFunc("GET /api/skills/{name}", h.handleGetSkill)
|
||||||
|
mux.HandleFunc("POST /api/skills/import", h.handleImportSkill)
|
||||||
|
mux.HandleFunc("DELETE /api/skills/{name}", h.handleDeleteSkill)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleListSkills(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
loader := newSkillsLoader(cfg.WorkspacePath())
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(skillSupportResponse{
|
||||||
|
Skills: loader.ListSkills(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleGetSkill(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
loader := newSkillsLoader(cfg.WorkspacePath())
|
||||||
|
name := r.PathValue("name")
|
||||||
|
allSkills := loader.ListSkills()
|
||||||
|
|
||||||
|
for _, skill := range allSkills {
|
||||||
|
if skill.Name != name {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
content, err := loadSkillContent(skill.Path)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Skill content not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(skillDetailResponse{
|
||||||
|
Name: skill.Name,
|
||||||
|
Path: skill.Path,
|
||||||
|
Source: skill.Source,
|
||||||
|
Description: skill.Description,
|
||||||
|
Content: content,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
http.Error(w, "Skill not found", http.StatusNotFound)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleImportSkill(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err = r.ParseMultipartForm(2 << 20)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid multipart form: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
uploadedFile, fileHeader, err := r.FormFile("file")
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "file is required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer uploadedFile.Close()
|
||||||
|
|
||||||
|
content, err := io.ReadAll(io.LimitReader(uploadedFile, (1<<20)+1))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to read file: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if len(content) > 1<<20 {
|
||||||
|
http.Error(w, "file exceeds 1MB limit", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
skillName, err := normalizeImportedSkillName(fileHeader.Filename, content)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
content = normalizeImportedSkillContent(content, skillName)
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
skillDir := filepath.Join(workspace, "skills", skillName)
|
||||||
|
skillFile := filepath.Join(skillDir, "SKILL.md")
|
||||||
|
if _, err := os.Stat(skillDir); err == nil {
|
||||||
|
http.Error(w, "skill already exists", http.StatusConflict)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.MkdirAll(skillDir, 0o755); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to create skill directory: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(skillFile, content, 0o644); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save skill: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
loader := newSkillsLoader(workspace)
|
||||||
|
for _, skill := range loader.ListSkills() {
|
||||||
|
if skill.Path == skillFile || (skill.Name == skillName && skill.Source == "workspace") {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(skill)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{
|
||||||
|
"name": skillName,
|
||||||
|
"path": skillFile,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleDeleteSkill(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
loader := newSkillsLoader(cfg.WorkspacePath())
|
||||||
|
name := r.PathValue("name")
|
||||||
|
for _, skill := range loader.ListSkills() {
|
||||||
|
if skill.Name != name {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if skill.Source != "workspace" {
|
||||||
|
http.Error(w, "only workspace skills can be deleted", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := os.RemoveAll(filepath.Dir(skill.Path)); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to delete skill: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
http.Error(w, "Skill not found", http.StatusNotFound)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSkillsLoader(workspace string) *skills.SkillsLoader {
|
||||||
|
return skills.NewSkillsLoader(
|
||||||
|
workspace,
|
||||||
|
filepath.Join(globalConfigDir(), "skills"),
|
||||||
|
builtinSkillsDir(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeImportedSkillName(filename string, content []byte) (string, error) {
|
||||||
|
rawContent := strings.ReplaceAll(string(content), "\r\n", "\n")
|
||||||
|
rawContent = strings.ReplaceAll(rawContent, "\r", "\n")
|
||||||
|
metadata, _ := extractImportedSkillMetadata(rawContent)
|
||||||
|
|
||||||
|
raw := strings.TrimSpace(metadata["name"])
|
||||||
|
if raw == "" {
|
||||||
|
raw = strings.TrimSpace(strings.TrimSuffix(filepath.Base(filename), filepath.Ext(filename)))
|
||||||
|
}
|
||||||
|
raw = strings.ToLower(raw)
|
||||||
|
raw = strings.ReplaceAll(raw, "_", "-")
|
||||||
|
raw = strings.ReplaceAll(raw, " ", "-")
|
||||||
|
raw = skillNameSanitizer.ReplaceAllString(raw, "-")
|
||||||
|
raw = strings.Trim(raw, "-")
|
||||||
|
raw = strings.Join(strings.FieldsFunc(raw, func(r rune) bool { return r == '-' }), "-")
|
||||||
|
|
||||||
|
if raw == "" {
|
||||||
|
return "", fmt.Errorf("skill name is required in frontmatter or filename")
|
||||||
|
}
|
||||||
|
if len(raw) > 64 {
|
||||||
|
return "", fmt.Errorf("skill name exceeds 64 characters")
|
||||||
|
}
|
||||||
|
matched, err := regexp.MatchString(`^[a-z0-9]+(-[a-z0-9]+)*$`, raw)
|
||||||
|
if err != nil || !matched {
|
||||||
|
return "", fmt.Errorf("skill name must be alphanumeric with hyphens")
|
||||||
|
}
|
||||||
|
return raw, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeImportedSkillContent(content []byte, skillName string) []byte {
|
||||||
|
raw := strings.ReplaceAll(string(content), "\r\n", "\n")
|
||||||
|
raw = strings.ReplaceAll(raw, "\r", "\n")
|
||||||
|
|
||||||
|
metadata, body := extractImportedSkillMetadata(raw)
|
||||||
|
description := strings.TrimSpace(metadata["description"])
|
||||||
|
if description == "" {
|
||||||
|
description = inferImportedSkillDescription(body)
|
||||||
|
}
|
||||||
|
if description == "" {
|
||||||
|
description = "Imported skill"
|
||||||
|
}
|
||||||
|
if len(description) > 1024 {
|
||||||
|
description = strings.TrimSpace(description[:1024])
|
||||||
|
}
|
||||||
|
|
||||||
|
body = strings.TrimLeft(body, "\n")
|
||||||
|
var builder strings.Builder
|
||||||
|
builder.WriteString("---\n")
|
||||||
|
builder.WriteString("name: ")
|
||||||
|
builder.WriteString(skillName)
|
||||||
|
builder.WriteString("\n")
|
||||||
|
builder.WriteString("description: ")
|
||||||
|
builder.WriteString(description)
|
||||||
|
builder.WriteString("\n")
|
||||||
|
builder.WriteString("---\n\n")
|
||||||
|
builder.WriteString(body)
|
||||||
|
if !strings.HasSuffix(builder.String(), "\n") {
|
||||||
|
builder.WriteString("\n")
|
||||||
|
}
|
||||||
|
return []byte(builder.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractImportedSkillMetadata(raw string) (map[string]string, string) {
|
||||||
|
matches := importedSkillFrontmatter.FindStringSubmatch(raw)
|
||||||
|
if len(matches) != 2 {
|
||||||
|
return map[string]string{}, raw
|
||||||
|
}
|
||||||
|
meta := parseImportedSkillYAML(matches[1])
|
||||||
|
body := importedSkillFrontmatter.ReplaceAllString(raw, "")
|
||||||
|
return meta, body
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseImportedSkillYAML(frontmatter string) map[string]string {
|
||||||
|
result := make(map[string]string)
|
||||||
|
for _, line := range strings.Split(frontmatter, "\n") {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line == "" || strings.HasPrefix(line, "#") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
key, value, ok := strings.Cut(line, ":")
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
result[strings.TrimSpace(key)] = strings.Trim(strings.TrimSpace(value), `"'`)
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func inferImportedSkillDescription(body string) string {
|
||||||
|
for _, line := range strings.Split(body, "\n") {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
line = strings.TrimLeft(line, "#-*0123456789. ")
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line != "" {
|
||||||
|
return line
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadSkillContent(path string) (string, error) {
|
||||||
|
content, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return skillFrontmatterStripper.ReplaceAllString(string(content), ""), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func globalConfigDir() string {
|
||||||
|
if home := os.Getenv("PICOCLAW_HOME"); home != "" {
|
||||||
|
return home
|
||||||
|
}
|
||||||
|
home, err := os.UserHomeDir()
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return filepath.Join(home, ".picoclaw")
|
||||||
|
}
|
||||||
|
|
||||||
|
func builtinSkillsDir() string {
|
||||||
|
if path := os.Getenv("PICOCLAW_BUILTIN_SKILLS"); path != "" {
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
wd, err := os.Getwd()
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return filepath.Join(wd, "skills")
|
||||||
|
}
|
||||||
336
web/backend/api/skills_test.go
Normal file
336
web/backend/api/skills_test.go
Normal file
|
|
@ -0,0 +1,336 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"mime/multipart"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestHandleListSkills(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := filepath.Join(t.TempDir(), "workspace")
|
||||||
|
cfg.Agents.Defaults.Workspace = workspace
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.MkdirAll(filepath.Join(workspace, "skills", "workspace-skill"), 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll(workspace skill) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(workspace, "skills", "workspace-skill", "SKILL.md"),
|
||||||
|
[]byte("---\nname: workspace-skill\ndescription: Workspace skill\n---\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile(workspace skill) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
globalSkillDir := filepath.Join(globalConfigDir(), "skills", "global-skill")
|
||||||
|
if err := os.MkdirAll(globalSkillDir, 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll(global skill) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(globalSkillDir, "SKILL.md"),
|
||||||
|
[]byte("---\nname: global-skill\ndescription: Global skill\n---\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile(global skill) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
builtinRoot := filepath.Join(t.TempDir(), "builtin-skills")
|
||||||
|
oldBuiltin := os.Getenv("PICOCLAW_BUILTIN_SKILLS")
|
||||||
|
if err := os.Setenv("PICOCLAW_BUILTIN_SKILLS", builtinRoot); err != nil {
|
||||||
|
t.Fatalf("Setenv(PICOCLAW_BUILTIN_SKILLS) error = %v", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if oldBuiltin == "" {
|
||||||
|
_ = os.Unsetenv("PICOCLAW_BUILTIN_SKILLS")
|
||||||
|
} else {
|
||||||
|
_ = os.Setenv("PICOCLAW_BUILTIN_SKILLS", oldBuiltin)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
builtinSkillDir := filepath.Join(builtinRoot, "builtin-skill")
|
||||||
|
if err := os.MkdirAll(builtinSkillDir, 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll(builtin skill) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(builtinSkillDir, "SKILL.md"),
|
||||||
|
[]byte("---\nname: builtin-skill\ndescription: Builtin skill\n---\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile(builtin skill) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/skills", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp skillSupportResponse
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(resp.Skills) != 3 {
|
||||||
|
t.Fatalf("skills count = %d, want 3", len(resp.Skills))
|
||||||
|
}
|
||||||
|
|
||||||
|
gotSkills := make(map[string]string, len(resp.Skills))
|
||||||
|
for _, skill := range resp.Skills {
|
||||||
|
gotSkills[skill.Name] = skill.Source
|
||||||
|
}
|
||||||
|
if gotSkills["workspace-skill"] != "workspace" {
|
||||||
|
t.Fatalf("workspace-skill source = %q, want workspace", gotSkills["workspace-skill"])
|
||||||
|
}
|
||||||
|
if gotSkills["global-skill"] != "global" {
|
||||||
|
t.Fatalf("global-skill source = %q, want global", gotSkills["global-skill"])
|
||||||
|
}
|
||||||
|
if gotSkills["builtin-skill"] != "builtin" {
|
||||||
|
t.Fatalf("builtin-skill source = %q, want builtin", gotSkills["builtin-skill"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleGetSkill(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := filepath.Join(t.TempDir(), "workspace")
|
||||||
|
cfg.Agents.Defaults.Workspace = workspace
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
skillDir := filepath.Join(workspace, "skills", "viewer-skill")
|
||||||
|
if err := os.MkdirAll(skillDir, 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(skillDir, "SKILL.md"),
|
||||||
|
[]byte(
|
||||||
|
"---\nname: viewer-skill\ndescription: Viewable skill\n---\n# Viewer Skill\n\nThis is visible content.\n",
|
||||||
|
),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/skills/viewer-skill", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp skillDetailResponse
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if resp.Name != "viewer-skill" || resp.Source != "workspace" || resp.Description != "Viewable skill" {
|
||||||
|
t.Fatalf("unexpected response: %#v", resp)
|
||||||
|
}
|
||||||
|
if resp.Content != "# Viewer Skill\n\nThis is visible content.\n" {
|
||||||
|
t.Fatalf("content = %q", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleGetSkillUsesResolvedPath(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := filepath.Join(t.TempDir(), "workspace")
|
||||||
|
cfg.Agents.Defaults.Workspace = workspace
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
skillDir := filepath.Join(workspace, "skills", "folder-name")
|
||||||
|
if err := os.MkdirAll(skillDir, 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(skillDir, "SKILL.md"),
|
||||||
|
[]byte("---\nname: display-name\ndescription: Mismatched path skill\n---\n# Display Name\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/skills/display-name", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp skillDetailResponse
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
if resp.Name != "display-name" {
|
||||||
|
t.Fatalf("resp.Name = %q, want display-name", resp.Name)
|
||||||
|
}
|
||||||
|
if resp.Content != "# Display Name\n" {
|
||||||
|
t.Fatalf("content = %q", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleImportSkill(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
workspace := filepath.Join(t.TempDir(), "workspace")
|
||||||
|
cfg.Agents.Defaults.Workspace = workspace
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var body bytes.Buffer
|
||||||
|
writer := multipart.NewWriter(&body)
|
||||||
|
part, err := writer.CreateFormFile("file", "Plain Skill.md")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateFormFile() error = %v", err)
|
||||||
|
}
|
||||||
|
_, err = io.WriteString(part, "# Plain Skill\n\nUse this skill to test imports.\n")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteString() error = %v", err)
|
||||||
|
}
|
||||||
|
err = writer.Close()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Close() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/skills/import", &body)
|
||||||
|
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
skillFile := filepath.Join(workspace, "skills", "plain-skill", "SKILL.md")
|
||||||
|
content, err := os.ReadFile(skillFile)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ReadFile() error = %v", err)
|
||||||
|
}
|
||||||
|
expected := "---\nname: plain-skill\ndescription: Plain Skill\n---\n\n# Plain Skill\n\nUse this skill to test imports.\n"
|
||||||
|
if string(content) != expected {
|
||||||
|
t.Fatalf("saved skill content mismatch:\n%s", string(content))
|
||||||
|
}
|
||||||
|
|
||||||
|
rec2 := httptest.NewRecorder()
|
||||||
|
req2 := httptest.NewRequest(http.MethodGet, "/api/skills", nil)
|
||||||
|
mux.ServeHTTP(rec2, req2)
|
||||||
|
if rec2.Code != http.StatusOK {
|
||||||
|
t.Fatalf("list status = %d, want %d, body=%s", rec2.Code, http.StatusOK, rec2.Body.String())
|
||||||
|
}
|
||||||
|
var listResp skillSupportResponse
|
||||||
|
if err := json.Unmarshal(rec2.Body.Bytes(), &listResp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal list response error = %v", err)
|
||||||
|
}
|
||||||
|
found := false
|
||||||
|
for _, skill := range listResp.Skills {
|
||||||
|
if skill.Name == "plain-skill" && skill.Source == "workspace" && skill.Description == "Plain Skill" {
|
||||||
|
found = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
t.Fatalf("plain-skill should be listed after import, got %#v", listResp.Skills)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleDeleteSkill(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
workspace := filepath.Join(t.TempDir(), "workspace")
|
||||||
|
cfg.Agents.Defaults.Workspace = workspace
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
skillDir := filepath.Join(workspace, "skills", "delete-me")
|
||||||
|
if err := os.MkdirAll(skillDir, 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(skillDir, "SKILL.md"),
|
||||||
|
[]byte("---\nname: delete-me\ndescription: delete me\n---\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodDelete, "/api/skills/delete-me", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(skillDir); !os.IsNotExist(err) {
|
||||||
|
t.Fatalf("skill directory should be removed, stat err=%v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
323
web/backend/api/tools.go
Normal file
323
web/backend/api/tools.go
Normal file
|
|
@ -0,0 +1,323 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"runtime"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
type toolCatalogEntry struct {
|
||||||
|
Name string
|
||||||
|
Description string
|
||||||
|
Category string
|
||||||
|
ConfigKey string
|
||||||
|
}
|
||||||
|
|
||||||
|
type toolSupportItem struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Category string `json:"category"`
|
||||||
|
ConfigKey string `json:"config_key"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
ReasonCode string `json:"reason_code,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type toolSupportResponse struct {
|
||||||
|
Tools []toolSupportItem `json:"tools"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type toolStateRequest struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
}
|
||||||
|
|
||||||
|
var toolCatalog = []toolCatalogEntry{
|
||||||
|
{
|
||||||
|
Name: "read_file",
|
||||||
|
Description: "Read file content from the workspace or explicitly allowed paths.",
|
||||||
|
Category: "filesystem",
|
||||||
|
ConfigKey: "read_file",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "write_file",
|
||||||
|
Description: "Create or overwrite files within the writable workspace scope.",
|
||||||
|
Category: "filesystem",
|
||||||
|
ConfigKey: "write_file",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "list_dir",
|
||||||
|
Description: "Inspect directories and enumerate files available to the agent.",
|
||||||
|
Category: "filesystem",
|
||||||
|
ConfigKey: "list_dir",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "edit_file",
|
||||||
|
Description: "Apply targeted edits to existing files without rewriting everything.",
|
||||||
|
Category: "filesystem",
|
||||||
|
ConfigKey: "edit_file",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "append_file",
|
||||||
|
Description: "Append content to the end of an existing file.",
|
||||||
|
Category: "filesystem",
|
||||||
|
ConfigKey: "append_file",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "exec",
|
||||||
|
Description: "Run shell commands inside the configured workspace sandbox.",
|
||||||
|
Category: "filesystem",
|
||||||
|
ConfigKey: "exec",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "cron",
|
||||||
|
Description: "Schedule one-time or recurring reminders, jobs, and shell commands.",
|
||||||
|
Category: "automation",
|
||||||
|
ConfigKey: "cron",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "web_search",
|
||||||
|
Description: "Search the web using the configured providers.",
|
||||||
|
Category: "web",
|
||||||
|
ConfigKey: "web",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "web_fetch",
|
||||||
|
Description: "Fetch and summarize the contents of a webpage.",
|
||||||
|
Category: "web",
|
||||||
|
ConfigKey: "web_fetch",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "message",
|
||||||
|
Description: "Send a follow-up message back to the active user or chat.",
|
||||||
|
Category: "communication",
|
||||||
|
ConfigKey: "message",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "send_file",
|
||||||
|
Description: "Send an outbound file or media attachment to the active chat.",
|
||||||
|
Category: "communication",
|
||||||
|
ConfigKey: "send_file",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "find_skills",
|
||||||
|
Description: "Search external skill registries for installable skills.",
|
||||||
|
Category: "skills",
|
||||||
|
ConfigKey: "find_skills",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "install_skill",
|
||||||
|
Description: "Install a skill into the current workspace from a registry.",
|
||||||
|
Category: "skills",
|
||||||
|
ConfigKey: "install_skill",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "spawn",
|
||||||
|
Description: "Launch a background subagent for long-running or delegated work.",
|
||||||
|
Category: "agents",
|
||||||
|
ConfigKey: "spawn",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "i2c",
|
||||||
|
Description: "Interact with I2C hardware devices exposed on the host.",
|
||||||
|
Category: "hardware",
|
||||||
|
ConfigKey: "i2c",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "spi",
|
||||||
|
Description: "Interact with SPI hardware devices exposed on the host.",
|
||||||
|
Category: "hardware",
|
||||||
|
ConfigKey: "spi",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "tool_search_tool_regex",
|
||||||
|
Description: "Discover hidden MCP tools by regex search when tool discovery is enabled.",
|
||||||
|
Category: "discovery",
|
||||||
|
ConfigKey: "mcp.discovery.use_regex",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "tool_search_tool_bm25",
|
||||||
|
Description: "Discover hidden MCP tools by semantic ranking when tool discovery is enabled.",
|
||||||
|
Category: "discovery",
|
||||||
|
ConfigKey: "mcp.discovery.use_bm25",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) registerToolRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/tools", h.handleListTools)
|
||||||
|
mux.HandleFunc("PUT /api/tools/{name}/state", h.handleUpdateToolState)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleListTools(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(toolSupportResponse{
|
||||||
|
Tools: buildToolSupport(cfg),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleUpdateToolState(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var req toolStateRequest
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := applyToolState(cfg, r.PathValue("name"), req.Enabled); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildToolSupport(cfg *config.Config) []toolSupportItem {
|
||||||
|
items := make([]toolSupportItem, 0, len(toolCatalog))
|
||||||
|
for _, entry := range toolCatalog {
|
||||||
|
status := "disabled"
|
||||||
|
reasonCode := ""
|
||||||
|
|
||||||
|
switch entry.Name {
|
||||||
|
case "find_skills", "install_skill":
|
||||||
|
if cfg.Tools.IsToolEnabled(entry.ConfigKey) {
|
||||||
|
if cfg.Tools.IsToolEnabled("skills") {
|
||||||
|
status = "enabled"
|
||||||
|
} else {
|
||||||
|
status = "blocked"
|
||||||
|
reasonCode = "requires_skills"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "spawn":
|
||||||
|
if cfg.Tools.IsToolEnabled(entry.ConfigKey) {
|
||||||
|
if cfg.Tools.IsToolEnabled("subagent") {
|
||||||
|
status = "enabled"
|
||||||
|
} else {
|
||||||
|
status = "blocked"
|
||||||
|
reasonCode = "requires_subagent"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "tool_search_tool_regex":
|
||||||
|
status, reasonCode = resolveDiscoveryToolSupport(cfg, cfg.Tools.MCP.Discovery.UseRegex)
|
||||||
|
case "tool_search_tool_bm25":
|
||||||
|
status, reasonCode = resolveDiscoveryToolSupport(cfg, cfg.Tools.MCP.Discovery.UseBM25)
|
||||||
|
case "i2c", "spi":
|
||||||
|
status, reasonCode = resolveHardwareToolSupport(cfg.Tools.IsToolEnabled(entry.ConfigKey))
|
||||||
|
default:
|
||||||
|
if cfg.Tools.IsToolEnabled(entry.ConfigKey) {
|
||||||
|
status = "enabled"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
items = append(items, toolSupportItem{
|
||||||
|
Name: entry.Name,
|
||||||
|
Description: entry.Description,
|
||||||
|
Category: entry.Category,
|
||||||
|
ConfigKey: entry.ConfigKey,
|
||||||
|
Status: status,
|
||||||
|
ReasonCode: reasonCode,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return items
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveHardwareToolSupport(enabled bool) (string, string) {
|
||||||
|
if !enabled {
|
||||||
|
return "disabled", ""
|
||||||
|
}
|
||||||
|
if runtime.GOOS != "linux" {
|
||||||
|
return "blocked", "requires_linux"
|
||||||
|
}
|
||||||
|
return "enabled", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveDiscoveryToolSupport(cfg *config.Config, methodEnabled bool) (string, string) {
|
||||||
|
if !cfg.Tools.IsToolEnabled("mcp") {
|
||||||
|
return "disabled", ""
|
||||||
|
}
|
||||||
|
if !cfg.Tools.MCP.Discovery.Enabled {
|
||||||
|
return "blocked", "requires_mcp_discovery"
|
||||||
|
}
|
||||||
|
if !methodEnabled {
|
||||||
|
return "disabled", ""
|
||||||
|
}
|
||||||
|
return "enabled", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func applyToolState(cfg *config.Config, toolName string, enabled bool) error {
|
||||||
|
switch toolName {
|
||||||
|
case "read_file":
|
||||||
|
cfg.Tools.ReadFile.Enabled = enabled
|
||||||
|
case "write_file":
|
||||||
|
cfg.Tools.WriteFile.Enabled = enabled
|
||||||
|
case "list_dir":
|
||||||
|
cfg.Tools.ListDir.Enabled = enabled
|
||||||
|
case "edit_file":
|
||||||
|
cfg.Tools.EditFile.Enabled = enabled
|
||||||
|
case "append_file":
|
||||||
|
cfg.Tools.AppendFile.Enabled = enabled
|
||||||
|
case "exec":
|
||||||
|
cfg.Tools.Exec.Enabled = enabled
|
||||||
|
case "cron":
|
||||||
|
cfg.Tools.Cron.Enabled = enabled
|
||||||
|
case "web_search":
|
||||||
|
cfg.Tools.Web.Enabled = enabled
|
||||||
|
case "web_fetch":
|
||||||
|
cfg.Tools.WebFetch.Enabled = enabled
|
||||||
|
case "message":
|
||||||
|
cfg.Tools.Message.Enabled = enabled
|
||||||
|
case "send_file":
|
||||||
|
cfg.Tools.SendFile.Enabled = enabled
|
||||||
|
case "find_skills":
|
||||||
|
cfg.Tools.FindSkills.Enabled = enabled
|
||||||
|
if enabled {
|
||||||
|
cfg.Tools.Skills.Enabled = true
|
||||||
|
}
|
||||||
|
case "install_skill":
|
||||||
|
cfg.Tools.InstallSkill.Enabled = enabled
|
||||||
|
if enabled {
|
||||||
|
cfg.Tools.Skills.Enabled = true
|
||||||
|
}
|
||||||
|
case "spawn":
|
||||||
|
cfg.Tools.Spawn.Enabled = enabled
|
||||||
|
if enabled {
|
||||||
|
cfg.Tools.Subagent.Enabled = true
|
||||||
|
}
|
||||||
|
case "i2c":
|
||||||
|
cfg.Tools.I2C.Enabled = enabled
|
||||||
|
case "spi":
|
||||||
|
cfg.Tools.SPI.Enabled = enabled
|
||||||
|
case "tool_search_tool_regex":
|
||||||
|
cfg.Tools.MCP.Discovery.UseRegex = enabled
|
||||||
|
if enabled {
|
||||||
|
cfg.Tools.MCP.Enabled = true
|
||||||
|
cfg.Tools.MCP.Discovery.Enabled = true
|
||||||
|
}
|
||||||
|
case "tool_search_tool_bm25":
|
||||||
|
cfg.Tools.MCP.Discovery.UseBM25 = enabled
|
||||||
|
if enabled {
|
||||||
|
cfg.Tools.MCP.Enabled = true
|
||||||
|
cfg.Tools.MCP.Discovery.Enabled = true
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("tool %q cannot be updated", toolName)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
198
web/backend/api/tools_test.go
Normal file
198
web/backend/api/tools_test.go
Normal file
|
|
@ -0,0 +1,198 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestHandleListTools(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.Tools.ReadFile.Enabled = true
|
||||||
|
cfg.Tools.WriteFile.Enabled = false
|
||||||
|
cfg.Tools.Cron.Enabled = true
|
||||||
|
cfg.Tools.FindSkills.Enabled = true
|
||||||
|
cfg.Tools.Skills.Enabled = true
|
||||||
|
cfg.Tools.Spawn.Enabled = true
|
||||||
|
cfg.Tools.Subagent.Enabled = false
|
||||||
|
cfg.Tools.MCP.Enabled = true
|
||||||
|
cfg.Tools.MCP.Discovery.Enabled = true
|
||||||
|
cfg.Tools.MCP.Discovery.UseRegex = true
|
||||||
|
cfg.Tools.MCP.Discovery.UseBM25 = false
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/tools", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp toolSupportResponse
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
gotTools := make(map[string]toolSupportItem, len(resp.Tools))
|
||||||
|
for _, tool := range resp.Tools {
|
||||||
|
gotTools[tool.Name] = tool
|
||||||
|
}
|
||||||
|
if gotTools["read_file"].Status != "enabled" {
|
||||||
|
t.Fatalf("read_file status = %q, want enabled", gotTools["read_file"].Status)
|
||||||
|
}
|
||||||
|
if gotTools["write_file"].Status != "disabled" {
|
||||||
|
t.Fatalf("write_file status = %q, want disabled", gotTools["write_file"].Status)
|
||||||
|
}
|
||||||
|
if gotTools["cron"].Status != "enabled" {
|
||||||
|
t.Fatalf("cron status = %q, want enabled", gotTools["cron"].Status)
|
||||||
|
}
|
||||||
|
if gotTools["spawn"].Status != "blocked" || gotTools["spawn"].ReasonCode != "requires_subagent" {
|
||||||
|
t.Fatalf("spawn = %#v, want blocked/requires_subagent", gotTools["spawn"])
|
||||||
|
}
|
||||||
|
if gotTools["find_skills"].Status != "enabled" {
|
||||||
|
t.Fatalf("find_skills status = %q, want enabled", gotTools["find_skills"].Status)
|
||||||
|
}
|
||||||
|
if gotTools["tool_search_tool_regex"].Status != "enabled" {
|
||||||
|
t.Fatalf("tool_search_tool_regex status = %q, want enabled", gotTools["tool_search_tool_regex"].Status)
|
||||||
|
}
|
||||||
|
if gotTools["tool_search_tool_regex"].ConfigKey != "mcp.discovery.use_regex" {
|
||||||
|
t.Fatalf(
|
||||||
|
"tool_search_tool_regex config_key = %q, want mcp.discovery.use_regex",
|
||||||
|
gotTools["tool_search_tool_regex"].ConfigKey,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if gotTools["tool_search_tool_bm25"].Status != "disabled" {
|
||||||
|
t.Fatalf("tool_search_tool_bm25 status = %q, want disabled", gotTools["tool_search_tool_bm25"].Status)
|
||||||
|
}
|
||||||
|
if gotTools["tool_search_tool_bm25"].ConfigKey != "mcp.discovery.use_bm25" {
|
||||||
|
t.Fatalf(
|
||||||
|
"tool_search_tool_bm25 config_key = %q, want mcp.discovery.use_bm25",
|
||||||
|
gotTools["tool_search_tool_bm25"].ConfigKey,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if runtime.GOOS == "linux" {
|
||||||
|
if gotTools["i2c"].Status != "disabled" {
|
||||||
|
t.Fatalf("i2c status = %q, want disabled on linux when config is off", gotTools["i2c"].Status)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
cfg.Tools.I2C.Enabled = true
|
||||||
|
cfg.Tools.SPI.Enabled = true
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
req = httptest.NewRequest(http.MethodGet, "/api/tools", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("Unmarshal() error = %v", err)
|
||||||
|
}
|
||||||
|
gotTools = make(map[string]toolSupportItem, len(resp.Tools))
|
||||||
|
for _, tool := range resp.Tools {
|
||||||
|
gotTools[tool.Name] = tool
|
||||||
|
}
|
||||||
|
|
||||||
|
if gotTools["i2c"].Status != "blocked" || gotTools["i2c"].ReasonCode != "requires_linux" {
|
||||||
|
t.Fatalf("i2c = %#v, want blocked/requires_linux", gotTools["i2c"])
|
||||||
|
}
|
||||||
|
if gotTools["spi"].Status != "blocked" || gotTools["spi"].ReasonCode != "requires_linux" {
|
||||||
|
t.Fatalf("spi = %#v, want blocked/requires_linux", gotTools["spi"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleUpdateToolState(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.Tools.Spawn.Enabled = false
|
||||||
|
cfg.Tools.Subagent.Enabled = false
|
||||||
|
cfg.Tools.Cron.Enabled = false
|
||||||
|
cfg.Tools.MCP.Enabled = false
|
||||||
|
cfg.Tools.MCP.Discovery.Enabled = false
|
||||||
|
cfg.Tools.MCP.Discovery.UseRegex = false
|
||||||
|
err = config.SaveConfig(configPath, cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPut,
|
||||||
|
"/api/tools/spawn/state",
|
||||||
|
bytes.NewBufferString(`{"enabled":true}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("spawn status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
rec2 := httptest.NewRecorder()
|
||||||
|
req2 := httptest.NewRequest(
|
||||||
|
http.MethodPut,
|
||||||
|
"/api/tools/tool_search_tool_regex/state",
|
||||||
|
bytes.NewBufferString(`{"enabled":true}`),
|
||||||
|
)
|
||||||
|
req2.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec2, req2)
|
||||||
|
if rec2.Code != http.StatusOK {
|
||||||
|
t.Fatalf("regex status = %d, want %d, body=%s", rec2.Code, http.StatusOK, rec2.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
rec3 := httptest.NewRecorder()
|
||||||
|
req3 := httptest.NewRequest(
|
||||||
|
http.MethodPut,
|
||||||
|
"/api/tools/cron/state",
|
||||||
|
bytes.NewBufferString(`{"enabled":true}`),
|
||||||
|
)
|
||||||
|
req3.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec3, req3)
|
||||||
|
if rec3.Code != http.StatusOK {
|
||||||
|
t.Fatalf("cron status = %d, want %d, body=%s", rec3.Code, http.StatusOK, rec3.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
updated, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig(updated) error = %v", err)
|
||||||
|
}
|
||||||
|
if !updated.Tools.Spawn.Enabled || !updated.Tools.Subagent.Enabled {
|
||||||
|
t.Fatalf("spawn/subagent should both be enabled: %#v", updated.Tools)
|
||||||
|
}
|
||||||
|
if !updated.Tools.MCP.Enabled || !updated.Tools.MCP.Discovery.Enabled || !updated.Tools.MCP.Discovery.UseRegex {
|
||||||
|
t.Fatalf("mcp regex discovery should be enabled: %#v", updated.Tools.MCP)
|
||||||
|
}
|
||||||
|
if !updated.Tools.Cron.Enabled {
|
||||||
|
t.Fatalf("cron should be enabled: %#v", updated.Tools.Cron)
|
||||||
|
}
|
||||||
|
}
|
||||||
1
web/backend/dist/.gitkeep
vendored
1
web/backend/dist/.gitkeep
vendored
|
|
@ -0,0 +1 @@
|
||||||
|
# Keep the embedded web backend dist directory in version control.
|
||||||
|
|
@ -25,6 +25,7 @@ import (
|
||||||
"github.com/sipeed/picoclaw/web/backend/api"
|
"github.com/sipeed/picoclaw/web/backend/api"
|
||||||
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
||||||
"github.com/sipeed/picoclaw/web/backend/middleware"
|
"github.com/sipeed/picoclaw/web/backend/middleware"
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
|
|
@ -51,7 +52,7 @@ func main() {
|
||||||
flag.Parse()
|
flag.Parse()
|
||||||
|
|
||||||
// Resolve config path
|
// Resolve config path
|
||||||
configPath := getDefaultConfigPath()
|
configPath := utils.GetDefaultConfigPath()
|
||||||
if flag.NArg() > 0 {
|
if flag.NArg() > 0 {
|
||||||
configPath = flag.Arg(0)
|
configPath = flag.Arg(0)
|
||||||
}
|
}
|
||||||
|
|
@ -60,6 +61,10 @@ func main() {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Fatalf("Failed to resolve config path: %v", err)
|
log.Fatalf("Failed to resolve config path: %v", err)
|
||||||
}
|
}
|
||||||
|
err = utils.EnsureOnboarded(absPath)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("Warning: Failed to initialize PicoClaw config automatically: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
var explicitPort bool
|
var explicitPort bool
|
||||||
var explicitPublic bool
|
var explicitPublic bool
|
||||||
|
|
@ -109,7 +114,7 @@ func main() {
|
||||||
|
|
||||||
// API Routes (e.g. /api/status)
|
// API Routes (e.g. /api/status)
|
||||||
apiHandler := api.NewHandler(absPath)
|
apiHandler := api.NewHandler(absPath)
|
||||||
apiHandler.SetServerOptions(portNum, effectivePublic, launcherCfg.AllowedCIDRs)
|
apiHandler.SetServerOptions(portNum, effectivePublic, explicitPublic, launcherCfg.AllowedCIDRs)
|
||||||
apiHandler.RegisterRoutes(mux)
|
apiHandler.RegisterRoutes(mux)
|
||||||
|
|
||||||
// Frontend Embedded Assets
|
// Frontend Embedded Assets
|
||||||
|
|
@ -128,13 +133,13 @@ func main() {
|
||||||
)
|
)
|
||||||
|
|
||||||
// Print startup banner
|
// Print startup banner
|
||||||
fmt.Print(banner)
|
fmt.Print(utils.Banner)
|
||||||
fmt.Println()
|
fmt.Println()
|
||||||
fmt.Println(" Open the following URL in your browser:")
|
fmt.Println(" Open the following URL in your browser:")
|
||||||
fmt.Println()
|
fmt.Println()
|
||||||
fmt.Printf(" >> http://localhost:%s <<\n", effectivePort)
|
fmt.Printf(" >> http://localhost:%s <<\n", effectivePort)
|
||||||
if effectivePublic {
|
if effectivePublic {
|
||||||
if ip := getLocalIP(); ip != "" {
|
if ip := utils.GetLocalIP(); ip != "" {
|
||||||
fmt.Printf(" >> http://%s:%s <<\n", ip, effectivePort)
|
fmt.Printf(" >> http://%s:%s <<\n", ip, effectivePort)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -145,7 +150,7 @@ func main() {
|
||||||
go func() {
|
go func() {
|
||||||
time.Sleep(500 * time.Millisecond)
|
time.Sleep(500 * time.Millisecond)
|
||||||
url := "http://localhost:" + effectivePort
|
url := "http://localhost:" + effectivePort
|
||||||
if err := openBrowser(url); err != nil {
|
if err := utils.OpenBrowser(url); err != nil {
|
||||||
log.Printf("Warning: Failed to auto-open browser: %v", err)
|
log.Printf("Warning: Failed to auto-open browser: %v", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
|
||||||
|
|
@ -1,19 +1,10 @@
|
||||||
package main
|
package utils
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"os"
|
|
||||||
"os/exec"
|
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
colorBlue = "\x1b[38;2;62;93;185m"
|
colorBlue = "\x1b[38;2;62;93;185m"
|
||||||
colorRed = "\x1b[38;2;213;70;70m"
|
colorRed = "\x1b[38;2;213;70;70m"
|
||||||
colorReset = "\x1b[0m"
|
colorReset = "\x1b[0m"
|
||||||
banner = "\r\n" +
|
Banner = "\r\n" +
|
||||||
colorBlue + "██████╗ ██╗ ██████╗ ██████╗ " + colorRed + " ██████╗██╗ █████╗ ██╗ ██╗\n" +
|
colorBlue + "██████╗ ██╗ ██████╗ ██████╗ " + colorRed + " ██████╗██╗ █████╗ ██╗ ██╗\n" +
|
||||||
colorBlue + "██╔══██╗██║██╔════╝██╔═══██╗" + colorRed + "██╔════╝██║ ██╔══██╗██║ ██║\n" +
|
colorBlue + "██╔══██╗██║██╔════╝██╔═══██╗" + colorRed + "██╔════╝██║ ██╔══██╗██║ ██║\n" +
|
||||||
colorBlue + "██████╔╝██║██║ ██║ ██║" + colorRed + "██║ ██║ ███████║██║ █╗ ██║\n" +
|
colorBlue + "██████╔╝██║██║ ██║ ██║" + colorRed + "██║ ██║ ███████║██║ █╗ ██║\n" +
|
||||||
|
|
@ -22,40 +13,3 @@ const (
|
||||||
colorBlue + "╚═╝ ╚═╝ ╚═════╝ ╚═════╝ " + colorRed + " ╚═════╝╚══════╝╚═╝ ╚═╝ ╚══╝╚══╝\n" +
|
colorBlue + "╚═╝ ╚═╝ ╚═════╝ ╚═════╝ " + colorRed + " ╚═════╝╚══════╝╚═╝ ╚═╝ ╚══╝╚══╝\n" +
|
||||||
colorReset
|
colorReset
|
||||||
)
|
)
|
||||||
|
|
||||||
// getDefaultConfigPath returns the default path to the picoclaw config file.
|
|
||||||
func getDefaultConfigPath() string {
|
|
||||||
home, err := os.UserHomeDir()
|
|
||||||
if err != nil {
|
|
||||||
return "config.json"
|
|
||||||
}
|
|
||||||
return filepath.Join(home, ".picoclaw", "config.json")
|
|
||||||
}
|
|
||||||
|
|
||||||
// getLocalIP returns the local IP address of the machine.
|
|
||||||
func getLocalIP() string {
|
|
||||||
addrs, err := net.InterfaceAddrs()
|
|
||||||
if err != nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
for _, a := range addrs {
|
|
||||||
if ipnet, ok := a.(*net.IPNet); ok && !ipnet.IP.IsLoopback() && ipnet.IP.To4() != nil {
|
|
||||||
return ipnet.IP.String()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// openBrowser automatically opens the given URL in the default browser.
|
|
||||||
func openBrowser(url string) error {
|
|
||||||
switch runtime.GOOS {
|
|
||||||
case "linux":
|
|
||||||
return exec.Command("xdg-open", url).Start()
|
|
||||||
case "windows":
|
|
||||||
return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start()
|
|
||||||
case "darwin":
|
|
||||||
return exec.Command("open", url).Start()
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("unsupported platform")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
42
web/backend/utils/onboard.go
Normal file
42
web/backend/utils/onboard.go
Normal file
|
|
@ -0,0 +1,42 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
var execCommand = exec.Command
|
||||||
|
|
||||||
|
func EnsureOnboarded(configPath string) error {
|
||||||
|
_, err := os.Stat(configPath)
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !os.IsNotExist(err) {
|
||||||
|
return fmt.Errorf("stat config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd := execCommand(FindPicoclawBinary(), "onboard")
|
||||||
|
cmd.Env = append(os.Environ(), "PICOCLAW_CONFIG="+configPath)
|
||||||
|
cmd.Stdin = strings.NewReader("n\n")
|
||||||
|
|
||||||
|
output, err := cmd.CombinedOutput()
|
||||||
|
if err != nil {
|
||||||
|
trimmed := strings.TrimSpace(string(output))
|
||||||
|
if trimmed == "" {
|
||||||
|
return fmt.Errorf("run onboard: %w", err)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("run onboard: %w: %s", err, trimmed)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err != nil {
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return fmt.Errorf("onboard completed but did not create config %s", configPath)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("verify config after onboard: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
101
web/backend/utils/onboard_test.go
Normal file
101
web/backend/utils/onboard_test.go
Normal file
|
|
@ -0,0 +1,101 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEnsureOnboardedSkipsWhenConfigExists(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
if err := os.WriteFile(configPath, []byte(`{}`), 0o644); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
origExecCommand := execCommand
|
||||||
|
defer func() { execCommand = origExecCommand }()
|
||||||
|
|
||||||
|
called := false
|
||||||
|
execCommand = func(name string, args ...string) *exec.Cmd {
|
||||||
|
called = true
|
||||||
|
return exec.Command("sh", "-c", "exit 1")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := EnsureOnboarded(configPath); err != nil {
|
||||||
|
t.Fatalf("EnsureOnboarded() error = %v", err)
|
||||||
|
}
|
||||||
|
if called {
|
||||||
|
t.Fatal("expected onboard command not to run when config already exists")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnsureOnboardedRunsOnboardWhenConfigMissing(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
t.Setenv("EXPECTED_CONFIG_PATH", configPath)
|
||||||
|
|
||||||
|
origExecCommand := execCommand
|
||||||
|
defer func() { execCommand = origExecCommand }()
|
||||||
|
|
||||||
|
var gotName string
|
||||||
|
var gotArgs []string
|
||||||
|
execCommand = func(name string, args ...string) *exec.Cmd {
|
||||||
|
gotName = name
|
||||||
|
gotArgs = append([]string(nil), args...)
|
||||||
|
return exec.Command(
|
||||||
|
"sh",
|
||||||
|
"-c",
|
||||||
|
`test "$PICOCLAW_CONFIG" = "$EXPECTED_CONFIG_PATH" &&
|
||||||
|
mkdir -p "$(dirname "$PICOCLAW_CONFIG")" &&
|
||||||
|
printf '{}' > "$PICOCLAW_CONFIG"`,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := EnsureOnboarded(configPath); err != nil {
|
||||||
|
t.Fatalf("EnsureOnboarded() error = %v", err)
|
||||||
|
}
|
||||||
|
if gotName == "" {
|
||||||
|
t.Fatal("expected onboard command to run")
|
||||||
|
}
|
||||||
|
if len(gotArgs) != 1 || gotArgs[0] != "onboard" {
|
||||||
|
t.Fatalf("command args = %#v, want []string{\"onboard\"}", gotArgs)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(configPath); err != nil {
|
||||||
|
t.Fatalf("expected config to be created: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnsureOnboardedFailsWhenOnboardDoesNotCreateConfig(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
|
||||||
|
origExecCommand := execCommand
|
||||||
|
defer func() { execCommand = origExecCommand }()
|
||||||
|
|
||||||
|
execCommand = func(name string, args ...string) *exec.Cmd {
|
||||||
|
return exec.Command("sh", "-c", "exit 0")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := EnsureOnboarded(configPath); err == nil {
|
||||||
|
t.Fatal("EnsureOnboarded() error = nil, want failure when onboard does not create config")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnsureOnboardedIncludesOnboardOutputOnFailure(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
|
||||||
|
origExecCommand := execCommand
|
||||||
|
defer func() { execCommand = origExecCommand }()
|
||||||
|
|
||||||
|
execCommand = func(name string, args ...string) *exec.Cmd {
|
||||||
|
return exec.Command("sh", "-c", "echo onboarding failed >&2; exit 2")
|
||||||
|
}
|
||||||
|
|
||||||
|
err := EnsureOnboarded(configPath)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("EnsureOnboarded() error = nil, want failure")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "onboarding failed") {
|
||||||
|
t.Fatalf("error = %q, want onboard output included", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue