Merge branch 'self_upgrade' of https://github.com/sky5454/picoclaw into self_upgrade
This commit is contained in:
commit
52d5765ff4
92 changed files with 5026 additions and 570 deletions
|
|
@ -61,6 +61,9 @@ linters:
|
|||
- usestdlibvars
|
||||
- usetesting
|
||||
settings:
|
||||
gomoddirectives:
|
||||
replace-allow-list:
|
||||
- github.com/bwmarrin/discordgo
|
||||
errcheck:
|
||||
check-type-assertions: true
|
||||
check-blank: true
|
||||
|
|
|
|||
|
|
@ -57,6 +57,8 @@
|
|||
|
||||
## 📢 Actualités
|
||||
|
||||
2026-03-31 📱 **Support Android !** PicoClaw fonctionne maintenant sur Android ! Téléchargez l'APK sur [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 publiée !** Refonte de l'architecture Agent (SubTurn, Hooks, Steering, EventBus), intégration WeChat/WeCom, renforcement de la sécurité (.security.yml, filtrage des données sensibles), nouveaux providers (AWS Bedrock, Azure, Xiaomi MiMo), et 35 corrections de bugs. PicoClaw a atteint **26K Stars** !
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 publiée !** Interface system tray (Windows & Linux), requête de statut des sous-agents (`spawn_status`), rechargement à chaud expérimental du Gateway, sécurisation Cron, et 2 correctifs de sécurité. PicoClaw a atteint **25K Stars** !
|
||||
|
|
@ -321,9 +323,9 @@ Suivez ensuite la section Terminal Launcher ci-dessous pour terminer la configur
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**Option 2 : Installation APK (bientôt disponible)**
|
||||
**Option 2 : Installation APK**
|
||||
|
||||
Un APK Android autonome avec WebUI intégré est en développement. Restez à l'écoute !
|
||||
Téléchargez l'APK depuis [picoclaw.io](https://picoclaw.io/download/) et installez-le directement. Pas besoin de Termux !
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher (pour les environnements à ressources limitées)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 Berita
|
||||
|
||||
2026-03-31 📱 **Dukungan Android!** PicoClaw sekarang berjalan di Android! Unduh APK di [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 Dirilis!** Perombakan arsitektur Agent (SubTurn, Hooks, Steering, EventBus), integrasi WeChat/WeCom, penguatan keamanan (.security.yml, penyaringan data sensitif), provider baru (AWS Bedrock, Azure, Xiaomi MiMo), dan 35 perbaikan bug. PicoClaw telah mencapai **26K Stars**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 Dirilis!** UI system tray (Windows & Linux), pelacakan status sub-agent (`spawn_status`), eksperimental Gateway hot-reload, gerbang keamanan Cron, dan 2 perbaikan keamanan. PicoClaw telah mencapai **25K Stars**!
|
||||
|
|
@ -318,9 +320,9 @@ Kemudian ikuti bagian Terminal Launcher di bawah untuk menyelesaikan konfigurasi
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**Opsi 2: Instal APK (segera hadir)**
|
||||
**Opsi 2: Instal APK**
|
||||
|
||||
APK Android mandiri dengan WebUI bawaan sedang dalam pengembangan. Pantau terus!
|
||||
Unduh APK dari [picoclaw.io](https://picoclaw.io/download/) dan instal langsung. Tanpa Termux!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher (untuk lingkungan dengan sumber daya terbatas)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 Novità
|
||||
|
||||
2026-03-31 📱 **Supporto Android!** PicoClaw ora funziona su Android! Scarica l'APK su [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 rilasciata!** Revisione dell'architettura Agent (SubTurn, Hooks, Steering, EventBus), integrazione WeChat/WeCom, rafforzamento della sicurezza (.security.yml, filtraggio dati sensibili), nuovi provider (AWS Bedrock, Azure, Xiaomi MiMo) e 35 correzioni di bug. PicoClaw raggiunge **26K Stars**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 rilasciata!** Interfaccia system tray (Windows & Linux), query sullo stato dei sub-agent (`spawn_status`), hot-reload sperimentale del Gateway, gate di sicurezza per Cron e 2 correzioni di sicurezza. PicoClaw raggiunge **25K Stars**!
|
||||
|
|
@ -318,9 +320,9 @@ Poi segui la sezione Terminal Launcher qui sotto per completare la configurazion
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**Opzione 2: APK Install (prossimamente)**
|
||||
**Opzione 2: Installazione APK**
|
||||
|
||||
Un APK Android standalone con WebUI integrato è in sviluppo. Resta sintonizzato!
|
||||
Scarica l'APK da [picoclaw.io](https://picoclaw.io/download/) e installa direttamente. Senza Termux!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher (per ambienti con risorse limitate)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 ニュース
|
||||
|
||||
2026-03-31 📱 **Android サポート!** PicoClawがAndroidで動作!APKは[picoclaw.io](https://picoclaw.io/download)からダウンロード
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 リリース!** Agent アーキテクチャ全面刷新(SubTurn、Hooks、Steering、EventBus)、WeChat/WeCom 統合、セキュリティ強化(.security.yml、機密データフィルタリング)、新プロバイダー(AWS Bedrock、Azure、Xiaomi MiMo)、35 件のバグ修正。PicoClaw **26K ⭐** 達成!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 リリース!** システムトレイ UI(Windows & Linux)、サブエージェントステータス追跡(`spawn_status`)、実験的 Gateway ホットリロード、cron セキュリティゲート、セキュリティ修正 2 件。PicoClaw **25K ⭐** 達成!
|
||||
|
|
@ -318,9 +320,9 @@ termux-chroot ./picoclaw onboard # chroot で標準的な Linux ファイル
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**オプション 2: APK インストール(近日公開)**
|
||||
**オプション 2: APK インストール**
|
||||
|
||||
内蔵 WebUI を備えたスタンドアロン Android APK を開発中です。お楽しみに!
|
||||
[picoclaw.io](https://picoclaw.io/download/) から APK をダウンロードして直接インストール。Termux 不要!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher(リソース制約環境向け)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 News
|
||||
|
||||
2026-03-31 📱 **Android Support!** PicoClaw now runs on Android! Download the APK at [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 Released!** Agent architecture overhaul (SubTurn, Hooks, Steering, EventBus), WeChat/WeCom integration, security hardening (.security.yml, sensitive data filtering), new providers (AWS Bedrock, Azure, Xiaomi MiMo), and 35 bug fixes. PicoClaw has reached **26K Stars**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 Released!** System tray UI (Windows & Linux), sub-agent status query (`spawn_status`), experimental Gateway hot-reload, Cron security gating, and 2 security fixes. PicoClaw has reached **25K Stars**!
|
||||
|
|
@ -318,9 +320,9 @@ Then follow the Terminal Launcher section below to complete configuration.
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**Option 2: APK Install (coming soon)**
|
||||
**Option 2: APK Install**
|
||||
|
||||
A standalone Android APK with built-in WebUI is in development. Stay tuned!
|
||||
Download the APK from [picoclaw.io](https://picoclaw.io/download/) and install directly. No Termux required!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher (for resource-constrained environments)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 Berita
|
||||
|
||||
2026-03-31 📱 **Sokongan Android!** PicoClaw sekarang berjalan di Android! Muat turun APK di [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 Dikeluarkan!** Penstrukturan semula seni bina Agent (SubTurn, Hooks, Steering, EventBus), integrasi WeChat/WeCom, penguatan keselamatan (.security.yml, penapisan data sensitif), penyedia baharu (AWS Bedrock, Azure, Xiaomi MiMo), dan 35 pembetulan pepijat. PicoClaw mencapai **26K Stars**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 Dikeluarkan!** UI dulang sistem (Windows & Linux), pertanyaan status sub-agent (`spawn_status`), muat semula panas Gateway eksperimental, kawalan keselamatan Cron, dan 2 pembetulan keselamatan. PicoClaw mencapai **25K Stars**!
|
||||
|
|
@ -315,9 +317,9 @@ Kemudian ikuti bahagian Pelancar Terminal di bawah untuk melengkapkan konfiguras
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw pada Termux" width="512">
|
||||
|
||||
**Pilihan 2: APK (akan datang)**
|
||||
**Pilihan 2: Pasang APK**
|
||||
|
||||
APK Android bebas dengan WebUI terbina dalam sedang dalam pembangunan. Nantikan!
|
||||
Muat turun APK dari [picoclaw.io](https://picoclaw.io/download/) dan pasang secara langsung. Tiada Termux diperlukan!
|
||||
|
||||
<details>
|
||||
<summary><b>Pelancar Terminal (untuk persekitaran terhad sumber)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 Novidades
|
||||
|
||||
2026-03-31 📱 **Suporte Android!** PicoClaw agora roda no Android! Baixe o APK em [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 Lançada!** Reformulação da arquitetura Agent (SubTurn, Hooks, Steering, EventBus), integração WeChat/WeCom, fortalecimento de segurança (.security.yml, filtragem de dados sensíveis), novos providers (AWS Bedrock, Azure, Xiaomi MiMo) e 35 correções de bugs. O PicoClaw atingiu **26K Stars**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 Lançada!** UI na bandeja do sistema (Windows e Linux), consulta de status de sub-agent (`spawn_status`), hot-reload experimental do Gateway, controle de segurança do Cron e 2 correções de segurança. O PicoClaw atingiu **25K Stars**!
|
||||
|
|
@ -318,9 +320,9 @@ Em seguida, siga a seção Terminal Launcher abaixo para concluir a configuraç
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**Opção 2: Instalação via APK (em breve)**
|
||||
**Opção 2: Instalação via APK**
|
||||
|
||||
Um APK Android independente com WebUI integrado está em desenvolvimento. Fique ligado!
|
||||
Baixe o APK de [picoclaw.io](https://picoclaw.io/download/) e instale diretamente. Sem necessidade de Termux!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher (para ambientes com recursos limitados)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 Tin tức
|
||||
|
||||
2026-03-31 📱 **Hỗ trợ Android!** PicoClaw giờ chạy trên Android! Tải APK tại [picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 đã phát hành!** Tái cấu trúc kiến trúc Agent (SubTurn, Hooks, Steering, EventBus), tích hợp WeChat/WeCom, tăng cường bảo mật (.security.yml, lọc dữ liệu nhạy cảm), provider mới (AWS Bedrock, Azure, Xiaomi MiMo) và 35 bản vá lỗi. PicoClaw đã đạt **26K Stars**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 đã phát hành!** Giao diện system tray (Windows & Linux), truy vấn trạng thái sub-agent (`spawn_status`), thử nghiệm Gateway hot-reload, bảo mật Cron, và 2 bản vá bảo mật. PicoClaw đã đạt **25K Stars**!
|
||||
|
|
@ -318,9 +320,9 @@ Sau đó làm theo phần Terminal Launcher bên dưới để hoàn tất cấu
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**Tùy chọn 2: Cài đặt APK (sắp ra mắt)**
|
||||
**Tùy chọn 2: Cài đặt APK**
|
||||
|
||||
Một APK Android độc lập với WebUI tích hợp đang được phát triển. Hãy đón chờ!
|
||||
Tải APK từ [picoclaw.io](https://picoclaw.io/download/) và cài đặt trực tiếp. Không cần Termux!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher (cho môi trường hạn chế tài nguyên)</b></summary>
|
||||
|
|
|
|||
|
|
@ -56,6 +56,8 @@
|
|||
|
||||
## 📢 新闻
|
||||
|
||||
2026-03-31 📱 **Android 支持!** PicoClaw 现可在 Android 上运行!APK 下载地址:[picoclaw.io](https://picoclaw.io/download)
|
||||
|
||||
2026-03-25 🚀 **v0.2.4 发布!** Agent 架构全面重构(SubTurn、Hook、Steering、EventBus)、微信/企业微信深度集成、安全体系升级(.security.yml、敏感数据过滤)、新增 Provider(AWS Bedrock、Azure、小米 MiMo),以及 35 项 Bug 修复。PicoClaw 已达 **26K ⭐**!
|
||||
|
||||
2026-03-17 🚀 **v0.2.3 发布!** 系统托盘 UI(Windows & Linux)、子 Agent 状态查询 (`spawn_status`)、实验性 Gateway 热重载、Cron 安全门控,以及 2 项安全修复。PicoClaw 已达 **25K ⭐**!
|
||||
|
|
@ -318,9 +320,9 @@ termux-chroot ./picoclaw onboard # chroot 提供标准 Linux 文件系统布
|
|||
|
||||
<img src="assets/termux.jpg" alt="PicoClaw on Termux" width="512">
|
||||
|
||||
**方式二:APK 安装(即将推出)**
|
||||
**方式二:APK 安装**
|
||||
|
||||
内置 WebUI 的独立 Android APK 正在开发中,敬请期待!
|
||||
从 [picoclaw.io](https://picoclaw.io/download/) 下载 APK 并直接安装,无需 Termux!
|
||||
|
||||
<details>
|
||||
<summary><b>Terminal Launcher(适用于资源受限环境)</b></summary>
|
||||
|
|
|
|||
|
|
@ -418,6 +418,9 @@
|
|||
"read_file": {
|
||||
"enabled": true
|
||||
},
|
||||
"send_tts": {
|
||||
"enabled": false
|
||||
},
|
||||
"spawn": {
|
||||
"enabled": true
|
||||
},
|
||||
|
|
|
|||
|
|
@ -13,18 +13,20 @@ Le canal Telegram utilise le long polling via l'API Bot Telegram pour une commun
|
|||
"enabled": true,
|
||||
"token": "123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
|
||||
"allow_from": ["123456789"],
|
||||
"proxy": ""
|
||||
"proxy": "",
|
||||
"use_markdown_v2": false
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| Champ | Type | Requis | Description |
|
||||
| ---------- | ------ | ------ | ------------------------------------------------------------------------ |
|
||||
| enabled | bool | Oui | Activer ou non le canal Telegram |
|
||||
| token | string | Oui | Token de l'API Bot Telegram |
|
||||
| allow_from | array | Non | Liste blanche d'identifiants utilisateur ; vide signifie tous les utilisateurs |
|
||||
| proxy | string | Non | URL du proxy pour se connecter à l'API Telegram (ex. http://127.0.0.1:7890) |
|
||||
| Champ | Type | Requis | Description |
|
||||
| --------------- | ------ | ------ | ------------------------------------------------------------------------ |
|
||||
| enabled | bool | Oui | Activer ou non le canal Telegram |
|
||||
| token | string | Oui | Token de l'API Bot Telegram |
|
||||
| allow_from | array | Non | Liste blanche d'identifiants utilisateur ; vide signifie tous les utilisateurs |
|
||||
| proxy | string | Non | URL du proxy pour se connecter à l'API Telegram (ex. http://127.0.0.1:7890) |
|
||||
| use_markdown_v2 | bool | Non | Activer le formatage Telegram MarkdownV2 |
|
||||
|
||||
## Configuration initiale
|
||||
|
||||
|
|
@ -33,3 +35,20 @@ Le canal Telegram utilise le long polling via l'API Bot Telegram pour une commun
|
|||
3. Obtenir le Token de l'API HTTP
|
||||
4. Renseigner le Token dans le fichier de configuration
|
||||
5. (Optionnel) Configurer `allow_from` pour restreindre les identifiants utilisateur autorisés à interagir (les IDs peuvent être obtenus via `@userinfobot`)
|
||||
|
||||
## Formatage avancées
|
||||
|
||||
Vous pouvez définir `use_markdown_v2: true` pour activer les options de formatage améliorées. Cela permet au bot d'utiliser toutes les fonctionnalités de Telegram MarkdownV2, y compris les styles imbriqués, les spoilers et les blocs de largeur fixe personnalisés.
|
||||
|
||||
```json
|
||||
{
|
||||
"channels": {
|
||||
"telegram": {
|
||||
"enabled": true,
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"use_markdown_v2": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
|
|
|||
|
|
@ -13,18 +13,20 @@ Telegram チャンネルは、Telegram Bot API を使用したロングポーリ
|
|||
"enabled": true,
|
||||
"token": "123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
|
||||
"allow_from": ["123456789"],
|
||||
"proxy": ""
|
||||
"proxy": "",
|
||||
"use_markdown_v2": false
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| フィールド | 型 | 必須 | 説明 |
|
||||
| ---------- | ------ | ---- | ----------------------------------------------------------------- |
|
||||
| enabled | bool | はい | Telegram チャンネルを有効にするかどうか |
|
||||
| token | string | はい | Telegram Bot API トークン |
|
||||
| allow_from | array | いいえ | 許可するユーザーIDのリスト。空の場合はすべてのユーザーを許可 |
|
||||
| proxy | string | いいえ | Telegram API への接続に使用するプロキシ URL (例: http://127.0.0.1:7890) |
|
||||
| フィールド | 型 | 必須 | 説明 |
|
||||
| --------------- | ------ | ---- | ----------------------------------------------------------------- |
|
||||
| enabled | bool | はい | Telegram チャンネルを有効にするかどうか |
|
||||
| token | string | はい | Telegram Bot API トークン |
|
||||
| allow_from | array | いいえ | 許可するユーザーIDのリスト。空の場合はすべてのユーザーを許可 |
|
||||
| proxy | string | いいえ | Telegram API への接続に使用するプロキシ URL (例: http://127.0.0.1:7890) |
|
||||
| use_markdown_v2 | bool | いいえ | Telegram MarkdownV2 フォーマットを有効にする |
|
||||
|
||||
## セットアップ手順
|
||||
|
||||
|
|
@ -33,3 +35,20 @@ Telegram チャンネルは、Telegram Bot API を使用したロングポーリ
|
|||
3. HTTP API トークンを取得する
|
||||
4. 設定ファイルにトークンを入力する
|
||||
5. (任意) `allow_from` を設定して、対話を許可するユーザー ID を制限する(ID は `@userinfobot` で取得可能)
|
||||
|
||||
## 高度なフォーマット
|
||||
|
||||
`use_markdown_v2: true` を設定することで、增强されたフォーマットオプションを有効にできます。これにより、ボットは Telegram MarkdownV2 の全機能(ネストされたスタイル、スポイラー、カスタム固定幅ブロックなど)を利用できます。
|
||||
|
||||
```json
|
||||
{
|
||||
"channels": {
|
||||
"telegram": {
|
||||
"enabled": true,
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"use_markdown_v2": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
|
|
|||
|
|
@ -13,18 +13,20 @@ The Telegram channel uses long polling via the Telegram Bot API for bot-based co
|
|||
"enabled": true,
|
||||
"token": "123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
|
||||
"allow_from": ["123456789"],
|
||||
"proxy": ""
|
||||
"proxy": "",
|
||||
"use_markdown_v2": false
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| Field | Type | Required | Description |
|
||||
| ---------- | ------ | -------- | ------------------------------------------------------------------ |
|
||||
| enabled | bool | Yes | Whether to enable the Telegram channel |
|
||||
| token | string | Yes | Telegram Bot API Token |
|
||||
| allow_from | array | No | Allowlist of user IDs; empty means all users are allowed |
|
||||
| proxy | string | No | Proxy URL for connecting to the Telegram API (e.g. http://127.0.0.1:7890) |
|
||||
| Field | Type | Required | Description |
|
||||
| ---------------- | ------ | -------- | ------------------------------------------------------------------ |
|
||||
| enabled | bool | Yes | Whether to enable the Telegram channel |
|
||||
| token | string | Yes | Telegram Bot API Token |
|
||||
| allow_from | array | No | Allowlist of user IDs; empty means all users are allowed |
|
||||
| proxy | string | No | Proxy URL for connecting to the Telegram API (e.g. http://127.0.0.1:7890) |
|
||||
| use_markdown_v2 | bool | No | Enable Telegram MarkdownV2 formatting |
|
||||
|
||||
## Setup
|
||||
|
||||
|
|
@ -53,3 +55,20 @@ Examples:
|
|||
/use git
|
||||
explain how to squash the last 3 commits
|
||||
```
|
||||
|
||||
## Advanced Formatting
|
||||
|
||||
You can set `use_markdown_v2: true` to enable enhanced formatting options. This allows the bot to utilize the full range of Telegram MarkdownV2 features, including nested styles, spoilers, and custom fixed-width blocks.
|
||||
|
||||
```json
|
||||
{
|
||||
"channels": {
|
||||
"telegram": {
|
||||
"enabled": true,
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"use_markdown_v2": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
|
|
|||
|
|
@ -13,18 +13,20 @@ O canal Telegram utiliza long polling via a API de Bot do Telegram para comunica
|
|||
"enabled": true,
|
||||
"token": "123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
|
||||
"allow_from": ["123456789"],
|
||||
"proxy": ""
|
||||
"proxy": "",
|
||||
"use_markdown_v2": false
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| Campo | Tipo | Obrigatório | Descrição |
|
||||
| ---------- | ------ | ----------- | -------------------------------------------------------------------------- |
|
||||
| enabled | bool | Sim | Se o canal Telegram deve ser habilitado |
|
||||
| token | string | Sim | Token da API de Bot do Telegram |
|
||||
| allow_from | array | Não | Lista de IDs de usuários permitidos; vazio significa todos os usuários |
|
||||
| proxy | string | Não | URL do proxy para conexão com a API do Telegram (ex. http://127.0.0.1:7890) |
|
||||
| Campo | Tipo | Obrigatório | Descrição |
|
||||
| --------------- | ------ | ----------- | -------------------------------------------------------------------------- |
|
||||
| enabled | bool | Sim | Se o canal Telegram deve ser habilitado |
|
||||
| token | string | Sim | Token da API de Bot do Telegram |
|
||||
| allow_from | array | Não | Lista de IDs de usuários permitidos; vazio significa todos os usuários |
|
||||
| proxy | string | Não | URL do proxy para conexão com a API do Telegram (ex. http://127.0.0.1:7890) |
|
||||
| use_markdown_v2 | bool | Não | Habilitar formatação Telegram MarkdownV2 |
|
||||
|
||||
## Configuração inicial
|
||||
|
||||
|
|
@ -33,3 +35,20 @@ O canal Telegram utiliza long polling via a API de Bot do Telegram para comunica
|
|||
3. Obtenha o Token da API HTTP
|
||||
4. Preencha o Token no arquivo de configuração
|
||||
5. (Opcional) Configure `allow_from` para restringir quais IDs de usuário podem interagir (os IDs podem ser obtidos via `@userinfobot`)
|
||||
|
||||
## Formatação Avançada
|
||||
|
||||
Você pode definir `use_markdown_v2: true` para habilitar opções de formatação aprimoradas. Isso permite que o bot utilize todos os recursos do Telegram MarkdownV2, incluindo estilos aninhados, spoilers e blocos de largura fixa personalizados.
|
||||
|
||||
```json
|
||||
{
|
||||
"channels": {
|
||||
"telegram": {
|
||||
"enabled": true,
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"use_markdown_v2": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
|
|
|||
|
|
@ -13,18 +13,20 @@ Kênh Telegram sử dụng long polling qua Telegram Bot API để giao tiếp d
|
|||
"enabled": true,
|
||||
"token": "123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
|
||||
"allow_from": ["123456789"],
|
||||
"proxy": ""
|
||||
"proxy": "",
|
||||
"use_markdown_v2": false
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| Trường | Kiểu | Bắt buộc | Mô tả |
|
||||
| ---------- | ------ | -------- | ------------------------------------------------------------------------ |
|
||||
| enabled | bool | Có | Có bật kênh Telegram hay không |
|
||||
| token | string | Có | Token API Bot Telegram |
|
||||
| allow_from | array | Không | Danh sách trắng ID người dùng; để trống nghĩa là cho phép tất cả |
|
||||
| proxy | string | Không | URL proxy để kết nối với Telegram API (ví dụ: http://127.0.0.1:7890) |
|
||||
| Trường | Kiểu | Bắt buộc | Mô tả |
|
||||
| -------------- | ------ | -------- | ------------------------------------------------------------------------ |
|
||||
| enabled | bool | Có | Có bật kênh Telegram hay không |
|
||||
| token | string | Có | Token API Bot Telegram |
|
||||
| allow_from | array | Không | Danh sách trắng ID người dùng; để trống nghĩa là cho phép tất cả |
|
||||
| proxy | string | Không | URL proxy để kết nối với Telegram API (ví dụ: http://127.0.0.1:7890) |
|
||||
| use_markdown_v2 | bool | Không | Bật định dạng Telegram MarkdownV2 |
|
||||
|
||||
## Hướng dẫn thiết lập
|
||||
|
||||
|
|
@ -33,3 +35,20 @@ Kênh Telegram sử dụng long polling qua Telegram Bot API để giao tiếp d
|
|||
3. Lấy Token API HTTP
|
||||
4. Điền Token vào file cấu hình
|
||||
5. (Tùy chọn) Cấu hình `allow_from` để giới hạn ID người dùng được phép tương tác (có thể lấy ID qua `@userinfobot`)
|
||||
|
||||
## Định dạng nâng cao
|
||||
|
||||
Bạn có thể đặt `use_markdown_v2: true` để bật các tùy chọn định dạng nâng cao. Điều này cho phép bot sử dụng toàn bộ các tính năng của Telegram MarkdownV2, bao gồm các kiểu lồng nhau, spoiler và các khối chiều rộng cố định tùy chỉnh.
|
||||
|
||||
```json
|
||||
{
|
||||
"channels": {
|
||||
"telegram": {
|
||||
"enabled": true,
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"use_markdown_v2": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
|
|
|||
|
|
@ -13,18 +13,20 @@ Telegram Channel 通过 Telegram 机器人 API 使用长轮询实现基于机器
|
|||
"enabled": true,
|
||||
"token": "123456789:ABCdefGHIjklMNOpqrsTUVwxyz",
|
||||
"allow_from": ["123456789"],
|
||||
"proxy": ""
|
||||
"proxy": "",
|
||||
"use_markdown_v2": false
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| 字段 | 类型 | 必填 | 描述 |
|
||||
| ---------- | ------ | ---- | --------------------------------------------------------- |
|
||||
| enabled | bool | 是 | 是否启用 Telegram 频道 |
|
||||
| token | string | 是 | Telegram 机器人 API Token |
|
||||
| allow_from | array | 否 | 用户ID白名单,空表示允许所有用户 |
|
||||
| proxy | string | 否 | 连接 Telegram API 的代理 URL (例如 http://127.0.0.1:7890) |
|
||||
| 字段 | 类型 | 必填 | 描述 |
|
||||
| ---------------- | ------ | ---- | --------------------------------------------------------- |
|
||||
| enabled | bool | 是 | 是否启用 Telegram 频道 |
|
||||
| token | string | 是 | Telegram 机器人 API Token |
|
||||
| allow_from | array | 否 | 用户ID白名单,空表示允许所有用户 |
|
||||
| proxy | string | 否 | 连接 Telegram API 的代理 URL (例如 http://127.0.0.1:7890) |
|
||||
| use_markdown_v2 | bool | 否 | 启用 Telegram MarkdownV2 格式化 |
|
||||
|
||||
## 设置流程
|
||||
|
||||
|
|
@ -50,6 +52,23 @@ Telegram 会在启动时自动注册 PicoClaw 的顶级 Bot 命令,包括 `/st
|
|||
```text
|
||||
/list skills
|
||||
/use git explain how to squash the last 3 commits
|
||||
/use italiapersonalfinance
|
||||
dammi le ultime news
|
||||
/use git
|
||||
explain how to squash the last 3 commits
|
||||
```
|
||||
|
||||
## 高级格式化
|
||||
|
||||
您可以设置 `use_markdown_v2: true` 来启用增强的格式化选项。这允许机器人使用 Telegram MarkdownV2 的全部功能,包括嵌套样式、剧透和自定义等宽代码块。
|
||||
|
||||
```json
|
||||
{
|
||||
"channels": {
|
||||
"telegram": {
|
||||
"enabled": true,
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"use_markdown_v2": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
|
|
|||
6
go.mod
6
go.mod
|
|
@ -28,6 +28,8 @@ require (
|
|||
github.com/mymmrac/telego v1.7.0
|
||||
github.com/open-dingtalk/dingtalk-stream-sdk-go v0.9.1
|
||||
github.com/openai/openai-go/v3 v3.22.0
|
||||
github.com/pion/rtp v1.8.7
|
||||
github.com/pion/webrtc/v3 v3.3.6
|
||||
github.com/rivo/tview v0.42.0
|
||||
github.com/rs/zerolog v1.34.0
|
||||
github.com/slack-go/slack v0.17.3
|
||||
|
|
@ -63,6 +65,7 @@ require (
|
|||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.9 // indirect
|
||||
github.com/aws/smithy-go v1.24.2 // indirect
|
||||
github.com/beeper/argo-go v1.1.2 // indirect
|
||||
github.com/cloudflare/circl v1.6.3 // indirect
|
||||
github.com/coder/websocket v1.8.14 // indirect
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
|
|
@ -78,6 +81,7 @@ require (
|
|||
github.com/mattn/go-sqlite3 v1.14.34 // indirect
|
||||
github.com/ncruces/go-strftime v1.0.0 // indirect
|
||||
github.com/petermattis/goid v0.0.0-20260226131333-17d1149c6ac6 // indirect
|
||||
github.com/pion/randutil v0.1.0 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
github.com/rivo/uniseg v0.4.7 // indirect
|
||||
|
|
@ -125,3 +129,5 @@ require (
|
|||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.42.0
|
||||
)
|
||||
|
||||
replace github.com/bwmarrin/discordgo => github.com/yeongaori/discordgo-fork v0.0.0-20260319072544-e8e546f5d532
|
||||
|
|
|
|||
12
go.sum
12
go.sum
|
|
@ -55,8 +55,6 @@ github.com/aws/smithy-go v1.24.2 h1:FzA3bu/nt/vDvmnkg+R8Xl46gmzEDam6mZ1hzmwXFng=
|
|||
github.com/aws/smithy-go v1.24.2/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc=
|
||||
github.com/beeper/argo-go v1.1.2 h1:UQI2G8F+NLfGTOmTUI0254pGKx/HUU/etbUGTJv91Fs=
|
||||
github.com/beeper/argo-go v1.1.2/go.mod h1:M+LJAnyowKVQ6Rdj6XYGEn+qcVFkb3R/MUpqkGR0hM4=
|
||||
github.com/bwmarrin/discordgo v0.29.0 h1:FmWeXFaKUwrcL3Cx65c20bTRW+vOb6k8AnaP+EgjDno=
|
||||
github.com/bwmarrin/discordgo v0.29.0/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY=
|
||||
github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M=
|
||||
github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM=
|
||||
github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE=
|
||||
|
|
@ -67,6 +65,8 @@ github.com/caarlos0/env/v11 v11.4.0 h1:Kcb6t5kIIr4XkoQC9AF2j+8E1Jsrl3Wz/hhm1LtoG
|
|||
github.com/caarlos0/env/v11 v11.4.0/go.mod h1:qupehSf/Y0TUTsxKywqRt/vJjN5nz6vauiYEUUr8P4U=
|
||||
github.com/cespare/xxhash/v2 v2.1.2/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/cloudflare/circl v1.6.3 h1:9GPOhQGF9MCYUeXyMYlqTR6a5gTrgR/fBLXvUgtVcg8=
|
||||
github.com/cloudflare/circl v1.6.3/go.mod h1:2eXP6Qfat4O/Yhh8BznvKnJ+uzEoTQ6jVKJRn81BiS4=
|
||||
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
|
||||
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
|
||||
github.com/coder/websocket v1.8.14 h1:9L0p0iKiNOibykf283eHkKUHHrpG7f65OE3BhhO7v9g=
|
||||
|
|
@ -208,6 +208,12 @@ github.com/openai/openai-go/v3 v3.22.0 h1:6MEoNoV8sbjOVmXdvhmuX3BjVbVdcExbVyGixi
|
|||
github.com/openai/openai-go/v3 v3.22.0/go.mod h1:cdufnVK14cWcT9qA1rRtrXx4FTRsgbDPW7Ia7SS5cZo=
|
||||
github.com/petermattis/goid v0.0.0-20260226131333-17d1149c6ac6 h1:rh2lKw/P/EqHa724vYH2+VVQ1YnW4u6EOXl0PMAovZE=
|
||||
github.com/petermattis/goid v0.0.0-20260226131333-17d1149c6ac6/go.mod h1:pxMtw7cyUw6B2bRH0ZBANSPg+AoSud1I1iyJHI69jH4=
|
||||
github.com/pion/randutil v0.1.0 h1:CFG1UdESneORglEsnimhUjf33Rwjubwj6xfiOXBa3mA=
|
||||
github.com/pion/randutil v0.1.0/go.mod h1:XcJrSMMbbMRhASFVOlj/5hQial/Y8oH/HVo7TBZq+j8=
|
||||
github.com/pion/rtp v1.8.7 h1:qslKkG8qxvQ7hqaxkmL7Pl0XcUm+/Er7nMnu6Vq+ZxM=
|
||||
github.com/pion/rtp v1.8.7/go.mod h1:pBGHaFt/yW7bf1jjWAoUjpSNoDnw98KTMg+jWWvziqU=
|
||||
github.com/pion/webrtc/v3 v3.3.6 h1:7XAh4RPtlY1Vul6/GmZrv7z+NnxKA6If0KStXBI2ZLE=
|
||||
github.com/pion/webrtc/v3 v3.3.6/go.mod h1:zyN7th4mZpV27eXybfR/cnUf3J2DRy8zw/mdjD9JTNM=
|
||||
github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA=
|
||||
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
|
|
@ -277,6 +283,8 @@ github.com/vektah/gqlparser/v2 v2.5.27 h1:RHPD3JOplpk5mP5JGX8RKZkt2/Vwj/PZv0HxTd
|
|||
github.com/vektah/gqlparser/v2 v2.5.27/go.mod h1:D1/VCZtV3LPnQrcPBeR/q5jkSQIPti0uYCP/RI0gIeo=
|
||||
github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU=
|
||||
github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E=
|
||||
github.com/yeongaori/discordgo-fork v0.0.0-20260319072544-e8e546f5d532 h1:gxFHYeUDGziRb0zXYEqBFohC+NJbIW9L0tddaXMWr2o=
|
||||
github.com/yeongaori/discordgo-fork v0.0.0-20260319072544-e8e546f5d532/go.mod h1:A0FcMFJKJ9fRjgSuZ2o+pIQ6mPS81SVuiLN2vYTa7Ao=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4=
|
||||
github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
|
||||
|
|
|
|||
|
|
@ -90,14 +90,29 @@ func findSafeBoundary(history []providers.Message, targetIndex int) int {
|
|||
// including Content, ReasoningContent, ToolCalls arguments, ToolCallID
|
||||
// metadata, and Media items. Uses a heuristic of 2.5 characters per token.
|
||||
func estimateMessageTokens(msg providers.Message) int {
|
||||
chars := utf8.RuneCountInString(msg.Content)
|
||||
contentChars := utf8.RuneCountInString(msg.Content)
|
||||
|
||||
// ReasoningContent (extended thinking / chain-of-thought) can be
|
||||
// substantial and is stored in session history via AddFullMessage.
|
||||
if msg.ReasoningContent != "" {
|
||||
chars += utf8.RuneCountInString(msg.ReasoningContent)
|
||||
// SystemParts are structured system blocks used for cache-aware adapters.
|
||||
// They carry the same content as Content, but in multiple blocks.
|
||||
// We estimate them as an alternative representation, not additive.
|
||||
systemPartsChars := 0
|
||||
if len(msg.SystemParts) > 0 {
|
||||
for _, part := range msg.SystemParts {
|
||||
systemPartsChars += utf8.RuneCountInString(part.Text)
|
||||
}
|
||||
// Per-part overhead for JSON structure (type, text, cache_control).
|
||||
const perPartOverhead = 20
|
||||
systemPartsChars += len(msg.SystemParts) * perPartOverhead
|
||||
}
|
||||
|
||||
// Use the larger of the two representations to stay conservative.
|
||||
chars := contentChars
|
||||
if systemPartsChars > chars {
|
||||
chars = systemPartsChars
|
||||
}
|
||||
|
||||
chars += utf8.RuneCountInString(msg.ReasoningContent)
|
||||
|
||||
for _, tc := range msg.ToolCalls {
|
||||
chars += len(tc.ID) + len(tc.Type)
|
||||
if tc.Function != nil {
|
||||
|
|
|
|||
|
|
@ -529,6 +529,26 @@ func TestEstimateMessageTokens_MediaItems(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestEstimateMessageTokens_SystemParts(t *testing.T) {
|
||||
plain := providers.Message{Role: "system", Content: "instructions"}
|
||||
withParts := providers.Message{
|
||||
Role: "system",
|
||||
Content: "instructions",
|
||||
SystemParts: []providers.ContentBlock{
|
||||
{Type: "text", Text: "some more system context"},
|
||||
{Type: "text", Text: "even more cached blocks"},
|
||||
},
|
||||
}
|
||||
|
||||
plainTokens := estimateMessageTokens(plain)
|
||||
partsTokens := estimateMessageTokens(withParts)
|
||||
|
||||
if partsTokens <= plainTokens {
|
||||
t.Errorf("system message with SystemParts (%d) should exceed plain message (%d)",
|
||||
partsTokens, plainTokens)
|
||||
}
|
||||
}
|
||||
|
||||
// --- estimateToolDefsTokens tests ---
|
||||
|
||||
func TestEstimateToolDefsTokens(t *testing.T) {
|
||||
|
|
|
|||
|
|
@ -18,6 +18,8 @@ import (
|
|||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio/asr"
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/commands"
|
||||
|
|
@ -31,7 +33,6 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/state"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
"github.com/sipeed/picoclaw/pkg/voice"
|
||||
)
|
||||
|
||||
type AgentLoop struct {
|
||||
|
|
@ -51,7 +52,7 @@ type AgentLoop struct {
|
|||
fallback *providers.FallbackChain
|
||||
channelManager *channels.Manager
|
||||
mediaStore media.MediaStore
|
||||
transcriber voice.Transcriber
|
||||
transcriber asr.Transcriber
|
||||
cmdRegistry *commands.Registry
|
||||
mcp mcpRuntime
|
||||
hookRuntime hookRuntime
|
||||
|
|
@ -159,6 +160,13 @@ func registerSharedTools(
|
|||
provider providers.LLMProvider,
|
||||
) {
|
||||
allowReadPaths := buildAllowReadPatterns(cfg)
|
||||
var ttsProvider tts.TTSProvider
|
||||
if cfg.Tools.IsToolEnabled("send_tts") {
|
||||
ttsProvider = tts.DetectTTS(cfg)
|
||||
if ttsProvider == nil {
|
||||
logger.WarnCF("voice-tts", "send_tts enabled but no TTS provider configured", nil)
|
||||
}
|
||||
}
|
||||
|
||||
for _, agentID := range registry.ListAgentIDs() {
|
||||
agent, ok := registry.GetAgent(agentID)
|
||||
|
|
@ -269,6 +277,10 @@ func registerSharedTools(
|
|||
agent.Tools.Register(sendFileTool)
|
||||
}
|
||||
|
||||
if ttsProvider != nil {
|
||||
agent.Tools.Register(tools.NewSendTTSTool(ttsProvider, nil))
|
||||
}
|
||||
|
||||
// Skill discovery and installation tools
|
||||
skills_enabled := cfg.Tools.IsToolEnabled("skills")
|
||||
find_skills_enable := cfg.Tools.IsToolEnabled("find_skills")
|
||||
|
|
@ -1059,10 +1071,15 @@ func (al *AgentLoop) SetMediaStore(s media.MediaStore) {
|
|||
agent.Tools.SetMediaStore(s)
|
||||
}
|
||||
}
|
||||
registry.ForEachTool("send_tts", func(t tools.Tool) {
|
||||
if st, ok := t.(*tools.SendTTSTool); ok {
|
||||
st.SetMediaStore(s)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// SetTranscriber injects a voice transcriber for agent-level audio transcription.
|
||||
func (al *AgentLoop) SetTranscriber(t voice.Transcriber) {
|
||||
func (al *AgentLoop) SetTranscriber(t asr.Transcriber) {
|
||||
al.transcriber = t
|
||||
}
|
||||
|
||||
|
|
@ -1083,19 +1100,23 @@ func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.Inbou
|
|||
|
||||
// Transcribe each audio media ref in order.
|
||||
var transcriptions []string
|
||||
var keptMedia []string
|
||||
for _, ref := range msg.Media {
|
||||
path, meta, err := al.mediaStore.ResolveWithMeta(ref)
|
||||
if err != nil {
|
||||
logger.WarnCF("voice", "Failed to resolve media ref", map[string]any{"ref": ref, "error": err})
|
||||
keptMedia = append(keptMedia, ref)
|
||||
continue
|
||||
}
|
||||
if !utils.IsAudioFile(meta.Filename, meta.ContentType) {
|
||||
keptMedia = append(keptMedia, ref)
|
||||
continue
|
||||
}
|
||||
result, err := al.transcriber.Transcribe(ctx, path)
|
||||
if err != nil {
|
||||
logger.WarnCF("voice", "Transcription failed", map[string]any{"ref": ref, "error": err})
|
||||
transcriptions = append(transcriptions, "")
|
||||
keptMedia = append(keptMedia, ref)
|
||||
continue
|
||||
}
|
||||
transcriptions = append(transcriptions, result.Text)
|
||||
|
|
@ -1115,15 +1136,21 @@ func (al *AgentLoop) transcribeAudioInMessage(ctx context.Context, msg bus.Inbou
|
|||
}
|
||||
text := transcriptions[idx]
|
||||
idx++
|
||||
if text == "" {
|
||||
return match
|
||||
}
|
||||
return "[voice: " + text + "]"
|
||||
})
|
||||
|
||||
// Append any remaining transcriptions not matched by an annotation.
|
||||
for ; idx < len(transcriptions); idx++ {
|
||||
newContent += "\n[voice: " + transcriptions[idx] + "]"
|
||||
if transcriptions[idx] != "" {
|
||||
newContent += "\n[voice: " + transcriptions[idx] + "]"
|
||||
}
|
||||
}
|
||||
|
||||
msg.Content = newContent
|
||||
msg.Media = keptMedia
|
||||
return msg, true
|
||||
}
|
||||
|
||||
|
|
@ -2464,6 +2491,28 @@ turnLoop:
|
|||
if toolResult == nil {
|
||||
toolResult = tools.ErrorResult("hook returned nil tool result")
|
||||
}
|
||||
|
||||
// Send ForUser if not silent and has content.
|
||||
// For ResponseHandled tools, send regardless of SendResponse setting,
|
||||
// since they've already handled the response (e.g., send_tts, send_file).
|
||||
shouldSendForUser := !toolResult.Silent && toolResult.ForUser != "" &&
|
||||
(ts.opts.SendResponse || toolResult.ResponseHandled)
|
||||
if shouldSendForUser {
|
||||
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Content: toolResult.ForUser,
|
||||
Metadata: map[string]string{
|
||||
"is_tool_call": "true",
|
||||
},
|
||||
})
|
||||
logger.DebugCF("agent", "Sent tool result to user",
|
||||
map[string]any{
|
||||
"tool": toolName,
|
||||
"content_len": len(toolResult.ForUser),
|
||||
})
|
||||
}
|
||||
|
||||
if len(toolResult.Media) > 0 && toolResult.ResponseHandled {
|
||||
parts := make([]bus.MediaPart, 0, len(toolResult.Media))
|
||||
for _, ref := range toolResult.Media {
|
||||
|
|
@ -2509,19 +2558,6 @@ turnLoop:
|
|||
allResponsesHandled = false
|
||||
}
|
||||
|
||||
if !toolResult.Silent && toolResult.ForUser != "" && ts.opts.SendResponse {
|
||||
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Content: toolResult.ForUser,
|
||||
})
|
||||
logger.DebugCF("agent", "Sent tool result to user",
|
||||
map[string]any{
|
||||
"tool": toolName,
|
||||
"content_len": len(toolResult.ForUser),
|
||||
})
|
||||
}
|
||||
|
||||
contentForLLM := toolResult.ContentForLLM()
|
||||
|
||||
// Filter sensitive data (API keys, tokens, secrets) before sending to LLM
|
||||
|
|
|
|||
166
pkg/audio/asr/README.md
Normal file
166
pkg/audio/asr/README.md
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
# ASR (Automatic Speech Recognition)
|
||||
|
||||
This package handles speech-to-text for PicoClaw voice input.
|
||||
|
||||
If you are new to ASR setup, the simplest mental model is:
|
||||
|
||||
1. Add one or more ASR-capable entries to `model_list`.
|
||||
2. Point `voice.model_name` at the one you want to use.
|
||||
3. Put the API key in `.security.yml`.
|
||||
|
||||
## Quick Recommendation
|
||||
|
||||
For most new users, start with one of these:
|
||||
|
||||
| Provider | Example model | Why start here |
|
||||
| --- | --- | --- |
|
||||
| [Groq](https://console.groq.com/keys) | `groq/whisper-large-v3-turbo` | Fast Whisper-style transcription and a straightforward OpenAI-compatible API. Groq currently advertises a free tier plan for 2000 reqs/day. |
|
||||
| [ElevenLabs](https://elevenlabs.io/pricing) | `elevenlabs/scribe_v1` | Easy setup and strong speech-to-text quality. ElevenLabs currently advertises a free plan that includes speech-to-text usage. |
|
||||
|
||||
Pricing and free-plan limits can change, so check the linked pricing pages before depending on them in production.
|
||||
|
||||
## How ASR Configuration Works
|
||||
|
||||
PicoClaw does not keep ASR API keys inside the `voice` section.
|
||||
|
||||
Instead:
|
||||
|
||||
- `voice.model_name` chooses a named entry from `model_list`.
|
||||
- The matching `model_list` entry describes the actual provider and model.
|
||||
- `.security.yml` stores the API key for that named model entry.
|
||||
|
||||
This is the recommended pattern because it is explicit, reusable, and consistent with the rest of PicoClaw's model configuration.
|
||||
|
||||
## Recommended Setup
|
||||
|
||||
### Option A: Groq Whisper
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"model_name": "groq-asr",
|
||||
"echo_transcription": true
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "groq-asr",
|
||||
"model": "groq/whisper-large-v3-turbo"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
groq-asr:
|
||||
api_keys:
|
||||
- "gsk_your_groq_key"
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- You can omit `api_base` and PicoClaw will use Groq's default API base automatically.
|
||||
- If you set `api_base` manually for Groq Whisper, both of these forms work:
|
||||
- `https://api.groq.com/openai/v1`
|
||||
- `https://api.groq.com/openai/v1/audio/transcriptions`
|
||||
- Any OpenAI-compatible Whisper model name containing `whisper` can use the Whisper transcription path, not only `whisper-large-v3-turbo`.
|
||||
|
||||
### Option B: ElevenLabs
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"model_name": "elevenlabs-asr",
|
||||
"echo_transcription": true
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "elevenlabs-asr",
|
||||
"model": "elevenlabs/scribe_v1"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
elevenlabs-asr:
|
||||
api_keys:
|
||||
- "sk-elevenlabs-your-key"
|
||||
```
|
||||
|
||||
### Option C: OpenAI Whisper
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"model_name": "openai-asr"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "openai-asr",
|
||||
"model": "openai/whisper-1"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
openai-asr:
|
||||
api_keys:
|
||||
- "sk-openai-your-key"
|
||||
```
|
||||
|
||||
## Other ASR-Capable Model Types
|
||||
|
||||
PicoClaw currently supports three main ASR routes:
|
||||
|
||||
| Route | Example models | Behavior |
|
||||
| --- | --- | --- |
|
||||
| ElevenLabs ASR | `elevenlabs/scribe_v1` | Uses the ElevenLabs transcription API. |
|
||||
| Whisper endpoint models | `openai/whisper-1`, `groq/whisper-large-v3` | Uses an OpenAI-compatible `/audio/transcriptions` endpoint. |
|
||||
| Audio-capable chat models **(Under construction)** | `openai/gpt-4o-audio-preview`, `gemini/gemini-2.5-flash` | Sends audio to a multimodal chat model and asks it to transcribe. |
|
||||
|
||||
If you are unsure which one to pick, choose Groq Whisper or ElevenLabs first.
|
||||
|
||||
## How PicoClaw Chooses a Transcriber
|
||||
|
||||
`DetectTranscriber` resolves ASR in this order:
|
||||
|
||||
1. **Preferred path**: resolve `voice.model_name` against `model_list`.
|
||||
2. If that resolved model is:
|
||||
- `elevenlabs/...`, PicoClaw uses the ElevenLabs transcriber.
|
||||
- an OpenAI-compatible Whisper model, PicoClaw uses the Whisper transcriber.
|
||||
- an audio-capable chat model, PicoClaw uses `AudioModelTranscriber`.
|
||||
3. **Fallback path**: if `voice.model_name` is not set, PicoClaw performs a compatibility scan through `model_list` for legacy auto-detected ASR entries.
|
||||
|
||||
Fallback scanning exists for backward compatibility. New configurations should set `voice.model_name` explicitly.
|
||||
|
||||
## Common Mistakes
|
||||
|
||||
- Defining an ASR model in `model_list` but forgetting to set `voice.model_name`.
|
||||
- Putting the API key in `voice` instead of `.security.yml`.
|
||||
- Using a non-ASR model and expecting Whisper-style transcription behavior.
|
||||
- Setting a custom `api_base` that points to the wrong provider endpoint.
|
||||
|
||||
## Minimal Checklist
|
||||
|
||||
Before testing voice input, make sure:
|
||||
|
||||
- `voice.model_name` matches a `model_list[].model_name`.
|
||||
- The matching `.security.yml` entry contains a valid API key.
|
||||
- The selected model is actually ASR-capable.
|
||||
- Voice input is enabled for the channel you are using.
|
||||
166
pkg/audio/asr/README_zh.md
Normal file
166
pkg/audio/asr/README_zh.md
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
# ASR(自动语音识别)
|
||||
|
||||
这个目录负责 PicoClaw 的语音转文字能力。
|
||||
|
||||
如果你是第一次配置 ASR,可以参考如下步骤:
|
||||
|
||||
1. 在 `model_list` 里添加一个或多个支持 ASR 的模型条目。
|
||||
2. 用 `voice.model_name` 指向你想使用的那个条目。
|
||||
3. 在 `.security.yml` 里配置对应的 API Key。
|
||||
|
||||
## 快速推荐
|
||||
|
||||
对于大多数新用户,建议先从下面两种开始:
|
||||
|
||||
| 提供商 | 示例模型 | 推荐理由 |
|
||||
| --- | --- | --- |
|
||||
| [Groq](https://console.groq.com/keys) | `groq/whisper-large-v3-turbo` | Whisper 风格转录速度快,并且提供 OpenAI 兼容接口,配置比较直接。Groq 目前官方提供2000请求每日的免费套餐。 |
|
||||
| [ElevenLabs](https://elevenlabs.io/pricing) | `elevenlabs/scribe_v1` | 上手简单,语音转文字质量也不错。ElevenLabs 目前官方免费套餐包含 STT 用量。 |
|
||||
|
||||
价格和免费额度可能会变化,正式使用前请以官网定价页为准。
|
||||
|
||||
## ASR 配置是如何工作的
|
||||
|
||||
PicoClaw 不会把 ASR 的 API Key 放在 `voice` 配置里。
|
||||
|
||||
推荐的方式是:
|
||||
|
||||
- `voice.model_name` 用来选择 `model_list` 里的某个命名模型。
|
||||
- `model_list` 条目描述真实的提供商和模型。
|
||||
- `.security.yml` 负责保存该模型条目的 API Key。
|
||||
|
||||
这种方式更明确、更安全,也和 PicoClaw 其他模型配置方式保持一致。
|
||||
|
||||
## 推荐配置方式
|
||||
|
||||
### 方案 A:Groq Whisper
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"model_name": "groq-asr",
|
||||
"echo_transcription": true
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "groq-asr",
|
||||
"model": "groq/whisper-large-v3-turbo"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
groq-asr:
|
||||
api_keys:
|
||||
- "gsk_your_groq_key"
|
||||
```
|
||||
|
||||
说明:
|
||||
|
||||
- 你可以不写 `api_base`,PicoClaw 会自动使用 Groq 默认接口地址。
|
||||
- 如果你手动设置 Groq Whisper 的 `api_base`,下面两种写法都可以:
|
||||
- `https://api.groq.com/openai/v1`
|
||||
- `https://api.groq.com/openai/v1/audio/transcriptions`
|
||||
- 只要是 OpenAI 兼容、并且模型名里包含 `whisper` 的模型,都可以走 Whisper 转录路径,不仅限于 `whisper-large-v3-turbo`。
|
||||
|
||||
### 方案 B:ElevenLabs
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"model_name": "elevenlabs-asr",
|
||||
"echo_transcription": true
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "elevenlabs-asr",
|
||||
"model": "elevenlabs/scribe_v1"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
elevenlabs-asr:
|
||||
api_keys:
|
||||
- "sk-elevenlabs-your-key"
|
||||
```
|
||||
|
||||
### 方案 C:OpenAI Whisper
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"model_name": "openai-asr"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "openai-asr",
|
||||
"model": "openai/whisper-1"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
openai-asr:
|
||||
api_keys:
|
||||
- "sk-openai-your-key"
|
||||
```
|
||||
|
||||
## 其他支持 ASR 的模型类型
|
||||
|
||||
PicoClaw 目前主要支持三种 ASR 路径:
|
||||
|
||||
| 路径 | 示例模型 | 行为说明 |
|
||||
| --- | --- | --- |
|
||||
| ElevenLabs ASR | `elevenlabs/scribe_v1` | 使用 ElevenLabs 的语音转录接口。 |
|
||||
| Whisper 接口模型 | `openai/whisper-1`、`groq/whisper-large-v3` | 使用 OpenAI 兼容的 `/audio/transcriptions` 接口。 |
|
||||
| 支持音频的聊天模型 **(重构中)** | `openai/gpt-4o-audio-preview`、`gemini/gemini-2.5-flash` | 把音频发给多模态聊天模型,并要求它返回转录结果。 |
|
||||
|
||||
如果你不确定该选哪种,建议优先使用 Groq Whisper 或 ElevenLabs。
|
||||
|
||||
## PicoClaw 如何选择转录器
|
||||
|
||||
`DetectTranscriber` 会按下面顺序选择 ASR:
|
||||
|
||||
1. **首选路径**:根据 `voice.model_name` 在 `model_list` 中找到对应模型。
|
||||
2. 如果找到的模型属于以下类型:
|
||||
- `elevenlabs/...`,则使用 ElevenLabs transcriber。
|
||||
- OpenAI 兼容的 Whisper 模型,则使用 Whisper transcriber。
|
||||
- 支持音频输入的聊天模型,则使用 `AudioModelTranscriber`。
|
||||
3. **回退路径**:如果没有设置 `voice.model_name`,PicoClaw 会为了兼容旧配置,扫描 `model_list` 中可自动识别的 ASR 条目。
|
||||
|
||||
回退扫描只是为了兼容旧行为。新配置建议始终显式设置 `voice.model_name`。
|
||||
|
||||
## 常见错误
|
||||
|
||||
- 在 `model_list` 里定义了 ASR 模型,但忘了设置 `voice.model_name`。
|
||||
- 把 API Key 写进了 `voice`,而不是 `.security.yml`。
|
||||
- 选择了不支持 ASR 的模型,却期望得到 Whisper 风格的转录结果。
|
||||
- 自定义了错误的 `api_base`,导致请求打到错误的接口地址。
|
||||
|
||||
## 最小检查清单
|
||||
|
||||
在测试语音输入前,请确认:
|
||||
|
||||
- `voice.model_name` 能正确匹配某个 `model_list[].model_name`。
|
||||
- `.security.yml` 中对应条目已经配置了有效 API Key。
|
||||
- 你选择的模型确实支持 ASR。
|
||||
- 你当前使用的频道已经启用了语音输入能力。
|
||||
252
pkg/audio/asr/agent.go
Normal file
252
pkg/audio/asr/agent.go
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/pion/rtp"
|
||||
"github.com/pion/webrtc/v3/pkg/media/oggwriter"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
type speechAccumulator struct {
|
||||
writer *oggwriter.OggWriter
|
||||
file string
|
||||
lastAudioAt time.Time
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
chatID string
|
||||
speakerID string
|
||||
sessionID string
|
||||
channel string
|
||||
}
|
||||
|
||||
func (a *speechAccumulator) Push(chunk bus.AudioChunk) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
if a.closed {
|
||||
return
|
||||
}
|
||||
|
||||
a.lastAudioAt = time.Now()
|
||||
|
||||
pkt := &rtp.Packet{
|
||||
Header: rtp.Header{
|
||||
SequenceNumber: uint16(chunk.Sequence),
|
||||
Timestamp: chunk.Timestamp,
|
||||
SSRC: 1, // Stable arbitrary dummy
|
||||
},
|
||||
Payload: chunk.Data,
|
||||
}
|
||||
|
||||
if err := a.writer.WriteRTP(pkt); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to write RTP", map[string]any{"error": err})
|
||||
}
|
||||
}
|
||||
|
||||
func (a *speechAccumulator) Close() {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
if !a.closed {
|
||||
a.writer.Close()
|
||||
a.closed = true
|
||||
}
|
||||
}
|
||||
|
||||
type Agent struct {
|
||||
bus *bus.MessageBus
|
||||
transcriber Transcriber
|
||||
|
||||
mu sync.Mutex
|
||||
sessions map[string]*speechAccumulator // keyed by sessionID_speakerID
|
||||
}
|
||||
|
||||
func NewAgent(mb *bus.MessageBus, t Transcriber) *Agent {
|
||||
return &Agent{
|
||||
bus: mb,
|
||||
transcriber: t,
|
||||
sessions: make(map[string]*speechAccumulator),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) Start(ctx context.Context) {
|
||||
logger.InfoCF("voice-agent", "Started Voice Agent orchestrator", nil)
|
||||
go a.listenChunks(ctx)
|
||||
go a.vadTick(ctx)
|
||||
|
||||
// Cleanup sessions on shutdown
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
a.mu.Lock()
|
||||
for key, acc := range a.sessions {
|
||||
acc.Close()
|
||||
os.Remove(acc.file)
|
||||
delete(a.sessions, key)
|
||||
}
|
||||
a.mu.Unlock()
|
||||
logger.InfoCF("voice-agent", "Cleaned up voice sessions on shutdown", nil)
|
||||
}()
|
||||
}
|
||||
|
||||
func (a *Agent) listenChunks(ctx context.Context) {
|
||||
chunks := a.bus.AudioChunksChan()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case chunk, ok := <-chunks:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
a.handleChunk(chunk)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) handleChunk(chunk bus.AudioChunk) {
|
||||
// Only accept Opus-encoded audio
|
||||
if chunk.Format != "opus" {
|
||||
logger.DebugCF("voice-agent", "Ignoring unsupported audio format", map[string]any{"format": chunk.Format})
|
||||
return
|
||||
}
|
||||
|
||||
key := fmt.Sprintf("%s_%s", chunk.SessionID, chunk.SpeakerID)
|
||||
|
||||
a.mu.Lock()
|
||||
acc, exists := a.sessions[key]
|
||||
if !exists {
|
||||
filename := filepath.Join(os.TempDir(), fmt.Sprintf("voice_%s_%d.ogg", key, time.Now().UnixNano()))
|
||||
writer, err := oggwriter.New(filename, uint32(chunk.SampleRate), uint16(chunk.Channels))
|
||||
if err != nil {
|
||||
a.mu.Unlock()
|
||||
logger.ErrorCF("voice-agent", "Failed to create OggWriter", map[string]any{"error": err})
|
||||
return
|
||||
}
|
||||
|
||||
acc = &speechAccumulator{
|
||||
writer: writer,
|
||||
file: filename,
|
||||
lastAudioAt: time.Now(),
|
||||
chatID: chunk.ChatID,
|
||||
speakerID: chunk.SpeakerID,
|
||||
sessionID: chunk.SessionID,
|
||||
channel: chunk.Channel,
|
||||
}
|
||||
a.sessions[key] = acc
|
||||
logger.DebugCF("voice-agent", "Started accumulating voice", map[string]any{"key": key, "file": filename})
|
||||
}
|
||||
a.mu.Unlock()
|
||||
|
||||
acc.Push(chunk)
|
||||
}
|
||||
|
||||
func (a *Agent) vadTick(ctx context.Context) {
|
||||
ticker := time.NewTicker(500 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
a.checkSilence(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) checkSilence(ctx context.Context) {
|
||||
a.mu.Lock()
|
||||
now := time.Now()
|
||||
var finished []*speechAccumulator
|
||||
|
||||
for key, acc := range a.sessions {
|
||||
acc.mu.Lock()
|
||||
last := acc.lastAudioAt
|
||||
acc.mu.Unlock()
|
||||
|
||||
if now.Sub(last) > 1500*time.Millisecond {
|
||||
acc.Close()
|
||||
delete(a.sessions, key)
|
||||
finished = append(finished, acc)
|
||||
}
|
||||
}
|
||||
a.mu.Unlock()
|
||||
|
||||
for _, acc := range finished {
|
||||
go a.processUtterance(ctx, acc)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) processUtterance(ctx context.Context, acc *speechAccumulator) {
|
||||
defer os.Remove(acc.file)
|
||||
|
||||
logger.InfoCF("voice-agent", "User finished speaking, transcribing...", map[string]any{"file": acc.file})
|
||||
|
||||
if a.transcriber == nil {
|
||||
logger.ErrorCF("voice-agent", "No STT configured!", nil)
|
||||
return
|
||||
}
|
||||
|
||||
res, err := a.transcriber.Transcribe(ctx, acc.file)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice-agent", "Transcription failed", map[string]any{"error": err})
|
||||
return
|
||||
}
|
||||
|
||||
if res.Text == "" {
|
||||
logger.DebugCF("voice-agent", "Ignored empty transcription", map[string]any{"file": acc.file})
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoCF("voice-agent", "Transcription result", map[string]any{"text": res.Text, "duration": res.Duration})
|
||||
|
||||
channelType := acc.channel
|
||||
if channelType == "" {
|
||||
channelType = "discord" // fallback for legacy chunks
|
||||
}
|
||||
|
||||
text := strings.ToLower(strings.TrimSpace(res.Text))
|
||||
if strings.Contains(text, "leave the voice channel") || strings.Contains(text, "leave voice") ||
|
||||
strings.Contains(text, "disconnect voice") || strings.Contains(text, "leave the channel") ||
|
||||
strings.Contains(text, "leave channel") {
|
||||
logger.InfoCF("voice-agent", "Voice command triggered: leave", nil)
|
||||
if err := a.bus.PublishVoiceControl(ctx, bus.VoiceControl{
|
||||
SessionID: acc.sessionID,
|
||||
Type: "command",
|
||||
Action: "leave",
|
||||
}); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to publish leave control", map[string]any{"error": err})
|
||||
}
|
||||
if err := a.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
Channel: channelType,
|
||||
ChatID: acc.chatID,
|
||||
Content: "Goodbye! Leaving the voice channel.",
|
||||
}); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to publish goodbye message", map[string]any{"error": err})
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
oralPrompt := "\n\n[SYSTEM]: The user just spoke this to you over voice chat. Please reply in a highly concise, conversational, oral style suitable for text-to-speech. Do not use markdown, emojis, asterisks, or code blocks. Speak naturally."
|
||||
|
||||
if err := a.bus.PublishInbound(ctx, bus.InboundMessage{
|
||||
Channel: channelType,
|
||||
SenderID: acc.speakerID,
|
||||
ChatID: acc.chatID,
|
||||
Content: res.Text + oralPrompt,
|
||||
Peer: bus.Peer{Kind: "channel", ID: acc.chatID},
|
||||
Metadata: map[string]string{
|
||||
"is_voice": "true",
|
||||
},
|
||||
}); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to publish inbound message", map[string]any{"error": err})
|
||||
}
|
||||
}
|
||||
196
pkg/audio/asr/agent_test.go
Normal file
196
pkg/audio/asr/agent_test.go
Normal file
|
|
@ -0,0 +1,196 @@
|
|||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/pion/webrtc/v3/pkg/media/oggwriter"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
)
|
||||
|
||||
type fakeTranscriber struct {
|
||||
text string
|
||||
err error
|
||||
lastPath string
|
||||
}
|
||||
|
||||
func (f *fakeTranscriber) Name() string { return "fake" }
|
||||
|
||||
func (f *fakeTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
||||
f.lastPath = audioFilePath
|
||||
if f.err != nil {
|
||||
return nil, f.err
|
||||
}
|
||||
return &TranscriptionResponse{Text: f.text}, nil
|
||||
}
|
||||
|
||||
func waitForFileRemoval(t *testing.T, path string, timeout time.Duration) {
|
||||
t.Helper()
|
||||
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if _, err := os.Stat(path); os.IsNotExist(err) {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
t.Fatalf("expected file to be removed: %s", path)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentHandleChunkCreatesSession(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
mb := bus.NewMessageBus()
|
||||
defer mb.Close()
|
||||
|
||||
agent := NewAgent(mb, &fakeTranscriber{})
|
||||
|
||||
chunk := bus.AudioChunk{
|
||||
SessionID: "sess",
|
||||
SpeakerID: "speaker",
|
||||
ChatID: "chat",
|
||||
Channel: "discord",
|
||||
Sequence: 1,
|
||||
Timestamp: 1,
|
||||
SampleRate: 48000,
|
||||
Channels: 2,
|
||||
Format: "opus",
|
||||
Data: []byte{0xF8, 0xFF, 0xFE},
|
||||
}
|
||||
|
||||
agent.handleChunk(chunk)
|
||||
|
||||
key := "sess_speaker"
|
||||
agent.mu.Lock()
|
||||
acc, ok := agent.sessions[key]
|
||||
agent.mu.Unlock()
|
||||
if !ok {
|
||||
t.Fatal("expected session to be created")
|
||||
}
|
||||
|
||||
acc.Close()
|
||||
_ = os.Remove(acc.file)
|
||||
}
|
||||
|
||||
func TestAgentHandleChunkIgnoresUnsupportedFormat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
mb := bus.NewMessageBus()
|
||||
defer mb.Close()
|
||||
|
||||
agent := NewAgent(mb, &fakeTranscriber{})
|
||||
|
||||
chunk := bus.AudioChunk{Format: "pcm"}
|
||||
agent.handleChunk(chunk)
|
||||
|
||||
agent.mu.Lock()
|
||||
count := len(agent.sessions)
|
||||
agent.mu.Unlock()
|
||||
if count != 0 {
|
||||
t.Fatalf("expected no sessions, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentProcessUtteranceLeaveCommand(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
mb := bus.NewMessageBus()
|
||||
defer mb.Close()
|
||||
|
||||
tr := &fakeTranscriber{text: "please leave the voice channel now"}
|
||||
agent := NewAgent(mb, tr)
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
filePath := filepath.Join(tmpDir, "voice.ogg")
|
||||
if err := os.WriteFile(filePath, []byte("data"), 0o600); err != nil {
|
||||
t.Fatalf("write temp file: %v", err)
|
||||
}
|
||||
|
||||
acc := &speechAccumulator{
|
||||
file: filePath,
|
||||
chatID: "chat",
|
||||
speakerID: "speaker",
|
||||
sessionID: "sess",
|
||||
channel: "discord",
|
||||
}
|
||||
|
||||
agent.processUtterance(context.Background(), acc)
|
||||
|
||||
select {
|
||||
case ctrl := <-mb.VoiceControlsChan():
|
||||
if ctrl.Action != "leave" || ctrl.Type != "command" || ctrl.SessionID != "sess" {
|
||||
t.Fatalf("unexpected voice control: %#v", ctrl)
|
||||
}
|
||||
case <-time.After(250 * time.Millisecond):
|
||||
t.Fatal("expected voice control publish")
|
||||
}
|
||||
|
||||
select {
|
||||
case out := <-mb.OutboundChan():
|
||||
if !strings.Contains(out.Content, "Leaving the voice channel") {
|
||||
t.Fatalf("unexpected outbound content: %q", out.Content)
|
||||
}
|
||||
case <-time.After(250 * time.Millisecond):
|
||||
t.Fatal("expected outbound publish")
|
||||
}
|
||||
|
||||
if _, err := os.Stat(filePath); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected temp file to be removed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentCheckSilencePublishesInboundAndCleansUp(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
mb := bus.NewMessageBus()
|
||||
defer mb.Close()
|
||||
|
||||
tr := &fakeTranscriber{text: "hello there"}
|
||||
agent := NewAgent(mb, tr)
|
||||
|
||||
filePath := filepath.Join(t.TempDir(), "voice.ogg")
|
||||
writer, err := oggwriter.New(filePath, 48000, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("create ogg writer: %v", err)
|
||||
}
|
||||
|
||||
acc := &speechAccumulator{
|
||||
writer: writer,
|
||||
file: filePath,
|
||||
lastAudioAt: time.Now().Add(-2 * time.Second),
|
||||
chatID: "chat",
|
||||
speakerID: "speaker",
|
||||
sessionID: "sess",
|
||||
channel: "slack",
|
||||
}
|
||||
|
||||
agent.mu.Lock()
|
||||
agent.sessions["sess_speaker"] = acc
|
||||
agent.mu.Unlock()
|
||||
|
||||
agent.checkSilence(context.Background())
|
||||
|
||||
select {
|
||||
case msg := <-mb.InboundChan():
|
||||
if msg.Channel != "slack" {
|
||||
t.Fatalf("unexpected inbound channel: %q", msg.Channel)
|
||||
}
|
||||
if !strings.Contains(msg.Content, "hello there") {
|
||||
t.Fatalf("unexpected inbound content: %q", msg.Content)
|
||||
}
|
||||
if msg.Metadata["is_voice"] != "true" {
|
||||
t.Fatalf("expected is_voice metadata, got %#v", msg.Metadata)
|
||||
}
|
||||
case <-time.After(500 * time.Millisecond):
|
||||
t.Fatal("expected inbound publish")
|
||||
}
|
||||
|
||||
waitForFileRemoval(t, filePath, 500*time.Millisecond)
|
||||
}
|
||||
131
pkg/audio/asr/asr.go
Normal file
131
pkg/audio/asr/asr.go
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
type Transcriber interface {
|
||||
Name() string
|
||||
Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error)
|
||||
}
|
||||
|
||||
type TranscriptionResponse struct {
|
||||
Text string `json:"text"`
|
||||
Language string `json:"language,omitempty"`
|
||||
Duration float64 `json:"duration,omitempty"`
|
||||
}
|
||||
|
||||
func supportsAudioTranscription(model string) bool {
|
||||
protocol, _ := providers.ExtractProtocol(model)
|
||||
|
||||
switch protocol {
|
||||
case "openai", "azure", "azure-openai",
|
||||
"litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"qwen-us", "dashscope-us", "mistral", "avian", "minimax", "longcat", "modelscope", "novita",
|
||||
"coding-plan", "alibaba-coding", "qwen-coding":
|
||||
// These protocols all go through the OpenAI-compatible or Azure provider path in
|
||||
// providers.CreateProviderFromConfig, so they are the only ones that can supply
|
||||
// the audio media payload shape expected by NewAudioModelTranscriber.
|
||||
|
||||
// TODO: Further restrict this by modelID, since not every model under these
|
||||
// protocols supports audio transcription.
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func supportsWhisperTranscription(model string) bool {
|
||||
protocol, _ := providers.ExtractProtocol(model)
|
||||
|
||||
switch protocol {
|
||||
case "openai", "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"qwen-us", "dashscope-us", "mistral", "avian", "minimax", "longcat", "modelscope", "novita",
|
||||
"coding-plan", "alibaba-coding", "qwen-coding", "mimo":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func whisperModelID(modelCfg *config.ModelConfig) string {
|
||||
if modelCfg == nil || modelCfg.APIKey() == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
if !supportsWhisperTranscription(modelCfg.Model) {
|
||||
return ""
|
||||
}
|
||||
|
||||
_, modelID := providers.ExtractProtocol(strings.TrimSpace(modelCfg.Model))
|
||||
if strings.Contains(strings.ToLower(modelID), "whisper") {
|
||||
return modelID
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func transcriberFromModelConfig(modelCfg *config.ModelConfig) Transcriber {
|
||||
if modelCfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
|
||||
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
|
||||
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
|
||||
}
|
||||
if modelID := whisperModelID(modelCfg); modelID != "" {
|
||||
return NewWhisperTranscriber(modelCfg)
|
||||
}
|
||||
if supportsAudioTranscription(modelCfg.Model) {
|
||||
return NewAudioModelTranscriber(modelCfg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func fallbackTranscriberFromModelConfig(modelCfg *config.ModelConfig) Transcriber {
|
||||
if modelCfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
|
||||
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
|
||||
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
|
||||
}
|
||||
if modelID := whisperModelID(modelCfg); modelID != "" {
|
||||
return NewWhisperTranscriber(modelCfg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DetectTranscriber inspects cfg and returns the appropriate Transcriber, or
|
||||
// nil if no supported transcription provider is configured.
|
||||
func DetectTranscriber(cfg *config.Config) Transcriber {
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if modelName := strings.TrimSpace(cfg.Voice.ModelName); modelName != "" {
|
||||
modelCfg, err := cfg.GetModelConfig(modelName)
|
||||
if err == nil {
|
||||
if tr := transcriberFromModelConfig(modelCfg); tr != nil {
|
||||
return tr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fall back to compatibility scanning for legacy auto-detected ASR providers.
|
||||
for _, mc := range cfg.ModelList {
|
||||
if tr := fallbackTranscriberFromModelConfig(mc); tr != nil {
|
||||
return tr
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
|
@ -33,26 +33,68 @@ func TestDetectTranscriber(t *testing.T) {
|
|||
wantName: "audio-model",
|
||||
},
|
||||
{
|
||||
name: "groq via model list",
|
||||
name: "voice model name alias selects elevenlabs transcriber",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "my-asr-model",
|
||||
Model: "elevenlabs/scribe_v1",
|
||||
APIKeys: config.SimpleSecureStrings("sk_elevenlabs_test"),
|
||||
},
|
||||
},
|
||||
},
|
||||
wantName: "elevenlabs",
|
||||
},
|
||||
{
|
||||
name: "voice model name alias selects whisper transcriber for groq",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "my-asr-model",
|
||||
Model: "groq/whisper-large-v3",
|
||||
APIKeys: config.SimpleSecureStrings("sk-groq-model"),
|
||||
},
|
||||
},
|
||||
},
|
||||
wantName: "whisper",
|
||||
},
|
||||
{
|
||||
name: "openai whisper alias selects whisper transcriber",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "my-asr-model",
|
||||
Model: "openai/whisper-1",
|
||||
APIKeys: config.SimpleSecureStrings("sk-openai-model"),
|
||||
},
|
||||
},
|
||||
},
|
||||
wantName: "whisper",
|
||||
},
|
||||
{
|
||||
name: "whisper via model list fallback",
|
||||
cfg: &config.Config{
|
||||
ModelList: []*config.ModelConfig{
|
||||
{ModelName: "openai", Model: "openai/gpt-4o", APIKeys: config.SimpleSecureStrings("sk-openai")},
|
||||
{
|
||||
ModelName: "groq",
|
||||
Model: "groq/llama-3.3-70b",
|
||||
Model: "groq/whisper-large-v3-turbo",
|
||||
APIKeys: config.SimpleSecureStrings("sk-groq-model"),
|
||||
},
|
||||
},
|
||||
},
|
||||
wantName: "groq",
|
||||
wantName: "whisper",
|
||||
},
|
||||
{
|
||||
name: "voice model name selects non-gemini audio model transcriber",
|
||||
name: "voice model name alias selects non-gemini audio model transcriber",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{ModelName: "voice-openai-audio"},
|
||||
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "voice-openai-audio",
|
||||
ModelName: "my-asr-model",
|
||||
Model: "openai/gpt-4o-audio-preview",
|
||||
APIKeys: config.SimpleSecureStrings("sk-openai"),
|
||||
},
|
||||
|
|
@ -92,7 +134,7 @@ func TestDetectTranscriber(t *testing.T) {
|
|||
name: "groq model list entry without key is skipped",
|
||||
cfg: &config.Config{
|
||||
ModelList: []*config.ModelConfig{
|
||||
{Model: "groq/llama-3.3-70b"},
|
||||
{Model: "groq/whisper-large-v3"},
|
||||
},
|
||||
},
|
||||
wantNil: true,
|
||||
|
|
@ -103,12 +145,12 @@ func TestDetectTranscriber(t *testing.T) {
|
|||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "groq",
|
||||
Model: "groq/llama-3.3-70b",
|
||||
Model: "groq/whisper-large-v3",
|
||||
APIKeys: config.SimpleSecureStrings("sk-groq-model"),
|
||||
},
|
||||
},
|
||||
},
|
||||
wantName: "groq",
|
||||
wantName: "whisper",
|
||||
},
|
||||
{
|
||||
name: "missing voice model name config returns nil",
|
||||
|
|
@ -127,15 +169,17 @@ func TestDetectTranscriber(t *testing.T) {
|
|||
{
|
||||
name: "elevenlabs voice config key",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{ElevenLabsAPIKey: "sk_elevenlabs_test"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{Model: "elevenlabs/scribe_v1", APIKeys: config.SimpleSecureStrings("sk_elevenlabs_test")},
|
||||
},
|
||||
},
|
||||
wantName: "elevenlabs",
|
||||
},
|
||||
{
|
||||
name: "elevenlabs takes priority over groq model list",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{ElevenLabsAPIKey: "sk_elevenlabs_test"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{Model: "elevenlabs/scribe_v1", APIKeys: config.SimpleSecureStrings("sk_elevenlabs_test")},
|
||||
{
|
||||
ModelName: "groq",
|
||||
Model: "groq/llama-3.3-70b",
|
||||
|
|
@ -149,10 +193,10 @@ func TestDetectTranscriber(t *testing.T) {
|
|||
name: "voice model name takes priority over elevenlabs",
|
||||
cfg: &config.Config{
|
||||
Voice: config.VoiceConfig{
|
||||
ModelName: "voice-gemini",
|
||||
ElevenLabsAPIKey: "sk_elevenlabs_test",
|
||||
ModelName: "voice-gemini",
|
||||
},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{Model: "elevenlabs", APIKeys: config.SimpleSecureStrings("sk_elevenlabs_test")},
|
||||
{
|
||||
ModelName: "voice-gemini",
|
||||
Model: "gemini/gemini-2.5-flash",
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
|
|
@ -23,12 +23,16 @@ type ElevenLabsTranscriber struct {
|
|||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewElevenLabsTranscriber(apiKey string) *ElevenLabsTranscriber {
|
||||
func NewElevenLabsTranscriber(apiKey, apiBase string) *ElevenLabsTranscriber {
|
||||
logger.DebugCF("voice", "Creating ElevenLabs transcriber", map[string]any{"has_api_key": apiKey != ""})
|
||||
|
||||
if apiBase == "" {
|
||||
apiBase = "https://api.elevenlabs.io"
|
||||
}
|
||||
|
||||
return &ElevenLabsTranscriber{
|
||||
apiKey: apiKey,
|
||||
apiBase: "https://api.elevenlabs.io",
|
||||
apiBase: apiBase,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 120 * time.Second,
|
||||
},
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
|
@ -14,7 +14,7 @@ import (
|
|||
var _ Transcriber = (*ElevenLabsTranscriber)(nil)
|
||||
|
||||
func TestElevenLabsTranscriberName(t *testing.T) {
|
||||
tr := NewElevenLabsTranscriber("sk_test")
|
||||
tr := NewElevenLabsTranscriber("sk_test", "")
|
||||
if got := tr.Name(); got != "elevenlabs" {
|
||||
t.Errorf("Name() = %q, want %q", got, "elevenlabs")
|
||||
}
|
||||
|
|
@ -43,7 +43,7 @@ func TestElevenLabsTranscribe(t *testing.T) {
|
|||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tr := NewElevenLabsTranscriber("sk_test")
|
||||
tr := NewElevenLabsTranscriber("sk_test", "")
|
||||
tr.apiBase = srv.URL
|
||||
|
||||
resp, err := tr.Transcribe(context.Background(), audioPath)
|
||||
|
|
@ -64,7 +64,7 @@ func TestElevenLabsTranscribe(t *testing.T) {
|
|||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tr := NewElevenLabsTranscriber("sk_bad")
|
||||
tr := NewElevenLabsTranscriber("sk_bad", "")
|
||||
tr.apiBase = srv.URL
|
||||
|
||||
_, err := tr.Transcribe(context.Background(), audioPath)
|
||||
|
|
@ -74,7 +74,7 @@ func TestElevenLabsTranscribe(t *testing.T) {
|
|||
})
|
||||
|
||||
t.Run("missing file", func(t *testing.T) {
|
||||
tr := NewElevenLabsTranscriber("sk_test")
|
||||
tr := NewElevenLabsTranscriber("sk_test", "")
|
||||
_, err := tr.Transcribe(context.Background(), filepath.Join(tmpDir, "nonexistent.ogg"))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing file, got nil")
|
||||
245
pkg/audio/asr/whisper_transcriber.go
Normal file
245
pkg/audio/asr/whisper_transcriber.go
Normal file
|
|
@ -0,0 +1,245 @@
|
|||
package asr
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
type WhisperTranscriber struct {
|
||||
apiKey string
|
||||
apiBase string
|
||||
modelID string
|
||||
providerName string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewWhisperTranscriber(modelCfg *config.ModelConfig) *WhisperTranscriber {
|
||||
if modelCfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
protocol, modelID := providers.ExtractProtocol(modelCfg.Model)
|
||||
if modelID == "" {
|
||||
modelID = strings.TrimSpace(modelCfg.Model)
|
||||
}
|
||||
|
||||
tr := newWhisperTranscriber(
|
||||
modelCfg.APIKey(),
|
||||
providers.ResolveAPIBase(modelCfg),
|
||||
modelID,
|
||||
protocol,
|
||||
)
|
||||
if tr == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "Creating whisper transcriber", map[string]any{
|
||||
"api_base": tr.apiBase,
|
||||
"has_key": tr.apiKey != "",
|
||||
"model": tr.modelID,
|
||||
"provider": tr.providerName,
|
||||
})
|
||||
return tr
|
||||
}
|
||||
|
||||
func NewGroqTranscriber(apiKey, modelID string) *WhisperTranscriber {
|
||||
return newWhisperTranscriber(apiKey, "https://api.groq.com/openai/v1", modelID, "groq")
|
||||
}
|
||||
|
||||
func newWhisperTranscriber(apiKey, apiBase, modelID, providerName string) *WhisperTranscriber {
|
||||
if modelID == "" {
|
||||
return nil
|
||||
}
|
||||
if providerName == "" {
|
||||
providerName = "whisper"
|
||||
}
|
||||
return &WhisperTranscriber{
|
||||
apiKey: apiKey,
|
||||
apiBase: strings.TrimRight(apiBase, "/"),
|
||||
modelID: modelID,
|
||||
providerName: providerName,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 60 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (t *WhisperTranscriber) transcriptionURL() string {
|
||||
base := strings.TrimRight(t.apiBase, "/")
|
||||
if strings.HasSuffix(base, "/audio/transcriptions") {
|
||||
return base
|
||||
}
|
||||
return base + "/audio/transcriptions"
|
||||
}
|
||||
|
||||
func (t *WhisperTranscriber) TranscribeData(
|
||||
ctx context.Context,
|
||||
data []byte,
|
||||
filename string,
|
||||
) (*TranscriptionResponse, error) {
|
||||
logger.InfoCF("voice", "Starting whisper transcription from memory", map[string]any{
|
||||
"bytes": len(data),
|
||||
"filename": filename,
|
||||
"model": t.modelID,
|
||||
"provider": t.providerName,
|
||||
})
|
||||
|
||||
var requestBody bytes.Buffer
|
||||
writer := multipart.NewWriter(&requestBody)
|
||||
|
||||
part, err := writer.CreateFormFile("file", filename)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to create whisper form file", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to create form file: %w", err)
|
||||
}
|
||||
|
||||
if _, copyErr := io.Copy(part, bytes.NewReader(data)); copyErr != nil {
|
||||
logger.ErrorCF("voice", "Failed to copy whisper file content", map[string]any{"error": copyErr})
|
||||
return nil, fmt.Errorf("failed to copy file content: %w", copyErr)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("model", t.modelID); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to write whisper model field", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to write model field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("response_format", "json"); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to write whisper response_format field", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to write response_format field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.Close(); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to close whisper multipart writer", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), int64(len(data)))
|
||||
}
|
||||
|
||||
func (t *WhisperTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
||||
logger.InfoCF("voice", "Starting whisper transcription", map[string]any{
|
||||
"audio_file": audioFilePath,
|
||||
"model": t.modelID,
|
||||
"provider": t.providerName,
|
||||
})
|
||||
|
||||
audioFile, err := os.Open(audioFilePath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open audio file %s: %w", audioFilePath, err)
|
||||
}
|
||||
defer audioFile.Close()
|
||||
|
||||
fileInfo, err := audioFile.Stat()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to stat audio file %s: %w", audioFilePath, err)
|
||||
}
|
||||
|
||||
var requestBody bytes.Buffer
|
||||
writer := multipart.NewWriter(&requestBody)
|
||||
|
||||
part, err := writer.CreateFormFile("file", filepath.Base(audioFilePath))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create form file: %w", err)
|
||||
}
|
||||
|
||||
if _, copyErr := io.Copy(part, audioFile); copyErr != nil {
|
||||
return nil, fmt.Errorf("failed to copy audio data: %w", copyErr)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("model", t.modelID); err != nil {
|
||||
return nil, fmt.Errorf("failed to write model field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("response_format", "json"); err != nil {
|
||||
return nil, fmt.Errorf("failed to write response_format field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.Close(); err != nil {
|
||||
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), fileInfo.Size())
|
||||
}
|
||||
|
||||
func (t *WhisperTranscriber) doRequest(
|
||||
ctx context.Context,
|
||||
requestBody *bytes.Buffer,
|
||||
contentType string,
|
||||
fileSize int64,
|
||||
) (*TranscriptionResponse, error) {
|
||||
url := t.transcriptionURL()
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", url, requestBody)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to create whisper request", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
if t.apiKey != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "Sending whisper transcription request", map[string]any{
|
||||
"file_size_bytes": fileSize,
|
||||
"model": t.modelID,
|
||||
"provider": t.providerName,
|
||||
"request_size_bytes": requestBody.Len(),
|
||||
"url": url,
|
||||
})
|
||||
|
||||
resp, err := t.httpClient.Do(req)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to send whisper request", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to read whisper response", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
logger.ErrorCF("voice", "Whisper API error", map[string]any{
|
||||
"provider": t.providerName,
|
||||
"response": string(body),
|
||||
"status_code": resp.StatusCode,
|
||||
})
|
||||
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var result TranscriptionResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to unmarshal whisper response", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
logger.InfoCF("voice", "Whisper transcription completed successfully", map[string]any{
|
||||
"duration_seconds": result.Duration,
|
||||
"language": result.Language,
|
||||
"provider": t.providerName,
|
||||
"text_length": len(result.Text),
|
||||
"transcription_preview": utils.Truncate(result.Text, 50),
|
||||
})
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func (t *WhisperTranscriber) Name() string {
|
||||
return "whisper"
|
||||
}
|
||||
102
pkg/audio/asr/whisper_transcriber_test.go
Normal file
102
pkg/audio/asr/whisper_transcriber_test.go
Normal file
|
|
@ -0,0 +1,102 @@
|
|||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
func TestWhisperTranscriberTranscribeDataUsesConfiguredModel(t *testing.T) {
|
||||
var gotModel string
|
||||
var gotPath string
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer sk-openai-test" {
|
||||
t.Errorf("Authorization = %q, want %q", got, "Bearer sk-openai-test")
|
||||
}
|
||||
|
||||
reader, err := r.MultipartReader()
|
||||
if err != nil {
|
||||
t.Fatalf("MultipartReader() error: %v", err)
|
||||
}
|
||||
|
||||
for {
|
||||
part, err := reader.NextPart()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("NextPart() error: %v", err)
|
||||
}
|
||||
|
||||
data, err := io.ReadAll(part)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadAll() error: %v", err)
|
||||
}
|
||||
|
||||
if part.FormName() == "model" {
|
||||
gotModel = string(data)
|
||||
}
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if err := json.NewEncoder(w).Encode(TranscriptionResponse{Text: "hello from whisper"}); err != nil {
|
||||
t.Fatalf("Encode() error: %v", err)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tr := NewWhisperTranscriber(&config.ModelConfig{
|
||||
Model: "openai/whisper-1",
|
||||
APIBase: server.URL,
|
||||
APIKeys: config.SimpleSecureStrings("sk-openai-test"),
|
||||
})
|
||||
tr.httpClient = server.Client()
|
||||
|
||||
resp, err := tr.TranscribeData(context.Background(), []byte("audio"), "clip.ogg")
|
||||
if err != nil {
|
||||
t.Fatalf("TranscribeData() error: %v", err)
|
||||
}
|
||||
if resp.Text != "hello from whisper" {
|
||||
t.Errorf("Text = %q, want %q", resp.Text, "hello from whisper")
|
||||
}
|
||||
if gotModel != "whisper-1" {
|
||||
t.Errorf("model field = %q, want %q", gotModel, "whisper-1")
|
||||
}
|
||||
if gotPath != "/audio/transcriptions" {
|
||||
t.Errorf("path = %q, want %q", gotPath, "/audio/transcriptions")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWhisperTranscriberUsesEndpointAPIBaseWithoutDoubleAppend(t *testing.T) {
|
||||
var gotPath string
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if err := json.NewEncoder(w).Encode(TranscriptionResponse{Text: "ok"}); err != nil {
|
||||
t.Fatalf("Encode() error: %v", err)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tr := NewWhisperTranscriber(&config.ModelConfig{
|
||||
Model: "groq/whisper-large-v3",
|
||||
APIBase: server.URL + "/audio/transcriptions",
|
||||
APIKeys: config.SimpleSecureStrings("sk-groq-test"),
|
||||
})
|
||||
tr.httpClient = server.Client()
|
||||
|
||||
if _, err := tr.TranscribeData(context.Background(), []byte("audio"), "clip.ogg"); err != nil {
|
||||
t.Fatalf("TranscribeData() error: %v", err)
|
||||
}
|
||||
if gotPath != "/audio/transcriptions" {
|
||||
t.Errorf("path = %q, want %q", gotPath, "/audio/transcriptions")
|
||||
}
|
||||
}
|
||||
57
pkg/audio/ogg.go
Normal file
57
pkg/audio/ogg.go
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
package audio
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
// DecodeOggOpus reads an Ogg format stream and extracts individual Opus payloads.
|
||||
// It calls onFrame for every complete Opus frame found in the stream.
|
||||
func DecodeOggOpus(r io.Reader, onFrame func([]byte) error) error {
|
||||
var packet bytes.Buffer
|
||||
header := make([]byte, 27)
|
||||
segment := make([]byte, 255)
|
||||
|
||||
for {
|
||||
if _, err := io.ReadFull(r, header); err != nil {
|
||||
if err == io.EOF || err == io.ErrUnexpectedEOF {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("failed to read ogg header: %w", err)
|
||||
}
|
||||
if string(header[:4]) != "OggS" {
|
||||
return fmt.Errorf("invalid ogg magic string")
|
||||
}
|
||||
|
||||
pageSegments := int(header[26])
|
||||
segmentTable := make([]byte, pageSegments)
|
||||
if _, err := io.ReadFull(r, segmentTable); err != nil {
|
||||
return fmt.Errorf("failed to read segment table: %w", err)
|
||||
}
|
||||
|
||||
for _, lacing := range segmentTable {
|
||||
if _, err := io.ReadFull(r, segment[:lacing]); err != nil {
|
||||
return fmt.Errorf("failed to read segment data: %w", err)
|
||||
}
|
||||
|
||||
packet.Write(segment[:lacing])
|
||||
|
||||
// If lacing is less than 255, the packet is complete
|
||||
if lacing < 255 {
|
||||
if packet.Len() > 0 {
|
||||
packetBytes := packet.Bytes()
|
||||
// Ignore Ogg Opus headers
|
||||
if !bytes.HasPrefix(packetBytes, []byte("OpusHead")) &&
|
||||
!bytes.HasPrefix(packetBytes, []byte("OpusTags")) {
|
||||
if err := onFrame(packetBytes); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// Start new packet
|
||||
packet.Reset()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
146
pkg/audio/ogg_test.go
Normal file
146
pkg/audio/ogg_test.go
Normal file
|
|
@ -0,0 +1,146 @@
|
|||
package audio
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildOggPage helper creates an Ogg page for testing.
|
||||
// lacingVals specifies the segment table, and data is the payload.
|
||||
func buildOggPage(lacingVals []byte, data []byte) []byte {
|
||||
var buf bytes.Buffer
|
||||
// 27-byte Ogg header
|
||||
header := make([]byte, 27)
|
||||
copy(header[:4], "OggS")
|
||||
header[5] = 0 // type flag
|
||||
// For testing, we only care about OggS magic and page_segments (byte 26)
|
||||
header[26] = byte(len(lacingVals))
|
||||
buf.Write(header)
|
||||
buf.Write(lacingVals)
|
||||
buf.Write(data)
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func TestDecodeOggOpus_ValidParsing(t *testing.T) {
|
||||
var b bytes.Buffer
|
||||
|
||||
// Packet 1: Single segment, length 50
|
||||
pkt1 := bytes.Repeat([]byte{1}, 50)
|
||||
// Packet 2: Multi-segment (255 + 10 = 265 bytes)
|
||||
pkt2Part1 := bytes.Repeat([]byte{2}, 255)
|
||||
pkt2Part2 := bytes.Repeat([]byte{2}, 10)
|
||||
// Packet 3: Continued across pages. Page 1 gets 255, Page 2 gets 20. Total 275 bytes.
|
||||
pkt3Part1 := bytes.Repeat([]byte{3}, 255)
|
||||
pkt3Part2 := bytes.Repeat([]byte{3}, 20)
|
||||
|
||||
// Page 1: OpusHead (skip), OpusTags (skip), pkt1, pkt2, pkt3Part1
|
||||
page1Lacing := []byte{8, 8, 50, 255, 10, 255}
|
||||
page1Data := bytes.Join([][]byte{
|
||||
[]byte("OpusHead"),
|
||||
[]byte("OpusTags"),
|
||||
pkt1,
|
||||
pkt2Part1, pkt2Part2,
|
||||
pkt3Part1,
|
||||
}, nil)
|
||||
|
||||
// Page 2: pkt3Part2, pkt4 (length 10)
|
||||
pkt4 := bytes.Repeat([]byte{4}, 10)
|
||||
page2Lacing := []byte{20, 10}
|
||||
page2Data := bytes.Join([][]byte{
|
||||
pkt3Part2,
|
||||
pkt4,
|
||||
}, nil)
|
||||
|
||||
b.Write(buildOggPage(page1Lacing, page1Data))
|
||||
b.Write(buildOggPage(page2Lacing, page2Data))
|
||||
|
||||
var frames [][]byte
|
||||
err := DecodeOggOpus(&b, func(frame []byte) error {
|
||||
// making a copy to store as DecodeOggOpus might reuse backing array
|
||||
cpy := make([]byte, len(frame))
|
||||
copy(cpy, frame)
|
||||
frames = append(frames, cpy)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
expectedFrames := [][]byte{
|
||||
pkt1,
|
||||
append(pkt2Part1, pkt2Part2...),
|
||||
append(pkt3Part1, pkt3Part2...),
|
||||
pkt4,
|
||||
}
|
||||
|
||||
if len(frames) != len(expectedFrames) {
|
||||
t.Fatalf("expected %d frames, got %d", len(expectedFrames), len(frames))
|
||||
}
|
||||
|
||||
for i, expected := range expectedFrames {
|
||||
if !reflect.DeepEqual(frames[i], expected) {
|
||||
t.Errorf("frame %d mismatch:\nexp: %v\ngot: %v", i, expected, frames[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeOggOpus_Errors(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
errContains string
|
||||
}{
|
||||
{
|
||||
name: "invalid magic string",
|
||||
data: []byte(
|
||||
"OggX\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00",
|
||||
),
|
||||
errContains: "invalid ogg magic string",
|
||||
},
|
||||
{
|
||||
name: "short header",
|
||||
data: []byte("Ogg"),
|
||||
errContains: "failed to read ogg header",
|
||||
},
|
||||
{
|
||||
name: "eof in segment table",
|
||||
data: func() []byte {
|
||||
h := make([]byte, 27)
|
||||
copy(h, "OggS")
|
||||
h[26] = 5 // expects 5 bytes of segment table, but none provided
|
||||
return h
|
||||
}(),
|
||||
errContains: "failed to read segment table",
|
||||
},
|
||||
{
|
||||
name: "eof in segment data",
|
||||
data: func() []byte {
|
||||
h := make([]byte, 27, 28)
|
||||
copy(h, "OggS")
|
||||
h[26] = 1
|
||||
return append(h, 100) // expects 100 bytes of data, but none provided
|
||||
}(),
|
||||
errContains: "failed to read segment data",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := DecodeOggOpus(bytes.NewReader(tt.data), func(b []byte) error { return nil })
|
||||
if tt.name == "short header" {
|
||||
if err != nil {
|
||||
t.Errorf("expected no error (io.EOF/ErrUnexpectedEOF swallowed), got %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil {
|
||||
t.Fatalf("expected error containing %q, got nil", tt.errContains)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tt.errContains) {
|
||||
t.Errorf("expected error to contain %q, got: %q", tt.errContains, err.Error())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
96
pkg/audio/sentence.go
Normal file
96
pkg/audio/sentence.go
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
package audio
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
// SplitSentences splits text into sentence-sized chunks suitable for TTS synthesis.
|
||||
// It splits on sentence-ending punctuation (.!?\n, as well as CJK 。, !, ?) while avoiding false splits
|
||||
// on decimal numbers. Very short fragments are merged with
|
||||
// the next sentence to prevent choppy playback.
|
||||
func SplitSentences(text string) []string {
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var sentences []string
|
||||
var current strings.Builder
|
||||
runes := []rune(text)
|
||||
|
||||
for i := 0; i < len(runes); i++ {
|
||||
r := runes[i]
|
||||
if r == '\n' {
|
||||
s := strings.TrimSpace(current.String())
|
||||
if s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
current.Reset()
|
||||
continue
|
||||
}
|
||||
|
||||
current.WriteRune(r)
|
||||
|
||||
if r == '.' || r == '!' || r == '?' || r == '。' || r == '!' || r == '?' {
|
||||
// Avoid splitting on decimal numbers like "3.14"
|
||||
if r == '.' && i > 0 && unicode.IsDigit(runes[i-1]) &&
|
||||
i+1 < len(runes) && unicode.IsDigit(runes[i+1]) {
|
||||
continue
|
||||
}
|
||||
|
||||
// Consume contiguous punctuation clusters (e.g., "..." or "?!").
|
||||
for i+1 < len(runes) && (runes[i+1] == '.' || runes[i+1] == '!' || runes[i+1] == '?' || runes[i+1] == '。' || runes[i+1] == '!' || runes[i+1] == '?') {
|
||||
i++
|
||||
current.WriteRune(runes[i])
|
||||
}
|
||||
|
||||
s := strings.TrimSpace(current.String())
|
||||
if s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
current.Reset()
|
||||
}
|
||||
}
|
||||
|
||||
// Flush remaining text
|
||||
if s := strings.TrimSpace(current.String()); s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
|
||||
// Merge very short fragments with the next sentence
|
||||
return mergeShorties(sentences, 15)
|
||||
}
|
||||
|
||||
// mergeShorties merges sentences shorter than minLen characters with the following sentence.
|
||||
func mergeShorties(sentences []string, minLen int) []string {
|
||||
if len(sentences) <= 1 {
|
||||
return sentences
|
||||
}
|
||||
|
||||
var merged []string
|
||||
var buf string
|
||||
|
||||
for _, s := range sentences {
|
||||
if buf != "" {
|
||||
buf += " " + s
|
||||
if len([]rune(buf)) >= minLen {
|
||||
merged = append(merged, buf)
|
||||
buf = ""
|
||||
}
|
||||
} else if len([]rune(s)) < minLen {
|
||||
buf = s
|
||||
} else {
|
||||
merged = append(merged, s)
|
||||
}
|
||||
}
|
||||
|
||||
if buf != "" {
|
||||
if len(merged) > 0 {
|
||||
merged[len(merged)-1] += " " + buf
|
||||
} else {
|
||||
merged = append(merged, buf)
|
||||
}
|
||||
}
|
||||
|
||||
return merged
|
||||
}
|
||||
69
pkg/audio/sentence_test.go
Normal file
69
pkg/audio/sentence_test.go
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
package audio
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSplitSentences(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "empty input",
|
||||
in: "",
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "single sentence",
|
||||
in: "Hello world.",
|
||||
want: []string{"Hello world."},
|
||||
},
|
||||
{
|
||||
name: "decimal numbers do not split",
|
||||
in: "The value is 3.14 today. Keep watching closely.",
|
||||
want: []string{"The value is 3.14 today.", "Keep watching closely."},
|
||||
},
|
||||
{
|
||||
name: "newline boundary",
|
||||
in: "This is line number one\nThis is line number two",
|
||||
want: []string{"This is line number one", "This is line number two"},
|
||||
},
|
||||
{
|
||||
name: "newline with surrounding spaces",
|
||||
in: " This is the first line \n This is the second line ",
|
||||
want: []string{"This is the first line", "This is the second line"},
|
||||
},
|
||||
{
|
||||
name: "trailing punctuation consumed",
|
||||
in: "Please wait a moment... What on earth?! That is perfectly fine.",
|
||||
want: []string{"Please wait a moment...", "What on earth?!", "That is perfectly fine."},
|
||||
},
|
||||
{
|
||||
name: "short leading fragment merges with next",
|
||||
in: "Hi. This is a longer sentence.",
|
||||
want: []string{"Hi. This is a longer sentence."},
|
||||
},
|
||||
{
|
||||
name: "consecutive short fragments keep merging",
|
||||
in: "A. B. C. This is the real sentence.",
|
||||
want: []string{"A. B. C. This is the real sentence."},
|
||||
},
|
||||
{
|
||||
name: "short trailing fragment merges back",
|
||||
in: "This sentence is long enough. End.",
|
||||
want: []string{"This sentence is long enough. End."},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := SplitSentences(tc.in)
|
||||
if !reflect.DeepEqual(got, tc.want) {
|
||||
t.Fatalf("SplitSentences(%q) = %#v, want %#v", tc.in, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
137
pkg/audio/tts/README.md
Normal file
137
pkg/audio/tts/README.md
Normal file
|
|
@ -0,0 +1,137 @@
|
|||
# TTS (Text-to-Speech)
|
||||
|
||||
This package handles speech synthesis for PicoClaw.
|
||||
|
||||
If you are new to TTS setup, the simplest workflow is:
|
||||
|
||||
1. Add a TTS-capable entry to `model_list`.
|
||||
2. Point `voice.tts_model_name` at that entry.
|
||||
3. Put the API key in `.security.yml`.
|
||||
|
||||
## Quick Recommendation
|
||||
|
||||
For most users, these are the best starting points:
|
||||
|
||||
| Provider | Why start here |
|
||||
| --- | --- |
|
||||
| [OpenAI](https://platform.openai.com/docs/guides/text-to-speech) | Best-supported path in PicoClaw today. The current TTS implementation is built around the OpenAI-compatible `/audio/speech` API shape, and OpenAI is the safest default. |
|
||||
| [Xiaomi MiMo](https://platform.xiaomimimo.com) | A good second option if you want an OpenAI-compatible provider endpoint and are already using MiMo models in the rest of your stack. |
|
||||
|
||||
## How TTS Configuration Works
|
||||
|
||||
PicoClaw does not keep TTS API keys inside `voice`.
|
||||
|
||||
Instead:
|
||||
|
||||
- `voice.tts_model_name` selects a named entry from `model_list`.
|
||||
- That `model_list` entry provides the provider, model ID, API base, and proxy settings.
|
||||
- `.security.yml` stores the API key for the same named model entry.
|
||||
|
||||
This is the recommended and supported configuration pattern.
|
||||
|
||||
## Recommended Setup
|
||||
|
||||
### Option A: OpenAI
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"tts_model_name": "openai-tts"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "openai-tts",
|
||||
"model": "openai/tts-1"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
openai-tts:
|
||||
api_keys:
|
||||
- "sk-openai-your-key"
|
||||
```
|
||||
|
||||
### Option B: Xiaomi MiMo
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"tts_model_name": "mimo-tts"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "mimo-tts",
|
||||
"model": "mimo/mimo-v2-tts"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
mimo-tts:
|
||||
api_keys:
|
||||
- "your-mimo-key"
|
||||
```
|
||||
|
||||
If you use a custom MiMo endpoint, you can also set `api_base` explicitly. Otherwise PicoClaw will use the provider default.
|
||||
|
||||
## What PicoClaw Sends Today
|
||||
|
||||
The current TTS runtime uses an OpenAI-compatible speech request with these defaults:
|
||||
|
||||
- Endpoint: `/audio/speech`
|
||||
- Response format: `opus`
|
||||
- Voice: `alloy`
|
||||
- Model: taken from the selected `model_list` entry
|
||||
|
||||
That means:
|
||||
|
||||
- `openai/tts-1` works naturally.
|
||||
- Other OpenAI-compatible providers can work if they accept the same request format.
|
||||
- PicoClaw currently does not expose a user-facing config field for changing the TTS voice from `alloy`.
|
||||
|
||||
## How PicoClaw Chooses a TTS Provider
|
||||
|
||||
`DetectTTS` resolves TTS in this order:
|
||||
|
||||
1. **Preferred path**: resolve `voice.tts_model_name` against `model_list`.
|
||||
2. If a matching model entry exists and has an API key, PicoClaw creates an OpenAI-compatible TTS provider using that model's settings.
|
||||
3. **Fallback path**: if `voice.tts_model_name` is not set or cannot be resolved, PicoClaw scans `model_list` for the first entry whose model string contains `tts` and has an API key.
|
||||
|
||||
Fallback scanning exists for compatibility. New configs should set `voice.tts_model_name` explicitly.
|
||||
|
||||
## Notes About API Base Handling
|
||||
|
||||
PicoClaw normalizes the configured base URL for TTS:
|
||||
|
||||
- For OpenAI, a base like `https://api.openai.com` or `https://api.openai.com/v1` becomes `https://api.openai.com/v1/audio/speech`.
|
||||
- For other OpenAI-compatible providers, PicoClaw preserves the configured base path and ensures it ends with `/audio/speech`.
|
||||
- If `api_base` is omitted, PicoClaw uses the provider default base when the model prefix is known.
|
||||
|
||||
## Common Mistakes
|
||||
|
||||
- Setting `voice.tts_model_name` to a name that does not exist in `model_list`.
|
||||
- Adding a TTS model but forgetting to put its API key in `.security.yml`.
|
||||
- Assuming PicoClaw will automatically use provider-specific custom voices.
|
||||
- Using a provider endpoint that is not compatible with the OpenAI `/audio/speech` request format.
|
||||
|
||||
## Minimal Checklist
|
||||
|
||||
Before testing `send_tts`, make sure:
|
||||
|
||||
- `voice.tts_model_name` matches a `model_list[].model_name`.
|
||||
- The matching `.security.yml` entry contains a valid API key.
|
||||
- The chosen provider supports an OpenAI-compatible speech synthesis endpoint.
|
||||
- Your selected model is actually a TTS-capable model.
|
||||
137
pkg/audio/tts/README_zh.md
Normal file
137
pkg/audio/tts/README_zh.md
Normal file
|
|
@ -0,0 +1,137 @@
|
|||
# TTS(文本转语音)
|
||||
|
||||
这个目录负责 PicoClaw 的语音合成能力。
|
||||
|
||||
如果你是第一次配置 TTS,可以参照下面这个流程:
|
||||
|
||||
1. 在 `model_list` 里添加一个支持 TTS 的模型。
|
||||
2. 用 `voice.tts_model_name` 指向这个模型。
|
||||
3. 在 `.security.yml` 里配置对应的 API Key。
|
||||
|
||||
## 快速推荐
|
||||
|
||||
对于大多数用户,建议优先从下面两种开始:
|
||||
|
||||
| 提供商 | 推荐理由 |
|
||||
| --- | --- |
|
||||
| [OpenAI](https://platform.openai.com/docs/guides/text-to-speech) | 这是 PicoClaw 当前最稳定、最直接的 TTS 路径。当前实现就是围绕 OpenAI 兼容的 `/audio/speech` 接口格式构建的,所以 OpenAI 是最稳妥的默认选择。 |
|
||||
| [Xiaomi MiMo](https://platform.xiaomimimo.com) | 由于响应速度和语音音色对于中国用户更友好,MiMo 是一个不错的第二选择。 |
|
||||
|
||||
## TTS 配置是如何工作的
|
||||
|
||||
PicoClaw 不会把 TTS 的 API Key 放在 `voice` 配置里。
|
||||
|
||||
推荐方式是:
|
||||
|
||||
- `voice.tts_model_name` 用来选择 `model_list` 里的某个命名模型。
|
||||
- 对应的 `model_list` 条目提供真实的 provider、model ID、`api_base` 和代理配置。
|
||||
- `.security.yml` 负责保存该模型条目的 API Key。
|
||||
|
||||
这是当前推荐且受支持的配置方式。
|
||||
|
||||
## 推荐配置方式
|
||||
|
||||
### 方案 A:OpenAI
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"tts_model_name": "openai-tts"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "openai-tts",
|
||||
"model": "openai/tts-1"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
openai-tts:
|
||||
api_keys:
|
||||
- "sk-openai-your-key"
|
||||
```
|
||||
|
||||
### 方案 B:Xiaomi MiMo
|
||||
|
||||
`config.json`
|
||||
|
||||
```json
|
||||
{
|
||||
"voice": {
|
||||
"tts_model_name": "mimo-tts"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "mimo-tts",
|
||||
"model": "mimo/mimo-v2-tts"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`.security.yml`
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
mimo-tts:
|
||||
api_keys:
|
||||
- "your-mimo-key"
|
||||
```
|
||||
|
||||
如果你使用自定义的 MiMo 接口地址,也可以显式设置 `api_base`。如果不设置,PicoClaw 会自动使用该 provider 的默认地址。
|
||||
|
||||
## PicoClaw 当前实际发送的 TTS 请求
|
||||
|
||||
当前 TTS 运行时使用的是 OpenAI 兼容的语音合成请求,并带有以下默认值:
|
||||
|
||||
- Endpoint:`/audio/speech`
|
||||
- 返回格式:`opus`
|
||||
- Voice:`alloy`
|
||||
- Model:来自你所选中的 `model_list` 条目
|
||||
|
||||
这意味着:
|
||||
|
||||
- `openai/tts-1` 可以自然工作。
|
||||
- 其他 OpenAI 兼容 provider 也可能可用,前提是它们接受相同的请求格式。
|
||||
- PicoClaw 目前还没有对用户暴露一个配置项来修改 TTS voice,当前固定为 `alloy`。
|
||||
|
||||
## PicoClaw 如何选择 TTS Provider
|
||||
|
||||
`DetectTTS` 会按下面顺序选择 TTS:
|
||||
|
||||
1. **首选路径**:根据 `voice.tts_model_name` 在 `model_list` 中找到对应模型。
|
||||
2. 如果找到了匹配条目,并且它有 API Key,PicoClaw 就会使用这个模型条目的配置创建一个 OpenAI 兼容的 TTS provider。
|
||||
3. **回退路径**:如果没有设置 `voice.tts_model_name`,或者该名字无法解析,PicoClaw 会扫描 `model_list`,选中第一个模型字符串里包含 `tts` 且带有 API Key 的条目。
|
||||
|
||||
回退扫描只是为了兼容旧行为。新配置建议始终显式设置 `voice.tts_model_name`。
|
||||
|
||||
## 关于 API Base 的处理方式
|
||||
|
||||
PicoClaw 会对 TTS 的 `api_base` 做规范化处理:
|
||||
|
||||
- 对 OpenAI 来说,像 `https://api.openai.com` 或 `https://api.openai.com/v1` 这样的地址,会自动变成 `https://api.openai.com/v1/audio/speech`。
|
||||
- 对其他 OpenAI 兼容 provider,PicoClaw 会尽量保留你提供的基础路径,只确保它最终以 `/audio/speech` 结尾。
|
||||
- 如果没有设置 `api_base`,并且模型前缀是已知 provider,PicoClaw 会自动使用该 provider 的默认地址。
|
||||
|
||||
## 常见错误
|
||||
|
||||
- `voice.tts_model_name` 指向了一个不存在的 `model_list` 名称。
|
||||
- 在 `model_list` 里定义了 TTS 模型,但忘了在 `.security.yml` 中配置对应 API Key。
|
||||
- 误以为 PicoClaw 会自动支持 provider 自定义 voice 参数。
|
||||
- 使用了不兼容 OpenAI `/audio/speech` 请求格式的接口地址。
|
||||
|
||||
## 最小检查清单
|
||||
|
||||
在测试 `send_tts` 之前,请确认:
|
||||
|
||||
- `voice.tts_model_name` 能正确匹配某个 `model_list[].model_name`。
|
||||
- `.security.yml` 中对应条目已经配置了有效 API Key。
|
||||
- 你所选的 provider 支持 OpenAI 兼容的语音合成接口。
|
||||
- 你选择的模型本身确实支持 TTS。
|
||||
162
pkg/audio/tts/mimo_tts.go
Normal file
162
pkg/audio/tts/mimo_tts.go
Normal file
|
|
@ -0,0 +1,162 @@
|
|||
package tts
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
type MimoTTSProvider struct {
|
||||
apiKey string
|
||||
apiBase string
|
||||
voice string
|
||||
format string
|
||||
model string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewMimoTTSProvider(apiKey string, apiBase string, model string, proxyURL string) *MimoTTSProvider {
|
||||
if apiBase == "" {
|
||||
apiBase = "https://api.xiaomimimo.com/v1/chat/completions"
|
||||
} else {
|
||||
if u, err := url.Parse(apiBase); err == nil && u.Scheme != "" && u.Host != "" {
|
||||
path := u.Path
|
||||
if u.Host == "api.xiaomimimo.com" {
|
||||
if path == "" || path == "/" || path == "/v1" || path == "/v1/" {
|
||||
path = "/v1/chat/completions"
|
||||
} else {
|
||||
if !strings.HasPrefix(path, "/") {
|
||||
path = "/" + path
|
||||
}
|
||||
if !strings.HasPrefix(path, "/v1/") {
|
||||
path = "/v1" + strings.TrimSuffix(path, "/")
|
||||
}
|
||||
if !strings.HasSuffix(path, "/chat/completions") {
|
||||
path = strings.TrimSuffix(path, "/") + "/chat/completions"
|
||||
}
|
||||
}
|
||||
} else {
|
||||
if !strings.HasSuffix(path, "/chat/completions") {
|
||||
path = strings.TrimSuffix(path, "/") + "/chat/completions"
|
||||
}
|
||||
}
|
||||
u.Path = path
|
||||
apiBase = u.String()
|
||||
} else {
|
||||
if apiBase == "https://api.xiaomimimo.com/v1" {
|
||||
apiBase = "https://api.xiaomimimo.com/v1/chat/completions"
|
||||
} else if !strings.HasSuffix(apiBase, "/chat/completions") {
|
||||
apiBase = strings.TrimSuffix(apiBase, "/") + "/chat/completions"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
model = strings.TrimSpace(model)
|
||||
if model == "" {
|
||||
model = "mimo-v2-tts"
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 60 * time.Second}
|
||||
if proxyURL != "" {
|
||||
if pURL, err := url.Parse(proxyURL); err == nil {
|
||||
client.Transport = &http.Transport{Proxy: http.ProxyURL(pURL)}
|
||||
} else {
|
||||
logger.WarnF(
|
||||
"NewMimoTTSProvider: invalid proxy URL; proceeding without proxy",
|
||||
map[string]any{"proxyURL": proxyURL, "error": err},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return &MimoTTSProvider{
|
||||
apiKey: apiKey,
|
||||
apiBase: apiBase,
|
||||
voice: "default_zh", // mimo_default now seems to be an alias for default_en, which is not working for Chinese TTS. default_zh seems to work fine with both English and Chinese, and is likely the intended default for TTS.
|
||||
format: "mp3",
|
||||
model: model,
|
||||
httpClient: client,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *MimoTTSProvider) Name() string {
|
||||
return "mimo-tts"
|
||||
}
|
||||
|
||||
func (t *MimoTTSProvider) Synthesize(ctx context.Context, text string) (io.ReadCloser, error) {
|
||||
logger.DebugCF("voice-tts", "Starting TTS synthesis", map[string]any{"text_len": len(text), "provider": t.Name()})
|
||||
|
||||
reqBody := map[string]any{
|
||||
"model": t.model,
|
||||
"messages": []map[string]string{
|
||||
{"role": "assistant", "content": text},
|
||||
},
|
||||
"audio": map[string]string{
|
||||
"format": t.format,
|
||||
"voice": t.voice,
|
||||
},
|
||||
"stream": false,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", t.apiBase, bytes.NewReader(jsonData))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Api-Key", t.apiKey)
|
||||
|
||||
resp, err := t.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var payload struct {
|
||||
Choices []struct {
|
||||
Message struct {
|
||||
Audio struct {
|
||||
Data string `json:"data"`
|
||||
} `json:"audio"`
|
||||
} `json:"message"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
|
||||
err = json.Unmarshal(body, &payload)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decode response: %w", err)
|
||||
}
|
||||
|
||||
if len(payload.Choices) == 0 || payload.Choices[0].Message.Audio.Data == "" {
|
||||
return nil, fmt.Errorf("invalid TTS response: missing audio data")
|
||||
}
|
||||
|
||||
audioBytes, err := base64.StdEncoding.DecodeString(payload.Choices[0].Message.Audio.Data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decode audio data: %w", err)
|
||||
}
|
||||
|
||||
return io.NopCloser(bytes.NewReader(audioBytes)), nil
|
||||
}
|
||||
126
pkg/audio/tts/openai_tts.go
Normal file
126
pkg/audio/tts/openai_tts.go
Normal file
|
|
@ -0,0 +1,126 @@
|
|||
package tts
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers/common"
|
||||
)
|
||||
|
||||
type OpenAITTSProvider struct {
|
||||
apiKey string
|
||||
apiBase string
|
||||
voice string
|
||||
model string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewOpenAITTSProvider(apiKey string, apiBase string, proxyURL string, model string) *OpenAITTSProvider {
|
||||
// Normalize apiBase to avoid malformed endpoints like
|
||||
// "https://api.openai.com/audio/speech" when "/v1" is required.
|
||||
if apiBase == "" {
|
||||
apiBase = "https://api.openai.com/v1/audio/speech"
|
||||
} else {
|
||||
if u, err := url.Parse(apiBase); err == nil && u.Scheme != "" && u.Host != "" {
|
||||
path := u.Path
|
||||
if u.Host == "api.openai.com" {
|
||||
// For the official OpenAI host, ensure exactly one /v1 prefix and
|
||||
// that the path ends with /audio/speech.
|
||||
if path == "" || path == "/" || path == "/v1" {
|
||||
path = "/v1/audio/speech"
|
||||
} else {
|
||||
if !strings.HasPrefix(path, "/") {
|
||||
path = "/" + path
|
||||
}
|
||||
if !strings.HasPrefix(path, "/v1/") {
|
||||
path = "/v1" + strings.TrimSuffix(path, "/")
|
||||
}
|
||||
if !strings.HasSuffix(path, "/audio/speech") {
|
||||
path = strings.TrimSuffix(path, "/") + "/audio/speech"
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// For non-OpenAI hosts (e.g., proxies), preserve the existing base
|
||||
// path and only ensure it ends with /audio/speech.
|
||||
if !strings.HasSuffix(path, "/audio/speech") {
|
||||
path = strings.TrimSuffix(path, "/") + "/audio/speech"
|
||||
}
|
||||
}
|
||||
u.Path = path
|
||||
apiBase = u.String()
|
||||
} else {
|
||||
// Fallback to the previous string-based behavior if parsing fails.
|
||||
if apiBase == "https://api.openai.com/v1" {
|
||||
apiBase = "https://api.openai.com/v1/audio/speech"
|
||||
} else if !strings.HasSuffix(apiBase, "/audio/speech") {
|
||||
// Just in case they provide openrouter base or standard base
|
||||
apiBase = strings.TrimSuffix(apiBase, "/") + "/audio/speech"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
client := common.NewHTTPClient(proxyURL)
|
||||
client.Timeout = 60 * time.Second
|
||||
|
||||
model = strings.TrimSpace(model)
|
||||
if model == "" {
|
||||
model = "tts-1"
|
||||
}
|
||||
|
||||
return &OpenAITTSProvider{
|
||||
apiKey: apiKey,
|
||||
apiBase: apiBase,
|
||||
voice: "alloy",
|
||||
model: model,
|
||||
httpClient: client,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *OpenAITTSProvider) Name() string {
|
||||
return "openai-tts"
|
||||
}
|
||||
|
||||
func (t *OpenAITTSProvider) Synthesize(ctx context.Context, text string) (io.ReadCloser, error) {
|
||||
logger.DebugCF("voice-tts", "Starting TTS synthesis", map[string]any{"text_len": len(text)})
|
||||
|
||||
reqBody := map[string]any{
|
||||
"model": t.model,
|
||||
"input": text,
|
||||
"voice": t.voice,
|
||||
"response_format": "opus",
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", t.apiBase, bytes.NewReader(jsonData))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
||||
|
||||
resp, err := t.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
return resp.Body, nil
|
||||
}
|
||||
151
pkg/audio/tts/tts.go
Normal file
151
pkg/audio/tts/tts.go
Normal file
|
|
@ -0,0 +1,151 @@
|
|||
package tts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
type TTSProvider interface {
|
||||
Name() string
|
||||
Synthesize(ctx context.Context, text string) (io.ReadCloser, error)
|
||||
}
|
||||
|
||||
func providerFromModelConfig(mc *config.ModelConfig) TTSProvider {
|
||||
if mc == nil || mc.APIKey() == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
protocol, modelID := providers.ExtractProtocol(mc.Model)
|
||||
if modelID == "" {
|
||||
modelID = strings.TrimSpace(mc.Model)
|
||||
}
|
||||
|
||||
switch protocol {
|
||||
case "mimo":
|
||||
return NewMimoTTSProvider(mc.APIKey(), providers.ResolveAPIBase(mc), modelID, mc.Proxy)
|
||||
default:
|
||||
return NewOpenAITTSProvider(mc.APIKey(), providers.ResolveAPIBase(mc), mc.Proxy, modelID)
|
||||
}
|
||||
}
|
||||
|
||||
func DetectTTS(cfg *config.Config) TTSProvider {
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if modelName := strings.TrimSpace(cfg.Voice.TTSModelName); modelName != "" {
|
||||
if mc, err := cfg.GetModelConfig(modelName); err == nil {
|
||||
if provider := providerFromModelConfig(mc); provider != nil {
|
||||
return provider
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, mc := range cfg.ModelList {
|
||||
if strings.Contains(strings.ToLower(mc.Model), "tts") && mc.APIKey() != "" {
|
||||
if provider := providerFromModelConfig(mc); provider != nil {
|
||||
return provider
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SynthesizeAndStore synthesizes text to speech and registers it in the media store, returning the media reference.
|
||||
func SynthesizeAndStore(
|
||||
ctx context.Context,
|
||||
provider TTSProvider,
|
||||
store media.MediaStore,
|
||||
text string,
|
||||
filename string,
|
||||
channel string,
|
||||
chatID string,
|
||||
) (string, error) {
|
||||
if provider == nil {
|
||||
return "", fmt.Errorf("tts provider is not configured")
|
||||
}
|
||||
if store == nil {
|
||||
return "", fmt.Errorf("media store not configured")
|
||||
}
|
||||
if channel == "" || chatID == "" {
|
||||
return "", fmt.Errorf("no target channel/chat available")
|
||||
}
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return "", fmt.Errorf("text is required")
|
||||
}
|
||||
|
||||
stream, err := provider.Synthesize(ctx, text)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("tts synthesize failed: %w", err)
|
||||
}
|
||||
defer stream.Close()
|
||||
|
||||
err = os.MkdirAll(media.TempDir(), 0o700)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create media temp dir: %w", err)
|
||||
}
|
||||
|
||||
fileExt := ".ogg"
|
||||
contentType := "audio/ogg"
|
||||
if provider.Name() == "mimo-tts" {
|
||||
fileExt = ".mp3"
|
||||
contentType = "audio/mpeg"
|
||||
}
|
||||
|
||||
file, err := os.CreateTemp(media.TempDir(), "tts-*"+fileExt)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create temp file: %w", err)
|
||||
}
|
||||
|
||||
removeTemp := true
|
||||
defer func() {
|
||||
if removeTemp {
|
||||
_ = os.Remove(file.Name())
|
||||
}
|
||||
}()
|
||||
|
||||
_, err = io.Copy(file, stream)
|
||||
if err != nil {
|
||||
file.Close()
|
||||
return "", fmt.Errorf("failed to write tts audio: %w", err)
|
||||
}
|
||||
|
||||
err = file.Close()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to close tts audio file: %w", err)
|
||||
}
|
||||
|
||||
filename = strings.TrimSpace(filename)
|
||||
if filename == "" {
|
||||
filename = fmt.Sprintf("tts-%d%s", time.Now().Unix(), fileExt)
|
||||
}
|
||||
|
||||
ext := strings.ToLower(filepath.Ext(filename))
|
||||
if ext == "" {
|
||||
filename += fileExt
|
||||
} else if ext != fileExt {
|
||||
filename = strings.TrimSuffix(filename, filepath.Ext(filename)) + fileExt
|
||||
}
|
||||
|
||||
scope := fmt.Sprintf("tool:send_tts:%s:%s:%d", channel, chatID, time.Now().UnixNano())
|
||||
ref, err := store.Store(file.Name(), media.MediaMeta{
|
||||
Filename: filename,
|
||||
ContentType: contentType,
|
||||
Source: "tool:send_tts",
|
||||
}, scope)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to register audio: %w", err)
|
||||
}
|
||||
removeTemp = false
|
||||
|
||||
return ref, nil
|
||||
}
|
||||
247
pkg/audio/tts/tts_test.go
Normal file
247
pkg/audio/tts/tts_test.go
Normal file
|
|
@ -0,0 +1,247 @@
|
|||
package tts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
)
|
||||
|
||||
func TestNewOpenAITTSProvider_APIBaseNormalization(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
input string
|
||||
expect string
|
||||
}{
|
||||
{
|
||||
name: "empty base",
|
||||
input: "",
|
||||
expect: "https://api.openai.com/v1/audio/speech",
|
||||
},
|
||||
{
|
||||
name: "official host no path",
|
||||
input: "https://api.openai.com",
|
||||
expect: "https://api.openai.com/v1/audio/speech",
|
||||
},
|
||||
{
|
||||
name: "official host v1",
|
||||
input: "https://api.openai.com/v1",
|
||||
expect: "https://api.openai.com/v1/audio/speech",
|
||||
},
|
||||
{
|
||||
name: "official host v1 slash",
|
||||
input: "https://api.openai.com/v1/",
|
||||
expect: "https://api.openai.com/v1/audio/speech",
|
||||
},
|
||||
{
|
||||
name: "non-openai host preserves base path",
|
||||
input: "https://proxy.example.com/base",
|
||||
expect: "https://proxy.example.com/base/audio/speech",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
provider := NewOpenAITTSProvider("key", tc.input, "", "")
|
||||
if provider.apiBase != tc.expect {
|
||||
t.Fatalf("apiBase mismatch: got %q, want %q", provider.apiBase, tc.expect)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAITTSProvider_SynthesizeSuccess(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var gotPath string
|
||||
var gotAuth string
|
||||
var gotContentType string
|
||||
var gotBody map[string]any
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
gotAuth = r.Header.Get("Authorization")
|
||||
gotContentType = r.Header.Get("Content-Type")
|
||||
|
||||
bodyBytes, _ := io.ReadAll(r.Body)
|
||||
_ = r.Body.Close()
|
||||
_ = json.Unmarshal(bodyBytes, &gotBody)
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte("audio-bytes"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewOpenAITTSProvider("k123", server.URL, "", "")
|
||||
stream, err := provider.Synthesize(context.Background(), "hello")
|
||||
if err != nil {
|
||||
t.Fatalf("Synthesize failed: %v", err)
|
||||
}
|
||||
defer stream.Close()
|
||||
|
||||
data, err := io.ReadAll(stream)
|
||||
if err != nil {
|
||||
t.Fatalf("read stream failed: %v", err)
|
||||
}
|
||||
|
||||
if gotPath != "/audio/speech" {
|
||||
t.Fatalf("request path mismatch: got %q", gotPath)
|
||||
}
|
||||
if gotAuth != "Bearer k123" {
|
||||
t.Fatalf("authorization mismatch: got %q", gotAuth)
|
||||
}
|
||||
if gotContentType != "application/json" {
|
||||
t.Fatalf("content-type mismatch: got %q", gotContentType)
|
||||
}
|
||||
if gotBody["model"] != "tts-1" || gotBody["voice"] != "alloy" || gotBody["response_format"] != "opus" ||
|
||||
gotBody["input"] != "hello" {
|
||||
bodyJSON, _ := json.Marshal(gotBody)
|
||||
t.Fatalf("request body mismatch: %s", string(bodyJSON))
|
||||
}
|
||||
if string(data) != "audio-bytes" {
|
||||
t.Fatalf("response body mismatch: got %q", string(data))
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAITTSProvider_SynthesizeNon200(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_, _ = w.Write([]byte("nope"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewOpenAITTSProvider("k123", server.URL, "", "")
|
||||
_, err := provider.Synthesize(context.Background(), "hello")
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "API error (status 500): nope") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewOpenAITTSProvider_UsesConfiguredModel(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
provider := NewOpenAITTSProvider("key", "https://api.xiaomimimo.com/v1", "", "mimo-v2-tts")
|
||||
if provider.model != "mimo-v2-tts" {
|
||||
t.Fatalf("model mismatch: got %q, want %q", provider.model, "mimo-v2-tts")
|
||||
}
|
||||
if provider.apiBase != "https://api.xiaomimimo.com/v1/audio/speech" {
|
||||
t.Fatalf("apiBase mismatch: got %q", provider.apiBase)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectTTS_UsesMimoProviderForMimoModels(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
provider := DetectTTS(&config.Config{
|
||||
Voice: config.VoiceConfig{TTSModelName: "mimo-tts"},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "mimo-tts",
|
||||
Model: "mimo/mimo-v2-tts",
|
||||
APIKeys: config.SimpleSecureStrings("sk-mimo"),
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
ttsProvider, ok := provider.(*MimoTTSProvider)
|
||||
if !ok {
|
||||
t.Fatalf("DetectTTS() type = %T, want *MimoTTSProvider", provider)
|
||||
}
|
||||
if ttsProvider.model != "mimo-v2-tts" {
|
||||
t.Fatalf("model mismatch: got %q, want %q", ttsProvider.model, "mimo-v2-tts")
|
||||
}
|
||||
if ttsProvider.apiBase != "https://api.xiaomimimo.com/v1/chat/completions" {
|
||||
t.Fatalf("apiBase mismatch: got %q", ttsProvider.apiBase)
|
||||
}
|
||||
}
|
||||
|
||||
type stubTTSProvider struct {
|
||||
name string
|
||||
}
|
||||
|
||||
func (s stubTTSProvider) Name() string {
|
||||
return s.name
|
||||
}
|
||||
|
||||
func (s stubTTSProvider) Synthesize(ctx context.Context, text string) (io.ReadCloser, error) {
|
||||
return io.NopCloser(strings.NewReader("audio")), nil
|
||||
}
|
||||
|
||||
func TestSynthesizeAndStore_UsesOggMetadataByDefault(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
store := media.NewFileMediaStore()
|
||||
ref, err := SynthesizeAndStore(
|
||||
context.Background(),
|
||||
stubTTSProvider{name: "openai-tts"},
|
||||
store,
|
||||
"hello",
|
||||
"",
|
||||
"discord",
|
||||
"chat123",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("SynthesizeAndStore failed: %v", err)
|
||||
}
|
||||
|
||||
path, meta, err := store.ResolveWithMeta(ref)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveWithMeta failed: %v", err)
|
||||
}
|
||||
if meta.ContentType != "audio/ogg" {
|
||||
t.Fatalf("ContentType = %q, want %q", meta.ContentType, "audio/ogg")
|
||||
}
|
||||
if filepath.Ext(path) != ".ogg" {
|
||||
t.Fatalf("stored file extension = %q, want %q", filepath.Ext(path), ".ogg")
|
||||
}
|
||||
if filepath.Ext(meta.Filename) != ".ogg" {
|
||||
t.Fatalf("filename extension = %q, want %q", filepath.Ext(meta.Filename), ".ogg")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSynthesizeAndStore_UsesMp3MetadataForMimo(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
store := media.NewFileMediaStore()
|
||||
ref, err := SynthesizeAndStore(
|
||||
context.Background(),
|
||||
stubTTSProvider{name: "mimo-tts"},
|
||||
store,
|
||||
"hello",
|
||||
"",
|
||||
"discord",
|
||||
"chat123",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("SynthesizeAndStore failed: %v", err)
|
||||
}
|
||||
|
||||
path, meta, err := store.ResolveWithMeta(ref)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveWithMeta failed: %v", err)
|
||||
}
|
||||
if meta.ContentType != "audio/mpeg" {
|
||||
t.Fatalf("ContentType = %q, want %q", meta.ContentType, "audio/mpeg")
|
||||
}
|
||||
if filepath.Ext(path) != ".mp3" {
|
||||
t.Fatalf("stored file extension = %q, want %q", filepath.Ext(path), ".mp3")
|
||||
}
|
||||
if filepath.Ext(meta.Filename) != ".mp3" {
|
||||
t.Fatalf("filename extension = %q, want %q", filepath.Ext(meta.Filename), ".mp3")
|
||||
}
|
||||
}
|
||||
|
|
@ -34,6 +34,8 @@ type MessageBus struct {
|
|||
inbound chan InboundMessage
|
||||
outbound chan OutboundMessage
|
||||
outboundMedia chan OutboundMediaMessage
|
||||
audioChunks chan AudioChunk
|
||||
voiceControls chan VoiceControl
|
||||
|
||||
closeOnce sync.Once
|
||||
done chan struct{}
|
||||
|
|
@ -47,6 +49,8 @@ func NewMessageBus() *MessageBus {
|
|||
inbound: make(chan InboundMessage, defaultBusBufferSize),
|
||||
outbound: make(chan OutboundMessage, defaultBusBufferSize),
|
||||
outboundMedia: make(chan OutboundMediaMessage, defaultBusBufferSize),
|
||||
audioChunks: make(chan AudioChunk, defaultBusBufferSize*4), // Audio chunks need more buffer
|
||||
voiceControls: make(chan VoiceControl, defaultBusBufferSize),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
|
@ -103,6 +107,22 @@ func (mb *MessageBus) OutboundMediaChan() <-chan OutboundMediaMessage {
|
|||
return mb.outboundMedia
|
||||
}
|
||||
|
||||
func (mb *MessageBus) PublishAudioChunk(ctx context.Context, chunk AudioChunk) error {
|
||||
return publish(ctx, mb, mb.audioChunks, chunk)
|
||||
}
|
||||
|
||||
func (mb *MessageBus) AudioChunksChan() <-chan AudioChunk {
|
||||
return mb.audioChunks
|
||||
}
|
||||
|
||||
func (mb *MessageBus) PublishVoiceControl(ctx context.Context, ctrl VoiceControl) error {
|
||||
return publish(ctx, mb, mb.voiceControls, ctrl)
|
||||
}
|
||||
|
||||
func (mb *MessageBus) VoiceControlsChan() <-chan VoiceControl {
|
||||
return mb.voiceControls
|
||||
}
|
||||
|
||||
// SetStreamDelegate registers a StreamDelegate (typically the channel Manager).
|
||||
func (mb *MessageBus) SetStreamDelegate(d StreamDelegate) {
|
||||
mb.streamDelegate.Store(d)
|
||||
|
|
@ -132,6 +152,8 @@ func (mb *MessageBus) Close() {
|
|||
close(mb.inbound)
|
||||
close(mb.outbound)
|
||||
close(mb.outboundMedia)
|
||||
close(mb.audioChunks)
|
||||
close(mb.voiceControls)
|
||||
|
||||
// clean up any remaining messages in channels
|
||||
drained := 0
|
||||
|
|
@ -144,6 +166,12 @@ func (mb *MessageBus) Close() {
|
|||
for range mb.outboundMedia {
|
||||
drained++
|
||||
}
|
||||
for range mb.audioChunks {
|
||||
drained++
|
||||
}
|
||||
for range mb.voiceControls {
|
||||
drained++
|
||||
}
|
||||
|
||||
if drained > 0 {
|
||||
logger.DebugCF("bus", "Drained buffered messages during close", map[string]any{
|
||||
|
|
|
|||
|
|
@ -30,10 +30,11 @@ type InboundMessage struct {
|
|||
}
|
||||
|
||||
type OutboundMessage struct {
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||
Metadata map[string]string `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// MediaPart describes a single media attachment to send.
|
||||
|
|
@ -51,3 +52,25 @@ type OutboundMediaMessage struct {
|
|||
ChatID string `json:"chat_id"`
|
||||
Parts []MediaPart `json:"parts"`
|
||||
}
|
||||
|
||||
// AudioChunk represents a chunk of streaming voice data.
|
||||
type AudioChunk struct {
|
||||
SessionID string `json:"session_id"`
|
||||
SpeakerID string `json:"speaker_id"` // User ID or SSRC
|
||||
ChatID string `json:"chat_id"` // Where to respond
|
||||
Channel string `json:"channel"` // Source channel type (e.g. "discord")
|
||||
Sequence uint64 `json:"sequence"`
|
||||
Timestamp uint32 `json:"timestamp"`
|
||||
SampleRate int `json:"sample_rate"`
|
||||
Channels int `json:"channels"`
|
||||
Format string `json:"format"` // "opus", "pcm", etc
|
||||
Data []byte `json:"data"`
|
||||
}
|
||||
|
||||
// VoiceControl represents state or commands for voice sessions.
|
||||
type VoiceControl struct {
|
||||
SessionID string `json:"session_id"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Type string `json:"type"` // "state", "command"
|
||||
Action string `json:"action"` // "idle", "listening", "start", "stop", "leave"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ package discord
|
|||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
|
|
@ -14,6 +15,8 @@ import (
|
|||
"github.com/bwmarrin/discordgo"
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio"
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
|
|
@ -42,6 +45,15 @@ type DiscordChannel struct {
|
|||
typingMu sync.Mutex
|
||||
typingStop map[string]chan struct{} // chatID → stop signal
|
||||
botUserID string // stored for mention checking
|
||||
bus *bus.MessageBus
|
||||
tts tts.TTSProvider
|
||||
voiceMu sync.RWMutex
|
||||
voiceSSRC map[string]map[uint32]string // guildID -> ssrc -> userID
|
||||
|
||||
// TTS interruption: cancel active playback when user speaks
|
||||
ttsMu sync.Mutex
|
||||
cancelTTS context.CancelFunc
|
||||
ttsPlayID uint64
|
||||
}
|
||||
|
||||
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
||||
|
|
@ -73,6 +85,8 @@ func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordC
|
|||
config: cfg,
|
||||
ctx: context.Background(),
|
||||
typingStop: make(map[string]chan struct{}),
|
||||
bus: bus,
|
||||
voiceSSRC: make(map[string]map[uint32]string),
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
|
@ -90,6 +104,8 @@ func (c *DiscordChannel) Start(ctx context.Context) error {
|
|||
|
||||
c.session.AddHandler(c.handleMessage)
|
||||
|
||||
go c.listenVoiceControl(c.ctx)
|
||||
|
||||
if err := c.session.Open(); err != nil {
|
||||
return fmt.Errorf("failed to open discord session: %w", err)
|
||||
}
|
||||
|
|
@ -142,6 +158,25 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]s
|
|||
return nil, nil
|
||||
}
|
||||
|
||||
if c.tts != nil {
|
||||
if ch, err := c.session.State.Channel(channelID); err == nil && ch.GuildID != "" {
|
||||
if vc, ok := c.session.VoiceConnections[ch.GuildID]; ok && vc != nil {
|
||||
// Cancel any previous TTS playback
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
}
|
||||
ttsCtx, ttsCancel := context.WithCancel(c.ctx)
|
||||
c.ttsPlayID++
|
||||
playID := c.ttsPlayID
|
||||
c.cancelTTS = ttsCancel
|
||||
c.ttsMu.Unlock()
|
||||
|
||||
go c.playTTS(ttsCtx, vc, msg.Content, playID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
msgID, err := c.sendChunk(ctx, channelID, msg.Content, msg.ReplyToMessageID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
|
@ -359,6 +394,10 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
|||
return
|
||||
}
|
||||
|
||||
if c.handleVoiceCommand(s, m) {
|
||||
return
|
||||
}
|
||||
|
||||
content := m.Content
|
||||
|
||||
// In guild (group) channels, apply unified group trigger filtering
|
||||
|
|
@ -630,3 +669,134 @@ func (c *DiscordChannel) stripBotMention(text string) string {
|
|||
text = strings.ReplaceAll(text, fmt.Sprintf("<@!%s>", c.botUserID), "")
|
||||
return strings.TrimSpace(text)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) listenVoiceControl(ctx context.Context) {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ctrl, ok := <-c.bus.VoiceControlsChan():
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if ctrl.Type == "command" && ctrl.Action == "leave" {
|
||||
if strings.HasPrefix(ctrl.SessionID, "discord_vc_") {
|
||||
guildID := strings.TrimPrefix(ctrl.SessionID, "discord_vc_")
|
||||
vc, exists := c.session.VoiceConnections[guildID]
|
||||
if exists && vc != nil {
|
||||
vc.Disconnect(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) playTTS(ctx context.Context, vc *discordgo.VoiceConnection, text string, playID uint64) {
|
||||
// Capture the cancel func associated with this playback (if any).
|
||||
// Clear cancelTTS when playback finishes (normal or interrupted),
|
||||
// but only if it still refers to this playback's cancel func.
|
||||
defer func() {
|
||||
c.ttsMu.Lock()
|
||||
if c.ttsPlayID == playID {
|
||||
c.cancelTTS = nil
|
||||
}
|
||||
c.ttsMu.Unlock()
|
||||
}()
|
||||
|
||||
sentences := audio.SplitSentences(text)
|
||||
if len(sentences) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoCF("discord", "Starting streamed TTS", map[string]any{"sentences": len(sentences)})
|
||||
|
||||
// Pipeline: prefetch next sentence's audio while playing current
|
||||
type ttResult struct {
|
||||
stream io.ReadCloser
|
||||
err error
|
||||
}
|
||||
|
||||
var prefetch chan ttResult
|
||||
|
||||
// Ensure any in-flight prefetch is drained on exit to prevent stream leaks,
|
||||
// but avoid blocking indefinitely if the prefetch goroutine is stuck or never sends.
|
||||
defer func() {
|
||||
if prefetch != nil {
|
||||
select {
|
||||
case result := <-prefetch:
|
||||
if result.stream != nil {
|
||||
result.stream.Close()
|
||||
}
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
// Timed out waiting for a prefetched result; avoid blocking on exit.
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for i, sentence := range sentences {
|
||||
// Check for cancellation (interruption)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
logger.InfoCF("discord", "TTS interrupted", map[string]any{"at_sentence": i})
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
// Start prefetching the NEXT sentence while we process the current one
|
||||
var nextPrefetch chan ttResult
|
||||
if i+1 < len(sentences) {
|
||||
nextPrefetch = make(chan ttResult, 1)
|
||||
nextSentence := sentences[i+1]
|
||||
go func() {
|
||||
s, e := c.tts.Synthesize(ctx, nextSentence)
|
||||
nextPrefetch <- ttResult{s, e}
|
||||
}()
|
||||
}
|
||||
|
||||
// Get the current sentence's audio
|
||||
var stream io.ReadCloser
|
||||
var err error
|
||||
|
||||
if prefetch != nil {
|
||||
// Use prefetched result from previous iteration, but be responsive to cancellation.
|
||||
var result ttResult
|
||||
select {
|
||||
case result = <-prefetch:
|
||||
stream, err = result.stream, result.err
|
||||
case <-ctx.Done():
|
||||
// Context canceled while waiting for prefetched audio; abort playback.
|
||||
logger.InfoCF(
|
||||
"discord",
|
||||
"TTS interrupted while waiting for prefetched audio",
|
||||
map[string]any{"at_sentence": i},
|
||||
)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
// First sentence: synthesize directly
|
||||
stream, err = c.tts.Synthesize(ctx, sentence)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
if stream != nil {
|
||||
stream.Close()
|
||||
}
|
||||
logger.ErrorCF("discord", "TTS synthesize failed", map[string]any{"error": err.Error(), "sentence": i})
|
||||
prefetch = nextPrefetch
|
||||
continue
|
||||
}
|
||||
|
||||
if err := streamOggOpusToDiscord(ctx, vc, stream); err != nil {
|
||||
logger.ErrorCF("discord", "TTS playback failed", map[string]any{"error": err.Error(), "sentence": i})
|
||||
}
|
||||
stream.Close()
|
||||
|
||||
prefetch = nextPrefetch
|
||||
}
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *DiscordChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package discord
|
||||
|
||||
import (
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
|
|
@ -8,6 +9,10 @@ import (
|
|||
|
||||
func init() {
|
||||
channels.RegisterFactory("discord", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||
return NewDiscordChannel(cfg.Channels.Discord, b)
|
||||
ch, err := NewDiscordChannel(cfg.Channels.Discord, b)
|
||||
if err == nil {
|
||||
ch.tts = tts.DetectTTS(cfg)
|
||||
}
|
||||
return ch, err
|
||||
})
|
||||
}
|
||||
|
|
|
|||
313
pkg/channels/discord/voice.go
Normal file
313
pkg/channels/discord/voice.go
Normal file
|
|
@ -0,0 +1,313 @@
|
|||
package discord
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/bwmarrin/discordgo"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/identity"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
func (c *DiscordChannel) setVoiceUserID(guildID string, ssrc uint32, userID string) {
|
||||
if userID == "" {
|
||||
return
|
||||
}
|
||||
|
||||
c.voiceMu.Lock()
|
||||
defer c.voiceMu.Unlock()
|
||||
|
||||
ssrcMap, ok := c.voiceSSRC[guildID]
|
||||
if !ok {
|
||||
ssrcMap = make(map[uint32]string)
|
||||
c.voiceSSRC[guildID] = ssrcMap
|
||||
}
|
||||
ssrcMap[ssrc] = userID
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) voiceUserID(guildID string, ssrc uint32) string {
|
||||
c.voiceMu.RLock()
|
||||
defer c.voiceMu.RUnlock()
|
||||
|
||||
ssrcMap, ok := c.voiceSSRC[guildID]
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return ssrcMap[ssrc]
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) handleVoiceCommand(s *discordgo.Session, m *discordgo.MessageCreate) bool {
|
||||
if m.Content == "!vc join" {
|
||||
vs, err := s.State.VoiceState(m.GuildID, m.Author.ID)
|
||||
if err != nil || vs == nil {
|
||||
if _, sendErr := s.ChannelMessageSend(
|
||||
m.ChannelID,
|
||||
"You need to be in a voice channel first!",
|
||||
); sendErr != nil {
|
||||
logger.InfoCF("discord", "Failed to send voice channel requirement message", map[string]any{
|
||||
"channel": m.ChannelID,
|
||||
"error": sendErr,
|
||||
})
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
logger.InfoCF("discord", "Joining voice channel", map[string]any{"channel": vs.ChannelID})
|
||||
vc, err := s.ChannelVoiceJoin(c.ctx, m.GuildID, vs.ChannelID, false, false)
|
||||
if err != nil {
|
||||
if _, sendErr := s.ChannelMessageSend(
|
||||
m.ChannelID,
|
||||
fmt.Sprintf("Failed to join voice channel: %v", err),
|
||||
); sendErr != nil {
|
||||
logger.InfoCF("discord", "Failed to send voice join error message", map[string]any{
|
||||
"channel": m.ChannelID,
|
||||
"error": sendErr,
|
||||
})
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
go c.receiveVoice(vc, m.GuildID, m.ChannelID)
|
||||
if _, sendErr := s.ChannelMessageSend(
|
||||
m.ChannelID,
|
||||
"Joined Voice Channel! Listening for audio...",
|
||||
); sendErr != nil {
|
||||
logger.InfoCF("discord", "Failed to send voice join success message", map[string]any{
|
||||
"channel": m.ChannelID,
|
||||
"error": sendErr,
|
||||
})
|
||||
}
|
||||
return true
|
||||
} else if m.Content == "!vc leave" {
|
||||
vc, exists := s.VoiceConnections[m.GuildID]
|
||||
if exists && vc != nil {
|
||||
if err := vc.Disconnect(c.ctx); err != nil {
|
||||
logger.InfoCF("discord", "Failed to disconnect from voice channel", map[string]any{
|
||||
"guild": m.GuildID,
|
||||
"error": err,
|
||||
})
|
||||
}
|
||||
if _, sendErr := s.ChannelMessageSend(m.ChannelID, "Left Voice Channel."); sendErr != nil {
|
||||
logger.InfoCF("discord", "Failed to send voice leave success message", map[string]any{
|
||||
"channel": m.ChannelID,
|
||||
"error": sendErr,
|
||||
})
|
||||
}
|
||||
} else {
|
||||
if _, sendErr := s.ChannelMessageSend(m.ChannelID, "Not in a voice channel."); sendErr != nil {
|
||||
logger.InfoCF("discord", "Failed to send voice not-in-channel message", map[string]any{
|
||||
"channel": m.ChannelID,
|
||||
"error": sendErr,
|
||||
})
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func VoiceReceiveActive(vc *discordgo.VoiceConnection) bool {
|
||||
return vc != nil && vc.OpusRecv != nil
|
||||
}
|
||||
|
||||
func streamOggOpusToDiscord(ctx context.Context, vc *discordgo.VoiceConnection, r io.Reader) (retErr error) {
|
||||
// Recover from panic if vc.OpusSend is closed mid-send (e.g. on disconnect)
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
retErr = fmt.Errorf("voice connection closed during playback")
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for the speaking transition to register
|
||||
vc.Speaking(true)
|
||||
defer vc.Speaking(false)
|
||||
|
||||
return audio.DecodeOggOpus(r, func(frame []byte) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case vc.OpusSend <- frame:
|
||||
return nil
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) receiveVoice(vc *discordgo.VoiceConnection, guildID string, chatID string) {
|
||||
logger.InfoCF("discord", "Started listening for voice", map[string]any{"guild": guildID})
|
||||
|
||||
vc.AddHandler(func(_ *discordgo.VoiceConnection, vs *discordgo.VoiceSpeakingUpdate) {
|
||||
if vs == nil {
|
||||
return
|
||||
}
|
||||
c.setVoiceUserID(guildID, uint32(vs.SSRC), vs.UserID)
|
||||
})
|
||||
|
||||
defer func() {
|
||||
c.voiceMu.Lock()
|
||||
delete(c.voiceSSRC, guildID)
|
||||
c.voiceMu.Unlock()
|
||||
}()
|
||||
|
||||
go func(ctx context.Context, vc *discordgo.VoiceConnection) {
|
||||
// Recover from potential panics if OpusSend is closed mid-send.
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.WarnCF("discord", "Recovered from panic while sending wake-up frames", map[string]any{
|
||||
"error": rec,
|
||||
"guild": guildID,
|
||||
})
|
||||
}
|
||||
}()
|
||||
|
||||
// If the voice connection or OpusSend are not available, nothing to do.
|
||||
if vc == nil || vc.OpusSend == nil {
|
||||
return
|
||||
}
|
||||
|
||||
time.Sleep(250 * time.Millisecond) // Wait a bit for connection to settle
|
||||
|
||||
// Abort if the context has already been canceled.
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
vc.Speaking(true)
|
||||
defer vc.Speaking(false)
|
||||
|
||||
silenceFrame := []byte{0xF8, 0xFF, 0xFE}
|
||||
for i := 0; i < 5; i++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case vc.OpusSend <- silenceFrame:
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
|
||||
logger.DebugCF("discord", "Sent wake-up silence frames", map[string]any{"guild": guildID})
|
||||
}(c.ctx, vc)
|
||||
sessionID := fmt.Sprintf("discord_vc_%s", guildID)
|
||||
|
||||
c.bus.PublishVoiceControl(c.ctx, bus.VoiceControl{
|
||||
SessionID: sessionID,
|
||||
Type: "state",
|
||||
Action: "listening",
|
||||
})
|
||||
|
||||
var sequence uint64 = 0
|
||||
var interruptCount int
|
||||
var lastInterruptAt time.Time
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return
|
||||
case p, ok := <-vc.OpusRecv:
|
||||
if !ok {
|
||||
logger.InfoCF("discord", "Voice channel closed", map[string]any{"guild": guildID})
|
||||
// Cancel any TTS that may still be playing
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
c.cancelTTS = nil
|
||||
}
|
||||
c.ttsMu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
if p == nil {
|
||||
logger.DebugCF("discord", "Received nil Opus packet", nil)
|
||||
continue
|
||||
}
|
||||
|
||||
if len(p.Opus) == 0 {
|
||||
logger.DebugCF("discord", "Received empty Opus packet", map[string]any{
|
||||
"seq": p.Sequence,
|
||||
"ssrc": p.SSRC,
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
logger.DebugCF("discord", "Received Opus packet", map[string]any{
|
||||
"seq": p.Sequence,
|
||||
"len": len(p.Opus),
|
||||
"ssrc": p.SSRC,
|
||||
})
|
||||
// Interruption detection: if user sends voice while TTS is playing,
|
||||
// cancel TTS after a short debounce (3 packets in 200ms)
|
||||
now := time.Now()
|
||||
if now.Sub(lastInterruptAt) > 500*time.Millisecond {
|
||||
interruptCount = 0
|
||||
}
|
||||
interruptCount++
|
||||
lastInterruptAt = now
|
||||
|
||||
if interruptCount >= 3 {
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
c.cancelTTS = nil
|
||||
logger.InfoCF("discord", "TTS interrupted by user voice", nil)
|
||||
}
|
||||
c.ttsMu.Unlock()
|
||||
interruptCount = 0
|
||||
}
|
||||
|
||||
userID := c.voiceUserID(guildID, p.SSRC)
|
||||
if userID == "" {
|
||||
logger.DebugCF("discord", "Dropping voice packet without user mapping", map[string]any{
|
||||
"ssrc": p.SSRC,
|
||||
"guild": guildID,
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
sender := bus.SenderInfo{
|
||||
Platform: "discord",
|
||||
PlatformID: userID,
|
||||
CanonicalID: identity.BuildCanonicalID("discord", userID),
|
||||
}
|
||||
if !c.IsAllowedSender(sender) {
|
||||
logger.DebugCF("discord", "Voice packet rejected by allowlist", map[string]any{
|
||||
"user_id": userID,
|
||||
"guild": guildID,
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
sequence++
|
||||
|
||||
chunk := bus.AudioChunk{
|
||||
SessionID: sessionID,
|
||||
SpeakerID: userID,
|
||||
ChatID: chatID,
|
||||
Channel: "discord",
|
||||
Sequence: sequence,
|
||||
Timestamp: p.Timestamp,
|
||||
SampleRate: 48000,
|
||||
Channels: 2,
|
||||
Format: "opus",
|
||||
Data: p.Opus,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(c.ctx, 100*time.Millisecond)
|
||||
err := c.bus.PublishAudioChunk(ctx, chunk)
|
||||
cancel()
|
||||
if err != nil {
|
||||
logger.ErrorCF("discord", "Failed to publish audio chunk", map[string]any{
|
||||
"guild": guildID,
|
||||
"sessionID": sessionID,
|
||||
"sequence": sequence,
|
||||
"error": err.Error(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -6,6 +6,8 @@ import (
|
|||
"strings"
|
||||
|
||||
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
)
|
||||
|
||||
// mentionPlaceholderRegex matches @_user_N placeholders inserted by Feishu for mentions.
|
||||
|
|
@ -145,3 +147,8 @@ func extractImageKeysRecursive(v any, feishuKeys, externalURLs *[]string) {
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *FeishuChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -684,3 +684,8 @@ func (c *LINEChannel) downloadContent(messageID, filename string) string {
|
|||
},
|
||||
})
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *LINEChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1300,3 +1300,8 @@ func stripUserMentionWithRegexp(text string, userID id.UserID, mentionR *regexp.
|
|||
cleaned = strings.TrimLeft(cleaned, ",:; ")
|
||||
return strings.TrimSpace(cleaned)
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *MatrixChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1104,3 +1104,8 @@ func truncate(s string, n int) string {
|
|||
}
|
||||
return string(runes[:n]) + "..."
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *OneBotChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1002,3 +1002,8 @@ func sanitizeURLs(text string) string {
|
|||
return scheme + domain + path
|
||||
})
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *QQChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1133,3 +1133,8 @@ func cryptoRandInt() int {
|
|||
_, _ = rand.Read(b[:])
|
||||
return int(binary.BigEndian.Uint32(b[:])) | 1 // ensure non-zero
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *TelegramChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
58
pkg/channels/voice_capabilities.go
Normal file
58
pkg/channels/voice_capabilities.go
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
package channels
|
||||
|
||||
// VoiceCapabilities describes whether ASR (speech-to-text) and TTS (text-to-speech)
|
||||
// are available for a channel under the current configuration.
|
||||
type VoiceCapabilities struct {
|
||||
ASR bool
|
||||
TTS bool
|
||||
}
|
||||
|
||||
// VoiceCapabilityProvider is an optional interface for channels that want to
|
||||
// explicitly declare their ASR/TTS support.
|
||||
type VoiceCapabilityProvider interface {
|
||||
VoiceCapabilities() VoiceCapabilities
|
||||
}
|
||||
|
||||
// Deprecated: Channels should implement VoiceCapabilityProvider instead.
|
||||
// To be removed once all existing capable channels conform to the interface.
|
||||
var asrCapableChannels = map[string]bool{
|
||||
"discord": true,
|
||||
"telegram": true,
|
||||
"matrix": true,
|
||||
"qq": true,
|
||||
"weixin": true,
|
||||
"line": true,
|
||||
"feishu": true,
|
||||
"onebot": true,
|
||||
}
|
||||
|
||||
// DetectVoiceCapabilities returns ASR/TTS availability for a channel, gated by
|
||||
// whether providers are configured.
|
||||
func DetectVoiceCapabilities(channelName string, ch Channel, asrAvailable bool, ttsAvailable bool) VoiceCapabilities {
|
||||
if ch == nil {
|
||||
return VoiceCapabilities{}
|
||||
}
|
||||
|
||||
if vcp, ok := ch.(VoiceCapabilityProvider); ok {
|
||||
caps := vcp.VoiceCapabilities()
|
||||
if !asrAvailable {
|
||||
caps.ASR = false
|
||||
}
|
||||
if !ttsAvailable {
|
||||
caps.TTS = false
|
||||
}
|
||||
return caps
|
||||
}
|
||||
|
||||
caps := VoiceCapabilities{}
|
||||
if asrAvailable {
|
||||
caps.ASR = asrCapableChannels[channelName]
|
||||
}
|
||||
if ttsAvailable {
|
||||
if _, ok := ch.(MediaSender); ok {
|
||||
caps.TTS = true
|
||||
}
|
||||
}
|
||||
|
||||
return caps
|
||||
}
|
||||
|
|
@ -402,3 +402,8 @@ func (c *WeixinChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]st
|
|||
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// VoiceCapabilities returns the voice capabilities of the channel.
|
||||
func (c *WeixinChannel) VoiceCapabilities() channels.VoiceCapabilities {
|
||||
return channels.VoiceCapabilities{ASR: true, TTS: true}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -558,9 +558,9 @@ type DevicesConfig struct {
|
|||
}
|
||||
|
||||
type VoiceConfig struct {
|
||||
ModelName string `json:"model_name,omitempty" env:"PICOCLAW_VOICE_MODEL_NAME"`
|
||||
EchoTranscription bool `json:"echo_transcription" env:"PICOCLAW_VOICE_ECHO_TRANSCRIPTION"`
|
||||
ElevenLabsAPIKey string `json:"elevenlabs_api_key,omitempty" env:"PICOCLAW_VOICE_ELEVENLABS_API_KEY"`
|
||||
ModelName string `json:"model_name,omitempty" env:"PICOCLAW_VOICE_MODEL_NAME"`
|
||||
TTSModelName string `json:"tts_model_name,omitempty" env:"PICOCLAW_VOICE_TTS_MODEL_NAME"`
|
||||
EchoTranscription bool `json:"echo_transcription" env:"PICOCLAW_VOICE_ECHO_TRANSCRIPTION"`
|
||||
}
|
||||
|
||||
// ModelConfig represents a model-centric provider configuration.
|
||||
|
|
@ -636,13 +636,6 @@ func (c *ModelConfig) SetAPIKey(value string) {
|
|||
}
|
||||
}
|
||||
|
||||
type GatewayConfig struct {
|
||||
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
|
||||
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
|
||||
}
|
||||
|
||||
type ToolDiscoveryConfig struct {
|
||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_DISCOVERY_ENABLED"`
|
||||
TTL int `json:"ttl" env:"PICOCLAW_TOOLS_DISCOVERY_TTL"`
|
||||
|
|
@ -836,6 +829,7 @@ type ToolsConfig struct {
|
|||
Message ToolConfig `json:"message" yaml:"-" envPrefix:"PICOCLAW_TOOLS_MESSAGE_"`
|
||||
ReadFile ReadFileToolConfig `json:"read_file" yaml:"-" envPrefix:"PICOCLAW_TOOLS_READ_FILE_"`
|
||||
SendFile ToolConfig `json:"send_file" yaml:"-" envPrefix:"PICOCLAW_TOOLS_SEND_FILE_"`
|
||||
SendTTS ToolConfig `json:"send_tts" yaml:"-" envPrefix:"PICOCLAW_TOOLS_SEND_TTS_"`
|
||||
Spawn ToolConfig `json:"spawn" yaml:"-" envPrefix:"PICOCLAW_TOOLS_SPAWN_"`
|
||||
SpawnStatus ToolConfig `json:"spawn_status" yaml:"-" envPrefix:"PICOCLAW_TOOLS_SPAWN_STATUS_"`
|
||||
SPI ToolConfig `json:"spi" yaml:"-" envPrefix:"PICOCLAW_TOOLS_SPI_"`
|
||||
|
|
@ -1288,6 +1282,8 @@ func (t *ToolsConfig) IsToolEnabled(name string) bool {
|
|||
return t.WebFetch.Enabled
|
||||
case "send_file":
|
||||
return t.SendFile.Enabled
|
||||
case "send_tts":
|
||||
return t.SendTTS.Enabled
|
||||
case "write_file":
|
||||
return t.WriteFile.Enabled
|
||||
case "mcp":
|
||||
|
|
|
|||
|
|
@ -1418,6 +1418,38 @@ func TestConfigLogLevelEmpty(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestResolveGatewayLogLevel(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
cfgPath := filepath.Join(dir, "config.json")
|
||||
data := `{"version":1,"gateway":{"log_level":"debug"}}`
|
||||
if err := os.WriteFile(cfgPath, []byte(data), 0o600); err != nil {
|
||||
t.Fatalf("setup: %v", err)
|
||||
}
|
||||
|
||||
if got := ResolveGatewayLogLevel(cfgPath); got != "debug" {
|
||||
t.Fatalf("ResolveGatewayLogLevel() = %q, want %q", got, "debug")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveGatewayLogLevel_UsesEnvOverrideAndNormalizesInvalid(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
cfgPath := filepath.Join(dir, "config.json")
|
||||
data := `{"version":1,"gateway":{"log_level":"debug"}}`
|
||||
if err := os.WriteFile(cfgPath, []byte(data), 0o600); err != nil {
|
||||
t.Fatalf("setup: %v", err)
|
||||
}
|
||||
|
||||
t.Setenv("PICOCLAW_LOG_LEVEL", "warning")
|
||||
if got := ResolveGatewayLogLevel(cfgPath); got != "warn" {
|
||||
t.Fatalf("ResolveGatewayLogLevel() with env override = %q, want %q", got, "warn")
|
||||
}
|
||||
|
||||
t.Setenv("PICOCLAW_LOG_LEVEL", "garbage")
|
||||
if got := ResolveGatewayLogLevel(cfgPath); got != DefaultGatewayLogLevel {
|
||||
t.Fatalf("ResolveGatewayLogLevel() with invalid env override = %q, want %q", got, DefaultGatewayLogLevel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelConfig_ExtraBodyRoundTrip(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
cfgPath := filepath.Join(dir, "config.json")
|
||||
|
|
|
|||
|
|
@ -347,7 +347,7 @@ func DefaultConfig() *Config {
|
|||
Host: "127.0.0.1",
|
||||
Port: 18790,
|
||||
HotReload: false,
|
||||
LogLevel: "warn",
|
||||
LogLevel: DefaultGatewayLogLevel,
|
||||
},
|
||||
Tools: ToolsConfig{
|
||||
FilterSensitiveData: true,
|
||||
|
|
@ -434,6 +434,9 @@ func DefaultConfig() *Config {
|
|||
SendFile: ToolConfig{
|
||||
Enabled: true,
|
||||
},
|
||||
SendTTS: ToolConfig{
|
||||
Enabled: false,
|
||||
},
|
||||
MCP: MCPConfig{
|
||||
ToolConfig: ToolConfig{
|
||||
Enabled: false,
|
||||
|
|
|
|||
72
pkg/config/gateway.go
Normal file
72
pkg/config/gateway.go
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
const DefaultGatewayLogLevel = "warn"
|
||||
|
||||
type GatewayConfig struct {
|
||||
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
|
||||
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
|
||||
}
|
||||
|
||||
func canonicalGatewayLogLevel(level logger.LogLevel) string {
|
||||
switch level {
|
||||
case logger.DEBUG:
|
||||
return "debug"
|
||||
case logger.INFO:
|
||||
return "info"
|
||||
case logger.WARN:
|
||||
return "warn"
|
||||
case logger.ERROR:
|
||||
return "error"
|
||||
case logger.FATAL:
|
||||
return "fatal"
|
||||
default:
|
||||
return DefaultGatewayLogLevel
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeGatewayLogLevel(logLevel string) string {
|
||||
if level, ok := logger.ParseLevel(logLevel); ok {
|
||||
return canonicalGatewayLogLevel(level)
|
||||
}
|
||||
return DefaultGatewayLogLevel
|
||||
}
|
||||
|
||||
// EffectiveGatewayLogLevel returns the normalized runtime log level from a loaded config.
|
||||
// Invalid or empty values fall back to the package default.
|
||||
func EffectiveGatewayLogLevel(cfg *Config) string {
|
||||
if cfg == nil {
|
||||
return DefaultGatewayLogLevel
|
||||
}
|
||||
return normalizeGatewayLogLevel(cfg.Gateway.LogLevel)
|
||||
}
|
||||
|
||||
// ResolveGatewayLogLevel reads the configured gateway log level without triggering
|
||||
// the full config loader, so startup code can apply logging before config load logs run.
|
||||
// The PICOCLAW_LOG_LEVEL environment variable overrides the file value.
|
||||
func ResolveGatewayLogLevel(path string) string {
|
||||
cfg := struct {
|
||||
Gateway GatewayConfig `json:"gateway"`
|
||||
}{
|
||||
Gateway: GatewayConfig{LogLevel: DefaultGatewayLogLevel},
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
_ = json.Unmarshal(data, &cfg)
|
||||
}
|
||||
|
||||
if envLevel := os.Getenv("PICOCLAW_LOG_LEVEL"); envLevel != "" {
|
||||
cfg.Gateway.LogLevel = envLevel
|
||||
}
|
||||
|
||||
return normalizeGatewayLogLevel(cfg.Gateway.LogLevel)
|
||||
}
|
||||
|
|
@ -6,6 +6,7 @@ import (
|
|||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
|
@ -13,6 +14,8 @@ import (
|
|||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/agent"
|
||||
"github.com/sipeed/picoclaw/pkg/audio/asr"
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
_ "github.com/sipeed/picoclaw/pkg/channels/dingtalk"
|
||||
|
|
@ -41,7 +44,6 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
"github.com/sipeed/picoclaw/pkg/state"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/voice"
|
||||
)
|
||||
|
||||
const (
|
||||
|
|
@ -61,6 +63,7 @@ type services struct {
|
|||
ChannelManager *channels.Manager
|
||||
DeviceService *devices.Service
|
||||
HealthServer *health.Server
|
||||
VoiceAgentCancel context.CancelFunc
|
||||
manualReloadChan chan struct{}
|
||||
reloading atomic.Bool
|
||||
authToken string
|
||||
|
|
@ -70,6 +73,27 @@ type startupBlockedProvider struct {
|
|||
reason string
|
||||
}
|
||||
|
||||
func logChannelVoiceCapabilities(cm *channels.Manager, asrAvailable bool, ttsAvailable bool) {
|
||||
if cm == nil {
|
||||
return
|
||||
}
|
||||
|
||||
names := cm.GetEnabledChannels()
|
||||
sort.Strings(names)
|
||||
for _, name := range names {
|
||||
ch, ok := cm.GetChannel(name)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
caps := channels.DetectVoiceCapabilities(name, ch, asrAvailable, ttsAvailable)
|
||||
logger.InfoCF("voice", "Channel voice capabilities", map[string]any{
|
||||
"channel": name,
|
||||
"asr": caps.ASR,
|
||||
"tts": caps.TTS,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (p *startupBlockedProvider) Chat(
|
||||
_ context.Context,
|
||||
_ []providers.Message,
|
||||
|
|
@ -98,6 +122,12 @@ func Run(debug bool, homePath, configPath string, allowEmptyStartup bool) error
|
|||
}
|
||||
defer logger.DisableFileLogging()
|
||||
|
||||
if debug {
|
||||
logger.SetLevel(logger.DEBUG)
|
||||
} else {
|
||||
logger.SetLevelFromString(config.ResolveGatewayLogLevel(configPath))
|
||||
}
|
||||
|
||||
cfg, err := config.LoadConfig(configPath)
|
||||
if err != nil {
|
||||
logger.Fatalf("error loading config: %v", err)
|
||||
|
|
@ -109,16 +139,17 @@ func Run(debug bool, homePath, configPath string, allowEmptyStartup bool) error
|
|||
|
||||
// Debug mode permanently overrides the config log level to DEBUG.
|
||||
if debug {
|
||||
logger.SetLevel(logger.DEBUG)
|
||||
fmt.Println("🔍 Debug mode enabled")
|
||||
} else {
|
||||
logger.SetLevelFromString(cfg.Gateway.LogLevel)
|
||||
logger.Infof("Log level set to %q", cfg.Gateway.LogLevel)
|
||||
effectiveLogLevel := config.EffectiveGatewayLogLevel(cfg)
|
||||
logger.SetLevelFromString(effectiveLogLevel)
|
||||
logger.Infof("Log level set to %q", effectiveLogLevel)
|
||||
}
|
||||
|
||||
// Enforce singleton: write PID file with generated token.
|
||||
pidData, err := pid.WritePidFile(homePath, cfg.Gateway.Host, cfg.Gateway.Port)
|
||||
if err != nil {
|
||||
logger.Warnf("write pid file failed: %v", err)
|
||||
return fmt.Errorf("singleton check failed: %w", err)
|
||||
}
|
||||
defer pid.RemovePidFile(homePath)
|
||||
|
|
@ -331,11 +362,14 @@ func setupAndStartServices(
|
|||
agentLoop.SetChannelManager(runningServices.ChannelManager)
|
||||
agentLoop.SetMediaStore(runningServices.MediaStore)
|
||||
|
||||
if transcriber := voice.DetectTranscriber(cfg); transcriber != nil {
|
||||
transcriber := asr.DetectTranscriber(cfg)
|
||||
if transcriber != nil {
|
||||
agentLoop.SetTranscriber(transcriber)
|
||||
logger.InfoCF("voice", "Transcription enabled (agent-level)", map[string]any{"provider": transcriber.Name()})
|
||||
}
|
||||
|
||||
ttsAvailable := tts.DetectTTS(cfg) != nil
|
||||
|
||||
enabledChannels := runningServices.ChannelManager.GetEnabledChannels()
|
||||
if len(enabledChannels) > 0 {
|
||||
fmt.Printf("✓ Channels enabled: %s\n", enabledChannels)
|
||||
|
|
@ -352,6 +386,16 @@ func setupAndStartServices(
|
|||
return nil, fmt.Errorf("error starting channels: %w", err)
|
||||
}
|
||||
|
||||
logChannelVoiceCapabilities(runningServices.ChannelManager, transcriber != nil, ttsAvailable)
|
||||
|
||||
if transcriber != nil {
|
||||
// Start Voice Agent Orchestrator after channels are ready.
|
||||
vaCtx, vaCancel := context.WithCancel(context.Background())
|
||||
runningServices.VoiceAgentCancel = vaCancel
|
||||
voiceAgent := asr.NewAgent(msgBus, transcriber)
|
||||
voiceAgent.Start(vaCtx)
|
||||
}
|
||||
|
||||
fmt.Printf(
|
||||
"✓ Health endpoints available at http://%s:%d/health, /ready and /reload (POST)\n",
|
||||
cfg.Gateway.Host,
|
||||
|
|
@ -381,6 +425,9 @@ func stopAndCleanupServices(runningServices *services, shutdownTimeout time.Dura
|
|||
if !isReload && runningServices.ChannelManager != nil {
|
||||
runningServices.ChannelManager.StopAll(shutdownCtx)
|
||||
}
|
||||
if runningServices.VoiceAgentCancel != nil {
|
||||
runningServices.VoiceAgentCancel()
|
||||
}
|
||||
if runningServices.DeviceService != nil {
|
||||
runningServices.DeviceService.Stop()
|
||||
}
|
||||
|
|
@ -476,8 +523,9 @@ func handleConfigReload(
|
|||
// Debug mode permanently overrides the config log level to DEBUG.
|
||||
if !debug {
|
||||
// Update log level last so that reload-related info/warn logs above are not suppressed.
|
||||
logger.SetLevelFromString(newCfg.Gateway.LogLevel)
|
||||
logger.Infof("Log level changing from current to %q", newCfg.Gateway.LogLevel)
|
||||
effectiveLogLevel := config.EffectiveGatewayLogLevel(newCfg)
|
||||
logger.SetLevelFromString(effectiveLogLevel)
|
||||
logger.Infof("Log level changing from current to %q", effectiveLogLevel)
|
||||
}
|
||||
|
||||
return nil
|
||||
|
|
@ -556,14 +604,22 @@ func restartServices(
|
|||
fmt.Println(" ✓ Device event service restarted")
|
||||
}
|
||||
|
||||
transcriber := voice.DetectTranscriber(cfg)
|
||||
transcriber := asr.DetectTranscriber(cfg)
|
||||
al.SetTranscriber(transcriber)
|
||||
if transcriber != nil {
|
||||
logger.InfoCF("voice", "Transcription re-enabled (agent-level)", map[string]any{"provider": transcriber.Name()})
|
||||
|
||||
// Start Voice Agent Orchestrator on reload
|
||||
vaCtx, vaCancel := context.WithCancel(context.Background())
|
||||
runningServices.VoiceAgentCancel = vaCancel
|
||||
voiceAgent := asr.NewAgent(msgBus, transcriber)
|
||||
voiceAgent.Start(vaCtx)
|
||||
} else {
|
||||
logger.InfoCF("voice", "Transcription disabled", nil)
|
||||
}
|
||||
|
||||
ttsAvailable := tts.DetectTTS(cfg) != nil
|
||||
logChannelVoiceCapabilities(runningServices.ChannelManager, transcriber != nil, ttsAvailable)
|
||||
// NOTE: PID file is written once at startup and not updated on reload.
|
||||
// Changing the gateway listen address requires a full restart.
|
||||
|
||||
|
|
|
|||
|
|
@ -94,6 +94,7 @@ func WritePidFile(homePath, host string, port int) (*PidFileData, error) {
|
|||
os.Remove(tmp)
|
||||
return nil, fmt.Errorf("failed to rename pid file: %w", err)
|
||||
}
|
||||
logger.Debugf("wrote pid file: %s success", pidPath)
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
|
@ -108,10 +109,12 @@ func ReadPidFileWithCheck(homePath string) *PidFileData {
|
|||
pidPath := pidFilePath(homePath)
|
||||
data, err := readPidFileUnlocked(pidPath)
|
||||
if err != nil {
|
||||
logger.Debugf("failed to read pid file: %s", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
if !isProcessRunning(data.PID) {
|
||||
logger.Debugf("process not running, remove pid file: %s", pidPath)
|
||||
os.Remove(pidPath)
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -98,6 +98,19 @@ func ExtractProtocol(model string) (protocol, modelID string) {
|
|||
return protocol, modelID
|
||||
}
|
||||
|
||||
// ResolveAPIBase returns the configured API base, or the protocol default when
|
||||
// the model uses an HTTP-based provider family with a known default endpoint.
|
||||
func ResolveAPIBase(cfg *config.ModelConfig) string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
}
|
||||
if apiBase := strings.TrimSpace(cfg.APIBase); apiBase != "" {
|
||||
return strings.TrimRight(apiBase, "/")
|
||||
}
|
||||
protocol, _ := ExtractProtocol(cfg.Model)
|
||||
return strings.TrimRight(getDefaultAPIBase(protocol), "/")
|
||||
}
|
||||
|
||||
// CreateProviderFromConfig creates a provider based on the ModelConfig.
|
||||
// It uses the protocol prefix in the Model field to determine which provider to create.
|
||||
// Supported protocol families include OpenAI-compatible prefixes (e.g., openai, openrouter, groq, gemini),
|
||||
|
|
|
|||
82
pkg/tools/tts_send.go
Normal file
82
pkg/tools/tts_send.go
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
)
|
||||
|
||||
type SendTTSTool struct {
|
||||
provider tts.TTSProvider
|
||||
mediaStore media.MediaStore
|
||||
}
|
||||
|
||||
func NewSendTTSTool(provider tts.TTSProvider, store media.MediaStore) *SendTTSTool {
|
||||
return &SendTTSTool{
|
||||
provider: provider,
|
||||
mediaStore: store,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *SendTTSTool) Name() string { return "send_tts" }
|
||||
|
||||
func (t *SendTTSTool) Description() string {
|
||||
return "Synthesize speech from text and send it as an audio file to the user."
|
||||
}
|
||||
|
||||
func (t *SendTTSTool) Parameters() map[string]any {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{
|
||||
"type": "string",
|
||||
"description": "The text to synthesize into speech. NOTE: Reply in a highly concise, conversational, oral style suitable for text-to-speech. Do not use markdown, emojis, asterisks, or code blocks. Speak naturally.",
|
||||
},
|
||||
"filename": map[string]any{
|
||||
"type": "string",
|
||||
"description": "Optional filename for the audio file (e.g., response.ogg).",
|
||||
},
|
||||
},
|
||||
"required": []string{"text"},
|
||||
}
|
||||
}
|
||||
|
||||
func (t *SendTTSTool) SetMediaStore(store media.MediaStore) {
|
||||
t.mediaStore = store
|
||||
}
|
||||
|
||||
func (t *SendTTSTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||
text, _ := args["text"].(string)
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return ErrorResult("text is required")
|
||||
}
|
||||
|
||||
channel := ToolChannel(ctx)
|
||||
chatID := ToolChatID(ctx)
|
||||
filename, _ := args["filename"].(string)
|
||||
|
||||
ref, err := tts.SynthesizeAndStore(
|
||||
ctx,
|
||||
t.provider,
|
||||
t.mediaStore,
|
||||
text,
|
||||
filename,
|
||||
channel,
|
||||
chatID,
|
||||
)
|
||||
if err != nil {
|
||||
return ErrorResult(err.Error()).WithError(err)
|
||||
}
|
||||
|
||||
// Return with ForUser set to original text, Media containing the audio ref,
|
||||
// and mark as ResponseHandled so the audio is sent immediately without LLM intervention.
|
||||
return &ToolResult{
|
||||
ForLLM: "TTS audio sent",
|
||||
ForUser: text,
|
||||
Media: []string{ref},
|
||||
ResponseHandled: true,
|
||||
}
|
||||
}
|
||||
|
|
@ -1,151 +0,0 @@
|
|||
package voice
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
type GroqTranscriber struct {
|
||||
apiKey string
|
||||
apiBase string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewGroqTranscriber(apiKey string) *GroqTranscriber {
|
||||
logger.DebugCF("voice", "Creating Groq transcriber", map[string]any{"has_api_key": apiKey != ""})
|
||||
|
||||
apiBase := "https://api.groq.com/openai/v1"
|
||||
return &GroqTranscriber{
|
||||
apiKey: apiKey,
|
||||
apiBase: apiBase,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 60 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
||||
logger.InfoCF("voice", "Starting transcription", map[string]any{"audio_file": audioFilePath})
|
||||
|
||||
audioFile, err := os.Open(audioFilePath)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to open audio file", map[string]any{"path": audioFilePath, "error": err})
|
||||
return nil, fmt.Errorf("failed to open audio file: %w", err)
|
||||
}
|
||||
defer audioFile.Close()
|
||||
|
||||
fileInfo, err := audioFile.Stat()
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to get file info", map[string]any{"path": audioFilePath, "error": err})
|
||||
return nil, fmt.Errorf("failed to get file info: %w", err)
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "Audio file details", map[string]any{
|
||||
"size_bytes": fileInfo.Size(),
|
||||
"file_name": filepath.Base(audioFilePath),
|
||||
})
|
||||
|
||||
var requestBody bytes.Buffer
|
||||
writer := multipart.NewWriter(&requestBody)
|
||||
|
||||
part, err := writer.CreateFormFile("file", filepath.Base(audioFilePath))
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to create form file", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to create form file: %w", err)
|
||||
}
|
||||
|
||||
copied, err := io.Copy(part, audioFile)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to copy file content", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to copy file content: %w", err)
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "File copied to request", map[string]any{"bytes_copied": copied})
|
||||
|
||||
if err = writer.WriteField("model", "whisper-large-v3"); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to write model field", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to write model field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("response_format", "json"); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to write response_format field", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to write response_format field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.Close(); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to close multipart writer", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
url := t.apiBase + "/audio/transcriptions"
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", url, &requestBody)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to create request", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
||||
|
||||
logger.DebugCF("voice", "Sending transcription request to Groq API", map[string]any{
|
||||
"url": url,
|
||||
"request_size_bytes": requestBody.Len(),
|
||||
"file_size_bytes": fileInfo.Size(),
|
||||
})
|
||||
|
||||
resp, err := t.httpClient.Do(req)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to send request", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to read response", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
logger.ErrorCF("voice", "API error", map[string]any{
|
||||
"status_code": resp.StatusCode,
|
||||
"response": string(body),
|
||||
})
|
||||
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "Received response from Groq API", map[string]any{
|
||||
"status_code": resp.StatusCode,
|
||||
"response_size_bytes": len(body),
|
||||
})
|
||||
|
||||
var result TranscriptionResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to unmarshal response", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
logger.InfoCF("voice", "Transcription completed successfully", map[string]any{
|
||||
"text_length": len(result.Text),
|
||||
"language": result.Language,
|
||||
"duration_seconds": result.Duration,
|
||||
"transcription_preview": utils.Truncate(result.Text, 50),
|
||||
})
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func (t *GroqTranscriber) Name() string {
|
||||
return "groq"
|
||||
}
|
||||
|
|
@ -1,84 +0,0 @@
|
|||
package voice
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
var _ Transcriber = (*GroqTranscriber)(nil)
|
||||
|
||||
func TestGroqTranscriberName(t *testing.T) {
|
||||
tr := NewGroqTranscriber("sk-test")
|
||||
if got := tr.Name(); got != "groq" {
|
||||
t.Errorf("Name() = %q, want %q", got, "groq")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGroqTranscribe(t *testing.T) {
|
||||
// Write a minimal fake audio file so the transcriber can open and send it.
|
||||
tmpDir := t.TempDir()
|
||||
audioPath := filepath.Join(tmpDir, "clip.ogg")
|
||||
if err := os.WriteFile(audioPath, []byte("fake-audio-data"), 0o644); err != nil {
|
||||
t.Fatalf("failed to write fake audio file: %v", err)
|
||||
}
|
||||
|
||||
t.Run("success", func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/audio/transcriptions" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
if r.Header.Get("Authorization") != "Bearer sk-test" {
|
||||
t.Errorf("unexpected Authorization header: %s", r.Header.Get("Authorization"))
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(TranscriptionResponse{
|
||||
Text: "hello world",
|
||||
Language: "en",
|
||||
Duration: 1.5,
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tr := NewGroqTranscriber("sk-test")
|
||||
tr.apiBase = srv.URL
|
||||
|
||||
resp, err := tr.Transcribe(context.Background(), audioPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Transcribe() error: %v", err)
|
||||
}
|
||||
if resp.Text != "hello world" {
|
||||
t.Errorf("Text = %q, want %q", resp.Text, "hello world")
|
||||
}
|
||||
if resp.Language != "en" {
|
||||
t.Errorf("Language = %q, want %q", resp.Language, "en")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("api error", func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, `{"error":"invalid_api_key"}`, http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tr := NewGroqTranscriber("sk-bad")
|
||||
tr.apiBase = srv.URL
|
||||
|
||||
_, err := tr.Transcribe(context.Background(), audioPath)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for non-200 response, got nil")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing file", func(t *testing.T) {
|
||||
tr := NewGroqTranscriber("sk-test")
|
||||
_, err := tr.Transcribe(context.Background(), filepath.Join(tmpDir, "nonexistent.ogg"))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing file, got nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
|
@ -1,68 +0,0 @@
|
|||
package voice
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
type Transcriber interface {
|
||||
Name() string
|
||||
Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error)
|
||||
}
|
||||
|
||||
type TranscriptionResponse struct {
|
||||
Text string `json:"text"`
|
||||
Language string `json:"language,omitempty"`
|
||||
Duration float64 `json:"duration,omitempty"`
|
||||
}
|
||||
|
||||
func supportsAudioTranscription(model string) bool {
|
||||
protocol, _ := providers.ExtractProtocol(model)
|
||||
|
||||
switch protocol {
|
||||
case "openai", "azure", "azure-openai",
|
||||
"litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"qwen-us", "dashscope-us", "mistral", "avian", "minimax", "longcat", "modelscope", "novita",
|
||||
"coding-plan", "alibaba-coding", "qwen-coding":
|
||||
// These protocols all go through the OpenAI-compatible or Azure provider path in
|
||||
// providers.CreateProviderFromConfig, so they are the only ones that can supply
|
||||
// the audio media payload shape expected by NewAudioModelTranscriber.
|
||||
|
||||
// TODO: Further restrict this by modelID, since not every model under these
|
||||
// protocols supports audio transcription.
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// DetectTranscriber inspects cfg and returns the appropriate Transcriber, or
|
||||
// nil if no supported transcription provider is configured.
|
||||
func DetectTranscriber(cfg *config.Config) Transcriber {
|
||||
if modelName := strings.TrimSpace(cfg.Voice.ModelName); modelName != "" {
|
||||
modelCfg, err := cfg.GetModelConfig(modelName)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if supportsAudioTranscription(modelCfg.Model) {
|
||||
return NewAudioModelTranscriber(modelCfg)
|
||||
}
|
||||
}
|
||||
|
||||
// ElevenLabs voice config (supports Scribe STT).
|
||||
if key := strings.TrimSpace(cfg.Voice.ElevenLabsAPIKey); key != "" {
|
||||
return NewElevenLabsTranscriber(key)
|
||||
}
|
||||
// Fall back to any model-list entry that uses the groq/ protocol.
|
||||
for _, mc := range cfg.ModelList {
|
||||
if strings.HasPrefix(mc.Model, "groq/") && mc.APIKey() != "" {
|
||||
return NewGroqTranscriber(mc.APIKey())
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
@ -20,6 +20,14 @@ func (h *Handler) registerConfigRoutes(mux *http.ServeMux) {
|
|||
mux.HandleFunc("POST /api/config/test-command-patterns", h.handleTestCommandPatterns)
|
||||
}
|
||||
|
||||
func (h *Handler) applyRuntimeLogLevel() {
|
||||
if h.debug {
|
||||
logger.SetLevel(logger.DEBUG)
|
||||
return
|
||||
}
|
||||
logger.SetLevelFromString(config.ResolveGatewayLogLevel(h.configPath))
|
||||
}
|
||||
|
||||
// handleGetConfig returns the complete system configuration.
|
||||
//
|
||||
// GET /api/config
|
||||
|
|
@ -80,8 +88,6 @@ func (h *Handler) handleUpdateConfig(w http.ResponseWriter, r *http.Request) {
|
|||
return
|
||||
}
|
||||
|
||||
logger.Infof("configuration updated successfully")
|
||||
|
||||
if err := config.SaveConfig(h.configPath, &cfg); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
|
|
@ -89,6 +95,8 @@ func (h *Handler) handleUpdateConfig(w http.ResponseWriter, r *http.Request) {
|
|||
|
||||
// Refresh cached pico token in case user changed it.
|
||||
refreshPicoToken(&cfg)
|
||||
h.applyRuntimeLogLevel()
|
||||
logger.Infof("configuration updated successfully")
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||
|
|
@ -133,7 +141,6 @@ func (h *Handler) handlePatchConfig(w http.ResponseWriter, r *http.Request) {
|
|||
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
existing, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
http.Error(w, "Failed to serialize current config", http.StatusInternalServerError)
|
||||
|
|
@ -187,6 +194,8 @@ func (h *Handler) handlePatchConfig(w http.ResponseWriter, r *http.Request) {
|
|||
|
||||
// Refresh cached pico token in case user changed it.
|
||||
refreshPicoToken(&newCfg)
|
||||
h.applyRuntimeLogLevel()
|
||||
logger.Infof("configuration updated successfully")
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||
|
|
|
|||
|
|
@ -9,8 +9,38 @@ import (
|
|||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
func assertGatewayLogLevelApplied(t *testing.T, method, body string, want logger.LogLevel) {
|
||||
t.Helper()
|
||||
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
initialLevel := logger.GetLevel()
|
||||
logger.SetLevel(logger.INFO)
|
||||
t.Cleanup(func() {
|
||||
logger.SetLevel(initialLevel)
|
||||
})
|
||||
|
||||
h := NewHandler(configPath)
|
||||
mux := http.NewServeMux()
|
||||
h.RegisterRoutes(mux)
|
||||
|
||||
req := httptest.NewRequest(method, "/api/config", bytes.NewBufferString(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
mux.ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("%s /api/config status = %d, want %d, body=%s", method, rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
if got := logger.GetLevel(); got != want {
|
||||
t.Fatalf("logger.GetLevel() = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleUpdateConfig_PreservesExecAllowRemoteDefaultWhenOmitted(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
|
|
@ -251,6 +281,68 @@ func TestHandlePatchConfig_SucceedsWhenPicoTokenInSecurityOnly(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestHandleUpdateConfig_AppliesGatewayLogLevel(t *testing.T) {
|
||||
assertGatewayLogLevelApplied(t, http.MethodPut, `{
|
||||
"version": 1,
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"workspace": "~/.picoclaw/workspace",
|
||||
"model_name": "custom-default"
|
||||
}
|
||||
},
|
||||
"gateway": {
|
||||
"log_level": "error"
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "custom-default",
|
||||
"model": "openai/gpt-4o",
|
||||
"api_keys": ["sk-default"]
|
||||
}
|
||||
]
|
||||
}`, logger.ERROR)
|
||||
}
|
||||
|
||||
func TestHandlePatchConfig_AppliesGatewayLogLevel(t *testing.T) {
|
||||
assertGatewayLogLevelApplied(t, http.MethodPatch, `{
|
||||
"gateway": {
|
||||
"log_level": "debug"
|
||||
}
|
||||
}`, logger.DEBUG)
|
||||
}
|
||||
|
||||
func TestHandlePatchConfig_PreservesDebugFlagOverride(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
initialLevel := logger.GetLevel()
|
||||
logger.SetLevel(logger.INFO)
|
||||
t.Cleanup(func() {
|
||||
logger.SetLevel(initialLevel)
|
||||
})
|
||||
|
||||
h := NewHandler(configPath)
|
||||
h.SetDebug(true)
|
||||
mux := http.NewServeMux()
|
||||
h.RegisterRoutes(mux)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPatch, "/api/config", bytes.NewBufferString(`{
|
||||
"gateway": {
|
||||
"log_level": "error"
|
||||
}
|
||||
}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
mux.ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("PATCH /api/config status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
if got := logger.GetLevel(); got != logger.DEBUG {
|
||||
t.Fatalf("logger.GetLevel() = %v, want %v", got, logger.DEBUG)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatchConfig_SavesDiscordTokenFromPayload(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
|
|
|
|||
|
|
@ -69,6 +69,14 @@ func ensurePicoTokenCachedLocked(configPath string) {
|
|||
refreshPicoTokensLocked(configPath)
|
||||
}
|
||||
|
||||
func (h *Handler) gatewayCommandArgs() []string {
|
||||
args := []string{"gateway", "-E"}
|
||||
if h.debug {
|
||||
args = append(args, "-d")
|
||||
}
|
||||
return args
|
||||
}
|
||||
|
||||
const (
|
||||
protocolKey = "Sec-Websocket-Protocol"
|
||||
tokenPrefix = "token."
|
||||
|
|
@ -531,7 +539,7 @@ func (h *Handler) startGatewayLocked(initialStatus string, existingPid int) (int
|
|||
execPath := utils.FindPicoclawBinary()
|
||||
logger.InfoC("gateway", fmt.Sprintf("Starting gateway process (%s)", execPath))
|
||||
|
||||
cmd = exec.Command(execPath, "gateway", "-E")
|
||||
cmd = exec.Command(execPath, h.gatewayCommandArgs()...)
|
||||
cmd.Env = os.Environ()
|
||||
// Forward the launcher's config path via the environment variable that
|
||||
// GetConfigPath() already reads, so the gateway sub-process uses the same
|
||||
|
|
@ -620,6 +628,7 @@ func (h *Handler) startGatewayLocked(initialStatus string, existingPid int) (int
|
|||
gateway.mu.Lock()
|
||||
if gateway.cmd == cmd {
|
||||
gateway.pidData = pd
|
||||
gateway.picoToken = cfg.Channels.Pico.Token.String()
|
||||
setGatewayRuntimeStatusLocked("running")
|
||||
}
|
||||
gateway.mu.Unlock()
|
||||
|
|
@ -914,34 +923,13 @@ func (h *Handler) gatewayStatusData() map[string]any {
|
|||
data["pid"] = pidData.PID
|
||||
gateway.mu.Unlock()
|
||||
} else {
|
||||
// Fallback: probe health endpoint to get pid and status
|
||||
_, statusCode, err := h.getGatewayHealth(cfg, 2*time.Second)
|
||||
if err != nil {
|
||||
gateway.mu.Lock()
|
||||
data["gateway_status"] = gatewayStatusWithoutHealthLocked()
|
||||
gateway.pidData = nil
|
||||
gateway.mu.Unlock()
|
||||
logger.ErrorC("gateway", fmt.Sprintf("Gateway health check failed: %v", err))
|
||||
} else {
|
||||
logger.InfoC("gateway", fmt.Sprintf("Gateway health status: %d", statusCode))
|
||||
if statusCode != http.StatusOK {
|
||||
gateway.mu.Lock()
|
||||
setGatewayRuntimeStatusLocked("error")
|
||||
gateway.pidData = nil
|
||||
gateway.mu.Unlock()
|
||||
data["gateway_status"] = "error"
|
||||
data["status_code"] = statusCode
|
||||
} else {
|
||||
gateway.mu.Lock()
|
||||
setGatewayRuntimeStatusLocked("running")
|
||||
bootDefaultModel := gateway.bootDefaultModel
|
||||
if bootDefaultModel != "" {
|
||||
data["boot_default_model"] = bootDefaultModel
|
||||
}
|
||||
data["gateway_status"] = "running"
|
||||
gateway.mu.Unlock()
|
||||
}
|
||||
}
|
||||
// Intentionally skip health probe here; the startup goroutine
|
||||
// (startGatewayLocked) already handles liveness detection via
|
||||
// pidFile polling and health fallback.
|
||||
gateway.mu.Lock()
|
||||
data["gateway_status"] = gatewayStatusWithoutHealthLocked()
|
||||
gateway.pidData = nil
|
||||
gateway.mu.Unlock()
|
||||
}
|
||||
|
||||
gatewayStatus, _ := data["gateway_status"].(string)
|
||||
|
|
|
|||
|
|
@ -15,8 +15,11 @@ import (
|
|||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/auth"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
ppid "github.com/sipeed/picoclaw/pkg/pid"
|
||||
"github.com/sipeed/picoclaw/web/backend/utils"
|
||||
)
|
||||
|
||||
|
|
@ -77,6 +80,8 @@ func resetGatewayTestState(t *testing.T) {
|
|||
|
||||
gateway.mu.Lock()
|
||||
gateway.cmd = nil
|
||||
gateway.pidData = nil
|
||||
gateway.owned = false
|
||||
gateway.bootDefaultModel = ""
|
||||
gateway.bootConfigSignature = ""
|
||||
setGatewayRuntimeStatusLocked("stopped")
|
||||
|
|
@ -166,6 +171,17 @@ func TestGatewayStartReady_DefaultModelWithoutCredential(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestGatewayCommandArgsIncludesDebugFlagWhenEnabled(t *testing.T) {
|
||||
h := NewHandler(filepath.Join(t.TempDir(), "config.json"))
|
||||
h.SetDebug(true)
|
||||
|
||||
args := h.gatewayCommandArgs()
|
||||
want := []string{"gateway", "-E", "-d"}
|
||||
if strings.Join(args, " ") != strings.Join(want, " ") {
|
||||
t.Fatalf("gatewayCommandArgs() = %v, want %v", args, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayStartReady_LocalModelWithoutAPIKey(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
|
|
@ -431,7 +447,7 @@ func TestGatewayStatusKeepsRunningWhenHealthProbeFailsAfterRunning(t *testing.T)
|
|||
}
|
||||
}
|
||||
|
||||
func TestGatewayStatusReportsRunningFromHealthProbe(t *testing.T) {
|
||||
func TestGatewayStatusReportsRunningFromPidProbe(t *testing.T) {
|
||||
resetGatewayTestState(t)
|
||||
|
||||
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||
|
|
@ -455,6 +471,9 @@ func TestGatewayStatusReportsRunningFromHealthProbe(t *testing.T) {
|
|||
return mockGatewayHealthResponse(http.StatusOK, cmd.Process.Pid), nil
|
||||
}
|
||||
|
||||
_, err := ppid.WritePidFile(globalConfigDir(), "localhost", 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/gateway/status", nil)
|
||||
mux.ServeHTTP(rec, req)
|
||||
|
|
@ -500,6 +519,8 @@ func TestGatewayStatusRequiresRestartAfterDefaultModelChange(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("FindProcess() error = %v", err)
|
||||
}
|
||||
_, err = ppid.WritePidFile(globalConfigDir(), "localhost", 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
bootSignature := computeConfigSignature(cfg)
|
||||
gateway.mu.Lock()
|
||||
|
|
|
|||
|
|
@ -1,26 +1,87 @@
|
|||
package api
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
const modelProbeTimeout = 800 * time.Millisecond
|
||||
const (
|
||||
modelProbeTimeout = 800 * time.Millisecond
|
||||
modelProbeSuccessBaseInterval = 2 * time.Second
|
||||
modelProbeSuccessMaxInterval = 60 * time.Second
|
||||
modelProbeFailureBaseInterval = 1 * time.Second
|
||||
modelProbeFailureMaxInterval = 30 * time.Second
|
||||
modelProbeBackoffMaxShift = 8
|
||||
modelProbeCacheMaxEntries = 1024
|
||||
modelProbeCacheEntryTTL = 30 * time.Minute
|
||||
modelProbeCacheTrimToEntries = modelProbeCacheMaxEntries * 8 / 10
|
||||
modelProbeTTLGCInterval = 1 * time.Minute
|
||||
)
|
||||
|
||||
const (
|
||||
modelStatusAvailable = "available"
|
||||
modelStatusUnconfigured = "unconfigured"
|
||||
modelStatusUnreachable = "unreachable"
|
||||
)
|
||||
|
||||
type modelConfigurationSummary struct {
|
||||
Available bool
|
||||
Status string
|
||||
}
|
||||
|
||||
var (
|
||||
probeTCPServiceFunc = probeTCPService
|
||||
probeOllamaModelFunc = probeOllamaModel
|
||||
probeOpenAICompatibleModelFunc = probeOpenAICompatibleModel
|
||||
modelProbeNowFunc = time.Now
|
||||
modelProbeState = newModelProbeCacheState()
|
||||
)
|
||||
|
||||
type modelProbeCacheState struct {
|
||||
mu sync.RWMutex
|
||||
cache map[string]*modelProbeCacheEntry
|
||||
group singleflight.Group
|
||||
nextTTLGCAt time.Time
|
||||
}
|
||||
|
||||
type modelProbeCacheEntry struct {
|
||||
lastResult bool
|
||||
hasResult bool
|
||||
successStreak int
|
||||
failureStreak int
|
||||
nextProbeAt time.Time
|
||||
updatedAt time.Time
|
||||
}
|
||||
|
||||
func newModelProbeCacheState() *modelProbeCacheState {
|
||||
return &modelProbeCacheState{cache: map[string]*modelProbeCacheEntry{}}
|
||||
}
|
||||
|
||||
func resetModelProbeCache() {
|
||||
modelProbeState.resetForTest()
|
||||
}
|
||||
|
||||
func (s *modelProbeCacheState) resetForTest() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.cache = map[string]*modelProbeCacheEntry{}
|
||||
s.nextTTLGCAt = time.Time{}
|
||||
}
|
||||
|
||||
func hasModelConfiguration(m *config.ModelConfig) bool {
|
||||
authMethod := strings.ToLower(strings.TrimSpace(m.AuthMethod))
|
||||
apiKey := strings.TrimSpace(m.APIKey())
|
||||
|
|
@ -43,16 +104,17 @@ func hasModelConfiguration(m *config.ModelConfig) bool {
|
|||
return apiKey != ""
|
||||
}
|
||||
|
||||
// isModelConfigured reports whether a model is currently available to use.
|
||||
// Local models must be reachable; remote/API-key models only need saved config.
|
||||
func isModelConfigured(m *config.ModelConfig) bool {
|
||||
func modelConfigurationStatus(m *config.ModelConfig) modelConfigurationSummary {
|
||||
if !hasModelConfiguration(m) {
|
||||
return false
|
||||
return modelConfigurationSummary{Available: false, Status: modelStatusUnconfigured}
|
||||
}
|
||||
if requiresRuntimeProbe(m) {
|
||||
return probeLocalModelAvailability(m)
|
||||
if probeLocalModelAvailability(m) {
|
||||
return modelConfigurationSummary{Available: true, Status: modelStatusAvailable}
|
||||
}
|
||||
return modelConfigurationSummary{Available: false, Status: modelStatusUnreachable}
|
||||
}
|
||||
return true
|
||||
return modelConfigurationSummary{Available: true, Status: modelStatusAvailable}
|
||||
}
|
||||
|
||||
func requiresRuntimeProbe(m *config.ModelConfig) bool {
|
||||
|
|
@ -81,6 +143,34 @@ func requiresRuntimeProbe(m *config.ModelConfig) bool {
|
|||
}
|
||||
|
||||
func probeLocalModelAvailability(m *config.ModelConfig) bool {
|
||||
cacheKey := modelProbeCacheKey(m)
|
||||
return modelProbeState.probe(cacheKey, func() bool {
|
||||
return runLocalModelProbe(m)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *modelProbeCacheState) probe(cacheKey string, probeFunc func() bool) bool {
|
||||
now := modelProbeNowFunc()
|
||||
if cachedResult, ok := s.getCachedResult(cacheKey, now); ok {
|
||||
return cachedResult
|
||||
}
|
||||
|
||||
v, _, _ := s.group.Do(cacheKey, func() (any, error) {
|
||||
now = modelProbeNowFunc()
|
||||
if cachedResult, ok := s.getCachedResult(cacheKey, now); ok {
|
||||
return cachedResult, nil
|
||||
}
|
||||
|
||||
result := probeFunc()
|
||||
s.setCachedResult(cacheKey, result, now)
|
||||
return result, nil
|
||||
})
|
||||
|
||||
result, _ := v.(bool)
|
||||
return result
|
||||
}
|
||||
|
||||
func runLocalModelProbe(m *config.ModelConfig) bool {
|
||||
apiBase := modelProbeAPIBase(m)
|
||||
protocol, modelID := splitModel(m.Model)
|
||||
switch protocol {
|
||||
|
|
@ -100,6 +190,195 @@ func probeLocalModelAvailability(m *config.ModelConfig) bool {
|
|||
}
|
||||
}
|
||||
|
||||
func modelProbeCacheKey(m *config.ModelConfig) string {
|
||||
protocol, modelID := splitModel(m.Model)
|
||||
|
||||
apiBaseRaw := modelProbeAPIBase(m)
|
||||
apiBase := strings.ToLower(strings.TrimRight(strings.TrimSpace(apiBaseRaw), "/"))
|
||||
apiKeyFingerprint := modelProbeAPIKeyFingerprint(m.APIKey())
|
||||
|
||||
var b strings.Builder
|
||||
b.Grow(len(protocol) + len(modelID) + len(apiBase) + len(apiKeyFingerprint) + 8)
|
||||
b.WriteString(protocol)
|
||||
b.WriteByte('|')
|
||||
b.WriteString(modelID)
|
||||
b.WriteByte('|')
|
||||
b.WriteString(apiBase)
|
||||
b.WriteByte('|')
|
||||
b.WriteString(apiKeyFingerprint)
|
||||
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func modelProbeAPIKeyFingerprint(raw string) string {
|
||||
apiKey := strings.TrimSpace(raw)
|
||||
if apiKey == "" {
|
||||
return "none"
|
||||
}
|
||||
|
||||
h := fnv.New64a()
|
||||
_, _ = h.Write([]byte(apiKey))
|
||||
return strconv.FormatUint(h.Sum64(), 36)
|
||||
}
|
||||
|
||||
func (s *modelProbeCacheState) getCachedResult(cacheKey string, now time.Time) (bool, bool) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
entry, ok := s.cache[cacheKey]
|
||||
if !ok || !entry.hasResult {
|
||||
return false, false
|
||||
}
|
||||
if now.Before(entry.nextProbeAt) {
|
||||
return entry.lastResult, true
|
||||
}
|
||||
return false, false
|
||||
}
|
||||
|
||||
func (s *modelProbeCacheState) setCachedResult(cacheKey string, result bool, now time.Time) {
|
||||
s.mu.Lock()
|
||||
|
||||
entry, ok := s.cache[cacheKey]
|
||||
if !ok {
|
||||
entry = &modelProbeCacheEntry{}
|
||||
s.cache[cacheKey] = entry
|
||||
}
|
||||
|
||||
entry.lastResult = result
|
||||
entry.hasResult = true
|
||||
entry.updatedAt = now
|
||||
|
||||
var delay time.Duration
|
||||
if result {
|
||||
entry.successStreak++
|
||||
entry.failureStreak = 0
|
||||
delay = modelProbeBackoffDelay(
|
||||
modelProbeSuccessBaseInterval,
|
||||
modelProbeSuccessMaxInterval,
|
||||
entry.successStreak,
|
||||
)
|
||||
} else {
|
||||
entry.failureStreak++
|
||||
entry.successStreak = 0
|
||||
delay = modelProbeBackoffDelay(
|
||||
modelProbeFailureBaseInterval,
|
||||
modelProbeFailureMaxInterval,
|
||||
entry.failureStreak,
|
||||
)
|
||||
}
|
||||
|
||||
entry.nextProbeAt = now.Add(delay)
|
||||
|
||||
shouldRunTTLGC := modelProbeCacheEntryTTL > 0 && (s.nextTTLGCAt.IsZero() || !now.Before(s.nextTTLGCAt))
|
||||
if shouldRunTTLGC {
|
||||
s.nextTTLGCAt = now.Add(modelProbeTTLGCInterval)
|
||||
}
|
||||
shouldRunSizeGC := len(s.cache) > modelProbeCacheMaxEntries
|
||||
s.mu.Unlock()
|
||||
|
||||
if shouldRunTTLGC || shouldRunSizeGC {
|
||||
s.gc(now, shouldRunTTLGC)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *modelProbeCacheState) gc(now time.Time, runTTL bool) {
|
||||
type evictionCandidate struct {
|
||||
key string
|
||||
updatedAt time.Time
|
||||
}
|
||||
|
||||
var expireBefore time.Time
|
||||
if runTTL && modelProbeCacheEntryTTL > 0 {
|
||||
expireBefore = now.Add(-modelProbeCacheEntryTTL)
|
||||
}
|
||||
|
||||
s.mu.RLock()
|
||||
cacheLen := len(s.cache)
|
||||
if cacheLen == 0 {
|
||||
s.mu.RUnlock()
|
||||
return
|
||||
}
|
||||
|
||||
expiredKeys := make([]string, 0)
|
||||
if !expireBefore.IsZero() {
|
||||
expiredKeys = make([]string, 0, min(cacheLen/8+1, 64))
|
||||
for key, entry := range s.cache {
|
||||
if entry.updatedAt.Before(expireBefore) {
|
||||
expiredKeys = append(expiredKeys, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
effectiveLen := cacheLen - len(expiredKeys)
|
||||
removeCount := max(effectiveLen-modelProbeCacheTrimToEntries, 0)
|
||||
|
||||
candidates := make([]evictionCandidate, 0)
|
||||
if removeCount > 0 {
|
||||
candidates = make([]evictionCandidate, 0, effectiveLen)
|
||||
for key, entry := range s.cache {
|
||||
if !expireBefore.IsZero() && entry.updatedAt.Before(expireBefore) {
|
||||
continue
|
||||
}
|
||||
candidates = append(candidates, evictionCandidate{key: key, updatedAt: entry.updatedAt})
|
||||
}
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
|
||||
if len(expiredKeys) == 0 && len(candidates) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
toEvict := map[string]time.Time{}
|
||||
for i := 0; i < removeCount && len(candidates) > 0; i++ {
|
||||
oldest := 0
|
||||
for j := 1; j < len(candidates); j++ {
|
||||
if candidates[j].updatedAt.Before(candidates[oldest].updatedAt) {
|
||||
oldest = j
|
||||
}
|
||||
}
|
||||
victim := candidates[oldest]
|
||||
toEvict[victim.key] = victim.updatedAt
|
||||
candidates[oldest] = candidates[len(candidates)-1]
|
||||
candidates = candidates[:len(candidates)-1]
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if !expireBefore.IsZero() {
|
||||
for _, key := range expiredKeys {
|
||||
entry, ok := s.cache[key]
|
||||
if ok && entry.updatedAt.Before(expireBefore) {
|
||||
delete(s.cache, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for key, victimUpdatedAt := range toEvict {
|
||||
entry, ok := s.cache[key]
|
||||
if ok && !entry.updatedAt.After(victimUpdatedAt) {
|
||||
delete(s.cache, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func modelProbeBackoffDelay(base, maxDelay time.Duration, streak int) time.Duration {
|
||||
if streak <= 0 {
|
||||
streak = 1
|
||||
}
|
||||
|
||||
shift := min(streak-1, modelProbeBackoffMaxShift)
|
||||
|
||||
delay := base * time.Duration(1<<shift)
|
||||
if maxDelay > 0 && (delay > maxDelay || delay < 0) {
|
||||
return maxDelay
|
||||
}
|
||||
if delay <= 0 {
|
||||
return base
|
||||
}
|
||||
return delay
|
||||
}
|
||||
|
||||
func modelProbeAPIBase(m *config.ModelConfig) string {
|
||||
if apiBase := strings.TrimSpace(m.APIBase); apiBase != "" {
|
||||
return normalizeModelProbeAPIBase(apiBase)
|
||||
|
|
@ -195,7 +474,11 @@ func probeTCPService(raw string) bool {
|
|||
return false
|
||||
}
|
||||
|
||||
conn, err := net.DialTimeout("tcp", hostPort, modelProbeTimeout)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), modelProbeTimeout)
|
||||
defer cancel()
|
||||
|
||||
dialer := &net.Dialer{}
|
||||
conn, err := dialer.DialContext(ctx, "tcp", hostPort)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
|
@ -250,7 +533,10 @@ func probeOpenAICompatibleModel(apiBase, modelID, apiKey string) bool {
|
|||
}
|
||||
|
||||
func getJSON(rawURL string, out any, apiKey string) error {
|
||||
req, err := http.NewRequest(http.MethodGet, rawURL, nil)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), modelProbeTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
|
@ -258,7 +544,7 @@ func getJSON(rawURL string, out any, apiKey string) error {
|
|||
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: modelProbeTimeout}
|
||||
client := &http.Client{}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
|
|
@ -324,10 +610,29 @@ func ollamaModelMatches(candidate, want string) bool {
|
|||
if candidate == "" || want == "" {
|
||||
return false
|
||||
}
|
||||
if strings.EqualFold(candidate, want) {
|
||||
return true
|
||||
|
||||
candidateBase, candidateTag := splitOllamaModel(candidate)
|
||||
wantBase, wantTag := splitOllamaModel(want)
|
||||
if candidateBase == "" || wantBase == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
base, _, _ := strings.Cut(candidate, ":")
|
||||
return strings.EqualFold(base, want)
|
||||
if candidateTag == "" {
|
||||
candidateTag = "latest"
|
||||
}
|
||||
if wantTag == "" {
|
||||
wantTag = "latest"
|
||||
}
|
||||
|
||||
return strings.EqualFold(candidateBase, wantBase) && strings.EqualFold(candidateTag, wantTag)
|
||||
}
|
||||
|
||||
func splitOllamaModel(raw string) (base, tag string) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return "", ""
|
||||
}
|
||||
|
||||
base, tag, _ = strings.Cut(raw, ":")
|
||||
return strings.TrimSpace(base), strings.TrimSpace(tag)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,10 @@ package api
|
|||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
|
@ -85,3 +88,307 @@ func TestProbeLocalModelAvailability_LMStudioUsesOpenAICompatibleProbe(t *testin
|
|||
t.Fatal("probeOpenAICompatibleModelFunc was not called for lmstudio")
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelProbeCacheKey_DifferentAPIKeysProduceDifferentKeys(t *testing.T) {
|
||||
base := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
AuthMethod: "local",
|
||||
ConnectMode: "",
|
||||
}
|
||||
|
||||
m1 := *base
|
||||
m1.SetAPIKey("key-a")
|
||||
m2 := *base
|
||||
m2.SetAPIKey("key-b")
|
||||
|
||||
k1 := modelProbeCacheKey(&m1)
|
||||
k2 := modelProbeCacheKey(&m2)
|
||||
if k1 == k2 {
|
||||
t.Fatal("modelProbeCacheKey() should differ when api key changes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelProbeCacheKey_NormalizesTrailingSlashInAPIBase(t *testing.T) {
|
||||
m1 := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
}
|
||||
m2 := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1/",
|
||||
}
|
||||
|
||||
k1 := modelProbeCacheKey(m1)
|
||||
k2 := modelProbeCacheKey(m2)
|
||||
if k1 != k2 {
|
||||
t.Fatalf("modelProbeCacheKey() mismatch for equivalent api_base values: %q vs %q", k1, k2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelProbeCacheKey_IgnoresDisplayAndConnectionFields(t *testing.T) {
|
||||
base := &config.ModelConfig{
|
||||
ModelName: "vllm-one",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
AuthMethod: "none",
|
||||
ConnectMode: "http",
|
||||
}
|
||||
changed := &config.ModelConfig{
|
||||
ModelName: "vllm-two",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
AuthMethod: "token",
|
||||
ConnectMode: "ws",
|
||||
}
|
||||
|
||||
k1 := modelProbeCacheKey(base)
|
||||
k2 := modelProbeCacheKey(changed)
|
||||
if k1 != k2 {
|
||||
t.Fatalf("modelProbeCacheKey() should ignore non-probe fields, got %q vs %q", k1, k2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeLocalModelAvailability_SuccessBackoff(t *testing.T) {
|
||||
resetModelProbeHooks(t)
|
||||
|
||||
now := time.Unix(1700000000, 0)
|
||||
modelProbeNowFunc = func() time.Time { return now }
|
||||
|
||||
calls := 0
|
||||
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
|
||||
calls++
|
||||
return true
|
||||
}
|
||||
|
||||
model := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
}
|
||||
|
||||
if !probeLocalModelAvailability(model) {
|
||||
t.Fatal("first probe result = false, want true")
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Fatalf("probe calls after first probe = %d, want 1", calls)
|
||||
}
|
||||
|
||||
if !probeLocalModelAvailability(model) {
|
||||
t.Fatal("cached probe result = false, want true")
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Fatalf("probe calls after immediate re-check = %d, want 1", calls)
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeSuccessBaseInterval)
|
||||
if !probeLocalModelAvailability(model) {
|
||||
t.Fatal("second probe result = false, want true")
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("probe calls after success backoff window = %d, want 2", calls)
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeSuccessBaseInterval)
|
||||
if !probeLocalModelAvailability(model) {
|
||||
t.Fatal("cached result after doubled backoff = false, want true")
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("probe calls before doubled backoff expires = %d, want 2", calls)
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeSuccessBaseInterval)
|
||||
if !probeLocalModelAvailability(model) {
|
||||
t.Fatal("third probe result = false, want true")
|
||||
}
|
||||
if calls != 3 {
|
||||
t.Fatalf("probe calls after doubled backoff expires = %d, want 3", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeLocalModelAvailability_FailureBackoff(t *testing.T) {
|
||||
resetModelProbeHooks(t)
|
||||
|
||||
now := time.Unix(1700000100, 0)
|
||||
modelProbeNowFunc = func() time.Time { return now }
|
||||
|
||||
calls := 0
|
||||
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
|
||||
calls++
|
||||
return false
|
||||
}
|
||||
|
||||
model := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
}
|
||||
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("first probe result = true, want false")
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Fatalf("probe calls after first failure = %d, want 1", calls)
|
||||
}
|
||||
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("cached failed probe result = true, want false")
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Fatalf("probe calls after immediate failed re-check = %d, want 1", calls)
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeFailureBaseInterval)
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("second failed probe result = true, want false")
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("probe calls after failure backoff window = %d, want 2", calls)
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeFailureBaseInterval)
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("cached failure after doubled backoff = true, want false")
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("probe calls before doubled failure backoff expires = %d, want 2", calls)
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeFailureBaseInterval)
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("third failed probe result = true, want false")
|
||||
}
|
||||
if calls != 3 {
|
||||
t.Fatalf("probe calls after doubled failure backoff expires = %d, want 3", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeLocalModelAvailability_ResultFlipResetsBackoff(t *testing.T) {
|
||||
resetModelProbeHooks(t)
|
||||
|
||||
now := time.Unix(1700000200, 0)
|
||||
modelProbeNowFunc = func() time.Time { return now }
|
||||
|
||||
results := []bool{true, false, false}
|
||||
index := 0
|
||||
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
|
||||
if index >= len(results) {
|
||||
return false
|
||||
}
|
||||
result := results[index]
|
||||
index++
|
||||
return result
|
||||
}
|
||||
|
||||
model := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
}
|
||||
|
||||
if !probeLocalModelAvailability(model) {
|
||||
t.Fatal("first probe result = false, want true")
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeSuccessBaseInterval)
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("second probe result = true, want false")
|
||||
}
|
||||
|
||||
now = now.Add(modelProbeFailureBaseInterval)
|
||||
if probeLocalModelAvailability(model) {
|
||||
t.Fatal("third probe result = true, want false")
|
||||
}
|
||||
|
||||
if index != 3 {
|
||||
t.Fatalf("probe invocations = %d, want 3", index)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeLocalModelAvailability_DeduplicatesInflightProbe(t *testing.T) {
|
||||
resetModelProbeHooks(t)
|
||||
|
||||
now := time.Unix(1700000300, 0)
|
||||
modelProbeNowFunc = func() time.Time { return now }
|
||||
|
||||
var calls int32
|
||||
probeStarted := make(chan struct{})
|
||||
releaseProbe := make(chan struct{})
|
||||
|
||||
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
|
||||
if atomic.AddInt32(&calls, 1) == 1 {
|
||||
close(probeStarted)
|
||||
}
|
||||
<-releaseProbe
|
||||
return true
|
||||
}
|
||||
|
||||
model := &config.ModelConfig{
|
||||
ModelName: "local-vllm",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
}
|
||||
|
||||
const workers = 8
|
||||
var wg sync.WaitGroup
|
||||
results := make(chan bool, workers)
|
||||
workerStarted := make(chan struct{}, workers)
|
||||
|
||||
for range workers {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
workerStarted <- struct{}{}
|
||||
results <- probeLocalModelAvailability(model)
|
||||
}()
|
||||
}
|
||||
|
||||
for range workers {
|
||||
<-workerStarted
|
||||
}
|
||||
|
||||
select {
|
||||
case <-probeStarted:
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
t.Fatal("probe did not start in time")
|
||||
}
|
||||
|
||||
if got := atomic.LoadInt32(&calls); got != 1 {
|
||||
t.Fatalf("concurrent probe calls = %d, want 1", got)
|
||||
}
|
||||
|
||||
close(releaseProbe)
|
||||
wg.Wait()
|
||||
close(results)
|
||||
|
||||
for result := range results {
|
||||
if !result {
|
||||
t.Fatal("deduplicated probe result = false, want true")
|
||||
}
|
||||
}
|
||||
|
||||
if got := atomic.LoadInt32(&calls); got != 1 {
|
||||
t.Fatalf("final probe calls = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOllamaModelMatches_WithTagRequiresExactTag(t *testing.T) {
|
||||
if ollamaModelMatches("llama3:8b", "llama3:7b") {
|
||||
t.Fatal("ollamaModelMatches() = true, want false for mismatched tags")
|
||||
}
|
||||
if !ollamaModelMatches("llama3:7b", "llama3:7b") {
|
||||
t.Fatal("ollamaModelMatches() = false, want true for exact tagged match")
|
||||
}
|
||||
if ollamaModelMatches("llama3:8b", "llama3") {
|
||||
t.Fatal("ollamaModelMatches() = true, want false when request omits tag (defaults to latest)")
|
||||
}
|
||||
if !ollamaModelMatches("llama3:latest", "llama3") {
|
||||
t.Fatal("ollamaModelMatches() = false, want true when request omits tag and candidate is latest")
|
||||
}
|
||||
if !ollamaModelMatches("llama3", "llama3") {
|
||||
t.Fatal("ollamaModelMatches() = false, want true when both candidate and request omit tag (latest)")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -40,10 +40,11 @@ type modelResponse struct {
|
|||
ThinkingLevel string `json:"thinking_level,omitempty"`
|
||||
ExtraBody map[string]any `json:"extra_body,omitempty"`
|
||||
// Meta
|
||||
Enabled bool `json:"enabled"`
|
||||
Configured bool `json:"configured"`
|
||||
IsDefault bool `json:"is_default"`
|
||||
IsVirtual bool `json:"is_virtual"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Available bool `json:"available"`
|
||||
Status string `json:"status"`
|
||||
IsDefault bool `json:"is_default"`
|
||||
IsVirtual bool `json:"is_virtual"`
|
||||
}
|
||||
|
||||
// handleListModels returns all model_list entries with masked API keys.
|
||||
|
|
@ -57,14 +58,14 @@ func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
|||
}
|
||||
|
||||
defaultModel := cfg.Agents.Defaults.GetModelName()
|
||||
configured := make([]bool, len(cfg.ModelList))
|
||||
modelStatuses := make([]modelConfigurationSummary, len(cfg.ModelList))
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(len(cfg.ModelList))
|
||||
for i, m := range cfg.ModelList {
|
||||
go func(i int, m *config.ModelConfig) {
|
||||
defer wg.Done()
|
||||
configured[i] = isModelConfigured(m)
|
||||
modelStatuses[i] = modelConfigurationStatus(m)
|
||||
}(i, m)
|
||||
}
|
||||
wg.Wait()
|
||||
|
|
@ -87,7 +88,8 @@ func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
|||
ThinkingLevel: m.ThinkingLevel,
|
||||
ExtraBody: m.ExtraBody,
|
||||
Enabled: m.Enabled,
|
||||
Configured: configured[i],
|
||||
Available: modelStatuses[i].Available,
|
||||
Status: modelStatuses[i].Status,
|
||||
IsDefault: m.ModelName == defaultModel,
|
||||
IsVirtual: m.IsVirtual(),
|
||||
})
|
||||
|
|
|
|||
|
|
@ -20,14 +20,18 @@ func resetModelProbeHooks(t *testing.T) {
|
|||
origTCPProbe := probeTCPServiceFunc
|
||||
origOllamaProbe := probeOllamaModelFunc
|
||||
origOpenAIProbe := probeOpenAICompatibleModelFunc
|
||||
origNow := modelProbeNowFunc
|
||||
resetModelProbeCache()
|
||||
t.Cleanup(func() {
|
||||
probeTCPServiceFunc = origTCPProbe
|
||||
probeOllamaModelFunc = origOllamaProbe
|
||||
probeOpenAICompatibleModelFunc = origOpenAIProbe
|
||||
modelProbeNowFunc = origNow
|
||||
resetModelProbeCache()
|
||||
})
|
||||
}
|
||||
|
||||
func TestHandleListModels_ConfiguredStatusUsesRuntimeProbesForLocalModels(t *testing.T) {
|
||||
func TestHandleListModels_AvailabilityUsesRuntimeProbesForLocalModels(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
resetOAuthHooks(t)
|
||||
|
|
@ -113,25 +117,42 @@ func TestHandleListModels_ConfiguredStatusUsesRuntimeProbesForLocalModels(t *tes
|
|||
t.Fatalf("Unmarshal() error = %v", err)
|
||||
}
|
||||
|
||||
got := make(map[string]bool, len(resp.Models))
|
||||
gotAvailable := make(map[string]bool, len(resp.Models))
|
||||
gotStatus := make(map[string]string, len(resp.Models))
|
||||
for _, model := range resp.Models {
|
||||
got[model.ModelName] = model.Configured
|
||||
gotAvailable[model.ModelName] = model.Available
|
||||
gotStatus[model.ModelName] = model.Status
|
||||
}
|
||||
|
||||
if got["openai-oauth"] {
|
||||
t.Fatalf("openai oauth model configured = true, want false without stored credential")
|
||||
if gotAvailable["openai-oauth"] {
|
||||
t.Fatalf("openai oauth model available = true, want false without stored credential")
|
||||
}
|
||||
if !got["vllm-local"] {
|
||||
t.Fatalf("vllm local model configured = false, want true when local probe succeeds")
|
||||
if !gotAvailable["vllm-local"] {
|
||||
t.Fatalf("vllm local model available = false, want true when local probe succeeds")
|
||||
}
|
||||
if !got["ollama-default"] {
|
||||
t.Fatalf("ollama default model configured = false, want true when default local probe succeeds")
|
||||
if !gotAvailable["ollama-default"] {
|
||||
t.Fatalf("ollama default model available = false, want true when default local probe succeeds")
|
||||
}
|
||||
if !got["vllm-remote"] {
|
||||
t.Fatalf("remote vllm model configured = false, want true with api_key")
|
||||
if !gotAvailable["vllm-remote"] {
|
||||
t.Fatalf("remote vllm model available = false, want true with api_key")
|
||||
}
|
||||
if !got["copilot-gpt-5.4"] {
|
||||
t.Fatalf("copilot model configured = false, want true when local bridge probe succeeds")
|
||||
if !gotAvailable["copilot-gpt-5.4"] {
|
||||
t.Fatalf("copilot model available = false, want true when local bridge probe succeeds")
|
||||
}
|
||||
if gotStatus["openai-oauth"] != modelStatusUnconfigured {
|
||||
t.Fatalf("openai oauth model status = %q, want %q", gotStatus["openai-oauth"], modelStatusUnconfigured)
|
||||
}
|
||||
if gotStatus["vllm-local"] != modelStatusAvailable {
|
||||
t.Fatalf("vllm local model status = %q, want %q", gotStatus["vllm-local"], modelStatusAvailable)
|
||||
}
|
||||
if gotStatus["ollama-default"] != modelStatusAvailable {
|
||||
t.Fatalf("ollama default model status = %q, want %q", gotStatus["ollama-default"], modelStatusAvailable)
|
||||
}
|
||||
if gotStatus["vllm-remote"] != modelStatusAvailable {
|
||||
t.Fatalf("remote vllm model status = %q, want %q", gotStatus["vllm-remote"], modelStatusAvailable)
|
||||
}
|
||||
if gotStatus["copilot-gpt-5.4"] != modelStatusAvailable {
|
||||
t.Fatalf("copilot model status = %q, want %q", gotStatus["copilot-gpt-5.4"], modelStatusAvailable)
|
||||
}
|
||||
if len(openAIProbes) != 1 || openAIProbes[0] != "http://127.0.0.1:8000/v1|custom-model|" {
|
||||
t.Fatalf("openAI probes = %#v, want only local vllm probe", openAIProbes)
|
||||
|
|
@ -144,7 +165,7 @@ func TestHandleListModels_ConfiguredStatusUsesRuntimeProbesForLocalModels(t *tes
|
|||
}
|
||||
}
|
||||
|
||||
func TestHandleListModels_ConfiguredStatusForOAuthModelWithCredential(t *testing.T) {
|
||||
func TestHandleListModels_AvailabilityForOAuthModelWithCredential(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
resetOAuthHooks(t)
|
||||
|
|
@ -193,8 +214,8 @@ func TestHandleListModels_ConfiguredStatusForOAuthModelWithCredential(t *testing
|
|||
if len(resp.Models) != 1 {
|
||||
t.Fatalf("len(models) = %d, want 1", len(resp.Models))
|
||||
}
|
||||
if !resp.Models[0].Configured {
|
||||
t.Fatalf("oauth model configured = false, want true with stored credential")
|
||||
if !resp.Models[0].Available {
|
||||
t.Fatalf("oauth model available = false, want true with stored credential")
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -306,14 +327,71 @@ func TestHandleListModels_NormalizesWildcardLocalAPIBaseForProbe(t *testing.T) {
|
|||
if len(resp.Models) != 1 {
|
||||
t.Fatalf("len(models) = %d, want 1", len(resp.Models))
|
||||
}
|
||||
if !resp.Models[0].Configured {
|
||||
t.Fatal("wildcard-bound local model configured = false, want true after probe host normalization")
|
||||
if !resp.Models[0].Available {
|
||||
t.Fatal("wildcard-bound local model available = false, want true after probe host normalization")
|
||||
}
|
||||
if gotProbe != "http://127.0.0.1:8000/v1|custom-model|" {
|
||||
t.Fatalf("probe api base = %q, want %q", gotProbe, "http://127.0.0.1:8000/v1|custom-model|")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleListModels_StatusMarksUnreachableLocalModel(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
resetOAuthHooks(t)
|
||||
resetModelProbeHooks(t)
|
||||
|
||||
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
cfg, err := config.LoadConfig(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConfig() error = %v", err)
|
||||
}
|
||||
cfg.ModelList = []*config.ModelConfig{{
|
||||
ModelName: "vllm-local-down",
|
||||
Model: "vllm/custom-model",
|
||||
APIBase: "http://127.0.0.1:8000/v1",
|
||||
APIKeys: config.SimpleSecureStrings("test-key"),
|
||||
}}
|
||||
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||
t.Fatalf("SaveConfig() error = %v", err)
|
||||
}
|
||||
|
||||
h := NewHandler(configPath)
|
||||
mux := http.NewServeMux()
|
||||
h.RegisterRoutes(mux)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/models", nil)
|
||||
mux.ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Models []modelResponse `json:"models"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("Unmarshal() error = %v", err)
|
||||
}
|
||||
if len(resp.Models) != 1 {
|
||||
t.Fatalf("len(models) = %d, want 1", len(resp.Models))
|
||||
}
|
||||
|
||||
if resp.Models[0].Available {
|
||||
t.Fatal("unreachable local model available = true, want false")
|
||||
}
|
||||
if resp.Models[0].Status != modelStatusUnreachable {
|
||||
t.Fatalf("unreachable local model status = %q, want %q", resp.Models[0].Status, modelStatusUnreachable)
|
||||
}
|
||||
if resp.Models[0].APIKey == "" {
|
||||
t.Fatal("masked API key preview should still be returned when API key is configured")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleAddModel_PersistsAPIKey(t *testing.T) {
|
||||
configPath, cleanup := setupOAuthTestEnv(t)
|
||||
defer cleanup()
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ type Handler struct {
|
|||
serverPublic bool
|
||||
serverPublicExplicit bool
|
||||
serverCIDRs []string
|
||||
debug bool
|
||||
oauthMu sync.Mutex
|
||||
oauthFlows map[string]*oauthFlow
|
||||
oauthState map[string]string
|
||||
|
|
@ -43,6 +44,10 @@ func (h *Handler) SetServerOptions(port int, public bool, publicExplicit bool, a
|
|||
h.serverCIDRs = append([]string(nil), allowedCIDRs...)
|
||||
}
|
||||
|
||||
func (h *Handler) SetDebug(debug bool) {
|
||||
h.debug = debug
|
||||
}
|
||||
|
||||
// RegisterRoutes binds all API endpoint handlers to the ServeMux.
|
||||
func (h *Handler) RegisterRoutes(mux *http.ServeMux) {
|
||||
// Config CRUD
|
||||
|
|
|
|||
|
|
@ -90,6 +90,9 @@ func (h *Handler) resolveLaunchCommand() (string, []string, error) {
|
|||
}
|
||||
|
||||
args := []string{"-no-browser"}
|
||||
if h.debug {
|
||||
args = append(args, "-d")
|
||||
}
|
||||
if h.configPath != "" {
|
||||
args = append(args, h.configPath)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -45,6 +45,29 @@ func TestResolveLaunchCommandUsesConfigFileDefaults(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestResolveLaunchCommandIncludesDebugFlagWhenEnabled(t *testing.T) {
|
||||
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||
h := NewHandler(configPath)
|
||||
h.SetDebug(true)
|
||||
|
||||
_, args, err := h.resolveLaunchCommand()
|
||||
if err != nil {
|
||||
t.Fatalf("resolveLaunchCommand() error = %v", err)
|
||||
}
|
||||
if len(args) != 3 {
|
||||
t.Fatalf("args len = %d, want 3 (got %v)", len(args), args)
|
||||
}
|
||||
if args[0] != "-no-browser" {
|
||||
t.Fatalf("args[0] = %q, want %q", args[0], "-no-browser")
|
||||
}
|
||||
if args[1] != "-d" {
|
||||
t.Fatalf("args[1] = %q, want %q", args[1], "-d")
|
||||
}
|
||||
if args[2] != configPath {
|
||||
t.Fatalf("args[2] = %q, want %q", args[2], configPath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildDarwinPlistIncludesRunAtLoad(t *testing.T) {
|
||||
plist := buildDarwinPlist("/tmp/picoclaw-web", []string{"-no-browser", "/tmp/config.json"})
|
||||
if !strings.Contains(plist, "<key>RunAtLoad</key>") {
|
||||
|
|
|
|||
|
|
@ -55,6 +55,10 @@ var (
|
|||
noBrowser *bool
|
||||
)
|
||||
|
||||
func shouldEnableLauncherFileLogging(enableConsole, debug bool) bool {
|
||||
return !enableConsole || debug
|
||||
}
|
||||
|
||||
func main() {
|
||||
port := flag.String("port", "18800", "Port to listen on")
|
||||
public := flag.Bool("public", false, "Listen on all interfaces (0.0.0.0) instead of localhost only")
|
||||
|
|
@ -62,21 +66,30 @@ func main() {
|
|||
lang := flag.String("lang", "", "Language: en (English) or zh (Chinese). Default: auto-detect from system locale")
|
||||
console := flag.Bool("console", false, "Console mode, no GUI")
|
||||
|
||||
var debug bool
|
||||
flag.BoolVar(&debug, "d", false, "Enable debug logging")
|
||||
flag.BoolVar(&debug, "debug", false, "Enable debug logging")
|
||||
|
||||
flag.Usage = func() {
|
||||
fmt.Fprintf(os.Stderr, "%s Launcher - A web-based configuration editor\n\n", appName)
|
||||
fmt.Fprintf(os.Stderr, "%s Launcher - Web console and gateway manager\n\n", appName)
|
||||
fmt.Fprintf(os.Stderr, "Usage: %s [options] [config.json]\n\n", os.Args[0])
|
||||
fmt.Fprintf(os.Stderr, "Arguments:\n")
|
||||
fmt.Fprintf(os.Stderr, " config.json Path to the configuration file (default: ~/.picoclaw/config.json)\n\n")
|
||||
fmt.Fprintf(os.Stderr, "Options:\n")
|
||||
flag.PrintDefaults()
|
||||
fmt.Fprintf(os.Stderr, "\nExamples:\n")
|
||||
fmt.Fprintf(os.Stderr, " %s Use default config path\n", os.Args[0])
|
||||
fmt.Fprintf(os.Stderr, " %s ./config.json Specify a config file\n", os.Args[0])
|
||||
fmt.Fprintf(os.Stderr, " %s\n", os.Args[0])
|
||||
fmt.Fprintf(os.Stderr, " Use default config path in GUI mode\n")
|
||||
fmt.Fprintf(os.Stderr, " %s ./config.json\n", os.Args[0])
|
||||
fmt.Fprintf(os.Stderr, " Specify a config file\n")
|
||||
fmt.Fprintf(
|
||||
os.Stderr,
|
||||
" %s -public ./config.json Allow access from other devices on the network\n",
|
||||
" %s -public ./config.json\n",
|
||||
os.Args[0],
|
||||
)
|
||||
fmt.Fprintf(os.Stderr, " Allow access from other devices on the local network\n")
|
||||
fmt.Fprintf(os.Stderr, " %s -console -d ./config.json\n", os.Args[0])
|
||||
fmt.Fprintf(os.Stderr, " Run in the terminal with debug logs enabled\n")
|
||||
}
|
||||
flag.Parse()
|
||||
|
||||
|
|
@ -90,12 +103,13 @@ func main() {
|
|||
}
|
||||
defer panicFunc()
|
||||
|
||||
// By default, detect terminal to decide console log behavior
|
||||
// If -console-logs flag is explicitly set, it overrides the detection
|
||||
enableConsole := *console
|
||||
if !enableConsole {
|
||||
// Disable console logging by setting level to Fatal (no output)
|
||||
logger.SetConsoleLevel(logger.FATAL)
|
||||
fileLoggingEnabled := shouldEnableLauncherFileLogging(enableConsole, debug)
|
||||
if fileLoggingEnabled {
|
||||
// GUI mode writes launcher logs to file. Debug mode keeps file logging enabled in console mode too.
|
||||
if !debug {
|
||||
logger.DisableConsole()
|
||||
}
|
||||
|
||||
f := filepath.Join(picoHome, logPath, logFile)
|
||||
if err = logger.EnableFileLogging(f); err != nil {
|
||||
|
|
@ -103,9 +117,9 @@ func main() {
|
|||
}
|
||||
defer logger.DisableFileLogging()
|
||||
}
|
||||
|
||||
logger.InfoC("web", fmt.Sprintf("%s launcher starting (version %s)...", appName, appVersion))
|
||||
logger.InfoC("web", fmt.Sprintf("%s Home: %s", appName, picoHome))
|
||||
if debug {
|
||||
logger.SetLevel(logger.DEBUG)
|
||||
}
|
||||
|
||||
// Set language from command line or auto-detect
|
||||
if *lang != "" {
|
||||
|
|
@ -126,6 +140,25 @@ func main() {
|
|||
if err != nil {
|
||||
logger.Errorf("Warning: Failed to initialize %s config automatically: %v", appName, err)
|
||||
}
|
||||
if !debug {
|
||||
logger.SetLevelFromString(config.ResolveGatewayLogLevel(absPath))
|
||||
}
|
||||
|
||||
logger.InfoC("web", fmt.Sprintf("%s launcher starting (version %s)...", appName, appVersion))
|
||||
logger.InfoC("web", fmt.Sprintf("%s Home: %s", appName, picoHome))
|
||||
if debug {
|
||||
logger.InfoC("web", "Debug mode enabled")
|
||||
logger.DebugC(
|
||||
"web",
|
||||
fmt.Sprintf(
|
||||
"Launcher flags: console=%t public=%t no_browser=%t config=%s",
|
||||
enableConsole,
|
||||
*public,
|
||||
*noBrowser,
|
||||
absPath,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
var explicitPort bool
|
||||
var explicitPublic bool
|
||||
|
|
@ -181,7 +214,7 @@ func main() {
|
|||
mux := http.NewServeMux()
|
||||
|
||||
tokenLogFileAbs := ""
|
||||
if !enableConsole {
|
||||
if fileLoggingEnabled {
|
||||
tokenLogFileAbs = filepath.Join(picoHome, logPath, logFile)
|
||||
}
|
||||
api.RegisterLauncherAuthRoutes(mux, api.LauncherAuthRouteOpts{
|
||||
|
|
@ -197,6 +230,7 @@ func main() {
|
|||
|
||||
// API Routes (e.g. /api/status)
|
||||
apiHandler = api.NewHandler(absPath)
|
||||
apiHandler.SetDebug(debug)
|
||||
if _, err = apiHandler.EnsurePicoChannel(""); err != nil {
|
||||
logger.ErrorC("web", fmt.Sprintf("Warning: failed to ensure pico channel on startup: %v", err))
|
||||
}
|
||||
|
|
@ -226,7 +260,7 @@ func main() {
|
|||
)
|
||||
|
||||
// Print startup banner and token (console mode only).
|
||||
if enableConsole {
|
||||
if enableConsole || debug {
|
||||
fmt.Print(utils.Banner)
|
||||
fmt.Println()
|
||||
fmt.Println(" Open the following URL in your browser:")
|
||||
|
|
|
|||
31
web/backend/main_test.go
Normal file
31
web/backend/main_test.go
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
package main
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestShouldEnableLauncherFileLogging(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
enableConsole bool
|
||||
debug bool
|
||||
want bool
|
||||
}{
|
||||
{name: "gui mode", enableConsole: false, debug: false, want: true},
|
||||
{name: "console mode", enableConsole: true, debug: false, want: false},
|
||||
{name: "debug gui mode", enableConsole: false, debug: true, want: true},
|
||||
{name: "debug console mode", enableConsole: true, debug: true, want: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := shouldEnableLauncherFileLogging(tt.enableConsole, tt.debug); got != tt.want {
|
||||
t.Fatalf(
|
||||
"shouldEnableLauncherFileLogging(%t, %t) = %t, want %t",
|
||||
tt.enableConsole,
|
||||
tt.debug,
|
||||
got,
|
||||
tt.want,
|
||||
)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -20,7 +20,8 @@ export interface ModelInfo {
|
|||
thinking_level?: string
|
||||
extra_body?: Record<string, unknown>
|
||||
// Meta
|
||||
configured: boolean
|
||||
available: boolean
|
||||
status: "available" | "unconfigured" | "unreachable"
|
||||
is_default: boolean
|
||||
is_virtual: boolean
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,19 +10,19 @@ import { useTranslation } from "react-i18next"
|
|||
import { Button } from "@/components/ui/button"
|
||||
|
||||
interface ChatEmptyStateProps {
|
||||
hasConfiguredModels: boolean
|
||||
hasAvailableModels: boolean
|
||||
defaultModelName: string
|
||||
isConnected: boolean
|
||||
}
|
||||
|
||||
export function ChatEmptyState({
|
||||
hasConfiguredModels,
|
||||
hasAvailableModels,
|
||||
defaultModelName,
|
||||
isConnected,
|
||||
}: ChatEmptyStateProps) {
|
||||
const { t } = useTranslation()
|
||||
|
||||
if (!hasConfiguredModels) {
|
||||
if (!hasAvailableModels) {
|
||||
return (
|
||||
<div className="flex flex-col items-center justify-center py-20 opacity-70">
|
||||
<div className="mb-6 flex h-16 w-16 items-center justify-center rounded-2xl bg-amber-500/10 text-amber-500">
|
||||
|
|
|
|||
|
|
@ -39,7 +39,7 @@ export function ChatPage() {
|
|||
|
||||
const {
|
||||
defaultModelName,
|
||||
hasConfiguredModels,
|
||||
hasAvailableModels,
|
||||
apiKeyModels,
|
||||
oauthModels,
|
||||
localModels,
|
||||
|
|
@ -94,7 +94,7 @@ export function ChatPage() {
|
|||
hasScrolled ? "shadow-sm" : "shadow-none"
|
||||
}`}
|
||||
titleExtra={
|
||||
hasConfiguredModels && (
|
||||
hasAvailableModels && (
|
||||
<ModelSelector
|
||||
defaultModelName={defaultModelName}
|
||||
apiKeyModels={apiKeyModels}
|
||||
|
|
@ -140,7 +140,7 @@ export function ChatPage() {
|
|||
<div className="mx-auto flex w-full max-w-250 flex-col gap-8 pb-8">
|
||||
{messages.length === 0 && !isTyping && (
|
||||
<ChatEmptyState
|
||||
hasConfiguredModels={hasConfiguredModels}
|
||||
hasAvailableModels={hasAvailableModels}
|
||||
defaultModelName={defaultModelName}
|
||||
isConnected={isGatewayRunning}
|
||||
/>
|
||||
|
|
|
|||
102
web/frontend/src/components/logs/log-level-select.tsx
Normal file
102
web/frontend/src/components/logs/log-level-select.tsx
Normal file
|
|
@ -0,0 +1,102 @@
|
|||
import { useQuery, useQueryClient } from "@tanstack/react-query"
|
||||
import { useEffect, useState } from "react"
|
||||
import { useTranslation } from "react-i18next"
|
||||
import { toast } from "sonner"
|
||||
|
||||
import { type AppConfig, getAppConfig, patchAppConfig } from "@/api/channels"
|
||||
import {
|
||||
Select,
|
||||
SelectContent,
|
||||
SelectItem,
|
||||
SelectTrigger,
|
||||
SelectValue,
|
||||
} from "@/components/ui/select"
|
||||
import { refreshGatewayState } from "@/store/gateway"
|
||||
|
||||
const LOG_LEVEL_OPTIONS = ["debug", "info", "warn", "error", "fatal"] as const
|
||||
type GatewayLogLevel = (typeof LOG_LEVEL_OPTIONS)[number]
|
||||
|
||||
const LOG_LEVEL_LABELS: Record<GatewayLogLevel, string> = {
|
||||
debug: "Debug",
|
||||
info: "Info",
|
||||
warn: "Warn",
|
||||
error: "Error",
|
||||
fatal: "Fatal",
|
||||
}
|
||||
|
||||
function getGatewayLogLevel(config: AppConfig | undefined): GatewayLogLevel {
|
||||
const gateway = config?.gateway
|
||||
if (typeof gateway === "object" && gateway !== null) {
|
||||
const logLevel = (gateway as Record<string, unknown>).log_level
|
||||
if (
|
||||
typeof logLevel === "string" &&
|
||||
LOG_LEVEL_OPTIONS.includes(logLevel as GatewayLogLevel)
|
||||
) {
|
||||
return logLevel as GatewayLogLevel
|
||||
}
|
||||
}
|
||||
return "warn"
|
||||
}
|
||||
|
||||
export function LogLevelSelect() {
|
||||
const { t } = useTranslation()
|
||||
const queryClient = useQueryClient()
|
||||
const [logLevel, setLogLevel] = useState<GatewayLogLevel>("warn")
|
||||
const [savingLogLevel, setSavingLogLevel] = useState(false)
|
||||
|
||||
const { data: configData } = useQuery({
|
||||
queryKey: ["config"],
|
||||
queryFn: getAppConfig,
|
||||
})
|
||||
|
||||
useEffect(() => {
|
||||
setLogLevel(getGatewayLogLevel(configData))
|
||||
}, [configData])
|
||||
|
||||
const handleLogLevelChange = async (nextValue: string) => {
|
||||
const nextLevel = nextValue as GatewayLogLevel
|
||||
const previousLevel = logLevel
|
||||
setLogLevel(nextLevel)
|
||||
setSavingLogLevel(true)
|
||||
|
||||
try {
|
||||
await patchAppConfig({
|
||||
gateway: {
|
||||
log_level: nextLevel,
|
||||
},
|
||||
})
|
||||
await queryClient.invalidateQueries({ queryKey: ["config"] })
|
||||
await refreshGatewayState({ force: true })
|
||||
} catch (error) {
|
||||
setLogLevel(previousLevel)
|
||||
toast.error(
|
||||
error instanceof Error
|
||||
? error.message
|
||||
: t("pages.logs.log_level_error"),
|
||||
)
|
||||
} finally {
|
||||
setSavingLogLevel(false)
|
||||
}
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="flex items-center gap-2">
|
||||
<Select
|
||||
value={logLevel}
|
||||
onValueChange={handleLogLevelChange}
|
||||
disabled={savingLogLevel}
|
||||
>
|
||||
<SelectTrigger size="sm" className="w-28">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent align="end">
|
||||
{LOG_LEVEL_OPTIONS.map((level) => (
|
||||
<SelectItem key={level} value={level}>
|
||||
{LOG_LEVEL_LABELS[level]}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
import { IconTrash } from "@tabler/icons-react"
|
||||
import { useTranslation } from "react-i18next"
|
||||
|
||||
import { LogLevelSelect } from "@/components/logs/log-level-select"
|
||||
import { LogsPanel } from "@/components/logs/logs-panel"
|
||||
import { PageHeader } from "@/components/page-header"
|
||||
import { Button } from "@/components/ui/button"
|
||||
|
|
@ -17,15 +18,19 @@ export function LogsPage() {
|
|||
<PageHeader
|
||||
title={t("navigation.logs")}
|
||||
children={
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
onClick={clearLogs}
|
||||
disabled={logs.length === 0 || clearing}
|
||||
>
|
||||
<IconTrash className="size-4" />
|
||||
{t("pages.logs.clear")}
|
||||
</Button>
|
||||
<>
|
||||
<LogLevelSelect />
|
||||
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
onClick={clearLogs}
|
||||
disabled={logs.length === 0 || clearing}
|
||||
>
|
||||
<IconTrash className="size-4" />
|
||||
{t("pages.logs.clear")}
|
||||
</Button>
|
||||
</>
|
||||
}
|
||||
/>
|
||||
|
||||
|
|
|
|||
|
|
@ -133,9 +133,10 @@ export function EditModelSheet({
|
|||
}
|
||||
|
||||
const isOAuth = model?.auth_method === "oauth"
|
||||
const apiKeyPlaceholder = model?.configured
|
||||
const hasSavedAPIKey = Boolean(model?.api_key)
|
||||
const apiKeyPlaceholder = hasSavedAPIKey
|
||||
? maskedSecretPlaceholder(
|
||||
model.api_key,
|
||||
model?.api_key ?? "",
|
||||
t("models.field.apiKeyPlaceholderSet"),
|
||||
)
|
||||
: t("models.field.apiKeyPlaceholder")
|
||||
|
|
@ -161,7 +162,7 @@ export function EditModelSheet({
|
|||
<Field
|
||||
label={t("models.field.apiKey")}
|
||||
hint={
|
||||
model?.configured ? t("models.edit.apiKeyHint") : undefined
|
||||
hasSavedAPIKey ? t("models.edit.apiKeyHint") : undefined
|
||||
}
|
||||
>
|
||||
<KeyInput
|
||||
|
|
|
|||
|
|
@ -28,14 +28,16 @@ export function ModelCard({
|
|||
}: ModelCardProps) {
|
||||
const { t } = useTranslation()
|
||||
const isOAuth = model.auth_method === "oauth"
|
||||
const status = model.status
|
||||
const statusLabel = t(`models.status.${status}`)
|
||||
const canSetDefault =
|
||||
model.configured && !model.is_default && !model.is_virtual
|
||||
model.available && !model.is_default && !model.is_virtual
|
||||
|
||||
return (
|
||||
<div
|
||||
className={[
|
||||
"group/card hover:bg-muted/30 relative flex w-full max-w-[36rem] flex-col gap-3 justify-self-start rounded-xl border p-4 transition-colors hover:shadow-xs",
|
||||
model.configured
|
||||
model.available
|
||||
? "border-border/60 bg-card"
|
||||
: "border-border/50 bg-card/60",
|
||||
].join(" ")}
|
||||
|
|
@ -47,15 +49,13 @@ export function ModelCard({
|
|||
"mt-0.5 h-2 w-2 shrink-0 rounded-full",
|
||||
model.is_default
|
||||
? "bg-green-400 shadow-[0_0_0_2px_rgba(74,222,128,0.35)]"
|
||||
: model.configured
|
||||
: status === "available"
|
||||
? "bg-green-500"
|
||||
: status === "unreachable"
|
||||
? "bg-amber-500"
|
||||
: "bg-muted-foreground/25",
|
||||
].join(" ")}
|
||||
title={
|
||||
model.configured
|
||||
? t("models.status.configured")
|
||||
: t("models.status.unconfigured")
|
||||
}
|
||||
title={statusLabel}
|
||||
/>
|
||||
<span className="text-foreground truncate text-sm font-semibold">
|
||||
{model.model_name}
|
||||
|
|
@ -127,14 +127,14 @@ export function ModelCard({
|
|||
<span className="text-muted-foreground bg-muted rounded px-1.5 py-0.5 text-[10px] font-medium">
|
||||
OAuth
|
||||
</span>
|
||||
) : model.configured && model.api_key ? (
|
||||
) : status === "available" && model.api_key ? (
|
||||
<span className="text-muted-foreground/70 flex items-center gap-1 font-mono text-[11px]">
|
||||
<IconKey className="size-3" />
|
||||
{model.api_key}
|
||||
</span>
|
||||
) : (
|
||||
<span className="text-muted-foreground/50 text-[11px]">
|
||||
{t("models.status.unconfigured")}
|
||||
{statusLabel}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ interface ProviderGroup {
|
|||
label: string
|
||||
models: ModelInfo[]
|
||||
hasDefault: boolean
|
||||
configuredCount: number
|
||||
availableCount: number
|
||||
}
|
||||
|
||||
export function ModelsPage() {
|
||||
|
|
@ -62,8 +62,8 @@ export function ModelsPage() {
|
|||
const sorted = [...data.models].sort((a, b) => {
|
||||
if (a.is_default && !b.is_default) return -1
|
||||
if (!a.is_default && b.is_default) return 1
|
||||
if (a.configured && !b.configured) return -1
|
||||
if (!a.configured && b.configured) return 1
|
||||
if (a.available && !b.available) return -1
|
||||
if (!a.available && b.available) return 1
|
||||
return a.model_name.localeCompare(b.model_name)
|
||||
})
|
||||
setModels(sorted)
|
||||
|
|
@ -107,23 +107,23 @@ export function ModelsPage() {
|
|||
|
||||
const providerGroups: ProviderGroup[] = Object.entries(grouped)
|
||||
.map(([key, group]) => {
|
||||
const configuredCount = group.models.filter(
|
||||
(model) => model.configured,
|
||||
const availableCount = group.models.filter(
|
||||
(model) => model.available,
|
||||
).length
|
||||
return {
|
||||
key,
|
||||
label: group.label,
|
||||
models: group.models,
|
||||
hasDefault: group.models.some((model) => model.is_default),
|
||||
configuredCount,
|
||||
availableCount,
|
||||
}
|
||||
})
|
||||
.sort((a, b) => {
|
||||
if (a.hasDefault && !b.hasDefault) return -1
|
||||
if (!a.hasDefault && b.hasDefault) return 1
|
||||
|
||||
if (a.configuredCount !== b.configuredCount) {
|
||||
return b.configuredCount - a.configuredCount
|
||||
if (a.availableCount !== b.availableCount) {
|
||||
return b.availableCount - a.availableCount
|
||||
}
|
||||
|
||||
const aPriority = PROVIDER_PRIORITY[a.key] ?? Number.MAX_SAFE_INTEGER
|
||||
|
|
|
|||
|
|
@ -65,32 +65,32 @@ export function useChatModels({ isConnected }: UseChatModelsOptions) {
|
|||
[defaultModelName],
|
||||
)
|
||||
|
||||
const hasConfiguredModels = useMemo(
|
||||
() => modelList.some((m) => m.configured),
|
||||
const hasAvailableModels = useMemo(
|
||||
() => modelList.some((m) => m.available),
|
||||
[modelList],
|
||||
)
|
||||
|
||||
const oauthModels = useMemo(
|
||||
() => modelList.filter((m) => m.configured && m.auth_method === "oauth"),
|
||||
() => modelList.filter((m) => m.available && m.auth_method === "oauth"),
|
||||
[modelList],
|
||||
)
|
||||
|
||||
const localModels = useMemo(
|
||||
() => modelList.filter((m) => m.configured && isLocalModel(m)),
|
||||
() => modelList.filter((m) => m.available && isLocalModel(m)),
|
||||
[modelList],
|
||||
)
|
||||
|
||||
const apiKeyModels = useMemo(
|
||||
() =>
|
||||
modelList.filter(
|
||||
(m) => m.configured && m.auth_method !== "oauth" && !isLocalModel(m),
|
||||
(m) => m.available && m.auth_method !== "oauth" && !isLocalModel(m),
|
||||
),
|
||||
[modelList],
|
||||
)
|
||||
|
||||
return {
|
||||
defaultModelName,
|
||||
hasConfiguredModels,
|
||||
hasAvailableModels,
|
||||
apiKeyModels,
|
||||
oauthModels,
|
||||
localModels,
|
||||
|
|
|
|||
|
|
@ -170,8 +170,9 @@
|
|||
"noDefaultHintPrefix": "No default model set yet. Click",
|
||||
"noDefaultHintSuffix": "to set one.",
|
||||
"status": {
|
||||
"configured": "Configured",
|
||||
"unconfigured": "Not configured"
|
||||
"available": "Available",
|
||||
"unconfigured": "Not configured",
|
||||
"unreachable": "Service unreachable"
|
||||
},
|
||||
"badge": {
|
||||
"default": "Default",
|
||||
|
|
@ -547,6 +548,7 @@
|
|||
"unsaved_changes": "You have unsaved changes."
|
||||
},
|
||||
"logs": {
|
||||
"log_level_error": "Failed to update log level.",
|
||||
"clear": "Clear logs",
|
||||
"empty": "Waiting for logs..."
|
||||
}
|
||||
|
|
|
|||
|
|
@ -170,8 +170,9 @@
|
|||
"noDefaultHintPrefix": "尚未设置默认模型,点击",
|
||||
"noDefaultHintSuffix": "设为默认。",
|
||||
"status": {
|
||||
"configured": "已配置",
|
||||
"unconfigured": "未配置"
|
||||
"available": "可用",
|
||||
"unconfigured": "未配置",
|
||||
"unreachable": "服务不可达"
|
||||
},
|
||||
"badge": {
|
||||
"default": "默认",
|
||||
|
|
@ -547,6 +548,7 @@
|
|||
"unsaved_changes": "您有未保存的更改。"
|
||||
},
|
||||
"logs": {
|
||||
"log_level_error": "更新日志等级失败。",
|
||||
"clear": "清空日志",
|
||||
"empty": "等待日志中..."
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue