Merge remote-tracking branch 'upstream/main' into Revised-search-mechanism
This commit is contained in:
commit
026175183a
88 changed files with 12638 additions and 2127 deletions
|
|
@ -5,6 +5,7 @@
|
||||||
# ANTHROPIC_API_KEY=sk-ant-xxx
|
# ANTHROPIC_API_KEY=sk-ant-xxx
|
||||||
# OPENAI_API_KEY=sk-xxx
|
# OPENAI_API_KEY=sk-xxx
|
||||||
# GEMINI_API_KEY=xxx
|
# GEMINI_API_KEY=xxx
|
||||||
|
# CEREBRAS_API_KEY=xxx
|
||||||
|
|
||||||
# ── Chat Channel ──────────────────────────
|
# ── Chat Channel ──────────────────────────
|
||||||
# TELEGRAM_BOT_TOKEN=123456:ABC...
|
# TELEGRAM_BOT_TOKEN=123456:ABC...
|
||||||
|
|
|
||||||
1038
README.fr.md
Normal file
1038
README.fr.md
Normal file
File diff suppressed because it is too large
Load diff
179
README.ja.md
179
README.ja.md
|
|
@ -12,7 +12,7 @@
|
||||||
<img src="https://img.shields.io/badge/license-MIT-green" alt="License">
|
<img src="https://img.shields.io/badge/license-MIT-green" alt="License">
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
[中文](README.zh.md) | **日本語** | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [English](README.md)
|
[中文](README.zh.md) | **日本語** | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | [English](README.md)
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
|
@ -209,7 +209,7 @@ picoclaw onboard
|
||||||
|
|
||||||
**3. API キーの取得**
|
**3. API キーの取得**
|
||||||
|
|
||||||
- **LLM プロバイダー**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
- **LLM プロバイダー**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys) · [Qwen](https://dashscope.console.aliyun.com)
|
||||||
- **Web 検索**(任意): [Brave Search](https://brave.com/search/api) - 無料枠あり(月 2000 リクエスト)
|
- **Web 検索**(任意): [Brave Search](https://brave.com/search/api) - 無料枠あり(月 2000 リクエスト)
|
||||||
|
|
||||||
> **注意**: 完全な設定テンプレートは `config.example.json` を参照してください。
|
> **注意**: 完全な設定テンプレートは `config.example.json` を参照してください。
|
||||||
|
|
@ -621,6 +621,22 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
- `PICOCLAW_HEARTBEAT_ENABLED=false` で無効化
|
- `PICOCLAW_HEARTBEAT_ENABLED=false` で無効化
|
||||||
- `PICOCLAW_HEARTBEAT_INTERVAL=60` で間隔変更
|
- `PICOCLAW_HEARTBEAT_INTERVAL=60` で間隔変更
|
||||||
|
|
||||||
|
### プロバイダー
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> Groq は Whisper による無料の音声文字起こしを提供しています。設定すると、Telegram の音声メッセージが自動的に文字起こしされます。
|
||||||
|
|
||||||
|
| プロバイダー | 用途 | API キー取得先 |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| `gemini` | LLM(Gemini 直接) | [aistudio.google.com](https://aistudio.google.com) |
|
||||||
|
| `zhipu` | LLM(Zhipu 直接) | [bigmodel.cn](https://bigmodel.cn) |
|
||||||
|
| `openrouter`(要テスト) | LLM(推奨、全モデルにアクセス可能) | [openrouter.ai](https://openrouter.ai) |
|
||||||
|
| `anthropic`(要テスト) | LLM(Claude 直接) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
|
| `openai`(要テスト) | LLM(GPT 直接) | [platform.openai.com](https://platform.openai.com) |
|
||||||
|
| `deepseek`(要テスト) | LLM(DeepSeek 直接) | [platform.deepseek.com](https://platform.deepseek.com) |
|
||||||
|
| `groq` | LLM + **音声文字起こし**(Whisper) | [console.groq.com](https://console.groq.com) |
|
||||||
|
| `cerebras` | LLM(Cerebras 直接) | [cerebras.ai](https://cerebras.ai) |
|
||||||
|
|
||||||
### 基本設定
|
### 基本設定
|
||||||
|
|
||||||
1. **設定ファイルの作成:**
|
1. **設定ファイルの作成:**
|
||||||
|
|
@ -714,6 +730,163 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
### モデル設定 (model_list)
|
||||||
|
|
||||||
|
> **新機能!** PicoClaw は現在 **モデル中心** の設定アプローチを採用しています。`ベンダー/モデル` 形式(例: `zhipu/glm-4.7`)を指定するだけで、新しいプロバイダーを追加できます—**コードの変更は一切不要!**
|
||||||
|
|
||||||
|
この設計は、柔軟なプロバイダー選択による **マルチエージェントサポート** も可能にします:
|
||||||
|
|
||||||
|
- **異なるエージェント、異なるプロバイダー** : 各エージェントは独自の LLM プロバイダーを使用可能
|
||||||
|
- **フォールバックモデル** : 耐障性のため、プライマリモデルとフォールバックモデルを設定可能
|
||||||
|
- **ロードバランシング** : 複数のエンドポイントにリクエストを分散
|
||||||
|
- **集中設定管理** : すべてのプロバイダーを一箇所で管理
|
||||||
|
|
||||||
|
#### 📋 サポートされているすべてのベンダー
|
||||||
|
|
||||||
|
| ベンダー | `model` プレフィックス | デフォルト API Base | プロトコル | API キー |
|
||||||
|
|-------------|-----------------|---------------------|----------|---------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [キーを取得](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [キーを取得](https://console.anthropic.com) |
|
||||||
|
| **Zhipu AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [キーを取得](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [キーを取得](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [キーを取得](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [キーを取得](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [キーを取得](https://platform.moonshot.cn) |
|
||||||
|
| **Qwen (Alibaba)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [キーを取得](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [キーを取得](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | ローカル(キー不要) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [キーを取得](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | ローカル |
|
||||||
|
| **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) |
|
||||||
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | カスタム | OAuthのみ |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### 基本設定
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### ベンダー別の例
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Zhipu AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (OAuth使用)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> OAuth認証を設定するには、`picoclaw auth login --provider anthropic` を実行してください。
|
||||||
|
|
||||||
|
#### ロードバランシング
|
||||||
|
|
||||||
|
同じモデル名で複数のエンドポイントを設定すると、PicoClaw が自動的にラウンドロビンで分散します:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 従来の `providers` 設定からの移行
|
||||||
|
|
||||||
|
古い `providers` 設定は**非推奨**ですが、後方互換性のためにサポートされています。
|
||||||
|
|
||||||
|
**旧設定(非推奨):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**新設定(推奨):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
詳細な移行ガイドは、[docs/migration/model-list-migration.md](docs/migration/model-list-migration.md) を参照してください。
|
||||||
|
|
||||||
## CLI リファレンス
|
## CLI リファレンス
|
||||||
|
|
||||||
| コマンド | 説明 |
|
| コマンド | 説明 |
|
||||||
|
|
@ -771,5 +944,7 @@ Web 検索を有効にするには:
|
||||||
|---------|--------|------------|
|
|---------|--------|------------|
|
||||||
| **OpenRouter** | 月 200K トークン | 複数モデル(Claude, GPT-4 など) |
|
| **OpenRouter** | 月 200K トークン | 複数モデル(Claude, GPT-4 など) |
|
||||||
| **Zhipu** | 月 200K トークン | 中国ユーザー向け最適 |
|
| **Zhipu** | 月 200K トークン | 中国ユーザー向け最適 |
|
||||||
|
| **Qwen** | 無料枠あり | 通義千問 (Qwen) |
|
||||||
| **Brave Search** | 月 2000 クエリ | Web 検索機能 |
|
| **Brave Search** | 月 2000 クエリ | Web 検索機能 |
|
||||||
| **Groq** | 無料枠あり | 高速推論(Llama, Mixtral) |
|
| **Groq** | 無料枠あり | 高速推論(Llama, Mixtral) |
|
||||||
|
| **Cerebras** | 無料枠あり | 高速推論(Llama, Qwen など) |
|
||||||
|
|
|
||||||
216
README.md
216
README.md
|
|
@ -14,7 +14,7 @@
|
||||||
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
[中文](README.zh.md) | [日本語](README.ja.md) | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | **English**
|
[中文](README.zh.md) | [日本語](README.ja.md) | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | **English**
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
@ -209,18 +209,24 @@ picoclaw onboard
|
||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"model": "glm-4.7",
|
"model": "gpt4",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
"model_list": [
|
||||||
"openrouter": {
|
{
|
||||||
"api_key": "xxx",
|
"model_name": "gpt4",
|
||||||
"api_base": "https://openrouter.ai/api/v1"
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "your-api-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "your-anthropic-key"
|
||||||
}
|
}
|
||||||
},
|
],
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": {
|
"web": {
|
||||||
"brave": {
|
"brave": {
|
||||||
|
|
@ -237,6 +243,8 @@ picoclaw onboard
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **New**: The `model_list` configuration format allows zero-code provider addition. See [Model Configuration](#-model-configuration) for details.
|
||||||
|
|
||||||
**3. Get API Keys**
|
**3. Get API Keys**
|
||||||
|
|
||||||
* **LLM Provider**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
* **LLM Provider**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
||||||
|
|
@ -326,7 +334,8 @@ picoclaw gateway
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allow_from": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"],
|
||||||
|
"mention_only": false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -339,6 +348,10 @@ picoclaw gateway
|
||||||
* Bot Permissions: `Send Messages`, `Read Message History`
|
* Bot Permissions: `Send Messages`, `Read Message History`
|
||||||
* Open the generated invite URL and add the bot to your server
|
* Open the generated invite URL and add the bot to your server
|
||||||
|
|
||||||
|
**Optional: Mention-only mode**
|
||||||
|
|
||||||
|
Set `"mention_only": true` to make the bot respond only when @-mentioned. Useful for shared servers where you want the bot to respond only when explicitly called.
|
||||||
|
|
||||||
**6. Run**
|
**6. Run**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
@ -677,7 +690,193 @@ The subagent has access to tools (message, web_search, etc.) and can communicate
|
||||||
| `anthropic(To be tested)` | LLM (Claude direct) | [console.anthropic.com](https://console.anthropic.com) |
|
| `anthropic(To be tested)` | LLM (Claude direct) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
| `openai(To be tested)` | LLM (GPT direct) | [platform.openai.com](https://platform.openai.com) |
|
| `openai(To be tested)` | LLM (GPT direct) | [platform.openai.com](https://platform.openai.com) |
|
||||||
| `deepseek(To be tested)` | LLM (DeepSeek direct) | [platform.deepseek.com](https://platform.deepseek.com) |
|
| `deepseek(To be tested)` | LLM (DeepSeek direct) | [platform.deepseek.com](https://platform.deepseek.com) |
|
||||||
|
| `qwen` | LLM (Qwen direct) | [dashscope.console.aliyun.com](https://dashscope.console.aliyun.com) |
|
||||||
| `groq` | LLM + **Voice transcription** (Whisper) | [console.groq.com](https://console.groq.com) |
|
| `groq` | LLM + **Voice transcription** (Whisper) | [console.groq.com](https://console.groq.com) |
|
||||||
|
| `cerebras` | LLM (Cerebras direct) | [cerebras.ai](https://cerebras.ai) |
|
||||||
|
|
||||||
|
### Model Configuration (model_list)
|
||||||
|
|
||||||
|
> **What's New?** PicoClaw now uses a **model-centric** configuration approach. Simply specify `vendor/model` format (e.g., `zhipu/glm-4.7`) to add new providers—**zero code changes required!**
|
||||||
|
|
||||||
|
This design also enables **multi-agent support** with flexible provider selection:
|
||||||
|
|
||||||
|
- **Different agents, different providers**: Each agent can use its own LLM provider
|
||||||
|
- **Model fallbacks**: Configure primary and fallback models for resilience
|
||||||
|
- **Load balancing**: Distribute requests across multiple endpoints
|
||||||
|
- **Centralized configuration**: Manage all providers in one place
|
||||||
|
|
||||||
|
#### 📋 All Supported Vendors
|
||||||
|
|
||||||
|
| Vendor | `model` Prefix | Default API Base | Protocol | API Key |
|
||||||
|
|--------|----------------|------------------|----------|---------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [Get Key](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Get Key](https://console.anthropic.com) |
|
||||||
|
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Get Key](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Get Key](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Get Key](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Get Key](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [Get Key](https://platform.moonshot.cn) |
|
||||||
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Get Key](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Get Key](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (no key needed) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||||
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Get Key](https://cerebras.ai) |
|
||||||
|
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://console.volcengine.com) |
|
||||||
|
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth only |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### Basic Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Vendor-Specific Examples
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**智谱 AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**DeepSeek**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (with OAuth)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> Run `picoclaw auth login --provider anthropic` to set up OAuth credentials.
|
||||||
|
|
||||||
|
**Ollama (local)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "llama3",
|
||||||
|
"model": "ollama/llama3"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Custom Proxy/API**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "my-custom-model",
|
||||||
|
"model": "openai/custom-model",
|
||||||
|
"api_base": "https://my-proxy.com/v1",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Load Balancing
|
||||||
|
|
||||||
|
Configure multiple endpoints for the same model name—PicoClaw will automatically round-robin between them:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Migration from Legacy `providers` Config
|
||||||
|
|
||||||
|
The old `providers` configuration is **deprecated** but still supported for backward compatibility.
|
||||||
|
|
||||||
|
**Old Config (deprecated):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**New Config (recommended):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
For detailed migration guide, see [docs/migration/model-list-migration.md](docs/migration/model-list-migration.md).
|
||||||
|
|
||||||
### Provider Architecture
|
### Provider Architecture
|
||||||
|
|
||||||
|
|
@ -883,3 +1082,4 @@ This happens when another instance of the bot is running. Make sure only one `pi
|
||||||
| **Zhipu** | 200K tokens/month | Best for Chinese users |
|
| **Zhipu** | 200K tokens/month | Best for Chinese users |
|
||||||
| **Brave Search** | 2000 queries/month | Web search functionality |
|
| **Brave Search** | 2000 queries/month | Web search functionality |
|
||||||
| **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.) |
|
||||||
|
|
|
||||||
159
README.pt-br.md
159
README.pt-br.md
|
|
@ -14,7 +14,7 @@
|
||||||
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
[中文](README.zh.md) | [日本語](README.ja.md) | [English](README.md) | **Português**
|
[中文](README.zh.md) | [日本語](README.ja.md) | **Português** | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | [English](README.md)
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
@ -795,6 +795,163 @@ picoclaw agent -m "Ola, como vai?"
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
### Configuração de Modelo (model_list)
|
||||||
|
|
||||||
|
> **Novidade!** PicoClaw agora usa uma abordagem de configuração **centrada no modelo**. Basta especificar o formato `fornecedor/modelo` (ex: `zhipu/glm-4.7`) para adicionar novos provedores—**nenhuma alteração de código necessária!**
|
||||||
|
|
||||||
|
Este design também possibilita o **suporte multi-agent** com seleção flexível de provedores:
|
||||||
|
|
||||||
|
- **Diferentes agentes, diferentes provedores** : Cada agente pode usar seu próprio provedor LLM
|
||||||
|
- **Modelos de fallback** : Configure modelos primários e de reserva para resiliência
|
||||||
|
- **Balanceamento de carga** : Distribua solicitações entre múltiplos endpoints
|
||||||
|
- **Configuração centralizada** : Gerencie todos os provedores em um só lugar
|
||||||
|
|
||||||
|
#### 📋 Todos os Fornecedores Suportados
|
||||||
|
|
||||||
|
| Fornecedor | Prefixo `model` | API Base Padrão | Protocolo | Chave API |
|
||||||
|
|-------------|-----------------|------------------|----------|-----------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [Obter Chave](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Obter Chave](https://console.anthropic.com) |
|
||||||
|
| **Zhipu AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Obter Chave](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Obter Chave](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Obter Chave](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Obter Chave](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [Obter Chave](https://platform.moonshot.cn) |
|
||||||
|
| **Qwen (Alibaba)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Obter Chave](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Obter Chave](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (sem chave necessária) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Obter Chave](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||||
|
| **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) |
|
||||||
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | Custom | Apenas OAuth |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### Configuração Básica
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Exemplos por Fornecedor
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Zhipu AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (com OAuth)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> Execute `picoclaw auth login --provider anthropic` para configurar credenciais OAuth.
|
||||||
|
|
||||||
|
#### Balanceamento de Carga
|
||||||
|
|
||||||
|
Configure vários endpoints para o mesmo nome de modelo—PicoClaw fará round-robin automaticamente entre eles:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Migração da Configuração Legada `providers`
|
||||||
|
|
||||||
|
A configuração antiga `providers` está **descontinuada** mas ainda é suportada para compatibilidade reversa.
|
||||||
|
|
||||||
|
**Configuração Antiga (descontinuada):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Nova Configuração (recomendada):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Para o guia de migração detalhado, consulte [docs/migration/model-list-migration.md](docs/migration/model-list-migration.md).
|
||||||
|
|
||||||
## Referência CLI
|
## Referência CLI
|
||||||
|
|
||||||
| Comando | Descrição |
|
| Comando | Descrição |
|
||||||
|
|
|
||||||
159
README.vi.md
159
README.vi.md
|
|
@ -14,7 +14,7 @@
|
||||||
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
**Tiếng Việt** | [中文](README.zh.md) | [日本語](README.ja.md) | [English](README.md)
|
[中文](README.zh.md) | [日本語](README.ja.md) | [Português](README.pt-br.md) | **Tiếng Việt** | [Français](README.fr.md) | [English](README.md)
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
@ -772,6 +772,163 @@ picoclaw agent -m "Xin chào"
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
### Cấu hình Mô hình (model_list)
|
||||||
|
|
||||||
|
> **Tính năng mới!** PicoClaw hiện sử dụng phương pháp cấu hình **đặt mô hình vào trung tâm**. Chỉ cần chỉ định dạng `nhà cung cấp/mô hình` (ví dụ: `zhipu/glm-4.7`) để thêm nhà cung cấp mới—**không cần thay đổi mã!**
|
||||||
|
|
||||||
|
Thiết kế này cũng cho phép **hỗ trợ đa tác nhân** với lựa chọn nhà cung cấp linh hoạt:
|
||||||
|
|
||||||
|
- **Tác nhân khác nhau, nhà cung cấp khác nhau** : Mỗi tác nhân có thể sử dụng nhà cung cấp LLM riêng
|
||||||
|
- **Mô hình dự phòng** : Cấu hình mô hình chính và dự phòng để tăng độ tin cậy
|
||||||
|
- **Cân bằng tải** : Phân phối yêu cầu trên nhiều endpoint khác nhau
|
||||||
|
- **Cấu hình tập trung** : Quản lý tất cả nhà cung cấp ở một nơi
|
||||||
|
|
||||||
|
#### 📋 Tất cả Nhà cung cấp được Hỗ trợ
|
||||||
|
|
||||||
|
| Nhà cung cấp | Prefix `model` | API Base Mặc định | Giao thức | Khóa API |
|
||||||
|
|-------------|----------------|-------------------|-----------|----------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [Lấy Khóa](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Lấy Khóa](https://console.anthropic.com) |
|
||||||
|
| **Zhipu AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Lấy Khóa](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Lấy Khóa](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Lấy Khóa](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Lấy Khóa](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [Lấy Khóa](https://platform.moonshot.cn) |
|
||||||
|
| **Qwen (Alibaba)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Lấy Khóa](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Lấy Khóa](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (không cần khóa) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Lấy Khóa](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||||
|
| **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) |
|
||||||
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | Tùy chỉnh | Chỉ OAuth |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### Cấu hình Cơ bản
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Ví dụ theo Nhà cung cấp
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Zhipu AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (với OAuth)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> Chạy `picoclaw auth login --provider anthropic` để thiết lập thông tin xác thực OAuth.
|
||||||
|
|
||||||
|
#### Cân bằng Tải tải
|
||||||
|
|
||||||
|
Định cấu hình nhiều endpoint cho cùng một tên mô hình—PicoClaw sẽ tự động phân phối round-robin giữa chúng:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Chuyển đổi từ Cấu hình `providers` Cũ
|
||||||
|
|
||||||
|
Cấu hình `providers` cũ đã **ngừng sử dụng** nhưng vẫn được hỗ trợ để tương thích ngược.
|
||||||
|
|
||||||
|
**Cấu hình Cũ (đã ngừng sử dụng):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Cấu hình Mới (khuyến nghị):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Xem hướng dẫn chuyển đổi chi tiết tại [docs/migration/model-list-migration.md](docs/migration/model-list-migration.md).
|
||||||
|
|
||||||
## Tham chiếu CLI
|
## Tham chiếu CLI
|
||||||
|
|
||||||
| Lệnh | Mô tả |
|
| Lệnh | Mô tả |
|
||||||
|
|
|
||||||
211
README.zh.md
211
README.zh.md
|
|
@ -14,7 +14,7 @@
|
||||||
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
**中文** | [日本語](README.ja.md) | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [English](README.md)
|
**中文** | [日本語](README.ja.md) | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | [English](README.md)
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
@ -218,18 +218,24 @@ picoclaw onboard
|
||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"model": "glm-4.7",
|
"model": "gpt4",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
"model_list": [
|
||||||
"openrouter": {
|
{
|
||||||
"api_key": "xxx",
|
"model_name": "gpt4",
|
||||||
"api_base": "https://openrouter.ai/api/v1"
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "your-api-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "your-anthropic-key"
|
||||||
}
|
}
|
||||||
},
|
],
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
|
|
@ -245,6 +251,8 @@ picoclaw onboard
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **新功能**: `model_list` 配置格式支持零代码添加 provider。详见[模型配置](#-模型配置-model_list)章节。
|
||||||
|
|
||||||
**3. 获取 API Key**
|
**3. 获取 API Key**
|
||||||
|
|
||||||
* **LLM 提供商**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
* **LLM 提供商**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
||||||
|
|
@ -554,7 +562,193 @@ Agent 读取 HEARTBEAT.md
|
||||||
| `anthropic(待测试)` | LLM (Claude 直连) | [console.anthropic.com](https://console.anthropic.com) |
|
| `anthropic(待测试)` | LLM (Claude 直连) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
| `openai(待测试)` | LLM (GPT 直连) | [platform.openai.com](https://platform.openai.com) |
|
| `openai(待测试)` | LLM (GPT 直连) | [platform.openai.com](https://platform.openai.com) |
|
||||||
| `deepseek(待测试)` | LLM (DeepSeek 直连) | [platform.deepseek.com](https://platform.deepseek.com) |
|
| `deepseek(待测试)` | LLM (DeepSeek 直连) | [platform.deepseek.com](https://platform.deepseek.com) |
|
||||||
|
| `qwen` | LLM (通义千问) | [dashscope.console.aliyun.com](https://dashscope.console.aliyun.com) |
|
||||||
| `groq` | LLM + **语音转录** (Whisper) | [console.groq.com](https://console.groq.com) |
|
| `groq` | LLM + **语音转录** (Whisper) | [console.groq.com](https://console.groq.com) |
|
||||||
|
| `cerebras` | LLM (Cerebras 直连) | [cerebras.ai](https://cerebras.ai) |
|
||||||
|
|
||||||
|
### 模型配置 (model_list)
|
||||||
|
|
||||||
|
> **新功能!** PicoClaw 现在采用**以模型为中心**的配置方式。只需使用 `厂商/模型` 格式(如 `zhipu/glm-4.7`)即可添加新的 provider——**无需修改任何代码!**
|
||||||
|
|
||||||
|
该设计同时支持**多 Agent 场景**,提供灵活的 Provider 选择:
|
||||||
|
|
||||||
|
- **不同 Agent 使用不同 Provider**:每个 Agent 可以使用自己的 LLM provider
|
||||||
|
- **模型回退(Fallback)**:配置主模型和备用模型,提高可靠性
|
||||||
|
- **负载均衡**:在多个 API 端点之间分配请求
|
||||||
|
- **集中化配置**:在一个地方管理所有 provider
|
||||||
|
|
||||||
|
#### 📋 所有支持的厂商
|
||||||
|
|
||||||
|
| 厂商 | `model` 前缀 | 默认 API Base | 协议 | 获取 API Key |
|
||||||
|
|------|-------------|---------------|------|--------------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [获取密钥](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [获取密钥](https://console.anthropic.com) |
|
||||||
|
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [获取密钥](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [获取密钥](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [获取密钥](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [获取密钥](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [获取密钥](https://platform.moonshot.cn) |
|
||||||
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [获取密钥](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [获取密钥](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | 本地(无需密钥) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取密钥](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
||||||
|
| **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) |
|
||||||
|
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### 基础配置示例
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 各厂商配置示例
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**智谱 AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**DeepSeek**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (使用 OAuth)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> 运行 `picoclaw auth login --provider anthropic` 来设置 OAuth 凭证。
|
||||||
|
|
||||||
|
**Ollama (本地)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "llama3",
|
||||||
|
"model": "ollama/llama3"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**自定义代理/API**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "my-custom-model",
|
||||||
|
"model": "openai/custom-model",
|
||||||
|
"api_base": "https://my-proxy.com/v1",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 负载均衡
|
||||||
|
|
||||||
|
为同一个模型名称配置多个端点——PicoClaw 会自动在它们之间轮询:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 从旧的 `providers` 配置迁移
|
||||||
|
|
||||||
|
旧的 `providers` 配置格式**已弃用**,但为向后兼容仍支持。
|
||||||
|
|
||||||
|
**旧配置(已弃用):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**新配置(推荐):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
详细的迁移指南请参考 [docs/migration/model-list-migration.md](docs/migration/model-list-migration.md)。
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>智谱 (Zhipu) 配置示例</b></summary>
|
<summary><b>智谱 (Zhipu) 配置示例</b></summary>
|
||||||
|
|
@ -741,4 +935,5 @@ Discord: [https://discord.gg/V4sAZ9XWpN](https://discord.gg/V4sAZ9XWpN)
|
||||||
| **OpenRouter** | 200K tokens/月 | 多模型聚合 (Claude, GPT-4 等) |
|
| **OpenRouter** | 200K tokens/月 | 多模型聚合 (Claude, GPT-4 等) |
|
||||||
| **智谱 (Zhipu)** | 200K tokens/月 | 最适合中国用户 |
|
| **智谱 (Zhipu)** | 200K tokens/月 | 最适合中国用户 |
|
||||||
| **Brave Search** | 2000 次查询/月 | 网络搜索功能 |
|
| **Brave Search** | 2000 次查询/月 | 网络搜索功能 |
|
||||||
| **Groq** | 提供免费层级 | 极速推理 (Llama, Mixtral) |
|
| **Groq** | 提供免费层级 | 极速推理 (Llama, Mixtral) |
|
||||||
|
| **Cerebras** | 提供免费层级 | 极速推理 (Llama, Qwen 等) |
|
||||||
181
cmd/picoclaw/cmd_agent.go
Normal file
181
cmd/picoclaw/cmd_agent.go
Normal file
|
|
@ -0,0 +1,181 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/chzyer/readline"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/agent"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func agentCmd() {
|
||||||
|
message := ""
|
||||||
|
sessionKey := "cli:default"
|
||||||
|
modelOverride := ""
|
||||||
|
|
||||||
|
args := os.Args[2:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--debug", "-d":
|
||||||
|
logger.SetLevel(logger.DEBUG)
|
||||||
|
fmt.Println("🔍 Debug mode enabled")
|
||||||
|
case "-m", "--message":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
message = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-s", "--session":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
sessionKey = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--model", "-model":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
modelOverride = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelOverride != "" {
|
||||||
|
cfg.Agents.Defaults.Model = modelOverride
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := providers.CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error creating provider: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
// Use the resolved model ID from provider creation
|
||||||
|
if modelID != "" {
|
||||||
|
cfg.Agents.Defaults.Model = modelID
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
// Print agent startup info (only for interactive mode)
|
||||||
|
startupInfo := agentLoop.GetStartupInfo()
|
||||||
|
logger.InfoCF("agent", "Agent initialized",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tools_count": startupInfo["tools"].(map[string]interface{})["count"],
|
||||||
|
"skills_total": startupInfo["skills"].(map[string]interface{})["total"],
|
||||||
|
"skills_available": startupInfo["skills"].(map[string]interface{})["available"],
|
||||||
|
})
|
||||||
|
|
||||||
|
if message != "" {
|
||||||
|
ctx := context.Background()
|
||||||
|
response, err := agentLoop.ProcessDirect(ctx, message, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
fmt.Printf("\n%s %s\n", logo, response)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("%s Interactive mode (Ctrl+C to exit)\n\n", logo)
|
||||||
|
interactiveMode(agentLoop, sessionKey)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func interactiveMode(agentLoop *agent.AgentLoop, sessionKey string) {
|
||||||
|
prompt := fmt.Sprintf("%s You: ", logo)
|
||||||
|
|
||||||
|
rl, err := readline.NewEx(&readline.Config{
|
||||||
|
Prompt: prompt,
|
||||||
|
HistoryFile: filepath.Join(os.TempDir(), ".picoclaw_history"),
|
||||||
|
HistoryLimit: 100,
|
||||||
|
InterruptPrompt: "^C",
|
||||||
|
EOFPrompt: "exit",
|
||||||
|
})
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error initializing readline: %v\n", err)
|
||||||
|
fmt.Println("Falling back to simple input mode...")
|
||||||
|
simpleInteractiveMode(agentLoop, sessionKey)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer rl.Close()
|
||||||
|
|
||||||
|
for {
|
||||||
|
line, err := rl.Readline()
|
||||||
|
if err != nil {
|
||||||
|
if err == readline.ErrInterrupt || err == io.EOF {
|
||||||
|
fmt.Println("\nGoodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fmt.Printf("Error reading input: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
input := strings.TrimSpace(line)
|
||||||
|
if input == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if input == "exit" || input == "quit" {
|
||||||
|
fmt.Println("Goodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
response, err := agentLoop.ProcessDirect(ctx, input, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\n%s %s\n\n", logo, response)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func simpleInteractiveMode(agentLoop *agent.AgentLoop, sessionKey string) {
|
||||||
|
reader := bufio.NewReader(os.Stdin)
|
||||||
|
for {
|
||||||
|
fmt.Print(fmt.Sprintf("%s You: ", logo))
|
||||||
|
line, err := reader.ReadString('\n')
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
fmt.Println("\nGoodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fmt.Printf("Error reading input: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
input := strings.TrimSpace(line)
|
||||||
|
if input == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if input == "exit" || input == "quit" {
|
||||||
|
fmt.Println("Goodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
response, err := agentLoop.ProcessDirect(ctx, input, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\n%s %s\n\n", logo, response)
|
||||||
|
}
|
||||||
|
}
|
||||||
512
cmd/picoclaw/cmd_auth.go
Normal file
512
cmd/picoclaw/cmd_auth.go
Normal file
|
|
@ -0,0 +1,512 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const supportedProvidersMsg = "Supported providers: openai, anthropic, google-antigravity"
|
||||||
|
|
||||||
|
func authCmd() {
|
||||||
|
if len(os.Args) < 3 {
|
||||||
|
authHelp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch os.Args[2] {
|
||||||
|
case "login":
|
||||||
|
authLoginCmd()
|
||||||
|
case "logout":
|
||||||
|
authLogoutCmd()
|
||||||
|
case "status":
|
||||||
|
authStatusCmd()
|
||||||
|
case "models":
|
||||||
|
authModelsCmd()
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unknown auth command: %s\n", os.Args[2])
|
||||||
|
authHelp()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authHelp() {
|
||||||
|
fmt.Println("\nAuth commands:")
|
||||||
|
fmt.Println(" login Login via OAuth or paste token")
|
||||||
|
fmt.Println(" logout Remove stored credentials")
|
||||||
|
fmt.Println(" status Show current auth status")
|
||||||
|
fmt.Println(" models List available Antigravity models")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Login options:")
|
||||||
|
fmt.Println(" --provider <name> Provider to login with (openai, anthropic, google-antigravity)")
|
||||||
|
fmt.Println(" --device-code Use device code flow (for headless environments)")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Examples:")
|
||||||
|
fmt.Println(" picoclaw auth login --provider openai")
|
||||||
|
fmt.Println(" picoclaw auth login --provider openai --device-code")
|
||||||
|
fmt.Println(" picoclaw auth login --provider anthropic")
|
||||||
|
fmt.Println(" picoclaw auth login --provider google-antigravity")
|
||||||
|
fmt.Println(" picoclaw auth models")
|
||||||
|
fmt.Println(" picoclaw auth logout --provider openai")
|
||||||
|
fmt.Println(" picoclaw auth status")
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginCmd() {
|
||||||
|
provider := ""
|
||||||
|
useDeviceCode := false
|
||||||
|
|
||||||
|
args := os.Args[3:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--provider", "-p":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
provider = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--device-code":
|
||||||
|
useDeviceCode = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider == "" {
|
||||||
|
fmt.Println("Error: --provider is required")
|
||||||
|
fmt.Println(supportedProvidersMsg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch provider {
|
||||||
|
case "openai":
|
||||||
|
authLoginOpenAI(useDeviceCode)
|
||||||
|
case "anthropic":
|
||||||
|
authLoginPasteToken(provider)
|
||||||
|
case "google-antigravity", "antigravity":
|
||||||
|
authLoginGoogleAntigravity()
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unsupported provider: %s\n", provider)
|
||||||
|
fmt.Println(supportedProvidersMsg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginOpenAI(useDeviceCode bool) {
|
||||||
|
cfg := auth.OpenAIOAuthConfig()
|
||||||
|
|
||||||
|
var cred *auth.AuthCredential
|
||||||
|
var err error
|
||||||
|
|
||||||
|
if useDeviceCode {
|
||||||
|
cred, err = auth.LoginDeviceCode(cfg)
|
||||||
|
} else {
|
||||||
|
cred, err = auth.LoginBrowser(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Login failed: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential("openai", cred); err != nil {
|
||||||
|
fmt.Printf("Failed to save credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Update Providers (legacy format)
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = "oauth"
|
||||||
|
|
||||||
|
// Update or add openai in ModelList
|
||||||
|
foundOpenAI := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||||
|
foundOpenAI = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no openai in ModelList, add it
|
||||||
|
if !foundOpenAI {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update default model to use OpenAI
|
||||||
|
appCfg.Agents.Defaults.Model = "gpt-5.2"
|
||||||
|
|
||||||
|
if err := config.SaveConfig(getConfigPath(), appCfg); err != nil {
|
||||||
|
fmt.Printf("Warning: could not update config: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Login successful!")
|
||||||
|
if cred.AccountID != "" {
|
||||||
|
fmt.Printf("Account: %s\n", cred.AccountID)
|
||||||
|
}
|
||||||
|
fmt.Println("Default model set to: gpt-5.2")
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginGoogleAntigravity() {
|
||||||
|
cfg := auth.GoogleAntigravityOAuthConfig()
|
||||||
|
|
||||||
|
cred, err := auth.LoginBrowser(cfg)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Login failed: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
cred.Provider = "google-antigravity"
|
||||||
|
|
||||||
|
// Fetch user email from Google userinfo
|
||||||
|
email, err := fetchGoogleUserEmail(cred.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Warning: could not fetch email: %v\n", err)
|
||||||
|
} else {
|
||||||
|
cred.Email = email
|
||||||
|
fmt.Printf("Email: %s\n", email)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fetch Cloud Code Assist project ID
|
||||||
|
projectID, err := providers.FetchAntigravityProjectID(cred.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Warning: could not fetch project ID: %v\n", err)
|
||||||
|
fmt.Println("You may need Google Cloud Code Assist enabled on your account.")
|
||||||
|
} else {
|
||||||
|
cred.ProjectID = projectID
|
||||||
|
fmt.Printf("Project: %s\n", projectID)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential("google-antigravity", cred); err != nil {
|
||||||
|
fmt.Printf("Failed to save credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Update Providers (legacy format, for backward compatibility)
|
||||||
|
appCfg.Providers.Antigravity.AuthMethod = "oauth"
|
||||||
|
|
||||||
|
// Update or add antigravity in ModelList
|
||||||
|
foundAntigravity := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isAntigravityModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||||
|
foundAntigravity = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no antigravity in ModelList, add it
|
||||||
|
if !foundAntigravity {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gemini-flash",
|
||||||
|
Model: "antigravity/gemini-3-flash",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update default model
|
||||||
|
appCfg.Agents.Defaults.Model = "gemini-flash"
|
||||||
|
|
||||||
|
if err := config.SaveConfig(getConfigPath(), appCfg); err != nil {
|
||||||
|
fmt.Printf("Warning: could not update config: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\n✓ Google Antigravity login successful!")
|
||||||
|
fmt.Println("Default model set to: gemini-flash")
|
||||||
|
fmt.Println("Try it: picoclaw agent -m \"Hello world\"")
|
||||||
|
}
|
||||||
|
|
||||||
|
func fetchGoogleUserEmail(accessToken string) (string, error) {
|
||||||
|
req, err := http.NewRequest("GET", "https://www.googleapis.com/oauth2/v2/userinfo", nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 10 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("userinfo request failed: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
var userInfo struct {
|
||||||
|
Email string `json:"email"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &userInfo); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return userInfo.Email, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginPasteToken(provider string) {
|
||||||
|
cred, err := auth.LoginPasteToken(provider, os.Stdin)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Login failed: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential(provider, cred); err != nil {
|
||||||
|
fmt.Printf("Failed to save credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
switch provider {
|
||||||
|
case "anthropic":
|
||||||
|
appCfg.Providers.Anthropic.AuthMethod = "token"
|
||||||
|
// Update ModelList
|
||||||
|
found := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "token"
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "claude-sonnet-4.6",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
AuthMethod: "token",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
// Update default model
|
||||||
|
appCfg.Agents.Defaults.Model = "claude-sonnet-4.6"
|
||||||
|
case "openai":
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = "token"
|
||||||
|
// Update ModelList
|
||||||
|
found := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "token"
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "token",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
// Update default model
|
||||||
|
appCfg.Agents.Defaults.Model = "gpt-5.2"
|
||||||
|
}
|
||||||
|
if err := config.SaveConfig(getConfigPath(), appCfg); err != nil {
|
||||||
|
fmt.Printf("Warning: could not update config: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Token saved for %s!\n", provider)
|
||||||
|
fmt.Printf("Default model set to: %s\n", appCfg.Agents.Defaults.Model)
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLogoutCmd() {
|
||||||
|
provider := ""
|
||||||
|
|
||||||
|
args := os.Args[3:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--provider", "-p":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
provider = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider != "" {
|
||||||
|
if err := auth.DeleteCredential(provider); err != nil {
|
||||||
|
fmt.Printf("Failed to remove credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Clear AuthMethod in ModelList
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
switch provider {
|
||||||
|
case "openai":
|
||||||
|
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
case "anthropic":
|
||||||
|
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
case "google-antigravity", "antigravity":
|
||||||
|
if isAntigravityModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Clear AuthMethod in Providers (legacy)
|
||||||
|
switch provider {
|
||||||
|
case "openai":
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = ""
|
||||||
|
case "anthropic":
|
||||||
|
appCfg.Providers.Anthropic.AuthMethod = ""
|
||||||
|
case "google-antigravity", "antigravity":
|
||||||
|
appCfg.Providers.Antigravity.AuthMethod = ""
|
||||||
|
}
|
||||||
|
config.SaveConfig(getConfigPath(), appCfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Logged out from %s\n", provider)
|
||||||
|
} else {
|
||||||
|
if err := auth.DeleteAllCredentials(); err != nil {
|
||||||
|
fmt.Printf("Failed to remove credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Clear all AuthMethods in ModelList
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
// Clear all AuthMethods in Providers (legacy)
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = ""
|
||||||
|
appCfg.Providers.Anthropic.AuthMethod = ""
|
||||||
|
appCfg.Providers.Antigravity.AuthMethod = ""
|
||||||
|
config.SaveConfig(getConfigPath(), appCfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Logged out from all providers")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authStatusCmd() {
|
||||||
|
store, err := auth.LoadStore()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading auth store: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(store.Credentials) == 0 {
|
||||||
|
fmt.Println("No authenticated providers.")
|
||||||
|
fmt.Println("Run: picoclaw auth login --provider <name>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\nAuthenticated Providers:")
|
||||||
|
fmt.Println("------------------------")
|
||||||
|
for provider, cred := range store.Credentials {
|
||||||
|
status := "active"
|
||||||
|
if cred.IsExpired() {
|
||||||
|
status = "expired"
|
||||||
|
} else if cred.NeedsRefresh() {
|
||||||
|
status = "needs refresh"
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf(" %s:\n", provider)
|
||||||
|
fmt.Printf(" Method: %s\n", cred.AuthMethod)
|
||||||
|
fmt.Printf(" Status: %s\n", status)
|
||||||
|
if cred.AccountID != "" {
|
||||||
|
fmt.Printf(" Account: %s\n", cred.AccountID)
|
||||||
|
}
|
||||||
|
if cred.Email != "" {
|
||||||
|
fmt.Printf(" Email: %s\n", cred.Email)
|
||||||
|
}
|
||||||
|
if cred.ProjectID != "" {
|
||||||
|
fmt.Printf(" Project: %s\n", cred.ProjectID)
|
||||||
|
}
|
||||||
|
if !cred.ExpiresAt.IsZero() {
|
||||||
|
fmt.Printf(" Expires: %s\n", cred.ExpiresAt.Format("2006-01-02 15:04"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authModelsCmd() {
|
||||||
|
cred, err := auth.GetCredential("google-antigravity")
|
||||||
|
if err != nil || cred == nil {
|
||||||
|
fmt.Println("Not logged in to Google Antigravity.")
|
||||||
|
fmt.Println("Run: picoclaw auth login --provider google-antigravity")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh token if needed
|
||||||
|
if cred.NeedsRefresh() && cred.RefreshToken != "" {
|
||||||
|
oauthCfg := auth.GoogleAntigravityOAuthConfig()
|
||||||
|
refreshed, refreshErr := auth.RefreshAccessToken(cred, oauthCfg)
|
||||||
|
if refreshErr == nil {
|
||||||
|
cred = refreshed
|
||||||
|
_ = auth.SetCredential("google-antigravity", cred)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
projectID := cred.ProjectID
|
||||||
|
if projectID == "" {
|
||||||
|
fmt.Println("No project ID stored. Try logging in again.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Fetching models for project: %s\n\n", projectID)
|
||||||
|
|
||||||
|
models, err := providers.FetchAntigravityModels(cred.AccessToken, projectID)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error fetching models: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(models) == 0 {
|
||||||
|
fmt.Println("No models available.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Available Antigravity Models:")
|
||||||
|
fmt.Println("-----------------------------")
|
||||||
|
for _, m := range models {
|
||||||
|
status := "✓"
|
||||||
|
if m.IsExhausted {
|
||||||
|
status = "✗ (quota exhausted)"
|
||||||
|
}
|
||||||
|
name := m.ID
|
||||||
|
if m.DisplayName != "" {
|
||||||
|
name = fmt.Sprintf("%s (%s)", m.ID, m.DisplayName)
|
||||||
|
}
|
||||||
|
fmt.Printf(" %s %s\n", status, name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isAntigravityModel checks if a model string belongs to antigravity provider
|
||||||
|
func isAntigravityModel(model string) bool {
|
||||||
|
return model == "antigravity" ||
|
||||||
|
model == "google-antigravity" ||
|
||||||
|
strings.HasPrefix(model, "antigravity/") ||
|
||||||
|
strings.HasPrefix(model, "google-antigravity/")
|
||||||
|
}
|
||||||
|
|
||||||
|
// isOpenAIModel checks if a model string belongs to openai provider
|
||||||
|
func isOpenAIModel(model string) bool {
|
||||||
|
return model == "openai" ||
|
||||||
|
strings.HasPrefix(model, "openai/")
|
||||||
|
}
|
||||||
|
|
||||||
|
// isAnthropicModel checks if a model string belongs to anthropic provider
|
||||||
|
func isAnthropicModel(model string) bool {
|
||||||
|
return model == "anthropic" ||
|
||||||
|
strings.HasPrefix(model, "anthropic/")
|
||||||
|
}
|
||||||
227
cmd/picoclaw/cmd_cron.go
Normal file
227
cmd/picoclaw/cmd_cron.go
Normal file
|
|
@ -0,0 +1,227 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
|
)
|
||||||
|
|
||||||
|
func cronCmd() {
|
||||||
|
if len(os.Args) < 3 {
|
||||||
|
cronHelp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
subcommand := os.Args[2]
|
||||||
|
|
||||||
|
// Load config to get workspace path
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cronStorePath := filepath.Join(cfg.WorkspacePath(), "cron", "jobs.json")
|
||||||
|
|
||||||
|
switch subcommand {
|
||||||
|
case "list":
|
||||||
|
cronListCmd(cronStorePath)
|
||||||
|
case "add":
|
||||||
|
cronAddCmd(cronStorePath)
|
||||||
|
case "remove":
|
||||||
|
if len(os.Args) < 4 {
|
||||||
|
fmt.Println("Usage: picoclaw cron remove <job_id>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cronRemoveCmd(cronStorePath, os.Args[3])
|
||||||
|
case "enable":
|
||||||
|
cronEnableCmd(cronStorePath, false)
|
||||||
|
case "disable":
|
||||||
|
cronEnableCmd(cronStorePath, true)
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unknown cron command: %s\n", subcommand)
|
||||||
|
cronHelp()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronHelp() {
|
||||||
|
fmt.Println("\nCron commands:")
|
||||||
|
fmt.Println(" list List all scheduled jobs")
|
||||||
|
fmt.Println(" add Add a new scheduled job")
|
||||||
|
fmt.Println(" remove <id> Remove a job by ID")
|
||||||
|
fmt.Println(" enable <id> Enable a job")
|
||||||
|
fmt.Println(" disable <id> Disable a job")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Add options:")
|
||||||
|
fmt.Println(" -n, --name Job name")
|
||||||
|
fmt.Println(" -m, --message Message for agent")
|
||||||
|
fmt.Println(" -e, --every Run every N seconds")
|
||||||
|
fmt.Println(" -c, --cron Cron expression (e.g. '0 9 * * *')")
|
||||||
|
fmt.Println(" -d, --deliver Deliver response to channel")
|
||||||
|
fmt.Println(" --to Recipient for delivery")
|
||||||
|
fmt.Println(" --channel Channel for delivery")
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronListCmd(storePath string) {
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
jobs := cs.ListJobs(true) // Show all jobs, including disabled
|
||||||
|
|
||||||
|
if len(jobs) == 0 {
|
||||||
|
fmt.Println("No scheduled jobs.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\nScheduled Jobs:")
|
||||||
|
fmt.Println("----------------")
|
||||||
|
for _, job := range jobs {
|
||||||
|
var schedule string
|
||||||
|
if job.Schedule.Kind == "every" && job.Schedule.EveryMS != nil {
|
||||||
|
schedule = fmt.Sprintf("every %ds", *job.Schedule.EveryMS/1000)
|
||||||
|
} else if job.Schedule.Kind == "cron" {
|
||||||
|
schedule = job.Schedule.Expr
|
||||||
|
} else {
|
||||||
|
schedule = "one-time"
|
||||||
|
}
|
||||||
|
|
||||||
|
nextRun := "scheduled"
|
||||||
|
if job.State.NextRunAtMS != nil {
|
||||||
|
nextTime := time.UnixMilli(*job.State.NextRunAtMS)
|
||||||
|
nextRun = nextTime.Format("2006-01-02 15:04")
|
||||||
|
}
|
||||||
|
|
||||||
|
status := "enabled"
|
||||||
|
if !job.Enabled {
|
||||||
|
status = "disabled"
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf(" %s (%s)\n", job.Name, job.ID)
|
||||||
|
fmt.Printf(" Schedule: %s\n", schedule)
|
||||||
|
fmt.Printf(" Status: %s\n", status)
|
||||||
|
fmt.Printf(" Next run: %s\n", nextRun)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronAddCmd(storePath string) {
|
||||||
|
name := ""
|
||||||
|
message := ""
|
||||||
|
var everySec *int64
|
||||||
|
cronExpr := ""
|
||||||
|
deliver := false
|
||||||
|
channel := ""
|
||||||
|
to := ""
|
||||||
|
|
||||||
|
args := os.Args[3:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "-n", "--name":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
name = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-m", "--message":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
message = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-e", "--every":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
var sec int64
|
||||||
|
fmt.Sscanf(args[i+1], "%d", &sec)
|
||||||
|
everySec = &sec
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-c", "--cron":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
cronExpr = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-d", "--deliver":
|
||||||
|
deliver = true
|
||||||
|
case "--to":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
to = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--channel":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
channel = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if name == "" {
|
||||||
|
fmt.Println("Error: --name is required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if message == "" {
|
||||||
|
fmt.Println("Error: --message is required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if everySec == nil && cronExpr == "" {
|
||||||
|
fmt.Println("Error: Either --every or --cron must be specified")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var schedule cron.CronSchedule
|
||||||
|
if everySec != nil {
|
||||||
|
everyMS := *everySec * 1000
|
||||||
|
schedule = cron.CronSchedule{
|
||||||
|
Kind: "every",
|
||||||
|
EveryMS: &everyMS,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
schedule = cron.CronSchedule{
|
||||||
|
Kind: "cron",
|
||||||
|
Expr: cronExpr,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
job, err := cs.AddJob(name, schedule, message, deliver, channel, to)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error adding job: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("✓ Added job '%s' (%s)\n", job.Name, job.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronRemoveCmd(storePath, jobID string) {
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
if cs.RemoveJob(jobID) {
|
||||||
|
fmt.Printf("✓ Removed job %s\n", jobID)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("✗ Job %s not found\n", jobID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronEnableCmd(storePath string, disable bool) {
|
||||||
|
if len(os.Args) < 4 {
|
||||||
|
fmt.Println("Usage: picoclaw cron enable/disable <job_id>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
jobID := os.Args[3]
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
enabled := !disable
|
||||||
|
|
||||||
|
job := cs.EnableJob(jobID, enabled)
|
||||||
|
if job != nil {
|
||||||
|
status := "enabled"
|
||||||
|
if disable {
|
||||||
|
status = "disabled"
|
||||||
|
}
|
||||||
|
fmt.Printf("✓ Job '%s' %s\n", job.Name, status)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("✗ Job %s not found\n", jobID)
|
||||||
|
}
|
||||||
|
}
|
||||||
223
cmd/picoclaw/cmd_gateway.go
Normal file
223
cmd/picoclaw/cmd_gateway.go
Normal file
|
|
@ -0,0 +1,223 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"os/signal"
|
||||||
|
"path/filepath"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/agent"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/devices"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/health"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/heartbeat"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/state"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/voice"
|
||||||
|
)
|
||||||
|
|
||||||
|
func gatewayCmd() {
|
||||||
|
// Check for --debug flag
|
||||||
|
args := os.Args[2:]
|
||||||
|
for _, arg := range args {
|
||||||
|
if arg == "--debug" || arg == "-d" {
|
||||||
|
logger.SetLevel(logger.DEBUG)
|
||||||
|
fmt.Println("🔍 Debug mode enabled")
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := providers.CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error creating provider: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
// Use the resolved model ID from provider creation
|
||||||
|
if modelID != "" {
|
||||||
|
cfg.Agents.Defaults.Model = modelID
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
// Print agent startup info
|
||||||
|
fmt.Println("\n📦 Agent Status:")
|
||||||
|
startupInfo := agentLoop.GetStartupInfo()
|
||||||
|
toolsInfo := startupInfo["tools"].(map[string]interface{})
|
||||||
|
skillsInfo := startupInfo["skills"].(map[string]interface{})
|
||||||
|
fmt.Printf(" • Tools: %d loaded\n", toolsInfo["count"])
|
||||||
|
fmt.Printf(" • Skills: %d/%d available\n",
|
||||||
|
skillsInfo["available"],
|
||||||
|
skillsInfo["total"])
|
||||||
|
|
||||||
|
// Log to file as well
|
||||||
|
logger.InfoCF("agent", "Agent initialized",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tools_count": toolsInfo["count"],
|
||||||
|
"skills_total": skillsInfo["total"],
|
||||||
|
"skills_available": skillsInfo["available"],
|
||||||
|
})
|
||||||
|
|
||||||
|
// Setup cron tool and service
|
||||||
|
execTimeout := time.Duration(cfg.Tools.Cron.ExecTimeoutMinutes) * time.Minute
|
||||||
|
cronService := setupCronTool(agentLoop, msgBus, cfg.WorkspacePath(), cfg.Agents.Defaults.RestrictToWorkspace, execTimeout, cfg)
|
||||||
|
|
||||||
|
heartbeatService := heartbeat.NewHeartbeatService(
|
||||||
|
cfg.WorkspacePath(),
|
||||||
|
cfg.Heartbeat.Interval,
|
||||||
|
cfg.Heartbeat.Enabled,
|
||||||
|
)
|
||||||
|
heartbeatService.SetBus(msgBus)
|
||||||
|
heartbeatService.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult {
|
||||||
|
// Use cli:direct as fallback if no valid channel
|
||||||
|
if channel == "" || chatID == "" {
|
||||||
|
channel, chatID = "cli", "direct"
|
||||||
|
}
|
||||||
|
// Use ProcessHeartbeat - no session history, each heartbeat is independent
|
||||||
|
response, err := agentLoop.ProcessHeartbeat(context.Background(), prompt, channel, chatID)
|
||||||
|
if err != nil {
|
||||||
|
return tools.ErrorResult(fmt.Sprintf("Heartbeat error: %v", err))
|
||||||
|
}
|
||||||
|
if response == "HEARTBEAT_OK" {
|
||||||
|
return tools.SilentResult("Heartbeat OK")
|
||||||
|
}
|
||||||
|
// For heartbeat, always return silent - the subagent result will be
|
||||||
|
// sent to user via processSystemMessage when the async task completes
|
||||||
|
return tools.SilentResult(response)
|
||||||
|
})
|
||||||
|
|
||||||
|
channelManager, err := channels.NewManager(cfg, msgBus)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error creating channel manager: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Inject channel manager into agent loop for command handling
|
||||||
|
agentLoop.SetChannelManager(channelManager)
|
||||||
|
|
||||||
|
var transcriber *voice.GroqTranscriber
|
||||||
|
if cfg.Providers.Groq.APIKey != "" {
|
||||||
|
transcriber = voice.NewGroqTranscriber(cfg.Providers.Groq.APIKey)
|
||||||
|
logger.InfoC("voice", "Groq voice transcription enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
if transcriber != nil {
|
||||||
|
if telegramChannel, ok := channelManager.GetChannel("telegram"); ok {
|
||||||
|
if tc, ok := telegramChannel.(*channels.TelegramChannel); ok {
|
||||||
|
tc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "Groq transcription attached to Telegram channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if discordChannel, ok := channelManager.GetChannel("discord"); ok {
|
||||||
|
if dc, ok := discordChannel.(*channels.DiscordChannel); ok {
|
||||||
|
dc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "Groq transcription attached to Discord channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if slackChannel, ok := channelManager.GetChannel("slack"); ok {
|
||||||
|
if sc, ok := slackChannel.(*channels.SlackChannel); ok {
|
||||||
|
sc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "Groq transcription attached to Slack channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
enabledChannels := channelManager.GetEnabledChannels()
|
||||||
|
if len(enabledChannels) > 0 {
|
||||||
|
fmt.Printf("✓ Channels enabled: %s\n", enabledChannels)
|
||||||
|
} else {
|
||||||
|
fmt.Println("⚠ Warning: No channels enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("✓ Gateway started on %s:%d\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
fmt.Println("Press Ctrl+C to stop")
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := cronService.Start(); err != nil {
|
||||||
|
fmt.Printf("Error starting cron service: %v\n", err)
|
||||||
|
}
|
||||||
|
fmt.Println("✓ Cron service started")
|
||||||
|
|
||||||
|
if err := heartbeatService.Start(); err != nil {
|
||||||
|
fmt.Printf("Error starting heartbeat service: %v\n", err)
|
||||||
|
}
|
||||||
|
fmt.Println("✓ Heartbeat service started")
|
||||||
|
|
||||||
|
stateManager := state.NewManager(cfg.WorkspacePath())
|
||||||
|
deviceService := devices.NewService(devices.Config{
|
||||||
|
Enabled: cfg.Devices.Enabled,
|
||||||
|
MonitorUSB: cfg.Devices.MonitorUSB,
|
||||||
|
}, stateManager)
|
||||||
|
deviceService.SetBus(msgBus)
|
||||||
|
if err := deviceService.Start(ctx); err != nil {
|
||||||
|
fmt.Printf("Error starting device service: %v\n", err)
|
||||||
|
} else if cfg.Devices.Enabled {
|
||||||
|
fmt.Println("✓ Device event service started")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := channelManager.StartAll(ctx); err != nil {
|
||||||
|
fmt.Printf("Error starting channels: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
healthServer := health.NewServer(cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
go func() {
|
||||||
|
if err := healthServer.Start(); err != nil && err != http.ErrServerClosed {
|
||||||
|
logger.ErrorCF("health", "Health server error", map[string]interface{}{"error": err.Error()})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
fmt.Printf("✓ Health endpoints available at http://%s:%d/health and /ready\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
|
||||||
|
go agentLoop.Run(ctx)
|
||||||
|
|
||||||
|
sigChan := make(chan os.Signal, 1)
|
||||||
|
signal.Notify(sigChan, os.Interrupt)
|
||||||
|
<-sigChan
|
||||||
|
|
||||||
|
fmt.Println("\nShutting down...")
|
||||||
|
cancel()
|
||||||
|
healthServer.Stop(context.Background())
|
||||||
|
deviceService.Stop()
|
||||||
|
heartbeatService.Stop()
|
||||||
|
cronService.Stop()
|
||||||
|
agentLoop.Stop()
|
||||||
|
channelManager.StopAll(ctx)
|
||||||
|
fmt.Println("✓ Gateway stopped")
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupCronTool(agentLoop *agent.AgentLoop, msgBus *bus.MessageBus, workspace string, restrict bool, execTimeout time.Duration, cfg *config.Config) *cron.CronService {
|
||||||
|
cronStorePath := filepath.Join(workspace, "cron", "jobs.json")
|
||||||
|
|
||||||
|
// Create cron service
|
||||||
|
cronService := cron.NewCronService(cronStorePath, nil)
|
||||||
|
|
||||||
|
// Create and register CronTool
|
||||||
|
cronTool := tools.NewCronTool(cronService, agentLoop, msgBus, workspace, restrict, execTimeout, cfg)
|
||||||
|
agentLoop.RegisterTool(cronTool)
|
||||||
|
|
||||||
|
// Set the onJob handler
|
||||||
|
cronService.SetOnJob(func(job *cron.CronJob) (string, error) {
|
||||||
|
result := cronTool.ExecuteJob(context.Background(), job)
|
||||||
|
return result, nil
|
||||||
|
})
|
||||||
|
|
||||||
|
return cronService
|
||||||
|
}
|
||||||
81
cmd/picoclaw/cmd_migrate.go
Normal file
81
cmd/picoclaw/cmd_migrate.go
Normal file
|
|
@ -0,0 +1,81 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/migrate"
|
||||||
|
)
|
||||||
|
|
||||||
|
func migrateCmd() {
|
||||||
|
if len(os.Args) > 2 && (os.Args[2] == "--help" || os.Args[2] == "-h") {
|
||||||
|
migrateHelp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := migrate.Options{}
|
||||||
|
|
||||||
|
args := os.Args[2:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--dry-run":
|
||||||
|
opts.DryRun = true
|
||||||
|
case "--config-only":
|
||||||
|
opts.ConfigOnly = true
|
||||||
|
case "--workspace-only":
|
||||||
|
opts.WorkspaceOnly = true
|
||||||
|
case "--force":
|
||||||
|
opts.Force = true
|
||||||
|
case "--refresh":
|
||||||
|
opts.Refresh = true
|
||||||
|
case "--openclaw-home":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
opts.OpenClawHome = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--picoclaw-home":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
opts.PicoClawHome = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unknown flag: %s\n", args[i])
|
||||||
|
migrateHelp()
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := migrate.Run(opts)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !opts.DryRun {
|
||||||
|
migrate.PrintSummary(result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func migrateHelp() {
|
||||||
|
fmt.Println("\nMigrate from OpenClaw to PicoClaw")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Usage: picoclaw migrate [options]")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Options:")
|
||||||
|
fmt.Println(" --dry-run Show what would be migrated without making changes")
|
||||||
|
fmt.Println(" --refresh Re-sync workspace files from OpenClaw (repeatable)")
|
||||||
|
fmt.Println(" --config-only Only migrate config, skip workspace files")
|
||||||
|
fmt.Println(" --workspace-only Only migrate workspace files, skip config")
|
||||||
|
fmt.Println(" --force Skip confirmation prompts")
|
||||||
|
fmt.Println(" --openclaw-home Override OpenClaw home directory (default: ~/.openclaw)")
|
||||||
|
fmt.Println(" --picoclaw-home Override PicoClaw home directory (default: ~/.picoclaw)")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Examples:")
|
||||||
|
fmt.Println(" picoclaw migrate Detect and migrate from OpenClaw")
|
||||||
|
fmt.Println(" picoclaw migrate --dry-run Show what would be migrated")
|
||||||
|
fmt.Println(" picoclaw migrate --refresh Re-sync workspace files")
|
||||||
|
fmt.Println(" picoclaw migrate --force Migrate without confirmation")
|
||||||
|
}
|
||||||
108
cmd/picoclaw/cmd_onboard.go
Normal file
108
cmd/picoclaw/cmd_onboard.go
Normal file
|
|
@ -0,0 +1,108 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"embed"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
//go:generate cp -r ../../workspace .
|
||||||
|
//go:embed workspace
|
||||||
|
var embeddedFiles embed.FS
|
||||||
|
|
||||||
|
func onboard() {
|
||||||
|
configPath := getConfigPath()
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err == nil {
|
||||||
|
fmt.Printf("Config already exists at %s\n", configPath)
|
||||||
|
fmt.Print("Overwrite? (y/n): ")
|
||||||
|
var response string
|
||||||
|
fmt.Scanln(&response)
|
||||||
|
if response != "y" {
|
||||||
|
fmt.Println("Aborted.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
fmt.Printf("Error saving config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
createWorkspaceTemplates(workspace)
|
||||||
|
|
||||||
|
fmt.Printf("%s picoclaw is ready!\n", logo)
|
||||||
|
fmt.Println("\nNext steps:")
|
||||||
|
fmt.Println(" 1. Add your API key to", configPath)
|
||||||
|
fmt.Println("")
|
||||||
|
fmt.Println(" Recommended:")
|
||||||
|
fmt.Println(" - OpenRouter: https://openrouter.ai/keys (access 100+ models)")
|
||||||
|
fmt.Println(" - Ollama: https://ollama.com (local, free)")
|
||||||
|
fmt.Println("")
|
||||||
|
fmt.Println(" See README.md for 17+ supported providers.")
|
||||||
|
fmt.Println("")
|
||||||
|
fmt.Println(" 2. Chat: picoclaw agent -m \"Hello!\"")
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyEmbeddedToTarget(targetDir string) error {
|
||||||
|
// Ensure target directory exists
|
||||||
|
if err := os.MkdirAll(targetDir, 0755); err != nil {
|
||||||
|
return fmt.Errorf("Failed to create target directory: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Walk through all files in embed.FS
|
||||||
|
err := fs.WalkDir(embeddedFiles, "workspace", func(path string, d fs.DirEntry, err error) error {
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip directories
|
||||||
|
if d.IsDir() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read embedded file
|
||||||
|
data, err := embeddedFiles.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("Failed to read embedded file %s: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
new_path, err := filepath.Rel("workspace", path)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("Failed to get relative path for %s: %v\n", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build target file path
|
||||||
|
targetPath := filepath.Join(targetDir, new_path)
|
||||||
|
|
||||||
|
// Ensure target file's directory exists
|
||||||
|
if err := os.MkdirAll(filepath.Dir(targetPath), 0755); err != nil {
|
||||||
|
return fmt.Errorf("Failed to create directory %s: %w", filepath.Dir(targetPath), err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write file
|
||||||
|
if err := os.WriteFile(targetPath, data, 0644); err != nil {
|
||||||
|
return fmt.Errorf("Failed to write file %s: %w", targetPath, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func createWorkspaceTemplates(workspace string) {
|
||||||
|
err := copyEmbeddedToTarget(workspace)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error copying workspace templates: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
305
cmd/picoclaw/cmd_skills.go
Normal file
305
cmd/picoclaw/cmd_skills.go
Normal file
|
|
@ -0,0 +1,305 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
func skillsHelp() {
|
||||||
|
fmt.Println("\nSkills commands:")
|
||||||
|
fmt.Println(" list List installed skills")
|
||||||
|
fmt.Println(" install <repo> Install skill from GitHub")
|
||||||
|
fmt.Println(" install-builtin Install all builtin skills to workspace")
|
||||||
|
fmt.Println(" list-builtin List available builtin skills")
|
||||||
|
fmt.Println(" remove <name> Remove installed skill")
|
||||||
|
fmt.Println(" search Search available skills")
|
||||||
|
fmt.Println(" show <name> Show skill details")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Examples:")
|
||||||
|
fmt.Println(" picoclaw skills list")
|
||||||
|
fmt.Println(" picoclaw skills install sipeed/picoclaw-skills/weather")
|
||||||
|
fmt.Println(" picoclaw skills install-builtin")
|
||||||
|
fmt.Println(" picoclaw skills list-builtin")
|
||||||
|
fmt.Println(" picoclaw skills remove weather")
|
||||||
|
fmt.Println(" picoclaw skills install --registry clawhub github")
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsListCmd(loader *skills.SkillsLoader) {
|
||||||
|
allSkills := loader.ListSkills()
|
||||||
|
|
||||||
|
if len(allSkills) == 0 {
|
||||||
|
fmt.Println("No skills installed.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\nInstalled Skills:")
|
||||||
|
fmt.Println("------------------")
|
||||||
|
for _, skill := range allSkills {
|
||||||
|
fmt.Printf(" ✓ %s (%s)\n", skill.Name, skill.Source)
|
||||||
|
if skill.Description != "" {
|
||||||
|
fmt.Printf(" %s\n", skill.Description)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsInstallCmd(installer *skills.SkillInstaller, cfg *config.Config) {
|
||||||
|
if len(os.Args) < 4 {
|
||||||
|
fmt.Println("Usage: picoclaw skills install <github-repo>")
|
||||||
|
fmt.Println(" picoclaw skills install --registry <name> <slug>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for --registry flag.
|
||||||
|
if os.Args[3] == "--registry" {
|
||||||
|
if len(os.Args) < 6 {
|
||||||
|
fmt.Println("Usage: picoclaw skills install --registry <name> <slug>")
|
||||||
|
fmt.Println("Example: picoclaw skills install --registry clawhub github")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
registryName := os.Args[4]
|
||||||
|
slug := os.Args[5]
|
||||||
|
skillsInstallFromRegistry(cfg, registryName, slug)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default: install from GitHub (backward compatible).
|
||||||
|
repo := os.Args[3]
|
||||||
|
fmt.Printf("Installing skill from %s...\n", repo)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := installer.InstallFromGitHub(ctx, repo); err != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to install skill: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\u2713 Skill '%s' installed successfully!\n", filepath.Base(repo))
|
||||||
|
}
|
||||||
|
|
||||||
|
// skillsInstallFromRegistry installs a skill from a named registry (e.g. clawhub).
|
||||||
|
func skillsInstallFromRegistry(cfg *config.Config, registryName, slug string) {
|
||||||
|
err := utils.ValidateSkillIdentifier(registryName)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("\u2717 Invalid registry name: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = utils.ValidateSkillIdentifier(slug)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("\u2717 Invalid slug: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Installing skill '%s' from %s registry...\n", slug, registryName)
|
||||||
|
|
||||||
|
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
|
||||||
|
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
|
||||||
|
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
|
||||||
|
})
|
||||||
|
|
||||||
|
registry := registryMgr.GetRegistry(registryName)
|
||||||
|
if registry == nil {
|
||||||
|
fmt.Printf("\u2717 Registry '%s' not found or not enabled. Check your config.json.\n", registryName)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
targetDir := filepath.Join(workspace, "skills", slug)
|
||||||
|
|
||||||
|
if _, err := os.Stat(targetDir); err == nil {
|
||||||
|
fmt.Printf("\u2717 Skill '%s' already installed at %s\n", slug, targetDir)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := os.MkdirAll(filepath.Join(workspace, "skills"), 0755); err != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to create skills directory: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := registry.DownloadAndInstall(ctx, slug, "", targetDir)
|
||||||
|
if err != nil {
|
||||||
|
rmErr := os.RemoveAll(targetDir)
|
||||||
|
if rmErr != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to remove partial install: %v\n", rmErr)
|
||||||
|
}
|
||||||
|
fmt.Printf("\u2717 Failed to install skill: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.IsMalwareBlocked {
|
||||||
|
rmErr := os.RemoveAll(targetDir)
|
||||||
|
if rmErr != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to remove partial install: %v\n", rmErr)
|
||||||
|
}
|
||||||
|
fmt.Printf("\u2717 Skill '%s' is flagged as malicious and cannot be installed.\n", slug)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.IsSuspicious {
|
||||||
|
fmt.Printf("\u26a0\ufe0f Warning: skill '%s' is flagged as suspicious.\n", slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\u2713 Skill '%s' v%s installed successfully!\n", slug, result.Version)
|
||||||
|
if result.Summary != "" {
|
||||||
|
fmt.Printf(" %s\n", result.Summary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsRemoveCmd(installer *skills.SkillInstaller, skillName string) {
|
||||||
|
fmt.Printf("Removing skill '%s'...\n", skillName)
|
||||||
|
|
||||||
|
if err := installer.Uninstall(skillName); err != nil {
|
||||||
|
fmt.Printf("✗ Failed to remove skill: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("✓ Skill '%s' removed successfully!\n", skillName)
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsInstallBuiltinCmd(workspace string) {
|
||||||
|
builtinSkillsDir := "./picoclaw/skills"
|
||||||
|
workspaceSkillsDir := filepath.Join(workspace, "skills")
|
||||||
|
|
||||||
|
fmt.Printf("Copying builtin skills to workspace...\n")
|
||||||
|
|
||||||
|
skillsToInstall := []string{
|
||||||
|
"weather",
|
||||||
|
"news",
|
||||||
|
"stock",
|
||||||
|
"calculator",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, skillName := range skillsToInstall {
|
||||||
|
builtinPath := filepath.Join(builtinSkillsDir, skillName)
|
||||||
|
workspacePath := filepath.Join(workspaceSkillsDir, skillName)
|
||||||
|
|
||||||
|
if _, err := os.Stat(builtinPath); err != nil {
|
||||||
|
fmt.Printf("⊘ Builtin skill '%s' not found: %v\n", skillName, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.MkdirAll(workspacePath, 0755); err != nil {
|
||||||
|
fmt.Printf("✗ Failed to create directory for %s: %v\n", skillName, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := copyDirectory(builtinPath, workspacePath); err != nil {
|
||||||
|
fmt.Printf("✗ Failed to copy %s: %v\n", skillName, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\n✓ All builtin skills installed!")
|
||||||
|
fmt.Println("Now you can use them in your workspace.")
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsListBuiltinCmd() {
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
builtinSkillsDir := filepath.Join(filepath.Dir(cfg.WorkspacePath()), "picoclaw", "skills")
|
||||||
|
|
||||||
|
fmt.Println("\nAvailable Builtin Skills:")
|
||||||
|
fmt.Println("-----------------------")
|
||||||
|
|
||||||
|
entries, err := os.ReadDir(builtinSkillsDir)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error reading builtin skills: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(entries) == 0 {
|
||||||
|
fmt.Println("No builtin skills available.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, entry := range entries {
|
||||||
|
if entry.IsDir() {
|
||||||
|
skillName := entry.Name()
|
||||||
|
skillFile := filepath.Join(builtinSkillsDir, skillName, "SKILL.md")
|
||||||
|
|
||||||
|
description := "No description"
|
||||||
|
if _, err := os.Stat(skillFile); err == nil {
|
||||||
|
data, err := os.ReadFile(skillFile)
|
||||||
|
if err == nil {
|
||||||
|
content := string(data)
|
||||||
|
if idx := strings.Index(content, "\n"); idx > 0 {
|
||||||
|
firstLine := content[:idx]
|
||||||
|
if strings.Contains(firstLine, "description:") {
|
||||||
|
descLine := strings.Index(content[idx:], "\n")
|
||||||
|
if descLine > 0 {
|
||||||
|
description = strings.TrimSpace(content[idx+descLine : idx+descLine])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
status := "✓"
|
||||||
|
fmt.Printf(" %s %s\n", status, entry.Name())
|
||||||
|
if description != "" {
|
||||||
|
fmt.Printf(" %s\n", description)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsSearchCmd(installer *skills.SkillInstaller) {
|
||||||
|
fmt.Println("Searching for available skills...")
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
availableSkills, err := installer.ListAvailableSkills(ctx)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("✗ Failed to fetch skills list: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(availableSkills) == 0 {
|
||||||
|
fmt.Println("No skills available.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\nAvailable Skills (%d):\n", len(availableSkills))
|
||||||
|
fmt.Println("--------------------")
|
||||||
|
for _, skill := range availableSkills {
|
||||||
|
fmt.Printf(" 📦 %s\n", skill.Name)
|
||||||
|
fmt.Printf(" %s\n", skill.Description)
|
||||||
|
fmt.Printf(" Repo: %s\n", skill.Repository)
|
||||||
|
if skill.Author != "" {
|
||||||
|
fmt.Printf(" Author: %s\n", skill.Author)
|
||||||
|
}
|
||||||
|
if len(skill.Tags) > 0 {
|
||||||
|
fmt.Printf(" Tags: %v\n", skill.Tags)
|
||||||
|
}
|
||||||
|
fmt.Println()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsShowCmd(loader *skills.SkillsLoader, skillName string) {
|
||||||
|
content, ok := loader.LoadSkill(skillName)
|
||||||
|
if !ok {
|
||||||
|
fmt.Printf("✗ Skill '%s' not found\n", skillName)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\n📦 Skill: %s\n", skillName)
|
||||||
|
fmt.Println("----------------------")
|
||||||
|
fmt.Println(content)
|
||||||
|
}
|
||||||
102
cmd/picoclaw/cmd_status.go
Normal file
102
cmd/picoclaw/cmd_status.go
Normal file
|
|
@ -0,0 +1,102 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
)
|
||||||
|
|
||||||
|
func statusCmd() {
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
configPath := getConfigPath()
|
||||||
|
|
||||||
|
fmt.Printf("%s picoclaw Status\n", logo)
|
||||||
|
fmt.Printf("Version: %s\n", formatVersion())
|
||||||
|
build, _ := formatBuildInfo()
|
||||||
|
if build != "" {
|
||||||
|
fmt.Printf("Build: %s\n", build)
|
||||||
|
}
|
||||||
|
fmt.Println()
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err == nil {
|
||||||
|
fmt.Println("Config:", configPath, "✓")
|
||||||
|
} else {
|
||||||
|
fmt.Println("Config:", configPath, "✗")
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
if _, err := os.Stat(workspace); err == nil {
|
||||||
|
fmt.Println("Workspace:", workspace, "✓")
|
||||||
|
} else {
|
||||||
|
fmt.Println("Workspace:", workspace, "✗")
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err == nil {
|
||||||
|
fmt.Printf("Model: %s\n", cfg.Agents.Defaults.Model)
|
||||||
|
|
||||||
|
hasOpenRouter := cfg.Providers.OpenRouter.APIKey != ""
|
||||||
|
hasAnthropic := cfg.Providers.Anthropic.APIKey != ""
|
||||||
|
hasOpenAI := cfg.Providers.OpenAI.APIKey != ""
|
||||||
|
hasGemini := cfg.Providers.Gemini.APIKey != ""
|
||||||
|
hasZhipu := cfg.Providers.Zhipu.APIKey != ""
|
||||||
|
hasQwen := cfg.Providers.Qwen.APIKey != ""
|
||||||
|
hasGroq := cfg.Providers.Groq.APIKey != ""
|
||||||
|
hasVLLM := cfg.Providers.VLLM.APIBase != ""
|
||||||
|
hasMoonshot := cfg.Providers.Moonshot.APIKey != ""
|
||||||
|
hasDeepSeek := cfg.Providers.DeepSeek.APIKey != ""
|
||||||
|
hasVolcEngine := cfg.Providers.VolcEngine.APIKey != ""
|
||||||
|
hasNvidia := cfg.Providers.Nvidia.APIKey != ""
|
||||||
|
hasOllama := cfg.Providers.Ollama.APIBase != ""
|
||||||
|
|
||||||
|
status := func(enabled bool) string {
|
||||||
|
if enabled {
|
||||||
|
return "✓"
|
||||||
|
}
|
||||||
|
return "not set"
|
||||||
|
}
|
||||||
|
fmt.Println("OpenRouter API:", status(hasOpenRouter))
|
||||||
|
fmt.Println("Anthropic API:", status(hasAnthropic))
|
||||||
|
fmt.Println("OpenAI API:", status(hasOpenAI))
|
||||||
|
fmt.Println("Gemini API:", status(hasGemini))
|
||||||
|
fmt.Println("Zhipu API:", status(hasZhipu))
|
||||||
|
fmt.Println("Qwen API:", status(hasQwen))
|
||||||
|
fmt.Println("Groq API:", status(hasGroq))
|
||||||
|
fmt.Println("Moonshot API:", status(hasMoonshot))
|
||||||
|
fmt.Println("DeepSeek API:", status(hasDeepSeek))
|
||||||
|
fmt.Println("VolcEngine API:", status(hasVolcEngine))
|
||||||
|
fmt.Println("Nvidia API:", status(hasNvidia))
|
||||||
|
if hasVLLM {
|
||||||
|
fmt.Printf("vLLM/Local: ✓ %s\n", cfg.Providers.VLLM.APIBase)
|
||||||
|
} else {
|
||||||
|
fmt.Println("vLLM/Local: not set")
|
||||||
|
}
|
||||||
|
if hasOllama {
|
||||||
|
fmt.Printf("Ollama: ✓ %s\n", cfg.Providers.Ollama.APIBase)
|
||||||
|
} else {
|
||||||
|
fmt.Println("Ollama: not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
store, _ := auth.LoadStore()
|
||||||
|
if store != nil && len(store.Credentials) > 0 {
|
||||||
|
fmt.Println("\nOAuth/Token Auth:")
|
||||||
|
for provider, cred := range store.Credentials {
|
||||||
|
status := "authenticated"
|
||||||
|
if cred.IsExpired() {
|
||||||
|
status = "expired"
|
||||||
|
} else if cred.NeedsRefresh() {
|
||||||
|
status = "needs refresh"
|
||||||
|
}
|
||||||
|
fmt.Printf(" %s (%s): %s\n", provider, cred.AuthMethod, status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
1230
cmd/picoclaw/main.go
1230
cmd/picoclaw/main.go
File diff suppressed because it is too large
Load diff
|
|
@ -3,12 +3,48 @@
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"restrict_to_workspace": true,
|
"restrict_to_workspace": true,
|
||||||
"model": "glm-4.7",
|
"model": "gpt4",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key",
|
||||||
|
"api_base": "https://api.anthropic.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gemini",
|
||||||
|
"model": "antigravity/gemini-2.0-flash",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "deepseek",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-your-deepseek-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "loadbalanced-gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key1",
|
||||||
|
"api_base": "https://api1.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "loadbalanced-gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key2",
|
||||||
|
"api_base": "https://api2.example.com/v1"
|
||||||
|
}
|
||||||
|
],
|
||||||
"channels": {
|
"channels": {
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
@ -21,6 +57,13 @@
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"token": "YOUR_DISCORD_BOT_TOKEN",
|
"token": "YOUR_DISCORD_BOT_TOKEN",
|
||||||
|
"allow_from": [],
|
||||||
|
"mention_only": false
|
||||||
|
},
|
||||||
|
"qq": {
|
||||||
|
"enabled": false,
|
||||||
|
"app_id": "YOUR_QQ_APP_ID",
|
||||||
|
"app_secret": "YOUR_QQ_APP_SECRET",
|
||||||
"allow_from": []
|
"allow_from": []
|
||||||
},
|
},
|
||||||
"maixcam": {
|
"maixcam": {
|
||||||
|
|
@ -73,6 +116,7 @@
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
"providers": {
|
||||||
|
"_comment": "DEPRECATED: Use model_list instead. This will be removed in a future version",
|
||||||
"anthropic": {
|
"anthropic": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": ""
|
"api_base": ""
|
||||||
|
|
@ -111,9 +155,21 @@
|
||||||
"api_key": "sk-xxx",
|
"api_key": "sk-xxx",
|
||||||
"api_base": ""
|
"api_base": ""
|
||||||
},
|
},
|
||||||
|
"qwen": {
|
||||||
|
"api_key": "sk-xxx",
|
||||||
|
"api_base": ""
|
||||||
|
},
|
||||||
"ollama": {
|
"ollama": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": "http://localhost:11434/v1"
|
"api_base": "http://localhost:11434/v1"
|
||||||
|
},
|
||||||
|
"cerebras": {
|
||||||
|
"api_key": "",
|
||||||
|
"api_base": ""
|
||||||
|
},
|
||||||
|
"volcengine": {
|
||||||
|
"api_key": "",
|
||||||
|
"api_base": ""
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"tools": {
|
"tools": {
|
||||||
|
|
@ -126,6 +182,10 @@
|
||||||
"query_param": "query for ollama, q for others",
|
"query_param": "query for ollama, q for others",
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
},
|
},
|
||||||
|
"duckduckgo": {
|
||||||
|
"enabled": true,
|
||||||
|
"max_results": 5
|
||||||
|
},
|
||||||
"perplexity": {
|
"perplexity": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "pplx-xxx",
|
"api_key": "pplx-xxx",
|
||||||
|
|
@ -134,6 +194,21 @@
|
||||||
},
|
},
|
||||||
"cron": {
|
"cron": {
|
||||||
"exec_timeout_minutes": 5
|
"exec_timeout_minutes": 5
|
||||||
|
},
|
||||||
|
"exec": {
|
||||||
|
"enable_deny_patterns": false,
|
||||||
|
"custom_deny_patterns": []
|
||||||
|
},
|
||||||
|
"skills": {
|
||||||
|
"registries": {
|
||||||
|
"clawhub": {
|
||||||
|
"enabled": true,
|
||||||
|
"base_url": "https://clawhub.ai",
|
||||||
|
"search_path": "/api/v1/search",
|
||||||
|
"skills_path": "/api/v1/skills",
|
||||||
|
"download_path": "/api/v1/download"
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"heartbeat": {
|
"heartbeat": {
|
||||||
|
|
|
||||||
1002
docs/ANTIGRAVITY_AUTH.md
Normal file
1002
docs/ANTIGRAVITY_AUTH.md
Normal file
File diff suppressed because it is too large
Load diff
72
docs/ANTIGRAVITY_USAGE.md
Normal file
72
docs/ANTIGRAVITY_USAGE.md
Normal file
|
|
@ -0,0 +1,72 @@
|
||||||
|
# Using Antigravity Provider in PicoClaw
|
||||||
|
|
||||||
|
This guide explains how to set up and use the **Antigravity** (Google Cloud Code Assist) provider in PicoClaw.
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
1. A Google account.
|
||||||
|
2. Google Cloud Code Assist enabled (usually available via the "Gemini for Google Cloud" onboarding).
|
||||||
|
|
||||||
|
## 1. Authentication
|
||||||
|
|
||||||
|
To authenticate with Antigravity, run the following command:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth login --provider antigravity
|
||||||
|
```
|
||||||
|
|
||||||
|
### Manual Authentication (Headless/VPS)
|
||||||
|
If you are running on a server (Coolify/Docker) and cannot reach `localhost`, follow these steps:
|
||||||
|
1. Run the command above.
|
||||||
|
2. Copy the URL provided and open it in your local browser.
|
||||||
|
3. Complete the login.
|
||||||
|
4. Your browser will redirect to a `localhost:51121` URL (which will fail to load).
|
||||||
|
5. **Copy that final URL** from your browser's address bar.
|
||||||
|
6. **Paste it back into the terminal** where PicoClaw is waiting.
|
||||||
|
|
||||||
|
PicoClaw will extract the authorization code and complete the process automatically.
|
||||||
|
|
||||||
|
## 2. Managing Models
|
||||||
|
|
||||||
|
### List Available Models
|
||||||
|
To see which models your project has access to and check their quotas:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth models
|
||||||
|
```
|
||||||
|
|
||||||
|
### Switch Models
|
||||||
|
You can change the default model in `~/.picoclaw/config.json` or override it via the CLI:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Override for a single command
|
||||||
|
picoclaw agent -m "Hello" --model claude-opus-4-6-thinking
|
||||||
|
```
|
||||||
|
|
||||||
|
## 3. Real-world Usage (Coolify/Docker)
|
||||||
|
|
||||||
|
If you are deploying via Coolify or Docker, follow these steps to test:
|
||||||
|
|
||||||
|
1. **Branch**: Use the `feat/antigravity-provider` branch.
|
||||||
|
2. **Environment Variables**:
|
||||||
|
* `PICOCLAW_AGENTS_DEFAULTS_PROVIDER=antigravity`
|
||||||
|
* `PICOCLAW_AGENTS_DEFAULTS_MODEL=gemini-3-flash`
|
||||||
|
3. **Authentication persistence**:
|
||||||
|
If you've logged in locally, you can copy your credentials to the server:
|
||||||
|
```bash
|
||||||
|
scp ~/.picoclaw/auth-profiles.json user@your-server:~/.picoclaw/
|
||||||
|
```
|
||||||
|
*Alternatively*, run the `auth login` command once on the server if you have terminal access.
|
||||||
|
|
||||||
|
## 4. Troubleshooting
|
||||||
|
|
||||||
|
* **Empty Response**: If a model returns an empty reply, it may be restricted for your project. Try `gemini-3-flash` or `claude-opus-4-6-thinking`.
|
||||||
|
* **429 Rate Limit**: Antigravity has strict quotas. PicoClaw will display the "reset time" in the error message if you hit a limit.
|
||||||
|
* **404 Not Found**: Ensure you are using a model ID from the `picoclaw auth models` list. Use the short ID (e.g., `gemini-3-flash`) not the full path.
|
||||||
|
|
||||||
|
## 5. Summary of Working Models
|
||||||
|
|
||||||
|
Based on testing, the following models are most reliable:
|
||||||
|
* `gemini-3-flash` (Fast, highly available)
|
||||||
|
* `gemini-2.5-flash-lite` (Lightweight)
|
||||||
|
* `claude-opus-4-6-thinking` (Powerful, includes reasoning)
|
||||||
179
docs/design/provider-refactoring-tests.md
Normal file
179
docs/design/provider-refactoring-tests.md
Normal file
|
|
@ -0,0 +1,179 @@
|
||||||
|
# Provider Architecture Refactoring - Test Suite Summary
|
||||||
|
|
||||||
|
> PRD: `tasks/prd-provider-refactoring.md`
|
||||||
|
|
||||||
|
This document summarizes the complete test suite designed for the Provider architecture refactoring.
|
||||||
|
|
||||||
|
## Test File Structure
|
||||||
|
|
||||||
|
```
|
||||||
|
pkg/
|
||||||
|
├── config/
|
||||||
|
│ ├── model_config_test.go # US-001, US-002: ModelConfig struct and GetModelConfig tests
|
||||||
|
│ └── migration_test.go # US-003: Backward compatibility and migration tests
|
||||||
|
├── providers/
|
||||||
|
│ ├── registry_test.go # US-006: Load balancing tests
|
||||||
|
│ ├── integration_test.go # E2E integration tests
|
||||||
|
│ └── factory/
|
||||||
|
│ └── factory_test.go # US-004, US-005: Provider factory tests
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Test Case Checklist
|
||||||
|
|
||||||
|
### 1. `pkg/config/model_config_test.go` - Configuration Parsing Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestModelConfig_Parsing` | Verify ModelConfig JSON parsing | US-001 |
|
||||||
|
| `TestModelConfig_ModelListInConfig` | Verify model_list parsing in Config | US-001 |
|
||||||
|
| `TestModelConfig_Validation` | Verify required field validation | US-001 |
|
||||||
|
| `TestConfig_GetModelConfig_Found` | Verify GetModelConfig finds model | US-002 |
|
||||||
|
| `TestConfig_GetModelConfig_NotFound` | Verify GetModelConfig returns error | US-002 |
|
||||||
|
| `TestConfig_GetModelConfig_EmptyModelList` | Verify empty model_list handling | US-002 |
|
||||||
|
| `TestConfig_BackwardCompatibility_ProvidersToModelList` | Verify old config conversion | US-003 |
|
||||||
|
| `TestConfig_DeprecationWarning` | Verify deprecation warning | US-003 |
|
||||||
|
| `TestModelConfig_ProtocolExtraction` | Verify protocol prefix extraction | US-004 |
|
||||||
|
| `TestConfig_ModelNameUniqueness` | Verify model_name uniqueness | US-001 |
|
||||||
|
|
||||||
|
### 2. `pkg/config/migration_test.go` - Migration Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestConvertProvidersToModelList_OpenAI` | OpenAI config conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_Anthropic` | Anthropic config conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_MultipleProviders` | Multiple provider conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_EmptyProviders` | Empty providers handling | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_GitHubCopilot` | GitHub Copilot conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_Antigravity` | Antigravity conversion | US-003 |
|
||||||
|
| `TestGenerateModelName_*` | Model name generation | US-003 |
|
||||||
|
| `TestHasProvidersConfig_*` | Detect old config existence | US-003 |
|
||||||
|
| `TestValidateMigration_*` | Migration validation | US-003 |
|
||||||
|
| `TestMigrateConfig_DryRun` | Dry run migration | US-003 |
|
||||||
|
| `TestMigrateConfig_Actual` | Actual migration | US-003 |
|
||||||
|
|
||||||
|
### 3. `pkg/providers/registry_test.go` - Load Balancing Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestModelRegistry_SingleConfig` | Single config returns same result | US-006 |
|
||||||
|
| `TestModelRegistry_RoundRobinSelection` | 3-config round-robin selection | US-006 |
|
||||||
|
| `TestModelRegistry_RoundRobinTwoConfigs` | 2-config round-robin selection | US-006 |
|
||||||
|
| `TestModelRegistry_ConcurrentAccess` | Concurrent access thread safety | US-006 |
|
||||||
|
| `TestModelRegistry_RaceDetection` | Data race detection | US-006 |
|
||||||
|
| `TestModelRegistry_ModelNotFound` | Model not found error | US-006 |
|
||||||
|
| `TestModelRegistry_EmptyRegistry` | Empty registry handling | US-006 |
|
||||||
|
| `TestModelRegistry_MultipleModels` | Multiple model registration | US-006 |
|
||||||
|
| `TestModelRegistry_MixedSingleAndMultiple` | Single/multiple config mix | US-006 |
|
||||||
|
| `TestModelRegistry_CaseSensitiveModelNames` | Case sensitivity | US-006 |
|
||||||
|
|
||||||
|
### 4. `pkg/providers/factory/factory_test.go` - Provider Factory Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestCreateProviderFromConfig_OpenAI` | Create OpenAI provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_OpenAIDefault` | Default openai protocol | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_Anthropic` | Create Anthropic provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_Antigravity` | Create Antigravity provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_ClaudeCLI` | Create Claude CLI provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_CodexCLI` | Create Codex CLI provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_GitHubCopilot` | Create GitHub Copilot provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_UnknownProtocol` | Unknown protocol error handling | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_MissingAPIKey` | Missing API key error | US-004 |
|
||||||
|
| `TestExtractProtocol` | Protocol prefix extraction | US-004 |
|
||||||
|
| `TestCreateProvider_UsesModelList` | Create using model_list | US-005 |
|
||||||
|
| `TestCreateProvider_FallbackToProviders` | Fallback to providers | US-005 |
|
||||||
|
| `TestCreateProvider_PriorityModelListOverProviders` | model_list priority | US-005 |
|
||||||
|
|
||||||
|
### 5. `pkg/providers/integration_test.go` - E2E Integration Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestE2E_OpenAICompatibleProvider_NoCodeChange` | Zero-code provider addition | Goal |
|
||||||
|
| `TestE2E_LoadBalancing_RoundRobin` | Load balancing actual effect | US-006 |
|
||||||
|
| `TestE2E_BackwardCompatibility_OldProvidersConfig` | Old config compatibility | US-003 |
|
||||||
|
| `TestE2E_ErrorHandling_ModelNotFound` | Model not found | FR-30 |
|
||||||
|
| `TestE2E_ErrorHandling_MissingAPIKey` | Missing API key | FR-31 |
|
||||||
|
| `TestE2E_ErrorHandling_InvalidAPIBase` | Invalid API base | FR-30 |
|
||||||
|
| `TestE2E_ToolCalls_OpenAICompatible` | Tool call support | - |
|
||||||
|
| `TestE2E_AntigravityProvider` | Antigravity provider | US-004 |
|
||||||
|
| `TestE2E_ClaudeCLIProvider` | Claude CLI provider | US-004 |
|
||||||
|
|
||||||
|
### 6. Performance Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose |
|
||||||
|
|-----------|---------|
|
||||||
|
| `BenchmarkCreateProviderFromConfig` | Provider creation performance |
|
||||||
|
| `BenchmarkGetModelConfig` | Model lookup performance |
|
||||||
|
| `BenchmarkGetModelConfigParallel` | Concurrent lookup performance |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Running Tests
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Run all tests
|
||||||
|
go test ./pkg/... -v
|
||||||
|
|
||||||
|
# Run with data race detection
|
||||||
|
go test ./pkg/... -race
|
||||||
|
|
||||||
|
# Run specific package tests
|
||||||
|
go test ./pkg/config -v
|
||||||
|
go test ./pkg/providers -v
|
||||||
|
go test ./pkg/providers/factory -v
|
||||||
|
|
||||||
|
# Run E2E tests
|
||||||
|
go test ./pkg/providers -run TestE2E -v
|
||||||
|
|
||||||
|
# Run performance tests
|
||||||
|
go test ./pkg/providers -bench=. -benchmem
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## PRD Acceptance Criteria Mapping
|
||||||
|
|
||||||
|
| PRD Acceptance Criteria | Test Cases |
|
||||||
|
|------------------------|------------|
|
||||||
|
| US-001: Add ModelConfig struct | `TestModelConfig_Parsing`, `TestModelConfig_Validation` |
|
||||||
|
| US-001: model_name unique | `TestConfig_ModelNameUniqueness` |
|
||||||
|
| US-002: GetModelConfig method | `TestConfig_GetModelConfig_*` |
|
||||||
|
| US-003: Auto-convert providers | `TestConvertProvidersToModelList_*` |
|
||||||
|
| US-003: Deprecation warning | `TestConfig_DeprecationWarning` |
|
||||||
|
| US-003: Existing tests pass | (existing test files unchanged) |
|
||||||
|
| US-004: Protocol prefix factory | `TestExtractProtocol`, `TestCreateProviderFromConfig_*` |
|
||||||
|
| US-004: Default prefix openai | `TestCreateProviderFromConfig_OpenAIDefault` |
|
||||||
|
| US-005: CreateProvider uses factory | `TestCreateProvider_*` |
|
||||||
|
| US-006: Round-robin selection | `TestModelRegistry_RoundRobin*` |
|
||||||
|
| US-006: Thread-safe atomic | `TestModelRegistry_RaceDetection` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Recommended Implementation Order
|
||||||
|
|
||||||
|
1. **Phase 1: Configuration Structure** (US-001, US-002)
|
||||||
|
- Implement `ModelConfig` struct
|
||||||
|
- Implement `GetModelConfig` method
|
||||||
|
- Run `model_config_test.go`
|
||||||
|
|
||||||
|
2. **Phase 2: Protocol Factory** (US-004)
|
||||||
|
- Implement `CreateProviderFromConfig`
|
||||||
|
- Implement `ExtractProtocol`
|
||||||
|
- Run `factory_test.go`
|
||||||
|
|
||||||
|
3. **Phase 3: Load Balancing** (US-006)
|
||||||
|
- Implement `ModelRegistry`
|
||||||
|
- Implement round-robin selection
|
||||||
|
- Run `registry_test.go` (with `-race`)
|
||||||
|
|
||||||
|
4. **Phase 4: Backward Compatibility** (US-003, US-005)
|
||||||
|
- Implement `ConvertProvidersToModelList`
|
||||||
|
- Refactor `CreateProvider`
|
||||||
|
- Run `migration_test.go`
|
||||||
|
- Verify existing tests pass
|
||||||
|
|
||||||
|
5. **Phase 5: E2E Verification**
|
||||||
|
- Run `integration_test.go`
|
||||||
|
- Manual testing with `config.example.json`
|
||||||
334
docs/design/provider-refactoring.md
Normal file
334
docs/design/provider-refactoring.md
Normal file
|
|
@ -0,0 +1,334 @@
|
||||||
|
# Provider Architecture Refactoring Design
|
||||||
|
|
||||||
|
> Issue: #283
|
||||||
|
> Discussion: #122
|
||||||
|
> Branch: feat/refactor-provider-by-protocol
|
||||||
|
|
||||||
|
## 1. Current Problems
|
||||||
|
|
||||||
|
### 1.1 Configuration Structure Issues
|
||||||
|
|
||||||
|
**Current State**: Each Provider requires a predefined field in `ProvidersConfig`
|
||||||
|
|
||||||
|
```go
|
||||||
|
type ProvidersConfig struct {
|
||||||
|
Anthropic ProviderConfig `json:"anthropic"`
|
||||||
|
OpenAI ProviderConfig `json:"openai"`
|
||||||
|
DeepSeek ProviderConfig `json:"deepseek"`
|
||||||
|
Qwen ProviderConfig `json:"qwen"`
|
||||||
|
Cerebras ProviderConfig `json:"cerebras"`
|
||||||
|
VolcEngine ProviderConfig `json:"volcengine"`
|
||||||
|
// ... every new provider requires changes here
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Problems**:
|
||||||
|
- Adding a new Provider requires modifying Go code (struct definition)
|
||||||
|
- `CreateProvider` function in `http_provider.go` has 200+ lines of switch-case
|
||||||
|
- Most Providers are OpenAI-compatible, but code is duplicated
|
||||||
|
|
||||||
|
### 1.2 Code Bloat Trend
|
||||||
|
|
||||||
|
Recent PRs demonstrate this issue:
|
||||||
|
|
||||||
|
| PR | Provider | Code Changes |
|
||||||
|
|----|----------|--------------|
|
||||||
|
| #365 | Qwen | +17 lines to http_provider.go |
|
||||||
|
| #333 | Cerebras | +17 lines to http_provider.go |
|
||||||
|
| #368 | Volcengine | +18 lines to http_provider.go |
|
||||||
|
|
||||||
|
Each OpenAI-compatible Provider requires:
|
||||||
|
1. Modify `config.go` to add configuration field
|
||||||
|
2. Modify `http_provider.go` to add switch case
|
||||||
|
3. Update documentation
|
||||||
|
|
||||||
|
### 1.3 Agent-Provider Coupling
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "deepseek", // need to know provider name
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Problem: Agent needs to know both `provider` and `model`, adding complexity.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. New Approach: model_list
|
||||||
|
|
||||||
|
### 2.1 Core Principles
|
||||||
|
|
||||||
|
Inspired by [LiteLLM](https://docs.litellm.ai/docs/proxy/configs) design:
|
||||||
|
|
||||||
|
1. **Model-centric**: Users care about models, not providers
|
||||||
|
2. **Protocol prefix**: Use `protocol/model_name` format, e.g., `openai/gpt-5.2`, `anthropic/claude-sonnet-4.6`
|
||||||
|
3. **Configuration-driven**: Adding new Providers only requires config changes, no code changes
|
||||||
|
|
||||||
|
### 2.2 New Configuration Structure
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "openai/deepseek-chat",
|
||||||
|
"api_base": "https://api.deepseek.com/v1",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gemini-3-flash",
|
||||||
|
"model": "antigravity/gemini-3-flash",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "my-company-llm",
|
||||||
|
"model": "openai/company-model-v1",
|
||||||
|
"api_base": "https://llm.company.com/v1",
|
||||||
|
"api_key": "xxx"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"max_tokens": 8192,
|
||||||
|
"temperature": 0.7
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.3 Go Struct Definition
|
||||||
|
|
||||||
|
```go
|
||||||
|
type Config struct {
|
||||||
|
ModelList []ModelConfig `json:"model_list"` // new
|
||||||
|
Providers ProvidersConfig `json:"providers"` // old, deprecated
|
||||||
|
|
||||||
|
Agents AgentsConfig `json:"agents"`
|
||||||
|
Channels ChannelsConfig `json:"channels"`
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
|
||||||
|
type ModelConfig struct {
|
||||||
|
// Required
|
||||||
|
ModelName string `json:"model_name"` // user-facing name (alias)
|
||||||
|
Model string `json:"model"` // protocol/model, e.g., openai/gpt-5.2
|
||||||
|
|
||||||
|
// Common config
|
||||||
|
APIBase string `json:"api_base,omitempty"`
|
||||||
|
APIKey string `json:"api_key,omitempty"`
|
||||||
|
Proxy string `json:"proxy,omitempty"`
|
||||||
|
|
||||||
|
// Special provider config
|
||||||
|
AuthMethod string `json:"auth_method,omitempty"` // oauth, token
|
||||||
|
ConnectMode string `json:"connect_mode,omitempty"` // stdio, grpc
|
||||||
|
|
||||||
|
// Optional optimizations
|
||||||
|
RPM int `json:"rpm,omitempty"` // rate limit
|
||||||
|
MaxTokensField string `json:"max_tokens_field,omitempty"` // max_tokens or max_completion_tokens
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.4 Protocol Recognition
|
||||||
|
|
||||||
|
Identify protocol via prefix in `model` field:
|
||||||
|
|
||||||
|
| Prefix | Protocol | Description |
|
||||||
|
|--------|----------|-------------|
|
||||||
|
| `openai/` | OpenAI-compatible | Most common, includes DeepSeek, Qwen, Groq, etc. |
|
||||||
|
| `anthropic/` | Anthropic | Claude series specific |
|
||||||
|
| `antigravity/` | Antigravity | Google Cloud Code Assist |
|
||||||
|
| `gemini/` | Gemini | Google Gemini native API (if needed) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Design Rationale
|
||||||
|
|
||||||
|
### 3.1 Problems Solved
|
||||||
|
|
||||||
|
| Problem | Old Approach | New Approach |
|
||||||
|
|---------|--------------|--------------|
|
||||||
|
| Add OpenAI-compatible Provider | Change 3 code locations | Add one config entry |
|
||||||
|
| Agent specifies model | Need provider + model | Only need model |
|
||||||
|
| Code duplication | Each Provider duplicates logic | Share protocol implementation |
|
||||||
|
| Multi-Agent support | Complex | Naturally compatible |
|
||||||
|
|
||||||
|
### 3.2 Multi-Agent Compatibility
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [...],
|
||||||
|
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
},
|
||||||
|
"coder": {
|
||||||
|
"model": "gpt-5.2",
|
||||||
|
"system_prompt": "You are a coding assistant..."
|
||||||
|
},
|
||||||
|
"translator": {
|
||||||
|
"model": "claude-sonnet-4.6"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Each Agent only needs to specify `model` (corresponds to `model_name` in `model_list`).
|
||||||
|
|
||||||
|
### 3.3 Industry Comparison
|
||||||
|
|
||||||
|
**LiteLLM** (most mature open-source LLM Proxy) uses similar design:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
- model_name: gpt-4o
|
||||||
|
litellm_params:
|
||||||
|
model: openai/gpt-5.2
|
||||||
|
api_key: xxx
|
||||||
|
- model_name: my-custom
|
||||||
|
litellm_params:
|
||||||
|
model: openai/custom-model
|
||||||
|
api_base: https://my-api.com/v1
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Migration Plan
|
||||||
|
|
||||||
|
### 4.1 Phase 1: Compatibility Period (v1.x)
|
||||||
|
|
||||||
|
Support both `providers` and `model_list`:
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (c *Config) GetModelConfig(modelName string) (*ModelConfig, error) {
|
||||||
|
// Prefer new config
|
||||||
|
if len(c.ModelList) > 0 {
|
||||||
|
return c.findModelByName(modelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Backward compatibility with old config
|
||||||
|
if !c.Providers.IsEmpty() {
|
||||||
|
logger.Warn("'providers' config is deprecated, please migrate to 'model_list'")
|
||||||
|
return c.convertFromProviders(modelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("model %s not found", modelName)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.2 Phase 2: Warning Period (late v1.x)
|
||||||
|
|
||||||
|
- Print more prominent warnings at startup
|
||||||
|
- Provide automatic migration script
|
||||||
|
- Mark `providers` as deprecated in documentation
|
||||||
|
|
||||||
|
### 4.3 Phase 3: Removal Period (v2.0)
|
||||||
|
|
||||||
|
- Completely remove `providers` support
|
||||||
|
- Remove `agents.defaults.provider` field
|
||||||
|
- Only support `model_list`
|
||||||
|
|
||||||
|
### 4.4 Configuration Migration Example
|
||||||
|
|
||||||
|
**Old Config**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"deepseek": {
|
||||||
|
"api_key": "sk-xxx",
|
||||||
|
"api_base": "https://api.deepseek.com/v1"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "deepseek",
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**New Config**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "openai/deepseek-chat",
|
||||||
|
"api_base": "https://api.deepseek.com/v1",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Implementation Checklist
|
||||||
|
|
||||||
|
### 5.1 Configuration Layer
|
||||||
|
|
||||||
|
- [ ] Add `ModelConfig` struct
|
||||||
|
- [ ] Add `Config.ModelList` field
|
||||||
|
- [ ] Implement `GetModelConfig(modelName)` method
|
||||||
|
- [ ] Implement old config compatibility conversion
|
||||||
|
- [ ] Add `model_name` uniqueness validation
|
||||||
|
|
||||||
|
### 5.2 Provider Layer
|
||||||
|
|
||||||
|
- [ ] Create `pkg/providers/factory/` directory
|
||||||
|
- [ ] Implement `CreateProviderFromModelConfig()`
|
||||||
|
- [ ] Refactor `http_provider.go` to `openai/provider.go`
|
||||||
|
- [ ] Maintain backward compatibility for old `CreateProvider()`
|
||||||
|
|
||||||
|
### 5.3 Testing
|
||||||
|
|
||||||
|
- [ ] New config unit tests
|
||||||
|
- [ ] Old config compatibility tests
|
||||||
|
- [ ] Integration tests
|
||||||
|
|
||||||
|
### 5.4 Documentation
|
||||||
|
|
||||||
|
- [ ] Update README
|
||||||
|
- [ ] Update config.example.json
|
||||||
|
- [ ] Write migration guide
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Risks and Mitigations
|
||||||
|
|
||||||
|
| Risk | Mitigation |
|
||||||
|
|------|------------|
|
||||||
|
| Breaking existing configs | Compatibility period keeps old config working |
|
||||||
|
| User migration cost | Provide automatic migration script |
|
||||||
|
| Special Provider incompatibility | Keep `auth_method` and other extension fields |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. References
|
||||||
|
|
||||||
|
- [LiteLLM Config Documentation](https://docs.litellm.ai/docs/proxy/configs)
|
||||||
|
- [One-API GitHub](https://github.com/songquanpeng/one-api)
|
||||||
|
- Discussion #122: Refactor Provider Architecture
|
||||||
211
docs/migration/model-list-migration.md
Normal file
211
docs/migration/model-list-migration.md
Normal file
|
|
@ -0,0 +1,211 @@
|
||||||
|
# Migration Guide: From `providers` to `model_list`
|
||||||
|
|
||||||
|
This guide explains how to migrate from the legacy `providers` configuration to the new `model_list` format.
|
||||||
|
|
||||||
|
## Why Migrate?
|
||||||
|
|
||||||
|
The new `model_list` configuration offers several advantages:
|
||||||
|
|
||||||
|
- **Zero-code provider addition**: Add OpenAI-compatible providers with configuration only
|
||||||
|
- **Load balancing**: Configure multiple endpoints for the same model
|
||||||
|
- **Protocol-based routing**: Use prefixes like `openai/`, `anthropic/`, etc.
|
||||||
|
- **Cleaner configuration**: Model-centric instead of vendor-centric
|
||||||
|
|
||||||
|
## Timeline
|
||||||
|
|
||||||
|
| Version | Status |
|
||||||
|
|---------|--------|
|
||||||
|
| v1.x | `model_list` introduced, `providers` deprecated but functional |
|
||||||
|
| v1.x+1 | Prominent deprecation warnings, migration tool available |
|
||||||
|
| v2.0 | `providers` configuration removed |
|
||||||
|
|
||||||
|
## Before and After
|
||||||
|
|
||||||
|
### Before: Legacy `providers` Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"openai": {
|
||||||
|
"api_key": "sk-your-openai-key",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
"anthropic": {
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
"deepseek": {
|
||||||
|
"api_key": "sk-your-deepseek-key"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "openai",
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### After: New `model_list` Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "deepseek",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-your-deepseek-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt4"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Protocol Prefixes
|
||||||
|
|
||||||
|
The `model` field uses a protocol prefix format: `[protocol/]model-identifier`
|
||||||
|
|
||||||
|
| Prefix | Description | Example |
|
||||||
|
|--------|-------------|---------|
|
||||||
|
| `openai/` | OpenAI API (default) | `openai/gpt-5.2` |
|
||||||
|
| `anthropic/` | Anthropic API | `anthropic/claude-opus-4` |
|
||||||
|
| `antigravity/` | Google via Antigravity OAuth | `antigravity/gemini-2.0-flash` |
|
||||||
|
| `claude-cli/` | Claude CLI (local) | `claude-cli/claude-sonnet-4.6` |
|
||||||
|
| `codex-cli/` | Codex CLI (local) | `codex-cli/codex-4` |
|
||||||
|
| `github-copilot/` | GitHub Copilot | `github-copilot/gpt-4o` |
|
||||||
|
| `openrouter/` | OpenRouter | `openrouter/anthropic/claude-sonnet-4.6` |
|
||||||
|
| `groq/` | Groq API | `groq/llama-3.1-70b` |
|
||||||
|
| `deepseek/` | DeepSeek API | `deepseek/deepseek-chat` |
|
||||||
|
| `cerebras/` | Cerebras API | `cerebras/llama-3.3-70b` |
|
||||||
|
| `qwen/` | Alibaba Qwen | `qwen/qwen-max` |
|
||||||
|
|
||||||
|
**Note**: If no prefix is specified, `openai/` is used as the default.
|
||||||
|
|
||||||
|
## ModelConfig Fields
|
||||||
|
|
||||||
|
| Field | Required | Description |
|
||||||
|
|-------|----------|-------------|
|
||||||
|
| `model_name` | Yes | User-facing alias for the model |
|
||||||
|
| `model` | Yes | Protocol and model identifier (e.g., `openai/gpt-5.2`) |
|
||||||
|
| `api_base` | No | API endpoint URL |
|
||||||
|
| `api_key` | No* | API authentication key |
|
||||||
|
| `proxy` | No | HTTP proxy URL |
|
||||||
|
| `auth_method` | No | Authentication method: `oauth`, `token` |
|
||||||
|
| `connect_mode` | No | Connection mode for CLI providers: `stdio`, `grpc` |
|
||||||
|
| `rpm` | No | Requests per minute limit |
|
||||||
|
| `max_tokens_field` | No | Field name for max tokens |
|
||||||
|
|
||||||
|
*`api_key` is required for HTTP-based protocols unless `api_base` points to a local server.
|
||||||
|
|
||||||
|
## Load Balancing
|
||||||
|
|
||||||
|
Configure multiple endpoints for the same model to distribute load:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key1",
|
||||||
|
"api_base": "https://api1.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key2",
|
||||||
|
"api_base": "https://api2.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key3",
|
||||||
|
"api_base": "https://api3.example.com/v1"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
When you request model `gpt4`, requests will be distributed across all three endpoints using round-robin selection.
|
||||||
|
|
||||||
|
## Adding a New OpenAI-Compatible Provider
|
||||||
|
|
||||||
|
With `model_list`, adding a new provider requires zero code changes:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "my-custom-llm",
|
||||||
|
"model": "openai/my-model-v1",
|
||||||
|
"api_key": "your-api-key",
|
||||||
|
"api_base": "https://api.your-provider.com/v1"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Just specify `openai/` as the protocol (or omit it for the default), and provide your provider's API base URL.
|
||||||
|
|
||||||
|
## Backward Compatibility
|
||||||
|
|
||||||
|
During the migration period, your existing `providers` configuration will continue to work:
|
||||||
|
|
||||||
|
1. If `model_list` is empty and `providers` has data, the system auto-converts internally
|
||||||
|
2. A deprecation warning is logged: `"providers config is deprecated, please migrate to model_list"`
|
||||||
|
3. All existing functionality remains unchanged
|
||||||
|
|
||||||
|
## Migration Checklist
|
||||||
|
|
||||||
|
- [ ] Identify all providers you're currently using
|
||||||
|
- [ ] Create `model_list` entries for each provider
|
||||||
|
- [ ] Use appropriate protocol prefixes
|
||||||
|
- [ ] Update `agents.defaults.model` to reference the new `model_name`
|
||||||
|
- [ ] Test that all models work correctly
|
||||||
|
- [ ] Remove or comment out the old `providers` section
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### Model not found error
|
||||||
|
|
||||||
|
```
|
||||||
|
model "xxx" not found in model_list or providers
|
||||||
|
```
|
||||||
|
|
||||||
|
**Solution**: Ensure the `model_name` in `model_list` matches the value in `agents.defaults.model`.
|
||||||
|
|
||||||
|
### Unknown protocol error
|
||||||
|
|
||||||
|
```
|
||||||
|
unknown protocol "xxx" in model "xxx/model-name"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Solution**: Use a supported protocol prefix. See the [Protocol Prefixes](#protocol-prefixes) table above.
|
||||||
|
|
||||||
|
### Missing API key error
|
||||||
|
|
||||||
|
```
|
||||||
|
api_key or api_base is required for HTTP-based protocol "xxx"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Solution**: Provide `api_key` and/or `api_base` for HTTP-based providers.
|
||||||
|
|
||||||
|
## Need Help?
|
||||||
|
|
||||||
|
- [GitHub Issues](https://github.com/sipeed/picoclaw/issues)
|
||||||
|
- [Discussion #122](https://github.com/sipeed/picoclaw/discussions/122): Original proposal
|
||||||
|
|
@ -189,16 +189,7 @@ func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary str
|
||||||
systemPrompt += "\n\n## Summary of Previous Conversation\n\n" + summary
|
systemPrompt += "\n\n## Summary of Previous Conversation\n\n" + summary
|
||||||
}
|
}
|
||||||
|
|
||||||
//This fix prevents the session memory from LLM failure due to elimination of toolu_IDs required from LLM
|
history = sanitizeHistoryForProvider(history)
|
||||||
// --- INICIO DEL FIX ---
|
|
||||||
//Diegox-17
|
|
||||||
for len(history) > 0 && (history[0].Role == "tool") {
|
|
||||||
logger.DebugCF("agent", "Removing orphaned tool message from history to prevent LLM error",
|
|
||||||
map[string]interface{}{"role": history[0].Role})
|
|
||||||
history = history[1:]
|
|
||||||
}
|
|
||||||
//Diegox-17
|
|
||||||
// --- FIN DEL FIX ---
|
|
||||||
|
|
||||||
messages = append(messages, providers.Message{
|
messages = append(messages, providers.Message{
|
||||||
Role: "system",
|
Role: "system",
|
||||||
|
|
@ -207,14 +198,58 @@ func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary str
|
||||||
|
|
||||||
messages = append(messages, history...)
|
messages = append(messages, history...)
|
||||||
|
|
||||||
messages = append(messages, providers.Message{
|
if strings.TrimSpace(currentMessage) != "" {
|
||||||
Role: "user",
|
messages = append(messages, providers.Message{
|
||||||
Content: currentMessage,
|
Role: "user",
|
||||||
})
|
Content: currentMessage,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
return messages
|
return messages
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func sanitizeHistoryForProvider(history []providers.Message) []providers.Message {
|
||||||
|
if len(history) == 0 {
|
||||||
|
return history
|
||||||
|
}
|
||||||
|
|
||||||
|
sanitized := make([]providers.Message, 0, len(history))
|
||||||
|
for _, msg := range history {
|
||||||
|
switch msg.Role {
|
||||||
|
case "tool":
|
||||||
|
if len(sanitized) == 0 {
|
||||||
|
logger.DebugCF("agent", "Dropping orphaned leading tool message", map[string]interface{}{})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
last := sanitized[len(sanitized)-1]
|
||||||
|
if last.Role != "assistant" || len(last.ToolCalls) == 0 {
|
||||||
|
logger.DebugCF("agent", "Dropping orphaned tool message", map[string]interface{}{})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
sanitized = append(sanitized, msg)
|
||||||
|
|
||||||
|
case "assistant":
|
||||||
|
if len(msg.ToolCalls) > 0 {
|
||||||
|
if len(sanitized) == 0 {
|
||||||
|
logger.DebugCF("agent", "Dropping assistant tool-call turn at history start", map[string]interface{}{})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
prev := sanitized[len(sanitized)-1]
|
||||||
|
if prev.Role != "user" && prev.Role != "tool" {
|
||||||
|
logger.DebugCF("agent", "Dropping assistant tool-call turn with invalid predecessor", map[string]interface{}{"prev_role": prev.Role})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sanitized = append(sanitized, msg)
|
||||||
|
|
||||||
|
default:
|
||||||
|
sanitized = append(sanitized, msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sanitized
|
||||||
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) AddToolResult(messages []providers.Message, toolCallID, toolName, result string) []providers.Message {
|
func (cb *ContextBuilder) AddToolResult(messages []providers.Message, toolCallID, toolName, result string) []providers.Message {
|
||||||
messages = append(messages, providers.Message{
|
messages = append(messages, providers.Message{
|
||||||
Role: "tool",
|
Role: "tool",
|
||||||
|
|
|
||||||
|
|
@ -21,6 +21,8 @@ type AgentInstance struct {
|
||||||
Fallbacks []string
|
Fallbacks []string
|
||||||
Workspace string
|
Workspace string
|
||||||
MaxIterations int
|
MaxIterations int
|
||||||
|
MaxTokens int
|
||||||
|
Temperature float64
|
||||||
ContextWindow int
|
ContextWindow int
|
||||||
Provider providers.LLMProvider
|
Provider providers.LLMProvider
|
||||||
Sessions *session.SessionManager
|
Sessions *session.SessionManager
|
||||||
|
|
@ -76,6 +78,16 @@ func NewAgentInstance(
|
||||||
maxIter = 20
|
maxIter = 20
|
||||||
}
|
}
|
||||||
|
|
||||||
|
maxTokens := defaults.MaxTokens
|
||||||
|
if maxTokens == 0 {
|
||||||
|
maxTokens = 8192
|
||||||
|
}
|
||||||
|
|
||||||
|
temperature := 0.7
|
||||||
|
if defaults.Temperature != nil {
|
||||||
|
temperature = *defaults.Temperature
|
||||||
|
}
|
||||||
|
|
||||||
// Resolve fallback candidates
|
// Resolve fallback candidates
|
||||||
modelCfg := providers.ModelConfig{
|
modelCfg := providers.ModelConfig{
|
||||||
Primary: model,
|
Primary: model,
|
||||||
|
|
@ -90,7 +102,9 @@ func NewAgentInstance(
|
||||||
Fallbacks: fallbacks,
|
Fallbacks: fallbacks,
|
||||||
Workspace: workspace,
|
Workspace: workspace,
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
ContextWindow: defaults.MaxTokens,
|
MaxTokens: maxTokens,
|
||||||
|
Temperature: temperature,
|
||||||
|
ContextWindow: maxTokens,
|
||||||
Provider: provider,
|
Provider: provider,
|
||||||
Sessions: sessionsManager,
|
Sessions: sessionsManager,
|
||||||
ContextBuilder: contextBuilder,
|
ContextBuilder: contextBuilder,
|
||||||
|
|
|
||||||
95
pkg/agent/instance_test.go
Normal file
95
pkg/agent/instance_test.go
Normal file
|
|
@ -0,0 +1,95 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewAgentInstance_UsesDefaultsTemperatureAndMaxTokens(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 1234,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
configuredTemp := 1.0
|
||||||
|
cfg.Agents.Defaults.Temperature = &configuredTemp
|
||||||
|
|
||||||
|
provider := &mockProvider{}
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
|
||||||
|
if agent.MaxTokens != 1234 {
|
||||||
|
t.Fatalf("MaxTokens = %d, want %d", agent.MaxTokens, 1234)
|
||||||
|
}
|
||||||
|
if agent.Temperature != 1.0 {
|
||||||
|
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 1.0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAgentInstance_DefaultsTemperatureWhenZero(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 1234,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
configuredTemp := 0.0
|
||||||
|
cfg.Agents.Defaults.Temperature = &configuredTemp
|
||||||
|
|
||||||
|
provider := &mockProvider{}
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
|
||||||
|
if agent.Temperature != 0.0 {
|
||||||
|
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAgentInstance_DefaultsTemperatureWhenUnset(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 1234,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := &mockProvider{}
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
|
||||||
|
if agent.Temperature != 0.7 {
|
||||||
|
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.7)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -23,6 +23,7 @@ import (
|
||||||
"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/routing"
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
"github.com/sipeed/picoclaw/pkg/state"
|
"github.com/sipeed/picoclaw/pkg/state"
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
|
@ -94,7 +95,19 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers
|
||||||
})
|
})
|
||||||
agent.Tools.Register(messageTool)
|
agent.Tools.Register(messageTool)
|
||||||
|
|
||||||
|
// Skill discovery and installation tools
|
||||||
|
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
|
||||||
|
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
|
||||||
|
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
|
||||||
|
})
|
||||||
|
searchCache := skills.NewSearchCache(cfg.Tools.Skills.SearchCache.MaxSize, time.Duration(cfg.Tools.Skills.SearchCache.TTLSeconds)*time.Second)
|
||||||
|
agent.Tools.Register(tools.NewFindSkillsTool(registryMgr, searchCache))
|
||||||
|
agent.Tools.Register(tools.NewInstallSkillTool(registryMgr, agent.Workspace))
|
||||||
|
|
||||||
|
// Spawn tool with allowlist checker
|
||||||
subagentManager := tools.NewSubagentManager(provider, agent.Model, agent.Workspace, msgBus)
|
subagentManager := tools.NewSubagentManager(provider, agent.Model, agent.Workspace, msgBus)
|
||||||
|
subagentManager.SetLLMOptions(agent.MaxTokens, agent.Temperature)
|
||||||
|
|
||||||
subagentTools := tools.NewToolRegistry()
|
subagentTools := tools.NewToolRegistry()
|
||||||
for _, toolName := range agent.Tools.List() {
|
for _, toolName := range agent.Tools.List() {
|
||||||
if tool, exists := agent.Tools.Get(toolName); exists {
|
if tool, exists := agent.Tools.Get(toolName); exists {
|
||||||
|
|
@ -102,7 +115,6 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
subagentManager.SetTools(subagentTools)
|
subagentManager.SetTools(subagentTools)
|
||||||
|
|
||||||
spawnTool := tools.NewSpawnTool(subagentManager)
|
spawnTool := tools.NewSpawnTool(subagentManager)
|
||||||
currentAgentID := agentID
|
currentAgentID := agentID
|
||||||
spawnTool.SetAllowlistChecker(func(targetAgentID string) bool {
|
spawnTool.SetAllowlistChecker(func(targetAgentID string) bool {
|
||||||
|
|
@ -465,8 +477,8 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
|
||||||
"model": agent.Model,
|
"model": agent.Model,
|
||||||
"messages_count": len(messages),
|
"messages_count": len(messages),
|
||||||
"tools_count": len(providerToolDefs),
|
"tools_count": len(providerToolDefs),
|
||||||
"max_tokens": 8192,
|
"max_tokens": agent.MaxTokens,
|
||||||
"temperature": 0.7,
|
"temperature": agent.Temperature,
|
||||||
"system_prompt_len": len(messages[0].Content),
|
"system_prompt_len": len(messages[0].Content),
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
@ -487,8 +499,8 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
|
||||||
fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates,
|
fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates,
|
||||||
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
|
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
|
||||||
return agent.Provider.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{
|
return agent.Provider.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{
|
||||||
"max_tokens": 8192,
|
"max_tokens": agent.MaxTokens,
|
||||||
"temperature": 0.7,
|
"temperature": agent.Temperature,
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
@ -503,8 +515,8 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
|
||||||
return fbResult.Response, nil
|
return fbResult.Response, nil
|
||||||
}
|
}
|
||||||
return agent.Provider.Chat(ctx, messages, providerToolDefs, agent.Model, map[string]interface{}{
|
return agent.Provider.Chat(ctx, messages, providerToolDefs, agent.Model, map[string]interface{}{
|
||||||
"max_tokens": 8192,
|
"max_tokens": agent.MaxTokens,
|
||||||
"temperature": 0.7,
|
"temperature": agent.Temperature,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -570,16 +582,21 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
// Log tool calls
|
normalizedToolCalls := make([]providers.ToolCall, 0, len(response.ToolCalls))
|
||||||
toolNames := make([]string, 0, len(response.ToolCalls))
|
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range response.ToolCalls {
|
||||||
|
normalizedToolCalls = append(normalizedToolCalls, providers.NormalizeToolCall(tc))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Log tool calls
|
||||||
|
toolNames := make([]string, 0, len(normalizedToolCalls))
|
||||||
|
for _, tc := range normalizedToolCalls {
|
||||||
toolNames = append(toolNames, tc.Name)
|
toolNames = append(toolNames, tc.Name)
|
||||||
}
|
}
|
||||||
logger.InfoCF("agent", "LLM requested tool calls",
|
logger.InfoCF("agent", "LLM requested tool calls",
|
||||||
map[string]interface{}{
|
map[string]interface{}{
|
||||||
"agent_id": agent.ID,
|
"agent_id": agent.ID,
|
||||||
"tools": toolNames,
|
"tools": toolNames,
|
||||||
"count": len(response.ToolCalls),
|
"count": len(normalizedToolCalls),
|
||||||
"iteration": iteration,
|
"iteration": iteration,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
@ -588,15 +605,26 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
|
||||||
Role: "assistant",
|
Role: "assistant",
|
||||||
Content: response.Content,
|
Content: response.Content,
|
||||||
}
|
}
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range normalizedToolCalls {
|
||||||
argumentsJSON, _ := json.Marshal(tc.Arguments)
|
argumentsJSON, _ := json.Marshal(tc.Arguments)
|
||||||
|
// Copy ExtraContent to ensure thought_signature is persisted for Gemini 3
|
||||||
|
extraContent := tc.ExtraContent
|
||||||
|
thoughtSignature := ""
|
||||||
|
if tc.Function != nil {
|
||||||
|
thoughtSignature = tc.Function.ThoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
assistantMsg.ToolCalls = append(assistantMsg.ToolCalls, providers.ToolCall{
|
assistantMsg.ToolCalls = append(assistantMsg.ToolCalls, providers.ToolCall{
|
||||||
ID: tc.ID,
|
ID: tc.ID,
|
||||||
Type: "function",
|
Type: "function",
|
||||||
|
Name: tc.Name,
|
||||||
Function: &providers.FunctionCall{
|
Function: &providers.FunctionCall{
|
||||||
Name: tc.Name,
|
Name: tc.Name,
|
||||||
Arguments: string(argumentsJSON),
|
Arguments: string(argumentsJSON),
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
},
|
},
|
||||||
|
ExtraContent: extraContent,
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
messages = append(messages, assistantMsg)
|
messages = append(messages, assistantMsg)
|
||||||
|
|
@ -605,7 +633,7 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
|
||||||
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
|
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
|
||||||
|
|
||||||
// Execute tool calls
|
// Execute tool calls
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range normalizedToolCalls {
|
||||||
argsJSON, _ := json.Marshal(tc.Arguments)
|
argsJSON, _ := json.Marshal(tc.Arguments)
|
||||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||||
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
||||||
|
|
@ -732,31 +760,21 @@ func (al *AgentLoop) forceCompression(agent *AgentInstance, sessionKey string) {
|
||||||
mid := len(conversation) / 2
|
mid := len(conversation) / 2
|
||||||
|
|
||||||
// New history structure:
|
// New history structure:
|
||||||
// 1. System Prompt
|
// 1. System Prompt (with compression note appended)
|
||||||
// 2. [Summary of dropped part] - synthesized
|
// 2. Second half of conversation
|
||||||
// 3. Second half of conversation
|
// 3. Last message
|
||||||
// 4. Last message
|
|
||||||
|
|
||||||
// Simplified approach for emergency: Drop first half of conversation
|
|
||||||
// and rely on existing summary if present, or create a placeholder.
|
|
||||||
|
|
||||||
droppedCount := mid
|
droppedCount := mid
|
||||||
keptConversation := conversation[mid:]
|
keptConversation := conversation[mid:]
|
||||||
|
|
||||||
newHistory := make([]providers.Message, 0)
|
newHistory := make([]providers.Message, 0)
|
||||||
newHistory = append(newHistory, history[0]) // System prompt
|
|
||||||
|
|
||||||
// Add a note about compression
|
// Append compression note to the original system prompt instead of adding a new system message
|
||||||
compressionNote := fmt.Sprintf("[System: Emergency compression dropped %d oldest messages due to context limit]", droppedCount)
|
// This avoids having two consecutive system messages which some APIs (like Zhipu) reject
|
||||||
// If there was an existing summary, we might lose it if it was in the dropped part (which is just messages).
|
compressionNote := fmt.Sprintf("\n\n[System Note: Emergency compression dropped %d oldest messages due to context limit]", droppedCount)
|
||||||
// The summary is stored separately in session.Summary, so it persists!
|
enhancedSystemPrompt := history[0]
|
||||||
// We just need to ensure the user knows there's a gap.
|
enhancedSystemPrompt.Content = enhancedSystemPrompt.Content + compressionNote
|
||||||
|
newHistory = append(newHistory, enhancedSystemPrompt)
|
||||||
// We only modify the messages list here
|
|
||||||
newHistory = append(newHistory, providers.Message{
|
|
||||||
Role: "system",
|
|
||||||
Content: compressionNote,
|
|
||||||
})
|
|
||||||
|
|
||||||
newHistory = append(newHistory, keptConversation...)
|
newHistory = append(newHistory, keptConversation...)
|
||||||
newHistory = append(newHistory, history[len(history)-1]) // Last message
|
newHistory = append(newHistory, history[len(history)-1]) // Last message
|
||||||
|
|
|
||||||
|
|
@ -14,20 +14,6 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
)
|
)
|
||||||
|
|
||||||
// mockProvider is a simple mock LLM provider for testing
|
|
||||||
type mockProvider struct{}
|
|
||||||
|
|
||||||
func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]interface{}) (*providers.LLMResponse, error) {
|
|
||||||
return &providers.LLMResponse{
|
|
||||||
Content: "Mock response",
|
|
||||||
ToolCalls: []providers.ToolCall{},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *mockProvider) GetDefaultModel() string {
|
|
||||||
return "mock-model"
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRecordLastChannel(t *testing.T) {
|
func TestRecordLastChannel(t *testing.T) {
|
||||||
// Create temp workspace
|
// Create temp workspace
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
|
@ -603,7 +589,6 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) {
|
||||||
// Call ProcessDirectWithChannel
|
// Call ProcessDirectWithChannel
|
||||||
// Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration
|
// Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration
|
||||||
response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat")
|
response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat")
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Expected success after retry, got error: %v", err)
|
t.Fatalf("Expected success after retry, got error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
20
pkg/agent/mock_provider_test.go
Normal file
20
pkg/agent/mock_provider_test.go
Normal file
|
|
@ -0,0 +1,20 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mockProvider struct{}
|
||||||
|
|
||||||
|
func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]interface{}) (*providers.LLMResponse, error) {
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: "Mock response",
|
||||||
|
ToolCalls: []providers.ToolCall{},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockProvider) GetDefaultModel() string {
|
||||||
|
return "mock-model"
|
||||||
|
}
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
|
|
@ -11,6 +12,7 @@ import (
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"runtime"
|
"runtime"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
@ -19,11 +21,13 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
type OAuthProviderConfig struct {
|
type OAuthProviderConfig struct {
|
||||||
Issuer string
|
Issuer string
|
||||||
ClientID string
|
ClientID string
|
||||||
Scopes string
|
ClientSecret string // Required for Google OAuth (confidential client)
|
||||||
Originator string
|
TokenURL string // Override token endpoint (Google uses a different URL than issuer)
|
||||||
Port int
|
Scopes string
|
||||||
|
Originator string
|
||||||
|
Port int
|
||||||
}
|
}
|
||||||
|
|
||||||
func OpenAIOAuthConfig() OAuthProviderConfig {
|
func OpenAIOAuthConfig() OAuthProviderConfig {
|
||||||
|
|
@ -36,6 +40,30 @@ func OpenAIOAuthConfig() OAuthProviderConfig {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GoogleAntigravityOAuthConfig returns the OAuth configuration for Google Cloud Code Assist (Antigravity).
|
||||||
|
// Client credentials are the same ones used by OpenCode/pi-ai for Cloud Code Assist access.
|
||||||
|
func GoogleAntigravityOAuthConfig() OAuthProviderConfig {
|
||||||
|
// These are the same client credentials used by the OpenCode antigravity plugin.
|
||||||
|
clientID := decodeBase64("MTA3MTAwNjA2MDU5MS10bWhzc2luMmgyMWxjcmUyMzV2dG9sb2poNGc0MDNlcC5hcHBzLmdvb2dsZXVzZXJjb250ZW50LmNvbQ==")
|
||||||
|
clientSecret := decodeBase64("R09DU1BYLUs1OEZXUjQ4NkxkTEoxbUxCOHNYQzR6NnFEQWY=")
|
||||||
|
return OAuthProviderConfig{
|
||||||
|
Issuer: "https://accounts.google.com/o/oauth2/v2",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: clientSecret,
|
||||||
|
Scopes: "https://www.googleapis.com/auth/cloud-platform https://www.googleapis.com/auth/userinfo.email https://www.googleapis.com/auth/userinfo.profile https://www.googleapis.com/auth/cclog https://www.googleapis.com/auth/experimentsandconfigs",
|
||||||
|
Port: 51121,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeBase64(s string) string {
|
||||||
|
data, err := base64.StdEncoding.DecodeString(s)
|
||||||
|
if err != nil {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return string(data)
|
||||||
|
}
|
||||||
|
|
||||||
func generateState() (string, error) {
|
func generateState() (string, error) {
|
||||||
buf := make([]byte, 32)
|
buf := make([]byte, 32)
|
||||||
if _, err := rand.Read(buf); err != nil {
|
if _, err := rand.Read(buf); err != nil {
|
||||||
|
|
@ -101,8 +129,17 @@ func LoginBrowser(cfg OAuthProviderConfig) (*AuthCredential, error) {
|
||||||
fmt.Printf("Could not open browser automatically.\nPlease open this URL manually:\n\n%s\n\n", authURL)
|
fmt.Printf("Could not open browser automatically.\nPlease open this URL manually:\n\n%s\n\n", authURL)
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Println("If you're running in a headless environment, use: picoclaw auth login --provider openai --device-code")
|
fmt.Printf("Wait! If you are in a headless environment (like Coolify/VPS) and cannot reach localhost:%d,\n", cfg.Port)
|
||||||
fmt.Println("Waiting for authentication in browser...")
|
fmt.Println("please complete the login in your local browser and then PASTE the final redirect URL (or just the code) here.")
|
||||||
|
fmt.Println("Waiting for authentication (browser or manual paste)...")
|
||||||
|
|
||||||
|
// Start manual input in a goroutine
|
||||||
|
manualCh := make(chan string)
|
||||||
|
go func() {
|
||||||
|
reader := bufio.NewReader(os.Stdin)
|
||||||
|
input, _ := reader.ReadString('\n')
|
||||||
|
manualCh <- strings.TrimSpace(input)
|
||||||
|
}()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case result := <-resultCh:
|
case result := <-resultCh:
|
||||||
|
|
@ -110,6 +147,22 @@ func LoginBrowser(cfg OAuthProviderConfig) (*AuthCredential, error) {
|
||||||
return nil, result.err
|
return nil, result.err
|
||||||
}
|
}
|
||||||
return exchangeCodeForTokens(cfg, result.code, pkce.CodeVerifier, redirectURI)
|
return exchangeCodeForTokens(cfg, result.code, pkce.CodeVerifier, redirectURI)
|
||||||
|
case manualInput := <-manualCh:
|
||||||
|
if manualInput == "" {
|
||||||
|
return nil, fmt.Errorf("manual input cancelled")
|
||||||
|
}
|
||||||
|
// Extract code from URL if it's a full URL
|
||||||
|
code := manualInput
|
||||||
|
if strings.Contains(manualInput, "?") {
|
||||||
|
u, err := url.Parse(manualInput)
|
||||||
|
if err == nil {
|
||||||
|
code = u.Query().Get("code")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if code == "" {
|
||||||
|
return nil, fmt.Errorf("could not find authorization code in input")
|
||||||
|
}
|
||||||
|
return exchangeCodeForTokens(cfg, code, pkce.CodeVerifier, redirectURI)
|
||||||
case <-time.After(5 * time.Minute):
|
case <-time.After(5 * time.Minute):
|
||||||
return nil, fmt.Errorf("authentication timed out after 5 minutes")
|
return nil, fmt.Errorf("authentication timed out after 5 minutes")
|
||||||
}
|
}
|
||||||
|
|
@ -269,8 +322,16 @@ func RefreshAccessToken(cred *AuthCredential, cfg OAuthProviderConfig) (*AuthCre
|
||||||
"refresh_token": {cred.RefreshToken},
|
"refresh_token": {cred.RefreshToken},
|
||||||
"scope": {"openid profile email"},
|
"scope": {"openid profile email"},
|
||||||
}
|
}
|
||||||
|
if cfg.ClientSecret != "" {
|
||||||
|
data.Set("client_secret", cfg.ClientSecret)
|
||||||
|
}
|
||||||
|
|
||||||
resp, err := http.PostForm(cfg.Issuer+"/oauth/token", data)
|
tokenURL := cfg.Issuer + "/oauth/token"
|
||||||
|
if cfg.TokenURL != "" {
|
||||||
|
tokenURL = cfg.TokenURL
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.PostForm(tokenURL, data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("refreshing token: %w", err)
|
return nil, fmt.Errorf("refreshing token: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -291,6 +352,12 @@ func RefreshAccessToken(cred *AuthCredential, cfg OAuthProviderConfig) (*AuthCre
|
||||||
if refreshed.AccountID == "" {
|
if refreshed.AccountID == "" {
|
||||||
refreshed.AccountID = cred.AccountID
|
refreshed.AccountID = cred.AccountID
|
||||||
}
|
}
|
||||||
|
if cred.Email != "" && refreshed.Email == "" {
|
||||||
|
refreshed.Email = cred.Email
|
||||||
|
}
|
||||||
|
if cred.ProjectID != "" && refreshed.ProjectID == "" {
|
||||||
|
refreshed.ProjectID = cred.ProjectID
|
||||||
|
}
|
||||||
return refreshed, nil
|
return refreshed, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -300,21 +367,35 @@ func BuildAuthorizeURL(cfg OAuthProviderConfig, pkce PKCECodes, state, redirectU
|
||||||
|
|
||||||
func buildAuthorizeURL(cfg OAuthProviderConfig, pkce PKCECodes, state, redirectURI string) string {
|
func buildAuthorizeURL(cfg OAuthProviderConfig, pkce PKCECodes, state, redirectURI string) string {
|
||||||
params := url.Values{
|
params := url.Values{
|
||||||
"response_type": {"code"},
|
"response_type": {"code"},
|
||||||
"client_id": {cfg.ClientID},
|
"client_id": {cfg.ClientID},
|
||||||
"redirect_uri": {redirectURI},
|
"redirect_uri": {redirectURI},
|
||||||
"scope": {cfg.Scopes},
|
"scope": {cfg.Scopes},
|
||||||
"code_challenge": {pkce.CodeChallenge},
|
"code_challenge": {pkce.CodeChallenge},
|
||||||
"code_challenge_method": {"S256"},
|
"code_challenge_method": {"S256"},
|
||||||
"id_token_add_organizations": {"true"},
|
"state": {state},
|
||||||
"codex_cli_simplified_flow": {"true"},
|
|
||||||
"state": {state},
|
|
||||||
}
|
}
|
||||||
if strings.Contains(strings.ToLower(cfg.Issuer), "auth.openai.com") {
|
|
||||||
params.Set("originator", "picoclaw")
|
isGoogle := strings.Contains(strings.ToLower(cfg.Issuer), "accounts.google.com")
|
||||||
|
if isGoogle {
|
||||||
|
// Google OAuth requires these for refresh token support
|
||||||
|
params.Set("access_type", "offline")
|
||||||
|
params.Set("prompt", "consent")
|
||||||
|
} else {
|
||||||
|
// OpenAI-specific parameters
|
||||||
|
params.Set("id_token_add_organizations", "true")
|
||||||
|
params.Set("codex_cli_simplified_flow", "true")
|
||||||
|
if strings.Contains(strings.ToLower(cfg.Issuer), "auth.openai.com") {
|
||||||
|
params.Set("originator", "picoclaw")
|
||||||
|
}
|
||||||
|
if cfg.Originator != "" {
|
||||||
|
params.Set("originator", cfg.Originator)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if cfg.Originator != "" {
|
|
||||||
params.Set("originator", cfg.Originator)
|
// Google uses /auth path, OpenAI uses /oauth/authorize
|
||||||
|
if isGoogle {
|
||||||
|
return cfg.Issuer + "/auth?" + params.Encode()
|
||||||
}
|
}
|
||||||
return cfg.Issuer + "/oauth/authorize?" + params.Encode()
|
return cfg.Issuer + "/oauth/authorize?" + params.Encode()
|
||||||
}
|
}
|
||||||
|
|
@ -327,8 +408,22 @@ func exchangeCodeForTokens(cfg OAuthProviderConfig, code, codeVerifier, redirect
|
||||||
"client_id": {cfg.ClientID},
|
"client_id": {cfg.ClientID},
|
||||||
"code_verifier": {codeVerifier},
|
"code_verifier": {codeVerifier},
|
||||||
}
|
}
|
||||||
|
if cfg.ClientSecret != "" {
|
||||||
|
data.Set("client_secret", cfg.ClientSecret)
|
||||||
|
}
|
||||||
|
|
||||||
resp, err := http.PostForm(cfg.Issuer+"/oauth/token", data)
|
tokenURL := cfg.Issuer + "/oauth/token"
|
||||||
|
if cfg.TokenURL != "" {
|
||||||
|
tokenURL = cfg.TokenURL
|
||||||
|
}
|
||||||
|
|
||||||
|
// Determine provider name from config
|
||||||
|
provider := "openai"
|
||||||
|
if cfg.TokenURL != "" && strings.Contains(cfg.TokenURL, "googleapis.com") {
|
||||||
|
provider = "google-antigravity"
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.PostForm(tokenURL, data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("exchanging code for tokens: %w", err)
|
return nil, fmt.Errorf("exchanging code for tokens: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -339,7 +434,7 @@ func exchangeCodeForTokens(cfg OAuthProviderConfig, code, codeVerifier, redirect
|
||||||
return nil, fmt.Errorf("token exchange failed: %s", string(body))
|
return nil, fmt.Errorf("token exchange failed: %s", string(body))
|
||||||
}
|
}
|
||||||
|
|
||||||
return parseTokenResponse(body, "openai")
|
return parseTokenResponse(body, provider)
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseTokenResponse(body []byte, provider string) (*AuthCredential, error) {
|
func parseTokenResponse(body []byte, provider string) (*AuthCredential, error) {
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,8 @@ type AuthCredential struct {
|
||||||
ExpiresAt time.Time `json:"expires_at,omitempty"`
|
ExpiresAt time.Time `json:"expires_at,omitempty"`
|
||||||
Provider string `json:"provider"`
|
Provider string `json:"provider"`
|
||||||
AuthMethod string `json:"auth_method"`
|
AuthMethod string `json:"auth_method"`
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
ProjectID string `json:"project_id,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuthStore struct {
|
type AuthStore struct {
|
||||||
|
|
|
||||||
|
|
@ -155,6 +155,14 @@ func (c *DingTalkChannel) onChatBotMessageReceived(ctx context.Context, data *ch
|
||||||
"session_webhook": data.SessionWebhook,
|
"session_webhook": data.SessionWebhook,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if data.ConversationType == "1" {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = data.ConversationId
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("dingtalk", "Received message", map[string]interface{}{
|
logger.DebugCF("dingtalk", "Received message", map[string]interface{}{
|
||||||
"sender_nick": senderNick,
|
"sender_nick": senderNick,
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bwmarrin/discordgo"
|
"github.com/bwmarrin/discordgo"
|
||||||
|
|
@ -26,6 +27,9 @@ type DiscordChannel struct {
|
||||||
config config.DiscordConfig
|
config config.DiscordConfig
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber *voice.GroqTranscriber
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
|
typingMu sync.Mutex
|
||||||
|
typingStop map[string]chan struct{} // chatID → stop signal
|
||||||
|
botUserID string // stored for mention checking
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
||||||
|
|
@ -42,6 +46,7 @@ func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordC
|
||||||
config: cfg,
|
config: cfg,
|
||||||
transcriber: nil,
|
transcriber: nil,
|
||||||
ctx: context.Background(),
|
ctx: context.Background(),
|
||||||
|
typingStop: make(map[string]chan struct{}),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -60,6 +65,14 @@ func (c *DiscordChannel) Start(ctx context.Context) error {
|
||||||
logger.InfoC("discord", "Starting Discord bot")
|
logger.InfoC("discord", "Starting Discord bot")
|
||||||
|
|
||||||
c.ctx = ctx
|
c.ctx = ctx
|
||||||
|
|
||||||
|
// Get bot user ID before opening session to avoid race condition
|
||||||
|
botUser, err := c.session.User("@me")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to get bot user: %w", err)
|
||||||
|
}
|
||||||
|
c.botUserID = botUser.ID
|
||||||
|
|
||||||
c.session.AddHandler(c.handleMessage)
|
c.session.AddHandler(c.handleMessage)
|
||||||
|
|
||||||
if err := c.session.Open(); err != nil {
|
if err := c.session.Open(); err != nil {
|
||||||
|
|
@ -68,10 +81,6 @@ func (c *DiscordChannel) Start(ctx context.Context) error {
|
||||||
|
|
||||||
c.setRunning(true)
|
c.setRunning(true)
|
||||||
|
|
||||||
botUser, err := c.session.User("@me")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to get bot user: %w", err)
|
|
||||||
}
|
|
||||||
logger.InfoCF("discord", "Discord bot connected", map[string]any{
|
logger.InfoCF("discord", "Discord bot connected", map[string]any{
|
||||||
"username": botUser.Username,
|
"username": botUser.Username,
|
||||||
"user_id": botUser.ID,
|
"user_id": botUser.ID,
|
||||||
|
|
@ -84,6 +93,14 @@ func (c *DiscordChannel) Stop(ctx context.Context) error {
|
||||||
logger.InfoC("discord", "Stopping Discord bot")
|
logger.InfoC("discord", "Stopping Discord bot")
|
||||||
c.setRunning(false)
|
c.setRunning(false)
|
||||||
|
|
||||||
|
// Stop all typing goroutines before closing session
|
||||||
|
c.typingMu.Lock()
|
||||||
|
for chatID, stop := range c.typingStop {
|
||||||
|
close(stop)
|
||||||
|
delete(c.typingStop, chatID)
|
||||||
|
}
|
||||||
|
c.typingMu.Unlock()
|
||||||
|
|
||||||
if err := c.session.Close(); err != nil {
|
if err := c.session.Close(); err != nil {
|
||||||
return fmt.Errorf("failed to close discord session: %w", err)
|
return fmt.Errorf("failed to close discord session: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -92,6 +109,8 @@ func (c *DiscordChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
c.stopTyping(msg.ChatID)
|
||||||
|
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return fmt.Errorf("discord bot not running")
|
return fmt.Errorf("discord bot not running")
|
||||||
}
|
}
|
||||||
|
|
@ -106,7 +125,7 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
chunks := splitMessage(msg.Content, 1500) // Discord has a limit of 2000 characters per message, leave 500 for natural split e.g. code blocks
|
chunks := utils.SplitMessage(msg.Content, 2000) // Split messages into chunks, Discord length limit: 2000 chars
|
||||||
|
|
||||||
for _, chunk := range chunks {
|
for _, chunk := range chunks {
|
||||||
if err := c.sendChunk(ctx, channelID, chunk); err != nil {
|
if err := c.sendChunk(ctx, channelID, chunk); err != nil {
|
||||||
|
|
@ -117,134 +136,8 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// splitMessage splits long messages into chunks, preserving code block integrity
|
|
||||||
// Uses natural boundaries (newlines, spaces) and extends messages slightly to avoid breaking code blocks
|
|
||||||
func splitMessage(content string, limit int) []string {
|
|
||||||
var messages []string
|
|
||||||
|
|
||||||
for len(content) > 0 {
|
|
||||||
if len(content) <= limit {
|
|
||||||
messages = append(messages, content)
|
|
||||||
break
|
|
||||||
}
|
|
||||||
|
|
||||||
msgEnd := limit
|
|
||||||
|
|
||||||
// Find natural split point within the limit
|
|
||||||
msgEnd = findLastNewline(content[:limit], 200)
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = findLastSpace(content[:limit], 100)
|
|
||||||
}
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = limit
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if this would end with an incomplete code block
|
|
||||||
candidate := content[:msgEnd]
|
|
||||||
unclosedIdx := findLastUnclosedCodeBlock(candidate)
|
|
||||||
|
|
||||||
if unclosedIdx >= 0 {
|
|
||||||
// Message would end with incomplete code block
|
|
||||||
// Try to extend to include the closing ``` (with some buffer)
|
|
||||||
extendedLimit := limit + 500 // Allow 500 char buffer for code blocks
|
|
||||||
if len(content) > extendedLimit {
|
|
||||||
closingIdx := findNextClosingCodeBlock(content, msgEnd)
|
|
||||||
if closingIdx > 0 && closingIdx <= extendedLimit {
|
|
||||||
// Extend to include the closing ```
|
|
||||||
msgEnd = closingIdx
|
|
||||||
} else {
|
|
||||||
// Can't find closing, split before the code block
|
|
||||||
msgEnd = findLastNewline(content[:unclosedIdx], 200)
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = findLastSpace(content[:unclosedIdx], 100)
|
|
||||||
}
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = unclosedIdx
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// Remaining content fits within extended limit
|
|
||||||
msgEnd = len(content)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = limit
|
|
||||||
}
|
|
||||||
|
|
||||||
messages = append(messages, content[:msgEnd])
|
|
||||||
content = strings.TrimSpace(content[msgEnd:])
|
|
||||||
}
|
|
||||||
|
|
||||||
return messages
|
|
||||||
}
|
|
||||||
|
|
||||||
// findLastUnclosedCodeBlock finds the last opening ``` that doesn't have a closing ```
|
|
||||||
// Returns the position of the opening ``` or -1 if all code blocks are complete
|
|
||||||
func findLastUnclosedCodeBlock(text string) int {
|
|
||||||
count := 0
|
|
||||||
lastOpenIdx := -1
|
|
||||||
|
|
||||||
for i := 0; i < len(text); i++ {
|
|
||||||
if i+2 < len(text) && text[i] == '`' && text[i+1] == '`' && text[i+2] == '`' {
|
|
||||||
if count == 0 {
|
|
||||||
lastOpenIdx = i
|
|
||||||
}
|
|
||||||
count++
|
|
||||||
i += 2
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// If odd number of ``` markers, last one is unclosed
|
|
||||||
if count%2 == 1 {
|
|
||||||
return lastOpenIdx
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// findNextClosingCodeBlock finds the next closing ``` starting from a position
|
|
||||||
// Returns the position after the closing ``` or -1 if not found
|
|
||||||
func findNextClosingCodeBlock(text string, startIdx int) int {
|
|
||||||
for i := startIdx; i < len(text); i++ {
|
|
||||||
if i+2 < len(text) && text[i] == '`' && text[i+1] == '`' && text[i+2] == '`' {
|
|
||||||
return i + 3
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// findLastNewline finds the last newline character within the last N characters
|
|
||||||
// Returns the position of the newline or -1 if not found
|
|
||||||
func findLastNewline(s string, searchWindow int) int {
|
|
||||||
searchStart := len(s) - searchWindow
|
|
||||||
if searchStart < 0 {
|
|
||||||
searchStart = 0
|
|
||||||
}
|
|
||||||
for i := len(s) - 1; i >= searchStart; i-- {
|
|
||||||
if s[i] == '\n' {
|
|
||||||
return i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// findLastSpace finds the last space character within the last N characters
|
|
||||||
// Returns the position of the space or -1 if not found
|
|
||||||
func findLastSpace(s string, searchWindow int) int {
|
|
||||||
searchStart := len(s) - searchWindow
|
|
||||||
if searchStart < 0 {
|
|
||||||
searchStart = 0
|
|
||||||
}
|
|
||||||
for i := len(s) - 1; i >= searchStart; i-- {
|
|
||||||
if s[i] == ' ' || s[i] == '\t' {
|
|
||||||
return i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content string) error {
|
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content string) error {
|
||||||
// 使用传入的 ctx 进行超时控制
|
// Use the passed ctx for timeout control
|
||||||
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
|
|
@ -265,7 +158,7 @@ func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content strin
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// appendContent 安全地追加内容到现有文本
|
// appendContent safely appends content to existing text
|
||||||
func appendContent(content, suffix string) string {
|
func appendContent(content, suffix string) string {
|
||||||
if content == "" {
|
if content == "" {
|
||||||
return suffix
|
return suffix
|
||||||
|
|
@ -282,13 +175,7 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.session.ChannelTyping(m.ChannelID); err != nil {
|
// Check allowlist first to avoid downloading attachments and transcribing for rejected users
|
||||||
logger.ErrorCF("discord", "Failed to send typing indicator", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// 检查白名单,避免为被拒绝的用户下载附件和转录
|
|
||||||
if !c.IsAllowed(m.Author.ID) {
|
if !c.IsAllowed(m.Author.ID) {
|
||||||
logger.DebugCF("discord", "Message rejected by allowlist", map[string]any{
|
logger.DebugCF("discord", "Message rejected by allowlist", map[string]any{
|
||||||
"user_id": m.Author.ID,
|
"user_id": m.Author.ID,
|
||||||
|
|
@ -296,6 +183,24 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// If configured to only respond to mentions, check if bot is mentioned
|
||||||
|
// Skip this check for DMs (GuildID is empty) - DMs should always be responded to
|
||||||
|
if c.config.MentionOnly && m.GuildID != "" {
|
||||||
|
isMentioned := false
|
||||||
|
for _, mention := range m.Mentions {
|
||||||
|
if mention.ID == c.botUserID {
|
||||||
|
isMentioned = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !isMentioned {
|
||||||
|
logger.DebugCF("discord", "Message ignored - bot not mentioned", map[string]any{
|
||||||
|
"user_id": m.Author.ID,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
senderID := m.Author.ID
|
senderID := m.Author.ID
|
||||||
senderName := m.Author.Username
|
senderName := m.Author.Username
|
||||||
if m.Author.Discriminator != "" && m.Author.Discriminator != "0" {
|
if m.Author.Discriminator != "" && m.Author.Discriminator != "0" {
|
||||||
|
|
@ -303,10 +208,11 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
}
|
}
|
||||||
|
|
||||||
content := m.Content
|
content := m.Content
|
||||||
|
content = c.stripBotMention(content)
|
||||||
mediaPaths := make([]string, 0, len(m.Attachments))
|
mediaPaths := make([]string, 0, len(m.Attachments))
|
||||||
localFiles := make([]string, 0, len(m.Attachments))
|
localFiles := make([]string, 0, len(m.Attachments))
|
||||||
|
|
||||||
// 确保临时文件在函数返回时被清理
|
// Ensure temp files are cleaned up when function returns
|
||||||
defer func() {
|
defer func() {
|
||||||
for _, file := range localFiles {
|
for _, file := range localFiles {
|
||||||
if err := os.Remove(file); err != nil {
|
if err := os.Remove(file); err != nil {
|
||||||
|
|
@ -330,7 +236,7 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
if c.transcriber != nil && c.transcriber.IsAvailable() {
|
if c.transcriber != nil && c.transcriber.IsAvailable() {
|
||||||
ctx, cancel := context.WithTimeout(c.getContext(), transcriptionTimeout)
|
ctx, cancel := context.WithTimeout(c.getContext(), transcriptionTimeout)
|
||||||
result, err := c.transcriber.Transcribe(ctx, localPath)
|
result, err := c.transcriber.Transcribe(ctx, localPath)
|
||||||
cancel() // 立即释放context资源,避免在for循环中泄漏
|
cancel() // Release context resources immediately to avoid leaks in for loop
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("discord", "Voice transcription failed", map[string]any{
|
logger.ErrorCF("discord", "Voice transcription failed", map[string]any{
|
||||||
|
|
@ -370,6 +276,9 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
content = "[media only]"
|
content = "[media only]"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Start typing after all early returns — guaranteed to have a matching Send()
|
||||||
|
c.startTyping(m.ChannelID)
|
||||||
|
|
||||||
logger.DebugCF("discord", "Received message", map[string]any{
|
logger.DebugCF("discord", "Received message", map[string]any{
|
||||||
"sender_name": senderName,
|
"sender_name": senderName,
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
|
|
@ -398,8 +307,66 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
c.HandleMessage(senderID, m.ChannelID, content, mediaPaths, metadata)
|
c.HandleMessage(senderID, m.ChannelID, content, mediaPaths, metadata)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// startTyping starts a continuous typing indicator loop for the given chatID.
|
||||||
|
// It stops any existing typing loop for that chatID before starting a new one.
|
||||||
|
func (c *DiscordChannel) startTyping(chatID string) {
|
||||||
|
c.typingMu.Lock()
|
||||||
|
// Stop existing loop for this chatID if any
|
||||||
|
if stop, ok := c.typingStop[chatID]; ok {
|
||||||
|
close(stop)
|
||||||
|
}
|
||||||
|
stop := make(chan struct{})
|
||||||
|
c.typingStop[chatID] = stop
|
||||||
|
c.typingMu.Unlock()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
if err := c.session.ChannelTyping(chatID); err != nil {
|
||||||
|
logger.DebugCF("discord", "ChannelTyping error", map[string]interface{}{"chatID": chatID, "err": err})
|
||||||
|
}
|
||||||
|
ticker := time.NewTicker(8 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
timeout := time.After(5 * time.Minute)
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-stop:
|
||||||
|
return
|
||||||
|
case <-timeout:
|
||||||
|
return
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
if err := c.session.ChannelTyping(chatID); err != nil {
|
||||||
|
logger.DebugCF("discord", "ChannelTyping error", map[string]interface{}{"chatID": chatID, "err": err})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopTyping stops the typing indicator loop for the given chatID.
|
||||||
|
func (c *DiscordChannel) stopTyping(chatID string) {
|
||||||
|
c.typingMu.Lock()
|
||||||
|
defer c.typingMu.Unlock()
|
||||||
|
if stop, ok := c.typingStop[chatID]; ok {
|
||||||
|
close(stop)
|
||||||
|
delete(c.typingStop, chatID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) downloadAttachment(url, filename string) string {
|
func (c *DiscordChannel) downloadAttachment(url, filename string) string {
|
||||||
return utils.DownloadFile(url, filename, utils.DownloadOptions{
|
return utils.DownloadFile(url, filename, utils.DownloadOptions{
|
||||||
LoggerPrefix: "discord",
|
LoggerPrefix: "discord",
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// stripBotMention removes the bot mention from the message content.
|
||||||
|
// Discord mentions have the format <@USER_ID> or <@!USER_ID> (with nickname).
|
||||||
|
func (c *DiscordChannel) stripBotMention(text string) string {
|
||||||
|
if c.botUserID == "" {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
// Remove both regular mention <@USER_ID> and nickname mention <@!USER_ID>
|
||||||
|
text = strings.ReplaceAll(text, fmt.Sprintf("<@%s>", c.botUserID), "")
|
||||||
|
text = strings.ReplaceAll(text, fmt.Sprintf("<@!%s>", c.botUserID), "")
|
||||||
|
return strings.TrimSpace(text)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -165,6 +165,15 @@ func (c *FeishuChannel) handleMessageReceive(_ context.Context, event *larkim.P2
|
||||||
metadata["tenant_key"] = *sender.TenantKey
|
metadata["tenant_key"] = *sender.TenantKey
|
||||||
}
|
}
|
||||||
|
|
||||||
|
chatType := stringValue(message.ChatType)
|
||||||
|
if chatType == "p2p" {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = chatID
|
||||||
|
}
|
||||||
|
|
||||||
logger.InfoCF("feishu", "Feishu message received", map[string]interface{}{
|
logger.InfoCF("feishu", "Feishu message received", map[string]interface{}{
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
"chat_id": chatID,
|
"chat_id": chatID,
|
||||||
|
|
|
||||||
|
|
@ -366,6 +366,14 @@ func (c *LINEChannel) processEvent(event lineEvent) {
|
||||||
"message_id": msg.ID,
|
"message_id": msg.ID,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if isGroup {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = chatID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("line", "Received message", map[string]interface{}{
|
logger.DebugCF("line", "Received message", map[string]interface{}{
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
"chat_id": chatID,
|
"chat_id": chatID,
|
||||||
|
|
|
||||||
|
|
@ -170,6 +170,8 @@ func (c *MaixCamChannel) handlePersonDetection(msg MaixCamMessage) {
|
||||||
"y": fmt.Sprintf("%.0f", y),
|
"y": fmt.Sprintf("%.0f", y),
|
||||||
"w": fmt.Sprintf("%.0f", w),
|
"w": fmt.Sprintf("%.0f", w),
|
||||||
"h": fmt.Sprintf("%.0f", h),
|
"h": fmt.Sprintf("%.0f", h),
|
||||||
|
"peer_kind": "channel",
|
||||||
|
"peer_id": "default",
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, []string{}, metadata)
|
c.HandleMessage(senderID, chatID, content, []string{}, metadata)
|
||||||
|
|
|
||||||
|
|
@ -4,9 +4,11 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"os"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
@ -14,20 +16,28 @@ 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/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/voice"
|
||||||
)
|
)
|
||||||
|
|
||||||
type OneBotChannel struct {
|
type OneBotChannel struct {
|
||||||
*BaseChannel
|
*BaseChannel
|
||||||
config config.OneBotConfig
|
config config.OneBotConfig
|
||||||
conn *websocket.Conn
|
conn *websocket.Conn
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
dedup map[string]struct{}
|
dedup map[string]struct{}
|
||||||
dedupRing []string
|
dedupRing []string
|
||||||
dedupIdx int
|
dedupIdx int
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
writeMu sync.Mutex
|
writeMu sync.Mutex
|
||||||
echoCounter int64
|
echoCounter int64
|
||||||
|
selfID int64
|
||||||
|
pending map[string]chan json.RawMessage
|
||||||
|
pendingMu sync.Mutex
|
||||||
|
transcriber *voice.GroqTranscriber
|
||||||
|
lastMessageID sync.Map
|
||||||
|
pendingEmojiMsg sync.Map
|
||||||
}
|
}
|
||||||
|
|
||||||
type oneBotRawEvent struct {
|
type oneBotRawEvent struct {
|
||||||
|
|
@ -43,9 +53,11 @@ type oneBotRawEvent struct {
|
||||||
SelfID json.RawMessage `json:"self_id"`
|
SelfID json.RawMessage `json:"self_id"`
|
||||||
Time json.RawMessage `json:"time"`
|
Time json.RawMessage `json:"time"`
|
||||||
MetaEventType string `json:"meta_event_type"`
|
MetaEventType string `json:"meta_event_type"`
|
||||||
|
NoticeType string `json:"notice_type"`
|
||||||
Echo string `json:"echo"`
|
Echo string `json:"echo"`
|
||||||
RetCode json.RawMessage `json:"retcode"`
|
RetCode json.RawMessage `json:"retcode"`
|
||||||
Status BotStatus `json:"status"`
|
Status json.RawMessage `json:"status"`
|
||||||
|
Data json.RawMessage `json:"data"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type BotStatus struct {
|
type BotStatus struct {
|
||||||
|
|
@ -53,42 +65,36 @@ type BotStatus struct {
|
||||||
Good bool `json:"good"`
|
Good bool `json:"good"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isAPIResponse(raw json.RawMessage) bool {
|
||||||
|
if len(raw) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
var s string
|
||||||
|
if json.Unmarshal(raw, &s) == nil {
|
||||||
|
return s == "ok" || s == "failed"
|
||||||
|
}
|
||||||
|
var bs BotStatus
|
||||||
|
if json.Unmarshal(raw, &bs) == nil {
|
||||||
|
return bs.Online || bs.Good
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
type oneBotSender struct {
|
type oneBotSender struct {
|
||||||
UserID json.RawMessage `json:"user_id"`
|
UserID json.RawMessage `json:"user_id"`
|
||||||
Nickname string `json:"nickname"`
|
Nickname string `json:"nickname"`
|
||||||
Card string `json:"card"`
|
Card string `json:"card"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type oneBotEvent struct {
|
|
||||||
PostType string
|
|
||||||
MessageType string
|
|
||||||
SubType string
|
|
||||||
MessageID string
|
|
||||||
UserID int64
|
|
||||||
GroupID int64
|
|
||||||
Content string
|
|
||||||
RawContent string
|
|
||||||
IsBotMentioned bool
|
|
||||||
Sender oneBotSender
|
|
||||||
SelfID int64
|
|
||||||
Time int64
|
|
||||||
MetaEventType string
|
|
||||||
}
|
|
||||||
|
|
||||||
type oneBotAPIRequest struct {
|
type oneBotAPIRequest struct {
|
||||||
Action string `json:"action"`
|
Action string `json:"action"`
|
||||||
Params interface{} `json:"params"`
|
Params interface{} `json:"params"`
|
||||||
Echo string `json:"echo,omitempty"`
|
Echo string `json:"echo,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type oneBotSendPrivateMsgParams struct {
|
type oneBotMessageSegment struct {
|
||||||
UserID int64 `json:"user_id"`
|
Type string `json:"type"`
|
||||||
Message string `json:"message"`
|
Data map[string]interface{} `json:"data"`
|
||||||
}
|
|
||||||
|
|
||||||
type oneBotSendGroupMsgParams struct {
|
|
||||||
GroupID int64 `json:"group_id"`
|
|
||||||
Message string `json:"message"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*OneBotChannel, error) {
|
func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*OneBotChannel, error) {
|
||||||
|
|
@ -101,9 +107,30 @@ func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*One
|
||||||
dedup: make(map[string]struct{}, dedupSize),
|
dedup: make(map[string]struct{}, dedupSize),
|
||||||
dedupRing: make([]string, dedupSize),
|
dedupRing: make([]string, dedupSize),
|
||||||
dedupIdx: 0,
|
dedupIdx: 0,
|
||||||
|
pending: make(map[string]chan json.RawMessage),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
|
||||||
|
c.transcriber = transcriber
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) setMsgEmojiLike(messageID string, emojiID int, set bool) {
|
||||||
|
go func() {
|
||||||
|
_, err := c.sendAPIRequest("set_msg_emoji_like", map[string]interface{}{
|
||||||
|
"message_id": messageID,
|
||||||
|
"emoji_id": emojiID,
|
||||||
|
"set": set,
|
||||||
|
}, 5*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
logger.DebugCF("onebot", "Failed to set emoji like", map[string]interface{}{
|
||||||
|
"message_id": messageID,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) Start(ctx context.Context) error {
|
func (c *OneBotChannel) Start(ctx context.Context) error {
|
||||||
if c.config.WSUrl == "" {
|
if c.config.WSUrl == "" {
|
||||||
return fmt.Errorf("OneBot ws_url not configured")
|
return fmt.Errorf("OneBot ws_url not configured")
|
||||||
|
|
@ -121,12 +148,12 @@ func (c *OneBotChannel) Start(ctx context.Context) error {
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
go c.listen()
|
go c.listen()
|
||||||
|
c.fetchSelfID()
|
||||||
}
|
}
|
||||||
|
|
||||||
if c.config.ReconnectInterval > 0 {
|
if c.config.ReconnectInterval > 0 {
|
||||||
go c.reconnectLoop()
|
go c.reconnectLoop()
|
||||||
} else {
|
} else {
|
||||||
// If reconnect is disabled but initial connection failed, we cannot recover
|
|
||||||
if c.conn == nil {
|
if c.conn == nil {
|
||||||
return fmt.Errorf("failed to connect to OneBot and reconnect is disabled")
|
return fmt.Errorf("failed to connect to OneBot and reconnect is disabled")
|
||||||
}
|
}
|
||||||
|
|
@ -152,14 +179,141 @@ func (c *OneBotChannel) connect() error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
conn.SetPongHandler(func(appData string) error {
|
||||||
|
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
c.conn = conn
|
c.conn = conn
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
go c.pinger(conn)
|
||||||
|
|
||||||
logger.InfoC("onebot", "WebSocket connected")
|
logger.InfoC("onebot", "WebSocket connected")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) pinger(conn *websocket.Conn) {
|
||||||
|
ticker := time.NewTicker(30 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
c.writeMu.Lock()
|
||||||
|
err := conn.WriteMessage(websocket.PingMessage, nil)
|
||||||
|
c.writeMu.Unlock()
|
||||||
|
if err != nil {
|
||||||
|
logger.DebugCF("onebot", "Ping write failed, stopping pinger", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) fetchSelfID() {
|
||||||
|
resp, err := c.sendAPIRequest("get_login_info", nil, 5*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("onebot", "Failed to get_login_info", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
type loginInfo struct {
|
||||||
|
UserID json.RawMessage `json:"user_id"`
|
||||||
|
Nickname string `json:"nickname"`
|
||||||
|
}
|
||||||
|
for _, extract := range []func() (*loginInfo, error){
|
||||||
|
func() (*loginInfo, error) {
|
||||||
|
var w struct {
|
||||||
|
Data loginInfo `json:"data"`
|
||||||
|
}
|
||||||
|
err := json.Unmarshal(resp, &w)
|
||||||
|
return &w.Data, err
|
||||||
|
},
|
||||||
|
func() (*loginInfo, error) {
|
||||||
|
var f loginInfo
|
||||||
|
err := json.Unmarshal(resp, &f)
|
||||||
|
return &f, err
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
info, err := extract()
|
||||||
|
if err != nil || len(info.UserID) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if uid, err := parseJSONInt64(info.UserID); err == nil && uid > 0 {
|
||||||
|
atomic.StoreInt64(&c.selfID, uid)
|
||||||
|
logger.InfoCF("onebot", "Bot self ID retrieved", map[string]interface{}{
|
||||||
|
"self_id": uid,
|
||||||
|
"nickname": info.Nickname,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.WarnCF("onebot", "Could not parse self ID from get_login_info response", map[string]interface{}{
|
||||||
|
"response": string(resp),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) sendAPIRequest(action string, params interface{}, timeout time.Duration) (json.RawMessage, error) {
|
||||||
|
c.mu.Lock()
|
||||||
|
conn := c.conn
|
||||||
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
if conn == nil {
|
||||||
|
return nil, fmt.Errorf("WebSocket not connected")
|
||||||
|
}
|
||||||
|
|
||||||
|
echo := fmt.Sprintf("api_%d_%d", time.Now().UnixNano(), atomic.AddInt64(&c.echoCounter, 1))
|
||||||
|
|
||||||
|
ch := make(chan json.RawMessage, 1)
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
c.pending[echo] = ch
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
delete(c.pending, echo)
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
}()
|
||||||
|
|
||||||
|
req := oneBotAPIRequest{
|
||||||
|
Action: action,
|
||||||
|
Params: params,
|
||||||
|
Echo: echo,
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal API request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.writeMu.Lock()
|
||||||
|
err = conn.WriteMessage(websocket.TextMessage, data)
|
||||||
|
c.writeMu.Unlock()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to write API request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case resp := <-ch:
|
||||||
|
return resp, nil
|
||||||
|
case <-time.After(timeout):
|
||||||
|
return nil, fmt.Errorf("API request %s timed out after %v", action, timeout)
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return nil, fmt.Errorf("context cancelled")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) reconnectLoop() {
|
func (c *OneBotChannel) reconnectLoop() {
|
||||||
interval := time.Duration(c.config.ReconnectInterval) * time.Second
|
interval := time.Duration(c.config.ReconnectInterval) * time.Second
|
||||||
if interval < 5*time.Second {
|
if interval < 5*time.Second {
|
||||||
|
|
@ -183,6 +337,7 @@ func (c *OneBotChannel) reconnectLoop() {
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
go c.listen()
|
go c.listen()
|
||||||
|
c.fetchSelfID()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -197,6 +352,13 @@ func (c *OneBotChannel) Stop(ctx context.Context) error {
|
||||||
c.cancel()
|
c.cancel()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
for echo, ch := range c.pending {
|
||||||
|
close(ch)
|
||||||
|
delete(c.pending, echo)
|
||||||
|
}
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
if c.conn != nil {
|
if c.conn != nil {
|
||||||
c.conn.Close()
|
c.conn.Close()
|
||||||
|
|
@ -225,10 +387,7 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
c.writeMu.Lock()
|
echo := fmt.Sprintf("send_%d", atomic.AddInt64(&c.echoCounter, 1))
|
||||||
c.echoCounter++
|
|
||||||
echo := fmt.Sprintf("send_%d", c.echoCounter)
|
|
||||||
c.writeMu.Unlock()
|
|
||||||
|
|
||||||
req := oneBotAPIRequest{
|
req := oneBotAPIRequest{
|
||||||
Action: action,
|
Action: action,
|
||||||
|
|
@ -252,67 +411,78 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if msgID, ok := c.pendingEmojiMsg.LoadAndDelete(msg.ChatID); ok {
|
||||||
|
if mid, ok := msgID.(string); ok && mid != "" {
|
||||||
|
c.setMsgEmojiLike(mid, 289, false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) buildMessageSegments(chatID, content string) []oneBotMessageSegment {
|
||||||
|
var segments []oneBotMessageSegment
|
||||||
|
|
||||||
|
if lastMsgID, ok := c.lastMessageID.Load(chatID); ok {
|
||||||
|
if msgID, ok := lastMsgID.(string); ok && msgID != "" {
|
||||||
|
segments = append(segments, oneBotMessageSegment{
|
||||||
|
Type: "reply",
|
||||||
|
Data: map[string]interface{}{"id": msgID},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
segments = append(segments, oneBotMessageSegment{
|
||||||
|
Type: "text",
|
||||||
|
Data: map[string]interface{}{"text": content},
|
||||||
|
})
|
||||||
|
|
||||||
|
return segments
|
||||||
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) buildSendRequest(msg bus.OutboundMessage) (string, interface{}, error) {
|
func (c *OneBotChannel) buildSendRequest(msg bus.OutboundMessage) (string, interface{}, error) {
|
||||||
chatID := msg.ChatID
|
chatID := msg.ChatID
|
||||||
|
segments := c.buildMessageSegments(chatID, msg.Content)
|
||||||
|
|
||||||
if len(chatID) > 6 && chatID[:6] == "group:" {
|
var action, idKey string
|
||||||
groupID, err := strconv.ParseInt(chatID[6:], 10, 64)
|
var rawID string
|
||||||
if err != nil {
|
if rest, ok := strings.CutPrefix(chatID, "group:"); ok {
|
||||||
return "", nil, fmt.Errorf("invalid group ID in chatID: %s", chatID)
|
action, idKey, rawID = "send_group_msg", "group_id", rest
|
||||||
}
|
} else if rest, ok := strings.CutPrefix(chatID, "private:"); ok {
|
||||||
return "send_group_msg", oneBotSendGroupMsgParams{
|
action, idKey, rawID = "send_private_msg", "user_id", rest
|
||||||
GroupID: groupID,
|
} else {
|
||||||
Message: msg.Content,
|
action, idKey, rawID = "send_private_msg", "user_id", chatID
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(chatID) > 8 && chatID[:8] == "private:" {
|
id, err := strconv.ParseInt(rawID, 10, 64)
|
||||||
userID, err := strconv.ParseInt(chatID[8:], 10, 64)
|
|
||||||
if err != nil {
|
|
||||||
return "", nil, fmt.Errorf("invalid user ID in chatID: %s", chatID)
|
|
||||||
}
|
|
||||||
return "send_private_msg", oneBotSendPrivateMsgParams{
|
|
||||||
UserID: userID,
|
|
||||||
Message: msg.Content,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
userID, err := strconv.ParseInt(chatID, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", nil, fmt.Errorf("invalid chatID for OneBot: %s", chatID)
|
return "", nil, fmt.Errorf("invalid %s in chatID: %s", idKey, chatID)
|
||||||
}
|
}
|
||||||
|
return action, map[string]interface{}{idKey: id, "message": segments}, nil
|
||||||
return "send_private_msg", oneBotSendPrivateMsgParams{
|
|
||||||
UserID: userID,
|
|
||||||
Message: msg.Content,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) listen() {
|
func (c *OneBotChannel) listen() {
|
||||||
|
c.mu.Lock()
|
||||||
|
conn := c.conn
|
||||||
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
if conn == nil {
|
||||||
|
logger.WarnC("onebot", "WebSocket connection is nil, listener exiting")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-c.ctx.Done():
|
case <-c.ctx.Done():
|
||||||
return
|
return
|
||||||
default:
|
default:
|
||||||
c.mu.Lock()
|
|
||||||
conn := c.conn
|
|
||||||
c.mu.Unlock()
|
|
||||||
|
|
||||||
if conn == nil {
|
|
||||||
logger.WarnC("onebot", "WebSocket connection is nil, listener exiting")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
_, message, err := conn.ReadMessage()
|
_, message, err := conn.ReadMessage()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("onebot", "WebSocket read error", map[string]interface{}{
|
logger.ErrorCF("onebot", "WebSocket read error", map[string]interface{}{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
if c.conn != nil {
|
if c.conn == conn {
|
||||||
c.conn.Close()
|
c.conn.Close()
|
||||||
c.conn = nil
|
c.conn = nil
|
||||||
}
|
}
|
||||||
|
|
@ -320,10 +490,7 @@ func (c *OneBotChannel) listen() {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("onebot", "Raw WebSocket message received", map[string]interface{}{
|
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||||
"length": len(message),
|
|
||||||
"payload": string(message),
|
|
||||||
})
|
|
||||||
|
|
||||||
var raw oneBotRawEvent
|
var raw oneBotRawEvent
|
||||||
if err := json.Unmarshal(message, &raw); err != nil {
|
if err := json.Unmarshal(message, &raw); err != nil {
|
||||||
|
|
@ -334,20 +501,37 @@ func (c *OneBotChannel) listen() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if raw.Echo != "" || raw.Status.Online || raw.Status.Good {
|
logger.DebugCF("onebot", "WebSocket event", map[string]interface{}{
|
||||||
logger.DebugCF("onebot", "Received API response, skipping", map[string]interface{}{
|
"length": len(message),
|
||||||
"echo": raw.Echo,
|
"post_type": raw.PostType,
|
||||||
"status": raw.Status,
|
"sub_type": raw.SubType,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
if raw.Echo != "" {
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
ch, ok := c.pending[raw.Echo]
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
|
||||||
|
if ok {
|
||||||
|
select {
|
||||||
|
case ch <- message:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.DebugCF("onebot", "Received API response (no waiter)", map[string]interface{}{
|
||||||
|
"echo": raw.Echo,
|
||||||
|
"status": string(raw.Status),
|
||||||
|
})
|
||||||
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("onebot", "Parsed raw event", map[string]interface{}{
|
if isAPIResponse(raw.Status) {
|
||||||
"post_type": raw.PostType,
|
logger.DebugCF("onebot", "Received API response without echo, skipping", map[string]interface{}{
|
||||||
"message_type": raw.MessageType,
|
"status": string(raw.Status),
|
||||||
"sub_type": raw.SubType,
|
})
|
||||||
"meta_event_type": raw.MetaEventType,
|
continue
|
||||||
})
|
}
|
||||||
|
|
||||||
c.handleRawEvent(&raw)
|
c.handleRawEvent(&raw)
|
||||||
}
|
}
|
||||||
|
|
@ -386,9 +570,12 @@ func parseJSONString(raw json.RawMessage) string {
|
||||||
type parseMessageResult struct {
|
type parseMessageResult struct {
|
||||||
Text string
|
Text string
|
||||||
IsBotMentioned bool
|
IsBotMentioned bool
|
||||||
|
Media []string
|
||||||
|
LocalFiles []string
|
||||||
|
ReplyTo string
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseMessageContentEx(raw json.RawMessage, selfID int64) parseMessageResult {
|
func (c *OneBotChannel) parseMessageSegments(raw json.RawMessage, selfID int64) parseMessageResult {
|
||||||
if len(raw) == 0 {
|
if len(raw) == 0 {
|
||||||
return parseMessageResult{}
|
return parseMessageResult{}
|
||||||
}
|
}
|
||||||
|
|
@ -408,60 +595,155 @@ func parseMessageContentEx(raw json.RawMessage, selfID int64) parseMessageResult
|
||||||
}
|
}
|
||||||
|
|
||||||
var segments []map[string]interface{}
|
var segments []map[string]interface{}
|
||||||
if err := json.Unmarshal(raw, &segments); err == nil {
|
if err := json.Unmarshal(raw, &segments); err != nil {
|
||||||
var text string
|
return parseMessageResult{}
|
||||||
mentioned := false
|
}
|
||||||
selfIDStr := strconv.FormatInt(selfID, 10)
|
|
||||||
for _, seg := range segments {
|
var textParts []string
|
||||||
segType, _ := seg["type"].(string)
|
mentioned := false
|
||||||
data, _ := seg["data"].(map[string]interface{})
|
selfIDStr := strconv.FormatInt(selfID, 10)
|
||||||
switch segType {
|
var media []string
|
||||||
case "text":
|
var localFiles []string
|
||||||
if data != nil {
|
var replyTo string
|
||||||
if t, ok := data["text"].(string); ok {
|
|
||||||
text += t
|
for _, seg := range segments {
|
||||||
}
|
segType, _ := seg["type"].(string)
|
||||||
|
data, _ := seg["data"].(map[string]interface{})
|
||||||
|
|
||||||
|
switch segType {
|
||||||
|
case "text":
|
||||||
|
if data != nil {
|
||||||
|
if t, ok := data["text"].(string); ok {
|
||||||
|
textParts = append(textParts, t)
|
||||||
}
|
}
|
||||||
case "at":
|
}
|
||||||
if data != nil && selfID > 0 {
|
|
||||||
qqVal := fmt.Sprintf("%v", data["qq"])
|
case "at":
|
||||||
if qqVal == selfIDStr || qqVal == "all" {
|
if data != nil && selfID > 0 {
|
||||||
mentioned = true
|
qqVal := fmt.Sprintf("%v", data["qq"])
|
||||||
|
if qqVal == selfIDStr || qqVal == "all" {
|
||||||
|
mentioned = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case "image", "video", "file":
|
||||||
|
if data != nil {
|
||||||
|
url, _ := data["url"].(string)
|
||||||
|
if url != "" {
|
||||||
|
defaults := map[string]string{"image": "image.jpg", "video": "video.mp4", "file": "file"}
|
||||||
|
filename := defaults[segType]
|
||||||
|
if f, ok := data["file"].(string); ok && f != "" {
|
||||||
|
filename = f
|
||||||
|
} else if n, ok := data["name"].(string); ok && n != "" {
|
||||||
|
filename = n
|
||||||
|
}
|
||||||
|
localPath := utils.DownloadFile(url, filename, utils.DownloadOptions{
|
||||||
|
LoggerPrefix: "onebot",
|
||||||
|
})
|
||||||
|
if localPath != "" {
|
||||||
|
media = append(media, localPath)
|
||||||
|
localFiles = append(localFiles, localPath)
|
||||||
|
textParts = append(textParts, fmt.Sprintf("[%s]", segType))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
case "record":
|
||||||
|
if data != nil {
|
||||||
|
url, _ := data["url"].(string)
|
||||||
|
if url != "" {
|
||||||
|
localPath := utils.DownloadFile(url, "voice.amr", utils.DownloadOptions{
|
||||||
|
LoggerPrefix: "onebot",
|
||||||
|
})
|
||||||
|
if localPath != "" {
|
||||||
|
localFiles = append(localFiles, localPath)
|
||||||
|
if c.transcriber != nil && c.transcriber.IsAvailable() {
|
||||||
|
tctx, tcancel := context.WithTimeout(c.ctx, 30*time.Second)
|
||||||
|
result, err := c.transcriber.Transcribe(tctx, localPath)
|
||||||
|
tcancel()
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("onebot", "Voice transcription failed", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
textParts = append(textParts, "[voice (transcription failed)]")
|
||||||
|
media = append(media, localPath)
|
||||||
|
} else {
|
||||||
|
textParts = append(textParts, fmt.Sprintf("[voice transcription: %s]", result.Text))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
textParts = append(textParts, "[voice]")
|
||||||
|
media = append(media, localPath)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case "reply":
|
||||||
|
if data != nil {
|
||||||
|
if id, ok := data["id"]; ok {
|
||||||
|
replyTo = fmt.Sprintf("%v", id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case "face":
|
||||||
|
if data != nil {
|
||||||
|
faceID, _ := data["id"]
|
||||||
|
textParts = append(textParts, fmt.Sprintf("[face:%v]", faceID))
|
||||||
|
}
|
||||||
|
|
||||||
|
case "forward":
|
||||||
|
textParts = append(textParts, "[forward message]")
|
||||||
|
|
||||||
|
default:
|
||||||
|
|
||||||
}
|
}
|
||||||
return parseMessageResult{Text: strings.TrimSpace(text), IsBotMentioned: mentioned}
|
|
||||||
}
|
}
|
||||||
return parseMessageResult{}
|
|
||||||
|
return parseMessageResult{
|
||||||
|
Text: strings.TrimSpace(strings.Join(textParts, "")),
|
||||||
|
IsBotMentioned: mentioned,
|
||||||
|
Media: media,
|
||||||
|
LocalFiles: localFiles,
|
||||||
|
ReplyTo: replyTo,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) handleRawEvent(raw *oneBotRawEvent) {
|
func (c *OneBotChannel) handleRawEvent(raw *oneBotRawEvent) {
|
||||||
switch raw.PostType {
|
switch raw.PostType {
|
||||||
case "message":
|
case "message":
|
||||||
evt, err := c.normalizeMessageEvent(raw)
|
if userID, err := parseJSONInt64(raw.UserID); err == nil && userID > 0 {
|
||||||
if err != nil {
|
if !c.IsAllowed(strconv.FormatInt(userID, 10)) {
|
||||||
logger.WarnCF("onebot", "Failed to normalize message event", map[string]interface{}{
|
logger.DebugCF("onebot", "Message rejected by allowlist", map[string]interface{}{
|
||||||
"error": err.Error(),
|
"user_id": userID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
c.handleMessage(evt)
|
c.handleMessage(raw)
|
||||||
|
|
||||||
|
case "message_sent":
|
||||||
|
logger.DebugCF("onebot", "Bot sent message event", map[string]interface{}{
|
||||||
|
"message_type": raw.MessageType,
|
||||||
|
"message_id": parseJSONString(raw.MessageID),
|
||||||
|
})
|
||||||
|
|
||||||
case "meta_event":
|
case "meta_event":
|
||||||
c.handleMetaEvent(raw)
|
c.handleMetaEvent(raw)
|
||||||
|
|
||||||
case "notice":
|
case "notice":
|
||||||
logger.DebugCF("onebot", "Notice event received", map[string]interface{}{
|
c.handleNoticeEvent(raw)
|
||||||
"sub_type": raw.SubType,
|
|
||||||
})
|
|
||||||
case "request":
|
case "request":
|
||||||
logger.DebugCF("onebot", "Request event received", map[string]interface{}{
|
logger.DebugCF("onebot", "Request event received", map[string]interface{}{
|
||||||
"sub_type": raw.SubType,
|
"sub_type": raw.SubType,
|
||||||
})
|
})
|
||||||
|
|
||||||
case "":
|
case "":
|
||||||
logger.DebugCF("onebot", "Event with empty post_type (possibly API response)", map[string]interface{}{
|
logger.DebugCF("onebot", "Event with empty post_type (possibly API response)", map[string]interface{}{
|
||||||
"echo": raw.Echo,
|
"echo": raw.Echo,
|
||||||
"status": raw.Status,
|
"status": raw.Status,
|
||||||
})
|
})
|
||||||
|
|
||||||
default:
|
default:
|
||||||
logger.DebugCF("onebot", "Unknown post_type", map[string]interface{}{
|
logger.DebugCF("onebot", "Unknown post_type", map[string]interface{}{
|
||||||
"post_type": raw.PostType,
|
"post_type": raw.PostType,
|
||||||
|
|
@ -469,18 +751,51 @@ func (c *OneBotChannel) handleRawEvent(raw *oneBotRawEvent) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) normalizeMessageEvent(raw *oneBotRawEvent) (*oneBotEvent, error) {
|
func (c *OneBotChannel) handleMetaEvent(raw *oneBotRawEvent) {
|
||||||
|
if raw.MetaEventType == "lifecycle" {
|
||||||
|
logger.InfoCF("onebot", "Lifecycle event", map[string]interface{}{"sub_type": raw.SubType})
|
||||||
|
} else if raw.MetaEventType != "heartbeat" {
|
||||||
|
logger.DebugCF("onebot", "Meta event: "+raw.MetaEventType, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) handleNoticeEvent(raw *oneBotRawEvent) {
|
||||||
|
fields := map[string]interface{}{
|
||||||
|
"notice_type": raw.NoticeType,
|
||||||
|
"sub_type": raw.SubType,
|
||||||
|
"group_id": parseJSONString(raw.GroupID),
|
||||||
|
"user_id": parseJSONString(raw.UserID),
|
||||||
|
"message_id": parseJSONString(raw.MessageID),
|
||||||
|
}
|
||||||
|
switch raw.NoticeType {
|
||||||
|
case "group_recall", "group_increase", "group_decrease",
|
||||||
|
"friend_add", "group_admin", "group_ban":
|
||||||
|
logger.InfoCF("onebot", "Notice: "+raw.NoticeType, fields)
|
||||||
|
default:
|
||||||
|
logger.DebugCF("onebot", "Notice: "+raw.NoticeType, fields)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) handleMessage(raw *oneBotRawEvent) {
|
||||||
|
// Parse fields from raw event
|
||||||
userID, err := parseJSONInt64(raw.UserID)
|
userID, err := parseJSONInt64(raw.UserID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("parse user_id: %w (raw: %s)", err, string(raw.UserID))
|
logger.WarnCF("onebot", "Failed to parse user_id", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
"raw": string(raw.UserID),
|
||||||
|
})
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
groupID, _ := parseJSONInt64(raw.GroupID)
|
groupID, _ := parseJSONInt64(raw.GroupID)
|
||||||
selfID, _ := parseJSONInt64(raw.SelfID)
|
selfID, _ := parseJSONInt64(raw.SelfID)
|
||||||
ts, _ := parseJSONInt64(raw.Time)
|
|
||||||
messageID := parseJSONString(raw.MessageID)
|
messageID := parseJSONString(raw.MessageID)
|
||||||
|
|
||||||
parsed := parseMessageContentEx(raw.Message, selfID)
|
if selfID == 0 {
|
||||||
|
selfID = atomic.LoadInt64(&c.selfID)
|
||||||
|
}
|
||||||
|
|
||||||
|
parsed := c.parseMessageSegments(raw.Message, selfID)
|
||||||
isBotMentioned := parsed.IsBotMentioned
|
isBotMentioned := parsed.IsBotMentioned
|
||||||
|
|
||||||
content := raw.RawMessage
|
content := raw.RawMessage
|
||||||
|
|
@ -495,6 +810,10 @@ func (c *OneBotChannel) normalizeMessageEvent(raw *oneBotRawEvent) (*oneBotEvent
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if parsed.Text != "" && content != parsed.Text && (len(parsed.Media) > 0 || parsed.ReplyTo != "") {
|
||||||
|
content = parsed.Text
|
||||||
|
}
|
||||||
|
|
||||||
var sender oneBotSender
|
var sender oneBotSender
|
||||||
if len(raw.Sender) > 0 {
|
if len(raw.Sender) > 0 {
|
||||||
if err := json.Unmarshal(raw.Sender, &sender); err != nil {
|
if err := json.Unmarshal(raw.Sender, &sender); err != nil {
|
||||||
|
|
@ -505,137 +824,111 @@ func (c *OneBotChannel) normalizeMessageEvent(raw *oneBotRawEvent) (*oneBotEvent
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("onebot", "Normalized message event", map[string]interface{}{
|
// Clean up temp files when done
|
||||||
"message_type": raw.MessageType,
|
if len(parsed.LocalFiles) > 0 {
|
||||||
"user_id": userID,
|
defer func() {
|
||||||
"group_id": groupID,
|
for _, f := range parsed.LocalFiles {
|
||||||
"message_id": messageID,
|
if err := os.Remove(f); err != nil {
|
||||||
"content_len": len(content),
|
logger.DebugCF("onebot", "Failed to remove temp file", map[string]interface{}{
|
||||||
"nickname": sender.Nickname,
|
"path": f,
|
||||||
})
|
"error": err.Error(),
|
||||||
|
})
|
||||||
return &oneBotEvent{
|
}
|
||||||
PostType: raw.PostType,
|
}
|
||||||
MessageType: raw.MessageType,
|
}()
|
||||||
SubType: raw.SubType,
|
|
||||||
MessageID: messageID,
|
|
||||||
UserID: userID,
|
|
||||||
GroupID: groupID,
|
|
||||||
Content: content,
|
|
||||||
RawContent: raw.RawMessage,
|
|
||||||
IsBotMentioned: isBotMentioned,
|
|
||||||
Sender: sender,
|
|
||||||
SelfID: selfID,
|
|
||||||
Time: ts,
|
|
||||||
MetaEventType: raw.MetaEventType,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *OneBotChannel) handleMetaEvent(raw *oneBotRawEvent) {
|
|
||||||
switch raw.MetaEventType {
|
|
||||||
case "lifecycle":
|
|
||||||
logger.InfoCF("onebot", "Lifecycle event", map[string]interface{}{
|
|
||||||
"sub_type": raw.SubType,
|
|
||||||
})
|
|
||||||
case "heartbeat":
|
|
||||||
logger.DebugC("onebot", "Heartbeat received")
|
|
||||||
default:
|
|
||||||
logger.DebugCF("onebot", "Unknown meta_event_type", map[string]interface{}{
|
|
||||||
"meta_event_type": raw.MetaEventType,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
func (c *OneBotChannel) handleMessage(evt *oneBotEvent) {
|
if c.isDuplicate(messageID) {
|
||||||
if c.isDuplicate(evt.MessageID) {
|
|
||||||
logger.DebugCF("onebot", "Duplicate message, skipping", map[string]interface{}{
|
logger.DebugCF("onebot", "Duplicate message, skipping", map[string]interface{}{
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
content := evt.Content
|
|
||||||
if content == "" {
|
if content == "" {
|
||||||
logger.DebugCF("onebot", "Received empty message, ignoring", map[string]interface{}{
|
logger.DebugCF("onebot", "Received empty message, ignoring", map[string]interface{}{
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
senderID := strconv.FormatInt(evt.UserID, 10)
|
senderID := strconv.FormatInt(userID, 10)
|
||||||
var chatID string
|
var chatID string
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
}
|
}
|
||||||
|
|
||||||
switch evt.MessageType {
|
if parsed.ReplyTo != "" {
|
||||||
|
metadata["reply_to_message_id"] = parsed.ReplyTo
|
||||||
|
}
|
||||||
|
|
||||||
|
switch raw.MessageType {
|
||||||
case "private":
|
case "private":
|
||||||
chatID = "private:" + senderID
|
chatID = "private:" + senderID
|
||||||
logger.InfoCF("onebot", "Received private message", map[string]interface{}{
|
metadata["peer_kind"] = "direct"
|
||||||
"sender": senderID,
|
metadata["peer_id"] = senderID
|
||||||
"message_id": evt.MessageID,
|
|
||||||
"length": len(content),
|
|
||||||
"content": truncate(content, 100),
|
|
||||||
})
|
|
||||||
|
|
||||||
case "group":
|
case "group":
|
||||||
groupIDStr := strconv.FormatInt(evt.GroupID, 10)
|
groupIDStr := strconv.FormatInt(groupID, 10)
|
||||||
chatID = "group:" + groupIDStr
|
chatID = "group:" + groupIDStr
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = groupIDStr
|
||||||
metadata["group_id"] = groupIDStr
|
metadata["group_id"] = groupIDStr
|
||||||
|
|
||||||
senderUserID, _ := parseJSONInt64(evt.Sender.UserID)
|
senderUserID, _ := parseJSONInt64(sender.UserID)
|
||||||
if senderUserID > 0 {
|
if senderUserID > 0 {
|
||||||
metadata["sender_user_id"] = strconv.FormatInt(senderUserID, 10)
|
metadata["sender_user_id"] = strconv.FormatInt(senderUserID, 10)
|
||||||
}
|
}
|
||||||
|
|
||||||
if evt.Sender.Card != "" {
|
if sender.Card != "" {
|
||||||
metadata["sender_name"] = evt.Sender.Card
|
metadata["sender_name"] = sender.Card
|
||||||
} else if evt.Sender.Nickname != "" {
|
} else if sender.Nickname != "" {
|
||||||
metadata["sender_name"] = evt.Sender.Nickname
|
metadata["sender_name"] = sender.Nickname
|
||||||
}
|
}
|
||||||
|
|
||||||
triggered, strippedContent := c.checkGroupTrigger(content, evt.IsBotMentioned)
|
triggered, strippedContent := c.checkGroupTrigger(content, isBotMentioned)
|
||||||
if !triggered {
|
if !triggered {
|
||||||
logger.DebugCF("onebot", "Group message ignored (no trigger)", map[string]interface{}{
|
logger.DebugCF("onebot", "Group message ignored (no trigger)", map[string]interface{}{
|
||||||
"sender": senderID,
|
"sender": senderID,
|
||||||
"group": groupIDStr,
|
"group": groupIDStr,
|
||||||
"is_mentioned": evt.IsBotMentioned,
|
"is_mentioned": isBotMentioned,
|
||||||
"content": truncate(content, 100),
|
"content": truncate(content, 100),
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
content = strippedContent
|
content = strippedContent
|
||||||
|
|
||||||
logger.InfoCF("onebot", "Received group message", map[string]interface{}{
|
|
||||||
"sender": senderID,
|
|
||||||
"group": groupIDStr,
|
|
||||||
"message_id": evt.MessageID,
|
|
||||||
"is_mentioned": evt.IsBotMentioned,
|
|
||||||
"length": len(content),
|
|
||||||
"content": truncate(content, 100),
|
|
||||||
})
|
|
||||||
|
|
||||||
default:
|
default:
|
||||||
logger.WarnCF("onebot", "Unknown message type, cannot route", map[string]interface{}{
|
logger.WarnCF("onebot", "Unknown message type, cannot route", map[string]interface{}{
|
||||||
"type": evt.MessageType,
|
"type": raw.MessageType,
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
"user_id": evt.UserID,
|
"user_id": userID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if evt.Sender.Nickname != "" {
|
logger.InfoCF("onebot", "Received "+raw.MessageType+" message", map[string]interface{}{
|
||||||
metadata["nickname"] = evt.Sender.Nickname
|
"sender": senderID,
|
||||||
}
|
"chat_id": chatID,
|
||||||
|
"message_id": messageID,
|
||||||
logger.DebugCF("onebot", "Forwarding message to bus", map[string]interface{}{
|
"length": len(content),
|
||||||
"sender_id": senderID,
|
"content": truncate(content, 100),
|
||||||
"chat_id": chatID,
|
"media_count": len(parsed.Media),
|
||||||
"content": truncate(content, 100),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, []string{}, metadata)
|
if sender.Nickname != "" {
|
||||||
|
metadata["nickname"] = sender.Nickname
|
||||||
|
}
|
||||||
|
|
||||||
|
c.lastMessageID.Store(chatID, messageID)
|
||||||
|
|
||||||
|
if raw.MessageType == "group" && messageID != "" && messageID != "0" {
|
||||||
|
c.setMsgEmojiLike(messageID, 289, true)
|
||||||
|
c.pendingEmojiMsg.Store(chatID, messageID)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.HandleMessage(senderID, chatID, content, parsed.Media, metadata)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) isDuplicate(messageID string) bool {
|
func (c *OneBotChannel) isDuplicate(messageID string) bool {
|
||||||
|
|
|
||||||
|
|
@ -165,6 +165,8 @@ func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
// 转发到消息总线
|
// 转发到消息总线
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": data.ID,
|
"message_id": data.ID,
|
||||||
|
"peer_kind": "direct",
|
||||||
|
"peer_id": senderID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, senderID, content, []string{}, metadata)
|
c.HandleMessage(senderID, senderID, content, []string{}, metadata)
|
||||||
|
|
@ -207,6 +209,8 @@ func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": data.ID,
|
"message_id": data.ID,
|
||||||
"group_id": data.GroupID,
|
"group_id": data.GroupID,
|
||||||
|
"peer_kind": "group",
|
||||||
|
"peer_id": data.GroupID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, data.GroupID, content, []string{}, metadata)
|
c.HandleMessage(senderID, data.GroupID, content, []string{}, metadata)
|
||||||
|
|
|
||||||
|
|
@ -59,6 +59,13 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
Proxy: http.ProxyURL(proxyURL),
|
||||||
},
|
},
|
||||||
}))
|
}))
|
||||||
|
} else if os.Getenv("HTTP_PROXY") != "" || os.Getenv("HTTPS_PROXY") != "" {
|
||||||
|
// Use environment proxy if configured
|
||||||
|
opts = append(opts, telego.WithHTTPClient(&http.Client{
|
||||||
|
Transport: &http.Transport{
|
||||||
|
Proxy: http.ProxyFromEnvironment,
|
||||||
|
},
|
||||||
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
||||||
|
|
|
||||||
|
|
@ -178,6 +178,14 @@ func (c *WhatsAppChannel) handleIncomingMessage(msg map[string]interface{}) {
|
||||||
metadata["user_name"] = userName
|
metadata["user_name"] = userName
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if chatID == senderID {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = chatID
|
||||||
|
}
|
||||||
|
|
||||||
log.Printf("WhatsApp message from %s: %s...", senderID, utils.Truncate(content, 50))
|
log.Printf("WhatsApp message from %s: %s...", senderID, utils.Truncate(content, 50))
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, mediaPaths, metadata)
|
c.HandleMessage(senderID, chatID, content, mediaPaths, metadata)
|
||||||
|
|
|
||||||
|
|
@ -6,11 +6,14 @@ import (
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/caarlos0/env/v11"
|
"github.com/caarlos0/env/v11"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// rrCounter is a global counter for round-robin load balancing across models.
|
||||||
|
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.
|
||||||
type FlexibleStringSlice []string
|
type FlexibleStringSlice []string
|
||||||
|
|
@ -49,12 +52,37 @@ type Config struct {
|
||||||
Bindings []AgentBinding `json:"bindings,omitempty"`
|
Bindings []AgentBinding `json:"bindings,omitempty"`
|
||||||
Session SessionConfig `json:"session,omitempty"`
|
Session SessionConfig `json:"session,omitempty"`
|
||||||
Channels ChannelsConfig `json:"channels"`
|
Channels ChannelsConfig `json:"channels"`
|
||||||
Providers ProvidersConfig `json:"providers"`
|
Providers ProvidersConfig `json:"providers,omitempty"`
|
||||||
|
ModelList []ModelConfig `json:"model_list"` // New model-centric provider configuration
|
||||||
Gateway GatewayConfig `json:"gateway"`
|
Gateway GatewayConfig `json:"gateway"`
|
||||||
Tools ToolsConfig `json:"tools"`
|
Tools ToolsConfig `json:"tools"`
|
||||||
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
||||||
Devices DevicesConfig `json:"devices"`
|
Devices DevicesConfig `json:"devices"`
|
||||||
mu sync.RWMutex
|
}
|
||||||
|
|
||||||
|
// MarshalJSON implements custom JSON marshaling for Config
|
||||||
|
// to omit providers section when empty and session when empty
|
||||||
|
func (c Config) MarshalJSON() ([]byte, error) {
|
||||||
|
type Alias Config
|
||||||
|
aux := &struct {
|
||||||
|
Providers *ProvidersConfig `json:"providers,omitempty"`
|
||||||
|
Session *SessionConfig `json:"session,omitempty"`
|
||||||
|
*Alias
|
||||||
|
}{
|
||||||
|
Alias: (*Alias)(&c),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only include providers if not empty
|
||||||
|
if !c.Providers.IsEmpty() {
|
||||||
|
aux.Providers = &c.Providers
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only include session if not empty
|
||||||
|
if c.Session.DMScope != "" || len(c.Session.IdentityLinks) > 0 {
|
||||||
|
aux.Session = &c.Session
|
||||||
|
}
|
||||||
|
|
||||||
|
return json.Marshal(aux)
|
||||||
}
|
}
|
||||||
|
|
||||||
type AgentsConfig struct {
|
type AgentsConfig struct {
|
||||||
|
|
@ -148,7 +176,7 @@ type AgentDefaults struct {
|
||||||
ImageModel string `json:"image_model,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_IMAGE_MODEL"`
|
ImageModel string `json:"image_model,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_IMAGE_MODEL"`
|
||||||
ImageModelFallbacks []string `json:"image_model_fallbacks,omitempty"`
|
ImageModelFallbacks []string `json:"image_model_fallbacks,omitempty"`
|
||||||
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
||||||
Temperature float64 `json:"temperature" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
||||||
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -188,9 +216,10 @@ type FeishuConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type DiscordConfig struct {
|
type DiscordConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_DISCORD_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_DISCORD_ENABLED"`
|
||||||
Token string `json:"token" env:"PICOCLAW_CHANNELS_DISCORD_TOKEN"`
|
Token string `json:"token" env:"PICOCLAW_CHANNELS_DISCORD_TOKEN"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_DISCORD_ALLOW_FROM"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_DISCORD_ALLOW_FROM"`
|
||||||
|
MentionOnly bool `json:"mention_only" env:"PICOCLAW_CHANNELS_DISCORD_MENTION_ONLY"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type MaixCamConfig struct {
|
type MaixCamConfig struct {
|
||||||
|
|
@ -263,7 +292,43 @@ type ProvidersConfig struct {
|
||||||
Moonshot ProviderConfig `json:"moonshot"`
|
Moonshot ProviderConfig `json:"moonshot"`
|
||||||
ShengSuanYun ProviderConfig `json:"shengsuanyun"`
|
ShengSuanYun ProviderConfig `json:"shengsuanyun"`
|
||||||
DeepSeek ProviderConfig `json:"deepseek"`
|
DeepSeek ProviderConfig `json:"deepseek"`
|
||||||
|
Cerebras ProviderConfig `json:"cerebras"`
|
||||||
|
VolcEngine ProviderConfig `json:"volcengine"`
|
||||||
GitHubCopilot ProviderConfig `json:"github_copilot"`
|
GitHubCopilot ProviderConfig `json:"github_copilot"`
|
||||||
|
Antigravity ProviderConfig `json:"antigravity"`
|
||||||
|
Qwen ProviderConfig `json:"qwen"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsEmpty checks if all provider configs are empty (no API keys or API bases set)
|
||||||
|
// Note: WebSearch is an optimization option and doesn't count as "non-empty"
|
||||||
|
func (p ProvidersConfig) IsEmpty() bool {
|
||||||
|
return p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" &&
|
||||||
|
p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" &&
|
||||||
|
p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" &&
|
||||||
|
p.Groq.APIKey == "" && p.Groq.APIBase == "" &&
|
||||||
|
p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" &&
|
||||||
|
p.VLLM.APIKey == "" && p.VLLM.APIBase == "" &&
|
||||||
|
p.Gemini.APIKey == "" && p.Gemini.APIBase == "" &&
|
||||||
|
p.Nvidia.APIKey == "" && p.Nvidia.APIBase == "" &&
|
||||||
|
p.Ollama.APIKey == "" && p.Ollama.APIBase == "" &&
|
||||||
|
p.Moonshot.APIKey == "" && p.Moonshot.APIBase == "" &&
|
||||||
|
p.ShengSuanYun.APIKey == "" && p.ShengSuanYun.APIBase == "" &&
|
||||||
|
p.DeepSeek.APIKey == "" && p.DeepSeek.APIBase == "" &&
|
||||||
|
p.Cerebras.APIKey == "" && p.Cerebras.APIBase == "" &&
|
||||||
|
p.VolcEngine.APIKey == "" && p.VolcEngine.APIBase == "" &&
|
||||||
|
p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" &&
|
||||||
|
p.Antigravity.APIKey == "" && p.Antigravity.APIBase == "" &&
|
||||||
|
p.Qwen.APIKey == "" && p.Qwen.APIBase == ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalJSON implements custom JSON marshaling for ProvidersConfig
|
||||||
|
// to omit the entire section when empty
|
||||||
|
func (p ProvidersConfig) MarshalJSON() ([]byte, error) {
|
||||||
|
if p.IsEmpty() {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
type Alias ProvidersConfig
|
||||||
|
return json.Marshal((*Alias)(&p))
|
||||||
}
|
}
|
||||||
|
|
||||||
type ProviderConfig struct {
|
type ProviderConfig struct {
|
||||||
|
|
@ -279,164 +344,122 @@ type OpenAIProviderConfig struct {
|
||||||
WebSearch bool `json:"web_search" env:"PICOCLAW_PROVIDERS_OPENAI_WEB_SEARCH"`
|
WebSearch bool `json:"web_search" env:"PICOCLAW_PROVIDERS_OPENAI_WEB_SEARCH"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ModelConfig represents a model-centric provider configuration.
|
||||||
|
// It allows adding new providers (especially OpenAI-compatible ones) via configuration only.
|
||||||
|
// The model field uses protocol prefix format: [protocol/]model-identifier
|
||||||
|
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot
|
||||||
|
// Default protocol is "openai" if no prefix is specified.
|
||||||
|
type ModelConfig struct {
|
||||||
|
// Required fields
|
||||||
|
ModelName string `json:"model_name"` // User-facing alias for the model
|
||||||
|
Model string `json:"model"` // Protocol/model-identifier (e.g., "openai/gpt-4o", "anthropic/claude-sonnet-4.6")
|
||||||
|
|
||||||
|
// HTTP-based providers
|
||||||
|
APIBase string `json:"api_base,omitempty"` // API endpoint URL
|
||||||
|
APIKey string `json:"api_key"` // API authentication key
|
||||||
|
Proxy string `json:"proxy,omitempty"` // HTTP proxy URL
|
||||||
|
|
||||||
|
// Special providers (CLI-based, OAuth, etc.)
|
||||||
|
AuthMethod string `json:"auth_method,omitempty"` // Authentication method: oauth, token
|
||||||
|
ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc
|
||||||
|
Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers
|
||||||
|
|
||||||
|
// Optional optimizations
|
||||||
|
RPM int `json:"rpm,omitempty"` // Requests per minute limit
|
||||||
|
MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate checks if the ModelConfig has all required fields.
|
||||||
|
func (c *ModelConfig) Validate() error {
|
||||||
|
if c.ModelName == "" {
|
||||||
|
return fmt.Errorf("model_name is required")
|
||||||
|
}
|
||||||
|
if c.Model == "" {
|
||||||
|
return fmt.Errorf("model is required")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
type GatewayConfig struct {
|
type GatewayConfig struct {
|
||||||
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// WebSearchConfig defines a single active web search provider
|
type BraveConfig struct {
|
||||||
// Falls back to DuckDuckGo if the primary provider fails
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
||||||
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
||||||
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type DuckDuckGoConfig struct {
|
||||||
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_DUCKDUCKGO_ENABLED"`
|
||||||
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_DUCKDUCKGO_MAX_RESULTS"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type PerplexityConfig struct {
|
||||||
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"`
|
||||||
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEY"`
|
||||||
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// WebSearchConfig is the unified web search provider config.
|
||||||
type WebSearchConfig struct {
|
type WebSearchConfig struct {
|
||||||
Provider string `json:"provider" env:"PICOCLAW_TOOLS_WEB_SEARCH_PROVIDER"` // "brave", "ollama", custom URL
|
Provider string `json:"provider" env:"PICOCLAW_TOOLS_WEB_SEARCH_PROVIDER"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_SEARCH_API_KEY"` // API key for Brave, etc.
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_SEARCH_API_KEY"`
|
||||||
Endpoint string `json:"endpoint" env:"PICOCLAW_TOOLS_WEB_SEARCH_ENDPOINT"` // Base URL for Ollama or custom
|
Endpoint string `json:"endpoint" env:"PICOCLAW_TOOLS_WEB_SEARCH_ENDPOINT"`
|
||||||
RestType string `json:"rest_type" env:"PICOCLAW_TOOLS_WEB_SEARCH_REST_TYPE"` // "GET" or "POST"
|
RestType string `json:"rest_type" env:"PICOCLAW_TOOLS_WEB_SEARCH_REST_TYPE"`
|
||||||
QueryParam string `json:"query_param" env:"PICOCLAW_TOOLS_WEB_SEARCH_QUERY_PARAM"` // "q", "query", etc.
|
QueryParam string `json:"query_param" env:"PICOCLAW_TOOLS_WEB_SEARCH_QUERY_PARAM"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_SEARCH_MAX_RESULTS"` // Default: 5
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_SEARCH_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type WebToolsConfig struct {
|
type WebToolsConfig struct {
|
||||||
Search WebSearchConfig `json:"search"`
|
Search WebSearchConfig `json:"search"`
|
||||||
|
Brave BraveConfig `json:"brave"`
|
||||||
|
DuckDuckGo DuckDuckGoConfig `json:"duckduckgo"`
|
||||||
|
Perplexity PerplexityConfig `json:"perplexity"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type CronToolsConfig struct {
|
type CronToolsConfig struct {
|
||||||
ExecTimeoutMinutes int `json:"exec_timeout_minutes" env:"PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES"`
|
ExecTimeoutMinutes int `json:"exec_timeout_minutes" env:"PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES"` // 0 means no timeout
|
||||||
}
|
}
|
||||||
|
|
||||||
type ExecConfig struct {
|
type ExecConfig struct {
|
||||||
EnableDenyPatterns bool `json:"enable_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS"`
|
EnableDenyPatterns bool `json:"enable_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS"`
|
||||||
CustomDenyPatterns []string `json:"custom_deny_patterns,omitempty" env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS"`
|
CustomDenyPatterns []string `json:"custom_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolsConfig struct {
|
type ToolsConfig struct {
|
||||||
Web WebToolsConfig `json:"web"`
|
Web WebToolsConfig `json:"web"`
|
||||||
Cron CronToolsConfig `json:"cron"`
|
Cron CronToolsConfig `json:"cron"`
|
||||||
Exec ExecConfig `json:"exec"`
|
Exec ExecConfig `json:"exec"`
|
||||||
|
Skills SkillsToolsConfig `json:"skills"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func DefaultConfig() *Config {
|
type SkillsToolsConfig struct {
|
||||||
return &Config{
|
Registries SkillsRegistriesConfig `json:"registries"`
|
||||||
Agents: AgentsConfig{
|
MaxConcurrentSearches int `json:"max_concurrent_searches" env:"PICOCLAW_SKILLS_MAX_CONCURRENT_SEARCHES"`
|
||||||
Defaults: AgentDefaults{
|
SearchCache SearchCacheConfig `json:"search_cache"`
|
||||||
Workspace: "~/.picoclaw/workspace",
|
}
|
||||||
RestrictToWorkspace: true,
|
|
||||||
Provider: "",
|
type SearchCacheConfig struct {
|
||||||
Model: "glm-4.7",
|
MaxSize int `json:"max_size" env:"PICOCLAW_SKILLS_SEARCH_CACHE_MAX_SIZE"`
|
||||||
MaxTokens: 8192,
|
TTLSeconds int `json:"ttl_seconds" env:"PICOCLAW_SKILLS_SEARCH_CACHE_TTL_SECONDS"`
|
||||||
Temperature: 0.7,
|
}
|
||||||
MaxToolIterations: 20,
|
|
||||||
},
|
type SkillsRegistriesConfig struct {
|
||||||
},
|
ClawHub ClawHubRegistryConfig `json:"clawhub"`
|
||||||
Channels: ChannelsConfig{
|
}
|
||||||
WhatsApp: WhatsAppConfig{
|
|
||||||
Enabled: false,
|
type ClawHubRegistryConfig struct {
|
||||||
BridgeURL: "ws://localhost:3001",
|
Enabled bool `json:"enabled" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_ENABLED"`
|
||||||
AllowFrom: FlexibleStringSlice{},
|
BaseURL string `json:"base_url" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_BASE_URL"`
|
||||||
},
|
AuthToken string `json:"auth_token" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_AUTH_TOKEN"`
|
||||||
Telegram: TelegramConfig{
|
SearchPath string `json:"search_path" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_SEARCH_PATH"`
|
||||||
Enabled: false,
|
SkillsPath string `json:"skills_path" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_SKILLS_PATH"`
|
||||||
Token: "",
|
DownloadPath string `json:"download_path" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_DOWNLOAD_PATH"`
|
||||||
AllowFrom: FlexibleStringSlice{},
|
Timeout int `json:"timeout" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_TIMEOUT"`
|
||||||
},
|
MaxZipSize int `json:"max_zip_size" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_MAX_ZIP_SIZE"`
|
||||||
Feishu: FeishuConfig{
|
MaxResponseSize int `json:"max_response_size" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_MAX_RESPONSE_SIZE"`
|
||||||
Enabled: false,
|
|
||||||
AppID: "",
|
|
||||||
AppSecret: "",
|
|
||||||
EncryptKey: "",
|
|
||||||
VerificationToken: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
Discord: DiscordConfig{
|
|
||||||
Enabled: false,
|
|
||||||
Token: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
MaixCam: MaixCamConfig{
|
|
||||||
Enabled: false,
|
|
||||||
Host: "0.0.0.0",
|
|
||||||
Port: 18790,
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
QQ: QQConfig{
|
|
||||||
Enabled: false,
|
|
||||||
AppID: "",
|
|
||||||
AppSecret: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
DingTalk: DingTalkConfig{
|
|
||||||
Enabled: false,
|
|
||||||
ClientID: "",
|
|
||||||
ClientSecret: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
Slack: SlackConfig{
|
|
||||||
Enabled: false,
|
|
||||||
BotToken: "",
|
|
||||||
AppToken: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
LINE: LINEConfig{
|
|
||||||
Enabled: false,
|
|
||||||
ChannelSecret: "",
|
|
||||||
ChannelAccessToken: "",
|
|
||||||
WebhookHost: "0.0.0.0",
|
|
||||||
WebhookPort: 18791,
|
|
||||||
WebhookPath: "/webhook/line",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
OneBot: OneBotConfig{
|
|
||||||
Enabled: false,
|
|
||||||
WSUrl: "ws://127.0.0.1:3001",
|
|
||||||
AccessToken: "",
|
|
||||||
ReconnectInterval: 5,
|
|
||||||
GroupTriggerPrefix: []string{},
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Providers: ProvidersConfig{
|
|
||||||
Anthropic: ProviderConfig{},
|
|
||||||
OpenAI: OpenAIProviderConfig{WebSearch: true},
|
|
||||||
OpenRouter: ProviderConfig{},
|
|
||||||
Groq: ProviderConfig{},
|
|
||||||
Zhipu: ProviderConfig{},
|
|
||||||
VLLM: ProviderConfig{},
|
|
||||||
Gemini: ProviderConfig{},
|
|
||||||
Nvidia: ProviderConfig{},
|
|
||||||
Moonshot: ProviderConfig{},
|
|
||||||
ShengSuanYun: ProviderConfig{},
|
|
||||||
},
|
|
||||||
Gateway: GatewayConfig{
|
|
||||||
Host: "0.0.0.0",
|
|
||||||
Port: 18790,
|
|
||||||
},
|
|
||||||
Tools: ToolsConfig{
|
|
||||||
Web: WebToolsConfig{
|
|
||||||
Search: WebSearchConfig{
|
|
||||||
Provider: "",
|
|
||||||
APIKey: "",
|
|
||||||
Endpoint: "",
|
|
||||||
RestType: "",
|
|
||||||
QueryParam: "",
|
|
||||||
MaxResults: 0,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Cron: CronToolsConfig{
|
|
||||||
ExecTimeoutMinutes: 5,
|
|
||||||
},
|
|
||||||
Exec: ExecConfig{
|
|
||||||
EnableDenyPatterns: true,
|
|
||||||
CustomDenyPatterns: []string{},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Heartbeat: HeartbeatConfig{
|
|
||||||
Enabled: true,
|
|
||||||
Interval: 30, // default 30 minutes
|
|
||||||
},
|
|
||||||
Devices: DevicesConfig{
|
|
||||||
Enabled: false,
|
|
||||||
MonitorUSB: true,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func LoadConfig(path string) (*Config, error) {
|
func LoadConfig(path string) (*Config, error) {
|
||||||
|
|
@ -462,13 +485,20 @@ func LoadConfig(path string) (*Config, error) {
|
||||||
return nil, fmt.Errorf("Please check new config for web search as config example")
|
return nil, fmt.Errorf("Please check new config for web search as config example")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Auto-migrate: if only legacy providers config exists, convert to model_list
|
||||||
|
if len(cfg.ModelList) == 0 && cfg.HasProvidersConfig() {
|
||||||
|
cfg.ModelList = ConvertProvidersToModelList(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate model_list for uniqueness and required fields
|
||||||
|
if err := cfg.ValidateModelList(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
return cfg, nil
|
return cfg, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func SaveConfig(path string, cfg *Config) error {
|
func SaveConfig(path string, cfg *Config) error {
|
||||||
cfg.mu.RLock()
|
|
||||||
defer cfg.mu.RUnlock()
|
|
||||||
|
|
||||||
data, err := json.MarshalIndent(cfg, "", " ")
|
data, err := json.MarshalIndent(cfg, "", " ")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
@ -483,14 +513,10 @@ func SaveConfig(path string, cfg *Config) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) WorkspacePath() string {
|
func (c *Config) WorkspacePath() string {
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
return expandHome(c.Agents.Defaults.Workspace)
|
return expandHome(c.Agents.Defaults.Workspace)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) GetAPIKey() string {
|
func (c *Config) GetAPIKey() string {
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
if c.Providers.OpenRouter.APIKey != "" {
|
if c.Providers.OpenRouter.APIKey != "" {
|
||||||
return c.Providers.OpenRouter.APIKey
|
return c.Providers.OpenRouter.APIKey
|
||||||
}
|
}
|
||||||
|
|
@ -515,12 +541,13 @@ func (c *Config) GetAPIKey() string {
|
||||||
if c.Providers.ShengSuanYun.APIKey != "" {
|
if c.Providers.ShengSuanYun.APIKey != "" {
|
||||||
return c.Providers.ShengSuanYun.APIKey
|
return c.Providers.ShengSuanYun.APIKey
|
||||||
}
|
}
|
||||||
|
if c.Providers.Cerebras.APIKey != "" {
|
||||||
|
return c.Providers.Cerebras.APIKey
|
||||||
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) GetAPIBase() string {
|
func (c *Config) GetAPIBase() string {
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
if c.Providers.OpenRouter.APIKey != "" {
|
if c.Providers.OpenRouter.APIKey != "" {
|
||||||
if c.Providers.OpenRouter.APIBase != "" {
|
if c.Providers.OpenRouter.APIBase != "" {
|
||||||
return c.Providers.OpenRouter.APIBase
|
return c.Providers.OpenRouter.APIBase
|
||||||
|
|
@ -536,32 +563,6 @@ func (c *Config) GetAPIBase() string {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// ModelConfig holds primary model and fallback list.
|
|
||||||
type ModelConfig struct {
|
|
||||||
Primary string
|
|
||||||
Fallbacks []string
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetModelConfig returns the text model configuration with fallbacks.
|
|
||||||
func (c *Config) GetModelConfig() ModelConfig {
|
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
return ModelConfig{
|
|
||||||
Primary: c.Agents.Defaults.Model,
|
|
||||||
Fallbacks: c.Agents.Defaults.ModelFallbacks,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetImageModelConfig returns the image model configuration with fallbacks.
|
|
||||||
func (c *Config) GetImageModelConfig() ModelConfig {
|
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
return ModelConfig{
|
|
||||||
Primary: c.Agents.Defaults.ImageModel,
|
|
||||||
Fallbacks: c.Agents.Defaults.ImageModelFallbacks,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func expandHome(path string) string {
|
func expandHome(path string) string {
|
||||||
if path == "" {
|
if path == "" {
|
||||||
return path
|
return path
|
||||||
|
|
@ -575,3 +576,65 @@ func expandHome(path string) string {
|
||||||
}
|
}
|
||||||
return path
|
return path
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetModelConfig returns the ModelConfig for the given model name.
|
||||||
|
// If multiple configs exist with the same model_name, it uses round-robin
|
||||||
|
// selection for load balancing. Returns an error if the model is not found.
|
||||||
|
func (c *Config) GetModelConfig(modelName string) (*ModelConfig, error) {
|
||||||
|
matches := c.findMatches(modelName)
|
||||||
|
if len(matches) == 0 {
|
||||||
|
return nil, fmt.Errorf("model %q not found in model_list or providers", modelName)
|
||||||
|
}
|
||||||
|
if len(matches) == 1 {
|
||||||
|
return &matches[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Multiple configs - use round-robin for load balancing
|
||||||
|
idx := rrCounter.Add(1) % uint64(len(matches))
|
||||||
|
return &matches[idx], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// findMatches finds all ModelConfig entries with the given model_name.
|
||||||
|
func (c *Config) findMatches(modelName string) []ModelConfig {
|
||||||
|
var matches []ModelConfig
|
||||||
|
for i := range c.ModelList {
|
||||||
|
if c.ModelList[i].ModelName == modelName {
|
||||||
|
matches = append(matches, c.ModelList[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return matches
|
||||||
|
}
|
||||||
|
|
||||||
|
// HasProvidersConfig checks if any provider in the old providers config has configuration.
|
||||||
|
func (c *Config) HasProvidersConfig() bool {
|
||||||
|
v := c.Providers
|
||||||
|
return v.Anthropic.APIKey != "" || v.Anthropic.APIBase != "" ||
|
||||||
|
v.OpenAI.APIKey != "" || v.OpenAI.APIBase != "" ||
|
||||||
|
v.OpenRouter.APIKey != "" || v.OpenRouter.APIBase != "" ||
|
||||||
|
v.Groq.APIKey != "" || v.Groq.APIBase != "" ||
|
||||||
|
v.Zhipu.APIKey != "" || v.Zhipu.APIBase != "" ||
|
||||||
|
v.VLLM.APIKey != "" || v.VLLM.APIBase != "" ||
|
||||||
|
v.Gemini.APIKey != "" || v.Gemini.APIBase != "" ||
|
||||||
|
v.Nvidia.APIKey != "" || v.Nvidia.APIBase != "" ||
|
||||||
|
v.Ollama.APIKey != "" || v.Ollama.APIBase != "" ||
|
||||||
|
v.Moonshot.APIKey != "" || v.Moonshot.APIBase != "" ||
|
||||||
|
v.ShengSuanYun.APIKey != "" || v.ShengSuanYun.APIBase != "" ||
|
||||||
|
v.DeepSeek.APIKey != "" || v.DeepSeek.APIBase != "" ||
|
||||||
|
v.Cerebras.APIKey != "" || v.Cerebras.APIBase != "" ||
|
||||||
|
v.VolcEngine.APIKey != "" || v.VolcEngine.APIBase != "" ||
|
||||||
|
v.GitHubCopilot.APIKey != "" || v.GitHubCopilot.APIBase != "" ||
|
||||||
|
v.Antigravity.APIKey != "" || v.Antigravity.APIBase != "" ||
|
||||||
|
v.Qwen.APIKey != "" || v.Qwen.APIBase != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateModelList validates all ModelConfig entries in the model_list.
|
||||||
|
// It checks that each model config is valid.
|
||||||
|
// Note: Multiple entries with the same model_name are allowed for load balancing.
|
||||||
|
func (c *Config) ValidateModelList() error {
|
||||||
|
for i := range c.ModelList {
|
||||||
|
if err := c.ModelList[i].Validate(); err != nil {
|
||||||
|
return fmt.Errorf("model_list[%d]: %w", i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -237,8 +237,8 @@ func TestDefaultConfig_MaxToolIterations(t *testing.T) {
|
||||||
func TestDefaultConfig_Temperature(t *testing.T) {
|
func TestDefaultConfig_Temperature(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
if cfg.Agents.Defaults.Temperature == 0 {
|
if cfg.Agents.Defaults.Temperature != nil {
|
||||||
t.Error("Temperature should not be zero")
|
t.Error("Temperature should be nil when not provided")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -339,8 +339,8 @@ func TestConfig_Complete(t *testing.T) {
|
||||||
if cfg.Agents.Defaults.Model == "" {
|
if cfg.Agents.Defaults.Model == "" {
|
||||||
t.Error("Model should not be empty")
|
t.Error("Model should not be empty")
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Temperature == 0 {
|
if cfg.Agents.Defaults.Temperature != nil {
|
||||||
t.Error("Temperature should have default value")
|
t.Error("Temperature should be nil when not provided")
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.MaxTokens == 0 {
|
if cfg.Agents.Defaults.MaxTokens == 0 {
|
||||||
t.Error("MaxTokens should not be zero")
|
t.Error("MaxTokens should not be zero")
|
||||||
|
|
|
||||||
292
pkg/config/defaults.go
Normal file
292
pkg/config/defaults.go
Normal file
|
|
@ -0,0 +1,292 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
// DefaultConfig returns the default configuration for PicoClaw.
|
||||||
|
func DefaultConfig() *Config {
|
||||||
|
return &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Workspace: "~/.picoclaw/workspace",
|
||||||
|
RestrictToWorkspace: true,
|
||||||
|
Provider: "",
|
||||||
|
Model: "glm-4.7",
|
||||||
|
MaxTokens: 8192,
|
||||||
|
Temperature: nil, // nil means use provider default
|
||||||
|
MaxToolIterations: 20,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Bindings: []AgentBinding{},
|
||||||
|
Session: SessionConfig{
|
||||||
|
DMScope: "main",
|
||||||
|
},
|
||||||
|
Channels: ChannelsConfig{
|
||||||
|
WhatsApp: WhatsAppConfig{
|
||||||
|
Enabled: false,
|
||||||
|
BridgeURL: "ws://localhost:3001",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Telegram: TelegramConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Token: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Feishu: FeishuConfig{
|
||||||
|
Enabled: false,
|
||||||
|
AppID: "",
|
||||||
|
AppSecret: "",
|
||||||
|
EncryptKey: "",
|
||||||
|
VerificationToken: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Discord: DiscordConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Token: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
MentionOnly: false,
|
||||||
|
},
|
||||||
|
MaixCam: MaixCamConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Host: "0.0.0.0",
|
||||||
|
Port: 18790,
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
QQ: QQConfig{
|
||||||
|
Enabled: false,
|
||||||
|
AppID: "",
|
||||||
|
AppSecret: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
DingTalk: DingTalkConfig{
|
||||||
|
Enabled: false,
|
||||||
|
ClientID: "",
|
||||||
|
ClientSecret: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Slack: SlackConfig{
|
||||||
|
Enabled: false,
|
||||||
|
BotToken: "",
|
||||||
|
AppToken: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
LINE: LINEConfig{
|
||||||
|
Enabled: false,
|
||||||
|
ChannelSecret: "",
|
||||||
|
ChannelAccessToken: "",
|
||||||
|
WebhookHost: "0.0.0.0",
|
||||||
|
WebhookPort: 18791,
|
||||||
|
WebhookPath: "/webhook/line",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
OneBot: OneBotConfig{
|
||||||
|
Enabled: false,
|
||||||
|
WSUrl: "ws://127.0.0.1:3001",
|
||||||
|
AccessToken: "",
|
||||||
|
ReconnectInterval: 5,
|
||||||
|
GroupTriggerPrefix: []string{},
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{WebSearch: true},
|
||||||
|
},
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
// ============================================
|
||||||
|
// Add your API key to the model you want to use
|
||||||
|
// ============================================
|
||||||
|
|
||||||
|
// Zhipu AI (智谱) - https://open.bigmodel.cn/usercenter/apikeys
|
||||||
|
{
|
||||||
|
ModelName: "glm-4.7",
|
||||||
|
Model: "zhipu/glm-4.7",
|
||||||
|
APIBase: "https://open.bigmodel.cn/api/paas/v4",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// OpenAI - https://platform.openai.com/api-keys
|
||||||
|
{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
APIBase: "https://api.openai.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Anthropic Claude - https://console.anthropic.com/settings/keys
|
||||||
|
{
|
||||||
|
ModelName: "claude-sonnet-4.6",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
APIBase: "https://api.anthropic.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// DeepSeek - https://platform.deepseek.com/
|
||||||
|
{
|
||||||
|
ModelName: "deepseek-chat",
|
||||||
|
Model: "deepseek/deepseek-chat",
|
||||||
|
APIBase: "https://api.deepseek.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Google Gemini - https://ai.google.dev/
|
||||||
|
{
|
||||||
|
ModelName: "gemini-2.0-flash",
|
||||||
|
Model: "gemini/gemini-2.0-flash-exp",
|
||||||
|
APIBase: "https://generativelanguage.googleapis.com/v1beta",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Qwen (通义千问) - https://dashscope.console.aliyun.com/apiKey
|
||||||
|
{
|
||||||
|
ModelName: "qwen-plus",
|
||||||
|
Model: "qwen/qwen-plus",
|
||||||
|
APIBase: "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Moonshot (月之暗面) - https://platform.moonshot.cn/console/api-keys
|
||||||
|
{
|
||||||
|
ModelName: "moonshot-v1-8k",
|
||||||
|
Model: "moonshot/moonshot-v1-8k",
|
||||||
|
APIBase: "https://api.moonshot.cn/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Groq - https://console.groq.com/keys
|
||||||
|
{
|
||||||
|
ModelName: "llama-3.3-70b",
|
||||||
|
Model: "groq/llama-3.3-70b-versatile",
|
||||||
|
APIBase: "https://api.groq.com/openai/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// OpenRouter (100+ models) - https://openrouter.ai/keys
|
||||||
|
{
|
||||||
|
ModelName: "openrouter-auto",
|
||||||
|
Model: "openrouter/auto",
|
||||||
|
APIBase: "https://openrouter.ai/api/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "openrouter-gpt-5.2",
|
||||||
|
Model: "openrouter/openai/gpt-5.2",
|
||||||
|
APIBase: "https://openrouter.ai/api/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// NVIDIA - https://build.nvidia.com/
|
||||||
|
{
|
||||||
|
ModelName: "nemotron-4-340b",
|
||||||
|
Model: "nvidia/nemotron-4-340b-instruct",
|
||||||
|
APIBase: "https://integrate.api.nvidia.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Cerebras - https://inference.cerebras.ai/
|
||||||
|
{
|
||||||
|
ModelName: "cerebras-llama-3.3-70b",
|
||||||
|
Model: "cerebras/llama-3.3-70b",
|
||||||
|
APIBase: "https://api.cerebras.ai/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Volcengine (火山引擎) - https://console.volcengine.com/ark
|
||||||
|
{
|
||||||
|
ModelName: "doubao-pro",
|
||||||
|
Model: "volcengine/doubao-pro-32k",
|
||||||
|
APIBase: "https://ark.cn-beijing.volces.com/api/v3",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// ShengsuanYun (神算云)
|
||||||
|
{
|
||||||
|
ModelName: "deepseek-v3",
|
||||||
|
Model: "shengsuanyun/deepseek-v3",
|
||||||
|
APIBase: "https://api.shengsuanyun.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Antigravity (Google Cloud Code Assist) - OAuth only
|
||||||
|
{
|
||||||
|
ModelName: "gemini-flash",
|
||||||
|
Model: "antigravity/gemini-3-flash",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
|
||||||
|
// GitHub Copilot - https://github.com/settings/tokens
|
||||||
|
{
|
||||||
|
ModelName: "copilot-gpt-5.2",
|
||||||
|
Model: "github-copilot/gpt-5.2",
|
||||||
|
APIBase: "http://localhost:4321",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Ollama (local) - https://ollama.com
|
||||||
|
{
|
||||||
|
ModelName: "llama3",
|
||||||
|
Model: "ollama/llama3",
|
||||||
|
APIBase: "http://localhost:11434/v1",
|
||||||
|
APIKey: "ollama",
|
||||||
|
},
|
||||||
|
|
||||||
|
// VLLM (local) - http://localhost:8000
|
||||||
|
{
|
||||||
|
ModelName: "local-model",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "http://localhost:8000/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Gateway: GatewayConfig{
|
||||||
|
Host: "0.0.0.0",
|
||||||
|
Port: 18790,
|
||||||
|
},
|
||||||
|
Tools: ToolsConfig{
|
||||||
|
Web: WebToolsConfig{
|
||||||
|
Brave: BraveConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
DuckDuckGo: DuckDuckGoConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
Perplexity: PerplexityConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Cron: CronToolsConfig{
|
||||||
|
ExecTimeoutMinutes: 5,
|
||||||
|
},
|
||||||
|
Exec: ExecConfig{
|
||||||
|
EnableDenyPatterns: true,
|
||||||
|
},
|
||||||
|
Skills: SkillsToolsConfig{
|
||||||
|
Registries: SkillsRegistriesConfig{
|
||||||
|
ClawHub: ClawHubRegistryConfig{
|
||||||
|
Enabled: true,
|
||||||
|
BaseURL: "https://clawhub.ai",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
MaxConcurrentSearches: 2,
|
||||||
|
SearchCache: SearchCacheConfig{
|
||||||
|
MaxSize: 50,
|
||||||
|
TTLSeconds: 300,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Heartbeat: HeartbeatConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Interval: 30,
|
||||||
|
},
|
||||||
|
Devices: DevicesConfig{
|
||||||
|
Enabled: false,
|
||||||
|
MonitorUSB: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
353
pkg/config/migration.go
Normal file
353
pkg/config/migration.go
Normal file
|
|
@ -0,0 +1,353 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// buildModelWithProtocol constructs a model string with protocol prefix.
|
||||||
|
// If the model already contains a "/" (indicating it has a protocol prefix), it is returned as-is.
|
||||||
|
// Otherwise, the protocol prefix is added.
|
||||||
|
func buildModelWithProtocol(protocol, model string) string {
|
||||||
|
if strings.Contains(model, "/") {
|
||||||
|
// Model already has a protocol prefix, return as-is
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
return protocol + "/" + model
|
||||||
|
}
|
||||||
|
|
||||||
|
// providerMigrationConfig defines how to migrate a provider from old config to new format.
|
||||||
|
type providerMigrationConfig struct {
|
||||||
|
// providerNames are the possible names used in agents.defaults.provider
|
||||||
|
providerNames []string
|
||||||
|
// protocol is the protocol prefix for the model field
|
||||||
|
protocol string
|
||||||
|
// buildConfig creates the ModelConfig from ProviderConfig
|
||||||
|
buildConfig func(p ProvidersConfig) (ModelConfig, bool)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConvertProvidersToModelList converts the old ProvidersConfig to a slice of ModelConfig.
|
||||||
|
// This enables backward compatibility with existing configurations.
|
||||||
|
// It preserves the user's configured model from agents.defaults.model when possible.
|
||||||
|
func ConvertProvidersToModelList(cfg *Config) []ModelConfig {
|
||||||
|
if cfg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get user's configured provider and model
|
||||||
|
userProvider := strings.ToLower(cfg.Agents.Defaults.Provider)
|
||||||
|
userModel := cfg.Agents.Defaults.Model
|
||||||
|
|
||||||
|
p := cfg.Providers
|
||||||
|
|
||||||
|
var result []ModelConfig
|
||||||
|
|
||||||
|
// Track if we've applied the legacy model name fix (only for first provider)
|
||||||
|
legacyModelNameApplied := false
|
||||||
|
|
||||||
|
// Define migration rules for each provider
|
||||||
|
migrations := []providerMigrationConfig{
|
||||||
|
{
|
||||||
|
providerNames: []string{"openai", "gpt"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "openai",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
APIKey: p.OpenAI.APIKey,
|
||||||
|
APIBase: p.OpenAI.APIBase,
|
||||||
|
Proxy: p.OpenAI.Proxy,
|
||||||
|
AuthMethod: p.OpenAI.AuthMethod,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"anthropic", "claude"},
|
||||||
|
protocol: "anthropic",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "anthropic",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
APIKey: p.Anthropic.APIKey,
|
||||||
|
APIBase: p.Anthropic.APIBase,
|
||||||
|
Proxy: p.Anthropic.Proxy,
|
||||||
|
AuthMethod: p.Anthropic.AuthMethod,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"openrouter"},
|
||||||
|
protocol: "openrouter",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "openrouter",
|
||||||
|
Model: "openrouter/auto",
|
||||||
|
APIKey: p.OpenRouter.APIKey,
|
||||||
|
APIBase: p.OpenRouter.APIBase,
|
||||||
|
Proxy: p.OpenRouter.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"groq"},
|
||||||
|
protocol: "groq",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Groq.APIKey == "" && p.Groq.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "groq",
|
||||||
|
Model: "groq/llama-3.1-70b-versatile",
|
||||||
|
APIKey: p.Groq.APIKey,
|
||||||
|
APIBase: p.Groq.APIBase,
|
||||||
|
Proxy: p.Groq.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"zhipu", "glm"},
|
||||||
|
protocol: "zhipu",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "zhipu",
|
||||||
|
Model: "zhipu/glm-4",
|
||||||
|
APIKey: p.Zhipu.APIKey,
|
||||||
|
APIBase: p.Zhipu.APIBase,
|
||||||
|
Proxy: p.Zhipu.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"vllm"},
|
||||||
|
protocol: "vllm",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.VLLM.APIKey == "" && p.VLLM.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "vllm",
|
||||||
|
Model: "vllm/auto",
|
||||||
|
APIKey: p.VLLM.APIKey,
|
||||||
|
APIBase: p.VLLM.APIBase,
|
||||||
|
Proxy: p.VLLM.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"gemini", "google"},
|
||||||
|
protocol: "gemini",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Gemini.APIKey == "" && p.Gemini.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "gemini",
|
||||||
|
Model: "gemini/gemini-pro",
|
||||||
|
APIKey: p.Gemini.APIKey,
|
||||||
|
APIBase: p.Gemini.APIBase,
|
||||||
|
Proxy: p.Gemini.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"nvidia"},
|
||||||
|
protocol: "nvidia",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Nvidia.APIKey == "" && p.Nvidia.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "nvidia",
|
||||||
|
Model: "nvidia/meta/llama-3.1-8b-instruct",
|
||||||
|
APIKey: p.Nvidia.APIKey,
|
||||||
|
APIBase: p.Nvidia.APIBase,
|
||||||
|
Proxy: p.Nvidia.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"ollama"},
|
||||||
|
protocol: "ollama",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Ollama.APIKey == "" && p.Ollama.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "ollama",
|
||||||
|
Model: "ollama/llama3",
|
||||||
|
APIKey: p.Ollama.APIKey,
|
||||||
|
APIBase: p.Ollama.APIBase,
|
||||||
|
Proxy: p.Ollama.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"moonshot", "kimi"},
|
||||||
|
protocol: "moonshot",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Moonshot.APIKey == "" && p.Moonshot.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "moonshot",
|
||||||
|
Model: "moonshot/kimi",
|
||||||
|
APIKey: p.Moonshot.APIKey,
|
||||||
|
APIBase: p.Moonshot.APIBase,
|
||||||
|
Proxy: p.Moonshot.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"shengsuanyun"},
|
||||||
|
protocol: "shengsuanyun",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.ShengSuanYun.APIKey == "" && p.ShengSuanYun.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "shengsuanyun",
|
||||||
|
Model: "shengsuanyun/auto",
|
||||||
|
APIKey: p.ShengSuanYun.APIKey,
|
||||||
|
APIBase: p.ShengSuanYun.APIBase,
|
||||||
|
Proxy: p.ShengSuanYun.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"deepseek"},
|
||||||
|
protocol: "deepseek",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.DeepSeek.APIKey == "" && p.DeepSeek.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "deepseek",
|
||||||
|
Model: "deepseek/deepseek-chat",
|
||||||
|
APIKey: p.DeepSeek.APIKey,
|
||||||
|
APIBase: p.DeepSeek.APIBase,
|
||||||
|
Proxy: p.DeepSeek.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"cerebras"},
|
||||||
|
protocol: "cerebras",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Cerebras.APIKey == "" && p.Cerebras.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "cerebras",
|
||||||
|
Model: "cerebras/llama-3.3-70b",
|
||||||
|
APIKey: p.Cerebras.APIKey,
|
||||||
|
APIBase: p.Cerebras.APIBase,
|
||||||
|
Proxy: p.Cerebras.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"volcengine", "doubao"},
|
||||||
|
protocol: "volcengine",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.VolcEngine.APIKey == "" && p.VolcEngine.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "volcengine",
|
||||||
|
Model: "volcengine/doubao-pro",
|
||||||
|
APIKey: p.VolcEngine.APIKey,
|
||||||
|
APIBase: p.VolcEngine.APIBase,
|
||||||
|
Proxy: p.VolcEngine.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"github_copilot", "copilot"},
|
||||||
|
protocol: "github-copilot",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" && p.GitHubCopilot.ConnectMode == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "github-copilot",
|
||||||
|
Model: "github-copilot/gpt-5.2",
|
||||||
|
APIBase: p.GitHubCopilot.APIBase,
|
||||||
|
ConnectMode: p.GitHubCopilot.ConnectMode,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"antigravity"},
|
||||||
|
protocol: "antigravity",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Antigravity.APIKey == "" && p.Antigravity.AuthMethod == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "antigravity",
|
||||||
|
Model: "antigravity/gemini-2.0-flash",
|
||||||
|
APIKey: p.Antigravity.APIKey,
|
||||||
|
AuthMethod: p.Antigravity.AuthMethod,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"qwen", "tongyi"},
|
||||||
|
protocol: "qwen",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Qwen.APIKey == "" && p.Qwen.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "qwen",
|
||||||
|
Model: "qwen/qwen-max",
|
||||||
|
APIKey: p.Qwen.APIKey,
|
||||||
|
APIBase: p.Qwen.APIBase,
|
||||||
|
Proxy: p.Qwen.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process each provider migration
|
||||||
|
for _, m := range migrations {
|
||||||
|
mc, ok := m.buildConfig(p)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is the user's configured provider
|
||||||
|
if slices.Contains(m.providerNames, userProvider) && userModel != "" {
|
||||||
|
// Use the user's configured model instead of default
|
||||||
|
mc.Model = buildModelWithProtocol(m.protocol, userModel)
|
||||||
|
} else if userProvider == "" && userModel != "" && !legacyModelNameApplied {
|
||||||
|
// Legacy config: no explicit provider field but model is specified
|
||||||
|
// Use userModel as ModelName for the FIRST provider so GetModelConfig(model) can find it
|
||||||
|
// This maintains backward compatibility with old configs that relied on implicit provider selection
|
||||||
|
mc.ModelName = userModel
|
||||||
|
mc.Model = buildModelWithProtocol(m.protocol, userModel)
|
||||||
|
legacyModelNameApplied = true
|
||||||
|
}
|
||||||
|
|
||||||
|
result = append(result, mc)
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
551
pkg/config/migration_test.go
Normal file
551
pkg/config/migration_test.go
Normal file
|
|
@ -0,0 +1,551 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_OpenAI(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{
|
||||||
|
ProviderConfig: ProviderConfig{
|
||||||
|
APIKey: "sk-test-key",
|
||||||
|
APIBase: "https://custom.api.com/v1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].ModelName != "openai" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "openai")
|
||||||
|
}
|
||||||
|
if result[0].Model != "openai/gpt-5.2" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "openai/gpt-5.2")
|
||||||
|
}
|
||||||
|
if result[0].APIKey != "sk-test-key" {
|
||||||
|
t.Errorf("APIKey = %q, want %q", result[0].APIKey, "sk-test-key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Anthropic(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Anthropic: ProviderConfig{
|
||||||
|
APIKey: "ant-key",
|
||||||
|
APIBase: "https://custom.anthropic.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].ModelName != "anthropic" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "anthropic")
|
||||||
|
}
|
||||||
|
if result[0].Model != "anthropic/claude-sonnet-4.6" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "anthropic/claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Multiple(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "openai-key"}},
|
||||||
|
Groq: ProviderConfig{APIKey: "groq-key"},
|
||||||
|
Zhipu: ProviderConfig{APIKey: "zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 3 {
|
||||||
|
t.Fatalf("len(result) = %d, want 3", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check that all providers are present
|
||||||
|
found := make(map[string]bool)
|
||||||
|
for _, mc := range result {
|
||||||
|
found[mc.ModelName] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range []string{"openai", "groq", "zhipu"} {
|
||||||
|
if !found[name] {
|
||||||
|
t.Errorf("Missing provider %q in result", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Empty(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 0 {
|
||||||
|
t.Errorf("len(result) = %d, want 0", len(result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Nil(t *testing.T) {
|
||||||
|
result := ConvertProvidersToModelList(nil)
|
||||||
|
|
||||||
|
if result != nil {
|
||||||
|
t.Errorf("result = %v, want nil", result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_AllProviders(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "key1"}},
|
||||||
|
Anthropic: ProviderConfig{APIKey: "key2"},
|
||||||
|
OpenRouter: ProviderConfig{APIKey: "key3"},
|
||||||
|
Groq: ProviderConfig{APIKey: "key4"},
|
||||||
|
Zhipu: ProviderConfig{APIKey: "key5"},
|
||||||
|
VLLM: ProviderConfig{APIKey: "key6"},
|
||||||
|
Gemini: ProviderConfig{APIKey: "key7"},
|
||||||
|
Nvidia: ProviderConfig{APIKey: "key8"},
|
||||||
|
Ollama: ProviderConfig{APIKey: "key9"},
|
||||||
|
Moonshot: ProviderConfig{APIKey: "key10"},
|
||||||
|
ShengSuanYun: ProviderConfig{APIKey: "key11"},
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "key12"},
|
||||||
|
Cerebras: ProviderConfig{APIKey: "key13"},
|
||||||
|
VolcEngine: ProviderConfig{APIKey: "key14"},
|
||||||
|
GitHubCopilot: ProviderConfig{ConnectMode: "grpc"},
|
||||||
|
Antigravity: ProviderConfig{AuthMethod: "oauth"},
|
||||||
|
Qwen: ProviderConfig{APIKey: "key17"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
// All 17 providers should be converted
|
||||||
|
if len(result) != 17 {
|
||||||
|
t.Errorf("len(result) = %d, want 17", len(result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Proxy(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{
|
||||||
|
ProviderConfig: ProviderConfig{
|
||||||
|
APIKey: "key",
|
||||||
|
Proxy: "http://proxy:8080",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Proxy != "http://proxy:8080" {
|
||||||
|
t.Errorf("Proxy = %q, want %q", result[0].Proxy, "http://proxy:8080")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_AuthMethod(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{
|
||||||
|
ProviderConfig: ProviderConfig{
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 0 {
|
||||||
|
t.Errorf("len(result) = %d, want 0 (AuthMethod alone should not create entry)", len(result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tests for preserving user's configured model during migration
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_DeepSeek(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "deepseek-reasoner",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use user's model, not default
|
||||||
|
if result[0].Model != "deepseek/deepseek-reasoner" {
|
||||||
|
t.Errorf("Model = %q, want %q (user's configured model)", result[0].Model, "deepseek/deepseek-reasoner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_OpenAI(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "openai",
|
||||||
|
Model: "gpt-4-turbo",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "sk-openai"}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "openai/gpt-4-turbo" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "openai/gpt-4-turbo")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_Anthropic(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "claude", // alternative name
|
||||||
|
Model: "claude-opus-4-20250514",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Anthropic: ProviderConfig{APIKey: "sk-ant"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "anthropic/claude-opus-4-20250514" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "anthropic/claude-opus-4-20250514")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_Qwen(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "qwen",
|
||||||
|
Model: "qwen-plus",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Qwen: ProviderConfig{APIKey: "sk-qwen"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "qwen/qwen-plus" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "qwen/qwen-plus")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_UsesDefaultWhenNoUserModel(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "", // no model specified
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use default model
|
||||||
|
if result[0].Model != "deepseek/deepseek-chat" {
|
||||||
|
t.Errorf("Model = %q, want %q (default)", result[0].Model, "deepseek/deepseek-chat")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_MultipleProviders_PreservesUserModel(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "deepseek-reasoner",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "sk-openai"}},
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 2 {
|
||||||
|
t.Fatalf("len(result) = %d, want 2", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find each provider and verify model
|
||||||
|
for _, mc := range result {
|
||||||
|
switch mc.ModelName {
|
||||||
|
case "openai":
|
||||||
|
if mc.Model != "openai/gpt-5.2" {
|
||||||
|
t.Errorf("OpenAI Model = %q, want %q (default)", mc.Model, "openai/gpt-5.2")
|
||||||
|
}
|
||||||
|
case "deepseek":
|
||||||
|
if mc.Model != "deepseek/deepseek-reasoner" {
|
||||||
|
t.Errorf("DeepSeek Model = %q, want %q (user's)", mc.Model, "deepseek/deepseek-reasoner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_ProviderNameAliases(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
providerAlias string
|
||||||
|
expectedModel string
|
||||||
|
provider ProviderConfig
|
||||||
|
}{
|
||||||
|
{"gpt", "openai/gpt-4-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"claude", "anthropic/claude-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"doubao", "volcengine/doubao-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"tongyi", "qwen/qwen-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"kimi", "moonshot/kimi-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.providerAlias, func(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: tt.providerAlias,
|
||||||
|
Model: strings.TrimPrefix(tt.expectedModel, tt.expectedModel[:strings.Index(tt.expectedModel, "/")+1]),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set the appropriate provider config
|
||||||
|
switch tt.providerAlias {
|
||||||
|
case "gpt":
|
||||||
|
cfg.Providers.OpenAI = OpenAIProviderConfig{ProviderConfig: tt.provider}
|
||||||
|
case "claude":
|
||||||
|
cfg.Providers.Anthropic = tt.provider
|
||||||
|
case "doubao":
|
||||||
|
cfg.Providers.VolcEngine = tt.provider
|
||||||
|
case "tongyi":
|
||||||
|
cfg.Providers.Qwen = tt.provider
|
||||||
|
case "kimi":
|
||||||
|
cfg.Providers.Moonshot = tt.provider
|
||||||
|
}
|
||||||
|
|
||||||
|
// Need to fix the model name in config
|
||||||
|
cfg.Agents.Defaults.Model = strings.TrimPrefix(tt.expectedModel, tt.expectedModel[:strings.Index(tt.expectedModel, "/")+1])
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract just the model ID part (after the first /)
|
||||||
|
expectedModelID := tt.expectedModel
|
||||||
|
if result[0].Model != expectedModelID {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, expectedModelID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test for backward compatibility: single provider without explicit provider field
|
||||||
|
// This matches the legacy config pattern where users only set model, not provider
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_NoProviderField_SingleProvider(t *testing.T) {
|
||||||
|
// This matches the user's actual config:
|
||||||
|
// - No provider field set
|
||||||
|
// - model = "glm-4.7"
|
||||||
|
// - Only zhipu has API key configured
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "", // Not set
|
||||||
|
Model: "glm-4.7",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Zhipu: ProviderConfig{APIKey: "test-zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelName should be the user's model value for backward compatibility
|
||||||
|
if result[0].ModelName != "glm-4.7" {
|
||||||
|
t.Errorf("ModelName = %q, want %q (user's model for backward compatibility)", result[0].ModelName, "glm-4.7")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Model should use the user's model with protocol prefix
|
||||||
|
if result[0].Model != "zhipu/glm-4.7" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "zhipu/glm-4.7")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_NoProviderField_MultipleProviders(t *testing.T) {
|
||||||
|
// When multiple providers are configured but no provider field is set,
|
||||||
|
// the FIRST provider (in migration order) will use userModel as ModelName
|
||||||
|
// for backward compatibility with legacy implicit provider selection
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "", // Not set
|
||||||
|
Model: "some-model",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "openai-key"}},
|
||||||
|
Zhipu: ProviderConfig{APIKey: "zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 2 {
|
||||||
|
t.Fatalf("len(result) = %d, want 2", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// The first provider (OpenAI in migration order) should use userModel as ModelName
|
||||||
|
// This ensures GetModelConfig("some-model") will find it
|
||||||
|
if result[0].ModelName != "some-model" {
|
||||||
|
t.Errorf("First provider ModelName = %q, want %q", result[0].ModelName, "some-model")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Other providers should use provider name as ModelName
|
||||||
|
if result[1].ModelName != "zhipu" {
|
||||||
|
t.Errorf("Second provider ModelName = %q, want %q", result[1].ModelName, "zhipu")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_NoProviderField_NoModel(t *testing.T) {
|
||||||
|
// Edge case: no provider, no model
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "",
|
||||||
|
Model: "",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Zhipu: ProviderConfig{APIKey: "zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use default provider name since no model is specified
|
||||||
|
if result[0].ModelName != "zhipu" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "zhipu")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tests for buildModelWithProtocol helper function
|
||||||
|
|
||||||
|
func TestBuildModelWithProtocol_NoPrefix(t *testing.T) {
|
||||||
|
result := buildModelWithProtocol("openai", "gpt-5.2")
|
||||||
|
if result != "openai/gpt-5.2" {
|
||||||
|
t.Errorf("buildModelWithProtocol(openai, gpt-5.2) = %q, want %q", result, "openai/gpt-5.2")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildModelWithProtocol_AlreadyHasPrefix(t *testing.T) {
|
||||||
|
result := buildModelWithProtocol("openrouter", "openrouter/auto")
|
||||||
|
if result != "openrouter/auto" {
|
||||||
|
t.Errorf("buildModelWithProtocol(openrouter, openrouter/auto) = %q, want %q", result, "openrouter/auto")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildModelWithProtocol_DifferentPrefix(t *testing.T) {
|
||||||
|
result := buildModelWithProtocol("anthropic", "openrouter/claude-sonnet-4.6")
|
||||||
|
if result != "openrouter/claude-sonnet-4.6" {
|
||||||
|
t.Errorf("buildModelWithProtocol(anthropic, openrouter/claude-sonnet-4.6) = %q, want %q", result, "openrouter/claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test for legacy config with protocol prefix in model name
|
||||||
|
func TestConvertProvidersToModelList_LegacyModelWithProtocolPrefix(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "", // No explicit provider
|
||||||
|
Model: "openrouter/auto", // Model already has protocol prefix
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenRouter: ProviderConfig{APIKey: "sk-or-test"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) < 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want at least 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// First provider should use userModel as ModelName for backward compatibility
|
||||||
|
if result[0].ModelName != "openrouter/auto" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "openrouter/auto")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Model should NOT have duplicated prefix
|
||||||
|
if result[0].Model != "openrouter/auto" {
|
||||||
|
t.Errorf("Model = %q, want %q (should not duplicate prefix)", result[0].Model, "openrouter/auto")
|
||||||
|
}
|
||||||
|
}
|
||||||
235
pkg/config/model_config_test.go
Normal file
235
pkg/config/model_config_test.go
Normal file
|
|
@ -0,0 +1,235 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGetModelConfig_Found(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test-model", Model: "openai/gpt-4o", APIKey: "key1"},
|
||||||
|
{ModelName: "other-model", Model: "anthropic/claude", APIKey: "key2"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := cfg.GetModelConfig("test-model")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetModelConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if result.Model != "openai/gpt-4o" {
|
||||||
|
t.Errorf("Model = %q, want %q", result.Model, "openai/gpt-4o")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_NotFound(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test-model", Model: "openai/gpt-4o", APIKey: "key1"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := cfg.GetModelConfig("nonexistent")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("GetModelConfig() expected error for nonexistent model")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_EmptyList(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := cfg.GetModelConfig("any-model")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("GetModelConfig() expected error for empty model list")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_RoundRobin(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "lb-model", Model: "openai/gpt-4o-1", APIKey: "key1"},
|
||||||
|
{ModelName: "lb-model", Model: "openai/gpt-4o-2", APIKey: "key2"},
|
||||||
|
{ModelName: "lb-model", Model: "openai/gpt-4o-3", APIKey: "key3"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test round-robin distribution
|
||||||
|
results := make(map[string]int)
|
||||||
|
for i := 0; i < 30; i++ {
|
||||||
|
result, err := cfg.GetModelConfig("lb-model")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetModelConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
results[result.Model]++
|
||||||
|
}
|
||||||
|
|
||||||
|
// Each model should appear roughly 10 times (30 calls / 3 models)
|
||||||
|
for model, count := range results {
|
||||||
|
if count < 5 || count > 15 {
|
||||||
|
t.Errorf("Model %s appeared %d times, expected ~10", model, count)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_Concurrent(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "concurrent-model", Model: "openai/gpt-4o-1", APIKey: "key1"},
|
||||||
|
{ModelName: "concurrent-model", Model: "openai/gpt-4o-2", APIKey: "key2"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
const goroutines = 100
|
||||||
|
const iterations = 10
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
errors := make(chan error, goroutines*iterations)
|
||||||
|
|
||||||
|
for i := 0; i < goroutines; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for j := 0; j < iterations; j++ {
|
||||||
|
_, err := cfg.GetModelConfig("concurrent-model")
|
||||||
|
if err != nil {
|
||||||
|
errors <- err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
close(errors)
|
||||||
|
|
||||||
|
for err := range errors {
|
||||||
|
t.Errorf("Concurrent GetModelConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelConfig_Validate(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
config ModelConfig
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "valid config",
|
||||||
|
config: ModelConfig{
|
||||||
|
ModelName: "test",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
},
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing model_name",
|
||||||
|
config: ModelConfig{
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
},
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing model",
|
||||||
|
config: ModelConfig{
|
||||||
|
ModelName: "test",
|
||||||
|
},
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty config",
|
||||||
|
config: ModelConfig{},
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := tt.config.Validate()
|
||||||
|
if (err != nil) != tt.wantErr {
|
||||||
|
t.Errorf("Validate() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfig_ValidateModelList(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
config *Config
|
||||||
|
wantErr bool
|
||||||
|
errMsg string // partial error message to check
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "valid list",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test1", Model: "openai/gpt-4o"},
|
||||||
|
{ModelName: "test2", Model: "anthropic/claude"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid entry",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test1", Model: "openai/gpt-4o"},
|
||||||
|
{ModelName: "", Model: "anthropic/claude"}, // missing model_name
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: true,
|
||||||
|
errMsg: "model_name is required",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty list",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{},
|
||||||
|
},
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Load balancing: multiple entries with same model_name are allowed
|
||||||
|
name: "duplicate model_name for load balancing",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "gpt-4", Model: "openai/gpt-4o", APIKey: "key1"},
|
||||||
|
{ModelName: "gpt-4", Model: "openai/gpt-4-turbo", APIKey: "key2"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: false, // Changed: duplicates are allowed for load balancing
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Load balancing: non-adjacent entries with same model_name are also allowed
|
||||||
|
name: "duplicate model_name non-adjacent for load balancing",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "model-a", Model: "openai/gpt-4o"},
|
||||||
|
{ModelName: "model-b", Model: "anthropic/claude"},
|
||||||
|
{ModelName: "model-a", Model: "openai/gpt-4-turbo"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: false, // Changed: duplicates are allowed for load balancing
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := tt.config.ValidateModelList()
|
||||||
|
if (err != nil) != tt.wantErr {
|
||||||
|
t.Errorf("ValidateModelList() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
}
|
||||||
|
if err != nil && tt.errMsg != "" {
|
||||||
|
if !strings.Contains(err.Error(), tt.errMsg) {
|
||||||
|
t.Errorf("ValidateModelList() error = %v, want error containing %q", err, tt.errMsg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,15 +1,16 @@
|
||||||
// Package constants provides shared constants across the codebase.
|
// Package constants provides shared constants across the codebase.
|
||||||
package constants
|
package constants
|
||||||
|
|
||||||
// InternalChannels defines channels that are used for internal communication
|
// internalChannels defines channels that are used for internal communication
|
||||||
// and should not be exposed to external users or recorded as last active channel.
|
// and should not be exposed to external users or recorded as last active channel.
|
||||||
var InternalChannels = map[string]bool{
|
var internalChannels = map[string]struct{}{
|
||||||
"cli": true,
|
"cli": {},
|
||||||
"system": true,
|
"system": {},
|
||||||
"subagent": true,
|
"subagent": {},
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsInternalChannel returns true if the channel is an internal channel.
|
// IsInternalChannel returns true if the channel is an internal channel.
|
||||||
func IsInternalChannel(channel string) bool {
|
func IsInternalChannel(channel string) bool {
|
||||||
return InternalChannels[channel]
|
_, found := internalChannels[channel]
|
||||||
|
return found
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,13 +12,16 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
var supportedProviders = map[string]bool{
|
var supportedProviders = map[string]bool{
|
||||||
"anthropic": true,
|
"anthropic": true,
|
||||||
"openai": true,
|
"openai": true,
|
||||||
"openrouter": true,
|
"openrouter": true,
|
||||||
"groq": true,
|
"groq": true,
|
||||||
"zhipu": true,
|
"zhipu": true,
|
||||||
"vllm": true,
|
"vllm": true,
|
||||||
"gemini": true,
|
"gemini": true,
|
||||||
|
"qwen": true,
|
||||||
|
"deepseek": true,
|
||||||
|
"github_copilot": true,
|
||||||
}
|
}
|
||||||
|
|
||||||
var supportedChannels = map[string]bool{
|
var supportedChannels = map[string]bool{
|
||||||
|
|
@ -76,7 +79,7 @@ func ConvertConfig(data map[string]interface{}) (*config.Config, []string, error
|
||||||
cfg.Agents.Defaults.MaxTokens = int(v)
|
cfg.Agents.Defaults.MaxTokens = int(v)
|
||||||
}
|
}
|
||||||
if v, ok := getFloat(defaults, "temperature"); ok {
|
if v, ok := getFloat(defaults, "temperature"); ok {
|
||||||
cfg.Agents.Defaults.Temperature = v
|
cfg.Agents.Defaults.Temperature = &v
|
||||||
}
|
}
|
||||||
if v, ok := getFloat(defaults, "max_tool_iterations"); ok {
|
if v, ok := getFloat(defaults, "max_tool_iterations"); ok {
|
||||||
cfg.Agents.Defaults.MaxToolIterations = int(v)
|
cfg.Agents.Defaults.MaxToolIterations = int(v)
|
||||||
|
|
@ -263,6 +266,15 @@ func MergeConfig(existing, incoming *config.Config) *config.Config {
|
||||||
if existing.Providers.Gemini.APIKey == "" {
|
if existing.Providers.Gemini.APIKey == "" {
|
||||||
existing.Providers.Gemini = incoming.Providers.Gemini
|
existing.Providers.Gemini = incoming.Providers.Gemini
|
||||||
}
|
}
|
||||||
|
if existing.Providers.DeepSeek.APIKey == "" {
|
||||||
|
existing.Providers.DeepSeek = incoming.Providers.DeepSeek
|
||||||
|
}
|
||||||
|
if existing.Providers.GitHubCopilot.APIBase == "" {
|
||||||
|
existing.Providers.GitHubCopilot = incoming.Providers.GitHubCopilot
|
||||||
|
}
|
||||||
|
if existing.Providers.Qwen.APIKey == "" {
|
||||||
|
existing.Providers.Qwen = incoming.Providers.Qwen
|
||||||
|
}
|
||||||
|
|
||||||
if !existing.Channels.Telegram.Enabled && incoming.Channels.Telegram.Enabled {
|
if !existing.Channels.Telegram.Enabled && incoming.Channels.Telegram.Enabled {
|
||||||
existing.Channels.Telegram = incoming.Channels.Telegram
|
existing.Channels.Telegram = incoming.Channels.Telegram
|
||||||
|
|
|
||||||
|
|
@ -180,8 +180,8 @@ func TestConvertConfig(t *testing.T) {
|
||||||
t.Run("unsupported provider warning", func(t *testing.T) {
|
t.Run("unsupported provider warning", func(t *testing.T) {
|
||||||
data := map[string]interface{}{
|
data := map[string]interface{}{
|
||||||
"providers": map[string]interface{}{
|
"providers": map[string]interface{}{
|
||||||
"deepseek": map[string]interface{}{
|
"unknown_provider": map[string]interface{}{
|
||||||
"api_key": "sk-deep-test",
|
"api_key": "sk-test",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
@ -193,7 +193,7 @@ func TestConvertConfig(t *testing.T) {
|
||||||
if len(warnings) != 1 {
|
if len(warnings) != 1 {
|
||||||
t.Fatalf("expected 1 warning, got %d", len(warnings))
|
t.Fatalf("expected 1 warning, got %d", len(warnings))
|
||||||
}
|
}
|
||||||
if warnings[0] != "Provider 'deepseek' not supported in PicoClaw, skipping" {
|
if warnings[0] != "Provider 'unknown_provider' not supported in PicoClaw, skipping" {
|
||||||
t.Errorf("unexpected warning: %s", warnings[0])
|
t.Errorf("unexpected warning: %s", warnings[0])
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
@ -275,8 +275,11 @@ func TestConvertConfig(t *testing.T) {
|
||||||
if cfg.Agents.Defaults.MaxTokens != 4096 {
|
if cfg.Agents.Defaults.MaxTokens != 4096 {
|
||||||
t.Errorf("MaxTokens = %d, want %d", cfg.Agents.Defaults.MaxTokens, 4096)
|
t.Errorf("MaxTokens = %d, want %d", cfg.Agents.Defaults.MaxTokens, 4096)
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Temperature != 0.5 {
|
if cfg.Agents.Defaults.Temperature == nil {
|
||||||
t.Errorf("Temperature = %f, want %f", cfg.Agents.Defaults.Temperature, 0.5)
|
t.Fatalf("Temperature is nil, want %f", 0.5)
|
||||||
|
}
|
||||||
|
if *cfg.Agents.Defaults.Temperature != 0.5 {
|
||||||
|
t.Errorf("Temperature = %f, want %f", *cfg.Agents.Defaults.Temperature, 0.5)
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Workspace != "~/.picoclaw/workspace" {
|
if cfg.Agents.Defaults.Workspace != "~/.picoclaw/workspace" {
|
||||||
t.Errorf("Workspace = %q, want %q", cfg.Agents.Defaults.Workspace, "~/.picoclaw/workspace")
|
t.Errorf("Workspace = %q, want %q", cfg.Agents.Defaults.Workspace, "~/.picoclaw/workspace")
|
||||||
|
|
|
||||||
|
|
@ -85,7 +85,7 @@ func (p *Provider) Chat(ctx context.Context, messages []Message, tools []ToolDef
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Provider) GetDefaultModel() string {
|
func (p *Provider) GetDefaultModel() string {
|
||||||
return "claude-sonnet-4-5-20250929"
|
return "claude-sonnet-4.6"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Provider) BaseURL() string {
|
func (p *Provider) BaseURL() string {
|
||||||
|
|
|
||||||
|
|
@ -15,14 +15,14 @@ func TestBuildParams_BasicMessage(t *testing.T) {
|
||||||
messages := []Message{
|
messages := []Message{
|
||||||
{Role: "user", Content: "Hello"},
|
{Role: "user", Content: "Hello"},
|
||||||
}
|
}
|
||||||
params, err := buildParams(messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{
|
params, err := buildParams(messages, nil, "claude-sonnet-4.6", map[string]interface{}{
|
||||||
"max_tokens": 1024,
|
"max_tokens": 1024,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("buildParams() error: %v", err)
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
}
|
}
|
||||||
if string(params.Model) != "claude-sonnet-4-5-20250929" {
|
if string(params.Model) != "claude-sonnet-4.6" {
|
||||||
t.Errorf("Model = %q, want %q", params.Model, "claude-sonnet-4-5-20250929")
|
t.Errorf("Model = %q, want %q", params.Model, "claude-sonnet-4.6")
|
||||||
}
|
}
|
||||||
if params.MaxTokens != 1024 {
|
if params.MaxTokens != 1024 {
|
||||||
t.Errorf("MaxTokens = %d, want 1024", params.MaxTokens)
|
t.Errorf("MaxTokens = %d, want 1024", params.MaxTokens)
|
||||||
|
|
@ -37,7 +37,7 @@ func TestBuildParams_SystemMessage(t *testing.T) {
|
||||||
{Role: "system", Content: "You are helpful"},
|
{Role: "system", Content: "You are helpful"},
|
||||||
{Role: "user", Content: "Hi"},
|
{Role: "user", Content: "Hi"},
|
||||||
}
|
}
|
||||||
params, err := buildParams(messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
params, err := buildParams(messages, nil, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("buildParams() error: %v", err)
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -68,7 +68,7 @@ func TestBuildParams_ToolCallMessage(t *testing.T) {
|
||||||
},
|
},
|
||||||
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
|
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
|
||||||
}
|
}
|
||||||
params, err := buildParams(messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
params, err := buildParams(messages, nil, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("buildParams() error: %v", err)
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -94,7 +94,7 @@ func TestBuildParams_WithTools(t *testing.T) {
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
params, err := buildParams([]Message{{Role: "user", Content: "Hi"}}, tools, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
params, err := buildParams([]Message{{Role: "user", Content: "Hi"}}, tools, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("buildParams() error: %v", err)
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -178,7 +178,7 @@ func TestProvider_ChatRoundTrip(t *testing.T) {
|
||||||
|
|
||||||
provider := NewProviderWithClient(createAnthropicTestClient(server.URL, "test-token"))
|
provider := NewProviderWithClient(createAnthropicTestClient(server.URL, "test-token"))
|
||||||
messages := []Message{{Role: "user", Content: "Hello"}}
|
messages := []Message{{Role: "user", Content: "Hello"}}
|
||||||
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{"max_tokens": 1024})
|
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4.6", map[string]interface{}{"max_tokens": 1024})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Chat() error: %v", err)
|
t.Fatalf("Chat() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -195,8 +195,8 @@ func TestProvider_ChatRoundTrip(t *testing.T) {
|
||||||
|
|
||||||
func TestProvider_GetDefaultModel(t *testing.T) {
|
func TestProvider_GetDefaultModel(t *testing.T) {
|
||||||
p := NewProvider("test-token")
|
p := NewProvider("test-token")
|
||||||
if got := p.GetDefaultModel(); got != "claude-sonnet-4-5-20250929" {
|
if got := p.GetDefaultModel(); got != "claude-sonnet-4.6" {
|
||||||
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4-5-20250929")
|
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4.6")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -247,7 +247,7 @@ func TestProvider_ChatUsesTokenSource(t *testing.T) {
|
||||||
return "refreshed-token", nil
|
return "refreshed-token", nil
|
||||||
}, server.URL)
|
}, server.URL)
|
||||||
|
|
||||||
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hello"}}, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hello"}}, nil, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Chat() error: %v", err)
|
t.Fatalf("Chat() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
827
pkg/providers/antigravity_provider.go
Normal file
827
pkg/providers/antigravity_provider.go
Normal file
|
|
@ -0,0 +1,827 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"math/rand"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
antigravityBaseURL = "https://cloudcode-pa.googleapis.com"
|
||||||
|
antigravityDefaultModel = "gemini-3-flash"
|
||||||
|
antigravityUserAgent = "antigravity"
|
||||||
|
antigravityXGoogClient = "google-cloud-sdk vscode_cloudshelleditor/0.1"
|
||||||
|
antigravityVersion = "1.15.8"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AntigravityProvider implements LLMProvider using Google's Cloud Code Assist (Antigravity) API.
|
||||||
|
// This provider authenticates via Google OAuth and provides access to models like Claude and Gemini
|
||||||
|
// through Google's infrastructure.
|
||||||
|
type AntigravityProvider struct {
|
||||||
|
tokenSource func() (string, string, error) // Returns (accessToken, projectID, error)
|
||||||
|
httpClient *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAntigravityProvider creates a new Antigravity provider using stored auth credentials.
|
||||||
|
func NewAntigravityProvider() *AntigravityProvider {
|
||||||
|
return &AntigravityProvider{
|
||||||
|
tokenSource: createAntigravityTokenSource(),
|
||||||
|
httpClient: &http.Client{
|
||||||
|
Timeout: 120 * time.Second,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Chat implements LLMProvider.Chat using the Cloud Code Assist v1internal API.
|
||||||
|
// The v1internal endpoint wraps the standard Gemini request in an envelope with
|
||||||
|
// project, model, request, requestType, userAgent, and requestId fields.
|
||||||
|
func (p *AntigravityProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
|
accessToken, projectID, err := p.tokenSource()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("antigravity auth: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if model == "" || model == "antigravity" || model == "google-antigravity" {
|
||||||
|
model = antigravityDefaultModel
|
||||||
|
}
|
||||||
|
// Strip provider prefixes if present
|
||||||
|
model = strings.TrimPrefix(model, "google-antigravity/")
|
||||||
|
model = strings.TrimPrefix(model, "antigravity/")
|
||||||
|
|
||||||
|
logger.DebugCF("provider.antigravity", "Starting chat", map[string]interface{}{
|
||||||
|
"model": model,
|
||||||
|
"project": projectID,
|
||||||
|
"requestId": fmt.Sprintf("agent-%d-%s", time.Now().UnixMilli(), randomString(9)),
|
||||||
|
})
|
||||||
|
|
||||||
|
// Build the inner Gemini-format request
|
||||||
|
innerRequest := p.buildRequest(messages, tools, model, options)
|
||||||
|
|
||||||
|
// Wrap in v1internal envelope (matches pi-ai SDK format)
|
||||||
|
envelope := map[string]interface{}{
|
||||||
|
"project": projectID,
|
||||||
|
"model": model,
|
||||||
|
"request": innerRequest,
|
||||||
|
"requestType": "agent",
|
||||||
|
"userAgent": antigravityUserAgent,
|
||||||
|
"requestId": fmt.Sprintf("agent-%d-%s", time.Now().UnixMilli(), randomString(9)),
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyBytes, err := json.Marshal(envelope)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("marshaling request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build API URL — uses Cloud Code Assist v1internal streaming endpoint
|
||||||
|
apiURL := fmt.Sprintf("%s/v1internal:streamGenerateContent?alt=sse", antigravityBaseURL)
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", apiURL, bytes.NewReader(bodyBytes))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Headers matching the pi-ai SDK antigravity format
|
||||||
|
clientMetadata, _ := json.Marshal(map[string]string{
|
||||||
|
"ideType": "IDE_UNSPECIFIED",
|
||||||
|
"platform": "PLATFORM_UNSPECIFIED",
|
||||||
|
"pluginType": "GEMINI",
|
||||||
|
})
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Accept", "text/event-stream")
|
||||||
|
req.Header.Set("User-Agent", fmt.Sprintf("antigravity/%s linux/amd64", antigravityVersion))
|
||||||
|
req.Header.Set("X-Goog-Api-Client", antigravityXGoogClient)
|
||||||
|
req.Header.Set("Client-Metadata", string(clientMetadata))
|
||||||
|
|
||||||
|
resp, err := p.httpClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("antigravity API call: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
respBody, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("reading response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
logger.ErrorCF("provider.antigravity", "API call failed", map[string]interface{}{
|
||||||
|
"status_code": resp.StatusCode,
|
||||||
|
"response": string(respBody),
|
||||||
|
"model": model,
|
||||||
|
})
|
||||||
|
|
||||||
|
return nil, p.parseAntigravityError(resp.StatusCode, respBody)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Response is always SSE from streamGenerateContent — each line is "data: {...}"
|
||||||
|
// with a "response" wrapper containing the standard Gemini response
|
||||||
|
llmResp, err := p.parseSSEResponse(string(respBody))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for empty response (some models might return valid success but empty text)
|
||||||
|
if llmResp.Content == "" && len(llmResp.ToolCalls) == 0 {
|
||||||
|
return nil, fmt.Errorf("antigravity: model returned an empty response (this model might be invalid or restricted)")
|
||||||
|
}
|
||||||
|
|
||||||
|
return llmResp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDefaultModel returns the default model identifier.
|
||||||
|
func (p *AntigravityProvider) GetDefaultModel() string {
|
||||||
|
return antigravityDefaultModel
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Request building ---
|
||||||
|
|
||||||
|
type antigravityRequest struct {
|
||||||
|
Contents []antigravityContent `json:"contents"`
|
||||||
|
Tools []antigravityTool `json:"tools,omitempty"`
|
||||||
|
SystemPrompt *antigravitySystemPrompt `json:"systemInstruction,omitempty"`
|
||||||
|
Config *antigravityGenConfig `json:"generationConfig,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityContent struct {
|
||||||
|
Role string `json:"role"`
|
||||||
|
Parts []antigravityPart `json:"parts"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityPart struct {
|
||||||
|
Text string `json:"text,omitempty"`
|
||||||
|
ThoughtSignature string `json:"thoughtSignature,omitempty"`
|
||||||
|
ThoughtSignatureSnake string `json:"thought_signature,omitempty"`
|
||||||
|
FunctionCall *antigravityFunctionCall `json:"functionCall,omitempty"`
|
||||||
|
FunctionResponse *antigravityFunctionResponse `json:"functionResponse,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityFunctionCall struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Args map[string]interface{} `json:"args"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityFunctionResponse struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Response map[string]interface{} `json:"response"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityTool struct {
|
||||||
|
FunctionDeclarations []antigravityFuncDecl `json:"functionDeclarations"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityFuncDecl struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description,omitempty"`
|
||||||
|
Parameters interface{} `json:"parameters,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravitySystemPrompt struct {
|
||||||
|
Parts []antigravityPart `json:"parts"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityGenConfig struct {
|
||||||
|
MaxOutputTokens int `json:"maxOutputTokens,omitempty"`
|
||||||
|
Temperature float64 `json:"temperature,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) buildRequest(messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) antigravityRequest {
|
||||||
|
req := antigravityRequest{}
|
||||||
|
toolCallNames := make(map[string]string)
|
||||||
|
|
||||||
|
// Build contents from messages
|
||||||
|
for _, msg := range messages {
|
||||||
|
switch msg.Role {
|
||||||
|
case "system":
|
||||||
|
req.SystemPrompt = &antigravitySystemPrompt{
|
||||||
|
Parts: []antigravityPart{{Text: msg.Content}},
|
||||||
|
}
|
||||||
|
case "user":
|
||||||
|
if msg.ToolCallID != "" {
|
||||||
|
toolName := resolveToolResponseName(msg.ToolCallID, toolCallNames)
|
||||||
|
// Tool result
|
||||||
|
req.Contents = append(req.Contents, antigravityContent{
|
||||||
|
Role: "user",
|
||||||
|
Parts: []antigravityPart{{
|
||||||
|
FunctionResponse: &antigravityFunctionResponse{
|
||||||
|
Name: toolName,
|
||||||
|
Response: map[string]interface{}{
|
||||||
|
"result": msg.Content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
req.Contents = append(req.Contents, antigravityContent{
|
||||||
|
Role: "user",
|
||||||
|
Parts: []antigravityPart{{Text: msg.Content}},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
case "assistant":
|
||||||
|
content := antigravityContent{
|
||||||
|
Role: "model",
|
||||||
|
}
|
||||||
|
if msg.Content != "" {
|
||||||
|
content.Parts = append(content.Parts, antigravityPart{Text: msg.Content})
|
||||||
|
}
|
||||||
|
for _, tc := range msg.ToolCalls {
|
||||||
|
toolName, toolArgs, thoughtSignature := normalizeStoredToolCall(tc)
|
||||||
|
if toolName == "" {
|
||||||
|
logger.WarnCF("provider.antigravity", "Skipping tool call with empty name in history", map[string]interface{}{
|
||||||
|
"tool_call_id": tc.ID,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if tc.ID != "" {
|
||||||
|
toolCallNames[tc.ID] = toolName
|
||||||
|
}
|
||||||
|
content.Parts = append(content.Parts, antigravityPart{
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
|
ThoughtSignatureSnake: thoughtSignature,
|
||||||
|
FunctionCall: &antigravityFunctionCall{
|
||||||
|
Name: toolName,
|
||||||
|
Args: toolArgs,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if len(content.Parts) > 0 {
|
||||||
|
req.Contents = append(req.Contents, content)
|
||||||
|
}
|
||||||
|
case "tool":
|
||||||
|
toolName := resolveToolResponseName(msg.ToolCallID, toolCallNames)
|
||||||
|
req.Contents = append(req.Contents, antigravityContent{
|
||||||
|
Role: "user",
|
||||||
|
Parts: []antigravityPart{{
|
||||||
|
FunctionResponse: &antigravityFunctionResponse{
|
||||||
|
Name: toolName,
|
||||||
|
Response: map[string]interface{}{
|
||||||
|
"result": msg.Content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build tools (sanitize schemas for Gemini compatibility)
|
||||||
|
if len(tools) > 0 {
|
||||||
|
var funcDecls []antigravityFuncDecl
|
||||||
|
for _, t := range tools {
|
||||||
|
if t.Type != "function" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
params := sanitizeSchemaForGemini(t.Function.Parameters)
|
||||||
|
funcDecls = append(funcDecls, antigravityFuncDecl{
|
||||||
|
Name: t.Function.Name,
|
||||||
|
Description: t.Function.Description,
|
||||||
|
Parameters: params,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if len(funcDecls) > 0 {
|
||||||
|
req.Tools = []antigravityTool{{FunctionDeclarations: funcDecls}}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generation config
|
||||||
|
config := &antigravityGenConfig{}
|
||||||
|
if val, ok := options["max_tokens"]; ok {
|
||||||
|
if maxTokens, ok := val.(int); ok && maxTokens > 0 {
|
||||||
|
config.MaxOutputTokens = maxTokens
|
||||||
|
} else if maxTokens, ok := val.(float64); ok && maxTokens > 0 {
|
||||||
|
config.MaxOutputTokens = int(maxTokens)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if temp, ok := options["temperature"].(float64); ok {
|
||||||
|
config.Temperature = temp
|
||||||
|
}
|
||||||
|
if config.MaxOutputTokens > 0 || config.Temperature > 0 {
|
||||||
|
req.Config = config
|
||||||
|
}
|
||||||
|
|
||||||
|
return req
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeStoredToolCall(tc ToolCall) (string, map[string]interface{}, string) {
|
||||||
|
name := tc.Name
|
||||||
|
args := tc.Arguments
|
||||||
|
thoughtSignature := ""
|
||||||
|
|
||||||
|
if name == "" && tc.Function != nil {
|
||||||
|
name = tc.Function.Name
|
||||||
|
thoughtSignature = tc.Function.ThoughtSignature
|
||||||
|
} else if tc.Function != nil {
|
||||||
|
thoughtSignature = tc.Function.ThoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
|
if args == nil {
|
||||||
|
args = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(args) == 0 && tc.Function != nil && tc.Function.Arguments != "" {
|
||||||
|
var parsed map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(tc.Function.Arguments), &parsed); err == nil && parsed != nil {
|
||||||
|
args = parsed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return name, args, thoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveToolResponseName(toolCallID string, toolCallNames map[string]string) string {
|
||||||
|
if toolCallID == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
if name, ok := toolCallNames[toolCallID]; ok && name != "" {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
|
||||||
|
return inferToolNameFromCallID(toolCallID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func inferToolNameFromCallID(toolCallID string) string {
|
||||||
|
if !strings.HasPrefix(toolCallID, "call_") {
|
||||||
|
return toolCallID
|
||||||
|
}
|
||||||
|
|
||||||
|
rest := strings.TrimPrefix(toolCallID, "call_")
|
||||||
|
if idx := strings.LastIndex(rest, "_"); idx > 0 {
|
||||||
|
candidate := rest[:idx]
|
||||||
|
if candidate != "" {
|
||||||
|
return candidate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return toolCallID
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Response parsing ---
|
||||||
|
|
||||||
|
type antigravityJSONResponse struct {
|
||||||
|
Candidates []struct {
|
||||||
|
Content struct {
|
||||||
|
Parts []struct {
|
||||||
|
Text string `json:"text,omitempty"`
|
||||||
|
ThoughtSignature string `json:"thoughtSignature,omitempty"`
|
||||||
|
ThoughtSignatureSnake string `json:"thought_signature,omitempty"`
|
||||||
|
FunctionCall *antigravityFunctionCall `json:"functionCall,omitempty"`
|
||||||
|
} `json:"parts"`
|
||||||
|
Role string `json:"role"`
|
||||||
|
} `json:"content"`
|
||||||
|
FinishReason string `json:"finishReason"`
|
||||||
|
} `json:"candidates"`
|
||||||
|
UsageMetadata struct {
|
||||||
|
PromptTokenCount int `json:"promptTokenCount"`
|
||||||
|
CandidatesTokenCount int `json:"candidatesTokenCount"`
|
||||||
|
TotalTokenCount int `json:"totalTokenCount"`
|
||||||
|
} `json:"usageMetadata"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) parseJSONResponse(body []byte) (*LLMResponse, error) {
|
||||||
|
var resp antigravityJSONResponse
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("parsing antigravity response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(resp.Candidates) == 0 {
|
||||||
|
return nil, fmt.Errorf("antigravity: no candidates in response")
|
||||||
|
}
|
||||||
|
|
||||||
|
candidate := resp.Candidates[0]
|
||||||
|
var contentParts []string
|
||||||
|
var toolCalls []ToolCall
|
||||||
|
|
||||||
|
for _, part := range candidate.Content.Parts {
|
||||||
|
if part.Text != "" {
|
||||||
|
contentParts = append(contentParts, part.Text)
|
||||||
|
}
|
||||||
|
if part.FunctionCall != nil {
|
||||||
|
argumentsJSON, _ := json.Marshal(part.FunctionCall.Args)
|
||||||
|
toolCalls = append(toolCalls, ToolCall{
|
||||||
|
ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()),
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: part.FunctionCall.Args,
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: string(argumentsJSON),
|
||||||
|
ThoughtSignature: extractPartThoughtSignature(part.ThoughtSignature, part.ThoughtSignatureSnake),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
finishReason := "stop"
|
||||||
|
if len(toolCalls) > 0 {
|
||||||
|
finishReason = "tool_calls"
|
||||||
|
}
|
||||||
|
if candidate.FinishReason == "MAX_TOKENS" {
|
||||||
|
finishReason = "length"
|
||||||
|
}
|
||||||
|
|
||||||
|
var usage *UsageInfo
|
||||||
|
if resp.UsageMetadata.TotalTokenCount > 0 {
|
||||||
|
usage = &UsageInfo{
|
||||||
|
PromptTokens: resp.UsageMetadata.PromptTokenCount,
|
||||||
|
CompletionTokens: resp.UsageMetadata.CandidatesTokenCount,
|
||||||
|
TotalTokens: resp.UsageMetadata.TotalTokenCount,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: strings.Join(contentParts, ""),
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: finishReason,
|
||||||
|
Usage: usage,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) parseSSEResponse(body string) (*LLMResponse, error) {
|
||||||
|
var contentParts []string
|
||||||
|
var toolCalls []ToolCall
|
||||||
|
var usage *UsageInfo
|
||||||
|
var finishReason string
|
||||||
|
|
||||||
|
scanner := bufio.NewScanner(strings.NewReader(body))
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Text()
|
||||||
|
if !strings.HasPrefix(line, "data: ") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
data := strings.TrimPrefix(line, "data: ")
|
||||||
|
if data == "[DONE]" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
// v1internal SSE wraps the Gemini response in a "response" field
|
||||||
|
var sseChunk struct {
|
||||||
|
Response antigravityJSONResponse `json:"response"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(data), &sseChunk); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
resp := sseChunk.Response
|
||||||
|
|
||||||
|
for _, candidate := range resp.Candidates {
|
||||||
|
for _, part := range candidate.Content.Parts {
|
||||||
|
if part.Text != "" {
|
||||||
|
contentParts = append(contentParts, part.Text)
|
||||||
|
}
|
||||||
|
if part.FunctionCall != nil {
|
||||||
|
argumentsJSON, _ := json.Marshal(part.FunctionCall.Args)
|
||||||
|
toolCalls = append(toolCalls, ToolCall{
|
||||||
|
ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()),
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: part.FunctionCall.Args,
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: string(argumentsJSON),
|
||||||
|
ThoughtSignature: extractPartThoughtSignature(part.ThoughtSignature, part.ThoughtSignatureSnake),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if candidate.FinishReason != "" {
|
||||||
|
finishReason = candidate.FinishReason
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.UsageMetadata.TotalTokenCount > 0 {
|
||||||
|
usage = &UsageInfo{
|
||||||
|
PromptTokens: resp.UsageMetadata.PromptTokenCount,
|
||||||
|
CompletionTokens: resp.UsageMetadata.CandidatesTokenCount,
|
||||||
|
TotalTokens: resp.UsageMetadata.TotalTokenCount,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mappedFinish := "stop"
|
||||||
|
if len(toolCalls) > 0 {
|
||||||
|
mappedFinish = "tool_calls"
|
||||||
|
}
|
||||||
|
if finishReason == "MAX_TOKENS" {
|
||||||
|
mappedFinish = "length"
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: strings.Join(contentParts, ""),
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: mappedFinish,
|
||||||
|
Usage: usage,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractPartThoughtSignature(thoughtSignature string, thoughtSignatureSnake string) string {
|
||||||
|
if thoughtSignature != "" {
|
||||||
|
return thoughtSignature
|
||||||
|
}
|
||||||
|
if thoughtSignatureSnake != "" {
|
||||||
|
return thoughtSignatureSnake
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Schema sanitization ---
|
||||||
|
|
||||||
|
// Google/Gemini doesn't support many JSON Schema keywords that other providers accept.
|
||||||
|
var geminiUnsupportedKeywords = map[string]bool{
|
||||||
|
"patternProperties": true,
|
||||||
|
"additionalProperties": true,
|
||||||
|
"$schema": true,
|
||||||
|
"$id": true,
|
||||||
|
"$ref": true,
|
||||||
|
"$defs": true,
|
||||||
|
"definitions": true,
|
||||||
|
"examples": true,
|
||||||
|
"minLength": true,
|
||||||
|
"maxLength": true,
|
||||||
|
"minimum": true,
|
||||||
|
"maximum": true,
|
||||||
|
"multipleOf": true,
|
||||||
|
"pattern": true,
|
||||||
|
"format": true,
|
||||||
|
"minItems": true,
|
||||||
|
"maxItems": true,
|
||||||
|
"uniqueItems": true,
|
||||||
|
"minProperties": true,
|
||||||
|
"maxProperties": true,
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeSchemaForGemini(schema map[string]interface{}) map[string]interface{} {
|
||||||
|
if schema == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
result := make(map[string]interface{})
|
||||||
|
for k, v := range schema {
|
||||||
|
if geminiUnsupportedKeywords[k] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Recursively sanitize nested objects
|
||||||
|
switch val := v.(type) {
|
||||||
|
case map[string]interface{}:
|
||||||
|
result[k] = sanitizeSchemaForGemini(val)
|
||||||
|
case []interface{}:
|
||||||
|
sanitized := make([]interface{}, len(val))
|
||||||
|
for i, item := range val {
|
||||||
|
if m, ok := item.(map[string]interface{}); ok {
|
||||||
|
sanitized[i] = sanitizeSchemaForGemini(m)
|
||||||
|
} else {
|
||||||
|
sanitized[i] = item
|
||||||
|
}
|
||||||
|
}
|
||||||
|
result[k] = sanitized
|
||||||
|
default:
|
||||||
|
result[k] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure top-level has type: "object" if properties are present
|
||||||
|
if _, hasProps := result["properties"]; hasProps {
|
||||||
|
if _, hasType := result["type"]; !hasType {
|
||||||
|
result["type"] = "object"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Token source ---
|
||||||
|
|
||||||
|
func createAntigravityTokenSource() func() (string, string, error) {
|
||||||
|
return func() (string, string, error) {
|
||||||
|
cred, err := auth.GetCredential("google-antigravity")
|
||||||
|
if err != nil {
|
||||||
|
return "", "", fmt.Errorf("loading auth credentials: %w", err)
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
return "", "", fmt.Errorf("no credentials for google-antigravity. Run: picoclaw auth login --provider google-antigravity")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh if needed
|
||||||
|
if cred.NeedsRefresh() && cred.RefreshToken != "" {
|
||||||
|
oauthCfg := auth.GoogleAntigravityOAuthConfig()
|
||||||
|
refreshed, err := auth.RefreshAccessToken(cred, oauthCfg)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", fmt.Errorf("refreshing token: %w", err)
|
||||||
|
}
|
||||||
|
refreshed.Email = cred.Email
|
||||||
|
if refreshed.ProjectID == "" {
|
||||||
|
refreshed.ProjectID = cred.ProjectID
|
||||||
|
}
|
||||||
|
if err := auth.SetCredential("google-antigravity", refreshed); err != nil {
|
||||||
|
return "", "", fmt.Errorf("saving refreshed token: %w", err)
|
||||||
|
}
|
||||||
|
cred = refreshed
|
||||||
|
}
|
||||||
|
|
||||||
|
if cred.IsExpired() {
|
||||||
|
return "", "", fmt.Errorf("antigravity credentials expired. Run: picoclaw auth login --provider google-antigravity")
|
||||||
|
}
|
||||||
|
|
||||||
|
projectID := cred.ProjectID
|
||||||
|
if projectID == "" {
|
||||||
|
// Try to fetch project ID from API
|
||||||
|
fetchedID, err := FetchAntigravityProjectID(cred.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("provider.antigravity", "Could not fetch project ID, using fallback", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
projectID = "rising-fact-p41fc" // Default fallback (same as OpenCode)
|
||||||
|
} else {
|
||||||
|
projectID = fetchedID
|
||||||
|
cred.ProjectID = projectID
|
||||||
|
_ = auth.SetCredential("google-antigravity", cred)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return cred.AccessToken, projectID, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// FetchAntigravityProjectID retrieves the Google Cloud project ID from the loadCodeAssist endpoint.
|
||||||
|
func FetchAntigravityProjectID(accessToken string) (string, error) {
|
||||||
|
reqBody, _ := json.Marshal(map[string]interface{}{
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"ideType": "IDE_UNSPECIFIED",
|
||||||
|
"platform": "PLATFORM_UNSPECIFIED",
|
||||||
|
"pluginType": "GEMINI",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
req, err := http.NewRequest("POST", antigravityBaseURL+"/v1internal:loadCodeAssist", bytes.NewReader(reqBody))
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("User-Agent", antigravityUserAgent)
|
||||||
|
req.Header.Set("X-Goog-Api-Client", antigravityXGoogClient)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 15 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("loadCodeAssist failed: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
var result struct {
|
||||||
|
CloudAICompanionProject string `json:"cloudaicompanionProject"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &result); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.CloudAICompanionProject == "" {
|
||||||
|
return "", fmt.Errorf("no project ID in loadCodeAssist response")
|
||||||
|
}
|
||||||
|
|
||||||
|
return result.CloudAICompanionProject, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FetchAntigravityModels fetches available models from the Cloud Code Assist API.
|
||||||
|
func FetchAntigravityModels(accessToken, projectID string) ([]AntigravityModelInfo, error) {
|
||||||
|
reqBody, _ := json.Marshal(map[string]interface{}{
|
||||||
|
"project": projectID,
|
||||||
|
})
|
||||||
|
|
||||||
|
req, err := http.NewRequest("POST", antigravityBaseURL+"/v1internal:fetchAvailableModels", bytes.NewReader(reqBody))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("User-Agent", antigravityUserAgent)
|
||||||
|
req.Header.Set("X-Goog-Api-Client", antigravityXGoogClient)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 15 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return nil, fmt.Errorf("fetchAvailableModels failed (HTTP %d): %s", resp.StatusCode, truncateString(string(body), 200))
|
||||||
|
}
|
||||||
|
|
||||||
|
var result struct {
|
||||||
|
Models map[string]struct {
|
||||||
|
DisplayName string `json:"displayName"`
|
||||||
|
QuotaInfo struct {
|
||||||
|
RemainingFraction interface{} `json:"remainingFraction"`
|
||||||
|
ResetTime string `json:"resetTime"`
|
||||||
|
IsExhausted bool `json:"isExhausted"`
|
||||||
|
} `json:"quotaInfo"`
|
||||||
|
} `json:"models"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("parsing models response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var models []AntigravityModelInfo
|
||||||
|
for id, info := range result.Models {
|
||||||
|
models = append(models, AntigravityModelInfo{
|
||||||
|
ID: id,
|
||||||
|
DisplayName: info.DisplayName,
|
||||||
|
IsExhausted: info.QuotaInfo.IsExhausted,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure gemini-3-flash-preview and gemini-3-flash are in the list if they aren't already
|
||||||
|
hasFlashPreview := false
|
||||||
|
hasFlash := false
|
||||||
|
for _, m := range models {
|
||||||
|
if m.ID == "gemini-3-flash-preview" {
|
||||||
|
hasFlashPreview = true
|
||||||
|
}
|
||||||
|
if m.ID == "gemini-3-flash" {
|
||||||
|
hasFlash = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !hasFlashPreview {
|
||||||
|
models = append(models, AntigravityModelInfo{
|
||||||
|
ID: "gemini-3-flash-preview",
|
||||||
|
DisplayName: "Gemini 3 Flash (Preview)",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if !hasFlash {
|
||||||
|
models = append(models, AntigravityModelInfo{
|
||||||
|
ID: "gemini-3-flash",
|
||||||
|
DisplayName: "Gemini 3 Flash",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return models, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type AntigravityModelInfo struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
IsExhausted bool `json:"is_exhausted"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Helpers ---
|
||||||
|
|
||||||
|
func truncateString(s string, maxLen int) string {
|
||||||
|
if len(s) <= maxLen {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return s[:maxLen] + "..."
|
||||||
|
}
|
||||||
|
|
||||||
|
func randomString(n int) string {
|
||||||
|
const letters = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||||
|
b := make([]byte, n)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = letters[rand.Intn(len(letters))]
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) parseAntigravityError(statusCode int, body []byte) error {
|
||||||
|
var errResp struct {
|
||||||
|
Error struct {
|
||||||
|
Code int `json:"code"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
Details []map[string]interface{} `json:"details"`
|
||||||
|
} `json:"error"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &errResp); err != nil {
|
||||||
|
return fmt.Errorf("antigravity API error (HTTP %d): %s", statusCode, truncateString(string(body), 500))
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := errResp.Error.Message
|
||||||
|
if statusCode == 429 {
|
||||||
|
// Try to extract quota reset info
|
||||||
|
for _, detail := range errResp.Error.Details {
|
||||||
|
if typeVal, ok := detail["@type"].(string); ok && strings.HasSuffix(typeVal, "ErrorInfo") {
|
||||||
|
if metadata, ok := detail["metadata"].(map[string]interface{}); ok {
|
||||||
|
if delay, ok := metadata["quotaResetDelay"].(string); ok {
|
||||||
|
return fmt.Errorf("antigravity rate limit exceeded: %s (reset in %s)", msg, delay)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("antigravity rate limit exceeded: %s", msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("antigravity API error (%s): %s", errResp.Error.Status, msg)
|
||||||
|
}
|
||||||
56
pkg/providers/antigravity_provider_test.go
Normal file
56
pkg/providers/antigravity_provider_test.go
Normal file
|
|
@ -0,0 +1,56 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestBuildRequestUsesFunctionFieldsWhenToolCallNameMissing(t *testing.T) {
|
||||||
|
p := &AntigravityProvider{}
|
||||||
|
|
||||||
|
messages := []Message{
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
ToolCalls: []ToolCall{{
|
||||||
|
ID: "call_read_file_123",
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: "read_file",
|
||||||
|
Arguments: `{"path":"README.md"}`,
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Role: "tool",
|
||||||
|
ToolCallID: "call_read_file_123",
|
||||||
|
Content: "ok",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
req := p.buildRequest(messages, nil, "", nil)
|
||||||
|
if len(req.Contents) != 2 {
|
||||||
|
t.Fatalf("expected 2 contents, got %d", len(req.Contents))
|
||||||
|
}
|
||||||
|
|
||||||
|
modelPart := req.Contents[0].Parts[0]
|
||||||
|
if modelPart.FunctionCall == nil {
|
||||||
|
t.Fatal("expected functionCall in assistant message")
|
||||||
|
}
|
||||||
|
if modelPart.FunctionCall.Name != "read_file" {
|
||||||
|
t.Fatalf("expected functionCall name read_file, got %q", modelPart.FunctionCall.Name)
|
||||||
|
}
|
||||||
|
if got := modelPart.FunctionCall.Args["path"]; got != "README.md" {
|
||||||
|
t.Fatalf("expected functionCall args[path] to be README.md, got %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
toolPart := req.Contents[1].Parts[0]
|
||||||
|
if toolPart.FunctionResponse == nil {
|
||||||
|
t.Fatal("expected functionResponse in tool message")
|
||||||
|
}
|
||||||
|
if toolPart.FunctionResponse.Name != "read_file" {
|
||||||
|
t.Fatalf("expected functionResponse name read_file, got %q", toolPart.FunctionResponse.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveToolResponseNameInfersNameFromGeneratedCallID(t *testing.T) {
|
||||||
|
got := resolveToolResponseName("call_search_docs_999", map[string]string{})
|
||||||
|
if got != "search_docs" {
|
||||||
|
t.Fatalf("expected inferred tool name search_docs, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -336,7 +336,7 @@ func TestChat_PassesModelFlag(t *testing.T) {
|
||||||
|
|
||||||
_, err := p.Chat(context.Background(), []Message{
|
_, err := p.Chat(context.Background(), []Message{
|
||||||
{Role: "user", Content: "Hi"},
|
{Role: "user", Content: "Hi"},
|
||||||
}, nil, "claude-sonnet-4-5-20250929", nil)
|
}, nil, "claude-sonnet-4.6", nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Chat() error = %v", err)
|
t.Fatalf("Chat() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -346,7 +346,7 @@ func TestChat_PassesModelFlag(t *testing.T) {
|
||||||
if !strings.Contains(args, "--model") {
|
if !strings.Contains(args, "--model") {
|
||||||
t.Errorf("CLI args missing --model, got: %s", args)
|
t.Errorf("CLI args missing --model, got: %s", args)
|
||||||
}
|
}
|
||||||
if !strings.Contains(args, "claude-sonnet-4-5-20250929") {
|
if !strings.Contains(args, "claude-sonnet-4.6") {
|
||||||
t.Errorf("CLI args missing model name, got: %s", args)
|
t.Errorf("CLI args missing model name, got: %s", args)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -416,10 +416,12 @@ func TestChat_EmptyWorkspaceDoesNotSetDir(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCli(t *testing.T) {
|
func TestCreateProvider_ClaudeCli(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claude-cli"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
cfg.Agents.Defaults.Workspace = "/test/ws"
|
{ModelName: "claude-sonnet-4.6", Model: "claude-cli/claude-sonnet-4.6", Workspace: "/test/ws"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claude-sonnet-4.6"
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider(claude-cli) error = %v", err)
|
t.Fatalf("CreateProvider(claude-cli) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -435,9 +437,12 @@ func TestCreateProvider_ClaudeCli(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCode(t *testing.T) {
|
func TestCreateProvider_ClaudeCode(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claude-code"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{ModelName: "claude-code", Model: "claude-cli/claude-code"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claude-code"
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider(claude-code) error = %v", err)
|
t.Fatalf("CreateProvider(claude-code) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -448,9 +453,12 @@ func TestCreateProvider_ClaudeCode(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCodec(t *testing.T) {
|
func TestCreateProvider_ClaudeCodec(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claudecode"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{ModelName: "claudecode", Model: "claude-cli/claudecode"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claudecode"
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider(claudecode) error = %v", err)
|
t.Fatalf("CreateProvider(claudecode) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -461,10 +469,13 @@ func TestCreateProvider_ClaudeCodec(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCliDefaultWorkspace(t *testing.T) {
|
func TestCreateProvider_ClaudeCliDefaultWorkspace(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claude-cli"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{ModelName: "claude-cli", Model: "claude-cli/claude-sonnet"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claude-cli"
|
||||||
cfg.Agents.Defaults.Workspace = ""
|
cfg.Agents.Defaults.Workspace = ""
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider error = %v", err)
|
t.Fatalf("CreateProvider error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -48,7 +48,7 @@ func TestClaudeProvider_ChatRoundTrip(t *testing.T) {
|
||||||
provider := newClaudeProviderWithDelegate(delegate)
|
provider := newClaudeProviderWithDelegate(delegate)
|
||||||
|
|
||||||
messages := []Message{{Role: "user", Content: "Hello"}}
|
messages := []Message{{Role: "user", Content: "Hello"}}
|
||||||
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{"max_tokens": 1024})
|
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4.6", map[string]interface{}{"max_tokens": 1024})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Chat() error: %v", err)
|
t.Fatalf("Chat() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -65,8 +65,8 @@ func TestClaudeProvider_ChatRoundTrip(t *testing.T) {
|
||||||
|
|
||||||
func TestClaudeProvider_GetDefaultModel(t *testing.T) {
|
func TestClaudeProvider_GetDefaultModel(t *testing.T) {
|
||||||
p := NewClaudeProvider("test-token")
|
p := NewClaudeProvider("test-token")
|
||||||
if got := p.GetDefaultModel(); got != "claude-sonnet-4-5-20250929" {
|
if got := p.GetDefaultModel(); got != "claude-sonnet-4.6" {
|
||||||
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4-5-20250929")
|
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4.6")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -35,33 +35,6 @@ type providerSelection struct {
|
||||||
enableWebSearch bool
|
enableWebSearch bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func createClaudeAuthProvider(apiBase string) (LLMProvider, error) {
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = defaultAnthropicAPIBase
|
|
||||||
}
|
|
||||||
cred, err := getCredential("anthropic")
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
|
||||||
}
|
|
||||||
if cred == nil {
|
|
||||||
return nil, fmt.Errorf("no credentials for anthropic. Run: picoclaw auth login --provider anthropic")
|
|
||||||
}
|
|
||||||
return NewClaudeProviderWithTokenSourceAndBaseURL(cred.AccessToken, createClaudeTokenSource(), apiBase), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func createCodexAuthProvider(enableWebSearch bool) (LLMProvider, error) {
|
|
||||||
cred, err := getCredential("openai")
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
|
||||||
}
|
|
||||||
if cred == nil {
|
|
||||||
return nil, fmt.Errorf("no credentials for openai. Run: picoclaw auth login --provider openai")
|
|
||||||
}
|
|
||||||
p := NewCodexProviderWithTokenSource(cred.AccessToken, cred.AccountID, createCodexTokenSource())
|
|
||||||
p.enableWebSearch = enableWebSearch
|
|
||||||
return p, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
model := cfg.Agents.Defaults.Model
|
model := cfg.Agents.Defaults.Model
|
||||||
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
||||||
|
|
@ -332,29 +305,3 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
|
|
||||||
return sel, nil
|
return sel, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
|
||||||
sel, err := resolveProviderSelection(cfg)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
switch sel.providerType {
|
|
||||||
case providerTypeClaudeAuth:
|
|
||||||
return createClaudeAuthProvider(sel.apiBase)
|
|
||||||
case providerTypeCodexAuth:
|
|
||||||
return createCodexAuthProvider(sel.enableWebSearch)
|
|
||||||
case providerTypeCodexCLIToken:
|
|
||||||
c := NewCodexProviderWithTokenSource("", "", CreateCodexCliTokenSource())
|
|
||||||
c.enableWebSearch = sel.enableWebSearch
|
|
||||||
return c, nil
|
|
||||||
case providerTypeClaudeCLI:
|
|
||||||
return NewClaudeCliProvider(sel.workspace), nil
|
|
||||||
case providerTypeCodexCLI:
|
|
||||||
return NewCodexCliProvider(sel.workspace), nil
|
|
||||||
case providerTypeGitHubCopilot:
|
|
||||||
return NewGitHubCopilotProvider(sel.apiBase, sel.connectMode, sel.model)
|
|
||||||
default:
|
|
||||||
return NewHTTPProvider(sel.apiKey, sel.apiBase, sel.proxy), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
|
||||||
192
pkg/providers/factory_provider.go
Normal file
192
pkg/providers/factory_provider.go
Normal file
|
|
@ -0,0 +1,192 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
||||||
|
func createClaudeAuthProvider() (LLMProvider, error) {
|
||||||
|
cred, err := getCredential("anthropic")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
return nil, fmt.Errorf("no credentials for anthropic. Run: picoclaw auth login --provider anthropic")
|
||||||
|
}
|
||||||
|
return NewClaudeProviderWithTokenSource(cred.AccessToken, createClaudeTokenSource()), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// createCodexAuthProvider creates a Codex provider using OAuth credentials from auth store.
|
||||||
|
func createCodexAuthProvider() (LLMProvider, error) {
|
||||||
|
cred, err := getCredential("openai")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
return nil, fmt.Errorf("no credentials for openai. Run: picoclaw auth login --provider openai")
|
||||||
|
}
|
||||||
|
return NewCodexProviderWithTokenSource(cred.AccessToken, cred.AccountID, createCodexTokenSource()), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtractProtocol extracts the protocol prefix and model identifier from a model string.
|
||||||
|
// If no prefix is specified, it defaults to "openai".
|
||||||
|
// Examples:
|
||||||
|
// - "openai/gpt-4o" -> ("openai", "gpt-4o")
|
||||||
|
// - "anthropic/claude-sonnet-4.6" -> ("anthropic", "claude-sonnet-4.6")
|
||||||
|
// - "gpt-4o" -> ("openai", "gpt-4o") // default protocol
|
||||||
|
func ExtractProtocol(model string) (protocol, modelID string) {
|
||||||
|
model = strings.TrimSpace(model)
|
||||||
|
protocol, modelID, found := strings.Cut(model, "/")
|
||||||
|
if !found {
|
||||||
|
return "openai", model
|
||||||
|
}
|
||||||
|
return protocol, modelID
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateProviderFromConfig creates a provider based on the ModelConfig.
|
||||||
|
// It uses the protocol prefix in the Model field to determine which provider to create.
|
||||||
|
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot
|
||||||
|
// Returns the provider, the model ID (without protocol prefix), and any error.
|
||||||
|
func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) {
|
||||||
|
if cfg == nil {
|
||||||
|
return nil, "", fmt.Errorf("config is nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Model == "" {
|
||||||
|
return nil, "", fmt.Errorf("model is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
protocol, modelID := ExtractProtocol(cfg.Model)
|
||||||
|
|
||||||
|
switch protocol {
|
||||||
|
case "openai":
|
||||||
|
// OpenAI with OAuth/token auth (Codex-style)
|
||||||
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
||||||
|
provider, err := createCodexAuthProvider()
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
return provider, modelID, nil
|
||||||
|
}
|
||||||
|
// OpenAI with API key
|
||||||
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
|
}
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = getDefaultAPIBase(protocol)
|
||||||
|
}
|
||||||
|
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil
|
||||||
|
|
||||||
|
case "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||||
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
|
"volcengine", "vllm", "qwen":
|
||||||
|
// All other OpenAI-compatible HTTP providers
|
||||||
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
|
}
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = getDefaultAPIBase(protocol)
|
||||||
|
}
|
||||||
|
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil
|
||||||
|
|
||||||
|
case "anthropic":
|
||||||
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
||||||
|
// Use OAuth credentials from auth store
|
||||||
|
provider, err := createClaudeAuthProvider()
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
return provider, modelID, nil
|
||||||
|
}
|
||||||
|
// Use API key with HTTP API
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = "https://api.anthropic.com/v1"
|
||||||
|
}
|
||||||
|
if cfg.APIKey == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model)
|
||||||
|
}
|
||||||
|
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil
|
||||||
|
|
||||||
|
case "antigravity":
|
||||||
|
return NewAntigravityProvider(), modelID, nil
|
||||||
|
|
||||||
|
case "claude-cli", "claudecli":
|
||||||
|
workspace := cfg.Workspace
|
||||||
|
if workspace == "" {
|
||||||
|
workspace = "."
|
||||||
|
}
|
||||||
|
return NewClaudeCliProvider(workspace), modelID, nil
|
||||||
|
|
||||||
|
case "codex-cli", "codexcli":
|
||||||
|
workspace := cfg.Workspace
|
||||||
|
if workspace == "" {
|
||||||
|
workspace = "."
|
||||||
|
}
|
||||||
|
return NewCodexCliProvider(workspace), modelID, nil
|
||||||
|
|
||||||
|
case "github-copilot", "copilot":
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = "localhost:4321"
|
||||||
|
}
|
||||||
|
connectMode := cfg.ConnectMode
|
||||||
|
if connectMode == "" {
|
||||||
|
connectMode = "grpc"
|
||||||
|
}
|
||||||
|
provider, err := NewGitHubCopilotProvider(apiBase, connectMode, modelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
return provider, modelID, nil
|
||||||
|
|
||||||
|
default:
|
||||||
|
return nil, "", fmt.Errorf("unknown protocol %q in model %q", protocol, cfg.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
||||||
|
func getDefaultAPIBase(protocol string) string {
|
||||||
|
switch protocol {
|
||||||
|
case "openai":
|
||||||
|
return "https://api.openai.com/v1"
|
||||||
|
case "openrouter":
|
||||||
|
return "https://openrouter.ai/api/v1"
|
||||||
|
case "groq":
|
||||||
|
return "https://api.groq.com/openai/v1"
|
||||||
|
case "zhipu":
|
||||||
|
return "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
case "gemini":
|
||||||
|
return "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
case "nvidia":
|
||||||
|
return "https://integrate.api.nvidia.com/v1"
|
||||||
|
case "ollama":
|
||||||
|
return "http://localhost:11434/v1"
|
||||||
|
case "moonshot":
|
||||||
|
return "https://api.moonshot.cn/v1"
|
||||||
|
case "shengsuanyun":
|
||||||
|
return "https://router.shengsuanyun.com/api/v1"
|
||||||
|
case "deepseek":
|
||||||
|
return "https://api.deepseek.com/v1"
|
||||||
|
case "cerebras":
|
||||||
|
return "https://api.cerebras.ai/v1"
|
||||||
|
case "volcengine":
|
||||||
|
return "https://ark.cn-beijing.volces.com/api/v3"
|
||||||
|
case "qwen":
|
||||||
|
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
|
case "vllm":
|
||||||
|
return "http://localhost:8000/v1"
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
249
pkg/providers/factory_provider_test.go
Normal file
249
pkg/providers/factory_provider_test.go
Normal file
|
|
@ -0,0 +1,249 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestExtractProtocol(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
model string
|
||||||
|
wantProtocol string
|
||||||
|
wantModelID string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "openai with prefix",
|
||||||
|
model: "openai/gpt-4o",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "gpt-4o",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "anthropic with prefix",
|
||||||
|
model: "anthropic/claude-sonnet-4.6",
|
||||||
|
wantProtocol: "anthropic",
|
||||||
|
wantModelID: "claude-sonnet-4.6",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no prefix - defaults to openai",
|
||||||
|
model: "gpt-4o",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "gpt-4o",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "groq with prefix",
|
||||||
|
model: "groq/llama-3.1-70b",
|
||||||
|
wantProtocol: "groq",
|
||||||
|
wantModelID: "llama-3.1-70b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty string",
|
||||||
|
model: "",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with whitespace",
|
||||||
|
model: " openai/gpt-4 ",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "gpt-4",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple slashes",
|
||||||
|
model: "nvidia/meta/llama-3.1-8b",
|
||||||
|
wantProtocol: "nvidia",
|
||||||
|
wantModelID: "meta/llama-3.1-8b",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
protocol, modelID := ExtractProtocol(tt.model)
|
||||||
|
if protocol != tt.wantProtocol {
|
||||||
|
t.Errorf("ExtractProtocol(%q) protocol = %q, want %q", tt.model, protocol, tt.wantProtocol)
|
||||||
|
}
|
||||||
|
if modelID != tt.wantModelID {
|
||||||
|
t.Errorf("ExtractProtocol(%q) modelID = %q, want %q", tt.model, modelID, tt.wantModelID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_OpenAI(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-openai",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
APIKey: "test-key",
|
||||||
|
APIBase: "https://api.example.com/v1",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "gpt-4o" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "gpt-4o")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
protocol string
|
||||||
|
}{
|
||||||
|
{"openai", "openai"},
|
||||||
|
{"groq", "groq"},
|
||||||
|
{"openrouter", "openrouter"},
|
||||||
|
{"cerebras", "cerebras"},
|
||||||
|
{"qwen", "qwen"},
|
||||||
|
{"vllm", "vllm"},
|
||||||
|
{"deepseek", "deepseek"},
|
||||||
|
{"ollama", "ollama"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-" + tt.protocol,
|
||||||
|
Model: tt.protocol + "/test-model",
|
||||||
|
APIKey: "test-key",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify we got an HTTPProvider for all these protocols
|
||||||
|
if _, ok := provider.(*HTTPProvider); !ok {
|
||||||
|
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_Anthropic(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-anthropic",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
APIKey: "test-key",
|
||||||
|
}
|
||||||
|
|
||||||
|
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 != "claude-sonnet-4.6" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_Antigravity(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-antigravity",
|
||||||
|
Model: "antigravity/gemini-2.0-flash",
|
||||||
|
}
|
||||||
|
|
||||||
|
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 != "gemini-2.0-flash" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "gemini-2.0-flash")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_ClaudeCLI(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-claude-cli",
|
||||||
|
Model: "claude-cli/claude-sonnet-4.6",
|
||||||
|
}
|
||||||
|
|
||||||
|
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 != "claude-sonnet-4.6" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_CodexCLI(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-codex-cli",
|
||||||
|
Model: "codex-cli/codex",
|
||||||
|
}
|
||||||
|
|
||||||
|
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 != "codex" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "codex")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_MissingAPIKey(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-no-key",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() expected error for missing API key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_UnknownProtocol(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-unknown",
|
||||||
|
Model: "unknown-protocol/model",
|
||||||
|
APIKey: "test-key",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() expected error for unknown protocol")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_NilConfig(t *testing.T) {
|
||||||
|
_, _, err := CreateProviderFromConfig(nil)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig(nil) expected error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_EmptyModel(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-empty",
|
||||||
|
Model: "",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() expected error for empty model")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -79,7 +79,7 @@ func TestResolveProviderSelection(t *testing.T) {
|
||||||
{
|
{
|
||||||
name: "anthropic oauth routes to claude auth provider",
|
name: "anthropic oauth routes to claude auth provider",
|
||||||
setup: func(cfg *config.Config) {
|
setup: func(cfg *config.Config) {
|
||||||
cfg.Agents.Defaults.Model = "claude-sonnet-4-5-20250929"
|
cfg.Agents.Defaults.Model = "claude-sonnet-4.6"
|
||||||
cfg.Providers.Anthropic.AuthMethod = "oauth"
|
cfg.Providers.Anthropic.AuthMethod = "oauth"
|
||||||
},
|
},
|
||||||
wantType: providerTypeClaudeAuth,
|
wantType: providerTypeClaudeAuth,
|
||||||
|
|
@ -196,10 +196,17 @@ func TestResolveProviderSelection(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProviderReturnsHTTPProviderForOpenRouter(t *testing.T) {
|
func TestCreateProviderReturnsHTTPProviderForOpenRouter(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Model = "openrouter/auto"
|
cfg.Agents.Defaults.Model = "test-openrouter"
|
||||||
cfg.Providers.OpenRouter.APIKey = "sk-or-test"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-openrouter",
|
||||||
|
Model: "openrouter/auto",
|
||||||
|
APIKey: "sk-or-test",
|
||||||
|
APIBase: "https://openrouter.ai/api/v1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider() error = %v", err)
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -211,9 +218,16 @@ func TestCreateProviderReturnsHTTPProviderForOpenRouter(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProviderReturnsCodexCliProviderForCodexCode(t *testing.T) {
|
func TestCreateProviderReturnsCodexCliProviderForCodexCode(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "codex-code"
|
cfg.Agents.Defaults.Model = "test-codex"
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-codex",
|
||||||
|
Model: "codex-cli/codex-model",
|
||||||
|
Workspace: "/tmp/workspace",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider() error = %v", err)
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -223,18 +237,24 @@ func TestCreateProviderReturnsCodexCliProviderForCodexCode(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCreateProviderReturnsCodexProviderForCodexCliAuthMethod(t *testing.T) {
|
func TestCreateProviderReturnsClaudeCliProviderForClaudeCli(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "openai"
|
cfg.Agents.Defaults.Model = "test-claude-cli"
|
||||||
cfg.Providers.OpenAI.AuthMethod = "codex-cli"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-claude-cli",
|
||||||
|
Model: "claude-cli/claude-sonnet",
|
||||||
|
Workspace: "/tmp/workspace",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider() error = %v", err)
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, ok := provider.(*CodexProvider); !ok {
|
if _, ok := provider.(*ClaudeCliProvider); !ok {
|
||||||
t.Fatalf("provider type = %T, want *CodexProvider", provider)
|
t.Fatalf("provider type = %T, want *ClaudeCliProvider", provider)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -252,48 +272,28 @@ func TestCreateProviderReturnsClaudeProviderForAnthropicOAuth(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "anthropic"
|
cfg.Agents.Defaults.Model = "test-claude-oauth"
|
||||||
cfg.Providers.Anthropic.AuthMethod = "oauth"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
cfg.Providers.Anthropic.APIBase = "https://proxy.example.com/v1"
|
{
|
||||||
|
ModelName: "test-claude-oauth",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider() error = %v", err)
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
claudeProvider, ok := provider.(*ClaudeProvider)
|
if _, ok := provider.(*ClaudeProvider); !ok {
|
||||||
if !ok {
|
|
||||||
t.Fatalf("provider type = %T, want *ClaudeProvider", provider)
|
t.Fatalf("provider type = %T, want *ClaudeProvider", provider)
|
||||||
}
|
}
|
||||||
if got := claudeProvider.delegate.BaseURL(); got != "https://proxy.example.com" {
|
// TODO: Test custom APIBase when createClaudeAuthProvider supports it
|
||||||
t.Fatalf("anthropic baseURL = %q, want %q", got, "https://proxy.example.com")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCreateProviderReturnsCodexProviderForOpenAIOAuth(t *testing.T) {
|
func TestCreateProviderReturnsCodexProviderForOpenAIOAuth(t *testing.T) {
|
||||||
originalGetCredential := getCredential
|
// TODO: This test requires openai protocol to support auth_method: "oauth"
|
||||||
t.Cleanup(func() { getCredential = originalGetCredential })
|
// which is not yet implemented in the new factory_provider.go
|
||||||
|
t.Skip("OpenAI OAuth via model_list not yet implemented")
|
||||||
getCredential = func(provider string) (*auth.AuthCredential, error) {
|
|
||||||
if provider != "openai" {
|
|
||||||
t.Fatalf("provider = %q, want openai", provider)
|
|
||||||
}
|
|
||||||
return &auth.AuthCredential{
|
|
||||||
AccessToken: "openai-token",
|
|
||||||
AccountID: "acct_123",
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
cfg := config.DefaultConfig()
|
|
||||||
cfg.Agents.Defaults.Provider = "openai"
|
|
||||||
cfg.Providers.OpenAI.AuthMethod = "oauth"
|
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("CreateProvider() error = %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, ok := provider.(*CodexProvider); !ok {
|
|
||||||
t.Fatalf("provider type = %T, want *CodexProvider", provider)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -22,6 +22,12 @@ func NewHTTPProvider(apiKey, apiBase, proxy string) *HTTPProvider {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func NewHTTPProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string) *HTTPProvider {
|
||||||
|
return &HTTPProvider{
|
||||||
|
delegate: openai_compat.NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
return p.delegate.Chat(ctx, messages, tools, model, options)
|
return p.delegate.Chat(ctx, messages, tools, model, options)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
49
pkg/providers/legacy_provider.go
Normal file
49
pkg/providers/legacy_provider.go
Normal file
|
|
@ -0,0 +1,49 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CreateProvider creates a provider based on the configuration.
|
||||||
|
// It uses the model_list configuration (new format) to create providers.
|
||||||
|
// The old providers config is automatically converted to model_list during config loading.
|
||||||
|
// Returns the provider, the model ID to use, and any error.
|
||||||
|
func CreateProvider(cfg *config.Config) (LLMProvider, string, error) {
|
||||||
|
model := cfg.Agents.Defaults.Model
|
||||||
|
|
||||||
|
// Ensure model_list is populated (should be done by LoadConfig, but handle edge cases)
|
||||||
|
if len(cfg.ModelList) == 0 && cfg.HasProvidersConfig() {
|
||||||
|
cfg.ModelList = config.ConvertProvidersToModelList(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Must have model_list at this point
|
||||||
|
if len(cfg.ModelList) == 0 {
|
||||||
|
return nil, "", fmt.Errorf("no providers configured. Please add entries to model_list in your config")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get model config from model_list
|
||||||
|
modelCfg, err := cfg.GetModelConfig(model)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", fmt.Errorf("model %q not found in model_list: %w", model, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Inject global workspace if not set in model config
|
||||||
|
if modelCfg.Workspace == "" {
|
||||||
|
modelCfg.Workspace = cfg.WorkspacePath()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use factory to create provider
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(modelCfg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", fmt.Errorf("failed to create provider for model %q: %w", model, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return provider, modelID, nil
|
||||||
|
}
|
||||||
|
|
@ -22,14 +22,21 @@ type UsageInfo = protocoltypes.UsageInfo
|
||||||
type Message = protocoltypes.Message
|
type Message = protocoltypes.Message
|
||||||
type ToolDefinition = protocoltypes.ToolDefinition
|
type ToolDefinition = protocoltypes.ToolDefinition
|
||||||
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
||||||
|
type ExtraContent = protocoltypes.ExtraContent
|
||||||
|
type GoogleExtra = protocoltypes.GoogleExtra
|
||||||
|
|
||||||
type Provider struct {
|
type Provider struct {
|
||||||
apiKey string
|
apiKey string
|
||||||
apiBase string
|
apiBase string
|
||||||
httpClient *http.Client
|
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
|
||||||
|
httpClient *http.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewProvider(apiKey, apiBase, proxy string) *Provider {
|
func NewProvider(apiKey, apiBase, proxy string) *Provider {
|
||||||
|
return NewProviderWithMaxTokensField(apiKey, apiBase, proxy, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string) *Provider {
|
||||||
client := &http.Client{
|
client := &http.Client{
|
||||||
Timeout: 120 * time.Second,
|
Timeout: 120 * time.Second,
|
||||||
}
|
}
|
||||||
|
|
@ -46,9 +53,10 @@ func NewProvider(apiKey, apiBase, proxy string) *Provider {
|
||||||
}
|
}
|
||||||
|
|
||||||
return &Provider{
|
return &Provider{
|
||||||
apiKey: apiKey,
|
apiKey: apiKey,
|
||||||
apiBase: strings.TrimRight(apiBase, "/"),
|
apiBase: strings.TrimRight(apiBase, "/"),
|
||||||
httpClient: client,
|
maxTokensField: maxTokensField,
|
||||||
|
httpClient: client,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -70,12 +78,18 @@ func (p *Provider) Chat(ctx context.Context, messages []Message, tools []ToolDef
|
||||||
}
|
}
|
||||||
|
|
||||||
if maxTokens, ok := asInt(options["max_tokens"]); ok {
|
if maxTokens, ok := asInt(options["max_tokens"]); ok {
|
||||||
lowerModel := strings.ToLower(model)
|
// Use configured maxTokensField if specified, otherwise fallback to model-based detection
|
||||||
if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") {
|
fieldName := p.maxTokensField
|
||||||
requestBody["max_completion_tokens"] = maxTokens
|
if fieldName == "" {
|
||||||
} else {
|
// Fallback: detect from model name for backward compatibility
|
||||||
requestBody["max_tokens"] = maxTokens
|
lowerModel := strings.ToLower(model)
|
||||||
|
if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") || strings.Contains(lowerModel, "gpt-5") {
|
||||||
|
fieldName = "max_completion_tokens"
|
||||||
|
} else {
|
||||||
|
fieldName = "max_tokens"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
requestBody[fieldName] = maxTokens
|
||||||
}
|
}
|
||||||
|
|
||||||
if temperature, ok := asFloat(options["temperature"]); ok {
|
if temperature, ok := asFloat(options["temperature"]); ok {
|
||||||
|
|
@ -133,6 +147,11 @@ func parseResponse(body []byte) (*LLMResponse, error) {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Arguments string `json:"arguments"`
|
Arguments string `json:"arguments"`
|
||||||
} `json:"function"`
|
} `json:"function"`
|
||||||
|
ExtraContent *struct {
|
||||||
|
Google *struct {
|
||||||
|
ThoughtSignature string `json:"thought_signature"`
|
||||||
|
} `json:"google"`
|
||||||
|
} `json:"extra_content"`
|
||||||
} `json:"tool_calls"`
|
} `json:"tool_calls"`
|
||||||
} `json:"message"`
|
} `json:"message"`
|
||||||
FinishReason string `json:"finish_reason"`
|
FinishReason string `json:"finish_reason"`
|
||||||
|
|
@ -157,6 +176,12 @@ func parseResponse(body []byte) (*LLMResponse, error) {
|
||||||
arguments := make(map[string]interface{})
|
arguments := make(map[string]interface{})
|
||||||
name := ""
|
name := ""
|
||||||
|
|
||||||
|
// Extract thought_signature from Gemini/Google-specific extra content
|
||||||
|
thoughtSignature := ""
|
||||||
|
if tc.ExtraContent != nil && tc.ExtraContent.Google != nil {
|
||||||
|
thoughtSignature = tc.ExtraContent.Google.ThoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
if tc.Function != nil {
|
if tc.Function != nil {
|
||||||
name = tc.Function.Name
|
name = tc.Function.Name
|
||||||
if tc.Function.Arguments != "" {
|
if tc.Function.Arguments != "" {
|
||||||
|
|
@ -167,11 +192,23 @@ func parseResponse(body []byte) (*LLMResponse, error) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
toolCalls = append(toolCalls, ToolCall{
|
// Build ToolCall with ExtraContent for Gemini 3 thought_signature persistence
|
||||||
ID: tc.ID,
|
toolCall := ToolCall{
|
||||||
Name: name,
|
ID: tc.ID,
|
||||||
Arguments: arguments,
|
Name: name,
|
||||||
})
|
Arguments: arguments,
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
|
}
|
||||||
|
|
||||||
|
if thoughtSignature != "" {
|
||||||
|
toolCall.ExtraContent = &ExtraContent{
|
||||||
|
Google: &GoogleExtra{
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
toolCalls = append(toolCalls, toolCall)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &LLMResponse{
|
return &LLMResponse{
|
||||||
|
|
|
||||||
|
|
@ -1,16 +1,27 @@
|
||||||
package protocoltypes
|
package protocoltypes
|
||||||
|
|
||||||
type ToolCall struct {
|
type ToolCall struct {
|
||||||
ID string `json:"id"`
|
ID string `json:"id"`
|
||||||
Type string `json:"type,omitempty"`
|
Type string `json:"type,omitempty"`
|
||||||
Function *FunctionCall `json:"function,omitempty"`
|
Function *FunctionCall `json:"function,omitempty"`
|
||||||
Name string `json:"name,omitempty"`
|
Name string `json:"name,omitempty"`
|
||||||
Arguments map[string]interface{} `json:"arguments,omitempty"`
|
Arguments map[string]interface{} `json:"arguments,omitempty"`
|
||||||
|
ThoughtSignature string `json:"-"` // Internal use only
|
||||||
|
ExtraContent *ExtraContent `json:"extra_content,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ExtraContent struct {
|
||||||
|
Google *GoogleExtra `json:"google,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type GoogleExtra struct {
|
||||||
|
ThoughtSignature string `json:"thought_signature,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type FunctionCall struct {
|
type FunctionCall struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Arguments string `json:"arguments"`
|
Arguments string `json:"arguments"`
|
||||||
|
ThoughtSignature string `json:"thought_signature,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type LLMResponse struct {
|
type LLMResponse struct {
|
||||||
|
|
|
||||||
54
pkg/providers/toolcall_utils.go
Normal file
54
pkg/providers/toolcall_utils.go
Normal file
|
|
@ -0,0 +1,54 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
// NormalizeToolCall normalizes a ToolCall to ensure all fields are properly populated.
|
||||||
|
// It handles cases where Name/Arguments might be in different locations (top-level vs Function)
|
||||||
|
// and ensures both are populated consistently.
|
||||||
|
func NormalizeToolCall(tc ToolCall) ToolCall {
|
||||||
|
normalized := tc
|
||||||
|
|
||||||
|
// Ensure Name is populated from Function if not set
|
||||||
|
if normalized.Name == "" && normalized.Function != nil {
|
||||||
|
normalized.Name = normalized.Function.Name
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure Arguments is not nil
|
||||||
|
if normalized.Arguments == nil {
|
||||||
|
normalized.Arguments = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse Arguments from Function.Arguments if not already set
|
||||||
|
if len(normalized.Arguments) == 0 && normalized.Function != nil && normalized.Function.Arguments != "" {
|
||||||
|
var parsed map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(normalized.Function.Arguments), &parsed); err == nil && parsed != nil {
|
||||||
|
normalized.Arguments = parsed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure Function is populated with consistent values
|
||||||
|
argsJSON, _ := json.Marshal(normalized.Arguments)
|
||||||
|
if normalized.Function == nil {
|
||||||
|
normalized.Function = &FunctionCall{
|
||||||
|
Name: normalized.Name,
|
||||||
|
Arguments: string(argsJSON),
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if normalized.Function.Name == "" {
|
||||||
|
normalized.Function.Name = normalized.Name
|
||||||
|
}
|
||||||
|
if normalized.Name == "" {
|
||||||
|
normalized.Name = normalized.Function.Name
|
||||||
|
}
|
||||||
|
if normalized.Function.Arguments == "" {
|
||||||
|
normalized.Function.Arguments = string(argsJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
|
@ -14,6 +14,8 @@ type UsageInfo = protocoltypes.UsageInfo
|
||||||
type Message = protocoltypes.Message
|
type Message = protocoltypes.Message
|
||||||
type ToolDefinition = protocoltypes.ToolDefinition
|
type ToolDefinition = protocoltypes.ToolDefinition
|
||||||
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
||||||
|
type ExtraContent = protocoltypes.ExtraContent
|
||||||
|
type GoogleExtra = protocoltypes.GoogleExtra
|
||||||
|
|
||||||
type LLMProvider interface {
|
type LLMProvider interface {
|
||||||
Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error)
|
Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error)
|
||||||
|
|
|
||||||
311
pkg/skills/clawhub_registry.go
Normal file
311
pkg/skills/clawhub_registry.go
Normal file
|
|
@ -0,0 +1,311 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultClawHubTimeout = 30 * time.Second
|
||||||
|
defaultMaxZipSize = 50 * 1024 * 1024 // 50 MB
|
||||||
|
defaultMaxResponseSize = 2 * 1024 * 1024 // 2 MB
|
||||||
|
)
|
||||||
|
|
||||||
|
// ClawHubRegistry implements SkillRegistry for the ClawHub platform.
|
||||||
|
type ClawHubRegistry struct {
|
||||||
|
baseURL string
|
||||||
|
authToken string // Optional - for elevated rate limits
|
||||||
|
searchPath string // Search API
|
||||||
|
skillsPath string // For retrieving skill metadata
|
||||||
|
downloadPath string // For fetching ZIP files for download
|
||||||
|
maxZipSize int
|
||||||
|
maxResponseSize int
|
||||||
|
client *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewClawHubRegistry creates a new ClawHub registry client from config.
|
||||||
|
func NewClawHubRegistry(cfg ClawHubConfig) *ClawHubRegistry {
|
||||||
|
baseURL := cfg.BaseURL
|
||||||
|
if baseURL == "" {
|
||||||
|
baseURL = "https://clawhub.ai"
|
||||||
|
}
|
||||||
|
searchPath := cfg.SearchPath
|
||||||
|
if searchPath == "" {
|
||||||
|
searchPath = "/api/v1/search"
|
||||||
|
}
|
||||||
|
skillsPath := cfg.SkillsPath
|
||||||
|
if skillsPath == "" {
|
||||||
|
skillsPath = "/api/v1/skills"
|
||||||
|
}
|
||||||
|
downloadPath := cfg.DownloadPath
|
||||||
|
if downloadPath == "" {
|
||||||
|
downloadPath = "/api/v1/download"
|
||||||
|
}
|
||||||
|
|
||||||
|
timeout := defaultClawHubTimeout
|
||||||
|
if cfg.Timeout > 0 {
|
||||||
|
timeout = time.Duration(cfg.Timeout) * time.Second
|
||||||
|
}
|
||||||
|
|
||||||
|
maxZip := defaultMaxZipSize
|
||||||
|
if cfg.MaxZipSize > 0 {
|
||||||
|
maxZip = cfg.MaxZipSize
|
||||||
|
}
|
||||||
|
|
||||||
|
maxResp := defaultMaxResponseSize
|
||||||
|
if cfg.MaxResponseSize > 0 {
|
||||||
|
maxResp = cfg.MaxResponseSize
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ClawHubRegistry{
|
||||||
|
baseURL: baseURL,
|
||||||
|
authToken: cfg.AuthToken,
|
||||||
|
searchPath: searchPath,
|
||||||
|
skillsPath: skillsPath,
|
||||||
|
downloadPath: downloadPath,
|
||||||
|
maxZipSize: maxZip,
|
||||||
|
maxResponseSize: maxResp,
|
||||||
|
client: &http.Client{
|
||||||
|
Timeout: timeout,
|
||||||
|
Transport: &http.Transport{
|
||||||
|
MaxIdleConns: 5,
|
||||||
|
IdleConnTimeout: 30 * time.Second,
|
||||||
|
TLSHandshakeTimeout: 10 * time.Second,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) Name() string {
|
||||||
|
return "clawhub"
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Search ---
|
||||||
|
|
||||||
|
type clawhubSearchResponse struct {
|
||||||
|
Results []clawhubSearchResult `json:"results"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type clawhubSearchResult struct {
|
||||||
|
Score float64 `json:"score"`
|
||||||
|
Slug *string `json:"slug"`
|
||||||
|
DisplayName *string `json:"displayName"`
|
||||||
|
Summary *string `json:"summary"`
|
||||||
|
Version *string `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) Search(ctx context.Context, query string, limit int) ([]SearchResult, error) {
|
||||||
|
u, err := url.Parse(c.baseURL + c.searchPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid base URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
q := u.Query()
|
||||||
|
q.Set("q", query)
|
||||||
|
if limit > 0 {
|
||||||
|
q.Set("limit", fmt.Sprintf("%d", limit))
|
||||||
|
}
|
||||||
|
u.RawQuery = q.Encode()
|
||||||
|
|
||||||
|
body, err := c.doGet(ctx, u.String())
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("search request failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp clawhubSearchResponse
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse search response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]SearchResult, 0, len(resp.Results))
|
||||||
|
for _, r := range resp.Results {
|
||||||
|
slug := utils.DerefStr(r.Slug, "")
|
||||||
|
if slug == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
summary := utils.DerefStr(r.Summary, "")
|
||||||
|
if summary == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
displayName := utils.DerefStr(r.DisplayName, "")
|
||||||
|
if displayName == "" {
|
||||||
|
displayName = slug
|
||||||
|
}
|
||||||
|
|
||||||
|
results = append(results, SearchResult{
|
||||||
|
Score: r.Score,
|
||||||
|
Slug: slug,
|
||||||
|
DisplayName: displayName,
|
||||||
|
Summary: summary,
|
||||||
|
Version: utils.DerefStr(r.Version, ""),
|
||||||
|
RegistryName: c.Name(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- GetSkillMeta ---
|
||||||
|
|
||||||
|
type clawhubSkillResponse struct {
|
||||||
|
Slug string `json:"slug"`
|
||||||
|
DisplayName string `json:"displayName"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
LatestVersion *clawhubVersionInfo `json:"latestVersion"`
|
||||||
|
Moderation *clawhubModerationInfo `json:"moderation"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type clawhubVersionInfo struct {
|
||||||
|
Version string `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type clawhubModerationInfo struct {
|
||||||
|
IsMalwareBlocked bool `json:"isMalwareBlocked"`
|
||||||
|
IsSuspicious bool `json:"isSuspicious"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) GetSkillMeta(ctx context.Context, slug string) (*SkillMeta, error) {
|
||||||
|
if err := utils.ValidateSkillIdentifier(slug); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid slug %q: error: %s", slug, err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
u := c.baseURL + c.skillsPath + "/" + url.PathEscape(slug)
|
||||||
|
|
||||||
|
body, err := c.doGet(ctx, u)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("skill metadata request failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp clawhubSkillResponse
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse skill metadata: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
meta := &SkillMeta{
|
||||||
|
Slug: resp.Slug,
|
||||||
|
DisplayName: resp.DisplayName,
|
||||||
|
Summary: resp.Summary,
|
||||||
|
RegistryName: c.Name(),
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.LatestVersion != nil {
|
||||||
|
meta.LatestVersion = resp.LatestVersion.Version
|
||||||
|
}
|
||||||
|
if resp.Moderation != nil {
|
||||||
|
meta.IsMalwareBlocked = resp.Moderation.IsMalwareBlocked
|
||||||
|
meta.IsSuspicious = resp.Moderation.IsSuspicious
|
||||||
|
}
|
||||||
|
|
||||||
|
return meta, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- DownloadAndInstall ---
|
||||||
|
|
||||||
|
// DownloadAndInstall fetches metadata (with fallback), resolves version,
|
||||||
|
// downloads the skill ZIP, and extracts it to targetDir.
|
||||||
|
// Returns an InstallResult for the caller to use for moderation decisions.
|
||||||
|
func (c *ClawHubRegistry) DownloadAndInstall(ctx context.Context, slug, version, targetDir string) (*InstallResult, error) {
|
||||||
|
if err := utils.ValidateSkillIdentifier(slug); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid slug %q: error: %s", slug, err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 1: Fetch metadata (with fallback).
|
||||||
|
result := &InstallResult{}
|
||||||
|
meta, err := c.GetSkillMeta(ctx, slug)
|
||||||
|
if err != nil {
|
||||||
|
// Fallback: proceed without metadata.
|
||||||
|
meta = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if meta != nil {
|
||||||
|
result.IsMalwareBlocked = meta.IsMalwareBlocked
|
||||||
|
result.IsSuspicious = meta.IsSuspicious
|
||||||
|
result.Summary = meta.Summary
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 2: Resolve version.
|
||||||
|
installVersion := version
|
||||||
|
if installVersion == "" && meta != nil {
|
||||||
|
installVersion = meta.LatestVersion
|
||||||
|
}
|
||||||
|
if installVersion == "" {
|
||||||
|
installVersion = "latest"
|
||||||
|
}
|
||||||
|
result.Version = installVersion
|
||||||
|
|
||||||
|
// Step 3: Download ZIP to temp file (streams in ~32KB chunks).
|
||||||
|
u, err := url.Parse(c.baseURL + c.downloadPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid base URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
q := u.Query()
|
||||||
|
q.Set("slug", slug)
|
||||||
|
if installVersion != "latest" {
|
||||||
|
q.Set("version", installVersion)
|
||||||
|
}
|
||||||
|
u.RawQuery = q.Encode()
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "GET", u.String(), nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
if c.authToken != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+c.authToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
tmpPath, err := utils.DownloadToFile(ctx, c.client, req, int64(c.maxZipSize))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("download failed: %w", err)
|
||||||
|
}
|
||||||
|
defer os.Remove(tmpPath)
|
||||||
|
|
||||||
|
// Step 4: Extract from file on disk.
|
||||||
|
if err := utils.ExtractZipFile(tmpPath, targetDir); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- HTTP helper ---
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) doGet(ctx context.Context, urlStr string) ([]byte, error) {
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "GET", urlStr, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Accept", "application/json")
|
||||||
|
if c.authToken != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+c.authToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := c.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
// Limit response body read to prevent memory issues.
|
||||||
|
body, err := io.ReadAll(io.LimitReader(resp.Body, int64(c.maxResponseSize)))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||||
|
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
return body, nil
|
||||||
|
}
|
||||||
256
pkg/skills/clawhub_registry_test.go
Normal file
256
pkg/skills/clawhub_registry_test.go
Normal file
|
|
@ -0,0 +1,256 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestRegistry(serverURL, authToken string) *ClawHubRegistry {
|
||||||
|
return NewClawHubRegistry(ClawHubConfig{
|
||||||
|
Enabled: true,
|
||||||
|
BaseURL: serverURL,
|
||||||
|
AuthToken: authToken,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistrySearch(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
assert.Equal(t, "/api/v1/search", r.URL.Path)
|
||||||
|
assert.Equal(t, "github", r.URL.Query().Get("q"))
|
||||||
|
|
||||||
|
slug := "github"
|
||||||
|
name := "GitHub Integration"
|
||||||
|
summary := "Interact with GitHub repos"
|
||||||
|
version := "1.0.0"
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(clawhubSearchResponse{
|
||||||
|
Results: []clawhubSearchResult{
|
||||||
|
{Score: 0.95, Slug: &slug, DisplayName: &name, Summary: &summary, Version: &version},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
reg := newTestRegistry(srv.URL, "")
|
||||||
|
results, err := reg.Search(context.Background(), "github", 5)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 1)
|
||||||
|
assert.Equal(t, "github", results[0].Slug)
|
||||||
|
assert.Equal(t, "GitHub Integration", results[0].DisplayName)
|
||||||
|
assert.InDelta(t, 0.95, results[0].Score, 0.001)
|
||||||
|
assert.Equal(t, "clawhub", results[0].RegistryName)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistryGetSkillMeta(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
assert.Equal(t, "/api/v1/skills/github", r.URL.Path)
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(clawhubSkillResponse{
|
||||||
|
Slug: "github",
|
||||||
|
DisplayName: "GitHub Integration",
|
||||||
|
Summary: "Full GitHub API integration",
|
||||||
|
LatestVersion: &clawhubVersionInfo{
|
||||||
|
Version: "2.1.0",
|
||||||
|
},
|
||||||
|
Moderation: &clawhubModerationInfo{
|
||||||
|
IsMalwareBlocked: false,
|
||||||
|
IsSuspicious: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
reg := newTestRegistry(srv.URL, "")
|
||||||
|
meta, err := reg.GetSkillMeta(context.Background(), "github")
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "github", meta.Slug)
|
||||||
|
assert.Equal(t, "2.1.0", meta.LatestVersion)
|
||||||
|
assert.False(t, meta.IsMalwareBlocked)
|
||||||
|
assert.True(t, meta.IsSuspicious)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistryGetSkillMetaUnsafeSlug(t *testing.T) {
|
||||||
|
reg := newTestRegistry("https://example.com", "")
|
||||||
|
_, err := reg.GetSkillMeta(context.Background(), "../etc/passwd")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "invalid slug")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistryDownloadAndInstall(t *testing.T) {
|
||||||
|
// Create a valid ZIP in memory.
|
||||||
|
zipBuf := createTestZip(t, map[string]string{
|
||||||
|
"SKILL.md": "---\nname: test-skill\ndescription: A test\n---\nHello skill",
|
||||||
|
"README.md": "# Test Skill\n",
|
||||||
|
})
|
||||||
|
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch r.URL.Path {
|
||||||
|
case "/api/v1/skills/test-skill":
|
||||||
|
// Metadata endpoint.
|
||||||
|
json.NewEncoder(w).Encode(clawhubSkillResponse{
|
||||||
|
Slug: "test-skill",
|
||||||
|
DisplayName: "Test Skill",
|
||||||
|
Summary: "A test skill",
|
||||||
|
LatestVersion: &clawhubVersionInfo{Version: "1.0.0"},
|
||||||
|
})
|
||||||
|
case "/api/v1/download":
|
||||||
|
assert.Equal(t, "test-skill", r.URL.Query().Get("slug"))
|
||||||
|
w.Header().Set("Content-Type", "application/zip")
|
||||||
|
w.Write(zipBuf)
|
||||||
|
default:
|
||||||
|
w.WriteHeader(http.StatusNotFound)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
targetDir := filepath.Join(tmpDir, "test-skill")
|
||||||
|
|
||||||
|
reg := newTestRegistry(srv.URL, "")
|
||||||
|
result, err := reg.DownloadAndInstall(context.Background(), "test-skill", "1.0.0", targetDir)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "1.0.0", result.Version)
|
||||||
|
assert.False(t, result.IsMalwareBlocked)
|
||||||
|
|
||||||
|
// Verify extracted files.
|
||||||
|
skillContent, err := os.ReadFile(filepath.Join(targetDir, "SKILL.md"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, string(skillContent), "Hello skill")
|
||||||
|
|
||||||
|
readmeContent, err := os.ReadFile(filepath.Join(targetDir, "README.md"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, string(readmeContent), "# Test Skill")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistryAuthToken(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
authHeader := r.Header.Get("Authorization")
|
||||||
|
assert.Equal(t, "Bearer test-token-123", authHeader)
|
||||||
|
json.NewEncoder(w).Encode(clawhubSearchResponse{Results: nil})
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
reg := newTestRegistry(srv.URL, "test-token-123")
|
||||||
|
_, _ = reg.Search(context.Background(), "test", 5)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractZipPathTraversal(t *testing.T) {
|
||||||
|
// Create a ZIP with a path traversal entry.
|
||||||
|
var buf bytes.Buffer
|
||||||
|
zw := zip.NewWriter(&buf)
|
||||||
|
|
||||||
|
// Malicious entry trying to escape directory.
|
||||||
|
w, err := zw.Create("../../etc/passwd")
|
||||||
|
require.NoError(t, err)
|
||||||
|
w.Write([]byte("malicious"))
|
||||||
|
|
||||||
|
zw.Close()
|
||||||
|
|
||||||
|
// Write to temp file for extractZipFile.
|
||||||
|
tmpZip := filepath.Join(t.TempDir(), "bad.zip")
|
||||||
|
require.NoError(t, os.WriteFile(tmpZip, buf.Bytes(), 0644))
|
||||||
|
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
err = utils.ExtractZipFile(tmpZip, tmpDir)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "unsafe path")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractZipWithSubdirectories(t *testing.T) {
|
||||||
|
zipBuf := createTestZip(t, map[string]string{
|
||||||
|
"SKILL.md": "root file",
|
||||||
|
"scripts/helper.sh": "#!/bin/bash\necho hello",
|
||||||
|
"examples/demo.yaml": "key: value",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Write to temp file for extractZipFile.
|
||||||
|
tmpZip := filepath.Join(t.TempDir(), "test.zip")
|
||||||
|
require.NoError(t, os.WriteFile(tmpZip, zipBuf, 0644))
|
||||||
|
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
targetDir := filepath.Join(tmpDir, "my-skill")
|
||||||
|
|
||||||
|
err := utils.ExtractZipFile(tmpZip, targetDir)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify nested file.
|
||||||
|
data, err := os.ReadFile(filepath.Join(targetDir, "scripts", "helper.sh"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, string(data), "#!/bin/bash")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistrySearchHTTPError(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
w.Write([]byte("Internal Server Error"))
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
reg := newTestRegistry(srv.URL, "")
|
||||||
|
_, err := reg.Search(context.Background(), "test", 5)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "500")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClawHubRegistrySearchNullableFields(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
validSlug := "valid-slug"
|
||||||
|
validSummary := "valid summary"
|
||||||
|
|
||||||
|
// Return results with various null/empty fields
|
||||||
|
json.NewEncoder(w).Encode(clawhubSearchResponse{
|
||||||
|
Results: []clawhubSearchResult{
|
||||||
|
// Case 1: Null Slug -> Skip
|
||||||
|
{Score: 0.1, Slug: nil, DisplayName: nil, Summary: nil, Version: nil},
|
||||||
|
// Case 2: Valid Slug, Null Summary -> Skip
|
||||||
|
{Score: 0.2, Slug: &validSlug, DisplayName: nil, Summary: nil, Version: nil},
|
||||||
|
// Case 3: Valid Slug, Valid Summary, Null Name -> Keep, Name=Slug
|
||||||
|
{Score: 0.8, Slug: &validSlug, DisplayName: nil, Summary: &validSummary, Version: nil},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
reg := newTestRegistry(srv.URL, "")
|
||||||
|
results, err := reg.Search(context.Background(), "test", 5)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 1, "should only return 1 valid result")
|
||||||
|
|
||||||
|
r := results[0]
|
||||||
|
assert.Equal(t, "valid-slug", r.Slug)
|
||||||
|
assert.Equal(t, "valid-slug", r.DisplayName, "should fallback name to slug")
|
||||||
|
assert.Equal(t, "valid summary", r.Summary)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- helpers ---
|
||||||
|
|
||||||
|
func createTestZip(t *testing.T, files map[string]string) []byte {
|
||||||
|
t.Helper()
|
||||||
|
var buf bytes.Buffer
|
||||||
|
zw := zip.NewWriter(&buf)
|
||||||
|
|
||||||
|
for name, content := range files {
|
||||||
|
w, err := zw.Create(name)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, err = w.Write([]byte(content))
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, zw.Close())
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
@ -8,7 +8,6 @@ import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -24,12 +23,6 @@ type AvailableSkill struct {
|
||||||
Tags []string `json:"tags"`
|
Tags []string `json:"tags"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type BuiltinSkill struct {
|
|
||||||
Name string `json:"name"`
|
|
||||||
Path string `json:"path"`
|
|
||||||
Enabled bool `json:"enabled"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewSkillInstaller(workspace string) *SkillInstaller {
|
func NewSkillInstaller(workspace string) *SkillInstaller {
|
||||||
return &SkillInstaller{
|
return &SkillInstaller{
|
||||||
workspace: workspace,
|
workspace: workspace,
|
||||||
|
|
@ -123,49 +116,3 @@ func (si *SkillInstaller) ListAvailableSkills(ctx context.Context) ([]AvailableS
|
||||||
|
|
||||||
return skills, nil
|
return skills, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (si *SkillInstaller) ListBuiltinSkills() []BuiltinSkill {
|
|
||||||
builtinSkillsDir := filepath.Join(filepath.Dir(si.workspace), "picoclaw", "skills")
|
|
||||||
|
|
||||||
entries, err := os.ReadDir(builtinSkillsDir)
|
|
||||||
if err != nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var skills []BuiltinSkill
|
|
||||||
for _, entry := range entries {
|
|
||||||
if entry.IsDir() {
|
|
||||||
_ = entry
|
|
||||||
skillName := entry.Name()
|
|
||||||
skillFile := filepath.Join(builtinSkillsDir, skillName, "SKILL.md")
|
|
||||||
|
|
||||||
data, err := os.ReadFile(skillFile)
|
|
||||||
description := ""
|
|
||||||
if err == nil {
|
|
||||||
content := string(data)
|
|
||||||
if idx := strings.Index(content, "\n"); idx > 0 {
|
|
||||||
firstLine := content[:idx]
|
|
||||||
if strings.Contains(firstLine, "description:") {
|
|
||||||
descLine := strings.Index(content[idx:], "\n")
|
|
||||||
if descLine > 0 {
|
|
||||||
description = strings.TrimSpace(content[idx+descLine : idx+descLine])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// skill := BuiltinSkill{
|
|
||||||
// Name: skillName,
|
|
||||||
// Path: description,
|
|
||||||
// Enabled: true,
|
|
||||||
// }
|
|
||||||
|
|
||||||
status := "✓"
|
|
||||||
fmt.Printf(" %s %s\n", status, entry.Name())
|
|
||||||
if description != "" {
|
|
||||||
fmt.Printf(" %s\n", description)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return skills
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,8 @@ import (
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
var namePattern = regexp.MustCompile(`^[a-zA-Z0-9]+(-[a-zA-Z0-9]+)*$`)
|
var namePattern = regexp.MustCompile(`^[a-zA-Z0-9]+(-[a-zA-Z0-9]+)*$`)
|
||||||
|
|
@ -251,6 +253,11 @@ func (sl *SkillsLoader) BuildSkillsSummary() string {
|
||||||
func (sl *SkillsLoader) getSkillMetadata(skillPath string) *SkillMetadata {
|
func (sl *SkillsLoader) getSkillMetadata(skillPath string) *SkillMetadata {
|
||||||
content, err := os.ReadFile(skillPath)
|
content, err := os.ReadFile(skillPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
logger.WarnCF("skills", "Failed to read skill metadata",
|
||||||
|
map[string]interface{}{
|
||||||
|
"skill_path": skillPath,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -283,10 +290,15 @@ func (sl *SkillsLoader) getSkillMetadata(skillPath string) *SkillMetadata {
|
||||||
|
|
||||||
// parseSimpleYAML parses simple key: value YAML format
|
// parseSimpleYAML parses simple key: value YAML format
|
||||||
// Example: name: github\n description: "..."
|
// Example: name: github\n description: "..."
|
||||||
|
// Normalizes line endings to handle \n (Unix), \r\n (Windows), and \r (classic Mac)
|
||||||
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)
|
||||||
|
|
||||||
for _, line := range strings.Split(content, "\n") {
|
// Normalize line endings: convert \r\n and \r to \n
|
||||||
|
normalized := strings.ReplaceAll(content, "\r\n", "\n")
|
||||||
|
normalized = strings.ReplaceAll(normalized, "\r", "\n")
|
||||||
|
|
||||||
|
for _, line := range strings.Split(normalized, "\n") {
|
||||||
line = strings.TrimSpace(line)
|
line = strings.TrimSpace(line)
|
||||||
if line == "" || strings.HasPrefix(line, "#") {
|
if line == "" || strings.HasPrefix(line, "#") {
|
||||||
continue
|
continue
|
||||||
|
|
@ -306,9 +318,10 @@ func (sl *SkillsLoader) parseSimpleYAML(content string) map[string]string {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sl *SkillsLoader) extractFrontmatter(content string) string {
|
func (sl *SkillsLoader) extractFrontmatter(content string) string {
|
||||||
// (?s) enables DOTALL mode so . matches newlines
|
// Support \n (Unix), \r\n (Windows), and \r (classic Mac) line endings for frontmatter blocks
|
||||||
// Match first ---, capture everything until next --- on its own line
|
// (?s) enables DOTALL so . matches newlines;
|
||||||
re := regexp.MustCompile(`(?s)^---\n(.*)\n---`)
|
// ^--- at start, then ... --- at start of line, honoring all three line ending types
|
||||||
|
re := regexp.MustCompile(`(?s)^---(?:\r\n|\n|\r)(.*?)(?:\r\n|\n|\r)---`)
|
||||||
match := re.FindStringSubmatch(content)
|
match := re.FindStringSubmatch(content)
|
||||||
if len(match) > 1 {
|
if len(match) > 1 {
|
||||||
return match[1]
|
return match[1]
|
||||||
|
|
@ -317,7 +330,11 @@ func (sl *SkillsLoader) extractFrontmatter(content string) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sl *SkillsLoader) stripFrontmatter(content string) string {
|
func (sl *SkillsLoader) stripFrontmatter(content string) string {
|
||||||
re := regexp.MustCompile(`^---\n.*?\n---\n`)
|
// Support \n (Unix), \r\n (Windows), and \r (classic Mac) line endings for frontmatter blocks
|
||||||
|
// (?s) enables DOTALL so . matches newlines;
|
||||||
|
// ^--- at start, then ... --- at start of line, honoring all three line ending types
|
||||||
|
// Match zero or more trailing line endings after closing --- (handles both with and without blank lines)
|
||||||
|
re := regexp.MustCompile(`(?s)^---(?:\r\n|\n|\r)(.*?)(?:\r\n|\n|\r)---(?:\r\n|\n|\r)*`)
|
||||||
return re.ReplaceAllString(content, "")
|
return re.ReplaceAllString(content, "")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -75,3 +75,105 @@ func TestSkillsInfoValidate(t *testing.T) {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestExtractFrontmatter(t *testing.T) {
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
|
||||||
|
testcases := []struct {
|
||||||
|
name string
|
||||||
|
content string
|
||||||
|
expectedName string
|
||||||
|
expectedDesc string
|
||||||
|
lineEndingType string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "unix-line-endings",
|
||||||
|
lineEndingType: "Unix (\\n)",
|
||||||
|
content: "---\nname: test-skill\ndescription: A test skill\n---\n\n# Skill Content",
|
||||||
|
expectedName: "test-skill",
|
||||||
|
expectedDesc: "A test skill",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "windows-line-endings",
|
||||||
|
lineEndingType: "Windows (\\r\\n)",
|
||||||
|
content: "---\r\nname: test-skill\r\ndescription: A test skill\r\n---\r\n\r\n# Skill Content",
|
||||||
|
expectedName: "test-skill",
|
||||||
|
expectedDesc: "A test skill",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "classic-mac-line-endings",
|
||||||
|
lineEndingType: "Classic Mac (\\r)",
|
||||||
|
content: "---\rname: test-skill\rdescription: A test skill\r---\r\r# Skill Content",
|
||||||
|
expectedName: "test-skill",
|
||||||
|
expectedDesc: "A test skill",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range testcases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
// Extract frontmatter
|
||||||
|
frontmatter := sl.extractFrontmatter(tc.content)
|
||||||
|
assert.NotEmpty(t, frontmatter, "Frontmatter should be extracted for %s line endings", tc.lineEndingType)
|
||||||
|
|
||||||
|
// Parse YAML to get name and description (parseSimpleYAML now handles all line ending types)
|
||||||
|
yamlMeta := sl.parseSimpleYAML(frontmatter)
|
||||||
|
assert.Equal(t, tc.expectedName, yamlMeta["name"], "Name should be correctly parsed from frontmatter with %s line endings", tc.lineEndingType)
|
||||||
|
assert.Equal(t, tc.expectedDesc, yamlMeta["description"], "Description should be correctly parsed from frontmatter with %s line endings", tc.lineEndingType)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStripFrontmatter(t *testing.T) {
|
||||||
|
sl := &SkillsLoader{}
|
||||||
|
|
||||||
|
testcases := []struct {
|
||||||
|
name string
|
||||||
|
content string
|
||||||
|
expectedContent string
|
||||||
|
lineEndingType string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "unix-line-endings",
|
||||||
|
lineEndingType: "Unix (\\n)",
|
||||||
|
content: "---\nname: test-skill\ndescription: A test skill\n---\n\n# Skill Content",
|
||||||
|
expectedContent: "# Skill Content",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "windows-line-endings",
|
||||||
|
lineEndingType: "Windows (\\r\\n)",
|
||||||
|
content: "---\r\nname: test-skill\r\ndescription: A test skill\r\n---\r\n\r\n# Skill Content",
|
||||||
|
expectedContent: "# Skill Content",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "classic-mac-line-endings",
|
||||||
|
lineEndingType: "Classic Mac (\\r)",
|
||||||
|
content: "---\rname: test-skill\rdescription: A test skill\r---\r\r# Skill Content",
|
||||||
|
expectedContent: "# Skill Content",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unix-line-endings-without-trailing-newline",
|
||||||
|
lineEndingType: "Unix (\\n) without trailing newline",
|
||||||
|
content: "---\nname: test-skill\ndescription: A test skill\n---\n# Skill Content",
|
||||||
|
expectedContent: "# Skill Content",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "windows-line-endings-without-trailing-newline",
|
||||||
|
lineEndingType: "Windows (\\r\\n) without trailing newline",
|
||||||
|
content: "---\r\nname: test-skill\r\ndescription: A test skill\r\n---\r\n# Skill Content",
|
||||||
|
expectedContent: "# Skill Content",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no-frontmatter",
|
||||||
|
lineEndingType: "No frontmatter",
|
||||||
|
content: "# Skill Content\n\nSome content here.",
|
||||||
|
expectedContent: "# Skill Content\n\nSome content here.",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range testcases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
result := sl.stripFrontmatter(tc.content)
|
||||||
|
assert.Equal(t, tc.expectedContent, result, "Frontmatter should be stripped correctly for %s", tc.lineEndingType)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
223
pkg/skills/registry.go
Normal file
223
pkg/skills/registry.go
Normal file
|
|
@ -0,0 +1,223 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultMaxConcurrentSearches = 2
|
||||||
|
)
|
||||||
|
|
||||||
|
// SearchResult represents a single result from a skill registry search.
|
||||||
|
type SearchResult struct {
|
||||||
|
Score float64 `json:"score"`
|
||||||
|
Slug string `json:"slug"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
Version string `json:"version"`
|
||||||
|
RegistryName string `json:"registry_name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SkillMeta holds metadata about a skill from a registry.
|
||||||
|
type SkillMeta struct {
|
||||||
|
Slug string `json:"slug"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
LatestVersion string `json:"latest_version"`
|
||||||
|
IsMalwareBlocked bool `json:"is_malware_blocked"`
|
||||||
|
IsSuspicious bool `json:"is_suspicious"`
|
||||||
|
RegistryName string `json:"registry_name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// InstallResult is returned by DownloadAndInstall to carry metadata
|
||||||
|
// back to the caller for moderation and user messaging.
|
||||||
|
type InstallResult struct {
|
||||||
|
Version string
|
||||||
|
IsMalwareBlocked bool
|
||||||
|
IsSuspicious bool
|
||||||
|
Summary string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SkillRegistry is the interface that all skill registries must implement.
|
||||||
|
// Each registry represents a different source of skills (e.g., clawhub.ai)
|
||||||
|
type SkillRegistry interface {
|
||||||
|
// Name returns the unique name of this registry (e.g., "clawhub").
|
||||||
|
Name() string
|
||||||
|
// Search searches the registry for skills matching the query.
|
||||||
|
Search(ctx context.Context, query string, limit int) ([]SearchResult, error)
|
||||||
|
// GetSkillMeta retrieves metadata for a specific skill by slug.
|
||||||
|
GetSkillMeta(ctx context.Context, slug string) (*SkillMeta, error)
|
||||||
|
// DownloadAndInstall fetches metadata, resolves the version, downloads and
|
||||||
|
// installs the skill to targetDir. Returns an InstallResult with metadata
|
||||||
|
// for the caller to use for moderation and user messaging.
|
||||||
|
DownloadAndInstall(ctx context.Context, slug, version, targetDir string) (*InstallResult, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegistryConfig holds configuration for all skill registries.
|
||||||
|
// This is the input to NewRegistryManagerFromConfig.
|
||||||
|
type RegistryConfig struct {
|
||||||
|
ClawHub ClawHubConfig
|
||||||
|
MaxConcurrentSearches int
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClawHubConfig configures the ClawHub registry.
|
||||||
|
type ClawHubConfig struct {
|
||||||
|
Enabled bool
|
||||||
|
BaseURL string
|
||||||
|
AuthToken string
|
||||||
|
SearchPath string // e.g. "/api/v1/search"
|
||||||
|
SkillsPath string // e.g. "/api/v1/skills"
|
||||||
|
DownloadPath string // e.g. "/api/v1/download"
|
||||||
|
Timeout int // seconds, 0 = default (30s)
|
||||||
|
MaxZipSize int // bytes, 0 = default (50MB)
|
||||||
|
MaxResponseSize int // bytes, 0 = default (2MB)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegistryManager coordinates multiple skill registries.
|
||||||
|
// It fans out search requests and routes installs to the correct registry.
|
||||||
|
type RegistryManager struct {
|
||||||
|
registries []SkillRegistry
|
||||||
|
maxConcurrent int
|
||||||
|
mu sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRegistryManager creates an empty RegistryManager.
|
||||||
|
func NewRegistryManager() *RegistryManager {
|
||||||
|
return &RegistryManager{
|
||||||
|
registries: make([]SkillRegistry, 0),
|
||||||
|
maxConcurrent: defaultMaxConcurrentSearches,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRegistryManagerFromConfig builds a RegistryManager from config,
|
||||||
|
// instantiating only the enabled registries.
|
||||||
|
func NewRegistryManagerFromConfig(cfg RegistryConfig) *RegistryManager {
|
||||||
|
rm := NewRegistryManager()
|
||||||
|
if cfg.MaxConcurrentSearches > 0 {
|
||||||
|
rm.maxConcurrent = cfg.MaxConcurrentSearches
|
||||||
|
}
|
||||||
|
if cfg.ClawHub.Enabled {
|
||||||
|
rm.AddRegistry(NewClawHubRegistry(cfg.ClawHub))
|
||||||
|
}
|
||||||
|
return rm
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddRegistry adds a registry to the manager.
|
||||||
|
func (rm *RegistryManager) AddRegistry(r SkillRegistry) {
|
||||||
|
rm.mu.Lock()
|
||||||
|
defer rm.mu.Unlock()
|
||||||
|
rm.registries = append(rm.registries, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRegistry returns a registry by name, or nil if not found.
|
||||||
|
func (rm *RegistryManager) GetRegistry(name string) SkillRegistry {
|
||||||
|
rm.mu.RLock()
|
||||||
|
defer rm.mu.RUnlock()
|
||||||
|
for _, r := range rm.registries {
|
||||||
|
if r.Name() == name {
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SearchAll fans out the query to all registries concurrently
|
||||||
|
// and merges results sorted by score descending.
|
||||||
|
func (rm *RegistryManager) SearchAll(ctx context.Context, query string, limit int) ([]SearchResult, error) {
|
||||||
|
rm.mu.RLock()
|
||||||
|
regs := make([]SkillRegistry, len(rm.registries))
|
||||||
|
copy(regs, rm.registries)
|
||||||
|
rm.mu.RUnlock()
|
||||||
|
|
||||||
|
if len(regs) == 0 {
|
||||||
|
return nil, fmt.Errorf("no registries configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
type regResult struct {
|
||||||
|
results []SearchResult
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Semaphore: limit concurrency.
|
||||||
|
sem := make(chan struct{}, rm.maxConcurrent)
|
||||||
|
resultsCh := make(chan regResult, len(regs))
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for _, reg := range regs {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(r SkillRegistry) {
|
||||||
|
defer wg.Done()
|
||||||
|
|
||||||
|
// Acquire semaphore slot.
|
||||||
|
select {
|
||||||
|
case sem <- struct{}{}:
|
||||||
|
defer func() { <-sem }()
|
||||||
|
case <-ctx.Done():
|
||||||
|
resultsCh <- regResult{err: ctx.Err()}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
searchCtx, cancel := context.WithTimeout(ctx, 1*time.Minute)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
results, err := r.Search(searchCtx, query, limit)
|
||||||
|
if err != nil {
|
||||||
|
slog.Warn("registry search failed", "registry", r.Name(), "error", err)
|
||||||
|
resultsCh <- regResult{err: err}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resultsCh <- regResult{results: results}
|
||||||
|
}(reg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close results channel after all goroutines complete.
|
||||||
|
go func() {
|
||||||
|
wg.Wait()
|
||||||
|
close(resultsCh)
|
||||||
|
}()
|
||||||
|
|
||||||
|
var merged []SearchResult
|
||||||
|
var lastErr error
|
||||||
|
|
||||||
|
var anyRegistrySucceeded bool
|
||||||
|
for rr := range resultsCh {
|
||||||
|
if rr.err != nil {
|
||||||
|
lastErr = rr.err
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
anyRegistrySucceeded = true
|
||||||
|
merged = append(merged, rr.results...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If all registries failed, return the last error.
|
||||||
|
if !anyRegistrySucceeded && lastErr != nil {
|
||||||
|
return nil, fmt.Errorf("all registries failed: %w", lastErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sort by score descending.
|
||||||
|
sortByScoreDesc(merged)
|
||||||
|
|
||||||
|
// Clamp to limit.
|
||||||
|
if limit > 0 && len(merged) > limit {
|
||||||
|
merged = merged[:limit]
|
||||||
|
}
|
||||||
|
|
||||||
|
return merged, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sortByScoreDesc sorts SearchResults by Score in descending order (insertion sort — small slices).
|
||||||
|
func sortByScoreDesc(results []SearchResult) {
|
||||||
|
for i := 1; i < len(results); i++ {
|
||||||
|
key := results[i]
|
||||||
|
j := i - 1
|
||||||
|
for j >= 0 && results[j].Score < key.Score {
|
||||||
|
results[j+1] = results[j]
|
||||||
|
j--
|
||||||
|
}
|
||||||
|
results[j+1] = key
|
||||||
|
}
|
||||||
|
}
|
||||||
179
pkg/skills/registry_test.go
Normal file
179
pkg/skills/registry_test.go
Normal file
|
|
@ -0,0 +1,179 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
// mockRegistry is a test double implementing SkillRegistry.
|
||||||
|
type mockRegistry struct {
|
||||||
|
name string
|
||||||
|
searchResults []SearchResult
|
||||||
|
searchErr error
|
||||||
|
meta *SkillMeta
|
||||||
|
metaErr error
|
||||||
|
installResult *InstallResult
|
||||||
|
installErr error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockRegistry) Name() string { return m.name }
|
||||||
|
|
||||||
|
func (m *mockRegistry) Search(_ context.Context, _ string, _ int) ([]SearchResult, error) {
|
||||||
|
return m.searchResults, m.searchErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockRegistry) GetSkillMeta(_ context.Context, _ string) (*SkillMeta, error) {
|
||||||
|
return m.meta, m.metaErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockRegistry) DownloadAndInstall(_ context.Context, _, _, _ string) (*InstallResult, error) {
|
||||||
|
return m.installResult, m.installErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllSingle(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "test",
|
||||||
|
searchResults: []SearchResult{
|
||||||
|
{Slug: "skill-a", Score: 0.9, RegistryName: "test"},
|
||||||
|
{Slug: "skill-b", Score: 0.5, RegistryName: "test"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
results, err := mgr.SearchAll(context.Background(), "test query", 10)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, results, 2)
|
||||||
|
assert.Equal(t, "skill-a", results[0].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllMultiple(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "alpha",
|
||||||
|
searchResults: []SearchResult{
|
||||||
|
{Slug: "skill-a", Score: 0.8, RegistryName: "alpha"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "beta",
|
||||||
|
searchResults: []SearchResult{
|
||||||
|
{Slug: "skill-b", Score: 0.95, RegistryName: "beta"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
results, err := mgr.SearchAll(context.Background(), "test query", 10)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, results, 2)
|
||||||
|
// Should be sorted by score descending
|
||||||
|
assert.Equal(t, "skill-b", results[0].Slug)
|
||||||
|
assert.Equal(t, "skill-a", results[1].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllOneFailsGracefully(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "failing",
|
||||||
|
searchErr: fmt.Errorf("network error"),
|
||||||
|
})
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "working",
|
||||||
|
searchResults: []SearchResult{
|
||||||
|
{Slug: "skill-a", Score: 0.8, RegistryName: "working"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
results, err := mgr.SearchAll(context.Background(), "test query", 10)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, results, 1)
|
||||||
|
assert.Equal(t, "skill-a", results[0].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllAllFail(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "fail-1",
|
||||||
|
searchErr: fmt.Errorf("error 1"),
|
||||||
|
})
|
||||||
|
|
||||||
|
_, err := mgr.SearchAll(context.Background(), "test query", 10)
|
||||||
|
assert.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllNoRegistries(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
_, err := mgr.SearchAll(context.Background(), "test query", 10)
|
||||||
|
assert.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerGetRegistry(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
mock := &mockRegistry{name: "clawhub"}
|
||||||
|
mgr.AddRegistry(mock)
|
||||||
|
|
||||||
|
got := mgr.GetRegistry("clawhub")
|
||||||
|
assert.NotNil(t, got)
|
||||||
|
assert.Equal(t, "clawhub", got.Name())
|
||||||
|
|
||||||
|
got = mgr.GetRegistry("nonexistent")
|
||||||
|
assert.Nil(t, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllRespectLimit(t *testing.T) {
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
results := make([]SearchResult, 20)
|
||||||
|
for i := range results {
|
||||||
|
results[i] = SearchResult{Slug: fmt.Sprintf("skill-%d", i), Score: float64(20 - i)}
|
||||||
|
}
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "test",
|
||||||
|
searchResults: results,
|
||||||
|
})
|
||||||
|
|
||||||
|
got, err := mgr.SearchAll(context.Background(), "test", 5)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, got, 5)
|
||||||
|
// Top scores first
|
||||||
|
assert.Equal(t, "skill-0", got[0].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryManagerSearchAllTimeout(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
time.Sleep(5 * time.Millisecond) // Let context expire.
|
||||||
|
|
||||||
|
mgr := NewRegistryManager()
|
||||||
|
mgr.AddRegistry(&mockRegistry{
|
||||||
|
name: "slow",
|
||||||
|
searchErr: fmt.Errorf("context deadline exceeded"),
|
||||||
|
})
|
||||||
|
|
||||||
|
_, err := mgr.SearchAll(ctx, "test", 5)
|
||||||
|
assert.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortByScoreDesc(t *testing.T) {
|
||||||
|
results := []SearchResult{
|
||||||
|
{Slug: "c", Score: 0.3},
|
||||||
|
{Slug: "a", Score: 0.9},
|
||||||
|
{Slug: "b", Score: 0.5},
|
||||||
|
}
|
||||||
|
sortByScoreDesc(results)
|
||||||
|
assert.Equal(t, "a", results[0].Slug)
|
||||||
|
assert.Equal(t, "b", results[1].Slug)
|
||||||
|
assert.Equal(t, "c", results[2].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsSafeSlug(t *testing.T) {
|
||||||
|
assert.NoError(t, utils.ValidateSkillIdentifier("github"))
|
||||||
|
assert.NoError(t, utils.ValidateSkillIdentifier("docker-compose"))
|
||||||
|
assert.Error(t, utils.ValidateSkillIdentifier(""))
|
||||||
|
assert.Error(t, utils.ValidateSkillIdentifier("../etc/passwd"))
|
||||||
|
assert.Error(t, utils.ValidateSkillIdentifier("path/traversal"))
|
||||||
|
assert.Error(t, utils.ValidateSkillIdentifier("path\\traversal"))
|
||||||
|
}
|
||||||
229
pkg/skills/search_cache.go
Normal file
229
pkg/skills/search_cache.go
Normal file
|
|
@ -0,0 +1,229 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SearchCache provides lightweight caching for search results.
|
||||||
|
// It uses trigram-based similarity to match similar queries to cached results,
|
||||||
|
// avoiding redundant API calls. Thread-safe for concurrent access.
|
||||||
|
type SearchCache struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
entries map[string]*cacheEntry
|
||||||
|
order []string // LRU order: oldest first.
|
||||||
|
maxEntries int
|
||||||
|
ttl time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
type cacheEntry struct {
|
||||||
|
query string
|
||||||
|
trigrams []uint32
|
||||||
|
results []SearchResult
|
||||||
|
createdAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// similarityThreshold is the minimum trigram Jaccard similarity for a cache hit.
|
||||||
|
const similarityThreshold = 0.7
|
||||||
|
|
||||||
|
// NewSearchCache creates a new search cache.
|
||||||
|
// maxEntries is the maximum number of cached queries (excess evicts LRU).
|
||||||
|
// ttl is how long each entry lives before expiration.
|
||||||
|
func NewSearchCache(maxEntries int, ttl time.Duration) *SearchCache {
|
||||||
|
if maxEntries <= 0 {
|
||||||
|
maxEntries = 50
|
||||||
|
}
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = 5 * time.Minute
|
||||||
|
}
|
||||||
|
return &SearchCache{
|
||||||
|
entries: make(map[string]*cacheEntry),
|
||||||
|
order: make([]string, 0),
|
||||||
|
maxEntries: maxEntries,
|
||||||
|
ttl: ttl,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get looks up results for a query. Returns cached results and true if found
|
||||||
|
// (either exact or similar match above threshold). Returns nil, false on miss.
|
||||||
|
func (sc *SearchCache) Get(query string) ([]SearchResult, bool) {
|
||||||
|
normalized := normalizeQuery(query)
|
||||||
|
if normalized == "" {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
sc.mu.Lock()
|
||||||
|
defer sc.mu.Unlock()
|
||||||
|
|
||||||
|
// Exact match first.
|
||||||
|
if entry, ok := sc.entries[normalized]; ok {
|
||||||
|
if time.Since(entry.createdAt) < sc.ttl {
|
||||||
|
sc.moveToEndLocked(normalized)
|
||||||
|
return copyResults(entry.results), true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Similarity match.
|
||||||
|
queryTrigrams := buildTrigrams(normalized)
|
||||||
|
var bestEntry *cacheEntry
|
||||||
|
var bestSim float64
|
||||||
|
|
||||||
|
for _, entry := range sc.entries {
|
||||||
|
if time.Since(entry.createdAt) >= sc.ttl {
|
||||||
|
continue // Skip expired.
|
||||||
|
}
|
||||||
|
sim := jaccardSimilarity(queryTrigrams, entry.trigrams)
|
||||||
|
if sim > bestSim {
|
||||||
|
bestSim = sim
|
||||||
|
bestEntry = entry
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if bestSim >= similarityThreshold && bestEntry != nil {
|
||||||
|
sc.moveToEndLocked(bestEntry.query)
|
||||||
|
return copyResults(bestEntry.results), true
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Put stores results for a query. Evicts the oldest entry if at capacity.
|
||||||
|
func (sc *SearchCache) Put(query string, results []SearchResult) {
|
||||||
|
normalized := normalizeQuery(query)
|
||||||
|
if normalized == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sc.mu.Lock()
|
||||||
|
defer sc.mu.Unlock()
|
||||||
|
|
||||||
|
// Evict expired entries first.
|
||||||
|
sc.evictExpiredLocked()
|
||||||
|
|
||||||
|
// If already exists, update.
|
||||||
|
if _, ok := sc.entries[normalized]; ok {
|
||||||
|
sc.entries[normalized] = &cacheEntry{
|
||||||
|
query: normalized,
|
||||||
|
trigrams: buildTrigrams(normalized),
|
||||||
|
results: copyResults(results),
|
||||||
|
createdAt: time.Now(),
|
||||||
|
}
|
||||||
|
// Move to end of LRU order.
|
||||||
|
sc.moveToEndLocked(normalized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Evict LRU if at capacity.
|
||||||
|
for len(sc.entries) >= sc.maxEntries && len(sc.order) > 0 {
|
||||||
|
oldest := sc.order[0]
|
||||||
|
sc.order = sc.order[1:]
|
||||||
|
delete(sc.entries, oldest)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Insert new entry.
|
||||||
|
sc.entries[normalized] = &cacheEntry{
|
||||||
|
query: normalized,
|
||||||
|
trigrams: buildTrigrams(normalized),
|
||||||
|
results: copyResults(results),
|
||||||
|
createdAt: time.Now(),
|
||||||
|
}
|
||||||
|
sc.order = append(sc.order, normalized)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Len returns the number of entries (for testing).
|
||||||
|
func (sc *SearchCache) Len() int {
|
||||||
|
sc.mu.RLock()
|
||||||
|
defer sc.mu.RUnlock()
|
||||||
|
return len(sc.entries)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- internal ---
|
||||||
|
|
||||||
|
func (sc *SearchCache) evictExpiredLocked() {
|
||||||
|
now := time.Now()
|
||||||
|
newOrder := make([]string, 0, len(sc.order))
|
||||||
|
for _, key := range sc.order {
|
||||||
|
entry, ok := sc.entries[key]
|
||||||
|
if !ok || now.Sub(entry.createdAt) >= sc.ttl {
|
||||||
|
delete(sc.entries, key)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
newOrder = append(newOrder, key)
|
||||||
|
}
|
||||||
|
sc.order = newOrder
|
||||||
|
}
|
||||||
|
|
||||||
|
func (sc *SearchCache) moveToEndLocked(key string) {
|
||||||
|
for i, k := range sc.order {
|
||||||
|
if k == key {
|
||||||
|
sc.order = append(sc.order[:i], sc.order[i+1:]...)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sc.order = append(sc.order, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeQuery(q string) string {
|
||||||
|
return strings.ToLower(strings.TrimSpace(q))
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTrigrams generates hash of trigrams from a string.
|
||||||
|
// Example: "hello" → {"hel", "ell", "llo"}
|
||||||
|
// "hel" -> 0x0068656c -> 4 bytes; compared to 16 bytes of a string
|
||||||
|
func buildTrigrams(s string) []uint32 {
|
||||||
|
if len(s) < 3 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
trigrams := make([]uint32, 0, len(s)-2)
|
||||||
|
for i := 0; i <= len(s)-3; i++ {
|
||||||
|
trigrams = append(trigrams, uint32(s[i])<<16|uint32(s[i+1])<<8|uint32(s[i+2]))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sort and Deduplication
|
||||||
|
sort.Slice(trigrams, func(i, j int) bool { return trigrams[i] < trigrams[j] })
|
||||||
|
n := 1
|
||||||
|
for i := 1; i < len(trigrams); i++ {
|
||||||
|
if trigrams[i] != trigrams[i-1] {
|
||||||
|
trigrams[n] = trigrams[i]
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return trigrams[:n]
|
||||||
|
}
|
||||||
|
|
||||||
|
// jaccardSimilarity computes |A ∩ B| / |A ∪ B|.
|
||||||
|
func jaccardSimilarity(a, b []uint32) float64 {
|
||||||
|
if len(a) == 0 && len(b) == 0 {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
i, j := 0, 0
|
||||||
|
intersection := 0
|
||||||
|
|
||||||
|
for i < len(a) && j < len(b) {
|
||||||
|
if a[i] == b[j] {
|
||||||
|
intersection++
|
||||||
|
i++
|
||||||
|
j++
|
||||||
|
} else if a[i] < b[j] {
|
||||||
|
i++
|
||||||
|
} else {
|
||||||
|
j++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
union := len(a) + len(b) - intersection
|
||||||
|
return float64(intersection) / float64(union)
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyResults(results []SearchResult) []SearchResult {
|
||||||
|
if results == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cp := make([]SearchResult, len(results))
|
||||||
|
copy(cp, results)
|
||||||
|
return cp
|
||||||
|
}
|
||||||
200
pkg/skills/search_cache_test.go
Normal file
200
pkg/skills/search_cache_test.go
Normal file
|
|
@ -0,0 +1,200 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSearchCacheExactHit(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 5*time.Minute)
|
||||||
|
|
||||||
|
results := []SearchResult{
|
||||||
|
{Slug: "github", Score: 0.9, RegistryName: "clawhub"},
|
||||||
|
{Slug: "docker", Score: 0.7, RegistryName: "clawhub"},
|
||||||
|
}
|
||||||
|
cache.Put("github integration", results)
|
||||||
|
|
||||||
|
got, hit := cache.Get("github integration")
|
||||||
|
assert.True(t, hit)
|
||||||
|
assert.Len(t, got, 2)
|
||||||
|
assert.Equal(t, "github", got[0].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheExactHitCaseInsensitive(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 5*time.Minute)
|
||||||
|
|
||||||
|
results := []SearchResult{{Slug: "github", Score: 0.9}}
|
||||||
|
cache.Put("GitHub Integration", results)
|
||||||
|
|
||||||
|
got, hit := cache.Get("github integration")
|
||||||
|
assert.True(t, hit)
|
||||||
|
assert.Len(t, got, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheSimilarHit(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 5*time.Minute)
|
||||||
|
|
||||||
|
results := []SearchResult{{Slug: "github", Score: 0.9}}
|
||||||
|
cache.Put("github integration tool", results)
|
||||||
|
|
||||||
|
// "github integration" is very similar to "github integration tool"
|
||||||
|
got, hit := cache.Get("github integration")
|
||||||
|
assert.True(t, hit)
|
||||||
|
assert.Len(t, got, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheDissimilarMiss(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 5*time.Minute)
|
||||||
|
|
||||||
|
results := []SearchResult{{Slug: "github", Score: 0.9}}
|
||||||
|
cache.Put("github integration", results)
|
||||||
|
|
||||||
|
// Completely unrelated query
|
||||||
|
_, hit := cache.Get("database management")
|
||||||
|
assert.False(t, hit)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheTTLExpiration(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 50*time.Millisecond)
|
||||||
|
|
||||||
|
results := []SearchResult{{Slug: "github", Score: 0.9}}
|
||||||
|
cache.Put("github integration", results)
|
||||||
|
|
||||||
|
// Immediately should hit
|
||||||
|
_, hit := cache.Get("github integration")
|
||||||
|
assert.True(t, hit)
|
||||||
|
|
||||||
|
// Wait for expiration
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
|
||||||
|
_, hit = cache.Get("github integration")
|
||||||
|
assert.False(t, hit)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheLRUEviction(t *testing.T) {
|
||||||
|
cache := NewSearchCache(3, 5*time.Minute)
|
||||||
|
|
||||||
|
cache.Put("query-1", []SearchResult{{Slug: "a"}})
|
||||||
|
cache.Put("query-2", []SearchResult{{Slug: "b"}})
|
||||||
|
cache.Put("query-3", []SearchResult{{Slug: "c"}})
|
||||||
|
|
||||||
|
assert.Equal(t, 3, cache.Len())
|
||||||
|
|
||||||
|
// Adding a 4th should evict query-1 (oldest)
|
||||||
|
cache.Put("query-4", []SearchResult{{Slug: "d"}})
|
||||||
|
assert.Equal(t, 3, cache.Len())
|
||||||
|
|
||||||
|
_, hit := cache.Get("query-1")
|
||||||
|
assert.False(t, hit, "oldest entry should be evicted")
|
||||||
|
|
||||||
|
got, hit := cache.Get("query-4")
|
||||||
|
assert.True(t, hit)
|
||||||
|
assert.Equal(t, "d", got[0].Slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheEmptyQuery(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 5*time.Minute)
|
||||||
|
|
||||||
|
_, hit := cache.Get("")
|
||||||
|
assert.False(t, hit)
|
||||||
|
|
||||||
|
_, hit = cache.Get(" ")
|
||||||
|
assert.False(t, hit)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheResultsCopied(t *testing.T) {
|
||||||
|
cache := NewSearchCache(10, 5*time.Minute)
|
||||||
|
|
||||||
|
original := []SearchResult{{Slug: "github", Score: 0.9}}
|
||||||
|
cache.Put("test", original)
|
||||||
|
|
||||||
|
// Mutate original after putting
|
||||||
|
original[0].Slug = "mutated"
|
||||||
|
|
||||||
|
got, hit := cache.Get("test")
|
||||||
|
assert.True(t, hit)
|
||||||
|
assert.Equal(t, "github", got[0].Slug, "cache should hold a copy, not a reference")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildTrigrams(t *testing.T) {
|
||||||
|
trigrams := buildTrigrams("hello")
|
||||||
|
assert.Contains(t, trigrams, uint32('h')<<16|uint32('e')<<8|uint32('l'))
|
||||||
|
assert.Contains(t, trigrams, uint32('e')<<16|uint32('l')<<8|uint32('l'))
|
||||||
|
assert.Contains(t, trigrams, uint32('l')<<16|uint32('l')<<8|uint32('o'))
|
||||||
|
assert.Len(t, trigrams, 3)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJaccardSimilarity(t *testing.T) {
|
||||||
|
a := buildTrigrams("github integration")
|
||||||
|
b := buildTrigrams("github integration tool")
|
||||||
|
|
||||||
|
sim := jaccardSimilarity(a, b)
|
||||||
|
assert.Greater(t, sim, 0.5, "similar strings should have high sim")
|
||||||
|
|
||||||
|
c := buildTrigrams("completely different query about databases")
|
||||||
|
sim2 := jaccardSimilarity(a, c)
|
||||||
|
assert.Less(t, sim2, 0.3, "dissimilar strings should have low sim")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJaccardSimilarityEdgeCases(t *testing.T) {
|
||||||
|
empty := buildTrigrams("")
|
||||||
|
nonempty := buildTrigrams("hello")
|
||||||
|
|
||||||
|
assert.Equal(t, 1.0, jaccardSimilarity(empty, empty))
|
||||||
|
assert.Equal(t, 0.0, jaccardSimilarity(empty, nonempty))
|
||||||
|
assert.Equal(t, 0.0, jaccardSimilarity(nonempty, empty))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheConcurrency(t *testing.T) {
|
||||||
|
cache := NewSearchCache(50, 5*time.Minute)
|
||||||
|
done := make(chan struct{})
|
||||||
|
|
||||||
|
// Concurrent writes
|
||||||
|
go func() {
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
cache.Put("query-write-"+string(rune('a'+i%26)), []SearchResult{{Slug: "x"}})
|
||||||
|
}
|
||||||
|
done <- struct{}{}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Concurrent reads
|
||||||
|
go func() {
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
cache.Get("query-write-a")
|
||||||
|
}
|
||||||
|
done <- struct{}{}
|
||||||
|
}()
|
||||||
|
|
||||||
|
<-done
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchCacheLRUUpdateOnGet(t *testing.T) {
|
||||||
|
// Capacity 3
|
||||||
|
cache := NewSearchCache(3, time.Hour)
|
||||||
|
|
||||||
|
// Fill cache: query-A, query-B, query-C
|
||||||
|
// Use longer strings to ensure trigrams are generated and avoid false positive similarity
|
||||||
|
cache.Put("query-A", []SearchResult{{Slug: "A"}})
|
||||||
|
cache.Put("query-B", []SearchResult{{Slug: "B"}})
|
||||||
|
cache.Put("query-C", []SearchResult{{Slug: "C"}})
|
||||||
|
|
||||||
|
// Access query-A (should make it most recently used)
|
||||||
|
if _, found := cache.Get("query-A"); !found {
|
||||||
|
t.Fatal("query-A should be in cache")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add query-D. Should evict query-B (LRU) instead of query-A (which was refreshed)
|
||||||
|
cache.Put("query-D", []SearchResult{{Slug: "D"}})
|
||||||
|
|
||||||
|
// Check if query-A is still there
|
||||||
|
if _, found := cache.Get("query-A"); !found {
|
||||||
|
t.Fatalf("query-A was evicted! valid LRU should have kept query-A and evicted query-B.")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if query-B is evicted
|
||||||
|
if _, found := cache.Get("query-B"); found {
|
||||||
|
t.Fatal("query-B should have been evicted")
|
||||||
|
}
|
||||||
|
}
|
||||||
199
pkg/tools/skills_install.go
Normal file
199
pkg/tools/skills_install.go
Normal file
|
|
@ -0,0 +1,199 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
// InstallSkillTool allows the LLM agent to install skills from registries.
|
||||||
|
// It shares the same RegistryManager that FindSkillsTool uses,
|
||||||
|
// so all registries configured in config are available for installation.
|
||||||
|
type InstallSkillTool struct {
|
||||||
|
registryMgr *skills.RegistryManager
|
||||||
|
workspace string
|
||||||
|
mu sync.Mutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewInstallSkillTool creates a new InstallSkillTool.
|
||||||
|
// registryMgr is the shared registry manager (same instance as FindSkillsTool).
|
||||||
|
// workspace is the root workspace directory; skills install to {workspace}/skills/{slug}/.
|
||||||
|
func NewInstallSkillTool(registryMgr *skills.RegistryManager, workspace string) *InstallSkillTool {
|
||||||
|
return &InstallSkillTool{
|
||||||
|
registryMgr: registryMgr,
|
||||||
|
workspace: workspace,
|
||||||
|
mu: sync.Mutex{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *InstallSkillTool) Name() string {
|
||||||
|
return "install_skill"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *InstallSkillTool) Description() string {
|
||||||
|
return "Install a skill from a registry by slug. Downloads and extracts the skill into the workspace. Use find_skills first to discover available skills."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *InstallSkillTool) Parameters() map[string]interface{} {
|
||||||
|
return map[string]interface{}{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]interface{}{
|
||||||
|
"slug": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "The unique slug of the skill to install (e.g., 'github', 'docker-compose')",
|
||||||
|
},
|
||||||
|
"version": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Specific version to install (optional, defaults to latest)",
|
||||||
|
},
|
||||||
|
"registry": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Registry to install from (required, e.g., 'clawhub')",
|
||||||
|
},
|
||||||
|
"force": map[string]interface{}{
|
||||||
|
"type": "boolean",
|
||||||
|
"description": "Force reinstall if skill already exists (default false)",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"slug", "registry"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *InstallSkillTool) Execute(ctx context.Context, args map[string]interface{}) *ToolResult {
|
||||||
|
// Install lock to prevent concurrent directory operations.
|
||||||
|
// Ideally this should be done at a `slug` level, currently, its at a `workspace` level.
|
||||||
|
t.mu.Lock()
|
||||||
|
defer t.mu.Unlock()
|
||||||
|
|
||||||
|
// Validate slug
|
||||||
|
slug, _ := args["slug"].(string)
|
||||||
|
if err := utils.ValidateSkillIdentifier(slug); err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("invalid slug %q: error: %s", slug, err.Error()))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate registry
|
||||||
|
registryName, _ := args["registry"].(string)
|
||||||
|
if err := utils.ValidateSkillIdentifier(registryName); err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("invalid registry %q: error: %s", registryName, err.Error()))
|
||||||
|
}
|
||||||
|
|
||||||
|
version, _ := args["version"].(string)
|
||||||
|
force, _ := args["force"].(bool)
|
||||||
|
|
||||||
|
// Check if already installed.
|
||||||
|
skillsDir := filepath.Join(t.workspace, "skills")
|
||||||
|
targetDir := filepath.Join(skillsDir, slug)
|
||||||
|
|
||||||
|
if !force {
|
||||||
|
if _, err := os.Stat(targetDir); err == nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("skill %q already installed at %s. Use force=true to reinstall.", slug, targetDir))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Force: remove existing if present.
|
||||||
|
os.RemoveAll(targetDir)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resolve which registry to use.
|
||||||
|
registry := t.registryMgr.GetRegistry(registryName)
|
||||||
|
if registry == nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("registry %q not found", registryName))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure skills directory exists.
|
||||||
|
if err := os.MkdirAll(skillsDir, 0755); err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to create skills directory: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Download and install (handles metadata, version resolution, extraction).
|
||||||
|
result, err := registry.DownloadAndInstall(ctx, slug, version, targetDir)
|
||||||
|
if err != nil {
|
||||||
|
// Clean up partial install.
|
||||||
|
rmErr := os.RemoveAll(targetDir)
|
||||||
|
if rmErr != nil {
|
||||||
|
logger.ErrorCF("tool", "Failed to remove partial install",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tool": "install_skill",
|
||||||
|
"target_dir": targetDir,
|
||||||
|
"error": rmErr.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to install %q: %v", slug, err))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Moderation: block malware.
|
||||||
|
if result.IsMalwareBlocked {
|
||||||
|
rmErr := os.RemoveAll(targetDir)
|
||||||
|
if rmErr != nil {
|
||||||
|
logger.ErrorCF("tool", "Failed to remove partial install",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tool": "install_skill",
|
||||||
|
"target_dir": targetDir,
|
||||||
|
"error": rmErr.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return ErrorResult(fmt.Sprintf("skill %q is flagged as malicious and cannot be installed", slug))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write origin metadata.
|
||||||
|
if err := writeOriginMeta(targetDir, registry.Name(), slug, result.Version); err != nil {
|
||||||
|
logger.ErrorCF("tool", "Failed to write origin metadata",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tool": "install_skill",
|
||||||
|
"error": err.Error(),
|
||||||
|
"target": targetDir,
|
||||||
|
"registry": registry.Name(),
|
||||||
|
"slug": slug,
|
||||||
|
"version": result.Version,
|
||||||
|
})
|
||||||
|
_ = err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build result with moderation warning if suspicious.
|
||||||
|
var output string
|
||||||
|
if result.IsSuspicious {
|
||||||
|
output = fmt.Sprintf("⚠️ Warning: skill %q is flagged as suspicious (may contain risky patterns).\n\n", slug)
|
||||||
|
}
|
||||||
|
output += fmt.Sprintf("Successfully installed skill %q v%s from %s registry.\nLocation: %s\n",
|
||||||
|
slug, result.Version, registry.Name(), targetDir)
|
||||||
|
|
||||||
|
if result.Summary != "" {
|
||||||
|
output += fmt.Sprintf("Description: %s\n", result.Summary)
|
||||||
|
}
|
||||||
|
output += "\nThe skill is now available and can be loaded in the current session."
|
||||||
|
|
||||||
|
return SilentResult(output)
|
||||||
|
}
|
||||||
|
|
||||||
|
// originMeta tracks which registry a skill was installed from.
|
||||||
|
type originMeta struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
Registry string `json:"registry"`
|
||||||
|
Slug string `json:"slug"`
|
||||||
|
InstalledVersion string `json:"installed_version"`
|
||||||
|
InstalledAt int64 `json:"installed_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeOriginMeta(targetDir, registryName, slug, version string) error {
|
||||||
|
meta := originMeta{
|
||||||
|
Version: 1,
|
||||||
|
Registry: registryName,
|
||||||
|
Slug: slug,
|
||||||
|
InstalledVersion: version,
|
||||||
|
InstalledAt: time.Now().UnixMilli(),
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.MarshalIndent(meta, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return os.WriteFile(filepath.Join(targetDir, ".skill-origin.json"), data, 0644)
|
||||||
|
}
|
||||||
103
pkg/tools/skills_install_test.go
Normal file
103
pkg/tools/skills_install_test.go
Normal file
|
|
@ -0,0 +1,103 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestInstallSkillToolName(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
||||||
|
assert.Equal(t, "install_skill", tool.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolMissingSlug(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "identifier is required and must be a non-empty string")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolEmptySlug(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"slug": " ",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "identifier is required and must be a non-empty string")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolUnsafeSlug(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
||||||
|
|
||||||
|
cases := []string{
|
||||||
|
"../etc/passwd",
|
||||||
|
"path/traversal",
|
||||||
|
"path\\traversal",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, slug := range cases {
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"slug": slug,
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError, "slug %q should be rejected", slug)
|
||||||
|
assert.Contains(t, result.ForLLM, "invalid slug")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolAlreadyExists(t *testing.T) {
|
||||||
|
workspace := t.TempDir()
|
||||||
|
skillDir := filepath.Join(workspace, "skills", "existing-skill")
|
||||||
|
require.NoError(t, os.MkdirAll(skillDir, 0755))
|
||||||
|
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), workspace)
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"slug": "existing-skill",
|
||||||
|
"registry": "clawhub",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "already installed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolRegistryNotFound(t *testing.T) {
|
||||||
|
workspace := t.TempDir()
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), workspace)
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"slug": "some-skill",
|
||||||
|
"registry": "nonexistent",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "registry")
|
||||||
|
assert.Contains(t, result.ForLLM, "not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolParameters(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
||||||
|
params := tool.Parameters()
|
||||||
|
|
||||||
|
props, ok := params["properties"].(map[string]interface{})
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Contains(t, props, "slug")
|
||||||
|
assert.Contains(t, props, "version")
|
||||||
|
assert.Contains(t, props, "registry")
|
||||||
|
assert.Contains(t, props, "force")
|
||||||
|
|
||||||
|
required, ok := params["required"].([]string)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Contains(t, required, "slug")
|
||||||
|
assert.Contains(t, required, "registry")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstallSkillToolMissingRegistry(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"slug": "some-skill",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "invalid registry")
|
||||||
|
}
|
||||||
119
pkg/tools/skills_search.go
Normal file
119
pkg/tools/skills_search.go
Normal file
|
|
@ -0,0 +1,119 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FindSkillsTool allows the LLM agent to search for installable skills from registries.
|
||||||
|
type FindSkillsTool struct {
|
||||||
|
registryMgr *skills.RegistryManager
|
||||||
|
cache *skills.SearchCache
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewFindSkillsTool creates a new FindSkillsTool.
|
||||||
|
// registryMgr is the shared registry manager (built from config in createToolRegistry).
|
||||||
|
// cache is the search cache for deduplicating similar queries.
|
||||||
|
func NewFindSkillsTool(registryMgr *skills.RegistryManager, cache *skills.SearchCache) *FindSkillsTool {
|
||||||
|
return &FindSkillsTool{
|
||||||
|
registryMgr: registryMgr,
|
||||||
|
cache: cache,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *FindSkillsTool) Name() string {
|
||||||
|
return "find_skills"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *FindSkillsTool) Description() string {
|
||||||
|
return "Search for installable skills from skill registries. Returns skill slugs, descriptions, versions, and relevance scores. Use this to discover skills before installing them with install_skill."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *FindSkillsTool) Parameters() map[string]interface{} {
|
||||||
|
return map[string]interface{}{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]interface{}{
|
||||||
|
"query": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Search query describing the desired skill capability (e.g., 'github integration', 'database management')",
|
||||||
|
},
|
||||||
|
"limit": map[string]interface{}{
|
||||||
|
"type": "integer",
|
||||||
|
"description": "Maximum number of results to return (1-20, default 5)",
|
||||||
|
"minimum": 1.0,
|
||||||
|
"maximum": 20.0,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"query"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *FindSkillsTool) Execute(ctx context.Context, args map[string]interface{}) *ToolResult {
|
||||||
|
query, ok := args["query"].(string)
|
||||||
|
query = strings.ToLower(strings.TrimSpace(query))
|
||||||
|
if !ok || query == "" {
|
||||||
|
return ErrorResult("query is required and must be a non-empty string")
|
||||||
|
}
|
||||||
|
|
||||||
|
limit := 5
|
||||||
|
if l, ok := args["limit"].(float64); ok {
|
||||||
|
li := int(l)
|
||||||
|
if li >= 1 && li <= 20 {
|
||||||
|
limit = li
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check cache first.
|
||||||
|
if t.cache != nil {
|
||||||
|
if cached, hit := t.cache.Get(query); hit {
|
||||||
|
return SilentResult(formatSearchResults(query, cached, true))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Search all registries.
|
||||||
|
results, err := t.registryMgr.SearchAll(ctx, query, limit)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("skill search failed: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cache the results.
|
||||||
|
if t.cache != nil && len(results) > 0 {
|
||||||
|
t.cache.Put(query, results)
|
||||||
|
}
|
||||||
|
|
||||||
|
return SilentResult(formatSearchResults(query, results, false))
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatSearchResults(query string, results []skills.SearchResult, cached bool) string {
|
||||||
|
if len(results) == 0 {
|
||||||
|
return fmt.Sprintf("No skills found for query: %q", query)
|
||||||
|
}
|
||||||
|
|
||||||
|
var sb strings.Builder
|
||||||
|
source := ""
|
||||||
|
if cached {
|
||||||
|
source = " (cached)"
|
||||||
|
}
|
||||||
|
sb.WriteString(fmt.Sprintf("Found %d skills for %q%s:\n\n", len(results), query, source))
|
||||||
|
|
||||||
|
for i, r := range results {
|
||||||
|
sb.WriteString(fmt.Sprintf("%d. **%s**", i+1, r.Slug))
|
||||||
|
if r.Version != "" {
|
||||||
|
sb.WriteString(fmt.Sprintf(" v%s", r.Version))
|
||||||
|
}
|
||||||
|
sb.WriteString(fmt.Sprintf(" (score: %.3f, registry: %s)\n", r.Score, r.RegistryName))
|
||||||
|
if r.DisplayName != "" && r.DisplayName != r.Slug {
|
||||||
|
sb.WriteString(fmt.Sprintf(" Name: %s\n", r.DisplayName))
|
||||||
|
}
|
||||||
|
if r.Summary != "" {
|
||||||
|
sb.WriteString(fmt.Sprintf(" %s\n", r.Summary))
|
||||||
|
}
|
||||||
|
sb.WriteString("\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
sb.WriteString("Use install_skill with the slug to install a skill.")
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
82
pkg/tools/skills_search_test.go
Normal file
82
pkg/tools/skills_search_test.go
Normal file
|
|
@ -0,0 +1,82 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestFindSkillsToolName(t *testing.T) {
|
||||||
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
||||||
|
assert.Equal(t, "find_skills", tool.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindSkillsToolMissingQuery(t *testing.T) {
|
||||||
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "query is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindSkillsToolEmptyQuery(t *testing.T) {
|
||||||
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"query": " ",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindSkillsToolCacheHit(t *testing.T) {
|
||||||
|
cache := skills.NewSearchCache(10, 5*60*1000*1000*1000) // 5 min
|
||||||
|
cache.Put("github", []skills.SearchResult{
|
||||||
|
{Slug: "github", Score: 0.9, RegistryName: "clawhub"},
|
||||||
|
})
|
||||||
|
|
||||||
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), cache)
|
||||||
|
result := tool.Execute(context.Background(), map[string]interface{}{
|
||||||
|
"query": "github",
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.False(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "github")
|
||||||
|
assert.Contains(t, result.ForLLM, "cached")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindSkillsToolParameters(t *testing.T) {
|
||||||
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
||||||
|
params := tool.Parameters()
|
||||||
|
|
||||||
|
props, ok := params["properties"].(map[string]interface{})
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Contains(t, props, "query")
|
||||||
|
assert.Contains(t, props, "limit")
|
||||||
|
|
||||||
|
required, ok := params["required"].([]string)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Contains(t, required, "query")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindSkillsToolDescription(t *testing.T) {
|
||||||
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
||||||
|
assert.NotEmpty(t, tool.Description())
|
||||||
|
assert.Contains(t, tool.Description(), "skill")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatSearchResultsEmpty(t *testing.T) {
|
||||||
|
result := formatSearchResults("test query", nil, false)
|
||||||
|
assert.Contains(t, result, "No skills found")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatSearchResultsWithData(t *testing.T) {
|
||||||
|
results := []skills.SearchResult{
|
||||||
|
{Slug: "github", Score: 0.95, DisplayName: "GitHub", Summary: "GitHub API integration", Version: "1.0.0", RegistryName: "clawhub"},
|
||||||
|
}
|
||||||
|
output := formatSearchResults("github", results, false)
|
||||||
|
assert.Contains(t, output, "github")
|
||||||
|
assert.Contains(t, output, "v1.0.0")
|
||||||
|
assert.Contains(t, output, "0.950")
|
||||||
|
assert.Contains(t, output, "clawhub")
|
||||||
|
assert.Contains(t, output, "install_skill")
|
||||||
|
}
|
||||||
|
|
@ -23,15 +23,19 @@ type SubagentTask struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type SubagentManager struct {
|
type SubagentManager struct {
|
||||||
tasks map[string]*SubagentTask
|
tasks map[string]*SubagentTask
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
provider providers.LLMProvider
|
provider providers.LLMProvider
|
||||||
defaultModel string
|
defaultModel string
|
||||||
bus *bus.MessageBus
|
bus *bus.MessageBus
|
||||||
workspace string
|
workspace string
|
||||||
tools *ToolRegistry
|
tools *ToolRegistry
|
||||||
maxIterations int
|
maxIterations int
|
||||||
nextID int
|
maxTokens int
|
||||||
|
temperature float64
|
||||||
|
hasMaxTokens bool
|
||||||
|
hasTemperature bool
|
||||||
|
nextID int
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewSubagentManager(provider providers.LLMProvider, defaultModel, workspace string, bus *bus.MessageBus) *SubagentManager {
|
func NewSubagentManager(provider providers.LLMProvider, defaultModel, workspace string, bus *bus.MessageBus) *SubagentManager {
|
||||||
|
|
@ -47,6 +51,16 @@ func NewSubagentManager(provider providers.LLMProvider, defaultModel, workspace
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetLLMOptions sets max tokens and temperature for subagent LLM calls.
|
||||||
|
func (sm *SubagentManager) SetLLMOptions(maxTokens int, temperature float64) {
|
||||||
|
sm.mu.Lock()
|
||||||
|
defer sm.mu.Unlock()
|
||||||
|
sm.maxTokens = maxTokens
|
||||||
|
sm.hasMaxTokens = true
|
||||||
|
sm.temperature = temperature
|
||||||
|
sm.hasTemperature = true
|
||||||
|
}
|
||||||
|
|
||||||
// SetTools sets the tool registry for subagent execution.
|
// SetTools sets the tool registry for subagent execution.
|
||||||
// If not set, subagent will have access to the provided tools.
|
// If not set, subagent will have access to the provided tools.
|
||||||
func (sm *SubagentManager) SetTools(tools *ToolRegistry) {
|
func (sm *SubagentManager) SetTools(tools *ToolRegistry) {
|
||||||
|
|
@ -125,17 +139,29 @@ After completing the task, provide a clear summary of what was done.`
|
||||||
sm.mu.RLock()
|
sm.mu.RLock()
|
||||||
tools := sm.tools
|
tools := sm.tools
|
||||||
maxIter := sm.maxIterations
|
maxIter := sm.maxIterations
|
||||||
|
maxTokens := sm.maxTokens
|
||||||
|
temperature := sm.temperature
|
||||||
|
hasMaxTokens := sm.hasMaxTokens
|
||||||
|
hasTemperature := sm.hasTemperature
|
||||||
sm.mu.RUnlock()
|
sm.mu.RUnlock()
|
||||||
|
|
||||||
|
var llmOptions map[string]any
|
||||||
|
if hasMaxTokens || hasTemperature {
|
||||||
|
llmOptions = map[string]any{}
|
||||||
|
if hasMaxTokens {
|
||||||
|
llmOptions["max_tokens"] = maxTokens
|
||||||
|
}
|
||||||
|
if hasTemperature {
|
||||||
|
llmOptions["temperature"] = temperature
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
loopResult, err := RunToolLoop(ctx, ToolLoopConfig{
|
loopResult, err := RunToolLoop(ctx, ToolLoopConfig{
|
||||||
Provider: sm.provider,
|
Provider: sm.provider,
|
||||||
Model: sm.defaultModel,
|
Model: sm.defaultModel,
|
||||||
Tools: tools,
|
Tools: tools,
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
LLMOptions: map[string]any{
|
LLMOptions: llmOptions,
|
||||||
"max_tokens": 4096,
|
|
||||||
"temperature": 0.7,
|
|
||||||
},
|
|
||||||
}, messages, task.OriginChannel, task.OriginChatID)
|
}, messages, task.OriginChannel, task.OriginChatID)
|
||||||
|
|
||||||
sm.mu.Lock()
|
sm.mu.Lock()
|
||||||
|
|
@ -283,19 +309,30 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
sm.mu.RLock()
|
sm.mu.RLock()
|
||||||
tools := sm.tools
|
tools := sm.tools
|
||||||
maxIter := sm.maxIterations
|
maxIter := sm.maxIterations
|
||||||
|
maxTokens := sm.maxTokens
|
||||||
|
temperature := sm.temperature
|
||||||
|
hasMaxTokens := sm.hasMaxTokens
|
||||||
|
hasTemperature := sm.hasTemperature
|
||||||
sm.mu.RUnlock()
|
sm.mu.RUnlock()
|
||||||
|
|
||||||
|
var llmOptions map[string]any
|
||||||
|
if hasMaxTokens || hasTemperature {
|
||||||
|
llmOptions = map[string]any{}
|
||||||
|
if hasMaxTokens {
|
||||||
|
llmOptions["max_tokens"] = maxTokens
|
||||||
|
}
|
||||||
|
if hasTemperature {
|
||||||
|
llmOptions["temperature"] = temperature
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
loopResult, err := RunToolLoop(ctx, ToolLoopConfig{
|
loopResult, err := RunToolLoop(ctx, ToolLoopConfig{
|
||||||
Provider: sm.provider,
|
Provider: sm.provider,
|
||||||
Model: sm.defaultModel,
|
Model: sm.defaultModel,
|
||||||
Tools: tools,
|
Tools: tools,
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
LLMOptions: map[string]any{
|
LLMOptions: llmOptions,
|
||||||
"max_tokens": 4096,
|
|
||||||
"temperature": 0.7,
|
|
||||||
},
|
|
||||||
}, messages, t.originChannel, t.originChatID)
|
}, messages, t.originChannel, t.originChatID)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -10,9 +10,12 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
// MockLLMProvider is a test implementation of LLMProvider
|
// MockLLMProvider is a test implementation of LLMProvider
|
||||||
type MockLLMProvider struct{}
|
type MockLLMProvider struct {
|
||||||
|
lastOptions map[string]interface{}
|
||||||
|
}
|
||||||
|
|
||||||
func (m *MockLLMProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, options map[string]interface{}) (*providers.LLMResponse, error) {
|
func (m *MockLLMProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, options map[string]interface{}) (*providers.LLMResponse, error) {
|
||||||
|
m.lastOptions = options
|
||||||
// Find the last user message to generate a response
|
// Find the last user message to generate a response
|
||||||
for i := len(messages) - 1; i >= 0; i-- {
|
for i := len(messages) - 1; i >= 0; i-- {
|
||||||
if messages[i].Role == "user" {
|
if messages[i].Role == "user" {
|
||||||
|
|
@ -36,6 +39,32 @@ func (m *MockLLMProvider) GetContextWindow() int {
|
||||||
return 4096
|
return 4096
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSubagentManager_SetLLMOptions_AppliesToRunToolLoop(t *testing.T) {
|
||||||
|
provider := &MockLLMProvider{}
|
||||||
|
manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil)
|
||||||
|
manager.SetLLMOptions(2048, 0.6)
|
||||||
|
tool := NewSubagentTool(manager)
|
||||||
|
tool.SetContext("cli", "direct")
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]interface{}{"task": "Do something"}
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
|
||||||
|
if result == nil || result.IsError {
|
||||||
|
t.Fatalf("Expected successful result, got: %+v", result)
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider.lastOptions == nil {
|
||||||
|
t.Fatal("Expected LLM options to be passed, got nil")
|
||||||
|
}
|
||||||
|
if provider.lastOptions["max_tokens"] != 2048 {
|
||||||
|
t.Fatalf("max_tokens = %v, want %d", provider.lastOptions["max_tokens"], 2048)
|
||||||
|
}
|
||||||
|
if provider.lastOptions["temperature"] != 0.6 {
|
||||||
|
t.Fatalf("temperature = %v, want %v", provider.lastOptions["temperature"], 0.6)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestSubagentTool_Name verifies tool name
|
// TestSubagentTool_Name verifies tool name
|
||||||
func TestSubagentTool_Name(t *testing.T) {
|
func TestSubagentTool_Name(t *testing.T) {
|
||||||
provider := &MockLLMProvider{}
|
provider := &MockLLMProvider{}
|
||||||
|
|
|
||||||
|
|
@ -55,12 +55,8 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
// 2. Set default LLM options
|
// 2. Set default LLM options
|
||||||
llmOpts := config.LLMOptions
|
llmOpts := config.LLMOptions
|
||||||
if llmOpts == nil {
|
if llmOpts == nil {
|
||||||
llmOpts = map[string]any{
|
llmOpts = map[string]any{}
|
||||||
"max_tokens": 4096,
|
|
||||||
"temperature": 0.7,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 3. Call LLM
|
// 3. Call LLM
|
||||||
response, err := config.Provider.Chat(ctx, messages, providerToolDefs, config.Model, llmOpts)
|
response, err := config.Provider.Chat(ctx, messages, providerToolDefs, config.Model, llmOpts)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -83,15 +79,20 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
// 5. Log tool calls
|
normalizedToolCalls := make([]providers.ToolCall, 0, len(response.ToolCalls))
|
||||||
toolNames := make([]string, 0, len(response.ToolCalls))
|
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range response.ToolCalls {
|
||||||
|
normalizedToolCalls = append(normalizedToolCalls, providers.NormalizeToolCall(tc))
|
||||||
|
}
|
||||||
|
|
||||||
|
// 5. Log tool calls
|
||||||
|
toolNames := make([]string, 0, len(normalizedToolCalls))
|
||||||
|
for _, tc := range normalizedToolCalls {
|
||||||
toolNames = append(toolNames, tc.Name)
|
toolNames = append(toolNames, tc.Name)
|
||||||
}
|
}
|
||||||
logger.InfoCF("toolloop", "LLM requested tool calls",
|
logger.InfoCF("toolloop", "LLM requested tool calls",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"tools": toolNames,
|
"tools": toolNames,
|
||||||
"count": len(response.ToolCalls),
|
"count": len(normalizedToolCalls),
|
||||||
"iteration": iteration,
|
"iteration": iteration,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
@ -100,11 +101,13 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
Role: "assistant",
|
Role: "assistant",
|
||||||
Content: response.Content,
|
Content: response.Content,
|
||||||
}
|
}
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range normalizedToolCalls {
|
||||||
argumentsJSON, _ := json.Marshal(tc.Arguments)
|
argumentsJSON, _ := json.Marshal(tc.Arguments)
|
||||||
assistantMsg.ToolCalls = append(assistantMsg.ToolCalls, providers.ToolCall{
|
assistantMsg.ToolCalls = append(assistantMsg.ToolCalls, providers.ToolCall{
|
||||||
ID: tc.ID,
|
ID: tc.ID,
|
||||||
Type: "function",
|
Type: "function",
|
||||||
|
Name: tc.Name,
|
||||||
|
Arguments: tc.Arguments,
|
||||||
Function: &providers.FunctionCall{
|
Function: &providers.FunctionCall{
|
||||||
Name: tc.Name,
|
Name: tc.Name,
|
||||||
Arguments: string(argumentsJSON),
|
Arguments: string(argumentsJSON),
|
||||||
|
|
@ -114,7 +117,7 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
messages = append(messages, assistantMsg)
|
messages = append(messages, assistantMsg)
|
||||||
|
|
||||||
// 7. Execute tool calls
|
// 7. Execute tool calls
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range normalizedToolCalls {
|
||||||
argsJSON, _ := json.Marshal(tc.Arguments)
|
argsJSON, _ := json.Marshal(tc.Arguments)
|
||||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||||
logger.InfoCF("toolloop", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
logger.InfoCF("toolloop", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
||||||
|
|
|
||||||
|
|
@ -755,8 +755,10 @@ func (t *WebFetchTool) extractText(htmlContent string) string {
|
||||||
|
|
||||||
result = strings.TrimSpace(result)
|
result = strings.TrimSpace(result)
|
||||||
|
|
||||||
re = regexp.MustCompile(`\s+`)
|
re = regexp.MustCompile(`[^\S\n]+`)
|
||||||
result = re.ReplaceAllLiteralString(result, " ")
|
result = re.ReplaceAllString(result, " ")
|
||||||
|
re = regexp.MustCompile(`\n{3,}`)
|
||||||
|
result = re.ReplaceAllString(result, "\n\n")
|
||||||
|
|
||||||
lines := strings.Split(result, "\n")
|
lines := strings.Split(result, "\n")
|
||||||
var cleanLines []string
|
var cleanLines []string
|
||||||
|
|
|
||||||
|
|
@ -361,6 +361,80 @@ func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestWebFetchTool_extractText verifies text extraction preserves newlines
|
||||||
|
func TestWebFetchTool_extractText(t *testing.T) {
|
||||||
|
tool := &WebFetchTool{}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
wantFunc func(t *testing.T, got string)
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "preserves newlines between block elements",
|
||||||
|
input: "<html><body><h1>Title</h1>\n<p>Paragraph 1</p>\n<p>Paragraph 2</p></body></html>",
|
||||||
|
wantFunc: func(t *testing.T, got string) {
|
||||||
|
lines := strings.Split(got, "\n")
|
||||||
|
if len(lines) < 2 {
|
||||||
|
t.Errorf("Expected multiple lines, got %d: %q", len(lines), got)
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, "Title") || !strings.Contains(got, "Paragraph 1") || !strings.Contains(got, "Paragraph 2") {
|
||||||
|
t.Errorf("Missing expected text: %q", got)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "removes script and style tags",
|
||||||
|
input: "<script>alert('x');</script><style>body{}</style><p>Keep this</p>",
|
||||||
|
wantFunc: func(t *testing.T, got string) {
|
||||||
|
if strings.Contains(got, "alert") || strings.Contains(got, "body{}") {
|
||||||
|
t.Errorf("Expected script/style content removed, got: %q", got)
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, "Keep this") {
|
||||||
|
t.Errorf("Expected 'Keep this' to remain, got: %q", got)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "collapses excessive blank lines",
|
||||||
|
input: "<p>A</p>\n\n\n\n\n<p>B</p>",
|
||||||
|
wantFunc: func(t *testing.T, got string) {
|
||||||
|
if strings.Contains(got, "\n\n\n") {
|
||||||
|
t.Errorf("Expected excessive blank lines collapsed, got: %q", got)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "collapses horizontal whitespace",
|
||||||
|
input: "<p>hello world</p>",
|
||||||
|
wantFunc: func(t *testing.T, got string) {
|
||||||
|
if strings.Contains(got, " ") {
|
||||||
|
t.Errorf("Expected spaces collapsed, got: %q", got)
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, "hello world") {
|
||||||
|
t.Errorf("Expected 'hello world', got: %q", got)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty input",
|
||||||
|
input: "",
|
||||||
|
wantFunc: func(t *testing.T, got string) {
|
||||||
|
if got != "" {
|
||||||
|
t.Errorf("Expected empty string, got: %q", got)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := tool.extractText(tt.input)
|
||||||
|
tt.wantFunc(t, got)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 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 := NewWebFetchTool(50000)
|
tool := NewWebFetchTool(50000)
|
||||||
|
|
|
||||||
93
pkg/utils/download.go
Normal file
93
pkg/utils/download.go
Normal file
|
|
@ -0,0 +1,93 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DownloadToFile streams an HTTP response body to a temporary file in small
|
||||||
|
// chunks (~32KB), keeping peak memory usage constant regardless of file size.
|
||||||
|
//
|
||||||
|
// Parameters:
|
||||||
|
// - ctx: context for cancellation/timeout
|
||||||
|
// - client: HTTP client to use (caller controls timeouts, transport, etc.)
|
||||||
|
// - req: fully prepared *http.Request (method, URL, headers, etc.)
|
||||||
|
// - maxBytes: maximum bytes to download; 0 means no limit
|
||||||
|
//
|
||||||
|
// Returns the path to the temporary file. The caller is responsible for
|
||||||
|
// removing it when done (defer os.Remove(path)).
|
||||||
|
//
|
||||||
|
// On any error the temp file is cleaned up automatically.
|
||||||
|
func DownloadToFile(ctx context.Context, client *http.Client, req *http.Request, maxBytes int64) (string, error) {
|
||||||
|
// Attach context.
|
||||||
|
req = req.WithContext(ctx)
|
||||||
|
|
||||||
|
logger.DebugCF("download", "Starting download", map[string]interface{}{
|
||||||
|
"url": req.URL.String(),
|
||||||
|
"max_bytes": maxBytes,
|
||||||
|
})
|
||||||
|
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("request failed: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||||
|
// Read a small amount for the error message.
|
||||||
|
errBody := make([]byte, 512)
|
||||||
|
n, _ := io.ReadFull(resp.Body, errBody)
|
||||||
|
return "", fmt.Errorf("HTTP %d: %s", resp.StatusCode, string(errBody[:n]))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create temp file.
|
||||||
|
tmpFile, err := os.CreateTemp("", "picoclaw-dl-*")
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create temp file: %w", err)
|
||||||
|
}
|
||||||
|
tmpPath := tmpFile.Name()
|
||||||
|
|
||||||
|
logger.DebugCF("download", "Streaming to temp file", map[string]interface{}{
|
||||||
|
"path": tmpPath,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Cleanup helper — removes the temp file on any error.
|
||||||
|
cleanup := func() {
|
||||||
|
_ = tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optionally limit the download size.
|
||||||
|
var src io.Reader = resp.Body
|
||||||
|
if maxBytes > 0 {
|
||||||
|
src = io.LimitReader(resp.Body, maxBytes+1) // +1 to detect overflow
|
||||||
|
}
|
||||||
|
|
||||||
|
written, err := io.Copy(tmpFile, src)
|
||||||
|
if err != nil {
|
||||||
|
cleanup()
|
||||||
|
return "", fmt.Errorf("download write failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if maxBytes > 0 && written > maxBytes {
|
||||||
|
cleanup()
|
||||||
|
return "", fmt.Errorf("download too large: %d bytes (max %d)", written, maxBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := tmpFile.Close(); err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Errorf("failed to close temp file: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.DebugCF("download", "Download complete", map[string]interface{}{
|
||||||
|
"path": tmpPath,
|
||||||
|
"bytes_written": written,
|
||||||
|
})
|
||||||
|
|
||||||
|
return tmpPath, nil
|
||||||
|
}
|
||||||
179
pkg/utils/message.go
Normal file
179
pkg/utils/message.go
Normal file
|
|
@ -0,0 +1,179 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SplitMessage splits long messages into chunks, preserving code block integrity.
|
||||||
|
// The function reserves a buffer (10% of maxLen, min 50) to leave room for closing code blocks,
|
||||||
|
// but may extend to maxLen when needed.
|
||||||
|
// Call SplitMessage with the full text content and the maximum allowed length of a single message;
|
||||||
|
// it returns a slice of message chunks that each respect maxLen and avoid splitting fenced code blocks.
|
||||||
|
func SplitMessage(content string, maxLen int) []string {
|
||||||
|
var messages []string
|
||||||
|
|
||||||
|
// Dynamic buffer: 10% of maxLen, but at least 50 chars if possible
|
||||||
|
codeBlockBuffer := maxLen / 10
|
||||||
|
if codeBlockBuffer < 50 {
|
||||||
|
codeBlockBuffer = 50
|
||||||
|
}
|
||||||
|
if codeBlockBuffer > maxLen/2 {
|
||||||
|
codeBlockBuffer = maxLen / 2
|
||||||
|
}
|
||||||
|
|
||||||
|
for len(content) > 0 {
|
||||||
|
if len(content) <= maxLen {
|
||||||
|
messages = append(messages, content)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
// Effective split point: maxLen minus buffer, to leave room for code blocks
|
||||||
|
effectiveLimit := maxLen - codeBlockBuffer
|
||||||
|
if effectiveLimit < maxLen/2 {
|
||||||
|
effectiveLimit = maxLen / 2
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find natural split point within the effective limit
|
||||||
|
msgEnd := findLastNewline(content[:effectiveLimit], 200)
|
||||||
|
if msgEnd <= 0 {
|
||||||
|
msgEnd = findLastSpace(content[:effectiveLimit], 100)
|
||||||
|
}
|
||||||
|
if msgEnd <= 0 {
|
||||||
|
msgEnd = effectiveLimit
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this would end with an incomplete code block
|
||||||
|
candidate := content[:msgEnd]
|
||||||
|
unclosedIdx := findLastUnclosedCodeBlock(candidate)
|
||||||
|
|
||||||
|
if unclosedIdx >= 0 {
|
||||||
|
// Message would end with incomplete code block
|
||||||
|
// Try to extend up to maxLen to include the closing ```
|
||||||
|
if len(content) > msgEnd {
|
||||||
|
closingIdx := findNextClosingCodeBlock(content, msgEnd)
|
||||||
|
if closingIdx > 0 && closingIdx <= maxLen {
|
||||||
|
// Extend to include the closing ```
|
||||||
|
msgEnd = closingIdx
|
||||||
|
} else {
|
||||||
|
// Code block is too long to fit in one chunk or missing closing fence.
|
||||||
|
// Try to split inside by injecting closing and reopening fences.
|
||||||
|
headerEnd := strings.Index(content[unclosedIdx:], "\n")
|
||||||
|
if headerEnd == -1 {
|
||||||
|
headerEnd = unclosedIdx + 3
|
||||||
|
} else {
|
||||||
|
headerEnd += unclosedIdx
|
||||||
|
}
|
||||||
|
header := strings.TrimSpace(content[unclosedIdx:headerEnd])
|
||||||
|
|
||||||
|
// If we have a reasonable amount of content after the header, split inside
|
||||||
|
if msgEnd > headerEnd+20 {
|
||||||
|
// Find a better split point closer to maxLen
|
||||||
|
innerLimit := maxLen - 5 // Leave room for "\n```"
|
||||||
|
betterEnd := findLastNewline(content[:innerLimit], 200)
|
||||||
|
if betterEnd > headerEnd {
|
||||||
|
msgEnd = betterEnd
|
||||||
|
} else {
|
||||||
|
msgEnd = innerLimit
|
||||||
|
}
|
||||||
|
messages = append(messages, strings.TrimRight(content[:msgEnd], " \t\n\r")+"\n```")
|
||||||
|
content = strings.TrimSpace(header + "\n" + content[msgEnd:])
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Otherwise, try to split before the code block starts
|
||||||
|
newEnd := findLastNewline(content[:unclosedIdx], 200)
|
||||||
|
if newEnd <= 0 {
|
||||||
|
newEnd = findLastSpace(content[:unclosedIdx], 100)
|
||||||
|
}
|
||||||
|
if newEnd > 0 {
|
||||||
|
msgEnd = newEnd
|
||||||
|
} else {
|
||||||
|
// If we can't split before, we MUST split inside (last resort)
|
||||||
|
if unclosedIdx > 20 {
|
||||||
|
msgEnd = unclosedIdx
|
||||||
|
} else {
|
||||||
|
msgEnd = maxLen - 5
|
||||||
|
messages = append(messages, strings.TrimRight(content[:msgEnd], " \t\n\r")+"\n```")
|
||||||
|
content = strings.TrimSpace(header + "\n" + content[msgEnd:])
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if msgEnd <= 0 {
|
||||||
|
msgEnd = effectiveLimit
|
||||||
|
}
|
||||||
|
|
||||||
|
messages = append(messages, content[:msgEnd])
|
||||||
|
content = strings.TrimSpace(content[msgEnd:])
|
||||||
|
}
|
||||||
|
|
||||||
|
return messages
|
||||||
|
}
|
||||||
|
|
||||||
|
// findLastUnclosedCodeBlock finds the last opening ``` that doesn't have a closing ```
|
||||||
|
// Returns the position of the opening ``` or -1 if all code blocks are complete
|
||||||
|
func findLastUnclosedCodeBlock(text string) int {
|
||||||
|
inCodeBlock := false
|
||||||
|
lastOpenIdx := -1
|
||||||
|
|
||||||
|
for i := 0; i < len(text); i++ {
|
||||||
|
if i+2 < len(text) && text[i] == '`' && text[i+1] == '`' && text[i+2] == '`' {
|
||||||
|
// Toggle code block state on each fence
|
||||||
|
if !inCodeBlock {
|
||||||
|
// Entering a code block: record this opening fence
|
||||||
|
lastOpenIdx = i
|
||||||
|
}
|
||||||
|
inCodeBlock = !inCodeBlock
|
||||||
|
i += 2
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if inCodeBlock {
|
||||||
|
return lastOpenIdx
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
|
// findNextClosingCodeBlock finds the next closing ``` starting from a position
|
||||||
|
// Returns the position after the closing ``` or -1 if not found
|
||||||
|
func findNextClosingCodeBlock(text string, startIdx int) int {
|
||||||
|
for i := startIdx; i < len(text); i++ {
|
||||||
|
if i+2 < len(text) && text[i] == '`' && text[i+1] == '`' && text[i+2] == '`' {
|
||||||
|
return i + 3
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
|
// findLastNewline finds the last newline character within the last N characters
|
||||||
|
// Returns the position of the newline or -1 if not found
|
||||||
|
func findLastNewline(s string, searchWindow int) int {
|
||||||
|
searchStart := len(s) - searchWindow
|
||||||
|
if searchStart < 0 {
|
||||||
|
searchStart = 0
|
||||||
|
}
|
||||||
|
for i := len(s) - 1; i >= searchStart; i-- {
|
||||||
|
if s[i] == '\n' {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
|
// findLastSpace finds the last space character within the last N characters
|
||||||
|
// Returns the position of the space or -1 if not found
|
||||||
|
func findLastSpace(s string, searchWindow int) int {
|
||||||
|
searchStart := len(s) - searchWindow
|
||||||
|
if searchStart < 0 {
|
||||||
|
searchStart = 0
|
||||||
|
}
|
||||||
|
for i := len(s) - 1; i >= searchStart; i-- {
|
||||||
|
if s[i] == ' ' || s[i] == '\t' {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
151
pkg/utils/message_test.go
Normal file
151
pkg/utils/message_test.go
Normal file
|
|
@ -0,0 +1,151 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSplitMessage(t *testing.T) {
|
||||||
|
longText := strings.Repeat("a", 2500)
|
||||||
|
longCode := "```go\n" + strings.Repeat("fmt.Println(\"hello\")\n", 100) + "```" // ~2100 chars
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
content string
|
||||||
|
maxLen int
|
||||||
|
expectChunks int // Check number of chunks
|
||||||
|
checkContent func(t *testing.T, chunks []string) // Custom validation
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "Empty message",
|
||||||
|
content: "",
|
||||||
|
maxLen: 2000,
|
||||||
|
expectChunks: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Short message fits in one chunk",
|
||||||
|
content: "Hello world",
|
||||||
|
maxLen: 2000,
|
||||||
|
expectChunks: 1,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Simple split regular text",
|
||||||
|
content: longText,
|
||||||
|
maxLen: 2000,
|
||||||
|
expectChunks: 2,
|
||||||
|
checkContent: func(t *testing.T, chunks []string) {
|
||||||
|
if len(chunks[0]) > 2000 {
|
||||||
|
t.Errorf("Chunk 0 too large: %d", len(chunks[0]))
|
||||||
|
}
|
||||||
|
if len(chunks[0])+len(chunks[1]) != len(longText) {
|
||||||
|
t.Errorf("Total length mismatch. Got %d, want %d", len(chunks[0])+len(chunks[1]), len(longText))
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Split at newline",
|
||||||
|
// 1750 chars then newline, then more chars.
|
||||||
|
// Dynamic buffer: 2000 / 10 = 200.
|
||||||
|
// Effective limit: 2000 - 200 = 1800.
|
||||||
|
// Split should happen at newline because it's at 1750 (< 1800).
|
||||||
|
// Total length must > 2000 to trigger split. 1750 + 1 + 300 = 2051.
|
||||||
|
content: strings.Repeat("a", 1750) + "\n" + strings.Repeat("b", 300),
|
||||||
|
maxLen: 2000,
|
||||||
|
expectChunks: 2,
|
||||||
|
checkContent: func(t *testing.T, chunks []string) {
|
||||||
|
if len(chunks[0]) != 1750 {
|
||||||
|
t.Errorf("Expected chunk 0 to be 1750 length (split at newline), got %d", len(chunks[0]))
|
||||||
|
}
|
||||||
|
if chunks[1] != strings.Repeat("b", 300) {
|
||||||
|
t.Errorf("Chunk 1 content mismatch. Len: %d", len(chunks[1]))
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Long code block split",
|
||||||
|
content: "Prefix\n" + longCode,
|
||||||
|
maxLen: 2000,
|
||||||
|
expectChunks: 2,
|
||||||
|
checkContent: func(t *testing.T, chunks []string) {
|
||||||
|
// Check that first chunk ends with closing fence
|
||||||
|
if !strings.HasSuffix(chunks[0], "\n```") {
|
||||||
|
t.Error("First chunk should end with injected closing fence")
|
||||||
|
}
|
||||||
|
// Check that second chunk starts with execution header
|
||||||
|
if !strings.HasPrefix(chunks[1], "```go") {
|
||||||
|
t.Error("Second chunk should start with injected code block header")
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Preserve Unicode characters",
|
||||||
|
content: strings.Repeat("\u4e16", 1000), // 3000 bytes
|
||||||
|
maxLen: 2000,
|
||||||
|
expectChunks: 2,
|
||||||
|
checkContent: func(t *testing.T, chunks []string) {
|
||||||
|
// Just verify we didn't panic and got valid strings.
|
||||||
|
// Go strings are UTF-8, if we split mid-rune it would be bad,
|
||||||
|
// but standard slicing might do that.
|
||||||
|
// Let's assume standard behavior is acceptable or check if it produces invalid rune?
|
||||||
|
if !strings.Contains(chunks[0], "\u4e16") {
|
||||||
|
t.Error("Chunk should contain unicode characters")
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
got := SplitMessage(tc.content, tc.maxLen)
|
||||||
|
|
||||||
|
if tc.expectChunks == 0 {
|
||||||
|
if len(got) != 0 {
|
||||||
|
t.Errorf("Expected 0 chunks, got %d", len(got))
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(got) != tc.expectChunks {
|
||||||
|
t.Errorf("Expected %d chunks, got %d", tc.expectChunks, len(got))
|
||||||
|
// Log sizes for debugging
|
||||||
|
for i, c := range got {
|
||||||
|
t.Logf("Chunk %d length: %d", i, len(c))
|
||||||
|
}
|
||||||
|
return // Stop further checks if count assumes specific split
|
||||||
|
}
|
||||||
|
|
||||||
|
if tc.checkContent != nil {
|
||||||
|
tc.checkContent(t, got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSplitMessage_CodeBlockIntegrity(t *testing.T) {
|
||||||
|
// Focused test for the core requirement: splitting inside a code block preserves syntax highlighting
|
||||||
|
|
||||||
|
// 60 chars total approximately
|
||||||
|
content := "```go\npackage main\n\nfunc main() {\n\tprintln(\"Hello\")\n}\n```"
|
||||||
|
maxLen := 40
|
||||||
|
|
||||||
|
chunks := SplitMessage(content, maxLen)
|
||||||
|
|
||||||
|
if len(chunks) != 2 {
|
||||||
|
t.Fatalf("Expected 2 chunks, got %d: %q", len(chunks), chunks)
|
||||||
|
}
|
||||||
|
|
||||||
|
// First chunk must end with "\n```"
|
||||||
|
if !strings.HasSuffix(chunks[0], "\n```") {
|
||||||
|
t.Errorf("First chunk should end with closing fence. Got: %q", chunks[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second chunk must start with the header "```go"
|
||||||
|
if !strings.HasPrefix(chunks[1], "```go") {
|
||||||
|
t.Errorf("Second chunk should start with code block header. Got: %q", chunks[1])
|
||||||
|
}
|
||||||
|
|
||||||
|
// First chunk should contain meaningful content
|
||||||
|
if len(chunks[0]) > 40 {
|
||||||
|
t.Errorf("First chunk exceeded maxLen: length %d", len(chunks[0]))
|
||||||
|
}
|
||||||
|
}
|
||||||
19
pkg/utils/skills.go
Normal file
19
pkg/utils/skills.go
Normal file
|
|
@ -0,0 +1,19 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ValidateSkillIdentifier validates that the given skill identifier (slug or registry name) is non-empty
|
||||||
|
// and does not contain path separators ("/", "\\") or ".." for security.
|
||||||
|
func ValidateSkillIdentifier(identifier string) error {
|
||||||
|
trimmed := strings.TrimSpace(identifier)
|
||||||
|
if trimmed == "" {
|
||||||
|
return fmt.Errorf("identifier is required and must be a non-empty string")
|
||||||
|
}
|
||||||
|
if strings.ContainsAny(trimmed, "/\\") || strings.Contains(trimmed, "..") {
|
||||||
|
return fmt.Errorf("identifier must not contain path separators or '..' to prevent directory traversal")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
@ -14,3 +14,12 @@ func Truncate(s string, maxLen int) string {
|
||||||
}
|
}
|
||||||
return string(runes[:maxLen-3]) + "..."
|
return string(runes[:maxLen-3]) + "..."
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DerefStr dereferences a pointer to a string and
|
||||||
|
// returns the value or a fallback if the pointer is nil.
|
||||||
|
func DerefStr(s *string, fallback string) string {
|
||||||
|
if s == nil {
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
return *s
|
||||||
|
}
|
||||||
|
|
|
||||||
120
pkg/utils/zip.go
Normal file
120
pkg/utils/zip.go
Normal file
|
|
@ -0,0 +1,120 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ExtractZipFile extracts a ZIP archive from disk to targetDir.
|
||||||
|
// It reads entries one at a time from disk, keeping memory usage minimal.
|
||||||
|
//
|
||||||
|
// Security: rejects path traversal attempts and symlinks.
|
||||||
|
func ExtractZipFile(zipPath string, targetDir string) error {
|
||||||
|
reader, err := zip.OpenReader(zipPath)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid ZIP: %w", err)
|
||||||
|
}
|
||||||
|
defer reader.Close()
|
||||||
|
|
||||||
|
logger.DebugCF("zip", "Extracting ZIP", map[string]interface{}{
|
||||||
|
"zip_path": zipPath,
|
||||||
|
"target_dir": targetDir,
|
||||||
|
"entries": len(reader.File),
|
||||||
|
})
|
||||||
|
|
||||||
|
if err := os.MkdirAll(targetDir, 0755); err != nil {
|
||||||
|
return fmt.Errorf("failed to create target dir: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, f := range reader.File {
|
||||||
|
// Path traversal protection.
|
||||||
|
cleanName := filepath.Clean(f.Name)
|
||||||
|
if strings.HasPrefix(cleanName, "..") || filepath.IsAbs(cleanName) {
|
||||||
|
return fmt.Errorf("zip entry has unsafe path: %q", f.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
destPath := filepath.Join(targetDir, cleanName)
|
||||||
|
|
||||||
|
// Double-check the resolved path is within target directory (defense-in-depth).
|
||||||
|
targetDirClean := filepath.Clean(targetDir)
|
||||||
|
if !strings.HasPrefix(filepath.Clean(destPath), targetDirClean+string(filepath.Separator)) && filepath.Clean(destPath) != targetDirClean {
|
||||||
|
return fmt.Errorf("zip entry escapes target dir: %q", f.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
mode := f.FileInfo().Mode()
|
||||||
|
|
||||||
|
// Reject any symlink.
|
||||||
|
if mode&os.ModeSymlink != 0 {
|
||||||
|
return fmt.Errorf("zip contains symlink %q; symlinks are not allowed", f.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
if f.FileInfo().IsDir() {
|
||||||
|
if err := os.MkdirAll(destPath, 0755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure parent directory exists.
|
||||||
|
if err := os.MkdirAll(filepath.Dir(destPath), 0755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := extractSingleFile(f, destPath); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractSingleFile extracts one zip.File entry to destPath, with a size check.
|
||||||
|
func extractSingleFile(f *zip.File, destPath string) error {
|
||||||
|
const maxFileSize = 5 * 1024 * 1024 // 5MB, adjust as appropriate
|
||||||
|
|
||||||
|
// Check the uncompressed size from the header, if available.
|
||||||
|
if f.UncompressedSize64 > maxFileSize {
|
||||||
|
return fmt.Errorf("zip entry %q is too large (%d bytes)", f.Name, f.UncompressedSize64)
|
||||||
|
}
|
||||||
|
|
||||||
|
rc, err := f.Open()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to open zip entry %q: %w", f.Name, err)
|
||||||
|
}
|
||||||
|
defer rc.Close()
|
||||||
|
|
||||||
|
outFile, err := os.Create(destPath)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to create file %q: %w", destPath, err)
|
||||||
|
}
|
||||||
|
// We don't return the close error via return, since it's not a named error return.
|
||||||
|
// Instead, we log to stderr and remove the partially written file as defensive cleanup.
|
||||||
|
defer func() {
|
||||||
|
if cerr := outFile.Close(); cerr != nil {
|
||||||
|
_ = os.Remove(destPath)
|
||||||
|
logger.ErrorCF("zip", "Failed to close file", map[string]interface{}{
|
||||||
|
"dest_path": destPath,
|
||||||
|
"error": cerr.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Streamed size check: prevent overruns and malicious/corrupt headers.
|
||||||
|
written, err := io.CopyN(outFile, rc, maxFileSize+1)
|
||||||
|
if err != nil && err != io.EOF {
|
||||||
|
_ = os.Remove(destPath)
|
||||||
|
return fmt.Errorf("failed to extract %q: %w", f.Name, err)
|
||||||
|
}
|
||||||
|
if written > maxFileSize {
|
||||||
|
_ = os.Remove(destPath)
|
||||||
|
return fmt.Errorf("zip entry %q exceeds max size (%d bytes)", f.Name, written)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue