Merge branch 'sipeed:main' into main
This commit is contained in:
commit
79c0ce4ba8
223 changed files with 16057 additions and 4600 deletions
60
.github/workflows/create-tag.yml
vendored
Normal file
60
.github/workflows/create-tag.yml
vendored
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
name: Create Tag
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Tag name (required, e.g. v0.2.0)"
|
||||
required: true
|
||||
type: string
|
||||
commit:
|
||||
description: "Target commit SHA (leave empty for latest main)"
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
|
||||
jobs:
|
||||
create-tag:
|
||||
name: Create Git Tag
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
ref: main
|
||||
|
||||
- name: Validate commit exists
|
||||
if: ${{ inputs.commit != '' }}
|
||||
shell: bash
|
||||
run: |
|
||||
if ! git cat-file -t "${{ inputs.commit }}" &>/dev/null; then
|
||||
echo "::error::Commit '${{ inputs.commit }}' does not exist."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Check tag does not already exist
|
||||
shell: bash
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
if gh api "repos/${{ github.repository }}/git/ref/tags/${{ inputs.tag }}" --silent 2>/dev/null; then
|
||||
echo "::error::Tag '${{ inputs.tag }}' already exists."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Create and push tag
|
||||
shell: bash
|
||||
run: |
|
||||
TARGET="${{ inputs.commit || 'HEAD' }}"
|
||||
COMMIT_SHA=$(git rev-parse "$TARGET")
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
git tag -a "${{ inputs.tag }}" "$COMMIT_SHA" -m "Release ${{ inputs.tag }}"
|
||||
git push origin "${{ inputs.tag }}"
|
||||
echo "### Tag Created" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "- **Tag:** \`${{ inputs.tag }}\`" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "- **Commit:** \`${COMMIT_SHA}\`" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "- **Branch:** \`$(git branch -r --contains "$COMMIT_SHA" | head -1 | xargs)\`" >> "$GITHUB_STEP_SUMMARY"
|
||||
36
.github/workflows/release.yml
vendored
36
.github/workflows/release.yml
vendored
|
|
@ -1,10 +1,10 @@
|
|||
name: Create Tag and Release
|
||||
name: Release
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Release tag (required, e.g. v0.2.0)"
|
||||
description: "Existing tag to release (e.g. v0.2.0)"
|
||||
required: true
|
||||
type: string
|
||||
prerelease:
|
||||
|
|
@ -24,35 +24,23 @@ on:
|
|||
default: true
|
||||
|
||||
jobs:
|
||||
create-tag:
|
||||
name: Create Git Tag
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Create and push tag
|
||||
shell: bash
|
||||
env:
|
||||
RELEASE_TAG: ${{ inputs.tag }}
|
||||
run: |
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
git tag -a "$RELEASE_TAG" -m "Release $RELEASE_TAG"
|
||||
git push origin "$RELEASE_TAG"
|
||||
|
||||
release:
|
||||
name: GoReleaser Release
|
||||
needs: create-tag
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
packages: write
|
||||
steps:
|
||||
- name: Verify tag exists
|
||||
shell: bash
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
if ! gh api "repos/${{ github.repository }}/git/ref/tags/${{ inputs.tag }}" --silent 2>/dev/null; then
|
||||
echo "::error::Tag '${{ inputs.tag }}' does not exist. Create it first using the 'Create Tag' workflow."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Checkout tag
|
||||
uses: actions/checkout@v6
|
||||
with:
|
||||
|
|
|
|||
Binary file not shown.
|
Before Width: | Height: | Size: 98 KiB After Width: | Height: | Size: 356 KiB |
|
|
@ -59,7 +59,7 @@ func authLoginOpenAI(useDeviceCode bool, noBrowser bool) error {
|
|||
// Update or add openai in ModelList
|
||||
foundOpenAI := false
|
||||
for i := range appCfg.ModelList {
|
||||
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||
if isOpenAIModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||
foundOpenAI = true
|
||||
break
|
||||
|
|
@ -130,7 +130,7 @@ func authLoginGoogleAntigravity(noBrowser bool) error {
|
|||
// Update or add antigravity in ModelList
|
||||
foundAntigravity := false
|
||||
for i := range appCfg.ModelList {
|
||||
if isAntigravityModel(appCfg.ModelList[i].Model) {
|
||||
if isAntigravityModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||
foundAntigravity = true
|
||||
break
|
||||
|
|
@ -206,7 +206,7 @@ func authLoginAnthropicSetupToken() error {
|
|||
if err == nil {
|
||||
found := false
|
||||
for i := range appCfg.ModelList {
|
||||
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||
if isAnthropicModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = "oauth"
|
||||
found = true
|
||||
break
|
||||
|
|
@ -282,7 +282,7 @@ func authLoginPasteToken(provider string) error {
|
|||
// Update ModelList
|
||||
found := false
|
||||
for i := range appCfg.ModelList {
|
||||
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||
if isAnthropicModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = "token"
|
||||
found = true
|
||||
break
|
||||
|
|
@ -300,7 +300,7 @@ func authLoginPasteToken(provider string) error {
|
|||
// Update ModelList
|
||||
found := false
|
||||
for i := range appCfg.ModelList {
|
||||
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||
if isOpenAIModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = "token"
|
||||
found = true
|
||||
break
|
||||
|
|
@ -342,15 +342,15 @@ func authLogoutCmd(provider string) error {
|
|||
for i := range appCfg.ModelList {
|
||||
switch provider {
|
||||
case "openai":
|
||||
if isOpenAIModel(appCfg.ModelList[i].Model) {
|
||||
if isOpenAIModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = ""
|
||||
}
|
||||
case "anthropic":
|
||||
if isAnthropicModel(appCfg.ModelList[i].Model) {
|
||||
if isAnthropicModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = ""
|
||||
}
|
||||
case "google-antigravity", "antigravity":
|
||||
if isAntigravityModel(appCfg.ModelList[i].Model) {
|
||||
if isAntigravityModel(appCfg.ModelList[i]) {
|
||||
appCfg.ModelList[i].AuthMethod = ""
|
||||
}
|
||||
}
|
||||
|
|
@ -484,22 +484,20 @@ func authModelsCmd() error {
|
|||
return nil
|
||||
}
|
||||
|
||||
// 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/")
|
||||
// isAntigravityModel checks if a model config belongs to an Antigravity provider.
|
||||
func isAntigravityModel(modelCfg *config.ModelConfig) bool {
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
return protocol == "antigravity" || protocol == "google-antigravity"
|
||||
}
|
||||
|
||||
// isOpenAIModel checks if a model string belongs to openai provider
|
||||
func isOpenAIModel(model string) bool {
|
||||
return model == "openai" ||
|
||||
strings.HasPrefix(model, "openai/")
|
||||
// isOpenAIModel checks if a model config belongs to the OpenAI provider.
|
||||
func isOpenAIModel(modelCfg *config.ModelConfig) bool {
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
return protocol == "openai"
|
||||
}
|
||||
|
||||
// isAnthropicModel checks if a model string belongs to anthropic provider
|
||||
func isAnthropicModel(model string) bool {
|
||||
return model == "anthropic" ||
|
||||
strings.HasPrefix(model, "anthropic/")
|
||||
// isAnthropicModel checks if a model config belongs to the Anthropic provider.
|
||||
func isAnthropicModel(modelCfg *config.ModelConfig) bool {
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
return protocol == "anthropic"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,53 @@
|
|||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
pkgauth "github.com/sipeed/picoclaw/pkg/auth"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
func captureAuthStdout(t *testing.T, fn func()) string {
|
||||
t.Helper()
|
||||
|
||||
oldStdout := os.Stdout
|
||||
r, w, err := os.Pipe()
|
||||
require.NoError(t, err)
|
||||
os.Stdout = w
|
||||
t.Cleanup(func() {
|
||||
os.Stdout = oldStdout
|
||||
})
|
||||
|
||||
fn()
|
||||
|
||||
require.NoError(t, w.Close())
|
||||
os.Stdout = oldStdout
|
||||
|
||||
var buf bytes.Buffer
|
||||
_, err = io.Copy(&buf, r)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, r.Close())
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func setAuthStatusTestHome(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv(config.EnvHome, filepath.Join(tmpDir, ".picoclaw"))
|
||||
return tmpDir
|
||||
}
|
||||
|
||||
func TestNewStatusSubcommand(t *testing.T) {
|
||||
cmd := newStatusCommand()
|
||||
|
||||
|
|
@ -16,3 +57,47 @@ func TestNewStatusSubcommand(t *testing.T) {
|
|||
|
||||
assert.False(t, cmd.HasFlags())
|
||||
}
|
||||
|
||||
func TestAuthStatusCmdShowsCanonicalGoogleAntigravityAfterLegacyRefresh(t *testing.T) {
|
||||
tmpDir := setAuthStatusTestHome(t)
|
||||
|
||||
legacyExpiry := time.Date(2026, 4, 16, 10, 0, 0, 0, time.UTC)
|
||||
legacyStore := map[string]any{
|
||||
"credentials": map[string]any{
|
||||
"antigravity": map[string]any{
|
||||
"access_token": "legacy-token",
|
||||
"expires_at": legacyExpiry.Format(time.RFC3339),
|
||||
"provider": "antigravity",
|
||||
"auth_method": "oauth",
|
||||
"project_id": "legacy-project",
|
||||
},
|
||||
},
|
||||
}
|
||||
data, err := json.Marshal(legacyStore)
|
||||
require.NoError(t, err)
|
||||
|
||||
authPath := filepath.Join(tmpDir, ".picoclaw", "auth.json")
|
||||
require.NoError(t, os.MkdirAll(filepath.Dir(authPath), 0o755))
|
||||
require.NoError(t, os.WriteFile(authPath, data, 0o600))
|
||||
|
||||
refreshedExpiry := time.Date(2026, 4, 16, 12, 30, 0, 0, time.UTC)
|
||||
err = pkgauth.SetCredential("google-antigravity", &pkgauth.AuthCredential{
|
||||
AccessToken: "fresh-token",
|
||||
ExpiresAt: refreshedExpiry,
|
||||
Provider: "google-antigravity",
|
||||
AuthMethod: "oauth",
|
||||
ProjectID: "fresh-project",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
output := captureAuthStdout(t, func() {
|
||||
require.NoError(t, authStatusCmd())
|
||||
})
|
||||
|
||||
assert.Contains(t, output, "\nAuthenticated Providers:")
|
||||
assert.Contains(t, output, "\n google-antigravity:\n")
|
||||
assert.NotContains(t, output, "\n antigravity:\n")
|
||||
assert.Contains(t, output, " Project: fresh-project")
|
||||
assert.Contains(t, output, " Expires: 2026-04-16 12:30")
|
||||
assert.Equal(t, 1, strings.Count(output, ":\n Method: oauth"))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ package auth
|
|||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
|
|
@ -19,6 +20,19 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
func newIPv4TestServer(t *testing.T, handler http.Handler) *httptest.Server {
|
||||
t.Helper()
|
||||
|
||||
server := httptest.NewUnstartedServer(handler)
|
||||
listener, err := net.Listen("tcp4", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
server.Listener = listener
|
||||
server.Start()
|
||||
t.Cleanup(server.Close)
|
||||
return server
|
||||
}
|
||||
|
||||
func TestNewWeComCommand(t *testing.T) {
|
||||
cmd := newWeComCommand()
|
||||
|
||||
|
|
@ -53,7 +67,7 @@ func TestBuildWeComQRCodePageURL(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestFetchWeComQRCode(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
server := newIPv4TestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "/generate", r.URL.Path)
|
||||
assert.Equal(t, wecomQRSourceID, r.URL.Query().Get("source"))
|
||||
assert.Equal(t, wecomQRSourceID, r.URL.Query().Get("sourceID"))
|
||||
|
|
@ -61,7 +75,6 @@ func TestFetchWeComQRCode(t *testing.T) {
|
|||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":{"scode":"scode-1","auth_url":"https://example.com/qr"}}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
opts := normalizeWeComQRFlowOptions(wecomQRFlowOptions{
|
||||
HTTPClient: server.Client(),
|
||||
|
|
@ -78,7 +91,7 @@ func TestFetchWeComQRCode(t *testing.T) {
|
|||
func TestPollWeComQRCodeResult(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
server := newIPv4TestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
call := calls.Add(1)
|
||||
assert.Equal(t, "/query", r.URL.Path)
|
||||
assert.Equal(t, "scode-1", r.URL.Query().Get("scode"))
|
||||
|
|
@ -92,7 +105,6 @@ func TestPollWeComQRCodeResult(t *testing.T) {
|
|||
_, _ = w.Write([]byte(`{"data":{"status":"success","bot_info":{"botid":"bot-1","secret":"secret-1"}}}`))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
var output bytes.Buffer
|
||||
opts := normalizeWeComQRFlowOptions(wecomQRFlowOptions{
|
||||
|
|
|
|||
|
|
@ -3,12 +3,12 @@ package status
|
|||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/cliui"
|
||||
"github.com/sipeed/picoclaw/pkg/auth"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
func statusCmd() {
|
||||
|
|
@ -44,12 +44,13 @@ func statusCmd() {
|
|||
// not depend on a legacy cfg.Providers field (which may not exist under some
|
||||
// build tags). We infer provider availability from model_list entries.
|
||||
hasProtocolKey := func(protocol string) bool {
|
||||
prefix := protocol + "/"
|
||||
want := providers.NormalizeProvider(protocol)
|
||||
for _, m := range cfg.ModelList {
|
||||
if m == nil {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(m.Model, prefix) && m.APIKey() != "" {
|
||||
got, _ := providers.ExtractProtocol(m)
|
||||
if got == want && m.APIKey() != "" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
|
@ -67,12 +68,13 @@ func statusCmd() {
|
|||
return "", false
|
||||
}
|
||||
findProtocolBase := func(protocol string) (string, bool) {
|
||||
prefix := protocol + "/"
|
||||
want := providers.NormalizeProvider(protocol)
|
||||
for _, m := range cfg.ModelList {
|
||||
if m == nil {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(m.Model, prefix) && m.APIBase != "" {
|
||||
got, _ := providers.ExtractProtocol(m)
|
||||
if got == want && m.APIBase != "" {
|
||||
return m.APIBase, true
|
||||
}
|
||||
}
|
||||
|
|
|
|||
89
cmd/picoclaw/internal/status/helpers_test.go
Normal file
89
cmd/picoclaw/internal/status/helpers_test.go
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
package status
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
func captureStdout(t *testing.T, fn func()) string {
|
||||
t.Helper()
|
||||
|
||||
oldStdout := os.Stdout
|
||||
r, w, err := os.Pipe()
|
||||
if err != nil {
|
||||
t.Fatalf("os.Pipe() error = %v", err)
|
||||
}
|
||||
os.Stdout = w
|
||||
|
||||
fn()
|
||||
|
||||
_ = w.Close()
|
||||
os.Stdout = oldStdout
|
||||
defer r.Close()
|
||||
|
||||
var buf bytes.Buffer
|
||||
if _, err := io.Copy(&buf, r); err != nil {
|
||||
t.Fatalf("io.Copy() error = %v", err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func TestStatusCmd_RecognizesProviderFieldWithoutModelPrefix(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
configPath := filepath.Join(tmpDir, "config.json")
|
||||
workspace := filepath.Join(tmpDir, "workspace")
|
||||
if err := os.MkdirAll(workspace, 0o755); err != nil {
|
||||
t.Fatalf("os.MkdirAll() error = %v", err)
|
||||
}
|
||||
|
||||
t.Setenv(config.EnvConfig, configPath)
|
||||
t.Setenv(config.EnvHome, tmpDir)
|
||||
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
ModelName: "gpt-5.4",
|
||||
Workspace: workspace,
|
||||
Provider: "openai",
|
||||
MaxTokens: 65536,
|
||||
Temperature: nil,
|
||||
},
|
||||
},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "gpt-5.4",
|
||||
Provider: "openai",
|
||||
Model: "gpt-5.4",
|
||||
APIBase: "https://api.openai.com/v1",
|
||||
APIKeys: config.SimpleSecureStrings("test-key"),
|
||||
Enabled: true,
|
||||
},
|
||||
{
|
||||
ModelName: "qwen-plus",
|
||||
Provider: "qwen",
|
||||
Model: "qwen-plus",
|
||||
APIBase: "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
APIKeys: config.SimpleSecureStrings("test-key"),
|
||||
Enabled: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||
t.Fatalf("config.SaveConfig() error = %v", err)
|
||||
}
|
||||
|
||||
output := captureStdout(t, statusCmd)
|
||||
|
||||
if !strings.Contains(output, "OpenAI API: \u2713") {
|
||||
t.Fatalf("status output missing OpenAI provider: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "Qwen API: \u2713") {
|
||||
t.Fatalf("status output missing Qwen provider: %s", output)
|
||||
}
|
||||
}
|
||||
100
docs/architecture/agent-refactor/agent-rename-plan.md
Normal file
100
docs/architecture/agent-refactor/agent-rename-plan.md
Normal file
|
|
@ -0,0 +1,100 @@
|
|||
# Agent File Rename Plan
|
||||
|
||||
## Goal
|
||||
|
||||
Unify `pkg/agent/` package file naming to resolve the `loop_*` prefix naming confusion and unclear responsibility boundaries.
|
||||
|
||||
## Change Overview
|
||||
|
||||
### File Renames (12 files)
|
||||
|
||||
| Original | New | Description |
|
||||
|----------|-----|-------------|
|
||||
| `loop.go` | `agent.go` | AgentLoop main body + lifecycle methods |
|
||||
| `loop_message.go` | `agent_message.go` | Message handling and routing |
|
||||
| `loop_outbound.go` | `agent_outbound.go` | Response publishing |
|
||||
| `loop_event.go` | `agent_event.go` | Event system |
|
||||
| `loop_command.go` | `agent_command.go` | Command processing |
|
||||
| `loop_steering.go` | `agent_steering.go` | Steering message handling |
|
||||
| `loop_transcribe.go` | `agent_transcribe.go` | Audio transcription |
|
||||
| `loop_media.go` | `agent_media.go` | Media processing |
|
||||
| `loop_mcp.go` | `agent_mcp.go` | MCP initialization |
|
||||
| `loop_utils.go` | `agent_utils.go` | Utility functions |
|
||||
| `loop_inject.go` | `agent_inject.go` | Dependency injection |
|
||||
| `loop_turn.go` | `turn_coord.go` | Turn coordinator |
|
||||
|
||||
### File Merges (2 → 1)
|
||||
|
||||
| Original | New | Description |
|
||||
|----------|-----|-------------|
|
||||
| `turn.go` + `turn_exec.go` | `turn_state.go` | Turn-related type definitions |
|
||||
|
||||
## Final File Structure
|
||||
|
||||
```
|
||||
pkg/agent/
|
||||
├── agent.go # AgentLoop + Run/Stop/Close lifecycle
|
||||
├── agent_message.go # Message processing
|
||||
├── agent_outbound.go # Response publishing
|
||||
├── agent_event.go # Event system
|
||||
├── agent_command.go # Command processing
|
||||
├── agent_steering.go # Steering
|
||||
├── agent_transcribe.go # Transcription
|
||||
├── agent_media.go # Media processing
|
||||
├── agent_mcp.go # MCP
|
||||
├── agent_utils.go # Utility functions
|
||||
├── agent_inject.go # Dependency injection
|
||||
├── turn_coord.go # runTurn + coordinator
|
||||
├── turn_state.go # turnState + turnExecution + Control + ToolControl + LLMPhase
|
||||
├── pipeline.go # Pipeline struct + NewPipeline
|
||||
├── pipeline_setup.go
|
||||
├── pipeline_llm.go
|
||||
├── pipeline_execute.go
|
||||
└── pipeline_finalize.go
|
||||
```
|
||||
|
||||
## Naming Convention
|
||||
|
||||
| Prefix | Content | Example |
|
||||
|--------|---------|---------|
|
||||
| `agent_*` | AgentLoop method files | `agent_message.go`, `agent_event.go` |
|
||||
| `turn_*` | Turn lifecycle related | `turn_coord.go`, `turn_state.go` |
|
||||
| `pipeline_*` | Pipeline methods | `pipeline_setup.go`, `pipeline_llm.go` |
|
||||
| `context_*` | Context management | `context_manager.go`, `context_legacy.go` |
|
||||
| `hook_*` | Hook system | `hook_process.go`, `hook_mount.go` |
|
||||
|
||||
## Architecture Layers
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ AgentLoop (agent.go) │
|
||||
│ - Message loop Run/Stop/Close │
|
||||
│ - Dependency injection (agent_inject.go) │
|
||||
│ - Message routing (agent_message.go) │
|
||||
│ - Response publishing (agent_outbound.go) │
|
||||
└─────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ Turn Coordinator (turn_coord.go) │
|
||||
│ - runTurn(): main coordinator │
|
||||
│ - abortTurn(): abort │
|
||||
│ - askSideQuestion(): side question │
|
||||
│ - selectCandidates(): model selection │
|
||||
└─────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ Pipeline (pipeline_*.go) │
|
||||
│ - SetupTurn(): initialization │
|
||||
│ - CallLLM(): LLM call │
|
||||
│ - ExecuteTools(): tool execution │
|
||||
│ - Finalize(): finalization │
|
||||
└─────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
## Verification Results
|
||||
|
||||
- ✅ `go build ./pkg/agent/...` - Pass
|
||||
- ✅ `go vet ./pkg/agent/...` - No warnings
|
||||
- ✅ `go test ./pkg/agent/... -skip "TestSeahorse|TestGlobalSkillFileContentChange"` - Pass
|
||||
100
docs/architecture/agent-refactor/agent-rename-plan.zh.md
Normal file
100
docs/architecture/agent-refactor/agent-rename-plan.zh.md
Normal file
|
|
@ -0,0 +1,100 @@
|
|||
# Agent 文件重命名计划
|
||||
|
||||
## 目标
|
||||
|
||||
统一 `pkg/agent/` 包的文件命名,解决 `loop_*` 前缀命名混乱、职责边界不清晰的问题。
|
||||
|
||||
## 变更概览
|
||||
|
||||
### 文件重命名(12 个)
|
||||
|
||||
| 原文件 | 新文件 | 说明 |
|
||||
|--------|--------|------|
|
||||
| `loop.go` | `agent.go` | AgentLoop 主体 + 生命周期方法 |
|
||||
| `loop_message.go` | `agent_message.go` | 消息处理和路由 |
|
||||
| `loop_outbound.go` | `agent_outbound.go` | 响应发布 |
|
||||
| `loop_event.go` | `agent_event.go` | 事件系统 |
|
||||
| `loop_command.go` | `agent_command.go` | 命令处理 |
|
||||
| `loop_steering.go` | `agent_steering.go` | Steering 消息处理 |
|
||||
| `loop_transcribe.go` | `agent_transcribe.go` | 音频转录 |
|
||||
| `loop_media.go` | `agent_media.go` | 媒体处理 |
|
||||
| `loop_mcp.go` | `agent_mcp.go` | MCP 初始化 |
|
||||
| `loop_utils.go` | `agent_utils.go` | 工具函数 |
|
||||
| `loop_inject.go` | `agent_inject.go` | 依赖注入 |
|
||||
| `loop_turn.go` | `turn_coord.go` | Turn 协调器 |
|
||||
|
||||
### 文件合并(2 → 1)
|
||||
|
||||
| 原文件 | 新文件 | 说明 |
|
||||
|--------|--------|------|
|
||||
| `turn.go` + `turn_exec.go` | `turn_state.go` | Turn 相关类型定义 |
|
||||
|
||||
## 最终文件结构
|
||||
|
||||
```
|
||||
pkg/agent/
|
||||
├── agent.go # AgentLoop + Run/Stop/Close 生命周期
|
||||
├── agent_message.go # 消息处理
|
||||
├── agent_outbound.go # 响应发布
|
||||
├── agent_event.go # 事件系统
|
||||
├── agent_command.go # 命令处理
|
||||
├── agent_steering.go # Steering
|
||||
├── agent_transcribe.go # 转录
|
||||
├── agent_media.go # 媒体处理
|
||||
├── agent_mcp.go # MCP
|
||||
├── agent_utils.go # 工具函数
|
||||
├── agent_inject.go # 依赖注入
|
||||
├── turn_coord.go # runTurn + 协调器
|
||||
├── turn_state.go # turnState + turnExecution + Control + ToolControl + LLMPhase
|
||||
├── pipeline.go # Pipeline struct + NewPipeline
|
||||
├── pipeline_setup.go
|
||||
├── pipeline_llm.go
|
||||
├── pipeline_execute.go
|
||||
└── pipeline_finalize.go
|
||||
```
|
||||
|
||||
## 命名约定
|
||||
|
||||
| 前缀 | 内容 | 示例 |
|
||||
|------|------|------|
|
||||
| `agent_*` | AgentLoop 的方法文件 | `agent_message.go`, `agent_event.go` |
|
||||
| `turn_*` | Turn 生命周期相关 | `turn_coord.go`, `turn_state.go` |
|
||||
| `pipeline_*` | Pipeline 方法 | `pipeline_setup.go`, `pipeline_llm.go` |
|
||||
| `context_*` | 上下文管理 | `context_manager.go`, `context_legacy.go` |
|
||||
| `hook_*` | Hook 系统 | `hook_process.go`, `hook_mount.go` |
|
||||
|
||||
## 架构层次
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ AgentLoop (agent.go) │
|
||||
│ - 消息循环 Run/Stop/Close │
|
||||
│ - 依赖注入 (agent_inject.go) │
|
||||
│ - 消息路由 (agent_message.go) │
|
||||
│ - 响应发布 (agent_outbound.go) │
|
||||
└─────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ Turn Coordinator (turn_coord.go) │
|
||||
│ - runTurn(): 主协调器 │
|
||||
│ - abortTurn(): 中止 │
|
||||
│ - askSideQuestion(): 侧问 │
|
||||
│ - selectCandidates(): 模型选择 │
|
||||
└─────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ Pipeline (pipeline_*.go) │
|
||||
│ - SetupTurn(): 初始化 │
|
||||
│ - CallLLM(): LLM 调用 │
|
||||
│ - ExecuteTools(): 工具执行 │
|
||||
│ - Finalize(): 终结 │
|
||||
└─────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
## 验证结果
|
||||
|
||||
- ✅ `go build ./pkg/agent/...` - 通过
|
||||
- ✅ `go vet ./pkg/agent/...` - 无警告
|
||||
- ✅ `go test ./pkg/agent/... -skip "TestSeahorse|TestGlobalSkillFileContentChange"` - 通过
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
# AgentLoop File Split
|
||||
|
||||
> **Note:** This document describes the file split that was completed in a previous phase. The `loop_*` naming has since been renamed to `agent_*` and `turn_*`. See [agent-rename-plan.md](./agent-rename-plan.md) for the current file structure.
|
||||
|
||||
## Overview
|
||||
|
||||
The `pkg/agent/loop.go` file (originally 4384 lines) has been split into 12 focused source files. This is a pure refactoring with no behavioral changes.
|
||||
|
|
@ -11,76 +13,65 @@ The `pkg/agent/loop.go` file (originally 4384 lines) has been split into 12 focu
|
|||
- Maintain all existing functionality and tests
|
||||
- Keep imports minimal per file
|
||||
|
||||
## File Map
|
||||
## Original File Map (Renamed in Phase 2)
|
||||
|
||||
| File | Lines | Responsibility |
|
||||
|------|-------|----------------|
|
||||
| `loop.go` | ~650 | Core `AgentLoop` struct, `Run`, `Stop`, `Close`, `ReloadProviderAndConfig`, `runAgentLoop` |
|
||||
| `loop_turn.go` | ~1880 | Turn execution: `runTurn`, `abortTurn`, `selectCandidates`, `askSideQuestion`, `isolatedSideQuestionProvider`, side question model config |
|
||||
| `loop_utils.go` | ~480 | Standalone utility functions: formatters, cloners, helpers (no receiver) |
|
||||
| `loop_init.go` | ~355 | `NewAgentLoop` constructor and `registerSharedTools` |
|
||||
| `loop_message.go` | ~300 | Message handling: `processMessage`, `processSystemMessage`, routing helpers, `ProcessDirect`, `ProcessHeartbeat` |
|
||||
| `loop_command.go` | ~265 | Command processing: `handleCommand`, `applyExplicitSkillCommand`, pending skills management |
|
||||
| `loop_mcp.go` | ~235 | MCP runtime: `ensureMCPInitialized`, server discovery, deferred server handling |
|
||||
| `loop_event.go` | ~205 | Event system helpers: `emitEvent`, `logEvent`, `hookAbortError`, `newTurnEventScope`, `MountHook`, `SubscribeEvents` |
|
||||
| `loop_media.go` | ~198 | Media resolution: `resolveMediaRefs`, artifact building, MIME detection |
|
||||
| `loop_outbound.go` | ~165 | Response publishing: `PublishResponseIfNeeded`, `publishPicoReasoning`, `handleReasoning` |
|
||||
| `loop_transcribe.go` | ~110 | Audio transcription: `transcribeAudioInMessage`, `sendTranscriptionFeedback` |
|
||||
| `loop_steering.go` | ~97 | Steering queue: `runTurnWithSteering`, `processMessageSync`, `resolveSteeringTarget` |
|
||||
| `loop_inject.go` | ~104 | Setter injection: `SetChannelManager`, `SetMediaStore`, `SetTranscriber`, `GetRegistry`, `GetConfig`, `RecordLastChannel` |
|
||||
| Old File | New File | Responsibility |
|
||||
|----------|----------|----------------|
|
||||
| `loop.go` | `agent.go` | Core `AgentLoop` struct, `Run`, `Stop`, `Close` |
|
||||
| `loop_turn.go` | `turn_coord.go` + `pipeline_*.go` | Turn execution: coordinator + Pipeline methods |
|
||||
| `loop_utils.go` | `agent_utils.go` | Standalone utility functions |
|
||||
| `loop_init.go` | `agent_init.go` | `NewAgentLoop` constructor and tool registration |
|
||||
| `loop_message.go` | `agent_message.go` | Message handling and routing |
|
||||
| `loop_command.go` | `agent_command.go` | Command processing |
|
||||
| `loop_mcp.go` | `agent_mcp.go` | MCP runtime |
|
||||
| `loop_event.go` | `agent_event.go` | Event system helpers |
|
||||
| `loop_media.go` | `agent_media.go` | Media resolution |
|
||||
| `loop_outbound.go` | `agent_outbound.go` | Response publishing |
|
||||
| `loop_transcribe.go` | `agent_transcribe.go` | Audio transcription |
|
||||
| `loop_steering.go` | `agent_steering.go` | Steering queue |
|
||||
| `loop_inject.go` | `agent_inject.go` | Setter injection |
|
||||
|
||||
## Current File Structure
|
||||
|
||||
See [agent-rename-plan.md](./agent-rename-plan.md) for the complete current file structure.
|
||||
|
||||
## Phase 2: Rename and Pipeline Restructuring
|
||||
|
||||
Phase 2 completed the following:
|
||||
|
||||
1. **File renaming**: All `loop_*` files renamed to `agent_*` or `turn_*`
|
||||
2. **Turn state merging**: `turn.go` + `turn_exec.go` → `turn_state.go`
|
||||
3. **Pipeline extraction**: Split large `runTurn` into Pipeline methods
|
||||
|
||||
### Pipeline Architecture
|
||||
|
||||
The Pipeline methods provide structured turn execution:
|
||||
|
||||
| Method | File | Responsibility |
|
||||
|--------|------|----------------|
|
||||
| `SetupTurn()` | `pipeline_setup.go` | History assembly, message building, candidate selection |
|
||||
| `CallLLM()` | `pipeline_llm.go` | PreLLM hooks, fallback, retry, AfterLLM hooks |
|
||||
| `ExecuteTools()` | `pipeline_execute.go` | Tool execution with hooks |
|
||||
| `Finalize()` | `pipeline_finalize.go` | Session persistence, compression |
|
||||
|
||||
## Core Principles Applied
|
||||
|
||||
### 1. Same Package, Independent Files
|
||||
All files belong to the `agent` package and compile together. This preserves the original visibility rules — no interface abstraction was introduced in this phase.
|
||||
All files belong to the `agent` package and compile together. This preserves the original visibility rules.
|
||||
|
||||
### 2. No Logic Changes
|
||||
All functions were moved verbatim (except updating import statements). The extraction script used the original `loop.go.backup` as source of truth to ensure no drift.
|
||||
All functions were moved verbatim. The extraction preserved behavioral equivalence.
|
||||
|
||||
### 3. Shared Types Remain in loop.go
|
||||
The `AgentLoop` struct, `processOptions`, `continuationTarget`, and all hook/event types stay in `loop.go` since they are referenced across files.
|
||||
|
||||
### 4. Turn State Is Central
|
||||
`loop_turn.go` is the largest file because the turn lifecycle (`runTurn`) is inherently large. It contains the core LLM interaction loop, tool execution, subturn spawning, and steering injection.
|
||||
|
||||
## What's Left in loop.go
|
||||
|
||||
```go
|
||||
// Core struct
|
||||
type AgentLoop struct { ... }
|
||||
|
||||
// Main lifecycle
|
||||
func (al *AgentLoop) Run(ctx context.Context) error
|
||||
func (al *AgentLoop) Stop()
|
||||
func (al *AgentLoop) Close()
|
||||
func (al *AgentLoop) ReloadProviderAndConfig(ctx, provider, cfg)
|
||||
|
||||
// Turn orchestration (calls into loop_turn.go)
|
||||
func (al *AgentLoop) runAgentLoop(ctx, agent, opts) (string, error)
|
||||
```
|
||||
|
||||
## Extraction Method
|
||||
|
||||
The split was done programmatically using Node.js to:
|
||||
1. Identify function boundaries using brace counting
|
||||
2. Extract each function to its target file
|
||||
3. Add necessary imports to each file
|
||||
4. Remove the extracted function from loop.go
|
||||
5. Run `go fmt` and `go vet` to verify
|
||||
### 3. Shared Types in turn_state.go
|
||||
The `turnState`, `turnExecution`, `Control`, `ToolControl`, and `LLMPhase` types are centralized in `turn_state.go`.
|
||||
|
||||
## Testing
|
||||
|
||||
All existing tests pass. The 5 failing tests (`TestGlobalSkillFileContentChange` and 4 Seahorse tests) are pre-existing failures unrelated to this refactor (database file locking issues on Windows).
|
||||
All existing tests pass. The 5 failing tests (`TestGlobalSkillFileContentChange` and 4 Seahorse tests) are pre-existing failures unrelated to this refactor.
|
||||
|
||||
Build status: `go build ./pkg/agent/...` passes with no errors.
|
||||
|
||||
## Phase 2: Dependency Inversion (Planned)
|
||||
|
||||
A future phase will introduce interface types to decouple `AgentLoop` from its dependencies, enabling:
|
||||
- Easier testing with mock dependencies
|
||||
- Alternative runtime configurations
|
||||
- Cleaner boundaries for MCP and other extensions
|
||||
|
||||
## See Also
|
||||
|
||||
- [agent-rename-plan.md](./agent-rename-plan.md) — Current file naming convention
|
||||
- [context.md](context.md) — context management and session handling
|
||||
|
|
|
|||
|
|
@ -0,0 +1,68 @@
|
|||
# Pipeline Restructuring Plan
|
||||
|
||||
## Goal
|
||||
|
||||
Split `agent/pipeline.go` (~1400 lines) into multiple logical files, organizing code by responsibility.
|
||||
|
||||
## Final File Structure
|
||||
|
||||
```
|
||||
pkg/agent/
|
||||
├── pipeline.go # Pipeline struct + NewPipeline (~39 lines)
|
||||
├── pipeline_setup.go # SetupTurn method (~115 lines)
|
||||
├── pipeline_llm.go # CallLLM method (~519 lines)
|
||||
├── pipeline_execute.go # ExecuteTools method (~693 lines)
|
||||
└── pipeline_finalize.go # Finalize method (~78 lines)
|
||||
```
|
||||
|
||||
## Actual Line Counts
|
||||
|
||||
| File | Lines |
|
||||
|------|-------|
|
||||
| `pipeline.go` | 39 |
|
||||
| `pipeline_setup.go` | 115 |
|
||||
| `pipeline_llm.go` | 519 |
|
||||
| `pipeline_execute.go` | 693 |
|
||||
| `pipeline_finalize.go` | 78 |
|
||||
| **Total** | **1444** |
|
||||
|
||||
## Responsibility Matrix
|
||||
|
||||
| File | Method | Responsibility |
|
||||
|------|--------|----------------|
|
||||
| `pipeline.go` | `Pipeline` struct, `NewPipeline()` | Pipeline dependency container |
|
||||
| `pipeline_setup.go` | `SetupTurn()` | Turn initialization: history assembly, message building, candidate selection |
|
||||
| `pipeline_llm.go` | `CallLLM()` | LLM call: PreLLM hooks, fallback, retry, AfterLLM hooks |
|
||||
| `pipeline_execute.go` | `ExecuteTools()` | Tool execution: BeforeTool/ApproveTool/AfterTool hooks, media sending, steering handling |
|
||||
| `pipeline_finalize.go` | `Finalize()` | Turn finalization: session save, compression, status setting |
|
||||
|
||||
## Relationship Between Pipeline and Turn Coordinator
|
||||
|
||||
```
|
||||
AgentLoop (agent.go)
|
||||
│
|
||||
├── runAgentLoop() ──────────────────┐
|
||||
│ │
|
||||
│ ┌───────────────────────────────▼───────────────────────────────┐
|
||||
│ │ Turn Coordinator (turn_coord.go) │
|
||||
│ │ │
|
||||
│ │ runTurn() { │
|
||||
│ │ exec = pipeline.SetupTurn() │
|
||||
│ │ loop { │
|
||||
│ │ ctrl = pipeline.CallLLM() ──► Pipeline (pipeline_*.go) │
|
||||
│ │ if ctrl == ToolLoop { │
|
||||
│ │ toolCtrl = pipeline.ExecuteTools() │
|
||||
│ │ } │
|
||||
│ │ } │
|
||||
│ │ return pipeline.Finalize() │
|
||||
│ │ } │
|
||||
│ └─────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
└── Publish response (agent_outbound.go)
|
||||
```
|
||||
|
||||
## Verification Results
|
||||
|
||||
- ✅ `go build ./pkg/agent/...` - Pass
|
||||
- ✅ `go vet ./pkg/agent/...` - No warnings
|
||||
- ✅ `go test ./pkg/agent/... -skip "TestSeahorse|TestGlobalSkillFileContentChange"` - Pass
|
||||
|
|
@ -0,0 +1,68 @@
|
|||
# Pipeline 重构文档
|
||||
|
||||
## 目标
|
||||
|
||||
将 `agent/pipeline.go` (1400行) 拆分为多个逻辑文件,代码按职责组织。
|
||||
|
||||
## 最终文件结构
|
||||
|
||||
```
|
||||
pkg/agent/
|
||||
├── pipeline.go # Pipeline struct + NewPipeline (~39行)
|
||||
├── pipeline_setup.go # SetupTurn 方法 (~115行)
|
||||
├── pipeline_llm.go # CallLLM 方法 (~519行)
|
||||
├── pipeline_execute.go # ExecuteTools 方法 (~693行)
|
||||
└── pipeline_finalize.go # Finalize 方法 (~78行)
|
||||
```
|
||||
|
||||
## 实际行数
|
||||
|
||||
| 文件 | 行数 |
|
||||
|------|------|
|
||||
| `pipeline.go` | 39 |
|
||||
| `pipeline_setup.go` | 115 |
|
||||
| `pipeline_llm.go` | 519 |
|
||||
| `pipeline_execute.go` | 693 |
|
||||
| `pipeline_finalize.go` | 78 |
|
||||
| **总计** | **1444** |
|
||||
|
||||
## 职责说明
|
||||
|
||||
| 文件 | 方法 | 职责 |
|
||||
|------|------|------|
|
||||
| `pipeline.go` | `Pipeline` struct, `NewPipeline()` | Pipeline 依赖容器 |
|
||||
| `pipeline_setup.go` | `SetupTurn()` | Turn 初始化:历史组装、消息构建、候选人选择 |
|
||||
| `pipeline_llm.go` | `CallLLM()` | LLM 调用:PreLLM hook、fallback、重试、AfterLLM hook |
|
||||
| `pipeline_execute.go` | `ExecuteTools()` | 工具执行:BeforeTool/ApproveTool/AfterTool hook、媒体发送、steering 处理 |
|
||||
| `pipeline_finalize.go` | `Finalize()` | Turn 终结:会话保存、压缩、状态设置 |
|
||||
|
||||
## Pipeline 与 Turn Coordinator 的关系
|
||||
|
||||
```
|
||||
AgentLoop (agent.go)
|
||||
│
|
||||
├── runAgentLoop() ──────────────────┐
|
||||
│ │
|
||||
│ ┌───────────────────────────────▼───────────────────────────────┐
|
||||
│ │ Turn Coordinator (turn_coord.go) │
|
||||
│ │ │
|
||||
│ │ runTurn() { │
|
||||
│ │ exec = pipeline.SetupTurn() │
|
||||
│ │ loop { │
|
||||
│ │ ctrl = pipeline.CallLLM() ──► Pipeline (pipeline_*.go) │
|
||||
│ │ if ctrl == ToolLoop { │
|
||||
│ │ toolCtrl = pipeline.ExecuteTools() │
|
||||
│ │ } │
|
||||
│ │ } │
|
||||
│ │ return pipeline.Finalize() │
|
||||
│ │ } │
|
||||
│ └─────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
└── 发布响应 (agent_outbound.go)
|
||||
```
|
||||
|
||||
## 验证结果
|
||||
|
||||
- ✅ `go build ./pkg/agent/...` - 通过
|
||||
- ✅ `go vet ./pkg/agent/...` - 无警告
|
||||
- ✅ `go test ./pkg/agent/... -skip "TestSeahorse|TestGlobalSkillFileContentChange"` - 通过
|
||||
|
|
@ -19,7 +19,7 @@ It does not describe the launcher's HTTP `ServeMux` routes or the frontend's Tan
|
|||
| Agent dispatch | `pkg/routing/route.go`, `pkg/routing/agent_id.go` | Choose the target agent for the inbound message. |
|
||||
| Session policy selection | `pkg/routing/route.go` | Decide which dimensions should define session isolation for that routed turn. |
|
||||
| Model routing | `pkg/routing/router.go`, `pkg/routing/features.go`, `pkg/routing/classifier.go` | Choose between the primary model and a configured light model based on message complexity. |
|
||||
| Runtime integration | `pkg/agent/registry.go`, `pkg/agent/loop_message.go`, `pkg/agent/loop_turn.go` | Apply the route result, allocate session scope, and select model candidates before provider execution. |
|
||||
| Runtime integration | `pkg/agent/registry.go`, `pkg/agent/agent_message.go`, `pkg/agent/turn_coord.go` | Apply the route result, allocate session scope, and select model candidates before provider execution. |
|
||||
|
||||
## End-To-End Flow
|
||||
|
||||
|
|
@ -242,8 +242,8 @@ That makes the following behavior intentional:
|
|||
Agent dispatch and model routing happen in different places:
|
||||
|
||||
- `pkg/agent/registry.go` owns `RouteResolver`
|
||||
- `pkg/agent/loop_message.go` resolves the route and allocates session scope
|
||||
- `pkg/agent/loop_turn.go:selectCandidates` calls `agent.Router.SelectModel(...)`
|
||||
- `pkg/agent/agent_message.go` resolves the route and allocates session scope
|
||||
- `pkg/agent/turn_coord.go:selectCandidates` calls `agent.Router.SelectModel(...)`
|
||||
|
||||
When the light model is selected, the agent loop swaps to `agent.LightCandidates`.
|
||||
When it is not selected, execution stays on the agent's primary provider candidate set.
|
||||
|
|
@ -252,7 +252,7 @@ When it is not selected, execution stays on the agent's primary provider candida
|
|||
|
||||
One nuance sits just outside `pkg/routing` but matters for the full routing story.
|
||||
|
||||
After a route is allocated, `pkg/agent/loop_utils.go:resolveScopeKey` preserves an explicit incoming session key when the caller already supplied:
|
||||
After a route is allocated, `pkg/agent/agent_utils.go:resolveScopeKey` preserves an explicit incoming session key when the caller already supplied:
|
||||
|
||||
- an opaque canonical key
|
||||
- a legacy `agent:...` key
|
||||
|
|
@ -278,5 +278,5 @@ They are separate from the runtime routing system described here.
|
|||
- `pkg/routing/agent_id.go`
|
||||
- `pkg/session/allocator.go`
|
||||
- `pkg/agent/registry.go`
|
||||
- `pkg/agent/loop_message.go`
|
||||
- `pkg/agent/loop_turn.go`
|
||||
- `pkg/agent/agent_message.go`
|
||||
- `pkg/agent/turn_coord.go`
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ The session system has four jobs:
|
|||
| Session adapter | `pkg/session/jsonl_backend.go` | Adapts `pkg/memory.Store` to `SessionStore`, including alias and scope metadata support. |
|
||||
| Durable storage | `pkg/memory/jsonl.go` | Append-only JSONL storage plus `.meta.json` sidecar metadata. |
|
||||
| Scope and key building | `pkg/session/scope.go`, `pkg/session/key.go`, `pkg/session/allocator.go` | Builds structured scopes, opaque canonical keys, and legacy aliases from routing results. |
|
||||
| Runtime integration | `pkg/agent/instance.go`, `pkg/agent/loop.go`, `pkg/agent/loop_message.go` | Initializes the store, allocates session scope, and persists metadata before turns run. |
|
||||
| Runtime integration | `pkg/agent/instance.go`, `pkg/agent/agent.go`, `pkg/agent/agent_message.go` | Initializes the store, allocates session scope, and persists metadata before turns run. |
|
||||
|
||||
## Session Data Model
|
||||
|
||||
|
|
@ -90,7 +90,7 @@ The agent loop also preserves explicit incoming session keys when the caller alr
|
|||
- opaque canonical key
|
||||
- legacy `agent:...` key
|
||||
|
||||
That behavior lives in `pkg/agent/loop_utils.go:resolveScopeKey`.
|
||||
That behavior lives in `pkg/agent/agent_utils.go:resolveScopeKey`.
|
||||
|
||||
## Allocation Flow
|
||||
|
||||
|
|
@ -108,7 +108,7 @@ InboundMessage
|
|||
|
||||
More concretely:
|
||||
|
||||
1. `pkg/agent/loop_message.go` resolves the agent route from normalized inbound context.
|
||||
1. `pkg/agent/agent_message.go` resolves the agent route from normalized inbound context.
|
||||
2. `session.AllocateRouteSession` converts the route's `SessionPolicy` plus inbound context into a structured `SessionScope`.
|
||||
3. The allocator builds:
|
||||
- `SessionKey`: canonical routed session key
|
||||
|
|
@ -251,5 +251,5 @@ The session system is consumed by more than the agent loop:
|
|||
- `pkg/session/allocator.go`
|
||||
- `pkg/memory/jsonl.go`
|
||||
- `pkg/agent/instance.go`
|
||||
- `pkg/agent/loop.go`
|
||||
- `pkg/agent/loop_message.go`
|
||||
- `pkg/agent/agent.go`
|
||||
- `pkg/agent/agent_message.go`
|
||||
|
|
|
|||
|
|
@ -8,26 +8,56 @@ Discord is a free voice, video, and text chat application designed for communiti
|
|||
|
||||
```json
|
||||
{
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"tool_feedback": {
|
||||
"enabled": true,
|
||||
"max_args_length": 300
|
||||
}
|
||||
}
|
||||
},
|
||||
"channel_list": {
|
||||
"discord": {
|
||||
"enabled": true,
|
||||
"type": "discord",
|
||||
"token": "YOUR_BOT_TOKEN",
|
||||
"allow_from": ["YOUR_USER_ID"],
|
||||
"placeholder": {
|
||||
"enabled": true,
|
||||
"text": ["Thinking... 💭"]
|
||||
},
|
||||
"group_trigger": {
|
||||
"mention_only": false
|
||||
}
|
||||
},
|
||||
"reasoning_channel_id": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| Field | Type | Required | Description |
|
||||
| ------------- | ------ | -------- | --------------------------------------------------------------------------- |
|
||||
| enabled | bool | Yes | Whether to enable the Discord channel |
|
||||
| token | string | Yes | Discord Bot Token |
|
||||
| allow_from | array | No | Allowlist of user IDs; empty means all users are allowed |
|
||||
| group_trigger | object | No | Group trigger settings (example: { "mention_only": false }) |
|
||||
| Field | Type | Required | Description |
|
||||
| -------------------- | ------ | -------- | --------------------------------------------------------------------------- |
|
||||
| enabled | bool | Yes | Whether to enable the Discord channel |
|
||||
| token | string | Yes | Discord Bot Token |
|
||||
| allow_from | array | No | Allowlist of user IDs; empty means all users are allowed |
|
||||
| placeholder | object | No | Placeholder message config shown while the agent is working |
|
||||
| group_trigger | object | No | Group trigger settings (example: { "mention_only": false }) |
|
||||
| reasoning_channel_id | string | No | Optional target channel ID for reasoning/thinking output |
|
||||
|
||||
## Visible Execution Feedback
|
||||
|
||||
Discord can show three different kinds of "working" feedback:
|
||||
|
||||
1. Typing indicator: automatic, no extra config needed.
|
||||
2. Placeholder message: enable `channel_list.discord.placeholder.enabled` to send a visible `Thinking...` message that is later edited into the final reply.
|
||||
3. Tool execution feedback: enable `agents.defaults.tool_feedback.enabled` to send a short message before each tool call, for example:
|
||||
|
||||
```text
|
||||
🔧 `web_search`
|
||||
Checking the latest PicoClaw release notes before I answer.
|
||||
```
|
||||
|
||||
If you only see `Bot is typing`, check that `placeholder.enabled` or `tool_feedback.enabled` is actually set in your runtime config.
|
||||
|
||||
## Setup
|
||||
|
||||
|
|
|
|||
|
|
@ -44,6 +44,8 @@ Telegram auto-registers PicoClaw's top-level bot commands at startup, including
|
|||
Skill-related commands:
|
||||
|
||||
- `/list skills` lists the installed skills visible to the current agent.
|
||||
- `/list mcp` lists configured MCP servers and whether they are deferred/connected.
|
||||
- `/show mcp <server>` lists the active tools for a connected MCP server.
|
||||
- `/use <skill> <message>` forces a skill for a single request.
|
||||
- `/use <skill>` arms the skill for your next message in the same chat.
|
||||
- `/use clear` clears a pending skill override.
|
||||
|
|
@ -52,6 +54,8 @@ Examples:
|
|||
|
||||
```text
|
||||
/list skills
|
||||
/list mcp
|
||||
/show mcp github
|
||||
/use git explain how to squash the last 3 commits
|
||||
/use git
|
||||
explain how to squash the last 3 commits
|
||||
|
|
|
|||
|
|
@ -154,7 +154,7 @@ Identify protocol via prefix in `model` field:
|
|||
| `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) |
|
||||
| `gemini/` | Gemini | Google Gemini native API |
|
||||
|
||||
---
|
||||
|
||||
|
|
|
|||
|
|
@ -67,9 +67,11 @@ Telegram command menu registration remains channel-local discovery UX; generic c
|
|||
|
||||
If command registration fails (network/API transient errors), the channel still starts and PicoClaw retries registration in the background.
|
||||
|
||||
You can also manage installed skills directly from Telegram:
|
||||
You can also inspect skills and MCP servers directly from Telegram:
|
||||
|
||||
- `/list skills`
|
||||
- `/list mcp`
|
||||
- `/show mcp <server>`
|
||||
- `/use <skill> <message>`
|
||||
- `/use <skill>` and then send the actual request in the next message
|
||||
- `/use clear`
|
||||
|
|
|
|||
|
|
@ -339,7 +339,7 @@ Répond HEARTBEAT_OK Utilisateur reçoit le résultat
|
|||
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Obtenir](https://console.anthropic.com) |
|
||||
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Obtenir](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Obtenir](https://platform.deepseek.com) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Obtenir](https://aistudio.google.com/api-keys) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [Obtenir](https://aistudio.google.com/api-keys) |
|
||||
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Obtenir](https://console.groq.com) |
|
||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Obtenir](https://dashscope.console.aliyun.com) |
|
||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (pas de clé) |
|
||||
|
|
@ -369,9 +369,12 @@ L'ancienne configuration `providers` est **dépréciée** et a été supprimée
|
|||
PicoClaw route les providers par famille de protocole :
|
||||
|
||||
- **Compatible OpenAI** : OpenRouter, Groq, Zhipu, endpoints vLLM et la plupart des autres.
|
||||
- **Gemini natif** : Google Gemini via les endpoints natifs `models/*:generateContent` et `models/*:streamGenerateContent`.
|
||||
- **Anthropic** : Comportement natif de l'API Claude.
|
||||
- **Codex/OAuth** : Route d'authentification OAuth/token OpenAI.
|
||||
|
||||
Cela maintient le runtime léger tout en faisant des nouveaux backends compatibles OpenAI principalement une opération de configuration (`api_base` + `api_keys`).
|
||||
|
||||
### Tâches Planifiées / Rappels
|
||||
|
||||
PicoClaw supporte les tâches planifiées via l'outil `cron`. L'agent peut définir, lister et annuler des rappels ou tâches récurrentes.
|
||||
|
|
|
|||
|
|
@ -340,7 +340,7 @@ HEARTBEAT_OK を返信 ユーザーが直接結果を受信
|
|||
| **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) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [取得](https://aistudio.google.com/api-keys) |
|
||||
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [取得](https://console.groq.com) |
|
||||
| **通義千問 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [取得](https://dashscope.console.aliyun.com) |
|
||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | ローカル(キー不要) |
|
||||
|
|
@ -370,9 +370,12 @@ HEARTBEAT_OK を返信 ユーザーが直接結果を受信
|
|||
PicoClaw はプロトコルファミリーで Provider をルーティングします:
|
||||
|
||||
- **OpenAI 互換**:OpenRouter、Groq、Zhipu、vLLM スタイルのエンドポイントなど。
|
||||
- **Gemini ネイティブ**:Google Gemini のネイティブ `models/*:generateContent` / `models/*:streamGenerateContent` エンドポイント。
|
||||
- **Anthropic**:Claude ネイティブ API の動作。
|
||||
- **Codex/OAuth**:OpenAI OAuth/トークン認証ルート。
|
||||
|
||||
これによりランタイムを軽量に保ちつつ、新しい OpenAI 互換バックエンドの追加をほぼ設定操作(`api_base` + `api_keys`)のみで実現します。
|
||||
|
||||
### スケジュールタスク / リマインダー
|
||||
|
||||
PicoClaw は `cron` ツールを通じて cron スタイルのスケジュールタスクをサポートします。
|
||||
|
|
|
|||
|
|
@ -71,15 +71,16 @@ PicoClaw stores data in your configured workspace (default: `~/.picoclaw/workspa
|
|||
|
||||
### Web launcher dashboard
|
||||
|
||||
**picoclaw-launcher** serves a browser UI that requires sign-in first. By default, the **dashboard token** and **session signing key** are **generated in memory on each start** (a new random token after every restart). Set **`PICOCLAW_LAUNCHER_TOKEN`** to pin a fixed token for that process (startup logs do not print the secret when this env var is used).
|
||||
|
||||
**Where to read the token**: In **console mode** (`-console`), it is printed at startup. In **tray / GUI mode**, use the tray action **Copy dashboard token**, and check **`$PICOCLAW_HOME/logs/launcher.log`** (typically `~/.picoclaw/logs/launcher.log` if `PICOCLAW_HOME` is unset) for the random token logged on startup. The login page shows hints that match how the launcher is running (including the absolute log path); **responses do not include the token itself**.
|
||||
**picoclaw-launcher** serves a browser UI that requires password sign-in first. On first run, open `/launcher-setup` to create the dashboard password. Later manual sign-ins use `/launcher-login`.
|
||||
|
||||
- **Config file**: Same directory as `config.json` (or the file pointed to by `PICOCLAW_CONFIG`). The launcher-specific file is `launcher-config.json`.
|
||||
- **Sign-in and links**: Enter the token on the login page, or open with `?token=` when the browser is launched automatically. All responses include **`Referrer-Policy: no-referrer`** to reduce leakage of `token` via the `Referer` header.
|
||||
- **Password storage**: On supported platforms, the password is stored as a bcrypt hash in `launcher-auth.db`. On platforms where the SQLite password store is unavailable, the bcrypt hash is stored in `launcher-config.json`.
|
||||
- **Legacy migration**: Older `launcher_token` values are migrated once into password login and removed from saved launcher config.
|
||||
- **Local auto-login**: When the launcher auto-opens a local browser after startup, it uses a one-shot loopback-only bootstrap endpoint to set the session cookie automatically.
|
||||
- **Unsupported auth paths**: URL token login (`?token=...`), `PICOCLAW_LAUNCHER_TOKEN`, and `Authorization: Bearer` dashboard auth are no longer supported.
|
||||
- **Sign-out**: Use **`POST /api/auth/logout`** with **`Content-Type: application/json`** (body may be `{}`). Do not rely on a GET URL for logout (CSRF-safe pattern).
|
||||
- **Brute-force**: **`POST /api/auth/login`** is **rate-limited per client IP per minute** (HTTP 429 when exceeded).
|
||||
- **Session lifetime**: The HttpOnly session cookie lasts about **7 days** by default; sign in again with the token after it expires.
|
||||
- **Session lifetime**: The HttpOnly session cookie lasts about **31 days** by default, but sessions are invalidated when the launcher process restarts.
|
||||
|
||||
### Skill Sources
|
||||
|
||||
|
|
@ -97,9 +98,11 @@ export PICOCLAW_BUILTIN_SKILLS=/path/to/skills
|
|||
|
||||
### Using Skills From Chat Channels
|
||||
|
||||
Once skills are installed, you can inspect and force them directly from a chat channel:
|
||||
Once skills are installed, and MCP servers are configured, you can inspect and force them directly from a chat channel:
|
||||
|
||||
- `/list skills` shows the installed skill names available to the current agent.
|
||||
- `/list mcp` shows configured MCP servers with enabled/deferred/connected status.
|
||||
- `/show mcp <server>` shows the active tools exposed by a connected MCP server.
|
||||
- `/use <skill> <message>` forces a specific skill for a single request.
|
||||
- `/use <skill>` arms that skill for your next message in the same chat session.
|
||||
- `/use clear` cancels a pending skill override created by `/use <skill>`.
|
||||
|
|
@ -109,6 +112,8 @@ Examples:
|
|||
|
||||
```text
|
||||
/list skills
|
||||
/list mcp
|
||||
/show mcp github
|
||||
/use git explain how to squash the last 3 commits
|
||||
/btw remind me what we already decided about the deploy plan
|
||||
/use italiapersonalfinance
|
||||
|
|
@ -493,7 +498,7 @@ The subagent has access to tools (message, web_search, etc.) and can communicate
|
|||
|
||||
### 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!**
|
||||
> **What's New?** PicoClaw now prefers explicit `provider` + native `model` configuration (for example `"provider": "zhipu", "model": "glm-4.7"`). The legacy single-field `provider/model` form remains supported for compatibility when `provider` is omitted.
|
||||
|
||||
This design also enables **multi-agent support** with flexible provider selection:
|
||||
|
||||
|
|
@ -546,7 +551,8 @@ chmod 600 ~/.picoclaw/.security.yml
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4"
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4"
|
||||
// api_key loaded from .security.yml
|
||||
}
|
||||
],
|
||||
|
|
@ -570,31 +576,31 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
|
||||
#### All Supported Vendors
|
||||
|
||||
| Vendor | `model` Prefix | Default API Base | Protocol | API Key |
|
||||
| Vendor | `provider` Value | 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) |
|
||||
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | Optional (local default: no key) |
|
||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Get Key](https://cerebras.ai) |
|
||||
| **VolcEngine (Doubao)** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | — |
|
||||
| **BytePlus** | `byteplus/` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [Get Key](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid/` | `https://api.vivgrid.com/v1` | OpenAI | [Get Key](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [Get Key](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)** | `modelscope/` | `https://api-inference.modelscope.cn/v1` | OpenAI | [Get Token](https://modelscope.cn/my/tokens) |
|
||||
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth only |
|
||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | — |
|
||||
| **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` | Gemini | [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) |
|
||||
| **LM Studio** | `lmstudio` | `http://localhost:1234/v1` | OpenAI | Optional (local default: no key) |
|
||||
| **OpenRouter** | `openrouter` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
||||
| **VLLM** | `vllm` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||
| **Cerebras** | `cerebras` | `https://api.cerebras.ai/v1` | OpenAI | [Get Key](https://cerebras.ai) |
|
||||
| **VolcEngine (Doubao)** | `volcengine` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun` | `https://router.shengsuanyun.com/api/v1` | OpenAI | — |
|
||||
| **BytePlus** | `byteplus` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [Get Key](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid` | `https://api.vivgrid.com/v1` | OpenAI | [Get Key](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat` | `https://api.longcat.chat/openai` | OpenAI | [Get Key](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)** | `modelscope` | `https://api-inference.modelscope.cn/v1` | OpenAI | [Get Token](https://modelscope.cn/my/tokens) |
|
||||
| **Antigravity** | `antigravity` | Google Cloud | Custom | OAuth only |
|
||||
| **GitHub Copilot** | `github-copilot` | `localhost:4321` | gRPC | — |
|
||||
|
||||
#### Basic Configuration
|
||||
|
||||
|
|
@ -603,22 +609,26 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-your-api-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-your-openai-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-zhipu-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -634,6 +644,13 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
>
|
||||
> **Note**: The `enabled` field can be set to `false` to disable a model entry without removing it. When omitted, it defaults to `true` during migration for models that have API keys.
|
||||
|
||||
Resolution rules:
|
||||
|
||||
- Prefer explicit `"provider": "openai", "model": "gpt-5.4"`.
|
||||
- If `provider` is set, PicoClaw sends `model` unchanged.
|
||||
- If `provider` is omitted, PicoClaw treats the first `/` segment in `model` as the provider and everything after that first `/` as the runtime model ID.
|
||||
- This means `"model": "openrouter/openai/gpt-5.4"` still works as a compatibility form and sends `openai/gpt-5.4` to OpenRouter.
|
||||
|
||||
#### Vendor-Specific Examples
|
||||
|
||||
> **Tip**: You can omit `api_key` fields and store them in `.security.yml` for better security. See [Security Configuration](#-security-configuration-recommended).
|
||||
|
|
@ -644,7 +661,8 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
```json
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4"
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4"
|
||||
// api_key: set in .security.yml
|
||||
}
|
||||
```
|
||||
|
|
@ -657,7 +675,8 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
```json
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest"
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest"
|
||||
// api_key: set in .security.yml
|
||||
}
|
||||
```
|
||||
|
|
@ -670,7 +689,8 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
```json
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7"
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7"
|
||||
// api_key: set in .security.yml
|
||||
}
|
||||
```
|
||||
|
|
@ -683,7 +703,8 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
```json
|
||||
{
|
||||
"model_name": "deepseek-chat",
|
||||
"model": "deepseek/deepseek-chat"
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat"
|
||||
// api_key: set in .security.yml
|
||||
}
|
||||
```
|
||||
|
|
@ -696,7 +717,8 @@ For complete documentation, see [`../security/security_configuration.md`](../sec
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6"
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6"
|
||||
// api_key: set in .security.yml
|
||||
}
|
||||
```
|
||||
|
|
@ -708,7 +730,8 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-opus-4-6",
|
||||
"model": "anthropic-messages/claude-opus-4-6",
|
||||
"provider": "anthropic-messages",
|
||||
"model": "claude-opus-4-6",
|
||||
"api_keys": ["sk-ant-your-key"],
|
||||
"api_base": "https://api.anthropic.com"
|
||||
}
|
||||
|
|
@ -724,7 +747,8 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
|||
```json
|
||||
{
|
||||
"model_name": "llama3",
|
||||
"model": "ollama/llama3"
|
||||
"provider": "ollama",
|
||||
"model": "llama3"
|
||||
}
|
||||
```
|
||||
|
||||
|
|
@ -736,12 +760,13 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
|||
```json
|
||||
{
|
||||
"model_name": "lmstudio-local",
|
||||
"model": "lmstudio/openai/gpt-oss-20b"
|
||||
"provider": "lmstudio",
|
||||
"model": "openai/gpt-oss-20b"
|
||||
}
|
||||
```
|
||||
|
||||
`api_base` defaults to `http://localhost:1234/v1`. API key is optional unless your LM Studio server enables authentication.<br/>
|
||||
PicoClaw sends OpenAI-compatible requests to LM Studio, and strips the `lmstudio/` prefix before sending requests, so `lmstudio/openai/gpt-oss-20b` sends `openai/gpt-oss-20b` to the LM Studio server.
|
||||
With explicit `provider`, PicoClaw sends `openai/gpt-oss-20b` unchanged to LM Studio. The legacy compatibility form `"model": "lmstudio/openai/gpt-oss-20b"` still resolves to the same upstream model ID when `provider` is omitted.
|
||||
|
||||
</details>
|
||||
|
||||
|
|
@ -751,13 +776,14 @@ PicoClaw sends OpenAI-compatible requests to LM Studio, and strips the `lmstudio
|
|||
```json
|
||||
{
|
||||
"model_name": "my-custom-model",
|
||||
"model": "openai/custom-model",
|
||||
"provider": "openai",
|
||||
"model": "custom-model",
|
||||
"api_base": "https://my-proxy.com/v1"
|
||||
// api_key: set in .security.yml
|
||||
}
|
||||
```
|
||||
|
||||
PicoClaw strips only the outer `litellm/` prefix before sending the request, so `litellm/lite-gpt4` sends `lite-gpt4`, while `litellm/openai/gpt-4o` sends `openai/gpt-4o`.
|
||||
With explicit `provider`, PicoClaw sends `model` unchanged. That means `"provider": "litellm", "model": "lite-gpt4"` sends `lite-gpt4`, while `"provider": "litellm", "model": "openai/gpt-4o"` sends `openai/gpt-4o`. The legacy compatibility forms `litellm/lite-gpt4` and `litellm/openai/gpt-4o` still resolve the same way when `provider` is omitted.
|
||||
|
||||
</details>
|
||||
|
||||
|
|
@ -782,7 +808,8 @@ model_list:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api.openai.com/v1"
|
||||
// api_keys loaded from .security.yml
|
||||
}
|
||||
|
|
@ -797,13 +824,15 @@ model_list:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api1.example.com/v1",
|
||||
"api_keys": ["sk-key1"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api2.example.com/v1",
|
||||
"api_keys": ["sk-key2"]
|
||||
}
|
||||
|
|
@ -820,6 +849,7 @@ The old `providers` configuration is **deprecated** and has been removed in V2.
|
|||
PicoClaw routes providers by protocol family:
|
||||
|
||||
- **OpenAI-compatible**: OpenRouter, Groq, Zhipu, vLLM-style endpoints, and most others.
|
||||
- **Gemini native**: Google Gemini via the native `models/*:generateContent` and `models/*:streamGenerateContent` endpoints.
|
||||
- **Anthropic**: Claude-native API behavior.
|
||||
- **Codex/OAuth**: OpenAI OAuth/token authentication route.
|
||||
|
||||
|
|
@ -862,7 +892,7 @@ This keeps the runtime lightweight while making new OpenAI-compatible backends m
|
|||
{
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"model": "anthropic/claude-opus-4-5"
|
||||
"model_name": "claude-opus-4-5"
|
||||
}
|
||||
},
|
||||
"session": {
|
||||
|
|
|
|||
|
|
@ -340,7 +340,7 @@ Responde HEARTBEAT_OK Usuário recebe resultado diretamente
|
|||
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Obter](https://console.anthropic.com) |
|
||||
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Obter](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Obter](https://platform.deepseek.com) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Obter](https://aistudio.google.com/api-keys) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [Obter](https://aistudio.google.com/api-keys) |
|
||||
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Obter](https://console.groq.com) |
|
||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Obter](https://dashscope.console.aliyun.com) |
|
||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (sem chave) |
|
||||
|
|
@ -370,9 +370,12 @@ A configuração antiga `providers` está **depreciada** e foi removida no V2. C
|
|||
PicoClaw roteia providers por família de protocolo:
|
||||
|
||||
- **Compatível com OpenAI**: OpenRouter, Groq, Zhipu, endpoints vLLM e a maioria dos outros.
|
||||
- **Gemini nativo**: Google Gemini via endpoints nativos `models/*:generateContent` e `models/*:streamGenerateContent`.
|
||||
- **Anthropic**: Comportamento nativo da API Claude.
|
||||
- **Codex/OAuth**: Rota de autenticação OAuth/token OpenAI.
|
||||
|
||||
Isso mantém o runtime leve enquanto torna novos backends compatíveis com OpenAI basicamente uma operação de configuração (`api_base` + `api_keys`).
|
||||
|
||||
### Tarefas Agendadas / Lembretes
|
||||
|
||||
PicoClaw suporta tarefas agendadas via ferramenta `cron`.
|
||||
|
|
|
|||
|
|
@ -340,7 +340,7 @@ Trả lời HEARTBEAT_OK Người dùng nhận kết quả trực tiếp
|
|||
| **Anthropic** | `anthropic/` | `https://api.anthropic.com/v1` | Anthropic | [Lấy](https://console.anthropic.com) |
|
||||
| **智谱 AI (GLM)** | `zhipu/` | `https://open.bigmodel.cn/api/paas/v4` | OpenAI | [Lấy](https://open.bigmodel.cn/usercenter/proj-mgmt/apikeys) |
|
||||
| **DeepSeek** | `deepseek/` | `https://api.deepseek.com/v1` | OpenAI | [Lấy](https://platform.deepseek.com) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | OpenAI | [Lấy](https://aistudio.google.com/api-keys) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [Lấy](https://aistudio.google.com/api-keys) |
|
||||
| **Groq** | `groq/` | `https://api.groq.com/openai/v1` | OpenAI | [Lấy](https://console.groq.com) |
|
||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Lấy](https://dashscope.console.aliyun.com) |
|
||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Cục bộ (không cần key) |
|
||||
|
|
@ -370,9 +370,12 @@ Cấu hình `providers` cũ đã **bị deprecated** và đã được loại b
|
|||
PicoClaw định tuyến provider theo họ giao thức:
|
||||
|
||||
- **Tương thích OpenAI**: OpenRouter, Groq, Zhipu, endpoint kiểu vLLM và hầu hết các provider khác.
|
||||
- **Gemini native**: Google Gemini qua các endpoint native `models/*:generateContent` và `models/*:streamGenerateContent`.
|
||||
- **Anthropic**: Hành vi API Claude gốc.
|
||||
- **Codex/OAuth**: Tuyến xác thực OAuth/token OpenAI.
|
||||
|
||||
Điều này giữ runtime nhẹ trong khi khiến backend OpenAI-compatible mới chủ yếu chỉ là thao tác cấu hình (`api_base` + `api_keys`).
|
||||
|
||||
### Tác Vụ Đã Lên Lịch / Nhắc Nhở
|
||||
|
||||
PicoClaw hỗ trợ tác vụ theo lịch qua công cụ `cron`.
|
||||
|
|
|
|||
|
|
@ -69,15 +69,16 @@ PicoClaw 将数据存储在您配置的工作区中(默认:`~/.picoclaw/work
|
|||
|
||||
### Web 启动器控制台
|
||||
|
||||
用 **picoclaw-launcher** 打开浏览器控制台前需要先登录。**访问口令**与 **会话签名密钥**默认在**每次启动时在内存中生成**(重启后随机口令会变)。若设置环境变量 **`PICOCLAW_LAUNCHER_TOKEN`**,则该进程使用固定口令(启动日志中不会打印具体口令值)。
|
||||
|
||||
**到哪里找口令**:**控制台模式**(`-console`)请看启动时的终端输出;**托盘 / GUI 模式**可使用托盘菜单中的「复制控制台口令」,并在 **`$PICOCLAW_HOME/logs/launcher.log`**(未设置 `PICOCLAW_HOME` 时一般为 `~/.picoclaw/logs/launcher.log`)中查看本次启动写入的随机口令。登录页在未登录时会根据当前运行方式展示提示(含日志文件绝对路径等;**接口与页面均不会返回口令本身**)。
|
||||
用 **picoclaw-launcher** 打开浏览器控制台前需要先使用密码登录。首次启动时打开 `/launcher-setup` 创建 dashboard 登录密码;后续手动登录使用 `/launcher-login`。
|
||||
|
||||
- **配置文件**:与 `config.json` 同一目录(若设置了 `PICOCLAW_CONFIG`,则与它所指的文件同目录)。启动器专用文件名为 `launcher-config.json`。
|
||||
- **登录与链接**:在登录页输入口令;自动打开浏览器时可在 URL 上使用 `?token=`。全站响应携带 **`Referrer-Policy: no-referrer`**,减轻 `token` 经 `Referer` 头泄露的风险。
|
||||
- **密码存储**:支持的平台会把 bcrypt 后的密码哈希存入 `launcher-auth.db`。如果当前平台不支持 SQLite 密码存储,则把 bcrypt 哈希存入 `launcher-config.json`。
|
||||
- **旧配置迁移**:旧版 `launcher_token` 会一次性迁移为密码登录,并从保存后的 launcher 配置中移除。
|
||||
- **本地自动登录**:launcher 启动后自动打开本地浏览器时,会使用仅允许 loopback 访问的一次性引导入口自动设置会话 Cookie。
|
||||
- **不再支持的鉴权方式**:不再支持 URL token 登录(`?token=...`)、`PICOCLAW_LAUNCHER_TOKEN` 和 `Authorization: Bearer` dashboard 鉴权。
|
||||
- **退出登录**:应使用 **`POST /api/auth/logout`**,且请求头为 **`Content-Type: application/json`**(请求体可为 `{}`),勿使用可被第三方页面触发的 GET 链接登出。
|
||||
- **暴力尝试**:`POST /api/auth/login` 对同一远程地址有 **每分钟尝试次数上限**(超限返回 HTTP 429)。
|
||||
- **会话时长**:登录后的 HttpOnly 会话 Cookie 默认约 **7 天**有效,到期需重新用口令登录。
|
||||
- **会话时长**:登录后的 HttpOnly 会话 Cookie 默认约 **31 天**有效,但 launcher 进程重启后已有会话会失效。
|
||||
|
||||
### 技能来源 (Skill Sources)
|
||||
|
||||
|
|
@ -424,7 +425,7 @@ Agent 读取 HEARTBEAT.md
|
|||
|
||||
### 模型配置 (model_list)
|
||||
|
||||
> **新特性:** PicoClaw 现在采用**以模型为中心**的配置方式。只需指定 `vendor/model` 格式(例如 `zhipu/glm-4.7`)即可接入新提供商——**无需修改任何代码!**
|
||||
> **新特性:** PicoClaw 现在优先推荐显式 `provider` + 原生 `model` 的配置方式,例如 `"provider": "zhipu", "model": "glm-4.7"`。如果未设置 `provider`,旧的单字段 `provider/model` 写法仍然兼容。
|
||||
|
||||
这一设计同时支持**多 Agent**场景,灵活选择提供商:
|
||||
|
||||
|
|
@ -435,31 +436,31 @@ Agent 读取 HEARTBEAT.md
|
|||
|
||||
#### 所有支持的厂商
|
||||
|
||||
| 厂商 | `model` 前缀 | 默认 API Base | 协议 | API Key |
|
||||
| 厂商 | `provider` 值 | 默认 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 | 本地(无需 Key) |
|
||||
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | 可选(本地默认无需密钥) |
|
||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理 Key |
|
||||
| **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://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | — |
|
||||
| **BytePlus** | `byteplus/` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [获取](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid/` | `https://api.vivgrid.com/v1` | OpenAI | [获取](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [获取](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)** | `modelscope/` | `https://api-inference.modelscope.cn/v1` | OpenAI | [获取](https://modelscope.cn/my/tokens) |
|
||||
| **Antigravity** | `antigravity/` | Google Cloud | Custom | 仅 OAuth |
|
||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | — |
|
||||
| **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` | Gemini | [获取](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 | 本地(无需 Key) |
|
||||
| **LM Studio** | `lmstudio` | `http://localhost:1234/v1` | OpenAI | 可选(本地默认无需密钥) |
|
||||
| **OpenRouter** | `openrouter` | `https://openrouter.ai/api/v1` | OpenAI | [获取](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理 Key |
|
||||
| **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://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun` | `https://router.shengsuanyun.com/api/v1` | OpenAI | — |
|
||||
| **BytePlus** | `byteplus` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [获取](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid` | `https://api.vivgrid.com/v1` | OpenAI | [获取](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat` | `https://api.longcat.chat/openai` | OpenAI | [获取](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)** | `modelscope` | `https://api-inference.modelscope.cn/v1` | OpenAI | [获取](https://modelscope.cn/my/tokens) |
|
||||
| **Antigravity** | `antigravity` | Google Cloud | Custom | 仅 OAuth |
|
||||
| **GitHub Copilot** | `github-copilot` | `localhost:4321` | gRPC | — |
|
||||
|
||||
#### 基础配置
|
||||
|
||||
|
|
@ -468,22 +469,26 @@ Agent 读取 HEARTBEAT.md
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-your-api-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-your-openai-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-zhipu-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -495,6 +500,13 @@ Agent 读取 HEARTBEAT.md
|
|||
}
|
||||
```
|
||||
|
||||
解析规则:
|
||||
|
||||
- 推荐显式写成 `"provider": "openai", "model": "gpt-5.4"`。
|
||||
- 如果设置了 `provider`,PicoClaw 会将 `model` 原样发送。
|
||||
- 如果未设置 `provider`,PicoClaw 会把 `model` 第一个 `/` 之前的字段当作 provider,并把第一个 `/` 之后的全部内容当作最终模型 ID。
|
||||
- 这意味着 `"model": "openrouter/openai/gpt-5.4"` 这样的兼容写法仍然可用,并会把 `openai/gpt-5.4` 发送给 OpenRouter。
|
||||
|
||||
#### 各厂商配置示例
|
||||
|
||||
<details>
|
||||
|
|
@ -503,7 +515,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -516,7 +529,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -529,7 +543,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-key"]
|
||||
}
|
||||
```
|
||||
|
|
@ -542,7 +557,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "deepseek-chat",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -555,7 +571,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
}
|
||||
```
|
||||
|
|
@ -567,7 +584,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-opus-4-6",
|
||||
"model": "anthropic-messages/claude-opus-4-6",
|
||||
"provider": "anthropic-messages",
|
||||
"model": "claude-opus-4-6",
|
||||
"api_keys": ["sk-ant-your-key"],
|
||||
"api_base": "https://api.anthropic.com"
|
||||
}
|
||||
|
|
@ -583,7 +601,8 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "llama3",
|
||||
"model": "ollama/llama3"
|
||||
"provider": "ollama",
|
||||
"model": "llama3"
|
||||
}
|
||||
```
|
||||
|
||||
|
|
@ -595,12 +614,13 @@ Agent 读取 HEARTBEAT.md
|
|||
```json
|
||||
{
|
||||
"model_name": "lmstudio-local",
|
||||
"model": "lmstudio/openai/gpt-oss-20b"
|
||||
"provider": "lmstudio",
|
||||
"model": "openai/gpt-oss-20b"
|
||||
}
|
||||
```
|
||||
|
||||
`api_base` 默认是 `http://localhost:1234/v1`。除非你在 LM Studio 侧启用了认证,否则不需要配置 API Key。
|
||||
PicoClaw 向 LM Studio 的 OpenAI 兼容终结点发送请求,且将移除首个 `lmstudio/` 前缀,因此 `lmstudio/openai/gpt-oss-20b` 会发送 `openai/gpt-oss-20b`。
|
||||
显式设置 `provider` 后,PicoClaw 会把 `openai/gpt-oss-20b` 原样发送给 LM Studio。旧的兼容写法 `"model": "lmstudio/openai/gpt-oss-20b"` 在未设置 `provider` 时也会解析成相同的上游模型 ID。
|
||||
|
||||
</details>
|
||||
|
||||
|
|
@ -610,13 +630,14 @@ PicoClaw 向 LM Studio 的 OpenAI 兼容终结点发送请求,且将移除首
|
|||
```json
|
||||
{
|
||||
"model_name": "my-custom-model",
|
||||
"model": "openai/custom-model",
|
||||
"provider": "openai",
|
||||
"model": "custom-model",
|
||||
"api_base": "https://my-proxy.com/v1",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
||||
PicoClaw 只剥离最外层的 `litellm/` 前缀再发送请求,因此 `litellm/lite-gpt4` 发送 `lite-gpt4`,而 `litellm/openai/gpt-4o` 发送 `openai/gpt-4o`。
|
||||
显式设置 `provider` 后,PicoClaw 会将 `model` 原样发送。因此 `"provider": "litellm", "model": "lite-gpt4"` 会发送 `lite-gpt4`,而 `"provider": "litellm", "model": "openai/gpt-4o"` 会发送 `openai/gpt-4o`。旧的兼容写法 `litellm/lite-gpt4` 和 `litellm/openai/gpt-4o` 在未设置 `provider` 时也会得到相同结果。
|
||||
|
||||
</details>
|
||||
|
||||
|
|
@ -629,13 +650,15 @@ PicoClaw 只剥离最外层的 `litellm/` 前缀再发送请求,因此 `litell
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api1.example.com/v1",
|
||||
"api_keys": ["sk-key1"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api2.example.com/v1",
|
||||
"api_keys": ["sk-key2"]
|
||||
}
|
||||
|
|
@ -652,10 +675,11 @@ PicoClaw 只剥离最外层的 `litellm/` 前缀再发送请求,因此 `litell
|
|||
PicoClaw 按协议族路由提供商:
|
||||
|
||||
- **OpenAI 兼容**:OpenRouter、Groq、智谱、vLLM 风格端点及大多数其他提供商。
|
||||
- **Gemini 原生**:Google Gemini 通过原生 `models/*:generateContent` 和 `models/*:streamGenerateContent` 端点接入。
|
||||
- **Anthropic**:Claude 原生 API 行为。
|
||||
- **Codex/OAuth**:OpenAI OAuth/Token 认证路由。
|
||||
|
||||
这使运行时保持轻量,同时让接入新的 OpenAI 兼容后端基本只需配置 `api_base` + `api_key`。
|
||||
这使运行时保持轻量,同时让接入新的 OpenAI 兼容后端基本只需配置 `api_base` + `api_keys`。
|
||||
|
||||
<details>
|
||||
<summary><b>智谱(旧版 providers 格式)</b></summary>
|
||||
|
|
@ -689,7 +713,7 @@ PicoClaw 按协议族路由提供商:
|
|||
{
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"model": "anthropic/claude-opus-4-5"
|
||||
"model_name": "claude-opus-4-5"
|
||||
}
|
||||
},
|
||||
"session": {
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
Ouvrez http://localhost:18800 dans votre navigateur. Le launcher gère automatiquement le processus gateway.
|
||||
|
||||
> [!WARNING]
|
||||
> La console web ne prend pas encore en charge l'authentification. Évitez de l'exposer sur Internet public.
|
||||
> La console web est protégée par un mot de passe de connexion au dashboard. Ne l'exposez pas à des réseaux non fiables ni à Internet public.
|
||||
|
||||
### Mode Agent (One-shot)
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
ブラウザで http://localhost:18800 を開いてください。Launcher が Gateway プロセスを自動管理します。
|
||||
|
||||
> [!WARNING]
|
||||
> Web コンソールはまだ認証をサポートしていません。公開インターネットに公開しないでください。
|
||||
> Web コンソールは dashboard ログインパスワードで保護されます。信頼できないネットワークや公開インターネットには公開しないでください。
|
||||
|
||||
### Agent モード (ワンショット)
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ docker compose -f docker/docker-compose.yml --profile gateway up -d
|
|||
> **Docker Users**: By default, the Gateway listens on `127.0.0.1` which is not accessible from the host. If you need to access the health endpoints or expose ports, set `PICOCLAW_GATEWAY_HOST=0.0.0.0` in your environment or update `config.json`.
|
||||
|
||||
> [!NOTE]
|
||||
> The `gateway` profile only serves the webhook handlers (including Pico when enabled) and health endpoints on the gateway port, so it does not expose generic REST chat endpoints such as `/chat` or `/a2a`. Launcher mode adds the browser UI plus `/api/pico/token` and a `/pico/ws` proxy on the launcher port, but `/pico/ws` is also available directly on the gateway whenever the Pico channel is enabled.
|
||||
> The `gateway` profile only serves the webhook handlers (including Pico when enabled) and health endpoints on the gateway port, so it does not expose generic REST chat endpoints such as `/chat` or `/a2a`. Launcher mode adds the browser UI plus `/api/pico/info` and an authenticated `/pico/ws` proxy on the launcher port, but `/pico/ws` is also available directly on the gateway whenever the Pico channel is enabled.
|
||||
|
||||
```bash
|
||||
# 5. Check logs
|
||||
|
|
@ -48,7 +48,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
Open http://localhost:18800 in your browser. The launcher manages the gateway process automatically.
|
||||
|
||||
> [!WARNING]
|
||||
> The web console uses a dashboard token (in-memory per run unless `PICOCLAW_LAUNCHER_TOKEN` is set). **Do not** expose the launcher to untrusted networks or the public internet. See [Web launcher dashboard](configuration.md#web-launcher-dashboard) in the Configuration Guide.
|
||||
> The web console is protected by dashboard password login. **Do not** expose the launcher to untrusted networks or the public internet. See [Web launcher dashboard](configuration.md#web-launcher-dashboard) in the Configuration Guide.
|
||||
|
||||
### Agent Mode (One-shot)
|
||||
|
||||
|
|
@ -94,19 +94,22 @@ picoclaw onboard
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-your-api-key"],
|
||||
"api_base":"https://ark.cn-beijing.volces.com/api/coding/v3"
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["your-api-key"],
|
||||
"request_timeout": 300
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["your-anthropic-key"]
|
||||
}
|
||||
],
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
Buka http://localhost:18800 dalam pelayar anda. Launcher mengurus proses gateway secara automatik.
|
||||
|
||||
> [!WARNING]
|
||||
> Konsol web belum menyokong autentikasi. Elakkan mendedahkannya ke internet awam.
|
||||
> Konsol web dilindungi oleh kata laluan log masuk dashboard. Jangan dedahkannya kepada rangkaian tidak dipercayai atau internet awam.
|
||||
|
||||
### Mod Agent (One-shot)
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
Abra http://localhost:18800 no seu navegador. O launcher gerencia o processo do gateway automaticamente.
|
||||
|
||||
> [!WARNING]
|
||||
> O console web ainda não suporta autenticação. Evite expô-lo na internet pública.
|
||||
> O console web é protegido por senha de login do dashboard. Não exponha o launcher a redes não confiáveis nem à internet pública.
|
||||
|
||||
### Modo Agent (One-shot)
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
Mở http://localhost:18800 trong trình duyệt. Launcher tự động quản lý tiến trình gateway.
|
||||
|
||||
> [!WARNING]
|
||||
> Web console chưa hỗ trợ xác thực. Tránh để lộ ra internet công cộng.
|
||||
> Web console được bảo vệ bằng mật khẩu đăng nhập dashboard. Không để lộ launcher ra mạng không tin cậy hoặc internet công cộng.
|
||||
|
||||
### Chế Độ Agent (One-shot)
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ docker compose -f docker/docker-compose.yml --profile launcher up -d
|
|||
在浏览器中打开 <http://localhost:18800>。Launcher 会自动管理 Gateway 进程。
|
||||
|
||||
> [!WARNING]
|
||||
> Web 控制台通过 dashboard 令牌鉴权(默认每次启动在内存中生成;可用 `PICOCLAW_LAUNCHER_TOKEN` 固定)。**不要**将启动器暴露到不可信网络或公网。完整说明见 [配置指南](configuration.md) 中的「Web 启动器控制台」一节。
|
||||
> Web 控制台通过 dashboard 登录密码保护。**不要**将启动器暴露到不可信网络或公网。完整说明见 [配置指南](configuration.md) 中的「Web 启动器控制台」一节。
|
||||
|
||||
### Agent 模式 (一次性运行)
|
||||
|
||||
|
|
@ -93,19 +93,22 @@ picoclaw onboard
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-your-api-key"],
|
||||
"api_base":"https://ark.cn-beijing.volces.com/api/coding/v3"
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["your-api-key"],
|
||||
"request_timeout": 300
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["your-anthropic-key"]
|
||||
}
|
||||
],
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ Cette conception permet également le **support multi-agents** avec une sélecti
|
|||
| **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) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [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) |
|
||||
|
|
@ -108,7 +108,7 @@ Cette conception permet également le **support multi-agents** avec une sélecti
|
|||
| `api_keys` | string[] | Oui* | Clé(s) API pour l'authentification. Plusieurs clés permettent la rotation par requête. Non requis pour les fournisseurs locaux (Ollama, LM Studio, VLLM) |
|
||||
| `api_base` | string | Non | Remplace l'URL de base API par défaut |
|
||||
| `proxy` | string | Non | URL du proxy HTTP pour cette entrée de modèle |
|
||||
| `user_agent` | string | Non | En-tête `User-Agent` personnalisé pour les requêtes API (supporté par les providers OpenAI-compatible, Anthropic et Azure) |
|
||||
| `user_agent` | string | Non | En-tête `User-Agent` personnalisé pour les requêtes API (supporté par les providers compatibles OpenAI, Gemini, Anthropic et Azure) |
|
||||
| `request_timeout` | int | Non | Délai d'expiration de la requête en secondes (la valeur par défaut varie selon le provider) |
|
||||
| `max_tokens_field` | string | Non | Remplace le nom du champ max tokens dans le corps de la requête (ex : `max_completion_tokens` pour les modèles o1) |
|
||||
| `thinking_level` | string | Non | Niveau de pensée étendue : `off`, `low`, `medium`, `high`, `xhigh` ou `adaptive` |
|
||||
|
|
@ -299,10 +299,11 @@ Pour un guide de migration détaillé, voir [migration/model-list-migration.md](
|
|||
PicoClaw route les fournisseurs par famille de protocoles :
|
||||
|
||||
- Protocole compatible OpenAI : OpenRouter, passerelles compatibles OpenAI, Groq, Zhipu et endpoints de type vLLM.
|
||||
- Protocole Gemini natif : Google Gemini via les endpoints natifs `models/*:generateContent` et `models/*:streamGenerateContent`.
|
||||
- Protocole Anthropic : Comportement natif de l'API Claude.
|
||||
- Chemin Codex/OAuth : Route d'authentification OAuth/token OpenAI.
|
||||
|
||||
Cela maintient le runtime léger tout en faisant des nouveaux backends compatibles OpenAI principalement une opération de configuration (`api_base` + `api_key`).
|
||||
Cela maintient le runtime léger tout en faisant des nouveaux backends compatibles OpenAI principalement une opération de configuration (`api_base` + `api_keys`).
|
||||
|
||||
<details>
|
||||
<summary><b>Zhipu</b></summary>
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@
|
|||
| **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) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [キーを取得](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) |
|
||||
|
|
@ -109,7 +109,7 @@
|
|||
| `api_keys` | string[] | はい* | 認証キー。複数キーでリクエストごとのローテーションが可能。ローカル provider(Ollama、LM Studio、VLLM)には不要 |
|
||||
| `api_base` | string | いいえ | デフォルトの API エンドポイント URL を上書き |
|
||||
| `proxy` | string | いいえ | このモデルエントリの HTTP プロキシ URL |
|
||||
| `user_agent` | string | いいえ | カスタム `User-Agent` リクエストヘッダー(OpenAI 互換、Anthropic、Azure provider で対応) |
|
||||
| `user_agent` | string | いいえ | カスタム `User-Agent` リクエストヘッダー(OpenAI 互換、Gemini、Anthropic、Azure provider で対応) |
|
||||
| `request_timeout` | int | いいえ | リクエストタイムアウト(秒)。デフォルト値は provider により異なる |
|
||||
| `max_tokens_field` | string | いいえ | リクエストボディの max tokens フィールド名を上書き(例:o1 モデルでは `max_completion_tokens`) |
|
||||
| `thinking_level` | string | いいえ | 拡張思考レベル:`off`、`low`、`medium`、`high`、`xhigh`、`adaptive` |
|
||||
|
|
@ -311,6 +311,7 @@ PicoClaw はリクエスト送信前に外側の `litellm/` プレフィック
|
|||
PicoClaw はプロトコルファミリーごとに Provider をルーティングします:
|
||||
|
||||
- OpenAI 互換プロトコル:OpenRouter、OpenAI 互換ゲートウェイ、Groq、Zhipu、vLLM スタイルのエンドポイント。
|
||||
- Gemini ネイティブプロトコル:Google Gemini のネイティブ `models/*:generateContent` / `models/*:streamGenerateContent` エンドポイント。
|
||||
- Anthropic プロトコル:Claude ネイティブ API 動作。
|
||||
- Codex/OAuth パス:OpenAI OAuth/Token 認証ルート。
|
||||
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@
|
|||
|
||||
### 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!**
|
||||
> **What's New?** PicoClaw now prefers explicit `provider` + native `model` configuration (for example `"provider": "zhipu", "model": "glm-4.7"`). The legacy single-field `provider/model` form remains supported for compatibility when `provider` is omitted.
|
||||
|
||||
For agent dispatch and light-model routing examples, see the [Routing Guide](routing-guide.md).
|
||||
|
||||
|
|
@ -46,35 +46,35 @@ This design also enables **multi-agent support** with flexible provider selectio
|
|||
|
||||
#### 📋 All Supported Vendors
|
||||
|
||||
| Vendor | `model` Prefix | Default API Base | Protocol | API Key |
|
||||
| Vendor | `provider` Value | Default API Base | Protocol | API Key |
|
||||
| ------------------- | ----------------- |-----------------------------------------------------| --------- | ---------------------------------------------------------------- |
|
||||
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [Get Key](https://platform.openai.com) |
|
||||
| **Venice AI** | `venice/` | `https://api.venice.ai/api/v1` | OpenAI | [Get Key](https://venice.ai) |
|
||||
| **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) |
|
||||
| **Z.AI Coding Plan** | `openai/` | `https://api.z.ai/api/coding/paas/v4` | OpenAI | [Get Key](https://z.ai/manage-apikey/apikey-list) |
|
||||
| **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) |
|
||||
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | Optional (local default: no key) |
|
||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [Get Key](https://cerebras.ai) |
|
||||
| **VolcEngine (Doubao)** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||
| **BytePlus** | `byteplus/` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [Get Key](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid/` | `https://api.vivgrid.com/v1` | OpenAI | [Get Key](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [Get Key](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)**| `modelscope/` | `https://api-inference.modelscope.cn/v1` | OpenAI | [Get Token](https://modelscope.cn/my/tokens) |
|
||||
| **Xiaomi MiMo** | `mimo/` | `https://api.xiaomimimo.com/v1` | OpenAI | [Get Key](https://platform.xiaomimimo.com) |
|
||||
| **Azure OpenAI** | `azure/` | `https://{resource}.openai.azure.com` | Azure | [Get Key](https://portal.azure.com) |
|
||||
| **Antigravity** | `antigravity/` | Google Cloud | Custom | OAuth only |
|
||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||
| **OpenAI** | `openai` | `https://api.openai.com/v1` | OpenAI | [Get Key](https://platform.openai.com) |
|
||||
| **Venice AI** | `venice` | `https://api.venice.ai/api/v1` | OpenAI | [Get Key](https://venice.ai) |
|
||||
| **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) |
|
||||
| **Z.AI Coding Plan** | `openai` | `https://api.z.ai/api/coding/paas/v4` | OpenAI | [Get Key](https://z.ai/manage-apikey/apikey-list) |
|
||||
| **DeepSeek** | `deepseek` | `https://api.deepseek.com/v1` | OpenAI | [Get Key](https://platform.deepseek.com) |
|
||||
| **Google Gemini** | `gemini` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [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) |
|
||||
| **LM Studio** | `lmstudio` | `http://localhost:1234/v1` | OpenAI | Optional (local default: no key) |
|
||||
| **OpenRouter** | `openrouter` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
||||
| **VLLM** | `vllm` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||
| **Cerebras** | `cerebras` | `https://api.cerebras.ai/v1` | OpenAI | [Get Key](https://cerebras.ai) |
|
||||
| **VolcEngine (Doubao)** | `volcengine` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [Get Key](https://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||
| **BytePlus** | `byteplus` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [Get Key](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid` | `https://api.vivgrid.com/v1` | OpenAI | [Get Key](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat` | `https://api.longcat.chat/openai` | OpenAI | [Get Key](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)**| `modelscope` | `https://api-inference.modelscope.cn/v1` | OpenAI | [Get Token](https://modelscope.cn/my/tokens) |
|
||||
| **Xiaomi MiMo** | `mimo` | `https://api.xiaomimimo.com/v1` | OpenAI | [Get Key](https://platform.xiaomimimo.com) |
|
||||
| **Azure OpenAI** | `azure` | `https://{resource}.openai.azure.com` | Azure | [Get Key](https://portal.azure.com) |
|
||||
| **Antigravity** | `antigravity` | Google Cloud | Custom | OAuth only |
|
||||
| **GitHub Copilot** | `github-copilot` | `localhost:4321` | gRPC | - |
|
||||
|
||||
#### Basic Configuration
|
||||
|
||||
|
|
@ -83,22 +83,26 @@ This design also enables **multi-agent support** with flexible provider selectio
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-your-api-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-your-openai-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-zhipu-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -115,11 +119,12 @@ This design also enables **multi-agent support** with flexible provider selectio
|
|||
| Field | Type | Required | Description |
|
||||
|-------|------|----------|-------------|
|
||||
| `model_name` | string | Yes | Unique name used to reference this model in agent config |
|
||||
| `model` | string | Yes | Vendor/model identifier (e.g., `openai/gpt-5.4`, `azure/gpt-5.4`, `anthropic/claude-sonnet-4.6`) |
|
||||
| `provider` | string | No | Preferred provider identifier. When present, PicoClaw sends `model` unchanged to that provider |
|
||||
| `model` | string | Yes | Native model ID when `provider` is set. If `provider` is omitted, the legacy `provider/model` form is still supported |
|
||||
| `api_keys` | string[] | Yes* | API key(s) for authentication. Multiple keys enable per-request rotation. Not required for local providers (Ollama, LM Studio, VLLM) |
|
||||
| `api_base` | string | No | Override the default API endpoint URL |
|
||||
| `proxy` | string | No | HTTP proxy URL for this model entry |
|
||||
| `user_agent` | string | No | Custom `User-Agent` header sent with API requests (supported by OpenAI-compatible, Anthropic, and Azure providers) |
|
||||
| `user_agent` | string | No | Custom `User-Agent` header sent with API requests (supported by OpenAI-compatible, Gemini, Anthropic, and Azure providers) |
|
||||
| `request_timeout` | int | No | Request timeout in seconds (default varies by provider) |
|
||||
| `max_tokens_field` | string | No | Override the max tokens field name in request body (e.g., `max_completion_tokens` for o1 models) |
|
||||
| `thinking_level` | string | No | Extended thinking level: `off`, `low`, `medium`, `high`, `xhigh`, or `adaptive` |
|
||||
|
|
@ -129,6 +134,22 @@ This design also enables **multi-agent support** with flexible provider selectio
|
|||
| `fallbacks` | string[] | No | Fallback model names for automatic failover |
|
||||
| `enabled` | bool | No | Whether this model entry is active (default: `true`) |
|
||||
|
||||
#### Provider / Model Resolution
|
||||
|
||||
PicoClaw resolves `provider` and the runtime model ID using these rules:
|
||||
|
||||
- If `provider` is set, `model` is used as-is.
|
||||
- If `provider` is omitted, PicoClaw treats the first `/` segment in `model` as the provider and everything after that first `/` as the runtime model ID.
|
||||
|
||||
Examples:
|
||||
|
||||
| Config | Resolved Provider | Model Sent Upstream |
|
||||
| --- | --- | --- |
|
||||
| `"provider": "openai", "model": "gpt-5.4"` | `openai` | `gpt-5.4` |
|
||||
| `"model": "openai/gpt-5.4"` | `openai` | `gpt-5.4` |
|
||||
| `"provider": "openrouter", "model": "openai/gpt-5.4"` | `openrouter` | `openai/gpt-5.4` |
|
||||
| `"model": "openrouter/openai/gpt-5.4"` | `openrouter` | `openai/gpt-5.4` |
|
||||
|
||||
#### Voice Transcription
|
||||
|
||||
You can configure a dedicated model for audio transcription with `voice.model_name`. This lets you reuse existing multimodal providers that support audio input instead of relying only on Groq.
|
||||
|
|
@ -140,7 +161,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "voice-gemini",
|
||||
"model": "gemini/gemini-2.5-flash",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.5-flash",
|
||||
"api_keys": ["your-gemini-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -163,7 +185,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
```json
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -173,7 +196,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
```json
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -183,7 +207,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
```json
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-key"]
|
||||
}
|
||||
```
|
||||
|
|
@ -193,7 +218,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
```json
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "openai/glm-4.7",
|
||||
"provider": "openai",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-z.ai-key"],
|
||||
"api_base": "https://api.z.ai/api/coding/paas/v4"
|
||||
}
|
||||
|
|
@ -204,7 +230,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
```json
|
||||
{
|
||||
"model_name": "deepseek-chat",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -214,7 +241,8 @@ If `voice.model_name` is not configured, PicoClaw will continue to fall back to
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
}
|
||||
```
|
||||
|
|
@ -228,7 +256,8 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-opus-4-6",
|
||||
"model": "anthropic-messages/claude-opus-4-6",
|
||||
"provider": "anthropic-messages",
|
||||
"model": "claude-opus-4-6",
|
||||
"api_keys": ["sk-ant-your-key"],
|
||||
"api_base": "https://api.anthropic.com"
|
||||
}
|
||||
|
|
@ -246,7 +275,8 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
|||
```json
|
||||
{
|
||||
"model_name": "llama3",
|
||||
"model": "ollama/llama3"
|
||||
"provider": "ollama",
|
||||
"model": "llama3"
|
||||
}
|
||||
```
|
||||
|
||||
|
|
@ -255,19 +285,21 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
|||
```json
|
||||
{
|
||||
"model_name": "lmstudio-local",
|
||||
"model": "lmstudio/openai/gpt-oss-20b"
|
||||
"provider": "lmstudio",
|
||||
"model": "openai/gpt-oss-20b"
|
||||
}
|
||||
```
|
||||
|
||||
`api_base` defaults to `http://localhost:1234/v1`. API key is optional unless your LM Studio server enables authentication.<br/>
|
||||
PicoClaw sends OpenAI-compatible requests to LM Studio, and strips the `lmstudio/` prefix before sending requests, so `lmstudio/openai/gpt-oss-20b` sends `openai/gpt-oss-20b` to the LM Studio server.
|
||||
With explicit `provider`, PicoClaw sends `openai/gpt-oss-20b` unchanged to the LM Studio server. The legacy compatibility form `"model": "lmstudio/openai/gpt-oss-20b"` still resolves to the same upstream model ID when `provider` is omitted.
|
||||
|
||||
**Custom Proxy/API**
|
||||
|
||||
```json
|
||||
{
|
||||
"model_name": "my-custom-model",
|
||||
"model": "openai/custom-model",
|
||||
"provider": "openai",
|
||||
"model": "custom-model",
|
||||
"api_base": "https://my-proxy.com/v1",
|
||||
"api_keys": ["sk-..."],
|
||||
"user_agent": "MyApp/1.0",
|
||||
|
|
@ -280,13 +312,14 @@ PicoClaw sends OpenAI-compatible requests to LM Studio, and strips the `lmstudio
|
|||
```json
|
||||
{
|
||||
"model_name": "lite-gpt4",
|
||||
"model": "litellm/lite-gpt4",
|
||||
"provider": "litellm",
|
||||
"model": "lite-gpt4",
|
||||
"api_base": "http://localhost:4000/v1",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
||||
PicoClaw strips only the outer `litellm/` prefix before sending the request, so proxy aliases like `litellm/lite-gpt4` send `lite-gpt4`, while `litellm/openai/gpt-4o` sends `openai/gpt-4o`.
|
||||
With explicit `provider`, PicoClaw sends `model` unchanged. That means `"provider": "litellm", "model": "lite-gpt4"` sends `lite-gpt4`, while `"provider": "litellm", "model": "openai/gpt-4o"` sends `openai/gpt-4o`. The legacy compatibility forms `litellm/lite-gpt4` and `litellm/openai/gpt-4o` still resolve the same way when `provider` is omitted.
|
||||
|
||||
**Z.AI Coding Plan**
|
||||
|
||||
|
|
@ -295,7 +328,8 @@ If the standard Zhipu endpoint (`https://open.bigmodel.cn/api/paas/v4`) returns
|
|||
```json
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "openai/glm-4.7",
|
||||
"provider": "openai",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-zhipu-api-key"],
|
||||
"api_base": "https://api.z.ai/api/coding/paas/v4"
|
||||
}
|
||||
|
|
@ -312,13 +346,15 @@ Configure multiple endpoints for the same model name—PicoClaw will automatical
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api1.example.com/v1",
|
||||
"api_keys": ["sk-key1"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api2.example.com/v1",
|
||||
"api_keys": ["sk-key2"]
|
||||
}
|
||||
|
|
@ -337,18 +373,21 @@ It also applies cooldown tracking per candidate to avoid immediately retrying a
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "qwen-main",
|
||||
"model": "openai/qwen3.5:cloud",
|
||||
"provider": "openai",
|
||||
"model": "qwen3.5:cloud",
|
||||
"api_base": "https://api.example.com/v1",
|
||||
"api_keys": ["sk-main"]
|
||||
},
|
||||
{
|
||||
"model_name": "deepseek-backup",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat",
|
||||
"api_keys": ["sk-backup-1"]
|
||||
},
|
||||
{
|
||||
"model_name": "gemini-backup",
|
||||
"model": "gemini/gemini-2.5-flash",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.5-flash",
|
||||
"api_keys": ["sk-backup-2"]
|
||||
}
|
||||
],
|
||||
|
|
@ -396,7 +435,8 @@ The old `providers` configuration is **deprecated** and has been removed in V2.
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -415,10 +455,11 @@ For detailed migration guide, see [migration/model-list-migration.md](../migrati
|
|||
PicoClaw routes providers by protocol family:
|
||||
|
||||
- OpenAI-compatible protocol: OpenRouter, OpenAI-compatible gateways, Groq, Zhipu, and vLLM-style endpoints.
|
||||
- Gemini native protocol: Google Gemini via the native `models/*:generateContent` and `models/*:streamGenerateContent` 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`).
|
||||
This keeps the runtime lightweight while making new OpenAI-compatible backends mostly a config operation (`api_base` + `api_keys`).
|
||||
|
||||
<details>
|
||||
<summary><b>Zhipu</b></summary>
|
||||
|
|
@ -464,7 +505,7 @@ picoclaw agent -m "Hello"
|
|||
{
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"model_name": "anthropic/claude-opus-4-5"
|
||||
"model_name": "claude-opus-4-5"
|
||||
}
|
||||
},
|
||||
"session": {
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ Este design também permite **suporte multi-agente** com seleção flexível de
|
|||
| **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) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [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) |
|
||||
|
|
@ -108,7 +108,7 @@ Este design também permite **suporte multi-agente** com seleção flexível de
|
|||
| `api_keys` | string[] | Sim* | Chave(s) API para autenticação. Múltiplas chaves permitem rotação por requisição. Não necessário para providers locais (Ollama, LM Studio, VLLM) |
|
||||
| `api_base` | string | Não | Substitui a URL base da API padrão |
|
||||
| `proxy` | string | Não | URL do proxy HTTP para esta entrada de modelo |
|
||||
| `user_agent` | string | Não | Cabeçalho `User-Agent` personalizado enviado com requisições API (suportado por providers OpenAI-compatible, Anthropic e Azure) |
|
||||
| `user_agent` | string | Não | Cabeçalho `User-Agent` personalizado enviado com requisições API (suportado por providers OpenAI-compatible, Gemini, Anthropic e Azure) |
|
||||
| `request_timeout` | int | Não | Timeout de requisição em segundos (o padrão varia por provider) |
|
||||
| `max_tokens_field` | string | Não | Substitui o nome do campo max tokens no corpo da requisição (ex: `max_completion_tokens` para modelos o1) |
|
||||
| `thinking_level` | string | Não | Nível de pensamento estendido: `off`, `low`, `medium`, `high`, `xhigh` ou `adaptive` |
|
||||
|
|
@ -299,6 +299,7 @@ Para guia de migração detalhado, veja [migration/model-list-migration.md](../m
|
|||
O PicoClaw roteia provedores por família de protocolo:
|
||||
|
||||
- Protocolo compatível com OpenAI: OpenRouter, gateways compatíveis com OpenAI, Groq, Zhipu e endpoints estilo vLLM.
|
||||
- Protocolo Gemini nativo: Google Gemini via endpoints nativos `models/*:generateContent` e `models/*:streamGenerateContent`.
|
||||
- Protocolo Anthropic: Comportamento nativo da API Claude.
|
||||
- Caminho Codex/OAuth: Rota de autenticação OAuth/token da OpenAI.
|
||||
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ Thiết kế này cũng cho phép **hỗ trợ đa agent** với lựa chọn pr
|
|||
| **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) |
|
||||
| **Google Gemini** | `gemini/` | `https://generativelanguage.googleapis.com/v1beta` | Gemini | [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) |
|
||||
|
|
@ -108,7 +108,7 @@ Thiết kế này cũng cho phép **hỗ trợ đa agent** với lựa chọn pr
|
|||
| `api_keys` | string[] | Có* | Khóa API xác thực. Nhiều khóa cho phép xoay vòng theo yêu cầu. Không cần thiết cho provider nội bộ (Ollama, LM Studio, VLLM) |
|
||||
| `api_base` | string | Không | Ghi đè URL endpoint API mặc định |
|
||||
| `proxy` | string | Không | URL proxy HTTP cho entry model này |
|
||||
| `user_agent` | string | Không | Header `User-Agent` tùy chỉnh gửi với yêu cầu API (được hỗ trợ bởi provider OpenAI-compatible, Anthropic và Azure) |
|
||||
| `user_agent` | string | Không | Header `User-Agent` tùy chỉnh gửi với yêu cầu API (được hỗ trợ bởi provider OpenAI-compatible, Gemini, Anthropic và Azure) |
|
||||
| `request_timeout` | int | Không | Timeout yêu cầu tính bằng giây (mặc định khác nhau tùy provider) |
|
||||
| `max_tokens_field` | string | Không | Ghi đè tên trường max tokens trong request body (ví dụ: `max_completion_tokens` cho model o1) |
|
||||
| `thinking_level` | string | Không | Mức độ tư duy mở rộng: `off`, `low`, `medium`, `high`, `xhigh` hoặc `adaptive` |
|
||||
|
|
@ -299,6 +299,7 @@ Cấu hình `providers` cũ đã **bị deprecated** và đã được loại b
|
|||
PicoClaw định tuyến provider theo họ giao thức:
|
||||
|
||||
- Giao thức tương thích OpenAI: OpenRouter, gateway tương thích OpenAI, Groq, Zhipu, và endpoint kiểu vLLM.
|
||||
- Giao thức Gemini native: Google Gemini qua các endpoint native `models/*:generateContent` và `models/*:streamGenerateContent`.
|
||||
- Giao thức Anthropic: Hành vi API native của Claude.
|
||||
- Đường dẫn Codex/OAuth: Tuyến xác thực OAuth/token của OpenAI.
|
||||
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@
|
|||
<a id="模型配置-model_list"></a>
|
||||
### 模型配置 (model_list)
|
||||
|
||||
> **新功能!** PicoClaw 现在采用**以模型为中心**的配置方式。只需使用 `厂商/模型` 格式(如 `zhipu/glm-4.7`)即可添加新的 provider——**无需修改任何代码!**
|
||||
> **新功能!** PicoClaw 现在优先推荐显式 `provider` + 原生 `model` 的配置方式,例如 `"provider": "zhipu", "model": "glm-4.7"`。如果未设置 `provider`,旧的单字段 `provider/model` 写法仍然兼容。
|
||||
|
||||
如果你想看 agent 分发和轻量模型路由的完整示例,请看 [路由使用指南](routing-guide.zh.md)。
|
||||
|
||||
|
|
@ -45,33 +45,33 @@
|
|||
|
||||
#### 📋 所有支持的厂商
|
||||
|
||||
| 厂商 | `model` 前缀 | 默认 API Base | 协议 | 获取 API Key |
|
||||
| 厂商 | `provider` 值 | 默认 API Base | 协议 | 获取 API Key |
|
||||
| ------------------- | ----------------- | --------------------------------------------------- | --------- | ----------------------------------------------------------------- |
|
||||
| **OpenAI** | `openai/` | `https://api.openai.com/v1` | OpenAI | [获取密钥](https://platform.openai.com) |
|
||||
| **Venice AI** | `venice/` | `https://api.venice.ai/api/v1` | OpenAI | [获取密钥](https://venice.ai) |
|
||||
| **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 | 本地(无需密钥) |
|
||||
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | 可选(本地默认无需密钥) |
|
||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取密钥](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理密钥 |
|
||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
||||
| **Cerebras** | `cerebras/` | `https://api.cerebras.ai/v1` | OpenAI | [获取密钥](https://cerebras.ai) |
|
||||
| **火山引擎(Doubao)** | `volcengine/` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [获取密钥](https://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||
| **BytePlus** | `byteplus/` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [获取密钥](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid/` | `https://api.vivgrid.com/v1` | OpenAI | [获取密钥](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat/` | `https://api.longcat.chat/openai` | OpenAI | [获取密钥](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)**| `modelscope/` | `https://api-inference.modelscope.cn/v1` | OpenAI | [获取 Token](https://modelscope.cn/my/tokens) |
|
||||
| **小米 MiMo** | `mimo/` | `https://api.xiaomimimo.com/v1` | OpenAI | [获取密钥](https://platform.xiaomimimo.com) |
|
||||
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
|
||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||
| **OpenAI** | `openai` | `https://api.openai.com/v1` | OpenAI | [获取密钥](https://platform.openai.com) |
|
||||
| **Venice AI** | `venice` | `https://api.venice.ai/api/v1` | OpenAI | [获取密钥](https://venice.ai) |
|
||||
| **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` | Gemini | [获取密钥](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 | 本地(无需密钥) |
|
||||
| **LM Studio** | `lmstudio` | `http://localhost:1234/v1` | OpenAI | 可选(本地默认无需密钥) |
|
||||
| **OpenRouter** | `openrouter` | `https://openrouter.ai/api/v1` | OpenAI | [获取密钥](https://openrouter.ai/keys) |
|
||||
| **LiteLLM Proxy** | `litellm` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理密钥 |
|
||||
| **VLLM** | `vllm` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
||||
| **Cerebras** | `cerebras` | `https://api.cerebras.ai/v1` | OpenAI | [获取密钥](https://cerebras.ai) |
|
||||
| **火山引擎(Doubao)** | `volcengine` | `https://ark.cn-beijing.volces.com/api/v3` | OpenAI | [获取密钥](https://www.volcengine.com/activity/codingplan?utm_campaign=PicoClaw&utm_content=PicoClaw&utm_medium=devrel&utm_source=OWO&utm_term=PicoClaw) |
|
||||
| **神算云** | `shengsuanyun` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||
| **BytePlus** | `byteplus` | `https://ark.ap-southeast.bytepluses.com/api/v3` | OpenAI | [获取密钥](https://www.byteplus.com) |
|
||||
| **Vivgrid** | `vivgrid` | `https://api.vivgrid.com/v1` | OpenAI | [获取密钥](https://vivgrid.com) |
|
||||
| **LongCat** | `longcat` | `https://api.longcat.chat/openai` | OpenAI | [获取密钥](https://longcat.chat/platform) |
|
||||
| **ModelScope (魔搭)**| `modelscope` | `https://api-inference.modelscope.cn/v1` | OpenAI | [获取 Token](https://modelscope.cn/my/tokens) |
|
||||
| **小米 MiMo** | `mimo` | `https://api.xiaomimimo.com/v1` | OpenAI | [获取密钥](https://platform.xiaomimimo.com) |
|
||||
| **Antigravity** | `antigravity` | Google Cloud | 自定义 | 仅 OAuth |
|
||||
| **GitHub Copilot** | `github-copilot` | `localhost:4321` | gRPC | - |
|
||||
|
||||
#### 基础配置示例
|
||||
|
||||
|
|
@ -80,22 +80,26 @@
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-your-api-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-your-openai-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-zhipu-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -112,11 +116,12 @@
|
|||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `model_name` | string | 是 | 在 agent 配置中引用此模型的唯一名称 |
|
||||
| `model` | string | 是 | 厂商/模型标识符(如 `openai/gpt-5.4`、`azure/gpt-5.4`、`anthropic/claude-sonnet-4.6`) |
|
||||
| `provider` | string | 否 | 推荐的 provider 标识。设置后,PicoClaw 会将 `model` 原样发送给该 provider |
|
||||
| `model` | string | 是 | 当设置 `provider` 时,这里填写 provider 原生模型 ID。若未设置 `provider`,仍兼容旧的 `provider/model` 写法 |
|
||||
| `api_keys` | string[] | 是* | 认证密钥。多个密钥可按请求轮换。本地 provider(Ollama、LM Studio、VLLM)不需要 |
|
||||
| `api_base` | string | 否 | 覆盖默认的 API 端点 URL |
|
||||
| `proxy` | string | 否 | 此模型条目的 HTTP 代理 URL |
|
||||
| `user_agent` | string | 否 | 自定义 `User-Agent` 请求头(支持 OpenAI 兼容、Anthropic 和 Azure provider) |
|
||||
| `user_agent` | string | 否 | 自定义 `User-Agent` 请求头(支持 OpenAI 兼容、Gemini、Anthropic 和 Azure provider) |
|
||||
| `request_timeout` | int | 否 | 请求超时时间(秒),默认值因 provider 而异 |
|
||||
| `max_tokens_field` | string | 否 | 覆盖请求体中 max tokens 的字段名(如 o1 模型使用 `max_completion_tokens`) |
|
||||
| `thinking_level` | string | 否 | 扩展思考级别:`off`、`low`、`medium`、`high`、`xhigh` 或 `adaptive` |
|
||||
|
|
@ -126,6 +131,22 @@
|
|||
| `fallbacks` | string[] | 否 | 自动故障转移的备用模型名称 |
|
||||
| `enabled` | bool | 否 | 是否启用此模型条目(默认:`true`) |
|
||||
|
||||
#### `provider` / `model` 解析规则
|
||||
|
||||
PicoClaw 按下面的规则解析 `provider` 和最终发给上游的模型 ID:
|
||||
|
||||
- 如果设置了 `provider`,则直接使用 `model`。
|
||||
- 如果未设置 `provider`,则把 `model` 中第一个 `/` 之前的字段当作 provider,第一个 `/` 之后的全部内容当作最终模型 ID。
|
||||
|
||||
示例:
|
||||
|
||||
| 配置 | 解析后的 Provider | 实际发送的模型 ID |
|
||||
| --- | --- | --- |
|
||||
| `"provider": "openai", "model": "gpt-5.4"` | `openai` | `gpt-5.4` |
|
||||
| `"model": "openai/gpt-5.4"` | `openai` | `gpt-5.4` |
|
||||
| `"provider": "openrouter", "model": "openai/gpt-5.4"` | `openrouter` | `openai/gpt-5.4` |
|
||||
| `"model": "openrouter/openai/gpt-5.4"` | `openrouter` | `openai/gpt-5.4` |
|
||||
|
||||
#### 语音转录
|
||||
|
||||
你可以通过 `voice.model_name` 为语音转录指定一个专用模型。这样可以直接复用已经配置好的、支持音频输入的多模态 provider,而不必只依赖 Groq。
|
||||
|
|
@ -137,7 +158,8 @@
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "voice-gemini",
|
||||
"model": "gemini/gemini-2.5-flash",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.5-flash",
|
||||
"api_keys": ["your-gemini-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -160,7 +182,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -170,7 +193,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "ark-code-latest",
|
||||
"model": "volcengine/ark-code-latest",
|
||||
"provider": "volcengine",
|
||||
"model": "ark-code-latest",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -180,7 +204,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-key"]
|
||||
}
|
||||
```
|
||||
|
|
@ -190,7 +215,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "deepseek-chat",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
|
@ -200,7 +226,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"auth_method": "oauth"
|
||||
}
|
||||
```
|
||||
|
|
@ -214,7 +241,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "claude-opus-4-6",
|
||||
"model": "anthropic-messages/claude-opus-4-6",
|
||||
"provider": "anthropic-messages",
|
||||
"model": "claude-opus-4-6",
|
||||
"api_keys": ["sk-ant-your-key"],
|
||||
"api_base": "https://api.anthropic.com"
|
||||
}
|
||||
|
|
@ -232,7 +260,8 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "llama3",
|
||||
"model": "ollama/llama3"
|
||||
"provider": "ollama",
|
||||
"model": "llama3"
|
||||
}
|
||||
```
|
||||
|
||||
|
|
@ -241,19 +270,21 @@
|
|||
```json
|
||||
{
|
||||
"model_name": "lmstudio-local",
|
||||
"model": "lmstudio/openai/gpt-oss-20b"
|
||||
"provider": "lmstudio",
|
||||
"model": "openai/gpt-oss-20b"
|
||||
}
|
||||
```
|
||||
|
||||
`api_base` 默认是 `http://localhost:1234/v1`。除非你在 LM Studio 侧启用了认证,否则不需要配置 API Key。
|
||||
PicoClaw 向 LM Studio 的 OpenAI 兼容终结点发送请求,且将移除首个 `lmstudio/` 前缀,因此 `lmstudio/openai/gpt-oss-20b` 会发送 `openai/gpt-oss-20b`。
|
||||
显式设置 `provider` 后,PicoClaw 会把 `openai/gpt-oss-20b` 原样发送给 LM Studio。旧的兼容写法 `"model": "lmstudio/openai/gpt-oss-20b"` 在未设置 `provider` 时也会解析成相同的上游模型 ID。
|
||||
|
||||
**自定义代理/API**
|
||||
|
||||
```json
|
||||
{
|
||||
"model_name": "my-custom-model",
|
||||
"model": "openai/custom-model",
|
||||
"provider": "openai",
|
||||
"model": "custom-model",
|
||||
"api_base": "https://my-proxy.com/v1",
|
||||
"api_keys": ["sk-..."],
|
||||
"user_agent": "MyApp/1.0",
|
||||
|
|
@ -266,13 +297,14 @@ PicoClaw 向 LM Studio 的 OpenAI 兼容终结点发送请求,且将移除首
|
|||
```json
|
||||
{
|
||||
"model_name": "lite-gpt4",
|
||||
"model": "litellm/lite-gpt4",
|
||||
"provider": "litellm",
|
||||
"model": "lite-gpt4",
|
||||
"api_base": "http://localhost:4000/v1",
|
||||
"api_keys": ["sk-..."]
|
||||
}
|
||||
```
|
||||
|
||||
PicoClaw 在发送请求前仅去除外层 `litellm/` 前缀,因此 `litellm/lite-gpt4` 会发送 `lite-gpt4`,而 `litellm/openai/gpt-4o` 会发送 `openai/gpt-4o`。
|
||||
显式设置 `provider` 后,PicoClaw 会将 `model` 原样发送。因此 `"provider": "litellm", "model": "lite-gpt4"` 会发送 `lite-gpt4`,而 `"provider": "litellm", "model": "openai/gpt-4o"` 会发送 `openai/gpt-4o`。旧的兼容写法 `litellm/lite-gpt4` 和 `litellm/openai/gpt-4o` 在未设置 `provider` 时也会得到相同结果。
|
||||
|
||||
#### 负载均衡
|
||||
|
||||
|
|
@ -283,13 +315,15 @@ PicoClaw 在发送请求前仅去除外层 `litellm/` 前缀,因此 `litellm/l
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api1.example.com/v1",
|
||||
"api_keys": ["sk-key1"]
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_base": "https://api2.example.com/v1",
|
||||
"api_keys": ["sk-key2"]
|
||||
}
|
||||
|
|
@ -308,18 +342,21 @@ PicoClaw 在发送请求前仅去除外层 `litellm/` 前缀,因此 `litellm/l
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "qwen-main",
|
||||
"model": "openai/qwen3.5:cloud",
|
||||
"provider": "openai",
|
||||
"model": "qwen3.5:cloud",
|
||||
"api_base": "https://api.example.com/v1",
|
||||
"api_keys": ["sk-main"]
|
||||
},
|
||||
{
|
||||
"model_name": "deepseek-backup",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat",
|
||||
"api_keys": ["sk-backup-1"]
|
||||
},
|
||||
{
|
||||
"model_name": "gemini-backup",
|
||||
"model": "gemini/gemini-2.5-flash",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.5-flash",
|
||||
"api_keys": ["sk-backup-2"]
|
||||
}
|
||||
],
|
||||
|
|
@ -367,7 +404,8 @@ PicoClaw 在发送请求前仅去除外层 `litellm/` 前缀,因此 `litellm/l
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "glm-4.7",
|
||||
"model": "zhipu/glm-4.7",
|
||||
"provider": "zhipu",
|
||||
"model": "glm-4.7",
|
||||
"api_keys": ["your-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -386,10 +424,11 @@ PicoClaw 在发送请求前仅去除外层 `litellm/` 前缀,因此 `litellm/l
|
|||
PicoClaw 按协议族路由 Provider:
|
||||
|
||||
- OpenAI 兼容协议:OpenRouter、OpenAI 兼容网关、Groq、智谱、vLLM 风格端点。
|
||||
- Gemini 原生协议:Google Gemini 通过原生 `models/*:generateContent` 和 `models/*:streamGenerateContent` 端点接入。
|
||||
- Anthropic 协议:Claude 原生 API 行为。
|
||||
- Codex/OAuth 路径:OpenAI OAuth/Token 认证路由。
|
||||
|
||||
这使得运行时保持轻量,同时让新的 OpenAI 兼容后端基本只需配置操作(`api_base` + `api_key`)。
|
||||
这使得运行时保持轻量,同时让新的 OpenAI 兼容后端基本只需配置操作(`api_base` + `api_keys`)。
|
||||
|
||||
<details>
|
||||
<summary><b>智谱 (Zhipu) 配置示例</b></summary>
|
||||
|
|
@ -435,7 +474,7 @@ picoclaw agent -m "你好"
|
|||
{
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"model_name": "anthropic/claude-opus-4-5"
|
||||
"model_name": "claude-opus-4-5"
|
||||
}
|
||||
},
|
||||
"session": {
|
||||
|
|
|
|||
|
|
@ -69,12 +69,14 @@ This guide explains how to configure both for real deployments.
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-main",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-main"]
|
||||
},
|
||||
{
|
||||
"model_name": "flash-light",
|
||||
"model": "gemini/gemini-2.0-flash-exp",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.0-flash-exp",
|
||||
"api_keys": ["sk-light"]
|
||||
}
|
||||
],
|
||||
|
|
|
|||
|
|
@ -69,12 +69,14 @@ PicoClaw 里用户能直接感知到的“路由”主要有两部分:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-main",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-main"]
|
||||
},
|
||||
{
|
||||
"model_name": "flash-light",
|
||||
"model": "gemini/gemini-2.0-flash-exp",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.0-flash-exp",
|
||||
"api_keys": ["sk-light"]
|
||||
}
|
||||
],
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ 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.
|
||||
- **Explicit provider resolution**: Prefer `provider` + native `model`, with legacy `provider/model` compatibility when needed
|
||||
- **Cleaner configuration**: Model-centric instead of vendor-centric
|
||||
|
||||
## Timeline
|
||||
|
|
@ -54,18 +54,21 @@ The new `model_list` configuration offers several advantages:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-your-openai-key"],
|
||||
"api_base": "https://api.openai.com/v1"
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4.6",
|
||||
"model": "anthropic/claude-sonnet-4.6",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4.6",
|
||||
"api_keys": ["sk-ant-your-key"]
|
||||
},
|
||||
{
|
||||
"model_name": "deepseek",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-chat",
|
||||
"api_keys": ["sk-your-deepseek-key"]
|
||||
}
|
||||
],
|
||||
|
|
@ -79,40 +82,46 @@ The new `model_list` configuration offers several advantages:
|
|||
|
||||
> **Note**: The `enabled` field can be omitted — during V1→V2 migration it is auto-inferred (models with API keys or the `local-model` name are enabled by default). For new configs, you can explicitly set `"enabled": false` to disable a model entry without removing it.
|
||||
|
||||
## Protocol Prefixes
|
||||
## Provider / Model Resolution
|
||||
|
||||
The `model` field uses a protocol prefix format: `[protocol/]model-identifier`
|
||||
Preferred format:
|
||||
|
||||
| Prefix | Description | Example |
|
||||
|--------|-------------|---------|
|
||||
| `openai/` | OpenAI API (default) | `openai/gpt-5.4` |
|
||||
| `anthropic/` | Anthropic API | `anthropic/claude-opus-4` |
|
||||
| `antigravity/` | Google via Antigravity OAuth | `antigravity/gemini-2.0-flash` |
|
||||
| `gemini/` | Google Gemini API | `gemini/gemini-2.0-flash-exp` |
|
||||
| `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` |
|
||||
| `zhipu/` | Zhipu AI | `zhipu/glm-4` |
|
||||
| `nvidia/` | NVIDIA NIM | `nvidia/llama-3.1-nemotron-70b` |
|
||||
| `ollama/` | Ollama (local) | `ollama/llama3` |
|
||||
| `vllm/` | vLLM (local) | `vllm/my-model` |
|
||||
| `moonshot/` | Moonshot AI | `moonshot/moonshot-v1-8k` |
|
||||
| `shengsuanyun/` | ShengSuanYun | `shengsuanyun/deepseek-v3` |
|
||||
| `volcengine/` | Volcengine | `volcengine/doubao-pro-32k` |
|
||||
```json
|
||||
{
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4"
|
||||
}
|
||||
```
|
||||
|
||||
**Note**: If no prefix is specified, `openai/` is used as the default.
|
||||
Legacy compatibility format:
|
||||
|
||||
```json
|
||||
{
|
||||
"model": "openai/gpt-5.4"
|
||||
}
|
||||
```
|
||||
|
||||
Resolution rules:
|
||||
|
||||
1. If `provider` is set, PicoClaw sends `model` unchanged.
|
||||
2. If `provider` is omitted, PicoClaw treats the first `/` segment in `model` as the provider and everything after that first `/` as the runtime model ID.
|
||||
|
||||
Examples:
|
||||
|
||||
| Config | Resolved Provider | Model Sent Upstream |
|
||||
|--------|-------------------|---------------------|
|
||||
| `"provider": "openai", "model": "gpt-5.4"` | `openai` | `gpt-5.4` |
|
||||
| `"model": "openai/gpt-5.4"` | `openai` | `gpt-5.4` |
|
||||
| `"provider": "openrouter", "model": "google/gemini-2.0-flash-exp:free"` | `openrouter` | `google/gemini-2.0-flash-exp:free` |
|
||||
| `"model": "openrouter/google/gemini-2.0-flash-exp:free"` | `openrouter` | `google/gemini-2.0-flash-exp:free` |
|
||||
|
||||
## 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.4`) |
|
||||
| `provider` | No | Preferred provider identifier. When set, `model` is sent unchanged |
|
||||
| `model` | Yes | Native model ID when `provider` is set, or legacy `provider/model` when `provider` is omitted |
|
||||
| `api_base` | No | API endpoint URL |
|
||||
| `api_keys` | No | API authentication keys (array; supports multiple keys for load balancing) |
|
||||
| `enabled` | No | Whether this model entry is active. Defaults to `true` during migration for models with API keys or named `local-model`. Set to `false` to disable. |
|
||||
|
|
@ -136,7 +145,8 @@ There are two ways to configure load balancing:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-key1", "sk-key2", "sk-key3"],
|
||||
"api_base": "https://api.openai.com/v1"
|
||||
}
|
||||
|
|
@ -162,19 +172,22 @@ model_list:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-key1"],
|
||||
"api_base": "https://api1.example.com/v1"
|
||||
},
|
||||
{
|
||||
"model_name": "gpt4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-key2"],
|
||||
"api_base": "https://api2.example.com/v1"
|
||||
},
|
||||
{
|
||||
"model_name": "gpt4",
|
||||
"model": "openai/gpt-5.4",
|
||||
"provider": "openai",
|
||||
"model": "gpt-5.4",
|
||||
"api_keys": ["sk-key3"],
|
||||
"api_base": "https://api3.example.com/v1"
|
||||
}
|
||||
|
|
@ -193,7 +206,8 @@ With `model_list`, adding a new provider requires zero code changes:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "my-custom-llm",
|
||||
"model": "openai/my-model-v1",
|
||||
"provider": "openai",
|
||||
"model": "my-model-v1",
|
||||
"api_keys": ["your-api-key"],
|
||||
"api_base": "https://api.your-provider.com/v1"
|
||||
}
|
||||
|
|
@ -201,7 +215,7 @@ With `model_list`, adding a new provider requires zero code changes:
|
|||
}
|
||||
```
|
||||
|
||||
Just specify `openai/` as the protocol (or omit it for the default), and provide your provider's API base URL.
|
||||
Just set `provider` to `openai` (or another supported provider), and provide your provider's API base URL.
|
||||
|
||||
## Backward Compatibility
|
||||
|
||||
|
|
@ -216,7 +230,7 @@ During the migration period, your existing V0/V1 config will be auto-migrated to
|
|||
|
||||
- [ ] Identify all providers you're currently using
|
||||
- [ ] Create `model_list` entries for each provider
|
||||
- [ ] Use appropriate protocol prefixes
|
||||
- [ ] Prefer explicit `provider` values and native model IDs
|
||||
- [ ] Update `agents.defaults.model_name` to reference the new `model_name`
|
||||
- [ ] Test that all models work correctly
|
||||
- [ ] Remove or comment out the old `providers` section
|
||||
|
|
@ -234,10 +248,10 @@ model "xxx" not found in model_list or providers
|
|||
### Unknown protocol error
|
||||
|
||||
```
|
||||
unknown protocol "xxx" in model "xxx/model-name"
|
||||
unknown provider "xxx" in model "xxx/model-name"
|
||||
```
|
||||
|
||||
**Solution**: Use a supported protocol prefix. See the [Protocol Prefixes](#protocol-prefixes) table above.
|
||||
**Solution**: Use a supported `provider` value, or use the legacy `provider/model` compatibility form correctly. See [Provider / Model Resolution](#provider--model-resolution).
|
||||
|
||||
### Missing API key error
|
||||
|
||||
|
|
|
|||
|
|
@ -7,16 +7,22 @@
|
|||
- `Error creating provider: model "openrouter/free" not found in model_list`
|
||||
- OpenRouter returns 400: `"free is not a valid model ID"`
|
||||
|
||||
**Cause:** The `model` field in your `model_list` entry is what gets sent to the API. For OpenRouter you must use the **full** model ID, not a shorthand.
|
||||
**Cause:** PicoClaw now resolves provider/model in two steps:
|
||||
|
||||
- **Wrong:** `"model": "free"` → OpenRouter receives `free` and rejects it.
|
||||
- **Right:** `"model": "openrouter/free"` → OpenRouter receives `openrouter/free` (auto free-tier routing).
|
||||
- If `provider` is set, the `model` field is sent to that provider unchanged.
|
||||
- If `provider` is omitted, PicoClaw infers the provider from the first `/` segment and sends everything after that first `/` as the runtime model ID.
|
||||
|
||||
For OpenRouter free-tier routing, the preferred config is explicit `provider`.
|
||||
|
||||
- **Wrong:** `"model": "free"` → no OpenRouter provider is selected, so `free` is not a valid OpenRouter model route.
|
||||
- **Right:** `"provider": "openrouter", "model": "free"` → OpenRouter receives `free`.
|
||||
- **Also supported:** `"model": "openrouter/free"` → provider resolves to `openrouter`, runtime model ID resolves to `free`.
|
||||
|
||||
**Fix:** In `~/.picoclaw/config.json` (or your config path):
|
||||
|
||||
1. **agents.defaults.model_name** must match a `model_name` in `model_list` (e.g. `"openrouter-free"`).
|
||||
2. That entry’s **model** must be a valid OpenRouter model ID, for example:
|
||||
- `"openrouter/free"` – auto free-tier
|
||||
2. That entry should preferably set **provider** to `openrouter`, and **model** should be a valid OpenRouter model ID, for example:
|
||||
- `"free"` – auto free-tier
|
||||
- `"google/gemini-2.0-flash-exp:free"`
|
||||
- `"meta-llama/llama-3.1-8b-instruct:free"`
|
||||
|
||||
|
|
@ -32,8 +38,9 @@ Example snippet:
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "openrouter-free",
|
||||
"model": "openrouter/free",
|
||||
"api_key": "sk-or-v1-YOUR_OPENROUTER_KEY",
|
||||
"provider": "openrouter",
|
||||
"model": "free",
|
||||
"api_keys": ["sk-or-v1-YOUR_OPENROUTER_KEY"],
|
||||
"api_base": "https://openrouter.ai/api/v1"
|
||||
}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -9,16 +9,22 @@
|
|||
- `Error creating provider: model "openrouter/free" not found in model_list`
|
||||
- OpenRouter 返回 400:`"free is not a valid model ID"`
|
||||
|
||||
**原因:** `model_list` 条目中的 `model` 字段是发送给 API 的内容。对于 OpenRouter,你必须使用**完整的**模型 ID,而不是简写。
|
||||
**原因:** PicoClaw 现在按两步解析 provider 和 model:
|
||||
|
||||
- **错误:** `"model": "free"` → OpenRouter 收到 `free` 并拒绝。
|
||||
- **正确:** `"model": "openrouter/free"` → OpenRouter 收到 `openrouter/free`(自动免费层路由)。
|
||||
- 如果设置了 `provider`,则会把 `model` 原样发送给该 provider。
|
||||
- 如果未设置 `provider`,则会把 `model` 第一个 `/` 之前的字段当作 provider,并把第一个 `/` 之后的全部内容当作最终发送的模型 ID。
|
||||
|
||||
对于 OpenRouter 免费层路由,推荐显式设置 `provider`。
|
||||
|
||||
- **错误:** `"model": "free"` → 不会选中 OpenRouter,`free` 也不是可直接路由的 OpenRouter 模型配置。
|
||||
- **正确:** `"provider": "openrouter", "model": "free"` → OpenRouter 收到 `free`。
|
||||
- **也兼容:** `"model": "openrouter/free"` → provider 解析为 `openrouter`,最终模型 ID 解析为 `free`。
|
||||
|
||||
**修复方法:** 在 `~/.picoclaw/config.json`(或你的配置路径)中:
|
||||
|
||||
1. **agents.defaults.model_name** 必须匹配 `model_list` 中的某个 `model_name`(例如 `"openrouter-free"`)。
|
||||
2. 该条目的 **model** 必须是有效的 OpenRouter 模型 ID,例如:
|
||||
- `"openrouter/free"` – 自动免费层
|
||||
2. 该条目推荐显式设置 **provider** 为 `openrouter`,并在 **model** 中填写有效的 OpenRouter 模型 ID,例如:
|
||||
- `"free"` – 自动免费层
|
||||
- `"google/gemini-2.0-flash-exp:free"`
|
||||
- `"meta-llama/llama-3.1-8b-instruct:free"`
|
||||
|
||||
|
|
@ -34,8 +40,9 @@
|
|||
"model_list": [
|
||||
{
|
||||
"model_name": "openrouter-free",
|
||||
"model": "openrouter/free",
|
||||
"api_key": "sk-or-v1-YOUR_OPENROUTER_KEY",
|
||||
"provider": "openrouter",
|
||||
"model": "free",
|
||||
"api_keys": ["sk-or-v1-YOUR_OPENROUTER_KEY"],
|
||||
"api_base": "https://openrouter.ai/api/v1"
|
||||
}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -39,20 +39,23 @@ Set `rpm` on any model in `model_list`:
|
|||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-4o-free
|
||||
model: openai/gpt-4o
|
||||
provider: openai
|
||||
model: gpt-4o
|
||||
api_base: https://api.openai.com/v1
|
||||
rpm: 3 # max 3 requests per minute
|
||||
api_keys:
|
||||
- sk-...
|
||||
|
||||
- model_name: claude-haiku
|
||||
model: anthropic/claude-haiku-4-5
|
||||
provider: anthropic
|
||||
model: claude-haiku-4-5
|
||||
rpm: 60 # 60 rpm (Anthropic free tier)
|
||||
api_keys:
|
||||
- sk-ant-...
|
||||
|
||||
- model_name: local-llm
|
||||
model: openai/llama3
|
||||
provider: ollama
|
||||
model: llama3
|
||||
api_base: http://localhost:11434/v1
|
||||
# no rpm → unrestricted
|
||||
```
|
||||
|
|
@ -68,7 +71,8 @@ When a model has fallbacks configured, each candidate is rate-limited **independ
|
|||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt4-with-fallback
|
||||
model: openai/gpt-4o
|
||||
provider: openai
|
||||
model: gpt-4o
|
||||
rpm: 5
|
||||
fallbacks:
|
||||
- gpt-4o-mini # must also be in model_list; its own rpm applies
|
||||
|
|
|
|||
45
pkg/agent/adapters/channelmanager.go
Normal file
45
pkg/agent/adapters/channelmanager.go
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package adapters
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/agent/interfaces"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
)
|
||||
|
||||
// channelManagerAdapter wraps *channels.Manager to implement interfaces.ChannelManager.
|
||||
type channelManagerAdapter struct {
|
||||
inner *channels.Manager
|
||||
}
|
||||
|
||||
// NewChannelManager creates an adapter for *channels.Manager.
|
||||
func NewChannelManager(inner *channels.Manager) interfaces.ChannelManager {
|
||||
return &channelManagerAdapter{inner: inner}
|
||||
}
|
||||
|
||||
func (a *channelManagerAdapter) GetChannel(name string) (channels.Channel, bool) {
|
||||
return a.inner.GetChannel(name)
|
||||
}
|
||||
|
||||
func (a *channelManagerAdapter) GetEnabledChannels() []string {
|
||||
return a.inner.GetEnabledChannels()
|
||||
}
|
||||
|
||||
func (a *channelManagerAdapter) InvokeTypingStop(channel, chatID string) {
|
||||
a.inner.InvokeTypingStop(channel, chatID)
|
||||
}
|
||||
|
||||
func (a *channelManagerAdapter) SendMessage(ctx context.Context, msg bus.OutboundMessage) error {
|
||||
return a.inner.SendMessage(ctx, msg)
|
||||
}
|
||||
|
||||
func (a *channelManagerAdapter) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||
return a.inner.SendMedia(ctx, msg)
|
||||
}
|
||||
|
||||
func (a *channelManagerAdapter) SendPlaceholder(ctx context.Context, channel, chatID string) bool {
|
||||
return a.inner.SendPlaceholder(ctx, channel, chatID)
|
||||
}
|
||||
36
pkg/agent/adapters/messagebus.go
Normal file
36
pkg/agent/adapters/messagebus.go
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package adapters
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/agent/interfaces"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
)
|
||||
|
||||
// messageBusAdapter wraps *bus.MessageBus to implement interfaces.MessageBus.
|
||||
type messageBusAdapter struct {
|
||||
inner *bus.MessageBus
|
||||
}
|
||||
|
||||
// NewMessageBus creates an adapter for *bus.MessageBus.
|
||||
func NewMessageBus(inner *bus.MessageBus) interfaces.MessageBus {
|
||||
return &messageBusAdapter{inner: inner}
|
||||
}
|
||||
|
||||
func (a *messageBusAdapter) PublishInbound(ctx context.Context, msg bus.InboundMessage) error {
|
||||
return a.inner.PublishInbound(ctx, msg)
|
||||
}
|
||||
|
||||
func (a *messageBusAdapter) PublishOutbound(ctx context.Context, msg bus.OutboundMessage) error {
|
||||
return a.inner.PublishOutbound(ctx, msg)
|
||||
}
|
||||
|
||||
func (a *messageBusAdapter) PublishOutboundMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||
return a.inner.PublishOutboundMedia(ctx, msg)
|
||||
}
|
||||
|
||||
func (a *messageBusAdapter) InboundChan() <-chan bus.InboundMessage {
|
||||
return a.inner.InboundChan()
|
||||
}
|
||||
|
|
@ -15,9 +15,9 @@ import (
|
|||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/agent/interfaces"
|
||||
"github.com/sipeed/picoclaw/pkg/audio/asr"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/commands"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/constants"
|
||||
|
|
@ -32,7 +32,7 @@ import (
|
|||
|
||||
type AgentLoop struct {
|
||||
// Core dependencies
|
||||
bus *bus.MessageBus
|
||||
bus interfaces.MessageBus
|
||||
cfg *config.Config
|
||||
registry *AgentRegistry
|
||||
state *state.Manager
|
||||
|
|
@ -45,7 +45,7 @@ type AgentLoop struct {
|
|||
running atomic.Bool
|
||||
contextManager ContextManager
|
||||
fallback *providers.FallbackChain
|
||||
channelManager *channels.Manager
|
||||
channelManager interfaces.ChannelManager
|
||||
mediaStore media.MediaStore
|
||||
transcriber asr.Transcriber
|
||||
cmdRegistry *commands.Registry
|
||||
|
|
@ -112,6 +112,7 @@ const (
|
|||
pendingTurnPrefix = "pending-"
|
||||
metadataKeyMessageKind = "message_kind"
|
||||
messageKindThought = "thought"
|
||||
messageKindToolFeedback = "tool_feedback"
|
||||
metadataKeyAccountID = "account_id"
|
||||
metadataKeyGuildID = "guild_id"
|
||||
metadataKeyTeamID = "team_id"
|
||||
|
|
@ -495,7 +496,8 @@ func (al *AgentLoop) runAgentLoop(
|
|||
newTurnContext(opts.Dispatch.InboundContext, opts.Dispatch.RouteResult, opts.Dispatch.SessionScope),
|
||||
)
|
||||
ts := newTurnState(agent, opts, turnScope)
|
||||
result, err := al.runTurn(ctx, ts)
|
||||
pipeline := NewPipeline(al)
|
||||
result, err := al.runTurn(ctx, ts, pipeline)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
|
@ -526,10 +528,11 @@ func (al *AgentLoop) runAgentLoop(
|
|||
opts.Dispatch.ChatID(),
|
||||
opts.Dispatch.ReplyToMessageID(),
|
||||
),
|
||||
AgentID: agentID,
|
||||
SessionKey: sessionKey,
|
||||
Scope: scope,
|
||||
Content: result.finalContent,
|
||||
AgentID: agentID,
|
||||
SessionKey: sessionKey,
|
||||
Scope: scope,
|
||||
Content: result.finalContent,
|
||||
ContextUsage: computeContextUsage(agent, opts.Dispatch.SessionKey),
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -4,11 +4,15 @@ package agent
|
|||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/commands"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
|
|
@ -133,6 +137,120 @@ func (al *AgentLoop) buildCommandsRuntime(
|
|||
Config: cfg,
|
||||
ListAgentIDs: registry.ListAgentIDs,
|
||||
ListDefinitions: al.cmdRegistry.Definitions,
|
||||
ListMCPServers: func(ctx context.Context) []commands.MCPServerInfo {
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(cfg.Tools.MCP.Servers) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := al.ensureMCPInitialized(ctx); err != nil {
|
||||
logger.WarnCF("agent", "Failed to refresh MCP status for command",
|
||||
map[string]any{
|
||||
"error": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
connected := make(map[string]int)
|
||||
if manager := al.mcp.getManager(); manager != nil {
|
||||
for serverName, conn := range manager.GetServers() {
|
||||
connected[serverName] = len(conn.Tools)
|
||||
}
|
||||
}
|
||||
|
||||
servers := make([]commands.MCPServerInfo, 0, len(cfg.Tools.MCP.Servers))
|
||||
for serverName, serverCfg := range cfg.Tools.MCP.Servers {
|
||||
toolCount, isConnected := connected[serverName]
|
||||
servers = append(servers, commands.MCPServerInfo{
|
||||
Name: serverName,
|
||||
Enabled: serverCfg.Enabled,
|
||||
Deferred: serverIsDeferred(cfg.Tools.MCP.Discovery.Enabled, serverCfg),
|
||||
Connected: isConnected,
|
||||
ToolCount: toolCount,
|
||||
})
|
||||
}
|
||||
|
||||
sort.Slice(servers, func(i, j int) bool {
|
||||
return strings.ToLower(servers[i].Name) < strings.ToLower(servers[j].Name)
|
||||
})
|
||||
|
||||
return servers
|
||||
},
|
||||
ListMCPTools: func(ctx context.Context, serverName string) ([]commands.MCPToolInfo, error) {
|
||||
if cfg == nil {
|
||||
return nil, fmt.Errorf("command unavailable: config not loaded")
|
||||
}
|
||||
|
||||
serverName = strings.TrimSpace(serverName)
|
||||
if serverName == "" {
|
||||
return nil, fmt.Errorf("server name is required")
|
||||
}
|
||||
|
||||
resolvedName := ""
|
||||
var serverCfg config.MCPServerConfig
|
||||
for name, candidate := range cfg.Tools.MCP.Servers {
|
||||
if strings.EqualFold(name, serverName) {
|
||||
resolvedName = name
|
||||
serverCfg = candidate
|
||||
break
|
||||
}
|
||||
}
|
||||
if resolvedName == "" {
|
||||
return nil, fmt.Errorf("MCP server '%s' is not configured", serverName)
|
||||
}
|
||||
if !serverCfg.Enabled {
|
||||
return nil, fmt.Errorf("MCP server '%s' is configured but disabled", resolvedName)
|
||||
}
|
||||
if !cfg.Tools.IsToolEnabled("mcp") {
|
||||
return nil, fmt.Errorf("MCP integration is disabled")
|
||||
}
|
||||
|
||||
if err := al.ensureMCPInitialized(ctx); err != nil {
|
||||
logger.WarnCF("agent", "Failed to initialize MCP runtime for command",
|
||||
map[string]any{
|
||||
"server": resolvedName,
|
||||
"error": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
manager := al.mcp.getManager()
|
||||
if manager == nil {
|
||||
return nil, fmt.Errorf("MCP server '%s' is configured but not connected", resolvedName)
|
||||
}
|
||||
|
||||
conn, ok := manager.GetServer(resolvedName)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("MCP server '%s' is configured but not connected", resolvedName)
|
||||
}
|
||||
|
||||
toolInfos := make([]commands.MCPToolInfo, 0, len(conn.Tools))
|
||||
for _, tool := range conn.Tools {
|
||||
if tool == nil {
|
||||
continue
|
||||
}
|
||||
name := strings.TrimSpace(tool.Name)
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
description := strings.TrimSpace(tool.Description)
|
||||
if description == "" {
|
||||
description = fmt.Sprintf("MCP tool from %s server", resolvedName)
|
||||
}
|
||||
|
||||
toolInfos = append(toolInfos, commands.MCPToolInfo{
|
||||
Name: name,
|
||||
Description: description,
|
||||
Parameters: summarizeMCPToolParameters(tool.InputSchema),
|
||||
})
|
||||
}
|
||||
sort.Slice(toolInfos, func(i, j int) bool {
|
||||
return toolInfos[i].Name < toolInfos[j].Name
|
||||
})
|
||||
return toolInfos, nil
|
||||
},
|
||||
GetEnabledChannels: func() []string {
|
||||
if al.channelManager == nil {
|
||||
return nil
|
||||
|
|
@ -214,10 +332,118 @@ func (al *AgentLoop) buildCommandsRuntime(
|
|||
rt.AskSideQuestion = func(ctx context.Context, question string) (string, error) {
|
||||
return al.askSideQuestion(ctx, agent, opts, question)
|
||||
}
|
||||
|
||||
rt.GetContextStats = func() *commands.ContextStats {
|
||||
if opts == nil || agent.Sessions == nil {
|
||||
return nil
|
||||
}
|
||||
usage := computeContextUsage(agent, opts.SessionKey)
|
||||
if usage == nil {
|
||||
return nil
|
||||
}
|
||||
history := agent.Sessions.GetHistory(opts.SessionKey)
|
||||
return &commands.ContextStats{
|
||||
UsedTokens: usage.UsedTokens,
|
||||
TotalTokens: usage.TotalTokens,
|
||||
CompressAtTokens: usage.CompressAtTokens,
|
||||
UsedPercent: usage.UsedPercent,
|
||||
MessageCount: len(history),
|
||||
}
|
||||
}
|
||||
}
|
||||
return rt
|
||||
}
|
||||
|
||||
func summarizeMCPToolParameters(schema any) []commands.MCPToolParameterInfo {
|
||||
schemaMap := normalizeMCPSchema(schema)
|
||||
properties, ok := schemaMap["properties"].(map[string]any)
|
||||
if !ok || len(properties) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
required := make(map[string]struct{})
|
||||
switch raw := schemaMap["required"].(type) {
|
||||
case []string:
|
||||
for _, name := range raw {
|
||||
required[name] = struct{}{}
|
||||
}
|
||||
case []any:
|
||||
for _, value := range raw {
|
||||
name, ok := value.(string)
|
||||
if ok {
|
||||
required[name] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(properties))
|
||||
for name := range properties {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
params := make([]commands.MCPToolParameterInfo, 0, len(names))
|
||||
for _, name := range names {
|
||||
param := commands.MCPToolParameterInfo{Name: name}
|
||||
if propMap, ok := properties[name].(map[string]any); ok {
|
||||
if typeName, ok := propMap["type"].(string); ok {
|
||||
param.Type = strings.TrimSpace(typeName)
|
||||
}
|
||||
if desc, ok := propMap["description"].(string); ok {
|
||||
param.Description = strings.TrimSpace(desc)
|
||||
}
|
||||
}
|
||||
_, param.Required = required[name]
|
||||
params = append(params, param)
|
||||
}
|
||||
return params
|
||||
}
|
||||
|
||||
func normalizeMCPSchema(schema any) map[string]any {
|
||||
if schema == nil {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{},
|
||||
"required": []string{},
|
||||
}
|
||||
}
|
||||
|
||||
if schemaMap, ok := schema.(map[string]any); ok {
|
||||
return schemaMap
|
||||
}
|
||||
|
||||
var jsonData []byte
|
||||
switch raw := schema.(type) {
|
||||
case json.RawMessage:
|
||||
jsonData = raw
|
||||
case []byte:
|
||||
jsonData = raw
|
||||
}
|
||||
|
||||
if jsonData == nil {
|
||||
var err error
|
||||
jsonData, err = json.Marshal(schema)
|
||||
if err != nil {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{},
|
||||
"required": []string{},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var result map[string]any
|
||||
if err := json.Unmarshal(jsonData, &result); err != nil {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{},
|
||||
"required": []string{},
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func (al *AgentLoop) setPendingSkills(sessionKey string, skillNames []string) {
|
||||
sessionKey = strings.TrimSpace(sessionKey)
|
||||
if sessionKey == "" || len(skillNames) == 0 {
|
||||
|
|
@ -48,24 +48,6 @@ func (al *AgentLoop) emitEvent(kind EventKind, meta EventMeta, payload any) {
|
|||
al.eventBus.Emit(evt)
|
||||
}
|
||||
|
||||
func (al *AgentLoop) hookAbortError(ts *turnState, stage string, decision HookDecision) error {
|
||||
reason := decision.Reason
|
||||
if reason == "" {
|
||||
reason = "hook requested turn abort"
|
||||
}
|
||||
|
||||
err := fmt.Errorf("hook aborted turn during %s: %s", stage, reason)
|
||||
al.emitEvent(
|
||||
EventKindError,
|
||||
ts.eventMeta("hooks", "turn.error"),
|
||||
ErrorPayload{
|
||||
Stage: "hook." + stage,
|
||||
Message: err.Error(),
|
||||
},
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (al *AgentLoop) logEvent(evt Event) {
|
||||
fields := map[string]any{
|
||||
"event_kind": evt.Kind.String(),
|
||||
|
|
@ -7,6 +7,7 @@ import (
|
|||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/agent/interfaces"
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
|
|
@ -79,7 +80,7 @@ func NewAgentLoop(
|
|||
func registerSharedTools(
|
||||
al *AgentLoop,
|
||||
cfg *config.Config,
|
||||
msgBus *bus.MessageBus,
|
||||
msgBus interfaces.MessageBus,
|
||||
registry *AgentRegistry,
|
||||
provider providers.LLMProvider,
|
||||
) {
|
||||
|
|
@ -99,33 +100,7 @@ func registerSharedTools(
|
|||
}
|
||||
|
||||
if cfg.Tools.IsToolEnabled("web") {
|
||||
searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
|
||||
BraveAPIKeys: cfg.Tools.Web.Brave.APIKeys.Values(),
|
||||
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
||||
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
||||
TavilyAPIKeys: cfg.Tools.Web.Tavily.APIKeys.Values(),
|
||||
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
||||
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
||||
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
||||
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
||||
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
||||
PerplexityAPIKeys: cfg.Tools.Web.Perplexity.APIKeys.Values(),
|
||||
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
||||
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
||||
SearXNGBaseURL: cfg.Tools.Web.SearXNG.BaseURL,
|
||||
SearXNGMaxResults: cfg.Tools.Web.SearXNG.MaxResults,
|
||||
SearXNGEnabled: cfg.Tools.Web.SearXNG.Enabled,
|
||||
GLMSearchAPIKey: cfg.Tools.Web.GLMSearch.APIKey.String(),
|
||||
GLMSearchBaseURL: cfg.Tools.Web.GLMSearch.BaseURL,
|
||||
GLMSearchEngine: cfg.Tools.Web.GLMSearch.SearchEngine,
|
||||
GLMSearchMaxResults: cfg.Tools.Web.GLMSearch.MaxResults,
|
||||
GLMSearchEnabled: cfg.Tools.Web.GLMSearch.Enabled,
|
||||
BaiduSearchAPIKey: cfg.Tools.Web.BaiduSearch.APIKey.String(),
|
||||
BaiduSearchBaseURL: cfg.Tools.Web.BaiduSearch.BaseURL,
|
||||
BaiduSearchMaxResults: cfg.Tools.Web.BaiduSearch.MaxResults,
|
||||
BaiduSearchEnabled: cfg.Tools.Web.BaiduSearch.Enabled,
|
||||
Proxy: cfg.Tools.Web.Proxy,
|
||||
})
|
||||
searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptionsFromConfig(cfg))
|
||||
if err != nil {
|
||||
logger.ErrorCF("agent", "Failed to create web search tool", map[string]any{"error": err.Error()})
|
||||
} else if searchTool != nil {
|
||||
|
|
@ -67,6 +67,12 @@ func (r *mcpRuntime) hasManager() bool {
|
|||
return r.manager != nil
|
||||
}
|
||||
|
||||
func (r *mcpRuntime) getManager() *mcp.Manager {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return r.manager
|
||||
}
|
||||
|
||||
// ensureMCPInitialized loads MCP servers/tools once so both Run() and direct
|
||||
// agent mode share the same initialization path.
|
||||
func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
|
||||
|
|
@ -100,6 +106,7 @@ func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
|
|||
}
|
||||
|
||||
if err := mcpManager.LoadFromMCPConfig(ctx, al.cfg.Tools.MCP, workspacePath); err != nil {
|
||||
al.mcp.setInitErr(fmt.Errorf("failed to load MCP servers: %w", err))
|
||||
logger.WarnCF("agent", "Failed to load MCP servers, MCP tools will not be available",
|
||||
map[string]any{
|
||||
"error": err.Error(),
|
||||
|
|
@ -9,6 +9,7 @@ package agent
|
|||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
|
|
@ -133,3 +134,48 @@ func TestServerIsDeferred(t *testing.T) {
|
|||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureMCPInitialized_LoadFailureSetsInitErr(t *testing.T) {
|
||||
al, cfg, _, _, cleanup := newTestAgentLoop(t)
|
||||
defer cleanup()
|
||||
defer al.Close()
|
||||
|
||||
cfg.Tools = config.ToolsConfig{
|
||||
MCP: config.MCPConfig{
|
||||
ToolConfig: config.ToolConfig{Enabled: true},
|
||||
Servers: map[string]config.MCPServerConfig{
|
||||
"broken": {
|
||||
Enabled: true,
|
||||
Command: "picoclaw-command-that-does-not-exist-for-mcp-tests",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := al.ensureMCPInitialized(context.Background())
|
||||
if err == nil {
|
||||
t.Fatal("ensureMCPInitialized() error = nil, want load failure")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed to load MCP servers") {
|
||||
t.Fatalf("ensureMCPInitialized() error = %q, want wrapped load failure", err.Error())
|
||||
}
|
||||
|
||||
initErr := al.mcp.getInitErr()
|
||||
if initErr == nil {
|
||||
t.Fatal("getInitErr() = nil, want cached load failure")
|
||||
}
|
||||
if !strings.Contains(initErr.Error(), "failed to load MCP servers") {
|
||||
t.Fatalf("getInitErr() = %q, want wrapped load failure", initErr.Error())
|
||||
}
|
||||
if al.mcp.getManager() != nil {
|
||||
t.Fatal("expected MCP manager to remain nil after load failure")
|
||||
}
|
||||
|
||||
err = al.ensureMCPInitialized(context.Background())
|
||||
if err == nil {
|
||||
t.Fatal("second ensureMCPInitialized() error = nil, want cached load failure")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed to load MCP servers") {
|
||||
t.Fatalf("second ensureMCPInitialized() error = %q, want wrapped load failure", err.Error())
|
||||
}
|
||||
}
|
||||
|
|
@ -105,6 +105,25 @@ func buildArtifactTags(store media.MediaStore, refs []string) []string {
|
|||
return tags
|
||||
}
|
||||
|
||||
func buildProviderAttachments(store media.MediaStore, refs []string) []providers.Attachment {
|
||||
if store == nil || len(refs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
attachments := make([]providers.Attachment, 0, len(refs))
|
||||
for _, ref := range refs {
|
||||
attachment := providers.Attachment{Ref: ref}
|
||||
if _, meta, err := store.ResolveWithMeta(ref); err == nil {
|
||||
attachment.Filename = meta.Filename
|
||||
attachment.ContentType = meta.ContentType
|
||||
attachment.Type = inferMediaType(meta.Filename, meta.ContentType)
|
||||
}
|
||||
attachments = append(attachments, attachment)
|
||||
}
|
||||
|
||||
return attachments
|
||||
}
|
||||
|
||||
// detectMIME determines the MIME type from metadata or magic-bytes detection.
|
||||
// Returns empty string if detection fails.
|
||||
func detectMIME(localPath string, meta media.MediaMeta) string {
|
||||
|
|
@ -60,10 +60,14 @@ func (al *AgentLoop) PublishResponseIfNeeded(ctx context.Context, channel, chatI
|
|||
return
|
||||
}
|
||||
|
||||
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
msg := bus.OutboundMessage{
|
||||
Context: bus.NewOutboundContext(channel, chatID, ""),
|
||||
Content: response,
|
||||
})
|
||||
}
|
||||
if sessionKey != "" {
|
||||
msg.ContextUsage = computeContextUsage(al.agentForSession(sessionKey), sessionKey)
|
||||
}
|
||||
al.bus.PublishOutbound(ctx, msg)
|
||||
logger.InfoCF("agent", "Published outbound response",
|
||||
map[string]any{
|
||||
"channel": channel,
|
||||
|
|
@ -24,6 +24,7 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/routing"
|
||||
"github.com/sipeed/picoclaw/pkg/session"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
type fakeChannel struct{ id string }
|
||||
|
|
@ -128,7 +129,7 @@ func useTestSideQuestionProvider(al *AgentLoop, provider providers.LLMProvider)
|
|||
al.providerFactory = func(mc *config.ModelConfig) (providers.LLMProvider, string, error) {
|
||||
model := provider.GetDefaultModel()
|
||||
if mc != nil {
|
||||
if _, modelID := providers.ExtractProtocol(mc.Model); modelID != "" {
|
||||
if _, modelID := providers.ExtractProtocol(mc); modelID != "" {
|
||||
model = modelID
|
||||
}
|
||||
}
|
||||
|
|
@ -160,6 +161,58 @@ func newTestAgentLoop(
|
|||
return al, cfg, msgBus, provider, func() { os.RemoveAll(tmpDir) }
|
||||
}
|
||||
|
||||
func TestNewAgentLoop_RegistersWebSearchTool(t *testing.T) {
|
||||
cfg := config.DefaultConfig()
|
||||
cfg.Agents.Defaults.Workspace = t.TempDir()
|
||||
|
||||
al := NewAgentLoop(cfg, bus.NewMessageBus(), &mockProvider{})
|
||||
|
||||
agent := al.registry.GetDefaultAgent()
|
||||
if agent == nil {
|
||||
t.Fatal("expected default agent")
|
||||
}
|
||||
if _, ok := agent.Tools.Get("web_search"); !ok {
|
||||
t.Fatal("expected web_search tool to be registered")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewAgentLoop_RegistersWebSearchTool_WhenExplicitProviderUnavailable(t *testing.T) {
|
||||
cfg := config.DefaultConfig()
|
||||
cfg.Agents.Defaults.Workspace = t.TempDir()
|
||||
cfg.Tools.Web.Provider = "brave"
|
||||
cfg.Tools.Web.Brave.Enabled = true
|
||||
cfg.Tools.Web.Sogou.Enabled = true
|
||||
|
||||
al := NewAgentLoop(cfg, bus.NewMessageBus(), &mockProvider{})
|
||||
|
||||
agent := al.registry.GetDefaultAgent()
|
||||
if agent == nil {
|
||||
t.Fatal("expected default agent")
|
||||
}
|
||||
if _, ok := agent.Tools.Get("web_search"); !ok {
|
||||
t.Fatal("expected web_search tool to fall back to auto provider selection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewAgentLoop_DoesNotRegisterWebSearchTool_WhenNoReadyProviders(t *testing.T) {
|
||||
cfg := config.DefaultConfig()
|
||||
cfg.Agents.Defaults.Workspace = t.TempDir()
|
||||
cfg.Tools.Web.Provider = "brave"
|
||||
cfg.Tools.Web.Brave.Enabled = true
|
||||
cfg.Tools.Web.Sogou.Enabled = false
|
||||
cfg.Tools.Web.DuckDuckGo.Enabled = false
|
||||
|
||||
al := NewAgentLoop(cfg, bus.NewMessageBus(), &mockProvider{})
|
||||
|
||||
agent := al.registry.GetDefaultAgent()
|
||||
if agent == nil {
|
||||
t.Fatal("expected default agent")
|
||||
}
|
||||
if _, ok := agent.Tools.Get("web_search"); ok {
|
||||
t.Fatal("expected web_search tool to be absent when no providers are ready")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_IncludesCurrentSenderInDynamicContext(t *testing.T) {
|
||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||
if err != nil {
|
||||
|
|
@ -1051,6 +1104,9 @@ func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.
|
|||
if last.Role != "assistant" || last.Content != "Requested output delivered via tool attachment." {
|
||||
t.Fatalf("expected handled assistant summary in history, got %+v", last)
|
||||
}
|
||||
if len(last.Attachments) != 1 {
|
||||
t.Fatalf("expected handled assistant summary attachments in history, got %+v", last.Attachments)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_HandledToolProcessesQueuedSteeringBeforeReturning(t *testing.T) {
|
||||
|
|
@ -1758,6 +1814,157 @@ func (m *toolFeedbackProvider) GetDefaultModel() string {
|
|||
return "heartbeat-tool-feedback-model"
|
||||
}
|
||||
|
||||
type toolFeedbackReasoningProvider struct {
|
||||
filePath string
|
||||
calls int
|
||||
}
|
||||
|
||||
func (m *toolFeedbackReasoningProvider) Chat(
|
||||
ctx context.Context,
|
||||
messages []providers.Message,
|
||||
tools []providers.ToolDefinition,
|
||||
model string,
|
||||
opts map[string]any,
|
||||
) (*providers.LLMResponse, error) {
|
||||
m.calls++
|
||||
if m.calls == 1 {
|
||||
return &providers.LLMResponse{
|
||||
ReasoningContent: "Read README.md first to confirm the context that needs to be changed.",
|
||||
ToolCalls: []providers.ToolCall{{
|
||||
ID: "call_reasoning_read_file",
|
||||
Type: "function",
|
||||
Name: "read_file",
|
||||
Arguments: map[string]any{"path": m.filePath},
|
||||
}},
|
||||
}, nil
|
||||
}
|
||||
|
||||
return &providers.LLMResponse{
|
||||
Content: "DONE",
|
||||
ToolCalls: []providers.ToolCall{},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *toolFeedbackReasoningProvider) GetDefaultModel() string {
|
||||
return "tool-feedback-reasoning-model"
|
||||
}
|
||||
|
||||
func TestToolFeedbackExplanationFromResponse_UsesCurrentContentFirst(t *testing.T) {
|
||||
response := &providers.LLMResponse{
|
||||
Content: "Read README.md first",
|
||||
ReasoningContent: "current reasoning fallback",
|
||||
}
|
||||
messages := []providers.Message{
|
||||
{Role: "user", Content: "check file"},
|
||||
{Role: "assistant", Content: "Previous turn explanation"},
|
||||
{Role: "tool", Content: "tool output", ToolCallID: "call_1"},
|
||||
}
|
||||
|
||||
got := toolFeedbackExplanationFromResponse(response, messages, 300)
|
||||
if got != "Read README.md first" {
|
||||
t.Fatalf("toolFeedbackExplanationFromResponse() = %q, want current content", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolFeedbackExplanationFromResponse_UsesExplicitToolCallExtraContent(t *testing.T) {
|
||||
response := &providers.LLMResponse{
|
||||
ToolCalls: []providers.ToolCall{{
|
||||
ID: "call_1",
|
||||
Name: "read_file",
|
||||
ExtraContent: &providers.ExtraContent{
|
||||
ToolFeedbackExplanation: "Read README.md first to confirm the current project structure.",
|
||||
},
|
||||
}},
|
||||
}
|
||||
messages := []providers.Message{
|
||||
{Role: "user", Content: "check file"},
|
||||
{Role: "assistant", Content: ""},
|
||||
{Role: "tool", Content: "tool output", ToolCallID: "call_1"},
|
||||
}
|
||||
|
||||
got := toolFeedbackExplanationFromResponse(response, messages, 300)
|
||||
if got != "Read README.md first to confirm the current project structure." {
|
||||
t.Fatalf("toolFeedbackExplanationFromResponse() = %q, want explicit tool feedback explanation", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolFeedbackExplanationForToolCall_PrefersToolSpecificExtraContent(t *testing.T) {
|
||||
response := &providers.LLMResponse{
|
||||
Content: "Shared explanation",
|
||||
ToolCalls: []providers.ToolCall{
|
||||
{
|
||||
ID: "call_1",
|
||||
Name: "read_file",
|
||||
ExtraContent: &providers.ExtraContent{
|
||||
ToolFeedbackExplanation: "Read README.md first.",
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "call_2",
|
||||
Name: "edit_file",
|
||||
ExtraContent: &providers.ExtraContent{
|
||||
ToolFeedbackExplanation: "Update config example after reading it.",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
got1 := toolFeedbackExplanationForToolCall(response, response.ToolCalls[0], nil, 300)
|
||||
got2 := toolFeedbackExplanationForToolCall(response, response.ToolCalls[1], nil, 300)
|
||||
if got1 != "Read README.md first." {
|
||||
t.Fatalf("toolFeedbackExplanationForToolCall() first = %q, want tool-specific explanation", got1)
|
||||
}
|
||||
if got2 != "Update config example after reading it." {
|
||||
t.Fatalf("toolFeedbackExplanationForToolCall() second = %q, want tool-specific explanation", got2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolFeedbackExplanationForToolCall_DoesNotReuseAnotherToolCallExplanation(t *testing.T) {
|
||||
response := &providers.LLMResponse{
|
||||
ToolCalls: []providers.ToolCall{
|
||||
{
|
||||
ID: "call_1",
|
||||
Name: "read_file",
|
||||
},
|
||||
{
|
||||
ID: "call_2",
|
||||
Name: "edit_file",
|
||||
ExtraContent: &providers.ExtraContent{
|
||||
ToolFeedbackExplanation: "Update config example after reading it.",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
messages := []providers.Message{
|
||||
{Role: "user", Content: "inspect the config and update the example"},
|
||||
}
|
||||
|
||||
got := toolFeedbackExplanationForToolCall(response, response.ToolCalls[0], messages, 300)
|
||||
want := utils.ToolFeedbackContinuationHint + ": inspect the config and update the example"
|
||||
if got != want {
|
||||
t.Fatalf("toolFeedbackExplanationForToolCall() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolFeedbackExplanationFromResponse_DoesNotUseReasoningContent(t *testing.T) {
|
||||
response := &providers.LLMResponse{
|
||||
Content: "",
|
||||
ReasoningContent: "hidden reasoning should not be shown",
|
||||
}
|
||||
messages := []providers.Message{
|
||||
{Role: "user", Content: "check file"},
|
||||
{Role: "assistant", Content: "Previous turn explanation"},
|
||||
{Role: "user", Content: "Inspect README.md and update the config example."},
|
||||
{Role: "tool", Content: "tool output", ToolCallID: "call_1"},
|
||||
}
|
||||
|
||||
got := toolFeedbackExplanationFromResponse(response, messages, 300)
|
||||
want := utils.ToolFeedbackContinuationHint + ": Inspect README.md and update the config example."
|
||||
if got != want {
|
||||
t.Fatalf("toolFeedbackExplanationFromResponse() = %q, want latest user content fallback", got)
|
||||
}
|
||||
}
|
||||
|
||||
type picoInterleavedContentProvider struct {
|
||||
calls int
|
||||
}
|
||||
|
|
@ -2269,6 +2476,75 @@ func TestProcessMessage_CommandOutcomes(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_MCPCommandsHandledWithoutLLMCall(t *testing.T) {
|
||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create temp dir: %v", err)
|
||||
}
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
deferred := true
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "test-model",
|
||||
MaxTokens: 4096,
|
||||
MaxToolIterations: 10,
|
||||
},
|
||||
},
|
||||
Session: config.SessionConfig{
|
||||
Dimensions: []string{"chat"},
|
||||
},
|
||||
Tools: config.ToolsConfig{
|
||||
MCP: config.MCPConfig{
|
||||
ToolConfig: config.ToolConfig{Enabled: true},
|
||||
Discovery: config.ToolDiscoveryConfig{Enabled: true},
|
||||
Servers: map[string]config.MCPServerConfig{
|
||||
"github": {
|
||||
Enabled: true,
|
||||
Deferred: &deferred,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
msgBus := bus.NewMessageBus()
|
||||
provider := &countingMockProvider{response: "LLM reply"}
|
||||
al := NewAgentLoop(cfg, msgBus, provider)
|
||||
helper := testHelper{al: al}
|
||||
|
||||
baseContext := bus.InboundContext{
|
||||
Channel: "whatsapp",
|
||||
ChatID: "chat1",
|
||||
ChatType: "direct",
|
||||
SenderID: "user1",
|
||||
}
|
||||
|
||||
listResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{
|
||||
Context: baseContext,
|
||||
Content: "/list mcp",
|
||||
})
|
||||
if !strings.Contains(listResp, "- `github`") || !strings.Contains(listResp, "Deferred: yes") {
|
||||
t.Fatalf("unexpected /list mcp reply: %q", listResp)
|
||||
}
|
||||
if provider.calls != 0 {
|
||||
t.Fatalf("LLM should not be called for /list mcp, calls=%d", provider.calls)
|
||||
}
|
||||
|
||||
showResp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{
|
||||
Context: baseContext,
|
||||
Content: "/show mcp github",
|
||||
})
|
||||
if showResp != "MCP server 'github' is configured but not connected" {
|
||||
t.Fatalf("unexpected /show mcp reply: %q", showResp)
|
||||
}
|
||||
if provider.calls != 0 {
|
||||
t.Fatalf("LLM should not be called for /show mcp, calls=%d", provider.calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_SwitchModelShowModelConsistency(t *testing.T) {
|
||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||
if err != nil {
|
||||
|
|
@ -3656,7 +3932,16 @@ func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) {
|
|||
t.Fatalf("unexpected tool feedback context: %+v", outbound.Context)
|
||||
}
|
||||
if !strings.Contains(outbound.Content, "`read_file`") {
|
||||
t.Fatalf("tool feedback content = %q, want read_file preview", outbound.Content)
|
||||
t.Fatalf("tool feedback content = %q, want read_file summary", outbound.Content)
|
||||
}
|
||||
if !strings.Contains(outbound.Content, utils.ToolFeedbackContinuationHint) {
|
||||
t.Fatalf("tool feedback content = %q, want continuation hint fallback", outbound.Content)
|
||||
}
|
||||
if !strings.Contains(outbound.Content, "check tool feedback") {
|
||||
t.Fatalf("tool feedback content = %q, want current user intent fallback", outbound.Content)
|
||||
}
|
||||
if strings.Contains(outbound.Content, "Previous turn explanation") {
|
||||
t.Fatalf("tool feedback content = %q, want no previous assistant fallback", outbound.Content)
|
||||
}
|
||||
if outbound.AgentID != "main" {
|
||||
t.Fatalf("tool feedback agent_id = %q, want main", outbound.AgentID)
|
||||
|
|
@ -3672,6 +3957,130 @@ func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_DoesNotLeakReasoningContentInToolFeedback(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
heartbeatFile := filepath.Join(tmpDir, "tool-feedback-reasoning.txt")
|
||||
if err := os.WriteFile(heartbeatFile, []byte("tool feedback task"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "test-model",
|
||||
MaxTokens: 4096,
|
||||
MaxToolIterations: 10,
|
||||
ToolFeedback: config.ToolFeedbackConfig{
|
||||
Enabled: true,
|
||||
MaxArgsLength: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
Tools: config.ToolsConfig{
|
||||
ReadFile: config.ReadFileToolConfig{
|
||||
Enabled: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
msgBus := bus.NewMessageBus()
|
||||
provider := &toolFeedbackReasoningProvider{filePath: heartbeatFile}
|
||||
al := NewAgentLoop(cfg, msgBus, provider)
|
||||
|
||||
response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{
|
||||
Channel: "telegram",
|
||||
SenderID: "user-1",
|
||||
ChatID: "chat-1",
|
||||
Content: "check reasoning fallback",
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("processMessage() error = %v", err)
|
||||
}
|
||||
if response != "DONE" {
|
||||
t.Fatalf("processMessage() response = %q, want %q", response, "DONE")
|
||||
}
|
||||
|
||||
select {
|
||||
case outbound := <-msgBus.OutboundChan():
|
||||
if !strings.Contains(outbound.Content, "`read_file`") {
|
||||
t.Fatalf("tool feedback content = %q, want read_file summary", outbound.Content)
|
||||
}
|
||||
if !strings.Contains(outbound.Content, utils.ToolFeedbackContinuationHint) {
|
||||
t.Fatalf("tool feedback content = %q, want continuation hint fallback", outbound.Content)
|
||||
}
|
||||
if !strings.Contains(outbound.Content, "check reasoning fallback") {
|
||||
t.Fatalf("tool feedback content = %q, want current user intent fallback", outbound.Content)
|
||||
}
|
||||
if strings.Contains(outbound.Content, "Read README.md first") {
|
||||
t.Fatalf("tool feedback content = %q, should not leak hidden reasoning", outbound.Content)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("expected outbound tool feedback without leaking reasoning")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_DoesNotPublishToolFeedbackForDiscordWhenDisabled(t *testing.T) {
|
||||
assertToolFeedbackNotPublishedWhenDisabled(t, "discord")
|
||||
}
|
||||
|
||||
func assertToolFeedbackNotPublishedWhenDisabled(t *testing.T, channel string) {
|
||||
t.Helper()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
heartbeatFile := filepath.Join(tmpDir, "tool-feedback-"+channel+".txt")
|
||||
if err := os.WriteFile(heartbeatFile, []byte("tool feedback task"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "test-model",
|
||||
MaxTokens: 4096,
|
||||
MaxToolIterations: 10,
|
||||
},
|
||||
},
|
||||
Tools: config.ToolsConfig{
|
||||
ReadFile: config.ReadFileToolConfig{
|
||||
Enabled: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
msgBus := bus.NewMessageBus()
|
||||
provider := &toolFeedbackProvider{filePath: heartbeatFile}
|
||||
al := NewAgentLoop(cfg, msgBus, provider)
|
||||
|
||||
response, err := al.processMessage(context.Background(), testInboundMessage(bus.InboundMessage{
|
||||
Channel: channel,
|
||||
SenderID: "user-1",
|
||||
ChatID: "chat-1",
|
||||
Content: "check tool feedback",
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("processMessage() error = %v", err)
|
||||
}
|
||||
if response != "HEARTBEAT_OK" {
|
||||
t.Fatalf("processMessage() response = %q, want %q", response, "HEARTBEAT_OK")
|
||||
}
|
||||
|
||||
select {
|
||||
case outbound := <-msgBus.OutboundChan():
|
||||
t.Fatalf("expected no outbound tool feedback for %s when disabled, got %+v", channel, outbound)
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessMessage_DoesNotPublishToolFeedbackForTelegramWhenDisabled(t *testing.T) {
|
||||
assertToolFeedbackNotPublishedWhenDisabled(t, "telegram")
|
||||
}
|
||||
|
||||
func TestProcessMessage_DoesNotPublishToolFeedbackForFeishuWhenDisabled(t *testing.T) {
|
||||
assertToolFeedbackNotPublishedWhenDisabled(t, "feishu")
|
||||
}
|
||||
|
||||
func TestProcessMessage_MessageToolPublishesOutboundWithTurnMetadata(t *testing.T) {
|
||||
cfg := config.DefaultConfig()
|
||||
cfg.Agents.Defaults.Workspace = t.TempDir()
|
||||
|
|
@ -3846,6 +4255,85 @@ func TestRunAgentLoop_PicoSkipsInterimPublishWhenNotAllowed(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestRun_PicoToolFeedbackSuppressesDuplicateInterimAssistantContent(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "test-model",
|
||||
MaxTokens: 4096,
|
||||
MaxToolIterations: 10,
|
||||
ToolFeedback: config.ToolFeedbackConfig{
|
||||
Enabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
msgBus := bus.NewMessageBus()
|
||||
provider := &picoInterleavedContentProvider{}
|
||||
al := NewAgentLoop(cfg, msgBus, provider)
|
||||
|
||||
agent := al.GetRegistry().GetDefaultAgent()
|
||||
if agent == nil {
|
||||
t.Fatal("expected default agent")
|
||||
}
|
||||
agent.Tools.Register(&toolLimitTestTool{})
|
||||
|
||||
runCtx, runCancel := context.WithCancel(context.Background())
|
||||
defer runCancel()
|
||||
|
||||
runDone := make(chan error, 1)
|
||||
go func() {
|
||||
runDone <- al.Run(runCtx)
|
||||
}()
|
||||
|
||||
if err := msgBus.PublishInbound(context.Background(), bus.InboundMessage{
|
||||
Channel: "pico",
|
||||
SenderID: "user-1",
|
||||
ChatID: "session-1",
|
||||
Content: "run with tools",
|
||||
}); err != nil {
|
||||
t.Fatalf("PublishInbound() error = %v", err)
|
||||
}
|
||||
|
||||
outputs := make([]string, 0, 2)
|
||||
deadline := time.After(2 * time.Second)
|
||||
for len(outputs) < 2 {
|
||||
select {
|
||||
case outbound := <-msgBus.OutboundChan():
|
||||
outputs = append(outputs, outbound.Content)
|
||||
case <-deadline:
|
||||
t.Fatalf("timed out waiting for pico outputs, got %v", outputs)
|
||||
}
|
||||
}
|
||||
|
||||
if outputs[0] != "🔧 `tool_limit_test_tool`\nintermediate model text" {
|
||||
t.Fatalf("first outbound content = %q, want tool feedback summary", outputs[0])
|
||||
}
|
||||
if outputs[1] != "final model text" {
|
||||
t.Fatalf("second outbound content = %q, want %q", outputs[1], "final model text")
|
||||
}
|
||||
|
||||
runCancel()
|
||||
select {
|
||||
case err := <-runDone:
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for Run() to exit")
|
||||
}
|
||||
|
||||
select {
|
||||
case outbound := <-msgBus.OutboundChan():
|
||||
t.Fatalf("unexpected extra pico output after tool feedback + final reply: %+v", outbound)
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveMediaRefs_ResolvesToBase64(t *testing.T) {
|
||||
store := media.NewFileMediaStore()
|
||||
dir := t.TempDir()
|
||||
|
|
@ -11,6 +11,7 @@ import (
|
|||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/commands"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
"github.com/sipeed/picoclaw/pkg/session"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
|
|
@ -84,6 +85,98 @@ func outboundMessageForTurn(ts *turnState, content string) bus.OutboundMessage {
|
|||
}
|
||||
}
|
||||
|
||||
func outboundMessageForTurnWithKind(ts *turnState, content, kind string) bus.OutboundMessage {
|
||||
msg := outboundMessageForTurn(ts, content)
|
||||
if strings.TrimSpace(kind) == "" {
|
||||
return msg
|
||||
}
|
||||
if msg.Context.Raw == nil {
|
||||
msg.Context.Raw = make(map[string]string, 1)
|
||||
}
|
||||
msg.Context.Raw[metadataKeyMessageKind] = kind
|
||||
return msg
|
||||
}
|
||||
|
||||
func latestUserContent(messages []providers.Message) string {
|
||||
for i := len(messages) - 1; i >= 0; i-- {
|
||||
msg := messages[i]
|
||||
if msg.Role != "user" {
|
||||
continue
|
||||
}
|
||||
if content := strings.TrimSpace(msg.Content); content != "" {
|
||||
return content
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func toolFeedbackExplanationFromResponse(
|
||||
response *providers.LLMResponse,
|
||||
messages []providers.Message,
|
||||
maxLen int,
|
||||
) string {
|
||||
if response == nil {
|
||||
return ""
|
||||
}
|
||||
explanation := strings.TrimSpace(response.Content)
|
||||
if explanation == "" {
|
||||
explanation = toolFeedbackExplanationFromToolCalls(response.ToolCalls)
|
||||
}
|
||||
if explanation == "" {
|
||||
explanation = toolFeedbackExplanationFromMessages(messages)
|
||||
}
|
||||
return utils.Truncate(explanation, maxLen)
|
||||
}
|
||||
|
||||
func toolFeedbackExplanationFromToolCalls(toolCalls []providers.ToolCall) string {
|
||||
for _, tc := range toolCalls {
|
||||
if tc.ExtraContent == nil {
|
||||
continue
|
||||
}
|
||||
if explanation := strings.TrimSpace(tc.ExtraContent.ToolFeedbackExplanation); explanation != "" {
|
||||
return explanation
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func toolFeedbackExplanationForToolCall(
|
||||
response *providers.LLMResponse,
|
||||
toolCall providers.ToolCall,
|
||||
messages []providers.Message,
|
||||
maxLen int,
|
||||
) string {
|
||||
if toolCall.ExtraContent != nil {
|
||||
if explanation := strings.TrimSpace(toolCall.ExtraContent.ToolFeedbackExplanation); explanation != "" {
|
||||
return utils.Truncate(explanation, maxLen)
|
||||
}
|
||||
}
|
||||
if response == nil {
|
||||
return utils.Truncate(toolFeedbackExplanationFromMessages(messages), maxLen)
|
||||
}
|
||||
|
||||
explanation := strings.TrimSpace(response.Content)
|
||||
if explanation == "" {
|
||||
explanation = toolFeedbackExplanationFromMessages(messages)
|
||||
}
|
||||
return utils.Truncate(explanation, maxLen)
|
||||
}
|
||||
|
||||
func toolFeedbackExplanationFromMessages(messages []providers.Message) string {
|
||||
explanation := latestUserContent(messages)
|
||||
if explanation != "" {
|
||||
return utils.ToolFeedbackContinuationHint + ": " + explanation
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func shouldPublishToolFeedback(cfg *config.Config, ts *turnState) bool {
|
||||
if ts == nil || ts.channel == "" || ts.opts.SuppressToolFeedback {
|
||||
return false
|
||||
}
|
||||
return cfg != nil && cfg.Agents.Defaults.IsToolFeedbackEnabled()
|
||||
}
|
||||
|
||||
func cloneEventArguments(args map[string]any) map[string]any {
|
||||
if len(args) == 0 {
|
||||
return nil
|
||||
|
|
@ -11,6 +11,7 @@ import (
|
|||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
|
|
@ -210,6 +211,36 @@ func (cb *ContextBuilder) BuildSystemPromptWithCache() string {
|
|||
return prompt
|
||||
}
|
||||
|
||||
// EstimateSystemTokens estimates the token count of the full system message
|
||||
// that would be sent to the LLM, mirroring the composition logic in BuildMessages.
|
||||
// It includes: static prompt, dynamic context, active skills, and summary with
|
||||
// wrapping prefixes and separators. This avoids needing all per-request parameters
|
||||
// that BuildMessages requires (media, channel, chatID, sender, etc.).
|
||||
func (cb *ContextBuilder) EstimateSystemTokens(summary string, activeSkills []string) int {
|
||||
staticPrompt := cb.BuildSystemPromptWithCache()
|
||||
|
||||
// Dynamic context is small and varies per request; use a representative estimate.
|
||||
// Actual buildDynamicContext produces ~200-400 chars of time/runtime/session info.
|
||||
const dynamicContextChars = 300
|
||||
|
||||
totalChars := utf8.RuneCountInString(staticPrompt) + dynamicContextChars
|
||||
|
||||
if skillsText := cb.buildActiveSkillsContext(activeSkills); skillsText != "" {
|
||||
totalChars += utf8.RuneCountInString(skillsText)
|
||||
totalChars += 7 // separator \n\n---\n\n
|
||||
}
|
||||
|
||||
if summary != "" {
|
||||
// Matches the CONTEXT_SUMMARY: prefix added in BuildMessages
|
||||
const summaryPrefix = "CONTEXT_SUMMARY: The following is an approximate summary of prior conversation " +
|
||||
"for reference only. It may be incomplete or outdated — always defer to explicit instructions.\n\n"
|
||||
totalChars += utf8.RuneCountInString(summaryPrefix) + utf8.RuneCountInString(summary)
|
||||
totalChars += 7 // separator
|
||||
}
|
||||
|
||||
return totalChars * 2 / 5 // same heuristic as tokenizer.EstimateMessageTokens
|
||||
}
|
||||
|
||||
// InvalidateCache clears the cached system prompt.
|
||||
// Normally not needed because the cache auto-invalidates via mtime checks,
|
||||
// but this is useful for tests or explicit reload commands.
|
||||
|
|
|
|||
78
pkg/agent/context_usage.go
Normal file
78
pkg/agent/context_usage.go
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
package agent
|
||||
|
||||
import (
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
)
|
||||
|
||||
// computeContextUsage estimates current context window consumption for the
|
||||
// given agent and session. Includes history, system prompt (with dynamic context,
|
||||
// summary, and skills — mirroring BuildMessages composition), and tool definitions.
|
||||
// The output reserve (MaxTokens) is not counted as "used" but reduces the
|
||||
// effective budget, matching isOverContextBudget's compression trigger:
|
||||
//
|
||||
// compress when: history + system + tools + maxTokens > contextWindow
|
||||
// equivalent to: history + system + tools > contextWindow - maxTokens
|
||||
//
|
||||
// Returns nil when the agent or session is unavailable.
|
||||
func computeContextUsage(agent *AgentInstance, sessionKey string) *bus.ContextUsage {
|
||||
if agent == nil || agent.Sessions == nil {
|
||||
return nil
|
||||
}
|
||||
contextWindow := agent.ContextWindow
|
||||
if contextWindow <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// History tokens
|
||||
history := agent.Sessions.GetHistory(sessionKey)
|
||||
historyTokens := 0
|
||||
for _, m := range history {
|
||||
historyTokens += EstimateMessageTokens(m)
|
||||
}
|
||||
|
||||
// System message tokens: uses EstimateSystemTokens which mirrors
|
||||
// the full system message composition in BuildMessages (static prompt,
|
||||
// dynamic context, active skills, summary with wrapping prefix).
|
||||
systemTokens := 0
|
||||
if agent.ContextBuilder != nil {
|
||||
summary := agent.Sessions.GetSummary(sessionKey)
|
||||
// Pass nil for active skills: skills are only injected when the user
|
||||
// explicitly activates them via /use, which is rare. Using nil matches
|
||||
// the common case and avoids over-counting all installed skills.
|
||||
systemTokens = agent.ContextBuilder.EstimateSystemTokens(summary, nil)
|
||||
}
|
||||
|
||||
// Tool definition tokens
|
||||
toolTokens := 0
|
||||
if agent.Tools != nil {
|
||||
toolTokens = EstimateToolDefsTokens(agent.Tools.ToProviderDefs())
|
||||
}
|
||||
|
||||
// Used = history + system (includes summary) + tools
|
||||
usedTokens := historyTokens + systemTokens + toolTokens
|
||||
|
||||
// Effective budget = contextWindow minus output reserve (maxTokens)
|
||||
effectiveWindow := contextWindow - agent.MaxTokens
|
||||
if effectiveWindow < 0 {
|
||||
effectiveWindow = contextWindow
|
||||
}
|
||||
|
||||
// compressAt = effectiveWindow: aligns with isOverContextBudget's
|
||||
// proactive trigger (msgTokens + toolTokens + maxTokens > contextWindow).
|
||||
compressAt := effectiveWindow
|
||||
|
||||
usedPercent := 0
|
||||
if compressAt > 0 {
|
||||
usedPercent = usedTokens * 100 / compressAt
|
||||
}
|
||||
if usedPercent > 100 {
|
||||
usedPercent = 100
|
||||
}
|
||||
|
||||
return &bus.ContextUsage{
|
||||
UsedTokens: usedTokens,
|
||||
TotalTokens: contextWindow,
|
||||
CompressAtTokens: compressAt,
|
||||
UsedPercent: usedPercent,
|
||||
}
|
||||
}
|
||||
|
|
@ -4,6 +4,7 @@ import (
|
|||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
|
@ -403,6 +404,24 @@ func (h *toolRewriteHook) AfterTool(
|
|||
return next, HookDecision{Action: HookActionModify}, nil
|
||||
}
|
||||
|
||||
type toolRenameHook struct{}
|
||||
|
||||
func (h *toolRenameHook) BeforeTool(
|
||||
ctx context.Context,
|
||||
call *ToolCallHookRequest,
|
||||
) (*ToolCallHookRequest, HookDecision, error) {
|
||||
next := call.Clone()
|
||||
next.Tool = "echo_text_rewritten"
|
||||
return next, HookDecision{Action: HookActionModify}, nil
|
||||
}
|
||||
|
||||
func (h *toolRenameHook) AfterTool(
|
||||
ctx context.Context,
|
||||
result *ToolResultHookResponse,
|
||||
) (*ToolResultHookResponse, HookDecision, error) {
|
||||
return result.Clone(), HookDecision{Action: HookActionContinue}, nil
|
||||
}
|
||||
|
||||
func TestAgentLoop_Hooks_ToolInterceptorCanRewrite(t *testing.T) {
|
||||
provider := &toolHookProvider{}
|
||||
al, agent, cleanup := newHookTestLoop(t, provider)
|
||||
|
|
@ -430,6 +449,75 @@ func TestAgentLoop_Hooks_ToolInterceptorCanRewrite(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
type echoTextRewrittenTool struct{}
|
||||
|
||||
func (t *echoTextRewrittenTool) Name() string {
|
||||
return "echo_text_rewritten"
|
||||
}
|
||||
|
||||
func (t *echoTextRewrittenTool) Description() string {
|
||||
return "echo a rewritten text argument"
|
||||
}
|
||||
|
||||
func (t *echoTextRewrittenTool) Parameters() map[string]any {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{
|
||||
"type": "string",
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (t *echoTextRewrittenTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
|
||||
text, _ := args["text"].(string)
|
||||
return tools.SilentResult("rewritten:" + text)
|
||||
}
|
||||
|
||||
func TestAgentLoop_Hooks_ToolFeedbackUsesRewrittenToolName(t *testing.T) {
|
||||
provider := &toolHookProvider{}
|
||||
al, agent, cleanup := newHookTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
al.cfg.Agents.Defaults.ToolFeedback.Enabled = true
|
||||
al.RegisterTool(&echoTextTool{})
|
||||
al.RegisterTool(&echoTextRewrittenTool{})
|
||||
if err := al.MountHook(NamedHook("tool-rename", &toolRenameHook{})); err != nil {
|
||||
t.Fatalf("MountHook failed: %v", err)
|
||||
}
|
||||
|
||||
_, err := al.runAgentLoop(context.Background(), agent, processOptions{
|
||||
SessionKey: "session-1",
|
||||
Channel: "cli",
|
||||
ChatID: "direct",
|
||||
UserMessage: "run tool",
|
||||
DefaultResponse: defaultResponse,
|
||||
EnableSummary: false,
|
||||
SendResponse: false,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("runAgentLoop failed: %v", err)
|
||||
}
|
||||
|
||||
msgBus, ok := al.bus.(*bus.MessageBus)
|
||||
if !ok {
|
||||
t.Fatalf("expected concrete MessageBus, got %T", al.bus)
|
||||
}
|
||||
|
||||
select {
|
||||
case outbound := <-msgBus.OutboundChan():
|
||||
if !strings.Contains(outbound.Content, "`echo_text_rewritten`") {
|
||||
t.Fatalf("tool feedback content = %q, want rewritten tool name", outbound.Content)
|
||||
}
|
||||
if strings.Contains(outbound.Content, "`echo_text`") {
|
||||
t.Fatalf("tool feedback content = %q, want no original tool name", outbound.Content)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("expected outbound tool feedback")
|
||||
}
|
||||
}
|
||||
|
||||
type denyApprovalHook struct{}
|
||||
|
||||
func (h *denyApprovalHook) ApproveTool(ctx context.Context, req *ToolApprovalRequest) (ApprovalDecision, error) {
|
||||
|
|
@ -709,9 +797,10 @@ func TestAgentLoop_HookRespond_MediaError(t *testing.T) {
|
|||
t.Fatalf("MountHook failed: %v", err)
|
||||
}
|
||||
|
||||
al.channelManager = newStartedTestChannelManager(t, al.bus, al.mediaStore, "discord", &errorMediaChannel{
|
||||
sendErr: errors.New("channel unavailable"),
|
||||
})
|
||||
al.channelManager = newStartedTestChannelManager(t,
|
||||
al.bus.(*bus.MessageBus), al.mediaStore, "discord", &errorMediaChannel{
|
||||
sendErr: errors.New("channel unavailable"),
|
||||
})
|
||||
|
||||
sub := al.SubscribeEvents(16)
|
||||
defer al.UnsubscribeEvents(sub.ID)
|
||||
|
|
@ -803,6 +892,77 @@ func TestAgentLoop_HookRespond_BusFallback(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestAgentLoop_HookRespond_ResponseHandledMediaPreservesOutboundContext(t *testing.T) {
|
||||
provider := &multiToolProvider{
|
||||
toolCalls: []providers.ToolCall{
|
||||
{ID: "call-1", Name: "media_tool", Arguments: map[string]any{}},
|
||||
},
|
||||
finalContent: "done",
|
||||
}
|
||||
al, agent, cleanup := newHookTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
hook := &respondWithMediaHook{
|
||||
respondTools: map[string]bool{"media_tool": true},
|
||||
media: []string{"media://test/image.png"},
|
||||
responseHandled: true,
|
||||
forLLM: "media sent successfully",
|
||||
}
|
||||
if err := al.MountHook(NamedHook("media-hook", hook)); err != nil {
|
||||
t.Fatalf("MountHook failed: %v", err)
|
||||
}
|
||||
|
||||
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
|
||||
al.channelManager = newStartedTestChannelManager(t,
|
||||
al.bus.(*bus.MessageBus), al.mediaStore, "telegram", telegramChannel)
|
||||
|
||||
_, err := al.runAgentLoop(context.Background(), agent, processOptions{
|
||||
Dispatch: DispatchRequest{
|
||||
SessionKey: "session-topic-media",
|
||||
SessionScope: &session.SessionScope{
|
||||
Version: session.ScopeVersionV1,
|
||||
AgentID: agent.ID,
|
||||
Channel: "telegram",
|
||||
Dimensions: []string{"chat"},
|
||||
Values: map[string]string{
|
||||
"chat": "forum:-100123/42",
|
||||
},
|
||||
},
|
||||
InboundContext: &bus.InboundContext{
|
||||
Channel: "telegram",
|
||||
ChatID: "-100123",
|
||||
TopicID: "42",
|
||||
ChatType: "group",
|
||||
SenderID: "user1",
|
||||
},
|
||||
UserMessage: "send media",
|
||||
},
|
||||
DefaultResponse: defaultResponse,
|
||||
EnableSummary: false,
|
||||
SendResponse: false,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("runAgentLoop failed: %v", err)
|
||||
}
|
||||
|
||||
if len(telegramChannel.sentMedia) != 1 {
|
||||
t.Fatalf("expected exactly 1 sent media message, got %d", len(telegramChannel.sentMedia))
|
||||
}
|
||||
sent := telegramChannel.sentMedia[0]
|
||||
if sent.Context.Channel != "telegram" || sent.Context.ChatID != "-100123" || sent.Context.TopicID != "42" {
|
||||
t.Fatalf("unexpected media context: %+v", sent.Context)
|
||||
}
|
||||
if sent.AgentID != agent.ID {
|
||||
t.Fatalf("sent media agent_id = %q, want %q", sent.AgentID, agent.ID)
|
||||
}
|
||||
if sent.SessionKey != "session-topic-media" {
|
||||
t.Fatalf("sent media session_key = %q, want session-topic-media", sent.SessionKey)
|
||||
}
|
||||
if sent.Scope == nil || sent.Scope.Values["chat"] != "forum:-100123/42" {
|
||||
t.Fatalf("unexpected sent media scope: %+v", sent.Scope)
|
||||
}
|
||||
}
|
||||
|
||||
type multiToolProvider struct {
|
||||
mu sync.Mutex
|
||||
callCount int
|
||||
|
|
@ -880,7 +1040,11 @@ func TestAgentLoop_HookRespond_InterruptSkipsRemaining(t *testing.T) {
|
|||
resultCh <- result{resp: resp, err: err}
|
||||
}()
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
select {
|
||||
case <-tool1ExecCh:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("timeout waiting for tool execution to start")
|
||||
}
|
||||
|
||||
if err := al.InterruptGraceful("stop now"); err != nil {
|
||||
t.Fatalf("InterruptGraceful failed: %v", err)
|
||||
|
|
|
|||
|
|
@ -270,8 +270,8 @@ func populateCandidateProvidersFromNames(
|
|||
map[string]any{"name": name, "error": err.Error()})
|
||||
continue
|
||||
}
|
||||
protocol, modelID := providers.ExtractProtocol(strings.TrimSpace(mc.Model))
|
||||
key := providers.ModelKey(providers.NormalizeProvider(protocol), modelID)
|
||||
protocol, modelID := providers.ExtractProtocol(mc)
|
||||
key := providers.ModelKey(protocol, modelID)
|
||||
if _, exists := out[key]; exists {
|
||||
continue
|
||||
}
|
||||
|
|
|
|||
|
|
@ -104,6 +104,7 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
|
|||
name string
|
||||
aliasName string
|
||||
modelName string
|
||||
provider string
|
||||
apiBase string
|
||||
wantProvider string
|
||||
wantModel string
|
||||
|
|
@ -124,6 +125,15 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
|
|||
wantProvider: "openai",
|
||||
wantModel: "glm-5",
|
||||
},
|
||||
{
|
||||
name: "explicit provider overrides model prefix",
|
||||
aliasName: "nvidia-gpt",
|
||||
modelName: "z-ai/glm-5.1",
|
||||
provider: "nvidia",
|
||||
apiBase: "https://integrate.api.nvidia.com/v1",
|
||||
wantProvider: "nvidia",
|
||||
wantModel: "z-ai/glm-5.1",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
|
|
@ -145,6 +155,7 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
|
|||
{
|
||||
ModelName: tt.aliasName,
|
||||
Model: tt.modelName,
|
||||
Provider: tt.provider,
|
||||
APIBase: tt.apiBase,
|
||||
},
|
||||
},
|
||||
|
|
@ -218,6 +229,43 @@ func TestNewAgentInstance_PreservesDistinctLimiterIdentityForSharedResolvedModel
|
|||
}
|
||||
}
|
||||
|
||||
func TestNewAgentInstance_PreservesConfigIdentityForExplicitProviderModelRef(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "nvidia/z-ai/glm-5.1",
|
||||
},
|
||||
},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "nvidia-glm",
|
||||
Provider: "nvidia",
|
||||
Model: "z-ai/glm-5.1",
|
||||
RPM: 7,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, &mockProvider{})
|
||||
if len(agent.Candidates) != 1 {
|
||||
t.Fatalf("len(Candidates) = %d, want 1", len(agent.Candidates))
|
||||
}
|
||||
|
||||
candidate := agent.Candidates[0]
|
||||
if candidate.Provider != "nvidia" || candidate.Model != "z-ai/glm-5.1" {
|
||||
t.Fatalf("candidate = %s/%s, want nvidia/z-ai/glm-5.1", candidate.Provider, candidate.Model)
|
||||
}
|
||||
if candidate.IdentityKey != "model_name:nvidia-glm" {
|
||||
t.Fatalf("identity key = %q, want %q", candidate.IdentityKey, "model_name:nvidia-glm")
|
||||
}
|
||||
if candidate.RPM != 7 {
|
||||
t.Fatalf("RPM = %d, want 7", candidate.RPM)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewAgentInstance_AllowsMediaTempDirForReadListAndExec(t *testing.T) {
|
||||
workspace := t.TempDir()
|
||||
mediaDir := media.TempDir()
|
||||
|
|
|
|||
47
pkg/agent/interfaces/interfaces.go
Normal file
47
pkg/agent/interfaces/interfaces.go
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package interfaces
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
)
|
||||
|
||||
// MessageBus publishes inbound and outbound messages.
|
||||
// It is the primary communication channel for the agent loop.
|
||||
type MessageBus interface {
|
||||
// PublishInbound sends an inbound message to be processed.
|
||||
PublishInbound(ctx context.Context, msg bus.InboundMessage) error
|
||||
|
||||
// PublishOutbound sends an outbound message to the appropriate channel.
|
||||
PublishOutbound(ctx context.Context, msg bus.OutboundMessage) error
|
||||
|
||||
// PublishOutboundMedia sends an outbound media message.
|
||||
PublishOutboundMedia(ctx context.Context, msg bus.OutboundMediaMessage) error
|
||||
|
||||
// InboundChan returns the channel for receiving inbound messages.
|
||||
InboundChan() <-chan bus.InboundMessage
|
||||
}
|
||||
|
||||
// ChannelManager manages channel lifecycle and provides channel access.
|
||||
type ChannelManager interface {
|
||||
// GetChannel returns the channel with the given name.
|
||||
GetChannel(name string) (channels.Channel, bool)
|
||||
|
||||
// GetEnabledChannels returns the list of enabled channel names.
|
||||
GetEnabledChannels() []string
|
||||
|
||||
// InvokeTypingStop signals that typing has stopped.
|
||||
InvokeTypingStop(channel, chatID string)
|
||||
|
||||
// SendMessage sends a text message to the specified channel and chat.
|
||||
SendMessage(ctx context.Context, msg bus.OutboundMessage) error
|
||||
|
||||
// SendMedia sends a media message to the specified channel and chat.
|
||||
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error
|
||||
|
||||
// SendPlaceholder sends a placeholder message (e.g., for audio transcription).
|
||||
SendPlaceholder(ctx context.Context, channel, chatID string) bool
|
||||
}
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -37,14 +37,14 @@ func candidateFromModelConfig(
|
|||
return providers.FallbackCandidate{}, false
|
||||
}
|
||||
|
||||
ref := providers.ParseModelRef(ensureProtocolModel(mc.Model), defaultProvider)
|
||||
if ref == nil {
|
||||
protocol, modelID := providers.ExtractProtocol(mc)
|
||||
if strings.TrimSpace(modelID) == "" {
|
||||
return providers.FallbackCandidate{}, false
|
||||
}
|
||||
|
||||
return providers.FallbackCandidate{
|
||||
Provider: ref.Provider,
|
||||
Model: ref.Model,
|
||||
Provider: protocol,
|
||||
Model: modelID,
|
||||
RPM: mc.RPM,
|
||||
IdentityKey: modelConfigIdentityKey(mc),
|
||||
}, true
|
||||
|
|
@ -60,6 +60,12 @@ func lookupModelConfigByRef(cfg *config.Config, raw string) *config.ModelConfig
|
|||
return mc
|
||||
}
|
||||
|
||||
rawRef := providers.ParseModelRef(raw, "")
|
||||
rawKey := ""
|
||||
if rawRef != nil && strings.TrimSpace(rawRef.Provider) != "" && strings.TrimSpace(rawRef.Model) != "" {
|
||||
rawKey = providers.ModelKey(rawRef.Provider, rawRef.Model)
|
||||
}
|
||||
|
||||
for i := range cfg.ModelList {
|
||||
mc := cfg.ModelList[i]
|
||||
if mc == nil {
|
||||
|
|
@ -72,10 +78,13 @@ func lookupModelConfigByRef(cfg *config.Config, raw string) *config.ModelConfig
|
|||
if fullModel == raw {
|
||||
return mc
|
||||
}
|
||||
_, modelID := providers.ExtractProtocol(fullModel)
|
||||
protocol, modelID := providers.ExtractProtocol(mc)
|
||||
if modelID == raw {
|
||||
return mc
|
||||
}
|
||||
if rawKey != "" && providers.ModelKey(protocol, modelID) == rawKey {
|
||||
return mc
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
|
|
|
|||
40
pkg/agent/pipeline.go
Normal file
40
pkg/agent/pipeline.go
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"github.com/sipeed/picoclaw/pkg/agent/interfaces"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// Pipeline holds the runtime dependencies used by Pipeline methods.
|
||||
// It is constructed by runTurn via NewPipeline and passed to sub-methods
|
||||
// so that the coordinator can delegate phase execution.
|
||||
type Pipeline struct {
|
||||
Bus interfaces.MessageBus
|
||||
Cfg *config.Config
|
||||
ContextManager ContextManager
|
||||
Hooks *HookManager
|
||||
Fallback *providers.FallbackChain
|
||||
ChannelManager interfaces.ChannelManager
|
||||
MediaStore media.MediaStore
|
||||
Steering any // TODO: *Steering
|
||||
al *AgentLoop
|
||||
}
|
||||
|
||||
// NewPipeline creates a Pipeline from an AgentLoop instance.
|
||||
func NewPipeline(al *AgentLoop) *Pipeline {
|
||||
return &Pipeline{
|
||||
Bus: al.bus,
|
||||
Cfg: al.GetConfig(),
|
||||
ContextManager: al.contextManager,
|
||||
Hooks: al.hooks,
|
||||
Fallback: al.fallback,
|
||||
ChannelManager: al.channelManager,
|
||||
MediaStore: al.mediaStore,
|
||||
Steering: al.steering,
|
||||
al: al,
|
||||
}
|
||||
}
|
||||
716
pkg/agent/pipeline_execute.go
Normal file
716
pkg/agent/pipeline_execute.go
Normal file
|
|
@ -0,0 +1,716 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/constants"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
// ExecuteTools executes the tool loop, handling BeforeTool/ApproveTool/AfterTool hooks,
|
||||
// tool execution with async callbacks, media delivery, and steering injection.
|
||||
// Returns ToolControl indicating what the coordinator should do next:
|
||||
// - ToolControlContinue: all tool results handled, pendingMessages or steering exists, continue turn
|
||||
// - ToolControlBreak: tool loop exited, proceed to coordinator's hardAbort/finalContent/finalize
|
||||
func (p *Pipeline) ExecuteTools(
|
||||
ctx context.Context,
|
||||
turnCtx context.Context,
|
||||
ts *turnState,
|
||||
exec *turnExecution,
|
||||
iteration int,
|
||||
) ToolControl {
|
||||
al := p.al
|
||||
normalizedToolCalls := exec.normalizedToolCalls
|
||||
|
||||
ts.setPhase(TurnPhaseTools)
|
||||
messages := exec.messages
|
||||
handledAttachments := make([]providers.Attachment, 0)
|
||||
|
||||
toolLoop:
|
||||
for i, tc := range normalizedToolCalls {
|
||||
if ts.hardAbortRequested() {
|
||||
exec.abortedByHardAbort = true
|
||||
return ToolControlBreak
|
||||
}
|
||||
|
||||
toolName := tc.Name
|
||||
toolArgs := cloneStringAnyMap(tc.Arguments)
|
||||
|
||||
if al.hooks != nil {
|
||||
toolReq, decision := al.hooks.BeforeTool(turnCtx, &ToolCallHookRequest{
|
||||
Meta: ts.eventMeta("runTurn", "turn.tool.before"),
|
||||
Context: cloneTurnContext(ts.turnCtx),
|
||||
Tool: toolName,
|
||||
Arguments: toolArgs,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if toolReq != nil {
|
||||
toolName = toolReq.Tool
|
||||
toolArgs = toolReq.Arguments
|
||||
}
|
||||
case HookActionRespond:
|
||||
if toolReq != nil && toolReq.HookResult != nil {
|
||||
hookResult := toolReq.HookResult
|
||||
|
||||
argsJSON, _ := json.Marshal(toolArgs)
|
||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||
logger.InfoCF("agent", fmt.Sprintf("Tool call (hook respond): %s(%s)", toolName, argsPreview),
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tool": toolName,
|
||||
"iteration": iteration,
|
||||
})
|
||||
|
||||
al.emitEvent(
|
||||
EventKindToolExecStart,
|
||||
ts.eventMeta("runTurn", "turn.tool.start"),
|
||||
ToolExecStartPayload{
|
||||
Tool: toolName,
|
||||
Arguments: cloneEventArguments(toolArgs),
|
||||
},
|
||||
)
|
||||
|
||||
if shouldPublishToolFeedback(al.cfg, ts) {
|
||||
toolFeedbackExplanation := toolFeedbackExplanationForToolCall(
|
||||
exec.response,
|
||||
tc,
|
||||
messages,
|
||||
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
||||
)
|
||||
feedbackMsg := utils.FormatToolFeedbackMessage(toolName, toolFeedbackExplanation)
|
||||
fbCtx, fbCancel := context.WithTimeout(turnCtx, 3*time.Second)
|
||||
_ = al.bus.PublishOutbound(fbCtx, outboundMessageForTurnWithKind(ts, feedbackMsg, messageKindToolFeedback))
|
||||
fbCancel()
|
||||
}
|
||||
|
||||
toolDuration := time.Duration(0)
|
||||
|
||||
shouldSendForUser := !hookResult.Silent && hookResult.ForUser != "" &&
|
||||
(ts.opts.SendResponse || hookResult.ResponseHandled)
|
||||
if shouldSendForUser {
|
||||
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
Context: bus.InboundContext{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Raw: map[string]string{
|
||||
"is_tool_call": "true",
|
||||
},
|
||||
},
|
||||
Content: hookResult.ForUser,
|
||||
})
|
||||
}
|
||||
|
||||
if len(hookResult.Media) > 0 && hookResult.ResponseHandled {
|
||||
parts := make([]bus.MediaPart, 0, len(hookResult.Media))
|
||||
for _, ref := range hookResult.Media {
|
||||
part := bus.MediaPart{Ref: ref}
|
||||
if al.mediaStore != nil {
|
||||
if _, meta, err := al.mediaStore.ResolveWithMeta(ref); err == nil {
|
||||
part.Filename = meta.Filename
|
||||
part.ContentType = meta.ContentType
|
||||
part.Type = inferMediaType(meta.Filename, meta.ContentType)
|
||||
}
|
||||
}
|
||||
parts = append(parts, part)
|
||||
}
|
||||
outboundMedia := bus.OutboundMediaMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Context: outboundContextFromInbound(
|
||||
ts.opts.Dispatch.InboundContext,
|
||||
ts.channel,
|
||||
ts.chatID,
|
||||
ts.opts.Dispatch.ReplyToMessageID(),
|
||||
),
|
||||
AgentID: ts.agent.ID,
|
||||
SessionKey: ts.sessionKey,
|
||||
Scope: outboundScopeFromSessionScope(ts.opts.Dispatch.SessionScope),
|
||||
Parts: parts,
|
||||
}
|
||||
if al.channelManager != nil && ts.channel != "" && !constants.IsInternalChannel(ts.channel) {
|
||||
if err := al.channelManager.SendMedia(ctx, outboundMedia); err != nil {
|
||||
logger.WarnCF("agent", "Failed to deliver hook media",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tool": toolName,
|
||||
"channel": ts.channel,
|
||||
"chat_id": ts.chatID,
|
||||
"error": err.Error(),
|
||||
})
|
||||
hookResult.IsError = true
|
||||
hookResult.ForLLM = fmt.Sprintf("failed to deliver attachment: %v", err)
|
||||
} else {
|
||||
handledAttachments = append(
|
||||
handledAttachments,
|
||||
buildProviderAttachments(al.mediaStore, hookResult.Media)...,
|
||||
)
|
||||
}
|
||||
} else if al.bus != nil {
|
||||
al.bus.PublishOutboundMedia(ctx, outboundMedia)
|
||||
hookResult.ResponseHandled = false
|
||||
}
|
||||
}
|
||||
|
||||
if !hookResult.ResponseHandled {
|
||||
exec.allResponsesHandled = false
|
||||
}
|
||||
|
||||
contentForLLM := hookResult.ContentForLLM()
|
||||
if al.cfg.Tools.IsFilterSensitiveDataEnabled() {
|
||||
contentForLLM = al.cfg.FilterSensitiveData(contentForLLM)
|
||||
}
|
||||
|
||||
toolResultMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: contentForLLM,
|
||||
ToolCallID: tc.ID,
|
||||
}
|
||||
|
||||
if len(hookResult.Media) > 0 && !hookResult.ResponseHandled {
|
||||
hookResult.ArtifactTags = buildArtifactTags(al.mediaStore, hookResult.Media)
|
||||
contentForLLM = hookResult.ContentForLLM()
|
||||
if al.cfg.Tools.IsFilterSensitiveDataEnabled() {
|
||||
contentForLLM = al.cfg.FilterSensitiveData(contentForLLM)
|
||||
}
|
||||
toolResultMsg.Content = contentForLLM
|
||||
toolResultMsg.Media = append(toolResultMsg.Media, hookResult.Media...)
|
||||
}
|
||||
|
||||
al.emitEvent(
|
||||
EventKindToolExecEnd,
|
||||
ts.eventMeta("runTurn", "turn.tool.end"),
|
||||
ToolExecEndPayload{
|
||||
Tool: toolName,
|
||||
Duration: toolDuration,
|
||||
ForLLMLen: len(contentForLLM),
|
||||
ForUserLen: len(hookResult.ForUser),
|
||||
IsError: hookResult.IsError,
|
||||
Async: hookResult.Async,
|
||||
},
|
||||
)
|
||||
|
||||
messages = append(messages, toolResultMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, toolResultMsg)
|
||||
ts.recordPersistedMessage(toolResultMsg)
|
||||
ts.ingestMessage(turnCtx, al, toolResultMsg)
|
||||
}
|
||||
|
||||
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 {
|
||||
exec.pendingMessages = append(exec.pendingMessages, steerMsgs...)
|
||||
}
|
||||
|
||||
skipReason := ""
|
||||
skipMessage := ""
|
||||
if len(exec.pendingMessages) > 0 {
|
||||
skipReason = "queued user steering message"
|
||||
skipMessage = "Skipped due to queued user message."
|
||||
} else if gracefulPending, _ := ts.gracefulInterruptRequested(); gracefulPending {
|
||||
skipReason = "graceful interrupt requested"
|
||||
skipMessage = "Skipped due to graceful interrupt."
|
||||
}
|
||||
|
||||
if skipReason != "" {
|
||||
remaining := len(normalizedToolCalls) - i - 1
|
||||
if remaining > 0 {
|
||||
logger.InfoCF("agent", "Turn checkpoint: skipping remaining tools after hook respond",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"completed": i + 1,
|
||||
"skipped": remaining,
|
||||
"reason": skipReason,
|
||||
})
|
||||
for j := i + 1; j < len(normalizedToolCalls); j++ {
|
||||
skippedTC := normalizedToolCalls[j]
|
||||
al.emitEvent(
|
||||
EventKindToolExecSkipped,
|
||||
ts.eventMeta("runTurn", "turn.tool.skipped"),
|
||||
ToolExecSkippedPayload{
|
||||
Tool: skippedTC.Name,
|
||||
Reason: skipReason,
|
||||
},
|
||||
)
|
||||
skippedMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: skipMessage,
|
||||
ToolCallID: skippedTC.ID,
|
||||
}
|
||||
messages = append(messages, skippedMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, skippedMsg)
|
||||
ts.recordPersistedMessage(skippedMsg)
|
||||
}
|
||||
}
|
||||
}
|
||||
break toolLoop
|
||||
}
|
||||
|
||||
if ts.pendingResults != nil {
|
||||
select {
|
||||
case result, ok := <-ts.pendingResults:
|
||||
if ok && result != nil && result.ForLLM != "" {
|
||||
content := al.cfg.FilterSensitiveData(result.ForLLM)
|
||||
msg := providers.Message{Role: "user", Content: fmt.Sprintf("[SubTurn Result] %s", content)}
|
||||
messages = append(messages, msg)
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, msg)
|
||||
}
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
continue
|
||||
}
|
||||
logger.WarnCF("agent", "Hook returned respond action but no HookResult provided",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tool": toolName,
|
||||
"action": "respond",
|
||||
})
|
||||
case HookActionDenyTool:
|
||||
exec.allResponsesHandled = false
|
||||
denyContent := hookDeniedToolContent("Tool execution denied by hook", decision.Reason)
|
||||
al.emitEvent(
|
||||
EventKindToolExecSkipped,
|
||||
ts.eventMeta("runTurn", "turn.tool.skipped"),
|
||||
ToolExecSkippedPayload{
|
||||
Tool: toolName,
|
||||
Reason: denyContent,
|
||||
},
|
||||
)
|
||||
deniedMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: denyContent,
|
||||
ToolCallID: tc.ID,
|
||||
}
|
||||
messages = append(messages, deniedMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, deniedMsg)
|
||||
ts.recordPersistedMessage(deniedMsg)
|
||||
}
|
||||
continue
|
||||
case HookActionAbortTurn:
|
||||
exec.abortedByHook = true
|
||||
return ToolControlBreak
|
||||
case HookActionHardAbort:
|
||||
_ = ts.requestHardAbort()
|
||||
exec.abortedByHardAbort = true
|
||||
return ToolControlBreak
|
||||
}
|
||||
}
|
||||
|
||||
if al.hooks != nil {
|
||||
approval := al.hooks.ApproveTool(turnCtx, &ToolApprovalRequest{
|
||||
Meta: ts.eventMeta("runTurn", "turn.tool.approve"),
|
||||
Context: cloneTurnContext(ts.turnCtx),
|
||||
Tool: toolName,
|
||||
Arguments: toolArgs,
|
||||
})
|
||||
if !approval.Approved {
|
||||
exec.allResponsesHandled = false
|
||||
denyContent := hookDeniedToolContent("Tool execution denied by approval hook", approval.Reason)
|
||||
al.emitEvent(
|
||||
EventKindToolExecSkipped,
|
||||
ts.eventMeta("runTurn", "turn.tool.skipped"),
|
||||
ToolExecSkippedPayload{
|
||||
Tool: toolName,
|
||||
Reason: denyContent,
|
||||
},
|
||||
)
|
||||
deniedMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: denyContent,
|
||||
ToolCallID: tc.ID,
|
||||
}
|
||||
messages = append(messages, deniedMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, deniedMsg)
|
||||
ts.recordPersistedMessage(deniedMsg)
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
argsJSON, _ := json.Marshal(toolArgs)
|
||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", toolName, argsPreview),
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tool": toolName,
|
||||
"iteration": iteration,
|
||||
})
|
||||
al.emitEvent(
|
||||
EventKindToolExecStart,
|
||||
ts.eventMeta("runTurn", "turn.tool.start"),
|
||||
ToolExecStartPayload{
|
||||
Tool: toolName,
|
||||
Arguments: cloneEventArguments(toolArgs),
|
||||
},
|
||||
)
|
||||
|
||||
if shouldPublishToolFeedback(al.cfg, ts) {
|
||||
toolFeedbackExplanation := toolFeedbackExplanationForToolCall(
|
||||
exec.response,
|
||||
tc,
|
||||
messages,
|
||||
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
||||
)
|
||||
feedbackMsg := utils.FormatToolFeedbackMessage(toolName, toolFeedbackExplanation)
|
||||
fbCtx, fbCancel := context.WithTimeout(turnCtx, 3*time.Second)
|
||||
_ = al.bus.PublishOutbound(fbCtx, outboundMessageForTurnWithKind(ts, feedbackMsg, messageKindToolFeedback))
|
||||
fbCancel()
|
||||
}
|
||||
|
||||
toolCallID := tc.ID
|
||||
asyncToolName := toolName
|
||||
asyncCallback := func(_ context.Context, result *tools.ToolResult) {
|
||||
if !result.Silent && result.ForUser != "" {
|
||||
outCtx, outCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer outCancel()
|
||||
_ = al.bus.PublishOutbound(outCtx, outboundMessageForTurn(ts, result.ForUser))
|
||||
}
|
||||
|
||||
content := result.ContentForLLM()
|
||||
if content == "" {
|
||||
return
|
||||
}
|
||||
|
||||
content = al.cfg.FilterSensitiveData(content)
|
||||
|
||||
logger.InfoCF("agent", "Async tool completed, publishing result",
|
||||
map[string]any{
|
||||
"tool": asyncToolName,
|
||||
"content_len": len(content),
|
||||
"channel": ts.channel,
|
||||
})
|
||||
al.emitEvent(
|
||||
EventKindFollowUpQueued,
|
||||
ts.scope.meta(iteration, "runTurn", "turn.follow_up.queued"),
|
||||
FollowUpQueuedPayload{
|
||||
SourceTool: asyncToolName,
|
||||
ContentLen: len(content),
|
||||
},
|
||||
)
|
||||
pubCtx, pubCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer pubCancel()
|
||||
_ = al.bus.PublishInbound(pubCtx, bus.InboundMessage{
|
||||
Context: bus.InboundContext{
|
||||
Channel: "system",
|
||||
ChatID: fmt.Sprintf("%s:%s", ts.channel, ts.chatID),
|
||||
ChatType: "direct",
|
||||
SenderID: fmt.Sprintf("async:%s", asyncToolName),
|
||||
},
|
||||
Content: content,
|
||||
})
|
||||
}
|
||||
|
||||
toolStart := time.Now()
|
||||
execCtx := tools.WithToolInboundContext(
|
||||
turnCtx,
|
||||
ts.channel,
|
||||
ts.chatID,
|
||||
ts.opts.Dispatch.MessageID(),
|
||||
ts.opts.Dispatch.ReplyToMessageID(),
|
||||
)
|
||||
execCtx = tools.WithToolSessionContext(
|
||||
execCtx,
|
||||
ts.agent.ID,
|
||||
ts.sessionKey,
|
||||
ts.opts.Dispatch.SessionScope,
|
||||
)
|
||||
toolResult := ts.agent.Tools.ExecuteWithContext(
|
||||
execCtx,
|
||||
toolName,
|
||||
toolArgs,
|
||||
ts.channel,
|
||||
ts.chatID,
|
||||
asyncCallback,
|
||||
)
|
||||
toolDuration := time.Since(toolStart)
|
||||
|
||||
if ts.hardAbortRequested() {
|
||||
exec.abortedByHardAbort = true
|
||||
return ToolControlBreak
|
||||
}
|
||||
|
||||
if al.hooks != nil {
|
||||
toolResp, decision := al.hooks.AfterTool(turnCtx, &ToolResultHookResponse{
|
||||
Meta: ts.eventMeta("runTurn", "turn.tool.after"),
|
||||
Context: cloneTurnContext(ts.turnCtx),
|
||||
Tool: toolName,
|
||||
Arguments: toolArgs,
|
||||
Result: toolResult,
|
||||
Duration: toolDuration,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if toolResp != nil {
|
||||
if toolResp.Tool != "" {
|
||||
toolName = toolResp.Tool
|
||||
}
|
||||
if toolResp.Result != nil {
|
||||
toolResult = toolResp.Result
|
||||
}
|
||||
}
|
||||
case HookActionAbortTurn:
|
||||
exec.abortedByHook = true
|
||||
return ToolControlBreak
|
||||
case HookActionHardAbort:
|
||||
_ = ts.requestHardAbort()
|
||||
exec.abortedByHardAbort = true
|
||||
return ToolControlBreak
|
||||
}
|
||||
}
|
||||
|
||||
if toolResult == nil {
|
||||
toolResult = tools.ErrorResult("hook returned nil tool result")
|
||||
}
|
||||
|
||||
if len(toolResult.Media) > 0 && toolResult.ResponseHandled {
|
||||
parts := make([]bus.MediaPart, 0, len(toolResult.Media))
|
||||
for _, ref := range toolResult.Media {
|
||||
part := bus.MediaPart{Ref: ref}
|
||||
if al.mediaStore != nil {
|
||||
if _, meta, err := al.mediaStore.ResolveWithMeta(ref); err == nil {
|
||||
part.Filename = meta.Filename
|
||||
part.ContentType = meta.ContentType
|
||||
part.Type = inferMediaType(meta.Filename, meta.ContentType)
|
||||
}
|
||||
}
|
||||
parts = append(parts, part)
|
||||
}
|
||||
outboundMedia := bus.OutboundMediaMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Context: outboundContextFromInbound(
|
||||
ts.opts.Dispatch.InboundContext,
|
||||
ts.channel,
|
||||
ts.chatID,
|
||||
ts.opts.Dispatch.ReplyToMessageID(),
|
||||
),
|
||||
AgentID: ts.agent.ID,
|
||||
SessionKey: ts.sessionKey,
|
||||
Scope: outboundScopeFromSessionScope(ts.opts.Dispatch.SessionScope),
|
||||
Parts: parts,
|
||||
}
|
||||
if al.channelManager != nil && ts.channel != "" && !constants.IsInternalChannel(ts.channel) {
|
||||
if err := al.channelManager.SendMedia(ctx, outboundMedia); err != nil {
|
||||
logger.WarnCF("agent", "Failed to deliver handled tool media",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tool": toolName,
|
||||
"channel": ts.channel,
|
||||
"chat_id": ts.chatID,
|
||||
"error": err.Error(),
|
||||
})
|
||||
toolResult = tools.ErrorResult(fmt.Sprintf("failed to deliver attachment: %v", err)).WithError(err)
|
||||
} else {
|
||||
handledAttachments = append(
|
||||
handledAttachments,
|
||||
buildProviderAttachments(al.mediaStore, toolResult.Media)...,
|
||||
)
|
||||
}
|
||||
} else if al.bus != nil {
|
||||
al.bus.PublishOutboundMedia(ctx, outboundMedia)
|
||||
toolResult.ResponseHandled = false
|
||||
}
|
||||
}
|
||||
|
||||
if len(toolResult.Media) > 0 && !toolResult.ResponseHandled {
|
||||
toolResult.ArtifactTags = buildArtifactTags(al.mediaStore, toolResult.Media)
|
||||
}
|
||||
|
||||
if !toolResult.ResponseHandled {
|
||||
exec.allResponsesHandled = false
|
||||
}
|
||||
|
||||
shouldSendForUser := !toolResult.Silent &&
|
||||
toolResult.ForUser != "" &&
|
||||
(ts.opts.SendResponse || toolResult.ResponseHandled)
|
||||
if shouldSendForUser {
|
||||
al.bus.PublishOutbound(ctx, outboundMessageForTurn(ts, toolResult.ForUser))
|
||||
logger.DebugCF("agent", "Sent tool result to user",
|
||||
map[string]any{
|
||||
"tool": toolName,
|
||||
"content_len": len(toolResult.ForUser),
|
||||
})
|
||||
}
|
||||
contentForLLM := toolResult.ContentForLLM()
|
||||
|
||||
if al.cfg.Tools.IsFilterSensitiveDataEnabled() {
|
||||
contentForLLM = al.cfg.FilterSensitiveData(contentForLLM)
|
||||
}
|
||||
|
||||
toolResultMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: contentForLLM,
|
||||
ToolCallID: toolCallID,
|
||||
}
|
||||
if len(toolResult.Media) > 0 && !toolResult.ResponseHandled {
|
||||
toolResultMsg.Media = append(toolResultMsg.Media, toolResult.Media...)
|
||||
}
|
||||
al.emitEvent(
|
||||
EventKindToolExecEnd,
|
||||
ts.eventMeta("runTurn", "turn.tool.end"),
|
||||
ToolExecEndPayload{
|
||||
Tool: toolName,
|
||||
Duration: toolDuration,
|
||||
ForLLMLen: len(contentForLLM),
|
||||
ForUserLen: len(toolResult.ForUser),
|
||||
IsError: toolResult.IsError,
|
||||
Async: toolResult.Async,
|
||||
},
|
||||
)
|
||||
messages = append(messages, toolResultMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, toolResultMsg)
|
||||
ts.recordPersistedMessage(toolResultMsg)
|
||||
ts.ingestMessage(turnCtx, al, toolResultMsg)
|
||||
}
|
||||
|
||||
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 {
|
||||
exec.pendingMessages = append(exec.pendingMessages, steerMsgs...)
|
||||
}
|
||||
|
||||
skipReason := ""
|
||||
skipMessage := ""
|
||||
if len(exec.pendingMessages) > 0 {
|
||||
skipReason = "queued user steering message"
|
||||
skipMessage = "Skipped due to queued user message."
|
||||
} else if gracefulPending, _ := ts.gracefulInterruptRequested(); gracefulPending {
|
||||
skipReason = "graceful interrupt requested"
|
||||
skipMessage = "Skipped due to graceful interrupt."
|
||||
}
|
||||
|
||||
if skipReason != "" {
|
||||
remaining := len(normalizedToolCalls) - i - 1
|
||||
if remaining > 0 {
|
||||
logger.InfoCF("agent", "Turn checkpoint: skipping remaining tools",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"completed": i + 1,
|
||||
"skipped": remaining,
|
||||
"reason": skipReason,
|
||||
})
|
||||
for j := i + 1; j < len(normalizedToolCalls); j++ {
|
||||
skippedTC := normalizedToolCalls[j]
|
||||
al.emitEvent(
|
||||
EventKindToolExecSkipped,
|
||||
ts.eventMeta("runTurn", "turn.tool.skipped"),
|
||||
ToolExecSkippedPayload{
|
||||
Tool: skippedTC.Name,
|
||||
Reason: skipReason,
|
||||
},
|
||||
)
|
||||
skippedMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: skipMessage,
|
||||
ToolCallID: skippedTC.ID,
|
||||
}
|
||||
messages = append(messages, skippedMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, skippedMsg)
|
||||
ts.recordPersistedMessage(skippedMsg)
|
||||
}
|
||||
}
|
||||
}
|
||||
break toolLoop
|
||||
}
|
||||
|
||||
if ts.pendingResults != nil {
|
||||
select {
|
||||
case result, ok := <-ts.pendingResults:
|
||||
if ok && result != nil && result.ForLLM != "" {
|
||||
content := al.cfg.FilterSensitiveData(result.ForLLM)
|
||||
msg := providers.Message{Role: "user", Content: fmt.Sprintf("[SubTurn Result] %s", content)}
|
||||
messages = append(messages, msg)
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, msg)
|
||||
}
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
exec.messages = messages
|
||||
|
||||
// Continue if pending steering exists (regardless of allResponsesHandled).
|
||||
// This covers the case where tools were partially executed and skipped due to steering,
|
||||
// but one tool had ResponseHandled=false (so allResponsesHandled=false).
|
||||
if len(exec.pendingMessages) > 0 {
|
||||
logger.InfoCF("agent", "Pending steering after partial tool execution; continuing turn",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"pending_count": len(exec.pendingMessages),
|
||||
"allResponsesHandled": exec.allResponsesHandled,
|
||||
})
|
||||
exec.allResponsesHandled = false
|
||||
return ToolControlContinue
|
||||
}
|
||||
|
||||
// Poll for newly arrived steering
|
||||
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 {
|
||||
logger.InfoCF("agent", "Steering arrived after tool delivery; continuing turn",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"steering_count": len(steerMsgs),
|
||||
})
|
||||
exec.pendingMessages = append(exec.pendingMessages, steerMsgs...)
|
||||
exec.allResponsesHandled = false
|
||||
return ToolControlContinue
|
||||
}
|
||||
|
||||
// No pending steering: finalize or break depending on allResponsesHandled
|
||||
if exec.allResponsesHandled {
|
||||
summaryMsg := providers.Message{
|
||||
Role: "assistant",
|
||||
Content: handledToolResponseSummary,
|
||||
Attachments: append([]providers.Attachment(nil), handledAttachments...),
|
||||
}
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, summaryMsg)
|
||||
ts.recordPersistedMessage(summaryMsg)
|
||||
ts.ingestMessage(turnCtx, al, summaryMsg)
|
||||
if err := ts.agent.Sessions.Save(ts.sessionKey); err != nil {
|
||||
logger.WarnCF("agent", "Failed to save session after tool delivery",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"error": err.Error(),
|
||||
})
|
||||
}
|
||||
}
|
||||
if ts.opts.EnableSummary {
|
||||
al.contextManager.Compact(turnCtx, &CompactRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Reason: ContextCompressReasonSummarize,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
})
|
||||
}
|
||||
ts.setPhase(TurnPhaseCompleted)
|
||||
ts.setFinalContent("")
|
||||
logger.InfoCF("agent", "Tool output satisfied delivery; ending turn without follow-up LLM",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"tool_count": len(normalizedToolCalls),
|
||||
})
|
||||
return ToolControlBreak
|
||||
}
|
||||
|
||||
// allResponsesHandled=false and no pending steering: continue so coordinator
|
||||
// makes another LLM call. The tool result is in messages and the LLM will
|
||||
// return it as finalContent in the next iteration.
|
||||
ts.agent.Tools.TickTTL()
|
||||
logger.DebugCF("agent", "TTL tick after tool execution", map[string]any{
|
||||
"agent_id": ts.agent.ID, "iteration": iteration,
|
||||
})
|
||||
return ToolControlContinue
|
||||
}
|
||||
77
pkg/agent/pipeline_finalize.go
Normal file
77
pkg/agent/pipeline_finalize.go
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// Finalize handles turn finalization, either:
|
||||
// - Early return when allResponsesHandled=true (ExecuteTools already finalized)
|
||||
// - Normal finalization for allResponsesHandled=false (sets finalContent, saves session, compact)
|
||||
func (p *Pipeline) Finalize(
|
||||
ctx context.Context,
|
||||
turnCtx context.Context,
|
||||
ts *turnState,
|
||||
exec *turnExecution,
|
||||
turnStatus TurnEndStatus,
|
||||
finalContent string,
|
||||
) (turnResult, error) {
|
||||
al := p.al
|
||||
|
||||
// When allResponsesHandled=true, ExecuteTools already finalized
|
||||
// (added handledToolResponseSummary, saved session, set phase to Completed).
|
||||
// But still check for hard abort - if requested, abort the turn.
|
||||
if exec.allResponsesHandled {
|
||||
if ts.hardAbortRequested() {
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
ts.setPhase(TurnPhaseCompleted)
|
||||
return turnResult{
|
||||
finalContent: finalContent,
|
||||
status: turnStatus,
|
||||
followUps: append([]bus.InboundMessage(nil), ts.followUps...),
|
||||
}, nil
|
||||
}
|
||||
|
||||
ts.setPhase(TurnPhaseFinalizing)
|
||||
ts.setFinalContent(finalContent)
|
||||
if !ts.opts.NoHistory {
|
||||
finalMsg := providers.Message{Role: "assistant", Content: finalContent}
|
||||
ts.agent.Sessions.AddMessage(ts.sessionKey, finalMsg.Role, finalMsg.Content)
|
||||
ts.recordPersistedMessage(finalMsg)
|
||||
ts.ingestMessage(turnCtx, al, finalMsg)
|
||||
if err := ts.agent.Sessions.Save(ts.sessionKey); err != nil {
|
||||
al.emitEvent(
|
||||
EventKindError,
|
||||
ts.eventMeta("runTurn", "turn.error"),
|
||||
ErrorPayload{
|
||||
Stage: "session_save",
|
||||
Message: err.Error(),
|
||||
},
|
||||
)
|
||||
return turnResult{status: TurnEndStatusError}, err
|
||||
}
|
||||
}
|
||||
|
||||
if ts.opts.EnableSummary {
|
||||
al.contextManager.Compact(
|
||||
turnCtx,
|
||||
&CompactRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Reason: ContextCompressReasonSummarize,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
ts.setPhase(TurnPhaseCompleted)
|
||||
return turnResult{
|
||||
finalContent: finalContent,
|
||||
status: turnStatus,
|
||||
followUps: append([]bus.InboundMessage(nil), ts.followUps...),
|
||||
}, nil
|
||||
}
|
||||
541
pkg/agent/pipeline_llm.go
Normal file
541
pkg/agent/pipeline_llm.go
Normal file
|
|
@ -0,0 +1,541 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/constants"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// CallLLM performs an LLM call with fallback support, hook invocation, and retry logic.
|
||||
// It handles PreLLM setup, the actual LLM invocation with retry, and AfterLLM processing.
|
||||
// Returns Control indicating what the coordinator should do next.
|
||||
func (p *Pipeline) CallLLM(
|
||||
ctx context.Context,
|
||||
turnCtx context.Context,
|
||||
ts *turnState,
|
||||
exec *turnExecution,
|
||||
iteration int,
|
||||
) (Control, error) {
|
||||
al := p.al
|
||||
maxMediaSize := p.Cfg.Agents.Defaults.GetMaxMediaSize()
|
||||
|
||||
// PreLLM: resolve media refs (except on iteration 1 where user media is already resolved)
|
||||
if iteration > 1 {
|
||||
exec.messages = resolveMediaRefs(exec.messages, p.MediaStore, maxMediaSize)
|
||||
}
|
||||
|
||||
// PreLLM: graceful terminal handling
|
||||
exec.gracefulTerminal, _ = ts.gracefulInterruptRequested()
|
||||
exec.providerToolDefs = ts.agent.Tools.ToProviderDefs()
|
||||
|
||||
// Native web search support
|
||||
webSearchEnabled := al.cfg.Tools.IsToolEnabled("web")
|
||||
exec.useNativeSearch = webSearchEnabled && al.cfg.Tools.Web.PreferNative &&
|
||||
func() bool {
|
||||
if ns, ok := ts.agent.Provider.(providers.NativeSearchCapable); ok {
|
||||
return ns.SupportsNativeSearch()
|
||||
}
|
||||
return false
|
||||
}()
|
||||
|
||||
if exec.useNativeSearch {
|
||||
filtered := make([]providers.ToolDefinition, 0, len(exec.providerToolDefs))
|
||||
for _, td := range exec.providerToolDefs {
|
||||
if td.Function.Name != "web_search" {
|
||||
filtered = append(filtered, td)
|
||||
}
|
||||
}
|
||||
exec.providerToolDefs = filtered
|
||||
}
|
||||
|
||||
exec.callMessages = exec.messages
|
||||
if exec.gracefulTerminal {
|
||||
exec.callMessages = append(append([]providers.Message(nil), exec.messages...), ts.interruptHintMessage())
|
||||
exec.providerToolDefs = nil
|
||||
ts.markGracefulTerminalUsed()
|
||||
}
|
||||
|
||||
exec.llmOpts = map[string]any{
|
||||
"max_tokens": ts.agent.MaxTokens,
|
||||
"temperature": ts.agent.Temperature,
|
||||
"prompt_cache_key": ts.agent.ID,
|
||||
}
|
||||
if exec.useNativeSearch {
|
||||
exec.llmOpts["native_search"] = true
|
||||
}
|
||||
if ts.agent.ThinkingLevel != ThinkingOff {
|
||||
if tc, ok := ts.agent.Provider.(providers.ThinkingCapable); ok && tc.SupportsThinking() {
|
||||
exec.llmOpts["thinking_level"] = string(ts.agent.ThinkingLevel)
|
||||
} else {
|
||||
logger.WarnCF("agent", "thinking_level is set but current provider does not support it, ignoring",
|
||||
map[string]any{"agent_id": ts.agent.ID, "thinking_level": string(ts.agent.ThinkingLevel)})
|
||||
}
|
||||
}
|
||||
|
||||
exec.llmModel = exec.activeModel
|
||||
|
||||
// BeforeLLM hook
|
||||
if p.Hooks != nil {
|
||||
llmReq, decision := p.Hooks.BeforeLLM(turnCtx, &LLMHookRequest{
|
||||
Meta: ts.eventMeta("runTurn", "turn.llm.request"),
|
||||
Context: cloneTurnContext(ts.turnCtx),
|
||||
Model: exec.llmModel,
|
||||
Messages: exec.callMessages,
|
||||
Tools: exec.providerToolDefs,
|
||||
Options: exec.llmOpts,
|
||||
GracefulTerminal: exec.gracefulTerminal,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if llmReq != nil {
|
||||
exec.llmModel = llmReq.Model
|
||||
exec.callMessages = llmReq.Messages
|
||||
exec.providerToolDefs = llmReq.Tools
|
||||
exec.llmOpts = llmReq.Options
|
||||
}
|
||||
case HookActionAbortTurn:
|
||||
exec.abortedByHook = true
|
||||
return ControlBreak, nil
|
||||
case HookActionHardAbort:
|
||||
_ = ts.requestHardAbort()
|
||||
exec.abortedByHardAbort = true
|
||||
return ControlBreak, nil
|
||||
}
|
||||
}
|
||||
|
||||
al.emitEvent(
|
||||
EventKindLLMRequest,
|
||||
ts.eventMeta("runTurn", "turn.llm.request"),
|
||||
LLMRequestPayload{
|
||||
Model: exec.llmModel,
|
||||
MessagesCount: len(exec.callMessages),
|
||||
ToolsCount: len(exec.providerToolDefs),
|
||||
MaxTokens: ts.agent.MaxTokens,
|
||||
Temperature: ts.agent.Temperature,
|
||||
},
|
||||
)
|
||||
|
||||
logger.DebugCF("agent", "LLM request",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"model": exec.llmModel,
|
||||
"messages_count": len(exec.callMessages),
|
||||
"tools_count": len(exec.providerToolDefs),
|
||||
"max_tokens": ts.agent.MaxTokens,
|
||||
"temperature": ts.agent.Temperature,
|
||||
"system_prompt_len": len(exec.callMessages[0].Content),
|
||||
})
|
||||
logger.DebugCF("agent", "Full LLM request",
|
||||
map[string]any{
|
||||
"iteration": iteration,
|
||||
"messages_json": formatMessagesForLog(exec.callMessages),
|
||||
"tools_json": formatToolsForLog(exec.providerToolDefs),
|
||||
})
|
||||
|
||||
// LLM call closure with fallback support
|
||||
callLLM := func(messagesForCall []providers.Message, toolDefsForCall []providers.ToolDefinition) (*providers.LLMResponse, error) {
|
||||
providerCtx, providerCancel := context.WithCancel(turnCtx)
|
||||
ts.setProviderCancel(providerCancel)
|
||||
defer func() {
|
||||
providerCancel()
|
||||
ts.clearProviderCancel(providerCancel)
|
||||
}()
|
||||
|
||||
al.activeRequests.Add(1)
|
||||
defer al.activeRequests.Done()
|
||||
|
||||
if len(exec.activeCandidates) > 1 && p.Fallback != nil {
|
||||
fbResult, fbErr := p.Fallback.Execute(
|
||||
providerCtx,
|
||||
exec.activeCandidates,
|
||||
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
|
||||
candidateProvider := exec.activeProvider
|
||||
if cp, ok := ts.agent.CandidateProviders[providers.ModelKey(provider, model)]; ok {
|
||||
candidateProvider = cp
|
||||
}
|
||||
return candidateProvider.Chat(ctx, messagesForCall, toolDefsForCall, model, exec.llmOpts)
|
||||
},
|
||||
)
|
||||
if fbErr != nil {
|
||||
return nil, fbErr
|
||||
}
|
||||
if fbResult.Provider != "" && len(fbResult.Attempts) > 0 {
|
||||
logger.InfoCF(
|
||||
"agent",
|
||||
fmt.Sprintf("Fallback: succeeded with %s/%s after %d attempts",
|
||||
fbResult.Provider, fbResult.Model, len(fbResult.Attempts)+1),
|
||||
map[string]any{"agent_id": ts.agent.ID, "iteration": iteration},
|
||||
)
|
||||
}
|
||||
return fbResult.Response, nil
|
||||
}
|
||||
return exec.activeProvider.Chat(providerCtx, messagesForCall, toolDefsForCall, exec.llmModel, exec.llmOpts)
|
||||
}
|
||||
|
||||
// Retry loop
|
||||
var err error
|
||||
maxRetries := 2
|
||||
for retry := 0; retry <= maxRetries; retry++ {
|
||||
exec.response, err = callLLM(exec.callMessages, exec.providerToolDefs)
|
||||
if err == nil {
|
||||
break
|
||||
}
|
||||
if ts.hardAbortRequested() && errors.Is(err, context.Canceled) {
|
||||
_ = ts.requestHardAbort()
|
||||
exec.abortedByHardAbort = true
|
||||
return ControlBreak, nil
|
||||
}
|
||||
|
||||
// Retry without media if vision is unsupported
|
||||
if hasMediaRefs(exec.callMessages) && isVisionUnsupportedError(err) && retry < maxRetries {
|
||||
al.emitEvent(
|
||||
EventKindLLMRetry,
|
||||
ts.eventMeta("runTurn", "turn.llm.retry"),
|
||||
LLMRetryPayload{
|
||||
Attempt: retry + 1,
|
||||
MaxRetries: maxRetries,
|
||||
Reason: "vision_unsupported",
|
||||
Error: err.Error(),
|
||||
Backoff: 0,
|
||||
},
|
||||
)
|
||||
logger.WarnCF("agent", "Vision unsupported, retrying without media", map[string]any{
|
||||
"error": err.Error(),
|
||||
"retry": retry,
|
||||
})
|
||||
exec.callMessages = stripMessageMedia(exec.callMessages)
|
||||
if !ts.opts.NoHistory {
|
||||
exec.history = stripMessageMedia(exec.history)
|
||||
ts.agent.Sessions.SetHistory(ts.sessionKey, exec.history)
|
||||
for i := range ts.persistedMessages {
|
||||
ts.persistedMessages[i].Media = nil
|
||||
}
|
||||
ts.refreshRestorePointFromSession(ts.agent)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
errMsg := strings.ToLower(err.Error())
|
||||
isTimeoutError := errors.Is(err, context.DeadlineExceeded) ||
|
||||
strings.Contains(errMsg, "deadline exceeded") ||
|
||||
strings.Contains(errMsg, "client.timeout") ||
|
||||
strings.Contains(errMsg, "timed out") ||
|
||||
strings.Contains(errMsg, "timeout exceeded")
|
||||
|
||||
isContextError := !isTimeoutError && (strings.Contains(errMsg, "context_length_exceeded") ||
|
||||
strings.Contains(errMsg, "context window") ||
|
||||
strings.Contains(errMsg, "context_window") ||
|
||||
strings.Contains(errMsg, "maximum context length") ||
|
||||
strings.Contains(errMsg, "token limit") ||
|
||||
strings.Contains(errMsg, "too many tokens") ||
|
||||
strings.Contains(errMsg, "max_tokens") ||
|
||||
strings.Contains(errMsg, "invalidparameter") ||
|
||||
strings.Contains(errMsg, "prompt is too long") ||
|
||||
strings.Contains(errMsg, "request too large"))
|
||||
|
||||
if isTimeoutError && retry < maxRetries {
|
||||
backoff := time.Duration(retry+1) * 5 * time.Second
|
||||
al.emitEvent(
|
||||
EventKindLLMRetry,
|
||||
ts.eventMeta("runTurn", "turn.llm.retry"),
|
||||
LLMRetryPayload{
|
||||
Attempt: retry + 1,
|
||||
MaxRetries: maxRetries,
|
||||
Reason: "timeout",
|
||||
Error: err.Error(),
|
||||
Backoff: backoff,
|
||||
},
|
||||
)
|
||||
logger.WarnCF("agent", "Timeout error, retrying after backoff", map[string]any{
|
||||
"error": err.Error(),
|
||||
"retry": retry,
|
||||
"backoff": backoff.String(),
|
||||
})
|
||||
if sleepErr := sleepWithContext(turnCtx, backoff); sleepErr != nil {
|
||||
if ts.hardAbortRequested() {
|
||||
_ = ts.requestHardAbort()
|
||||
return ControlBreak, nil
|
||||
}
|
||||
err = sleepErr
|
||||
break
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if isContextError && retry < maxRetries && !ts.opts.NoHistory {
|
||||
al.emitEvent(
|
||||
EventKindLLMRetry,
|
||||
ts.eventMeta("runTurn", "turn.llm.retry"),
|
||||
LLMRetryPayload{
|
||||
Attempt: retry + 1,
|
||||
MaxRetries: maxRetries,
|
||||
Reason: "context_limit",
|
||||
Error: err.Error(),
|
||||
},
|
||||
)
|
||||
logger.WarnCF(
|
||||
"agent",
|
||||
"Context window error detected, attempting compression",
|
||||
map[string]any{
|
||||
"error": err.Error(),
|
||||
"retry": retry,
|
||||
},
|
||||
)
|
||||
|
||||
if retry == 0 && !constants.IsInternalChannel(ts.channel) {
|
||||
al.bus.PublishOutbound(ctx, outboundMessageForTurn(
|
||||
ts,
|
||||
"Context window exceeded. Compressing history and retrying...",
|
||||
))
|
||||
}
|
||||
|
||||
if compactErr := p.ContextManager.Compact(ctx, &CompactRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Reason: ContextCompressReasonRetry,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
}); compactErr != nil {
|
||||
logger.WarnCF("agent", "Context overflow compact failed", map[string]any{
|
||||
"session_key": ts.sessionKey,
|
||||
"error": compactErr.Error(),
|
||||
})
|
||||
}
|
||||
ts.refreshRestorePointFromSession(ts.agent)
|
||||
if asmResp, asmErr := p.ContextManager.Assemble(ctx, &AssembleRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
MaxTokens: ts.agent.MaxTokens,
|
||||
}); asmErr == nil && asmResp != nil {
|
||||
exec.history = asmResp.History
|
||||
exec.summary = asmResp.Summary
|
||||
}
|
||||
exec.messages = ts.agent.ContextBuilder.BuildMessages(
|
||||
exec.history, exec.summary, "",
|
||||
nil, ts.channel, ts.chatID, ts.opts.Dispatch.SenderID(), ts.opts.SenderDisplayName,
|
||||
activeSkillNames(ts.agent, ts.opts)...,
|
||||
)
|
||||
exec.callMessages = exec.messages
|
||||
if exec.gracefulTerminal {
|
||||
msgs := append([]providers.Message(nil), exec.messages...)
|
||||
exec.callMessages = append(msgs, ts.interruptHintMessage())
|
||||
}
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
al.emitEvent(
|
||||
EventKindError,
|
||||
ts.eventMeta("runTurn", "turn.error"),
|
||||
ErrorPayload{
|
||||
Stage: "llm",
|
||||
Message: err.Error(),
|
||||
},
|
||||
)
|
||||
logger.ErrorCF("agent", "LLM call failed",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"model": exec.llmModel,
|
||||
"error": err.Error(),
|
||||
})
|
||||
return ControlBreak, fmt.Errorf("LLM call failed after retries: %w", err)
|
||||
}
|
||||
|
||||
// AfterLLM hook
|
||||
if p.Hooks != nil {
|
||||
llmResp, decision := p.Hooks.AfterLLM(turnCtx, &LLMHookResponse{
|
||||
Meta: ts.eventMeta("runTurn", "turn.llm.response"),
|
||||
Context: cloneTurnContext(ts.turnCtx),
|
||||
Model: exec.llmModel,
|
||||
Response: exec.response,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if llmResp != nil && llmResp.Response != nil {
|
||||
exec.response = llmResp.Response
|
||||
}
|
||||
case HookActionAbortTurn:
|
||||
exec.abortedByHook = true
|
||||
return ControlBreak, nil
|
||||
case HookActionHardAbort:
|
||||
_ = ts.requestHardAbort()
|
||||
exec.abortedByHardAbort = true
|
||||
return ControlBreak, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Save finishReason to turnState for SubTurn truncation detection
|
||||
if innerTS := turnStateFromContext(ctx); innerTS != nil {
|
||||
innerTS.SetLastFinishReason(exec.response.FinishReason)
|
||||
if exec.response.Usage != nil {
|
||||
innerTS.SetLastUsage(exec.response.Usage)
|
||||
}
|
||||
}
|
||||
|
||||
reasoningContent := exec.response.Reasoning
|
||||
if reasoningContent == "" {
|
||||
reasoningContent = exec.response.ReasoningContent
|
||||
}
|
||||
if ts.channel == "pico" {
|
||||
go al.publishPicoReasoning(turnCtx, reasoningContent, ts.chatID)
|
||||
} else {
|
||||
go al.handleReasoning(
|
||||
turnCtx,
|
||||
reasoningContent,
|
||||
ts.channel,
|
||||
al.targetReasoningChannelID(ts.channel),
|
||||
)
|
||||
}
|
||||
al.emitEvent(
|
||||
EventKindLLMResponse,
|
||||
ts.eventMeta("runTurn", "turn.llm.response"),
|
||||
LLMResponsePayload{
|
||||
ContentLen: len(exec.response.Content),
|
||||
ToolCalls: len(exec.response.ToolCalls),
|
||||
HasReasoning: exec.response.Reasoning != "" || exec.response.ReasoningContent != "",
|
||||
},
|
||||
)
|
||||
|
||||
llmResponseFields := map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"content_chars": len(exec.response.Content),
|
||||
"tool_calls": len(exec.response.ToolCalls),
|
||||
"reasoning": exec.response.Reasoning,
|
||||
"target_channel": al.targetReasoningChannelID(ts.channel),
|
||||
"channel": ts.channel,
|
||||
}
|
||||
if exec.response.Usage != nil {
|
||||
llmResponseFields["prompt_tokens"] = exec.response.Usage.PromptTokens
|
||||
llmResponseFields["completion_tokens"] = exec.response.Usage.CompletionTokens
|
||||
llmResponseFields["total_tokens"] = exec.response.Usage.TotalTokens
|
||||
}
|
||||
logger.DebugCF("agent", "LLM response", llmResponseFields)
|
||||
|
||||
if al.bus != nil &&
|
||||
ts.channel == "pico" &&
|
||||
len(exec.response.ToolCalls) > 0 &&
|
||||
ts.opts.AllowInterimPicoPublish &&
|
||||
!shouldPublishToolFeedback(al.cfg, ts) {
|
||||
if strings.TrimSpace(exec.response.Content) != "" {
|
||||
outCtx, outCancel := context.WithTimeout(turnCtx, 3*time.Second)
|
||||
publishErr := al.bus.PublishOutbound(outCtx, bus.OutboundMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Content: exec.response.Content,
|
||||
})
|
||||
outCancel()
|
||||
if publishErr != nil {
|
||||
logger.WarnCF("agent", "Failed to publish pico interim tool-call content", map[string]any{
|
||||
"error": publishErr.Error(),
|
||||
"channel": ts.channel,
|
||||
"chat_id": ts.chatID,
|
||||
"iteration": iteration,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// No-tool-call path: steering check and direct response
|
||||
if len(exec.response.ToolCalls) == 0 || exec.gracefulTerminal {
|
||||
responseContent := exec.response.Content
|
||||
if responseContent == "" && exec.response.ReasoningContent != "" && ts.channel != "pico" {
|
||||
responseContent = exec.response.ReasoningContent
|
||||
}
|
||||
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 {
|
||||
logger.InfoCF("agent", "Steering arrived after direct LLM response; continuing turn",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"steering_count": len(steerMsgs),
|
||||
})
|
||||
exec.pendingMessages = append(exec.pendingMessages, steerMsgs...)
|
||||
return ControlContinue, nil
|
||||
}
|
||||
exec.finalContent = responseContent
|
||||
logger.InfoCF("agent", "LLM response without tool calls (direct answer)",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"content_chars": len(exec.finalContent),
|
||||
})
|
||||
return ControlBreak, nil
|
||||
}
|
||||
|
||||
// Tool-call path: normalize and prepare for tool execution
|
||||
exec.normalizedToolCalls = make([]providers.ToolCall, 0, len(exec.response.ToolCalls))
|
||||
for _, tc := range exec.response.ToolCalls {
|
||||
exec.normalizedToolCalls = append(exec.normalizedToolCalls, providers.NormalizeToolCall(tc))
|
||||
}
|
||||
|
||||
toolNames := make([]string, 0, len(exec.normalizedToolCalls))
|
||||
for _, tc := range exec.normalizedToolCalls {
|
||||
toolNames = append(toolNames, tc.Name)
|
||||
}
|
||||
logger.InfoCF("agent", "LLM requested tool calls",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tools": toolNames,
|
||||
"count": len(exec.normalizedToolCalls),
|
||||
"iteration": iteration,
|
||||
})
|
||||
|
||||
exec.allResponsesHandled = len(exec.normalizedToolCalls) > 0
|
||||
assistantMsg := providers.Message{
|
||||
Role: "assistant",
|
||||
Content: exec.response.Content,
|
||||
ReasoningContent: exec.response.ReasoningContent,
|
||||
}
|
||||
for _, tc := range exec.normalizedToolCalls {
|
||||
argumentsJSON, _ := json.Marshal(tc.Arguments)
|
||||
toolFeedbackExplanation := toolFeedbackExplanationForToolCall(
|
||||
exec.response,
|
||||
tc,
|
||||
exec.messages,
|
||||
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
||||
)
|
||||
extraContent := tc.ExtraContent
|
||||
if strings.TrimSpace(toolFeedbackExplanation) != "" {
|
||||
if extraContent == nil {
|
||||
extraContent = &providers.ExtraContent{}
|
||||
}
|
||||
extraContent.ToolFeedbackExplanation = toolFeedbackExplanation
|
||||
}
|
||||
thoughtSignature := ""
|
||||
if tc.Function != nil {
|
||||
thoughtSignature = tc.Function.ThoughtSignature
|
||||
}
|
||||
assistantMsg.ToolCalls = append(assistantMsg.ToolCalls, providers.ToolCall{
|
||||
ID: tc.ID,
|
||||
Type: "function",
|
||||
Name: tc.Name,
|
||||
Function: &providers.FunctionCall{
|
||||
Name: tc.Name,
|
||||
Arguments: string(argumentsJSON),
|
||||
ThoughtSignature: thoughtSignature,
|
||||
},
|
||||
ExtraContent: extraContent,
|
||||
ThoughtSignature: thoughtSignature,
|
||||
})
|
||||
}
|
||||
exec.messages = append(exec.messages, assistantMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, assistantMsg)
|
||||
ts.recordPersistedMessage(assistantMsg)
|
||||
ts.ingestMessage(turnCtx, al, assistantMsg)
|
||||
}
|
||||
|
||||
return ControlToolLoop, nil
|
||||
}
|
||||
116
pkg/agent/pipeline_setup.go
Normal file
116
pkg/agent/pipeline_setup.go
Normal file
|
|
@ -0,0 +1,116 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// SetupTurn extracts the one-time initialization phase, returning a
|
||||
// turnExecution populated with history, messages, and candidate selection.
|
||||
// It replaces lines 56-145 of the original runTurn.
|
||||
func (p *Pipeline) SetupTurn(ctx context.Context, ts *turnState) (*turnExecution, error) {
|
||||
cfg := p.Cfg
|
||||
maxMediaSize := cfg.Agents.Defaults.GetMaxMediaSize()
|
||||
|
||||
var history []providers.Message
|
||||
var summary string
|
||||
if !ts.opts.NoHistory {
|
||||
if resp, err := p.ContextManager.Assemble(ctx, &AssembleRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
MaxTokens: ts.agent.MaxTokens,
|
||||
}); err == nil && resp != nil {
|
||||
history = resp.History
|
||||
summary = resp.Summary
|
||||
}
|
||||
}
|
||||
ts.captureRestorePoint(history, summary)
|
||||
|
||||
messages := ts.agent.ContextBuilder.BuildMessages(
|
||||
history,
|
||||
summary,
|
||||
ts.userMessage,
|
||||
ts.media,
|
||||
ts.channel,
|
||||
ts.chatID,
|
||||
ts.opts.Dispatch.SenderID(),
|
||||
ts.opts.SenderDisplayName,
|
||||
activeSkillNames(ts.agent, ts.opts)...,
|
||||
)
|
||||
|
||||
messages = resolveMediaRefs(messages, p.MediaStore, maxMediaSize)
|
||||
|
||||
if !ts.opts.NoHistory {
|
||||
toolDefs := ts.agent.Tools.ToProviderDefs()
|
||||
if isOverContextBudget(ts.agent.ContextWindow, messages, toolDefs, ts.agent.MaxTokens) {
|
||||
logger.WarnCF("agent", "Proactive compression: context budget exceeded before LLM call",
|
||||
map[string]any{"session_key": ts.sessionKey})
|
||||
if err := p.ContextManager.Compact(ctx, &CompactRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Reason: ContextCompressReasonProactive,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
}); err != nil {
|
||||
logger.WarnCF("agent", "Proactive compact failed", map[string]any{
|
||||
"session_key": ts.sessionKey,
|
||||
"error": err.Error(),
|
||||
})
|
||||
}
|
||||
ts.refreshRestorePointFromSession(ts.agent)
|
||||
if resp, err := p.ContextManager.Assemble(ctx, &AssembleRequest{
|
||||
SessionKey: ts.sessionKey,
|
||||
Budget: ts.agent.ContextWindow,
|
||||
MaxTokens: ts.agent.MaxTokens,
|
||||
}); err == nil && resp != nil {
|
||||
history = resp.History
|
||||
summary = resp.Summary
|
||||
}
|
||||
messages = ts.agent.ContextBuilder.BuildMessages(
|
||||
history, summary, ts.userMessage,
|
||||
ts.media, ts.channel, ts.chatID,
|
||||
ts.opts.Dispatch.SenderID(), ts.opts.SenderDisplayName,
|
||||
activeSkillNames(ts.agent, ts.opts)...,
|
||||
)
|
||||
messages = resolveMediaRefs(messages, p.MediaStore, maxMediaSize)
|
||||
}
|
||||
}
|
||||
|
||||
if !ts.opts.NoHistory && (strings.TrimSpace(ts.userMessage) != "" || len(ts.media) > 0) {
|
||||
rootMsg := providers.Message{
|
||||
Role: "user",
|
||||
Content: ts.userMessage,
|
||||
Media: append([]string(nil), ts.media...),
|
||||
}
|
||||
if len(rootMsg.Media) > 0 {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, rootMsg)
|
||||
} else {
|
||||
ts.agent.Sessions.AddMessage(ts.sessionKey, rootMsg.Role, rootMsg.Content)
|
||||
}
|
||||
ts.recordPersistedMessage(rootMsg)
|
||||
ts.ingestMessage(ctx, p.al, rootMsg)
|
||||
}
|
||||
|
||||
activeCandidates, activeModel, usedLight := p.al.selectCandidates(ts.agent, ts.userMessage, messages)
|
||||
activeProvider := ts.agent.Provider
|
||||
if usedLight && ts.agent.LightProvider != nil {
|
||||
activeProvider = ts.agent.LightProvider
|
||||
}
|
||||
|
||||
exec := newTurnExecution(
|
||||
ts.agent,
|
||||
ts.opts,
|
||||
history,
|
||||
summary,
|
||||
messages,
|
||||
)
|
||||
exec.activeCandidates = activeCandidates
|
||||
exec.activeModel = activeModel
|
||||
exec.activeProvider = activeProvider
|
||||
exec.usedLight = usedLight
|
||||
|
||||
return exec, nil
|
||||
}
|
||||
|
|
@ -462,7 +462,8 @@ func spawnSubTurn(
|
|||
}()
|
||||
|
||||
// 8. Execute sub-turn via the real agent loop.
|
||||
turnRes, turnErr := al.runTurn(childCtx, childTS)
|
||||
pipeline := NewPipeline(al)
|
||||
turnRes, turnErr := al.runTurn(childCtx, childTS, pipeline)
|
||||
|
||||
// Release the concurrency semaphore immediately after runTurn completes,
|
||||
// before the cleanup defer runs. This prevents a deadlock where:
|
||||
|
|
|
|||
|
|
@ -1650,6 +1650,38 @@ func TestGrandchildAbort_CascadingCancellation(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestNestedSubTurn_GracefulFinishSignalsDirectChildren(t *testing.T) {
|
||||
parentCtx := context.Background()
|
||||
parentTS := &turnState{
|
||||
ctx: parentCtx,
|
||||
turnID: "parent-graceful",
|
||||
depth: 1,
|
||||
pendingResults: make(chan *tools.ToolResult, 16),
|
||||
}
|
||||
parentTS.ctx, parentTS.cancelFunc = context.WithCancel(parentCtx)
|
||||
|
||||
childTS := &turnState{
|
||||
ctx: context.Background(),
|
||||
turnID: "child-graceful",
|
||||
depth: 2,
|
||||
parentTurnState: parentTS,
|
||||
pendingResults: make(chan *tools.ToolResult, 16),
|
||||
}
|
||||
|
||||
if childTS.IsParentEnded() {
|
||||
t.Fatal("IsParentEnded should be false before parent finishes")
|
||||
}
|
||||
|
||||
parentTS.Finish(false)
|
||||
|
||||
if !parentTS.parentEnded.Load() {
|
||||
t.Fatal("parentEnded should be true after graceful finish")
|
||||
}
|
||||
if !childTS.IsParentEnded() {
|
||||
t.Fatal("nested child should observe parent graceful finish")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSpawnDuringAbort_RaceCondition verifies behavior when trying to spawn
|
||||
// a sub-turn while the parent is being aborted.
|
||||
func TestSpawnDuringAbort_RaceCondition(t *testing.T) {
|
||||
|
|
|
|||
624
pkg/agent/turn_coord.go
Normal file
624
pkg/agent/turn_coord.go
Normal file
|
|
@ -0,0 +1,624 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState, pipeline *Pipeline) (turnResult, error) {
|
||||
turnCtx, turnCancel := context.WithCancel(ctx)
|
||||
defer turnCancel()
|
||||
ts.setTurnCancel(turnCancel)
|
||||
|
||||
// Inject turnState and AgentLoop into context so tools (e.g. spawn) can retrieve them.
|
||||
turnCtx = withTurnState(turnCtx, ts)
|
||||
turnCtx = WithAgentLoop(turnCtx, al)
|
||||
|
||||
al.registerActiveTurn(ts)
|
||||
defer al.clearActiveTurn(ts)
|
||||
|
||||
turnStatus := TurnEndStatusCompleted
|
||||
defer func() {
|
||||
al.emitEvent(
|
||||
EventKindTurnEnd,
|
||||
ts.eventMeta("runTurn", "turn.end"),
|
||||
TurnEndPayload{
|
||||
Status: turnStatus,
|
||||
Iterations: ts.currentIteration(),
|
||||
Duration: time.Since(ts.startedAt),
|
||||
FinalContentLen: ts.finalContentLen(),
|
||||
},
|
||||
)
|
||||
}()
|
||||
|
||||
al.emitEvent(
|
||||
EventKindTurnStart,
|
||||
ts.eventMeta("runTurn", "turn.start"),
|
||||
TurnStartPayload{
|
||||
UserMessage: ts.userMessage,
|
||||
MediaCount: len(ts.media),
|
||||
},
|
||||
)
|
||||
|
||||
// SetupTurn extracts the one-time initialization phase.
|
||||
exec, err := pipeline.SetupTurn(turnCtx, ts)
|
||||
if err != nil {
|
||||
return turnResult{}, err
|
||||
}
|
||||
|
||||
// Convenience references to exec fields used throughout the turn loop.
|
||||
messages := exec.messages
|
||||
pendingMessages := exec.pendingMessages
|
||||
maxMediaSize := pipeline.Cfg.Agents.Defaults.GetMaxMediaSize()
|
||||
finalContent := exec.finalContent
|
||||
|
||||
for ts.currentIteration() < ts.agent.MaxIterations || len(exec.pendingMessages) > 0 || func() bool {
|
||||
graceful, _ := ts.gracefulInterruptRequested()
|
||||
return graceful
|
||||
}() {
|
||||
if ts.hardAbortRequested() {
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
|
||||
iteration := ts.currentIteration() + 1
|
||||
ts.setIteration(iteration)
|
||||
ts.setPhase(TurnPhaseRunning)
|
||||
|
||||
if iteration > 1 {
|
||||
// For subsequent iterations, read from exec.pendingMessages which
|
||||
// is where ExecuteTools (or initial poll) deposits steering.
|
||||
// We do NOT call dequeueSteeringMessagesForScope here because
|
||||
// steering was already consumed from al.steering by ExecuteTools.
|
||||
if len(exec.pendingMessages) > 0 {
|
||||
pendingMessages = append(pendingMessages, exec.pendingMessages...)
|
||||
exec.pendingMessages = nil
|
||||
}
|
||||
} else if !ts.opts.SkipInitialSteeringPoll {
|
||||
if steerMsgs := al.dequeueSteeringMessagesForScopeWithFallback(ts.sessionKey); len(steerMsgs) > 0 {
|
||||
pendingMessages = append(pendingMessages, steerMsgs...)
|
||||
}
|
||||
}
|
||||
|
||||
// Check if parent turn has ended (SubTurn support from HEAD)
|
||||
if ts.parentTurnState != nil && ts.IsParentEnded() {
|
||||
if !ts.critical {
|
||||
logger.InfoCF("agent", "Parent turn ended, non-critical SubTurn exiting gracefully", map[string]any{
|
||||
"agent_id": ts.agentID,
|
||||
"iteration": iteration,
|
||||
"turn_id": ts.turnID,
|
||||
})
|
||||
break
|
||||
}
|
||||
logger.InfoCF("agent", "Parent turn ended, critical SubTurn continues running", map[string]any{
|
||||
"agent_id": ts.agentID,
|
||||
"iteration": iteration,
|
||||
"turn_id": ts.turnID,
|
||||
})
|
||||
}
|
||||
|
||||
// Poll for pending SubTurn results (from HEAD)
|
||||
if ts.pendingResults != nil {
|
||||
select {
|
||||
case result, ok := <-ts.pendingResults:
|
||||
if ok && result != nil && result.ForLLM != "" {
|
||||
content := al.cfg.FilterSensitiveData(result.ForLLM)
|
||||
msg := providers.Message{Role: "user", Content: fmt.Sprintf("[SubTurn Result] %s", content)}
|
||||
pendingMessages = append(pendingMessages, msg)
|
||||
}
|
||||
default:
|
||||
// No results available
|
||||
}
|
||||
}
|
||||
|
||||
// Inject pending steering messages
|
||||
if len(pendingMessages) > 0 {
|
||||
resolvedPending := resolveMediaRefs(pendingMessages, al.mediaStore, maxMediaSize)
|
||||
totalContentLen := 0
|
||||
for i, pm := range pendingMessages {
|
||||
messages = append(messages, resolvedPending[i])
|
||||
totalContentLen += len(pm.Content)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, pm)
|
||||
ts.recordPersistedMessage(pm)
|
||||
ts.ingestMessage(turnCtx, al, pm)
|
||||
}
|
||||
logger.InfoCF("agent", "Injected steering message into context",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"content_len": len(pm.Content),
|
||||
"media_count": len(pm.Media),
|
||||
})
|
||||
}
|
||||
al.emitEvent(
|
||||
EventKindSteeringInjected,
|
||||
ts.eventMeta("runTurn", "turn.steering.injected"),
|
||||
SteeringInjectedPayload{
|
||||
Count: len(pendingMessages),
|
||||
TotalContentLen: totalContentLen,
|
||||
},
|
||||
)
|
||||
// Clear exec.pendingMessages after injection so InitialSteeringMessages
|
||||
// are not re-injected on subsequent iterations (Issue 2 fix).
|
||||
exec.pendingMessages = nil
|
||||
}
|
||||
// Always sync messages into exec.messages so CallLLM sees the updated state
|
||||
exec.messages = messages
|
||||
|
||||
logger.DebugCF("agent", "LLM iteration",
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"iteration": iteration,
|
||||
"max": ts.agent.MaxIterations,
|
||||
})
|
||||
|
||||
// Execute LLM call via Pipeline
|
||||
ts.setPhase(TurnPhaseRunning)
|
||||
ctrl, callErr := pipeline.CallLLM(ctx, turnCtx, ts, exec, iteration)
|
||||
if callErr != nil {
|
||||
turnStatus = TurnEndStatusError
|
||||
return turnResult{}, callErr
|
||||
}
|
||||
messages = exec.messages
|
||||
pendingMessages = exec.pendingMessages
|
||||
finalContent = exec.finalContent
|
||||
|
||||
switch ctrl {
|
||||
case ControlContinue:
|
||||
continue
|
||||
case ControlBreak:
|
||||
// Hard abort: delegate to abortTurn (sets TurnEndStatusAborted)
|
||||
if exec.abortedByHardAbort {
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
// Hook abort (HookActionAbortTurn): sets TurnEndStatusError, returns error
|
||||
if exec.abortedByHook {
|
||||
turnStatus = TurnEndStatusError
|
||||
return turnResult{}, fmt.Errorf("hook requested turn abort")
|
||||
}
|
||||
// Ensure empty response falls back to DefaultResponse
|
||||
if finalContent == "" {
|
||||
finalContent = ts.opts.DefaultResponse
|
||||
}
|
||||
return pipeline.Finalize(ctx, turnCtx, ts, exec, turnStatus, finalContent)
|
||||
case ControlToolLoop:
|
||||
// Execute tools via Pipeline
|
||||
toolCtrl := pipeline.ExecuteTools(ctx, turnCtx, ts, exec, iteration)
|
||||
switch toolCtrl {
|
||||
case ToolControlContinue:
|
||||
// Re-read exec.messages since ExecuteTools may have updated it
|
||||
// (added tool results/skipped messages) before returning ControlContinue
|
||||
messages = exec.messages
|
||||
continue
|
||||
case ToolControlBreak:
|
||||
// Hard abort: delegate to abortTurn (sets TurnEndStatusAborted)
|
||||
if exec.abortedByHardAbort {
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
// Hook abort (HookActionAbortTurn): sets TurnEndStatusError, returns error
|
||||
if exec.abortedByHook {
|
||||
turnStatus = TurnEndStatusError
|
||||
return turnResult{}, fmt.Errorf("hook requested turn abort")
|
||||
}
|
||||
// ExecuteTools returned ControlBreak:
|
||||
// - allResponsesHandled=true: finalize without DefaultResponse (exec.finalContent empty)
|
||||
// - allResponsesHandled=false: coordinator applies DefaultResponse before finalize
|
||||
if exec.allResponsesHandled {
|
||||
finalContent = ""
|
||||
}
|
||||
return pipeline.Finalize(ctx, turnCtx, ts, exec, turnStatus, finalContent)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if ts.hardAbortRequested() {
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
|
||||
if finalContent == "" {
|
||||
if ts.currentIteration() >= ts.agent.MaxIterations && ts.agent.MaxIterations > 0 {
|
||||
finalContent = toolLimitResponse
|
||||
} else {
|
||||
finalContent = ts.opts.DefaultResponse
|
||||
}
|
||||
}
|
||||
|
||||
// Check hard abort before finalizing (may have been set during tool execution)
|
||||
if ts.hardAbortRequested() {
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
|
||||
return pipeline.Finalize(ctx, turnCtx, ts, exec, turnStatus, finalContent)
|
||||
}
|
||||
|
||||
func (al *AgentLoop) abortTurn(ts *turnState) (turnResult, error) {
|
||||
ts.setPhase(TurnPhaseAborted)
|
||||
if !ts.opts.NoHistory {
|
||||
if err := ts.restoreSession(ts.agent); err != nil {
|
||||
al.emitEvent(
|
||||
EventKindError,
|
||||
ts.eventMeta("abortTurn", "turn.error"),
|
||||
ErrorPayload{
|
||||
Stage: "session_restore",
|
||||
Message: err.Error(),
|
||||
},
|
||||
)
|
||||
return turnResult{}, err
|
||||
}
|
||||
}
|
||||
return turnResult{status: TurnEndStatusAborted}, nil
|
||||
}
|
||||
|
||||
func (al *AgentLoop) selectCandidates(
|
||||
agent *AgentInstance,
|
||||
userMsg string,
|
||||
history []providers.Message,
|
||||
) (candidates []providers.FallbackCandidate, model string, usedLight bool) {
|
||||
if agent.Router == nil || len(agent.LightCandidates) == 0 {
|
||||
return agent.Candidates, resolvedCandidateModel(agent.Candidates, agent.Model), false
|
||||
}
|
||||
|
||||
_, usedLight, score := agent.Router.SelectModel(userMsg, history, agent.Model)
|
||||
if !usedLight {
|
||||
logger.DebugCF("agent", "Model routing: primary model selected",
|
||||
map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"score": score,
|
||||
"threshold": agent.Router.Threshold(),
|
||||
})
|
||||
return agent.Candidates, resolvedCandidateModel(agent.Candidates, agent.Model), false
|
||||
}
|
||||
|
||||
logger.InfoCF("agent", "Model routing: light model selected",
|
||||
map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"light_model": agent.Router.LightModel(),
|
||||
"score": score,
|
||||
"threshold": agent.Router.Threshold(),
|
||||
})
|
||||
return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true
|
||||
}
|
||||
|
||||
func (al *AgentLoop) resolveContextManager() ContextManager {
|
||||
name := al.cfg.Agents.Defaults.ContextManager
|
||||
if name == "" || name == "legacy" {
|
||||
return &legacyContextManager{al: al}
|
||||
}
|
||||
factory, ok := lookupContextManager(name)
|
||||
if !ok {
|
||||
logger.WarnCF("agent", "Unknown context manager, falling back to legacy", map[string]any{
|
||||
"name": name,
|
||||
})
|
||||
return &legacyContextManager{al: al}
|
||||
}
|
||||
cm, err := factory(al.cfg.Agents.Defaults.ContextManagerConfig, al)
|
||||
if err != nil {
|
||||
logger.WarnCF("agent", "Failed to create context manager, falling back to legacy", map[string]any{
|
||||
"name": name,
|
||||
"error": err.Error(),
|
||||
})
|
||||
return &legacyContextManager{al: al}
|
||||
}
|
||||
return cm
|
||||
}
|
||||
|
||||
func (al *AgentLoop) askSideQuestion(
|
||||
ctx context.Context,
|
||||
agent *AgentInstance,
|
||||
opts *processOptions,
|
||||
question string,
|
||||
) (string, error) {
|
||||
if agent == nil {
|
||||
return "", fmt.Errorf("askSideQuestion: no agent available for /btw")
|
||||
}
|
||||
|
||||
question = strings.TrimSpace(question)
|
||||
if question == "" {
|
||||
return "", fmt.Errorf("askSideQuestion: %w", fmt.Errorf("Usage: /btw <question>"))
|
||||
}
|
||||
|
||||
if opts != nil {
|
||||
normalizeProcessOptionsInPlace(opts)
|
||||
}
|
||||
|
||||
var media []string
|
||||
var channel, chatID, senderID, senderDisplayName string
|
||||
if opts != nil {
|
||||
media = opts.Media
|
||||
channel = opts.Channel
|
||||
chatID = opts.ChatID
|
||||
senderID = opts.SenderID
|
||||
senderDisplayName = opts.SenderDisplayName
|
||||
}
|
||||
|
||||
// Build messages with context but WITHOUT adding to session history
|
||||
var history []providers.Message
|
||||
var summary string
|
||||
if opts != nil && !opts.NoHistory {
|
||||
if resp, err := al.contextManager.Assemble(ctx, &AssembleRequest{
|
||||
SessionKey: opts.SessionKey,
|
||||
Budget: agent.ContextWindow,
|
||||
MaxTokens: agent.MaxTokens,
|
||||
}); err == nil && resp != nil {
|
||||
history = resp.History
|
||||
summary = resp.Summary
|
||||
}
|
||||
}
|
||||
|
||||
messages := agent.ContextBuilder.BuildMessages(
|
||||
history,
|
||||
summary,
|
||||
question,
|
||||
media,
|
||||
channel,
|
||||
chatID,
|
||||
senderID,
|
||||
senderDisplayName,
|
||||
)
|
||||
|
||||
maxMediaSize := al.GetConfig().Agents.Defaults.GetMaxMediaSize()
|
||||
messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize)
|
||||
|
||||
activeCandidates, activeModel, usedLight := al.selectCandidates(agent, question, messages)
|
||||
selectedModelName := sideQuestionModelName(agent, usedLight)
|
||||
|
||||
llmOpts := map[string]any{
|
||||
"max_tokens": agent.MaxTokens,
|
||||
"temperature": agent.Temperature,
|
||||
"prompt_cache_key": agent.ID + ":btw",
|
||||
}
|
||||
|
||||
hookModelChanged := false
|
||||
callProvider := func(
|
||||
ctx context.Context,
|
||||
candidate providers.FallbackCandidate,
|
||||
model string,
|
||||
forceModel bool,
|
||||
callMessages []providers.Message,
|
||||
) (*providers.LLMResponse, error) {
|
||||
provider, providerModel, cleanup, err := al.isolatedSideQuestionProvider(agent, selectedModelName, candidate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer cleanup()
|
||||
if !forceModel || strings.TrimSpace(model) == "" {
|
||||
model = providerModel
|
||||
}
|
||||
callOpts := llmOpts
|
||||
if _, exists := callOpts["thinking_level"]; !exists && agent.ThinkingLevel != ThinkingOff {
|
||||
if tc, ok := provider.(providers.ThinkingCapable); ok && tc.SupportsThinking() {
|
||||
callOpts = shallowCloneLLMOptions(llmOpts)
|
||||
callOpts["thinking_level"] = string(agent.ThinkingLevel)
|
||||
}
|
||||
}
|
||||
return provider.Chat(ctx, callMessages, nil, model, callOpts)
|
||||
}
|
||||
|
||||
turnCtx := newTurnContext(nil, nil, nil)
|
||||
if opts != nil {
|
||||
turnCtx = newTurnContext(opts.Dispatch.InboundContext, opts.Dispatch.RouteResult, opts.Dispatch.SessionScope)
|
||||
}
|
||||
llmModel := activeModel
|
||||
if al.hooks != nil {
|
||||
llmReq, decision := al.hooks.BeforeLLM(ctx, &LLMHookRequest{
|
||||
Meta: EventMeta{
|
||||
Source: "askSideQuestion",
|
||||
TracePath: "turn.llm.request",
|
||||
turnContext: cloneTurnContext(turnCtx),
|
||||
},
|
||||
Context: cloneTurnContext(turnCtx),
|
||||
Model: llmModel,
|
||||
Messages: messages,
|
||||
Tools: nil,
|
||||
Options: llmOpts,
|
||||
GracefulTerminal: false,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if llmReq != nil {
|
||||
if strings.TrimSpace(llmReq.Model) != "" && llmReq.Model != llmModel {
|
||||
hookModelChanged = true
|
||||
}
|
||||
llmModel = llmReq.Model
|
||||
messages = llmReq.Messages
|
||||
llmOpts = llmReq.Options
|
||||
}
|
||||
case HookActionAbortTurn:
|
||||
reason := decision.Reason
|
||||
if reason == "" {
|
||||
reason = "hook requested turn abort"
|
||||
}
|
||||
return "", fmt.Errorf("hook aborted turn during before_llm: %s", reason)
|
||||
case HookActionHardAbort:
|
||||
reason := decision.Reason
|
||||
if reason == "" {
|
||||
reason = "hook requested turn abort"
|
||||
}
|
||||
return "", fmt.Errorf("hook aborted turn during before_llm: %s", reason)
|
||||
}
|
||||
}
|
||||
if hookModelChanged {
|
||||
// Hook-selected models must not continue through the pre-hook fallback
|
||||
// candidate list, otherwise fallback execution would call the original
|
||||
// candidate model and silently ignore the hook decision.
|
||||
activeCandidates = nil
|
||||
}
|
||||
|
||||
callSideLLM := func(callMessages []providers.Message) (*providers.LLMResponse, error) {
|
||||
if len(activeCandidates) > 1 && al.fallback != nil {
|
||||
fbResult, err := al.fallback.Execute(
|
||||
ctx,
|
||||
activeCandidates,
|
||||
func(ctx context.Context, providerName, model string) (*providers.LLMResponse, error) {
|
||||
candidate := providers.FallbackCandidate{Provider: providerName, Model: model}
|
||||
for _, activeCandidate := range activeCandidates {
|
||||
if activeCandidate.Provider == providerName && activeCandidate.Model == model {
|
||||
candidate = activeCandidate
|
||||
break
|
||||
}
|
||||
}
|
||||
return callProvider(ctx, candidate, model, false, callMessages)
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return fbResult.Response, nil
|
||||
}
|
||||
|
||||
var candidate providers.FallbackCandidate
|
||||
if len(activeCandidates) > 0 {
|
||||
candidate = activeCandidates[0]
|
||||
}
|
||||
return callProvider(ctx, candidate, llmModel, hookModelChanged, callMessages)
|
||||
}
|
||||
|
||||
// Retry without media if vision is unsupported
|
||||
// Note: Vision retry is only applied to the initial call. If fallback chain
|
||||
// is used, vision errors from fallback providers will not trigger retry.
|
||||
var resp *providers.LLMResponse
|
||||
var err error
|
||||
resp, err = callSideLLM(messages)
|
||||
if err != nil && hasMediaRefs(messages) && isVisionUnsupportedError(err) {
|
||||
al.emitEvent(
|
||||
EventKindLLMRetry,
|
||||
EventMeta{
|
||||
Source: "askSideQuestion",
|
||||
TracePath: "turn.llm.retry",
|
||||
turnContext: cloneTurnContext(turnCtx),
|
||||
},
|
||||
LLMRetryPayload{
|
||||
Attempt: 1,
|
||||
MaxRetries: 1,
|
||||
Reason: "vision_unsupported",
|
||||
Error: err.Error(),
|
||||
Backoff: 0,
|
||||
},
|
||||
)
|
||||
messagesWithoutMedia := stripMessageMedia(messages)
|
||||
resp, err = callSideLLM(messagesWithoutMedia)
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if resp == nil {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// Apply after_llm hooks
|
||||
if al.hooks != nil {
|
||||
llmResp, decision := al.hooks.AfterLLM(ctx, &LLMHookResponse{
|
||||
Meta: EventMeta{
|
||||
Source: "askSideQuestion",
|
||||
TracePath: "turn.llm.response",
|
||||
turnContext: cloneTurnContext(turnCtx),
|
||||
},
|
||||
Context: cloneTurnContext(turnCtx),
|
||||
Model: llmModel,
|
||||
Response: resp,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if llmResp != nil && llmResp.Response != nil {
|
||||
resp = llmResp.Response
|
||||
}
|
||||
case HookActionAbortTurn, HookActionHardAbort:
|
||||
reason := decision.Reason
|
||||
if reason == "" {
|
||||
reason = "hook requested turn abort"
|
||||
}
|
||||
return "", fmt.Errorf("hook aborted turn during after_llm: %s", reason)
|
||||
}
|
||||
}
|
||||
|
||||
return sideQuestionResponseContent(resp), nil
|
||||
}
|
||||
|
||||
func (al *AgentLoop) isolatedSideQuestionProvider(
|
||||
agent *AgentInstance,
|
||||
baseModelName string,
|
||||
candidate providers.FallbackCandidate,
|
||||
) (providers.LLMProvider, string, func(), error) {
|
||||
if agent == nil {
|
||||
return nil, "", func() {}, fmt.Errorf("isolatedSideQuestionProvider: no agent available for /btw")
|
||||
}
|
||||
|
||||
modelCfg, err := al.sideQuestionModelConfig(agent, baseModelName, candidate)
|
||||
if err != nil {
|
||||
return nil, "", func() {}, fmt.Errorf("isolatedSideQuestionProvider: %w", err)
|
||||
}
|
||||
|
||||
factory := al.providerFactory
|
||||
if factory == nil {
|
||||
factory = providers.CreateProviderFromConfig
|
||||
}
|
||||
provider, modelID, err := factory(modelCfg)
|
||||
if err != nil {
|
||||
return nil, "", func() {}, fmt.Errorf("isolatedSideQuestionProvider: %w", err)
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
closeProviderIfStateful(provider)
|
||||
}
|
||||
return provider, modelID, cleanup, nil
|
||||
}
|
||||
|
||||
func (al *AgentLoop) sideQuestionModelConfig(
|
||||
agent *AgentInstance,
|
||||
baseModelName string,
|
||||
candidate providers.FallbackCandidate,
|
||||
) (*config.ModelConfig, error) {
|
||||
if agent == nil {
|
||||
return nil, fmt.Errorf("sideQuestionModelConfig: no agent available for /btw")
|
||||
}
|
||||
|
||||
// If candidate has an identity key, use that
|
||||
if name := modelNameFromIdentityKey(candidate.IdentityKey); name != "" {
|
||||
modelCfg, err := resolvedModelConfig(al.GetConfig(), name, agent.Workspace)
|
||||
if err == nil {
|
||||
return modelCfg, nil
|
||||
}
|
||||
// Fallback: create a minimal config if lookup fails
|
||||
}
|
||||
|
||||
// Otherwise, clean up the base model name and use it
|
||||
baseModelName = strings.TrimSpace(baseModelName)
|
||||
modelCfg, err := resolvedModelConfig(al.GetConfig(), baseModelName, agent.Workspace)
|
||||
if err != nil {
|
||||
// Fallback: create a minimal config for test scenarios
|
||||
model := strings.TrimSpace(baseModelName)
|
||||
if candidate.Model != "" {
|
||||
model = candidate.Model
|
||||
}
|
||||
if candidate.Provider != "" && candidate.Model != "" {
|
||||
model = providers.NormalizeProvider(candidate.Provider) + "/" + candidate.Model
|
||||
} else {
|
||||
model = ensureProtocolModel(model)
|
||||
}
|
||||
return &config.ModelConfig{
|
||||
ModelName: baseModelName,
|
||||
Model: model,
|
||||
Workspace: agent.Workspace,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// If candidate specifies a different provider/model, override
|
||||
clone := *modelCfg
|
||||
if candidate.Provider != "" && candidate.Model != "" {
|
||||
clone.Model = providers.NormalizeProvider(candidate.Provider) + "/" + candidate.Model
|
||||
}
|
||||
return &clone, nil
|
||||
}
|
||||
615
pkg/agent/turn_coord_test.go
Normal file
615
pkg/agent/turn_coord_test.go
Normal file
|
|
@ -0,0 +1,615 @@
|
|||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// =============================================================================
|
||||
// Mock Providers for turn_coord Tests
|
||||
// =============================================================================
|
||||
|
||||
// simpleConvProvider returns a simple text response without tools
|
||||
type simpleConvProvider struct{}
|
||||
|
||||
func (p *simpleConvProvider) Chat(
|
||||
ctx context.Context,
|
||||
messages []providers.Message,
|
||||
tools []providers.ToolDefinition,
|
||||
model string,
|
||||
opts map[string]any,
|
||||
) (*providers.LLMResponse, error) {
|
||||
return &providers.LLMResponse{
|
||||
Content: "Hello! How can I help you today?",
|
||||
FinishReason: "stop",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *simpleConvProvider) GetDefaultModel() string {
|
||||
return "simple-model"
|
||||
}
|
||||
|
||||
type nativeSearchCaptureProvider struct {
|
||||
lastOpts map[string]any
|
||||
}
|
||||
|
||||
func (p *nativeSearchCaptureProvider) Chat(
|
||||
ctx context.Context,
|
||||
messages []providers.Message,
|
||||
tools []providers.ToolDefinition,
|
||||
model string,
|
||||
opts map[string]any,
|
||||
) (*providers.LLMResponse, error) {
|
||||
p.lastOpts = make(map[string]any, len(opts))
|
||||
for k, v := range opts {
|
||||
p.lastOpts[k] = v
|
||||
}
|
||||
return &providers.LLMResponse{
|
||||
Content: "Using native search",
|
||||
FinishReason: "stop",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *nativeSearchCaptureProvider) GetDefaultModel() string {
|
||||
return "native-search-model"
|
||||
}
|
||||
|
||||
func (p *nativeSearchCaptureProvider) SupportsNativeSearch() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// toolCallRespProvider returns a tool call response
|
||||
type toolCallRespProvider struct {
|
||||
toolName string
|
||||
toolArgs map[string]any
|
||||
response string
|
||||
callCount int
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func (p *toolCallRespProvider) Chat(
|
||||
ctx context.Context,
|
||||
messages []providers.Message,
|
||||
tools []providers.ToolDefinition,
|
||||
model string,
|
||||
opts map[string]any,
|
||||
) (*providers.LLMResponse, error) {
|
||||
p.mu.Lock()
|
||||
p.callCount++
|
||||
count := p.callCount
|
||||
p.mu.Unlock()
|
||||
|
||||
// First call returns a tool call, subsequent calls return final response
|
||||
if count == 1 {
|
||||
return &providers.LLMResponse{
|
||||
Content: "Let me search for that information.",
|
||||
ToolCalls: []providers.ToolCall{
|
||||
{
|
||||
ID: "call_1",
|
||||
Name: p.toolName,
|
||||
Arguments: p.toolArgs,
|
||||
},
|
||||
},
|
||||
FinishReason: "tool_calls",
|
||||
}, nil
|
||||
}
|
||||
return &providers.LLMResponse{
|
||||
Content: p.response,
|
||||
FinishReason: "stop",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *toolCallRespProvider) GetDefaultModel() string {
|
||||
return "tool-model"
|
||||
}
|
||||
|
||||
// errorProvider simulates various error conditions
|
||||
type errorProvider struct {
|
||||
errType string
|
||||
callCount int
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func (p *errorProvider) Chat(
|
||||
ctx context.Context,
|
||||
messages []providers.Message,
|
||||
tools []providers.ToolDefinition,
|
||||
model string,
|
||||
opts map[string]any,
|
||||
) (*providers.LLMResponse, error) {
|
||||
p.mu.Lock()
|
||||
p.callCount++
|
||||
p.mu.Unlock()
|
||||
|
||||
switch p.errType {
|
||||
case "timeout":
|
||||
return nil, context.DeadlineExceeded
|
||||
case "context_length":
|
||||
return nil, errors.New("context_length_exceeded")
|
||||
case "vision":
|
||||
return nil, errors.New("vision_unsupported")
|
||||
default:
|
||||
return nil, errors.New("unknown error")
|
||||
}
|
||||
}
|
||||
|
||||
func (p *errorProvider) GetDefaultModel() string {
|
||||
return "error-model"
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Test Helper Functions
|
||||
// =============================================================================
|
||||
|
||||
func newTurnCoordTestLoop(t *testing.T, provider providers.LLMProvider) (*AgentLoop, *AgentInstance, func()) {
|
||||
t.Helper()
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "test-model",
|
||||
MaxTokens: 4096,
|
||||
MaxToolIterations: 10,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
msgBus := bus.NewMessageBus()
|
||||
al := NewAgentLoop(cfg, msgBus, provider)
|
||||
agent := al.registry.GetDefaultAgent()
|
||||
if agent == nil {
|
||||
t.Fatal("expected default agent")
|
||||
}
|
||||
|
||||
return al, agent, func() {
|
||||
al.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func makeTestProcessOpts(sessionKey string) processOptions {
|
||||
return processOptions{
|
||||
SessionKey: sessionKey,
|
||||
Channel: "cli",
|
||||
ChatID: "test-chat",
|
||||
UserMessage: "test message",
|
||||
DefaultResponse: "I couldn't process your request.",
|
||||
EnableSummary: false,
|
||||
SendResponse: false,
|
||||
NoHistory: false,
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Pipeline Method Tests: SetupTurn
|
||||
// =============================================================================
|
||||
|
||||
func TestPipeline_SetupTurn_BasicInitialization(t *testing.T) {
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, &simpleConvProvider{})
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
if exec == nil {
|
||||
t.Fatal("expected non-nil turnExecution")
|
||||
}
|
||||
if len(exec.messages) == 0 {
|
||||
t.Error("expected messages to be populated")
|
||||
}
|
||||
if exec.iteration != 0 {
|
||||
t.Errorf("expected iteration 0, got %d", exec.iteration)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Pipeline Method Tests: CallLLM
|
||||
// =============================================================================
|
||||
|
||||
func TestPipeline_CallLLM_SimpleResponse(t *testing.T) {
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, &simpleConvProvider{})
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
|
||||
ctrl, err := pipeline.CallLLM(context.Background(), context.Background(), ts, exec, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CallLLM failed: %v", err)
|
||||
}
|
||||
if ctrl != ControlBreak {
|
||||
t.Errorf("expected ControlBreak, got %v", ctrl)
|
||||
}
|
||||
if exec.response == nil {
|
||||
t.Fatal("expected non-nil response")
|
||||
}
|
||||
if exec.response.Content == "" {
|
||||
t.Error("expected non-empty content")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPipeline_CallLLM_WithToolCall(t *testing.T) {
|
||||
provider := &toolCallRespProvider{
|
||||
toolName: "web_search",
|
||||
toolArgs: map[string]any{"query": "test"},
|
||||
response: "Found information about test.",
|
||||
}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
|
||||
ctrl, err := pipeline.CallLLM(context.Background(), context.Background(), ts, exec, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CallLLM failed: %v", err)
|
||||
}
|
||||
if ctrl != ControlToolLoop {
|
||||
t.Errorf("expected ControlToolLoop, got %v", ctrl)
|
||||
}
|
||||
if len(exec.normalizedToolCalls) == 0 {
|
||||
t.Fatal("expected tool calls")
|
||||
}
|
||||
if exec.normalizedToolCalls[0].Name != "web_search" {
|
||||
t.Errorf("expected tool name 'web_search', got %q", exec.normalizedToolCalls[0].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPipeline_CallLLM_UsesNativeSearchWithoutClientWebSearchTool(t *testing.T) {
|
||||
provider := &nativeSearchCaptureProvider{}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
if _, ok := agent.Tools.Get("web_search"); ok {
|
||||
t.Fatal("expected no client-side web_search tool to be registered")
|
||||
}
|
||||
|
||||
al.cfg.Tools.Web.Enabled = true
|
||||
al.cfg.Tools.Web.PreferNative = true
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
|
||||
ctrl, err := pipeline.CallLLM(context.Background(), context.Background(), ts, exec, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CallLLM failed: %v", err)
|
||||
}
|
||||
if ctrl != ControlBreak {
|
||||
t.Fatalf("expected ControlBreak, got %v", ctrl)
|
||||
}
|
||||
if got, _ := provider.lastOpts["native_search"].(bool); !got {
|
||||
t.Fatalf("expected native_search=true, got %#v", provider.lastOpts["native_search"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestPipeline_CallLLM_TimeoutRetry(t *testing.T) {
|
||||
errorPrv := &errorProvider{errType: "timeout"}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, errorPrv)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
|
||||
// Should retry and eventually fail after max retries
|
||||
_, err = pipeline.CallLLM(context.Background(), context.Background(), ts, exec, 1)
|
||||
if err == nil {
|
||||
t.Error("expected error after retries")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPipeline_CallLLM_ContextLengthError(t *testing.T) {
|
||||
errorPrv := &errorProvider{errType: "context_length"}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, errorPrv)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
|
||||
// Should trigger context compression and retry
|
||||
_, err = pipeline.CallLLM(context.Background(), context.Background(), ts, exec, 1)
|
||||
// May succeed after compression or fail - either is acceptable
|
||||
t.Logf("CallLLM result after context error: err=%v", err)
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Pipeline Method Tests: ExecuteTools
|
||||
// =============================================================================
|
||||
|
||||
func TestPipeline_ExecuteTools_NoTools(t *testing.T) {
|
||||
// Provider returns no tool calls, so ExecuteTools should not be called
|
||||
// This test verifies the ControlBreak path from CallLLM
|
||||
provider := &simpleConvProvider{}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
ts := newTurnState(agent, makeTestProcessOpts("test-session"), turnEventScope{
|
||||
turnID: "turn-1",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
exec, err := pipeline.SetupTurn(context.Background(), ts)
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTurn failed: %v", err)
|
||||
}
|
||||
|
||||
// First CallLLM returns ControlBreak (no tools)
|
||||
ctrl, err := pipeline.CallLLM(context.Background(), context.Background(), ts, exec, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CallLLM failed: %v", err)
|
||||
}
|
||||
|
||||
if ctrl != ControlBreak {
|
||||
t.Fatalf("expected ControlBreak, got %v", ctrl)
|
||||
}
|
||||
// No tools to execute, Finalize should be called directly
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// runTurn Integration Tests
|
||||
// =============================================================================
|
||||
|
||||
func TestRunTurn_SimpleConversation(t *testing.T) {
|
||||
provider := &simpleConvProvider{}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
opts := makeTestProcessOpts("test-session-simple")
|
||||
|
||||
ts := newTurnState(agent, opts, turnEventScope{
|
||||
turnID: "turn-simple",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
result, err := al.runTurn(context.Background(), ts, pipeline)
|
||||
if err != nil {
|
||||
t.Fatalf("runTurn failed: %v", err)
|
||||
}
|
||||
if result.status != TurnEndStatusCompleted {
|
||||
t.Errorf("expected status Completed, got %v", result.status)
|
||||
}
|
||||
if result.finalContent == "" {
|
||||
t.Error("expected non-empty finalContent")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunTurn_MaxIterations(t *testing.T) {
|
||||
// Provider always returns tool calls, should hit max iterations
|
||||
provider := &toolCallRespProvider{
|
||||
toolName: "search",
|
||||
toolArgs: map[string]any{"q": "x"},
|
||||
response: "done",
|
||||
}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
// Override max iterations to 2
|
||||
agent.MaxIterations = 2
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
opts := makeTestProcessOpts("test-session-maxiter")
|
||||
|
||||
ts := newTurnState(agent, opts, turnEventScope{
|
||||
turnID: "turn-maxiter",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
result, err := al.runTurn(context.Background(), ts, pipeline)
|
||||
if err != nil {
|
||||
t.Fatalf("runTurn failed: %v", err)
|
||||
}
|
||||
// Should complete due to max iterations
|
||||
if result.status != TurnEndStatusCompleted {
|
||||
t.Errorf("expected status Completed, got %v", result.status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunTurn_HardAbort(t *testing.T) {
|
||||
// Provider simulates a slow response, but we'll abort mid-turn
|
||||
slowProvider := &slowMockProvider{delay: 10 * time.Second}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, slowProvider)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
opts := makeTestProcessOpts("test-session-abort")
|
||||
|
||||
ts := newTurnState(agent, opts, turnEventScope{
|
||||
turnID: "turn-abort",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
// Run in goroutine with abort after short delay
|
||||
done := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
al.runTurn(context.Background(), ts, pipeline)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Give it a moment to start
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Request hard abort
|
||||
ts.requestHardAbort()
|
||||
|
||||
// Wait for runTurn to complete
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("runTurn did not complete after abort")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunTurn_SteeringMessageInjection(t *testing.T) {
|
||||
provider := &simpleConvProvider{}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
opts := makeTestProcessOpts("test-session-steering")
|
||||
|
||||
ts := newTurnState(agent, opts, turnEventScope{
|
||||
turnID: "turn-steering",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
// Enqueue steering message before runTurn
|
||||
steeringMsg := providers.Message{
|
||||
Role: "user",
|
||||
Content: "Steering message",
|
||||
}
|
||||
al.Steer(steeringMsg)
|
||||
|
||||
result, err := al.runTurn(context.Background(), ts, pipeline)
|
||||
if err != nil {
|
||||
t.Fatalf("runTurn failed: %v", err)
|
||||
}
|
||||
if result.status != TurnEndStatusCompleted {
|
||||
t.Errorf("expected status Completed, got %v", result.status)
|
||||
}
|
||||
// Steering message should have been injected
|
||||
}
|
||||
|
||||
func TestRunTurn_GracefulInterrupt(t *testing.T) {
|
||||
provider := &toolCallRespProvider{
|
||||
toolName: "search",
|
||||
toolArgs: map[string]any{"q": "test"},
|
||||
response: "Final response after interrupt",
|
||||
}
|
||||
al, agent, cleanup := newTurnCoordTestLoop(t, provider)
|
||||
defer cleanup()
|
||||
|
||||
pipeline := NewPipeline(al)
|
||||
opts := makeTestProcessOpts("test-session-graceful")
|
||||
|
||||
ts := newTurnState(agent, opts, turnEventScope{
|
||||
turnID: "turn-graceful",
|
||||
context: newTurnContext(nil, nil, nil),
|
||||
})
|
||||
|
||||
// Run in goroutine with graceful interrupt after first iteration
|
||||
done := make(chan struct{})
|
||||
var result turnResult
|
||||
|
||||
go func() {
|
||||
result, _ = al.runTurn(context.Background(), ts, pipeline)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Give it a moment to start first iteration
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Request graceful interrupt
|
||||
ts.requestGracefulInterrupt("Please stop")
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("runTurn did not complete after graceful interrupt")
|
||||
}
|
||||
|
||||
// Should complete gracefully
|
||||
if result.status != TurnEndStatusCompleted {
|
||||
t.Errorf("expected status Completed, got %v", result.status)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// turnState Tests
|
||||
// =============================================================================
|
||||
|
||||
func TestTurnState_GracefulInterruptRequested(t *testing.T) {
|
||||
ts := &turnState{
|
||||
gracefulInterrupt: false,
|
||||
gracefulInterruptHint: "",
|
||||
}
|
||||
|
||||
// Initially should not be requested
|
||||
requested, _ := ts.gracefulInterruptRequested()
|
||||
if requested {
|
||||
t.Error("expected no interrupt initially")
|
||||
}
|
||||
|
||||
// Request interrupt
|
||||
ts.requestGracefulInterrupt("test hint")
|
||||
|
||||
requested, hint := ts.gracefulInterruptRequested()
|
||||
if !requested {
|
||||
t.Error("expected interrupt to be requested")
|
||||
}
|
||||
if hint != "test hint" {
|
||||
t.Errorf("expected hint 'test hint', got %q", hint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTurnState_HardAbortRequested(t *testing.T) {
|
||||
ts := &turnState{
|
||||
hardAbort: false,
|
||||
}
|
||||
|
||||
if ts.hardAbortRequested() {
|
||||
t.Error("expected no hard abort initially")
|
||||
}
|
||||
|
||||
ts.requestHardAbort()
|
||||
|
||||
if !ts.hardAbortRequested() {
|
||||
t.Error("expected hard abort to be requested")
|
||||
}
|
||||
}
|
||||
|
|
@ -1,3 +1,5 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
|
|
@ -14,6 +16,10 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
)
|
||||
|
||||
// =============================================================================
|
||||
// TurnPhase - represents the current phase of a turn
|
||||
// =============================================================================
|
||||
|
||||
type TurnPhase string
|
||||
|
||||
const (
|
||||
|
|
@ -25,6 +31,65 @@ const (
|
|||
TurnPhaseAborted TurnPhase = "aborted"
|
||||
)
|
||||
|
||||
// =============================================================================
|
||||
// Control signals - returned from Pipeline methods to drive runTurn's coordinator loop
|
||||
// =============================================================================
|
||||
|
||||
type Control int
|
||||
|
||||
const (
|
||||
// ControlContinue tells the coordinator to jump back to the top of the turn loop
|
||||
// (equivalent to the original "goto turnLoop").
|
||||
ControlContinue Control = iota
|
||||
// ControlBreak tells the coordinator to exit the turn loop and proceed to Finalize.
|
||||
ControlBreak
|
||||
// ControlToolLoop tells the coordinator to execute the tool loop.
|
||||
ControlToolLoop
|
||||
)
|
||||
|
||||
// ToolControl signals returned from ExecuteTools to drive tool loop iteration.
|
||||
type ToolControl int
|
||||
|
||||
const (
|
||||
// ToolControlContinue tells the tool loop to jump to the next iteration
|
||||
// (pendingMessages arrived, SubTurn results, etc.).
|
||||
ToolControlContinue ToolControl = iota
|
||||
// ToolControlBreak tells the tool loop to exit and return to the coordinator.
|
||||
ToolControlBreak
|
||||
// ToolControlFinalize tells the coordinator that all tool responses were
|
||||
// handled and the turn should finalize without another LLM call.
|
||||
ToolControlFinalize
|
||||
)
|
||||
|
||||
// LLMPhase indicates which phase the turn is executing in.
|
||||
type LLMPhase int
|
||||
|
||||
const (
|
||||
LLMPhaseSetup LLMPhase = iota
|
||||
LLMPhasePreLLM
|
||||
LLMPhaseLLMCall
|
||||
LLMPhaseProcessing
|
||||
LLMPhaseToolLoop
|
||||
LLMPhaseTools
|
||||
LLMPhaseFinalizing
|
||||
LLMPhaseCompleted
|
||||
LLMPhaseAborted
|
||||
)
|
||||
|
||||
// =============================================================================
|
||||
// turnResult - returned from runTurn
|
||||
// =============================================================================
|
||||
|
||||
type turnResult struct {
|
||||
finalContent string
|
||||
status TurnEndStatus
|
||||
followUps []bus.InboundMessage
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ActiveTurnInfo - public info about an active turn
|
||||
// =============================================================================
|
||||
|
||||
type ActiveTurnInfo struct {
|
||||
TurnID string
|
||||
AgentID string
|
||||
|
|
@ -40,12 +105,70 @@ type ActiveTurnInfo struct {
|
|||
ChildTurnIDs []string
|
||||
}
|
||||
|
||||
type turnResult struct {
|
||||
// =============================================================================
|
||||
// turnExecution - mutable state that persists across turn loop iterations
|
||||
// =============================================================================
|
||||
|
||||
type turnExecution struct {
|
||||
// Core message state (accumulates throughout the turn)
|
||||
messages []providers.Message // built from ContextBuilder, grows per-iteration
|
||||
pendingMessages []providers.Message // steering/SubTurn messages awaiting injection
|
||||
history []providers.Message // from ContextManager.Assemble
|
||||
summary string
|
||||
|
||||
// Turn output
|
||||
finalContent string
|
||||
status TurnEndStatus
|
||||
followUps []bus.InboundMessage
|
||||
|
||||
// Iteration tracking
|
||||
iteration int
|
||||
|
||||
// Per-iteration state set by Pipeline.PreLLM
|
||||
activeCandidates []providers.FallbackCandidate
|
||||
activeModel string
|
||||
activeProvider providers.LLMProvider
|
||||
usedLight bool
|
||||
|
||||
// LLM call per-iteration state
|
||||
response *providers.LLMResponse
|
||||
normalizedToolCalls []providers.ToolCall
|
||||
allResponsesHandled bool
|
||||
callMessages []providers.Message
|
||||
providerToolDefs []providers.ToolDefinition
|
||||
llmModel string
|
||||
llmOpts map[string]any
|
||||
gracefulTerminal bool
|
||||
useNativeSearch bool
|
||||
|
||||
// Phase tracking
|
||||
phase LLMPhase
|
||||
|
||||
// Abort signaling for coordinator (set by Pipeline methods)
|
||||
abortedByHardAbort bool // true when hard abort triggered during LLM/tools
|
||||
abortedByHook bool // true when HookActionAbortTurn triggered
|
||||
}
|
||||
|
||||
// newTurnExecution creates a turnExecution initialized from turnState and options.
|
||||
func newTurnExecution(
|
||||
agent *AgentInstance,
|
||||
opts processOptions,
|
||||
history []providers.Message,
|
||||
summary string,
|
||||
messages []providers.Message,
|
||||
) *turnExecution {
|
||||
return &turnExecution{
|
||||
history: history,
|
||||
summary: summary,
|
||||
messages: messages,
|
||||
pendingMessages: append([]providers.Message(nil), opts.InitialSteeringMessages...),
|
||||
iteration: 0,
|
||||
phase: LLMPhaseSetup,
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// turnState - the full state for a turn, constructed once per turn
|
||||
// =============================================================================
|
||||
|
||||
type turnState struct {
|
||||
mu sync.RWMutex
|
||||
|
||||
|
|
@ -109,6 +232,10 @@ type turnState struct {
|
|||
al *AgentLoop
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// turnState constructors and active turn management
|
||||
// =============================================================================
|
||||
|
||||
func newTurnState(agent *AgentInstance, opts processOptions, scope turnEventScope) *turnState {
|
||||
ts := &turnState{
|
||||
agent: agent,
|
||||
|
|
@ -194,6 +321,10 @@ func (al *AgentLoop) GetActiveTurnBySession(sessionKey string) *ActiveTurnInfo {
|
|||
return &info
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// turnState - getters and setters
|
||||
// =============================================================================
|
||||
|
||||
func (ts *turnState) snapshot() ActiveTurnInfo {
|
||||
ts.mu.RLock()
|
||||
defer ts.mu.RUnlock()
|
||||
|
|
@ -402,7 +533,9 @@ func (ts *turnState) interruptHintMessage() providers.Message {
|
|||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// SubTurn-related methods
|
||||
// =============================================================================
|
||||
|
||||
// Finish marks the turn as finished and closes the pendingResults channel
|
||||
func (ts *turnState) Finish(isHardAbort bool) {
|
||||
|
|
@ -421,9 +554,9 @@ func (ts *turnState) Finish(isHardAbort bool) {
|
|||
ts.mu.Unlock()
|
||||
})
|
||||
|
||||
// If this is a graceful finish (not hard abort), signal to children
|
||||
if !isHardAbort && ts.parentTurnState == nil {
|
||||
// This is a root turn finishing gracefully
|
||||
// Any graceful finish must signal direct children so nested SubTurns can
|
||||
// observe parent completion and decide whether to stop or continue.
|
||||
if !isHardAbort {
|
||||
ts.parentEnded.Store(true)
|
||||
}
|
||||
|
||||
|
|
@ -493,7 +626,9 @@ func (ts *turnState) SetLastUsage(usage *providers.UsageInfo) {
|
|||
ts.lastUsage = usage
|
||||
}
|
||||
|
||||
// Context helper functions for SubTurn
|
||||
// =============================================================================
|
||||
// Context helper functions for turnState
|
||||
// =============================================================================
|
||||
|
||||
type turnStateKeyType struct{}
|
||||
|
||||
|
|
@ -19,16 +19,16 @@ type TranscriptionResponse struct {
|
|||
Duration float64 `json:"duration,omitempty"`
|
||||
}
|
||||
|
||||
func supportsAudioTranscription(model string) bool {
|
||||
protocol, _ := providers.ExtractProtocol(model)
|
||||
func supportsAudioTranscription(modelCfg *config.ModelConfig) bool {
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
|
||||
switch protocol {
|
||||
case "openai", "azure", "azure-openai",
|
||||
"litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-portal", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"qwen-us", "dashscope-us", "mistral", "avian", "minimax", "longcat", "modelscope", "novita",
|
||||
"coding-plan", "alibaba-coding", "qwen-coding":
|
||||
"coding-plan", "alibaba-coding", "qwen-coding", "zai":
|
||||
// These protocols all go through the OpenAI-compatible or Azure provider path in
|
||||
// providers.CreateProviderFromConfig, so they are the only ones that can supply
|
||||
// the audio media payload shape expected by NewAudioModelTranscriber.
|
||||
|
|
@ -41,15 +41,15 @@ func supportsAudioTranscription(model string) bool {
|
|||
}
|
||||
}
|
||||
|
||||
func supportsWhisperTranscription(model string) bool {
|
||||
protocol, _ := providers.ExtractProtocol(model)
|
||||
func supportsWhisperTranscription(modelCfg *config.ModelConfig) bool {
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
|
||||
switch protocol {
|
||||
case "openai", "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-portal", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||
"qwen-us", "dashscope-us", "mistral", "avian", "minimax", "longcat", "modelscope", "novita",
|
||||
"coding-plan", "alibaba-coding", "qwen-coding", "mimo":
|
||||
"coding-plan", "alibaba-coding", "qwen-coding", "zai", "mimo":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
|
|
@ -61,11 +61,11 @@ func whisperModelID(modelCfg *config.ModelConfig) string {
|
|||
return ""
|
||||
}
|
||||
|
||||
if !supportsWhisperTranscription(modelCfg.Model) {
|
||||
if !supportsWhisperTranscription(modelCfg) {
|
||||
return ""
|
||||
}
|
||||
|
||||
_, modelID := providers.ExtractProtocol(strings.TrimSpace(modelCfg.Model))
|
||||
_, modelID := providers.ExtractProtocol(modelCfg)
|
||||
if strings.Contains(strings.ToLower(modelID), "whisper") {
|
||||
return modelID
|
||||
}
|
||||
|
|
@ -77,14 +77,14 @@ func transcriberFromModelConfig(modelCfg *config.ModelConfig) Transcriber {
|
|||
return nil
|
||||
}
|
||||
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
|
||||
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
|
||||
}
|
||||
if modelID := whisperModelID(modelCfg); modelID != "" {
|
||||
return NewWhisperTranscriber(modelCfg)
|
||||
}
|
||||
if supportsAudioTranscription(modelCfg.Model) {
|
||||
if supportsAudioTranscription(modelCfg) {
|
||||
return NewAudioModelTranscriber(modelCfg)
|
||||
}
|
||||
return nil
|
||||
|
|
@ -95,7 +95,7 @@ func fallbackTranscriberFromModelConfig(modelCfg *config.ModelConfig) Transcribe
|
|||
return nil
|
||||
}
|
||||
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
|
||||
protocol, _ := providers.ExtractProtocol(modelCfg)
|
||||
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
|
||||
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ func NewWhisperTranscriber(modelCfg *config.ModelConfig) *WhisperTranscriber {
|
|||
return nil
|
||||
}
|
||||
|
||||
protocol, modelID := providers.ExtractProtocol(modelCfg.Model)
|
||||
protocol, modelID := providers.ExtractProtocol(modelCfg)
|
||||
if modelID == "" {
|
||||
modelID = strings.TrimSpace(modelCfg.Model)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ func providerFromModelConfig(mc *config.ModelConfig) TTSProvider {
|
|||
return nil
|
||||
}
|
||||
|
||||
protocol, modelID := providers.ExtractProtocol(mc.Model)
|
||||
protocol, modelID := providers.ExtractProtocol(mc)
|
||||
if modelID == "" {
|
||||
modelID = strings.TrimSpace(mc.Model)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import (
|
|||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
|
|
@ -25,6 +26,11 @@ type AuthStore struct {
|
|||
Credentials map[string]*AuthCredential `json:"credentials"`
|
||||
}
|
||||
|
||||
const (
|
||||
providerGoogleAntigravity = "google-antigravity"
|
||||
providerAntigravityAlias = "antigravity"
|
||||
)
|
||||
|
||||
func (c *AuthCredential) IsExpired() bool {
|
||||
if c.ExpiresAt.IsZero() {
|
||||
return false
|
||||
|
|
@ -43,6 +49,125 @@ func authFilePath() string {
|
|||
return filepath.Join(config.GetHome(), "auth.json")
|
||||
}
|
||||
|
||||
func canonicalProvider(provider string) string {
|
||||
normalized := strings.ToLower(strings.TrimSpace(provider))
|
||||
switch normalized {
|
||||
case providerAntigravityAlias:
|
||||
return providerGoogleAntigravity
|
||||
default:
|
||||
return normalized
|
||||
}
|
||||
}
|
||||
|
||||
func cloneCredential(cred *AuthCredential) *AuthCredential {
|
||||
if cred == nil {
|
||||
return nil
|
||||
}
|
||||
cp := *cred
|
||||
return &cp
|
||||
}
|
||||
|
||||
func mergeCredentials(primary, secondary *AuthCredential) *AuthCredential {
|
||||
if primary == nil {
|
||||
return cloneCredential(secondary)
|
||||
}
|
||||
|
||||
merged := *primary
|
||||
if secondary == nil {
|
||||
return &merged
|
||||
}
|
||||
if merged.AccessToken == "" {
|
||||
merged.AccessToken = secondary.AccessToken
|
||||
}
|
||||
if merged.RefreshToken == "" {
|
||||
merged.RefreshToken = secondary.RefreshToken
|
||||
}
|
||||
if merged.AccountID == "" {
|
||||
merged.AccountID = secondary.AccountID
|
||||
}
|
||||
if merged.ExpiresAt.IsZero() {
|
||||
merged.ExpiresAt = secondary.ExpiresAt
|
||||
}
|
||||
if merged.Provider == "" {
|
||||
merged.Provider = secondary.Provider
|
||||
}
|
||||
if merged.AuthMethod == "" {
|
||||
merged.AuthMethod = secondary.AuthMethod
|
||||
}
|
||||
if merged.Email == "" {
|
||||
merged.Email = secondary.Email
|
||||
}
|
||||
if merged.ProjectID == "" {
|
||||
merged.ProjectID = secondary.ProjectID
|
||||
}
|
||||
|
||||
return &merged
|
||||
}
|
||||
|
||||
func shouldPreferCredential(
|
||||
candidate *AuthCredential,
|
||||
candidateCanonical bool,
|
||||
current *AuthCredential,
|
||||
currentCanonical bool,
|
||||
) bool {
|
||||
if candidate == nil {
|
||||
return false
|
||||
}
|
||||
if current == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
switch {
|
||||
case candidate.ExpiresAt.After(current.ExpiresAt):
|
||||
return true
|
||||
case current.ExpiresAt.After(candidate.ExpiresAt):
|
||||
return false
|
||||
case candidateCanonical != currentCanonical:
|
||||
return candidateCanonical
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeStore(store *AuthStore) {
|
||||
if store == nil {
|
||||
return
|
||||
}
|
||||
if store.Credentials == nil {
|
||||
store.Credentials = make(map[string]*AuthCredential)
|
||||
return
|
||||
}
|
||||
|
||||
normalized := make(map[string]*AuthCredential, len(store.Credentials))
|
||||
canonicalFlags := make(map[string]bool, len(store.Credentials))
|
||||
|
||||
for provider, cred := range store.Credentials {
|
||||
normalizedProvider := strings.ToLower(strings.TrimSpace(provider))
|
||||
canonical := canonicalProvider(provider)
|
||||
normalizedCred := cloneCredential(cred)
|
||||
if normalizedCred != nil {
|
||||
normalizedCred.Provider = canonicalProvider(normalizedCred.Provider)
|
||||
if normalizedCred.Provider == "" {
|
||||
normalizedCred.Provider = canonical
|
||||
}
|
||||
}
|
||||
|
||||
current := normalized[canonical]
|
||||
currentCanonical := canonicalFlags[canonical]
|
||||
candidateCanonical := normalizedProvider == canonical
|
||||
|
||||
if shouldPreferCredential(normalizedCred, candidateCanonical, current, currentCanonical) {
|
||||
normalized[canonical] = mergeCredentials(normalizedCred, current)
|
||||
canonicalFlags[canonical] = candidateCanonical
|
||||
continue
|
||||
}
|
||||
|
||||
normalized[canonical] = mergeCredentials(current, normalizedCred)
|
||||
}
|
||||
|
||||
store.Credentials = normalized
|
||||
}
|
||||
|
||||
func LoadStore() (*AuthStore, error) {
|
||||
path := authFilePath()
|
||||
data, err := os.ReadFile(path)
|
||||
|
|
@ -57,9 +182,7 @@ func LoadStore() (*AuthStore, error) {
|
|||
if err := json.Unmarshal(data, &store); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if store.Credentials == nil {
|
||||
store.Credentials = make(map[string]*AuthCredential)
|
||||
}
|
||||
normalizeStore(&store)
|
||||
return &store, nil
|
||||
}
|
||||
|
||||
|
|
@ -79,7 +202,7 @@ func GetCredential(provider string) (*AuthCredential, error) {
|
|||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cred, ok := store.Credentials[provider]
|
||||
cred, ok := store.Credentials[canonicalProvider(provider)]
|
||||
if !ok {
|
||||
return nil, nil
|
||||
}
|
||||
|
|
@ -91,7 +214,17 @@ func SetCredential(provider string, cred *AuthCredential) error {
|
|||
if err != nil {
|
||||
return err
|
||||
}
|
||||
store.Credentials[provider] = cred
|
||||
|
||||
canonical := canonicalProvider(provider)
|
||||
normalized := cloneCredential(cred)
|
||||
if normalized != nil {
|
||||
normalized.Provider = canonicalProvider(normalized.Provider)
|
||||
if normalized.Provider == "" {
|
||||
normalized.Provider = canonical
|
||||
}
|
||||
}
|
||||
|
||||
store.Credentials[canonical] = normalized
|
||||
return SaveStore(store)
|
||||
}
|
||||
|
||||
|
|
@ -100,7 +233,7 @@ func DeleteCredential(provider string) error {
|
|||
if err != nil {
|
||||
return err
|
||||
}
|
||||
delete(store.Credentials, provider)
|
||||
delete(store.Credentials, canonicalProvider(provider))
|
||||
return SaveStore(store)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,12 +1,24 @@
|
|||
package auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
func setTestAuthHome(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv(config.EnvHome, filepath.Join(tmpDir, ".picoclaw"))
|
||||
return tmpDir
|
||||
}
|
||||
|
||||
func TestAuthCredentialIsExpired(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
|
@ -51,10 +63,7 @@ func TestAuthCredentialNeedsRefresh(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestStoreRoundtrip(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
origHome := os.Getenv("HOME")
|
||||
t.Setenv("HOME", tmpDir)
|
||||
defer os.Setenv("HOME", origHome)
|
||||
setTestAuthHome(t)
|
||||
|
||||
cred := &AuthCredential{
|
||||
AccessToken: "test-access-token",
|
||||
|
|
@ -88,10 +97,7 @@ func TestStoreRoundtrip(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestStoreFilePermissions(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
origHome := os.Getenv("HOME")
|
||||
t.Setenv("HOME", tmpDir)
|
||||
defer os.Setenv("HOME", origHome)
|
||||
tmpDir := setTestAuthHome(t)
|
||||
|
||||
cred := &AuthCredential{
|
||||
AccessToken: "secret-token",
|
||||
|
|
@ -108,16 +114,16 @@ func TestStoreFilePermissions(t *testing.T) {
|
|||
t.Fatalf("Stat() error: %v", err)
|
||||
}
|
||||
perm := info.Mode().Perm()
|
||||
if runtime.GOOS == "windows" {
|
||||
return
|
||||
}
|
||||
if perm != 0o600 {
|
||||
t.Errorf("file permissions = %o, want 0600", perm)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStoreMultiProvider(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
origHome := os.Getenv("HOME")
|
||||
t.Setenv("HOME", tmpDir)
|
||||
defer os.Setenv("HOME", origHome)
|
||||
setTestAuthHome(t)
|
||||
|
||||
openaiCred := &AuthCredential{AccessToken: "openai-token", Provider: "openai", AuthMethod: "oauth"}
|
||||
anthropicCred := &AuthCredential{AccessToken: "anthropic-token", Provider: "anthropic", AuthMethod: "token"}
|
||||
|
|
@ -147,10 +153,7 @@ func TestStoreMultiProvider(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestDeleteCredential(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
origHome := os.Getenv("HOME")
|
||||
t.Setenv("HOME", tmpDir)
|
||||
defer os.Setenv("HOME", origHome)
|
||||
setTestAuthHome(t)
|
||||
|
||||
cred := &AuthCredential{AccessToken: "to-delete", Provider: "openai", AuthMethod: "oauth"}
|
||||
if err := SetCredential("openai", cred); err != nil {
|
||||
|
|
@ -171,10 +174,7 @@ func TestDeleteCredential(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestLoadStoreEmpty(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
origHome := os.Getenv("HOME")
|
||||
t.Setenv("HOME", tmpDir)
|
||||
defer os.Setenv("HOME", origHome)
|
||||
setTestAuthHome(t)
|
||||
|
||||
store, err := LoadStore()
|
||||
if err != nil {
|
||||
|
|
@ -187,3 +187,319 @@ func TestLoadStoreEmpty(t *testing.T) {
|
|||
t.Errorf("expected empty credentials, got %d", len(store.Credentials))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetCredentialCanonicalizesLegacyAntigravityProvider(t *testing.T) {
|
||||
tmpDir := setTestAuthHome(t)
|
||||
|
||||
expiresAt := time.Date(2026, 4, 16, 10, 0, 0, 0, time.UTC)
|
||||
store := map[string]any{
|
||||
"credentials": map[string]any{
|
||||
"antigravity": map[string]any{
|
||||
"access_token": "legacy-token",
|
||||
"expires_at": expiresAt.Format(time.RFC3339),
|
||||
"provider": "antigravity",
|
||||
"auth_method": "oauth",
|
||||
"project_id": "project-1",
|
||||
},
|
||||
},
|
||||
}
|
||||
data, err := json.Marshal(store)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal() error: %v", err)
|
||||
}
|
||||
path := filepath.Join(tmpDir, ".picoclaw", "auth.json")
|
||||
err = os.MkdirAll(filepath.Dir(path), 0o755)
|
||||
if err != nil {
|
||||
t.Fatalf("MkdirAll() error: %v", err)
|
||||
}
|
||||
err = os.WriteFile(path, data, 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteFile() error: %v", err)
|
||||
}
|
||||
|
||||
cred, err := GetCredential("google-antigravity")
|
||||
if err != nil {
|
||||
t.Fatalf("GetCredential() error: %v", err)
|
||||
}
|
||||
if cred == nil {
|
||||
t.Fatal("GetCredential() returned nil")
|
||||
}
|
||||
if cred.Provider != "google-antigravity" {
|
||||
t.Fatalf("Provider = %q, want %q", cred.Provider, "google-antigravity")
|
||||
}
|
||||
if !cred.ExpiresAt.Equal(expiresAt) {
|
||||
t.Fatalf("ExpiresAt = %v, want %v", cred.ExpiresAt, expiresAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadStoreMergesAntigravityAliasesPreferringNewerExpiry(t *testing.T) {
|
||||
tmpDir := setTestAuthHome(t)
|
||||
|
||||
legacyExpiry := time.Date(2026, 4, 16, 10, 0, 0, 0, time.UTC)
|
||||
refreshedExpiry := time.Date(2026, 4, 16, 12, 0, 0, 0, time.UTC)
|
||||
store := map[string]any{
|
||||
"credentials": map[string]any{
|
||||
"antigravity": map[string]any{
|
||||
"access_token": "legacy-token",
|
||||
"refresh_token": "legacy-refresh",
|
||||
"expires_at": legacyExpiry.Format(time.RFC3339),
|
||||
"provider": "antigravity",
|
||||
"auth_method": "oauth",
|
||||
"email": "legacy@example.com",
|
||||
},
|
||||
"google-antigravity": map[string]any{
|
||||
"access_token": "fresh-token",
|
||||
"expires_at": refreshedExpiry.Format(time.RFC3339),
|
||||
"provider": "google-antigravity",
|
||||
"auth_method": "oauth",
|
||||
"project_id": "project-2",
|
||||
},
|
||||
},
|
||||
}
|
||||
data, err := json.Marshal(store)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal() error: %v", err)
|
||||
}
|
||||
path := filepath.Join(tmpDir, ".picoclaw", "auth.json")
|
||||
err = os.MkdirAll(filepath.Dir(path), 0o755)
|
||||
if err != nil {
|
||||
t.Fatalf("MkdirAll() error: %v", err)
|
||||
}
|
||||
err = os.WriteFile(path, data, 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteFile() error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := LoadStore()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadStore() error: %v", err)
|
||||
}
|
||||
if len(loaded.Credentials) != 1 {
|
||||
t.Fatalf("credential count = %d, want 1", len(loaded.Credentials))
|
||||
}
|
||||
|
||||
cred := loaded.Credentials["google-antigravity"]
|
||||
if cred == nil {
|
||||
t.Fatal("google-antigravity credential missing")
|
||||
}
|
||||
if cred.AccessToken != "fresh-token" {
|
||||
t.Fatalf("AccessToken = %q, want %q", cred.AccessToken, "fresh-token")
|
||||
}
|
||||
if cred.RefreshToken != "legacy-refresh" {
|
||||
t.Fatalf("RefreshToken = %q, want %q", cred.RefreshToken, "legacy-refresh")
|
||||
}
|
||||
if cred.Email != "legacy@example.com" {
|
||||
t.Fatalf("Email = %q, want %q", cred.Email, "legacy@example.com")
|
||||
}
|
||||
if cred.ProjectID != "project-2" {
|
||||
t.Fatalf("ProjectID = %q, want %q", cred.ProjectID, "project-2")
|
||||
}
|
||||
if !cred.ExpiresAt.Equal(refreshedExpiry) {
|
||||
t.Fatalf("ExpiresAt = %v, want %v", cred.ExpiresAt, refreshedExpiry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadStorePrefersCanonicalKeyWhenExpiryMatchesAlias(t *testing.T) {
|
||||
tmpDir := setTestAuthHome(t)
|
||||
|
||||
expiresAt := time.Date(2026, 4, 16, 12, 0, 0, 0, time.UTC)
|
||||
store := map[string]any{
|
||||
"credentials": map[string]any{
|
||||
"antigravity": map[string]any{
|
||||
"access_token": "legacy-token",
|
||||
"refresh_token": "legacy-refresh",
|
||||
"expires_at": expiresAt.Format(time.RFC3339),
|
||||
"provider": "antigravity",
|
||||
"auth_method": "oauth",
|
||||
"email": "legacy@example.com",
|
||||
},
|
||||
" Google-Antigravity ": map[string]any{
|
||||
"access_token": "fresh-token",
|
||||
"expires_at": expiresAt.Format(time.RFC3339),
|
||||
"provider": " Google-Antigravity ",
|
||||
"auth_method": "oauth",
|
||||
"project_id": "project-2",
|
||||
},
|
||||
},
|
||||
}
|
||||
data, err := json.Marshal(store)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal() error: %v", err)
|
||||
}
|
||||
path := filepath.Join(tmpDir, ".picoclaw", "auth.json")
|
||||
err = os.MkdirAll(filepath.Dir(path), 0o755)
|
||||
if err != nil {
|
||||
t.Fatalf("MkdirAll() error: %v", err)
|
||||
}
|
||||
err = os.WriteFile(path, data, 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteFile() error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := LoadStore()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadStore() error: %v", err)
|
||||
}
|
||||
if len(loaded.Credentials) != 1 {
|
||||
t.Fatalf("credential count = %d, want 1", len(loaded.Credentials))
|
||||
}
|
||||
|
||||
cred := loaded.Credentials["google-antigravity"]
|
||||
if cred == nil {
|
||||
t.Fatal("google-antigravity credential missing")
|
||||
}
|
||||
if cred.AccessToken != "fresh-token" {
|
||||
t.Fatalf("AccessToken = %q, want %q", cred.AccessToken, "fresh-token")
|
||||
}
|
||||
if cred.RefreshToken != "legacy-refresh" {
|
||||
t.Fatalf("RefreshToken = %q, want %q", cred.RefreshToken, "legacy-refresh")
|
||||
}
|
||||
if cred.Email != "legacy@example.com" {
|
||||
t.Fatalf("Email = %q, want %q", cred.Email, "legacy@example.com")
|
||||
}
|
||||
if cred.ProjectID != "project-2" {
|
||||
t.Fatalf("ProjectID = %q, want %q", cred.ProjectID, "project-2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetCredentialReplacesLegacyAntigravityEntry(t *testing.T) {
|
||||
tmpDir := setTestAuthHome(t)
|
||||
|
||||
legacyStore := map[string]any{
|
||||
"credentials": map[string]any{
|
||||
"antigravity": map[string]any{
|
||||
"access_token": "legacy-token",
|
||||
"expires_at": time.Date(2026, 4, 16, 10, 0, 0, 0, time.UTC).Format(time.RFC3339),
|
||||
"provider": "antigravity",
|
||||
"auth_method": "oauth",
|
||||
},
|
||||
},
|
||||
}
|
||||
data, err := json.Marshal(legacyStore)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal() error: %v", err)
|
||||
}
|
||||
path := filepath.Join(tmpDir, ".picoclaw", "auth.json")
|
||||
err = os.MkdirAll(filepath.Dir(path), 0o755)
|
||||
if err != nil {
|
||||
t.Fatalf("MkdirAll() error: %v", err)
|
||||
}
|
||||
err = os.WriteFile(path, data, 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteFile() error: %v", err)
|
||||
}
|
||||
|
||||
refreshedExpiry := time.Date(2026, 4, 16, 12, 30, 0, 0, time.UTC)
|
||||
err = SetCredential("google-antigravity", &AuthCredential{
|
||||
AccessToken: "fresh-token",
|
||||
ExpiresAt: refreshedExpiry,
|
||||
Provider: "google-antigravity",
|
||||
AuthMethod: "oauth",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SetCredential() error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := LoadStore()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadStore() error: %v", err)
|
||||
}
|
||||
if len(loaded.Credentials) != 1 {
|
||||
t.Fatalf("credential count = %d, want 1", len(loaded.Credentials))
|
||||
}
|
||||
|
||||
cred := loaded.Credentials["google-antigravity"]
|
||||
if cred == nil {
|
||||
t.Fatal("google-antigravity credential missing")
|
||||
}
|
||||
if cred.AccessToken != "fresh-token" {
|
||||
t.Fatalf("AccessToken = %q, want %q", cred.AccessToken, "fresh-token")
|
||||
}
|
||||
if !cred.ExpiresAt.Equal(refreshedExpiry) {
|
||||
t.Fatalf("ExpiresAt = %v, want %v", cred.ExpiresAt, refreshedExpiry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteCredentialRemovesLegacyAntigravityAlias(t *testing.T) {
|
||||
tmpDir := setTestAuthHome(t)
|
||||
|
||||
legacyStore := map[string]any{
|
||||
"credentials": map[string]any{
|
||||
"antigravity": map[string]any{
|
||||
"access_token": "legacy-token",
|
||||
"provider": "antigravity",
|
||||
"auth_method": "oauth",
|
||||
},
|
||||
},
|
||||
}
|
||||
data, err := json.Marshal(legacyStore)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal() error: %v", err)
|
||||
}
|
||||
path := filepath.Join(tmpDir, ".picoclaw", "auth.json")
|
||||
err = os.MkdirAll(filepath.Dir(path), 0o755)
|
||||
if err != nil {
|
||||
t.Fatalf("MkdirAll() error: %v", err)
|
||||
}
|
||||
err = os.WriteFile(path, data, 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteFile() error: %v", err)
|
||||
}
|
||||
|
||||
err = DeleteCredential(" google-antigravity ")
|
||||
if err != nil {
|
||||
t.Fatalf("DeleteCredential() error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := LoadStore()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadStore() error: %v", err)
|
||||
}
|
||||
if len(loaded.Credentials) != 0 {
|
||||
t.Fatalf("credential count = %d, want 0", len(loaded.Credentials))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetCredentialCanonicalizesTrimmedMixedCaseProvider(t *testing.T) {
|
||||
setTestAuthHome(t)
|
||||
|
||||
expiresAt := time.Date(2026, 4, 16, 13, 0, 0, 0, time.UTC)
|
||||
if err := SetCredential(" AnTiGrAvItY ", &AuthCredential{
|
||||
AccessToken: "fresh-token",
|
||||
ExpiresAt: expiresAt,
|
||||
Provider: " AnTiGrAvItY ",
|
||||
AuthMethod: "oauth",
|
||||
}); err != nil {
|
||||
t.Fatalf("SetCredential() error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := LoadStore()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadStore() error: %v", err)
|
||||
}
|
||||
if len(loaded.Credentials) != 1 {
|
||||
t.Fatalf("credential count = %d, want 1", len(loaded.Credentials))
|
||||
}
|
||||
|
||||
cred := loaded.Credentials["google-antigravity"]
|
||||
if cred == nil {
|
||||
t.Fatal("google-antigravity credential missing")
|
||||
}
|
||||
if cred.Provider != "google-antigravity" {
|
||||
t.Fatalf("Provider = %q, want %q", cred.Provider, "google-antigravity")
|
||||
}
|
||||
if !cred.ExpiresAt.Equal(expiresAt) {
|
||||
t.Fatalf("ExpiresAt = %v, want %v", cred.ExpiresAt, expiresAt)
|
||||
}
|
||||
|
||||
got, err := GetCredential(" GoOgLe-AnTiGrAvItY ")
|
||||
if err != nil {
|
||||
t.Fatalf("GetCredential() error: %v", err)
|
||||
}
|
||||
if got == nil {
|
||||
t.Fatal("GetCredential() returned nil")
|
||||
}
|
||||
if got.Provider != "google-antigravity" {
|
||||
t.Fatalf("GetCredential provider = %q, want %q", got.Provider, "google-antigravity")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -61,6 +61,15 @@ type OutboundScope struct {
|
|||
Values map[string]string `json:"values,omitempty"`
|
||||
}
|
||||
|
||||
// ContextUsage describes how much of the model's context window the current
|
||||
// session consumes, and how far it is from triggering compression.
|
||||
type ContextUsage struct {
|
||||
UsedTokens int `json:"used_tokens"`
|
||||
TotalTokens int `json:"total_tokens"` // model context window
|
||||
CompressAtTokens int `json:"compress_at_tokens"` // threshold that triggers compression
|
||||
UsedPercent int `json:"used_percent"` // 0-100
|
||||
}
|
||||
|
||||
type OutboundMessage struct {
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
|
|
@ -70,6 +79,7 @@ type OutboundMessage struct {
|
|||
Scope *OutboundScope `json:"scope,omitempty"`
|
||||
Content string `json:"content"`
|
||||
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||
ContextUsage *ContextUsage `json:"context_usage,omitempty"`
|
||||
}
|
||||
|
||||
// MediaPart describes a single media attachment to send.
|
||||
|
|
|
|||
|
|
@ -45,9 +45,12 @@ type DiscordChannel struct {
|
|||
cancel context.CancelFunc
|
||||
typingMu sync.Mutex
|
||||
typingStop map[string]chan struct{} // chatID → stop signal
|
||||
botUserID string // stored for mention checking
|
||||
progress *channels.ToolFeedbackAnimator
|
||||
botUserID string // stored for mention checking
|
||||
bus *bus.MessageBus
|
||||
tts tts.TTSProvider
|
||||
playTTSFn func(context.Context, *discordgo.VoiceConnection, string, uint64)
|
||||
ttsVoiceFn func(string) (*discordgo.VoiceConnection, bool)
|
||||
voiceMu sync.RWMutex
|
||||
voiceSSRC map[string]map[uint32]string // guildID -> ssrc -> userID
|
||||
|
||||
|
|
@ -84,7 +87,7 @@ func NewDiscordChannel(
|
|||
channels.WithReasoningChannelID(bc.ReasoningChannelID),
|
||||
)
|
||||
|
||||
return &DiscordChannel{
|
||||
ch := &DiscordChannel{
|
||||
BaseChannel: base,
|
||||
bc: bc,
|
||||
session: session,
|
||||
|
|
@ -93,7 +96,11 @@ func NewDiscordChannel(
|
|||
typingStop: make(map[string]chan struct{}),
|
||||
bus: bus,
|
||||
voiceSSRC: make(map[string]map[uint32]string),
|
||||
}, nil
|
||||
}
|
||||
ch.playTTSFn = ch.playTTS
|
||||
ch.ttsVoiceFn = ch.voiceConnectionForTTS
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) Start(ctx context.Context) error {
|
||||
|
|
@ -142,6 +149,9 @@ func (c *DiscordChannel) Stop(ctx context.Context) error {
|
|||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
if c.progress != nil {
|
||||
c.progress.StopAll()
|
||||
}
|
||||
|
||||
if err := c.session.Close(); err != nil {
|
||||
return fmt.Errorf("failed to close discord session: %w", err)
|
||||
|
|
@ -164,32 +174,88 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]s
|
|||
return nil, nil
|
||||
}
|
||||
|
||||
if c.tts != nil {
|
||||
if ch, err := c.session.State.Channel(channelID); err == nil && ch.GuildID != "" {
|
||||
if vc, ok := c.session.VoiceConnections[ch.GuildID]; ok && vc != nil {
|
||||
// Cancel any previous TTS playback
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
}
|
||||
ttsCtx, ttsCancel := context.WithCancel(c.ctx)
|
||||
c.ttsPlayID++
|
||||
playID := c.ttsPlayID
|
||||
c.cancelTTS = ttsCancel
|
||||
c.ttsMu.Unlock()
|
||||
|
||||
go c.playTTS(ttsCtx, vc, msg.Content, playID)
|
||||
isToolFeedback := outboundMessageIsToolFeedback(msg)
|
||||
if isToolFeedback {
|
||||
if msgID, handled, err := c.progress.Update(ctx, channelID, msg.Content); handled {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(channelID)
|
||||
c.maybeStartTTS(channelID, msg.Content, isToolFeedback)
|
||||
if !isToolFeedback {
|
||||
if msgIDs, handled := c.FinalizeToolFeedbackMessage(ctx, msg); handled {
|
||||
return msgIDs, nil
|
||||
}
|
||||
}
|
||||
|
||||
msgID, err := c.sendChunk(ctx, channelID, msg.Content, msg.ReplyToMessageID)
|
||||
content := msg.Content
|
||||
if isToolFeedback {
|
||||
content = channels.InitialAnimatedToolFeedbackContent(msg.Content)
|
||||
}
|
||||
msgID, err := c.sendChunk(ctx, channelID, content, msg.ReplyToMessageID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if isToolFeedback {
|
||||
c.RecordToolFeedbackMessage(channelID, msgID, msg.Content)
|
||||
} else if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, channelID, trackedMsgID)
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) maybeStartTTS(channelID, content string, isToolFeedback bool) {
|
||||
if c.tts == nil || isToolFeedback {
|
||||
return
|
||||
}
|
||||
|
||||
voiceFn := c.ttsVoiceFn
|
||||
if voiceFn == nil {
|
||||
voiceFn = c.voiceConnectionForTTS
|
||||
}
|
||||
vc, ok := voiceFn(channelID)
|
||||
if !ok || vc == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Cancel any previous TTS playback.
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
}
|
||||
ttsCtx, ttsCancel := context.WithCancel(c.ctx)
|
||||
c.ttsPlayID++
|
||||
playID := c.ttsPlayID
|
||||
c.cancelTTS = ttsCancel
|
||||
playFn := c.playTTSFn
|
||||
c.ttsMu.Unlock()
|
||||
|
||||
if playFn == nil {
|
||||
playFn = c.playTTS
|
||||
}
|
||||
go playFn(ttsCtx, vc, content, playID)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) voiceConnectionForTTS(channelID string) (*discordgo.VoiceConnection, bool) {
|
||||
if c.session == nil || c.session.State == nil {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
ch, err := c.session.State.Channel(channelID)
|
||||
if err != nil || ch == nil || ch.GuildID == "" {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
vc, ok := c.session.VoiceConnections[ch.GuildID]
|
||||
if !ok || vc == nil {
|
||||
return nil, false
|
||||
}
|
||||
return vc, true
|
||||
}
|
||||
|
||||
// SendMedia implements the channels.MediaSender interface.
|
||||
func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||
if !c.IsRunning() {
|
||||
|
|
@ -200,6 +266,7 @@ func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMes
|
|||
if channelID == "" {
|
||||
return nil, fmt.Errorf("channel ID is empty")
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(channelID)
|
||||
|
||||
store := c.GetMediaStore()
|
||||
if store == nil {
|
||||
|
|
@ -281,6 +348,9 @@ func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMes
|
|||
if r.err != nil {
|
||||
return nil, fmt.Errorf("discord send media: %w", channels.ErrTemporary)
|
||||
}
|
||||
if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, channelID, trackedMsgID)
|
||||
}
|
||||
return []string{r.id}, nil
|
||||
case <-sendCtx.Done():
|
||||
// Close all file readers
|
||||
|
|
@ -295,10 +365,15 @@ func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMes
|
|||
|
||||
// EditMessage implements channels.MessageEditor.
|
||||
func (c *DiscordChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
||||
_, err := c.session.ChannelMessageEdit(chatID, messageID, content)
|
||||
_, err := c.session.ChannelMessageEdit(chatID, messageID, content, discordgo.WithContext(ctx))
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteMessage implements channels.MessageDeleter.
|
||||
func (c *DiscordChannel) DeleteMessage(ctx context.Context, chatID string, messageID string) error {
|
||||
return c.session.ChannelMessageDelete(chatID, messageID, discordgo.WithContext(ctx))
|
||||
}
|
||||
|
||||
// SendPlaceholder implements channels.PlaceholderCapable.
|
||||
// It sends a placeholder message that will later be edited to the actual
|
||||
// response via EditMessage (channels.MessageEditor).
|
||||
|
|
@ -317,6 +392,81 @@ func (c *DiscordChannel) SendPlaceholder(ctx context.Context, chatID string) (st
|
|||
return msg.ID, nil
|
||||
}
|
||||
|
||||
func outboundMessageIsToolFeedback(msg bus.OutboundMessage) bool {
|
||||
if len(msg.Context.Raw) == 0 {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), "tool_feedback")
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) currentToolFeedbackMessage(chatID string) (string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", false
|
||||
}
|
||||
return c.progress.Current(chatID)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) takeToolFeedbackMessage(chatID string) (string, string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", "", false
|
||||
}
|
||||
return c.progress.Take(chatID)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) RecordToolFeedbackMessage(chatID, messageID, content string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Record(chatID, messageID, content)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) ClearToolFeedbackMessage(chatID string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Clear(chatID)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) DismissToolFeedbackMessage(ctx context.Context, chatID string) {
|
||||
msgID, ok := c.currentToolFeedbackMessage(chatID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, chatID, msgID)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) dismissTrackedToolFeedbackMessage(ctx context.Context, chatID, messageID string) {
|
||||
if strings.TrimSpace(chatID) == "" || strings.TrimSpace(messageID) == "" {
|
||||
return
|
||||
}
|
||||
c.ClearToolFeedbackMessage(chatID)
|
||||
_ = c.DeleteMessage(ctx, chatID, messageID)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) finalizeTrackedToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
content string,
|
||||
editFn func(context.Context, string, string, string) error,
|
||||
) ([]string, bool) {
|
||||
msgID, baseContent, ok := c.takeToolFeedbackMessage(chatID)
|
||||
if !ok || editFn == nil {
|
||||
return nil, false
|
||||
}
|
||||
if err := editFn(ctx, chatID, msgID, content); err != nil {
|
||||
c.RecordToolFeedbackMessage(chatID, msgID, baseContent)
|
||||
return nil, false
|
||||
}
|
||||
return []string{msgID}, true
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) FinalizeToolFeedbackMessage(ctx context.Context, msg bus.OutboundMessage) ([]string, bool) {
|
||||
if outboundMessageIsToolFeedback(msg) {
|
||||
return nil, false
|
||||
}
|
||||
return c.finalizeTrackedToolFeedbackMessage(ctx, msg.ChatID, msg.Content, c.EditMessage)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content, replyToID string) (string, error) {
|
||||
// Use the passed ctx for timeout control
|
||||
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||
|
|
|
|||
|
|
@ -1,13 +1,37 @@
|
|||
package discord
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bwmarrin/discordgo"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
)
|
||||
|
||||
type stubTTSProvider struct{}
|
||||
|
||||
func (stubTTSProvider) Name() string { return "stub-tts" }
|
||||
|
||||
func (stubTTSProvider) Synthesize(context.Context, string) (io.ReadCloser, error) {
|
||||
return io.NopCloser(&noopReader{}), nil
|
||||
}
|
||||
|
||||
type noopReader struct{}
|
||||
|
||||
func (*noopReader) Read(p []byte) (int, error) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
func TestApplyDiscordProxy_CustomProxy(t *testing.T) {
|
||||
session, err := discordgo.New("Bot test-token")
|
||||
if err != nil {
|
||||
|
|
@ -89,3 +113,224 @@ func TestApplyDiscordProxy_InvalidProxyURL(t *testing.T) {
|
|||
t.Fatal("applyDiscordProxy() expected error for invalid proxy URL, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSend_NonToolFeedbackDeletesTrackedProgressMessage(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
requests []string
|
||||
)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
mu.Lock()
|
||||
requests = append(requests, r.Method+" "+r.URL.Path)
|
||||
mu.Unlock()
|
||||
|
||||
switch {
|
||||
case r.Method == http.MethodPatch && r.URL.Path == "/channels/chat-1/messages/prog-1":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = io.WriteString(w, `{"id":"prog-1"}`)
|
||||
default:
|
||||
t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
origChannels := discordgo.EndpointChannels
|
||||
discordgo.EndpointChannels = server.URL + "/channels/"
|
||||
defer func() {
|
||||
discordgo.EndpointChannels = origChannels
|
||||
}()
|
||||
|
||||
session, err := discordgo.New("Bot test-token")
|
||||
if err != nil {
|
||||
t.Fatalf("discordgo.New() error: %v", err)
|
||||
}
|
||||
session.Client = server.Client()
|
||||
|
||||
ch := &DiscordChannel{
|
||||
BaseChannel: channels.NewBaseChannel("discord", nil, bus.NewMessageBus(), nil),
|
||||
session: session,
|
||||
ctx: context.Background(),
|
||||
typingStop: make(map[string]chan struct{}),
|
||||
voiceSSRC: make(map[string]map[uint32]string),
|
||||
}
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
ch.SetRunning(true)
|
||||
ch.RecordToolFeedbackMessage("chat-1", "prog-1", "🔧 `read_file`")
|
||||
|
||||
ids, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "chat-1",
|
||||
Content: "final reply",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "discord",
|
||||
ChatID: "chat-1",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Send() error = %v", err)
|
||||
}
|
||||
if got, want := ids, []string{"prog-1"}; !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("Send() ids = %v, want %v", got, want)
|
||||
}
|
||||
if _, ok := ch.currentToolFeedbackMessage("chat-1"); ok {
|
||||
t.Fatal("expected tracked tool feedback message to be cleared")
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
wantRequests := []string{
|
||||
"PATCH /channels/chat-1/messages/prog-1",
|
||||
}
|
||||
if !reflect.DeepEqual(requests, wantRequests) {
|
||||
t.Fatalf("requests = %v, want %v", requests, wantRequests)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditMessage_UsesContextCancellation(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
select {
|
||||
case <-r.Context().Done():
|
||||
return
|
||||
case <-time.After(time.Second):
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = io.WriteString(w, `{"id":"msg-1"}`)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
origChannels := discordgo.EndpointChannels
|
||||
discordgo.EndpointChannels = server.URL + "/channels/"
|
||||
defer func() {
|
||||
discordgo.EndpointChannels = origChannels
|
||||
}()
|
||||
|
||||
session, err := discordgo.New("Bot test-token")
|
||||
if err != nil {
|
||||
t.Fatalf("discordgo.New() error: %v", err)
|
||||
}
|
||||
session.Client = server.Client()
|
||||
|
||||
ch := &DiscordChannel{
|
||||
BaseChannel: channels.NewBaseChannel("discord", nil, bus.NewMessageBus(), nil),
|
||||
session: session,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
start := time.Now()
|
||||
err = ch.EditMessage(ctx, "chat-1", "msg-1", "still running")
|
||||
elapsed := time.Since(start)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected EditMessage() to fail when context times out")
|
||||
}
|
||||
if elapsed >= 500*time.Millisecond {
|
||||
t.Fatalf("EditMessage() ignored context timeout, elapsed=%v", elapsed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_StopsTrackingBeforeEdit(t *testing.T) {
|
||||
ch := &DiscordChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"chat-1",
|
||||
"final reply",
|
||||
func(_ context.Context, chatID, messageID, content string) error {
|
||||
if _, ok := ch.currentToolFeedbackMessage(chatID); ok {
|
||||
t.Fatal("expected tracked tool feedback to be stopped before edit")
|
||||
}
|
||||
if chatID != "chat-1" || messageID != "msg-1" || content != "final reply" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if !handled {
|
||||
t.Fatal("expected finalizeTrackedToolFeedbackMessage to handle tracked message")
|
||||
}
|
||||
if got, want := msgIDs, []string{"msg-1"}; !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("finalizeTrackedToolFeedbackMessage() ids = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSend_NonToolFeedbackFinalizerStillStartsTTS(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
requests []string
|
||||
)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
mu.Lock()
|
||||
requests = append(requests, r.Method+" "+r.URL.Path)
|
||||
mu.Unlock()
|
||||
|
||||
switch {
|
||||
case r.Method == http.MethodPatch && r.URL.Path == "/channels/chat-1/messages/prog-1":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = io.WriteString(w, `{"id":"prog-1"}`)
|
||||
default:
|
||||
t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
origChannels := discordgo.EndpointChannels
|
||||
discordgo.EndpointChannels = server.URL + "/channels/"
|
||||
defer func() {
|
||||
discordgo.EndpointChannels = origChannels
|
||||
}()
|
||||
|
||||
session, err := discordgo.New("Bot test-token")
|
||||
if err != nil {
|
||||
t.Fatalf("discordgo.New() error: %v", err)
|
||||
}
|
||||
session.Client = server.Client()
|
||||
|
||||
ttsStarted := make(chan string, 1)
|
||||
ch := &DiscordChannel{
|
||||
BaseChannel: channels.NewBaseChannel("discord", nil, bus.NewMessageBus(), nil),
|
||||
session: session,
|
||||
ctx: context.Background(),
|
||||
typingStop: make(map[string]chan struct{}),
|
||||
voiceSSRC: make(map[string]map[uint32]string),
|
||||
tts: tts.TTSProvider(stubTTSProvider{}),
|
||||
}
|
||||
ch.ttsVoiceFn = func(string) (*discordgo.VoiceConnection, bool) {
|
||||
return &discordgo.VoiceConnection{}, true
|
||||
}
|
||||
ch.playTTSFn = func(_ context.Context, _ *discordgo.VoiceConnection, text string, _ uint64) {
|
||||
ttsStarted <- text
|
||||
}
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
ch.SetRunning(true)
|
||||
ch.RecordToolFeedbackMessage("chat-1", "prog-1", "🔧 `read_file`")
|
||||
|
||||
ids, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "chat-1",
|
||||
Content: "final reply",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "discord",
|
||||
ChatID: "chat-1",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Send() error = %v", err)
|
||||
}
|
||||
if got, want := ids, []string{"prog-1"}; !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("Send() ids = %v, want %v", got, want)
|
||||
}
|
||||
|
||||
select {
|
||||
case got := <-ttsStarted:
|
||||
if got != "final reply" {
|
||||
t.Fatalf("TTS content = %q, want final reply", got)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("expected TTS to start for finalized tracked tool feedback reply")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -49,6 +49,9 @@ type FeishuChannel struct {
|
|||
|
||||
mu sync.Mutex
|
||||
cancel context.CancelFunc
|
||||
|
||||
progress *channels.ToolFeedbackAnimator
|
||||
deleteMessageFn func(context.Context, string, string) error
|
||||
}
|
||||
|
||||
type cachedMessage struct {
|
||||
|
|
@ -74,6 +77,8 @@ func NewFeishuChannel(bc *config.Channel, cfg *config.FeishuSettings, bus *bus.M
|
|||
tokenCache: tc,
|
||||
client: lark.NewClient(cfg.AppID, cfg.AppSecret.String(), opts...),
|
||||
}
|
||||
ch.deleteMessageFn = ch.deleteMessageAPI
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
ch.SetOwner(ch)
|
||||
return ch, nil
|
||||
}
|
||||
|
|
@ -132,6 +137,9 @@ func (c *FeishuChannel) Stop(ctx context.Context) error {
|
|||
}
|
||||
c.wsClient = nil
|
||||
c.mu.Unlock()
|
||||
if c.progress != nil {
|
||||
c.progress.StopAll()
|
||||
}
|
||||
|
||||
c.SetRunning(false)
|
||||
logger.InfoC("feishu", "Feishu channel stopped")
|
||||
|
|
@ -149,17 +157,55 @@ func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]st
|
|||
return nil, fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
||||
}
|
||||
|
||||
isToolFeedback := outboundMessageIsToolFeedback(msg)
|
||||
if isToolFeedback {
|
||||
if msgID, handled, err := c.progress.Update(ctx, msg.ChatID, msg.Content); handled {
|
||||
if err != nil {
|
||||
// Feishu can fall back to plain text for a previous progress
|
||||
// message, and those messages cannot be patched through the card
|
||||
// edit API. Drop the stale tracker and recreate the progress
|
||||
// message so later tool feedback is not blocked.
|
||||
c.resetTrackedToolFeedbackAfterEditFailure(ctx, msg.ChatID)
|
||||
} else {
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
}
|
||||
} else {
|
||||
if msgIDs, handled := c.FinalizeToolFeedbackMessage(ctx, msg); handled {
|
||||
return msgIDs, nil
|
||||
}
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
|
||||
|
||||
// Build interactive card with markdown content
|
||||
cardContent, err := buildMarkdownCard(msg.Content)
|
||||
sendContent := msg.Content
|
||||
if isToolFeedback {
|
||||
sendContent = channels.InitialAnimatedToolFeedbackContent(msg.Content)
|
||||
}
|
||||
cardContent, err := buildMarkdownCard(sendContent)
|
||||
if err != nil {
|
||||
// If card build fails, fall back to plain text
|
||||
return nil, c.sendText(ctx, msg.ChatID, msg.Content)
|
||||
msgID, sendErr := c.sendText(ctx, msg.ChatID, sendContent)
|
||||
if sendErr != nil {
|
||||
return nil, sendErr
|
||||
}
|
||||
if isToolFeedback {
|
||||
c.RecordToolFeedbackMessage(msg.ChatID, msgID, msg.Content)
|
||||
} else if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
|
||||
// First attempt: try sending as interactive card
|
||||
err = c.sendCard(ctx, msg.ChatID, cardContent)
|
||||
msgID, err := c.sendCard(ctx, msg.ChatID, cardContent)
|
||||
if err == nil {
|
||||
return nil, nil
|
||||
if isToolFeedback {
|
||||
c.RecordToolFeedbackMessage(msg.ChatID, msgID, msg.Content)
|
||||
} else if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
|
||||
// Check if error is due to card table limit (error code 11310)
|
||||
|
|
@ -174,9 +220,14 @@ func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]st
|
|||
})
|
||||
|
||||
// Second attempt: fall back to plain text message
|
||||
textErr := c.sendText(ctx, msg.ChatID, msg.Content)
|
||||
msgID, textErr := c.sendText(ctx, msg.ChatID, sendContent)
|
||||
if textErr == nil {
|
||||
return nil, nil
|
||||
if isToolFeedback {
|
||||
c.RecordToolFeedbackMessage(msg.ChatID, msgID, msg.Content)
|
||||
} else if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
// If text also fails, return the text error
|
||||
return nil, textErr
|
||||
|
|
@ -210,6 +261,31 @@ func (c *FeishuChannel) EditMessage(ctx context.Context, chatID, messageID, cont
|
|||
return nil
|
||||
}
|
||||
|
||||
// DeleteMessage implements channels.MessageDeleter.
|
||||
func (c *FeishuChannel) DeleteMessage(ctx context.Context, chatID, messageID string) error {
|
||||
deleteFn := c.deleteMessageFn
|
||||
if deleteFn == nil {
|
||||
deleteFn = c.deleteMessageAPI
|
||||
}
|
||||
return deleteFn(ctx, chatID, messageID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) deleteMessageAPI(ctx context.Context, chatID, messageID string) error {
|
||||
req := larkim.NewDeleteMessageReqBuilder().
|
||||
MessageId(messageID).
|
||||
Build()
|
||||
|
||||
resp, err := c.client.Im.V1.Message.Delete(ctx, req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("feishu delete: %w", err)
|
||||
}
|
||||
if !resp.Success() {
|
||||
c.invalidateTokenOnAuthError(resp.Code)
|
||||
return fmt.Errorf("feishu delete api error (code=%d msg=%s)", resp.Code, resp.Msg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SendPlaceholder implements channels.PlaceholderCapable.
|
||||
// Sends an interactive card with placeholder text and returns its message ID.
|
||||
func (c *FeishuChannel) SendPlaceholder(ctx context.Context, chatID string) (string, error) {
|
||||
|
|
@ -251,6 +327,93 @@ func (c *FeishuChannel) SendPlaceholder(ctx context.Context, chatID string) (str
|
|||
return "", nil
|
||||
}
|
||||
|
||||
func outboundMessageIsToolFeedback(msg bus.OutboundMessage) bool {
|
||||
if len(msg.Context.Raw) == 0 {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), "tool_feedback")
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) currentToolFeedbackMessage(chatID string) (string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", false
|
||||
}
|
||||
return c.progress.Current(chatID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) takeToolFeedbackMessage(chatID string) (string, string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", "", false
|
||||
}
|
||||
return c.progress.Take(chatID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) RecordToolFeedbackMessage(chatID, messageID, content string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Record(chatID, messageID, content)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) ClearToolFeedbackMessage(chatID string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Clear(chatID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) DismissToolFeedbackMessage(ctx context.Context, chatID string) {
|
||||
msgID, ok := c.currentToolFeedbackMessage(chatID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, chatID, msgID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) resetTrackedToolFeedbackAfterEditFailure(ctx context.Context, chatID string) {
|
||||
msgID, ok := c.currentToolFeedbackMessage(chatID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, chatID, msgID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) dismissTrackedToolFeedbackMessage(ctx context.Context, chatID, messageID string) {
|
||||
if strings.TrimSpace(chatID) == "" || strings.TrimSpace(messageID) == "" {
|
||||
return
|
||||
}
|
||||
c.ClearToolFeedbackMessage(chatID)
|
||||
deleteFn := c.deleteMessageFn
|
||||
if deleteFn == nil {
|
||||
deleteFn = c.deleteMessageAPI
|
||||
}
|
||||
_ = deleteFn(ctx, chatID, messageID)
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) finalizeTrackedToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
content string,
|
||||
editFn func(context.Context, string, string, string) error,
|
||||
) ([]string, bool) {
|
||||
msgID, baseContent, ok := c.takeToolFeedbackMessage(chatID)
|
||||
if !ok || editFn == nil {
|
||||
return nil, false
|
||||
}
|
||||
if err := editFn(ctx, chatID, msgID, content); err != nil {
|
||||
c.RecordToolFeedbackMessage(chatID, msgID, baseContent)
|
||||
return nil, false
|
||||
}
|
||||
return []string{msgID}, true
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) FinalizeToolFeedbackMessage(ctx context.Context, msg bus.OutboundMessage) ([]string, bool) {
|
||||
if outboundMessageIsToolFeedback(msg) {
|
||||
return nil, false
|
||||
}
|
||||
return c.finalizeTrackedToolFeedbackMessage(ctx, msg.ChatID, msg.Content, c.EditMessage)
|
||||
}
|
||||
|
||||
// ReactToMessage implements channels.ReactionCapable.
|
||||
// Adds a reaction (randomly chosen from config) and returns an undo function to remove it.
|
||||
func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID string) (func(), error) {
|
||||
|
|
@ -323,6 +486,7 @@ func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
|||
if !c.IsRunning() {
|
||||
return nil, channels.ErrNotRunning
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
|
||||
|
||||
if msg.ChatID == "" {
|
||||
return nil, fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
||||
|
|
@ -339,6 +503,10 @@ func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
|||
}
|
||||
}
|
||||
|
||||
if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
|
|
@ -801,7 +969,7 @@ func appendMediaTags(content, messageType string, mediaRefs []string) string {
|
|||
}
|
||||
|
||||
// sendCard sends an interactive card message to a chat.
|
||||
func (c *FeishuChannel) sendCard(ctx context.Context, chatID, cardContent string) error {
|
||||
func (c *FeishuChannel) sendCard(ctx context.Context, chatID, cardContent string) (string, error) {
|
||||
req := larkim.NewCreateMessageReqBuilder().
|
||||
ReceiveIdType(larkim.ReceiveIdTypeChatId).
|
||||
Body(larkim.NewCreateMessageReqBodyBuilder().
|
||||
|
|
@ -813,23 +981,26 @@ func (c *FeishuChannel) sendCard(ctx context.Context, chatID, cardContent string
|
|||
|
||||
resp, err := c.client.Im.V1.Message.Create(ctx, req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("feishu send card: %w", channels.ErrTemporary)
|
||||
return "", fmt.Errorf("feishu send card: %w", channels.ErrTemporary)
|
||||
}
|
||||
|
||||
if !resp.Success() {
|
||||
c.invalidateTokenOnAuthError(resp.Code)
|
||||
return fmt.Errorf("feishu api error (code=%d msg=%s): %w", resp.Code, resp.Msg, channels.ErrTemporary)
|
||||
return "", fmt.Errorf("feishu api error (code=%d msg=%s): %w", resp.Code, resp.Msg, channels.ErrTemporary)
|
||||
}
|
||||
|
||||
logger.DebugCF("feishu", "Feishu card message sent", map[string]any{
|
||||
"chat_id": chatID,
|
||||
})
|
||||
|
||||
return nil
|
||||
if resp.Data != nil && resp.Data.MessageId != nil {
|
||||
return *resp.Data.MessageId, nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// sendText sends a plain text message to a chat (fallback when card fails).
|
||||
func (c *FeishuChannel) sendText(ctx context.Context, chatID, text string) error {
|
||||
func (c *FeishuChannel) sendText(ctx context.Context, chatID, text string) (string, error) {
|
||||
content, _ := json.Marshal(map[string]string{"text": text})
|
||||
|
||||
req := larkim.NewCreateMessageReqBuilder().
|
||||
|
|
@ -843,18 +1014,21 @@ func (c *FeishuChannel) sendText(ctx context.Context, chatID, text string) error
|
|||
|
||||
resp, err := c.client.Im.V1.Message.Create(ctx, req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("feishu send text: %w", channels.ErrTemporary)
|
||||
return "", fmt.Errorf("feishu send text: %w", channels.ErrTemporary)
|
||||
}
|
||||
|
||||
if !resp.Success() {
|
||||
return fmt.Errorf("feishu text api error (code=%d msg=%s): %w", resp.Code, resp.Msg, channels.ErrTemporary)
|
||||
return "", fmt.Errorf("feishu text api error (code=%d msg=%s): %w", resp.Code, resp.Msg, channels.ErrTemporary)
|
||||
}
|
||||
|
||||
logger.DebugCF("feishu", "Feishu text message sent (fallback)", map[string]any{
|
||||
"chat_id": chatID,
|
||||
})
|
||||
|
||||
return nil
|
||||
if resp.Data != nil && resp.Data.MessageId != nil {
|
||||
return *resp.Data.MessageId, nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// sendImage uploads an image and sends it as a message.
|
||||
|
|
|
|||
|
|
@ -3,9 +3,13 @@
|
|||
package feishu
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
)
|
||||
|
||||
func TestExtractContent(t *testing.T) {
|
||||
|
|
@ -279,3 +283,110 @@ func TestExtractFeishuSenderID(t *testing.T) {
|
|||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_ClearAfterSuccessfulEdit(t *testing.T) {
|
||||
ch := &FeishuChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"chat-1",
|
||||
"final reply",
|
||||
func(_ context.Context, chatID, messageID, content string) error {
|
||||
if chatID != "chat-1" || messageID != "msg-1" || content != "final reply" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if !handled {
|
||||
t.Fatal("expected finalizeTrackedToolFeedbackMessage to handle tracked message")
|
||||
}
|
||||
if len(msgIDs) != 1 || msgIDs[0] != "msg-1" {
|
||||
t.Fatalf("unexpected msgIDs: %v", msgIDs)
|
||||
}
|
||||
if _, ok := ch.currentToolFeedbackMessage("chat-1"); ok {
|
||||
t.Fatal("expected tracked tool feedback to be cleared after successful edit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_StopsTrackingBeforeEdit(t *testing.T) {
|
||||
ch := &FeishuChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"chat-1",
|
||||
"final reply",
|
||||
func(_ context.Context, chatID, messageID, content string) error {
|
||||
if _, ok := ch.currentToolFeedbackMessage(chatID); ok {
|
||||
t.Fatal("expected tracked tool feedback to be stopped before edit")
|
||||
}
|
||||
if chatID != "chat-1" || messageID != "msg-1" || content != "final reply" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if !handled {
|
||||
t.Fatal("expected finalizeTrackedToolFeedbackMessage to handle tracked message")
|
||||
}
|
||||
if len(msgIDs) != 1 || msgIDs[0] != "msg-1" {
|
||||
t.Fatalf("unexpected msgIDs: %v", msgIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_EditFailureKeepsTrackedMessage(t *testing.T) {
|
||||
ch := &FeishuChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"chat-1",
|
||||
"final reply",
|
||||
func(context.Context, string, string, string) error {
|
||||
return errors.New("edit failed")
|
||||
},
|
||||
)
|
||||
if handled {
|
||||
t.Fatal("expected finalizeTrackedToolFeedbackMessage to report unhandled on edit failure")
|
||||
}
|
||||
if len(msgIDs) != 0 {
|
||||
t.Fatalf("unexpected msgIDs: %v", msgIDs)
|
||||
}
|
||||
if msgID, ok := ch.currentToolFeedbackMessage("chat-1"); !ok || msgID != "msg-1" {
|
||||
t.Fatalf("expected tracked tool feedback to remain after failed edit, got (%q, %v)", msgID, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetTrackedToolFeedbackAfterEditFailure_DismissesTrackedMessage(t *testing.T) {
|
||||
var (
|
||||
deletedChatID string
|
||||
deletedMsgID string
|
||||
)
|
||||
|
||||
ch := &FeishuChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
deleteMessageFn: func(_ context.Context, chatID, messageID string) error {
|
||||
deletedChatID = chatID
|
||||
deletedMsgID = messageID
|
||||
return nil
|
||||
},
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
ch.resetTrackedToolFeedbackAfterEditFailure(context.Background(), "chat-1")
|
||||
|
||||
if deletedChatID != "chat-1" || deletedMsgID != "msg-1" {
|
||||
t.Fatalf("unexpected delete target: chat=%q msg=%q", deletedChatID, deletedMsgID)
|
||||
}
|
||||
if _, ok := ch.currentToolFeedbackMessage("chat-1"); ok {
|
||||
t.Fatal("expected tracked tool feedback to be cleared after edit failure reset")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import (
|
|||
"net"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
|
|
@ -25,6 +26,7 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/health"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
const (
|
||||
|
|
@ -96,6 +98,23 @@ type Manager struct {
|
|||
channelHashes map[string]string // channel name → config hash
|
||||
}
|
||||
|
||||
type toolFeedbackMessageTracker interface {
|
||||
RecordToolFeedbackMessage(chatID, messageID, content string)
|
||||
ClearToolFeedbackMessage(chatID string)
|
||||
}
|
||||
|
||||
type toolFeedbackMessageCleaner interface {
|
||||
DismissToolFeedbackMessage(ctx context.Context, chatID string)
|
||||
}
|
||||
|
||||
type toolFeedbackMessageTargetResolver interface {
|
||||
ToolFeedbackMessageChatID(chatID string, outboundCtx *bus.InboundContext) string
|
||||
}
|
||||
|
||||
type toolFeedbackMessageContentPreparer interface {
|
||||
PrepareToolFeedbackMessageContent(content string) string
|
||||
}
|
||||
|
||||
type asyncTask struct {
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
|
@ -108,6 +127,13 @@ func outboundMessageChatID(msg bus.OutboundMessage) string {
|
|||
return msg.ChatID
|
||||
}
|
||||
|
||||
func outboundMessageIsToolFeedback(msg bus.OutboundMessage) bool {
|
||||
if len(msg.Context.Raw) == 0 {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), "tool_feedback")
|
||||
}
|
||||
|
||||
func outboundMediaChannel(msg bus.OutboundMediaMessage) string {
|
||||
return msg.Context.Channel
|
||||
}
|
||||
|
|
@ -116,6 +142,47 @@ func outboundMediaChatID(msg bus.OutboundMediaMessage) string {
|
|||
return msg.ChatID
|
||||
}
|
||||
|
||||
func trackedToolFeedbackMessageChatID(ch Channel, chatID string, outboundCtx *bus.InboundContext) string {
|
||||
if resolver, ok := ch.(toolFeedbackMessageTargetResolver); ok {
|
||||
if resolved := strings.TrimSpace(resolver.ToolFeedbackMessageChatID(chatID, outboundCtx)); resolved != "" {
|
||||
return resolved
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(chatID)
|
||||
}
|
||||
|
||||
func dismissTrackedToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
ch Channel,
|
||||
chatID string,
|
||||
outboundCtx *bus.InboundContext,
|
||||
) {
|
||||
trackedChatID := trackedToolFeedbackMessageChatID(ch, chatID, outboundCtx)
|
||||
if trackedChatID == "" {
|
||||
return
|
||||
}
|
||||
if cleaner, ok := ch.(toolFeedbackMessageCleaner); ok {
|
||||
cleaner.DismissToolFeedbackMessage(ctx, trackedChatID)
|
||||
return
|
||||
}
|
||||
if tracker, ok := ch.(toolFeedbackMessageTracker); ok {
|
||||
tracker.ClearToolFeedbackMessage(trackedChatID)
|
||||
}
|
||||
}
|
||||
|
||||
func prepareToolFeedbackMessageContent(ch Channel, content string) string {
|
||||
prepared := strings.TrimSpace(content)
|
||||
if prepared == "" {
|
||||
return ""
|
||||
}
|
||||
if preparer, ok := ch.(toolFeedbackMessageContentPreparer); ok {
|
||||
if candidate := strings.TrimSpace(preparer.PrepareToolFeedbackMessageContent(prepared)); candidate != "" {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
return prepared
|
||||
}
|
||||
|
||||
// RecordPlaceholder registers a placeholder message for later editing.
|
||||
// Implements PlaceholderRecorder.
|
||||
func (m *Manager) RecordPlaceholder(channel, chatID, placeholderID string) {
|
||||
|
|
@ -196,7 +263,19 @@ func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMess
|
|||
}
|
||||
}
|
||||
|
||||
// 3. If a stream already finalized this message, delete the placeholder and skip send
|
||||
isToolFeedback := outboundMessageIsToolFeedback(msg)
|
||||
|
||||
// 3. If a stream already finalized this chat, stale tool feedback must be
|
||||
// dropped without consuming the final-response marker. Streaming finalization
|
||||
// bypasses the worker queue, so older queued feedback can arrive before the
|
||||
// normal final outbound message that cleans up the marker and placeholder.
|
||||
if isToolFeedback {
|
||||
if _, loaded := m.streamActive.Load(key); loaded {
|
||||
return nil, true
|
||||
}
|
||||
}
|
||||
|
||||
// 4. If a stream already finalized this message, delete the placeholder and skip send
|
||||
if _, loaded := m.streamActive.LoadAndDelete(key); loaded {
|
||||
if v, loaded := m.placeholders.LoadAndDelete(key); loaded {
|
||||
if entry, ok := v.(placeholderEntry); ok && entry.id != "" {
|
||||
|
|
@ -208,14 +287,29 @@ func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMess
|
|||
}
|
||||
}
|
||||
}
|
||||
if !isToolFeedback {
|
||||
dismissTrackedToolFeedbackMessage(ctx, ch, chatID, &msg.Context)
|
||||
}
|
||||
return nil, true
|
||||
}
|
||||
|
||||
// 4. Try editing placeholder
|
||||
// 5. Try editing placeholder
|
||||
if v, loaded := m.placeholders.LoadAndDelete(key); loaded {
|
||||
if entry, ok := v.(placeholderEntry); ok && entry.id != "" {
|
||||
if editor, ok := ch.(MessageEditor); ok {
|
||||
if err := editor.EditMessage(ctx, chatID, entry.id, msg.Content); err == nil {
|
||||
content := msg.Content
|
||||
trackedContent := msg.Content
|
||||
if isToolFeedback {
|
||||
trackedContent = prepareToolFeedbackMessageContent(ch, msg.Content)
|
||||
content = InitialAnimatedToolFeedbackContent(trackedContent)
|
||||
}
|
||||
if err := editor.EditMessage(ctx, chatID, entry.id, content); err == nil {
|
||||
trackedChatID := trackedToolFeedbackMessageChatID(ch, chatID, &msg.Context)
|
||||
if tracker, ok := ch.(toolFeedbackMessageTracker); ok && isToolFeedback {
|
||||
tracker.RecordToolFeedbackMessage(trackedChatID, entry.id, trackedContent)
|
||||
} else if !isToolFeedback {
|
||||
dismissTrackedToolFeedbackMessage(ctx, ch, chatID, &msg.Context)
|
||||
}
|
||||
return []string{entry.id}, true
|
||||
}
|
||||
// edit failed → fall through to normal Send
|
||||
|
|
@ -312,22 +406,35 @@ func (m *Manager) GetStreamer(ctx context.Context, channelName, chatID string) (
|
|||
// Mark streamActive on Finalize so preSend knows to clean up the placeholder
|
||||
key := channelName + ":" + chatID
|
||||
return &finalizeHookStreamer{
|
||||
Streamer: streamer,
|
||||
onFinalize: func() { m.streamActive.Store(key, true) },
|
||||
Streamer: streamer,
|
||||
onFinalize: func(finalizeCtx context.Context) {
|
||||
dismissTrackedToolFeedbackMessage(
|
||||
finalizeCtx,
|
||||
ch,
|
||||
chatID,
|
||||
&bus.InboundContext{
|
||||
Channel: channelName,
|
||||
ChatID: chatID,
|
||||
},
|
||||
)
|
||||
m.streamActive.Store(key, true)
|
||||
},
|
||||
}, true
|
||||
}
|
||||
|
||||
// finalizeHookStreamer wraps a Streamer to run a hook on Finalize.
|
||||
type finalizeHookStreamer struct {
|
||||
Streamer
|
||||
onFinalize func()
|
||||
onFinalize func(context.Context)
|
||||
}
|
||||
|
||||
func (s *finalizeHookStreamer) Finalize(ctx context.Context, content string) error {
|
||||
if err := s.Streamer.Finalize(ctx, content); err != nil {
|
||||
return err
|
||||
}
|
||||
s.onFinalize()
|
||||
if s.onFinalize != nil {
|
||||
s.onFinalize(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
|
@ -769,18 +876,21 @@ func (m *Manager) runWorker(ctx context.Context, name string, w *channelWorker)
|
|||
// Collect all message chunks to send
|
||||
var chunks []string
|
||||
|
||||
// Step 1: Try marker-based splitting if enabled
|
||||
if m.config != nil && m.config.Agents.Defaults.SplitOnMarker {
|
||||
// Step 1: Try marker-based splitting if enabled.
|
||||
// Tool feedback must stay a single message, so it skips marker splitting.
|
||||
if m.config != nil && m.config.Agents.Defaults.SplitOnMarker && !outboundMessageIsToolFeedback(msg) {
|
||||
if markerChunks := SplitByMarker(msg.Content); len(markerChunks) > 1 {
|
||||
for _, chunk := range markerChunks {
|
||||
chunks = append(chunks, splitByLength(chunk, maxLen)...)
|
||||
chunkMsg := msg
|
||||
chunkMsg.Content = chunk
|
||||
chunks = append(chunks, splitOutboundMessageContent(chunkMsg, maxLen)...)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Step 2: Fallback to length-based splitting if no chunks from marker
|
||||
if len(chunks) == 0 {
|
||||
chunks = splitByLength(msg.Content, maxLen)
|
||||
chunks = splitOutboundMessageContent(msg, maxLen)
|
||||
}
|
||||
|
||||
// Step 3: Send all chunks
|
||||
|
|
@ -795,12 +905,25 @@ func (m *Manager) runWorker(ctx context.Context, name string, w *channelWorker)
|
|||
}
|
||||
}
|
||||
|
||||
// splitByLength splits content by maxLen if needed, otherwise returns single chunk.
|
||||
func splitByLength(content string, maxLen int) []string {
|
||||
if maxLen > 0 && len([]rune(content)) > maxLen {
|
||||
return SplitMessage(content, maxLen)
|
||||
// splitOutboundMessageContent splits regular outbound content by maxLen, but
|
||||
// keeps tool feedback in a single message by truncating the explanation body.
|
||||
func splitOutboundMessageContent(msg bus.OutboundMessage, maxLen int) []string {
|
||||
if maxLen > 0 {
|
||||
if outboundMessageIsToolFeedback(msg) {
|
||||
animationSafeLen := maxLen - MaxToolFeedbackAnimationFrameLength()
|
||||
if animationSafeLen <= 0 {
|
||||
animationSafeLen = maxLen
|
||||
}
|
||||
if len([]rune(msg.Content)) > animationSafeLen {
|
||||
return []string{utils.FitToolFeedbackMessage(msg.Content, animationSafeLen)}
|
||||
}
|
||||
return []string{msg.Content}
|
||||
}
|
||||
if len([]rune(msg.Content)) > maxLen {
|
||||
return SplitMessage(msg.Content, maxLen)
|
||||
}
|
||||
}
|
||||
return []string{content}
|
||||
return []string{msg.Content}
|
||||
}
|
||||
|
||||
// sendWithRetry sends a message through the channel with rate limiting and
|
||||
|
|
@ -1264,13 +1387,16 @@ func (m *Manager) SendMessage(ctx context.Context, msg bus.OutboundMessage) erro
|
|||
if mlp, ok := w.ch.(MessageLengthProvider); ok {
|
||||
maxLen = mlp.MaxMessageLength()
|
||||
}
|
||||
if maxLen > 0 && len([]rune(msg.Content)) > maxLen {
|
||||
for _, chunk := range SplitMessage(msg.Content, maxLen) {
|
||||
if chunks := splitOutboundMessageContent(msg, maxLen); len(chunks) > 1 {
|
||||
for _, chunk := range chunks {
|
||||
chunkMsg := msg
|
||||
chunkMsg.Content = chunk
|
||||
m.sendWithRetry(ctx, channelName, w, chunkMsg)
|
||||
}
|
||||
} else {
|
||||
if len(chunks) == 1 {
|
||||
msg.Content = chunks[0]
|
||||
}
|
||||
m.sendWithRetry(ctx, channelName, w, msg)
|
||||
}
|
||||
return nil
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ import (
|
|||
"golang.org/x/time/rate"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
// mockChannel is a test double that delegates Send to a configurable function.
|
||||
|
|
@ -76,8 +78,9 @@ func (m *mockMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaM
|
|||
|
||||
type mockDeletingMediaChannel struct {
|
||||
mockMediaChannel
|
||||
deleteCalls int
|
||||
lastDeleted struct {
|
||||
deleteCalls int
|
||||
dismissedChatID string
|
||||
lastDeleted struct {
|
||||
chatID string
|
||||
messageID string
|
||||
}
|
||||
|
|
@ -94,6 +97,48 @@ func (m *mockDeletingMediaChannel) DeleteMessage(
|
|||
return nil
|
||||
}
|
||||
|
||||
func (m *mockDeletingMediaChannel) DismissToolFeedbackMessage(_ context.Context, chatID string) {
|
||||
m.dismissedChatID = chatID
|
||||
}
|
||||
|
||||
type mockStreamer struct {
|
||||
finalizeFn func(context.Context, string) error
|
||||
}
|
||||
|
||||
func (m *mockStreamer) Update(context.Context, string) error { return nil }
|
||||
|
||||
func (m *mockStreamer) Finalize(ctx context.Context, content string) error {
|
||||
if m.finalizeFn != nil {
|
||||
return m.finalizeFn(ctx, content)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockStreamer) Cancel(context.Context) {}
|
||||
|
||||
type mockStreamingChannel struct {
|
||||
mockMessageEditor
|
||||
streamer Streamer
|
||||
resolveChatIDFn func(chatID string, outboundCtx *bus.InboundContext) string
|
||||
}
|
||||
|
||||
func (m *mockStreamingChannel) BeginStream(context.Context, string) (Streamer, error) {
|
||||
if m.streamer == nil {
|
||||
return nil, errors.New("missing streamer")
|
||||
}
|
||||
return m.streamer, nil
|
||||
}
|
||||
|
||||
func (m *mockStreamingChannel) ToolFeedbackMessageChatID(
|
||||
chatID string,
|
||||
outboundCtx *bus.InboundContext,
|
||||
) string {
|
||||
if m.resolveChatIDFn != nil {
|
||||
return m.resolveChatIDFn(chatID, outboundCtx)
|
||||
}
|
||||
return chatID
|
||||
}
|
||||
|
||||
// newTestManager creates a minimal Manager suitable for unit tests.
|
||||
func newTestManager() *Manager {
|
||||
return &Manager{
|
||||
|
|
@ -715,13 +760,72 @@ func TestSendWithRetry_ExponentialBackoff(t *testing.T) {
|
|||
// mockMessageEditor is a channel that supports MessageEditor.
|
||||
type mockMessageEditor struct {
|
||||
mockChannel
|
||||
editFn func(ctx context.Context, chatID, messageID, content string) error
|
||||
editFn func(ctx context.Context, chatID, messageID, content string) error
|
||||
finalizeFn func(ctx context.Context, msg bus.OutboundMessage) ([]string, bool)
|
||||
finalizeCalled bool
|
||||
recordedChatID string
|
||||
recordedMessageID string
|
||||
recordedContent string
|
||||
clearedChatID string
|
||||
dismissedChatID string
|
||||
}
|
||||
|
||||
func (m *mockMessageEditor) EditMessage(ctx context.Context, chatID, messageID, content string) error {
|
||||
return m.editFn(ctx, chatID, messageID, content)
|
||||
}
|
||||
|
||||
func (m *mockMessageEditor) RecordToolFeedbackMessage(chatID, messageID, content string) {
|
||||
m.recordedChatID = chatID
|
||||
m.recordedMessageID = messageID
|
||||
m.recordedContent = content
|
||||
}
|
||||
|
||||
func (m *mockMessageEditor) ClearToolFeedbackMessage(chatID string) {
|
||||
m.clearedChatID = chatID
|
||||
}
|
||||
|
||||
func (m *mockMessageEditor) DismissToolFeedbackMessage(_ context.Context, chatID string) {
|
||||
m.dismissedChatID = chatID
|
||||
}
|
||||
|
||||
func (m *mockMessageEditor) FinalizeToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
msg bus.OutboundMessage,
|
||||
) ([]string, bool) {
|
||||
m.finalizeCalled = true
|
||||
if m.finalizeFn == nil {
|
||||
return nil, false
|
||||
}
|
||||
return m.finalizeFn(ctx, msg)
|
||||
}
|
||||
|
||||
type mockResolvedToolFeedbackEditor struct {
|
||||
mockMessageEditor
|
||||
resolveChatIDFn func(chatID string, outboundCtx *bus.InboundContext) string
|
||||
}
|
||||
|
||||
func (m *mockResolvedToolFeedbackEditor) ToolFeedbackMessageChatID(
|
||||
chatID string,
|
||||
outboundCtx *bus.InboundContext,
|
||||
) string {
|
||||
if m.resolveChatIDFn != nil {
|
||||
return m.resolveChatIDFn(chatID, outboundCtx)
|
||||
}
|
||||
return chatID
|
||||
}
|
||||
|
||||
type mockPreparedToolFeedbackEditor struct {
|
||||
mockMessageEditor
|
||||
prepareFn func(content string) string
|
||||
}
|
||||
|
||||
func (m *mockPreparedToolFeedbackEditor) PrepareToolFeedbackMessageContent(content string) string {
|
||||
if m.prepareFn != nil {
|
||||
return m.prepareFn(content)
|
||||
}
|
||||
return content
|
||||
}
|
||||
|
||||
func TestPreSend_PlaceholderEditSuccess(t *testing.T) {
|
||||
m := newTestManager()
|
||||
var sendCalled bool
|
||||
|
|
@ -766,6 +870,539 @@ func TestPreSend_PlaceholderEditSuccess(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestPreSend_ToolFeedbackPlaceholderEditRecordsTrackedMessage(t *testing.T) {
|
||||
m := newTestManager()
|
||||
|
||||
ch := &mockMessageEditor{
|
||||
editFn: func(_ context.Context, chatID, messageID, content string) error {
|
||||
if chatID != "123" || messageID != "456" || content != "hello" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
m.RecordPlaceholder("test", "123", "456")
|
||||
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "hello",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||
if !edited {
|
||||
t.Fatal("expected preSend to edit placeholder")
|
||||
}
|
||||
if ch.recordedChatID != "123" || ch.recordedMessageID != "456" {
|
||||
t.Fatalf("expected tracked message 123/456, got %q/%q", ch.recordedChatID, ch.recordedMessageID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_ToolFeedbackPlaceholderEditUsesResolvedTrackedChatID(t *testing.T) {
|
||||
m := newTestManager()
|
||||
|
||||
ch := &mockResolvedToolFeedbackEditor{
|
||||
mockMessageEditor: mockMessageEditor{
|
||||
editFn: func(_ context.Context, chatID, messageID, content string) error {
|
||||
if chatID != "-100123" || messageID != "456" || content != "hello" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
},
|
||||
resolveChatIDFn: func(chatID string, outboundCtx *bus.InboundContext) string {
|
||||
if chatID != "-100123" {
|
||||
t.Fatalf("expected raw chat ID, got %q", chatID)
|
||||
}
|
||||
if outboundCtx == nil || outboundCtx.TopicID != "42" {
|
||||
t.Fatalf("expected topic-aware outbound context, got %+v", outboundCtx)
|
||||
}
|
||||
return chatID + "/" + outboundCtx.TopicID
|
||||
},
|
||||
}
|
||||
|
||||
m.RecordPlaceholder("test", "-100123", "456")
|
||||
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "-100123",
|
||||
Content: "hello",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "-100123",
|
||||
TopicID: "42",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||
if !edited {
|
||||
t.Fatal("expected preSend to edit placeholder")
|
||||
}
|
||||
if ch.recordedChatID != "-100123/42" || ch.recordedMessageID != "456" {
|
||||
t.Fatalf("expected resolved tracked message -100123/42/456, got %q/%q",
|
||||
ch.recordedChatID, ch.recordedMessageID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_ToolFeedbackPlaceholderEditUsesPreparedContent(t *testing.T) {
|
||||
m := newTestManager()
|
||||
|
||||
const rawContent = "🔧 `read_file`\n" + "<raw>"
|
||||
const preparedContent = "🔧 `read_file`\n<raw>"
|
||||
|
||||
ch := &mockPreparedToolFeedbackEditor{
|
||||
mockMessageEditor: mockMessageEditor{
|
||||
editFn: func(_ context.Context, chatID, messageID, content string) error {
|
||||
if chatID != "123" || messageID != "456" {
|
||||
t.Fatalf("unexpected edit target: %s/%s", chatID, messageID)
|
||||
}
|
||||
if content != InitialAnimatedToolFeedbackContent(preparedContent) {
|
||||
t.Fatalf("unexpected prepared content: %q", content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
},
|
||||
prepareFn: func(content string) string {
|
||||
if content != rawContent {
|
||||
t.Fatalf("unexpected raw tool feedback: %q", content)
|
||||
}
|
||||
return preparedContent
|
||||
},
|
||||
}
|
||||
|
||||
m.RecordPlaceholder("test", "123", "456")
|
||||
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: rawContent,
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||
if !edited {
|
||||
t.Fatal("expected preSend to edit placeholder")
|
||||
}
|
||||
if ch.recordedContent != preparedContent {
|
||||
t.Fatalf("expected tracked content %q, got %q", preparedContent, ch.recordedContent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_NonToolFeedbackLeavesTrackedMessageForChannelSend(t *testing.T) {
|
||||
m := newTestManager()
|
||||
ch := &mockMessageEditor{}
|
||||
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "final reply",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
},
|
||||
})
|
||||
|
||||
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||
if edited {
|
||||
t.Fatal("expected preSend to fall through when no placeholder exists")
|
||||
}
|
||||
if ch.dismissedChatID != "" {
|
||||
t.Fatalf("expected tracked tool feedback cleanup to be deferred to channel send, got %q", ch.dismissedChatID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_NonToolFeedbackDefersTrackedMessageFinalizationToChannelSend(t *testing.T) {
|
||||
m := newTestManager()
|
||||
ch := &mockMessageEditor{
|
||||
finalizeFn: func(_ context.Context, msg bus.OutboundMessage) ([]string, bool) {
|
||||
if msg.ChatID != "123" || msg.Content != "final reply" {
|
||||
t.Fatalf("unexpected finalize msg: %+v", msg)
|
||||
}
|
||||
return []string{"tool-msg-1"}, true
|
||||
},
|
||||
}
|
||||
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "final reply",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
},
|
||||
})
|
||||
|
||||
msgIDs, handled := m.preSend(context.Background(), "test", msg, ch)
|
||||
if handled {
|
||||
t.Fatalf("expected preSend to defer to channel Send, got msgIDs=%v", msgIDs)
|
||||
}
|
||||
if len(msgIDs) != 0 {
|
||||
t.Fatalf("expected no msgIDs from preSend, got %v", msgIDs)
|
||||
}
|
||||
if ch.dismissedChatID != "" {
|
||||
t.Fatalf("expected tracked cleanup to remain in channel Send, got %q", ch.dismissedChatID)
|
||||
}
|
||||
if ch.finalizeCalled {
|
||||
t.Fatal("expected preSend to skip channel tool feedback finalization")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_StaleToolFeedbackDoesNotConsumeStreamActiveMarker(t *testing.T) {
|
||||
m := newTestManager()
|
||||
m.streamActive.Store("test:123", true)
|
||||
m.RecordPlaceholder("test", "123", "placeholder-1")
|
||||
|
||||
var editedContent string
|
||||
ch := &mockMessageEditor{
|
||||
editFn: func(_ context.Context, chatID, messageID, content string) error {
|
||||
if chatID != "123" || messageID != "placeholder-1" {
|
||||
t.Fatalf("unexpected edit target: %s/%s", chatID, messageID)
|
||||
}
|
||||
editedContent = content
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
toolFeedback := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "🔧 `read_file`\nReading config",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
msgIDs, handled := m.preSend(context.Background(), "test", toolFeedback, ch)
|
||||
if !handled {
|
||||
t.Fatal("expected stale tool feedback to be dropped after stream finalize")
|
||||
}
|
||||
if len(msgIDs) != 0 {
|
||||
t.Fatalf("expected no delivered message IDs for stale feedback, got %v", msgIDs)
|
||||
}
|
||||
if _, ok := m.streamActive.Load("test:123"); !ok {
|
||||
t.Fatal("expected streamActive marker to remain for the final outbound message")
|
||||
}
|
||||
if _, ok := m.placeholders.Load("test:123"); !ok {
|
||||
t.Fatal("expected placeholder cleanup to remain deferred to the final outbound message")
|
||||
}
|
||||
if ch.editedMessages != 0 {
|
||||
t.Fatalf("expected no placeholder edit for stale feedback, got %d edits", ch.editedMessages)
|
||||
}
|
||||
|
||||
finalMsg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "final streamed reply",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
},
|
||||
})
|
||||
|
||||
_, handled = m.preSend(context.Background(), "test", finalMsg, ch)
|
||||
if !handled {
|
||||
t.Fatal("expected final outbound message to consume streamActive marker")
|
||||
}
|
||||
if _, ok := m.streamActive.Load("test:123"); ok {
|
||||
t.Fatal("expected streamActive marker to be cleared by final outbound message")
|
||||
}
|
||||
if _, ok := m.placeholders.Load("test:123"); ok {
|
||||
t.Fatal("expected placeholder to be cleaned up by final outbound message")
|
||||
}
|
||||
if editedContent != "final streamed reply" {
|
||||
t.Fatalf("editedContent = %q, want final streamed reply", editedContent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSendMedia_LeavesTrackedMessageForChannelSend(t *testing.T) {
|
||||
m := newTestManager()
|
||||
ch := &mockDeletingMediaChannel{}
|
||||
|
||||
m.preSendMedia(context.Background(), "test", bus.OutboundMediaMessage{
|
||||
ChatID: "123",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
},
|
||||
}, ch)
|
||||
|
||||
if ch.dismissedChatID != "" {
|
||||
t.Fatalf(
|
||||
"expected tracked tool feedback cleanup to be deferred to channel media send, got %q",
|
||||
ch.dismissedChatID,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitOutboundMessageContent_ToolFeedbackTruncatesInsteadOfSplitting(t *testing.T) {
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "\U0001f527 `read_file`\nRead README.md first to confirm the current project structure before editing the config example.",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
chunks := splitOutboundMessageContent(msg, 40)
|
||||
if len(chunks) != 1 {
|
||||
t.Fatalf("len(chunks) = %d, want 1", len(chunks))
|
||||
}
|
||||
want := utils.FitToolFeedbackMessage(msg.Content, 40-MaxToolFeedbackAnimationFrameLength())
|
||||
if chunks[0] != want {
|
||||
t.Fatalf("chunk = %q, want %q", chunks[0], want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitOutboundMessageContent_ToolFeedbackReservesAnimationFrame(t *testing.T) {
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: "🔧 `read_file`\n1234567890",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
chunks := splitOutboundMessageContent(msg, len([]rune(msg.Content)))
|
||||
if len(chunks) != 1 {
|
||||
t.Fatalf("len(chunks) = %d, want 1", len(chunks))
|
||||
}
|
||||
|
||||
animated := formatAnimatedToolFeedbackContent(chunks[0], strings.Repeat(".", MaxToolFeedbackAnimationFrameLength()))
|
||||
if got, maxLen := len([]rune(animated)), len([]rune(msg.Content)); got > maxLen {
|
||||
t.Fatalf("animated len = %d, want <= %d; content=%q", got, maxLen, animated)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStreamer_FinalizeDismissesTrackedToolFeedback(t *testing.T) {
|
||||
m := newTestManager()
|
||||
ch := &mockStreamingChannel{
|
||||
mockMessageEditor: mockMessageEditor{},
|
||||
streamer: &mockStreamer{
|
||||
finalizeFn: func(_ context.Context, content string) error {
|
||||
if content != "final reply" {
|
||||
t.Fatalf("unexpected finalize content: %q", content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
},
|
||||
}
|
||||
m.channels["test"] = ch
|
||||
|
||||
streamer, ok := m.GetStreamer(context.Background(), "test", "123")
|
||||
if !ok {
|
||||
t.Fatal("expected streamer to be available")
|
||||
}
|
||||
if err := streamer.Finalize(context.Background(), "final reply"); err != nil {
|
||||
t.Fatalf("Finalize() error = %v", err)
|
||||
}
|
||||
if ch.dismissedChatID != "123" {
|
||||
t.Fatalf("expected tracked tool feedback to be dismissed for chat 123, got %q", ch.dismissedChatID)
|
||||
}
|
||||
if _, ok := m.streamActive.Load("test:123"); !ok {
|
||||
t.Fatal("expected streamActive marker to be recorded after finalize")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStreamer_FinalizeDismissesResolvedTrackedToolFeedback(t *testing.T) {
|
||||
m := newTestManager()
|
||||
ch := &mockStreamingChannel{
|
||||
mockMessageEditor: mockMessageEditor{},
|
||||
streamer: &mockStreamer{
|
||||
finalizeFn: func(_ context.Context, content string) error {
|
||||
if content != "final reply" {
|
||||
t.Fatalf("unexpected finalize content: %q", content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
},
|
||||
resolveChatIDFn: func(chatID string, outboundCtx *bus.InboundContext) string {
|
||||
if outboundCtx == nil {
|
||||
t.Fatal("expected outbound context during stream finalize")
|
||||
}
|
||||
if outboundCtx.ChatID != "-100123/42" {
|
||||
t.Fatalf("unexpected outbound context: %+v", outboundCtx)
|
||||
}
|
||||
return outboundCtx.ChatID
|
||||
},
|
||||
}
|
||||
m.channels["test"] = ch
|
||||
|
||||
streamer, ok := m.GetStreamer(context.Background(), "test", "-100123/42")
|
||||
if !ok {
|
||||
t.Fatal("expected streamer to be available")
|
||||
}
|
||||
if err := streamer.Finalize(context.Background(), "final reply"); err != nil {
|
||||
t.Fatalf("Finalize() error = %v", err)
|
||||
}
|
||||
if ch.dismissedChatID != "-100123/42" {
|
||||
t.Fatalf("expected resolved tracked tool feedback dismissal, got %q", ch.dismissedChatID)
|
||||
}
|
||||
if _, ok := m.streamActive.Load("test:-100123/42"); !ok {
|
||||
t.Fatal("expected streamActive marker to be recorded after finalize")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_PlaceholderEditSuccessDismissesResolvedTrackedToolFeedback(t *testing.T) {
|
||||
m := newTestManager()
|
||||
|
||||
ch := &mockResolvedToolFeedbackEditor{
|
||||
mockMessageEditor: mockMessageEditor{
|
||||
editFn: func(_ context.Context, chatID, messageID, content string) error {
|
||||
if chatID != "-100123" || messageID != "456" || content != "done" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
},
|
||||
resolveChatIDFn: func(chatID string, outboundCtx *bus.InboundContext) string {
|
||||
if outboundCtx == nil || outboundCtx.TopicID != "42" {
|
||||
t.Fatalf("expected topic-aware outbound context, got %+v", outboundCtx)
|
||||
}
|
||||
return chatID + "/" + outboundCtx.TopicID
|
||||
},
|
||||
}
|
||||
|
||||
m.RecordPlaceholder("test", "-100123", "456")
|
||||
|
||||
msg := testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "-100123",
|
||||
Content: "done",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "-100123",
|
||||
TopicID: "42",
|
||||
},
|
||||
})
|
||||
|
||||
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||
if !edited {
|
||||
t.Fatal("expected preSend to edit placeholder")
|
||||
}
|
||||
if ch.dismissedChatID != "-100123/42" {
|
||||
t.Fatalf("expected resolved tracked dismissal, got %q", ch.dismissedChatID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStreamer_FinalizeFailureDoesNotDismissTrackedToolFeedback(t *testing.T) {
|
||||
m := newTestManager()
|
||||
ch := &mockStreamingChannel{
|
||||
mockMessageEditor: mockMessageEditor{},
|
||||
streamer: &mockStreamer{
|
||||
finalizeFn: func(context.Context, string) error {
|
||||
return errors.New("finalize failed")
|
||||
},
|
||||
},
|
||||
}
|
||||
m.channels["test"] = ch
|
||||
|
||||
streamer, ok := m.GetStreamer(context.Background(), "test", "123")
|
||||
if !ok {
|
||||
t.Fatal("expected streamer to be available")
|
||||
}
|
||||
if err := streamer.Finalize(context.Background(), "final reply"); err == nil {
|
||||
t.Fatal("expected Finalize() to fail")
|
||||
}
|
||||
if ch.dismissedChatID != "" {
|
||||
t.Fatalf("expected no tool feedback dismissal on finalize failure, got %q", ch.dismissedChatID)
|
||||
}
|
||||
if _, ok := m.streamActive.Load("test:123"); ok {
|
||||
t.Fatal("expected no streamActive marker after finalize failure")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunWorker_ToolFeedbackSkipsMarkerSplitting(t *testing.T) {
|
||||
m := newTestManager()
|
||||
m.config = &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
SplitOnMarker: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
var (
|
||||
mu sync.Mutex
|
||||
received []string
|
||||
)
|
||||
ch := &mockChannelWithLength{
|
||||
mockChannel: mockChannel{
|
||||
sendFn: func(_ context.Context, msg bus.OutboundMessage) error {
|
||||
mu.Lock()
|
||||
received = append(received, msg.Content)
|
||||
mu.Unlock()
|
||||
return nil
|
||||
},
|
||||
},
|
||||
maxLen: 200,
|
||||
}
|
||||
|
||||
w := &channelWorker{
|
||||
ch: ch,
|
||||
queue: make(chan bus.OutboundMessage, 1),
|
||||
done: make(chan struct{}),
|
||||
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
go m.runWorker(ctx, "test", w)
|
||||
|
||||
content := "🔧 `read_file`\nRead current config first.<|[SPLIT]|>Then update the example."
|
||||
w.queue <- testOutboundMessage(bus.OutboundMessage{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Content: content,
|
||||
Context: bus.InboundContext{
|
||||
Channel: "test",
|
||||
ChatID: "123",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if len(received) != 1 {
|
||||
t.Fatalf("len(received) = %d, want 1", len(received))
|
||||
}
|
||||
if received[0] != content {
|
||||
t.Fatalf("received[0] = %q, want %q", received[0], content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreSend_PlaceholderEditFails_FallsThrough(t *testing.T) {
|
||||
m := newTestManager()
|
||||
|
||||
|
|
|
|||
|
|
@ -46,6 +46,13 @@ const (
|
|||
|
||||
var matrixMentionHrefRegexp = regexp.MustCompile(`(?i)<a[^>]+href=["']([^"']+)["']`)
|
||||
|
||||
func outboundMessageIsToolFeedback(msg bus.OutboundMessage) bool {
|
||||
if len(msg.Context.Raw) == 0 {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), "tool_feedback")
|
||||
}
|
||||
|
||||
type roomKindCacheEntry struct {
|
||||
isGroup bool
|
||||
expiresAt time.Time
|
||||
|
|
@ -192,6 +199,7 @@ type MatrixChannel struct {
|
|||
|
||||
cryptoHelper *cryptohelper.CryptoHelper
|
||||
cryptoDbPath string
|
||||
progress *channels.ToolFeedbackAnimator
|
||||
}
|
||||
|
||||
func NewMatrixChannel(
|
||||
|
|
@ -236,7 +244,7 @@ func NewMatrixChannel(
|
|||
channels.WithReasoningChannelID(bc.ReasoningChannelID),
|
||||
)
|
||||
|
||||
return &MatrixChannel{
|
||||
ch := &MatrixChannel{
|
||||
BaseChannel: base,
|
||||
bc: bc,
|
||||
client: client,
|
||||
|
|
@ -248,7 +256,9 @@ func NewMatrixChannel(
|
|||
localpartMentionR: localpartMentionRegexp(matrixLocalpart(client.UserID)),
|
||||
typingMu: sync.Mutex{},
|
||||
cryptoDbPath: cryptoDatabasePath,
|
||||
}, nil
|
||||
}
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) Start(ctx context.Context) error {
|
||||
|
|
@ -297,6 +307,9 @@ func (c *MatrixChannel) Stop(ctx context.Context) error {
|
|||
c.cancel()
|
||||
}
|
||||
c.stopTypingSessions(ctx)
|
||||
if c.progress != nil {
|
||||
c.progress.StopAll()
|
||||
}
|
||||
|
||||
// Close crypto helper if initialized
|
||||
if c.cryptoHelper != nil {
|
||||
|
|
@ -398,11 +411,36 @@ func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]st
|
|||
return nil, nil
|
||||
}
|
||||
|
||||
isToolFeedback := outboundMessageIsToolFeedback(msg)
|
||||
if isToolFeedback {
|
||||
if msgID, handled, err := c.progress.Update(ctx, msg.ChatID, content); handled {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
|
||||
if !isToolFeedback {
|
||||
if msgIDs, handled := c.FinalizeToolFeedbackMessage(ctx, msg); handled {
|
||||
return msgIDs, nil
|
||||
}
|
||||
}
|
||||
if isToolFeedback {
|
||||
content = channels.InitialAnimatedToolFeedbackContent(content)
|
||||
}
|
||||
|
||||
resp, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, c.messageContent(content))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("matrix send: %w", channels.ErrTemporary)
|
||||
}
|
||||
return []string{resp.EventID.String()}, nil
|
||||
msgID := resp.EventID.String()
|
||||
if isToolFeedback {
|
||||
c.RecordToolFeedbackMessage(msg.ChatID, msgID, msg.Content)
|
||||
} else if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) messageContent(text string) *event.MessageEventContent {
|
||||
|
|
@ -419,6 +457,8 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
|||
if !c.IsRunning() {
|
||||
return nil, channels.ErrNotRunning
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
|
||||
|
||||
sendCtx := ctx
|
||||
if sendCtx == nil {
|
||||
sendCtx = context.Background()
|
||||
|
|
@ -529,6 +569,10 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
|||
}
|
||||
}
|
||||
|
||||
if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
|
||||
return eventIDs, nil
|
||||
}
|
||||
|
||||
|
|
@ -612,6 +656,89 @@ func (c *MatrixChannel) EditMessage(ctx context.Context, chatID string, messageI
|
|||
return err
|
||||
}
|
||||
|
||||
// DeleteMessage implements channels.MessageDeleter.
|
||||
func (c *MatrixChannel) DeleteMessage(ctx context.Context, chatID string, messageID string) error {
|
||||
roomID := id.RoomID(strings.TrimSpace(chatID))
|
||||
if roomID == "" {
|
||||
return fmt.Errorf("matrix room ID is empty")
|
||||
}
|
||||
eventID := id.EventID(strings.TrimSpace(messageID))
|
||||
if eventID == "" {
|
||||
return fmt.Errorf("matrix message ID is empty")
|
||||
}
|
||||
|
||||
_, err := c.client.RedactEvent(ctx, roomID, eventID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) currentToolFeedbackMessage(chatID string) (string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", false
|
||||
}
|
||||
return c.progress.Current(chatID)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) takeToolFeedbackMessage(chatID string) (string, string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", "", false
|
||||
}
|
||||
return c.progress.Take(chatID)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) RecordToolFeedbackMessage(chatID, messageID, content string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Record(chatID, messageID, content)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) ClearToolFeedbackMessage(chatID string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Clear(chatID)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) DismissToolFeedbackMessage(ctx context.Context, chatID string) {
|
||||
msgID, ok := c.currentToolFeedbackMessage(chatID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, chatID, msgID)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) dismissTrackedToolFeedbackMessage(ctx context.Context, chatID, messageID string) {
|
||||
if strings.TrimSpace(chatID) == "" || strings.TrimSpace(messageID) == "" {
|
||||
return
|
||||
}
|
||||
c.ClearToolFeedbackMessage(chatID)
|
||||
_ = c.DeleteMessage(ctx, chatID, messageID)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) finalizeTrackedToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
content string,
|
||||
editFn func(context.Context, string, string, string) error,
|
||||
) ([]string, bool) {
|
||||
msgID, baseContent, ok := c.takeToolFeedbackMessage(chatID)
|
||||
if !ok || editFn == nil {
|
||||
return nil, false
|
||||
}
|
||||
if err := editFn(ctx, chatID, msgID, content); err != nil {
|
||||
c.RecordToolFeedbackMessage(chatID, msgID, baseContent)
|
||||
return nil, false
|
||||
}
|
||||
return []string{msgID}, true
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) FinalizeToolFeedbackMessage(ctx context.Context, msg bus.OutboundMessage) ([]string, bool) {
|
||||
if outboundMessageIsToolFeedback(msg) {
|
||||
return nil, false
|
||||
}
|
||||
return c.finalizeTrackedToolFeedbackMessage(ctx, msg.ChatID, msg.Content, c.EditMessage)
|
||||
}
|
||||
|
||||
func (c *MatrixChannel) handleMemberEvent(ctx context.Context, evt *event.Event) {
|
||||
if !c.config.JoinOnInvite {
|
||||
return
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import (
|
|||
"maunium.net/go/mautrix/event"
|
||||
"maunium.net/go/mautrix/id"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
)
|
||||
|
|
@ -41,6 +42,34 @@ func TestMatrixLocalpartMentionRegexp(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_StopsTrackingBeforeEdit(t *testing.T) {
|
||||
ch := &MatrixChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("!room:matrix.org", "$event1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"!room:matrix.org",
|
||||
"final reply",
|
||||
func(_ context.Context, chatID, messageID, content string) error {
|
||||
if _, ok := ch.currentToolFeedbackMessage(chatID); ok {
|
||||
t.Fatal("expected tracked tool feedback to be stopped before edit")
|
||||
}
|
||||
if chatID != "!room:matrix.org" || messageID != "$event1" || content != "final reply" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if !handled {
|
||||
t.Fatal("expected finalizeTrackedToolFeedbackMessage to handle tracked message")
|
||||
}
|
||||
if len(msgIDs) != 1 || msgIDs[0] != "$event1" {
|
||||
t.Fatalf("finalizeTrackedToolFeedbackMessage() ids = %v, want [$event1]", msgIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripUserMention(t *testing.T) {
|
||||
userID := id.UserID("@picoclaw:matrix.org")
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,11 @@ import (
|
|||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
|
@ -46,6 +50,17 @@ func outboundMessageIsThought(msg bus.OutboundMessage) bool {
|
|||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), MessageKindThought)
|
||||
}
|
||||
|
||||
func outboundMessageIsToolFeedback(msg bus.OutboundMessage) bool {
|
||||
if len(msg.Context.Raw) == 0 {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), "tool_feedback")
|
||||
}
|
||||
|
||||
func outboundMessageFinalizesTrackedToolFeedback(msg bus.OutboundMessage) bool {
|
||||
return !outboundMessageIsToolFeedback(msg) && !outboundMessageIsThought(msg)
|
||||
}
|
||||
|
||||
// writeJSON sends a JSON message to the connection with write locking.
|
||||
func (pc *picoConn) writeJSON(v any) error {
|
||||
if pc.closed.Load() {
|
||||
|
|
@ -78,6 +93,8 @@ type PicoChannel struct {
|
|||
connsMu sync.RWMutex
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
progress *channels.ToolFeedbackAnimator
|
||||
deleteMessageFn func(context.Context, string, string) error
|
||||
}
|
||||
|
||||
// NewPicoChannel creates a new Pico Protocol channel.
|
||||
|
|
@ -106,7 +123,7 @@ func NewPicoChannel(
|
|||
return false
|
||||
}
|
||||
|
||||
return &PicoChannel{
|
||||
ch := &PicoChannel{
|
||||
BaseChannel: base,
|
||||
bc: bc,
|
||||
config: cfg,
|
||||
|
|
@ -117,7 +134,10 @@ func NewPicoChannel(
|
|||
},
|
||||
connections: make(map[string]*picoConn),
|
||||
sessionConnections: make(map[string]map[string]*picoConn),
|
||||
}, nil
|
||||
}
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
ch.deleteMessageFn = ch.DeleteMessage
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
// createAndAddConnection checks MaxConnections and registers a connection atomically.
|
||||
|
|
@ -235,6 +255,9 @@ func (c *PicoChannel) Stop(ctx context.Context) error {
|
|||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
if c.progress != nil {
|
||||
c.progress.StopAll()
|
||||
}
|
||||
|
||||
logger.InfoC("pico", "Pico Protocol channel stopped")
|
||||
return nil
|
||||
|
|
@ -251,6 +274,10 @@ func (c *PicoChannel) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||
case "/ws", "/ws/":
|
||||
c.handleWebSocket(w, r)
|
||||
default:
|
||||
if strings.HasPrefix(path, "/media/") {
|
||||
c.handleMediaDownload(w, r)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}
|
||||
|
|
@ -261,24 +288,133 @@ func (c *PicoChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]stri
|
|||
return nil, channels.ErrNotRunning
|
||||
}
|
||||
isThought := outboundMessageIsThought(msg)
|
||||
isToolFeedback := outboundMessageIsToolFeedback(msg)
|
||||
if isToolFeedback {
|
||||
if msgID, handled, err := c.progress.Update(ctx, msg.ChatID, msg.Content); handled {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
|
||||
if outboundMessageFinalizesTrackedToolFeedback(msg) {
|
||||
if msgIDs, handled := c.FinalizeToolFeedbackMessage(ctx, msg); handled {
|
||||
return msgIDs, nil
|
||||
}
|
||||
}
|
||||
|
||||
outMsg := newMessage(TypeMessageCreate, map[string]any{
|
||||
PayloadKeyContent: msg.Content,
|
||||
content := msg.Content
|
||||
if isToolFeedback {
|
||||
content = channels.InitialAnimatedToolFeedbackContent(msg.Content)
|
||||
}
|
||||
msgID := uuid.New().String()
|
||||
|
||||
payload := map[string]any{
|
||||
PayloadKeyContent: content,
|
||||
PayloadKeyThought: isThought,
|
||||
})
|
||||
"message_id": msgID,
|
||||
}
|
||||
setContextUsagePayload(payload, msg.ContextUsage)
|
||||
outMsg := newMessage(TypeMessageCreate, payload)
|
||||
|
||||
return nil, c.broadcastToSession(msg.ChatID, outMsg)
|
||||
if err := c.broadcastToSession(msg.ChatID, outMsg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if isToolFeedback {
|
||||
c.RecordToolFeedbackMessage(msg.ChatID, msgID, msg.Content)
|
||||
} else if hasTrackedMsg && outboundMessageFinalizesTrackedToolFeedback(msg) {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
|
||||
// EditMessage implements channels.MessageEditor.
|
||||
func (c *PicoChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
||||
outMsg := newMessage(TypeMessageUpdate, map[string]any{
|
||||
return c.editMessage(ctx, chatID, messageID, content, nil)
|
||||
}
|
||||
|
||||
// DeleteMessage implements channels.MessageDeleter.
|
||||
func (c *PicoChannel) DeleteMessage(ctx context.Context, chatID string, messageID string) error {
|
||||
outMsg := newMessage(TypeMessageDelete, map[string]any{
|
||||
"message_id": messageID,
|
||||
"content": content,
|
||||
})
|
||||
return c.broadcastToSession(chatID, outMsg)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) currentToolFeedbackMessage(chatID string) (string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", false
|
||||
}
|
||||
return c.progress.Current(chatID)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) takeToolFeedbackMessage(chatID string) (string, string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", "", false
|
||||
}
|
||||
return c.progress.Take(chatID)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) RecordToolFeedbackMessage(chatID, messageID, content string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Record(chatID, messageID, content)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) ClearToolFeedbackMessage(chatID string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Clear(chatID)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) DismissToolFeedbackMessage(ctx context.Context, chatID string) {
|
||||
msgID, ok := c.currentToolFeedbackMessage(chatID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, chatID, msgID)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) dismissTrackedToolFeedbackMessage(ctx context.Context, chatID, messageID string) {
|
||||
if strings.TrimSpace(chatID) == "" || strings.TrimSpace(messageID) == "" {
|
||||
return
|
||||
}
|
||||
c.ClearToolFeedbackMessage(chatID)
|
||||
deleteFn := c.deleteMessageFn
|
||||
if deleteFn == nil {
|
||||
deleteFn = c.DeleteMessage
|
||||
}
|
||||
_ = deleteFn(ctx, chatID, messageID)
|
||||
}
|
||||
|
||||
func (c *PicoChannel) finalizeTrackedToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
content string,
|
||||
editFn func(context.Context, string, string, string, *bus.ContextUsage) error,
|
||||
contextUsage *bus.ContextUsage,
|
||||
) ([]string, bool) {
|
||||
msgID, baseContent, ok := c.takeToolFeedbackMessage(chatID)
|
||||
if !ok || editFn == nil {
|
||||
return nil, false
|
||||
}
|
||||
if err := editFn(ctx, chatID, msgID, content, contextUsage); err != nil {
|
||||
c.RecordToolFeedbackMessage(chatID, msgID, baseContent)
|
||||
return nil, false
|
||||
}
|
||||
return []string{msgID}, true
|
||||
}
|
||||
|
||||
func (c *PicoChannel) FinalizeToolFeedbackMessage(ctx context.Context, msg bus.OutboundMessage) ([]string, bool) {
|
||||
if !outboundMessageFinalizesTrackedToolFeedback(msg) {
|
||||
return nil, false
|
||||
}
|
||||
return c.finalizeTrackedToolFeedbackMessage(ctx, msg.ChatID, msg.Content, c.editMessage, msg.ContextUsage)
|
||||
}
|
||||
|
||||
// StartTyping implements channels.TypingCapable.
|
||||
func (c *PicoChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
||||
startMsg := newMessage(TypeTypingStart, nil)
|
||||
|
|
@ -315,6 +451,210 @@ func (c *PicoChannel) SendPlaceholder(ctx context.Context, chatID string) (strin
|
|||
return msgID, nil
|
||||
}
|
||||
|
||||
// SendMedia implements channels.MediaSender for the Pico web UI.
|
||||
// Media is delivered as a normal assistant message carrying structured
|
||||
// attachments plus an authenticated same-origin download URL.
|
||||
func (c *PicoChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||
if !c.IsRunning() {
|
||||
return nil, channels.ErrNotRunning
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
|
||||
|
||||
store := c.GetMediaStore()
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||
}
|
||||
|
||||
attachments := make([]map[string]any, 0, len(msg.Parts))
|
||||
caption := ""
|
||||
|
||||
for _, part := range msg.Parts {
|
||||
localPath, meta, err := store.ResolveWithMeta(part.Ref)
|
||||
if err != nil {
|
||||
logger.ErrorCF("pico", "Failed to resolve media ref", map[string]any{
|
||||
"ref": part.Ref,
|
||||
"error": err.Error(),
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
filename := strings.TrimSpace(part.Filename)
|
||||
if filename == "" {
|
||||
filename = strings.TrimSpace(meta.Filename)
|
||||
}
|
||||
if filename == "" {
|
||||
filename = filepath.Base(localPath)
|
||||
}
|
||||
|
||||
contentType := strings.TrimSpace(part.ContentType)
|
||||
if contentType == "" {
|
||||
contentType = strings.TrimSpace(meta.ContentType)
|
||||
}
|
||||
if contentType == "" {
|
||||
contentType = "application/octet-stream"
|
||||
}
|
||||
|
||||
attachmentType := strings.TrimSpace(part.Type)
|
||||
if attachmentType == "" {
|
||||
attachmentType = picoInferAttachmentType(filename, contentType)
|
||||
}
|
||||
|
||||
attachmentURL, err := picoDownloadURLForRef(part.Ref)
|
||||
if err != nil {
|
||||
logger.ErrorCF("pico", "Failed to build media download URL", map[string]any{
|
||||
"ref": part.Ref,
|
||||
"error": err.Error(),
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
attachments = append(attachments, map[string]any{
|
||||
"type": attachmentType,
|
||||
"url": attachmentURL,
|
||||
"filename": filename,
|
||||
"content_type": contentType,
|
||||
})
|
||||
|
||||
if caption == "" && strings.TrimSpace(part.Caption) != "" {
|
||||
caption = strings.TrimSpace(part.Caption)
|
||||
}
|
||||
}
|
||||
|
||||
if len(attachments) == 0 {
|
||||
return nil, fmt.Errorf("no deliverable media parts: %w", channels.ErrSendFailed)
|
||||
}
|
||||
|
||||
msgID := uuid.New().String()
|
||||
outMsg := newMessage(TypeMessageCreate, map[string]any{
|
||||
PayloadKeyContent: caption,
|
||||
"attachments": attachments,
|
||||
"message_id": msgID,
|
||||
})
|
||||
|
||||
if err := c.broadcastToSession(msg.ChatID, outMsg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
|
||||
}
|
||||
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
|
||||
func picoDownloadURLForRef(ref string) (string, error) {
|
||||
refID, err := picoMediaRefID(ref)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "/pico/media/" + url.PathEscape(refID), nil
|
||||
}
|
||||
|
||||
func picoMediaRefID(ref string) (string, error) {
|
||||
refID := strings.TrimSpace(strings.TrimPrefix(ref, "media://"))
|
||||
if refID == "" || strings.Contains(refID, "/") {
|
||||
return "", fmt.Errorf("invalid media ref %q", ref)
|
||||
}
|
||||
return refID, nil
|
||||
}
|
||||
|
||||
func picoInferAttachmentType(filename, contentType string) string {
|
||||
contentType = strings.ToLower(strings.TrimSpace(contentType))
|
||||
filename = strings.ToLower(strings.TrimSpace(filename))
|
||||
|
||||
switch {
|
||||
case strings.HasPrefix(contentType, "image/"):
|
||||
return "image"
|
||||
case strings.HasPrefix(contentType, "audio/"):
|
||||
return "audio"
|
||||
case strings.HasPrefix(contentType, "video/"):
|
||||
return "video"
|
||||
}
|
||||
|
||||
switch ext := filepath.Ext(filename); ext {
|
||||
case ".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp", ".svg":
|
||||
return "image"
|
||||
case ".mp3", ".wav", ".ogg", ".m4a", ".flac", ".aac", ".wma", ".opus":
|
||||
return "audio"
|
||||
case ".mp4", ".avi", ".mov", ".webm", ".mkv":
|
||||
return "video"
|
||||
default:
|
||||
return "file"
|
||||
}
|
||||
}
|
||||
|
||||
func picoAllowsInlineDisplay(filename, contentType string) bool {
|
||||
contentType = strings.ToLower(strings.TrimSpace(contentType))
|
||||
filename = strings.ToLower(strings.TrimSpace(filename))
|
||||
|
||||
if strings.Contains(contentType, "svg") || filepath.Ext(filename) == ".svg" {
|
||||
return false
|
||||
}
|
||||
|
||||
return picoInferAttachmentType(filename, contentType) == "image"
|
||||
}
|
||||
|
||||
func (c *PicoChannel) handleMediaDownload(w http.ResponseWriter, r *http.Request) {
|
||||
if !c.IsRunning() {
|
||||
http.Error(w, "channel not running", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
if !c.authenticate(r) {
|
||||
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
refID := strings.TrimSpace(strings.TrimPrefix(strings.TrimPrefix(r.URL.Path, "/pico/media/"), "/"))
|
||||
if refID == "" {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
store := c.GetMediaStore()
|
||||
if store == nil {
|
||||
http.Error(w, "media store unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
|
||||
localPath, meta, err := store.ResolveWithMeta("media://" + refID)
|
||||
if err != nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
file, err := os.Open(localPath)
|
||||
if err != nil {
|
||||
http.Error(w, "failed to open media", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
http.Error(w, "failed to stat media", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
filename := strings.TrimSpace(meta.Filename)
|
||||
if filename == "" {
|
||||
filename = filepath.Base(localPath)
|
||||
}
|
||||
contentType := strings.TrimSpace(meta.ContentType)
|
||||
if contentType == "" {
|
||||
contentType = "application/octet-stream"
|
||||
}
|
||||
|
||||
dispositionType := "attachment"
|
||||
if picoAllowsInlineDisplay(filename, contentType) {
|
||||
dispositionType = "inline"
|
||||
}
|
||||
|
||||
if cd := mime.FormatMediaType(dispositionType, map[string]string{"filename": filename}); cd != "" {
|
||||
w.Header().Set("Content-Disposition", cd)
|
||||
}
|
||||
w.Header().Set("Content-Type", contentType)
|
||||
http.ServeContent(w, r, filename, info.ModTime(), file)
|
||||
}
|
||||
|
||||
// broadcastToSession sends a message to all connections with a matching session.
|
||||
func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error {
|
||||
// chatID format: "pico:<sessionID>"
|
||||
|
|
@ -716,3 +1056,32 @@ func validateInlineImageDataURL(mediaURL string) error {
|
|||
|
||||
return nil
|
||||
}
|
||||
|
||||
// setContextUsagePayload adds context window usage stats to a pico payload.
|
||||
func setContextUsagePayload(payload map[string]any, u *bus.ContextUsage) {
|
||||
if u == nil {
|
||||
return
|
||||
}
|
||||
payload["context_usage"] = map[string]any{
|
||||
"used_tokens": u.UsedTokens,
|
||||
"total_tokens": u.TotalTokens,
|
||||
"compress_at_tokens": u.CompressAtTokens,
|
||||
"used_percent": u.UsedPercent,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *PicoChannel) editMessage(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
messageID string,
|
||||
content string,
|
||||
contextUsage *bus.ContextUsage,
|
||||
) error {
|
||||
payload := map[string]any{
|
||||
"message_id": messageID,
|
||||
"content": content,
|
||||
}
|
||||
setContextUsagePayload(payload, contextUsage)
|
||||
outMsg := newMessage(TypeMessageUpdate, payload)
|
||||
return c.broadcastToSession(chatID, outMsg)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,12 +4,21 @@ import (
|
|||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
)
|
||||
|
||||
func newTestPicoChannel(t *testing.T) *PicoChannel {
|
||||
|
|
@ -27,6 +36,163 @@ func newTestPicoChannel(t *testing.T) *PicoChannel {
|
|||
return ch
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_StopsTrackingBeforeEdit(t *testing.T) {
|
||||
ch := &PicoChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("pico:chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"pico:chat-1",
|
||||
"final reply",
|
||||
func(_ context.Context, chatID, messageID, content string, contextUsage *bus.ContextUsage) error {
|
||||
if _, ok := ch.currentToolFeedbackMessage(chatID); ok {
|
||||
t.Fatal("expected tracked tool feedback to be stopped before edit")
|
||||
}
|
||||
if chatID != "pico:chat-1" || messageID != "msg-1" || content != "final reply" {
|
||||
t.Fatalf("unexpected edit args: %s %s %s", chatID, messageID, content)
|
||||
}
|
||||
if contextUsage != nil {
|
||||
t.Fatalf("unexpected context usage: %+v", contextUsage)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
nil,
|
||||
)
|
||||
if !handled {
|
||||
t.Fatal("expected finalizeTrackedToolFeedbackMessage to handle tracked message")
|
||||
}
|
||||
if len(msgIDs) != 1 || msgIDs[0] != "msg-1" {
|
||||
t.Fatalf("finalizeTrackedToolFeedbackMessage() ids = %v, want [msg-1]", msgIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDismissTrackedToolFeedbackMessage_DeletesProgressMessage(t *testing.T) {
|
||||
ch := &PicoChannel{
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
ch.RecordToolFeedbackMessage("pico:chat-1", "msg-1", "🔧 `read_file`")
|
||||
|
||||
var deleted struct {
|
||||
chatID string
|
||||
messageID string
|
||||
}
|
||||
ch.deleteMessageFn = func(_ context.Context, chatID string, messageID string) error {
|
||||
deleted.chatID = chatID
|
||||
deleted.messageID = messageID
|
||||
return nil
|
||||
}
|
||||
|
||||
ch.DismissToolFeedbackMessage(context.Background(), "pico:chat-1")
|
||||
|
||||
if deleted.chatID != "pico:chat-1" || deleted.messageID != "msg-1" {
|
||||
t.Fatalf("unexpected delete target: %+v", deleted)
|
||||
}
|
||||
if _, ok := ch.currentToolFeedbackMessage("pico:chat-1"); ok {
|
||||
t.Fatal("expected tracked tool feedback to be cleared after dismissal")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSend_ThoughtMessageDoesNotFinalizeTrackedToolFeedback(t *testing.T) {
|
||||
ch := newTestPicoChannel(t)
|
||||
|
||||
if err := ch.Start(context.Background()); err != nil {
|
||||
t.Fatalf("Start() error = %v", err)
|
||||
}
|
||||
defer ch.Stop(context.Background())
|
||||
|
||||
clientConn, received, cleanup := newTestPicoWebSocket(t)
|
||||
defer cleanup()
|
||||
ch.addConnForTest(&picoConn{id: "conn-1", conn: clientConn, sessionID: "sess-1"})
|
||||
|
||||
ch.RecordToolFeedbackMessage("pico:sess-1", "msg-progress", "🔧 `read_file`\nReading config")
|
||||
|
||||
if _, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "pico:sess-1",
|
||||
Content: "thinking trace",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "pico",
|
||||
ChatID: "pico:sess-1",
|
||||
Raw: map[string]string{
|
||||
"message_kind": MessageKindThought,
|
||||
},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("Send(thought) error = %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case msg := <-received:
|
||||
if msg.Type != TypeMessageCreate {
|
||||
t.Fatalf("thought message type = %q, want %q", msg.Type, TypeMessageCreate)
|
||||
}
|
||||
payload := msg.Payload
|
||||
if got := payload[PayloadKeyContent]; got != "thinking trace" {
|
||||
t.Fatalf("thought content = %#v, want %q", got, "thinking trace")
|
||||
}
|
||||
if got := payload[PayloadKeyThought]; got != true {
|
||||
t.Fatalf("thought flag = %#v, want true", got)
|
||||
}
|
||||
if got := payload["message_id"]; got == "msg-progress" || got == nil || got == "" {
|
||||
t.Fatalf("thought message_id = %#v, want new non-progress id", got)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected thought message to be delivered")
|
||||
}
|
||||
|
||||
if msgID, ok := ch.currentToolFeedbackMessage("pico:sess-1"); !ok || msgID != "msg-progress" {
|
||||
t.Fatalf("tracked tool feedback = (%q, %v), want (msg-progress, true)", msgID, ok)
|
||||
}
|
||||
|
||||
if _, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "pico:sess-1",
|
||||
Content: "final reply",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "pico",
|
||||
ChatID: "pico:sess-1",
|
||||
},
|
||||
ContextUsage: &bus.ContextUsage{
|
||||
UsedTokens: 321,
|
||||
TotalTokens: 4096,
|
||||
CompressAtTokens: 3072,
|
||||
UsedPercent: 8,
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("Send(final) error = %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case msg := <-received:
|
||||
if msg.Type != TypeMessageUpdate {
|
||||
t.Fatalf("final message type = %q, want %q", msg.Type, TypeMessageUpdate)
|
||||
}
|
||||
payload := msg.Payload
|
||||
if got := payload["message_id"]; got != "msg-progress" {
|
||||
t.Fatalf("final message_id = %#v, want %q", got, "msg-progress")
|
||||
}
|
||||
if got := payload[PayloadKeyContent]; got != "final reply" {
|
||||
t.Fatalf("final content = %#v, want %q", got, "final reply")
|
||||
}
|
||||
rawUsage, ok := payload["context_usage"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("final context_usage = %#v, want map payload", payload["context_usage"])
|
||||
}
|
||||
if got, ok := rawUsage["used_tokens"].(float64); !ok || got != 321 {
|
||||
t.Fatalf("used_tokens = %#v, want 321", rawUsage["used_tokens"])
|
||||
}
|
||||
if got, ok := rawUsage["total_tokens"].(float64); !ok || got != 4096 {
|
||||
t.Fatalf("total_tokens = %#v, want 4096", rawUsage["total_tokens"])
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected final reply to finalize tracked tool feedback")
|
||||
}
|
||||
|
||||
if _, ok := ch.currentToolFeedbackMessage("pico:sess-1"); ok {
|
||||
t.Fatal("expected tracked tool feedback to be cleared after final reply")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateAndAddConnection_RespectsMaxConnectionsConcurrently(t *testing.T) {
|
||||
ch := newTestPicoChannel(t)
|
||||
|
||||
|
|
@ -123,6 +289,167 @@ func TestBroadcastToSession_TargetsOnlyRequestedSession(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestSendMedia_ResolvesMediaBeforeDelivery(t *testing.T) {
|
||||
ch := newTestPicoChannel(t)
|
||||
store := media.NewFileMediaStore()
|
||||
ch.SetMediaStore(store)
|
||||
|
||||
if err := ch.Start(context.Background()); err != nil {
|
||||
t.Fatalf("Start() error = %v", err)
|
||||
}
|
||||
defer ch.Stop(context.Background())
|
||||
|
||||
localPath := filepath.Join(t.TempDir(), "report.txt")
|
||||
if err := os.WriteFile(localPath, []byte("attachment body"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
ref, err := store.Store(localPath, media.MediaMeta{
|
||||
Filename: "report.txt",
|
||||
ContentType: "text/plain",
|
||||
}, "test-scope")
|
||||
if err != nil {
|
||||
t.Fatalf("Store() error = %v", err)
|
||||
}
|
||||
|
||||
closedConn := &picoConn{id: "closed", sessionID: "sess-1"}
|
||||
closedConn.closed.Store(true)
|
||||
ch.addConnForTest(closedConn)
|
||||
|
||||
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||
ChatID: "pico:sess-1",
|
||||
Parts: []bus.MediaPart{{
|
||||
Ref: ref,
|
||||
Type: "file",
|
||||
Filename: "report.txt",
|
||||
ContentType: "text/plain",
|
||||
}},
|
||||
})
|
||||
if !errors.Is(err, channels.ErrSendFailed) {
|
||||
t.Fatalf("SendMedia() error = %v, want ErrSendFailed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendMedia_DismissesTrackedToolFeedbackMessage(t *testing.T) {
|
||||
ch := newTestPicoChannel(t)
|
||||
store := media.NewFileMediaStore()
|
||||
ch.SetMediaStore(store)
|
||||
|
||||
if err := ch.Start(context.Background()); err != nil {
|
||||
t.Fatalf("Start() error = %v", err)
|
||||
}
|
||||
defer ch.Stop(context.Background())
|
||||
|
||||
clientConn, received, cleanup := newTestPicoWebSocket(t)
|
||||
defer cleanup()
|
||||
ch.addConnForTest(&picoConn{id: "conn-1", conn: clientConn, sessionID: "sess-1"})
|
||||
|
||||
localPath := filepath.Join(t.TempDir(), "report.txt")
|
||||
if err := os.WriteFile(localPath, []byte("attachment body"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
ref, err := store.Store(localPath, media.MediaMeta{
|
||||
Filename: "report.txt",
|
||||
ContentType: "text/plain",
|
||||
}, "test-scope")
|
||||
if err != nil {
|
||||
t.Fatalf("Store() error = %v", err)
|
||||
}
|
||||
|
||||
ch.RecordToolFeedbackMessage("pico:sess-1", "msg-progress", "🔧 `read_file`")
|
||||
|
||||
var deleted struct {
|
||||
chatID string
|
||||
messageID string
|
||||
}
|
||||
ch.deleteMessageFn = func(_ context.Context, chatID string, messageID string) error {
|
||||
deleted.chatID = chatID
|
||||
deleted.messageID = messageID
|
||||
return nil
|
||||
}
|
||||
|
||||
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||
ChatID: "pico:sess-1",
|
||||
Parts: []bus.MediaPart{{
|
||||
Ref: ref,
|
||||
Type: "file",
|
||||
Filename: "report.txt",
|
||||
ContentType: "text/plain",
|
||||
}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMedia() error = %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case msg := <-received:
|
||||
if msg.Type != TypeMessageCreate {
|
||||
t.Fatalf("message type = %q, want %q", msg.Type, TypeMessageCreate)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected media message to be delivered")
|
||||
}
|
||||
|
||||
if deleted.chatID != "pico:sess-1" || deleted.messageID != "msg-progress" {
|
||||
t.Fatalf("unexpected delete target: %+v", deleted)
|
||||
}
|
||||
if _, ok := ch.currentToolFeedbackMessage("pico:sess-1"); ok {
|
||||
t.Fatal("expected tracked tool feedback to be cleared after media delivery")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPicoDownloadURLForRef(t *testing.T) {
|
||||
got, err := picoDownloadURLForRef("media://attachment-1")
|
||||
if err != nil {
|
||||
t.Fatalf("picoDownloadURLForRef() error = %v", err)
|
||||
}
|
||||
if got != "/pico/media/attachment-1" {
|
||||
t.Fatalf("picoDownloadURLForRef() = %q, want %q", got, "/pico/media/attachment-1")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleMediaDownload_ServesStoredFile(t *testing.T) {
|
||||
ch := newTestPicoChannel(t)
|
||||
store := media.NewFileMediaStore()
|
||||
ch.SetMediaStore(store)
|
||||
|
||||
if err := ch.Start(context.Background()); err != nil {
|
||||
t.Fatalf("Start() error = %v", err)
|
||||
}
|
||||
defer ch.Stop(context.Background())
|
||||
|
||||
localPath := filepath.Join(t.TempDir(), "report.txt")
|
||||
if err := os.WriteFile(localPath, []byte("downloadable"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
ref, err := store.Store(localPath, media.MediaMeta{
|
||||
Filename: "report.txt",
|
||||
ContentType: "text/plain",
|
||||
}, "test-scope")
|
||||
if err != nil {
|
||||
t.Fatalf("Store() error = %v", err)
|
||||
}
|
||||
|
||||
refID := strings.TrimPrefix(ref, "media://")
|
||||
req := httptest.NewRequest("GET", "/pico/media/"+refID, nil)
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
ch.ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != 200 {
|
||||
t.Fatalf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
if body := rec.Body.String(); body != "downloadable" {
|
||||
t.Fatalf("body = %q, want %q", body, "downloadable")
|
||||
}
|
||||
if got := rec.Header().Get("Content-Type"); got != "text/plain" {
|
||||
t.Fatalf("Content-Type = %q, want %q", got, "text/plain")
|
||||
}
|
||||
}
|
||||
|
||||
func (c *PicoChannel) addConnForTest(pc *picoConn) {
|
||||
c.connsMu.Lock()
|
||||
defer c.connsMu.Unlock()
|
||||
|
|
@ -143,3 +470,39 @@ func (c *PicoChannel) addConnForTest(pc *picoConn) {
|
|||
}
|
||||
bySession[pc.id] = pc
|
||||
}
|
||||
|
||||
func newTestPicoWebSocket(t *testing.T) (*websocket.Conn, <-chan PicoMessage, func()) {
|
||||
t.Helper()
|
||||
|
||||
received := make(chan PicoMessage, 4)
|
||||
upgrader := websocket.Upgrader{}
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
t.Errorf("Upgrade() error = %v", err)
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
for {
|
||||
var msg PicoMessage
|
||||
if err := conn.ReadJSON(&msg); err != nil {
|
||||
return
|
||||
}
|
||||
received <- msg
|
||||
}
|
||||
}))
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
clientConn, resp, err := websocket.DefaultDialer.Dial(wsURL, nil)
|
||||
if err != nil {
|
||||
server.Close()
|
||||
t.Fatalf("Dial() error = %v", err)
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
clientConn.Close()
|
||||
server.Close()
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
return clientConn, received, cleanup
|
||||
}
|
||||
|
|
|
|||
|
|
@ -12,14 +12,13 @@ const (
|
|||
// TypeMessageCreate is sent from server to client.
|
||||
TypeMessageCreate = "message.create"
|
||||
TypeMessageUpdate = "message.update"
|
||||
TypeMessageDelete = "message.delete"
|
||||
TypeMediaCreate = "media.create"
|
||||
TypeTypingStart = "typing.start"
|
||||
TypeTypingStop = "typing.stop"
|
||||
TypeError = "error"
|
||||
TypePong = "pong"
|
||||
|
||||
PicoTokenPrefix = "pico-"
|
||||
|
||||
PayloadKeyContent = "content"
|
||||
PayloadKeyThought = "thought"
|
||||
|
||||
|
|
|
|||
|
|
@ -66,6 +66,10 @@ func (c *TelegramChannel) startCommandRegistration(ctx context.Context, defs []c
|
|||
if register == nil {
|
||||
register = c.RegisterCommands
|
||||
}
|
||||
delayFn := c.commandRegDelayFn
|
||||
if delayFn == nil {
|
||||
delayFn = commandRegistrationDelay
|
||||
}
|
||||
|
||||
regCtx, cancel := context.WithCancel(ctx)
|
||||
c.commandRegCancel = cancel
|
||||
|
|
@ -91,7 +95,7 @@ func (c *TelegramChannel) startCommandRegistration(ctx context.Context, defs []c
|
|||
return
|
||||
}
|
||||
|
||||
delay := commandRegistrationDelay(attempt)
|
||||
delay := delayFn(attempt)
|
||||
logger.WarnCF("telegram", "Telegram command registration failed; will retry", map[string]any{
|
||||
"error": err.Error(),
|
||||
"retry_after": delay.String(),
|
||||
|
|
|
|||
|
|
@ -31,14 +31,12 @@ func TestStartCommandRegistration_DoesNotBlock(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestStartCommandRegistration_RetriesUntilSuccessThenStops(t *testing.T) {
|
||||
ch := &TelegramChannel{}
|
||||
ch := &TelegramChannel{
|
||||
commandRegDelayFn: func(int) time.Duration { return 5 * time.Millisecond },
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
origBackoff := commandRegistrationBackoff
|
||||
commandRegistrationBackoff = []time.Duration{5 * time.Millisecond}
|
||||
defer func() { commandRegistrationBackoff = origBackoff }()
|
||||
|
||||
var attempts atomic.Int32
|
||||
ch.registerFunc = func(context.Context, []commands.Definition) error {
|
||||
n := attempts.Add(1)
|
||||
|
|
@ -69,12 +67,10 @@ func TestStartCommandRegistration_RetriesUntilSuccessThenStops(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestStartCommandRegistration_StopsAfterCancel(t *testing.T) {
|
||||
ch := &TelegramChannel{}
|
||||
ch := &TelegramChannel{
|
||||
commandRegDelayFn: func(int) time.Duration { return 5 * time.Millisecond },
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
origBackoff := commandRegistrationBackoff
|
||||
commandRegistrationBackoff = []time.Duration{5 * time.Millisecond}
|
||||
defer func() { commandRegistrationBackoff = origBackoff }()
|
||||
defer cancel()
|
||||
|
||||
var attempts atomic.Int32
|
||||
|
|
|
|||
|
|
@ -45,16 +45,18 @@ var (
|
|||
|
||||
type TelegramChannel struct {
|
||||
*channels.BaseChannel
|
||||
bot *telego.Bot
|
||||
bh *th.BotHandler
|
||||
bc *config.Channel
|
||||
chatIDs map[string]int64
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
tgCfg *config.TelegramSettings
|
||||
bot *telego.Bot
|
||||
bh *th.BotHandler
|
||||
bc *config.Channel
|
||||
chatIDs map[string]int64
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
tgCfg *config.TelegramSettings
|
||||
progress *channels.ToolFeedbackAnimator
|
||||
|
||||
registerFunc func(context.Context, []commands.Definition) error
|
||||
commandRegCancel context.CancelFunc
|
||||
registerFunc func(context.Context, []commands.Definition) error
|
||||
commandRegDelayFn func(int) time.Duration
|
||||
commandRegCancel context.CancelFunc
|
||||
}
|
||||
|
||||
func NewTelegramChannel(
|
||||
|
|
@ -104,13 +106,15 @@ func NewTelegramChannel(
|
|||
channels.WithReasoningChannelID(bc.ReasoningChannelID),
|
||||
)
|
||||
|
||||
return &TelegramChannel{
|
||||
ch := &TelegramChannel{
|
||||
BaseChannel: base,
|
||||
bot: bot,
|
||||
bc: bc,
|
||||
chatIDs: make(map[string]int64),
|
||||
tgCfg: telegramCfg,
|
||||
}, nil
|
||||
}
|
||||
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) Start(ctx context.Context) error {
|
||||
|
|
@ -168,6 +172,9 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
|
|||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
if c.progress != nil {
|
||||
c.progress.StopAll()
|
||||
}
|
||||
if c.commandRegCancel != nil {
|
||||
c.commandRegCancel()
|
||||
}
|
||||
|
|
@ -191,12 +198,36 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]
|
|||
return nil, nil
|
||||
}
|
||||
|
||||
isToolFeedback := outboundMessageIsToolFeedback(msg)
|
||||
toolFeedbackContent := msg.Content
|
||||
if isToolFeedback {
|
||||
toolFeedbackContent = fitToolFeedbackForTelegram(msg.Content, useMarkdownV2, 4096)
|
||||
}
|
||||
trackedChatID := telegramToolFeedbackChatKey(msg.ChatID, &msg.Context)
|
||||
if isToolFeedback {
|
||||
if msgID, handled, err := c.progress.Update(ctx, trackedChatID, toolFeedbackContent); handled {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{msgID}, nil
|
||||
}
|
||||
}
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(trackedChatID)
|
||||
if !isToolFeedback {
|
||||
if msgIDs, handled := c.finalizeToolFeedbackMessageForChat(ctx, trackedChatID, msg); handled {
|
||||
return msgIDs, nil
|
||||
}
|
||||
}
|
||||
|
||||
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
||||
// so msg.Content is guaranteed to be within that limit. We still need to
|
||||
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
||||
replyToID := msg.ReplyToMessageID
|
||||
var messageIDs []string
|
||||
queue := []string{msg.Content}
|
||||
if isToolFeedback {
|
||||
queue = []string{channels.InitialAnimatedToolFeedbackContent(toolFeedbackContent)}
|
||||
}
|
||||
for len(queue) > 0 {
|
||||
chunk := queue[0]
|
||||
queue = queue[1:]
|
||||
|
|
@ -204,6 +235,13 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]
|
|||
content := parseContent(chunk, useMarkdownV2)
|
||||
|
||||
if len([]rune(content)) > 4096 {
|
||||
if isToolFeedback {
|
||||
fittedChunk := fitToolFeedbackForTelegram(chunk, useMarkdownV2, 4096)
|
||||
if fittedChunk != "" && fittedChunk != chunk {
|
||||
queue = append([]string{fittedChunk}, queue...)
|
||||
continue
|
||||
}
|
||||
}
|
||||
runeChunk := []rune(chunk)
|
||||
ratio := float64(len(runeChunk)) / float64(len([]rune(content)))
|
||||
smallerLen := int(float64(4096) * ratio * 0.95) // 5% safety margin
|
||||
|
|
@ -270,6 +308,12 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]
|
|||
replyToID = ""
|
||||
}
|
||||
|
||||
if isToolFeedback && len(messageIDs) > 0 {
|
||||
c.RecordToolFeedbackMessage(trackedChatID, messageIDs[0], toolFeedbackContent)
|
||||
} else if !isToolFeedback && hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, trackedChatID, trackedMsgID)
|
||||
}
|
||||
|
||||
return messageIDs, nil
|
||||
}
|
||||
|
||||
|
|
@ -437,6 +481,89 @@ func (c *TelegramChannel) DeleteMessage(ctx context.Context, chatID string, mess
|
|||
})
|
||||
}
|
||||
|
||||
func outboundMessageIsToolFeedback(msg bus.OutboundMessage) bool {
|
||||
if len(msg.Context.Raw) == 0 {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(strings.TrimSpace(msg.Context.Raw["message_kind"]), "tool_feedback")
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) currentToolFeedbackMessage(chatID string) (string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", false
|
||||
}
|
||||
return c.progress.Current(chatID)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) takeToolFeedbackMessage(chatID string) (string, string, bool) {
|
||||
if c.progress == nil {
|
||||
return "", "", false
|
||||
}
|
||||
return c.progress.Take(chatID)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) RecordToolFeedbackMessage(chatID, messageID, content string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Record(chatID, messageID, content)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) ClearToolFeedbackMessage(chatID string) {
|
||||
if c.progress == nil {
|
||||
return
|
||||
}
|
||||
c.progress.Clear(chatID)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) DismissToolFeedbackMessage(ctx context.Context, chatID string) {
|
||||
msgID, ok := c.currentToolFeedbackMessage(chatID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, chatID, msgID)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) dismissTrackedToolFeedbackMessage(ctx context.Context, chatID, messageID string) {
|
||||
if strings.TrimSpace(chatID) == "" || strings.TrimSpace(messageID) == "" {
|
||||
return
|
||||
}
|
||||
c.ClearToolFeedbackMessage(chatID)
|
||||
_ = c.DeleteMessage(ctx, chatID, messageID)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) finalizeTrackedToolFeedbackMessage(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
content string,
|
||||
editFn func(context.Context, string, string, string) error,
|
||||
) ([]string, bool) {
|
||||
msgID, baseContent, ok := c.takeToolFeedbackMessage(chatID)
|
||||
if !ok || editFn == nil {
|
||||
return nil, false
|
||||
}
|
||||
if err := editFn(ctx, chatID, msgID, content); err != nil {
|
||||
c.RecordToolFeedbackMessage(chatID, msgID, baseContent)
|
||||
return nil, false
|
||||
}
|
||||
return []string{msgID}, true
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) FinalizeToolFeedbackMessage(ctx context.Context, msg bus.OutboundMessage) ([]string, bool) {
|
||||
if outboundMessageIsToolFeedback(msg) {
|
||||
return nil, false
|
||||
}
|
||||
return c.finalizeToolFeedbackMessageForChat(ctx, telegramToolFeedbackChatKey(msg.ChatID, &msg.Context), msg)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) finalizeToolFeedbackMessageForChat(
|
||||
ctx context.Context,
|
||||
chatID string,
|
||||
msg bus.OutboundMessage,
|
||||
) ([]string, bool) {
|
||||
return c.finalizeTrackedToolFeedbackMessage(ctx, chatID, msg.Content, c.EditMessage)
|
||||
}
|
||||
|
||||
// SendPlaceholder implements channels.PlaceholderCapable.
|
||||
// It sends a placeholder message (e.g. "Thinking... 💭") that will later be
|
||||
// edited to the actual response via EditMessage (channels.MessageEditor).
|
||||
|
|
@ -468,6 +595,8 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
|||
if !c.IsRunning() {
|
||||
return nil, channels.ErrNotRunning
|
||||
}
|
||||
trackedChatID := telegramToolFeedbackChatKey(msg.ChatID, &msg.Context)
|
||||
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(trackedChatID)
|
||||
|
||||
chatID, threadID, err := resolveTelegramOutboundTarget(msg.ChatID, &msg.Context)
|
||||
if err != nil {
|
||||
|
|
@ -576,6 +705,10 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
|||
}
|
||||
}
|
||||
|
||||
if hasTrackedMsg {
|
||||
c.dismissTrackedToolFeedbackMessage(ctx, trackedChatID, trackedMsgID)
|
||||
}
|
||||
|
||||
return messageIDs, nil
|
||||
}
|
||||
|
||||
|
|
@ -947,6 +1080,60 @@ func parseContent(text string, useMarkdownV2 bool) string {
|
|||
return markdownToTelegramHTML(text)
|
||||
}
|
||||
|
||||
func fitToolFeedbackForTelegram(content string, useMarkdownV2 bool, maxParsedLen int) string {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" || maxParsedLen <= 0 {
|
||||
return ""
|
||||
}
|
||||
animationSafeLen := maxParsedLen - channels.MaxToolFeedbackAnimationFrameLength()
|
||||
if animationSafeLen <= 0 {
|
||||
animationSafeLen = maxParsedLen
|
||||
}
|
||||
if len([]rune(parseContent(content, useMarkdownV2))) <= animationSafeLen {
|
||||
return content
|
||||
}
|
||||
|
||||
low := 1
|
||||
high := len([]rune(content))
|
||||
best := utils.Truncate(content, 1)
|
||||
|
||||
for low <= high {
|
||||
mid := (low + high) / 2
|
||||
candidate := utils.FitToolFeedbackMessage(content, mid)
|
||||
if candidate == "" {
|
||||
high = mid - 1
|
||||
continue
|
||||
}
|
||||
if len([]rune(parseContent(candidate, useMarkdownV2))) <= animationSafeLen {
|
||||
best = candidate
|
||||
low = mid + 1
|
||||
continue
|
||||
}
|
||||
high = mid - 1
|
||||
}
|
||||
|
||||
return best
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) PrepareToolFeedbackMessageContent(content string) string {
|
||||
if c == nil || c.tgCfg == nil {
|
||||
return strings.TrimSpace(content)
|
||||
}
|
||||
return fitToolFeedbackForTelegram(content, c.tgCfg.UseMarkdownV2, 4096)
|
||||
}
|
||||
|
||||
func telegramToolFeedbackChatKey(chatID string, outboundCtx *bus.InboundContext) string {
|
||||
resolvedChatID, threadID, err := resolveTelegramOutboundTarget(chatID, outboundCtx)
|
||||
if err != nil || threadID == 0 {
|
||||
return strings.TrimSpace(chatID)
|
||||
}
|
||||
return fmt.Sprintf("%d/%d", resolvedChatID, threadID)
|
||||
}
|
||||
|
||||
func (c *TelegramChannel) ToolFeedbackMessageChatID(chatID string, outboundCtx *bus.InboundContext) string {
|
||||
return telegramToolFeedbackChatKey(chatID, outboundCtx)
|
||||
}
|
||||
|
||||
// parseTelegramChatID splits "chatID/threadID" into its components.
|
||||
// Returns threadID=0 when no "/" is present (non-forum messages).
|
||||
func parseTelegramChatID(chatID string) (int64, int, error) {
|
||||
|
|
@ -1097,7 +1284,7 @@ func (c *TelegramChannel) BeginStream(ctx context.Context, chatID string) (chann
|
|||
return nil, fmt.Errorf("streaming disabled in config")
|
||||
}
|
||||
|
||||
cid, _, err := parseTelegramChatID(chatID)
|
||||
cid, threadID, err := parseTelegramChatID(chatID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -1106,6 +1293,7 @@ func (c *TelegramChannel) BeginStream(ctx context.Context, chatID string) (chann
|
|||
return &telegramStreamer{
|
||||
bot: c.bot,
|
||||
chatID: cid,
|
||||
threadID: threadID,
|
||||
draftID: cryptoRandInt(),
|
||||
throttleInterval: time.Duration(streamCfg.ThrottleSeconds) * time.Second,
|
||||
minGrowth: streamCfg.MinGrowthChars,
|
||||
|
|
@ -1118,6 +1306,7 @@ func (c *TelegramChannel) BeginStream(ctx context.Context, chatID string) (chann
|
|||
type telegramStreamer struct {
|
||||
bot *telego.Bot
|
||||
chatID int64
|
||||
threadID int
|
||||
draftID int
|
||||
throttleInterval time.Duration
|
||||
minGrowth int
|
||||
|
|
@ -1145,10 +1334,11 @@ func (s *telegramStreamer) Update(ctx context.Context, content string) error {
|
|||
htmlContent := markdownToTelegramHTML(content)
|
||||
|
||||
err := s.bot.SendMessageDraft(ctx, &telego.SendMessageDraftParams{
|
||||
ChatID: s.chatID,
|
||||
DraftID: s.draftID,
|
||||
Text: htmlContent,
|
||||
ParseMode: telego.ModeHTML,
|
||||
ChatID: s.chatID,
|
||||
MessageThreadID: s.threadID,
|
||||
DraftID: s.draftID,
|
||||
Text: htmlContent,
|
||||
ParseMode: telego.ModeHTML,
|
||||
})
|
||||
if err != nil {
|
||||
// First error → degrade silently (e.g. no forum mode)
|
||||
|
|
@ -1167,6 +1357,7 @@ func (s *telegramStreamer) Update(ctx context.Context, content string) error {
|
|||
func (s *telegramStreamer) Finalize(ctx context.Context, content string) error {
|
||||
htmlContent := markdownToTelegramHTML(content)
|
||||
tgMsg := tu.Message(tu.ID(s.chatID), htmlContent)
|
||||
tgMsg.MessageThreadID = s.threadID
|
||||
tgMsg.ParseMode = telego.ModeHTML
|
||||
|
||||
if _, err := s.bot.SendMessage(ctx, tgMsg); err != nil {
|
||||
|
|
|
|||
|
|
@ -108,7 +108,7 @@ func TestHandleMessage_GroupMentionOnly_BotCommandEntity(t *testing.T) {
|
|||
t.Fatalf("handleMessage error: %v", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Microsecond)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||
defer cancel()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
|
|
|
|||
|
|
@ -98,8 +98,12 @@ func (s *multipartRecordingConstructor) MultipartRequest(
|
|||
|
||||
// successResponse returns a ta.Response that telego will treat as a successful SendMessage.
|
||||
func successResponse(t *testing.T) *ta.Response {
|
||||
return successResponseWithMessageID(t, 1)
|
||||
}
|
||||
|
||||
func successResponseWithMessageID(t *testing.T, messageID int) *ta.Response {
|
||||
t.Helper()
|
||||
msg := &telego.Message{MessageID: 1}
|
||||
msg := &telego.Message{MessageID: messageID}
|
||||
b, err := json.Marshal(msg)
|
||||
require.NoError(t, err)
|
||||
return &ta.Response{Ok: true, Result: b}
|
||||
|
|
@ -142,6 +146,7 @@ func newTestChannelWithConstructor(
|
|||
chatIDs: make(map[string]int64),
|
||||
bc: &config.Channel{Type: config.ChannelTelegram, Enabled: true},
|
||||
tgCfg: &config.TelegramSettings{},
|
||||
progress: channels.NewToolFeedbackAnimator(nil),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -266,6 +271,176 @@ func TestSend_ShortMessage_SingleCall(t *testing.T) {
|
|||
assert.Len(t, caller.calls, 1, "short message should result in exactly one SendMessage call")
|
||||
}
|
||||
|
||||
func TestSend_NonToolFeedbackDeletesTrackedProgressMessage(t *testing.T) {
|
||||
caller := &stubCaller{
|
||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||
switch {
|
||||
case strings.Contains(url, "editMessageText"):
|
||||
return successResponseWithMessageID(t, 1), nil
|
||||
default:
|
||||
t.Fatalf("unexpected API call: %s", url)
|
||||
return nil, nil
|
||||
}
|
||||
},
|
||||
}
|
||||
ch := newTestChannel(t, caller)
|
||||
ch.RecordToolFeedbackMessage("12345", "1", "🔧 `read_file`")
|
||||
|
||||
ids, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "12345",
|
||||
Content: "final reply",
|
||||
})
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, []string{"1"}, ids)
|
||||
require.Len(t, caller.calls, 1)
|
||||
assert.Contains(t, caller.calls[0].URL, "editMessageText")
|
||||
_, ok := ch.currentToolFeedbackMessage("12345")
|
||||
assert.False(t, ok, "tracked tool feedback should be cleared after final reply")
|
||||
}
|
||||
|
||||
func TestSend_ToolFeedbackTrackingIsTopicScoped(t *testing.T) {
|
||||
nextMessageID := 0
|
||||
caller := &stubCaller{
|
||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||
nextMessageID++
|
||||
return successResponseWithMessageID(t, nextMessageID), nil
|
||||
},
|
||||
}
|
||||
ch := newTestChannel(t, caller)
|
||||
|
||||
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "-1001234567890",
|
||||
Content: "🔧 `read_file`",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "telegram",
|
||||
ChatID: "-1001234567890",
|
||||
TopicID: "42",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, ok := ch.currentToolFeedbackMessage("-1001234567890")
|
||||
assert.False(t, ok, "base chat should not track topic-specific tool feedback")
|
||||
|
||||
msgID, ok := ch.currentToolFeedbackMessage("-1001234567890/42")
|
||||
require.True(t, ok, "topic chat should track tool feedback")
|
||||
assert.Equal(t, "1", msgID)
|
||||
}
|
||||
|
||||
func TestSend_TopicReplyDoesNotFinalizeDifferentTopicToolFeedback(t *testing.T) {
|
||||
nextMessageID := 0
|
||||
caller := &stubCaller{
|
||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||
nextMessageID++
|
||||
return successResponseWithMessageID(t, nextMessageID), nil
|
||||
},
|
||||
}
|
||||
ch := newTestChannel(t, caller)
|
||||
|
||||
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "-1001234567890",
|
||||
Content: "🔧 `read_file`",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "telegram",
|
||||
ChatID: "-1001234567890",
|
||||
TopicID: "42",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
ids, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "-1001234567890",
|
||||
Content: "final reply in another topic",
|
||||
Context: bus.InboundContext{
|
||||
Channel: "telegram",
|
||||
ChatID: "-1001234567890",
|
||||
TopicID: "43",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, caller.calls, 2)
|
||||
assert.Equal(t, []string{"2"}, ids)
|
||||
assert.Contains(t, caller.calls[1].URL, "sendMessage")
|
||||
assert.NotContains(t, caller.calls[1].URL, "editMessageText")
|
||||
|
||||
_, ok := ch.currentToolFeedbackMessage("-1001234567890/42")
|
||||
assert.True(t, ok, "tool feedback in the original topic should remain tracked")
|
||||
}
|
||||
|
||||
func TestFinalizeTrackedToolFeedbackMessage_StopsTrackingBeforeEdit(t *testing.T) {
|
||||
ch := newTestChannel(t, &stubCaller{
|
||||
callFn: func(context.Context, string, *ta.RequestData) (*ta.Response, error) {
|
||||
t.Fatal("unexpected API call")
|
||||
return nil, nil
|
||||
},
|
||||
})
|
||||
ch.RecordToolFeedbackMessage("12345", "1", "🔧 `read_file`")
|
||||
|
||||
msgIDs, handled := ch.finalizeTrackedToolFeedbackMessage(
|
||||
context.Background(),
|
||||
"12345",
|
||||
"final reply",
|
||||
func(_ context.Context, chatID, messageID, content string) error {
|
||||
_, ok := ch.currentToolFeedbackMessage(chatID)
|
||||
assert.False(t, ok, "tracked tool feedback should be stopped before edit")
|
||||
assert.Equal(t, "12345", chatID)
|
||||
assert.Equal(t, "1", messageID)
|
||||
assert.Equal(t, "final reply", content)
|
||||
return nil
|
||||
},
|
||||
)
|
||||
|
||||
assert.True(t, handled)
|
||||
assert.Equal(t, []string{"1"}, msgIDs)
|
||||
}
|
||||
|
||||
func TestSend_ToolFeedbackStaysSingleMessageAfterHTMLExpansion(t *testing.T) {
|
||||
caller := &stubCaller{
|
||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||
return successResponse(t), nil
|
||||
},
|
||||
}
|
||||
ch := newTestChannel(t, caller)
|
||||
|
||||
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
ChatID: "12345",
|
||||
Content: "🔧 `read_file`\n" + strings.Repeat("<", 2000),
|
||||
Context: bus.InboundContext{
|
||||
Channel: "telegram",
|
||||
ChatID: "12345",
|
||||
Raw: map[string]string{
|
||||
"message_kind": "tool_feedback",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Len(t, caller.calls, 1, "tool feedback should stay a single Telegram message after HTML escaping")
|
||||
}
|
||||
|
||||
func TestFitToolFeedbackForTelegram_ReservesAnimationFrame(t *testing.T) {
|
||||
content := "🔧 `read_file`\n" + strings.Repeat("a", 4096)
|
||||
|
||||
fitted := fitToolFeedbackForTelegram(content, false, 4096)
|
||||
animated := strings.Replace(
|
||||
fitted,
|
||||
"`\n",
|
||||
strings.Repeat(".", channels.MaxToolFeedbackAnimationFrameLength())+"`\n",
|
||||
1,
|
||||
)
|
||||
|
||||
if got := len([]rune(parseContent(animated, false))); got > 4096 {
|
||||
t.Fatalf("animated parsed length = %d, want <= 4096", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSend_LongMessage_SingleCall(t *testing.T) {
|
||||
// With WithMaxMessageLength(4000), the Manager pre-splits messages before
|
||||
// they reach Send(). A message at exactly 4000 chars should go through
|
||||
|
|
@ -560,6 +735,58 @@ func TestSend_UsesContextTopicIDWhenChatIDDoesNotIncludeThread(t *testing.T) {
|
|||
assert.Equal(t, "Hello from topic context", params.Text)
|
||||
}
|
||||
|
||||
func TestBeginStream_UpdateUsesForumThreadID(t *testing.T) {
|
||||
caller := &stubCaller{
|
||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||
return &ta.Response{Ok: true, Result: []byte("true")}, nil
|
||||
},
|
||||
}
|
||||
ch := newTestChannel(t, caller)
|
||||
ch.tgCfg.Streaming.Enabled = true
|
||||
|
||||
streamer, err := ch.BeginStream(context.Background(), "-1001234567890/42")
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, streamer.Update(context.Background(), "partial"))
|
||||
require.Len(t, caller.calls, 1)
|
||||
assert.Contains(t, caller.calls[0].URL, "sendMessageDraft")
|
||||
|
||||
var params struct {
|
||||
ChatID int64 `json:"chat_id"`
|
||||
MessageThreadID int `json:"message_thread_id"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(caller.calls[0].Data.BodyRaw, ¶ms))
|
||||
assert.Equal(t, int64(-1001234567890), params.ChatID)
|
||||
assert.Equal(t, 42, params.MessageThreadID)
|
||||
assert.Equal(t, "partial", params.Text)
|
||||
}
|
||||
|
||||
func TestBeginStream_FinalizeUsesForumThreadID(t *testing.T) {
|
||||
caller := &stubCaller{
|
||||
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||
return successResponse(t), nil
|
||||
},
|
||||
}
|
||||
ch := newTestChannel(t, caller)
|
||||
ch.tgCfg.Streaming.Enabled = true
|
||||
|
||||
streamer, err := ch.BeginStream(context.Background(), "-1001234567890/42")
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, streamer.Finalize(context.Background(), "final"))
|
||||
require.Len(t, caller.calls, 1)
|
||||
assert.Contains(t, caller.calls[0].URL, "sendMessage")
|
||||
|
||||
var params struct {
|
||||
ChatID int64 `json:"chat_id"`
|
||||
MessageThreadID int `json:"message_thread_id"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(caller.calls[0].Data.BodyRaw, ¶ms))
|
||||
assert.Equal(t, int64(-1001234567890), params.ChatID)
|
||||
assert.Equal(t, 42, params.MessageThreadID)
|
||||
assert.Equal(t, "final", params.Text)
|
||||
}
|
||||
|
||||
func TestHandleMessage_ForumTopic_SetsMetadata(t *testing.T) {
|
||||
messageBus := bus.NewMessageBus()
|
||||
ch := &TelegramChannel{
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue