Merge branch 'sipeed:main' into main
This commit is contained in:
commit
b7f5343e23
126 changed files with 21113 additions and 3193 deletions
|
|
@ -5,6 +5,7 @@
|
||||||
# ANTHROPIC_API_KEY=sk-ant-xxx
|
# ANTHROPIC_API_KEY=sk-ant-xxx
|
||||||
# OPENAI_API_KEY=sk-xxx
|
# OPENAI_API_KEY=sk-xxx
|
||||||
# GEMINI_API_KEY=xxx
|
# GEMINI_API_KEY=xxx
|
||||||
|
# CEREBRAS_API_KEY=xxx
|
||||||
|
|
||||||
# ── Chat Channel ──────────────────────────
|
# ── Chat Channel ──────────────────────────
|
||||||
# TELEGRAM_BOT_TOKEN=123456:ABC...
|
# TELEGRAM_BOT_TOKEN=123456:ABC...
|
||||||
|
|
|
||||||
30
.github/pull_request_template.md
vendored
30
.github/pull_request_template.md
vendored
|
|
@ -1,4 +1,7 @@
|
||||||
## 📝 Description
|
## 📝 Description
|
||||||
|
|
||||||
|
<!-- Please briefly describe the changes and purpose of this PR -->
|
||||||
|
|
||||||
## 🗣️ Type of Change
|
## 🗣️ Type of Change
|
||||||
- [ ] 🐞 Bug fix (non-breaking change which fixes an issue)
|
- [ ] 🐞 Bug fix (non-breaking change which fixes an issue)
|
||||||
- [ ] ✨ New feature (non-breaking change which adds functionality)
|
- [ ] ✨ New feature (non-breaking change which adds functionality)
|
||||||
|
|
@ -11,25 +14,28 @@
|
||||||
- [ ] 👨💻 Mostly Human-written (Human lead, AI assisted or none)
|
- [ ] 👨💻 Mostly Human-written (Human lead, AI assisted or none)
|
||||||
|
|
||||||
|
|
||||||
## 🔗 Linked Issue
|
## 🔗 Related Issue
|
||||||
|
|
||||||
|
<!-- Please link the related issue(s) (e.g., Fixes #123, Closes #456) -->
|
||||||
|
|
||||||
## 📚 Technical Context (Skip for Docs)
|
## 📚 Technical Context (Skip for Docs)
|
||||||
* **Reference:** [URL]
|
- **Reference URL:**
|
||||||
* **Reasoning:** ...
|
- **Reasoning:**
|
||||||
|
|
||||||
|
## 🧪 Test Environment
|
||||||
|
- **Hardware:** <!-- e.g. Raspberry Pi 5, Orange Pi, PC-->
|
||||||
|
- **OS:** <!-- e.g. Debian 12, Ubuntu 22.04 -->
|
||||||
|
- **Model/Provider:** <!-- e.g. OpenAI GPT-4o, Kimi k2, DeepSeek-V3 -->
|
||||||
|
- **Channels:** <!-- e.g. Discord, Telegram, Feishu, ... -->
|
||||||
|
|
||||||
|
|
||||||
## 🧪 Test Environment & Hardware
|
## 📸 Evidence (Optional)
|
||||||
- **Hardware:** [e.g. Raspberry Pi 5, Orange Pi, PC]
|
|
||||||
- **OS:** [e.g. Debian 12, Ubuntu 22.04]
|
|
||||||
- **Model/Provider:** [e.g. OpenAI GPT-4o, Kimi k2, DeepSeek-V3]
|
|
||||||
- **Channels:** [e.g. Discord, Telegram, Feishu, ...]
|
|
||||||
|
|
||||||
|
|
||||||
## 📸 Proof of Work (Optional for Docs)
|
|
||||||
<details>
|
<details>
|
||||||
<summary>Click to view Logs/Screenshots</summary>
|
<summary>Click to view Logs/Screenshots</summary>
|
||||||
|
|
||||||
</details>
|
<!-- Please paste relevant screenshots or logs here -->
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
## ☑️ Checklist
|
## ☑️ Checklist
|
||||||
- [ ] My code/docs follow the style of this project.
|
- [ ] My code/docs follow the style of this project.
|
||||||
|
|
|
||||||
4
.github/workflows/build.yml
vendored
4
.github/workflows/build.yml
vendored
|
|
@ -9,10 +9,10 @@ jobs:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Checkout
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Setup Go
|
- name: Setup Go
|
||||||
uses: actions/setup-go@v5
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
|
|
||||||
2
.github/workflows/docker-build.yml
vendored
2
.github/workflows/docker-build.yml
vendored
|
|
@ -25,7 +25,7 @@ jobs:
|
||||||
steps:
|
steps:
|
||||||
# ── Checkout ──────────────────────────────
|
# ── Checkout ──────────────────────────────
|
||||||
- name: 📥 Checkout repository
|
- name: 📥 Checkout repository
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
with:
|
with:
|
||||||
ref: ${{ inputs.tag }}
|
ref: ${{ inputs.tag }}
|
||||||
|
|
||||||
|
|
|
||||||
44
.github/workflows/pr.yml
vendored
44
.github/workflows/pr.yml
vendored
|
|
@ -1,17 +1,39 @@
|
||||||
name: pr-check
|
name: PR
|
||||||
|
|
||||||
on:
|
on:
|
||||||
pull_request:
|
pull_request: { }
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
fmt-check:
|
lint:
|
||||||
|
name: Linter
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Checkout
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Setup Go
|
- name: Setup Go
|
||||||
uses: actions/setup-go@v5
|
uses: actions/setup-go@v6
|
||||||
|
with:
|
||||||
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
- name: Run go generate
|
||||||
|
run: go generate ./...
|
||||||
|
|
||||||
|
- name: Golangci Lint
|
||||||
|
uses: golangci/golangci-lint-action@v9
|
||||||
|
with:
|
||||||
|
version: v2.10.1
|
||||||
|
|
||||||
|
# TODO: Remove once linter is properly configured
|
||||||
|
fmt-check:
|
||||||
|
name: Formatting
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- name: Checkout
|
||||||
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
|
- name: Setup Go
|
||||||
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
|
@ -20,15 +42,17 @@ jobs:
|
||||||
make fmt
|
make fmt
|
||||||
git diff --exit-code || (echo "::error::Code is not formatted. Run 'make fmt' and commit the changes." && exit 1)
|
git diff --exit-code || (echo "::error::Code is not formatted. Run 'make fmt' and commit the changes." && exit 1)
|
||||||
|
|
||||||
|
# TODO: Remove once linter is properly configured
|
||||||
vet:
|
vet:
|
||||||
|
name: Vet
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
needs: fmt-check
|
needs: fmt-check
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Checkout
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Setup Go
|
- name: Setup Go
|
||||||
uses: actions/setup-go@v5
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
|
@ -39,14 +63,15 @@ jobs:
|
||||||
run: go vet ./...
|
run: go vet ./...
|
||||||
|
|
||||||
test:
|
test:
|
||||||
|
name: Tests
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
needs: fmt-check
|
needs: fmt-check
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Checkout
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Setup Go
|
- name: Setup Go
|
||||||
uses: actions/setup-go@v5
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
|
@ -55,4 +80,3 @@ jobs:
|
||||||
|
|
||||||
- name: Run go test
|
- name: Run go test
|
||||||
run: go test ./...
|
run: go test ./...
|
||||||
|
|
||||||
|
|
|
||||||
8
.github/workflows/release.yml
vendored
8
.github/workflows/release.yml
vendored
|
|
@ -26,7 +26,7 @@ jobs:
|
||||||
contents: write
|
contents: write
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Checkout
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
with:
|
with:
|
||||||
fetch-depth: 0
|
fetch-depth: 0
|
||||||
|
|
||||||
|
|
@ -49,13 +49,14 @@ jobs:
|
||||||
packages: write
|
packages: write
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout tag
|
- name: Checkout tag
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v6
|
||||||
with:
|
with:
|
||||||
fetch-depth: 0
|
fetch-depth: 0
|
||||||
ref: ${{ inputs.tag }}
|
ref: ${{ inputs.tag }}
|
||||||
|
|
||||||
- name: Setup Go from go.mod
|
- name: Setup Go from go.mod
|
||||||
uses: actions/setup-go@v5
|
id: setup-go
|
||||||
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
|
@ -89,6 +90,7 @@ jobs:
|
||||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
GITHUB_REPOSITORY_OWNER: ${{ github.repository_owner }}
|
GITHUB_REPOSITORY_OWNER: ${{ github.repository_owner }}
|
||||||
DOCKERHUB_IMAGE_NAME: ${{ vars.DOCKERHUB_REPOSITORY }}
|
DOCKERHUB_IMAGE_NAME: ${{ vars.DOCKERHUB_REPOSITORY }}
|
||||||
|
GOVERSION: ${{ steps.setup-go.outputs.go-version }}
|
||||||
|
|
||||||
- name: Apply release flags
|
- name: Apply release flags
|
||||||
shell: bash
|
shell: bash
|
||||||
|
|
|
||||||
184
.golangci.yaml
Normal file
184
.golangci.yaml
Normal file
|
|
@ -0,0 +1,184 @@
|
||||||
|
version: "2"
|
||||||
|
|
||||||
|
linters:
|
||||||
|
default: all
|
||||||
|
disable:
|
||||||
|
# TODO: Tweak for current project needs
|
||||||
|
- containedctx
|
||||||
|
- cyclop
|
||||||
|
- depguard
|
||||||
|
- dupl
|
||||||
|
- dupword
|
||||||
|
- err113
|
||||||
|
- exhaustruct
|
||||||
|
- funcorder
|
||||||
|
- gochecknoglobals
|
||||||
|
- godot
|
||||||
|
- intrange
|
||||||
|
- ireturn
|
||||||
|
- nlreturn
|
||||||
|
- noctx
|
||||||
|
- noinlineerr
|
||||||
|
- nonamedreturns
|
||||||
|
- tagliatelle
|
||||||
|
- testpackage
|
||||||
|
- varnamelen
|
||||||
|
- wrapcheck
|
||||||
|
- wsl
|
||||||
|
- wsl_v5
|
||||||
|
|
||||||
|
# TODO: Disabled, because they are failing at the moment, we should fix them and enable (step by step)
|
||||||
|
- bodyclose
|
||||||
|
- contextcheck
|
||||||
|
- dogsled
|
||||||
|
- embeddedstructfieldcheck
|
||||||
|
- errcheck
|
||||||
|
- errchkjson
|
||||||
|
- errorlint
|
||||||
|
- exhaustive
|
||||||
|
- forbidigo
|
||||||
|
- forcetypeassert
|
||||||
|
- funlen
|
||||||
|
- gochecknoinits
|
||||||
|
- gocognit
|
||||||
|
- goconst
|
||||||
|
- gocritic
|
||||||
|
- gocyclo
|
||||||
|
- godox
|
||||||
|
- goprintffuncname
|
||||||
|
- gosec
|
||||||
|
- govet
|
||||||
|
- ineffassign
|
||||||
|
- lll
|
||||||
|
- maintidx
|
||||||
|
- misspell
|
||||||
|
- mnd
|
||||||
|
- modernize
|
||||||
|
- nakedret
|
||||||
|
- nestif
|
||||||
|
- nilnil
|
||||||
|
- paralleltest
|
||||||
|
- perfsprint
|
||||||
|
- prealloc
|
||||||
|
- predeclared
|
||||||
|
- revive
|
||||||
|
- staticcheck
|
||||||
|
- tagalign
|
||||||
|
- testifylint
|
||||||
|
- thelper
|
||||||
|
- unparam
|
||||||
|
- unused
|
||||||
|
- usestdlibvars
|
||||||
|
- usetesting
|
||||||
|
- wastedassign
|
||||||
|
- whitespace
|
||||||
|
settings:
|
||||||
|
errcheck:
|
||||||
|
check-type-assertions: true
|
||||||
|
check-blank: true
|
||||||
|
exhaustive:
|
||||||
|
default-signifies-exhaustive: true
|
||||||
|
funlen:
|
||||||
|
lines: 120
|
||||||
|
statements: 40
|
||||||
|
gocognit:
|
||||||
|
min-complexity: 25
|
||||||
|
gocyclo:
|
||||||
|
min-complexity: 20
|
||||||
|
govet:
|
||||||
|
enable-all: true
|
||||||
|
disable:
|
||||||
|
- fieldalignment
|
||||||
|
lll:
|
||||||
|
line-length: 120
|
||||||
|
tab-width: 4
|
||||||
|
misspell:
|
||||||
|
locale: US
|
||||||
|
mnd:
|
||||||
|
checks:
|
||||||
|
- argument
|
||||||
|
- assign
|
||||||
|
- case
|
||||||
|
- condition
|
||||||
|
- operation
|
||||||
|
- return
|
||||||
|
nakedret:
|
||||||
|
max-func-lines: 3
|
||||||
|
revive:
|
||||||
|
enable-all-rules: true
|
||||||
|
rules:
|
||||||
|
- name: add-constant
|
||||||
|
disabled: true
|
||||||
|
- name: argument-limit
|
||||||
|
arguments:
|
||||||
|
- 7
|
||||||
|
severity: warning
|
||||||
|
- name: banned-characters
|
||||||
|
disabled: true
|
||||||
|
- name: cognitive-complexity
|
||||||
|
disabled: true
|
||||||
|
- name: comment-spacings
|
||||||
|
arguments:
|
||||||
|
- nolint
|
||||||
|
severity: warning
|
||||||
|
- name: cyclomatic
|
||||||
|
disabled: true
|
||||||
|
- name: file-header
|
||||||
|
disabled: true
|
||||||
|
- name: function-result-limit
|
||||||
|
arguments:
|
||||||
|
- 3
|
||||||
|
severity: warning
|
||||||
|
- name: function-length
|
||||||
|
disabled: true
|
||||||
|
- name: line-length-limit
|
||||||
|
disabled: true
|
||||||
|
- name: max-public-structs
|
||||||
|
disabled: true
|
||||||
|
- name: modifies-value-receiver
|
||||||
|
disabled: true
|
||||||
|
- name: package-comments
|
||||||
|
disabled: true
|
||||||
|
- name: unused-receiver
|
||||||
|
disabled: true
|
||||||
|
exclusions:
|
||||||
|
generated: lax
|
||||||
|
rules:
|
||||||
|
- linters:
|
||||||
|
- lll
|
||||||
|
source: '^//go:generate '
|
||||||
|
- linters:
|
||||||
|
- funlen
|
||||||
|
- maintidx
|
||||||
|
- gocognit
|
||||||
|
- gocyclo
|
||||||
|
path: _test\.go$
|
||||||
|
|
||||||
|
issues:
|
||||||
|
max-issues-per-linter: 0
|
||||||
|
max-same-issues: 0
|
||||||
|
|
||||||
|
formatters:
|
||||||
|
enable:
|
||||||
|
- goimports
|
||||||
|
# TODO: Disabled, because they are failing at the moment, we should fix them and enable (step by step)
|
||||||
|
# - gci
|
||||||
|
# - gofmt
|
||||||
|
# - gofumpt
|
||||||
|
# - golines
|
||||||
|
settings:
|
||||||
|
gci:
|
||||||
|
sections:
|
||||||
|
- standard
|
||||||
|
- default
|
||||||
|
- localmodule
|
||||||
|
custom-order: true
|
||||||
|
gofmt:
|
||||||
|
simplify: true
|
||||||
|
rewrite-rules:
|
||||||
|
- pattern: "interface{}"
|
||||||
|
replacement: "any"
|
||||||
|
- pattern: "a[b:len(a)]"
|
||||||
|
replacement: "a[b:]"
|
||||||
|
golines:
|
||||||
|
max-len: 120
|
||||||
|
|
@ -11,6 +11,14 @@ builds:
|
||||||
- id: picoclaw
|
- id: picoclaw
|
||||||
env:
|
env:
|
||||||
- CGO_ENABLED=0
|
- CGO_ENABLED=0
|
||||||
|
tags:
|
||||||
|
- stdjson
|
||||||
|
ldflags:
|
||||||
|
- -s -w
|
||||||
|
- -X main.version={{ .Version }}
|
||||||
|
- -X main.gitCommit={{ .ShortCommit }}
|
||||||
|
- -X main.buildTime={{ .Date }}
|
||||||
|
- -X main.goVersion={{ .Env.GOVERSION }}
|
||||||
goos:
|
goos:
|
||||||
- linux
|
- linux
|
||||||
- windows
|
- windows
|
||||||
|
|
|
||||||
|
|
@ -29,7 +29,14 @@ HEALTHCHECK --interval=30s --timeout=3s --start-period=5s --retries=3 \
|
||||||
# Copy binary
|
# Copy binary
|
||||||
COPY --from=builder /src/build/picoclaw /usr/local/bin/picoclaw
|
COPY --from=builder /src/build/picoclaw /usr/local/bin/picoclaw
|
||||||
|
|
||||||
# Create picoclaw home directory
|
# Create non-root user and group
|
||||||
|
RUN addgroup -g 1000 picoclaw && \
|
||||||
|
adduser -D -u 1000 -G picoclaw picoclaw
|
||||||
|
|
||||||
|
# Switch to non-root user
|
||||||
|
USER picoclaw
|
||||||
|
|
||||||
|
# Run onboard to create initial directories and config
|
||||||
RUN /usr/local/bin/picoclaw onboard
|
RUN /usr/local/bin/picoclaw onboard
|
||||||
|
|
||||||
ENTRYPOINT ["picoclaw"]
|
ENTRYPOINT ["picoclaw"]
|
||||||
|
|
|
||||||
4
Makefile
4
Makefile
|
|
@ -11,11 +11,11 @@ VERSION?=$(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
|
||||||
GIT_COMMIT=$(shell git rev-parse --short=8 HEAD 2>/dev/null || echo "dev")
|
GIT_COMMIT=$(shell git rev-parse --short=8 HEAD 2>/dev/null || echo "dev")
|
||||||
BUILD_TIME=$(shell date +%FT%T%z)
|
BUILD_TIME=$(shell date +%FT%T%z)
|
||||||
GO_VERSION=$(shell $(GO) version | awk '{print $$3}')
|
GO_VERSION=$(shell $(GO) version | awk '{print $$3}')
|
||||||
LDFLAGS=-ldflags "-X main.version=$(VERSION) -X main.gitCommit=$(GIT_COMMIT) -X main.buildTime=$(BUILD_TIME) -X main.goVersion=$(GO_VERSION)"
|
LDFLAGS=-ldflags "-X main.version=$(VERSION) -X main.gitCommit=$(GIT_COMMIT) -X main.buildTime=$(BUILD_TIME) -X main.goVersion=$(GO_VERSION) -s -w"
|
||||||
|
|
||||||
# Go variables
|
# Go variables
|
||||||
GO?=go
|
GO?=go
|
||||||
GOFLAGS?=-v
|
GOFLAGS?=-v -tags stdjson
|
||||||
|
|
||||||
# Installation
|
# Installation
|
||||||
INSTALL_PREFIX?=$(HOME)/.local
|
INSTALL_PREFIX?=$(HOME)/.local
|
||||||
|
|
|
||||||
1038
README.fr.md
Normal file
1038
README.fr.md
Normal file
File diff suppressed because it is too large
Load diff
193
README.ja.md
193
README.ja.md
|
|
@ -3,7 +3,7 @@
|
||||||
|
|
||||||
<h1>PicoClaw: Go で書かれた超効率 AI アシスタント</h1>
|
<h1>PicoClaw: Go で書かれた超効率 AI アシスタント</h1>
|
||||||
|
|
||||||
<h3>$10 ハードウェア · 10MB RAM · 1秒起動 · 皮皮虾,我们走!</h3>
|
<h3>$10 ハードウェア · 10MB RAM · 1秒起動 · 行くぜ、シャコ!</h3>
|
||||||
<h3></h3>
|
<h3></h3>
|
||||||
|
|
||||||
<p>
|
<p>
|
||||||
|
|
@ -12,7 +12,7 @@
|
||||||
<img src="https://img.shields.io/badge/license-MIT-green" alt="License">
|
<img src="https://img.shields.io/badge/license-MIT-green" alt="License">
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
**日本語** | [English](README.md)
|
[中文](README.zh.md) | **日本語** | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | [English](README.md)
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
|
@ -39,7 +39,7 @@
|
||||||
</table>
|
</table>
|
||||||
|
|
||||||
## 📢 ニュース
|
## 📢 ニュース
|
||||||
2026-02-09 🎉 PicoClaw リリース!$10 ハードウェアで 10MB 未満の RAM で動く AI エージェントを 1 日で構築。🦐 皮皮虾,我们走!
|
2026-02-09 🎉 PicoClaw リリース!$10 ハードウェアで 10MB 未満の RAM で動く AI エージェントを 1 日で構築。🦐 行くぜ、シャコ!
|
||||||
|
|
||||||
## ✨ 特徴
|
## ✨ 特徴
|
||||||
|
|
||||||
|
|
@ -209,7 +209,7 @@ picoclaw onboard
|
||||||
|
|
||||||
**3. API キーの取得**
|
**3. API キーの取得**
|
||||||
|
|
||||||
- **LLM プロバイダー**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
- **LLM プロバイダー**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys) · [Qwen](https://dashscope.console.aliyun.com)
|
||||||
- **Web 検索**(任意): [Brave Search](https://brave.com/search/api) - 無料枠あり(月 2000 リクエスト)
|
- **Web 検索**(任意): [Brave Search](https://brave.com/search/api) - 無料枠あり(月 2000 リクエスト)
|
||||||
|
|
||||||
> **注意**: 完全な設定テンプレートは `config.example.json` を参照してください。
|
> **注意**: 完全な設定テンプレートは `config.example.json` を参照してください。
|
||||||
|
|
@ -253,7 +253,7 @@ Telegram、Discord、QQ、DingTalk、LINE で PicoClaw と会話できます
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allowFrom": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -293,7 +293,7 @@ picoclaw gateway
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allowFrom": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -621,6 +621,22 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
- `PICOCLAW_HEARTBEAT_ENABLED=false` で無効化
|
- `PICOCLAW_HEARTBEAT_ENABLED=false` で無効化
|
||||||
- `PICOCLAW_HEARTBEAT_INTERVAL=60` で間隔変更
|
- `PICOCLAW_HEARTBEAT_INTERVAL=60` で間隔変更
|
||||||
|
|
||||||
|
### プロバイダー
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> Groq は Whisper による無料の音声文字起こしを提供しています。設定すると、Telegram の音声メッセージが自動的に文字起こしされます。
|
||||||
|
|
||||||
|
| プロバイダー | 用途 | API キー取得先 |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| `gemini` | LLM(Gemini 直接) | [aistudio.google.com](https://aistudio.google.com) |
|
||||||
|
| `zhipu` | LLM(Zhipu 直接) | [bigmodel.cn](https://bigmodel.cn) |
|
||||||
|
| `openrouter`(要テスト) | LLM(推奨、全モデルにアクセス可能) | [openrouter.ai](https://openrouter.ai) |
|
||||||
|
| `anthropic`(要テスト) | LLM(Claude 直接) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
|
| `openai`(要テスト) | LLM(GPT 直接) | [platform.openai.com](https://platform.openai.com) |
|
||||||
|
| `deepseek`(要テスト) | LLM(DeepSeek 直接) | [platform.deepseek.com](https://platform.deepseek.com) |
|
||||||
|
| `groq` | LLM + **音声文字起こし**(Whisper) | [console.groq.com](https://console.groq.com) |
|
||||||
|
| `cerebras` | LLM(Cerebras 直接) | [cerebras.ai](https://cerebras.ai) |
|
||||||
|
|
||||||
### 基本設定
|
### 基本設定
|
||||||
|
|
||||||
1. **設定ファイルの作成:**
|
1. **設定ファイルの作成:**
|
||||||
|
|
@ -676,7 +692,7 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "123456:ABC...",
|
"token": "123456:ABC...",
|
||||||
"allowFrom": ["123456789"]
|
"allow_from": ["123456789"]
|
||||||
},
|
},
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
|
|
@ -692,7 +708,7 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
"appSecret": "xxx",
|
"appSecret": "xxx",
|
||||||
"encryptKey": "",
|
"encryptKey": "",
|
||||||
"verificationToken": "",
|
"verificationToken": "",
|
||||||
"allowFrom": []
|
"allow_from": []
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"tools": {
|
"tools": {
|
||||||
|
|
@ -714,6 +730,163 @@ HEARTBEAT_OK 応答 ユーザーが直接結果を受け取る
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
### モデル設定 (model_list)
|
||||||
|
|
||||||
|
> **新機能!** PicoClaw は現在 **モデル中心** の設定アプローチを採用しています。`ベンダー/モデル` 形式(例: `zhipu/glm-4.7`)を指定するだけで、新しいプロバイダーを追加できます—**コードの変更は一切不要!**
|
||||||
|
|
||||||
|
この設計は、柔軟なプロバイダー選択による **マルチエージェントサポート** も可能にします:
|
||||||
|
|
||||||
|
- **異なるエージェント、異なるプロバイダー** : 各エージェントは独自の LLM プロバイダーを使用可能
|
||||||
|
- **フォールバックモデル** : 耐障性のため、プライマリモデルとフォールバックモデルを設定可能
|
||||||
|
- **ロードバランシング** : 複数のエンドポイントにリクエストを分散
|
||||||
|
- **集中設定管理** : すべてのプロバイダーを一箇所で管理
|
||||||
|
|
||||||
|
#### 📋 サポートされているすべてのベンダー
|
||||||
|
|
||||||
|
| ベンダー | `model` プレフィックス | デフォルト API Base | プロトコル | API キー |
|
||||||
|
|-------------|-----------------|---------------------|----------|---------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [キーを取得](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [キーを取得](https://console.anthropic.com) |
|
||||||
|
| **Zhipu AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [キーを取得](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [キーを取得](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [キーを取得](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [キーを取得](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [キーを取得](https://platform.moonshot.cn) |
|
||||||
|
| **Qwen (Alibaba)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [キーを取得](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [キーを取得](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | ローカル(キー不要) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [キーを取得](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | ローカル |
|
||||||
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [キーを取得](https://cerebras.ai) |
|
||||||
|
| **Volcengine** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [キーを取得](https://console.volcengine.com) |
|
||||||
|
| **ShengsuanYun** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | カスタム | OAuthのみ |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### 基本設定
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### ベンダー別の例
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Zhipu AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (OAuth使用)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> OAuth認証を設定するには、`picoclaw auth login --provider anthropic` を実行してください。
|
||||||
|
|
||||||
|
#### ロードバランシング
|
||||||
|
|
||||||
|
同じモデル名で複数のエンドポイントを設定すると、PicoClaw が自動的にラウンドロビンで分散します:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 従来の `providers` 設定からの移行
|
||||||
|
|
||||||
|
古い `providers` 設定は**非推奨**ですが、後方互換性のためにサポートされています。
|
||||||
|
|
||||||
|
**旧設定(非推奨):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**新設定(推奨):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
詳細な移行ガイドは、[docs/migration/model-list-migration.md](docs/migration/model-list-migration.md) を参照してください。
|
||||||
|
|
||||||
## CLI リファレンス
|
## CLI リファレンス
|
||||||
|
|
||||||
| コマンド | 説明 |
|
| コマンド | 説明 |
|
||||||
|
|
@ -735,7 +908,7 @@ Discord: https://discord.gg/V4sAZ9XWpN
|
||||||
|
|
||||||
## 🐛 トラブルシューティング
|
## 🐛 トラブルシューティング
|
||||||
|
|
||||||
### Web 検索で「API 配置问题」と表示される
|
### Web 検索で「API 設定の問題」と表示される
|
||||||
|
|
||||||
検索 API キーをまだ設定していない場合、これは正常です。PicoClaw は手動検索用の便利なリンクを提供します。
|
検索 API キーをまだ設定していない場合、これは正常です。PicoClaw は手動検索用の便利なリンクを提供します。
|
||||||
|
|
||||||
|
|
@ -771,5 +944,7 @@ Web 検索を有効にするには:
|
||||||
|---------|--------|------------|
|
|---------|--------|------------|
|
||||||
| **OpenRouter** | 月 200K トークン | 複数モデル(Claude, GPT-4 など) |
|
| **OpenRouter** | 月 200K トークン | 複数モデル(Claude, GPT-4 など) |
|
||||||
| **Zhipu** | 月 200K トークン | 中国ユーザー向け最適 |
|
| **Zhipu** | 月 200K トークン | 中国ユーザー向け最適 |
|
||||||
|
| **Qwen** | 無料枠あり | 通義千問 (Qwen) |
|
||||||
| **Brave Search** | 月 2000 クエリ | Web 検索機能 |
|
| **Brave Search** | 月 2000 クエリ | Web 検索機能 |
|
||||||
| **Groq** | 無料枠あり | 高速推論(Llama, Mixtral) |
|
| **Groq** | 無料枠あり | 高速推論(Llama, Mixtral) |
|
||||||
|
| **Cerebras** | 無料枠あり | 高速推論(Llama, Qwen など) |
|
||||||
|
|
|
||||||
228
README.md
228
README.md
|
|
@ -14,7 +14,7 @@
|
||||||
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
[中文](README.zh.md) | [日本語](README.ja.md) | **English**
|
[中文](README.zh.md) | [日本語](README.ja.md) | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | **English**
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
@ -209,18 +209,24 @@ picoclaw onboard
|
||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"model": "glm-4.7",
|
"model": "gpt4",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
"model_list": [
|
||||||
"openrouter": {
|
{
|
||||||
"api_key": "xxx",
|
"model_name": "gpt4",
|
||||||
"api_base": "https://openrouter.ai/api/v1"
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "your-api-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "your-anthropic-key"
|
||||||
}
|
}
|
||||||
},
|
],
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": {
|
"web": {
|
||||||
"brave": {
|
"brave": {
|
||||||
|
|
@ -237,6 +243,8 @@ picoclaw onboard
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **New**: The `model_list` configuration format allows zero-code provider addition. See [Model Configuration](#-model-configuration) for details.
|
||||||
|
|
||||||
**3. Get API Keys**
|
**3. Get API Keys**
|
||||||
|
|
||||||
* **LLM Provider**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
* **LLM Provider**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
||||||
|
|
@ -283,7 +291,7 @@ Talk to your picoclaw through Telegram, Discord, DingTalk, or LINE
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allowFrom": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -326,7 +334,8 @@ picoclaw gateway
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allowFrom": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"],
|
||||||
|
"mention_only": false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -339,6 +348,10 @@ picoclaw gateway
|
||||||
* Bot Permissions: `Send Messages`, `Read Message History`
|
* Bot Permissions: `Send Messages`, `Read Message History`
|
||||||
* Open the generated invite URL and add the bot to your server
|
* Open the generated invite URL and add the bot to your server
|
||||||
|
|
||||||
|
**Optional: Mention-only mode**
|
||||||
|
|
||||||
|
Set `"mention_only": true` to make the bot respond only when @-mentioned. Useful for shared servers where you want the bot to respond only when explicitly called.
|
||||||
|
|
||||||
**6. Run**
|
**6. Run**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
@ -677,7 +690,203 @@ The subagent has access to tools (message, web_search, etc.) and can communicate
|
||||||
| `anthropic(To be tested)` | LLM (Claude direct) | [console.anthropic.com](https://console.anthropic.com) |
|
| `anthropic(To be tested)` | LLM (Claude direct) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
| `openai(To be tested)` | LLM (GPT direct) | [platform.openai.com](https://platform.openai.com) |
|
| `openai(To be tested)` | LLM (GPT direct) | [platform.openai.com](https://platform.openai.com) |
|
||||||
| `deepseek(To be tested)` | LLM (DeepSeek direct) | [platform.deepseek.com](https://platform.deepseek.com) |
|
| `deepseek(To be tested)` | LLM (DeepSeek direct) | [platform.deepseek.com](https://platform.deepseek.com) |
|
||||||
|
| `qwen` | LLM (Qwen direct) | [dashscope.console.aliyun.com](https://dashscope.console.aliyun.com) |
|
||||||
| `groq` | LLM + **Voice transcription** (Whisper) | [console.groq.com](https://console.groq.com) |
|
| `groq` | LLM + **Voice transcription** (Whisper) | [console.groq.com](https://console.groq.com) |
|
||||||
|
| `cerebras` | LLM (Cerebras direct) | [cerebras.ai](https://cerebras.ai) |
|
||||||
|
|
||||||
|
### Model Configuration (model_list)
|
||||||
|
|
||||||
|
> **What's New?** PicoClaw now uses a **model-centric** configuration approach. Simply specify `vendor/model` format (e.g., `zhipu/glm-4.7`) to add new providers—**zero code changes required!**
|
||||||
|
|
||||||
|
This design also enables **multi-agent support** with flexible provider selection:
|
||||||
|
|
||||||
|
- **Different agents, different providers**: Each agent can use its own LLM provider
|
||||||
|
- **Model fallbacks**: Configure primary and fallback models for resilience
|
||||||
|
- **Load balancing**: Distribute requests across multiple endpoints
|
||||||
|
- **Centralized configuration**: Manage all providers in one place
|
||||||
|
|
||||||
|
#### 📋 All Supported Vendors
|
||||||
|
|
||||||
|
| Vendor | `model` Prefix | Default API Base | Protocol | API Key |
|
||||||
|
|--------|----------------|------------------|----------|---------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [Get Key](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Get Key](https://console.anthropic.com) |
|
||||||
|
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Get Key](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Get Key](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Get Key](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Get Key](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [Get Key](https://platform.moonshot.cn) |
|
||||||
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Get Key](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Get Key](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (no key needed) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||||
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Get Key](https://cerebras.ai) |
|
||||||
|
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://console.volcengine.com) |
|
||||||
|
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth only |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### Basic Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Vendor-Specific Examples
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**智谱 AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**DeepSeek**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (with OAuth)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> Run `picoclaw auth login --provider anthropic` to set up OAuth credentials.
|
||||||
|
|
||||||
|
**Ollama (local)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "llama3",
|
||||||
|
"model": "ollama/llama3"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Custom Proxy/API**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "my-custom-model",
|
||||||
|
"model": "openai/custom-model",
|
||||||
|
"api_base": "https://my-proxy.com/v1",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Load Balancing
|
||||||
|
|
||||||
|
Configure multiple endpoints for the same model name—PicoClaw will automatically round-robin between them:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Migration from Legacy `providers` Config
|
||||||
|
|
||||||
|
The old `providers` configuration is **deprecated** but still supported for backward compatibility.
|
||||||
|
|
||||||
|
**Old Config (deprecated):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**New Config (recommended):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
For detailed migration guide, see [docs/migration/model-list-migration.md](docs/migration/model-list-migration.md).
|
||||||
|
|
||||||
|
### Provider Architecture
|
||||||
|
|
||||||
|
PicoClaw routes providers by protocol family:
|
||||||
|
|
||||||
|
- OpenAI-compatible protocol: OpenRouter, OpenAI-compatible gateways, Groq, Zhipu, and vLLM-style endpoints.
|
||||||
|
- Anthropic protocol: Claude-native API behavior.
|
||||||
|
- Codex/OAuth path: OpenAI OAuth/token authentication route.
|
||||||
|
|
||||||
|
This keeps the runtime lightweight while making new OpenAI-compatible backends mostly a config operation (`api_base` + `api_key`).
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Zhipu</b></summary>
|
<summary><b>Zhipu</b></summary>
|
||||||
|
|
@ -873,3 +1082,4 @@ This happens when another instance of the bot is running. Make sure only one `pi
|
||||||
| **Zhipu** | 200K tokens/month | Best for Chinese users |
|
| **Zhipu** | 200K tokens/month | Best for Chinese users |
|
||||||
| **Brave Search** | 2000 queries/month | Web search functionality |
|
| **Brave Search** | 2000 queries/month | Web search functionality |
|
||||||
| **Groq** | Free tier available | Fast inference (Llama, Mixtral) |
|
| **Groq** | Free tier available | Fast inference (Llama, Mixtral) |
|
||||||
|
| **Cerebras** | Free tier available | Fast inference (Llama, Qwen, etc.) |
|
||||||
|
|
|
||||||
1039
README.pt-br.md
Normal file
1039
README.pt-br.md
Normal file
File diff suppressed because it is too large
Load diff
1016
README.vi.md
Normal file
1016
README.vi.md
Normal file
File diff suppressed because it is too large
Load diff
215
README.zh.md
215
README.zh.md
|
|
@ -14,7 +14,7 @@
|
||||||
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
<a href="https://x.com/SipeedIO"><img src="https://img.shields.io/badge/X_(Twitter)-SipeedIO-black?style=flat&logo=x&logoColor=white" alt="Twitter"></a>
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
**中文** | [日本語](README.ja.md) | [English](README.md)
|
**中文** | [日本語](README.ja.md) | [Português](README.pt-br.md) | [Tiếng Việt](README.vi.md) | [Français](README.fr.md) | [English](README.md)
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
@ -218,18 +218,24 @@ picoclaw onboard
|
||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"model": "glm-4.7",
|
"model": "gpt4",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
"model_list": [
|
||||||
"openrouter": {
|
{
|
||||||
"api_key": "xxx",
|
"model_name": "gpt4",
|
||||||
"api_base": "https://openrouter.ai/api/v1"
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "your-api-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "your-anthropic-key"
|
||||||
}
|
}
|
||||||
},
|
],
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
|
|
@ -245,6 +251,8 @@ picoclaw onboard
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **新功能**: `model_list` 配置格式支持零代码添加 provider。详见[模型配置](#-模型配置-model_list)章节。
|
||||||
|
|
||||||
**3. 获取 API Key**
|
**3. 获取 API Key**
|
||||||
|
|
||||||
* **LLM 提供商**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
* **LLM 提供商**: [OpenRouter](https://openrouter.ai/keys) · [Zhipu](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) · [Anthropic](https://console.anthropic.com) · [OpenAI](https://platform.openai.com) · [Gemini](https://aistudio.google.com/api-keys)
|
||||||
|
|
@ -291,7 +299,7 @@ picoclaw agent -m "2+2 等于几?"
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allowFrom": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -336,7 +344,7 @@ picoclaw gateway
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_BOT_TOKEN",
|
"token": "YOUR_BOT_TOKEN",
|
||||||
"allowFrom": ["YOUR_USER_ID"]
|
"allow_from": ["YOUR_USER_ID"]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -554,7 +562,193 @@ Agent 读取 HEARTBEAT.md
|
||||||
| `anthropic(待测试)` | LLM (Claude 直连) | [console.anthropic.com](https://console.anthropic.com) |
|
| `anthropic(待测试)` | LLM (Claude 直连) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
| `openai(待测试)` | LLM (GPT 直连) | [platform.openai.com](https://platform.openai.com) |
|
| `openai(待测试)` | LLM (GPT 直连) | [platform.openai.com](https://platform.openai.com) |
|
||||||
| `deepseek(待测试)` | LLM (DeepSeek 直连) | [platform.deepseek.com](https://platform.deepseek.com) |
|
| `deepseek(待测试)` | LLM (DeepSeek 直连) | [platform.deepseek.com](https://platform.deepseek.com) |
|
||||||
|
| `qwen` | LLM (通义千问) | [dashscope.console.aliyun.com](https://dashscope.console.aliyun.com) |
|
||||||
| `groq` | LLM + **语音转录** (Whisper) | [console.groq.com](https://console.groq.com) |
|
| `groq` | LLM + **语音转录** (Whisper) | [console.groq.com](https://console.groq.com) |
|
||||||
|
| `cerebras` | LLM (Cerebras 直连) | [cerebras.ai](https://cerebras.ai) |
|
||||||
|
|
||||||
|
### 模型配置 (model_list)
|
||||||
|
|
||||||
|
> **新功能!** PicoClaw 现在采用**以模型为中心**的配置方式。只需使用 `厂商/模型` 格式(如 `zhipu/glm-4.7`)即可添加新的 provider——**无需修改任何代码!**
|
||||||
|
|
||||||
|
该设计同时支持**多 Agent 场景**,提供灵活的 Provider 选择:
|
||||||
|
|
||||||
|
- **不同 Agent 使用不同 Provider**:每个 Agent 可以使用自己的 LLM provider
|
||||||
|
- **模型回退(Fallback)**:配置主模型和备用模型,提高可靠性
|
||||||
|
- **负载均衡**:在多个 API 端点之间分配请求
|
||||||
|
- **集中化配置**:在一个地方管理所有 provider
|
||||||
|
|
||||||
|
#### 📋 所有支持的厂商
|
||||||
|
|
||||||
|
| 厂商 | `model` 前缀 | 默认 API Base | 协议 | 获取 API Key |
|
||||||
|
|------|-------------|---------------|------|--------------|
|
||||||
|
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [获取密钥](https://platform.openai.com) |
|
||||||
|
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [获取密钥](https://console.anthropic.com) |
|
||||||
|
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [获取密钥](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||||
|
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [获取密钥](https://platform.deepseek.com) |
|
||||||
|
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [获取密钥](https://aistudio.google.com/api-keys) |
|
||||||
|
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [获取密钥](https://console.groq.com) |
|
||||||
|
| **Moonshot** | `moonshot/` | `https://api.moonshot.cn/v1` | OpenAI | [获取密钥](https://platform.moonshot.cn) |
|
||||||
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [获取密钥](https://dashscope.console.aliyun.com) |
|
||||||
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [获取密钥](https://build.nvidia.com) |
|
||||||
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | 本地(无需密钥) |
|
||||||
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取密钥](https://openrouter.ai/keys) |
|
||||||
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
||||||
|
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [获取密钥](https://cerebras.ai) |
|
||||||
|
| **火山引擎** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [获取密钥](https://console.volcengine.com) |
|
||||||
|
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||||
|
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
|
||||||
|
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||||
|
|
||||||
|
#### 基础配置示例
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-zhipu-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 各厂商配置示例
|
||||||
|
|
||||||
|
**OpenAI**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**智谱 AI (GLM)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**DeepSeek**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Anthropic (使用 OAuth)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
> 运行 `picoclaw auth login --provider anthropic` 来设置 OAuth 凭证。
|
||||||
|
|
||||||
|
**Ollama (本地)**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "llama3",
|
||||||
|
"model": "ollama/llama3"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**自定义代理/API**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "my-custom-model",
|
||||||
|
"model": "openai/custom-model",
|
||||||
|
"api_base": "https://my-proxy.com/v1",
|
||||||
|
"api_key": "sk-..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 负载均衡
|
||||||
|
|
||||||
|
为同一个模型名称配置多个端点——PicoClaw 会自动在它们之间轮询:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api1.example.com/v1",
|
||||||
|
"api_key": "sk-key1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_base": "https://api2.example.com/v1",
|
||||||
|
"api_key": "sk-key2"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 从旧的 `providers` 配置迁移
|
||||||
|
|
||||||
|
旧的 `providers` 配置格式**已弃用**,但为向后兼容仍支持。
|
||||||
|
|
||||||
|
**旧配置(已弃用):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"zhipu": {
|
||||||
|
"api_key": "your-key",
|
||||||
|
"api_base": "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "zhipu",
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**新配置(推荐):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "glm-4.7",
|
||||||
|
"model": "zhipu/glm-4.7",
|
||||||
|
"api_key": "your-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "glm-4.7"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
详细的迁移指南请参考 [docs/migration/model-list-migration.md](docs/migration/model-list-migration.md)。
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>智谱 (Zhipu) 配置示例</b></summary>
|
<summary><b>智谱 (Zhipu) 配置示例</b></summary>
|
||||||
|
|
@ -741,4 +935,5 @@ Discord: [https://discord.gg/V4sAZ9XWpN](https://discord.gg/V4sAZ9XWpN)
|
||||||
| **OpenRouter** | 200K tokens/月 | 多模型聚合 (Claude, GPT-4 等) |
|
| **OpenRouter** | 200K tokens/月 | 多模型聚合 (Claude, GPT-4 等) |
|
||||||
| **智谱 (Zhipu)** | 200K tokens/月 | 最适合中国用户 |
|
| **智谱 (Zhipu)** | 200K tokens/月 | 最适合中国用户 |
|
||||||
| **Brave Search** | 2000 次查询/月 | 网络搜索功能 |
|
| **Brave Search** | 2000 次查询/月 | 网络搜索功能 |
|
||||||
| **Groq** | 提供免费层级 | 极速推理 (Llama, Mixtral) |
|
| **Groq** | 提供免费层级 | 极速推理 (Llama, Mixtral) |
|
||||||
|
| **Cerebras** | 提供免费层级 | 极速推理 (Llama, Qwen 等) |
|
||||||
Binary file not shown.
|
Before Width: | Height: | Size: 142 KiB After Width: | Height: | Size: 141 KiB |
181
cmd/picoclaw/cmd_agent.go
Normal file
181
cmd/picoclaw/cmd_agent.go
Normal file
|
|
@ -0,0 +1,181 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/chzyer/readline"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/agent"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func agentCmd() {
|
||||||
|
message := ""
|
||||||
|
sessionKey := "cli:default"
|
||||||
|
modelOverride := ""
|
||||||
|
|
||||||
|
args := os.Args[2:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--debug", "-d":
|
||||||
|
logger.SetLevel(logger.DEBUG)
|
||||||
|
fmt.Println("🔍 Debug mode enabled")
|
||||||
|
case "-m", "--message":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
message = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-s", "--session":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
sessionKey = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--model", "-model":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
modelOverride = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelOverride != "" {
|
||||||
|
cfg.Agents.Defaults.Model = modelOverride
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := providers.CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error creating provider: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
// Use the resolved model ID from provider creation
|
||||||
|
if modelID != "" {
|
||||||
|
cfg.Agents.Defaults.Model = modelID
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
// Print agent startup info (only for interactive mode)
|
||||||
|
startupInfo := agentLoop.GetStartupInfo()
|
||||||
|
logger.InfoCF("agent", "Agent initialized",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tools_count": startupInfo["tools"].(map[string]interface{})["count"],
|
||||||
|
"skills_total": startupInfo["skills"].(map[string]interface{})["total"],
|
||||||
|
"skills_available": startupInfo["skills"].(map[string]interface{})["available"],
|
||||||
|
})
|
||||||
|
|
||||||
|
if message != "" {
|
||||||
|
ctx := context.Background()
|
||||||
|
response, err := agentLoop.ProcessDirect(ctx, message, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
fmt.Printf("\n%s %s\n", logo, response)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("%s Interactive mode (Ctrl+C to exit)\n\n", logo)
|
||||||
|
interactiveMode(agentLoop, sessionKey)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func interactiveMode(agentLoop *agent.AgentLoop, sessionKey string) {
|
||||||
|
prompt := fmt.Sprintf("%s You: ", logo)
|
||||||
|
|
||||||
|
rl, err := readline.NewEx(&readline.Config{
|
||||||
|
Prompt: prompt,
|
||||||
|
HistoryFile: filepath.Join(os.TempDir(), ".picoclaw_history"),
|
||||||
|
HistoryLimit: 100,
|
||||||
|
InterruptPrompt: "^C",
|
||||||
|
EOFPrompt: "exit",
|
||||||
|
})
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error initializing readline: %v\n", err)
|
||||||
|
fmt.Println("Falling back to simple input mode...")
|
||||||
|
simpleInteractiveMode(agentLoop, sessionKey)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer rl.Close()
|
||||||
|
|
||||||
|
for {
|
||||||
|
line, err := rl.Readline()
|
||||||
|
if err != nil {
|
||||||
|
if err == readline.ErrInterrupt || err == io.EOF {
|
||||||
|
fmt.Println("\nGoodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fmt.Printf("Error reading input: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
input := strings.TrimSpace(line)
|
||||||
|
if input == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if input == "exit" || input == "quit" {
|
||||||
|
fmt.Println("Goodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
response, err := agentLoop.ProcessDirect(ctx, input, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\n%s %s\n\n", logo, response)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func simpleInteractiveMode(agentLoop *agent.AgentLoop, sessionKey string) {
|
||||||
|
reader := bufio.NewReader(os.Stdin)
|
||||||
|
for {
|
||||||
|
fmt.Print(fmt.Sprintf("%s You: ", logo))
|
||||||
|
line, err := reader.ReadString('\n')
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
fmt.Println("\nGoodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fmt.Printf("Error reading input: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
input := strings.TrimSpace(line)
|
||||||
|
if input == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if input == "exit" || input == "quit" {
|
||||||
|
fmt.Println("Goodbye!")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
response, err := agentLoop.ProcessDirect(ctx, input, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\n%s %s\n\n", logo, response)
|
||||||
|
}
|
||||||
|
}
|
||||||
512
cmd/picoclaw/cmd_auth.go
Normal file
512
cmd/picoclaw/cmd_auth.go
Normal file
|
|
@ -0,0 +1,512 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const supportedProvidersMsg = "Supported providers: openai, anthropic, google-antigravity"
|
||||||
|
|
||||||
|
func authCmd() {
|
||||||
|
if len(os.Args) < 3 {
|
||||||
|
authHelp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch os.Args[2] {
|
||||||
|
case "login":
|
||||||
|
authLoginCmd()
|
||||||
|
case "logout":
|
||||||
|
authLogoutCmd()
|
||||||
|
case "status":
|
||||||
|
authStatusCmd()
|
||||||
|
case "models":
|
||||||
|
authModelsCmd()
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unknown auth command: %s\n", os.Args[2])
|
||||||
|
authHelp()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authHelp() {
|
||||||
|
fmt.Println("\nAuth commands:")
|
||||||
|
fmt.Println(" login Login via OAuth or paste token")
|
||||||
|
fmt.Println(" logout Remove stored credentials")
|
||||||
|
fmt.Println(" status Show current auth status")
|
||||||
|
fmt.Println(" models List available Antigravity models")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Login options:")
|
||||||
|
fmt.Println(" --provider <name> Provider to login with (openai, anthropic, google-antigravity)")
|
||||||
|
fmt.Println(" --device-code Use device code flow (for headless environments)")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Examples:")
|
||||||
|
fmt.Println(" picoclaw auth login --provider openai")
|
||||||
|
fmt.Println(" picoclaw auth login --provider openai --device-code")
|
||||||
|
fmt.Println(" picoclaw auth login --provider anthropic")
|
||||||
|
fmt.Println(" picoclaw auth login --provider google-antigravity")
|
||||||
|
fmt.Println(" picoclaw auth models")
|
||||||
|
fmt.Println(" picoclaw auth logout --provider openai")
|
||||||
|
fmt.Println(" picoclaw auth status")
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginCmd() {
|
||||||
|
provider := ""
|
||||||
|
useDeviceCode := false
|
||||||
|
|
||||||
|
args := os.Args[3:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--provider", "-p":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
provider = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--device-code":
|
||||||
|
useDeviceCode = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider == "" {
|
||||||
|
fmt.Println("Error: --provider is required")
|
||||||
|
fmt.Println(supportedProvidersMsg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch provider {
|
||||||
|
case "openai":
|
||||||
|
authLoginOpenAI(useDeviceCode)
|
||||||
|
case "anthropic":
|
||||||
|
authLoginPasteToken(provider)
|
||||||
|
case "google-antigravity", "antigravity":
|
||||||
|
authLoginGoogleAntigravity()
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unsupported provider: %s\n", provider)
|
||||||
|
fmt.Println(supportedProvidersMsg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginOpenAI(useDeviceCode bool) {
|
||||||
|
cfg := auth.OpenAIOAuthConfig()
|
||||||
|
|
||||||
|
var cred *auth.AuthCredential
|
||||||
|
var err error
|
||||||
|
|
||||||
|
if useDeviceCode {
|
||||||
|
cred, err = auth.LoginDeviceCode(cfg)
|
||||||
|
} else {
|
||||||
|
cred, err = auth.LoginBrowser(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Login failed: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential("openai", cred); err != nil {
|
||||||
|
fmt.Printf("Failed to save credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Update Providers (legacy format)
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = "oauth"
|
||||||
|
|
||||||
|
// Update or add openai in ModelList
|
||||||
|
foundOpenAI := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||||
|
foundOpenAI = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no openai in ModelList, add it
|
||||||
|
if !foundOpenAI {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update default model to use OpenAI
|
||||||
|
appCfg.Agents.Defaults.Model = "gpt-5.2"
|
||||||
|
|
||||||
|
if err := config.SaveConfig(getConfigPath(), appCfg); err != nil {
|
||||||
|
fmt.Printf("Warning: could not update config: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Login successful!")
|
||||||
|
if cred.AccountID != "" {
|
||||||
|
fmt.Printf("Account: %s\n", cred.AccountID)
|
||||||
|
}
|
||||||
|
fmt.Println("Default model set to: gpt-5.2")
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginGoogleAntigravity() {
|
||||||
|
cfg := auth.GoogleAntigravityOAuthConfig()
|
||||||
|
|
||||||
|
cred, err := auth.LoginBrowser(cfg)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Login failed: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
cred.Provider = "google-antigravity"
|
||||||
|
|
||||||
|
// Fetch user email from Google userinfo
|
||||||
|
email, err := fetchGoogleUserEmail(cred.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Warning: could not fetch email: %v\n", err)
|
||||||
|
} else {
|
||||||
|
cred.Email = email
|
||||||
|
fmt.Printf("Email: %s\n", email)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fetch Cloud Code Assist project ID
|
||||||
|
projectID, err := providers.FetchAntigravityProjectID(cred.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Warning: could not fetch project ID: %v\n", err)
|
||||||
|
fmt.Println("You may need Google Cloud Code Assist enabled on your account.")
|
||||||
|
} else {
|
||||||
|
cred.ProjectID = projectID
|
||||||
|
fmt.Printf("Project: %s\n", projectID)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential("google-antigravity", cred); err != nil {
|
||||||
|
fmt.Printf("Failed to save credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Update Providers (legacy format, for backward compatibility)
|
||||||
|
appCfg.Providers.Antigravity.AuthMethod = "oauth"
|
||||||
|
|
||||||
|
// Update or add antigravity in ModelList
|
||||||
|
foundAntigravity := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isAntigravityModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||||
|
foundAntigravity = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no antigravity in ModelList, add it
|
||||||
|
if !foundAntigravity {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gemini-flash",
|
||||||
|
Model: "antigravity/gemini-3-flash",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update default model
|
||||||
|
appCfg.Agents.Defaults.Model = "gemini-flash"
|
||||||
|
|
||||||
|
if err := config.SaveConfig(getConfigPath(), appCfg); err != nil {
|
||||||
|
fmt.Printf("Warning: could not update config: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\n✓ Google Antigravity login successful!")
|
||||||
|
fmt.Println("Default model set to: gemini-flash")
|
||||||
|
fmt.Println("Try it: picoclaw agent -m \"Hello world\"")
|
||||||
|
}
|
||||||
|
|
||||||
|
func fetchGoogleUserEmail(accessToken string) (string, error) {
|
||||||
|
req, err := http.NewRequest("GET", "https://www.googleapis.com/oauth2/v2/userinfo", nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 10 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("userinfo request failed: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
var userInfo struct {
|
||||||
|
Email string `json:"email"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &userInfo); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return userInfo.Email, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLoginPasteToken(provider string) {
|
||||||
|
cred, err := auth.LoginPasteToken(provider, os.Stdin)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Login failed: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := auth.SetCredential(provider, cred); err != nil {
|
||||||
|
fmt.Printf("Failed to save credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
switch provider {
|
||||||
|
case "anthropic":
|
||||||
|
appCfg.Providers.Anthropic.AuthMethod = "token"
|
||||||
|
// Update ModelList
|
||||||
|
found := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "token"
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "claude-sonnet-4.6",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
AuthMethod: "token",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
// Update default model
|
||||||
|
appCfg.Agents.Defaults.Model = "claude-sonnet-4.6"
|
||||||
|
case "openai":
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = "token"
|
||||||
|
// Update ModelList
|
||||||
|
found := false
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = "token"
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
appCfg.ModelList = append(appCfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "token",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
// Update default model
|
||||||
|
appCfg.Agents.Defaults.Model = "gpt-5.2"
|
||||||
|
}
|
||||||
|
if err := config.SaveConfig(getConfigPath(), appCfg); err != nil {
|
||||||
|
fmt.Printf("Warning: could not update config: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Token saved for %s!\n", provider)
|
||||||
|
fmt.Printf("Default model set to: %s\n", appCfg.Agents.Defaults.Model)
|
||||||
|
}
|
||||||
|
|
||||||
|
func authLogoutCmd() {
|
||||||
|
provider := ""
|
||||||
|
|
||||||
|
args := os.Args[3:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--provider", "-p":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
provider = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider != "" {
|
||||||
|
if err := auth.DeleteCredential(provider); err != nil {
|
||||||
|
fmt.Printf("Failed to remove credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Clear AuthMethod in ModelList
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
switch provider {
|
||||||
|
case "openai":
|
||||||
|
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
case "anthropic":
|
||||||
|
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
case "google-antigravity", "antigravity":
|
||||||
|
if isAntigravityModel(appCfg.ModelList[i].Model) {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Clear AuthMethod in Providers (legacy)
|
||||||
|
switch provider {
|
||||||
|
case "openai":
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = ""
|
||||||
|
case "anthropic":
|
||||||
|
appCfg.Providers.Anthropic.AuthMethod = ""
|
||||||
|
case "google-antigravity", "antigravity":
|
||||||
|
appCfg.Providers.Antigravity.AuthMethod = ""
|
||||||
|
}
|
||||||
|
config.SaveConfig(getConfigPath(), appCfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Logged out from %s\n", provider)
|
||||||
|
} else {
|
||||||
|
if err := auth.DeleteAllCredentials(); err != nil {
|
||||||
|
fmt.Printf("Failed to remove credentials: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
appCfg, err := loadConfig()
|
||||||
|
if err == nil {
|
||||||
|
// Clear all AuthMethods in ModelList
|
||||||
|
for i := range appCfg.ModelList {
|
||||||
|
appCfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
// Clear all AuthMethods in Providers (legacy)
|
||||||
|
appCfg.Providers.OpenAI.AuthMethod = ""
|
||||||
|
appCfg.Providers.Anthropic.AuthMethod = ""
|
||||||
|
appCfg.Providers.Antigravity.AuthMethod = ""
|
||||||
|
config.SaveConfig(getConfigPath(), appCfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Logged out from all providers")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authStatusCmd() {
|
||||||
|
store, err := auth.LoadStore()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading auth store: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(store.Credentials) == 0 {
|
||||||
|
fmt.Println("No authenticated providers.")
|
||||||
|
fmt.Println("Run: picoclaw auth login --provider <name>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\nAuthenticated Providers:")
|
||||||
|
fmt.Println("------------------------")
|
||||||
|
for provider, cred := range store.Credentials {
|
||||||
|
status := "active"
|
||||||
|
if cred.IsExpired() {
|
||||||
|
status = "expired"
|
||||||
|
} else if cred.NeedsRefresh() {
|
||||||
|
status = "needs refresh"
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf(" %s:\n", provider)
|
||||||
|
fmt.Printf(" Method: %s\n", cred.AuthMethod)
|
||||||
|
fmt.Printf(" Status: %s\n", status)
|
||||||
|
if cred.AccountID != "" {
|
||||||
|
fmt.Printf(" Account: %s\n", cred.AccountID)
|
||||||
|
}
|
||||||
|
if cred.Email != "" {
|
||||||
|
fmt.Printf(" Email: %s\n", cred.Email)
|
||||||
|
}
|
||||||
|
if cred.ProjectID != "" {
|
||||||
|
fmt.Printf(" Project: %s\n", cred.ProjectID)
|
||||||
|
}
|
||||||
|
if !cred.ExpiresAt.IsZero() {
|
||||||
|
fmt.Printf(" Expires: %s\n", cred.ExpiresAt.Format("2006-01-02 15:04"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func authModelsCmd() {
|
||||||
|
cred, err := auth.GetCredential("google-antigravity")
|
||||||
|
if err != nil || cred == nil {
|
||||||
|
fmt.Println("Not logged in to Google Antigravity.")
|
||||||
|
fmt.Println("Run: picoclaw auth login --provider google-antigravity")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh token if needed
|
||||||
|
if cred.NeedsRefresh() && cred.RefreshToken != "" {
|
||||||
|
oauthCfg := auth.GoogleAntigravityOAuthConfig()
|
||||||
|
refreshed, refreshErr := auth.RefreshAccessToken(cred, oauthCfg)
|
||||||
|
if refreshErr == nil {
|
||||||
|
cred = refreshed
|
||||||
|
_ = auth.SetCredential("google-antigravity", cred)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
projectID := cred.ProjectID
|
||||||
|
if projectID == "" {
|
||||||
|
fmt.Println("No project ID stored. Try logging in again.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Fetching models for project: %s\n\n", projectID)
|
||||||
|
|
||||||
|
models, err := providers.FetchAntigravityModels(cred.AccessToken, projectID)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error fetching models: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(models) == 0 {
|
||||||
|
fmt.Println("No models available.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Available Antigravity Models:")
|
||||||
|
fmt.Println("-----------------------------")
|
||||||
|
for _, m := range models {
|
||||||
|
status := "✓"
|
||||||
|
if m.IsExhausted {
|
||||||
|
status = "✗ (quota exhausted)"
|
||||||
|
}
|
||||||
|
name := m.ID
|
||||||
|
if m.DisplayName != "" {
|
||||||
|
name = fmt.Sprintf("%s (%s)", m.ID, m.DisplayName)
|
||||||
|
}
|
||||||
|
fmt.Printf(" %s %s\n", status, name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isAntigravityModel checks if a model string belongs to antigravity provider
|
||||||
|
func isAntigravityModel(model string) bool {
|
||||||
|
return model == "antigravity" ||
|
||||||
|
model == "google-antigravity" ||
|
||||||
|
strings.HasPrefix(model, "antigravity/") ||
|
||||||
|
strings.HasPrefix(model, "google-antigravity/")
|
||||||
|
}
|
||||||
|
|
||||||
|
// isOpenAIModel checks if a model string belongs to openai provider
|
||||||
|
func isOpenAIModel(model string) bool {
|
||||||
|
return model == "openai" ||
|
||||||
|
strings.HasPrefix(model, "openai/")
|
||||||
|
}
|
||||||
|
|
||||||
|
// isAnthropicModel checks if a model string belongs to anthropic provider
|
||||||
|
func isAnthropicModel(model string) bool {
|
||||||
|
return model == "anthropic" ||
|
||||||
|
strings.HasPrefix(model, "anthropic/")
|
||||||
|
}
|
||||||
227
cmd/picoclaw/cmd_cron.go
Normal file
227
cmd/picoclaw/cmd_cron.go
Normal file
|
|
@ -0,0 +1,227 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
|
)
|
||||||
|
|
||||||
|
func cronCmd() {
|
||||||
|
if len(os.Args) < 3 {
|
||||||
|
cronHelp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
subcommand := os.Args[2]
|
||||||
|
|
||||||
|
// Load config to get workspace path
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cronStorePath := filepath.Join(cfg.WorkspacePath(), "cron", "jobs.json")
|
||||||
|
|
||||||
|
switch subcommand {
|
||||||
|
case "list":
|
||||||
|
cronListCmd(cronStorePath)
|
||||||
|
case "add":
|
||||||
|
cronAddCmd(cronStorePath)
|
||||||
|
case "remove":
|
||||||
|
if len(os.Args) < 4 {
|
||||||
|
fmt.Println("Usage: picoclaw cron remove <job_id>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cronRemoveCmd(cronStorePath, os.Args[3])
|
||||||
|
case "enable":
|
||||||
|
cronEnableCmd(cronStorePath, false)
|
||||||
|
case "disable":
|
||||||
|
cronEnableCmd(cronStorePath, true)
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unknown cron command: %s\n", subcommand)
|
||||||
|
cronHelp()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronHelp() {
|
||||||
|
fmt.Println("\nCron commands:")
|
||||||
|
fmt.Println(" list List all scheduled jobs")
|
||||||
|
fmt.Println(" add Add a new scheduled job")
|
||||||
|
fmt.Println(" remove <id> Remove a job by ID")
|
||||||
|
fmt.Println(" enable <id> Enable a job")
|
||||||
|
fmt.Println(" disable <id> Disable a job")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Add options:")
|
||||||
|
fmt.Println(" -n, --name Job name")
|
||||||
|
fmt.Println(" -m, --message Message for agent")
|
||||||
|
fmt.Println(" -e, --every Run every N seconds")
|
||||||
|
fmt.Println(" -c, --cron Cron expression (e.g. '0 9 * * *')")
|
||||||
|
fmt.Println(" -d, --deliver Deliver response to channel")
|
||||||
|
fmt.Println(" --to Recipient for delivery")
|
||||||
|
fmt.Println(" --channel Channel for delivery")
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronListCmd(storePath string) {
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
jobs := cs.ListJobs(true) // Show all jobs, including disabled
|
||||||
|
|
||||||
|
if len(jobs) == 0 {
|
||||||
|
fmt.Println("No scheduled jobs.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\nScheduled Jobs:")
|
||||||
|
fmt.Println("----------------")
|
||||||
|
for _, job := range jobs {
|
||||||
|
var schedule string
|
||||||
|
if job.Schedule.Kind == "every" && job.Schedule.EveryMS != nil {
|
||||||
|
schedule = fmt.Sprintf("every %ds", *job.Schedule.EveryMS/1000)
|
||||||
|
} else if job.Schedule.Kind == "cron" {
|
||||||
|
schedule = job.Schedule.Expr
|
||||||
|
} else {
|
||||||
|
schedule = "one-time"
|
||||||
|
}
|
||||||
|
|
||||||
|
nextRun := "scheduled"
|
||||||
|
if job.State.NextRunAtMS != nil {
|
||||||
|
nextTime := time.UnixMilli(*job.State.NextRunAtMS)
|
||||||
|
nextRun = nextTime.Format("2006-01-02 15:04")
|
||||||
|
}
|
||||||
|
|
||||||
|
status := "enabled"
|
||||||
|
if !job.Enabled {
|
||||||
|
status = "disabled"
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf(" %s (%s)\n", job.Name, job.ID)
|
||||||
|
fmt.Printf(" Schedule: %s\n", schedule)
|
||||||
|
fmt.Printf(" Status: %s\n", status)
|
||||||
|
fmt.Printf(" Next run: %s\n", nextRun)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronAddCmd(storePath string) {
|
||||||
|
name := ""
|
||||||
|
message := ""
|
||||||
|
var everySec *int64
|
||||||
|
cronExpr := ""
|
||||||
|
deliver := false
|
||||||
|
channel := ""
|
||||||
|
to := ""
|
||||||
|
|
||||||
|
args := os.Args[3:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "-n", "--name":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
name = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-m", "--message":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
message = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-e", "--every":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
var sec int64
|
||||||
|
fmt.Sscanf(args[i+1], "%d", &sec)
|
||||||
|
everySec = &sec
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-c", "--cron":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
cronExpr = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "-d", "--deliver":
|
||||||
|
deliver = true
|
||||||
|
case "--to":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
to = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--channel":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
channel = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if name == "" {
|
||||||
|
fmt.Println("Error: --name is required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if message == "" {
|
||||||
|
fmt.Println("Error: --message is required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if everySec == nil && cronExpr == "" {
|
||||||
|
fmt.Println("Error: Either --every or --cron must be specified")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var schedule cron.CronSchedule
|
||||||
|
if everySec != nil {
|
||||||
|
everyMS := *everySec * 1000
|
||||||
|
schedule = cron.CronSchedule{
|
||||||
|
Kind: "every",
|
||||||
|
EveryMS: &everyMS,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
schedule = cron.CronSchedule{
|
||||||
|
Kind: "cron",
|
||||||
|
Expr: cronExpr,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
job, err := cs.AddJob(name, schedule, message, deliver, channel, to)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error adding job: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("✓ Added job '%s' (%s)\n", job.Name, job.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronRemoveCmd(storePath, jobID string) {
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
if cs.RemoveJob(jobID) {
|
||||||
|
fmt.Printf("✓ Removed job %s\n", jobID)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("✗ Job %s not found\n", jobID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cronEnableCmd(storePath string, disable bool) {
|
||||||
|
if len(os.Args) < 4 {
|
||||||
|
fmt.Println("Usage: picoclaw cron enable/disable <job_id>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
jobID := os.Args[3]
|
||||||
|
cs := cron.NewCronService(storePath, nil)
|
||||||
|
enabled := !disable
|
||||||
|
|
||||||
|
job := cs.EnableJob(jobID, enabled)
|
||||||
|
if job != nil {
|
||||||
|
status := "enabled"
|
||||||
|
if disable {
|
||||||
|
status = "disabled"
|
||||||
|
}
|
||||||
|
fmt.Printf("✓ Job '%s' %s\n", job.Name, status)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("✗ Job %s not found\n", jobID)
|
||||||
|
}
|
||||||
|
}
|
||||||
223
cmd/picoclaw/cmd_gateway.go
Normal file
223
cmd/picoclaw/cmd_gateway.go
Normal file
|
|
@ -0,0 +1,223 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"os/signal"
|
||||||
|
"path/filepath"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/agent"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/devices"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/health"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/heartbeat"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/state"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/voice"
|
||||||
|
)
|
||||||
|
|
||||||
|
func gatewayCmd() {
|
||||||
|
// Check for --debug flag
|
||||||
|
args := os.Args[2:]
|
||||||
|
for _, arg := range args {
|
||||||
|
if arg == "--debug" || arg == "-d" {
|
||||||
|
logger.SetLevel(logger.DEBUG)
|
||||||
|
fmt.Println("🔍 Debug mode enabled")
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := providers.CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error creating provider: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
// Use the resolved model ID from provider creation
|
||||||
|
if modelID != "" {
|
||||||
|
cfg.Agents.Defaults.Model = modelID
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
// Print agent startup info
|
||||||
|
fmt.Println("\n📦 Agent Status:")
|
||||||
|
startupInfo := agentLoop.GetStartupInfo()
|
||||||
|
toolsInfo := startupInfo["tools"].(map[string]interface{})
|
||||||
|
skillsInfo := startupInfo["skills"].(map[string]interface{})
|
||||||
|
fmt.Printf(" • Tools: %d loaded\n", toolsInfo["count"])
|
||||||
|
fmt.Printf(" • Skills: %d/%d available\n",
|
||||||
|
skillsInfo["available"],
|
||||||
|
skillsInfo["total"])
|
||||||
|
|
||||||
|
// Log to file as well
|
||||||
|
logger.InfoCF("agent", "Agent initialized",
|
||||||
|
map[string]interface{}{
|
||||||
|
"tools_count": toolsInfo["count"],
|
||||||
|
"skills_total": skillsInfo["total"],
|
||||||
|
"skills_available": skillsInfo["available"],
|
||||||
|
})
|
||||||
|
|
||||||
|
// Setup cron tool and service
|
||||||
|
execTimeout := time.Duration(cfg.Tools.Cron.ExecTimeoutMinutes) * time.Minute
|
||||||
|
cronService := setupCronTool(agentLoop, msgBus, cfg.WorkspacePath(), cfg.Agents.Defaults.RestrictToWorkspace, execTimeout, cfg)
|
||||||
|
|
||||||
|
heartbeatService := heartbeat.NewHeartbeatService(
|
||||||
|
cfg.WorkspacePath(),
|
||||||
|
cfg.Heartbeat.Interval,
|
||||||
|
cfg.Heartbeat.Enabled,
|
||||||
|
)
|
||||||
|
heartbeatService.SetBus(msgBus)
|
||||||
|
heartbeatService.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult {
|
||||||
|
// Use cli:direct as fallback if no valid channel
|
||||||
|
if channel == "" || chatID == "" {
|
||||||
|
channel, chatID = "cli", "direct"
|
||||||
|
}
|
||||||
|
// Use ProcessHeartbeat - no session history, each heartbeat is independent
|
||||||
|
response, err := agentLoop.ProcessHeartbeat(context.Background(), prompt, channel, chatID)
|
||||||
|
if err != nil {
|
||||||
|
return tools.ErrorResult(fmt.Sprintf("Heartbeat error: %v", err))
|
||||||
|
}
|
||||||
|
if response == "HEARTBEAT_OK" {
|
||||||
|
return tools.SilentResult("Heartbeat OK")
|
||||||
|
}
|
||||||
|
// For heartbeat, always return silent - the subagent result will be
|
||||||
|
// sent to user via processSystemMessage when the async task completes
|
||||||
|
return tools.SilentResult(response)
|
||||||
|
})
|
||||||
|
|
||||||
|
channelManager, err := channels.NewManager(cfg, msgBus)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error creating channel manager: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Inject channel manager into agent loop for command handling
|
||||||
|
agentLoop.SetChannelManager(channelManager)
|
||||||
|
|
||||||
|
var transcriber *voice.GroqTranscriber
|
||||||
|
if cfg.Providers.Groq.APIKey != "" {
|
||||||
|
transcriber = voice.NewGroqTranscriber(cfg.Providers.Groq.APIKey)
|
||||||
|
logger.InfoC("voice", "Groq voice transcription enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
if transcriber != nil {
|
||||||
|
if telegramChannel, ok := channelManager.GetChannel("telegram"); ok {
|
||||||
|
if tc, ok := telegramChannel.(*channels.TelegramChannel); ok {
|
||||||
|
tc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "Groq transcription attached to Telegram channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if discordChannel, ok := channelManager.GetChannel("discord"); ok {
|
||||||
|
if dc, ok := discordChannel.(*channels.DiscordChannel); ok {
|
||||||
|
dc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "Groq transcription attached to Discord channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if slackChannel, ok := channelManager.GetChannel("slack"); ok {
|
||||||
|
if sc, ok := slackChannel.(*channels.SlackChannel); ok {
|
||||||
|
sc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "Groq transcription attached to Slack channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
enabledChannels := channelManager.GetEnabledChannels()
|
||||||
|
if len(enabledChannels) > 0 {
|
||||||
|
fmt.Printf("✓ Channels enabled: %s\n", enabledChannels)
|
||||||
|
} else {
|
||||||
|
fmt.Println("⚠ Warning: No channels enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("✓ Gateway started on %s:%d\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
fmt.Println("Press Ctrl+C to stop")
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := cronService.Start(); err != nil {
|
||||||
|
fmt.Printf("Error starting cron service: %v\n", err)
|
||||||
|
}
|
||||||
|
fmt.Println("✓ Cron service started")
|
||||||
|
|
||||||
|
if err := heartbeatService.Start(); err != nil {
|
||||||
|
fmt.Printf("Error starting heartbeat service: %v\n", err)
|
||||||
|
}
|
||||||
|
fmt.Println("✓ Heartbeat service started")
|
||||||
|
|
||||||
|
stateManager := state.NewManager(cfg.WorkspacePath())
|
||||||
|
deviceService := devices.NewService(devices.Config{
|
||||||
|
Enabled: cfg.Devices.Enabled,
|
||||||
|
MonitorUSB: cfg.Devices.MonitorUSB,
|
||||||
|
}, stateManager)
|
||||||
|
deviceService.SetBus(msgBus)
|
||||||
|
if err := deviceService.Start(ctx); err != nil {
|
||||||
|
fmt.Printf("Error starting device service: %v\n", err)
|
||||||
|
} else if cfg.Devices.Enabled {
|
||||||
|
fmt.Println("✓ Device event service started")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := channelManager.StartAll(ctx); err != nil {
|
||||||
|
fmt.Printf("Error starting channels: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
healthServer := health.NewServer(cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
go func() {
|
||||||
|
if err := healthServer.Start(); err != nil && err != http.ErrServerClosed {
|
||||||
|
logger.ErrorCF("health", "Health server error", map[string]interface{}{"error": err.Error()})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
fmt.Printf("✓ Health endpoints available at http://%s:%d/health and /ready\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
|
||||||
|
go agentLoop.Run(ctx)
|
||||||
|
|
||||||
|
sigChan := make(chan os.Signal, 1)
|
||||||
|
signal.Notify(sigChan, os.Interrupt)
|
||||||
|
<-sigChan
|
||||||
|
|
||||||
|
fmt.Println("\nShutting down...")
|
||||||
|
cancel()
|
||||||
|
healthServer.Stop(context.Background())
|
||||||
|
deviceService.Stop()
|
||||||
|
heartbeatService.Stop()
|
||||||
|
cronService.Stop()
|
||||||
|
agentLoop.Stop()
|
||||||
|
channelManager.StopAll(ctx)
|
||||||
|
fmt.Println("✓ Gateway stopped")
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupCronTool(agentLoop *agent.AgentLoop, msgBus *bus.MessageBus, workspace string, restrict bool, execTimeout time.Duration, cfg *config.Config) *cron.CronService {
|
||||||
|
cronStorePath := filepath.Join(workspace, "cron", "jobs.json")
|
||||||
|
|
||||||
|
// Create cron service
|
||||||
|
cronService := cron.NewCronService(cronStorePath, nil)
|
||||||
|
|
||||||
|
// Create and register CronTool
|
||||||
|
cronTool := tools.NewCronTool(cronService, agentLoop, msgBus, workspace, restrict, execTimeout, cfg)
|
||||||
|
agentLoop.RegisterTool(cronTool)
|
||||||
|
|
||||||
|
// Set the onJob handler
|
||||||
|
cronService.SetOnJob(func(job *cron.CronJob) (string, error) {
|
||||||
|
result := cronTool.ExecuteJob(context.Background(), job)
|
||||||
|
return result, nil
|
||||||
|
})
|
||||||
|
|
||||||
|
return cronService
|
||||||
|
}
|
||||||
81
cmd/picoclaw/cmd_migrate.go
Normal file
81
cmd/picoclaw/cmd_migrate.go
Normal file
|
|
@ -0,0 +1,81 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/migrate"
|
||||||
|
)
|
||||||
|
|
||||||
|
func migrateCmd() {
|
||||||
|
if len(os.Args) > 2 && (os.Args[2] == "--help" || os.Args[2] == "-h") {
|
||||||
|
migrateHelp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := migrate.Options{}
|
||||||
|
|
||||||
|
args := os.Args[2:]
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--dry-run":
|
||||||
|
opts.DryRun = true
|
||||||
|
case "--config-only":
|
||||||
|
opts.ConfigOnly = true
|
||||||
|
case "--workspace-only":
|
||||||
|
opts.WorkspaceOnly = true
|
||||||
|
case "--force":
|
||||||
|
opts.Force = true
|
||||||
|
case "--refresh":
|
||||||
|
opts.Refresh = true
|
||||||
|
case "--openclaw-home":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
opts.OpenClawHome = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
case "--picoclaw-home":
|
||||||
|
if i+1 < len(args) {
|
||||||
|
opts.PicoClawHome = args[i+1]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
fmt.Printf("Unknown flag: %s\n", args[i])
|
||||||
|
migrateHelp()
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := migrate.Run(opts)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !opts.DryRun {
|
||||||
|
migrate.PrintSummary(result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func migrateHelp() {
|
||||||
|
fmt.Println("\nMigrate from OpenClaw to PicoClaw")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Usage: picoclaw migrate [options]")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Options:")
|
||||||
|
fmt.Println(" --dry-run Show what would be migrated without making changes")
|
||||||
|
fmt.Println(" --refresh Re-sync workspace files from OpenClaw (repeatable)")
|
||||||
|
fmt.Println(" --config-only Only migrate config, skip workspace files")
|
||||||
|
fmt.Println(" --workspace-only Only migrate workspace files, skip config")
|
||||||
|
fmt.Println(" --force Skip confirmation prompts")
|
||||||
|
fmt.Println(" --openclaw-home Override OpenClaw home directory (default: ~/.openclaw)")
|
||||||
|
fmt.Println(" --picoclaw-home Override PicoClaw home directory (default: ~/.picoclaw)")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Examples:")
|
||||||
|
fmt.Println(" picoclaw migrate Detect and migrate from OpenClaw")
|
||||||
|
fmt.Println(" picoclaw migrate --dry-run Show what would be migrated")
|
||||||
|
fmt.Println(" picoclaw migrate --refresh Re-sync workspace files")
|
||||||
|
fmt.Println(" picoclaw migrate --force Migrate without confirmation")
|
||||||
|
}
|
||||||
108
cmd/picoclaw/cmd_onboard.go
Normal file
108
cmd/picoclaw/cmd_onboard.go
Normal file
|
|
@ -0,0 +1,108 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"embed"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
//go:generate cp -r ../../workspace .
|
||||||
|
//go:embed workspace
|
||||||
|
var embeddedFiles embed.FS
|
||||||
|
|
||||||
|
func onboard() {
|
||||||
|
configPath := getConfigPath()
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err == nil {
|
||||||
|
fmt.Printf("Config already exists at %s\n", configPath)
|
||||||
|
fmt.Print("Overwrite? (y/n): ")
|
||||||
|
var response string
|
||||||
|
fmt.Scanln(&response)
|
||||||
|
if response != "y" {
|
||||||
|
fmt.Println("Aborted.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
fmt.Printf("Error saving config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
createWorkspaceTemplates(workspace)
|
||||||
|
|
||||||
|
fmt.Printf("%s picoclaw is ready!\n", logo)
|
||||||
|
fmt.Println("\nNext steps:")
|
||||||
|
fmt.Println(" 1. Add your API key to", configPath)
|
||||||
|
fmt.Println("")
|
||||||
|
fmt.Println(" Recommended:")
|
||||||
|
fmt.Println(" - OpenRouter: https://openrouter.ai/keys (access 100+ models)")
|
||||||
|
fmt.Println(" - Ollama: https://ollama.com (local, free)")
|
||||||
|
fmt.Println("")
|
||||||
|
fmt.Println(" See README.md for 17+ supported providers.")
|
||||||
|
fmt.Println("")
|
||||||
|
fmt.Println(" 2. Chat: picoclaw agent -m \"Hello!\"")
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyEmbeddedToTarget(targetDir string) error {
|
||||||
|
// Ensure target directory exists
|
||||||
|
if err := os.MkdirAll(targetDir, 0755); err != nil {
|
||||||
|
return fmt.Errorf("Failed to create target directory: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Walk through all files in embed.FS
|
||||||
|
err := fs.WalkDir(embeddedFiles, "workspace", func(path string, d fs.DirEntry, err error) error {
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip directories
|
||||||
|
if d.IsDir() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read embedded file
|
||||||
|
data, err := embeddedFiles.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("Failed to read embedded file %s: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
new_path, err := filepath.Rel("workspace", path)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("Failed to get relative path for %s: %v\n", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build target file path
|
||||||
|
targetPath := filepath.Join(targetDir, new_path)
|
||||||
|
|
||||||
|
// Ensure target file's directory exists
|
||||||
|
if err := os.MkdirAll(filepath.Dir(targetPath), 0755); err != nil {
|
||||||
|
return fmt.Errorf("Failed to create directory %s: %w", filepath.Dir(targetPath), err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write file
|
||||||
|
if err := os.WriteFile(targetPath, data, 0644); err != nil {
|
||||||
|
return fmt.Errorf("Failed to write file %s: %w", targetPath, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func createWorkspaceTemplates(workspace string) {
|
||||||
|
err := copyEmbeddedToTarget(workspace)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error copying workspace templates: %v\n", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
305
cmd/picoclaw/cmd_skills.go
Normal file
305
cmd/picoclaw/cmd_skills.go
Normal file
|
|
@ -0,0 +1,305 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
func skillsHelp() {
|
||||||
|
fmt.Println("\nSkills commands:")
|
||||||
|
fmt.Println(" list List installed skills")
|
||||||
|
fmt.Println(" install <repo> Install skill from GitHub")
|
||||||
|
fmt.Println(" install-builtin Install all builtin skills to workspace")
|
||||||
|
fmt.Println(" list-builtin List available builtin skills")
|
||||||
|
fmt.Println(" remove <name> Remove installed skill")
|
||||||
|
fmt.Println(" search Search available skills")
|
||||||
|
fmt.Println(" show <name> Show skill details")
|
||||||
|
fmt.Println()
|
||||||
|
fmt.Println("Examples:")
|
||||||
|
fmt.Println(" picoclaw skills list")
|
||||||
|
fmt.Println(" picoclaw skills install sipeed/picoclaw-skills/weather")
|
||||||
|
fmt.Println(" picoclaw skills install-builtin")
|
||||||
|
fmt.Println(" picoclaw skills list-builtin")
|
||||||
|
fmt.Println(" picoclaw skills remove weather")
|
||||||
|
fmt.Println(" picoclaw skills install --registry clawhub github")
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsListCmd(loader *skills.SkillsLoader) {
|
||||||
|
allSkills := loader.ListSkills()
|
||||||
|
|
||||||
|
if len(allSkills) == 0 {
|
||||||
|
fmt.Println("No skills installed.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\nInstalled Skills:")
|
||||||
|
fmt.Println("------------------")
|
||||||
|
for _, skill := range allSkills {
|
||||||
|
fmt.Printf(" ✓ %s (%s)\n", skill.Name, skill.Source)
|
||||||
|
if skill.Description != "" {
|
||||||
|
fmt.Printf(" %s\n", skill.Description)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsInstallCmd(installer *skills.SkillInstaller, cfg *config.Config) {
|
||||||
|
if len(os.Args) < 4 {
|
||||||
|
fmt.Println("Usage: picoclaw skills install <github-repo>")
|
||||||
|
fmt.Println(" picoclaw skills install --registry <name> <slug>")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for --registry flag.
|
||||||
|
if os.Args[3] == "--registry" {
|
||||||
|
if len(os.Args) < 6 {
|
||||||
|
fmt.Println("Usage: picoclaw skills install --registry <name> <slug>")
|
||||||
|
fmt.Println("Example: picoclaw skills install --registry clawhub github")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
registryName := os.Args[4]
|
||||||
|
slug := os.Args[5]
|
||||||
|
skillsInstallFromRegistry(cfg, registryName, slug)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default: install from GitHub (backward compatible).
|
||||||
|
repo := os.Args[3]
|
||||||
|
fmt.Printf("Installing skill from %s...\n", repo)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := installer.InstallFromGitHub(ctx, repo); err != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to install skill: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\u2713 Skill '%s' installed successfully!\n", filepath.Base(repo))
|
||||||
|
}
|
||||||
|
|
||||||
|
// skillsInstallFromRegistry installs a skill from a named registry (e.g. clawhub).
|
||||||
|
func skillsInstallFromRegistry(cfg *config.Config, registryName, slug string) {
|
||||||
|
err := utils.ValidateSkillIdentifier(registryName)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("\u2717 Invalid registry name: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = utils.ValidateSkillIdentifier(slug)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("\u2717 Invalid slug: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Installing skill '%s' from %s registry...\n", slug, registryName)
|
||||||
|
|
||||||
|
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
|
||||||
|
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
|
||||||
|
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
|
||||||
|
})
|
||||||
|
|
||||||
|
registry := registryMgr.GetRegistry(registryName)
|
||||||
|
if registry == nil {
|
||||||
|
fmt.Printf("\u2717 Registry '%s' not found or not enabled. Check your config.json.\n", registryName)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
targetDir := filepath.Join(workspace, "skills", slug)
|
||||||
|
|
||||||
|
if _, err := os.Stat(targetDir); err == nil {
|
||||||
|
fmt.Printf("\u2717 Skill '%s' already installed at %s\n", slug, targetDir)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := os.MkdirAll(filepath.Join(workspace, "skills"), 0755); err != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to create skills directory: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := registry.DownloadAndInstall(ctx, slug, "", targetDir)
|
||||||
|
if err != nil {
|
||||||
|
rmErr := os.RemoveAll(targetDir)
|
||||||
|
if rmErr != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to remove partial install: %v\n", rmErr)
|
||||||
|
}
|
||||||
|
fmt.Printf("\u2717 Failed to install skill: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.IsMalwareBlocked {
|
||||||
|
rmErr := os.RemoveAll(targetDir)
|
||||||
|
if rmErr != nil {
|
||||||
|
fmt.Printf("\u2717 Failed to remove partial install: %v\n", rmErr)
|
||||||
|
}
|
||||||
|
fmt.Printf("\u2717 Skill '%s' is flagged as malicious and cannot be installed.\n", slug)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.IsSuspicious {
|
||||||
|
fmt.Printf("\u26a0\ufe0f Warning: skill '%s' is flagged as suspicious.\n", slug)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\u2713 Skill '%s' v%s installed successfully!\n", slug, result.Version)
|
||||||
|
if result.Summary != "" {
|
||||||
|
fmt.Printf(" %s\n", result.Summary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsRemoveCmd(installer *skills.SkillInstaller, skillName string) {
|
||||||
|
fmt.Printf("Removing skill '%s'...\n", skillName)
|
||||||
|
|
||||||
|
if err := installer.Uninstall(skillName); err != nil {
|
||||||
|
fmt.Printf("✗ Failed to remove skill: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("✓ Skill '%s' removed successfully!\n", skillName)
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsInstallBuiltinCmd(workspace string) {
|
||||||
|
builtinSkillsDir := "./picoclaw/skills"
|
||||||
|
workspaceSkillsDir := filepath.Join(workspace, "skills")
|
||||||
|
|
||||||
|
fmt.Printf("Copying builtin skills to workspace...\n")
|
||||||
|
|
||||||
|
skillsToInstall := []string{
|
||||||
|
"weather",
|
||||||
|
"news",
|
||||||
|
"stock",
|
||||||
|
"calculator",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, skillName := range skillsToInstall {
|
||||||
|
builtinPath := filepath.Join(builtinSkillsDir, skillName)
|
||||||
|
workspacePath := filepath.Join(workspaceSkillsDir, skillName)
|
||||||
|
|
||||||
|
if _, err := os.Stat(builtinPath); err != nil {
|
||||||
|
fmt.Printf("⊘ Builtin skill '%s' not found: %v\n", skillName, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.MkdirAll(workspacePath, 0755); err != nil {
|
||||||
|
fmt.Printf("✗ Failed to create directory for %s: %v\n", skillName, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := copyDirectory(builtinPath, workspacePath); err != nil {
|
||||||
|
fmt.Printf("✗ Failed to copy %s: %v\n", skillName, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("\n✓ All builtin skills installed!")
|
||||||
|
fmt.Println("Now you can use them in your workspace.")
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsListBuiltinCmd() {
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
builtinSkillsDir := filepath.Join(filepath.Dir(cfg.WorkspacePath()), "picoclaw", "skills")
|
||||||
|
|
||||||
|
fmt.Println("\nAvailable Builtin Skills:")
|
||||||
|
fmt.Println("-----------------------")
|
||||||
|
|
||||||
|
entries, err := os.ReadDir(builtinSkillsDir)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error reading builtin skills: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(entries) == 0 {
|
||||||
|
fmt.Println("No builtin skills available.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, entry := range entries {
|
||||||
|
if entry.IsDir() {
|
||||||
|
skillName := entry.Name()
|
||||||
|
skillFile := filepath.Join(builtinSkillsDir, skillName, "SKILL.md")
|
||||||
|
|
||||||
|
description := "No description"
|
||||||
|
if _, err := os.Stat(skillFile); err == nil {
|
||||||
|
data, err := os.ReadFile(skillFile)
|
||||||
|
if err == nil {
|
||||||
|
content := string(data)
|
||||||
|
if idx := strings.Index(content, "\n"); idx > 0 {
|
||||||
|
firstLine := content[:idx]
|
||||||
|
if strings.Contains(firstLine, "description:") {
|
||||||
|
descLine := strings.Index(content[idx:], "\n")
|
||||||
|
if descLine > 0 {
|
||||||
|
description = strings.TrimSpace(content[idx+descLine : idx+descLine])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
status := "✓"
|
||||||
|
fmt.Printf(" %s %s\n", status, entry.Name())
|
||||||
|
if description != "" {
|
||||||
|
fmt.Printf(" %s\n", description)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsSearchCmd(installer *skills.SkillInstaller) {
|
||||||
|
fmt.Println("Searching for available skills...")
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
availableSkills, err := installer.ListAvailableSkills(ctx)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("✗ Failed to fetch skills list: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(availableSkills) == 0 {
|
||||||
|
fmt.Println("No skills available.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\nAvailable Skills (%d):\n", len(availableSkills))
|
||||||
|
fmt.Println("--------------------")
|
||||||
|
for _, skill := range availableSkills {
|
||||||
|
fmt.Printf(" 📦 %s\n", skill.Name)
|
||||||
|
fmt.Printf(" %s\n", skill.Description)
|
||||||
|
fmt.Printf(" Repo: %s\n", skill.Repository)
|
||||||
|
if skill.Author != "" {
|
||||||
|
fmt.Printf(" Author: %s\n", skill.Author)
|
||||||
|
}
|
||||||
|
if len(skill.Tags) > 0 {
|
||||||
|
fmt.Printf(" Tags: %v\n", skill.Tags)
|
||||||
|
}
|
||||||
|
fmt.Println()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillsShowCmd(loader *skills.SkillsLoader, skillName string) {
|
||||||
|
content, ok := loader.LoadSkill(skillName)
|
||||||
|
if !ok {
|
||||||
|
fmt.Printf("✗ Skill '%s' not found\n", skillName)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("\n📦 Skill: %s\n", skillName)
|
||||||
|
fmt.Println("----------------------")
|
||||||
|
fmt.Println(content)
|
||||||
|
}
|
||||||
102
cmd/picoclaw/cmd_status.go
Normal file
102
cmd/picoclaw/cmd_status.go
Normal file
|
|
@ -0,0 +1,102 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
)
|
||||||
|
|
||||||
|
func statusCmd() {
|
||||||
|
cfg, err := loadConfig()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
configPath := getConfigPath()
|
||||||
|
|
||||||
|
fmt.Printf("%s picoclaw Status\n", logo)
|
||||||
|
fmt.Printf("Version: %s\n", formatVersion())
|
||||||
|
build, _ := formatBuildInfo()
|
||||||
|
if build != "" {
|
||||||
|
fmt.Printf("Build: %s\n", build)
|
||||||
|
}
|
||||||
|
fmt.Println()
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err == nil {
|
||||||
|
fmt.Println("Config:", configPath, "✓")
|
||||||
|
} else {
|
||||||
|
fmt.Println("Config:", configPath, "✗")
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
if _, err := os.Stat(workspace); err == nil {
|
||||||
|
fmt.Println("Workspace:", workspace, "✓")
|
||||||
|
} else {
|
||||||
|
fmt.Println("Workspace:", workspace, "✗")
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.Stat(configPath); err == nil {
|
||||||
|
fmt.Printf("Model: %s\n", cfg.Agents.Defaults.Model)
|
||||||
|
|
||||||
|
hasOpenRouter := cfg.Providers.OpenRouter.APIKey != ""
|
||||||
|
hasAnthropic := cfg.Providers.Anthropic.APIKey != ""
|
||||||
|
hasOpenAI := cfg.Providers.OpenAI.APIKey != ""
|
||||||
|
hasGemini := cfg.Providers.Gemini.APIKey != ""
|
||||||
|
hasZhipu := cfg.Providers.Zhipu.APIKey != ""
|
||||||
|
hasQwen := cfg.Providers.Qwen.APIKey != ""
|
||||||
|
hasGroq := cfg.Providers.Groq.APIKey != ""
|
||||||
|
hasVLLM := cfg.Providers.VLLM.APIBase != ""
|
||||||
|
hasMoonshot := cfg.Providers.Moonshot.APIKey != ""
|
||||||
|
hasDeepSeek := cfg.Providers.DeepSeek.APIKey != ""
|
||||||
|
hasVolcEngine := cfg.Providers.VolcEngine.APIKey != ""
|
||||||
|
hasNvidia := cfg.Providers.Nvidia.APIKey != ""
|
||||||
|
hasOllama := cfg.Providers.Ollama.APIBase != ""
|
||||||
|
|
||||||
|
status := func(enabled bool) string {
|
||||||
|
if enabled {
|
||||||
|
return "✓"
|
||||||
|
}
|
||||||
|
return "not set"
|
||||||
|
}
|
||||||
|
fmt.Println("OpenRouter API:", status(hasOpenRouter))
|
||||||
|
fmt.Println("Anthropic API:", status(hasAnthropic))
|
||||||
|
fmt.Println("OpenAI API:", status(hasOpenAI))
|
||||||
|
fmt.Println("Gemini API:", status(hasGemini))
|
||||||
|
fmt.Println("Zhipu API:", status(hasZhipu))
|
||||||
|
fmt.Println("Qwen API:", status(hasQwen))
|
||||||
|
fmt.Println("Groq API:", status(hasGroq))
|
||||||
|
fmt.Println("Moonshot API:", status(hasMoonshot))
|
||||||
|
fmt.Println("DeepSeek API:", status(hasDeepSeek))
|
||||||
|
fmt.Println("VolcEngine API:", status(hasVolcEngine))
|
||||||
|
fmt.Println("Nvidia API:", status(hasNvidia))
|
||||||
|
if hasVLLM {
|
||||||
|
fmt.Printf("vLLM/Local: ✓ %s\n", cfg.Providers.VLLM.APIBase)
|
||||||
|
} else {
|
||||||
|
fmt.Println("vLLM/Local: not set")
|
||||||
|
}
|
||||||
|
if hasOllama {
|
||||||
|
fmt.Printf("Ollama: ✓ %s\n", cfg.Providers.Ollama.APIBase)
|
||||||
|
} else {
|
||||||
|
fmt.Println("Ollama: not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
store, _ := auth.LoadStore()
|
||||||
|
if store != nil && len(store.Credentials) > 0 {
|
||||||
|
fmt.Println("\nOAuth/Token Auth:")
|
||||||
|
for provider, cred := range store.Credentials {
|
||||||
|
status := "authenticated"
|
||||||
|
if cred.IsExpired() {
|
||||||
|
status = "expired"
|
||||||
|
} else if cred.NeedsRefresh() {
|
||||||
|
status = "needs refresh"
|
||||||
|
}
|
||||||
|
fmt.Printf(" %s (%s): %s\n", provider, cred.AuthMethod, status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
1244
cmd/picoclaw/main.go
1244
cmd/picoclaw/main.go
File diff suppressed because it is too large
Load diff
|
|
@ -3,12 +3,48 @@
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"restrict_to_workspace": true,
|
"restrict_to_workspace": true,
|
||||||
"model": "glm-4.7",
|
"model": "gpt4",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key",
|
||||||
|
"api_base": "https://api.anthropic.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gemini",
|
||||||
|
"model": "antigravity/gemini-2.0-flash",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "deepseek",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-your-deepseek-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "loadbalanced-gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key1",
|
||||||
|
"api_base": "https://api1.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "loadbalanced-gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key2",
|
||||||
|
"api_base": "https://api2.example.com/v1"
|
||||||
|
}
|
||||||
|
],
|
||||||
"channels": {
|
"channels": {
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
@ -21,6 +57,13 @@
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"token": "YOUR_DISCORD_BOT_TOKEN",
|
"token": "YOUR_DISCORD_BOT_TOKEN",
|
||||||
|
"allow_from": [],
|
||||||
|
"mention_only": false
|
||||||
|
},
|
||||||
|
"qq": {
|
||||||
|
"enabled": false,
|
||||||
|
"app_id": "YOUR_QQ_APP_ID",
|
||||||
|
"app_secret": "YOUR_QQ_APP_SECRET",
|
||||||
"allow_from": []
|
"allow_from": []
|
||||||
},
|
},
|
||||||
"maixcam": {
|
"maixcam": {
|
||||||
|
|
@ -73,13 +116,15 @@
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
"providers": {
|
||||||
|
"_comment": "DEPRECATED: Use model_list instead. This will be removed in a future version",
|
||||||
"anthropic": {
|
"anthropic": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": ""
|
"api_base": ""
|
||||||
},
|
},
|
||||||
"openai": {
|
"openai": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": ""
|
"api_base": "",
|
||||||
|
"web_search": true
|
||||||
},
|
},
|
||||||
"openrouter": {
|
"openrouter": {
|
||||||
"api_key": "sk-or-v1-xxx",
|
"api_key": "sk-or-v1-xxx",
|
||||||
|
|
@ -110,10 +155,18 @@
|
||||||
"api_key": "sk-xxx",
|
"api_key": "sk-xxx",
|
||||||
"api_base": ""
|
"api_base": ""
|
||||||
},
|
},
|
||||||
|
"qwen": {
|
||||||
|
"api_key": "sk-xxx",
|
||||||
|
"api_base": ""
|
||||||
|
},
|
||||||
"ollama": {
|
"ollama": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": "http://localhost:11434/v1"
|
"api_base": "http://localhost:11434/v1"
|
||||||
},
|
},
|
||||||
|
"cerebras": {
|
||||||
|
"api_key": "",
|
||||||
|
"api_base": ""
|
||||||
|
},
|
||||||
"volcengine": {
|
"volcengine": {
|
||||||
"api_key": "",
|
"api_key": "",
|
||||||
"api_base": ""
|
"api_base": ""
|
||||||
|
|
@ -126,6 +179,10 @@
|
||||||
"api_key": "YOUR_BRAVE_API_KEY",
|
"api_key": "YOUR_BRAVE_API_KEY",
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
},
|
},
|
||||||
|
"duckduckgo": {
|
||||||
|
"enabled": true,
|
||||||
|
"max_results": 5
|
||||||
|
},
|
||||||
"perplexity": {
|
"perplexity": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "pplx-xxx",
|
"api_key": "pplx-xxx",
|
||||||
|
|
@ -134,6 +191,21 @@
|
||||||
},
|
},
|
||||||
"cron": {
|
"cron": {
|
||||||
"exec_timeout_minutes": 5
|
"exec_timeout_minutes": 5
|
||||||
|
},
|
||||||
|
"exec": {
|
||||||
|
"enable_deny_patterns": false,
|
||||||
|
"custom_deny_patterns": []
|
||||||
|
},
|
||||||
|
"skills": {
|
||||||
|
"registries": {
|
||||||
|
"clawhub": {
|
||||||
|
"enabled": true,
|
||||||
|
"base_url": "https://clawhub.ai",
|
||||||
|
"search_path": "/api/v1/search",
|
||||||
|
"skills_path": "/api/v1/skills",
|
||||||
|
"download_path": "/api/v1/download"
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"heartbeat": {
|
"heartbeat": {
|
||||||
|
|
@ -148,4 +220,4 @@
|
||||||
"host": "0.0.0.0",
|
"host": "0.0.0.0",
|
||||||
"port": 18790
|
"port": 18790
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -11,8 +11,8 @@ services:
|
||||||
profiles:
|
profiles:
|
||||||
- agent
|
- agent
|
||||||
volumes:
|
volumes:
|
||||||
- ./config/config.json:/root/.picoclaw/config.json:ro
|
- ./config/config.json:/home/picoclaw/.picoclaw/config.json:ro
|
||||||
- picoclaw-workspace:/root/.picoclaw/workspace
|
- picoclaw-workspace:/home/picoclaw/.picoclaw/workspace
|
||||||
entrypoint: ["picoclaw", "agent"]
|
entrypoint: ["picoclaw", "agent"]
|
||||||
stdin_open: true
|
stdin_open: true
|
||||||
tty: true
|
tty: true
|
||||||
|
|
@ -31,9 +31,9 @@ services:
|
||||||
- gateway
|
- gateway
|
||||||
volumes:
|
volumes:
|
||||||
# Configuration file
|
# Configuration file
|
||||||
- ./config/config.json:/root/.picoclaw/config.json:ro
|
- ./config/config.json:/home/picoclaw/.picoclaw/config.json:ro
|
||||||
# Persistent workspace (sessions, memory, logs)
|
# Persistent workspace (sessions, memory, logs)
|
||||||
- picoclaw-workspace:/root/.picoclaw/workspace
|
- picoclaw-workspace:/home/picoclaw/.picoclaw/workspace
|
||||||
command: ["gateway"]
|
command: ["gateway"]
|
||||||
|
|
||||||
volumes:
|
volumes:
|
||||||
|
|
|
||||||
1002
docs/ANTIGRAVITY_AUTH.md
Normal file
1002
docs/ANTIGRAVITY_AUTH.md
Normal file
File diff suppressed because it is too large
Load diff
72
docs/ANTIGRAVITY_USAGE.md
Normal file
72
docs/ANTIGRAVITY_USAGE.md
Normal file
|
|
@ -0,0 +1,72 @@
|
||||||
|
# Using Antigravity Provider in PicoClaw
|
||||||
|
|
||||||
|
This guide explains how to set up and use the **Antigravity** (Google Cloud Code Assist) provider in PicoClaw.
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
1. A Google account.
|
||||||
|
2. Google Cloud Code Assist enabled (usually available via the "Gemini for Google Cloud" onboarding).
|
||||||
|
|
||||||
|
## 1. Authentication
|
||||||
|
|
||||||
|
To authenticate with Antigravity, run the following command:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth login --provider antigravity
|
||||||
|
```
|
||||||
|
|
||||||
|
### Manual Authentication (Headless/VPS)
|
||||||
|
If you are running on a server (Coolify/Docker) and cannot reach `localhost`, follow these steps:
|
||||||
|
1. Run the command above.
|
||||||
|
2. Copy the URL provided and open it in your local browser.
|
||||||
|
3. Complete the login.
|
||||||
|
4. Your browser will redirect to a `localhost:51121` URL (which will fail to load).
|
||||||
|
5. **Copy that final URL** from your browser's address bar.
|
||||||
|
6. **Paste it back into the terminal** where PicoClaw is waiting.
|
||||||
|
|
||||||
|
PicoClaw will extract the authorization code and complete the process automatically.
|
||||||
|
|
||||||
|
## 2. Managing Models
|
||||||
|
|
||||||
|
### List Available Models
|
||||||
|
To see which models your project has access to and check their quotas:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth models
|
||||||
|
```
|
||||||
|
|
||||||
|
### Switch Models
|
||||||
|
You can change the default model in `~/.picoclaw/config.json` or override it via the CLI:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Override for a single command
|
||||||
|
picoclaw agent -m "Hello" --model claude-opus-4-6-thinking
|
||||||
|
```
|
||||||
|
|
||||||
|
## 3. Real-world Usage (Coolify/Docker)
|
||||||
|
|
||||||
|
If you are deploying via Coolify or Docker, follow these steps to test:
|
||||||
|
|
||||||
|
1. **Branch**: Use the `feat/antigravity-provider` branch.
|
||||||
|
2. **Environment Variables**:
|
||||||
|
* `PICOCLAW_AGENTS_DEFAULTS_PROVIDER=antigravity`
|
||||||
|
* `PICOCLAW_AGENTS_DEFAULTS_MODEL=gemini-3-flash`
|
||||||
|
3. **Authentication persistence**:
|
||||||
|
If you've logged in locally, you can copy your credentials to the server:
|
||||||
|
```bash
|
||||||
|
scp ~/.picoclaw/auth-profiles.json user@your-server:~/.picoclaw/
|
||||||
|
```
|
||||||
|
*Alternatively*, run the `auth login` command once on the server if you have terminal access.
|
||||||
|
|
||||||
|
## 4. Troubleshooting
|
||||||
|
|
||||||
|
* **Empty Response**: If a model returns an empty reply, it may be restricted for your project. Try `gemini-3-flash` or `claude-opus-4-6-thinking`.
|
||||||
|
* **429 Rate Limit**: Antigravity has strict quotas. PicoClaw will display the "reset time" in the error message if you hit a limit.
|
||||||
|
* **404 Not Found**: Ensure you are using a model ID from the `picoclaw auth models` list. Use the short ID (e.g., `gemini-3-flash`) not the full path.
|
||||||
|
|
||||||
|
## 5. Summary of Working Models
|
||||||
|
|
||||||
|
Based on testing, the following models are most reliable:
|
||||||
|
* `gemini-3-flash` (Fast, highly available)
|
||||||
|
* `gemini-2.5-flash-lite` (Lightweight)
|
||||||
|
* `claude-opus-4-6-thinking` (Powerful, includes reasoning)
|
||||||
179
docs/design/provider-refactoring-tests.md
Normal file
179
docs/design/provider-refactoring-tests.md
Normal file
|
|
@ -0,0 +1,179 @@
|
||||||
|
# Provider Architecture Refactoring - Test Suite Summary
|
||||||
|
|
||||||
|
> PRD: `tasks/prd-provider-refactoring.md`
|
||||||
|
|
||||||
|
This document summarizes the complete test suite designed for the Provider architecture refactoring.
|
||||||
|
|
||||||
|
## Test File Structure
|
||||||
|
|
||||||
|
```
|
||||||
|
pkg/
|
||||||
|
├── config/
|
||||||
|
│ ├── model_config_test.go # US-001, US-002: ModelConfig struct and GetModelConfig tests
|
||||||
|
│ └── migration_test.go # US-003: Backward compatibility and migration tests
|
||||||
|
├── providers/
|
||||||
|
│ ├── registry_test.go # US-006: Load balancing tests
|
||||||
|
│ ├── integration_test.go # E2E integration tests
|
||||||
|
│ └── factory/
|
||||||
|
│ └── factory_test.go # US-004, US-005: Provider factory tests
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Test Case Checklist
|
||||||
|
|
||||||
|
### 1. `pkg/config/model_config_test.go` - Configuration Parsing Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestModelConfig_Parsing` | Verify ModelConfig JSON parsing | US-001 |
|
||||||
|
| `TestModelConfig_ModelListInConfig` | Verify model_list parsing in Config | US-001 |
|
||||||
|
| `TestModelConfig_Validation` | Verify required field validation | US-001 |
|
||||||
|
| `TestConfig_GetModelConfig_Found` | Verify GetModelConfig finds model | US-002 |
|
||||||
|
| `TestConfig_GetModelConfig_NotFound` | Verify GetModelConfig returns error | US-002 |
|
||||||
|
| `TestConfig_GetModelConfig_EmptyModelList` | Verify empty model_list handling | US-002 |
|
||||||
|
| `TestConfig_BackwardCompatibility_ProvidersToModelList` | Verify old config conversion | US-003 |
|
||||||
|
| `TestConfig_DeprecationWarning` | Verify deprecation warning | US-003 |
|
||||||
|
| `TestModelConfig_ProtocolExtraction` | Verify protocol prefix extraction | US-004 |
|
||||||
|
| `TestConfig_ModelNameUniqueness` | Verify model_name uniqueness | US-001 |
|
||||||
|
|
||||||
|
### 2. `pkg/config/migration_test.go` - Migration Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestConvertProvidersToModelList_OpenAI` | OpenAI config conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_Anthropic` | Anthropic config conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_MultipleProviders` | Multiple provider conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_EmptyProviders` | Empty providers handling | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_GitHubCopilot` | GitHub Copilot conversion | US-003 |
|
||||||
|
| `TestConvertProvidersToModelList_Antigravity` | Antigravity conversion | US-003 |
|
||||||
|
| `TestGenerateModelName_*` | Model name generation | US-003 |
|
||||||
|
| `TestHasProvidersConfig_*` | Detect old config existence | US-003 |
|
||||||
|
| `TestValidateMigration_*` | Migration validation | US-003 |
|
||||||
|
| `TestMigrateConfig_DryRun` | Dry run migration | US-003 |
|
||||||
|
| `TestMigrateConfig_Actual` | Actual migration | US-003 |
|
||||||
|
|
||||||
|
### 3. `pkg/providers/registry_test.go` - Load Balancing Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestModelRegistry_SingleConfig` | Single config returns same result | US-006 |
|
||||||
|
| `TestModelRegistry_RoundRobinSelection` | 3-config round-robin selection | US-006 |
|
||||||
|
| `TestModelRegistry_RoundRobinTwoConfigs` | 2-config round-robin selection | US-006 |
|
||||||
|
| `TestModelRegistry_ConcurrentAccess` | Concurrent access thread safety | US-006 |
|
||||||
|
| `TestModelRegistry_RaceDetection` | Data race detection | US-006 |
|
||||||
|
| `TestModelRegistry_ModelNotFound` | Model not found error | US-006 |
|
||||||
|
| `TestModelRegistry_EmptyRegistry` | Empty registry handling | US-006 |
|
||||||
|
| `TestModelRegistry_MultipleModels` | Multiple model registration | US-006 |
|
||||||
|
| `TestModelRegistry_MixedSingleAndMultiple` | Single/multiple config mix | US-006 |
|
||||||
|
| `TestModelRegistry_CaseSensitiveModelNames` | Case sensitivity | US-006 |
|
||||||
|
|
||||||
|
### 4. `pkg/providers/factory/factory_test.go` - Provider Factory Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestCreateProviderFromConfig_OpenAI` | Create OpenAI provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_OpenAIDefault` | Default openai protocol | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_Anthropic` | Create Anthropic provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_Antigravity` | Create Antigravity provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_ClaudeCLI` | Create Claude CLI provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_CodexCLI` | Create Codex CLI provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_GitHubCopilot` | Create GitHub Copilot provider | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_UnknownProtocol` | Unknown protocol error handling | US-004 |
|
||||||
|
| `TestCreateProviderFromConfig_MissingAPIKey` | Missing API key error | US-004 |
|
||||||
|
| `TestExtractProtocol` | Protocol prefix extraction | US-004 |
|
||||||
|
| `TestCreateProvider_UsesModelList` | Create using model_list | US-005 |
|
||||||
|
| `TestCreateProvider_FallbackToProviders` | Fallback to providers | US-005 |
|
||||||
|
| `TestCreateProvider_PriorityModelListOverProviders` | model_list priority | US-005 |
|
||||||
|
|
||||||
|
### 5. `pkg/providers/integration_test.go` - E2E Integration Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose | PRD Reference |
|
||||||
|
|-----------|---------|---------------|
|
||||||
|
| `TestE2E_OpenAICompatibleProvider_NoCodeChange` | Zero-code provider addition | Goal |
|
||||||
|
| `TestE2E_LoadBalancing_RoundRobin` | Load balancing actual effect | US-006 |
|
||||||
|
| `TestE2E_BackwardCompatibility_OldProvidersConfig` | Old config compatibility | US-003 |
|
||||||
|
| `TestE2E_ErrorHandling_ModelNotFound` | Model not found | FR-30 |
|
||||||
|
| `TestE2E_ErrorHandling_MissingAPIKey` | Missing API key | FR-31 |
|
||||||
|
| `TestE2E_ErrorHandling_InvalidAPIBase` | Invalid API base | FR-30 |
|
||||||
|
| `TestE2E_ToolCalls_OpenAICompatible` | Tool call support | - |
|
||||||
|
| `TestE2E_AntigravityProvider` | Antigravity provider | US-004 |
|
||||||
|
| `TestE2E_ClaudeCLIProvider` | Claude CLI provider | US-004 |
|
||||||
|
|
||||||
|
### 6. Performance Tests
|
||||||
|
|
||||||
|
| Test Name | Purpose |
|
||||||
|
|-----------|---------|
|
||||||
|
| `BenchmarkCreateProviderFromConfig` | Provider creation performance |
|
||||||
|
| `BenchmarkGetModelConfig` | Model lookup performance |
|
||||||
|
| `BenchmarkGetModelConfigParallel` | Concurrent lookup performance |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Running Tests
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Run all tests
|
||||||
|
go test ./pkg/... -v
|
||||||
|
|
||||||
|
# Run with data race detection
|
||||||
|
go test ./pkg/... -race
|
||||||
|
|
||||||
|
# Run specific package tests
|
||||||
|
go test ./pkg/config -v
|
||||||
|
go test ./pkg/providers -v
|
||||||
|
go test ./pkg/providers/factory -v
|
||||||
|
|
||||||
|
# Run E2E tests
|
||||||
|
go test ./pkg/providers -run TestE2E -v
|
||||||
|
|
||||||
|
# Run performance tests
|
||||||
|
go test ./pkg/providers -bench=. -benchmem
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## PRD Acceptance Criteria Mapping
|
||||||
|
|
||||||
|
| PRD Acceptance Criteria | Test Cases |
|
||||||
|
|------------------------|------------|
|
||||||
|
| US-001: Add ModelConfig struct | `TestModelConfig_Parsing`, `TestModelConfig_Validation` |
|
||||||
|
| US-001: model_name unique | `TestConfig_ModelNameUniqueness` |
|
||||||
|
| US-002: GetModelConfig method | `TestConfig_GetModelConfig_*` |
|
||||||
|
| US-003: Auto-convert providers | `TestConvertProvidersToModelList_*` |
|
||||||
|
| US-003: Deprecation warning | `TestConfig_DeprecationWarning` |
|
||||||
|
| US-003: Existing tests pass | (existing test files unchanged) |
|
||||||
|
| US-004: Protocol prefix factory | `TestExtractProtocol`, `TestCreateProviderFromConfig_*` |
|
||||||
|
| US-004: Default prefix openai | `TestCreateProviderFromConfig_OpenAIDefault` |
|
||||||
|
| US-005: CreateProvider uses factory | `TestCreateProvider_*` |
|
||||||
|
| US-006: Round-robin selection | `TestModelRegistry_RoundRobin*` |
|
||||||
|
| US-006: Thread-safe atomic | `TestModelRegistry_RaceDetection` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Recommended Implementation Order
|
||||||
|
|
||||||
|
1. **Phase 1: Configuration Structure** (US-001, US-002)
|
||||||
|
- Implement `ModelConfig` struct
|
||||||
|
- Implement `GetModelConfig` method
|
||||||
|
- Run `model_config_test.go`
|
||||||
|
|
||||||
|
2. **Phase 2: Protocol Factory** (US-004)
|
||||||
|
- Implement `CreateProviderFromConfig`
|
||||||
|
- Implement `ExtractProtocol`
|
||||||
|
- Run `factory_test.go`
|
||||||
|
|
||||||
|
3. **Phase 3: Load Balancing** (US-006)
|
||||||
|
- Implement `ModelRegistry`
|
||||||
|
- Implement round-robin selection
|
||||||
|
- Run `registry_test.go` (with `-race`)
|
||||||
|
|
||||||
|
4. **Phase 4: Backward Compatibility** (US-003, US-005)
|
||||||
|
- Implement `ConvertProvidersToModelList`
|
||||||
|
- Refactor `CreateProvider`
|
||||||
|
- Run `migration_test.go`
|
||||||
|
- Verify existing tests pass
|
||||||
|
|
||||||
|
5. **Phase 5: E2E Verification**
|
||||||
|
- Run `integration_test.go`
|
||||||
|
- Manual testing with `config.example.json`
|
||||||
334
docs/design/provider-refactoring.md
Normal file
334
docs/design/provider-refactoring.md
Normal file
|
|
@ -0,0 +1,334 @@
|
||||||
|
# Provider Architecture Refactoring Design
|
||||||
|
|
||||||
|
> Issue: #283
|
||||||
|
> Discussion: #122
|
||||||
|
> Branch: feat/refactor-provider-by-protocol
|
||||||
|
|
||||||
|
## 1. Current Problems
|
||||||
|
|
||||||
|
### 1.1 Configuration Structure Issues
|
||||||
|
|
||||||
|
**Current State**: Each Provider requires a predefined field in `ProvidersConfig`
|
||||||
|
|
||||||
|
```go
|
||||||
|
type ProvidersConfig struct {
|
||||||
|
Anthropic ProviderConfig `json:"anthropic"`
|
||||||
|
OpenAI ProviderConfig `json:"openai"`
|
||||||
|
DeepSeek ProviderConfig `json:"deepseek"`
|
||||||
|
Qwen ProviderConfig `json:"qwen"`
|
||||||
|
Cerebras ProviderConfig `json:"cerebras"`
|
||||||
|
VolcEngine ProviderConfig `json:"volcengine"`
|
||||||
|
// ... every new provider requires changes here
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Problems**:
|
||||||
|
- Adding a new Provider requires modifying Go code (struct definition)
|
||||||
|
- `CreateProvider` function in `http_provider.go` has 200+ lines of switch-case
|
||||||
|
- Most Providers are OpenAI-compatible, but code is duplicated
|
||||||
|
|
||||||
|
### 1.2 Code Bloat Trend
|
||||||
|
|
||||||
|
Recent PRs demonstrate this issue:
|
||||||
|
|
||||||
|
| PR | Provider | Code Changes |
|
||||||
|
|----|----------|--------------|
|
||||||
|
| #365 | Qwen | +17 lines to http_provider.go |
|
||||||
|
| #333 | Cerebras | +17 lines to http_provider.go |
|
||||||
|
| #368 | Volcengine | +18 lines to http_provider.go |
|
||||||
|
|
||||||
|
Each OpenAI-compatible Provider requires:
|
||||||
|
1. Modify `config.go` to add configuration field
|
||||||
|
2. Modify `http_provider.go` to add switch case
|
||||||
|
3. Update documentation
|
||||||
|
|
||||||
|
### 1.3 Agent-Provider Coupling
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "deepseek", // need to know provider name
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Problem: Agent needs to know both `provider` and `model`, adding complexity.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. New Approach: model_list
|
||||||
|
|
||||||
|
### 2.1 Core Principles
|
||||||
|
|
||||||
|
Inspired by [LiteLLM](https://docs.litellm.ai/docs/proxy/configs) design:
|
||||||
|
|
||||||
|
1. **Model-centric**: Users care about models, not providers
|
||||||
|
2. **Protocol prefix**: Use `protocol/model_name` format, e.g., `openai/gpt-5.2`, `anthropic/claude-sonnet-4.6`
|
||||||
|
3. **Configuration-driven**: Adding new Providers only requires config changes, no code changes
|
||||||
|
|
||||||
|
### 2.2 New Configuration Structure
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "openai/deepseek-chat",
|
||||||
|
"api_base": "https://api.deepseek.com/v1",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.2",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gemini-3-flash",
|
||||||
|
"model": "antigravity/gemini-3-flash",
|
||||||
|
"auth_method": "oauth"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "my-company-llm",
|
||||||
|
"model": "openai/company-model-v1",
|
||||||
|
"api_base": "https://llm.company.com/v1",
|
||||||
|
"api_key": "xxx"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"max_tokens": 8192,
|
||||||
|
"temperature": 0.7
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.3 Go Struct Definition
|
||||||
|
|
||||||
|
```go
|
||||||
|
type Config struct {
|
||||||
|
ModelList []ModelConfig `json:"model_list"` // new
|
||||||
|
Providers ProvidersConfig `json:"providers"` // old, deprecated
|
||||||
|
|
||||||
|
Agents AgentsConfig `json:"agents"`
|
||||||
|
Channels ChannelsConfig `json:"channels"`
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
|
||||||
|
type ModelConfig struct {
|
||||||
|
// Required
|
||||||
|
ModelName string `json:"model_name"` // user-facing name (alias)
|
||||||
|
Model string `json:"model"` // protocol/model, e.g., openai/gpt-5.2
|
||||||
|
|
||||||
|
// Common config
|
||||||
|
APIBase string `json:"api_base,omitempty"`
|
||||||
|
APIKey string `json:"api_key,omitempty"`
|
||||||
|
Proxy string `json:"proxy,omitempty"`
|
||||||
|
|
||||||
|
// Special provider config
|
||||||
|
AuthMethod string `json:"auth_method,omitempty"` // oauth, token
|
||||||
|
ConnectMode string `json:"connect_mode,omitempty"` // stdio, grpc
|
||||||
|
|
||||||
|
// Optional optimizations
|
||||||
|
RPM int `json:"rpm,omitempty"` // rate limit
|
||||||
|
MaxTokensField string `json:"max_tokens_field,omitempty"` // max_tokens or max_completion_tokens
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.4 Protocol Recognition
|
||||||
|
|
||||||
|
Identify protocol via prefix in `model` field:
|
||||||
|
|
||||||
|
| Prefix | Protocol | Description |
|
||||||
|
|--------|----------|-------------|
|
||||||
|
| `openai/` | OpenAI-compatible | Most common, includes DeepSeek, Qwen, Groq, etc. |
|
||||||
|
| `anthropic/` | Anthropic | Claude series specific |
|
||||||
|
| `antigravity/` | Antigravity | Google Cloud Code Assist |
|
||||||
|
| `gemini/` | Gemini | Google Gemini native API (if needed) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Design Rationale
|
||||||
|
|
||||||
|
### 3.1 Problems Solved
|
||||||
|
|
||||||
|
| Problem | Old Approach | New Approach |
|
||||||
|
|---------|--------------|--------------|
|
||||||
|
| Add OpenAI-compatible Provider | Change 3 code locations | Add one config entry |
|
||||||
|
| Agent specifies model | Need provider + model | Only need model |
|
||||||
|
| Code duplication | Each Provider duplicates logic | Share protocol implementation |
|
||||||
|
| Multi-Agent support | Complex | Naturally compatible |
|
||||||
|
|
||||||
|
### 3.2 Multi-Agent Compatibility
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [...],
|
||||||
|
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
},
|
||||||
|
"coder": {
|
||||||
|
"model": "gpt-5.2",
|
||||||
|
"system_prompt": "You are a coding assistant..."
|
||||||
|
},
|
||||||
|
"translator": {
|
||||||
|
"model": "claude-sonnet-4.6"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Each Agent only needs to specify `model` (corresponds to `model_name` in `model_list`).
|
||||||
|
|
||||||
|
### 3.3 Industry Comparison
|
||||||
|
|
||||||
|
**LiteLLM** (most mature open-source LLM Proxy) uses similar design:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
- model_name: gpt-4o
|
||||||
|
litellm_params:
|
||||||
|
model: openai/gpt-5.2
|
||||||
|
api_key: xxx
|
||||||
|
- model_name: my-custom
|
||||||
|
litellm_params:
|
||||||
|
model: openai/custom-model
|
||||||
|
api_base: https://my-api.com/v1
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Migration Plan
|
||||||
|
|
||||||
|
### 4.1 Phase 1: Compatibility Period (v1.x)
|
||||||
|
|
||||||
|
Support both `providers` and `model_list`:
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (c *Config) GetModelConfig(modelName string) (*ModelConfig, error) {
|
||||||
|
// Prefer new config
|
||||||
|
if len(c.ModelList) > 0 {
|
||||||
|
return c.findModelByName(modelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Backward compatibility with old config
|
||||||
|
if !c.Providers.IsEmpty() {
|
||||||
|
logger.Warn("'providers' config is deprecated, please migrate to 'model_list'")
|
||||||
|
return c.convertFromProviders(modelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("model %s not found", modelName)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.2 Phase 2: Warning Period (late v1.x)
|
||||||
|
|
||||||
|
- Print more prominent warnings at startup
|
||||||
|
- Provide automatic migration script
|
||||||
|
- Mark `providers` as deprecated in documentation
|
||||||
|
|
||||||
|
### 4.3 Phase 3: Removal Period (v2.0)
|
||||||
|
|
||||||
|
- Completely remove `providers` support
|
||||||
|
- Remove `agents.defaults.provider` field
|
||||||
|
- Only support `model_list`
|
||||||
|
|
||||||
|
### 4.4 Configuration Migration Example
|
||||||
|
|
||||||
|
**Old Config**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"deepseek": {
|
||||||
|
"api_key": "sk-xxx",
|
||||||
|
"api_base": "https://api.deepseek.com/v1"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "deepseek",
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**New Config**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "deepseek-chat",
|
||||||
|
"model": "openai/deepseek-chat",
|
||||||
|
"api_base": "https://api.deepseek.com/v1",
|
||||||
|
"api_key": "sk-xxx"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Implementation Checklist
|
||||||
|
|
||||||
|
### 5.1 Configuration Layer
|
||||||
|
|
||||||
|
- [ ] Add `ModelConfig` struct
|
||||||
|
- [ ] Add `Config.ModelList` field
|
||||||
|
- [ ] Implement `GetModelConfig(modelName)` method
|
||||||
|
- [ ] Implement old config compatibility conversion
|
||||||
|
- [ ] Add `model_name` uniqueness validation
|
||||||
|
|
||||||
|
### 5.2 Provider Layer
|
||||||
|
|
||||||
|
- [ ] Create `pkg/providers/factory/` directory
|
||||||
|
- [ ] Implement `CreateProviderFromModelConfig()`
|
||||||
|
- [ ] Refactor `http_provider.go` to `openai/provider.go`
|
||||||
|
- [ ] Maintain backward compatibility for old `CreateProvider()`
|
||||||
|
|
||||||
|
### 5.3 Testing
|
||||||
|
|
||||||
|
- [ ] New config unit tests
|
||||||
|
- [ ] Old config compatibility tests
|
||||||
|
- [ ] Integration tests
|
||||||
|
|
||||||
|
### 5.4 Documentation
|
||||||
|
|
||||||
|
- [ ] Update README
|
||||||
|
- [ ] Update config.example.json
|
||||||
|
- [ ] Write migration guide
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Risks and Mitigations
|
||||||
|
|
||||||
|
| Risk | Mitigation |
|
||||||
|
|------|------------|
|
||||||
|
| Breaking existing configs | Compatibility period keeps old config working |
|
||||||
|
| User migration cost | Provide automatic migration script |
|
||||||
|
| Special Provider incompatibility | Keep `auth_method` and other extension fields |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. References
|
||||||
|
|
||||||
|
- [LiteLLM Config Documentation](https://docs.litellm.ai/docs/proxy/configs)
|
||||||
|
- [One-API GitHub](https://github.com/songquanpeng/one-api)
|
||||||
|
- Discussion #122: Refactor Provider Architecture
|
||||||
211
docs/migration/model-list-migration.md
Normal file
211
docs/migration/model-list-migration.md
Normal file
|
|
@ -0,0 +1,211 @@
|
||||||
|
# Migration Guide: From `providers` to `model_list`
|
||||||
|
|
||||||
|
This guide explains how to migrate from the legacy `providers` configuration to the new `model_list` format.
|
||||||
|
|
||||||
|
## Why Migrate?
|
||||||
|
|
||||||
|
The new `model_list` configuration offers several advantages:
|
||||||
|
|
||||||
|
- **Zero-code provider addition**: Add OpenAI-compatible providers with configuration only
|
||||||
|
- **Load balancing**: Configure multiple endpoints for the same model
|
||||||
|
- **Protocol-based routing**: Use prefixes like `openai/`, `anthropic/`, etc.
|
||||||
|
- **Cleaner configuration**: Model-centric instead of vendor-centric
|
||||||
|
|
||||||
|
## Timeline
|
||||||
|
|
||||||
|
| Version | Status |
|
||||||
|
|---------|--------|
|
||||||
|
| v1.x | `model_list` introduced, `providers` deprecated but functional |
|
||||||
|
| v1.x+1 | Prominent deprecation warnings, migration tool available |
|
||||||
|
| v2.0 | `providers` configuration removed |
|
||||||
|
|
||||||
|
## Before and After
|
||||||
|
|
||||||
|
### Before: Legacy `providers` Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"openai": {
|
||||||
|
"api_key": "sk-your-openai-key",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
"anthropic": {
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
"deepseek": {
|
||||||
|
"api_key": "sk-your-deepseek-key"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "openai",
|
||||||
|
"model": "gpt-5.2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### After: New `model_list` Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-your-openai-key",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_key": "sk-ant-your-key"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "deepseek",
|
||||||
|
"model": "deepseek/deepseek-chat",
|
||||||
|
"api_key": "sk-your-deepseek-key"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "gpt4"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Protocol Prefixes
|
||||||
|
|
||||||
|
The `model` field uses a protocol prefix format: `[protocol/]model-identifier`
|
||||||
|
|
||||||
|
| Prefix | Description | Example |
|
||||||
|
|--------|-------------|---------|
|
||||||
|
| `openai/` | OpenAI API (default) | `openai/gpt-5.2` |
|
||||||
|
| `anthropic/` | Anthropic API | `anthropic/claude-opus-4` |
|
||||||
|
| `antigravity/` | Google via Antigravity OAuth | `antigravity/gemini-2.0-flash` |
|
||||||
|
| `claude-cli/` | Claude CLI (local) | `claude-cli/claude-sonnet-4.6` |
|
||||||
|
| `codex-cli/` | Codex CLI (local) | `codex-cli/codex-4` |
|
||||||
|
| `github-copilot/` | GitHub Copilot | `github-copilot/gpt-4o` |
|
||||||
|
| `openrouter/` | OpenRouter | `openrouter/anthropic/claude-sonnet-4.6` |
|
||||||
|
| `groq/` | Groq API | `groq/llama-3.1-70b` |
|
||||||
|
| `deepseek/` | DeepSeek API | `deepseek/deepseek-chat` |
|
||||||
|
| `cerebras/` | Cerebras API | `cerebras/llama-3.3-70b` |
|
||||||
|
| `qwen/` | Alibaba Qwen | `qwen/qwen-max` |
|
||||||
|
|
||||||
|
**Note**: If no prefix is specified, `openai/` is used as the default.
|
||||||
|
|
||||||
|
## ModelConfig Fields
|
||||||
|
|
||||||
|
| Field | Required | Description |
|
||||||
|
|-------|----------|-------------|
|
||||||
|
| `model_name` | Yes | User-facing alias for the model |
|
||||||
|
| `model` | Yes | Protocol and model identifier (e.g., `openai/gpt-5.2`) |
|
||||||
|
| `api_base` | No | API endpoint URL |
|
||||||
|
| `api_key` | No* | API authentication key |
|
||||||
|
| `proxy` | No | HTTP proxy URL |
|
||||||
|
| `auth_method` | No | Authentication method: `oauth`, `token` |
|
||||||
|
| `connect_mode` | No | Connection mode for CLI providers: `stdio`, `grpc` |
|
||||||
|
| `rpm` | No | Requests per minute limit |
|
||||||
|
| `max_tokens_field` | No | Field name for max tokens |
|
||||||
|
|
||||||
|
*`api_key` is required for HTTP-based protocols unless `api_base` points to a local server.
|
||||||
|
|
||||||
|
## Load Balancing
|
||||||
|
|
||||||
|
Configure multiple endpoints for the same model to distribute load:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key1",
|
||||||
|
"api_base": "https://api1.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key2",
|
||||||
|
"api_base": "https://api2.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "gpt4",
|
||||||
|
"model": "openai/gpt-5.2",
|
||||||
|
"api_key": "sk-key3",
|
||||||
|
"api_base": "https://api3.example.com/v1"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
When you request model `gpt4`, requests will be distributed across all three endpoints using round-robin selection.
|
||||||
|
|
||||||
|
## Adding a New OpenAI-Compatible Provider
|
||||||
|
|
||||||
|
With `model_list`, adding a new provider requires zero code changes:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "my-custom-llm",
|
||||||
|
"model": "openai/my-model-v1",
|
||||||
|
"api_key": "your-api-key",
|
||||||
|
"api_base": "https://api.your-provider.com/v1"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Just specify `openai/` as the protocol (or omit it for the default), and provide your provider's API base URL.
|
||||||
|
|
||||||
|
## Backward Compatibility
|
||||||
|
|
||||||
|
During the migration period, your existing `providers` configuration will continue to work:
|
||||||
|
|
||||||
|
1. If `model_list` is empty and `providers` has data, the system auto-converts internally
|
||||||
|
2. A deprecation warning is logged: `"providers config is deprecated, please migrate to model_list"`
|
||||||
|
3. All existing functionality remains unchanged
|
||||||
|
|
||||||
|
## Migration Checklist
|
||||||
|
|
||||||
|
- [ ] Identify all providers you're currently using
|
||||||
|
- [ ] Create `model_list` entries for each provider
|
||||||
|
- [ ] Use appropriate protocol prefixes
|
||||||
|
- [ ] Update `agents.defaults.model` to reference the new `model_name`
|
||||||
|
- [ ] Test that all models work correctly
|
||||||
|
- [ ] Remove or comment out the old `providers` section
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### Model not found error
|
||||||
|
|
||||||
|
```
|
||||||
|
model "xxx" not found in model_list or providers
|
||||||
|
```
|
||||||
|
|
||||||
|
**Solution**: Ensure the `model_name` in `model_list` matches the value in `agents.defaults.model`.
|
||||||
|
|
||||||
|
### Unknown protocol error
|
||||||
|
|
||||||
|
```
|
||||||
|
unknown protocol "xxx" in model "xxx/model-name"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Solution**: Use a supported protocol prefix. See the [Protocol Prefixes](#protocol-prefixes) table above.
|
||||||
|
|
||||||
|
### Missing API key error
|
||||||
|
|
||||||
|
```
|
||||||
|
api_key or api_base is required for HTTP-based protocol "xxx"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Solution**: Provide `api_key` and/or `api_base` for HTTP-based providers.
|
||||||
|
|
||||||
|
## Need Help?
|
||||||
|
|
||||||
|
- [GitHub Issues](https://github.com/sipeed/picoclaw/issues)
|
||||||
|
- [Discussion #122](https://github.com/sipeed/picoclaw/discussions/122): Original proposal
|
||||||
122
docs/tools_configuration.md
Normal file
122
docs/tools_configuration.md
Normal file
|
|
@ -0,0 +1,122 @@
|
||||||
|
# Tools Configuration
|
||||||
|
|
||||||
|
PicoClaw's tools configuration is located in the `tools` field of `config.json`.
|
||||||
|
|
||||||
|
## Directory Structure
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"web": { ... },
|
||||||
|
"exec": { ... },
|
||||||
|
"approval": { ... },
|
||||||
|
"cron": { ... }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Web Tools
|
||||||
|
|
||||||
|
Web tools are used for web search and fetching.
|
||||||
|
|
||||||
|
### Brave
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|--------|------|---------|-------------|
|
||||||
|
| `enabled` | bool | false | Enable Brave search |
|
||||||
|
| `api_key` | string | - | Brave Search API key |
|
||||||
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
||||||
|
### DuckDuckGo
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|--------|------|---------|-------------|
|
||||||
|
| `enabled` | bool | true | Enable DuckDuckGo search |
|
||||||
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
||||||
|
### Perplexity
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|--------|------|---------|-------------|
|
||||||
|
| `enabled` | bool | false | Enable Perplexity search |
|
||||||
|
| `api_key` | string | - | Perplexity API key |
|
||||||
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
||||||
|
## Exec Tool
|
||||||
|
|
||||||
|
The exec tool is used to execute shell commands.
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|--------|------|---------|-------------|
|
||||||
|
| `enable_deny_patterns` | bool | true | Enable default dangerous command blocking |
|
||||||
|
| `custom_deny_patterns` | array | [] | Custom deny patterns (regular expressions) |
|
||||||
|
|
||||||
|
### Functionality
|
||||||
|
|
||||||
|
- **`enable_deny_patterns`**: Set to `false` to completely disable the default dangerous command blocking patterns
|
||||||
|
- **`custom_deny_patterns`**: Add custom deny regex patterns; commands matching these will be blocked
|
||||||
|
|
||||||
|
### Default Blocked Command Patterns
|
||||||
|
|
||||||
|
By default, PicoClaw blocks the following dangerous commands:
|
||||||
|
|
||||||
|
- Delete commands: `rm -rf`, `del /f/q`, `rmdir /s`
|
||||||
|
- Disk operations: `format`, `mkfs`, `diskpart`, `dd if=`, writing to `/dev/sd*`
|
||||||
|
- System operations: `shutdown`, `reboot`, `poweroff`
|
||||||
|
- Command substitution: `$()`, `${}`, backticks
|
||||||
|
- Pipe to shell: `| sh`, `| bash`
|
||||||
|
- Privilege escalation: `sudo`, `chmod`, `chown`
|
||||||
|
- Process control: `pkill`, `killall`, `kill -9`
|
||||||
|
- Remote operations: `curl | sh`, `wget | sh`, `ssh`
|
||||||
|
- Package management: `apt`, `yum`, `dnf`, `npm install -g`, `pip install --user`
|
||||||
|
- Containers: `docker run`, `docker exec`
|
||||||
|
- Git: `git push`, `git force`
|
||||||
|
- Other: `eval`, `source *.sh`
|
||||||
|
|
||||||
|
### Configuration Example
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"enable_deny_patterns": true,
|
||||||
|
"custom_deny_patterns": [
|
||||||
|
"\\brm\\s+-r\\b",
|
||||||
|
"\\bkillall\\s+python"
|
||||||
|
],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Approval Tool
|
||||||
|
|
||||||
|
The approval tool controls permissions for dangerous operations.
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|--------|------|---------|-------------|
|
||||||
|
| `enabled` | bool | true | Enable approval functionality |
|
||||||
|
| `write_file` | bool | true | Require approval for file writes |
|
||||||
|
| `edit_file` | bool | true | Require approval for file edits |
|
||||||
|
| `append_file` | bool | true | Require approval for file appends |
|
||||||
|
| `exec` | bool | true | Require approval for command execution |
|
||||||
|
| `timeout_minutes` | int | 5 | Approval timeout in minutes |
|
||||||
|
|
||||||
|
## Cron Tool
|
||||||
|
|
||||||
|
The cron tool is used for scheduling periodic tasks.
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|--------|------|---------|-------------|
|
||||||
|
| `exec_timeout_minutes` | int | 5 | Execution timeout in minutes, 0 means no limit |
|
||||||
|
|
||||||
|
## Environment Variables
|
||||||
|
|
||||||
|
All configuration options can be overridden via environment variables with the format `PICOCLAW_TOOLS_<SECTION>_<KEY>`:
|
||||||
|
|
||||||
|
For example:
|
||||||
|
- `PICOCLAW_TOOLS_WEB_BRAVE_ENABLED=true`
|
||||||
|
- `PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS=false`
|
||||||
|
- `PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES=10`
|
||||||
|
|
||||||
|
Note: Array-type environment variables are not currently supported and must be set via the config file.
|
||||||
|
|
@ -189,16 +189,7 @@ func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary str
|
||||||
systemPrompt += "\n\n## Summary of Previous Conversation\n\n" + summary
|
systemPrompt += "\n\n## Summary of Previous Conversation\n\n" + summary
|
||||||
}
|
}
|
||||||
|
|
||||||
//This fix prevents the session memory from LLM failure due to elimination of toolu_IDs required from LLM
|
history = sanitizeHistoryForProvider(history)
|
||||||
// --- INICIO DEL FIX ---
|
|
||||||
//Diegox-17
|
|
||||||
for len(history) > 0 && (history[0].Role == "tool") {
|
|
||||||
logger.DebugCF("agent", "Removing orphaned tool message from history to prevent LLM error",
|
|
||||||
map[string]interface{}{"role": history[0].Role})
|
|
||||||
history = history[1:]
|
|
||||||
}
|
|
||||||
//Diegox-17
|
|
||||||
// --- FIN DEL FIX ---
|
|
||||||
|
|
||||||
messages = append(messages, providers.Message{
|
messages = append(messages, providers.Message{
|
||||||
Role: "system",
|
Role: "system",
|
||||||
|
|
@ -207,14 +198,58 @@ func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary str
|
||||||
|
|
||||||
messages = append(messages, history...)
|
messages = append(messages, history...)
|
||||||
|
|
||||||
messages = append(messages, providers.Message{
|
if strings.TrimSpace(currentMessage) != "" {
|
||||||
Role: "user",
|
messages = append(messages, providers.Message{
|
||||||
Content: currentMessage,
|
Role: "user",
|
||||||
})
|
Content: currentMessage,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
return messages
|
return messages
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func sanitizeHistoryForProvider(history []providers.Message) []providers.Message {
|
||||||
|
if len(history) == 0 {
|
||||||
|
return history
|
||||||
|
}
|
||||||
|
|
||||||
|
sanitized := make([]providers.Message, 0, len(history))
|
||||||
|
for _, msg := range history {
|
||||||
|
switch msg.Role {
|
||||||
|
case "tool":
|
||||||
|
if len(sanitized) == 0 {
|
||||||
|
logger.DebugCF("agent", "Dropping orphaned leading tool message", map[string]interface{}{})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
last := sanitized[len(sanitized)-1]
|
||||||
|
if last.Role != "assistant" || len(last.ToolCalls) == 0 {
|
||||||
|
logger.DebugCF("agent", "Dropping orphaned tool message", map[string]interface{}{})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
sanitized = append(sanitized, msg)
|
||||||
|
|
||||||
|
case "assistant":
|
||||||
|
if len(msg.ToolCalls) > 0 {
|
||||||
|
if len(sanitized) == 0 {
|
||||||
|
logger.DebugCF("agent", "Dropping assistant tool-call turn at history start", map[string]interface{}{})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
prev := sanitized[len(sanitized)-1]
|
||||||
|
if prev.Role != "user" && prev.Role != "tool" {
|
||||||
|
logger.DebugCF("agent", "Dropping assistant tool-call turn with invalid predecessor", map[string]interface{}{"prev_role": prev.Role})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sanitized = append(sanitized, msg)
|
||||||
|
|
||||||
|
default:
|
||||||
|
sanitized = append(sanitized, msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sanitized
|
||||||
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) AddToolResult(messages []providers.Message, toolCallID, toolName, result string) []providers.Message {
|
func (cb *ContextBuilder) AddToolResult(messages []providers.Message, toolCallID, toolName, result string) []providers.Message {
|
||||||
messages = append(messages, providers.Message{
|
messages = append(messages, providers.Message{
|
||||||
Role: "tool",
|
Role: "tool",
|
||||||
|
|
|
||||||
159
pkg/agent/instance.go
Normal file
159
pkg/agent/instance.go
Normal file
|
|
@ -0,0 +1,159 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/session"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AgentInstance represents a fully configured agent with its own workspace,
|
||||||
|
// session manager, context builder, and tool registry.
|
||||||
|
type AgentInstance struct {
|
||||||
|
ID string
|
||||||
|
Name string
|
||||||
|
Model string
|
||||||
|
Fallbacks []string
|
||||||
|
Workspace string
|
||||||
|
MaxIterations int
|
||||||
|
MaxTokens int
|
||||||
|
Temperature float64
|
||||||
|
ContextWindow int
|
||||||
|
Provider providers.LLMProvider
|
||||||
|
Sessions *session.SessionManager
|
||||||
|
ContextBuilder *ContextBuilder
|
||||||
|
Tools *tools.ToolRegistry
|
||||||
|
Subagents *config.SubagentsConfig
|
||||||
|
SkillsFilter []string
|
||||||
|
Candidates []providers.FallbackCandidate
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAgentInstance creates an agent instance from config.
|
||||||
|
func NewAgentInstance(
|
||||||
|
agentCfg *config.AgentConfig,
|
||||||
|
defaults *config.AgentDefaults,
|
||||||
|
cfg *config.Config,
|
||||||
|
provider providers.LLMProvider,
|
||||||
|
) *AgentInstance {
|
||||||
|
workspace := resolveAgentWorkspace(agentCfg, defaults)
|
||||||
|
os.MkdirAll(workspace, 0755)
|
||||||
|
|
||||||
|
model := resolveAgentModel(agentCfg, defaults)
|
||||||
|
fallbacks := resolveAgentFallbacks(agentCfg, defaults)
|
||||||
|
|
||||||
|
restrict := defaults.RestrictToWorkspace
|
||||||
|
toolsRegistry := tools.NewToolRegistry()
|
||||||
|
toolsRegistry.Register(tools.NewReadFileTool(workspace, restrict))
|
||||||
|
toolsRegistry.Register(tools.NewWriteFileTool(workspace, restrict))
|
||||||
|
toolsRegistry.Register(tools.NewListDirTool(workspace, restrict))
|
||||||
|
toolsRegistry.Register(tools.NewExecToolWithConfig(workspace, restrict, cfg))
|
||||||
|
toolsRegistry.Register(tools.NewEditFileTool(workspace, restrict))
|
||||||
|
toolsRegistry.Register(tools.NewAppendFileTool(workspace, restrict))
|
||||||
|
|
||||||
|
sessionsDir := filepath.Join(workspace, "sessions")
|
||||||
|
sessionsManager := session.NewSessionManager(sessionsDir)
|
||||||
|
|
||||||
|
contextBuilder := NewContextBuilder(workspace)
|
||||||
|
contextBuilder.SetToolsRegistry(toolsRegistry)
|
||||||
|
|
||||||
|
agentID := routing.DefaultAgentID
|
||||||
|
agentName := ""
|
||||||
|
var subagents *config.SubagentsConfig
|
||||||
|
var skillsFilter []string
|
||||||
|
|
||||||
|
if agentCfg != nil {
|
||||||
|
agentID = routing.NormalizeAgentID(agentCfg.ID)
|
||||||
|
agentName = agentCfg.Name
|
||||||
|
subagents = agentCfg.Subagents
|
||||||
|
skillsFilter = agentCfg.Skills
|
||||||
|
}
|
||||||
|
|
||||||
|
maxIter := defaults.MaxToolIterations
|
||||||
|
if maxIter == 0 {
|
||||||
|
maxIter = 20
|
||||||
|
}
|
||||||
|
|
||||||
|
maxTokens := defaults.MaxTokens
|
||||||
|
if maxTokens == 0 {
|
||||||
|
maxTokens = 8192
|
||||||
|
}
|
||||||
|
|
||||||
|
temperature := 0.7
|
||||||
|
if defaults.Temperature != nil {
|
||||||
|
temperature = *defaults.Temperature
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resolve fallback candidates
|
||||||
|
modelCfg := providers.ModelConfig{
|
||||||
|
Primary: model,
|
||||||
|
Fallbacks: fallbacks,
|
||||||
|
}
|
||||||
|
candidates := providers.ResolveCandidates(modelCfg, defaults.Provider)
|
||||||
|
|
||||||
|
return &AgentInstance{
|
||||||
|
ID: agentID,
|
||||||
|
Name: agentName,
|
||||||
|
Model: model,
|
||||||
|
Fallbacks: fallbacks,
|
||||||
|
Workspace: workspace,
|
||||||
|
MaxIterations: maxIter,
|
||||||
|
MaxTokens: maxTokens,
|
||||||
|
Temperature: temperature,
|
||||||
|
ContextWindow: maxTokens,
|
||||||
|
Provider: provider,
|
||||||
|
Sessions: sessionsManager,
|
||||||
|
ContextBuilder: contextBuilder,
|
||||||
|
Tools: toolsRegistry,
|
||||||
|
Subagents: subagents,
|
||||||
|
SkillsFilter: skillsFilter,
|
||||||
|
Candidates: candidates,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveAgentWorkspace determines the workspace directory for an agent.
|
||||||
|
func resolveAgentWorkspace(agentCfg *config.AgentConfig, defaults *config.AgentDefaults) string {
|
||||||
|
if agentCfg != nil && strings.TrimSpace(agentCfg.Workspace) != "" {
|
||||||
|
return expandHome(strings.TrimSpace(agentCfg.Workspace))
|
||||||
|
}
|
||||||
|
if agentCfg == nil || agentCfg.Default || agentCfg.ID == "" || routing.NormalizeAgentID(agentCfg.ID) == "main" {
|
||||||
|
return expandHome(defaults.Workspace)
|
||||||
|
}
|
||||||
|
home, _ := os.UserHomeDir()
|
||||||
|
id := routing.NormalizeAgentID(agentCfg.ID)
|
||||||
|
return filepath.Join(home, ".picoclaw", "workspace-"+id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveAgentModel resolves the primary model for an agent.
|
||||||
|
func resolveAgentModel(agentCfg *config.AgentConfig, defaults *config.AgentDefaults) string {
|
||||||
|
if agentCfg != nil && agentCfg.Model != nil && strings.TrimSpace(agentCfg.Model.Primary) != "" {
|
||||||
|
return strings.TrimSpace(agentCfg.Model.Primary)
|
||||||
|
}
|
||||||
|
return defaults.Model
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveAgentFallbacks resolves the fallback models for an agent.
|
||||||
|
func resolveAgentFallbacks(agentCfg *config.AgentConfig, defaults *config.AgentDefaults) []string {
|
||||||
|
if agentCfg != nil && agentCfg.Model != nil && agentCfg.Model.Fallbacks != nil {
|
||||||
|
return agentCfg.Model.Fallbacks
|
||||||
|
}
|
||||||
|
return defaults.ModelFallbacks
|
||||||
|
}
|
||||||
|
|
||||||
|
func expandHome(path string) string {
|
||||||
|
if path == "" {
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
if path[0] == '~' {
|
||||||
|
home, _ := os.UserHomeDir()
|
||||||
|
if len(path) > 1 && path[1] == '/' {
|
||||||
|
return home + path[1:]
|
||||||
|
}
|
||||||
|
return home
|
||||||
|
}
|
||||||
|
return path
|
||||||
|
}
|
||||||
95
pkg/agent/instance_test.go
Normal file
95
pkg/agent/instance_test.go
Normal file
|
|
@ -0,0 +1,95 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewAgentInstance_UsesDefaultsTemperatureAndMaxTokens(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 1234,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
configuredTemp := 1.0
|
||||||
|
cfg.Agents.Defaults.Temperature = &configuredTemp
|
||||||
|
|
||||||
|
provider := &mockProvider{}
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
|
||||||
|
if agent.MaxTokens != 1234 {
|
||||||
|
t.Fatalf("MaxTokens = %d, want %d", agent.MaxTokens, 1234)
|
||||||
|
}
|
||||||
|
if agent.Temperature != 1.0 {
|
||||||
|
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 1.0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAgentInstance_DefaultsTemperatureWhenZero(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 1234,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
configuredTemp := 0.0
|
||||||
|
cfg.Agents.Defaults.Temperature = &configuredTemp
|
||||||
|
|
||||||
|
provider := &mockProvider{}
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
|
||||||
|
if agent.Temperature != 0.0 {
|
||||||
|
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAgentInstance_DefaultsTemperatureWhenUnset(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
Model: "test-model",
|
||||||
|
MaxTokens: 1234,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := &mockProvider{}
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
|
||||||
|
if agent.Temperature != 0.7 {
|
||||||
|
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.7)
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load diff
|
|
@ -14,20 +14,6 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
)
|
)
|
||||||
|
|
||||||
// mockProvider is a simple mock LLM provider for testing
|
|
||||||
type mockProvider struct{}
|
|
||||||
|
|
||||||
func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]interface{}) (*providers.LLMResponse, error) {
|
|
||||||
return &providers.LLMResponse{
|
|
||||||
Content: "Mock response",
|
|
||||||
ToolCalls: []providers.ToolCall{},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *mockProvider) GetDefaultModel() string {
|
|
||||||
return "mock-model"
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRecordLastChannel(t *testing.T) {
|
func TestRecordLastChannel(t *testing.T) {
|
||||||
// Create temp workspace
|
// Create temp workspace
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
|
@ -594,12 +580,15 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) {
|
||||||
{Role: "assistant", Content: "Old response 2"},
|
{Role: "assistant", Content: "Old response 2"},
|
||||||
{Role: "user", Content: "Trigger message"},
|
{Role: "user", Content: "Trigger message"},
|
||||||
}
|
}
|
||||||
al.sessions.SetHistory(sessionKey, history)
|
defaultAgent := al.registry.GetDefaultAgent()
|
||||||
|
if defaultAgent == nil {
|
||||||
|
t.Fatal("No default agent found")
|
||||||
|
}
|
||||||
|
defaultAgent.Sessions.SetHistory(sessionKey, history)
|
||||||
|
|
||||||
// Call ProcessDirectWithChannel
|
// Call ProcessDirectWithChannel
|
||||||
// Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration
|
// Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration
|
||||||
response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat")
|
response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat")
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Expected success after retry, got error: %v", err)
|
t.Fatalf("Expected success after retry, got error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -614,7 +603,7 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check final history length
|
// Check final history length
|
||||||
finalHistory := al.sessions.GetHistory(sessionKey)
|
finalHistory := defaultAgent.Sessions.GetHistory(sessionKey)
|
||||||
// We verify that the history has been modified (compressed)
|
// We verify that the history has been modified (compressed)
|
||||||
// Original length: 6
|
// Original length: 6
|
||||||
// Expected behavior: compression drops ~50% of history (mid slice)
|
// Expected behavior: compression drops ~50% of history (mid slice)
|
||||||
|
|
|
||||||
20
pkg/agent/mock_provider_test.go
Normal file
20
pkg/agent/mock_provider_test.go
Normal file
|
|
@ -0,0 +1,20 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mockProvider struct{}
|
||||||
|
|
||||||
|
func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]interface{}) (*providers.LLMResponse, error) {
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: "Mock response",
|
||||||
|
ToolCalls: []providers.ToolCall{},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockProvider) GetDefaultModel() string {
|
||||||
|
return "mock-model"
|
||||||
|
}
|
||||||
114
pkg/agent/registry.go
Normal file
114
pkg/agent/registry.go
Normal file
|
|
@ -0,0 +1,114 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AgentRegistry manages multiple agent instances and routes messages to them.
|
||||||
|
type AgentRegistry struct {
|
||||||
|
agents map[string]*AgentInstance
|
||||||
|
resolver *routing.RouteResolver
|
||||||
|
mu sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAgentRegistry creates a registry from config, instantiating all agents.
|
||||||
|
func NewAgentRegistry(
|
||||||
|
cfg *config.Config,
|
||||||
|
provider providers.LLMProvider,
|
||||||
|
) *AgentRegistry {
|
||||||
|
registry := &AgentRegistry{
|
||||||
|
agents: make(map[string]*AgentInstance),
|
||||||
|
resolver: routing.NewRouteResolver(cfg),
|
||||||
|
}
|
||||||
|
|
||||||
|
agentConfigs := cfg.Agents.List
|
||||||
|
if len(agentConfigs) == 0 {
|
||||||
|
implicitAgent := &config.AgentConfig{
|
||||||
|
ID: "main",
|
||||||
|
Default: true,
|
||||||
|
}
|
||||||
|
instance := NewAgentInstance(implicitAgent, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
registry.agents["main"] = instance
|
||||||
|
logger.InfoCF("agent", "Created implicit main agent (no agents.list configured)", nil)
|
||||||
|
} else {
|
||||||
|
for i := range agentConfigs {
|
||||||
|
ac := &agentConfigs[i]
|
||||||
|
id := routing.NormalizeAgentID(ac.ID)
|
||||||
|
instance := NewAgentInstance(ac, &cfg.Agents.Defaults, cfg, provider)
|
||||||
|
registry.agents[id] = instance
|
||||||
|
logger.InfoCF("agent", "Registered agent",
|
||||||
|
map[string]interface{}{
|
||||||
|
"agent_id": id,
|
||||||
|
"name": ac.Name,
|
||||||
|
"workspace": instance.Workspace,
|
||||||
|
"model": instance.Model,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return registry
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetAgent returns the agent instance for a given ID.
|
||||||
|
func (r *AgentRegistry) GetAgent(agentID string) (*AgentInstance, bool) {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
id := routing.NormalizeAgentID(agentID)
|
||||||
|
agent, ok := r.agents[id]
|
||||||
|
return agent, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveRoute determines which agent handles the message.
|
||||||
|
func (r *AgentRegistry) ResolveRoute(input routing.RouteInput) routing.ResolvedRoute {
|
||||||
|
return r.resolver.ResolveRoute(input)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListAgentIDs returns all registered agent IDs.
|
||||||
|
func (r *AgentRegistry) ListAgentIDs() []string {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
ids := make([]string, 0, len(r.agents))
|
||||||
|
for id := range r.agents {
|
||||||
|
ids = append(ids, id)
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
|
// CanSpawnSubagent checks if parentAgentID is allowed to spawn targetAgentID.
|
||||||
|
func (r *AgentRegistry) CanSpawnSubagent(parentAgentID, targetAgentID string) bool {
|
||||||
|
parent, ok := r.GetAgent(parentAgentID)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if parent.Subagents == nil || parent.Subagents.AllowAgents == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
targetNorm := routing.NormalizeAgentID(targetAgentID)
|
||||||
|
for _, allowed := range parent.Subagents.AllowAgents {
|
||||||
|
if allowed == "*" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if routing.NormalizeAgentID(allowed) == targetNorm {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDefaultAgent returns the default agent instance.
|
||||||
|
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
if agent, ok := r.agents["main"]; ok {
|
||||||
|
return agent
|
||||||
|
}
|
||||||
|
for _, agent := range r.agents {
|
||||||
|
return agent
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
199
pkg/agent/registry_test.go
Normal file
199
pkg/agent/registry_test.go
Normal file
|
|
@ -0,0 +1,199 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mockRegistryProvider struct{}
|
||||||
|
|
||||||
|
func (m *mockRegistryProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, options map[string]interface{}) (*providers.LLMResponse, error) {
|
||||||
|
return &providers.LLMResponse{Content: "mock", FinishReason: "stop"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockRegistryProvider) GetDefaultModel() string {
|
||||||
|
return "mock-model"
|
||||||
|
}
|
||||||
|
|
||||||
|
func testCfg(agents []config.AgentConfig) *config.Config {
|
||||||
|
return &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: "/tmp/picoclaw-test-registry",
|
||||||
|
Model: "gpt-4",
|
||||||
|
MaxTokens: 8192,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
List: agents,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAgentRegistry_ImplicitMain(t *testing.T) {
|
||||||
|
cfg := testCfg(nil)
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
ids := registry.ListAgentIDs()
|
||||||
|
if len(ids) != 1 || ids[0] != "main" {
|
||||||
|
t.Errorf("expected implicit main agent, got %v", ids)
|
||||||
|
}
|
||||||
|
|
||||||
|
agent, ok := registry.GetAgent("main")
|
||||||
|
if !ok || agent == nil {
|
||||||
|
t.Fatal("expected to find 'main' agent")
|
||||||
|
}
|
||||||
|
if agent.ID != "main" {
|
||||||
|
t.Errorf("agent.ID = %q, want 'main'", agent.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAgentRegistry_ExplicitAgents(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "sales", Default: true, Name: "Sales Bot"},
|
||||||
|
{ID: "support", Name: "Support Bot"},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
ids := registry.ListAgentIDs()
|
||||||
|
if len(ids) != 2 {
|
||||||
|
t.Fatalf("expected 2 agents, got %d: %v", len(ids), ids)
|
||||||
|
}
|
||||||
|
|
||||||
|
sales, ok := registry.GetAgent("sales")
|
||||||
|
if !ok || sales == nil {
|
||||||
|
t.Fatal("expected to find 'sales' agent")
|
||||||
|
}
|
||||||
|
if sales.Name != "Sales Bot" {
|
||||||
|
t.Errorf("sales.Name = %q, want 'Sales Bot'", sales.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
support, ok := registry.GetAgent("support")
|
||||||
|
if !ok || support == nil {
|
||||||
|
t.Fatal("expected to find 'support' agent")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentRegistry_GetAgent_Normalize(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "my-agent", Default: true},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
agent, ok := registry.GetAgent("My-Agent")
|
||||||
|
if !ok || agent == nil {
|
||||||
|
t.Fatal("expected to find agent with normalized ID")
|
||||||
|
}
|
||||||
|
if agent.ID != "my-agent" {
|
||||||
|
t.Errorf("agent.ID = %q, want 'my-agent'", agent.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentRegistry_GetDefaultAgent(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "alpha"},
|
||||||
|
{ID: "beta", Default: true},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
// GetDefaultAgent first checks for "main", then returns any
|
||||||
|
agent := registry.GetDefaultAgent()
|
||||||
|
if agent == nil {
|
||||||
|
t.Fatal("expected a default agent")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentRegistry_CanSpawnSubagent(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{
|
||||||
|
ID: "parent",
|
||||||
|
Default: true,
|
||||||
|
Subagents: &config.SubagentsConfig{
|
||||||
|
AllowAgents: []string{"child1", "child2"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{ID: "child1"},
|
||||||
|
{ID: "child2"},
|
||||||
|
{ID: "restricted"},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
if !registry.CanSpawnSubagent("parent", "child1") {
|
||||||
|
t.Error("expected parent to be allowed to spawn child1")
|
||||||
|
}
|
||||||
|
if !registry.CanSpawnSubagent("parent", "child2") {
|
||||||
|
t.Error("expected parent to be allowed to spawn child2")
|
||||||
|
}
|
||||||
|
if registry.CanSpawnSubagent("parent", "restricted") {
|
||||||
|
t.Error("expected parent to NOT be allowed to spawn restricted")
|
||||||
|
}
|
||||||
|
if registry.CanSpawnSubagent("child1", "child2") {
|
||||||
|
t.Error("expected child1 to NOT be allowed to spawn (no subagents config)")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentRegistry_CanSpawnSubagent_Wildcard(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{
|
||||||
|
ID: "admin",
|
||||||
|
Default: true,
|
||||||
|
Subagents: &config.SubagentsConfig{
|
||||||
|
AllowAgents: []string{"*"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{ID: "any-agent"},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
if !registry.CanSpawnSubagent("admin", "any-agent") {
|
||||||
|
t.Error("expected wildcard to allow spawning any agent")
|
||||||
|
}
|
||||||
|
if !registry.CanSpawnSubagent("admin", "nonexistent") {
|
||||||
|
t.Error("expected wildcard to allow spawning even nonexistent agents")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentInstance_Model(t *testing.T) {
|
||||||
|
model := &config.AgentModelConfig{Primary: "claude-opus"}
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "custom", Default: true, Model: model},
|
||||||
|
})
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
agent, _ := registry.GetAgent("custom")
|
||||||
|
if agent.Model != "claude-opus" {
|
||||||
|
t.Errorf("agent.Model = %q, want 'claude-opus'", agent.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentInstance_FallbackInheritance(t *testing.T) {
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "inherit", Default: true},
|
||||||
|
})
|
||||||
|
cfg.Agents.Defaults.ModelFallbacks = []string{"openai/gpt-4o-mini", "anthropic/haiku"}
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
agent, _ := registry.GetAgent("inherit")
|
||||||
|
if len(agent.Fallbacks) != 2 {
|
||||||
|
t.Errorf("expected 2 fallbacks inherited from defaults, got %d", len(agent.Fallbacks))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentInstance_FallbackExplicitEmpty(t *testing.T) {
|
||||||
|
model := &config.AgentModelConfig{
|
||||||
|
Primary: "gpt-4",
|
||||||
|
Fallbacks: []string{}, // explicitly empty = disable
|
||||||
|
}
|
||||||
|
cfg := testCfg([]config.AgentConfig{
|
||||||
|
{ID: "no-fallback", Default: true, Model: model},
|
||||||
|
})
|
||||||
|
cfg.Agents.Defaults.ModelFallbacks = []string{"should-not-inherit"}
|
||||||
|
registry := NewAgentRegistry(cfg, &mockRegistryProvider{})
|
||||||
|
|
||||||
|
agent, _ := registry.GetAgent("no-fallback")
|
||||||
|
if len(agent.Fallbacks) != 0 {
|
||||||
|
t.Errorf("expected 0 fallbacks (explicit empty), got %d: %v", len(agent.Fallbacks), agent.Fallbacks)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
|
|
@ -11,6 +12,7 @@ import (
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"runtime"
|
"runtime"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
@ -19,11 +21,13 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
type OAuthProviderConfig struct {
|
type OAuthProviderConfig struct {
|
||||||
Issuer string
|
Issuer string
|
||||||
ClientID string
|
ClientID string
|
||||||
Scopes string
|
ClientSecret string // Required for Google OAuth (confidential client)
|
||||||
Originator string
|
TokenURL string // Override token endpoint (Google uses a different URL than issuer)
|
||||||
Port int
|
Scopes string
|
||||||
|
Originator string
|
||||||
|
Port int
|
||||||
}
|
}
|
||||||
|
|
||||||
func OpenAIOAuthConfig() OAuthProviderConfig {
|
func OpenAIOAuthConfig() OAuthProviderConfig {
|
||||||
|
|
@ -36,6 +40,30 @@ func OpenAIOAuthConfig() OAuthProviderConfig {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GoogleAntigravityOAuthConfig returns the OAuth configuration for Google Cloud Code Assist (Antigravity).
|
||||||
|
// Client credentials are the same ones used by OpenCode/pi-ai for Cloud Code Assist access.
|
||||||
|
func GoogleAntigravityOAuthConfig() OAuthProviderConfig {
|
||||||
|
// These are the same client credentials used by the OpenCode antigravity plugin.
|
||||||
|
clientID := decodeBase64("MTA3MTAwNjA2MDU5MS10bWhzc2luMmgyMWxjcmUyMzV2dG9sb2poNGc0MDNlcC5hcHBzLmdvb2dsZXVzZXJjb250ZW50LmNvbQ==")
|
||||||
|
clientSecret := decodeBase64("R09DU1BYLUs1OEZXUjQ4NkxkTEoxbUxCOHNYQzR6NnFEQWY=")
|
||||||
|
return OAuthProviderConfig{
|
||||||
|
Issuer: "https://accounts.google.com/o/oauth2/v2",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: clientSecret,
|
||||||
|
Scopes: "https://www.googleapis.com/auth/cloud-platform https://www.googleapis.com/auth/userinfo.email https://www.googleapis.com/auth/userinfo.profile https://www.googleapis.com/auth/cclog https://www.googleapis.com/auth/experimentsandconfigs",
|
||||||
|
Port: 51121,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeBase64(s string) string {
|
||||||
|
data, err := base64.StdEncoding.DecodeString(s)
|
||||||
|
if err != nil {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return string(data)
|
||||||
|
}
|
||||||
|
|
||||||
func generateState() (string, error) {
|
func generateState() (string, error) {
|
||||||
buf := make([]byte, 32)
|
buf := make([]byte, 32)
|
||||||
if _, err := rand.Read(buf); err != nil {
|
if _, err := rand.Read(buf); err != nil {
|
||||||
|
|
@ -101,8 +129,17 @@ func LoginBrowser(cfg OAuthProviderConfig) (*AuthCredential, error) {
|
||||||
fmt.Printf("Could not open browser automatically.\nPlease open this URL manually:\n\n%s\n\n", authURL)
|
fmt.Printf("Could not open browser automatically.\nPlease open this URL manually:\n\n%s\n\n", authURL)
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Println("If you're running in a headless environment, use: picoclaw auth login --provider openai --device-code")
|
fmt.Printf("Wait! If you are in a headless environment (like Coolify/VPS) and cannot reach localhost:%d,\n", cfg.Port)
|
||||||
fmt.Println("Waiting for authentication in browser...")
|
fmt.Println("please complete the login in your local browser and then PASTE the final redirect URL (or just the code) here.")
|
||||||
|
fmt.Println("Waiting for authentication (browser or manual paste)...")
|
||||||
|
|
||||||
|
// Start manual input in a goroutine
|
||||||
|
manualCh := make(chan string)
|
||||||
|
go func() {
|
||||||
|
reader := bufio.NewReader(os.Stdin)
|
||||||
|
input, _ := reader.ReadString('\n')
|
||||||
|
manualCh <- strings.TrimSpace(input)
|
||||||
|
}()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case result := <-resultCh:
|
case result := <-resultCh:
|
||||||
|
|
@ -110,6 +147,22 @@ func LoginBrowser(cfg OAuthProviderConfig) (*AuthCredential, error) {
|
||||||
return nil, result.err
|
return nil, result.err
|
||||||
}
|
}
|
||||||
return exchangeCodeForTokens(cfg, result.code, pkce.CodeVerifier, redirectURI)
|
return exchangeCodeForTokens(cfg, result.code, pkce.CodeVerifier, redirectURI)
|
||||||
|
case manualInput := <-manualCh:
|
||||||
|
if manualInput == "" {
|
||||||
|
return nil, fmt.Errorf("manual input cancelled")
|
||||||
|
}
|
||||||
|
// Extract code from URL if it's a full URL
|
||||||
|
code := manualInput
|
||||||
|
if strings.Contains(manualInput, "?") {
|
||||||
|
u, err := url.Parse(manualInput)
|
||||||
|
if err == nil {
|
||||||
|
code = u.Query().Get("code")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if code == "" {
|
||||||
|
return nil, fmt.Errorf("could not find authorization code in input")
|
||||||
|
}
|
||||||
|
return exchangeCodeForTokens(cfg, code, pkce.CodeVerifier, redirectURI)
|
||||||
case <-time.After(5 * time.Minute):
|
case <-time.After(5 * time.Minute):
|
||||||
return nil, fmt.Errorf("authentication timed out after 5 minutes")
|
return nil, fmt.Errorf("authentication timed out after 5 minutes")
|
||||||
}
|
}
|
||||||
|
|
@ -269,8 +322,16 @@ func RefreshAccessToken(cred *AuthCredential, cfg OAuthProviderConfig) (*AuthCre
|
||||||
"refresh_token": {cred.RefreshToken},
|
"refresh_token": {cred.RefreshToken},
|
||||||
"scope": {"openid profile email"},
|
"scope": {"openid profile email"},
|
||||||
}
|
}
|
||||||
|
if cfg.ClientSecret != "" {
|
||||||
|
data.Set("client_secret", cfg.ClientSecret)
|
||||||
|
}
|
||||||
|
|
||||||
resp, err := http.PostForm(cfg.Issuer+"/oauth/token", data)
|
tokenURL := cfg.Issuer + "/oauth/token"
|
||||||
|
if cfg.TokenURL != "" {
|
||||||
|
tokenURL = cfg.TokenURL
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.PostForm(tokenURL, data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("refreshing token: %w", err)
|
return nil, fmt.Errorf("refreshing token: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -291,6 +352,12 @@ func RefreshAccessToken(cred *AuthCredential, cfg OAuthProviderConfig) (*AuthCre
|
||||||
if refreshed.AccountID == "" {
|
if refreshed.AccountID == "" {
|
||||||
refreshed.AccountID = cred.AccountID
|
refreshed.AccountID = cred.AccountID
|
||||||
}
|
}
|
||||||
|
if cred.Email != "" && refreshed.Email == "" {
|
||||||
|
refreshed.Email = cred.Email
|
||||||
|
}
|
||||||
|
if cred.ProjectID != "" && refreshed.ProjectID == "" {
|
||||||
|
refreshed.ProjectID = cred.ProjectID
|
||||||
|
}
|
||||||
return refreshed, nil
|
return refreshed, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -300,21 +367,35 @@ func BuildAuthorizeURL(cfg OAuthProviderConfig, pkce PKCECodes, state, redirectU
|
||||||
|
|
||||||
func buildAuthorizeURL(cfg OAuthProviderConfig, pkce PKCECodes, state, redirectURI string) string {
|
func buildAuthorizeURL(cfg OAuthProviderConfig, pkce PKCECodes, state, redirectURI string) string {
|
||||||
params := url.Values{
|
params := url.Values{
|
||||||
"response_type": {"code"},
|
"response_type": {"code"},
|
||||||
"client_id": {cfg.ClientID},
|
"client_id": {cfg.ClientID},
|
||||||
"redirect_uri": {redirectURI},
|
"redirect_uri": {redirectURI},
|
||||||
"scope": {cfg.Scopes},
|
"scope": {cfg.Scopes},
|
||||||
"code_challenge": {pkce.CodeChallenge},
|
"code_challenge": {pkce.CodeChallenge},
|
||||||
"code_challenge_method": {"S256"},
|
"code_challenge_method": {"S256"},
|
||||||
"id_token_add_organizations": {"true"},
|
"state": {state},
|
||||||
"codex_cli_simplified_flow": {"true"},
|
|
||||||
"state": {state},
|
|
||||||
}
|
}
|
||||||
if strings.Contains(strings.ToLower(cfg.Issuer), "auth.openai.com") {
|
|
||||||
params.Set("originator", "picoclaw")
|
isGoogle := strings.Contains(strings.ToLower(cfg.Issuer), "accounts.google.com")
|
||||||
|
if isGoogle {
|
||||||
|
// Google OAuth requires these for refresh token support
|
||||||
|
params.Set("access_type", "offline")
|
||||||
|
params.Set("prompt", "consent")
|
||||||
|
} else {
|
||||||
|
// OpenAI-specific parameters
|
||||||
|
params.Set("id_token_add_organizations", "true")
|
||||||
|
params.Set("codex_cli_simplified_flow", "true")
|
||||||
|
if strings.Contains(strings.ToLower(cfg.Issuer), "auth.openai.com") {
|
||||||
|
params.Set("originator", "picoclaw")
|
||||||
|
}
|
||||||
|
if cfg.Originator != "" {
|
||||||
|
params.Set("originator", cfg.Originator)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if cfg.Originator != "" {
|
|
||||||
params.Set("originator", cfg.Originator)
|
// Google uses /auth path, OpenAI uses /oauth/authorize
|
||||||
|
if isGoogle {
|
||||||
|
return cfg.Issuer + "/auth?" + params.Encode()
|
||||||
}
|
}
|
||||||
return cfg.Issuer + "/oauth/authorize?" + params.Encode()
|
return cfg.Issuer + "/oauth/authorize?" + params.Encode()
|
||||||
}
|
}
|
||||||
|
|
@ -327,8 +408,22 @@ func exchangeCodeForTokens(cfg OAuthProviderConfig, code, codeVerifier, redirect
|
||||||
"client_id": {cfg.ClientID},
|
"client_id": {cfg.ClientID},
|
||||||
"code_verifier": {codeVerifier},
|
"code_verifier": {codeVerifier},
|
||||||
}
|
}
|
||||||
|
if cfg.ClientSecret != "" {
|
||||||
|
data.Set("client_secret", cfg.ClientSecret)
|
||||||
|
}
|
||||||
|
|
||||||
resp, err := http.PostForm(cfg.Issuer+"/oauth/token", data)
|
tokenURL := cfg.Issuer + "/oauth/token"
|
||||||
|
if cfg.TokenURL != "" {
|
||||||
|
tokenURL = cfg.TokenURL
|
||||||
|
}
|
||||||
|
|
||||||
|
// Determine provider name from config
|
||||||
|
provider := "openai"
|
||||||
|
if cfg.TokenURL != "" && strings.Contains(cfg.TokenURL, "googleapis.com") {
|
||||||
|
provider = "google-antigravity"
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.PostForm(tokenURL, data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("exchanging code for tokens: %w", err)
|
return nil, fmt.Errorf("exchanging code for tokens: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -339,7 +434,7 @@ func exchangeCodeForTokens(cfg OAuthProviderConfig, code, codeVerifier, redirect
|
||||||
return nil, fmt.Errorf("token exchange failed: %s", string(body))
|
return nil, fmt.Errorf("token exchange failed: %s", string(body))
|
||||||
}
|
}
|
||||||
|
|
||||||
return parseTokenResponse(body, "openai")
|
return parseTokenResponse(body, provider)
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseTokenResponse(body []byte, provider string) (*AuthCredential, error) {
|
func parseTokenResponse(body []byte, provider string) (*AuthCredential, error) {
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,8 @@ type AuthCredential struct {
|
||||||
ExpiresAt time.Time `json:"expires_at,omitempty"`
|
ExpiresAt time.Time `json:"expires_at,omitempty"`
|
||||||
Provider string `json:"provider"`
|
Provider string `json:"provider"`
|
||||||
AuthMethod string `json:"auth_method"`
|
AuthMethod string `json:"auth_method"`
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
ProjectID string `json:"project_id,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuthStore struct {
|
type AuthStore struct {
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,6 @@ package channels
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
|
@ -87,17 +86,13 @@ func (c *BaseChannel) HandleMessage(senderID, chatID, content string, media []st
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build session key: channel:chatID
|
|
||||||
sessionKey := fmt.Sprintf("%s:%s", c.name, chatID)
|
|
||||||
|
|
||||||
msg := bus.InboundMessage{
|
msg := bus.InboundMessage{
|
||||||
Channel: c.name,
|
Channel: c.name,
|
||||||
SenderID: senderID,
|
SenderID: senderID,
|
||||||
ChatID: chatID,
|
ChatID: chatID,
|
||||||
Content: content,
|
Content: content,
|
||||||
Media: media,
|
Media: media,
|
||||||
SessionKey: sessionKey,
|
Metadata: metadata,
|
||||||
Metadata: metadata,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
c.bus.PublishInbound(msg)
|
c.bus.PublishInbound(msg)
|
||||||
|
|
|
||||||
|
|
@ -155,6 +155,14 @@ func (c *DingTalkChannel) onChatBotMessageReceived(ctx context.Context, data *ch
|
||||||
"session_webhook": data.SessionWebhook,
|
"session_webhook": data.SessionWebhook,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if data.ConversationType == "1" {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = data.ConversationId
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("dingtalk", "Received message", map[string]interface{}{
|
logger.DebugCF("dingtalk", "Received message", map[string]interface{}{
|
||||||
"sender_nick": senderNick,
|
"sender_nick": senderNick,
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bwmarrin/discordgo"
|
"github.com/bwmarrin/discordgo"
|
||||||
|
|
@ -26,6 +27,9 @@ type DiscordChannel struct {
|
||||||
config config.DiscordConfig
|
config config.DiscordConfig
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber *voice.GroqTranscriber
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
|
typingMu sync.Mutex
|
||||||
|
typingStop map[string]chan struct{} // chatID → stop signal
|
||||||
|
botUserID string // stored for mention checking
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
||||||
|
|
@ -42,6 +46,7 @@ func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordC
|
||||||
config: cfg,
|
config: cfg,
|
||||||
transcriber: nil,
|
transcriber: nil,
|
||||||
ctx: context.Background(),
|
ctx: context.Background(),
|
||||||
|
typingStop: make(map[string]chan struct{}),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -60,6 +65,14 @@ func (c *DiscordChannel) Start(ctx context.Context) error {
|
||||||
logger.InfoC("discord", "Starting Discord bot")
|
logger.InfoC("discord", "Starting Discord bot")
|
||||||
|
|
||||||
c.ctx = ctx
|
c.ctx = ctx
|
||||||
|
|
||||||
|
// Get bot user ID before opening session to avoid race condition
|
||||||
|
botUser, err := c.session.User("@me")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to get bot user: %w", err)
|
||||||
|
}
|
||||||
|
c.botUserID = botUser.ID
|
||||||
|
|
||||||
c.session.AddHandler(c.handleMessage)
|
c.session.AddHandler(c.handleMessage)
|
||||||
|
|
||||||
if err := c.session.Open(); err != nil {
|
if err := c.session.Open(); err != nil {
|
||||||
|
|
@ -68,10 +81,6 @@ func (c *DiscordChannel) Start(ctx context.Context) error {
|
||||||
|
|
||||||
c.setRunning(true)
|
c.setRunning(true)
|
||||||
|
|
||||||
botUser, err := c.session.User("@me")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to get bot user: %w", err)
|
|
||||||
}
|
|
||||||
logger.InfoCF("discord", "Discord bot connected", map[string]any{
|
logger.InfoCF("discord", "Discord bot connected", map[string]any{
|
||||||
"username": botUser.Username,
|
"username": botUser.Username,
|
||||||
"user_id": botUser.ID,
|
"user_id": botUser.ID,
|
||||||
|
|
@ -84,6 +93,14 @@ func (c *DiscordChannel) Stop(ctx context.Context) error {
|
||||||
logger.InfoC("discord", "Stopping Discord bot")
|
logger.InfoC("discord", "Stopping Discord bot")
|
||||||
c.setRunning(false)
|
c.setRunning(false)
|
||||||
|
|
||||||
|
// Stop all typing goroutines before closing session
|
||||||
|
c.typingMu.Lock()
|
||||||
|
for chatID, stop := range c.typingStop {
|
||||||
|
close(stop)
|
||||||
|
delete(c.typingStop, chatID)
|
||||||
|
}
|
||||||
|
c.typingMu.Unlock()
|
||||||
|
|
||||||
if err := c.session.Close(); err != nil {
|
if err := c.session.Close(); err != nil {
|
||||||
return fmt.Errorf("failed to close discord session: %w", err)
|
return fmt.Errorf("failed to close discord session: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -92,6 +109,8 @@ func (c *DiscordChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
c.stopTyping(msg.ChatID)
|
||||||
|
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return fmt.Errorf("discord bot not running")
|
return fmt.Errorf("discord bot not running")
|
||||||
}
|
}
|
||||||
|
|
@ -106,7 +125,7 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
chunks := splitMessage(msg.Content, 1500) // Discord has a limit of 2000 characters per message, leave 500 for natural split e.g. code blocks
|
chunks := utils.SplitMessage(msg.Content, 2000) // Split messages into chunks, Discord length limit: 2000 chars
|
||||||
|
|
||||||
for _, chunk := range chunks {
|
for _, chunk := range chunks {
|
||||||
if err := c.sendChunk(ctx, channelID, chunk); err != nil {
|
if err := c.sendChunk(ctx, channelID, chunk); err != nil {
|
||||||
|
|
@ -117,134 +136,8 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// splitMessage splits long messages into chunks, preserving code block integrity
|
|
||||||
// Uses natural boundaries (newlines, spaces) and extends messages slightly to avoid breaking code blocks
|
|
||||||
func splitMessage(content string, limit int) []string {
|
|
||||||
var messages []string
|
|
||||||
|
|
||||||
for len(content) > 0 {
|
|
||||||
if len(content) <= limit {
|
|
||||||
messages = append(messages, content)
|
|
||||||
break
|
|
||||||
}
|
|
||||||
|
|
||||||
msgEnd := limit
|
|
||||||
|
|
||||||
// Find natural split point within the limit
|
|
||||||
msgEnd = findLastNewline(content[:limit], 200)
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = findLastSpace(content[:limit], 100)
|
|
||||||
}
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = limit
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if this would end with an incomplete code block
|
|
||||||
candidate := content[:msgEnd]
|
|
||||||
unclosedIdx := findLastUnclosedCodeBlock(candidate)
|
|
||||||
|
|
||||||
if unclosedIdx >= 0 {
|
|
||||||
// Message would end with incomplete code block
|
|
||||||
// Try to extend to include the closing ``` (with some buffer)
|
|
||||||
extendedLimit := limit + 500 // Allow 500 char buffer for code blocks
|
|
||||||
if len(content) > extendedLimit {
|
|
||||||
closingIdx := findNextClosingCodeBlock(content, msgEnd)
|
|
||||||
if closingIdx > 0 && closingIdx <= extendedLimit {
|
|
||||||
// Extend to include the closing ```
|
|
||||||
msgEnd = closingIdx
|
|
||||||
} else {
|
|
||||||
// Can't find closing, split before the code block
|
|
||||||
msgEnd = findLastNewline(content[:unclosedIdx], 200)
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = findLastSpace(content[:unclosedIdx], 100)
|
|
||||||
}
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = unclosedIdx
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// Remaining content fits within extended limit
|
|
||||||
msgEnd = len(content)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if msgEnd <= 0 {
|
|
||||||
msgEnd = limit
|
|
||||||
}
|
|
||||||
|
|
||||||
messages = append(messages, content[:msgEnd])
|
|
||||||
content = strings.TrimSpace(content[msgEnd:])
|
|
||||||
}
|
|
||||||
|
|
||||||
return messages
|
|
||||||
}
|
|
||||||
|
|
||||||
// findLastUnclosedCodeBlock finds the last opening ``` that doesn't have a closing ```
|
|
||||||
// Returns the position of the opening ``` or -1 if all code blocks are complete
|
|
||||||
func findLastUnclosedCodeBlock(text string) int {
|
|
||||||
count := 0
|
|
||||||
lastOpenIdx := -1
|
|
||||||
|
|
||||||
for i := 0; i < len(text); i++ {
|
|
||||||
if i+2 < len(text) && text[i] == '`' && text[i+1] == '`' && text[i+2] == '`' {
|
|
||||||
if count == 0 {
|
|
||||||
lastOpenIdx = i
|
|
||||||
}
|
|
||||||
count++
|
|
||||||
i += 2
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// If odd number of ``` markers, last one is unclosed
|
|
||||||
if count%2 == 1 {
|
|
||||||
return lastOpenIdx
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// findNextClosingCodeBlock finds the next closing ``` starting from a position
|
|
||||||
// Returns the position after the closing ``` or -1 if not found
|
|
||||||
func findNextClosingCodeBlock(text string, startIdx int) int {
|
|
||||||
for i := startIdx; i < len(text); i++ {
|
|
||||||
if i+2 < len(text) && text[i] == '`' && text[i+1] == '`' && text[i+2] == '`' {
|
|
||||||
return i + 3
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// findLastNewline finds the last newline character within the last N characters
|
|
||||||
// Returns the position of the newline or -1 if not found
|
|
||||||
func findLastNewline(s string, searchWindow int) int {
|
|
||||||
searchStart := len(s) - searchWindow
|
|
||||||
if searchStart < 0 {
|
|
||||||
searchStart = 0
|
|
||||||
}
|
|
||||||
for i := len(s) - 1; i >= searchStart; i-- {
|
|
||||||
if s[i] == '\n' {
|
|
||||||
return i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// findLastSpace finds the last space character within the last N characters
|
|
||||||
// Returns the position of the space or -1 if not found
|
|
||||||
func findLastSpace(s string, searchWindow int) int {
|
|
||||||
searchStart := len(s) - searchWindow
|
|
||||||
if searchStart < 0 {
|
|
||||||
searchStart = 0
|
|
||||||
}
|
|
||||||
for i := len(s) - 1; i >= searchStart; i-- {
|
|
||||||
if s[i] == ' ' || s[i] == '\t' {
|
|
||||||
return i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content string) error {
|
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content string) error {
|
||||||
// 使用传入的 ctx 进行超时控制
|
// Use the passed ctx for timeout control
|
||||||
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
|
|
@ -265,7 +158,7 @@ func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content strin
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// appendContent 安全地追加内容到现有文本
|
// appendContent safely appends content to existing text
|
||||||
func appendContent(content, suffix string) string {
|
func appendContent(content, suffix string) string {
|
||||||
if content == "" {
|
if content == "" {
|
||||||
return suffix
|
return suffix
|
||||||
|
|
@ -282,13 +175,7 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.session.ChannelTyping(m.ChannelID); err != nil {
|
// Check allowlist first to avoid downloading attachments and transcribing for rejected users
|
||||||
logger.ErrorCF("discord", "Failed to send typing indicator", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// 检查白名单,避免为被拒绝的用户下载附件和转录
|
|
||||||
if !c.IsAllowed(m.Author.ID) {
|
if !c.IsAllowed(m.Author.ID) {
|
||||||
logger.DebugCF("discord", "Message rejected by allowlist", map[string]any{
|
logger.DebugCF("discord", "Message rejected by allowlist", map[string]any{
|
||||||
"user_id": m.Author.ID,
|
"user_id": m.Author.ID,
|
||||||
|
|
@ -296,6 +183,24 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// If configured to only respond to mentions, check if bot is mentioned
|
||||||
|
// Skip this check for DMs (GuildID is empty) - DMs should always be responded to
|
||||||
|
if c.config.MentionOnly && m.GuildID != "" {
|
||||||
|
isMentioned := false
|
||||||
|
for _, mention := range m.Mentions {
|
||||||
|
if mention.ID == c.botUserID {
|
||||||
|
isMentioned = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !isMentioned {
|
||||||
|
logger.DebugCF("discord", "Message ignored - bot not mentioned", map[string]any{
|
||||||
|
"user_id": m.Author.ID,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
senderID := m.Author.ID
|
senderID := m.Author.ID
|
||||||
senderName := m.Author.Username
|
senderName := m.Author.Username
|
||||||
if m.Author.Discriminator != "" && m.Author.Discriminator != "0" {
|
if m.Author.Discriminator != "" && m.Author.Discriminator != "0" {
|
||||||
|
|
@ -303,10 +208,11 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
}
|
}
|
||||||
|
|
||||||
content := m.Content
|
content := m.Content
|
||||||
|
content = c.stripBotMention(content)
|
||||||
mediaPaths := make([]string, 0, len(m.Attachments))
|
mediaPaths := make([]string, 0, len(m.Attachments))
|
||||||
localFiles := make([]string, 0, len(m.Attachments))
|
localFiles := make([]string, 0, len(m.Attachments))
|
||||||
|
|
||||||
// 确保临时文件在函数返回时被清理
|
// Ensure temp files are cleaned up when function returns
|
||||||
defer func() {
|
defer func() {
|
||||||
for _, file := range localFiles {
|
for _, file := range localFiles {
|
||||||
if err := os.Remove(file); err != nil {
|
if err := os.Remove(file); err != nil {
|
||||||
|
|
@ -330,7 +236,7 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
if c.transcriber != nil && c.transcriber.IsAvailable() {
|
if c.transcriber != nil && c.transcriber.IsAvailable() {
|
||||||
ctx, cancel := context.WithTimeout(c.getContext(), transcriptionTimeout)
|
ctx, cancel := context.WithTimeout(c.getContext(), transcriptionTimeout)
|
||||||
result, err := c.transcriber.Transcribe(ctx, localPath)
|
result, err := c.transcriber.Transcribe(ctx, localPath)
|
||||||
cancel() // 立即释放context资源,避免在for循环中泄漏
|
cancel() // Release context resources immediately to avoid leaks in for loop
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("discord", "Voice transcription failed", map[string]any{
|
logger.ErrorCF("discord", "Voice transcription failed", map[string]any{
|
||||||
|
|
@ -370,12 +276,22 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
content = "[media only]"
|
content = "[media only]"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Start typing after all early returns — guaranteed to have a matching Send()
|
||||||
|
c.startTyping(m.ChannelID)
|
||||||
|
|
||||||
logger.DebugCF("discord", "Received message", map[string]any{
|
logger.DebugCF("discord", "Received message", map[string]any{
|
||||||
"sender_name": senderName,
|
"sender_name": senderName,
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
"preview": utils.Truncate(content, 50),
|
"preview": utils.Truncate(content, 50),
|
||||||
})
|
})
|
||||||
|
|
||||||
|
peerKind := "channel"
|
||||||
|
peerID := m.ChannelID
|
||||||
|
if m.GuildID == "" {
|
||||||
|
peerKind = "direct"
|
||||||
|
peerID = senderID
|
||||||
|
}
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": m.ID,
|
"message_id": m.ID,
|
||||||
"user_id": senderID,
|
"user_id": senderID,
|
||||||
|
|
@ -384,13 +300,73 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
||||||
"guild_id": m.GuildID,
|
"guild_id": m.GuildID,
|
||||||
"channel_id": m.ChannelID,
|
"channel_id": m.ChannelID,
|
||||||
"is_dm": fmt.Sprintf("%t", m.GuildID == ""),
|
"is_dm": fmt.Sprintf("%t", m.GuildID == ""),
|
||||||
|
"peer_kind": peerKind,
|
||||||
|
"peer_id": peerID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, m.ChannelID, content, mediaPaths, metadata)
|
c.HandleMessage(senderID, m.ChannelID, content, mediaPaths, metadata)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// startTyping starts a continuous typing indicator loop for the given chatID.
|
||||||
|
// It stops any existing typing loop for that chatID before starting a new one.
|
||||||
|
func (c *DiscordChannel) startTyping(chatID string) {
|
||||||
|
c.typingMu.Lock()
|
||||||
|
// Stop existing loop for this chatID if any
|
||||||
|
if stop, ok := c.typingStop[chatID]; ok {
|
||||||
|
close(stop)
|
||||||
|
}
|
||||||
|
stop := make(chan struct{})
|
||||||
|
c.typingStop[chatID] = stop
|
||||||
|
c.typingMu.Unlock()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
if err := c.session.ChannelTyping(chatID); err != nil {
|
||||||
|
logger.DebugCF("discord", "ChannelTyping error", map[string]interface{}{"chatID": chatID, "err": err})
|
||||||
|
}
|
||||||
|
ticker := time.NewTicker(8 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
timeout := time.After(5 * time.Minute)
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-stop:
|
||||||
|
return
|
||||||
|
case <-timeout:
|
||||||
|
return
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
if err := c.session.ChannelTyping(chatID); err != nil {
|
||||||
|
logger.DebugCF("discord", "ChannelTyping error", map[string]interface{}{"chatID": chatID, "err": err})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopTyping stops the typing indicator loop for the given chatID.
|
||||||
|
func (c *DiscordChannel) stopTyping(chatID string) {
|
||||||
|
c.typingMu.Lock()
|
||||||
|
defer c.typingMu.Unlock()
|
||||||
|
if stop, ok := c.typingStop[chatID]; ok {
|
||||||
|
close(stop)
|
||||||
|
delete(c.typingStop, chatID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) downloadAttachment(url, filename string) string {
|
func (c *DiscordChannel) downloadAttachment(url, filename string) string {
|
||||||
return utils.DownloadFile(url, filename, utils.DownloadOptions{
|
return utils.DownloadFile(url, filename, utils.DownloadOptions{
|
||||||
LoggerPrefix: "discord",
|
LoggerPrefix: "discord",
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// stripBotMention removes the bot mention from the message content.
|
||||||
|
// Discord mentions have the format <@USER_ID> or <@!USER_ID> (with nickname).
|
||||||
|
func (c *DiscordChannel) stripBotMention(text string) string {
|
||||||
|
if c.botUserID == "" {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
// Remove both regular mention <@USER_ID> and nickname mention <@!USER_ID>
|
||||||
|
text = strings.ReplaceAll(text, fmt.Sprintf("<@%s>", c.botUserID), "")
|
||||||
|
text = strings.ReplaceAll(text, fmt.Sprintf("<@!%s>", c.botUserID), "")
|
||||||
|
return strings.TrimSpace(text)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -165,6 +165,15 @@ func (c *FeishuChannel) handleMessageReceive(_ context.Context, event *larkim.P2
|
||||||
metadata["tenant_key"] = *sender.TenantKey
|
metadata["tenant_key"] = *sender.TenantKey
|
||||||
}
|
}
|
||||||
|
|
||||||
|
chatType := stringValue(message.ChatType)
|
||||||
|
if chatType == "p2p" {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = chatID
|
||||||
|
}
|
||||||
|
|
||||||
logger.InfoCF("feishu", "Feishu message received", map[string]interface{}{
|
logger.InfoCF("feishu", "Feishu message received", map[string]interface{}{
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
"chat_id": chatID,
|
"chat_id": chatID,
|
||||||
|
|
|
||||||
|
|
@ -366,6 +366,14 @@ func (c *LINEChannel) processEvent(event lineEvent) {
|
||||||
"message_id": msg.ID,
|
"message_id": msg.ID,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if isGroup {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = chatID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("line", "Received message", map[string]interface{}{
|
logger.DebugCF("line", "Received message", map[string]interface{}{
|
||||||
"sender_id": senderID,
|
"sender_id": senderID,
|
||||||
"chat_id": chatID,
|
"chat_id": chatID,
|
||||||
|
|
|
||||||
|
|
@ -18,7 +18,6 @@ type MaixCamChannel struct {
|
||||||
listener net.Listener
|
listener net.Listener
|
||||||
clients map[net.Conn]bool
|
clients map[net.Conn]bool
|
||||||
clientsMux sync.RWMutex
|
clientsMux sync.RWMutex
|
||||||
running bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type MaixCamMessage struct {
|
type MaixCamMessage struct {
|
||||||
|
|
@ -35,7 +34,6 @@ func NewMaixCamChannel(cfg config.MaixCamConfig, bus *bus.MessageBus) (*MaixCamC
|
||||||
BaseChannel: base,
|
BaseChannel: base,
|
||||||
config: cfg,
|
config: cfg,
|
||||||
clients: make(map[net.Conn]bool),
|
clients: make(map[net.Conn]bool),
|
||||||
running: false,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -172,6 +170,8 @@ func (c *MaixCamChannel) handlePersonDetection(msg MaixCamMessage) {
|
||||||
"y": fmt.Sprintf("%.0f", y),
|
"y": fmt.Sprintf("%.0f", y),
|
||||||
"w": fmt.Sprintf("%.0f", w),
|
"w": fmt.Sprintf("%.0f", w),
|
||||||
"h": fmt.Sprintf("%.0f", h),
|
"h": fmt.Sprintf("%.0f", h),
|
||||||
|
"peer_kind": "channel",
|
||||||
|
"peer_id": "default",
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, []string{}, metadata)
|
c.HandleMessage(senderID, chatID, content, []string{}, metadata)
|
||||||
|
|
|
||||||
|
|
@ -4,9 +4,11 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"os"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
@ -14,20 +16,28 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/voice"
|
||||||
)
|
)
|
||||||
|
|
||||||
type OneBotChannel struct {
|
type OneBotChannel struct {
|
||||||
*BaseChannel
|
*BaseChannel
|
||||||
config config.OneBotConfig
|
config config.OneBotConfig
|
||||||
conn *websocket.Conn
|
conn *websocket.Conn
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
dedup map[string]struct{}
|
dedup map[string]struct{}
|
||||||
dedupRing []string
|
dedupRing []string
|
||||||
dedupIdx int
|
dedupIdx int
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
writeMu sync.Mutex
|
writeMu sync.Mutex
|
||||||
echoCounter int64
|
echoCounter int64
|
||||||
|
selfID int64
|
||||||
|
pending map[string]chan json.RawMessage
|
||||||
|
pendingMu sync.Mutex
|
||||||
|
transcriber *voice.GroqTranscriber
|
||||||
|
lastMessageID sync.Map
|
||||||
|
pendingEmojiMsg sync.Map
|
||||||
}
|
}
|
||||||
|
|
||||||
type oneBotRawEvent struct {
|
type oneBotRawEvent struct {
|
||||||
|
|
@ -43,9 +53,11 @@ type oneBotRawEvent struct {
|
||||||
SelfID json.RawMessage `json:"self_id"`
|
SelfID json.RawMessage `json:"self_id"`
|
||||||
Time json.RawMessage `json:"time"`
|
Time json.RawMessage `json:"time"`
|
||||||
MetaEventType string `json:"meta_event_type"`
|
MetaEventType string `json:"meta_event_type"`
|
||||||
|
NoticeType string `json:"notice_type"`
|
||||||
Echo string `json:"echo"`
|
Echo string `json:"echo"`
|
||||||
RetCode json.RawMessage `json:"retcode"`
|
RetCode json.RawMessage `json:"retcode"`
|
||||||
Status BotStatus `json:"status"`
|
Status json.RawMessage `json:"status"`
|
||||||
|
Data json.RawMessage `json:"data"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type BotStatus struct {
|
type BotStatus struct {
|
||||||
|
|
@ -53,42 +65,36 @@ type BotStatus struct {
|
||||||
Good bool `json:"good"`
|
Good bool `json:"good"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isAPIResponse(raw json.RawMessage) bool {
|
||||||
|
if len(raw) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
var s string
|
||||||
|
if json.Unmarshal(raw, &s) == nil {
|
||||||
|
return s == "ok" || s == "failed"
|
||||||
|
}
|
||||||
|
var bs BotStatus
|
||||||
|
if json.Unmarshal(raw, &bs) == nil {
|
||||||
|
return bs.Online || bs.Good
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
type oneBotSender struct {
|
type oneBotSender struct {
|
||||||
UserID json.RawMessage `json:"user_id"`
|
UserID json.RawMessage `json:"user_id"`
|
||||||
Nickname string `json:"nickname"`
|
Nickname string `json:"nickname"`
|
||||||
Card string `json:"card"`
|
Card string `json:"card"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type oneBotEvent struct {
|
|
||||||
PostType string
|
|
||||||
MessageType string
|
|
||||||
SubType string
|
|
||||||
MessageID string
|
|
||||||
UserID int64
|
|
||||||
GroupID int64
|
|
||||||
Content string
|
|
||||||
RawContent string
|
|
||||||
IsBotMentioned bool
|
|
||||||
Sender oneBotSender
|
|
||||||
SelfID int64
|
|
||||||
Time int64
|
|
||||||
MetaEventType string
|
|
||||||
}
|
|
||||||
|
|
||||||
type oneBotAPIRequest struct {
|
type oneBotAPIRequest struct {
|
||||||
Action string `json:"action"`
|
Action string `json:"action"`
|
||||||
Params interface{} `json:"params"`
|
Params interface{} `json:"params"`
|
||||||
Echo string `json:"echo,omitempty"`
|
Echo string `json:"echo,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type oneBotSendPrivateMsgParams struct {
|
type oneBotMessageSegment struct {
|
||||||
UserID int64 `json:"user_id"`
|
Type string `json:"type"`
|
||||||
Message string `json:"message"`
|
Data map[string]interface{} `json:"data"`
|
||||||
}
|
|
||||||
|
|
||||||
type oneBotSendGroupMsgParams struct {
|
|
||||||
GroupID int64 `json:"group_id"`
|
|
||||||
Message string `json:"message"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*OneBotChannel, error) {
|
func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*OneBotChannel, error) {
|
||||||
|
|
@ -101,9 +107,30 @@ func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*One
|
||||||
dedup: make(map[string]struct{}, dedupSize),
|
dedup: make(map[string]struct{}, dedupSize),
|
||||||
dedupRing: make([]string, dedupSize),
|
dedupRing: make([]string, dedupSize),
|
||||||
dedupIdx: 0,
|
dedupIdx: 0,
|
||||||
|
pending: make(map[string]chan json.RawMessage),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
|
||||||
|
c.transcriber = transcriber
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) setMsgEmojiLike(messageID string, emojiID int, set bool) {
|
||||||
|
go func() {
|
||||||
|
_, err := c.sendAPIRequest("set_msg_emoji_like", map[string]interface{}{
|
||||||
|
"message_id": messageID,
|
||||||
|
"emoji_id": emojiID,
|
||||||
|
"set": set,
|
||||||
|
}, 5*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
logger.DebugCF("onebot", "Failed to set emoji like", map[string]interface{}{
|
||||||
|
"message_id": messageID,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) Start(ctx context.Context) error {
|
func (c *OneBotChannel) Start(ctx context.Context) error {
|
||||||
if c.config.WSUrl == "" {
|
if c.config.WSUrl == "" {
|
||||||
return fmt.Errorf("OneBot ws_url not configured")
|
return fmt.Errorf("OneBot ws_url not configured")
|
||||||
|
|
@ -121,12 +148,12 @@ func (c *OneBotChannel) Start(ctx context.Context) error {
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
go c.listen()
|
go c.listen()
|
||||||
|
c.fetchSelfID()
|
||||||
}
|
}
|
||||||
|
|
||||||
if c.config.ReconnectInterval > 0 {
|
if c.config.ReconnectInterval > 0 {
|
||||||
go c.reconnectLoop()
|
go c.reconnectLoop()
|
||||||
} else {
|
} else {
|
||||||
// If reconnect is disabled but initial connection failed, we cannot recover
|
|
||||||
if c.conn == nil {
|
if c.conn == nil {
|
||||||
return fmt.Errorf("failed to connect to OneBot and reconnect is disabled")
|
return fmt.Errorf("failed to connect to OneBot and reconnect is disabled")
|
||||||
}
|
}
|
||||||
|
|
@ -152,14 +179,141 @@ func (c *OneBotChannel) connect() error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
conn.SetPongHandler(func(appData string) error {
|
||||||
|
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
c.conn = conn
|
c.conn = conn
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
go c.pinger(conn)
|
||||||
|
|
||||||
logger.InfoC("onebot", "WebSocket connected")
|
logger.InfoC("onebot", "WebSocket connected")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) pinger(conn *websocket.Conn) {
|
||||||
|
ticker := time.NewTicker(30 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
c.writeMu.Lock()
|
||||||
|
err := conn.WriteMessage(websocket.PingMessage, nil)
|
||||||
|
c.writeMu.Unlock()
|
||||||
|
if err != nil {
|
||||||
|
logger.DebugCF("onebot", "Ping write failed, stopping pinger", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) fetchSelfID() {
|
||||||
|
resp, err := c.sendAPIRequest("get_login_info", nil, 5*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("onebot", "Failed to get_login_info", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
type loginInfo struct {
|
||||||
|
UserID json.RawMessage `json:"user_id"`
|
||||||
|
Nickname string `json:"nickname"`
|
||||||
|
}
|
||||||
|
for _, extract := range []func() (*loginInfo, error){
|
||||||
|
func() (*loginInfo, error) {
|
||||||
|
var w struct {
|
||||||
|
Data loginInfo `json:"data"`
|
||||||
|
}
|
||||||
|
err := json.Unmarshal(resp, &w)
|
||||||
|
return &w.Data, err
|
||||||
|
},
|
||||||
|
func() (*loginInfo, error) {
|
||||||
|
var f loginInfo
|
||||||
|
err := json.Unmarshal(resp, &f)
|
||||||
|
return &f, err
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
info, err := extract()
|
||||||
|
if err != nil || len(info.UserID) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if uid, err := parseJSONInt64(info.UserID); err == nil && uid > 0 {
|
||||||
|
atomic.StoreInt64(&c.selfID, uid)
|
||||||
|
logger.InfoCF("onebot", "Bot self ID retrieved", map[string]interface{}{
|
||||||
|
"self_id": uid,
|
||||||
|
"nickname": info.Nickname,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.WarnCF("onebot", "Could not parse self ID from get_login_info response", map[string]interface{}{
|
||||||
|
"response": string(resp),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) sendAPIRequest(action string, params interface{}, timeout time.Duration) (json.RawMessage, error) {
|
||||||
|
c.mu.Lock()
|
||||||
|
conn := c.conn
|
||||||
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
if conn == nil {
|
||||||
|
return nil, fmt.Errorf("WebSocket not connected")
|
||||||
|
}
|
||||||
|
|
||||||
|
echo := fmt.Sprintf("api_%d_%d", time.Now().UnixNano(), atomic.AddInt64(&c.echoCounter, 1))
|
||||||
|
|
||||||
|
ch := make(chan json.RawMessage, 1)
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
c.pending[echo] = ch
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
delete(c.pending, echo)
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
}()
|
||||||
|
|
||||||
|
req := oneBotAPIRequest{
|
||||||
|
Action: action,
|
||||||
|
Params: params,
|
||||||
|
Echo: echo,
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal API request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.writeMu.Lock()
|
||||||
|
err = conn.WriteMessage(websocket.TextMessage, data)
|
||||||
|
c.writeMu.Unlock()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to write API request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case resp := <-ch:
|
||||||
|
return resp, nil
|
||||||
|
case <-time.After(timeout):
|
||||||
|
return nil, fmt.Errorf("API request %s timed out after %v", action, timeout)
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return nil, fmt.Errorf("context cancelled")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) reconnectLoop() {
|
func (c *OneBotChannel) reconnectLoop() {
|
||||||
interval := time.Duration(c.config.ReconnectInterval) * time.Second
|
interval := time.Duration(c.config.ReconnectInterval) * time.Second
|
||||||
if interval < 5*time.Second {
|
if interval < 5*time.Second {
|
||||||
|
|
@ -183,6 +337,7 @@ func (c *OneBotChannel) reconnectLoop() {
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
go c.listen()
|
go c.listen()
|
||||||
|
c.fetchSelfID()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -197,6 +352,13 @@ func (c *OneBotChannel) Stop(ctx context.Context) error {
|
||||||
c.cancel()
|
c.cancel()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
for echo, ch := range c.pending {
|
||||||
|
close(ch)
|
||||||
|
delete(c.pending, echo)
|
||||||
|
}
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
if c.conn != nil {
|
if c.conn != nil {
|
||||||
c.conn.Close()
|
c.conn.Close()
|
||||||
|
|
@ -225,10 +387,7 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
c.writeMu.Lock()
|
echo := fmt.Sprintf("send_%d", atomic.AddInt64(&c.echoCounter, 1))
|
||||||
c.echoCounter++
|
|
||||||
echo := fmt.Sprintf("send_%d", c.echoCounter)
|
|
||||||
c.writeMu.Unlock()
|
|
||||||
|
|
||||||
req := oneBotAPIRequest{
|
req := oneBotAPIRequest{
|
||||||
Action: action,
|
Action: action,
|
||||||
|
|
@ -252,67 +411,78 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if msgID, ok := c.pendingEmojiMsg.LoadAndDelete(msg.ChatID); ok {
|
||||||
|
if mid, ok := msgID.(string); ok && mid != "" {
|
||||||
|
c.setMsgEmojiLike(mid, 289, false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) buildMessageSegments(chatID, content string) []oneBotMessageSegment {
|
||||||
|
var segments []oneBotMessageSegment
|
||||||
|
|
||||||
|
if lastMsgID, ok := c.lastMessageID.Load(chatID); ok {
|
||||||
|
if msgID, ok := lastMsgID.(string); ok && msgID != "" {
|
||||||
|
segments = append(segments, oneBotMessageSegment{
|
||||||
|
Type: "reply",
|
||||||
|
Data: map[string]interface{}{"id": msgID},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
segments = append(segments, oneBotMessageSegment{
|
||||||
|
Type: "text",
|
||||||
|
Data: map[string]interface{}{"text": content},
|
||||||
|
})
|
||||||
|
|
||||||
|
return segments
|
||||||
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) buildSendRequest(msg bus.OutboundMessage) (string, interface{}, error) {
|
func (c *OneBotChannel) buildSendRequest(msg bus.OutboundMessage) (string, interface{}, error) {
|
||||||
chatID := msg.ChatID
|
chatID := msg.ChatID
|
||||||
|
segments := c.buildMessageSegments(chatID, msg.Content)
|
||||||
|
|
||||||
if len(chatID) > 6 && chatID[:6] == "group:" {
|
var action, idKey string
|
||||||
groupID, err := strconv.ParseInt(chatID[6:], 10, 64)
|
var rawID string
|
||||||
if err != nil {
|
if rest, ok := strings.CutPrefix(chatID, "group:"); ok {
|
||||||
return "", nil, fmt.Errorf("invalid group ID in chatID: %s", chatID)
|
action, idKey, rawID = "send_group_msg", "group_id", rest
|
||||||
}
|
} else if rest, ok := strings.CutPrefix(chatID, "private:"); ok {
|
||||||
return "send_group_msg", oneBotSendGroupMsgParams{
|
action, idKey, rawID = "send_private_msg", "user_id", rest
|
||||||
GroupID: groupID,
|
} else {
|
||||||
Message: msg.Content,
|
action, idKey, rawID = "send_private_msg", "user_id", chatID
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(chatID) > 8 && chatID[:8] == "private:" {
|
id, err := strconv.ParseInt(rawID, 10, 64)
|
||||||
userID, err := strconv.ParseInt(chatID[8:], 10, 64)
|
|
||||||
if err != nil {
|
|
||||||
return "", nil, fmt.Errorf("invalid user ID in chatID: %s", chatID)
|
|
||||||
}
|
|
||||||
return "send_private_msg", oneBotSendPrivateMsgParams{
|
|
||||||
UserID: userID,
|
|
||||||
Message: msg.Content,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
userID, err := strconv.ParseInt(chatID, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", nil, fmt.Errorf("invalid chatID for OneBot: %s", chatID)
|
return "", nil, fmt.Errorf("invalid %s in chatID: %s", idKey, chatID)
|
||||||
}
|
}
|
||||||
|
return action, map[string]interface{}{idKey: id, "message": segments}, nil
|
||||||
return "send_private_msg", oneBotSendPrivateMsgParams{
|
|
||||||
UserID: userID,
|
|
||||||
Message: msg.Content,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) listen() {
|
func (c *OneBotChannel) listen() {
|
||||||
|
c.mu.Lock()
|
||||||
|
conn := c.conn
|
||||||
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
if conn == nil {
|
||||||
|
logger.WarnC("onebot", "WebSocket connection is nil, listener exiting")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-c.ctx.Done():
|
case <-c.ctx.Done():
|
||||||
return
|
return
|
||||||
default:
|
default:
|
||||||
c.mu.Lock()
|
|
||||||
conn := c.conn
|
|
||||||
c.mu.Unlock()
|
|
||||||
|
|
||||||
if conn == nil {
|
|
||||||
logger.WarnC("onebot", "WebSocket connection is nil, listener exiting")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
_, message, err := conn.ReadMessage()
|
_, message, err := conn.ReadMessage()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("onebot", "WebSocket read error", map[string]interface{}{
|
logger.ErrorCF("onebot", "WebSocket read error", map[string]interface{}{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
if c.conn != nil {
|
if c.conn == conn {
|
||||||
c.conn.Close()
|
c.conn.Close()
|
||||||
c.conn = nil
|
c.conn = nil
|
||||||
}
|
}
|
||||||
|
|
@ -320,10 +490,7 @@ func (c *OneBotChannel) listen() {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("onebot", "Raw WebSocket message received", map[string]interface{}{
|
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||||
"length": len(message),
|
|
||||||
"payload": string(message),
|
|
||||||
})
|
|
||||||
|
|
||||||
var raw oneBotRawEvent
|
var raw oneBotRawEvent
|
||||||
if err := json.Unmarshal(message, &raw); err != nil {
|
if err := json.Unmarshal(message, &raw); err != nil {
|
||||||
|
|
@ -334,20 +501,37 @@ func (c *OneBotChannel) listen() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if raw.Echo != "" || raw.Status.Online || raw.Status.Good {
|
logger.DebugCF("onebot", "WebSocket event", map[string]interface{}{
|
||||||
logger.DebugCF("onebot", "Received API response, skipping", map[string]interface{}{
|
"length": len(message),
|
||||||
"echo": raw.Echo,
|
"post_type": raw.PostType,
|
||||||
"status": raw.Status,
|
"sub_type": raw.SubType,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
if raw.Echo != "" {
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
ch, ok := c.pending[raw.Echo]
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
|
||||||
|
if ok {
|
||||||
|
select {
|
||||||
|
case ch <- message:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.DebugCF("onebot", "Received API response (no waiter)", map[string]interface{}{
|
||||||
|
"echo": raw.Echo,
|
||||||
|
"status": string(raw.Status),
|
||||||
|
})
|
||||||
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("onebot", "Parsed raw event", map[string]interface{}{
|
if isAPIResponse(raw.Status) {
|
||||||
"post_type": raw.PostType,
|
logger.DebugCF("onebot", "Received API response without echo, skipping", map[string]interface{}{
|
||||||
"message_type": raw.MessageType,
|
"status": string(raw.Status),
|
||||||
"sub_type": raw.SubType,
|
})
|
||||||
"meta_event_type": raw.MetaEventType,
|
continue
|
||||||
})
|
}
|
||||||
|
|
||||||
c.handleRawEvent(&raw)
|
c.handleRawEvent(&raw)
|
||||||
}
|
}
|
||||||
|
|
@ -386,9 +570,12 @@ func parseJSONString(raw json.RawMessage) string {
|
||||||
type parseMessageResult struct {
|
type parseMessageResult struct {
|
||||||
Text string
|
Text string
|
||||||
IsBotMentioned bool
|
IsBotMentioned bool
|
||||||
|
Media []string
|
||||||
|
LocalFiles []string
|
||||||
|
ReplyTo string
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseMessageContentEx(raw json.RawMessage, selfID int64) parseMessageResult {
|
func (c *OneBotChannel) parseMessageSegments(raw json.RawMessage, selfID int64) parseMessageResult {
|
||||||
if len(raw) == 0 {
|
if len(raw) == 0 {
|
||||||
return parseMessageResult{}
|
return parseMessageResult{}
|
||||||
}
|
}
|
||||||
|
|
@ -408,60 +595,155 @@ func parseMessageContentEx(raw json.RawMessage, selfID int64) parseMessageResult
|
||||||
}
|
}
|
||||||
|
|
||||||
var segments []map[string]interface{}
|
var segments []map[string]interface{}
|
||||||
if err := json.Unmarshal(raw, &segments); err == nil {
|
if err := json.Unmarshal(raw, &segments); err != nil {
|
||||||
var text string
|
return parseMessageResult{}
|
||||||
mentioned := false
|
}
|
||||||
selfIDStr := strconv.FormatInt(selfID, 10)
|
|
||||||
for _, seg := range segments {
|
var textParts []string
|
||||||
segType, _ := seg["type"].(string)
|
mentioned := false
|
||||||
data, _ := seg["data"].(map[string]interface{})
|
selfIDStr := strconv.FormatInt(selfID, 10)
|
||||||
switch segType {
|
var media []string
|
||||||
case "text":
|
var localFiles []string
|
||||||
if data != nil {
|
var replyTo string
|
||||||
if t, ok := data["text"].(string); ok {
|
|
||||||
text += t
|
for _, seg := range segments {
|
||||||
}
|
segType, _ := seg["type"].(string)
|
||||||
|
data, _ := seg["data"].(map[string]interface{})
|
||||||
|
|
||||||
|
switch segType {
|
||||||
|
case "text":
|
||||||
|
if data != nil {
|
||||||
|
if t, ok := data["text"].(string); ok {
|
||||||
|
textParts = append(textParts, t)
|
||||||
}
|
}
|
||||||
case "at":
|
}
|
||||||
if data != nil && selfID > 0 {
|
|
||||||
qqVal := fmt.Sprintf("%v", data["qq"])
|
case "at":
|
||||||
if qqVal == selfIDStr || qqVal == "all" {
|
if data != nil && selfID > 0 {
|
||||||
mentioned = true
|
qqVal := fmt.Sprintf("%v", data["qq"])
|
||||||
|
if qqVal == selfIDStr || qqVal == "all" {
|
||||||
|
mentioned = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case "image", "video", "file":
|
||||||
|
if data != nil {
|
||||||
|
url, _ := data["url"].(string)
|
||||||
|
if url != "" {
|
||||||
|
defaults := map[string]string{"image": "image.jpg", "video": "video.mp4", "file": "file"}
|
||||||
|
filename := defaults[segType]
|
||||||
|
if f, ok := data["file"].(string); ok && f != "" {
|
||||||
|
filename = f
|
||||||
|
} else if n, ok := data["name"].(string); ok && n != "" {
|
||||||
|
filename = n
|
||||||
|
}
|
||||||
|
localPath := utils.DownloadFile(url, filename, utils.DownloadOptions{
|
||||||
|
LoggerPrefix: "onebot",
|
||||||
|
})
|
||||||
|
if localPath != "" {
|
||||||
|
media = append(media, localPath)
|
||||||
|
localFiles = append(localFiles, localPath)
|
||||||
|
textParts = append(textParts, fmt.Sprintf("[%s]", segType))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
case "record":
|
||||||
|
if data != nil {
|
||||||
|
url, _ := data["url"].(string)
|
||||||
|
if url != "" {
|
||||||
|
localPath := utils.DownloadFile(url, "voice.amr", utils.DownloadOptions{
|
||||||
|
LoggerPrefix: "onebot",
|
||||||
|
})
|
||||||
|
if localPath != "" {
|
||||||
|
localFiles = append(localFiles, localPath)
|
||||||
|
if c.transcriber != nil && c.transcriber.IsAvailable() {
|
||||||
|
tctx, tcancel := context.WithTimeout(c.ctx, 30*time.Second)
|
||||||
|
result, err := c.transcriber.Transcribe(tctx, localPath)
|
||||||
|
tcancel()
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("onebot", "Voice transcription failed", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
textParts = append(textParts, "[voice (transcription failed)]")
|
||||||
|
media = append(media, localPath)
|
||||||
|
} else {
|
||||||
|
textParts = append(textParts, fmt.Sprintf("[voice transcription: %s]", result.Text))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
textParts = append(textParts, "[voice]")
|
||||||
|
media = append(media, localPath)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case "reply":
|
||||||
|
if data != nil {
|
||||||
|
if id, ok := data["id"]; ok {
|
||||||
|
replyTo = fmt.Sprintf("%v", id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case "face":
|
||||||
|
if data != nil {
|
||||||
|
faceID, _ := data["id"]
|
||||||
|
textParts = append(textParts, fmt.Sprintf("[face:%v]", faceID))
|
||||||
|
}
|
||||||
|
|
||||||
|
case "forward":
|
||||||
|
textParts = append(textParts, "[forward message]")
|
||||||
|
|
||||||
|
default:
|
||||||
|
|
||||||
}
|
}
|
||||||
return parseMessageResult{Text: strings.TrimSpace(text), IsBotMentioned: mentioned}
|
|
||||||
}
|
}
|
||||||
return parseMessageResult{}
|
|
||||||
|
return parseMessageResult{
|
||||||
|
Text: strings.TrimSpace(strings.Join(textParts, "")),
|
||||||
|
IsBotMentioned: mentioned,
|
||||||
|
Media: media,
|
||||||
|
LocalFiles: localFiles,
|
||||||
|
ReplyTo: replyTo,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) handleRawEvent(raw *oneBotRawEvent) {
|
func (c *OneBotChannel) handleRawEvent(raw *oneBotRawEvent) {
|
||||||
switch raw.PostType {
|
switch raw.PostType {
|
||||||
case "message":
|
case "message":
|
||||||
evt, err := c.normalizeMessageEvent(raw)
|
if userID, err := parseJSONInt64(raw.UserID); err == nil && userID > 0 {
|
||||||
if err != nil {
|
if !c.IsAllowed(strconv.FormatInt(userID, 10)) {
|
||||||
logger.WarnCF("onebot", "Failed to normalize message event", map[string]interface{}{
|
logger.DebugCF("onebot", "Message rejected by allowlist", map[string]interface{}{
|
||||||
"error": err.Error(),
|
"user_id": userID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
c.handleMessage(evt)
|
c.handleMessage(raw)
|
||||||
|
|
||||||
|
case "message_sent":
|
||||||
|
logger.DebugCF("onebot", "Bot sent message event", map[string]interface{}{
|
||||||
|
"message_type": raw.MessageType,
|
||||||
|
"message_id": parseJSONString(raw.MessageID),
|
||||||
|
})
|
||||||
|
|
||||||
case "meta_event":
|
case "meta_event":
|
||||||
c.handleMetaEvent(raw)
|
c.handleMetaEvent(raw)
|
||||||
|
|
||||||
case "notice":
|
case "notice":
|
||||||
logger.DebugCF("onebot", "Notice event received", map[string]interface{}{
|
c.handleNoticeEvent(raw)
|
||||||
"sub_type": raw.SubType,
|
|
||||||
})
|
|
||||||
case "request":
|
case "request":
|
||||||
logger.DebugCF("onebot", "Request event received", map[string]interface{}{
|
logger.DebugCF("onebot", "Request event received", map[string]interface{}{
|
||||||
"sub_type": raw.SubType,
|
"sub_type": raw.SubType,
|
||||||
})
|
})
|
||||||
|
|
||||||
case "":
|
case "":
|
||||||
logger.DebugCF("onebot", "Event with empty post_type (possibly API response)", map[string]interface{}{
|
logger.DebugCF("onebot", "Event with empty post_type (possibly API response)", map[string]interface{}{
|
||||||
"echo": raw.Echo,
|
"echo": raw.Echo,
|
||||||
"status": raw.Status,
|
"status": raw.Status,
|
||||||
})
|
})
|
||||||
|
|
||||||
default:
|
default:
|
||||||
logger.DebugCF("onebot", "Unknown post_type", map[string]interface{}{
|
logger.DebugCF("onebot", "Unknown post_type", map[string]interface{}{
|
||||||
"post_type": raw.PostType,
|
"post_type": raw.PostType,
|
||||||
|
|
@ -469,18 +751,51 @@ func (c *OneBotChannel) handleRawEvent(raw *oneBotRawEvent) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) normalizeMessageEvent(raw *oneBotRawEvent) (*oneBotEvent, error) {
|
func (c *OneBotChannel) handleMetaEvent(raw *oneBotRawEvent) {
|
||||||
|
if raw.MetaEventType == "lifecycle" {
|
||||||
|
logger.InfoCF("onebot", "Lifecycle event", map[string]interface{}{"sub_type": raw.SubType})
|
||||||
|
} else if raw.MetaEventType != "heartbeat" {
|
||||||
|
logger.DebugCF("onebot", "Meta event: "+raw.MetaEventType, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) handleNoticeEvent(raw *oneBotRawEvent) {
|
||||||
|
fields := map[string]interface{}{
|
||||||
|
"notice_type": raw.NoticeType,
|
||||||
|
"sub_type": raw.SubType,
|
||||||
|
"group_id": parseJSONString(raw.GroupID),
|
||||||
|
"user_id": parseJSONString(raw.UserID),
|
||||||
|
"message_id": parseJSONString(raw.MessageID),
|
||||||
|
}
|
||||||
|
switch raw.NoticeType {
|
||||||
|
case "group_recall", "group_increase", "group_decrease",
|
||||||
|
"friend_add", "group_admin", "group_ban":
|
||||||
|
logger.InfoCF("onebot", "Notice: "+raw.NoticeType, fields)
|
||||||
|
default:
|
||||||
|
logger.DebugCF("onebot", "Notice: "+raw.NoticeType, fields)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *OneBotChannel) handleMessage(raw *oneBotRawEvent) {
|
||||||
|
// Parse fields from raw event
|
||||||
userID, err := parseJSONInt64(raw.UserID)
|
userID, err := parseJSONInt64(raw.UserID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("parse user_id: %w (raw: %s)", err, string(raw.UserID))
|
logger.WarnCF("onebot", "Failed to parse user_id", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
"raw": string(raw.UserID),
|
||||||
|
})
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
groupID, _ := parseJSONInt64(raw.GroupID)
|
groupID, _ := parseJSONInt64(raw.GroupID)
|
||||||
selfID, _ := parseJSONInt64(raw.SelfID)
|
selfID, _ := parseJSONInt64(raw.SelfID)
|
||||||
ts, _ := parseJSONInt64(raw.Time)
|
|
||||||
messageID := parseJSONString(raw.MessageID)
|
messageID := parseJSONString(raw.MessageID)
|
||||||
|
|
||||||
parsed := parseMessageContentEx(raw.Message, selfID)
|
if selfID == 0 {
|
||||||
|
selfID = atomic.LoadInt64(&c.selfID)
|
||||||
|
}
|
||||||
|
|
||||||
|
parsed := c.parseMessageSegments(raw.Message, selfID)
|
||||||
isBotMentioned := parsed.IsBotMentioned
|
isBotMentioned := parsed.IsBotMentioned
|
||||||
|
|
||||||
content := raw.RawMessage
|
content := raw.RawMessage
|
||||||
|
|
@ -495,6 +810,10 @@ func (c *OneBotChannel) normalizeMessageEvent(raw *oneBotRawEvent) (*oneBotEvent
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if parsed.Text != "" && content != parsed.Text && (len(parsed.Media) > 0 || parsed.ReplyTo != "") {
|
||||||
|
content = parsed.Text
|
||||||
|
}
|
||||||
|
|
||||||
var sender oneBotSender
|
var sender oneBotSender
|
||||||
if len(raw.Sender) > 0 {
|
if len(raw.Sender) > 0 {
|
||||||
if err := json.Unmarshal(raw.Sender, &sender); err != nil {
|
if err := json.Unmarshal(raw.Sender, &sender); err != nil {
|
||||||
|
|
@ -505,137 +824,111 @@ func (c *OneBotChannel) normalizeMessageEvent(raw *oneBotRawEvent) (*oneBotEvent
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("onebot", "Normalized message event", map[string]interface{}{
|
// Clean up temp files when done
|
||||||
"message_type": raw.MessageType,
|
if len(parsed.LocalFiles) > 0 {
|
||||||
"user_id": userID,
|
defer func() {
|
||||||
"group_id": groupID,
|
for _, f := range parsed.LocalFiles {
|
||||||
"message_id": messageID,
|
if err := os.Remove(f); err != nil {
|
||||||
"content_len": len(content),
|
logger.DebugCF("onebot", "Failed to remove temp file", map[string]interface{}{
|
||||||
"nickname": sender.Nickname,
|
"path": f,
|
||||||
})
|
"error": err.Error(),
|
||||||
|
})
|
||||||
return &oneBotEvent{
|
}
|
||||||
PostType: raw.PostType,
|
}
|
||||||
MessageType: raw.MessageType,
|
}()
|
||||||
SubType: raw.SubType,
|
|
||||||
MessageID: messageID,
|
|
||||||
UserID: userID,
|
|
||||||
GroupID: groupID,
|
|
||||||
Content: content,
|
|
||||||
RawContent: raw.RawMessage,
|
|
||||||
IsBotMentioned: isBotMentioned,
|
|
||||||
Sender: sender,
|
|
||||||
SelfID: selfID,
|
|
||||||
Time: ts,
|
|
||||||
MetaEventType: raw.MetaEventType,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *OneBotChannel) handleMetaEvent(raw *oneBotRawEvent) {
|
|
||||||
switch raw.MetaEventType {
|
|
||||||
case "lifecycle":
|
|
||||||
logger.InfoCF("onebot", "Lifecycle event", map[string]interface{}{
|
|
||||||
"sub_type": raw.SubType,
|
|
||||||
})
|
|
||||||
case "heartbeat":
|
|
||||||
logger.DebugC("onebot", "Heartbeat received")
|
|
||||||
default:
|
|
||||||
logger.DebugCF("onebot", "Unknown meta_event_type", map[string]interface{}{
|
|
||||||
"meta_event_type": raw.MetaEventType,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
func (c *OneBotChannel) handleMessage(evt *oneBotEvent) {
|
if c.isDuplicate(messageID) {
|
||||||
if c.isDuplicate(evt.MessageID) {
|
|
||||||
logger.DebugCF("onebot", "Duplicate message, skipping", map[string]interface{}{
|
logger.DebugCF("onebot", "Duplicate message, skipping", map[string]interface{}{
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
content := evt.Content
|
|
||||||
if content == "" {
|
if content == "" {
|
||||||
logger.DebugCF("onebot", "Received empty message, ignoring", map[string]interface{}{
|
logger.DebugCF("onebot", "Received empty message, ignoring", map[string]interface{}{
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
senderID := strconv.FormatInt(evt.UserID, 10)
|
senderID := strconv.FormatInt(userID, 10)
|
||||||
var chatID string
|
var chatID string
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
}
|
}
|
||||||
|
|
||||||
switch evt.MessageType {
|
if parsed.ReplyTo != "" {
|
||||||
|
metadata["reply_to_message_id"] = parsed.ReplyTo
|
||||||
|
}
|
||||||
|
|
||||||
|
switch raw.MessageType {
|
||||||
case "private":
|
case "private":
|
||||||
chatID = "private:" + senderID
|
chatID = "private:" + senderID
|
||||||
logger.InfoCF("onebot", "Received private message", map[string]interface{}{
|
metadata["peer_kind"] = "direct"
|
||||||
"sender": senderID,
|
metadata["peer_id"] = senderID
|
||||||
"message_id": evt.MessageID,
|
|
||||||
"length": len(content),
|
|
||||||
"content": truncate(content, 100),
|
|
||||||
})
|
|
||||||
|
|
||||||
case "group":
|
case "group":
|
||||||
groupIDStr := strconv.FormatInt(evt.GroupID, 10)
|
groupIDStr := strconv.FormatInt(groupID, 10)
|
||||||
chatID = "group:" + groupIDStr
|
chatID = "group:" + groupIDStr
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = groupIDStr
|
||||||
metadata["group_id"] = groupIDStr
|
metadata["group_id"] = groupIDStr
|
||||||
|
|
||||||
senderUserID, _ := parseJSONInt64(evt.Sender.UserID)
|
senderUserID, _ := parseJSONInt64(sender.UserID)
|
||||||
if senderUserID > 0 {
|
if senderUserID > 0 {
|
||||||
metadata["sender_user_id"] = strconv.FormatInt(senderUserID, 10)
|
metadata["sender_user_id"] = strconv.FormatInt(senderUserID, 10)
|
||||||
}
|
}
|
||||||
|
|
||||||
if evt.Sender.Card != "" {
|
if sender.Card != "" {
|
||||||
metadata["sender_name"] = evt.Sender.Card
|
metadata["sender_name"] = sender.Card
|
||||||
} else if evt.Sender.Nickname != "" {
|
} else if sender.Nickname != "" {
|
||||||
metadata["sender_name"] = evt.Sender.Nickname
|
metadata["sender_name"] = sender.Nickname
|
||||||
}
|
}
|
||||||
|
|
||||||
triggered, strippedContent := c.checkGroupTrigger(content, evt.IsBotMentioned)
|
triggered, strippedContent := c.checkGroupTrigger(content, isBotMentioned)
|
||||||
if !triggered {
|
if !triggered {
|
||||||
logger.DebugCF("onebot", "Group message ignored (no trigger)", map[string]interface{}{
|
logger.DebugCF("onebot", "Group message ignored (no trigger)", map[string]interface{}{
|
||||||
"sender": senderID,
|
"sender": senderID,
|
||||||
"group": groupIDStr,
|
"group": groupIDStr,
|
||||||
"is_mentioned": evt.IsBotMentioned,
|
"is_mentioned": isBotMentioned,
|
||||||
"content": truncate(content, 100),
|
"content": truncate(content, 100),
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
content = strippedContent
|
content = strippedContent
|
||||||
|
|
||||||
logger.InfoCF("onebot", "Received group message", map[string]interface{}{
|
|
||||||
"sender": senderID,
|
|
||||||
"group": groupIDStr,
|
|
||||||
"message_id": evt.MessageID,
|
|
||||||
"is_mentioned": evt.IsBotMentioned,
|
|
||||||
"length": len(content),
|
|
||||||
"content": truncate(content, 100),
|
|
||||||
})
|
|
||||||
|
|
||||||
default:
|
default:
|
||||||
logger.WarnCF("onebot", "Unknown message type, cannot route", map[string]interface{}{
|
logger.WarnCF("onebot", "Unknown message type, cannot route", map[string]interface{}{
|
||||||
"type": evt.MessageType,
|
"type": raw.MessageType,
|
||||||
"message_id": evt.MessageID,
|
"message_id": messageID,
|
||||||
"user_id": evt.UserID,
|
"user_id": userID,
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if evt.Sender.Nickname != "" {
|
logger.InfoCF("onebot", "Received "+raw.MessageType+" message", map[string]interface{}{
|
||||||
metadata["nickname"] = evt.Sender.Nickname
|
"sender": senderID,
|
||||||
}
|
"chat_id": chatID,
|
||||||
|
"message_id": messageID,
|
||||||
logger.DebugCF("onebot", "Forwarding message to bus", map[string]interface{}{
|
"length": len(content),
|
||||||
"sender_id": senderID,
|
"content": truncate(content, 100),
|
||||||
"chat_id": chatID,
|
"media_count": len(parsed.Media),
|
||||||
"content": truncate(content, 100),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, []string{}, metadata)
|
if sender.Nickname != "" {
|
||||||
|
metadata["nickname"] = sender.Nickname
|
||||||
|
}
|
||||||
|
|
||||||
|
c.lastMessageID.Store(chatID, messageID)
|
||||||
|
|
||||||
|
if raw.MessageType == "group" && messageID != "" && messageID != "0" {
|
||||||
|
c.setMsgEmojiLike(messageID, 289, true)
|
||||||
|
c.pendingEmojiMsg.Store(chatID, messageID)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.HandleMessage(senderID, chatID, content, parsed.Media, metadata)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) isDuplicate(messageID string) bool {
|
func (c *OneBotChannel) isDuplicate(messageID string) bool {
|
||||||
|
|
|
||||||
|
|
@ -165,6 +165,8 @@ func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
// 转发到消息总线
|
// 转发到消息总线
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": data.ID,
|
"message_id": data.ID,
|
||||||
|
"peer_kind": "direct",
|
||||||
|
"peer_id": senderID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, senderID, content, []string{}, metadata)
|
c.HandleMessage(senderID, senderID, content, []string{}, metadata)
|
||||||
|
|
@ -207,6 +209,8 @@ func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": data.ID,
|
"message_id": data.ID,
|
||||||
"group_id": data.GroupID,
|
"group_id": data.GroupID,
|
||||||
|
"peer_kind": "group",
|
||||||
|
"peer_id": data.GroupID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, data.GroupID, content, []string{}, metadata)
|
c.HandleMessage(senderID, data.GroupID, content, []string{}, metadata)
|
||||||
|
|
|
||||||
|
|
@ -25,6 +25,7 @@ type SlackChannel struct {
|
||||||
api *slack.Client
|
api *slack.Client
|
||||||
socketClient *socketmode.Client
|
socketClient *socketmode.Client
|
||||||
botUserID string
|
botUserID string
|
||||||
|
teamID string
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber *voice.GroqTranscriber
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
|
|
@ -72,6 +73,7 @@ func (c *SlackChannel) Start(ctx context.Context) error {
|
||||||
return fmt.Errorf("slack auth test failed: %w", err)
|
return fmt.Errorf("slack auth test failed: %w", err)
|
||||||
}
|
}
|
||||||
c.botUserID = authResp.UserID
|
c.botUserID = authResp.UserID
|
||||||
|
c.teamID = authResp.TeamID
|
||||||
|
|
||||||
logger.InfoCF("slack", "Slack bot connected", map[string]interface{}{
|
logger.InfoCF("slack", "Slack bot connected", map[string]interface{}{
|
||||||
"bot_user_id": c.botUserID,
|
"bot_user_id": c.botUserID,
|
||||||
|
|
@ -274,11 +276,21 @@ func (c *SlackChannel) handleMessageEvent(ev *slackevents.MessageEvent) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
peerKind := "channel"
|
||||||
|
peerID := channelID
|
||||||
|
if strings.HasPrefix(channelID, "D") {
|
||||||
|
peerKind = "direct"
|
||||||
|
peerID = senderID
|
||||||
|
}
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_ts": messageTS,
|
"message_ts": messageTS,
|
||||||
"channel_id": channelID,
|
"channel_id": channelID,
|
||||||
"thread_ts": threadTS,
|
"thread_ts": threadTS,
|
||||||
"platform": "slack",
|
"platform": "slack",
|
||||||
|
"peer_kind": peerKind,
|
||||||
|
"peer_id": peerID,
|
||||||
|
"team_id": c.teamID,
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("slack", "Received message", map[string]interface{}{
|
logger.DebugCF("slack", "Received message", map[string]interface{}{
|
||||||
|
|
@ -331,12 +343,22 @@ func (c *SlackChannel) handleAppMention(ev *slackevents.AppMentionEvent) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
mentionPeerKind := "channel"
|
||||||
|
mentionPeerID := channelID
|
||||||
|
if strings.HasPrefix(channelID, "D") {
|
||||||
|
mentionPeerKind = "direct"
|
||||||
|
mentionPeerID = senderID
|
||||||
|
}
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_ts": messageTS,
|
"message_ts": messageTS,
|
||||||
"channel_id": channelID,
|
"channel_id": channelID,
|
||||||
"thread_ts": threadTS,
|
"thread_ts": threadTS,
|
||||||
"platform": "slack",
|
"platform": "slack",
|
||||||
"is_mention": "true",
|
"is_mention": "true",
|
||||||
|
"peer_kind": mentionPeerKind,
|
||||||
|
"peer_id": mentionPeerID,
|
||||||
|
"team_id": c.teamID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, nil, metadata)
|
c.HandleMessage(senderID, chatID, content, nil, metadata)
|
||||||
|
|
@ -373,6 +395,9 @@ func (c *SlackChannel) handleSlashCommand(event socketmode.Event) {
|
||||||
"platform": "slack",
|
"platform": "slack",
|
||||||
"is_command": "true",
|
"is_command": "true",
|
||||||
"trigger_id": cmd.TriggerID,
|
"trigger_id": cmd.TriggerID,
|
||||||
|
"peer_kind": "channel",
|
||||||
|
"peer_id": channelID,
|
||||||
|
"team_id": c.teamID,
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("slack", "Slash command received", map[string]interface{}{
|
logger.DebugCF("slack", "Slash command received", map[string]interface{}{
|
||||||
|
|
|
||||||
|
|
@ -59,6 +59,13 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
Proxy: http.ProxyURL(proxyURL),
|
||||||
},
|
},
|
||||||
}))
|
}))
|
||||||
|
} else if os.Getenv("HTTP_PROXY") != "" || os.Getenv("HTTPS_PROXY") != "" {
|
||||||
|
// Use environment proxy if configured
|
||||||
|
opts = append(opts, telego.WithHTTPClient(&http.Client{
|
||||||
|
Transport: &http.Transport{
|
||||||
|
Proxy: http.ProxyFromEnvironment,
|
||||||
|
},
|
||||||
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
||||||
|
|
@ -347,12 +354,21 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
c.placeholders.Store(chatIDStr, pID)
|
c.placeholders.Store(chatIDStr, pID)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
peerKind := "direct"
|
||||||
|
peerID := fmt.Sprintf("%d", user.ID)
|
||||||
|
if message.Chat.Type != "private" {
|
||||||
|
peerKind = "group"
|
||||||
|
peerID = fmt.Sprintf("%d", chatID)
|
||||||
|
}
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": fmt.Sprintf("%d", message.MessageID),
|
"message_id": fmt.Sprintf("%d", message.MessageID),
|
||||||
"user_id": fmt.Sprintf("%d", user.ID),
|
"user_id": fmt.Sprintf("%d", user.ID),
|
||||||
"username": user.Username,
|
"username": user.Username,
|
||||||
"first_name": user.FirstName,
|
"first_name": user.FirstName,
|
||||||
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
|
"peer_kind": peerKind,
|
||||||
|
"peer_id": peerID,
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(fmt.Sprintf("%d", user.ID), fmt.Sprintf("%d", chatID), content, mediaPaths, metadata)
|
c.HandleMessage(fmt.Sprintf("%d", user.ID), fmt.Sprintf("%d", chatID), content, mediaPaths, metadata)
|
||||||
|
|
|
||||||
|
|
@ -178,6 +178,14 @@ func (c *WhatsAppChannel) handleIncomingMessage(msg map[string]interface{}) {
|
||||||
metadata["user_name"] = userName
|
metadata["user_name"] = userName
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if chatID == senderID {
|
||||||
|
metadata["peer_kind"] = "direct"
|
||||||
|
metadata["peer_id"] = senderID
|
||||||
|
} else {
|
||||||
|
metadata["peer_kind"] = "group"
|
||||||
|
metadata["peer_id"] = chatID
|
||||||
|
}
|
||||||
|
|
||||||
log.Printf("WhatsApp message from %s: %s...", senderID, utils.Truncate(content, 50))
|
log.Printf("WhatsApp message from %s: %s...", senderID, utils.Truncate(content, 50))
|
||||||
|
|
||||||
c.HandleMessage(senderID, chatID, content, mediaPaths, metadata)
|
c.HandleMessage(senderID, chatID, content, mediaPaths, metadata)
|
||||||
|
|
|
||||||
|
|
@ -5,11 +5,14 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/caarlos0/env/v11"
|
"github.com/caarlos0/env/v11"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// rrCounter is a global counter for round-robin load balancing across models.
|
||||||
|
var rrCounter atomic.Uint64
|
||||||
|
|
||||||
// FlexibleStringSlice is a []string that also accepts JSON numbers,
|
// FlexibleStringSlice is a []string that also accepts JSON numbers,
|
||||||
// so allow_from can contain both "123" and 123.
|
// so allow_from can contain both "123" and 123.
|
||||||
type FlexibleStringSlice []string
|
type FlexibleStringSlice []string
|
||||||
|
|
@ -45,27 +48,135 @@ func (f *FlexibleStringSlice) UnmarshalJSON(data []byte) error {
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Agents AgentsConfig `json:"agents"`
|
Agents AgentsConfig `json:"agents"`
|
||||||
|
Bindings []AgentBinding `json:"bindings,omitempty"`
|
||||||
|
Session SessionConfig `json:"session,omitempty"`
|
||||||
Channels ChannelsConfig `json:"channels"`
|
Channels ChannelsConfig `json:"channels"`
|
||||||
Providers ProvidersConfig `json:"providers"`
|
Providers ProvidersConfig `json:"providers,omitempty"`
|
||||||
|
ModelList []ModelConfig `json:"model_list"` // New model-centric provider configuration
|
||||||
Gateway GatewayConfig `json:"gateway"`
|
Gateway GatewayConfig `json:"gateway"`
|
||||||
Tools ToolsConfig `json:"tools"`
|
Tools ToolsConfig `json:"tools"`
|
||||||
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
||||||
Devices DevicesConfig `json:"devices"`
|
Devices DevicesConfig `json:"devices"`
|
||||||
mu sync.RWMutex
|
}
|
||||||
|
|
||||||
|
// MarshalJSON implements custom JSON marshaling for Config
|
||||||
|
// to omit providers section when empty and session when empty
|
||||||
|
func (c Config) MarshalJSON() ([]byte, error) {
|
||||||
|
type Alias Config
|
||||||
|
aux := &struct {
|
||||||
|
Providers *ProvidersConfig `json:"providers,omitempty"`
|
||||||
|
Session *SessionConfig `json:"session,omitempty"`
|
||||||
|
*Alias
|
||||||
|
}{
|
||||||
|
Alias: (*Alias)(&c),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only include providers if not empty
|
||||||
|
if !c.Providers.IsEmpty() {
|
||||||
|
aux.Providers = &c.Providers
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only include session if not empty
|
||||||
|
if c.Session.DMScope != "" || len(c.Session.IdentityLinks) > 0 {
|
||||||
|
aux.Session = &c.Session
|
||||||
|
}
|
||||||
|
|
||||||
|
return json.Marshal(aux)
|
||||||
}
|
}
|
||||||
|
|
||||||
type AgentsConfig struct {
|
type AgentsConfig struct {
|
||||||
Defaults AgentDefaults `json:"defaults"`
|
Defaults AgentDefaults `json:"defaults"`
|
||||||
|
List []AgentConfig `json:"list,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// AgentModelConfig supports both string and structured model config.
|
||||||
|
// String format: "gpt-4" (just primary, no fallbacks)
|
||||||
|
// Object format: {"primary": "gpt-4", "fallbacks": ["claude-haiku"]}
|
||||||
|
type AgentModelConfig struct {
|
||||||
|
Primary string `json:"primary,omitempty"`
|
||||||
|
Fallbacks []string `json:"fallbacks,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *AgentModelConfig) UnmarshalJSON(data []byte) error {
|
||||||
|
var s string
|
||||||
|
if err := json.Unmarshal(data, &s); err == nil {
|
||||||
|
m.Primary = s
|
||||||
|
m.Fallbacks = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
type raw struct {
|
||||||
|
Primary string `json:"primary"`
|
||||||
|
Fallbacks []string `json:"fallbacks"`
|
||||||
|
}
|
||||||
|
var r raw
|
||||||
|
if err := json.Unmarshal(data, &r); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
m.Primary = r.Primary
|
||||||
|
m.Fallbacks = r.Fallbacks
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m AgentModelConfig) MarshalJSON() ([]byte, error) {
|
||||||
|
if len(m.Fallbacks) == 0 && m.Primary != "" {
|
||||||
|
return json.Marshal(m.Primary)
|
||||||
|
}
|
||||||
|
type raw struct {
|
||||||
|
Primary string `json:"primary,omitempty"`
|
||||||
|
Fallbacks []string `json:"fallbacks,omitempty"`
|
||||||
|
}
|
||||||
|
return json.Marshal(raw{Primary: m.Primary, Fallbacks: m.Fallbacks})
|
||||||
|
}
|
||||||
|
|
||||||
|
type AgentConfig struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Default bool `json:"default,omitempty"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
Workspace string `json:"workspace,omitempty"`
|
||||||
|
Model *AgentModelConfig `json:"model,omitempty"`
|
||||||
|
Skills []string `json:"skills,omitempty"`
|
||||||
|
Subagents *SubagentsConfig `json:"subagents,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type SubagentsConfig struct {
|
||||||
|
AllowAgents []string `json:"allow_agents,omitempty"`
|
||||||
|
Model *AgentModelConfig `json:"model,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type PeerMatch struct {
|
||||||
|
Kind string `json:"kind"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type BindingMatch struct {
|
||||||
|
Channel string `json:"channel"`
|
||||||
|
AccountID string `json:"account_id,omitempty"`
|
||||||
|
Peer *PeerMatch `json:"peer,omitempty"`
|
||||||
|
GuildID string `json:"guild_id,omitempty"`
|
||||||
|
TeamID string `json:"team_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type AgentBinding struct {
|
||||||
|
AgentID string `json:"agent_id"`
|
||||||
|
Match BindingMatch `json:"match"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type SessionConfig struct {
|
||||||
|
DMScope string `json:"dm_scope,omitempty"`
|
||||||
|
IdentityLinks map[string][]string `json:"identity_links,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type AgentDefaults struct {
|
type AgentDefaults struct {
|
||||||
Workspace string `json:"workspace" env:"PICOCLAW_AGENTS_DEFAULTS_WORKSPACE"`
|
Workspace string `json:"workspace" env:"PICOCLAW_AGENTS_DEFAULTS_WORKSPACE"`
|
||||||
RestrictToWorkspace bool `json:"restrict_to_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_RESTRICT_TO_WORKSPACE"`
|
RestrictToWorkspace bool `json:"restrict_to_workspace" env:"PICOCLAW_AGENTS_DEFAULTS_RESTRICT_TO_WORKSPACE"`
|
||||||
Provider string `json:"provider" env:"PICOCLAW_AGENTS_DEFAULTS_PROVIDER"`
|
Provider string `json:"provider" env:"PICOCLAW_AGENTS_DEFAULTS_PROVIDER"`
|
||||||
Model string `json:"model" env:"PICOCLAW_AGENTS_DEFAULTS_MODEL"`
|
Model string `json:"model" env:"PICOCLAW_AGENTS_DEFAULTS_MODEL"`
|
||||||
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
ModelFallbacks []string `json:"model_fallbacks,omitempty"`
|
||||||
Temperature float64 `json:"temperature" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
ImageModel string `json:"image_model,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_IMAGE_MODEL"`
|
||||||
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
ImageModelFallbacks []string `json:"image_model_fallbacks,omitempty"`
|
||||||
|
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
||||||
|
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
||||||
|
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ChannelsConfig struct {
|
type ChannelsConfig struct {
|
||||||
|
|
@ -104,9 +215,10 @@ type FeishuConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type DiscordConfig struct {
|
type DiscordConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_DISCORD_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_DISCORD_ENABLED"`
|
||||||
Token string `json:"token" env:"PICOCLAW_CHANNELS_DISCORD_TOKEN"`
|
Token string `json:"token" env:"PICOCLAW_CHANNELS_DISCORD_TOKEN"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_DISCORD_ALLOW_FROM"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_DISCORD_ALLOW_FROM"`
|
||||||
|
MentionOnly bool `json:"mention_only" env:"PICOCLAW_CHANNELS_DISCORD_MENTION_ONLY"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type MaixCamConfig struct {
|
type MaixCamConfig struct {
|
||||||
|
|
@ -167,20 +279,55 @@ type DevicesConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type ProvidersConfig struct {
|
type ProvidersConfig struct {
|
||||||
Anthropic ProviderConfig `json:"anthropic"`
|
Anthropic ProviderConfig `json:"anthropic"`
|
||||||
OpenAI ProviderConfig `json:"openai"`
|
OpenAI OpenAIProviderConfig `json:"openai"`
|
||||||
OpenRouter ProviderConfig `json:"openrouter"`
|
OpenRouter ProviderConfig `json:"openrouter"`
|
||||||
Groq ProviderConfig `json:"groq"`
|
Groq ProviderConfig `json:"groq"`
|
||||||
Zhipu ProviderConfig `json:"zhipu"`
|
Zhipu ProviderConfig `json:"zhipu"`
|
||||||
VLLM ProviderConfig `json:"vllm"`
|
VLLM ProviderConfig `json:"vllm"`
|
||||||
Gemini ProviderConfig `json:"gemini"`
|
Gemini ProviderConfig `json:"gemini"`
|
||||||
Nvidia ProviderConfig `json:"nvidia"`
|
Nvidia ProviderConfig `json:"nvidia"`
|
||||||
Ollama ProviderConfig `json:"ollama"`
|
Ollama ProviderConfig `json:"ollama"`
|
||||||
Moonshot ProviderConfig `json:"moonshot"`
|
Moonshot ProviderConfig `json:"moonshot"`
|
||||||
ShengSuanYun ProviderConfig `json:"shengsuanyun"`
|
ShengSuanYun ProviderConfig `json:"shengsuanyun"`
|
||||||
DeepSeek ProviderConfig `json:"deepseek"`
|
DeepSeek ProviderConfig `json:"deepseek"`
|
||||||
VolcEngine ProviderConfig `json:"volcengine"`
|
Cerebras ProviderConfig `json:"cerebras"`
|
||||||
GitHubCopilot ProviderConfig `json:"github_copilot"`
|
VolcEngine ProviderConfig `json:"volcengine"`
|
||||||
|
GitHubCopilot ProviderConfig `json:"github_copilot"`
|
||||||
|
Antigravity ProviderConfig `json:"antigravity"`
|
||||||
|
Qwen ProviderConfig `json:"qwen"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsEmpty checks if all provider configs are empty (no API keys or API bases set)
|
||||||
|
// Note: WebSearch is an optimization option and doesn't count as "non-empty"
|
||||||
|
func (p ProvidersConfig) IsEmpty() bool {
|
||||||
|
return p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" &&
|
||||||
|
p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" &&
|
||||||
|
p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" &&
|
||||||
|
p.Groq.APIKey == "" && p.Groq.APIBase == "" &&
|
||||||
|
p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" &&
|
||||||
|
p.VLLM.APIKey == "" && p.VLLM.APIBase == "" &&
|
||||||
|
p.Gemini.APIKey == "" && p.Gemini.APIBase == "" &&
|
||||||
|
p.Nvidia.APIKey == "" && p.Nvidia.APIBase == "" &&
|
||||||
|
p.Ollama.APIKey == "" && p.Ollama.APIBase == "" &&
|
||||||
|
p.Moonshot.APIKey == "" && p.Moonshot.APIBase == "" &&
|
||||||
|
p.ShengSuanYun.APIKey == "" && p.ShengSuanYun.APIBase == "" &&
|
||||||
|
p.DeepSeek.APIKey == "" && p.DeepSeek.APIBase == "" &&
|
||||||
|
p.Cerebras.APIKey == "" && p.Cerebras.APIBase == "" &&
|
||||||
|
p.VolcEngine.APIKey == "" && p.VolcEngine.APIBase == "" &&
|
||||||
|
p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" &&
|
||||||
|
p.Antigravity.APIKey == "" && p.Antigravity.APIBase == "" &&
|
||||||
|
p.Qwen.APIKey == "" && p.Qwen.APIBase == ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalJSON implements custom JSON marshaling for ProvidersConfig
|
||||||
|
// to omit the entire section when empty
|
||||||
|
func (p ProvidersConfig) MarshalJSON() ([]byte, error) {
|
||||||
|
if p.IsEmpty() {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
type Alias ProvidersConfig
|
||||||
|
return json.Marshal((*Alias)(&p))
|
||||||
}
|
}
|
||||||
|
|
||||||
type ProviderConfig struct {
|
type ProviderConfig struct {
|
||||||
|
|
@ -191,6 +338,47 @@ type ProviderConfig struct {
|
||||||
ConnectMode string `json:"connect_mode,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_CONNECT_MODE"` //only for Github Copilot, `stdio` or `grpc`
|
ConnectMode string `json:"connect_mode,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_CONNECT_MODE"` //only for Github Copilot, `stdio` or `grpc`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type OpenAIProviderConfig struct {
|
||||||
|
ProviderConfig
|
||||||
|
WebSearch bool `json:"web_search" env:"PICOCLAW_PROVIDERS_OPENAI_WEB_SEARCH"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelConfig represents a model-centric provider configuration.
|
||||||
|
// It allows adding new providers (especially OpenAI-compatible ones) via configuration only.
|
||||||
|
// The model field uses protocol prefix format: [protocol/]model-identifier
|
||||||
|
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot
|
||||||
|
// Default protocol is "openai" if no prefix is specified.
|
||||||
|
type ModelConfig struct {
|
||||||
|
// Required fields
|
||||||
|
ModelName string `json:"model_name"` // User-facing alias for the model
|
||||||
|
Model string `json:"model"` // Protocol/model-identifier (e.g., "openai/gpt-4o", "anthropic/claude-sonnet-4.6")
|
||||||
|
|
||||||
|
// HTTP-based providers
|
||||||
|
APIBase string `json:"api_base,omitempty"` // API endpoint URL
|
||||||
|
APIKey string `json:"api_key"` // API authentication key
|
||||||
|
Proxy string `json:"proxy,omitempty"` // HTTP proxy URL
|
||||||
|
|
||||||
|
// Special providers (CLI-based, OAuth, etc.)
|
||||||
|
AuthMethod string `json:"auth_method,omitempty"` // Authentication method: oauth, token
|
||||||
|
ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc
|
||||||
|
Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers
|
||||||
|
|
||||||
|
// Optional optimizations
|
||||||
|
RPM int `json:"rpm,omitempty"` // Requests per minute limit
|
||||||
|
MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate checks if the ModelConfig has all required fields.
|
||||||
|
func (c *ModelConfig) Validate() error {
|
||||||
|
if c.ModelName == "" {
|
||||||
|
return fmt.Errorf("model_name is required")
|
||||||
|
}
|
||||||
|
if c.Model == "" {
|
||||||
|
return fmt.Errorf("model is required")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
type GatewayConfig struct {
|
type GatewayConfig struct {
|
||||||
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
|
|
@ -223,137 +411,43 @@ type CronToolsConfig struct {
|
||||||
ExecTimeoutMinutes int `json:"exec_timeout_minutes" env:"PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES"` // 0 means no timeout
|
ExecTimeoutMinutes int `json:"exec_timeout_minutes" env:"PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES"` // 0 means no timeout
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolsConfig struct {
|
type ExecConfig struct {
|
||||||
Web WebToolsConfig `json:"web"`
|
EnableDenyPatterns bool `json:"enable_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_ENABLE_DENY_PATTERNS"`
|
||||||
Cron CronToolsConfig `json:"cron"`
|
CustomDenyPatterns []string `json:"custom_deny_patterns" env:"PICOCLAW_TOOLS_EXEC_CUSTOM_DENY_PATTERNS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func DefaultConfig() *Config {
|
type ToolsConfig struct {
|
||||||
return &Config{
|
Web WebToolsConfig `json:"web"`
|
||||||
Agents: AgentsConfig{
|
Cron CronToolsConfig `json:"cron"`
|
||||||
Defaults: AgentDefaults{
|
Exec ExecConfig `json:"exec"`
|
||||||
Workspace: "~/.picoclaw/workspace",
|
Skills SkillsToolsConfig `json:"skills"`
|
||||||
RestrictToWorkspace: true,
|
}
|
||||||
Provider: "",
|
|
||||||
Model: "glm-4.7",
|
type SkillsToolsConfig struct {
|
||||||
MaxTokens: 8192,
|
Registries SkillsRegistriesConfig `json:"registries"`
|
||||||
Temperature: 0.7,
|
MaxConcurrentSearches int `json:"max_concurrent_searches" env:"PICOCLAW_SKILLS_MAX_CONCURRENT_SEARCHES"`
|
||||||
MaxToolIterations: 20,
|
SearchCache SearchCacheConfig `json:"search_cache"`
|
||||||
},
|
}
|
||||||
},
|
|
||||||
Channels: ChannelsConfig{
|
type SearchCacheConfig struct {
|
||||||
WhatsApp: WhatsAppConfig{
|
MaxSize int `json:"max_size" env:"PICOCLAW_SKILLS_SEARCH_CACHE_MAX_SIZE"`
|
||||||
Enabled: false,
|
TTLSeconds int `json:"ttl_seconds" env:"PICOCLAW_SKILLS_SEARCH_CACHE_TTL_SECONDS"`
|
||||||
BridgeURL: "ws://localhost:3001",
|
}
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
type SkillsRegistriesConfig struct {
|
||||||
Telegram: TelegramConfig{
|
ClawHub ClawHubRegistryConfig `json:"clawhub"`
|
||||||
Enabled: false,
|
}
|
||||||
Token: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
type ClawHubRegistryConfig struct {
|
||||||
},
|
Enabled bool `json:"enabled" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_ENABLED"`
|
||||||
Feishu: FeishuConfig{
|
BaseURL string `json:"base_url" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_BASE_URL"`
|
||||||
Enabled: false,
|
AuthToken string `json:"auth_token" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_AUTH_TOKEN"`
|
||||||
AppID: "",
|
SearchPath string `json:"search_path" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_SEARCH_PATH"`
|
||||||
AppSecret: "",
|
SkillsPath string `json:"skills_path" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_SKILLS_PATH"`
|
||||||
EncryptKey: "",
|
DownloadPath string `json:"download_path" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_DOWNLOAD_PATH"`
|
||||||
VerificationToken: "",
|
Timeout int `json:"timeout" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_TIMEOUT"`
|
||||||
AllowFrom: FlexibleStringSlice{},
|
MaxZipSize int `json:"max_zip_size" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_MAX_ZIP_SIZE"`
|
||||||
},
|
MaxResponseSize int `json:"max_response_size" env:"PICOCLAW_SKILLS_REGISTRIES_CLAWHUB_MAX_RESPONSE_SIZE"`
|
||||||
Discord: DiscordConfig{
|
|
||||||
Enabled: false,
|
|
||||||
Token: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
MaixCam: MaixCamConfig{
|
|
||||||
Enabled: false,
|
|
||||||
Host: "0.0.0.0",
|
|
||||||
Port: 18790,
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
QQ: QQConfig{
|
|
||||||
Enabled: false,
|
|
||||||
AppID: "",
|
|
||||||
AppSecret: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
DingTalk: DingTalkConfig{
|
|
||||||
Enabled: false,
|
|
||||||
ClientID: "",
|
|
||||||
ClientSecret: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
Slack: SlackConfig{
|
|
||||||
Enabled: false,
|
|
||||||
BotToken: "",
|
|
||||||
AppToken: "",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
LINE: LINEConfig{
|
|
||||||
Enabled: false,
|
|
||||||
ChannelSecret: "",
|
|
||||||
ChannelAccessToken: "",
|
|
||||||
WebhookHost: "0.0.0.0",
|
|
||||||
WebhookPort: 18791,
|
|
||||||
WebhookPath: "/webhook/line",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
OneBot: OneBotConfig{
|
|
||||||
Enabled: false,
|
|
||||||
WSUrl: "ws://127.0.0.1:3001",
|
|
||||||
AccessToken: "",
|
|
||||||
ReconnectInterval: 5,
|
|
||||||
GroupTriggerPrefix: []string{},
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Providers: ProvidersConfig{
|
|
||||||
Anthropic: ProviderConfig{},
|
|
||||||
OpenAI: ProviderConfig{},
|
|
||||||
OpenRouter: ProviderConfig{},
|
|
||||||
Groq: ProviderConfig{},
|
|
||||||
Zhipu: ProviderConfig{},
|
|
||||||
VLLM: ProviderConfig{},
|
|
||||||
Gemini: ProviderConfig{},
|
|
||||||
Nvidia: ProviderConfig{},
|
|
||||||
Moonshot: ProviderConfig{},
|
|
||||||
ShengSuanYun: ProviderConfig{},
|
|
||||||
VolcEngine: ProviderConfig{},
|
|
||||||
},
|
|
||||||
Gateway: GatewayConfig{
|
|
||||||
Host: "0.0.0.0",
|
|
||||||
Port: 18790,
|
|
||||||
},
|
|
||||||
Tools: ToolsConfig{
|
|
||||||
Web: WebToolsConfig{
|
|
||||||
Brave: BraveConfig{
|
|
||||||
Enabled: false,
|
|
||||||
APIKey: "",
|
|
||||||
MaxResults: 5,
|
|
||||||
},
|
|
||||||
DuckDuckGo: DuckDuckGoConfig{
|
|
||||||
Enabled: true,
|
|
||||||
MaxResults: 5,
|
|
||||||
},
|
|
||||||
Perplexity: PerplexityConfig{
|
|
||||||
Enabled: false,
|
|
||||||
APIKey: "",
|
|
||||||
MaxResults: 5,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Cron: CronToolsConfig{
|
|
||||||
ExecTimeoutMinutes: 5, // default 5 minutes for LLM operations
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Heartbeat: HeartbeatConfig{
|
|
||||||
Enabled: true,
|
|
||||||
Interval: 30, // default 30 minutes
|
|
||||||
},
|
|
||||||
Devices: DevicesConfig{
|
|
||||||
Enabled: false,
|
|
||||||
MonitorUSB: true,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func LoadConfig(path string) (*Config, error) {
|
func LoadConfig(path string) (*Config, error) {
|
||||||
|
|
@ -375,13 +469,20 @@ func LoadConfig(path string) (*Config, error) {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Auto-migrate: if only legacy providers config exists, convert to model_list
|
||||||
|
if len(cfg.ModelList) == 0 && cfg.HasProvidersConfig() {
|
||||||
|
cfg.ModelList = ConvertProvidersToModelList(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate model_list for uniqueness and required fields
|
||||||
|
if err := cfg.ValidateModelList(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
return cfg, nil
|
return cfg, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func SaveConfig(path string, cfg *Config) error {
|
func SaveConfig(path string, cfg *Config) error {
|
||||||
cfg.mu.RLock()
|
|
||||||
defer cfg.mu.RUnlock()
|
|
||||||
|
|
||||||
data, err := json.MarshalIndent(cfg, "", " ")
|
data, err := json.MarshalIndent(cfg, "", " ")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
@ -396,14 +497,10 @@ func SaveConfig(path string, cfg *Config) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) WorkspacePath() string {
|
func (c *Config) WorkspacePath() string {
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
return expandHome(c.Agents.Defaults.Workspace)
|
return expandHome(c.Agents.Defaults.Workspace)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) GetAPIKey() string {
|
func (c *Config) GetAPIKey() string {
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
if c.Providers.OpenRouter.APIKey != "" {
|
if c.Providers.OpenRouter.APIKey != "" {
|
||||||
return c.Providers.OpenRouter.APIKey
|
return c.Providers.OpenRouter.APIKey
|
||||||
}
|
}
|
||||||
|
|
@ -428,12 +525,13 @@ func (c *Config) GetAPIKey() string {
|
||||||
if c.Providers.ShengSuanYun.APIKey != "" {
|
if c.Providers.ShengSuanYun.APIKey != "" {
|
||||||
return c.Providers.ShengSuanYun.APIKey
|
return c.Providers.ShengSuanYun.APIKey
|
||||||
}
|
}
|
||||||
|
if c.Providers.Cerebras.APIKey != "" {
|
||||||
|
return c.Providers.Cerebras.APIKey
|
||||||
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) GetAPIBase() string {
|
func (c *Config) GetAPIBase() string {
|
||||||
c.mu.RLock()
|
|
||||||
defer c.mu.RUnlock()
|
|
||||||
if c.Providers.OpenRouter.APIKey != "" {
|
if c.Providers.OpenRouter.APIKey != "" {
|
||||||
if c.Providers.OpenRouter.APIBase != "" {
|
if c.Providers.OpenRouter.APIBase != "" {
|
||||||
return c.Providers.OpenRouter.APIBase
|
return c.Providers.OpenRouter.APIBase
|
||||||
|
|
@ -462,3 +560,65 @@ func expandHome(path string) string {
|
||||||
}
|
}
|
||||||
return path
|
return path
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetModelConfig returns the ModelConfig for the given model name.
|
||||||
|
// If multiple configs exist with the same model_name, it uses round-robin
|
||||||
|
// selection for load balancing. Returns an error if the model is not found.
|
||||||
|
func (c *Config) GetModelConfig(modelName string) (*ModelConfig, error) {
|
||||||
|
matches := c.findMatches(modelName)
|
||||||
|
if len(matches) == 0 {
|
||||||
|
return nil, fmt.Errorf("model %q not found in model_list or providers", modelName)
|
||||||
|
}
|
||||||
|
if len(matches) == 1 {
|
||||||
|
return &matches[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Multiple configs - use round-robin for load balancing
|
||||||
|
idx := rrCounter.Add(1) % uint64(len(matches))
|
||||||
|
return &matches[idx], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// findMatches finds all ModelConfig entries with the given model_name.
|
||||||
|
func (c *Config) findMatches(modelName string) []ModelConfig {
|
||||||
|
var matches []ModelConfig
|
||||||
|
for i := range c.ModelList {
|
||||||
|
if c.ModelList[i].ModelName == modelName {
|
||||||
|
matches = append(matches, c.ModelList[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return matches
|
||||||
|
}
|
||||||
|
|
||||||
|
// HasProvidersConfig checks if any provider in the old providers config has configuration.
|
||||||
|
func (c *Config) HasProvidersConfig() bool {
|
||||||
|
v := c.Providers
|
||||||
|
return v.Anthropic.APIKey != "" || v.Anthropic.APIBase != "" ||
|
||||||
|
v.OpenAI.APIKey != "" || v.OpenAI.APIBase != "" ||
|
||||||
|
v.OpenRouter.APIKey != "" || v.OpenRouter.APIBase != "" ||
|
||||||
|
v.Groq.APIKey != "" || v.Groq.APIBase != "" ||
|
||||||
|
v.Zhipu.APIKey != "" || v.Zhipu.APIBase != "" ||
|
||||||
|
v.VLLM.APIKey != "" || v.VLLM.APIBase != "" ||
|
||||||
|
v.Gemini.APIKey != "" || v.Gemini.APIBase != "" ||
|
||||||
|
v.Nvidia.APIKey != "" || v.Nvidia.APIBase != "" ||
|
||||||
|
v.Ollama.APIKey != "" || v.Ollama.APIBase != "" ||
|
||||||
|
v.Moonshot.APIKey != "" || v.Moonshot.APIBase != "" ||
|
||||||
|
v.ShengSuanYun.APIKey != "" || v.ShengSuanYun.APIBase != "" ||
|
||||||
|
v.DeepSeek.APIKey != "" || v.DeepSeek.APIBase != "" ||
|
||||||
|
v.Cerebras.APIKey != "" || v.Cerebras.APIBase != "" ||
|
||||||
|
v.VolcEngine.APIKey != "" || v.VolcEngine.APIBase != "" ||
|
||||||
|
v.GitHubCopilot.APIKey != "" || v.GitHubCopilot.APIBase != "" ||
|
||||||
|
v.Antigravity.APIKey != "" || v.Antigravity.APIBase != "" ||
|
||||||
|
v.Qwen.APIKey != "" || v.Qwen.APIBase != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateModelList validates all ModelConfig entries in the model_list.
|
||||||
|
// It checks that each model config is valid.
|
||||||
|
// Note: Multiple entries with the same model_name are allowed for load balancing.
|
||||||
|
func (c *Config) ValidateModelList() error {
|
||||||
|
for i := range c.ModelList {
|
||||||
|
if err := c.ModelList[i].Validate(); err != nil {
|
||||||
|
return fmt.Errorf("model_list[%d]: %w", i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,12 +1,193 @@
|
||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
"runtime"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func TestAgentModelConfig_UnmarshalString(t *testing.T) {
|
||||||
|
var m AgentModelConfig
|
||||||
|
if err := json.Unmarshal([]byte(`"gpt-4"`), &m); err != nil {
|
||||||
|
t.Fatalf("unmarshal string: %v", err)
|
||||||
|
}
|
||||||
|
if m.Primary != "gpt-4" {
|
||||||
|
t.Errorf("Primary = %q, want 'gpt-4'", m.Primary)
|
||||||
|
}
|
||||||
|
if m.Fallbacks != nil {
|
||||||
|
t.Errorf("Fallbacks = %v, want nil", m.Fallbacks)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentModelConfig_UnmarshalObject(t *testing.T) {
|
||||||
|
var m AgentModelConfig
|
||||||
|
data := `{"primary": "claude-opus", "fallbacks": ["gpt-4o-mini", "haiku"]}`
|
||||||
|
if err := json.Unmarshal([]byte(data), &m); err != nil {
|
||||||
|
t.Fatalf("unmarshal object: %v", err)
|
||||||
|
}
|
||||||
|
if m.Primary != "claude-opus" {
|
||||||
|
t.Errorf("Primary = %q, want 'claude-opus'", m.Primary)
|
||||||
|
}
|
||||||
|
if len(m.Fallbacks) != 2 {
|
||||||
|
t.Fatalf("Fallbacks len = %d, want 2", len(m.Fallbacks))
|
||||||
|
}
|
||||||
|
if m.Fallbacks[0] != "gpt-4o-mini" || m.Fallbacks[1] != "haiku" {
|
||||||
|
t.Errorf("Fallbacks = %v", m.Fallbacks)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentModelConfig_MarshalString(t *testing.T) {
|
||||||
|
m := AgentModelConfig{Primary: "gpt-4"}
|
||||||
|
data, err := json.Marshal(m)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != `"gpt-4"` {
|
||||||
|
t.Errorf("marshal = %s, want '\"gpt-4\"'", string(data))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentModelConfig_MarshalObject(t *testing.T) {
|
||||||
|
m := AgentModelConfig{Primary: "claude-opus", Fallbacks: []string{"haiku"}}
|
||||||
|
data, err := json.Marshal(m)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal: %v", err)
|
||||||
|
}
|
||||||
|
var result map[string]interface{}
|
||||||
|
json.Unmarshal(data, &result)
|
||||||
|
if result["primary"] != "claude-opus" {
|
||||||
|
t.Errorf("primary = %v", result["primary"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentConfig_FullParse(t *testing.T) {
|
||||||
|
jsonData := `{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"workspace": "~/.picoclaw/workspace",
|
||||||
|
"model": "glm-4.7",
|
||||||
|
"max_tokens": 8192,
|
||||||
|
"max_tool_iterations": 20
|
||||||
|
},
|
||||||
|
"list": [
|
||||||
|
{
|
||||||
|
"id": "sales",
|
||||||
|
"default": true,
|
||||||
|
"name": "Sales Bot",
|
||||||
|
"model": "gpt-4"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "support",
|
||||||
|
"name": "Support Bot",
|
||||||
|
"model": {
|
||||||
|
"primary": "claude-opus",
|
||||||
|
"fallbacks": ["haiku"]
|
||||||
|
},
|
||||||
|
"subagents": {
|
||||||
|
"allow_agents": ["sales"]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"bindings": [
|
||||||
|
{
|
||||||
|
"agent_id": "support",
|
||||||
|
"match": {
|
||||||
|
"channel": "telegram",
|
||||||
|
"account_id": "*",
|
||||||
|
"peer": {"kind": "direct", "id": "user123"}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"session": {
|
||||||
|
"dm_scope": "per-peer",
|
||||||
|
"identity_links": {
|
||||||
|
"john": ["telegram:123", "discord:john#1234"]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}`
|
||||||
|
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
if err := json.Unmarshal([]byte(jsonData), cfg); err != nil {
|
||||||
|
t.Fatalf("unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cfg.Agents.List) != 2 {
|
||||||
|
t.Fatalf("agents.list len = %d, want 2", len(cfg.Agents.List))
|
||||||
|
}
|
||||||
|
|
||||||
|
sales := cfg.Agents.List[0]
|
||||||
|
if sales.ID != "sales" || !sales.Default || sales.Name != "Sales Bot" {
|
||||||
|
t.Errorf("sales = %+v", sales)
|
||||||
|
}
|
||||||
|
if sales.Model == nil || sales.Model.Primary != "gpt-4" {
|
||||||
|
t.Errorf("sales.Model = %+v", sales.Model)
|
||||||
|
}
|
||||||
|
|
||||||
|
support := cfg.Agents.List[1]
|
||||||
|
if support.ID != "support" || support.Name != "Support Bot" {
|
||||||
|
t.Errorf("support = %+v", support)
|
||||||
|
}
|
||||||
|
if support.Model == nil || support.Model.Primary != "claude-opus" {
|
||||||
|
t.Errorf("support.Model = %+v", support.Model)
|
||||||
|
}
|
||||||
|
if len(support.Model.Fallbacks) != 1 || support.Model.Fallbacks[0] != "haiku" {
|
||||||
|
t.Errorf("support.Model.Fallbacks = %v", support.Model.Fallbacks)
|
||||||
|
}
|
||||||
|
if support.Subagents == nil || len(support.Subagents.AllowAgents) != 1 {
|
||||||
|
t.Errorf("support.Subagents = %+v", support.Subagents)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cfg.Bindings) != 1 {
|
||||||
|
t.Fatalf("bindings len = %d, want 1", len(cfg.Bindings))
|
||||||
|
}
|
||||||
|
binding := cfg.Bindings[0]
|
||||||
|
if binding.AgentID != "support" || binding.Match.Channel != "telegram" {
|
||||||
|
t.Errorf("binding = %+v", binding)
|
||||||
|
}
|
||||||
|
if binding.Match.Peer == nil || binding.Match.Peer.Kind != "direct" || binding.Match.Peer.ID != "user123" {
|
||||||
|
t.Errorf("binding.Match.Peer = %+v", binding.Match.Peer)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Session.DMScope != "per-peer" {
|
||||||
|
t.Errorf("Session.DMScope = %q", cfg.Session.DMScope)
|
||||||
|
}
|
||||||
|
if len(cfg.Session.IdentityLinks) != 1 {
|
||||||
|
t.Errorf("Session.IdentityLinks = %v", cfg.Session.IdentityLinks)
|
||||||
|
}
|
||||||
|
links := cfg.Session.IdentityLinks["john"]
|
||||||
|
if len(links) != 2 {
|
||||||
|
t.Errorf("john links = %v", links)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfig_BackwardCompat_NoAgentsList(t *testing.T) {
|
||||||
|
jsonData := `{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"workspace": "~/.picoclaw/workspace",
|
||||||
|
"model": "glm-4.7",
|
||||||
|
"max_tokens": 8192,
|
||||||
|
"max_tool_iterations": 20
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}`
|
||||||
|
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
if err := json.Unmarshal([]byte(jsonData), cfg); err != nil {
|
||||||
|
t.Fatalf("unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cfg.Agents.List) != 0 {
|
||||||
|
t.Errorf("agents.list should be empty for backward compat, got %d", len(cfg.Agents.List))
|
||||||
|
}
|
||||||
|
if len(cfg.Bindings) != 0 {
|
||||||
|
t.Errorf("bindings should be empty, got %d", len(cfg.Bindings))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestDefaultConfig_HeartbeatEnabled verifies heartbeat is enabled by default
|
// TestDefaultConfig_HeartbeatEnabled verifies heartbeat is enabled by default
|
||||||
func TestDefaultConfig_HeartbeatEnabled(t *testing.T) {
|
func TestDefaultConfig_HeartbeatEnabled(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
@ -20,8 +201,6 @@ func TestDefaultConfig_HeartbeatEnabled(t *testing.T) {
|
||||||
func TestDefaultConfig_WorkspacePath(t *testing.T) {
|
func TestDefaultConfig_WorkspacePath(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
// Just verify the workspace is set, don't compare exact paths
|
|
||||||
// since expandHome behavior may differ based on environment
|
|
||||||
if cfg.Agents.Defaults.Workspace == "" {
|
if cfg.Agents.Defaults.Workspace == "" {
|
||||||
t.Error("Workspace should not be empty")
|
t.Error("Workspace should not be empty")
|
||||||
}
|
}
|
||||||
|
|
@ -58,8 +237,8 @@ func TestDefaultConfig_MaxToolIterations(t *testing.T) {
|
||||||
func TestDefaultConfig_Temperature(t *testing.T) {
|
func TestDefaultConfig_Temperature(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
if cfg.Agents.Defaults.Temperature == 0 {
|
if cfg.Agents.Defaults.Temperature != nil {
|
||||||
t.Error("Temperature should not be zero")
|
t.Error("Temperature should be nil when not provided")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -79,7 +258,6 @@ func TestDefaultConfig_Gateway(t *testing.T) {
|
||||||
func TestDefaultConfig_Providers(t *testing.T) {
|
func TestDefaultConfig_Providers(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
// Verify all providers are empty by default
|
|
||||||
if cfg.Providers.Anthropic.APIKey != "" {
|
if cfg.Providers.Anthropic.APIKey != "" {
|
||||||
t.Error("Anthropic API key should be empty by default")
|
t.Error("Anthropic API key should be empty by default")
|
||||||
}
|
}
|
||||||
|
|
@ -89,46 +267,18 @@ func TestDefaultConfig_Providers(t *testing.T) {
|
||||||
if cfg.Providers.OpenRouter.APIKey != "" {
|
if cfg.Providers.OpenRouter.APIKey != "" {
|
||||||
t.Error("OpenRouter API key should be empty by default")
|
t.Error("OpenRouter API key should be empty by default")
|
||||||
}
|
}
|
||||||
if cfg.Providers.Groq.APIKey != "" {
|
|
||||||
t.Error("Groq API key should be empty by default")
|
|
||||||
}
|
|
||||||
if cfg.Providers.Zhipu.APIKey != "" {
|
|
||||||
t.Error("Zhipu API key should be empty by default")
|
|
||||||
}
|
|
||||||
if cfg.Providers.VLLM.APIKey != "" {
|
|
||||||
t.Error("VLLM API key should be empty by default")
|
|
||||||
}
|
|
||||||
if cfg.Providers.Gemini.APIKey != "" {
|
|
||||||
t.Error("Gemini API key should be empty by default")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestDefaultConfig_Channels verifies channels are disabled by default
|
// TestDefaultConfig_Channels verifies channels are disabled by default
|
||||||
func TestDefaultConfig_Channels(t *testing.T) {
|
func TestDefaultConfig_Channels(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
// Verify all channels are disabled by default
|
|
||||||
if cfg.Channels.WhatsApp.Enabled {
|
|
||||||
t.Error("WhatsApp should be disabled by default")
|
|
||||||
}
|
|
||||||
if cfg.Channels.Telegram.Enabled {
|
if cfg.Channels.Telegram.Enabled {
|
||||||
t.Error("Telegram should be disabled by default")
|
t.Error("Telegram should be disabled by default")
|
||||||
}
|
}
|
||||||
if cfg.Channels.Feishu.Enabled {
|
|
||||||
t.Error("Feishu should be disabled by default")
|
|
||||||
}
|
|
||||||
if cfg.Channels.Discord.Enabled {
|
if cfg.Channels.Discord.Enabled {
|
||||||
t.Error("Discord should be disabled by default")
|
t.Error("Discord should be disabled by default")
|
||||||
}
|
}
|
||||||
if cfg.Channels.MaixCam.Enabled {
|
|
||||||
t.Error("MaixCam should be disabled by default")
|
|
||||||
}
|
|
||||||
if cfg.Channels.QQ.Enabled {
|
|
||||||
t.Error("QQ should be disabled by default")
|
|
||||||
}
|
|
||||||
if cfg.Channels.DingTalk.Enabled {
|
|
||||||
t.Error("DingTalk should be disabled by default")
|
|
||||||
}
|
|
||||||
if cfg.Channels.Slack.Enabled {
|
if cfg.Channels.Slack.Enabled {
|
||||||
t.Error("Slack should be disabled by default")
|
t.Error("Slack should be disabled by default")
|
||||||
}
|
}
|
||||||
|
|
@ -178,15 +328,14 @@ func TestSaveConfig_FilePermissions(t *testing.T) {
|
||||||
func TestConfig_Complete(t *testing.T) {
|
func TestConfig_Complete(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
// Verify complete config structure
|
|
||||||
if cfg.Agents.Defaults.Workspace == "" {
|
if cfg.Agents.Defaults.Workspace == "" {
|
||||||
t.Error("Workspace should not be empty")
|
t.Error("Workspace should not be empty")
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Model == "" {
|
if cfg.Agents.Defaults.Model == "" {
|
||||||
t.Error("Model should not be empty")
|
t.Error("Model should not be empty")
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Temperature == 0 {
|
if cfg.Agents.Defaults.Temperature != nil {
|
||||||
t.Error("Temperature should have default value")
|
t.Error("Temperature should be nil when not provided")
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.MaxTokens == 0 {
|
if cfg.Agents.Defaults.MaxTokens == 0 {
|
||||||
t.Error("MaxTokens should not be zero")
|
t.Error("MaxTokens should not be zero")
|
||||||
|
|
@ -204,3 +353,42 @@ func TestConfig_Complete(t *testing.T) {
|
||||||
t.Error("Heartbeat should be enabled by default")
|
t.Error("Heartbeat should be enabled by default")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDefaultConfig_OpenAIWebSearchEnabled(t *testing.T) {
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
if !cfg.Providers.OpenAI.WebSearch {
|
||||||
|
t.Fatal("DefaultConfig().Providers.OpenAI.WebSearch should be true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadConfig_OpenAIWebSearchDefaultsTrueWhenUnset(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
configPath := filepath.Join(dir, "config.json")
|
||||||
|
if err := os.WriteFile(configPath, []byte(`{"providers":{"openai":{"api_base":""}}}`), 0o600); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
if !cfg.Providers.OpenAI.WebSearch {
|
||||||
|
t.Fatal("OpenAI codex web search should remain true when unset in config file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadConfig_OpenAIWebSearchCanBeDisabled(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
configPath := filepath.Join(dir, "config.json")
|
||||||
|
if err := os.WriteFile(configPath, []byte(`{"providers":{"openai":{"web_search":false}}}`), 0o600); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.Providers.OpenAI.WebSearch {
|
||||||
|
t.Fatal("OpenAI codex web search should be false when disabled in config file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
292
pkg/config/defaults.go
Normal file
292
pkg/config/defaults.go
Normal file
|
|
@ -0,0 +1,292 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
// DefaultConfig returns the default configuration for PicoClaw.
|
||||||
|
func DefaultConfig() *Config {
|
||||||
|
return &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Workspace: "~/.picoclaw/workspace",
|
||||||
|
RestrictToWorkspace: true,
|
||||||
|
Provider: "",
|
||||||
|
Model: "glm-4.7",
|
||||||
|
MaxTokens: 8192,
|
||||||
|
Temperature: nil, // nil means use provider default
|
||||||
|
MaxToolIterations: 20,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Bindings: []AgentBinding{},
|
||||||
|
Session: SessionConfig{
|
||||||
|
DMScope: "main",
|
||||||
|
},
|
||||||
|
Channels: ChannelsConfig{
|
||||||
|
WhatsApp: WhatsAppConfig{
|
||||||
|
Enabled: false,
|
||||||
|
BridgeURL: "ws://localhost:3001",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Telegram: TelegramConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Token: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Feishu: FeishuConfig{
|
||||||
|
Enabled: false,
|
||||||
|
AppID: "",
|
||||||
|
AppSecret: "",
|
||||||
|
EncryptKey: "",
|
||||||
|
VerificationToken: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Discord: DiscordConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Token: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
MentionOnly: false,
|
||||||
|
},
|
||||||
|
MaixCam: MaixCamConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Host: "0.0.0.0",
|
||||||
|
Port: 18790,
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
QQ: QQConfig{
|
||||||
|
Enabled: false,
|
||||||
|
AppID: "",
|
||||||
|
AppSecret: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
DingTalk: DingTalkConfig{
|
||||||
|
Enabled: false,
|
||||||
|
ClientID: "",
|
||||||
|
ClientSecret: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
Slack: SlackConfig{
|
||||||
|
Enabled: false,
|
||||||
|
BotToken: "",
|
||||||
|
AppToken: "",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
LINE: LINEConfig{
|
||||||
|
Enabled: false,
|
||||||
|
ChannelSecret: "",
|
||||||
|
ChannelAccessToken: "",
|
||||||
|
WebhookHost: "0.0.0.0",
|
||||||
|
WebhookPort: 18791,
|
||||||
|
WebhookPath: "/webhook/line",
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
OneBot: OneBotConfig{
|
||||||
|
Enabled: false,
|
||||||
|
WSUrl: "ws://127.0.0.1:3001",
|
||||||
|
AccessToken: "",
|
||||||
|
ReconnectInterval: 5,
|
||||||
|
GroupTriggerPrefix: []string{},
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{WebSearch: true},
|
||||||
|
},
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
// ============================================
|
||||||
|
// Add your API key to the model you want to use
|
||||||
|
// ============================================
|
||||||
|
|
||||||
|
// Zhipu AI (智谱) - https://open.bigmodel.cn/usercenter/apikeys
|
||||||
|
{
|
||||||
|
ModelName: "glm-4.7",
|
||||||
|
Model: "zhipu/glm-4.7",
|
||||||
|
APIBase: "https://open.bigmodel.cn/api/paas/v4",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// OpenAI - https://platform.openai.com/api-keys
|
||||||
|
{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
APIBase: "https://api.openai.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Anthropic Claude - https://console.anthropic.com/settings/keys
|
||||||
|
{
|
||||||
|
ModelName: "claude-sonnet-4.6",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
APIBase: "https://api.anthropic.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// DeepSeek - https://platform.deepseek.com/
|
||||||
|
{
|
||||||
|
ModelName: "deepseek-chat",
|
||||||
|
Model: "deepseek/deepseek-chat",
|
||||||
|
APIBase: "https://api.deepseek.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Google Gemini - https://ai.google.dev/
|
||||||
|
{
|
||||||
|
ModelName: "gemini-2.0-flash",
|
||||||
|
Model: "gemini/gemini-2.0-flash-exp",
|
||||||
|
APIBase: "https://generativelanguage.googleapis.com/v1beta",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Qwen (通义千问) - https://dashscope.console.aliyun.com/apiKey
|
||||||
|
{
|
||||||
|
ModelName: "qwen-plus",
|
||||||
|
Model: "qwen/qwen-plus",
|
||||||
|
APIBase: "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Moonshot (月之暗面) - https://platform.moonshot.cn/console/api-keys
|
||||||
|
{
|
||||||
|
ModelName: "moonshot-v1-8k",
|
||||||
|
Model: "moonshot/moonshot-v1-8k",
|
||||||
|
APIBase: "https://api.moonshot.cn/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Groq - https://console.groq.com/keys
|
||||||
|
{
|
||||||
|
ModelName: "llama-3.3-70b",
|
||||||
|
Model: "groq/llama-3.3-70b-versatile",
|
||||||
|
APIBase: "https://api.groq.com/openai/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// OpenRouter (100+ models) - https://openrouter.ai/keys
|
||||||
|
{
|
||||||
|
ModelName: "openrouter-auto",
|
||||||
|
Model: "openrouter/auto",
|
||||||
|
APIBase: "https://openrouter.ai/api/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
ModelName: "openrouter-gpt-5.2",
|
||||||
|
Model: "openrouter/openai/gpt-5.2",
|
||||||
|
APIBase: "https://openrouter.ai/api/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// NVIDIA - https://build.nvidia.com/
|
||||||
|
{
|
||||||
|
ModelName: "nemotron-4-340b",
|
||||||
|
Model: "nvidia/nemotron-4-340b-instruct",
|
||||||
|
APIBase: "https://integrate.api.nvidia.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Cerebras - https://inference.cerebras.ai/
|
||||||
|
{
|
||||||
|
ModelName: "cerebras-llama-3.3-70b",
|
||||||
|
Model: "cerebras/llama-3.3-70b",
|
||||||
|
APIBase: "https://api.cerebras.ai/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Volcengine (火山引擎) - https://console.volcengine.com/ark
|
||||||
|
{
|
||||||
|
ModelName: "doubao-pro",
|
||||||
|
Model: "volcengine/doubao-pro-32k",
|
||||||
|
APIBase: "https://ark.cn-beijing.volces.com/api/v3",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// ShengsuanYun (神算云)
|
||||||
|
{
|
||||||
|
ModelName: "deepseek-v3",
|
||||||
|
Model: "shengsuanyun/deepseek-v3",
|
||||||
|
APIBase: "https://api.shengsuanyun.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Antigravity (Google Cloud Code Assist) - OAuth only
|
||||||
|
{
|
||||||
|
ModelName: "gemini-flash",
|
||||||
|
Model: "antigravity/gemini-3-flash",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
|
||||||
|
// GitHub Copilot - https://github.com/settings/tokens
|
||||||
|
{
|
||||||
|
ModelName: "copilot-gpt-5.2",
|
||||||
|
Model: "github-copilot/gpt-5.2",
|
||||||
|
APIBase: "http://localhost:4321",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
|
||||||
|
// Ollama (local) - https://ollama.com
|
||||||
|
{
|
||||||
|
ModelName: "llama3",
|
||||||
|
Model: "ollama/llama3",
|
||||||
|
APIBase: "http://localhost:11434/v1",
|
||||||
|
APIKey: "ollama",
|
||||||
|
},
|
||||||
|
|
||||||
|
// VLLM (local) - http://localhost:8000
|
||||||
|
{
|
||||||
|
ModelName: "local-model",
|
||||||
|
Model: "vllm/custom-model",
|
||||||
|
APIBase: "http://localhost:8000/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Gateway: GatewayConfig{
|
||||||
|
Host: "0.0.0.0",
|
||||||
|
Port: 18790,
|
||||||
|
},
|
||||||
|
Tools: ToolsConfig{
|
||||||
|
Web: WebToolsConfig{
|
||||||
|
Brave: BraveConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
DuckDuckGo: DuckDuckGoConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
Perplexity: PerplexityConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Cron: CronToolsConfig{
|
||||||
|
ExecTimeoutMinutes: 5,
|
||||||
|
},
|
||||||
|
Exec: ExecConfig{
|
||||||
|
EnableDenyPatterns: true,
|
||||||
|
},
|
||||||
|
Skills: SkillsToolsConfig{
|
||||||
|
Registries: SkillsRegistriesConfig{
|
||||||
|
ClawHub: ClawHubRegistryConfig{
|
||||||
|
Enabled: true,
|
||||||
|
BaseURL: "https://clawhub.ai",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
MaxConcurrentSearches: 2,
|
||||||
|
SearchCache: SearchCacheConfig{
|
||||||
|
MaxSize: 50,
|
||||||
|
TTLSeconds: 300,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Heartbeat: HeartbeatConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Interval: 30,
|
||||||
|
},
|
||||||
|
Devices: DevicesConfig{
|
||||||
|
Enabled: false,
|
||||||
|
MonitorUSB: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
353
pkg/config/migration.go
Normal file
353
pkg/config/migration.go
Normal file
|
|
@ -0,0 +1,353 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// buildModelWithProtocol constructs a model string with protocol prefix.
|
||||||
|
// If the model already contains a "/" (indicating it has a protocol prefix), it is returned as-is.
|
||||||
|
// Otherwise, the protocol prefix is added.
|
||||||
|
func buildModelWithProtocol(protocol, model string) string {
|
||||||
|
if strings.Contains(model, "/") {
|
||||||
|
// Model already has a protocol prefix, return as-is
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
return protocol + "/" + model
|
||||||
|
}
|
||||||
|
|
||||||
|
// providerMigrationConfig defines how to migrate a provider from old config to new format.
|
||||||
|
type providerMigrationConfig struct {
|
||||||
|
// providerNames are the possible names used in agents.defaults.provider
|
||||||
|
providerNames []string
|
||||||
|
// protocol is the protocol prefix for the model field
|
||||||
|
protocol string
|
||||||
|
// buildConfig creates the ModelConfig from ProviderConfig
|
||||||
|
buildConfig func(p ProvidersConfig) (ModelConfig, bool)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConvertProvidersToModelList converts the old ProvidersConfig to a slice of ModelConfig.
|
||||||
|
// This enables backward compatibility with existing configurations.
|
||||||
|
// It preserves the user's configured model from agents.defaults.model when possible.
|
||||||
|
func ConvertProvidersToModelList(cfg *Config) []ModelConfig {
|
||||||
|
if cfg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get user's configured provider and model
|
||||||
|
userProvider := strings.ToLower(cfg.Agents.Defaults.Provider)
|
||||||
|
userModel := cfg.Agents.Defaults.Model
|
||||||
|
|
||||||
|
p := cfg.Providers
|
||||||
|
|
||||||
|
var result []ModelConfig
|
||||||
|
|
||||||
|
// Track if we've applied the legacy model name fix (only for first provider)
|
||||||
|
legacyModelNameApplied := false
|
||||||
|
|
||||||
|
// Define migration rules for each provider
|
||||||
|
migrations := []providerMigrationConfig{
|
||||||
|
{
|
||||||
|
providerNames: []string{"openai", "gpt"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "openai",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
APIKey: p.OpenAI.APIKey,
|
||||||
|
APIBase: p.OpenAI.APIBase,
|
||||||
|
Proxy: p.OpenAI.Proxy,
|
||||||
|
AuthMethod: p.OpenAI.AuthMethod,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"anthropic", "claude"},
|
||||||
|
protocol: "anthropic",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "anthropic",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
APIKey: p.Anthropic.APIKey,
|
||||||
|
APIBase: p.Anthropic.APIBase,
|
||||||
|
Proxy: p.Anthropic.Proxy,
|
||||||
|
AuthMethod: p.Anthropic.AuthMethod,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"openrouter"},
|
||||||
|
protocol: "openrouter",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "openrouter",
|
||||||
|
Model: "openrouter/auto",
|
||||||
|
APIKey: p.OpenRouter.APIKey,
|
||||||
|
APIBase: p.OpenRouter.APIBase,
|
||||||
|
Proxy: p.OpenRouter.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"groq"},
|
||||||
|
protocol: "groq",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Groq.APIKey == "" && p.Groq.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "groq",
|
||||||
|
Model: "groq/llama-3.1-70b-versatile",
|
||||||
|
APIKey: p.Groq.APIKey,
|
||||||
|
APIBase: p.Groq.APIBase,
|
||||||
|
Proxy: p.Groq.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"zhipu", "glm"},
|
||||||
|
protocol: "zhipu",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "zhipu",
|
||||||
|
Model: "zhipu/glm-4",
|
||||||
|
APIKey: p.Zhipu.APIKey,
|
||||||
|
APIBase: p.Zhipu.APIBase,
|
||||||
|
Proxy: p.Zhipu.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"vllm"},
|
||||||
|
protocol: "vllm",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.VLLM.APIKey == "" && p.VLLM.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "vllm",
|
||||||
|
Model: "vllm/auto",
|
||||||
|
APIKey: p.VLLM.APIKey,
|
||||||
|
APIBase: p.VLLM.APIBase,
|
||||||
|
Proxy: p.VLLM.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"gemini", "google"},
|
||||||
|
protocol: "gemini",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Gemini.APIKey == "" && p.Gemini.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "gemini",
|
||||||
|
Model: "gemini/gemini-pro",
|
||||||
|
APIKey: p.Gemini.APIKey,
|
||||||
|
APIBase: p.Gemini.APIBase,
|
||||||
|
Proxy: p.Gemini.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"nvidia"},
|
||||||
|
protocol: "nvidia",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Nvidia.APIKey == "" && p.Nvidia.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "nvidia",
|
||||||
|
Model: "nvidia/meta/llama-3.1-8b-instruct",
|
||||||
|
APIKey: p.Nvidia.APIKey,
|
||||||
|
APIBase: p.Nvidia.APIBase,
|
||||||
|
Proxy: p.Nvidia.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"ollama"},
|
||||||
|
protocol: "ollama",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Ollama.APIKey == "" && p.Ollama.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "ollama",
|
||||||
|
Model: "ollama/llama3",
|
||||||
|
APIKey: p.Ollama.APIKey,
|
||||||
|
APIBase: p.Ollama.APIBase,
|
||||||
|
Proxy: p.Ollama.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"moonshot", "kimi"},
|
||||||
|
protocol: "moonshot",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Moonshot.APIKey == "" && p.Moonshot.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "moonshot",
|
||||||
|
Model: "moonshot/kimi",
|
||||||
|
APIKey: p.Moonshot.APIKey,
|
||||||
|
APIBase: p.Moonshot.APIBase,
|
||||||
|
Proxy: p.Moonshot.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"shengsuanyun"},
|
||||||
|
protocol: "shengsuanyun",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.ShengSuanYun.APIKey == "" && p.ShengSuanYun.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "shengsuanyun",
|
||||||
|
Model: "shengsuanyun/auto",
|
||||||
|
APIKey: p.ShengSuanYun.APIKey,
|
||||||
|
APIBase: p.ShengSuanYun.APIBase,
|
||||||
|
Proxy: p.ShengSuanYun.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"deepseek"},
|
||||||
|
protocol: "deepseek",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.DeepSeek.APIKey == "" && p.DeepSeek.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "deepseek",
|
||||||
|
Model: "deepseek/deepseek-chat",
|
||||||
|
APIKey: p.DeepSeek.APIKey,
|
||||||
|
APIBase: p.DeepSeek.APIBase,
|
||||||
|
Proxy: p.DeepSeek.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"cerebras"},
|
||||||
|
protocol: "cerebras",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Cerebras.APIKey == "" && p.Cerebras.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "cerebras",
|
||||||
|
Model: "cerebras/llama-3.3-70b",
|
||||||
|
APIKey: p.Cerebras.APIKey,
|
||||||
|
APIBase: p.Cerebras.APIBase,
|
||||||
|
Proxy: p.Cerebras.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"volcengine", "doubao"},
|
||||||
|
protocol: "volcengine",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.VolcEngine.APIKey == "" && p.VolcEngine.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "volcengine",
|
||||||
|
Model: "volcengine/doubao-pro",
|
||||||
|
APIKey: p.VolcEngine.APIKey,
|
||||||
|
APIBase: p.VolcEngine.APIBase,
|
||||||
|
Proxy: p.VolcEngine.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"github_copilot", "copilot"},
|
||||||
|
protocol: "github-copilot",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" && p.GitHubCopilot.ConnectMode == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "github-copilot",
|
||||||
|
Model: "github-copilot/gpt-5.2",
|
||||||
|
APIBase: p.GitHubCopilot.APIBase,
|
||||||
|
ConnectMode: p.GitHubCopilot.ConnectMode,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"antigravity"},
|
||||||
|
protocol: "antigravity",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Antigravity.APIKey == "" && p.Antigravity.AuthMethod == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "antigravity",
|
||||||
|
Model: "antigravity/gemini-2.0-flash",
|
||||||
|
APIKey: p.Antigravity.APIKey,
|
||||||
|
AuthMethod: p.Antigravity.AuthMethod,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"qwen", "tongyi"},
|
||||||
|
protocol: "qwen",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Qwen.APIKey == "" && p.Qwen.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
|
ModelName: "qwen",
|
||||||
|
Model: "qwen/qwen-max",
|
||||||
|
APIKey: p.Qwen.APIKey,
|
||||||
|
APIBase: p.Qwen.APIBase,
|
||||||
|
Proxy: p.Qwen.Proxy,
|
||||||
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process each provider migration
|
||||||
|
for _, m := range migrations {
|
||||||
|
mc, ok := m.buildConfig(p)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is the user's configured provider
|
||||||
|
if slices.Contains(m.providerNames, userProvider) && userModel != "" {
|
||||||
|
// Use the user's configured model instead of default
|
||||||
|
mc.Model = buildModelWithProtocol(m.protocol, userModel)
|
||||||
|
} else if userProvider == "" && userModel != "" && !legacyModelNameApplied {
|
||||||
|
// Legacy config: no explicit provider field but model is specified
|
||||||
|
// Use userModel as ModelName for the FIRST provider so GetModelConfig(model) can find it
|
||||||
|
// This maintains backward compatibility with old configs that relied on implicit provider selection
|
||||||
|
mc.ModelName = userModel
|
||||||
|
mc.Model = buildModelWithProtocol(m.protocol, userModel)
|
||||||
|
legacyModelNameApplied = true
|
||||||
|
}
|
||||||
|
|
||||||
|
result = append(result, mc)
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
551
pkg/config/migration_test.go
Normal file
551
pkg/config/migration_test.go
Normal file
|
|
@ -0,0 +1,551 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_OpenAI(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{
|
||||||
|
ProviderConfig: ProviderConfig{
|
||||||
|
APIKey: "sk-test-key",
|
||||||
|
APIBase: "https://custom.api.com/v1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].ModelName != "openai" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "openai")
|
||||||
|
}
|
||||||
|
if result[0].Model != "openai/gpt-5.2" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "openai/gpt-5.2")
|
||||||
|
}
|
||||||
|
if result[0].APIKey != "sk-test-key" {
|
||||||
|
t.Errorf("APIKey = %q, want %q", result[0].APIKey, "sk-test-key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Anthropic(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Anthropic: ProviderConfig{
|
||||||
|
APIKey: "ant-key",
|
||||||
|
APIBase: "https://custom.anthropic.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].ModelName != "anthropic" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "anthropic")
|
||||||
|
}
|
||||||
|
if result[0].Model != "anthropic/claude-sonnet-4.6" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "anthropic/claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Multiple(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "openai-key"}},
|
||||||
|
Groq: ProviderConfig{APIKey: "groq-key"},
|
||||||
|
Zhipu: ProviderConfig{APIKey: "zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 3 {
|
||||||
|
t.Fatalf("len(result) = %d, want 3", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check that all providers are present
|
||||||
|
found := make(map[string]bool)
|
||||||
|
for _, mc := range result {
|
||||||
|
found[mc.ModelName] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range []string{"openai", "groq", "zhipu"} {
|
||||||
|
if !found[name] {
|
||||||
|
t.Errorf("Missing provider %q in result", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Empty(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 0 {
|
||||||
|
t.Errorf("len(result) = %d, want 0", len(result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Nil(t *testing.T) {
|
||||||
|
result := ConvertProvidersToModelList(nil)
|
||||||
|
|
||||||
|
if result != nil {
|
||||||
|
t.Errorf("result = %v, want nil", result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_AllProviders(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "key1"}},
|
||||||
|
Anthropic: ProviderConfig{APIKey: "key2"},
|
||||||
|
OpenRouter: ProviderConfig{APIKey: "key3"},
|
||||||
|
Groq: ProviderConfig{APIKey: "key4"},
|
||||||
|
Zhipu: ProviderConfig{APIKey: "key5"},
|
||||||
|
VLLM: ProviderConfig{APIKey: "key6"},
|
||||||
|
Gemini: ProviderConfig{APIKey: "key7"},
|
||||||
|
Nvidia: ProviderConfig{APIKey: "key8"},
|
||||||
|
Ollama: ProviderConfig{APIKey: "key9"},
|
||||||
|
Moonshot: ProviderConfig{APIKey: "key10"},
|
||||||
|
ShengSuanYun: ProviderConfig{APIKey: "key11"},
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "key12"},
|
||||||
|
Cerebras: ProviderConfig{APIKey: "key13"},
|
||||||
|
VolcEngine: ProviderConfig{APIKey: "key14"},
|
||||||
|
GitHubCopilot: ProviderConfig{ConnectMode: "grpc"},
|
||||||
|
Antigravity: ProviderConfig{AuthMethod: "oauth"},
|
||||||
|
Qwen: ProviderConfig{APIKey: "key17"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
// All 17 providers should be converted
|
||||||
|
if len(result) != 17 {
|
||||||
|
t.Errorf("len(result) = %d, want 17", len(result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_Proxy(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{
|
||||||
|
ProviderConfig: ProviderConfig{
|
||||||
|
APIKey: "key",
|
||||||
|
Proxy: "http://proxy:8080",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Proxy != "http://proxy:8080" {
|
||||||
|
t.Errorf("Proxy = %q, want %q", result[0].Proxy, "http://proxy:8080")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_AuthMethod(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{
|
||||||
|
ProviderConfig: ProviderConfig{
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 0 {
|
||||||
|
t.Errorf("len(result) = %d, want 0 (AuthMethod alone should not create entry)", len(result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tests for preserving user's configured model during migration
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_DeepSeek(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "deepseek-reasoner",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use user's model, not default
|
||||||
|
if result[0].Model != "deepseek/deepseek-reasoner" {
|
||||||
|
t.Errorf("Model = %q, want %q (user's configured model)", result[0].Model, "deepseek/deepseek-reasoner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_OpenAI(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "openai",
|
||||||
|
Model: "gpt-4-turbo",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "sk-openai"}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "openai/gpt-4-turbo" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "openai/gpt-4-turbo")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_Anthropic(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "claude", // alternative name
|
||||||
|
Model: "claude-opus-4-20250514",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Anthropic: ProviderConfig{APIKey: "sk-ant"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "anthropic/claude-opus-4-20250514" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "anthropic/claude-opus-4-20250514")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_Qwen(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "qwen",
|
||||||
|
Model: "qwen-plus",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Qwen: ProviderConfig{APIKey: "sk-qwen"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "qwen/qwen-plus" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "qwen/qwen-plus")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_UsesDefaultWhenNoUserModel(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "", // no model specified
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use default model
|
||||||
|
if result[0].Model != "deepseek/deepseek-chat" {
|
||||||
|
t.Errorf("Model = %q, want %q (default)", result[0].Model, "deepseek/deepseek-chat")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_MultipleProviders_PreservesUserModel(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "deepseek-reasoner",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "sk-openai"}},
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 2 {
|
||||||
|
t.Fatalf("len(result) = %d, want 2", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find each provider and verify model
|
||||||
|
for _, mc := range result {
|
||||||
|
switch mc.ModelName {
|
||||||
|
case "openai":
|
||||||
|
if mc.Model != "openai/gpt-5.2" {
|
||||||
|
t.Errorf("OpenAI Model = %q, want %q (default)", mc.Model, "openai/gpt-5.2")
|
||||||
|
}
|
||||||
|
case "deepseek":
|
||||||
|
if mc.Model != "deepseek/deepseek-reasoner" {
|
||||||
|
t.Errorf("DeepSeek Model = %q, want %q (user's)", mc.Model, "deepseek/deepseek-reasoner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_ProviderNameAliases(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
providerAlias string
|
||||||
|
expectedModel string
|
||||||
|
provider ProviderConfig
|
||||||
|
}{
|
||||||
|
{"gpt", "openai/gpt-4-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"claude", "anthropic/claude-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"doubao", "volcengine/doubao-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"tongyi", "qwen/qwen-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"kimi", "moonshot/kimi-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.providerAlias, func(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: tt.providerAlias,
|
||||||
|
Model: strings.TrimPrefix(tt.expectedModel, tt.expectedModel[:strings.Index(tt.expectedModel, "/")+1]),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set the appropriate provider config
|
||||||
|
switch tt.providerAlias {
|
||||||
|
case "gpt":
|
||||||
|
cfg.Providers.OpenAI = OpenAIProviderConfig{ProviderConfig: tt.provider}
|
||||||
|
case "claude":
|
||||||
|
cfg.Providers.Anthropic = tt.provider
|
||||||
|
case "doubao":
|
||||||
|
cfg.Providers.VolcEngine = tt.provider
|
||||||
|
case "tongyi":
|
||||||
|
cfg.Providers.Qwen = tt.provider
|
||||||
|
case "kimi":
|
||||||
|
cfg.Providers.Moonshot = tt.provider
|
||||||
|
}
|
||||||
|
|
||||||
|
// Need to fix the model name in config
|
||||||
|
cfg.Agents.Defaults.Model = strings.TrimPrefix(tt.expectedModel, tt.expectedModel[:strings.Index(tt.expectedModel, "/")+1])
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract just the model ID part (after the first /)
|
||||||
|
expectedModelID := tt.expectedModel
|
||||||
|
if result[0].Model != expectedModelID {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, expectedModelID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test for backward compatibility: single provider without explicit provider field
|
||||||
|
// This matches the legacy config pattern where users only set model, not provider
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_NoProviderField_SingleProvider(t *testing.T) {
|
||||||
|
// This matches the user's actual config:
|
||||||
|
// - No provider field set
|
||||||
|
// - model = "glm-4.7"
|
||||||
|
// - Only zhipu has API key configured
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "", // Not set
|
||||||
|
Model: "glm-4.7",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Zhipu: ProviderConfig{APIKey: "test-zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelName should be the user's model value for backward compatibility
|
||||||
|
if result[0].ModelName != "glm-4.7" {
|
||||||
|
t.Errorf("ModelName = %q, want %q (user's model for backward compatibility)", result[0].ModelName, "glm-4.7")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Model should use the user's model with protocol prefix
|
||||||
|
if result[0].Model != "zhipu/glm-4.7" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "zhipu/glm-4.7")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_NoProviderField_MultipleProviders(t *testing.T) {
|
||||||
|
// When multiple providers are configured but no provider field is set,
|
||||||
|
// the FIRST provider (in migration order) will use userModel as ModelName
|
||||||
|
// for backward compatibility with legacy implicit provider selection
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "", // Not set
|
||||||
|
Model: "some-model",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: OpenAIProviderConfig{ProviderConfig: ProviderConfig{APIKey: "openai-key"}},
|
||||||
|
Zhipu: ProviderConfig{APIKey: "zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 2 {
|
||||||
|
t.Fatalf("len(result) = %d, want 2", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// The first provider (OpenAI in migration order) should use userModel as ModelName
|
||||||
|
// This ensures GetModelConfig("some-model") will find it
|
||||||
|
if result[0].ModelName != "some-model" {
|
||||||
|
t.Errorf("First provider ModelName = %q, want %q", result[0].ModelName, "some-model")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Other providers should use provider name as ModelName
|
||||||
|
if result[1].ModelName != "zhipu" {
|
||||||
|
t.Errorf("Second provider ModelName = %q, want %q", result[1].ModelName, "zhipu")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_NoProviderField_NoModel(t *testing.T) {
|
||||||
|
// Edge case: no provider, no model
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "",
|
||||||
|
Model: "",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Zhipu: ProviderConfig{APIKey: "zhipu-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use default provider name since no model is specified
|
||||||
|
if result[0].ModelName != "zhipu" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "zhipu")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tests for buildModelWithProtocol helper function
|
||||||
|
|
||||||
|
func TestBuildModelWithProtocol_NoPrefix(t *testing.T) {
|
||||||
|
result := buildModelWithProtocol("openai", "gpt-5.2")
|
||||||
|
if result != "openai/gpt-5.2" {
|
||||||
|
t.Errorf("buildModelWithProtocol(openai, gpt-5.2) = %q, want %q", result, "openai/gpt-5.2")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildModelWithProtocol_AlreadyHasPrefix(t *testing.T) {
|
||||||
|
result := buildModelWithProtocol("openrouter", "openrouter/auto")
|
||||||
|
if result != "openrouter/auto" {
|
||||||
|
t.Errorf("buildModelWithProtocol(openrouter, openrouter/auto) = %q, want %q", result, "openrouter/auto")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildModelWithProtocol_DifferentPrefix(t *testing.T) {
|
||||||
|
result := buildModelWithProtocol("anthropic", "openrouter/claude-sonnet-4.6")
|
||||||
|
if result != "openrouter/claude-sonnet-4.6" {
|
||||||
|
t.Errorf("buildModelWithProtocol(anthropic, openrouter/claude-sonnet-4.6) = %q, want %q", result, "openrouter/claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test for legacy config with protocol prefix in model name
|
||||||
|
func TestConvertProvidersToModelList_LegacyModelWithProtocolPrefix(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "", // No explicit provider
|
||||||
|
Model: "openrouter/auto", // Model already has protocol prefix
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenRouter: ProviderConfig{APIKey: "sk-or-test"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) < 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want at least 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// First provider should use userModel as ModelName for backward compatibility
|
||||||
|
if result[0].ModelName != "openrouter/auto" {
|
||||||
|
t.Errorf("ModelName = %q, want %q", result[0].ModelName, "openrouter/auto")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Model should NOT have duplicated prefix
|
||||||
|
if result[0].Model != "openrouter/auto" {
|
||||||
|
t.Errorf("Model = %q, want %q (should not duplicate prefix)", result[0].Model, "openrouter/auto")
|
||||||
|
}
|
||||||
|
}
|
||||||
235
pkg/config/model_config_test.go
Normal file
235
pkg/config/model_config_test.go
Normal file
|
|
@ -0,0 +1,235 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGetModelConfig_Found(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test-model", Model: "openai/gpt-4o", APIKey: "key1"},
|
||||||
|
{ModelName: "other-model", Model: "anthropic/claude", APIKey: "key2"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := cfg.GetModelConfig("test-model")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetModelConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if result.Model != "openai/gpt-4o" {
|
||||||
|
t.Errorf("Model = %q, want %q", result.Model, "openai/gpt-4o")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_NotFound(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test-model", Model: "openai/gpt-4o", APIKey: "key1"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := cfg.GetModelConfig("nonexistent")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("GetModelConfig() expected error for nonexistent model")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_EmptyList(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := cfg.GetModelConfig("any-model")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("GetModelConfig() expected error for empty model list")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_RoundRobin(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "lb-model", Model: "openai/gpt-4o-1", APIKey: "key1"},
|
||||||
|
{ModelName: "lb-model", Model: "openai/gpt-4o-2", APIKey: "key2"},
|
||||||
|
{ModelName: "lb-model", Model: "openai/gpt-4o-3", APIKey: "key3"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test round-robin distribution
|
||||||
|
results := make(map[string]int)
|
||||||
|
for i := 0; i < 30; i++ {
|
||||||
|
result, err := cfg.GetModelConfig("lb-model")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetModelConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
results[result.Model]++
|
||||||
|
}
|
||||||
|
|
||||||
|
// Each model should appear roughly 10 times (30 calls / 3 models)
|
||||||
|
for model, count := range results {
|
||||||
|
if count < 5 || count > 15 {
|
||||||
|
t.Errorf("Model %s appeared %d times, expected ~10", model, count)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetModelConfig_Concurrent(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "concurrent-model", Model: "openai/gpt-4o-1", APIKey: "key1"},
|
||||||
|
{ModelName: "concurrent-model", Model: "openai/gpt-4o-2", APIKey: "key2"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
const goroutines = 100
|
||||||
|
const iterations = 10
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
errors := make(chan error, goroutines*iterations)
|
||||||
|
|
||||||
|
for i := 0; i < goroutines; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for j := 0; j < iterations; j++ {
|
||||||
|
_, err := cfg.GetModelConfig("concurrent-model")
|
||||||
|
if err != nil {
|
||||||
|
errors <- err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
close(errors)
|
||||||
|
|
||||||
|
for err := range errors {
|
||||||
|
t.Errorf("Concurrent GetModelConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelConfig_Validate(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
config ModelConfig
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "valid config",
|
||||||
|
config: ModelConfig{
|
||||||
|
ModelName: "test",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
},
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing model_name",
|
||||||
|
config: ModelConfig{
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
},
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing model",
|
||||||
|
config: ModelConfig{
|
||||||
|
ModelName: "test",
|
||||||
|
},
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty config",
|
||||||
|
config: ModelConfig{},
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := tt.config.Validate()
|
||||||
|
if (err != nil) != tt.wantErr {
|
||||||
|
t.Errorf("Validate() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfig_ValidateModelList(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
config *Config
|
||||||
|
wantErr bool
|
||||||
|
errMsg string // partial error message to check
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "valid list",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test1", Model: "openai/gpt-4o"},
|
||||||
|
{ModelName: "test2", Model: "anthropic/claude"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid entry",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "test1", Model: "openai/gpt-4o"},
|
||||||
|
{ModelName: "", Model: "anthropic/claude"}, // missing model_name
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: true,
|
||||||
|
errMsg: "model_name is required",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty list",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{},
|
||||||
|
},
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Load balancing: multiple entries with same model_name are allowed
|
||||||
|
name: "duplicate model_name for load balancing",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "gpt-4", Model: "openai/gpt-4o", APIKey: "key1"},
|
||||||
|
{ModelName: "gpt-4", Model: "openai/gpt-4-turbo", APIKey: "key2"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: false, // Changed: duplicates are allowed for load balancing
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Load balancing: non-adjacent entries with same model_name are also allowed
|
||||||
|
name: "duplicate model_name non-adjacent for load balancing",
|
||||||
|
config: &Config{
|
||||||
|
ModelList: []ModelConfig{
|
||||||
|
{ModelName: "model-a", Model: "openai/gpt-4o"},
|
||||||
|
{ModelName: "model-b", Model: "anthropic/claude"},
|
||||||
|
{ModelName: "model-a", Model: "openai/gpt-4-turbo"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
wantErr: false, // Changed: duplicates are allowed for load balancing
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := tt.config.ValidateModelList()
|
||||||
|
if (err != nil) != tt.wantErr {
|
||||||
|
t.Errorf("ValidateModelList() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
}
|
||||||
|
if err != nil && tt.errMsg != "" {
|
||||||
|
if !strings.Contains(err.Error(), tt.errMsg) {
|
||||||
|
t.Errorf("ValidateModelList() error = %v, want error containing %q", err, tt.errMsg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,15 +1,16 @@
|
||||||
// Package constants provides shared constants across the codebase.
|
// Package constants provides shared constants across the codebase.
|
||||||
package constants
|
package constants
|
||||||
|
|
||||||
// InternalChannels defines channels that are used for internal communication
|
// internalChannels defines channels that are used for internal communication
|
||||||
// and should not be exposed to external users or recorded as last active channel.
|
// and should not be exposed to external users or recorded as last active channel.
|
||||||
var InternalChannels = map[string]bool{
|
var internalChannels = map[string]struct{}{
|
||||||
"cli": true,
|
"cli": {},
|
||||||
"system": true,
|
"system": {},
|
||||||
"subagent": true,
|
"subagent": {},
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsInternalChannel returns true if the channel is an internal channel.
|
// IsInternalChannel returns true if the channel is an internal channel.
|
||||||
func IsInternalChannel(channel string) bool {
|
func IsInternalChannel(channel string) bool {
|
||||||
return InternalChannels[channel]
|
_, found := internalChannels[channel]
|
||||||
|
return found
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,13 +12,16 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
var supportedProviders = map[string]bool{
|
var supportedProviders = map[string]bool{
|
||||||
"anthropic": true,
|
"anthropic": true,
|
||||||
"openai": true,
|
"openai": true,
|
||||||
"openrouter": true,
|
"openrouter": true,
|
||||||
"groq": true,
|
"groq": true,
|
||||||
"zhipu": true,
|
"zhipu": true,
|
||||||
"vllm": true,
|
"vllm": true,
|
||||||
"gemini": true,
|
"gemini": true,
|
||||||
|
"qwen": true,
|
||||||
|
"deepseek": true,
|
||||||
|
"github_copilot": true,
|
||||||
}
|
}
|
||||||
|
|
||||||
var supportedChannels = map[string]bool{
|
var supportedChannels = map[string]bool{
|
||||||
|
|
@ -76,7 +79,7 @@ func ConvertConfig(data map[string]interface{}) (*config.Config, []string, error
|
||||||
cfg.Agents.Defaults.MaxTokens = int(v)
|
cfg.Agents.Defaults.MaxTokens = int(v)
|
||||||
}
|
}
|
||||||
if v, ok := getFloat(defaults, "temperature"); ok {
|
if v, ok := getFloat(defaults, "temperature"); ok {
|
||||||
cfg.Agents.Defaults.Temperature = v
|
cfg.Agents.Defaults.Temperature = &v
|
||||||
}
|
}
|
||||||
if v, ok := getFloat(defaults, "max_tool_iterations"); ok {
|
if v, ok := getFloat(defaults, "max_tool_iterations"); ok {
|
||||||
cfg.Agents.Defaults.MaxToolIterations = int(v)
|
cfg.Agents.Defaults.MaxToolIterations = int(v)
|
||||||
|
|
@ -108,7 +111,10 @@ func ConvertConfig(data map[string]interface{}) (*config.Config, []string, error
|
||||||
case "anthropic":
|
case "anthropic":
|
||||||
cfg.Providers.Anthropic = pc
|
cfg.Providers.Anthropic = pc
|
||||||
case "openai":
|
case "openai":
|
||||||
cfg.Providers.OpenAI = pc
|
cfg.Providers.OpenAI = config.OpenAIProviderConfig{
|
||||||
|
ProviderConfig: pc,
|
||||||
|
WebSearch: getBoolOrDefault(pMap, "web_search", true),
|
||||||
|
}
|
||||||
case "openrouter":
|
case "openrouter":
|
||||||
cfg.Providers.OpenRouter = pc
|
cfg.Providers.OpenRouter = pc
|
||||||
case "groq":
|
case "groq":
|
||||||
|
|
@ -253,6 +259,15 @@ func MergeConfig(existing, incoming *config.Config) *config.Config {
|
||||||
if existing.Providers.Gemini.APIKey == "" {
|
if existing.Providers.Gemini.APIKey == "" {
|
||||||
existing.Providers.Gemini = incoming.Providers.Gemini
|
existing.Providers.Gemini = incoming.Providers.Gemini
|
||||||
}
|
}
|
||||||
|
if existing.Providers.DeepSeek.APIKey == "" {
|
||||||
|
existing.Providers.DeepSeek = incoming.Providers.DeepSeek
|
||||||
|
}
|
||||||
|
if existing.Providers.GitHubCopilot.APIBase == "" {
|
||||||
|
existing.Providers.GitHubCopilot = incoming.Providers.GitHubCopilot
|
||||||
|
}
|
||||||
|
if existing.Providers.Qwen.APIKey == "" {
|
||||||
|
existing.Providers.Qwen = incoming.Providers.Qwen
|
||||||
|
}
|
||||||
|
|
||||||
if !existing.Channels.Telegram.Enabled && incoming.Channels.Telegram.Enabled {
|
if !existing.Channels.Telegram.Enabled && incoming.Channels.Telegram.Enabled {
|
||||||
existing.Channels.Telegram = incoming.Channels.Telegram
|
existing.Channels.Telegram = incoming.Channels.Telegram
|
||||||
|
|
@ -363,6 +378,13 @@ func getBool(data map[string]interface{}, key string) (bool, bool) {
|
||||||
return b, ok
|
return b, ok
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func getBoolOrDefault(data map[string]interface{}, key string, defaultVal bool) bool {
|
||||||
|
if v, ok := getBool(data, key); ok {
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
return defaultVal
|
||||||
|
}
|
||||||
|
|
||||||
func getStringSlice(data map[string]interface{}, key string) []string {
|
func getStringSlice(data map[string]interface{}, key string) []string {
|
||||||
v, ok := data[key]
|
v, ok := data[key]
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|
|
||||||
|
|
@ -180,8 +180,8 @@ func TestConvertConfig(t *testing.T) {
|
||||||
t.Run("unsupported provider warning", func(t *testing.T) {
|
t.Run("unsupported provider warning", func(t *testing.T) {
|
||||||
data := map[string]interface{}{
|
data := map[string]interface{}{
|
||||||
"providers": map[string]interface{}{
|
"providers": map[string]interface{}{
|
||||||
"deepseek": map[string]interface{}{
|
"unknown_provider": map[string]interface{}{
|
||||||
"api_key": "sk-deep-test",
|
"api_key": "sk-test",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
@ -193,7 +193,7 @@ func TestConvertConfig(t *testing.T) {
|
||||||
if len(warnings) != 1 {
|
if len(warnings) != 1 {
|
||||||
t.Fatalf("expected 1 warning, got %d", len(warnings))
|
t.Fatalf("expected 1 warning, got %d", len(warnings))
|
||||||
}
|
}
|
||||||
if warnings[0] != "Provider 'deepseek' not supported in PicoClaw, skipping" {
|
if warnings[0] != "Provider 'unknown_provider' not supported in PicoClaw, skipping" {
|
||||||
t.Errorf("unexpected warning: %s", warnings[0])
|
t.Errorf("unexpected warning: %s", warnings[0])
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
@ -275,8 +275,11 @@ func TestConvertConfig(t *testing.T) {
|
||||||
if cfg.Agents.Defaults.MaxTokens != 4096 {
|
if cfg.Agents.Defaults.MaxTokens != 4096 {
|
||||||
t.Errorf("MaxTokens = %d, want %d", cfg.Agents.Defaults.MaxTokens, 4096)
|
t.Errorf("MaxTokens = %d, want %d", cfg.Agents.Defaults.MaxTokens, 4096)
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Temperature != 0.5 {
|
if cfg.Agents.Defaults.Temperature == nil {
|
||||||
t.Errorf("Temperature = %f, want %f", cfg.Agents.Defaults.Temperature, 0.5)
|
t.Fatalf("Temperature is nil, want %f", 0.5)
|
||||||
|
}
|
||||||
|
if *cfg.Agents.Defaults.Temperature != 0.5 {
|
||||||
|
t.Errorf("Temperature = %f, want %f", *cfg.Agents.Defaults.Temperature, 0.5)
|
||||||
}
|
}
|
||||||
if cfg.Agents.Defaults.Workspace != "~/.picoclaw/workspace" {
|
if cfg.Agents.Defaults.Workspace != "~/.picoclaw/workspace" {
|
||||||
t.Errorf("Workspace = %q, want %q", cfg.Agents.Defaults.Workspace, "~/.picoclaw/workspace")
|
t.Errorf("Workspace = %q, want %q", cfg.Agents.Defaults.Workspace, "~/.picoclaw/workspace")
|
||||||
|
|
@ -299,6 +302,24 @@ func TestConvertConfig(t *testing.T) {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSupportedProvidersCompatibility(t *testing.T) {
|
||||||
|
expected := []string{
|
||||||
|
"anthropic",
|
||||||
|
"openai",
|
||||||
|
"openrouter",
|
||||||
|
"groq",
|
||||||
|
"zhipu",
|
||||||
|
"vllm",
|
||||||
|
"gemini",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, provider := range expected {
|
||||||
|
if !supportedProviders[provider] {
|
||||||
|
t.Fatalf("supportedProviders missing expected key %q", provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestMergeConfig(t *testing.T) {
|
func TestMergeConfig(t *testing.T) {
|
||||||
t.Run("fills empty fields", func(t *testing.T) {
|
t.Run("fills empty fields", func(t *testing.T) {
|
||||||
existing := config.DefaultConfig()
|
existing := config.DefaultConfig()
|
||||||
|
|
|
||||||
248
pkg/providers/anthropic/provider.go
Normal file
248
pkg/providers/anthropic/provider.go
Normal file
|
|
@ -0,0 +1,248 @@
|
||||||
|
package anthropicprovider
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/anthropics/anthropic-sdk-go"
|
||||||
|
"github.com/anthropics/anthropic-sdk-go/option"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ToolCall = protocoltypes.ToolCall
|
||||||
|
type FunctionCall = protocoltypes.FunctionCall
|
||||||
|
type LLMResponse = protocoltypes.LLMResponse
|
||||||
|
type UsageInfo = protocoltypes.UsageInfo
|
||||||
|
type Message = protocoltypes.Message
|
||||||
|
type ToolDefinition = protocoltypes.ToolDefinition
|
||||||
|
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
||||||
|
|
||||||
|
const defaultBaseURL = "https://api.anthropic.com"
|
||||||
|
|
||||||
|
type Provider struct {
|
||||||
|
client *anthropic.Client
|
||||||
|
tokenSource func() (string, error)
|
||||||
|
baseURL string
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProvider(token string) *Provider {
|
||||||
|
return NewProviderWithBaseURL(token, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProviderWithBaseURL(token, apiBase string) *Provider {
|
||||||
|
baseURL := normalizeBaseURL(apiBase)
|
||||||
|
client := anthropic.NewClient(
|
||||||
|
option.WithAuthToken(token),
|
||||||
|
option.WithBaseURL(baseURL),
|
||||||
|
)
|
||||||
|
return &Provider{
|
||||||
|
client: &client,
|
||||||
|
baseURL: baseURL,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProviderWithClient(client *anthropic.Client) *Provider {
|
||||||
|
return &Provider{
|
||||||
|
client: client,
|
||||||
|
baseURL: defaultBaseURL,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProviderWithTokenSource(token string, tokenSource func() (string, error)) *Provider {
|
||||||
|
return NewProviderWithTokenSourceAndBaseURL(token, tokenSource, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProviderWithTokenSourceAndBaseURL(token string, tokenSource func() (string, error), apiBase string) *Provider {
|
||||||
|
p := NewProviderWithBaseURL(token, apiBase)
|
||||||
|
p.tokenSource = tokenSource
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Provider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
|
var opts []option.RequestOption
|
||||||
|
if p.tokenSource != nil {
|
||||||
|
tok, err := p.tokenSource()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("refreshing token: %w", err)
|
||||||
|
}
|
||||||
|
opts = append(opts, option.WithAuthToken(tok))
|
||||||
|
}
|
||||||
|
|
||||||
|
params, err := buildParams(messages, tools, model, options)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := p.client.Messages.New(ctx, params, opts...)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("claude API call: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return parseResponse(resp), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Provider) GetDefaultModel() string {
|
||||||
|
return "claude-sonnet-4.6"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Provider) BaseURL() string {
|
||||||
|
return p.baseURL
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildParams(messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (anthropic.MessageNewParams, error) {
|
||||||
|
var system []anthropic.TextBlockParam
|
||||||
|
var anthropicMessages []anthropic.MessageParam
|
||||||
|
|
||||||
|
for _, msg := range messages {
|
||||||
|
switch msg.Role {
|
||||||
|
case "system":
|
||||||
|
system = append(system, anthropic.TextBlockParam{Text: msg.Content})
|
||||||
|
case "user":
|
||||||
|
if msg.ToolCallID != "" {
|
||||||
|
anthropicMessages = append(anthropicMessages,
|
||||||
|
anthropic.NewUserMessage(anthropic.NewToolResultBlock(msg.ToolCallID, msg.Content, false)),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
anthropicMessages = append(anthropicMessages,
|
||||||
|
anthropic.NewUserMessage(anthropic.NewTextBlock(msg.Content)),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
case "assistant":
|
||||||
|
if len(msg.ToolCalls) > 0 {
|
||||||
|
var blocks []anthropic.ContentBlockParamUnion
|
||||||
|
if msg.Content != "" {
|
||||||
|
blocks = append(blocks, anthropic.NewTextBlock(msg.Content))
|
||||||
|
}
|
||||||
|
for _, tc := range msg.ToolCalls {
|
||||||
|
blocks = append(blocks, anthropic.NewToolUseBlock(tc.ID, tc.Arguments, tc.Name))
|
||||||
|
}
|
||||||
|
anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...))
|
||||||
|
} else {
|
||||||
|
anthropicMessages = append(anthropicMessages,
|
||||||
|
anthropic.NewAssistantMessage(anthropic.NewTextBlock(msg.Content)),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
case "tool":
|
||||||
|
anthropicMessages = append(anthropicMessages,
|
||||||
|
anthropic.NewUserMessage(anthropic.NewToolResultBlock(msg.ToolCallID, msg.Content, false)),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
maxTokens := int64(4096)
|
||||||
|
if mt, ok := options["max_tokens"].(int); ok {
|
||||||
|
maxTokens = int64(mt)
|
||||||
|
}
|
||||||
|
|
||||||
|
params := anthropic.MessageNewParams{
|
||||||
|
Model: anthropic.Model(model),
|
||||||
|
Messages: anthropicMessages,
|
||||||
|
MaxTokens: maxTokens,
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(system) > 0 {
|
||||||
|
params.System = system
|
||||||
|
}
|
||||||
|
|
||||||
|
if temp, ok := options["temperature"].(float64); ok {
|
||||||
|
params.Temperature = anthropic.Float(temp)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(tools) > 0 {
|
||||||
|
params.Tools = translateTools(tools)
|
||||||
|
}
|
||||||
|
|
||||||
|
return params, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func translateTools(tools []ToolDefinition) []anthropic.ToolUnionParam {
|
||||||
|
result := make([]anthropic.ToolUnionParam, 0, len(tools))
|
||||||
|
for _, t := range tools {
|
||||||
|
tool := anthropic.ToolParam{
|
||||||
|
Name: t.Function.Name,
|
||||||
|
InputSchema: anthropic.ToolInputSchemaParam{
|
||||||
|
Properties: t.Function.Parameters["properties"],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if desc := t.Function.Description; desc != "" {
|
||||||
|
tool.Description = anthropic.String(desc)
|
||||||
|
}
|
||||||
|
if req, ok := t.Function.Parameters["required"].([]interface{}); ok {
|
||||||
|
required := make([]string, 0, len(req))
|
||||||
|
for _, r := range req {
|
||||||
|
if s, ok := r.(string); ok {
|
||||||
|
required = append(required, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tool.InputSchema.Required = required
|
||||||
|
}
|
||||||
|
result = append(result, anthropic.ToolUnionParam{OfTool: &tool})
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseResponse(resp *anthropic.Message) *LLMResponse {
|
||||||
|
var content string
|
||||||
|
var toolCalls []ToolCall
|
||||||
|
|
||||||
|
for _, block := range resp.Content {
|
||||||
|
switch block.Type {
|
||||||
|
case "text":
|
||||||
|
tb := block.AsText()
|
||||||
|
content += tb.Text
|
||||||
|
case "tool_use":
|
||||||
|
tu := block.AsToolUse()
|
||||||
|
var args map[string]interface{}
|
||||||
|
if err := json.Unmarshal(tu.Input, &args); err != nil {
|
||||||
|
log.Printf("anthropic: failed to decode tool call input for %q: %v", tu.Name, err)
|
||||||
|
args = map[string]interface{}{"raw": string(tu.Input)}
|
||||||
|
}
|
||||||
|
toolCalls = append(toolCalls, ToolCall{
|
||||||
|
ID: tu.ID,
|
||||||
|
Name: tu.Name,
|
||||||
|
Arguments: args,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
finishReason := "stop"
|
||||||
|
switch resp.StopReason {
|
||||||
|
case anthropic.StopReasonToolUse:
|
||||||
|
finishReason = "tool_calls"
|
||||||
|
case anthropic.StopReasonMaxTokens:
|
||||||
|
finishReason = "length"
|
||||||
|
case anthropic.StopReasonEndTurn:
|
||||||
|
finishReason = "stop"
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: content,
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: finishReason,
|
||||||
|
Usage: &UsageInfo{
|
||||||
|
PromptTokens: int(resp.Usage.InputTokens),
|
||||||
|
CompletionTokens: int(resp.Usage.OutputTokens),
|
||||||
|
TotalTokens: int(resp.Usage.InputTokens + resp.Usage.OutputTokens),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeBaseURL(apiBase string) string {
|
||||||
|
base := strings.TrimSpace(apiBase)
|
||||||
|
if base == "" {
|
||||||
|
return defaultBaseURL
|
||||||
|
}
|
||||||
|
|
||||||
|
base = strings.TrimRight(base, "/")
|
||||||
|
if strings.HasSuffix(base, "/v1") {
|
||||||
|
base = strings.TrimSuffix(base, "/v1")
|
||||||
|
}
|
||||||
|
if base == "" {
|
||||||
|
return defaultBaseURL
|
||||||
|
}
|
||||||
|
|
||||||
|
return base
|
||||||
|
}
|
||||||
265
pkg/providers/anthropic/provider_test.go
Normal file
265
pkg/providers/anthropic/provider_test.go
Normal file
|
|
@ -0,0 +1,265 @@
|
||||||
|
package anthropicprovider
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/anthropics/anthropic-sdk-go"
|
||||||
|
anthropicoption "github.com/anthropics/anthropic-sdk-go/option"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBuildParams_BasicMessage(t *testing.T) {
|
||||||
|
messages := []Message{
|
||||||
|
{Role: "user", Content: "Hello"},
|
||||||
|
}
|
||||||
|
params, err := buildParams(messages, nil, "claude-sonnet-4.6", map[string]interface{}{
|
||||||
|
"max_tokens": 1024,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
|
}
|
||||||
|
if string(params.Model) != "claude-sonnet-4.6" {
|
||||||
|
t.Errorf("Model = %q, want %q", params.Model, "claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
if params.MaxTokens != 1024 {
|
||||||
|
t.Errorf("MaxTokens = %d, want 1024", params.MaxTokens)
|
||||||
|
}
|
||||||
|
if len(params.Messages) != 1 {
|
||||||
|
t.Fatalf("len(Messages) = %d, want 1", len(params.Messages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildParams_SystemMessage(t *testing.T) {
|
||||||
|
messages := []Message{
|
||||||
|
{Role: "system", Content: "You are helpful"},
|
||||||
|
{Role: "user", Content: "Hi"},
|
||||||
|
}
|
||||||
|
params, err := buildParams(messages, nil, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
|
}
|
||||||
|
if len(params.System) != 1 {
|
||||||
|
t.Fatalf("len(System) = %d, want 1", len(params.System))
|
||||||
|
}
|
||||||
|
if params.System[0].Text != "You are helpful" {
|
||||||
|
t.Errorf("System[0].Text = %q, want %q", params.System[0].Text, "You are helpful")
|
||||||
|
}
|
||||||
|
if len(params.Messages) != 1 {
|
||||||
|
t.Fatalf("len(Messages) = %d, want 1", len(params.Messages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildParams_ToolCallMessage(t *testing.T) {
|
||||||
|
messages := []Message{
|
||||||
|
{Role: "user", Content: "What's the weather?"},
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: "",
|
||||||
|
ToolCalls: []ToolCall{
|
||||||
|
{
|
||||||
|
ID: "call_1",
|
||||||
|
Name: "get_weather",
|
||||||
|
Arguments: map[string]interface{}{"city": "SF"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
|
||||||
|
}
|
||||||
|
params, err := buildParams(messages, nil, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
|
}
|
||||||
|
if len(params.Messages) != 3 {
|
||||||
|
t.Fatalf("len(Messages) = %d, want 3", len(params.Messages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildParams_WithTools(t *testing.T) {
|
||||||
|
tools := []ToolDefinition{
|
||||||
|
{
|
||||||
|
Type: "function",
|
||||||
|
Function: ToolFunctionDefinition{
|
||||||
|
Name: "get_weather",
|
||||||
|
Description: "Get weather for a city",
|
||||||
|
Parameters: map[string]interface{}{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]interface{}{
|
||||||
|
"city": map[string]interface{}{"type": "string"},
|
||||||
|
},
|
||||||
|
"required": []interface{}{"city"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
params, err := buildParams([]Message{{Role: "user", Content: "Hi"}}, tools, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("buildParams() error: %v", err)
|
||||||
|
}
|
||||||
|
if len(params.Tools) != 1 {
|
||||||
|
t.Fatalf("len(Tools) = %d, want 1", len(params.Tools))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseResponse_TextOnly(t *testing.T) {
|
||||||
|
resp := &anthropic.Message{
|
||||||
|
Content: []anthropic.ContentBlockUnion{},
|
||||||
|
Usage: anthropic.Usage{
|
||||||
|
InputTokens: 10,
|
||||||
|
OutputTokens: 20,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result := parseResponse(resp)
|
||||||
|
if result.Usage.PromptTokens != 10 {
|
||||||
|
t.Errorf("PromptTokens = %d, want 10", result.Usage.PromptTokens)
|
||||||
|
}
|
||||||
|
if result.Usage.CompletionTokens != 20 {
|
||||||
|
t.Errorf("CompletionTokens = %d, want 20", result.Usage.CompletionTokens)
|
||||||
|
}
|
||||||
|
if result.FinishReason != "stop" {
|
||||||
|
t.Errorf("FinishReason = %q, want %q", result.FinishReason, "stop")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseResponse_StopReasons(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
stopReason anthropic.StopReason
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{anthropic.StopReasonEndTurn, "stop"},
|
||||||
|
{anthropic.StopReasonMaxTokens, "length"},
|
||||||
|
{anthropic.StopReasonToolUse, "tool_calls"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
resp := &anthropic.Message{
|
||||||
|
StopReason: tt.stopReason,
|
||||||
|
}
|
||||||
|
result := parseResponse(resp)
|
||||||
|
if result.FinishReason != tt.want {
|
||||||
|
t.Errorf("StopReason %q: FinishReason = %q, want %q", tt.stopReason, result.FinishReason, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProvider_ChatRoundTrip(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.URL.Path != "/v1/messages" {
|
||||||
|
http.Error(w, "not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if r.Header.Get("Authorization") != "Bearer test-token" {
|
||||||
|
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var reqBody map[string]interface{}
|
||||||
|
json.NewDecoder(r.Body).Decode(&reqBody)
|
||||||
|
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"id": "msg_test",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": reqBody["model"],
|
||||||
|
"stop_reason": "end_turn",
|
||||||
|
"content": []map[string]interface{}{
|
||||||
|
{"type": "text", "text": "Hello! How can I help you?"},
|
||||||
|
},
|
||||||
|
"usage": map[string]interface{}{
|
||||||
|
"input_tokens": 15,
|
||||||
|
"output_tokens": 8,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
provider := NewProviderWithClient(createAnthropicTestClient(server.URL, "test-token"))
|
||||||
|
messages := []Message{{Role: "user", Content: "Hello"}}
|
||||||
|
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4.6", map[string]interface{}{"max_tokens": 1024})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error: %v", err)
|
||||||
|
}
|
||||||
|
if resp.Content != "Hello! How can I help you?" {
|
||||||
|
t.Errorf("Content = %q, want %q", resp.Content, "Hello! How can I help you?")
|
||||||
|
}
|
||||||
|
if resp.FinishReason != "stop" {
|
||||||
|
t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "stop")
|
||||||
|
}
|
||||||
|
if resp.Usage.PromptTokens != 15 {
|
||||||
|
t.Errorf("PromptTokens = %d, want 15", resp.Usage.PromptTokens)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProvider_GetDefaultModel(t *testing.T) {
|
||||||
|
p := NewProvider("test-token")
|
||||||
|
if got := p.GetDefaultModel(); got != "claude-sonnet-4.6" {
|
||||||
|
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProvider_NewProviderWithBaseURL_NormalizesV1Suffix(t *testing.T) {
|
||||||
|
p := NewProviderWithBaseURL("token", "https://api.anthropic.com/v1/")
|
||||||
|
if got := p.BaseURL(); got != "https://api.anthropic.com" {
|
||||||
|
t.Fatalf("BaseURL() = %q, want %q", got, "https://api.anthropic.com")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProvider_ChatUsesTokenSource(t *testing.T) {
|
||||||
|
var requests int32
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.URL.Path != "/v1/messages" {
|
||||||
|
http.Error(w, "not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
atomic.AddInt32(&requests, 1)
|
||||||
|
|
||||||
|
if got := r.Header.Get("Authorization"); got != "Bearer refreshed-token" {
|
||||||
|
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var reqBody map[string]interface{}
|
||||||
|
json.NewDecoder(r.Body).Decode(&reqBody)
|
||||||
|
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"id": "msg_test",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": reqBody["model"],
|
||||||
|
"stop_reason": "end_turn",
|
||||||
|
"content": []map[string]interface{}{
|
||||||
|
{"type": "text", "text": "ok"},
|
||||||
|
},
|
||||||
|
"usage": map[string]interface{}{
|
||||||
|
"input_tokens": 1,
|
||||||
|
"output_tokens": 1,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProviderWithTokenSourceAndBaseURL("stale-token", func() (string, error) {
|
||||||
|
return "refreshed-token", nil
|
||||||
|
}, server.URL)
|
||||||
|
|
||||||
|
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hello"}}, nil, "claude-sonnet-4.6", map[string]interface{}{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error: %v", err)
|
||||||
|
}
|
||||||
|
if got := atomic.LoadInt32(&requests); got != 1 {
|
||||||
|
t.Fatalf("requests = %d, want 1", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func createAnthropicTestClient(baseURL, token string) *anthropic.Client {
|
||||||
|
c := anthropic.NewClient(
|
||||||
|
anthropicoption.WithAuthToken(token),
|
||||||
|
anthropicoption.WithBaseURL(baseURL),
|
||||||
|
)
|
||||||
|
return &c
|
||||||
|
}
|
||||||
827
pkg/providers/antigravity_provider.go
Normal file
827
pkg/providers/antigravity_provider.go
Normal file
|
|
@ -0,0 +1,827 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"math/rand"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
antigravityBaseURL = "https://cloudcode-pa.googleapis.com"
|
||||||
|
antigravityDefaultModel = "gemini-3-flash"
|
||||||
|
antigravityUserAgent = "antigravity"
|
||||||
|
antigravityXGoogClient = "google-cloud-sdk vscode_cloudshelleditor/0.1"
|
||||||
|
antigravityVersion = "1.15.8"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AntigravityProvider implements LLMProvider using Google's Cloud Code Assist (Antigravity) API.
|
||||||
|
// This provider authenticates via Google OAuth and provides access to models like Claude and Gemini
|
||||||
|
// through Google's infrastructure.
|
||||||
|
type AntigravityProvider struct {
|
||||||
|
tokenSource func() (string, string, error) // Returns (accessToken, projectID, error)
|
||||||
|
httpClient *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAntigravityProvider creates a new Antigravity provider using stored auth credentials.
|
||||||
|
func NewAntigravityProvider() *AntigravityProvider {
|
||||||
|
return &AntigravityProvider{
|
||||||
|
tokenSource: createAntigravityTokenSource(),
|
||||||
|
httpClient: &http.Client{
|
||||||
|
Timeout: 120 * time.Second,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Chat implements LLMProvider.Chat using the Cloud Code Assist v1internal API.
|
||||||
|
// The v1internal endpoint wraps the standard Gemini request in an envelope with
|
||||||
|
// project, model, request, requestType, userAgent, and requestId fields.
|
||||||
|
func (p *AntigravityProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
|
accessToken, projectID, err := p.tokenSource()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("antigravity auth: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if model == "" || model == "antigravity" || model == "google-antigravity" {
|
||||||
|
model = antigravityDefaultModel
|
||||||
|
}
|
||||||
|
// Strip provider prefixes if present
|
||||||
|
model = strings.TrimPrefix(model, "google-antigravity/")
|
||||||
|
model = strings.TrimPrefix(model, "antigravity/")
|
||||||
|
|
||||||
|
logger.DebugCF("provider.antigravity", "Starting chat", map[string]interface{}{
|
||||||
|
"model": model,
|
||||||
|
"project": projectID,
|
||||||
|
"requestId": fmt.Sprintf("agent-%d-%s", time.Now().UnixMilli(), randomString(9)),
|
||||||
|
})
|
||||||
|
|
||||||
|
// Build the inner Gemini-format request
|
||||||
|
innerRequest := p.buildRequest(messages, tools, model, options)
|
||||||
|
|
||||||
|
// Wrap in v1internal envelope (matches pi-ai SDK format)
|
||||||
|
envelope := map[string]interface{}{
|
||||||
|
"project": projectID,
|
||||||
|
"model": model,
|
||||||
|
"request": innerRequest,
|
||||||
|
"requestType": "agent",
|
||||||
|
"userAgent": antigravityUserAgent,
|
||||||
|
"requestId": fmt.Sprintf("agent-%d-%s", time.Now().UnixMilli(), randomString(9)),
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyBytes, err := json.Marshal(envelope)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("marshaling request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build API URL — uses Cloud Code Assist v1internal streaming endpoint
|
||||||
|
apiURL := fmt.Sprintf("%s/v1internal:streamGenerateContent?alt=sse", antigravityBaseURL)
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", apiURL, bytes.NewReader(bodyBytes))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Headers matching the pi-ai SDK antigravity format
|
||||||
|
clientMetadata, _ := json.Marshal(map[string]string{
|
||||||
|
"ideType": "IDE_UNSPECIFIED",
|
||||||
|
"platform": "PLATFORM_UNSPECIFIED",
|
||||||
|
"pluginType": "GEMINI",
|
||||||
|
})
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Accept", "text/event-stream")
|
||||||
|
req.Header.Set("User-Agent", fmt.Sprintf("antigravity/%s linux/amd64", antigravityVersion))
|
||||||
|
req.Header.Set("X-Goog-Api-Client", antigravityXGoogClient)
|
||||||
|
req.Header.Set("Client-Metadata", string(clientMetadata))
|
||||||
|
|
||||||
|
resp, err := p.httpClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("antigravity API call: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
respBody, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("reading response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
logger.ErrorCF("provider.antigravity", "API call failed", map[string]interface{}{
|
||||||
|
"status_code": resp.StatusCode,
|
||||||
|
"response": string(respBody),
|
||||||
|
"model": model,
|
||||||
|
})
|
||||||
|
|
||||||
|
return nil, p.parseAntigravityError(resp.StatusCode, respBody)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Response is always SSE from streamGenerateContent — each line is "data: {...}"
|
||||||
|
// with a "response" wrapper containing the standard Gemini response
|
||||||
|
llmResp, err := p.parseSSEResponse(string(respBody))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for empty response (some models might return valid success but empty text)
|
||||||
|
if llmResp.Content == "" && len(llmResp.ToolCalls) == 0 {
|
||||||
|
return nil, fmt.Errorf("antigravity: model returned an empty response (this model might be invalid or restricted)")
|
||||||
|
}
|
||||||
|
|
||||||
|
return llmResp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDefaultModel returns the default model identifier.
|
||||||
|
func (p *AntigravityProvider) GetDefaultModel() string {
|
||||||
|
return antigravityDefaultModel
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Request building ---
|
||||||
|
|
||||||
|
type antigravityRequest struct {
|
||||||
|
Contents []antigravityContent `json:"contents"`
|
||||||
|
Tools []antigravityTool `json:"tools,omitempty"`
|
||||||
|
SystemPrompt *antigravitySystemPrompt `json:"systemInstruction,omitempty"`
|
||||||
|
Config *antigravityGenConfig `json:"generationConfig,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityContent struct {
|
||||||
|
Role string `json:"role"`
|
||||||
|
Parts []antigravityPart `json:"parts"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityPart struct {
|
||||||
|
Text string `json:"text,omitempty"`
|
||||||
|
ThoughtSignature string `json:"thoughtSignature,omitempty"`
|
||||||
|
ThoughtSignatureSnake string `json:"thought_signature,omitempty"`
|
||||||
|
FunctionCall *antigravityFunctionCall `json:"functionCall,omitempty"`
|
||||||
|
FunctionResponse *antigravityFunctionResponse `json:"functionResponse,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityFunctionCall struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Args map[string]interface{} `json:"args"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityFunctionResponse struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Response map[string]interface{} `json:"response"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityTool struct {
|
||||||
|
FunctionDeclarations []antigravityFuncDecl `json:"functionDeclarations"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityFuncDecl struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description,omitempty"`
|
||||||
|
Parameters interface{} `json:"parameters,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravitySystemPrompt struct {
|
||||||
|
Parts []antigravityPart `json:"parts"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type antigravityGenConfig struct {
|
||||||
|
MaxOutputTokens int `json:"maxOutputTokens,omitempty"`
|
||||||
|
Temperature float64 `json:"temperature,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) buildRequest(messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) antigravityRequest {
|
||||||
|
req := antigravityRequest{}
|
||||||
|
toolCallNames := make(map[string]string)
|
||||||
|
|
||||||
|
// Build contents from messages
|
||||||
|
for _, msg := range messages {
|
||||||
|
switch msg.Role {
|
||||||
|
case "system":
|
||||||
|
req.SystemPrompt = &antigravitySystemPrompt{
|
||||||
|
Parts: []antigravityPart{{Text: msg.Content}},
|
||||||
|
}
|
||||||
|
case "user":
|
||||||
|
if msg.ToolCallID != "" {
|
||||||
|
toolName := resolveToolResponseName(msg.ToolCallID, toolCallNames)
|
||||||
|
// Tool result
|
||||||
|
req.Contents = append(req.Contents, antigravityContent{
|
||||||
|
Role: "user",
|
||||||
|
Parts: []antigravityPart{{
|
||||||
|
FunctionResponse: &antigravityFunctionResponse{
|
||||||
|
Name: toolName,
|
||||||
|
Response: map[string]interface{}{
|
||||||
|
"result": msg.Content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
req.Contents = append(req.Contents, antigravityContent{
|
||||||
|
Role: "user",
|
||||||
|
Parts: []antigravityPart{{Text: msg.Content}},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
case "assistant":
|
||||||
|
content := antigravityContent{
|
||||||
|
Role: "model",
|
||||||
|
}
|
||||||
|
if msg.Content != "" {
|
||||||
|
content.Parts = append(content.Parts, antigravityPart{Text: msg.Content})
|
||||||
|
}
|
||||||
|
for _, tc := range msg.ToolCalls {
|
||||||
|
toolName, toolArgs, thoughtSignature := normalizeStoredToolCall(tc)
|
||||||
|
if toolName == "" {
|
||||||
|
logger.WarnCF("provider.antigravity", "Skipping tool call with empty name in history", map[string]interface{}{
|
||||||
|
"tool_call_id": tc.ID,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if tc.ID != "" {
|
||||||
|
toolCallNames[tc.ID] = toolName
|
||||||
|
}
|
||||||
|
content.Parts = append(content.Parts, antigravityPart{
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
|
ThoughtSignatureSnake: thoughtSignature,
|
||||||
|
FunctionCall: &antigravityFunctionCall{
|
||||||
|
Name: toolName,
|
||||||
|
Args: toolArgs,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if len(content.Parts) > 0 {
|
||||||
|
req.Contents = append(req.Contents, content)
|
||||||
|
}
|
||||||
|
case "tool":
|
||||||
|
toolName := resolveToolResponseName(msg.ToolCallID, toolCallNames)
|
||||||
|
req.Contents = append(req.Contents, antigravityContent{
|
||||||
|
Role: "user",
|
||||||
|
Parts: []antigravityPart{{
|
||||||
|
FunctionResponse: &antigravityFunctionResponse{
|
||||||
|
Name: toolName,
|
||||||
|
Response: map[string]interface{}{
|
||||||
|
"result": msg.Content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build tools (sanitize schemas for Gemini compatibility)
|
||||||
|
if len(tools) > 0 {
|
||||||
|
var funcDecls []antigravityFuncDecl
|
||||||
|
for _, t := range tools {
|
||||||
|
if t.Type != "function" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
params := sanitizeSchemaForGemini(t.Function.Parameters)
|
||||||
|
funcDecls = append(funcDecls, antigravityFuncDecl{
|
||||||
|
Name: t.Function.Name,
|
||||||
|
Description: t.Function.Description,
|
||||||
|
Parameters: params,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if len(funcDecls) > 0 {
|
||||||
|
req.Tools = []antigravityTool{{FunctionDeclarations: funcDecls}}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generation config
|
||||||
|
config := &antigravityGenConfig{}
|
||||||
|
if val, ok := options["max_tokens"]; ok {
|
||||||
|
if maxTokens, ok := val.(int); ok && maxTokens > 0 {
|
||||||
|
config.MaxOutputTokens = maxTokens
|
||||||
|
} else if maxTokens, ok := val.(float64); ok && maxTokens > 0 {
|
||||||
|
config.MaxOutputTokens = int(maxTokens)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if temp, ok := options["temperature"].(float64); ok {
|
||||||
|
config.Temperature = temp
|
||||||
|
}
|
||||||
|
if config.MaxOutputTokens > 0 || config.Temperature > 0 {
|
||||||
|
req.Config = config
|
||||||
|
}
|
||||||
|
|
||||||
|
return req
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeStoredToolCall(tc ToolCall) (string, map[string]interface{}, string) {
|
||||||
|
name := tc.Name
|
||||||
|
args := tc.Arguments
|
||||||
|
thoughtSignature := ""
|
||||||
|
|
||||||
|
if name == "" && tc.Function != nil {
|
||||||
|
name = tc.Function.Name
|
||||||
|
thoughtSignature = tc.Function.ThoughtSignature
|
||||||
|
} else if tc.Function != nil {
|
||||||
|
thoughtSignature = tc.Function.ThoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
|
if args == nil {
|
||||||
|
args = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(args) == 0 && tc.Function != nil && tc.Function.Arguments != "" {
|
||||||
|
var parsed map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(tc.Function.Arguments), &parsed); err == nil && parsed != nil {
|
||||||
|
args = parsed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return name, args, thoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveToolResponseName(toolCallID string, toolCallNames map[string]string) string {
|
||||||
|
if toolCallID == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
if name, ok := toolCallNames[toolCallID]; ok && name != "" {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
|
||||||
|
return inferToolNameFromCallID(toolCallID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func inferToolNameFromCallID(toolCallID string) string {
|
||||||
|
if !strings.HasPrefix(toolCallID, "call_") {
|
||||||
|
return toolCallID
|
||||||
|
}
|
||||||
|
|
||||||
|
rest := strings.TrimPrefix(toolCallID, "call_")
|
||||||
|
if idx := strings.LastIndex(rest, "_"); idx > 0 {
|
||||||
|
candidate := rest[:idx]
|
||||||
|
if candidate != "" {
|
||||||
|
return candidate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return toolCallID
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Response parsing ---
|
||||||
|
|
||||||
|
type antigravityJSONResponse struct {
|
||||||
|
Candidates []struct {
|
||||||
|
Content struct {
|
||||||
|
Parts []struct {
|
||||||
|
Text string `json:"text,omitempty"`
|
||||||
|
ThoughtSignature string `json:"thoughtSignature,omitempty"`
|
||||||
|
ThoughtSignatureSnake string `json:"thought_signature,omitempty"`
|
||||||
|
FunctionCall *antigravityFunctionCall `json:"functionCall,omitempty"`
|
||||||
|
} `json:"parts"`
|
||||||
|
Role string `json:"role"`
|
||||||
|
} `json:"content"`
|
||||||
|
FinishReason string `json:"finishReason"`
|
||||||
|
} `json:"candidates"`
|
||||||
|
UsageMetadata struct {
|
||||||
|
PromptTokenCount int `json:"promptTokenCount"`
|
||||||
|
CandidatesTokenCount int `json:"candidatesTokenCount"`
|
||||||
|
TotalTokenCount int `json:"totalTokenCount"`
|
||||||
|
} `json:"usageMetadata"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) parseJSONResponse(body []byte) (*LLMResponse, error) {
|
||||||
|
var resp antigravityJSONResponse
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("parsing antigravity response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(resp.Candidates) == 0 {
|
||||||
|
return nil, fmt.Errorf("antigravity: no candidates in response")
|
||||||
|
}
|
||||||
|
|
||||||
|
candidate := resp.Candidates[0]
|
||||||
|
var contentParts []string
|
||||||
|
var toolCalls []ToolCall
|
||||||
|
|
||||||
|
for _, part := range candidate.Content.Parts {
|
||||||
|
if part.Text != "" {
|
||||||
|
contentParts = append(contentParts, part.Text)
|
||||||
|
}
|
||||||
|
if part.FunctionCall != nil {
|
||||||
|
argumentsJSON, _ := json.Marshal(part.FunctionCall.Args)
|
||||||
|
toolCalls = append(toolCalls, ToolCall{
|
||||||
|
ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()),
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: part.FunctionCall.Args,
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: string(argumentsJSON),
|
||||||
|
ThoughtSignature: extractPartThoughtSignature(part.ThoughtSignature, part.ThoughtSignatureSnake),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
finishReason := "stop"
|
||||||
|
if len(toolCalls) > 0 {
|
||||||
|
finishReason = "tool_calls"
|
||||||
|
}
|
||||||
|
if candidate.FinishReason == "MAX_TOKENS" {
|
||||||
|
finishReason = "length"
|
||||||
|
}
|
||||||
|
|
||||||
|
var usage *UsageInfo
|
||||||
|
if resp.UsageMetadata.TotalTokenCount > 0 {
|
||||||
|
usage = &UsageInfo{
|
||||||
|
PromptTokens: resp.UsageMetadata.PromptTokenCount,
|
||||||
|
CompletionTokens: resp.UsageMetadata.CandidatesTokenCount,
|
||||||
|
TotalTokens: resp.UsageMetadata.TotalTokenCount,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: strings.Join(contentParts, ""),
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: finishReason,
|
||||||
|
Usage: usage,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) parseSSEResponse(body string) (*LLMResponse, error) {
|
||||||
|
var contentParts []string
|
||||||
|
var toolCalls []ToolCall
|
||||||
|
var usage *UsageInfo
|
||||||
|
var finishReason string
|
||||||
|
|
||||||
|
scanner := bufio.NewScanner(strings.NewReader(body))
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Text()
|
||||||
|
if !strings.HasPrefix(line, "data: ") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
data := strings.TrimPrefix(line, "data: ")
|
||||||
|
if data == "[DONE]" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
// v1internal SSE wraps the Gemini response in a "response" field
|
||||||
|
var sseChunk struct {
|
||||||
|
Response antigravityJSONResponse `json:"response"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(data), &sseChunk); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
resp := sseChunk.Response
|
||||||
|
|
||||||
|
for _, candidate := range resp.Candidates {
|
||||||
|
for _, part := range candidate.Content.Parts {
|
||||||
|
if part.Text != "" {
|
||||||
|
contentParts = append(contentParts, part.Text)
|
||||||
|
}
|
||||||
|
if part.FunctionCall != nil {
|
||||||
|
argumentsJSON, _ := json.Marshal(part.FunctionCall.Args)
|
||||||
|
toolCalls = append(toolCalls, ToolCall{
|
||||||
|
ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()),
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: part.FunctionCall.Args,
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: string(argumentsJSON),
|
||||||
|
ThoughtSignature: extractPartThoughtSignature(part.ThoughtSignature, part.ThoughtSignatureSnake),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if candidate.FinishReason != "" {
|
||||||
|
finishReason = candidate.FinishReason
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.UsageMetadata.TotalTokenCount > 0 {
|
||||||
|
usage = &UsageInfo{
|
||||||
|
PromptTokens: resp.UsageMetadata.PromptTokenCount,
|
||||||
|
CompletionTokens: resp.UsageMetadata.CandidatesTokenCount,
|
||||||
|
TotalTokens: resp.UsageMetadata.TotalTokenCount,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mappedFinish := "stop"
|
||||||
|
if len(toolCalls) > 0 {
|
||||||
|
mappedFinish = "tool_calls"
|
||||||
|
}
|
||||||
|
if finishReason == "MAX_TOKENS" {
|
||||||
|
mappedFinish = "length"
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: strings.Join(contentParts, ""),
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: mappedFinish,
|
||||||
|
Usage: usage,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractPartThoughtSignature(thoughtSignature string, thoughtSignatureSnake string) string {
|
||||||
|
if thoughtSignature != "" {
|
||||||
|
return thoughtSignature
|
||||||
|
}
|
||||||
|
if thoughtSignatureSnake != "" {
|
||||||
|
return thoughtSignatureSnake
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Schema sanitization ---
|
||||||
|
|
||||||
|
// Google/Gemini doesn't support many JSON Schema keywords that other providers accept.
|
||||||
|
var geminiUnsupportedKeywords = map[string]bool{
|
||||||
|
"patternProperties": true,
|
||||||
|
"additionalProperties": true,
|
||||||
|
"$schema": true,
|
||||||
|
"$id": true,
|
||||||
|
"$ref": true,
|
||||||
|
"$defs": true,
|
||||||
|
"definitions": true,
|
||||||
|
"examples": true,
|
||||||
|
"minLength": true,
|
||||||
|
"maxLength": true,
|
||||||
|
"minimum": true,
|
||||||
|
"maximum": true,
|
||||||
|
"multipleOf": true,
|
||||||
|
"pattern": true,
|
||||||
|
"format": true,
|
||||||
|
"minItems": true,
|
||||||
|
"maxItems": true,
|
||||||
|
"uniqueItems": true,
|
||||||
|
"minProperties": true,
|
||||||
|
"maxProperties": true,
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeSchemaForGemini(schema map[string]interface{}) map[string]interface{} {
|
||||||
|
if schema == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
result := make(map[string]interface{})
|
||||||
|
for k, v := range schema {
|
||||||
|
if geminiUnsupportedKeywords[k] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Recursively sanitize nested objects
|
||||||
|
switch val := v.(type) {
|
||||||
|
case map[string]interface{}:
|
||||||
|
result[k] = sanitizeSchemaForGemini(val)
|
||||||
|
case []interface{}:
|
||||||
|
sanitized := make([]interface{}, len(val))
|
||||||
|
for i, item := range val {
|
||||||
|
if m, ok := item.(map[string]interface{}); ok {
|
||||||
|
sanitized[i] = sanitizeSchemaForGemini(m)
|
||||||
|
} else {
|
||||||
|
sanitized[i] = item
|
||||||
|
}
|
||||||
|
}
|
||||||
|
result[k] = sanitized
|
||||||
|
default:
|
||||||
|
result[k] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure top-level has type: "object" if properties are present
|
||||||
|
if _, hasProps := result["properties"]; hasProps {
|
||||||
|
if _, hasType := result["type"]; !hasType {
|
||||||
|
result["type"] = "object"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Token source ---
|
||||||
|
|
||||||
|
func createAntigravityTokenSource() func() (string, string, error) {
|
||||||
|
return func() (string, string, error) {
|
||||||
|
cred, err := auth.GetCredential("google-antigravity")
|
||||||
|
if err != nil {
|
||||||
|
return "", "", fmt.Errorf("loading auth credentials: %w", err)
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
return "", "", fmt.Errorf("no credentials for google-antigravity. Run: picoclaw auth login --provider google-antigravity")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh if needed
|
||||||
|
if cred.NeedsRefresh() && cred.RefreshToken != "" {
|
||||||
|
oauthCfg := auth.GoogleAntigravityOAuthConfig()
|
||||||
|
refreshed, err := auth.RefreshAccessToken(cred, oauthCfg)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", fmt.Errorf("refreshing token: %w", err)
|
||||||
|
}
|
||||||
|
refreshed.Email = cred.Email
|
||||||
|
if refreshed.ProjectID == "" {
|
||||||
|
refreshed.ProjectID = cred.ProjectID
|
||||||
|
}
|
||||||
|
if err := auth.SetCredential("google-antigravity", refreshed); err != nil {
|
||||||
|
return "", "", fmt.Errorf("saving refreshed token: %w", err)
|
||||||
|
}
|
||||||
|
cred = refreshed
|
||||||
|
}
|
||||||
|
|
||||||
|
if cred.IsExpired() {
|
||||||
|
return "", "", fmt.Errorf("antigravity credentials expired. Run: picoclaw auth login --provider google-antigravity")
|
||||||
|
}
|
||||||
|
|
||||||
|
projectID := cred.ProjectID
|
||||||
|
if projectID == "" {
|
||||||
|
// Try to fetch project ID from API
|
||||||
|
fetchedID, err := FetchAntigravityProjectID(cred.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("provider.antigravity", "Could not fetch project ID, using fallback", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
projectID = "rising-fact-p41fc" // Default fallback (same as OpenCode)
|
||||||
|
} else {
|
||||||
|
projectID = fetchedID
|
||||||
|
cred.ProjectID = projectID
|
||||||
|
_ = auth.SetCredential("google-antigravity", cred)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return cred.AccessToken, projectID, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// FetchAntigravityProjectID retrieves the Google Cloud project ID from the loadCodeAssist endpoint.
|
||||||
|
func FetchAntigravityProjectID(accessToken string) (string, error) {
|
||||||
|
reqBody, _ := json.Marshal(map[string]interface{}{
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"ideType": "IDE_UNSPECIFIED",
|
||||||
|
"platform": "PLATFORM_UNSPECIFIED",
|
||||||
|
"pluginType": "GEMINI",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
req, err := http.NewRequest("POST", antigravityBaseURL+"/v1internal:loadCodeAssist", bytes.NewReader(reqBody))
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("User-Agent", antigravityUserAgent)
|
||||||
|
req.Header.Set("X-Goog-Api-Client", antigravityXGoogClient)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 15 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("loadCodeAssist failed: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
var result struct {
|
||||||
|
CloudAICompanionProject string `json:"cloudaicompanionProject"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &result); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.CloudAICompanionProject == "" {
|
||||||
|
return "", fmt.Errorf("no project ID in loadCodeAssist response")
|
||||||
|
}
|
||||||
|
|
||||||
|
return result.CloudAICompanionProject, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FetchAntigravityModels fetches available models from the Cloud Code Assist API.
|
||||||
|
func FetchAntigravityModels(accessToken, projectID string) ([]AntigravityModelInfo, error) {
|
||||||
|
reqBody, _ := json.Marshal(map[string]interface{}{
|
||||||
|
"project": projectID,
|
||||||
|
})
|
||||||
|
|
||||||
|
req, err := http.NewRequest("POST", antigravityBaseURL+"/v1internal:fetchAvailableModels", bytes.NewReader(reqBody))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("User-Agent", antigravityUserAgent)
|
||||||
|
req.Header.Set("X-Goog-Api-Client", antigravityXGoogClient)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 15 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return nil, fmt.Errorf("fetchAvailableModels failed (HTTP %d): %s", resp.StatusCode, truncateString(string(body), 200))
|
||||||
|
}
|
||||||
|
|
||||||
|
var result struct {
|
||||||
|
Models map[string]struct {
|
||||||
|
DisplayName string `json:"displayName"`
|
||||||
|
QuotaInfo struct {
|
||||||
|
RemainingFraction interface{} `json:"remainingFraction"`
|
||||||
|
ResetTime string `json:"resetTime"`
|
||||||
|
IsExhausted bool `json:"isExhausted"`
|
||||||
|
} `json:"quotaInfo"`
|
||||||
|
} `json:"models"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("parsing models response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var models []AntigravityModelInfo
|
||||||
|
for id, info := range result.Models {
|
||||||
|
models = append(models, AntigravityModelInfo{
|
||||||
|
ID: id,
|
||||||
|
DisplayName: info.DisplayName,
|
||||||
|
IsExhausted: info.QuotaInfo.IsExhausted,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure gemini-3-flash-preview and gemini-3-flash are in the list if they aren't already
|
||||||
|
hasFlashPreview := false
|
||||||
|
hasFlash := false
|
||||||
|
for _, m := range models {
|
||||||
|
if m.ID == "gemini-3-flash-preview" {
|
||||||
|
hasFlashPreview = true
|
||||||
|
}
|
||||||
|
if m.ID == "gemini-3-flash" {
|
||||||
|
hasFlash = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !hasFlashPreview {
|
||||||
|
models = append(models, AntigravityModelInfo{
|
||||||
|
ID: "gemini-3-flash-preview",
|
||||||
|
DisplayName: "Gemini 3 Flash (Preview)",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if !hasFlash {
|
||||||
|
models = append(models, AntigravityModelInfo{
|
||||||
|
ID: "gemini-3-flash",
|
||||||
|
DisplayName: "Gemini 3 Flash",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return models, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type AntigravityModelInfo struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
IsExhausted bool `json:"is_exhausted"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Helpers ---
|
||||||
|
|
||||||
|
func truncateString(s string, maxLen int) string {
|
||||||
|
if len(s) <= maxLen {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return s[:maxLen] + "..."
|
||||||
|
}
|
||||||
|
|
||||||
|
func randomString(n int) string {
|
||||||
|
const letters = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||||
|
b := make([]byte, n)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = letters[rand.Intn(len(letters))]
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *AntigravityProvider) parseAntigravityError(statusCode int, body []byte) error {
|
||||||
|
var errResp struct {
|
||||||
|
Error struct {
|
||||||
|
Code int `json:"code"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
Details []map[string]interface{} `json:"details"`
|
||||||
|
} `json:"error"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &errResp); err != nil {
|
||||||
|
return fmt.Errorf("antigravity API error (HTTP %d): %s", statusCode, truncateString(string(body), 500))
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := errResp.Error.Message
|
||||||
|
if statusCode == 429 {
|
||||||
|
// Try to extract quota reset info
|
||||||
|
for _, detail := range errResp.Error.Details {
|
||||||
|
if typeVal, ok := detail["@type"].(string); ok && strings.HasSuffix(typeVal, "ErrorInfo") {
|
||||||
|
if metadata, ok := detail["metadata"].(map[string]interface{}); ok {
|
||||||
|
if delay, ok := metadata["quotaResetDelay"].(string); ok {
|
||||||
|
return fmt.Errorf("antigravity rate limit exceeded: %s (reset in %s)", msg, delay)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("antigravity rate limit exceeded: %s", msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("antigravity API error (%s): %s", errResp.Error.Status, msg)
|
||||||
|
}
|
||||||
56
pkg/providers/antigravity_provider_test.go
Normal file
56
pkg/providers/antigravity_provider_test.go
Normal file
|
|
@ -0,0 +1,56 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestBuildRequestUsesFunctionFieldsWhenToolCallNameMissing(t *testing.T) {
|
||||||
|
p := &AntigravityProvider{}
|
||||||
|
|
||||||
|
messages := []Message{
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
ToolCalls: []ToolCall{{
|
||||||
|
ID: "call_read_file_123",
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: "read_file",
|
||||||
|
Arguments: `{"path":"README.md"}`,
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Role: "tool",
|
||||||
|
ToolCallID: "call_read_file_123",
|
||||||
|
Content: "ok",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
req := p.buildRequest(messages, nil, "", nil)
|
||||||
|
if len(req.Contents) != 2 {
|
||||||
|
t.Fatalf("expected 2 contents, got %d", len(req.Contents))
|
||||||
|
}
|
||||||
|
|
||||||
|
modelPart := req.Contents[0].Parts[0]
|
||||||
|
if modelPart.FunctionCall == nil {
|
||||||
|
t.Fatal("expected functionCall in assistant message")
|
||||||
|
}
|
||||||
|
if modelPart.FunctionCall.Name != "read_file" {
|
||||||
|
t.Fatalf("expected functionCall name read_file, got %q", modelPart.FunctionCall.Name)
|
||||||
|
}
|
||||||
|
if got := modelPart.FunctionCall.Args["path"]; got != "README.md" {
|
||||||
|
t.Fatalf("expected functionCall args[path] to be README.md, got %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
toolPart := req.Contents[1].Parts[0]
|
||||||
|
if toolPart.FunctionResponse == nil {
|
||||||
|
t.Fatal("expected functionResponse in tool message")
|
||||||
|
}
|
||||||
|
if toolPart.FunctionResponse.Name != "read_file" {
|
||||||
|
t.Fatalf("expected functionResponse name read_file, got %q", toolPart.FunctionResponse.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveToolResponseNameInfersNameFromGeneratedCallID(t *testing.T) {
|
||||||
|
got := resolveToolResponseName("call_search_docs_999", map[string]string{})
|
||||||
|
if got != "search_docs" {
|
||||||
|
t.Fatalf("expected inferred tool name search_docs, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -336,7 +336,7 @@ func TestChat_PassesModelFlag(t *testing.T) {
|
||||||
|
|
||||||
_, err := p.Chat(context.Background(), []Message{
|
_, err := p.Chat(context.Background(), []Message{
|
||||||
{Role: "user", Content: "Hi"},
|
{Role: "user", Content: "Hi"},
|
||||||
}, nil, "claude-sonnet-4-5-20250929", nil)
|
}, nil, "claude-sonnet-4.6", nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Chat() error = %v", err)
|
t.Fatalf("Chat() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -346,7 +346,7 @@ func TestChat_PassesModelFlag(t *testing.T) {
|
||||||
if !strings.Contains(args, "--model") {
|
if !strings.Contains(args, "--model") {
|
||||||
t.Errorf("CLI args missing --model, got: %s", args)
|
t.Errorf("CLI args missing --model, got: %s", args)
|
||||||
}
|
}
|
||||||
if !strings.Contains(args, "claude-sonnet-4-5-20250929") {
|
if !strings.Contains(args, "claude-sonnet-4.6") {
|
||||||
t.Errorf("CLI args missing model name, got: %s", args)
|
t.Errorf("CLI args missing model name, got: %s", args)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -416,10 +416,12 @@ func TestChat_EmptyWorkspaceDoesNotSetDir(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCli(t *testing.T) {
|
func TestCreateProvider_ClaudeCli(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claude-cli"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
cfg.Agents.Defaults.Workspace = "/test/ws"
|
{ModelName: "claude-sonnet-4.6", Model: "claude-cli/claude-sonnet-4.6", Workspace: "/test/ws"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claude-sonnet-4.6"
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider(claude-cli) error = %v", err)
|
t.Fatalf("CreateProvider(claude-cli) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -435,9 +437,12 @@ func TestCreateProvider_ClaudeCli(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCode(t *testing.T) {
|
func TestCreateProvider_ClaudeCode(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claude-code"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{ModelName: "claude-code", Model: "claude-cli/claude-code"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claude-code"
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider(claude-code) error = %v", err)
|
t.Fatalf("CreateProvider(claude-code) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -448,9 +453,12 @@ func TestCreateProvider_ClaudeCode(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCodec(t *testing.T) {
|
func TestCreateProvider_ClaudeCodec(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claudecode"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{ModelName: "claudecode", Model: "claude-cli/claudecode"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claudecode"
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider(claudecode) error = %v", err)
|
t.Fatalf("CreateProvider(claudecode) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -461,10 +469,13 @@ func TestCreateProvider_ClaudeCodec(t *testing.T) {
|
||||||
|
|
||||||
func TestCreateProvider_ClaudeCliDefaultWorkspace(t *testing.T) {
|
func TestCreateProvider_ClaudeCliDefaultWorkspace(t *testing.T) {
|
||||||
cfg := config.DefaultConfig()
|
cfg := config.DefaultConfig()
|
||||||
cfg.Agents.Defaults.Provider = "claude-cli"
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{ModelName: "claude-cli", Model: "claude-cli/claude-sonnet"},
|
||||||
|
}
|
||||||
|
cfg.Agents.Defaults.Model = "claude-cli"
|
||||||
cfg.Agents.Defaults.Workspace = ""
|
cfg.Agents.Defaults.Workspace = ""
|
||||||
|
|
||||||
provider, err := CreateProvider(cfg)
|
provider, _, err := CreateProvider(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CreateProvider error = %v", err)
|
t.Fatalf("CreateProvider error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,200 +2,58 @@ package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
"github.com/anthropics/anthropic-sdk-go"
|
anthropicprovider "github.com/sipeed/picoclaw/pkg/providers/anthropic"
|
||||||
"github.com/anthropics/anthropic-sdk-go/option"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type ClaudeProvider struct {
|
type ClaudeProvider struct {
|
||||||
client *anthropic.Client
|
delegate *anthropicprovider.Provider
|
||||||
tokenSource func() (string, error)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewClaudeProvider(token string) *ClaudeProvider {
|
func NewClaudeProvider(token string) *ClaudeProvider {
|
||||||
client := anthropic.NewClient(
|
return &ClaudeProvider{
|
||||||
option.WithAuthToken(token),
|
delegate: anthropicprovider.NewProvider(token),
|
||||||
option.WithBaseURL("https://api.anthropic.com"),
|
}
|
||||||
)
|
}
|
||||||
return &ClaudeProvider{client: &client}
|
|
||||||
|
func NewClaudeProviderWithBaseURL(token, apiBase string) *ClaudeProvider {
|
||||||
|
return &ClaudeProvider{
|
||||||
|
delegate: anthropicprovider.NewProviderWithBaseURL(token, apiBase),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewClaudeProviderWithTokenSource(token string, tokenSource func() (string, error)) *ClaudeProvider {
|
func NewClaudeProviderWithTokenSource(token string, tokenSource func() (string, error)) *ClaudeProvider {
|
||||||
p := NewClaudeProvider(token)
|
return &ClaudeProvider{
|
||||||
p.tokenSource = tokenSource
|
delegate: anthropicprovider.NewProviderWithTokenSource(token, tokenSource),
|
||||||
return p
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewClaudeProviderWithTokenSourceAndBaseURL(token string, tokenSource func() (string, error), apiBase string) *ClaudeProvider {
|
||||||
|
return &ClaudeProvider{
|
||||||
|
delegate: anthropicprovider.NewProviderWithTokenSourceAndBaseURL(token, tokenSource, apiBase),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newClaudeProviderWithDelegate(delegate *anthropicprovider.Provider) *ClaudeProvider {
|
||||||
|
return &ClaudeProvider{delegate: delegate}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *ClaudeProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
func (p *ClaudeProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
var opts []option.RequestOption
|
resp, err := p.delegate.Chat(ctx, messages, tools, model, options)
|
||||||
if p.tokenSource != nil {
|
|
||||||
tok, err := p.tokenSource()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("refreshing token: %w", err)
|
|
||||||
}
|
|
||||||
opts = append(opts, option.WithAuthToken(tok))
|
|
||||||
}
|
|
||||||
|
|
||||||
params, err := buildClaudeParams(messages, tools, model, options)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
return resp, nil
|
||||||
resp, err := p.client.Messages.New(ctx, params, opts...)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("claude API call: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return parseClaudeResponse(resp), nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *ClaudeProvider) GetDefaultModel() string {
|
func (p *ClaudeProvider) GetDefaultModel() string {
|
||||||
return "claude-sonnet-4-5-20250929"
|
return p.delegate.GetDefaultModel()
|
||||||
}
|
|
||||||
|
|
||||||
func buildClaudeParams(messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (anthropic.MessageNewParams, error) {
|
|
||||||
var system []anthropic.TextBlockParam
|
|
||||||
var anthropicMessages []anthropic.MessageParam
|
|
||||||
|
|
||||||
for _, msg := range messages {
|
|
||||||
switch msg.Role {
|
|
||||||
case "system":
|
|
||||||
system = append(system, anthropic.TextBlockParam{Text: msg.Content})
|
|
||||||
case "user":
|
|
||||||
if msg.ToolCallID != "" {
|
|
||||||
anthropicMessages = append(anthropicMessages,
|
|
||||||
anthropic.NewUserMessage(anthropic.NewToolResultBlock(msg.ToolCallID, msg.Content, false)),
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
anthropicMessages = append(anthropicMessages,
|
|
||||||
anthropic.NewUserMessage(anthropic.NewTextBlock(msg.Content)),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
case "assistant":
|
|
||||||
if len(msg.ToolCalls) > 0 {
|
|
||||||
var blocks []anthropic.ContentBlockParamUnion
|
|
||||||
if msg.Content != "" {
|
|
||||||
blocks = append(blocks, anthropic.NewTextBlock(msg.Content))
|
|
||||||
}
|
|
||||||
for _, tc := range msg.ToolCalls {
|
|
||||||
blocks = append(blocks, anthropic.NewToolUseBlock(tc.ID, tc.Arguments, tc.Name))
|
|
||||||
}
|
|
||||||
anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...))
|
|
||||||
} else {
|
|
||||||
anthropicMessages = append(anthropicMessages,
|
|
||||||
anthropic.NewAssistantMessage(anthropic.NewTextBlock(msg.Content)),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
case "tool":
|
|
||||||
anthropicMessages = append(anthropicMessages,
|
|
||||||
anthropic.NewUserMessage(anthropic.NewToolResultBlock(msg.ToolCallID, msg.Content, false)),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
maxTokens := int64(4096)
|
|
||||||
if mt, ok := options["max_tokens"].(int); ok {
|
|
||||||
maxTokens = int64(mt)
|
|
||||||
}
|
|
||||||
|
|
||||||
params := anthropic.MessageNewParams{
|
|
||||||
Model: anthropic.Model(model),
|
|
||||||
Messages: anthropicMessages,
|
|
||||||
MaxTokens: maxTokens,
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(system) > 0 {
|
|
||||||
params.System = system
|
|
||||||
}
|
|
||||||
|
|
||||||
if temp, ok := options["temperature"].(float64); ok {
|
|
||||||
params.Temperature = anthropic.Float(temp)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(tools) > 0 {
|
|
||||||
params.Tools = translateToolsForClaude(tools)
|
|
||||||
}
|
|
||||||
|
|
||||||
return params, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func translateToolsForClaude(tools []ToolDefinition) []anthropic.ToolUnionParam {
|
|
||||||
result := make([]anthropic.ToolUnionParam, 0, len(tools))
|
|
||||||
for _, t := range tools {
|
|
||||||
tool := anthropic.ToolParam{
|
|
||||||
Name: t.Function.Name,
|
|
||||||
InputSchema: anthropic.ToolInputSchemaParam{
|
|
||||||
Properties: t.Function.Parameters["properties"],
|
|
||||||
},
|
|
||||||
}
|
|
||||||
if desc := t.Function.Description; desc != "" {
|
|
||||||
tool.Description = anthropic.String(desc)
|
|
||||||
}
|
|
||||||
if req, ok := t.Function.Parameters["required"].([]interface{}); ok {
|
|
||||||
required := make([]string, 0, len(req))
|
|
||||||
for _, r := range req {
|
|
||||||
if s, ok := r.(string); ok {
|
|
||||||
required = append(required, s)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
tool.InputSchema.Required = required
|
|
||||||
}
|
|
||||||
result = append(result, anthropic.ToolUnionParam{OfTool: &tool})
|
|
||||||
}
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseClaudeResponse(resp *anthropic.Message) *LLMResponse {
|
|
||||||
var content string
|
|
||||||
var toolCalls []ToolCall
|
|
||||||
|
|
||||||
for _, block := range resp.Content {
|
|
||||||
switch block.Type {
|
|
||||||
case "text":
|
|
||||||
tb := block.AsText()
|
|
||||||
content += tb.Text
|
|
||||||
case "tool_use":
|
|
||||||
tu := block.AsToolUse()
|
|
||||||
var args map[string]interface{}
|
|
||||||
if err := json.Unmarshal(tu.Input, &args); err != nil {
|
|
||||||
args = map[string]interface{}{"raw": string(tu.Input)}
|
|
||||||
}
|
|
||||||
toolCalls = append(toolCalls, ToolCall{
|
|
||||||
ID: tu.ID,
|
|
||||||
Name: tu.Name,
|
|
||||||
Arguments: args,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
finishReason := "stop"
|
|
||||||
switch resp.StopReason {
|
|
||||||
case anthropic.StopReasonToolUse:
|
|
||||||
finishReason = "tool_calls"
|
|
||||||
case anthropic.StopReasonMaxTokens:
|
|
||||||
finishReason = "length"
|
|
||||||
case anthropic.StopReasonEndTurn:
|
|
||||||
finishReason = "stop"
|
|
||||||
}
|
|
||||||
|
|
||||||
return &LLMResponse{
|
|
||||||
Content: content,
|
|
||||||
ToolCalls: toolCalls,
|
|
||||||
FinishReason: finishReason,
|
|
||||||
Usage: &UsageInfo{
|
|
||||||
PromptTokens: int(resp.Usage.InputTokens),
|
|
||||||
CompletionTokens: int(resp.Usage.OutputTokens),
|
|
||||||
TotalTokens: int(resp.Usage.InputTokens + resp.Usage.OutputTokens),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func createClaudeTokenSource() func() (string, error) {
|
func createClaudeTokenSource() func() (string, error) {
|
||||||
return func() (string, error) {
|
return func() (string, error) {
|
||||||
cred, err := auth.GetCredential("anthropic")
|
cred, err := getCredential("anthropic")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("loading auth credentials: %w", err)
|
return "", fmt.Errorf("loading auth credentials: %w", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -8,140 +8,9 @@ import (
|
||||||
|
|
||||||
"github.com/anthropics/anthropic-sdk-go"
|
"github.com/anthropics/anthropic-sdk-go"
|
||||||
anthropicoption "github.com/anthropics/anthropic-sdk-go/option"
|
anthropicoption "github.com/anthropics/anthropic-sdk-go/option"
|
||||||
|
anthropicprovider "github.com/sipeed/picoclaw/pkg/providers/anthropic"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestBuildClaudeParams_BasicMessage(t *testing.T) {
|
|
||||||
messages := []Message{
|
|
||||||
{Role: "user", Content: "Hello"},
|
|
||||||
}
|
|
||||||
params, err := buildClaudeParams(messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{
|
|
||||||
"max_tokens": 1024,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("buildClaudeParams() error: %v", err)
|
|
||||||
}
|
|
||||||
if string(params.Model) != "claude-sonnet-4-5-20250929" {
|
|
||||||
t.Errorf("Model = %q, want %q", params.Model, "claude-sonnet-4-5-20250929")
|
|
||||||
}
|
|
||||||
if params.MaxTokens != 1024 {
|
|
||||||
t.Errorf("MaxTokens = %d, want 1024", params.MaxTokens)
|
|
||||||
}
|
|
||||||
if len(params.Messages) != 1 {
|
|
||||||
t.Fatalf("len(Messages) = %d, want 1", len(params.Messages))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildClaudeParams_SystemMessage(t *testing.T) {
|
|
||||||
messages := []Message{
|
|
||||||
{Role: "system", Content: "You are helpful"},
|
|
||||||
{Role: "user", Content: "Hi"},
|
|
||||||
}
|
|
||||||
params, err := buildClaudeParams(messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("buildClaudeParams() error: %v", err)
|
|
||||||
}
|
|
||||||
if len(params.System) != 1 {
|
|
||||||
t.Fatalf("len(System) = %d, want 1", len(params.System))
|
|
||||||
}
|
|
||||||
if params.System[0].Text != "You are helpful" {
|
|
||||||
t.Errorf("System[0].Text = %q, want %q", params.System[0].Text, "You are helpful")
|
|
||||||
}
|
|
||||||
if len(params.Messages) != 1 {
|
|
||||||
t.Fatalf("len(Messages) = %d, want 1", len(params.Messages))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildClaudeParams_ToolCallMessage(t *testing.T) {
|
|
||||||
messages := []Message{
|
|
||||||
{Role: "user", Content: "What's the weather?"},
|
|
||||||
{
|
|
||||||
Role: "assistant",
|
|
||||||
Content: "",
|
|
||||||
ToolCalls: []ToolCall{
|
|
||||||
{
|
|
||||||
ID: "call_1",
|
|
||||||
Name: "get_weather",
|
|
||||||
Arguments: map[string]interface{}{"city": "SF"},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
|
|
||||||
}
|
|
||||||
params, err := buildClaudeParams(messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("buildClaudeParams() error: %v", err)
|
|
||||||
}
|
|
||||||
if len(params.Messages) != 3 {
|
|
||||||
t.Fatalf("len(Messages) = %d, want 3", len(params.Messages))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildClaudeParams_WithTools(t *testing.T) {
|
|
||||||
tools := []ToolDefinition{
|
|
||||||
{
|
|
||||||
Type: "function",
|
|
||||||
Function: ToolFunctionDefinition{
|
|
||||||
Name: "get_weather",
|
|
||||||
Description: "Get weather for a city",
|
|
||||||
Parameters: map[string]interface{}{
|
|
||||||
"type": "object",
|
|
||||||
"properties": map[string]interface{}{
|
|
||||||
"city": map[string]interface{}{"type": "string"},
|
|
||||||
},
|
|
||||||
"required": []interface{}{"city"},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
params, err := buildClaudeParams([]Message{{Role: "user", Content: "Hi"}}, tools, "claude-sonnet-4-5-20250929", map[string]interface{}{})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("buildClaudeParams() error: %v", err)
|
|
||||||
}
|
|
||||||
if len(params.Tools) != 1 {
|
|
||||||
t.Fatalf("len(Tools) = %d, want 1", len(params.Tools))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseClaudeResponse_TextOnly(t *testing.T) {
|
|
||||||
resp := &anthropic.Message{
|
|
||||||
Content: []anthropic.ContentBlockUnion{},
|
|
||||||
Usage: anthropic.Usage{
|
|
||||||
InputTokens: 10,
|
|
||||||
OutputTokens: 20,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
result := parseClaudeResponse(resp)
|
|
||||||
if result.Usage.PromptTokens != 10 {
|
|
||||||
t.Errorf("PromptTokens = %d, want 10", result.Usage.PromptTokens)
|
|
||||||
}
|
|
||||||
if result.Usage.CompletionTokens != 20 {
|
|
||||||
t.Errorf("CompletionTokens = %d, want 20", result.Usage.CompletionTokens)
|
|
||||||
}
|
|
||||||
if result.FinishReason != "stop" {
|
|
||||||
t.Errorf("FinishReason = %q, want %q", result.FinishReason, "stop")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseClaudeResponse_StopReasons(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
stopReason anthropic.StopReason
|
|
||||||
want string
|
|
||||||
}{
|
|
||||||
{anthropic.StopReasonEndTurn, "stop"},
|
|
||||||
{anthropic.StopReasonMaxTokens, "length"},
|
|
||||||
{anthropic.StopReasonToolUse, "tool_calls"},
|
|
||||||
}
|
|
||||||
for _, tt := range tests {
|
|
||||||
resp := &anthropic.Message{
|
|
||||||
StopReason: tt.stopReason,
|
|
||||||
}
|
|
||||||
result := parseClaudeResponse(resp)
|
|
||||||
if result.FinishReason != tt.want {
|
|
||||||
t.Errorf("StopReason %q: FinishReason = %q, want %q", tt.stopReason, result.FinishReason, tt.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClaudeProvider_ChatRoundTrip(t *testing.T) {
|
func TestClaudeProvider_ChatRoundTrip(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.URL.Path != "/v1/messages" {
|
if r.URL.Path != "/v1/messages" {
|
||||||
|
|
@ -175,11 +44,11 @@ func TestClaudeProvider_ChatRoundTrip(t *testing.T) {
|
||||||
}))
|
}))
|
||||||
defer server.Close()
|
defer server.Close()
|
||||||
|
|
||||||
provider := NewClaudeProvider("test-token")
|
delegate := anthropicprovider.NewProviderWithClient(createAnthropicTestClient(server.URL, "test-token"))
|
||||||
provider.client = createAnthropicTestClient(server.URL, "test-token")
|
provider := newClaudeProviderWithDelegate(delegate)
|
||||||
|
|
||||||
messages := []Message{{Role: "user", Content: "Hello"}}
|
messages := []Message{{Role: "user", Content: "Hello"}}
|
||||||
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4-5-20250929", map[string]interface{}{"max_tokens": 1024})
|
resp, err := provider.Chat(t.Context(), messages, nil, "claude-sonnet-4.6", map[string]interface{}{"max_tokens": 1024})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Chat() error: %v", err)
|
t.Fatalf("Chat() error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -196,8 +65,8 @@ func TestClaudeProvider_ChatRoundTrip(t *testing.T) {
|
||||||
|
|
||||||
func TestClaudeProvider_GetDefaultModel(t *testing.T) {
|
func TestClaudeProvider_GetDefaultModel(t *testing.T) {
|
||||||
p := NewClaudeProvider("test-token")
|
p := NewClaudeProvider("test-token")
|
||||||
if got := p.GetDefaultModel(); got != "claude-sonnet-4-5-20250929" {
|
if got := p.GetDefaultModel(); got != "claude-sonnet-4.6" {
|
||||||
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4-5-20250929")
|
t.Errorf("GetDefaultModel() = %q, want %q", got, "claude-sonnet-4.6")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
119
pkg/providers/codex_cli_provider_integration_test.go
Normal file
119
pkg/providers/codex_cli_provider_integration_test.go
Normal file
|
|
@ -0,0 +1,119 @@
|
||||||
|
//go:build integration
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
exec "os/exec"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestIntegration_RealCodexCLI tests the CodexCliProvider with a real codex CLI.
|
||||||
|
// Run with: go test -tags=integration ./pkg/providers/...
|
||||||
|
func TestIntegration_RealCodexCLI(t *testing.T) {
|
||||||
|
path, err := exec.LookPath("codex")
|
||||||
|
if err != nil {
|
||||||
|
t.Skip("codex CLI not found in PATH, skipping integration test")
|
||||||
|
}
|
||||||
|
t.Logf("Using codex CLI at: %s", path)
|
||||||
|
|
||||||
|
p := NewCodexCliProvider(t.TempDir())
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 90*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
resp, err := p.Chat(ctx, []Message{
|
||||||
|
{Role: "user", Content: "Respond with only the word 'pong'. Nothing else."},
|
||||||
|
}, nil, "", nil)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() with real CLI error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.Content == "" {
|
||||||
|
t.Error("Content is empty")
|
||||||
|
}
|
||||||
|
if resp.FinishReason != "stop" {
|
||||||
|
t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "stop")
|
||||||
|
}
|
||||||
|
if resp.Usage != nil {
|
||||||
|
t.Logf("Usage: prompt=%d, completion=%d, total=%d",
|
||||||
|
resp.Usage.PromptTokens, resp.Usage.CompletionTokens, resp.Usage.TotalTokens)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("Response content: %q", resp.Content)
|
||||||
|
|
||||||
|
if !strings.Contains(strings.ToLower(resp.Content), "pong") {
|
||||||
|
t.Errorf("Content = %q, expected to contain 'pong'", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIntegration_RealCodexCLI_WithSystemPrompt(t *testing.T) {
|
||||||
|
if _, err := exec.LookPath("codex"); err != nil {
|
||||||
|
t.Skip("codex CLI not found in PATH")
|
||||||
|
}
|
||||||
|
|
||||||
|
p := NewCodexCliProvider(t.TempDir())
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 90*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
resp, err := p.Chat(ctx, []Message{
|
||||||
|
{Role: "system", Content: "You are a calculator. Only respond with numbers. No text."},
|
||||||
|
{Role: "user", Content: "What is 2+2?"},
|
||||||
|
}, nil, "", nil)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("Response: %q", resp.Content)
|
||||||
|
|
||||||
|
if !strings.Contains(resp.Content, "4") {
|
||||||
|
t.Errorf("Content = %q, expected to contain '4'", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIntegration_RealCodexCLI_ParsesRealJSONL(t *testing.T) {
|
||||||
|
if _, err := exec.LookPath("codex"); err != nil {
|
||||||
|
t.Skip("codex CLI not found in PATH")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run codex directly and verify our parser handles real output
|
||||||
|
cmd := exec.Command("codex", "exec",
|
||||||
|
"--json",
|
||||||
|
"--dangerously-bypass-approvals-and-sandbox",
|
||||||
|
"--skip-git-repo-check",
|
||||||
|
"--color", "never",
|
||||||
|
"-C", t.TempDir(),
|
||||||
|
"-")
|
||||||
|
cmd.Stdin = strings.NewReader("Say hi")
|
||||||
|
|
||||||
|
output, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
// codex may write diagnostic noise to stderr but still produce valid output
|
||||||
|
if len(output) == 0 {
|
||||||
|
t.Fatalf("codex CLI failed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("Raw CLI output (first 500 chars): %s", string(output[:min(len(output), 500)]))
|
||||||
|
|
||||||
|
// Verify our parser can handle real output
|
||||||
|
p := NewCodexCliProvider("")
|
||||||
|
resp, err := p.parseJSONLEvents(string(output))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parseJSONLEvents() failed on real CLI output: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.Content == "" {
|
||||||
|
t.Error("parsed Content is empty")
|
||||||
|
}
|
||||||
|
if resp.FinishReason != "stop" {
|
||||||
|
t.Errorf("FinishReason = %q, want stop", resp.FinishReason)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("Parsed: content=%q, finish=%s, usage=%+v", resp.Content, resp.FinishReason, resp.Usage)
|
||||||
|
}
|
||||||
|
|
@ -18,9 +18,10 @@ const codexDefaultModel = "gpt-5.2"
|
||||||
const codexDefaultInstructions = "You are Codex, a coding assistant."
|
const codexDefaultInstructions = "You are Codex, a coding assistant."
|
||||||
|
|
||||||
type CodexProvider struct {
|
type CodexProvider struct {
|
||||||
client *openai.Client
|
client *openai.Client
|
||||||
accountID string
|
accountID string
|
||||||
tokenSource func() (string, string, error)
|
tokenSource func() (string, string, error)
|
||||||
|
enableWebSearch bool
|
||||||
}
|
}
|
||||||
|
|
||||||
const defaultCodexInstructions = "You are Codex, a coding assistant."
|
const defaultCodexInstructions = "You are Codex, a coding assistant."
|
||||||
|
|
@ -37,8 +38,9 @@ func NewCodexProvider(token, accountID string) *CodexProvider {
|
||||||
}
|
}
|
||||||
client := openai.NewClient(opts...)
|
client := openai.NewClient(opts...)
|
||||||
return &CodexProvider{
|
return &CodexProvider{
|
||||||
client: &client,
|
client: &client,
|
||||||
accountID: accountID,
|
accountID: accountID,
|
||||||
|
enableWebSearch: true,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -78,7 +80,7 @@ func (p *CodexProvider) Chat(ctx context.Context, messages []Message, tools []To
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
params := buildCodexParams(messages, tools, resolvedModel, options)
|
params := buildCodexParams(messages, tools, resolvedModel, options, p.enableWebSearch)
|
||||||
|
|
||||||
stream := p.client.Responses.NewStreaming(ctx, params, opts...)
|
stream := p.client.Responses.NewStreaming(ctx, params, opts...)
|
||||||
defer stream.Close()
|
defer stream.Close()
|
||||||
|
|
@ -182,7 +184,7 @@ func resolveCodexModel(model string) (string, string) {
|
||||||
return codexDefaultModel, "unsupported model family"
|
return codexDefaultModel, "unsupported model family"
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildCodexParams(messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) responses.ResponseNewParams {
|
func buildCodexParams(messages []Message, tools []ToolDefinition, model string, options map[string]interface{}, enableWebSearch bool) responses.ResponseNewParams {
|
||||||
var inputItems responses.ResponseInputParam
|
var inputItems responses.ResponseInputParam
|
||||||
var instructions string
|
var instructions string
|
||||||
|
|
||||||
|
|
@ -217,12 +219,18 @@ func buildCodexParams(messages []Message, tools []ToolDefinition, model string,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
for _, tc := range msg.ToolCalls {
|
for _, tc := range msg.ToolCalls {
|
||||||
argsJSON, _ := json.Marshal(tc.Arguments)
|
name, args, ok := resolveCodexToolCall(tc)
|
||||||
|
if !ok {
|
||||||
|
logger.WarnCF("provider.codex", "Skipping invalid tool call in history", map[string]interface{}{
|
||||||
|
"call_id": tc.ID,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
inputItems = append(inputItems, responses.ResponseInputItemUnionParam{
|
inputItems = append(inputItems, responses.ResponseInputItemUnionParam{
|
||||||
OfFunctionCall: &responses.ResponseFunctionToolCallParam{
|
OfFunctionCall: &responses.ResponseFunctionToolCallParam{
|
||||||
CallID: tc.ID,
|
CallID: tc.ID,
|
||||||
Name: tc.Name,
|
Name: name,
|
||||||
Arguments: string(argsJSON),
|
Arguments: args,
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
@ -260,20 +268,50 @@ func buildCodexParams(messages []Message, tools []ToolDefinition, model string,
|
||||||
params.Instructions = openai.Opt(defaultCodexInstructions)
|
params.Instructions = openai.Opt(defaultCodexInstructions)
|
||||||
}
|
}
|
||||||
|
|
||||||
if maxTokens, ok := options["max_tokens"].(int); ok {
|
if len(tools) > 0 || enableWebSearch {
|
||||||
params.MaxOutputTokens = openai.Opt(int64(maxTokens))
|
params.Tools = translateToolsForCodex(tools, enableWebSearch)
|
||||||
}
|
|
||||||
|
|
||||||
if len(tools) > 0 {
|
|
||||||
params.Tools = translateToolsForCodex(tools)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return params
|
return params
|
||||||
}
|
}
|
||||||
|
|
||||||
func translateToolsForCodex(tools []ToolDefinition) []responses.ToolUnionParam {
|
func resolveCodexToolCall(tc ToolCall) (name string, arguments string, ok bool) {
|
||||||
result := make([]responses.ToolUnionParam, 0, len(tools))
|
name = tc.Name
|
||||||
|
if name == "" && tc.Function != nil {
|
||||||
|
name = tc.Function.Name
|
||||||
|
}
|
||||||
|
if name == "" {
|
||||||
|
return "", "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(tc.Arguments) > 0 {
|
||||||
|
argsJSON, err := json.Marshal(tc.Arguments)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", false
|
||||||
|
}
|
||||||
|
return name, string(argsJSON), true
|
||||||
|
}
|
||||||
|
|
||||||
|
if tc.Function != nil && tc.Function.Arguments != "" {
|
||||||
|
return name, tc.Function.Arguments, true
|
||||||
|
}
|
||||||
|
|
||||||
|
return name, "{}", true
|
||||||
|
}
|
||||||
|
|
||||||
|
func translateToolsForCodex(tools []ToolDefinition, enableWebSearch bool) []responses.ToolUnionParam {
|
||||||
|
capHint := len(tools)
|
||||||
|
if enableWebSearch {
|
||||||
|
capHint++
|
||||||
|
}
|
||||||
|
result := make([]responses.ToolUnionParam, 0, capHint)
|
||||||
for _, t := range tools {
|
for _, t := range tools {
|
||||||
|
if t.Type != "function" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if enableWebSearch && strings.EqualFold(t.Function.Name, "web_search") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
ft := responses.FunctionToolParam{
|
ft := responses.FunctionToolParam{
|
||||||
Name: t.Function.Name,
|
Name: t.Function.Name,
|
||||||
Parameters: t.Function.Parameters,
|
Parameters: t.Function.Parameters,
|
||||||
|
|
@ -284,6 +322,9 @@ func translateToolsForCodex(tools []ToolDefinition) []responses.ToolUnionParam {
|
||||||
}
|
}
|
||||||
result = append(result, responses.ToolUnionParam{OfFunction: &ft})
|
result = append(result, responses.ToolUnionParam{OfFunction: &ft})
|
||||||
}
|
}
|
||||||
|
if enableWebSearch {
|
||||||
|
result = append(result, responses.ToolParamOfWebSearch(responses.WebSearchToolTypeWebSearch))
|
||||||
|
}
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,7 @@ func TestBuildCodexParams_BasicMessage(t *testing.T) {
|
||||||
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{
|
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{
|
||||||
"max_tokens": 2048,
|
"max_tokens": 2048,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
})
|
}, true)
|
||||||
if params.Model != "gpt-4o" {
|
if params.Model != "gpt-4o" {
|
||||||
t.Errorf("Model = %q, want %q", params.Model, "gpt-4o")
|
t.Errorf("Model = %q, want %q", params.Model, "gpt-4o")
|
||||||
}
|
}
|
||||||
|
|
@ -29,6 +29,9 @@ func TestBuildCodexParams_BasicMessage(t *testing.T) {
|
||||||
if params.Instructions.Or("") != defaultCodexInstructions {
|
if params.Instructions.Or("") != defaultCodexInstructions {
|
||||||
t.Errorf("Instructions = %q, want %q", params.Instructions.Or(""), defaultCodexInstructions)
|
t.Errorf("Instructions = %q, want %q", params.Instructions.Or(""), defaultCodexInstructions)
|
||||||
}
|
}
|
||||||
|
if params.MaxOutputTokens.Valid() {
|
||||||
|
t.Fatalf("MaxOutputTokens should not be set for Codex backend")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBuildCodexParams_SystemAsInstructions(t *testing.T) {
|
func TestBuildCodexParams_SystemAsInstructions(t *testing.T) {
|
||||||
|
|
@ -36,7 +39,7 @@ func TestBuildCodexParams_SystemAsInstructions(t *testing.T) {
|
||||||
{Role: "system", Content: "You are helpful"},
|
{Role: "system", Content: "You are helpful"},
|
||||||
{Role: "user", Content: "Hi"},
|
{Role: "user", Content: "Hi"},
|
||||||
}
|
}
|
||||||
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{})
|
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{}, true)
|
||||||
if !params.Instructions.Valid() {
|
if !params.Instructions.Valid() {
|
||||||
t.Fatal("Instructions should be set")
|
t.Fatal("Instructions should be set")
|
||||||
}
|
}
|
||||||
|
|
@ -56,7 +59,7 @@ func TestBuildCodexParams_ToolCallConversation(t *testing.T) {
|
||||||
},
|
},
|
||||||
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
|
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
|
||||||
}
|
}
|
||||||
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{})
|
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{}, false)
|
||||||
if params.Input.OfInputItemList == nil {
|
if params.Input.OfInputItemList == nil {
|
||||||
t.Fatal("Input.OfInputItemList should not be nil")
|
t.Fatal("Input.OfInputItemList should not be nil")
|
||||||
}
|
}
|
||||||
|
|
@ -65,6 +68,45 @@ func TestBuildCodexParams_ToolCallConversation(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBuildCodexParams_ToolCallFunctionFallback(t *testing.T) {
|
||||||
|
messages := []Message{
|
||||||
|
{Role: "user", Content: "Read a file"},
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
ToolCalls: []ToolCall{
|
||||||
|
{
|
||||||
|
ID: "call_1",
|
||||||
|
Type: "function",
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: "read_file",
|
||||||
|
Arguments: `{"path":"README.md"}`,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{Role: "tool", Content: "ok", ToolCallID: "call_1"},
|
||||||
|
}
|
||||||
|
|
||||||
|
params := buildCodexParams(messages, nil, "gpt-4o", map[string]interface{}{}, false)
|
||||||
|
if params.Input.OfInputItemList == nil {
|
||||||
|
t.Fatal("Input.OfInputItemList should not be nil")
|
||||||
|
}
|
||||||
|
if len(params.Input.OfInputItemList) != 3 {
|
||||||
|
t.Fatalf("len(Input items) = %d, want 3", len(params.Input.OfInputItemList))
|
||||||
|
}
|
||||||
|
|
||||||
|
fc := params.Input.OfInputItemList[1].OfFunctionCall
|
||||||
|
if fc == nil {
|
||||||
|
t.Fatal("assistant tool call should be converted to function_call input item")
|
||||||
|
}
|
||||||
|
if fc.Name != "read_file" {
|
||||||
|
t.Errorf("Function call name = %q, want %q", fc.Name, "read_file")
|
||||||
|
}
|
||||||
|
if fc.Arguments != `{"path":"README.md"}` {
|
||||||
|
t.Errorf("Function call arguments = %q, want %q", fc.Arguments, `{"path":"README.md"}`)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestBuildCodexParams_WithTools(t *testing.T) {
|
func TestBuildCodexParams_WithTools(t *testing.T) {
|
||||||
tools := []ToolDefinition{
|
tools := []ToolDefinition{
|
||||||
{
|
{
|
||||||
|
|
@ -81,7 +123,7 @@ func TestBuildCodexParams_WithTools(t *testing.T) {
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
params := buildCodexParams([]Message{{Role: "user", Content: "Hi"}}, tools, "gpt-4o", map[string]interface{}{})
|
params := buildCodexParams([]Message{{Role: "user", Content: "Hi"}}, tools, "gpt-4o", map[string]interface{}{}, false)
|
||||||
if len(params.Tools) != 1 {
|
if len(params.Tools) != 1 {
|
||||||
t.Fatalf("len(Tools) = %d, want 1", len(params.Tools))
|
t.Fatalf("len(Tools) = %d, want 1", len(params.Tools))
|
||||||
}
|
}
|
||||||
|
|
@ -94,12 +136,61 @@ func TestBuildCodexParams_WithTools(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBuildCodexParams_StoreIsFalse(t *testing.T) {
|
func TestBuildCodexParams_StoreIsFalse(t *testing.T) {
|
||||||
params := buildCodexParams([]Message{{Role: "user", Content: "Hi"}}, nil, "gpt-4o", map[string]interface{}{})
|
params := buildCodexParams([]Message{{Role: "user", Content: "Hi"}}, nil, "gpt-4o", map[string]interface{}{}, false)
|
||||||
if !params.Store.Valid() || params.Store.Or(true) != false {
|
if !params.Store.Valid() || params.Store.Or(true) != false {
|
||||||
t.Error("Store should be explicitly set to false")
|
t.Error("Store should be explicitly set to false")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBuildCodexParams_DefaultWebSearchEnabled(t *testing.T) {
|
||||||
|
params := buildCodexParams([]Message{{Role: "user", Content: "Hi"}}, nil, "gpt-4o", map[string]interface{}{}, true)
|
||||||
|
if len(params.Tools) != 1 {
|
||||||
|
t.Fatalf("len(Tools) = %d, want 1", len(params.Tools))
|
||||||
|
}
|
||||||
|
if params.Tools[0].OfWebSearch == nil {
|
||||||
|
t.Fatal("Tool should include built-in web_search")
|
||||||
|
}
|
||||||
|
if params.Tools[0].OfWebSearch.Type != responses.WebSearchToolTypeWebSearch {
|
||||||
|
t.Errorf("Web search tool type = %q, want %q", params.Tools[0].OfWebSearch.Type, responses.WebSearchToolTypeWebSearch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildCodexParams_WebSearchFunctionReplacedWithBuiltin(t *testing.T) {
|
||||||
|
tools := []ToolDefinition{
|
||||||
|
{
|
||||||
|
Type: "function",
|
||||||
|
Function: ToolFunctionDefinition{
|
||||||
|
Name: "web_search",
|
||||||
|
Description: "local web search",
|
||||||
|
Parameters: map[string]interface{}{
|
||||||
|
"type": "object",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: "function",
|
||||||
|
Function: ToolFunctionDefinition{
|
||||||
|
Name: "read_file",
|
||||||
|
Description: "read file",
|
||||||
|
Parameters: map[string]interface{}{
|
||||||
|
"type": "object",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
params := buildCodexParams([]Message{{Role: "user", Content: "Hi"}}, tools, "gpt-4o", map[string]interface{}{}, true)
|
||||||
|
if len(params.Tools) != 2 {
|
||||||
|
t.Fatalf("len(Tools) = %d, want 2", len(params.Tools))
|
||||||
|
}
|
||||||
|
if params.Tools[0].OfFunction == nil || params.Tools[0].OfFunction.Name != "read_file" {
|
||||||
|
t.Fatalf("first tool should be function read_file, got %#v", params.Tools[0])
|
||||||
|
}
|
||||||
|
if params.Tools[1].OfWebSearch == nil {
|
||||||
|
t.Fatalf("second tool should be built-in web_search, got %#v", params.Tools[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestParseCodexResponse_TextOutput(t *testing.T) {
|
func TestParseCodexResponse_TextOutput(t *testing.T) {
|
||||||
respJSON := `{
|
respJSON := `{
|
||||||
"id": "resp_test",
|
"id": "resp_test",
|
||||||
|
|
@ -214,6 +305,20 @@ func TestCodexProvider_ChatRoundTrip(t *testing.T) {
|
||||||
http.Error(w, "stream must be true", http.StatusBadRequest)
|
http.Error(w, "stream must be true", http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if _, ok := reqBody["max_output_tokens"]; ok {
|
||||||
|
http.Error(w, "max_output_tokens is not supported", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
toolsAny, ok := reqBody["tools"].([]interface{})
|
||||||
|
if !ok || len(toolsAny) != 1 {
|
||||||
|
http.Error(w, "missing default web search tool", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
toolObj, ok := toolsAny[0].(map[string]interface{})
|
||||||
|
if !ok || toolObj["type"] != "web_search" {
|
||||||
|
http.Error(w, "expected web_search tool", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
resp := map[string]interface{}{
|
resp := map[string]interface{}{
|
||||||
"id": "resp_test",
|
"id": "resp_test",
|
||||||
|
|
@ -261,6 +366,64 @@ func TestCodexProvider_ChatRoundTrip(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCodexProvider_ChatRoundTrip_WebSearchDisabled(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.URL.Path != "/responses" {
|
||||||
|
http.Error(w, "not found: "+r.URL.Path, http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var reqBody map[string]interface{}
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&reqBody); err != nil {
|
||||||
|
http.Error(w, "invalid json", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if _, ok := reqBody["tools"]; ok {
|
||||||
|
http.Error(w, "tools should be absent when web search disabled", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"id": "resp_test",
|
||||||
|
"object": "response",
|
||||||
|
"status": "completed",
|
||||||
|
"output": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"id": "msg_1",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"status": "completed",
|
||||||
|
"content": []map[string]interface{}{
|
||||||
|
{"type": "output_text", "text": "Hi from Codex!"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"usage": map[string]interface{}{
|
||||||
|
"input_tokens": 4,
|
||||||
|
"output_tokens": 3,
|
||||||
|
"total_tokens": 7,
|
||||||
|
"input_tokens_details": map[string]interface{}{"cached_tokens": 0},
|
||||||
|
"output_tokens_details": map[string]interface{}{"reasoning_tokens": 0},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
writeCompletedSSE(w, resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
provider := NewCodexProvider("test-token", "acc-123")
|
||||||
|
provider.enableWebSearch = false
|
||||||
|
provider.client = createOpenAITestClient(server.URL, "test-token", "acc-123")
|
||||||
|
|
||||||
|
messages := []Message{{Role: "user", Content: "Hello"}}
|
||||||
|
resp, err := provider.Chat(t.Context(), messages, nil, "gpt-4o", map[string]interface{}{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error: %v", err)
|
||||||
|
}
|
||||||
|
if resp.Content != "Hi from Codex!" {
|
||||||
|
t.Errorf("Content = %q, want %q", resp.Content, "Hi from Codex!")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCodexProvider_ChatRoundTrip_TokenSourceFallbackAccountID(t *testing.T) {
|
func TestCodexProvider_ChatRoundTrip_TokenSourceFallbackAccountID(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.URL.Path != "/responses" {
|
if r.URL.Path != "/responses" {
|
||||||
|
|
@ -293,6 +456,10 @@ func TestCodexProvider_ChatRoundTrip_TokenSourceFallbackAccountID(t *testing.T)
|
||||||
http.Error(w, "temperature is not supported", http.StatusBadRequest)
|
http.Error(w, "temperature is not supported", http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if _, ok := reqBody["max_output_tokens"]; ok {
|
||||||
|
http.Error(w, "max_output_tokens is not supported", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
if reqBody["stream"] != true {
|
if reqBody["stream"] != true {
|
||||||
http.Error(w, "stream must be true", http.StatusBadRequest)
|
http.Error(w, "stream must be true", http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
|
|
|
||||||
207
pkg/providers/cooldown.go
Normal file
207
pkg/providers/cooldown.go
Normal file
|
|
@ -0,0 +1,207 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultFailureWindow = 24 * time.Hour
|
||||||
|
)
|
||||||
|
|
||||||
|
// CooldownTracker manages per-provider cooldown state for the fallback chain.
|
||||||
|
// Thread-safe via sync.RWMutex. In-memory only (resets on restart).
|
||||||
|
type CooldownTracker struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
entries map[string]*cooldownEntry
|
||||||
|
failureWindow time.Duration
|
||||||
|
nowFunc func() time.Time // for testing
|
||||||
|
}
|
||||||
|
|
||||||
|
type cooldownEntry struct {
|
||||||
|
ErrorCount int
|
||||||
|
FailureCounts map[FailoverReason]int
|
||||||
|
CooldownEnd time.Time // standard cooldown expiry
|
||||||
|
DisabledUntil time.Time // billing-specific disable expiry
|
||||||
|
DisabledReason FailoverReason // reason for disable (billing)
|
||||||
|
LastFailure time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCooldownTracker creates a tracker with default 24h failure window.
|
||||||
|
func NewCooldownTracker() *CooldownTracker {
|
||||||
|
return &CooldownTracker{
|
||||||
|
entries: make(map[string]*cooldownEntry),
|
||||||
|
failureWindow: defaultFailureWindow,
|
||||||
|
nowFunc: time.Now,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkFailure records a failure for a provider and sets appropriate cooldown.
|
||||||
|
// Resets error counts if last failure was more than failureWindow ago.
|
||||||
|
func (ct *CooldownTracker) MarkFailure(provider string, reason FailoverReason) {
|
||||||
|
ct.mu.Lock()
|
||||||
|
defer ct.mu.Unlock()
|
||||||
|
|
||||||
|
now := ct.nowFunc()
|
||||||
|
entry := ct.getOrCreate(provider)
|
||||||
|
|
||||||
|
// 24h failure window reset: if no failure in failureWindow, reset counters.
|
||||||
|
if !entry.LastFailure.IsZero() && now.Sub(entry.LastFailure) > ct.failureWindow {
|
||||||
|
entry.ErrorCount = 0
|
||||||
|
entry.FailureCounts = make(map[FailoverReason]int)
|
||||||
|
}
|
||||||
|
|
||||||
|
entry.ErrorCount++
|
||||||
|
entry.FailureCounts[reason]++
|
||||||
|
entry.LastFailure = now
|
||||||
|
|
||||||
|
if reason == FailoverBilling {
|
||||||
|
billingCount := entry.FailureCounts[FailoverBilling]
|
||||||
|
entry.DisabledUntil = now.Add(calculateBillingCooldown(billingCount))
|
||||||
|
entry.DisabledReason = FailoverBilling
|
||||||
|
} else {
|
||||||
|
entry.CooldownEnd = now.Add(calculateStandardCooldown(entry.ErrorCount))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkSuccess resets all counters and cooldowns for a provider.
|
||||||
|
func (ct *CooldownTracker) MarkSuccess(provider string) {
|
||||||
|
ct.mu.Lock()
|
||||||
|
defer ct.mu.Unlock()
|
||||||
|
|
||||||
|
entry := ct.entries[provider]
|
||||||
|
if entry == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
entry.ErrorCount = 0
|
||||||
|
entry.FailureCounts = make(map[FailoverReason]int)
|
||||||
|
entry.CooldownEnd = time.Time{}
|
||||||
|
entry.DisabledUntil = time.Time{}
|
||||||
|
entry.DisabledReason = ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsAvailable returns true if the provider is not in cooldown or disabled.
|
||||||
|
func (ct *CooldownTracker) IsAvailable(provider string) bool {
|
||||||
|
ct.mu.RLock()
|
||||||
|
defer ct.mu.RUnlock()
|
||||||
|
|
||||||
|
entry := ct.entries[provider]
|
||||||
|
if entry == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
now := ct.nowFunc()
|
||||||
|
|
||||||
|
// Billing disable takes precedence (longer cooldown).
|
||||||
|
if !entry.DisabledUntil.IsZero() && now.Before(entry.DisabledUntil) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Standard cooldown.
|
||||||
|
if !entry.CooldownEnd.IsZero() && now.Before(entry.CooldownEnd) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// CooldownRemaining returns how long until the provider becomes available.
|
||||||
|
// Returns 0 if already available.
|
||||||
|
func (ct *CooldownTracker) CooldownRemaining(provider string) time.Duration {
|
||||||
|
ct.mu.RLock()
|
||||||
|
defer ct.mu.RUnlock()
|
||||||
|
|
||||||
|
entry := ct.entries[provider]
|
||||||
|
if entry == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
now := ct.nowFunc()
|
||||||
|
var remaining time.Duration
|
||||||
|
|
||||||
|
if !entry.DisabledUntil.IsZero() && now.Before(entry.DisabledUntil) {
|
||||||
|
d := entry.DisabledUntil.Sub(now)
|
||||||
|
if d > remaining {
|
||||||
|
remaining = d
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !entry.CooldownEnd.IsZero() && now.Before(entry.CooldownEnd) {
|
||||||
|
d := entry.CooldownEnd.Sub(now)
|
||||||
|
if d > remaining {
|
||||||
|
remaining = d
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return remaining
|
||||||
|
}
|
||||||
|
|
||||||
|
// ErrorCount returns the current error count for a provider.
|
||||||
|
func (ct *CooldownTracker) ErrorCount(provider string) int {
|
||||||
|
ct.mu.RLock()
|
||||||
|
defer ct.mu.RUnlock()
|
||||||
|
|
||||||
|
entry := ct.entries[provider]
|
||||||
|
if entry == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return entry.ErrorCount
|
||||||
|
}
|
||||||
|
|
||||||
|
// FailureCount returns the failure count for a specific reason.
|
||||||
|
func (ct *CooldownTracker) FailureCount(provider string, reason FailoverReason) int {
|
||||||
|
ct.mu.RLock()
|
||||||
|
defer ct.mu.RUnlock()
|
||||||
|
|
||||||
|
entry := ct.entries[provider]
|
||||||
|
if entry == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return entry.FailureCounts[reason]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ct *CooldownTracker) getOrCreate(provider string) *cooldownEntry {
|
||||||
|
entry := ct.entries[provider]
|
||||||
|
if entry == nil {
|
||||||
|
entry = &cooldownEntry{
|
||||||
|
FailureCounts: make(map[FailoverReason]int),
|
||||||
|
}
|
||||||
|
ct.entries[provider] = entry
|
||||||
|
}
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
|
||||||
|
// calculateStandardCooldown computes standard exponential backoff.
|
||||||
|
// Formula from OpenClaw: min(1h, 1min * 5^min(n-1, 3))
|
||||||
|
//
|
||||||
|
// 1 error → 1 min
|
||||||
|
// 2 errors → 5 min
|
||||||
|
// 3 errors → 25 min
|
||||||
|
// 4+ errors → 1 hour (cap)
|
||||||
|
func calculateStandardCooldown(errorCount int) time.Duration {
|
||||||
|
n := max(1, errorCount)
|
||||||
|
exp := min(n-1, 3)
|
||||||
|
ms := 60_000 * int(math.Pow(5, float64(exp)))
|
||||||
|
ms = min(3_600_000, ms) // cap at 1 hour
|
||||||
|
return time.Duration(ms) * time.Millisecond
|
||||||
|
}
|
||||||
|
|
||||||
|
// calculateBillingCooldown computes billing-specific exponential backoff.
|
||||||
|
// Formula from OpenClaw: min(24h, 5h * 2^min(n-1, 10))
|
||||||
|
//
|
||||||
|
// 1 error → 5 hours
|
||||||
|
// 2 errors → 10 hours
|
||||||
|
// 3 errors → 20 hours
|
||||||
|
// 4+ errors → 24 hours (cap)
|
||||||
|
func calculateBillingCooldown(billingErrorCount int) time.Duration {
|
||||||
|
const baseMs = 5 * 60 * 60 * 1000 // 5 hours
|
||||||
|
const maxMs = 24 * 60 * 60 * 1000 // 24 hours
|
||||||
|
|
||||||
|
n := max(1, billingErrorCount)
|
||||||
|
exp := min(n-1, 10)
|
||||||
|
raw := float64(baseMs) * math.Pow(2, float64(exp))
|
||||||
|
ms := int(math.Min(float64(maxMs), raw))
|
||||||
|
return time.Duration(ms) * time.Millisecond
|
||||||
|
}
|
||||||
269
pkg/providers/cooldown_test.go
Normal file
269
pkg/providers/cooldown_test.go
Normal file
|
|
@ -0,0 +1,269 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestTracker(now time.Time) (*CooldownTracker, *time.Time) {
|
||||||
|
current := now
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
ct.nowFunc = func() time.Time { return current }
|
||||||
|
return ct, ¤t
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_InitiallyAvailable(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("new provider should be available")
|
||||||
|
}
|
||||||
|
if ct.ErrorCount("openai") != 0 {
|
||||||
|
t.Error("new provider should have 0 errors")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_StandardEscalation(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
ct, current := newTestTracker(now)
|
||||||
|
|
||||||
|
// 1st error → 1 min cooldown
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
if ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be in cooldown after 1st error")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Advance 61 seconds → available
|
||||||
|
*current = now.Add(61 * time.Second)
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be available after 1 min cooldown")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2nd error → 5 min cooldown
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
*current = now.Add(61*time.Second + 4*time.Minute)
|
||||||
|
if ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be in cooldown (5 min) after 2nd error")
|
||||||
|
}
|
||||||
|
*current = now.Add(61*time.Second + 6*time.Minute)
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be available after 5 min cooldown")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_StandardCap(t *testing.T) {
|
||||||
|
// Verify formula: 1m, 5m, 25m, 1h, 1h, 1h...
|
||||||
|
expected := []time.Duration{
|
||||||
|
1 * time.Minute,
|
||||||
|
5 * time.Minute,
|
||||||
|
25 * time.Minute,
|
||||||
|
1 * time.Hour,
|
||||||
|
1 * time.Hour,
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, want := range expected {
|
||||||
|
got := calculateStandardCooldown(i + 1)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("calculateStandardCooldown(%d) = %v, want %v", i+1, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_BillingEscalation(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
ct, current := newTestTracker(now)
|
||||||
|
|
||||||
|
// 1st billing error → 5h cooldown
|
||||||
|
ct.MarkFailure("openai", FailoverBilling)
|
||||||
|
if ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be disabled after billing error")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Advance 4h → still disabled
|
||||||
|
*current = now.Add(4 * time.Hour)
|
||||||
|
if ct.IsAvailable("openai") {
|
||||||
|
t.Error("should still be disabled (5h cooldown)")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Advance 5h + 1s → available
|
||||||
|
*current = now.Add(5*time.Hour + 1*time.Second)
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be available after 5h billing cooldown")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_BillingCap(t *testing.T) {
|
||||||
|
expected := []time.Duration{
|
||||||
|
5 * time.Hour,
|
||||||
|
10 * time.Hour,
|
||||||
|
20 * time.Hour,
|
||||||
|
24 * time.Hour,
|
||||||
|
24 * time.Hour,
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, want := range expected {
|
||||||
|
got := calculateBillingCooldown(i + 1)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("calculateBillingCooldown(%d) = %v, want %v", i+1, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_SuccessReset(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
ct.MarkFailure("openai", FailoverBilling)
|
||||||
|
if ct.ErrorCount("openai") != 2 {
|
||||||
|
t.Errorf("error count = %d, want 2", ct.ErrorCount("openai"))
|
||||||
|
}
|
||||||
|
|
||||||
|
ct.MarkSuccess("openai")
|
||||||
|
if ct.ErrorCount("openai") != 0 {
|
||||||
|
t.Errorf("error count after success = %d, want 0", ct.ErrorCount("openai"))
|
||||||
|
}
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be available after success")
|
||||||
|
}
|
||||||
|
if ct.FailureCount("openai", FailoverRateLimit) != 0 {
|
||||||
|
t.Error("failure counts should be reset after success")
|
||||||
|
}
|
||||||
|
if ct.FailureCount("openai", FailoverBilling) != 0 {
|
||||||
|
t.Error("billing failure count should be reset after success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_FailureWindowReset(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
ct, current := newTestTracker(now)
|
||||||
|
|
||||||
|
// 4 errors → 1h cooldown
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
*current = current.Add(2 * time.Second) // small advance between errors
|
||||||
|
}
|
||||||
|
if ct.ErrorCount("openai") != 4 {
|
||||||
|
t.Errorf("error count = %d, want 4", ct.ErrorCount("openai"))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Advance 25 hours (past 24h failure window)
|
||||||
|
*current = now.Add(25 * time.Hour)
|
||||||
|
|
||||||
|
// Next error should reset counters first, then increment to 1
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
if ct.ErrorCount("openai") != 1 {
|
||||||
|
t.Errorf("error count after window reset = %d, want 1 (reset + 1)", ct.ErrorCount("openai"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_PerReasonTracking(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
ct.MarkFailure("openai", FailoverBilling)
|
||||||
|
ct.MarkFailure("openai", FailoverAuth)
|
||||||
|
|
||||||
|
if ct.FailureCount("openai", FailoverRateLimit) != 2 {
|
||||||
|
t.Errorf("rate_limit count = %d, want 2", ct.FailureCount("openai", FailoverRateLimit))
|
||||||
|
}
|
||||||
|
if ct.FailureCount("openai", FailoverBilling) != 1 {
|
||||||
|
t.Errorf("billing count = %d, want 1", ct.FailureCount("openai", FailoverBilling))
|
||||||
|
}
|
||||||
|
if ct.FailureCount("openai", FailoverAuth) != 1 {
|
||||||
|
t.Errorf("auth count = %d, want 1", ct.FailureCount("openai", FailoverAuth))
|
||||||
|
}
|
||||||
|
if ct.ErrorCount("openai") != 4 {
|
||||||
|
t.Errorf("total error count = %d, want 4", ct.ErrorCount("openai"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_BillingTakesPrecedence(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
ct, current := newTestTracker(now)
|
||||||
|
|
||||||
|
// Standard cooldown (1 min) + billing disable (5h)
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit) // 1 min cooldown
|
||||||
|
ct.MarkFailure("openai", FailoverBilling) // 5h disable
|
||||||
|
|
||||||
|
// After 2 min: standard cooldown expired but billing still active
|
||||||
|
*current = now.Add(2 * time.Minute)
|
||||||
|
if ct.IsAvailable("openai") {
|
||||||
|
t.Error("billing disable should take precedence over standard cooldown")
|
||||||
|
}
|
||||||
|
|
||||||
|
// After 5h + 1s: both expired
|
||||||
|
*current = now.Add(5*time.Hour + 1*time.Second)
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("should be available after all cooldowns expire")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_CooldownRemaining(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
ct, current := newTestTracker(now)
|
||||||
|
|
||||||
|
// No failures → 0 remaining
|
||||||
|
if ct.CooldownRemaining("openai") != 0 {
|
||||||
|
t.Error("expected 0 remaining for new provider")
|
||||||
|
}
|
||||||
|
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
|
||||||
|
*current = now.Add(30 * time.Second)
|
||||||
|
remaining := ct.CooldownRemaining("openai")
|
||||||
|
if remaining <= 0 || remaining > 1*time.Minute {
|
||||||
|
t.Errorf("remaining = %v, expected ~30s", remaining)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_SuccessOnUnknownProvider(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
// Should not panic
|
||||||
|
ct.MarkSuccess("nonexistent")
|
||||||
|
if !ct.IsAvailable("nonexistent") {
|
||||||
|
t.Error("nonexistent provider should be available")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_ConcurrentAccess(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
wg.Add(3)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
}()
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
ct.IsAvailable("openai")
|
||||||
|
}()
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
ct.MarkSuccess("openai")
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
// If we got here without panic, concurrent access is safe
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCooldown_MultipleProviders(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
ct.MarkFailure("anthropic", FailoverBilling)
|
||||||
|
|
||||||
|
if ct.IsAvailable("openai") {
|
||||||
|
t.Error("openai should be in cooldown")
|
||||||
|
}
|
||||||
|
if ct.IsAvailable("anthropic") {
|
||||||
|
t.Error("anthropic should be in cooldown")
|
||||||
|
}
|
||||||
|
// groq was never touched
|
||||||
|
if !ct.IsAvailable("groq") {
|
||||||
|
t.Error("groq should be available")
|
||||||
|
}
|
||||||
|
}
|
||||||
253
pkg/providers/error_classifier.go
Normal file
253
pkg/providers/error_classifier.go
Normal file
|
|
@ -0,0 +1,253 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// errorPattern defines a single pattern (string or regex) for error classification.
|
||||||
|
type errorPattern struct {
|
||||||
|
substring string
|
||||||
|
regex *regexp.Regexp
|
||||||
|
}
|
||||||
|
|
||||||
|
func substr(s string) errorPattern { return errorPattern{substring: s} }
|
||||||
|
func rxp(r string) errorPattern { return errorPattern{regex: regexp.MustCompile("(?i)" + r)} }
|
||||||
|
|
||||||
|
// Error patterns organized by FailoverReason, matching OpenClaw production (~40 patterns).
|
||||||
|
var (
|
||||||
|
rateLimitPatterns = []errorPattern{
|
||||||
|
rxp(`rate[_ ]limit`),
|
||||||
|
substr("too many requests"),
|
||||||
|
substr("429"),
|
||||||
|
substr("exceeded your current quota"),
|
||||||
|
rxp(`exceeded.*quota`),
|
||||||
|
rxp(`resource has been exhausted`),
|
||||||
|
rxp(`resource.*exhausted`),
|
||||||
|
substr("resource_exhausted"),
|
||||||
|
substr("quota exceeded"),
|
||||||
|
substr("usage limit"),
|
||||||
|
}
|
||||||
|
|
||||||
|
overloadedPatterns = []errorPattern{
|
||||||
|
rxp(`overloaded_error`),
|
||||||
|
rxp(`"type"\s*:\s*"overloaded_error"`),
|
||||||
|
substr("overloaded"),
|
||||||
|
}
|
||||||
|
|
||||||
|
timeoutPatterns = []errorPattern{
|
||||||
|
substr("timeout"),
|
||||||
|
substr("timed out"),
|
||||||
|
substr("deadline exceeded"),
|
||||||
|
substr("context deadline exceeded"),
|
||||||
|
}
|
||||||
|
|
||||||
|
billingPatterns = []errorPattern{
|
||||||
|
rxp(`\b402\b`),
|
||||||
|
substr("payment required"),
|
||||||
|
substr("insufficient credits"),
|
||||||
|
substr("credit balance"),
|
||||||
|
substr("plans & billing"),
|
||||||
|
substr("insufficient balance"),
|
||||||
|
}
|
||||||
|
|
||||||
|
authPatterns = []errorPattern{
|
||||||
|
rxp(`invalid[_ ]?api[_ ]?key`),
|
||||||
|
substr("incorrect api key"),
|
||||||
|
substr("invalid token"),
|
||||||
|
substr("authentication"),
|
||||||
|
substr("re-authenticate"),
|
||||||
|
substr("oauth token refresh failed"),
|
||||||
|
substr("unauthorized"),
|
||||||
|
substr("forbidden"),
|
||||||
|
substr("access denied"),
|
||||||
|
substr("expired"),
|
||||||
|
substr("token has expired"),
|
||||||
|
rxp(`\b401\b`),
|
||||||
|
rxp(`\b403\b`),
|
||||||
|
substr("no credentials found"),
|
||||||
|
substr("no api key found"),
|
||||||
|
}
|
||||||
|
|
||||||
|
formatPatterns = []errorPattern{
|
||||||
|
substr("string should match pattern"),
|
||||||
|
substr("tool_use.id"),
|
||||||
|
substr("tool_use_id"),
|
||||||
|
substr("messages.1.content.1.tool_use.id"),
|
||||||
|
substr("invalid request format"),
|
||||||
|
}
|
||||||
|
|
||||||
|
imageDimensionPatterns = []errorPattern{
|
||||||
|
rxp(`image dimensions exceed max`),
|
||||||
|
}
|
||||||
|
|
||||||
|
imageSizePatterns = []errorPattern{
|
||||||
|
rxp(`image exceeds.*mb`),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Transient HTTP status codes that map to timeout (server-side failures).
|
||||||
|
transientStatusCodes = map[int]bool{
|
||||||
|
500: true, 502: true, 503: true,
|
||||||
|
521: true, 522: true, 523: true, 524: true,
|
||||||
|
529: true,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
// ClassifyError classifies an error into a FailoverError with reason.
|
||||||
|
// Returns nil if the error is not classifiable (unknown errors should not trigger fallback).
|
||||||
|
func ClassifyError(err error, provider, model string) *FailoverError {
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Context cancellation: user abort, never fallback.
|
||||||
|
if err == context.Canceled {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Context deadline exceeded: treat as timeout, always fallback.
|
||||||
|
if err == context.DeadlineExceeded {
|
||||||
|
return &FailoverError{
|
||||||
|
Reason: FailoverTimeout,
|
||||||
|
Provider: provider,
|
||||||
|
Model: model,
|
||||||
|
Wrapped: err,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := strings.ToLower(err.Error())
|
||||||
|
|
||||||
|
// Image dimension/size errors: non-retriable, non-fallback.
|
||||||
|
if IsImageDimensionError(msg) || IsImageSizeError(msg) {
|
||||||
|
return &FailoverError{
|
||||||
|
Reason: FailoverFormat,
|
||||||
|
Provider: provider,
|
||||||
|
Model: model,
|
||||||
|
Wrapped: err,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try HTTP status code extraction first.
|
||||||
|
if status := extractHTTPStatus(msg); status > 0 {
|
||||||
|
if reason := classifyByStatus(status); reason != "" {
|
||||||
|
return &FailoverError{
|
||||||
|
Reason: reason,
|
||||||
|
Provider: provider,
|
||||||
|
Model: model,
|
||||||
|
Status: status,
|
||||||
|
Wrapped: err,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Message pattern matching (priority order from OpenClaw).
|
||||||
|
if reason := classifyByMessage(msg); reason != "" {
|
||||||
|
return &FailoverError{
|
||||||
|
Reason: reason,
|
||||||
|
Provider: provider,
|
||||||
|
Model: model,
|
||||||
|
Wrapped: err,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyByStatus maps HTTP status codes to FailoverReason.
|
||||||
|
func classifyByStatus(status int) FailoverReason {
|
||||||
|
switch {
|
||||||
|
case status == 401 || status == 403:
|
||||||
|
return FailoverAuth
|
||||||
|
case status == 402:
|
||||||
|
return FailoverBilling
|
||||||
|
case status == 408:
|
||||||
|
return FailoverTimeout
|
||||||
|
case status == 429:
|
||||||
|
return FailoverRateLimit
|
||||||
|
case status == 400:
|
||||||
|
return FailoverFormat
|
||||||
|
case transientStatusCodes[status]:
|
||||||
|
return FailoverTimeout
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyByMessage matches error messages against patterns.
|
||||||
|
// Priority order matters (from OpenClaw classifyFailoverReason).
|
||||||
|
func classifyByMessage(msg string) FailoverReason {
|
||||||
|
if matchesAny(msg, rateLimitPatterns) {
|
||||||
|
return FailoverRateLimit
|
||||||
|
}
|
||||||
|
if matchesAny(msg, overloadedPatterns) {
|
||||||
|
return FailoverRateLimit // Overloaded treated as rate_limit
|
||||||
|
}
|
||||||
|
if matchesAny(msg, billingPatterns) {
|
||||||
|
return FailoverBilling
|
||||||
|
}
|
||||||
|
if matchesAny(msg, timeoutPatterns) {
|
||||||
|
return FailoverTimeout
|
||||||
|
}
|
||||||
|
if matchesAny(msg, authPatterns) {
|
||||||
|
return FailoverAuth
|
||||||
|
}
|
||||||
|
if matchesAny(msg, formatPatterns) {
|
||||||
|
return FailoverFormat
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractHTTPStatus extracts an HTTP status code from an error message.
|
||||||
|
// Looks for patterns like "status: 429", "status 429", "HTTP 429", or standalone "429".
|
||||||
|
func extractHTTPStatus(msg string) int {
|
||||||
|
// Common patterns in Go HTTP error messages
|
||||||
|
patterns := []*regexp.Regexp{
|
||||||
|
regexp.MustCompile(`status[:\s]+(\d{3})`),
|
||||||
|
regexp.MustCompile(`HTTP[/\s]+\d*\.?\d*\s+(\d{3})`),
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, p := range patterns {
|
||||||
|
if m := p.FindStringSubmatch(msg); len(m) > 1 {
|
||||||
|
return parseDigits(m[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsImageDimensionError returns true if the message indicates an image dimension error.
|
||||||
|
func IsImageDimensionError(msg string) bool {
|
||||||
|
return matchesAny(msg, imageDimensionPatterns)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsImageSizeError returns true if the message indicates an image file size error.
|
||||||
|
func IsImageSizeError(msg string) bool {
|
||||||
|
return matchesAny(msg, imageSizePatterns)
|
||||||
|
}
|
||||||
|
|
||||||
|
// matchesAny checks if msg matches any of the patterns.
|
||||||
|
func matchesAny(msg string, patterns []errorPattern) bool {
|
||||||
|
for _, p := range patterns {
|
||||||
|
if p.regex != nil {
|
||||||
|
if p.regex.MatchString(msg) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
} else if p.substring != "" {
|
||||||
|
if strings.Contains(msg, p.substring) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseDigits converts a string of digits to an int.
|
||||||
|
func parseDigits(s string) int {
|
||||||
|
n := 0
|
||||||
|
for _, c := range s {
|
||||||
|
if c >= '0' && c <= '9' {
|
||||||
|
n = n*10 + int(c-'0')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
337
pkg/providers/error_classifier_test.go
Normal file
337
pkg/providers/error_classifier_test.go
Normal file
|
|
@ -0,0 +1,337 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestClassifyError_Nil(t *testing.T) {
|
||||||
|
result := ClassifyError(nil, "openai", "gpt-4")
|
||||||
|
if result != nil {
|
||||||
|
t.Errorf("expected nil for nil error, got %+v", result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_ContextCanceled(t *testing.T) {
|
||||||
|
result := ClassifyError(context.Canceled, "openai", "gpt-4")
|
||||||
|
if result != nil {
|
||||||
|
t.Errorf("expected nil for context.Canceled (user abort), got %+v", result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_ContextDeadlineExceeded(t *testing.T) {
|
||||||
|
result := ClassifyError(context.DeadlineExceeded, "openai", "gpt-4")
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("expected non-nil for deadline exceeded")
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverTimeout {
|
||||||
|
t.Errorf("reason = %q, want timeout", result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_StatusCodes(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
status int
|
||||||
|
reason FailoverReason
|
||||||
|
}{
|
||||||
|
{401, FailoverAuth},
|
||||||
|
{403, FailoverAuth},
|
||||||
|
{402, FailoverBilling},
|
||||||
|
{408, FailoverTimeout},
|
||||||
|
{429, FailoverRateLimit},
|
||||||
|
{400, FailoverFormat},
|
||||||
|
{500, FailoverTimeout},
|
||||||
|
{502, FailoverTimeout},
|
||||||
|
{503, FailoverTimeout},
|
||||||
|
{521, FailoverTimeout},
|
||||||
|
{522, FailoverTimeout},
|
||||||
|
{523, FailoverTimeout},
|
||||||
|
{524, FailoverTimeout},
|
||||||
|
{529, FailoverTimeout},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
err := fmt.Errorf("API error: status: %d something went wrong", tt.status)
|
||||||
|
result := ClassifyError(err, "test", "model")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("status %d: expected non-nil", tt.status)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if result.Reason != tt.reason {
|
||||||
|
t.Errorf("status %d: reason = %q, want %q", tt.status, result.Reason, tt.reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_RateLimitPatterns(t *testing.T) {
|
||||||
|
patterns := []string{
|
||||||
|
"rate limit exceeded",
|
||||||
|
"rate_limit reached",
|
||||||
|
"too many requests",
|
||||||
|
"exceeded your current quota",
|
||||||
|
"resource has been exhausted",
|
||||||
|
"resource_exhausted",
|
||||||
|
"quota exceeded",
|
||||||
|
"usage limit reached",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range patterns {
|
||||||
|
err := errors.New(msg)
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("pattern %q: expected non-nil", msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverRateLimit {
|
||||||
|
t.Errorf("pattern %q: reason = %q, want rate_limit", msg, result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_OverloadedPatterns(t *testing.T) {
|
||||||
|
patterns := []string{
|
||||||
|
"overloaded_error",
|
||||||
|
`{"type": "overloaded_error"}`,
|
||||||
|
"server is overloaded",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range patterns {
|
||||||
|
err := errors.New(msg)
|
||||||
|
result := ClassifyError(err, "anthropic", "claude")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("pattern %q: expected non-nil", msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Overloaded is treated as rate_limit
|
||||||
|
if result.Reason != FailoverRateLimit {
|
||||||
|
t.Errorf("pattern %q: reason = %q, want rate_limit", msg, result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_BillingPatterns(t *testing.T) {
|
||||||
|
patterns := []string{
|
||||||
|
"payment required",
|
||||||
|
"insufficient credits",
|
||||||
|
"credit balance too low",
|
||||||
|
"plans & billing page",
|
||||||
|
"insufficient balance",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range patterns {
|
||||||
|
err := errors.New(msg)
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("pattern %q: expected non-nil", msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverBilling {
|
||||||
|
t.Errorf("pattern %q: reason = %q, want billing", msg, result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_TimeoutPatterns(t *testing.T) {
|
||||||
|
patterns := []string{
|
||||||
|
"request timeout",
|
||||||
|
"connection timed out",
|
||||||
|
"deadline exceeded",
|
||||||
|
"context deadline exceeded",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range patterns {
|
||||||
|
err := errors.New(msg)
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("pattern %q: expected non-nil", msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverTimeout {
|
||||||
|
t.Errorf("pattern %q: reason = %q, want timeout", msg, result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_AuthPatterns(t *testing.T) {
|
||||||
|
patterns := []string{
|
||||||
|
"invalid api key",
|
||||||
|
"invalid_api_key",
|
||||||
|
"incorrect api key",
|
||||||
|
"invalid token",
|
||||||
|
"authentication failed",
|
||||||
|
"re-authenticate",
|
||||||
|
"oauth token refresh failed",
|
||||||
|
"unauthorized access",
|
||||||
|
"forbidden",
|
||||||
|
"access denied",
|
||||||
|
"expired",
|
||||||
|
"token has expired",
|
||||||
|
"no credentials found",
|
||||||
|
"no api key found",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range patterns {
|
||||||
|
err := errors.New(msg)
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("pattern %q: expected non-nil", msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverAuth {
|
||||||
|
t.Errorf("pattern %q: reason = %q, want auth", msg, result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_FormatPatterns(t *testing.T) {
|
||||||
|
patterns := []string{
|
||||||
|
"string should match pattern",
|
||||||
|
"tool_use.id is required",
|
||||||
|
"invalid tool_use_id",
|
||||||
|
"messages.1.content.1.tool_use.id must be valid",
|
||||||
|
"invalid request format",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range patterns {
|
||||||
|
err := errors.New(msg)
|
||||||
|
result := ClassifyError(err, "anthropic", "claude")
|
||||||
|
if result == nil {
|
||||||
|
t.Errorf("pattern %q: expected non-nil", msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverFormat {
|
||||||
|
t.Errorf("pattern %q: reason = %q, want format", msg, result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_ImageDimensionError(t *testing.T) {
|
||||||
|
err := errors.New("image dimensions exceed max allowed 2048x2048")
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4o")
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("expected non-nil for image dimension error")
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverFormat {
|
||||||
|
t.Errorf("reason = %q, want format", result.Reason)
|
||||||
|
}
|
||||||
|
if result.IsRetriable() {
|
||||||
|
t.Error("image dimension error should not be retriable")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_ImageSizeError(t *testing.T) {
|
||||||
|
err := errors.New("image exceeds 20 mb limit")
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4o")
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("expected non-nil for image size error")
|
||||||
|
}
|
||||||
|
if result.Reason != FailoverFormat {
|
||||||
|
t.Errorf("reason = %q, want format", result.Reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_UnknownError(t *testing.T) {
|
||||||
|
err := errors.New("some completely random error")
|
||||||
|
result := ClassifyError(err, "openai", "gpt-4")
|
||||||
|
if result != nil {
|
||||||
|
t.Errorf("expected nil for unknown error, got %+v", result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassifyError_ProviderModelPropagation(t *testing.T) {
|
||||||
|
err := errors.New("rate limit exceeded")
|
||||||
|
result := ClassifyError(err, "my-provider", "my-model")
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("expected non-nil")
|
||||||
|
}
|
||||||
|
if result.Provider != "my-provider" {
|
||||||
|
t.Errorf("provider = %q, want my-provider", result.Provider)
|
||||||
|
}
|
||||||
|
if result.Model != "my-model" {
|
||||||
|
t.Errorf("model = %q, want my-model", result.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFailoverError_IsRetriable(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
reason FailoverReason
|
||||||
|
retriable bool
|
||||||
|
}{
|
||||||
|
{FailoverAuth, true},
|
||||||
|
{FailoverRateLimit, true},
|
||||||
|
{FailoverBilling, true},
|
||||||
|
{FailoverTimeout, true},
|
||||||
|
{FailoverOverloaded, true},
|
||||||
|
{FailoverFormat, false},
|
||||||
|
{FailoverUnknown, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
fe := &FailoverError{Reason: tt.reason}
|
||||||
|
if fe.IsRetriable() != tt.retriable {
|
||||||
|
t.Errorf("IsRetriable(%q) = %v, want %v", tt.reason, fe.IsRetriable(), tt.retriable)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFailoverError_ErrorString(t *testing.T) {
|
||||||
|
fe := &FailoverError{
|
||||||
|
Reason: FailoverRateLimit,
|
||||||
|
Provider: "openai",
|
||||||
|
Model: "gpt-4",
|
||||||
|
Status: 429,
|
||||||
|
Wrapped: errors.New("too many requests"),
|
||||||
|
}
|
||||||
|
s := fe.Error()
|
||||||
|
if s == "" {
|
||||||
|
t.Error("expected non-empty error string")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFailoverError_Unwrap(t *testing.T) {
|
||||||
|
inner := errors.New("inner error")
|
||||||
|
fe := &FailoverError{Reason: FailoverTimeout, Wrapped: inner}
|
||||||
|
if fe.Unwrap() != inner {
|
||||||
|
t.Error("Unwrap should return wrapped error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractHTTPStatus(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
msg string
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{"status: 429 rate limited", 429},
|
||||||
|
{"status 401 unauthorized", 401},
|
||||||
|
{"HTTP/1.1 502 Bad Gateway", 502},
|
||||||
|
{"no status code here", 0},
|
||||||
|
{"random number 12345", 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := extractHTTPStatus(tt.msg)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("extractHTTPStatus(%q) = %d, want %d", tt.msg, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsImageDimensionError(t *testing.T) {
|
||||||
|
if !IsImageDimensionError("image dimensions exceed max 4096x4096") {
|
||||||
|
t.Error("should match image dimensions exceed max")
|
||||||
|
}
|
||||||
|
if IsImageDimensionError("normal error message") {
|
||||||
|
t.Error("should not match normal error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsImageSizeError(t *testing.T) {
|
||||||
|
if !IsImageSizeError("image exceeds 20 mb") {
|
||||||
|
t.Error("should match image exceeds mb")
|
||||||
|
}
|
||||||
|
if IsImageSizeError("normal error message") {
|
||||||
|
t.Error("should not match normal error")
|
||||||
|
}
|
||||||
|
}
|
||||||
307
pkg/providers/factory.go
Normal file
307
pkg/providers/factory.go
Normal file
|
|
@ -0,0 +1,307 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
const defaultAnthropicAPIBase = "https://api.anthropic.com/v1"
|
||||||
|
|
||||||
|
var getCredential = auth.GetCredential
|
||||||
|
|
||||||
|
type providerType int
|
||||||
|
|
||||||
|
const (
|
||||||
|
providerTypeHTTPCompat providerType = iota
|
||||||
|
providerTypeClaudeAuth
|
||||||
|
providerTypeCodexAuth
|
||||||
|
providerTypeCodexCLIToken
|
||||||
|
providerTypeClaudeCLI
|
||||||
|
providerTypeCodexCLI
|
||||||
|
providerTypeGitHubCopilot
|
||||||
|
)
|
||||||
|
|
||||||
|
type providerSelection struct {
|
||||||
|
providerType providerType
|
||||||
|
apiKey string
|
||||||
|
apiBase string
|
||||||
|
proxy string
|
||||||
|
model string
|
||||||
|
workspace string
|
||||||
|
connectMode string
|
||||||
|
enableWebSearch bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
|
model := cfg.Agents.Defaults.Model
|
||||||
|
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
||||||
|
lowerModel := strings.ToLower(model)
|
||||||
|
|
||||||
|
sel := providerSelection{
|
||||||
|
providerType: providerTypeHTTPCompat,
|
||||||
|
model: model,
|
||||||
|
}
|
||||||
|
|
||||||
|
// First, prefer explicit provider configuration.
|
||||||
|
if providerName != "" {
|
||||||
|
switch providerName {
|
||||||
|
case "groq":
|
||||||
|
if cfg.Providers.Groq.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.Groq.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Groq.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Groq.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.groq.com/openai/v1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "openai", "gpt":
|
||||||
|
if cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != "" {
|
||||||
|
sel.enableWebSearch = cfg.Providers.OpenAI.WebSearch
|
||||||
|
if cfg.Providers.OpenAI.AuthMethod == "codex-cli" {
|
||||||
|
sel.providerType = providerTypeCodexCLIToken
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" {
|
||||||
|
sel.providerType = providerTypeCodexAuth
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
sel.apiKey = cfg.Providers.OpenAI.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.OpenAI.APIBase
|
||||||
|
sel.proxy = cfg.Providers.OpenAI.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.openai.com/v1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "anthropic", "claude":
|
||||||
|
if cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != "" {
|
||||||
|
if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" {
|
||||||
|
sel.apiBase = cfg.Providers.Anthropic.APIBase
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = defaultAnthropicAPIBase
|
||||||
|
}
|
||||||
|
sel.providerType = providerTypeClaudeAuth
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
sel.apiKey = cfg.Providers.Anthropic.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Anthropic.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Anthropic.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = defaultAnthropicAPIBase
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "openrouter":
|
||||||
|
if cfg.Providers.OpenRouter.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.OpenRouter.APIKey
|
||||||
|
sel.proxy = cfg.Providers.OpenRouter.Proxy
|
||||||
|
if cfg.Providers.OpenRouter.APIBase != "" {
|
||||||
|
sel.apiBase = cfg.Providers.OpenRouter.APIBase
|
||||||
|
} else {
|
||||||
|
sel.apiBase = "https://openrouter.ai/api/v1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "zhipu", "glm":
|
||||||
|
if cfg.Providers.Zhipu.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.Zhipu.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Zhipu.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Zhipu.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "gemini", "google":
|
||||||
|
if cfg.Providers.Gemini.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.Gemini.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Gemini.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Gemini.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "vllm":
|
||||||
|
if cfg.Providers.VLLM.APIBase != "" {
|
||||||
|
sel.apiKey = cfg.Providers.VLLM.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.VLLM.APIBase
|
||||||
|
sel.proxy = cfg.Providers.VLLM.Proxy
|
||||||
|
}
|
||||||
|
case "shengsuanyun":
|
||||||
|
if cfg.Providers.ShengSuanYun.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.ShengSuanYun.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.ShengSuanYun.APIBase
|
||||||
|
sel.proxy = cfg.Providers.ShengSuanYun.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://router.shengsuanyun.com/api/v1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "nvidia":
|
||||||
|
if cfg.Providers.Nvidia.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.Nvidia.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Nvidia.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Nvidia.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://integrate.api.nvidia.com/v1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "claude-cli", "claude-code", "claudecode":
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
if workspace == "" {
|
||||||
|
workspace = "."
|
||||||
|
}
|
||||||
|
sel.providerType = providerTypeClaudeCLI
|
||||||
|
sel.workspace = workspace
|
||||||
|
return sel, nil
|
||||||
|
case "codex-cli", "codex-code":
|
||||||
|
workspace := cfg.WorkspacePath()
|
||||||
|
if workspace == "" {
|
||||||
|
workspace = "."
|
||||||
|
}
|
||||||
|
sel.providerType = providerTypeCodexCLI
|
||||||
|
sel.workspace = workspace
|
||||||
|
return sel, nil
|
||||||
|
case "deepseek":
|
||||||
|
if cfg.Providers.DeepSeek.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.DeepSeek.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.DeepSeek.APIBase
|
||||||
|
sel.proxy = cfg.Providers.DeepSeek.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.deepseek.com/v1"
|
||||||
|
}
|
||||||
|
if model != "deepseek-chat" && model != "deepseek-reasoner" {
|
||||||
|
sel.model = "deepseek-chat"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case "github_copilot", "copilot":
|
||||||
|
sel.providerType = providerTypeGitHubCopilot
|
||||||
|
if cfg.Providers.GitHubCopilot.APIBase != "" {
|
||||||
|
sel.apiBase = cfg.Providers.GitHubCopilot.APIBase
|
||||||
|
} else {
|
||||||
|
sel.apiBase = "localhost:4321"
|
||||||
|
}
|
||||||
|
sel.connectMode = cfg.Providers.GitHubCopilot.ConnectMode
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: infer provider from model and configured keys.
|
||||||
|
if sel.apiKey == "" && sel.apiBase == "" {
|
||||||
|
switch {
|
||||||
|
case (strings.Contains(lowerModel, "kimi") || strings.Contains(lowerModel, "moonshot") || strings.HasPrefix(model, "moonshot/")) && cfg.Providers.Moonshot.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Moonshot.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Moonshot.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Moonshot.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.moonshot.cn/v1"
|
||||||
|
}
|
||||||
|
case strings.HasPrefix(model, "openrouter/") ||
|
||||||
|
strings.HasPrefix(model, "anthropic/") ||
|
||||||
|
strings.HasPrefix(model, "openai/") ||
|
||||||
|
strings.HasPrefix(model, "meta-llama/") ||
|
||||||
|
strings.HasPrefix(model, "deepseek/") ||
|
||||||
|
strings.HasPrefix(model, "google/"):
|
||||||
|
sel.apiKey = cfg.Providers.OpenRouter.APIKey
|
||||||
|
sel.proxy = cfg.Providers.OpenRouter.Proxy
|
||||||
|
if cfg.Providers.OpenRouter.APIBase != "" {
|
||||||
|
sel.apiBase = cfg.Providers.OpenRouter.APIBase
|
||||||
|
} else {
|
||||||
|
sel.apiBase = "https://openrouter.ai/api/v1"
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "claude") || strings.HasPrefix(model, "anthropic/")) &&
|
||||||
|
(cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != ""):
|
||||||
|
if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" {
|
||||||
|
sel.apiBase = cfg.Providers.Anthropic.APIBase
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = defaultAnthropicAPIBase
|
||||||
|
}
|
||||||
|
sel.providerType = providerTypeClaudeAuth
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
sel.apiKey = cfg.Providers.Anthropic.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Anthropic.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Anthropic.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = defaultAnthropicAPIBase
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "gpt") || strings.HasPrefix(model, "openai/")) &&
|
||||||
|
(cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != ""):
|
||||||
|
sel.enableWebSearch = cfg.Providers.OpenAI.WebSearch
|
||||||
|
if cfg.Providers.OpenAI.AuthMethod == "codex-cli" {
|
||||||
|
sel.providerType = providerTypeCodexCLIToken
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" {
|
||||||
|
sel.providerType = providerTypeCodexAuth
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
|
sel.apiKey = cfg.Providers.OpenAI.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.OpenAI.APIBase
|
||||||
|
sel.proxy = cfg.Providers.OpenAI.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.openai.com/v1"
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "gemini") || strings.HasPrefix(model, "google/")) && cfg.Providers.Gemini.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Gemini.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Gemini.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Gemini.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "zhipu") || strings.Contains(lowerModel, "zai")) && cfg.Providers.Zhipu.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Zhipu.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Zhipu.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Zhipu.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "groq") || strings.HasPrefix(model, "groq/")) && cfg.Providers.Groq.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Groq.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Groq.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Groq.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.groq.com/openai/v1"
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "nvidia") || strings.HasPrefix(model, "nvidia/")) && cfg.Providers.Nvidia.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Nvidia.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Nvidia.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Nvidia.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://integrate.api.nvidia.com/v1"
|
||||||
|
}
|
||||||
|
case (strings.Contains(lowerModel, "ollama") || strings.HasPrefix(model, "ollama/")) && cfg.Providers.Ollama.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Ollama.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Ollama.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Ollama.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "http://localhost:11434/v1"
|
||||||
|
}
|
||||||
|
case cfg.Providers.VLLM.APIBase != "":
|
||||||
|
sel.apiKey = cfg.Providers.VLLM.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.VLLM.APIBase
|
||||||
|
sel.proxy = cfg.Providers.VLLM.Proxy
|
||||||
|
default:
|
||||||
|
if cfg.Providers.OpenRouter.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.OpenRouter.APIKey
|
||||||
|
sel.proxy = cfg.Providers.OpenRouter.Proxy
|
||||||
|
if cfg.Providers.OpenRouter.APIBase != "" {
|
||||||
|
sel.apiBase = cfg.Providers.OpenRouter.APIBase
|
||||||
|
} else {
|
||||||
|
sel.apiBase = "https://openrouter.ai/api/v1"
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
return providerSelection{}, fmt.Errorf("no API key configured for model: %s", model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if sel.providerType == providerTypeHTTPCompat {
|
||||||
|
if sel.apiKey == "" && !strings.HasPrefix(model, "bedrock/") {
|
||||||
|
return providerSelection{}, fmt.Errorf("no API key configured for provider (model: %s)", model)
|
||||||
|
}
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
return providerSelection{}, fmt.Errorf("no API base configured for provider (model: %s)", model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sel, nil
|
||||||
|
}
|
||||||
192
pkg/providers/factory_provider.go
Normal file
192
pkg/providers/factory_provider.go
Normal file
|
|
@ -0,0 +1,192 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
||||||
|
func createClaudeAuthProvider() (LLMProvider, error) {
|
||||||
|
cred, err := getCredential("anthropic")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
return nil, fmt.Errorf("no credentials for anthropic. Run: picoclaw auth login --provider anthropic")
|
||||||
|
}
|
||||||
|
return NewClaudeProviderWithTokenSource(cred.AccessToken, createClaudeTokenSource()), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// createCodexAuthProvider creates a Codex provider using OAuth credentials from auth store.
|
||||||
|
func createCodexAuthProvider() (LLMProvider, error) {
|
||||||
|
cred, err := getCredential("openai")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
return nil, fmt.Errorf("no credentials for openai. Run: picoclaw auth login --provider openai")
|
||||||
|
}
|
||||||
|
return NewCodexProviderWithTokenSource(cred.AccessToken, cred.AccountID, createCodexTokenSource()), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtractProtocol extracts the protocol prefix and model identifier from a model string.
|
||||||
|
// If no prefix is specified, it defaults to "openai".
|
||||||
|
// Examples:
|
||||||
|
// - "openai/gpt-4o" -> ("openai", "gpt-4o")
|
||||||
|
// - "anthropic/claude-sonnet-4.6" -> ("anthropic", "claude-sonnet-4.6")
|
||||||
|
// - "gpt-4o" -> ("openai", "gpt-4o") // default protocol
|
||||||
|
func ExtractProtocol(model string) (protocol, modelID string) {
|
||||||
|
model = strings.TrimSpace(model)
|
||||||
|
protocol, modelID, found := strings.Cut(model, "/")
|
||||||
|
if !found {
|
||||||
|
return "openai", model
|
||||||
|
}
|
||||||
|
return protocol, modelID
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateProviderFromConfig creates a provider based on the ModelConfig.
|
||||||
|
// It uses the protocol prefix in the Model field to determine which provider to create.
|
||||||
|
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot
|
||||||
|
// Returns the provider, the model ID (without protocol prefix), and any error.
|
||||||
|
func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) {
|
||||||
|
if cfg == nil {
|
||||||
|
return nil, "", fmt.Errorf("config is nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Model == "" {
|
||||||
|
return nil, "", fmt.Errorf("model is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
protocol, modelID := ExtractProtocol(cfg.Model)
|
||||||
|
|
||||||
|
switch protocol {
|
||||||
|
case "openai":
|
||||||
|
// OpenAI with OAuth/token auth (Codex-style)
|
||||||
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
||||||
|
provider, err := createCodexAuthProvider()
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
return provider, modelID, nil
|
||||||
|
}
|
||||||
|
// OpenAI with API key
|
||||||
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
|
}
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = getDefaultAPIBase(protocol)
|
||||||
|
}
|
||||||
|
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil
|
||||||
|
|
||||||
|
case "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||||
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
|
"volcengine", "vllm", "qwen":
|
||||||
|
// All other OpenAI-compatible HTTP providers
|
||||||
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
|
}
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = getDefaultAPIBase(protocol)
|
||||||
|
}
|
||||||
|
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil
|
||||||
|
|
||||||
|
case "anthropic":
|
||||||
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
||||||
|
// Use OAuth credentials from auth store
|
||||||
|
provider, err := createClaudeAuthProvider()
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
return provider, modelID, nil
|
||||||
|
}
|
||||||
|
// Use API key with HTTP API
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = "https://api.anthropic.com/v1"
|
||||||
|
}
|
||||||
|
if cfg.APIKey == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model)
|
||||||
|
}
|
||||||
|
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil
|
||||||
|
|
||||||
|
case "antigravity":
|
||||||
|
return NewAntigravityProvider(), modelID, nil
|
||||||
|
|
||||||
|
case "claude-cli", "claudecli":
|
||||||
|
workspace := cfg.Workspace
|
||||||
|
if workspace == "" {
|
||||||
|
workspace = "."
|
||||||
|
}
|
||||||
|
return NewClaudeCliProvider(workspace), modelID, nil
|
||||||
|
|
||||||
|
case "codex-cli", "codexcli":
|
||||||
|
workspace := cfg.Workspace
|
||||||
|
if workspace == "" {
|
||||||
|
workspace = "."
|
||||||
|
}
|
||||||
|
return NewCodexCliProvider(workspace), modelID, nil
|
||||||
|
|
||||||
|
case "github-copilot", "copilot":
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = "localhost:4321"
|
||||||
|
}
|
||||||
|
connectMode := cfg.ConnectMode
|
||||||
|
if connectMode == "" {
|
||||||
|
connectMode = "grpc"
|
||||||
|
}
|
||||||
|
provider, err := NewGitHubCopilotProvider(apiBase, connectMode, modelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
return provider, modelID, nil
|
||||||
|
|
||||||
|
default:
|
||||||
|
return nil, "", fmt.Errorf("unknown protocol %q in model %q", protocol, cfg.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
||||||
|
func getDefaultAPIBase(protocol string) string {
|
||||||
|
switch protocol {
|
||||||
|
case "openai":
|
||||||
|
return "https://api.openai.com/v1"
|
||||||
|
case "openrouter":
|
||||||
|
return "https://openrouter.ai/api/v1"
|
||||||
|
case "groq":
|
||||||
|
return "https://api.groq.com/openai/v1"
|
||||||
|
case "zhipu":
|
||||||
|
return "https://open.bigmodel.cn/api/paas/v4"
|
||||||
|
case "gemini":
|
||||||
|
return "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
case "nvidia":
|
||||||
|
return "https://integrate.api.nvidia.com/v1"
|
||||||
|
case "ollama":
|
||||||
|
return "http://localhost:11434/v1"
|
||||||
|
case "moonshot":
|
||||||
|
return "https://api.moonshot.cn/v1"
|
||||||
|
case "shengsuanyun":
|
||||||
|
return "https://router.shengsuanyun.com/api/v1"
|
||||||
|
case "deepseek":
|
||||||
|
return "https://api.deepseek.com/v1"
|
||||||
|
case "cerebras":
|
||||||
|
return "https://api.cerebras.ai/v1"
|
||||||
|
case "volcengine":
|
||||||
|
return "https://ark.cn-beijing.volces.com/api/v3"
|
||||||
|
case "qwen":
|
||||||
|
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
|
case "vllm":
|
||||||
|
return "http://localhost:8000/v1"
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
249
pkg/providers/factory_provider_test.go
Normal file
249
pkg/providers/factory_provider_test.go
Normal file
|
|
@ -0,0 +1,249 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestExtractProtocol(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
model string
|
||||||
|
wantProtocol string
|
||||||
|
wantModelID string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "openai with prefix",
|
||||||
|
model: "openai/gpt-4o",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "gpt-4o",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "anthropic with prefix",
|
||||||
|
model: "anthropic/claude-sonnet-4.6",
|
||||||
|
wantProtocol: "anthropic",
|
||||||
|
wantModelID: "claude-sonnet-4.6",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no prefix - defaults to openai",
|
||||||
|
model: "gpt-4o",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "gpt-4o",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "groq with prefix",
|
||||||
|
model: "groq/llama-3.1-70b",
|
||||||
|
wantProtocol: "groq",
|
||||||
|
wantModelID: "llama-3.1-70b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty string",
|
||||||
|
model: "",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with whitespace",
|
||||||
|
model: " openai/gpt-4 ",
|
||||||
|
wantProtocol: "openai",
|
||||||
|
wantModelID: "gpt-4",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple slashes",
|
||||||
|
model: "nvidia/meta/llama-3.1-8b",
|
||||||
|
wantProtocol: "nvidia",
|
||||||
|
wantModelID: "meta/llama-3.1-8b",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
protocol, modelID := ExtractProtocol(tt.model)
|
||||||
|
if protocol != tt.wantProtocol {
|
||||||
|
t.Errorf("ExtractProtocol(%q) protocol = %q, want %q", tt.model, protocol, tt.wantProtocol)
|
||||||
|
}
|
||||||
|
if modelID != tt.wantModelID {
|
||||||
|
t.Errorf("ExtractProtocol(%q) modelID = %q, want %q", tt.model, modelID, tt.wantModelID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_OpenAI(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-openai",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
APIKey: "test-key",
|
||||||
|
APIBase: "https://api.example.com/v1",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "gpt-4o" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "gpt-4o")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
protocol string
|
||||||
|
}{
|
||||||
|
{"openai", "openai"},
|
||||||
|
{"groq", "groq"},
|
||||||
|
{"openrouter", "openrouter"},
|
||||||
|
{"cerebras", "cerebras"},
|
||||||
|
{"qwen", "qwen"},
|
||||||
|
{"vllm", "vllm"},
|
||||||
|
{"deepseek", "deepseek"},
|
||||||
|
{"ollama", "ollama"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-" + tt.protocol,
|
||||||
|
Model: tt.protocol + "/test-model",
|
||||||
|
APIKey: "test-key",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify we got an HTTPProvider for all these protocols
|
||||||
|
if _, ok := provider.(*HTTPProvider); !ok {
|
||||||
|
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_Anthropic(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-anthropic",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
APIKey: "test-key",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "claude-sonnet-4.6" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_Antigravity(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-antigravity",
|
||||||
|
Model: "antigravity/gemini-2.0-flash",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "gemini-2.0-flash" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "gemini-2.0-flash")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_ClaudeCLI(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-claude-cli",
|
||||||
|
Model: "claude-cli/claude-sonnet-4.6",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "claude-sonnet-4.6" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "claude-sonnet-4.6")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_CodexCLI(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-codex-cli",
|
||||||
|
Model: "codex-cli/codex",
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != "codex" {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, "codex")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_MissingAPIKey(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-no-key",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() expected error for missing API key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_UnknownProtocol(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-unknown",
|
||||||
|
Model: "unknown-protocol/model",
|
||||||
|
APIKey: "test-key",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() expected error for unknown protocol")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_NilConfig(t *testing.T) {
|
||||||
|
_, _, err := CreateProviderFromConfig(nil)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig(nil) expected error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_EmptyModel(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: "test-empty",
|
||||||
|
Model: "",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() expected error for empty model")
|
||||||
|
}
|
||||||
|
}
|
||||||
299
pkg/providers/factory_test.go
Normal file
299
pkg/providers/factory_test.go
Normal file
|
|
@ -0,0 +1,299 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestResolveProviderSelection(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
setup func(*config.Config)
|
||||||
|
wantType providerType
|
||||||
|
wantAPIBase string
|
||||||
|
wantProxy string
|
||||||
|
wantErrSubstr string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "explicit claude-cli provider routes to cli provider type",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "claude-cli"
|
||||||
|
cfg.Agents.Defaults.Workspace = "/tmp/ws"
|
||||||
|
},
|
||||||
|
wantType: providerTypeClaudeCLI,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "explicit copilot provider routes to github copilot type",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "copilot"
|
||||||
|
},
|
||||||
|
wantType: providerTypeGitHubCopilot,
|
||||||
|
wantAPIBase: "localhost:4321",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "explicit deepseek provider uses deepseek defaults",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "deepseek"
|
||||||
|
cfg.Agents.Defaults.Model = "deepseek/deepseek-chat"
|
||||||
|
cfg.Providers.DeepSeek.APIKey = "deepseek-key"
|
||||||
|
cfg.Providers.DeepSeek.Proxy = "http://127.0.0.1:7890"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://api.deepseek.com/v1",
|
||||||
|
wantProxy: "http://127.0.0.1:7890",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "explicit shengsuanyun provider uses defaults",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "shengsuanyun"
|
||||||
|
cfg.Providers.ShengSuanYun.APIKey = "ssy-key"
|
||||||
|
cfg.Providers.ShengSuanYun.Proxy = "http://127.0.0.1:7890"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://router.shengsuanyun.com/api/v1",
|
||||||
|
wantProxy: "http://127.0.0.1:7890",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "explicit nvidia provider uses defaults",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "nvidia"
|
||||||
|
cfg.Providers.Nvidia.APIKey = "nvapi-test"
|
||||||
|
cfg.Providers.Nvidia.Proxy = "http://127.0.0.1:7890"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://integrate.api.nvidia.com/v1",
|
||||||
|
wantProxy: "http://127.0.0.1:7890",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "openrouter model uses openrouter defaults",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "openrouter/auto"
|
||||||
|
cfg.Providers.OpenRouter.APIKey = "sk-or-test"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://openrouter.ai/api/v1",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "anthropic oauth routes to claude auth provider",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "claude-sonnet-4.6"
|
||||||
|
cfg.Providers.Anthropic.AuthMethod = "oauth"
|
||||||
|
},
|
||||||
|
wantType: providerTypeClaudeAuth,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "openai oauth routes to codex auth provider",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "gpt-4o"
|
||||||
|
cfg.Providers.OpenAI.AuthMethod = "oauth"
|
||||||
|
},
|
||||||
|
wantType: providerTypeCodexAuth,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "openai codex-cli auth routes to codex cli token provider",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "gpt-4o"
|
||||||
|
cfg.Providers.OpenAI.AuthMethod = "codex-cli"
|
||||||
|
},
|
||||||
|
wantType: providerTypeCodexCLIToken,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "explicit codex-code provider routes to codex cli provider type",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Provider = "codex-code"
|
||||||
|
cfg.Agents.Defaults.Workspace = "/tmp/ws"
|
||||||
|
},
|
||||||
|
wantType: providerTypeCodexCLI,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "zhipu model uses zhipu base default",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "glm-4.7"
|
||||||
|
cfg.Providers.Zhipu.APIKey = "zhipu-key"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://open.bigmodel.cn/api/paas/v4",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "groq model uses groq base default",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "groq/llama-3.3-70b"
|
||||||
|
cfg.Providers.Groq.APIKey = "gsk-key"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://api.groq.com/openai/v1",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ollama model uses ollama base default",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "ollama/qwen2.5:14b"
|
||||||
|
cfg.Providers.Ollama.APIKey = "ollama-key"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "http://localhost:11434/v1",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "moonshot model keeps proxy and default base",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "moonshot/kimi-k2.5"
|
||||||
|
cfg.Providers.Moonshot.APIKey = "moonshot-key"
|
||||||
|
cfg.Providers.Moonshot.Proxy = "http://127.0.0.1:7890"
|
||||||
|
},
|
||||||
|
wantType: providerTypeHTTPCompat,
|
||||||
|
wantAPIBase: "https://api.moonshot.cn/v1",
|
||||||
|
wantProxy: "http://127.0.0.1:7890",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing keys returns model config error",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "custom-model"
|
||||||
|
},
|
||||||
|
wantErrSubstr: "no API key configured for model",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "openrouter prefix without key returns provider key error",
|
||||||
|
setup: func(cfg *config.Config) {
|
||||||
|
cfg.Agents.Defaults.Model = "openrouter/auto"
|
||||||
|
},
|
||||||
|
wantErrSubstr: "no API key configured for provider",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
tt.setup(cfg)
|
||||||
|
|
||||||
|
got, err := resolveProviderSelection(cfg)
|
||||||
|
if tt.wantErrSubstr != "" {
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected error containing %q, got nil", tt.wantErrSubstr)
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), tt.wantErrSubstr) {
|
||||||
|
t.Fatalf("error = %q, want substring %q", err.Error(), tt.wantErrSubstr)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolveProviderSelection() error = %v", err)
|
||||||
|
}
|
||||||
|
if got.providerType != tt.wantType {
|
||||||
|
t.Fatalf("providerType = %v, want %v", got.providerType, tt.wantType)
|
||||||
|
}
|
||||||
|
if tt.wantAPIBase != "" && got.apiBase != tt.wantAPIBase {
|
||||||
|
t.Fatalf("apiBase = %q, want %q", got.apiBase, tt.wantAPIBase)
|
||||||
|
}
|
||||||
|
if tt.wantProxy != "" && got.proxy != tt.wantProxy {
|
||||||
|
t.Fatalf("proxy = %q, want %q", got.proxy, tt.wantProxy)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderReturnsHTTPProviderForOpenRouter(t *testing.T) {
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Model = "test-openrouter"
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-openrouter",
|
||||||
|
Model: "openrouter/auto",
|
||||||
|
APIKey: "sk-or-test",
|
||||||
|
APIBase: "https://openrouter.ai/api/v1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, _, err := CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := provider.(*HTTPProvider); !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *HTTPProvider", provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderReturnsCodexCliProviderForCodexCode(t *testing.T) {
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Model = "test-codex"
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-codex",
|
||||||
|
Model: "codex-cli/codex-model",
|
||||||
|
Workspace: "/tmp/workspace",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, _, err := CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := provider.(*CodexCliProvider); !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *CodexCliProvider", provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderReturnsClaudeCliProviderForClaudeCli(t *testing.T) {
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Model = "test-claude-cli"
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-claude-cli",
|
||||||
|
Model: "claude-cli/claude-sonnet",
|
||||||
|
Workspace: "/tmp/workspace",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, _, err := CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := provider.(*ClaudeCliProvider); !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *ClaudeCliProvider", provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderReturnsClaudeProviderForAnthropicOAuth(t *testing.T) {
|
||||||
|
originalGetCredential := getCredential
|
||||||
|
t.Cleanup(func() { getCredential = originalGetCredential })
|
||||||
|
|
||||||
|
getCredential = func(provider string) (*auth.AuthCredential, error) {
|
||||||
|
if provider != "anthropic" {
|
||||||
|
t.Fatalf("provider = %q, want anthropic", provider)
|
||||||
|
}
|
||||||
|
return &auth.AuthCredential{
|
||||||
|
AccessToken: "anthropic-token",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Model = "test-claude-oauth"
|
||||||
|
cfg.ModelList = []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "test-claude-oauth",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, _, err := CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := provider.(*ClaudeProvider); !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *ClaudeProvider", provider)
|
||||||
|
}
|
||||||
|
// TODO: Test custom APIBase when createClaudeAuthProvider supports it
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderReturnsCodexProviderForOpenAIOAuth(t *testing.T) {
|
||||||
|
// TODO: This test requires openai protocol to support auth_method: "oauth"
|
||||||
|
// which is not yet implemented in the new factory_provider.go
|
||||||
|
t.Skip("OpenAI OAuth via model_list not yet implemented")
|
||||||
|
}
|
||||||
283
pkg/providers/fallback.go
Normal file
283
pkg/providers/fallback.go
Normal file
|
|
@ -0,0 +1,283 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FallbackChain orchestrates model fallback across multiple candidates.
|
||||||
|
type FallbackChain struct {
|
||||||
|
cooldown *CooldownTracker
|
||||||
|
}
|
||||||
|
|
||||||
|
// FallbackCandidate represents one model/provider to try.
|
||||||
|
type FallbackCandidate struct {
|
||||||
|
Provider string
|
||||||
|
Model string
|
||||||
|
}
|
||||||
|
|
||||||
|
// FallbackResult contains the successful response and metadata about all attempts.
|
||||||
|
type FallbackResult struct {
|
||||||
|
Response *LLMResponse
|
||||||
|
Provider string
|
||||||
|
Model string
|
||||||
|
Attempts []FallbackAttempt
|
||||||
|
}
|
||||||
|
|
||||||
|
// FallbackAttempt records one attempt in the fallback chain.
|
||||||
|
type FallbackAttempt struct {
|
||||||
|
Provider string
|
||||||
|
Model string
|
||||||
|
Error error
|
||||||
|
Reason FailoverReason
|
||||||
|
Duration time.Duration
|
||||||
|
Skipped bool // true if skipped due to cooldown
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewFallbackChain creates a new fallback chain with the given cooldown tracker.
|
||||||
|
func NewFallbackChain(cooldown *CooldownTracker) *FallbackChain {
|
||||||
|
return &FallbackChain{cooldown: cooldown}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveCandidates parses model config into a deduplicated candidate list.
|
||||||
|
func ResolveCandidates(cfg ModelConfig, defaultProvider string) []FallbackCandidate {
|
||||||
|
seen := make(map[string]bool)
|
||||||
|
var candidates []FallbackCandidate
|
||||||
|
|
||||||
|
addCandidate := func(raw string) {
|
||||||
|
ref := ParseModelRef(raw, defaultProvider)
|
||||||
|
if ref == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
key := ModelKey(ref.Provider, ref.Model)
|
||||||
|
if seen[key] {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
seen[key] = true
|
||||||
|
candidates = append(candidates, FallbackCandidate{
|
||||||
|
Provider: ref.Provider,
|
||||||
|
Model: ref.Model,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Primary first.
|
||||||
|
addCandidate(cfg.Primary)
|
||||||
|
|
||||||
|
// Then fallbacks.
|
||||||
|
for _, fb := range cfg.Fallbacks {
|
||||||
|
addCandidate(fb)
|
||||||
|
}
|
||||||
|
|
||||||
|
return candidates
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute runs the fallback chain for text/chat requests.
|
||||||
|
// It tries each candidate in order, respecting cooldowns and error classification.
|
||||||
|
//
|
||||||
|
// Behavior:
|
||||||
|
// - Candidates in cooldown are skipped (logged as skipped attempt).
|
||||||
|
// - context.Canceled aborts immediately (user abort, no fallback).
|
||||||
|
// - Non-retriable errors (format) abort immediately.
|
||||||
|
// - Retriable errors trigger fallback to next candidate.
|
||||||
|
// - Success marks provider as good (resets cooldown).
|
||||||
|
// - If all fail, returns aggregate error with all attempts.
|
||||||
|
func (fc *FallbackChain) Execute(
|
||||||
|
ctx context.Context,
|
||||||
|
candidates []FallbackCandidate,
|
||||||
|
run func(ctx context.Context, provider, model string) (*LLMResponse, error),
|
||||||
|
) (*FallbackResult, error) {
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return nil, fmt.Errorf("fallback: no candidates configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
result := &FallbackResult{
|
||||||
|
Attempts: make([]FallbackAttempt, 0, len(candidates)),
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, candidate := range candidates {
|
||||||
|
// Check context before each attempt.
|
||||||
|
if ctx.Err() == context.Canceled {
|
||||||
|
return nil, context.Canceled
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check cooldown.
|
||||||
|
if !fc.cooldown.IsAvailable(candidate.Provider) {
|
||||||
|
remaining := fc.cooldown.CooldownRemaining(candidate.Provider)
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Skipped: true,
|
||||||
|
Reason: FailoverRateLimit,
|
||||||
|
Error: fmt.Errorf("provider %s in cooldown (%s remaining)", candidate.Provider, remaining.Round(time.Second)),
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute the run function.
|
||||||
|
start := time.Now()
|
||||||
|
resp, err := run(ctx, candidate.Provider, candidate.Model)
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
// Success.
|
||||||
|
fc.cooldown.MarkSuccess(candidate.Provider)
|
||||||
|
result.Response = resp
|
||||||
|
result.Provider = candidate.Provider
|
||||||
|
result.Model = candidate.Model
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Context cancellation: abort immediately, no fallback.
|
||||||
|
if ctx.Err() == context.Canceled {
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: err,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
return nil, context.Canceled
|
||||||
|
}
|
||||||
|
|
||||||
|
// Classify the error.
|
||||||
|
failErr := ClassifyError(err, candidate.Provider, candidate.Model)
|
||||||
|
|
||||||
|
if failErr == nil {
|
||||||
|
// Unclassifiable error: do not fallback, return immediately.
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: err,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
return nil, fmt.Errorf("fallback: unclassified error from %s/%s: %w",
|
||||||
|
candidate.Provider, candidate.Model, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Non-retriable error: abort immediately.
|
||||||
|
if !failErr.IsRetriable() {
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: failErr,
|
||||||
|
Reason: failErr.Reason,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
return nil, failErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Retriable error: mark failure and continue to next candidate.
|
||||||
|
fc.cooldown.MarkFailure(candidate.Provider, failErr.Reason)
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: failErr,
|
||||||
|
Reason: failErr.Reason,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
|
||||||
|
// If this was the last candidate, return aggregate error.
|
||||||
|
if i == len(candidates)-1 {
|
||||||
|
return nil, &FallbackExhaustedError{Attempts: result.Attempts}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// All candidates were skipped (all in cooldown).
|
||||||
|
return nil, &FallbackExhaustedError{Attempts: result.Attempts}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExecuteImage runs the fallback chain for image/vision requests.
|
||||||
|
// Simpler than Execute: no cooldown checks (image endpoints have different rate limits).
|
||||||
|
// Image dimension/size errors abort immediately (non-retriable).
|
||||||
|
func (fc *FallbackChain) ExecuteImage(
|
||||||
|
ctx context.Context,
|
||||||
|
candidates []FallbackCandidate,
|
||||||
|
run func(ctx context.Context, provider, model string) (*LLMResponse, error),
|
||||||
|
) (*FallbackResult, error) {
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return nil, fmt.Errorf("image fallback: no candidates configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
result := &FallbackResult{
|
||||||
|
Attempts: make([]FallbackAttempt, 0, len(candidates)),
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, candidate := range candidates {
|
||||||
|
if ctx.Err() == context.Canceled {
|
||||||
|
return nil, context.Canceled
|
||||||
|
}
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
resp, err := run(ctx, candidate.Provider, candidate.Model)
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
result.Response = resp
|
||||||
|
result.Provider = candidate.Provider
|
||||||
|
result.Model = candidate.Model
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if ctx.Err() == context.Canceled {
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: err,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
return nil, context.Canceled
|
||||||
|
}
|
||||||
|
|
||||||
|
// Image dimension/size errors are non-retriable.
|
||||||
|
errMsg := strings.ToLower(err.Error())
|
||||||
|
if IsImageDimensionError(errMsg) || IsImageSizeError(errMsg) {
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: err,
|
||||||
|
Reason: FailoverFormat,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
return nil, &FailoverError{
|
||||||
|
Reason: FailoverFormat,
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Wrapped: err,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Any other error: record and try next.
|
||||||
|
result.Attempts = append(result.Attempts, FallbackAttempt{
|
||||||
|
Provider: candidate.Provider,
|
||||||
|
Model: candidate.Model,
|
||||||
|
Error: err,
|
||||||
|
Duration: elapsed,
|
||||||
|
})
|
||||||
|
|
||||||
|
if i == len(candidates)-1 {
|
||||||
|
return nil, &FallbackExhaustedError{Attempts: result.Attempts}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, &FallbackExhaustedError{Attempts: result.Attempts}
|
||||||
|
}
|
||||||
|
|
||||||
|
// FallbackExhaustedError indicates all fallback candidates were tried and failed.
|
||||||
|
type FallbackExhaustedError struct {
|
||||||
|
Attempts []FallbackAttempt
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *FallbackExhaustedError) Error() string {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString(fmt.Sprintf("fallback: all %d candidates failed:", len(e.Attempts)))
|
||||||
|
for i, a := range e.Attempts {
|
||||||
|
if a.Skipped {
|
||||||
|
sb.WriteString(fmt.Sprintf("\n [%d] %s/%s: skipped (cooldown)", i+1, a.Provider, a.Model))
|
||||||
|
} else {
|
||||||
|
sb.WriteString(fmt.Sprintf("\n [%d] %s/%s: %v (reason=%s, %s)",
|
||||||
|
i+1, a.Provider, a.Model, a.Error, a.Reason, a.Duration.Round(time.Millisecond)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
473
pkg/providers/fallback_test.go
Normal file
473
pkg/providers/fallback_test.go
Normal file
|
|
@ -0,0 +1,473 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func makeCandidate(provider, model string) FallbackCandidate {
|
||||||
|
return FallbackCandidate{Provider: provider, Model: model}
|
||||||
|
}
|
||||||
|
|
||||||
|
func successRun(content string) func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
return func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
return &LLMResponse{Content: content, FinishReason: "stop"}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func failRun(err error) func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
return func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_SingleCandidate_Success(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{makeCandidate("openai", "gpt-4")}
|
||||||
|
result, err := fc.Execute(context.Background(), candidates, successRun("hello"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Response.Content != "hello" {
|
||||||
|
t.Errorf("content = %q, want hello", result.Response.Content)
|
||||||
|
}
|
||||||
|
if result.Provider != "openai" || result.Model != "gpt-4" {
|
||||||
|
t.Errorf("provider/model = %s/%s, want openai/gpt-4", result.Provider, result.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_SecondCandidateSuccess(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude-opus"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
if attempt == 1 {
|
||||||
|
return nil, errors.New("rate limit exceeded")
|
||||||
|
}
|
||||||
|
return &LLMResponse{Content: "from claude", FinishReason: "stop"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := fc.Execute(context.Background(), candidates, run)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Provider != "anthropic" {
|
||||||
|
t.Errorf("provider = %q, want anthropic", result.Provider)
|
||||||
|
}
|
||||||
|
if result.Response.Content != "from claude" {
|
||||||
|
t.Errorf("content = %q, want 'from claude'", result.Response.Content)
|
||||||
|
}
|
||||||
|
if len(result.Attempts) != 1 {
|
||||||
|
t.Errorf("attempts = %d, want 1 (failed attempt recorded)", len(result.Attempts))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_AllFail(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
makeCandidate("groq", "llama"),
|
||||||
|
}
|
||||||
|
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
return nil, errors.New("rate limit exceeded")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.Execute(context.Background(), candidates, run)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error when all candidates fail")
|
||||||
|
}
|
||||||
|
var exhausted *FallbackExhaustedError
|
||||||
|
if !errors.As(err, &exhausted) {
|
||||||
|
t.Errorf("expected FallbackExhaustedError, got %T: %v", err, err)
|
||||||
|
}
|
||||||
|
if len(exhausted.Attempts) != 3 {
|
||||||
|
t.Errorf("attempts = %d, want 3", len(exhausted.Attempts))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_ContextCanceled(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
if attempt == 1 {
|
||||||
|
cancel() // cancel context
|
||||||
|
return nil, context.Canceled
|
||||||
|
}
|
||||||
|
t.Error("should not reach second candidate after cancel")
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.Execute(ctx, candidates, run)
|
||||||
|
if err != context.Canceled {
|
||||||
|
t.Errorf("expected context.Canceled, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_NonRetriableError(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
return nil, errors.New("string should match pattern")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.Execute(context.Background(), candidates, run)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error for non-retriable")
|
||||||
|
}
|
||||||
|
var fe *FailoverError
|
||||||
|
if !errors.As(err, &fe) {
|
||||||
|
t.Fatalf("expected FailoverError, got %T", err)
|
||||||
|
}
|
||||||
|
if fe.Reason != FailoverFormat {
|
||||||
|
t.Errorf("reason = %q, want format", fe.Reason)
|
||||||
|
}
|
||||||
|
if attempt != 1 {
|
||||||
|
t.Errorf("attempt = %d, want 1 (non-retriable should not try next)", attempt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_CooldownSkip(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
ct, _ := newTestTracker(now)
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
// Put openai in cooldown
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
if provider == "openai" {
|
||||||
|
t.Error("should not call openai (in cooldown)")
|
||||||
|
}
|
||||||
|
return &LLMResponse{Content: "claude response", FinishReason: "stop"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := fc.Execute(context.Background(), candidates, run)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Provider != "anthropic" {
|
||||||
|
t.Errorf("provider = %q, want anthropic", result.Provider)
|
||||||
|
}
|
||||||
|
// Should have 1 skipped attempt
|
||||||
|
skipped := 0
|
||||||
|
for _, a := range result.Attempts {
|
||||||
|
if a.Skipped {
|
||||||
|
skipped++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if skipped != 1 {
|
||||||
|
t.Errorf("skipped = %d, want 1", skipped)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_AllInCooldown(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
// Put all providers in cooldown
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit)
|
||||||
|
ct.MarkFailure("anthropic", FailoverBilling)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.Execute(context.Background(), candidates,
|
||||||
|
func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
t.Error("should not call any provider (all in cooldown)")
|
||||||
|
return nil, nil
|
||||||
|
})
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error when all in cooldown")
|
||||||
|
}
|
||||||
|
var exhausted *FallbackExhaustedError
|
||||||
|
if !errors.As(err, &exhausted) {
|
||||||
|
t.Fatalf("expected FallbackExhaustedError, got %T", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_NoCandidates(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
_, err := fc.Execute(context.Background(), nil, successRun("ok"))
|
||||||
|
if err == nil {
|
||||||
|
t.Error("expected error for empty candidates")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_EmptyFallbacks(t *testing.T) {
|
||||||
|
// Single primary, no fallbacks: should work like direct call
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{makeCandidate("openai", "gpt-4")}
|
||||||
|
result, err := fc.Execute(context.Background(), candidates, successRun("ok"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Response.Content != "ok" {
|
||||||
|
t.Error("expected success with single candidate")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_UnclassifiedError(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
return nil, errors.New("completely unknown internal error")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.Execute(context.Background(), candidates, run)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error for unclassified error")
|
||||||
|
}
|
||||||
|
if attempt != 1 {
|
||||||
|
t.Errorf("attempt = %d, want 1 (should not fallback on unclassified)", attempt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallback_SuccessResetsCooldown(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{makeCandidate("openai", "gpt-4")}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
if attempt == 1 {
|
||||||
|
ct.MarkFailure("openai", FailoverRateLimit) // simulate failure tracked elsewhere
|
||||||
|
}
|
||||||
|
return &LLMResponse{Content: "ok", FinishReason: "stop"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.Execute(context.Background(), candidates, run)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if !ct.IsAvailable("openai") {
|
||||||
|
t.Error("success should reset cooldown")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Image Fallback Tests ---
|
||||||
|
|
||||||
|
func TestImageFallback_Success(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{makeCandidate("openai", "gpt-4o")}
|
||||||
|
result, err := fc.ExecuteImage(context.Background(), candidates, successRun("image result"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Response.Content != "image result" {
|
||||||
|
t.Error("expected image result")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImageFallback_DimensionError(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4o"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
return nil, errors.New("image dimensions exceed max 4096x4096")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.ExecuteImage(context.Background(), candidates, run)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error for image dimension error")
|
||||||
|
}
|
||||||
|
if attempt != 1 {
|
||||||
|
t.Errorf("attempt = %d, want 1 (image dimension error should not retry)", attempt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImageFallback_SizeError(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4o"),
|
||||||
|
makeCandidate("anthropic", "claude"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
return nil, errors.New("image exceeds 20 mb")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := fc.ExecuteImage(context.Background(), candidates, run)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error for image size error")
|
||||||
|
}
|
||||||
|
if attempt != 1 {
|
||||||
|
t.Errorf("attempt = %d, want 1 (image size error should not retry)", attempt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImageFallback_RetryOnOtherErrors(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
candidates := []FallbackCandidate{
|
||||||
|
makeCandidate("openai", "gpt-4o"),
|
||||||
|
makeCandidate("anthropic", "claude-sonnet"),
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt := 0
|
||||||
|
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
||||||
|
attempt++
|
||||||
|
if attempt == 1 {
|
||||||
|
return nil, errors.New("rate limit exceeded")
|
||||||
|
}
|
||||||
|
return &LLMResponse{Content: "image ok", FinishReason: "stop"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := fc.ExecuteImage(context.Background(), candidates, run)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Provider != "anthropic" {
|
||||||
|
t.Errorf("provider = %q, want anthropic", result.Provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImageFallback_NoCandidates(t *testing.T) {
|
||||||
|
ct := NewCooldownTracker()
|
||||||
|
fc := NewFallbackChain(ct)
|
||||||
|
|
||||||
|
_, err := fc.ExecuteImage(context.Background(), nil, successRun("ok"))
|
||||||
|
if err == nil {
|
||||||
|
t.Error("expected error for empty candidates")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- ResolveCandidates Tests ---
|
||||||
|
|
||||||
|
func TestResolveCandidates_Simple(t *testing.T) {
|
||||||
|
cfg := ModelConfig{
|
||||||
|
Primary: "gpt-4",
|
||||||
|
Fallbacks: []string{"anthropic/claude-opus", "groq/llama-3"},
|
||||||
|
}
|
||||||
|
|
||||||
|
candidates := ResolveCandidates(cfg, "openai")
|
||||||
|
if len(candidates) != 3 {
|
||||||
|
t.Fatalf("candidates = %d, want 3", len(candidates))
|
||||||
|
}
|
||||||
|
|
||||||
|
if candidates[0].Provider != "openai" || candidates[0].Model != "gpt-4" {
|
||||||
|
t.Errorf("candidate[0] = %s/%s, want openai/gpt-4", candidates[0].Provider, candidates[0].Model)
|
||||||
|
}
|
||||||
|
if candidates[1].Provider != "anthropic" || candidates[1].Model != "claude-opus" {
|
||||||
|
t.Errorf("candidate[1] = %s/%s, want anthropic/claude-opus", candidates[1].Provider, candidates[1].Model)
|
||||||
|
}
|
||||||
|
if candidates[2].Provider != "groq" || candidates[2].Model != "llama-3" {
|
||||||
|
t.Errorf("candidate[2] = %s/%s, want groq/llama-3", candidates[2].Provider, candidates[2].Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveCandidates_Deduplication(t *testing.T) {
|
||||||
|
cfg := ModelConfig{
|
||||||
|
Primary: "openai/gpt-4",
|
||||||
|
Fallbacks: []string{"openai/gpt-4", "anthropic/claude"},
|
||||||
|
}
|
||||||
|
|
||||||
|
candidates := ResolveCandidates(cfg, "default")
|
||||||
|
if len(candidates) != 2 {
|
||||||
|
t.Errorf("candidates = %d, want 2 (duplicate removed)", len(candidates))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveCandidates_EmptyFallbacks(t *testing.T) {
|
||||||
|
cfg := ModelConfig{
|
||||||
|
Primary: "gpt-4",
|
||||||
|
Fallbacks: nil,
|
||||||
|
}
|
||||||
|
|
||||||
|
candidates := ResolveCandidates(cfg, "openai")
|
||||||
|
if len(candidates) != 1 {
|
||||||
|
t.Errorf("candidates = %d, want 1", len(candidates))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveCandidates_EmptyPrimary(t *testing.T) {
|
||||||
|
cfg := ModelConfig{
|
||||||
|
Primary: "",
|
||||||
|
Fallbacks: []string{"anthropic/claude"},
|
||||||
|
}
|
||||||
|
|
||||||
|
candidates := ResolveCandidates(cfg, "openai")
|
||||||
|
if len(candidates) != 1 {
|
||||||
|
t.Errorf("candidates = %d, want 1", len(candidates))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallbackExhaustedError_Message(t *testing.T) {
|
||||||
|
e := &FallbackExhaustedError{
|
||||||
|
Attempts: []FallbackAttempt{
|
||||||
|
{Provider: "openai", Model: "gpt-4", Error: errors.New("rate limited"), Reason: FailoverRateLimit, Duration: 500 * time.Millisecond},
|
||||||
|
{Provider: "anthropic", Model: "claude", Skipped: true},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
msg := e.Error()
|
||||||
|
if msg == "" {
|
||||||
|
t.Error("expected non-empty error message")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -7,462 +7,31 @@
|
||||||
package providers
|
package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
"github.com/sipeed/picoclaw/pkg/providers/openai_compat"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type HTTPProvider struct {
|
type HTTPProvider struct {
|
||||||
apiKey string
|
delegate *openai_compat.Provider
|
||||||
apiBase string
|
|
||||||
httpClient *http.Client
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewHTTPProvider(apiKey, apiBase, proxy string) *HTTPProvider {
|
func NewHTTPProvider(apiKey, apiBase, proxy string) *HTTPProvider {
|
||||||
client := &http.Client{
|
|
||||||
Timeout: 120 * time.Second,
|
|
||||||
}
|
|
||||||
|
|
||||||
if proxy != "" {
|
|
||||||
proxyURL, err := url.Parse(proxy)
|
|
||||||
if err == nil {
|
|
||||||
client.Transport = &http.Transport{
|
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return &HTTPProvider{
|
return &HTTPProvider{
|
||||||
apiKey: apiKey,
|
delegate: openai_compat.NewProvider(apiKey, apiBase, proxy),
|
||||||
apiBase: strings.TrimRight(apiBase, "/"),
|
}
|
||||||
httpClient: client,
|
}
|
||||||
|
|
||||||
|
func NewHTTPProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string) *HTTPProvider {
|
||||||
|
return &HTTPProvider{
|
||||||
|
delegate: openai_compat.NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
if p.apiBase == "" {
|
return p.delegate.Chat(ctx, messages, tools, model, options)
|
||||||
return nil, fmt.Errorf("API base not configured")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Strip provider prefix from model name (e.g., moonshot/kimi-k2.5 -> kimi-k2.5, groq/openai/gpt-oss-120b -> openai/gpt-oss-120b, ollama/qwen2.5:14b -> qwen2.5:14b)
|
|
||||||
if idx := strings.Index(model, "/"); idx != -1 {
|
|
||||||
prefix := model[:idx]
|
|
||||||
if prefix == "moonshot" || prefix == "nvidia" || prefix == "groq" || prefix == "ollama" {
|
|
||||||
model = model[idx+1:]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
requestBody := map[string]interface{}{
|
|
||||||
"model": model,
|
|
||||||
"messages": messages,
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(tools) > 0 {
|
|
||||||
requestBody["tools"] = tools
|
|
||||||
requestBody["tool_choice"] = "auto"
|
|
||||||
}
|
|
||||||
|
|
||||||
if maxTokens, ok := options["max_tokens"].(int); ok {
|
|
||||||
lowerModel := strings.ToLower(model)
|
|
||||||
if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") {
|
|
||||||
requestBody["max_completion_tokens"] = maxTokens
|
|
||||||
} else {
|
|
||||||
requestBody["max_tokens"] = maxTokens
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if temperature, ok := options["temperature"].(float64); ok {
|
|
||||||
lowerModel := strings.ToLower(model)
|
|
||||||
// Kimi k2 models only support temperature=1
|
|
||||||
if strings.Contains(lowerModel, "kimi") && strings.Contains(lowerModel, "k2") {
|
|
||||||
requestBody["temperature"] = 1.0
|
|
||||||
} else {
|
|
||||||
requestBody["temperature"] = temperature
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
jsonData, err := json.Marshal(requestBody)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/chat/completions", bytes.NewReader(jsonData))
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
if p.apiKey != "" {
|
|
||||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp, err := p.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 request failed:\n Status: %d\n Body: %s", resp.StatusCode, string(body))
|
|
||||||
}
|
|
||||||
|
|
||||||
return p.parseResponse(body)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *HTTPProvider) parseResponse(body []byte) (*LLMResponse, error) {
|
|
||||||
var apiResponse struct {
|
|
||||||
Choices []struct {
|
|
||||||
Message struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
ToolCalls []struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
Type string `json:"type"`
|
|
||||||
Function *struct {
|
|
||||||
Name string `json:"name"`
|
|
||||||
Arguments string `json:"arguments"`
|
|
||||||
} `json:"function"`
|
|
||||||
} `json:"tool_calls"`
|
|
||||||
} `json:"message"`
|
|
||||||
FinishReason string `json:"finish_reason"`
|
|
||||||
} `json:"choices"`
|
|
||||||
Usage *UsageInfo `json:"usage"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := json.Unmarshal(body, &apiResponse); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(apiResponse.Choices) == 0 {
|
|
||||||
return &LLMResponse{
|
|
||||||
Content: "",
|
|
||||||
FinishReason: "stop",
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
choice := apiResponse.Choices[0]
|
|
||||||
|
|
||||||
toolCalls := make([]ToolCall, 0, len(choice.Message.ToolCalls))
|
|
||||||
for _, tc := range choice.Message.ToolCalls {
|
|
||||||
arguments := make(map[string]interface{})
|
|
||||||
name := ""
|
|
||||||
|
|
||||||
// Handle OpenAI format with nested function object
|
|
||||||
if tc.Type == "function" && tc.Function != nil {
|
|
||||||
name = tc.Function.Name
|
|
||||||
if tc.Function.Arguments != "" {
|
|
||||||
if err := json.Unmarshal([]byte(tc.Function.Arguments), &arguments); err != nil {
|
|
||||||
arguments["raw"] = tc.Function.Arguments
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else if tc.Function != nil {
|
|
||||||
// Legacy format without type field
|
|
||||||
name = tc.Function.Name
|
|
||||||
if tc.Function.Arguments != "" {
|
|
||||||
if err := json.Unmarshal([]byte(tc.Function.Arguments), &arguments); err != nil {
|
|
||||||
arguments["raw"] = tc.Function.Arguments
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
toolCalls = append(toolCalls, ToolCall{
|
|
||||||
ID: tc.ID,
|
|
||||||
Name: name,
|
|
||||||
Arguments: arguments,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
return &LLMResponse{
|
|
||||||
Content: choice.Message.Content,
|
|
||||||
ToolCalls: toolCalls,
|
|
||||||
FinishReason: choice.FinishReason,
|
|
||||||
Usage: apiResponse.Usage,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *HTTPProvider) GetDefaultModel() string {
|
func (p *HTTPProvider) GetDefaultModel() string {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func createClaudeAuthProvider() (LLMProvider, error) {
|
|
||||||
cred, err := auth.GetCredential("anthropic")
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
|
||||||
}
|
|
||||||
if cred == nil {
|
|
||||||
return nil, fmt.Errorf("no credentials for anthropic. Run: picoclaw auth login --provider anthropic")
|
|
||||||
}
|
|
||||||
return NewClaudeProviderWithTokenSource(cred.AccessToken, createClaudeTokenSource()), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func createCodexAuthProvider() (LLMProvider, error) {
|
|
||||||
cred, err := auth.GetCredential("openai")
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
|
||||||
}
|
|
||||||
if cred == nil {
|
|
||||||
return nil, fmt.Errorf("no credentials for openai. Run: picoclaw auth login --provider openai")
|
|
||||||
}
|
|
||||||
return NewCodexProviderWithTokenSource(cred.AccessToken, cred.AccountID, createCodexTokenSource()), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
|
||||||
model := cfg.Agents.Defaults.Model
|
|
||||||
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
|
||||||
|
|
||||||
var apiKey, apiBase, proxy string
|
|
||||||
|
|
||||||
lowerModel := strings.ToLower(model)
|
|
||||||
|
|
||||||
// First, try to use explicitly configured provider
|
|
||||||
if providerName != "" {
|
|
||||||
switch providerName {
|
|
||||||
case "groq":
|
|
||||||
if cfg.Providers.Groq.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.Groq.APIKey
|
|
||||||
apiBase = cfg.Providers.Groq.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.groq.com/openai/v1"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "openai", "gpt":
|
|
||||||
if cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != "" {
|
|
||||||
if cfg.Providers.OpenAI.AuthMethod == "codex-cli" {
|
|
||||||
return NewCodexProviderWithTokenSource("", "", CreateCodexCliTokenSource()), nil
|
|
||||||
}
|
|
||||||
if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" {
|
|
||||||
return createCodexAuthProvider()
|
|
||||||
}
|
|
||||||
apiKey = cfg.Providers.OpenAI.APIKey
|
|
||||||
apiBase = cfg.Providers.OpenAI.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.openai.com/v1"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "anthropic", "claude":
|
|
||||||
if cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != "" {
|
|
||||||
if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" {
|
|
||||||
return createClaudeAuthProvider()
|
|
||||||
}
|
|
||||||
apiKey = cfg.Providers.Anthropic.APIKey
|
|
||||||
apiBase = cfg.Providers.Anthropic.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.anthropic.com/v1"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "openrouter":
|
|
||||||
if cfg.Providers.OpenRouter.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.OpenRouter.APIKey
|
|
||||||
if cfg.Providers.OpenRouter.APIBase != "" {
|
|
||||||
apiBase = cfg.Providers.OpenRouter.APIBase
|
|
||||||
} else {
|
|
||||||
apiBase = "https://openrouter.ai/api/v1"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "zhipu", "glm":
|
|
||||||
if cfg.Providers.Zhipu.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.Zhipu.APIKey
|
|
||||||
apiBase = cfg.Providers.Zhipu.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://open.bigmodel.cn/api/paas/v4"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "gemini", "google":
|
|
||||||
if cfg.Providers.Gemini.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.Gemini.APIKey
|
|
||||||
apiBase = cfg.Providers.Gemini.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "vllm":
|
|
||||||
if cfg.Providers.VLLM.APIBase != "" {
|
|
||||||
apiKey = cfg.Providers.VLLM.APIKey
|
|
||||||
apiBase = cfg.Providers.VLLM.APIBase
|
|
||||||
}
|
|
||||||
case "shengsuanyun":
|
|
||||||
if cfg.Providers.ShengSuanYun.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.ShengSuanYun.APIKey
|
|
||||||
apiBase = cfg.Providers.ShengSuanYun.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://router.shengsuanyun.com/api/v1"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "claude-cli", "claudecode", "claude-code":
|
|
||||||
workspace := cfg.WorkspacePath()
|
|
||||||
if workspace == "" {
|
|
||||||
workspace = "."
|
|
||||||
}
|
|
||||||
return NewClaudeCliProvider(workspace), nil
|
|
||||||
case "codex-cli", "codex-code":
|
|
||||||
workspace := cfg.WorkspacePath()
|
|
||||||
if workspace == "" {
|
|
||||||
workspace = "."
|
|
||||||
}
|
|
||||||
return NewCodexCliProvider(workspace), nil
|
|
||||||
case "deepseek":
|
|
||||||
if cfg.Providers.DeepSeek.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.DeepSeek.APIKey
|
|
||||||
apiBase = cfg.Providers.DeepSeek.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.deepseek.com/v1"
|
|
||||||
}
|
|
||||||
if model != "deepseek-chat" && model != "deepseek-reasoner" {
|
|
||||||
model = "deepseek-chat"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "github_copilot", "copilot":
|
|
||||||
if cfg.Providers.GitHubCopilot.APIBase != "" {
|
|
||||||
apiBase = cfg.Providers.GitHubCopilot.APIBase
|
|
||||||
} else {
|
|
||||||
apiBase = "localhost:4321"
|
|
||||||
}
|
|
||||||
return NewGitHubCopilotProvider(apiBase, cfg.Providers.GitHubCopilot.ConnectMode, model)
|
|
||||||
|
|
||||||
case "volcengine", "doubao":
|
|
||||||
if cfg.Providers.VolcEngine.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.VolcEngine.APIKey
|
|
||||||
apiBase = cfg.Providers.VolcEngine.APIBase
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://ark.cn-beijing.volces.com/api/v3"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// Fallback: detect provider from model name
|
|
||||||
if apiKey == "" && apiBase == "" {
|
|
||||||
switch {
|
|
||||||
case (strings.Contains(lowerModel, "kimi") || strings.Contains(lowerModel, "moonshot") || strings.HasPrefix(model, "moonshot/")) && cfg.Providers.Moonshot.APIKey != "":
|
|
||||||
apiKey = cfg.Providers.Moonshot.APIKey
|
|
||||||
apiBase = cfg.Providers.Moonshot.APIBase
|
|
||||||
proxy = cfg.Providers.Moonshot.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.moonshot.cn/v1"
|
|
||||||
}
|
|
||||||
|
|
||||||
case strings.HasPrefix(model, "openrouter/") || strings.HasPrefix(model, "anthropic/") || strings.HasPrefix(model, "openai/") || strings.HasPrefix(model, "meta-llama/") || strings.HasPrefix(model, "deepseek/") || strings.HasPrefix(model, "google/"):
|
|
||||||
apiKey = cfg.Providers.OpenRouter.APIKey
|
|
||||||
proxy = cfg.Providers.OpenRouter.Proxy
|
|
||||||
if cfg.Providers.OpenRouter.APIBase != "" {
|
|
||||||
apiBase = cfg.Providers.OpenRouter.APIBase
|
|
||||||
} else {
|
|
||||||
apiBase = "https://openrouter.ai/api/v1"
|
|
||||||
}
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "claude") || strings.HasPrefix(model, "anthropic/")) && (cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != ""):
|
|
||||||
if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" {
|
|
||||||
return createClaudeAuthProvider()
|
|
||||||
}
|
|
||||||
apiKey = cfg.Providers.Anthropic.APIKey
|
|
||||||
apiBase = cfg.Providers.Anthropic.APIBase
|
|
||||||
proxy = cfg.Providers.Anthropic.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.anthropic.com/v1"
|
|
||||||
}
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "gpt") || strings.HasPrefix(model, "openai/")) && (cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != ""):
|
|
||||||
if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" {
|
|
||||||
return createCodexAuthProvider()
|
|
||||||
}
|
|
||||||
apiKey = cfg.Providers.OpenAI.APIKey
|
|
||||||
apiBase = cfg.Providers.OpenAI.APIBase
|
|
||||||
proxy = cfg.Providers.OpenAI.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.openai.com/v1"
|
|
||||||
}
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "gemini") || strings.HasPrefix(model, "google/")) && cfg.Providers.Gemini.APIKey != "":
|
|
||||||
apiKey = cfg.Providers.Gemini.APIKey
|
|
||||||
apiBase = cfg.Providers.Gemini.APIBase
|
|
||||||
proxy = cfg.Providers.Gemini.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
|
||||||
}
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "zhipu") || strings.Contains(lowerModel, "zai")) && cfg.Providers.Zhipu.APIKey != "":
|
|
||||||
apiKey = cfg.Providers.Zhipu.APIKey
|
|
||||||
apiBase = cfg.Providers.Zhipu.APIBase
|
|
||||||
proxy = cfg.Providers.Zhipu.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://open.bigmodel.cn/api/paas/v4"
|
|
||||||
}
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "groq") || strings.HasPrefix(model, "groq/")) && cfg.Providers.Groq.APIKey != "":
|
|
||||||
apiKey = cfg.Providers.Groq.APIKey
|
|
||||||
apiBase = cfg.Providers.Groq.APIBase
|
|
||||||
proxy = cfg.Providers.Groq.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://api.groq.com/openai/v1"
|
|
||||||
}
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "nvidia") || strings.HasPrefix(model, "nvidia/")) && cfg.Providers.Nvidia.APIKey != "":
|
|
||||||
apiKey = cfg.Providers.Nvidia.APIKey
|
|
||||||
apiBase = cfg.Providers.Nvidia.APIBase
|
|
||||||
proxy = cfg.Providers.Nvidia.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://integrate.api.nvidia.com/v1"
|
|
||||||
}
|
|
||||||
case (strings.Contains(lowerModel, "ollama") || strings.HasPrefix(model, "ollama/")) && cfg.Providers.Ollama.APIKey != "":
|
|
||||||
fmt.Println("Ollama provider selected based on model name prefix")
|
|
||||||
apiKey = cfg.Providers.Ollama.APIKey
|
|
||||||
apiBase = cfg.Providers.Ollama.APIBase
|
|
||||||
proxy = cfg.Providers.Ollama.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "http://localhost:11434/v1"
|
|
||||||
}
|
|
||||||
fmt.Println("Ollama apiBase:", apiBase)
|
|
||||||
|
|
||||||
case (strings.Contains(lowerModel, "doubao") || strings.HasPrefix(lowerModel, "doubao") || strings.Contains(lowerModel, "volcengine")) && cfg.Providers.VolcEngine.APIKey != "":
|
|
||||||
apiKey = cfg.Providers.VolcEngine.APIKey
|
|
||||||
apiBase = cfg.Providers.VolcEngine.APIBase
|
|
||||||
proxy = cfg.Providers.VolcEngine.Proxy
|
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "https://ark.cn-beijing.volces.com/api/v3"
|
|
||||||
}
|
|
||||||
|
|
||||||
case cfg.Providers.VLLM.APIBase != "":
|
|
||||||
apiKey = cfg.Providers.VLLM.APIKey
|
|
||||||
apiBase = cfg.Providers.VLLM.APIBase
|
|
||||||
proxy = cfg.Providers.VLLM.Proxy
|
|
||||||
|
|
||||||
default:
|
|
||||||
if cfg.Providers.OpenRouter.APIKey != "" {
|
|
||||||
apiKey = cfg.Providers.OpenRouter.APIKey
|
|
||||||
proxy = cfg.Providers.OpenRouter.Proxy
|
|
||||||
if cfg.Providers.OpenRouter.APIBase != "" {
|
|
||||||
apiBase = cfg.Providers.OpenRouter.APIBase
|
|
||||||
} else {
|
|
||||||
apiBase = "https://openrouter.ai/api/v1"
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
return nil, fmt.Errorf("no API key configured for model: %s", model)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if apiKey == "" && !strings.HasPrefix(model, "bedrock/") {
|
|
||||||
return nil, fmt.Errorf("no API key configured for provider (model: %s)", model)
|
|
||||||
}
|
|
||||||
|
|
||||||
if apiBase == "" {
|
|
||||||
return nil, fmt.Errorf("no API base configured for provider (model: %s)", model)
|
|
||||||
}
|
|
||||||
|
|
||||||
return NewHTTPProvider(apiKey, apiBase, proxy), nil
|
|
||||||
}
|
|
||||||
|
|
|
||||||
49
pkg/providers/legacy_provider.go
Normal file
49
pkg/providers/legacy_provider.go
Normal file
|
|
@ -0,0 +1,49 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CreateProvider creates a provider based on the configuration.
|
||||||
|
// It uses the model_list configuration (new format) to create providers.
|
||||||
|
// The old providers config is automatically converted to model_list during config loading.
|
||||||
|
// Returns the provider, the model ID to use, and any error.
|
||||||
|
func CreateProvider(cfg *config.Config) (LLMProvider, string, error) {
|
||||||
|
model := cfg.Agents.Defaults.Model
|
||||||
|
|
||||||
|
// Ensure model_list is populated (should be done by LoadConfig, but handle edge cases)
|
||||||
|
if len(cfg.ModelList) == 0 && cfg.HasProvidersConfig() {
|
||||||
|
cfg.ModelList = config.ConvertProvidersToModelList(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Must have model_list at this point
|
||||||
|
if len(cfg.ModelList) == 0 {
|
||||||
|
return nil, "", fmt.Errorf("no providers configured. Please add entries to model_list in your config")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get model config from model_list
|
||||||
|
modelCfg, err := cfg.GetModelConfig(model)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", fmt.Errorf("model %q not found in model_list: %w", model, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Inject global workspace if not set in model config
|
||||||
|
if modelCfg.Workspace == "" {
|
||||||
|
modelCfg.Workspace = cfg.WorkspacePath()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use factory to create provider
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(modelCfg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", fmt.Errorf("failed to create provider for model %q: %w", model, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return provider, modelID, nil
|
||||||
|
}
|
||||||
64
pkg/providers/model_ref.go
Normal file
64
pkg/providers/model_ref.go
Normal file
|
|
@ -0,0 +1,64 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "strings"
|
||||||
|
|
||||||
|
// ModelRef represents a parsed model reference with provider and model name.
|
||||||
|
type ModelRef struct {
|
||||||
|
Provider string
|
||||||
|
Model string
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseModelRef parses "anthropic/claude-opus" into {Provider: "anthropic", Model: "claude-opus"}.
|
||||||
|
// If no slash present, uses defaultProvider.
|
||||||
|
// Returns nil for empty input.
|
||||||
|
func ParseModelRef(raw string, defaultProvider string) *ModelRef {
|
||||||
|
raw = strings.TrimSpace(raw)
|
||||||
|
if raw == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if idx := strings.Index(raw, "/"); idx > 0 {
|
||||||
|
provider := NormalizeProvider(raw[:idx])
|
||||||
|
model := strings.TrimSpace(raw[idx+1:])
|
||||||
|
if model == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &ModelRef{Provider: provider, Model: model}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ModelRef{
|
||||||
|
Provider: NormalizeProvider(defaultProvider),
|
||||||
|
Model: raw,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeProvider normalizes provider identifiers to canonical form.
|
||||||
|
func NormalizeProvider(provider string) string {
|
||||||
|
p := strings.ToLower(strings.TrimSpace(provider))
|
||||||
|
|
||||||
|
switch p {
|
||||||
|
case "z.ai", "z-ai":
|
||||||
|
return "zai"
|
||||||
|
case "opencode-zen":
|
||||||
|
return "opencode"
|
||||||
|
case "qwen":
|
||||||
|
return "qwen-portal"
|
||||||
|
case "kimi-code":
|
||||||
|
return "kimi-coding"
|
||||||
|
case "gpt":
|
||||||
|
return "openai"
|
||||||
|
case "claude":
|
||||||
|
return "anthropic"
|
||||||
|
case "glm":
|
||||||
|
return "zhipu"
|
||||||
|
case "google":
|
||||||
|
return "gemini"
|
||||||
|
}
|
||||||
|
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelKey returns a canonical "provider/model" key for deduplication.
|
||||||
|
func ModelKey(provider, model string) string {
|
||||||
|
return NormalizeProvider(provider) + "/" + strings.ToLower(strings.TrimSpace(model))
|
||||||
|
}
|
||||||
125
pkg/providers/model_ref_test.go
Normal file
125
pkg/providers/model_ref_test.go
Normal file
|
|
@ -0,0 +1,125 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestParseModelRef_WithSlash(t *testing.T) {
|
||||||
|
ref := ParseModelRef("anthropic/claude-opus", "openai")
|
||||||
|
if ref == nil {
|
||||||
|
t.Fatal("expected non-nil ref")
|
||||||
|
}
|
||||||
|
if ref.Provider != "anthropic" {
|
||||||
|
t.Errorf("provider = %q, want anthropic", ref.Provider)
|
||||||
|
}
|
||||||
|
if ref.Model != "claude-opus" {
|
||||||
|
t.Errorf("model = %q, want claude-opus", ref.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseModelRef_WithoutSlash(t *testing.T) {
|
||||||
|
ref := ParseModelRef("gpt-4", "openai")
|
||||||
|
if ref == nil {
|
||||||
|
t.Fatal("expected non-nil ref")
|
||||||
|
}
|
||||||
|
if ref.Provider != "openai" {
|
||||||
|
t.Errorf("provider = %q, want openai", ref.Provider)
|
||||||
|
}
|
||||||
|
if ref.Model != "gpt-4" {
|
||||||
|
t.Errorf("model = %q, want gpt-4", ref.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseModelRef_Empty(t *testing.T) {
|
||||||
|
ref := ParseModelRef("", "openai")
|
||||||
|
if ref != nil {
|
||||||
|
t.Errorf("expected nil for empty string, got %+v", ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseModelRef_EmptyModelAfterSlash(t *testing.T) {
|
||||||
|
ref := ParseModelRef("openai/", "default")
|
||||||
|
if ref != nil {
|
||||||
|
t.Errorf("expected nil for empty model, got %+v", ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseModelRef_WhitespaceHandling(t *testing.T) {
|
||||||
|
ref := ParseModelRef(" anthropic / claude-opus ", "openai")
|
||||||
|
if ref == nil {
|
||||||
|
t.Fatal("expected non-nil ref")
|
||||||
|
}
|
||||||
|
if ref.Provider != "anthropic" {
|
||||||
|
t.Errorf("provider = %q, want anthropic", ref.Provider)
|
||||||
|
}
|
||||||
|
if ref.Model != "claude-opus" {
|
||||||
|
t.Errorf("model = %q, want claude-opus", ref.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeProvider(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"OpenAI", "openai"},
|
||||||
|
{"ANTHROPIC", "anthropic"},
|
||||||
|
{"z.ai", "zai"},
|
||||||
|
{"z-ai", "zai"},
|
||||||
|
{"Z.AI", "zai"},
|
||||||
|
{"opencode-zen", "opencode"},
|
||||||
|
{"qwen", "qwen-portal"},
|
||||||
|
{"kimi-code", "kimi-coding"},
|
||||||
|
{"gpt", "openai"},
|
||||||
|
{"claude", "anthropic"},
|
||||||
|
{"glm", "zhipu"},
|
||||||
|
{"google", "gemini"},
|
||||||
|
{"groq", "groq"},
|
||||||
|
{"", ""},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := NormalizeProvider(tt.input)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("NormalizeProvider(%q) = %q, want %q", tt.input, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelKey(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
provider string
|
||||||
|
model string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"openai", "gpt-4", "openai/gpt-4"},
|
||||||
|
{"Anthropic", "Claude-Opus", "anthropic/claude-opus"},
|
||||||
|
{"claude", "sonnet", "anthropic/sonnet"},
|
||||||
|
{"z.ai", "Model-X", "zai/model-x"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := ModelKey(tt.provider, tt.model)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("ModelKey(%q, %q) = %q, want %q", tt.provider, tt.model, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseModelRef_ProviderNormalization(t *testing.T) {
|
||||||
|
ref := ParseModelRef("Z.AI/model-x", "default")
|
||||||
|
if ref == nil {
|
||||||
|
t.Fatal("expected non-nil ref")
|
||||||
|
}
|
||||||
|
if ref.Provider != "zai" {
|
||||||
|
t.Errorf("provider = %q, want zai", ref.Provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseModelRef_DefaultProviderNormalization(t *testing.T) {
|
||||||
|
ref := ParseModelRef("gpt-4o", "GPT")
|
||||||
|
if ref == nil {
|
||||||
|
t.Fatal("expected non-nil ref")
|
||||||
|
}
|
||||||
|
if ref.Provider != "openai" {
|
||||||
|
t.Errorf("provider = %q, want openai (normalized from GPT)", ref.Provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
269
pkg/providers/openai_compat/provider.go
Normal file
269
pkg/providers/openai_compat/provider.go
Normal file
|
|
@ -0,0 +1,269 @@
|
||||||
|
package openai_compat
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ToolCall = protocoltypes.ToolCall
|
||||||
|
type FunctionCall = protocoltypes.FunctionCall
|
||||||
|
type LLMResponse = protocoltypes.LLMResponse
|
||||||
|
type UsageInfo = protocoltypes.UsageInfo
|
||||||
|
type Message = protocoltypes.Message
|
||||||
|
type ToolDefinition = protocoltypes.ToolDefinition
|
||||||
|
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
||||||
|
type ExtraContent = protocoltypes.ExtraContent
|
||||||
|
type GoogleExtra = protocoltypes.GoogleExtra
|
||||||
|
|
||||||
|
type Provider struct {
|
||||||
|
apiKey string
|
||||||
|
apiBase string
|
||||||
|
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
|
||||||
|
httpClient *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProvider(apiKey, apiBase, proxy string) *Provider {
|
||||||
|
return NewProviderWithMaxTokensField(apiKey, apiBase, proxy, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string) *Provider {
|
||||||
|
client := &http.Client{
|
||||||
|
Timeout: 120 * time.Second,
|
||||||
|
}
|
||||||
|
|
||||||
|
if proxy != "" {
|
||||||
|
parsed, err := url.Parse(proxy)
|
||||||
|
if err == nil {
|
||||||
|
client.Transport = &http.Transport{
|
||||||
|
Proxy: http.ProxyURL(parsed),
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
log.Printf("openai_compat: invalid proxy URL %q: %v", proxy, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Provider{
|
||||||
|
apiKey: apiKey,
|
||||||
|
apiBase: strings.TrimRight(apiBase, "/"),
|
||||||
|
maxTokensField: maxTokensField,
|
||||||
|
httpClient: client,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Provider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
|
if p.apiBase == "" {
|
||||||
|
return nil, fmt.Errorf("API base not configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
model = normalizeModel(model, p.apiBase)
|
||||||
|
|
||||||
|
requestBody := map[string]interface{}{
|
||||||
|
"model": model,
|
||||||
|
"messages": messages,
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(tools) > 0 {
|
||||||
|
requestBody["tools"] = tools
|
||||||
|
requestBody["tool_choice"] = "auto"
|
||||||
|
}
|
||||||
|
|
||||||
|
if maxTokens, ok := asInt(options["max_tokens"]); ok {
|
||||||
|
// Use configured maxTokensField if specified, otherwise fallback to model-based detection
|
||||||
|
fieldName := p.maxTokensField
|
||||||
|
if fieldName == "" {
|
||||||
|
// Fallback: detect from model name for backward compatibility
|
||||||
|
lowerModel := strings.ToLower(model)
|
||||||
|
if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") || strings.Contains(lowerModel, "gpt-5") {
|
||||||
|
fieldName = "max_completion_tokens"
|
||||||
|
} else {
|
||||||
|
fieldName = "max_tokens"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
requestBody[fieldName] = maxTokens
|
||||||
|
}
|
||||||
|
|
||||||
|
if temperature, ok := asFloat(options["temperature"]); ok {
|
||||||
|
lowerModel := strings.ToLower(model)
|
||||||
|
// Kimi k2 models only support temperature=1.
|
||||||
|
if strings.Contains(lowerModel, "kimi") && strings.Contains(lowerModel, "k2") {
|
||||||
|
requestBody["temperature"] = 1.0
|
||||||
|
} else {
|
||||||
|
requestBody["temperature"] = temperature
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
jsonData, err := json.Marshal(requestBody)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/chat/completions", bytes.NewReader(jsonData))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
if p.apiKey != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := p.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 request failed:\n Status: %d\n Body: %s", resp.StatusCode, string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
return parseResponse(body)
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseResponse(body []byte) (*LLMResponse, error) {
|
||||||
|
var apiResponse struct {
|
||||||
|
Choices []struct {
|
||||||
|
Message struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
ToolCalls []struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
Function *struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Arguments string `json:"arguments"`
|
||||||
|
} `json:"function"`
|
||||||
|
ExtraContent *struct {
|
||||||
|
Google *struct {
|
||||||
|
ThoughtSignature string `json:"thought_signature"`
|
||||||
|
} `json:"google"`
|
||||||
|
} `json:"extra_content"`
|
||||||
|
} `json:"tool_calls"`
|
||||||
|
} `json:"message"`
|
||||||
|
FinishReason string `json:"finish_reason"`
|
||||||
|
} `json:"choices"`
|
||||||
|
Usage *UsageInfo `json:"usage"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &apiResponse); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(apiResponse.Choices) == 0 {
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: "",
|
||||||
|
FinishReason: "stop",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
choice := apiResponse.Choices[0]
|
||||||
|
toolCalls := make([]ToolCall, 0, len(choice.Message.ToolCalls))
|
||||||
|
for _, tc := range choice.Message.ToolCalls {
|
||||||
|
arguments := make(map[string]interface{})
|
||||||
|
name := ""
|
||||||
|
|
||||||
|
// Extract thought_signature from Gemini/Google-specific extra content
|
||||||
|
thoughtSignature := ""
|
||||||
|
if tc.ExtraContent != nil && tc.ExtraContent.Google != nil {
|
||||||
|
thoughtSignature = tc.ExtraContent.Google.ThoughtSignature
|
||||||
|
}
|
||||||
|
|
||||||
|
if tc.Function != nil {
|
||||||
|
name = tc.Function.Name
|
||||||
|
if tc.Function.Arguments != "" {
|
||||||
|
if err := json.Unmarshal([]byte(tc.Function.Arguments), &arguments); err != nil {
|
||||||
|
log.Printf("openai_compat: failed to decode tool call arguments for %q: %v", name, err)
|
||||||
|
arguments["raw"] = tc.Function.Arguments
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build ToolCall with ExtraContent for Gemini 3 thought_signature persistence
|
||||||
|
toolCall := ToolCall{
|
||||||
|
ID: tc.ID,
|
||||||
|
Name: name,
|
||||||
|
Arguments: arguments,
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
|
}
|
||||||
|
|
||||||
|
if thoughtSignature != "" {
|
||||||
|
toolCall.ExtraContent = &ExtraContent{
|
||||||
|
Google: &GoogleExtra{
|
||||||
|
ThoughtSignature: thoughtSignature,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
toolCalls = append(toolCalls, toolCall)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: choice.Message.Content,
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: choice.FinishReason,
|
||||||
|
Usage: apiResponse.Usage,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeModel(model, apiBase string) string {
|
||||||
|
idx := strings.Index(model, "/")
|
||||||
|
if idx == -1 {
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(strings.ToLower(apiBase), "openrouter.ai") {
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
|
||||||
|
prefix := strings.ToLower(model[:idx])
|
||||||
|
switch prefix {
|
||||||
|
case "moonshot", "nvidia", "groq", "ollama", "deepseek", "google", "openrouter", "zhipu":
|
||||||
|
return model[idx+1:]
|
||||||
|
default:
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func asInt(v interface{}) (int, bool) {
|
||||||
|
switch val := v.(type) {
|
||||||
|
case int:
|
||||||
|
return val, true
|
||||||
|
case int64:
|
||||||
|
return int(val), true
|
||||||
|
case float64:
|
||||||
|
return int(val), true
|
||||||
|
case float32:
|
||||||
|
return int(val), true
|
||||||
|
default:
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func asFloat(v interface{}) (float64, bool) {
|
||||||
|
switch val := v.(type) {
|
||||||
|
case float64:
|
||||||
|
return val, true
|
||||||
|
case float32:
|
||||||
|
return float64(val), true
|
||||||
|
case int:
|
||||||
|
return float64(val), true
|
||||||
|
case int64:
|
||||||
|
return float64(val), true
|
||||||
|
default:
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
}
|
||||||
277
pkg/providers/openai_compat/provider_test.go
Normal file
277
pkg/providers/openai_compat/provider_test.go
Normal file
|
|
@ -0,0 +1,277 @@
|
||||||
|
package openai_compat
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"net/url"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestProviderChat_UsesMaxCompletionTokensForGLM(t *testing.T) {
|
||||||
|
var requestBody map[string]interface{}
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.URL.Path != "/chat/completions" {
|
||||||
|
http.Error(w, "not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"choices": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"message": map[string]interface{}{"content": "ok"},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "glm-4.7", map[string]interface{}{"max_tokens": 1234})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := requestBody["max_completion_tokens"]; !ok {
|
||||||
|
t.Fatalf("expected max_completion_tokens in request body")
|
||||||
|
}
|
||||||
|
if _, ok := requestBody["max_tokens"]; ok {
|
||||||
|
t.Fatalf("did not expect max_tokens key for glm model")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_ParsesToolCalls(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"choices": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"message": map[string]interface{}{
|
||||||
|
"content": "",
|
||||||
|
"tool_calls": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"id": "call_1",
|
||||||
|
"type": "function",
|
||||||
|
"function": map[string]interface{}{
|
||||||
|
"name": "get_weather",
|
||||||
|
"arguments": "{\"city\":\"SF\"}",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"finish_reason": "tool_calls",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"usage": map[string]interface{}{
|
||||||
|
"prompt_tokens": 10,
|
||||||
|
"completion_tokens": 5,
|
||||||
|
"total_tokens": 15,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "gpt-4o", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
if len(out.ToolCalls) != 1 {
|
||||||
|
t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
|
||||||
|
}
|
||||||
|
if out.ToolCalls[0].Name != "get_weather" {
|
||||||
|
t.Fatalf("ToolCalls[0].Name = %q, want %q", out.ToolCalls[0].Name, "get_weather")
|
||||||
|
}
|
||||||
|
if out.ToolCalls[0].Arguments["city"] != "SF" {
|
||||||
|
t.Fatalf("ToolCalls[0].Arguments[city] = %v, want SF", out.ToolCalls[0].Arguments["city"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_HTTPError(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
http.Error(w, "bad request", http.StatusBadRequest)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "gpt-4o", nil)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_StripsMoonshotPrefixAndNormalizesKimiTemperature(t *testing.T) {
|
||||||
|
var requestBody map[string]interface{}
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"choices": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"message": map[string]interface{}{"content": "ok"},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
_, err := p.Chat(
|
||||||
|
t.Context(),
|
||||||
|
[]Message{{Role: "user", Content: "hi"}},
|
||||||
|
nil,
|
||||||
|
"moonshot/kimi-k2.5",
|
||||||
|
map[string]interface{}{"temperature": 0.3},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if requestBody["model"] != "kimi-k2.5" {
|
||||||
|
t.Fatalf("model = %v, want kimi-k2.5", requestBody["model"])
|
||||||
|
}
|
||||||
|
if requestBody["temperature"] != 1.0 {
|
||||||
|
t.Fatalf("temperature = %v, want 1.0", requestBody["temperature"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_StripsGroqAndOllamaPrefixes(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
wantModel string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "strips groq prefix and keeps nested model",
|
||||||
|
input: "groq/openai/gpt-oss-120b",
|
||||||
|
wantModel: "openai/gpt-oss-120b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "strips ollama prefix",
|
||||||
|
input: "ollama/qwen2.5:14b",
|
||||||
|
wantModel: "qwen2.5:14b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "strips deepseek prefix",
|
||||||
|
input: "deepseek/deepseek-chat",
|
||||||
|
wantModel: "deepseek-chat",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var requestBody map[string]interface{}
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"choices": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"message": map[string]interface{}{"content": "ok"},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, tt.input, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if requestBody["model"] != tt.wantModel {
|
||||||
|
t.Fatalf("model = %v, want %s", requestBody["model"], tt.wantModel)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProvider_ProxyConfigured(t *testing.T) {
|
||||||
|
proxyURL := "http://127.0.0.1:8080"
|
||||||
|
p := NewProvider("key", "https://example.com", proxyURL)
|
||||||
|
|
||||||
|
transport, ok := p.httpClient.Transport.(*http.Transport)
|
||||||
|
if !ok || transport == nil {
|
||||||
|
t.Fatalf("expected http transport with proxy, got %T", p.httpClient.Transport)
|
||||||
|
}
|
||||||
|
|
||||||
|
req := &http.Request{URL: &url.URL{Scheme: "https", Host: "api.example.com"}}
|
||||||
|
gotProxy, err := transport.Proxy(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("proxy function returned error: %v", err)
|
||||||
|
}
|
||||||
|
if gotProxy == nil || gotProxy.String() != proxyURL {
|
||||||
|
t.Fatalf("proxy = %v, want %s", gotProxy, proxyURL)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_AcceptsNumericOptionTypes(t *testing.T) {
|
||||||
|
var requestBody map[string]interface{}
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"choices": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"message": map[string]interface{}{"content": "ok"},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("key", server.URL, "")
|
||||||
|
_, err := p.Chat(
|
||||||
|
t.Context(),
|
||||||
|
[]Message{{Role: "user", Content: "hi"}},
|
||||||
|
nil,
|
||||||
|
"gpt-4o",
|
||||||
|
map[string]interface{}{"max_tokens": float64(512), "temperature": 1},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if requestBody["max_tokens"] != float64(512) {
|
||||||
|
t.Fatalf("max_tokens = %v, want 512", requestBody["max_tokens"])
|
||||||
|
}
|
||||||
|
if requestBody["temperature"] != float64(1) {
|
||||||
|
t.Fatalf("temperature = %v, want 1", requestBody["temperature"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeModel_UsesAPIBase(t *testing.T) {
|
||||||
|
if got := normalizeModel("deepseek/deepseek-chat", "https://api.deepseek.com/v1"); got != "deepseek-chat" {
|
||||||
|
t.Fatalf("normalizeModel(deepseek) = %q, want %q", got, "deepseek-chat")
|
||||||
|
}
|
||||||
|
if got := normalizeModel("openrouter/auto", "https://openrouter.ai/api/v1"); got != "openrouter/auto" {
|
||||||
|
t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto")
|
||||||
|
}
|
||||||
|
}
|
||||||
56
pkg/providers/protocoltypes/types.go
Normal file
56
pkg/providers/protocoltypes/types.go
Normal file
|
|
@ -0,0 +1,56 @@
|
||||||
|
package protocoltypes
|
||||||
|
|
||||||
|
type ToolCall struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Type string `json:"type,omitempty"`
|
||||||
|
Function *FunctionCall `json:"function,omitempty"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
Arguments map[string]interface{} `json:"arguments,omitempty"`
|
||||||
|
ThoughtSignature string `json:"-"` // Internal use only
|
||||||
|
ExtraContent *ExtraContent `json:"extra_content,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ExtraContent struct {
|
||||||
|
Google *GoogleExtra `json:"google,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type GoogleExtra struct {
|
||||||
|
ThoughtSignature string `json:"thought_signature,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type FunctionCall struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Arguments string `json:"arguments"`
|
||||||
|
ThoughtSignature string `json:"thought_signature,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type LLMResponse struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
|
||||||
|
FinishReason string `json:"finish_reason"`
|
||||||
|
Usage *UsageInfo `json:"usage,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type UsageInfo struct {
|
||||||
|
PromptTokens int `json:"prompt_tokens"`
|
||||||
|
CompletionTokens int `json:"completion_tokens"`
|
||||||
|
TotalTokens int `json:"total_tokens"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type Message struct {
|
||||||
|
Role string `json:"role"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
|
||||||
|
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ToolDefinition struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Function ToolFunctionDefinition `json:"function"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ToolFunctionDefinition struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Parameters map[string]interface{} `json:"parameters"`
|
||||||
|
}
|
||||||
54
pkg/providers/toolcall_utils.go
Normal file
54
pkg/providers/toolcall_utils.go
Normal file
|
|
@ -0,0 +1,54 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
// NormalizeToolCall normalizes a ToolCall to ensure all fields are properly populated.
|
||||||
|
// It handles cases where Name/Arguments might be in different locations (top-level vs Function)
|
||||||
|
// and ensures both are populated consistently.
|
||||||
|
func NormalizeToolCall(tc ToolCall) ToolCall {
|
||||||
|
normalized := tc
|
||||||
|
|
||||||
|
// Ensure Name is populated from Function if not set
|
||||||
|
if normalized.Name == "" && normalized.Function != nil {
|
||||||
|
normalized.Name = normalized.Function.Name
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure Arguments is not nil
|
||||||
|
if normalized.Arguments == nil {
|
||||||
|
normalized.Arguments = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse Arguments from Function.Arguments if not already set
|
||||||
|
if len(normalized.Arguments) == 0 && normalized.Function != nil && normalized.Function.Arguments != "" {
|
||||||
|
var parsed map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(normalized.Function.Arguments), &parsed); err == nil && parsed != nil {
|
||||||
|
normalized.Arguments = parsed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure Function is populated with consistent values
|
||||||
|
argsJSON, _ := json.Marshal(normalized.Arguments)
|
||||||
|
if normalized.Function == nil {
|
||||||
|
normalized.Function = &FunctionCall{
|
||||||
|
Name: normalized.Name,
|
||||||
|
Arguments: string(argsJSON),
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if normalized.Function.Name == "" {
|
||||||
|
normalized.Function.Name = normalized.Name
|
||||||
|
}
|
||||||
|
if normalized.Name == "" {
|
||||||
|
normalized.Name = normalized.Function.Name
|
||||||
|
}
|
||||||
|
if normalized.Function.Arguments == "" {
|
||||||
|
normalized.Function.Arguments = string(argsJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
|
@ -1,52 +1,66 @@
|
||||||
package providers
|
package providers
|
||||||
|
|
||||||
import "context"
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
type ToolCall struct {
|
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
|
||||||
ID string `json:"id"`
|
)
|
||||||
Type string `json:"type,omitempty"`
|
|
||||||
Function *FunctionCall `json:"function,omitempty"`
|
|
||||||
Name string `json:"name,omitempty"`
|
|
||||||
Arguments map[string]interface{} `json:"arguments,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type FunctionCall struct {
|
type ToolCall = protocoltypes.ToolCall
|
||||||
Name string `json:"name"`
|
type FunctionCall = protocoltypes.FunctionCall
|
||||||
Arguments string `json:"arguments"`
|
type LLMResponse = protocoltypes.LLMResponse
|
||||||
}
|
type UsageInfo = protocoltypes.UsageInfo
|
||||||
|
type Message = protocoltypes.Message
|
||||||
type LLMResponse struct {
|
type ToolDefinition = protocoltypes.ToolDefinition
|
||||||
Content string `json:"content"`
|
type ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
|
||||||
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
|
type ExtraContent = protocoltypes.ExtraContent
|
||||||
FinishReason string `json:"finish_reason"`
|
type GoogleExtra = protocoltypes.GoogleExtra
|
||||||
Usage *UsageInfo `json:"usage,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type UsageInfo struct {
|
|
||||||
PromptTokens int `json:"prompt_tokens"`
|
|
||||||
CompletionTokens int `json:"completion_tokens"`
|
|
||||||
TotalTokens int `json:"total_tokens"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type Message struct {
|
|
||||||
Role string `json:"role"`
|
|
||||||
Content string `json:"content"`
|
|
||||||
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
|
|
||||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type LLMProvider interface {
|
type LLMProvider interface {
|
||||||
Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error)
|
Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error)
|
||||||
GetDefaultModel() string
|
GetDefaultModel() string
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolDefinition struct {
|
// FailoverReason classifies why an LLM request failed for fallback decisions.
|
||||||
Type string `json:"type"`
|
type FailoverReason string
|
||||||
Function ToolFunctionDefinition `json:"function"`
|
|
||||||
|
const (
|
||||||
|
FailoverAuth FailoverReason = "auth"
|
||||||
|
FailoverRateLimit FailoverReason = "rate_limit"
|
||||||
|
FailoverBilling FailoverReason = "billing"
|
||||||
|
FailoverTimeout FailoverReason = "timeout"
|
||||||
|
FailoverFormat FailoverReason = "format"
|
||||||
|
FailoverOverloaded FailoverReason = "overloaded"
|
||||||
|
FailoverUnknown FailoverReason = "unknown"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FailoverError wraps an LLM provider error with classification metadata.
|
||||||
|
type FailoverError struct {
|
||||||
|
Reason FailoverReason
|
||||||
|
Provider string
|
||||||
|
Model string
|
||||||
|
Status int
|
||||||
|
Wrapped error
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolFunctionDefinition struct {
|
func (e *FailoverError) Error() string {
|
||||||
Name string `json:"name"`
|
return fmt.Sprintf("failover(%s): provider=%s model=%s status=%d: %v",
|
||||||
Description string `json:"description"`
|
e.Reason, e.Provider, e.Model, e.Status, e.Wrapped)
|
||||||
Parameters map[string]interface{} `json:"parameters"`
|
}
|
||||||
|
|
||||||
|
func (e *FailoverError) Unwrap() error {
|
||||||
|
return e.Wrapped
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsRetriable returns true if this error should trigger fallback to next candidate.
|
||||||
|
// Non-retriable: Format errors (bad request structure, image dimension/size).
|
||||||
|
func (e *FailoverError) IsRetriable() bool {
|
||||||
|
return e.Reason != FailoverFormat
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelConfig holds primary model and fallback list.
|
||||||
|
type ModelConfig struct {
|
||||||
|
Primary string
|
||||||
|
Fallbacks []string
|
||||||
}
|
}
|
||||||
|
|
|
||||||
66
pkg/routing/agent_id.go
Normal file
66
pkg/routing/agent_id.go
Normal file
|
|
@ -0,0 +1,66 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import (
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
DefaultAgentID = "main"
|
||||||
|
DefaultMainKey = "main"
|
||||||
|
DefaultAccountID = "default"
|
||||||
|
MaxAgentIDLength = 64
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
validIDRe = regexp.MustCompile(`^[a-z0-9][a-z0-9_-]{0,63}$`)
|
||||||
|
invalidCharsRe = regexp.MustCompile(`[^a-z0-9_-]+`)
|
||||||
|
leadingDashRe = regexp.MustCompile(`^-+`)
|
||||||
|
trailingDashRe = regexp.MustCompile(`-+$`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// NormalizeAgentID sanitizes an agent ID to [a-z0-9][a-z0-9_-]{0,63}.
|
||||||
|
// Invalid characters are collapsed to "-". Leading/trailing dashes stripped.
|
||||||
|
// Empty input returns DefaultAgentID ("main").
|
||||||
|
func NormalizeAgentID(id string) string {
|
||||||
|
trimmed := strings.TrimSpace(id)
|
||||||
|
if trimmed == "" {
|
||||||
|
return DefaultAgentID
|
||||||
|
}
|
||||||
|
lower := strings.ToLower(trimmed)
|
||||||
|
if validIDRe.MatchString(lower) {
|
||||||
|
return lower
|
||||||
|
}
|
||||||
|
result := invalidCharsRe.ReplaceAllString(lower, "-")
|
||||||
|
result = leadingDashRe.ReplaceAllString(result, "")
|
||||||
|
result = trailingDashRe.ReplaceAllString(result, "")
|
||||||
|
if len(result) > MaxAgentIDLength {
|
||||||
|
result = result[:MaxAgentIDLength]
|
||||||
|
}
|
||||||
|
if result == "" {
|
||||||
|
return DefaultAgentID
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeAccountID sanitizes an account ID. Empty returns DefaultAccountID.
|
||||||
|
func NormalizeAccountID(id string) string {
|
||||||
|
trimmed := strings.TrimSpace(id)
|
||||||
|
if trimmed == "" {
|
||||||
|
return DefaultAccountID
|
||||||
|
}
|
||||||
|
lower := strings.ToLower(trimmed)
|
||||||
|
if validIDRe.MatchString(lower) {
|
||||||
|
return lower
|
||||||
|
}
|
||||||
|
result := invalidCharsRe.ReplaceAllString(lower, "-")
|
||||||
|
result = leadingDashRe.ReplaceAllString(result, "")
|
||||||
|
result = trailingDashRe.ReplaceAllString(result, "")
|
||||||
|
if len(result) > MaxAgentIDLength {
|
||||||
|
result = result[:MaxAgentIDLength]
|
||||||
|
}
|
||||||
|
if result == "" {
|
||||||
|
return DefaultAccountID
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
86
pkg/routing/agent_id_test.go
Normal file
86
pkg/routing/agent_id_test.go
Normal file
|
|
@ -0,0 +1,86 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestNormalizeAgentID_Empty(t *testing.T) {
|
||||||
|
if got := NormalizeAgentID(""); got != DefaultAgentID {
|
||||||
|
t.Errorf("NormalizeAgentID('') = %q, want %q", got, DefaultAgentID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAgentID_Whitespace(t *testing.T) {
|
||||||
|
if got := NormalizeAgentID(" "); got != DefaultAgentID {
|
||||||
|
t.Errorf("NormalizeAgentID(' ') = %q, want %q", got, DefaultAgentID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAgentID_Valid(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input, want string
|
||||||
|
}{
|
||||||
|
{"main", "main"},
|
||||||
|
{"Main", "main"},
|
||||||
|
{"SALES", "sales"},
|
||||||
|
{"support-bot", "support-bot"},
|
||||||
|
{"agent_1", "agent_1"},
|
||||||
|
{"a", "a"},
|
||||||
|
{"0test", "0test"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := NormalizeAgentID(tt.input); got != tt.want {
|
||||||
|
t.Errorf("NormalizeAgentID(%q) = %q, want %q", tt.input, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAgentID_InvalidChars(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input, want string
|
||||||
|
}{
|
||||||
|
{"Hello World", "hello-world"},
|
||||||
|
{"agent@123", "agent-123"},
|
||||||
|
{"foo.bar.baz", "foo-bar-baz"},
|
||||||
|
{"--leading", "leading"},
|
||||||
|
{"--both--", "both"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := NormalizeAgentID(tt.input); got != tt.want {
|
||||||
|
t.Errorf("NormalizeAgentID(%q) = %q, want %q", tt.input, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAgentID_AllInvalid(t *testing.T) {
|
||||||
|
if got := NormalizeAgentID("@@@"); got != DefaultAgentID {
|
||||||
|
t.Errorf("NormalizeAgentID('@@@') = %q, want %q", got, DefaultAgentID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAgentID_TruncatesAt64(t *testing.T) {
|
||||||
|
long := ""
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
long += "a"
|
||||||
|
}
|
||||||
|
got := NormalizeAgentID(long)
|
||||||
|
if len(got) > MaxAgentIDLength {
|
||||||
|
t.Errorf("length = %d, want <= %d", len(got), MaxAgentIDLength)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAccountID_Empty(t *testing.T) {
|
||||||
|
if got := NormalizeAccountID(""); got != DefaultAccountID {
|
||||||
|
t.Errorf("NormalizeAccountID('') = %q, want %q", got, DefaultAccountID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAccountID_Valid(t *testing.T) {
|
||||||
|
if got := NormalizeAccountID("MyBot"); got != "mybot" {
|
||||||
|
t.Errorf("NormalizeAccountID('MyBot') = %q, want 'mybot'", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAccountID_InvalidChars(t *testing.T) {
|
||||||
|
if got := NormalizeAccountID("bot@home"); got != "bot-home" {
|
||||||
|
t.Errorf("NormalizeAccountID('bot@home') = %q, want 'bot-home'", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
252
pkg/routing/route.go
Normal file
252
pkg/routing/route.go
Normal file
|
|
@ -0,0 +1,252 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RouteInput contains the routing context from an inbound message.
|
||||||
|
type RouteInput struct {
|
||||||
|
Channel string
|
||||||
|
AccountID string
|
||||||
|
Peer *RoutePeer
|
||||||
|
ParentPeer *RoutePeer
|
||||||
|
GuildID string
|
||||||
|
TeamID string
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolvedRoute is the result of agent routing.
|
||||||
|
type ResolvedRoute struct {
|
||||||
|
AgentID string
|
||||||
|
Channel string
|
||||||
|
AccountID string
|
||||||
|
SessionKey string
|
||||||
|
MainSessionKey string
|
||||||
|
MatchedBy string // "binding.peer", "binding.peer.parent", "binding.guild", "binding.team", "binding.account", "binding.channel", "default"
|
||||||
|
}
|
||||||
|
|
||||||
|
// RouteResolver determines which agent handles a message based on config bindings.
|
||||||
|
type RouteResolver struct {
|
||||||
|
cfg *config.Config
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRouteResolver creates a new route resolver.
|
||||||
|
func NewRouteResolver(cfg *config.Config) *RouteResolver {
|
||||||
|
return &RouteResolver{cfg: cfg}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveRoute determines which agent handles the message and constructs session keys.
|
||||||
|
// Implements the 7-level priority cascade:
|
||||||
|
// peer > parent_peer > guild > team > account > channel_wildcard > default
|
||||||
|
func (r *RouteResolver) ResolveRoute(input RouteInput) ResolvedRoute {
|
||||||
|
channel := strings.ToLower(strings.TrimSpace(input.Channel))
|
||||||
|
accountID := NormalizeAccountID(input.AccountID)
|
||||||
|
peer := input.Peer
|
||||||
|
|
||||||
|
dmScope := DMScope(r.cfg.Session.DMScope)
|
||||||
|
if dmScope == "" {
|
||||||
|
dmScope = DMScopeMain
|
||||||
|
}
|
||||||
|
identityLinks := r.cfg.Session.IdentityLinks
|
||||||
|
|
||||||
|
bindings := r.filterBindings(channel, accountID)
|
||||||
|
|
||||||
|
choose := func(agentID string, matchedBy string) ResolvedRoute {
|
||||||
|
resolvedAgentID := r.pickAgentID(agentID)
|
||||||
|
sessionKey := strings.ToLower(BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: resolvedAgentID,
|
||||||
|
Channel: channel,
|
||||||
|
AccountID: accountID,
|
||||||
|
Peer: peer,
|
||||||
|
DMScope: dmScope,
|
||||||
|
IdentityLinks: identityLinks,
|
||||||
|
}))
|
||||||
|
mainSessionKey := strings.ToLower(BuildAgentMainSessionKey(resolvedAgentID))
|
||||||
|
return ResolvedRoute{
|
||||||
|
AgentID: resolvedAgentID,
|
||||||
|
Channel: channel,
|
||||||
|
AccountID: accountID,
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
MainSessionKey: mainSessionKey,
|
||||||
|
MatchedBy: matchedBy,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 1: Peer binding
|
||||||
|
if peer != nil && strings.TrimSpace(peer.ID) != "" {
|
||||||
|
if match := r.findPeerMatch(bindings, peer); match != nil {
|
||||||
|
return choose(match.AgentID, "binding.peer")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 2: Parent peer binding
|
||||||
|
parentPeer := input.ParentPeer
|
||||||
|
if parentPeer != nil && strings.TrimSpace(parentPeer.ID) != "" {
|
||||||
|
if match := r.findPeerMatch(bindings, parentPeer); match != nil {
|
||||||
|
return choose(match.AgentID, "binding.peer.parent")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 3: Guild binding
|
||||||
|
guildID := strings.TrimSpace(input.GuildID)
|
||||||
|
if guildID != "" {
|
||||||
|
if match := r.findGuildMatch(bindings, guildID); match != nil {
|
||||||
|
return choose(match.AgentID, "binding.guild")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 4: Team binding
|
||||||
|
teamID := strings.TrimSpace(input.TeamID)
|
||||||
|
if teamID != "" {
|
||||||
|
if match := r.findTeamMatch(bindings, teamID); match != nil {
|
||||||
|
return choose(match.AgentID, "binding.team")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 5: Account binding
|
||||||
|
if match := r.findAccountMatch(bindings); match != nil {
|
||||||
|
return choose(match.AgentID, "binding.account")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 6: Channel wildcard binding
|
||||||
|
if match := r.findChannelWildcardMatch(bindings); match != nil {
|
||||||
|
return choose(match.AgentID, "binding.channel")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Priority 7: Default agent
|
||||||
|
return choose(r.resolveDefaultAgentID(), "default")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) filterBindings(channel, accountID string) []config.AgentBinding {
|
||||||
|
var filtered []config.AgentBinding
|
||||||
|
for _, b := range r.cfg.Bindings {
|
||||||
|
matchChannel := strings.ToLower(strings.TrimSpace(b.Match.Channel))
|
||||||
|
if matchChannel == "" || matchChannel != channel {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !matchesAccountID(b.Match.AccountID, accountID) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
filtered = append(filtered, b)
|
||||||
|
}
|
||||||
|
return filtered
|
||||||
|
}
|
||||||
|
|
||||||
|
func matchesAccountID(matchAccountID, actual string) bool {
|
||||||
|
trimmed := strings.TrimSpace(matchAccountID)
|
||||||
|
if trimmed == "" {
|
||||||
|
return actual == DefaultAccountID
|
||||||
|
}
|
||||||
|
if trimmed == "*" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return strings.ToLower(trimmed) == strings.ToLower(actual)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) findPeerMatch(bindings []config.AgentBinding, peer *RoutePeer) *config.AgentBinding {
|
||||||
|
for i := range bindings {
|
||||||
|
b := &bindings[i]
|
||||||
|
if b.Match.Peer == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
peerKind := strings.ToLower(strings.TrimSpace(b.Match.Peer.Kind))
|
||||||
|
peerID := strings.TrimSpace(b.Match.Peer.ID)
|
||||||
|
if peerKind == "" || peerID == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if peerKind == strings.ToLower(peer.Kind) && peerID == peer.ID {
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) findGuildMatch(bindings []config.AgentBinding, guildID string) *config.AgentBinding {
|
||||||
|
for i := range bindings {
|
||||||
|
b := &bindings[i]
|
||||||
|
matchGuild := strings.TrimSpace(b.Match.GuildID)
|
||||||
|
if matchGuild != "" && matchGuild == guildID {
|
||||||
|
return &bindings[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) findTeamMatch(bindings []config.AgentBinding, teamID string) *config.AgentBinding {
|
||||||
|
for i := range bindings {
|
||||||
|
b := &bindings[i]
|
||||||
|
matchTeam := strings.TrimSpace(b.Match.TeamID)
|
||||||
|
if matchTeam != "" && matchTeam == teamID {
|
||||||
|
return &bindings[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) findAccountMatch(bindings []config.AgentBinding) *config.AgentBinding {
|
||||||
|
for i := range bindings {
|
||||||
|
b := &bindings[i]
|
||||||
|
accountID := strings.TrimSpace(b.Match.AccountID)
|
||||||
|
if accountID == "*" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if b.Match.Peer != nil || b.Match.GuildID != "" || b.Match.TeamID != "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return &bindings[i]
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) findChannelWildcardMatch(bindings []config.AgentBinding) *config.AgentBinding {
|
||||||
|
for i := range bindings {
|
||||||
|
b := &bindings[i]
|
||||||
|
accountID := strings.TrimSpace(b.Match.AccountID)
|
||||||
|
if accountID != "*" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if b.Match.Peer != nil || b.Match.GuildID != "" || b.Match.TeamID != "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return &bindings[i]
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) pickAgentID(agentID string) string {
|
||||||
|
trimmed := strings.TrimSpace(agentID)
|
||||||
|
if trimmed == "" {
|
||||||
|
return NormalizeAgentID(r.resolveDefaultAgentID())
|
||||||
|
}
|
||||||
|
normalized := NormalizeAgentID(trimmed)
|
||||||
|
agents := r.cfg.Agents.List
|
||||||
|
if len(agents) == 0 {
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
for _, a := range agents {
|
||||||
|
if NormalizeAgentID(a.ID) == normalized {
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return NormalizeAgentID(r.resolveDefaultAgentID())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *RouteResolver) resolveDefaultAgentID() string {
|
||||||
|
agents := r.cfg.Agents.List
|
||||||
|
if len(agents) == 0 {
|
||||||
|
return DefaultAgentID
|
||||||
|
}
|
||||||
|
for _, a := range agents {
|
||||||
|
if a.Default {
|
||||||
|
id := strings.TrimSpace(a.ID)
|
||||||
|
if id != "" {
|
||||||
|
return NormalizeAgentID(id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if id := strings.TrimSpace(agents[0].ID); id != "" {
|
||||||
|
return NormalizeAgentID(id)
|
||||||
|
}
|
||||||
|
return DefaultAgentID
|
||||||
|
}
|
||||||
297
pkg/routing/route_test.go
Normal file
297
pkg/routing/route_test.go
Normal file
|
|
@ -0,0 +1,297 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func testConfig(agents []config.AgentConfig, bindings []config.AgentBinding) *config.Config {
|
||||||
|
return &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: "/tmp/picoclaw-test",
|
||||||
|
Model: "gpt-4",
|
||||||
|
},
|
||||||
|
List: agents,
|
||||||
|
},
|
||||||
|
Bindings: bindings,
|
||||||
|
Session: config.SessionConfig{
|
||||||
|
DMScope: "per-peer",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_DefaultAgent_NoBindings(t *testing.T) {
|
||||||
|
cfg := testConfig(nil, nil)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user1"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != DefaultAgentID {
|
||||||
|
t.Errorf("AgentID = %q, want %q", route.AgentID, DefaultAgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "default" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'default'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_PeerBinding(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "sales", Default: true},
|
||||||
|
{ID: "support"},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "support",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "telegram",
|
||||||
|
AccountID: "*",
|
||||||
|
Peer: &config.PeerMatch{Kind: "direct", ID: "user123"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user123"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "support" {
|
||||||
|
t.Errorf("AgentID = %q, want 'support'", route.AgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "binding.peer" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'binding.peer'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_GuildBinding(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "general", Default: true},
|
||||||
|
{ID: "gaming"},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "gaming",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "discord",
|
||||||
|
AccountID: "*",
|
||||||
|
GuildID: "guild-abc",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "discord",
|
||||||
|
GuildID: "guild-abc",
|
||||||
|
Peer: &RoutePeer{Kind: "channel", ID: "ch1"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "gaming" {
|
||||||
|
t.Errorf("AgentID = %q, want 'gaming'", route.AgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "binding.guild" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'binding.guild'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_TeamBinding(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "general", Default: true},
|
||||||
|
{ID: "work"},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "work",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "slack",
|
||||||
|
AccountID: "*",
|
||||||
|
TeamID: "T12345",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "slack",
|
||||||
|
TeamID: "T12345",
|
||||||
|
Peer: &RoutePeer{Kind: "channel", ID: "C001"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "work" {
|
||||||
|
t.Errorf("AgentID = %q, want 'work'", route.AgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "binding.team" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'binding.team'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_AccountBinding(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "default-agent", Default: true},
|
||||||
|
{ID: "premium"},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "premium",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "telegram",
|
||||||
|
AccountID: "bot2",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "telegram",
|
||||||
|
AccountID: "bot2",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user1"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "premium" {
|
||||||
|
t.Errorf("AgentID = %q, want 'premium'", route.AgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "binding.account" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'binding.account'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_ChannelWildcard(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "main", Default: true},
|
||||||
|
{ID: "telegram-bot"},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "telegram-bot",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "telegram",
|
||||||
|
AccountID: "*",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user1"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "telegram-bot" {
|
||||||
|
t.Errorf("AgentID = %q, want 'telegram-bot'", route.AgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "binding.channel" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'binding.channel'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_PriorityOrder_PeerBeatsGuild(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "general", Default: true},
|
||||||
|
{ID: "vip"},
|
||||||
|
{ID: "gaming"},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "vip",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "discord",
|
||||||
|
AccountID: "*",
|
||||||
|
Peer: &config.PeerMatch{Kind: "direct", ID: "user-vip"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
AgentID: "gaming",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "discord",
|
||||||
|
AccountID: "*",
|
||||||
|
GuildID: "guild-1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "discord",
|
||||||
|
GuildID: "guild-1",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user-vip"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "vip" {
|
||||||
|
t.Errorf("AgentID = %q, want 'vip' (peer should beat guild)", route.AgentID)
|
||||||
|
}
|
||||||
|
if route.MatchedBy != "binding.peer" {
|
||||||
|
t.Errorf("MatchedBy = %q, want 'binding.peer'", route.MatchedBy)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_InvalidAgentFallsToDefault(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "main", Default: true},
|
||||||
|
}
|
||||||
|
bindings := []config.AgentBinding{
|
||||||
|
{
|
||||||
|
AgentID: "nonexistent",
|
||||||
|
Match: config.BindingMatch{
|
||||||
|
Channel: "telegram",
|
||||||
|
AccountID: "*",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, bindings)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "telegram",
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "main" {
|
||||||
|
t.Errorf("AgentID = %q, want 'main' (invalid agent should fall to default)", route.AgentID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_DefaultAgentSelection(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "alpha"},
|
||||||
|
{ID: "beta", Default: true},
|
||||||
|
{ID: "gamma"},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, nil)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "cli",
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "beta" {
|
||||||
|
t.Errorf("AgentID = %q, want 'beta' (marked as default)", route.AgentID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveRoute_NoDefaultUsesFirst(t *testing.T) {
|
||||||
|
agents := []config.AgentConfig{
|
||||||
|
{ID: "alpha"},
|
||||||
|
{ID: "beta"},
|
||||||
|
}
|
||||||
|
cfg := testConfig(agents, nil)
|
||||||
|
r := NewRouteResolver(cfg)
|
||||||
|
|
||||||
|
route := r.ResolveRoute(RouteInput{
|
||||||
|
Channel: "cli",
|
||||||
|
})
|
||||||
|
|
||||||
|
if route.AgentID != "alpha" {
|
||||||
|
t.Errorf("AgentID = %q, want 'alpha' (first in list)", route.AgentID)
|
||||||
|
}
|
||||||
|
}
|
||||||
183
pkg/routing/session_key.go
Normal file
183
pkg/routing/session_key.go
Normal file
|
|
@ -0,0 +1,183 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DMScope controls DM session isolation granularity.
|
||||||
|
type DMScope string
|
||||||
|
|
||||||
|
const (
|
||||||
|
DMScopeMain DMScope = "main"
|
||||||
|
DMScopePerPeer DMScope = "per-peer"
|
||||||
|
DMScopePerChannelPeer DMScope = "per-channel-peer"
|
||||||
|
DMScopePerAccountChannelPeer DMScope = "per-account-channel-peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RoutePeer represents a chat peer with kind and ID.
|
||||||
|
type RoutePeer struct {
|
||||||
|
Kind string // "direct", "group", "channel"
|
||||||
|
ID string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SessionKeyParams holds all inputs for session key construction.
|
||||||
|
type SessionKeyParams struct {
|
||||||
|
AgentID string
|
||||||
|
Channel string
|
||||||
|
AccountID string
|
||||||
|
Peer *RoutePeer
|
||||||
|
DMScope DMScope
|
||||||
|
IdentityLinks map[string][]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParsedSessionKey is the result of parsing an agent-scoped session key.
|
||||||
|
type ParsedSessionKey struct {
|
||||||
|
AgentID string
|
||||||
|
Rest string
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildAgentMainSessionKey returns "agent:<agentId>:main".
|
||||||
|
func BuildAgentMainSessionKey(agentID string) string {
|
||||||
|
return fmt.Sprintf("agent:%s:%s", NormalizeAgentID(agentID), DefaultMainKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildAgentPeerSessionKey constructs a session key based on agent, channel, peer, and DM scope.
|
||||||
|
func BuildAgentPeerSessionKey(params SessionKeyParams) string {
|
||||||
|
agentID := NormalizeAgentID(params.AgentID)
|
||||||
|
|
||||||
|
peer := params.Peer
|
||||||
|
if peer == nil {
|
||||||
|
peer = &RoutePeer{Kind: "direct"}
|
||||||
|
}
|
||||||
|
peerKind := strings.TrimSpace(peer.Kind)
|
||||||
|
if peerKind == "" {
|
||||||
|
peerKind = "direct"
|
||||||
|
}
|
||||||
|
|
||||||
|
if peerKind == "direct" {
|
||||||
|
dmScope := params.DMScope
|
||||||
|
if dmScope == "" {
|
||||||
|
dmScope = DMScopeMain
|
||||||
|
}
|
||||||
|
peerID := strings.TrimSpace(peer.ID)
|
||||||
|
|
||||||
|
// Resolve identity links (cross-platform collapse)
|
||||||
|
if dmScope != DMScopeMain && peerID != "" {
|
||||||
|
if linked := resolveLinkedPeerID(params.IdentityLinks, params.Channel, peerID); linked != "" {
|
||||||
|
peerID = linked
|
||||||
|
}
|
||||||
|
}
|
||||||
|
peerID = strings.ToLower(peerID)
|
||||||
|
|
||||||
|
switch dmScope {
|
||||||
|
case DMScopePerAccountChannelPeer:
|
||||||
|
if peerID != "" {
|
||||||
|
channel := normalizeChannel(params.Channel)
|
||||||
|
accountID := NormalizeAccountID(params.AccountID)
|
||||||
|
return fmt.Sprintf("agent:%s:%s:%s:direct:%s", agentID, channel, accountID, peerID)
|
||||||
|
}
|
||||||
|
case DMScopePerChannelPeer:
|
||||||
|
if peerID != "" {
|
||||||
|
channel := normalizeChannel(params.Channel)
|
||||||
|
return fmt.Sprintf("agent:%s:%s:direct:%s", agentID, channel, peerID)
|
||||||
|
}
|
||||||
|
case DMScopePerPeer:
|
||||||
|
if peerID != "" {
|
||||||
|
return fmt.Sprintf("agent:%s:direct:%s", agentID, peerID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return BuildAgentMainSessionKey(agentID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Group/channel peers always get per-peer sessions
|
||||||
|
channel := normalizeChannel(params.Channel)
|
||||||
|
peerID := strings.ToLower(strings.TrimSpace(peer.ID))
|
||||||
|
if peerID == "" {
|
||||||
|
peerID = "unknown"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("agent:%s:%s:%s:%s", agentID, channel, peerKind, peerID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseAgentSessionKey extracts agentId and rest from "agent:<agentId>:<rest>".
|
||||||
|
func ParseAgentSessionKey(sessionKey string) *ParsedSessionKey {
|
||||||
|
raw := strings.TrimSpace(sessionKey)
|
||||||
|
if raw == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
parts := strings.SplitN(raw, ":", 3)
|
||||||
|
if len(parts) < 3 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if parts[0] != "agent" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
agentID := strings.TrimSpace(parts[1])
|
||||||
|
rest := parts[2]
|
||||||
|
if agentID == "" || rest == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &ParsedSessionKey{AgentID: agentID, Rest: rest}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsSubagentSessionKey returns true if the session key represents a subagent.
|
||||||
|
func IsSubagentSessionKey(sessionKey string) bool {
|
||||||
|
raw := strings.TrimSpace(sessionKey)
|
||||||
|
if raw == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(strings.ToLower(raw), "subagent:") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
parsed := ParseAgentSessionKey(raw)
|
||||||
|
if parsed == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.HasPrefix(strings.ToLower(parsed.Rest), "subagent:")
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeChannel(channel string) string {
|
||||||
|
c := strings.TrimSpace(strings.ToLower(channel))
|
||||||
|
if c == "" {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveLinkedPeerID(identityLinks map[string][]string, channel, peerID string) string {
|
||||||
|
if len(identityLinks) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
peerID = strings.TrimSpace(peerID)
|
||||||
|
if peerID == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
candidates := make(map[string]bool)
|
||||||
|
rawCandidate := strings.ToLower(peerID)
|
||||||
|
if rawCandidate != "" {
|
||||||
|
candidates[rawCandidate] = true
|
||||||
|
}
|
||||||
|
channel = strings.ToLower(strings.TrimSpace(channel))
|
||||||
|
if channel != "" {
|
||||||
|
scopedCandidate := fmt.Sprintf("%s:%s", channel, strings.ToLower(peerID))
|
||||||
|
candidates[scopedCandidate] = true
|
||||||
|
}
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
for canonical, ids := range identityLinks {
|
||||||
|
canonicalName := strings.TrimSpace(canonical)
|
||||||
|
if canonicalName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, id := range ids {
|
||||||
|
normalized := strings.ToLower(strings.TrimSpace(id))
|
||||||
|
if normalized != "" && candidates[normalized] {
|
||||||
|
return canonicalName
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
162
pkg/routing/session_key_test.go
Normal file
162
pkg/routing/session_key_test.go
Normal file
|
|
@ -0,0 +1,162 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestBuildAgentMainSessionKey(t *testing.T) {
|
||||||
|
got := BuildAgentMainSessionKey("sales")
|
||||||
|
want := "agent:sales:main"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("BuildAgentMainSessionKey('sales') = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentMainSessionKey_Normalizes(t *testing.T) {
|
||||||
|
got := BuildAgentMainSessionKey("Sales Bot")
|
||||||
|
want := "agent:sales-bot:main"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("BuildAgentMainSessionKey('Sales Bot') = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_DMScopeMain(t *testing.T) {
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user123"},
|
||||||
|
DMScope: DMScopeMain,
|
||||||
|
})
|
||||||
|
want := "agent:main:main"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("DMScopeMain = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_DMScopePerPeer(t *testing.T) {
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user123"},
|
||||||
|
DMScope: DMScopePerPeer,
|
||||||
|
})
|
||||||
|
want := "agent:main:direct:user123"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("DMScopePerPeer = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_DMScopePerChannelPeer(t *testing.T) {
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user123"},
|
||||||
|
DMScope: DMScopePerChannelPeer,
|
||||||
|
})
|
||||||
|
want := "agent:main:telegram:direct:user123"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("DMScopePerChannelPeer = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_DMScopePerAccountChannelPeer(t *testing.T) {
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
AccountID: "bot1",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "User123"},
|
||||||
|
DMScope: DMScopePerAccountChannelPeer,
|
||||||
|
})
|
||||||
|
want := "agent:main:telegram:bot1:direct:user123"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("DMScopePerAccountChannelPeer = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_GroupPeer(t *testing.T) {
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "group", ID: "chat456"},
|
||||||
|
DMScope: DMScopePerPeer,
|
||||||
|
})
|
||||||
|
want := "agent:main:telegram:group:chat456"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("GroupPeer = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_NilPeer(t *testing.T) {
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: nil,
|
||||||
|
DMScope: DMScopePerPeer,
|
||||||
|
})
|
||||||
|
// nil peer defaults to direct with empty ID, falls to main
|
||||||
|
want := "agent:main:main"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("NilPeer = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildAgentPeerSessionKey_IdentityLink(t *testing.T) {
|
||||||
|
links := map[string][]string{
|
||||||
|
"john": {"telegram:user123", "discord:john#1234"},
|
||||||
|
}
|
||||||
|
got := BuildAgentPeerSessionKey(SessionKeyParams{
|
||||||
|
AgentID: "main",
|
||||||
|
Channel: "telegram",
|
||||||
|
Peer: &RoutePeer{Kind: "direct", ID: "user123"},
|
||||||
|
DMScope: DMScopePerPeer,
|
||||||
|
IdentityLinks: links,
|
||||||
|
})
|
||||||
|
want := "agent:main:direct:john"
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("IdentityLink = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseAgentSessionKey_Valid(t *testing.T) {
|
||||||
|
parsed := ParseAgentSessionKey("agent:sales:telegram:direct:user123")
|
||||||
|
if parsed == nil {
|
||||||
|
t.Fatal("expected non-nil result")
|
||||||
|
}
|
||||||
|
if parsed.AgentID != "sales" {
|
||||||
|
t.Errorf("AgentID = %q, want 'sales'", parsed.AgentID)
|
||||||
|
}
|
||||||
|
if parsed.Rest != "telegram:direct:user123" {
|
||||||
|
t.Errorf("Rest = %q, want 'telegram:direct:user123'", parsed.Rest)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseAgentSessionKey_Invalid(t *testing.T) {
|
||||||
|
tests := []string{
|
||||||
|
"",
|
||||||
|
"foo:bar",
|
||||||
|
"notprefix:sales:main",
|
||||||
|
"agent::main",
|
||||||
|
"agent:sales:",
|
||||||
|
}
|
||||||
|
for _, input := range tests {
|
||||||
|
if got := ParseAgentSessionKey(input); got != nil {
|
||||||
|
t.Errorf("ParseAgentSessionKey(%q) = %+v, want nil", input, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsSubagentSessionKey(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"subagent:task-1", true},
|
||||||
|
{"agent:main:subagent:task-1", true},
|
||||||
|
{"agent:main:main", false},
|
||||||
|
{"agent:main:telegram:direct:user123", false},
|
||||||
|
{"", false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := IsSubagentSessionKey(tt.input); got != tt.want {
|
||||||
|
t.Errorf("IsSubagentSessionKey(%q) = %v, want %v", tt.input, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
311
pkg/skills/clawhub_registry.go
Normal file
311
pkg/skills/clawhub_registry.go
Normal file
|
|
@ -0,0 +1,311 @@
|
||||||
|
package skills
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultClawHubTimeout = 30 * time.Second
|
||||||
|
defaultMaxZipSize = 50 * 1024 * 1024 // 50 MB
|
||||||
|
defaultMaxResponseSize = 2 * 1024 * 1024 // 2 MB
|
||||||
|
)
|
||||||
|
|
||||||
|
// ClawHubRegistry implements SkillRegistry for the ClawHub platform.
|
||||||
|
type ClawHubRegistry struct {
|
||||||
|
baseURL string
|
||||||
|
authToken string // Optional - for elevated rate limits
|
||||||
|
searchPath string // Search API
|
||||||
|
skillsPath string // For retrieving skill metadata
|
||||||
|
downloadPath string // For fetching ZIP files for download
|
||||||
|
maxZipSize int
|
||||||
|
maxResponseSize int
|
||||||
|
client *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewClawHubRegistry creates a new ClawHub registry client from config.
|
||||||
|
func NewClawHubRegistry(cfg ClawHubConfig) *ClawHubRegistry {
|
||||||
|
baseURL := cfg.BaseURL
|
||||||
|
if baseURL == "" {
|
||||||
|
baseURL = "https://clawhub.ai"
|
||||||
|
}
|
||||||
|
searchPath := cfg.SearchPath
|
||||||
|
if searchPath == "" {
|
||||||
|
searchPath = "/api/v1/search"
|
||||||
|
}
|
||||||
|
skillsPath := cfg.SkillsPath
|
||||||
|
if skillsPath == "" {
|
||||||
|
skillsPath = "/api/v1/skills"
|
||||||
|
}
|
||||||
|
downloadPath := cfg.DownloadPath
|
||||||
|
if downloadPath == "" {
|
||||||
|
downloadPath = "/api/v1/download"
|
||||||
|
}
|
||||||
|
|
||||||
|
timeout := defaultClawHubTimeout
|
||||||
|
if cfg.Timeout > 0 {
|
||||||
|
timeout = time.Duration(cfg.Timeout) * time.Second
|
||||||
|
}
|
||||||
|
|
||||||
|
maxZip := defaultMaxZipSize
|
||||||
|
if cfg.MaxZipSize > 0 {
|
||||||
|
maxZip = cfg.MaxZipSize
|
||||||
|
}
|
||||||
|
|
||||||
|
maxResp := defaultMaxResponseSize
|
||||||
|
if cfg.MaxResponseSize > 0 {
|
||||||
|
maxResp = cfg.MaxResponseSize
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ClawHubRegistry{
|
||||||
|
baseURL: baseURL,
|
||||||
|
authToken: cfg.AuthToken,
|
||||||
|
searchPath: searchPath,
|
||||||
|
skillsPath: skillsPath,
|
||||||
|
downloadPath: downloadPath,
|
||||||
|
maxZipSize: maxZip,
|
||||||
|
maxResponseSize: maxResp,
|
||||||
|
client: &http.Client{
|
||||||
|
Timeout: timeout,
|
||||||
|
Transport: &http.Transport{
|
||||||
|
MaxIdleConns: 5,
|
||||||
|
IdleConnTimeout: 30 * time.Second,
|
||||||
|
TLSHandshakeTimeout: 10 * time.Second,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) Name() string {
|
||||||
|
return "clawhub"
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Search ---
|
||||||
|
|
||||||
|
type clawhubSearchResponse struct {
|
||||||
|
Results []clawhubSearchResult `json:"results"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type clawhubSearchResult struct {
|
||||||
|
Score float64 `json:"score"`
|
||||||
|
Slug *string `json:"slug"`
|
||||||
|
DisplayName *string `json:"displayName"`
|
||||||
|
Summary *string `json:"summary"`
|
||||||
|
Version *string `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) Search(ctx context.Context, query string, limit int) ([]SearchResult, error) {
|
||||||
|
u, err := url.Parse(c.baseURL + c.searchPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid base URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
q := u.Query()
|
||||||
|
q.Set("q", query)
|
||||||
|
if limit > 0 {
|
||||||
|
q.Set("limit", fmt.Sprintf("%d", limit))
|
||||||
|
}
|
||||||
|
u.RawQuery = q.Encode()
|
||||||
|
|
||||||
|
body, err := c.doGet(ctx, u.String())
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("search request failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp clawhubSearchResponse
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse search response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]SearchResult, 0, len(resp.Results))
|
||||||
|
for _, r := range resp.Results {
|
||||||
|
slug := utils.DerefStr(r.Slug, "")
|
||||||
|
if slug == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
summary := utils.DerefStr(r.Summary, "")
|
||||||
|
if summary == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
displayName := utils.DerefStr(r.DisplayName, "")
|
||||||
|
if displayName == "" {
|
||||||
|
displayName = slug
|
||||||
|
}
|
||||||
|
|
||||||
|
results = append(results, SearchResult{
|
||||||
|
Score: r.Score,
|
||||||
|
Slug: slug,
|
||||||
|
DisplayName: displayName,
|
||||||
|
Summary: summary,
|
||||||
|
Version: utils.DerefStr(r.Version, ""),
|
||||||
|
RegistryName: c.Name(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- GetSkillMeta ---
|
||||||
|
|
||||||
|
type clawhubSkillResponse struct {
|
||||||
|
Slug string `json:"slug"`
|
||||||
|
DisplayName string `json:"displayName"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
LatestVersion *clawhubVersionInfo `json:"latestVersion"`
|
||||||
|
Moderation *clawhubModerationInfo `json:"moderation"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type clawhubVersionInfo struct {
|
||||||
|
Version string `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type clawhubModerationInfo struct {
|
||||||
|
IsMalwareBlocked bool `json:"isMalwareBlocked"`
|
||||||
|
IsSuspicious bool `json:"isSuspicious"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) GetSkillMeta(ctx context.Context, slug string) (*SkillMeta, error) {
|
||||||
|
if err := utils.ValidateSkillIdentifier(slug); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid slug %q: error: %s", slug, err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
u := c.baseURL + c.skillsPath + "/" + url.PathEscape(slug)
|
||||||
|
|
||||||
|
body, err := c.doGet(ctx, u)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("skill metadata request failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp clawhubSkillResponse
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse skill metadata: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
meta := &SkillMeta{
|
||||||
|
Slug: resp.Slug,
|
||||||
|
DisplayName: resp.DisplayName,
|
||||||
|
Summary: resp.Summary,
|
||||||
|
RegistryName: c.Name(),
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.LatestVersion != nil {
|
||||||
|
meta.LatestVersion = resp.LatestVersion.Version
|
||||||
|
}
|
||||||
|
if resp.Moderation != nil {
|
||||||
|
meta.IsMalwareBlocked = resp.Moderation.IsMalwareBlocked
|
||||||
|
meta.IsSuspicious = resp.Moderation.IsSuspicious
|
||||||
|
}
|
||||||
|
|
||||||
|
return meta, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- DownloadAndInstall ---
|
||||||
|
|
||||||
|
// DownloadAndInstall fetches metadata (with fallback), resolves version,
|
||||||
|
// downloads the skill ZIP, and extracts it to targetDir.
|
||||||
|
// Returns an InstallResult for the caller to use for moderation decisions.
|
||||||
|
func (c *ClawHubRegistry) DownloadAndInstall(ctx context.Context, slug, version, targetDir string) (*InstallResult, error) {
|
||||||
|
if err := utils.ValidateSkillIdentifier(slug); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid slug %q: error: %s", slug, err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 1: Fetch metadata (with fallback).
|
||||||
|
result := &InstallResult{}
|
||||||
|
meta, err := c.GetSkillMeta(ctx, slug)
|
||||||
|
if err != nil {
|
||||||
|
// Fallback: proceed without metadata.
|
||||||
|
meta = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if meta != nil {
|
||||||
|
result.IsMalwareBlocked = meta.IsMalwareBlocked
|
||||||
|
result.IsSuspicious = meta.IsSuspicious
|
||||||
|
result.Summary = meta.Summary
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 2: Resolve version.
|
||||||
|
installVersion := version
|
||||||
|
if installVersion == "" && meta != nil {
|
||||||
|
installVersion = meta.LatestVersion
|
||||||
|
}
|
||||||
|
if installVersion == "" {
|
||||||
|
installVersion = "latest"
|
||||||
|
}
|
||||||
|
result.Version = installVersion
|
||||||
|
|
||||||
|
// Step 3: Download ZIP to temp file (streams in ~32KB chunks).
|
||||||
|
u, err := url.Parse(c.baseURL + c.downloadPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid base URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
q := u.Query()
|
||||||
|
q.Set("slug", slug)
|
||||||
|
if installVersion != "latest" {
|
||||||
|
q.Set("version", installVersion)
|
||||||
|
}
|
||||||
|
u.RawQuery = q.Encode()
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "GET", u.String(), nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
if c.authToken != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+c.authToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
tmpPath, err := utils.DownloadToFile(ctx, c.client, req, int64(c.maxZipSize))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("download failed: %w", err)
|
||||||
|
}
|
||||||
|
defer os.Remove(tmpPath)
|
||||||
|
|
||||||
|
// Step 4: Extract from file on disk.
|
||||||
|
if err := utils.ExtractZipFile(tmpPath, targetDir); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- HTTP helper ---
|
||||||
|
|
||||||
|
func (c *ClawHubRegistry) doGet(ctx context.Context, urlStr string) ([]byte, error) {
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "GET", urlStr, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Accept", "application/json")
|
||||||
|
if c.authToken != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+c.authToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := c.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
// Limit response body read to prevent memory issues.
|
||||||
|
body, err := io.ReadAll(io.LimitReader(resp.Body, int64(c.maxResponseSize)))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||||
|
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
return body, nil
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue