Files
multica/server/pkg/agent/claude_test.go
Bohan Jiang ce447c7f06 feat(agent): add custom CLI arguments support (#986)
* feat(agent): add custom CLI arguments support

Allow users to configure custom CLI arguments per agent that get
appended to the agent subprocess command at launch time. This enables
use cases like specifying different models (--model o3), max turns,
or other provider-specific flags without needing separate runtimes.

Changes:
- Add custom_args JSONB column to agent table (migration 041)
- Update API handler to accept/return custom_args in create/update
- Pass custom_args through claim endpoint to daemon
- Append custom_args to CLI commands for all agent backends
- Add ExecOptions.CustomArgs field in agent package
- Add Custom Args tab in agent detail UI
- Add --custom-args flag to CLI agent create/update commands

Closes MUL-802

* fix(agent): filter protocol-critical flags from custom_args

Add per-backend filtering of custom_args to prevent users from
accidentally overriding flags that the daemon hardcodes for its
communication protocol (e.g. --output-format, --input-format,
--permission-mode for Claude).

This follows the same pattern as custom_env's isBlockedEnvKey: we
only block the small, stable set of flags that would break the
daemon↔agent protocol — not every possible dangerous flag. Workspace
members are trusted for everything else.

Each backend defines its own blocked set:
- Claude: -p, --output-format, --input-format, --permission-mode
- Gemini: -p, --yolo, -o
- Codex: --listen
- OpenCode: --format
- OpenClaw: --local, --json, --session-id, --message
- Hermes: none (ACP is positional)

Includes unit tests for the filtering logic.

* fix(agent): address code review nits for custom_args

- Replace module-level `nextArgId` counter with `crypto.randomUUID()`
  in custom-args-tab.tsx to avoid SSR ID conflicts
- Add unit tests for custom args passthrough and blocked-arg filtering
  in both Claude and Gemini arg builders
2026-04-15 14:58:53 +08:00

420 lines
10 KiB
Go

package agent
import (
"bytes"
"encoding/json"
"log/slog"
"strings"
"testing"
)
func TestClaudeHandleAssistantText(t *testing.T) {
t.Parallel()
b := &claudeBackend{cfg: Config{Logger: slog.Default()}}
ch := make(chan Message, 10)
var output strings.Builder
msg := claudeSDKMessage{
Type: "assistant",
Message: mustMarshal(t, claudeMessageContent{
Role: "assistant",
Content: []claudeContentBlock{
{Type: "text", Text: "Hello world"},
},
}),
}
b.handleAssistant(msg, ch, &output, make(map[string]TokenUsage))
if output.String() != "Hello world" {
t.Fatalf("expected output 'Hello world', got %q", output.String())
}
select {
case m := <-ch:
if m.Type != MessageText || m.Content != "Hello world" {
t.Fatalf("unexpected message: %+v", m)
}
default:
t.Fatal("expected message on channel")
}
}
func TestClaudeHandleAssistantToolUse(t *testing.T) {
t.Parallel()
b := &claudeBackend{cfg: Config{Logger: slog.Default()}}
ch := make(chan Message, 10)
var output strings.Builder
msg := claudeSDKMessage{
Type: "assistant",
Message: mustMarshal(t, claudeMessageContent{
Role: "assistant",
Content: []claudeContentBlock{
{
Type: "tool_use",
ID: "call-1",
Name: "Read",
Input: mustMarshal(t, map[string]any{"path": "/tmp/foo"}),
},
},
}),
}
b.handleAssistant(msg, ch, &output, make(map[string]TokenUsage))
if output.String() != "" {
t.Fatalf("tool_use should not add to output, got %q", output.String())
}
select {
case m := <-ch:
if m.Type != MessageToolUse || m.Tool != "Read" || m.CallID != "call-1" {
t.Fatalf("unexpected message: %+v", m)
}
if m.Input["path"] != "/tmp/foo" {
t.Fatalf("expected input path /tmp/foo, got %v", m.Input["path"])
}
default:
t.Fatal("expected message on channel")
}
}
func TestClaudeHandleUserToolResult(t *testing.T) {
t.Parallel()
b := &claudeBackend{cfg: Config{Logger: slog.Default()}}
ch := make(chan Message, 10)
msg := claudeSDKMessage{
Type: "user",
Message: mustMarshal(t, claudeMessageContent{
Role: "user",
Content: []claudeContentBlock{
{
Type: "tool_result",
ToolUseID: "call-1",
Content: mustMarshal(t, "file contents here"),
},
},
}),
}
b.handleUser(msg, ch)
select {
case m := <-ch:
if m.Type != MessageToolResult || m.CallID != "call-1" {
t.Fatalf("unexpected message: %+v", m)
}
default:
t.Fatal("expected message on channel")
}
}
func TestClaudeHandleControlRequestAutoApproves(t *testing.T) {
t.Parallel()
b := &claudeBackend{cfg: Config{Logger: slog.Default()}}
var written bytes.Buffer
msg := claudeSDKMessage{
Type: "control_request",
RequestID: "req-42",
Request: mustMarshal(t, claudeControlRequestPayload{
Subtype: "tool_use",
ToolName: "Bash",
Input: mustMarshal(t, map[string]any{"command": "ls"}),
}),
}
b.handleControlRequest(msg, &written)
var resp map[string]any
if err := json.Unmarshal(bytes.TrimSpace(written.Bytes()), &resp); err != nil {
t.Fatalf("unmarshal response: %v", err)
}
if resp["type"] != "control_response" {
t.Fatalf("expected type control_response, got %v", resp["type"])
}
respInner := resp["response"].(map[string]any)
if respInner["request_id"] != "req-42" {
t.Fatalf("expected request_id req-42, got %v", respInner["request_id"])
}
innerResp := respInner["response"].(map[string]any)
if innerResp["behavior"] != "allow" {
t.Fatalf("expected behavior allow, got %v", innerResp["behavior"])
}
}
func TestClaudeHandleAssistantInvalidJSON(t *testing.T) {
t.Parallel()
b := &claudeBackend{cfg: Config{Logger: slog.Default()}}
ch := make(chan Message, 10)
var output strings.Builder
msg := claudeSDKMessage{
Type: "assistant",
Message: json.RawMessage(`invalid json`),
}
// Should not panic
b.handleAssistant(msg, ch, &output, make(map[string]TokenUsage))
if output.String() != "" {
t.Fatalf("expected empty output for invalid JSON, got %q", output.String())
}
select {
case m := <-ch:
t.Fatalf("expected no message, got %+v", m)
default:
}
}
func TestTrySendDropsWhenFull(t *testing.T) {
t.Parallel()
ch := make(chan Message, 1)
// Fill the channel
trySend(ch, Message{Type: MessageText, Content: "first"})
// This should not block
trySend(ch, Message{Type: MessageText, Content: "second"})
m := <-ch
if m.Content != "first" {
t.Fatalf("expected 'first', got %q", m.Content)
}
select {
case m := <-ch:
t.Fatalf("expected empty channel, got %+v", m)
default:
}
}
func TestBuildClaudeArgsIncludesStrictMCPConfig(t *testing.T) {
t.Parallel()
args := buildClaudeArgs(ExecOptions{}, slog.Default())
expected := []string{
"-p",
"--output-format", "stream-json",
"--input-format", "stream-json",
"--verbose",
"--strict-mcp-config",
"--permission-mode", "bypassPermissions",
}
if len(args) != len(expected) {
t.Fatalf("expected %d args, got %d: %v", len(expected), len(args), args)
}
for i, want := range expected {
if args[i] != want {
t.Fatalf("expected args[%d] = %q, got %q", i, want, args[i])
}
}
}
func TestFilterCustomArgsBlocksProtocolFlags(t *testing.T) {
t.Parallel()
blocked := map[string]blockedArgMode{
"--output-format": blockedWithValue,
"--permission-mode": blockedWithValue,
"-p": blockedStandalone,
}
logger := slog.Default()
// Blocks flag with separate value
result := filterCustomArgs([]string{"--output-format", "text", "--model", "o3"}, blocked, logger)
if len(result) != 2 || result[0] != "--model" || result[1] != "o3" {
t.Fatalf("expected [--model o3], got %v", result)
}
// Blocks flag=value form
result = filterCustomArgs([]string{"--permission-mode=plan", "--verbose"}, blocked, logger)
if len(result) != 1 || result[0] != "--verbose" {
t.Fatalf("expected [--verbose], got %v", result)
}
// Blocks standalone short flags without consuming next arg
result = filterCustomArgs([]string{"-p", "--max-turns", "10"}, blocked, logger)
if len(result) != 2 || result[0] != "--max-turns" || result[1] != "10" {
t.Fatalf("expected [--max-turns 10], got %v", result)
}
// Passes through non-blocked args
result = filterCustomArgs([]string{"--model", "o3", "--max-turns", "50"}, blocked, logger)
if len(result) != 4 {
t.Fatalf("expected all 4 args to pass through, got %v", result)
}
// Handles nil blocked map
result = filterCustomArgs([]string{"--anything"}, nil, logger)
if len(result) != 1 {
t.Fatalf("expected args to pass through with nil blocked map, got %v", result)
}
// Handles empty args
result = filterCustomArgs(nil, blocked, logger)
if result != nil {
t.Fatalf("expected nil for nil input, got %v", result)
}
}
func TestBuildClaudeArgsPassesThroughCustomArgs(t *testing.T) {
t.Parallel()
args := buildClaudeArgs(ExecOptions{
CustomArgs: []string{"--max-turns", "50", "--verbose"},
}, slog.Default())
// Custom args should appear at the end
found := 0
for i, a := range args {
if a == "--max-turns" && i+1 < len(args) && args[i+1] == "50" {
found++
}
}
if found != 1 {
t.Fatalf("expected --max-turns 50 in args: %v", args)
}
}
func TestBuildClaudeArgsFiltersBlockedCustomArgs(t *testing.T) {
t.Parallel()
args := buildClaudeArgs(ExecOptions{
CustomArgs: []string{"--output-format", "text", "--model", "o3"},
}, slog.Default())
// --output-format text should be stripped
for _, a := range args[len(args)-2:] {
if a == "text" {
// "text" should not be in the last args since --output-format was blocked
// The actual --output-format stream-json is earlier in the list
}
}
// --model o3 should pass through
foundModel := false
for i, a := range args {
if a == "--model" && i+1 < len(args) && args[i+1] == "o3" {
foundModel = true
}
// Verify no duplicate --output-format with value "text"
if a == "--output-format" && i+1 < len(args) && args[i+1] == "text" {
t.Fatalf("blocked --output-format text should have been filtered: %v", args)
}
}
if !foundModel {
t.Fatalf("expected --model o3 in args but it was missing: %v", args)
}
}
func TestBuildClaudeInputEncodesUserMessage(t *testing.T) {
t.Parallel()
data, err := buildClaudeInput("say pong")
if err != nil {
t.Fatalf("buildClaudeInput: %v", err)
}
if len(data) == 0 || data[len(data)-1] != '\n' {
t.Fatalf("expected newline-terminated payload, got %q", data)
}
var payload map[string]any
if err := json.Unmarshal(bytes.TrimSpace(data), &payload); err != nil {
t.Fatalf("unmarshal payload: %v", err)
}
if payload["type"] != "user" {
t.Fatalf("expected type user, got %v", payload["type"])
}
message, ok := payload["message"].(map[string]any)
if !ok {
t.Fatalf("expected message object, got %T", payload["message"])
}
if message["role"] != "user" {
t.Fatalf("expected role user, got %v", message["role"])
}
content, ok := message["content"].([]any)
if !ok || len(content) != 1 {
t.Fatalf("expected one content block, got %v", message["content"])
}
block, ok := content[0].(map[string]any)
if !ok {
t.Fatalf("expected content block object, got %T", content[0])
}
if block["type"] != "text" || block["text"] != "say pong" {
t.Fatalf("unexpected content block: %v", block)
}
}
func TestMergeEnvFiltersClaudeCodeVars(t *testing.T) {
t.Parallel()
env := mergeEnv([]string{
"PATH=/usr/bin",
"CLAUDECODE=1",
"CLAUDE_CODE_ENTRYPOINT=cli",
"CLAUDECODEX=keep-me",
}, map[string]string{"FOO": "bar"})
for _, entry := range env {
if entry == "CLAUDECODE=1" || entry == "CLAUDE_CODE_ENTRYPOINT=cli" {
t.Fatalf("expected CLAUDECODE vars to be filtered, got %v", env)
}
}
found := map[string]bool{}
for _, entry := range env {
found[entry] = true
}
if !found["PATH=/usr/bin"] {
t.Fatalf("expected PATH to be preserved, got %v", env)
}
if !found["CLAUDECODEX=keep-me"] {
t.Fatalf("expected unrelated env vars to be preserved, got %v", env)
}
if !found["FOO=bar"] {
t.Fatalf("expected extra env var to be appended, got %v", env)
}
}
func TestBuildEnvAppendsExtras(t *testing.T) {
t.Parallel()
env := buildEnv(map[string]string{"FOO": "bar", "BAZ": "qux"})
found := 0
for _, e := range env {
if e == "FOO=bar" || e == "BAZ=qux" {
found++
}
}
if found != 2 {
t.Fatalf("expected 2 extra env vars, found %d", found)
}
}
func TestBuildEnvNilExtras(t *testing.T) {
t.Parallel()
env := buildEnv(nil)
if len(env) == 0 {
t.Fatal("expected at least system env vars")
}
}
func mustMarshal(t *testing.T, v any) json.RawMessage {
t.Helper()
data, err := json.Marshal(v)
if err != nil {
t.Fatalf("json.Marshal: %v", err)
}
return data
}