Merge remote-tracking branch 'origin/main' into feat/telegram-use-md2
This commit is contained in:
commit
1000daa9e2
230 changed files with 30087 additions and 4809 deletions
|
|
@ -9,6 +9,10 @@
|
||||||
# ── Chat Channel ──────────────────────────
|
# ── Chat Channel ──────────────────────────
|
||||||
# TELEGRAM_BOT_TOKEN=123456:ABC...
|
# TELEGRAM_BOT_TOKEN=123456:ABC...
|
||||||
# DISCORD_BOT_TOKEN=xxx
|
# DISCORD_BOT_TOKEN=xxx
|
||||||
|
# Feishu (飞书)
|
||||||
|
# PICOCLAW_CHANNELS_FEISHU_APP_ID=cli_xxx
|
||||||
|
# PICOCLAW_CHANNELS_FEISHU_APP_SECRET=xxx
|
||||||
|
# PICOCLAW_CHANNELS_FEISHU_RANDOM_REACTION_EMOJI=Typing,OneSecond
|
||||||
|
|
||||||
# ── Web Search (optional) ────────────────
|
# ── Web Search (optional) ────────────────
|
||||||
# BRAVE_SEARCH_API_KEY=BSA...
|
# BRAVE_SEARCH_API_KEY=BSA...
|
||||||
|
|
|
||||||
204
.github/workflows/nightly.yml
vendored
Normal file
204
.github/workflows/nightly.yml
vendored
Normal file
|
|
@ -0,0 +1,204 @@
|
||||||
|
name: Nightly Build
|
||||||
|
|
||||||
|
on:
|
||||||
|
schedule:
|
||||||
|
- cron: '0 0 * * *'
|
||||||
|
workflow_dispatch:
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
create-tag:
|
||||||
|
name: Create Git Tag
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
permissions:
|
||||||
|
contents: write
|
||||||
|
outputs:
|
||||||
|
version: ${{ steps.version.outputs.version }}
|
||||||
|
tag: ${{ steps.version.outputs.tag }}
|
||||||
|
changelog: ${{ steps.version.outputs.changelog }}
|
||||||
|
steps:
|
||||||
|
- name: Checkout
|
||||||
|
uses: actions/checkout@v6
|
||||||
|
with:
|
||||||
|
fetch-depth: 0
|
||||||
|
|
||||||
|
- name: Generate and push tag
|
||||||
|
id: version
|
||||||
|
run: |
|
||||||
|
DATE=$(date -u +%Y%m%d)
|
||||||
|
SHA=$(git rev-parse --short=8 HEAD)
|
||||||
|
BASE_VERSION=$(git describe --tags --match "v*" --exclude "*nightly*" --abbrev=0 2>/dev/null || true)
|
||||||
|
if [ -z "$BASE_VERSION" ] || [ "$BASE_VERSION" = "v0.0.0" ]; then
|
||||||
|
TAG="v0.0.0-nightly.${DATE}.${SHA}"
|
||||||
|
else
|
||||||
|
TAG="${BASE_VERSION}-nightly.${DATE}.${SHA}"
|
||||||
|
fi
|
||||||
|
VERSION=$TAG
|
||||||
|
git config user.name "github-actions[bot]"
|
||||||
|
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||||
|
if git rev-parse -q --verify "refs/tags/$TAG" >/dev/null; then
|
||||||
|
echo "Tag $TAG already exists, reusing existing tag"
|
||||||
|
else
|
||||||
|
git tag -a "$TAG" -m "Nightly build $VERSION"
|
||||||
|
fi
|
||||||
|
git push origin "$TAG"
|
||||||
|
|
||||||
|
COMPARE_URL="https://github.com/${{ github.repository }}/commits/${TAG}"
|
||||||
|
if [ -n "$BASE_VERSION" ] && [ "$BASE_VERSION" != "v0.0.0" ]; then
|
||||||
|
COMPARE_URL="https://github.com/${{ github.repository }}/compare/${BASE_VERSION}...${TAG}"
|
||||||
|
fi
|
||||||
|
echo "changelog=**Full Changelog**: $COMPARE_URL" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
|
echo "version=${VERSION}" >> "$GITHUB_OUTPUT"
|
||||||
|
echo "tag=${TAG}" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
|
release:
|
||||||
|
name: GoReleaser Release
|
||||||
|
needs: create-tag
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
permissions:
|
||||||
|
contents: write
|
||||||
|
packages: write
|
||||||
|
steps:
|
||||||
|
- name: Checkout tag
|
||||||
|
uses: actions/checkout@v6
|
||||||
|
with:
|
||||||
|
fetch-depth: 0
|
||||||
|
ref: ${{ needs.create-tag.outputs.tag }}
|
||||||
|
|
||||||
|
- name: Setup Go from go.mod
|
||||||
|
id: setup-go
|
||||||
|
uses: actions/setup-go@v6
|
||||||
|
with:
|
||||||
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
- name: Setup Node.js
|
||||||
|
uses: actions/setup-node@v4
|
||||||
|
with:
|
||||||
|
node-version: 22
|
||||||
|
|
||||||
|
- name: Setup pnpm
|
||||||
|
run: corepack enable && corepack prepare pnpm@latest --activate
|
||||||
|
|
||||||
|
- name: Set up QEMU
|
||||||
|
uses: docker/setup-qemu-action@v3
|
||||||
|
|
||||||
|
- name: Set up Docker Buildx
|
||||||
|
uses: docker/setup-buildx-action@v3
|
||||||
|
|
||||||
|
- name: Login to GitHub Container Registry
|
||||||
|
uses: docker/login-action@v3
|
||||||
|
with:
|
||||||
|
registry: ghcr.io
|
||||||
|
username: ${{ github.actor }}
|
||||||
|
password: ${{ secrets.GITHUB_TOKEN }}
|
||||||
|
|
||||||
|
- name: Run GoReleaser
|
||||||
|
uses: goreleaser/goreleaser-action@v6
|
||||||
|
with:
|
||||||
|
distribution: goreleaser
|
||||||
|
version: ~> v2
|
||||||
|
args: release --clean
|
||||||
|
env:
|
||||||
|
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
|
GITHUB_REPOSITORY_OWNER: ${{ github.repository_owner }}
|
||||||
|
DOCKERHUB_IMAGE_NAME: ${{ vars.DOCKERHUB_REPOSITORY }}
|
||||||
|
GOVERSION: ${{ steps.setup-go.outputs.go-version }}
|
||||||
|
NIGHTLY_BUILD: "true"
|
||||||
|
MACOS_SIGN_P12: ${{ secrets.MACOS_SIGN_P12 }}
|
||||||
|
MACOS_SIGN_PASSWORD: ${{ secrets.MACOS_SIGN_PASSWORD }}
|
||||||
|
MACOS_NOTARY_ISSUER_ID: ${{ secrets.MACOS_NOTARY_ISSUER_ID }}
|
||||||
|
MACOS_NOTARY_KEY_ID: ${{ secrets.MACOS_NOTARY_KEY_ID }}
|
||||||
|
MACOS_NOTARY_KEY: ${{ secrets.MACOS_NOTARY_KEY }}
|
||||||
|
|
||||||
|
update-rolling:
|
||||||
|
name: Update Rolling Nightly
|
||||||
|
needs: [create-tag, release]
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
permissions:
|
||||||
|
contents: write
|
||||||
|
packages: write
|
||||||
|
steps:
|
||||||
|
- name: Checkout
|
||||||
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
|
- name: Update nightly release
|
||||||
|
env:
|
||||||
|
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
|
TAG: ${{ needs.create-tag.outputs.tag }}
|
||||||
|
TITLE: ${{ needs.create-tag.outputs.version }}
|
||||||
|
run: |
|
||||||
|
CHANGELOG='${{ needs.create-tag.outputs.changelog }}'
|
||||||
|
NOTES=$(cat <<EOF
|
||||||
|
Nightly build for **${TITLE}**
|
||||||
|
|
||||||
|
This is an automated build and may be unstable. Use with caution.
|
||||||
|
|
||||||
|
${CHANGELOG}
|
||||||
|
EOF
|
||||||
|
)
|
||||||
|
|
||||||
|
# Download assets from the newly created release if it exists,
|
||||||
|
# otherwise fall back to using locally built dist/ artifacts.
|
||||||
|
mkdir -p build
|
||||||
|
if gh release view "$TAG" >/dev/null 2>&1; then
|
||||||
|
echo "Downloading assets from GitHub release for $TAG..."
|
||||||
|
gh release download "$TAG" --dir build
|
||||||
|
else
|
||||||
|
echo "GitHub release for $TAG not found; falling back to local dist/ artifacts..."
|
||||||
|
if [ -d "dist" ]; then
|
||||||
|
cp -R dist/* build/
|
||||||
|
else
|
||||||
|
echo "Error: no GitHub release for $TAG and no local dist/ directory found." >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Delete existing nightly release and tag to avoid conflicts
|
||||||
|
echo "Deleting existing nightly release and tag..."
|
||||||
|
gh release delete nightly --cleanup-tag -y || true
|
||||||
|
git push origin :refs/tags/nightly || true
|
||||||
|
|
||||||
|
gh release create nightly \
|
||||||
|
--title "Nightly Build" \
|
||||||
|
--notes "$NOTES" \
|
||||||
|
--target "${{ github.sha }}" \
|
||||||
|
--prerelease \
|
||||||
|
build/*
|
||||||
|
|
||||||
|
echo "Cleaning up old nightly releases (keeping only the most recent)..."
|
||||||
|
gh release list --limit 100 --json tagName -q '.[].tagName | select(contains("-nightly."))' | tail -n +2 | while read -r old_tag; do
|
||||||
|
if [ -n "$old_tag" ] && [ "$old_tag" != "$TAG" ]; then
|
||||||
|
echo "Deleting old nightly release: $old_tag"
|
||||||
|
gh release delete "$old_tag" --cleanup-tag -y || true
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
|
||||||
|
echo "Cleaning up old 'vX.X.X-nightly...' Docker images on GHCR..."
|
||||||
|
OWNER="${{ github.repository_owner }}"
|
||||||
|
PACKAGE_NAME="${{ github.event.repository.name }}"
|
||||||
|
|
||||||
|
# Check if owner is an organization or user
|
||||||
|
ORG_TEST=$(gh api -H "Accept: application/vnd.github+json" /orgs/$OWNER 2>/dev/null || true)
|
||||||
|
if echo "$ORG_TEST" | grep -q '"login"'; then
|
||||||
|
ACCOUNT_TYPE="orgs"
|
||||||
|
else
|
||||||
|
ACCOUNT_TYPE="users"
|
||||||
|
fi
|
||||||
|
|
||||||
|
PACKAGE_URL="/${ACCOUNT_TYPE}/${OWNER}/packages/container/${PACKAGE_NAME}/versions"
|
||||||
|
OLD_NIGHTLY_VERSIONS=$(gh api --paginate -H "Accept: application/vnd.github+json" \
|
||||||
|
-H "X-GitHub-Api-Version: 2022-11-28" \
|
||||||
|
"$PACKAGE_URL" \
|
||||||
|
--jq ". | map(select(any(.metadata.container.tags[]; contains(\"-nightly.\") and (. != \"nightly\") and (. != \"$TAG\")))) | .[].id" 2>/dev/null || true)
|
||||||
|
|
||||||
|
for version_id in $OLD_NIGHTLY_VERSIONS; do
|
||||||
|
if [ -n "$version_id" ]; then
|
||||||
|
echo "Deleting Docker image version ID: $version_id"
|
||||||
|
gh api -X DELETE -H "Accept: application/vnd.github+json" \
|
||||||
|
-H "X-GitHub-Api-Version: 2022-11-28" \
|
||||||
|
"/${ACCOUNT_TYPE}/${OWNER}/packages/container/${PACKAGE_NAME}/versions/$version_id" || true
|
||||||
|
fi
|
||||||
|
done
|
||||||
13
.github/workflows/release.yml
vendored
13
.github/workflows/release.yml
vendored
|
|
@ -65,6 +65,14 @@ jobs:
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
|
|
||||||
|
- name: Setup Node.js
|
||||||
|
uses: actions/setup-node@v4
|
||||||
|
with:
|
||||||
|
node-version: 22
|
||||||
|
|
||||||
|
- name: Setup pnpm
|
||||||
|
run: corepack enable && corepack prepare pnpm@latest --activate
|
||||||
|
|
||||||
- name: Set up QEMU
|
- name: Set up QEMU
|
||||||
uses: docker/setup-qemu-action@v3
|
uses: docker/setup-qemu-action@v3
|
||||||
|
|
||||||
|
|
@ -96,6 +104,11 @@ jobs:
|
||||||
GITHUB_REPOSITORY_OWNER: ${{ github.repository_owner }}
|
GITHUB_REPOSITORY_OWNER: ${{ github.repository_owner }}
|
||||||
DOCKERHUB_IMAGE_NAME: ${{ vars.DOCKERHUB_REPOSITORY }}
|
DOCKERHUB_IMAGE_NAME: ${{ vars.DOCKERHUB_REPOSITORY }}
|
||||||
GOVERSION: ${{ steps.setup-go.outputs.go-version }}
|
GOVERSION: ${{ steps.setup-go.outputs.go-version }}
|
||||||
|
MACOS_SIGN_P12: ${{ secrets.MACOS_SIGN_P12 }}
|
||||||
|
MACOS_SIGN_PASSWORD: ${{ secrets.MACOS_SIGN_PASSWORD }}
|
||||||
|
MACOS_NOTARY_ISSUER_ID: ${{ secrets.MACOS_NOTARY_ISSUER_ID }}
|
||||||
|
MACOS_NOTARY_KEY_ID: ${{ secrets.MACOS_NOTARY_KEY_ID }}
|
||||||
|
MACOS_NOTARY_KEY: ${{ secrets.MACOS_NOTARY_KEY }}
|
||||||
|
|
||||||
- name: Apply release flags
|
- name: Apply release flags
|
||||||
shell: bash
|
shell: bash
|
||||||
|
|
|
||||||
9
.gitignore
vendored
9
.gitignore
vendored
|
|
@ -47,9 +47,16 @@ docs/plans/
|
||||||
|
|
||||||
# Added by goreleaser init:
|
# Added by goreleaser init:
|
||||||
dist/
|
dist/
|
||||||
|
*.vite/
|
||||||
|
|
||||||
# Windows Application Icon/Resource
|
# Windows Application Icon/Resource
|
||||||
*.syso
|
*.syso
|
||||||
|
|
||||||
# Test channels
|
# Test telegram integration
|
||||||
cmd/telegram/
|
cmd/telegram/
|
||||||
|
|
||||||
|
# Keep embedded backend dist directory placeholder in VCS
|
||||||
|
!web/backend/dist/
|
||||||
|
web/backend/dist/*
|
||||||
|
!web/backend/dist/.gitkeep
|
||||||
|
>>>>>>> origin/main
|
||||||
|
|
|
||||||
|
|
@ -6,8 +6,9 @@ before:
|
||||||
hooks:
|
hooks:
|
||||||
- go mod tidy
|
- go mod tidy
|
||||||
- go generate ./...
|
- go generate ./...
|
||||||
|
- sh -c 'cd web/frontend && pnpm install && pnpm build:backend'
|
||||||
- go install github.com/tc-hib/go-winres@latest
|
- go install github.com/tc-hib/go-winres@latest
|
||||||
- go-winres make --in cmd/picoclaw-launcher/winres/winres.json --out cmd/picoclaw-launcher/rsrc --product-version={{ .Version }} --file-version={{ .Version }}
|
- go-winres make --in web/backend/winres/winres.json --out web/backend/rsrc --product-version={{ .Version }} --file-version={{ .Version }}
|
||||||
|
|
||||||
builds:
|
builds:
|
||||||
- id: picoclaw
|
- id: picoclaw
|
||||||
|
|
@ -17,10 +18,10 @@ builds:
|
||||||
- stdjson
|
- stdjson
|
||||||
ldflags:
|
ldflags:
|
||||||
- -s -w
|
- -s -w
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.version={{ .Version }}
|
- -X github.com/sipeed/picoclaw/pkg/config.Version={{ .Version }}
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.gitCommit={{ .ShortCommit }}
|
- -X github.com/sipeed/picoclaw/pkg/config.GitCommit={{ .ShortCommit }}
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.buildTime={{ .Date }}
|
- -X github.com/sipeed/picoclaw/pkg/config.BuildTime={{ .Date }}
|
||||||
- -X github.com/sipeed/picoclaw/cmd/picoclaw/internal.goVersion={{ .Env.GOVERSION }}
|
- -X github.com/sipeed/picoclaw/pkg/config.GoVersion={{ .Env.GOVERSION }}
|
||||||
goos:
|
goos:
|
||||||
- linux
|
- linux
|
||||||
- windows
|
- windows
|
||||||
|
|
@ -32,9 +33,13 @@ builds:
|
||||||
- riscv64
|
- riscv64
|
||||||
- loong64
|
- loong64
|
||||||
- arm
|
- arm
|
||||||
|
- s390x
|
||||||
|
- mipsle
|
||||||
goarm:
|
goarm:
|
||||||
- "6"
|
- "6"
|
||||||
- "7"
|
- "7"
|
||||||
|
gomips:
|
||||||
|
- softfloat
|
||||||
main: ./cmd/picoclaw
|
main: ./cmd/picoclaw
|
||||||
ignore:
|
ignore:
|
||||||
- goos: windows
|
- goos: windows
|
||||||
|
|
@ -59,10 +64,14 @@ builds:
|
||||||
- riscv64
|
- riscv64
|
||||||
- loong64
|
- loong64
|
||||||
- arm
|
- arm
|
||||||
|
- s390x
|
||||||
|
- mipsle
|
||||||
goarm:
|
goarm:
|
||||||
- "6"
|
- "6"
|
||||||
- "7"
|
- "7"
|
||||||
main: ./cmd/picoclaw-launcher
|
gomips:
|
||||||
|
- softfloat
|
||||||
|
main: ./web/backend
|
||||||
ignore:
|
ignore:
|
||||||
- goos: windows
|
- goos: windows
|
||||||
goarch: arm
|
goarch: arm
|
||||||
|
|
@ -86,9 +95,13 @@ builds:
|
||||||
- riscv64
|
- riscv64
|
||||||
- loong64
|
- loong64
|
||||||
- arm
|
- arm
|
||||||
|
- s390x
|
||||||
|
- mipsle
|
||||||
goarm:
|
goarm:
|
||||||
- "6"
|
- "6"
|
||||||
- "7"
|
- "7"
|
||||||
|
gomips:
|
||||||
|
- softfloat
|
||||||
main: ./cmd/picoclaw-launcher-tui
|
main: ./cmd/picoclaw-launcher-tui
|
||||||
ignore:
|
ignore:
|
||||||
- goos: windows
|
- goos: windows
|
||||||
|
|
@ -103,15 +116,49 @@ dockers_v2:
|
||||||
- picoclaw
|
- picoclaw
|
||||||
images:
|
images:
|
||||||
- "ghcr.io/{{ .Env.GITHUB_REPOSITORY_OWNER }}/picoclaw"
|
- "ghcr.io/{{ .Env.GITHUB_REPOSITORY_OWNER }}/picoclaw"
|
||||||
- "docker.io/{{ .Env.DOCKERHUB_IMAGE_NAME }}"
|
- '{{ if not (isEnvSet "NIGHTLY_BUILD") }}docker.io/{{ .Env.DOCKERHUB_IMAGE_NAME }}{{ end }}'
|
||||||
tags:
|
tags:
|
||||||
- "{{ .Tag }}"
|
- "{{ .Tag }}"
|
||||||
- "latest"
|
- '{{ if isEnvSet "NIGHTLY_BUILD" }}nightly{{ else }}latest{{ end }}'
|
||||||
platforms:
|
platforms:
|
||||||
- linux/amd64
|
- linux/amd64
|
||||||
- linux/arm64
|
- linux/arm64
|
||||||
- linux/riscv64
|
- linux/riscv64
|
||||||
|
|
||||||
|
- id: picoclaw-launcher
|
||||||
|
dockerfile: docker/Dockerfile.goreleaser.launcher
|
||||||
|
ids:
|
||||||
|
- picoclaw
|
||||||
|
- picoclaw-launcher
|
||||||
|
- picoclaw-launcher-tui
|
||||||
|
images:
|
||||||
|
- "ghcr.io/{{ .Env.GITHUB_REPOSITORY_OWNER }}/picoclaw"
|
||||||
|
- '{{ if not (isEnvSet "NIGHTLY_BUILD") }}docker.io/{{ .Env.DOCKERHUB_IMAGE_NAME }}{{ end }}'
|
||||||
|
tags:
|
||||||
|
- "{{ .Tag }}-launcher"
|
||||||
|
- '{{ if isEnvSet "NIGHTLY_BUILD" }}nightly-launcher{{ else }}launcher{{ end }}'
|
||||||
|
platforms:
|
||||||
|
- linux/amd64
|
||||||
|
- linux/arm64
|
||||||
|
- linux/riscv64
|
||||||
|
|
||||||
|
notarize:
|
||||||
|
macos:
|
||||||
|
- enabled: '{{ isEnvSet "MACOS_SIGN_P12" }}'
|
||||||
|
ids:
|
||||||
|
- picoclaw
|
||||||
|
- picoclaw-launcher
|
||||||
|
- picoclaw-launcher-tui
|
||||||
|
sign:
|
||||||
|
certificate: "{{.Env.MACOS_SIGN_P12}}"
|
||||||
|
password: "{{.Env.MACOS_SIGN_PASSWORD}}"
|
||||||
|
notarize:
|
||||||
|
issuer_id: "{{.Env.MACOS_NOTARY_ISSUER_ID}}"
|
||||||
|
key_id: "{{.Env.MACOS_NOTARY_KEY_ID}}"
|
||||||
|
key: "{{.Env.MACOS_NOTARY_KEY}}"
|
||||||
|
wait: true
|
||||||
|
timeout: 20m
|
||||||
|
|
||||||
archives:
|
archives:
|
||||||
- formats: [tar.gz]
|
- formats: [tar.gz]
|
||||||
# this name template makes the OS and Arch compatible with the results of `uname`.
|
# this name template makes the OS and Arch compatible with the results of `uname`.
|
||||||
|
|
@ -129,7 +176,7 @@ archives:
|
||||||
|
|
||||||
nfpms:
|
nfpms:
|
||||||
- id: picoclaw
|
- id: picoclaw
|
||||||
builds:
|
ids:
|
||||||
- picoclaw
|
- picoclaw
|
||||||
- picoclaw-launcher
|
- picoclaw-launcher
|
||||||
- picoclaw-launcher-tui
|
- picoclaw-launcher-tui
|
||||||
|
|
@ -149,6 +196,11 @@ nfpms:
|
||||||
- rpm
|
- rpm
|
||||||
- deb
|
- deb
|
||||||
bindir: /usr/bin
|
bindir: /usr/bin
|
||||||
|
contents:
|
||||||
|
- src: web/picoclaw-launcher.desktop
|
||||||
|
dst: /usr/share/applications/picoclaw-launcher.desktop
|
||||||
|
- src: web/picoclaw-launcher.png
|
||||||
|
dst: /usr/share/icons/hicolor/512x512/apps/picoclaw-launcher.png
|
||||||
|
|
||||||
changelog:
|
changelog:
|
||||||
sort: asc
|
sort: asc
|
||||||
|
|
|
||||||
16
Makefile
16
Makefile
|
|
@ -11,8 +11,8 @@ VERSION?=$(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
|
||||||
GIT_COMMIT=$(shell git rev-parse --short=8 HEAD 2>/dev/null || echo "dev")
|
GIT_COMMIT=$(shell git rev-parse --short=8 HEAD 2>/dev/null || echo "dev")
|
||||||
BUILD_TIME=$(shell date +%FT%T%z)
|
BUILD_TIME=$(shell date +%FT%T%z)
|
||||||
GO_VERSION=$(shell $(GO) version | awk '{print $$3}')
|
GO_VERSION=$(shell $(GO) version | awk '{print $$3}')
|
||||||
INTERNAL=github.com/sipeed/picoclaw/cmd/picoclaw/internal
|
CONFIG_PKG=github.com/sipeed/picoclaw/pkg/config
|
||||||
LDFLAGS=-ldflags "-X $(INTERNAL).version=$(VERSION) -X $(INTERNAL).gitCommit=$(GIT_COMMIT) -X $(INTERNAL).buildTime=$(BUILD_TIME) -X $(INTERNAL).goVersion=$(GO_VERSION) -s -w"
|
LDFLAGS=-ldflags "-X $(CONFIG_PKG).Version=$(VERSION) -X $(CONFIG_PKG).GitCommit=$(GIT_COMMIT) -X $(CONFIG_PKG).BuildTime=$(BUILD_TIME) -X $(CONFIG_PKG).GoVersion=$(GO_VERSION) -s -w"
|
||||||
|
|
||||||
# Go variables
|
# Go variables
|
||||||
GO?=CGO_ENABLED=0 go
|
GO?=CGO_ENABLED=0 go
|
||||||
|
|
@ -111,6 +111,18 @@ build: generate
|
||||||
@echo "Build complete: $(BINARY_PATH)"
|
@echo "Build complete: $(BINARY_PATH)"
|
||||||
@ln -sf $(BINARY_NAME)-$(PLATFORM)-$(ARCH) $(BUILD_DIR)/$(BINARY_NAME)
|
@ln -sf $(BINARY_NAME)-$(PLATFORM)-$(ARCH) $(BUILD_DIR)/$(BINARY_NAME)
|
||||||
|
|
||||||
|
## build-launcher: Build the picoclaw-launcher (web console) binary
|
||||||
|
build-launcher:
|
||||||
|
@echo "Building picoclaw-launcher for $(PLATFORM)/$(ARCH)..."
|
||||||
|
@mkdir -p $(BUILD_DIR)
|
||||||
|
@if [ ! -f web/backend/dist/index.html ]; then \
|
||||||
|
echo "Building frontend..."; \
|
||||||
|
cd web/frontend && pnpm install && pnpm build:backend; \
|
||||||
|
fi
|
||||||
|
@$(GO) build $(GOFLAGS) -o $(BUILD_DIR)/picoclaw-launcher-$(PLATFORM)-$(ARCH) ./web/backend
|
||||||
|
@ln -sf picoclaw-launcher-$(PLATFORM)-$(ARCH) $(BUILD_DIR)/picoclaw-launcher
|
||||||
|
@echo "Build complete: $(BUILD_DIR)/picoclaw-launcher"
|
||||||
|
|
||||||
## build-whatsapp-native: Build with WhatsApp native (whatsmeow) support; larger binary
|
## build-whatsapp-native: Build with WhatsApp native (whatsmeow) support; larger binary
|
||||||
build-whatsapp-native: generate
|
build-whatsapp-native: generate
|
||||||
## @echo "Building $(BINARY_NAME) with WhatsApp native for $(PLATFORM)/$(ARCH)..."
|
## @echo "Building $(BINARY_NAME) with WhatsApp native for $(PLATFORM)/$(ARCH)..."
|
||||||
|
|
|
||||||
50
README.md
50
README.md
|
|
@ -194,6 +194,19 @@ docker compose -f docker/docker-compose.yml logs -f picoclaw-gateway
|
||||||
docker compose -f docker/docker-compose.yml --profile gateway down
|
docker compose -f docker/docker-compose.yml --profile gateway down
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Launcher Mode (Web Console)
|
||||||
|
|
||||||
|
The `launcher` image includes all three binaries (`picoclaw`, `picoclaw-launcher`, `picoclaw-launcher-tui`) and starts the web console by default, which provides a browser-based UI for configuration and chat.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker compose -f docker/docker-compose.yml --profile launcher up -d
|
||||||
|
```
|
||||||
|
|
||||||
|
Open http://localhost:18800 in your browser. The launcher manages the gateway process automatically.
|
||||||
|
|
||||||
|
> [!WARNING]
|
||||||
|
> The web console does not yet support authentication. Avoid exposing it to the public internet.
|
||||||
|
|
||||||
### Agent Mode (One-shot)
|
### Agent Mode (One-shot)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
@ -308,7 +321,7 @@ That's it! You have a working AI assistant in 2 minutes.
|
||||||
|
|
||||||
## 💬 Chat Apps
|
## 💬 Chat Apps
|
||||||
|
|
||||||
Talk to your picoclaw through Telegram, Discord, WhatsApp, DingTalk, LINE, or WeCom
|
Talk to your picoclaw through Telegram, Discord, WhatsApp, Matrix, QQ, DingTalk, LINE, or WeCom
|
||||||
|
|
||||||
> **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**: 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.
|
||||||
|
|
||||||
|
|
@ -317,6 +330,7 @@ Talk to your picoclaw through Telegram, Discord, WhatsApp, DingTalk, LINE, or We
|
||||||
| **Telegram** | Easy (just a token) |
|
| **Telegram** | Easy (just a token) |
|
||||||
| **Discord** | Easy (bot token + intents) |
|
| **Discord** | Easy (bot token + intents) |
|
||||||
| **WhatsApp** | Easy (native: QR scan; or bridge URL) |
|
| **WhatsApp** | Easy (native: QR scan; or bridge URL) |
|
||||||
|
| **Matrix** | Medium (homeserver + bot access token) |
|
||||||
| **QQ** | Easy (AppID + AppSecret) |
|
| **QQ** | Easy (AppID + AppSecret) |
|
||||||
| **DingTalk** | Medium (app credentials) |
|
| **DingTalk** | Medium (app credentials) |
|
||||||
| **LINE** | Medium (credentials + webhook URL) |
|
| **LINE** | Medium (credentials + webhook URL) |
|
||||||
|
|
@ -530,6 +544,40 @@ picoclaw gateway
|
||||||
```
|
```
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>Matrix</b></summary>
|
||||||
|
|
||||||
|
**1. Prepare bot account**
|
||||||
|
|
||||||
|
* Use your preferred homeserver (e.g. `https://matrix.org` or self-hosted)
|
||||||
|
* Create a bot user and obtain its access token
|
||||||
|
|
||||||
|
**2. Configure**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"matrix": {
|
||||||
|
"enabled": true,
|
||||||
|
"homeserver": "https://matrix.org",
|
||||||
|
"user_id": "@your-bot:matrix.org",
|
||||||
|
"access_token": "YOUR_MATRIX_ACCESS_TOKEN",
|
||||||
|
"allow_from": []
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**3. Run**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
picoclaw gateway
|
||||||
|
```
|
||||||
|
|
||||||
|
For full options (`device_id`, `join_on_invite`, `group_trigger`, `placeholder`, `reasoning_channel_id`), see [Matrix Channel Configuration Guide](docs/channels/matrix/README.md).
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>LINE</b></summary>
|
<summary><b>LINE</b></summary>
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -299,6 +299,7 @@ PicoClaw 支持多种聊天平台,使您的 Agent 能够连接到任何地方
|
||||||
| **Telegram** | ⭐ 简单 | 推荐,支持语音转文字,长轮询无需公网 | [查看文档](docs/channels/telegram/README.zh.md) |
|
| **Telegram** | ⭐ 简单 | 推荐,支持语音转文字,长轮询无需公网 | [查看文档](docs/channels/telegram/README.zh.md) |
|
||||||
| **Discord** | ⭐ 简单 | Socket Mode,支持群组/私信,Bot 生态成熟 | [查看文档](docs/channels/discord/README.zh.md) |
|
| **Discord** | ⭐ 简单 | Socket Mode,支持群组/私信,Bot 生态成熟 | [查看文档](docs/channels/discord/README.zh.md) |
|
||||||
| **Slack** | ⭐ 简单 | **Socket Mode** (无需公网 IP),企业级支持 | [查看文档](docs/channels/slack/README.zh.md) |
|
| **Slack** | ⭐ 简单 | **Socket Mode** (无需公网 IP),企业级支持 | [查看文档](docs/channels/slack/README.zh.md) |
|
||||||
|
| **Matrix** | ⭐⭐ 中等 | 联邦协议,支持自建 homeserver 与公开服务器 | [查看文档](docs/channels/matrix/README.zh.md) |
|
||||||
| **QQ** | ⭐⭐ 中等 | 官方机器人 API,适合国内社群 | [查看文档](docs/channels/qq/README.zh.md) |
|
| **QQ** | ⭐⭐ 中等 | 官方机器人 API,适合国内社群 | [查看文档](docs/channels/qq/README.zh.md) |
|
||||||
| **钉钉 (DingTalk)** | ⭐⭐ 中等 | Stream 模式无需公网,企业办公首选 | [查看文档](docs/channels/dingtalk/README.zh.md) |
|
| **钉钉 (DingTalk)** | ⭐⭐ 中等 | Stream 模式无需公网,企业办公首选 | [查看文档](docs/channels/dingtalk/README.zh.md) |
|
||||||
| **企业微信 (WeCom)** | ⭐⭐⭐ 较难 | 支持群机器人(Webhook)、自建应用(API)和智能机器人(AI Bot) | [Bot 文档](docs/channels/wecom/wecom_bot/README.zh.md) / [App 文档](docs/channels/wecom/wecom_app/README.zh.md) / [AI Bot 文档](docs/channels/wecom/wecom_aibot/README.zh.md) |
|
| **企业微信 (WeCom)** | ⭐⭐⭐ 较难 | 支持群机器人(Webhook)、自建应用(API)和智能机器人(AI Bot) | [Bot 文档](docs/channels/wecom/wecom_bot/README.zh.md) / [App 文档](docs/channels/wecom/wecom_app/README.zh.md) / [AI Bot 文档](docs/channels/wecom/wecom_aibot/README.zh.md) |
|
||||||
|
|
|
||||||
Binary file not shown.
|
Before Width: | Height: | Size: 386 KiB After Width: | Height: | Size: 348 KiB |
|
|
@ -1,6 +1,7 @@
|
||||||
package ui
|
package ui
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
|
@ -67,6 +68,7 @@ func Run() error {
|
||||||
root := tview.NewFlex().SetDirection(tview.FlexRow)
|
root := tview.NewFlex().SetDirection(tview.FlexRow)
|
||||||
root.AddItem(bannerView(), 6, 0, false)
|
root.AddItem(bannerView(), 6, 0, false)
|
||||||
root.AddItem(state.pages, 0, 1, true)
|
root.AddItem(state.pages, 0, 1, true)
|
||||||
|
root.AddItem(footerView(), 1, 0, false)
|
||||||
|
|
||||||
if err := state.app.SetRoot(root, true).EnableMouse(false).Run(); err != nil {
|
if err := state.app.SetRoot(root, true).EnableMouse(false).Run(); err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
@ -102,7 +104,7 @@ func (s *appState) pop() {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *appState) mainMenu() tview.Primitive {
|
func (s *appState) mainMenu() tview.Primitive {
|
||||||
menu := NewMenu("Config Menu", nil)
|
menu := NewMenu("Menu", nil)
|
||||||
refreshMainMenu(menu, s)
|
refreshMainMenu(menu, s)
|
||||||
menu.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
menu.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
switch event.Key() {
|
switch event.Key() {
|
||||||
|
|
@ -110,10 +112,7 @@ func (s *appState) mainMenu() tview.Primitive {
|
||||||
s.requestExit()
|
s.requestExit()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if event.Rune() == 'q' {
|
|
||||||
s.requestExit()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return event
|
return event
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
@ -131,6 +130,32 @@ func (s *appState) refreshMenu(name string, menu *Menu) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *appState) countChannels() (enabled int, total int) {
|
||||||
|
c := s.config.Channels
|
||||||
|
entries := []bool{
|
||||||
|
c.Telegram.Enabled,
|
||||||
|
c.Discord.Enabled,
|
||||||
|
c.QQ.Enabled,
|
||||||
|
c.MaixCam.Enabled,
|
||||||
|
c.WhatsApp.Enabled,
|
||||||
|
c.Feishu.Enabled,
|
||||||
|
c.DingTalk.Enabled,
|
||||||
|
c.Slack.Enabled,
|
||||||
|
c.Matrix.Enabled,
|
||||||
|
c.LINE.Enabled,
|
||||||
|
c.OneBot.Enabled,
|
||||||
|
c.WeCom.Enabled,
|
||||||
|
c.WeComApp.Enabled,
|
||||||
|
}
|
||||||
|
total = len(entries)
|
||||||
|
for _, v := range entries {
|
||||||
|
if v {
|
||||||
|
enabled++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return enabled, total
|
||||||
|
}
|
||||||
|
|
||||||
func refreshMainMenuIfPresent(s *appState) {
|
func refreshMainMenuIfPresent(s *appState) {
|
||||||
if menu, ok := s.menus["main"]; ok {
|
if menu, ok := s.menus["main"]; ok {
|
||||||
refreshMainMenu(menu, s)
|
refreshMainMenu(menu, s)
|
||||||
|
|
@ -141,6 +166,7 @@ func refreshMainMenu(menu *Menu, s *appState) {
|
||||||
selectedModel := s.selectedModelName()
|
selectedModel := s.selectedModelName()
|
||||||
modelReady := selectedModel != ""
|
modelReady := selectedModel != ""
|
||||||
channelReady := s.hasEnabledChannel()
|
channelReady := s.hasEnabledChannel()
|
||||||
|
enabledCount, totalChannels := s.countChannels()
|
||||||
gatewayRunning := s.gatewayCmd != nil || s.isGatewayRunning()
|
gatewayRunning := s.gatewayCmd != nil || s.isGatewayRunning()
|
||||||
|
|
||||||
gatewayLabel := "Start Gateway"
|
gatewayLabel := "Start Gateway"
|
||||||
|
|
@ -153,7 +179,7 @@ func refreshMainMenu(menu *Menu, s *appState) {
|
||||||
items := []MenuItem{
|
items := []MenuItem{
|
||||||
{
|
{
|
||||||
Label: rootModelLabel(selectedModel),
|
Label: rootModelLabel(selectedModel),
|
||||||
Description: rootModelDescription(selectedModel),
|
Description: rootModelDescription(),
|
||||||
Action: func() {
|
Action: func() {
|
||||||
s.push("model", s.modelMenu())
|
s.push("model", s.modelMenu())
|
||||||
},
|
},
|
||||||
|
|
@ -167,7 +193,7 @@ func refreshMainMenu(menu *Menu, s *appState) {
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
Label: rootChannelLabel(channelReady),
|
Label: rootChannelLabel(channelReady),
|
||||||
Description: rootChannelDescription(channelReady),
|
Description: fmt.Sprintf("%d/%d enabled", enabledCount, totalChannels),
|
||||||
Action: func() {
|
Action: func() {
|
||||||
s.push("channel", s.channelMenu())
|
s.push("channel", s.channelMenu())
|
||||||
},
|
},
|
||||||
|
|
@ -311,16 +337,13 @@ func (s *appState) selectedModelName() string {
|
||||||
|
|
||||||
func rootModelLabel(selected string) string {
|
func rootModelLabel(selected string) string {
|
||||||
if selected == "" {
|
if selected == "" {
|
||||||
return "Model (no model selected)"
|
return "Model (None)"
|
||||||
}
|
}
|
||||||
return "Model (" + selected + ")"
|
return "Model (" + selected + ")"
|
||||||
}
|
}
|
||||||
|
|
||||||
func rootModelDescription(selected string) string {
|
func rootModelDescription() string {
|
||||||
if selected == "" {
|
return "Using SPACE to choose your model"
|
||||||
return "no model selected"
|
|
||||||
}
|
|
||||||
return "selected"
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func rootChannelLabel(valid bool) string {
|
func rootChannelLabel(valid bool) string {
|
||||||
|
|
@ -330,13 +353,6 @@ func rootChannelLabel(valid bool) string {
|
||||||
return "Channel"
|
return "Channel"
|
||||||
}
|
}
|
||||||
|
|
||||||
func rootChannelDescription(valid bool) string {
|
|
||||||
if !valid {
|
|
||||||
return "no channel enabled"
|
|
||||||
}
|
|
||||||
return "enabled"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *appState) startTalk() {
|
func (s *appState) startTalk() {
|
||||||
if !s.isActiveModelValid() {
|
if !s.isActiveModelValid() {
|
||||||
s.showMessage("Model required", "Select a valid model before starting talk")
|
s.showMessage("Model required", "Select a valid model before starting talk")
|
||||||
|
|
@ -423,7 +439,7 @@ func (s *appState) hasEnabledChannel() bool {
|
||||||
c := s.config.Channels
|
c := s.config.Channels
|
||||||
return c.Telegram.Enabled || c.Discord.Enabled || c.QQ.Enabled || c.MaixCam.Enabled ||
|
return c.Telegram.Enabled || c.Discord.Enabled || c.QQ.Enabled || c.MaixCam.Enabled ||
|
||||||
c.WhatsApp.Enabled || c.Feishu.Enabled || c.DingTalk.Enabled || c.Slack.Enabled ||
|
c.WhatsApp.Enabled || c.Feishu.Enabled || c.DingTalk.Enabled || c.Slack.Enabled ||
|
||||||
c.LINE.Enabled || c.OneBot.Enabled || c.WeCom.Enabled || c.WeComApp.Enabled
|
c.Matrix.Enabled || c.LINE.Enabled || c.OneBot.Enabled || c.WeCom.Enabled || c.WeComApp.Enabled
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *appState) confirmApplyOrDiscard(onApply func(), onDiscard func()) {
|
func (s *appState) confirmApplyOrDiscard(onApply func(), onDiscard func()) {
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,6 @@ import (
|
||||||
|
|
||||||
func (s *appState) buildChannelMenuItems() []MenuItem {
|
func (s *appState) buildChannelMenuItems() []MenuItem {
|
||||||
return []MenuItem{
|
return []MenuItem{
|
||||||
{Label: "Back", Description: "Return to main menu", Action: func() { s.pop() }},
|
|
||||||
channelItem(
|
channelItem(
|
||||||
"Telegram",
|
"Telegram",
|
||||||
"Telegram bot settings",
|
"Telegram bot settings",
|
||||||
|
|
@ -61,6 +60,12 @@ func (s *appState) buildChannelMenuItems() []MenuItem {
|
||||||
s.config.Channels.Slack.Enabled,
|
s.config.Channels.Slack.Enabled,
|
||||||
func() { s.push("channel-slack", s.slackForm()) },
|
func() { s.push("channel-slack", s.slackForm()) },
|
||||||
),
|
),
|
||||||
|
channelItem(
|
||||||
|
"Matrix",
|
||||||
|
"Matrix bot settings",
|
||||||
|
s.config.Channels.Matrix.Enabled,
|
||||||
|
func() { s.push("channel-matrix", s.matrixForm()) },
|
||||||
|
),
|
||||||
channelItem(
|
channelItem(
|
||||||
"LINE",
|
"LINE",
|
||||||
"LINE bot settings",
|
"LINE bot settings",
|
||||||
|
|
@ -95,10 +100,6 @@ func (s *appState) channelMenu() tview.Primitive {
|
||||||
s.pop()
|
s.pop()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if event.Rune() == 'q' {
|
|
||||||
s.pop()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return event
|
return event
|
||||||
})
|
})
|
||||||
return menu
|
return menu
|
||||||
|
|
@ -233,6 +234,28 @@ func (s *appState) lineForm() tview.Primitive {
|
||||||
return wrapWithBack(form, s)
|
return wrapWithBack(form, s)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *appState) matrixForm() tview.Primitive {
|
||||||
|
cfg := &s.config.Channels.Matrix
|
||||||
|
form := baseChannelForm("Matrix", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
|
||||||
|
form.AddInputField("Homeserver", cfg.Homeserver, 128, nil, func(text string) {
|
||||||
|
cfg.Homeserver = strings.TrimSpace(text)
|
||||||
|
})
|
||||||
|
form.AddInputField("User ID", cfg.UserID, 128, nil, func(text string) {
|
||||||
|
cfg.UserID = strings.TrimSpace(text)
|
||||||
|
})
|
||||||
|
form.AddInputField("Access Token", cfg.AccessToken, 128, nil, func(text string) {
|
||||||
|
cfg.AccessToken = strings.TrimSpace(text)
|
||||||
|
})
|
||||||
|
form.AddInputField("Device ID", cfg.DeviceID, 128, nil, func(text string) {
|
||||||
|
cfg.DeviceID = strings.TrimSpace(text)
|
||||||
|
})
|
||||||
|
form.AddCheckbox("Join On Invite", cfg.JoinOnInvite, func(checked bool) {
|
||||||
|
cfg.JoinOnInvite = checked
|
||||||
|
})
|
||||||
|
addAllowFromField(form, &cfg.AllowFrom)
|
||||||
|
return wrapWithBack(form, s)
|
||||||
|
}
|
||||||
|
|
||||||
func (s *appState) onebotForm() tview.Primitive {
|
func (s *appState) onebotForm() tview.Primitive {
|
||||||
cfg := &s.config.Channels.OneBot
|
cfg := &s.config.Channels.OneBot
|
||||||
form := baseChannelForm("OneBot", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
|
form := baseChannelForm("OneBot", cfg.Enabled, s.makeChannelOnEnabled(&cfg.Enabled))
|
||||||
|
|
|
||||||
|
|
@ -14,23 +14,7 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
func (s *appState) modelMenu() tview.Primitive {
|
func (s *appState) modelMenu() tview.Primitive {
|
||||||
items := make([]MenuItem, 0, 2+len(s.config.ModelList))
|
items := make([]MenuItem, 0, 1+len(s.config.ModelList))
|
||||||
items = append(items,
|
|
||||||
MenuItem{Label: "Back", Description: "Return to main menu", Action: func() { s.pop() }},
|
|
||||||
MenuItem{
|
|
||||||
Label: "Add model",
|
|
||||||
Description: "Append a new model entry",
|
|
||||||
Action: func() {
|
|
||||||
s.addModel(
|
|
||||||
picoclawconfig.ModelConfig{ModelName: "new-model", Model: "openai/gpt-5.2"},
|
|
||||||
)
|
|
||||||
s.push(
|
|
||||||
fmt.Sprintf("model-%d", len(s.config.ModelList)-1),
|
|
||||||
s.modelForm(len(s.config.ModelList)-1),
|
|
||||||
)
|
|
||||||
},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
currentModel := strings.TrimSpace(s.config.Agents.Defaults.Model)
|
currentModel := strings.TrimSpace(s.config.Agents.Defaults.Model)
|
||||||
for i := range s.config.ModelList {
|
for i := range s.config.ModelList {
|
||||||
index := i
|
index := i
|
||||||
|
|
@ -57,6 +41,23 @@ func (s *appState) modelMenu() tview.Primitive {
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
// Add model entry appended at the end so the models map to rows 1..N
|
||||||
|
items = append(items,
|
||||||
|
MenuItem{
|
||||||
|
Label: "**Add model**",
|
||||||
|
Description: "Append a new model entry",
|
||||||
|
Action: func() {
|
||||||
|
newName := s.nextAvailableModelName("new-model")
|
||||||
|
s.addModel(
|
||||||
|
picoclawconfig.ModelConfig{ModelName: newName, Model: "openai/gpt-5.2"},
|
||||||
|
)
|
||||||
|
s.push(
|
||||||
|
fmt.Sprintf("model-%d", len(s.config.ModelList)-1),
|
||||||
|
s.modelForm(len(s.config.ModelList)-1),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
menu := NewMenu("Models", items)
|
menu := NewMenu("Models", items)
|
||||||
menu.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
menu.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
|
@ -64,14 +65,11 @@ func (s *appState) modelMenu() tview.Primitive {
|
||||||
s.pop()
|
s.pop()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if event.Rune() == 'q' {
|
|
||||||
s.pop()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if event.Rune() == ' ' {
|
if event.Rune() == ' ' {
|
||||||
row, _ := menu.GetSelection()
|
row, _ := menu.GetSelection()
|
||||||
if row > 0 && row <= len(s.config.ModelList) {
|
if row >= 0 && row < len(s.config.ModelList) {
|
||||||
model := s.config.ModelList[row-1]
|
model := s.config.ModelList[row]
|
||||||
if !isModelValid(model) {
|
if !isModelValid(model) {
|
||||||
s.showMessage(
|
s.showMessage(
|
||||||
"Invalid model",
|
"Invalid model",
|
||||||
|
|
@ -95,12 +93,23 @@ func (s *appState) modelForm(index int) tview.Primitive {
|
||||||
model := &s.config.ModelList[index]
|
model := &s.config.ModelList[index]
|
||||||
form := tview.NewForm()
|
form := tview.NewForm()
|
||||||
form.SetBorder(true).SetTitle(fmt.Sprintf("Model: %s", model.ModelName))
|
form.SetBorder(true).SetTitle(fmt.Sprintf("Model: %s", model.ModelName))
|
||||||
form.SetButtonBackgroundColor(tcell.NewRGBColor(80, 250, 123))
|
|
||||||
form.SetButtonTextColor(tcell.NewRGBColor(12, 13, 22))
|
|
||||||
|
|
||||||
addInput(form, "Model Name", model.ModelName, func(value string) {
|
addInput(form, "Model Name", model.ModelName, func(value string) {
|
||||||
|
if value == "" {
|
||||||
|
s.showMessage("Invalid model name", "Model Name cannot be empty")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.modelNameExists(value, index) {
|
||||||
|
s.showMessage("Duplicate model name", fmt.Sprintf("Model Name '%s' already exists", value))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
oldName := model.ModelName
|
||||||
model.ModelName = value
|
model.ModelName = value
|
||||||
|
if s.config.Agents.Defaults.Model == oldName {
|
||||||
|
s.config.Agents.Defaults.Model = value
|
||||||
|
}
|
||||||
s.dirty = true
|
s.dirty = true
|
||||||
|
form.SetTitle(fmt.Sprintf("Model: %s", model.ModelName))
|
||||||
refreshMainMenuIfPresent(s)
|
refreshMainMenuIfPresent(s)
|
||||||
if menu, ok := s.menus["model"]; ok {
|
if menu, ok := s.menus["model"]; ok {
|
||||||
refreshModelMenuFromState(menu, s)
|
refreshModelMenuFromState(menu, s)
|
||||||
|
|
@ -158,7 +167,21 @@ func (s *appState) modelForm(index int) tview.Primitive {
|
||||||
})
|
})
|
||||||
|
|
||||||
form.AddButton("Delete", func() {
|
form.AddButton("Delete", func() {
|
||||||
s.deleteModel(index)
|
pageName := "confirm-delete-model"
|
||||||
|
if s.pages.HasPage(pageName) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
modal := tview.NewModal().
|
||||||
|
SetText("Are you sure you want to delete this model?").
|
||||||
|
AddButtons([]string{"Cancel", "Delete"}).
|
||||||
|
SetDoneFunc(func(buttonIndex int, buttonLabel string) {
|
||||||
|
s.pages.RemovePage(pageName)
|
||||||
|
if buttonLabel == "Delete" {
|
||||||
|
s.deleteModel(index)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
modal.SetTitle("Confirm Delete").SetBorder(true)
|
||||||
|
s.pages.AddPage(pageName, modal, true, true)
|
||||||
})
|
})
|
||||||
form.AddButton("Test", func() {
|
form.AddButton("Test", func() {
|
||||||
s.testModel(model)
|
s.testModel(model)
|
||||||
|
|
@ -215,7 +238,7 @@ func modelStatusColor(valid bool, selected bool) *tcell.Color {
|
||||||
|
|
||||||
func refreshModelMenu(menu *Menu, currentModel string, models []picoclawconfig.ModelConfig) {
|
func refreshModelMenu(menu *Menu, currentModel string, models []picoclawconfig.ModelConfig) {
|
||||||
for i, model := range models {
|
for i, model := range models {
|
||||||
row := i + 1
|
row := i
|
||||||
label := fmt.Sprintf("%s (%s)", model.ModelName, model.Model)
|
label := fmt.Sprintf("%s (%s)", model.ModelName, model.Model)
|
||||||
isValid := isModelValid(model)
|
isValid := isModelValid(model)
|
||||||
if model.ModelName == currentModel && currentModel != "" {
|
if model.ModelName == currentModel && currentModel != "" {
|
||||||
|
|
@ -234,23 +257,7 @@ func refreshModelMenu(menu *Menu, currentModel string, models []picoclawconfig.M
|
||||||
}
|
}
|
||||||
|
|
||||||
func refreshModelMenuFromState(menu *Menu, s *appState) {
|
func refreshModelMenuFromState(menu *Menu, s *appState) {
|
||||||
items := make([]MenuItem, 0, 2+len(s.config.ModelList))
|
items := make([]MenuItem, 0, 1+len(s.config.ModelList))
|
||||||
items = append(items,
|
|
||||||
MenuItem{Label: "Back", Description: "Return to main menu", Action: func() { s.pop() }},
|
|
||||||
MenuItem{
|
|
||||||
Label: "Add model",
|
|
||||||
Description: "Append a new model entry",
|
|
||||||
Action: func() {
|
|
||||||
s.addModel(
|
|
||||||
picoclawconfig.ModelConfig{ModelName: "new-model", Model: "openai/gpt-5.2"},
|
|
||||||
)
|
|
||||||
s.push(
|
|
||||||
fmt.Sprintf("model-%d", len(s.config.ModelList)-1),
|
|
||||||
s.modelForm(len(s.config.ModelList)-1),
|
|
||||||
)
|
|
||||||
},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
currentModel := strings.TrimSpace(s.config.Agents.Defaults.Model)
|
currentModel := strings.TrimSpace(s.config.Agents.Defaults.Model)
|
||||||
for i := range s.config.ModelList {
|
for i := range s.config.ModelList {
|
||||||
index := i
|
index := i
|
||||||
|
|
@ -277,6 +284,19 @@ func refreshModelMenuFromState(menu *Menu, s *appState) {
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
items = append(items,
|
||||||
|
MenuItem{
|
||||||
|
Label: "**Add Model**",
|
||||||
|
Description: "Append a new model entry",
|
||||||
|
Action: func() {
|
||||||
|
newName := s.nextAvailableModelName("new-model")
|
||||||
|
s.addModel(
|
||||||
|
picoclawconfig.ModelConfig{ModelName: newName, Model: "openai/gpt-5.2"},
|
||||||
|
)
|
||||||
|
s.push(fmt.Sprintf("model-%d", len(s.config.ModelList)-1), s.modelForm(len(s.config.ModelList)-1))
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
menu.applyItems(items)
|
menu.applyItems(items)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -287,6 +307,38 @@ func isModelValid(model picoclawconfig.ModelConfig) bool {
|
||||||
return hasKey && hasModel
|
return hasKey && hasModel
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *appState) modelNameExists(name string, excludeIndex int) bool {
|
||||||
|
target := strings.TrimSpace(name)
|
||||||
|
if target == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for i := range s.config.ModelList {
|
||||||
|
if i == excludeIndex {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(s.config.ModelList[i].ModelName) == target {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *appState) nextAvailableModelName(base string) string {
|
||||||
|
name := strings.TrimSpace(base)
|
||||||
|
if name == "" {
|
||||||
|
name = "new-model"
|
||||||
|
}
|
||||||
|
if !s.modelNameExists(name, -1) {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
for i := 2; ; i++ {
|
||||||
|
candidate := fmt.Sprintf("%s-%d", name, i)
|
||||||
|
if !s.modelNameExists(candidate, -1) {
|
||||||
|
return candidate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (s *appState) testModel(model *picoclawconfig.ModelConfig) {
|
func (s *appState) testModel(model *picoclawconfig.ModelConfig) {
|
||||||
if model == nil {
|
if model == nil {
|
||||||
return
|
return
|
||||||
|
|
|
||||||
|
|
@ -41,3 +41,15 @@ func bannerView() *tview.TextView {
|
||||||
text.SetBorder(false)
|
text.SetBorder(false)
|
||||||
return text
|
return text
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const footerText = "Esc: Back/Exit | Enter: Enter | ←↓↑→ : Move | Space: Select | Tab/Shift+Tab: Switch"
|
||||||
|
|
||||||
|
func footerView() *tview.TextView {
|
||||||
|
text := tview.NewTextView()
|
||||||
|
text.SetTextAlign(tview.AlignCenter)
|
||||||
|
text.SetText(footerText)
|
||||||
|
text.SetBackgroundColor(tview.Styles.MoreContrastBackgroundColor)
|
||||||
|
text.SetTextColor(tview.Styles.PrimaryTextColor)
|
||||||
|
text.SetBorder(false)
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,290 +0,0 @@
|
||||||
# PicoClaw Launcher
|
|
||||||
|
|
||||||
> [!WARNING]
|
|
||||||
> This project is a temporary solution and will be refactored in the future to provide a complete web service. Therefore, the APIs in this directory are not stable.
|
|
||||||
|
|
||||||
A standalone launcher for PicoClaw, providing visual JSON editing and OAuth provider authentication management.
|
|
||||||
|
|
||||||
## Features
|
|
||||||
|
|
||||||
- 📝 **Config Editor** — Sidebar-based settings UI with model management, channel configuration forms, and a raw JSON editor
|
|
||||||
- 🤖 **Model Management** — Model card grid with availability status (grayed out without API key), primary model selection, add/edit/delete with required/optional field separation
|
|
||||||
- 📡 **Channel Configuration** — Form-based settings for 12 channel types (Telegram, Discord, Slack, WeCom, DingTalk, Feishu, LINE, WhatsApp, QQ, OneBot, MaixCAM, etc.) with documentation links
|
|
||||||
- 🔐 **Provider Auth** — Login to OpenAI (Device Code), Anthropic (API Token), Google Antigravity (Browser OAuth)
|
|
||||||
- 🌐 **Embedded Frontend** — Compiles to a single binary with no external dependencies
|
|
||||||
- 🌍 **i18n** — Chinese/English language switching with browser auto-detection
|
|
||||||
- 🎨 **Theme** — Light / Dark / System theme toggle with localStorage persistence
|
|
||||||
|
|
||||||
## Quick Start
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Build
|
|
||||||
go build -o picoclaw-launcher ./cmd/picoclaw-launcher/
|
|
||||||
|
|
||||||
# Run with default config path (~/.picoclaw/config.json)
|
|
||||||
./picoclaw-launcher
|
|
||||||
|
|
||||||
# Specify a config file
|
|
||||||
./picoclaw-launcher ./config.json
|
|
||||||
|
|
||||||
# Allow LAN access
|
|
||||||
./picoclaw-launcher -public
|
|
||||||
```
|
|
||||||
|
|
||||||
Open `http://localhost:18800` in your browser.
|
|
||||||
|
|
||||||
## CLI Options
|
|
||||||
|
|
||||||
```
|
|
||||||
Usage: picoclaw-config [options] [config.json]
|
|
||||||
|
|
||||||
Arguments:
|
|
||||||
config.json Path to the configuration file (default: ~/.picoclaw/config.json)
|
|
||||||
|
|
||||||
Options:
|
|
||||||
-public Listen on all interfaces (0.0.0.0), allowing access from other devices
|
|
||||||
```
|
|
||||||
|
|
||||||
## API Reference
|
|
||||||
|
|
||||||
Base URL: `http://localhost:18800`
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Static Files
|
|
||||||
|
|
||||||
#### GET /
|
|
||||||
|
|
||||||
Serves the embedded frontend (`index.html`).
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Config API
|
|
||||||
|
|
||||||
#### GET /api/config
|
|
||||||
|
|
||||||
Reads the current configuration file.
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"config": { ... },
|
|
||||||
"path": "/Users/xiao/.picoclaw/config.json"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### PUT /api/config
|
|
||||||
|
|
||||||
Saves the configuration. The request body must be a complete Config JSON object.
|
|
||||||
|
|
||||||
**Request Body** — `application/json`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"agents": { "defaults": { "model_name": "gpt-5.2" } },
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.2",
|
|
||||||
"model": "openai/gpt-5.2",
|
|
||||||
"auth_method": "oauth"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "status": "ok" }
|
|
||||||
```
|
|
||||||
|
|
||||||
**Error** `400 Bad Request` — Invalid JSON
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Auth API
|
|
||||||
|
|
||||||
#### GET /api/auth/status
|
|
||||||
|
|
||||||
Returns the authentication status of all providers and any in-progress device code login.
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"providers": [
|
|
||||||
{
|
|
||||||
"provider": "openai",
|
|
||||||
"auth_method": "oauth",
|
|
||||||
"status": "active",
|
|
||||||
"account_id": "user-xxx",
|
|
||||||
"expires_at": "2026-03-01T00:00:00Z"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"pending_device": {
|
|
||||||
"provider": "openai",
|
|
||||||
"status": "pending",
|
|
||||||
"device_url": "https://auth.openai.com/activate",
|
|
||||||
"user_code": "ABCD-1234"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
`status` values: `active` | `expired` | `needs_refresh`
|
|
||||||
|
|
||||||
`pending_device` is only present when a device code login is in progress.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/auth/login
|
|
||||||
|
|
||||||
Initiates a provider login.
|
|
||||||
|
|
||||||
**Request Body** — `application/json`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "provider": "openai" }
|
|
||||||
```
|
|
||||||
|
|
||||||
Supported `provider` values: `openai` | `anthropic` | `google-antigravity`
|
|
||||||
|
|
||||||
##### OpenAI (Device Code Flow)
|
|
||||||
|
|
||||||
Returns device code info. The server polls for completion in the background.
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "pending",
|
|
||||||
"device_url": "https://auth.openai.com/activate",
|
|
||||||
"user_code": "ABCD-1234",
|
|
||||||
"message": "Open the URL and enter the code to authenticate."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
The user opens `device_url` in a browser and enters `user_code`. Once authenticated, `GET /api/auth/status` will show `pending_device.status` as `success`.
|
|
||||||
|
|
||||||
##### Anthropic (API Token)
|
|
||||||
|
|
||||||
Requires a `token` field in the request:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "provider": "anthropic", "token": "sk-ant-xxx" }
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response:**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "status": "success", "message": "Anthropic token saved" }
|
|
||||||
```
|
|
||||||
|
|
||||||
##### Google Antigravity (Browser OAuth)
|
|
||||||
|
|
||||||
Returns an authorization URL for the frontend to open in a new tab:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "redirect",
|
|
||||||
"auth_url": "https://accounts.google.com/o/oauth2/auth?...",
|
|
||||||
"message": "Open the URL to authenticate with Google."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
After authentication, Google redirects to `GET /auth/callback`, which saves the credentials and redirects back to the picoclaw-config UI.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/auth/logout
|
|
||||||
|
|
||||||
Logs out from a provider.
|
|
||||||
|
|
||||||
**Request Body** — `application/json`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "provider": "openai" }
|
|
||||||
```
|
|
||||||
|
|
||||||
Omit or leave `provider` empty to log out from all providers.
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "status": "ok" }
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### GET /auth/callback
|
|
||||||
|
|
||||||
OAuth browser callback endpoint (used by Google Antigravity). Called by the OAuth provider's redirect — **not invoked directly by the frontend**.
|
|
||||||
|
|
||||||
**Query Parameters:**
|
|
||||||
- `state` — OAuth state for CSRF validation
|
|
||||||
- `code` — Authorization code
|
|
||||||
|
|
||||||
On success, redirects to `/#auth`.
|
|
||||||
|
|
||||||
|
|
||||||
### Process API
|
|
||||||
|
|
||||||
#### GET /api/process/status
|
|
||||||
|
|
||||||
Gets the running status of the `picoclaw gateway` process.
|
|
||||||
|
|
||||||
**Response** `200 OK` (Running)
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"process_status": "running",
|
|
||||||
"status": "ok",
|
|
||||||
"uptime": "1.010814s"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response** `200 OK` (Stopped)
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"process_status": "stopped",
|
|
||||||
"error": "Get \"http://localhost:18790/health\": dial tcp [::1]:18790: connect: connection refused"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/process/start
|
|
||||||
|
|
||||||
Starts the `picoclaw gateway` process in the background.
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "ok",
|
|
||||||
"pid": 12345
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/process/stop
|
|
||||||
|
|
||||||
Stops the running `picoclaw gateway` process.
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "ok"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Testing
|
|
||||||
|
|
||||||
```bash
|
|
||||||
go test -v ./cmd/picoclaw-launcher/
|
|
||||||
```
|
|
||||||
|
|
@ -1,287 +0,0 @@
|
||||||
# PicoClaw Launcher
|
|
||||||
|
|
||||||
> [!WARNING]
|
|
||||||
> 该项目属于临时解决方案,后续会重构并提供完整的 Web 服务,因此该目录下的接口并不稳定。
|
|
||||||
|
|
||||||
PicoClaw 的独立启动器,提供可视化 JSON 配置编辑和 OAuth Provider 认证管理。
|
|
||||||
|
|
||||||
## 功能
|
|
||||||
|
|
||||||
- 📝 **配置编辑** — 侧边栏式设置 UI,支持模型管理、通道配置表单和原始 JSON 编辑器
|
|
||||||
- 🤖 **模型管理** — 模型卡片网格,可用性状态显示(无 API Key 时灰色),主模型选择,增删改查,必填/选填字段分离
|
|
||||||
- 📡 **通道配置** — 12 种通道类型(Telegram、Discord、Slack、企业微信、钉钉、飞书、LINE、WhatsApp、QQ、OneBot、MaixCAM 等)的表单化配置,附带文档链接
|
|
||||||
- 🔐 **Provider 认证** — 支持 OpenAI (Device Code)、Anthropic (API Token)、Google Antigravity (Browser OAuth) 登录
|
|
||||||
- 🌐 **嵌入式前端** — 编译为单一二进制文件,无需额外依赖
|
|
||||||
- 🌍 **国际化** — 中英文切换,首次访问自动检测浏览器语言
|
|
||||||
- 🎨 **主题** — 亮色 / 暗色 / 跟随系统,偏好保存在 localStorage
|
|
||||||
|
|
||||||
## 快速开始
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# 编译
|
|
||||||
go build -o picoclaw-launcher ./cmd/picoclaw-launcher/
|
|
||||||
|
|
||||||
# 运行(使用默认配置路径 ~/.picoclaw/config.json)
|
|
||||||
./picoclaw-launcher
|
|
||||||
|
|
||||||
# 指定配置文件
|
|
||||||
./picoclaw-launcher ./config.json
|
|
||||||
|
|
||||||
# 允许局域网访问
|
|
||||||
./picoclaw-launcher -public
|
|
||||||
```
|
|
||||||
|
|
||||||
启动后在浏览器中打开 `http://localhost:18800`。
|
|
||||||
|
|
||||||
## 命令行参数
|
|
||||||
|
|
||||||
```
|
|
||||||
Usage: picoclaw-launcher [options] [config.json]
|
|
||||||
|
|
||||||
Arguments:
|
|
||||||
config.json 配置文件路径(默认: ~/.picoclaw/config.json)
|
|
||||||
|
|
||||||
Options:
|
|
||||||
-public 监听所有网络接口(0.0.0.0),允许局域网设备访问
|
|
||||||
```
|
|
||||||
|
|
||||||
## API 文档
|
|
||||||
|
|
||||||
Base URL: `http://localhost:18800`
|
|
||||||
|
|
||||||
### 静态文件
|
|
||||||
|
|
||||||
#### GET /
|
|
||||||
|
|
||||||
提供嵌入式前端页面(`index.html`)。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Config API
|
|
||||||
|
|
||||||
#### GET /api/config
|
|
||||||
|
|
||||||
读取当前配置文件内容。
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"config": { ... },
|
|
||||||
"path": "/Users/xiao/.picoclaw/config.json"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### PUT /api/config
|
|
||||||
|
|
||||||
保存配置。请求体为完整的 Config JSON。
|
|
||||||
|
|
||||||
**Request Body** — `application/json`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"agents": { "defaults": { "model_name": "gpt-5.2" } },
|
|
||||||
"model_list": [
|
|
||||||
{
|
|
||||||
"model_name": "gpt-5.2",
|
|
||||||
"model": "openai/gpt-5.2",
|
|
||||||
"auth_method": "oauth"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "status": "ok" }
|
|
||||||
```
|
|
||||||
|
|
||||||
**Error** `400 Bad Request` — 无效 JSON
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Auth API
|
|
||||||
|
|
||||||
#### GET /api/auth/status
|
|
||||||
|
|
||||||
获取所有 Provider 的认证状态和进行中的 Device Code 登录信息。
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"providers": [
|
|
||||||
{
|
|
||||||
"provider": "openai",
|
|
||||||
"auth_method": "oauth",
|
|
||||||
"status": "active",
|
|
||||||
"account_id": "user-xxx",
|
|
||||||
"expires_at": "2026-03-01T00:00:00Z"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"pending_device": {
|
|
||||||
"provider": "openai",
|
|
||||||
"status": "pending",
|
|
||||||
"device_url": "https://auth.openai.com/activate",
|
|
||||||
"user_code": "ABCD-1234"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
`status` 可选值: `active` | `expired` | `needs_refresh`
|
|
||||||
|
|
||||||
`pending_device` 仅在有进行中的 Device Code 登录时返回。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/auth/login
|
|
||||||
|
|
||||||
发起 Provider 登录。
|
|
||||||
|
|
||||||
**Request Body** — `application/json`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "provider": "openai" }
|
|
||||||
```
|
|
||||||
|
|
||||||
支持的 `provider` 值: `openai` | `anthropic` | `google-antigravity`
|
|
||||||
|
|
||||||
##### OpenAI (Device Code Flow)
|
|
||||||
|
|
||||||
返回 Device Code 信息,后台自动轮询认证结果:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "pending",
|
|
||||||
"device_url": "https://auth.openai.com/activate",
|
|
||||||
"user_code": "ABCD-1234",
|
|
||||||
"message": "Open the URL and enter the code to authenticate."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
用户在浏览器中打开 `device_url` 并输入 `user_code`。认证完成后通过 `GET /api/auth/status` 的 `pending_device.status` 变为 `success` 通知前端。
|
|
||||||
|
|
||||||
##### Anthropic (API Token)
|
|
||||||
|
|
||||||
需在请求中附带 token:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "provider": "anthropic", "token": "sk-ant-xxx" }
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response:**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "status": "success", "message": "Anthropic token saved" }
|
|
||||||
```
|
|
||||||
|
|
||||||
##### Google Antigravity (Browser OAuth)
|
|
||||||
|
|
||||||
返回授权 URL,前端打开新标签页:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "redirect",
|
|
||||||
"auth_url": "https://accounts.google.com/o/oauth2/auth?...",
|
|
||||||
"message": "Open the URL to authenticate with Google."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
认证完成后 Google 回调至 `GET /auth/callback`,自动保存凭据并重定向回 picoclaw-config 页面。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/auth/logout
|
|
||||||
|
|
||||||
登出 Provider。
|
|
||||||
|
|
||||||
**Request Body** — `application/json`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "provider": "openai" }
|
|
||||||
```
|
|
||||||
|
|
||||||
传空字符串或省略 `provider` 则登出所有 Provider。
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{ "status": "ok" }
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### GET /auth/callback
|
|
||||||
|
|
||||||
OAuth Browser 回调端点(Google Antigravity 专用),由 OAuth Provider 重定向调用,**非前端直接使用**。
|
|
||||||
|
|
||||||
**Query Parameters:**
|
|
||||||
- `state` — OAuth state 校验
|
|
||||||
- `code` — 授权码
|
|
||||||
|
|
||||||
认证成功后重定向到 `/#auth`。
|
|
||||||
|
|
||||||
### Process API
|
|
||||||
|
|
||||||
#### GET /api/process/status
|
|
||||||
|
|
||||||
获取 `picoclaw gateway` 进程的运行状态。
|
|
||||||
|
|
||||||
**Response** `200 OK` (运行中)
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"process_status": "running",
|
|
||||||
"status": "ok",
|
|
||||||
"uptime": "1.010814s"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response** `200 OK` (未运行)
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"process_status": "stopped",
|
|
||||||
"error": "Get \"http://localhost:18790/health\": dial tcp [::1]:18790: connect: connection refused"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/process/start
|
|
||||||
|
|
||||||
在后台启动 `picoclaw gateway` 进程。
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "ok",
|
|
||||||
"pid": 12345
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
#### POST /api/process/stop
|
|
||||||
|
|
||||||
停止正在运行的 `picoclaw gateway` 进程。
|
|
||||||
|
|
||||||
**Response** `200 OK`
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"status": "ok"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 测试
|
|
||||||
|
|
||||||
```bash
|
|
||||||
go test -v ./cmd/picoclaw-launcher/
|
|
||||||
```
|
|
||||||
|
|
@ -1,147 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"log"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// updateConfigAfterLogin updates config.json after a successful provider login.
|
|
||||||
func updateConfigAfterLogin(configPath, provider string, cred *auth.AuthCredential) {
|
|
||||||
cfg, err := config.LoadConfig(configPath)
|
|
||||||
if err != nil {
|
|
||||||
log.Printf("Warning: could not load config to update auth_method: %v", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
switch provider {
|
|
||||||
case "openai":
|
|
||||||
cfg.Providers.OpenAI.AuthMethod = "oauth"
|
|
||||||
found := false
|
|
||||||
for i := range cfg.ModelList {
|
|
||||||
if isOpenAIModel(cfg.ModelList[i].Model) {
|
|
||||||
cfg.ModelList[i].AuthMethod = "oauth"
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
cfg.ModelList = append(cfg.ModelList, config.ModelConfig{
|
|
||||||
ModelName: "gpt-5.2",
|
|
||||||
Model: "openai/gpt-5.2",
|
|
||||||
AuthMethod: "oauth",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
cfg.Agents.Defaults.ModelName = "gpt-5.2"
|
|
||||||
|
|
||||||
case "anthropic":
|
|
||||||
cfg.Providers.Anthropic.AuthMethod = "token"
|
|
||||||
found := false
|
|
||||||
for i := range cfg.ModelList {
|
|
||||||
if isAnthropicModel(cfg.ModelList[i].Model) {
|
|
||||||
cfg.ModelList[i].AuthMethod = "token"
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
cfg.ModelList = append(cfg.ModelList, config.ModelConfig{
|
|
||||||
ModelName: "claude-sonnet-4.6",
|
|
||||||
Model: "anthropic/claude-sonnet-4.6",
|
|
||||||
AuthMethod: "token",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
cfg.Agents.Defaults.ModelName = "claude-sonnet-4.6"
|
|
||||||
|
|
||||||
case "google-antigravity":
|
|
||||||
cfg.Providers.Antigravity.AuthMethod = "oauth"
|
|
||||||
found := false
|
|
||||||
for i := range cfg.ModelList {
|
|
||||||
if isAntigravityModel(cfg.ModelList[i].Model) {
|
|
||||||
cfg.ModelList[i].AuthMethod = "oauth"
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
cfg.ModelList = append(cfg.ModelList, config.ModelConfig{
|
|
||||||
ModelName: "gemini-flash",
|
|
||||||
Model: "antigravity/gemini-3-flash",
|
|
||||||
AuthMethod: "oauth",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
cfg.Agents.Defaults.ModelName = "gemini-flash"
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := config.SaveConfig(configPath, cfg); err != nil {
|
|
||||||
log.Printf("Warning: could not update config: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// clearAuthMethodInConfig clears auth_method for a specific provider in config.json.
|
|
||||||
func clearAuthMethodInConfig(configPath, provider string) {
|
|
||||||
cfg, err := config.LoadConfig(configPath)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range cfg.ModelList {
|
|
||||||
switch provider {
|
|
||||||
case "openai":
|
|
||||||
if isOpenAIModel(cfg.ModelList[i].Model) {
|
|
||||||
cfg.ModelList[i].AuthMethod = ""
|
|
||||||
}
|
|
||||||
case "anthropic":
|
|
||||||
if isAnthropicModel(cfg.ModelList[i].Model) {
|
|
||||||
cfg.ModelList[i].AuthMethod = ""
|
|
||||||
}
|
|
||||||
case "google-antigravity", "antigravity":
|
|
||||||
if isAntigravityModel(cfg.ModelList[i].Model) {
|
|
||||||
cfg.ModelList[i].AuthMethod = ""
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
switch provider {
|
|
||||||
case "openai":
|
|
||||||
cfg.Providers.OpenAI.AuthMethod = ""
|
|
||||||
case "anthropic":
|
|
||||||
cfg.Providers.Anthropic.AuthMethod = ""
|
|
||||||
case "google-antigravity", "antigravity":
|
|
||||||
cfg.Providers.Antigravity.AuthMethod = ""
|
|
||||||
}
|
|
||||||
|
|
||||||
config.SaveConfig(configPath, cfg)
|
|
||||||
}
|
|
||||||
|
|
||||||
// clearAllAuthMethodsInConfig clears auth_method for all providers in config.json.
|
|
||||||
func clearAllAuthMethodsInConfig(configPath string) {
|
|
||||||
cfg, err := config.LoadConfig(configPath)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
for i := range cfg.ModelList {
|
|
||||||
cfg.ModelList[i].AuthMethod = ""
|
|
||||||
}
|
|
||||||
cfg.Providers.OpenAI.AuthMethod = ""
|
|
||||||
cfg.Providers.Anthropic.AuthMethod = ""
|
|
||||||
cfg.Providers.Antigravity.AuthMethod = ""
|
|
||||||
config.SaveConfig(configPath, cfg)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ── Model identification helpers ─────────────────────────────────
|
|
||||||
|
|
||||||
func isOpenAIModel(model string) bool {
|
|
||||||
return model == "openai" || strings.HasPrefix(model, "openai/")
|
|
||||||
}
|
|
||||||
|
|
||||||
func isAnthropicModel(model string) bool {
|
|
||||||
return model == "anthropic" || strings.HasPrefix(model, "anthropic/")
|
|
||||||
}
|
|
||||||
|
|
||||||
func isAntigravityModel(model string) bool {
|
|
||||||
return model == "antigravity" || model == "google-antigravity" ||
|
|
||||||
strings.HasPrefix(model, "antigravity/") || strings.HasPrefix(model, "google-antigravity/")
|
|
||||||
}
|
|
||||||
|
|
@ -1,222 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ── Model identification helpers ─────────────────────────────────
|
|
||||||
|
|
||||||
func TestIsOpenAIModel(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
model string
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{"openai", true},
|
|
||||||
{"openai/gpt-4o", true},
|
|
||||||
{"openai/gpt-5.2", true},
|
|
||||||
{"anthropic", false},
|
|
||||||
{"anthropic/claude-sonnet-4.6", false},
|
|
||||||
{"openai-compatible", false},
|
|
||||||
{"", false},
|
|
||||||
}
|
|
||||||
for _, tt := range tests {
|
|
||||||
if got := isOpenAIModel(tt.model); got != tt.want {
|
|
||||||
t.Errorf("isOpenAIModel(%q) = %v, want %v", tt.model, got, tt.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIsAnthropicModel(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
model string
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{"anthropic", true},
|
|
||||||
{"anthropic/claude-sonnet-4.6", true},
|
|
||||||
{"openai", false},
|
|
||||||
{"openai/gpt-4o", false},
|
|
||||||
{"", false},
|
|
||||||
}
|
|
||||||
for _, tt := range tests {
|
|
||||||
if got := isAnthropicModel(tt.model); got != tt.want {
|
|
||||||
t.Errorf("isAnthropicModel(%q) = %v, want %v", tt.model, got, tt.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIsAntigravityModel(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
model string
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{"antigravity", true},
|
|
||||||
{"google-antigravity", true},
|
|
||||||
{"antigravity/gemini-3-flash", true},
|
|
||||||
{"google-antigravity/gemini-3-flash", true},
|
|
||||||
{"openai", false},
|
|
||||||
{"antigravity-custom", false},
|
|
||||||
{"", false},
|
|
||||||
}
|
|
||||||
for _, tt := range tests {
|
|
||||||
if got := isAntigravityModel(tt.model); got != tt.want {
|
|
||||||
t.Errorf("isAntigravityModel(%q) = %v, want %v", tt.model, got, tt.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ── Config update helpers ────────────────────────────────────────
|
|
||||||
|
|
||||||
func writeTempConfigViaSave(t *testing.T, cfg *config.Config) string {
|
|
||||||
t.Helper()
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "config.json")
|
|
||||||
if err := config.SaveConfig(path, cfg); err != nil {
|
|
||||||
t.Fatalf("save config: %v", err)
|
|
||||||
}
|
|
||||||
return path
|
|
||||||
}
|
|
||||||
|
|
||||||
func loadTempConfig(t *testing.T, path string) *config.Config {
|
|
||||||
t.Helper()
|
|
||||||
cfg, err := config.LoadConfig(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load config: %v", err)
|
|
||||||
}
|
|
||||||
return cfg
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUpdateConfigAfterLogin_OpenAI_ExistingModel(t *testing.T) {
|
|
||||||
cfg := &config.Config{
|
|
||||||
ModelList: []config.ModelConfig{
|
|
||||||
{ModelName: "gpt-4o", Model: "openai/gpt-4o"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
path := writeTempConfigViaSave(t, cfg)
|
|
||||||
|
|
||||||
cred := &auth.AuthCredential{AuthMethod: "oauth"}
|
|
||||||
updateConfigAfterLogin(path, "openai", cred)
|
|
||||||
|
|
||||||
result := loadTempConfig(t, path)
|
|
||||||
|
|
||||||
// Model-level auth_method persists through serialization
|
|
||||||
if len(result.ModelList) != 1 {
|
|
||||||
t.Fatalf("expected 1 model, got %d", len(result.ModelList))
|
|
||||||
}
|
|
||||||
if result.ModelList[0].AuthMethod != "oauth" {
|
|
||||||
t.Errorf("expected model auth_method=oauth, got %q", result.ModelList[0].AuthMethod)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUpdateConfigAfterLogin_OpenAI_NoExistingModel(t *testing.T) {
|
|
||||||
cfg := &config.Config{
|
|
||||||
ModelList: []config.ModelConfig{
|
|
||||||
{ModelName: "claude", Model: "anthropic/claude-sonnet-4.6"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
path := writeTempConfigViaSave(t, cfg)
|
|
||||||
|
|
||||||
cred := &auth.AuthCredential{AuthMethod: "oauth"}
|
|
||||||
updateConfigAfterLogin(path, "openai", cred)
|
|
||||||
|
|
||||||
result := loadTempConfig(t, path)
|
|
||||||
|
|
||||||
if len(result.ModelList) != 2 {
|
|
||||||
t.Fatalf("expected 2 models (original + added), got %d", len(result.ModelList))
|
|
||||||
}
|
|
||||||
if result.ModelList[1].Model != "openai/gpt-5.2" {
|
|
||||||
t.Errorf("expected added model openai/gpt-5.2, got %q", result.ModelList[1].Model)
|
|
||||||
}
|
|
||||||
if result.Agents.Defaults.ModelName != "gpt-5.2" {
|
|
||||||
t.Errorf("expected default model_name=gpt-5.2, got %q", result.Agents.Defaults.ModelName)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUpdateConfigAfterLogin_Anthropic(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
path := writeTempConfigViaSave(t, cfg)
|
|
||||||
|
|
||||||
cred := &auth.AuthCredential{AuthMethod: "token"}
|
|
||||||
updateConfigAfterLogin(path, "anthropic", cred)
|
|
||||||
|
|
||||||
result := loadTempConfig(t, path)
|
|
||||||
|
|
||||||
// Model should be added with correct auth_method
|
|
||||||
if len(result.ModelList) != 1 {
|
|
||||||
t.Fatalf("expected 1 model added, got %d", len(result.ModelList))
|
|
||||||
}
|
|
||||||
if result.ModelList[0].Model != "anthropic/claude-sonnet-4.6" {
|
|
||||||
t.Errorf("expected model anthropic/claude-sonnet-4.6, got %q", result.ModelList[0].Model)
|
|
||||||
}
|
|
||||||
if result.ModelList[0].AuthMethod != "token" {
|
|
||||||
t.Errorf("expected model auth_method=token, got %q", result.ModelList[0].AuthMethod)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUpdateConfigAfterLogin_GoogleAntigravity(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
path := writeTempConfigViaSave(t, cfg)
|
|
||||||
|
|
||||||
cred := &auth.AuthCredential{AuthMethod: "oauth"}
|
|
||||||
updateConfigAfterLogin(path, "google-antigravity", cred)
|
|
||||||
|
|
||||||
result := loadTempConfig(t, path)
|
|
||||||
|
|
||||||
// Model should be added with correct auth_method
|
|
||||||
if len(result.ModelList) != 1 {
|
|
||||||
t.Fatalf("expected 1 model added, got %d", len(result.ModelList))
|
|
||||||
}
|
|
||||||
if result.ModelList[0].Model != "antigravity/gemini-3-flash" {
|
|
||||||
t.Errorf("expected model antigravity/gemini-3-flash, got %q", result.ModelList[0].Model)
|
|
||||||
}
|
|
||||||
if result.ModelList[0].AuthMethod != "oauth" {
|
|
||||||
t.Errorf("expected model auth_method=oauth, got %q", result.ModelList[0].AuthMethod)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClearAuthMethodInConfig(t *testing.T) {
|
|
||||||
cfg := &config.Config{
|
|
||||||
ModelList: []config.ModelConfig{
|
|
||||||
{ModelName: "gpt-4o", Model: "openai/gpt-4o", AuthMethod: "oauth"},
|
|
||||||
{ModelName: "claude", Model: "anthropic/claude-sonnet-4.6", AuthMethod: "token"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
path := writeTempConfigViaSave(t, cfg)
|
|
||||||
|
|
||||||
clearAuthMethodInConfig(path, "openai")
|
|
||||||
|
|
||||||
result := loadTempConfig(t, path)
|
|
||||||
|
|
||||||
// Openai model auth_method should be cleared
|
|
||||||
if result.ModelList[0].AuthMethod != "" {
|
|
||||||
t.Errorf("expected openai model auth_method cleared, got %q", result.ModelList[0].AuthMethod)
|
|
||||||
}
|
|
||||||
// Anthropic model should be unchanged
|
|
||||||
if result.ModelList[1].AuthMethod != "token" {
|
|
||||||
t.Errorf("expected anthropic model auth_method unchanged, got %q", result.ModelList[1].AuthMethod)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClearAllAuthMethodsInConfig(t *testing.T) {
|
|
||||||
cfg := &config.Config{
|
|
||||||
ModelList: []config.ModelConfig{
|
|
||||||
{ModelName: "gpt-4o", Model: "openai/gpt-4o", AuthMethod: "oauth"},
|
|
||||||
{ModelName: "claude", Model: "anthropic/claude-sonnet-4.6", AuthMethod: "token"},
|
|
||||||
{ModelName: "gemini", Model: "antigravity/gemini-3-flash", AuthMethod: "oauth"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
path := writeTempConfigViaSave(t, cfg)
|
|
||||||
|
|
||||||
clearAllAuthMethodsInConfig(path)
|
|
||||||
|
|
||||||
result := loadTempConfig(t, path)
|
|
||||||
|
|
||||||
for i, m := range result.ModelList {
|
|
||||||
if m.AuthMethod != "" {
|
|
||||||
t.Errorf("model[%d] auth_method not cleared, got %q", i, m.AuthMethod)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,315 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"log"
|
|
||||||
"net/http"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
|
||||||
)
|
|
||||||
|
|
||||||
// oauthSession stores in-flight OAuth state for browser-based flows.
|
|
||||||
type oauthSession struct {
|
|
||||||
Provider string
|
|
||||||
PKCE auth.PKCECodes
|
|
||||||
State string
|
|
||||||
RedirectURI string
|
|
||||||
OAuthCfg auth.OAuthProviderConfig
|
|
||||||
ConfigPath string
|
|
||||||
}
|
|
||||||
|
|
||||||
// deviceCodeSession stores in-flight device code flow state.
|
|
||||||
type deviceCodeSession struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
Provider string
|
|
||||||
Info *auth.DeviceCodeInfo
|
|
||||||
OAuthCfg auth.OAuthProviderConfig
|
|
||||||
ConfigPath string
|
|
||||||
Status string // "pending", "success", "error"
|
|
||||||
Error string
|
|
||||||
Done bool
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
oauthSessions = map[string]*oauthSession{} // keyed by state
|
|
||||||
oauthSessionsMu sync.Mutex
|
|
||||||
|
|
||||||
activeDeviceSession *deviceCodeSession
|
|
||||||
activeDeviceSessionMu sync.Mutex
|
|
||||||
)
|
|
||||||
|
|
||||||
// handleOpenAILogin starts the OpenAI device code flow and returns device code info to the frontend.
|
|
||||||
func handleOpenAILogin(w http.ResponseWriter, configPath string) {
|
|
||||||
// Check if there's already a pending device code session
|
|
||||||
activeDeviceSessionMu.Lock()
|
|
||||||
if activeDeviceSession != nil {
|
|
||||||
activeDeviceSession.mu.Lock()
|
|
||||||
if !activeDeviceSession.Done {
|
|
||||||
resp := map[string]any{
|
|
||||||
"status": "pending",
|
|
||||||
"device_url": activeDeviceSession.Info.VerifyURL,
|
|
||||||
"user_code": activeDeviceSession.Info.UserCode,
|
|
||||||
"message": "Device code flow already in progress. Enter the code in your browser.",
|
|
||||||
}
|
|
||||||
activeDeviceSession.mu.Unlock()
|
|
||||||
activeDeviceSessionMu.Unlock()
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(resp)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
activeDeviceSession.mu.Unlock()
|
|
||||||
}
|
|
||||||
activeDeviceSessionMu.Unlock()
|
|
||||||
|
|
||||||
// Request a device code
|
|
||||||
oauthCfg := auth.OpenAIOAuthConfig()
|
|
||||||
info, err := auth.RequestDeviceCode(oauthCfg)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to request device code: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
session := &deviceCodeSession{
|
|
||||||
Provider: "openai",
|
|
||||||
Info: info,
|
|
||||||
OAuthCfg: oauthCfg,
|
|
||||||
ConfigPath: configPath,
|
|
||||||
Status: "pending",
|
|
||||||
}
|
|
||||||
|
|
||||||
activeDeviceSessionMu.Lock()
|
|
||||||
activeDeviceSession = session
|
|
||||||
activeDeviceSessionMu.Unlock()
|
|
||||||
|
|
||||||
// Start background polling
|
|
||||||
go func() {
|
|
||||||
deadline := time.After(15 * time.Minute)
|
|
||||||
ticker := time.NewTicker(time.Duration(info.Interval) * time.Second)
|
|
||||||
defer ticker.Stop()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-deadline:
|
|
||||||
session.mu.Lock()
|
|
||||||
session.Status = "error"
|
|
||||||
session.Error = "Authentication timed out after 15 minutes"
|
|
||||||
session.Done = true
|
|
||||||
session.mu.Unlock()
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
cred, err := auth.PollDeviceCodeOnce(oauthCfg, info.DeviceAuthID, info.UserCode)
|
|
||||||
if err != nil {
|
|
||||||
continue // Still pending
|
|
||||||
}
|
|
||||||
if cred != nil {
|
|
||||||
if saveErr := auth.SetCredential("openai", cred); saveErr != nil {
|
|
||||||
session.mu.Lock()
|
|
||||||
session.Status = "error"
|
|
||||||
session.Error = saveErr.Error()
|
|
||||||
session.Done = true
|
|
||||||
session.mu.Unlock()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
updateConfigAfterLogin(configPath, "openai", cred)
|
|
||||||
session.mu.Lock()
|
|
||||||
session.Status = "success"
|
|
||||||
session.Done = true
|
|
||||||
session.mu.Unlock()
|
|
||||||
log.Printf("OpenAI device code login successful (account: %s)", cred.AccountID)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Return device code info to frontend
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
|
||||||
"status": "pending",
|
|
||||||
"device_url": info.VerifyURL,
|
|
||||||
"user_code": info.UserCode,
|
|
||||||
"message": "Open the URL and enter the code to authenticate.",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleAnthropicLogin saves a pasted API token for Anthropic.
|
|
||||||
func handleAnthropicLogin(w http.ResponseWriter, token, configPath string) {
|
|
||||||
if token == "" {
|
|
||||||
http.Error(w, "Token is required for Anthropic login", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
cred := &auth.AuthCredential{
|
|
||||||
AccessToken: token,
|
|
||||||
Provider: "anthropic",
|
|
||||||
AuthMethod: "token",
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := auth.SetCredential("anthropic", cred); err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to save credentials: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
updateConfigAfterLogin(configPath, "anthropic", cred)
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{
|
|
||||||
"status": "success",
|
|
||||||
"message": "Anthropic token saved",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleGoogleAntigravityLogin generates a PKCE + auth URL and returns it to the frontend.
|
|
||||||
func handleGoogleAntigravityLogin(w http.ResponseWriter, r *http.Request, configPath string) {
|
|
||||||
oauthCfg := auth.GoogleAntigravityOAuthConfig()
|
|
||||||
|
|
||||||
pkce, err := auth.GeneratePKCE()
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to generate PKCE: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
state, err := auth.GenerateState()
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to generate state: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build redirect URI pointing to picoclaw-launcher's own callback
|
|
||||||
scheme := "http"
|
|
||||||
redirectURI := fmt.Sprintf("%s://%s/auth/callback", scheme, r.Host)
|
|
||||||
|
|
||||||
authURL := auth.BuildAuthorizeURL(oauthCfg, pkce, state, redirectURI)
|
|
||||||
|
|
||||||
// Store session for callback
|
|
||||||
oauthSessionsMu.Lock()
|
|
||||||
oauthSessions[state] = &oauthSession{
|
|
||||||
Provider: "google-antigravity",
|
|
||||||
PKCE: pkce,
|
|
||||||
State: state,
|
|
||||||
RedirectURI: redirectURI,
|
|
||||||
OAuthCfg: oauthCfg,
|
|
||||||
ConfigPath: configPath,
|
|
||||||
}
|
|
||||||
oauthSessionsMu.Unlock()
|
|
||||||
|
|
||||||
// Clean up stale sessions after 10 minutes
|
|
||||||
go func() {
|
|
||||||
time.Sleep(10 * time.Minute)
|
|
||||||
oauthSessionsMu.Lock()
|
|
||||||
delete(oauthSessions, state)
|
|
||||||
oauthSessionsMu.Unlock()
|
|
||||||
}()
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{
|
|
||||||
"status": "redirect",
|
|
||||||
"auth_url": authURL,
|
|
||||||
"message": "Open the URL to authenticate with Google.",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleOAuthCallback processes the OAuth callback from Google Antigravity.
|
|
||||||
func handleOAuthCallback(w http.ResponseWriter, r *http.Request) {
|
|
||||||
state := r.URL.Query().Get("state")
|
|
||||||
code := r.URL.Query().Get("code")
|
|
||||||
|
|
||||||
oauthSessionsMu.Lock()
|
|
||||||
session, ok := oauthSessions[state]
|
|
||||||
if ok {
|
|
||||||
delete(oauthSessions, state)
|
|
||||||
}
|
|
||||||
oauthSessionsMu.Unlock()
|
|
||||||
|
|
||||||
if !ok {
|
|
||||||
http.Error(w, "Invalid or expired OAuth state", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if code == "" {
|
|
||||||
errMsg := r.URL.Query().Get("error")
|
|
||||||
w.Header().Set("Content-Type", "text/html")
|
|
||||||
fmt.Fprintf(
|
|
||||||
w,
|
|
||||||
`<html><body><h2>Authentication failed</h2><p>%s</p><p>You can close this window.</p></body></html>`,
|
|
||||||
errMsg,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
cred, err := auth.ExchangeCodeForTokens(session.OAuthCfg, code, session.PKCE.CodeVerifier, session.RedirectURI)
|
|
||||||
if err != nil {
|
|
||||||
w.Header().Set("Content-Type", "text/html")
|
|
||||||
fmt.Fprintf(
|
|
||||||
w,
|
|
||||||
`<html><body><h2>Authentication failed</h2><p>%s</p><p>You can close this window.</p></body></html>`,
|
|
||||||
err.Error(),
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
cred.Provider = session.Provider
|
|
||||||
|
|
||||||
// Fetch user info for Google Antigravity
|
|
||||||
if session.Provider == "google-antigravity" {
|
|
||||||
if email, err := fetchGoogleUserEmail(cred.AccessToken); err == nil {
|
|
||||||
cred.Email = email
|
|
||||||
}
|
|
||||||
if projectID, err := providers.FetchAntigravityProjectID(cred.AccessToken); err == nil {
|
|
||||||
cred.ProjectID = projectID
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := auth.SetCredential(session.Provider, cred); err != nil {
|
|
||||||
w.Header().Set("Content-Type", "text/html")
|
|
||||||
fmt.Fprintf(w, `<html><body><h2>Failed to save credentials</h2><p>%s</p></body></html>`, err.Error())
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
updateConfigAfterLogin(session.ConfigPath, session.Provider, cred)
|
|
||||||
|
|
||||||
// Redirect back to picoclaw-launcher UI
|
|
||||||
w.Header().Set("Content-Type", "text/html")
|
|
||||||
fmt.Fprintf(w, `<html><body>
|
|
||||||
<h2>Authentication successful!</h2>
|
|
||||||
<p>Redirecting back to Config Editor...</p>
|
|
||||||
<script>setTimeout(function(){ window.location.href = '/#auth'; }, 1000);</script>
|
|
||||||
</body></html>`)
|
|
||||||
}
|
|
||||||
|
|
||||||
// fetchGoogleUserEmail retrieves the user's email from Google's userinfo endpoint.
|
|
||||||
func fetchGoogleUserEmail(accessToken string) (string, error) {
|
|
||||||
req, err := http.NewRequest("GET", "https://www.googleapis.com/oauth2/v2/userinfo", nil)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
|
||||||
|
|
||||||
client := &http.Client{Timeout: 10 * time.Second}
|
|
||||||
resp, err := client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("reading userinfo response: %w", err)
|
|
||||||
}
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
return "", fmt.Errorf("userinfo request failed: %s", string(body))
|
|
||||||
}
|
|
||||||
|
|
||||||
var userInfo struct {
|
|
||||||
Email string `json:"email"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(body, &userInfo); err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
return userInfo.Email, nil
|
|
||||||
}
|
|
||||||
|
|
@ -1,116 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestLogBuffer_Basic(t *testing.T) {
|
|
||||||
buf := NewLogBuffer(5)
|
|
||||||
|
|
||||||
// Empty buffer
|
|
||||||
lines, total, runID := buf.LinesSince(0)
|
|
||||||
assert.Nil(t, lines)
|
|
||||||
assert.Equal(t, 0, total)
|
|
||||||
assert.Equal(t, 0, runID)
|
|
||||||
|
|
||||||
// Append some lines
|
|
||||||
buf.Append("line1")
|
|
||||||
buf.Append("line2")
|
|
||||||
buf.Append("line3")
|
|
||||||
|
|
||||||
lines, total, runID = buf.LinesSince(0)
|
|
||||||
assert.Equal(t, []string{"line1", "line2", "line3"}, lines)
|
|
||||||
assert.Equal(t, 3, total)
|
|
||||||
assert.Equal(t, 0, runID)
|
|
||||||
|
|
||||||
// Incremental read
|
|
||||||
lines, total, _ = buf.LinesSince(2)
|
|
||||||
assert.Equal(t, []string{"line3"}, lines)
|
|
||||||
assert.Equal(t, 3, total)
|
|
||||||
|
|
||||||
// No new lines
|
|
||||||
lines, total, _ = buf.LinesSince(3)
|
|
||||||
assert.Nil(t, lines)
|
|
||||||
assert.Equal(t, 3, total)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLogBuffer_Wrap(t *testing.T) {
|
|
||||||
buf := NewLogBuffer(3)
|
|
||||||
|
|
||||||
buf.Append("a")
|
|
||||||
buf.Append("b")
|
|
||||||
buf.Append("c")
|
|
||||||
buf.Append("d") // evicts "a"
|
|
||||||
buf.Append("e") // evicts "b"
|
|
||||||
|
|
||||||
lines, total, _ := buf.LinesSince(0)
|
|
||||||
assert.Equal(t, []string{"c", "d", "e"}, lines)
|
|
||||||
assert.Equal(t, 5, total)
|
|
||||||
|
|
||||||
// Incremental after wrap
|
|
||||||
lines, total, _ = buf.LinesSince(3)
|
|
||||||
assert.Equal(t, []string{"d", "e"}, lines)
|
|
||||||
assert.Equal(t, 5, total)
|
|
||||||
|
|
||||||
// Offset too old (before buffer start), get all buffered
|
|
||||||
lines, total, _ = buf.LinesSince(1)
|
|
||||||
assert.Equal(t, []string{"c", "d", "e"}, lines)
|
|
||||||
assert.Equal(t, 5, total)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLogBuffer_Reset(t *testing.T) {
|
|
||||||
buf := NewLogBuffer(5)
|
|
||||||
|
|
||||||
buf.Append("before")
|
|
||||||
assert.Equal(t, 0, buf.RunID())
|
|
||||||
|
|
||||||
buf.Reset()
|
|
||||||
assert.Equal(t, 1, buf.RunID())
|
|
||||||
assert.Equal(t, 0, buf.Total())
|
|
||||||
|
|
||||||
lines, total, runID := buf.LinesSince(0)
|
|
||||||
assert.Nil(t, lines)
|
|
||||||
assert.Equal(t, 0, total)
|
|
||||||
assert.Equal(t, 1, runID)
|
|
||||||
|
|
||||||
buf.Append("after")
|
|
||||||
lines, total, runID = buf.LinesSince(0)
|
|
||||||
assert.Equal(t, []string{"after"}, lines)
|
|
||||||
assert.Equal(t, 1, total)
|
|
||||||
assert.Equal(t, 1, runID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLogBuffer_Concurrent(t *testing.T) {
|
|
||||||
buf := NewLogBuffer(100)
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
|
|
||||||
// 10 writers
|
|
||||||
for i := range 10 {
|
|
||||||
wg.Add(1)
|
|
||||||
go func(id int) {
|
|
||||||
defer wg.Done()
|
|
||||||
for j := range 50 {
|
|
||||||
buf.Append(fmt.Sprintf("writer-%d-line-%d", id, j))
|
|
||||||
}
|
|
||||||
}(i)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 5 readers
|
|
||||||
for range 5 {
|
|
||||||
wg.Add(1)
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
for range 100 {
|
|
||||||
buf.LinesSince(0)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
|
|
||||||
wg.Wait()
|
|
||||||
|
|
||||||
assert.Equal(t, 500, buf.Total())
|
|
||||||
}
|
|
||||||
|
|
@ -1,232 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"log"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"os"
|
|
||||||
"os/exec"
|
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
|
||||||
"strconv"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// gatewayLogs stores captured stdout/stderr from the gateway process launched by the launcher.
|
|
||||||
var gatewayLogs = NewLogBuffer(200)
|
|
||||||
|
|
||||||
// RegisterProcessAPI registers endpoints to start, stop and check status of the picoclaw gateway.
|
|
||||||
func RegisterProcessAPI(mux *http.ServeMux, absPath string) {
|
|
||||||
mux.HandleFunc("GET /api/process/status", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
handleStatusGateway(w, r, absPath)
|
|
||||||
})
|
|
||||||
mux.HandleFunc("POST /api/process/start", handleStartGateway)
|
|
||||||
mux.HandleFunc("POST /api/process/stop", handleStopGateway)
|
|
||||||
}
|
|
||||||
|
|
||||||
func handleStartGateway(w http.ResponseWriter, r *http.Request) {
|
|
||||||
// Locate picoclaw executable:
|
|
||||||
// 1. Try same directory as current executable
|
|
||||||
// 2. Fallback to just "picoclaw" (relies on $PATH)
|
|
||||||
execPath := "picoclaw"
|
|
||||||
|
|
||||||
if exe, err := os.Executable(); err == nil {
|
|
||||||
dir := filepath.Dir(exe)
|
|
||||||
candidate := filepath.Join(dir, "picoclaw")
|
|
||||||
if runtime.GOOS == "windows" {
|
|
||||||
candidate += ".exe"
|
|
||||||
}
|
|
||||||
|
|
||||||
if info, err := os.Stat(candidate); err == nil && !info.IsDir() {
|
|
||||||
execPath = candidate
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
cmd := exec.Command(execPath, "gateway")
|
|
||||||
|
|
||||||
stdoutPipe, err := cmd.StdoutPipe()
|
|
||||||
if err != nil {
|
|
||||||
log.Printf("Failed to create stdout pipe: %v\n", err)
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to start gateway: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
stderrPipe, err := cmd.StderrPipe()
|
|
||||||
if err != nil {
|
|
||||||
log.Printf("Failed to create stderr pipe: %v\n", err)
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to start gateway: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Clear old logs and increment runID before starting
|
|
||||||
gatewayLogs.Reset()
|
|
||||||
|
|
||||||
if err := cmd.Start(); err != nil {
|
|
||||||
log.Printf("Failed to start picoclaw gateway: %v\n", err)
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to start gateway: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read stdout and stderr into the log buffer
|
|
||||||
go scanPipe(stdoutPipe, gatewayLogs)
|
|
||||||
go scanPipe(stderrPipe, gatewayLogs)
|
|
||||||
|
|
||||||
// Wait for the process to exit in the background to avoid zombies
|
|
||||||
go func() {
|
|
||||||
if err := cmd.Wait(); err != nil {
|
|
||||||
log.Printf("Gateway process exited: %v\n", err)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
log.Printf("Started picoclaw gateway (PID: %d) from %s\n", cmd.Process.Pid, execPath)
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
|
||||||
"status": "ok",
|
|
||||||
"pid": cmd.Process.Pid,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// scanPipe reads lines from r and appends them to buf. It returns when r reaches EOF.
|
|
||||||
func scanPipe(r io.Reader, buf *LogBuffer) {
|
|
||||||
scanner := bufio.NewScanner(r)
|
|
||||||
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) // up to 1MB per line
|
|
||||||
|
|
||||||
for scanner.Scan() {
|
|
||||||
buf.Append(scanner.Text())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func handleStopGateway(w http.ResponseWriter, r *http.Request) {
|
|
||||||
var err error
|
|
||||||
if runtime.GOOS == "windows" {
|
|
||||||
// Kill via taskkill finding picoclaw.exe (though it might kill this config tool if it's named picoclaw-launcher.exe...? No, /IM does exact match usually, but just to be safe let's stop exactly picoclaw.exe)
|
|
||||||
// Alternatively, we use powershell to kill processes with commandline containing 'gateway'
|
|
||||||
psCmd := `Get-WmiObject Win32_Process | Where-Object { $_.CommandLine -match 'picoclaw.*gateway' } | ForEach-Object { Stop-Process $_.ProcessId -Force }`
|
|
||||||
err = exec.Command("powershell", "-Command", psCmd).Run()
|
|
||||||
} else {
|
|
||||||
// Linux/macOS
|
|
||||||
err = exec.Command("pkill", "-f", "picoclaw gateway").Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
log.Printf("Warning: Failed to stop gateway (perhaps not running?): %v\n", err)
|
|
||||||
// We still return 200 OK because pkill returns an error if no process was found
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
|
||||||
"status": "ok", // or "not_found"
|
|
||||||
"msg": "Stop command executed, but returned error (process might not be running).",
|
|
||||||
"error": err.Error(),
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Printf("Stopped picoclaw gateway processes.\n")
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{
|
|
||||||
"status": "ok",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func handleStatusGateway(w http.ResponseWriter, r *http.Request, absPath string) {
|
|
||||||
cfg, cfgErr := config.LoadConfig(absPath)
|
|
||||||
host := "127.0.0.1"
|
|
||||||
port := 18790
|
|
||||||
if cfgErr == nil && cfg != nil {
|
|
||||||
if cfg.Gateway.Host != "" && cfg.Gateway.Host != "0.0.0.0" {
|
|
||||||
host = cfg.Gateway.Host
|
|
||||||
}
|
|
||||||
if cfg.Gateway.Port != 0 {
|
|
||||||
port = cfg.Gateway.Port
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
url := fmt.Sprintf("http://%s/health", net.JoinHostPort(host, strconv.Itoa(port)))
|
|
||||||
client := http.Client{Timeout: 2 * time.Second}
|
|
||||||
resp, err := client.Get(url)
|
|
||||||
|
|
||||||
// Build the response data map
|
|
||||||
data := map[string]any{}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
data["process_status"] = "stopped"
|
|
||||||
data["error"] = err.Error()
|
|
||||||
} else {
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
data["process_status"] = "error"
|
|
||||||
data["status_code"] = resp.StatusCode
|
|
||||||
} else {
|
|
||||||
var healthData map[string]any
|
|
||||||
if decErr := json.NewDecoder(resp.Body).Decode(&healthData); decErr != nil {
|
|
||||||
data["process_status"] = "error"
|
|
||||||
data["error"] = "invalid response from gateway"
|
|
||||||
} else {
|
|
||||||
// Gateway is running and responded properly — merge health data
|
|
||||||
for k, v := range healthData {
|
|
||||||
data[k] = v
|
|
||||||
}
|
|
||||||
data["process_status"] = "running"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Append log data from the buffer
|
|
||||||
appendLogData(r, data)
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(data)
|
|
||||||
}
|
|
||||||
|
|
||||||
// appendLogData reads log_offset and log_run_id query params from the request and
|
|
||||||
// populates the response data map with incremental log lines.
|
|
||||||
func appendLogData(r *http.Request, data map[string]any) {
|
|
||||||
clientOffset := 0
|
|
||||||
clientRunID := -1
|
|
||||||
|
|
||||||
if v := r.URL.Query().Get("log_offset"); v != "" {
|
|
||||||
if n, err := strconv.Atoi(v); err == nil {
|
|
||||||
clientOffset = n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if v := r.URL.Query().Get("log_run_id"); v != "" {
|
|
||||||
if n, err := strconv.Atoi(v); err == nil {
|
|
||||||
clientRunID = n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
runID := gatewayLogs.RunID()
|
|
||||||
|
|
||||||
// If runID is 0 (never reset = never launched from this launcher), report no source
|
|
||||||
if runID == 0 {
|
|
||||||
data["logs"] = []string{}
|
|
||||||
data["log_total"] = 0
|
|
||||||
data["log_run_id"] = 0
|
|
||||||
data["log_source"] = "none"
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// If the client's runID doesn't match, send all buffered lines (gateway restarted)
|
|
||||||
offset := clientOffset
|
|
||||||
if clientRunID != runID {
|
|
||||||
offset = 0
|
|
||||||
}
|
|
||||||
|
|
||||||
lines, total, runID := gatewayLogs.LinesSince(offset)
|
|
||||||
if lines == nil {
|
|
||||||
lines = []string{}
|
|
||||||
}
|
|
||||||
|
|
||||||
data["logs"] = lines
|
|
||||||
data["log_total"] = total
|
|
||||||
data["log_run_id"] = runID
|
|
||||||
data["log_source"] = "launcher"
|
|
||||||
}
|
|
||||||
|
|
@ -1,196 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"log"
|
|
||||||
"net/http"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
const DefaultPort = "18800"
|
|
||||||
|
|
||||||
// providerStatus represents the auth status of a single provider in API responses.
|
|
||||||
type providerStatus struct {
|
|
||||||
Provider string `json:"provider"`
|
|
||||||
AuthMethod string `json:"auth_method"`
|
|
||||||
Status string `json:"status"`
|
|
||||||
AccountID string `json:"account_id,omitempty"`
|
|
||||||
Email string `json:"email,omitempty"`
|
|
||||||
ProjectID string `json:"project_id,omitempty"`
|
|
||||||
ExpiresAt string `json:"expires_at,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// ── Route registration ───────────────────────────────────────────
|
|
||||||
|
|
||||||
func RegisterConfigAPI(mux *http.ServeMux, absPath string) {
|
|
||||||
// GET /api/config — read config
|
|
||||||
mux.HandleFunc("GET /api/config", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
cfg, err := config.LoadConfig(absPath)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
resp := map[string]any{
|
|
||||||
"config": cfg,
|
|
||||||
"path": absPath,
|
|
||||||
}
|
|
||||||
enc := json.NewEncoder(w)
|
|
||||||
enc.SetIndent("", " ")
|
|
||||||
if err := enc.Encode(resp); err != nil {
|
|
||||||
log.Printf("Failed to encode response: %v", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
// PUT /api/config — save config
|
|
||||||
mux.HandleFunc("PUT /api/config", func(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 cfg config.Config
|
|
||||||
if err := json.Unmarshal(body, &cfg); err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := config.SaveConfig(absPath, &cfg); err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func RegisterAuthAPI(mux *http.ServeMux, absPath string) {
|
|
||||||
// GET /api/auth/status — all authenticated providers + pending login state
|
|
||||||
mux.HandleFunc("GET /api/auth/status", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
store, err := auth.LoadStore()
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to load auth store: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
result := []providerStatus{}
|
|
||||||
for name, cred := range store.Credentials {
|
|
||||||
status := "active"
|
|
||||||
if cred.IsExpired() {
|
|
||||||
status = "expired"
|
|
||||||
} else if cred.NeedsRefresh() {
|
|
||||||
status = "needs_refresh"
|
|
||||||
}
|
|
||||||
ps := providerStatus{
|
|
||||||
Provider: name,
|
|
||||||
AuthMethod: cred.AuthMethod,
|
|
||||||
Status: status,
|
|
||||||
AccountID: cred.AccountID,
|
|
||||||
Email: cred.Email,
|
|
||||||
ProjectID: cred.ProjectID,
|
|
||||||
}
|
|
||||||
if !cred.ExpiresAt.IsZero() {
|
|
||||||
ps.ExpiresAt = cred.ExpiresAt.Format(time.RFC3339)
|
|
||||||
}
|
|
||||||
result = append(result, ps)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Include pending device code state
|
|
||||||
var pendingDevice map[string]any
|
|
||||||
activeDeviceSessionMu.Lock()
|
|
||||||
if activeDeviceSession != nil {
|
|
||||||
activeDeviceSession.mu.Lock()
|
|
||||||
pendingDevice = map[string]any{
|
|
||||||
"provider": activeDeviceSession.Provider,
|
|
||||||
"status": activeDeviceSession.Status,
|
|
||||||
"device_url": activeDeviceSession.Info.VerifyURL,
|
|
||||||
"user_code": activeDeviceSession.Info.UserCode,
|
|
||||||
}
|
|
||||||
if activeDeviceSession.Error != "" {
|
|
||||||
pendingDevice["error"] = activeDeviceSession.Error
|
|
||||||
}
|
|
||||||
if activeDeviceSession.Done {
|
|
||||||
activeDeviceSession.mu.Unlock()
|
|
||||||
activeDeviceSession = nil
|
|
||||||
} else {
|
|
||||||
activeDeviceSession.mu.Unlock()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
activeDeviceSessionMu.Unlock()
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]any{
|
|
||||||
"providers": result,
|
|
||||||
"pending_device": pendingDevice,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
|
|
||||||
// POST /api/auth/login — initiate provider login
|
|
||||||
mux.HandleFunc("POST /api/auth/login", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
var req struct {
|
|
||||||
Provider string `json:"provider"`
|
|
||||||
Token string `json:"token,omitempty"`
|
|
||||||
}
|
|
||||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
||||||
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
switch req.Provider {
|
|
||||||
case "openai":
|
|
||||||
handleOpenAILogin(w, absPath)
|
|
||||||
case "anthropic":
|
|
||||||
handleAnthropicLogin(w, req.Token, absPath)
|
|
||||||
case "google-antigravity", "antigravity":
|
|
||||||
handleGoogleAntigravityLogin(w, r, absPath)
|
|
||||||
default:
|
|
||||||
http.Error(
|
|
||||||
w,
|
|
||||||
fmt.Sprintf(
|
|
||||||
"Unsupported provider: %s (supported: openai, anthropic, google-antigravity)",
|
|
||||||
req.Provider,
|
|
||||||
),
|
|
||||||
http.StatusBadRequest,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
// POST /api/auth/logout — logout a provider
|
|
||||||
mux.HandleFunc("POST /api/auth/logout", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
var req struct {
|
|
||||||
Provider string `json:"provider"`
|
|
||||||
}
|
|
||||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
||||||
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if req.Provider == "" {
|
|
||||||
if err := auth.DeleteAllCredentials(); err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to logout: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
clearAllAuthMethodsInConfig(absPath)
|
|
||||||
} else {
|
|
||||||
if err := auth.DeleteCredential(req.Provider); err != nil {
|
|
||||||
http.Error(w, fmt.Sprintf("Failed to logout: %v", err), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
clearAuthMethodInConfig(absPath, req.Provider)
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// GET /auth/callback — OAuth browser callback for Google Antigravity
|
|
||||||
mux.HandleFunc("GET /auth/callback", handleOAuthCallback)
|
|
||||||
}
|
|
||||||
|
|
@ -1,247 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ── Config API tests ─────────────────────────────────────────────
|
|
||||||
|
|
||||||
func setupConfigMux(t *testing.T, cfg *config.Config) (*http.ServeMux, string) {
|
|
||||||
t.Helper()
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "config.json")
|
|
||||||
data, err := json.MarshalIndent(cfg, "", " ")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("marshal config: %v", err)
|
|
||||||
}
|
|
||||||
if err := os.WriteFile(path, data, 0o600); err != nil {
|
|
||||||
t.Fatalf("write config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
RegisterConfigAPI(mux, path)
|
|
||||||
RegisterAuthAPI(mux, path)
|
|
||||||
return mux, path
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetConfig(t *testing.T) {
|
|
||||||
cfg := &config.Config{
|
|
||||||
ModelList: []config.ModelConfig{
|
|
||||||
{ModelName: "gpt-4o", Model: "openai/gpt-4o"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
mux, path := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("GET", "/api/config", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Fatalf("GET /api/config: expected 200, got %d: %s", w.Code, w.Body.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
var resp struct {
|
|
||||||
Config config.Config `json:"config"`
|
|
||||||
Path string `json:"path"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
|
||||||
t.Fatalf("decode response: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.Path != path {
|
|
||||||
t.Errorf("expected path %q, got %q", path, resp.Path)
|
|
||||||
}
|
|
||||||
if len(resp.Config.ModelList) != 1 {
|
|
||||||
t.Errorf("expected 1 model, got %d", len(resp.Config.ModelList))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetConfig_MissingFile_ReturnsDefault(t *testing.T) {
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
RegisterConfigAPI(mux, "/tmp/nonexistent-picoclaw-launcher-test/config.json")
|
|
||||||
|
|
||||||
req := httptest.NewRequest("GET", "/api/config", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
// LoadConfig returns a default empty config when file is missing
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Errorf("expected 200 for missing file (default config), got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPutConfig(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, path := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
newCfg := config.Config{
|
|
||||||
ModelList: []config.ModelConfig{
|
|
||||||
{ModelName: "claude", Model: "anthropic/claude-sonnet-4.6", AuthMethod: "token"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
body, _ := json.Marshal(newCfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("PUT", "/api/config", strings.NewReader(string(body)))
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Fatalf("PUT /api/config: expected 200, got %d: %s", w.Code, w.Body.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
saved, err := config.LoadConfig(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load saved config: %v", err)
|
|
||||||
}
|
|
||||||
if len(saved.ModelList) != 1 {
|
|
||||||
t.Fatalf("expected 1 model saved, got %d", len(saved.ModelList))
|
|
||||||
}
|
|
||||||
if saved.ModelList[0].Model != "anthropic/claude-sonnet-4.6" {
|
|
||||||
t.Errorf("expected model anthropic/claude-sonnet-4.6, got %q", saved.ModelList[0].Model)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPutConfig_InvalidJSON(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("PUT", "/api/config", strings.NewReader("{invalid"))
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("expected 400 for invalid JSON, got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ── Auth API tests ───────────────────────────────────────────────
|
|
||||||
|
|
||||||
func TestAuthStatus(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("GET", "/api/auth/status", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusOK {
|
|
||||||
t.Fatalf("GET /api/auth/status: expected 200, got %d: %s", w.Code, w.Body.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
var resp struct {
|
|
||||||
Providers []providerStatus `json:"providers"`
|
|
||||||
PendingDevice map[string]any `json:"pending_device"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
|
||||||
t.Fatalf("decode response: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// providers should be a non-nil list (could be empty)
|
|
||||||
if resp.Providers == nil {
|
|
||||||
t.Error("providers should not be nil")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAuthLogin_UnsupportedProvider(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
body := `{"provider": "unsupported"}`
|
|
||||||
req := httptest.NewRequest("POST", "/api/auth/login", strings.NewReader(body))
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("expected 400 for unsupported provider, got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAuthLogin_AnthropicNoToken(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
body := `{"provider": "anthropic"}`
|
|
||||||
req := httptest.NewRequest("POST", "/api/auth/login", strings.NewReader(body))
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("expected 400 for anthropic without token, got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAuthLogin_InvalidBody(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("POST", "/api/auth/login", strings.NewReader("{bad"))
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("expected 400 for invalid JSON body, got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAuthLogout_InvalidBody(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("POST", "/api/auth/logout", strings.NewReader("{bad"))
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("expected 400 for invalid body, got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOAuthCallback_InvalidState(t *testing.T) {
|
|
||||||
cfg := &config.Config{}
|
|
||||||
mux, _ := setupConfigMux(t, cfg)
|
|
||||||
|
|
||||||
req := httptest.NewRequest("GET", "/auth/callback?state=invalid&code=test", nil)
|
|
||||||
w := httptest.NewRecorder()
|
|
||||||
mux.ServeHTTP(w, req)
|
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
|
||||||
t.Errorf("expected 400 for invalid state, got %d", w.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ── Utility tests ────────────────────────────────────────────────
|
|
||||||
|
|
||||||
func TestDefaultConfigPath(t *testing.T) {
|
|
||||||
path := DefaultConfigPath()
|
|
||||||
if path == "" {
|
|
||||||
t.Error("defaultConfigPath should not return empty")
|
|
||||||
}
|
|
||||||
if !strings.HasSuffix(path, filepath.Join(".picoclaw", "config.json")) {
|
|
||||||
t.Errorf("expected path ending with .picoclaw/config.json, got %q", path)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetLocalIP(t *testing.T) {
|
|
||||||
// Just ensure it doesn't panic; IP may or may not be available
|
|
||||||
ip := GetLocalIP()
|
|
||||||
if ip != "" {
|
|
||||||
// If returned, should look like an IP
|
|
||||||
if !strings.Contains(ip, ".") {
|
|
||||||
t.Errorf("getLocalIP returned non-IPv4 looking string: %q", ip)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,28 +0,0 @@
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
)
|
|
||||||
|
|
||||||
func DefaultConfigPath() string {
|
|
||||||
home, err := os.UserHomeDir()
|
|
||||||
if err != nil {
|
|
||||||
return "config.json"
|
|
||||||
}
|
|
||||||
return filepath.Join(home, ".picoclaw", "config.json")
|
|
||||||
}
|
|
||||||
|
|
||||||
func GetLocalIP() string {
|
|
||||||
addrs, err := net.InterfaceAddrs()
|
|
||||||
if err != nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
for _, a := range addrs {
|
|
||||||
if ipnet, ok := a.(*net.IPNet); ok && !ipnet.IP.IsLoopback() && ipnet.IP.To4() != nil {
|
|
||||||
return ipnet.IP.String()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,127 +0,0 @@
|
||||||
// PicoClaw Launcher - Standalone HTTP service
|
|
||||||
//
|
|
||||||
// Provides a web-based JSON editor for picoclaw config files,
|
|
||||||
// with OAuth provider authentication support.
|
|
||||||
//
|
|
||||||
// Usage:
|
|
||||||
//
|
|
||||||
// go build -o picoclaw-launcher ./cmd/picoclaw-launcher/
|
|
||||||
// ./picoclaw-launcher [config.json]
|
|
||||||
// ./picoclaw-launcher -public config.json
|
|
||||||
|
|
||||||
package main
|
|
||||||
|
|
||||||
import (
|
|
||||||
"embed"
|
|
||||||
"flag"
|
|
||||||
"fmt"
|
|
||||||
"io/fs"
|
|
||||||
"log"
|
|
||||||
"net/http"
|
|
||||||
"os"
|
|
||||||
"os/exec"
|
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw-launcher/internal/server"
|
|
||||||
)
|
|
||||||
|
|
||||||
//go:embed internal/ui/index.html
|
|
||||||
var staticFiles embed.FS
|
|
||||||
|
|
||||||
func main() {
|
|
||||||
public := flag.Bool("public", false, "Listen on all interfaces (0.0.0.0) instead of localhost only")
|
|
||||||
flag.Usage = func() {
|
|
||||||
fmt.Fprintf(os.Stderr, "PicoClaw Launcher - A web-based configuration editor\n\n")
|
|
||||||
fmt.Fprintf(os.Stderr, "Usage: %s [options] [config.json]\n\n", os.Args[0])
|
|
||||||
fmt.Fprintf(os.Stderr, "Arguments:\n")
|
|
||||||
fmt.Fprintf(os.Stderr, " config.json Path to the configuration file (default: ~/.picoclaw/config.json)\n\n")
|
|
||||||
fmt.Fprintf(os.Stderr, "Options:\n")
|
|
||||||
flag.PrintDefaults()
|
|
||||||
fmt.Fprintf(os.Stderr, "\nExamples:\n")
|
|
||||||
fmt.Fprintf(os.Stderr, " %s Use default config path\n", os.Args[0])
|
|
||||||
fmt.Fprintf(os.Stderr, " %s ./config.json Specify a config file\n", os.Args[0])
|
|
||||||
fmt.Fprintf(
|
|
||||||
os.Stderr,
|
|
||||||
" %s -public ./config.json Allow access from other devices on the network\n",
|
|
||||||
os.Args[0],
|
|
||||||
)
|
|
||||||
}
|
|
||||||
flag.Parse()
|
|
||||||
|
|
||||||
configPath := server.DefaultConfigPath()
|
|
||||||
if flag.NArg() > 0 {
|
|
||||||
configPath = flag.Arg(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
absPath, err := filepath.Abs(configPath)
|
|
||||||
if err != nil {
|
|
||||||
log.Fatalf("Failed to resolve config path: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var addr string
|
|
||||||
if *public {
|
|
||||||
addr = "0.0.0.0:" + server.DefaultPort
|
|
||||||
} else {
|
|
||||||
addr = "127.0.0.1:" + server.DefaultPort
|
|
||||||
}
|
|
||||||
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
server.RegisterConfigAPI(mux, absPath)
|
|
||||||
server.RegisterAuthAPI(mux, absPath)
|
|
||||||
server.RegisterProcessAPI(mux, absPath)
|
|
||||||
|
|
||||||
staticFS, err := fs.Sub(staticFiles, "internal/ui")
|
|
||||||
if err != nil {
|
|
||||||
log.Fatalf("Failed to create sub filesystem: %v", err)
|
|
||||||
}
|
|
||||||
mux.Handle("/", http.FileServer(http.FS(staticFS)))
|
|
||||||
|
|
||||||
// Print startup banner
|
|
||||||
fmt.Println("=============================================")
|
|
||||||
fmt.Println(" PicoClaw Launcher")
|
|
||||||
fmt.Println("=============================================")
|
|
||||||
fmt.Printf(" Config file : %s\n", absPath)
|
|
||||||
fmt.Printf(" Listen addr : %s\n\n", addr)
|
|
||||||
fmt.Println(" Open the following URL in your browser")
|
|
||||||
fmt.Println(" to view and edit the configuration:")
|
|
||||||
fmt.Println()
|
|
||||||
fmt.Printf(" >> http://localhost:%s <<\n", server.DefaultPort)
|
|
||||||
if *public {
|
|
||||||
if ip := server.GetLocalIP(); ip != "" {
|
|
||||||
fmt.Printf(" >> http://%s:%s <<\n", ip, server.DefaultPort)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
fmt.Println()
|
|
||||||
// fmt.Println("=============================================")
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
// Wait briefly to ensure the server is ready before opening the browser
|
|
||||||
time.Sleep(500 * time.Millisecond)
|
|
||||||
url := "http://localhost:" + server.DefaultPort
|
|
||||||
if err := openBrowser(url); err != nil {
|
|
||||||
log.Printf("Warning: Failed to auto-open browser: %v\n", err)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if err := http.ListenAndServe(addr, mux); err != nil {
|
|
||||||
log.Fatalf("Server failed: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// openBrowser automatically opens the given URL in the default browser.
|
|
||||||
func openBrowser(url string) error {
|
|
||||||
var err error
|
|
||||||
switch runtime.GOOS {
|
|
||||||
case "linux":
|
|
||||||
err = exec.Command("xdg-open", url).Start()
|
|
||||||
case "windows":
|
|
||||||
err = exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start()
|
|
||||||
case "darwin":
|
|
||||||
err = exec.Command("open", url).Start()
|
|
||||||
default:
|
|
||||||
err = fmt.Errorf("unsupported platform")
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
@ -50,6 +50,7 @@ func agentCmd(message, sessionKey, model string, debug bool) error {
|
||||||
msgBus := bus.NewMessageBus()
|
msgBus := bus.NewMessageBus()
|
||||||
defer msgBus.Close()
|
defer msgBus.Close()
|
||||||
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
defer agentLoop.Close()
|
||||||
|
|
||||||
// Print agent startup info (only for interactive mode)
|
// Print agent startup info (only for interactive mode)
|
||||||
startupInfo := agentLoop.GetStartupInfo()
|
startupInfo := agentLoop.GetStartupInfo()
|
||||||
|
|
|
||||||
|
|
@ -19,6 +19,7 @@ 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"
|
||||||
|
|
@ -213,6 +214,7 @@ func gatewayCmd(debug bool) error {
|
||||||
cronService.Stop()
|
cronService.Stop()
|
||||||
mediaStore.Stop()
|
mediaStore.Stop()
|
||||||
agentLoop.Stop()
|
agentLoop.Stop()
|
||||||
|
agentLoop.Close()
|
||||||
fmt.Println("✓ Gateway stopped")
|
fmt.Println("✓ Gateway stopped")
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|
|
||||||
|
|
@ -1,23 +1,14 @@
|
||||||
package internal
|
package internal
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
const Logo = "🦞"
|
const Logo = "🦞"
|
||||||
|
|
||||||
var (
|
|
||||||
version = "dev"
|
|
||||||
gitCommit string
|
|
||||||
buildTime string
|
|
||||||
goVersion string
|
|
||||||
)
|
|
||||||
|
|
||||||
// GetPicoclawHome returns the picoclaw home directory.
|
// GetPicoclawHome returns the picoclaw home directory.
|
||||||
// Priority: $PICOCLAW_HOME > ~/.picoclaw
|
// Priority: $PICOCLAW_HOME > ~/.picoclaw
|
||||||
func GetPicoclawHome() string {
|
func GetPicoclawHome() string {
|
||||||
|
|
@ -40,25 +31,19 @@ func LoadConfig() (*config.Config, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// FormatVersion returns the version string with optional git commit
|
// FormatVersion returns the version string with optional git commit
|
||||||
|
// Deprecated: Use pkg/config.FormatVersion instead
|
||||||
func FormatVersion() string {
|
func FormatVersion() string {
|
||||||
v := version
|
return config.FormatVersion()
|
||||||
if gitCommit != "" {
|
|
||||||
v += fmt.Sprintf(" (git: %s)", gitCommit)
|
|
||||||
}
|
|
||||||
return v
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// FormatBuildInfo returns build time and go version info
|
// FormatBuildInfo returns build time and go version info
|
||||||
|
// Deprecated: Use pkg/config.FormatBuildInfo instead
|
||||||
func FormatBuildInfo() (string, string) {
|
func FormatBuildInfo() (string, string) {
|
||||||
build := buildTime
|
return config.FormatBuildInfo()
|
||||||
goVer := goVersion
|
|
||||||
if goVer == "" {
|
|
||||||
goVer = runtime.Version()
|
|
||||||
}
|
|
||||||
return build, goVer
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetVersion returns the version string
|
// GetVersion returns the version string
|
||||||
|
// Deprecated: Use pkg/config.GetVersion instead
|
||||||
func GetVersion() string {
|
func GetVersion() string {
|
||||||
return version
|
return config.GetVersion()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -40,65 +40,6 @@ func TestGetConfigPath_WithPICOCLAW_CONFIG(t *testing.T) {
|
||||||
assert.Equal(t, want, got)
|
assert.Equal(t, want, got)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFormatVersion_NoGitCommit(t *testing.T) {
|
|
||||||
oldVersion, oldGit := version, gitCommit
|
|
||||||
t.Cleanup(func() { version, gitCommit = oldVersion, oldGit })
|
|
||||||
|
|
||||||
version = "1.2.3"
|
|
||||||
gitCommit = ""
|
|
||||||
|
|
||||||
assert.Equal(t, "1.2.3", FormatVersion())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatVersion_WithGitCommit(t *testing.T) {
|
|
||||||
oldVersion, oldGit := version, gitCommit
|
|
||||||
t.Cleanup(func() { version, gitCommit = oldVersion, oldGit })
|
|
||||||
|
|
||||||
version = "1.2.3"
|
|
||||||
gitCommit = "abc123"
|
|
||||||
|
|
||||||
assert.Equal(t, "1.2.3 (git: abc123)", FormatVersion())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatBuildInfo_UsesBuildTimeAndGoVersion_WhenSet(t *testing.T) {
|
|
||||||
oldBuildTime, oldGoVersion := buildTime, goVersion
|
|
||||||
t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion })
|
|
||||||
|
|
||||||
buildTime = "2026-02-20T00:00:00Z"
|
|
||||||
goVersion = "go1.23.0"
|
|
||||||
|
|
||||||
build, goVer := FormatBuildInfo()
|
|
||||||
|
|
||||||
assert.Equal(t, buildTime, build)
|
|
||||||
assert.Equal(t, goVersion, goVer)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatBuildInfo_EmptyBuildTime_ReturnsEmptyBuild(t *testing.T) {
|
|
||||||
oldBuildTime, oldGoVersion := buildTime, goVersion
|
|
||||||
t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion })
|
|
||||||
|
|
||||||
buildTime = ""
|
|
||||||
goVersion = "go1.23.0"
|
|
||||||
|
|
||||||
build, goVer := FormatBuildInfo()
|
|
||||||
|
|
||||||
assert.Empty(t, build)
|
|
||||||
assert.Equal(t, goVersion, goVer)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormatBuildInfo_EmptyGoVersion_FallsBackToRuntimeVersion(t *testing.T) {
|
|
||||||
oldBuildTime, oldGoVersion := buildTime, goVersion
|
|
||||||
t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion })
|
|
||||||
|
|
||||||
buildTime = "x"
|
|
||||||
goVersion = ""
|
|
||||||
|
|
||||||
build, goVer := FormatBuildInfo()
|
|
||||||
|
|
||||||
assert.Equal(t, "x", build)
|
|
||||||
assert.Equal(t, runtime.Version(), goVer)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetConfigPath_Windows(t *testing.T) {
|
func TestGetConfigPath_Windows(t *testing.T) {
|
||||||
if runtime.GOOS != "windows" {
|
if runtime.GOOS != "windows" {
|
||||||
t.Skip("windows-specific HOME behavior varies; run on windows")
|
t.Skip("windows-specific HOME behavior varies; run on windows")
|
||||||
|
|
@ -112,17 +53,3 @@ func TestGetConfigPath_Windows(t *testing.T) {
|
||||||
|
|
||||||
require.True(t, strings.EqualFold(got, want), "GetConfigPath() = %q, want %q", got, want)
|
require.True(t, strings.EqualFold(got, want), "GetConfigPath() = %q, want %q", got, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGetVersion(t *testing.T) {
|
|
||||||
assert.Equal(t, "dev", GetVersion())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetConfigPath_WithEnv(t *testing.T) {
|
|
||||||
t.Setenv("PICOCLAW_CONFIG", "/tmp/custom/config.json")
|
|
||||||
t.Setenv("HOME", "/tmp/home") // Also set home to ensure env is preferred
|
|
||||||
|
|
||||||
got := GetConfigPath()
|
|
||||||
want := "/tmp/custom/config.json"
|
|
||||||
|
|
||||||
assert.Equal(t, want, got)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
"github.com/sipeed/picoclaw/pkg/auth"
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func statusCmd() {
|
func statusCmd() {
|
||||||
|
|
@ -18,8 +19,8 @@ func statusCmd() {
|
||||||
configPath := internal.GetConfigPath()
|
configPath := internal.GetConfigPath()
|
||||||
|
|
||||||
fmt.Printf("%s picoclaw Status\n", internal.Logo)
|
fmt.Printf("%s picoclaw Status\n", internal.Logo)
|
||||||
fmt.Printf("Version: %s\n", internal.FormatVersion())
|
fmt.Printf("Version: %s\n", config.FormatVersion())
|
||||||
build, _ := internal.FormatBuildInfo()
|
build, _ := config.FormatBuildInfo()
|
||||||
if build != "" {
|
if build != "" {
|
||||||
fmt.Printf("Build: %s\n", build)
|
fmt.Printf("Build: %s\n", build)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewVersionCommand() *cobra.Command {
|
func NewVersionCommand() *cobra.Command {
|
||||||
|
|
@ -22,8 +23,8 @@ func NewVersionCommand() *cobra.Command {
|
||||||
}
|
}
|
||||||
|
|
||||||
func printVersion() {
|
func printVersion() {
|
||||||
fmt.Printf("%s picoclaw %s\n", internal.Logo, internal.FormatVersion())
|
fmt.Printf("%s picoclaw %s\n", internal.Logo, config.FormatVersion())
|
||||||
build, goVer := internal.FormatBuildInfo()
|
build, goVer := config.FormatBuildInfo()
|
||||||
if build != "" {
|
if build != "" {
|
||||||
fmt.Printf(" Build: %s\n", build)
|
fmt.Printf(" Build: %s\n", build)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -22,15 +22,16 @@ import (
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/skills"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/skills"
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/status"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/status"
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/version"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal/version"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewPicoclawCommand() *cobra.Command {
|
func NewPicoclawCommand() *cobra.Command {
|
||||||
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, internal.GetVersion())
|
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, config.GetVersion())
|
||||||
|
|
||||||
cmd := &cobra.Command{
|
cmd := &cobra.Command{
|
||||||
Use: "picoclaw",
|
Use: "picoclaw",
|
||||||
Short: short,
|
Short: short,
|
||||||
Example: "picoclaw list",
|
Example: "picoclaw version",
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd.AddCommand(
|
cmd.AddCommand(
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ import (
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
"github.com/sipeed/picoclaw/cmd/picoclaw/internal"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewPicoclawCommand(t *testing.T) {
|
func TestNewPicoclawCommand(t *testing.T) {
|
||||||
|
|
@ -16,7 +17,7 @@ func TestNewPicoclawCommand(t *testing.T) {
|
||||||
|
|
||||||
require.NotNil(t, cmd)
|
require.NotNil(t, cmd)
|
||||||
|
|
||||||
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, internal.GetVersion())
|
short := fmt.Sprintf("%s picoclaw - Personal AI Assistant v%s\n\n", internal.Logo, config.GetVersion())
|
||||||
|
|
||||||
assert.Equal(t, "picoclaw", cmd.Use)
|
assert.Equal(t, "picoclaw", cmd.Use)
|
||||||
assert.Equal(t, short, cmd.Short)
|
assert.Equal(t, short, cmd.Short)
|
||||||
|
|
|
||||||
|
|
@ -97,7 +97,8 @@
|
||||||
"encrypt_key": "",
|
"encrypt_key": "",
|
||||||
"verification_token": "",
|
"verification_token": "",
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": "",
|
||||||
|
"random_reaction_emoji": []
|
||||||
},
|
},
|
||||||
"dingtalk": {
|
"dingtalk": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
@ -113,6 +114,23 @@
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"reasoning_channel_id": ""
|
"reasoning_channel_id": ""
|
||||||
},
|
},
|
||||||
|
"matrix": {
|
||||||
|
"enabled": false,
|
||||||
|
"homeserver": "https://matrix.org",
|
||||||
|
"user_id": "@your-bot:matrix.org",
|
||||||
|
"access_token": "YOUR_MATRIX_ACCESS_TOKEN",
|
||||||
|
"device_id": "",
|
||||||
|
"join_on_invite": true,
|
||||||
|
"allow_from": [],
|
||||||
|
"group_trigger": {
|
||||||
|
"mention_only": true
|
||||||
|
},
|
||||||
|
"placeholder": {
|
||||||
|
"enabled": true,
|
||||||
|
"text": "Thinking... 💭"
|
||||||
|
},
|
||||||
|
"reasoning_channel_id": ""
|
||||||
|
},
|
||||||
"line": {
|
"line": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"channel_secret": "YOUR_LINE_CHANNEL_SECRET",
|
"channel_secret": "YOUR_LINE_CHANNEL_SECRET",
|
||||||
|
|
@ -175,8 +193,13 @@
|
||||||
"nickserv_password": "",
|
"nickserv_password": "",
|
||||||
"sasl_user": "",
|
"sasl_user": "",
|
||||||
"sasl_password": "",
|
"sasl_password": "",
|
||||||
"channels": ["#mychannel"],
|
"channels": [
|
||||||
"request_caps": ["server-time", "message-tags"],
|
"#mychannel"
|
||||||
|
],
|
||||||
|
"request_caps": [
|
||||||
|
"server-time",
|
||||||
|
"message-tags"
|
||||||
|
],
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"group_trigger": {
|
"group_trigger": {
|
||||||
"mention_only": true
|
"mention_only": true
|
||||||
|
|
@ -260,6 +283,9 @@
|
||||||
"brave": {
|
"brave": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "YOUR_BRAVE_API_KEY",
|
"api_key": "YOUR_BRAVE_API_KEY",
|
||||||
|
"api_keys": [
|
||||||
|
"YOUR_BRAVE_API_KEY"
|
||||||
|
],
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
},
|
},
|
||||||
"tavily": {
|
"tavily": {
|
||||||
|
|
@ -274,7 +300,10 @@
|
||||||
},
|
},
|
||||||
"perplexity": {
|
"perplexity": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "",
|
"api_key": "pplx-xxx",
|
||||||
|
"api_keys": [
|
||||||
|
"pplx-xxx"
|
||||||
|
],
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
},
|
},
|
||||||
"searxng": {
|
"searxng": {
|
||||||
|
|
@ -297,6 +326,13 @@
|
||||||
},
|
},
|
||||||
"mcp": {
|
"mcp": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
"discovery": {
|
||||||
|
"enabled": false,
|
||||||
|
"ttl": 5,
|
||||||
|
"max_search_results": 5,
|
||||||
|
"use_bm25": true,
|
||||||
|
"use_regex": false
|
||||||
|
},
|
||||||
"servers": {
|
"servers": {
|
||||||
"context7": {
|
"context7": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
|
||||||
12
docker/Dockerfile.goreleaser.launcher
Normal file
12
docker/Dockerfile.goreleaser.launcher
Normal file
|
|
@ -0,0 +1,12 @@
|
||||||
|
FROM alpine:3.21
|
||||||
|
|
||||||
|
ARG TARGETPLATFORM
|
||||||
|
|
||||||
|
RUN apk add --no-cache ca-certificates tzdata
|
||||||
|
|
||||||
|
COPY $TARGETPLATFORM/picoclaw /usr/local/bin/picoclaw
|
||||||
|
COPY $TARGETPLATFORM/picoclaw-launcher /usr/local/bin/picoclaw-launcher
|
||||||
|
COPY $TARGETPLATFORM/picoclaw-launcher-tui /usr/local/bin/picoclaw-launcher-tui
|
||||||
|
|
||||||
|
ENTRYPOINT ["picoclaw-launcher"]
|
||||||
|
CMD ["-public", "-no-browser"]
|
||||||
|
|
@ -19,7 +19,7 @@ services:
|
||||||
|
|
||||||
# ─────────────────────────────────────────────
|
# ─────────────────────────────────────────────
|
||||||
# PicoClaw Gateway (Long-running Bot)
|
# PicoClaw Gateway (Long-running Bot)
|
||||||
# docker compose -f docker/docker-compose.yml up picoclaw-gateway
|
# docker compose -f docker/docker-compose.yml --profile gateway up
|
||||||
# ─────────────────────────────────────────────
|
# ─────────────────────────────────────────────
|
||||||
picoclaw-gateway:
|
picoclaw-gateway:
|
||||||
image: docker.io/sipeed/picoclaw:latest
|
image: docker.io/sipeed/picoclaw:latest
|
||||||
|
|
@ -32,3 +32,21 @@ services:
|
||||||
# - "host.docker.internal:host-gateway"
|
# - "host.docker.internal:host-gateway"
|
||||||
volumes:
|
volumes:
|
||||||
- ./data:/root/.picoclaw
|
- ./data:/root/.picoclaw
|
||||||
|
|
||||||
|
# ─────────────────────────────────────────────
|
||||||
|
# PicoClaw Launcher (Web Console + Gateway)
|
||||||
|
# docker compose -f docker/docker-compose.yml --profile launcher up
|
||||||
|
# ─────────────────────────────────────────────
|
||||||
|
picoclaw-launcher:
|
||||||
|
image: docker.io/sipeed/picoclaw:launcher
|
||||||
|
container_name: picoclaw-launcher
|
||||||
|
restart: on-failure
|
||||||
|
profiles:
|
||||||
|
- launcher
|
||||||
|
environment:
|
||||||
|
- PICOCLAW_GATEWAY_HOST=0.0.0.0
|
||||||
|
ports:
|
||||||
|
- "127.0.0.1:18800:18800"
|
||||||
|
- "127.0.0.1:18790:18790"
|
||||||
|
volumes:
|
||||||
|
- ./data:/root/.picoclaw
|
||||||
|
|
|
||||||
|
|
@ -26,7 +26,8 @@
|
||||||
| app_secret | string | 是 | 飞书应用的 App Secret |
|
| app_secret | string | 是 | 飞书应用的 App Secret |
|
||||||
| encrypt_key | string | 否 | 事件回调加密密钥 |
|
| encrypt_key | string | 否 | 事件回调加密密钥 |
|
||||||
| verification_token | string | 否 | 用于Webhook事件验证的Token |
|
| verification_token | string | 否 | 用于Webhook事件验证的Token |
|
||||||
| allow_from | array | 否 | 用户ID白名单,空表示允许所有用户 |
|
| allow_from | array | 否 | 用户ID白名单,空表示所有用户 |
|
||||||
|
| random_reaction_emoji | array | 否 | 随机添加的表情列表,空则使用默认 "Pin" |
|
||||||
|
|
||||||
## 设置流程
|
## 设置流程
|
||||||
|
|
||||||
|
|
@ -35,3 +36,4 @@
|
||||||
3. 配置事件订阅和Webhook URL
|
3. 配置事件订阅和Webhook URL
|
||||||
4. 设置加密(可选,生产环境建议启用)
|
4. 设置加密(可选,生产环境建议启用)
|
||||||
5. 将 App ID、App Secret、Encrypt Key 和 Verification Token(如果启用加密) 填入配置文件中
|
5. 将 App ID、App Secret、Encrypt Key 和 Verification Token(如果启用加密) 填入配置文件中
|
||||||
|
6. 自定义你希望 PicoClaw react 你消息时的表情(可选, Reference URL: [Feishu Emoji List](https://open.larkoffice.com/document/server-docs/im-v1/message-reaction/emojis-introduce))
|
||||||
|
|
|
||||||
59
docs/channels/matrix/README.md
Normal file
59
docs/channels/matrix/README.md
Normal file
|
|
@ -0,0 +1,59 @@
|
||||||
|
# Matrix Channel Configuration Guide
|
||||||
|
|
||||||
|
## 1. Example Configuration
|
||||||
|
|
||||||
|
Add this to `config.json`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"matrix": {
|
||||||
|
"enabled": true,
|
||||||
|
"homeserver": "https://matrix.org",
|
||||||
|
"user_id": "@your-bot:matrix.org",
|
||||||
|
"access_token": "YOUR_MATRIX_ACCESS_TOKEN",
|
||||||
|
"device_id": "",
|
||||||
|
"join_on_invite": true,
|
||||||
|
"allow_from": [],
|
||||||
|
"group_trigger": {
|
||||||
|
"mention_only": true
|
||||||
|
},
|
||||||
|
"placeholder": {
|
||||||
|
"enabled": true,
|
||||||
|
"text": "Thinking..."
|
||||||
|
},
|
||||||
|
"reasoning_channel_id": ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 2. Field Reference
|
||||||
|
|
||||||
|
| Field | Type | Required | Description |
|
||||||
|
|----------------------|----------|----------|-------------|
|
||||||
|
| enabled | bool | Yes | Enable or disable the Matrix channel |
|
||||||
|
| homeserver | string | Yes | Matrix homeserver URL (for example `https://matrix.org`) |
|
||||||
|
| user_id | string | Yes | Bot Matrix user ID (for example `@bot:matrix.org`) |
|
||||||
|
| access_token | string | Yes | Bot access token |
|
||||||
|
| device_id | string | No | Optional Matrix device ID |
|
||||||
|
| join_on_invite | bool | No | Auto-join invited rooms |
|
||||||
|
| allow_from | []string | No | User whitelist (Matrix user IDs) |
|
||||||
|
| group_trigger | object | No | Group trigger strategy (`mention_only` / `prefixes`) |
|
||||||
|
| placeholder | object | No | Placeholder message config |
|
||||||
|
| reasoning_channel_id | string | No | Target channel for reasoning output |
|
||||||
|
|
||||||
|
## 3. Currently Supported
|
||||||
|
|
||||||
|
- Text message send/receive
|
||||||
|
- Incoming image/audio/video/file download (MediaStore first, local path fallback)
|
||||||
|
- Incoming audio normalization into existing transcription flow (`[audio: ...]`)
|
||||||
|
- Outgoing image/audio/video/file upload and send
|
||||||
|
- Group trigger rules (including mention-only mode)
|
||||||
|
- Typing state (`m.typing`)
|
||||||
|
- Placeholder message + final reply replacement
|
||||||
|
- Auto-join invited rooms (can be disabled)
|
||||||
|
|
||||||
|
## 4. TODO
|
||||||
|
|
||||||
|
- Rich media metadata improvements (for example image/video size and thumbnails)
|
||||||
59
docs/channels/matrix/README.zh.md
Normal file
59
docs/channels/matrix/README.zh.md
Normal file
|
|
@ -0,0 +1,59 @@
|
||||||
|
# Matrix 通道配置指南
|
||||||
|
|
||||||
|
## 1. 配置示例
|
||||||
|
|
||||||
|
在 `config.json` 中添加:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"matrix": {
|
||||||
|
"enabled": true,
|
||||||
|
"homeserver": "https://matrix.org",
|
||||||
|
"user_id": "@your-bot:matrix.org",
|
||||||
|
"access_token": "YOUR_MATRIX_ACCESS_TOKEN",
|
||||||
|
"device_id": "",
|
||||||
|
"join_on_invite": true,
|
||||||
|
"allow_from": [],
|
||||||
|
"group_trigger": {
|
||||||
|
"mention_only": true
|
||||||
|
},
|
||||||
|
"placeholder": {
|
||||||
|
"enabled": true,
|
||||||
|
"text": "Thinking... 💭"
|
||||||
|
},
|
||||||
|
"reasoning_channel_id": ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 2. 参数说明
|
||||||
|
|
||||||
|
| 字段 | 类型 | 必填 | 说明 |
|
||||||
|
|----------------------|----------|------|------|
|
||||||
|
| enabled | bool | 是 | 是否启用 Matrix 通道 |
|
||||||
|
| homeserver | string | 是 | Matrix 服务器地址(例如 `https://matrix.org`) |
|
||||||
|
| user_id | string | 是 | 机器人 Matrix 用户 ID(例如 `@bot:matrix.org`) |
|
||||||
|
| access_token | string | 是 | 机器人 access token |
|
||||||
|
| device_id | string | 否 | 设备 ID(可选) |
|
||||||
|
| join_on_invite | bool | 否 | 是否自动加入邀请房间 |
|
||||||
|
| allow_from | []string | 否 | 白名单用户(Matrix 用户 ID) |
|
||||||
|
| group_trigger | object | 否 | 群聊触发策略(支持 `mention_only` / `prefixes`) |
|
||||||
|
| placeholder | object | 否 | 占位消息配置 |
|
||||||
|
| reasoning_channel_id | string | 否 | 思维链输出目标通道 |
|
||||||
|
|
||||||
|
## 3. 当前支持
|
||||||
|
|
||||||
|
- 文本消息收发
|
||||||
|
- 图片/音频/视频/文件消息入站下载(写入 MediaStore / 本地路径回退)
|
||||||
|
- 音频消息按统一标记进入现有转写流程(`[audio: ...]`)
|
||||||
|
- 图片/音频/视频/文件消息出站发送(上传到 Matrix 媒体库后发送)
|
||||||
|
- 群聊触发规则(支持仅 @ 提及时响应)
|
||||||
|
- Typing 状态(`m.typing`)
|
||||||
|
- 占位消息(`Thinking... 💭`)+ 最终回复替换
|
||||||
|
- 自动加入邀请房间(可关闭)
|
||||||
|
|
||||||
|
## 4. TODO
|
||||||
|
|
||||||
|
- 富媒体细节增强(如 image/video 的尺寸、缩略图等 metadata)
|
||||||
|
|
@ -7,11 +7,21 @@ PicoClaw's tools configuration is located in the `tools` field of `config.json`.
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": { ... },
|
"web": {
|
||||||
"mcp": { ... },
|
...
|
||||||
"exec": { ... },
|
},
|
||||||
"cron": { ... },
|
"mcp": {
|
||||||
"skills": { ... }
|
...
|
||||||
|
},
|
||||||
|
"exec": {
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"cron": {
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"skills": {
|
||||||
|
...
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
@ -23,7 +33,7 @@ Web tools are used for web search and fetching.
|
||||||
### Brave
|
### Brave
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ------------- | ------ | ------- | ------------------------- |
|
|---------------|--------|---------|---------------------------|
|
||||||
| `enabled` | bool | false | Enable Brave search |
|
| `enabled` | bool | false | Enable Brave search |
|
||||||
| `api_key` | string | - | Brave Search API key |
|
| `api_key` | string | - | Brave Search API key |
|
||||||
| `max_results` | int | 5 | Maximum number of results |
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
@ -31,14 +41,14 @@ Web tools are used for web search and fetching.
|
||||||
### DuckDuckGo
|
### DuckDuckGo
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ------------- | ---- | ------- | ------------------------- |
|
|---------------|------|---------|---------------------------|
|
||||||
| `enabled` | bool | true | Enable DuckDuckGo search |
|
| `enabled` | bool | true | Enable DuckDuckGo search |
|
||||||
| `max_results` | int | 5 | Maximum number of results |
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
||||||
### Perplexity
|
### Perplexity
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ------------- | ------ | ------- | ------------------------- |
|
|---------------|--------|---------|---------------------------|
|
||||||
| `enabled` | bool | false | Enable Perplexity search |
|
| `enabled` | bool | false | Enable Perplexity search |
|
||||||
| `api_key` | string | - | Perplexity API key |
|
| `api_key` | string | - | Perplexity API key |
|
||||||
| `max_results` | int | 5 | Maximum number of results |
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
@ -48,7 +58,7 @@ Web tools are used for web search and fetching.
|
||||||
The exec tool is used to execute shell commands.
|
The exec tool is used to execute shell commands.
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ---------------------- | ----- | ------- | ------------------------------------------ |
|
|------------------------|-------|---------|--------------------------------------------|
|
||||||
| `enable_deny_patterns` | bool | true | Enable default dangerous command blocking |
|
| `enable_deny_patterns` | bool | true | Enable default dangerous command blocking |
|
||||||
| `custom_deny_patterns` | array | [] | Custom deny patterns (regular expressions) |
|
| `custom_deny_patterns` | array | [] | Custom deny patterns (regular expressions) |
|
||||||
|
|
||||||
|
|
@ -81,7 +91,10 @@ By default, PicoClaw blocks the following dangerous commands:
|
||||||
"tools": {
|
"tools": {
|
||||||
"exec": {
|
"exec": {
|
||||||
"enable_deny_patterns": true,
|
"enable_deny_patterns": true,
|
||||||
"custom_deny_patterns": ["\\brm\\s+-r\\b", "\\bkillall\\s+python"]
|
"custom_deny_patterns": [
|
||||||
|
"\\brm\\s+-r\\b",
|
||||||
|
"\\bkillall\\s+python"
|
||||||
|
]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -92,24 +105,47 @@ By default, PicoClaw blocks the following dangerous commands:
|
||||||
The cron tool is used for scheduling periodic tasks.
|
The cron tool is used for scheduling periodic tasks.
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ---------------------- | ---- | ------- | ---------------------------------------------- |
|
|------------------------|------|---------|------------------------------------------------|
|
||||||
| `exec_timeout_minutes` | int | 5 | Execution timeout in minutes, 0 means no limit |
|
| `exec_timeout_minutes` | int | 5 | Execution timeout in minutes, 0 means no limit |
|
||||||
|
|
||||||
## MCP Tool
|
## MCP Tool
|
||||||
|
|
||||||
The MCP tool enables integration with external Model Context Protocol servers.
|
The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
|
|
||||||
|
### Tool Discovery (Lazy Loading)
|
||||||
|
|
||||||
|
When connecting to multiple MCP servers, exposing hundreds of tools simultaneously can exhaust the LLM's context window
|
||||||
|
and increase API costs. The **Discovery** feature solves this by keeping MCP tools *hidden* by default.
|
||||||
|
|
||||||
|
Instead of loading all tools, the LLM is provided with a lightweight search tool (using BM25 keyword matching or Regex).
|
||||||
|
When the LLM needs a specific capability, it searches the hidden library. Matching tools are then temporarily "unlocked"
|
||||||
|
and injected into the context for a configured number of turns (`ttl`).
|
||||||
|
|
||||||
### Global Config
|
### Global Config
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| --------- | ------ | ------- | ----------------------------------- |
|
|-------------|--------|---------|----------------------------------------------|
|
||||||
| `enabled` | bool | false | Enable MCP integration globally |
|
| `enabled` | bool | false | Enable MCP integration globally |
|
||||||
| `servers` | object | `{}` | Map of server name to server config |
|
| `discovery` | object | `{}` | Configuration for Tool Discovery (see below) |
|
||||||
|
| `servers` | object | `{}` | Map of server name to server config |
|
||||||
|
|
||||||
|
### Discovery Config (`discovery`)
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|----------------------|------|---------|-----------------------------------------------------------------------------------------------------------------------------------|
|
||||||
|
| `enabled` | bool | false | If true, MCP tools are hidden and loaded on-demand via search. If false, all tools are loaded |
|
||||||
|
| `ttl` | int | 5 | Number of conversational turns a discovered tool remains unlocked |
|
||||||
|
| `max_search_results` | int | 5 | Maximum number of tools returned per search query |
|
||||||
|
| `use_bm25` | bool | true | Enable the natural language/keyword search tool (`tool_search_tool_bm25`). **Warning**: consumes more resources than regex search |
|
||||||
|
| `use_regex` | bool | false | Enable the regex pattern search tool (`tool_search_tool_regex`) |
|
||||||
|
|
||||||
|
> **Note:** If `discovery.enabled` is `true`, you MUST enable at least one search engine (`use_bm25` or `use_regex`),
|
||||||
|
> otherwise the application will fail to start.
|
||||||
|
|
||||||
### Per-Server Config
|
### Per-Server Config
|
||||||
|
|
||||||
| Config | Type | Required | Description |
|
| Config | Type | Required | Description |
|
||||||
| ---------- | ------ | -------- | ------------------------------------------ |
|
|------------|--------|----------|--------------------------------------------|
|
||||||
| `enabled` | bool | yes | Enable this MCP server |
|
| `enabled` | bool | yes | Enable this MCP server |
|
||||||
| `type` | string | no | Transport type: `stdio`, `sse`, `http` |
|
| `type` | string | no | Transport type: `stdio`, `sse`, `http` |
|
||||||
| `command` | string | stdio | Executable command for stdio transport |
|
| `command` | string | stdio | Executable command for stdio transport |
|
||||||
|
|
@ -122,8 +158,8 @@ The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
### Transport Behavior
|
### Transport Behavior
|
||||||
|
|
||||||
- If `type` is omitted, transport is auto-detected:
|
- If `type` is omitted, transport is auto-detected:
|
||||||
- `url` is set → `sse`
|
- `url` is set → `sse`
|
||||||
- `command` is set → `stdio`
|
- `command` is set → `stdio`
|
||||||
- `http` and `sse` both use `url` + optional `headers`.
|
- `http` and `sse` both use `url` + optional `headers`.
|
||||||
- `env` and `env_file` are only applied to `stdio` servers.
|
- `env` and `env_file` are only applied to `stdio` servers.
|
||||||
|
|
||||||
|
|
@ -140,7 +176,11 @@ The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
"filesystem": {
|
"filesystem": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"command": "npx",
|
"command": "npx",
|
||||||
"args": ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"]
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-filesystem",
|
||||||
|
"/tmp"
|
||||||
|
]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -170,20 +210,76 @@ The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### 3) Massive MCP setup with Tool Discovery enabled
|
||||||
|
|
||||||
|
*In this example, the LLM will only see the `tool_search_tool_bm25`. It will search and unlock Github or Postgres tools
|
||||||
|
dynamically only when requested by the user.*
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"mcp": {
|
||||||
|
"enabled": true,
|
||||||
|
"discovery": {
|
||||||
|
"enabled": true,
|
||||||
|
"ttl": 5,
|
||||||
|
"max_search_results": 5,
|
||||||
|
"use_bm25": true,
|
||||||
|
"use_regex": false
|
||||||
|
},
|
||||||
|
"servers": {
|
||||||
|
"github": {
|
||||||
|
"enabled": true,
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-github"
|
||||||
|
],
|
||||||
|
"env": {
|
||||||
|
"GITHUB_PERSONAL_ACCESS_TOKEN": "YOUR_GITHUB_TOKEN"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"postgres": {
|
||||||
|
"enabled": true,
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-postgres",
|
||||||
|
"postgresql://user:password@localhost/dbname"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"slack": {
|
||||||
|
"enabled": true,
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-slack"
|
||||||
|
],
|
||||||
|
"env": {
|
||||||
|
"SLACK_BOT_TOKEN": "YOUR_SLACK_BOT_TOKEN",
|
||||||
|
"SLACK_TEAM_ID": "YOUR_SLACK_TEAM_ID"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
## Skills Tool
|
## Skills Tool
|
||||||
|
|
||||||
The skills tool configures skill discovery and installation via registries like ClawHub.
|
The skills tool configures skill discovery and installation via registries like ClawHub.
|
||||||
|
|
||||||
### Registries
|
### Registries
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ---------------------------------- | ------ | -------------------- | ----------------------- |
|
|------------------------------------|--------|----------------------|----------------------------------------------|
|
||||||
| `registries.clawhub.enabled` | bool | true | Enable ClawHub registry |
|
| `registries.clawhub.enabled` | bool | true | Enable ClawHub registry |
|
||||||
| `registries.clawhub.base_url` | string | `https://clawhub.ai` | ClawHub base URL |
|
| `registries.clawhub.base_url` | string | `https://clawhub.ai` | ClawHub base URL |
|
||||||
| `registries.clawhub.auth_token` | string | `""` | Optional Bearer token for higher rate limits |
|
| `registries.clawhub.auth_token` | string | `""` | Optional Bearer token for higher rate limits |
|
||||||
| `registries.clawhub.search_path` | string | `/api/v1/search` | Search API path |
|
| `registries.clawhub.search_path` | string | `/api/v1/search` | Search API path |
|
||||||
| `registries.clawhub.skills_path` | string | `/api/v1/skills` | Skills API path |
|
| `registries.clawhub.skills_path` | string | `/api/v1/skills` | Skills API path |
|
||||||
| `registries.clawhub.download_path` | string | `/api/v1/download` | Download API path |
|
| `registries.clawhub.download_path` | string | `/api/v1/download` | Download API path |
|
||||||
|
|
||||||
### Configuration Example
|
### Configuration Example
|
||||||
|
|
||||||
|
|
@ -217,4 +313,5 @@ For example:
|
||||||
- `PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES=10`
|
- `PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES=10`
|
||||||
- `PICOCLAW_TOOLS_MCP_ENABLED=true`
|
- `PICOCLAW_TOOLS_MCP_ENABLED=true`
|
||||||
|
|
||||||
Note: Nested map-style config (for example `tools.mcp.servers.<name>.*`) is configured in `config.json` rather than environment variables.
|
Note: Nested map-style config (for example `tools.mcp.servers.<name>.*`) is configured in `config.json` rather than
|
||||||
|
environment variables.
|
||||||
|
|
|
||||||
5
go.mod
5
go.mod
|
|
@ -8,6 +8,7 @@ require (
|
||||||
github.com/bwmarrin/discordgo v0.29.0
|
github.com/bwmarrin/discordgo v0.29.0
|
||||||
github.com/caarlos0/env/v11 v11.3.1
|
github.com/caarlos0/env/v11 v11.3.1
|
||||||
github.com/chzyer/readline v1.5.1
|
github.com/chzyer/readline v1.5.1
|
||||||
|
github.com/ergochat/irc-go v0.5.0
|
||||||
github.com/gdamore/tcell/v2 v2.13.8
|
github.com/gdamore/tcell/v2 v2.13.8
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
|
|
@ -27,6 +28,7 @@ require (
|
||||||
golang.org/x/oauth2 v0.35.0
|
golang.org/x/oauth2 v0.35.0
|
||||||
golang.org/x/time v0.14.0
|
golang.org/x/time v0.14.0
|
||||||
google.golang.org/protobuf v1.36.11
|
google.golang.org/protobuf v1.36.11
|
||||||
|
maunium.net/go/mautrix v0.26.3
|
||||||
modernc.org/sqlite v1.46.1
|
modernc.org/sqlite v1.46.1
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -37,7 +39,6 @@ require (
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect
|
github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect
|
||||||
github.com/ergochat/irc-go v0.5.0 // indirect
|
|
||||||
github.com/gdamore/encoding v1.0.1 // indirect
|
github.com/gdamore/encoding v1.0.1 // indirect
|
||||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||||
github.com/lucasb-eyer/go-colorful v1.3.0 // indirect
|
github.com/lucasb-eyer/go-colorful v1.3.0 // indirect
|
||||||
|
|
@ -89,7 +90,7 @@ require (
|
||||||
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
||||||
golang.org/x/arch v0.24.0 // indirect
|
golang.org/x/arch v0.24.0 // indirect
|
||||||
golang.org/x/crypto v0.48.0 // indirect
|
golang.org/x/crypto v0.48.0 // indirect
|
||||||
golang.org/x/net v0.50.0 // indirect
|
golang.org/x/net v0.51.0 // indirect
|
||||||
golang.org/x/sync v0.19.0 // indirect
|
golang.org/x/sync v0.19.0 // indirect
|
||||||
golang.org/x/sys v0.41.0 // indirect
|
golang.org/x/sys v0.41.0 // indirect
|
||||||
)
|
)
|
||||||
|
|
|
||||||
4
go.sum
4
go.sum
|
|
@ -271,6 +271,8 @@ golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
|
||||||
golang.org/x/net v0.19.0/go.mod h1:CfAk/cbD4CthTvqiEl8NpboMuiuOYsAr/7NOjZJtv1U=
|
golang.org/x/net v0.19.0/go.mod h1:CfAk/cbD4CthTvqiEl8NpboMuiuOYsAr/7NOjZJtv1U=
|
||||||
golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60=
|
golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60=
|
||||||
golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM=
|
golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM=
|
||||||
|
golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo=
|
||||||
|
golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y=
|
||||||
golang.org/x/oauth2 v0.23.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
|
golang.org/x/oauth2 v0.23.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
|
||||||
golang.org/x/oauth2 v0.35.0 h1:Mv2mzuHuZuY2+bkyWXIHMfhNdJAdwW3FuWeCPYN5GVQ=
|
golang.org/x/oauth2 v0.35.0 h1:Mv2mzuHuZuY2+bkyWXIHMfhNdJAdwW3FuWeCPYN5GVQ=
|
||||||
golang.org/x/oauth2 v0.35.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
|
golang.org/x/oauth2 v0.35.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
|
||||||
|
|
@ -361,6 +363,8 @@ gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ=
|
||||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
|
maunium.net/go/mautrix v0.26.3 h1:tWZih6Vjw0qGTWuPmg9JUrQPzViTNDPGQLVc5UXC4nk=
|
||||||
|
maunium.net/go/mautrix v0.26.3/go.mod h1:v5ZdDoCwUpNqEj5OrhEoUa3L1kEddKPaAya9TgGXN38=
|
||||||
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
||||||
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
||||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
||||||
|
|
|
||||||
|
|
@ -12,15 +12,18 @@ import (
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/skills"
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ContextBuilder struct {
|
type ContextBuilder struct {
|
||||||
workspace string
|
workspace string
|
||||||
skillsLoader *skills.SkillsLoader
|
skillsLoader *skills.SkillsLoader
|
||||||
memory *MemoryStore
|
memory *MemoryStore
|
||||||
|
toolDiscoveryBM25 bool
|
||||||
|
toolDiscoveryRegex bool
|
||||||
|
|
||||||
// Cache for system prompt to avoid rebuilding on every call.
|
// Cache for system prompt to avoid rebuilding on every call.
|
||||||
// This fixes issue #607: repeated reprocessing of the entire context.
|
// This fixes issue #607: repeated reprocessing of the entire context.
|
||||||
|
|
@ -41,6 +44,12 @@ type ContextBuilder struct {
|
||||||
skillFilesAtCache map[string]time.Time
|
skillFilesAtCache map[string]time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) WithToolDiscovery(useBM25, useRegex bool) *ContextBuilder {
|
||||||
|
cb.toolDiscoveryBM25 = useBM25
|
||||||
|
cb.toolDiscoveryRegex = useRegex
|
||||||
|
return cb
|
||||||
|
}
|
||||||
|
|
||||||
func getGlobalConfigDir() string {
|
func getGlobalConfigDir() string {
|
||||||
if home := os.Getenv("PICOCLAW_HOME"); home != "" {
|
if home := os.Getenv("PICOCLAW_HOME"); home != "" {
|
||||||
return home
|
return home
|
||||||
|
|
@ -71,8 +80,11 @@ func NewContextBuilder(workspace string) *ContextBuilder {
|
||||||
|
|
||||||
func (cb *ContextBuilder) getIdentity() string {
|
func (cb *ContextBuilder) getIdentity() string {
|
||||||
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
||||||
|
toolDiscovery := cb.getDiscoveryRule()
|
||||||
|
version := config.FormatVersion()
|
||||||
|
|
||||||
return fmt.Sprintf(`# picoclaw 🦞
|
return fmt.Sprintf(
|
||||||
|
`# picoclaw 🦞 (%s)
|
||||||
|
|
||||||
You are picoclaw, a helpful AI assistant.
|
You are picoclaw, a helpful AI assistant.
|
||||||
|
|
||||||
|
|
@ -90,8 +102,29 @@ Your workspace is at: %s
|
||||||
|
|
||||||
3. **Memory** - When interacting with me if something seems memorable, update %s/memory/MEMORY.md
|
3. **Memory** - When interacting with me if something seems memorable, update %s/memory/MEMORY.md
|
||||||
|
|
||||||
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.`,
|
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.
|
||||||
workspacePath, workspacePath, workspacePath, workspacePath, workspacePath)
|
|
||||||
|
%s`,
|
||||||
|
version, workspacePath, workspacePath, workspacePath, workspacePath, workspacePath, toolDiscovery)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) getDiscoveryRule() string {
|
||||||
|
if !cb.toolDiscoveryBM25 && !cb.toolDiscoveryRegex {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var toolNames []string
|
||||||
|
if cb.toolDiscoveryBM25 {
|
||||||
|
toolNames = append(toolNames, `"tool_search_tool_bm25"`)
|
||||||
|
}
|
||||||
|
if cb.toolDiscoveryRegex {
|
||||||
|
toolNames = append(toolNames, `"tool_search_tool_regex"`)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf(
|
||||||
|
`5. **Tool Discovery** - Your visible tools are limited to save memory, but a vast hidden library exists. If you lack the right tool for a task, BEFORE giving up, you MUST search using the %s tool. Do not refuse a request unless the search returns nothing. Found tools will temporarily unlock for your next turn.`,
|
||||||
|
strings.Join(toolNames, " or "),
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) BuildSystemPrompt() string {
|
func (cb *ContextBuilder) BuildSystemPrompt() string {
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
package agent
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
|
|
@ -9,6 +10,7 @@ import (
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/memory"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/routing"
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
"github.com/sipeed/picoclaw/pkg/session"
|
"github.com/sipeed/picoclaw/pkg/session"
|
||||||
|
|
@ -31,7 +33,7 @@ type AgentInstance struct {
|
||||||
SummarizeMessageThreshold int
|
SummarizeMessageThreshold int
|
||||||
SummarizeTokenPercent int
|
SummarizeTokenPercent int
|
||||||
Provider providers.LLMProvider
|
Provider providers.LLMProvider
|
||||||
Sessions *session.SessionManager
|
Sessions session.SessionStore
|
||||||
ContextBuilder *ContextBuilder
|
ContextBuilder *ContextBuilder
|
||||||
Tools *tools.ToolRegistry
|
Tools *tools.ToolRegistry
|
||||||
Subagents *config.SubagentsConfig
|
Subagents *config.SubagentsConfig
|
||||||
|
|
@ -70,7 +72,8 @@ func NewAgentInstance(
|
||||||
toolsRegistry := tools.NewToolRegistry()
|
toolsRegistry := tools.NewToolRegistry()
|
||||||
|
|
||||||
if cfg.Tools.IsToolEnabled("read_file") {
|
if cfg.Tools.IsToolEnabled("read_file") {
|
||||||
toolsRegistry.Register(tools.NewReadFileTool(workspace, readRestrict, allowReadPaths))
|
maxReadFileSize := cfg.Tools.ReadFile.MaxReadFileSize
|
||||||
|
toolsRegistry.Register(tools.NewReadFileTool(workspace, readRestrict, maxReadFileSize, allowReadPaths))
|
||||||
}
|
}
|
||||||
if cfg.Tools.IsToolEnabled("write_file") {
|
if cfg.Tools.IsToolEnabled("write_file") {
|
||||||
toolsRegistry.Register(tools.NewWriteFileTool(workspace, restrict, allowWritePaths))
|
toolsRegistry.Register(tools.NewWriteFileTool(workspace, restrict, allowWritePaths))
|
||||||
|
|
@ -94,9 +97,13 @@ func NewAgentInstance(
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionsDir := filepath.Join(workspace, "sessions")
|
sessionsDir := filepath.Join(workspace, "sessions")
|
||||||
sessionsManager := session.NewSessionManager(sessionsDir)
|
sessions := initSessionStore(sessionsDir)
|
||||||
|
|
||||||
contextBuilder := NewContextBuilder(workspace)
|
mcpDiscoveryActive := cfg.Tools.MCP.Enabled && cfg.Tools.MCP.Discovery.Enabled
|
||||||
|
contextBuilder := NewContextBuilder(workspace).WithToolDiscovery(
|
||||||
|
mcpDiscoveryActive && cfg.Tools.MCP.Discovery.UseBM25,
|
||||||
|
mcpDiscoveryActive && cfg.Tools.MCP.Discovery.UseRegex,
|
||||||
|
)
|
||||||
|
|
||||||
agentID := routing.DefaultAgentID
|
agentID := routing.DefaultAgentID
|
||||||
agentName := ""
|
agentName := ""
|
||||||
|
|
@ -221,7 +228,7 @@ func NewAgentInstance(
|
||||||
SummarizeMessageThreshold: summarizeMessageThreshold,
|
SummarizeMessageThreshold: summarizeMessageThreshold,
|
||||||
SummarizeTokenPercent: summarizeTokenPercent,
|
SummarizeTokenPercent: summarizeTokenPercent,
|
||||||
Provider: provider,
|
Provider: provider,
|
||||||
Sessions: sessionsManager,
|
Sessions: sessions,
|
||||||
ContextBuilder: contextBuilder,
|
ContextBuilder: contextBuilder,
|
||||||
Tools: toolsRegistry,
|
Tools: toolsRegistry,
|
||||||
Subagents: subagents,
|
Subagents: subagents,
|
||||||
|
|
@ -275,6 +282,39 @@ func compilePatterns(patterns []string) []*regexp.Regexp {
|
||||||
return compiled
|
return compiled
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Close releases resources held by the agent's session store.
|
||||||
|
func (a *AgentInstance) Close() error {
|
||||||
|
if a.Sessions != nil {
|
||||||
|
return a.Sessions.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// initSessionStore creates the session persistence backend.
|
||||||
|
// It uses the JSONL store by default and auto-migrates legacy JSON sessions.
|
||||||
|
// Falls back to SessionManager if the JSONL store cannot be initialized or
|
||||||
|
// if migration fails (which indicates the store cannot write reliably).
|
||||||
|
func initSessionStore(dir string) session.SessionStore {
|
||||||
|
store, err := memory.NewJSONLStore(dir)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("memory: init store: %v; using json sessions", err)
|
||||||
|
return session.NewSessionManager(dir)
|
||||||
|
}
|
||||||
|
|
||||||
|
if n, merr := memory.MigrateFromJSON(context.Background(), dir, store); merr != nil {
|
||||||
|
// Migration failure means the store could not write data.
|
||||||
|
// Fall back to SessionManager to avoid a split state where
|
||||||
|
// some sessions are in JSONL and others remain in JSON.
|
||||||
|
log.Printf("memory: migration failed: %v; falling back to json sessions", merr)
|
||||||
|
store.Close()
|
||||||
|
return session.NewSessionManager(dir)
|
||||||
|
} else if n > 0 {
|
||||||
|
log.Printf("memory: migrated %d session(s) to jsonl", n)
|
||||||
|
}
|
||||||
|
|
||||||
|
return session.NewJSONLBackend(store)
|
||||||
|
}
|
||||||
|
|
||||||
func expandHome(path string) string {
|
func expandHome(path string) string {
|
||||||
if path == "" {
|
if path == "" {
|
||||||
return path
|
return path
|
||||||
|
|
|
||||||
|
|
@ -120,19 +120,21 @@ func registerSharedTools(
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Web tools
|
|
||||||
if cfg.Tools.IsToolEnabled("web") {
|
if cfg.Tools.IsToolEnabled("web") {
|
||||||
searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
|
searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
|
||||||
BraveAPIKey: cfg.Tools.Web.Brave.APIKey,
|
BraveAPIKeys: config.MergeAPIKeys(cfg.Tools.Web.Brave.APIKey, cfg.Tools.Web.Brave.APIKeys),
|
||||||
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
||||||
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
||||||
TavilyAPIKey: cfg.Tools.Web.Tavily.APIKey,
|
TavilyAPIKeys: config.MergeAPIKeys(cfg.Tools.Web.Tavily.APIKey, cfg.Tools.Web.Tavily.APIKeys),
|
||||||
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
||||||
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
||||||
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
||||||
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
||||||
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
||||||
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
|
PerplexityAPIKeys: config.MergeAPIKeys(
|
||||||
|
cfg.Tools.Web.Perplexity.APIKey,
|
||||||
|
cfg.Tools.Web.Perplexity.APIKeys,
|
||||||
|
),
|
||||||
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
||||||
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
||||||
SearXNGBaseURL: cfg.Tools.Web.SearXNG.BaseURL,
|
SearXNGBaseURL: cfg.Tools.Web.SearXNG.BaseURL,
|
||||||
|
|
@ -283,7 +285,13 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
|
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
|
||||||
agent.Tools.Register(mcpTool)
|
|
||||||
|
if al.cfg.Tools.MCP.Discovery.Enabled {
|
||||||
|
agent.Tools.RegisterHidden(mcpTool)
|
||||||
|
} else {
|
||||||
|
agent.Tools.Register(mcpTool)
|
||||||
|
}
|
||||||
|
|
||||||
totalRegistrations++
|
totalRegistrations++
|
||||||
logger.DebugCF("agent", "Registered MCP tool",
|
logger.DebugCF("agent", "Registered MCP tool",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
|
|
@ -302,6 +310,47 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
"total_registrations": totalRegistrations,
|
"total_registrations": totalRegistrations,
|
||||||
"agent_count": agentCount,
|
"agent_count": agentCount,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// Initializes Discovery Tools only if enabled by configuration
|
||||||
|
if al.cfg.Tools.MCP.Enabled && al.cfg.Tools.MCP.Discovery.Enabled {
|
||||||
|
useBM25 := al.cfg.Tools.MCP.Discovery.UseBM25
|
||||||
|
useRegex := al.cfg.Tools.MCP.Discovery.UseRegex
|
||||||
|
|
||||||
|
// Fail fast: If discovery is enabled but no search method is turned on
|
||||||
|
if !useBM25 && !useRegex {
|
||||||
|
return fmt.Errorf(
|
||||||
|
"tool discovery is enabled but neither 'use_bm25' nor 'use_regex' is set to true in the configuration",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
ttl := al.cfg.Tools.MCP.Discovery.TTL
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = 5 // Default value
|
||||||
|
}
|
||||||
|
|
||||||
|
maxSearchResults := al.cfg.Tools.MCP.Discovery.MaxSearchResults
|
||||||
|
if maxSearchResults <= 0 {
|
||||||
|
maxSearchResults = 5 // Default value
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.InfoCF("agent", "Initializing tool discovery", map[string]any{
|
||||||
|
"bm25": useBM25, "regex": useRegex, "ttl": ttl, "max_results": maxSearchResults,
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, agentID := range agentIDs {
|
||||||
|
agent, ok := al.registry.GetAgent(agentID)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if useRegex {
|
||||||
|
agent.Tools.Register(tools.NewRegexSearchTool(agent.Tools, ttl, maxSearchResults))
|
||||||
|
}
|
||||||
|
if useBM25 {
|
||||||
|
agent.Tools.Register(tools.NewBM25SearchTool(agent.Tools, ttl, maxSearchResults))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -380,6 +429,11 @@ func (al *AgentLoop) Stop() {
|
||||||
al.running.Store(false)
|
al.running.Store(false)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Close releases resources held by agent session stores. Call after Stop.
|
||||||
|
func (al *AgentLoop) Close() {
|
||||||
|
al.registry.Close()
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) RegisterTool(tool tools.Tool) {
|
func (al *AgentLoop) RegisterTool(tool tools.Tool) {
|
||||||
for _, agentID := range al.registry.ListAgentIDs() {
|
for _, agentID := range al.registry.ListAgentIDs() {
|
||||||
if agent, ok := al.registry.GetAgent(agentID); ok {
|
if agent, ok := al.registry.GetAgent(agentID); ok {
|
||||||
|
|
@ -581,15 +635,6 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
|
||||||
}
|
}
|
||||||
|
|
||||||
route, agent, routeErr := al.resolveMessageRoute(msg)
|
route, agent, routeErr := al.resolveMessageRoute(msg)
|
||||||
|
|
||||||
// Commands are checked before requiring a successful route.
|
|
||||||
// Global commands (/help, /show, /switch) work even when routing fails;
|
|
||||||
// context-dependent commands check their own Runtime fields and report
|
|
||||||
// "unavailable" when the required capability is nil.
|
|
||||||
if response, handled := al.handleCommand(ctx, msg, agent); handled {
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if routeErr != nil {
|
if routeErr != nil {
|
||||||
return "", routeErr
|
return "", routeErr
|
||||||
}
|
}
|
||||||
|
|
@ -615,7 +660,7 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
|
||||||
"route_channel": route.Channel,
|
"route_channel": route.Channel,
|
||||||
})
|
})
|
||||||
|
|
||||||
return al.runAgentLoop(ctx, agent, processOptions{
|
opts := processOptions{
|
||||||
SessionKey: sessionKey,
|
SessionKey: sessionKey,
|
||||||
Channel: msg.Channel,
|
Channel: msg.Channel,
|
||||||
ChatID: msg.ChatID,
|
ChatID: msg.ChatID,
|
||||||
|
|
@ -624,7 +669,15 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
|
||||||
DefaultResponse: defaultResponse,
|
DefaultResponse: defaultResponse,
|
||||||
EnableSummary: true,
|
EnableSummary: true,
|
||||||
SendResponse: false,
|
SendResponse: false,
|
||||||
})
|
}
|
||||||
|
|
||||||
|
// context-dependent commands check their own Runtime fields and report
|
||||||
|
// "unavailable" when the required capability is nil.
|
||||||
|
if response, handled := al.handleCommand(ctx, msg, agent, &opts); handled {
|
||||||
|
return response, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return al.runAgentLoop(ctx, agent, opts)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) resolveMessageRoute(msg bus.InboundMessage) (routing.ResolvedRoute, *AgentInstance, error) {
|
func (al *AgentLoop) resolveMessageRoute(msg bus.InboundMessage) (routing.ResolvedRoute, *AgentInstance, error) {
|
||||||
|
|
@ -1255,6 +1308,17 @@ func (al *AgentLoop) runLLMIteration(
|
||||||
// Save tool result message to session
|
// Save tool result message to session
|
||||||
agent.Sessions.AddFullMessage(opts.SessionKey, toolResultMsg)
|
agent.Sessions.AddFullMessage(opts.SessionKey, toolResultMsg)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Tick down TTL of discovered tools after processing tool results.
|
||||||
|
// Only reached when tool calls were made (the loop continues);
|
||||||
|
// the break on no-tool-call responses skips this.
|
||||||
|
// NOTE: This is safe because processMessage is sequential per agent.
|
||||||
|
// If per-agent concurrency is added, TTL consistency between
|
||||||
|
// ToProviderDefs and Get must be re-evaluated.
|
||||||
|
agent.Tools.TickTTL()
|
||||||
|
logger.DebugCF("agent", "TTL tick after tool execution", map[string]any{
|
||||||
|
"agent_id": agent.ID, "iteration": iteration,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
return finalContent, iteration, nil
|
return finalContent, iteration, nil
|
||||||
|
|
@ -1492,10 +1556,20 @@ func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxSummarizationMessages = 10
|
||||||
|
llmMaxRetries = 3
|
||||||
|
llmTemperature = 0.3
|
||||||
|
fallbackMaxContentLength = 200
|
||||||
|
)
|
||||||
|
|
||||||
// Multi-Part Summarization
|
// Multi-Part Summarization
|
||||||
var finalSummary string
|
var finalSummary string
|
||||||
if len(validMessages) > 10 {
|
if len(validMessages) > maxSummarizationMessages {
|
||||||
mid := len(validMessages) / 2
|
mid := len(validMessages) / 2
|
||||||
|
|
||||||
|
mid = al.findNearestUserMessage(validMessages, mid)
|
||||||
|
|
||||||
part1 := validMessages[:mid]
|
part1 := validMessages[:mid]
|
||||||
part2 := validMessages[mid:]
|
part2 := validMessages[mid:]
|
||||||
|
|
||||||
|
|
@ -1507,18 +1581,9 @@ func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string) {
|
||||||
s1,
|
s1,
|
||||||
s2,
|
s2,
|
||||||
)
|
)
|
||||||
resp, err := agent.Provider.Chat(
|
|
||||||
ctx,
|
resp, err := al.retryLLMCall(ctx, agent, mergePrompt, llmMaxRetries)
|
||||||
[]providers.Message{{Role: "user", Content: mergePrompt}},
|
if err == nil && resp.Content != "" {
|
||||||
nil,
|
|
||||||
agent.Model,
|
|
||||||
map[string]any{
|
|
||||||
"max_tokens": 1024,
|
|
||||||
"temperature": 0.3,
|
|
||||||
"prompt_cache_key": agent.ID,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
if err == nil {
|
|
||||||
finalSummary = resp.Content
|
finalSummary = resp.Content
|
||||||
} else {
|
} else {
|
||||||
finalSummary = s1 + " " + s2
|
finalSummary = s1 + " " + s2
|
||||||
|
|
@ -1538,6 +1603,68 @@ func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// findNearestUserMessage finds the nearest user message to the given index.
|
||||||
|
// It searches backward first, then forward if no user message is found.
|
||||||
|
func (al *AgentLoop) findNearestUserMessage(messages []providers.Message, mid int) int {
|
||||||
|
originalMid := mid
|
||||||
|
|
||||||
|
for mid > 0 && messages[mid].Role != "user" {
|
||||||
|
mid--
|
||||||
|
}
|
||||||
|
|
||||||
|
if messages[mid].Role == "user" {
|
||||||
|
return mid
|
||||||
|
}
|
||||||
|
|
||||||
|
mid = originalMid
|
||||||
|
for mid < len(messages) && messages[mid].Role != "user" {
|
||||||
|
mid++
|
||||||
|
}
|
||||||
|
|
||||||
|
if mid < len(messages) {
|
||||||
|
return mid
|
||||||
|
}
|
||||||
|
|
||||||
|
return originalMid
|
||||||
|
}
|
||||||
|
|
||||||
|
// retryLLMCall calls the LLM with retry logic.
|
||||||
|
func (al *AgentLoop) retryLLMCall(
|
||||||
|
ctx context.Context,
|
||||||
|
agent *AgentInstance,
|
||||||
|
prompt string,
|
||||||
|
maxRetries int,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
const (
|
||||||
|
llmTemperature = 0.3
|
||||||
|
)
|
||||||
|
|
||||||
|
var resp *providers.LLMResponse
|
||||||
|
var err error
|
||||||
|
|
||||||
|
for attempt := 0; attempt < maxRetries; attempt++ {
|
||||||
|
resp, err = agent.Provider.Chat(
|
||||||
|
ctx,
|
||||||
|
[]providers.Message{{Role: "user", Content: prompt}},
|
||||||
|
nil,
|
||||||
|
agent.Model,
|
||||||
|
map[string]any{
|
||||||
|
"max_tokens": agent.MaxTokens,
|
||||||
|
"temperature": llmTemperature,
|
||||||
|
"prompt_cache_key": agent.ID,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if err == nil && resp != nil && resp.Content != "" {
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
if attempt < maxRetries-1 {
|
||||||
|
time.Sleep(time.Duration(attempt+1) * 100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return resp, err
|
||||||
|
}
|
||||||
|
|
||||||
// summarizeBatch summarizes a batch of messages.
|
// summarizeBatch summarizes a batch of messages.
|
||||||
func (al *AgentLoop) summarizeBatch(
|
func (al *AgentLoop) summarizeBatch(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
|
|
@ -1545,6 +1672,13 @@ func (al *AgentLoop) summarizeBatch(
|
||||||
batch []providers.Message,
|
batch []providers.Message,
|
||||||
existingSummary string,
|
existingSummary string,
|
||||||
) (string, error) {
|
) (string, error) {
|
||||||
|
const (
|
||||||
|
llmMaxRetries = 3
|
||||||
|
llmTemperature = 0.3
|
||||||
|
fallbackMinContentLength = 200
|
||||||
|
fallbackMaxContentPercent = 10
|
||||||
|
)
|
||||||
|
|
||||||
var sb strings.Builder
|
var sb strings.Builder
|
||||||
sb.WriteString(
|
sb.WriteString(
|
||||||
"Provide a concise summary of this conversation segment, preserving core context and key points.\n",
|
"Provide a concise summary of this conversation segment, preserving core context and key points.\n",
|
||||||
|
|
@ -1560,21 +1694,40 @@ func (al *AgentLoop) summarizeBatch(
|
||||||
}
|
}
|
||||||
prompt := sb.String()
|
prompt := sb.String()
|
||||||
|
|
||||||
response, err := agent.Provider.Chat(
|
response, err := al.retryLLMCall(ctx, agent, prompt, llmMaxRetries)
|
||||||
ctx,
|
if err == nil && response.Content != "" {
|
||||||
[]providers.Message{{Role: "user", Content: prompt}},
|
return strings.TrimSpace(response.Content), nil
|
||||||
nil,
|
|
||||||
agent.Model,
|
|
||||||
map[string]any{
|
|
||||||
"max_tokens": 1024,
|
|
||||||
"temperature": 0.3,
|
|
||||||
"prompt_cache_key": agent.ID,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
}
|
||||||
return response.Content, nil
|
|
||||||
|
var fallback strings.Builder
|
||||||
|
fallback.WriteString("Conversation summary: ")
|
||||||
|
for i, m := range batch {
|
||||||
|
if i > 0 {
|
||||||
|
fallback.WriteString(" | ")
|
||||||
|
}
|
||||||
|
content := strings.TrimSpace(m.Content)
|
||||||
|
runes := []rune(content)
|
||||||
|
if len(runes) == 0 {
|
||||||
|
fallback.WriteString(fmt.Sprintf("%s: ", m.Role))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
keepLength := len(runes) * fallbackMaxContentPercent / 100
|
||||||
|
if keepLength < fallbackMinContentLength {
|
||||||
|
keepLength = fallbackMinContentLength
|
||||||
|
}
|
||||||
|
|
||||||
|
if keepLength > len(runes) {
|
||||||
|
keepLength = len(runes)
|
||||||
|
}
|
||||||
|
|
||||||
|
content = string(runes[:keepLength])
|
||||||
|
if keepLength < len(runes) {
|
||||||
|
content += "..."
|
||||||
|
}
|
||||||
|
fallback.WriteString(fmt.Sprintf("%s: %s", m.Role, content))
|
||||||
|
}
|
||||||
|
return fallback.String(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// estimateTokens estimates the number of tokens in a message list.
|
// estimateTokens estimates the number of tokens in a message list.
|
||||||
|
|
@ -1593,6 +1746,7 @@ func (al *AgentLoop) handleCommand(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
msg bus.InboundMessage,
|
msg bus.InboundMessage,
|
||||||
agent *AgentInstance,
|
agent *AgentInstance,
|
||||||
|
opts *processOptions,
|
||||||
) (string, bool) {
|
) (string, bool) {
|
||||||
if !commands.HasCommandPrefix(msg.Content) {
|
if !commands.HasCommandPrefix(msg.Content) {
|
||||||
return "", false
|
return "", false
|
||||||
|
|
@ -1602,7 +1756,7 @@ func (al *AgentLoop) handleCommand(
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
rt := al.buildCommandsRuntime(agent)
|
rt := al.buildCommandsRuntime(agent, opts)
|
||||||
executor := commands.NewExecutor(al.cmdRegistry, rt)
|
executor := commands.NewExecutor(al.cmdRegistry, rt)
|
||||||
|
|
||||||
var commandReply string
|
var commandReply string
|
||||||
|
|
@ -1631,7 +1785,7 @@ func (al *AgentLoop) handleCommand(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance) *commands.Runtime {
|
func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, opts *processOptions) *commands.Runtime {
|
||||||
rt := &commands.Runtime{
|
rt := &commands.Runtime{
|
||||||
Config: al.cfg,
|
Config: al.cfg,
|
||||||
ListAgentIDs: al.registry.ListAgentIDs,
|
ListAgentIDs: al.registry.ListAgentIDs,
|
||||||
|
|
@ -1661,6 +1815,20 @@ func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance) *commands.Runtim
|
||||||
agent.Model = value
|
agent.Model = value
|
||||||
return oldModel, nil
|
return oldModel, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
rt.ClearHistory = func() error {
|
||||||
|
if opts == nil {
|
||||||
|
return fmt.Errorf("process options not available")
|
||||||
|
}
|
||||||
|
if agent.Sessions == nil {
|
||||||
|
return fmt.Errorf("sessions not initialized for agent")
|
||||||
|
}
|
||||||
|
|
||||||
|
agent.Sessions.SetHistory(opts.SessionKey, make([]providers.Message, 0))
|
||||||
|
agent.Sessions.SetSummary(opts.SessionKey, "")
|
||||||
|
agent.Sessions.Save(opts.SessionKey)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return rt
|
return rt
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -114,6 +114,18 @@ func (r *AgentRegistry) ForEachTool(name string, fn func(tools.Tool)) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Close releases resources held by all registered agents.
|
||||||
|
func (r *AgentRegistry) Close() {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
for _, agent := range r.agents {
|
||||||
|
if err := agent.Close(); err != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to close agent",
|
||||||
|
map[string]any{"agent_id": agent.ID, "error": err.Error()})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// GetDefaultAgent returns the default agent instance.
|
// GetDefaultAgent returns the default agent instance.
|
||||||
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
func (r *AgentRegistry) GetDefaultAgent() *AgentInstance {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"math/rand"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
|
@ -195,18 +196,30 @@ func (c *FeishuChannel) SendPlaceholder(ctx context.Context, chatID string) (str
|
||||||
}
|
}
|
||||||
|
|
||||||
// ReactToMessage implements channels.ReactionCapable.
|
// ReactToMessage implements channels.ReactionCapable.
|
||||||
// Adds an "Pin" reaction and returns an undo function to remove it.
|
// Adds a reaction (randomly chosen from config) and returns an undo function to remove it.
|
||||||
func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID string) (func(), error) {
|
func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID string) (func(), error) {
|
||||||
|
// Get emoji list from config
|
||||||
|
emojiList := c.config.RandomReactionEmoji
|
||||||
|
var chosenEmoji string
|
||||||
|
if len(emojiList) == 0 {
|
||||||
|
// Default to "Pin" if no config
|
||||||
|
chosenEmoji = "Pin"
|
||||||
|
} else {
|
||||||
|
idx := rand.Intn(len(emojiList))
|
||||||
|
chosenEmoji = emojiList[idx]
|
||||||
|
}
|
||||||
|
|
||||||
req := larkim.NewCreateMessageReactionReqBuilder().
|
req := larkim.NewCreateMessageReactionReqBuilder().
|
||||||
MessageId(messageID).
|
MessageId(messageID).
|
||||||
Body(larkim.NewCreateMessageReactionReqBodyBuilder().
|
Body(larkim.NewCreateMessageReactionReqBodyBuilder().
|
||||||
ReactionType(larkim.NewEmojiBuilder().EmojiType("Pin").Build()).
|
ReactionType(larkim.NewEmojiBuilder().EmojiType(chosenEmoji).Build()).
|
||||||
Build()).
|
Build()).
|
||||||
Build()
|
Build()
|
||||||
|
|
||||||
resp, err := c.client.Im.V1.MessageReaction.Create(ctx, req)
|
resp, err := c.client.Im.V1.MessageReaction.Create(ctx, req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("feishu", "Failed to add reaction", map[string]any{
|
logger.ErrorCF("feishu", "Failed to add reaction", map[string]any{
|
||||||
|
"emoji": chosenEmoji,
|
||||||
"message_id": messageID,
|
"message_id": messageID,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
|
|
@ -214,6 +227,7 @@ func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID st
|
||||||
}
|
}
|
||||||
if !resp.Success() {
|
if !resp.Success() {
|
||||||
logger.ErrorCF("feishu", "Reaction API error", map[string]any{
|
logger.ErrorCF("feishu", "Reaction API error", map[string]any{
|
||||||
|
"emoji": chosenEmoji,
|
||||||
"message_id": messageID,
|
"message_id": messageID,
|
||||||
"code": resp.Code,
|
"code": resp.Code,
|
||||||
"msg": resp.Msg,
|
"msg": resp.Msg,
|
||||||
|
|
|
||||||
|
|
@ -61,7 +61,9 @@ var channelRateConfig = map[string]float64{
|
||||||
"telegram": 20,
|
"telegram": 20,
|
||||||
"discord": 1,
|
"discord": 1,
|
||||||
"slack": 1,
|
"slack": 1,
|
||||||
|
"matrix": 2,
|
||||||
"line": 10,
|
"line": 10,
|
||||||
|
"qq": 5,
|
||||||
"irc": 2,
|
"irc": 2,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -244,6 +246,13 @@ func (m *Manager) initChannels() error {
|
||||||
m.initChannel("slack", "Slack")
|
m.initChannel("slack", "Slack")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if m.config.Channels.Matrix.Enabled &&
|
||||||
|
m.config.Channels.Matrix.Homeserver != "" &&
|
||||||
|
m.config.Channels.Matrix.UserID != "" &&
|
||||||
|
m.config.Channels.Matrix.AccessToken != "" {
|
||||||
|
m.initChannel("matrix", "Matrix")
|
||||||
|
}
|
||||||
|
|
||||||
if m.config.Channels.LINE.Enabled && m.config.Channels.LINE.ChannelAccessToken != "" {
|
if m.config.Channels.LINE.Enabled && m.config.Channels.LINE.ChannelAccessToken != "" {
|
||||||
m.initChannel("line", "LINE")
|
m.initChannel("line", "LINE")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
13
pkg/channels/matrix/init.go
Normal file
13
pkg/channels/matrix/init.go
Normal file
|
|
@ -0,0 +1,13 @@
|
||||||
|
package matrix
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
channels.RegisterFactory("matrix", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||||
|
return NewMatrixChannel(cfg.Channels.Matrix, b)
|
||||||
|
})
|
||||||
|
}
|
||||||
1115
pkg/channels/matrix/matrix.go
Normal file
1115
pkg/channels/matrix/matrix.go
Normal file
File diff suppressed because it is too large
Load diff
291
pkg/channels/matrix/matrix_test.go
Normal file
291
pkg/channels/matrix/matrix_test.go
Normal file
|
|
@ -0,0 +1,291 @@
|
||||||
|
package matrix
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"maunium.net/go/mautrix"
|
||||||
|
"maunium.net/go/mautrix/event"
|
||||||
|
"maunium.net/go/mautrix/id"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMatrixLocalpartMentionRegexp(t *testing.T) {
|
||||||
|
re := localpartMentionRegexp("picoclaw")
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
text string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{text: "@picoclaw hello", want: true},
|
||||||
|
{text: "hi @picoclaw:matrix.org", want: true},
|
||||||
|
{
|
||||||
|
text: "\u6b22\u8fce\u4e00\u4e0bpicoclaw\u5c0f\u9f99\u867e",
|
||||||
|
want: false, // historical false-positive case in PR #356
|
||||||
|
},
|
||||||
|
{text: "mail test@example.com", want: false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
if got := re.MatchString(tc.text); got != tc.want {
|
||||||
|
t.Fatalf("text=%q match=%v want=%v", tc.text, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStripUserMention(t *testing.T) {
|
||||||
|
userID := id.UserID("@picoclaw:matrix.org")
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
in string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{in: "@picoclaw:matrix.org hello", want: "hello"},
|
||||||
|
{in: "@picoclaw, hello", want: "hello"},
|
||||||
|
{in: "no mention here", want: "no mention here"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
if got := stripUserMention(tc.in, userID); got != tc.want {
|
||||||
|
t.Fatalf("stripUserMention(%q)=%q want=%q", tc.in, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsBotMentioned(t *testing.T) {
|
||||||
|
ch := &MatrixChannel{
|
||||||
|
client: &mautrix.Client{
|
||||||
|
UserID: id.UserID("@picoclaw:matrix.org"),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
msg event.MessageEventContent
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "mentions field",
|
||||||
|
msg: event.MessageEventContent{
|
||||||
|
Body: "hello",
|
||||||
|
Mentions: &event.Mentions{
|
||||||
|
UserIDs: []id.UserID{id.UserID("@picoclaw:matrix.org")},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "full user id in body",
|
||||||
|
msg: event.MessageEventContent{
|
||||||
|
Body: "@picoclaw:matrix.org hello",
|
||||||
|
},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "localpart with at sign",
|
||||||
|
msg: event.MessageEventContent{
|
||||||
|
Body: "@picoclaw hello",
|
||||||
|
},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "localpart without at sign should not match",
|
||||||
|
msg: event.MessageEventContent{
|
||||||
|
Body: "\u6b22\u8fce\u4e00\u4e0bpicoclaw\u5c0f\u9f99\u867e",
|
||||||
|
},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "formatted mention href matrix.to plain",
|
||||||
|
msg: event.MessageEventContent{
|
||||||
|
Body: "hello bot",
|
||||||
|
FormattedBody: `<a href="https://matrix.to/#/@picoclaw:matrix.org">PicoClaw</a> hello`,
|
||||||
|
},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "formatted mention href matrix.to encoded",
|
||||||
|
msg: event.MessageEventContent{
|
||||||
|
Body: "hello bot",
|
||||||
|
FormattedBody: `<a href="https://matrix.to/#/%40picoclaw%3Amatrix.org">PicoClaw</a> hello`,
|
||||||
|
},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
if got := ch.isBotMentioned(&tc.msg); got != tc.want {
|
||||||
|
t.Fatalf("%s: got=%v want=%v", tc.name, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRoomKindCache_ExpiresEntries(t *testing.T) {
|
||||||
|
cache := newRoomKindCache(4, 5*time.Second)
|
||||||
|
now := time.Unix(100, 0)
|
||||||
|
cache.set("!room:matrix.org", true, now)
|
||||||
|
|
||||||
|
if got, ok := cache.get("!room:matrix.org", now.Add(2*time.Second)); !ok || !got {
|
||||||
|
t.Fatalf("expected cached group room before ttl, got ok=%v group=%v", ok, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := cache.get("!room:matrix.org", now.Add(6*time.Second)); ok {
|
||||||
|
t.Fatal("expected cache miss after ttl expiry")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRoomKindCache_EvictsOldestWhenFull(t *testing.T) {
|
||||||
|
cache := newRoomKindCache(2, time.Minute)
|
||||||
|
now := time.Unix(200, 0)
|
||||||
|
|
||||||
|
cache.set("!room1:matrix.org", false, now)
|
||||||
|
cache.set("!room2:matrix.org", false, now.Add(1*time.Second))
|
||||||
|
cache.set("!room3:matrix.org", true, now.Add(2*time.Second))
|
||||||
|
|
||||||
|
if _, ok := cache.get("!room1:matrix.org", now.Add(2*time.Second)); ok {
|
||||||
|
t.Fatal("expected oldest cache entry to be evicted")
|
||||||
|
}
|
||||||
|
if got, ok := cache.get("!room2:matrix.org", now.Add(2*time.Second)); !ok || got {
|
||||||
|
t.Fatalf("expected room2 to remain and be direct, got ok=%v group=%v", ok, got)
|
||||||
|
}
|
||||||
|
if got, ok := cache.get("!room3:matrix.org", now.Add(2*time.Second)); !ok || !got {
|
||||||
|
t.Fatalf("expected room3 to remain and be group, got ok=%v group=%v", ok, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatrixMediaTempDir(t *testing.T) {
|
||||||
|
dir, err := matrixMediaTempDir()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("matrixMediaTempDir failed: %v", err)
|
||||||
|
}
|
||||||
|
if filepath.Base(dir) != matrixMediaTempDirName {
|
||||||
|
t.Fatalf("unexpected media dir base: %q", filepath.Base(dir))
|
||||||
|
}
|
||||||
|
|
||||||
|
info, err := os.Stat(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("media dir not created: %v", err)
|
||||||
|
}
|
||||||
|
if !info.IsDir() {
|
||||||
|
t.Fatalf("expected directory, got mode=%v", info.Mode())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatrixMediaExt(t *testing.T) {
|
||||||
|
if got := matrixMediaExt("photo.png", "", "image"); got != ".png" {
|
||||||
|
t.Fatalf("filename extension mismatch: got=%q", got)
|
||||||
|
}
|
||||||
|
if got := matrixMediaExt("", "image/webp", "image"); got != ".webp" {
|
||||||
|
t.Fatalf("content-type extension mismatch: got=%q", got)
|
||||||
|
}
|
||||||
|
if got := matrixMediaExt("", "", "image"); got != ".jpg" {
|
||||||
|
t.Fatalf("default image extension mismatch: got=%q", got)
|
||||||
|
}
|
||||||
|
if got := matrixMediaExt("", "", "audio"); got != ".ogg" {
|
||||||
|
t.Fatalf("default audio extension mismatch: got=%q", got)
|
||||||
|
}
|
||||||
|
if got := matrixMediaExt("", "", "video"); got != ".mp4" {
|
||||||
|
t.Fatalf("default video extension mismatch: got=%q", got)
|
||||||
|
}
|
||||||
|
if got := matrixMediaExt("", "", "file"); got != ".bin" {
|
||||||
|
t.Fatalf("default file extension mismatch: got=%q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractInboundContent_ImageNoURLFallback(t *testing.T) {
|
||||||
|
ch := &MatrixChannel{}
|
||||||
|
msg := &event.MessageEventContent{
|
||||||
|
MsgType: event.MsgImage,
|
||||||
|
Body: "test.png",
|
||||||
|
}
|
||||||
|
|
||||||
|
content, mediaRefs, ok := ch.extractInboundContent(context.Background(), msg, "matrix:room:event")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok for image fallback")
|
||||||
|
}
|
||||||
|
if content != "[image: test.png]" {
|
||||||
|
t.Fatalf("unexpected content: %q", content)
|
||||||
|
}
|
||||||
|
if len(mediaRefs) != 0 {
|
||||||
|
t.Fatalf("expected no media refs, got %d", len(mediaRefs))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractInboundContent_AudioNoURLFallback(t *testing.T) {
|
||||||
|
ch := &MatrixChannel{}
|
||||||
|
msg := &event.MessageEventContent{
|
||||||
|
MsgType: event.MsgAudio,
|
||||||
|
FileName: "voice.ogg",
|
||||||
|
Body: "please transcribe",
|
||||||
|
}
|
||||||
|
|
||||||
|
content, mediaRefs, ok := ch.extractInboundContent(context.Background(), msg, "matrix:room:event")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok for audio fallback")
|
||||||
|
}
|
||||||
|
if content != "please transcribe\n[audio: voice.ogg]" {
|
||||||
|
t.Fatalf("unexpected content: %q", content)
|
||||||
|
}
|
||||||
|
if len(mediaRefs) != 0 {
|
||||||
|
t.Fatalf("expected no media refs, got %d", len(mediaRefs))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatrixOutboundMsgType(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
partType string
|
||||||
|
filename string
|
||||||
|
contentType string
|
||||||
|
want event.MessageType
|
||||||
|
}{
|
||||||
|
{name: "explicit image", partType: "image", want: event.MsgImage},
|
||||||
|
{name: "explicit audio", partType: "audio", want: event.MsgAudio},
|
||||||
|
{name: "mime fallback video", contentType: "video/mp4", want: event.MsgVideo},
|
||||||
|
{name: "extension fallback audio", filename: "voice.ogg", want: event.MsgAudio},
|
||||||
|
{name: "unknown defaults file", filename: "report.txt", want: event.MsgFile},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
if got := matrixOutboundMsgType(tc.partType, tc.filename, tc.contentType); got != tc.want {
|
||||||
|
t.Fatalf("%s: got=%q want=%q", tc.name, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatrixOutboundContent(t *testing.T) {
|
||||||
|
content := matrixOutboundContent(
|
||||||
|
"please review",
|
||||||
|
"voice.ogg",
|
||||||
|
event.MsgAudio,
|
||||||
|
"audio/ogg",
|
||||||
|
1234,
|
||||||
|
id.ContentURIString("mxc://matrix.org/abc"),
|
||||||
|
)
|
||||||
|
if content.Body != "please review" {
|
||||||
|
t.Fatalf("unexpected body: %q", content.Body)
|
||||||
|
}
|
||||||
|
if content.FileName != "voice.ogg" {
|
||||||
|
t.Fatalf("unexpected filename: %q", content.FileName)
|
||||||
|
}
|
||||||
|
if content.Info == nil || content.Info.MimeType != "audio/ogg" {
|
||||||
|
t.Fatalf("unexpected content type: %+v", content.Info)
|
||||||
|
}
|
||||||
|
if content.Info == nil || content.Info.Size != 1234 {
|
||||||
|
t.Fatalf("unexpected size: %+v", content.Info)
|
||||||
|
}
|
||||||
|
|
||||||
|
noCaption := matrixOutboundContent(
|
||||||
|
"",
|
||||||
|
"image.png",
|
||||||
|
event.MsgImage,
|
||||||
|
"image/png",
|
||||||
|
0,
|
||||||
|
id.ContentURIString("mxc://matrix.org/def"),
|
||||||
|
)
|
||||||
|
if noCaption.Body != "image.png" {
|
||||||
|
t.Fatalf("unexpected fallback body: %q", noCaption.Body)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -3,7 +3,10 @@ package qq
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/tencent-connect/botgo"
|
"github.com/tencent-connect/botgo"
|
||||||
|
|
@ -20,6 +23,14 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
dedupTTL = 5 * time.Minute
|
||||||
|
dedupInterval = 60 * time.Second
|
||||||
|
dedupMaxSize = 10000 // hard cap on dedup map entries
|
||||||
|
typingResend = 8 * time.Second
|
||||||
|
typingSeconds = 10
|
||||||
|
)
|
||||||
|
|
||||||
type QQChannel struct {
|
type QQChannel struct {
|
||||||
*channels.BaseChannel
|
*channels.BaseChannel
|
||||||
config config.QQConfig
|
config config.QQConfig
|
||||||
|
|
@ -28,20 +39,37 @@ type QQChannel struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
sessionManager botgo.SessionManager
|
sessionManager botgo.SessionManager
|
||||||
processedIDs map[string]bool
|
|
||||||
mu sync.RWMutex
|
// Chat routing: track whether a chatID is group or direct.
|
||||||
|
chatType sync.Map // chatID → "group" | "direct"
|
||||||
|
|
||||||
|
// Passive reply: store last inbound message ID per chat.
|
||||||
|
lastMsgID sync.Map // chatID → string
|
||||||
|
|
||||||
|
// msg_seq: per-chat atomic counter for multi-part replies.
|
||||||
|
msgSeqCounters sync.Map // chatID → *atomic.Uint64
|
||||||
|
|
||||||
|
// Time-based dedup replacing the unbounded map.
|
||||||
|
dedup map[string]time.Time
|
||||||
|
muDedup sync.Mutex
|
||||||
|
|
||||||
|
// done is closed on Stop to shut down the dedup janitor.
|
||||||
|
done chan struct{}
|
||||||
|
stopOnce sync.Once
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewQQChannel(cfg config.QQConfig, messageBus *bus.MessageBus) (*QQChannel, error) {
|
func NewQQChannel(cfg config.QQConfig, messageBus *bus.MessageBus) (*QQChannel, error) {
|
||||||
base := channels.NewBaseChannel("qq", cfg, messageBus, cfg.AllowFrom,
|
base := channels.NewBaseChannel("qq", cfg, messageBus, cfg.AllowFrom,
|
||||||
|
channels.WithMaxMessageLength(cfg.MaxMessageLength),
|
||||||
channels.WithGroupTrigger(cfg.GroupTrigger),
|
channels.WithGroupTrigger(cfg.GroupTrigger),
|
||||||
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
||||||
)
|
)
|
||||||
|
|
||||||
return &QQChannel{
|
return &QQChannel{
|
||||||
BaseChannel: base,
|
BaseChannel: base,
|
||||||
config: cfg,
|
config: cfg,
|
||||||
processedIDs: make(map[string]bool),
|
dedup: make(map[string]time.Time),
|
||||||
|
done: make(chan struct{}),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -52,6 +80,10 @@ func (c *QQChannel) Start(ctx context.Context) error {
|
||||||
|
|
||||||
logger.InfoC("qq", "Starting QQ bot (WebSocket mode)")
|
logger.InfoC("qq", "Starting QQ bot (WebSocket mode)")
|
||||||
|
|
||||||
|
// Reinitialize shutdown signal for clean restart.
|
||||||
|
c.done = make(chan struct{})
|
||||||
|
c.stopOnce = sync.Once{}
|
||||||
|
|
||||||
// create token source
|
// create token source
|
||||||
credentials := &token.QQBotCredentials{
|
credentials := &token.QQBotCredentials{
|
||||||
AppID: c.config.AppID,
|
AppID: c.config.AppID,
|
||||||
|
|
@ -99,6 +131,15 @@ func (c *QQChannel) Start(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
// start dedup janitor goroutine
|
||||||
|
go c.dedupJanitor()
|
||||||
|
|
||||||
|
// Pre-register reasoning_channel_id as group chat if configured,
|
||||||
|
// so outbound-only destinations are routed correctly.
|
||||||
|
if c.config.ReasoningChannelID != "" {
|
||||||
|
c.chatType.Store(c.config.ReasoningChannelID, "group")
|
||||||
|
}
|
||||||
|
|
||||||
c.SetRunning(true)
|
c.SetRunning(true)
|
||||||
logger.InfoC("qq", "QQ bot started successfully")
|
logger.InfoC("qq", "QQ bot started successfully")
|
||||||
|
|
||||||
|
|
@ -109,6 +150,9 @@ func (c *QQChannel) Stop(ctx context.Context) error {
|
||||||
logger.InfoC("qq", "Stopping QQ bot")
|
logger.InfoC("qq", "Stopping QQ bot")
|
||||||
c.SetRunning(false)
|
c.SetRunning(false)
|
||||||
|
|
||||||
|
// Signal the dedup janitor to stop (idempotent).
|
||||||
|
c.stopOnce.Do(func() { close(c.done) })
|
||||||
|
|
||||||
if c.cancel != nil {
|
if c.cancel != nil {
|
||||||
c.cancel()
|
c.cancel()
|
||||||
}
|
}
|
||||||
|
|
@ -116,21 +160,82 @@ func (c *QQChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// getChatKind returns the chat type for a given chatID ("group" or "direct").
|
||||||
|
// Unknown chatIDs default to "group" and log a warning, since QQ group IDs are
|
||||||
|
// more common as outbound-only destinations (e.g. reasoning_channel_id).
|
||||||
|
func (c *QQChannel) getChatKind(chatID string) string {
|
||||||
|
if v, ok := c.chatType.Load(chatID); ok {
|
||||||
|
if k, ok := v.(string); ok {
|
||||||
|
return k
|
||||||
|
}
|
||||||
|
}
|
||||||
|
logger.DebugCF("qq", "Unknown chat type for chatID, defaulting to group", map[string]any{
|
||||||
|
"chat_id": chatID,
|
||||||
|
})
|
||||||
|
return "group"
|
||||||
|
}
|
||||||
|
|
||||||
func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// construct message
|
chatKind := c.getChatKind(msg.ChatID)
|
||||||
|
|
||||||
|
// Build message with content.
|
||||||
msgToCreate := &dto.MessageToCreate{
|
msgToCreate := &dto.MessageToCreate{
|
||||||
Content: msg.Content,
|
Content: msg.Content,
|
||||||
|
MsgType: dto.TextMsg,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use Markdown message type if enabled in config.
|
||||||
|
if c.config.SendMarkdown {
|
||||||
|
msgToCreate.MsgType = dto.MarkdownMsg
|
||||||
|
msgToCreate.Markdown = &dto.Markdown{
|
||||||
|
Content: msg.Content,
|
||||||
|
}
|
||||||
|
// Clear plain content to avoid sending duplicate text.
|
||||||
|
msgToCreate.Content = ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Attach passive reply msg_id and msg_seq if available.
|
||||||
|
if v, ok := c.lastMsgID.Load(msg.ChatID); ok {
|
||||||
|
if msgID, ok := v.(string); ok && msgID != "" {
|
||||||
|
msgToCreate.MsgID = msgID
|
||||||
|
|
||||||
|
// Increment msg_seq atomically for multi-part replies.
|
||||||
|
if counterVal, ok := c.msgSeqCounters.Load(msg.ChatID); ok {
|
||||||
|
if counter, ok := counterVal.(*atomic.Uint64); ok {
|
||||||
|
seq := counter.Add(1)
|
||||||
|
msgToCreate.MsgSeq = uint32(seq)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sanitize URLs in group messages to avoid QQ's URL blacklist rejection.
|
||||||
|
if chatKind == "group" {
|
||||||
|
if msgToCreate.Content != "" {
|
||||||
|
msgToCreate.Content = sanitizeURLs(msgToCreate.Content)
|
||||||
|
}
|
||||||
|
if msgToCreate.Markdown != nil && msgToCreate.Markdown.Content != "" {
|
||||||
|
msgToCreate.Markdown.Content = sanitizeURLs(msgToCreate.Markdown.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Route to group or C2C.
|
||||||
|
var err error
|
||||||
|
if chatKind == "group" {
|
||||||
|
_, err = c.api.PostGroupMessage(ctx, msg.ChatID, msgToCreate)
|
||||||
|
} else {
|
||||||
|
_, err = c.api.PostC2CMessage(ctx, msg.ChatID, msgToCreate)
|
||||||
}
|
}
|
||||||
|
|
||||||
// send C2C message
|
|
||||||
_, err := c.api.PostC2CMessage(ctx, msg.ChatID, msgToCreate)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("qq", "Failed to send C2C message", map[string]any{
|
logger.ErrorCF("qq", "Failed to send message", map[string]any{
|
||||||
"error": err.Error(),
|
"chat_id": msg.ChatID,
|
||||||
|
"chat_kind": chatKind,
|
||||||
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("qq send: %w", channels.ErrTemporary)
|
return fmt.Errorf("qq send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
@ -138,7 +243,150 @@ func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleC2CMessage handles QQ private messages
|
// StartTyping implements channels.TypingCapable.
|
||||||
|
// It sends an InputNotify (msg_type=6) immediately and re-sends every 8 seconds.
|
||||||
|
// The returned stop function is idempotent and cancels the goroutine.
|
||||||
|
func (c *QQChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
||||||
|
// We need a stored msg_id for passive InputNotify; skip if none available.
|
||||||
|
v, ok := c.lastMsgID.Load(chatID)
|
||||||
|
if !ok {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
msgID, ok := v.(string)
|
||||||
|
if !ok || msgID == "" {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
chatKind := c.getChatKind(chatID)
|
||||||
|
|
||||||
|
sendTyping := func(sendCtx context.Context) {
|
||||||
|
typingMsg := &dto.MessageToCreate{
|
||||||
|
MsgType: dto.InputNotifyMsg,
|
||||||
|
MsgID: msgID,
|
||||||
|
InputNotify: &dto.InputNotify{
|
||||||
|
InputType: 1,
|
||||||
|
InputSecond: typingSeconds,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
var err error
|
||||||
|
if chatKind == "group" {
|
||||||
|
_, err = c.api.PostGroupMessage(sendCtx, chatID, typingMsg)
|
||||||
|
} else {
|
||||||
|
_, err = c.api.PostC2CMessage(sendCtx, chatID, typingMsg)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
logger.DebugCF("qq", "Failed to send typing indicator", map[string]any{
|
||||||
|
"chat_id": chatID,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send immediately.
|
||||||
|
sendTyping(c.ctx)
|
||||||
|
|
||||||
|
typingCtx, cancel := context.WithCancel(c.ctx)
|
||||||
|
go func() {
|
||||||
|
ticker := time.NewTicker(typingResend)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-typingCtx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
sendTyping(typingCtx)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return cancel, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
|
// QQ RichMediaMessage requires an HTTP/HTTPS URL — local file paths are not supported.
|
||||||
|
// If part.Ref is already an http(s) URL it is used directly; otherwise we try
|
||||||
|
// the media store, and skip with a warning if the resolved path is not an HTTP URL.
|
||||||
|
func (c *QQChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
|
||||||
|
chatKind := c.getChatKind(msg.ChatID)
|
||||||
|
|
||||||
|
for _, part := range msg.Parts {
|
||||||
|
// If the ref is already an HTTP(S) URL, use it directly.
|
||||||
|
mediaURL := part.Ref
|
||||||
|
if !isHTTPURL(mediaURL) {
|
||||||
|
// Try resolving through media store.
|
||||||
|
store := c.GetMediaStore()
|
||||||
|
if store == nil {
|
||||||
|
logger.WarnCF("qq", "QQ media requires HTTP/HTTPS URL, no media store available", map[string]any{
|
||||||
|
"ref": part.Ref,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
resolved, err := store.Resolve(part.Ref)
|
||||||
|
if err != nil {
|
||||||
|
logger.ErrorCF("qq", "Failed to resolve media ref", map[string]any{
|
||||||
|
"ref": part.Ref,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isHTTPURL(resolved) {
|
||||||
|
logger.WarnCF("qq", "QQ media requires HTTP/HTTPS URL, local files not supported", map[string]any{
|
||||||
|
"ref": part.Ref,
|
||||||
|
"resolved": resolved,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
mediaURL = resolved
|
||||||
|
}
|
||||||
|
|
||||||
|
// Map part type to QQ file type: 1=image, 2=video, 3=audio, 4=file.
|
||||||
|
var fileType uint64
|
||||||
|
switch part.Type {
|
||||||
|
case "image":
|
||||||
|
fileType = 1
|
||||||
|
case "video":
|
||||||
|
fileType = 2
|
||||||
|
case "audio":
|
||||||
|
fileType = 3
|
||||||
|
default:
|
||||||
|
fileType = 4 // file
|
||||||
|
}
|
||||||
|
|
||||||
|
richMedia := &dto.RichMediaMessage{
|
||||||
|
FileType: fileType,
|
||||||
|
URL: mediaURL,
|
||||||
|
SrvSendMsg: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
var sendErr error
|
||||||
|
if chatKind == "group" {
|
||||||
|
_, sendErr = c.api.PostGroupMessage(ctx, msg.ChatID, richMedia)
|
||||||
|
} else {
|
||||||
|
_, sendErr = c.api.PostC2CMessage(ctx, msg.ChatID, richMedia)
|
||||||
|
}
|
||||||
|
|
||||||
|
if sendErr != nil {
|
||||||
|
logger.ErrorCF("qq", "Failed to send media", map[string]any{
|
||||||
|
"type": part.Type,
|
||||||
|
"chat_id": msg.ChatID,
|
||||||
|
"error": sendErr.Error(),
|
||||||
|
})
|
||||||
|
return fmt.Errorf("qq send media: %w", channels.ErrTemporary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleC2CMessage handles QQ private messages.
|
||||||
func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
return func(event *dto.WSPayload, data *dto.WSC2CMessageData) error {
|
return func(event *dto.WSPayload, data *dto.WSC2CMessageData) error {
|
||||||
// deduplication check
|
// deduplication check
|
||||||
|
|
@ -167,7 +415,13 @@ func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
"length": len(content),
|
"length": len(content),
|
||||||
})
|
})
|
||||||
|
|
||||||
// 转发到消息总线
|
// Store chat routing context.
|
||||||
|
c.chatType.Store(senderID, "direct")
|
||||||
|
c.lastMsgID.Store(senderID, data.ID)
|
||||||
|
|
||||||
|
// Reset msg_seq counter for new inbound message.
|
||||||
|
c.msgSeqCounters.Store(senderID, new(atomic.Uint64))
|
||||||
|
|
||||||
metadata := map[string]string{}
|
metadata := map[string]string{}
|
||||||
|
|
||||||
sender := bus.SenderInfo{
|
sender := bus.SenderInfo{
|
||||||
|
|
@ -195,7 +449,7 @@ func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleGroupATMessage handles QQ group @ messages
|
// handleGroupATMessage handles QQ group @ messages.
|
||||||
func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
return func(event *dto.WSPayload, data *dto.WSGroupATMessageData) error {
|
return func(event *dto.WSPayload, data *dto.WSGroupATMessageData) error {
|
||||||
// deduplication check
|
// deduplication check
|
||||||
|
|
@ -232,7 +486,13 @@ func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
"length": len(content),
|
"length": len(content),
|
||||||
})
|
})
|
||||||
|
|
||||||
// 转发到消息总线(使用 GroupID 作为 ChatID)
|
// Store chat routing context using GroupID as chatID.
|
||||||
|
c.chatType.Store(data.GroupID, "group")
|
||||||
|
c.lastMsgID.Store(data.GroupID, data.ID)
|
||||||
|
|
||||||
|
// Reset msg_seq counter for new inbound message.
|
||||||
|
c.msgSeqCounters.Store(data.GroupID, new(atomic.Uint64))
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"group_id": data.GroupID,
|
"group_id": data.GroupID,
|
||||||
}
|
}
|
||||||
|
|
@ -262,29 +522,102 @@ func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// isDuplicate 检查消息是否重复
|
// isDuplicate checks whether a message has been seen within the TTL window.
|
||||||
|
// It also enforces a hard cap on map size by evicting oldest entries.
|
||||||
func (c *QQChannel) isDuplicate(messageID string) bool {
|
func (c *QQChannel) isDuplicate(messageID string) bool {
|
||||||
c.mu.Lock()
|
c.muDedup.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.muDedup.Unlock()
|
||||||
|
|
||||||
if c.processedIDs[messageID] {
|
if ts, exists := c.dedup[messageID]; exists && time.Since(ts) < dedupTTL {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
c.processedIDs[messageID] = true
|
// Enforce hard cap: evict oldest entries when at capacity.
|
||||||
|
if len(c.dedup) >= dedupMaxSize {
|
||||||
// 简单清理:限制 map 大小
|
var oldestID string
|
||||||
if len(c.processedIDs) > 10000 {
|
var oldestTS time.Time
|
||||||
// 清空一半
|
for id, ts := range c.dedup {
|
||||||
count := 0
|
if oldestID == "" || ts.Before(oldestTS) {
|
||||||
for id := range c.processedIDs {
|
oldestID = id
|
||||||
if count >= 5000 {
|
oldestTS = ts
|
||||||
break
|
|
||||||
}
|
}
|
||||||
delete(c.processedIDs, id)
|
}
|
||||||
count++
|
if oldestID != "" {
|
||||||
|
delete(c.dedup, oldestID)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.dedup[messageID] = time.Now()
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// dedupJanitor periodically evicts expired entries from the dedup map.
|
||||||
|
func (c *QQChannel) dedupJanitor() {
|
||||||
|
ticker := time.NewTicker(dedupInterval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-c.done:
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
// Collect expired keys under read-like scan.
|
||||||
|
c.muDedup.Lock()
|
||||||
|
now := time.Now()
|
||||||
|
var expired []string
|
||||||
|
for id, ts := range c.dedup {
|
||||||
|
if now.Sub(ts) >= dedupTTL {
|
||||||
|
expired = append(expired, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, id := range expired {
|
||||||
|
delete(c.dedup, id)
|
||||||
|
}
|
||||||
|
c.muDedup.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isHTTPURL returns true if s starts with http:// or https://.
|
||||||
|
func isHTTPURL(s string) bool {
|
||||||
|
return strings.HasPrefix(s, "http://") || strings.HasPrefix(s, "https://")
|
||||||
|
}
|
||||||
|
|
||||||
|
// urlPattern matches URLs with explicit http(s):// scheme.
|
||||||
|
// Only scheme-prefixed URLs are matched to avoid false positives on bare text
|
||||||
|
// like version numbers (e.g., "1.2.3") or domain-like fragments.
|
||||||
|
var urlPattern = regexp.MustCompile(
|
||||||
|
`(?i)` +
|
||||||
|
`https?://` + // required scheme
|
||||||
|
`(?:[a-zA-Z0-9](?:[a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)+` + // domain parts
|
||||||
|
`[a-zA-Z]{2,}` + // TLD
|
||||||
|
`(?:[/?#]\S*)?`, // optional path/query/fragment
|
||||||
|
)
|
||||||
|
|
||||||
|
// sanitizeURLs replaces dots in URL domains with "。" (fullwidth period)
|
||||||
|
// to prevent QQ's URL blacklist from rejecting the message.
|
||||||
|
func sanitizeURLs(text string) string {
|
||||||
|
return urlPattern.ReplaceAllStringFunc(text, func(match string) string {
|
||||||
|
// Split into scheme + rest (scheme is always present).
|
||||||
|
idx := strings.Index(match, "://")
|
||||||
|
scheme := match[:idx+3]
|
||||||
|
rest := match[idx+3:]
|
||||||
|
|
||||||
|
// Find where the domain ends (first / ? or #).
|
||||||
|
domainEnd := len(rest)
|
||||||
|
for i, ch := range rest {
|
||||||
|
if ch == '/' || ch == '?' || ch == '#' {
|
||||||
|
domainEnd = i
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
domain := rest[:domainEnd]
|
||||||
|
path := rest[domainEnd:]
|
||||||
|
|
||||||
|
// Replace dots in domain only.
|
||||||
|
domain = strings.ReplaceAll(domain, ".", "。")
|
||||||
|
|
||||||
|
return scheme + domain + path
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -170,7 +170,7 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
|
|
||||||
useMarkdownV2 := c.config.Channels.Telegram.UseMarkdownV2
|
useMarkdownV2 := c.config.Channels.Telegram.UseMarkdownV2
|
||||||
|
|
||||||
chatID, err := parseChatID(msg.ChatID)
|
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
@ -202,7 +202,7 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendChunk(ctx, chatID, content, chunk, useMarkdownV2); err != nil {
|
if err := c.sendChunk(ctx, chatID, threadID, content, chunk, useMarkdownV2); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -215,10 +215,12 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
func (c *TelegramChannel) sendChunk(
|
func (c *TelegramChannel) sendChunk(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
chatID int64,
|
chatID int64,
|
||||||
|
threadID int,
|
||||||
content, mdFallback string,
|
content, mdFallback string,
|
||||||
useMarkdownV2 bool,
|
useMarkdownV2 bool,
|
||||||
) error {
|
) error {
|
||||||
tgMsg := tu.Message(tu.ID(chatID), content)
|
tgMsg := tu.Message(tu.ID(chatID), content)
|
||||||
|
tgMsg.MessageThreadID = threadID
|
||||||
if useMarkdownV2 {
|
if useMarkdownV2 {
|
||||||
tgMsg.WithParseMode(telego.ModeMarkdownV2)
|
tgMsg.WithParseMode(telego.ModeMarkdownV2)
|
||||||
} else {
|
} else {
|
||||||
|
|
@ -243,13 +245,16 @@ func (c *TelegramChannel) sendChunk(
|
||||||
// (Telegram's typing indicator expires after ~5s) in a background goroutine.
|
// (Telegram's typing indicator expires after ~5s) in a background goroutine.
|
||||||
// The returned stop function is idempotent and cancels the goroutine.
|
// The returned stop function is idempotent and cancels the goroutine.
|
||||||
func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
||||||
cid, err := parseChatID(chatID)
|
cid, threadID, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return func() {}, err
|
return func() {}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
action := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
||||||
|
action.MessageThreadID = threadID
|
||||||
|
|
||||||
// Send the first typing action immediately
|
// Send the first typing action immediately
|
||||||
_ = c.bot.SendChatAction(ctx, tu.ChatAction(tu.ID(cid), telego.ChatActionTyping))
|
_ = c.bot.SendChatAction(ctx, action)
|
||||||
|
|
||||||
typingCtx, cancel := context.WithCancel(ctx)
|
typingCtx, cancel := context.WithCancel(ctx)
|
||||||
go func() {
|
go func() {
|
||||||
|
|
@ -260,7 +265,9 @@ func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(
|
||||||
case <-typingCtx.Done():
|
case <-typingCtx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
_ = c.bot.SendChatAction(typingCtx, tu.ChatAction(tu.ID(cid), telego.ChatActionTyping))
|
a := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
||||||
|
a.MessageThreadID = threadID
|
||||||
|
_ = c.bot.SendChatAction(typingCtx, a)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
@ -271,7 +278,7 @@ func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(
|
||||||
// EditMessage implements channels.MessageEditor.
|
// EditMessage implements channels.MessageEditor.
|
||||||
func (c *TelegramChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
func (c *TelegramChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
||||||
useMarkdownV2 := c.config.Channels.Telegram.UseMarkdownV2
|
useMarkdownV2 := c.config.Channels.Telegram.UseMarkdownV2
|
||||||
cid, err := parseChatID(chatID)
|
cid, _, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
@ -305,12 +312,14 @@ func (c *TelegramChannel) SendPlaceholder(ctx context.Context, chatID string) (s
|
||||||
text = "Thinking... 💭"
|
text = "Thinking... 💭"
|
||||||
}
|
}
|
||||||
|
|
||||||
cid, err := parseChatID(chatID)
|
cid, threadID, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
pMsg, err := c.bot.SendMessage(ctx, tu.Message(tu.ID(cid), text))
|
phMsg := tu.Message(tu.ID(cid), text)
|
||||||
|
phMsg.MessageThreadID = threadID
|
||||||
|
pMsg, err := c.bot.SendMessage(ctx, phMsg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
@ -324,7 +333,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, err := parseChatID(msg.ChatID)
|
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
@ -356,30 +365,34 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
switch part.Type {
|
switch part.Type {
|
||||||
case "image":
|
case "image":
|
||||||
params := &telego.SendPhotoParams{
|
params := &telego.SendPhotoParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
Photo: telego.InputFile{File: file},
|
MessageThreadID: threadID,
|
||||||
Caption: part.Caption,
|
Photo: telego.InputFile{File: file},
|
||||||
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendPhoto(ctx, params)
|
_, err = c.bot.SendPhoto(ctx, params)
|
||||||
case "audio":
|
case "audio":
|
||||||
params := &telego.SendAudioParams{
|
params := &telego.SendAudioParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
Audio: telego.InputFile{File: file},
|
MessageThreadID: threadID,
|
||||||
Caption: part.Caption,
|
Audio: telego.InputFile{File: file},
|
||||||
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendAudio(ctx, params)
|
_, err = c.bot.SendAudio(ctx, params)
|
||||||
case "video":
|
case "video":
|
||||||
params := &telego.SendVideoParams{
|
params := &telego.SendVideoParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
Video: telego.InputFile{File: file},
|
MessageThreadID: threadID,
|
||||||
Caption: part.Caption,
|
Video: telego.InputFile{File: file},
|
||||||
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendVideo(ctx, params)
|
_, err = c.bot.SendVideo(ctx, params)
|
||||||
default: // "file" or unknown types
|
default: // "file" or unknown types
|
||||||
params := &telego.SendDocumentParams{
|
params := &telego.SendDocumentParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
Document: telego.InputFile{File: file},
|
MessageThreadID: threadID,
|
||||||
Caption: part.Caption,
|
Document: telego.InputFile{File: file},
|
||||||
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendDocument(ctx, params)
|
_, err = c.bot.SendDocument(ctx, params)
|
||||||
}
|
}
|
||||||
|
|
@ -523,19 +536,28 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
content = cleaned
|
content = cleaned
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// For forum topics, embed the thread ID as "chatID/threadID" so replies
|
||||||
|
// route to the correct topic and each topic gets its own session.
|
||||||
|
// Only forum groups (IsForum) are handled; regular group reply threads
|
||||||
|
// must share one session per group.
|
||||||
|
compositeChatID := fmt.Sprintf("%d", chatID)
|
||||||
|
threadID := message.MessageThreadID
|
||||||
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
compositeChatID = fmt.Sprintf("%d/%d", chatID, threadID)
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("telegram", "Received message", map[string]any{
|
logger.DebugCF("telegram", "Received message", map[string]any{
|
||||||
"sender_id": sender.CanonicalID,
|
"sender_id": sender.CanonicalID,
|
||||||
"chat_id": fmt.Sprintf("%d", chatID),
|
"chat_id": compositeChatID,
|
||||||
|
"thread_id": threadID,
|
||||||
"preview": utils.Truncate(content, 50),
|
"preview": utils.Truncate(content, 50),
|
||||||
})
|
})
|
||||||
|
|
||||||
// Placeholder is now auto-triggered by BaseChannel.HandleMessage via PlaceholderCapable
|
|
||||||
|
|
||||||
peerKind := "direct"
|
peerKind := "direct"
|
||||||
peerID := fmt.Sprintf("%d", user.ID)
|
peerID := fmt.Sprintf("%d", user.ID)
|
||||||
if message.Chat.Type != "private" {
|
if message.Chat.Type != "private" {
|
||||||
peerKind = "group"
|
peerKind = "group"
|
||||||
peerID = fmt.Sprintf("%d", chatID)
|
peerID = compositeChatID
|
||||||
}
|
}
|
||||||
|
|
||||||
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
||||||
|
|
@ -548,11 +570,17 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Set parent_peer metadata for per-topic agent binding.
|
||||||
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
metadata["parent_peer_kind"] = "topic"
|
||||||
|
metadata["parent_peer_id"] = fmt.Sprintf("%d", threadID)
|
||||||
|
}
|
||||||
|
|
||||||
c.HandleMessage(c.ctx,
|
c.HandleMessage(c.ctx,
|
||||||
peer,
|
peer,
|
||||||
messageID,
|
messageID,
|
||||||
platformID,
|
platformID,
|
||||||
fmt.Sprintf("%d", chatID),
|
compositeChatID,
|
||||||
content,
|
content,
|
||||||
mediaPaths,
|
mediaPaths,
|
||||||
metadata,
|
metadata,
|
||||||
|
|
@ -608,6 +636,25 @@ func parseContent(text string, useMarkdownV2 bool) string {
|
||||||
return markdownToTelegramHTML(text)
|
return markdownToTelegramHTML(text)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseTelegramChatID splits "chatID/threadID" into its components.
|
||||||
|
// Returns threadID=0 when no "/" is present (non-forum messages).
|
||||||
|
func parseTelegramChatID(chatID string) (int64, int, error) {
|
||||||
|
idx := strings.Index(chatID, "/")
|
||||||
|
if idx == -1 {
|
||||||
|
cid, err := strconv.ParseInt(chatID, 10, 64)
|
||||||
|
return cid, 0, err
|
||||||
|
}
|
||||||
|
cid, err := strconv.ParseInt(chatID[:idx], 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, err
|
||||||
|
}
|
||||||
|
tid, err := strconv.Atoi(chatID[idx+1:])
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, fmt.Errorf("invalid thread ID in chat ID %q: %w", chatID, err)
|
||||||
|
}
|
||||||
|
return cid, tid, nil
|
||||||
|
}
|
||||||
|
|
||||||
func logParseFailed(err error, useMarkdownV2 bool) {
|
func logParseFailed(err error, useMarkdownV2 bool) {
|
||||||
parsingName := "HTML"
|
parsingName := "HTML"
|
||||||
if useMarkdownV2 {
|
if useMarkdownV2 {
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/mymmrac/telego"
|
"github.com/mymmrac/telego"
|
||||||
ta "github.com/mymmrac/telego/telegoapi"
|
ta "github.com/mymmrac/telego/telegoapi"
|
||||||
|
|
@ -273,3 +274,191 @@ func TestSend_InvalidChatID(t *testing.T) {
|
||||||
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
|
assert.True(t, errors.Is(err, channels.ErrSendFailed), "error should wrap ErrSendFailed")
|
||||||
assert.Empty(t, caller.calls)
|
assert.Empty(t, caller.calls)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_Plain(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("12345")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(12345), cid)
|
||||||
|
assert.Equal(t, 0, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_NegativeGroup(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("-1001234567890")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(-1001234567890), cid)
|
||||||
|
assert.Equal(t, 0, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_WithThreadID(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("-1001234567890/42")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(-1001234567890), cid)
|
||||||
|
assert.Equal(t, 42, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_GeneralTopic(t *testing.T) {
|
||||||
|
cid, tid, err := parseTelegramChatID("-100123/1")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(-100123), cid)
|
||||||
|
assert.Equal(t, 1, tid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_Invalid(t *testing.T) {
|
||||||
|
_, _, err := parseTelegramChatID("not-a-number")
|
||||||
|
assert.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTelegramChatID_InvalidThreadID(t *testing.T) {
|
||||||
|
_, _, err := parseTelegramChatID("-100123/not-a-thread")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "invalid thread ID")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSend_WithForumThreadID(t *testing.T) {
|
||||||
|
caller := &stubCaller{
|
||||||
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
|
return successResponse(t), nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
|
err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
|
ChatID: "-1001234567890/42",
|
||||||
|
Content: "Hello from topic",
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Len(t, caller.calls, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ForumTopic_SetsMetadata(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "hello from topic",
|
||||||
|
MessageID: 10,
|
||||||
|
MessageThreadID: 42,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: -1001234567890,
|
||||||
|
Type: "supergroup",
|
||||||
|
IsForum: true,
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 7,
|
||||||
|
FirstName: "Alice",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
inbound, ok := messageBus.ConsumeInbound(ctx)
|
||||||
|
require.True(t, ok, "expected inbound message")
|
||||||
|
|
||||||
|
// Composite chatID should include thread ID
|
||||||
|
assert.Equal(t, "-1001234567890/42", inbound.ChatID)
|
||||||
|
|
||||||
|
// Peer ID should include thread ID for session key isolation
|
||||||
|
assert.Equal(t, "group", inbound.Peer.Kind)
|
||||||
|
assert.Equal(t, "-1001234567890/42", inbound.Peer.ID)
|
||||||
|
|
||||||
|
// Parent peer metadata should be set for agent binding
|
||||||
|
assert.Equal(t, "topic", inbound.Metadata["parent_peer_kind"])
|
||||||
|
assert.Equal(t, "42", inbound.Metadata["parent_peer_id"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_NoForum_NoThreadMetadata(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "regular group message",
|
||||||
|
MessageID: 11,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: -100999,
|
||||||
|
Type: "group",
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 8,
|
||||||
|
FirstName: "Bob",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
inbound, ok := messageBus.ConsumeInbound(ctx)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
// Plain chatID without thread suffix
|
||||||
|
assert.Equal(t, "-100999", inbound.ChatID)
|
||||||
|
|
||||||
|
// Peer ID should be raw chat ID (no thread suffix)
|
||||||
|
assert.Equal(t, "group", inbound.Peer.Kind)
|
||||||
|
assert.Equal(t, "-100999", inbound.Peer.ID)
|
||||||
|
|
||||||
|
// No parent peer metadata
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_kind"])
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ReplyThread_NonForum_NoIsolation(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// In regular groups, reply threads set MessageThreadID to the original
|
||||||
|
// message ID. This should NOT trigger per-thread session isolation.
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "reply in thread",
|
||||||
|
MessageID: 20,
|
||||||
|
MessageThreadID: 15,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: -100999,
|
||||||
|
Type: "supergroup",
|
||||||
|
IsForum: false,
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 9,
|
||||||
|
FirstName: "Carol",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
inbound, ok := messageBus.ConsumeInbound(ctx)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
// chatID should NOT include thread suffix for non-forum groups
|
||||||
|
assert.Equal(t, "-100999", inbound.ChatID)
|
||||||
|
|
||||||
|
// Peer ID should be raw chat ID (shared session for whole group)
|
||||||
|
assert.Equal(t, "group", inbound.Peer.Kind)
|
||||||
|
assert.Equal(t, "-100999", inbound.Peer.ID)
|
||||||
|
|
||||||
|
// No parent peer metadata
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_kind"])
|
||||||
|
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,5 +12,6 @@ func BuiltinDefinitions() []Definition {
|
||||||
listCommand(),
|
listCommand(),
|
||||||
switchCommand(),
|
switchCommand(),
|
||||||
checkCommand(),
|
checkCommand(),
|
||||||
|
clearCommand(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
20
pkg/commands/cmd_clear.go
Normal file
20
pkg/commands/cmd_clear.go
Normal file
|
|
@ -0,0 +1,20 @@
|
||||||
|
package commands
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
func clearCommand() Definition {
|
||||||
|
return Definition{
|
||||||
|
Name: "clear",
|
||||||
|
Description: "Clear the chat history",
|
||||||
|
Usage: "/clear",
|
||||||
|
Handler: func(_ context.Context, req Request, rt *Runtime) error {
|
||||||
|
if rt == nil || rt.ClearHistory == nil {
|
||||||
|
return req.Reply(unavailableMsg)
|
||||||
|
}
|
||||||
|
if err := rt.ClearHistory(); err != nil {
|
||||||
|
return req.Reply("Failed to clear chat history: " + err.Error())
|
||||||
|
}
|
||||||
|
return req.Reply("Chat history cleared!")
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -13,4 +13,5 @@ type Runtime struct {
|
||||||
GetEnabledChannels func() []string
|
GetEnabledChannels func() []string
|
||||||
SwitchModel func(value string) (oldModel string, err error)
|
SwitchModel func(value string) (oldModel string, err error)
|
||||||
SwitchChannel func(value string) error
|
SwitchChannel func(value string) error
|
||||||
|
ClearHistory func() error
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/caarlos0/env/v11"
|
"github.com/caarlos0/env/v11"
|
||||||
|
|
@ -58,6 +59,16 @@ type Config struct {
|
||||||
Tools ToolsConfig `json:"tools"`
|
Tools ToolsConfig `json:"tools"`
|
||||||
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
||||||
Devices DevicesConfig `json:"devices"`
|
Devices DevicesConfig `json:"devices"`
|
||||||
|
// BuildInfo contains build-time version information
|
||||||
|
BuildInfo BuildInfo `json:"build_info,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildInfo contains build-time version information
|
||||||
|
type BuildInfo struct {
|
||||||
|
Version string `json:"version"`
|
||||||
|
GitCommit string `json:"git_commit"`
|
||||||
|
BuildTime string `json:"build_time"`
|
||||||
|
GoVersion string `json:"go_version"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// MarshalJSON implements custom JSON marshaling for Config
|
// MarshalJSON implements custom JSON marshaling for Config
|
||||||
|
|
@ -225,6 +236,7 @@ type ChannelsConfig struct {
|
||||||
QQ QQConfig `json:"qq"`
|
QQ QQConfig `json:"qq"`
|
||||||
DingTalk DingTalkConfig `json:"dingtalk"`
|
DingTalk DingTalkConfig `json:"dingtalk"`
|
||||||
Slack SlackConfig `json:"slack"`
|
Slack SlackConfig `json:"slack"`
|
||||||
|
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"`
|
||||||
|
|
@ -274,15 +286,16 @@ type TelegramConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type FeishuConfig struct {
|
type FeishuConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_FEISHU_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_FEISHU_ENABLED"`
|
||||||
AppID string `json:"app_id" env:"PICOCLAW_CHANNELS_FEISHU_APP_ID"`
|
AppID string `json:"app_id" env:"PICOCLAW_CHANNELS_FEISHU_APP_ID"`
|
||||||
AppSecret string `json:"app_secret" env:"PICOCLAW_CHANNELS_FEISHU_APP_SECRET"`
|
AppSecret string `json:"app_secret" env:"PICOCLAW_CHANNELS_FEISHU_APP_SECRET"`
|
||||||
EncryptKey string `json:"encrypt_key" env:"PICOCLAW_CHANNELS_FEISHU_ENCRYPT_KEY"`
|
EncryptKey string `json:"encrypt_key" env:"PICOCLAW_CHANNELS_FEISHU_ENCRYPT_KEY"`
|
||||||
VerificationToken string `json:"verification_token" env:"PICOCLAW_CHANNELS_FEISHU_VERIFICATION_TOKEN"`
|
VerificationToken string `json:"verification_token" env:"PICOCLAW_CHANNELS_FEISHU_VERIFICATION_TOKEN"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_FEISHU_ALLOW_FROM"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_FEISHU_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_FEISHU_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_FEISHU_REASONING_CHANNEL_ID"`
|
||||||
|
RandomReactionEmoji FlexibleStringSlice `json:"random_reaction_emoji" env:"PICOCLAW_CHANNELS_FEISHU_RANDOM_REACTION_EMOJI"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type DiscordConfig struct {
|
type DiscordConfig struct {
|
||||||
|
|
@ -311,6 +324,8 @@ type QQConfig struct {
|
||||||
AppSecret string `json:"app_secret" env:"PICOCLAW_CHANNELS_QQ_APP_SECRET"`
|
AppSecret string `json:"app_secret" env:"PICOCLAW_CHANNELS_QQ_APP_SECRET"`
|
||||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_QQ_ALLOW_FROM"`
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_QQ_ALLOW_FROM"`
|
||||||
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
||||||
|
MaxMessageLength int `json:"max_message_length" env:"PICOCLAW_CHANNELS_QQ_MAX_MESSAGE_LENGTH"`
|
||||||
|
SendMarkdown bool `json:"send_markdown" env:"PICOCLAW_CHANNELS_QQ_SEND_MARKDOWN"`
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_QQ_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_QQ_REASONING_CHANNEL_ID"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -334,6 +349,19 @@ type SlackConfig struct {
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_SLACK_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_SLACK_REASONING_CHANNEL_ID"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type MatrixConfig struct {
|
||||||
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_MATRIX_ENABLED"`
|
||||||
|
Homeserver string `json:"homeserver" env:"PICOCLAW_CHANNELS_MATRIX_HOMESERVER"`
|
||||||
|
UserID string `json:"user_id" env:"PICOCLAW_CHANNELS_MATRIX_USER_ID"`
|
||||||
|
AccessToken string `json:"access_token" env:"PICOCLAW_CHANNELS_MATRIX_ACCESS_TOKEN"`
|
||||||
|
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"`
|
||||||
|
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_MATRIX_ALLOW_FROM"`
|
||||||
|
GroupTrigger GroupTriggerConfig `json:"group_trigger,omitempty"`
|
||||||
|
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
|
||||||
|
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_MATRIX_REASONING_CHANNEL_ID"`
|
||||||
|
}
|
||||||
|
|
||||||
type LINEConfig struct {
|
type LINEConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_LINE_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_LINE_ENABLED"`
|
||||||
ChannelSecret string `json:"channel_secret" env:"PICOCLAW_CHANNELS_LINE_CHANNEL_SECRET"`
|
ChannelSecret string `json:"channel_secret" env:"PICOCLAW_CHANNELS_LINE_CHANNEL_SECRET"`
|
||||||
|
|
@ -467,6 +495,7 @@ type ProvidersConfig struct {
|
||||||
Qwen ProviderConfig `json:"qwen"`
|
Qwen ProviderConfig `json:"qwen"`
|
||||||
Mistral ProviderConfig `json:"mistral"`
|
Mistral ProviderConfig `json:"mistral"`
|
||||||
Avian ProviderConfig `json:"avian"`
|
Avian ProviderConfig `json:"avian"`
|
||||||
|
Minimax ProviderConfig `json:"minimax"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsEmpty checks if all provider configs are empty (no API keys or API bases set)
|
// IsEmpty checks if all provider configs are empty (no API keys or API bases set)
|
||||||
|
|
@ -492,7 +521,8 @@ func (p ProvidersConfig) IsEmpty() bool {
|
||||||
p.Antigravity.APIKey == "" && p.Antigravity.APIBase == "" &&
|
p.Antigravity.APIKey == "" && p.Antigravity.APIBase == "" &&
|
||||||
p.Qwen.APIKey == "" && p.Qwen.APIBase == "" &&
|
p.Qwen.APIKey == "" && p.Qwen.APIBase == "" &&
|
||||||
p.Mistral.APIKey == "" && p.Mistral.APIBase == "" &&
|
p.Mistral.APIKey == "" && p.Mistral.APIBase == "" &&
|
||||||
p.Avian.APIKey == "" && p.Avian.APIBase == ""
|
p.Avian.APIKey == "" && p.Avian.APIBase == "" &&
|
||||||
|
p.Minimax.APIKey == "" && p.Minimax.APIBase == ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// MarshalJSON implements custom JSON marshaling for ProvidersConfig
|
// MarshalJSON implements custom JSON marshaling for ProvidersConfig
|
||||||
|
|
@ -562,21 +592,31 @@ type GatewayConfig struct {
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ToolDiscoveryConfig struct {
|
||||||
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_DISCOVERY_ENABLED"`
|
||||||
|
TTL int `json:"ttl" env:"PICOCLAW_TOOLS_DISCOVERY_TTL"`
|
||||||
|
MaxSearchResults int `json:"max_search_results" env:"PICOCLAW_MAX_SEARCH_RESULTS"`
|
||||||
|
UseBM25 bool `json:"use_bm25" env:"PICOCLAW_TOOLS_DISCOVERY_USE_BM25"`
|
||||||
|
UseRegex bool `json:"use_regex" env:"PICOCLAW_TOOLS_DISCOVERY_USE_REGEX"`
|
||||||
|
}
|
||||||
|
|
||||||
type ToolConfig struct {
|
type ToolConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"ENABLED"`
|
Enabled bool `json:"enabled" env:"ENABLED"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type BraveConfig struct {
|
type BraveConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"`
|
APIKeys []string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEYS"`
|
||||||
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type TavilyConfig struct {
|
type TavilyConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_TAVILY_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_TAVILY_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEY"`
|
||||||
BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_TAVILY_BASE_URL"`
|
APIKeys []string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEYS"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"`
|
BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_TAVILY_BASE_URL"`
|
||||||
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type DuckDuckGoConfig struct {
|
type DuckDuckGoConfig struct {
|
||||||
|
|
@ -585,9 +625,10 @@ type DuckDuckGoConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type PerplexityConfig struct {
|
type PerplexityConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEY"`
|
||||||
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
|
APIKeys []string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEYS"`
|
||||||
|
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type SearXNGConfig struct {
|
type SearXNGConfig struct {
|
||||||
|
|
@ -646,6 +687,11 @@ type MediaCleanupConfig struct {
|
||||||
Interval int ` env:"PICOCLAW_MEDIA_CLEANUP_INTERVAL" json:"interval_minutes"`
|
Interval int ` env:"PICOCLAW_MEDIA_CLEANUP_INTERVAL" json:"interval_minutes"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ReadFileToolConfig struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
MaxReadFileSize int `json:"max_read_file_size"`
|
||||||
|
}
|
||||||
|
|
||||||
type ToolsConfig struct {
|
type ToolsConfig struct {
|
||||||
AllowReadPaths []string `json:"allow_read_paths" env:"PICOCLAW_TOOLS_ALLOW_READ_PATHS"`
|
AllowReadPaths []string `json:"allow_read_paths" env:"PICOCLAW_TOOLS_ALLOW_READ_PATHS"`
|
||||||
AllowWritePaths []string `json:"allow_write_paths" env:"PICOCLAW_TOOLS_ALLOW_WRITE_PATHS"`
|
AllowWritePaths []string `json:"allow_write_paths" env:"PICOCLAW_TOOLS_ALLOW_WRITE_PATHS"`
|
||||||
|
|
@ -662,7 +708,7 @@ type ToolsConfig struct {
|
||||||
InstallSkill ToolConfig `json:"install_skill" envPrefix:"PICOCLAW_TOOLS_INSTALL_SKILL_"`
|
InstallSkill ToolConfig `json:"install_skill" envPrefix:"PICOCLAW_TOOLS_INSTALL_SKILL_"`
|
||||||
ListDir ToolConfig `json:"list_dir" envPrefix:"PICOCLAW_TOOLS_LIST_DIR_"`
|
ListDir ToolConfig `json:"list_dir" envPrefix:"PICOCLAW_TOOLS_LIST_DIR_"`
|
||||||
Message ToolConfig `json:"message" envPrefix:"PICOCLAW_TOOLS_MESSAGE_"`
|
Message ToolConfig `json:"message" envPrefix:"PICOCLAW_TOOLS_MESSAGE_"`
|
||||||
ReadFile ToolConfig `json:"read_file" envPrefix:"PICOCLAW_TOOLS_READ_FILE_"`
|
ReadFile ReadFileToolConfig `json:"read_file" envPrefix:"PICOCLAW_TOOLS_READ_FILE_"`
|
||||||
SendFile ToolConfig `json:"send_file" envPrefix:"PICOCLAW_TOOLS_SEND_FILE_"`
|
SendFile ToolConfig `json:"send_file" envPrefix:"PICOCLAW_TOOLS_SEND_FILE_"`
|
||||||
Spawn ToolConfig `json:"spawn" envPrefix:"PICOCLAW_TOOLS_SPAWN_"`
|
Spawn ToolConfig `json:"spawn" envPrefix:"PICOCLAW_TOOLS_SPAWN_"`
|
||||||
SPI ToolConfig `json:"spi" envPrefix:"PICOCLAW_TOOLS_SPI_"`
|
SPI ToolConfig `json:"spi" envPrefix:"PICOCLAW_TOOLS_SPI_"`
|
||||||
|
|
@ -714,7 +760,8 @@ type MCPServerConfig struct {
|
||||||
|
|
||||||
// MCPConfig defines configuration for all MCP servers
|
// MCPConfig defines configuration for all MCP servers
|
||||||
type MCPConfig struct {
|
type MCPConfig struct {
|
||||||
ToolConfig `envPrefix:"PICOCLAW_TOOLS_MCP_"`
|
ToolConfig ` envPrefix:"PICOCLAW_TOOLS_MCP_"`
|
||||||
|
Discovery ToolDiscoveryConfig ` json:"discovery"`
|
||||||
// Servers is a map of server name to server configuration
|
// Servers is a map of server name to server configuration
|
||||||
Servers map[string]MCPServerConfig `json:"servers,omitempty"`
|
Servers map[string]MCPServerConfig `json:"servers,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
@ -901,6 +948,29 @@ func (c *Config) ValidateModelList() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func MergeAPIKeys(apiKey string, apiKeys []string) []string {
|
||||||
|
seen := make(map[string]struct{})
|
||||||
|
var all []string
|
||||||
|
|
||||||
|
if k := strings.TrimSpace(apiKey); k != "" {
|
||||||
|
if _, exists := seen[k]; !exists {
|
||||||
|
seen[k] = struct{}{}
|
||||||
|
all = append(all, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, k := range apiKeys {
|
||||||
|
if trimmed := strings.TrimSpace(k); trimmed != "" {
|
||||||
|
if _, exists := seen[trimmed]; !exists {
|
||||||
|
seen[trimmed] = struct{}{}
|
||||||
|
all = append(all, trimmed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return all
|
||||||
|
}
|
||||||
|
|
||||||
func (t *ToolsConfig) IsToolEnabled(name string) bool {
|
func (t *ToolsConfig) IsToolEnabled(name string) bool {
|
||||||
switch name {
|
switch name {
|
||||||
case "web":
|
case "web":
|
||||||
|
|
|
||||||
|
|
@ -283,6 +283,9 @@ func TestDefaultConfig_Channels(t *testing.T) {
|
||||||
if cfg.Channels.Slack.Enabled {
|
if cfg.Channels.Slack.Enabled {
|
||||||
t.Error("Slack should be disabled by default")
|
t.Error("Slack should be disabled by default")
|
||||||
}
|
}
|
||||||
|
if cfg.Channels.Matrix.Enabled {
|
||||||
|
t.Error("Matrix should be disabled by default")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestDefaultConfig_WebTools verifies web tools config
|
// TestDefaultConfig_WebTools verifies web tools config
|
||||||
|
|
@ -293,7 +296,7 @@ func TestDefaultConfig_WebTools(t *testing.T) {
|
||||||
if cfg.Tools.Web.Brave.MaxResults != 5 {
|
if cfg.Tools.Web.Brave.MaxResults != 5 {
|
||||||
t.Error("Expected Brave MaxResults 5, got ", cfg.Tools.Web.Brave.MaxResults)
|
t.Error("Expected Brave MaxResults 5, got ", cfg.Tools.Web.Brave.MaxResults)
|
||||||
}
|
}
|
||||||
if cfg.Tools.Web.Brave.APIKey != "" {
|
if len(cfg.Tools.Web.Brave.APIKeys) != 0 {
|
||||||
t.Error("Brave API key should be empty by default")
|
t.Error("Brave API key should be empty by default")
|
||||||
}
|
}
|
||||||
if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 {
|
if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 {
|
||||||
|
|
|
||||||
|
|
@ -81,10 +81,11 @@ func DefaultConfig() *Config {
|
||||||
AllowFrom: FlexibleStringSlice{},
|
AllowFrom: FlexibleStringSlice{},
|
||||||
},
|
},
|
||||||
QQ: QQConfig{
|
QQ: QQConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
AppID: "",
|
AppID: "",
|
||||||
AppSecret: "",
|
AppSecret: "",
|
||||||
AllowFrom: FlexibleStringSlice{},
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
MaxMessageLength: 2000,
|
||||||
},
|
},
|
||||||
DingTalk: DingTalkConfig{
|
DingTalk: DingTalkConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
|
|
@ -98,6 +99,22 @@ func DefaultConfig() *Config {
|
||||||
AppToken: "",
|
AppToken: "",
|
||||||
AllowFrom: FlexibleStringSlice{},
|
AllowFrom: FlexibleStringSlice{},
|
||||||
},
|
},
|
||||||
|
Matrix: MatrixConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Homeserver: "https://matrix.org",
|
||||||
|
UserID: "",
|
||||||
|
AccessToken: "",
|
||||||
|
DeviceID: "",
|
||||||
|
JoinOnInvite: true,
|
||||||
|
AllowFrom: FlexibleStringSlice{},
|
||||||
|
GroupTrigger: GroupTriggerConfig{
|
||||||
|
MentionOnly: true,
|
||||||
|
},
|
||||||
|
Placeholder: PlaceholderConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Text: "Thinking... 💭",
|
||||||
|
},
|
||||||
|
},
|
||||||
LINE: LINEConfig{
|
LINE: LINEConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
ChannelSecret: "",
|
ChannelSecret: "",
|
||||||
|
|
@ -331,6 +348,14 @@ func DefaultConfig() *Config {
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
},
|
},
|
||||||
|
|
||||||
|
// Minimax - https://api.minimaxi.com/
|
||||||
|
{
|
||||||
|
ModelName: "MiniMax-M2.5",
|
||||||
|
Model: "minimax/MiniMax-M2.5",
|
||||||
|
APIBase: "https://api.minimaxi.com/v1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
|
||||||
// VLLM (local) - http://localhost:8000
|
// VLLM (local) - http://localhost:8000
|
||||||
{
|
{
|
||||||
ModelName: "local-model",
|
ModelName: "local-model",
|
||||||
|
|
@ -360,6 +385,13 @@ func DefaultConfig() *Config {
|
||||||
Brave: BraveConfig{
|
Brave: BraveConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
|
APIKeys: nil,
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
Tavily: TavilyConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
APIKeys: nil,
|
||||||
MaxResults: 5,
|
MaxResults: 5,
|
||||||
},
|
},
|
||||||
DuckDuckGo: DuckDuckGoConfig{
|
DuckDuckGo: DuckDuckGoConfig{
|
||||||
|
|
@ -369,6 +401,7 @@ func DefaultConfig() *Config {
|
||||||
Perplexity: PerplexityConfig{
|
Perplexity: PerplexityConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
|
APIKeys: nil,
|
||||||
MaxResults: 5,
|
MaxResults: 5,
|
||||||
},
|
},
|
||||||
SearXNG: SearXNGConfig{
|
SearXNG: SearXNGConfig{
|
||||||
|
|
@ -420,6 +453,13 @@ func DefaultConfig() *Config {
|
||||||
ToolConfig: ToolConfig{
|
ToolConfig: ToolConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
},
|
},
|
||||||
|
Discovery: ToolDiscoveryConfig{
|
||||||
|
Enabled: false,
|
||||||
|
TTL: 5,
|
||||||
|
MaxSearchResults: 5,
|
||||||
|
UseBM25: true,
|
||||||
|
UseRegex: false,
|
||||||
|
},
|
||||||
Servers: map[string]MCPServerConfig{},
|
Servers: map[string]MCPServerConfig{},
|
||||||
},
|
},
|
||||||
AppendFile: ToolConfig{
|
AppendFile: ToolConfig{
|
||||||
|
|
@ -443,8 +483,9 @@ func DefaultConfig() *Config {
|
||||||
Message: ToolConfig{
|
Message: ToolConfig{
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
},
|
},
|
||||||
ReadFile: ToolConfig{
|
ReadFile: ReadFileToolConfig{
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
|
MaxReadFileSize: 64 * 1024, // 64KB
|
||||||
},
|
},
|
||||||
Spawn: ToolConfig{
|
Spawn: ToolConfig{
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
|
|
@ -470,5 +511,11 @@ func DefaultConfig() *Config {
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
MonitorUSB: true,
|
MonitorUSB: true,
|
||||||
},
|
},
|
||||||
|
BuildInfo: BuildInfo{
|
||||||
|
Version: Version,
|
||||||
|
GitCommit: GitCommit,
|
||||||
|
BuildTime: BuildTime,
|
||||||
|
GoVersion: GoVersion,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
44
pkg/config/version.go
Normal file
44
pkg/config/version.go
Normal file
|
|
@ -0,0 +1,44 @@
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"runtime"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Build-time variables injected via ldflags during build process.
|
||||||
|
// These are set by the Makefile or .goreleaser.yaml using the -X flag:
|
||||||
|
//
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.Version=<version>
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.GitCommit=<commit>
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.BuildTime=<timestamp>
|
||||||
|
// -X github.com/sipeed/picoclaw/pkg/config.GoVersion=<go-version>
|
||||||
|
var (
|
||||||
|
Version = "dev" // Default value when not built with ldflags
|
||||||
|
GitCommit string // Git commit SHA (short)
|
||||||
|
BuildTime string // Build timestamp in RFC3339 format
|
||||||
|
GoVersion string // Go version used for building
|
||||||
|
)
|
||||||
|
|
||||||
|
// FormatVersion returns the version string with optional git commit
|
||||||
|
func FormatVersion() string {
|
||||||
|
v := Version
|
||||||
|
if GitCommit != "" {
|
||||||
|
v += fmt.Sprintf(" (git: %s)", GitCommit)
|
||||||
|
}
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatBuildInfo returns build time and go version info
|
||||||
|
func FormatBuildInfo() (string, string) {
|
||||||
|
build := BuildTime
|
||||||
|
goVer := GoVersion
|
||||||
|
if goVer == "" {
|
||||||
|
goVer = runtime.Version()
|
||||||
|
}
|
||||||
|
return build, goVer
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetVersion returns the version string
|
||||||
|
func GetVersion() string {
|
||||||
|
return Version
|
||||||
|
}
|
||||||
92
pkg/config/version_test.go
Normal file
92
pkg/config/version_test.go
Normal file
|
|
@ -0,0 +1,92 @@
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestFormatVersion_NoGitCommit(t *testing.T) {
|
||||||
|
oldVersion, oldGit := Version, GitCommit
|
||||||
|
t.Cleanup(func() { Version, GitCommit = oldVersion, oldGit })
|
||||||
|
|
||||||
|
Version = "1.2.3"
|
||||||
|
GitCommit = ""
|
||||||
|
|
||||||
|
assert.Equal(t, "1.2.3", FormatVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatVersion_WithGitCommit(t *testing.T) {
|
||||||
|
oldVersion, oldGit := Version, GitCommit
|
||||||
|
t.Cleanup(func() { Version, GitCommit = oldVersion, oldGit })
|
||||||
|
|
||||||
|
Version = "1.2.3"
|
||||||
|
GitCommit = "abc123"
|
||||||
|
|
||||||
|
assert.Equal(t, "1.2.3 (git: abc123)", FormatVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatBuildInfo_UsesBuildTimeAndGoVersion_WhenSet(t *testing.T) {
|
||||||
|
oldBuildTime, oldGoVersion := BuildTime, GoVersion
|
||||||
|
t.Cleanup(func() { BuildTime, GoVersion = oldBuildTime, oldGoVersion })
|
||||||
|
|
||||||
|
BuildTime = "2026-02-20T00:00:00Z"
|
||||||
|
GoVersion = "go1.23.0"
|
||||||
|
|
||||||
|
build, goVer := FormatBuildInfo()
|
||||||
|
|
||||||
|
assert.Equal(t, BuildTime, build)
|
||||||
|
assert.Equal(t, GoVersion, goVer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatBuildInfo_EmptyBuildTime_ReturnsEmptyBuild(t *testing.T) {
|
||||||
|
oldBuildTime, oldGoVersion := BuildTime, GoVersion
|
||||||
|
t.Cleanup(func() { BuildTime, GoVersion = oldBuildTime, oldGoVersion })
|
||||||
|
|
||||||
|
BuildTime = ""
|
||||||
|
GoVersion = "go1.23.0"
|
||||||
|
|
||||||
|
build, goVer := FormatBuildInfo()
|
||||||
|
|
||||||
|
assert.Empty(t, build)
|
||||||
|
assert.Equal(t, GoVersion, goVer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatBuildInfo_EmptyGoVersion_FallsBackToRuntimeVersion(t *testing.T) {
|
||||||
|
oldBuildTime, oldGoVersion := BuildTime, GoVersion
|
||||||
|
t.Cleanup(func() { BuildTime, GoVersion = oldBuildTime, oldGoVersion })
|
||||||
|
|
||||||
|
BuildTime = "x"
|
||||||
|
GoVersion = ""
|
||||||
|
|
||||||
|
build, goVer := FormatBuildInfo()
|
||||||
|
|
||||||
|
assert.Equal(t, "x", build)
|
||||||
|
assert.Equal(t, runtime.Version(), goVer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetVersion(t *testing.T) {
|
||||||
|
oldVersion := Version
|
||||||
|
t.Cleanup(func() { Version = oldVersion })
|
||||||
|
|
||||||
|
Version = "dev"
|
||||||
|
assert.Equal(t, "dev", GetVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetVersion_Custom(t *testing.T) {
|
||||||
|
oldVersion := Version
|
||||||
|
t.Cleanup(func() { Version = oldVersion })
|
||||||
|
|
||||||
|
Version = "v1.0.0"
|
||||||
|
assert.Equal(t, "v1.0.0", GetVersion())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestVersion_DefaultIsDev(t *testing.T) {
|
||||||
|
// Reset to default values
|
||||||
|
oldVersion := Version
|
||||||
|
Version = "dev"
|
||||||
|
t.Cleanup(func() { Version = oldVersion })
|
||||||
|
|
||||||
|
assert.Equal(t, "dev", Version)
|
||||||
|
}
|
||||||
|
|
@ -22,6 +22,7 @@ var supportedChannels = map[string]bool{
|
||||||
"qq": true,
|
"qq": true,
|
||||||
"dingtalk": true,
|
"dingtalk": true,
|
||||||
"slack": true,
|
"slack": true,
|
||||||
|
"matrix": true,
|
||||||
"line": true,
|
"line": true,
|
||||||
"onebot": true,
|
"onebot": true,
|
||||||
"wecom": true,
|
"wecom": true,
|
||||||
|
|
|
||||||
|
|
@ -372,6 +372,8 @@ func (c *OpenClawConfig) IsChannelEnabled(name string) bool {
|
||||||
return c.Channels.Discord == nil || c.Channels.Discord.Enabled == nil || *c.Channels.Discord.Enabled
|
return c.Channels.Discord == nil || c.Channels.Discord.Enabled == nil || *c.Channels.Discord.Enabled
|
||||||
case "slack":
|
case "slack":
|
||||||
return c.Channels.Slack == nil || c.Channels.Slack.Enabled == nil || *c.Channels.Slack.Enabled
|
return c.Channels.Slack == nil || c.Channels.Slack.Enabled == nil || *c.Channels.Slack.Enabled
|
||||||
|
case "matrix":
|
||||||
|
return c.Channels.Matrix == nil || c.Channels.Matrix.Enabled == nil || *c.Channels.Matrix.Enabled
|
||||||
case "whatsapp":
|
case "whatsapp":
|
||||||
return c.Channels.WhatsApp == nil || c.Channels.WhatsApp.Enabled == nil || *c.Channels.WhatsApp.Enabled
|
return c.Channels.WhatsApp == nil || c.Channels.WhatsApp.Enabled == nil || *c.Channels.WhatsApp.Enabled
|
||||||
case "feishu":
|
case "feishu":
|
||||||
|
|
@ -398,6 +400,11 @@ func GetChannelAllowFrom(ch any) []string {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return c.AllowFrom
|
return c.AllowFrom
|
||||||
|
case *OpenClawMatrixConfig:
|
||||||
|
if c == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.AllowFrom
|
||||||
case *OpenClawWhatsAppConfig:
|
case *OpenClawWhatsAppConfig:
|
||||||
if c == nil {
|
if c == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|
@ -628,6 +635,7 @@ type ChannelsConfig struct {
|
||||||
QQ QQConfig `json:"qq"`
|
QQ QQConfig `json:"qq"`
|
||||||
DingTalk DingTalkConfig `json:"dingtalk"`
|
DingTalk DingTalkConfig `json:"dingtalk"`
|
||||||
Slack SlackConfig `json:"slack"`
|
Slack SlackConfig `json:"slack"`
|
||||||
|
Matrix MatrixConfig `json:"matrix"`
|
||||||
LINE LINEConfig `json:"line"`
|
LINE LINEConfig `json:"line"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -689,6 +697,14 @@ type SlackConfig struct {
|
||||||
AllowFrom []string `json:"allow_from"`
|
AllowFrom []string `json:"allow_from"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type MatrixConfig struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
Homeserver string `json:"homeserver"`
|
||||||
|
UserID string `json:"user_id"`
|
||||||
|
AccessToken string `json:"access_token"`
|
||||||
|
AllowFrom []string `json:"allow_from"`
|
||||||
|
}
|
||||||
|
|
||||||
type LINEConfig struct {
|
type LINEConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
ChannelSecret string `json:"channel_secret"`
|
ChannelSecret string `json:"channel_secret"`
|
||||||
|
|
@ -719,16 +735,18 @@ type WebToolsConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type BraveConfig struct {
|
type BraveConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
APIKey string `json:"api_key"`
|
APIKey string `json:"api_key"`
|
||||||
MaxResults int `json:"max_results"`
|
APIKeys []string `json:"api_keys"`
|
||||||
|
MaxResults int `json:"max_results"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type TavilyConfig struct {
|
type TavilyConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
APIKey string `json:"api_key"`
|
APIKey string `json:"api_key"`
|
||||||
BaseURL string `json:"base_url"`
|
APIKeys []string `json:"api_keys"`
|
||||||
MaxResults int `json:"max_results"`
|
BaseURL string `json:"base_url"`
|
||||||
|
MaxResults int `json:"max_results"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type DuckDuckGoConfig struct {
|
type DuckDuckGoConfig struct {
|
||||||
|
|
@ -737,9 +755,10 @@ type DuckDuckGoConfig struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type PerplexityConfig struct {
|
type PerplexityConfig struct {
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
APIKey string `json:"api_key"`
|
APIKey string `json:"api_key"`
|
||||||
MaxResults int `json:"max_results"`
|
APIKeys []string `json:"api_keys"`
|
||||||
|
MaxResults int `json:"max_results"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type CronConfig struct {
|
type CronConfig struct {
|
||||||
|
|
@ -866,12 +885,26 @@ func (c *OpenClawConfig) convertChannels(warnings *[]string) ChannelsConfig {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if c.Channels.Matrix != nil && supportedChannels["matrix"] {
|
||||||
|
enabled := c.Channels.Matrix.Enabled == nil || *c.Channels.Matrix.Enabled
|
||||||
|
channels.Matrix = MatrixConfig{
|
||||||
|
Enabled: enabled,
|
||||||
|
AllowFrom: c.Channels.Matrix.AllowFrom,
|
||||||
|
}
|
||||||
|
if c.Channels.Matrix.Homeserver != nil {
|
||||||
|
channels.Matrix.Homeserver = *c.Channels.Matrix.Homeserver
|
||||||
|
}
|
||||||
|
if c.Channels.Matrix.UserID != nil {
|
||||||
|
channels.Matrix.UserID = *c.Channels.Matrix.UserID
|
||||||
|
}
|
||||||
|
if c.Channels.Matrix.AccessToken != nil {
|
||||||
|
channels.Matrix.AccessToken = *c.Channels.Matrix.AccessToken
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if c.Channels.Signal != nil {
|
if c.Channels.Signal != nil {
|
||||||
*warnings = append(*warnings, "Channel 'signal': No PicoClaw adapter available")
|
*warnings = append(*warnings, "Channel 'signal': No PicoClaw adapter available")
|
||||||
}
|
}
|
||||||
if c.Channels.Matrix != nil {
|
|
||||||
*warnings = append(*warnings, "Channel 'matrix': No PicoClaw adapter available")
|
|
||||||
}
|
|
||||||
if c.Channels.IRC != nil {
|
if c.Channels.IRC != nil {
|
||||||
*warnings = append(*warnings, "Channel 'irc': No PicoClaw adapter available")
|
*warnings = append(*warnings, "Channel 'irc': No PicoClaw adapter available")
|
||||||
}
|
}
|
||||||
|
|
@ -1024,6 +1057,14 @@ func (c ChannelsConfig) ToStandardChannels() config.ChannelsConfig {
|
||||||
BotToken: c.Slack.BotToken,
|
BotToken: c.Slack.BotToken,
|
||||||
AppToken: c.Slack.AppToken,
|
AppToken: c.Slack.AppToken,
|
||||||
},
|
},
|
||||||
|
Matrix: config.MatrixConfig{
|
||||||
|
Enabled: c.Matrix.Enabled,
|
||||||
|
Homeserver: c.Matrix.Homeserver,
|
||||||
|
UserID: c.Matrix.UserID,
|
||||||
|
AccessToken: c.Matrix.AccessToken,
|
||||||
|
AllowFrom: c.Matrix.AllowFrom,
|
||||||
|
JoinOnInvite: true,
|
||||||
|
},
|
||||||
LINE: config.LINEConfig{
|
LINE: config.LINEConfig{
|
||||||
Enabled: c.LINE.Enabled,
|
Enabled: c.LINE.Enabled,
|
||||||
ChannelSecret: c.LINE.ChannelSecret,
|
ChannelSecret: c.LINE.ChannelSecret,
|
||||||
|
|
@ -1048,6 +1089,7 @@ func (c ToolsConfig) ToStandardTools() config.ToolsConfig {
|
||||||
Brave: config.BraveConfig{
|
Brave: config.BraveConfig{
|
||||||
Enabled: c.Web.Brave.Enabled,
|
Enabled: c.Web.Brave.Enabled,
|
||||||
APIKey: c.Web.Brave.APIKey,
|
APIKey: c.Web.Brave.APIKey,
|
||||||
|
APIKeys: c.Web.Brave.APIKeys,
|
||||||
MaxResults: c.Web.Brave.MaxResults,
|
MaxResults: c.Web.Brave.MaxResults,
|
||||||
},
|
},
|
||||||
Tavily: config.TavilyConfig{
|
Tavily: config.TavilyConfig{
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -375,6 +376,96 @@ func TestConvertToPicoClawWithQQAndDingTalk(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestConvertToPicoClawWithMatrix(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
configPath := filepath.Join(tmpDir, "openclaw.json")
|
||||||
|
|
||||||
|
testConfig := `{
|
||||||
|
"channels": {
|
||||||
|
"matrix": {
|
||||||
|
"enabled": true,
|
||||||
|
"homeserver": "https://matrix.example.com",
|
||||||
|
"userId": "@bot:matrix.example.com",
|
||||||
|
"accessToken": "syt_test_token",
|
||||||
|
"allowFrom": ["@alice:matrix.example.com"]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}`
|
||||||
|
|
||||||
|
err := os.WriteFile(configPath, []byte(testConfig), 0o644)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to write test config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := LoadOpenClawConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to load config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
picoCfg, warnings, err := cfg.ConvertToPicoClaw("")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to convert config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !picoCfg.Channels.Matrix.Enabled {
|
||||||
|
t.Error("matrix should be enabled")
|
||||||
|
}
|
||||||
|
if picoCfg.Channels.Matrix.Homeserver != "https://matrix.example.com" {
|
||||||
|
t.Errorf("expected matrix homeserver, got %q", picoCfg.Channels.Matrix.Homeserver)
|
||||||
|
}
|
||||||
|
if picoCfg.Channels.Matrix.UserID != "@bot:matrix.example.com" {
|
||||||
|
t.Errorf("expected matrix user_id, got %q", picoCfg.Channels.Matrix.UserID)
|
||||||
|
}
|
||||||
|
if picoCfg.Channels.Matrix.AccessToken != "syt_test_token" {
|
||||||
|
t.Errorf("expected matrix access_token, got %q", picoCfg.Channels.Matrix.AccessToken)
|
||||||
|
}
|
||||||
|
if len(picoCfg.Channels.Matrix.AllowFrom) != 1 ||
|
||||||
|
picoCfg.Channels.Matrix.AllowFrom[0] != "@alice:matrix.example.com" {
|
||||||
|
t.Errorf("unexpected matrix allow_from: %#v", picoCfg.Channels.Matrix.AllowFrom)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, w := range warnings {
|
||||||
|
if strings.Contains(w, "Channel 'matrix'") {
|
||||||
|
t.Fatalf("matrix should no longer be reported as unsupported, warning=%q", w)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertToPicoClawWithMatrixDisabled(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
configPath := filepath.Join(tmpDir, "openclaw.json")
|
||||||
|
|
||||||
|
testConfig := `{
|
||||||
|
"channels": {
|
||||||
|
"matrix": {
|
||||||
|
"enabled": false,
|
||||||
|
"homeserver": "https://matrix.example.com",
|
||||||
|
"userId": "@bot:matrix.example.com",
|
||||||
|
"accessToken": "syt_test_token"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}`
|
||||||
|
|
||||||
|
err := os.WriteFile(configPath, []byte(testConfig), 0o644)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to write test config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := LoadOpenClawConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to load config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
picoCfg, _, err := cfg.ConvertToPicoClaw("")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to convert config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if picoCfg.Channels.Matrix.Enabled {
|
||||||
|
t.Error("matrix should respect enabled=false from source config")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestOpenClawAgentModel(t *testing.T) {
|
func TestOpenClawAgentModel(t *testing.T) {
|
||||||
model := &OpenClawAgentModel{
|
model := &OpenClawAgentModel{
|
||||||
Primary: strPtr("anthropic/claude-3-opus"),
|
Primary: strPtr("anthropic/claude-3-opus"),
|
||||||
|
|
@ -425,6 +516,9 @@ func TestChannelEnabled(t *testing.T) {
|
||||||
if !cfg.IsChannelEnabled("slack") {
|
if !cfg.IsChannelEnabled("slack") {
|
||||||
t.Error("slack should be enabled (explicitly set)")
|
t.Error("slack should be enabled (explicitly set)")
|
||||||
}
|
}
|
||||||
|
if !cfg.IsChannelEnabled("matrix") {
|
||||||
|
t.Error("matrix should be enabled (nil config defaults to enabled)")
|
||||||
|
}
|
||||||
if cfg.IsChannelEnabled("line") {
|
if cfg.IsChannelEnabled("line") {
|
||||||
t.Error("line should return false (not in switch cases)")
|
t.Error("line should return false (not in switch cases)")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -208,6 +208,15 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
sel.apiBase = "https://api.mistral.ai/v1"
|
sel.apiBase = "https://api.mistral.ai/v1"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
case "minimax":
|
||||||
|
if cfg.Providers.Minimax.APIKey != "" {
|
||||||
|
sel.apiKey = cfg.Providers.Minimax.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Minimax.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Minimax.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.minimaxi.com/v1"
|
||||||
|
}
|
||||||
|
}
|
||||||
case "github_copilot", "copilot":
|
case "github_copilot", "copilot":
|
||||||
sel.providerType = providerTypeGitHubCopilot
|
sel.providerType = providerTypeGitHubCopilot
|
||||||
if cfg.Providers.GitHubCopilot.APIBase != "" {
|
if cfg.Providers.GitHubCopilot.APIBase != "" {
|
||||||
|
|
@ -325,6 +334,13 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
|
||||||
if sel.apiBase == "" {
|
if sel.apiBase == "" {
|
||||||
sel.apiBase = "https://api.mistral.ai/v1"
|
sel.apiBase = "https://api.mistral.ai/v1"
|
||||||
}
|
}
|
||||||
|
case (strings.Contains(lowerModel, "minimax") || strings.HasPrefix(model, "minimax/")) && cfg.Providers.Minimax.APIKey != "":
|
||||||
|
sel.apiKey = cfg.Providers.Minimax.APIKey
|
||||||
|
sel.apiBase = cfg.Providers.Minimax.APIBase
|
||||||
|
sel.proxy = cfg.Providers.Minimax.Proxy
|
||||||
|
if sel.apiBase == "" {
|
||||||
|
sel.apiBase = "https://api.minimaxi.com/v1"
|
||||||
|
}
|
||||||
case strings.HasPrefix(model, "avian/") && cfg.Providers.Avian.APIKey != "":
|
case strings.HasPrefix(model, "avian/") && cfg.Providers.Avian.APIKey != "":
|
||||||
sel.apiKey = cfg.Providers.Avian.APIKey
|
sel.apiKey = cfg.Providers.Avian.APIKey
|
||||||
sel.apiBase = cfg.Providers.Avian.APIBase
|
sel.apiBase = cfg.Providers.Avian.APIBase
|
||||||
|
|
|
||||||
|
|
@ -94,7 +94,8 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
|
|
||||||
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
"vivgrid", "volcengine", "vllm", "qwen", "mistral", "avian":
|
"vivgrid", "volcengine", "vllm", "qwen", "mistral", "avian",
|
||||||
|
"minimax":
|
||||||
// All other OpenAI-compatible HTTP providers
|
// All other OpenAI-compatible HTTP providers
|
||||||
if cfg.APIKey == "" && cfg.APIBase == "" {
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
||||||
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
|
|
@ -212,6 +213,8 @@ func getDefaultAPIBase(protocol string) string {
|
||||||
return "https://api.mistral.ai/v1"
|
return "https://api.mistral.ai/v1"
|
||||||
case "avian":
|
case "avian":
|
||||||
return "https://api.avian.io/v1"
|
return "https://api.avian.io/v1"
|
||||||
|
case "minimax":
|
||||||
|
return "https://api.minimaxi.com/v1"
|
||||||
default:
|
default:
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -440,7 +440,7 @@ func normalizeModel(model, apiBase string) string {
|
||||||
prefix := strings.ToLower(before)
|
prefix := strings.ToLower(before)
|
||||||
switch prefix {
|
switch prefix {
|
||||||
case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google",
|
case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google",
|
||||||
"openrouter", "zhipu", "mistral", "vivgrid":
|
"openrouter", "zhipu", "mistral", "vivgrid", "minimax":
|
||||||
return after
|
return after
|
||||||
default:
|
default:
|
||||||
return model
|
return model
|
||||||
|
|
|
||||||
81
pkg/session/jsonl_backend.go
Normal file
81
pkg/session/jsonl_backend.go
Normal file
|
|
@ -0,0 +1,81 @@
|
||||||
|
package session
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/memory"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
// JSONLBackend adapts a memory.Store into the SessionStore interface.
|
||||||
|
// Write errors are logged rather than returned, matching the fire-and-forget
|
||||||
|
// contract of SessionManager that the agent loop relies on.
|
||||||
|
type JSONLBackend struct {
|
||||||
|
store memory.Store
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewJSONLBackend wraps a memory.Store for use as a SessionStore.
|
||||||
|
func NewJSONLBackend(store memory.Store) *JSONLBackend {
|
||||||
|
return &JSONLBackend{store: store}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) AddMessage(sessionKey, role, content string) {
|
||||||
|
if err := b.store.AddMessage(context.Background(), sessionKey, role, content); err != nil {
|
||||||
|
log.Printf("session: add message: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) AddFullMessage(sessionKey string, msg providers.Message) {
|
||||||
|
if err := b.store.AddFullMessage(context.Background(), sessionKey, msg); err != nil {
|
||||||
|
log.Printf("session: add full message: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) GetHistory(key string) []providers.Message {
|
||||||
|
msgs, err := b.store.GetHistory(context.Background(), key)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("session: get history: %v", err)
|
||||||
|
return []providers.Message{}
|
||||||
|
}
|
||||||
|
return msgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) GetSummary(key string) string {
|
||||||
|
summary, err := b.store.GetSummary(context.Background(), key)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("session: get summary: %v", err)
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return summary
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) SetSummary(key, summary string) {
|
||||||
|
if err := b.store.SetSummary(context.Background(), key, summary); err != nil {
|
||||||
|
log.Printf("session: set summary: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) SetHistory(key string, history []providers.Message) {
|
||||||
|
if err := b.store.SetHistory(context.Background(), key, history); err != nil {
|
||||||
|
log.Printf("session: set history: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *JSONLBackend) TruncateHistory(key string, keepLast int) {
|
||||||
|
if err := b.store.TruncateHistory(context.Background(), key, keepLast); err != nil {
|
||||||
|
log.Printf("session: truncate history: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save persists session state. Since the JSONL store fsyncs every write
|
||||||
|
// immediately, the data is already durable. Save runs compaction to reclaim
|
||||||
|
// space from logically truncated messages (no-op when there are none).
|
||||||
|
func (b *JSONLBackend) Save(key string) error {
|
||||||
|
return b.store.Compact(context.Background(), key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close releases resources held by the underlying store.
|
||||||
|
func (b *JSONLBackend) Close() error {
|
||||||
|
return b.store.Close()
|
||||||
|
}
|
||||||
179
pkg/session/jsonl_backend_test.go
Normal file
179
pkg/session/jsonl_backend_test.go
Normal file
|
|
@ -0,0 +1,179 @@
|
||||||
|
package session_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/memory"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Compile-time interface satisfaction checks.
|
||||||
|
var (
|
||||||
|
_ session.SessionStore = (*session.SessionManager)(nil)
|
||||||
|
_ session.SessionStore = (*session.JSONLBackend)(nil)
|
||||||
|
)
|
||||||
|
|
||||||
|
func newBackend(t *testing.T) *session.JSONLBackend {
|
||||||
|
t.Helper()
|
||||||
|
store, err := memory.NewJSONLStore(t.TempDir())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { store.Close() })
|
||||||
|
return session.NewJSONLBackend(store)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_AddAndGetHistory(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
|
||||||
|
b.AddMessage("s1", "user", "hello")
|
||||||
|
b.AddMessage("s1", "assistant", "hi")
|
||||||
|
|
||||||
|
history := b.GetHistory("s1")
|
||||||
|
if len(history) != 2 {
|
||||||
|
t.Fatalf("got %d messages, want 2", len(history))
|
||||||
|
}
|
||||||
|
if history[0].Role != "user" || history[0].Content != "hello" {
|
||||||
|
t.Errorf("msg[0] = %+v", history[0])
|
||||||
|
}
|
||||||
|
if history[1].Role != "assistant" || history[1].Content != "hi" {
|
||||||
|
t.Errorf("msg[1] = %+v", history[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_AddFullMessage(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
|
||||||
|
msg := providers.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: "done",
|
||||||
|
ToolCalls: []providers.ToolCall{
|
||||||
|
{ID: "tc1", Function: &providers.FunctionCall{Name: "read_file", Arguments: `{"path":"x"}`}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
b.AddFullMessage("s1", msg)
|
||||||
|
|
||||||
|
history := b.GetHistory("s1")
|
||||||
|
if len(history) != 1 {
|
||||||
|
t.Fatalf("got %d, want 1", len(history))
|
||||||
|
}
|
||||||
|
if len(history[0].ToolCalls) != 1 || history[0].ToolCalls[0].ID != "tc1" {
|
||||||
|
t.Errorf("tool calls = %+v", history[0].ToolCalls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_Summary(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
|
||||||
|
if got := b.GetSummary("s1"); got != "" {
|
||||||
|
t.Errorf("got %q, want empty", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
b.SetSummary("s1", "test summary")
|
||||||
|
if got := b.GetSummary("s1"); got != "test summary" {
|
||||||
|
t.Errorf("got %q, want %q", got, "test summary")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_TruncateAndSave(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
b.AddMessage("s1", "user", fmt.Sprintf("msg %d", i))
|
||||||
|
}
|
||||||
|
b.TruncateHistory("s1", 3)
|
||||||
|
|
||||||
|
history := b.GetHistory("s1")
|
||||||
|
if len(history) != 3 {
|
||||||
|
t.Fatalf("got %d, want 3", len(history))
|
||||||
|
}
|
||||||
|
if history[0].Content != "msg 7" {
|
||||||
|
t.Errorf("got %q, want %q", history[0].Content, "msg 7")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save triggers compaction.
|
||||||
|
if err := b.Save("s1"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Messages still accessible after compaction.
|
||||||
|
history = b.GetHistory("s1")
|
||||||
|
if len(history) != 3 {
|
||||||
|
t.Fatalf("after save: got %d, want 3", len(history))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_SetHistory(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
b.AddMessage("s1", "user", "old")
|
||||||
|
|
||||||
|
b.SetHistory("s1", []providers.Message{
|
||||||
|
{Role: "user", Content: "new1"},
|
||||||
|
{Role: "assistant", Content: "new2"},
|
||||||
|
})
|
||||||
|
|
||||||
|
history := b.GetHistory("s1")
|
||||||
|
if len(history) != 2 {
|
||||||
|
t.Fatalf("got %d, want 2", len(history))
|
||||||
|
}
|
||||||
|
if history[0].Content != "new1" {
|
||||||
|
t.Errorf("got %q, want %q", history[0].Content, "new1")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_EmptySession(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
|
||||||
|
history := b.GetHistory("nonexistent")
|
||||||
|
if history == nil {
|
||||||
|
t.Fatal("got nil, want empty slice")
|
||||||
|
}
|
||||||
|
if len(history) != 0 {
|
||||||
|
t.Errorf("got %d, want 0", len(history))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_SessionIsolation(t *testing.T) {
|
||||||
|
b := newBackend(t)
|
||||||
|
b.AddMessage("s1", "user", "session1")
|
||||||
|
b.AddMessage("s2", "user", "session2")
|
||||||
|
|
||||||
|
h1 := b.GetHistory("s1")
|
||||||
|
h2 := b.GetHistory("s2")
|
||||||
|
|
||||||
|
if len(h1) != 1 || h1[0].Content != "session1" {
|
||||||
|
t.Errorf("s1: %+v", h1)
|
||||||
|
}
|
||||||
|
if len(h2) != 1 || h2[0].Content != "session2" {
|
||||||
|
t.Errorf("s2: %+v", h2)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONLBackend_SummarizeFlow(t *testing.T) {
|
||||||
|
// Simulates the real summarization flow in the agent loop:
|
||||||
|
// SetSummary → TruncateHistory → Save
|
||||||
|
b := newBackend(t)
|
||||||
|
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
b.AddMessage("s1", "user", fmt.Sprintf("msg %d", i))
|
||||||
|
}
|
||||||
|
|
||||||
|
b.SetSummary("s1", "conversation about testing")
|
||||||
|
b.TruncateHistory("s1", 4)
|
||||||
|
if err := b.Save("s1"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := b.GetSummary("s1"); got != "conversation about testing" {
|
||||||
|
t.Errorf("summary = %q", got)
|
||||||
|
}
|
||||||
|
history := b.GetHistory("s1")
|
||||||
|
if len(history) != 4 {
|
||||||
|
t.Fatalf("got %d messages, want 4", len(history))
|
||||||
|
}
|
||||||
|
if history[0].Content != "msg 16" {
|
||||||
|
t.Errorf("first message = %q, want %q", history[0].Content, "msg 16")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -265,6 +265,12 @@ func (sm *SessionManager) loadSessions() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Close is a no-op for the in-memory SessionManager; it satisfies the
|
||||||
|
// SessionStore interface so callers can release resources uniformly.
|
||||||
|
func (sm *SessionManager) Close() error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// SetHistory updates the messages of a session.
|
// SetHistory updates the messages of a session.
|
||||||
func (sm *SessionManager) SetHistory(key string, history []providers.Message) {
|
func (sm *SessionManager) SetHistory(key string, history []providers.Message) {
|
||||||
sm.mu.Lock()
|
sm.mu.Lock()
|
||||||
|
|
|
||||||
32
pkg/session/session_store.go
Normal file
32
pkg/session/session_store.go
Normal file
|
|
@ -0,0 +1,32 @@
|
||||||
|
package session
|
||||||
|
|
||||||
|
import "github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
|
||||||
|
// SessionStore defines the persistence operations used by the agent loop.
|
||||||
|
// Both SessionManager (legacy JSON backend) and JSONLBackend satisfy this
|
||||||
|
// interface, allowing the storage layer to be swapped without touching the
|
||||||
|
// agent loop code.
|
||||||
|
//
|
||||||
|
// Write methods (Add*, Set*, Truncate*) are fire-and-forget: they do not
|
||||||
|
// return errors. Implementations should log failures internally. This
|
||||||
|
// matches the original SessionManager contract that the agent loop relies on.
|
||||||
|
type SessionStore interface {
|
||||||
|
// AddMessage appends a simple role/content message to the session.
|
||||||
|
AddMessage(sessionKey, role, content string)
|
||||||
|
// AddFullMessage appends a complete message including tool calls.
|
||||||
|
AddFullMessage(sessionKey string, msg providers.Message)
|
||||||
|
// GetHistory returns the full message history for the session.
|
||||||
|
GetHistory(key string) []providers.Message
|
||||||
|
// GetSummary returns the conversation summary, or "" if none.
|
||||||
|
GetSummary(key string) string
|
||||||
|
// SetSummary replaces the conversation summary.
|
||||||
|
SetSummary(key, summary string)
|
||||||
|
// SetHistory replaces the full message history.
|
||||||
|
SetHistory(key string, history []providers.Message)
|
||||||
|
// TruncateHistory keeps only the last keepLast messages.
|
||||||
|
TruncateHistory(key string, keepLast int)
|
||||||
|
// Save persists any pending state to durable storage.
|
||||||
|
Save(key string) error
|
||||||
|
// Close releases resources held by the store.
|
||||||
|
Close() error
|
||||||
|
}
|
||||||
|
|
@ -2,17 +2,24 @@ package tools
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
|
"math"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/fileutil"
|
"github.com/sipeed/picoclaw/pkg/fileutil"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const MaxReadFileSize = 64 * 1024 // 64KB limit to avoid context overflow
|
||||||
|
|
||||||
// validatePath ensures the given path is within the workspace if restrict is true.
|
// validatePath ensures the given path is within the workspace if restrict is true.
|
||||||
func validatePath(path, workspace string, restrict bool) (string, error) {
|
func validatePath(path, workspace string, restrict bool) (string, error) {
|
||||||
if workspace == "" {
|
if workspace == "" {
|
||||||
|
|
@ -85,15 +92,30 @@ func isWithinWorkspace(candidate, workspace string) bool {
|
||||||
}
|
}
|
||||||
|
|
||||||
type ReadFileTool struct {
|
type ReadFileTool struct {
|
||||||
fs fileSystem
|
fs fileSystem
|
||||||
|
maxSize int64
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewReadFileTool(workspace string, restrict bool, allowPaths ...[]*regexp.Regexp) *ReadFileTool {
|
func NewReadFileTool(
|
||||||
|
workspace string,
|
||||||
|
restrict bool,
|
||||||
|
maxReadFileSize int,
|
||||||
|
allowPaths ...[]*regexp.Regexp,
|
||||||
|
) *ReadFileTool {
|
||||||
var patterns []*regexp.Regexp
|
var patterns []*regexp.Regexp
|
||||||
if len(allowPaths) > 0 {
|
if len(allowPaths) > 0 {
|
||||||
patterns = allowPaths[0]
|
patterns = allowPaths[0]
|
||||||
}
|
}
|
||||||
return &ReadFileTool{fs: buildFs(workspace, restrict, patterns)}
|
|
||||||
|
maxSize := int64(maxReadFileSize)
|
||||||
|
if maxSize <= 0 {
|
||||||
|
maxSize = MaxReadFileSize
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ReadFileTool{
|
||||||
|
fs: buildFs(workspace, restrict, patterns),
|
||||||
|
maxSize: maxSize,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *ReadFileTool) Name() string {
|
func (t *ReadFileTool) Name() string {
|
||||||
|
|
@ -101,7 +123,7 @@ func (t *ReadFileTool) Name() string {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *ReadFileTool) Description() string {
|
func (t *ReadFileTool) Description() string {
|
||||||
return "Read the contents of a file"
|
return "Read the contents of a file. Supports pagination via `offset` and `length`."
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *ReadFileTool) Parameters() map[string]any {
|
func (t *ReadFileTool) Parameters() map[string]any {
|
||||||
|
|
@ -110,7 +132,17 @@ func (t *ReadFileTool) Parameters() map[string]any {
|
||||||
"properties": map[string]any{
|
"properties": map[string]any{
|
||||||
"path": map[string]any{
|
"path": map[string]any{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Path to the file to read",
|
"description": "Path to the file to read.",
|
||||||
|
},
|
||||||
|
"offset": map[string]any{
|
||||||
|
"type": "integer",
|
||||||
|
"description": "Byte offset to start reading from.",
|
||||||
|
"default": 0,
|
||||||
|
},
|
||||||
|
"length": map[string]any{
|
||||||
|
"type": "integer",
|
||||||
|
"description": "Maximum number of bytes to read.",
|
||||||
|
"default": t.maxSize,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
"required": []string{"path"},
|
"required": []string{"path"},
|
||||||
|
|
@ -123,11 +155,171 @@ func (t *ReadFileTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
return ErrorResult("path is required")
|
return ErrorResult("path is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
content, err := t.fs.ReadFile(path)
|
// offset (optional, default 0)
|
||||||
|
offset, err := getInt64Arg(args, "offset", 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(err.Error())
|
return ErrorResult(err.Error())
|
||||||
}
|
}
|
||||||
return NewToolResult(string(content))
|
if offset < 0 {
|
||||||
|
return ErrorResult("offset must be >= 0")
|
||||||
|
}
|
||||||
|
|
||||||
|
// length (optional, capped at MaxReadFileSize)
|
||||||
|
length, err := getInt64Arg(args, "length", t.maxSize)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(err.Error())
|
||||||
|
}
|
||||||
|
if length <= 0 {
|
||||||
|
return ErrorResult("length must be > 0")
|
||||||
|
}
|
||||||
|
if length > t.maxSize {
|
||||||
|
length = t.maxSize
|
||||||
|
}
|
||||||
|
|
||||||
|
file, err := t.fs.Open(path)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(err.Error())
|
||||||
|
}
|
||||||
|
defer file.Close()
|
||||||
|
|
||||||
|
// measure total size
|
||||||
|
totalSize := int64(-1) // -1 means unknown
|
||||||
|
if info, statErr := file.Stat(); statErr == nil {
|
||||||
|
totalSize = info.Size()
|
||||||
|
}
|
||||||
|
|
||||||
|
// sniff the first 512 bytes to detect binary content before loading
|
||||||
|
// it into the LLM context. Seeking back to 0 afterwards restores state.
|
||||||
|
sniff := make([]byte, 512)
|
||||||
|
sniffN, _ := file.Read(sniff)
|
||||||
|
|
||||||
|
// Reset read position to beginning before applying the caller's offset.
|
||||||
|
if seeker, ok := file.(io.Seeker); ok {
|
||||||
|
_, err = seeker.Seek(0, io.SeekStart)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to reset file position after sniff: %v", err))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Non-seekable: we consumed sniffN bytes above; account for them when
|
||||||
|
// discarding to reach the requested offset below.
|
||||||
|
// If offset < sniffN the data we already read covers it, which we
|
||||||
|
// cannot replay on a non-seekable stream — return a clear error.
|
||||||
|
if offset < int64(sniffN) && offset > 0 {
|
||||||
|
return ErrorResult(
|
||||||
|
"non-seekable file: cannot seek to an offset within the first 512 bytes after binary detection",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Seek to the requested offset.
|
||||||
|
if seeker, ok := file.(io.Seeker); ok {
|
||||||
|
_, err = seeker.Seek(offset, io.SeekStart)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to seek to offset %d: %v", offset, err))
|
||||||
|
}
|
||||||
|
} else if offset > 0 {
|
||||||
|
// Fallback for non-seekable streams: discard leading bytes.
|
||||||
|
// sniffN bytes were already consumed above, so subtract them.
|
||||||
|
remaining := offset - int64(sniffN)
|
||||||
|
if remaining > 0 {
|
||||||
|
_, err = io.CopyN(io.Discard, file, remaining)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to advance to offset %d: %v", offset, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// read length+1 bytes to reliably detect whether more content exists
|
||||||
|
// without relying on totalSize (which may be -1 for non-seekable streams).
|
||||||
|
// This avoids the false-positive TRUNCATED message on the last page.
|
||||||
|
probe := make([]byte, length+1)
|
||||||
|
n, err := io.ReadFull(file, probe)
|
||||||
|
// FIX: io.ReadFull returns io.ErrUnexpectedEOF for partial reads (0 < n < len),
|
||||||
|
// and io.EOF only when n == 0. Both are normal terminal conditions — only
|
||||||
|
// other errors are genuine failures.
|
||||||
|
if err != nil && err != io.EOF && !errors.Is(err, io.ErrUnexpectedEOF) {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to read file content: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
// hasMore is true only when we actually got the extra probe byte.
|
||||||
|
hasMore := int64(n) > length
|
||||||
|
data := probe[:min(int64(n), length)]
|
||||||
|
|
||||||
|
if len(data) == 0 {
|
||||||
|
return NewToolResult("[END OF FILE - no content at this offset]")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build metadata header.
|
||||||
|
// use filepath.Base(path) instead of the raw path to avoid leaking
|
||||||
|
// internal filesystem structure into the LLM context.
|
||||||
|
readEnd := offset + int64(len(data))
|
||||||
|
// use ASCII hyphen-minus instead of en-dash (U+2013) to keep the
|
||||||
|
// header parseable by downstream tools and log processors.
|
||||||
|
readRange := fmt.Sprintf("bytes %d-%d", offset, readEnd-1)
|
||||||
|
|
||||||
|
displayPath := filepath.Base(path)
|
||||||
|
var header string
|
||||||
|
if totalSize >= 0 {
|
||||||
|
header = fmt.Sprintf(
|
||||||
|
"[file: %s | total: %d bytes | read: %s]",
|
||||||
|
displayPath, totalSize, readRange,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
header = fmt.Sprintf(
|
||||||
|
"[file: %s | read: %s | total size unknown]",
|
||||||
|
displayPath, readRange,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if hasMore {
|
||||||
|
header += fmt.Sprintf(
|
||||||
|
"\n[TRUNCATED - file has more content. Call read_file again with offset=%d to continue.]",
|
||||||
|
readEnd,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
header += "\n[END OF FILE - no further content.]"
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.DebugCF("tool", "ReadFileTool execution completed successfully",
|
||||||
|
map[string]any{
|
||||||
|
"path": path,
|
||||||
|
"bytes_read": len(data),
|
||||||
|
"has_more": hasMore,
|
||||||
|
})
|
||||||
|
|
||||||
|
return NewToolResult(header + "\n\n" + string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// getInt64Arg extracts an integer argument from the args map, returning the
|
||||||
|
// provided default if the key is absent.
|
||||||
|
func getInt64Arg(args map[string]any, key string, defaultVal int64) (int64, error) {
|
||||||
|
raw, exists := args[key]
|
||||||
|
if !exists {
|
||||||
|
return defaultVal, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v := raw.(type) {
|
||||||
|
case float64:
|
||||||
|
if v != math.Trunc(v) {
|
||||||
|
return 0, fmt.Errorf("%s must be an integer, got float %v", key, v)
|
||||||
|
}
|
||||||
|
if v > math.MaxInt64 || v < math.MinInt64 {
|
||||||
|
return 0, fmt.Errorf("%s value %v overflows int64", key, v)
|
||||||
|
}
|
||||||
|
return int64(v), nil
|
||||||
|
case int:
|
||||||
|
return int64(v), nil
|
||||||
|
case int64:
|
||||||
|
return v, nil
|
||||||
|
case string:
|
||||||
|
parsed, err := strconv.ParseInt(v, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("invalid integer format for %s parameter: %w", key, err)
|
||||||
|
}
|
||||||
|
return parsed, nil
|
||||||
|
default:
|
||||||
|
return 0, fmt.Errorf("unsupported type %T for %s parameter", raw, key)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
type WriteFileTool struct {
|
type WriteFileTool struct {
|
||||||
|
|
@ -249,6 +441,7 @@ type fileSystem interface {
|
||||||
ReadFile(path string) ([]byte, error)
|
ReadFile(path string) ([]byte, error)
|
||||||
WriteFile(path string, data []byte) error
|
WriteFile(path string, data []byte) error
|
||||||
ReadDir(path string) ([]os.DirEntry, error)
|
ReadDir(path string) ([]os.DirEntry, error)
|
||||||
|
Open(path string) (fs.File, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// hostFs is an unrestricted fileReadWriter that operates directly on the host filesystem.
|
// hostFs is an unrestricted fileReadWriter that operates directly on the host filesystem.
|
||||||
|
|
@ -278,6 +471,20 @@ func (h *hostFs) WriteFile(path string, data []byte) error {
|
||||||
return fileutil.WriteFileAtomic(path, data, 0o600)
|
return fileutil.WriteFileAtomic(path, data, 0o600)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (h *hostFs) Open(path string) (fs.File, error) {
|
||||||
|
f, err := os.Open(path)
|
||||||
|
if err != nil {
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return nil, fmt.Errorf("failed to open file: file not found: %w", err)
|
||||||
|
}
|
||||||
|
if os.IsPermission(err) {
|
||||||
|
return nil, fmt.Errorf("failed to open file: access denied: %w", err)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to open file: %w", err)
|
||||||
|
}
|
||||||
|
return f, nil
|
||||||
|
}
|
||||||
|
|
||||||
// sandboxFs is a sandboxed fileSystem that operates within a strictly defined workspace using os.Root.
|
// sandboxFs is a sandboxed fileSystem that operates within a strictly defined workspace using os.Root.
|
||||||
type sandboxFs struct {
|
type sandboxFs struct {
|
||||||
workspace string
|
workspace string
|
||||||
|
|
@ -389,6 +596,26 @@ func (r *sandboxFs) ReadDir(path string) ([]os.DirEntry, error) {
|
||||||
return entries, err
|
return entries, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *sandboxFs) Open(path string) (fs.File, error) {
|
||||||
|
var f fs.File
|
||||||
|
err := r.execute(path, func(root *os.Root, relPath string) error {
|
||||||
|
file, err := root.Open(relPath)
|
||||||
|
if err != nil {
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return fmt.Errorf("failed to open file: file not found: %w", err)
|
||||||
|
}
|
||||||
|
if os.IsPermission(err) || strings.Contains(err.Error(), "escapes from parent") ||
|
||||||
|
strings.Contains(err.Error(), "permission denied") {
|
||||||
|
return fmt.Errorf("failed to open file: access denied: %w", err)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to open file: %w", err)
|
||||||
|
}
|
||||||
|
f = file
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
return f, err
|
||||||
|
}
|
||||||
|
|
||||||
// whitelistFs wraps a sandboxFs and allows access to specific paths outside
|
// whitelistFs wraps a sandboxFs and allows access to specific paths outside
|
||||||
// the workspace when they match any of the provided patterns.
|
// the workspace when they match any of the provided patterns.
|
||||||
type whitelistFs struct {
|
type whitelistFs struct {
|
||||||
|
|
@ -427,6 +654,13 @@ func (w *whitelistFs) ReadDir(path string) ([]os.DirEntry, error) {
|
||||||
return w.sandbox.ReadDir(path)
|
return w.sandbox.ReadDir(path)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (w *whitelistFs) Open(path string) (fs.File, error) {
|
||||||
|
if w.matches(path) {
|
||||||
|
return w.host.Open(path)
|
||||||
|
}
|
||||||
|
return w.sandbox.Open(path)
|
||||||
|
}
|
||||||
|
|
||||||
// buildFs returns the appropriate fileSystem implementation based on restriction
|
// buildFs returns the appropriate fileSystem implementation based on restriction
|
||||||
// settings and optional path whitelist patterns.
|
// settings and optional path whitelist patterns.
|
||||||
func buildFs(workspace string, restrict bool, patterns []*regexp.Regexp) fileSystem {
|
func buildFs(workspace string, restrict bool, patterns []*regexp.Regexp) fileSystem {
|
||||||
|
|
|
||||||
|
|
@ -18,7 +18,7 @@ func TestFilesystemTool_ReadFile_Success(t *testing.T) {
|
||||||
testFile := filepath.Join(tmpDir, "test.txt")
|
testFile := filepath.Join(tmpDir, "test.txt")
|
||||||
os.WriteFile(testFile, []byte("test content"), 0o644)
|
os.WriteFile(testFile, []byte("test content"), 0o644)
|
||||||
|
|
||||||
tool := NewReadFileTool("", false)
|
tool := NewReadFileTool("", false, MaxReadFileSize)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
args := map[string]any{
|
args := map[string]any{
|
||||||
"path": testFile,
|
"path": testFile,
|
||||||
|
|
@ -45,7 +45,7 @@ func TestFilesystemTool_ReadFile_Success(t *testing.T) {
|
||||||
|
|
||||||
// TestFilesystemTool_ReadFile_NotFound verifies error handling for missing file
|
// TestFilesystemTool_ReadFile_NotFound verifies error handling for missing file
|
||||||
func TestFilesystemTool_ReadFile_NotFound(t *testing.T) {
|
func TestFilesystemTool_ReadFile_NotFound(t *testing.T) {
|
||||||
tool := NewReadFileTool("", false)
|
tool := NewReadFileTool("", false, MaxReadFileSize)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
args := map[string]any{
|
args := map[string]any{
|
||||||
"path": "/nonexistent_file_12345.txt",
|
"path": "/nonexistent_file_12345.txt",
|
||||||
|
|
@ -59,7 +59,7 @@ func TestFilesystemTool_ReadFile_NotFound(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Should contain error message
|
// Should contain error message
|
||||||
if !strings.Contains(result.ForLLM, "failed to read") && !strings.Contains(result.ForUser, "failed to read") {
|
if !strings.Contains(result.ForLLM, "failed to open file") && !strings.Contains(result.ForUser, "failed to read") {
|
||||||
t.Errorf("Expected error message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
|
t.Errorf("Expected error message, got ForLLM: %s, ForUser: %s", result.ForLLM, result.ForUser)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -271,7 +271,7 @@ func TestFilesystemTool_ReadFile_RejectsSymlinkEscape(t *testing.T) {
|
||||||
t.Skipf("symlink not supported in this environment: %v", err)
|
t.Skipf("symlink not supported in this environment: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
tool := NewReadFileTool(workspace, true)
|
tool := NewReadFileTool(workspace, true, MaxReadFileSize)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"path": link,
|
"path": link,
|
||||||
})
|
})
|
||||||
|
|
@ -289,7 +289,7 @@ func TestFilesystemTool_ReadFile_RejectsSymlinkEscape(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFilesystemTool_EmptyWorkspace_AccessDenied(t *testing.T) {
|
func TestFilesystemTool_EmptyWorkspace_AccessDenied(t *testing.T) {
|
||||||
tool := NewReadFileTool("", true) // restrict=true but workspace=""
|
tool := NewReadFileTool("", true, MaxReadFileSize) // restrict=true but workspace=""
|
||||||
|
|
||||||
// Try to read a sensitive file (simulated by a temp file outside workspace)
|
// Try to read a sensitive file (simulated by a temp file outside workspace)
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
@ -499,7 +499,7 @@ func TestWhitelistFs_AllowsMatchingPaths(t *testing.T) {
|
||||||
// Pattern allows access to the outsideDir.
|
// Pattern allows access to the outsideDir.
|
||||||
patterns := []*regexp.Regexp{regexp.MustCompile(`^` + regexp.QuoteMeta(outsideDir))}
|
patterns := []*regexp.Regexp{regexp.MustCompile(`^` + regexp.QuoteMeta(outsideDir))}
|
||||||
|
|
||||||
tool := NewReadFileTool(workspace, true, patterns)
|
tool := NewReadFileTool(workspace, true, MaxReadFileSize, patterns)
|
||||||
|
|
||||||
// Read from whitelisted path should succeed.
|
// Read from whitelisted path should succeed.
|
||||||
result := tool.Execute(context.Background(), map[string]any{"path": outsideFile})
|
result := tool.Execute(context.Background(), map[string]any{"path": outsideFile})
|
||||||
|
|
@ -520,3 +520,127 @@ func TestWhitelistFs_AllowsMatchingPaths(t *testing.T) {
|
||||||
t.Errorf("expected non-whitelisted path to be blocked, got: %s", result.ForLLM)
|
t.Errorf("expected non-whitelisted path to be blocked, got: %s", result.ForLLM)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestReadFileTool_ChunkedReading verifies the pagination logic of the tool
|
||||||
|
// by reading a file in multiple chunks using 'offset' and 'length'.
|
||||||
|
func TestReadFileTool_ChunkedReading(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
testFile := filepath.Join(tmpDir, "pagination_test.txt")
|
||||||
|
|
||||||
|
// Create a test file with exactly 26 bytes of content
|
||||||
|
fullContent := "abcdefghijklmnopqrstuvwxyz"
|
||||||
|
err := os.WriteFile(testFile, []byte(fullContent), 0o644)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to write test file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tool := NewReadFileTool(tmpDir, false, MaxReadFileSize)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// --- Step 1: Read the first chunk (10 bytes) ---
|
||||||
|
args1 := map[string]any{
|
||||||
|
"path": testFile,
|
||||||
|
"offset": 0,
|
||||||
|
"length": 10,
|
||||||
|
}
|
||||||
|
result1 := tool.Execute(ctx, args1)
|
||||||
|
|
||||||
|
if result1.IsError {
|
||||||
|
t.Fatalf("Chunk 1 failed: %s", result1.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expect the first 10 characters
|
||||||
|
if !strings.Contains(result1.ForLLM, "abcdefghij") {
|
||||||
|
t.Errorf("Chunk 1 should contain 'abcdefghij', got: %s", result1.ForLLM)
|
||||||
|
}
|
||||||
|
// Expect the header to indicate the file is truncated
|
||||||
|
if !strings.Contains(result1.ForLLM, "[TRUNCATED") {
|
||||||
|
t.Errorf("Chunk 1 header should indicate truncation, got: %s", result1.ForLLM)
|
||||||
|
}
|
||||||
|
// Expect the header to suggest the next offset (10)
|
||||||
|
if !strings.Contains(result1.ForLLM, "offset=10") {
|
||||||
|
t.Errorf("Chunk 1 header should suggest next offset=10, got: %s", result1.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 2: Read the second chunk (10 bytes) ---
|
||||||
|
args2 := map[string]any{
|
||||||
|
"path": testFile,
|
||||||
|
"offset": 10,
|
||||||
|
"length": 10,
|
||||||
|
}
|
||||||
|
result2 := tool.Execute(ctx, args2)
|
||||||
|
|
||||||
|
if result2.IsError {
|
||||||
|
t.Fatalf("Chunk 2 failed: %s", result2.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expect the next 10 characters
|
||||||
|
if !strings.Contains(result2.ForLLM, "klmnopqrst") {
|
||||||
|
t.Errorf("Chunk 2 should contain 'klmnopqrst', got: %s", result2.ForLLM)
|
||||||
|
}
|
||||||
|
// Expect the header to suggest the next offset (20)
|
||||||
|
if !strings.Contains(result2.ForLLM, "offset=20") {
|
||||||
|
t.Errorf("Chunk 2 header should suggest next offset=20, got: %s", result2.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 3: Read the final chunk (remaining 6 bytes) ---
|
||||||
|
// We ask for 10 bytes, but only 6 are left in the file
|
||||||
|
args3 := map[string]any{
|
||||||
|
"path": testFile,
|
||||||
|
"offset": 20,
|
||||||
|
"length": 10,
|
||||||
|
}
|
||||||
|
result3 := tool.Execute(ctx, args3)
|
||||||
|
|
||||||
|
if result3.IsError {
|
||||||
|
t.Fatalf("Chunk 3 failed: %s", result3.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expect the last 6 characters
|
||||||
|
if !strings.Contains(result3.ForLLM, "uvwxyz") {
|
||||||
|
t.Errorf("Chunk 3 should contain 'uvwxyz', got: %s", result3.ForLLM)
|
||||||
|
}
|
||||||
|
// Expect the header to indicate the end of the file
|
||||||
|
if !strings.Contains(result3.ForLLM, "[END OF FILE") {
|
||||||
|
t.Errorf("Chunk 3 header should indicate end of file, got: %s", result3.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure no TRUNCATED message is present in the final chunk
|
||||||
|
if strings.Contains(result3.ForLLM, "[TRUNCATED") {
|
||||||
|
t.Errorf("Chunk 3 header should NOT indicate truncation, got: %s", result3.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestReadFileTool_OffsetBeyondEOF checks the behavior when requesting
|
||||||
|
// An offset that exceeds the total file size.
|
||||||
|
func TestReadFileTool_OffsetBeyondEOF(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
testFile := filepath.Join(tmpDir, "short.txt")
|
||||||
|
|
||||||
|
// create a file of only 5 bytes
|
||||||
|
err := os.WriteFile(testFile, []byte("12345"), 0o644)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to write test file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tool := NewReadFileTool(tmpDir, false, MaxReadFileSize)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
args := map[string]any{
|
||||||
|
"path": testFile,
|
||||||
|
"offset": int64(100), // Offset beyond the end of the file
|
||||||
|
}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
|
||||||
|
// It should not be classified as a tool execution error
|
||||||
|
if result.IsError {
|
||||||
|
t.Errorf("A mistake was not expected, obtained IsError=true: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Must return EXACTLY the string provided in the code
|
||||||
|
expectedMsg := "[END OF FILE - no content at this offset]"
|
||||||
|
if result.ForLLM != expectedMsg {
|
||||||
|
t.Errorf("The message %q was expected, obtained: %q", expectedMsg, result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,20 +5,28 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type ToolEntry struct {
|
||||||
|
Tool Tool
|
||||||
|
IsCore bool
|
||||||
|
TTL int
|
||||||
|
}
|
||||||
|
|
||||||
type ToolRegistry struct {
|
type ToolRegistry struct {
|
||||||
tools map[string]Tool
|
tools map[string]*ToolEntry
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
version atomic.Uint64 // incremented on Register/RegisterHidden for cache invalidation
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewToolRegistry() *ToolRegistry {
|
func NewToolRegistry() *ToolRegistry {
|
||||||
return &ToolRegistry{
|
return &ToolRegistry{
|
||||||
tools: make(map[string]Tool),
|
tools: make(map[string]*ToolEntry),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -30,14 +38,116 @@ func (r *ToolRegistry) Register(tool Tool) {
|
||||||
logger.WarnCF("tools", "Tool registration overwrites existing tool",
|
logger.WarnCF("tools", "Tool registration overwrites existing tool",
|
||||||
map[string]any{"name": name})
|
map[string]any{"name": name})
|
||||||
}
|
}
|
||||||
r.tools[name] = tool
|
r.tools[name] = &ToolEntry{
|
||||||
|
Tool: tool,
|
||||||
|
IsCore: true,
|
||||||
|
TTL: 0, // Core tools do not use TTL
|
||||||
|
}
|
||||||
|
r.version.Add(1)
|
||||||
|
logger.DebugCF("tools", "Registered core tool", map[string]any{"name": name})
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterHidden saves hidden tools (visible only via TTL)
|
||||||
|
func (r *ToolRegistry) RegisterHidden(tool Tool) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
name := tool.Name()
|
||||||
|
if _, exists := r.tools[name]; exists {
|
||||||
|
logger.WarnCF("tools", "Hidden tool registration overwrites existing tool",
|
||||||
|
map[string]any{"name": name})
|
||||||
|
}
|
||||||
|
r.tools[name] = &ToolEntry{
|
||||||
|
Tool: tool,
|
||||||
|
IsCore: false,
|
||||||
|
TTL: 0,
|
||||||
|
}
|
||||||
|
r.version.Add(1)
|
||||||
|
logger.DebugCF("tools", "Registered hidden tool", map[string]any{"name": name})
|
||||||
|
}
|
||||||
|
|
||||||
|
// PromoteTools atomically sets the TTL for multiple non-core tools.
|
||||||
|
// This prevents a concurrent TickTTL from decrementing between promotions.
|
||||||
|
func (r *ToolRegistry) PromoteTools(names []string, ttl int) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
promoted := 0
|
||||||
|
for _, name := range names {
|
||||||
|
if entry, exists := r.tools[name]; exists {
|
||||||
|
if !entry.IsCore {
|
||||||
|
entry.TTL = ttl
|
||||||
|
promoted++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
logger.DebugCF(
|
||||||
|
"tools",
|
||||||
|
"PromoteTools completed",
|
||||||
|
map[string]any{"requested": len(names), "promoted": promoted, "ttl": ttl},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TickTTL decreases TTL only for non-core tools
|
||||||
|
func (r *ToolRegistry) TickTTL() {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
for _, entry := range r.tools {
|
||||||
|
if !entry.IsCore && entry.TTL > 0 {
|
||||||
|
entry.TTL--
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Version returns the current registry version (atomically).
|
||||||
|
func (r *ToolRegistry) Version() uint64 {
|
||||||
|
return r.version.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// HiddenToolSnapshot holds a consistent snapshot of hidden tools and the
|
||||||
|
// registry version at which it was taken. Used by BM25SearchTool cache.
|
||||||
|
type HiddenToolSnapshot struct {
|
||||||
|
Docs []HiddenToolDoc
|
||||||
|
Version uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
// HiddenToolDoc is a lightweight representation of a hidden tool for search indexing.
|
||||||
|
type HiddenToolDoc struct {
|
||||||
|
Name string
|
||||||
|
Description string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SnapshotHiddenTools returns all non-core tools and the current registry
|
||||||
|
// version under a single read-lock, guaranteeing consistency between the
|
||||||
|
// two values.
|
||||||
|
func (r *ToolRegistry) SnapshotHiddenTools() HiddenToolSnapshot {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
docs := make([]HiddenToolDoc, 0, len(r.tools))
|
||||||
|
for name, entry := range r.tools {
|
||||||
|
if !entry.IsCore {
|
||||||
|
docs = append(docs, HiddenToolDoc{
|
||||||
|
Name: name,
|
||||||
|
Description: entry.Tool.Description(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return HiddenToolSnapshot{
|
||||||
|
Docs: docs,
|
||||||
|
Version: r.version.Load(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *ToolRegistry) Get(name string) (Tool, bool) {
|
func (r *ToolRegistry) Get(name string) (Tool, bool) {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
defer r.mu.RUnlock()
|
defer r.mu.RUnlock()
|
||||||
tool, ok := r.tools[name]
|
entry, ok := r.tools[name]
|
||||||
return tool, ok
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
// Hidden tools with expired TTL are not callable.
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return entry.Tool, true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string]any) *ToolResult {
|
func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string]any) *ToolResult {
|
||||||
|
|
@ -135,7 +245,13 @@ func (r *ToolRegistry) GetDefinitions() []map[string]any {
|
||||||
sorted := r.sortedToolNames()
|
sorted := r.sortedToolNames()
|
||||||
definitions := make([]map[string]any, 0, len(sorted))
|
definitions := make([]map[string]any, 0, len(sorted))
|
||||||
for _, name := range sorted {
|
for _, name := range sorted {
|
||||||
definitions = append(definitions, ToolToSchema(r.tools[name]))
|
entry := r.tools[name]
|
||||||
|
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
definitions = append(definitions, ToolToSchema(r.tools[name].Tool))
|
||||||
}
|
}
|
||||||
return definitions
|
return definitions
|
||||||
}
|
}
|
||||||
|
|
@ -149,8 +265,13 @@ func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition {
|
||||||
sorted := r.sortedToolNames()
|
sorted := r.sortedToolNames()
|
||||||
definitions := make([]providers.ToolDefinition, 0, len(sorted))
|
definitions := make([]providers.ToolDefinition, 0, len(sorted))
|
||||||
for _, name := range sorted {
|
for _, name := range sorted {
|
||||||
tool := r.tools[name]
|
entry := r.tools[name]
|
||||||
schema := ToolToSchema(tool)
|
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := ToolToSchema(entry.Tool)
|
||||||
|
|
||||||
// Safely extract nested values with type checks
|
// Safely extract nested values with type checks
|
||||||
fn, ok := schema["function"].(map[string]any)
|
fn, ok := schema["function"].(map[string]any)
|
||||||
|
|
@ -198,8 +319,13 @@ func (r *ToolRegistry) GetSummaries() []string {
|
||||||
sorted := r.sortedToolNames()
|
sorted := r.sortedToolNames()
|
||||||
summaries := make([]string, 0, len(sorted))
|
summaries := make([]string, 0, len(sorted))
|
||||||
for _, name := range sorted {
|
for _, name := range sorted {
|
||||||
tool := r.tools[name]
|
entry := r.tools[name]
|
||||||
summaries = append(summaries, fmt.Sprintf("- `%s` - %s", tool.Name(), tool.Description()))
|
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
summaries = append(summaries, fmt.Sprintf("- `%s` - %s", entry.Tool.Name(), entry.Tool.Description()))
|
||||||
}
|
}
|
||||||
return summaries
|
return summaries
|
||||||
}
|
}
|
||||||
|
|
|
||||||
304
pkg/tools/search_tool.go
Normal file
304
pkg/tools/search_tool.go
Normal file
|
|
@ -0,0 +1,304 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
MaxRegexPatternLength = 200
|
||||||
|
)
|
||||||
|
|
||||||
|
type RegexSearchTool struct {
|
||||||
|
registry *ToolRegistry
|
||||||
|
ttl int
|
||||||
|
maxSearchResults int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewRegexSearchTool(r *ToolRegistry, ttl int, maxSearchResults int) *RegexSearchTool {
|
||||||
|
return &RegexSearchTool{registry: r, ttl: ttl, maxSearchResults: maxSearchResults}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Name() string {
|
||||||
|
return "tool_search_tool_regex"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Description() string {
|
||||||
|
return "Search available hidden tools on-demand using a regex pattern. Returns JSON schemas of discovered tools."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"pattern": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Regex pattern to match tool name or description",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"pattern"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
pattern, ok := args["pattern"].(string)
|
||||||
|
if !ok || strings.TrimSpace(pattern) == "" {
|
||||||
|
// An empty string regex (?i) will match every hidden tool,
|
||||||
|
// dumping massive payloads into the context and burning tokens.
|
||||||
|
return ErrorResult("Missing or invalid 'pattern' argument. Must be a non-empty string.")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(pattern) > MaxRegexPatternLength {
|
||||||
|
logger.WarnCF("discovery", "Regex pattern rejected (too long)", map[string]any{"len": len(pattern)})
|
||||||
|
return ErrorResult(fmt.Sprintf("Pattern too long: max %d characters allowed", MaxRegexPatternLength))
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.DebugCF("discovery", "Regex search", map[string]any{"pattern": pattern})
|
||||||
|
|
||||||
|
res, err := t.registry.SearchRegex(pattern, t.maxSearchResults)
|
||||||
|
if err != nil {
|
||||||
|
logger.WarnCF("discovery", "Invalid regex pattern", map[string]any{"pattern": pattern, "error": err.Error()})
|
||||||
|
return ErrorResult(fmt.Sprintf("Invalid regex pattern syntax: %v. Please fix your regex and try again.", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.InfoCF("discovery", "Regex search completed", map[string]any{"pattern": pattern, "results": len(res)})
|
||||||
|
return formatDiscoveryResponse(t.registry, res, t.ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
type BM25SearchTool struct {
|
||||||
|
registry *ToolRegistry
|
||||||
|
ttl int
|
||||||
|
maxSearchResults int
|
||||||
|
|
||||||
|
// Cache: rebuilt only when the registry version changes.
|
||||||
|
cacheMu sync.Mutex
|
||||||
|
cachedEngine *bm25CachedEngine
|
||||||
|
cacheVersion uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewBM25SearchTool(r *ToolRegistry, ttl int, maxSearchResults int) *BM25SearchTool {
|
||||||
|
return &BM25SearchTool{registry: r, ttl: ttl, maxSearchResults: maxSearchResults}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Name() string {
|
||||||
|
return "tool_search_tool_bm25"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Description() string {
|
||||||
|
return "Search available hidden tools on-demand using natural language query describing the action you need to perform. Returns JSON schemas of discovered tools."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"query": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Search query",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"query"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
query, ok := args["query"].(string)
|
||||||
|
if !ok || strings.TrimSpace(query) == "" {
|
||||||
|
// An empty string query will match every hidden tool,
|
||||||
|
// dumping massive payloads into the context and burning tokens.
|
||||||
|
return ErrorResult("Missing or invalid 'query' argument. Must be a non-empty string.")
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.DebugCF("discovery", "BM25 search", map[string]any{"query": query})
|
||||||
|
|
||||||
|
cached := t.getOrBuildEngine()
|
||||||
|
if cached == nil {
|
||||||
|
logger.DebugCF("discovery", "BM25 search: no hidden tools available", nil)
|
||||||
|
return SilentResult("No tools found matching the query.")
|
||||||
|
}
|
||||||
|
|
||||||
|
ranked := cached.engine.Search(query, t.maxSearchResults)
|
||||||
|
if len(ranked) == 0 {
|
||||||
|
logger.DebugCF("discovery", "BM25 search: no matches", map[string]any{"query": query})
|
||||||
|
return SilentResult("No tools found matching the query.")
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]ToolSearchResult, len(ranked))
|
||||||
|
for i, r := range ranked {
|
||||||
|
results[i] = ToolSearchResult{
|
||||||
|
Name: r.Document.Name,
|
||||||
|
Description: r.Document.Description,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.InfoCF("discovery", "BM25 search completed", map[string]any{"query": query, "results": len(results)})
|
||||||
|
return formatDiscoveryResponse(t.registry, results, t.ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolSearchResult represents the result returned to the LLM.
|
||||||
|
// Parameters are omitted from the JSON response to save context tokens;
|
||||||
|
// the LLM will see full schemas via ToProviderDefs after promotion.
|
||||||
|
type ToolSearchResult struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *ToolRegistry) SearchRegex(pattern string, maxSearchResults int) ([]ToolSearchResult, error) {
|
||||||
|
if maxSearchResults <= 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
regex, err := regexp.Compile("(?i)" + pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to compile regex pattern %q: %w", pattern, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
|
||||||
|
var results []ToolSearchResult
|
||||||
|
|
||||||
|
// Iterate in sorted order for deterministic results across calls.
|
||||||
|
for _, name := range r.sortedToolNames() {
|
||||||
|
entry := r.tools[name]
|
||||||
|
// Search only among the hidden tools (Core tools are already visible)
|
||||||
|
if !entry.IsCore {
|
||||||
|
// Directly call interface methods! No reflection/unmarshalling needed.
|
||||||
|
desc := entry.Tool.Description()
|
||||||
|
|
||||||
|
if regex.MatchString(name) || regex.MatchString(desc) {
|
||||||
|
results = append(results, ToolSearchResult{
|
||||||
|
Name: name,
|
||||||
|
Description: desc,
|
||||||
|
})
|
||||||
|
if len(results) >= maxSearchResults {
|
||||||
|
break // Stop searching once we hit the max! Saves CPU.
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatDiscoveryResponse(registry *ToolRegistry, results []ToolSearchResult, ttl int) *ToolResult {
|
||||||
|
if len(results) == 0 {
|
||||||
|
return SilentResult("No tools found matching the query.")
|
||||||
|
}
|
||||||
|
|
||||||
|
names := make([]string, len(results))
|
||||||
|
for i, r := range results {
|
||||||
|
names[i] = r.Name
|
||||||
|
}
|
||||||
|
registry.PromoteTools(names, ttl)
|
||||||
|
logger.InfoCF("discovery", "Promoted tools", map[string]any{"tools": names, "ttl": ttl})
|
||||||
|
|
||||||
|
b, err := json.Marshal(results)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult("Failed to format search results: " + err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := fmt.Sprintf(
|
||||||
|
"Found %d tools:\n%s\n\nSUCCESS: These tools have been temporarily UNLOCKED as native tools! In your next response, you can call them directly just like any normal tool",
|
||||||
|
len(results),
|
||||||
|
string(b),
|
||||||
|
)
|
||||||
|
|
||||||
|
return SilentResult(msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Lightweight internal type used as corpus document for BM25.
|
||||||
|
type searchDoc struct {
|
||||||
|
Name string
|
||||||
|
Description string
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25CachedEngine wraps a BM25Engine with its corpus snapshot.
|
||||||
|
type bm25CachedEngine struct {
|
||||||
|
engine *utils.BM25Engine[searchDoc]
|
||||||
|
}
|
||||||
|
|
||||||
|
// snapshotToSearchDocs converts a HiddenToolSnapshot to BM25 searchDoc slice.
|
||||||
|
func snapshotToSearchDocs(snap HiddenToolSnapshot) []searchDoc {
|
||||||
|
docs := make([]searchDoc, len(snap.Docs))
|
||||||
|
for i, d := range snap.Docs {
|
||||||
|
docs[i] = searchDoc{Name: d.Name, Description: d.Description}
|
||||||
|
}
|
||||||
|
return docs
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildBM25Engine creates a BM25Engine from a slice of searchDocs.
|
||||||
|
func buildBM25Engine(docs []searchDoc) *utils.BM25Engine[searchDoc] {
|
||||||
|
return utils.NewBM25Engine(
|
||||||
|
docs,
|
||||||
|
func(doc searchDoc) string {
|
||||||
|
return doc.Name + " " + doc.Description
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// getOrBuildEngine returns a cached BM25 engine, rebuilding it only when
|
||||||
|
// the registry version has changed (new tools registered).
|
||||||
|
func (t *BM25SearchTool) getOrBuildEngine() *bm25CachedEngine {
|
||||||
|
// Fast path: optimistic check without locking.
|
||||||
|
if t.cachedEngine != nil && t.cacheVersion == t.registry.Version() {
|
||||||
|
return t.cachedEngine
|
||||||
|
}
|
||||||
|
|
||||||
|
t.cacheMu.Lock()
|
||||||
|
defer t.cacheMu.Unlock()
|
||||||
|
|
||||||
|
// Snapshot + version are read under a single registry RLock,
|
||||||
|
// guaranteeing consistency (no TOCTOU).
|
||||||
|
snap := t.registry.SnapshotHiddenTools()
|
||||||
|
|
||||||
|
// Re-check: another goroutine may have rebuilt while we waited for cacheMu.
|
||||||
|
if t.cachedEngine != nil && t.cacheVersion == snap.Version {
|
||||||
|
return t.cachedEngine
|
||||||
|
}
|
||||||
|
|
||||||
|
docs := snapshotToSearchDocs(snap)
|
||||||
|
if len(docs) == 0 {
|
||||||
|
t.cachedEngine = nil
|
||||||
|
t.cacheVersion = snap.Version
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cached := &bm25CachedEngine{engine: buildBM25Engine(docs)}
|
||||||
|
t.cachedEngine = cached
|
||||||
|
t.cacheVersion = snap.Version
|
||||||
|
logger.DebugCF("discovery", "BM25 engine rebuilt", map[string]any{"docs": len(docs), "version": snap.Version})
|
||||||
|
return cached
|
||||||
|
}
|
||||||
|
|
||||||
|
// SearchBM25 ranks hidden tools against query using BM25 via utils.BM25Engine.
|
||||||
|
// This non-cached variant rebuilds the engine on every call. Used by tests
|
||||||
|
// and any code that doesn't hold a BM25SearchTool instance.
|
||||||
|
func (r *ToolRegistry) SearchBM25(query string, maxSearchResults int) []ToolSearchResult {
|
||||||
|
snap := r.SnapshotHiddenTools()
|
||||||
|
docs := snapshotToSearchDocs(snap)
|
||||||
|
if len(docs) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
ranked := buildBM25Engine(docs).Search(query, maxSearchResults)
|
||||||
|
if len(ranked) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
out := make([]ToolSearchResult, len(ranked))
|
||||||
|
for i, r := range ranked {
|
||||||
|
out[i] = ToolSearchResult{
|
||||||
|
Name: r.Document.Name,
|
||||||
|
Description: r.Document.Description,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
339
pkg/tools/search_tools_test.go
Normal file
339
pkg/tools/search_tools_test.go
Normal file
|
|
@ -0,0 +1,339 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Dummy tool to fill the registry in our tests.
|
||||||
|
type mockSearchableTool struct {
|
||||||
|
name string
|
||||||
|
desc string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSearchableTool) Name() string { return m.name }
|
||||||
|
func (m *mockSearchableTool) Description() string { return m.desc }
|
||||||
|
func (m *mockSearchableTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{"type": "object"}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSearchableTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
return SilentResult("mock executed: " + m.name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper to initialize a populated ToolRegistry
|
||||||
|
func setupPopulatedRegistry() *ToolRegistry {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
|
||||||
|
// A core tool (NOT to be found by searches)
|
||||||
|
reg.Register(&mockSearchableTool{
|
||||||
|
name: "core_search",
|
||||||
|
desc: "I am a visible core tool for searching files",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hidden tools (must be found by searches)
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: "mcp_read_file",
|
||||||
|
desc: "Read the contents of a system file",
|
||||||
|
})
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: "mcp_list_dir",
|
||||||
|
desc: "List directories and files in the system",
|
||||||
|
})
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: "mcp_fetch_net",
|
||||||
|
desc: "Fetch data from a network database",
|
||||||
|
})
|
||||||
|
|
||||||
|
return reg
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexSearchTool_Execute(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
tool := NewRegexSearchTool(reg, 5, 10)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("Empty Pattern Error", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Missing or invalid 'pattern'") {
|
||||||
|
t.Errorf("Expected missing pattern error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Invalid Regex Syntax", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": "[unclosed"})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Invalid regex pattern syntax") {
|
||||||
|
t.Errorf("Expected regex syntax error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("No Match Found", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": "alien"})
|
||||||
|
if res.IsError || !strings.Contains(res.ForLLM, "No tools found matching") {
|
||||||
|
t.Errorf("Expected 'no tools found' message, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Successful Match & Promotion", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": "system"})
|
||||||
|
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatalf("Unexpected error: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "SUCCESS: These tools have been temporarily UNLOCKED") {
|
||||||
|
t.Errorf("Expected success string, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "mcp_read_file") {
|
||||||
|
t.Errorf("Expected 'mcp_read_file' in results")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify that the TTL has been updated for the tools found
|
||||||
|
reg.mu.RLock()
|
||||||
|
defer reg.mu.RUnlock()
|
||||||
|
if reg.tools["mcp_read_file"].TTL != 5 {
|
||||||
|
t.Errorf("Expected TTL of 'mcp_read_file' to be promoted to 5, got %d", reg.tools["mcp_read_file"].TTL)
|
||||||
|
}
|
||||||
|
if reg.tools["mcp_fetch_net"].TTL != 0 {
|
||||||
|
t.Errorf("Expected 'mcp_fetch_net' to NOT be promoted (TTL=0)")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25SearchTool_Execute(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
tool := NewBM25SearchTool(reg, 3, 10)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("Empty Query Error", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": " "})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Missing or invalid 'query'") {
|
||||||
|
t.Errorf("Expected missing query error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("No Match Found", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": "aliens spaceships"})
|
||||||
|
if res.IsError || !strings.Contains(res.ForLLM, "No tools found matching") {
|
||||||
|
t.Errorf("Expected 'no tools found', got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Successful Match & Promotion", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": "read files"})
|
||||||
|
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatalf("Unexpected error: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "mcp_read_file") {
|
||||||
|
t.Errorf("Expected 'mcp_read_file' in BM25 results")
|
||||||
|
}
|
||||||
|
|
||||||
|
reg.mu.RLock()
|
||||||
|
defer reg.mu.RUnlock()
|
||||||
|
if reg.tools["mcp_read_file"].TTL != 3 {
|
||||||
|
t.Errorf("Expected TTL of 'mcp_read_file' to be promoted to 3")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexSearchTool_PatternTooLong(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
tool := NewRegexSearchTool(reg, 5, 10)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
longPattern := strings.Repeat("a", MaxRegexPatternLength+1)
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": longPattern})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Pattern too long") {
|
||||||
|
t.Errorf("Expected pattern too long error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchRegex_ZeroMaxResults(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
|
||||||
|
res, err := reg.SearchRegex("mcp", 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SearchRegex failed: %v", err)
|
||||||
|
}
|
||||||
|
if len(res) != 0 {
|
||||||
|
t.Errorf("Expected 0 results with maxSearchResults=0, got %d", len(res))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchBM25_ZeroMaxResults(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
|
||||||
|
res := reg.SearchBM25("read file", 0)
|
||||||
|
if len(res) != 0 {
|
||||||
|
t.Errorf("Expected 0 results with maxSearchResults=0, got %d", len(res))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearchRegex_DeterministicOrder(t *testing.T) {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: fmt.Sprintf("tool_%02d", i),
|
||||||
|
desc: "searchable tool",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run the same search multiple times and verify order is stable
|
||||||
|
var firstRun []string
|
||||||
|
for attempt := 0; attempt < 10; attempt++ {
|
||||||
|
res, err := reg.SearchRegex("searchable", 20)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SearchRegex failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
names := make([]string, len(res))
|
||||||
|
for i, r := range res {
|
||||||
|
names[i] = r.Name
|
||||||
|
}
|
||||||
|
|
||||||
|
if attempt == 0 {
|
||||||
|
firstRun = names
|
||||||
|
} else {
|
||||||
|
for i, name := range names {
|
||||||
|
if name != firstRun[i] {
|
||||||
|
t.Fatalf("Non-deterministic order at attempt %d, index %d: got %q, want %q",
|
||||||
|
attempt, i, name, firstRun[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_SearchLimitsAndCoreFiltering(t *testing.T) {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
|
||||||
|
// Add 1 Core and 10 Hidden, all containing the word "match"
|
||||||
|
reg.Register(&mockSearchableTool{"core_match", "I am core with match"})
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: fmt.Sprintf("hidden_match_%d", i),
|
||||||
|
desc: "this has a match",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("Regex limits and core filtering", func(t *testing.T) {
|
||||||
|
// Search with Regex and a limit of maxSearchResults = 4
|
||||||
|
res, err := reg.SearchRegex("match", 4)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SearchRegex failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(res) != 4 {
|
||||||
|
t.Errorf("Expected exactly 4 results due to limit, got %d", len(res))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range res {
|
||||||
|
if r.Name == "core_match" {
|
||||||
|
t.Errorf("SearchRegex returned a Core tool, which should be excluded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("BM25 limits and core filtering", func(t *testing.T) {
|
||||||
|
// Search with BM25 and a limit of maxSearchResults = 3
|
||||||
|
res := reg.SearchBM25("match", 3)
|
||||||
|
|
||||||
|
if len(res) != 3 {
|
||||||
|
t.Errorf("Expected exactly 3 results due to limit, got %d", len(res))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range res {
|
||||||
|
if r.Name == "core_match" {
|
||||||
|
t.Errorf("SearchBM25 returned a Core tool, which should be excluded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGet_HiddenToolTTLLifecycle(t *testing.T) {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{name: "hidden_tool", desc: "test"})
|
||||||
|
|
||||||
|
// TTL=0 at registration → not gettable
|
||||||
|
_, ok := reg.Get("hidden_tool")
|
||||||
|
if ok {
|
||||||
|
t.Error("Expected hidden tool with TTL=0 to NOT be gettable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Promote → gettable
|
||||||
|
reg.PromoteTools([]string{"hidden_tool"}, 3)
|
||||||
|
_, ok = reg.Get("hidden_tool")
|
||||||
|
if !ok {
|
||||||
|
t.Error("Expected promoted hidden tool to be gettable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tick down to 0 → not gettable again
|
||||||
|
reg.TickTTL() // 3→2
|
||||||
|
reg.TickTTL() // 2→1
|
||||||
|
reg.TickTTL() // 1→0
|
||||||
|
_, ok = reg.Get("hidden_tool")
|
||||||
|
if ok {
|
||||||
|
t.Error("Expected hidden tool with TTL ticked to 0 to NOT be gettable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Core tools remain always gettable
|
||||||
|
reg.Register(&mockSearchableTool{name: "core_tool", desc: "core"})
|
||||||
|
_, ok = reg.Get("core_tool")
|
||||||
|
if !ok {
|
||||||
|
t.Error("Expected core tool to always be gettable")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25CacheInvalidation(t *testing.T) {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{name: "tool_alpha", desc: "alpha functionality"})
|
||||||
|
|
||||||
|
tool := NewBM25SearchTool(reg, 5, 10)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// First search should find tool_alpha
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": "alpha"})
|
||||||
|
if !strings.Contains(res.ForLLM, "tool_alpha") {
|
||||||
|
t.Fatalf("Expected 'tool_alpha' in first search, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register a new hidden tool
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{name: "tool_beta", desc: "beta functionality"})
|
||||||
|
|
||||||
|
// Cache should be invalidated; new tool should be findable
|
||||||
|
res = tool.Execute(ctx, map[string]any{"query": "beta"})
|
||||||
|
if !strings.Contains(res.ForLLM, "tool_beta") {
|
||||||
|
t.Errorf("Expected 'tool_beta' after cache invalidation, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPromoteTools_ConcurrentWithTickTTL(t *testing.T) {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: fmt.Sprintf("concurrent_tool_%d", i),
|
||||||
|
desc: "concurrent test tool",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
names := make([]string, 20)
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
names[i] = fmt.Sprintf("concurrent_tool_%d", i)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hammer PromoteTools and TickTTL concurrently to detect races
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
for i := 0; i < 1000; i++ {
|
||||||
|
reg.PromoteTools(names, 5)
|
||||||
|
}
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
for i := 0; i < 1000; i++ {
|
||||||
|
reg.TickTTL()
|
||||||
|
}
|
||||||
|
<-done
|
||||||
|
}
|
||||||
483
pkg/tools/web.go
483
pkg/tools/web.go
|
|
@ -11,6 +11,7 @@ import (
|
||||||
"net/url"
|
"net/url"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -76,81 +77,140 @@ func createHTTPClient(proxyURL string, timeout time.Duration) (*http.Client, err
|
||||||
return client, nil
|
return client, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type APIKeyPool struct {
|
||||||
|
keys []string
|
||||||
|
current uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewAPIKeyPool(keys []string) *APIKeyPool {
|
||||||
|
return &APIKeyPool{
|
||||||
|
keys: keys,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type APIKeyIterator struct {
|
||||||
|
pool *APIKeyPool
|
||||||
|
startIdx uint32
|
||||||
|
attempt uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *APIKeyPool) NewIterator() *APIKeyIterator {
|
||||||
|
if len(p.keys) == 0 {
|
||||||
|
return &APIKeyIterator{pool: p}
|
||||||
|
}
|
||||||
|
idx := atomic.AddUint32(&p.current, 1) - 1
|
||||||
|
return &APIKeyIterator{
|
||||||
|
pool: p,
|
||||||
|
startIdx: idx,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (it *APIKeyIterator) Next() (string, bool) {
|
||||||
|
length := uint32(len(it.pool.keys))
|
||||||
|
if length == 0 || it.attempt >= length {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
key := it.pool.keys[(it.startIdx+it.attempt)%length]
|
||||||
|
it.attempt++
|
||||||
|
return key, true
|
||||||
|
}
|
||||||
|
|
||||||
type SearchProvider interface {
|
type SearchProvider interface {
|
||||||
Search(ctx context.Context, query string, count int) (string, error)
|
Search(ctx context.Context, query string, count int) (string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type BraveSearchProvider struct {
|
type BraveSearchProvider struct {
|
||||||
apiKey string
|
keyPool *APIKeyPool
|
||||||
proxy string
|
proxy string
|
||||||
client *http.Client
|
client *http.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
||||||
searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d",
|
searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d",
|
||||||
url.QueryEscape(query), count)
|
url.QueryEscape(query), count)
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil)
|
var lastErr error
|
||||||
if err != nil {
|
iter := p.keyPool.NewIterator()
|
||||||
return "", fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req.Header.Set("Accept", "application/json")
|
for {
|
||||||
req.Header.Set("X-Subscription-Token", p.apiKey)
|
apiKey, ok := iter.Next()
|
||||||
|
if !ok {
|
||||||
resp, err := p.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
return "", fmt.Errorf("brave api error (status %d): %s", resp.StatusCode, string(body))
|
|
||||||
}
|
|
||||||
|
|
||||||
var searchResp struct {
|
|
||||||
Web struct {
|
|
||||||
Results []struct {
|
|
||||||
Title string `json:"title"`
|
|
||||||
URL string `json:"url"`
|
|
||||||
Description string `json:"description"`
|
|
||||||
} `json:"results"`
|
|
||||||
} `json:"web"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := json.Unmarshal(body, &searchResp); err != nil {
|
|
||||||
// Log error body for debugging
|
|
||||||
fmt.Printf("Brave API Error Body: %s\n", string(body))
|
|
||||||
return "", fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
results := searchResp.Web.Results
|
|
||||||
if len(results) == 0 {
|
|
||||||
return fmt.Sprintf("No results for: %s", query), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var lines []string
|
|
||||||
lines = append(lines, fmt.Sprintf("Results for: %s", query))
|
|
||||||
for i, item := range results {
|
|
||||||
if i >= count {
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
lines = append(lines, fmt.Sprintf("%d. %s\n %s", i+1, item.Title, item.URL))
|
|
||||||
if item.Description != "" {
|
req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil)
|
||||||
lines = append(lines, fmt.Sprintf(" %s", item.Description))
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create request: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Accept", "application/json")
|
||||||
|
req.Header.Set("X-Subscription-Token", apiKey)
|
||||||
|
|
||||||
|
resp, err := p.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = fmt.Errorf("request failed: %w", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
lastErr = fmt.Errorf("failed to read response: %w", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
lastErr = fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||||
|
if resp.StatusCode == http.StatusTooManyRequests ||
|
||||||
|
resp.StatusCode == http.StatusUnauthorized ||
|
||||||
|
resp.StatusCode == http.StatusForbidden ||
|
||||||
|
resp.StatusCode >= 500 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return "", lastErr
|
||||||
|
}
|
||||||
|
|
||||||
|
var searchResp struct {
|
||||||
|
Web struct {
|
||||||
|
Results []struct {
|
||||||
|
Title string `json:"title"`
|
||||||
|
URL string `json:"url"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
} `json:"results"`
|
||||||
|
} `json:"web"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &searchResp); err != nil {
|
||||||
|
// Log error body for debugging
|
||||||
|
return "", fmt.Errorf("failed to parse response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
results := searchResp.Web.Results
|
||||||
|
if len(results) == 0 {
|
||||||
|
return fmt.Sprintf("No results for: %s", query), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var lines []string
|
||||||
|
lines = append(lines, fmt.Sprintf("Results for: %s", query))
|
||||||
|
for i, item := range results {
|
||||||
|
if i >= count {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
lines = append(lines, fmt.Sprintf("%d. %s\n %s", i+1, item.Title, item.URL))
|
||||||
|
if item.Description != "" {
|
||||||
|
lines = append(lines, fmt.Sprintf(" %s", item.Description))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(lines, "\n"), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return strings.Join(lines, "\n"), nil
|
return "", fmt.Errorf("all api keys failed, last error: %w", lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
type TavilySearchProvider struct {
|
type TavilySearchProvider struct {
|
||||||
apiKey string
|
keyPool *APIKeyPool
|
||||||
baseURL string
|
baseURL string
|
||||||
proxy string
|
proxy string
|
||||||
client *http.Client
|
client *http.Client
|
||||||
|
|
@ -162,74 +222,96 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
|
||||||
searchURL = "https://api.tavily.com/search"
|
searchURL = "https://api.tavily.com/search"
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := map[string]any{
|
var lastErr error
|
||||||
"api_key": p.apiKey,
|
iter := p.keyPool.NewIterator()
|
||||||
"query": query,
|
|
||||||
"search_depth": "advanced",
|
|
||||||
"include_answer": false,
|
|
||||||
"include_images": false,
|
|
||||||
"include_raw_content": false,
|
|
||||||
"max_results": count,
|
|
||||||
}
|
|
||||||
|
|
||||||
bodyBytes, err := json.Marshal(payload)
|
for {
|
||||||
if err != nil {
|
apiKey, ok := iter.Next()
|
||||||
return "", fmt.Errorf("failed to marshal payload: %w", err)
|
if !ok {
|
||||||
}
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, "POST", searchURL, bytes.NewBuffer(bodyBytes))
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
req.Header.Set("User-Agent", userAgent)
|
|
||||||
|
|
||||||
resp, err := p.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
return "", fmt.Errorf("tavily api error (status %d): %s", resp.StatusCode, string(body))
|
|
||||||
}
|
|
||||||
|
|
||||||
var searchResp struct {
|
|
||||||
Results []struct {
|
|
||||||
Title string `json:"title"`
|
|
||||||
URL string `json:"url"`
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"results"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := json.Unmarshal(body, &searchResp); err != nil {
|
|
||||||
return "", fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
results := searchResp.Results
|
|
||||||
if len(results) == 0 {
|
|
||||||
return fmt.Sprintf("No results for: %s", query), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var lines []string
|
|
||||||
lines = append(lines, fmt.Sprintf("Results for: %s (via Tavily)", query))
|
|
||||||
for i, item := range results {
|
|
||||||
if i >= count {
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
lines = append(lines, fmt.Sprintf("%d. %s\n %s", i+1, item.Title, item.URL))
|
|
||||||
if item.Content != "" {
|
payload := map[string]any{
|
||||||
lines = append(lines, fmt.Sprintf(" %s", item.Content))
|
"api_key": apiKey,
|
||||||
|
"query": query,
|
||||||
|
"search_depth": "advanced",
|
||||||
|
"include_answer": false,
|
||||||
|
"include_images": false,
|
||||||
|
"include_raw_content": false,
|
||||||
|
"max_results": count,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
bodyBytes, err := json.Marshal(payload)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to marshal payload: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", searchURL, bytes.NewBuffer(bodyBytes))
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
|
resp, err := p.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = fmt.Errorf("request failed: %w", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
lastErr = fmt.Errorf("failed to read response: %w", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
lastErr = fmt.Errorf("tavily api error (status %d): %s", resp.StatusCode, string(body))
|
||||||
|
if resp.StatusCode == http.StatusTooManyRequests ||
|
||||||
|
resp.StatusCode == http.StatusUnauthorized ||
|
||||||
|
resp.StatusCode == http.StatusForbidden ||
|
||||||
|
resp.StatusCode >= 500 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return "", lastErr
|
||||||
|
}
|
||||||
|
|
||||||
|
var searchResp struct {
|
||||||
|
Results []struct {
|
||||||
|
Title string `json:"title"`
|
||||||
|
URL string `json:"url"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"results"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &searchResp); err != nil {
|
||||||
|
return "", fmt.Errorf("failed to parse response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
results := searchResp.Results
|
||||||
|
if len(results) == 0 {
|
||||||
|
return fmt.Sprintf("No results for: %s", query), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var lines []string
|
||||||
|
lines = append(lines, fmt.Sprintf("Results for: %s (via Tavily)", query))
|
||||||
|
for i, item := range results {
|
||||||
|
if i >= count {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
lines = append(lines, fmt.Sprintf("%d. %s\n %s", i+1, item.Title, item.URL))
|
||||||
|
if item.Content != "" {
|
||||||
|
lines = append(lines, fmt.Sprintf(" %s", item.Content))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(lines, "\n"), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return strings.Join(lines, "\n"), nil
|
return "", fmt.Errorf("all api keys failed, last error: %w", lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
type DuckDuckGoSearchProvider struct {
|
type DuckDuckGoSearchProvider struct {
|
||||||
|
|
@ -324,75 +406,97 @@ func stripTags(content string) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
type PerplexitySearchProvider struct {
|
type PerplexitySearchProvider struct {
|
||||||
apiKey string
|
keyPool *APIKeyPool
|
||||||
proxy string
|
proxy string
|
||||||
client *http.Client
|
client *http.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
||||||
searchURL := "https://api.perplexity.ai/chat/completions"
|
searchURL := "https://api.perplexity.ai/chat/completions"
|
||||||
|
|
||||||
payload := map[string]any{
|
var lastErr error
|
||||||
"model": "sonar",
|
iter := p.keyPool.NewIterator()
|
||||||
"messages": []map[string]string{
|
|
||||||
{
|
for {
|
||||||
"role": "system",
|
apiKey, ok := iter.Next()
|
||||||
"content": "You are a search assistant. Provide concise search results with titles, URLs, and brief descriptions in the following format:\n1. Title\n URL\n Description\n\nDo not add extra commentary.",
|
if !ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := map[string]any{
|
||||||
|
"model": "sonar",
|
||||||
|
"messages": []map[string]string{
|
||||||
|
{
|
||||||
|
"role": "system",
|
||||||
|
"content": "You are a search assistant. Provide concise search results with titles, URLs, and brief descriptions in the following format:\n1. Title\n URL\n Description\n\nDo not add extra commentary.",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": fmt.Sprintf("Search for: %s. Provide up to %d relevant results.", query, count),
|
||||||
|
},
|
||||||
},
|
},
|
||||||
{
|
"max_tokens": 1000,
|
||||||
"role": "user",
|
}
|
||||||
"content": fmt.Sprintf("Search for: %s. Provide up to %d relevant results.", query, count),
|
|
||||||
},
|
payloadBytes, err := json.Marshal(payload)
|
||||||
},
|
if err != nil {
|
||||||
"max_tokens": 1000,
|
return "", fmt.Errorf("failed to marshal request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", searchURL, strings.NewReader(string(payloadBytes)))
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||||
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
|
resp, err := p.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = fmt.Errorf("request failed: %w", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
lastErr = fmt.Errorf("failed to read response: %w", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
lastErr = fmt.Errorf("Perplexity API error: %s", string(body))
|
||||||
|
if resp.StatusCode == http.StatusTooManyRequests ||
|
||||||
|
resp.StatusCode == http.StatusUnauthorized ||
|
||||||
|
resp.StatusCode == http.StatusForbidden ||
|
||||||
|
resp.StatusCode >= 500 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return "", lastErr
|
||||||
|
}
|
||||||
|
|
||||||
|
var searchResp struct {
|
||||||
|
Choices []struct {
|
||||||
|
Message struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"message"`
|
||||||
|
} `json:"choices"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &searchResp); err != nil {
|
||||||
|
return "", fmt.Errorf("failed to parse response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(searchResp.Choices) == 0 {
|
||||||
|
return fmt.Sprintf("No results for: %s", query), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf("Results for: %s (via Perplexity)\n%s", query, searchResp.Choices[0].Message.Content), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
payloadBytes, err := json.Marshal(payload)
|
return "", fmt.Errorf("all api keys failed, last error: %w", lastErr)
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to marshal request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, "POST", searchURL, strings.NewReader(string(payloadBytes)))
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
|
||||||
req.Header.Set("User-Agent", userAgent)
|
|
||||||
|
|
||||||
resp, err := p.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
return "", fmt.Errorf("Perplexity API error: %s", string(body))
|
|
||||||
}
|
|
||||||
|
|
||||||
var searchResp struct {
|
|
||||||
Choices []struct {
|
|
||||||
Message struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"message"`
|
|
||||||
} `json:"choices"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := json.Unmarshal(body, &searchResp); err != nil {
|
|
||||||
return "", fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(searchResp.Choices) == 0 {
|
|
||||||
return fmt.Sprintf("No results for: %s", query), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return fmt.Sprintf("Results for: %s (via Perplexity)\n%s", query, searchResp.Choices[0].Message.Content), nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type SearXNGSearchProvider struct {
|
type SearXNGSearchProvider struct {
|
||||||
|
|
@ -545,16 +649,16 @@ type WebSearchTool struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type WebSearchToolOptions struct {
|
type WebSearchToolOptions struct {
|
||||||
BraveAPIKey string
|
BraveAPIKeys []string
|
||||||
BraveMaxResults int
|
BraveMaxResults int
|
||||||
BraveEnabled bool
|
BraveEnabled bool
|
||||||
TavilyAPIKey string
|
TavilyAPIKeys []string
|
||||||
TavilyBaseURL string
|
TavilyBaseURL string
|
||||||
TavilyMaxResults int
|
TavilyMaxResults int
|
||||||
TavilyEnabled bool
|
TavilyEnabled bool
|
||||||
DuckDuckGoMaxResults int
|
DuckDuckGoMaxResults int
|
||||||
DuckDuckGoEnabled bool
|
DuckDuckGoEnabled bool
|
||||||
PerplexityAPIKey string
|
PerplexityAPIKeys []string
|
||||||
PerplexityMaxResults int
|
PerplexityMaxResults int
|
||||||
PerplexityEnabled bool
|
PerplexityEnabled bool
|
||||||
SearXNGBaseURL string
|
SearXNGBaseURL string
|
||||||
|
|
@ -571,23 +675,26 @@ type WebSearchToolOptions struct {
|
||||||
func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
|
func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
|
||||||
var provider SearchProvider
|
var provider SearchProvider
|
||||||
maxResults := 5
|
maxResults := 5
|
||||||
|
|
||||||
// Priority: Perplexity > Brave > SearXNG > Tavily > DuckDuckGo > GLM Search
|
// Priority: Perplexity > Brave > SearXNG > Tavily > DuckDuckGo > GLM Search
|
||||||
if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" {
|
if opts.PerplexityEnabled && len(opts.PerplexityAPIKeys) > 0 {
|
||||||
client, err := createHTTPClient(opts.Proxy, perplexityTimeout)
|
client, err := createHTTPClient(opts.Proxy, perplexityTimeout)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for Perplexity: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for Perplexity: %w", err)
|
||||||
}
|
}
|
||||||
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey, proxy: opts.Proxy, client: client}
|
provider = &PerplexitySearchProvider{
|
||||||
|
keyPool: NewAPIKeyPool(opts.PerplexityAPIKeys),
|
||||||
|
proxy: opts.Proxy,
|
||||||
|
client: client,
|
||||||
|
}
|
||||||
if opts.PerplexityMaxResults > 0 {
|
if opts.PerplexityMaxResults > 0 {
|
||||||
maxResults = opts.PerplexityMaxResults
|
maxResults = opts.PerplexityMaxResults
|
||||||
}
|
}
|
||||||
} else if opts.BraveEnabled && opts.BraveAPIKey != "" {
|
} else if opts.BraveEnabled && len(opts.BraveAPIKeys) > 0 {
|
||||||
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for Brave: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for Brave: %w", err)
|
||||||
}
|
}
|
||||||
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey, proxy: opts.Proxy, client: client}
|
provider = &BraveSearchProvider{keyPool: NewAPIKeyPool(opts.BraveAPIKeys), proxy: opts.Proxy, client: client}
|
||||||
if opts.BraveMaxResults > 0 {
|
if opts.BraveMaxResults > 0 {
|
||||||
maxResults = opts.BraveMaxResults
|
maxResults = opts.BraveMaxResults
|
||||||
}
|
}
|
||||||
|
|
@ -596,13 +703,13 @@ func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
|
||||||
if opts.SearXNGMaxResults > 0 {
|
if opts.SearXNGMaxResults > 0 {
|
||||||
maxResults = opts.SearXNGMaxResults
|
maxResults = opts.SearXNGMaxResults
|
||||||
}
|
}
|
||||||
} else if opts.TavilyEnabled && opts.TavilyAPIKey != "" {
|
} else if opts.TavilyEnabled && len(opts.TavilyAPIKeys) > 0 {
|
||||||
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
client, err := createHTTPClient(opts.Proxy, searchTimeout)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create HTTP client for Tavily: %w", err)
|
return nil, fmt.Errorf("failed to create HTTP client for Tavily: %w", err)
|
||||||
}
|
}
|
||||||
provider = &TavilySearchProvider{
|
provider = &TavilySearchProvider{
|
||||||
apiKey: opts.TavilyAPIKey,
|
keyPool: NewAPIKeyPool(opts.TavilyAPIKeys),
|
||||||
baseURL: opts.TavilyBaseURL,
|
baseURL: opts.TavilyBaseURL,
|
||||||
proxy: opts.Proxy,
|
proxy: opts.Proxy,
|
||||||
client: client,
|
client: client,
|
||||||
|
|
|
||||||
|
|
@ -249,7 +249,7 @@ func TestWebFetchTool_PayloadTooLarge(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
|
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
|
||||||
func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""})
|
tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKeys: nil})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Unexpected error: %v", err)
|
t.Fatalf("Unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -269,7 +269,11 @@ func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
||||||
|
|
||||||
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
|
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
|
||||||
func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5})
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
BraveEnabled: true,
|
||||||
|
BraveAPIKeys: []string{"test-key"},
|
||||||
|
BraveMaxResults: 5,
|
||||||
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Unexpected error: %v", err)
|
t.Fatalf("Unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -553,7 +557,7 @@ func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
|
||||||
t.Run("perplexity", func(t *testing.T) {
|
t.Run("perplexity", func(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
PerplexityEnabled: true,
|
PerplexityEnabled: true,
|
||||||
PerplexityAPIKey: "k",
|
PerplexityAPIKeys: []string{"k"},
|
||||||
PerplexityMaxResults: 3,
|
PerplexityMaxResults: 3,
|
||||||
Proxy: "http://127.0.0.1:7890",
|
Proxy: "http://127.0.0.1:7890",
|
||||||
})
|
})
|
||||||
|
|
@ -572,7 +576,7 @@ func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
|
||||||
t.Run("brave", func(t *testing.T) {
|
t.Run("brave", func(t *testing.T) {
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
BraveEnabled: true,
|
BraveEnabled: true,
|
||||||
BraveAPIKey: "k",
|
BraveAPIKeys: []string{"k"},
|
||||||
BraveMaxResults: 3,
|
BraveMaxResults: 3,
|
||||||
Proxy: "http://127.0.0.1:7890",
|
Proxy: "http://127.0.0.1:7890",
|
||||||
})
|
})
|
||||||
|
|
@ -650,7 +654,7 @@ func TestWebTool_TavilySearch_Success(t *testing.T) {
|
||||||
|
|
||||||
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
TavilyEnabled: true,
|
TavilyEnabled: true,
|
||||||
TavilyAPIKey: "test-key",
|
TavilyAPIKeys: []string{"test-key"},
|
||||||
TavilyBaseURL: server.URL,
|
TavilyBaseURL: server.URL,
|
||||||
TavilyMaxResults: 5,
|
TavilyMaxResults: 5,
|
||||||
})
|
})
|
||||||
|
|
@ -682,6 +686,121 @@ func TestWebTool_TavilySearch_Success(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAPIKeyPool(t *testing.T) {
|
||||||
|
pool := NewAPIKeyPool([]string{"key1", "key2", "key3"})
|
||||||
|
if len(pool.keys) != 3 {
|
||||||
|
t.Fatalf("expected 3 keys, got %d", len(pool.keys))
|
||||||
|
}
|
||||||
|
if pool.keys[0] != "key1" || pool.keys[1] != "key2" || pool.keys[2] != "key3" {
|
||||||
|
t.Fatalf("unexpected keys: %v", pool.keys)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test Iterator: each iterator should cover all keys exactly once
|
||||||
|
iter := pool.NewIterator()
|
||||||
|
expected := []string{"key1", "key2", "key3"}
|
||||||
|
for i, want := range expected {
|
||||||
|
k, ok := iter.Next()
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("iter.Next() returned false at step %d", i)
|
||||||
|
}
|
||||||
|
if k != want {
|
||||||
|
t.Errorf("step %d: expected %s, got %s", i, want, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Should be exhausted
|
||||||
|
if _, ok := iter.Next(); ok {
|
||||||
|
t.Errorf("expected iterator exhausted after all keys")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second iterator starts at next position (load balancing)
|
||||||
|
iter2 := pool.NewIterator()
|
||||||
|
k, ok := iter2.Next()
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("iter2.Next() returned false")
|
||||||
|
}
|
||||||
|
if k != "key2" {
|
||||||
|
t.Errorf("expected key2 (round-robin), got %s", k)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Empty pool
|
||||||
|
emptyPool := NewAPIKeyPool([]string{})
|
||||||
|
emptyIter := emptyPool.NewIterator()
|
||||||
|
if _, ok := emptyIter.Next(); ok {
|
||||||
|
t.Errorf("expected false for empty pool")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single key pool
|
||||||
|
singlePool := NewAPIKeyPool([]string{"single"})
|
||||||
|
singleIter := singlePool.NewIterator()
|
||||||
|
if k, ok := singleIter.Next(); !ok || k != "single" {
|
||||||
|
t.Errorf("expected single, got %s (ok=%v)", k, ok)
|
||||||
|
}
|
||||||
|
if _, ok := singleIter.Next(); ok {
|
||||||
|
t.Errorf("expected exhausted after single key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebTool_TavilySearch_Failover(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var payload map[string]any
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
|
||||||
|
t.Fatalf("failed to decode payload: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
apiKey := payload["api_key"].(string)
|
||||||
|
|
||||||
|
if apiKey == "key1" {
|
||||||
|
w.WriteHeader(http.StatusTooManyRequests)
|
||||||
|
w.Write([]byte("Rate limited"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if apiKey == "key2" {
|
||||||
|
// Success
|
||||||
|
response := map[string]any{
|
||||||
|
"results": []map[string]any{
|
||||||
|
{
|
||||||
|
"title": "Success Result",
|
||||||
|
"url": "https://example.com/success",
|
||||||
|
"content": "Success content",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
json.NewEncoder(w).Encode(response)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
tool, err := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
TavilyEnabled: true,
|
||||||
|
TavilyAPIKeys: []string{"key1", "key2"},
|
||||||
|
TavilyBaseURL: server.URL,
|
||||||
|
TavilyMaxResults: 5,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewWebSearchTool() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]any{
|
||||||
|
"query": "test query",
|
||||||
|
}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Errorf("Expected success, got Error: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForUser, "Success Result") {
|
||||||
|
t.Errorf("Expected failover to second key and success result, got: %s", result.ForUser)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestWebTool_GLMSearch_Success(t *testing.T) {
|
func TestWebTool_GLMSearch_Success(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.Method != "POST" {
|
if r.Method != "POST" {
|
||||||
|
|
|
||||||
272
pkg/utils/bm25.go
Normal file
272
pkg/utils/bm25.go
Normal file
|
|
@ -0,0 +1,272 @@
|
||||||
|
// Package utils provides shared, reusable algorithms.
|
||||||
|
// This file implements a generic BM25 search engine.
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// type MyDoc struct { ID string; Body string }
|
||||||
|
//
|
||||||
|
// corpus := []MyDoc{...}
|
||||||
|
// engine := bm25.New(corpus, func(d MyDoc) string {
|
||||||
|
// return d.ID + " " + d.Body
|
||||||
|
// })
|
||||||
|
// results := engine.Search("my query", 5)
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ── Tuning defaults ───────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
const (
|
||||||
|
// DefaultBM25K1 is the term-frequency saturation factor (typical range 1.2–2.0).
|
||||||
|
// Higher values give more weight to repeated terms.
|
||||||
|
DefaultBM25K1 = 1.2
|
||||||
|
|
||||||
|
// DefaultBM25B is the document-length normalization factor (0 = none, 1 = full).
|
||||||
|
DefaultBM25B = 0.75
|
||||||
|
)
|
||||||
|
|
||||||
|
// BM25Engine is a query-time BM25 search engine over a generic corpus.
|
||||||
|
// T is the document type; the caller supplies a TextFunc that extracts the
|
||||||
|
// searchable text from each document.
|
||||||
|
//
|
||||||
|
// The engine is stateless between queries: no caching, no invalidation logic.
|
||||||
|
// All indexing work is performed inside Search() on every call, making it
|
||||||
|
// safe to use on corpora that change frequently.
|
||||||
|
type BM25Engine[T any] struct {
|
||||||
|
corpus []T
|
||||||
|
textFunc func(T) string
|
||||||
|
k1 float64
|
||||||
|
b float64
|
||||||
|
}
|
||||||
|
|
||||||
|
// BM25Option is a functional option to configure a BM25Engine.
|
||||||
|
type BM25Option func(*bm25Config)
|
||||||
|
|
||||||
|
type bm25Config struct {
|
||||||
|
k1 float64
|
||||||
|
b float64
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithK1 overrides the term-frequency saturation constant (default 1.2).
|
||||||
|
func WithK1(k1 float64) BM25Option {
|
||||||
|
return func(c *bm25Config) { c.k1 = k1 }
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithB overrides the document-length normalization factor (default 0.75).
|
||||||
|
func WithB(b float64) BM25Option {
|
||||||
|
return func(c *bm25Config) { c.b = b }
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBM25Engine creates a BM25Engine for the given corpus.
|
||||||
|
//
|
||||||
|
// - corpus : slice of documents of any type T.
|
||||||
|
// - textFunc : function that returns the searchable text for a document.
|
||||||
|
// - opts : optional tuning (WithK1, WithB).
|
||||||
|
//
|
||||||
|
// The corpus slice is referenced, not copied. Callers must not mutate it
|
||||||
|
// concurrently with Search().
|
||||||
|
func NewBM25Engine[T any](corpus []T, textFunc func(T) string, opts ...BM25Option) *BM25Engine[T] {
|
||||||
|
cfg := bm25Config{k1: DefaultBM25K1, b: DefaultBM25B}
|
||||||
|
for _, o := range opts {
|
||||||
|
o(&cfg)
|
||||||
|
}
|
||||||
|
return &BM25Engine[T]{
|
||||||
|
corpus: corpus,
|
||||||
|
textFunc: textFunc,
|
||||||
|
k1: cfg.k1,
|
||||||
|
b: cfg.b,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BM25Result is a single ranked result from a Search call.
|
||||||
|
type BM25Result[T any] struct {
|
||||||
|
Document T
|
||||||
|
Score float32
|
||||||
|
}
|
||||||
|
|
||||||
|
// Search ranks the corpus against query and returns the top-k results.
|
||||||
|
// Returns an empty slice (not nil) when there are no matches.
|
||||||
|
//
|
||||||
|
// Complexity: O(N×L) for indexing + O(|Q|×avgPostingLen) for scoring,
|
||||||
|
// where N = corpus size, L = average document length, Q = query terms.
|
||||||
|
// Top-k extraction uses a fixed-size min-heap: O(candidates × log k).
|
||||||
|
func (e *BM25Engine[T]) Search(query string, topK int) []BM25Result[T] {
|
||||||
|
if topK <= 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
queryTerms := bm25Tokenize(query)
|
||||||
|
if len(queryTerms) == 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
N := len(e.corpus)
|
||||||
|
if N == 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 1: build per-document tf + raw doc lengths
|
||||||
|
type docEntry struct {
|
||||||
|
tf map[string]uint32
|
||||||
|
rawLen int
|
||||||
|
}
|
||||||
|
|
||||||
|
entries := make([]docEntry, N)
|
||||||
|
df := make(map[string]int, 64)
|
||||||
|
totalLen := 0
|
||||||
|
|
||||||
|
for i, doc := range e.corpus {
|
||||||
|
tokens := bm25Tokenize(e.textFunc(doc))
|
||||||
|
totalLen += len(tokens)
|
||||||
|
|
||||||
|
tf := make(map[string]uint32, len(tokens))
|
||||||
|
for _, t := range tokens {
|
||||||
|
tf[t]++
|
||||||
|
}
|
||||||
|
// df: each term counts once per document (iterate the map, keys are unique)
|
||||||
|
for t := range tf {
|
||||||
|
df[t]++
|
||||||
|
}
|
||||||
|
|
||||||
|
entries[i] = docEntry{tf: tf, rawLen: len(tokens)}
|
||||||
|
}
|
||||||
|
|
||||||
|
avgDocLen := float64(totalLen) / float64(N)
|
||||||
|
|
||||||
|
// Step 2: pre-compute IDF and per-doc length normalization
|
||||||
|
// IDF (Robertson smoothing): log( (N - df(t) + 0.5) / (df(t) + 0.5) + 1 )
|
||||||
|
idf := make(map[string]float32, len(df))
|
||||||
|
for term, freq := range df {
|
||||||
|
idf[term] = float32(math.Log(
|
||||||
|
(float64(N)-float64(freq)+0.5)/(float64(freq)+0.5) + 1,
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
// docLenNorm[i] = k1 * (1 - b + b * |doc_i| / avgDocLen)
|
||||||
|
// Stored as float32 — sufficient precision for ranking.
|
||||||
|
docLenNorm := make([]float32, N)
|
||||||
|
for i, entry := range entries {
|
||||||
|
docLenNorm[i] = float32(e.k1 * (1 - e.b + e.b*float64(entry.rawLen)/avgDocLen))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 3: build inverted index (posting lists)
|
||||||
|
// Iterate the tf map directly — map keys are already unique, no seen-set needed.
|
||||||
|
posting := make(map[string][]int32, len(df))
|
||||||
|
for i, entry := range entries {
|
||||||
|
for term := range entry.tf {
|
||||||
|
posting[term] = append(posting[term], int32(i))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 4: score via posting lists
|
||||||
|
// Deduplicate query terms to avoid double-weighting the same term.
|
||||||
|
unique := bm25Dedupe(queryTerms)
|
||||||
|
|
||||||
|
scores := make(map[int32]float32)
|
||||||
|
for _, term := range unique {
|
||||||
|
termIDF, ok := idf[term]
|
||||||
|
if !ok {
|
||||||
|
continue // term not in vocabulary → zero contribution
|
||||||
|
}
|
||||||
|
for _, docID := range posting[term] {
|
||||||
|
freq := float32(entries[docID].tf[term])
|
||||||
|
// TF_norm = freq * (k1+1) / (freq + docLenNorm)
|
||||||
|
tfNorm := freq * float32(e.k1+1) / (freq + docLenNorm[docID])
|
||||||
|
scores[docID] += termIDF * tfNorm
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(scores) == 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 5: top-K via fixed-size min-heap
|
||||||
|
heap := make([]bm25ScoredDoc, 0, topK)
|
||||||
|
|
||||||
|
for docID, sc := range scores {
|
||||||
|
switch {
|
||||||
|
case len(heap) < topK:
|
||||||
|
heap = append(heap, bm25ScoredDoc{docID: docID, score: sc})
|
||||||
|
if len(heap) == topK {
|
||||||
|
bm25MinHeapify(heap)
|
||||||
|
}
|
||||||
|
case sc > heap[0].score:
|
||||||
|
heap[0] = bm25ScoredDoc{docID: docID, score: sc}
|
||||||
|
bm25SiftDown(heap, 0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sort.Slice(heap, func(i, j int) bool { return heap[i].score > heap[j].score })
|
||||||
|
|
||||||
|
out := make([]BM25Result[T], len(heap))
|
||||||
|
for i, h := range heap {
|
||||||
|
out[i] = BM25Result[T]{
|
||||||
|
Document: e.corpus[h.docID],
|
||||||
|
Score: h.score,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25Tokenize splits s into lowercase tokens, stripping edge punctuation.
|
||||||
|
func bm25Tokenize(s string) []string {
|
||||||
|
raw := strings.Fields(strings.ToLower(s))
|
||||||
|
out := raw[:0] // reuse backing array to avoid extra allocation
|
||||||
|
for _, t := range raw {
|
||||||
|
t = strings.Trim(t, ".,;:!?\"'()/\\-_")
|
||||||
|
if t != "" {
|
||||||
|
out = append(out, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25Dedupe returns a new slice with duplicate tokens removed,
|
||||||
|
// preserving first-occurrence order.
|
||||||
|
func bm25Dedupe(tokens []string) []string {
|
||||||
|
seen := make(map[string]struct{}, len(tokens))
|
||||||
|
out := make([]string, 0, len(tokens))
|
||||||
|
for _, t := range tokens {
|
||||||
|
if _, ok := seen[t]; !ok {
|
||||||
|
seen[t] = struct{}{}
|
||||||
|
out = append(out, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
type bm25ScoredDoc struct {
|
||||||
|
docID int32
|
||||||
|
score float32
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25MinHeapify builds a min-heap in-place using Floyd's algorithm: O(k).
|
||||||
|
func bm25MinHeapify(h []bm25ScoredDoc) {
|
||||||
|
for i := len(h)/2 - 1; i >= 0; i-- {
|
||||||
|
bm25SiftDown(h, i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25SiftDown restores the min-heap property starting at node i: O(log k).
|
||||||
|
func bm25SiftDown(h []bm25ScoredDoc, i int) {
|
||||||
|
n := len(h)
|
||||||
|
for {
|
||||||
|
smallest := i
|
||||||
|
l, r := 2*i+1, 2*i+2
|
||||||
|
if l < n && h[l].score < h[smallest].score {
|
||||||
|
smallest = l
|
||||||
|
}
|
||||||
|
if r < n && h[r].score < h[smallest].score {
|
||||||
|
smallest = r
|
||||||
|
}
|
||||||
|
if smallest == i {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
h[i], h[smallest] = h[smallest], h[i]
|
||||||
|
i = smallest
|
||||||
|
}
|
||||||
|
}
|
||||||
175
pkg/utils/bm25_test.go
Normal file
175
pkg/utils/bm25_test.go
Normal file
|
|
@ -0,0 +1,175 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testDoc is a generic structure for use in tests.
|
||||||
|
type testDoc struct {
|
||||||
|
ID int
|
||||||
|
Text string
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractText(d testDoc) string {
|
||||||
|
return d.Text
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_EdgeCases(t *testing.T) {
|
||||||
|
corpus := []testDoc{
|
||||||
|
{1, "hello world"},
|
||||||
|
{2, "foo bar"},
|
||||||
|
}
|
||||||
|
engine := NewBM25Engine(corpus, extractText)
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
query string
|
||||||
|
topK int
|
||||||
|
}{
|
||||||
|
{"Zero topK", "hello", 0},
|
||||||
|
{"Negative topK", "hello", -1},
|
||||||
|
{"Empty query", "", 5},
|
||||||
|
{"Query with only punctuation", "...,,,!!!", 5},
|
||||||
|
{"No matches found", "golang", 5},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
results := engine.Search(tt.query, tt.topK)
|
||||||
|
if len(results) != 0 {
|
||||||
|
t.Errorf("expected 0 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
// Check that it never returns nil, but an empty slice
|
||||||
|
if results == nil {
|
||||||
|
t.Errorf("expected empty slice, got nil")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_EmptyCorpus(t *testing.T) {
|
||||||
|
engine := NewBM25Engine([]testDoc{}, extractText)
|
||||||
|
results := engine.Search("hello", 5)
|
||||||
|
if len(results) != 0 || results == nil {
|
||||||
|
t.Errorf("expected empty slice from empty corpus, got %v", results)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_RankingLogic(t *testing.T) {
|
||||||
|
corpus := []testDoc{
|
||||||
|
{1, "the quick brown fox jumps over the lazy dog"},
|
||||||
|
{2, "quick fox"},
|
||||||
|
{3, "quick quick quick fox"}, // High Term Frequency (TF)
|
||||||
|
{4, "completely irrelevant document here"},
|
||||||
|
}
|
||||||
|
engine := NewBM25Engine(corpus, extractText)
|
||||||
|
|
||||||
|
t.Run("Term Frequency (TF) boosts score", func(t *testing.T) {
|
||||||
|
results := engine.Search("quick", 5)
|
||||||
|
if len(results) < 3 {
|
||||||
|
t.Fatalf("expected at least 3 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
// Doc 3 has the word "quick" repeated 3 times, it should beat Doc 2
|
||||||
|
if results[0].Document.ID != 3 {
|
||||||
|
t.Errorf("expected doc 3 to rank first due to high TF, got doc %d", results[0].Document.ID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Document Length penalty", func(t *testing.T) {
|
||||||
|
results := engine.Search("fox", 5)
|
||||||
|
if len(results) < 3 {
|
||||||
|
t.Fatalf("expected at least 3 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
// Doc 2 ("quick fox") is much shorter than Doc 1 ("the quick brown fox..."),
|
||||||
|
// so, with equal Term Frequency for the word "fox" (1 time), Doc 2 wins.
|
||||||
|
if results[0].Document.ID != 2 {
|
||||||
|
t.Errorf("expected doc 2 to rank first due to shorter length, got doc %d", results[0].Document.ID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("TopK limits results", func(t *testing.T) {
|
||||||
|
results := engine.Search("quick", 2)
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Errorf("expected exactly 2 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Tokenize(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input string
|
||||||
|
expected []string
|
||||||
|
}{
|
||||||
|
{"Hello World", []string{"hello", "world"}},
|
||||||
|
{" spaces everywhere ", []string{"spaces", "everywhere"}},
|
||||||
|
{"punctuation... test!!!", []string{"punctuation", "test"}},
|
||||||
|
{"(parentheses) and-hyphens", []string{"parentheses", "and-hyphens"}}, // hyphens trimmed from edges
|
||||||
|
{"internal-hyphen is kept", []string{"internal-hyphen", "is", "kept"}},
|
||||||
|
{".,;?!", []string{}}, // Becomes empty after trim
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.input, func(t *testing.T) {
|
||||||
|
got := bm25Tokenize(tt.input)
|
||||||
|
if len(got) == 0 && len(tt.expected) == 0 {
|
||||||
|
return // Both empty
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(got, tt.expected) {
|
||||||
|
t.Errorf("bm25Tokenize(%q) = %v, want %v", tt.input, got, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Dedupe(t *testing.T) {
|
||||||
|
input := []string{"apple", "banana", "apple", "orange", "banana"}
|
||||||
|
expected := []string{"apple", "banana", "orange"}
|
||||||
|
|
||||||
|
got := bm25Dedupe(input)
|
||||||
|
if !reflect.DeepEqual(got, expected) {
|
||||||
|
t.Errorf("bm25Dedupe() = %v, want %v", got, expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Options(t *testing.T) {
|
||||||
|
corpus := []testDoc{{1, "test"}}
|
||||||
|
|
||||||
|
engine := NewBM25Engine(
|
||||||
|
corpus,
|
||||||
|
extractText,
|
||||||
|
WithK1(2.5),
|
||||||
|
WithB(0.9),
|
||||||
|
)
|
||||||
|
|
||||||
|
if engine.k1 != 2.5 {
|
||||||
|
t.Errorf("expected k1 to be 2.5, got %v", engine.k1)
|
||||||
|
}
|
||||||
|
if engine.b != 0.9 {
|
||||||
|
t.Errorf("expected b to be 0.9, got %v", engine.b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_SortingStability(t *testing.T) {
|
||||||
|
// Ensure that sorting by heap returns in correct descending order
|
||||||
|
corpus := []testDoc{
|
||||||
|
{1, "golang is good"},
|
||||||
|
{2, "golang golang"},
|
||||||
|
{3, "golang golang golang"},
|
||||||
|
{4, "golang golang golang golang"},
|
||||||
|
}
|
||||||
|
engine := NewBM25Engine(corpus, extractText)
|
||||||
|
results := engine.Search("golang", 10)
|
||||||
|
|
||||||
|
if len(results) != 4 {
|
||||||
|
t.Fatalf("expected 4 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Score should be strictly decreasing
|
||||||
|
for i := 1; i < len(results); i++ {
|
||||||
|
if results[i].Score > results[i-1].Score {
|
||||||
|
t.Errorf("results not sorted correctly: result %d score (%v) > result %d score (%v)",
|
||||||
|
i, results[i].Score, i-1, results[i-1].Score)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
38
web/Makefile
Normal file
38
web/Makefile
Normal file
|
|
@ -0,0 +1,38 @@
|
||||||
|
.PHONY: dev dev-frontend dev-backend build test lint clean
|
||||||
|
|
||||||
|
# Run both frontend and backend dev servers
|
||||||
|
dev:
|
||||||
|
@if [ ! -f backend/picoclaw-web ] || [ ! -d backend/dist ]; then \
|
||||||
|
echo "Build artifacts not found, building..."; \
|
||||||
|
$(MAKE) build; \
|
||||||
|
fi
|
||||||
|
@echo "Starting backend and frontend dev servers..."
|
||||||
|
@$(MAKE) dev-backend & $(MAKE) dev-frontend
|
||||||
|
|
||||||
|
# Start frontend dev server (Vite, with proxy to backend)
|
||||||
|
dev-frontend:
|
||||||
|
cd frontend && pnpm dev
|
||||||
|
|
||||||
|
# Start backend dev server
|
||||||
|
dev-backend:
|
||||||
|
cd backend && go run .
|
||||||
|
|
||||||
|
# Build frontend and embed into Go binary
|
||||||
|
build:
|
||||||
|
cd frontend && pnpm build:backend
|
||||||
|
cd backend && go build -o picoclaw-web .
|
||||||
|
|
||||||
|
# Run all tests
|
||||||
|
test:
|
||||||
|
cd backend && go test ./...
|
||||||
|
cd frontend && pnpm lint
|
||||||
|
|
||||||
|
# Lint and format
|
||||||
|
lint:
|
||||||
|
cd backend && go vet ./...
|
||||||
|
cd frontend && pnpm check
|
||||||
|
|
||||||
|
# Clean build artifacts
|
||||||
|
clean:
|
||||||
|
rm -rf frontend/dist backend/dist backend/picoclaw-web
|
||||||
|
mkdir -p backend/dist && touch backend/dist/.gitkeep
|
||||||
51
web/README.md
Normal file
51
web/README.md
Normal file
|
|
@ -0,0 +1,51 @@
|
||||||
|
# Picoclaw Web
|
||||||
|
|
||||||
|
This directory contains the standalone web service for `picoclaw`.
|
||||||
|
It provides a complete unified web interface, acting as a dashboard, configuration center, and interactive console (channel client) for the core `picoclaw` engine.
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
|
The service is structured as a monorepo containing both the backend and frontend code to ensure high cohesion and simplify deployment.
|
||||||
|
|
||||||
|
* **`backend/`**: The Go-based web server. It provides RESTful APIs, manages WebSocket connections for chat, and handles the lifecycle of the `picoclaw` process. It eventually embeds the compiled frontend assets into a single executable.
|
||||||
|
* **`frontend/`**: The Vite + React + TanStack Router single-page application (SPA). It provides the interactive user interface.
|
||||||
|
|
||||||
|
## Getting Started
|
||||||
|
|
||||||
|
### Prerequisites
|
||||||
|
|
||||||
|
* Go 1.25+
|
||||||
|
* Node.js 20+ with pnpm
|
||||||
|
|
||||||
|
### Development
|
||||||
|
|
||||||
|
Run both the frontend dev server and the Go backend simultaneously:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
make dev
|
||||||
|
```
|
||||||
|
|
||||||
|
Or run them separately:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
make dev-frontend # Vite dev server
|
||||||
|
make dev-backend # Go backend
|
||||||
|
```
|
||||||
|
|
||||||
|
### Build
|
||||||
|
|
||||||
|
Build the frontend and embed it into a single Go binary:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
make build
|
||||||
|
```
|
||||||
|
|
||||||
|
The output binary is `backend/picoclaw-web`.
|
||||||
|
|
||||||
|
### Other Commands
|
||||||
|
|
||||||
|
```bash
|
||||||
|
make test # Run backend tests and frontend lint
|
||||||
|
make lint # Run go vet and prettier/eslint
|
||||||
|
make clean # Remove all build artifacts
|
||||||
|
```
|
||||||
19
web/backend/.gitignore
vendored
Normal file
19
web/backend/.gitignore
vendored
Normal file
|
|
@ -0,0 +1,19 @@
|
||||||
|
# Go build output
|
||||||
|
*.exe
|
||||||
|
*.dll
|
||||||
|
*.so
|
||||||
|
*.dylib
|
||||||
|
*.test
|
||||||
|
*.out
|
||||||
|
picoclaw-web
|
||||||
|
|
||||||
|
# Frontend build artifacts (embedded by Go)
|
||||||
|
dist/*
|
||||||
|
!dist/.gitkeep
|
||||||
|
|
||||||
|
# OS
|
||||||
|
.DS_Store
|
||||||
|
|
||||||
|
# Editors
|
||||||
|
.vscode/
|
||||||
|
.idea/
|
||||||
47
web/backend/api/channels.go
Normal file
47
web/backend/api/channels.go
Normal file
|
|
@ -0,0 +1,47 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
type channelCatalogItem struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
ConfigKey string `json:"config_key"`
|
||||||
|
Variant string `json:"variant,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
var channelCatalog = []channelCatalogItem{
|
||||||
|
{Name: "telegram", ConfigKey: "telegram"},
|
||||||
|
{Name: "discord", ConfigKey: "discord"},
|
||||||
|
{Name: "slack", ConfigKey: "slack"},
|
||||||
|
{Name: "feishu", ConfigKey: "feishu"},
|
||||||
|
{Name: "dingtalk", ConfigKey: "dingtalk"},
|
||||||
|
{Name: "line", ConfigKey: "line"},
|
||||||
|
{Name: "qq", ConfigKey: "qq"},
|
||||||
|
{Name: "onebot", ConfigKey: "onebot"},
|
||||||
|
{Name: "wecom", ConfigKey: "wecom"},
|
||||||
|
{Name: "wecom_app", ConfigKey: "wecom_app"},
|
||||||
|
{Name: "wecom_aibot", ConfigKey: "wecom_aibot"},
|
||||||
|
{Name: "whatsapp", ConfigKey: "whatsapp", Variant: "bridge"},
|
||||||
|
{Name: "whatsapp_native", ConfigKey: "whatsapp", Variant: "native"},
|
||||||
|
{Name: "pico", ConfigKey: "pico"},
|
||||||
|
{Name: "maixcam", ConfigKey: "maixcam"},
|
||||||
|
{Name: "matrix", ConfigKey: "matrix"},
|
||||||
|
{Name: "irc", ConfigKey: "irc"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// registerChannelRoutes binds read-only channel catalog endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerChannelRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/channels/catalog", h.handleListChannelCatalog)
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleListChannelCatalog returns the channels supported by backend.
|
||||||
|
//
|
||||||
|
// GET /api/channels/catalog
|
||||||
|
func (h *Handler) handleListChannelCatalog(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"channels": channelCatalog,
|
||||||
|
})
|
||||||
|
}
|
||||||
221
web/backend/api/config.go
Normal file
221
web/backend/api/config.go
Normal file
|
|
@ -0,0 +1,221 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// registerConfigRoutes binds configuration management endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerConfigRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/config", h.handleGetConfig)
|
||||||
|
mux.HandleFunc("PUT /api/config", h.handleUpdateConfig)
|
||||||
|
mux.HandleFunc("PATCH /api/config", h.handlePatchConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
// loadFilteredConfig loads the configuration and filters out default placeholder credentials
|
||||||
|
// (like API limits/keys) if the configuration file has not been created yet by the user.
|
||||||
|
func (h *Handler) loadFilteredConfig() (*config.Config, error) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
configExists := false
|
||||||
|
if h.configPath != "" {
|
||||||
|
if _, err := os.Stat(h.configPath); err == nil {
|
||||||
|
configExists = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !configExists {
|
||||||
|
for i := range cfg.ModelList {
|
||||||
|
cfg.ModelList[i].APIKey = ""
|
||||||
|
cfg.ModelList[i].AuthMethod = ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return cfg, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGetConfig returns the complete system configuration.
|
||||||
|
//
|
||||||
|
// GET /api/config
|
||||||
|
func (h *Handler) handleGetConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := h.loadFilteredConfig()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
if err := json.NewEncoder(w).Encode(cfg); err != nil {
|
||||||
|
http.Error(w, "Failed to encode response", http.StatusInternalServerError)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleUpdateConfig updates the complete system configuration.
|
||||||
|
//
|
||||||
|
// PUT /api/config
|
||||||
|
func (h *Handler) handleUpdateConfig(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 cfg config.Config
|
||||||
|
if err := json.Unmarshal(body, &cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if errs := validateConfig(&cfg); len(errs) > 0 {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "validation_error",
|
||||||
|
"errors": errs,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, &cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handlePatchConfig partially updates the system configuration using JSON Merge Patch (RFC 7396).
|
||||||
|
// Only the fields present in the request body will be updated; all other fields remain unchanged.
|
||||||
|
//
|
||||||
|
// PATCH /api/config
|
||||||
|
func (h *Handler) handlePatchConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
|
patchBody, 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()
|
||||||
|
|
||||||
|
// Validate the patch is valid JSON
|
||||||
|
var patch map[string]any
|
||||||
|
if err = json.Unmarshal(patchBody, &patch); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load existing config and marshal to a map for merging
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
existing, err := json.Marshal(cfg)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Failed to serialize current config", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var base map[string]any
|
||||||
|
if err = json.Unmarshal(existing, &base); err != nil {
|
||||||
|
http.Error(w, "Failed to parse current config", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Recursively merge patch into base
|
||||||
|
mergeMap(base, patch)
|
||||||
|
|
||||||
|
// Convert merged map back to Config struct
|
||||||
|
merged, err := json.Marshal(base)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Failed to serialize merged config", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var newCfg config.Config
|
||||||
|
if err := json.Unmarshal(merged, &newCfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Merged config is invalid: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if errs := validateConfig(&newCfg); len(errs) > 0 {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "validation_error",
|
||||||
|
"errors": errs,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, &newCfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateConfig checks the config for common errors before saving.
|
||||||
|
// Returns a list of human-readable error strings; empty means valid.
|
||||||
|
func validateConfig(cfg *config.Config) []string {
|
||||||
|
var errs []string
|
||||||
|
|
||||||
|
// Validate model_list entries
|
||||||
|
if err := cfg.ValidateModelList(); err != nil {
|
||||||
|
errs = append(errs, err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Gateway port range
|
||||||
|
if cfg.Gateway.Port != 0 && (cfg.Gateway.Port < 1 || cfg.Gateway.Port > 65535) {
|
||||||
|
errs = append(errs, fmt.Sprintf("gateway.port %d is out of valid range (1-65535)", cfg.Gateway.Port))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pico channel: token required when enabled
|
||||||
|
if cfg.Channels.Pico.Enabled && cfg.Channels.Pico.Token == "" {
|
||||||
|
errs = append(errs, "channels.pico.token is required when pico channel is enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Telegram: token required when enabled
|
||||||
|
if cfg.Channels.Telegram.Enabled && cfg.Channels.Telegram.Token == "" {
|
||||||
|
errs = append(errs, "channels.telegram.token is required when telegram channel is enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Discord: token required when enabled
|
||||||
|
if cfg.Channels.Discord.Enabled && cfg.Channels.Discord.Token == "" {
|
||||||
|
errs = append(errs, "channels.discord.token is required when discord channel is enabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
return errs
|
||||||
|
}
|
||||||
|
|
||||||
|
// mergeMap recursively merges src into dst (JSON Merge Patch semantics).
|
||||||
|
// - If a key in src has a null value, it is deleted from dst.
|
||||||
|
// - If both dst and src have a nested object for the same key, merge recursively.
|
||||||
|
// - Otherwise the value from src overwrites dst.
|
||||||
|
func mergeMap(dst, src map[string]any) {
|
||||||
|
for key, srcVal := range src {
|
||||||
|
if srcVal == nil {
|
||||||
|
delete(dst, key)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
srcMap, srcIsMap := srcVal.(map[string]any)
|
||||||
|
dstMap, dstIsMap := dst[key].(map[string]any)
|
||||||
|
if srcIsMap && dstIsMap {
|
||||||
|
mergeMap(dstMap, srcMap)
|
||||||
|
} else {
|
||||||
|
dst[key] = srcVal
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
62
web/backend/api/events.go
Normal file
62
web/backend/api/events.go
Normal file
|
|
@ -0,0 +1,62 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GatewayEvent represents a state change event for the gateway process.
|
||||||
|
type GatewayEvent struct {
|
||||||
|
Status string `json:"gateway_status"` // "running", "starting", "stopped", "error"
|
||||||
|
PID int `json:"pid,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventBroadcaster manages SSE client subscriptions and broadcasts events.
|
||||||
|
type EventBroadcaster struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
clients map[chan string]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewEventBroadcaster creates a new broadcaster.
|
||||||
|
func NewEventBroadcaster() *EventBroadcaster {
|
||||||
|
return &EventBroadcaster{
|
||||||
|
clients: make(map[chan string]struct{}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Subscribe adds a new listener channel and returns it.
|
||||||
|
// The caller must call Unsubscribe when done.
|
||||||
|
func (b *EventBroadcaster) Subscribe() chan string {
|
||||||
|
ch := make(chan string, 8)
|
||||||
|
b.mu.Lock()
|
||||||
|
b.clients[ch] = struct{}{}
|
||||||
|
b.mu.Unlock()
|
||||||
|
return ch
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unsubscribe removes a listener channel and closes it.
|
||||||
|
func (b *EventBroadcaster) Unsubscribe(ch chan string) {
|
||||||
|
b.mu.Lock()
|
||||||
|
delete(b.clients, ch)
|
||||||
|
b.mu.Unlock()
|
||||||
|
close(ch)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Broadcast sends a GatewayEvent to all connected SSE clients.
|
||||||
|
func (b *EventBroadcaster) Broadcast(event GatewayEvent) {
|
||||||
|
data, err := json.Marshal(event)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
b.mu.RLock()
|
||||||
|
defer b.mu.RUnlock()
|
||||||
|
|
||||||
|
for ch := range b.clients {
|
||||||
|
// Non-blocking send; drop event if client is slow
|
||||||
|
select {
|
||||||
|
case ch <- string(data):
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
555
web/backend/api/gateway.go
Normal file
555
web/backend/api/gateway.go
Normal file
|
|
@ -0,0 +1,555 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// gateway holds the state for the managed gateway process.
|
||||||
|
var gateway = struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
cmd *exec.Cmd
|
||||||
|
logs *LogBuffer
|
||||||
|
events *EventBroadcaster
|
||||||
|
}{
|
||||||
|
logs: NewLogBuffer(200),
|
||||||
|
events: NewEventBroadcaster(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// registerGatewayRoutes binds gateway lifecycle endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerGatewayRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/gateway/status", h.handleGatewayStatus)
|
||||||
|
mux.HandleFunc("GET /api/gateway/events", h.handleGatewayEvents)
|
||||||
|
mux.HandleFunc("POST /api/gateway/start", h.handleGatewayStart)
|
||||||
|
mux.HandleFunc("POST /api/gateway/stop", h.handleGatewayStop)
|
||||||
|
mux.HandleFunc("POST /api/gateway/restart", h.handleGatewayRestart)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TryAutoStartGateway checks whether gateway start preconditions are met and
|
||||||
|
// starts it when possible. Intended to be called by the backend at startup.
|
||||||
|
func (h *Handler) TryAutoStartGateway() {
|
||||||
|
gateway.mu.Lock()
|
||||||
|
defer gateway.mu.Unlock()
|
||||||
|
|
||||||
|
if isGatewayProcessAliveLocked() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if gateway.cmd != nil && gateway.cmd.Process != nil {
|
||||||
|
gateway.cmd = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("Skip auto-starting gateway: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
log.Printf("Skip auto-starting gateway: %s", reason)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pid, err := h.startGatewayLocked()
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("Failed to auto-start gateway: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Printf("Gateway auto-started (PID: %d)", pid)
|
||||||
|
}
|
||||||
|
|
||||||
|
// gatewayStartReady validates whether current config can start the gateway.
|
||||||
|
func (h *Handler) gatewayStartReady() (bool, string, error) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
return false, "", fmt.Errorf("failed to load config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
modelName := strings.TrimSpace(cfg.Agents.Defaults.GetModelName())
|
||||||
|
if modelName == "" {
|
||||||
|
return false, "no default model configured", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
modelCfg := lookupModelConfig(cfg, modelName)
|
||||||
|
if modelCfg == nil {
|
||||||
|
return false, fmt.Sprintf("default model %q is invalid", modelName), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
hasCredential := strings.TrimSpace(modelCfg.APIKey) != "" ||
|
||||||
|
strings.TrimSpace(modelCfg.AuthMethod) != ""
|
||||||
|
if !hasCredential {
|
||||||
|
return false, fmt.Sprintf("default model %q has no credentials configured", modelName), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return true, "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func lookupModelConfig(cfg *config.Config, modelName string) *config.ModelConfig {
|
||||||
|
modelCfg, err := cfg.GetModelConfig(modelName)
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return modelCfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func isGatewayProcessAliveLocked() bool {
|
||||||
|
return isCmdProcessAliveLocked(gateway.cmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
func isCmdProcessAliveLocked(cmd *exec.Cmd) bool {
|
||||||
|
if cmd == nil || cmd.Process == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait() sets ProcessState when the process exits; use it when available.
|
||||||
|
if cmd.ProcessState != nil && cmd.ProcessState.Exited() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Windows does not support Signal(0) probing. If we still own cmd and it
|
||||||
|
// has not reported exit, treat it as alive.
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return cmd.Process.Signal(syscall.Signal(0)) == nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) startGatewayLocked() (int, error) {
|
||||||
|
// Locate the picoclaw executable
|
||||||
|
execPath := findPicoclawBinary()
|
||||||
|
|
||||||
|
cmd := exec.Command(execPath, "gateway")
|
||||||
|
|
||||||
|
stdoutPipe, err := cmd.StdoutPipe()
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to create stdout pipe: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
stderrPipe, err := cmd.StderrPipe()
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to create stderr pipe: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear old logs for this new run
|
||||||
|
gateway.logs.Reset()
|
||||||
|
|
||||||
|
// Ensure Pico Channel is configured before starting gateway
|
||||||
|
if _, err := h.ensurePicoChannel(); err != nil {
|
||||||
|
log.Printf("Warning: failed to ensure pico channel: %v", err)
|
||||||
|
// Non-fatal: gateway can still start without pico channel
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := cmd.Start(); err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to start gateway: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
gateway.cmd = cmd
|
||||||
|
pid := cmd.Process.Pid
|
||||||
|
log.Printf("Started picoclaw gateway (PID: %d) from %s", pid, execPath)
|
||||||
|
|
||||||
|
// Broadcast starting event
|
||||||
|
gateway.events.Broadcast(GatewayEvent{Status: "starting", PID: pid})
|
||||||
|
|
||||||
|
// Capture stdout/stderr in background
|
||||||
|
go scanPipe(stdoutPipe, gateway.logs)
|
||||||
|
go scanPipe(stderrPipe, gateway.logs)
|
||||||
|
|
||||||
|
// Wait for exit in background and clean up
|
||||||
|
go func() {
|
||||||
|
if err := cmd.Wait(); err != nil {
|
||||||
|
log.Printf("Gateway process exited: %v", err)
|
||||||
|
} else {
|
||||||
|
log.Printf("Gateway process exited normally")
|
||||||
|
}
|
||||||
|
|
||||||
|
gateway.mu.Lock()
|
||||||
|
if gateway.cmd == cmd {
|
||||||
|
gateway.cmd = nil
|
||||||
|
}
|
||||||
|
gateway.mu.Unlock()
|
||||||
|
|
||||||
|
// Broadcast stopped event
|
||||||
|
gateway.events.Broadcast(GatewayEvent{Status: "stopped"})
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Start a goroutine to probe health and broadcast "running" once ready
|
||||||
|
go func() {
|
||||||
|
for i := 0; i < 30; i++ { // try for up to 15 seconds
|
||||||
|
time.Sleep(500 * time.Millisecond)
|
||||||
|
gateway.mu.Lock()
|
||||||
|
stillOurs := gateway.cmd == cmd
|
||||||
|
gateway.mu.Unlock()
|
||||||
|
if !stillOurs {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
healthHost := "127.0.0.1"
|
||||||
|
if cfg.Gateway.Host != "" && cfg.Gateway.Host != "0.0.0.0" {
|
||||||
|
healthHost = cfg.Gateway.Host
|
||||||
|
}
|
||||||
|
healthPort := cfg.Gateway.Port
|
||||||
|
if healthPort == 0 {
|
||||||
|
healthPort = 18790
|
||||||
|
}
|
||||||
|
healthURL := fmt.Sprintf("http://%s/health", net.JoinHostPort(healthHost, strconv.Itoa(healthPort)))
|
||||||
|
client := http.Client{Timeout: 1 * time.Second}
|
||||||
|
resp, err := client.Get(healthURL)
|
||||||
|
if err == nil {
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode == http.StatusOK {
|
||||||
|
gateway.events.Broadcast(GatewayEvent{Status: "running", PID: pid})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return pid, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGatewayStart starts the picoclaw gateway subprocess.
|
||||||
|
//
|
||||||
|
// POST /api/gateway/start
|
||||||
|
func (h *Handler) handleGatewayStart(w http.ResponseWriter, r *http.Request) {
|
||||||
|
gateway.mu.Lock()
|
||||||
|
defer gateway.mu.Unlock()
|
||||||
|
|
||||||
|
// Prevent duplicate starts
|
||||||
|
if isGatewayProcessAliveLocked() {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusConflict)
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "already_running",
|
||||||
|
"pid": gateway.cmd.Process.Pid,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if gateway.cmd != nil && gateway.cmd.Process != nil {
|
||||||
|
gateway.cmd = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(
|
||||||
|
w,
|
||||||
|
fmt.Sprintf("Failed to validate gateway start conditions: %v", err),
|
||||||
|
http.StatusInternalServerError,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "precondition_failed",
|
||||||
|
"message": reason,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pid, err := h.startGatewayLocked()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to start gateway: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"pid": pid,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGatewayStop stops the running gateway subprocess gracefully.
|
||||||
|
//
|
||||||
|
// POST /api/gateway/stop
|
||||||
|
func (h *Handler) handleGatewayStop(w http.ResponseWriter, r *http.Request) {
|
||||||
|
gateway.mu.Lock()
|
||||||
|
defer gateway.mu.Unlock()
|
||||||
|
|
||||||
|
if gateway.cmd == nil || gateway.cmd.Process == nil {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "not_running",
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pid := gateway.cmd.Process.Pid
|
||||||
|
|
||||||
|
// Send SIGTERM for graceful shutdown (SIGKILL on Windows)
|
||||||
|
var sigErr error
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
sigErr = gateway.cmd.Process.Kill()
|
||||||
|
} else {
|
||||||
|
sigErr = gateway.cmd.Process.Signal(syscall.SIGTERM)
|
||||||
|
}
|
||||||
|
|
||||||
|
if sigErr != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to stop gateway (PID %d): %v", pid, sigErr), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Printf("Sent stop signal to gateway (PID: %d)", pid)
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"pid": pid,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGatewayRestart stops the gateway (if running) and starts a new instance.
|
||||||
|
//
|
||||||
|
// POST /api/gateway/restart
|
||||||
|
func (h *Handler) handleGatewayRestart(w http.ResponseWriter, r *http.Request) {
|
||||||
|
gateway.mu.Lock()
|
||||||
|
|
||||||
|
// Stop existing process if running
|
||||||
|
if gateway.cmd != nil && gateway.cmd.Process != nil {
|
||||||
|
if isCmdProcessAliveLocked(gateway.cmd) {
|
||||||
|
// Process is alive, send SIGTERM
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
gateway.cmd.Process.Kill()
|
||||||
|
} else {
|
||||||
|
gateway.cmd.Process.Signal(syscall.SIGTERM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait briefly for it to exit
|
||||||
|
gateway.mu.Unlock()
|
||||||
|
time.Sleep(2 * time.Second)
|
||||||
|
gateway.mu.Lock()
|
||||||
|
}
|
||||||
|
gateway.cmd = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
gateway.mu.Unlock()
|
||||||
|
|
||||||
|
// Start fresh via the existing handler
|
||||||
|
h.handleGatewayStart(w, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGatewayStatus returns the gateway run status, health info, and logs.
|
||||||
|
//
|
||||||
|
// GET /api/gateway/status
|
||||||
|
func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) {
|
||||||
|
data := map[string]any{}
|
||||||
|
|
||||||
|
// Check process state
|
||||||
|
gateway.mu.Lock()
|
||||||
|
processAlive := isGatewayProcessAliveLocked()
|
||||||
|
if processAlive {
|
||||||
|
data["pid"] = gateway.cmd.Process.Pid
|
||||||
|
}
|
||||||
|
gateway.mu.Unlock()
|
||||||
|
|
||||||
|
if !processAlive {
|
||||||
|
data["gateway_status"] = "stopped"
|
||||||
|
} else {
|
||||||
|
// Process is alive — probe its health endpoint
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
host := "127.0.0.1"
|
||||||
|
port := 18790
|
||||||
|
if err == nil && cfg != nil {
|
||||||
|
if cfg.Gateway.Host != "" && cfg.Gateway.Host != "0.0.0.0" {
|
||||||
|
host = cfg.Gateway.Host
|
||||||
|
}
|
||||||
|
if cfg.Gateway.Port != 0 {
|
||||||
|
port = cfg.Gateway.Port
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
url := fmt.Sprintf("http://%s/health", net.JoinHostPort(host, strconv.Itoa(port)))
|
||||||
|
client := http.Client{Timeout: 2 * time.Second}
|
||||||
|
resp, err := client.Get(url)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
data["gateway_status"] = "starting"
|
||||||
|
} else {
|
||||||
|
defer resp.Body.Close()
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
data["gateway_status"] = "error"
|
||||||
|
data["status_code"] = resp.StatusCode
|
||||||
|
} else {
|
||||||
|
var healthData map[string]any
|
||||||
|
if decErr := json.NewDecoder(resp.Body).Decode(&healthData); decErr != nil {
|
||||||
|
data["gateway_status"] = "error"
|
||||||
|
} else {
|
||||||
|
for k, v := range healthData {
|
||||||
|
data[k] = v
|
||||||
|
}
|
||||||
|
data["gateway_status"] = "running"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ready, reason, readyErr := h.gatewayStartReady()
|
||||||
|
if readyErr != nil {
|
||||||
|
data["gateway_start_allowed"] = false
|
||||||
|
data["gateway_start_reason"] = readyErr.Error()
|
||||||
|
} else {
|
||||||
|
data["gateway_start_allowed"] = ready
|
||||||
|
if !ready {
|
||||||
|
data["gateway_start_reason"] = reason
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Append incremental log data
|
||||||
|
appendGatewayLogs(r, data)
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// appendGatewayLogs reads log_offset and log_run_id query params from the request
|
||||||
|
// and populates the response data map with incremental log lines.
|
||||||
|
func appendGatewayLogs(r *http.Request, data map[string]any) {
|
||||||
|
clientOffset := 0
|
||||||
|
clientRunID := -1
|
||||||
|
|
||||||
|
if v := r.URL.Query().Get("log_offset"); v != "" {
|
||||||
|
if n, err := strconv.Atoi(v); err == nil {
|
||||||
|
clientOffset = n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if v := r.URL.Query().Get("log_run_id"); v != "" {
|
||||||
|
if n, err := strconv.Atoi(v); err == nil {
|
||||||
|
clientRunID = n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
runID := gateway.logs.RunID()
|
||||||
|
|
||||||
|
if runID == 0 {
|
||||||
|
data["logs"] = []string{}
|
||||||
|
data["log_total"] = 0
|
||||||
|
data["log_run_id"] = 0
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// If runID changed, reset offset to get all logs from new run
|
||||||
|
offset := clientOffset
|
||||||
|
if clientRunID != runID {
|
||||||
|
offset = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
lines, total, runID := gateway.logs.LinesSince(offset)
|
||||||
|
if lines == nil {
|
||||||
|
lines = []string{}
|
||||||
|
}
|
||||||
|
|
||||||
|
data["logs"] = lines
|
||||||
|
data["log_total"] = total
|
||||||
|
data["log_run_id"] = runID
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGatewayEvents serves an SSE stream of gateway state change events.
|
||||||
|
//
|
||||||
|
// GET /api/gateway/events
|
||||||
|
func (h *Handler) handleGatewayEvents(w http.ResponseWriter, r *http.Request) {
|
||||||
|
flusher, ok := w.(http.Flusher)
|
||||||
|
if !ok {
|
||||||
|
http.Error(w, "SSE not supported", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "text/event-stream")
|
||||||
|
w.Header().Set("Cache-Control", "no-cache")
|
||||||
|
w.Header().Set("Connection", "keep-alive")
|
||||||
|
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||||
|
|
||||||
|
// Subscribe to gateway events
|
||||||
|
ch := gateway.events.Subscribe()
|
||||||
|
defer gateway.events.Unsubscribe(ch)
|
||||||
|
|
||||||
|
// Send initial status so the client doesn't start blank
|
||||||
|
initial := h.currentGatewayStatus()
|
||||||
|
fmt.Fprintf(w, "data: %s\n\n", initial)
|
||||||
|
flusher.Flush()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-r.Context().Done():
|
||||||
|
return
|
||||||
|
case data, ok := <-ch:
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fmt.Fprintf(w, "data: %s\n\n", data)
|
||||||
|
flusher.Flush()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// currentGatewayStatus returns the current gateway status as a JSON string.
|
||||||
|
func (h *Handler) currentGatewayStatus() string {
|
||||||
|
gateway.mu.Lock()
|
||||||
|
defer gateway.mu.Unlock()
|
||||||
|
|
||||||
|
data := map[string]any{
|
||||||
|
"gateway_status": "stopped",
|
||||||
|
}
|
||||||
|
if isGatewayProcessAliveLocked() {
|
||||||
|
data["gateway_status"] = "running"
|
||||||
|
data["pid"] = gateway.cmd.Process.Pid
|
||||||
|
}
|
||||||
|
|
||||||
|
ready, reason, readyErr := h.gatewayStartReady()
|
||||||
|
if readyErr != nil {
|
||||||
|
data["gateway_start_allowed"] = false
|
||||||
|
data["gateway_start_reason"] = readyErr.Error()
|
||||||
|
} else {
|
||||||
|
data["gateway_start_allowed"] = ready
|
||||||
|
if !ready {
|
||||||
|
data["gateway_start_reason"] = reason
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
encoded, _ := json.Marshal(data)
|
||||||
|
return string(encoded)
|
||||||
|
}
|
||||||
|
|
||||||
|
// findPicoclawBinary locates the picoclaw executable.
|
||||||
|
// Tries the same directory as the current executable first, then falls back to $PATH.
|
||||||
|
func findPicoclawBinary() string {
|
||||||
|
if exe, err := os.Executable(); err == nil {
|
||||||
|
dir := filepath.Dir(exe)
|
||||||
|
candidate := filepath.Join(dir, "picoclaw")
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
candidate += ".exe"
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(candidate); err == nil && !info.IsDir() {
|
||||||
|
return candidate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "picoclaw"
|
||||||
|
}
|
||||||
|
|
||||||
|
// scanPipe reads lines from r and appends them to buf. Returns when r reaches EOF.
|
||||||
|
func scanPipe(r io.Reader, buf *LogBuffer) {
|
||||||
|
scanner := bufio.NewScanner(r)
|
||||||
|
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024)
|
||||||
|
for scanner.Scan() {
|
||||||
|
buf.Append(scanner.Text())
|
||||||
|
}
|
||||||
|
}
|
||||||
122
web/backend/api/gateway_test.go
Normal file
122
web/backend/api/gateway_test.go
Normal file
|
|
@ -0,0 +1,122 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGatewayStartReady_NoDefaultModel(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = true, want false")
|
||||||
|
}
|
||||||
|
if reason != "no default model configured" {
|
||||||
|
t.Fatalf("gatewayStartReady() reason = %q, want %q", reason, "no default model configured")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_InvalidDefaultModel(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Model = "missing-model"
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = true, want false")
|
||||||
|
}
|
||||||
|
if reason == "" {
|
||||||
|
t.Fatalf("gatewayStartReady() reason is empty")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_ValidDefaultModel(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
|
||||||
|
cfg.ModelList[0].APIKey = "test-key"
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = false, want true (reason=%q)", reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStartReady_DefaultModelWithoutCredential(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
|
||||||
|
cfg.ModelList[0].APIKey = ""
|
||||||
|
cfg.ModelList[0].AuthMethod = ""
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
ready, reason, err := h.gatewayStartReady()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gatewayStartReady() error = %v", err)
|
||||||
|
}
|
||||||
|
if ready {
|
||||||
|
t.Fatalf("gatewayStartReady() ready = true, want false")
|
||||||
|
}
|
||||||
|
if !strings.Contains(reason, "no credentials configured") {
|
||||||
|
t.Fatalf("gatewayStartReady() reason = %q, want contains %q", reason, "no credentials configured")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGatewayStatusIncludesStartConditionWhenNotReady(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/gateway/status", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
var body map[string]any
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil {
|
||||||
|
t.Fatalf("unmarshal response: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
allowed, ok := body["gateway_start_allowed"].(bool)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("gateway_start_allowed missing or not bool: %#v", body["gateway_start_allowed"])
|
||||||
|
}
|
||||||
|
if allowed {
|
||||||
|
t.Fatalf("gateway_start_allowed = true, want false")
|
||||||
|
}
|
||||||
|
if _, ok := body["gateway_start_reason"].(string); !ok {
|
||||||
|
t.Fatalf("gateway_start_reason missing or not string: %#v", body["gateway_start_reason"])
|
||||||
|
}
|
||||||
|
}
|
||||||
85
web/backend/api/launcher_config.go
Normal file
85
web/backend/api/launcher_config.go
Normal file
|
|
@ -0,0 +1,85 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
||||||
|
)
|
||||||
|
|
||||||
|
type launcherConfigPayload struct {
|
||||||
|
Port int `json:"port"`
|
||||||
|
Public bool `json:"public"`
|
||||||
|
AllowedCIDRs []string `json:"allowed_cidrs"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) registerLauncherConfigRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/system/launcher-config", h.handleGetLauncherConfig)
|
||||||
|
mux.HandleFunc("PUT /api/system/launcher-config", h.handleUpdateLauncherConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) launcherConfigPath() string {
|
||||||
|
return launcherconfig.PathForAppConfig(h.configPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) launcherFallbackConfig() launcherconfig.Config {
|
||||||
|
port := h.serverPort
|
||||||
|
if port <= 0 {
|
||||||
|
port = launcherconfig.DefaultPort
|
||||||
|
}
|
||||||
|
return launcherconfig.Config{
|
||||||
|
Port: port,
|
||||||
|
Public: h.serverPublic,
|
||||||
|
AllowedCIDRs: append([]string(nil), h.serverCIDRs...),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) loadLauncherConfig() (launcherconfig.Config, error) {
|
||||||
|
return launcherconfig.Load(h.launcherConfigPath(), h.launcherFallbackConfig())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleGetLauncherConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := h.loadLauncherConfig()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load launcher config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(launcherConfigPayload{
|
||||||
|
Port: cfg.Port,
|
||||||
|
Public: cfg.Public,
|
||||||
|
AllowedCIDRs: append([]string(nil), cfg.AllowedCIDRs...),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleUpdateLauncherConfig(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var payload launcherConfigPayload
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := launcherconfig.Config{
|
||||||
|
Port: payload.Port,
|
||||||
|
Public: payload.Public,
|
||||||
|
AllowedCIDRs: append([]string(nil), payload.AllowedCIDRs...),
|
||||||
|
}
|
||||||
|
if err := launcherconfig.Validate(cfg); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := launcherconfig.Save(h.launcherConfigPath(), cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save launcher config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(launcherConfigPayload{
|
||||||
|
Port: cfg.Port,
|
||||||
|
Public: cfg.Public,
|
||||||
|
AllowedCIDRs: append([]string(nil), cfg.AllowedCIDRs...),
|
||||||
|
})
|
||||||
|
}
|
||||||
115
web/backend/api/launcher_config_test.go
Normal file
115
web/backend/api/launcher_config_test.go
Normal file
|
|
@ -0,0 +1,115 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGetLauncherConfigUsesRuntimeFallback(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
h.SetServerOptions(19999, true, []string{"192.168.1.0/24"})
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/system/launcher-config", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var got launcherConfigPayload
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||||
|
t.Fatalf("unmarshal response: %v", err)
|
||||||
|
}
|
||||||
|
if got.Port != 19999 || !got.Public {
|
||||||
|
t.Fatalf("response = %+v, want port=19999 public=true", got)
|
||||||
|
}
|
||||||
|
if len(got.AllowedCIDRs) != 1 || got.AllowedCIDRs[0] != "192.168.1.0/24" {
|
||||||
|
t.Fatalf("response allowed_cidrs = %v, want [192.168.1.0/24]", got.AllowedCIDRs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPutLauncherConfigPersists(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPut,
|
||||||
|
"/api/system/launcher-config",
|
||||||
|
strings.NewReader(`{"port":18080,"public":true,"allowed_cidrs":["192.168.1.0/24"]}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
path := launcherconfig.PathForAppConfig(configPath)
|
||||||
|
cfg, err := launcherconfig.Load(path, launcherconfig.Default())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("launcherconfig.Load() error = %v", err)
|
||||||
|
}
|
||||||
|
if cfg.Port != 18080 || !cfg.Public {
|
||||||
|
t.Fatalf("saved config = %+v, want port=18080 public=true", cfg)
|
||||||
|
}
|
||||||
|
if len(cfg.AllowedCIDRs) != 1 || cfg.AllowedCIDRs[0] != "192.168.1.0/24" {
|
||||||
|
t.Fatalf("saved config allowed_cidrs = %v, want [192.168.1.0/24]", cfg.AllowedCIDRs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPutLauncherConfigRejectsInvalidPort(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPut,
|
||||||
|
"/api/system/launcher-config",
|
||||||
|
strings.NewReader(`{"port":70000,"public":false}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPutLauncherConfigRejectsInvalidCIDR(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPut,
|
||||||
|
"/api/system/launcher-config",
|
||||||
|
strings.NewReader(`{"port":18080,"public":false,"allowed_cidrs":["bad-cidr"]}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package server
|
package api
|
||||||
|
|
||||||
import "sync"
|
import "sync"
|
||||||
|
|
||||||
|
|
@ -89,11 +89,3 @@ func (b *LogBuffer) RunID() int {
|
||||||
|
|
||||||
return b.runID
|
return b.runID
|
||||||
}
|
}
|
||||||
|
|
||||||
// Total returns the total number of lines appended in the current run.
|
|
||||||
func (b *LogBuffer) Total() int {
|
|
||||||
b.mu.RLock()
|
|
||||||
defer b.mu.RUnlock()
|
|
||||||
|
|
||||||
return b.total
|
|
||||||
}
|
|
||||||
298
web/backend/api/models.go
Normal file
298
web/backend/api/models.go
Normal file
|
|
@ -0,0 +1,298 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// registerModelRoutes binds model list management endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerModelRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/models", h.handleListModels)
|
||||||
|
mux.HandleFunc("POST /api/models", h.handleAddModel)
|
||||||
|
mux.HandleFunc("POST /api/models/default", h.handleSetDefaultModel)
|
||||||
|
mux.HandleFunc("PUT /api/models/{index}", h.handleUpdateModel)
|
||||||
|
mux.HandleFunc("DELETE /api/models/{index}", h.handleDeleteModel)
|
||||||
|
}
|
||||||
|
|
||||||
|
// modelResponse is the JSON structure returned for each model in the list.
|
||||||
|
// All ModelConfig fields are included so the frontend can display and edit them.
|
||||||
|
type modelResponse struct {
|
||||||
|
Index int `json:"index"`
|
||||||
|
ModelName string `json:"model_name"`
|
||||||
|
Model string `json:"model"`
|
||||||
|
APIBase string `json:"api_base,omitempty"`
|
||||||
|
APIKey string `json:"api_key"`
|
||||||
|
Proxy string `json:"proxy,omitempty"`
|
||||||
|
AuthMethod string `json:"auth_method,omitempty"`
|
||||||
|
// Advanced fields
|
||||||
|
ConnectMode string `json:"connect_mode,omitempty"`
|
||||||
|
Workspace string `json:"workspace,omitempty"`
|
||||||
|
RPM int `json:"rpm,omitempty"`
|
||||||
|
MaxTokensField string `json:"max_tokens_field,omitempty"`
|
||||||
|
RequestTimeout int `json:"request_timeout,omitempty"`
|
||||||
|
ThinkingLevel string `json:"thinking_level,omitempty"`
|
||||||
|
// Meta
|
||||||
|
Configured bool `json:"configured"`
|
||||||
|
IsDefault bool `json:"is_default"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleListModels returns all model_list entries with masked API keys.
|
||||||
|
//
|
||||||
|
// GET /api/models
|
||||||
|
func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := h.loadFilteredConfig()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
defaultModel := cfg.Agents.Defaults.GetModelName()
|
||||||
|
|
||||||
|
models := make([]modelResponse, 0, len(cfg.ModelList))
|
||||||
|
for i, m := range cfg.ModelList {
|
||||||
|
models = append(models, modelResponse{
|
||||||
|
Index: i,
|
||||||
|
ModelName: m.ModelName,
|
||||||
|
Model: m.Model,
|
||||||
|
APIBase: m.APIBase,
|
||||||
|
APIKey: maskAPIKey(m.APIKey),
|
||||||
|
Proxy: m.Proxy,
|
||||||
|
AuthMethod: m.AuthMethod,
|
||||||
|
ConnectMode: m.ConnectMode,
|
||||||
|
Workspace: m.Workspace,
|
||||||
|
RPM: m.RPM,
|
||||||
|
MaxTokensField: m.MaxTokensField,
|
||||||
|
RequestTimeout: m.RequestTimeout,
|
||||||
|
ThinkingLevel: m.ThinkingLevel,
|
||||||
|
Configured: m.APIKey != "" || m.AuthMethod != "",
|
||||||
|
IsDefault: m.ModelName == defaultModel,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"models": models,
|
||||||
|
"total": len(models),
|
||||||
|
"default_model": defaultModel,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleAddModel appends a new model configuration entry.
|
||||||
|
//
|
||||||
|
// POST /api/models
|
||||||
|
func (h *Handler) handleAddModel(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 mc config.ModelConfig
|
||||||
|
if err = json.Unmarshal(body, &mc); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = mc.Validate(); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Validation error: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg.ModelList = append(cfg.ModelList, mc)
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"index": len(cfg.ModelList) - 1,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleUpdateModel replaces a model configuration entry at the given index.
|
||||||
|
// If the request body omits api_key (or sends an empty string), the existing
|
||||||
|
// stored key is preserved so callers can update only api_base / proxy without
|
||||||
|
// exposing or clearing the secret.
|
||||||
|
//
|
||||||
|
// PUT /api/models/{index}
|
||||||
|
func (h *Handler) handleUpdateModel(w http.ResponseWriter, r *http.Request) {
|
||||||
|
idx, err := strconv.Atoi(r.PathValue("index"))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Invalid index", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
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 mc config.ModelConfig
|
||||||
|
if err = json.Unmarshal(body, &mc); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = mc.Validate(); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Validation error: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if idx < 0 || idx >= len(cfg.ModelList) {
|
||||||
|
http.Error(w, fmt.Sprintf("Index %d out of range (0-%d)", idx, len(cfg.ModelList)-1), http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Preserve the existing API key when the caller omits it (empty string).
|
||||||
|
// This lets the UI update api_base / proxy without clearing the stored secret.
|
||||||
|
if mc.APIKey == "" {
|
||||||
|
mc.APIKey = cfg.ModelList[idx].APIKey
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg.ModelList[idx] = mc
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleDeleteModel removes a model configuration entry at the given index.
|
||||||
|
//
|
||||||
|
// DELETE /api/models/{index}
|
||||||
|
func (h *Handler) handleDeleteModel(w http.ResponseWriter, r *http.Request) {
|
||||||
|
idx, err := strconv.Atoi(r.PathValue("index"))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Invalid index", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if idx < 0 || idx >= len(cfg.ModelList) {
|
||||||
|
http.Error(w, fmt.Sprintf("Index %d out of range (0-%d)", idx, len(cfg.ModelList)-1), http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
deletedModelName := cfg.ModelList[idx].ModelName
|
||||||
|
|
||||||
|
cfg.ModelList = append(cfg.ModelList[:idx], cfg.ModelList[idx+1:]...)
|
||||||
|
|
||||||
|
// If the deleted model was the default, clear it.
|
||||||
|
if cfg.Agents.Defaults.ModelName == deletedModelName {
|
||||||
|
cfg.Agents.Defaults.ModelName = ""
|
||||||
|
}
|
||||||
|
if cfg.Agents.Defaults.Model == deletedModelName {
|
||||||
|
cfg.Agents.Defaults.Model = ""
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleSetDefaultModel sets the default model for all agents.
|
||||||
|
//
|
||||||
|
// POST /api/models/default
|
||||||
|
func (h *Handler) handleSetDefaultModel(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 {
|
||||||
|
ModelName string `json:"model_name"`
|
||||||
|
}
|
||||||
|
if err = json.Unmarshal(body, &req); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if req.ModelName == "" {
|
||||||
|
http.Error(w, "model_name is required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify the model_name exists in model_list
|
||||||
|
found := false
|
||||||
|
for _, m := range cfg.ModelList {
|
||||||
|
if m.ModelName == req.ModelName {
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
http.Error(w, fmt.Sprintf("Model %q not found in model_list", req.ModelName), http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg.Agents.Defaults.ModelName = req.ModelName
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{
|
||||||
|
"status": "ok",
|
||||||
|
"default_model": req.ModelName,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// maskAPIKey returns a masked version of an API key for safe display.
|
||||||
|
// Keys longer than 8 chars show prefix + last 4 chars: "sk-****abcd"
|
||||||
|
// Shorter keys are fully masked as "****".
|
||||||
|
// Empty keys return empty string.
|
||||||
|
func maskAPIKey(key string) string {
|
||||||
|
if key == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if len(key) <= 8 {
|
||||||
|
return "****"
|
||||||
|
}
|
||||||
|
// Show first 3 chars and last 4 chars
|
||||||
|
return key[:3] + "****" + key[len(key)-4:]
|
||||||
|
}
|
||||||
844
web/backend/api/oauth.go
Normal file
844
web/backend/api/oauth.go
Normal file
|
|
@ -0,0 +1,844 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"html"
|
||||||
|
"io"
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
oauthProviderOpenAI = "openai"
|
||||||
|
oauthProviderAnthropic = "anthropic"
|
||||||
|
oauthProviderGoogleAntigravity = "google-antigravity"
|
||||||
|
|
||||||
|
oauthMethodBrowser = "browser"
|
||||||
|
oauthMethodDeviceCode = "device_code"
|
||||||
|
oauthMethodToken = "token"
|
||||||
|
|
||||||
|
oauthFlowPending = "pending"
|
||||||
|
oauthFlowSuccess = "success"
|
||||||
|
oauthFlowError = "error"
|
||||||
|
oauthFlowExpired = "expired"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
oauthBrowserFlowTTL = 10 * time.Minute
|
||||||
|
oauthDeviceCodeFlowTTL = 15 * time.Minute
|
||||||
|
oauthTerminalFlowGC = 30 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
var oauthProviderOrder = []string{
|
||||||
|
oauthProviderOpenAI,
|
||||||
|
oauthProviderAnthropic,
|
||||||
|
oauthProviderGoogleAntigravity,
|
||||||
|
}
|
||||||
|
|
||||||
|
var oauthProviderMethods = map[string][]string{
|
||||||
|
oauthProviderOpenAI: {oauthMethodBrowser, oauthMethodDeviceCode, oauthMethodToken},
|
||||||
|
oauthProviderAnthropic: {oauthMethodToken},
|
||||||
|
oauthProviderGoogleAntigravity: {oauthMethodBrowser},
|
||||||
|
}
|
||||||
|
|
||||||
|
var oauthProviderLabels = map[string]string{
|
||||||
|
oauthProviderOpenAI: "OpenAI",
|
||||||
|
oauthProviderAnthropic: "Anthropic",
|
||||||
|
oauthProviderGoogleAntigravity: "Google Antigravity",
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
oauthNow = time.Now
|
||||||
|
oauthGeneratePKCE = auth.GeneratePKCE
|
||||||
|
oauthGenerateState = auth.GenerateState
|
||||||
|
oauthBuildAuthorizeURL = auth.BuildAuthorizeURL
|
||||||
|
oauthRequestDeviceCode = auth.RequestDeviceCode
|
||||||
|
oauthPollDeviceCodeOnce = auth.PollDeviceCodeOnce
|
||||||
|
oauthExchangeCodeForTokens = auth.ExchangeCodeForTokens
|
||||||
|
oauthGetCredential = auth.GetCredential
|
||||||
|
oauthSetCredential = auth.SetCredential
|
||||||
|
oauthDeleteCredential = auth.DeleteCredential
|
||||||
|
oauthLoadConfig = config.LoadConfig
|
||||||
|
oauthSaveConfig = config.SaveConfig
|
||||||
|
oauthFetchAntigravityProject = providers.FetchAntigravityProjectID
|
||||||
|
oauthFetchGoogleUserEmailFunc = fetchGoogleUserEmail
|
||||||
|
)
|
||||||
|
|
||||||
|
type oauthFlow struct {
|
||||||
|
ID string
|
||||||
|
Provider string
|
||||||
|
Method string
|
||||||
|
Status string
|
||||||
|
CreatedAt time.Time
|
||||||
|
UpdatedAt time.Time
|
||||||
|
ExpiresAt time.Time
|
||||||
|
Error string
|
||||||
|
CodeVerifier string
|
||||||
|
OAuthState string
|
||||||
|
RedirectURI string
|
||||||
|
DeviceAuthID string
|
||||||
|
UserCode string
|
||||||
|
VerifyURL string
|
||||||
|
Interval int
|
||||||
|
}
|
||||||
|
|
||||||
|
type oauthProviderStatus struct {
|
||||||
|
Provider string `json:"provider"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
Methods []string `json:"methods"`
|
||||||
|
LoggedIn bool `json:"logged_in"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
AuthMethod string `json:"auth_method,omitempty"`
|
||||||
|
ExpiresAt string `json:"expires_at,omitempty"`
|
||||||
|
AccountID string `json:"account_id,omitempty"`
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
ProjectID string `json:"project_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type oauthFlowResponse struct {
|
||||||
|
FlowID string `json:"flow_id"`
|
||||||
|
Provider string `json:"provider"`
|
||||||
|
Method string `json:"method"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
ExpiresAt string `json:"expires_at,omitempty"`
|
||||||
|
Error string `json:"error,omitempty"`
|
||||||
|
UserCode string `json:"user_code,omitempty"`
|
||||||
|
VerifyURL string `json:"verify_url,omitempty"`
|
||||||
|
Interval int `json:"interval,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// registerOAuthRoutes binds OAuth login/logout endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerOAuthRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/oauth/providers", h.handleListOAuthProviders)
|
||||||
|
mux.HandleFunc("POST /api/oauth/login", h.handleOAuthLogin)
|
||||||
|
mux.HandleFunc("GET /api/oauth/flows/{id}", h.handleGetOAuthFlow)
|
||||||
|
mux.HandleFunc("POST /api/oauth/flows/{id}/poll", h.handlePollOAuthFlow)
|
||||||
|
mux.HandleFunc("POST /api/oauth/logout", h.handleOAuthLogout)
|
||||||
|
mux.HandleFunc("GET /oauth/callback", h.handleOAuthCallback)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleListOAuthProviders(w http.ResponseWriter, r *http.Request) {
|
||||||
|
providersResp := make([]oauthProviderStatus, 0, len(oauthProviderOrder))
|
||||||
|
|
||||||
|
for _, provider := range oauthProviderOrder {
|
||||||
|
cred, err := oauthGetCredential(provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("failed to load credentials: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
item := oauthProviderStatus{
|
||||||
|
Provider: provider,
|
||||||
|
DisplayName: oauthProviderLabels[provider],
|
||||||
|
Methods: oauthProviderMethods[provider],
|
||||||
|
Status: "not_logged_in",
|
||||||
|
}
|
||||||
|
if cred != nil {
|
||||||
|
item.LoggedIn = true
|
||||||
|
item.AuthMethod = cred.AuthMethod
|
||||||
|
item.AccountID = cred.AccountID
|
||||||
|
item.Email = cred.Email
|
||||||
|
item.ProjectID = cred.ProjectID
|
||||||
|
if !cred.ExpiresAt.IsZero() {
|
||||||
|
item.ExpiresAt = cred.ExpiresAt.Format(time.RFC3339)
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case cred.IsExpired():
|
||||||
|
item.Status = "expired"
|
||||||
|
case cred.NeedsRefresh():
|
||||||
|
item.Status = "needs_refresh"
|
||||||
|
default:
|
||||||
|
item.Status = "connected"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
providersResp = append(providersResp, item)
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"providers": providersResp,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleOAuthLogin(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 {
|
||||||
|
Provider string `json:"provider"`
|
||||||
|
Method string `json:"method"`
|
||||||
|
Token string `json:"token"`
|
||||||
|
}
|
||||||
|
if err = json.Unmarshal(body, &req); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, err := normalizeOAuthProvider(req.Provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
method := strings.ToLower(strings.TrimSpace(req.Method))
|
||||||
|
if !isOAuthMethodSupported(provider, method) {
|
||||||
|
http.Error(
|
||||||
|
w,
|
||||||
|
fmt.Sprintf("unsupported login method %q for provider %q", method, provider),
|
||||||
|
http.StatusBadRequest,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch method {
|
||||||
|
case oauthMethodToken:
|
||||||
|
token := strings.TrimSpace(req.Token)
|
||||||
|
if token == "" {
|
||||||
|
http.Error(w, "token is required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cred := &auth.AuthCredential{
|
||||||
|
AccessToken: token,
|
||||||
|
Provider: provider,
|
||||||
|
AuthMethod: oauthMethodToken,
|
||||||
|
}
|
||||||
|
if err := h.persistCredentialAndConfig(provider, oauthMethodToken, cred); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("token login failed: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"provider": provider,
|
||||||
|
"method": method,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
|
||||||
|
case oauthMethodDeviceCode:
|
||||||
|
cfg := auth.OpenAIOAuthConfig()
|
||||||
|
info, err := oauthRequestDeviceCode(cfg)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("failed to request device code: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
now := oauthNow()
|
||||||
|
flow := &oauthFlow{
|
||||||
|
ID: newOAuthFlowID(),
|
||||||
|
Provider: provider,
|
||||||
|
Method: method,
|
||||||
|
Status: oauthFlowPending,
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
ExpiresAt: now.Add(oauthDeviceCodeFlowTTL),
|
||||||
|
DeviceAuthID: info.DeviceAuthID,
|
||||||
|
UserCode: info.UserCode,
|
||||||
|
VerifyURL: info.VerifyURL,
|
||||||
|
Interval: info.Interval,
|
||||||
|
}
|
||||||
|
h.storeOAuthFlow(flow)
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"provider": provider,
|
||||||
|
"method": method,
|
||||||
|
"flow_id": flow.ID,
|
||||||
|
"user_code": flow.UserCode,
|
||||||
|
"verify_url": flow.VerifyURL,
|
||||||
|
"interval": flow.Interval,
|
||||||
|
"expires_at": flow.ExpiresAt.Format(time.RFC3339),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
|
||||||
|
case oauthMethodBrowser:
|
||||||
|
cfg, err := oauthConfigForProvider(provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pkce, err := oauthGeneratePKCE()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("failed to generate PKCE: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
state, err := oauthGenerateState()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("failed to generate state: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
redirectURI := buildOAuthRedirectURI(r)
|
||||||
|
authURL := oauthBuildAuthorizeURL(cfg, pkce, state, redirectURI)
|
||||||
|
|
||||||
|
now := oauthNow()
|
||||||
|
flow := &oauthFlow{
|
||||||
|
ID: newOAuthFlowID(),
|
||||||
|
Provider: provider,
|
||||||
|
Method: method,
|
||||||
|
Status: oauthFlowPending,
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
ExpiresAt: now.Add(oauthBrowserFlowTTL),
|
||||||
|
CodeVerifier: pkce.CodeVerifier,
|
||||||
|
OAuthState: state,
|
||||||
|
RedirectURI: redirectURI,
|
||||||
|
}
|
||||||
|
h.storeOAuthFlow(flow)
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"provider": provider,
|
||||||
|
"method": method,
|
||||||
|
"flow_id": flow.ID,
|
||||||
|
"auth_url": authURL,
|
||||||
|
"expires_at": flow.ExpiresAt.Format(time.RFC3339),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
http.Error(w, "unsupported login method", http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleGetOAuthFlow(w http.ResponseWriter, r *http.Request) {
|
||||||
|
flowID := strings.TrimSpace(r.PathValue("id"))
|
||||||
|
if flowID == "" {
|
||||||
|
http.Error(w, "missing flow id", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
flow, ok := h.getOAuthFlow(flowID)
|
||||||
|
if !ok {
|
||||||
|
http.Error(w, "flow not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(flow))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handlePollOAuthFlow(w http.ResponseWriter, r *http.Request) {
|
||||||
|
flowID := strings.TrimSpace(r.PathValue("id"))
|
||||||
|
if flowID == "" {
|
||||||
|
http.Error(w, "missing flow id", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
flow, ok := h.getOAuthFlow(flowID)
|
||||||
|
if !ok {
|
||||||
|
http.Error(w, "flow not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if flow.Method != oauthMethodDeviceCode {
|
||||||
|
http.Error(w, "flow does not support polling", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if flow.Status != oauthFlowPending {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(flow))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := auth.OpenAIOAuthConfig()
|
||||||
|
cred, err := oauthPollDeviceCodeOnce(cfg, flow.DeviceAuthID, flow.UserCode)
|
||||||
|
if err != nil {
|
||||||
|
if strings.Contains(strings.ToLower(err.Error()), "pending") {
|
||||||
|
updated, _ := h.getOAuthFlow(flowID)
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(updated))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.setOAuthFlowError(flowID, fmt.Sprintf("device code poll failed: %v", err))
|
||||||
|
updated, _ := h.getOAuthFlow(flowID)
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(updated))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if cred == nil {
|
||||||
|
updated, _ := h.getOAuthFlow(flowID)
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(updated))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.persistCredentialAndConfig(flow.Provider, oauthMethodTokenOrOAuth(flow.Method), cred); err != nil {
|
||||||
|
h.setOAuthFlowError(flowID, fmt.Sprintf("failed to save credential: %v", err))
|
||||||
|
updated, _ := h.getOAuthFlow(flowID)
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(updated))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
h.setOAuthFlowSuccess(flowID)
|
||||||
|
updated, _ := h.getOAuthFlow(flowID)
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(flowToResponse(updated))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleOAuthCallback(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state := strings.TrimSpace(r.URL.Query().Get("state"))
|
||||||
|
if state == "" {
|
||||||
|
renderOAuthCallbackPage(w, "", oauthFlowError, "Missing state", "missing_state")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
flow, ok := h.getOAuthFlowByState(state)
|
||||||
|
if !ok {
|
||||||
|
renderOAuthCallbackPage(w, "", oauthFlowError, "OAuth flow not found", "flow_not_found")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if flow.Status != oauthFlowPending {
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, flow.Status, "Flow already completed", flow.Error)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if errMsg := strings.TrimSpace(r.URL.Query().Get("error")); errMsg != "" {
|
||||||
|
if desc := strings.TrimSpace(r.URL.Query().Get("error_description")); desc != "" {
|
||||||
|
errMsg += ": " + desc
|
||||||
|
}
|
||||||
|
h.setOAuthFlowError(flow.ID, errMsg)
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, oauthFlowError, "Authorization failed", errMsg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
code := strings.TrimSpace(r.URL.Query().Get("code"))
|
||||||
|
if code == "" {
|
||||||
|
h.setOAuthFlowError(flow.ID, "missing authorization code")
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, oauthFlowError, "Missing authorization code", "missing_code")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := oauthConfigForProvider(flow.Provider)
|
||||||
|
if err != nil {
|
||||||
|
h.setOAuthFlowError(flow.ID, err.Error())
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, oauthFlowError, "Unsupported provider", err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cred, err := oauthExchangeCodeForTokens(cfg, code, flow.CodeVerifier, flow.RedirectURI)
|
||||||
|
if err != nil {
|
||||||
|
h.setOAuthFlowError(flow.ID, fmt.Sprintf("token exchange failed: %v", err))
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, oauthFlowError, "Token exchange failed", err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.persistCredentialAndConfig(flow.Provider, oauthMethodTokenOrOAuth(flow.Method), cred); err != nil {
|
||||||
|
h.setOAuthFlowError(flow.ID, fmt.Sprintf("failed to save credential: %v", err))
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, oauthFlowError, "Failed to save credential", err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
h.setOAuthFlowSuccess(flow.ID)
|
||||||
|
renderOAuthCallbackPage(w, flow.ID, oauthFlowSuccess, "Authentication successful", "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleOAuthLogout(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 {
|
||||||
|
Provider string `json:"provider"`
|
||||||
|
}
|
||||||
|
if err = json.Unmarshal(body, &req); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, err := normalizeOAuthProvider(req.Provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := oauthDeleteCredential(provider); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("failed to delete credential: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := h.syncProviderAuthMethod(provider, ""); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("failed to update config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"status": "ok",
|
||||||
|
"provider": provider,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func renderOAuthCallbackPage(w http.ResponseWriter, flowID, status, title, errMsg string) {
|
||||||
|
payload := map[string]string{
|
||||||
|
"type": "picoclaw-oauth-result",
|
||||||
|
"flowId": flowID,
|
||||||
|
"status": status,
|
||||||
|
}
|
||||||
|
if errMsg != "" {
|
||||||
|
payload["error"] = errMsg
|
||||||
|
}
|
||||||
|
payloadJSON, _ := json.Marshal(payload)
|
||||||
|
|
||||||
|
message := title
|
||||||
|
if errMsg != "" {
|
||||||
|
message = fmt.Sprintf("%s: %s", title, errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||||
|
if status == oauthFlowSuccess {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
} else {
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _ = fmt.Fprintf(
|
||||||
|
w,
|
||||||
|
"<!doctype html><html><head><meta charset=\"utf-8\"><title>PicoClaw OAuth</title></head><body><script>(function(){var payload=%s;var hasOpener=false;try{if(window.opener&&!window.opener.closed){window.opener.postMessage(payload,window.location.origin);hasOpener=true}}catch(e){}var target='/credentials?oauth_flow_id='+encodeURIComponent(payload.flowId||'')+'&oauth_status='+encodeURIComponent(payload.status||'');setTimeout(function(){if(hasOpener){window.close();return}window.location.replace(target)},800)})();</script><div style=\"font-family:Inter,system-ui,sans-serif;padding:24px\"><h2>%s</h2><p>%s</p><p>You can close this window.</p></div></body></html>",
|
||||||
|
string(payloadJSON),
|
||||||
|
html.EscapeString(title),
|
||||||
|
html.EscapeString(message),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeOAuthProvider(raw string) (string, error) {
|
||||||
|
provider := strings.ToLower(strings.TrimSpace(raw))
|
||||||
|
switch provider {
|
||||||
|
case "antigravity":
|
||||||
|
return oauthProviderGoogleAntigravity, nil
|
||||||
|
case oauthProviderOpenAI, oauthProviderAnthropic, oauthProviderGoogleAntigravity:
|
||||||
|
return provider, nil
|
||||||
|
default:
|
||||||
|
return "", fmt.Errorf("unsupported provider %q", raw)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isOAuthMethodSupported(provider, method string) bool {
|
||||||
|
methods := oauthProviderMethods[provider]
|
||||||
|
for _, m := range methods {
|
||||||
|
if m == method {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func oauthConfigForProvider(provider string) (auth.OAuthProviderConfig, error) {
|
||||||
|
switch provider {
|
||||||
|
case oauthProviderOpenAI:
|
||||||
|
return auth.OpenAIOAuthConfig(), nil
|
||||||
|
case oauthProviderGoogleAntigravity:
|
||||||
|
return auth.GoogleAntigravityOAuthConfig(), nil
|
||||||
|
default:
|
||||||
|
return auth.OAuthProviderConfig{}, fmt.Errorf("provider %q does not support browser oauth", provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func oauthMethodTokenOrOAuth(method string) string {
|
||||||
|
if method == oauthMethodToken {
|
||||||
|
return oauthMethodToken
|
||||||
|
}
|
||||||
|
return "oauth"
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildOAuthRedirectURI(r *http.Request) string {
|
||||||
|
scheme := "http"
|
||||||
|
if r.TLS != nil {
|
||||||
|
scheme = "https"
|
||||||
|
}
|
||||||
|
if forwarded := strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")); forwarded != "" {
|
||||||
|
scheme = strings.Split(forwarded, ",")[0]
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s://%s/oauth/callback", scheme, r.Host)
|
||||||
|
}
|
||||||
|
|
||||||
|
func flowToResponse(flow *oauthFlow) oauthFlowResponse {
|
||||||
|
resp := oauthFlowResponse{
|
||||||
|
FlowID: flow.ID,
|
||||||
|
Provider: flow.Provider,
|
||||||
|
Method: flow.Method,
|
||||||
|
Status: flow.Status,
|
||||||
|
Error: flow.Error,
|
||||||
|
}
|
||||||
|
if !flow.ExpiresAt.IsZero() {
|
||||||
|
resp.ExpiresAt = flow.ExpiresAt.Format(time.RFC3339)
|
||||||
|
}
|
||||||
|
if flow.Method == oauthMethodDeviceCode {
|
||||||
|
resp.UserCode = flow.UserCode
|
||||||
|
resp.VerifyURL = flow.VerifyURL
|
||||||
|
resp.Interval = flow.Interval
|
||||||
|
}
|
||||||
|
return resp
|
||||||
|
}
|
||||||
|
|
||||||
|
func newOAuthFlowID() string {
|
||||||
|
buf := make([]byte, 16)
|
||||||
|
if _, err := rand.Read(buf); err != nil {
|
||||||
|
return fmt.Sprintf("oauth_%d", time.Now().UnixNano())
|
||||||
|
}
|
||||||
|
return hex.EncodeToString(buf)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) storeOAuthFlow(flow *oauthFlow) {
|
||||||
|
now := oauthNow()
|
||||||
|
h.oauthMu.Lock()
|
||||||
|
defer h.oauthMu.Unlock()
|
||||||
|
|
||||||
|
h.gcOAuthFlowsLocked(now)
|
||||||
|
h.oauthFlows[flow.ID] = flow
|
||||||
|
if flow.OAuthState != "" {
|
||||||
|
h.oauthState[flow.OAuthState] = flow.ID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getOAuthFlow(flowID string) (*oauthFlow, bool) {
|
||||||
|
now := oauthNow()
|
||||||
|
h.oauthMu.Lock()
|
||||||
|
defer h.oauthMu.Unlock()
|
||||||
|
|
||||||
|
h.gcOAuthFlowsLocked(now)
|
||||||
|
flow, ok := h.oauthFlows[flowID]
|
||||||
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
cp := *flow
|
||||||
|
return &cp, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getOAuthFlowByState(state string) (*oauthFlow, bool) {
|
||||||
|
now := oauthNow()
|
||||||
|
h.oauthMu.Lock()
|
||||||
|
defer h.oauthMu.Unlock()
|
||||||
|
|
||||||
|
h.gcOAuthFlowsLocked(now)
|
||||||
|
flowID, ok := h.oauthState[state]
|
||||||
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
flow, ok := h.oauthFlows[flowID]
|
||||||
|
if !ok {
|
||||||
|
delete(h.oauthState, state)
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
cp := *flow
|
||||||
|
return &cp, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) setOAuthFlowSuccess(flowID string) {
|
||||||
|
now := oauthNow()
|
||||||
|
h.oauthMu.Lock()
|
||||||
|
defer h.oauthMu.Unlock()
|
||||||
|
|
||||||
|
flow, ok := h.oauthFlows[flowID]
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
flow.Status = oauthFlowSuccess
|
||||||
|
flow.Error = ""
|
||||||
|
flow.UpdatedAt = now
|
||||||
|
if flow.OAuthState != "" {
|
||||||
|
delete(h.oauthState, flow.OAuthState)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) setOAuthFlowError(flowID, errMsg string) {
|
||||||
|
now := oauthNow()
|
||||||
|
h.oauthMu.Lock()
|
||||||
|
defer h.oauthMu.Unlock()
|
||||||
|
|
||||||
|
flow, ok := h.oauthFlows[flowID]
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
flow.Status = oauthFlowError
|
||||||
|
flow.Error = errMsg
|
||||||
|
flow.UpdatedAt = now
|
||||||
|
if flow.OAuthState != "" {
|
||||||
|
delete(h.oauthState, flow.OAuthState)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) gcOAuthFlowsLocked(now time.Time) {
|
||||||
|
for id, flow := range h.oauthFlows {
|
||||||
|
if flow.Status == oauthFlowPending && !flow.ExpiresAt.IsZero() && now.After(flow.ExpiresAt) {
|
||||||
|
flow.Status = oauthFlowExpired
|
||||||
|
flow.Error = "flow expired"
|
||||||
|
flow.UpdatedAt = now
|
||||||
|
if flow.OAuthState != "" {
|
||||||
|
delete(h.oauthState, flow.OAuthState)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if flow.Status != oauthFlowPending && now.Sub(flow.UpdatedAt) > oauthTerminalFlowGC {
|
||||||
|
if flow.OAuthState != "" {
|
||||||
|
delete(h.oauthState, flow.OAuthState)
|
||||||
|
}
|
||||||
|
delete(h.oauthFlows, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) persistCredentialAndConfig(provider, authMethod string, cred *auth.AuthCredential) error {
|
||||||
|
if cred == nil {
|
||||||
|
return fmt.Errorf("empty credential")
|
||||||
|
}
|
||||||
|
|
||||||
|
cp := *cred
|
||||||
|
cp.Provider = provider
|
||||||
|
if cp.AuthMethod == "" {
|
||||||
|
cp.AuthMethod = authMethod
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider == oauthProviderGoogleAntigravity {
|
||||||
|
if cp.Email == "" {
|
||||||
|
email, err := oauthFetchGoogleUserEmailFunc(cp.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("oauth warning: could not fetch google email: %v", err)
|
||||||
|
} else {
|
||||||
|
cp.Email = email
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if cp.ProjectID == "" {
|
||||||
|
projectID, err := oauthFetchAntigravityProject(cp.AccessToken)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("oauth warning: could not fetch antigravity project id: %v", err)
|
||||||
|
} else {
|
||||||
|
cp.ProjectID = projectID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := oauthSetCredential(provider, &cp); err != nil {
|
||||||
|
return fmt.Errorf("saving credential: %w", err)
|
||||||
|
}
|
||||||
|
if err := h.syncProviderAuthMethod(provider, authMethod); err != nil {
|
||||||
|
return fmt.Errorf("syncing provider auth config: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) syncProviderAuthMethod(provider, authMethod string) error {
|
||||||
|
cfg, err := oauthLoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch provider {
|
||||||
|
case oauthProviderOpenAI:
|
||||||
|
cfg.Providers.OpenAI.AuthMethod = authMethod
|
||||||
|
case oauthProviderAnthropic:
|
||||||
|
cfg.Providers.Anthropic.AuthMethod = authMethod
|
||||||
|
case oauthProviderGoogleAntigravity:
|
||||||
|
cfg.Providers.Antigravity.AuthMethod = authMethod
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported provider %q", provider)
|
||||||
|
}
|
||||||
|
|
||||||
|
found := false
|
||||||
|
for i := range cfg.ModelList {
|
||||||
|
if modelBelongsToProvider(provider, cfg.ModelList[i].Model) {
|
||||||
|
cfg.ModelList[i].AuthMethod = authMethod
|
||||||
|
found = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !found && authMethod != "" {
|
||||||
|
cfg.ModelList = append(cfg.ModelList, defaultModelConfigForProvider(provider, authMethod))
|
||||||
|
}
|
||||||
|
|
||||||
|
return oauthSaveConfig(h.configPath, cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func modelBelongsToProvider(provider, model string) bool {
|
||||||
|
lower := strings.ToLower(strings.TrimSpace(model))
|
||||||
|
switch provider {
|
||||||
|
case oauthProviderOpenAI:
|
||||||
|
return lower == "openai" || strings.HasPrefix(lower, "openai/")
|
||||||
|
case oauthProviderAnthropic:
|
||||||
|
return lower == "anthropic" || strings.HasPrefix(lower, "anthropic/")
|
||||||
|
case oauthProviderGoogleAntigravity:
|
||||||
|
return lower == "antigravity" ||
|
||||||
|
lower == "google-antigravity" ||
|
||||||
|
strings.HasPrefix(lower, "antigravity/") ||
|
||||||
|
strings.HasPrefix(lower, "google-antigravity/")
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultModelConfigForProvider(provider, authMethod string) config.ModelConfig {
|
||||||
|
switch provider {
|
||||||
|
case oauthProviderOpenAI:
|
||||||
|
return config.ModelConfig{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: authMethod,
|
||||||
|
}
|
||||||
|
case oauthProviderAnthropic:
|
||||||
|
return config.ModelConfig{
|
||||||
|
ModelName: "claude-sonnet-4.6",
|
||||||
|
Model: "anthropic/claude-sonnet-4.6",
|
||||||
|
AuthMethod: authMethod,
|
||||||
|
}
|
||||||
|
case oauthProviderGoogleAntigravity:
|
||||||
|
return config.ModelConfig{
|
||||||
|
ModelName: "gemini-flash",
|
||||||
|
Model: "antigravity/gemini-3-flash",
|
||||||
|
AuthMethod: authMethod,
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return config.ModelConfig{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func fetchGoogleUserEmail(accessToken string) (string, error) {
|
||||||
|
req, err := http.NewRequest(http.MethodGet, "https://www.googleapis.com/oauth2/v2/userinfo", nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||||
|
|
||||||
|
client := &http.Client{Timeout: 10 * time.Second}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return "", fmt.Errorf("userinfo request failed: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
var userInfo struct {
|
||||||
|
Email string `json:"email"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &userInfo); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if userInfo.Email == "" {
|
||||||
|
return "", fmt.Errorf("empty email in userinfo response")
|
||||||
|
}
|
||||||
|
return userInfo.Email, nil
|
||||||
|
}
|
||||||
293
web/backend/api/oauth_test.go
Normal file
293
web/backend/api/oauth_test.go
Normal file
|
|
@ -0,0 +1,293 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/auth"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestOAuthLoginRejectsUnsupportedMethod(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPost,
|
||||||
|
"/api/oauth/login",
|
||||||
|
strings.NewReader(`{"provider":"anthropic","method":"browser"}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthBrowserFlowCreatedAndQueried(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
|
||||||
|
oauthGeneratePKCE = func() (auth.PKCECodes, error) {
|
||||||
|
return auth.PKCECodes{CodeVerifier: "verifier-1", CodeChallenge: "challenge-1"}, nil
|
||||||
|
}
|
||||||
|
oauthGenerateState = func() (string, error) { return "state-1", nil }
|
||||||
|
oauthBuildAuthorizeURL = func(cfg auth.OAuthProviderConfig, pkce auth.PKCECodes, state, redirectURI string) string {
|
||||||
|
return "https://example.com/authorize?state=" + state
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPost,
|
||||||
|
"/api/oauth/login",
|
||||||
|
strings.NewReader(`{"provider":"openai","method":"browser"}`),
|
||||||
|
)
|
||||||
|
req.Host = "localhost:18800"
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
var loginResp map[string]any
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &loginResp); err != nil {
|
||||||
|
t.Fatalf("unmarshal login response: %v", err)
|
||||||
|
}
|
||||||
|
flowID, _ := loginResp["flow_id"].(string)
|
||||||
|
if flowID == "" {
|
||||||
|
t.Fatalf("flow_id is empty: %v", loginResp)
|
||||||
|
}
|
||||||
|
if loginResp["auth_url"] != "https://example.com/authorize?state=state-1" {
|
||||||
|
t.Fatalf("unexpected auth_url: %v", loginResp["auth_url"])
|
||||||
|
}
|
||||||
|
|
||||||
|
rec2 := httptest.NewRecorder()
|
||||||
|
req2 := httptest.NewRequest(http.MethodGet, "/api/oauth/flows/"+flowID, nil)
|
||||||
|
mux.ServeHTTP(rec2, req2)
|
||||||
|
if rec2.Code != http.StatusOK {
|
||||||
|
t.Fatalf("flow status code = %d, want %d, body=%s", rec2.Code, http.StatusOK, rec2.Body.String())
|
||||||
|
}
|
||||||
|
var flowResp oauthFlowResponse
|
||||||
|
if err := json.Unmarshal(rec2.Body.Bytes(), &flowResp); err != nil {
|
||||||
|
t.Fatalf("unmarshal flow response: %v", err)
|
||||||
|
}
|
||||||
|
if flowResp.Status != oauthFlowPending {
|
||||||
|
t.Fatalf("flow status = %q, want %q", flowResp.Status, oauthFlowPending)
|
||||||
|
}
|
||||||
|
if flowResp.Method != oauthMethodBrowser {
|
||||||
|
t.Fatalf("flow method = %q, want %q", flowResp.Method, oauthMethodBrowser)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthFlowExpiresWhenQueried(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
|
||||||
|
now := time.Date(2026, 3, 6, 12, 0, 0, 0, time.UTC)
|
||||||
|
oauthNow = func() time.Time { return now }
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
h.storeOAuthFlow(&oauthFlow{
|
||||||
|
ID: "expired-flow",
|
||||||
|
Provider: oauthProviderOpenAI,
|
||||||
|
Method: oauthMethodBrowser,
|
||||||
|
Status: oauthFlowPending,
|
||||||
|
CreatedAt: now.Add(-20 * time.Minute),
|
||||||
|
UpdatedAt: now.Add(-20 * time.Minute),
|
||||||
|
ExpiresAt: now.Add(-1 * time.Minute),
|
||||||
|
})
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/oauth/flows/expired-flow", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
var flowResp oauthFlowResponse
|
||||||
|
if err := json.Unmarshal(rec.Body.Bytes(), &flowResp); err != nil {
|
||||||
|
t.Fatalf("unmarshal flow response: %v", err)
|
||||||
|
}
|
||||||
|
if flowResp.Status != oauthFlowExpired {
|
||||||
|
t.Fatalf("flow status = %q, want %q", flowResp.Status, oauthFlowExpired)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthCallbackUnknownState(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/oauth/callback?state=unknown&code=abc", nil)
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
if !strings.Contains(rec.Body.String(), "OAuth flow not found") {
|
||||||
|
t.Fatalf("unexpected body: %s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthLogoutClearsCredentialAndConfig(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
resetOAuthHooks(t)
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig error: %v", err)
|
||||||
|
}
|
||||||
|
cfg.Providers.OpenAI.AuthMethod = "oauth"
|
||||||
|
cfg.ModelList = append(cfg.ModelList, config.ModelConfig{
|
||||||
|
ModelName: "gpt-5.2",
|
||||||
|
Model: "openai/gpt-5.2",
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
})
|
||||||
|
if err = config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig error: %v", err)
|
||||||
|
}
|
||||||
|
if err = auth.SetCredential(oauthProviderOpenAI, &auth.AuthCredential{
|
||||||
|
AccessToken: "token-before-logout",
|
||||||
|
Provider: oauthProviderOpenAI,
|
||||||
|
AuthMethod: "oauth",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("SetCredential error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/oauth/logout", bytes.NewBufferString(`{"provider":"openai"}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
cred, err := auth.GetCredential(oauthProviderOpenAI)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetCredential error: %v", err)
|
||||||
|
}
|
||||||
|
if cred != nil {
|
||||||
|
t.Fatalf("expected credential deleted, got %#v", cred)
|
||||||
|
}
|
||||||
|
|
||||||
|
updated, err := config.LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig error: %v", err)
|
||||||
|
}
|
||||||
|
if updated.Providers.OpenAI.AuthMethod != "" {
|
||||||
|
t.Fatalf("providers.openai.auth_method = %q, want empty", updated.Providers.OpenAI.AuthMethod)
|
||||||
|
}
|
||||||
|
for _, m := range updated.ModelList {
|
||||||
|
if strings.HasPrefix(m.Model, "openai/") && m.AuthMethod != "" {
|
||||||
|
t.Fatalf("openai model auth_method = %q, want empty", m.AuthMethod)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupOAuthTestEnv(t *testing.T) (string, func()) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
tmp := t.TempDir()
|
||||||
|
oldHome := os.Getenv("HOME")
|
||||||
|
oldPicoHome := os.Getenv("PICOCLAW_HOME")
|
||||||
|
|
||||||
|
if err := os.Setenv("HOME", tmp); err != nil {
|
||||||
|
t.Fatalf("set HOME: %v", err)
|
||||||
|
}
|
||||||
|
if err := os.Setenv("PICOCLAW_HOME", filepath.Join(tmp, ".picoclaw")); err != nil {
|
||||||
|
t.Fatalf("set PICOCLAW_HOME: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.ModelList = []config.ModelConfig{{
|
||||||
|
ModelName: "custom-default",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
APIKey: "sk-default",
|
||||||
|
}}
|
||||||
|
cfg.Agents.Defaults.ModelName = "custom-default"
|
||||||
|
|
||||||
|
configPath := filepath.Join(tmp, "config.json")
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cleanup := func() {
|
||||||
|
_ = os.Setenv("HOME", oldHome)
|
||||||
|
if oldPicoHome == "" {
|
||||||
|
_ = os.Unsetenv("PICOCLAW_HOME")
|
||||||
|
} else {
|
||||||
|
_ = os.Setenv("PICOCLAW_HOME", oldPicoHome)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return configPath, cleanup
|
||||||
|
}
|
||||||
|
|
||||||
|
func resetOAuthHooks(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
origNow := oauthNow
|
||||||
|
origGeneratePKCE := oauthGeneratePKCE
|
||||||
|
origGenerateState := oauthGenerateState
|
||||||
|
origBuildAuthorizeURL := oauthBuildAuthorizeURL
|
||||||
|
origRequestDeviceCode := oauthRequestDeviceCode
|
||||||
|
origPollDeviceCodeOnce := oauthPollDeviceCodeOnce
|
||||||
|
origExchangeCodeForTokens := oauthExchangeCodeForTokens
|
||||||
|
origGetCredential := oauthGetCredential
|
||||||
|
origSetCredential := oauthSetCredential
|
||||||
|
origDeleteCredential := oauthDeleteCredential
|
||||||
|
origLoadConfig := oauthLoadConfig
|
||||||
|
origSaveConfig := oauthSaveConfig
|
||||||
|
origFetchProject := oauthFetchAntigravityProject
|
||||||
|
origFetchGoogleEmail := oauthFetchGoogleUserEmailFunc
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
oauthNow = origNow
|
||||||
|
oauthGeneratePKCE = origGeneratePKCE
|
||||||
|
oauthGenerateState = origGenerateState
|
||||||
|
oauthBuildAuthorizeURL = origBuildAuthorizeURL
|
||||||
|
oauthRequestDeviceCode = origRequestDeviceCode
|
||||||
|
oauthPollDeviceCodeOnce = origPollDeviceCodeOnce
|
||||||
|
oauthExchangeCodeForTokens = origExchangeCodeForTokens
|
||||||
|
oauthGetCredential = origGetCredential
|
||||||
|
oauthSetCredential = origSetCredential
|
||||||
|
oauthDeleteCredential = origDeleteCredential
|
||||||
|
oauthLoadConfig = origLoadConfig
|
||||||
|
oauthSaveConfig = origSaveConfig
|
||||||
|
oauthFetchAntigravityProject = origFetchProject
|
||||||
|
oauthFetchGoogleUserEmailFunc = origFetchGoogleEmail
|
||||||
|
})
|
||||||
|
}
|
||||||
161
web/backend/api/pico.go
Normal file
161
web/backend/api/pico.go
Normal file
|
|
@ -0,0 +1,161 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// registerPicoRoutes binds Pico Channel management endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerPicoRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/pico/token", h.handleGetPicoToken)
|
||||||
|
mux.HandleFunc("POST /api/pico/token", h.handleRegenPicoToken)
|
||||||
|
mux.HandleFunc("POST /api/pico/setup", h.handlePicoSetup)
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGetPicoToken returns the current WS token and URL for the frontend.
|
||||||
|
//
|
||||||
|
// GET /api/pico/token
|
||||||
|
func (h *Handler) handleGetPicoToken(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
wsURL := buildWsURL(r, cfg)
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"token": cfg.Channels.Pico.Token,
|
||||||
|
"ws_url": wsURL,
|
||||||
|
"enabled": cfg.Channels.Pico.Enabled,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleRegenPicoToken generates a new Pico WebSocket token and saves it.
|
||||||
|
//
|
||||||
|
// POST /api/pico/token
|
||||||
|
func (h *Handler) handleRegenPicoToken(w http.ResponseWriter, r *http.Request) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
token := generateSecureToken()
|
||||||
|
cfg.Channels.Pico.Token = token
|
||||||
|
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to save config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
wsURL := fmt.Sprintf("ws://%s/pico/ws", net.JoinHostPort(cfg.Gateway.Host, strconv.Itoa(cfg.Gateway.Port)))
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"token": token,
|
||||||
|
"ws_url": wsURL,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensurePicoChannel checks if the Pico Channel is properly configured and
|
||||||
|
// enables it with sensible defaults if not. Returns true if config was changed.
|
||||||
|
func (h *Handler) ensurePicoChannel() (bool, error) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("failed to load config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
changed := false
|
||||||
|
|
||||||
|
if !cfg.Channels.Pico.Enabled {
|
||||||
|
cfg.Channels.Pico.Enabled = true
|
||||||
|
changed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Channels.Pico.Token == "" {
|
||||||
|
cfg.Channels.Pico.Token = generateSecureToken()
|
||||||
|
changed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if !cfg.Channels.Pico.AllowTokenQuery {
|
||||||
|
cfg.Channels.Pico.AllowTokenQuery = true
|
||||||
|
changed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Make sure origins are allowed (frontend might be running on a different port like 5173 during dev)
|
||||||
|
if len(cfg.Channels.Pico.AllowOrigins) == 0 {
|
||||||
|
cfg.Channels.Pico.AllowOrigins = []string{"*"}
|
||||||
|
changed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if changed {
|
||||||
|
if err := config.SaveConfig(h.configPath, cfg); err != nil {
|
||||||
|
return false, fmt.Errorf("failed to save config: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return changed, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// handlePicoSetup automatically configures everything needed for the Pico Channel to work.
|
||||||
|
//
|
||||||
|
// POST /api/pico/setup
|
||||||
|
func (h *Handler) handlePicoSetup(w http.ResponseWriter, r *http.Request) {
|
||||||
|
changed, err := h.ensurePicoChannel()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to load config: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
wsURL := buildWsURL(r, cfg)
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"token": cfg.Channels.Pico.Token,
|
||||||
|
"ws_url": wsURL,
|
||||||
|
"enabled": true,
|
||||||
|
"changed": changed,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildWsURL creates a WebSocket URL for the Pico Channel.
|
||||||
|
// When the gateway host is "0.0.0.0" or empty, it uses the hostname from the
|
||||||
|
// incoming HTTP request so the browser gets a connectable address.
|
||||||
|
func buildWsURL(r *http.Request, cfg *config.Config) string {
|
||||||
|
host := cfg.Gateway.Host
|
||||||
|
if host == "" || host == "0.0.0.0" {
|
||||||
|
// Use the hostname the browser used to reach this backend
|
||||||
|
reqHost, _, err := net.SplitHostPort(r.Host)
|
||||||
|
if err != nil {
|
||||||
|
reqHost = r.Host // r.Host might not have a port
|
||||||
|
}
|
||||||
|
host = reqHost
|
||||||
|
}
|
||||||
|
return "ws://" + net.JoinHostPort(host, strconv.Itoa(cfg.Gateway.Port)) + "/pico/ws"
|
||||||
|
}
|
||||||
|
|
||||||
|
// generateSecureToken creates a random 32-character hex string.
|
||||||
|
func generateSecureToken() string {
|
||||||
|
b := make([]byte, 16)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
// Fallback to something pseudo-random if crypto/rand fails
|
||||||
|
return fmt.Sprintf("pico_%x", time.Now().UnixNano())
|
||||||
|
}
|
||||||
|
return hex.EncodeToString(b)
|
||||||
|
}
|
||||||
66
web/backend/api/router.go
Normal file
66
web/backend/api/router.go
Normal file
|
|
@ -0,0 +1,66 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Handler serves HTTP API requests.
|
||||||
|
type Handler struct {
|
||||||
|
configPath string
|
||||||
|
serverPort int
|
||||||
|
serverPublic bool
|
||||||
|
serverCIDRs []string
|
||||||
|
oauthMu sync.Mutex
|
||||||
|
oauthFlows map[string]*oauthFlow
|
||||||
|
oauthState map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewHandler creates an instance of the API handler.
|
||||||
|
func NewHandler(configPath string) *Handler {
|
||||||
|
return &Handler{
|
||||||
|
configPath: configPath,
|
||||||
|
serverPort: launcherconfig.DefaultPort,
|
||||||
|
oauthFlows: make(map[string]*oauthFlow),
|
||||||
|
oauthState: make(map[string]string),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetServerOptions stores current backend listen options for fallback behavior.
|
||||||
|
func (h *Handler) SetServerOptions(port int, public bool, allowedCIDRs []string) {
|
||||||
|
h.serverPort = port
|
||||||
|
h.serverPublic = public
|
||||||
|
h.serverCIDRs = append([]string(nil), allowedCIDRs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterRoutes binds all API endpoint handlers to the ServeMux.
|
||||||
|
func (h *Handler) RegisterRoutes(mux *http.ServeMux) {
|
||||||
|
// Config CRUD
|
||||||
|
h.registerConfigRoutes(mux)
|
||||||
|
|
||||||
|
// Pico Channel (WebSocket chat)
|
||||||
|
h.registerPicoRoutes(mux)
|
||||||
|
|
||||||
|
// Gateway process lifecycle
|
||||||
|
h.registerGatewayRoutes(mux)
|
||||||
|
|
||||||
|
// Session history
|
||||||
|
h.registerSessionRoutes(mux)
|
||||||
|
|
||||||
|
// OAuth login and credential management
|
||||||
|
h.registerOAuthRoutes(mux)
|
||||||
|
|
||||||
|
// Model list management
|
||||||
|
h.registerModelRoutes(mux)
|
||||||
|
|
||||||
|
// Channel catalog (for frontend navigation/config pages)
|
||||||
|
h.registerChannelRoutes(mux)
|
||||||
|
|
||||||
|
// OS startup / launch-at-login
|
||||||
|
h.registerStartupRoutes(mux)
|
||||||
|
|
||||||
|
// Launcher service parameters (port/public)
|
||||||
|
h.registerLauncherConfigRoutes(mux)
|
||||||
|
}
|
||||||
286
web/backend/api/session.go
Normal file
286
web/backend/api/session.go
Normal file
|
|
@ -0,0 +1,286 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
// registerSessionRoutes binds session list and detail endpoints to the ServeMux.
|
||||||
|
func (h *Handler) registerSessionRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/sessions", h.handleListSessions)
|
||||||
|
mux.HandleFunc("GET /api/sessions/{id}", h.handleGetSession)
|
||||||
|
mux.HandleFunc("DELETE /api/sessions/{id}", h.handleDeleteSession)
|
||||||
|
}
|
||||||
|
|
||||||
|
// sessionFile mirrors the on-disk session JSON structure from pkg/session.
|
||||||
|
type sessionFile struct {
|
||||||
|
Key string `json:"key"`
|
||||||
|
Messages []providers.Message `json:"messages"`
|
||||||
|
Summary string `json:"summary,omitempty"`
|
||||||
|
Created time.Time `json:"created"`
|
||||||
|
Updated time.Time `json:"updated"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// sessionListItem is a lightweight summary returned by GET /api/sessions.
|
||||||
|
type sessionListItem struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Preview string `json:"preview"`
|
||||||
|
MessageCount int `json:"message_count"`
|
||||||
|
Created string `json:"created"`
|
||||||
|
Updated string `json:"updated"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// picoSessionPrefix is the key prefix used by the gateway's routing for Pico
|
||||||
|
// channel sessions. The full key format is:
|
||||||
|
//
|
||||||
|
// agent:main:pico:direct:pico:<session-uuid>
|
||||||
|
//
|
||||||
|
// The sanitized filename replaces ':' with '_', so on disk it becomes:
|
||||||
|
//
|
||||||
|
// agent_main_pico_direct_pico_<session-uuid>.json
|
||||||
|
const picoSessionPrefix = "agent:main:pico:direct:pico:"
|
||||||
|
|
||||||
|
// extractPicoSessionID extracts the session UUID from a full session key.
|
||||||
|
// Returns the UUID and true if the key matches the Pico session pattern.
|
||||||
|
func extractPicoSessionID(key string) (string, bool) {
|
||||||
|
if strings.HasPrefix(key, picoSessionPrefix) {
|
||||||
|
return strings.TrimPrefix(key, picoSessionPrefix), true
|
||||||
|
}
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
// sessionsDir resolves the path to the gateway's session storage directory.
|
||||||
|
// It reads the workspace from config, falling back to ~/.picoclaw/workspace.
|
||||||
|
func (h *Handler) sessionsDir() (string, error) {
|
||||||
|
cfg, err := config.LoadConfig(h.configPath)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
workspace := cfg.Agents.Defaults.Workspace
|
||||||
|
if workspace == "" {
|
||||||
|
home, _ := os.UserHomeDir()
|
||||||
|
workspace = filepath.Join(home, ".picoclaw", "workspace")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expand ~ prefix
|
||||||
|
if len(workspace) > 0 && workspace[0] == '~' {
|
||||||
|
home, _ := os.UserHomeDir()
|
||||||
|
if len(workspace) > 1 && workspace[1] == '/' {
|
||||||
|
workspace = home + workspace[1:]
|
||||||
|
} else {
|
||||||
|
workspace = home
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return filepath.Join(workspace, "sessions"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleListSessions returns a list of Pico session summaries.
|
||||||
|
//
|
||||||
|
// GET /api/sessions
|
||||||
|
func (h *Handler) handleListSessions(w http.ResponseWriter, r *http.Request) {
|
||||||
|
dir, err := h.sessionsDir()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "failed to resolve sessions directory", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
entries, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
// Directory doesn't exist yet = no sessions
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode([]sessionListItem{})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
items := []sessionListItem{}
|
||||||
|
|
||||||
|
for _, entry := range entries {
|
||||||
|
if entry.IsDir() || filepath.Ext(entry.Name()) != ".json" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := os.ReadFile(filepath.Join(dir, entry.Name()))
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var sess sessionFile
|
||||||
|
if err := json.Unmarshal(data, &sess); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only include Pico channel sessions
|
||||||
|
sessionID, ok := extractPicoSessionID(sess.Key)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build a preview from the first user message
|
||||||
|
preview := ""
|
||||||
|
for _, msg := range sess.Messages {
|
||||||
|
if msg.Role == "user" && strings.TrimSpace(msg.Content) != "" {
|
||||||
|
preview = msg.Content
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len([]rune(preview)) > 60 {
|
||||||
|
preview = string([]rune(preview)[:60]) + "..."
|
||||||
|
}
|
||||||
|
if preview == "" {
|
||||||
|
preview = "(empty)"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only count non-empty user and assistant messages
|
||||||
|
validMessageCount := 0
|
||||||
|
for _, msg := range sess.Messages {
|
||||||
|
if (msg.Role == "user" || msg.Role == "assistant") && strings.TrimSpace(msg.Content) != "" {
|
||||||
|
validMessageCount++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
items = append(items, sessionListItem{
|
||||||
|
ID: sessionID,
|
||||||
|
Preview: preview,
|
||||||
|
MessageCount: validMessageCount,
|
||||||
|
Created: sess.Created.Format(time.RFC3339),
|
||||||
|
Updated: sess.Updated.Format(time.RFC3339),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sort by updated descending (most recent first)
|
||||||
|
sort.Slice(items, func(i, j int) bool {
|
||||||
|
return items[i].Updated > items[j].Updated
|
||||||
|
})
|
||||||
|
|
||||||
|
// Pagination parameters
|
||||||
|
offsetStr := r.URL.Query().Get("offset")
|
||||||
|
limitStr := r.URL.Query().Get("limit")
|
||||||
|
|
||||||
|
offset := 0
|
||||||
|
limit := 20 // Default limit
|
||||||
|
|
||||||
|
if val, err := strconv.Atoi(offsetStr); err == nil && val >= 0 {
|
||||||
|
offset = val
|
||||||
|
}
|
||||||
|
if val, err := strconv.Atoi(limitStr); err == nil && val > 0 {
|
||||||
|
limit = val
|
||||||
|
}
|
||||||
|
|
||||||
|
totalItems := len(items)
|
||||||
|
|
||||||
|
end := offset + limit
|
||||||
|
if offset >= totalItems {
|
||||||
|
items = []sessionListItem{} // Out of bounds, return empty
|
||||||
|
} else {
|
||||||
|
if end > totalItems {
|
||||||
|
end = totalItems
|
||||||
|
}
|
||||||
|
items = items[offset:end]
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(items)
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleGetSession returns the full message history for a specific session.
|
||||||
|
//
|
||||||
|
// GET /api/sessions/{id}
|
||||||
|
func (h *Handler) handleGetSession(w http.ResponseWriter, r *http.Request) {
|
||||||
|
sessionID := r.PathValue("id")
|
||||||
|
if sessionID == "" {
|
||||||
|
http.Error(w, "missing session id", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
dir, err := h.sessionsDir()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "failed to resolve sessions directory", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// The sanitized filename replaces ':' with '_':
|
||||||
|
// agent:main:pico:direct:pico:<uuid> -> agent_main_pico_direct_pico_<uuid>.json
|
||||||
|
filename := strings.ReplaceAll(picoSessionPrefix+sessionID, ":", "_") + ".json"
|
||||||
|
|
||||||
|
data, err := os.ReadFile(filepath.Join(dir, filename))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "session not found", http.StatusNotFound)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var sess sessionFile
|
||||||
|
if err := json.Unmarshal(data, &sess); err != nil {
|
||||||
|
http.Error(w, "failed to parse session", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to a simpler format for the frontend
|
||||||
|
type chatMessage struct {
|
||||||
|
Role string `json:"role"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
}
|
||||||
|
|
||||||
|
messages := make([]chatMessage, 0, len(sess.Messages))
|
||||||
|
for _, msg := range sess.Messages {
|
||||||
|
// Only include user and assistant messages that have actual content
|
||||||
|
if (msg.Role == "user" || msg.Role == "assistant") && strings.TrimSpace(msg.Content) != "" {
|
||||||
|
messages = append(messages, chatMessage{
|
||||||
|
Role: msg.Role,
|
||||||
|
Content: msg.Content,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"id": sessionID,
|
||||||
|
"messages": messages,
|
||||||
|
"summary": sess.Summary,
|
||||||
|
"created": sess.Created.Format(time.RFC3339),
|
||||||
|
"updated": sess.Updated.Format(time.RFC3339),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleDeleteSession deletes a specific session.
|
||||||
|
//
|
||||||
|
// DELETE /api/sessions/{id}
|
||||||
|
func (h *Handler) handleDeleteSession(w http.ResponseWriter, r *http.Request) {
|
||||||
|
sessionID := r.PathValue("id")
|
||||||
|
if sessionID == "" {
|
||||||
|
http.Error(w, "missing session id", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
dir, err := h.sessionsDir()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "failed to resolve sessions directory", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// The sanitized filename replaces ':' with '_':
|
||||||
|
// agent:main:pico:direct:pico:<uuid> -> agent_main_pico_direct_pico_<uuid>.json
|
||||||
|
filename := strings.ReplaceAll(picoSessionPrefix+sessionID, ":", "_") + ".json"
|
||||||
|
filePath := filepath.Join(dir, filename)
|
||||||
|
|
||||||
|
if err := os.Remove(filePath); err != nil {
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
http.Error(w, "session not found", http.StatusNotFound)
|
||||||
|
} else {
|
||||||
|
http.Error(w, "failed to delete session", http.StatusInternalServerError)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
}
|
||||||
305
web/backend/api/startup.go
Normal file
305
web/backend/api/startup.go
Normal file
|
|
@ -0,0 +1,305 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
autoStartEntryName = "PicoClawLauncher"
|
||||||
|
launchAgentLabel = "io.picoclaw.launcher"
|
||||||
|
)
|
||||||
|
|
||||||
|
type autoStartRequest struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type autoStartResponse struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
Supported bool `json:"supported"`
|
||||||
|
Platform string `json:"platform"`
|
||||||
|
Message string `json:"message,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
var errAutoStartUnsupported = errors.New("autostart is not supported on this platform")
|
||||||
|
|
||||||
|
func (h *Handler) registerStartupRoutes(mux *http.ServeMux) {
|
||||||
|
mux.HandleFunc("GET /api/system/autostart", h.handleGetAutoStart)
|
||||||
|
mux.HandleFunc("PUT /api/system/autostart", h.handleSetAutoStart)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleGetAutoStart(w http.ResponseWriter, r *http.Request) {
|
||||||
|
enabled, supported, message, err := h.getAutoStartStatus()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to read startup setting: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(autoStartResponse{
|
||||||
|
Enabled: enabled,
|
||||||
|
Supported: supported,
|
||||||
|
Platform: runtime.GOOS,
|
||||||
|
Message: message,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleSetAutoStart(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req autoStartRequest
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Invalid JSON: %v", err), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.setAutoStart(req.Enabled); err != nil {
|
||||||
|
if errors.Is(err, errAutoStartUnsupported) {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to update startup setting: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
enabled, supported, message, err := h.getAutoStartStatus()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("Failed to verify startup setting: %v", err), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(autoStartResponse{
|
||||||
|
Enabled: enabled,
|
||||||
|
Supported: supported,
|
||||||
|
Platform: runtime.GOOS,
|
||||||
|
Message: message,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) resolveLaunchCommand() (string, []string, error) {
|
||||||
|
exePath, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
args := []string{"-no-browser"}
|
||||||
|
if h.configPath != "" {
|
||||||
|
args = append(args, h.configPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
return exePath, args, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getAutoStartStatus() (enabled bool, supported bool, message string, err error) {
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "darwin":
|
||||||
|
exists, err := fileExists(macLaunchAgentPath())
|
||||||
|
return exists, true, "Changes apply on next login.", err
|
||||||
|
case "linux":
|
||||||
|
exists, err := fileExists(linuxAutoStartPath())
|
||||||
|
return exists, true, "Changes apply on next login.", err
|
||||||
|
case "windows":
|
||||||
|
exists, err := windowsRunKeyExists()
|
||||||
|
return exists, true, "Changes apply on next login.", err
|
||||||
|
default:
|
||||||
|
return false, false, "Current platform does not support launch at login.", nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) setAutoStart(enabled bool) error {
|
||||||
|
exePath, args, err := h.resolveLaunchCommand()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "darwin":
|
||||||
|
return setDarwinAutoStart(enabled, exePath, args)
|
||||||
|
case "linux":
|
||||||
|
return setLinuxAutoStart(enabled, exePath, args)
|
||||||
|
case "windows":
|
||||||
|
return setWindowsAutoStart(enabled, exePath, args)
|
||||||
|
default:
|
||||||
|
return errAutoStartUnsupported
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func fileExists(path string) (bool, error) {
|
||||||
|
_, err := os.Stat(path)
|
||||||
|
if err == nil {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func macLaunchAgentPath() string {
|
||||||
|
home, _ := os.UserHomeDir()
|
||||||
|
return filepath.Join(home, "Library", "LaunchAgents", launchAgentLabel+".plist")
|
||||||
|
}
|
||||||
|
|
||||||
|
func setDarwinAutoStart(enabled bool, exePath string, args []string) error {
|
||||||
|
plistPath := macLaunchAgentPath()
|
||||||
|
if enabled {
|
||||||
|
if err := os.MkdirAll(filepath.Dir(plistPath), 0o755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
content := buildDarwinPlist(exePath, args)
|
||||||
|
return os.WriteFile(plistPath, []byte(content), 0o644)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.Remove(plistPath); err != nil && !os.IsNotExist(err) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func xmlEscape(s string) string {
|
||||||
|
var b bytes.Buffer
|
||||||
|
for _, r := range s {
|
||||||
|
switch r {
|
||||||
|
case '&':
|
||||||
|
b.WriteString("&")
|
||||||
|
case '<':
|
||||||
|
b.WriteString("<")
|
||||||
|
case '>':
|
||||||
|
b.WriteString(">")
|
||||||
|
case '"':
|
||||||
|
b.WriteString(""")
|
||||||
|
case '\'':
|
||||||
|
b.WriteString("'")
|
||||||
|
default:
|
||||||
|
b.WriteRune(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildDarwinPlist(exePath string, args []string) string {
|
||||||
|
programArgs := make([]string, 0, len(args)+1)
|
||||||
|
programArgs = append(programArgs, exePath)
|
||||||
|
programArgs = append(programArgs, args...)
|
||||||
|
|
||||||
|
var b strings.Builder
|
||||||
|
b.WriteString(`<?xml version="1.0" encoding="UTF-8"?>` + "\n")
|
||||||
|
b.WriteString(
|
||||||
|
`<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">` + "\n",
|
||||||
|
)
|
||||||
|
b.WriteString(`<plist version="1.0">` + "\n")
|
||||||
|
b.WriteString(`<dict>` + "\n")
|
||||||
|
b.WriteString(` <key>Label</key>` + "\n")
|
||||||
|
b.WriteString(` <string>` + launchAgentLabel + `</string>` + "\n")
|
||||||
|
b.WriteString(` <key>ProgramArguments</key>` + "\n")
|
||||||
|
b.WriteString(` <array>` + "\n")
|
||||||
|
for _, arg := range programArgs {
|
||||||
|
b.WriteString(` <string>` + xmlEscape(arg) + `</string>` + "\n")
|
||||||
|
}
|
||||||
|
b.WriteString(` </array>` + "\n")
|
||||||
|
b.WriteString(` <key>RunAtLoad</key>` + "\n")
|
||||||
|
b.WriteString(` <true/>` + "\n")
|
||||||
|
b.WriteString(` <key>ProcessType</key>` + "\n")
|
||||||
|
b.WriteString(` <string>Background</string>` + "\n")
|
||||||
|
b.WriteString(`</dict>` + "\n")
|
||||||
|
b.WriteString(`</plist>` + "\n")
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func linuxAutoStartPath() string {
|
||||||
|
home, _ := os.UserHomeDir()
|
||||||
|
return filepath.Join(home, ".config", "autostart", "picoclaw-web.desktop")
|
||||||
|
}
|
||||||
|
|
||||||
|
func shellQuote(s string) string {
|
||||||
|
if s == "" {
|
||||||
|
return "''"
|
||||||
|
}
|
||||||
|
if !strings.ContainsAny(s, " \t\n'\"\\$`") {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return "'" + strings.ReplaceAll(s, "'", "'\"'\"'") + "'"
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildLinuxExecLine(exePath string, args []string) string {
|
||||||
|
parts := make([]string, 0, len(args)+1)
|
||||||
|
parts = append(parts, shellQuote(exePath))
|
||||||
|
for _, arg := range args {
|
||||||
|
parts = append(parts, shellQuote(arg))
|
||||||
|
}
|
||||||
|
return strings.Join(parts, " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func setLinuxAutoStart(enabled bool, exePath string, args []string) error {
|
||||||
|
desktopPath := linuxAutoStartPath()
|
||||||
|
if enabled {
|
||||||
|
if err := os.MkdirAll(filepath.Dir(desktopPath), 0o755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
content := strings.Join([]string{
|
||||||
|
"[Desktop Entry]",
|
||||||
|
"Type=Application",
|
||||||
|
"Version=1.0",
|
||||||
|
"Name=PicoClaw Web",
|
||||||
|
"Comment=Start PicoClaw Web on login",
|
||||||
|
"Exec=" + buildLinuxExecLine(exePath, args),
|
||||||
|
"Terminal=false",
|
||||||
|
"X-GNOME-Autostart-enabled=true",
|
||||||
|
"NoDisplay=true",
|
||||||
|
"",
|
||||||
|
}, "\n")
|
||||||
|
return os.WriteFile(desktopPath, []byte(content), 0o644)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.Remove(desktopPath); err != nil && !os.IsNotExist(err) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func windowsCommandLine(exePath string, args []string) string {
|
||||||
|
parts := make([]string, 0, len(args)+1)
|
||||||
|
parts = append(parts, fmt.Sprintf("%q", exePath))
|
||||||
|
for _, arg := range args {
|
||||||
|
parts = append(parts, fmt.Sprintf("%q", arg))
|
||||||
|
}
|
||||||
|
return strings.Join(parts, " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func windowsRunKeyExists() (bool, error) {
|
||||||
|
cmd := exec.Command("reg", "query", `HKCU\Software\Microsoft\Windows\CurrentVersion\Run`, "/v", autoStartEntryName)
|
||||||
|
if err := cmd.Run(); err != nil {
|
||||||
|
var exitErr *exec.ExitError
|
||||||
|
if errors.As(err, &exitErr) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func setWindowsAutoStart(enabled bool, exePath string, args []string) error {
|
||||||
|
key := `HKCU\Software\Microsoft\Windows\CurrentVersion\Run`
|
||||||
|
if enabled {
|
||||||
|
commandLine := windowsCommandLine(exePath, args)
|
||||||
|
cmd := exec.Command("reg", "add", key, "/v", autoStartEntryName, "/t", "REG_SZ", "/d", commandLine, "/f")
|
||||||
|
return cmd.Run()
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd := exec.Command("reg", "delete", key, "/v", autoStartEntryName, "/f")
|
||||||
|
if err := cmd.Run(); err != nil {
|
||||||
|
var exitErr *exec.ExitError
|
||||||
|
if errors.As(err, &exitErr) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
56
web/backend/api/startup_test.go
Normal file
56
web/backend/api/startup_test.go
Normal file
|
|
@ -0,0 +1,56 @@
|
||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/web/backend/launcherconfig"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestResolveLaunchCommandUsesConfigFileDefaults(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
|
||||||
|
// Persist non-default launcher options to ensure resolveLaunchCommand does not
|
||||||
|
// pin them into autostart args.
|
||||||
|
launcherPath := launcherconfig.PathForAppConfig(configPath)
|
||||||
|
if err := launcherconfig.Save(launcherPath, launcherconfig.Config{
|
||||||
|
Port: 19999,
|
||||||
|
Public: true,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("launcherconfig.Save() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
exePath, args, err := h.resolveLaunchCommand()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolveLaunchCommand() error = %v", err)
|
||||||
|
}
|
||||||
|
if exePath == "" {
|
||||||
|
t.Fatal("resolveLaunchCommand() returned empty executable path")
|
||||||
|
}
|
||||||
|
if len(args) != 2 {
|
||||||
|
t.Fatalf("args len = %d, want 2 (got %v)", len(args), args)
|
||||||
|
}
|
||||||
|
if args[0] != "-no-browser" {
|
||||||
|
t.Fatalf("args[0] = %q, want %q", args[0], "-no-browser")
|
||||||
|
}
|
||||||
|
if args[1] != configPath {
|
||||||
|
t.Fatalf("args[1] = %q, want %q", args[1], configPath)
|
||||||
|
}
|
||||||
|
for _, arg := range args {
|
||||||
|
if arg == "-port" || arg == "-public" {
|
||||||
|
t.Fatalf("autostart args should not pin network flags, got %v", args)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildDarwinPlistIncludesRunAtLoad(t *testing.T) {
|
||||||
|
plist := buildDarwinPlist("/tmp/picoclaw-web", []string{"-no-browser", "/tmp/config.json"})
|
||||||
|
if !strings.Contains(plist, "<key>RunAtLoad</key>") {
|
||||||
|
t.Fatalf("plist missing RunAtLoad key:\n%s", plist)
|
||||||
|
}
|
||||||
|
if !strings.Contains(plist, "<true/>") {
|
||||||
|
t.Fatalf("plist missing RunAtLoad true value:\n%s", plist)
|
||||||
|
}
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue