Merge branch 'main' into fix/config_channel_issue
This commit is contained in:
commit
5748fe2d22
95 changed files with 8306 additions and 8524 deletions
5
.github/workflows/pr.yml
vendored
5
.github/workflows/pr.yml
vendored
|
|
@ -23,10 +23,13 @@ jobs:
|
||||||
uses: golangci/golangci-lint-action@v9
|
uses: golangci/golangci-lint-action@v9
|
||||||
with:
|
with:
|
||||||
version: v2.10.1
|
version: v2.10.1
|
||||||
|
args: --build-tags=goolm,stdjson
|
||||||
|
|
||||||
vuln_check:
|
vuln_check:
|
||||||
name: Security Check
|
name: Security Check
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
|
env:
|
||||||
|
GOFLAGS: -tags=goolm,stdjson
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Checkout
|
||||||
uses: actions/checkout@v6
|
uses: actions/checkout@v6
|
||||||
|
|
@ -59,4 +62,4 @@ jobs:
|
||||||
run: go generate ./...
|
run: go generate ./...
|
||||||
|
|
||||||
- name: Run go test
|
- name: Run go test
|
||||||
run: go test ./...
|
run: go test -tags goolm,stdjson ./...
|
||||||
|
|
|
||||||
|
|
@ -15,6 +15,7 @@ builds:
|
||||||
env:
|
env:
|
||||||
- CGO_ENABLED=0
|
- CGO_ENABLED=0
|
||||||
tags:
|
tags:
|
||||||
|
- goolm
|
||||||
- stdjson
|
- stdjson
|
||||||
ldflags:
|
ldflags:
|
||||||
- -s -w
|
- -s -w
|
||||||
|
|
@ -57,6 +58,7 @@ builds:
|
||||||
env:
|
env:
|
||||||
- CGO_ENABLED=0
|
- CGO_ENABLED=0
|
||||||
tags:
|
tags:
|
||||||
|
- goolm
|
||||||
- stdjson
|
- stdjson
|
||||||
ldflags:
|
ldflags:
|
||||||
- -s -w
|
- -s -w
|
||||||
|
|
@ -95,6 +97,7 @@ builds:
|
||||||
env:
|
env:
|
||||||
- CGO_ENABLED=0
|
- CGO_ENABLED=0
|
||||||
tags:
|
tags:
|
||||||
|
- goolm
|
||||||
- stdjson
|
- stdjson
|
||||||
ldflags:
|
ldflags:
|
||||||
- -s -w
|
- -s -w
|
||||||
|
|
|
||||||
62
Makefile
62
Makefile
|
|
@ -17,7 +17,13 @@ LDFLAGS=-X $(CONFIG_PKG).Version=$(VERSION) -X $(CONFIG_PKG).GitCommit=$(GIT_COM
|
||||||
# Go variables
|
# Go variables
|
||||||
GO?=CGO_ENABLED=0 go
|
GO?=CGO_ENABLED=0 go
|
||||||
WEB_GO?=$(GO)
|
WEB_GO?=$(GO)
|
||||||
GOFLAGS?=-v -tags stdjson
|
GO_BUILD_TAGS?=goolm,stdjson
|
||||||
|
GOFLAGS?=-v -tags $(GO_BUILD_TAGS)
|
||||||
|
comma:=,
|
||||||
|
empty:=
|
||||||
|
space:=$(empty) $(empty)
|
||||||
|
GO_BUILD_TAGS_NO_GOOLM:=$(subst $(space),$(comma),$(strip $(filter-out goolm,$(subst $(comma),$(space),$(GO_BUILD_TAGS)))))
|
||||||
|
GOFLAGS_NO_GOOLM?=-v -tags $(GO_BUILD_TAGS_NO_GOOLM)
|
||||||
|
|
||||||
# Patch MIPS LE ELF e_flags (offset 36) for NaN2008-only kernels (e.g. Ingenic X2600).
|
# Patch MIPS LE ELF e_flags (offset 36) for NaN2008-only kernels (e.g. Ingenic X2600).
|
||||||
#
|
#
|
||||||
|
|
@ -130,15 +136,15 @@ build-whatsapp-native: generate
|
||||||
## @echo "Building $(BINARY_NAME) with WhatsApp native for $(PLATFORM)/$(ARCH)..."
|
## @echo "Building $(BINARY_NAME) with WhatsApp native for $(PLATFORM)/$(ARCH)..."
|
||||||
@echo "Building for multiple platforms..."
|
@echo "Building for multiple platforms..."
|
||||||
@mkdir -p $(BUILD_DIR)
|
@mkdir -p $(BUILD_DIR)
|
||||||
GOOS=linux GOARCH=amd64 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-amd64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=amd64 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-amd64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=arm GOARM=7 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm GOARM=7 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=arm64 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm64 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=loong64 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-loong64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=loong64 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-loong64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=riscv64 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-riscv64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=riscv64 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-riscv64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=mipsle GOMIPS=softfloat $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle ./$(CMD_DIR)
|
GOOS=linux GOARCH=mipsle GOMIPS=softfloat $(GO) build -tags $(GO_BUILD_TAGS_NO_GOOLM),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle ./$(CMD_DIR)
|
||||||
$(call PATCH_MIPS_FLAGS,$(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle)
|
$(call PATCH_MIPS_FLAGS,$(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle)
|
||||||
GOOS=darwin GOARCH=arm64 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-darwin-arm64 ./$(CMD_DIR)
|
GOOS=darwin GOARCH=arm64 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-darwin-arm64 ./$(CMD_DIR)
|
||||||
GOOS=windows GOARCH=amd64 $(GO) build -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-windows-amd64.exe ./$(CMD_DIR)
|
GOOS=windows GOARCH=amd64 $(GO) build -tags $(GO_BUILD_TAGS),whatsapp_native -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-windows-amd64.exe ./$(CMD_DIR)
|
||||||
## @$(GO) build $(GOFLAGS) -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BINARY_PATH) ./$(CMD_DIR)
|
## @$(GO) build $(GOFLAGS) -tags whatsapp_native -ldflags "$(LDFLAGS)" -o $(BINARY_PATH) ./$(CMD_DIR)
|
||||||
@echo "Build complete"
|
@echo "Build complete"
|
||||||
## @ln -sf $(BINARY_NAME)-$(PLATFORM)-$(ARCH) $(BUILD_DIR)/$(BINARY_NAME)
|
## @ln -sf $(BINARY_NAME)-$(PLATFORM)-$(ARCH) $(BUILD_DIR)/$(BINARY_NAME)
|
||||||
|
|
@ -147,21 +153,21 @@ build-whatsapp-native: generate
|
||||||
build-linux-arm: generate
|
build-linux-arm: generate
|
||||||
@echo "Building for linux/arm (GOARM=7)..."
|
@echo "Building for linux/arm (GOARM=7)..."
|
||||||
@mkdir -p $(BUILD_DIR)
|
@mkdir -p $(BUILD_DIR)
|
||||||
GOOS=linux GOARCH=arm GOARM=7 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm GOARM=7 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm ./$(CMD_DIR)
|
||||||
@echo "Build complete: $(BUILD_DIR)/$(BINARY_NAME)-linux-arm"
|
@echo "Build complete: $(BUILD_DIR)/$(BINARY_NAME)-linux-arm"
|
||||||
|
|
||||||
## build-linux-arm64: Build for Linux ARM64 (e.g. Raspberry Pi Zero 2 W 64-bit)
|
## build-linux-arm64: Build for Linux ARM64 (e.g. Raspberry Pi Zero 2 W 64-bit)
|
||||||
build-linux-arm64: generate
|
build-linux-arm64: generate
|
||||||
@echo "Building for linux/arm64..."
|
@echo "Building for linux/arm64..."
|
||||||
@mkdir -p $(BUILD_DIR)
|
@mkdir -p $(BUILD_DIR)
|
||||||
GOOS=linux GOARCH=arm64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64 ./$(CMD_DIR)
|
||||||
@echo "Build complete: $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64"
|
@echo "Build complete: $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64"
|
||||||
|
|
||||||
## build-linux-mipsle: Build for Linux MIPS32 LE
|
## build-linux-mipsle: Build for Linux MIPS32 LE
|
||||||
build-linux-mipsle: generate
|
build-linux-mipsle: generate
|
||||||
@echo "Building for linux/mipsle (softfloat)..."
|
@echo "Building for linux/mipsle (softfloat)..."
|
||||||
@mkdir -p $(BUILD_DIR)
|
@mkdir -p $(BUILD_DIR)
|
||||||
GOOS=linux GOARCH=mipsle GOMIPS=softfloat $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle ./$(CMD_DIR)
|
GOOS=linux GOARCH=mipsle GOMIPS=softfloat $(GO) build $(GOFLAGS_NO_GOOLM) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle ./$(CMD_DIR)
|
||||||
$(call PATCH_MIPS_FLAGS,$(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle)
|
$(call PATCH_MIPS_FLAGS,$(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle)
|
||||||
@echo "Build complete: $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle"
|
@echo "Build complete: $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle"
|
||||||
|
|
||||||
|
|
@ -173,18 +179,18 @@ build-pi-zero: build-linux-arm build-linux-arm64
|
||||||
build-all: generate
|
build-all: generate
|
||||||
@echo "Building for multiple platforms..."
|
@echo "Building for multiple platforms..."
|
||||||
@mkdir -p $(BUILD_DIR)
|
@mkdir -p $(BUILD_DIR)
|
||||||
GOOS=linux GOARCH=amd64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-amd64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=amd64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-amd64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=arm GOARM=7 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm GOARM=7 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=arm64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-arm64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=loong64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-loong64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=loong64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-loong64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=riscv64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-riscv64 ./$(CMD_DIR)
|
GOOS=linux GOARCH=riscv64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-riscv64 ./$(CMD_DIR)
|
||||||
GOOS=linux GOARCH=mipsle GOMIPS=softfloat $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle ./$(CMD_DIR)
|
GOOS=linux GOARCH=mipsle GOMIPS=softfloat $(GO) build $(GOFLAGS_NO_GOOLM) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle ./$(CMD_DIR)
|
||||||
$(call PATCH_MIPS_FLAGS,$(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle)
|
$(call PATCH_MIPS_FLAGS,$(BUILD_DIR)/$(BINARY_NAME)-linux-mipsle)
|
||||||
GOOS=linux GOARCH=arm GOARM=7 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-armv7 ./$(CMD_DIR)
|
GOOS=linux GOARCH=arm GOARM=7 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-linux-armv7 ./$(CMD_DIR)
|
||||||
GOOS=darwin GOARCH=arm64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-darwin-arm64 ./$(CMD_DIR)
|
GOOS=darwin GOARCH=arm64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-darwin-arm64 ./$(CMD_DIR)
|
||||||
GOOS=windows GOARCH=amd64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-windows-amd64.exe ./$(CMD_DIR)
|
GOOS=windows GOARCH=amd64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-windows-amd64.exe ./$(CMD_DIR)
|
||||||
GOOS=netbsd GOARCH=amd64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-netbsd-amd64 ./$(CMD_DIR)
|
GOOS=netbsd GOARCH=amd64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-netbsd-amd64 ./$(CMD_DIR)
|
||||||
GOOS=netbsd GOARCH=arm64 $(GO) build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-netbsd-arm64 ./$(CMD_DIR)
|
GOOS=netbsd GOARCH=arm64 $(GO) build $(GOFLAGS) -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY_NAME)-netbsd-arm64 ./$(CMD_DIR)
|
||||||
@echo "All builds complete"
|
@echo "All builds complete"
|
||||||
|
|
||||||
## install: Install picoclaw to system and copy builtin skills
|
## install: Install picoclaw to system and copy builtin skills
|
||||||
|
|
@ -221,13 +227,13 @@ clean:
|
||||||
|
|
||||||
## vet: Run go vet for static analysis
|
## vet: Run go vet for static analysis
|
||||||
vet: generate
|
vet: generate
|
||||||
@packages="$$(go list ./...)" && \
|
@packages="$$($(GO) list $(GOFLAGS) ./...)" && \
|
||||||
$(GO) vet $$(printf '%s\n' "$$packages" | grep -v '^github.com/sipeed/picoclaw/web/')
|
$(GO) vet $(GOFLAGS) $$(printf '%s\n' "$$packages" | grep -v '^github.com/sipeed/picoclaw/web/')
|
||||||
@cd web/backend && $(WEB_GO) vet ./...
|
@cd web/backend && $(WEB_GO) vet ./...
|
||||||
|
|
||||||
## test: Test Go code
|
## test: Test Go code
|
||||||
test: generate
|
test: generate
|
||||||
@$(GO) test $$(go list ./... | grep -v github.com/sipeed/picoclaw/web/)
|
@$(GO) test $(GOFLAGS) $$($(GO) list $(GOFLAGS) ./... | grep -v github.com/sipeed/picoclaw/web/)
|
||||||
@cd web && make test
|
@cd web && make test
|
||||||
|
|
||||||
## fmt: Format Go code
|
## fmt: Format Go code
|
||||||
|
|
@ -236,11 +242,11 @@ fmt:
|
||||||
|
|
||||||
## lint: Run linters
|
## lint: Run linters
|
||||||
lint:
|
lint:
|
||||||
@$(GOLANGCI_LINT) run
|
@$(GOLANGCI_LINT) run --build-tags $(GO_BUILD_TAGS)
|
||||||
|
|
||||||
## fix: Fix linting issues
|
## fix: Fix linting issues
|
||||||
fix:
|
fix:
|
||||||
@$(GOLANGCI_LINT) run --fix
|
@$(GOLANGCI_LINT) run --fix --build-tags $(GO_BUILD_TAGS)
|
||||||
|
|
||||||
## deps: Download dependencies
|
## deps: Download dependencies
|
||||||
deps:
|
deps:
|
||||||
|
|
|
||||||
|
|
@ -322,14 +322,17 @@ This creates `~/.picoclaw/config.json` and the workspace directory.
|
||||||
"model_list": [
|
"model_list": [
|
||||||
{
|
{
|
||||||
"model_name": "gpt-5.4",
|
"model_name": "gpt-5.4",
|
||||||
"model": "openai/gpt-5.4",
|
"model": "openai/gpt-5.4"
|
||||||
"api_key": "sk-your-api-key"
|
// api_key is now loaded from .security.yml
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
> See `config/config.example.json` in the repo for a complete configuration template with all available options.
|
> See `config/config.example.json` in the repo for a complete configuration template with all available options.
|
||||||
|
>
|
||||||
|
> Please note: config.example.json format is version 0, with sensitive codes in it, and will be auto migrated to version 1+, then, the config.json will only store insensitive data, the sensitive codes will be stored in .security.yml, if you need manually modify the codes, please see `docs/security_configuration.md` for more details.
|
||||||
|
|
||||||
|
|
||||||
**3. Chat**
|
**3. Chat**
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -17,6 +17,7 @@ func NewAuthCommand() *cobra.Command {
|
||||||
newStatusCommand(),
|
newStatusCommand(),
|
||||||
newModelsCommand(),
|
newModelsCommand(),
|
||||||
newWeixinCommand(),
|
newWeixinCommand(),
|
||||||
|
newWeComCommand(),
|
||||||
)
|
)
|
||||||
|
|
||||||
return cmd
|
return cmd
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,7 @@ func TestNewAuthCommand(t *testing.T) {
|
||||||
"status",
|
"status",
|
||||||
"models",
|
"models",
|
||||||
"weixin",
|
"weixin",
|
||||||
|
"wecom",
|
||||||
}
|
}
|
||||||
|
|
||||||
subcommands := cmd.Commands()
|
subcommands := cmd.Commands()
|
||||||
|
|
|
||||||
407
cmd/picoclaw/internal/auth/wecom.go
Normal file
407
cmd/picoclaw/internal/auth/wecom.go
Normal file
|
|
@ -0,0 +1,407 @@
|
||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"runtime"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/mdp/qrterminal/v3"
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
wecomQRSourceID = "picoclaw"
|
||||||
|
wecomQRGenerateEndpoint = "https://work.weixin.qq.com/ai/qc/generate"
|
||||||
|
wecomQRQueryEndpoint = "https://work.weixin.qq.com/ai/qc/query_result"
|
||||||
|
wecomQRPageEndpoint = "https://work.weixin.qq.com/ai/qc/gen"
|
||||||
|
wecomQRHTTPTimeout = 15 * time.Second
|
||||||
|
wecomQRPollInterval = 3 * time.Second
|
||||||
|
wecomQRPollTimeout = 5 * time.Minute
|
||||||
|
wecomDefaultWebSocketURL = "wss://openws.work.weixin.qq.com"
|
||||||
|
)
|
||||||
|
|
||||||
|
type wecomQRScanner func(context.Context, wecomQRFlowOptions) (wecomQRBotInfo, error)
|
||||||
|
|
||||||
|
type wecomQRFlowOptions struct {
|
||||||
|
HTTPClient *http.Client
|
||||||
|
GenerateURL string
|
||||||
|
QueryURL string
|
||||||
|
QRCodePageURL string
|
||||||
|
SourceID string
|
||||||
|
PollInterval time.Duration
|
||||||
|
PollTimeout time.Duration
|
||||||
|
Writer io.Writer
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomQRBotInfo struct {
|
||||||
|
BotID string
|
||||||
|
Secret string
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomQRSession struct {
|
||||||
|
SCode string
|
||||||
|
AuthURL string
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomQRGenerateResponse struct {
|
||||||
|
ErrCode int `json:"errcode,omitempty"`
|
||||||
|
ErrMsg string `json:"errmsg,omitempty"`
|
||||||
|
Data struct {
|
||||||
|
SCode string `json:"scode"`
|
||||||
|
AuthURL string `json:"auth_url"`
|
||||||
|
} `json:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomQRQueryResponse struct {
|
||||||
|
ErrCode int `json:"errcode,omitempty"`
|
||||||
|
ErrMsg string `json:"errmsg,omitempty"`
|
||||||
|
Data struct {
|
||||||
|
Status string `json:"status"`
|
||||||
|
BotInfo struct {
|
||||||
|
BotID string `json:"botid"`
|
||||||
|
Secret string `json:"secret"`
|
||||||
|
} `json:"bot_info"`
|
||||||
|
} `json:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func newWeComCommand() *cobra.Command {
|
||||||
|
var timeout time.Duration
|
||||||
|
|
||||||
|
cmd := &cobra.Command{
|
||||||
|
Use: "wecom",
|
||||||
|
Short: "Scan a WeCom QR code and configure channels.wecom",
|
||||||
|
Args: cobra.NoArgs,
|
||||||
|
RunE: func(_ *cobra.Command, _ []string) error {
|
||||||
|
return authWeComCmd(timeout)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd.Flags().DurationVar(&timeout, "timeout", wecomQRPollTimeout, "How long to wait for QR confirmation")
|
||||||
|
|
||||||
|
return cmd
|
||||||
|
}
|
||||||
|
|
||||||
|
func authWeComCmd(timeout time.Duration) error {
|
||||||
|
return authWeComCmdWithScanner(context.Background(), os.Stdout, timeout, scanWeComQRCodeInteractive)
|
||||||
|
}
|
||||||
|
|
||||||
|
func authWeComCmdWithScanner(
|
||||||
|
ctx context.Context,
|
||||||
|
writer io.Writer,
|
||||||
|
timeout time.Duration,
|
||||||
|
scanner wecomQRScanner,
|
||||||
|
) error {
|
||||||
|
if scanner == nil {
|
||||||
|
return fmt.Errorf("wecom QR scanner is nil")
|
||||||
|
}
|
||||||
|
if writer == nil {
|
||||||
|
writer = os.Stdout
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := internal.LoadConfig()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to load config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := defaultWeComQRFlowOptions(timeout)
|
||||||
|
opts.Writer = writer
|
||||||
|
|
||||||
|
botInfo, err := scanner(ctx, opts)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
applyWeComAuthResult(cfg, botInfo)
|
||||||
|
|
||||||
|
if saveErr := config.SaveConfig(internal.GetConfigPath(), cfg); saveErr != nil {
|
||||||
|
return fmt.Errorf("failed to save config: %w", saveErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintln(writer)
|
||||||
|
fmt.Fprintln(writer, "WeCom connected.")
|
||||||
|
fmt.Fprintf(writer, "Bot ID: %s\n", botInfo.BotID)
|
||||||
|
fmt.Fprintf(writer, "Config: %s\n", internal.GetConfigPath())
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultWeComQRFlowOptions(timeout time.Duration) wecomQRFlowOptions {
|
||||||
|
if timeout <= 0 {
|
||||||
|
timeout = wecomQRPollTimeout
|
||||||
|
}
|
||||||
|
|
||||||
|
return wecomQRFlowOptions{
|
||||||
|
HTTPClient: &http.Client{Timeout: wecomQRHTTPTimeout},
|
||||||
|
GenerateURL: wecomQRGenerateEndpoint,
|
||||||
|
QueryURL: wecomQRQueryEndpoint,
|
||||||
|
QRCodePageURL: wecomQRPageEndpoint,
|
||||||
|
SourceID: wecomQRSourceID,
|
||||||
|
PollInterval: wecomQRPollInterval,
|
||||||
|
PollTimeout: timeout,
|
||||||
|
Writer: os.Stdout,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func applyWeComAuthResult(cfg *config.Config, botInfo wecomQRBotInfo) {
|
||||||
|
cfg.Channels.WeCom.Enabled = true
|
||||||
|
cfg.Channels.WeCom.BotID = botInfo.BotID
|
||||||
|
cfg.Channels.WeCom.SetSecret(botInfo.Secret)
|
||||||
|
if strings.TrimSpace(cfg.Channels.WeCom.WebSocketURL) == "" {
|
||||||
|
cfg.Channels.WeCom.WebSocketURL = wecomDefaultWebSocketURL
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func scanWeComQRCodeInteractive(ctx context.Context, opts wecomQRFlowOptions) (wecomQRBotInfo, error) {
|
||||||
|
opts = normalizeWeComQRFlowOptions(opts)
|
||||||
|
|
||||||
|
fmt.Fprintln(opts.Writer, "Requesting WeCom QR code...")
|
||||||
|
|
||||||
|
session, err := fetchWeComQRCode(ctx, opts)
|
||||||
|
if err != nil {
|
||||||
|
return wecomQRBotInfo{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintln(opts.Writer)
|
||||||
|
fmt.Fprintln(opts.Writer, "=======================================================")
|
||||||
|
fmt.Fprintln(opts.Writer, "Please scan the following QR code with WeCom:")
|
||||||
|
fmt.Fprintln(opts.Writer, "=======================================================")
|
||||||
|
fmt.Fprintln(opts.Writer)
|
||||||
|
|
||||||
|
qrterminal.GenerateWithConfig(session.AuthURL, qrterminal.Config{
|
||||||
|
Level: qrterminal.L,
|
||||||
|
Writer: opts.Writer,
|
||||||
|
HalfBlocks: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
pageURL, err := buildWeComQRCodePageURL(opts.QRCodePageURL, opts.SourceID, session.SCode)
|
||||||
|
if err != nil {
|
||||||
|
return wecomQRBotInfo{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintln(opts.Writer)
|
||||||
|
fmt.Fprintf(opts.Writer, "QR Code Link: %s\n", pageURL)
|
||||||
|
fmt.Fprintln(opts.Writer)
|
||||||
|
fmt.Fprintln(opts.Writer, "Waiting for scan...")
|
||||||
|
|
||||||
|
return pollWeComQRCodeResult(ctx, opts, session.SCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeWeComQRFlowOptions(opts wecomQRFlowOptions) wecomQRFlowOptions {
|
||||||
|
if opts.HTTPClient == nil {
|
||||||
|
opts.HTTPClient = &http.Client{Timeout: wecomQRHTTPTimeout}
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(opts.GenerateURL) == "" {
|
||||||
|
opts.GenerateURL = wecomQRGenerateEndpoint
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(opts.QueryURL) == "" {
|
||||||
|
opts.QueryURL = wecomQRQueryEndpoint
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(opts.QRCodePageURL) == "" {
|
||||||
|
opts.QRCodePageURL = wecomQRPageEndpoint
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(opts.SourceID) == "" {
|
||||||
|
opts.SourceID = wecomQRSourceID
|
||||||
|
}
|
||||||
|
if opts.PollInterval <= 0 {
|
||||||
|
opts.PollInterval = wecomQRPollInterval
|
||||||
|
}
|
||||||
|
if opts.PollTimeout <= 0 {
|
||||||
|
opts.PollTimeout = wecomQRPollTimeout
|
||||||
|
}
|
||||||
|
if opts.Writer == nil {
|
||||||
|
opts.Writer = os.Stdout
|
||||||
|
}
|
||||||
|
|
||||||
|
return opts
|
||||||
|
}
|
||||||
|
|
||||||
|
func fetchWeComQRCode(ctx context.Context, opts wecomQRFlowOptions) (wecomQRSession, error) {
|
||||||
|
generateURL, err := buildWeComQRGenerateURL(opts.GenerateURL, opts.SourceID, wecomPlatformCode())
|
||||||
|
if err != nil {
|
||||||
|
return wecomQRSession{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp wecomQRGenerateResponse
|
||||||
|
if err := doWeComJSONGet(ctx, opts.HTTPClient, generateURL, &resp); err != nil {
|
||||||
|
return wecomQRSession{}, fmt.Errorf("failed to get WeCom QR code: %w", err)
|
||||||
|
}
|
||||||
|
if resp.ErrCode != 0 {
|
||||||
|
return wecomQRSession{}, fmt.Errorf(
|
||||||
|
"failed to get WeCom QR code: errcode=%d errmsg=%s",
|
||||||
|
resp.ErrCode,
|
||||||
|
resp.ErrMsg,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if resp.Data.SCode == "" || resp.Data.AuthURL == "" {
|
||||||
|
return wecomQRSession{}, fmt.Errorf("failed to get WeCom QR code: response missing scode or auth_url")
|
||||||
|
}
|
||||||
|
|
||||||
|
return wecomQRSession{
|
||||||
|
SCode: resp.Data.SCode,
|
||||||
|
AuthURL: resp.Data.AuthURL,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func pollWeComQRCodeResult(ctx context.Context, opts wecomQRFlowOptions, scode string) (wecomQRBotInfo, error) {
|
||||||
|
if strings.TrimSpace(scode) == "" {
|
||||||
|
return wecomQRBotInfo{}, fmt.Errorf("missing WeCom QR scode")
|
||||||
|
}
|
||||||
|
|
||||||
|
timeoutCtx, cancel := context.WithTimeout(ctx, opts.PollTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
var scannedPrinted bool
|
||||||
|
|
||||||
|
for {
|
||||||
|
status, err := queryWeComQRCodeStatus(timeoutCtx, opts, scode)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
|
||||||
|
return wecomQRBotInfo{}, fmt.Errorf("WeCom QR scan timed out after %s", opts.PollTimeout)
|
||||||
|
}
|
||||||
|
return wecomQRBotInfo{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch strings.ToLower(status.Data.Status) {
|
||||||
|
case "success":
|
||||||
|
if status.Data.BotInfo.BotID == "" || status.Data.BotInfo.Secret == "" {
|
||||||
|
return wecomQRBotInfo{}, fmt.Errorf("WeCom QR scan succeeded but bot credentials are missing")
|
||||||
|
}
|
||||||
|
return wecomQRBotInfo{
|
||||||
|
BotID: status.Data.BotInfo.BotID,
|
||||||
|
Secret: status.Data.BotInfo.Secret,
|
||||||
|
}, nil
|
||||||
|
case "expired":
|
||||||
|
return wecomQRBotInfo{}, fmt.Errorf("WeCom QR code expired, please retry")
|
||||||
|
case "scaned", "scanned":
|
||||||
|
if !scannedPrinted {
|
||||||
|
fmt.Fprintln(opts.Writer, "QR code scanned. Confirm the login in WeCom.")
|
||||||
|
scannedPrinted = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-timeoutCtx.Done():
|
||||||
|
if errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
|
||||||
|
return wecomQRBotInfo{}, fmt.Errorf("WeCom QR scan timed out after %s", opts.PollTimeout)
|
||||||
|
}
|
||||||
|
return wecomQRBotInfo{}, timeoutCtx.Err()
|
||||||
|
case <-time.After(opts.PollInterval):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func queryWeComQRCodeStatus(ctx context.Context, opts wecomQRFlowOptions, scode string) (wecomQRQueryResponse, error) {
|
||||||
|
queryURL, err := buildWeComQRQueryURL(opts.QueryURL, scode)
|
||||||
|
if err != nil {
|
||||||
|
return wecomQRQueryResponse{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp wecomQRQueryResponse
|
||||||
|
if err := doWeComJSONGet(ctx, opts.HTTPClient, queryURL, &resp); err != nil {
|
||||||
|
return wecomQRQueryResponse{}, fmt.Errorf("failed to query WeCom QR result: %w", err)
|
||||||
|
}
|
||||||
|
if resp.ErrCode != 0 {
|
||||||
|
return wecomQRQueryResponse{}, fmt.Errorf(
|
||||||
|
"failed to query WeCom QR result: errcode=%d errmsg=%s",
|
||||||
|
resp.ErrCode,
|
||||||
|
resp.ErrMsg,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildWeComQRGenerateURL(baseURL, sourceID string, platformCode int) (string, error) {
|
||||||
|
u, err := url.Parse(baseURL)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("invalid WeCom QR generate URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
query := u.Query()
|
||||||
|
query.Set("source", sourceID)
|
||||||
|
query.Set("sourceID", sourceID)
|
||||||
|
query.Set("plat", strconv.Itoa(platformCode))
|
||||||
|
u.RawQuery = query.Encode()
|
||||||
|
|
||||||
|
return u.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildWeComQRQueryURL(baseURL, scode string) (string, error) {
|
||||||
|
u, err := url.Parse(baseURL)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("invalid WeCom QR query URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
query := u.Query()
|
||||||
|
query.Set("scode", scode)
|
||||||
|
u.RawQuery = query.Encode()
|
||||||
|
|
||||||
|
return u.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildWeComQRCodePageURL(baseURL, sourceID, scode string) (string, error) {
|
||||||
|
u, err := url.Parse(baseURL)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("invalid WeCom QR page URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
query := u.Query()
|
||||||
|
query.Set("source", sourceID)
|
||||||
|
query.Set("sourceID", sourceID)
|
||||||
|
query.Set("scode", scode)
|
||||||
|
u.RawQuery = query.Encode()
|
||||||
|
|
||||||
|
return u.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func doWeComJSONGet(ctx context.Context, client *http.Client, targetURL string, out any) error {
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, targetURL, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
body, readErr := io.ReadAll(io.LimitReader(resp.Body, 8192))
|
||||||
|
if readErr != nil {
|
||||||
|
return fmt.Errorf("unexpected status %s", resp.Status)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("unexpected status %s: %s", resp.Status, strings.TrimSpace(string(body)))
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.NewDecoder(resp.Body).Decode(out); err != nil {
|
||||||
|
return fmt.Errorf("decode JSON response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func wecomPlatformCode() int {
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "darwin":
|
||||||
|
return 1
|
||||||
|
case "windows":
|
||||||
|
return 2
|
||||||
|
case "linux":
|
||||||
|
return 3
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
}
|
||||||
157
cmd/picoclaw/internal/auth/wecom_test.go
Normal file
157
cmd/picoclaw/internal/auth/wecom_test.go
Normal file
|
|
@ -0,0 +1,157 @@
|
||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"net/url"
|
||||||
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewWeComCommand(t *testing.T) {
|
||||||
|
cmd := newWeComCommand()
|
||||||
|
|
||||||
|
require.NotNil(t, cmd)
|
||||||
|
assert.Equal(t, "wecom", cmd.Use)
|
||||||
|
assert.Equal(t, "Scan a WeCom QR code and configure channels.wecom", cmd.Short)
|
||||||
|
assert.NotNil(t, cmd.Flags().Lookup("timeout"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildWeComQRGenerateURL(t *testing.T) {
|
||||||
|
rawURL, err := buildWeComQRGenerateURL("https://example.com/ai/qc/generate", wecomQRSourceID, 3)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
parsed, err := url.Parse(rawURL)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, wecomQRSourceID, parsed.Query().Get("source"))
|
||||||
|
assert.Equal(t, wecomQRSourceID, parsed.Query().Get("sourceID"))
|
||||||
|
assert.Equal(t, "3", parsed.Query().Get("plat"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildWeComQRCodePageURL(t *testing.T) {
|
||||||
|
rawURL, err := buildWeComQRCodePageURL("https://example.com/ai/qc/gen", wecomQRSourceID, "scode-1")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
parsed, err := url.Parse(rawURL)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, wecomQRSourceID, parsed.Query().Get("source"))
|
||||||
|
assert.Equal(t, wecomQRSourceID, parsed.Query().Get("sourceID"))
|
||||||
|
assert.Equal(t, "scode-1", parsed.Query().Get("scode"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFetchWeComQRCode(t *testing.T) {
|
||||||
|
server := httptest.NewServer(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"))
|
||||||
|
assert.Equal(t, strconv.Itoa(wecomPlatformCode()), r.URL.Query().Get("plat"))
|
||||||
|
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(),
|
||||||
|
GenerateURL: server.URL + "/generate",
|
||||||
|
Writer: bytes.NewBuffer(nil),
|
||||||
|
})
|
||||||
|
|
||||||
|
session, err := fetchWeComQRCode(context.Background(), opts)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "scode-1", session.SCode)
|
||||||
|
assert.Equal(t, "https://example.com/qr", session.AuthURL)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPollWeComQRCodeResult(t *testing.T) {
|
||||||
|
var calls atomic.Int32
|
||||||
|
|
||||||
|
server := httptest.NewServer(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"))
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
switch call {
|
||||||
|
case 1:
|
||||||
|
_, _ = w.Write([]byte(`{"data":{"status":"wait"}}`))
|
||||||
|
case 2:
|
||||||
|
_, _ = w.Write([]byte(`{"data":{"status":"scaned"}}`))
|
||||||
|
default:
|
||||||
|
_, _ = w.Write([]byte(`{"data":{"status":"success","bot_info":{"botid":"bot-1","secret":"secret-1"}}}`))
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
var output bytes.Buffer
|
||||||
|
opts := normalizeWeComQRFlowOptions(wecomQRFlowOptions{
|
||||||
|
HTTPClient: server.Client(),
|
||||||
|
QueryURL: server.URL + "/query",
|
||||||
|
PollInterval: time.Millisecond,
|
||||||
|
PollTimeout: time.Second,
|
||||||
|
Writer: &output,
|
||||||
|
})
|
||||||
|
|
||||||
|
botInfo, err := pollWeComQRCodeResult(context.Background(), opts, "scode-1")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "bot-1", botInfo.BotID)
|
||||||
|
assert.Equal(t, "secret-1", botInfo.Secret)
|
||||||
|
assert.Contains(t, output.String(), "QR code scanned. Confirm the login in WeCom.")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyWeComAuthResult(t *testing.T) {
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Channels.WeCom.WebSocketURL = ""
|
||||||
|
|
||||||
|
applyWeComAuthResult(cfg, wecomQRBotInfo{
|
||||||
|
BotID: "bot-1",
|
||||||
|
Secret: "secret-1",
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.True(t, cfg.Channels.WeCom.Enabled)
|
||||||
|
assert.Equal(t, "bot-1", cfg.Channels.WeCom.BotID)
|
||||||
|
assert.Equal(t, "secret-1", cfg.Channels.WeCom.Secret())
|
||||||
|
assert.Equal(t, wecomDefaultWebSocketURL, cfg.Channels.WeCom.WebSocketURL)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthWeComCmdWithScanner(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
configPath := filepath.Join(tmpDir, "config.json")
|
||||||
|
|
||||||
|
t.Setenv(config.EnvHome, tmpDir)
|
||||||
|
t.Setenv(config.EnvConfig, configPath)
|
||||||
|
|
||||||
|
var output bytes.Buffer
|
||||||
|
err := authWeComCmdWithScanner(
|
||||||
|
context.Background(),
|
||||||
|
&output,
|
||||||
|
time.Second,
|
||||||
|
func(_ context.Context, opts wecomQRFlowOptions) (wecomQRBotInfo, error) {
|
||||||
|
assert.Equal(t, wecomQRSourceID, opts.SourceID)
|
||||||
|
return wecomQRBotInfo{
|
||||||
|
BotID: "bot-1",
|
||||||
|
Secret: "secret-1",
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(internal.GetConfigPath())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, cfg.Channels.WeCom.Enabled)
|
||||||
|
assert.Equal(t, "bot-1", cfg.Channels.WeCom.BotID)
|
||||||
|
assert.Equal(t, "secret-1", cfg.Channels.WeCom.Secret())
|
||||||
|
assert.Equal(t, wecomDefaultWebSocketURL, cfg.Channels.WeCom.WebSocketURL)
|
||||||
|
assert.Contains(t, output.String(), "WeCom connected.")
|
||||||
|
}
|
||||||
|
|
@ -162,7 +162,9 @@
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"text": "Thinking... 💭"
|
"text": "Thinking... 💭"
|
||||||
},
|
},
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": "",
|
||||||
|
"crypto_database_path": "",
|
||||||
|
"crypto_passphrase": "YOUR_MATRIX_CRYPTO_PICKLE_KEY"
|
||||||
},
|
},
|
||||||
"line": {
|
"line": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
@ -182,39 +184,13 @@
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": ""
|
||||||
},
|
},
|
||||||
"wecom": {
|
"wecom": {
|
||||||
"_comment": "WeCom Bot - Easier setup, supports group chats",
|
"_comment": "WeCom AI Bot over WebSocket.",
|
||||||
"enabled": false,
|
|
||||||
"token": "YOUR_TOKEN",
|
|
||||||
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
|
|
||||||
"webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=YOUR_KEY",
|
|
||||||
"webhook_path": "/webhook/wecom",
|
|
||||||
"allow_from": [],
|
|
||||||
"reply_timeout": 5,
|
|
||||||
"reasoning_channel_id": ""
|
|
||||||
},
|
|
||||||
"wecom_app": {
|
|
||||||
"_comment": "WeCom App (自建应用) - More features, proactive messaging, private chat only.",
|
|
||||||
"enabled": false,
|
|
||||||
"corp_id": "YOUR_CORP_ID",
|
|
||||||
"corp_secret": "YOUR_CORP_SECRET",
|
|
||||||
"agent_id": 1000002,
|
|
||||||
"token": "YOUR_TOKEN",
|
|
||||||
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
|
|
||||||
"webhook_path": "/webhook/wecom-app",
|
|
||||||
"allow_from": [],
|
|
||||||
"reply_timeout": 5,
|
|
||||||
"reasoning_channel_id": ""
|
|
||||||
},
|
|
||||||
"wecom_aibot": {
|
|
||||||
"_comment": "WeCom AI Bot (智能机器人) - Official WeCom AI Bot integration, supports proactive messaging and private chats.",
|
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"bot_id": "YOUR_BOT_ID",
|
"bot_id": "YOUR_BOT_ID",
|
||||||
"secret": "YOUR_SECRET",
|
"secret": "YOUR_SECRET",
|
||||||
"token": "YOUR_TOKEN",
|
"websocket_url": "wss://openws.work.weixin.qq.com",
|
||||||
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
|
"send_thinking_message": true,
|
||||||
"webhook_path": "/webhook/wecom-aibot",
|
"allow_from": [],
|
||||||
"max_steps": 10,
|
|
||||||
"welcome_message": "Hello! I'm your AI assistant. How can I help you today?",
|
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": ""
|
||||||
},
|
},
|
||||||
"pico": {
|
"pico": {
|
||||||
|
|
@ -264,79 +240,6 @@
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": ""
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"providers": {
|
|
||||||
"_comment": "DEPRECATED: Use model_list instead. This will be removed in a future version",
|
|
||||||
"anthropic": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"openai": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": "",
|
|
||||||
"web_search": true
|
|
||||||
},
|
|
||||||
"openrouter": {
|
|
||||||
"api_key": "sk-or-v1-xxx",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"groq": {
|
|
||||||
"api_key": "gsk_xxx",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"zhipu": {
|
|
||||||
"api_key": "YOUR_ZHIPU_API_KEY",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"gemini": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"vllm": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"nvidia": {
|
|
||||||
"api_key": "nvapi-xxx",
|
|
||||||
"api_base": "",
|
|
||||||
"proxy": "http://127.0.0.1:7890"
|
|
||||||
},
|
|
||||||
"moonshot": {
|
|
||||||
"api_key": "sk-xxx",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"qwen": {
|
|
||||||
"api_key": "sk-xxx",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"ollama": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": "http://localhost:11434/v1"
|
|
||||||
},
|
|
||||||
"cerebras": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"volcengine": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": ""
|
|
||||||
},
|
|
||||||
"mistral": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": "https://api.mistral.ai/v1"
|
|
||||||
},
|
|
||||||
"avian": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": "https://api.avian.io/v1"
|
|
||||||
},
|
|
||||||
"longcat": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": "https://api.longcat.chat/openai"
|
|
||||||
},
|
|
||||||
"modelscope": {
|
|
||||||
"api_key": "",
|
|
||||||
"api_base": "https://api-inference.modelscope.cn/v1"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"tools": {
|
"tools": {
|
||||||
"allow_read_paths": null,
|
"allow_read_paths": null,
|
||||||
"allow_write_paths": null,
|
"allow_write_paths": null,
|
||||||
|
|
|
||||||
|
|
@ -25,7 +25,9 @@ Add this to `config.json`:
|
||||||
"text": "Thinking..."
|
"text": "Thinking..."
|
||||||
},
|
},
|
||||||
"reasoning_channel_id": "",
|
"reasoning_channel_id": "",
|
||||||
"message_format": "richtext"
|
"message_format": "richtext",
|
||||||
|
"crypto_database_path": "",
|
||||||
|
"crypto_passphrase": "YOUR_MATRIX_CRYPTO_PICKLE_KEY"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -46,6 +48,8 @@ Add this to `config.json`:
|
||||||
| placeholder | object | No | Placeholder message config |
|
| placeholder | object | No | Placeholder message config |
|
||||||
| reasoning_channel_id | string | No | Target channel for reasoning output |
|
| reasoning_channel_id | string | No | Target channel for reasoning output |
|
||||||
| message_format | string | No | Output format: `"richtext"` (default) renders markdown as HTML; `"plain"` sends plain text only |
|
| message_format | string | No | Output format: `"richtext"` (default) renders markdown as HTML; `"plain"` sends plain text only |
|
||||||
|
| crypto_database_path | string | No | Path to store the crypto database (uses workspace path `~/.picoclaw/workspace` if empty) |
|
||||||
|
| crypto_passphrase | string | No | Serialization key for encrypting session keys in the database; must remain unchanged once set |
|
||||||
|
|
||||||
## 3. Currently Supported
|
## 3. Currently Supported
|
||||||
|
|
||||||
|
|
@ -58,6 +62,7 @@ Add this to `config.json`:
|
||||||
- Typing state (`m.typing`)
|
- Typing state (`m.typing`)
|
||||||
- Placeholder message + final reply replacement
|
- Placeholder message + final reply replacement
|
||||||
- Auto-join invited rooms (can be disabled)
|
- Auto-join invited rooms (can be disabled)
|
||||||
|
- End-to-end encryption (E2EE) support for encrypted messages
|
||||||
|
|
||||||
## 4. TODO
|
## 4. TODO
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -24,7 +24,10 @@
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"text": "Thinking... 💭"
|
"text": "Thinking... 💭"
|
||||||
},
|
},
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": "",
|
||||||
|
"message_format": "richtext",
|
||||||
|
"crypto_database_path": "",
|
||||||
|
"crypto_passphrase": "YOUR_MATRIX_CRYPTO_PICKLE_KEY"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -45,6 +48,8 @@
|
||||||
| placeholder | object | 否 | 占位消息配置 |
|
| placeholder | object | 否 | 占位消息配置 |
|
||||||
| reasoning_channel_id | string | 否 | 思维链输出目标通道 |
|
| reasoning_channel_id | string | 否 | 思维链输出目标通道 |
|
||||||
| message_format | string | 否 | 消息格式:`richtext`(富文本)或 `plain`(纯文本) |
|
| message_format | string | 否 | 消息格式:`richtext`(富文本)或 `plain`(纯文本) |
|
||||||
|
| crypto_database_path | string | 否 | 加密数据库存储路径(为空时使用工作空间路径 `~/.picoclaw/workspace`) |
|
||||||
|
| crypto_passphrase | string | 否 | 加密数据库中 session key 的序列化密钥;设置后不能更改 |
|
||||||
|
|
||||||
## 3. 当前支持
|
## 3. 当前支持
|
||||||
|
|
||||||
|
|
@ -56,6 +61,7 @@
|
||||||
- Typing 状态(`m.typing`)
|
- Typing 状态(`m.typing`)
|
||||||
- 占位消息(`Thinking... 💭`)+ 最终回复替换
|
- 占位消息(`Thinking... 💭`)+ 最终回复替换
|
||||||
- 自动加入邀请房间(可关闭)
|
- 自动加入邀请房间(可关闭)
|
||||||
|
- 端对端加密(E2EE)消息支持
|
||||||
|
|
||||||
## 4. TODO
|
## 4. TODO
|
||||||
|
|
||||||
|
|
|
||||||
104
docs/channels/wecom/README.md
Normal file
104
docs/channels/wecom/README.md
Normal file
|
|
@ -0,0 +1,104 @@
|
||||||
|
> Back to [README](../../../README.md)
|
||||||
|
|
||||||
|
# WeCom
|
||||||
|
|
||||||
|
PicoClaw now exposes WeCom as a single `channels.wecom` channel built on the official WeCom AI Bot WebSocket API.
|
||||||
|
This replaces the legacy `wecom`, `wecom_app`, and `wecom_aibot` split with one configuration model.
|
||||||
|
|
||||||
|
## What This Channel Supports
|
||||||
|
|
||||||
|
- Direct chat and group chat delivery
|
||||||
|
- Channel-side streaming replies over WeCom's AI Bot protocol
|
||||||
|
- Incoming text, voice, image, file, video, and mixed messages
|
||||||
|
- Outbound text and media replies (`image`, `file`, `voice`, `video`)
|
||||||
|
- QR-based CLI onboarding with `picoclaw auth wecom`
|
||||||
|
- Shared allowlist and `reasoning_channel_id` routing
|
||||||
|
|
||||||
|
> No public webhook callback URL is required for this channel. PicoClaw opens an outbound WebSocket connection to WeCom.
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
### Option 1: QR Login From CLI
|
||||||
|
|
||||||
|
Run:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth wecom
|
||||||
|
```
|
||||||
|
|
||||||
|
The command prints a QR code in the terminal, waits for confirmation in WeCom, and then writes the resulting
|
||||||
|
`bot_id` and `secret` into `channels.wecom`.
|
||||||
|
|
||||||
|
Use `--timeout` if you want to wait longer:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth wecom --timeout 10m
|
||||||
|
```
|
||||||
|
|
||||||
|
### Option 2: Configure Manually
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"wecom": {
|
||||||
|
"enabled": true,
|
||||||
|
"bot_id": "YOUR_BOT_ID",
|
||||||
|
"secret": "YOUR_SECRET",
|
||||||
|
"websocket_url": "wss://openws.work.weixin.qq.com",
|
||||||
|
"send_thinking_message": true,
|
||||||
|
"allow_from": [],
|
||||||
|
"reasoning_channel_id": ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
| Field | Type | Required | Description |
|
||||||
|
| ----- | ---- | -------- | ----------- |
|
||||||
|
| `enabled` | bool | No | Enables the WeCom channel. |
|
||||||
|
| `bot_id` | string | Yes | WeCom AI Bot identifier. Required when the channel is enabled. |
|
||||||
|
| `secret` | string | Yes | WeCom AI Bot secret. Required when the channel is enabled. |
|
||||||
|
| `websocket_url` | string | No | WebSocket endpoint. Defaults to `wss://openws.work.weixin.qq.com`. |
|
||||||
|
| `send_thinking_message` | bool | No | Sends an initial `Processing...` chunk before the final streamed reply. Defaults to `true`. |
|
||||||
|
| `allow_from` | array | No | Sender allowlist. Empty means allow all senders. |
|
||||||
|
| `reasoning_channel_id` | string | No | Optional destination for reasoning/thinking output. |
|
||||||
|
|
||||||
|
## Runtime Behavior
|
||||||
|
|
||||||
|
- PicoClaw keeps the active WeCom turn so normal replies can continue the same stream when possible.
|
||||||
|
- If streaming is no longer available, replies fall back to active push delivery to the resolved chat route.
|
||||||
|
- Incoming media is downloaded into the media store before being handed to the agent.
|
||||||
|
- Outbound media is uploaded to WeCom in temporary chunks and then sent as a regular media message.
|
||||||
|
|
||||||
|
## Migration Notes
|
||||||
|
|
||||||
|
This branch removes the old multi-channel WeCom model.
|
||||||
|
|
||||||
|
| Previous config | Now |
|
||||||
|
| --------------- | --- |
|
||||||
|
| `channels.wecom` webhook bot | Replace with `channels.wecom` using `bot_id` + `secret`. |
|
||||||
|
| `channels.wecom_app` | Remove it and use `channels.wecom`. |
|
||||||
|
| `channels.wecom_aibot` | Move the config to `channels.wecom`. |
|
||||||
|
| `token`, `encoding_aes_key`, `webhook_url`, `webhook_path` | No longer used by the WeCom channel. |
|
||||||
|
| `corp_id`, `corp_secret`, `agent_id` | No longer used by the WeCom channel. |
|
||||||
|
| `welcome_message`, `processing_message`, `max_steps` under WeCom | No longer part of the WeCom channel config. |
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### `picoclaw auth wecom` times out
|
||||||
|
|
||||||
|
- Re-run with a larger `--timeout`.
|
||||||
|
- Make sure the QR code was confirmed inside WeCom, not only scanned.
|
||||||
|
|
||||||
|
### WebSocket connection fails
|
||||||
|
|
||||||
|
- Verify `bot_id` and `secret`.
|
||||||
|
- Confirm the host can reach `wss://openws.work.weixin.qq.com`.
|
||||||
|
|
||||||
|
### Replies do not arrive
|
||||||
|
|
||||||
|
- Check whether `allow_from` blocks the sender.
|
||||||
|
- Check launcher or startup validation for missing `channels.wecom.bot_id` / `channels.wecom.secret`.
|
||||||
|
|
||||||
104
docs/channels/wecom/README.zh.md
Normal file
104
docs/channels/wecom/README.zh.md
Normal file
|
|
@ -0,0 +1,104 @@
|
||||||
|
> 返回 [README](../../../README.zh.md)
|
||||||
|
|
||||||
|
# 企业微信
|
||||||
|
|
||||||
|
PicoClaw 现在将企业微信统一为一个 `channels.wecom` 渠道,并基于企业微信官方 AI Bot WebSocket 协议实现。
|
||||||
|
这取代了旧的 `wecom`、`wecom_app`、`wecom_aibot` 三套配置模型。
|
||||||
|
|
||||||
|
## 当前渠道能力
|
||||||
|
|
||||||
|
- 支持私聊和群聊
|
||||||
|
- 支持企业微信侧流式回复
|
||||||
|
- 支持接收文本、语音、图片、文件、视频和 mixed 消息
|
||||||
|
- 支持发送文本与媒体消息(`image`、`file`、`voice`、`video`)
|
||||||
|
- 支持通过 `picoclaw auth wecom` 扫码写入配置
|
||||||
|
- 支持统一白名单与 `reasoning_channel_id`
|
||||||
|
|
||||||
|
> 这个渠道不再需要公网 webhook 回调地址。PicoClaw 会主动向企业微信发起 WebSocket 连接。
|
||||||
|
|
||||||
|
## 快速开始
|
||||||
|
|
||||||
|
### 方式 1:命令行扫码登录
|
||||||
|
|
||||||
|
运行:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth wecom
|
||||||
|
```
|
||||||
|
|
||||||
|
该命令会在终端打印二维码,等待你在企业微信中确认,然后把生成的 `bot_id` 和 `secret` 写入
|
||||||
|
`channels.wecom`。
|
||||||
|
|
||||||
|
如果需要更长等待时间,可以加 `--timeout`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw auth wecom --timeout 10m
|
||||||
|
```
|
||||||
|
|
||||||
|
### 方式 2:手动配置
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"wecom": {
|
||||||
|
"enabled": true,
|
||||||
|
"bot_id": "YOUR_BOT_ID",
|
||||||
|
"secret": "YOUR_SECRET",
|
||||||
|
"websocket_url": "wss://openws.work.weixin.qq.com",
|
||||||
|
"send_thinking_message": true,
|
||||||
|
"allow_from": [],
|
||||||
|
"reasoning_channel_id": ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 配置字段
|
||||||
|
|
||||||
|
| 字段 | 类型 | 必填 | 说明 |
|
||||||
|
| ---- | ---- | ---- | ---- |
|
||||||
|
| `enabled` | bool | 否 | 是否启用企业微信渠道。 |
|
||||||
|
| `bot_id` | string | 是 | 企业微信 AI Bot 标识。渠道启用时必填。 |
|
||||||
|
| `secret` | string | 是 | 企业微信 AI Bot 密钥。渠道启用时必填。 |
|
||||||
|
| `websocket_url` | string | 否 | WebSocket 地址,默认 `wss://openws.work.weixin.qq.com`。 |
|
||||||
|
| `send_thinking_message` | bool | 否 | 是否在流式最终回复前先发送一段 `Processing...` 开场消息,默认 `true`。 |
|
||||||
|
| `allow_from` | array | 否 | 发送者白名单;空数组表示允许所有发送者。 |
|
||||||
|
| `reasoning_channel_id` | string | 否 | 可选的 reasoning/thinking 输出目标。 |
|
||||||
|
|
||||||
|
## 运行时行为
|
||||||
|
|
||||||
|
- PicoClaw 会保留当前会话对应的企业微信 turn,优先继续同一个流式回复。
|
||||||
|
- 如果流式上下文已经失效,回复会自动回退到主动推送消息。
|
||||||
|
- 收到的媒体会先下载到 media store,再交给 Agent 处理。
|
||||||
|
- 发出的媒体会先按分片上传到企业微信,再作为普通媒体消息发送。
|
||||||
|
|
||||||
|
## 迁移说明
|
||||||
|
|
||||||
|
这个分支移除了旧的多通道企业微信模型。
|
||||||
|
|
||||||
|
| 旧配置 | 现在怎么做 |
|
||||||
|
| ------ | ---------- |
|
||||||
|
| `channels.wecom` webhook 机器人 | 改为使用 `bot_id` + `secret` 的 `channels.wecom`。 |
|
||||||
|
| `channels.wecom_app` | 删除,统一迁移到 `channels.wecom`。 |
|
||||||
|
| `channels.wecom_aibot` | 配置迁移到 `channels.wecom`。 |
|
||||||
|
| `token`、`encoding_aes_key`、`webhook_url`、`webhook_path` | 企业微信渠道不再使用这些字段。 |
|
||||||
|
| `corp_id`、`corp_secret`、`agent_id` | 企业微信渠道不再使用这些字段。 |
|
||||||
|
| 企业微信下的 `welcome_message`、`processing_message`、`max_steps` | 不再属于企业微信渠道配置。 |
|
||||||
|
|
||||||
|
## 常见问题
|
||||||
|
|
||||||
|
### `picoclaw auth wecom` 超时
|
||||||
|
|
||||||
|
- 用更大的 `--timeout` 重新执行。
|
||||||
|
- 确认是在企业微信里完成了确认,而不只是扫描二维码。
|
||||||
|
|
||||||
|
### WebSocket 连接失败
|
||||||
|
|
||||||
|
- 检查 `bot_id` 和 `secret` 是否正确。
|
||||||
|
- 确认运行环境可以访问 `wss://openws.work.weixin.qq.com`。
|
||||||
|
|
||||||
|
### 消息没有回到企业微信
|
||||||
|
|
||||||
|
- 检查 `allow_from` 是否拦截了发送者。
|
||||||
|
- 检查启动日志或 launcher 校验,确认 `channels.wecom.bot_id` / `channels.wecom.secret` 已填写。
|
||||||
|
|
||||||
|
|
@ -6,7 +6,7 @@
|
||||||
|
|
||||||
Talk to your picoclaw through Telegram, Discord, WhatsApp, Matrix, QQ, DingTalk, LINE, WeCom, Feishu, Slack, IRC, OneBot, MaixCam, or Pico (native protocol)
|
Talk to your picoclaw through Telegram, Discord, WhatsApp, Matrix, QQ, DingTalk, LINE, WeCom, Feishu, Slack, IRC, OneBot, MaixCam, or Pico (native protocol)
|
||||||
|
|
||||||
> **Note**: All webhook-based channels (LINE, WeCom, etc.) are served on a single shared Gateway HTTP server (`gateway.host`:`gateway.port`, default `127.0.0.1:18790`). There are no per-channel ports to configure. Note: Feishu uses WebSocket/SDK mode and does not use the shared HTTP webhook server.
|
> **Note**: Channels that rely on HTTP callbacks share a single Gateway HTTP server (`gateway.host`:`gateway.port`, default `127.0.0.1:18790`). Socket/stream-based channels such as Feishu, DingTalk, and WeCom do not rely on the shared webhook server for inbound delivery.
|
||||||
|
|
||||||
| Channel | Difficulty | Description | Documentation |
|
| Channel | Difficulty | Description | Documentation |
|
||||||
| -------------------- | ------------------ | ----------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------- |
|
| -------------------- | ------------------ | ----------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------- |
|
||||||
|
|
@ -19,7 +19,7 @@ Talk to your picoclaw through Telegram, Discord, WhatsApp, Matrix, QQ, DingTalk,
|
||||||
| **QQ** | ⭐⭐ Medium | Official bot API, Chinese community | [Docs](channels/qq/README.md) |
|
| **QQ** | ⭐⭐ Medium | Official bot API, Chinese community | [Docs](channels/qq/README.md) |
|
||||||
| **DingTalk** | ⭐⭐ Medium | Stream mode (no public IP needed), enterprise | [Docs](channels/dingtalk/README.md) |
|
| **DingTalk** | ⭐⭐ Medium | Stream mode (no public IP needed), enterprise | [Docs](channels/dingtalk/README.md) |
|
||||||
| **LINE** | ⭐⭐⭐ Advanced | HTTPS Webhook required | [Docs](channels/line/README.md) |
|
| **LINE** | ⭐⭐⭐ Advanced | HTTPS Webhook required | [Docs](channels/line/README.md) |
|
||||||
| **WeCom (企业微信)** | ⭐⭐⭐ Advanced | Group Bot (Webhook), custom App (API), AI Bot | [Bot](channels/wecom/wecom_bot/README.md) / [App](channels/wecom/wecom_app/README.md) / [AI Bot](channels/wecom/wecom_aibot/README.md) |
|
| **WeCom (企业微信)** | ⭐⭐⭐ Advanced | Official AI Bot over WebSocket, streaming + media | [Docs](channels/wecom/README.md) |
|
||||||
| **Feishu (飞书)** | ⭐⭐⭐ Advanced | Enterprise collaboration, feature-rich | [Docs](channels/feishu/README.md) |
|
| **Feishu (飞书)** | ⭐⭐⭐ Advanced | Enterprise collaboration, feature-rich | [Docs](channels/feishu/README.md) |
|
||||||
| **IRC** | ⭐⭐ Medium | Server + TLS configuration | [Docs](#irc) |
|
| **IRC** | ⭐⭐ Medium | Server + TLS configuration | [Docs](#irc) |
|
||||||
| **OneBot** | ⭐⭐ Medium | NapCat/Go-CQHTTP compatible, community ecosystem | [Docs](channels/onebot/README.md) |
|
| **OneBot** | ⭐⭐ Medium | NapCat/Go-CQHTTP compatible, community ecosystem | [Docs](channels/onebot/README.md) |
|
||||||
|
|
@ -380,102 +380,34 @@ picoclaw gateway
|
||||||
<details>
|
<details>
|
||||||
<summary><b>WeCom (企业微信)</b></summary>
|
<summary><b>WeCom (企业微信)</b></summary>
|
||||||
|
|
||||||
PicoClaw supports three types of WeCom integration:
|
PicoClaw now exposes WeCom as a single AI Bot channel over WebSocket.
|
||||||
|
No public webhook callback URL is required.
|
||||||
|
|
||||||
**Option 1: WeCom Bot (Bot)** - Easier setup, supports group chats
|
See [WeCom Configuration Guide](channels/wecom/README.md) for the full configuration reference and migration notes.
|
||||||
**Option 2: WeCom App (Custom App)** - More features, proactive messaging, private chat only
|
|
||||||
**Option 3: WeCom AI Bot (AI Bot)** - Official AI Bot, streaming replies, supports group & private chat
|
|
||||||
|
|
||||||
See [WeCom AI Bot Configuration Guide](channels/wecom/wecom_aibot/README.md) for detailed setup instructions.
|
**Quick Setup - Recommended**
|
||||||
|
|
||||||
**Quick Setup - WeCom Bot:**
|
**1. Authenticate**
|
||||||
|
|
||||||
**1. Create a bot**
|
```bash
|
||||||
|
picoclaw auth wecom
|
||||||
|
```
|
||||||
|
|
||||||
* Go to WeCom Admin Console → Group Chat → Add Group Bot
|
This command shows a QR code, waits for approval in WeCom, and writes `bot_id` + `secret` into `channels.wecom`.
|
||||||
* Copy the webhook URL (format: `https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=xxx`)
|
|
||||||
|
|
||||||
**2. Configure**
|
**2. Configure manually if needed**
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"channels": {
|
"channels": {
|
||||||
"wecom": {
|
"wecom": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_TOKEN",
|
"bot_id": "YOUR_BOT_ID",
|
||||||
"encoding_aes_key": "YOUR_ENCODING_AES_KEY",
|
"secret": "YOUR_SECRET",
|
||||||
"webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=YOUR_KEY",
|
"websocket_url": "wss://openws.work.weixin.qq.com",
|
||||||
"webhook_path": "/webhook/wecom",
|
"send_thinking_message": true,
|
||||||
"allow_from": []
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
> WeCom webhook is served on the shared Gateway server (`gateway.host`:`gateway.port`, default `127.0.0.1:18790`).
|
|
||||||
|
|
||||||
**Quick Setup - WeCom App:**
|
|
||||||
|
|
||||||
**1. Create an app**
|
|
||||||
|
|
||||||
* Go to WeCom Admin Console → App Management → Create App
|
|
||||||
* Copy **AgentId** and **Secret**
|
|
||||||
* Go to "My Company" page, copy **CorpID**
|
|
||||||
|
|
||||||
**2. Configure receive message**
|
|
||||||
|
|
||||||
* In App details, click "Receive Message" → "Set API"
|
|
||||||
* Set URL to `http://your-server:18790/webhook/wecom-app`
|
|
||||||
* Generate **Token** and **EncodingAESKey**
|
|
||||||
|
|
||||||
**3. Configure**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"channels": {
|
|
||||||
"wecom_app": {
|
|
||||||
"enabled": true,
|
|
||||||
"corp_id": "wwxxxxxxxxxxxxxxxx",
|
|
||||||
"corp_secret": "YOUR_CORP_SECRET",
|
|
||||||
"agent_id": 1000002,
|
|
||||||
"token": "YOUR_TOKEN",
|
|
||||||
"encoding_aes_key": "YOUR_ENCODING_AES_KEY",
|
|
||||||
"webhook_path": "/webhook/wecom-app",
|
|
||||||
"allow_from": []
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**4. Run**
|
|
||||||
|
|
||||||
```bash
|
|
||||||
picoclaw gateway
|
|
||||||
```
|
|
||||||
|
|
||||||
> **Note**: WeCom webhook callbacks are served on the Gateway port (default 18790). Use a reverse proxy for HTTPS.
|
|
||||||
|
|
||||||
**Quick Setup - WeCom AI Bot:**
|
|
||||||
|
|
||||||
**1. Create an AI Bot**
|
|
||||||
|
|
||||||
* Go to WeCom Admin Console → App Management → AI Bot
|
|
||||||
* In the AI Bot settings, configure callback URL: `http://your-server:18790/webhook/wecom-aibot`
|
|
||||||
* Copy **Token** and click "Random Generate" for **EncodingAESKey**
|
|
||||||
|
|
||||||
**2. Configure**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"channels": {
|
|
||||||
"wecom_aibot": {
|
|
||||||
"enabled": true,
|
|
||||||
"token": "YOUR_TOKEN",
|
|
||||||
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
|
|
||||||
"webhook_path": "/webhook/wecom-aibot",
|
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"welcome_message": "Hello! How can I help you?",
|
"reasoning_channel_id": ""
|
||||||
"processing_message": "⏳ Processing, please wait. The results will be sent shortly."
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -487,7 +419,7 @@ picoclaw gateway
|
||||||
picoclaw gateway
|
picoclaw gateway
|
||||||
```
|
```
|
||||||
|
|
||||||
> **Note**: WeCom AI Bot uses streaming pull protocol — no reply timeout concerns. Long tasks (>30 seconds) automatically switch to `response_url` push delivery.
|
> Legacy `wecom_app` and `wecom_aibot` entries are replaced by the unified `channels.wecom` config in this branch.
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -454,6 +454,70 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
- **Load balancing**: Distribute requests across multiple endpoints
|
- **Load balancing**: Distribute requests across multiple endpoints
|
||||||
- **Centralized configuration**: Manage all providers in one place
|
- **Centralized configuration**: Manage all providers in one place
|
||||||
|
|
||||||
|
#### 🔒 Security Configuration (Recommended)
|
||||||
|
|
||||||
|
PicoClaw supports separating sensitive data (API keys, tokens, secrets) from your main configuration by storing them in a `.security.yml` file.
|
||||||
|
|
||||||
|
**Key Benefits:**
|
||||||
|
- **Security**: Sensitive data is never in your main config file
|
||||||
|
- **Easy sharing**: Share config.json without exposing API keys
|
||||||
|
- **Version control**: Add `.security.yml` to `.gitignore`
|
||||||
|
- **Flexible deployment**: Different environments can use different security files
|
||||||
|
|
||||||
|
**Quick Setup:**
|
||||||
|
|
||||||
|
1. Create `~/.picoclaw/.security.yml` with your API keys:
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "sk-proj-your-actual-openai-key"
|
||||||
|
claude-sonnet-4.6:
|
||||||
|
api_keys:
|
||||||
|
- "sk-ant-your-actual-anthropic-key"
|
||||||
|
channels:
|
||||||
|
telegram:
|
||||||
|
token: "your-telegram-bot-token"
|
||||||
|
web:
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "BSAyour-brave-api-key"
|
||||||
|
glm_search:
|
||||||
|
api_key: "your-glm-search-api-key"
|
||||||
|
```
|
||||||
|
|
||||||
|
2. Set proper permissions:
|
||||||
|
```bash
|
||||||
|
chmod 600 ~/.picoclaw/.security.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
3. Remove sensitive fields from `config.json` (recommended):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4",
|
||||||
|
"model": "openai/gpt-5.4"
|
||||||
|
// api_key loaded from .security.yml
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"channels": {
|
||||||
|
"telegram": {
|
||||||
|
"enabled": true"
|
||||||
|
// token loaded from .security.yml
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**How it works:**
|
||||||
|
- Values from `.security.yml` are automatically mapped to config fields
|
||||||
|
- No special syntax needed — just omit sensitive fields from config.json
|
||||||
|
- If a field exists in both files, `.security.yml` value takes precedence
|
||||||
|
- You can mix direct values in config.json with security values
|
||||||
|
|
||||||
|
For complete documentation, see [`security_configuration.md`](security_configuration.md).
|
||||||
|
|
||||||
#### All Supported Vendors
|
#### All Supported Vendors
|
||||||
|
|
||||||
| Vendor | `model` Prefix | Default API Base | Protocol | API Key |
|
| Vendor | `model` Prefix | Default API Base | Protocol | API Key |
|
||||||
|
|
@ -515,16 +579,20 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **Security Note**: You can remove `api_key` fields from your config and store them in `.security.yml` instead. See [Security Configuration](#-security-configuration-recommended) above for details.
|
||||||
|
|
||||||
#### Vendor-Specific Examples
|
#### 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).
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>OpenAI</b></summary>
|
<summary><b>OpenAI</b></summary>
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"model_name": "gpt-5.4",
|
"model_name": "gpt-5.4",
|
||||||
"model": "openai/gpt-5.4",
|
"model": "openai/gpt-5.4"
|
||||||
"api_key": "sk-..."
|
// api_key: set in .security.yml
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -536,8 +604,8 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"model_name": "ark-code-latest",
|
"model_name": "ark-code-latest",
|
||||||
"model": "volcengine/ark-code-latest",
|
"model": "volcengine/ark-code-latest"
|
||||||
"api_key": "sk-..."
|
// api_key: set in .security.yml
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -549,8 +617,8 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"model_name": "glm-4.7",
|
"model_name": "glm-4.7",
|
||||||
"model": "zhipu/glm-4.7",
|
"model": "zhipu/glm-4.7"
|
||||||
"api_key": "your-key"
|
// api_key: set in .security.yml
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -562,8 +630,8 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"model_name": "deepseek-chat",
|
"model_name": "deepseek-chat",
|
||||||
"model": "deepseek/deepseek-chat",
|
"model": "deepseek/deepseek-chat"
|
||||||
"api_key": "sk-..."
|
// api_key: set in .security.yml
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -575,8 +643,8 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"model_name": "claude-sonnet-4.6",
|
"model_name": "claude-sonnet-4.6",
|
||||||
"model": "anthropic/claude-sonnet-4.6",
|
"model": "anthropic/claude-sonnet-4.6"
|
||||||
"api_key": "sk-ant-your-key"
|
// api_key: set in .security.yml
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -616,8 +684,8 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
||||||
{
|
{
|
||||||
"model_name": "my-custom-model",
|
"model_name": "my-custom-model",
|
||||||
"model": "openai/custom-model",
|
"model": "openai/custom-model",
|
||||||
"api_base": "https://my-proxy.com/v1",
|
"api_base": "https://my-proxy.com/v1"
|
||||||
"api_key": "sk-..."
|
// api_key: set in .security.yml
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -629,6 +697,33 @@ PicoClaw strips only the outer `litellm/` prefix before sending the request, so
|
||||||
|
|
||||||
Configure multiple endpoints for the same model name — PicoClaw will automatically round-robin between them:
|
Configure multiple endpoints for the same model name — PicoClaw will automatically round-robin between them:
|
||||||
|
|
||||||
|
**Option 1: Multiple API Keys in .security.yml (Recommended)**
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# .security.yml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "sk-proj-key-1"
|
||||||
|
- "sk-proj-key-2"
|
||||||
|
```
|
||||||
|
|
||||||
|
```json
|
||||||
|
// config.json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4",
|
||||||
|
"model": "openai/gpt-5.4",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
// api_keys loaded from .security.yml
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Option 2: Multiple Model Entries**
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"model_list": [
|
"model_list": [
|
||||||
|
|
@ -685,6 +780,8 @@ This keeps the runtime lightweight while making new OpenAI-compatible backends m
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **Note**: The `providers` format is deprecated. Use the new `model_list` format with `.security.yml` for better security.
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
|
|
@ -701,18 +798,10 @@ This keeps the runtime lightweight while making new OpenAI-compatible backends m
|
||||||
"dm_scope": "per-channel-peer",
|
"dm_scope": "per-channel-peer",
|
||||||
"backlog_limit": 20
|
"backlog_limit": 20
|
||||||
},
|
},
|
||||||
"providers": {
|
|
||||||
"openrouter": {
|
|
||||||
"api_key": "sk-or-v1-xxx"
|
|
||||||
},
|
|
||||||
"groq": {
|
|
||||||
"api_key": "gsk_xxx"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"channels": {
|
"channels": {
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": true,
|
"enabled": true"
|
||||||
"token": "123456:ABC...",
|
// token: set in .security.yml
|
||||||
"allow_from": ["123456789"]
|
"allow_from": ["123456789"]
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
@ -731,6 +820,8 @@ This keeps the runtime lightweight while making new OpenAI-compatible backends m
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> **Note**: Sensitive fields (`api_key`, `token`, etc.) can be omitted and stored in `.security.yml` for better security.
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
### Scheduled Tasks / Reminders
|
### Scheduled Tasks / Reminders
|
||||||
|
|
@ -754,6 +845,7 @@ Scheduled tasks persist across restarts and are stored in `~/.picoclaw/workspace
|
||||||
|
|
||||||
| Topic | Description |
|
| Topic | Description |
|
||||||
| ----- | ----------- |
|
| ----- | ----------- |
|
||||||
|
| [Security Configuration](security_configuration.md) | Store API keys and secrets in separate `.security.yml` file |
|
||||||
| [Sensitive Data Filtering](sensitive_data_filtering.md) | Filter API keys and tokens from tool results before sending to LLM |
|
| [Sensitive Data Filtering](sensitive_data_filtering.md) | Filter API keys and tokens from tool results before sending to LLM |
|
||||||
| [Hook System](hooks/README.md) | Event-driven hooks: observers, interceptors, approval hooks |
|
| [Hook System](hooks/README.md) | Event-driven hooks: observers, interceptors, approval hooks |
|
||||||
| [Steering](steering.md) | Inject messages into a running agent loop between tool calls |
|
| [Steering](steering.md) | Inject messages into a running agent loop between tool calls |
|
||||||
|
|
|
||||||
|
|
@ -31,7 +31,7 @@ enc://AAAA...base64...
|
||||||
{
|
{
|
||||||
"model_name": "gpt-4o",
|
"model_name": "gpt-4o",
|
||||||
"model": "openai/gpt-4o",
|
"model": "openai/gpt-4o",
|
||||||
"api_key": "enc://AAAA...base64...",
|
// "api_key": "enc://AAAA...base64..." move to .security.yml
|
||||||
"api_base": "https://api.openai.com/v1"
|
"api_base": "https://api.openai.com/v1"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
|
|
|
||||||
644
docs/security_configuration.md
Normal file
644
docs/security_configuration.md
Normal file
|
|
@ -0,0 +1,644 @@
|
||||||
|
# Security Configuration
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
PicoClaw supports separating sensitive data (API keys, tokens, secrets, passwords) from the main configuration by storing them in a `.security.yml` file. This improves security by:
|
||||||
|
|
||||||
|
1. **Separation of concerns**: Configuration settings and secrets are in separate files
|
||||||
|
2. **Easier sharing**: The main config can be shared without exposing sensitive data
|
||||||
|
3. **Better version control**: `.security.yml` should be added to `.gitignore`
|
||||||
|
4. **Flexible deployment**: Different environments can use different security files
|
||||||
|
|
||||||
|
## File Structure
|
||||||
|
|
||||||
|
```
|
||||||
|
~/.picoclaw/
|
||||||
|
├── config.json # Main configuration (safe to share)
|
||||||
|
└── .security.yml # Security data (never share)
|
||||||
|
```
|
||||||
|
|
||||||
|
## How It Works
|
||||||
|
|
||||||
|
The security configuration works through **direct field mapping**, NOT through `ref:` string references. The system automatically loads values from `.security.yml` and applies them to the corresponding fields in `config.json`.
|
||||||
|
|
||||||
|
### Key Points:
|
||||||
|
|
||||||
|
- Values in `.security.yml` are automatically mapped to corresponding fields in the config
|
||||||
|
- The mapping is based on field names and structure, not on reference strings
|
||||||
|
- If a value exists in `.security.yml`, it **overrides** the value in `config.json`
|
||||||
|
- You can omit sensitive fields from `config.json` entirely (recommended)
|
||||||
|
|
||||||
|
## Security Configuration Structure
|
||||||
|
|
||||||
|
### Complete Example: .security.yml
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# Model API Keys
|
||||||
|
# All models MUST use `api_keys` (plural) array format
|
||||||
|
# Even a single key must be provided as an array with one element
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "sk-proj-your-actual-openai-key-1"
|
||||||
|
- "sk-proj-your-actual-openai-key-2" # Optional: Multiple keys for failover
|
||||||
|
claude-sonnet-4.6:
|
||||||
|
api_keys:
|
||||||
|
- "sk-ant-your-actual-anthropic-key" # Single key in array format
|
||||||
|
|
||||||
|
# Channel Tokens
|
||||||
|
channels:
|
||||||
|
telegram:
|
||||||
|
token: "your-telegram-bot-token"
|
||||||
|
feishu:
|
||||||
|
app_secret: "your-feishu-app-secret"
|
||||||
|
encrypt_key: "your-feishu-encrypt-key"
|
||||||
|
verification_token: "your-feishu-verification-token"
|
||||||
|
discord:
|
||||||
|
token: "your-discord-bot-token"
|
||||||
|
weixin:
|
||||||
|
token: "your-weixin-token"
|
||||||
|
qq:
|
||||||
|
app_secret: "your-qq-app-secret"
|
||||||
|
dingtalk:
|
||||||
|
client_secret: "your-dingtalk-client-secret"
|
||||||
|
slack:
|
||||||
|
bot_token: "your-slack-bot-token"
|
||||||
|
app_token: "your-slack-app-token"
|
||||||
|
matrix:
|
||||||
|
access_token: "your-matrix-access-token"
|
||||||
|
line:
|
||||||
|
channel_secret: "your-line-channel-secret"
|
||||||
|
channel_access_token: "your-line-channel-access-token"
|
||||||
|
onebot:
|
||||||
|
access_token: "your-onebot-access-token"
|
||||||
|
wecom:
|
||||||
|
token: "your-wecom-token"
|
||||||
|
encoding_aes_key: "your-wecom-encoding-aes-key"
|
||||||
|
wecom_app:
|
||||||
|
corp_secret: "your-wecom-app-corp-secret"
|
||||||
|
token: "your-wecom-app-token"
|
||||||
|
encoding_aes_key: "your-wecom-app-encoding-aes-key"
|
||||||
|
wecom_aibot:
|
||||||
|
secret: "your-wecom-aibot-secret"
|
||||||
|
token: "your-wecom-aibot-token"
|
||||||
|
encoding_aes_key: "your-wecom-aibot-encoding-aes-key"
|
||||||
|
pico:
|
||||||
|
token: "your-pico-token"
|
||||||
|
irc:
|
||||||
|
password: "your-irc-password"
|
||||||
|
nickserv_password: "your-irc-nickserv-password"
|
||||||
|
sasl_password: "your-irc-sasl-password"
|
||||||
|
|
||||||
|
# Web Tool API Keys
|
||||||
|
web:
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "BSAyour-brave-api-key-1"
|
||||||
|
- "BSAyour-brave-api-key-2" # Optional: Multiple keys for failover
|
||||||
|
tavily:
|
||||||
|
api_keys:
|
||||||
|
- "tvly-your-tavily-api-key" # Single key in array format
|
||||||
|
perplexity:
|
||||||
|
api_keys:
|
||||||
|
- "pplx-your-perplexity-api-key" # Single key in array format
|
||||||
|
glm_search:
|
||||||
|
api_key: "your-glm-search-api-key" # GLMSearch uses single key format (not array)
|
||||||
|
baidu_search:
|
||||||
|
api_key: "your-baidu-search-api-key"
|
||||||
|
|
||||||
|
# Skills Registry Tokens
|
||||||
|
skills:
|
||||||
|
github:
|
||||||
|
token: "your-github-token"
|
||||||
|
clawhub:
|
||||||
|
auth_token: "your-clawhub-auth-token"
|
||||||
|
```
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
### Step 1: Create .security.yml
|
||||||
|
|
||||||
|
Create or copy the security file:
|
||||||
|
```bash
|
||||||
|
cp security.example.yml ~/.picoclaw/.security.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 2: Fill in your actual values
|
||||||
|
|
||||||
|
Edit `~/.picoclaw/.security.yml` and replace placeholder values with your actual API keys and tokens.
|
||||||
|
|
||||||
|
### Step 3: Set proper permissions
|
||||||
|
|
||||||
|
```bash
|
||||||
|
chmod 600 ~/.picoclaw/.security.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 4: Simplify config.json (Recommended)
|
||||||
|
|
||||||
|
You can now remove sensitive fields from `config.json` since they're loaded from `.security.yml`:
|
||||||
|
|
||||||
|
**Before:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4",
|
||||||
|
"model": "openai/gpt-5.4",
|
||||||
|
"api_base": "https://api.openai.com/v1",
|
||||||
|
"api_key": "sk-your-actual-api-key-here"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"channels": {
|
||||||
|
"telegram": {
|
||||||
|
"enabled": true,
|
||||||
|
"token": "1234567890:ABCdefGHIjklMNOpqrsTUVwxyz"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**After:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4",
|
||||||
|
"model": "openai/gpt-5.4",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
// api_key is now loaded from .security.yml
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"channels": {
|
||||||
|
"telegram": {
|
||||||
|
"enabled": true"
|
||||||
|
// token is now loaded from .security.yml
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 5: Verify
|
||||||
|
|
||||||
|
Restart PicoClaw and verify it loads correctly:
|
||||||
|
```bash
|
||||||
|
picoclaw --version
|
||||||
|
```
|
||||||
|
|
||||||
|
## Field Mapping Rules
|
||||||
|
|
||||||
|
### Models
|
||||||
|
|
||||||
|
**In .security.yml:**
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
<model_name>:
|
||||||
|
api_keys:
|
||||||
|
- "key-1"
|
||||||
|
- "key-2"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Mapping:**
|
||||||
|
- Field `api_keys` (array) maps to the model's API keys
|
||||||
|
- The `<model_name>` must match the `model_name` field in `config.json`
|
||||||
|
- Supports indexed names (e.g., "gpt-5.4:0") - the system will also try the base name ("gpt-5.4")
|
||||||
|
|
||||||
|
### Channels
|
||||||
|
|
||||||
|
Each channel maps its fields directly:
|
||||||
|
|
||||||
|
**In .security.yml:**
|
||||||
|
```yaml
|
||||||
|
channels:
|
||||||
|
telegram:
|
||||||
|
token: "value"
|
||||||
|
feishu:
|
||||||
|
app_secret: "value"
|
||||||
|
encrypt_key: "value"
|
||||||
|
verification_token: "value"
|
||||||
|
discord:
|
||||||
|
token: "value"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Mapping:**
|
||||||
|
- `channels.telegram.token` → `config.channels.telegram.token`
|
||||||
|
- `channels.feishu.app_secret` → `config.channels.feishu.app_secret`
|
||||||
|
- etc.
|
||||||
|
|
||||||
|
### Web Tools
|
||||||
|
|
||||||
|
**Brave, Tavily, Perplexity:**
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "key-1"
|
||||||
|
- "key-2"
|
||||||
|
```
|
||||||
|
- Use `api_keys` (plural) array format
|
||||||
|
|
||||||
|
**GLMSearch:**
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
glm_search:
|
||||||
|
api_key: "single-key-here"
|
||||||
|
```
|
||||||
|
- Use `api_key` (singular) single string format
|
||||||
|
|
||||||
|
**BaiduSearch:**
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
baidu_search:
|
||||||
|
api_key: "your-key"
|
||||||
|
```
|
||||||
|
- Use `api_key` (singular) single string format
|
||||||
|
|
||||||
|
### Skills
|
||||||
|
|
||||||
|
**In .security.yml:**
|
||||||
|
```yaml
|
||||||
|
skills:
|
||||||
|
github:
|
||||||
|
token: "value"
|
||||||
|
clawhub:
|
||||||
|
auth_token: "value"
|
||||||
|
```
|
||||||
|
|
||||||
|
## API Key Formats
|
||||||
|
|
||||||
|
### Models - Single key
|
||||||
|
|
||||||
|
Use array format with one element:
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "sk-your-key"
|
||||||
|
```
|
||||||
|
|
||||||
|
### Models - Multiple keys (Load Balancing & Failover)
|
||||||
|
|
||||||
|
Use array format with multiple elements:
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "sk-your-key-1"
|
||||||
|
- "sk-your-key-2"
|
||||||
|
- "sk-your-key-3"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Benefits:**
|
||||||
|
- **Load balancing**: Requests are distributed across multiple keys
|
||||||
|
- **Failover**: Automatic switching to another key if one fails
|
||||||
|
- **Rate limit management**: Distribute usage across multiple keys
|
||||||
|
- **High availability**: Reduce downtime during API provider issues
|
||||||
|
|
||||||
|
### Web Tools (Brave/Tavily/Perplexity) - Single key
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "BSA-your-key"
|
||||||
|
```
|
||||||
|
|
||||||
|
### Web Tools (Brave/Tavily/Perplexity) - Multiple keys
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "BSA-key-1"
|
||||||
|
- "BSA-key-2"
|
||||||
|
```
|
||||||
|
|
||||||
|
### Web Tool (GLMSearch/BaiduSearch) - Single key only
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
glm_search:
|
||||||
|
api_key: "your-glm-key" # Single string (NOT array)
|
||||||
|
baidu_search:
|
||||||
|
api_key: "your-baidu-key" # Single string (NOT array)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Model Name Matching
|
||||||
|
|
||||||
|
The system supports intelligent model name matching in `.security.yml`:
|
||||||
|
|
||||||
|
### Example 1: Exact Match
|
||||||
|
|
||||||
|
**config.json:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4:0"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**.security.yml (exact match with index):**
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:0:
|
||||||
|
api_keys: ["key-1"]
|
||||||
|
```
|
||||||
|
|
||||||
|
### Example 2: Base Name Match
|
||||||
|
|
||||||
|
**config.json:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4:0"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**.security.yml (base name without index):**
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys: ["key-1", "key-2"]
|
||||||
|
```
|
||||||
|
|
||||||
|
Both methods work. The base name match allows you to use simpler keys in `.security.yml` even when your config uses indexed model names for load balancing.
|
||||||
|
|
||||||
|
## Backward Compatibility
|
||||||
|
|
||||||
|
The system maintains full backward compatibility:
|
||||||
|
|
||||||
|
1. **Direct values**: You can still use direct values in `config.json` (not recommended for production)
|
||||||
|
2. **Mixed usage**: You can have some fields in `.security.yml` and others in `config.json`
|
||||||
|
3. **Optional security file**: If `.security.yml` doesn't exist, the system will only use values from `config.json`
|
||||||
|
4. **Override behavior**: If a field exists in both files, `.security.yml` value takes precedence
|
||||||
|
|
||||||
|
## Environment Variables
|
||||||
|
|
||||||
|
You can override any security value using environment variables:
|
||||||
|
|
||||||
|
**For models:**
|
||||||
|
```bash
|
||||||
|
export PICOCLAW_CHANNELS_TELEGRAM_TOKEN="token-from-env"
|
||||||
|
```
|
||||||
|
|
||||||
|
**For channels:**
|
||||||
|
```bash
|
||||||
|
export PICOCLAW_CHANNELS_TELEGRAM_TOKEN="token-from-env"
|
||||||
|
export PICOCLAW_CHANNELS_FEISHU_APP_SECRET="secret-from-env"
|
||||||
|
```
|
||||||
|
|
||||||
|
**For web tools:**
|
||||||
|
```bash
|
||||||
|
export PICOCLAW_TOOLS_WEB_BRAVE_API_KEY="key-from-env"
|
||||||
|
export PICOCLAW_TOOLS_WEB_BAIDU_API_KEY="baidu-key-from-env"
|
||||||
|
```
|
||||||
|
|
||||||
|
Environment variables have the highest priority and will override both `config.json` and `.security.yml` values.
|
||||||
|
|
||||||
|
The pattern is: `PICOCLAW_<SECTION>_<KEY>_<FIELD>` with underscores separating path segments and converted to uppercase.
|
||||||
|
|
||||||
|
## Security Best Practices
|
||||||
|
|
||||||
|
1. **Never commit `.security.yml`** to version control
|
||||||
|
2. **Add to .gitignore**: Ensure `.security.yml` is in your `.gitignore` file
|
||||||
|
3. **Set file permissions**: `chmod 600 ~/.picoclaw/.security.yml`
|
||||||
|
4. **Use different keys** for different environments (dev, staging, production)
|
||||||
|
5. **Rotate keys regularly** and update `.security.yml`
|
||||||
|
6. **Backup securely**: Encrypt backups containing `.security.yml`
|
||||||
|
7. **Review access**: Ensure only authorized users have read access to the file
|
||||||
|
|
||||||
|
## API
|
||||||
|
|
||||||
|
### loadSecurityConfig
|
||||||
|
|
||||||
|
```go
|
||||||
|
func loadSecurityConfig(securityPath string) (*SecurityConfig, error)
|
||||||
|
```
|
||||||
|
|
||||||
|
Loads the security configuration from `.security.yml`. Returns an empty `SecurityConfig` if the file doesn't exist.
|
||||||
|
|
||||||
|
### saveSecurityConfig
|
||||||
|
|
||||||
|
```go
|
||||||
|
func saveSecurityConfig(securityPath string, sec *SecurityConfig) error
|
||||||
|
```
|
||||||
|
|
||||||
|
Saves the security configuration to `.security.yml` with `0o600` permissions.
|
||||||
|
|
||||||
|
### applySecurityConfig
|
||||||
|
|
||||||
|
```go
|
||||||
|
func applySecurityConfig(cfg *Config, sec *SecurityConfig) error
|
||||||
|
```
|
||||||
|
|
||||||
|
Applies security configuration to the main config by copying values from `.security.yml` to the corresponding fields in the config.
|
||||||
|
|
||||||
|
### securityPath
|
||||||
|
|
||||||
|
```go
|
||||||
|
func securityPath(configPath string) string
|
||||||
|
```
|
||||||
|
|
||||||
|
Returns the path to `.security.yml` relative to the config file.
|
||||||
|
|
||||||
|
## Example: Complete Configuration
|
||||||
|
|
||||||
|
### config.json
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"version": 1,
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"workspace": "~/picoclaw-workspace",
|
||||||
|
"model_name": "gpt-5.4"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"model_list": [
|
||||||
|
{
|
||||||
|
"model_name": "gpt-5.4",
|
||||||
|
"model": "openai/gpt-5.4",
|
||||||
|
"api_base": "https://api.openai.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "claude-sonnet-4.6",
|
||||||
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
|
"api_base": "https://api.anthropic.com/v1"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"channels": {
|
||||||
|
"telegram": {
|
||||||
|
"enabled": true
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"web": {
|
||||||
|
"brave": {
|
||||||
|
"enabled": true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### .security.yml
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "sk-proj-actual-openai-key-1"
|
||||||
|
- "sk-proj-actual-openai-key-2"
|
||||||
|
claude-sonnet-4.6:
|
||||||
|
api_keys:
|
||||||
|
- "sk-ant-actual-anthropic-key"
|
||||||
|
|
||||||
|
channels:
|
||||||
|
telegram:
|
||||||
|
token: "1234567890:ABCdefGHIjklMNOpqrsTUVwxyz"
|
||||||
|
|
||||||
|
web:
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "BSAactualbravekey-1"
|
||||||
|
- "BSAactualbravekey-2"
|
||||||
|
tavily:
|
||||||
|
api_keys:
|
||||||
|
- "tvly-your-tavily-key"
|
||||||
|
glm_search:
|
||||||
|
api_key: "your-glm-key"
|
||||||
|
baidu_search:
|
||||||
|
api_key: "your-baidu-key"
|
||||||
|
```
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
Run the security configuration tests:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
go test ./pkg/config -run TestSecurityConfig
|
||||||
|
```
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### Error: "failed to load security config"
|
||||||
|
|
||||||
|
- Verify `.security.yml` exists in the same directory as `config.json`
|
||||||
|
- Check the YAML syntax is valid (use a YAML validator)
|
||||||
|
- Ensure file permissions allow reading
|
||||||
|
|
||||||
|
### Error: "model security entry not found"
|
||||||
|
|
||||||
|
- Ensure the model name in `config.json` matches exactly in `.security.yml`
|
||||||
|
- Check that the `model_list` section exists in `.security.yml`
|
||||||
|
- For models with indexed names (e.g., "gpt-5.4:0"), ensure the exact name is used or check the base name without index
|
||||||
|
- Verify the YAML structure is correct (proper indentation)
|
||||||
|
|
||||||
|
### Multiple API Keys Not Working
|
||||||
|
|
||||||
|
- Ensure you're using `api_keys` (plural) in `.security.yml` for models and web tools (except GLMSearch/BaiduSearch)
|
||||||
|
- Check that the array format is correct in YAML (proper indentation with dashes)
|
||||||
|
- Remember: Models, Brave, Tavily, Perplexity MUST use `api_keys` (array format)
|
||||||
|
- GLMSearch and BaiduSearch MUST use `api_key` (single string format)
|
||||||
|
|
||||||
|
### Load Balancing/Failover Issues
|
||||||
|
|
||||||
|
- Verify all API keys in the `api_keys` array are valid
|
||||||
|
- Check that all keys have the same rate limits and permissions
|
||||||
|
- Monitor logs to see which keys are being used and failing
|
||||||
|
- Ensure the `api_keys` array is properly formatted in YAML
|
||||||
|
|
||||||
|
### Keys Not Being Applied
|
||||||
|
|
||||||
|
- Check that `.security.yml` is in the same directory as `config.json`
|
||||||
|
- Verify the file permissions allow reading (`chmod 600 ~/.picoclaw/.security.yml`)
|
||||||
|
- Ensure the YAML structure matches the expected format
|
||||||
|
- Check for typos in field names (case-sensitive)
|
||||||
|
- Verify the model/channel names match exactly (case-sensitive)
|
||||||
|
|
||||||
|
## Migration Guide
|
||||||
|
|
||||||
|
### Step 1: Backup your config
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cp ~/.picoclaw/config.json ~/.picoclaw/config.json.backup
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 2: Create .security.yml
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cp security.example.yml ~/.picoclaw/.security.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 3: Fill in your API keys
|
||||||
|
|
||||||
|
Edit `~/.picoclaw/.security.yml` and replace placeholder values with your actual keys.
|
||||||
|
|
||||||
|
### Step 4: Remove sensitive fields from config.json
|
||||||
|
|
||||||
|
Remove or comment out sensitive fields from `config.json`:
|
||||||
|
- `api_key` fields from `model_list` entries
|
||||||
|
- `token` fields from `channels`
|
||||||
|
- `api_key` fields from `tools.web`
|
||||||
|
- `token`/`auth_token` fields from `tools.skills`
|
||||||
|
|
||||||
|
### Step 5: Set proper permissions
|
||||||
|
|
||||||
|
```bash
|
||||||
|
chmod 600 ~/.picoclaw/.security.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 6: Test
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw --version
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 7: Verify functionality
|
||||||
|
|
||||||
|
Test your models and channels to ensure everything works correctly.
|
||||||
|
|
||||||
|
### Step 8: Clean up (optional)
|
||||||
|
|
||||||
|
If everything works, you can delete the backup:
|
||||||
|
```bash
|
||||||
|
rm ~/.picoclaw/config.json.backup
|
||||||
|
```
|
||||||
|
|
||||||
|
## Advanced: Encrypted API Keys
|
||||||
|
|
||||||
|
PicoClaw supports encrypting API keys in the security file for additional protection.
|
||||||
|
|
||||||
|
### Setup
|
||||||
|
|
||||||
|
1. Set a passphrase via environment variable:
|
||||||
|
```bash
|
||||||
|
export PICOCLAW_CREDENTIAL_PASSPHRASE="your-secure-passphrase"
|
||||||
|
```
|
||||||
|
|
||||||
|
2. When saving config, API keys will be encrypted automatically:
|
||||||
|
```go
|
||||||
|
SaveConfig(path, config)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Encrypted Format
|
||||||
|
|
||||||
|
Encrypted keys are stored as:
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
gpt-5.4:
|
||||||
|
api_keys:
|
||||||
|
- "enc://encrypted-base64-string"
|
||||||
|
```
|
||||||
|
|
||||||
|
The system automatically decrypts keys at runtime when loading the configuration.
|
||||||
|
|
||||||
|
### Benefits
|
||||||
|
|
||||||
|
- Additional layer of security
|
||||||
|
- Keys are encrypted at rest
|
||||||
|
- Passphrase can be managed separately from the config file
|
||||||
|
|
||||||
|
### Important Notes
|
||||||
|
|
||||||
|
- Always backup your passphrase securely
|
||||||
|
- If you lose the passphrase, you'll lose access to encrypted keys
|
||||||
|
- Use a strong, unique passphrase
|
||||||
|
- Never commit the passphrase to version control
|
||||||
|
|
@ -6,7 +6,7 @@
|
||||||
|
|
||||||
PicoClaw 支持多种聊天平台,使您的 Agent 能够连接到任何地方。
|
PicoClaw 支持多种聊天平台,使您的 Agent 能够连接到任何地方。
|
||||||
|
|
||||||
> **注意**: 所有 Webhook 类渠道(LINE、WeCom 等)均挂载在同一个 Gateway HTTP 服务器上(`gateway.host`:`gateway.port`,默认 `127.0.0.1:18790`),无需为每个渠道单独配置端口。注意:飞书(Feishu)使用 WebSocket/SDK 模式,不通过该共享 HTTP webhook 服务器接收消息。
|
> **注意**: 依赖 HTTP 回调的渠道共用同一个 Gateway HTTP 服务器(`gateway.host`:`gateway.port`,默认 `127.0.0.1:18790`),无需为每个渠道单独配置端口。飞书、钉钉、企业微信这类 Socket/Stream 模式渠道不依赖共享 webhook 服务器来接收入站消息。
|
||||||
|
|
||||||
### 核心渠道
|
### 核心渠道
|
||||||
|
|
||||||
|
|
@ -21,7 +21,7 @@ PicoClaw 支持多种聊天平台,使您的 Agent 能够连接到任何地方
|
||||||
| **QQ** | ⭐⭐ 中等 | 官方机器人 API,适合国内社群 | [查看文档](../channels/qq/README.zh.md) |
|
| **QQ** | ⭐⭐ 中等 | 官方机器人 API,适合国内社群 | [查看文档](../channels/qq/README.zh.md) |
|
||||||
| **钉钉 (DingTalk)** | ⭐⭐ 中等 | Stream 模式无需公网,企业办公首选 | [查看文档](../channels/dingtalk/README.zh.md) |
|
| **钉钉 (DingTalk)** | ⭐⭐ 中等 | Stream 模式无需公网,企业办公首选 | [查看文档](../channels/dingtalk/README.zh.md) |
|
||||||
| **LINE** | ⭐⭐⭐ 较难 | 需要 HTTPS Webhook | [查看文档](../channels/line/README.zh.md) |
|
| **LINE** | ⭐⭐⭐ 较难 | 需要 HTTPS Webhook | [查看文档](../channels/line/README.zh.md) |
|
||||||
| **企业微信 (WeCom)** | ⭐⭐⭐ 较难 | 支持群机器人(Webhook)、自建应用(API)和智能机器人(AI Bot) | [Bot 文档](../channels/wecom/wecom_bot/README.zh.md) / [App 文档](../channels/wecom/wecom_app/README.zh.md) / [AI Bot 文档](../channels/wecom/wecom_aibot/README.zh.md) |
|
| **企业微信 (WeCom)** | ⭐⭐⭐ 较难 | 官方 AI Bot WebSocket 接入,支持流式回复和媒体消息 | [查看文档](../channels/wecom/README.zh.md) |
|
||||||
| **飞书 (Feishu)** | ⭐⭐⭐ 较难 | 企业级协作,功能丰富 | [查看文档](../channels/feishu/README.zh.md) |
|
| **飞书 (Feishu)** | ⭐⭐⭐ 较难 | 企业级协作,功能丰富 | [查看文档](../channels/feishu/README.zh.md) |
|
||||||
| **IRC** | ⭐⭐ 中等 | 服务器 + TLS 配置 | [查看文档](#irc) |
|
| **IRC** | ⭐⭐ 中等 | 服务器 + TLS 配置 | [查看文档](#irc) |
|
||||||
| **OneBot** | ⭐⭐ 中等 | 兼容 NapCat/Go-CQHTTP,社区生态丰富 | [查看文档](../channels/onebot/README.zh.md) |
|
| **OneBot** | ⭐⭐ 中等 | 兼容 NapCat/Go-CQHTTP,社区生态丰富 | [查看文档](../channels/onebot/README.zh.md) |
|
||||||
|
|
@ -492,102 +492,34 @@ picoclaw gateway
|
||||||
<details>
|
<details>
|
||||||
<summary><b>企业微信 (WeCom)</b></summary>
|
<summary><b>企业微信 (WeCom)</b></summary>
|
||||||
|
|
||||||
PicoClaw 支持三种企业微信集成方式:
|
PicoClaw 现在将企业微信统一为一个基于 WebSocket 的 AI Bot 渠道。
|
||||||
|
它不再需要公网 webhook 回调地址。
|
||||||
|
|
||||||
**方式 1: 群机器人 (Bot)** — 设置简单,支持群聊
|
完整配置说明和迁移说明请参考 [企业微信配置指南](../channels/wecom/README.zh.md)。
|
||||||
**方式 2: 自建应用 (App)** — 功能更多,支持主动推送,仅私聊
|
|
||||||
**方式 3: 智能机器人 (AI Bot)** — 官方 AI Bot,流式回复,支持群聊和私聊
|
|
||||||
|
|
||||||
详细设置请参考 [企业微信 AI Bot 配置指南](../channels/wecom/wecom_aibot/README.zh.md)。
|
**推荐快速接入**
|
||||||
|
|
||||||
**快速设置 — 群机器人:**
|
**1. 认证**
|
||||||
|
|
||||||
**1. 创建 Bot**
|
```bash
|
||||||
|
picoclaw auth wecom
|
||||||
|
```
|
||||||
|
|
||||||
* 企业微信管理后台 → 群聊 → 添加群机器人
|
该命令会显示二维码,等待你在企业微信里确认,然后把 `bot_id` 和 `secret` 写入 `channels.wecom`。
|
||||||
* 复制 Webhook URL(格式:`https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=xxx`)
|
|
||||||
|
|
||||||
**2. 配置**
|
**2. 如需手动配置**
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"channels": {
|
"channels": {
|
||||||
"wecom": {
|
"wecom": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"token": "YOUR_TOKEN",
|
"bot_id": "YOUR_BOT_ID",
|
||||||
"encoding_aes_key": "YOUR_ENCODING_AES_KEY",
|
"secret": "YOUR_SECRET",
|
||||||
"webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=YOUR_KEY",
|
"websocket_url": "wss://openws.work.weixin.qq.com",
|
||||||
"webhook_path": "/webhook/wecom",
|
"send_thinking_message": true,
|
||||||
"allow_from": []
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
> WeCom Webhook 挂载在共享 Gateway 服务器上(`gateway.host`:`gateway.port`,默认 `127.0.0.1:18790`)。
|
|
||||||
|
|
||||||
**快速设置 — 自建应用:**
|
|
||||||
|
|
||||||
**1. 创建应用**
|
|
||||||
|
|
||||||
* 企业微信管理后台 → 应用管理 → 创建应用
|
|
||||||
* 复制 **AgentId** 和 **Secret**
|
|
||||||
* 前往"我的企业"页面,复制 **CorpID**
|
|
||||||
|
|
||||||
**2. 配置接收消息**
|
|
||||||
|
|
||||||
* 在应用详情中,点击"接收消息" → "设置 API"
|
|
||||||
* 设置 URL 为 `http://your-server:18790/webhook/wecom-app`
|
|
||||||
* 生成 **Token** 和 **EncodingAESKey**
|
|
||||||
|
|
||||||
**3. 配置**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"channels": {
|
|
||||||
"wecom_app": {
|
|
||||||
"enabled": true,
|
|
||||||
"corp_id": "wwxxxxxxxxxxxxxxxx",
|
|
||||||
"corp_secret": "YOUR_CORP_SECRET",
|
|
||||||
"agent_id": 1000002,
|
|
||||||
"token": "YOUR_TOKEN",
|
|
||||||
"encoding_aes_key": "YOUR_ENCODING_AES_KEY",
|
|
||||||
"webhook_path": "/webhook/wecom-app",
|
|
||||||
"allow_from": []
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**4. 运行**
|
|
||||||
|
|
||||||
```bash
|
|
||||||
picoclaw gateway
|
|
||||||
```
|
|
||||||
|
|
||||||
> **注意**: WeCom Webhook 回调挂载在 Gateway 端口(默认 18790)。使用反向代理配置 HTTPS。
|
|
||||||
|
|
||||||
**快速设置 — 智能机器人 (AI Bot):**
|
|
||||||
|
|
||||||
**1. 创建 AI Bot**
|
|
||||||
|
|
||||||
* 企业微信管理后台 → 应用管理 → AI Bot
|
|
||||||
* 在 AI Bot 设置中配置回调 URL:`http://your-server:18790/webhook/wecom-aibot`
|
|
||||||
* 复制 **Token** 并点击"随机生成" **EncodingAESKey**
|
|
||||||
|
|
||||||
**2. 配置**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"channels": {
|
|
||||||
"wecom_aibot": {
|
|
||||||
"enabled": true,
|
|
||||||
"token": "YOUR_TOKEN",
|
|
||||||
"encoding_aes_key": "YOUR_43_CHAR_ENCODING_AES_KEY",
|
|
||||||
"webhook_path": "/webhook/wecom-aibot",
|
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"welcome_message": "你好!有什么可以帮你的?",
|
"reasoning_channel_id": ""
|
||||||
"processing_message": "⏳ Processing, please wait. The results will be sent shortly."
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -599,7 +531,7 @@ picoclaw gateway
|
||||||
picoclaw gateway
|
picoclaw gateway
|
||||||
```
|
```
|
||||||
|
|
||||||
> **注意**: 企业微信 AI Bot 使用流式拉取协议,无回复超时问题。长任务(>30 秒)会自动切换到 `response_url` 推送投递。
|
> 这个分支中旧的 `wecom_app` 和 `wecom_aibot` 配置已经被统一的 `channels.wecom` 替代。
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
|
|
||||||
3
go.mod
3
go.mod
|
|
@ -20,6 +20,7 @@ require (
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
github.com/h2non/filetype v1.1.3
|
github.com/h2non/filetype v1.1.3
|
||||||
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
github.com/larksuite/oapi-sdk-go/v3 v3.5.3
|
||||||
|
github.com/mattn/go-sqlite3 v1.14.34
|
||||||
github.com/mdp/qrterminal/v3 v3.2.1
|
github.com/mdp/qrterminal/v3 v3.2.1
|
||||||
github.com/modelcontextprotocol/go-sdk v1.4.1
|
github.com/modelcontextprotocol/go-sdk v1.4.1
|
||||||
github.com/mymmrac/telego v1.7.0
|
github.com/mymmrac/telego v1.7.0
|
||||||
|
|
@ -31,6 +32,7 @@ require (
|
||||||
github.com/spf13/cobra v1.10.2
|
github.com/spf13/cobra v1.10.2
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/tencent-connect/botgo v0.2.1
|
github.com/tencent-connect/botgo v0.2.1
|
||||||
|
go.mau.fi/util v0.9.7
|
||||||
go.mau.fi/whatsmeow v0.0.0-20260219150138-7ae702b1eed4
|
go.mau.fi/whatsmeow v0.0.0-20260219150138-7ae702b1eed4
|
||||||
golang.org/x/oauth2 v0.36.0
|
golang.org/x/oauth2 v0.36.0
|
||||||
golang.org/x/term v0.41.0
|
golang.org/x/term v0.41.0
|
||||||
|
|
@ -77,7 +79,6 @@ require (
|
||||||
github.com/spf13/pflag v1.0.10 // indirect
|
github.com/spf13/pflag v1.0.10 // indirect
|
||||||
github.com/vektah/gqlparser/v2 v2.5.27 // indirect
|
github.com/vektah/gqlparser/v2 v2.5.27 // indirect
|
||||||
go.mau.fi/libsignal v0.2.1 // indirect
|
go.mau.fi/libsignal v0.2.1 // indirect
|
||||||
go.mau.fi/util v0.9.7 // indirect
|
|
||||||
golang.org/x/exp v0.0.0-20260312153236-7ab1446f8b90 // indirect
|
golang.org/x/exp v0.0.0-20260312153236-7ab1446f8b90 // indirect
|
||||||
golang.org/x/text v0.35.0 // indirect
|
golang.org/x/text v0.35.0 // indirect
|
||||||
modernc.org/libc v1.67.6 // indirect
|
modernc.org/libc v1.67.6 // indirect
|
||||||
|
|
|
||||||
|
|
@ -85,6 +85,7 @@ type processOptions struct {
|
||||||
DefaultResponse string // Response when LLM returns empty
|
DefaultResponse string // Response when LLM returns empty
|
||||||
EnableSummary bool // Whether to trigger summarization
|
EnableSummary bool // Whether to trigger summarization
|
||||||
SendResponse bool // Whether to send response via bus
|
SendResponse bool // Whether to send response via bus
|
||||||
|
SuppressToolFeedback bool // Whether to suppress inline tool feedback messages
|
||||||
NoHistory bool // If true, don't load session history (for heartbeat)
|
NoHistory bool // If true, don't load session history (for heartbeat)
|
||||||
SkipInitialSteeringPoll bool // If true, skip the steering poll at loop start (used by Continue)
|
SkipInitialSteeringPoll bool // If true, skip the steering poll at loop start (used by Continue)
|
||||||
}
|
}
|
||||||
|
|
@ -96,14 +97,15 @@ type continuationTarget struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
defaultResponse = "The model returned an empty response. This may indicate a provider error or token limit."
|
defaultResponse = "The model returned an empty response. This may indicate a provider error or token limit."
|
||||||
toolLimitResponse = "I've reached `max_tool_iterations` without a final response. Increase `max_tool_iterations` in config.json if this task needs more tool steps."
|
toolLimitResponse = "I've reached `max_tool_iterations` without a final response. Increase `max_tool_iterations` in config.json if this task needs more tool steps."
|
||||||
sessionKeyAgentPrefix = "agent:"
|
handledToolResponseSummary = "Requested output delivered via tool attachment."
|
||||||
metadataKeyAccountID = "account_id"
|
sessionKeyAgentPrefix = "agent:"
|
||||||
metadataKeyGuildID = "guild_id"
|
metadataKeyAccountID = "account_id"
|
||||||
metadataKeyTeamID = "team_id"
|
metadataKeyGuildID = "guild_id"
|
||||||
metadataKeyParentPeerKind = "parent_peer_kind"
|
metadataKeyTeamID = "team_id"
|
||||||
metadataKeyParentPeerID = "parent_peer_id"
|
metadataKeyParentPeerKind = "parent_peer_kind"
|
||||||
|
metadataKeyParentPeerID = "parent_peer_id"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewAgentLoop(
|
func NewAgentLoop(
|
||||||
|
|
@ -1030,13 +1032,13 @@ func (al *AgentLoop) GetConfig() *config.Config {
|
||||||
func (al *AgentLoop) SetMediaStore(s media.MediaStore) {
|
func (al *AgentLoop) SetMediaStore(s media.MediaStore) {
|
||||||
al.mediaStore = s
|
al.mediaStore = s
|
||||||
|
|
||||||
// Propagate store to send_file tools in all agents.
|
// Propagate store to all registered tools that can emit media.
|
||||||
registry := al.GetRegistry()
|
registry := al.GetRegistry()
|
||||||
registry.ForEachTool("send_file", func(t tools.Tool) {
|
for _, agentID := range registry.ListAgentIDs() {
|
||||||
if sf, ok := t.(*tools.SendFileTool); ok {
|
if agent, ok := registry.GetAgent(agentID); ok {
|
||||||
sf.SetMediaStore(s)
|
agent.Tools.SetMediaStore(s)
|
||||||
}
|
}
|
||||||
})
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetTranscriber injects a voice transcriber for agent-level audio transcription.
|
// SetTranscriber injects a voice transcriber for agent-level audio transcription.
|
||||||
|
|
@ -1241,14 +1243,15 @@ func (al *AgentLoop) ProcessHeartbeat(
|
||||||
return "", fmt.Errorf("no default agent for heartbeat")
|
return "", fmt.Errorf("no default agent for heartbeat")
|
||||||
}
|
}
|
||||||
return al.runAgentLoop(ctx, agent, processOptions{
|
return al.runAgentLoop(ctx, agent, processOptions{
|
||||||
SessionKey: "heartbeat",
|
SessionKey: "heartbeat",
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
ChatID: chatID,
|
ChatID: chatID,
|
||||||
UserMessage: content,
|
UserMessage: content,
|
||||||
DefaultResponse: defaultResponse,
|
DefaultResponse: defaultResponse,
|
||||||
EnableSummary: false,
|
EnableSummary: false,
|
||||||
SendResponse: false,
|
SendResponse: false,
|
||||||
NoHistory: true, // Don't load session history for heartbeat
|
SuppressToolFeedback: true,
|
||||||
|
NoHistory: true, // Don't load session history for heartbeat
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -2165,6 +2168,7 @@ turnLoop:
|
||||||
"iteration": iteration,
|
"iteration": iteration,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
allResponsesHandled := len(normalizedToolCalls) > 0
|
||||||
assistantMsg := providers.Message{
|
assistantMsg := providers.Message{
|
||||||
Role: "assistant",
|
Role: "assistant",
|
||||||
Content: response.Content,
|
Content: response.Content,
|
||||||
|
|
@ -2221,6 +2225,7 @@ turnLoop:
|
||||||
toolArgs = toolReq.Arguments
|
toolArgs = toolReq.Arguments
|
||||||
}
|
}
|
||||||
case HookActionDenyTool:
|
case HookActionDenyTool:
|
||||||
|
allResponsesHandled = false
|
||||||
denyContent := hookDeniedToolContent("Tool execution denied by hook", decision.Reason)
|
denyContent := hookDeniedToolContent("Tool execution denied by hook", decision.Reason)
|
||||||
al.emitEvent(
|
al.emitEvent(
|
||||||
EventKindToolExecSkipped,
|
EventKindToolExecSkipped,
|
||||||
|
|
@ -2260,6 +2265,7 @@ turnLoop:
|
||||||
ChatID: ts.chatID,
|
ChatID: ts.chatID,
|
||||||
})
|
})
|
||||||
if !approval.Approved {
|
if !approval.Approved {
|
||||||
|
allResponsesHandled = false
|
||||||
denyContent := hookDeniedToolContent("Tool execution denied by approval hook", approval.Reason)
|
denyContent := hookDeniedToolContent("Tool execution denied by approval hook", approval.Reason)
|
||||||
al.emitEvent(
|
al.emitEvent(
|
||||||
EventKindToolExecSkipped,
|
EventKindToolExecSkipped,
|
||||||
|
|
@ -2301,7 +2307,9 @@ turnLoop:
|
||||||
)
|
)
|
||||||
|
|
||||||
// Send tool feedback to chat channel if enabled (from HEAD)
|
// Send tool feedback to chat channel if enabled (from HEAD)
|
||||||
if al.cfg.Agents.Defaults.IsToolFeedbackEnabled() && ts.channel != "" {
|
if al.cfg.Agents.Defaults.IsToolFeedbackEnabled() &&
|
||||||
|
ts.channel != "" &&
|
||||||
|
!ts.opts.SuppressToolFeedback {
|
||||||
feedbackPreview := utils.Truncate(
|
feedbackPreview := utils.Truncate(
|
||||||
string(argsJSON),
|
string(argsJSON),
|
||||||
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
||||||
|
|
@ -2333,10 +2341,7 @@ turnLoop:
|
||||||
}
|
}
|
||||||
|
|
||||||
// Determine content for the agent loop (ForLLM or error).
|
// Determine content for the agent loop (ForLLM or error).
|
||||||
content := result.ForLLM
|
content := result.ContentForLLM()
|
||||||
if content == "" && result.Err != nil {
|
|
||||||
content = result.Err.Error()
|
|
||||||
}
|
|
||||||
if content == "" {
|
if content == "" {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
@ -2420,6 +2425,50 @@ turnLoop:
|
||||||
if toolResult == nil {
|
if toolResult == nil {
|
||||||
toolResult = tools.ErrorResult("hook returned nil tool result")
|
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,
|
||||||
|
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 if al.bus != nil {
|
||||||
|
al.bus.PublishOutboundMedia(ctx, outboundMedia)
|
||||||
|
// Queuing media is only best-effort; it has not been delivered yet.
|
||||||
|
toolResult.ResponseHandled = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(toolResult.Media) > 0 && !toolResult.ResponseHandled {
|
||||||
|
toolResult.ArtifactTags = buildArtifactTags(al.mediaStore, toolResult.Media)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !toolResult.ResponseHandled {
|
||||||
|
allResponsesHandled = false
|
||||||
|
}
|
||||||
|
|
||||||
if !toolResult.Silent && toolResult.ForUser != "" && ts.opts.SendResponse {
|
if !toolResult.Silent && toolResult.ForUser != "" && ts.opts.SendResponse {
|
||||||
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||||
|
|
@ -2434,30 +2483,7 @@ turnLoop:
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(toolResult.Media) > 0 {
|
contentForLLM := toolResult.ContentForLLM()
|
||||||
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)
|
|
||||||
}
|
|
||||||
al.bus.PublishOutboundMedia(ctx, bus.OutboundMediaMessage{
|
|
||||||
Channel: ts.channel,
|
|
||||||
ChatID: ts.chatID,
|
|
||||||
Parts: parts,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
contentForLLM := toolResult.ForLLM
|
|
||||||
if contentForLLM == "" && toolResult.Err != nil {
|
|
||||||
contentForLLM = toolResult.Err.Error()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Filter sensitive data (API keys, tokens, secrets) before sending to LLM
|
// Filter sensitive data (API keys, tokens, secrets) before sending to LLM
|
||||||
if al.cfg.Tools.IsFilterSensitiveDataEnabled() {
|
if al.cfg.Tools.IsFilterSensitiveDataEnabled() {
|
||||||
|
|
@ -2552,6 +2578,70 @@ turnLoop:
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if allResponsesHandled {
|
||||||
|
if len(pendingMessages) > 0 {
|
||||||
|
logger.InfoCF("agent", "Pending steering exists after handled tool delivery; continuing turn before finalizing",
|
||||||
|
map[string]any{
|
||||||
|
"agent_id": ts.agent.ID,
|
||||||
|
"steering_count": len(pendingMessages),
|
||||||
|
"session_key": ts.sessionKey,
|
||||||
|
})
|
||||||
|
finalContent = ""
|
||||||
|
goto turnLoop
|
||||||
|
}
|
||||||
|
|
||||||
|
if steerMsgs := al.dequeueSteeringMessagesForScope(ts.sessionKey); len(steerMsgs) > 0 {
|
||||||
|
logger.InfoCF("agent", "Steering arrived after handled tool delivery; continuing turn before finalizing",
|
||||||
|
map[string]any{
|
||||||
|
"agent_id": ts.agent.ID,
|
||||||
|
"steering_count": len(steerMsgs),
|
||||||
|
"session_key": ts.sessionKey,
|
||||||
|
})
|
||||||
|
pendingMessages = append(pendingMessages, steerMsgs...)
|
||||||
|
finalContent = ""
|
||||||
|
goto turnLoop
|
||||||
|
}
|
||||||
|
|
||||||
|
summaryMsg := providers.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: handledToolResponseSummary,
|
||||||
|
}
|
||||||
|
|
||||||
|
if !ts.opts.NoHistory {
|
||||||
|
ts.agent.Sessions.AddMessage(ts.sessionKey, summaryMsg.Role, summaryMsg.Content)
|
||||||
|
ts.recordPersistedMessage(summaryMsg)
|
||||||
|
if err := ts.agent.Sessions.Save(ts.sessionKey); err != nil {
|
||||||
|
turnStatus = TurnEndStatusError
|
||||||
|
al.emitEvent(
|
||||||
|
EventKindError,
|
||||||
|
ts.eventMeta("runTurn", "turn.error"),
|
||||||
|
ErrorPayload{
|
||||||
|
Stage: "session_save",
|
||||||
|
Message: err.Error(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return turnResult{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if ts.opts.EnableSummary {
|
||||||
|
al.maybeSummarize(ts.agent, ts.sessionKey, ts.scope)
|
||||||
|
}
|
||||||
|
|
||||||
|
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 turnResult{
|
||||||
|
finalContent: "",
|
||||||
|
status: turnStatus,
|
||||||
|
followUps: append([]bus.InboundMessage(nil), ts.followUps...),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
ts.agent.Tools.TickTTL()
|
ts.agent.Tools.TickTTL()
|
||||||
logger.DebugCF("agent", "TTL tick after tool execution", map[string]any{
|
logger.DebugCF("agent", "TTL tick after tool execution", map[string]any{
|
||||||
"agent_id": ts.agent.ID, "iteration": iteration,
|
"agent_id": ts.agent.ID, "iteration": iteration,
|
||||||
|
|
@ -3159,6 +3249,97 @@ func (al *AgentLoop) handleCommand(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func activeSkillNames(agent *AgentInstance, opts processOptions) []string {
|
||||||
|
if agent == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
combined := make([]string, 0, len(agent.SkillsFilter)+len(opts.ForcedSkills))
|
||||||
|
combined = append(combined, agent.SkillsFilter...)
|
||||||
|
combined = append(combined, opts.ForcedSkills...)
|
||||||
|
if len(combined) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var resolved []string
|
||||||
|
seen := make(map[string]struct{}, len(combined))
|
||||||
|
for _, name := range combined {
|
||||||
|
name = strings.TrimSpace(name)
|
||||||
|
if name == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if agent.ContextBuilder != nil {
|
||||||
|
if canonical, ok := agent.ContextBuilder.ResolveSkillName(name); ok {
|
||||||
|
name = canonical
|
||||||
|
}
|
||||||
|
}
|
||||||
|
key := strings.ToLower(name)
|
||||||
|
if _, ok := seen[key]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[key] = struct{}{}
|
||||||
|
resolved = append(resolved, name)
|
||||||
|
}
|
||||||
|
|
||||||
|
return resolved
|
||||||
|
}
|
||||||
|
|
||||||
|
func (al *AgentLoop) applyExplicitSkillCommand(
|
||||||
|
raw string,
|
||||||
|
agent *AgentInstance,
|
||||||
|
opts *processOptions,
|
||||||
|
) (matched bool, handled bool, reply string) {
|
||||||
|
cmdName, ok := commands.CommandName(raw)
|
||||||
|
if !ok || cmdName != "use" {
|
||||||
|
return false, false, ""
|
||||||
|
}
|
||||||
|
|
||||||
|
if agent == nil || agent.ContextBuilder == nil {
|
||||||
|
return true, true, commandsUnavailableSkillMessage()
|
||||||
|
}
|
||||||
|
|
||||||
|
parts := strings.Fields(strings.TrimSpace(raw))
|
||||||
|
if len(parts) < 2 {
|
||||||
|
return true, true, buildUseCommandHelp(agent)
|
||||||
|
}
|
||||||
|
|
||||||
|
arg := strings.TrimSpace(parts[1])
|
||||||
|
if strings.EqualFold(arg, "clear") || strings.EqualFold(arg, "off") {
|
||||||
|
if opts != nil {
|
||||||
|
al.clearPendingSkills(opts.SessionKey)
|
||||||
|
}
|
||||||
|
return true, true, "Cleared pending skill override."
|
||||||
|
}
|
||||||
|
|
||||||
|
skillName, ok := agent.ContextBuilder.ResolveSkillName(arg)
|
||||||
|
if !ok {
|
||||||
|
return true, true, fmt.Sprintf("Unknown skill: %s\nUse /list skills to see installed skills.", arg)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(parts) < 3 {
|
||||||
|
if opts == nil || strings.TrimSpace(opts.SessionKey) == "" {
|
||||||
|
return true, true, commandsUnavailableSkillMessage()
|
||||||
|
}
|
||||||
|
al.setPendingSkills(opts.SessionKey, []string{skillName})
|
||||||
|
return true, true, fmt.Sprintf(
|
||||||
|
"Skill %q is armed for your next message. Send your next prompt normally, or use /use clear to cancel.",
|
||||||
|
skillName,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
message := strings.TrimSpace(strings.Join(parts[2:], " "))
|
||||||
|
if message == "" {
|
||||||
|
return true, true, buildUseCommandHelp(agent)
|
||||||
|
}
|
||||||
|
|
||||||
|
if opts != nil {
|
||||||
|
opts.ForcedSkills = append(opts.ForcedSkills, skillName)
|
||||||
|
opts.UserMessage = message
|
||||||
|
}
|
||||||
|
|
||||||
|
return true, false, ""
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, opts *processOptions) *commands.Runtime {
|
func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, opts *processOptions) *commands.Runtime {
|
||||||
registry := al.GetRegistry()
|
registry := al.GetRegistry()
|
||||||
cfg := al.GetConfig()
|
cfg := al.GetConfig()
|
||||||
|
|
@ -3199,6 +3380,9 @@ func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, opts *processOpt
|
||||||
return al.reloadFunc()
|
return al.reloadFunc()
|
||||||
}
|
}
|
||||||
if agent != nil {
|
if agent != nil {
|
||||||
|
if agent.ContextBuilder != nil {
|
||||||
|
rt.ListSkillNames = agent.ContextBuilder.ListSkillNames
|
||||||
|
}
|
||||||
rt.GetModelInfo = func() (string, string) {
|
rt.GetModelInfo = func() (string, string) {
|
||||||
return agent.Model, resolvedCandidateProvider(agent.Candidates, cfg.Agents.Defaults.Provider)
|
return agent.Model, resolvedCandidateProvider(agent.Candidates, cfg.Agents.Defaults.Provider)
|
||||||
}
|
}
|
||||||
|
|
@ -3251,79 +3435,6 @@ func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, opts *processOpt
|
||||||
return rt
|
return rt
|
||||||
}
|
}
|
||||||
|
|
||||||
func activeSkillNames(agent *AgentInstance, opts processOptions) []string {
|
|
||||||
var out []string
|
|
||||||
seen := make(map[string]struct{})
|
|
||||||
|
|
||||||
appendNames := func(names []string) {
|
|
||||||
for _, name := range names {
|
|
||||||
name = strings.TrimSpace(name)
|
|
||||||
if name == "" {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if _, exists := seen[name]; exists {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seen[name] = struct{}{}
|
|
||||||
out = append(out, name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if agent != nil {
|
|
||||||
appendNames(agent.SkillsFilter)
|
|
||||||
}
|
|
||||||
appendNames(opts.ForcedSkills)
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func (al *AgentLoop) applyExplicitSkillCommand(
|
|
||||||
raw string,
|
|
||||||
agent *AgentInstance,
|
|
||||||
opts *processOptions,
|
|
||||||
) (matched bool, handled bool, reply string) {
|
|
||||||
commandName, ok := commands.CommandName(raw)
|
|
||||||
if !ok || commandName != "use" {
|
|
||||||
return false, false, ""
|
|
||||||
}
|
|
||||||
|
|
||||||
if agent == nil || agent.ContextBuilder == nil {
|
|
||||||
return true, true, commandsUnavailableSkillMessage()
|
|
||||||
}
|
|
||||||
|
|
||||||
fields := strings.Fields(strings.TrimSpace(raw))
|
|
||||||
if len(fields) < 2 {
|
|
||||||
return true, true, buildUseCommandHelp(agent)
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.EqualFold(fields[1], "clear") || strings.EqualFold(fields[1], "off") {
|
|
||||||
al.clearPendingSkills(opts.SessionKey)
|
|
||||||
return true, true, "Cleared pending skill override."
|
|
||||||
}
|
|
||||||
|
|
||||||
canonicalSkill, ok := agent.ContextBuilder.ResolveSkillName(fields[1])
|
|
||||||
if !ok {
|
|
||||||
return true, true, fmt.Sprintf("Unknown skill: %s\nUse /list skills to see installed skills.", fields[1])
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(fields) == 2 {
|
|
||||||
al.setPendingSkills(opts.SessionKey, []string{canonicalSkill})
|
|
||||||
return true, true, fmt.Sprintf(
|
|
||||||
"Skill %q is armed for your next message.\nSend your next request normally, or use /use clear to cancel.",
|
|
||||||
canonicalSkill,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
message := strings.TrimSpace(strings.Join(fields[2:], " "))
|
|
||||||
if message == "" {
|
|
||||||
return true, true, buildUseCommandHelp(agent)
|
|
||||||
}
|
|
||||||
|
|
||||||
opts.UserMessage = message
|
|
||||||
opts.ForcedSkills = append(opts.ForcedSkills, canonicalSkill)
|
|
||||||
return true, false, ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func commandsUnavailableSkillMessage() string {
|
func commandsUnavailableSkillMessage() string {
|
||||||
return "Skill selection is unavailable in the current context."
|
return "Skill selection is unavailable in the current context."
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -87,6 +87,24 @@ func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxS
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func buildArtifactTags(store media.MediaStore, refs []string) []string {
|
||||||
|
if store == nil || len(refs) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
tags := make([]string, 0, len(refs))
|
||||||
|
for _, ref := range refs {
|
||||||
|
localPath, meta, err := store.ResolveWithMeta(ref)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
mime := detectMIME(localPath, meta)
|
||||||
|
tags = append(tags, buildPathTag(mime, localPath))
|
||||||
|
}
|
||||||
|
|
||||||
|
return tags
|
||||||
|
}
|
||||||
|
|
||||||
// detectMIME determines the MIME type from metadata or magic-bytes detection.
|
// detectMIME determines the MIME type from metadata or magic-bytes detection.
|
||||||
// Returns empty string if detection fails.
|
// Returns empty string if detection fails.
|
||||||
func detectMIME(localPath string, meta media.MediaMeta) string {
|
func detectMIME(localPath string, meta media.MediaMeta) string {
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,41 @@ func (f *fakeChannel) IsAllowed(string) bool {
|
||||||
func (f *fakeChannel) IsAllowedSender(sender bus.SenderInfo) bool { return true }
|
func (f *fakeChannel) IsAllowedSender(sender bus.SenderInfo) bool { return true }
|
||||||
func (f *fakeChannel) ReasoningChannelID() string { return f.id }
|
func (f *fakeChannel) ReasoningChannelID() string { return f.id }
|
||||||
|
|
||||||
|
type fakeMediaChannel struct {
|
||||||
|
fakeChannel
|
||||||
|
sentMedia []bus.OutboundMediaMessage
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
|
f.sentMedia = append(f.sentMedia, msg)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newStartedTestChannelManager(
|
||||||
|
t *testing.T,
|
||||||
|
msgBus *bus.MessageBus,
|
||||||
|
store media.MediaStore,
|
||||||
|
name string,
|
||||||
|
ch channels.Channel,
|
||||||
|
) *channels.Manager {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
cm, err := channels.NewManager(&config.Config{}, msgBus, store)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewManager() error = %v", err)
|
||||||
|
}
|
||||||
|
cm.RegisterChannel(name, ch)
|
||||||
|
if err := cm.StartAll(context.Background()); err != nil {
|
||||||
|
t.Fatalf("StartAll() error = %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := cm.StopAll(context.Background()); err != nil {
|
||||||
|
t.Fatalf("StopAll() error = %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return cm
|
||||||
|
}
|
||||||
|
|
||||||
type recordingProvider struct {
|
type recordingProvider struct {
|
||||||
lastMessages []providers.Message
|
lastMessages []providers.Message
|
||||||
}
|
}
|
||||||
|
|
@ -289,6 +324,86 @@ func TestProcessMessage_UseCommandArmsSkillForNextMessage(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestApplyExplicitSkillCommand_ArmsSkillForNextMessage(t *testing.T) {
|
||||||
|
al, cfg, _, _, cleanup := newTestAgentLoop(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
if err := os.MkdirAll(filepath.Join(cfg.Agents.Defaults.Workspace, "skills", "finance-news"), 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll(skill) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(cfg.Agents.Defaults.Workspace, "skills", "finance-news", "SKILL.md"),
|
||||||
|
[]byte("# Finance News\n\nUse web tools for current finance updates.\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile(SKILL.md) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
agent := al.GetRegistry().GetDefaultAgent()
|
||||||
|
if agent == nil {
|
||||||
|
t.Fatal("expected default agent")
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := &processOptions{SessionKey: "agent:main:test"}
|
||||||
|
matched, handled, reply := al.applyExplicitSkillCommand("/use finance-news", agent, opts)
|
||||||
|
if !matched {
|
||||||
|
t.Fatal("expected /use command to match")
|
||||||
|
}
|
||||||
|
if !handled {
|
||||||
|
t.Fatal("expected /use without inline message to be handled immediately")
|
||||||
|
}
|
||||||
|
if !strings.Contains(reply, `Skill "finance-news" is armed for your next message`) {
|
||||||
|
t.Fatalf("unexpected reply: %q", reply)
|
||||||
|
}
|
||||||
|
|
||||||
|
pending := al.takePendingSkills(opts.SessionKey)
|
||||||
|
if len(pending) != 1 || pending[0] != "finance-news" {
|
||||||
|
t.Fatalf("pending skills = %#v, want [finance-news]", pending)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyExplicitSkillCommand_InlineMessageMutatesOptions(t *testing.T) {
|
||||||
|
al, cfg, _, _, cleanup := newTestAgentLoop(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
if err := os.MkdirAll(filepath.Join(cfg.Agents.Defaults.Workspace, "skills", "finance-news"), 0o755); err != nil {
|
||||||
|
t.Fatalf("MkdirAll(skill) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(
|
||||||
|
filepath.Join(cfg.Agents.Defaults.Workspace, "skills", "finance-news", "SKILL.md"),
|
||||||
|
[]byte("# Finance News\n\nUse web tools for current finance updates.\n"),
|
||||||
|
0o644,
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("WriteFile(SKILL.md) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
agent := al.GetRegistry().GetDefaultAgent()
|
||||||
|
if agent == nil {
|
||||||
|
t.Fatal("expected default agent")
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := &processOptions{
|
||||||
|
SessionKey: "agent:main:test",
|
||||||
|
UserMessage: "/use finance-news dammi le ultime news",
|
||||||
|
}
|
||||||
|
matched, handled, reply := al.applyExplicitSkillCommand(opts.UserMessage, agent, opts)
|
||||||
|
if !matched {
|
||||||
|
t.Fatal("expected /use command to match")
|
||||||
|
}
|
||||||
|
if handled {
|
||||||
|
t.Fatal("expected /use with inline message to fall through into normal agent execution")
|
||||||
|
}
|
||||||
|
if reply != "" {
|
||||||
|
t.Fatalf("unexpected reply: %q", reply)
|
||||||
|
}
|
||||||
|
if opts.UserMessage != "dammi le ultime news" {
|
||||||
|
t.Fatalf("opts.UserMessage = %q, want %q", opts.UserMessage, "dammi le ultime news")
|
||||||
|
}
|
||||||
|
if len(opts.ForcedSkills) != 1 || opts.ForcedSkills[0] != "finance-news" {
|
||||||
|
t.Fatalf("opts.ForcedSkills = %#v, want [finance-news]", opts.ForcedSkills)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestRecordLastChannel(t *testing.T) {
|
func TestRecordLastChannel(t *testing.T) {
|
||||||
al, cfg, msgBus, provider, cleanup := newTestAgentLoop(t)
|
al, cfg, msgBus, provider, cleanup := newTestAgentLoop(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
@ -455,6 +570,217 @@ func TestToolRegistry_GetDefinitions(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProcessMessage_MediaToolHandledSkipsFollowUpLLMAndFinalText(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
ModelName: "test-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &handledMediaProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
al.SetMediaStore(store)
|
||||||
|
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
|
||||||
|
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
|
||||||
|
|
||||||
|
imagePath := filepath.Join(tmpDir, "screen.png")
|
||||||
|
if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil {
|
||||||
|
t.Fatalf("WriteFile(imagePath) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
al.RegisterTool(&handledMediaTool{
|
||||||
|
store: store,
|
||||||
|
path: imagePath,
|
||||||
|
})
|
||||||
|
|
||||||
|
response, err := al.processMessage(context.Background(), bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat1",
|
||||||
|
SenderID: "user1",
|
||||||
|
Content: "take a screenshot of the screen and send it to me",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
if response != "" {
|
||||||
|
t.Fatalf("expected no final response when media tool already handled delivery, got %q", response)
|
||||||
|
}
|
||||||
|
if provider.calls != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 LLM call, got %d", provider.calls)
|
||||||
|
}
|
||||||
|
if len(provider.toolCounts) != 1 {
|
||||||
|
t.Fatalf("expected tool counts for 1 provider call, got %d", len(provider.toolCounts))
|
||||||
|
}
|
||||||
|
if provider.toolCounts[0] == 0 {
|
||||||
|
t.Fatal("expected tools to be available on the first LLM call")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(telegramChannel.sentMedia) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia))
|
||||||
|
}
|
||||||
|
if telegramChannel.sentMedia[0].Channel != "telegram" || telegramChannel.sentMedia[0].ChatID != "chat1" {
|
||||||
|
t.Fatalf("unexpected sent media target: %+v", telegramChannel.sentMedia[0])
|
||||||
|
}
|
||||||
|
if len(telegramChannel.sentMedia[0].Parts) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 sent media part, got %d", len(telegramChannel.sentMedia[0].Parts))
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case extra := <-msgBus.OutboundMediaChan():
|
||||||
|
t.Fatalf("expected handled media to bypass async queue, got %+v", extra)
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
defaultAgent := al.GetRegistry().GetDefaultAgent()
|
||||||
|
if defaultAgent == nil {
|
||||||
|
t.Fatal("expected default agent")
|
||||||
|
}
|
||||||
|
route, _, err := al.resolveMessageRoute(bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat1",
|
||||||
|
SenderID: "user1",
|
||||||
|
Content: "take a screenshot of the screen and send it to me",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolveMessageRoute() error = %v", err)
|
||||||
|
}
|
||||||
|
sessionKey := resolveScopeKey(route, "")
|
||||||
|
history := defaultAgent.Sessions.GetHistory(sessionKey)
|
||||||
|
if len(history) == 0 {
|
||||||
|
t.Fatal("expected session history to be saved")
|
||||||
|
}
|
||||||
|
last := history[len(history)-1]
|
||||||
|
if last.Role != "assistant" || last.Content != "Requested output delivered via tool attachment." {
|
||||||
|
t.Fatalf("expected handled assistant summary in history, got %+v", last)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessMessage_HandledToolProcessesQueuedSteeringBeforeReturning(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
ModelName: "test-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &handledMediaWithSteeringProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
al.SetMediaStore(store)
|
||||||
|
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
|
||||||
|
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
|
||||||
|
|
||||||
|
imagePath := filepath.Join(tmpDir, "screen-steering.png")
|
||||||
|
if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil {
|
||||||
|
t.Fatalf("WriteFile(imagePath) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
al.RegisterTool(&handledMediaWithSteeringTool{
|
||||||
|
store: store,
|
||||||
|
path: imagePath,
|
||||||
|
loop: al,
|
||||||
|
})
|
||||||
|
|
||||||
|
response, err := al.processMessage(context.Background(), bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat1",
|
||||||
|
SenderID: "user1",
|
||||||
|
Content: "take a screenshot of the screen and send it to me",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
if response != "Handled the queued steering message." {
|
||||||
|
t.Fatalf("response = %q, want queued steering response", response)
|
||||||
|
}
|
||||||
|
if provider.calls != 2 {
|
||||||
|
t.Fatalf("expected 2 LLM calls after queued steering, got %d", provider.calls)
|
||||||
|
}
|
||||||
|
if len(telegramChannel.sentMedia) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessMessage_MediaArtifactCanBeForwardedBySendFile(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Workspace = tmpDir
|
||||||
|
cfg.Agents.Defaults.ModelName = "test-model"
|
||||||
|
cfg.Agents.Defaults.MaxTokens = 4096
|
||||||
|
cfg.Agents.Defaults.MaxToolIterations = 10
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &artifactThenSendProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
al.SetMediaStore(store)
|
||||||
|
telegramChannel := &fakeMediaChannel{fakeChannel: fakeChannel{id: "rid-telegram"}}
|
||||||
|
al.SetChannelManager(newStartedTestChannelManager(t, msgBus, store, "telegram", telegramChannel))
|
||||||
|
|
||||||
|
mediaDir := media.TempDir()
|
||||||
|
if err := os.MkdirAll(mediaDir, 0o700); err != nil {
|
||||||
|
t.Fatalf("MkdirAll(mediaDir) error = %v", err)
|
||||||
|
}
|
||||||
|
imagePath := filepath.Join(mediaDir, "artifact-screen.png")
|
||||||
|
if err := os.WriteFile(imagePath, []byte("fake screenshot"), 0o644); err != nil {
|
||||||
|
t.Fatalf("WriteFile(imagePath) error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
al.RegisterTool(&mediaArtifactTool{
|
||||||
|
store: store,
|
||||||
|
path: imagePath,
|
||||||
|
})
|
||||||
|
|
||||||
|
response, err := al.processMessage(context.Background(), bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat1",
|
||||||
|
SenderID: "user1",
|
||||||
|
Content: "take a screenshot of the screen and send it to me",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processMessage() error = %v", err)
|
||||||
|
}
|
||||||
|
if response != "" {
|
||||||
|
t.Fatalf("expected no final response after send_file handled delivery, got %q", response)
|
||||||
|
}
|
||||||
|
if provider.calls != 2 {
|
||||||
|
t.Fatalf("expected 2 LLM calls (artifact + send_file), got %d", provider.calls)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(telegramChannel.sentMedia) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 synchronously sent media message, got %d", len(telegramChannel.sentMedia))
|
||||||
|
}
|
||||||
|
if telegramChannel.sentMedia[0].Channel != "telegram" || telegramChannel.sentMedia[0].ChatID != "chat1" {
|
||||||
|
t.Fatalf("unexpected sent media target: %+v", telegramChannel.sentMedia[0])
|
||||||
|
}
|
||||||
|
if len(telegramChannel.sentMedia[0].Parts) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 sent media part, got %d", len(telegramChannel.sentMedia[0].Parts))
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case extra := <-msgBus.OutboundMediaChan():
|
||||||
|
t.Fatalf("expected synchronous send_file delivery to bypass async queue, got %+v", extra)
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestAgentLoop_GetStartupInfo verifies startup info contains tools
|
// TestAgentLoop_GetStartupInfo verifies startup info contains tools
|
||||||
func TestAgentLoop_GetStartupInfo(t *testing.T) {
|
func TestAgentLoop_GetStartupInfo(t *testing.T) {
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
|
@ -600,6 +926,132 @@ func (m *countingMockProvider) GetDefaultModel() string {
|
||||||
return "counting-mock-model"
|
return "counting-mock-model"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type handledMediaProvider struct {
|
||||||
|
calls int
|
||||||
|
toolCounts []int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaProvider) Chat(
|
||||||
|
ctx context.Context,
|
||||||
|
messages []providers.Message,
|
||||||
|
tools []providers.ToolDefinition,
|
||||||
|
model string,
|
||||||
|
opts map[string]any,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
m.calls++
|
||||||
|
m.toolCounts = append(m.toolCounts, len(tools))
|
||||||
|
if m.calls == 1 {
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: "Taking the screenshot now.",
|
||||||
|
ToolCalls: []providers.ToolCall{{
|
||||||
|
ID: "call_handled_media",
|
||||||
|
Type: "function",
|
||||||
|
Name: "handled_media_tool",
|
||||||
|
Arguments: map[string]any{},
|
||||||
|
}},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
return &providers.LLMResponse{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaProvider) GetDefaultModel() string {
|
||||||
|
return "handled-media-model"
|
||||||
|
}
|
||||||
|
|
||||||
|
type artifactThenSendProvider struct {
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *artifactThenSendProvider) 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{
|
||||||
|
Content: "Taking the screenshot now.",
|
||||||
|
ToolCalls: []providers.ToolCall{{
|
||||||
|
ID: "call_artifact_media",
|
||||||
|
Type: "function",
|
||||||
|
Name: "media_artifact_tool",
|
||||||
|
Arguments: map[string]any{},
|
||||||
|
}},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var artifactPath string
|
||||||
|
for i := len(messages) - 1; i >= 0; i-- {
|
||||||
|
if messages[i].Role != "tool" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
start := strings.Index(messages[i].Content, "[file:")
|
||||||
|
if start < 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rest := messages[i].Content[start+len("[file:"):]
|
||||||
|
end := strings.Index(rest, "]")
|
||||||
|
if end < 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
artifactPath = rest[:end]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if artifactPath == "" {
|
||||||
|
return nil, fmt.Errorf("provider did not receive artifact path in tool result")
|
||||||
|
}
|
||||||
|
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: "",
|
||||||
|
ToolCalls: []providers.ToolCall{{
|
||||||
|
ID: "call_send_file",
|
||||||
|
Type: "function",
|
||||||
|
Name: "send_file",
|
||||||
|
Arguments: map[string]any{"path": artifactPath},
|
||||||
|
}},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *artifactThenSendProvider) GetDefaultModel() string {
|
||||||
|
return "artifact-then-send-model"
|
||||||
|
}
|
||||||
|
|
||||||
|
type toolFeedbackProvider struct {
|
||||||
|
filePath string
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *toolFeedbackProvider) 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{
|
||||||
|
ToolCalls: []providers.ToolCall{{
|
||||||
|
ID: "call_heartbeat_read_file",
|
||||||
|
Type: "function",
|
||||||
|
Name: "read_file",
|
||||||
|
Arguments: map[string]any{"path": m.filePath},
|
||||||
|
}},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
Content: "HEARTBEAT_OK",
|
||||||
|
ToolCalls: []providers.ToolCall{},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *toolFeedbackProvider) GetDefaultModel() string {
|
||||||
|
return "heartbeat-tool-feedback-model"
|
||||||
|
}
|
||||||
|
|
||||||
type toolLimitOnlyProvider struct{}
|
type toolLimitOnlyProvider struct{}
|
||||||
|
|
||||||
func (m *toolLimitOnlyProvider) Chat(
|
func (m *toolLimitOnlyProvider) Chat(
|
||||||
|
|
@ -646,6 +1098,135 @@ func (m *mockCustomTool) Execute(ctx context.Context, args map[string]any) *tool
|
||||||
return tools.SilentResult("Custom tool executed")
|
return tools.SilentResult("Custom tool executed")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type handledMediaTool struct {
|
||||||
|
store media.MediaStore
|
||||||
|
path string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaTool) Name() string { return "handled_media_tool" }
|
||||||
|
func (m *handledMediaTool) Description() string {
|
||||||
|
return "Returns a media attachment and fully handles the user response"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
|
||||||
|
ref, err := m.store.Store(m.path, media.MediaMeta{
|
||||||
|
Filename: filepath.Base(m.path),
|
||||||
|
ContentType: "image/png",
|
||||||
|
Source: "test:handled_media_tool",
|
||||||
|
}, "test:handled_media")
|
||||||
|
if err != nil {
|
||||||
|
return tools.ErrorResult(err.Error()).WithError(err)
|
||||||
|
}
|
||||||
|
return tools.MediaResult("Attachment delivered by tool.", []string{ref}).WithResponseHandled()
|
||||||
|
}
|
||||||
|
|
||||||
|
type handledMediaWithSteeringProvider struct {
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaWithSteeringProvider) 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{
|
||||||
|
Content: "Taking the screenshot now.",
|
||||||
|
ToolCalls: []providers.ToolCall{{
|
||||||
|
ID: "call_handled_media_steering",
|
||||||
|
Type: "function",
|
||||||
|
Name: "handled_media_with_steering_tool",
|
||||||
|
Arguments: map[string]any{},
|
||||||
|
}},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, msg := range messages {
|
||||||
|
if msg.Role == "user" && msg.Content == "what about this instead?" {
|
||||||
|
return &providers.LLMResponse{Content: "Handled the queued steering message."}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("provider did not receive queued steering message")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaWithSteeringProvider) GetDefaultModel() string {
|
||||||
|
return "handled-media-with-steering-model"
|
||||||
|
}
|
||||||
|
|
||||||
|
type handledMediaWithSteeringTool struct {
|
||||||
|
store media.MediaStore
|
||||||
|
path string
|
||||||
|
loop *AgentLoop
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaWithSteeringTool) Name() string { return "handled_media_with_steering_tool" }
|
||||||
|
func (m *handledMediaWithSteeringTool) Description() string {
|
||||||
|
return "Returns handled media and enqueues a steering message during execution"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaWithSteeringTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *handledMediaWithSteeringTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
|
||||||
|
if err := m.loop.Steer(providers.Message{Role: "user", Content: "what about this instead?"}); err != nil {
|
||||||
|
return tools.ErrorResult(err.Error()).WithError(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ref, err := m.store.Store(m.path, media.MediaMeta{
|
||||||
|
Filename: filepath.Base(m.path),
|
||||||
|
ContentType: "image/png",
|
||||||
|
Source: "test:handled_media_with_steering_tool",
|
||||||
|
}, "test:handled_media_with_steering")
|
||||||
|
if err != nil {
|
||||||
|
return tools.ErrorResult(err.Error()).WithError(err)
|
||||||
|
}
|
||||||
|
return tools.MediaResult("Attachment delivered by tool.", []string{ref}).WithResponseHandled()
|
||||||
|
}
|
||||||
|
|
||||||
|
type mediaArtifactTool struct {
|
||||||
|
store media.MediaStore
|
||||||
|
path string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mediaArtifactTool) Name() string { return "media_artifact_tool" }
|
||||||
|
func (m *mediaArtifactTool) Description() string {
|
||||||
|
return "Returns a media artifact that the agent can forward or save later"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mediaArtifactTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mediaArtifactTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
|
||||||
|
ref, err := m.store.Store(m.path, media.MediaMeta{
|
||||||
|
Filename: filepath.Base(m.path),
|
||||||
|
ContentType: "image/png",
|
||||||
|
Source: "test:media_artifact_tool",
|
||||||
|
}, "test:media_artifact")
|
||||||
|
if err != nil {
|
||||||
|
return tools.ErrorResult(err.Error()).WithError(err)
|
||||||
|
}
|
||||||
|
return tools.MediaResult("Artifact created.", []string{ref})
|
||||||
|
}
|
||||||
|
|
||||||
type toolLimitTestTool struct{}
|
type toolLimitTestTool struct{}
|
||||||
|
|
||||||
func (m *toolLimitTestTool) Name() string {
|
func (m *toolLimitTestTool) Name() string {
|
||||||
|
|
@ -1495,18 +2076,17 @@ func TestTargetReasoningChannelID_AllChannels(t *testing.T) {
|
||||||
t.Fatalf("Failed to create channel manager: %v", err)
|
t.Fatalf("Failed to create channel manager: %v", err)
|
||||||
}
|
}
|
||||||
for name, id := range map[string]string{
|
for name, id := range map[string]string{
|
||||||
"whatsapp": "rid-whatsapp",
|
"whatsapp": "rid-whatsapp",
|
||||||
"telegram": "rid-telegram",
|
"telegram": "rid-telegram",
|
||||||
"feishu": "rid-feishu",
|
"feishu": "rid-feishu",
|
||||||
"discord": "rid-discord",
|
"discord": "rid-discord",
|
||||||
"maixcam": "rid-maixcam",
|
"maixcam": "rid-maixcam",
|
||||||
"qq": "rid-qq",
|
"qq": "rid-qq",
|
||||||
"dingtalk": "rid-dingtalk",
|
"dingtalk": "rid-dingtalk",
|
||||||
"slack": "rid-slack",
|
"slack": "rid-slack",
|
||||||
"line": "rid-line",
|
"line": "rid-line",
|
||||||
"onebot": "rid-onebot",
|
"onebot": "rid-onebot",
|
||||||
"wecom": "rid-wecom",
|
"wecom": "rid-wecom",
|
||||||
"wecom_app": "rid-wecom-app",
|
|
||||||
} {
|
} {
|
||||||
chManager.RegisterChannel(name, &fakeChannel{id: id})
|
chManager.RegisterChannel(name, &fakeChannel{id: id})
|
||||||
}
|
}
|
||||||
|
|
@ -1526,7 +2106,6 @@ func TestTargetReasoningChannelID_AllChannels(t *testing.T) {
|
||||||
{channel: "line", wantID: "rid-line"},
|
{channel: "line", wantID: "rid-line"},
|
||||||
{channel: "onebot", wantID: "rid-onebot"},
|
{channel: "onebot", wantID: "rid-onebot"},
|
||||||
{channel: "wecom", wantID: "rid-wecom"},
|
{channel: "wecom", wantID: "rid-wecom"},
|
||||||
{channel: "wecom_app", wantID: "rid-wecom-app"},
|
|
||||||
{channel: "unknown", wantID: ""},
|
{channel: "unknown", wantID: ""},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1768,6 +2347,112 @@ func TestProcessMessage_PublishesReasoningContentToReasoningChannel(t *testing.T
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProcessHeartbeat_DoesNotPublishToolFeedback(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
heartbeatFile := filepath.Join(tmpDir, "heartbeat-task.txt")
|
||||||
|
if err := os.WriteFile(heartbeatFile, []byte("heartbeat 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 := &toolFeedbackProvider{filePath: heartbeatFile}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
response, err := al.ProcessHeartbeat(context.Background(), "check heartbeat tasks", "telegram", "chat-1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessHeartbeat() error = %v", err)
|
||||||
|
}
|
||||||
|
if response != "HEARTBEAT_OK" {
|
||||||
|
t.Fatalf("ProcessHeartbeat() response = %q, want %q", response, "HEARTBEAT_OK")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case outbound := <-msgBus.OutboundChan():
|
||||||
|
t.Fatalf("expected no outbound tool feedback during heartbeat, got %+v", outbound)
|
||||||
|
case <-time.After(200 * time.Millisecond):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
heartbeatFile := filepath.Join(tmpDir, "tool-feedback.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 := &toolFeedbackProvider{filePath: heartbeatFile}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
|
response, err := al.processMessage(context.Background(), bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
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():
|
||||||
|
if outbound.Channel != "telegram" {
|
||||||
|
t.Fatalf("tool feedback channel = %q, want %q", outbound.Channel, "telegram")
|
||||||
|
}
|
||||||
|
if outbound.ChatID != "chat-1" {
|
||||||
|
t.Fatalf("tool feedback chatID = %q, want %q", outbound.ChatID, "chat-1")
|
||||||
|
}
|
||||||
|
if !strings.Contains(outbound.Content, "`read_file`") {
|
||||||
|
t.Fatalf("tool feedback content = %q, want read_file preview", outbound.Content)
|
||||||
|
}
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("expected outbound tool feedback for regular messages")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestResolveMediaRefs_ResolvesToBase64(t *testing.T) {
|
func TestResolveMediaRefs_ResolvesToBase64(t *testing.T) {
|
||||||
store := media.NewFileMediaStore()
|
store := media.NewFileMediaStore()
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
|
|
|
||||||
|
|
@ -1255,8 +1255,7 @@ make test # Full test suite
|
||||||
| `pkg/channels/onebot/` | `"onebot"` | ReactionCapable, MediaSender |
|
| `pkg/channels/onebot/` | `"onebot"` | ReactionCapable, MediaSender |
|
||||||
| `pkg/channels/dingtalk/` | `"dingtalk"` | — |
|
| `pkg/channels/dingtalk/` | `"dingtalk"` | — |
|
||||||
| `pkg/channels/feishu/` | `"feishu"` | — (architecture-specific build tags: `feishu_32.go` / `feishu_64.go`) |
|
| `pkg/channels/feishu/` | `"feishu"` | — (architecture-specific build tags: `feishu_32.go` / `feishu_64.go`) |
|
||||||
| `pkg/channels/wecom/` | `"wecom"` | WebhookHandler, HealthChecker |
|
| `pkg/channels/wecom/` | `"wecom"` | MediaSender |
|
||||||
| `pkg/channels/wecom/` | `"wecom_app"` | MediaSender, WebhookHandler, HealthChecker |
|
|
||||||
| `pkg/channels/qq/` | `"qq"` | — |
|
| `pkg/channels/qq/` | `"qq"` | — |
|
||||||
| `pkg/channels/whatsapp/` | `"whatsapp"` | — (Bridge mode) |
|
| `pkg/channels/whatsapp/` | `"whatsapp"` | — (Bridge mode) |
|
||||||
| `pkg/channels/whatsapp_native/` | `"whatsapp_native"` | — (Native whatsmeow mode) |
|
| `pkg/channels/whatsapp_native/` | `"whatsapp_native"` | — (Native whatsmeow mode) |
|
||||||
|
|
@ -1371,7 +1370,7 @@ agentLoop.Stop() // Stop Agent
|
||||||
|
|
||||||
2. **Feishu architecture-specific compilation**: The Feishu channel uses build tags to distinguish 32-bit and 64-bit architectures (`feishu_32.go` / `feishu_64.go`). Feishu uses the SDK's WebSocket mode (not HTTP webhook), so it does not implement `WebhookHandler`.
|
2. **Feishu architecture-specific compilation**: The Feishu channel uses build tags to distinguish 32-bit and 64-bit architectures (`feishu_32.go` / `feishu_64.go`). Feishu uses the SDK's WebSocket mode (not HTTP webhook), so it does not implement `WebhookHandler`.
|
||||||
|
|
||||||
3. **WeCom has two factories**: `"wecom"` (Bot mode, webhook only) and `"wecom_app"` (App mode, supports MediaSender) are registered separately. Both implement `WebhookHandler` and `HealthChecker`.
|
3. **WeCom is now a single channel**: `"wecom"` is implemented as a WebSocket-based AI Bot channel with route persistence. Access control uses the shared channel allowlist mechanism. It no longer exposes the legacy webhook/app split.
|
||||||
|
|
||||||
4. **Pico Protocol**: `pkg/channels/pico/` implements a custom PicoClaw native protocol channel that receives messages via WebSocket webhook (`/pico/ws`).
|
4. **Pico Protocol**: `pkg/channels/pico/` implements a custom PicoClaw native protocol channel that receives messages via WebSocket webhook (`/pico/ws`).
|
||||||
|
|
||||||
|
|
@ -1381,4 +1380,4 @@ agentLoop.Stop() // Stop Agent
|
||||||
|
|
||||||
7. **PlaceholderConfig vs implementation**: `PlaceholderConfig` appears in 6 channel configs (Telegram, Discord, Slack, LINE, OneBot, Pico), but only channels that implement both `PlaceholderCapable` + `MessageEditor` (Telegram, Discord, Pico) can actually use placeholder message editing. The rest are reserved fields.
|
7. **PlaceholderConfig vs implementation**: `PlaceholderConfig` appears in 6 channel configs (Telegram, Discord, Slack, LINE, OneBot, Pico), but only channels that implement both `PlaceholderCapable` + `MessageEditor` (Telegram, Discord, Pico) can actually use placeholder message editing. The rest are reserved fields.
|
||||||
|
|
||||||
8. **ReasoningChannelID**: Most channel configs include a `reasoning_channel_id` field to route LLM reasoning/thinking output to a designated channel (WhatsApp, Telegram, Feishu, Discord, MaixCam, QQ, DingTalk, Slack, LINE, OneBot, WeCom, WeComApp). Note: `PicoConfig` does not currently expose this field. `BaseChannel` exposes this via the `WithReasoningChannelID` option and `ReasoningChannelID()` method.
|
8. **ReasoningChannelID**: Most channel configs include a `reasoning_channel_id` field to route LLM reasoning/thinking output to a designated channel (WhatsApp, Telegram, Feishu, Discord, MaixCam, QQ, DingTalk, Slack, LINE, OneBot, WeCom). Note: `PicoConfig` does not currently expose this field. `BaseChannel` exposes this via the `WithReasoningChannelID` option and `ReasoningChannelID()` method.
|
||||||
|
|
|
||||||
|
|
@ -1254,8 +1254,7 @@ make test # 全量测试
|
||||||
| `pkg/channels/onebot/` | `"onebot"` | ReactionCapable, MediaSender |
|
| `pkg/channels/onebot/` | `"onebot"` | ReactionCapable, MediaSender |
|
||||||
| `pkg/channels/dingtalk/` | `"dingtalk"` | — |
|
| `pkg/channels/dingtalk/` | `"dingtalk"` | — |
|
||||||
| `pkg/channels/feishu/` | `"feishu"` | — (架构特定 build tags: `feishu_32.go` / `feishu_64.go`) |
|
| `pkg/channels/feishu/` | `"feishu"` | — (架构特定 build tags: `feishu_32.go` / `feishu_64.go`) |
|
||||||
| `pkg/channels/wecom/` | `"wecom"` | WebhookHandler, HealthChecker |
|
| `pkg/channels/wecom/` | `"wecom"` | MediaSender |
|
||||||
| `pkg/channels/wecom/` | `"wecom_app"` | MediaSender, WebhookHandler, HealthChecker |
|
|
||||||
| `pkg/channels/qq/` | `"qq"` | — |
|
| `pkg/channels/qq/` | `"qq"` | — |
|
||||||
| `pkg/channels/whatsapp/` | `"whatsapp"` | — (Bridge 模式) |
|
| `pkg/channels/whatsapp/` | `"whatsapp"` | — (Bridge 模式) |
|
||||||
| `pkg/channels/whatsapp_native/` | `"whatsapp_native"` | — (原生 whatsmeow 模式) |
|
| `pkg/channels/whatsapp_native/` | `"whatsapp_native"` | — (原生 whatsmeow 模式) |
|
||||||
|
|
@ -1370,7 +1369,7 @@ agentLoop.Stop() // 停止 Agent
|
||||||
|
|
||||||
2. **Feishu 架构特定编译**:Feishu channel 使用 build tags 区分 32 位和 64 位架构(`feishu_32.go` / `feishu_64.go`)。Feishu 使用 SDK 的 WebSocket 模式(非 HTTP webhook),因此不实现 `WebhookHandler`。
|
2. **Feishu 架构特定编译**:Feishu channel 使用 build tags 区分 32 位和 64 位架构(`feishu_32.go` / `feishu_64.go`)。Feishu 使用 SDK 的 WebSocket 模式(非 HTTP webhook),因此不实现 `WebhookHandler`。
|
||||||
|
|
||||||
3. **WeCom 有两个工厂**:`"wecom"`(Bot 模式,纯 webhook)和 `"wecom_app"`(应用模式,支持 MediaSender)分别注册。两者都实现了 `WebhookHandler` 和 `HealthChecker`。
|
3. **WeCom 现在只有一个 channel**:`"wecom"` 采用 WebSocket AI Bot 实现,带路由持久化;访问控制走统一的 channel 白名单机制,不再保留旧的 webhook/app 双分支。
|
||||||
|
|
||||||
4. **Pico Protocol**:`pkg/channels/pico/` 实现了一个自定义的 PicoClaw 原生协议 channel,通过 WebSocket webhook (`/pico/ws`) 接收消息。
|
4. **Pico Protocol**:`pkg/channels/pico/` 实现了一个自定义的 PicoClaw 原生协议 channel,通过 WebSocket webhook (`/pico/ws`) 接收消息。
|
||||||
|
|
||||||
|
|
@ -1380,4 +1379,4 @@ agentLoop.Stop() // 停止 Agent
|
||||||
|
|
||||||
7. **PlaceholderConfig 的配置与实现**:`PlaceholderConfig` 出现在 6 个 channel config 中(Telegram、Discord、Slack、LINE、OneBot、Pico),但只有实现了 `PlaceholderCapable` + `MessageEditor` 的 channel(Telegram、Discord、Pico)能真正使用占位消息编辑功能。其余 channel 的 `PlaceholderConfig` 为预留字段。
|
7. **PlaceholderConfig 的配置与实现**:`PlaceholderConfig` 出现在 6 个 channel config 中(Telegram、Discord、Slack、LINE、OneBot、Pico),但只有实现了 `PlaceholderCapable` + `MessageEditor` 的 channel(Telegram、Discord、Pico)能真正使用占位消息编辑功能。其余 channel 的 `PlaceholderConfig` 为预留字段。
|
||||||
|
|
||||||
8. **ReasoningChannelID**:大多数 channel config 都包含 `reasoning_channel_id` 字段,用于将 LLM 的思维链(reasoning/thinking)路由到指定 channel(WhatsApp、Telegram、Feishu、Discord、MaixCam、QQ、DingTalk、Slack、LINE、OneBot、WeCom、WeComApp)。注意:`PicoConfig` 目前不包含该字段。`BaseChannel` 通过 `WithReasoningChannelID` 选项和 `ReasoningChannelID()` 方法暴露此配置。
|
8. **ReasoningChannelID**:大多数 channel config 都包含 `reasoning_channel_id` 字段,用于将 LLM 的思维链(reasoning/thinking)路由到指定 channel(WhatsApp、Telegram、Feishu、Discord、MaixCam、QQ、DingTalk、Slack、LINE、OneBot、WeCom)。注意:`PicoConfig` 目前不包含该字段。`BaseChannel` 通过 `WithReasoningChannelID` 选项和 `ReasoningChannelID()` 方法暴露此配置。
|
||||||
|
|
|
||||||
|
|
@ -206,6 +206,40 @@ func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMess
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// preSendMedia handles typing stop, reaction undo, and placeholder cleanup
|
||||||
|
// before sending media attachments. Unlike preSend for text messages, media
|
||||||
|
// delivery never edits the placeholder because there is no text payload to
|
||||||
|
// replace it with; it only attempts to delete the placeholder when possible.
|
||||||
|
func (m *Manager) preSendMedia(ctx context.Context, name string, msg bus.OutboundMediaMessage, ch Channel) {
|
||||||
|
key := name + ":" + msg.ChatID
|
||||||
|
|
||||||
|
// 1. Stop typing
|
||||||
|
if v, loaded := m.typingStops.LoadAndDelete(key); loaded {
|
||||||
|
if entry, ok := v.(typingEntry); ok {
|
||||||
|
entry.stop() // idempotent, safe
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Undo reaction
|
||||||
|
if v, loaded := m.reactionUndos.LoadAndDelete(key); loaded {
|
||||||
|
if entry, ok := v.(reactionEntry); ok {
|
||||||
|
entry.undo() // idempotent, safe
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. Clear any finalized stream marker for this chat before media delivery.
|
||||||
|
m.streamActive.LoadAndDelete(key)
|
||||||
|
|
||||||
|
// 4. Delete placeholder if present.
|
||||||
|
if v, loaded := m.placeholders.LoadAndDelete(key); loaded {
|
||||||
|
if entry, ok := v.(placeholderEntry); ok && entry.id != "" {
|
||||||
|
if deleter, ok := ch.(MessageDeleter); ok {
|
||||||
|
deleter.DeleteMessage(ctx, msg.ChatID, entry.id) // best effort
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func NewManager(cfg *config.Config, messageBus *bus.MessageBus, store media.MediaStore) (*Manager, error) {
|
func NewManager(cfg *config.Config, messageBus *bus.MessageBus, store media.MediaStore) (*Manager, error) {
|
||||||
m := &Manager{
|
m := &Manager{
|
||||||
channels: make(map[string]Channel),
|
channels: make(map[string]Channel),
|
||||||
|
|
@ -371,19 +405,10 @@ func (m *Manager) initChannels(channels *config.ChannelsConfig) error {
|
||||||
m.initChannel("onebot", "OneBot")
|
m.initChannel("onebot", "OneBot")
|
||||||
}
|
}
|
||||||
|
|
||||||
if channels.WeCom.Enabled && channels.WeCom.Token() != "" {
|
if channels.WeCom.Enabled && channels.WeCom.BotID != "" && channels.WeCom.Secret() != "" {
|
||||||
m.initChannel("wecom", "WeCom")
|
m.initChannel("wecom", "WeCom")
|
||||||
}
|
}
|
||||||
|
|
||||||
if channels.WeComAIBot.Enabled && (channels.WeComAIBot.Token() != "" ||
|
|
||||||
(channels.WeComAIBot.Secret() != "" && channels.WeComAIBot.BotID != "")) {
|
|
||||||
m.initChannel("wecom_aibot", "WeCom AI Bot")
|
|
||||||
}
|
|
||||||
|
|
||||||
if channels.WeComApp.Enabled && channels.WeComApp.CorpID != "" {
|
|
||||||
m.initChannel("wecom_app", "WeCom App")
|
|
||||||
}
|
|
||||||
|
|
||||||
if channels.Weixin.Enabled && channels.Weixin.Token() != "" {
|
if channels.Weixin.Enabled && channels.Weixin.Token() != "" {
|
||||||
m.initChannel("weixin", "Weixin")
|
m.initChannel("weixin", "Weixin")
|
||||||
}
|
}
|
||||||
|
|
@ -774,7 +799,7 @@ func (m *Manager) runMediaWorker(ctx context.Context, name string, w *channelWor
|
||||||
if !ok {
|
if !ok {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
m.sendMediaWithRetry(ctx, name, w, msg)
|
_ = m.sendMediaWithRetry(ctx, name, w, msg)
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
@ -782,26 +807,37 @@ func (m *Manager) runMediaWorker(ctx context.Context, name string, w *channelWor
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendMediaWithRetry sends a media message through the channel with rate limiting and
|
// sendMediaWithRetry sends a media message through the channel with rate limiting and
|
||||||
// retry logic. If the channel does not implement MediaSender, it silently skips.
|
// retry logic. It returns nil on success, or the last error after retries,
|
||||||
func (m *Manager) sendMediaWithRetry(ctx context.Context, name string, w *channelWorker, msg bus.OutboundMediaMessage) {
|
// including when the channel does not support MediaSender.
|
||||||
|
func (m *Manager) sendMediaWithRetry(
|
||||||
|
ctx context.Context,
|
||||||
|
name string,
|
||||||
|
w *channelWorker,
|
||||||
|
msg bus.OutboundMediaMessage,
|
||||||
|
) error {
|
||||||
ms, ok := w.ch.(MediaSender)
|
ms, ok := w.ch.(MediaSender)
|
||||||
if !ok {
|
if !ok {
|
||||||
logger.DebugCF("channels", "Channel does not support MediaSender, skipping media", map[string]any{
|
err := fmt.Errorf("channel %q does not support media sending", name)
|
||||||
|
logger.WarnCF("channels", "Channel does not support MediaSender", map[string]any{
|
||||||
"channel": name,
|
"channel": name,
|
||||||
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Rate limit: wait for token
|
// Rate limit: wait for token
|
||||||
if err := w.limiter.Wait(ctx); err != nil {
|
if err := w.limiter.Wait(ctx); err != nil {
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Pre-send: stop typing and clean up any placeholder before sending media.
|
||||||
|
m.preSendMedia(ctx, name, msg, w.ch)
|
||||||
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
for attempt := 0; attempt <= maxRetries; attempt++ {
|
for attempt := 0; attempt <= maxRetries; attempt++ {
|
||||||
lastErr = ms.SendMedia(ctx, msg)
|
lastErr = ms.SendMedia(ctx, msg)
|
||||||
if lastErr == nil {
|
if lastErr == nil {
|
||||||
return
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Permanent failures — don't retry
|
// Permanent failures — don't retry
|
||||||
|
|
@ -820,7 +856,7 @@ func (m *Manager) sendMediaWithRetry(ctx context.Context, name string, w *channe
|
||||||
case <-time.After(rateLimitDelay):
|
case <-time.After(rateLimitDelay):
|
||||||
continue
|
continue
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return ctx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -829,7 +865,7 @@ func (m *Manager) sendMediaWithRetry(ctx context.Context, name string, w *channe
|
||||||
select {
|
select {
|
||||||
case <-time.After(backoff):
|
case <-time.After(backoff):
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return ctx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -840,6 +876,7 @@ func (m *Manager) sendMediaWithRetry(ctx context.Context, name string, w *channe
|
||||||
"error": lastErr.Error(),
|
"error": lastErr.Error(),
|
||||||
"retries": maxRetries,
|
"retries": maxRetries,
|
||||||
})
|
})
|
||||||
|
return lastErr
|
||||||
}
|
}
|
||||||
|
|
||||||
// runTTLJanitor periodically scans the typingStops and placeholders maps
|
// runTTLJanitor periodically scans the typingStops and placeholders maps
|
||||||
|
|
@ -1032,6 +1069,26 @@ func (m *Manager) SendMessage(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SendMedia sends outbound media synchronously through the channel worker's
|
||||||
|
// rate limiter and retry logic. It blocks until the media is delivered (or all
|
||||||
|
// retries are exhausted), which preserves ordering when later agent behavior
|
||||||
|
// depends on actual media delivery.
|
||||||
|
func (m *Manager) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
|
m.mu.RLock()
|
||||||
|
_, exists := m.channels[msg.Channel]
|
||||||
|
w, wExists := m.workers[msg.Channel]
|
||||||
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
if !exists {
|
||||||
|
return fmt.Errorf("channel %s not found", msg.Channel)
|
||||||
|
}
|
||||||
|
if !wExists || w == nil {
|
||||||
|
return fmt.Errorf("channel %s has no active worker", msg.Channel)
|
||||||
|
}
|
||||||
|
|
||||||
|
return m.sendMediaWithRetry(ctx, msg.Channel, w, msg)
|
||||||
|
}
|
||||||
|
|
||||||
func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, content string) error {
|
func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, content string) error {
|
||||||
m.mu.RLock()
|
m.mu.RLock()
|
||||||
_, exists := m.channels[channelName]
|
_, exists := m.channels[channelName]
|
||||||
|
|
|
||||||
|
|
@ -49,15 +49,7 @@ func hiddenValues(key string, value map[string]any, ch config.ChannelsConfig) {
|
||||||
value["token"] = ch.LINE.ChannelAccessToken()
|
value["token"] = ch.LINE.ChannelAccessToken()
|
||||||
value["secret"] = ch.LINE.ChannelSecret()
|
value["secret"] = ch.LINE.ChannelSecret()
|
||||||
case "wecom":
|
case "wecom":
|
||||||
value["token"] = ch.WeCom.Token()
|
value["secret"] = ch.WeCom.Secret()
|
||||||
value["key"] = ch.WeCom.EncodingAESKey()
|
|
||||||
case "wecom_app":
|
|
||||||
value["token"] = ch.WeComApp.Token()
|
|
||||||
value["secret"] = ch.WeComApp.CorpSecret()
|
|
||||||
case "wecom_aibot":
|
|
||||||
value["token"] = ch.WeComAIBot.Token()
|
|
||||||
value["key"] = ch.WeComAIBot.EncodingAESKey()
|
|
||||||
value["secret"] = ch.WeComAIBot.Secret()
|
|
||||||
case "dingtalk":
|
case "dingtalk":
|
||||||
value["secret"] = ch.QQ.AppSecret()
|
value["secret"] = ch.QQ.AppSecret()
|
||||||
case "qq":
|
case "qq":
|
||||||
|
|
@ -156,16 +148,7 @@ func updateKeys(newcfg, old *config.ChannelsConfig) {
|
||||||
newcfg.LINE.SetChannelSecret(old.LINE.ChannelSecret())
|
newcfg.LINE.SetChannelSecret(old.LINE.ChannelSecret())
|
||||||
}
|
}
|
||||||
if newcfg.WeCom.Enabled {
|
if newcfg.WeCom.Enabled {
|
||||||
newcfg.WeCom.SetToken(old.WeCom.Token())
|
newcfg.WeCom.SetSecret(old.WeCom.Secret())
|
||||||
newcfg.WeCom.SetEncodingAESKey(old.WeCom.EncodingAESKey())
|
|
||||||
}
|
|
||||||
if newcfg.WeComApp.Enabled {
|
|
||||||
newcfg.WeComApp.SetToken(old.WeComApp.Token())
|
|
||||||
newcfg.WeComApp.SetCorpSecret(old.WeComApp.CorpSecret())
|
|
||||||
}
|
|
||||||
if newcfg.WeComAIBot.Enabled {
|
|
||||||
newcfg.WeComAIBot.SetToken(old.WeComAIBot.Token())
|
|
||||||
newcfg.WeComAIBot.SetEncodingAESKey(old.WeComAIBot.EncodingAESKey())
|
|
||||||
}
|
}
|
||||||
if newcfg.DingTalk.Enabled {
|
if newcfg.DingTalk.Enabled {
|
||||||
newcfg.DingTalk.SetClientSecret(old.DingTalk.ClientSecret())
|
newcfg.DingTalk.SetClientSecret(old.DingTalk.ClientSecret())
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -43,6 +44,40 @@ func (m *mockChannel) EditMessage(ctx context.Context, chatID, messageID, conten
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type mockMediaChannel struct {
|
||||||
|
mockChannel
|
||||||
|
sendMediaFn func(ctx context.Context, msg bus.OutboundMediaMessage) error
|
||||||
|
sentMediaMessages []bus.OutboundMediaMessage
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
|
m.sentMediaMessages = append(m.sentMediaMessages, msg)
|
||||||
|
if m.sendMediaFn != nil {
|
||||||
|
return m.sendMediaFn(ctx, msg)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type mockDeletingMediaChannel struct {
|
||||||
|
mockMediaChannel
|
||||||
|
deleteCalls int
|
||||||
|
lastDeleted struct {
|
||||||
|
chatID string
|
||||||
|
messageID string
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockDeletingMediaChannel) DeleteMessage(
|
||||||
|
_ context.Context,
|
||||||
|
chatID string,
|
||||||
|
messageID string,
|
||||||
|
) error {
|
||||||
|
m.deleteCalls++
|
||||||
|
m.lastDeleted.chatID = chatID
|
||||||
|
m.lastDeleted.messageID = messageID
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// newTestManager creates a minimal Manager suitable for unit tests.
|
// newTestManager creates a minimal Manager suitable for unit tests.
|
||||||
func newTestManager() *Manager {
|
func newTestManager() *Manager {
|
||||||
return &Manager{
|
return &Manager{
|
||||||
|
|
@ -208,6 +243,125 @@ func TestSendWithRetry_MaxRetriesExhausted(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_Success(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
var callCount int
|
||||||
|
ch := &mockMediaChannel{
|
||||||
|
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) error {
|
||||||
|
callCount++
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
err := m.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Parts: []bus.MediaPart{{Ref: "media://abc"}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendMedia() error = %v", err)
|
||||||
|
}
|
||||||
|
if callCount != 1 {
|
||||||
|
t.Fatalf("expected 1 SendMedia call, got %d", callCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_PropagatesFailure(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
ch := &mockMediaChannel{
|
||||||
|
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) error {
|
||||||
|
return fmt.Errorf("bad upload: %w", ErrSendFailed)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
err := m.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Parts: []bus.MediaPart{{Ref: "media://abc"}},
|
||||||
|
})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected SendMedia to return error")
|
||||||
|
}
|
||||||
|
if !errors.Is(err, ErrSendFailed) {
|
||||||
|
t.Fatalf("expected ErrSendFailed, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_UnsupportedChannelReturnsError(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
ch := &mockChannel{
|
||||||
|
sendFn: func(_ context.Context, _ bus.OutboundMessage) error {
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
|
||||||
|
err := m.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Parts: []bus.MediaPart{{Ref: "media://abc"}},
|
||||||
|
})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected SendMedia to return error for unsupported channel")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "does not support media sending") {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_DeletesPlaceholderBeforeSending(t *testing.T) {
|
||||||
|
m := newTestManager()
|
||||||
|
ch := &mockDeletingMediaChannel{
|
||||||
|
mockMediaChannel: mockMediaChannel{
|
||||||
|
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) error {
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w := &channelWorker{
|
||||||
|
ch: ch,
|
||||||
|
limiter: rate.NewLimiter(rate.Inf, 1),
|
||||||
|
}
|
||||||
|
m.channels["test"] = ch
|
||||||
|
m.workers["test"] = w
|
||||||
|
m.RecordPlaceholder("test", "chat1", "placeholder-1")
|
||||||
|
|
||||||
|
err := m.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "test",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Parts: []bus.MediaPart{{Ref: "media://abc"}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendMedia() error = %v", err)
|
||||||
|
}
|
||||||
|
if ch.deleteCalls != 1 {
|
||||||
|
t.Fatalf("expected placeholder delete to be called once, got %d", ch.deleteCalls)
|
||||||
|
}
|
||||||
|
if ch.lastDeleted.chatID != "chat1" || ch.lastDeleted.messageID != "placeholder-1" {
|
||||||
|
t.Fatalf("unexpected placeholder deletion target: %+v", ch.lastDeleted)
|
||||||
|
}
|
||||||
|
if len(ch.sentMediaMessages) != 1 {
|
||||||
|
t.Fatalf("expected media to be sent once, got %d", len(ch.sentMediaMessages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestSendWithRetry_UnknownError(t *testing.T) {
|
func TestSendWithRetry_UnknownError(t *testing.T) {
|
||||||
m := newTestManager()
|
m := newTestManager()
|
||||||
var callCount int
|
var callCount int
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package matrix
|
package matrix
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
|
@ -8,6 +10,11 @@ import (
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
channels.RegisterFactory("matrix", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
channels.RegisterFactory("matrix", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||||
return NewMatrixChannel(cfg.Channels.Matrix, b)
|
matrixCfg := cfg.Channels.Matrix
|
||||||
|
cryptoDatabasePath := matrixCfg.CryptoDatabasePath
|
||||||
|
if cryptoDatabasePath == "" {
|
||||||
|
cryptoDatabasePath = filepath.Join(cfg.WorkspacePath(), "matrix")
|
||||||
|
}
|
||||||
|
return NewMatrixChannel(matrixCfg, b, cryptoDatabasePath)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package matrix
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"html"
|
"html"
|
||||||
"io"
|
"io"
|
||||||
|
|
@ -17,9 +18,12 @@ import (
|
||||||
"github.com/gomarkdown/markdown"
|
"github.com/gomarkdown/markdown"
|
||||||
mdhtml "github.com/gomarkdown/markdown/html"
|
mdhtml "github.com/gomarkdown/markdown/html"
|
||||||
"github.com/gomarkdown/markdown/parser"
|
"github.com/gomarkdown/markdown/parser"
|
||||||
|
"go.mau.fi/util/dbutil"
|
||||||
"maunium.net/go/mautrix"
|
"maunium.net/go/mautrix"
|
||||||
|
"maunium.net/go/mautrix/crypto/cryptohelper"
|
||||||
"maunium.net/go/mautrix/event"
|
"maunium.net/go/mautrix/event"
|
||||||
"maunium.net/go/mautrix/id"
|
"maunium.net/go/mautrix/id"
|
||||||
|
_ "modernc.org/sqlite"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
|
@ -30,6 +34,9 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
sqliteDriver = "sqlite"
|
||||||
|
dbName = "store.db"
|
||||||
|
|
||||||
typingRefreshInterval = 20 * time.Second
|
typingRefreshInterval = 20 * time.Second
|
||||||
typingServerTTL = 30 * time.Second
|
typingServerTTL = 30 * time.Second
|
||||||
roomKindCacheTTL = 5 * time.Minute
|
roomKindCacheTTL = 5 * time.Minute
|
||||||
|
|
@ -181,9 +188,16 @@ type MatrixChannel struct {
|
||||||
|
|
||||||
roomKindCache *roomKindCache
|
roomKindCache *roomKindCache
|
||||||
localpartMentionR *regexp.Regexp
|
localpartMentionR *regexp.Regexp
|
||||||
|
|
||||||
|
cryptoHelper *cryptohelper.CryptoHelper
|
||||||
|
cryptoDbPath string
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMatrixChannel(cfg config.MatrixConfig, messageBus *bus.MessageBus) (*MatrixChannel, error) {
|
func NewMatrixChannel(
|
||||||
|
cfg config.MatrixConfig,
|
||||||
|
messageBus *bus.MessageBus,
|
||||||
|
cryptoDatabasePath string,
|
||||||
|
) (*MatrixChannel, error) {
|
||||||
homeserver := strings.TrimSpace(cfg.Homeserver)
|
homeserver := strings.TrimSpace(cfg.Homeserver)
|
||||||
userID := strings.TrimSpace(cfg.UserID)
|
userID := strings.TrimSpace(cfg.UserID)
|
||||||
accessToken := strings.TrimSpace(cfg.AccessToken())
|
accessToken := strings.TrimSpace(cfg.AccessToken())
|
||||||
|
|
@ -230,6 +244,7 @@ func NewMatrixChannel(cfg config.MatrixConfig, messageBus *bus.MessageBus) (*Mat
|
||||||
roomKindCache: newRoomKindCache(roomKindCacheMaxEntries, roomKindCacheTTL),
|
roomKindCache: newRoomKindCache(roomKindCacheMaxEntries, roomKindCacheTTL),
|
||||||
localpartMentionR: localpartMentionRegexp(matrixLocalpart(client.UserID)),
|
localpartMentionR: localpartMentionRegexp(matrixLocalpart(client.UserID)),
|
||||||
typingMu: sync.Mutex{},
|
typingMu: sync.Mutex{},
|
||||||
|
cryptoDbPath: cryptoDatabasePath,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -239,7 +254,21 @@ func (c *MatrixChannel) Start(ctx context.Context) error {
|
||||||
c.ctx, c.cancel = context.WithCancel(ctx)
|
c.ctx, c.cancel = context.WithCancel(ctx)
|
||||||
c.startTime = time.Now()
|
c.startTime = time.Now()
|
||||||
|
|
||||||
|
// Initialize crypto helper if database and passphrase are configured
|
||||||
|
if c.cryptoDbPath != "" && c.config.CryptoPassphrase != "" {
|
||||||
|
if err := c.initCrypto(ctx); err != nil {
|
||||||
|
logger.WarnCF(
|
||||||
|
"matrix",
|
||||||
|
"Failed to initialize crypto, continuing without encryption support",
|
||||||
|
map[string]any{
|
||||||
|
"error": err.Error(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
c.syncer.OnEventType(event.EventMessage, c.handleMessageEvent)
|
c.syncer.OnEventType(event.EventMessage, c.handleMessageEvent)
|
||||||
|
c.syncer.OnEventType(event.EventEncrypted, c.handleMessageEvent)
|
||||||
c.syncer.OnEventType(event.StateMember, c.handleMemberEvent)
|
c.syncer.OnEventType(event.StateMember, c.handleMemberEvent)
|
||||||
|
|
||||||
c.SetRunning(true)
|
c.SetRunning(true)
|
||||||
|
|
@ -266,10 +295,84 @@ func (c *MatrixChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
c.stopTypingSessions(ctx)
|
c.stopTypingSessions(ctx)
|
||||||
|
|
||||||
|
// Close crypto helper if initialized
|
||||||
|
if c.cryptoHelper != nil {
|
||||||
|
c.cryptoHelper.Close()
|
||||||
|
c.cryptoHelper = nil
|
||||||
|
c.client.Crypto = nil
|
||||||
|
}
|
||||||
|
|
||||||
logger.InfoC("matrix", "Matrix channel stopped")
|
logger.InfoC("matrix", "Matrix channel stopped")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *MatrixChannel) initCrypto(ctx context.Context) error {
|
||||||
|
logger.InfoC("matrix", "Initializing crypto helper")
|
||||||
|
|
||||||
|
// Ensure the crypto database directory exists
|
||||||
|
if err := os.MkdirAll(c.cryptoDbPath, 0o700); err != nil {
|
||||||
|
return fmt.Errorf("create crypto database directory: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create database with sqlite driver (modernc.org/sqlite)
|
||||||
|
dbPath := filepath.Join(c.cryptoDbPath, dbName)
|
||||||
|
connStr := "file:" + dbPath + "?_foreign_keys=on"
|
||||||
|
|
||||||
|
db, err := sql.Open(sqliteDriver, connStr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("open crypto database: %w", err)
|
||||||
|
}
|
||||||
|
db.SetMaxOpenConns(1)
|
||||||
|
db.SetMaxIdleConns(1)
|
||||||
|
|
||||||
|
// Execute PRAGMA statements
|
||||||
|
// This is equivalent to the "sqlite3-fk-wal" dialect used by cryptohelper
|
||||||
|
pragmaStmts := []string{
|
||||||
|
"PRAGMA foreign_keys = ON",
|
||||||
|
"PRAGMA journal_mode = WAL",
|
||||||
|
"PRAGMA synchronous = NORMAL",
|
||||||
|
"PRAGMA busy_timeout = 5000",
|
||||||
|
}
|
||||||
|
for _, pragma := range pragmaStmts {
|
||||||
|
if _, err = db.ExecContext(ctx, pragma); err != nil {
|
||||||
|
_ = db.Close()
|
||||||
|
return fmt.Errorf("execute %s: %w", pragma, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wrap with dbutil for dialect support
|
||||||
|
wrappedDB, err := dbutil.NewWithDB(db, sqliteDriver)
|
||||||
|
if err != nil {
|
||||||
|
_ = db.Close()
|
||||||
|
return fmt.Errorf("wrap database: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cryptoHelper, err := cryptohelper.NewCryptoHelper(c.client, []byte(c.config.CryptoPassphrase), wrappedDB)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("create crypto helper: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.client.DeviceID == "" {
|
||||||
|
resp, whoamiErr := c.client.Whoami(ctx)
|
||||||
|
if whoamiErr != nil {
|
||||||
|
_ = db.Close()
|
||||||
|
return fmt.Errorf("get device ID via whoami: %w", whoamiErr)
|
||||||
|
}
|
||||||
|
c.client.DeviceID = resp.DeviceID
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = cryptoHelper.Init(ctx); err != nil {
|
||||||
|
cryptoHelper.Close()
|
||||||
|
return fmt.Errorf("init crypto helper: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.client.Crypto = cryptoHelper
|
||||||
|
c.cryptoHelper = cryptoHelper
|
||||||
|
|
||||||
|
logger.InfoC("matrix", "Crypto helper initialized successfully")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func markdownToHTML(md string) string {
|
func markdownToHTML(md string) string {
|
||||||
p := parser.NewWithExtensions(parser.CommonExtensions | parser.AutoHeadingIDs)
|
p := parser.NewWithExtensions(parser.CommonExtensions | parser.AutoHeadingIDs)
|
||||||
renderer := mdhtml.NewRenderer(mdhtml.RendererOptions{Flags: mdhtml.CommonFlags})
|
renderer := mdhtml.NewRenderer(mdhtml.RendererOptions{Flags: mdhtml.CommonFlags})
|
||||||
|
|
@ -548,9 +651,26 @@ func (c *MatrixChannel) handleMessageEvent(ctx context.Context, evt *event.Event
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
msgEvt := evt.Content.AsMessage()
|
var msgEvt *event.MessageEventContent
|
||||||
if msgEvt == nil {
|
switch evt.Type {
|
||||||
return
|
case event.EventMessage:
|
||||||
|
// When crypto is enabled, events marked WasEncrypted=true are
|
||||||
|
// re-dispatched by c.cryptoHelper after decryption and will be
|
||||||
|
// processed again in the EventEncrypted branch. Skip to avoid duplication.
|
||||||
|
if c.client.Crypto != nil && evt.Mautrix.WasEncrypted {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
msgEvt = evt.Content.AsMessage()
|
||||||
|
if msgEvt == nil || msgEvt.MsgType == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
case event.EventEncrypted:
|
||||||
|
var ok bool
|
||||||
|
msgEvt, ok = c.decryptEvent(ctx, evt)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Ignore edits.
|
// Ignore edits.
|
||||||
|
|
@ -642,6 +762,36 @@ func (c *MatrixChannel) handleMessageEvent(ctx context.Context, evt *event.Event
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// decryptEvent decrypts an encrypted event and returns the decrypted message event content.
|
||||||
|
// It returns the decrypted content and a boolean indicating whether decryption was successful.
|
||||||
|
func (c *MatrixChannel) decryptEvent(ctx context.Context, evt *event.Event) (*event.MessageEventContent, bool) {
|
||||||
|
if c.client.Crypto == nil {
|
||||||
|
logger.DebugCF("matrix", "Received encrypted message but crypto is not enabled", map[string]any{
|
||||||
|
"room_id": evt.RoomID.String(),
|
||||||
|
})
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
decrypted, err := c.client.Crypto.Decrypt(ctx, evt)
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("matrix", "Failed to decrypt message", map[string]any{
|
||||||
|
"room_id": evt.RoomID.String(),
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
if decrypted.Type != event.EventMessage {
|
||||||
|
logger.DebugCF("matrix", "Decrypted event is not a message event", map[string]any{
|
||||||
|
"room_id": evt.RoomID.String(),
|
||||||
|
"type": decrypted.Type.String(),
|
||||||
|
})
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return decrypted.Content.AsMessage(), true
|
||||||
|
}
|
||||||
|
|
||||||
func (c *MatrixChannel) extractInboundContent(
|
func (c *MatrixChannel) extractInboundContent(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
msgEvt *event.MessageEventContent,
|
msgEvt *event.MessageEventContent,
|
||||||
|
|
|
||||||
|
|
@ -54,12 +54,13 @@ func (pc *picoConn) close() {
|
||||||
// It serves as the reference implementation for all optional capability interfaces.
|
// It serves as the reference implementation for all optional capability interfaces.
|
||||||
type PicoChannel struct {
|
type PicoChannel struct {
|
||||||
*channels.BaseChannel
|
*channels.BaseChannel
|
||||||
config config.PicoConfig
|
config config.PicoConfig
|
||||||
upgrader websocket.Upgrader
|
upgrader websocket.Upgrader
|
||||||
connections sync.Map // connID → *picoConn
|
connections map[string]*picoConn // connID -> *picoConn
|
||||||
connCount atomic.Int32
|
sessionConnections map[string]map[string]*picoConn // sessionID -> connID -> *picoConn
|
||||||
ctx context.Context
|
connsMu sync.RWMutex
|
||||||
cancel context.CancelFunc
|
ctx context.Context
|
||||||
|
cancel context.CancelFunc
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPicoChannel creates a new Pico Protocol channel.
|
// NewPicoChannel creates a new Pico Protocol channel.
|
||||||
|
|
@ -92,9 +93,104 @@ func NewPicoChannel(cfg config.PicoConfig, messageBus *bus.MessageBus) (*PicoCha
|
||||||
ReadBufferSize: 1024,
|
ReadBufferSize: 1024,
|
||||||
WriteBufferSize: 1024,
|
WriteBufferSize: 1024,
|
||||||
},
|
},
|
||||||
|
connections: make(map[string]*picoConn),
|
||||||
|
sessionConnections: make(map[string]map[string]*picoConn),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// createAndAddConnection checks MaxConnections and registers a connection atomically.
|
||||||
|
func (c *PicoChannel) createAndAddConnection(conn *websocket.Conn, sessionID string, maxConns int) (*picoConn, error) {
|
||||||
|
c.connsMu.Lock()
|
||||||
|
defer c.connsMu.Unlock()
|
||||||
|
if len(c.connections) >= maxConns {
|
||||||
|
return nil, channels.ErrTemporary
|
||||||
|
}
|
||||||
|
|
||||||
|
var connID string
|
||||||
|
for {
|
||||||
|
connID = uuid.New().String()
|
||||||
|
if _, exists := c.connections[connID]; !exists {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pc := &picoConn{
|
||||||
|
id: connID,
|
||||||
|
conn: conn,
|
||||||
|
sessionID: sessionID,
|
||||||
|
}
|
||||||
|
|
||||||
|
c.connections[pc.id] = pc
|
||||||
|
bySession, ok := c.sessionConnections[pc.sessionID]
|
||||||
|
if !ok {
|
||||||
|
bySession = make(map[string]*picoConn)
|
||||||
|
c.sessionConnections[pc.sessionID] = bySession
|
||||||
|
}
|
||||||
|
bySession[pc.id] = pc
|
||||||
|
|
||||||
|
return pc, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeConnection deletes a connection from indexes and returns it when found.
|
||||||
|
func (c *PicoChannel) removeConnection(connID string) *picoConn {
|
||||||
|
c.connsMu.Lock()
|
||||||
|
defer c.connsMu.Unlock()
|
||||||
|
|
||||||
|
pc, ok := c.connections[connID]
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(c.connections, connID)
|
||||||
|
if bySession, ok := c.sessionConnections[pc.sessionID]; ok {
|
||||||
|
delete(bySession, connID)
|
||||||
|
if len(bySession) == 0 {
|
||||||
|
delete(c.sessionConnections, pc.sessionID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return pc
|
||||||
|
}
|
||||||
|
|
||||||
|
// takeAllConnections snapshots and clears all connection indexes.
|
||||||
|
func (c *PicoChannel) takeAllConnections() []*picoConn {
|
||||||
|
c.connsMu.Lock()
|
||||||
|
defer c.connsMu.Unlock()
|
||||||
|
|
||||||
|
all := make([]*picoConn, 0, len(c.connections))
|
||||||
|
for _, pc := range c.connections {
|
||||||
|
all = append(all, pc)
|
||||||
|
}
|
||||||
|
clear(c.connections)
|
||||||
|
clear(c.sessionConnections)
|
||||||
|
|
||||||
|
return all
|
||||||
|
}
|
||||||
|
|
||||||
|
// sessionConnectionsSnapshot returns all active connections for a session.
|
||||||
|
func (c *PicoChannel) sessionConnectionsSnapshot(sessionID string) []*picoConn {
|
||||||
|
c.connsMu.RLock()
|
||||||
|
defer c.connsMu.RUnlock()
|
||||||
|
|
||||||
|
bySession, ok := c.sessionConnections[sessionID]
|
||||||
|
if !ok || len(bySession) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
conns := make([]*picoConn, 0, len(bySession))
|
||||||
|
for _, pc := range bySession {
|
||||||
|
conns = append(conns, pc)
|
||||||
|
}
|
||||||
|
return conns
|
||||||
|
}
|
||||||
|
|
||||||
|
// currentConnCount returns a lock-protected snapshot of active connection count.
|
||||||
|
func (c *PicoChannel) currentConnCount() int {
|
||||||
|
c.connsMu.RLock()
|
||||||
|
defer c.connsMu.RUnlock()
|
||||||
|
return len(c.connections)
|
||||||
|
}
|
||||||
|
|
||||||
// Start implements Channel.
|
// Start implements Channel.
|
||||||
func (c *PicoChannel) Start(ctx context.Context) error {
|
func (c *PicoChannel) Start(ctx context.Context) error {
|
||||||
logger.InfoC("pico", "Starting Pico Protocol channel")
|
logger.InfoC("pico", "Starting Pico Protocol channel")
|
||||||
|
|
@ -110,13 +206,9 @@ func (c *PicoChannel) Stop(ctx context.Context) error {
|
||||||
c.SetRunning(false)
|
c.SetRunning(false)
|
||||||
|
|
||||||
// Close all connections
|
// Close all connections
|
||||||
c.connections.Range(func(key, value any) bool {
|
for _, pc := range c.takeAllConnections() {
|
||||||
if pc, ok := value.(*picoConn); ok {
|
pc.close()
|
||||||
pc.close()
|
}
|
||||||
}
|
|
||||||
c.connections.Delete(key)
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
|
|
||||||
if c.cancel != nil {
|
if c.cancel != nil {
|
||||||
c.cancel()
|
c.cancel()
|
||||||
|
|
@ -133,8 +225,8 @@ func (c *PicoChannel) WebhookPath() string { return "/pico/" }
|
||||||
func (c *PicoChannel) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
func (c *PicoChannel) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
path := strings.TrimPrefix(r.URL.Path, "/pico")
|
path := strings.TrimPrefix(r.URL.Path, "/pico")
|
||||||
|
|
||||||
switch {
|
switch path {
|
||||||
case path == "/ws" || path == "/ws/":
|
case "/ws", "/ws/":
|
||||||
c.handleWebSocket(w, r)
|
c.handleWebSocket(w, r)
|
||||||
default:
|
default:
|
||||||
http.NotFound(w, r)
|
http.NotFound(w, r)
|
||||||
|
|
@ -208,23 +300,16 @@ func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error {
|
||||||
msg.SessionID = sessionID
|
msg.SessionID = sessionID
|
||||||
|
|
||||||
var sent bool
|
var sent bool
|
||||||
c.connections.Range(func(key, value any) bool {
|
for _, pc := range c.sessionConnectionsSnapshot(sessionID) {
|
||||||
pc, ok := value.(*picoConn)
|
if err := pc.writeJSON(msg); err != nil {
|
||||||
if !ok {
|
logger.DebugCF("pico", "Write to connection failed", map[string]any{
|
||||||
return true
|
"conn_id": pc.id,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
sent = true
|
||||||
}
|
}
|
||||||
if pc.sessionID == sessionID {
|
}
|
||||||
if err := pc.writeJSON(msg); err != nil {
|
|
||||||
logger.DebugCF("pico", "Write to connection failed", map[string]any{
|
|
||||||
"conn_id": pc.id,
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
} else {
|
|
||||||
sent = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
|
|
||||||
if !sent {
|
if !sent {
|
||||||
return fmt.Errorf("no active connections for session %s: %w", sessionID, channels.ErrSendFailed)
|
return fmt.Errorf("no active connections for session %s: %w", sessionID, channels.ErrSendFailed)
|
||||||
|
|
@ -250,7 +335,7 @@ func (c *PicoChannel) handleWebSocket(w http.ResponseWriter, r *http.Request) {
|
||||||
if maxConns <= 0 {
|
if maxConns <= 0 {
|
||||||
maxConns = 100
|
maxConns = 100
|
||||||
}
|
}
|
||||||
if int(c.connCount.Load()) >= maxConns {
|
if c.currentConnCount() >= maxConns {
|
||||||
http.Error(w, "too many connections", http.StatusServiceUnavailable)
|
http.Error(w, "too many connections", http.StatusServiceUnavailable)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
@ -275,15 +360,17 @@ func (c *PicoChannel) handleWebSocket(w http.ResponseWriter, r *http.Request) {
|
||||||
sessionID = uuid.New().String()
|
sessionID = uuid.New().String()
|
||||||
}
|
}
|
||||||
|
|
||||||
pc := &picoConn{
|
pc, err := c.createAndAddConnection(conn, sessionID, maxConns)
|
||||||
id: uuid.New().String(),
|
if err != nil {
|
||||||
conn: conn,
|
_ = conn.WriteControl(
|
||||||
sessionID: sessionID,
|
websocket.CloseMessage,
|
||||||
|
websocket.FormatCloseMessage(websocket.CloseTryAgainLater, "too many connections"),
|
||||||
|
time.Now().Add(2*time.Second),
|
||||||
|
)
|
||||||
|
_ = conn.Close()
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.connections.Store(pc.id, pc)
|
|
||||||
c.connCount.Add(1)
|
|
||||||
|
|
||||||
logger.InfoCF("pico", "WebSocket client connected", map[string]any{
|
logger.InfoCF("pico", "WebSocket client connected", map[string]any{
|
||||||
"conn_id": pc.id,
|
"conn_id": pc.id,
|
||||||
"session_id": sessionID,
|
"session_id": sessionID,
|
||||||
|
|
@ -341,12 +428,12 @@ func (c *PicoChannel) matchedSubprotocol(r *http.Request) string {
|
||||||
func (c *PicoChannel) readLoop(pc *picoConn) {
|
func (c *PicoChannel) readLoop(pc *picoConn) {
|
||||||
defer func() {
|
defer func() {
|
||||||
pc.close()
|
pc.close()
|
||||||
c.connections.Delete(pc.id)
|
if removed := c.removeConnection(pc.id); removed != nil {
|
||||||
c.connCount.Add(-1)
|
logger.InfoCF("pico", "WebSocket client disconnected", map[string]any{
|
||||||
logger.InfoCF("pico", "WebSocket client disconnected", map[string]any{
|
"conn_id": removed.id,
|
||||||
"conn_id": pc.id,
|
"session_id": removed.sessionID,
|
||||||
"session_id": pc.sessionID,
|
})
|
||||||
})
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
readTimeout := time.Duration(c.config.ReadTimeout) * time.Second
|
readTimeout := time.Duration(c.config.ReadTimeout) * time.Second
|
||||||
|
|
|
||||||
144
pkg/channels/pico/pico_test.go
Normal file
144
pkg/channels/pico/pico_test.go
Normal file
|
|
@ -0,0 +1,144 @@
|
||||||
|
package pico
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestPicoChannel(t *testing.T) *PicoChannel {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
cfg := config.PicoConfig{}
|
||||||
|
cfg.SetToken("test-token")
|
||||||
|
ch, err := NewPicoChannel(cfg, bus.NewMessageBus())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewPicoChannel: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ch.ctx = context.Background()
|
||||||
|
return ch
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAndAddConnection_RespectsMaxConnectionsConcurrently(t *testing.T) {
|
||||||
|
ch := newTestPicoChannel(t)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxConns = 5
|
||||||
|
goroutines = 64
|
||||||
|
sessionID = "session-a"
|
||||||
|
)
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
var mu sync.Mutex
|
||||||
|
successCount := 0
|
||||||
|
errCount := 0
|
||||||
|
|
||||||
|
wg.Add(goroutines)
|
||||||
|
for i := 0; i < goroutines; i++ {
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
|
||||||
|
pc, err := ch.createAndAddConnection(nil, sessionID, maxConns)
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
successCount++
|
||||||
|
if pc == nil {
|
||||||
|
t.Errorf("pc is nil on success")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !errors.Is(err, channels.ErrTemporary) {
|
||||||
|
t.Errorf("unexpected error: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
errCount++
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
if successCount > maxConns {
|
||||||
|
t.Fatalf("successCount=%d > maxConns=%d", successCount, maxConns)
|
||||||
|
}
|
||||||
|
if successCount+errCount != goroutines {
|
||||||
|
t.Fatalf("success=%d err=%d total=%d want=%d", successCount, errCount, successCount+errCount, goroutines)
|
||||||
|
}
|
||||||
|
if got := ch.currentConnCount(); got != maxConns {
|
||||||
|
t.Fatalf("currentConnCount=%d want=%d", got, maxConns)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveConnection_CleansBothIndexes(t *testing.T) {
|
||||||
|
ch := newTestPicoChannel(t)
|
||||||
|
|
||||||
|
pc, err := ch.createAndAddConnection(nil, "session-cleanup", 10)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("createAndAddConnection: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
removed := ch.removeConnection(pc.id)
|
||||||
|
if removed == nil {
|
||||||
|
t.Fatal("removeConnection returned nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
ch.connsMu.RLock()
|
||||||
|
defer ch.connsMu.RUnlock()
|
||||||
|
|
||||||
|
if _, ok := ch.connections[pc.id]; ok {
|
||||||
|
t.Fatalf("connID %s still exists in connections", pc.id)
|
||||||
|
}
|
||||||
|
if _, ok := ch.sessionConnections[pc.sessionID]; ok {
|
||||||
|
t.Fatalf("session %s still exists in sessionConnections", pc.sessionID)
|
||||||
|
}
|
||||||
|
if got := len(ch.connections); got != 0 {
|
||||||
|
t.Fatalf("len(connections)=%d want=0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBroadcastToSession_TargetsOnlyRequestedSession(t *testing.T) {
|
||||||
|
ch := newTestPicoChannel(t)
|
||||||
|
|
||||||
|
target := &picoConn{id: "target", sessionID: "s-target"}
|
||||||
|
target.closed.Store(true)
|
||||||
|
ch.addConnForTest(target)
|
||||||
|
|
||||||
|
other := &picoConn{id: "other", sessionID: "s-other"}
|
||||||
|
ch.addConnForTest(other)
|
||||||
|
|
||||||
|
err := ch.broadcastToSession("pico:s-target", newMessage(TypeMessageCreate, map[string]any{"content": "hello"}))
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected send failure due to closed target connection")
|
||||||
|
}
|
||||||
|
if !errors.Is(err, channels.ErrSendFailed) {
|
||||||
|
t.Fatalf("expected ErrSendFailed, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PicoChannel) addConnForTest(pc *picoConn) {
|
||||||
|
c.connsMu.Lock()
|
||||||
|
defer c.connsMu.Unlock()
|
||||||
|
if c.connections == nil {
|
||||||
|
c.connections = make(map[string]*picoConn)
|
||||||
|
}
|
||||||
|
if c.sessionConnections == nil {
|
||||||
|
c.sessionConnections = make(map[string]map[string]*picoConn)
|
||||||
|
}
|
||||||
|
if _, exists := c.connections[pc.id]; exists {
|
||||||
|
panic(fmt.Sprintf("duplicate conn id in test: %s", pc.id))
|
||||||
|
}
|
||||||
|
c.connections[pc.id] = pc
|
||||||
|
bySession, ok := c.sessionConnections[pc.sessionID]
|
||||||
|
if !ok {
|
||||||
|
bySession = make(map[string]*picoConn)
|
||||||
|
c.sessionConnections[pc.sessionID] = bySession
|
||||||
|
}
|
||||||
|
bySession[pc.id] = pc
|
||||||
|
}
|
||||||
|
|
@ -642,8 +642,12 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if content == "" && len(mediaPaths) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
if content == "" {
|
if content == "" {
|
||||||
content = "[empty message]"
|
content = "[media only]"
|
||||||
}
|
}
|
||||||
|
|
||||||
// In group chats, apply unified group trigger filtering
|
// In group chats, apply unified group trigger filtering
|
||||||
|
|
|
||||||
|
|
@ -641,3 +641,35 @@ func TestHandleMessage_ReplyThread_NonForum_NoIsolation(t *testing.T) {
|
||||||
assert.Empty(t, inbound.Metadata["parent_peer_kind"])
|
assert.Empty(t, inbound.Metadata["parent_peer_kind"])
|
||||||
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_EmptyContent_Ignored(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Service message with no text/caption/media (like ForumTopicCreated)
|
||||||
|
msg := &telego.Message{
|
||||||
|
MessageID: 123,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: 456,
|
||||||
|
Type: "group",
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 789,
|
||||||
|
FirstName: "User",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Should NOT publish to message bus
|
||||||
|
select {
|
||||||
|
case <-messageBus.InboundChan():
|
||||||
|
t.Fatal("Empty message should not be published to message bus")
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,559 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ---- Webhook mode tests ----
|
|
||||||
|
|
||||||
func TestNewWeComAIBotChannel_WebhookMode(t *testing.T) {
|
|
||||||
t.Run("success with valid config", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
cfg.WebhookPath = "/webhook/test"
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Expected no error, got %v", err)
|
|
||||||
}
|
|
||||||
if ch == nil {
|
|
||||||
t.Fatal("Expected channel to be created")
|
|
||||||
}
|
|
||||||
if ch.Name() != "wecom_aibot" {
|
|
||||||
t.Errorf("Expected name 'wecom_aibot', got '%s'", ch.Name())
|
|
||||||
}
|
|
||||||
// Webhook mode must implement WebhookHandler.
|
|
||||||
if _, ok := ch.(channels.WebhookHandler); !ok {
|
|
||||||
t.Error("Webhook mode channel should implement WebhookHandler")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("error with missing token", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
_, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("Expected error for missing token, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("error with missing encoding key", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
_, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("Expected error for missing encoding key, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComAIBotWebhookChannelStartStop(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to create channel: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
if err := ch.Start(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to start channel: %v", err)
|
|
||||||
}
|
|
||||||
if !ch.IsRunning() {
|
|
||||||
t.Error("Expected channel to be running after Start")
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := ch.Stop(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to stop channel: %v", err)
|
|
||||||
}
|
|
||||||
if ch.IsRunning() {
|
|
||||||
t.Error("Expected channel to be stopped after Stop")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComAIBotChannelWebhookPath(t *testing.T) {
|
|
||||||
t.Run("default path", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
|
|
||||||
wh, ok := ch.(channels.WebhookHandler)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("Expected channel to implement WebhookHandler")
|
|
||||||
}
|
|
||||||
expectedPath := "/webhook/wecom-aibot"
|
|
||||||
if wh.WebhookPath() != expectedPath {
|
|
||||||
t.Errorf("Expected webhook path '%s', got '%s'", expectedPath, wh.WebhookPath())
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("custom path", func(t *testing.T) {
|
|
||||||
customPath := "/custom/webhook"
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
cfg.WebhookPath = customPath
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
|
|
||||||
wh, ok := ch.(channels.WebhookHandler)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("Expected channel to implement WebhookHandler")
|
|
||||||
}
|
|
||||||
if wh.WebhookPath() != customPath {
|
|
||||||
t.Errorf("Expected webhook path '%s', got '%s'", customPath, wh.WebhookPath())
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComAIBotChannelGetStreamResponseProcessingMessage(t *testing.T) {
|
|
||||||
validAESKey := "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
|
|
||||||
|
|
||||||
t.Run("uses default processing message", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey(validAESKey)
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
channel, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to create channel: %v", err)
|
|
||||||
}
|
|
||||||
ch, ok := channel.(*WeComAIBotChannel)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("Expected webhook mode channel")
|
|
||||||
}
|
|
||||||
|
|
||||||
task := &streamTask{
|
|
||||||
StreamID: "stream-default",
|
|
||||||
ChatID: "chat-default",
|
|
||||||
Deadline: time.Now().Add(-time.Second),
|
|
||||||
}
|
|
||||||
ch.streamTasks[task.StreamID] = task
|
|
||||||
ch.chatTasks[task.ChatID] = []*streamTask{task}
|
|
||||||
|
|
||||||
resp := decodeStreamResponse(t, ch, ch.getStreamResponse(task, "1234567890", "nonce"))
|
|
||||||
|
|
||||||
if !resp.Stream.Finish {
|
|
||||||
t.Fatal("Expected finished stream response after deadline")
|
|
||||||
}
|
|
||||||
if resp.Stream.Content != config.DefaultWeComAIBotProcessingMessage {
|
|
||||||
t.Fatalf("Expected default processing message %q, got %q",
|
|
||||||
config.DefaultWeComAIBotProcessingMessage, resp.Stream.Content)
|
|
||||||
}
|
|
||||||
if !task.StreamClosed {
|
|
||||||
t.Fatal("Expected task stream to be marked closed")
|
|
||||||
}
|
|
||||||
if _, ok := ch.streamTasks[task.StreamID]; ok {
|
|
||||||
t.Fatal("Expected closed stream task to be removed from streamTasks")
|
|
||||||
}
|
|
||||||
if len(ch.chatTasks[task.ChatID]) != 1 {
|
|
||||||
t.Fatalf("Expected task to remain queued for response_url delivery, got %d entries",
|
|
||||||
len(ch.chatTasks[task.ChatID]))
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("uses custom processing message", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
ProcessingMessage: "Please wait a moment. The result will be delivered in a follow-up message.",
|
|
||||||
}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey(validAESKey)
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
channel, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to create channel: %v", err)
|
|
||||||
}
|
|
||||||
ch, ok := channel.(*WeComAIBotChannel)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("Expected webhook mode channel")
|
|
||||||
}
|
|
||||||
|
|
||||||
task := &streamTask{
|
|
||||||
StreamID: "stream-custom",
|
|
||||||
ChatID: "chat-custom",
|
|
||||||
Deadline: time.Now().Add(-time.Second),
|
|
||||||
}
|
|
||||||
|
|
||||||
resp := decodeStreamResponse(t, ch, ch.getStreamResponse(task, "1234567890", "nonce"))
|
|
||||||
|
|
||||||
if resp.Stream.Content != cfg.ProcessingMessage {
|
|
||||||
t.Fatalf("Expected custom processing message %q, got %q", cfg.ProcessingMessage, resp.Stream.Content)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGenerateStreamID(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
webhookCh, ok := ch.(*WeComAIBotChannel)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("Expected webhook mode channel")
|
|
||||||
}
|
|
||||||
|
|
||||||
ids := make(map[string]bool)
|
|
||||||
for i := 0; i < 100; i++ {
|
|
||||||
id := webhookCh.generateStreamID()
|
|
||||||
if len(id) != 10 {
|
|
||||||
t.Errorf("Expected stream ID length 10, got %d", len(id))
|
|
||||||
}
|
|
||||||
if ids[id] {
|
|
||||||
t.Errorf("Duplicate stream ID generated: %s", id)
|
|
||||||
}
|
|
||||||
ids[id] = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestEncryptDecrypt(t *testing.T) {
|
|
||||||
// Use a valid 43-character base64 key (企业微信标准格式)
|
|
||||||
cfg := config.WeComAIBotConfig{}
|
|
||||||
cfg.Enabled = true
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG") // 43 characters
|
|
||||||
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, _ := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
webhookCh, ok := ch.(*WeComAIBotChannel)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("Expected webhook mode channel")
|
|
||||||
}
|
|
||||||
|
|
||||||
plaintext := "Hello, World!"
|
|
||||||
receiveid := ""
|
|
||||||
|
|
||||||
encrypted, err := webhookCh.encryptMessage(plaintext, receiveid)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to encrypt message: %v", err)
|
|
||||||
}
|
|
||||||
if encrypted == "" {
|
|
||||||
t.Fatal("Encrypted message is empty")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Decrypt
|
|
||||||
decrypted, err := decryptMessageWithVerify(encrypted, cfg.EncodingAESKey(), receiveid)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to decrypt message: %v", err)
|
|
||||||
}
|
|
||||||
if decrypted != plaintext {
|
|
||||||
t.Errorf("Expected decrypted message '%s', got '%s'", plaintext, decrypted)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGenerateSignature(t *testing.T) {
|
|
||||||
token := "test_token"
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
encrypt := "encrypted_msg"
|
|
||||||
|
|
||||||
signature := computeSignature(token, timestamp, nonce, encrypt)
|
|
||||||
if signature == "" {
|
|
||||||
t.Error("Generated signature is empty")
|
|
||||||
}
|
|
||||||
if !verifySignature(token, signature, timestamp, nonce, encrypt) {
|
|
||||||
t.Error("Generated signature does not verify correctly")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func decodeStreamResponse(t *testing.T, ch *WeComAIBotChannel, encryptedResponse string) WeComAIBotStreamResponse {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var wrapped WeComAIBotEncryptedResponse
|
|
||||||
if err := json.Unmarshal([]byte(encryptedResponse), &wrapped); err != nil {
|
|
||||||
t.Fatalf("Failed to unmarshal encrypted response: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
plaintext, err := decryptMessageWithVerify(wrapped.Encrypt, ch.config.EncodingAESKey(), "")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to decrypt response: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var resp WeComAIBotStreamResponse
|
|
||||||
if err := json.Unmarshal([]byte(plaintext), &resp); err != nil {
|
|
||||||
t.Fatalf("Failed to unmarshal decrypted response: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return resp
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- WebSocket long-connection mode tests ----
|
|
||||||
|
|
||||||
func TestNewWeComAIBotChannel_WSMode(t *testing.T) {
|
|
||||||
t.Run("success with bot_id and secret", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
BotID: "test_bot_id",
|
|
||||||
}
|
|
||||||
cfg.SetSecret("test_secret")
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Expected no error, got %v", err)
|
|
||||||
}
|
|
||||||
if ch == nil {
|
|
||||||
t.Fatal("Expected channel to be created")
|
|
||||||
}
|
|
||||||
if ch.Name() != "wecom_aibot" {
|
|
||||||
t.Errorf("Expected name 'wecom_aibot', got '%s'", ch.Name())
|
|
||||||
}
|
|
||||||
// WebSocket mode must NOT implement WebhookHandler.
|
|
||||||
if _, ok := ch.(channels.WebhookHandler); ok {
|
|
||||||
t.Error("WebSocket mode channel should NOT implement WebhookHandler")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("ws mode takes priority over webhook fields", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
BotID: "test_bot_id",
|
|
||||||
}
|
|
||||||
cfg.SetSecret("test_secret")
|
|
||||||
cfg.SetToken("also_set")
|
|
||||||
cfg.SetEncodingAESKey("testkey1234567890123456789012345678901234567")
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Expected no error, got %v", err)
|
|
||||||
}
|
|
||||||
if _, ok := ch.(*WeComAIBotWSChannel); !ok {
|
|
||||||
t.Error("Expected WebSocket mode channel when both BotID+secret and Token+Key are set")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("error with missing bot_id", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
}
|
|
||||||
cfg.SetSecret("test_secret")
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
_, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
// Missing bot_id alone means neither WS mode nor webhook mode is fully configured.
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("Expected error for missing bot_id, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("error with missing secret", func(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
BotID: "test_bot_id",
|
|
||||||
}
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
_, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("Expected error for missing secret, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComAIBotWSChannelStartStop(t *testing.T) {
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
BotID: "test_bot_id",
|
|
||||||
}
|
|
||||||
cfg.SetSecret("test_secret")
|
|
||||||
messageBus := bus.NewMessageBus()
|
|
||||||
ch, err := NewWeComAIBotChannel(cfg, messageBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to create channel: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// Start launches a background goroutine; it should not block or return an error.
|
|
||||||
if err := ch.Start(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to start channel: %v", err)
|
|
||||||
}
|
|
||||||
if !ch.IsRunning() {
|
|
||||||
t.Error("Expected channel to be running after Start")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stop should work regardless of whether the WebSocket actually connected.
|
|
||||||
if err := ch.Stop(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to stop channel: %v", err)
|
|
||||||
}
|
|
||||||
if ch.IsRunning() {
|
|
||||||
t.Error("Expected channel to be stopped after Stop")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGenerateRandomID(t *testing.T) {
|
|
||||||
ids := make(map[string]bool)
|
|
||||||
for i := 0; i < 200; i++ {
|
|
||||||
id := generateRandomID(10)
|
|
||||||
if len(id) != 10 {
|
|
||||||
t.Errorf("Expected ID length 10, got %d", len(id))
|
|
||||||
}
|
|
||||||
if ids[id] {
|
|
||||||
t.Errorf("Duplicate ID generated: %s", id)
|
|
||||||
}
|
|
||||||
ids[id] = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWSGenerateID(t *testing.T) {
|
|
||||||
ids := make(map[string]bool)
|
|
||||||
for i := 0; i < 200; i++ {
|
|
||||||
id := wsGenerateID()
|
|
||||||
if len(id) != 10 {
|
|
||||||
t.Errorf("Expected ID length 10, got %d", len(id))
|
|
||||||
}
|
|
||||||
if ids[id] {
|
|
||||||
t.Errorf("Duplicate wsGenerateID result: %s", id)
|
|
||||||
}
|
|
||||||
ids[id] = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- Webhook streaming fallback tests ----
|
|
||||||
|
|
||||||
// makeWebhookChannel creates a started WeComAIBotChannel for testing.
|
|
||||||
func makeWebhookChannel(t *testing.T) *WeComAIBotChannel {
|
|
||||||
t.Helper()
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey("abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG")
|
|
||||||
ch, err := NewWeComAIBotChannel(cfg, bus.NewMessageBus())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create channel: %v", err)
|
|
||||||
}
|
|
||||||
wc := ch.(*WeComAIBotChannel)
|
|
||||||
wc.ctx, wc.cancel = context.WithCancel(context.Background())
|
|
||||||
return wc
|
|
||||||
}
|
|
||||||
|
|
||||||
// makeStreamTask creates and registers a streamTask for testing.
|
|
||||||
func makeStreamTask(t *testing.T, ch *WeComAIBotChannel, streamID, chatID string, deadline time.Time) *streamTask {
|
|
||||||
t.Helper()
|
|
||||||
task := &streamTask{
|
|
||||||
StreamID: streamID,
|
|
||||||
ChatID: chatID,
|
|
||||||
Deadline: deadline,
|
|
||||||
answerCh: make(chan string, 1),
|
|
||||||
}
|
|
||||||
task.ctx, task.cancel = context.WithCancel(ch.ctx)
|
|
||||||
ch.taskMu.Lock()
|
|
||||||
ch.streamTasks[streamID] = task
|
|
||||||
ch.chatTasks[chatID] = append(ch.chatTasks[chatID], task)
|
|
||||||
ch.taskMu.Unlock()
|
|
||||||
return task
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestGetStreamResponse_ImmediateAnswer verifies that when the agent has already
|
|
||||||
// placed its answer in answerCh, getStreamResponse returns a finish=true response
|
|
||||||
// and fully removes the task.
|
|
||||||
func TestGetStreamResponse_ImmediateAnswer(t *testing.T) {
|
|
||||||
ch := makeWebhookChannel(t)
|
|
||||||
defer ch.cancel()
|
|
||||||
|
|
||||||
task := makeStreamTask(t, ch, "stream-1", "chat-1", time.Now().Add(30*time.Second))
|
|
||||||
task.answerCh <- "hello from agent"
|
|
||||||
|
|
||||||
result := ch.getStreamResponse(task, "ts123", "nonce123")
|
|
||||||
if result == "" {
|
|
||||||
t.Fatal("expected non-empty encrypted response")
|
|
||||||
}
|
|
||||||
|
|
||||||
ch.taskMu.RLock()
|
|
||||||
_, exists := ch.streamTasks["stream-1"]
|
|
||||||
ch.taskMu.RUnlock()
|
|
||||||
if exists {
|
|
||||||
t.Error("task should have been removed from streamTasks after normal finish")
|
|
||||||
}
|
|
||||||
if !task.Finished {
|
|
||||||
t.Error("task.Finished should be true after normal finish")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestGetStreamResponse_DeadlinePassed verifies that when the stream deadline has
|
|
||||||
// elapsed (no agent reply yet), getStreamResponse closes the stream but keeps the
|
|
||||||
// task alive so the response_url fallback can still deliver the answer.
|
|
||||||
func TestGetStreamResponse_DeadlinePassed(t *testing.T) {
|
|
||||||
ch := makeWebhookChannel(t)
|
|
||||||
defer ch.cancel()
|
|
||||||
|
|
||||||
task := makeStreamTask(t, ch, "stream-2", "chat-2", time.Now().Add(-time.Millisecond))
|
|
||||||
|
|
||||||
result := ch.getStreamResponse(task, "ts456", "nonce456")
|
|
||||||
if result == "" {
|
|
||||||
t.Fatal("expected non-empty encrypted response")
|
|
||||||
}
|
|
||||||
|
|
||||||
ch.taskMu.RLock()
|
|
||||||
_, stillStreaming := ch.streamTasks["stream-2"]
|
|
||||||
ch.taskMu.RUnlock()
|
|
||||||
if stillStreaming {
|
|
||||||
t.Error("task should have been removed from streamTasks after deadline")
|
|
||||||
}
|
|
||||||
if !task.StreamClosed {
|
|
||||||
t.Error("task.StreamClosed should be true after deadline")
|
|
||||||
}
|
|
||||||
if task.Finished {
|
|
||||||
t.Error("task.Finished must remain false: agent reply still expected via response_url")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestGetStreamResponse_StillPending verifies that when neither the agent has
|
|
||||||
// replied nor the deadline has passed, getStreamResponse returns without altering
|
|
||||||
// task state (client should poll again).
|
|
||||||
func TestGetStreamResponse_StillPending(t *testing.T) {
|
|
||||||
ch := makeWebhookChannel(t)
|
|
||||||
defer ch.cancel()
|
|
||||||
|
|
||||||
task := makeStreamTask(t, ch, "stream-3", "chat-3", time.Now().Add(30*time.Second))
|
|
||||||
|
|
||||||
result := ch.getStreamResponse(task, "ts789", "nonce789")
|
|
||||||
if result == "" {
|
|
||||||
t.Fatal("expected non-empty encrypted response")
|
|
||||||
}
|
|
||||||
|
|
||||||
ch.taskMu.RLock()
|
|
||||||
_, exists := ch.streamTasks["stream-3"]
|
|
||||||
ch.taskMu.RUnlock()
|
|
||||||
if !exists {
|
|
||||||
t.Error("pending task should still be in streamTasks")
|
|
||||||
}
|
|
||||||
if task.Finished || task.StreamClosed {
|
|
||||||
t.Error("pending task should not be finished or stream-closed")
|
|
||||||
}
|
|
||||||
// Cleanup.
|
|
||||||
ch.removeTask(task)
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,295 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"os"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/media"
|
|
||||||
)
|
|
||||||
|
|
||||||
// newTestWSChannel creates a WeComAIBotWSChannel ready for unit testing.
|
|
||||||
func newTestWSChannel(t *testing.T) *WeComAIBotWSChannel {
|
|
||||||
t.Helper()
|
|
||||||
cfg := config.WeComAIBotConfig{
|
|
||||||
Enabled: true,
|
|
||||||
BotID: "test_bot_id",
|
|
||||||
}
|
|
||||||
cfg.SetSecret("test_secret")
|
|
||||||
ch, err := newWeComAIBotWSChannel(cfg, bus.NewMessageBus())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create WS channel: %v", err)
|
|
||||||
}
|
|
||||||
return ch
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestStoreWSMedia_NilStore verifies that storeWSMedia returns an error when no
|
|
||||||
// MediaStore has been injected.
|
|
||||||
func TestStoreWSMedia_NilStore(t *testing.T) {
|
|
||||||
ch := newTestWSChannel(t)
|
|
||||||
_, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", "http://any", "", ".jpg")
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected error when no MediaStore is set")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestStoreWSMedia_HTTPError verifies that storeWSMedia propagates HTTP errors
|
|
||||||
// from the media server.
|
|
||||||
func TestStoreWSMedia_HTTPError(t *testing.T) {
|
|
||||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
http.Error(w, "not found", http.StatusNotFound)
|
|
||||||
}))
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
ch := newTestWSChannel(t)
|
|
||||||
ch.SetMediaStore(media.NewFileMediaStore())
|
|
||||||
|
|
||||||
_, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", srv.URL, "", ".jpg")
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected error for HTTP 404")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestStoreWSMedia_ServerUnavailable verifies that storeWSMedia returns a clear
|
|
||||||
// error when the media server cannot be reached.
|
|
||||||
func TestStoreWSMedia_ServerUnavailable(t *testing.T) {
|
|
||||||
ch := newTestWSChannel(t)
|
|
||||||
ch.SetMediaStore(media.NewFileMediaStore())
|
|
||||||
|
|
||||||
// Port 1 is reserved and will refuse the connection immediately.
|
|
||||||
_, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", "http://127.0.0.1:1", "", ".jpg")
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected error for unreachable server")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestStoreWSMedia_Success_NoAES verifies the happy path: the media is downloaded,
|
|
||||||
// a media ref is returned, and the file persists and is readable via Resolve until
|
|
||||||
// ReleaseAll is called. The server returns no Content-Type, so the defaultExt is used.
|
|
||||||
func TestStoreWSMedia_Success_NoAES(t *testing.T) {
|
|
||||||
imageData := bytes.Repeat([]byte("x"), 256)
|
|
||||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
_, _ = w.Write(imageData)
|
|
||||||
}))
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
ch := newTestWSChannel(t)
|
|
||||||
store := media.NewFileMediaStore()
|
|
||||||
ch.SetMediaStore(store)
|
|
||||||
|
|
||||||
ref, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", srv.URL, "", ".jpg")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("expected no error, got %v", err)
|
|
||||||
}
|
|
||||||
if ref == "" {
|
|
||||||
t.Fatal("expected non-empty ref")
|
|
||||||
}
|
|
||||||
|
|
||||||
// File must be accessible after storeWSMedia returns (no premature deletion).
|
|
||||||
path, err := store.Resolve(ref)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ref should resolve: %v", err)
|
|
||||||
}
|
|
||||||
got, err := os.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("file should exist at %s: %v", path, err)
|
|
||||||
}
|
|
||||||
if !bytes.Equal(got, imageData) {
|
|
||||||
t.Errorf("content mismatch: got len=%d, want len=%d", len(got), len(imageData))
|
|
||||||
}
|
|
||||||
|
|
||||||
// ReleaseAll must delete the file (store owns lifecycle).
|
|
||||||
scope := channels.BuildMediaScope("wecom_aibot", "chat1", "msg1")
|
|
||||||
if err := store.ReleaseAll(scope); err != nil {
|
|
||||||
t.Fatalf("ReleaseAll failed: %v", err)
|
|
||||||
}
|
|
||||||
if _, err := os.Stat(path); !os.IsNotExist(err) {
|
|
||||||
t.Errorf("file should have been deleted by ReleaseAll, stat err: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestStoreWSMedia_MultipleMessages verifies that concurrent media messages with
|
|
||||||
// different msgIDs do not collide and each resolve to distinct files.
|
|
||||||
func TestStoreWSMedia_MultipleMessages(t *testing.T) {
|
|
||||||
imageA := bytes.Repeat([]byte("a"), 64)
|
|
||||||
imageB := bytes.Repeat([]byte("b"), 64)
|
|
||||||
|
|
||||||
srvA := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
_, _ = w.Write(imageA)
|
|
||||||
}))
|
|
||||||
defer srvA.Close()
|
|
||||||
srvB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
_, _ = w.Write(imageB)
|
|
||||||
}))
|
|
||||||
defer srvB.Close()
|
|
||||||
|
|
||||||
ch := newTestWSChannel(t)
|
|
||||||
store := media.NewFileMediaStore()
|
|
||||||
ch.SetMediaStore(store)
|
|
||||||
|
|
||||||
refA, err := ch.storeWSMedia(context.Background(), "chat1", "msgA", srvA.URL, "", ".jpg")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("storeWSMedia A: %v", err)
|
|
||||||
}
|
|
||||||
refB, err := ch.storeWSMedia(context.Background(), "chat1", "msgB", srvB.URL, "", ".jpg")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("storeWSMedia B: %v", err)
|
|
||||||
}
|
|
||||||
if refA == refB {
|
|
||||||
t.Fatal("distinct messages must produce distinct refs")
|
|
||||||
}
|
|
||||||
|
|
||||||
pathA, _ := store.Resolve(refA)
|
|
||||||
pathB, _ := store.Resolve(refB)
|
|
||||||
if pathA == pathB {
|
|
||||||
t.Fatal("distinct messages must be stored at distinct paths")
|
|
||||||
}
|
|
||||||
|
|
||||||
gotA, _ := os.ReadFile(pathA)
|
|
||||||
gotB, _ := os.ReadFile(pathB)
|
|
||||||
if !bytes.Equal(gotA, imageA) {
|
|
||||||
t.Errorf("content mismatch for message A")
|
|
||||||
}
|
|
||||||
if !bytes.Equal(gotB, imageB) {
|
|
||||||
t.Errorf("content mismatch for message B")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestStoreWSMedia_ContentTypeExt verifies that the file extension is inferred
|
|
||||||
// from the HTTP Content-Type header and the defaultExt fallback is used when the
|
|
||||||
// type is absent or unrecognized.
|
|
||||||
func TestStoreWSMedia_ContentTypeExt(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
contentType string
|
|
||||||
wantExt string
|
|
||||||
}{
|
|
||||||
{"image/jpeg", ".jpg"},
|
|
||||||
{"image/png", ".png"},
|
|
||||||
{"video/mp4", ".mp4"},
|
|
||||||
{"application/pdf", ".pdf"},
|
|
||||||
{"application/zip", ".zip"},
|
|
||||||
// With parameters stripped.
|
|
||||||
{"video/mp4; codecs=avc1", ".mp4"},
|
|
||||||
// Unknown type → falls back to defaultExt.
|
|
||||||
{"", ""},
|
|
||||||
{"application/octet-stream", ""},
|
|
||||||
}
|
|
||||||
for _, tc := range tests {
|
|
||||||
got := wsMediaExtFromContentType(tc.contentType)
|
|
||||||
if got != tc.wantExt {
|
|
||||||
t.Errorf("wsMediaExtFromContentType(%q) = %q, want %q", tc.contentType, got, tc.wantExt)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// End-to-end: server returns Content-Type: video/mp4, defaultExt is .bin.
|
|
||||||
// The stored file should carry the .mp4 extension, not .bin.
|
|
||||||
payload := bytes.Repeat([]byte("v"), 128)
|
|
||||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "video/mp4")
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
_, _ = w.Write(payload)
|
|
||||||
}))
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
ch := newTestWSChannel(t)
|
|
||||||
store := media.NewFileMediaStore()
|
|
||||||
ch.SetMediaStore(store)
|
|
||||||
|
|
||||||
ref, err := ch.storeWSMedia(context.Background(), "chat1", "vid1", srv.URL, "", ".bin")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("storeWSMedia: %v", err)
|
|
||||||
}
|
|
||||||
path, err := store.Resolve(ref)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("resolve: %v", err)
|
|
||||||
}
|
|
||||||
if ext := path[len(path)-4:]; ext != ".mp4" {
|
|
||||||
t.Errorf("expected .mp4 extension from Content-Type, got %q", ext)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSplitWSContent verifies byte-aware splitting of stream content.
|
|
||||||
func TestSplitWSContent(t *testing.T) {
|
|
||||||
t.Run("short content is not split", func(t *testing.T) {
|
|
||||||
chunks := splitWSContent("hello", 20480)
|
|
||||||
if len(chunks) != 1 || chunks[0] != "hello" {
|
|
||||||
t.Fatalf("unexpected chunks: %v", chunks)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("ASCII content split at byte boundary", func(t *testing.T) {
|
|
||||||
// Build a string just over the limit.
|
|
||||||
content := strings.Repeat("a", 20481)
|
|
||||||
chunks := splitWSContent(content, 20480)
|
|
||||||
if len(chunks) < 2 {
|
|
||||||
t.Fatalf("expected >= 2 chunks, got %d", len(chunks))
|
|
||||||
}
|
|
||||||
for i, c := range chunks {
|
|
||||||
if len(c) > 20480 {
|
|
||||||
t.Errorf("chunk %d has %d bytes, want <= 20480", i, len(c))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// Reassembled content must equal the original (possibly without leading
|
|
||||||
// whitespace that splitWSContent trims between chunks).
|
|
||||||
joined := strings.Join(chunks, "")
|
|
||||||
if len(joined) < len(content)-len(chunks) {
|
|
||||||
t.Errorf("joined length %d too short (original %d)", len(joined), len(content))
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("CJK content split within byte limit", func(t *testing.T) {
|
|
||||||
// Each CJK rune is 3 bytes in UTF-8.
|
|
||||||
// 7000 CJK chars = 21000 bytes, which exceeds 20480.
|
|
||||||
content := strings.Repeat("\u4e2d", 7000)
|
|
||||||
chunks := splitWSContent(content, 20480)
|
|
||||||
if len(chunks) < 2 {
|
|
||||||
t.Fatalf("expected >= 2 chunks for 21000-byte CJK content, got %d", len(chunks))
|
|
||||||
}
|
|
||||||
for i, c := range chunks {
|
|
||||||
if len(c) > 20480 {
|
|
||||||
t.Errorf("chunk %d has %d bytes, want <= 20480", i, len(c))
|
|
||||||
}
|
|
||||||
// Every chunk must be valid UTF-8.
|
|
||||||
if !strings.ContainsRune(c, '\u4e2d') && len(c) > 0 {
|
|
||||||
// quick plausibility check — content was pure CJK
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSplitAtByteBoundary verifies the last-resort byte-boundary splitter.
|
|
||||||
func TestSplitAtByteBoundary(t *testing.T) {
|
|
||||||
t.Run("ASCII fits in one chunk", func(t *testing.T) {
|
|
||||||
parts := splitAtByteBoundary("hello world", 100)
|
|
||||||
if len(parts) != 1 {
|
|
||||||
t.Fatalf("expected 1 part, got %d", len(parts))
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("splits at byte boundary, never mid-rune", func(t *testing.T) {
|
|
||||||
// 10 CJK characters = 30 bytes; split at 20 bytes.
|
|
||||||
s := strings.Repeat("\u6587", 10) // 10 × 3 bytes = 30 bytes
|
|
||||||
parts := splitAtByteBoundary(s, 20)
|
|
||||||
for i, p := range parts {
|
|
||||||
if len(p) > 20 {
|
|
||||||
t.Errorf("part %d has %d bytes, want <= 20", i, len(p))
|
|
||||||
}
|
|
||||||
// Must be valid UTF-8 (no torn multi-byte sequences).
|
|
||||||
for j, r := range p {
|
|
||||||
if r == '\uFFFD' {
|
|
||||||
t.Errorf("part %d has replacement rune at position %d: torn UTF-8", i, j)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
@ -1,756 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"encoding/xml"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"mime/multipart"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/identity"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
wecomAPIBase = "https://qyapi.weixin.qq.com"
|
|
||||||
)
|
|
||||||
|
|
||||||
// WeComAppChannel implements the Channel interface for WeCom App (企业微信自建应用)
|
|
||||||
type WeComAppChannel struct {
|
|
||||||
*channels.BaseChannel
|
|
||||||
config config.WeComAppConfig
|
|
||||||
client *http.Client
|
|
||||||
accessToken string
|
|
||||||
tokenExpiry time.Time
|
|
||||||
tokenMu sync.RWMutex
|
|
||||||
ctx context.Context
|
|
||||||
cancel context.CancelFunc
|
|
||||||
processedMsgs *MessageDeduplicator
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComXMLMessage represents the XML message structure from WeCom
|
|
||||||
type WeComXMLMessage struct {
|
|
||||||
XMLName xml.Name `xml:"xml"`
|
|
||||||
ToUserName string `xml:"ToUserName"`
|
|
||||||
FromUserName string `xml:"FromUserName"`
|
|
||||||
CreateTime int64 `xml:"CreateTime"`
|
|
||||||
MsgType string `xml:"MsgType"`
|
|
||||||
Content string `xml:"Content"`
|
|
||||||
MsgId int64 `xml:"MsgId"`
|
|
||||||
AgentID int64 `xml:"AgentID"`
|
|
||||||
PicUrl string `xml:"PicUrl"`
|
|
||||||
MediaId string `xml:"MediaId"`
|
|
||||||
Format string `xml:"Format"`
|
|
||||||
ThumbMediaId string `xml:"ThumbMediaId"`
|
|
||||||
LocationX float64 `xml:"Location_X"`
|
|
||||||
LocationY float64 `xml:"Location_Y"`
|
|
||||||
Scale int `xml:"Scale"`
|
|
||||||
Label string `xml:"Label"`
|
|
||||||
Title string `xml:"Title"`
|
|
||||||
Description string `xml:"Description"`
|
|
||||||
Url string `xml:"Url"`
|
|
||||||
Event string `xml:"Event"`
|
|
||||||
EventKey string `xml:"EventKey"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComTextMessage represents text message for sending
|
|
||||||
type WeComTextMessage struct {
|
|
||||||
ToUser string `json:"touser"`
|
|
||||||
MsgType string `json:"msgtype"`
|
|
||||||
AgentID int64 `json:"agentid"`
|
|
||||||
Text struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"text"`
|
|
||||||
Safe int `json:"safe,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComMarkdownMessage represents markdown message for sending
|
|
||||||
type WeComMarkdownMessage struct {
|
|
||||||
ToUser string `json:"touser"`
|
|
||||||
MsgType string `json:"msgtype"`
|
|
||||||
AgentID int64 `json:"agentid"`
|
|
||||||
Markdown struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"markdown"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComImageMessage represents image message for sending
|
|
||||||
type WeComImageMessage struct {
|
|
||||||
ToUser string `json:"touser"`
|
|
||||||
MsgType string `json:"msgtype"`
|
|
||||||
AgentID int64 `json:"agentid"`
|
|
||||||
Image struct {
|
|
||||||
MediaID string `json:"media_id"`
|
|
||||||
} `json:"image"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComAccessTokenResponse represents the access token API response
|
|
||||||
type WeComAccessTokenResponse struct {
|
|
||||||
ErrCode int `json:"errcode"`
|
|
||||||
ErrMsg string `json:"errmsg"`
|
|
||||||
AccessToken string `json:"access_token"`
|
|
||||||
ExpiresIn int `json:"expires_in"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComSendMessageResponse represents the send message API response
|
|
||||||
type WeComSendMessageResponse struct {
|
|
||||||
ErrCode int `json:"errcode"`
|
|
||||||
ErrMsg string `json:"errmsg"`
|
|
||||||
InvalidUser string `json:"invaliduser"`
|
|
||||||
InvalidParty string `json:"invalidparty"`
|
|
||||||
InvalidTag string `json:"invalidtag"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// PKCS7Padding adds PKCS7 padding
|
|
||||||
type PKCS7Padding struct{}
|
|
||||||
|
|
||||||
// NewWeComAppChannel creates a new WeCom App channel instance
|
|
||||||
func NewWeComAppChannel(cfg config.WeComAppConfig, messageBus *bus.MessageBus) (*WeComAppChannel, error) {
|
|
||||||
if cfg.CorpID == "" || cfg.CorpSecret() == "" || cfg.AgentID == 0 {
|
|
||||||
return nil, fmt.Errorf("wecom_app corp_id, corp_secret and agent_id are required")
|
|
||||||
}
|
|
||||||
|
|
||||||
base := channels.NewBaseChannel("wecom_app", cfg, messageBus, cfg.AllowFrom,
|
|
||||||
channels.WithMaxMessageLength(2048),
|
|
||||||
channels.WithGroupTrigger(cfg.GroupTrigger),
|
|
||||||
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
|
||||||
)
|
|
||||||
|
|
||||||
// Client timeout must be >= the configured ReplyTimeout so the
|
|
||||||
// per-request context deadline is always the effective limit.
|
|
||||||
clientTimeout := 30 * time.Second
|
|
||||||
if d := time.Duration(cfg.ReplyTimeout) * time.Second; d > clientTimeout {
|
|
||||||
clientTimeout = d
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
return &WeComAppChannel{
|
|
||||||
BaseChannel: base,
|
|
||||||
config: cfg,
|
|
||||||
client: &http.Client{Timeout: clientTimeout},
|
|
||||||
ctx: ctx,
|
|
||||||
cancel: cancel,
|
|
||||||
processedMsgs: NewMessageDeduplicator(wecomMaxProcessedMessages),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Name returns the channel name
|
|
||||||
func (c *WeComAppChannel) Name() string {
|
|
||||||
return "wecom_app"
|
|
||||||
}
|
|
||||||
|
|
||||||
// Start initializes the WeCom App channel
|
|
||||||
func (c *WeComAppChannel) Start(ctx context.Context) error {
|
|
||||||
logger.InfoC("wecom_app", "Starting WeCom App channel...")
|
|
||||||
|
|
||||||
// Cancel the context created in the constructor to avoid a resource leak.
|
|
||||||
if c.cancel != nil {
|
|
||||||
c.cancel()
|
|
||||||
}
|
|
||||||
c.ctx, c.cancel = context.WithCancel(ctx)
|
|
||||||
|
|
||||||
// Get initial access token
|
|
||||||
if err := c.refreshAccessToken(); err != nil {
|
|
||||||
logger.WarnCF("wecom_app", "Failed to get initial access token", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Start token refresh goroutine
|
|
||||||
go c.tokenRefreshLoop()
|
|
||||||
|
|
||||||
c.SetRunning(true)
|
|
||||||
logger.InfoC("wecom_app", "WeCom App channel started")
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stop gracefully stops the WeCom App channel
|
|
||||||
func (c *WeComAppChannel) Stop(ctx context.Context) error {
|
|
||||||
logger.InfoC("wecom_app", "Stopping WeCom App channel...")
|
|
||||||
|
|
||||||
if c.cancel != nil {
|
|
||||||
c.cancel()
|
|
||||||
}
|
|
||||||
|
|
||||||
c.SetRunning(false)
|
|
||||||
logger.InfoC("wecom_app", "WeCom App channel stopped")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Send sends a message to WeCom user proactively using access token
|
|
||||||
func (c *WeComAppChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
|
||||||
if !c.IsRunning() {
|
|
||||||
return channels.ErrNotRunning
|
|
||||||
}
|
|
||||||
|
|
||||||
accessToken := c.getAccessToken()
|
|
||||||
if accessToken == "" {
|
|
||||||
return fmt.Errorf("no valid access token available")
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.DebugCF("wecom_app", "Sending message", map[string]any{
|
|
||||||
"chat_id": msg.ChatID,
|
|
||||||
"preview": utils.Truncate(msg.Content, 100),
|
|
||||||
})
|
|
||||||
|
|
||||||
return c.sendTextMessage(ctx, accessToken, msg.ChatID, msg.Content)
|
|
||||||
}
|
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
|
||||||
func (c *WeComAppChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
|
||||||
if !c.IsRunning() {
|
|
||||||
return channels.ErrNotRunning
|
|
||||||
}
|
|
||||||
|
|
||||||
accessToken := c.getAccessToken()
|
|
||||||
if accessToken == "" {
|
|
||||||
return fmt.Errorf("no valid access token available: %w", channels.ErrTemporary)
|
|
||||||
}
|
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
|
||||||
if store == nil {
|
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, part := range msg.Parts {
|
|
||||||
localPath, err := store.Resolve(part.Ref)
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to resolve media ref", map[string]any{
|
|
||||||
"ref": part.Ref,
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Map part type to WeCom media type
|
|
||||||
var mediaType string
|
|
||||||
switch part.Type {
|
|
||||||
case "image":
|
|
||||||
mediaType = "image"
|
|
||||||
case "audio":
|
|
||||||
mediaType = "voice"
|
|
||||||
case "video":
|
|
||||||
mediaType = "video"
|
|
||||||
default:
|
|
||||||
mediaType = "file"
|
|
||||||
}
|
|
||||||
|
|
||||||
// Upload media to get media_id
|
|
||||||
mediaID, err := c.uploadMedia(ctx, accessToken, mediaType, localPath)
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to upload media", map[string]any{
|
|
||||||
"type": mediaType,
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
// Fallback: send caption as text
|
|
||||||
if part.Caption != "" {
|
|
||||||
_ = c.sendTextMessage(ctx, accessToken, msg.ChatID, part.Caption)
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Send media message using the media_id
|
|
||||||
if mediaType == "image" {
|
|
||||||
err = c.sendImageMessage(ctx, accessToken, msg.ChatID, mediaID)
|
|
||||||
} else {
|
|
||||||
// For non-image types, send as text fallback with caption
|
|
||||||
caption := part.Caption
|
|
||||||
if caption == "" {
|
|
||||||
caption = fmt.Sprintf("[%s: %s]", part.Type, part.Filename)
|
|
||||||
}
|
|
||||||
err = c.sendTextMessage(ctx, accessToken, msg.ChatID, caption)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// uploadMedia uploads a local file to WeCom temporary media storage.
|
|
||||||
func (c *WeComAppChannel) uploadMedia(ctx context.Context, accessToken, mediaType, localPath string) (string, error) {
|
|
||||||
apiURL := fmt.Sprintf("%s/cgi-bin/media/upload?access_token=%s&type=%s",
|
|
||||||
wecomAPIBase, url.QueryEscape(accessToken), url.QueryEscape(mediaType))
|
|
||||||
|
|
||||||
file, err := os.Open(localPath)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to open file: %w", err)
|
|
||||||
}
|
|
||||||
defer file.Close()
|
|
||||||
|
|
||||||
body := &bytes.Buffer{}
|
|
||||||
writer := multipart.NewWriter(body)
|
|
||||||
|
|
||||||
filename := filepath.Base(localPath)
|
|
||||||
formFile, err := writer.CreateFormFile("media", filename)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to create form file: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err = io.Copy(formFile, file); err != nil {
|
|
||||||
return "", fmt.Errorf("failed to copy file content: %w", err)
|
|
||||||
}
|
|
||||||
writer.Close()
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, body)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
|
||||||
|
|
||||||
resp, err := c.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", channels.ClassifyNetError(err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
respBody, readErr := io.ReadAll(resp.Body)
|
|
||||||
if readErr != nil {
|
|
||||||
return "", channels.ClassifySendError(
|
|
||||||
resp.StatusCode,
|
|
||||||
fmt.Errorf("reading wecom upload error response: %w", readErr),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return "", channels.ClassifySendError(
|
|
||||||
resp.StatusCode,
|
|
||||||
fmt.Errorf("wecom upload error: %s", string(respBody)),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
var result struct {
|
|
||||||
ErrCode int `json:"errcode"`
|
|
||||||
ErrMsg string `json:"errmsg"`
|
|
||||||
MediaID string `json:"media_id"`
|
|
||||||
}
|
|
||||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
|
||||||
return "", fmt.Errorf("failed to parse upload response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.ErrCode != 0 {
|
|
||||||
return "", fmt.Errorf("upload API error: %s (code: %d)", result.ErrMsg, result.ErrCode)
|
|
||||||
}
|
|
||||||
|
|
||||||
return result.MediaID, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendWeComMessage marshals payload and POSTs it to the WeCom message API.
|
|
||||||
func (c *WeComAppChannel) sendWeComMessage(ctx context.Context, accessToken string, payload any) error {
|
|
||||||
apiURL := fmt.Sprintf("%s/cgi-bin/message/send?access_token=%s", wecomAPIBase, accessToken)
|
|
||||||
|
|
||||||
jsonData, err := json.Marshal(payload)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to marshal message: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
timeout := c.config.ReplyTimeout
|
|
||||||
if timeout <= 0 {
|
|
||||||
timeout = 5
|
|
||||||
}
|
|
||||||
|
|
||||||
reqCtx, cancel := context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
|
|
||||||
resp, err := c.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return channels.ClassifyNetError(err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
respBody, readErr := io.ReadAll(resp.Body)
|
|
||||||
if readErr != nil {
|
|
||||||
return channels.ClassifySendError(
|
|
||||||
resp.StatusCode,
|
|
||||||
fmt.Errorf("reading wecom_app error response: %w", readErr),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return channels.ClassifySendError(
|
|
||||||
resp.StatusCode,
|
|
||||||
fmt.Errorf("wecom_app API error: %s", string(respBody)),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
respBody, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var sendResp WeComSendMessageResponse
|
|
||||||
if err := json.Unmarshal(respBody, &sendResp); err != nil {
|
|
||||||
return fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if sendResp.ErrCode != 0 {
|
|
||||||
return fmt.Errorf("API error: %s (code: %d)", sendResp.ErrMsg, sendResp.ErrCode)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendImageMessage sends an image message using a media_id.
|
|
||||||
func (c *WeComAppChannel) sendImageMessage(ctx context.Context, accessToken, userID, mediaID string) error {
|
|
||||||
msg := WeComImageMessage{
|
|
||||||
ToUser: userID,
|
|
||||||
MsgType: "image",
|
|
||||||
AgentID: c.config.AgentID,
|
|
||||||
}
|
|
||||||
msg.Image.MediaID = mediaID
|
|
||||||
return c.sendWeComMessage(ctx, accessToken, msg)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WebhookPath returns the path for registering on the shared HTTP server.
|
|
||||||
func (c *WeComAppChannel) WebhookPath() string {
|
|
||||||
if c.config.WebhookPath != "" {
|
|
||||||
return c.config.WebhookPath
|
|
||||||
}
|
|
||||||
return "/webhook/wecom-app"
|
|
||||||
}
|
|
||||||
|
|
||||||
// ServeHTTP implements http.Handler for the shared HTTP server.
|
|
||||||
func (c *WeComAppChannel) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
||||||
c.handleWebhook(w, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
// HealthPath returns the health check endpoint path.
|
|
||||||
func (c *WeComAppChannel) HealthPath() string {
|
|
||||||
return "/health/wecom-app"
|
|
||||||
}
|
|
||||||
|
|
||||||
// HealthHandler handles health check requests.
|
|
||||||
func (c *WeComAppChannel) HealthHandler(w http.ResponseWriter, r *http.Request) {
|
|
||||||
c.handleHealth(w, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleWebhook handles incoming webhook requests from WeCom
|
|
||||||
func (c *WeComAppChannel) handleWebhook(w http.ResponseWriter, r *http.Request) {
|
|
||||||
ctx := r.Context()
|
|
||||||
|
|
||||||
// Log all incoming requests for debugging
|
|
||||||
logger.DebugCF("wecom_app", "Received webhook request", map[string]any{
|
|
||||||
"method": r.Method,
|
|
||||||
"url": r.URL.String(),
|
|
||||||
"path": r.URL.Path,
|
|
||||||
"query": r.URL.RawQuery,
|
|
||||||
})
|
|
||||||
|
|
||||||
if r.Method == http.MethodGet {
|
|
||||||
// Handle verification request
|
|
||||||
c.handleVerification(ctx, w, r)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if r.Method == http.MethodPost {
|
|
||||||
// Handle message callback
|
|
||||||
c.handleMessageCallback(ctx, w, r)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.WarnCF("wecom_app", "Method not allowed", map[string]any{
|
|
||||||
"method": r.Method,
|
|
||||||
})
|
|
||||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleVerification handles the URL verification request from WeCom
|
|
||||||
func (c *WeComAppChannel) handleVerification(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
|
||||||
query := r.URL.Query()
|
|
||||||
msgSignature := query.Get("msg_signature")
|
|
||||||
timestamp := query.Get("timestamp")
|
|
||||||
nonce := query.Get("nonce")
|
|
||||||
echostr := query.Get("echostr")
|
|
||||||
|
|
||||||
logger.DebugCF("wecom_app", "Handling verification request", map[string]any{
|
|
||||||
"msg_signature": msgSignature,
|
|
||||||
"timestamp": timestamp,
|
|
||||||
"nonce": nonce,
|
|
||||||
"echostr": echostr,
|
|
||||||
"corp_id": c.config.CorpID,
|
|
||||||
})
|
|
||||||
|
|
||||||
if msgSignature == "" || timestamp == "" || nonce == "" || echostr == "" {
|
|
||||||
logger.ErrorC("wecom_app", "Missing parameters in verification request")
|
|
||||||
http.Error(w, "Missing parameters", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify signature
|
|
||||||
if !verifySignature(c.config.Token(), msgSignature, timestamp, nonce, echostr) {
|
|
||||||
logger.WarnCF("wecom_app", "Signature verification failed", map[string]any{
|
|
||||||
"token": c.config.Token(),
|
|
||||||
"msg_signature": msgSignature,
|
|
||||||
"timestamp": timestamp,
|
|
||||||
"nonce": nonce,
|
|
||||||
})
|
|
||||||
http.Error(w, "Invalid signature", http.StatusForbidden)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.DebugC("wecom_app", "Signature verification passed")
|
|
||||||
|
|
||||||
// Decrypt echostr with CorpID verification
|
|
||||||
// For WeCom App (自建应用), receiveid should be corp_id
|
|
||||||
logger.DebugCF("wecom_app", "Attempting to decrypt echostr", map[string]any{
|
|
||||||
"encoding_aes_key": c.config.EncodingAESKey(),
|
|
||||||
"corp_id": c.config.CorpID,
|
|
||||||
})
|
|
||||||
decryptedEchoStr, err := decryptMessageWithVerify(echostr, c.config.EncodingAESKey(), c.config.CorpID)
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to decrypt echostr", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
"encoding_aes_key": c.config.EncodingAESKey,
|
|
||||||
"corp_id": c.config.CorpID,
|
|
||||||
})
|
|
||||||
http.Error(w, "Decryption failed", http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.DebugCF("wecom_app", "Successfully decrypted echostr", map[string]any{
|
|
||||||
"decrypted": decryptedEchoStr,
|
|
||||||
})
|
|
||||||
|
|
||||||
// Remove BOM and whitespace as per WeCom documentation
|
|
||||||
// The response must be plain text without quotes, BOM, or newlines
|
|
||||||
decryptedEchoStr = strings.TrimSpace(decryptedEchoStr)
|
|
||||||
decryptedEchoStr = strings.TrimPrefix(decryptedEchoStr, "\xef\xbb\xbf") // Remove UTF-8 BOM
|
|
||||||
w.Write([]byte(decryptedEchoStr))
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleMessageCallback handles incoming messages from WeCom
|
|
||||||
func (c *WeComAppChannel) handleMessageCallback(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
|
||||||
query := r.URL.Query()
|
|
||||||
msgSignature := query.Get("msg_signature")
|
|
||||||
timestamp := query.Get("timestamp")
|
|
||||||
nonce := query.Get("nonce")
|
|
||||||
|
|
||||||
if msgSignature == "" || timestamp == "" || nonce == "" {
|
|
||||||
http.Error(w, "Missing parameters", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read request body
|
|
||||||
body, err := io.ReadAll(r.Body)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, "Failed to read body", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
defer r.Body.Close()
|
|
||||||
|
|
||||||
// Parse XML to get encrypted message
|
|
||||||
var encryptedMsg struct {
|
|
||||||
XMLName xml.Name `xml:"xml"`
|
|
||||||
ToUserName string `xml:"ToUserName"`
|
|
||||||
Encrypt string `xml:"Encrypt"`
|
|
||||||
AgentID string `xml:"AgentID"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err = xml.Unmarshal(body, &encryptedMsg); err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to parse XML", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Invalid XML", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify signature
|
|
||||||
if !verifySignature(c.config.Token(), msgSignature, timestamp, nonce, encryptedMsg.Encrypt) {
|
|
||||||
logger.WarnC("wecom_app", "Message signature verification failed")
|
|
||||||
http.Error(w, "Invalid signature", http.StatusForbidden)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Decrypt message with CorpID verification
|
|
||||||
// For WeCom App (自建应用), receiveid should be corp_id
|
|
||||||
decryptedMsg, err := decryptMessageWithVerify(encryptedMsg.Encrypt, c.config.EncodingAESKey(), c.config.CorpID)
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to decrypt message", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Decryption failed", http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Parse decrypted XML message
|
|
||||||
var msg WeComXMLMessage
|
|
||||||
if err := xml.Unmarshal([]byte(decryptedMsg), &msg); err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to parse decrypted message", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Invalid message format", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Process the message with the channel's long-lived context (not the HTTP
|
|
||||||
// request context, which is canceled as soon as we return the response).
|
|
||||||
go c.processMessage(c.ctx, msg)
|
|
||||||
|
|
||||||
// Return success response immediately
|
|
||||||
// WeCom App requires response within configured timeout (default 5 seconds)
|
|
||||||
w.Write([]byte("success"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// processMessage processes the received message
|
|
||||||
func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessage) {
|
|
||||||
// Skip non-text messages for now (can be extended)
|
|
||||||
if msg.MsgType != "text" && msg.MsgType != "image" && msg.MsgType != "voice" {
|
|
||||||
logger.DebugCF("wecom_app", "Skipping non-supported message type", map[string]any{
|
|
||||||
"msg_type": msg.MsgType,
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Message deduplication: Use msg_id to prevent duplicate processing
|
|
||||||
// As per WeCom documentation, use msg_id for deduplication
|
|
||||||
msgID := fmt.Sprintf("%d", msg.MsgId)
|
|
||||||
if !c.processedMsgs.MarkMessageProcessed(msgID) {
|
|
||||||
logger.DebugCF("wecom_app", "Skipping duplicate message", map[string]any{
|
|
||||||
"msg_id": msgID,
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
senderID := msg.FromUserName
|
|
||||||
chatID := senderID // WeCom App uses user ID as chat ID for direct messages
|
|
||||||
|
|
||||||
// Build metadata
|
|
||||||
// WeCom App only supports direct messages (private chat)
|
|
||||||
peer := bus.Peer{Kind: "direct", ID: senderID}
|
|
||||||
messageID := fmt.Sprintf("%d", msg.MsgId)
|
|
||||||
|
|
||||||
metadata := map[string]string{
|
|
||||||
"msg_type": msg.MsgType,
|
|
||||||
"msg_id": fmt.Sprintf("%d", msg.MsgId),
|
|
||||||
"agent_id": fmt.Sprintf("%d", msg.AgentID),
|
|
||||||
"platform": "wecom_app",
|
|
||||||
"media_id": msg.MediaId,
|
|
||||||
"create_time": fmt.Sprintf("%d", msg.CreateTime),
|
|
||||||
}
|
|
||||||
|
|
||||||
content := msg.Content
|
|
||||||
|
|
||||||
logger.DebugCF("wecom_app", "Received message", map[string]any{
|
|
||||||
"sender_id": senderID,
|
|
||||||
"msg_type": msg.MsgType,
|
|
||||||
"preview": utils.Truncate(content, 50),
|
|
||||||
})
|
|
||||||
|
|
||||||
// Build sender info
|
|
||||||
appSender := bus.SenderInfo{
|
|
||||||
Platform: "wecom",
|
|
||||||
PlatformID: senderID,
|
|
||||||
CanonicalID: identity.BuildCanonicalID("wecom", senderID),
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handle the message through the base channel
|
|
||||||
c.HandleMessage(ctx, peer, messageID, senderID, chatID, content, nil, metadata, appSender)
|
|
||||||
}
|
|
||||||
|
|
||||||
// tokenRefreshLoop periodically refreshes the access token
|
|
||||||
func (c *WeComAppChannel) tokenRefreshLoop() {
|
|
||||||
ticker := time.NewTicker(5 * time.Minute)
|
|
||||||
defer ticker.Stop()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-c.ctx.Done():
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
if err := c.refreshAccessToken(); err != nil {
|
|
||||||
logger.ErrorCF("wecom_app", "Failed to refresh access token", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// refreshAccessToken gets a new access token from WeCom API
|
|
||||||
func (c *WeComAppChannel) refreshAccessToken() error {
|
|
||||||
apiURL := fmt.Sprintf("%s/cgi-bin/gettoken?corpid=%s&corpsecret=%s",
|
|
||||||
wecomAPIBase, url.QueryEscape(c.config.CorpID), url.QueryEscape(c.config.CorpSecret()))
|
|
||||||
|
|
||||||
resp, err := http.Get(apiURL)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to request access token: %w", err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var tokenResp WeComAccessTokenResponse
|
|
||||||
if err := json.Unmarshal(body, &tokenResp); err != nil {
|
|
||||||
return fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if tokenResp.ErrCode != 0 {
|
|
||||||
return fmt.Errorf("API error: %s (code: %d)", tokenResp.ErrMsg, tokenResp.ErrCode)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.tokenMu.Lock()
|
|
||||||
c.accessToken = tokenResp.AccessToken
|
|
||||||
c.tokenExpiry = time.Now().Add(time.Duration(tokenResp.ExpiresIn-300) * time.Second) // Refresh 5 minutes early
|
|
||||||
c.tokenMu.Unlock()
|
|
||||||
|
|
||||||
logger.DebugC("wecom_app", "Access token refreshed successfully")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// getAccessToken returns the current valid access token
|
|
||||||
func (c *WeComAppChannel) getAccessToken() string {
|
|
||||||
c.tokenMu.RLock()
|
|
||||||
defer c.tokenMu.RUnlock()
|
|
||||||
|
|
||||||
if time.Now().After(c.tokenExpiry) {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.accessToken
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendTextMessage sends a text message to a user.
|
|
||||||
func (c *WeComAppChannel) sendTextMessage(ctx context.Context, accessToken, userID, content string) error {
|
|
||||||
msg := WeComTextMessage{
|
|
||||||
ToUser: userID,
|
|
||||||
MsgType: "text",
|
|
||||||
AgentID: c.config.AgentID,
|
|
||||||
}
|
|
||||||
msg.Text.Content = content
|
|
||||||
return c.sendWeComMessage(ctx, accessToken, msg)
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleHealth handles health check requests
|
|
||||||
func (c *WeComAppChannel) handleHealth(w http.ResponseWriter, r *http.Request) {
|
|
||||||
status := map[string]any{
|
|
||||||
"status": "ok",
|
|
||||||
"running": c.IsRunning(),
|
|
||||||
"has_token": c.getAccessToken() != "",
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(status)
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,499 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"encoding/xml"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/channels"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/identity"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
|
||||||
)
|
|
||||||
|
|
||||||
// WeComBotChannel implements the Channel interface for WeCom Bot (企业微信智能机器人)
|
|
||||||
// Uses webhook callback mode - simpler than WeCom App but only supports passive replies
|
|
||||||
type WeComBotChannel struct {
|
|
||||||
*channels.BaseChannel
|
|
||||||
config config.WeComConfig
|
|
||||||
client *http.Client
|
|
||||||
ctx context.Context
|
|
||||||
cancel context.CancelFunc
|
|
||||||
processedMsgs *MessageDeduplicator
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComBotMessage represents the JSON message structure from WeCom Bot (AIBOT)
|
|
||||||
type WeComBotMessage struct {
|
|
||||||
MsgID string `json:"msgid"`
|
|
||||||
AIBotID string `json:"aibotid"`
|
|
||||||
ChatID string `json:"chatid"` // Session ID, only present for group chats
|
|
||||||
ChatType string `json:"chattype"` // "single" for DM, "group" for group chat
|
|
||||||
From struct {
|
|
||||||
UserID string `json:"userid"`
|
|
||||||
} `json:"from"`
|
|
||||||
ResponseURL string `json:"response_url"`
|
|
||||||
MsgType string `json:"msgtype"` // text, image, voice, file, mixed
|
|
||||||
Text struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"text"`
|
|
||||||
Image struct {
|
|
||||||
URL string `json:"url"`
|
|
||||||
} `json:"image"`
|
|
||||||
Voice struct {
|
|
||||||
Content string `json:"content"` // Voice to text content
|
|
||||||
} `json:"voice"`
|
|
||||||
File struct {
|
|
||||||
URL string `json:"url"`
|
|
||||||
} `json:"file"`
|
|
||||||
Mixed struct {
|
|
||||||
MsgItem []struct {
|
|
||||||
MsgType string `json:"msgtype"`
|
|
||||||
Text struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"text"`
|
|
||||||
Image struct {
|
|
||||||
URL string `json:"url"`
|
|
||||||
} `json:"image"`
|
|
||||||
} `json:"msg_item"`
|
|
||||||
} `json:"mixed"`
|
|
||||||
Quote struct {
|
|
||||||
MsgType string `json:"msgtype"`
|
|
||||||
Text struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"text"`
|
|
||||||
} `json:"quote"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// WeComBotReplyMessage represents the reply message structure
|
|
||||||
type WeComBotReplyMessage struct {
|
|
||||||
MsgType string `json:"msgtype"`
|
|
||||||
Text struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"text,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewWeComBotChannel creates a new WeCom Bot channel instance
|
|
||||||
func NewWeComBotChannel(cfg config.WeComConfig, messageBus *bus.MessageBus) (*WeComBotChannel, error) {
|
|
||||||
if cfg.Token() == "" || cfg.WebhookURL == "" {
|
|
||||||
return nil, fmt.Errorf("wecom token and webhook_url are required")
|
|
||||||
}
|
|
||||||
|
|
||||||
base := channels.NewBaseChannel("wecom", cfg, messageBus, cfg.AllowFrom,
|
|
||||||
channels.WithMaxMessageLength(2048),
|
|
||||||
channels.WithGroupTrigger(cfg.GroupTrigger),
|
|
||||||
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
|
||||||
)
|
|
||||||
|
|
||||||
// Client timeout must be >= the configured ReplyTimeout so the
|
|
||||||
// per-request context deadline is always the effective limit.
|
|
||||||
clientTimeout := 30 * time.Second
|
|
||||||
if d := time.Duration(cfg.ReplyTimeout) * time.Second; d > clientTimeout {
|
|
||||||
clientTimeout = d
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
return &WeComBotChannel{
|
|
||||||
BaseChannel: base,
|
|
||||||
config: cfg,
|
|
||||||
client: &http.Client{Timeout: clientTimeout},
|
|
||||||
ctx: ctx,
|
|
||||||
cancel: cancel,
|
|
||||||
processedMsgs: NewMessageDeduplicator(wecomMaxProcessedMessages),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Name returns the channel name
|
|
||||||
func (c *WeComBotChannel) Name() string {
|
|
||||||
return "wecom"
|
|
||||||
}
|
|
||||||
|
|
||||||
// Start initializes the WeCom Bot channel
|
|
||||||
func (c *WeComBotChannel) Start(ctx context.Context) error {
|
|
||||||
logger.InfoC("wecom", "Starting WeCom Bot channel...")
|
|
||||||
|
|
||||||
// Cancel the context created in the constructor to avoid a resource leak.
|
|
||||||
if c.cancel != nil {
|
|
||||||
c.cancel()
|
|
||||||
}
|
|
||||||
c.ctx, c.cancel = context.WithCancel(ctx)
|
|
||||||
|
|
||||||
c.SetRunning(true)
|
|
||||||
logger.InfoC("wecom", "WeCom Bot channel started")
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stop gracefully stops the WeCom Bot channel
|
|
||||||
func (c *WeComBotChannel) Stop(ctx context.Context) error {
|
|
||||||
logger.InfoC("wecom", "Stopping WeCom Bot channel...")
|
|
||||||
|
|
||||||
if c.cancel != nil {
|
|
||||||
c.cancel()
|
|
||||||
}
|
|
||||||
|
|
||||||
c.SetRunning(false)
|
|
||||||
logger.InfoC("wecom", "WeCom Bot channel stopped")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Send sends a message to WeCom user via webhook API
|
|
||||||
// Note: WeCom Bot can only reply within the configured timeout (default 5 seconds) of receiving a message
|
|
||||||
// For delayed responses, we use the webhook URL
|
|
||||||
func (c *WeComBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
|
||||||
if !c.IsRunning() {
|
|
||||||
return channels.ErrNotRunning
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.DebugCF("wecom", "Sending message via webhook", map[string]any{
|
|
||||||
"chat_id": msg.ChatID,
|
|
||||||
"preview": utils.Truncate(msg.Content, 100),
|
|
||||||
})
|
|
||||||
|
|
||||||
return c.sendWebhookReply(ctx, msg.ChatID, msg.Content)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WebhookPath returns the path for registering on the shared HTTP server.
|
|
||||||
func (c *WeComBotChannel) WebhookPath() string {
|
|
||||||
if c.config.WebhookPath != "" {
|
|
||||||
return c.config.WebhookPath
|
|
||||||
}
|
|
||||||
return "/webhook/wecom"
|
|
||||||
}
|
|
||||||
|
|
||||||
// ServeHTTP implements http.Handler for the shared HTTP server.
|
|
||||||
func (c *WeComBotChannel) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
||||||
c.handleWebhook(w, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
// HealthPath returns the health check endpoint path.
|
|
||||||
func (c *WeComBotChannel) HealthPath() string {
|
|
||||||
return "/health/wecom"
|
|
||||||
}
|
|
||||||
|
|
||||||
// HealthHandler handles health check requests.
|
|
||||||
func (c *WeComBotChannel) HealthHandler(w http.ResponseWriter, r *http.Request) {
|
|
||||||
c.handleHealth(w, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleWebhook handles incoming webhook requests from WeCom
|
|
||||||
func (c *WeComBotChannel) handleWebhook(w http.ResponseWriter, r *http.Request) {
|
|
||||||
ctx := r.Context()
|
|
||||||
|
|
||||||
if r.Method == http.MethodGet {
|
|
||||||
// Handle verification request
|
|
||||||
c.handleVerification(ctx, w, r)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if r.Method == http.MethodPost {
|
|
||||||
// Handle message callback
|
|
||||||
c.handleMessageCallback(ctx, w, r)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleVerification handles the URL verification request from WeCom
|
|
||||||
func (c *WeComBotChannel) handleVerification(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
|
||||||
query := r.URL.Query()
|
|
||||||
msgSignature := query.Get("msg_signature")
|
|
||||||
timestamp := query.Get("timestamp")
|
|
||||||
nonce := query.Get("nonce")
|
|
||||||
echostr := query.Get("echostr")
|
|
||||||
|
|
||||||
if msgSignature == "" || timestamp == "" || nonce == "" || echostr == "" {
|
|
||||||
http.Error(w, "Missing parameters", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify signature
|
|
||||||
if !verifySignature(c.config.Token(), msgSignature, timestamp, nonce, echostr) {
|
|
||||||
logger.WarnC("wecom", "Signature verification failed")
|
|
||||||
http.Error(w, "Invalid signature", http.StatusForbidden)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Decrypt echostr
|
|
||||||
// For AIBOT (智能机器人), receiveid should be empty string ""
|
|
||||||
// Reference: https://developer.work.weixin.qq.com/document/path/101033
|
|
||||||
decryptedEchoStr, err := decryptMessageWithVerify(echostr, c.config.EncodingAESKey(), "")
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorCF("wecom", "Failed to decrypt echostr", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Decryption failed", http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Remove BOM and whitespace as per WeCom documentation
|
|
||||||
// The response must be plain text without quotes, BOM, or newlines
|
|
||||||
decryptedEchoStr = strings.TrimSpace(decryptedEchoStr)
|
|
||||||
decryptedEchoStr = strings.TrimPrefix(decryptedEchoStr, "\xef\xbb\xbf") // Remove UTF-8 BOM
|
|
||||||
w.Write([]byte(decryptedEchoStr))
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleMessageCallback handles incoming messages from WeCom
|
|
||||||
func (c *WeComBotChannel) handleMessageCallback(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
|
||||||
query := r.URL.Query()
|
|
||||||
msgSignature := query.Get("msg_signature")
|
|
||||||
timestamp := query.Get("timestamp")
|
|
||||||
nonce := query.Get("nonce")
|
|
||||||
|
|
||||||
if msgSignature == "" || timestamp == "" || nonce == "" {
|
|
||||||
http.Error(w, "Missing parameters", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read request body
|
|
||||||
body, err := io.ReadAll(r.Body)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, "Failed to read body", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
defer r.Body.Close()
|
|
||||||
|
|
||||||
// Parse XML to get encrypted message
|
|
||||||
var encryptedMsg struct {
|
|
||||||
XMLName xml.Name `xml:"xml"`
|
|
||||||
ToUserName string `xml:"ToUserName"`
|
|
||||||
Encrypt string `xml:"Encrypt"`
|
|
||||||
AgentID string `xml:"AgentID"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err = xml.Unmarshal(body, &encryptedMsg); err != nil {
|
|
||||||
logger.ErrorCF("wecom", "Failed to parse XML", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Invalid XML", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify signature
|
|
||||||
if !verifySignature(c.config.Token(), msgSignature, timestamp, nonce, encryptedMsg.Encrypt) {
|
|
||||||
logger.WarnC("wecom", "Message signature verification failed")
|
|
||||||
http.Error(w, "Invalid signature", http.StatusForbidden)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Decrypt message
|
|
||||||
// For AIBOT (智能机器人), receiveid should be empty string ""
|
|
||||||
// Reference: https://developer.work.weixin.qq.com/document/path/101033
|
|
||||||
decryptedMsg, err := decryptMessageWithVerify(encryptedMsg.Encrypt, c.config.EncodingAESKey(), "")
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorCF("wecom", "Failed to decrypt message", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Decryption failed", http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Parse decrypted JSON message (AIBOT uses JSON format)
|
|
||||||
var msg WeComBotMessage
|
|
||||||
if err := json.Unmarshal([]byte(decryptedMsg), &msg); err != nil {
|
|
||||||
logger.ErrorCF("wecom", "Failed to parse decrypted message", map[string]any{
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
http.Error(w, "Invalid message format", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Process the message with the channel's long-lived context (not the HTTP
|
|
||||||
// request context, which is canceled as soon as we return the response).
|
|
||||||
go c.processMessage(c.ctx, msg)
|
|
||||||
|
|
||||||
// Return success response immediately
|
|
||||||
// WeCom Bot requires response within configured timeout (default 5 seconds)
|
|
||||||
w.Write([]byte("success"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// processMessage processes the received message
|
|
||||||
func (c *WeComBotChannel) processMessage(ctx context.Context, msg WeComBotMessage) {
|
|
||||||
// Skip unsupported message types
|
|
||||||
if msg.MsgType != "text" && msg.MsgType != "image" && msg.MsgType != "voice" && msg.MsgType != "file" &&
|
|
||||||
msg.MsgType != "mixed" {
|
|
||||||
logger.DebugCF("wecom", "Skipping non-supported message type", map[string]any{
|
|
||||||
"msg_type": msg.MsgType,
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Message deduplication: Use msg_id to prevent duplicate processing
|
|
||||||
msgID := msg.MsgID
|
|
||||||
if !c.processedMsgs.MarkMessageProcessed(msgID) {
|
|
||||||
logger.DebugCF("wecom", "Skipping duplicate message", map[string]any{
|
|
||||||
"msg_id": msgID,
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
senderID := msg.From.UserID
|
|
||||||
|
|
||||||
// Determine if this is a group chat or direct message
|
|
||||||
// ChatType: "single" for DM, "group" for group chat
|
|
||||||
isGroupChat := msg.ChatType == "group"
|
|
||||||
|
|
||||||
var chatID, peerKind, peerID string
|
|
||||||
if isGroupChat {
|
|
||||||
// Group chat: use ChatID as chatID and peer_id
|
|
||||||
chatID = msg.ChatID
|
|
||||||
peerKind = "group"
|
|
||||||
peerID = msg.ChatID
|
|
||||||
} else {
|
|
||||||
// Direct message: use senderID as chatID and peer_id
|
|
||||||
chatID = senderID
|
|
||||||
peerKind = "direct"
|
|
||||||
peerID = senderID
|
|
||||||
}
|
|
||||||
|
|
||||||
// Extract content based on message type
|
|
||||||
var content string
|
|
||||||
switch msg.MsgType {
|
|
||||||
case "text":
|
|
||||||
content = msg.Text.Content
|
|
||||||
case "voice":
|
|
||||||
content = msg.Voice.Content // Voice to text content
|
|
||||||
case "mixed":
|
|
||||||
// For mixed messages, concatenate text items
|
|
||||||
for _, item := range msg.Mixed.MsgItem {
|
|
||||||
if item.MsgType == "text" {
|
|
||||||
content += item.Text.Content
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "image", "file":
|
|
||||||
// For image and file, we don't have text content
|
|
||||||
content = ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build metadata
|
|
||||||
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
|
||||||
|
|
||||||
// In group chats, apply unified group trigger filtering
|
|
||||||
if isGroupChat {
|
|
||||||
respond, cleaned := c.ShouldRespondInGroup(false, content)
|
|
||||||
if !respond {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
content = cleaned
|
|
||||||
}
|
|
||||||
|
|
||||||
metadata := map[string]string{
|
|
||||||
"msg_type": msg.MsgType,
|
|
||||||
"msg_id": msg.MsgID,
|
|
||||||
"platform": "wecom",
|
|
||||||
"response_url": msg.ResponseURL,
|
|
||||||
}
|
|
||||||
if isGroupChat {
|
|
||||||
metadata["chat_id"] = msg.ChatID
|
|
||||||
metadata["sender_id"] = senderID
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.DebugCF("wecom", "Received message", map[string]any{
|
|
||||||
"sender_id": senderID,
|
|
||||||
"msg_type": msg.MsgType,
|
|
||||||
"peer_kind": peerKind,
|
|
||||||
"is_group_chat": isGroupChat,
|
|
||||||
"preview": utils.Truncate(content, 50),
|
|
||||||
})
|
|
||||||
|
|
||||||
// Build sender info
|
|
||||||
sender := bus.SenderInfo{
|
|
||||||
Platform: "wecom",
|
|
||||||
PlatformID: senderID,
|
|
||||||
CanonicalID: identity.BuildCanonicalID("wecom", senderID),
|
|
||||||
}
|
|
||||||
|
|
||||||
if !c.IsAllowedSender(sender) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handle the message through the base channel
|
|
||||||
c.HandleMessage(ctx, peer, msg.MsgID, senderID, chatID, content, nil, metadata, sender)
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendWebhookReply sends a reply using the webhook URL
|
|
||||||
func (c *WeComBotChannel) sendWebhookReply(ctx context.Context, userID, content string) error {
|
|
||||||
reply := WeComBotReplyMessage{
|
|
||||||
MsgType: "text",
|
|
||||||
}
|
|
||||||
reply.Text.Content = content
|
|
||||||
|
|
||||||
jsonData, err := json.Marshal(reply)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to marshal reply: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Use configurable timeout (default 5 seconds)
|
|
||||||
timeout := c.config.ReplyTimeout
|
|
||||||
if timeout <= 0 {
|
|
||||||
timeout = 5
|
|
||||||
}
|
|
||||||
|
|
||||||
reqCtx, cancel := context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, c.config.WebhookURL, bytes.NewBuffer(jsonData))
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
|
|
||||||
resp, err := c.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return channels.ClassifyNetError(err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
body, readErr := io.ReadAll(resp.Body)
|
|
||||||
if readErr != nil {
|
|
||||||
return channels.ClassifySendError(
|
|
||||||
resp.StatusCode,
|
|
||||||
fmt.Errorf("reading webhook error response: %w", readErr),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return channels.ClassifySendError(
|
|
||||||
resp.StatusCode,
|
|
||||||
fmt.Errorf("webhook API error: %s", string(body)),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check response
|
|
||||||
var result struct {
|
|
||||||
ErrCode int `json:"errcode"`
|
|
||||||
ErrMsg string `json:"errmsg"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(body, &result); err != nil {
|
|
||||||
return fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.ErrCode != 0 {
|
|
||||||
return fmt.Errorf("webhook API error: %s (code: %d)", result.ErrMsg, result.ErrCode)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleHealth handles health check requests
|
|
||||||
func (c *WeComBotChannel) handleHealth(w http.ResponseWriter, r *http.Request) {
|
|
||||||
status := map[string]any{
|
|
||||||
"status": "ok",
|
|
||||||
"running": c.IsRunning(),
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(status)
|
|
||||||
}
|
|
||||||
|
|
@ -1,734 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"crypto/aes"
|
|
||||||
"crypto/cipher"
|
|
||||||
"crypto/sha1"
|
|
||||||
"encoding/base64"
|
|
||||||
"encoding/binary"
|
|
||||||
"encoding/json"
|
|
||||||
"encoding/xml"
|
|
||||||
"fmt"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"sort"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// generateTestAESKey generates a valid test AES key
|
|
||||||
func generateTestAESKey() string {
|
|
||||||
// AES key needs to be 32 bytes (256 bits) for AES-256
|
|
||||||
key := make([]byte, 32)
|
|
||||||
for i := range key {
|
|
||||||
key[i] = byte(i)
|
|
||||||
}
|
|
||||||
// Return base64 encoded key without padding
|
|
||||||
return base64.StdEncoding.EncodeToString(key)[:43]
|
|
||||||
}
|
|
||||||
|
|
||||||
// encryptTestMessage encrypts a message for testing (AIBOT JSON format)
|
|
||||||
func encryptTestMessage(message, aesKey string) (string, error) {
|
|
||||||
// Decode AES key
|
|
||||||
key, err := base64.StdEncoding.DecodeString(aesKey + "=")
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Prepare message: random(16) + msg_len(4) + msg + receiveid
|
|
||||||
random := make([]byte, 0, 16)
|
|
||||||
for i := range 16 {
|
|
||||||
random = append(random, byte(i))
|
|
||||||
}
|
|
||||||
|
|
||||||
msgBytes := []byte(message)
|
|
||||||
receiveID := []byte("test_aibot_id")
|
|
||||||
|
|
||||||
msgLen := uint32(len(msgBytes))
|
|
||||||
lenBytes := make([]byte, 4)
|
|
||||||
binary.BigEndian.PutUint32(lenBytes, msgLen)
|
|
||||||
|
|
||||||
plainText := append(random, lenBytes...)
|
|
||||||
plainText = append(plainText, msgBytes...)
|
|
||||||
plainText = append(plainText, receiveID...)
|
|
||||||
|
|
||||||
// PKCS7 padding
|
|
||||||
blockSize := aes.BlockSize
|
|
||||||
padding := blockSize - len(plainText)%blockSize
|
|
||||||
padText := bytes.Repeat([]byte{byte(padding)}, padding)
|
|
||||||
plainText = append(plainText, padText...)
|
|
||||||
|
|
||||||
// Encrypt
|
|
||||||
block, err := aes.NewCipher(key)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
mode := cipher.NewCBCEncrypter(block, key[:aes.BlockSize])
|
|
||||||
cipherText := make([]byte, len(plainText))
|
|
||||||
mode.CryptBlocks(cipherText, plainText)
|
|
||||||
|
|
||||||
return base64.StdEncoding.EncodeToString(cipherText), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// generateSignature generates a signature for testing
|
|
||||||
func generateSignature(token, timestamp, nonce, msgEncrypt string) string {
|
|
||||||
params := []string{token, timestamp, nonce, msgEncrypt}
|
|
||||||
sort.Strings(params)
|
|
||||||
str := strings.Join(params, "")
|
|
||||||
hash := sha1.Sum([]byte(str))
|
|
||||||
return fmt.Sprintf("%x", hash)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewWeComBotChannel(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
|
|
||||||
t.Run("missing token", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
_, err := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for missing token, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("missing webhook_url", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = ""
|
|
||||||
_, err := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for missing webhook_url, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("valid config", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.AllowFrom = []string{"user1", "user2"}
|
|
||||||
ch, err := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
if ch.Name() != "wecom" {
|
|
||||||
t.Errorf("Name() = %q, want %q", ch.Name(), "wecom")
|
|
||||||
}
|
|
||||||
if ch.IsRunning() {
|
|
||||||
t.Error("new channel should not be running")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotChannelIsAllowed(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
|
|
||||||
t.Run("empty allowlist allows all", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.AllowFrom = []string{}
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
if !ch.IsAllowed("any_user") {
|
|
||||||
t.Error("empty allowlist should allow all users")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("allowlist restricts users", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.AllowFrom = []string{"allowed_user"}
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
if !ch.IsAllowed("allowed_user") {
|
|
||||||
t.Error("allowed user should pass allowlist check")
|
|
||||||
}
|
|
||||||
if ch.IsAllowed("blocked_user") {
|
|
||||||
t.Error("non-allowed user should be blocked")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotVerifySignature(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
t.Run("valid signature", func(t *testing.T) {
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
msgEncrypt := "test_message"
|
|
||||||
expectedSig := generateSignature("test_token", timestamp, nonce, msgEncrypt)
|
|
||||||
|
|
||||||
if !verifySignature(ch.config.Token(), expectedSig, timestamp, nonce, msgEncrypt) {
|
|
||||||
t.Error("valid signature should pass verification")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid signature", func(t *testing.T) {
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
msgEncrypt := "test_message"
|
|
||||||
|
|
||||||
if verifySignature(ch.config.Token(), "invalid_sig", timestamp, nonce, msgEncrypt) {
|
|
||||||
t.Error("invalid signature should fail verification")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("empty token rejects verification (fail-closed)", func(t *testing.T) {
|
|
||||||
cfgEmpty := config.WeComConfig{}
|
|
||||||
cfgEmpty.SetToken("")
|
|
||||||
cfgEmpty.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
chEmpty := &WeComBotChannel{
|
|
||||||
config: cfgEmpty,
|
|
||||||
}
|
|
||||||
|
|
||||||
if verifySignature(chEmpty.config.Token(), "any_sig", "any_ts", "any_nonce", "any_msg") {
|
|
||||||
t.Error("empty token should reject verification (fail-closed)")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotDecryptMessage(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
|
|
||||||
t.Run("decrypt without AES key", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.SetEncodingAESKey("")
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
// Without AES key, message should be base64 decoded only
|
|
||||||
plainText := "hello world"
|
|
||||||
encoded := base64.StdEncoding.EncodeToString([]byte(plainText))
|
|
||||||
|
|
||||||
result, err := decryptMessage(encoded, ch.config.EncodingAESKey())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
if result != plainText {
|
|
||||||
t.Errorf("decryptMessage() = %q, want %q", result, plainText)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("decrypt with AES key", func(t *testing.T) {
|
|
||||||
aesKey := generateTestAESKey()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.SetEncodingAESKey(aesKey)
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
originalMsg := "<xml><Content>Hello</Content></xml>"
|
|
||||||
encrypted, err := encryptTestMessage(originalMsg, aesKey)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to encrypt test message: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err := decryptMessage(encrypted, ch.config.EncodingAESKey())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
if result != originalMsg {
|
|
||||||
t.Errorf("WeComDecryptMessage() = %q, want %q", result, originalMsg)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid base64", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.SetEncodingAESKey("")
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
_, err := decryptMessage("invalid_base64!!!", ch.config.EncodingAESKey())
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for invalid base64, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid AES key", func(t *testing.T) {
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
cfg.SetEncodingAESKey("invalid_key")
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
_, err := decryptMessage(base64.StdEncoding.EncodeToString([]byte("test")), ch.config.EncodingAESKey())
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for invalid AES key, got nil")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotPKCS7Unpad(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
input []byte
|
|
||||||
expected []byte
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "empty input",
|
|
||||||
input: []byte{},
|
|
||||||
expected: []byte{},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "valid padding 3 bytes",
|
|
||||||
input: append([]byte("hello"), bytes.Repeat([]byte{3}, 3)...),
|
|
||||||
expected: []byte("hello"),
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "valid padding 16 bytes (full block)",
|
|
||||||
input: append([]byte("123456789012345"), bytes.Repeat([]byte{16}, 16)...),
|
|
||||||
expected: []byte("123456789012345"),
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "invalid padding larger than data",
|
|
||||||
input: []byte{20},
|
|
||||||
expected: nil, // should return error
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "invalid padding zero",
|
|
||||||
input: append([]byte("test"), byte(0)),
|
|
||||||
expected: nil, // should return error
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
result, err := pkcs7Unpad(tt.input)
|
|
||||||
if tt.expected == nil {
|
|
||||||
// This case should return an error
|
|
||||||
if err == nil {
|
|
||||||
t.Errorf("pkcs7Unpad() expected error for invalid padding, got result: %v", result)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
t.Errorf("pkcs7Unpad() unexpected error: %v", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if !bytes.Equal(result, tt.expected) {
|
|
||||||
t.Errorf("pkcs7Unpad() = %v, want %v", result, tt.expected)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotHandleVerification(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
aesKey := generateTestAESKey()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey(aesKey)
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
t.Run("valid verification request", func(t *testing.T) {
|
|
||||||
echostr := "test_echostr_123"
|
|
||||||
encryptedEchostr, _ := encryptTestMessage(echostr, aesKey)
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
signature := generateSignature("test_token", timestamp, nonce, encryptedEchostr)
|
|
||||||
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodGet,
|
|
||||||
"/webhook/wecom?msg_signature="+signature+"×tamp="+timestamp+"&nonce="+nonce+"&echostr="+encryptedEchostr,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleVerification(context.Background(), w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
|
|
||||||
}
|
|
||||||
if w.Body.String() != echostr {
|
|
||||||
t.Errorf("response body = %q, want %q", w.Body.String(), echostr)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("missing parameters", func(t *testing.T) {
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/webhook/wecom?msg_signature=sig×tamp=ts", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleVerification(context.Background(), w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusBadRequest)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid signature", func(t *testing.T) {
|
|
||||||
echostr := "test_echostr"
|
|
||||||
encryptedEchostr, _ := encryptTestMessage(echostr, aesKey)
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodGet,
|
|
||||||
"/webhook/wecom?msg_signature=invalid_sig×tamp="+timestamp+"&nonce="+nonce+"&echostr="+encryptedEchostr,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleVerification(context.Background(), w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusForbidden {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusForbidden)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotHandleMessageCallback(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
aesKey := generateTestAESKey()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.SetEncodingAESKey(aesKey)
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
runBotMessageCallback := func(t *testing.T, jsonMsg string) *httptest.ResponseRecorder {
|
|
||||||
t.Helper()
|
|
||||||
encrypted, _ := encryptTestMessage(jsonMsg, aesKey)
|
|
||||||
encryptedWrapper := struct {
|
|
||||||
XMLName xml.Name `xml:"xml"`
|
|
||||||
Encrypt string `xml:"Encrypt"`
|
|
||||||
}{
|
|
||||||
Encrypt: encrypted,
|
|
||||||
}
|
|
||||||
wrapperData, _ := xml.Marshal(encryptedWrapper)
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
signature := generateSignature("test_token", timestamp, nonce, encrypted)
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodPost,
|
|
||||||
"/webhook/wecom?msg_signature="+signature+"×tamp="+timestamp+"&nonce="+nonce,
|
|
||||||
bytes.NewReader(wrapperData),
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
ch.handleMessageCallback(context.Background(), w, req)
|
|
||||||
return w
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Run("valid direct message callback", func(t *testing.T) {
|
|
||||||
w := runBotMessageCallback(t, `{
|
|
||||||
"msgid": "test_msg_id_123",
|
|
||||||
"aibotid": "test_aibot_id",
|
|
||||||
"chattype": "single",
|
|
||||||
"from": {"userid": "user123"},
|
|
||||||
"response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
"msgtype": "text",
|
|
||||||
"text": {"content": "Hello World"}
|
|
||||||
}`)
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
|
|
||||||
}
|
|
||||||
if w.Body.String() != "success" {
|
|
||||||
t.Errorf("response body = %q, want %q", w.Body.String(), "success")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("valid group message callback", func(t *testing.T) {
|
|
||||||
w := runBotMessageCallback(t, `{
|
|
||||||
"msgid": "test_msg_id_456",
|
|
||||||
"aibotid": "test_aibot_id",
|
|
||||||
"chatid": "group_chat_id_123",
|
|
||||||
"chattype": "group",
|
|
||||||
"from": {"userid": "user456"},
|
|
||||||
"response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
"msgtype": "text",
|
|
||||||
"text": {"content": "Hello Group"}
|
|
||||||
}`)
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
|
|
||||||
}
|
|
||||||
if w.Body.String() != "success" {
|
|
||||||
t.Errorf("response body = %q, want %q", w.Body.String(), "success")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("missing parameters", func(t *testing.T) {
|
|
||||||
req := httptest.NewRequest(http.MethodPost, "/webhook/wecom?msg_signature=sig", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleMessageCallback(context.Background(), w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusBadRequest)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid XML", func(t *testing.T) {
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
signature := generateSignature("test_token", timestamp, nonce, "")
|
|
||||||
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodPost,
|
|
||||||
"/webhook/wecom?msg_signature="+signature+"×tamp="+timestamp+"&nonce="+nonce,
|
|
||||||
strings.NewReader("invalid xml"),
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleMessageCallback(context.Background(), w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusBadRequest)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid signature", func(t *testing.T) {
|
|
||||||
encryptedWrapper := struct {
|
|
||||||
XMLName xml.Name `xml:"xml"`
|
|
||||||
Encrypt string `xml:"Encrypt"`
|
|
||||||
}{
|
|
||||||
Encrypt: "encrypted_data",
|
|
||||||
}
|
|
||||||
wrapperData, _ := xml.Marshal(encryptedWrapper)
|
|
||||||
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodPost,
|
|
||||||
"/webhook/wecom?msg_signature=invalid_sig×tamp="+timestamp+"&nonce="+nonce,
|
|
||||||
bytes.NewReader(wrapperData),
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleMessageCallback(context.Background(), w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusForbidden {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusForbidden)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotProcessMessage(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
t.Run("process direct text message", func(t *testing.T) {
|
|
||||||
msg := WeComBotMessage{
|
|
||||||
MsgID: "test_msg_id_123",
|
|
||||||
AIBotID: "test_aibot_id",
|
|
||||||
ChatType: "single",
|
|
||||||
ResponseURL: "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
MsgType: "text",
|
|
||||||
}
|
|
||||||
msg.From.UserID = "user123"
|
|
||||||
msg.Text.Content = "Hello World"
|
|
||||||
|
|
||||||
// Should not panic
|
|
||||||
ch.processMessage(context.Background(), msg)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("process group text message", func(t *testing.T) {
|
|
||||||
msg := WeComBotMessage{
|
|
||||||
MsgID: "test_msg_id_456",
|
|
||||||
AIBotID: "test_aibot_id",
|
|
||||||
ChatID: "group_chat_id_123",
|
|
||||||
ChatType: "group",
|
|
||||||
ResponseURL: "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
MsgType: "text",
|
|
||||||
}
|
|
||||||
msg.From.UserID = "user456"
|
|
||||||
msg.Text.Content = "Hello Group"
|
|
||||||
|
|
||||||
// Should not panic
|
|
||||||
ch.processMessage(context.Background(), msg)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("process voice message", func(t *testing.T) {
|
|
||||||
msg := WeComBotMessage{
|
|
||||||
MsgID: "test_msg_id_789",
|
|
||||||
AIBotID: "test_aibot_id",
|
|
||||||
ChatType: "single",
|
|
||||||
ResponseURL: "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
MsgType: "voice",
|
|
||||||
}
|
|
||||||
msg.From.UserID = "user123"
|
|
||||||
msg.Voice.Content = "Voice message text"
|
|
||||||
|
|
||||||
// Should not panic
|
|
||||||
ch.processMessage(context.Background(), msg)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("skip unsupported message type", func(t *testing.T) {
|
|
||||||
msg := WeComBotMessage{
|
|
||||||
MsgID: "test_msg_id_000",
|
|
||||||
AIBotID: "test_aibot_id",
|
|
||||||
ChatType: "single",
|
|
||||||
ResponseURL: "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
MsgType: "video",
|
|
||||||
}
|
|
||||||
msg.From.UserID = "user123"
|
|
||||||
|
|
||||||
// Should not panic
|
|
||||||
ch.processMessage(context.Background(), msg)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotHandleWebhook(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
t.Run("GET request calls verification", func(t *testing.T) {
|
|
||||||
echostr := "test_echostr"
|
|
||||||
encoded := base64.StdEncoding.EncodeToString([]byte(echostr))
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
signature := generateSignature("test_token", timestamp, nonce, encoded)
|
|
||||||
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodGet,
|
|
||||||
"/webhook/wecom?msg_signature="+signature+"×tamp="+timestamp+"&nonce="+nonce+"&echostr="+encoded,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleWebhook(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("POST request calls message callback", func(t *testing.T) {
|
|
||||||
encryptedWrapper := struct {
|
|
||||||
XMLName xml.Name `xml:"xml"`
|
|
||||||
Encrypt string `xml:"Encrypt"`
|
|
||||||
}{
|
|
||||||
Encrypt: base64.StdEncoding.EncodeToString([]byte("test")),
|
|
||||||
}
|
|
||||||
wrapperData, _ := xml.Marshal(encryptedWrapper)
|
|
||||||
|
|
||||||
timestamp := "1234567890"
|
|
||||||
nonce := "test_nonce"
|
|
||||||
signature := generateSignature("test_token", timestamp, nonce, encryptedWrapper.Encrypt)
|
|
||||||
|
|
||||||
req := httptest.NewRequest(
|
|
||||||
http.MethodPost,
|
|
||||||
"/webhook/wecom?msg_signature="+signature+"×tamp="+timestamp+"&nonce="+nonce,
|
|
||||||
bytes.NewReader(wrapperData),
|
|
||||||
)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleWebhook(w, req)
|
|
||||||
|
|
||||||
// Should not be method not allowed
|
|
||||||
if w.Code == http.StatusMethodNotAllowed {
|
|
||||||
t.Error("POST request should not return Method Not Allowed")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("unsupported method", func(t *testing.T) {
|
|
||||||
req := httptest.NewRequest(http.MethodPut, "/webhook/wecom", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleWebhook(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusMethodNotAllowed {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusMethodNotAllowed)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotHandleHealth(t *testing.T) {
|
|
||||||
msgBus := bus.NewMessageBus()
|
|
||||||
cfg := config.WeComConfig{}
|
|
||||||
cfg.SetToken("test_token")
|
|
||||||
cfg.WebhookURL = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test"
|
|
||||||
ch, _ := NewWeComBotChannel(cfg, msgBus)
|
|
||||||
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/health/wecom", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
|
|
||||||
ch.handleHealth(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Errorf("status code = %d, want %d", w.Code, http.StatusOK)
|
|
||||||
}
|
|
||||||
|
|
||||||
contentType := w.Header().Get("Content-Type")
|
|
||||||
if contentType != "application/json" {
|
|
||||||
t.Errorf("Content-Type = %q, want %q", contentType, "application/json")
|
|
||||||
}
|
|
||||||
|
|
||||||
body := w.Body.String()
|
|
||||||
if !strings.Contains(body, "status") || !strings.Contains(body, "running") {
|
|
||||||
t.Errorf("response body should contain status and running fields, got: %s", body)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotReplyMessage(t *testing.T) {
|
|
||||||
msg := WeComBotReplyMessage{
|
|
||||||
MsgType: "text",
|
|
||||||
}
|
|
||||||
msg.Text.Content = "Hello World"
|
|
||||||
|
|
||||||
if msg.MsgType != "text" {
|
|
||||||
t.Errorf("MsgType = %q, want %q", msg.MsgType, "text")
|
|
||||||
}
|
|
||||||
if msg.Text.Content != "Hello World" {
|
|
||||||
t.Errorf("Text.Content = %q, want %q", msg.Text.Content, "Hello World")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWeComBotMessageStructure(t *testing.T) {
|
|
||||||
jsonData := `{
|
|
||||||
"msgid": "test_msg_id_123",
|
|
||||||
"aibotid": "test_aibot_id",
|
|
||||||
"chatid": "group_chat_id_123",
|
|
||||||
"chattype": "group",
|
|
||||||
"from": {"userid": "user123"},
|
|
||||||
"response_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=test",
|
|
||||||
"msgtype": "text",
|
|
||||||
"text": {"content": "Hello World"}
|
|
||||||
}`
|
|
||||||
|
|
||||||
var msg WeComBotMessage
|
|
||||||
err := json.Unmarshal([]byte(jsonData), &msg)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to unmarshal JSON: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if msg.MsgID != "test_msg_id_123" {
|
|
||||||
t.Errorf("MsgID = %q, want %q", msg.MsgID, "test_msg_id_123")
|
|
||||||
}
|
|
||||||
if msg.AIBotID != "test_aibot_id" {
|
|
||||||
t.Errorf("AIBotID = %q, want %q", msg.AIBotID, "test_aibot_id")
|
|
||||||
}
|
|
||||||
if msg.ChatID != "group_chat_id_123" {
|
|
||||||
t.Errorf("ChatID = %q, want %q", msg.ChatID, "group_chat_id_123")
|
|
||||||
}
|
|
||||||
if msg.ChatType != "group" {
|
|
||||||
t.Errorf("ChatType = %q, want %q", msg.ChatType, "group")
|
|
||||||
}
|
|
||||||
if msg.From.UserID != "user123" {
|
|
||||||
t.Errorf("From.UserID = %q, want %q", msg.From.UserID, "user123")
|
|
||||||
}
|
|
||||||
if msg.MsgType != "text" {
|
|
||||||
t.Errorf("MsgType = %q, want %q", msg.MsgType, "text")
|
|
||||||
}
|
|
||||||
if msg.Text.Content != "Hello World" {
|
|
||||||
t.Errorf("Text.Content = %q, want %q", msg.Text.Content, "Hello World")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,199 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/aes"
|
|
||||||
"crypto/cipher"
|
|
||||||
"crypto/rand"
|
|
||||||
"crypto/sha1"
|
|
||||||
"encoding/base64"
|
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
|
||||||
"math/big"
|
|
||||||
"sort"
|
|
||||||
"strings"
|
|
||||||
)
|
|
||||||
|
|
||||||
// blockSize is the PKCS7 block size used by WeCom (32)
|
|
||||||
const blockSize = 32
|
|
||||||
|
|
||||||
// computeSignature computes the WeCom message signature from the given parameters.
|
|
||||||
// It sorts [token, timestamp, nonce, encrypt], concatenates them and returns the SHA1 hex digest.
|
|
||||||
func computeSignature(token, timestamp, nonce, encrypt string) string {
|
|
||||||
params := []string{token, timestamp, nonce, encrypt}
|
|
||||||
sort.Strings(params)
|
|
||||||
str := strings.Join(params, "")
|
|
||||||
hash := sha1.Sum([]byte(str))
|
|
||||||
return fmt.Sprintf("%x", hash)
|
|
||||||
}
|
|
||||||
|
|
||||||
// verifySignature verifies the message signature for WeCom
|
|
||||||
// This is a common function used by both WeCom Bot and WeCom App
|
|
||||||
func verifySignature(token, msgSignature, timestamp, nonce, msgEncrypt string) bool {
|
|
||||||
if token == "" {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return computeSignature(token, timestamp, nonce, msgEncrypt) == msgSignature
|
|
||||||
}
|
|
||||||
|
|
||||||
// decryptMessage decrypts the encrypted message using AES
|
|
||||||
// For AIBOT, receiveid should be the aibotid; for other apps, it should be corp_id
|
|
||||||
func decryptMessage(encryptedMsg, encodingAESKey string) (string, error) {
|
|
||||||
return decryptMessageWithVerify(encryptedMsg, encodingAESKey, "")
|
|
||||||
}
|
|
||||||
|
|
||||||
// decryptMessageWithVerify decrypts the encrypted message and optionally verifies receiveid
|
|
||||||
// receiveid: for AIBOT use aibotid, for WeCom App use corp_id. If empty, skip verification.
|
|
||||||
func decryptMessageWithVerify(encryptedMsg, encodingAESKey, receiveid string) (string, error) {
|
|
||||||
if encodingAESKey == "" {
|
|
||||||
// No encryption, return as is (base64 decode)
|
|
||||||
decoded, err := base64.StdEncoding.DecodeString(encryptedMsg)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
return string(decoded), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
aesKey, err := decodeWeComAESKey(encodingAESKey)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
cipherText, err := base64.StdEncoding.DecodeString(encryptedMsg)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to decode message: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
plainText, err := decryptAESCBC(aesKey, cipherText)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
return unpackWeComFrame(plainText, receiveid)
|
|
||||||
}
|
|
||||||
|
|
||||||
// decodeWeComAESKey base64-decodes the 43-character EncodingAESKey (trailing "=" is
|
|
||||||
// appended automatically) and validates that the result is exactly 32 bytes.
|
|
||||||
// It is the single place that handles this repeated pattern in both encrypt and decrypt paths.
|
|
||||||
func decodeWeComAESKey(encodingAESKey string) ([]byte, error) {
|
|
||||||
aesKey, err := base64.StdEncoding.DecodeString(encodingAESKey + "=")
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to decode AES key: %w", err)
|
|
||||||
}
|
|
||||||
if len(aesKey) != 32 {
|
|
||||||
return nil, fmt.Errorf("invalid AES key length: %d", len(aesKey))
|
|
||||||
}
|
|
||||||
return aesKey, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// encryptAESCBC encrypts plaintext using AES-CBC with the given key, mirroring
|
|
||||||
// decryptAESCBC. IV = aesKey[:aes.BlockSize]. The caller must PKCS7-pad the
|
|
||||||
// plaintext to a multiple of aes.BlockSize before calling.
|
|
||||||
func encryptAESCBC(aesKey, plaintext []byte) ([]byte, error) {
|
|
||||||
block, err := aes.NewCipher(aesKey)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create cipher: %w", err)
|
|
||||||
}
|
|
||||||
iv := aesKey[:aes.BlockSize]
|
|
||||||
ciphertext := make([]byte, len(plaintext))
|
|
||||||
cipher.NewCBCEncrypter(block, iv).CryptBlocks(ciphertext, plaintext)
|
|
||||||
return ciphertext, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// packWeComFrame builds the WeCom wire format:
|
|
||||||
//
|
|
||||||
// random(16 ASCII digits) + msg_len(4, big-endian) + msg + receiveid
|
|
||||||
func packWeComFrame(msg, receiveid string) ([]byte, error) {
|
|
||||||
randomBytes := make([]byte, 16)
|
|
||||||
for i := range 16 {
|
|
||||||
n, err := rand.Int(rand.Reader, big.NewInt(10))
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to generate random: %w", err)
|
|
||||||
}
|
|
||||||
randomBytes[i] = byte('0' + n.Int64())
|
|
||||||
}
|
|
||||||
msgBytes := []byte(msg)
|
|
||||||
msgLenBytes := make([]byte, 4)
|
|
||||||
binary.BigEndian.PutUint32(msgLenBytes, uint32(len(msgBytes)))
|
|
||||||
var buf bytes.Buffer
|
|
||||||
buf.Write(randomBytes)
|
|
||||||
buf.Write(msgLenBytes)
|
|
||||||
buf.Write(msgBytes)
|
|
||||||
buf.WriteString(receiveid)
|
|
||||||
return buf.Bytes(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// unpackWeComFrame parses the WeCom wire format produced by packWeComFrame.
|
|
||||||
// If receiveid is non-empty it verifies the frame's trailing receiveid field.
|
|
||||||
func unpackWeComFrame(data []byte, receiveid string) (string, error) {
|
|
||||||
if len(data) < 20 {
|
|
||||||
return "", fmt.Errorf("decrypted frame too short: %d bytes", len(data))
|
|
||||||
}
|
|
||||||
msgLen := binary.BigEndian.Uint32(data[16:20])
|
|
||||||
if int(msgLen) > len(data)-20 {
|
|
||||||
return "", fmt.Errorf("invalid message length: %d", msgLen)
|
|
||||||
}
|
|
||||||
msg := data[20 : 20+msgLen]
|
|
||||||
if receiveid != "" && len(data) > 20+int(msgLen) {
|
|
||||||
actualReceiveID := string(data[20+msgLen:])
|
|
||||||
if actualReceiveID != receiveid {
|
|
||||||
return "", fmt.Errorf("receiveid mismatch: expected %s, got %s", receiveid, actualReceiveID)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return string(msg), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// decryptAESCBC decrypts ciphertext using AES-CBC with the given key.
|
|
||||||
// IV = aesKey[:aes.BlockSize]. PKCS7 padding is stripped from the returned plaintext.
|
|
||||||
func decryptAESCBC(aesKey, ciphertext []byte) ([]byte, error) {
|
|
||||||
if len(ciphertext) == 0 {
|
|
||||||
return nil, fmt.Errorf("ciphertext is empty")
|
|
||||||
}
|
|
||||||
if len(ciphertext)%aes.BlockSize != 0 {
|
|
||||||
return nil, fmt.Errorf("ciphertext length %d is not a multiple of block size", len(ciphertext))
|
|
||||||
}
|
|
||||||
block, err := aes.NewCipher(aesKey)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create cipher: %w", err)
|
|
||||||
}
|
|
||||||
iv := aesKey[:aes.BlockSize]
|
|
||||||
plaintext := make([]byte, len(ciphertext))
|
|
||||||
cipher.NewCBCDecrypter(block, iv).CryptBlocks(plaintext, ciphertext)
|
|
||||||
plaintext, err = pkcs7Unpad(plaintext)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to unpad: %w", err)
|
|
||||||
}
|
|
||||||
return plaintext, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// pkcs7Pad adds PKCS7 padding
|
|
||||||
func pkcs7Pad(data []byte, blockSize int) []byte {
|
|
||||||
padding := blockSize - (len(data) % blockSize)
|
|
||||||
if padding == 0 {
|
|
||||||
padding = blockSize
|
|
||||||
}
|
|
||||||
padText := bytes.Repeat([]byte{byte(padding)}, padding)
|
|
||||||
return append(data, padText...)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pkcs7Unpad removes PKCS7 padding with validation
|
|
||||||
func pkcs7Unpad(data []byte) ([]byte, error) {
|
|
||||||
if len(data) == 0 {
|
|
||||||
return data, nil
|
|
||||||
}
|
|
||||||
padding := int(data[len(data)-1])
|
|
||||||
// WeCom uses 32-byte block size for PKCS7 padding
|
|
||||||
if padding == 0 || padding > blockSize {
|
|
||||||
return nil, fmt.Errorf("invalid padding size: %d", padding)
|
|
||||||
}
|
|
||||||
if padding > len(data) {
|
|
||||||
return nil, fmt.Errorf("padding size larger than data")
|
|
||||||
}
|
|
||||||
// Verify all padding bytes
|
|
||||||
for i := range padding {
|
|
||||||
if data[len(data)-1-i] != byte(padding) {
|
|
||||||
return nil, fmt.Errorf("invalid padding byte at position %d", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return data[:len(data)-padding], nil
|
|
||||||
}
|
|
||||||
|
|
@ -1,54 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import "sync"
|
|
||||||
|
|
||||||
const wecomMaxProcessedMessages = 1000
|
|
||||||
|
|
||||||
// MessageDeduplicator provides thread-safe message deduplication using a circular queue (ring buffer)
|
|
||||||
// combined with a hash map. This ensures fast O(1) lookups while naturally evicting the oldest
|
|
||||||
// messages without causing "amnesia cliffs" when the limit is reached.
|
|
||||||
type MessageDeduplicator struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
msgs map[string]bool
|
|
||||||
ring []string
|
|
||||||
idx int
|
|
||||||
max int
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewMessageDeduplicator creates a new deduplicator with the specified capacity.
|
|
||||||
func NewMessageDeduplicator(maxEntries int) *MessageDeduplicator {
|
|
||||||
if maxEntries <= 0 {
|
|
||||||
maxEntries = wecomMaxProcessedMessages
|
|
||||||
}
|
|
||||||
return &MessageDeduplicator{
|
|
||||||
msgs: make(map[string]bool, maxEntries),
|
|
||||||
ring: make([]string, maxEntries),
|
|
||||||
max: maxEntries,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// MarkMessageProcessed marks msgID as processed and returns false for duplicates.
|
|
||||||
func (d *MessageDeduplicator) MarkMessageProcessed(msgID string) bool {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
|
|
||||||
// 1. Check for duplicate
|
|
||||||
if d.msgs[msgID] {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// 2. Evict the oldest message at our current ring position (if any)
|
|
||||||
oldestID := d.ring[d.idx]
|
|
||||||
if oldestID != "" {
|
|
||||||
delete(d.msgs, oldestID)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 3. Store the new message
|
|
||||||
d.msgs[msgID] = true
|
|
||||||
d.ring[d.idx] = msgID
|
|
||||||
|
|
||||||
// 4. Advance the circle queue index
|
|
||||||
d.idx = (d.idx + 1) % d.max
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
@ -1,83 +0,0 @@
|
||||||
package wecom
|
|
||||||
|
|
||||||
import (
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMessageDeduplicator_DuplicateDetection(t *testing.T) {
|
|
||||||
d := NewMessageDeduplicator(wecomMaxProcessedMessages)
|
|
||||||
|
|
||||||
if ok := d.MarkMessageProcessed("msg-1"); !ok {
|
|
||||||
t.Fatalf("first message should be accepted")
|
|
||||||
}
|
|
||||||
|
|
||||||
if ok := d.MarkMessageProcessed("msg-1"); ok {
|
|
||||||
t.Fatalf("duplicate message should be rejected")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMessageDeduplicator_ConcurrentSameMessage(t *testing.T) {
|
|
||||||
d := NewMessageDeduplicator(wecomMaxProcessedMessages)
|
|
||||||
|
|
||||||
const goroutines = 64
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
wg.Add(goroutines)
|
|
||||||
|
|
||||||
results := make(chan bool, goroutines)
|
|
||||||
for i := 0; i < goroutines; i++ {
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
results <- d.MarkMessageProcessed("msg-concurrent")
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
|
|
||||||
wg.Wait()
|
|
||||||
close(results)
|
|
||||||
|
|
||||||
successes := 0
|
|
||||||
for ok := range results {
|
|
||||||
if ok {
|
|
||||||
successes++
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if successes != 1 {
|
|
||||||
t.Fatalf("expected exactly 1 successful mark, got %d", successes)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMessageDeduplicator_CircularQueueEviction(t *testing.T) {
|
|
||||||
// Create a deduplicator with a very small capacity to test eviction easily.
|
|
||||||
capacity := 3
|
|
||||||
d := NewMessageDeduplicator(capacity)
|
|
||||||
|
|
||||||
// Fill the queue.
|
|
||||||
d.MarkMessageProcessed("msg-1")
|
|
||||||
d.MarkMessageProcessed("msg-2")
|
|
||||||
d.MarkMessageProcessed("msg-3")
|
|
||||||
|
|
||||||
// At this point, the queue is full. msg-1 is the oldest.
|
|
||||||
if len(d.msgs) != 3 {
|
|
||||||
t.Fatalf("expected map size to be 3, got %d", len(d.msgs))
|
|
||||||
}
|
|
||||||
|
|
||||||
// This should evict msg-1 and add msg-4.
|
|
||||||
if ok := d.MarkMessageProcessed("msg-4"); !ok {
|
|
||||||
t.Fatalf("msg-4 should be accepted")
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(d.msgs) != 3 {
|
|
||||||
t.Fatalf("expected map size to remain at max capacity (3), got %d", len(d.msgs))
|
|
||||||
}
|
|
||||||
|
|
||||||
// msg-1 should now be forgotten (evicted).
|
|
||||||
if ok := d.MarkMessageProcessed("msg-1"); !ok {
|
|
||||||
t.Fatalf("msg-1 should be accepted again because it was evicted")
|
|
||||||
}
|
|
||||||
|
|
||||||
// msg-2 should have been evicted when we added msg-1 back.
|
|
||||||
if ok := d.MarkMessageProcessed("msg-2"); !ok {
|
|
||||||
t.Fatalf("msg-2 should be accepted again because it was evicted")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -8,12 +8,6 @@ import (
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
channels.RegisterFactory("wecom", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
channels.RegisterFactory("wecom", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||||
return NewWeComBotChannel(cfg.Channels.WeCom, b)
|
return NewChannel(cfg.Channels.WeCom, b)
|
||||||
})
|
|
||||||
channels.RegisterFactory("wecom_app", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
|
||||||
return NewWeComAppChannel(cfg.Channels.WeComApp, b)
|
|
||||||
})
|
|
||||||
channels.RegisterFactory("wecom_aibot", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
|
||||||
return NewWeComAIBotChannel(cfg.Channels.WeComAIBot, b)
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
802
pkg/channels/wecom/media.go
Normal file
802
pkg/channels/wecom/media.go
Normal file
|
|
@ -0,0 +1,802 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/aes"
|
||||||
|
"crypto/cipher"
|
||||||
|
"crypto/md5"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"mime"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/h2non/filetype"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
wecomOutboundMediaMaxBytes = 20 << 20
|
||||||
|
wecomOutboundImageMaxBytes = 2 << 20
|
||||||
|
wecomOutboundVoiceMaxBytes = 2 << 20
|
||||||
|
wecomOutboundVideoMaxBytes = 10 << 20
|
||||||
|
wecomUploadChunkMaxBytes = 512 << 10
|
||||||
|
wecomUploadMaxChunks = 100
|
||||||
|
wecomUploadMinBytes = 5
|
||||||
|
)
|
||||||
|
|
||||||
|
type wecomOutboundMedia struct {
|
||||||
|
MsgType string
|
||||||
|
MediaID string
|
||||||
|
Title string
|
||||||
|
Description string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *wecomOutboundMedia) respondBody() wecomRespondMsgBody {
|
||||||
|
body := wecomRespondMsgBody{MsgType: m.MsgType}
|
||||||
|
switch m.MsgType {
|
||||||
|
case "file":
|
||||||
|
body.File = &wecomMediaRefContent{MediaID: m.MediaID}
|
||||||
|
case "image":
|
||||||
|
body.Image = &wecomMediaRefContent{MediaID: m.MediaID}
|
||||||
|
case "voice":
|
||||||
|
body.Voice = &wecomMediaRefContent{MediaID: m.MediaID}
|
||||||
|
case "video":
|
||||||
|
body.Video = &wecomVideoContent{
|
||||||
|
MediaID: m.MediaID,
|
||||||
|
Title: m.Title,
|
||||||
|
Description: m.Description,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return body
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *wecomOutboundMedia) sendBody(chatID string, chatType uint32) wecomSendMsgBody {
|
||||||
|
body := wecomSendMsgBody{
|
||||||
|
ChatID: chatID,
|
||||||
|
ChatType: chatType,
|
||||||
|
MsgType: m.MsgType,
|
||||||
|
}
|
||||||
|
switch m.MsgType {
|
||||||
|
case "file":
|
||||||
|
body.File = &wecomMediaRefContent{MediaID: m.MediaID}
|
||||||
|
case "image":
|
||||||
|
body.Image = &wecomMediaRefContent{MediaID: m.MediaID}
|
||||||
|
case "voice":
|
||||||
|
body.Voice = &wecomMediaRefContent{MediaID: m.MediaID}
|
||||||
|
case "video":
|
||||||
|
body.Video = &wecomVideoContent{
|
||||||
|
MediaID: m.MediaID,
|
||||||
|
Title: m.Title,
|
||||||
|
Description: m.Description,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return body
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeMediaAESKey(value string) ([]byte, error) {
|
||||||
|
if value == "" {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
key, err := base64.StdEncoding.DecodeString(value)
|
||||||
|
if err == nil && len(key) == 32 {
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
key, err = base64.StdEncoding.DecodeString(value + "=")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("decode AES key: %w", err)
|
||||||
|
}
|
||||||
|
if len(key) != 32 {
|
||||||
|
return nil, fmt.Errorf("invalid AES key length %d", len(key))
|
||||||
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func decryptAESCBC(key, ciphertext []byte) ([]byte, error) {
|
||||||
|
if len(ciphertext) == 0 {
|
||||||
|
return nil, fmt.Errorf("ciphertext is empty")
|
||||||
|
}
|
||||||
|
if len(ciphertext)%aes.BlockSize != 0 {
|
||||||
|
return nil, fmt.Errorf("ciphertext length %d is not a multiple of block size", len(ciphertext))
|
||||||
|
}
|
||||||
|
block, err := aes.NewCipher(key)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create cipher: %w", err)
|
||||||
|
}
|
||||||
|
plaintext := make([]byte, len(ciphertext))
|
||||||
|
iv := key[:aes.BlockSize]
|
||||||
|
cipher.NewCBCDecrypter(block, iv).CryptBlocks(plaintext, ciphertext)
|
||||||
|
return pkcs7Unpad(plaintext)
|
||||||
|
}
|
||||||
|
|
||||||
|
func pkcs7Unpad(data []byte) ([]byte, error) {
|
||||||
|
if len(data) == 0 {
|
||||||
|
return nil, fmt.Errorf("empty plaintext")
|
||||||
|
}
|
||||||
|
padding := int(data[len(data)-1])
|
||||||
|
if padding == 0 || padding > 32 || padding > len(data) {
|
||||||
|
return nil, fmt.Errorf("invalid padding size %d", padding)
|
||||||
|
}
|
||||||
|
for i := 0; i < padding; i++ {
|
||||||
|
if data[len(data)-1-i] != byte(padding) {
|
||||||
|
return nil, fmt.Errorf("invalid padding byte")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return data[:len(data)-padding], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func inferMediaExt(contentType, fallback string) string {
|
||||||
|
contentType = normalizeWeComContentType(contentType)
|
||||||
|
switch contentType {
|
||||||
|
case "image/jpeg", "image/jpg":
|
||||||
|
return ".jpg"
|
||||||
|
case "image/png":
|
||||||
|
return ".png"
|
||||||
|
case "image/gif":
|
||||||
|
return ".gif"
|
||||||
|
case "image/webp":
|
||||||
|
return ".webp"
|
||||||
|
case "application/pdf":
|
||||||
|
return ".pdf"
|
||||||
|
case "video/mp4":
|
||||||
|
return ".mp4"
|
||||||
|
default:
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeWeComContentType(value string) string {
|
||||||
|
value = strings.ToLower(strings.TrimSpace(value))
|
||||||
|
if idx := strings.Index(value, ";"); idx >= 0 {
|
||||||
|
value = strings.TrimSpace(value[:idx])
|
||||||
|
}
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
func isGenericWeComContentType(value string) bool {
|
||||||
|
switch normalizeWeComContentType(value) {
|
||||||
|
case "", "application/octet-stream", "binary/octet-stream", "application/unknown", "application/binary":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeWeComFilename(name string) string {
|
||||||
|
name = filepath.Base(strings.TrimSpace(name))
|
||||||
|
if name == "." || name == "/" || name == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
|
||||||
|
func candidateWeComFilename(resourceURL, contentDisposition, fallbackName string) string {
|
||||||
|
if _, params, err := mime.ParseMediaType(contentDisposition); err == nil {
|
||||||
|
if name := sanitizeWeComFilename(params["filename"]); name != "" {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
if name := sanitizeWeComFilename(params["filename*"]); name != "" {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsed, err := url.Parse(resourceURL); err == nil {
|
||||||
|
query := parsed.Query()
|
||||||
|
for _, key := range []string{"filename", "file_name", "name"} {
|
||||||
|
if name := sanitizeWeComFilename(query.Get(key)); name != "" {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if name := sanitizeWeComFilename(parsed.Path); name != "" {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sanitizeWeComFilename(fallbackName)
|
||||||
|
}
|
||||||
|
|
||||||
|
func detectWeComFiletype(data []byte) (string, string) {
|
||||||
|
kind, err := filetype.Match(data)
|
||||||
|
if err != nil || kind == filetype.Unknown {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
ext := ""
|
||||||
|
if kind.Extension != "" {
|
||||||
|
ext = "." + strings.ToLower(kind.Extension)
|
||||||
|
}
|
||||||
|
return normalizeWeComContentType(kind.MIME.Value), ext
|
||||||
|
}
|
||||||
|
|
||||||
|
func detectWeComMediaMetadata(
|
||||||
|
data []byte,
|
||||||
|
fallbackName, fallbackContentType, resourceURL, contentDisposition string,
|
||||||
|
) (string, string) {
|
||||||
|
filename := candidateWeComFilename(resourceURL, contentDisposition, fallbackName)
|
||||||
|
if filename == "" {
|
||||||
|
filename = "media"
|
||||||
|
}
|
||||||
|
|
||||||
|
ext := strings.ToLower(filepath.Ext(filename))
|
||||||
|
contentType := normalizeWeComContentType(fallbackContentType)
|
||||||
|
detectedType, detectedExt := detectWeComFiletype(data)
|
||||||
|
|
||||||
|
if ext != "" && isGenericWeComContentType(contentType) {
|
||||||
|
if byExt := normalizeWeComContentType(mime.TypeByExtension(ext)); byExt != "" {
|
||||||
|
contentType = byExt
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if detectedType != "" {
|
||||||
|
switch {
|
||||||
|
case contentType == "":
|
||||||
|
contentType = detectedType
|
||||||
|
case isGenericWeComContentType(contentType):
|
||||||
|
contentType = detectedType
|
||||||
|
case strings.HasPrefix(detectedType, "image/") && !strings.HasPrefix(contentType, "image/"):
|
||||||
|
contentType = detectedType
|
||||||
|
case strings.HasPrefix(detectedType, "audio/") && !strings.HasPrefix(contentType, "audio/"):
|
||||||
|
contentType = detectedType
|
||||||
|
case strings.HasPrefix(detectedType, "video/") && !strings.HasPrefix(contentType, "video/"):
|
||||||
|
contentType = detectedType
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if contentType == "" && ext != "" {
|
||||||
|
contentType = normalizeWeComContentType(mime.TypeByExtension(ext))
|
||||||
|
}
|
||||||
|
if contentType == "" {
|
||||||
|
contentType = normalizeWeComContentType(http.DetectContentType(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
if ext == "" {
|
||||||
|
ext = detectedExt
|
||||||
|
}
|
||||||
|
if ext == "" && contentType != "" {
|
||||||
|
if exts, err := mime.ExtensionsByType(contentType); err == nil && len(exts) > 0 {
|
||||||
|
ext = strings.ToLower(exts[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if filepath.Ext(filename) == "" && ext != "" {
|
||||||
|
filename += ext
|
||||||
|
}
|
||||||
|
return filename, contentType
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) storeRemoteMedia(
|
||||||
|
ctx context.Context,
|
||||||
|
scope, msgID, resourceURL, aesKey, fallbackExt string,
|
||||||
|
) (string, error) {
|
||||||
|
store := c.GetMediaStore()
|
||||||
|
if store == nil {
|
||||||
|
return "", fmt.Errorf("no media store available")
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, resourceURL, nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("create request: %w", err)
|
||||||
|
}
|
||||||
|
resp, err := c.mediaClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("download media: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("download media returned HTTP %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := io.ReadAll(io.LimitReader(resp.Body, wecomOutboundMediaMaxBytes+1))
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("read media: %w", err)
|
||||||
|
}
|
||||||
|
if len(data) > wecomOutboundMediaMaxBytes {
|
||||||
|
return "", fmt.Errorf("media too large")
|
||||||
|
}
|
||||||
|
|
||||||
|
if aesKey != "" {
|
||||||
|
key, keyErr := decodeMediaAESKey(aesKey)
|
||||||
|
if keyErr != nil {
|
||||||
|
return "", keyErr
|
||||||
|
}
|
||||||
|
data, err = decryptAESCBC(key, data)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("decrypt media: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
filename, contentType := detectWeComMediaMetadata(
|
||||||
|
data,
|
||||||
|
msgID+fallbackExt,
|
||||||
|
resp.Header.Get("Content-Type"),
|
||||||
|
resourceURL,
|
||||||
|
resp.Header.Get("Content-Disposition"),
|
||||||
|
)
|
||||||
|
ext := filepath.Ext(filename)
|
||||||
|
if ext == "" {
|
||||||
|
ext = inferMediaExt(contentType, fallbackExt)
|
||||||
|
}
|
||||||
|
mediaDir := filepath.Join(os.TempDir(), "picoclaw_media")
|
||||||
|
if mkdirErr := os.MkdirAll(mediaDir, 0o700); mkdirErr != nil {
|
||||||
|
return "", fmt.Errorf("mkdir media dir: %w", mkdirErr)
|
||||||
|
}
|
||||||
|
tmpFile, err := os.CreateTemp(mediaDir, msgID+"-*"+ext)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("create temp file: %w", err)
|
||||||
|
}
|
||||||
|
tmpPath := tmpFile.Name()
|
||||||
|
if _, writeErr := tmpFile.Write(data); writeErr != nil {
|
||||||
|
tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Errorf("write temp file: %w", writeErr)
|
||||||
|
}
|
||||||
|
if closeErr := tmpFile.Close(); closeErr != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Errorf("close temp file: %w", closeErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
ref, err := store.Store(tmpPath, media.MediaMeta{
|
||||||
|
Filename: filename,
|
||||||
|
ContentType: contentType,
|
||||||
|
Source: "wecom",
|
||||||
|
CleanupPolicy: media.CleanupPolicyDeleteOnCleanup,
|
||||||
|
}, scope)
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return ref, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func detectLocalWeComContentType(localPath, hint string) string {
|
||||||
|
contentType := normalizeWeComContentType(hint)
|
||||||
|
if !isGenericWeComContentType(contentType) {
|
||||||
|
return contentType
|
||||||
|
}
|
||||||
|
|
||||||
|
if kind, err := filetype.MatchFile(localPath); err == nil && kind != filetype.Unknown {
|
||||||
|
return normalizeWeComContentType(kind.MIME.Value)
|
||||||
|
}
|
||||||
|
|
||||||
|
if ext := strings.ToLower(filepath.Ext(localPath)); ext != "" {
|
||||||
|
if byExt := normalizeWeComContentType(mime.TypeByExtension(ext)); byExt != "" {
|
||||||
|
return byExt
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
file, err := os.Open(localPath)
|
||||||
|
if err != nil {
|
||||||
|
return contentType
|
||||||
|
}
|
||||||
|
defer file.Close()
|
||||||
|
|
||||||
|
buf := make([]byte, 512)
|
||||||
|
n, err := file.Read(buf)
|
||||||
|
if err != nil && err != io.EOF {
|
||||||
|
return contentType
|
||||||
|
}
|
||||||
|
if n == 0 {
|
||||||
|
return contentType
|
||||||
|
}
|
||||||
|
return normalizeWeComContentType(http.DetectContentType(buf[:n]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeWeComTempFile(prefix, filename string, data []byte) (string, error) {
|
||||||
|
mediaDir := filepath.Join(os.TempDir(), "picoclaw_media")
|
||||||
|
if err := os.MkdirAll(mediaDir, 0o700); err != nil {
|
||||||
|
return "", fmt.Errorf("mkdir media dir: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ext := strings.ToLower(filepath.Ext(filename))
|
||||||
|
tmpFile, err := os.CreateTemp(mediaDir, prefix+"-*"+ext)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("create temp file: %w", err)
|
||||||
|
}
|
||||||
|
tmpPath := tmpFile.Name()
|
||||||
|
|
||||||
|
if _, err := tmpFile.Write(data); err != nil {
|
||||||
|
_ = tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Errorf("write temp file: %w", err)
|
||||||
|
}
|
||||||
|
if err := tmpFile.Close(); err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Errorf("close temp file: %w", err)
|
||||||
|
}
|
||||||
|
return tmpPath, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) downloadRemoteMediaToTemp(
|
||||||
|
ctx context.Context,
|
||||||
|
resourceURL, fallbackName string,
|
||||||
|
) (string, string, string, error) {
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, resourceURL, nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", fmt.Errorf("create request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := c.mediaClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", fmt.Errorf("download media: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||||
|
return "", "", "", fmt.Errorf("download media returned HTTP %d: %s", resp.StatusCode, string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := io.ReadAll(io.LimitReader(resp.Body, wecomOutboundMediaMaxBytes+1))
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", fmt.Errorf("read media: %w", err)
|
||||||
|
}
|
||||||
|
if len(data) > wecomOutboundMediaMaxBytes {
|
||||||
|
return "", "", "", fmt.Errorf("media too large")
|
||||||
|
}
|
||||||
|
|
||||||
|
filename, contentType := detectWeComMediaMetadata(
|
||||||
|
data,
|
||||||
|
fallbackName,
|
||||||
|
resp.Header.Get("Content-Type"),
|
||||||
|
resourceURL,
|
||||||
|
resp.Header.Get("Content-Disposition"),
|
||||||
|
)
|
||||||
|
tmpPath, err := writeWeComTempFile("wecom-outbound", filename, data)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", err
|
||||||
|
}
|
||||||
|
return tmpPath, filename, contentType, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) resolveOutboundPart(
|
||||||
|
ctx context.Context,
|
||||||
|
part bus.MediaPart,
|
||||||
|
) (string, string, string, func(), error) {
|
||||||
|
cleanup := func() {}
|
||||||
|
filename := sanitizeWeComFilename(part.Filename)
|
||||||
|
contentType := normalizeWeComContentType(part.ContentType)
|
||||||
|
ref := strings.TrimSpace(part.Ref)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case ref == "":
|
||||||
|
return "", filename, contentType, cleanup, nil
|
||||||
|
|
||||||
|
case strings.HasPrefix(ref, "http://") || strings.HasPrefix(ref, "https://"):
|
||||||
|
localPath, name, ct, err := c.downloadRemoteMediaToTemp(ctx, ref, filename)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
return localPath, name, ct, func() { _ = os.Remove(localPath) }, nil
|
||||||
|
|
||||||
|
case strings.HasPrefix(ref, "media://"):
|
||||||
|
store := c.GetMediaStore()
|
||||||
|
if store == nil {
|
||||||
|
return "", "", "", cleanup, fmt.Errorf("no media store available")
|
||||||
|
}
|
||||||
|
|
||||||
|
localPath, meta, err := store.ResolveWithMeta(ref)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
if filename == "" {
|
||||||
|
filename = sanitizeWeComFilename(meta.Filename)
|
||||||
|
}
|
||||||
|
if contentType == "" {
|
||||||
|
contentType = normalizeWeComContentType(meta.ContentType)
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(localPath, "http://") || strings.HasPrefix(localPath, "https://") {
|
||||||
|
tmpPath, name, ct, err := c.downloadRemoteMediaToTemp(ctx, localPath, filename)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
return tmpPath, name, ct, func() { _ = os.Remove(tmpPath) }, nil
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(localPath); err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
if filename == "" {
|
||||||
|
filename = sanitizeWeComFilename(filepath.Base(localPath))
|
||||||
|
}
|
||||||
|
if contentType == "" {
|
||||||
|
contentType = detectLocalWeComContentType(localPath, "")
|
||||||
|
}
|
||||||
|
return localPath, filename, contentType, cleanup, nil
|
||||||
|
|
||||||
|
case strings.HasPrefix(ref, "file://"):
|
||||||
|
u, err := url.Parse(ref)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
localPath := u.Path
|
||||||
|
if _, err := os.Stat(localPath); err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
if filename == "" {
|
||||||
|
filename = sanitizeWeComFilename(filepath.Base(localPath))
|
||||||
|
}
|
||||||
|
if contentType == "" {
|
||||||
|
contentType = detectLocalWeComContentType(localPath, "")
|
||||||
|
}
|
||||||
|
return localPath, filename, contentType, cleanup, nil
|
||||||
|
|
||||||
|
default:
|
||||||
|
if _, err := os.Stat(ref); err != nil {
|
||||||
|
return "", "", "", cleanup, err
|
||||||
|
}
|
||||||
|
if filename == "" {
|
||||||
|
filename = sanitizeWeComFilename(filepath.Base(ref))
|
||||||
|
}
|
||||||
|
if contentType == "" {
|
||||||
|
contentType = detectLocalWeComContentType(ref, "")
|
||||||
|
}
|
||||||
|
return ref, filename, contentType, cleanup, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func canWeComSendImage(contentType, ext string, size int64) bool {
|
||||||
|
if size > wecomOutboundImageMaxBytes {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch normalizeWeComContentType(contentType) {
|
||||||
|
case "image/jpeg", "image/jpg", "image/png", "image/gif":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
switch strings.ToLower(ext) {
|
||||||
|
case ".jpg", ".jpeg", ".png", ".gif":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func canWeComSendVoice(contentType, ext string, size int64) bool {
|
||||||
|
if size > wecomOutboundVoiceMaxBytes {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
contentType = normalizeWeComContentType(contentType)
|
||||||
|
return strings.Contains(contentType, "amr") || strings.EqualFold(ext, ".amr")
|
||||||
|
}
|
||||||
|
|
||||||
|
func canWeComSendVideo(contentType, ext string, size int64) bool {
|
||||||
|
if size > wecomOutboundVideoMaxBytes {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return normalizeWeComContentType(contentType) == "video/mp4" || strings.EqualFold(ext, ".mp4")
|
||||||
|
}
|
||||||
|
|
||||||
|
func outboundWeComMediaKind(partType, filename, contentType string, size int64) string {
|
||||||
|
if size < wecomUploadMinBytes {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
partType = strings.ToLower(strings.TrimSpace(partType))
|
||||||
|
contentType = normalizeWeComContentType(contentType)
|
||||||
|
ext := strings.ToLower(filepath.Ext(filename))
|
||||||
|
|
||||||
|
if partType == "file" {
|
||||||
|
if size <= wecomOutboundMediaMaxBytes {
|
||||||
|
return "file"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
if (partType == "image" || partType == "") && canWeComSendImage(contentType, ext, size) {
|
||||||
|
return "image"
|
||||||
|
}
|
||||||
|
if (partType == "audio" || partType == "voice" || partType == "") && canWeComSendVoice(contentType, ext, size) {
|
||||||
|
return "voice"
|
||||||
|
}
|
||||||
|
if (partType == "video" || partType == "") && canWeComSendVideo(contentType, ext, size) {
|
||||||
|
return "video"
|
||||||
|
}
|
||||||
|
if size <= wecomOutboundMediaMaxBytes {
|
||||||
|
return "file"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func trimWeComBytes(value string, limit int) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
if limit <= 0 || len(value) <= limit {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
size := 0
|
||||||
|
var out strings.Builder
|
||||||
|
for _, r := range value {
|
||||||
|
width := len(string(r))
|
||||||
|
if size+width > limit {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
size += width
|
||||||
|
out.WriteRune(r)
|
||||||
|
}
|
||||||
|
return out.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func ensureWeComOutboundFilename(filename, localPath, contentType string) string {
|
||||||
|
filename = sanitizeWeComFilename(filename)
|
||||||
|
if filename == "" {
|
||||||
|
filename = sanitizeWeComFilename(filepath.Base(localPath))
|
||||||
|
}
|
||||||
|
if filename == "" {
|
||||||
|
filename = "media"
|
||||||
|
}
|
||||||
|
if filepath.Ext(filename) == "" {
|
||||||
|
fallbackExt := inferMediaExt(contentType, strings.ToLower(filepath.Ext(localPath)))
|
||||||
|
if fallbackExt != "" {
|
||||||
|
filename += fallbackExt
|
||||||
|
}
|
||||||
|
}
|
||||||
|
filename = trimWeComBytes(filename, 256)
|
||||||
|
if filename == "" {
|
||||||
|
return "media"
|
||||||
|
}
|
||||||
|
return filename
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildWeComVideoContent(mediaID, filename, description string) *wecomVideoContent {
|
||||||
|
title := strings.TrimSuffix(filename, filepath.Ext(filename))
|
||||||
|
title = trimWeComBytes(title, 64)
|
||||||
|
if title == "" {
|
||||||
|
title = "video"
|
||||||
|
}
|
||||||
|
description = trimWeComBytes(description, 512)
|
||||||
|
return &wecomVideoContent{
|
||||||
|
MediaID: mediaID,
|
||||||
|
Title: title,
|
||||||
|
Description: description,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeWeComEnvelopeBody[T any](env wecomEnvelope) (T, error) {
|
||||||
|
var out T
|
||||||
|
if len(env.Body) == 0 {
|
||||||
|
return out, fmt.Errorf("wecom response body is empty")
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(env.Body, &out); err != nil {
|
||||||
|
return out, fmt.Errorf("decode wecom response body: %w", err)
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) uploadOutboundMedia(
|
||||||
|
ctx context.Context,
|
||||||
|
localPath, filename, contentType string,
|
||||||
|
part bus.MediaPart,
|
||||||
|
) (*wecomOutboundMedia, error) {
|
||||||
|
_ = ctx
|
||||||
|
|
||||||
|
contentType = detectLocalWeComContentType(localPath, contentType)
|
||||||
|
filename = ensureWeComOutboundFilename(filename, localPath, contentType)
|
||||||
|
|
||||||
|
data, err := os.ReadFile(localPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("read media file: %w", err)
|
||||||
|
}
|
||||||
|
size := int64(len(data))
|
||||||
|
kind := outboundWeComMediaKind(part.Type, filename, contentType, size)
|
||||||
|
if kind == "" {
|
||||||
|
return nil, fmt.Errorf("unsupported wecom media type or size for %q", filename)
|
||||||
|
}
|
||||||
|
|
||||||
|
totalChunks := (len(data) + wecomUploadChunkMaxBytes - 1) / wecomUploadChunkMaxBytes
|
||||||
|
if totalChunks <= 0 || totalChunks > wecomUploadMaxChunks {
|
||||||
|
return nil, fmt.Errorf("wecom upload requires 1-%d chunks, got %d", wecomUploadMaxChunks, totalChunks)
|
||||||
|
}
|
||||||
|
|
||||||
|
sum := md5.Sum(data)
|
||||||
|
initEnv, err := c.sendCommandAck(wecomCommand{
|
||||||
|
Cmd: wecomCmdUploadMediaInit,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
Body: wecomUploadMediaInitBody{
|
||||||
|
Type: kind,
|
||||||
|
Filename: filename,
|
||||||
|
TotalSize: size,
|
||||||
|
TotalChunks: totalChunks,
|
||||||
|
MD5: hex.EncodeToString(sum[:]),
|
||||||
|
},
|
||||||
|
}, wecomUploadTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
initResp, err := decodeWeComEnvelopeBody[wecomUploadMediaInitResponse](initEnv)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(initResp.UploadID) == "" {
|
||||||
|
return nil, fmt.Errorf("wecom upload init returned empty upload_id")
|
||||||
|
}
|
||||||
|
|
||||||
|
for idx, offset := 0, 0; offset < len(data); idx, offset = idx+1, offset+wecomUploadChunkMaxBytes {
|
||||||
|
end := offset + wecomUploadChunkMaxBytes
|
||||||
|
if end > len(data) {
|
||||||
|
end = len(data)
|
||||||
|
}
|
||||||
|
sendErr := c.sendCommand(wecomCommand{
|
||||||
|
Cmd: wecomCmdUploadMediaChunk,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
Body: wecomUploadMediaChunkBody{
|
||||||
|
UploadID: initResp.UploadID,
|
||||||
|
ChunkIndex: idx,
|
||||||
|
Base64Data: base64.StdEncoding.EncodeToString(data[offset:end]),
|
||||||
|
},
|
||||||
|
}, wecomUploadTimeout)
|
||||||
|
if sendErr != nil {
|
||||||
|
return nil, sendErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
finishEnv, err := c.sendCommandAck(wecomCommand{
|
||||||
|
Cmd: wecomCmdUploadMediaEnd,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
Body: wecomUploadMediaFinishBody{
|
||||||
|
UploadID: initResp.UploadID,
|
||||||
|
},
|
||||||
|
}, wecomUploadTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
finishResp, err := decodeWeComEnvelopeBody[wecomUploadMediaFinishResponse](finishEnv)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(finishResp.MediaID) == "" {
|
||||||
|
return nil, fmt.Errorf("wecom upload finish returned empty media_id")
|
||||||
|
}
|
||||||
|
|
||||||
|
uploaded := &wecomOutboundMedia{
|
||||||
|
MsgType: kind,
|
||||||
|
MediaID: finishResp.MediaID,
|
||||||
|
}
|
||||||
|
if kind == "video" {
|
||||||
|
video := buildWeComVideoContent(finishResp.MediaID, filename, part.Caption)
|
||||||
|
uploaded.Title = video.Title
|
||||||
|
uploaded.Description = video.Description
|
||||||
|
}
|
||||||
|
return uploaded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func fallbackWeComMediaText(part bus.MediaPart, kind, filename string) string {
|
||||||
|
var lines []string
|
||||||
|
if caption := strings.TrimSpace(part.Caption); caption != "" {
|
||||||
|
lines = append(lines, caption)
|
||||||
|
}
|
||||||
|
|
||||||
|
label := kind
|
||||||
|
if label == "" {
|
||||||
|
label = "media"
|
||||||
|
}
|
||||||
|
if filename != "" {
|
||||||
|
lines = append(lines, fmt.Sprintf("[%s: %s]", label, filename))
|
||||||
|
} else {
|
||||||
|
lines = append(lines, fmt.Sprintf("[%s attachment]", label))
|
||||||
|
}
|
||||||
|
|
||||||
|
ref := strings.TrimSpace(part.Ref)
|
||||||
|
if strings.HasPrefix(ref, "http://") || strings.HasPrefix(ref, "https://") {
|
||||||
|
lines = append(lines, ref)
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(lines, "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) resolveMediaRoute(chatID string) (wecomTurn, uint32, bool) {
|
||||||
|
if turn, ok := c.getTurn(chatID); ok {
|
||||||
|
if time.Since(turn.CreatedAt) <= wecomStreamMaxDuration {
|
||||||
|
return turn, turn.ChatType, true
|
||||||
|
}
|
||||||
|
c.deleteTurn(chatID)
|
||||||
|
}
|
||||||
|
if route, ok := c.routes.Get(chatID); ok {
|
||||||
|
return wecomTurn{ChatID: route.ChatID, ChatType: route.ChatType}, route.ChatType, false
|
||||||
|
}
|
||||||
|
return wecomTurn{ChatID: chatID}, 0, false
|
||||||
|
}
|
||||||
180
pkg/channels/wecom/media_test.go
Normal file
180
pkg/channels/wecom/media_test.go
Normal file
|
|
@ -0,0 +1,180 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
basechannels "github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestStoreRemoteMedia_DetectsJPEGContentTypeFromBody(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const jpegBase64 = "/9j/4AAQSkZJRgABAQAAAQABAAD/2wBDAP//////////////////////////////////////////////////////////////////////////////////////" +
|
||||||
|
"//////////////////////////////////////////////////////////////////////////////////////////////2wBDAf//////////////////////////////////////////////////////////////////////////////////////" +
|
||||||
|
"//////////////////////////////////////////////////////////////////////////////////////////////wAARCAABAAEDASIAAhEBAxEB/8QAFQABAQAAAAAAAAAAAAAAAAAAAAb/xAAVEQEBAAAAAAAAAAAAAAAAAAAABf/aAAwDAQACEAMQAAAB6A//xAAVEAEBAAAAAAAAAAAAAAAAAAAAEf/aAAgBAQABBQJf/8QAFBEBAAAAAAAAAAAAAAAAAAAAEP/aAAgBAwEBPwF//8QAFBEBAAAAAAAAAAAAAAAAAAAAEP/aAAgBAgEBPwF//8QAFBABAAAAAAAAAAAAAAAAAAAAEP/aAAgBAQAGPwJf/8QAFBABAAAAAAAAAAAAAAAAAAAAEP/aAAgBAQABPyFf/9k="
|
||||||
|
|
||||||
|
jpegData := decodeTestBase64(t, jpegBase64)
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
ch := &WeComChannel{
|
||||||
|
BaseChannel: basechannels.NewBaseChannel("wecom", nil, nil, nil),
|
||||||
|
mediaClient: &http.Client{
|
||||||
|
Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||||
|
return &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Header: http.Header{"Content-Type": []string{"application/octet-stream"}},
|
||||||
|
Body: io.NopCloser(bytes.NewReader(jpegData)),
|
||||||
|
}, nil
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch.SetMediaStore(store)
|
||||||
|
|
||||||
|
ref, err := ch.storeRemoteMedia(context.Background(), "test-scope", "msg-1", "https://wecom.example/media", "", "")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("storeRemoteMedia returned error: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = store.ReleaseAll("test-scope")
|
||||||
|
})
|
||||||
|
|
||||||
|
_, meta, err := store.ResolveWithMeta(ref)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolve media ref: %v", err)
|
||||||
|
}
|
||||||
|
if meta.ContentType != "image/jpeg" {
|
||||||
|
t.Fatalf("expected image/jpeg content type, got %q", meta.ContentType)
|
||||||
|
}
|
||||||
|
if !strings.HasSuffix(meta.Filename, ".jpg") && !strings.HasSuffix(meta.Filename, ".jpeg") {
|
||||||
|
t.Fatalf("expected jpeg filename, got %q", meta.Filename)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetectWeComMediaMetadata_UsesFallbackExtensionWhenBodyUnknown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
filename, contentType := detectWeComMediaMetadata([]byte("not a real image"), "msg-2.pdf", "", "", "")
|
||||||
|
if filename != "msg-2.pdf" {
|
||||||
|
t.Fatalf("expected fallback filename to be preserved, got %q", filename)
|
||||||
|
}
|
||||||
|
if contentType != "application/pdf" {
|
||||||
|
t.Fatalf("expected application/pdf from fallback extension, got %q", contentType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStoreRemoteMedia_PreservesSuffixFromURL(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
docxLikeData := []byte("PK\x03\x04fake office payload")
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
ch := &WeComChannel{
|
||||||
|
BaseChannel: basechannels.NewBaseChannel("wecom", nil, nil, nil),
|
||||||
|
mediaClient: &http.Client{
|
||||||
|
Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||||
|
return &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Header: http.Header{"Content-Type": []string{"application/octet-stream"}},
|
||||||
|
Body: io.NopCloser(bytes.NewReader(docxLikeData)),
|
||||||
|
}, nil
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch.SetMediaStore(store)
|
||||||
|
|
||||||
|
ref, err := ch.storeRemoteMedia(
|
||||||
|
context.Background(),
|
||||||
|
"test-scope",
|
||||||
|
"msg-docx",
|
||||||
|
"https://wecom.example/media/report.docx?signature=1",
|
||||||
|
"",
|
||||||
|
".bin",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("storeRemoteMedia returned error: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = store.ReleaseAll("test-scope")
|
||||||
|
})
|
||||||
|
|
||||||
|
localPath, meta, err := store.ResolveWithMeta(ref)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolve media ref: %v", err)
|
||||||
|
}
|
||||||
|
if !strings.HasSuffix(meta.Filename, ".docx") {
|
||||||
|
t.Fatalf("expected docx filename, got %q", meta.Filename)
|
||||||
|
}
|
||||||
|
if !strings.HasSuffix(strings.ToLower(localPath), ".docx") {
|
||||||
|
t.Fatalf("expected docx temp path, got %q", localPath)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStoreRemoteMedia_PreservesSuffixFromContentDisposition(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
pptxLikeData := []byte("PK\x03\x04fake office payload")
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
ch := &WeComChannel{
|
||||||
|
BaseChannel: basechannels.NewBaseChannel("wecom", nil, nil, nil),
|
||||||
|
mediaClient: &http.Client{
|
||||||
|
Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||||
|
return &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Header: http.Header{
|
||||||
|
"Content-Type": []string{"application/octet-stream"},
|
||||||
|
"Content-Disposition": []string{`attachment; filename="slides.pptx"`},
|
||||||
|
},
|
||||||
|
Body: io.NopCloser(bytes.NewReader(pptxLikeData)),
|
||||||
|
}, nil
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch.SetMediaStore(store)
|
||||||
|
|
||||||
|
ref, err := ch.storeRemoteMedia(
|
||||||
|
context.Background(),
|
||||||
|
"test-scope",
|
||||||
|
"msg-pptx",
|
||||||
|
"https://wecom.example/media/download",
|
||||||
|
"",
|
||||||
|
".bin",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("storeRemoteMedia returned error: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = store.ReleaseAll("test-scope")
|
||||||
|
})
|
||||||
|
|
||||||
|
localPath, meta, err := store.ResolveWithMeta(ref)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolve media ref: %v", err)
|
||||||
|
}
|
||||||
|
if !strings.HasSuffix(meta.Filename, ".pptx") {
|
||||||
|
t.Fatalf("expected pptx filename, got %q", meta.Filename)
|
||||||
|
}
|
||||||
|
if !strings.HasSuffix(strings.ToLower(localPath), ".pptx") {
|
||||||
|
t.Fatalf("expected pptx temp path, got %q", localPath)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeTestBase64(t *testing.T, value string) []byte {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
data, err := io.ReadAll(base64.NewDecoder(base64.StdEncoding, strings.NewReader(value)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("decode base64 fixture: %v", err)
|
||||||
|
}
|
||||||
|
return data
|
||||||
|
}
|
||||||
|
|
||||||
|
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||||
|
|
||||||
|
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
return f(req)
|
||||||
|
}
|
||||||
173
pkg/channels/wecom/protocol.go
Normal file
173
pkg/channels/wecom/protocol.go
Normal file
|
|
@ -0,0 +1,173 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
const (
|
||||||
|
wecomDefaultWebSocketURL = "wss://openws.work.weixin.qq.com"
|
||||||
|
wecomCmdSubscribe = "aibot_subscribe"
|
||||||
|
wecomCmdPing = "ping"
|
||||||
|
wecomCmdMsgCallback = "aibot_msg_callback"
|
||||||
|
wecomCmdEventCallback = "aibot_event_callback"
|
||||||
|
wecomCmdRespondMsg = "aibot_respond_msg"
|
||||||
|
wecomCmdSendMsg = "aibot_send_msg"
|
||||||
|
wecomCmdUploadMediaInit = "aibot_upload_media_init"
|
||||||
|
wecomCmdUploadMediaChunk = "aibot_upload_media_chunk"
|
||||||
|
wecomCmdUploadMediaEnd = "aibot_upload_media_finish"
|
||||||
|
)
|
||||||
|
|
||||||
|
type wecomEnvelope struct {
|
||||||
|
Cmd string `json:"cmd,omitempty"`
|
||||||
|
Headers wecomHeaders `json:"headers"`
|
||||||
|
Body json.RawMessage `json:"body,omitempty"`
|
||||||
|
ErrCode int `json:"errcode,omitempty"`
|
||||||
|
ErrMsg string `json:"errmsg,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomHeaders struct {
|
||||||
|
ReqID string `json:"req_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomCommand struct {
|
||||||
|
Cmd string `json:"cmd"`
|
||||||
|
Headers wecomHeaders `json:"headers"`
|
||||||
|
Body any `json:"body,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomSendMsgBody struct {
|
||||||
|
ChatID string `json:"chatid"`
|
||||||
|
ChatType uint32 `json:"chat_type,omitempty"`
|
||||||
|
MsgType string `json:"msgtype"`
|
||||||
|
Markdown *wecomMarkdownContent `json:"markdown,omitempty"`
|
||||||
|
File *wecomMediaRefContent `json:"file,omitempty"`
|
||||||
|
Image *wecomMediaRefContent `json:"image,omitempty"`
|
||||||
|
Voice *wecomMediaRefContent `json:"voice,omitempty"`
|
||||||
|
Video *wecomVideoContent `json:"video,omitempty"`
|
||||||
|
TemplateCard map[string]any `json:"template_card,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomRespondMsgBody struct {
|
||||||
|
MsgType string `json:"msgtype"`
|
||||||
|
Stream *wecomStreamContent `json:"stream,omitempty"`
|
||||||
|
Markdown *wecomMarkdownContent `json:"markdown,omitempty"`
|
||||||
|
File *wecomMediaRefContent `json:"file,omitempty"`
|
||||||
|
Image *wecomMediaRefContent `json:"image,omitempty"`
|
||||||
|
Voice *wecomMediaRefContent `json:"voice,omitempty"`
|
||||||
|
Video *wecomVideoContent `json:"video,omitempty"`
|
||||||
|
TemplateCard map[string]any `json:"template_card,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomStreamContent struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Finish bool `json:"finish"`
|
||||||
|
Content string `json:"content,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomMarkdownContent struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomMediaRefContent struct {
|
||||||
|
MediaID string `json:"media_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomVideoContent struct {
|
||||||
|
MediaID string `json:"media_id"`
|
||||||
|
Title string `json:"title,omitempty"`
|
||||||
|
Description string `json:"description,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomUploadMediaInitBody struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Filename string `json:"filename"`
|
||||||
|
TotalSize int64 `json:"total_size"`
|
||||||
|
TotalChunks int `json:"total_chunks"`
|
||||||
|
MD5 string `json:"md5,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomUploadMediaInitResponse struct {
|
||||||
|
UploadID string `json:"upload_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomUploadMediaChunkBody struct {
|
||||||
|
UploadID string `json:"upload_id"`
|
||||||
|
ChunkIndex int `json:"chunk_index"`
|
||||||
|
Base64Data string `json:"base64_data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomUploadMediaFinishBody struct {
|
||||||
|
UploadID string `json:"upload_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomUploadMediaFinishResponse struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
MediaID string `json:"media_id"`
|
||||||
|
CreatedAt json.RawMessage `json:"created_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomIncomingMessage struct {
|
||||||
|
MsgID string `json:"msgid"`
|
||||||
|
AIBotID string `json:"aibotid"`
|
||||||
|
ChatID string `json:"chatid,omitempty"`
|
||||||
|
ChatType string `json:"chattype,omitempty"`
|
||||||
|
From struct {
|
||||||
|
UserID string `json:"userid"`
|
||||||
|
} `json:"from"`
|
||||||
|
MsgType string `json:"msgtype"`
|
||||||
|
Text *struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"text,omitempty"`
|
||||||
|
Image *struct {
|
||||||
|
URL string `json:"url"`
|
||||||
|
AESKey string `json:"aeskey,omitempty"`
|
||||||
|
} `json:"image,omitempty"`
|
||||||
|
File *struct {
|
||||||
|
URL string `json:"url"`
|
||||||
|
AESKey string `json:"aeskey,omitempty"`
|
||||||
|
} `json:"file,omitempty"`
|
||||||
|
Video *struct {
|
||||||
|
URL string `json:"url"`
|
||||||
|
AESKey string `json:"aeskey,omitempty"`
|
||||||
|
} `json:"video,omitempty"`
|
||||||
|
Voice *struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"voice,omitempty"`
|
||||||
|
Mixed *struct {
|
||||||
|
MsgItem []struct {
|
||||||
|
MsgType string `json:"msgtype"`
|
||||||
|
Text *struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"text,omitempty"`
|
||||||
|
Image *struct {
|
||||||
|
URL string `json:"url"`
|
||||||
|
AESKey string `json:"aeskey,omitempty"`
|
||||||
|
} `json:"image,omitempty"`
|
||||||
|
File *struct {
|
||||||
|
URL string `json:"url"`
|
||||||
|
AESKey string `json:"aeskey,omitempty"`
|
||||||
|
} `json:"file,omitempty"`
|
||||||
|
} `json:"msg_item"`
|
||||||
|
} `json:"mixed,omitempty"`
|
||||||
|
Quote *struct {
|
||||||
|
MsgType string `json:"msgtype"`
|
||||||
|
Text *struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"text,omitempty"`
|
||||||
|
} `json:"quote,omitempty"`
|
||||||
|
Event *struct {
|
||||||
|
EventType string `json:"eventtype"`
|
||||||
|
} `json:"event,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func incomingChatID(msg wecomIncomingMessage) string {
|
||||||
|
if msg.ChatID != "" {
|
||||||
|
return msg.ChatID
|
||||||
|
}
|
||||||
|
return msg.From.UserID
|
||||||
|
}
|
||||||
|
|
||||||
|
func incomingChatTypeCode(kind string) uint32 {
|
||||||
|
if kind == "group" {
|
||||||
|
return 2
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
113
pkg/channels/wecom/reqid_store.go
Normal file
113
pkg/channels/wecom/reqid_store.go
Normal file
|
|
@ -0,0 +1,113 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type wecomRoute struct {
|
||||||
|
ReqID string `json:"req_id"`
|
||||||
|
ChatID string `json:"chat_id"`
|
||||||
|
ChatType uint32 `json:"chat_type"`
|
||||||
|
ExpiresAt time.Time `json:"expires_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type reqIDStore struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
path string
|
||||||
|
routes map[string]wecomRoute
|
||||||
|
}
|
||||||
|
|
||||||
|
func newReqIDStore(path string) *reqIDStore {
|
||||||
|
if path == "" {
|
||||||
|
path = defaultReqIDStorePath()
|
||||||
|
}
|
||||||
|
s := &reqIDStore{
|
||||||
|
path: path,
|
||||||
|
routes: make(map[string]wecomRoute),
|
||||||
|
}
|
||||||
|
_ = s.load()
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultReqIDStorePath() string {
|
||||||
|
if home, err := os.UserHomeDir(); err == nil && home != "" {
|
||||||
|
return filepath.Join(home, ".picoclaw", "wecom", "reqid-store.json")
|
||||||
|
}
|
||||||
|
return filepath.Join(os.TempDir(), "picoclaw-wecom-reqid-store.json")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *reqIDStore) Put(chatID, reqID string, chatType uint32, ttl time.Duration) error {
|
||||||
|
if reqID == "" || chatID == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.deleteExpiredLocked(time.Now())
|
||||||
|
s.routes[chatID] = wecomRoute{
|
||||||
|
ReqID: reqID,
|
||||||
|
ChatID: chatID,
|
||||||
|
ChatType: chatType,
|
||||||
|
ExpiresAt: time.Now().Add(ttl),
|
||||||
|
}
|
||||||
|
return s.saveLocked()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *reqIDStore) Get(chatID string) (wecomRoute, bool) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.deleteExpiredLocked(time.Now())
|
||||||
|
route, ok := s.routes[chatID]
|
||||||
|
return route, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *reqIDStore) Delete(chatID string) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
delete(s.routes, chatID)
|
||||||
|
return s.saveLocked()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *reqIDStore) load() error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
|
data, err := os.ReadFile(s.path)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, os.ErrNotExist) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
var routes map[string]wecomRoute
|
||||||
|
if err := json.Unmarshal(data, &routes); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.routes = routes
|
||||||
|
s.deleteExpiredLocked(time.Now())
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *reqIDStore) deleteExpiredLocked(now time.Time) {
|
||||||
|
for chatID, route := range s.routes {
|
||||||
|
if !route.ExpiresAt.IsZero() && now.After(route.ExpiresAt) {
|
||||||
|
delete(s.routes, chatID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *reqIDStore) saveLocked() error {
|
||||||
|
if err := os.MkdirAll(filepath.Dir(s.path), 0o700); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
data, err := json.MarshalIndent(s.routes, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return os.WriteFile(s.path, data, 0o600)
|
||||||
|
}
|
||||||
24
pkg/channels/wecom/reqid_store_test.go
Normal file
24
pkg/channels/wecom/reqid_store_test.go
Normal file
|
|
@ -0,0 +1,24 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestReqIDStorePersistsRoutes(t *testing.T) {
|
||||||
|
storePath := filepath.Join(t.TempDir(), "reqids.json")
|
||||||
|
store := newReqIDStore(storePath)
|
||||||
|
if err := store.Put("chat-1", "req-1", 2, time.Hour); err != nil {
|
||||||
|
t.Fatalf("Put() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
reloaded := newReqIDStore(storePath)
|
||||||
|
route, ok := reloaded.Get("chat-1")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected persisted route to be loaded")
|
||||||
|
}
|
||||||
|
if route.ChatID != "chat-1" || route.ReqID != "req-1" || route.ChatType != 2 {
|
||||||
|
t.Fatalf("loaded route = %+v", route)
|
||||||
|
}
|
||||||
|
}
|
||||||
970
pkg/channels/wecom/wecom.go
Normal file
970
pkg/channels/wecom/wecom.go
Normal file
|
|
@ -0,0 +1,970 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"math/big"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"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/identity"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
wecomConnectTimeout = 15 * time.Second
|
||||||
|
wecomCommandTimeout = 10 * time.Second
|
||||||
|
wecomUploadTimeout = 30 * time.Second
|
||||||
|
wecomHeartbeatInterval = 30 * time.Second
|
||||||
|
wecomStreamMaxDuration = 5*time.Minute + 30*time.Second
|
||||||
|
wecomStreamMinInterval = 500 * time.Millisecond
|
||||||
|
wecomRouteTTL = 30 * time.Minute
|
||||||
|
wecomMediaTimeout = 30 * time.Second
|
||||||
|
wecomRecentMessageMax = 1000
|
||||||
|
)
|
||||||
|
|
||||||
|
type WeComChannel struct {
|
||||||
|
*channels.BaseChannel
|
||||||
|
config config.WeComConfig
|
||||||
|
|
||||||
|
ctx context.Context
|
||||||
|
cancel context.CancelFunc
|
||||||
|
|
||||||
|
conn *websocket.Conn
|
||||||
|
connMu sync.Mutex
|
||||||
|
|
||||||
|
pendingMu sync.Mutex
|
||||||
|
pending map[string]chan wecomEnvelope
|
||||||
|
|
||||||
|
turnsMu sync.Mutex
|
||||||
|
turns map[string][]wecomTurn
|
||||||
|
|
||||||
|
recent *recentMessageSet
|
||||||
|
routes *reqIDStore
|
||||||
|
mediaClient *http.Client
|
||||||
|
commandSend func(wecomCommand, time.Duration) (wecomEnvelope, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomTurn struct {
|
||||||
|
ReqID string
|
||||||
|
ChatID string
|
||||||
|
ChatType uint32
|
||||||
|
StreamID string
|
||||||
|
CreatedAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type wecomStreamer struct {
|
||||||
|
channel *WeComChannel
|
||||||
|
chatID string
|
||||||
|
turn wecomTurn
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
closed bool
|
||||||
|
lastSentAt time.Time
|
||||||
|
content string
|
||||||
|
}
|
||||||
|
|
||||||
|
type recentMessageSet struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
seen map[string]struct{}
|
||||||
|
ring []string
|
||||||
|
idx int
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRecentMessageSet(capacity int) *recentMessageSet {
|
||||||
|
if capacity <= 0 {
|
||||||
|
capacity = wecomRecentMessageMax
|
||||||
|
}
|
||||||
|
return &recentMessageSet{
|
||||||
|
seen: make(map[string]struct{}, capacity),
|
||||||
|
ring: make([]string, capacity),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *recentMessageSet) Mark(id string) bool {
|
||||||
|
if id == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
if _, ok := s.seen[id]; ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if old := s.ring[s.idx]; old != "" {
|
||||||
|
delete(s.seen, old)
|
||||||
|
}
|
||||||
|
s.ring[s.idx] = id
|
||||||
|
s.idx = (s.idx + 1) % len(s.ring)
|
||||||
|
s.seen[id] = struct{}{}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewChannel(cfg config.WeComConfig, messageBus *bus.MessageBus) (*WeComChannel, error) {
|
||||||
|
if cfg.BotID == "" || cfg.Secret() == "" {
|
||||||
|
return nil, fmt.Errorf("wecom bot_id and secret are required")
|
||||||
|
}
|
||||||
|
if cfg.WebSocketURL == "" {
|
||||||
|
cfg.WebSocketURL = wecomDefaultWebSocketURL
|
||||||
|
}
|
||||||
|
|
||||||
|
base := channels.NewBaseChannel(
|
||||||
|
"wecom",
|
||||||
|
cfg,
|
||||||
|
messageBus,
|
||||||
|
cfg.AllowFrom,
|
||||||
|
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
||||||
|
)
|
||||||
|
|
||||||
|
ch := &WeComChannel{
|
||||||
|
BaseChannel: base,
|
||||||
|
config: cfg,
|
||||||
|
pending: make(map[string]chan wecomEnvelope),
|
||||||
|
turns: make(map[string][]wecomTurn),
|
||||||
|
recent: newRecentMessageSet(wecomRecentMessageMax),
|
||||||
|
routes: newReqIDStore(""),
|
||||||
|
mediaClient: &http.Client{Timeout: wecomMediaTimeout},
|
||||||
|
}
|
||||||
|
ch.SetOwner(ch)
|
||||||
|
return ch, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) Name() string { return "wecom" }
|
||||||
|
|
||||||
|
func (c *WeComChannel) Start(ctx context.Context) error {
|
||||||
|
logger.InfoC("wecom", "Starting WeCom channel...")
|
||||||
|
c.ctx, c.cancel = context.WithCancel(ctx)
|
||||||
|
c.SetRunning(true)
|
||||||
|
go c.connectLoop()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) Stop(_ context.Context) error {
|
||||||
|
logger.InfoC("wecom", "Stopping WeCom channel...")
|
||||||
|
if c.cancel != nil {
|
||||||
|
c.cancel()
|
||||||
|
}
|
||||||
|
c.connMu.Lock()
|
||||||
|
if c.conn != nil {
|
||||||
|
_ = c.conn.Close()
|
||||||
|
c.conn = nil
|
||||||
|
}
|
||||||
|
c.connMu.Unlock()
|
||||||
|
c.clearTurns()
|
||||||
|
c.SetRunning(false)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) BeginStream(_ context.Context, chatID string) (channels.Streamer, error) {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return nil, channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
|
||||||
|
turn, ok := c.getTurn(chatID)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("wecom streaming unavailable: no active turn")
|
||||||
|
}
|
||||||
|
if time.Since(turn.CreatedAt) > wecomStreamMaxDuration {
|
||||||
|
c.consumeTurn(chatID, turn)
|
||||||
|
return nil, fmt.Errorf("wecom streaming unavailable: turn expired")
|
||||||
|
}
|
||||||
|
|
||||||
|
return &wecomStreamer{
|
||||||
|
channel: c,
|
||||||
|
chatID: chatID,
|
||||||
|
turn: turn,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
content := strings.TrimSpace(msg.Content)
|
||||||
|
if content == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if turn, ok := c.getTurn(msg.ChatID); ok {
|
||||||
|
if time.Since(turn.CreatedAt) <= wecomStreamMaxDuration {
|
||||||
|
if err := c.sendStreamReply(turn, content); err == nil {
|
||||||
|
c.consumeTurn(msg.ChatID, turn)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
c.consumeTurn(msg.ChatID, turn)
|
||||||
|
}
|
||||||
|
|
||||||
|
if route, ok := c.routes.Get(msg.ChatID); ok {
|
||||||
|
if err := c.sendActivePush(route.ChatID, route.ChatType, content); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := c.sendActivePush(msg.ChatID, 0, content); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
|
||||||
|
route, chatType, hasTurn := c.resolveMediaRoute(msg.ChatID)
|
||||||
|
chatID := route.ChatID
|
||||||
|
if chatID == "" {
|
||||||
|
chatID = msg.ChatID
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, part := range msg.Parts {
|
||||||
|
if strings.TrimSpace(part.Ref) == "" {
|
||||||
|
if caption := strings.TrimSpace(part.Caption); caption != "" {
|
||||||
|
if err := c.sendActivePush(chatID, chatType, caption); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
localPath, filename, contentType, cleanup, err := c.resolveOutboundPart(ctx, part)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("wecom resolve media %q: %v: %w", part.Ref, err, channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func() {
|
||||||
|
if cleanup != nil {
|
||||||
|
defer cleanup()
|
||||||
|
}
|
||||||
|
|
||||||
|
uploaded, uploadErr := c.uploadOutboundMedia(ctx, localPath, filename, contentType, part)
|
||||||
|
if uploadErr != nil {
|
||||||
|
logger.WarnCF("wecom", "Falling back to placeholder after media upload failure", map[string]any{
|
||||||
|
"chat_id": chatID,
|
||||||
|
"ref": part.Ref,
|
||||||
|
"filename": filename,
|
||||||
|
"content_type": contentType,
|
||||||
|
"error": uploadErr.Error(),
|
||||||
|
})
|
||||||
|
if hasTurn {
|
||||||
|
if finishErr := c.sendStreamChunk(route, true, ""); finishErr != nil {
|
||||||
|
err = finishErr
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.deleteTurn(msg.ChatID)
|
||||||
|
hasTurn = false
|
||||||
|
}
|
||||||
|
err = c.sendActivePush(chatID, chatType, fallbackWeComMediaText(part, "", filename))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if hasTurn {
|
||||||
|
err = c.sendTurnMedia(route, uploaded)
|
||||||
|
c.deleteTurn(msg.ChatID)
|
||||||
|
hasTurn = false
|
||||||
|
} else {
|
||||||
|
err = c.sendActiveMedia(chatID, chatType, uploaded)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if caption := strings.TrimSpace(part.Caption); caption != "" {
|
||||||
|
err = c.sendActivePush(chatID, chatType, caption)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) connectLoop() {
|
||||||
|
backoff := time.Second
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := c.runConnection(); err != nil {
|
||||||
|
logger.WarnCF("wecom", "WeCom connection lost", map[string]any{
|
||||||
|
"error": err.Error(),
|
||||||
|
"backoff": backoff.String(),
|
||||||
|
})
|
||||||
|
select {
|
||||||
|
case <-time.After(backoff):
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if backoff < time.Minute {
|
||||||
|
backoff *= 2
|
||||||
|
if backoff > time.Minute {
|
||||||
|
backoff = time.Minute
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) runConnection() error {
|
||||||
|
dialCtx, cancel := context.WithTimeout(c.ctx, wecomConnectTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
conn, resp, err := websocket.DefaultDialer.DialContext(dialCtx, c.config.WebSocketURL, nil)
|
||||||
|
if resp != nil {
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.connMu.Lock()
|
||||||
|
c.conn = conn
|
||||||
|
c.connMu.Unlock()
|
||||||
|
defer func() {
|
||||||
|
c.connMu.Lock()
|
||||||
|
if c.conn == conn {
|
||||||
|
c.conn = nil
|
||||||
|
}
|
||||||
|
c.connMu.Unlock()
|
||||||
|
_ = conn.Close()
|
||||||
|
c.clearTurns()
|
||||||
|
}()
|
||||||
|
|
||||||
|
readErrCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
readErrCh <- c.readLoop(conn)
|
||||||
|
}()
|
||||||
|
|
||||||
|
if writeErr := c.writeAndWait(conn, wecomCommand{
|
||||||
|
Cmd: wecomCmdSubscribe,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
Body: map[string]string{
|
||||||
|
"bot_id": c.config.BotID,
|
||||||
|
"secret": c.config.Secret(),
|
||||||
|
},
|
||||||
|
}, wecomCommandTimeout); writeErr != nil {
|
||||||
|
return writeErr
|
||||||
|
}
|
||||||
|
|
||||||
|
heartbeatDone := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(heartbeatDone)
|
||||||
|
c.heartbeatLoop(conn)
|
||||||
|
}()
|
||||||
|
|
||||||
|
err = <-readErrCh
|
||||||
|
_ = conn.Close()
|
||||||
|
<-heartbeatDone
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) heartbeatLoop(conn *websocket.Conn) {
|
||||||
|
ticker := time.NewTicker(wecomHeartbeatInterval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ticker.C:
|
||||||
|
if err := c.writeAndWait(conn, wecomCommand{
|
||||||
|
Cmd: wecomCmdPing,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
}, wecomCommandTimeout); err != nil {
|
||||||
|
logger.WarnCF("wecom", "Heartbeat failed", map[string]any{"error": err.Error()})
|
||||||
|
_ = conn.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) readLoop(conn *websocket.Conn) error {
|
||||||
|
for {
|
||||||
|
_, raw, err := conn.ReadMessage()
|
||||||
|
if err != nil {
|
||||||
|
select {
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return nil
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var env wecomEnvelope
|
||||||
|
if err := json.Unmarshal(raw, &env); err != nil {
|
||||||
|
logger.WarnCF("wecom", "Failed to parse WebSocket message", map[string]any{"error": err.Error()})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if env.Cmd == "" && env.Headers.ReqID != "" {
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
ch, ok := c.pending[env.Headers.ReqID]
|
||||||
|
if ok {
|
||||||
|
delete(c.pending, env.Headers.ReqID)
|
||||||
|
}
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
if ok {
|
||||||
|
ch <- env
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
go c.handleEnvelope(env)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) handleEnvelope(env wecomEnvelope) {
|
||||||
|
switch env.Cmd {
|
||||||
|
case wecomCmdMsgCallback:
|
||||||
|
c.handleMessageCallback(env)
|
||||||
|
case wecomCmdEventCallback:
|
||||||
|
c.handleEventCallback(env)
|
||||||
|
default:
|
||||||
|
logger.DebugCF("wecom", "Ignoring unsupported WeCom command", map[string]any{"cmd": env.Cmd})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) handleEventCallback(env wecomEnvelope) {
|
||||||
|
var msg wecomIncomingMessage
|
||||||
|
if err := json.Unmarshal(env.Body, &msg); err != nil {
|
||||||
|
logger.WarnCF("wecom", "Failed to parse WeCom event callback", map[string]any{"error": err.Error()})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) handleMessageCallback(env wecomEnvelope) {
|
||||||
|
var msg wecomIncomingMessage
|
||||||
|
if err := json.Unmarshal(env.Body, &msg); err != nil {
|
||||||
|
logger.WarnCF("wecom", "Failed to parse WeCom message callback", map[string]any{"error": err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !c.recent.Mark(msg.MsgID) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
reqID := env.Headers.ReqID
|
||||||
|
if reqID == "" {
|
||||||
|
logger.WarnC("wecom", "WeCom message callback missing req_id")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if msg.Event != nil && msg.Event.EventType != "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := c.dispatchIncoming(reqID, msg); err != nil {
|
||||||
|
logger.WarnCF("wecom", "Failed to dispatch WeCom message", map[string]any{
|
||||||
|
"req_id": reqID,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
_ = c.respondImmediate(reqID, "The WeCom message could not be processed.")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) dispatchIncoming(reqID string, msg wecomIncomingMessage) error {
|
||||||
|
senderID := msg.From.UserID
|
||||||
|
if senderID == "" {
|
||||||
|
senderID = "unknown"
|
||||||
|
}
|
||||||
|
actualChatID := incomingChatID(msg)
|
||||||
|
chatType := incomingChatTypeCode(msg.ChatType)
|
||||||
|
peerKind := "direct"
|
||||||
|
if msg.ChatType == "group" {
|
||||||
|
peerKind = "group"
|
||||||
|
}
|
||||||
|
|
||||||
|
sender := bus.SenderInfo{
|
||||||
|
Platform: "wecom",
|
||||||
|
PlatformID: senderID,
|
||||||
|
CanonicalID: identity.BuildCanonicalID("wecom", senderID),
|
||||||
|
DisplayName: senderID,
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
content string
|
||||||
|
quoteText string
|
||||||
|
mediaRefs []string
|
||||||
|
err error
|
||||||
|
)
|
||||||
|
scope := channels.BuildMediaScope("wecom", actualChatID, msg.MsgID)
|
||||||
|
switch msg.MsgType {
|
||||||
|
case "text":
|
||||||
|
if msg.Text != nil {
|
||||||
|
content = strings.TrimSpace(msg.Text.Content)
|
||||||
|
}
|
||||||
|
case "voice":
|
||||||
|
if msg.Voice != nil {
|
||||||
|
content = strings.TrimSpace(msg.Voice.Content)
|
||||||
|
}
|
||||||
|
case "image":
|
||||||
|
content = "[image]"
|
||||||
|
mediaRefs, err = c.collectSingleMedia(c.ctx, scope, msg.MsgID, &mediaPayload{
|
||||||
|
url: msg.Image.URL,
|
||||||
|
aesKey: msg.Image.AESKey,
|
||||||
|
}, "image", ".jpg")
|
||||||
|
case "file":
|
||||||
|
content = "[file]"
|
||||||
|
mediaRefs, err = c.collectSingleMedia(c.ctx, scope, msg.MsgID, &mediaPayload{
|
||||||
|
url: msg.File.URL,
|
||||||
|
aesKey: msg.File.AESKey,
|
||||||
|
}, "file", ".bin")
|
||||||
|
case "video":
|
||||||
|
content = "[video]"
|
||||||
|
mediaRefs, err = c.collectSingleMedia(c.ctx, scope, msg.MsgID, &mediaPayload{
|
||||||
|
url: msg.Video.URL,
|
||||||
|
aesKey: msg.Video.AESKey,
|
||||||
|
}, "video", ".mp4")
|
||||||
|
case "mixed":
|
||||||
|
content, mediaRefs, err = c.collectMixedMedia(c.ctx, scope, msg)
|
||||||
|
default:
|
||||||
|
return c.respondImmediate(reqID, "Unsupported WeCom message type: "+msg.MsgType)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if msg.Quote != nil && msg.Quote.Text != nil {
|
||||||
|
quoteText = strings.TrimSpace(msg.Quote.Text.Content)
|
||||||
|
if content == "" {
|
||||||
|
content = quoteText
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if content == "" && len(mediaRefs) == 0 {
|
||||||
|
return c.respondImmediate(reqID, "The WeCom message did not contain usable content.")
|
||||||
|
}
|
||||||
|
|
||||||
|
turn := wecomTurn{
|
||||||
|
ReqID: reqID,
|
||||||
|
ChatID: actualChatID,
|
||||||
|
ChatType: chatType,
|
||||||
|
StreamID: randomID(10),
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
}
|
||||||
|
c.queueTurn(actualChatID, turn)
|
||||||
|
if err := c.routes.Put(actualChatID, reqID, chatType, wecomRouteTTL); err != nil {
|
||||||
|
logger.WarnCF("wecom", "Failed to persist req_id route", map[string]any{
|
||||||
|
"chat_id": actualChatID,
|
||||||
|
"req_id": reqID,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
opening := ""
|
||||||
|
if c.config.SendThinkingMessage {
|
||||||
|
opening = "Processing..."
|
||||||
|
}
|
||||||
|
if err := c.sendStreamChunk(turn, false, opening); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
peer := bus.Peer{Kind: peerKind, ID: actualChatID}
|
||||||
|
metadata := map[string]string{
|
||||||
|
"channel": "wecom",
|
||||||
|
"req_id": reqID,
|
||||||
|
"chat_id": actualChatID,
|
||||||
|
"chat_type": msg.ChatType,
|
||||||
|
"msg_id": msg.MsgID,
|
||||||
|
"msg_type": msg.MsgType,
|
||||||
|
}
|
||||||
|
if quoteText != "" {
|
||||||
|
metadata["quote_text"] = quoteText
|
||||||
|
}
|
||||||
|
|
||||||
|
c.HandleMessage(c.ctx, peer, msg.MsgID, senderID, actualChatID, content, mediaRefs, metadata, sender)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) collectSingleMedia(
|
||||||
|
ctx context.Context,
|
||||||
|
scope, msgID string,
|
||||||
|
payload interface {
|
||||||
|
GetURL() string
|
||||||
|
GetAESKey() string
|
||||||
|
},
|
||||||
|
label, fallbackExt string,
|
||||||
|
) ([]string, error) {
|
||||||
|
if payload == nil || payload.GetURL() == "" {
|
||||||
|
return nil, fmt.Errorf("%s payload is empty", label)
|
||||||
|
}
|
||||||
|
ref, err := c.storeRemoteMedia(ctx, scope, msgID, payload.GetURL(), payload.GetAESKey(), fallbackExt)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return []string{ref}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type mediaPayload struct {
|
||||||
|
url string
|
||||||
|
aesKey string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *mediaPayload) GetURL() string { return p.url }
|
||||||
|
func (p *mediaPayload) GetAESKey() string { return p.aesKey }
|
||||||
|
|
||||||
|
func (c *WeComChannel) collectMixedMedia(
|
||||||
|
ctx context.Context,
|
||||||
|
scope string,
|
||||||
|
msg wecomIncomingMessage,
|
||||||
|
) (string, []string, error) {
|
||||||
|
if msg.Mixed == nil {
|
||||||
|
return "", nil, fmt.Errorf("mixed message is empty")
|
||||||
|
}
|
||||||
|
|
||||||
|
var textParts []string
|
||||||
|
var refs []string
|
||||||
|
for idx, item := range msg.Mixed.MsgItem {
|
||||||
|
switch item.MsgType {
|
||||||
|
case "text":
|
||||||
|
if item.Text != nil && strings.TrimSpace(item.Text.Content) != "" {
|
||||||
|
textParts = append(textParts, strings.TrimSpace(item.Text.Content))
|
||||||
|
}
|
||||||
|
case "image":
|
||||||
|
if item.Image != nil && item.Image.URL != "" {
|
||||||
|
ref, err := c.storeRemoteMedia(
|
||||||
|
ctx,
|
||||||
|
scope,
|
||||||
|
fmt.Sprintf("%s-%d", msg.MsgID, idx),
|
||||||
|
item.Image.URL,
|
||||||
|
item.Image.AESKey,
|
||||||
|
".jpg",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
refs = append(refs, ref)
|
||||||
|
}
|
||||||
|
case "file":
|
||||||
|
if item.File != nil && item.File.URL != "" {
|
||||||
|
ref, err := c.storeRemoteMedia(
|
||||||
|
ctx,
|
||||||
|
scope,
|
||||||
|
fmt.Sprintf("%s-%d", msg.MsgID, idx),
|
||||||
|
item.File.URL,
|
||||||
|
item.File.AESKey,
|
||||||
|
".bin",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
refs = append(refs, ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
content := strings.Join(textParts, "\n")
|
||||||
|
if content == "" && len(refs) > 0 {
|
||||||
|
content = "[media]"
|
||||||
|
}
|
||||||
|
return content, refs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) respondImmediate(reqID, content string) error {
|
||||||
|
turn := wecomTurn{
|
||||||
|
ReqID: reqID,
|
||||||
|
StreamID: randomID(10),
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
}
|
||||||
|
return c.sendStreamChunk(turn, true, content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendStreamReply(turn wecomTurn, content string) error {
|
||||||
|
return c.sendStreamChunk(turn, true, content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendStreamChunk(turn wecomTurn, finish bool, content string) error {
|
||||||
|
return c.sendCommand(wecomCommand{
|
||||||
|
Cmd: wecomCmdRespondMsg,
|
||||||
|
Headers: wecomHeaders{ReqID: turn.ReqID},
|
||||||
|
Body: wecomRespondMsgBody{
|
||||||
|
MsgType: "stream",
|
||||||
|
Stream: &wecomStreamContent{
|
||||||
|
ID: turn.StreamID,
|
||||||
|
Finish: finish,
|
||||||
|
Content: content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}, wecomCommandTimeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendTurnMedia(turn wecomTurn, uploaded *wecomOutboundMedia) error {
|
||||||
|
if uploaded == nil {
|
||||||
|
return fmt.Errorf("wecom outbound media is nil: %w", channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
if err := c.sendCommand(wecomCommand{
|
||||||
|
Cmd: wecomCmdRespondMsg,
|
||||||
|
Headers: wecomHeaders{ReqID: turn.ReqID},
|
||||||
|
Body: uploaded.respondBody(),
|
||||||
|
}, wecomCommandTimeout); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return c.sendStreamChunk(turn, true, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendActivePush(chatID string, chatType uint32, content string) error {
|
||||||
|
if strings.TrimSpace(chatID) == "" {
|
||||||
|
return fmt.Errorf("empty chat ID: %w", channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
return c.sendCommand(wecomCommand{
|
||||||
|
Cmd: wecomCmdSendMsg,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
Body: wecomSendMsgBody{
|
||||||
|
ChatID: chatID,
|
||||||
|
ChatType: chatType,
|
||||||
|
MsgType: "markdown",
|
||||||
|
Markdown: &wecomMarkdownContent{Content: content},
|
||||||
|
},
|
||||||
|
}, wecomCommandTimeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendActiveMedia(chatID string, chatType uint32, uploaded *wecomOutboundMedia) error {
|
||||||
|
if strings.TrimSpace(chatID) == "" {
|
||||||
|
return fmt.Errorf("empty chat ID: %w", channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
if uploaded == nil {
|
||||||
|
return fmt.Errorf("wecom outbound media is nil: %w", channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
return c.sendCommand(wecomCommand{
|
||||||
|
Cmd: wecomCmdSendMsg,
|
||||||
|
Headers: wecomHeaders{ReqID: randomID(10)},
|
||||||
|
Body: uploaded.sendBody(chatID, chatType),
|
||||||
|
}, wecomCommandTimeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendCommand(cmd wecomCommand, timeout time.Duration) error {
|
||||||
|
_, err := c.sendCommandAck(cmd, timeout)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) sendCommandAck(cmd wecomCommand, timeout time.Duration) (wecomEnvelope, error) {
|
||||||
|
if c.commandSend != nil {
|
||||||
|
return c.commandSend(cmd, timeout)
|
||||||
|
}
|
||||||
|
return c.writeCurrentAck(cmd, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) writeCurrentAck(cmd wecomCommand, timeout time.Duration) (wecomEnvelope, error) {
|
||||||
|
c.connMu.Lock()
|
||||||
|
conn := c.conn
|
||||||
|
c.connMu.Unlock()
|
||||||
|
if conn == nil {
|
||||||
|
return wecomEnvelope{}, fmt.Errorf("wecom websocket not connected: %w", channels.ErrTemporary)
|
||||||
|
}
|
||||||
|
return c.writeAndWaitAck(conn, cmd, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) writeAndWait(conn *websocket.Conn, cmd wecomCommand, timeout time.Duration) error {
|
||||||
|
_, err := c.writeAndWaitAck(conn, cmd, timeout)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) writeAndWaitAck(
|
||||||
|
conn *websocket.Conn,
|
||||||
|
cmd wecomCommand,
|
||||||
|
timeout time.Duration,
|
||||||
|
) (wecomEnvelope, error) {
|
||||||
|
if cmd.Headers.ReqID == "" {
|
||||||
|
cmd.Headers.ReqID = randomID(10)
|
||||||
|
}
|
||||||
|
waitCh := make(chan wecomEnvelope, 1)
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
c.pending[cmd.Headers.ReqID] = waitCh
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
defer func() {
|
||||||
|
c.pendingMu.Lock()
|
||||||
|
delete(c.pending, cmd.Headers.ReqID)
|
||||||
|
c.pendingMu.Unlock()
|
||||||
|
}()
|
||||||
|
|
||||||
|
data, err := json.Marshal(cmd)
|
||||||
|
if err != nil {
|
||||||
|
return wecomEnvelope{}, fmt.Errorf("%w: %v", channels.ErrSendFailed, err)
|
||||||
|
}
|
||||||
|
c.connMu.Lock()
|
||||||
|
err = conn.WriteMessage(websocket.TextMessage, data)
|
||||||
|
c.connMu.Unlock()
|
||||||
|
if err != nil {
|
||||||
|
return wecomEnvelope{}, fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
timer := time.NewTimer(timeout)
|
||||||
|
defer timer.Stop()
|
||||||
|
select {
|
||||||
|
case env := <-waitCh:
|
||||||
|
if env.ErrCode != 0 {
|
||||||
|
return wecomEnvelope{}, fmt.Errorf(
|
||||||
|
"%w: wecom errcode=%d errmsg=%s",
|
||||||
|
channels.ErrTemporary,
|
||||||
|
env.ErrCode,
|
||||||
|
env.ErrMsg,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return env, nil
|
||||||
|
case <-timer.C:
|
||||||
|
return wecomEnvelope{}, fmt.Errorf("%w: timeout waiting for WeCom ack", channels.ErrTemporary)
|
||||||
|
case <-c.ctx.Done():
|
||||||
|
return wecomEnvelope{}, c.ctx.Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) getTurn(chatID string) (wecomTurn, bool) {
|
||||||
|
c.turnsMu.Lock()
|
||||||
|
defer c.turnsMu.Unlock()
|
||||||
|
queue := c.turns[chatID]
|
||||||
|
if len(queue) == 0 {
|
||||||
|
return wecomTurn{}, false
|
||||||
|
}
|
||||||
|
return queue[0], true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) deleteTurn(chatID string) {
|
||||||
|
c.turnsMu.Lock()
|
||||||
|
defer c.turnsMu.Unlock()
|
||||||
|
queue := c.turns[chatID]
|
||||||
|
if len(queue) <= 1 {
|
||||||
|
delete(c.turns, chatID)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.turns[chatID] = queue[1:]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) queueTurn(chatID string, turn wecomTurn) {
|
||||||
|
c.turnsMu.Lock()
|
||||||
|
defer c.turnsMu.Unlock()
|
||||||
|
c.turns[chatID] = append(c.turns[chatID], turn)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) consumeTurn(chatID string, turn wecomTurn) bool {
|
||||||
|
c.turnsMu.Lock()
|
||||||
|
defer c.turnsMu.Unlock()
|
||||||
|
|
||||||
|
queue := c.turns[chatID]
|
||||||
|
if len(queue) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
current := queue[0]
|
||||||
|
if current.ReqID != turn.ReqID || current.StreamID != turn.StreamID {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if len(queue) == 1 {
|
||||||
|
delete(c.turns, chatID)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
c.turns[chatID] = queue[1:]
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *WeComChannel) clearTurns() {
|
||||||
|
c.turnsMu.Lock()
|
||||||
|
c.turns = make(map[string][]wecomTurn)
|
||||||
|
c.turnsMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func randomID(n int) string {
|
||||||
|
const alphabet = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
|
||||||
|
if n <= 0 {
|
||||||
|
n = 10
|
||||||
|
}
|
||||||
|
buf := make([]byte, n)
|
||||||
|
for i := range buf {
|
||||||
|
v, _ := rand.Int(rand.Reader, big.NewInt(int64(len(alphabet))))
|
||||||
|
buf[i] = alphabet[v.Int64()]
|
||||||
|
}
|
||||||
|
return string(buf)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *wecomStreamer) Update(ctx context.Context, content string) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
|
if s.closed {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err := s.validateActiveTurn(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if !s.lastSentAt.IsZero() {
|
||||||
|
wait := time.Until(s.lastSentAt.Add(wecomStreamMinInterval))
|
||||||
|
if wait > 0 {
|
||||||
|
timer := time.NewTimer(wait)
|
||||||
|
defer timer.Stop()
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return ctx.Err()
|
||||||
|
case <-timer.C:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := s.channel.sendStreamChunk(s.turn, false, content); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.content = content
|
||||||
|
s.lastSentAt = time.Now()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *wecomStreamer) Finalize(ctx context.Context, content string) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
|
if s.closed {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err := s.validateActiveTurn(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := s.channel.sendStreamChunk(s.turn, true, content); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.content = content
|
||||||
|
s.closed = true
|
||||||
|
s.channel.consumeTurn(s.chatID, s.turn)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *wecomStreamer) Cancel(_ context.Context) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
|
if s.closed {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.validateActiveTurn() == nil {
|
||||||
|
_ = s.channel.sendStreamChunk(s.turn, true, s.content)
|
||||||
|
s.channel.consumeTurn(s.chatID, s.turn)
|
||||||
|
}
|
||||||
|
s.closed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *wecomStreamer) validateActiveTurn() error {
|
||||||
|
if time.Since(s.turn.CreatedAt) > wecomStreamMaxDuration {
|
||||||
|
s.channel.consumeTurn(s.chatID, s.turn)
|
||||||
|
return fmt.Errorf("wecom streaming unavailable: turn expired")
|
||||||
|
}
|
||||||
|
current, ok := s.channel.getTurn(s.chatID)
|
||||||
|
if !ok || current.ReqID != s.turn.ReqID || current.StreamID != s.turn.StreamID {
|
||||||
|
return fmt.Errorf("wecom streaming unavailable: turn no longer active")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
660
pkg/channels/wecom/wecom_test.go
Normal file
660
pkg/channels/wecom/wecom_test.go
Normal file
|
|
@ -0,0 +1,660 @@
|
||||||
|
package wecom
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDispatchIncoming_UsesActualChatIDAndStoresReqIDRoute(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := newTestWeComChannel(t, messageBus)
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := wecomIncomingMessage{
|
||||||
|
MsgID: "msg-1",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
ChatType: "direct",
|
||||||
|
MsgType: "text",
|
||||||
|
Text: &struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
}{Content: "hello"},
|
||||||
|
}
|
||||||
|
msg.From.UserID = "user-1"
|
||||||
|
|
||||||
|
if err := ch.dispatchIncoming("req-1", msg); err != nil {
|
||||||
|
t.Fatalf("dispatchIncoming() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case inbound := <-messageBus.InboundChan():
|
||||||
|
if inbound.ChatID != "chat-1" {
|
||||||
|
t.Fatalf("inbound ChatID = %q, want chat-1", inbound.ChatID)
|
||||||
|
}
|
||||||
|
if inbound.MessageID != "msg-1" {
|
||||||
|
t.Fatalf("inbound MessageID = %q, want msg-1", inbound.MessageID)
|
||||||
|
}
|
||||||
|
if inbound.Peer.ID != "chat-1" {
|
||||||
|
t.Fatalf("inbound Peer.ID = %q, want chat-1", inbound.Peer.ID)
|
||||||
|
}
|
||||||
|
if inbound.Metadata["req_id"] != "req-1" {
|
||||||
|
t.Fatalf("inbound req_id = %q, want req-1", inbound.Metadata["req_id"])
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
t.Fatal("expected inbound message to be published")
|
||||||
|
}
|
||||||
|
|
||||||
|
turn, ok := ch.getTurn("chat-1")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected queued turn for chat-1")
|
||||||
|
}
|
||||||
|
if turn.ReqID != "req-1" {
|
||||||
|
t.Fatalf("turn.ReqID = %q, want req-1", turn.ReqID)
|
||||||
|
}
|
||||||
|
|
||||||
|
route, ok := ch.routes.Get("chat-1")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected persisted route for chat-1")
|
||||||
|
}
|
||||||
|
if route.ReqID != "req-1" || route.ChatType != 1 {
|
||||||
|
t.Fatalf("route = %+v", route)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 1 {
|
||||||
|
t.Fatalf("expected 1 opening command, got %d", len(commands))
|
||||||
|
}
|
||||||
|
if commands[0].Cmd != wecomCmdRespondMsg {
|
||||||
|
t.Fatalf("opening command = %q, want %q", commands[0].Cmd, wecomCmdRespondMsg)
|
||||||
|
}
|
||||||
|
if commands[0].Headers.ReqID != "req-1" {
|
||||||
|
t.Fatalf("opening req_id = %q, want req-1", commands[0].Headers.ReqID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewChannel_DoesNotRegisterMessageSplitLimit(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
if got := ch.MaxMessageLength(); got != 0 {
|
||||||
|
t.Fatalf("MaxMessageLength() = %d, want 0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBeginStream_UpdateAndFinalize(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
ch.queueTurn("chat-1", wecomTurn{
|
||||||
|
ReqID: "req-1",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
ChatType: 1,
|
||||||
|
StreamID: "stream-1",
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
})
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
streamer, err := ch.BeginStream(context.Background(), "chat-1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("BeginStream() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := streamer.Update(context.Background(), "draft"); err != nil {
|
||||||
|
t.Fatalf("Update() error = %v", err)
|
||||||
|
}
|
||||||
|
if err := streamer.Finalize(context.Background(), "final"); err != nil {
|
||||||
|
t.Fatalf("Finalize() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 2 {
|
||||||
|
t.Fatalf("expected 2 commands, got %d", len(commands))
|
||||||
|
}
|
||||||
|
for i, wantFinish := range []bool{false, true} {
|
||||||
|
if commands[i].Cmd != wecomCmdRespondMsg {
|
||||||
|
t.Fatalf("command[%d].Cmd = %q, want %q", i, commands[i].Cmd, wecomCmdRespondMsg)
|
||||||
|
}
|
||||||
|
body, ok := commands[i].Body.(wecomRespondMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("command[%d] body type = %T", i, commands[i].Body)
|
||||||
|
}
|
||||||
|
if body.Stream == nil {
|
||||||
|
t.Fatalf("command[%d] missing stream body", i)
|
||||||
|
}
|
||||||
|
if body.Stream.ID != "stream-1" {
|
||||||
|
t.Fatalf("command[%d] stream id = %q, want stream-1", i, body.Stream.ID)
|
||||||
|
}
|
||||||
|
if body.Stream.Finish != wantFinish {
|
||||||
|
t.Fatalf("command[%d] finish = %v, want %v", i, body.Stream.Finish, wantFinish)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if body := commands[0].Body.(wecomRespondMsgBody); body.Stream.Content != "draft" {
|
||||||
|
t.Fatalf("update content = %q, want draft", body.Stream.Content)
|
||||||
|
}
|
||||||
|
if body := commands[1].Body.(wecomRespondMsgBody); body.Stream.Content != "final" {
|
||||||
|
t.Fatalf("final content = %q, want final", body.Stream.Content)
|
||||||
|
}
|
||||||
|
if _, ok := ch.getTurn("chat-1"); ok {
|
||||||
|
t.Fatal("expected turn to be consumed after Finalize")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSend_StreamFailureFallsBackToActualChatID(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
ch.queueTurn("chat-1", wecomTurn{
|
||||||
|
ReqID: "req-1",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
ChatType: 1,
|
||||||
|
StreamID: "stream-1",
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
})
|
||||||
|
ch.queueTurn("chat-1", wecomTurn{
|
||||||
|
ReqID: "req-2",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
ChatType: 1,
|
||||||
|
StreamID: "stream-2",
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
})
|
||||||
|
if err := ch.routes.Put("chat-1", "req-2", 1, time.Hour); err != nil {
|
||||||
|
t.Fatalf("Put() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
if len(commands) == 1 && cmd.Cmd == wecomCmdRespondMsg {
|
||||||
|
return wecomEnvelope{}, errors.New("stream send failed")
|
||||||
|
}
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
|
Channel: "wecom",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Content: "hello",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("Send() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 2 {
|
||||||
|
t.Fatalf("expected 2 commands, got %d", len(commands))
|
||||||
|
}
|
||||||
|
if commands[0].Cmd != wecomCmdRespondMsg || commands[0].Headers.ReqID != "req-1" {
|
||||||
|
t.Fatalf("first command = %+v", commands[0])
|
||||||
|
}
|
||||||
|
if commands[1].Cmd != wecomCmdSendMsg {
|
||||||
|
t.Fatalf("second command = %q, want %q", commands[1].Cmd, wecomCmdSendMsg)
|
||||||
|
}
|
||||||
|
body, ok := commands[1].Body.(wecomSendMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected send body type %T", commands[1].Body)
|
||||||
|
}
|
||||||
|
if body.ChatID != "chat-1" {
|
||||||
|
t.Fatalf("send chatid = %q, want chat-1", body.ChatID)
|
||||||
|
}
|
||||||
|
if body.ChatType != 1 {
|
||||||
|
t.Fatalf("send chat_type = %d, want 1", body.ChatType)
|
||||||
|
}
|
||||||
|
|
||||||
|
nextTurn, ok := ch.getTurn("chat-1")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected second turn to remain queued")
|
||||||
|
}
|
||||||
|
if nextTurn.ReqID != "req-2" {
|
||||||
|
t.Fatalf("next queued req_id = %q, want req-2", nextTurn.ReqID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSend_DoesNotSplitStreamReply(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
ch.queueTurn("chat-1", wecomTurn{
|
||||||
|
ReqID: "req-1",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
ChatType: 1,
|
||||||
|
StreamID: "stream-1",
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
})
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
content := strings.Repeat("\u4e2d", 30000)
|
||||||
|
if err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
|
Channel: "wecom",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Content: content,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("Send() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 1 {
|
||||||
|
t.Fatalf("expected 1 stream command, got %d", len(commands))
|
||||||
|
}
|
||||||
|
body, ok := commands[0].Body.(wecomRespondMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected body type %T", commands[0].Body)
|
||||||
|
}
|
||||||
|
if body.Stream == nil || !body.Stream.Finish {
|
||||||
|
t.Fatalf("stream body = %+v", body.Stream)
|
||||||
|
}
|
||||||
|
if body.Stream.Content != content {
|
||||||
|
t.Fatalf("stream content length = %d, want %d", len(body.Stream.Content), len(content))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSend_DoesNotSplitActivePush(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
content := strings.Repeat("a", 30000)
|
||||||
|
if err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
|
Channel: "wecom",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Content: content,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("Send() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 1 {
|
||||||
|
t.Fatalf("expected 1 send command, got %d", len(commands))
|
||||||
|
}
|
||||||
|
if commands[0].Cmd != wecomCmdSendMsg {
|
||||||
|
t.Fatalf("command = %q, want %q", commands[0].Cmd, wecomCmdSendMsg)
|
||||||
|
}
|
||||||
|
body, ok := commands[0].Body.(wecomSendMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected body type %T", commands[0].Body)
|
||||||
|
}
|
||||||
|
if body.Markdown == nil || body.Markdown.Content != content {
|
||||||
|
t.Fatalf("markdown content length = %d, want %d", len(body.Markdown.Content), len(content))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_SendsActiveImage(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
ch.SetMediaStore(store)
|
||||||
|
|
||||||
|
imageData := wecomTestJPEGData(t)
|
||||||
|
imagePath := filepath.Join(t.TempDir(), "photo.jpg")
|
||||||
|
if err := os.WriteFile(imagePath, imageData, 0o600); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
ref, err := store.Store(imagePath, media.MediaMeta{
|
||||||
|
Filename: "photo.jpg",
|
||||||
|
ContentType: "image/jpeg",
|
||||||
|
Source: "test",
|
||||||
|
CleanupPolicy: media.CleanupPolicyForgetOnly,
|
||||||
|
}, "scope-1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Store() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
switch cmd.Cmd {
|
||||||
|
case wecomCmdUploadMediaInit:
|
||||||
|
return wecomTestAck(wecomUploadMediaInitResponse{UploadID: "upload-1"}), nil
|
||||||
|
case wecomCmdUploadMediaEnd:
|
||||||
|
return wecomTestAck(wecomUploadMediaFinishResponse{
|
||||||
|
Type: "image",
|
||||||
|
MediaID: "media-1",
|
||||||
|
}), nil
|
||||||
|
default:
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "wecom",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Parts: []bus.MediaPart{{
|
||||||
|
Ref: ref,
|
||||||
|
Type: "image",
|
||||||
|
Filename: "photo.jpg",
|
||||||
|
ContentType: "image/jpeg",
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendMedia() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 4 {
|
||||||
|
t.Fatalf("expected 4 commands, got %d", len(commands))
|
||||||
|
}
|
||||||
|
if commands[0].Cmd != wecomCmdUploadMediaInit {
|
||||||
|
t.Fatalf("first command = %q, want %q", commands[0].Cmd, wecomCmdUploadMediaInit)
|
||||||
|
}
|
||||||
|
initBody, ok := commands[0].Body.(wecomUploadMediaInitBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected init body type %T", commands[0].Body)
|
||||||
|
}
|
||||||
|
if initBody.Type != "image" || initBody.Filename != "photo.jpg" || initBody.TotalChunks != 1 {
|
||||||
|
t.Fatalf("init body = %+v", initBody)
|
||||||
|
}
|
||||||
|
if commands[1].Cmd != wecomCmdUploadMediaChunk {
|
||||||
|
t.Fatalf("second command = %q, want %q", commands[1].Cmd, wecomCmdUploadMediaChunk)
|
||||||
|
}
|
||||||
|
chunkBody, ok := commands[1].Body.(wecomUploadMediaChunkBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected chunk body type %T", commands[1].Body)
|
||||||
|
}
|
||||||
|
if chunkBody.UploadID != "upload-1" || chunkBody.ChunkIndex != 0 || chunkBody.Base64Data == "" {
|
||||||
|
t.Fatalf("chunk body = %+v", chunkBody)
|
||||||
|
}
|
||||||
|
if commands[2].Cmd != wecomCmdUploadMediaEnd {
|
||||||
|
t.Fatalf("third command = %q, want %q", commands[2].Cmd, wecomCmdUploadMediaEnd)
|
||||||
|
}
|
||||||
|
if commands[3].Cmd != wecomCmdSendMsg {
|
||||||
|
t.Fatalf("fourth command = %q, want %q", commands[3].Cmd, wecomCmdSendMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
body, ok := commands[3].Body.(wecomSendMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected send body type %T", commands[3].Body)
|
||||||
|
}
|
||||||
|
if body.MsgType != "image" || body.Image == nil {
|
||||||
|
t.Fatalf("send body = %+v", body)
|
||||||
|
}
|
||||||
|
if body.ChatID != "chat-1" {
|
||||||
|
t.Fatalf("send chatid = %q, want chat-1", body.ChatID)
|
||||||
|
}
|
||||||
|
if body.Image.MediaID != "media-1" {
|
||||||
|
t.Fatalf("image media_id = %q, want media-1", body.Image.MediaID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_UsesTurnImageAndFinishesStream(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
ch.SetMediaStore(store)
|
||||||
|
|
||||||
|
imageData := wecomTestJPEGData(t)
|
||||||
|
imagePath := filepath.Join(t.TempDir(), "reply.jpg")
|
||||||
|
if err := os.WriteFile(imagePath, imageData, 0o600); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
ref, err := store.Store(imagePath, media.MediaMeta{
|
||||||
|
Filename: "reply.jpg",
|
||||||
|
ContentType: "image/jpeg",
|
||||||
|
Source: "test",
|
||||||
|
CleanupPolicy: media.CleanupPolicyForgetOnly,
|
||||||
|
}, "scope-2")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Store() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ch.queueTurn("chat-1", wecomTurn{
|
||||||
|
ReqID: "req-1",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
ChatType: 1,
|
||||||
|
StreamID: "stream-1",
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
})
|
||||||
|
putErr := ch.routes.Put("chat-1", "req-1", 1, time.Hour)
|
||||||
|
if putErr != nil {
|
||||||
|
t.Fatalf("Put() error = %v", putErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
switch cmd.Cmd {
|
||||||
|
case wecomCmdUploadMediaInit:
|
||||||
|
return wecomTestAck(wecomUploadMediaInitResponse{UploadID: "upload-2"}), nil
|
||||||
|
case wecomCmdUploadMediaEnd:
|
||||||
|
return wecomTestAck(wecomUploadMediaFinishResponse{
|
||||||
|
Type: "image",
|
||||||
|
MediaID: "media-2",
|
||||||
|
}), nil
|
||||||
|
default:
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "wecom",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Parts: []bus.MediaPart{{
|
||||||
|
Ref: ref,
|
||||||
|
Type: "image",
|
||||||
|
Filename: "reply.jpg",
|
||||||
|
ContentType: "image/jpeg",
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendMedia() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 5 {
|
||||||
|
t.Fatalf("expected 5 commands, got %d", len(commands))
|
||||||
|
}
|
||||||
|
if commands[0].Cmd != wecomCmdUploadMediaInit {
|
||||||
|
t.Fatalf("first command = %+v", commands[0])
|
||||||
|
}
|
||||||
|
if commands[1].Cmd != wecomCmdUploadMediaChunk {
|
||||||
|
t.Fatalf("second command = %+v", commands[1])
|
||||||
|
}
|
||||||
|
if commands[2].Cmd != wecomCmdUploadMediaEnd {
|
||||||
|
t.Fatalf("third command = %+v", commands[2])
|
||||||
|
}
|
||||||
|
if commands[3].Cmd != wecomCmdRespondMsg || commands[3].Headers.ReqID != "req-1" {
|
||||||
|
t.Fatalf("fourth command = %+v", commands[3])
|
||||||
|
}
|
||||||
|
if commands[4].Cmd != wecomCmdRespondMsg || commands[4].Headers.ReqID != "req-1" {
|
||||||
|
t.Fatalf("fifth command = %+v", commands[4])
|
||||||
|
}
|
||||||
|
|
||||||
|
imageBody, ok := commands[3].Body.(wecomRespondMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected image body type %T", commands[3].Body)
|
||||||
|
}
|
||||||
|
if imageBody.MsgType != "image" || imageBody.Image == nil {
|
||||||
|
t.Fatalf("image body = %+v", imageBody)
|
||||||
|
}
|
||||||
|
if imageBody.Image.MediaID != "media-2" {
|
||||||
|
t.Fatalf("image media_id = %q, want media-2", imageBody.Image.MediaID)
|
||||||
|
}
|
||||||
|
|
||||||
|
streamBody, ok := commands[4].Body.(wecomRespondMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected finish body type %T", commands[4].Body)
|
||||||
|
}
|
||||||
|
if streamBody.MsgType != "stream" || streamBody.Stream == nil || !streamBody.Stream.Finish {
|
||||||
|
t.Fatalf("finish body = %+v", streamBody)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := ch.getTurn("chat-1"); ok {
|
||||||
|
t.Fatal("expected turn to be removed after media send")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendMedia_SendsActiveFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ch := newTestWeComChannel(t, bus.NewMessageBus())
|
||||||
|
ch.SetRunning(true)
|
||||||
|
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
ch.SetMediaStore(store)
|
||||||
|
|
||||||
|
filePath := filepath.Join(t.TempDir(), "report.pdf")
|
||||||
|
if err := os.WriteFile(filePath, []byte("%PDF-1.4"), 0o600); err != nil {
|
||||||
|
t.Fatalf("WriteFile() error = %v", err)
|
||||||
|
}
|
||||||
|
ref, err := store.Store(filePath, media.MediaMeta{
|
||||||
|
Filename: "report.pdf",
|
||||||
|
ContentType: "application/pdf",
|
||||||
|
Source: "test",
|
||||||
|
CleanupPolicy: media.CleanupPolicyForgetOnly,
|
||||||
|
}, "scope-3")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Store() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var commands []wecomCommand
|
||||||
|
ch.commandSend = func(cmd wecomCommand, _ time.Duration) (wecomEnvelope, error) {
|
||||||
|
commands = append(commands, cmd)
|
||||||
|
switch cmd.Cmd {
|
||||||
|
case wecomCmdUploadMediaInit:
|
||||||
|
return wecomTestAck(wecomUploadMediaInitResponse{UploadID: "upload-3"}), nil
|
||||||
|
case wecomCmdUploadMediaEnd:
|
||||||
|
return wecomTestAck(wecomUploadMediaFinishResponse{
|
||||||
|
Type: "file",
|
||||||
|
MediaID: "media-3",
|
||||||
|
}), nil
|
||||||
|
default:
|
||||||
|
return wecomTestAck(nil), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
|
Channel: "wecom",
|
||||||
|
ChatID: "chat-2",
|
||||||
|
Parts: []bus.MediaPart{{
|
||||||
|
Ref: ref,
|
||||||
|
Type: "file",
|
||||||
|
Filename: "report.pdf",
|
||||||
|
ContentType: "application/pdf",
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendMedia() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commands) != 4 {
|
||||||
|
t.Fatalf("expected 4 commands, got %d", len(commands))
|
||||||
|
}
|
||||||
|
if commands[0].Cmd != wecomCmdUploadMediaInit {
|
||||||
|
t.Fatalf("first command = %q, want %q", commands[0].Cmd, wecomCmdUploadMediaInit)
|
||||||
|
}
|
||||||
|
initBody, ok := commands[0].Body.(wecomUploadMediaInitBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected init body type %T", commands[0].Body)
|
||||||
|
}
|
||||||
|
if initBody.Type != "file" || initBody.Filename != "report.pdf" {
|
||||||
|
t.Fatalf("init body = %+v", initBody)
|
||||||
|
}
|
||||||
|
if commands[1].Cmd != wecomCmdUploadMediaChunk {
|
||||||
|
t.Fatalf("second command = %q, want %q", commands[1].Cmd, wecomCmdUploadMediaChunk)
|
||||||
|
}
|
||||||
|
if commands[2].Cmd != wecomCmdUploadMediaEnd {
|
||||||
|
t.Fatalf("third command = %q, want %q", commands[2].Cmd, wecomCmdUploadMediaEnd)
|
||||||
|
}
|
||||||
|
if commands[3].Cmd != wecomCmdSendMsg {
|
||||||
|
t.Fatalf("fourth command = %q, want %q", commands[3].Cmd, wecomCmdSendMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
body, ok := commands[3].Body.(wecomSendMsgBody)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("unexpected body type %T", commands[3].Body)
|
||||||
|
}
|
||||||
|
if body.MsgType != "file" || body.File == nil {
|
||||||
|
t.Fatalf("body = %+v", body)
|
||||||
|
}
|
||||||
|
if body.File.MediaID != "media-3" {
|
||||||
|
t.Fatalf("file media_id = %q, want media-3", body.File.MediaID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestWeComChannel(t *testing.T, messageBus *bus.MessageBus) *WeComChannel {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
cfg := config.WeComConfig{BotID: "bot-1"}
|
||||||
|
cfg.SetSecret("secret-1")
|
||||||
|
ch, err := NewChannel(cfg, messageBus)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewChannel() error = %v", err)
|
||||||
|
}
|
||||||
|
ch.ctx = context.Background()
|
||||||
|
ch.routes = newReqIDStore(filepath.Join(t.TempDir(), "reqids.json"))
|
||||||
|
return ch
|
||||||
|
}
|
||||||
|
|
||||||
|
func wecomTestJPEGData(t *testing.T) []byte {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
const jpegBase64 = "/9j/4AAQSkZJRgABAQAAAQABAAD/2wBDAP//////////////////////////////////////////////////////////////////////////////////////" +
|
||||||
|
"//////////////////////////////////////////////////////////////////////////////////////////////2wBDAf//////////////////////////////////////////////////////////////////////////////////////" +
|
||||||
|
"//////////////////////////////////////////////////////////////////////////////////////////////wAARCAABAAEDASIAAhEBAxEB/8QAFQABAQAAAAAAAAAAAAAAAAAAAAb/xAAVEQEBAAAAAAAAAAAAAAAAAAAABf/aAAwDAQACEAMQAAAB6A//xAAVEAEBAAAAAAAAAAAAAAAAAAAAEf/aAAgBAQABBQJf/8QAFBEBAAAAAAAAAAAAAAAAAAAAEP/aAAgBAwEBPwF//8QAFBEBAAAAAAAAAAAAAAAAAAAAEP/aAAgBAgEBPwF//8QAFBABAAAAAAAAAAAAAAAAAAAAEP/aAAgBAQAGPwJf/8QAFBABAAAAAAAAAAAAAAAAAAAAEP/aAAgBAQABPyFf/9k="
|
||||||
|
|
||||||
|
return decodeTestBase64(t, jpegBase64)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeWeComUploadFinish_AcceptsNumericCreatedAt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
resp, err := decodeWeComEnvelopeBody[wecomUploadMediaFinishResponse](wecomEnvelope{
|
||||||
|
Body: json.RawMessage(`{"type":"file","media_id":"media-1","created_at":1380000000}`),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("decodeWeComEnvelopeBody() error = %v", err)
|
||||||
|
}
|
||||||
|
if resp.Type != "file" || resp.MediaID != "media-1" {
|
||||||
|
t.Fatalf("response = %+v", resp)
|
||||||
|
}
|
||||||
|
if string(resp.CreatedAt) != "1380000000" {
|
||||||
|
t.Fatalf("created_at = %s, want 1380000000", string(resp.CreatedAt))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func wecomTestAck(body any) wecomEnvelope {
|
||||||
|
var raw []byte
|
||||||
|
if body != nil {
|
||||||
|
encoded, err := json.Marshal(body)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
raw = encoded
|
||||||
|
}
|
||||||
|
return wecomEnvelope{
|
||||||
|
ErrCode: 0,
|
||||||
|
ErrMsg: "ok",
|
||||||
|
Body: raw,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,606 +0,0 @@
|
||||||
# Security Configuration Refactoring
|
|
||||||
|
|
||||||
## Overview
|
|
||||||
|
|
||||||
This refactoring introduces a `.security.yml` file to store all sensitive data (API keys, tokens, secrets, passwords) separately from the main configuration. This improves security by:
|
|
||||||
|
|
||||||
1. **Separation of concerns**: Configuration settings and secrets are in separate files
|
|
||||||
2. **Easier sharing**: The main config can be shared without exposing sensitive data
|
|
||||||
3. **Better version control**: `.security.yml` can be added to `.gitignore`
|
|
||||||
4. **Flexible deployment**: Different environments can use different security files
|
|
||||||
|
|
||||||
## File Structure
|
|
||||||
|
|
||||||
```
|
|
||||||
~/.picoclaw/
|
|
||||||
├── config.json # Main configuration (safe to share)
|
|
||||||
└── .security.yml # Security data (never share)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Usage
|
|
||||||
|
|
||||||
### Basic Configuration
|
|
||||||
|
|
||||||
In your `config.json`, use `ref:` references to point to values in `.security.yml`:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"version": 1,
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.4",
|
|
||||||
"model": "openai/gpt-5.4",
|
|
||||||
"api_base": "https://api.openai.com/v1",
|
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"channels": {
|
|
||||||
"telegram": {
|
|
||||||
"enabled": true,
|
|
||||||
"token": "ref:channels.telegram.token"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Security Configuration
|
|
||||||
|
|
||||||
In your `.security.yml`, store the actual values:
|
|
||||||
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-your-actual-api-key-1"
|
|
||||||
- "sk-your-actual-api-key-2" # Optional: Multiple keys for failover
|
|
||||||
claude-sonnet-4.6:
|
|
||||||
api_keys:
|
|
||||||
- "sk-your-actual-anthropic-key" # Single key in array format
|
|
||||||
|
|
||||||
channels:
|
|
||||||
telegram:
|
|
||||||
token: "your-telegram-bot-token"
|
|
||||||
|
|
||||||
web:
|
|
||||||
brave:
|
|
||||||
api_keys:
|
|
||||||
- "BSAyour-brave-api-key-1"
|
|
||||||
- "BSAyour-brave-api-key-2" # Optional: Multiple keys for failover
|
|
||||||
tavily:
|
|
||||||
api_keys:
|
|
||||||
- "tvly-your-tavily-api-key" # Single key in array format
|
|
||||||
glm_search:
|
|
||||||
api_key: "your-glm-search-api-key" # GLMSearch uses single key format
|
|
||||||
```
|
|
||||||
|
|
||||||
## Reference Format
|
|
||||||
|
|
||||||
### Model API Keys
|
|
||||||
|
|
||||||
Format: `ref:model_list.<model_name>.api_key`
|
|
||||||
|
|
||||||
Example: `ref:model_list.gpt-5.4.api_key`
|
|
||||||
|
|
||||||
### Channel Tokens/Secrets
|
|
||||||
|
|
||||||
Format: `ref:channels.<channel_name>.<field>`
|
|
||||||
|
|
||||||
Examples:
|
|
||||||
- `ref:channels.telegram.token`
|
|
||||||
- `ref:channels.feishu.app_secret`
|
|
||||||
- `ref:channels.feishu.encrypt_key`
|
|
||||||
- `ref:channels.feishu.verification_token`
|
|
||||||
- `ref:channels.discord.token`
|
|
||||||
- `ref:channels.qq.app_secret`
|
|
||||||
- `ref:channels.dingtalk.client_secret`
|
|
||||||
- `ref:channels.slack.bot_token`
|
|
||||||
- `ref:channels.slack.app_token`
|
|
||||||
- `ref:channels.matrix.access_token`
|
|
||||||
- `ref:channels.line.channel_secret`
|
|
||||||
- `ref:channels.line.channel_access_token`
|
|
||||||
- `ref:channels.onebot.access_token`
|
|
||||||
- `ref:channels.wecom.token`
|
|
||||||
- `ref:channels.wecom.encoding_aes_key`
|
|
||||||
- `ref:channels.wecom_app.corp_secret`
|
|
||||||
- `ref:channels.wecom_app.token`
|
|
||||||
- `ref:channels.wecom_app.encoding_aes_key`
|
|
||||||
- `ref:channels.wecom_aibot.token`
|
|
||||||
- `ref:channels.wecom_aibot.encoding_aes_key`
|
|
||||||
- `ref:channels.pico.token`
|
|
||||||
- `ref:channels.irc.password`
|
|
||||||
- `ref:channels.irc.nickserv_password`
|
|
||||||
- `ref:channels.irc.sasl_password`
|
|
||||||
|
|
||||||
### Web Tool API Keys
|
|
||||||
|
|
||||||
Format: `ref:web.<provider>.<field>`
|
|
||||||
|
|
||||||
Examples:
|
|
||||||
- `ref:web.brave.api_key`
|
|
||||||
- `ref:web.tavily.api_key`
|
|
||||||
- `ref:web.perplexity.api_key`
|
|
||||||
- `ref:web.glm_search.api_key`
|
|
||||||
|
|
||||||
### Skills Registry Tokens
|
|
||||||
|
|
||||||
Format: `ref:skills.<registry>.<field>`
|
|
||||||
|
|
||||||
Examples:
|
|
||||||
- `ref:skills.github.token`
|
|
||||||
- `ref:skills.clawhub.auth_token`
|
|
||||||
|
|
||||||
## Backward Compatibility
|
|
||||||
|
|
||||||
The refactoring maintains full backward compatibility:
|
|
||||||
|
|
||||||
1. **Direct values**: You can still use direct values in `config.json` (not recommended for production)
|
|
||||||
2. **Mixed usage**: You can mix `ref:` references and direct values
|
|
||||||
3. **Optional security file**: If `.security.yml` doesn't exist, all references will fail (but direct values still work)
|
|
||||||
|
|
||||||
## Configuration Precedence
|
|
||||||
|
|
||||||
When both `config.json` and `.security.yml` contain security configurations, PicoClaw uses the following precedence rules:
|
|
||||||
|
|
||||||
### Priority Order (Highest to Lowest)
|
|
||||||
|
|
||||||
1. **Security settings in `config.json`** (highest priority)
|
|
||||||
- Direct values in `config.json`
|
|
||||||
- These settings override any conflicting values in `.security.yml`
|
|
||||||
|
|
||||||
2. **Security settings in `.security.yml`**
|
|
||||||
- Used when no conflicting setting exists in `config.json`
|
|
||||||
- Provides default/fallback security values
|
|
||||||
|
|
||||||
### Practical Example
|
|
||||||
|
|
||||||
**Scenario**: You have API keys defined in both files.
|
|
||||||
|
|
||||||
**.security.yml:**
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-4o:
|
|
||||||
api_keys:
|
|
||||||
- "sk-default-key-from-security-yml"
|
|
||||||
```
|
|
||||||
|
|
||||||
**config.json:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-4o",
|
|
||||||
"api_key": "sk-custom-key-from-config-json"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**Result**: The API key `"sk-custom-key-from-config-json"` from `config.json` takes precedence.
|
|
||||||
|
|
||||||
### Use Cases
|
|
||||||
|
|
||||||
This precedence system enables several useful patterns:
|
|
||||||
|
|
||||||
1. **Environment-specific overrides**: Keep default keys in `.security.yml`, override per-environment keys in `config.json`
|
|
||||||
2. **Temporary testing**: Quickly test a new API key in `config.json` without modifying `.security.yml`
|
|
||||||
3. **Team sharing**: Share common keys via `.security.yml` while allowing individual developers to override in their local `config.json`
|
|
||||||
|
|
||||||
### Migration Behavior
|
|
||||||
|
|
||||||
When migrating from config v0 to v1:
|
|
||||||
- Security values extracted from legacy `config.json` take precedence
|
|
||||||
- Existing `.security.yml` values serve as fallback
|
|
||||||
- No data loss: all values are preserved and merged appropriately
|
|
||||||
|
|
||||||
### API Key Formats in .security.yml
|
|
||||||
|
|
||||||
**Models (gpt-5.4, claude-sonnet-4.6, etc.):**
|
|
||||||
- Must use `api_keys` (array) format
|
|
||||||
- Both single and multiple keys use array format
|
|
||||||
|
|
||||||
**Web Tools (Brave, Tavily, Perplexity):**
|
|
||||||
- Must use `api_keys` (array) format
|
|
||||||
- Both single and multiple keys use array format
|
|
||||||
|
|
||||||
**Web Tools (GLMSearch):**
|
|
||||||
- Must use `api_key` (single string) format
|
|
||||||
- Does NOT support array format
|
|
||||||
|
|
||||||
**Channels (Telegram, Discord, etc.):**
|
|
||||||
- Use single field names (e.g., `token`, `app_secret`)
|
|
||||||
- Each channel uses its specific field names
|
|
||||||
|
|
||||||
### Single Key (Models)
|
|
||||||
|
|
||||||
Use array format with one element:
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-your-key"
|
|
||||||
```
|
|
||||||
|
|
||||||
In `config.json`:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Single Key (GLMSearch)
|
|
||||||
|
|
||||||
Use single string format:
|
|
||||||
```yaml
|
|
||||||
web:
|
|
||||||
glm_search:
|
|
||||||
api_key: "your-glm-key"
|
|
||||||
```
|
|
||||||
|
|
||||||
In `config.json`:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"api_key": "ref:web.glm_search.api_key"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
## Migration Guide
|
|
||||||
|
|
||||||
### Step 1: Create .security.yml
|
|
||||||
|
|
||||||
Copy the example template:
|
|
||||||
```bash
|
|
||||||
cp security.example.yml ~/.picoclaw/.security.yml
|
|
||||||
```
|
|
||||||
|
|
||||||
### Step 2: Fill in your actual values
|
|
||||||
|
|
||||||
Edit `~/.picoclaw/.security.yml` and replace placeholder values with your actual API keys and tokens.
|
|
||||||
|
|
||||||
### Step 3: Update config.json
|
|
||||||
|
|
||||||
Replace sensitive values in `~/.picoclaw/config.json` with `ref:` references:
|
|
||||||
|
|
||||||
**Before:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.4",
|
|
||||||
"model": "openai/gpt-5.4",
|
|
||||||
"api_key": "sk-your-actual-api-key-here"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**After:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.4",
|
|
||||||
"model": "openai/gpt-5.4",
|
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Step 4: Verify
|
|
||||||
|
|
||||||
Restart PicoClaw and verify it loads correctly:
|
|
||||||
```bash
|
|
||||||
picoclaw --version
|
|
||||||
```
|
|
||||||
|
|
||||||
## Security Best Practices
|
|
||||||
|
|
||||||
1. **Never commit `.security.yml`** to version control
|
|
||||||
2. **Set file permissions**: `chmod 600 ~/.picoclaw/.security.yml`
|
|
||||||
3. **Use different keys** for different environments (dev, staging, production)
|
|
||||||
4. **Rotate keys regularly** and update `.security.yml`
|
|
||||||
5. **Backup securely**: Encrypt backups containing `.security.yml`
|
|
||||||
|
|
||||||
## API
|
|
||||||
|
|
||||||
### LoadSecurityConfig
|
|
||||||
|
|
||||||
```go
|
|
||||||
func LoadSecurityConfig(securityPath string) (*SecurityConfig, error)
|
|
||||||
```
|
|
||||||
|
|
||||||
Loads the security configuration from `.security.yml`. Returns an empty `SecurityConfig` if the file doesn't exist.
|
|
||||||
|
|
||||||
### SaveSecurityConfig
|
|
||||||
|
|
||||||
```go
|
|
||||||
func SaveSecurityConfig(securityPath string, sec *SecurityConfig) error
|
|
||||||
```
|
|
||||||
|
|
||||||
Saves the security configuration to `.security.yml` with `0o600` permissions.
|
|
||||||
|
|
||||||
### ResolveReference
|
|
||||||
|
|
||||||
```go
|
|
||||||
func (sec *SecurityConfig) ResolveReference(ref string) (string, error)
|
|
||||||
```
|
|
||||||
|
|
||||||
Resolves a reference string (e.g., `"ref:model_list.test.api_key"`) and returns the actual value.
|
|
||||||
|
|
||||||
### SecurityPath
|
|
||||||
|
|
||||||
```go
|
|
||||||
func SecurityPath(configPath string) string
|
|
||||||
```
|
|
||||||
|
|
||||||
Returns the path to `.security.yml` relative to the config file.
|
|
||||||
|
|
||||||
## Example: Complete Configuration
|
|
||||||
|
|
||||||
### config.json
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"version": 1,
|
|
||||||
"agents": {
|
|
||||||
"defaults": {
|
|
||||||
"workspace": "~/picoclaw-workspace",
|
|
||||||
"model_name": "gpt-5.4"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.4",
|
|
||||||
"model": "openai/gpt-5.4",
|
|
||||||
"api_base": "https://api.openai.com/v1",
|
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"model_name": "claude-sonnet-4.6",
|
|
||||||
"model": "anthropic/claude-sonnet-4.6",
|
|
||||||
"api_base": "https://api.anthropic.com/v1",
|
|
||||||
"api_key": "ref:model_list.claude-sonnet-4.6.api_key"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"channels": {
|
|
||||||
"telegram": {
|
|
||||||
"enabled": true,
|
|
||||||
"token": "ref:channels.telegram.token"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"tools": {
|
|
||||||
"web": {
|
|
||||||
"brave": {
|
|
||||||
"enabled": true,
|
|
||||||
"api_key": "ref:web.brave.api_key"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### .security.yml
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-proj-actual-openai-key-1"
|
|
||||||
- "sk-proj-actual-openai-key-2"
|
|
||||||
claude-sonnet-4.6:
|
|
||||||
api_keys:
|
|
||||||
- "sk-ant-actual-anthropic-key" # Single key in array format
|
|
||||||
|
|
||||||
channels:
|
|
||||||
telegram:
|
|
||||||
token: "1234567890:ABCdefGHIjklMNOpqrsTUVwxyz"
|
|
||||||
|
|
||||||
web:
|
|
||||||
brave:
|
|
||||||
api_keys:
|
|
||||||
- "BSAactualbravekey-1"
|
|
||||||
- "BSAactualbravekey-2"
|
|
||||||
tavily:
|
|
||||||
api_keys:
|
|
||||||
- "tvly-your-tavily-key" # Single key in array format
|
|
||||||
glm_search:
|
|
||||||
api_key: "your-glm-key" # GLMSearch uses single key format
|
|
||||||
```
|
|
||||||
|
|
||||||
## Testing
|
|
||||||
|
|
||||||
The refactoring includes comprehensive tests:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
go test ./pkg/config -run TestSecurityConfig
|
|
||||||
```
|
|
||||||
|
|
||||||
## Troubleshooting
|
|
||||||
|
|
||||||
### Error: "model security entry not found"
|
|
||||||
|
|
||||||
- Ensure the model name in your reference matches exactly in `.security.yml`
|
|
||||||
- Check that the `model_list` section exists in `.security.yml`
|
|
||||||
- For models with indexed names (e.g., "gpt-5.4:0"), ensure the exact name is used or check the base name without index
|
|
||||||
|
|
||||||
### Error: "failed to load security config"
|
|
||||||
|
|
||||||
- Verify `.security.yml` exists in the same directory as `config.json`
|
|
||||||
- Check the YAML syntax is valid (use a YAML validator)
|
|
||||||
- Ensure file permissions allow reading
|
|
||||||
|
|
||||||
### Error: "unknown reference path"
|
|
||||||
|
|
||||||
- Verify the reference format is correct
|
|
||||||
- Check the path structure matches the examples above
|
|
||||||
- Ensure all required sections exist in `.security.yml`
|
|
||||||
|
|
||||||
## Advanced Features
|
|
||||||
|
|
||||||
### Multiple API Keys (Load Balancing & Failover)
|
|
||||||
|
|
||||||
Both models and web tools support multiple API keys for improved reliability:
|
|
||||||
|
|
||||||
**Benefits:**
|
|
||||||
- **Load balancing**: Requests are distributed across multiple keys
|
|
||||||
- **Failover**: Automatic switching to another key if one fails
|
|
||||||
- **Rate limit management**: Distribute usage across multiple keys
|
|
||||||
- **High availability**: Reduce downtime during API provider issues
|
|
||||||
|
|
||||||
#### Example: Model with Multiple Keys
|
|
||||||
|
|
||||||
**.security.yml:**
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-proj-key-1"
|
|
||||||
- "sk-proj-key-2"
|
|
||||||
- "sk-proj-key-3"
|
|
||||||
```
|
|
||||||
|
|
||||||
**config.json:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.4",
|
|
||||||
"model": "openai/gpt-5.4",
|
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### Example: Web Tool with Multiple Keys
|
|
||||||
|
|
||||||
**.security.yml:**
|
|
||||||
```yaml
|
|
||||||
web:
|
|
||||||
brave:
|
|
||||||
api_keys:
|
|
||||||
- "BSA-key-1"
|
|
||||||
- "BSA-key-2"
|
|
||||||
tavily:
|
|
||||||
api_keys:
|
|
||||||
- "tvly-your-key" # Single key in array format
|
|
||||||
glm_search:
|
|
||||||
api_key: "your-glm-key" # GLMSearch uses single key format
|
|
||||||
```
|
|
||||||
|
|
||||||
**config.json:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"tools": {
|
|
||||||
"web": {
|
|
||||||
"brave": {
|
|
||||||
"enabled": true,
|
|
||||||
"api_key": "ref:web.brave.api_key"
|
|
||||||
},
|
|
||||||
"tavily": {
|
|
||||||
"enabled": true,
|
|
||||||
"api_key": "ref:web.tavily.api_key"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### Supported Formats
|
|
||||||
|
|
||||||
**Models - Single key:**
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-your-key" # Array with one element
|
|
||||||
```
|
|
||||||
|
|
||||||
**Models - Multiple keys:**
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-your-key-1"
|
|
||||||
- "sk-your-key-2"
|
|
||||||
- "sk-your-key-3"
|
|
||||||
```
|
|
||||||
|
|
||||||
**Web Tools (Brave/Tavily/Perplexity) - Single key:**
|
|
||||||
```yaml
|
|
||||||
web:
|
|
||||||
brave:
|
|
||||||
api_keys:
|
|
||||||
- "BSA-your-key" # Array with one element
|
|
||||||
```
|
|
||||||
|
|
||||||
**Web Tools (Brave/Tavily/Perplexity) - Multiple keys:**
|
|
||||||
```yaml
|
|
||||||
web:
|
|
||||||
brave:
|
|
||||||
api_keys:
|
|
||||||
- "BSA-key-1"
|
|
||||||
- "BSA-key-2"
|
|
||||||
```
|
|
||||||
|
|
||||||
**Web Tool (GLMSearch) - Single key only:**
|
|
||||||
```yaml
|
|
||||||
web:
|
|
||||||
glm_search:
|
|
||||||
api_key: "your-glm-key" # Single string (NOT array)
|
|
||||||
```
|
|
||||||
|
|
||||||
All formats work identically in `config.json` - you always use the same reference format:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Model Indexing for Load Balancing
|
|
||||||
|
|
||||||
When you have multiple models with the same base name but different API keys, you can use indexed names:
|
|
||||||
|
|
||||||
**.security.yml:**
|
|
||||||
```yaml
|
|
||||||
model_list:
|
|
||||||
gpt-5.4:
|
|
||||||
api_keys:
|
|
||||||
- "sk-proj-key-1"
|
|
||||||
- "sk-proj-key-2"
|
|
||||||
```
|
|
||||||
|
|
||||||
The system will automatically expand this into multiple model entries with fallback support.
|
|
||||||
|
|
||||||
### Environment Variables
|
|
||||||
|
|
||||||
You can override any security value using environment variables:
|
|
||||||
|
|
||||||
**For models:**
|
|
||||||
```bash
|
|
||||||
export PICOCLAW_MODEL_LIST_GPT-5.4_API_KEY="sk-from-env"
|
|
||||||
```
|
|
||||||
|
|
||||||
**For channels:**
|
|
||||||
```bash
|
|
||||||
export PICOCLAW_CHANNELS_TELEGRAM_TOKEN="token-from-env"
|
|
||||||
```
|
|
||||||
|
|
||||||
**For web tools:**
|
|
||||||
```bash
|
|
||||||
export PICOCLAW_WEB_BRAVE_API_KEY="key-from-env"
|
|
||||||
```
|
|
||||||
|
|
||||||
Environment variables follow this pattern: `PICOCLAW_<SECTION>_<KEY1>_<KEY2>_<FIELD>` with dots replaced by underscores and converted to uppercase.
|
|
||||||
|
|
||||||
### Multiple API Keys Not Working
|
|
||||||
|
|
||||||
- Ensure you're using `api_keys` (plural) in `.security.yml` for models and web tools (except GLMSearch)
|
|
||||||
- Check that the array format is correct in YAML (proper indentation)
|
|
||||||
- Remember: Models, Brave, Tavily, Perplexity MUST use `api_keys` (array format)
|
|
||||||
- GLMSearch MUST use `api_key` (single string format)
|
|
||||||
- The reference in `config.json` is the same regardless of single or multiple keys
|
|
||||||
|
|
||||||
### Load Balancing/Failover Issues
|
|
||||||
|
|
||||||
- Verify all API keys in the `api_keys` array are valid
|
|
||||||
- Check that all keys have the same rate limits and permissions
|
|
||||||
- Monitor logs to see which keys are being used and failing
|
|
||||||
|
|
@ -321,10 +321,7 @@ type AgentDefaults struct {
|
||||||
ToolFeedback ToolFeedbackConfig `json:"tool_feedback,omitempty"`
|
ToolFeedback ToolFeedbackConfig `json:"tool_feedback,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const DefaultMaxMediaSize = 20 * 1024 * 1024 // 20 MB
|
||||||
DefaultMaxMediaSize = 20 * 1024 * 1024 // 20 MB
|
|
||||||
DefaultWeComAIBotProcessingMessage = "⏳ Processing, please wait. The results will be sent shortly."
|
|
||||||
)
|
|
||||||
|
|
||||||
func (d *AgentDefaults) GetMaxMediaSize() int {
|
func (d *AgentDefaults) GetMaxMediaSize() int {
|
||||||
if d.MaxMediaSize > 0 {
|
if d.MaxMediaSize > 0 {
|
||||||
|
|
@ -364,9 +361,7 @@ type ChannelsConfig struct {
|
||||||
Matrix MatrixConfig `json:"matrix"`
|
Matrix MatrixConfig `json:"matrix"`
|
||||||
LINE LINEConfig `json:"line"`
|
LINE LINEConfig `json:"line"`
|
||||||
OneBot OneBotConfig `json:"onebot"`
|
OneBot OneBotConfig `json:"onebot"`
|
||||||
WeCom WeComConfig `json:"wecom"`
|
WeCom WeComConfig `json:"wecom" envPrefix:"PICOCLAW_CHANNELS_WECOM_"`
|
||||||
WeComApp WeComAppConfig `json:"wecom_app"`
|
|
||||||
WeComAIBot WeComAIBotConfig `json:"wecom_aibot"`
|
|
||||||
Weixin WeixinConfig `json:"weixin"`
|
Weixin WeixinConfig `json:"weixin"`
|
||||||
Pico PicoConfig `json:"pico"`
|
Pico PicoConfig `json:"pico"`
|
||||||
PicoClient PicoClientConfig `json:"pico_client"`
|
PicoClient PicoClientConfig `json:"pico_client"`
|
||||||
|
|
@ -386,7 +381,7 @@ type TypingConfig struct {
|
||||||
|
|
||||||
// PlaceholderConfig controls placeholder message behavior (Phase 10).
|
// PlaceholderConfig controls placeholder message behavior (Phase 10).
|
||||||
type PlaceholderConfig struct {
|
type PlaceholderConfig struct {
|
||||||
Enabled bool `json:"enabled,omitempty"`
|
Enabled bool `json:"enabled"`
|
||||||
Text string `json:"text,omitempty"`
|
Text string `json:"text,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -591,18 +586,20 @@ func (c *SlackConfig) SetAppToken(token string) {
|
||||||
}
|
}
|
||||||
|
|
||||||
type MatrixConfig struct {
|
type MatrixConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_MATRIX_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_MATRIX_ENABLED"`
|
||||||
Homeserver string `json:"homeserver" env:"PICOCLAW_CHANNELS_MATRIX_HOMESERVER"`
|
Homeserver string `json:"homeserver" env:"PICOCLAW_CHANNELS_MATRIX_HOMESERVER"`
|
||||||
UserID string `json:"user_id" env:"PICOCLAW_CHANNELS_MATRIX_USER_ID"`
|
UserID string `json:"user_id" env:"PICOCLAW_CHANNELS_MATRIX_USER_ID"`
|
||||||
accessToken string
|
accessToken string
|
||||||
DeviceID string `json:"device_id,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_DEVICE_ID"`
|
DeviceID string `json:"device_id,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_DEVICE_ID"`
|
||||||
JoinOnInvite bool `json:"join_on_invite" env:"PICOCLAW_CHANNELS_MATRIX_JOIN_ON_INVITE"`
|
JoinOnInvite bool `json:"join_on_invite" env:"PICOCLAW_CHANNELS_MATRIX_JOIN_ON_INVITE"`
|
||||||
MessageFormat string `json:"message_format,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_MESSAGE_FORMAT"`
|
MessageFormat string `json:"message_format,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_MESSAGE_FORMAT"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_MATRIX_ALLOW_FROM"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_MATRIX_ALLOW_FROM"`
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
||||||
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
|
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_MATRIX_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_MATRIX_REASONING_CHANNEL_ID"`
|
||||||
secDirty bool
|
secDirty bool
|
||||||
|
CryptoDatabasePath string `json:"crypto_database_path,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_CRYPTO_DATABASE_PATH"`
|
||||||
|
CryptoPassphrase string `json:"crypto_passphrase,omitempty" env:"PICOCLAW_CHANNELS_MATRIX_CRYPTO_PASSPHRASE"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// AccessToken returns the Matrix access token
|
// AccessToken returns the Matrix access token
|
||||||
|
|
@ -678,136 +675,28 @@ func (c *OneBotConfig) SetAccessToken(token string) {
|
||||||
c.secDirty = true
|
c.secDirty = true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type WeComGroupConfig struct {
|
||||||
|
AllowFrom FlexibleStringSlice `json:"allow_from,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
type WeComConfig struct {
|
type WeComConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_ENABLED"`
|
Enabled bool `json:"enabled" env:"ENABLED"`
|
||||||
token string
|
BotID string `json:"bot_id" env:"BOT_ID"`
|
||||||
encodingAESKey string
|
secret string
|
||||||
WebhookURL string `json:"webhook_url" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_URL"`
|
WebSocketURL string `json:"websocket_url,omitempty" env:"WEBSOCKET_URL"`
|
||||||
WebhookHost string `json:"webhook_host" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_HOST"`
|
SendThinkingMessage bool `json:"send_thinking_message" env:"SEND_THINKING_MESSAGE"`
|
||||||
WebhookPort int `json:"webhook_port" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_PORT"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"ALLOW_FROM"`
|
||||||
WebhookPath string `json:"webhook_path" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_PATH"`
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"REASONING_CHANNEL_ID"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_ALLOW_FROM"`
|
secDirty bool
|
||||||
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_REPLY_TIMEOUT"`
|
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_REASONING_CHANNEL_ID"`
|
|
||||||
secDirty bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Token returns the WeCom token
|
// Secret returns the WeCom bot secret.
|
||||||
func (c *WeComConfig) Token() string {
|
func (c *WeComConfig) Secret() string {
|
||||||
return c.token
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetToken sets the WeCom token
|
|
||||||
func (c *WeComConfig) SetToken(token string) {
|
|
||||||
c.token = token
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// EncodingAESKey returns the WeCom encoding AES key
|
|
||||||
func (c *WeComConfig) EncodingAESKey() string {
|
|
||||||
return c.encodingAESKey
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetEncodingAESKey sets the WeCom encoding AES key
|
|
||||||
func (c *WeComConfig) SetEncodingAESKey(key string) {
|
|
||||||
c.encodingAESKey = key
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
type WeComAppConfig struct {
|
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_APP_ENABLED"`
|
|
||||||
CorpID string `json:"corp_id" env:"PICOCLAW_CHANNELS_WECOM_APP_CORP_ID"`
|
|
||||||
corpSecret string
|
|
||||||
AgentID int64 `json:"agent_id" env:"PICOCLAW_CHANNELS_WECOM_APP_AGENT_ID"`
|
|
||||||
token string
|
|
||||||
encodingAESKey string
|
|
||||||
WebhookHost string `json:"webhook_host" env:"PICOCLAW_CHANNELS_WECOM_APP_WEBHOOK_HOST"`
|
|
||||||
WebhookPort int `json:"webhook_port" env:"PICOCLAW_CHANNELS_WECOM_APP_WEBHOOK_PORT"`
|
|
||||||
WebhookPath string `json:"webhook_path" env:"PICOCLAW_CHANNELS_WECOM_APP_WEBHOOK_PATH"`
|
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_APP_ALLOW_FROM"`
|
|
||||||
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_APP_REPLY_TIMEOUT"`
|
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_APP_REASONING_CHANNEL_ID"`
|
|
||||||
secDirty bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// CorpSecret returns the corporate secret for WeCom app
|
|
||||||
func (c *WeComAppConfig) CorpSecret() string {
|
|
||||||
return c.corpSecret
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetCorpSecret sets the corporate secret for WeCom app
|
|
||||||
func (c *WeComAppConfig) SetCorpSecret(secret string) {
|
|
||||||
c.corpSecret = secret
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Token returns the webhook token for WeCom app
|
|
||||||
func (c *WeComAppConfig) Token() string {
|
|
||||||
return c.token
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetToken sets the webhook token for WeCom app
|
|
||||||
func (c *WeComAppConfig) SetToken(token string) {
|
|
||||||
c.token = token
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// EncodingAESKey returns the encoding AES key for WeCom app
|
|
||||||
func (c *WeComAppConfig) EncodingAESKey() string {
|
|
||||||
return c.encodingAESKey
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetEncodingAESKey sets the encoding AES key for WeCom app
|
|
||||||
func (c *WeComAppConfig) SetEncodingAESKey(key string) {
|
|
||||||
c.encodingAESKey = key
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
type WeComAIBotConfig struct {
|
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ENABLED"`
|
|
||||||
BotID string `json:"bot_id,omitempty" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_BOT_ID"`
|
|
||||||
secret string
|
|
||||||
token string
|
|
||||||
encodingAESKey string
|
|
||||||
WebhookPath string `json:"webhook_path,omitempty" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_WEBHOOK_PATH"`
|
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ALLOW_FROM"`
|
|
||||||
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_REPLY_TIMEOUT"`
|
|
||||||
MaxSteps int `json:"max_steps" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_MAX_STEPS"` // Maximum streaming steps
|
|
||||||
WelcomeMessage string `json:"welcome_message" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_WELCOME_MESSAGE"` // Sent on enter_chat event; empty = no welcome
|
|
||||||
ProcessingMessage string `json:"processing_message,omitempty" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_PROCESSING_MESSAGE"`
|
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_REASONING_CHANNEL_ID"`
|
|
||||||
secDirty bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// Token returns the webhook token for WeCom AI bot
|
|
||||||
func (c *WeComAIBotConfig) Token() string {
|
|
||||||
return c.token
|
|
||||||
}
|
|
||||||
|
|
||||||
// EncodingAESKey returns the encoding AES key for WeCom AI bot
|
|
||||||
func (c *WeComAIBotConfig) EncodingAESKey() string {
|
|
||||||
return c.encodingAESKey
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetToken sets the token for WeCom AI bot
|
|
||||||
func (c *WeComAIBotConfig) SetToken(token string) {
|
|
||||||
c.token = token
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetEncodingAESKey sets the encoding AES key for WeCom AI bot
|
|
||||||
func (c *WeComAIBotConfig) SetEncodingAESKey(key string) {
|
|
||||||
c.encodingAESKey = key
|
|
||||||
c.secDirty = true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *WeComAIBotConfig) Secret() string {
|
|
||||||
return c.secret
|
return c.secret
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WeComAIBotConfig) SetSecret(secret string) {
|
// SetSecret sets the WeCom bot secret.
|
||||||
|
func (c *WeComConfig) SetSecret(secret string) {
|
||||||
c.secret = secret
|
c.secret = secret
|
||||||
c.secDirty = true
|
c.secDirty = true
|
||||||
}
|
}
|
||||||
|
|
@ -968,6 +857,10 @@ type ModelConfig struct {
|
||||||
secModelName string
|
secModelName string
|
||||||
apiKeys []string
|
apiKeys []string
|
||||||
secDirty bool
|
secDirty bool
|
||||||
|
|
||||||
|
// isVirtual marks this model as a virtual model generated from multi-key expansion.
|
||||||
|
// Virtual models should not be persisted to config files.
|
||||||
|
isVirtual bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// APIKey returns the first API key from apiKeys
|
// APIKey returns the first API key from apiKeys
|
||||||
|
|
@ -978,6 +871,11 @@ func (c *ModelConfig) APIKey() string {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// IsVirtual returns true if this model was generated from multi-key expansion.
|
||||||
|
func (c *ModelConfig) IsVirtual() bool {
|
||||||
|
return c.isVirtual
|
||||||
|
}
|
||||||
|
|
||||||
// Validate checks if the ModelConfig has all required fields.
|
// Validate checks if the ModelConfig has all required fields.
|
||||||
func (c *ModelConfig) Validate() error {
|
func (c *ModelConfig) Validate() error {
|
||||||
if c.ModelName == "" {
|
if c.ModelName == "" {
|
||||||
|
|
@ -1635,39 +1533,10 @@ func applySecurityConfig(cfg *Config, sec *SecurityConfig) error {
|
||||||
cfg.Channels.OneBot.accessToken = sec.Channels.OneBot.AccessToken
|
cfg.Channels.OneBot.accessToken = sec.Channels.OneBot.AccessToken
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handle WeCom token and encoding key
|
// Handle WeCom bot secret
|
||||||
if sec.Channels.WeCom != nil {
|
if sec.Channels.WeCom != nil {
|
||||||
if sec.Channels.WeCom.Token != "" {
|
if sec.Channels.WeCom.Secret != "" {
|
||||||
cfg.Channels.WeCom.token = sec.Channels.WeCom.Token
|
cfg.Channels.WeCom.secret = sec.Channels.WeCom.Secret
|
||||||
}
|
|
||||||
if sec.Channels.WeCom.EncodingAESKey != "" {
|
|
||||||
cfg.Channels.WeCom.encodingAESKey = sec.Channels.WeCom.EncodingAESKey
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handle WeCom App credentials
|
|
||||||
if sec.Channels.WeComApp != nil {
|
|
||||||
if sec.Channels.WeComApp.CorpSecret != "" {
|
|
||||||
cfg.Channels.WeComApp.corpSecret = sec.Channels.WeComApp.CorpSecret
|
|
||||||
}
|
|
||||||
if sec.Channels.WeComApp.Token != "" {
|
|
||||||
cfg.Channels.WeComApp.token = sec.Channels.WeComApp.Token
|
|
||||||
}
|
|
||||||
if sec.Channels.WeComApp.EncodingAESKey != "" {
|
|
||||||
cfg.Channels.WeComApp.encodingAESKey = sec.Channels.WeComApp.EncodingAESKey
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handle WeCom AI Bot credentials
|
|
||||||
if sec.Channels.WeComAIBot != nil {
|
|
||||||
if sec.Channels.WeComAIBot.Token != "" {
|
|
||||||
cfg.Channels.WeComAIBot.token = sec.Channels.WeComAIBot.Token
|
|
||||||
}
|
|
||||||
if sec.Channels.WeComAIBot.EncodingAESKey != "" {
|
|
||||||
cfg.Channels.WeComAIBot.encodingAESKey = sec.Channels.WeComAIBot.EncodingAESKey
|
|
||||||
}
|
|
||||||
if sec.Channels.WeComAIBot.Secret != "" {
|
|
||||||
cfg.Channels.WeComAIBot.secret = sec.Channels.WeComAIBot.Secret
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1913,27 +1782,10 @@ func SaveConfig(path string, cfg *Config) error {
|
||||||
}
|
}
|
||||||
if cfg.Channels.WeCom.secDirty {
|
if cfg.Channels.WeCom.secDirty {
|
||||||
cfg.security.Channels.WeCom = &WeComSecurity{
|
cfg.security.Channels.WeCom = &WeComSecurity{
|
||||||
Token: cfg.Channels.WeCom.Token(),
|
Secret: cfg.Channels.WeCom.Secret(),
|
||||||
EncodingAESKey: cfg.Channels.WeCom.EncodingAESKey(),
|
|
||||||
}
|
}
|
||||||
cfg.Channels.WeCom.secDirty = false
|
cfg.Channels.WeCom.secDirty = false
|
||||||
}
|
}
|
||||||
if cfg.Channels.WeComApp.secDirty {
|
|
||||||
cfg.security.Channels.WeComApp = &WeComAppSecurity{
|
|
||||||
CorpSecret: cfg.Channels.WeComApp.CorpSecret(),
|
|
||||||
Token: cfg.Channels.WeComApp.Token(),
|
|
||||||
EncodingAESKey: cfg.Channels.WeComApp.EncodingAESKey(),
|
|
||||||
}
|
|
||||||
cfg.Channels.WeComApp.secDirty = false
|
|
||||||
}
|
|
||||||
if cfg.Channels.WeComAIBot.secDirty {
|
|
||||||
cfg.security.Channels.WeComAIBot = &WeComAIBotSecurity{
|
|
||||||
Token: cfg.Channels.WeComAIBot.Token(),
|
|
||||||
EncodingAESKey: cfg.Channels.WeComAIBot.EncodingAESKey(),
|
|
||||||
Secret: cfg.Channels.WeComAIBot.Secret(),
|
|
||||||
}
|
|
||||||
cfg.Channels.WeComAIBot.secDirty = false
|
|
||||||
}
|
|
||||||
if cfg.Tools.Web.Brave.secDirty {
|
if cfg.Tools.Web.Brave.secDirty {
|
||||||
cfg.security.Web.Brave = &BraveSecurity{
|
cfg.security.Web.Brave = &BraveSecurity{
|
||||||
APIKeys: cfg.Tools.Web.Brave.APIKeys(),
|
APIKeys: cfg.Tools.Web.Brave.APIKeys(),
|
||||||
|
|
@ -1991,7 +1843,20 @@ func SaveConfig(path string, cfg *Config) error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Filter out virtual models before serializing to config file
|
||||||
|
nonVirtualModels := make([]*ModelConfig, 0, len(cfg.ModelList))
|
||||||
|
for _, m := range cfg.ModelList {
|
||||||
|
if !m.isVirtual {
|
||||||
|
nonVirtualModels = append(nonVirtualModels, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Temporarily replace ModelList with filtered version for serialization
|
||||||
|
originalModelList := cfg.ModelList
|
||||||
|
cfg.ModelList = nonVirtualModels
|
||||||
|
|
||||||
data, err := json.MarshalIndent(cfg, "", " ")
|
data, err := json.MarshalIndent(cfg, "", " ")
|
||||||
|
// Restore original ModelList after serialization
|
||||||
|
cfg.ModelList = originalModelList
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
@ -2209,6 +2074,7 @@ func expandMultiKeyModels(models []*ModelConfig) []*ModelConfig {
|
||||||
RequestTimeout: m.RequestTimeout,
|
RequestTimeout: m.RequestTimeout,
|
||||||
ThinkingLevel: m.ThinkingLevel,
|
ThinkingLevel: m.ThinkingLevel,
|
||||||
ExtraBody: m.ExtraBody,
|
ExtraBody: m.ExtraBody,
|
||||||
|
isVirtual: true,
|
||||||
}
|
}
|
||||||
expanded = append(expanded, additionalEntry)
|
expanded = append(expanded, additionalEntry)
|
||||||
fallbackNames = append(fallbackNames, expandedName)
|
fallbackNames = append(fallbackNames, expandedName)
|
||||||
|
|
|
||||||
|
|
@ -85,23 +85,21 @@ type toolsConfigV0 struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type channelsConfigV0 struct {
|
type channelsConfigV0 struct {
|
||||||
WhatsApp WhatsAppConfig `json:"whatsapp"`
|
WhatsApp WhatsAppConfig `json:"whatsapp"`
|
||||||
Telegram telegramConfigV0 `json:"telegram"`
|
Telegram telegramConfigV0 `json:"telegram"`
|
||||||
Feishu feishuConfigV0 `json:"feishu"`
|
Feishu feishuConfigV0 `json:"feishu"`
|
||||||
Discord discordConfigV0 `json:"discord"`
|
Discord discordConfigV0 `json:"discord"`
|
||||||
MaixCam maixcamConfigV0 `json:"maixcam"`
|
MaixCam maixcamConfigV0 `json:"maixcam"`
|
||||||
Weixin weixinConfigV0 `json:"weixin"`
|
Weixin weixinConfigV0 `json:"weixin"`
|
||||||
QQ qqConfigV0 `json:"qq"`
|
QQ qqConfigV0 `json:"qq"`
|
||||||
DingTalk dingtalkConfigV0 `json:"dingtalk"`
|
DingTalk dingtalkConfigV0 `json:"dingtalk"`
|
||||||
Slack slackConfigV0 `json:"slack"`
|
Slack slackConfigV0 `json:"slack"`
|
||||||
Matrix matrixConfigV0 `json:"matrix"`
|
Matrix matrixConfigV0 `json:"matrix"`
|
||||||
LINE lineConfigV0 `json:"line"`
|
LINE lineConfigV0 `json:"line"`
|
||||||
OneBot onebotConfigV0 `json:"onebot"`
|
OneBot onebotConfigV0 `json:"onebot"`
|
||||||
WeCom wecomConfigV0 `json:"wecom"`
|
WeCom wecomConfigV0 `json:"wecom" envPrefix:"PICOCLAW_CHANNELS_WECOM_"`
|
||||||
WeComApp wecomappConfigV0 `json:"wecom_app"`
|
Pico picoConfigV0 `json:"pico"`
|
||||||
WeComAIBot wecomaibotConfigV0 `json:"wecom_aibot"`
|
IRC ircConfigV0 `json:"irc"`
|
||||||
Pico picoConfigV0 `json:"pico"`
|
|
||||||
IRC ircConfigV0 `json:"irc"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *channelsConfigV0) ToChannelsConfig() (ChannelsConfig, ChannelsSecurity) {
|
func (v *channelsConfigV0) ToChannelsConfig() (ChannelsConfig, ChannelsSecurity) {
|
||||||
|
|
@ -117,45 +115,39 @@ func (v *channelsConfigV0) ToChannelsConfig() (ChannelsConfig, ChannelsSecurity)
|
||||||
line, lineSecurity := v.LINE.ToLINEConfig()
|
line, lineSecurity := v.LINE.ToLINEConfig()
|
||||||
onebot, onebotSecurity := v.OneBot.ToOneBotConfig()
|
onebot, onebotSecurity := v.OneBot.ToOneBotConfig()
|
||||||
wecom, wecomSecurity := v.WeCom.ToWeComConfig()
|
wecom, wecomSecurity := v.WeCom.ToWeComConfig()
|
||||||
wecomapp, wecomappSecurity := v.WeComApp.ToWeComAppConfig()
|
|
||||||
wecomaibot, wecomaibotSecurity := v.WeComAIBot.ToWeComAIBotConfig()
|
|
||||||
pico, picoSecurity := v.Pico.ToPicoConfig()
|
pico, picoSecurity := v.Pico.ToPicoConfig()
|
||||||
irc, ircSecurity := v.IRC.ToIRCConfig()
|
irc, ircSecurity := v.IRC.ToIRCConfig()
|
||||||
|
|
||||||
return ChannelsConfig{
|
return ChannelsConfig{
|
||||||
WhatsApp: v.WhatsApp,
|
WhatsApp: v.WhatsApp,
|
||||||
Telegram: telegram,
|
Telegram: telegram,
|
||||||
Feishu: feishu,
|
Feishu: feishu,
|
||||||
Discord: discord,
|
Discord: discord,
|
||||||
MaixCam: maixcam,
|
MaixCam: maixcam,
|
||||||
QQ: qq,
|
QQ: qq,
|
||||||
Weixin: weixin,
|
Weixin: weixin,
|
||||||
DingTalk: dingtalk,
|
DingTalk: dingtalk,
|
||||||
Slack: slack,
|
Slack: slack,
|
||||||
Matrix: matrix,
|
Matrix: matrix,
|
||||||
LINE: line,
|
LINE: line,
|
||||||
OneBot: onebot,
|
OneBot: onebot,
|
||||||
WeCom: wecom,
|
WeCom: wecom,
|
||||||
WeComApp: wecomapp,
|
Pico: pico,
|
||||||
WeComAIBot: wecomaibot,
|
IRC: irc,
|
||||||
Pico: pico,
|
|
||||||
IRC: irc,
|
|
||||||
}, ChannelsSecurity{
|
}, ChannelsSecurity{
|
||||||
Telegram: telegramSecurity,
|
Telegram: telegramSecurity,
|
||||||
Feishu: feishuSecurity,
|
Feishu: feishuSecurity,
|
||||||
Discord: discordSecurity,
|
Discord: discordSecurity,
|
||||||
QQ: qqSecurity,
|
QQ: qqSecurity,
|
||||||
Weixin: weixinSecurity,
|
Weixin: weixinSecurity,
|
||||||
DingTalk: dingtalkSecurity,
|
DingTalk: dingtalkSecurity,
|
||||||
Slack: slackSecurity,
|
Slack: slackSecurity,
|
||||||
Matrix: matrixSecurity,
|
Matrix: matrixSecurity,
|
||||||
LINE: lineSecurity,
|
LINE: lineSecurity,
|
||||||
OneBot: onebotSecurity,
|
OneBot: onebotSecurity,
|
||||||
WeCom: wecomSecurity,
|
WeCom: wecomSecurity,
|
||||||
WeComApp: wecomappSecurity,
|
Pico: picoSecurity,
|
||||||
WeComAIBot: wecomaibotSecurity,
|
IRC: ircSecurity,
|
||||||
Pico: picoSecurity,
|
|
||||||
IRC: ircSecurity,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -473,39 +465,32 @@ func (v *onebotConfigV0) ToOneBotConfig() (OneBotConfig, *OneBotSecurity) {
|
||||||
}
|
}
|
||||||
|
|
||||||
type wecomConfigV0 struct {
|
type wecomConfigV0 struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_ENABLED"`
|
Enabled bool `json:"enabled" env:"ENABLED"`
|
||||||
Token string `json:"token" env:"PICOCLAW_CHANNELS_WECOM_TOKEN"`
|
BotID string `json:"bot_id" env:"BOT_ID"`
|
||||||
EncodingAESKey string `json:"encoding_aes_key" env:"PICOCLAW_CHANNELS_WECOM_ENCODING_AES_KEY"`
|
Secret string `json:"secret" env:"SECRET"`
|
||||||
WebhookURL string `json:"webhook_url" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_URL"`
|
WebSocketURL string `json:"websocket_url,omitempty" env:"WEBSOCKET_URL"`
|
||||||
WebhookHost string `json:"webhook_host" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_HOST"`
|
SendThinkingMessage bool `json:"send_thinking_message" env:"SEND_THINKING_MESSAGE"`
|
||||||
WebhookPort int `json:"webhook_port" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_PORT"`
|
DMPolicy string `json:"dm_policy,omitempty" env:"DM_POLICY"`
|
||||||
WebhookPath string `json:"webhook_path" env:"PICOCLAW_CHANNELS_WECOM_WEBHOOK_PATH"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"ALLOW_FROM"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_ALLOW_FROM"`
|
GroupPolicy string `json:"group_policy,omitempty" env:"GROUP_POLICY"`
|
||||||
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_REPLY_TIMEOUT"`
|
GroupAllowFrom FlexibleStringSlice `json:"group_allow_from,omitempty" env:"GROUP_ALLOW_FROM"`
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
Groups map[string]WeComGroupConfig `json:"groups,omitempty"`
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"REASONING_CHANNEL_ID"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *wecomConfigV0) ToWeComConfig() (WeComConfig, *WeComSecurity) {
|
func (v *wecomConfigV0) ToWeComConfig() (WeComConfig, *WeComSecurity) {
|
||||||
var sec *WeComSecurity
|
var sec *WeComSecurity
|
||||||
if v.Token != "" || v.EncodingAESKey != "" {
|
if v.Secret != "" {
|
||||||
sec = &WeComSecurity{
|
sec = &WeComSecurity{Secret: v.Secret}
|
||||||
Token: v.Token,
|
|
||||||
EncodingAESKey: v.EncodingAESKey,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return WeComConfig{
|
return WeComConfig{
|
||||||
Enabled: v.Enabled,
|
Enabled: v.Enabled,
|
||||||
token: v.Token,
|
BotID: v.BotID,
|
||||||
encodingAESKey: v.EncodingAESKey,
|
secret: v.Secret,
|
||||||
WebhookURL: v.WebhookURL,
|
WebSocketURL: v.WebSocketURL,
|
||||||
WebhookHost: v.WebhookHost,
|
SendThinkingMessage: v.SendThinkingMessage,
|
||||||
WebhookPort: v.WebhookPort,
|
AllowFrom: v.AllowFrom,
|
||||||
WebhookPath: v.WebhookPath,
|
ReasoningChannelID: v.ReasoningChannelID,
|
||||||
AllowFrom: v.AllowFrom,
|
|
||||||
ReplyTimeout: v.ReplyTimeout,
|
|
||||||
GroupTrigger: v.GroupTrigger,
|
|
||||||
ReasoningChannelID: v.ReasoningChannelID,
|
|
||||||
}, sec
|
}, sec
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -537,81 +522,6 @@ func (v *weixinConfigV0) ToWeiXinConfig() (WeixinConfig, *WeixinSecurity) {
|
||||||
}, sec
|
}, sec
|
||||||
}
|
}
|
||||||
|
|
||||||
type wecomappConfigV0 struct {
|
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_APP_ENABLED"`
|
|
||||||
CorpID string `json:"corp_id" env:"PICOCLAW_CHANNELS_WECOM_APP_CORP_ID"`
|
|
||||||
CorpSecret string `json:"corp_secret" env:"PICOCLAW_CHANNELS_WECOM_APP_CORP_SECRET"`
|
|
||||||
AgentID int64 `json:"agent_id" env:"PICOCLAW_CHANNELS_WECOM_APP_AGENT_ID"`
|
|
||||||
Token string `json:"token" env:"PICOCLAW_CHANNELS_WECOM_APP_TOKEN"`
|
|
||||||
EncodingAESKey string `json:"encoding_aes_key" env:"PICOCLAW_CHANNELS_WECOM_APP_ENCODING_AES_KEY"`
|
|
||||||
WebhookHost string `json:"webhook_host" env:"PICOCLAW_CHANNELS_WECOM_APP_WEBHOOK_HOST"`
|
|
||||||
WebhookPort int `json:"webhook_port" env:"PICOCLAW_CHANNELS_WECOM_APP_WEBHOOK_PORT"`
|
|
||||||
WebhookPath string `json:"webhook_path" env:"PICOCLAW_CHANNELS_WECOM_APP_WEBHOOK_PATH"`
|
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_APP_ALLOW_FROM"`
|
|
||||||
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_APP_REPLY_TIMEOUT"`
|
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_APP_REASONING_CHANNEL_ID"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func (v *wecomappConfigV0) ToWeComAppConfig() (WeComAppConfig, *WeComAppSecurity) {
|
|
||||||
var sec *WeComAppSecurity
|
|
||||||
if v.CorpSecret != "" || v.Token != "" || v.EncodingAESKey != "" {
|
|
||||||
sec = &WeComAppSecurity{
|
|
||||||
CorpSecret: v.CorpSecret,
|
|
||||||
Token: v.Token,
|
|
||||||
EncodingAESKey: v.EncodingAESKey,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return WeComAppConfig{
|
|
||||||
Enabled: v.Enabled,
|
|
||||||
CorpID: v.CorpID,
|
|
||||||
corpSecret: v.CorpSecret,
|
|
||||||
AgentID: v.AgentID,
|
|
||||||
token: v.Token,
|
|
||||||
encodingAESKey: v.EncodingAESKey,
|
|
||||||
WebhookHost: v.WebhookHost,
|
|
||||||
WebhookPort: v.WebhookPort,
|
|
||||||
WebhookPath: v.WebhookPath,
|
|
||||||
AllowFrom: v.AllowFrom,
|
|
||||||
ReplyTimeout: v.ReplyTimeout,
|
|
||||||
GroupTrigger: v.GroupTrigger,
|
|
||||||
ReasoningChannelID: v.ReasoningChannelID,
|
|
||||||
}, sec
|
|
||||||
}
|
|
||||||
|
|
||||||
type wecomaibotConfigV0 struct {
|
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ENABLED"`
|
|
||||||
Token string `json:"token" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_TOKEN"`
|
|
||||||
Secret string `json:"secret" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_SECRET"`
|
|
||||||
EncodingAESKey string `json:"encoding_aes_key" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ENCODING_AES_KEY"`
|
|
||||||
WebhookPath string `json:"webhook_path" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_WEBHOOK_PATH"`
|
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ALLOW_FROM"`
|
|
||||||
ReplyTimeout int `json:"reply_timeout" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_REPLY_TIMEOUT"`
|
|
||||||
MaxSteps int `json:"max_steps" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_MAX_STEPS"`
|
|
||||||
WelcomeMessage string `json:"welcome_message" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_WELCOME_MESSAGE"`
|
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_REASONING_CHANNEL_ID"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func (v *wecomaibotConfigV0) ToWeComAIBotConfig() (WeComAIBotConfig, *WeComAIBotSecurity) {
|
|
||||||
var sec *WeComAIBotSecurity
|
|
||||||
if v.Token != "" || v.Secret != "" || v.EncodingAESKey != "" {
|
|
||||||
sec = &WeComAIBotSecurity{
|
|
||||||
Token: v.Token,
|
|
||||||
Secret: v.Secret,
|
|
||||||
EncodingAESKey: v.EncodingAESKey,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return WeComAIBotConfig{
|
|
||||||
Enabled: v.Enabled,
|
|
||||||
WebhookPath: v.WebhookPath,
|
|
||||||
AllowFrom: v.AllowFrom,
|
|
||||||
ReplyTimeout: v.ReplyTimeout,
|
|
||||||
MaxSteps: v.MaxSteps,
|
|
||||||
WelcomeMessage: v.WelcomeMessage,
|
|
||||||
ReasoningChannelID: v.ReasoningChannelID,
|
|
||||||
}, sec
|
|
||||||
}
|
|
||||||
|
|
||||||
type picoConfigV0 struct {
|
type picoConfigV0 struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_PICO_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_PICO_ENABLED"`
|
||||||
Token string `json:"token" env:"PICOCLAW_CHANNELS_PICO_TOKEN"`
|
Token string `json:"token" env:"PICOCLAW_CHANNELS_PICO_TOKEN"`
|
||||||
|
|
|
||||||
|
|
@ -360,6 +360,96 @@ func TestSaveConfig_IncludesEmptyLegacyModelField(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSaveConfig_PreservesDisabledTelegramPlaceholder(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
path := filepath.Join(tmpDir, "config.json")
|
||||||
|
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
cfg.Channels.Telegram.Placeholder.Enabled = false
|
||||||
|
|
||||||
|
if err := SaveConfig(path, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ReadFile failed: %v", err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(string(data), `"placeholder": {`) {
|
||||||
|
t.Fatalf("saved config should include telegram placeholder config, got: %s", string(data))
|
||||||
|
}
|
||||||
|
if !strings.Contains(string(data), `"enabled": false`) {
|
||||||
|
t.Fatalf("saved config should persist placeholder.enabled=false, got: %s", string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
loaded, err := LoadConfig(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig failed: %v", err)
|
||||||
|
}
|
||||||
|
if loaded.Channels.Telegram.Placeholder.Enabled {
|
||||||
|
t.Fatal("telegram placeholder should remain disabled after SaveConfig/LoadConfig round-trip")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSaveConfig_FiltersVirtualModels verifies that SaveConfig does not write
|
||||||
|
// virtual models (generated by expandMultiKeyModels) to the config file.
|
||||||
|
func TestSaveConfig_FiltersVirtualModels(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
path := filepath.Join(tmpDir, "config.json")
|
||||||
|
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
|
// Manually add a virtual model to ModelList (simulating what expandMultiKeyModels does)
|
||||||
|
primaryModel := &ModelConfig{
|
||||||
|
ModelName: "gpt-4",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
apiKeys: []string{"key1"},
|
||||||
|
}
|
||||||
|
virtualModel := &ModelConfig{
|
||||||
|
ModelName: "gpt-4__key_1",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
apiKeys: []string{"key2"},
|
||||||
|
isVirtual: true,
|
||||||
|
}
|
||||||
|
cfg.ModelList = []*ModelConfig{primaryModel, virtualModel}
|
||||||
|
|
||||||
|
// SaveConfig should filter out virtual models
|
||||||
|
if err := SaveConfig(path, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reload and verify
|
||||||
|
reloaded, err := LoadConfig(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should only have the primary model, not the virtual one
|
||||||
|
if len(reloaded.ModelList) != 1 {
|
||||||
|
t.Fatalf("expected 1 model after reload, got %d", len(reloaded.ModelList))
|
||||||
|
}
|
||||||
|
|
||||||
|
if reloaded.ModelList[0].ModelName != "gpt-4" {
|
||||||
|
t.Errorf("expected model_name 'gpt-4', got %q", reloaded.ModelList[0].ModelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify virtual model was not persisted
|
||||||
|
for _, m := range reloaded.ModelList {
|
||||||
|
if m.ModelName == "gpt-4__key_1" {
|
||||||
|
t.Errorf("virtual model gpt-4__key_1 should not have been saved")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify the saved file does not contain the virtual model name
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ReadFile failed: %v", err)
|
||||||
|
}
|
||||||
|
if strings.Contains(string(data), "gpt-4__key_1") {
|
||||||
|
t.Errorf("saved config should not contain virtual model name 'gpt-4__key_1'")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestConfig_Complete verifies all config fields are set
|
// TestConfig_Complete verifies all config fields are set
|
||||||
func TestConfig_Complete(t *testing.T) {
|
func TestConfig_Complete(t *testing.T) {
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
@ -1372,8 +1462,7 @@ func TestFilterSensitiveData_AllTokenTypes(t *testing.T) {
|
||||||
Feishu: &FeishuSecurity{AppSecret: "feishu-app-secret-123", EncryptKey: "feishu-encrypt-key"},
|
Feishu: &FeishuSecurity{AppSecret: "feishu-app-secret-123", EncryptKey: "feishu-encrypt-key"},
|
||||||
DingTalk: &DingTalkSecurity{ClientSecret: "dingtalk-client-secret"},
|
DingTalk: &DingTalkSecurity{ClientSecret: "dingtalk-client-secret"},
|
||||||
OneBot: &OneBotSecurity{AccessToken: "onebot-access-token"},
|
OneBot: &OneBotSecurity{AccessToken: "onebot-access-token"},
|
||||||
WeCom: &WeComSecurity{Token: "wecom-token", EncodingAESKey: "wecom-aes-key"},
|
WeCom: &WeComSecurity{Secret: "wecom-secret"},
|
||||||
WeComApp: &WeComAppSecurity{CorpSecret: "wecom-app-secret", Token: "wecom-app-token"},
|
|
||||||
Pico: &PicoSecurity{Token: "pico-token-abc123"},
|
Pico: &PicoSecurity{Token: "pico-token-abc123"},
|
||||||
IRC: &IRCSecurity{
|
IRC: &IRCSecurity{
|
||||||
Password: "irc-password",
|
Password: "irc-password",
|
||||||
|
|
|
||||||
|
|
@ -113,6 +113,8 @@ func DefaultConfig() *Config {
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
Text: "Thinking... 💭",
|
Text: "Thinking... 💭",
|
||||||
},
|
},
|
||||||
|
CryptoDatabasePath: "",
|
||||||
|
CryptoPassphrase: "",
|
||||||
},
|
},
|
||||||
LINE: LINEConfig{
|
LINE: LINEConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
|
|
@ -129,32 +131,11 @@ func DefaultConfig() *Config {
|
||||||
AllowFrom: FlexibleStringSlice{},
|
AllowFrom: FlexibleStringSlice{},
|
||||||
},
|
},
|
||||||
WeCom: WeComConfig{
|
WeCom: WeComConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
WebhookURL: "",
|
BotID: "",
|
||||||
WebhookHost: "0.0.0.0",
|
WebSocketURL: "wss://openws.work.weixin.qq.com",
|
||||||
WebhookPort: 18793,
|
SendThinkingMessage: true,
|
||||||
WebhookPath: "/webhook/wecom",
|
AllowFrom: FlexibleStringSlice{},
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
ReplyTimeout: 5,
|
|
||||||
},
|
|
||||||
WeComApp: WeComAppConfig{
|
|
||||||
Enabled: false,
|
|
||||||
CorpID: "",
|
|
||||||
AgentID: 0,
|
|
||||||
WebhookHost: "0.0.0.0",
|
|
||||||
WebhookPort: 18792,
|
|
||||||
WebhookPath: "/webhook/wecom-app",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
ReplyTimeout: 5,
|
|
||||||
},
|
|
||||||
WeComAIBot: WeComAIBotConfig{
|
|
||||||
Enabled: false,
|
|
||||||
WebhookPath: "/webhook/wecom-aibot",
|
|
||||||
AllowFrom: FlexibleStringSlice{},
|
|
||||||
ReplyTimeout: 5,
|
|
||||||
MaxSteps: 10,
|
|
||||||
WelcomeMessage: "Hello! I'm your AI assistant. How can I help you today?",
|
|
||||||
ProcessingMessage: DefaultWeComAIBotProcessingMessage,
|
|
||||||
},
|
},
|
||||||
Weixin: WeixinConfig{
|
Weixin: WeixinConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
|
|
|
||||||
|
|
@ -11,20 +11,33 @@ Package config
|
||||||
|
|
||||||
# Example: Using Security Configuration
|
# Example: Using Security Configuration
|
||||||
|
|
||||||
## 1. Create security.yml
|
## Overview
|
||||||
|
|
||||||
File: ~/.picoclaw/security.yml
|
The security configuration feature allows you to separate sensitive data (API keys,
|
||||||
|
tokens, secrets, passwords) from your main configuration. The system automatically
|
||||||
|
loads values from `.security.yml` and applies them to the corresponding fields in
|
||||||
|
your config.
|
||||||
|
|
||||||
|
**Key Points:**
|
||||||
|
- Values from `.security.yml` are automatically mapped to config fields
|
||||||
|
- No `ref:` syntax is needed - just omit sensitive fields from config.json
|
||||||
|
- If a field exists in both files, `.security.yml` value takes precedence
|
||||||
|
- You can mix direct values in config.json with security values
|
||||||
|
|
||||||
|
## 1. Create .security.yml
|
||||||
|
|
||||||
|
File: ~/.picoclaw/.security.yml
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
# Model API Keys
|
# Model API Keys
|
||||||
# Note: Use 'api_keys' array for multiple keys (load balancing/failover)
|
# All models MUST use 'api_keys' (plural) array format
|
||||||
# Single key should be provided as an array with one element
|
# Even a single key must be provided as an array with one element
|
||||||
model_list:
|
model_list:
|
||||||
|
|
||||||
gpt-5.4:
|
gpt-5.4:
|
||||||
api_keys:
|
api_keys:
|
||||||
- "sk-proj-your-actual-openai-key-1"
|
- "sk-proj-your-actual-openai-key-1"
|
||||||
- "sk-proj-your-actual-openai-key-2" # Failover key
|
- "sk-proj-your-actual-openai-key-2" # Optional: Multiple keys for failover
|
||||||
claude-sonnet-4.6:
|
claude-sonnet-4.6:
|
||||||
api_keys:
|
api_keys:
|
||||||
- "sk-ant-your-actual-anthropic-key" # Single key in array format
|
- "sk-ant-your-actual-anthropic-key" # Single key in array format
|
||||||
|
|
@ -38,80 +51,95 @@ channels:
|
||||||
token: "your-discord-bot-token"
|
token: "your-discord-bot-token"
|
||||||
|
|
||||||
# Web Tool Keys
|
# Web Tool Keys
|
||||||
# Note: Use 'api_keys' array for multiple keys (load balancing/failover)
|
# Brave, Tavily, Perplexity: Use 'api_keys' array
|
||||||
# For GLMSearch, use 'api_key' (single string)
|
# GLMSearch, BaiduSearch: Use 'api_key' single string
|
||||||
web:
|
web:
|
||||||
|
|
||||||
brave:
|
brave:
|
||||||
api_keys:
|
api_keys:
|
||||||
- "BSAyour-brave-api-key-1"
|
- "BSAyour-brave-api-key-1"
|
||||||
- "BSAyour-brave-api-key-2" # Failover key
|
- "BSAyour-brave-api-key-2" # Optional: Multiple keys for failover
|
||||||
tavily:
|
tavily:
|
||||||
api_keys:
|
api_keys:
|
||||||
- "tvly-your-tavily-api-key" # Single key in array format
|
- "tvly-your-tavily-api-key" # Single key in array format
|
||||||
|
perplexity:
|
||||||
|
api_keys:
|
||||||
|
- "pplx-your-perplexity-api-key" # Single key in array format
|
||||||
glm_search:
|
glm_search:
|
||||||
api_key: "your-glm-search-api-key" # Single key (not array)
|
api_key: "your-glm-search-api-key" # Single key (not array)
|
||||||
|
baidu_search:
|
||||||
|
api_key: "your-baidu-search-api-key" # Single key (not array)
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## 2. Update config.json to use references
|
## 2. Simplify config.json
|
||||||
|
|
||||||
File: ~/.picoclaw/config.json
|
File: ~/.picoclaw/config.json
|
||||||
|
|
||||||
|
Note: Sensitive fields are omitted because they're loaded from .security.yml
|
||||||
|
|
||||||
```json
|
```json
|
||||||
|
|
||||||
{
|
{
|
||||||
"version": 1,
|
"version": 1,
|
||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"workspace": "~/picoclaw-workspace",
|
"workspace": "~/picoclaw-workspace",
|
||||||
"model_name": "gpt-5.4"
|
"model_name": "gpt-5.4"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"model_list": [
|
"model_list": [
|
||||||
{
|
{
|
||||||
"model_name": "gpt-5.4",
|
"model_name": "gpt-5.4",
|
||||||
"model": "openai/gpt-5.4",
|
"model": "openai/gpt-5.4",
|
||||||
"api_base": "https://api.openai.com/v1",
|
"api_base": "https://api.openai.com/v1"
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
// api_key is automatically loaded from .security.yml
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"model_name": "claude-sonnet-4.6",
|
"model_name": "claude-sonnet-4.6",
|
||||||
"model": "anthropic/claude-sonnet-4.6",
|
"model": "anthropic/claude-sonnet-4.6",
|
||||||
"api_base": "https://api.anthropic.com/v1",
|
"api_base": "https://api.anthropic.com/v1"
|
||||||
"api_key": "ref:model_list.claude-sonnet-4.6.api_key"
|
// api_key is automatically loaded from .security.yml
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"channels": {
|
"channels": {
|
||||||
"telegram": {
|
"telegram": {
|
||||||
"enabled": true,
|
"enabled": true
|
||||||
"token": "ref:channels.telegram.token"
|
// token is automatically loaded from .security.yml
|
||||||
},
|
},
|
||||||
"discord": {
|
"discord": {
|
||||||
"enabled": true,
|
"enabled": true
|
||||||
"token": "ref:channels.discord.token"
|
// token is automatically loaded from .security.yml
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": {
|
"web": {
|
||||||
"brave": {
|
"brave": {
|
||||||
"enabled": true,
|
"enabled": true
|
||||||
"api_key": "ref:web.brave.api_key"
|
// api_key is automatically loaded from .security.yml
|
||||||
},
|
},
|
||||||
"tavily": {
|
"tavily": {
|
||||||
"enabled": true,
|
"enabled": true
|
||||||
"api_key": "ref:web.tavily.api_key"
|
// api_key is automatically loaded from .security.yml
|
||||||
|
},
|
||||||
|
"glm_search": {
|
||||||
|
"enabled": true
|
||||||
|
// api_key is automatically loaded from .security.yml
|
||||||
|
},
|
||||||
|
"baidu_search": {
|
||||||
|
"enabled": true
|
||||||
|
// api_key is automatically loaded from .security.yml
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## 3. Set proper permissions
|
## 3. Set proper permissions
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
chmod 600 ~/.picoclaw/security.yml
|
chmod 600 ~/.picoclaw/.security.yml
|
||||||
```
|
```
|
||||||
|
|
||||||
## 4. Add to .gitignore
|
## 4. Add to .gitignore
|
||||||
|
|
@ -127,57 +155,131 @@ chmod 600 ~/.picoclaw/security.yml
|
||||||
picoclaw --version
|
picoclaw --version
|
||||||
```
|
```
|
||||||
|
|
||||||
# Available Reference Paths
|
# Supported Fields in .security.yml
|
||||||
|
|
||||||
## Model API Keys
|
## Model API Keys
|
||||||
- ref:model_list.<model_name>.api_key
|
|
||||||
|
All models MUST use the `api_keys` (plural) array format in .security.yml.
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
model_list:
|
||||||
|
|
||||||
|
<model_name>:
|
||||||
|
api_keys:
|
||||||
|
- "key-1"
|
||||||
|
- "key-2" # Optional: Multiple keys for failover
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
- ref:model_list.gpt-5.4.api_key
|
```yaml
|
||||||
- ref:model_list.claude-sonnet-4.6.api_key
|
model_list:
|
||||||
|
|
||||||
**Note:** In .security.yml, use `api_keys` (array) format for models.
|
gpt-5.4:
|
||||||
Both single and multiple keys should use the array format.
|
api_keys:
|
||||||
|
- "sk-proj-key-1"
|
||||||
|
- "sk-proj-key-2"
|
||||||
|
claude-sonnet-4.6:
|
||||||
|
api_keys:
|
||||||
|
- "sk-ant-key"
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
**Important:**
|
||||||
|
- Always use `api_keys` (plural) for models
|
||||||
|
- Even a single key must be in an array format
|
||||||
|
- The model_name in .security.yml must match the model_name in config.json
|
||||||
|
|
||||||
## Channel Tokens/Secrets
|
## Channel Tokens/Secrets
|
||||||
- ref:channels.telegram.token
|
|
||||||
- ref:channels.feishu.app_secret
|
```yaml
|
||||||
- ref:channels.feishu.encrypt_key
|
channels:
|
||||||
- ref:channels.feishu.verification_token
|
|
||||||
- ref:channels.discord.token
|
telegram:
|
||||||
- ref:channels.qq.app_secret
|
token: "value"
|
||||||
- ref:channels.dingtalk.client_secret
|
feishu:
|
||||||
- ref:channels.slack.bot_token
|
app_secret: "value"
|
||||||
- ref:channels.slack.app_token
|
encrypt_key: "value"
|
||||||
- ref:channels.matrix.access_token
|
verification_token: "value"
|
||||||
- ref:channels.line.channel_secret
|
discord:
|
||||||
- ref:channels.line.channel_access_token
|
token: "value"
|
||||||
- ref:channels.onebot.access_token
|
weixin:
|
||||||
- ref:channels.wecom.token
|
token: "value"
|
||||||
- ref:channels.wecom.encoding_aes_key
|
qq:
|
||||||
- ref:channels.wecom_app.corp_secret
|
app_secret: "value"
|
||||||
- ref:channels.wecom_app.token
|
dingtalk:
|
||||||
- ref:channels.wecom_app.encoding_aes_key
|
client_secret: "value"
|
||||||
- ref:channels.wecom_aibot.token
|
slack:
|
||||||
- ref:channels.wecom_aibot.encoding_aes_key
|
bot_token: "value"
|
||||||
- ref:channels.pico.token
|
app_token: "value"
|
||||||
- ref:channels.irc.password
|
matrix:
|
||||||
- ref:channels.irc.nickserv_password
|
access_token: "value"
|
||||||
- ref:channels.irc.sasl_password
|
line:
|
||||||
|
channel_secret: "value"
|
||||||
|
channel_access_token: "value"
|
||||||
|
onebot:
|
||||||
|
access_token: "value"
|
||||||
|
wecom:
|
||||||
|
token: "value"
|
||||||
|
encoding_aes_key: "value"
|
||||||
|
wecom_app:
|
||||||
|
corp_secret: "value"
|
||||||
|
token: "value"
|
||||||
|
encoding_aes_key: "value"
|
||||||
|
wecom_aibot:
|
||||||
|
secret: "value"
|
||||||
|
token: "value"
|
||||||
|
encoding_aes_key: "value"
|
||||||
|
pico:
|
||||||
|
token: "value"
|
||||||
|
irc:
|
||||||
|
password: "value"
|
||||||
|
nickserv_password: "value"
|
||||||
|
sasl_password: "value"
|
||||||
|
|
||||||
## Web Tool API Keys
|
## Web Tool API Keys
|
||||||
- ref:web.brave.api_key
|
|
||||||
- ref:web.tavily.api_key
|
|
||||||
- ref:web.perplexity.api_key
|
|
||||||
- ref:web.glm_search.api_key
|
|
||||||
|
|
||||||
**Note:**
|
**Brave, Tavily, Perplexity:**
|
||||||
- Brave, Tavily, Perplexity: Use `api_keys` (array) format in .security.yml
|
```yaml
|
||||||
- GLMSearch: Use `api_key` (single string) format in .security.yml
|
web:
|
||||||
|
|
||||||
|
brave:
|
||||||
|
api_keys:
|
||||||
|
- "BSA-key-1"
|
||||||
|
- "BSA-key-2"
|
||||||
|
tavily:
|
||||||
|
api_keys:
|
||||||
|
- "tvly-key"
|
||||||
|
perplexity:
|
||||||
|
api_keys:
|
||||||
|
- "pplx-key"
|
||||||
|
|
||||||
|
```
|
||||||
|
Use `api_keys` (plural) array format.
|
||||||
|
|
||||||
|
**GLMSearch, BaiduSearch:**
|
||||||
|
```yaml
|
||||||
|
web:
|
||||||
|
|
||||||
|
glm_search:
|
||||||
|
api_key: "your-glm-key"
|
||||||
|
baidu_search:
|
||||||
|
api_key: "your-baidu-key"
|
||||||
|
|
||||||
|
```
|
||||||
|
Use `api_key` (singular) single string format.
|
||||||
|
|
||||||
## Skills Registry Tokens
|
## Skills Registry Tokens
|
||||||
- ref:skills.github.token
|
|
||||||
- ref:skills.clawhub.auth_token
|
```yaml
|
||||||
|
skills:
|
||||||
|
|
||||||
|
github:
|
||||||
|
token: "value"
|
||||||
|
clawhub:
|
||||||
|
auth_token: "value"
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
# Backward Compatibility
|
# Backward Compatibility
|
||||||
|
|
||||||
|
|
@ -191,14 +293,14 @@ You can still use direct values in config.json if needed:
|
||||||
"model_name": "local-model",
|
"model_name": "local-model",
|
||||||
"model": "ollama/llama3",
|
"model": "ollama/llama3",
|
||||||
"api_base": "http://localhost:11434/v1",
|
"api_base": "http://localhost:11434/v1",
|
||||||
"api_key": "ollama" // Direct value (no reference)
|
"api_key": "ollama" // Direct value (works fine)
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
You can also mix references and direct values:
|
You can also mix security values and direct values:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
|
|
||||||
|
|
@ -206,10 +308,12 @@ You can also mix references and direct values:
|
||||||
"model_list": [
|
"model_list": [
|
||||||
{
|
{
|
||||||
"model_name": "cloud-model",
|
"model_name": "cloud-model",
|
||||||
"api_key": "ref:model_list.cloud-model.api_key" // From .security.yml
|
// api_key loaded from .security.yml
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"model_name": "local-model",
|
"model_name": "local-model",
|
||||||
|
"model": "ollama/llama3",
|
||||||
|
"api_base": "http://localhost:11434/v1",
|
||||||
"api_key": "ollama" // Direct value
|
"api_key": "ollama" // Direct value
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
|
|
@ -217,6 +321,11 @@ You can also mix references and direct values:
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Priority Order:**
|
||||||
|
1. Environment variables (highest priority)
|
||||||
|
2. .security.yml values
|
||||||
|
3. config.json direct values (lowest priority)
|
||||||
|
|
||||||
# Migration from Old Config
|
# Migration from Old Config
|
||||||
|
|
||||||
## Step 1: Backup your config
|
## Step 1: Backup your config
|
||||||
|
|
@ -224,7 +333,7 @@ You can also mix references and direct values:
|
||||||
cp ~/.picoclaw/config.json ~/.picoclaw/config.json.backup
|
cp ~/.picoclaw/config.json ~/.picoclaw/config.json.backup
|
||||||
```
|
```
|
||||||
|
|
||||||
## Step 2: Copy the example security file
|
## Step 2: Create .security.yml
|
||||||
```bash
|
```bash
|
||||||
cp security.example.yml ~/.picoclaw/.security.yml
|
cp security.example.yml ~/.picoclaw/.security.yml
|
||||||
```
|
```
|
||||||
|
|
@ -232,10 +341,19 @@ cp security.example.yml ~/.picoclaw/.security.yml
|
||||||
## Step 3: Fill in your API keys
|
## Step 3: Fill in your API keys
|
||||||
Edit ~/.picoclaw/.security.yml and replace placeholders with your actual keys.
|
Edit ~/.picoclaw/.security.yml and replace placeholders with your actual keys.
|
||||||
|
|
||||||
## Step 4: Update config.json references
|
## Step 4: Simplify config.json (Recommended)
|
||||||
Replace sensitive values in ~/.picoclaw/config.json with ref: references.
|
Remove sensitive fields from ~/.picoclaw/config.json:
|
||||||
|
- `api_key` fields from model_list entries
|
||||||
|
- `token` fields from channels
|
||||||
|
- `api_key` fields from tools.web
|
||||||
|
- `token`/`auth_token` fields from tools.skills
|
||||||
|
|
||||||
## Step 5: Test
|
## Step 5: Set permissions
|
||||||
|
```bash
|
||||||
|
chmod 600 ~/.picoclaw/.security.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
## Step 6: Test
|
||||||
```bash
|
```bash
|
||||||
picoclaw --version
|
picoclaw --version
|
||||||
```
|
```
|
||||||
|
|
@ -249,9 +367,11 @@ rm ~/.picoclaw/config.json.backup
|
||||||
|
|
||||||
## Multiple API Keys (Load Balancing & Failover)
|
## Multiple API Keys (Load Balancing & Failover)
|
||||||
|
|
||||||
You can configure multiple API keys for both models and web tools to enable:
|
You can configure multiple API keys for models and web tools to enable:
|
||||||
- **Load balancing**: Requests are distributed across multiple keys
|
- **Load balancing**: Requests are distributed across multiple keys
|
||||||
- **Failover**: If a key fails, the system automatically switches to another key
|
- **Failover**: If a key fails, the system automatically switches to another key
|
||||||
|
- **Rate limit management**: Distribute usage across multiple keys
|
||||||
|
- **High availability**: Reduce downtime during API provider issues
|
||||||
|
|
||||||
### Example: Model with Multiple Keys
|
### Example: Model with Multiple Keys
|
||||||
|
|
||||||
|
|
@ -275,7 +395,7 @@ model_list:
|
||||||
{
|
{
|
||||||
"model_name": "gpt-5.4",
|
"model_name": "gpt-5.4",
|
||||||
"model": "openai/gpt-5.4",
|
"model": "openai/gpt-5.4",
|
||||||
"api_key": "ref:model_list.gpt-5.4.api_key"
|
"api_base": "https://api.openai.com/v1"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|
@ -307,8 +427,13 @@ web:
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": {
|
"web": {
|
||||||
"brave": {
|
"brave": {
|
||||||
"enabled": true,
|
"enabled": true
|
||||||
"api_key": "ref:web.brave.api_key"
|
},
|
||||||
|
"tavily": {
|
||||||
|
"enabled": true
|
||||||
|
},
|
||||||
|
"glm_search": {
|
||||||
|
"enabled": true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -316,9 +441,9 @@ web:
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Single Key
|
## Single Key Format
|
||||||
|
|
||||||
Use array format with one element:
|
**Models, Brave, Tavily, Perplexity:**
|
||||||
```yaml
|
```yaml
|
||||||
model_list:
|
model_list:
|
||||||
|
|
||||||
|
|
@ -328,36 +453,32 @@ model_list:
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Multiple Keys (Load Balancing & Failover)
|
**GLMSearch, BaiduSearch:**
|
||||||
|
|
||||||
Use array format with multiple elements:
|
|
||||||
```yaml
|
```yaml
|
||||||
model_list:
|
web:
|
||||||
|
|
||||||
gpt-5.4:
|
glm_search:
|
||||||
api_keys:
|
api_key: "your-glm-key" # Single key (not array)
|
||||||
- "sk-proj-key-1"
|
|
||||||
- "sk-proj-key-2"
|
|
||||||
- "sk-proj-key-3"
|
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Important:** All model keys in .security.yml must use the `api_keys` (plural) array format.
|
## Model Name Matching
|
||||||
The single `api_key` (singular) format is NOT supported for models.
|
|
||||||
|
|
||||||
### Model Index Matching
|
|
||||||
|
|
||||||
The system supports intelligent model name matching in .security.yml:
|
The system supports intelligent model name matching in .security.yml:
|
||||||
|
|
||||||
**Example 1: Exact Match**
|
### Example 1: Exact Match
|
||||||
```yaml
|
|
||||||
# config.json
|
**config.json:**
|
||||||
|
```json
|
||||||
|
|
||||||
{
|
{
|
||||||
"model_name": "gpt-5.4:0"
|
"model_name": "gpt-5.4:0"
|
||||||
}
|
}
|
||||||
|
|
||||||
# .security.yml (exact match with index)
|
```
|
||||||
|
|
||||||
|
**.security.yml (exact match with index):**
|
||||||
|
```yaml
|
||||||
model_list:
|
model_list:
|
||||||
|
|
||||||
gpt-5.4:0:
|
gpt-5.4:0:
|
||||||
|
|
@ -365,26 +486,30 @@ model_list:
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Example 2: Base Name Match**
|
### Example 2: Base Name Match
|
||||||
```yaml
|
|
||||||
# config.json
|
**config.json:**
|
||||||
|
```json
|
||||||
|
|
||||||
{
|
{
|
||||||
"model_name": "gpt-5.4:0"
|
"model_name": "gpt-5.4:0"
|
||||||
}
|
}
|
||||||
|
|
||||||
# .security.yml (base name without index)
|
```
|
||||||
|
|
||||||
|
**.security.yml (base name without index):**
|
||||||
|
```yaml
|
||||||
model_list:
|
model_list:
|
||||||
|
|
||||||
gpt-5.4:
|
gpt-5.4:
|
||||||
api_keys: ["key-1"]
|
api_keys: ["key-1", "key-2"]
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
Both methods work. The base name match allows you to use simpler keys in .security.yml
|
Both methods work. The base name match allows you to use simpler keys in .security.yml
|
||||||
even when your config uses indexed model names for load balancing.
|
even when your config uses indexed model names for load balancing.
|
||||||
|
|
||||||
### Security File Permissions
|
## Security File Permissions
|
||||||
|
|
||||||
The security file should have restricted permissions:
|
The security file should have restricted permissions:
|
||||||
|
|
||||||
|
|
@ -397,26 +522,64 @@ This ensures only the owner can read and write the file.
|
||||||
# Security Best Practices
|
# Security Best Practices
|
||||||
|
|
||||||
1. Never commit .security.yml to version control
|
1. Never commit .security.yml to version control
|
||||||
2. Set file permissions: chmod 600 ~/.picoclaw/.security.yml
|
2. Add .security.yml to your .gitignore file
|
||||||
3. Use different keys for different environments
|
3. Set file permissions: chmod 600 ~/.picoclaw/.security.yml
|
||||||
4. Rotate keys regularly and update .security.yml
|
4. Use different keys for different environments (dev, staging, production)
|
||||||
5. Encrypt backups containing .security.yml
|
5. Rotate keys regularly and update .security.yml
|
||||||
|
6. Encrypt backups containing .security.yml
|
||||||
|
7. Review access regularly
|
||||||
|
|
||||||
|
# Environment Variables
|
||||||
|
|
||||||
|
You can override any security value using environment variables:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Channels
|
||||||
|
export PICOCLAW_CHANNELS_TELEGRAM_TOKEN="token-from-env"
|
||||||
|
export PICOCLAW_CHANNELS_DISCORD_TOKEN="discord-token-from-env"
|
||||||
|
|
||||||
|
# Web Tools
|
||||||
|
export PICOCLAW_TOOLS_WEB_BRAVE_API_KEY="brave-key-from-env"
|
||||||
|
export PICOCLAW_TOOLS_WEB_BAIDU_API_KEY="baidu-key-from-env"
|
||||||
|
|
||||||
|
# Skills
|
||||||
|
export PICOCLAW_TOOLS_SKILLS_GITHUB_TOKEN="github-token-from-env"
|
||||||
|
```
|
||||||
|
|
||||||
|
Environment variables have the highest priority and will override both config.json
|
||||||
|
and .security.yml values.
|
||||||
|
|
||||||
# Troubleshooting
|
# Troubleshooting
|
||||||
|
|
||||||
|
## Error: "failed to load security config"
|
||||||
|
- Ensure .security.yml exists in the same directory as config.json
|
||||||
|
- Check YAML syntax is valid (use a YAML validator)
|
||||||
|
- Verify file permissions allow reading
|
||||||
|
|
||||||
## Error: "model security entry not found"
|
## Error: "model security entry not found"
|
||||||
- Check that the model name in config.json matches exactly in .security.yml
|
- Check that the model name in config.json matches exactly in .security.yml
|
||||||
- Verify the model_list section exists in .security.yml
|
- Verify the model_list section exists in .security.yml
|
||||||
|
- For indexed names (e.g., "gpt-5.4:0"), check both exact match and base name match
|
||||||
|
- Ensure the YAML structure is correct (proper indentation)
|
||||||
|
|
||||||
## Error: "failed to load security config"
|
## Multiple API Keys Not Working
|
||||||
- Ensure .security.yml exists in the same directory as config.json
|
- Ensure you're using `api_keys` (plural) in .security.yml for models and web tools (except GLMSearch/BaiduSearch)
|
||||||
- Check YAML syntax is valid
|
- Check that the array format is correct in YAML (proper indentation with dashes)
|
||||||
- Verify file permissions allow reading
|
- Remember: Models, Brave, Tavily, Perplexity MUST use `api_keys` (array format)
|
||||||
|
- GLMSearch and BaiduSearch MUST use `api_key` (single string format)
|
||||||
|
|
||||||
## Error: "unknown reference path"
|
## Keys Not Being Applied
|
||||||
- Verify the reference format is correct
|
- Check that .security.yml is in the same directory as config.json
|
||||||
- Check the path structure matches the examples above
|
- Verify the file permissions allow reading (chmod 600 ~/.picoclaw/.security.yml)
|
||||||
- Ensure all required sections exist in .security.yml
|
- Ensure the YAML structure matches the expected format
|
||||||
|
- Check for typos in field names (case-sensitive)
|
||||||
|
- Verify the model/channel names match exactly (case-sensitive)
|
||||||
|
|
||||||
|
## Load Balancing/Failover Issues
|
||||||
|
- Verify all API keys in the api_keys array are valid
|
||||||
|
- Check that all keys have the same rate limits and permissions
|
||||||
|
- Monitor logs to see which keys are being used and failing
|
||||||
|
- Ensure the api_keys array is properly formatted in YAML
|
||||||
*/
|
*/
|
||||||
package config
|
package config
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -232,6 +232,78 @@ func TestExpandMultiKeyModels_PreservesOtherFields(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestExpandMultiKeyModels_IsVirtualFlag(t *testing.T) {
|
||||||
|
models := []*ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "gpt-4",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
apiKeys: []string{"key1", "key2", "key3"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := expandMultiKeyModels(models)
|
||||||
|
|
||||||
|
// Should expand to 3 models
|
||||||
|
if len(result) != 3 {
|
||||||
|
t.Fatalf("expected 3 models, got %d", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Primary model should NOT be virtual
|
||||||
|
primary := result[2]
|
||||||
|
if primary.isVirtual {
|
||||||
|
t.Errorf("primary model should not be virtual")
|
||||||
|
}
|
||||||
|
if primary.ModelName != "gpt-4" {
|
||||||
|
t.Errorf("expected primary model_name 'gpt-4', got %q", primary.ModelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Virtual models should have isVirtual = true
|
||||||
|
virtual1 := result[0]
|
||||||
|
if !virtual1.isVirtual {
|
||||||
|
t.Errorf("gpt-4__key_1 should be virtual")
|
||||||
|
}
|
||||||
|
if virtual1.ModelName != "gpt-4__key_1" {
|
||||||
|
t.Errorf("expected virtual model_name 'gpt-4__key_1', got %q", virtual1.ModelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
virtual2 := result[1]
|
||||||
|
if !virtual2.isVirtual {
|
||||||
|
t.Errorf("gpt-4__key_2 should be virtual")
|
||||||
|
}
|
||||||
|
if virtual2.ModelName != "gpt-4__key_2" {
|
||||||
|
t.Errorf("expected virtual model_name 'gpt-4__key_2', got %q", virtual2.ModelName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsVirtual() method should work
|
||||||
|
if !virtual1.IsVirtual() {
|
||||||
|
t.Errorf("IsVirtual() should return true for virtual model")
|
||||||
|
}
|
||||||
|
if primary.IsVirtual() {
|
||||||
|
t.Errorf("IsVirtual() should return false for primary model")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExpandMultiKeyModels_SingleKey_NotVirtual(t *testing.T) {
|
||||||
|
models := []*ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "gpt-4",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
apiKeys: []string{"single-key"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := expandMultiKeyModels(models)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("expected 1 model, got %d", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single key model should NOT be virtual
|
||||||
|
if result[0].isVirtual {
|
||||||
|
t.Errorf("single key model should not be virtual")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestMergeAPIKeys(t *testing.T) {
|
func TestMergeAPIKeys(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|
|
||||||
|
|
@ -69,21 +69,19 @@ type ModelSecurityEntry struct {
|
||||||
|
|
||||||
// ChannelsSecurity stores channel-related security data
|
// ChannelsSecurity stores channel-related security data
|
||||||
type ChannelsSecurity struct {
|
type ChannelsSecurity struct {
|
||||||
Telegram *TelegramSecurity `yaml:"telegram,omitempty"`
|
Telegram *TelegramSecurity `yaml:"telegram,omitempty"`
|
||||||
Feishu *FeishuSecurity `yaml:"feishu,omitempty"`
|
Feishu *FeishuSecurity `yaml:"feishu,omitempty"`
|
||||||
Discord *DiscordSecurity `yaml:"discord,omitempty"`
|
Discord *DiscordSecurity `yaml:"discord,omitempty"`
|
||||||
Weixin *WeixinSecurity `yaml:"weixin,omitempty"`
|
Weixin *WeixinSecurity `yaml:"weixin,omitempty"`
|
||||||
QQ *QQSecurity `yaml:"qq,omitempty"`
|
QQ *QQSecurity `yaml:"qq,omitempty"`
|
||||||
DingTalk *DingTalkSecurity `yaml:"dingtalk,omitempty"`
|
DingTalk *DingTalkSecurity `yaml:"dingtalk,omitempty"`
|
||||||
Slack *SlackSecurity `yaml:"slack,omitempty"`
|
Slack *SlackSecurity `yaml:"slack,omitempty"`
|
||||||
Matrix *MatrixSecurity `yaml:"matrix,omitempty"`
|
Matrix *MatrixSecurity `yaml:"matrix,omitempty"`
|
||||||
LINE *LINESecurity `yaml:"line,omitempty"`
|
LINE *LINESecurity `yaml:"line,omitempty"`
|
||||||
OneBot *OneBotSecurity `yaml:"onebot,omitempty"`
|
OneBot *OneBotSecurity `yaml:"onebot,omitempty"`
|
||||||
WeCom *WeComSecurity `yaml:"wecom,omitempty"`
|
WeCom *WeComSecurity `yaml:"wecom,omitempty"`
|
||||||
WeComApp *WeComAppSecurity `yaml:"wecom_app,omitempty"`
|
Pico *PicoSecurity `yaml:"pico,omitempty"`
|
||||||
WeComAIBot *WeComAIBotSecurity `yaml:"wecom_aibot,omitempty"`
|
IRC *IRCSecurity `yaml:"irc,omitempty"`
|
||||||
Pico *PicoSecurity `yaml:"pico,omitempty"`
|
|
||||||
IRC *IRCSecurity `yaml:"irc,omitempty"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type TelegramSecurity struct {
|
type TelegramSecurity struct {
|
||||||
|
|
@ -131,20 +129,7 @@ type OneBotSecurity struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type WeComSecurity struct {
|
type WeComSecurity struct {
|
||||||
Token string `yaml:"token,omitempty" env:"PICOCLAW_CHANNELS_WECOM_TOKEN"`
|
Secret string `yaml:"secret,omitempty" env:"PICOCLAW_CHANNELS_WECOM_SECRET"`
|
||||||
EncodingAESKey string `yaml:"encoding_aes_key,omitempty" env:"PICOCLAW_CHANNELS_WECOM_ENCODING_AES_KEY"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type WeComAppSecurity struct {
|
|
||||||
CorpSecret string `yaml:"corp_secret,omitempty" env:"PICOCLAW_CHANNELS_WECOM_APP_CORP_SECRET"`
|
|
||||||
Token string `yaml:"token,omitempty" env:"PICOCLAW_CHANNELS_WECOM_APP_TOKEN"`
|
|
||||||
EncodingAESKey string `yaml:"encoding_aes_key,omitempty" env:"PICOCLAW_CHANNELS_WECOM_APP_ENCODING_AES_KEY"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type WeComAIBotSecurity struct {
|
|
||||||
Secret string `yaml:"secret,omitempty" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_SECRET"`
|
|
||||||
Token string `yaml:"token,omitempty" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_TOKEN"`
|
|
||||||
EncodingAESKey string `yaml:"encoding_aes_key,omitempty" env:"PICOCLAW_CHANNELS_WECOM_AIBOT_ENCODING_AES_KEY"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type PicoSecurity struct {
|
type PicoSecurity struct {
|
||||||
|
|
@ -334,17 +319,9 @@ func mergeChannelsSecurity(dst, src *ChannelsSecurity) {
|
||||||
if src.OneBot != nil && src.OneBot.AccessToken != "" {
|
if src.OneBot != nil && src.OneBot.AccessToken != "" {
|
||||||
dst.OneBot = src.OneBot
|
dst.OneBot = src.OneBot
|
||||||
}
|
}
|
||||||
if src.WeCom != nil && (src.WeCom.Token != "" || src.WeCom.EncodingAESKey != "") {
|
if src.WeCom != nil && src.WeCom.Secret != "" {
|
||||||
dst.WeCom = src.WeCom
|
dst.WeCom = src.WeCom
|
||||||
}
|
}
|
||||||
if src.WeComApp != nil &&
|
|
||||||
(src.WeComApp.CorpSecret != "" || src.WeComApp.Token != "" || src.WeComApp.EncodingAESKey != "") {
|
|
||||||
dst.WeComApp = src.WeComApp
|
|
||||||
}
|
|
||||||
if src.WeComAIBot != nil &&
|
|
||||||
(src.WeComAIBot.Secret != "" || src.WeComAIBot.Token != "" || src.WeComAIBot.EncodingAESKey != "") {
|
|
||||||
dst.WeComAIBot = src.WeComAIBot
|
|
||||||
}
|
|
||||||
if src.Pico != nil && src.Pico.Token != "" {
|
if src.Pico != nil && src.Pico.Token != "" {
|
||||||
dst.Pico = src.Pico
|
dst.Pico = src.Pico
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -242,15 +242,7 @@ func TestAllSecurityKeysAccessible(t *testing.T) {
|
||||||
},
|
},
|
||||||
"wecom": {
|
"wecom": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook"
|
"bot_id": "test_wecom_bot_id"
|
||||||
},
|
|
||||||
"wecom_app": {
|
|
||||||
"enabled": true,
|
|
||||||
"corp_id": "test_corp_id",
|
|
||||||
"agent_id": 123456
|
|
||||||
},
|
|
||||||
"wecom_aibot": {
|
|
||||||
"enabled": true
|
|
||||||
},
|
},
|
||||||
"pico": {
|
"pico": {
|
||||||
"enabled": true
|
"enabled": true
|
||||||
|
|
@ -317,15 +309,7 @@ channels:
|
||||||
onebot:
|
onebot:
|
||||||
access_token: "onebot_test_access_token"
|
access_token: "onebot_test_access_token"
|
||||||
wecom:
|
wecom:
|
||||||
token: "wecom_test_webhook_token"
|
secret: "wecom_test_secret"
|
||||||
encoding_aes_key: "wecom_test_aes_key"
|
|
||||||
wecom_app:
|
|
||||||
corp_secret: "wecom_app_test_corp_secret"
|
|
||||||
token: "wecom_app_test_token"
|
|
||||||
encoding_aes_key: "wecom_app_test_aes_key"
|
|
||||||
wecom_aibot:
|
|
||||||
token: "wecom_aibot_test_token"
|
|
||||||
encoding_aes_key: "wecom_aibot_test_aes_key"
|
|
||||||
pico:
|
pico:
|
||||||
token: "pico_test_token"
|
token: "pico_test_token"
|
||||||
irc:
|
irc:
|
||||||
|
|
@ -411,24 +395,10 @@ skills:
|
||||||
t.Logf("OneBot AccessToken(): %s", cfg.Channels.OneBot.AccessToken())
|
t.Logf("OneBot AccessToken(): %s", cfg.Channels.OneBot.AccessToken())
|
||||||
|
|
||||||
// WeCom
|
// WeCom
|
||||||
assert.Equal(t, "wecom_test_webhook_token", cfg.Channels.WeCom.Token())
|
assert.Equal(t, "test_wecom_bot_id", cfg.Channels.WeCom.BotID)
|
||||||
assert.Equal(t, "wecom_test_aes_key", cfg.Channels.WeCom.EncodingAESKey())
|
assert.Equal(t, "wecom_test_secret", cfg.Channels.WeCom.Secret())
|
||||||
t.Logf("WeCom Token(): %s", cfg.Channels.WeCom.Token())
|
t.Logf("WeCom BotID: %s", cfg.Channels.WeCom.BotID)
|
||||||
t.Logf("WeCom EncodingAESKey(): %s", cfg.Channels.WeCom.EncodingAESKey())
|
t.Logf("WeCom Secret(): %s", cfg.Channels.WeCom.Secret())
|
||||||
|
|
||||||
// WeCom App
|
|
||||||
assert.Equal(t, "wecom_app_test_corp_secret", cfg.Channels.WeComApp.CorpSecret())
|
|
||||||
assert.Equal(t, "wecom_app_test_token", cfg.Channels.WeComApp.Token())
|
|
||||||
assert.Equal(t, "wecom_app_test_aes_key", cfg.Channels.WeComApp.EncodingAESKey())
|
|
||||||
t.Logf("WeComApp CorpSecret(): %s", cfg.Channels.WeComApp.CorpSecret())
|
|
||||||
t.Logf("WeComApp Token(): %s", cfg.Channels.WeComApp.Token())
|
|
||||||
t.Logf("WeComApp EncodingAESKey(): %s", cfg.Channels.WeComApp.EncodingAESKey())
|
|
||||||
|
|
||||||
// WeCom AI Bot
|
|
||||||
assert.Equal(t, "wecom_aibot_test_token", cfg.Channels.WeComAIBot.Token())
|
|
||||||
assert.Equal(t, "wecom_aibot_test_aes_key", cfg.Channels.WeComAIBot.EncodingAESKey())
|
|
||||||
t.Logf("WeComAIBot Token(): %s", cfg.Channels.WeComAIBot.Token())
|
|
||||||
t.Logf("WeComAIBot EncodingAESKey(): %s", cfg.Channels.WeComAIBot.EncodingAESKey())
|
|
||||||
|
|
||||||
// Pico
|
// Pico
|
||||||
assert.Equal(t, "pico_test_token", cfg.Channels.Pico.Token())
|
assert.Equal(t, "pico_test_token", cfg.Channels.Pico.Token())
|
||||||
|
|
|
||||||
21
pkg/gateway/channel_matrix.go
Normal file
21
pkg/gateway/channel_matrix.go
Normal file
|
|
@ -0,0 +1,21 @@
|
||||||
|
//go:build !mipsle && !netbsd
|
||||||
|
|
||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
// Matrix currently pulls in mautrix crypto and modernc sqlite transitively.
|
||||||
|
//
|
||||||
|
// We exclude it on:
|
||||||
|
// - linux/mipsle: mautrix crypto falls back to libolm when the `goolm` build
|
||||||
|
// tag is unavailable, and modernc.org/sqlite/modernc.org/libc also lacks a
|
||||||
|
// working build path for our mipsle + softfloat target.
|
||||||
|
// - netbsd/*: modernc.org/sqlite v1.46.1 fails to compile due to broken
|
||||||
|
// generated mutex code on NetBSD (for example sqlite_netbsd_amd64.go calls
|
||||||
|
// mu.enter/mu.leave, but the generated mutex type does not define them).
|
||||||
|
//
|
||||||
|
// This means Matrix is currently unavailable on those targets. The proper
|
||||||
|
// long-term fix is to split Matrix basic support from its E2EE/sqlite-backed
|
||||||
|
// crypto path, or to upgrade/replace the upstream sqlite dependency once the
|
||||||
|
// affected targets are supported.
|
||||||
|
_ "github.com/sipeed/picoclaw/pkg/channels/matrix"
|
||||||
|
)
|
||||||
|
|
@ -20,7 +20,6 @@ import (
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/irc"
|
_ "github.com/sipeed/picoclaw/pkg/channels/irc"
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/line"
|
_ "github.com/sipeed/picoclaw/pkg/channels/line"
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/maixcam"
|
_ "github.com/sipeed/picoclaw/pkg/channels/maixcam"
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/matrix"
|
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/onebot"
|
_ "github.com/sipeed/picoclaw/pkg/channels/onebot"
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/pico"
|
_ "github.com/sipeed/picoclaw/pkg/channels/pico"
|
||||||
_ "github.com/sipeed/picoclaw/pkg/channels/qq"
|
_ "github.com/sipeed/picoclaw/pkg/channels/qq"
|
||||||
|
|
|
||||||
|
|
@ -13,17 +13,16 @@ var migrateableDirs = []string{
|
||||||
}
|
}
|
||||||
|
|
||||||
var supportedChannels = map[string]bool{
|
var supportedChannels = map[string]bool{
|
||||||
"whatsapp": true,
|
"whatsapp": true,
|
||||||
"telegram": true,
|
"telegram": true,
|
||||||
"feishu": true,
|
"feishu": true,
|
||||||
"discord": true,
|
"discord": true,
|
||||||
"maixcam": true,
|
"maixcam": true,
|
||||||
"qq": true,
|
"qq": true,
|
||||||
"dingtalk": true,
|
"dingtalk": true,
|
||||||
"slack": true,
|
"slack": true,
|
||||||
"matrix": true,
|
"matrix": true,
|
||||||
"line": true,
|
"line": true,
|
||||||
"onebot": true,
|
"onebot": true,
|
||||||
"wecom": true,
|
"wecom": true,
|
||||||
"wecom_app": true,
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,9 +5,13 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"hash/fnv"
|
"hash/fnv"
|
||||||
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
)
|
)
|
||||||
|
|
||||||
// MCPManager defines the interface for MCP manager operations
|
// MCPManager defines the interface for MCP manager operations
|
||||||
|
|
@ -25,6 +29,7 @@ type MCPTool struct {
|
||||||
manager MCPManager
|
manager MCPManager
|
||||||
serverName string
|
serverName string
|
||||||
tool *mcp.Tool
|
tool *mcp.Tool
|
||||||
|
mediaStore media.MediaStore
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewMCPTool creates a new MCP tool wrapper
|
// NewMCPTool creates a new MCP tool wrapper
|
||||||
|
|
@ -36,6 +41,10 @@ func NewMCPTool(manager MCPManager, serverName string, tool *mcp.Tool) *MCPTool
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *MCPTool) SetMediaStore(store media.MediaStore) {
|
||||||
|
t.mediaStore = store
|
||||||
|
}
|
||||||
|
|
||||||
// sanitizeIdentifierComponent normalizes a string so it can be safely used
|
// sanitizeIdentifierComponent normalizes a string so it can be safely used
|
||||||
// as part of a tool/function identifier for downstream providers.
|
// as part of a tool/function identifier for downstream providers.
|
||||||
// It:
|
// It:
|
||||||
|
|
@ -218,13 +227,7 @@ func (t *MCPTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
||||||
WithError(fmt.Errorf("MCP tool error: %s", errMsg))
|
WithError(fmt.Errorf("MCP tool error: %s", errMsg))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Extract text content from result
|
return t.normalizeResultContent(ctx, result.Content)
|
||||||
output := extractContentText(result.Content)
|
|
||||||
|
|
||||||
return &ToolResult{
|
|
||||||
ForLLM: output,
|
|
||||||
IsError: false,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// extractContentText extracts text from MCP content array
|
// extractContentText extracts text from MCP content array
|
||||||
|
|
@ -233,14 +236,269 @@ func extractContentText(content []mcp.Content) string {
|
||||||
for _, c := range content {
|
for _, c := range content {
|
||||||
switch v := c.(type) {
|
switch v := c.(type) {
|
||||||
case *mcp.TextContent:
|
case *mcp.TextContent:
|
||||||
parts = append(parts, v.Text)
|
parts = append(parts, sanitizeToolLLMContent(v.Text))
|
||||||
case *mcp.ImageContent:
|
case *mcp.ImageContent:
|
||||||
// For images, just indicate that an image was returned
|
parts = append(parts, fmt.Sprintf("[Image: %s]", normalizedMIMEType(v.MIMEType)))
|
||||||
parts = append(parts, fmt.Sprintf("[Image: %s]", v.MIMEType))
|
case *mcp.AudioContent:
|
||||||
|
parts = append(parts, fmt.Sprintf("[Audio: %s]", normalizedMIMEType(v.MIMEType)))
|
||||||
|
case *mcp.ResourceLink:
|
||||||
|
parts = append(parts, summarizeResourceLink(v))
|
||||||
|
case *mcp.EmbeddedResource:
|
||||||
|
parts = append(parts, summarizeEmbeddedResource(v))
|
||||||
default:
|
default:
|
||||||
// For other content types, use string representation
|
// For other content types, use string representation
|
||||||
parts = append(parts, fmt.Sprintf("[Content: %T]", v))
|
parts = append(parts, fmt.Sprintf("[Content: %T]", v))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return strings.Join(parts, "\n")
|
return sanitizeToolLLMContent(strings.Join(parts, "\n"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *MCPTool) normalizeResultContent(ctx context.Context, content []mcp.Content) *ToolResult {
|
||||||
|
llmParts := make([]string, 0, len(content))
|
||||||
|
mediaRefs := make([]string, 0, len(content))
|
||||||
|
|
||||||
|
for _, c := range content {
|
||||||
|
switch v := c.(type) {
|
||||||
|
case *mcp.TextContent:
|
||||||
|
text := strings.TrimSpace(sanitizeToolLLMContent(v.Text))
|
||||||
|
if text != "" {
|
||||||
|
llmParts = append(llmParts, text)
|
||||||
|
}
|
||||||
|
case *mcp.ImageContent:
|
||||||
|
ref, note := t.storeBinaryContent(
|
||||||
|
ctx,
|
||||||
|
"image",
|
||||||
|
normalizedMIMEType(v.MIMEType),
|
||||||
|
v.Data,
|
||||||
|
v.Annotations,
|
||||||
|
)
|
||||||
|
if ref != "" {
|
||||||
|
mediaRefs = append(mediaRefs, ref)
|
||||||
|
}
|
||||||
|
if note != "" {
|
||||||
|
llmParts = append(llmParts, note)
|
||||||
|
}
|
||||||
|
case *mcp.AudioContent:
|
||||||
|
ref, note := t.storeBinaryContent(
|
||||||
|
ctx,
|
||||||
|
"audio",
|
||||||
|
normalizedMIMEType(v.MIMEType),
|
||||||
|
v.Data,
|
||||||
|
v.Annotations,
|
||||||
|
)
|
||||||
|
if ref != "" {
|
||||||
|
mediaRefs = append(mediaRefs, ref)
|
||||||
|
}
|
||||||
|
if note != "" {
|
||||||
|
llmParts = append(llmParts, note)
|
||||||
|
}
|
||||||
|
case *mcp.ResourceLink:
|
||||||
|
llmParts = append(llmParts, summarizeResourceLink(v))
|
||||||
|
case *mcp.EmbeddedResource:
|
||||||
|
ref, note := t.storeEmbeddedResource(ctx, v)
|
||||||
|
if ref != "" {
|
||||||
|
mediaRefs = append(mediaRefs, ref)
|
||||||
|
}
|
||||||
|
if note != "" {
|
||||||
|
llmParts = append(llmParts, note)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
llmParts = append(llmParts, fmt.Sprintf("[MCP returned unsupported content type %T]", v))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
result := &ToolResult{
|
||||||
|
ForLLM: strings.Join(compactStrings(llmParts), "\n"),
|
||||||
|
Media: mediaRefs,
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *MCPTool) storeEmbeddedResource(ctx context.Context, content *mcp.EmbeddedResource) (string, string) {
|
||||||
|
if content == nil || content.Resource == nil {
|
||||||
|
return "", "[MCP returned an embedded resource without data.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
resource := content.Resource
|
||||||
|
if len(resource.Blob) > 0 {
|
||||||
|
return t.storeBinaryContent(
|
||||||
|
ctx,
|
||||||
|
"resource",
|
||||||
|
normalizedMIMEType(resource.MIMEType),
|
||||||
|
resource.Blob,
|
||||||
|
content.Annotations,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.TrimSpace(resource.Text) != "" {
|
||||||
|
return "", sanitizeToolLLMContent(resource.Text)
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", summarizeEmbeddedResource(content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *MCPTool) storeBinaryContent(
|
||||||
|
ctx context.Context,
|
||||||
|
kind string,
|
||||||
|
mimeType string,
|
||||||
|
data []byte,
|
||||||
|
annotations *mcp.Annotations,
|
||||||
|
) (string, string) {
|
||||||
|
if len(data) == 0 {
|
||||||
|
return "", fmt.Sprintf("[MCP returned %s content (%s) but it was empty.]", kind, mimeType)
|
||||||
|
}
|
||||||
|
if !annotationsAllowUser(annotations) {
|
||||||
|
return "", fmt.Sprintf(
|
||||||
|
"[MCP returned %s content (%s) for non-user audience; omitted from model context.]",
|
||||||
|
kind,
|
||||||
|
mimeType,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if t.mediaStore == nil {
|
||||||
|
return "", fmt.Sprintf(
|
||||||
|
"[MCP returned %s content (%s); omitted from model context because media delivery is unavailable.]",
|
||||||
|
kind,
|
||||||
|
mimeType,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
channel := ToolChannel(ctx)
|
||||||
|
chatID := ToolChatID(ctx)
|
||||||
|
if channel == "" || chatID == "" {
|
||||||
|
return "", fmt.Sprintf(
|
||||||
|
"[MCP returned %s content (%s); omitted from model context because no target chat was available.]",
|
||||||
|
kind,
|
||||||
|
mimeType,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
dir := media.TempDir()
|
||||||
|
if err := os.MkdirAll(dir, 0o700); err != nil {
|
||||||
|
return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
ext := extensionForMIMEType(mimeType)
|
||||||
|
tmpFile, err := os.CreateTemp(dir, "mcp-*"+ext)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
|
||||||
|
}
|
||||||
|
tmpPath := tmpFile.Name()
|
||||||
|
if _, err = tmpFile.Write(data); err != nil {
|
||||||
|
_ = tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
|
||||||
|
}
|
||||||
|
if err = tmpFile.Close(); err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Sprintf("[MCP returned %s content (%s) but it could not be stored.]", kind, mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
scope := fmt.Sprintf(
|
||||||
|
"tool:mcp:%s:%s:%s:%d",
|
||||||
|
sanitizeIdentifierComponent(t.serverName),
|
||||||
|
channel,
|
||||||
|
chatID,
|
||||||
|
time.Now().UnixNano(),
|
||||||
|
)
|
||||||
|
filename := fmt.Sprintf(
|
||||||
|
"%s_%s%s",
|
||||||
|
sanitizeIdentifierComponent(t.serverName),
|
||||||
|
sanitizeIdentifierComponent(t.tool.Name),
|
||||||
|
ext,
|
||||||
|
)
|
||||||
|
|
||||||
|
ref, err := t.mediaStore.Store(tmpPath, media.MediaMeta{
|
||||||
|
Filename: filename,
|
||||||
|
ContentType: mimeType,
|
||||||
|
Source: fmt.Sprintf(
|
||||||
|
"tool:mcp:%s:%s",
|
||||||
|
sanitizeIdentifierComponent(t.serverName),
|
||||||
|
sanitizeIdentifierComponent(t.tool.Name),
|
||||||
|
),
|
||||||
|
}, scope)
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Sprintf(
|
||||||
|
"[MCP returned %s content (%s) but it could not be registered as media.]",
|
||||||
|
kind,
|
||||||
|
mimeType,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
return ref, fmt.Sprintf(
|
||||||
|
"[MCP returned %s content (%s); omitted from model context and stored as a local media artifact.]",
|
||||||
|
kind,
|
||||||
|
mimeType,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func summarizeResourceLink(content *mcp.ResourceLink) string {
|
||||||
|
if content == nil {
|
||||||
|
return "[MCP returned an empty resource link.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
parts := []string{"[MCP returned resource link"}
|
||||||
|
if content.Name != "" {
|
||||||
|
parts = append(parts, fmt.Sprintf("name=%q", content.Name))
|
||||||
|
}
|
||||||
|
if content.URI != "" {
|
||||||
|
parts = append(parts, fmt.Sprintf("uri=%q", content.URI))
|
||||||
|
}
|
||||||
|
if content.MIMEType != "" {
|
||||||
|
parts = append(parts, fmt.Sprintf("mime=%q", content.MIMEType))
|
||||||
|
}
|
||||||
|
if content.Description != "" {
|
||||||
|
desc := strings.TrimSpace(content.Description)
|
||||||
|
if len(desc) > 200 {
|
||||||
|
desc = desc[:200] + "..."
|
||||||
|
}
|
||||||
|
parts = append(parts, fmt.Sprintf("description=%q", desc))
|
||||||
|
}
|
||||||
|
return strings.Join(parts, ", ") + "]"
|
||||||
|
}
|
||||||
|
|
||||||
|
func summarizeEmbeddedResource(content *mcp.EmbeddedResource) string {
|
||||||
|
if content == nil || content.Resource == nil {
|
||||||
|
return "[MCP returned an embedded resource.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
resource := content.Resource
|
||||||
|
if resource.URI != "" {
|
||||||
|
return fmt.Sprintf(
|
||||||
|
"[MCP returned embedded resource %q (%s).]",
|
||||||
|
resource.URI,
|
||||||
|
normalizedMIMEType(resource.MIMEType),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("[MCP returned embedded resource (%s).]", normalizedMIMEType(resource.MIMEType))
|
||||||
|
}
|
||||||
|
|
||||||
|
func annotationsAllowUser(annotations *mcp.Annotations) bool {
|
||||||
|
if annotations == nil || len(annotations.Audience) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, audience := range annotations.Audience {
|
||||||
|
if strings.EqualFold(string(audience), "user") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizedMIMEType(mimeType string) string {
|
||||||
|
if strings.TrimSpace(mimeType) == "" {
|
||||||
|
return "application/octet-stream"
|
||||||
|
}
|
||||||
|
return mimeType
|
||||||
|
}
|
||||||
|
|
||||||
|
func compactStrings(parts []string) []string {
|
||||||
|
compact := make([]string, 0, len(parts))
|
||||||
|
for _, part := range parts {
|
||||||
|
if strings.TrimSpace(part) == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
compact = append(compact, part)
|
||||||
|
}
|
||||||
|
return compact
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -3,10 +3,14 @@ package tools
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
)
|
)
|
||||||
|
|
||||||
// MockMCPManager is a mock implementation of MCPManager interface for testing
|
// MockMCPManager is a mock implementation of MCPManager interface for testing
|
||||||
|
|
@ -490,3 +494,143 @@ func TestMCPTool_Parameters_MapSchema(t *testing.T) {
|
||||||
t.Errorf("Name type should be 'string', got '%v'", nameParam["type"])
|
t.Errorf("Name type should be 'string', got '%v'", nameParam["type"])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMCPTool_Execute_ImageContentStoredAsMedia(t *testing.T) {
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
manager := &MockMCPManager{
|
||||||
|
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
|
||||||
|
return &mcp.CallToolResult{
|
||||||
|
Content: []mcp.Content{
|
||||||
|
&mcp.ImageContent{
|
||||||
|
Data: []byte("fake-image-bytes"),
|
||||||
|
MIMEType: "image/png",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpTool := NewMCPTool(manager, "screenshoto", &mcp.Tool{Name: "take_screenshot"})
|
||||||
|
mcpTool.SetMediaStore(store)
|
||||||
|
|
||||||
|
result := mcpTool.Execute(WithToolContext(context.Background(), "telegram", "chat-42"), nil)
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("expected success, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
if len(result.Media) != 1 {
|
||||||
|
t.Fatalf("expected 1 media ref, got %d", len(result.Media))
|
||||||
|
}
|
||||||
|
if result.ResponseHandled {
|
||||||
|
t.Fatal("expected MCP image artifact not to mark response as handled")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "stored as a local media artifact") {
|
||||||
|
t.Fatalf("expected local media artifact note, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
path, meta, err := store.ResolveWithMeta(result.Media[0])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected stored media ref to resolve: %v", err)
|
||||||
|
}
|
||||||
|
if meta.ContentType != "image/png" {
|
||||||
|
t.Fatalf("expected image/png content type, got %q", meta.ContentType)
|
||||||
|
}
|
||||||
|
if filepath.Ext(path) != ".png" {
|
||||||
|
t.Fatalf("expected png temp file, got %q", path)
|
||||||
|
}
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected stored media file to be readable: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "fake-image-bytes" {
|
||||||
|
t.Fatalf("expected stored media bytes to match input, got %q", string(data))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPTool_Execute_EmbeddedResourceBlobStoredAsMedia(t *testing.T) {
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
manager := &MockMCPManager{
|
||||||
|
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
|
||||||
|
return &mcp.CallToolResult{
|
||||||
|
Content: []mcp.Content{
|
||||||
|
&mcp.EmbeddedResource{
|
||||||
|
Resource: &mcp.ResourceContents{
|
||||||
|
URI: "file:///tmp/report.png",
|
||||||
|
MIMEType: "image/png",
|
||||||
|
Blob: []byte("blob-bytes"),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpTool := NewMCPTool(manager, "grafana", &mcp.Tool{Name: "get_dashboard_image"})
|
||||||
|
mcpTool.SetMediaStore(store)
|
||||||
|
|
||||||
|
result := mcpTool.Execute(WithToolContext(context.Background(), "telegram", "chat-42"), nil)
|
||||||
|
|
||||||
|
if len(result.Media) != 1 {
|
||||||
|
t.Fatalf("expected embedded resource blob to be stored as media, got %d refs", len(result.Media))
|
||||||
|
}
|
||||||
|
path, _, err := store.ResolveWithMeta(result.Media[0])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected stored media ref to resolve: %v", err)
|
||||||
|
}
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected stored media file to be readable: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "blob-bytes" {
|
||||||
|
t.Fatalf("expected stored blob bytes to match input, got %q", string(data))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPTool_Execute_RespectsUserAudienceForBinaryContent(t *testing.T) {
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
manager := &MockMCPManager{
|
||||||
|
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
|
||||||
|
return &mcp.CallToolResult{
|
||||||
|
Content: []mcp.Content{
|
||||||
|
&mcp.ImageContent{
|
||||||
|
Data: []byte("assistant-only"),
|
||||||
|
MIMEType: "image/png",
|
||||||
|
Annotations: &mcp.Annotations{Audience: []mcp.Role{"assistant"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpTool := NewMCPTool(manager, "screenshoto", &mcp.Tool{Name: "take_screenshot"})
|
||||||
|
mcpTool.SetMediaStore(store)
|
||||||
|
|
||||||
|
result := mcpTool.Execute(WithToolContext(context.Background(), "telegram", "chat-42"), nil)
|
||||||
|
|
||||||
|
if len(result.Media) != 0 {
|
||||||
|
t.Fatalf("expected no media ref for non-user audience, got %d", len(result.Media))
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "non-user audience") {
|
||||||
|
t.Fatalf("expected audience note, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPTool_Execute_LargeBase64TextIsOmittedFromContext(t *testing.T) {
|
||||||
|
manager := &MockMCPManager{
|
||||||
|
callToolFunc: func(ctx context.Context, serverName, toolName string, arguments map[string]any) (*mcp.CallToolResult, error) {
|
||||||
|
return &mcp.CallToolResult{
|
||||||
|
Content: []mcp.Content{
|
||||||
|
&mcp.TextContent{Text: strings.Repeat("QUJD", 400)},
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpTool := NewMCPTool(manager, "test_server", &mcp.Tool{Name: "dump_payload"})
|
||||||
|
|
||||||
|
result := mcpTool.Execute(context.Background(), nil)
|
||||||
|
|
||||||
|
if result.ForLLM != largeBase64OmittedMessage {
|
||||||
|
t.Fatalf("expected sanitized large base64 note, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
292
pkg/tools/normalization.go
Normal file
292
pkg/tools/normalization.go
Normal file
|
|
@ -0,0 +1,292 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/base64"
|
||||||
|
"fmt"
|
||||||
|
"mime"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
"unicode"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
largeBase64OmittedMessage = "[Tool returned a large base64-like payload; omitted from model context.]"
|
||||||
|
inlineMediaOmittedMessage = "[Tool returned inline media content; omitted from model context.]"
|
||||||
|
inlineMediaStoredMessage = "[Tool returned inline media content (%s); omitted from model context and registered as a media attachment.]"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
inlineMarkdownDataURLRe = regexp.MustCompile(`!\[[^\]]*\]\((data:[^)]+)\)`)
|
||||||
|
inlineRawDataURLRe = regexp.MustCompile(`data:[^;\s]+;base64,[A-Za-z0-9+/=\r\n]+`)
|
||||||
|
)
|
||||||
|
|
||||||
|
func normalizeToolResult(
|
||||||
|
result *ToolResult,
|
||||||
|
toolName string,
|
||||||
|
store media.MediaStore,
|
||||||
|
channel string,
|
||||||
|
chatID string,
|
||||||
|
) *ToolResult {
|
||||||
|
if result == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
notes := make([]string, 0, 2)
|
||||||
|
seen := make(map[string]struct{})
|
||||||
|
|
||||||
|
if store != nil && channel != "" && chatID != "" {
|
||||||
|
var refs []string
|
||||||
|
var extractedNotes []string
|
||||||
|
|
||||||
|
result.ForLLM, refs, extractedNotes = extractInlineMediaRefs(
|
||||||
|
result.ForLLM,
|
||||||
|
toolName,
|
||||||
|
store,
|
||||||
|
channel,
|
||||||
|
chatID,
|
||||||
|
seen,
|
||||||
|
)
|
||||||
|
result.Media = append(result.Media, refs...)
|
||||||
|
notes = append(notes, extractedNotes...)
|
||||||
|
|
||||||
|
result.ForUser, refs, extractedNotes = extractInlineMediaRefs(
|
||||||
|
result.ForUser,
|
||||||
|
toolName,
|
||||||
|
store,
|
||||||
|
channel,
|
||||||
|
chatID,
|
||||||
|
seen,
|
||||||
|
)
|
||||||
|
result.Media = append(result.Media, refs...)
|
||||||
|
notes = append(notes, extractedNotes...)
|
||||||
|
}
|
||||||
|
|
||||||
|
result.ForLLM = sanitizeToolLLMContent(result.ForLLM)
|
||||||
|
|
||||||
|
if len(result.Media) > 0 && len(notes) > 0 {
|
||||||
|
if strings.TrimSpace(result.ForLLM) == "" {
|
||||||
|
result.ForLLM = strings.Join(notes, "\n")
|
||||||
|
} else {
|
||||||
|
result.ForLLM = strings.TrimSpace(result.ForLLM) + "\n" + strings.Join(notes, "\n")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(result.Media) > 0 && strings.TrimSpace(result.ForLLM) == "" {
|
||||||
|
result.ForLLM = "[Tool returned media content; omitted from model context and registered as a media attachment.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeToolLLMContent(text string) string {
|
||||||
|
trimmed := strings.TrimSpace(text)
|
||||||
|
if trimmed == "" {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
if inlineMarkdownDataURLRe.MatchString(trimmed) || inlineRawDataURLRe.MatchString(trimmed) {
|
||||||
|
cleaned := inlineMarkdownDataURLRe.ReplaceAllString(trimmed, "")
|
||||||
|
cleaned = inlineRawDataURLRe.ReplaceAllString(cleaned, "")
|
||||||
|
cleaned = strings.TrimSpace(cleaned)
|
||||||
|
if cleaned == "" {
|
||||||
|
return inlineMediaOmittedMessage
|
||||||
|
}
|
||||||
|
return cleaned + "\n" + inlineMediaOmittedMessage
|
||||||
|
}
|
||||||
|
if looksLikeLargeBase64Payload(trimmed) {
|
||||||
|
return largeBase64OmittedMessage
|
||||||
|
}
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
|
||||||
|
func looksLikeLargeBase64Payload(text string) bool {
|
||||||
|
trimmed := strings.TrimSpace(text)
|
||||||
|
if len(trimmed) < 1024 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
nonSpace := 0
|
||||||
|
base64Like := 0
|
||||||
|
spaceCount := 0
|
||||||
|
|
||||||
|
for _, r := range trimmed {
|
||||||
|
if unicode.IsSpace(r) {
|
||||||
|
spaceCount++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
nonSpace++
|
||||||
|
if (r >= 'A' && r <= 'Z') ||
|
||||||
|
(r >= 'a' && r <= 'z') ||
|
||||||
|
(r >= '0' && r <= '9') ||
|
||||||
|
r == '+' || r == '/' || r == '=' {
|
||||||
|
base64Like++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if nonSpace == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
ratio := float64(base64Like) / float64(nonSpace)
|
||||||
|
return ratio >= 0.97 && spaceCount <= len(trimmed)/128
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractInlineMediaRefs(
|
||||||
|
text string,
|
||||||
|
toolName string,
|
||||||
|
store media.MediaStore,
|
||||||
|
channel string,
|
||||||
|
chatID string,
|
||||||
|
seen map[string]struct{},
|
||||||
|
) (cleaned string, refs []string, notes []string) {
|
||||||
|
cleaned = text
|
||||||
|
|
||||||
|
matches := inlineMarkdownDataURLRe.FindAllStringSubmatch(cleaned, -1)
|
||||||
|
for _, match := range matches {
|
||||||
|
if len(match) < 2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
dataURL := match[1]
|
||||||
|
ref, note := storeInlineDataURL(toolName, store, channel, chatID, dataURL, seen)
|
||||||
|
if ref != "" {
|
||||||
|
refs = append(refs, ref)
|
||||||
|
}
|
||||||
|
if note != "" {
|
||||||
|
notes = append(notes, note)
|
||||||
|
}
|
||||||
|
cleaned = strings.ReplaceAll(cleaned, match[0], "")
|
||||||
|
}
|
||||||
|
|
||||||
|
rawMatches := inlineRawDataURLRe.FindAllString(cleaned, -1)
|
||||||
|
for _, dataURL := range rawMatches {
|
||||||
|
ref, note := storeInlineDataURL(toolName, store, channel, chatID, dataURL, seen)
|
||||||
|
if ref != "" {
|
||||||
|
refs = append(refs, ref)
|
||||||
|
}
|
||||||
|
if note != "" {
|
||||||
|
notes = append(notes, note)
|
||||||
|
}
|
||||||
|
cleaned = strings.ReplaceAll(cleaned, dataURL, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.TrimSpace(cleaned), refs, notes
|
||||||
|
}
|
||||||
|
|
||||||
|
func storeInlineDataURL(
|
||||||
|
toolName string,
|
||||||
|
store media.MediaStore,
|
||||||
|
channel string,
|
||||||
|
chatID string,
|
||||||
|
dataURL string,
|
||||||
|
seen map[string]struct{},
|
||||||
|
) (ref string, note string) {
|
||||||
|
dataURL = strings.TrimSpace(dataURL)
|
||||||
|
if _, ok := seen[dataURL]; ok {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
seen[dataURL] = struct{}{}
|
||||||
|
|
||||||
|
if !strings.HasPrefix(strings.ToLower(dataURL), "data:") {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
comma := strings.IndexByte(dataURL, ',')
|
||||||
|
if comma <= 5 {
|
||||||
|
return "", "[Tool returned inline media content that could not be parsed.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
metaPart := dataURL[:comma]
|
||||||
|
payload := dataURL[comma+1:]
|
||||||
|
if !strings.Contains(strings.ToLower(metaPart), ";base64") {
|
||||||
|
return "", "[Tool returned inline media content that was not base64-encoded.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
mimeType := strings.TrimSpace(strings.TrimPrefix(metaPart, "data:"))
|
||||||
|
if semi := strings.IndexByte(mimeType, ';'); semi >= 0 {
|
||||||
|
mimeType = mimeType[:semi]
|
||||||
|
}
|
||||||
|
if mimeType == "" {
|
||||||
|
mimeType = "application/octet-stream"
|
||||||
|
}
|
||||||
|
|
||||||
|
payload = strings.NewReplacer("\n", "", "\r", "", "\t", "", " ", "").Replace(payload)
|
||||||
|
decoded, err := base64.StdEncoding.DecodeString(payload)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Sprintf("[Tool returned inline media content (%s) that could not be decoded.]", mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
dir := media.TempDir()
|
||||||
|
if err = os.MkdirAll(dir, 0o700); err != nil {
|
||||||
|
return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
ext := extensionForMIMEType(mimeType)
|
||||||
|
tmpFile, err := os.CreateTemp(dir, "tool-inline-*"+ext)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
|
||||||
|
}
|
||||||
|
tmpPath := tmpFile.Name()
|
||||||
|
if _, err = tmpFile.Write(decoded); err != nil {
|
||||||
|
tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
|
||||||
|
}
|
||||||
|
if err = tmpFile.Close(); err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be stored.]", mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
filename := sanitizeIdentifierComponent(toolName) + ext
|
||||||
|
scope := fmt.Sprintf(
|
||||||
|
"tool:inline:%s:%s:%s:%d",
|
||||||
|
sanitizeIdentifierComponent(toolName),
|
||||||
|
channel,
|
||||||
|
chatID,
|
||||||
|
time.Now().UnixNano(),
|
||||||
|
)
|
||||||
|
|
||||||
|
ref, err = store.Store(tmpPath, media.MediaMeta{
|
||||||
|
Filename: filename,
|
||||||
|
ContentType: mimeType,
|
||||||
|
Source: fmt.Sprintf("tool:inline:%s", sanitizeIdentifierComponent(toolName)),
|
||||||
|
}, scope)
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
return "", fmt.Sprintf("[Tool returned inline media content (%s) but it could not be registered.]", mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
return ref, fmt.Sprintf(inlineMediaStoredMessage, mimeType)
|
||||||
|
}
|
||||||
|
|
||||||
|
func extensionForMIMEType(mimeType string) string {
|
||||||
|
if mimeType == "" {
|
||||||
|
return ".bin"
|
||||||
|
}
|
||||||
|
if exts, err := mime.ExtensionsByType(mimeType); err == nil && len(exts) > 0 {
|
||||||
|
return exts[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
switch strings.ToLower(mimeType) {
|
||||||
|
case "image/jpeg":
|
||||||
|
return ".jpg"
|
||||||
|
case "image/png":
|
||||||
|
return ".png"
|
||||||
|
case "image/gif":
|
||||||
|
return ".gif"
|
||||||
|
case "image/webp":
|
||||||
|
return ".webp"
|
||||||
|
case "audio/wav", "audio/x-wav":
|
||||||
|
return ".wav"
|
||||||
|
case "audio/mpeg":
|
||||||
|
return ".mp3"
|
||||||
|
case "audio/ogg":
|
||||||
|
return ".ogg"
|
||||||
|
case "video/mp4":
|
||||||
|
return ".mp4"
|
||||||
|
default:
|
||||||
|
return filepath.Ext(mimeType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -9,6 +9,7 @@ import (
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -19,9 +20,14 @@ type ToolEntry struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolRegistry struct {
|
type ToolRegistry struct {
|
||||||
tools map[string]*ToolEntry
|
tools map[string]*ToolEntry
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
version atomic.Uint64 // incremented on Register/RegisterHidden for cache invalidation
|
version atomic.Uint64 // incremented on Register/RegisterHidden for cache invalidation
|
||||||
|
mediaStore media.MediaStore
|
||||||
|
}
|
||||||
|
|
||||||
|
type mediaStoreAware interface {
|
||||||
|
SetMediaStore(store media.MediaStore)
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewToolRegistry() *ToolRegistry {
|
func NewToolRegistry() *ToolRegistry {
|
||||||
|
|
@ -43,6 +49,9 @@ func (r *ToolRegistry) Register(tool Tool) {
|
||||||
IsCore: true,
|
IsCore: true,
|
||||||
TTL: 0, // Core tools do not use TTL
|
TTL: 0, // Core tools do not use TTL
|
||||||
}
|
}
|
||||||
|
if aware, ok := tool.(mediaStoreAware); ok && r.mediaStore != nil {
|
||||||
|
aware.SetMediaStore(r.mediaStore)
|
||||||
|
}
|
||||||
r.version.Add(1)
|
r.version.Add(1)
|
||||||
logger.DebugCF("tools", "Registered core tool", map[string]any{"name": name})
|
logger.DebugCF("tools", "Registered core tool", map[string]any{"name": name})
|
||||||
}
|
}
|
||||||
|
|
@ -61,10 +70,27 @@ func (r *ToolRegistry) RegisterHidden(tool Tool) {
|
||||||
IsCore: false,
|
IsCore: false,
|
||||||
TTL: 0,
|
TTL: 0,
|
||||||
}
|
}
|
||||||
|
if aware, ok := tool.(mediaStoreAware); ok && r.mediaStore != nil {
|
||||||
|
aware.SetMediaStore(r.mediaStore)
|
||||||
|
}
|
||||||
r.version.Add(1)
|
r.version.Add(1)
|
||||||
logger.DebugCF("tools", "Registered hidden tool", map[string]any{"name": name})
|
logger.DebugCF("tools", "Registered hidden tool", map[string]any{"name": name})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetMediaStore injects a MediaStore into all registered tools that can
|
||||||
|
// consume it, and remembers it for future registrations.
|
||||||
|
func (r *ToolRegistry) SetMediaStore(store media.MediaStore) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
|
r.mediaStore = store
|
||||||
|
for _, entry := range r.tools {
|
||||||
|
if aware, ok := entry.Tool.(mediaStoreAware); ok {
|
||||||
|
aware.SetMediaStore(store)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// PromoteTools atomically sets the TTL for multiple non-core tools.
|
// PromoteTools atomically sets the TTL for multiple non-core tools.
|
||||||
// This prevents a concurrent TickTTL from decrementing between promotions.
|
// This prevents a concurrent TickTTL from decrementing between promotions.
|
||||||
func (r *ToolRegistry) PromoteTools(names []string, ttl int) {
|
func (r *ToolRegistry) PromoteTools(names []string, ttl int) {
|
||||||
|
|
@ -238,6 +264,8 @@ func (r *ToolRegistry) ExecuteWithContext(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
result = normalizeToolResult(result, name, r.mediaStore, channel, chatID)
|
||||||
|
|
||||||
duration := time.Since(start)
|
duration := time.Since(start)
|
||||||
|
|
||||||
// Log based on result type
|
// Log based on result type
|
||||||
|
|
@ -259,7 +287,7 @@ func (r *ToolRegistry) ExecuteWithContext(
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"tool": name,
|
"tool": name,
|
||||||
"duration_ms": duration.Milliseconds(),
|
"duration_ms": duration.Milliseconds(),
|
||||||
"result_length": len(result.ForLLM),
|
"result_length": len(result.ContentForLLM()),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -354,7 +382,8 @@ func (r *ToolRegistry) Clone() *ToolRegistry {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
defer r.mu.RUnlock()
|
defer r.mu.RUnlock()
|
||||||
clone := &ToolRegistry{
|
clone := &ToolRegistry{
|
||||||
tools: make(map[string]*ToolEntry, len(r.tools)),
|
tools: make(map[string]*ToolEntry, len(r.tools)),
|
||||||
|
mediaStore: r.mediaStore,
|
||||||
}
|
}
|
||||||
for name, entry := range r.tools {
|
for name, entry := range r.tools {
|
||||||
clone.tools[name] = &ToolEntry{
|
clone.tools[name] = &ToolEntry{
|
||||||
|
|
|
||||||
|
|
@ -3,10 +3,13 @@ package tools
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -46,6 +49,15 @@ func (m *mockAsyncRegistryTool) ExecuteAsync(_ context.Context, args map[string]
|
||||||
return m.result
|
return m.result
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type mockMediaStoreAwareTool struct {
|
||||||
|
mockRegistryTool
|
||||||
|
store media.MediaStore
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockMediaStoreAwareTool) SetMediaStore(store media.MediaStore) {
|
||||||
|
m.store = store
|
||||||
|
}
|
||||||
|
|
||||||
// --- helpers ---
|
// --- helpers ---
|
||||||
|
|
||||||
func newMockTool(name, desc string) *mockRegistryTool {
|
func newMockTool(name, desc string) *mockRegistryTool {
|
||||||
|
|
@ -621,3 +633,102 @@ func TestToolRegistry_Execute_PanicDoesNotAffectOtherTools(t *testing.T) {
|
||||||
t.Errorf("expected 'success', got %q", result2.ForLLM)
|
t.Errorf("expected 'success', got %q", result2.ForLLM)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_SetMediaStore_PropagatesToExistingAndNewTools(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
|
||||||
|
existing := &mockMediaStoreAwareTool{
|
||||||
|
mockRegistryTool: *newMockTool("existing", "existing tool"),
|
||||||
|
}
|
||||||
|
r.Register(existing)
|
||||||
|
|
||||||
|
r.SetMediaStore(store)
|
||||||
|
if existing.store != store {
|
||||||
|
t.Fatal("expected existing tool to receive media store")
|
||||||
|
}
|
||||||
|
|
||||||
|
later := &mockMediaStoreAwareTool{
|
||||||
|
mockRegistryTool: *newMockTool("later", "later tool"),
|
||||||
|
}
|
||||||
|
r.Register(later)
|
||||||
|
|
||||||
|
if later.store != store {
|
||||||
|
t.Fatal("expected newly registered tool to inherit media store")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_ExecuteWithContext_SanitizesLargeBase64Payload(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
payload := strings.Repeat("QUJD", 400)
|
||||||
|
r.Register(&mockRegistryTool{
|
||||||
|
name: "base64_tool",
|
||||||
|
desc: "returns huge base64",
|
||||||
|
params: map[string]any{},
|
||||||
|
result: SilentResult(payload),
|
||||||
|
})
|
||||||
|
|
||||||
|
result := r.ExecuteWithContext(context.Background(), "base64_tool", nil, "telegram", "chat-1", nil)
|
||||||
|
|
||||||
|
if result.ForLLM != largeBase64OmittedMessage {
|
||||||
|
t.Fatalf("expected sanitized payload, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_ExecuteWithContext_ExtractsInlineMediaDataURL(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
store := media.NewFileMediaStore()
|
||||||
|
r.SetMediaStore(store)
|
||||||
|
|
||||||
|
payload := ""
|
||||||
|
r.Register(&mockRegistryTool{
|
||||||
|
name: "inline_media_tool",
|
||||||
|
desc: "returns inline data url",
|
||||||
|
params: map[string]any{},
|
||||||
|
result: SilentResult(payload),
|
||||||
|
})
|
||||||
|
|
||||||
|
result := r.ExecuteWithContext(context.Background(), "inline_media_tool", nil, "telegram", "chat-42", nil)
|
||||||
|
|
||||||
|
if len(result.Media) != 1 {
|
||||||
|
t.Fatalf("expected 1 media ref, got %d", len(result.Media))
|
||||||
|
}
|
||||||
|
if strings.Contains(result.ForLLM, "data:image/png;base64") {
|
||||||
|
t.Fatalf("expected inline data URL to be stripped from ForLLM, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "registered as a media attachment") {
|
||||||
|
t.Fatalf("expected delivery note in ForLLM, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
path, err := store.Resolve(result.Media[0])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected stored media ref to resolve: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(path); err != nil {
|
||||||
|
t.Fatalf("expected stored media file to exist: %v", err)
|
||||||
|
}
|
||||||
|
if filepath.Ext(path) != ".png" {
|
||||||
|
t.Fatalf("expected stored inline media to use png extension, got %q", path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_ExecuteWithContext_SanitizesInlineMediaWithoutStore(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
|
||||||
|
payload := "before  after"
|
||||||
|
r.Register(&mockRegistryTool{
|
||||||
|
name: "inline_media_no_store",
|
||||||
|
desc: "returns inline data url without store",
|
||||||
|
params: map[string]any{},
|
||||||
|
result: SilentResult(payload),
|
||||||
|
})
|
||||||
|
|
||||||
|
result := r.ExecuteWithContext(context.Background(), "inline_media_no_store", nil, "telegram", "chat-42", nil)
|
||||||
|
|
||||||
|
if strings.Contains(result.ForLLM, "data:image/png;base64") {
|
||||||
|
t.Fatalf("expected inline data URL to be removed from ForLLM, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, inlineMediaOmittedMessage) {
|
||||||
|
t.Fatalf("expected inline media omission note, got %q", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,10 +2,16 @@ package tools
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
handledToolLLMNote = "The requested output has already been delivered to the user in the current chat. Do not call send_file or any other delivery tool again. If you reply, provide only a brief confirmation."
|
||||||
|
artifactPathsLLMNote = "Use `send_file` with one of these paths to send it to the user, or use file/exec tools to save it inside the workspace if requested."
|
||||||
|
)
|
||||||
|
|
||||||
// ToolResult represents the structured return value from tool execution.
|
// ToolResult represents the structured return value from tool execution.
|
||||||
// It provides clear semantics for different types of results and supports
|
// It provides clear semantics for different types of results and supports
|
||||||
// async operations, user-facing messages, and error handling.
|
// async operations, user-facing messages, and error handling.
|
||||||
|
|
@ -43,6 +49,48 @@ type ToolResult struct {
|
||||||
// Only populated by SubTurn executions; used by evaluator_optimizer
|
// Only populated by SubTurn executions; used by evaluator_optimizer
|
||||||
// to carry stateful worker context across evaluation iterations.
|
// to carry stateful worker context across evaluation iterations.
|
||||||
Messages []providers.Message `json:"-"`
|
Messages []providers.Message `json:"-"`
|
||||||
|
|
||||||
|
// ArtifactTags exposes local artifact paths back to the LLM in a structured
|
||||||
|
// form, e.g. "[file:/tmp/example.png]". This is used when a tool produced a
|
||||||
|
// reusable local artifact but did not deliver it to the user yet.
|
||||||
|
ArtifactTags []string `json:"artifact_tags,omitempty"`
|
||||||
|
|
||||||
|
// ResponseHandled indicates that this tool execution already satisfied the
|
||||||
|
// user's request at the channel/output level, so the agent loop can stop
|
||||||
|
// without a follow-up assistant response.
|
||||||
|
ResponseHandled bool `json:"response_handled,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ContentForLLM returns the normalized textual content to append to the
|
||||||
|
// conversation after a tool call. Errors fall back to Err when ForLLM is empty.
|
||||||
|
func (tr *ToolResult) ContentForLLM() string {
|
||||||
|
if tr == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
content := tr.ForLLM
|
||||||
|
if content == "" && tr.Err != nil {
|
||||||
|
content = tr.Err.Error()
|
||||||
|
}
|
||||||
|
if tr.ResponseHandled {
|
||||||
|
if content == "" {
|
||||||
|
return handledToolLLMNote
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, handledToolLLMNote) {
|
||||||
|
content += "\n" + handledToolLLMNote
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(tr.ArtifactTags) > 0 {
|
||||||
|
artifactNote := "Local artifact paths: " + strings.Join(tr.ArtifactTags, " ") + "\n" + artifactPathsLLMNote
|
||||||
|
if content == "" {
|
||||||
|
content = artifactNote
|
||||||
|
} else if !strings.Contains(content, artifactNote) {
|
||||||
|
content += "\n" + artifactNote
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if content != "" {
|
||||||
|
return content
|
||||||
|
}
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewToolResult creates a basic ToolResult with content for the LLM.
|
// NewToolResult creates a basic ToolResult with content for the LLM.
|
||||||
|
|
@ -167,3 +215,9 @@ func (tr *ToolResult) WithError(err error) *ToolResult {
|
||||||
tr.Err = err
|
tr.Err = err
|
||||||
return tr
|
return tr
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WithResponseHandled marks the tool result as already delivered to the user.
|
||||||
|
func (tr *ToolResult) WithResponseHandled() *ToolResult {
|
||||||
|
tr.ResponseHandled = true
|
||||||
|
return tr
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,7 @@ package tools
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -227,3 +228,41 @@ func TestToolResultJSONStructure(t *testing.T) {
|
||||||
t.Errorf("Expected silent false, got %v", parsed["silent"])
|
t.Errorf("Expected silent false, got %v", parsed["silent"])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToolResultContentForLLM_AppendsHandledDeliveryNote(t *testing.T) {
|
||||||
|
result := MediaResult("Screenshot attached.", []string{"media://example"}).WithResponseHandled()
|
||||||
|
|
||||||
|
content := result.ContentForLLM()
|
||||||
|
if !strings.Contains(content, "Screenshot attached.") {
|
||||||
|
t.Fatalf("expected original content in ContentForLLM, got %q", content)
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, handledToolLLMNote) {
|
||||||
|
t.Fatalf("expected handled delivery note in ContentForLLM, got %q", content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolResultContentForLLM_UsesHandledDeliveryNoteWhenEmpty(t *testing.T) {
|
||||||
|
result := (&ToolResult{}).WithResponseHandled()
|
||||||
|
|
||||||
|
if got := result.ContentForLLM(); got != handledToolLLMNote {
|
||||||
|
t.Fatalf("ContentForLLM() = %q, want %q", got, handledToolLLMNote)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolResultContentForLLM_AppendsArtifactPaths(t *testing.T) {
|
||||||
|
result := &ToolResult{
|
||||||
|
ForLLM: "Artifact created.",
|
||||||
|
ArtifactTags: []string{"[file:/tmp/example.png]"},
|
||||||
|
}
|
||||||
|
|
||||||
|
content := result.ContentForLLM()
|
||||||
|
if !strings.Contains(content, "Artifact created.") {
|
||||||
|
t.Fatalf("expected original content in ContentForLLM, got %q", content)
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, "Local artifact paths: [file:/tmp/example.png]") {
|
||||||
|
t.Fatalf("expected artifact path note in ContentForLLM, got %q", content)
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, artifactPathsLLMNote) {
|
||||||
|
t.Fatalf("expected artifact guidance note in ContentForLLM, got %q", content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -142,7 +142,7 @@ func (t *SendFileTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
return ErrorResult(fmt.Sprintf("failed to register media: %v", err))
|
return ErrorResult(fmt.Sprintf("failed to register media: %v", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
return MediaResult(fmt.Sprintf("File %q sent to user", filename), []string{ref})
|
return MediaResult(fmt.Sprintf("File %q sent to user", filename), []string{ref}).WithResponseHandled()
|
||||||
}
|
}
|
||||||
|
|
||||||
// detectMediaType determines the MIME type of a file.
|
// detectMediaType determines the MIME type of a file.
|
||||||
|
|
|
||||||
|
|
@ -104,6 +104,9 @@ func TestSendFileTool_Success(t *testing.T) {
|
||||||
if result.Media[0][:8] != "media://" {
|
if result.Media[0][:8] != "media://" {
|
||||||
t.Errorf("expected media:// ref, got %q", result.Media[0])
|
t.Errorf("expected media:// ref, got %q", result.Media[0])
|
||||||
}
|
}
|
||||||
|
if !result.ResponseHandled {
|
||||||
|
t.Fatal("expected send_file success to mark response handled")
|
||||||
|
}
|
||||||
|
|
||||||
_, meta, err := store.ResolveWithMeta(result.Media[0])
|
_, meta, err := store.ResolveWithMeta(result.Media[0])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
|
||||||
|
|
@ -159,10 +159,7 @@ func RunToolLoop(
|
||||||
|
|
||||||
// Append results in original order
|
// Append results in original order
|
||||||
for _, r := range results {
|
for _, r := range results {
|
||||||
contentForLLM := r.result.ForLLM
|
contentForLLM := r.result.ContentForLLM()
|
||||||
if contentForLLM == "" && r.result.Err != nil {
|
|
||||||
contentForLLM = r.result.Err.Error()
|
|
||||||
}
|
|
||||||
|
|
||||||
messages = append(messages, providers.Message{
|
messages = append(messages, providers.Message{
|
||||||
Role: "tool",
|
Role: "tool",
|
||||||
|
|
|
||||||
|
|
@ -22,8 +22,6 @@ var channelCatalog = []channelCatalogItem{
|
||||||
{Name: "qq", ConfigKey: "qq"},
|
{Name: "qq", ConfigKey: "qq"},
|
||||||
{Name: "onebot", ConfigKey: "onebot"},
|
{Name: "onebot", ConfigKey: "onebot"},
|
||||||
{Name: "wecom", ConfigKey: "wecom"},
|
{Name: "wecom", ConfigKey: "wecom"},
|
||||||
{Name: "wecom_app", ConfigKey: "wecom_app"},
|
|
||||||
{Name: "wecom_aibot", ConfigKey: "wecom_aibot"},
|
|
||||||
{Name: "whatsapp", ConfigKey: "whatsapp", Variant: "bridge"},
|
{Name: "whatsapp", ConfigKey: "whatsapp", Variant: "bridge"},
|
||||||
{Name: "whatsapp_native", ConfigKey: "whatsapp", Variant: "native"},
|
{Name: "whatsapp_native", ConfigKey: "whatsapp", Variant: "native"},
|
||||||
{Name: "pico", ConfigKey: "pico"},
|
{Name: "pico", ConfigKey: "pico"},
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
|
@ -16,6 +17,7 @@ func (h *Handler) registerConfigRoutes(mux *http.ServeMux) {
|
||||||
mux.HandleFunc("GET /api/config", h.handleGetConfig)
|
mux.HandleFunc("GET /api/config", h.handleGetConfig)
|
||||||
mux.HandleFunc("PUT /api/config", h.handleUpdateConfig)
|
mux.HandleFunc("PUT /api/config", h.handleUpdateConfig)
|
||||||
mux.HandleFunc("PATCH /api/config", h.handlePatchConfig)
|
mux.HandleFunc("PATCH /api/config", h.handlePatchConfig)
|
||||||
|
mux.HandleFunc("POST /api/config/test-command-patterns", h.handleTestCommandPatterns)
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleGetConfig returns the complete system configuration.
|
// handleGetConfig returns the complete system configuration.
|
||||||
|
|
@ -179,6 +181,70 @@ func (h *Handler) handlePatchConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// handleTestCommandPatterns tests a command against whitelist and blacklist patterns.
|
||||||
|
//
|
||||||
|
// POST /api/config/test-command-patterns
|
||||||
|
func (h *Handler) handleTestCommandPatterns(w http.ResponseWriter, r *http.Request) {
|
||||||
|
body, err := io.ReadAll(io.LimitReader(r.Body, 1<<20))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Failed to read request body", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer r.Body.Close()
|
||||||
|
|
||||||
|
var req struct {
|
||||||
|
AllowPatterns []string `json:"allow_patterns"`
|
||||||
|
DenyPatterns []string `json:"deny_patterns"`
|
||||||
|
Command string `json:"command"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &req); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
lower := strings.ToLower(strings.TrimSpace(req.Command))
|
||||||
|
|
||||||
|
type result struct {
|
||||||
|
Allowed bool `json:"allowed"`
|
||||||
|
Blocked bool `json:"blocked"`
|
||||||
|
MatchedWhitelist *string `json:"matched_whitelist,omitempty"`
|
||||||
|
MatchedBlacklist *string `json:"matched_blacklist,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
resp := result{Allowed: false, Blocked: false}
|
||||||
|
|
||||||
|
// Check whitelist first
|
||||||
|
for _, pattern := range req.AllowPatterns {
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
continue // skip invalid patterns
|
||||||
|
}
|
||||||
|
if re.MatchString(lower) {
|
||||||
|
resp.Allowed = true
|
||||||
|
resp.MatchedWhitelist = &pattern
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check blacklist
|
||||||
|
for _, pattern := range req.DenyPatterns {
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if re.MatchString(lower) {
|
||||||
|
resp.Blocked = true
|
||||||
|
resp.MatchedBlacklist = &pattern
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}
|
||||||
|
|
||||||
// validateConfig checks the config for common errors before saving.
|
// validateConfig checks the config for common errors before saving.
|
||||||
// Returns a list of human-readable error strings; empty means valid.
|
// Returns a list of human-readable error strings; empty means valid.
|
||||||
func validateConfig(cfg *config.Config) []string {
|
func validateConfig(cfg *config.Config) []string {
|
||||||
|
|
@ -209,6 +275,15 @@ func validateConfig(cfg *config.Config) []string {
|
||||||
errs = append(errs, "channels.discord.token is required when discord channel is enabled")
|
errs = append(errs, "channels.discord.token is required when discord channel is enabled")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if cfg.Channels.WeCom.Enabled {
|
||||||
|
if cfg.Channels.WeCom.BotID == "" {
|
||||||
|
errs = append(errs, "channels.wecom.bot_id is required when wecom channel is enabled")
|
||||||
|
}
|
||||||
|
if cfg.Channels.WeCom.Secret() == "" {
|
||||||
|
errs = append(errs, "channels.wecom.secret is required when wecom channel is enabled")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if cfg.Tools.Exec.Enabled {
|
if cfg.Tools.Exec.Enabled {
|
||||||
if cfg.Tools.Exec.EnableDenyPatterns {
|
if cfg.Tools.Exec.EnableDenyPatterns {
|
||||||
errs = append(
|
errs = append(
|
||||||
|
|
|
||||||
|
|
@ -282,3 +282,170 @@ func TestHandlePatchConfig_AllowsInvalidDenyRegexPatternsWhenDenyPatternsDisable
|
||||||
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// testCommandPatterns is a helper that sets up a handler and sends a test-command-patterns request.
|
||||||
|
func testCommandPatterns(t *testing.T, configPath string, body string) *httptest.ResponseRecorder {
|
||||||
|
t.Helper()
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/config/test-command-patterns", bytes.NewBufferString(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
return rec
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_MatchesWhitelist(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": ["^echo\\s+hello"],
|
||||||
|
"deny_patterns": ["^rm\\s+-rf"],
|
||||||
|
"command": "echo hello world"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte(`"allowed":true`)) {
|
||||||
|
t.Fatalf("expected allowed=true, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
if bytes.Contains(rec.Body.Bytes(), []byte(`"blocked":true`)) {
|
||||||
|
t.Fatalf("expected blocked=false when whitelist matches, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_MatchesBlacklistNotWhitelist(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": ["^echo\\s+hello"],
|
||||||
|
"deny_patterns": ["^rm\\s+-rf"],
|
||||||
|
"command": "rm -rf /tmp"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte(`"blocked":true`)) {
|
||||||
|
t.Fatalf("expected blocked=true, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
if bytes.Contains(rec.Body.Bytes(), []byte(`"allowed":true`)) {
|
||||||
|
t.Fatalf("expected allowed=false when blacklist matches but not whitelist, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_MatchesNeither(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": ["^echo\\s+hello"],
|
||||||
|
"deny_patterns": ["^rm\\s+-rf"],
|
||||||
|
"command": "ls -la"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if bytes.Contains(rec.Body.Bytes(), []byte(`"allowed":true`)) {
|
||||||
|
t.Fatalf("expected allowed=false, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
if bytes.Contains(rec.Body.Bytes(), []byte(`"blocked":true`)) {
|
||||||
|
t.Fatalf("expected blocked=false, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_CaseInsensitiveWithGoFlag(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": ["(?i)^ECHO"],
|
||||||
|
"deny_patterns": [],
|
||||||
|
"command": "echo hello"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte(`"allowed":true`)) {
|
||||||
|
t.Fatalf("expected allowed=true with Go (?i) flag, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_EmptyPatterns(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": [],
|
||||||
|
"deny_patterns": [],
|
||||||
|
"command": "rm -rf /tmp"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if bytes.Contains(rec.Body.Bytes(), []byte(`"allowed":true`)) {
|
||||||
|
t.Fatalf("expected allowed=false with empty patterns, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
if bytes.Contains(rec.Body.Bytes(), []byte(`"blocked":true`)) {
|
||||||
|
t.Fatalf("expected blocked=false with empty patterns, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_InvalidRegexSkipped(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": ["([[", "^echo"],
|
||||||
|
"deny_patterns": [],
|
||||||
|
"command": "echo hello"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte(`"allowed":true`)) {
|
||||||
|
t.Fatalf("expected allowed=true, invalid pattern skipped and valid one matched, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_ReturnsMatchedPattern(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
rec := testCommandPatterns(t, configPath, `{
|
||||||
|
"allow_patterns": [],
|
||||||
|
"deny_patterns": ["\\$(?i)[a-zA-Z_]*(SECRET|KEY|PASSWORD|TOKEN|AUTH)[a-zA-Z0-9_]*"],
|
||||||
|
"command": "echo $GITHUB_API_KEY"
|
||||||
|
}`)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte(`"blocked":true`)) {
|
||||||
|
t.Fatalf("expected blocked=true, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte(`matched_blacklist`)) {
|
||||||
|
t.Fatalf("expected matched_blacklist field, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleTestCommandPatterns_InvalidJSON(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPost,
|
||||||
|
"/api/config/test-command-patterns",
|
||||||
|
bytes.NewBufferString(`{invalid json}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -42,6 +42,7 @@ type modelResponse struct {
|
||||||
// Meta
|
// Meta
|
||||||
Configured bool `json:"configured"`
|
Configured bool `json:"configured"`
|
||||||
IsDefault bool `json:"is_default"`
|
IsDefault bool `json:"is_default"`
|
||||||
|
IsVirtual bool `json:"is_virtual"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleListModels returns all model_list entries with masked API keys.
|
// handleListModels returns all model_list entries with masked API keys.
|
||||||
|
|
@ -86,6 +87,7 @@ func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
||||||
ExtraBody: m.ExtraBody,
|
ExtraBody: m.ExtraBody,
|
||||||
Configured: configured[i],
|
Configured: configured[i],
|
||||||
IsDefault: m.ModelName == defaultModel,
|
IsDefault: m.ModelName == defaultModel,
|
||||||
|
IsVirtual: m.IsVirtual(),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -202,8 +204,13 @@ func (h *Handler) handleUpdateModel(w http.ResponseWriter, r *http.Request) {
|
||||||
} else {
|
} else {
|
||||||
mc.ModelConfig.SetAPIKey(mc.APIKey)
|
mc.ModelConfig.SetAPIKey(mc.APIKey)
|
||||||
}
|
}
|
||||||
|
// Preserve existing ExtraBody when omitted (nil), but clear it when
|
||||||
|
// the frontend sends an empty object {} to indicate the field should
|
||||||
|
// be removed.
|
||||||
if mc.ExtraBody == nil {
|
if mc.ExtraBody == nil {
|
||||||
mc.ExtraBody = cfg.ModelList[idx].ExtraBody
|
mc.ExtraBody = cfg.ModelList[idx].ExtraBody
|
||||||
|
} else if len(mc.ExtraBody) == 0 {
|
||||||
|
mc.ExtraBody = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg.ModelList[idx] = &mc.ModelConfig
|
cfg.ModelList[idx] = &mc.ModelConfig
|
||||||
|
|
@ -288,11 +295,13 @@ func (h *Handler) handleSetDefaultModel(w http.ResponseWriter, r *http.Request)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify the model_name exists in model_list
|
// Verify the model_name exists in model_list and is not a virtual model
|
||||||
found := false
|
found := false
|
||||||
|
isVirtual := false
|
||||||
for _, m := range cfg.ModelList {
|
for _, m := range cfg.ModelList {
|
||||||
if m.ModelName == req.ModelName {
|
if m.ModelName == req.ModelName {
|
||||||
found = true
|
found = true
|
||||||
|
isVirtual = m.IsVirtual()
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -300,6 +309,10 @@ func (h *Handler) handleSetDefaultModel(w http.ResponseWriter, r *http.Request)
|
||||||
http.Error(w, fmt.Sprintf("Model %q not found in model_list", req.ModelName), http.StatusNotFound)
|
http.Error(w, fmt.Sprintf("Model %q not found in model_list", req.ModelName), http.StatusNotFound)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if isVirtual {
|
||||||
|
http.Error(w, fmt.Sprintf("Cannot set virtual model %q as default", req.ModelName), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
cfg.Agents.Defaults.ModelName = req.ModelName
|
cfg.Agents.Defaults.ModelName = req.ModelName
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -356,6 +356,46 @@ func TestHandleAddModel_PersistsAPIKey(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestHandleSetDefaultModel_RejectsNonexistentModel tests that setting a non-existent
|
||||||
|
// model as default returns 404. This covers the case where virtual models (which are
|
||||||
|
// filtered by SaveConfig) cannot be set as default.
|
||||||
|
func TestHandleSetDefaultModel_RejectsNonexistentModel(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
// First save a valid config with a primary model
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
cfg.ModelList = []*config.ModelConfig{
|
||||||
|
{ModelName: "gpt-4", Model: "openai/gpt-4o"},
|
||||||
|
}
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to set a non-existent model (like a virtual model name) as default
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/models/default", bytes.NewBufferString(`{
|
||||||
|
"model_name": "gpt-4__key_1"
|
||||||
|
}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
// Should return 404 because the virtual model doesn't exist in the persisted config
|
||||||
|
if rec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusNotFound, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !strings.Contains(rec.Body.String(), "not found") {
|
||||||
|
t.Fatalf("error message should mention 'not found', got: %s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestMaskAPIKey(t *testing.T) {
|
func TestMaskAPIKey(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|
|
||||||
|
|
@ -98,7 +98,7 @@ func main() {
|
||||||
defer logger.DisableFileLogging()
|
defer logger.DisableFileLogging()
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.InfoC("web", "PicoClaw Launcher starting...")
|
logger.InfoC("web", fmt.Sprintf("%s Launcher %s starting...", appName, appVersion))
|
||||||
logger.InfoC("web", fmt.Sprintf("PicoClaw Home: %s", picoHome))
|
logger.InfoC("web", fmt.Sprintf("PicoClaw Home: %s", picoHome))
|
||||||
|
|
||||||
// Set language from command line or auto-detect
|
// Set language from command line or auto-detect
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,6 @@
|
||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
_ "embed"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
"fyne.io/systray"
|
"fyne.io/systray"
|
||||||
|
|
@ -93,8 +92,3 @@ func onReady() {
|
||||||
func onExit() {
|
func onExit() {
|
||||||
logger.Info(T(Exiting))
|
logger.Info(T(Exiting))
|
||||||
}
|
}
|
||||||
|
|
||||||
// getIcon returns the system tray icon
|
|
||||||
func getIcon() []byte {
|
|
||||||
return iconData
|
|
||||||
}
|
|
||||||
|
|
|
||||||
12
web/backend/systray_icon_nonwindows.go
Normal file
12
web/backend/systray_icon_nonwindows.go
Normal file
|
|
@ -0,0 +1,12 @@
|
||||||
|
//go:build !windows && ((!darwin && !freebsd) || cgo)
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import _ "embed"
|
||||||
|
|
||||||
|
//go:embed icon.png
|
||||||
|
var iconPNG []byte
|
||||||
|
|
||||||
|
func getIcon() []byte {
|
||||||
|
return iconPNG
|
||||||
|
}
|
||||||
|
|
@ -5,4 +5,8 @@ package main
|
||||||
import _ "embed"
|
import _ "embed"
|
||||||
|
|
||||||
//go:embed icon.ico
|
//go:embed icon.ico
|
||||||
var iconData []byte
|
var iconICO []byte
|
||||||
|
|
||||||
|
func getIcon() []byte {
|
||||||
|
return iconICO
|
||||||
|
}
|
||||||
|
|
@ -13,6 +13,7 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// runTray falls back to a headless mode on platforms where systray requires cgo.
|
||||||
func runTray() {
|
func runTray() {
|
||||||
logger.Infof("System tray is unavailable in %s builds without cgo; running without tray", runtime.GOOS)
|
logger.Infof("System tray is unavailable in %s builds without cgo; running without tray", runtime.GOOS)
|
||||||
|
|
||||||
|
|
@ -1,8 +0,0 @@
|
||||||
//go:build !windows
|
|
||||||
|
|
||||||
package main
|
|
||||||
|
|
||||||
import _ "embed"
|
|
||||||
|
|
||||||
//go:embed icon.png
|
|
||||||
var iconData []byte
|
|
||||||
|
|
@ -21,6 +21,7 @@ export interface ModelInfo {
|
||||||
// Meta
|
// Meta
|
||||||
configured: boolean
|
configured: boolean
|
||||||
is_default: boolean
|
is_default: boolean
|
||||||
|
is_virtual: boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
interface ModelsListResponse {
|
interface ModelsListResponse {
|
||||||
|
|
|
||||||
|
|
@ -145,13 +145,7 @@ function isConfigured(
|
||||||
case "weixin":
|
case "weixin":
|
||||||
return asString(config.account_id) !== ""
|
return asString(config.account_id) !== ""
|
||||||
case "wecom":
|
case "wecom":
|
||||||
return asString(config.token) !== ""
|
return asString(config.bot_id) !== ""
|
||||||
case "wecom_app":
|
|
||||||
return (
|
|
||||||
asString(config.corp_id) !== "" && asString(config.corp_secret) !== ""
|
|
||||||
)
|
|
||||||
case "wecom_aibot":
|
|
||||||
return asString(config.token) !== ""
|
|
||||||
case "whatsapp":
|
case "whatsapp":
|
||||||
return asString(config.bridge_url) !== ""
|
return asString(config.bridge_url) !== ""
|
||||||
case "whatsapp_native":
|
case "whatsapp_native":
|
||||||
|
|
@ -192,11 +186,7 @@ function getRequiredFieldKeys(channelName: string): string[] {
|
||||||
case "onebot":
|
case "onebot":
|
||||||
return ["ws_url"]
|
return ["ws_url"]
|
||||||
case "wecom":
|
case "wecom":
|
||||||
return ["token"]
|
return ["bot_id", "secret"]
|
||||||
case "wecom_app":
|
|
||||||
return ["corp_id", "corp_secret"]
|
|
||||||
case "wecom_aibot":
|
|
||||||
return ["token"]
|
|
||||||
case "whatsapp":
|
case "whatsapp":
|
||||||
return ["bridge_url"]
|
return ["bridge_url"]
|
||||||
case "pico":
|
case "pico":
|
||||||
|
|
|
||||||
|
|
@ -28,6 +28,7 @@ const SECRET_FIELDS = new Set([
|
||||||
"encoding_aes_key",
|
"encoding_aes_key",
|
||||||
"encrypt_key",
|
"encrypt_key",
|
||||||
"verification_token",
|
"verification_token",
|
||||||
|
"secret",
|
||||||
"password",
|
"password",
|
||||||
"nickserv_password",
|
"nickserv_password",
|
||||||
"sasl_password",
|
"sasl_password",
|
||||||
|
|
@ -44,6 +45,7 @@ const OBJECT_FIELDS = new Set([
|
||||||
"allow_token_query",
|
"allow_token_query",
|
||||||
"allow_from",
|
"allow_from",
|
||||||
"allow_origins",
|
"allow_origins",
|
||||||
|
"groups",
|
||||||
])
|
])
|
||||||
|
|
||||||
function formatLabel(key: string): string {
|
function formatLabel(key: string): string {
|
||||||
|
|
@ -118,6 +120,14 @@ export function GenericForm({
|
||||||
app_id: t("channels.form.desc.appId"),
|
app_id: t("channels.form.desc.appId"),
|
||||||
client_id: t("channels.form.desc.clientId"),
|
client_id: t("channels.form.desc.clientId"),
|
||||||
corp_id: t("channels.form.desc.corpId"),
|
corp_id: t("channels.form.desc.corpId"),
|
||||||
|
bot_id: t("channels.form.desc.appId"),
|
||||||
|
websocket_url: t("channels.form.desc.wsUrl"),
|
||||||
|
dm_policy: t("channels.form.desc.genericField", { field: "DM policy" }),
|
||||||
|
group_policy: t("channels.form.desc.genericField", { field: "group policy" }),
|
||||||
|
group_allow_from: t("channels.form.desc.allowFrom"),
|
||||||
|
send_thinking_message: t("channels.form.desc.genericField", {
|
||||||
|
field: "thinking message behavior",
|
||||||
|
}),
|
||||||
agent_id: t("channels.form.desc.agentId"),
|
agent_id: t("channels.form.desc.agentId"),
|
||||||
webhook_url: t("channels.form.desc.webhookUrl"),
|
webhook_url: t("channels.form.desc.webhookUrl"),
|
||||||
webhook_host: t("channels.form.desc.webhookHost"),
|
webhook_host: t("channels.form.desc.webhookHost"),
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,4 @@
|
||||||
|
import { useState } from "react"
|
||||||
import type { ReactNode } from "react"
|
import type { ReactNode } from "react"
|
||||||
import { useTranslation } from "react-i18next"
|
import { useTranslation } from "react-i18next"
|
||||||
|
|
||||||
|
|
@ -7,6 +8,7 @@ import {
|
||||||
type LauncherForm,
|
type LauncherForm,
|
||||||
} from "@/components/config/form-model"
|
} from "@/components/config/form-model"
|
||||||
import { Field, SwitchCardField } from "@/components/shared-form"
|
import { Field, SwitchCardField } from "@/components/shared-form"
|
||||||
|
import { Button } from "@/components/ui/button"
|
||||||
import {
|
import {
|
||||||
Card,
|
Card,
|
||||||
CardContent,
|
CardContent,
|
||||||
|
|
@ -201,6 +203,56 @@ interface ExecSectionProps {
|
||||||
|
|
||||||
export function ExecSection({ form, onFieldChange }: ExecSectionProps) {
|
export function ExecSection({ form, onFieldChange }: ExecSectionProps) {
|
||||||
const { t } = useTranslation()
|
const { t } = useTranslation()
|
||||||
|
const [testCommand, setTestCommand] = useState("")
|
||||||
|
const [testResult, setTestResult] = useState<{
|
||||||
|
allowed: boolean
|
||||||
|
blocked: boolean
|
||||||
|
matchedWhitelist: string | null
|
||||||
|
matchedBlacklist: string | null
|
||||||
|
} | null>(null)
|
||||||
|
const [isLoading, setIsLoading] = useState(false)
|
||||||
|
|
||||||
|
const testPatterns = async () => {
|
||||||
|
if (!testCommand.trim()) {
|
||||||
|
setTestResult(null)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
const allowPatterns = form.customAllowPatternsText
|
||||||
|
.split("\n")
|
||||||
|
.map((p) => p.trim())
|
||||||
|
.filter((p) => p.length > 0)
|
||||||
|
const denyPatterns = form.enableDenyPatterns
|
||||||
|
? form.customDenyPatternsText
|
||||||
|
.split("\n")
|
||||||
|
.map((p) => p.trim())
|
||||||
|
.filter((p) => p.length > 0)
|
||||||
|
: []
|
||||||
|
|
||||||
|
setIsLoading(true)
|
||||||
|
try {
|
||||||
|
const res = await fetch("/api/config/test-command-patterns", {
|
||||||
|
method: "POST",
|
||||||
|
headers: { "Content-Type": "application/json" },
|
||||||
|
body: JSON.stringify({
|
||||||
|
allow_patterns: allowPatterns,
|
||||||
|
deny_patterns: denyPatterns,
|
||||||
|
command: testCommand,
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
const data = await res.json()
|
||||||
|
setTestResult({
|
||||||
|
allowed: data.allowed,
|
||||||
|
blocked: data.blocked,
|
||||||
|
matchedWhitelist: data.matched_whitelist ?? null,
|
||||||
|
matchedBlacklist: data.matched_blacklist ?? null,
|
||||||
|
})
|
||||||
|
} catch {
|
||||||
|
setTestResult(null)
|
||||||
|
} finally {
|
||||||
|
setIsLoading(false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return (
|
return (
|
||||||
<ConfigSectionCard title={t("pages.config.sections.exec")}>
|
<ConfigSectionCard title={t("pages.config.sections.exec")}>
|
||||||
|
|
@ -266,6 +318,50 @@ export function ExecSection({ form, onFieldChange }: ExecSectionProps) {
|
||||||
/>
|
/>
|
||||||
</Field>
|
</Field>
|
||||||
|
|
||||||
|
<Field
|
||||||
|
label={t("pages.config.pattern_detector_title")}
|
||||||
|
hint={t("pages.config.pattern_detector_hint")}
|
||||||
|
layout="setting-row"
|
||||||
|
controlClassName="md:max-w-md"
|
||||||
|
>
|
||||||
|
<div className="flex w-full flex-col gap-2">
|
||||||
|
<div className="flex gap-2">
|
||||||
|
<Input
|
||||||
|
value={testCommand}
|
||||||
|
placeholder={t(
|
||||||
|
"pages.config.pattern_detector_input_placeholder",
|
||||||
|
)}
|
||||||
|
onChange={(e) => setTestCommand(e.target.value)}
|
||||||
|
onKeyDown={(e) => {
|
||||||
|
if (e.key === "Enter") {
|
||||||
|
testPatterns()
|
||||||
|
}
|
||||||
|
}}
|
||||||
|
/>
|
||||||
|
<Button onClick={testPatterns} disabled={isLoading}>
|
||||||
|
{t("pages.config.pattern_detector_test_button")}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
{testResult && (
|
||||||
|
<div
|
||||||
|
className={`rounded-md p-2 text-sm ${
|
||||||
|
testResult.allowed
|
||||||
|
? "bg-green-500/10 text-green-600"
|
||||||
|
: testResult.blocked
|
||||||
|
? "bg-red-500/10 text-red-600"
|
||||||
|
: "bg-muted text-muted-foreground"
|
||||||
|
}`}
|
||||||
|
>
|
||||||
|
{testResult.allowed
|
||||||
|
? `${t("pages.config.pattern_detector_result_allowed")}${testResult.matchedWhitelist ? ` (${testResult.matchedWhitelist})` : ""}`
|
||||||
|
: testResult.blocked
|
||||||
|
? `${t("pages.config.pattern_detector_result_blocked")}${testResult.matchedBlacklist ? ` (${testResult.matchedBlacklist})` : ""}`
|
||||||
|
: t("pages.config.pattern_detector_result_no_match")}
|
||||||
|
</div>
|
||||||
|
)}
|
||||||
|
</div>
|
||||||
|
</Field>
|
||||||
|
|
||||||
<Field
|
<Field
|
||||||
label={t("pages.config.exec_timeout_seconds")}
|
label={t("pages.config.exec_timeout_seconds")}
|
||||||
hint={t("pages.config.exec_timeout_seconds_hint")}
|
hint={t("pages.config.exec_timeout_seconds_hint")}
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import {
|
||||||
} from "@/components/shared-form"
|
} from "@/components/shared-form"
|
||||||
import { Button } from "@/components/ui/button"
|
import { Button } from "@/components/ui/button"
|
||||||
import { Input } from "@/components/ui/input"
|
import { Input } from "@/components/ui/input"
|
||||||
|
import { Textarea } from "@/components/ui/textarea"
|
||||||
import {
|
import {
|
||||||
Sheet,
|
Sheet,
|
||||||
SheetContent,
|
SheetContent,
|
||||||
|
|
@ -34,6 +35,7 @@ interface AddForm {
|
||||||
maxTokensField: string
|
maxTokensField: string
|
||||||
requestTimeout: string
|
requestTimeout: string
|
||||||
thinkingLevel: string
|
thinkingLevel: string
|
||||||
|
extraBody: string
|
||||||
}
|
}
|
||||||
|
|
||||||
const EMPTY_ADD_FORM: AddForm = {
|
const EMPTY_ADD_FORM: AddForm = {
|
||||||
|
|
@ -49,6 +51,7 @@ const EMPTY_ADD_FORM: AddForm = {
|
||||||
maxTokensField: "",
|
maxTokensField: "",
|
||||||
requestTimeout: "",
|
requestTimeout: "",
|
||||||
thinkingLevel: "",
|
thinkingLevel: "",
|
||||||
|
extraBody: "",
|
||||||
}
|
}
|
||||||
|
|
||||||
interface AddModelSheetProps {
|
interface AddModelSheetProps {
|
||||||
|
|
@ -100,7 +103,7 @@ export function AddModelSheet({
|
||||||
}
|
}
|
||||||
|
|
||||||
const setField =
|
const setField =
|
||||||
(key: keyof AddForm) => (e: React.ChangeEvent<HTMLInputElement>) => {
|
(key: keyof AddForm) => (e: React.ChangeEvent<HTMLInputElement | HTMLTextAreaElement>) => {
|
||||||
setForm((f) => ({ ...f, [key]: e.target.value }))
|
setForm((f) => ({ ...f, [key]: e.target.value }))
|
||||||
if (fieldErrors[key]) {
|
if (fieldErrors[key]) {
|
||||||
setFieldErrors((prev) => ({ ...prev, [key]: undefined }))
|
setFieldErrors((prev) => ({ ...prev, [key]: undefined }))
|
||||||
|
|
@ -129,6 +132,9 @@ export function AddModelSheet({
|
||||||
? Number(form.requestTimeout)
|
? Number(form.requestTimeout)
|
||||||
: undefined,
|
: undefined,
|
||||||
thinking_level: form.thinkingLevel.trim() || undefined,
|
thinking_level: form.thinkingLevel.trim() || undefined,
|
||||||
|
extra_body: form.extraBody.trim()
|
||||||
|
? JSON.parse(form.extraBody.trim())
|
||||||
|
: undefined,
|
||||||
})
|
})
|
||||||
if (setAsDefault) {
|
if (setAsDefault) {
|
||||||
await setDefaultModel(modelName)
|
await setDefaultModel(modelName)
|
||||||
|
|
@ -305,6 +311,18 @@ export function AddModelSheet({
|
||||||
placeholder="max_completion_tokens"
|
placeholder="max_completion_tokens"
|
||||||
/>
|
/>
|
||||||
</Field>
|
</Field>
|
||||||
|
|
||||||
|
<Field
|
||||||
|
label={t("models.field.extraBody")}
|
||||||
|
hint={t("models.field.extraBodyHint")}
|
||||||
|
>
|
||||||
|
<Textarea
|
||||||
|
value={form.extraBody}
|
||||||
|
onChange={setField("extraBody")}
|
||||||
|
placeholder='{"key": "value"}'
|
||||||
|
rows={3}
|
||||||
|
/>
|
||||||
|
</Field>
|
||||||
</AdvancedSection>
|
</AdvancedSection>
|
||||||
|
|
||||||
{serverError && (
|
{serverError && (
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import {
|
||||||
} from "@/components/shared-form"
|
} from "@/components/shared-form"
|
||||||
import { Button } from "@/components/ui/button"
|
import { Button } from "@/components/ui/button"
|
||||||
import { Input } from "@/components/ui/input"
|
import { Input } from "@/components/ui/input"
|
||||||
|
import { Textarea } from "@/components/ui/textarea"
|
||||||
import {
|
import {
|
||||||
Sheet,
|
Sheet,
|
||||||
SheetContent,
|
SheetContent,
|
||||||
|
|
@ -32,6 +33,7 @@ interface EditForm {
|
||||||
maxTokensField: string
|
maxTokensField: string
|
||||||
requestTimeout: string
|
requestTimeout: string
|
||||||
thinkingLevel: string
|
thinkingLevel: string
|
||||||
|
extraBody: string
|
||||||
}
|
}
|
||||||
|
|
||||||
interface EditModelSheetProps {
|
interface EditModelSheetProps {
|
||||||
|
|
@ -59,6 +61,7 @@ export function EditModelSheet({
|
||||||
maxTokensField: "",
|
maxTokensField: "",
|
||||||
requestTimeout: "",
|
requestTimeout: "",
|
||||||
thinkingLevel: "",
|
thinkingLevel: "",
|
||||||
|
extraBody: "",
|
||||||
})
|
})
|
||||||
const [saving, setSaving] = useState(false)
|
const [saving, setSaving] = useState(false)
|
||||||
const [setAsDefault, setSetAsDefault] = useState(false)
|
const [setAsDefault, setSetAsDefault] = useState(false)
|
||||||
|
|
@ -79,6 +82,9 @@ export function EditModelSheet({
|
||||||
? String(model.request_timeout)
|
? String(model.request_timeout)
|
||||||
: "",
|
: "",
|
||||||
thinkingLevel: model.thinking_level ?? "",
|
thinkingLevel: model.thinking_level ?? "",
|
||||||
|
extraBody: model.extra_body
|
||||||
|
? JSON.stringify(model.extra_body, null, 2)
|
||||||
|
: "",
|
||||||
})
|
})
|
||||||
setSetAsDefault(model.is_default)
|
setSetAsDefault(model.is_default)
|
||||||
setError("")
|
setError("")
|
||||||
|
|
@ -86,7 +92,7 @@ export function EditModelSheet({
|
||||||
}, [model])
|
}, [model])
|
||||||
|
|
||||||
const setField =
|
const setField =
|
||||||
(key: keyof EditForm) => (e: React.ChangeEvent<HTMLInputElement>) =>
|
(key: keyof EditForm) => (e: React.ChangeEvent<HTMLInputElement | HTMLTextAreaElement>) =>
|
||||||
setForm((f) => ({ ...f, [key]: e.target.value }))
|
setForm((f) => ({ ...f, [key]: e.target.value }))
|
||||||
|
|
||||||
const handleSave = async () => {
|
const handleSave = async () => {
|
||||||
|
|
@ -109,6 +115,9 @@ export function EditModelSheet({
|
||||||
? Number(form.requestTimeout)
|
? Number(form.requestTimeout)
|
||||||
: undefined,
|
: undefined,
|
||||||
thinking_level: form.thinkingLevel || undefined,
|
thinking_level: form.thinkingLevel || undefined,
|
||||||
|
extra_body: form.extraBody.trim()
|
||||||
|
? JSON.parse(form.extraBody.trim())
|
||||||
|
: {},
|
||||||
})
|
})
|
||||||
if (setAsDefault && !model.is_default) {
|
if (setAsDefault && !model.is_default) {
|
||||||
await setDefaultModel(model.model_name)
|
await setDefaultModel(model.model_name)
|
||||||
|
|
@ -273,6 +282,18 @@ export function EditModelSheet({
|
||||||
placeholder="max_completion_tokens"
|
placeholder="max_completion_tokens"
|
||||||
/>
|
/>
|
||||||
</Field>
|
</Field>
|
||||||
|
|
||||||
|
<Field
|
||||||
|
label={t("models.field.extraBody")}
|
||||||
|
hint={t("models.field.extraBodyHint")}
|
||||||
|
>
|
||||||
|
<Textarea
|
||||||
|
value={form.extraBody}
|
||||||
|
onChange={setField("extraBody")}
|
||||||
|
placeholder='{"key": "value"}'
|
||||||
|
rows={3}
|
||||||
|
/>
|
||||||
|
</Field>
|
||||||
</AdvancedSection>
|
</AdvancedSection>
|
||||||
|
|
||||||
{error && (
|
{error && (
|
||||||
|
|
|
||||||
|
|
@ -28,7 +28,7 @@ export function ModelCard({
|
||||||
}: ModelCardProps) {
|
}: ModelCardProps) {
|
||||||
const { t } = useTranslation()
|
const { t } = useTranslation()
|
||||||
const isOAuth = model.auth_method === "oauth"
|
const isOAuth = model.auth_method === "oauth"
|
||||||
const canSetDefault = model.configured && !model.is_default
|
const canSetDefault = model.configured && !model.is_default && !model.is_virtual
|
||||||
|
|
||||||
return (
|
return (
|
||||||
<div
|
<div
|
||||||
|
|
@ -64,6 +64,11 @@ export function ModelCard({
|
||||||
{t("models.badge.default")}
|
{t("models.badge.default")}
|
||||||
</span>
|
</span>
|
||||||
)}
|
)}
|
||||||
|
{model.is_virtual && (
|
||||||
|
<span className="bg-muted text-muted-foreground shrink-0 rounded px-1.5 py-0.5 text-[10px] leading-none font-medium">
|
||||||
|
{t("models.badge.virtual")}
|
||||||
|
</span>
|
||||||
|
)}
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div className="flex shrink-0 items-center gap-0.5">
|
<div className="flex shrink-0 items-center gap-0.5">
|
||||||
|
|
|
||||||
|
|
@ -32,8 +32,6 @@ const CHANNEL_IMPORTANCE_TAIL = [
|
||||||
"slack",
|
"slack",
|
||||||
"line",
|
"line",
|
||||||
"wecom",
|
"wecom",
|
||||||
"wecom_app",
|
|
||||||
"wecom_aibot",
|
|
||||||
"dingtalk",
|
"dingtalk",
|
||||||
"qq",
|
"qq",
|
||||||
"onebot",
|
"onebot",
|
||||||
|
|
@ -78,8 +76,6 @@ const CHANNEL_ICON_MAP: Record<
|
||||||
qq: IconBrandQq,
|
qq: IconBrandQq,
|
||||||
weixin: IconBrandWechat,
|
weixin: IconBrandWechat,
|
||||||
wecom: IconBrandWechat,
|
wecom: IconBrandWechat,
|
||||||
wecom_app: IconBrandWechat,
|
|
||||||
wecom_aibot: IconBrandWechat,
|
|
||||||
whatsapp: IconBrandWhatsapp,
|
whatsapp: IconBrandWhatsapp,
|
||||||
whatsapp_native: IconBrandWhatsapp,
|
whatsapp_native: IconBrandWhatsapp,
|
||||||
matrix: IconBrandMatrix,
|
matrix: IconBrandMatrix,
|
||||||
|
|
|
||||||
|
|
@ -154,7 +154,8 @@
|
||||||
"unconfigured": "Not configured"
|
"unconfigured": "Not configured"
|
||||||
},
|
},
|
||||||
"badge": {
|
"badge": {
|
||||||
"default": "Default"
|
"default": "Default",
|
||||||
|
"virtual": "Virtual"
|
||||||
},
|
},
|
||||||
"action": {
|
"action": {
|
||||||
"edit": "Edit API key",
|
"edit": "Edit API key",
|
||||||
|
|
@ -208,7 +209,9 @@
|
||||||
"thinkingLevel": "Thinking Level",
|
"thinkingLevel": "Thinking Level",
|
||||||
"thinkingLevelHint": "Extended thinking budget: off, low, medium, high, xhigh, adaptive.",
|
"thinkingLevelHint": "Extended thinking budget: off, low, medium, high, xhigh, adaptive.",
|
||||||
"maxTokensField": "Max Tokens Field",
|
"maxTokensField": "Max Tokens Field",
|
||||||
"maxTokensFieldHint": "Override the request field name for max tokens, e.g. max_completion_tokens."
|
"maxTokensFieldHint": "Override the request field name for max tokens, e.g. max_completion_tokens.",
|
||||||
|
"extraBody": "Extra Body",
|
||||||
|
"extraBodyHint": "Additional JSON fields to inject into the request body, e.g. {\"reasoning_split\": true}."
|
||||||
},
|
},
|
||||||
"edit": {
|
"edit": {
|
||||||
"title": "Configure {{name}}",
|
"title": "Configure {{name}}",
|
||||||
|
|
@ -233,8 +236,6 @@
|
||||||
"qq": "QQ",
|
"qq": "QQ",
|
||||||
"onebot": "OneBot",
|
"onebot": "OneBot",
|
||||||
"wecom": "WeCom",
|
"wecom": "WeCom",
|
||||||
"wecom_app": "WeCom App",
|
|
||||||
"wecom_aibot": "WeCom AI Bot",
|
|
||||||
"whatsapp": "WhatsApp",
|
"whatsapp": "WhatsApp",
|
||||||
"whatsapp_native": "WhatsApp Native",
|
"whatsapp_native": "WhatsApp Native",
|
||||||
"pico": "Web",
|
"pico": "Web",
|
||||||
|
|
@ -434,6 +435,13 @@
|
||||||
"custom_allow_patterns": "Command Whitelist",
|
"custom_allow_patterns": "Command Whitelist",
|
||||||
"custom_allow_patterns_hint": "Add extra command-allow rules, one regular expression per line. A command matching any rule here skips blacklist matching, but other safety limits still apply.",
|
"custom_allow_patterns_hint": "Add extra command-allow rules, one regular expression per line. A command matching any rule here skips blacklist matching, but other safety limits still apply.",
|
||||||
"custom_patterns_placeholder": "^rm\\s+-rf\\b\n^git\\s+push\\b",
|
"custom_patterns_placeholder": "^rm\\s+-rf\\b\n^git\\s+push\\b",
|
||||||
|
"pattern_detector_title": "Pattern Detection Tool",
|
||||||
|
"pattern_detector_hint": "Enter a command to test if it matches any blacklist or whitelist patterns.",
|
||||||
|
"pattern_detector_input_placeholder": "Enter a command to test, e.g., rm -rf /tmp",
|
||||||
|
"pattern_detector_test_button": "Test",
|
||||||
|
"pattern_detector_result_allowed": "Allowed (matches whitelist)",
|
||||||
|
"pattern_detector_result_blocked": "Blocked (matches blacklist)",
|
||||||
|
"pattern_detector_result_no_match": "No match (will use default rules)",
|
||||||
"allow_shell_execution": "Allow Scheduled Commands",
|
"allow_shell_execution": "Allow Scheduled Commands",
|
||||||
"allow_shell_execution_hint": "Allow scheduled tasks to run commands by default. When disabled, users must pass command_confirm=true to schedule a command task.",
|
"allow_shell_execution_hint": "Allow scheduled tasks to run commands by default. When disabled, users must pass command_confirm=true to schedule a command task.",
|
||||||
"cron_exec_timeout": "Scheduled Command Timeout (minutes)",
|
"cron_exec_timeout": "Scheduled Command Timeout (minutes)",
|
||||||
|
|
|
||||||
|
|
@ -154,7 +154,8 @@
|
||||||
"unconfigured": "未配置"
|
"unconfigured": "未配置"
|
||||||
},
|
},
|
||||||
"badge": {
|
"badge": {
|
||||||
"default": "默认"
|
"default": "默认",
|
||||||
|
"virtual": "虚拟"
|
||||||
},
|
},
|
||||||
"action": {
|
"action": {
|
||||||
"edit": "编辑 API Key",
|
"edit": "编辑 API Key",
|
||||||
|
|
@ -208,7 +209,9 @@
|
||||||
"thinkingLevel": "思考级别",
|
"thinkingLevel": "思考级别",
|
||||||
"thinkingLevelHint": "扩展思考预算:off、low、medium、high、xhigh、adaptive。",
|
"thinkingLevelHint": "扩展思考预算:off、low、medium、high、xhigh、adaptive。",
|
||||||
"maxTokensField": "Max Tokens 字段名",
|
"maxTokensField": "Max Tokens 字段名",
|
||||||
"maxTokensFieldHint": "覆盖请求中 max_tokens 的字段名,例如 max_completion_tokens。"
|
"maxTokensFieldHint": "覆盖请求中 max_tokens 的字段名,例如 max_completion_tokens。",
|
||||||
|
"extraBody": "Extra Body",
|
||||||
|
"extraBodyHint": "要注入到请求体中的额外 JSON 字段,例如 {\"reasoning_split\": true}。"
|
||||||
},
|
},
|
||||||
"edit": {
|
"edit": {
|
||||||
"title": "配置 {{name}}",
|
"title": "配置 {{name}}",
|
||||||
|
|
@ -233,8 +236,6 @@
|
||||||
"qq": "QQ",
|
"qq": "QQ",
|
||||||
"onebot": "OneBot",
|
"onebot": "OneBot",
|
||||||
"wecom": "企业微信",
|
"wecom": "企业微信",
|
||||||
"wecom_app": "企业微信应用",
|
|
||||||
"wecom_aibot": "企业微信 AI 机器人",
|
|
||||||
"whatsapp": "WhatsApp",
|
"whatsapp": "WhatsApp",
|
||||||
"whatsapp_native": "WhatsApp Native",
|
"whatsapp_native": "WhatsApp Native",
|
||||||
"pico": "Web",
|
"pico": "Web",
|
||||||
|
|
@ -434,6 +435,13 @@
|
||||||
"custom_allow_patterns": "命令白名单",
|
"custom_allow_patterns": "命令白名单",
|
||||||
"custom_allow_patterns_hint": "用于补充额外的命令放行规则,每行一个正则表达式。命中任意一条规则的命令会跳过黑名单检查,但仍受其他安全限制约束。",
|
"custom_allow_patterns_hint": "用于补充额外的命令放行规则,每行一个正则表达式。命中任意一条规则的命令会跳过黑名单检查,但仍受其他安全限制约束。",
|
||||||
"custom_patterns_placeholder": "^rm\\s+-rf\\b\n^git\\s+push\\b",
|
"custom_patterns_placeholder": "^rm\\s+-rf\\b\n^git\\s+push\\b",
|
||||||
|
"pattern_detector_title": "规则检测工具",
|
||||||
|
"pattern_detector_hint": "输入命令以检测其是否匹配黑名单或白名单规则。",
|
||||||
|
"pattern_detector_input_placeholder": "输入要检测的命令,例如 rm -rf /tmp",
|
||||||
|
"pattern_detector_test_button": "检测",
|
||||||
|
"pattern_detector_result_allowed": "允许(匹配白名单)",
|
||||||
|
"pattern_detector_result_blocked": "阻止(匹配黑名单)",
|
||||||
|
"pattern_detector_result_no_match": "无匹配(将使用默认规则)",
|
||||||
"allow_shell_execution": "允许定时任务运行命令",
|
"allow_shell_execution": "允许定时任务运行命令",
|
||||||
"allow_shell_execution_hint": "开启后,定时任务默认允许运行命令。关闭后,必须显式传入 command_confirm=true 才能创建运行命令的定时任务。",
|
"allow_shell_execution_hint": "开启后,定时任务默认允许运行命令。关闭后,必须显式传入 command_confirm=true 才能创建运行命令的定时任务。",
|
||||||
"cron_exec_timeout": "定时命令超时(分钟)",
|
"cron_exec_timeout": "定时命令超时(分钟)",
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue