%s", escaped))
@@ -92,6 +108,11 @@ type codeBlockMatch struct {
codes []string
}
+type rawURLMatch struct {
+ text string
+ urls []string
+}
+
func extractCodeBlocks(text string) codeBlockMatch {
matches := reCodeBlock.FindAllStringSubmatch(text, -1)
@@ -110,6 +131,24 @@ func extractCodeBlocks(text string) codeBlockMatch {
return codeBlockMatch{text: text, codes: codes}
}
+func extractRawURLs(text string) rawURLMatch {
+ matches := reRawURL.FindAllString(text, -1)
+
+ urls := make([]string, 0, len(matches))
+ for _, match := range matches {
+ urls = append(urls, match)
+ }
+
+ i := 0
+ text = reRawURL.ReplaceAllStringFunc(text, func(string) string {
+ placeholder := fmt.Sprintf("\x00RU%d\x00", i)
+ i++
+ return placeholder
+ })
+
+ return rawURLMatch{text: text, urls: urls}
+}
+
type inlineCodeMatch struct {
text string
codes []string
@@ -139,3 +178,7 @@ func escapeHTML(text string) string {
text = strings.ReplaceAll(text, ">", ">")
return text
}
+
+func escapeHTMLAttr(text string) string {
+ return html.EscapeString(text)
+}
diff --git a/pkg/channels/telegram/parser_markdown_to_html_test.go b/pkg/channels/telegram/parser_markdown_to_html_test.go
index 7754ee076..a54a1c2c7 100644
--- a/pkg/channels/telegram/parser_markdown_to_html_test.go
+++ b/pkg/channels/telegram/parser_markdown_to_html_test.go
@@ -32,6 +32,11 @@ func Test_markdownToTelegramHTML(t *testing.T) {
input: "[click here](https://example.com/path)",
expected: `click here`,
},
+ {
+ name: "raw oauth url with underscores survives",
+ input: "Apri https://accounts.google.com/o/oauth2/auth?response_type=code&client_id=test-client&redirect_uri=http%3A%2F%2Flocalhost%3A8001%2Foauth2callback&code_challenge=abc_def&code_challenge_method=S256",
+ expected: `Apri https://accounts.google.com/o/oauth2/auth?response_type=code&client_id=test-client&redirect_uri=http%3A%2F%2Flocalhost%3A8001%2Foauth2callback&code_challenge=abc_def&code_challenge_method=S256`,
+ },
{
name: "link with underscores in URL is not corrupted by italic regex",
// Google Flights URLs use URL-safe base64 with underscores in the tfs param.
@@ -45,6 +50,11 @@ func Test_markdownToTelegramHTML(t *testing.T) {
input: "[first](https://a.com/path_one) and [second](https://b.com/path_two_x)",
expected: `first and second`,
},
+ {
+ name: "markdown link query params are escaped in href",
+ input: "[oauth](https://example.com/cb?response_type=code&client_id=test-client)",
+ expected: `oauth`,
+ },
{
name: "link label with HTML special chars is escaped",
input: "[a & b](https://example.com)",
@@ -55,6 +65,11 @@ func Test_markdownToTelegramHTML(t *testing.T) {
input: "a & b < c > d",
expected: "a & b < c > d",
},
+ {
+ name: "code block with language",
+ input: "```json\n{\n \"path\": \"README.md\"\n}\n```",
+ expected: "{\n \"path\": \"README.md\"\n}\n",
+ },
}
for _, tc := range cases {
diff --git a/pkg/config/config.go b/pkg/config/config.go
index 161108638..6bb8d3ce6 100644
--- a/pkg/config/config.go
+++ b/pkg/config/config.go
@@ -247,8 +247,9 @@ type SubTurnConfig struct {
}
type ToolFeedbackConfig struct {
- Enabled bool `json:"enabled" env:"PICOCLAW_AGENTS_DEFAULTS_TOOL_FEEDBACK_ENABLED"`
- MaxArgsLength int `json:"max_args_length" env:"PICOCLAW_AGENTS_DEFAULTS_TOOL_FEEDBACK_MAX_ARGS_LENGTH"`
+ Enabled bool `json:"enabled" env:"PICOCLAW_AGENTS_DEFAULTS_TOOL_FEEDBACK_ENABLED"`
+ MaxArgsLength int `json:"max_args_length" env:"PICOCLAW_AGENTS_DEFAULTS_TOOL_FEEDBACK_MAX_ARGS_LENGTH"`
+ SeparateMessages bool `json:"separate_messages" env:"PICOCLAW_AGENTS_DEFAULTS_TOOL_FEEDBACK_SEPARATE_MESSAGES"`
}
type AgentDefaults struct {
@@ -299,6 +300,13 @@ func (d *AgentDefaults) IsToolFeedbackEnabled() bool {
return d.ToolFeedback.Enabled
}
+// IsToolFeedbackSeparateMessagesEnabled returns true when each tool feedback
+// update should be sent as its own chat message instead of editing a single
+// in-place progress message.
+func (d *AgentDefaults) IsToolFeedbackSeparateMessagesEnabled() bool {
+ return d.ToolFeedback.SeparateMessages
+}
+
// GetModelName returns the effective model name for the agent defaults.
// It prefers the new "model_name" field but falls back to "model" for backward compatibility.
func (d *AgentDefaults) GetModelName() string {
diff --git a/pkg/config/config_test.go b/pkg/config/config_test.go
index 624cc7305..2be1bcc67 100644
--- a/pkg/config/config_test.go
+++ b/pkg/config/config_test.go
@@ -787,6 +787,9 @@ func TestDefaultConfig_ToolFeedbackDisabled(t *testing.T) {
if cfg.Agents.Defaults.ToolFeedback.Enabled {
t.Fatal("DefaultConfig().Agents.Defaults.ToolFeedback.Enabled should be false")
}
+ if cfg.Agents.Defaults.ToolFeedback.SeparateMessages {
+ t.Fatal("DefaultConfig().Agents.Defaults.ToolFeedback.SeparateMessages should be false")
+ }
}
func TestLoadConfig_ToolFeedbackDefaultsFalseWhenUnset(t *testing.T) {
@@ -807,6 +810,9 @@ func TestLoadConfig_ToolFeedbackDefaultsFalseWhenUnset(t *testing.T) {
if cfg.Agents.Defaults.ToolFeedback.Enabled {
t.Fatal("agents.defaults.tool_feedback.enabled should remain false when unset in config file")
}
+ if cfg.Agents.Defaults.ToolFeedback.SeparateMessages {
+ t.Fatal("agents.defaults.tool_feedback.separate_messages should remain false when unset in config file")
+ }
}
func TestLoadConfig_WebPreferNativeDefaultsTrueWhenUnset(t *testing.T) {
diff --git a/pkg/config/defaults.go b/pkg/config/defaults.go
index 35ef7cdd8..f3aaca7ab 100644
--- a/pkg/config/defaults.go
+++ b/pkg/config/defaults.go
@@ -35,8 +35,9 @@ func DefaultConfig() *Config {
SummarizeTokenPercent: 75,
SteeringMode: "one-at-a-time",
ToolFeedback: ToolFeedbackConfig{
- Enabled: false,
- MaxArgsLength: 300,
+ Enabled: false,
+ MaxArgsLength: 300,
+ SeparateMessages: false,
},
SplitOnMarker: false,
},
diff --git a/pkg/isolation/platform_windows.go b/pkg/isolation/platform_windows.go
index 9b39c85cf..1b3be8bd3 100644
--- a/pkg/isolation/platform_windows.go
+++ b/pkg/isolation/platform_windows.go
@@ -102,7 +102,7 @@ func postStartPlatformIsolation(cmd *exec.Cmd, isolation config.IsolationConfig,
return fmt.Errorf("open process for job assignment: %w", err)
}
- if err := windows.AssignProcessToJobObject(job, proc); err != nil {
+ if err = windows.AssignProcessToJobObject(job, proc); err != nil {
_ = windows.CloseHandle(proc)
_ = windows.CloseHandle(job)
if resources.token != 0 {
diff --git a/pkg/mcp/manager.go b/pkg/mcp/manager.go
index f589f82a9..92ea426a6 100644
--- a/pkg/mcp/manager.go
+++ b/pkg/mcp/manager.go
@@ -25,6 +25,24 @@ type headerTransport struct {
headers map[string]string
}
+func expandHomeCommandPath(command string) string {
+ if command == "" || command[0] != '~' {
+ return command
+ }
+
+ home, err := os.UserHomeDir()
+ if err != nil {
+ return command
+ }
+ if command == "~" {
+ return home
+ }
+ if strings.HasPrefix(command, "~/") || strings.HasPrefix(command, "~\\") {
+ return filepath.Join(home, command[2:])
+ }
+ return command
+}
+
func (t *headerTransport) RoundTrip(req *http.Request) (*http.Response, error) {
// Clone the request to avoid modifying the original
req = req.Clone(req.Context())
@@ -99,10 +117,12 @@ func loadEnvFile(path string) (map[string]string, error) {
// ServerConnection represents a connection to an MCP server
type ServerConnection struct {
- Name string
- Client *mcp.Client
- Session *mcp.ClientSession
- Tools []*mcp.Tool
+ Name string
+ Config config.MCPServerConfig
+ Client *mcp.Client
+ Session *mcp.ClientSession
+ Tools []*mcp.Tool
+ reconnectMu sync.Mutex
}
// Manager manages multiple MCP server connections
@@ -113,6 +133,8 @@ type Manager struct {
wg sync.WaitGroup // tracks in-flight CallTool calls
}
+var connectServerFunc = connectServer
+
// NewManager creates a new MCP manager
func NewManager() *Manager {
return &Manager{
@@ -242,6 +264,28 @@ func (m *Manager) ConnectServer(
name string,
cfg config.MCPServerConfig,
) error {
+ conn, err := connectServerFunc(ctx, name, cfg)
+ if err != nil {
+ return err
+ }
+
+ m.mu.Lock()
+ defer m.mu.Unlock()
+
+ if m.closed.Load() {
+ _ = conn.Session.Close()
+ return fmt.Errorf("manager is closed")
+ }
+
+ m.servers[name] = conn
+ return nil
+}
+
+func connectServer(
+ ctx context.Context,
+ name string,
+ cfg config.MCPServerConfig,
+) (*ServerConnection, error) {
logger.InfoCF("mcp", "Connecting to MCP server",
map[string]any{
"server": name,
@@ -267,14 +311,14 @@ func (m *Manager) ConnectServer(
} else if cfg.Command != "" {
transportType = "stdio"
} else {
- return fmt.Errorf("either URL or command must be provided")
+ return nil, fmt.Errorf("either URL or command must be provided")
}
}
switch transportType {
case "sse", "http":
if cfg.URL == "" {
- return fmt.Errorf("URL is required for SSE/HTTP transport")
+ return nil, fmt.Errorf("URL is required for SSE/HTTP transport")
}
// Configure DisableStandaloneSSE based on transport type.
@@ -316,7 +360,7 @@ func (m *Manager) ConnectServer(
transport = sseTransport
case "stdio":
if cfg.Command == "" {
- return fmt.Errorf("command is required for stdio transport")
+ return nil, fmt.Errorf("command is required for stdio transport")
}
logger.DebugCF("mcp", "Using stdio transport",
map[string]any{
@@ -324,7 +368,7 @@ func (m *Manager) ConnectServer(
"command": cfg.Command,
})
// Create command with context
- cmd := exec.CommandContext(ctx, cfg.Command, cfg.Args...)
+ cmd := exec.CommandContext(ctx, expandHomeCommandPath(cfg.Command), cfg.Args...)
// Build environment variables with proper override semantics
// Use a map to ensure config variables override file variables
@@ -341,7 +385,7 @@ func (m *Manager) ConnectServer(
if cfg.EnvFile != "" {
envVars, err := loadEnvFile(cfg.EnvFile)
if err != nil {
- return fmt.Errorf("failed to load env file %s: %w", cfg.EnvFile, err)
+ return nil, fmt.Errorf("failed to load env file %s: %w", cfg.EnvFile, err)
}
for k, v := range envVars {
envMap[k] = v
@@ -367,7 +411,7 @@ func (m *Manager) ConnectServer(
cmd.Env = env
transport = &isolatedCommandTransport{Command: cmd}
default:
- return fmt.Errorf(
+ return nil, fmt.Errorf(
"unsupported transport type: %s (supported: stdio, sse, http)",
transportType,
)
@@ -376,7 +420,7 @@ func (m *Manager) ConnectServer(
// Connect to server
session, err := client.Connect(ctx, transport, nil)
if err != nil {
- return fmt.Errorf("failed to connect: %w", err)
+ return nil, fmt.Errorf("failed to connect: %w", err)
}
// Get server info
@@ -390,38 +434,19 @@ func (m *Manager) ConnectServer(
})
// List available tools if supported
- var tools []*mcp.Tool
- if initResult.Capabilities.Tools != nil {
- for tool, err := range session.Tools(ctx, nil) {
- if err != nil {
- logger.WarnCF("mcp", "Error listing tool",
- map[string]any{
- "server": name,
- "error": err.Error(),
- })
- continue
- }
- tools = append(tools, tool)
- }
-
- logger.InfoCF("mcp", "Listed tools from MCP server",
- map[string]any{
- "server": name,
- "toolCount": len(tools),
- })
+ tools, err := listServerTools(ctx, name, session, initResult)
+ if err != nil {
+ _ = session.Close()
+ return nil, err
}
- // Store connection
- m.mu.Lock()
- m.servers[name] = &ServerConnection{
+ return &ServerConnection{
Name: name,
+ Config: cfg,
Client: client,
Session: session,
Tools: tools,
- }
- m.mu.Unlock()
-
- return nil
+ }, nil
}
// GetServers returns all connected servers
@@ -480,12 +505,131 @@ func (m *Manager) CallTool(
result, err := conn.Session.CallTool(ctx, params)
if err != nil {
+ if shouldReconnectCallError(err) {
+ logger.WarnCF("mcp", "MCP server session was lost during tool call, reconnecting",
+ map[string]any{
+ "server": serverName,
+ "tool": toolName,
+ "error": err.Error(),
+ })
+
+ reconnectedConn, reconnectErr := m.reconnectServer(ctx, serverName, conn)
+ if reconnectErr != nil {
+ return nil, fmt.Errorf("failed to recover lost MCP session: %w", reconnectErr)
+ }
+
+ result, err = reconnectedConn.Session.CallTool(ctx, params)
+ if err == nil {
+ return result, nil
+ }
+ }
+
return nil, fmt.Errorf("failed to call tool: %w", err)
}
return result, nil
}
+func listServerTools(
+ ctx context.Context,
+ name string,
+ session *mcp.ClientSession,
+ initResult *mcp.InitializeResult,
+) ([]*mcp.Tool, error) {
+ var tools []*mcp.Tool
+ if initResult.Capabilities.Tools == nil {
+ return tools, nil
+ }
+
+ for tool, err := range session.Tools(ctx, nil) {
+ if err != nil {
+ logger.WarnCF("mcp", "Error listing tool",
+ map[string]any{
+ "server": name,
+ "error": err.Error(),
+ })
+ continue
+ }
+ tools = append(tools, tool)
+ }
+
+ logger.InfoCF("mcp", "Listed tools from MCP server",
+ map[string]any{
+ "server": name,
+ "toolCount": len(tools),
+ })
+
+ return tools, nil
+}
+
+func shouldReconnectCallError(err error) bool {
+ if err == nil {
+ return false
+ }
+ if errors.Is(err, mcp.ErrSessionMissing) {
+ return true
+ }
+ return strings.Contains(strings.ToLower(err.Error()), mcp.ErrSessionMissing.Error())
+}
+
+func (m *Manager) reconnectServer(
+ ctx context.Context,
+ serverName string,
+ staleConn *ServerConnection,
+) (*ServerConnection, error) {
+ if staleConn == nil {
+ return nil, fmt.Errorf("server %s not found", serverName)
+ }
+
+ staleConn.reconnectMu.Lock()
+ defer staleConn.reconnectMu.Unlock()
+
+ if m.closed.Load() {
+ return nil, fmt.Errorf("manager is closed")
+ }
+
+ m.mu.RLock()
+ currentConn, ok := m.servers[serverName]
+ m.mu.RUnlock()
+ if !ok {
+ return nil, fmt.Errorf("server %s not found", serverName)
+ }
+ if currentConn != staleConn {
+ return currentConn, nil
+ }
+
+ freshConn, err := connectServerFunc(ctx, serverName, staleConn.Config)
+ if err != nil {
+ return nil, err
+ }
+
+ m.mu.Lock()
+ if m.closed.Load() {
+ m.mu.Unlock()
+ _ = freshConn.Session.Close()
+ return nil, fmt.Errorf("manager is closed")
+ }
+
+ currentConn, ok = m.servers[serverName]
+ if !ok {
+ m.mu.Unlock()
+ _ = freshConn.Session.Close()
+ return nil, fmt.Errorf("server %s not found", serverName)
+ }
+
+ if currentConn == staleConn {
+ m.servers[serverName] = freshConn
+ staleToClose := staleConn
+ m.mu.Unlock()
+ _ = staleToClose.Session.Close()
+ return freshConn, nil
+ }
+
+ m.mu.Unlock()
+ _ = freshConn.Session.Close()
+ return currentConn, nil
+}
+
// Close closes all server connections
func (m *Manager) Close() error {
// Use Swap to atomically set closed=true and get the previous value
diff --git a/pkg/mcp/manager_test.go b/pkg/mcp/manager_test.go
index f353942ab..682d4c346 100644
--- a/pkg/mcp/manager_test.go
+++ b/pkg/mcp/manager_test.go
@@ -2,11 +2,16 @@ package mcp
import (
"context"
+ "encoding/json"
+ "fmt"
+ "io"
"os"
"path/filepath"
"strings"
+ "sync"
"testing"
+ "github.com/modelcontextprotocol/go-sdk/jsonrpc"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/sipeed/picoclaw/pkg/config"
@@ -136,6 +141,22 @@ func TestLoadEnvFileNotFound(t *testing.T) {
}
}
+func TestExpandHomeCommandPath(t *testing.T) {
+ homeDir := t.TempDir()
+ t.Setenv("HOME", homeDir)
+ t.Setenv("USERPROFILE", homeDir)
+
+ want := filepath.Join(homeDir, "bin", "my-mcp")
+ got := expandHomeCommandPath("~" + string(os.PathSeparator) + filepath.Join("bin", "my-mcp"))
+ if got != want {
+ t.Fatalf("expandHomeCommandPath() = %q, want %q", got, want)
+ }
+
+ if got := expandHomeCommandPath("npx"); got != "npx" {
+ t.Fatalf("expandHomeCommandPath() should leave bare commands unchanged, got %q", got)
+ }
+}
+
func TestEnvFilePriority(t *testing.T) {
// Create a temporary .env file
tmpDir := t.TempDir()
@@ -296,6 +317,81 @@ func TestCallTool_ErrorsForClosedOrMissingServer(t *testing.T) {
})
}
+func TestCallTool_ReconnectsWhenHTTPServerLosesSession(t *testing.T) {
+ originalConnectServerFunc := connectServerFunc
+ t.Cleanup(func() {
+ connectServerFunc = originalConnectServerFunc
+ })
+
+ staleConn, staleTransport, err := newScriptedServerConnection(
+ "session-1",
+ nil,
+ fmt.Errorf(`sending "tools/call": failed to connect (session ID: session-1): %w`, sdkmcp.ErrSessionMissing),
+ )
+ if err != nil {
+ t.Fatalf("newScriptedServerConnection(stale) error = %v", err)
+ }
+ freshConn, freshTransport, err := newScriptedServerConnection(
+ "session-2",
+ &sdkmcp.CallToolResult{
+ Content: []sdkmcp.Content{
+ &sdkmcp.TextContent{Text: "reconnected"},
+ },
+ },
+ nil,
+ )
+ if err != nil {
+ t.Fatalf("newScriptedServerConnection(fresh) error = %v", err)
+ }
+
+ connectCalls := 0
+ connectServerFunc = func(ctx context.Context, name string, cfg config.MCPServerConfig) (*ServerConnection, error) {
+ connectCalls++
+ if connectCalls == 1 {
+ return freshConn, nil
+ }
+ return nil, fmt.Errorf("unexpected reconnect attempt %d", connectCalls)
+ }
+
+ mgr := NewManager()
+ mgr.servers["flaky"] = staleConn
+
+ result, err := mgr.CallTool(context.Background(), "flaky", "echo", map[string]any{
+ "query": "hello",
+ })
+ if err != nil {
+ t.Fatalf("CallTool() error = %v", err)
+ }
+ if result == nil || len(result.Content) != 1 {
+ t.Fatalf("CallTool() returned unexpected content: %#v", result)
+ }
+
+ text, ok := result.Content[0].(*sdkmcp.TextContent)
+ if !ok {
+ t.Fatalf("CallTool() content type = %T, want *sdkmcp.TextContent", result.Content[0])
+ }
+ if text.Text != "reconnected" {
+ t.Fatalf("CallTool() text = %q, want %q", text.Text, "reconnected")
+ }
+
+ conn, ok := mgr.GetServer("flaky")
+ if !ok {
+ t.Fatal("expected flaky server to remain connected after reconnect")
+ }
+ if conn.Session.ID() != "session-2" {
+ t.Fatalf("Session.ID() = %q, want %q", conn.Session.ID(), "session-2")
+ }
+ if connectCalls != 1 {
+ t.Fatalf("connectCalls = %d, want 1", connectCalls)
+ }
+ if staleTransport.toolCallCalls != 1 {
+ t.Fatalf("stale toolCallCalls = %d, want 1", staleTransport.toolCallCalls)
+ }
+ if freshTransport.toolCallCalls != 1 {
+ t.Fatalf("fresh toolCallCalls = %d, want 1", freshTransport.toolCallCalls)
+ }
+}
+
func TestClose_IdempotentOnEmptyManager(t *testing.T) {
mgr := NewManager()
@@ -306,3 +402,138 @@ func TestClose_IdempotentOnEmptyManager(t *testing.T) {
t.Fatalf("second close should be idempotent, got: %v", err)
}
}
+
+func newScriptedServerConnection(
+ sessionID string,
+ toolCallResult *sdkmcp.CallToolResult,
+ toolCallErr error,
+) (*ServerConnection, *scriptedTransport, error) {
+ transport := &scriptedTransport{
+ sessionID: sessionID,
+ toolCallResult: toolCallResult,
+ toolCallErr: toolCallErr,
+ }
+
+ client := sdkmcp.NewClient(&sdkmcp.Implementation{
+ Name: "picoclaw-test",
+ Version: "1.0.0",
+ }, nil)
+ session, err := client.Connect(context.Background(), transport, nil)
+ if err != nil {
+ return nil, nil, err
+ }
+
+ return &ServerConnection{
+ Name: "flaky",
+ Config: config.MCPServerConfig{Enabled: true, Type: "http", URL: "https://example.invalid/mcp"},
+ Client: client,
+ Session: session,
+ Tools: []*sdkmcp.Tool{
+ {
+ Name: "echo",
+ Description: "Echo test tool",
+ InputSchema: map[string]any{"type": "object"},
+ },
+ },
+ }, transport, nil
+}
+
+type scriptedTransport struct {
+ sessionID string
+ toolCallResult *sdkmcp.CallToolResult
+ toolCallErr error
+
+ mu sync.Mutex
+ toolCallCalls int
+ closed bool
+ incoming chan jsonrpc.Message
+}
+
+func (t *scriptedTransport) Connect(context.Context) (sdkmcp.Connection, error) {
+ if t.incoming == nil {
+ t.incoming = make(chan jsonrpc.Message, 4)
+ }
+ return t, nil
+}
+
+func (t *scriptedTransport) Read(ctx context.Context) (jsonrpc.Message, error) {
+ select {
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ case msg, ok := <-t.incoming:
+ if !ok {
+ return nil, io.EOF
+ }
+ return msg, nil
+ }
+}
+
+func (t *scriptedTransport) Write(ctx context.Context, msg jsonrpc.Message) error {
+ req, ok := msg.(*jsonrpc.Request)
+ if !ok {
+ return nil
+ }
+
+ switch req.Method {
+ case "initialize":
+ payload, err := json.Marshal(&sdkmcp.InitializeResult{
+ ProtocolVersion: "2025-11-25",
+ ServerInfo: &sdkmcp.Implementation{
+ Name: "scripted-test-server",
+ Version: "1.0.0",
+ },
+ Capabilities: &sdkmcp.ServerCapabilities{
+ Tools: &sdkmcp.ToolCapabilities{},
+ },
+ })
+ if err != nil {
+ return err
+ }
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ case t.incoming <- &jsonrpc.Response{ID: req.ID, Result: payload}:
+ return nil
+ }
+
+ case "notifications/initialized":
+ return nil
+
+ case "tools/call":
+ t.mu.Lock()
+ t.toolCallCalls++
+ t.mu.Unlock()
+
+ if t.toolCallErr != nil {
+ return t.toolCallErr
+ }
+
+ payload, err := json.Marshal(t.toolCallResult)
+ if err != nil {
+ return err
+ }
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ case t.incoming <- &jsonrpc.Response{ID: req.ID, Result: payload}:
+ return nil
+ }
+ }
+
+ return fmt.Errorf("unexpected method %q", req.Method)
+}
+
+func (t *scriptedTransport) Close() error {
+ t.mu.Lock()
+ defer t.mu.Unlock()
+ if t.closed {
+ return nil
+ }
+ t.closed = true
+ close(t.incoming)
+ return nil
+}
+
+func (t *scriptedTransport) SessionID() string {
+ return t.sessionID
+}
diff --git a/pkg/memory/jsonl.go b/pkg/memory/jsonl.go
index 8d3320f3f..492205114 100644
--- a/pkg/memory/jsonl.go
+++ b/pkg/memory/jsonl.go
@@ -10,12 +10,14 @@ import (
"log"
"os"
"path/filepath"
+ "sort"
"strings"
"sync"
"time"
"github.com/sipeed/picoclaw/pkg/fileutil"
"github.com/sipeed/picoclaw/pkg/providers"
+ "github.com/sipeed/picoclaw/pkg/providers/messageutil"
)
const (
@@ -405,12 +407,9 @@ func (s *JSONLStore) promoteAliasHistoryLocked(
}
func (s *JSONLStore) sessionHasVisibleContentLocked(sessionKey string, meta SessionMeta) (bool, error) {
- if meta.Count-meta.Skip > 0 || strings.TrimSpace(meta.Summary) != "" {
+ if strings.TrimSpace(meta.Summary) != "" {
return true, nil
}
- if meta.Count != 0 || meta.Skip != 0 {
- return false, nil
- }
history, err := readMessages(s.jsonlPath(sessionKey), meta.Skip)
if err != nil {
return false, err
@@ -482,6 +481,9 @@ func readMessages(path string, skip int) ([]providers.Message, error) {
lineNum, filepath.Base(path), err)
continue
}
+ if messageutil.IsTransientAssistantThoughtMessage(msg) {
+ continue
+ }
msgs = append(msgs, msg)
}
if scanner.Err() != nil {
@@ -494,28 +496,44 @@ func readMessages(path string, skip int) ([]providers.Message, error) {
return msgs, nil
}
-// countLines counts the total number of non-empty lines in a .jsonl file.
-// Used by TruncateHistory to reconcile a stale meta.Count without
-// the overhead of unmarshaling every message.
-func countLines(path string) (int, error) {
+// scanRetainedMessageLines returns the total number of non-empty raw JSONL
+// lines plus the raw line numbers that survive readMessages filtering.
+// TruncateHistory uses this to compute keepLast against retained messages
+// while preserving the raw-line skip offset stored in metadata.
+func scanRetainedMessageLines(path string) (int, []int, error) {
f, err := os.Open(path)
if os.IsNotExist(err) {
- return 0, nil
+ return 0, []int{}, nil
}
if err != nil {
- return 0, fmt.Errorf("memory: open jsonl: %w", err)
+ return 0, nil, fmt.Errorf("memory: open jsonl: %w", err)
}
defer f.Close()
- n := 0
+ rawCount := 0
+ retained := make([]int, 0)
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize)
for scanner.Scan() {
- if len(scanner.Bytes()) > 0 {
- n++
+ line := scanner.Bytes()
+ if len(line) == 0 {
+ continue
}
+ rawCount++
+
+ var msg providers.Message
+ if err := json.Unmarshal(line, &msg); err != nil {
+ continue
+ }
+ if messageutil.IsTransientAssistantThoughtMessage(msg) {
+ continue
+ }
+ retained = append(retained, rawCount)
}
- return n, scanner.Err()
+ if err := scanner.Err(); err != nil {
+ return 0, nil, err
+ }
+ return rawCount, retained, nil
}
func (s *JSONLStore) AddMessage(
@@ -535,6 +553,10 @@ func (s *JSONLStore) AddFullMessage(
// addMsg is the shared implementation for AddMessage and AddFullMessage.
func (s *JSONLStore) addMsg(sessionKey string, msg providers.Message) error {
+ if messageutil.IsTransientAssistantThoughtMessage(msg) {
+ return nil
+ }
+
l := s.sessionLock(sessionKey)
l.Lock()
defer l.Unlock()
@@ -655,24 +677,26 @@ func (s *JSONLStore) TruncateHistory(
return err
}
- // Always reconcile meta.Count with the actual line count on disk.
- // A crash between the JSONL append and the meta update in addMsg
- // leaves meta.Count stale (e.g. file has 101 lines but meta says
- // 100). Counting lines is cheap — no unmarshal, just a scan — and
- // TruncateHistory is not a hot path, so always re-count.
- n, countErr := countLines(s.jsonlPath(sessionKey))
- if countErr != nil {
- return countErr
+ rawCount, retainedRawLines, scanErr := scanRetainedMessageLines(s.jsonlPath(sessionKey))
+ if scanErr != nil {
+ return scanErr
}
- meta.Count = n
-
- if keepLast <= 0 {
+ meta.Count = rawCount
+ if meta.Skip > meta.Count {
meta.Skip = meta.Count
- } else {
- effective := meta.Count - meta.Skip
- if keepLast < effective {
- meta.Skip = meta.Count - keepLast
- }
+ }
+
+ activeStart := sort.Search(len(retainedRawLines), func(i int) bool {
+ return retainedRawLines[i] > meta.Skip
+ })
+ activeRetainedCount := len(retainedRawLines) - activeStart
+
+ switch {
+ case keepLast <= 0 || activeRetainedCount == 0:
+ meta.Skip = meta.Count
+ case keepLast < activeRetainedCount:
+ activeRawLines := retainedRawLines[activeStart:]
+ meta.Skip = activeRawLines[activeRetainedCount-keepLast-1]
}
meta.UpdatedAt = time.Now()
@@ -684,6 +708,8 @@ func (s *JSONLStore) SetHistory(
sessionKey string,
history []providers.Message,
) error {
+ history = messageutil.FilterInvalidHistoryMessages(history)
+
l := s.sessionLock(sessionKey)
l.Lock()
defer l.Unlock()
@@ -762,6 +788,8 @@ func (s *JSONLStore) Compact(
func (s *JSONLStore) rewriteJSONL(
sessionKey string, msgs []providers.Message,
) error {
+ msgs = messageutil.FilterInvalidHistoryMessages(msgs)
+
var buf bytes.Buffer
for i, msg := range msgs {
line, err := json.Marshal(msg)
diff --git a/pkg/memory/jsonl_test.go b/pkg/memory/jsonl_test.go
index b64c1b25f..3a7b98130 100644
--- a/pkg/memory/jsonl_test.go
+++ b/pkg/memory/jsonl_test.go
@@ -6,8 +6,10 @@ import (
"os"
"path/filepath"
"reflect"
+ "strings"
"sync"
"testing"
+ "time"
"github.com/sipeed/picoclaw/pkg/providers"
)
@@ -155,6 +157,27 @@ func TestAddFullMessage_ToolCallID(t *testing.T) {
}
}
+func TestAddFullMessage_DropsTransientAssistantThought(t *testing.T) {
+ store := newTestStore(t)
+ ctx := context.Background()
+
+ err := store.AddFullMessage(ctx, "transient-thought", providers.Message{
+ Role: "assistant",
+ ReasoningContent: "internal chain of thought",
+ })
+ if err != nil {
+ t.Fatalf("AddFullMessage: %v", err)
+ }
+
+ history, err := store.GetHistory(ctx, "transient-thought")
+ if err != nil {
+ t.Fatalf("GetHistory: %v", err)
+ }
+ if len(history) != 0 {
+ t.Fatalf("expected transient thought to be discarded, got %d messages", len(history))
+ }
+}
+
func TestGetHistory_EmptySession(t *testing.T) {
store := newTestStore(t)
ctx := context.Background()
@@ -243,6 +266,46 @@ func TestSetSummary_GetSummary(t *testing.T) {
}
}
+func TestSetHistory_DropsTransientAssistantThought(t *testing.T) {
+ store := newTestStore(t)
+ ctx := context.Background()
+
+ newHistory := []providers.Message{
+ {Role: "user", Content: "hello"},
+ {Role: "assistant", ReasoningContent: "internal chain of thought"},
+ {Role: "assistant", Content: "visible answer", ReasoningContent: "visible thought"},
+ }
+
+ err := store.SetHistory(ctx, "replace", newHistory)
+ if err != nil {
+ t.Fatalf("SetHistory: %v", err)
+ }
+
+ history, err := store.GetHistory(ctx, "replace")
+ if err != nil {
+ t.Fatalf("GetHistory: %v", err)
+ }
+ if len(history) != 2 {
+ t.Fatalf("expected transient thought to be removed, got %d messages", len(history))
+ }
+ if history[0].Role != "user" || history[0].Content != "hello" {
+ t.Fatalf("history[0] = %+v, want user/hello", history[0])
+ }
+ if history[1].Role != "assistant" || history[1].Content != "visible answer" ||
+ history[1].ReasoningContent != "visible thought" {
+ t.Fatalf("history[1] = %+v, want assistant visible answer with reasoning", history[1])
+ }
+
+ data, err := os.ReadFile(store.jsonlPath("replace"))
+ if err != nil {
+ t.Fatalf("ReadFile(jsonl): %v", err)
+ }
+ lines := strings.Split(strings.TrimSpace(string(data)), "\n")
+ if len(lines) != 2 {
+ t.Fatalf("jsonl line count = %d, want 2", len(lines))
+ }
+}
+
func TestSessionMetaScopeAndAliasesPersist(t *testing.T) {
store := newTestStore(t)
ctx := context.Background()
@@ -733,6 +796,56 @@ func TestTruncateHistory_StaleMetaCount(t *testing.T) {
}
}
+func TestTruncateHistory_IgnoresTransientThoughtForKeepLast(t *testing.T) {
+ store := newTestStore(t)
+ ctx := context.Background()
+ sessionKey := "transient-keep-last"
+ now := time.Now()
+
+ rawJSONL := strings.Join([]string{
+ `{"role":"user","content":"a"}`,
+ `{"role":"assistant","content":"b"}`,
+ `{"role":"assistant","content":"","reasoning_content":"dangling thought"}`,
+ `{"role":"user","content":"c"}`,
+ `{"role":"assistant","content":"d"}`,
+ }, "\n") + "\n"
+ if err := os.WriteFile(store.jsonlPath(sessionKey), []byte(rawJSONL), 0o644); err != nil {
+ t.Fatalf("WriteFile(jsonl): %v", err)
+ }
+ if err := store.writeMeta(sessionKey, SessionMeta{
+ Key: sessionKey,
+ Count: 5,
+ Skip: 0,
+ CreatedAt: now,
+ UpdatedAt: now,
+ }); err != nil {
+ t.Fatalf("writeMeta: %v", err)
+ }
+
+ if err := store.TruncateHistory(ctx, sessionKey, 2); err != nil {
+ t.Fatalf("TruncateHistory: %v", err)
+ }
+
+ history, err := store.GetHistory(ctx, sessionKey)
+ if err != nil {
+ t.Fatalf("GetHistory: %v", err)
+ }
+ if len(history) != 2 {
+ t.Fatalf("expected 2 retained messages, got %d", len(history))
+ }
+ if history[0].Content != "c" || history[1].Content != "d" {
+ t.Fatalf("kept history = %+v, want c,d", history)
+ }
+
+ meta, err := store.readMeta(sessionKey)
+ if err != nil {
+ t.Fatalf("readMeta: %v", err)
+ }
+ if meta.Skip != 2 {
+ t.Fatalf("meta.Skip = %d, want 2 raw lines skipped", meta.Skip)
+ }
+}
+
func TestCrashRecovery_PartialLine(t *testing.T) {
store := newTestStore(t)
ctx := context.Background()
diff --git a/pkg/pid/pidfile.go b/pkg/pid/pidfile.go
index f7c1f42b2..00601195f 100644
--- a/pkg/pid/pidfile.go
+++ b/pkg/pid/pidfile.go
@@ -58,7 +58,12 @@ func WritePidFile(homePath, host string, port int) (*PidFileData, error) {
if data, err := readPidFileUnlocked(pidPath); err == nil {
if os.Getpid() != data.PID {
logger.Infof("found pid file (PID: %d, version: %s)", data.PID, data.Version)
- if isProcessRunning(data.PID) {
+ // PID 1 is typically init/systemd on the host or the entrypoint
+ // inside a container. When a container stops and leaves behind a
+ // PID file on a shared volume, the host's PID 1 (init) would
+ // pass the isProcessRunning check, blocking new gateway starts.
+ // Treat recorded PID 1 as always stale.
+ if data.PID != 1 && isProcessRunning(data.PID) {
return nil, fmt.Errorf("gateway is already running (PID: %d, version: %s)", data.PID, data.Version)
}
logger.Warnf("not running (PID: %d) so will remove the pid file: %s", data.PID, pidPath)
@@ -124,6 +129,14 @@ func ReadPidFileWithCheck(homePath string) *PidFileData {
return nil
}
+ // Treat PID 1 as stale when we are not PID 1 ourselves (container
+ // leftover on a shared volume — host PID 1 is init, not gateway).
+ if data.PID == 1 && os.Getpid() != 1 {
+ logger.Debugf("stale container PID 1, remove pid file: %s", pidPath)
+ os.Remove(pidPath)
+ return nil
+ }
+
if !isProcessRunning(data.PID) {
logger.Debugf("process not running, remove pid file: %s", pidPath)
os.Remove(pidPath)
diff --git a/pkg/pid/pidfile_test.go b/pkg/pid/pidfile_test.go
index 2da44bbbc..2d3c11f63 100644
--- a/pkg/pid/pidfile_test.go
+++ b/pkg/pid/pidfile_test.go
@@ -278,6 +278,46 @@ func TestRemovePidFileIfPIDMismatch(t *testing.T) {
}
}
+// TestWritePidFileContainerPID1 verifies that a leftover PID file with PID 1
+// (typical container entrypoint) is treated as stale and overwritten.
+func TestWritePidFileContainerPID1(t *testing.T) {
+ dir := tmpDir(t)
+
+ stale := PidFileData{PID: 1, Token: "deadbeef12345678deadbeef12345678"}
+ raw, _ := json.MarshalIndent(stale, "", " ")
+ os.WriteFile(filepath.Join(dir, pidFileName), raw, 0o600)
+
+ data, err := WritePidFile(dir, "127.0.0.1", 18790)
+ if err != nil {
+ t.Fatalf("WritePidFile should treat PID 1 as stale, got error: %v", err)
+ }
+ if data.PID != os.Getpid() {
+ t.Errorf("PID = %d, want %d", data.PID, os.Getpid())
+ }
+}
+
+// TestReadPidFileWithCheckContainerPID1 verifies that a leftover PID file
+// with PID 1 is treated as stale and cleaned up.
+func TestReadPidFileWithCheckContainerPID1(t *testing.T) {
+ if os.Getpid() == 1 {
+ t.Skip("test not meaningful when running as PID 1")
+ }
+ dir := tmpDir(t)
+
+ stale := PidFileData{PID: 1, Token: "deadbeef12345678deadbeef12345678"}
+ raw, _ := json.MarshalIndent(stale, "", " ")
+ os.WriteFile(filepath.Join(dir, pidFileName), raw, 0o600)
+
+ data := ReadPidFileWithCheck(dir)
+ if data != nil {
+ t.Error("expected nil for PID 1 leftover")
+ }
+
+ if _, err := os.Stat(filepath.Join(dir, pidFileName)); !os.IsNotExist(err) {
+ t.Error("PID 1 leftover file should be removed")
+ }
+}
+
// TestReadPidFileUnlockedInvalidJSON returns error for malformed content.
func TestReadPidFileUnlockedInvalidJSON(t *testing.T) {
dir := tmpDir(t)
diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go
index 86d009811..ce83c6c54 100644
--- a/pkg/providers/factory_provider.go
+++ b/pkg/providers/factory_provider.go
@@ -178,7 +178,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if apiBase == "" {
apiBase = getDefaultAPIBase(protocol)
}
- return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
+ provider := NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
cfg.APIKey(),
apiBase,
cfg.Proxy,
@@ -187,7 +187,9 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
cfg.ExtraBody,
cfg.CustomHeaders,
- ), modelID, nil
+ )
+ provider.SetProviderName(protocol)
+ return provider, modelID, nil
case "azure", "azure-openai":
// Azure OpenAI uses deployment-based URLs, api-key header auth,
@@ -257,7 +259,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if apiBase == "" {
apiBase = getDefaultAPIBase(protocol)
}
- return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
+ provider := NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
cfg.APIKey(),
apiBase,
cfg.Proxy,
@@ -266,7 +268,9 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
cfg.ExtraBody,
cfg.CustomHeaders,
- ), modelID, nil
+ )
+ provider.SetProviderName(protocol)
+ return provider, modelID, nil
case "gemini":
if cfg.APIKey() == "" && cfg.APIBase == "" {
@@ -302,7 +306,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if _, ok := extraBody["reasoning_split"]; !ok {
extraBody["reasoning_split"] = true
}
- return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
+ provider := NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
cfg.APIKey(),
apiBase,
cfg.Proxy,
@@ -311,7 +315,9 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
extraBody,
cfg.CustomHeaders,
- ), modelID, nil
+ )
+ provider.SetProviderName(protocol)
+ return provider, modelID, nil
case "anthropic":
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
@@ -330,7 +336,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if cfg.APIKey() == "" {
return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model)
}
- return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
+ provider := NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
cfg.APIKey(),
apiBase,
cfg.Proxy,
@@ -339,7 +345,9 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
cfg.ExtraBody,
cfg.CustomHeaders,
- ), modelID, nil
+ )
+ provider.SetProviderName(protocol)
+ return provider, modelID, nil
case "anthropic-messages":
// Anthropic Messages API with native format (HTTP-based, no SDK)
diff --git a/pkg/providers/httpapi/http_provider.go b/pkg/providers/httpapi/http_provider.go
index a84962622..90f389cc8 100644
--- a/pkg/providers/httpapi/http_provider.go
+++ b/pkg/providers/httpapi/http_provider.go
@@ -77,3 +77,10 @@ func (p *HTTPProvider) GetDefaultModel() string {
func (p *HTTPProvider) SupportsNativeSearch() bool {
return p.delegate.SupportsNativeSearch()
}
+
+func (p *HTTPProvider) SetProviderName(providerName string) {
+ if p == nil || p.delegate == nil {
+ return
+ }
+ p.delegate.SetProviderName(providerName)
+}
diff --git a/pkg/providers/messageutil/messageutil.go b/pkg/providers/messageutil/messageutil.go
new file mode 100644
index 000000000..c4382d894
--- /dev/null
+++ b/pkg/providers/messageutil/messageutil.go
@@ -0,0 +1,38 @@
+package messageutil
+
+import (
+ "strings"
+
+ "github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
+)
+
+// IsTransientAssistantThoughtMessage reports whether msg is an invalid
+// reasoning-only assistant history record. These "hanging" thought messages
+// are not a canonical persisted format and should be discarded instead of
+// replayed or reconstructed.
+func IsTransientAssistantThoughtMessage(msg protocoltypes.Message) bool {
+ return msg.Role == "assistant" &&
+ strings.TrimSpace(msg.Content) == "" &&
+ strings.TrimSpace(msg.ReasoningContent) != "" &&
+ len(msg.ToolCalls) == 0 &&
+ len(msg.Media) == 0 &&
+ len(msg.Attachments) == 0 &&
+ strings.TrimSpace(msg.ToolCallID) == ""
+}
+
+// FilterInvalidHistoryMessages removes invalid persisted history records such
+// as transient assistant thought-only messages.
+func FilterInvalidHistoryMessages(history []protocoltypes.Message) []protocoltypes.Message {
+ if len(history) == 0 {
+ return []protocoltypes.Message{}
+ }
+
+ filtered := make([]protocoltypes.Message, 0, len(history))
+ for _, msg := range history {
+ if IsTransientAssistantThoughtMessage(msg) {
+ continue
+ }
+ filtered = append(filtered, msg)
+ }
+ return filtered
+}
diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go
index 29667cd31..c3733ce3a 100644
--- a/pkg/providers/openai_compat/provider.go
+++ b/pkg/providers/openai_compat/provider.go
@@ -15,6 +15,7 @@ import (
"time"
"github.com/sipeed/picoclaw/pkg/providers/common"
+ "github.com/sipeed/picoclaw/pkg/providers/messageutil"
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
)
@@ -34,6 +35,7 @@ type (
type Provider struct {
apiKey string
apiBase string
+ providerName string
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
httpClient *http.Client
extraBody map[string]any // Additional fields to inject into request body
@@ -95,6 +97,12 @@ func WithCustomHeaders(customHeaders map[string]string) Option {
}
}
+func WithProviderName(providerName string) Option {
+ return func(p *Provider) {
+ p.providerName = strings.ToLower(strings.TrimSpace(providerName))
+ }
+}
+
func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider {
p := &Provider{
apiKey: apiKey,
@@ -136,7 +144,7 @@ func (p *Provider) buildRequestBody(
requestBody := map[string]any{
"model": model,
- "messages": common.SerializeMessages(messages),
+ "messages": common.SerializeMessages(p.prepareMessagesForRequest(messages)),
}
// When fallback uses a different provider (e.g. DeepSeek), that provider must not inject web_search_preview.
@@ -196,6 +204,111 @@ func (p *Provider) applyCustomHeaders(req *http.Request) {
}
}
+func (p *Provider) SetProviderName(providerName string) {
+ p.providerName = strings.ToLower(strings.TrimSpace(providerName))
+}
+
+func (p *Provider) prepareMessagesForRequest(messages []Message) []Message {
+ if len(messages) == 0 {
+ return nil
+ }
+
+ if p.isDeepSeekReasoningProvider() {
+ return filterDeepSeekReasoningMessages(messages)
+ }
+ return stripReasoningMessages(messages)
+}
+
+func (p *Provider) isDeepSeekReasoningProvider() bool {
+ return p.providerName == "deepseek" || isDeepSeekHost(p.apiBase)
+}
+
+func isDeepSeekHost(apiBase string) bool {
+ parsed, err := url.Parse(strings.TrimSpace(apiBase))
+ if err != nil {
+ return false
+ }
+ host := strings.ToLower(strings.TrimSpace(parsed.Hostname()))
+ return host == "deepseek.com" || strings.HasSuffix(host, ".deepseek.com")
+}
+
+func filterDeepSeekReasoningMessages(messages []Message) []Message {
+ out := make([]Message, 0, len(messages))
+ start := 0
+
+ flush := func(end int) {
+ if end <= start {
+ return
+ }
+ out = append(out, filterDeepSeekReasoningTurn(messages[start:end])...)
+ start = end
+ }
+
+ for i := 1; i < len(messages); i++ {
+ if messages[i].Role == "user" {
+ flush(i)
+ }
+ }
+ flush(len(messages))
+
+ return out
+}
+
+func filterDeepSeekReasoningTurn(messages []Message) []Message {
+ hasToolInteraction := false
+ for _, msg := range messages {
+ if msg.Role == "tool" || (msg.Role == "assistant" && len(msg.ToolCalls) > 0) {
+ hasToolInteraction = true
+ break
+ }
+ }
+
+ out := make([]Message, 0, len(messages))
+ for _, msg := range messages {
+ if messageutil.IsTransientAssistantThoughtMessage(msg) {
+ continue
+ }
+
+ cloned := msg
+ if cloned.Role == "assistant" && strings.TrimSpace(cloned.ReasoningContent) != "" && !hasToolInteraction {
+ cloned.ReasoningContent = ""
+ }
+ if assistantMessageEmpty(cloned) {
+ continue
+ }
+ out = append(out, cloned)
+ }
+
+ return out
+}
+
+func stripReasoningMessages(messages []Message) []Message {
+ out := make([]Message, 0, len(messages))
+ for _, msg := range messages {
+ if messageutil.IsTransientAssistantThoughtMessage(msg) {
+ continue
+ }
+
+ cloned := msg
+ cloned.ReasoningContent = ""
+ if assistantMessageEmpty(cloned) {
+ continue
+ }
+ out = append(out, cloned)
+ }
+ return out
+}
+
+func assistantMessageEmpty(msg Message) bool {
+ return msg.Role == "assistant" &&
+ strings.TrimSpace(msg.Content) == "" &&
+ strings.TrimSpace(msg.ReasoningContent) == "" &&
+ len(msg.ToolCalls) == 0 &&
+ len(msg.Media) == 0 &&
+ len(msg.Attachments) == 0 &&
+ strings.TrimSpace(msg.ToolCallID) == ""
+}
+
func (p *Provider) Chat(
ctx context.Context,
messages []Message,
diff --git a/pkg/providers/openai_compat/provider_test.go b/pkg/providers/openai_compat/provider_test.go
index d140d63d6..594048ea5 100644
--- a/pkg/providers/openai_compat/provider_test.go
+++ b/pkg/providers/openai_compat/provider_test.go
@@ -202,7 +202,7 @@ func TestProviderChat_ParsesReasoningContent(t *testing.T) {
}
}
-func TestProviderChat_PreservesReasoningContentInHistory(t *testing.T) {
+func TestProviderChat_StripsReasoningContentForNonDeepSeekHistory(t *testing.T) {
var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -225,8 +225,6 @@ func TestProviderChat_PreservesReasoningContentInHistory(t *testing.T) {
p := NewProvider("key", server.URL, "")
- // Simulate a multi-turn conversation where the assistant's previous
- // reply included reasoning_content (e.g. from kimi-k2.5).
messages := []Message{
{Role: "user", Content: "What is 1+1?"},
{Role: "assistant", Content: "2", ReasoningContent: "Let me think... 1+1=2"},
@@ -238,7 +236,6 @@ func TestProviderChat_PreservesReasoningContentInHistory(t *testing.T) {
t.Fatalf("Chat() error = %v", err)
}
- // Verify reasoning_content is preserved in the serialized request.
reqMessages, ok := requestBody["messages"].([]any)
if !ok {
t.Fatalf("messages is not []any: %T", requestBody["messages"])
@@ -247,11 +244,288 @@ func TestProviderChat_PreservesReasoningContentInHistory(t *testing.T) {
if !ok {
t.Fatalf("assistant message is not map[string]any: %T", reqMessages[1])
}
- if assistantMsg["reasoning_content"] != "Let me think... 1+1=2" {
- t.Errorf("reasoning_content not preserved in request, got %v", assistantMsg["reasoning_content"])
+ if _, exists := assistantMsg["reasoning_content"]; exists {
+ t.Fatalf(
+ "reasoning_content should be stripped for non-DeepSeek providers, got %v",
+ assistantMsg["reasoning_content"],
+ )
}
}
+func TestProviderChat_DeepSeekOmitsReasoningContentForNonToolTurnHistory(t *testing.T) {
+ var requestBody map[string]any
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
+ http.Error(w, err.Error(), http.StatusBadRequest)
+ return
+ }
+ resp := map[string]any{
+ "choices": []map[string]any{
+ {
+ "message": map[string]any{"content": "ok"},
+ "finish_reason": "stop",
+ },
+ },
+ }
+ w.Header().Set("Content-Type", "application/json")
+ json.NewEncoder(w).Encode(resp)
+ }))
+ defer server.Close()
+
+ p := NewProvider("key", server.URL, "")
+ p.apiBase = "https://api.deepseek.com/v1"
+ p.httpClient = &http.Client{
+ Transport: roundTripperFunc(func(r *http.Request) (*http.Response, error) {
+ r.URL, _ = url.Parse(server.URL + r.URL.Path)
+ return http.DefaultTransport.RoundTrip(r)
+ }),
+ }
+
+ messages := []Message{
+ {Role: "user", Content: "What is 1+1?"},
+ {Role: "assistant", Content: "2", ReasoningContent: "Let me think... 1+1=2"},
+ {Role: "user", Content: "What about 2+2?"},
+ }
+
+ _, err := p.Chat(t.Context(), messages, nil, "deepseek-v4-flash", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ reqMessages, ok := requestBody["messages"].([]any)
+ if !ok {
+ t.Fatalf("messages is not []any: %T", requestBody["messages"])
+ }
+ assistantMsg, ok := reqMessages[1].(map[string]any)
+ if !ok {
+ t.Fatalf("assistant message is not map[string]any: %T", reqMessages[1])
+ }
+ if _, exists := assistantMsg["reasoning_content"]; exists {
+ t.Fatalf(
+ "reasoning_content should be omitted for DeepSeek non-tool turns, got %v",
+ assistantMsg["reasoning_content"],
+ )
+ }
+}
+
+func TestProviderChat_DeepSeekPreservesReasoningContentForToolTurnHistory(t *testing.T) {
+ var requestBody map[string]any
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
+ http.Error(w, err.Error(), http.StatusBadRequest)
+ return
+ }
+ resp := map[string]any{
+ "choices": []map[string]any{
+ {
+ "message": map[string]any{"content": "ok"},
+ "finish_reason": "stop",
+ },
+ },
+ }
+ w.Header().Set("Content-Type", "application/json")
+ json.NewEncoder(w).Encode(resp)
+ }))
+ defer server.Close()
+
+ p := NewProvider("key", server.URL, "")
+ p.SetProviderName("deepseek")
+
+ messages := []Message{
+ {Role: "user", Content: "How's the weather tomorrow?"},
+ {
+ Role: "assistant",
+ Content: "Let me check the date first.",
+ ReasoningContent: "I need tomorrow's date before checking the weather.",
+ ToolCalls: []ToolCall{{
+ ID: "call_1",
+ Type: "function",
+ Function: &FunctionCall{
+ Name: "get_date",
+ Arguments: "{}",
+ },
+ }},
+ },
+ {Role: "tool", ToolCallID: "call_1", Content: "2026-04-24"},
+ {
+ Role: "assistant",
+ Content: "Tomorrow is 2026-04-25.",
+ ReasoningContent: "Now I can share the final answer.",
+ },
+ {Role: "user", Content: "What about Guangzhou?"},
+ }
+
+ _, err := p.Chat(t.Context(), messages, nil, "deepseek-v4-flash", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ reqMessages, ok := requestBody["messages"].([]any)
+ if !ok {
+ t.Fatalf("messages is not []any: %T", requestBody["messages"])
+ }
+ if len(reqMessages) != len(messages) {
+ t.Fatalf("len(messages) = %d, want %d", len(reqMessages), len(messages))
+ }
+
+ firstAssistant, ok := reqMessages[1].(map[string]any)
+ if !ok {
+ t.Fatalf("first assistant message is not map[string]any: %T", reqMessages[1])
+ }
+ if firstAssistant["reasoning_content"] != "I need tomorrow's date before checking the weather." {
+ t.Fatalf("first assistant reasoning_content = %v, want preserved", firstAssistant["reasoning_content"])
+ }
+
+ finalAssistant, ok := reqMessages[3].(map[string]any)
+ if !ok {
+ t.Fatalf("final assistant message is not map[string]any: %T", reqMessages[3])
+ }
+ if finalAssistant["reasoning_content"] != "Now I can share the final answer." {
+ t.Fatalf("final assistant reasoning_content = %v, want preserved", finalAssistant["reasoning_content"])
+ }
+}
+
+func TestProviderChat_HistoryCanonicalizationMatrix(t *testing.T) {
+ baseMessages := []Message{
+ {Role: "user", Content: "turn1"},
+ {Role: "assistant", Content: "plain visible", ReasoningContent: "plain thought"},
+ {Role: "user", Content: "turn2"},
+ {
+ Role: "assistant",
+ Content: "",
+ ReasoningContent: "tool thought",
+ ToolCalls: []ToolCall{{
+ ID: "call_read_file",
+ Type: "function",
+ Function: &FunctionCall{
+ Name: "read_file",
+ Arguments: `{"path":"README.md"}`,
+ },
+ }},
+ },
+ {Role: "tool", ToolCallID: "call_read_file", Content: "file content"},
+ {Role: "user", Content: "turn3"},
+ {
+ Role: "assistant",
+ Content: "tool visible only",
+ ToolCalls: []ToolCall{{
+ ID: "call_list_dir",
+ Type: "function",
+ Function: &FunctionCall{
+ Name: "list_dir",
+ Arguments: `{"path":"."}`,
+ },
+ }},
+ },
+ {Role: "tool", ToolCallID: "call_list_dir", Content: "dir listing"},
+ {Role: "user", Content: "turn4"},
+ {
+ Role: "assistant",
+ Content: "tool visible and thought",
+ ReasoningContent: "tool mixed thought",
+ ToolCalls: []ToolCall{{
+ ID: "call_exec",
+ Type: "function",
+ Function: &FunctionCall{
+ Name: "exec",
+ Arguments: `{"command":"pwd"}`,
+ },
+ }},
+ },
+ {Role: "tool", ToolCallID: "call_exec", Content: "pwd output"},
+ {Role: "user", Content: "current turn"},
+ }
+
+ captureRequestMessages := func(t *testing.T, providerName string) []map[string]any {
+ t.Helper()
+
+ var requestBody map[string]any
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
+ http.Error(w, err.Error(), http.StatusBadRequest)
+ return
+ }
+ resp := map[string]any{
+ "choices": []map[string]any{
+ {
+ "message": map[string]any{"content": "ok"},
+ "finish_reason": "stop",
+ },
+ },
+ }
+ w.Header().Set("Content-Type", "application/json")
+ json.NewEncoder(w).Encode(resp)
+ }))
+ defer server.Close()
+
+ p := NewProvider("key", server.URL, "")
+ if providerName != "" {
+ p.SetProviderName(providerName)
+ }
+
+ _, err := p.Chat(t.Context(), baseMessages, nil, "test-model", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ rawMessages, ok := requestBody["messages"].([]any)
+ if !ok {
+ t.Fatalf("messages is not []any: %T", requestBody["messages"])
+ }
+
+ out := make([]map[string]any, 0, len(rawMessages))
+ for i, raw := range rawMessages {
+ msg, ok := raw.(map[string]any)
+ if !ok {
+ t.Fatalf("messages[%d] is %T, want map[string]any", i, raw)
+ }
+ out = append(out, msg)
+ }
+ return out
+ }
+
+ t.Run("deepseek", func(t *testing.T) {
+ msgs := captureRequestMessages(t, "deepseek")
+ if len(msgs) != len(baseMessages) {
+ t.Fatalf("len(messages) = %d, want %d", len(msgs), len(baseMessages))
+ }
+
+ if _, ok := msgs[1]["reasoning_content"]; ok {
+ t.Fatalf(
+ "turn1 reasoning_content should be stripped for DeepSeek non-tool turn, got %v",
+ msgs[1]["reasoning_content"],
+ )
+ }
+ if msgs[3]["reasoning_content"] != "tool thought" {
+ t.Fatalf("turn2 reasoning_content = %v, want preserved", msgs[3]["reasoning_content"])
+ }
+ if _, ok := msgs[6]["reasoning_content"]; ok {
+ t.Fatalf("turn3 reasoning_content should be absent, got %v", msgs[6]["reasoning_content"])
+ }
+ if msgs[9]["reasoning_content"] != "tool mixed thought" {
+ t.Fatalf("turn4 reasoning_content = %v, want preserved", msgs[9]["reasoning_content"])
+ }
+ if msgs[9]["content"] != "tool visible and thought" {
+ t.Fatalf("turn4 content = %v, want preserved", msgs[9]["content"])
+ }
+ })
+
+ t.Run("non-deepseek", func(t *testing.T) {
+ msgs := captureRequestMessages(t, "")
+ for i, msg := range msgs {
+ if _, ok := msg["reasoning_content"]; ok {
+ t.Fatalf(
+ "messages[%d] reasoning_content should be stripped for non-DeepSeek providers, got %v",
+ i,
+ msg["reasoning_content"],
+ )
+ }
+ }
+ })
+}
+
func TestProviderChat_HTTPError(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "bad request", http.StatusBadRequest)
diff --git a/pkg/providers/protocoltypes/types.go b/pkg/providers/protocoltypes/types.go
index f3553f8b0..bab4433e7 100644
--- a/pkg/providers/protocoltypes/types.go
+++ b/pkg/providers/protocoltypes/types.go
@@ -61,6 +61,13 @@ type ContentBlock struct {
Type string `json:"type"` // "text"
Text string `json:"text"`
CacheControl *CacheControl `json:"cache_control,omitempty"`
+
+ // Prompt metadata is internal to the agent runtime. It records which
+ // structured prompt segment produced this block without changing provider
+ // JSON.
+ PromptLayer string `json:"-"`
+ PromptSlot string `json:"-"`
+ PromptSource string `json:"-"`
}
type Attachment struct {
@@ -80,11 +87,24 @@ type Message struct {
SystemParts []ContentBlock `json:"system_parts,omitempty"` // structured system blocks for cache-aware adapters
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
+
+ // Prompt metadata is internal to the agent runtime. It records where a
+ // message or system part came from without changing provider/session JSON.
+ PromptLayer string `json:"-"`
+ PromptSlot string `json:"-"`
+ PromptSource string `json:"-"`
}
type ToolDefinition struct {
Type string `json:"type"`
Function ToolFunctionDefinition `json:"function"`
+
+ // Prompt metadata is internal to the agent runtime. Tool definitions are
+ // model-visible capability prompts even though providers send them outside
+ // the system message.
+ PromptLayer string `json:"-"`
+ PromptSlot string `json:"-"`
+ PromptSource string `json:"-"`
}
type ToolFunctionDefinition struct {
diff --git a/pkg/session/manager.go b/pkg/session/manager.go
index 7f87d460a..1d6fa3106 100644
--- a/pkg/session/manager.go
+++ b/pkg/session/manager.go
@@ -9,6 +9,7 @@ import (
"time"
"github.com/sipeed/picoclaw/pkg/providers"
+ "github.com/sipeed/picoclaw/pkg/providers/messageutil"
)
type Session struct {
@@ -69,6 +70,10 @@ func (sm *SessionManager) AddMessage(sessionKey, role, content string) {
// AddFullMessage adds a complete message with tool calls and tool call ID to the session.
// This is used to save the full conversation flow including tool calls and tool results.
func (sm *SessionManager) AddFullMessage(sessionKey string, msg providers.Message) {
+ if messageutil.IsTransientAssistantThoughtMessage(msg) {
+ return
+ }
+
sm.mu.Lock()
defer sm.mu.Unlock()
@@ -196,8 +201,7 @@ func (sm *SessionManager) Save(key string) error {
Updated: stored.Updated,
}
if len(stored.Messages) > 0 {
- snapshot.Messages = make([]providers.Message, len(stored.Messages))
- copy(snapshot.Messages, stored.Messages)
+ snapshot.Messages = messageutil.FilterInvalidHistoryMessages(stored.Messages)
} else {
snapshot.Messages = []providers.Message{}
}
@@ -270,6 +274,7 @@ func (sm *SessionManager) loadSessions() error {
if err := json.Unmarshal(data, &session); err != nil {
continue
}
+ session.Messages = messageutil.FilterInvalidHistoryMessages(session.Messages)
sm.sessions[session.Key] = &session
}
@@ -290,6 +295,7 @@ func (sm *SessionManager) SetHistory(key string, history []providers.Message) {
session, ok := sm.sessions[key]
if ok {
+ history = messageutil.FilterInvalidHistoryMessages(history)
// Create a deep copy to strictly isolate internal state
// from the caller's slice.
msgs := make([]providers.Message, len(history))
diff --git a/pkg/tools/integration/mcp_tool.go b/pkg/tools/integration/mcp_tool.go
index 340bb9e8e..78c348316 100644
--- a/pkg/tools/integration/mcp_tool.go
+++ b/pkg/tools/integration/mcp_tool.go
@@ -15,6 +15,7 @@ import (
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/media"
+ toolshared "github.com/sipeed/picoclaw/pkg/tools/shared"
)
// MCPManager defines the interface for MCP manager operations
@@ -161,6 +162,14 @@ func (t *MCPTool) Description() string {
return fmt.Sprintf("[MCP:%s] %s", t.serverName, desc)
}
+func (t *MCPTool) PromptMetadata() toolshared.PromptMetadata {
+ return toolshared.PromptMetadata{
+ Layer: toolshared.ToolPromptLayerCapability,
+ Slot: toolshared.ToolPromptSlotMCP,
+ Source: "mcp:" + sanitizeIdentifierComponent(t.serverName),
+ }
+}
+
// Parameters returns the tool parameters schema
func (t *MCPTool) Parameters() map[string]any {
// The InputSchema is already a JSON Schema object
diff --git a/pkg/tools/integration/mcp_tool_test.go b/pkg/tools/integration/mcp_tool_test.go
index e5c54abb6..7b0b2cd5a 100644
--- a/pkg/tools/integration/mcp_tool_test.go
+++ b/pkg/tools/integration/mcp_tool_test.go
@@ -11,6 +11,7 @@ import (
"github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/sipeed/picoclaw/pkg/media"
+ toolshared "github.com/sipeed/picoclaw/pkg/tools/shared"
)
// MockMCPManager is a mock implementation of MCPManager interface for testing
@@ -104,6 +105,22 @@ func TestMCPTool_Name(t *testing.T) {
}
}
+func TestMCPTool_PromptMetadata(t *testing.T) {
+ manager := &MockMCPManager{}
+ tool := NewMCPTool(manager, "GitHub Server", &mcp.Tool{Name: "create_issue"})
+
+ metadata := tool.PromptMetadata()
+ if metadata.Layer != toolshared.ToolPromptLayerCapability {
+ t.Fatalf("metadata.Layer = %q, want %q", metadata.Layer, toolshared.ToolPromptLayerCapability)
+ }
+ if metadata.Slot != toolshared.ToolPromptSlotMCP {
+ t.Fatalf("metadata.Slot = %q, want %q", metadata.Slot, toolshared.ToolPromptSlotMCP)
+ }
+ if metadata.Source != "mcp:github_server" {
+ t.Fatalf("metadata.Source = %q, want mcp:github_server", metadata.Source)
+ }
+}
+
// TestMCPTool_Description verifies tool description generation
func TestMCPTool_Description(t *testing.T) {
tests := []struct {
diff --git a/pkg/tools/integration/web.go b/pkg/tools/integration/web.go
index 56663ecda..75821e40d 100644
--- a/pkg/tools/integration/web.go
+++ b/pkg/tools/integration/web.go
@@ -58,8 +58,6 @@ var (
reSogouRealURL = regexp.MustCompile(`url=([^&]+)`)
)
-var preferredWebSearchLanguage atomic.Value
-
type APIKeyPool struct {
keys []string
current uint32
@@ -250,27 +248,6 @@ func mapBaiduRecencyFilter(rangeCode string) string {
}
}
-func normalizePreferredWebSearchLanguage(lang string) string {
- lang = strings.ToLower(strings.TrimSpace(lang))
- switch {
- case strings.HasPrefix(lang, "zh"), lang == "chinese":
- return "zh"
- case strings.HasPrefix(lang, "en"), lang == "english":
- return "en"
- default:
- return ""
- }
-}
-
-func SetPreferredWebSearchLanguage(lang string) {
- preferredWebSearchLanguage.Store(normalizePreferredWebSearchLanguage(lang))
-}
-
-func GetPreferredWebSearchLanguage() string {
- lang, _ := preferredWebSearchLanguage.Load().(string)
- return lang
-}
-
type BraveSearchProvider struct {
keyPool *APIKeyPool
proxy string
@@ -1420,7 +1397,7 @@ func containsLatinLetter(text string) bool {
func prefersDuckDuckGoQuery(text string) bool {
trimmed := strings.TrimSpace(text)
if trimmed == "" {
- return GetPreferredWebSearchLanguage() == "en"
+ return false
}
if containsHan(trimmed) {
return false
@@ -1428,7 +1405,7 @@ func prefersDuckDuckGoQuery(text string) bool {
if containsLatinLetter(trimmed) {
return true
}
- return GetPreferredWebSearchLanguage() == "en"
+ return false
}
func (opts WebSearchToolOptions) buildProviderResolver() (func(query string) (SearchProvider, int), error) {
diff --git a/pkg/tools/integration/web_test.go b/pkg/tools/integration/web_test.go
index d47d8e7c9..ba6b3da45 100644
--- a/pkg/tools/integration/web_test.go
+++ b/pkg/tools/integration/web_test.go
@@ -1778,11 +1778,6 @@ func TestApplySogouRangeHint(t *testing.T) {
}
func TestPrefersDuckDuckGoQuery(t *testing.T) {
- SetPreferredWebSearchLanguage("")
- t.Cleanup(func() {
- SetPreferredWebSearchLanguage("")
- })
-
tests := []struct {
name string
query string
@@ -1805,19 +1800,9 @@ func TestPrefersDuckDuckGoQuery(t *testing.T) {
}
}
-func TestPrefersDuckDuckGoQuery_FallsBackToPreferredLanguage(t *testing.T) {
- SetPreferredWebSearchLanguage("en")
- t.Cleanup(func() {
- SetPreferredWebSearchLanguage("")
- })
-
- if !prefersDuckDuckGoQuery("2026 04 15") {
- t.Fatal("numeric query should prefer DuckDuckGo when preferred language is English")
- }
-
- SetPreferredWebSearchLanguage("zh")
+func TestPrefersDuckDuckGoQuery_DoesNotUseGlobalLanguageFallback(t *testing.T) {
if prefersDuckDuckGoQuery("2026 04 15") {
- t.Fatal("numeric query should prefer Sogou when preferred language is Chinese")
+ t.Fatal("numeric query should default to Sogou when no script-specific hint is present")
}
}
diff --git a/pkg/tools/integration_facade.go b/pkg/tools/integration_facade.go
index b05a22fe2..193ecd6f5 100644
--- a/pkg/tools/integration_facade.go
+++ b/pkg/tools/integration_facade.go
@@ -65,14 +65,6 @@ func NewAPIKeyPool(keys []string) *APIKeyPool {
return integrationtools.NewAPIKeyPool(keys)
}
-func SetPreferredWebSearchLanguage(lang string) {
- integrationtools.SetPreferredWebSearchLanguage(lang)
-}
-
-func GetPreferredWebSearchLanguage() string {
- return integrationtools.GetPreferredWebSearchLanguage()
-}
-
func WebSearchToolOptionsFromConfig(cfg *config.Config) WebSearchToolOptions {
return integrationtools.WebSearchToolOptionsFromConfig(cfg)
}
diff --git a/pkg/tools/registry.go b/pkg/tools/registry.go
index e51dff71a..0ff9293a3 100644
--- a/pkg/tools/registry.go
+++ b/pkg/tools/registry.go
@@ -352,6 +352,7 @@ func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition {
name, _ := fn["name"].(string)
desc, _ := fn["description"].(string)
params, _ := fn["parameters"].(map[string]any)
+ metadata := promptMetadataForTool(entry.Tool)
definitions = append(definitions, providers.ToolDefinition{
Type: "function",
@@ -360,11 +361,35 @@ func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition {
Description: desc,
Parameters: params,
},
+ PromptLayer: metadata.Layer,
+ PromptSlot: metadata.Slot,
+ PromptSource: metadata.Source,
})
}
return definitions
}
+func promptMetadataForTool(tool Tool) PromptMetadata {
+ metadata := PromptMetadata{
+ Layer: ToolPromptLayerCapability,
+ Slot: ToolPromptSlotTooling,
+ Source: ToolPromptSourceRegistry,
+ }
+ if provider, ok := tool.(PromptMetadataProvider); ok {
+ provided := provider.PromptMetadata()
+ if provided.Layer != "" {
+ metadata.Layer = provided.Layer
+ }
+ if provided.Slot != "" {
+ metadata.Slot = provided.Slot
+ }
+ if provided.Source != "" {
+ metadata.Source = provided.Source
+ }
+ }
+ return metadata
+}
+
// List returns a list of all registered tool names.
func (r *ToolRegistry) List() []string {
r.mu.RLock()
diff --git a/pkg/tools/registry_test.go b/pkg/tools/registry_test.go
index 16bd30928..eac96382f 100644
--- a/pkg/tools/registry_test.go
+++ b/pkg/tools/registry_test.go
@@ -39,6 +39,15 @@ func (m *mockContextAwareTool) Execute(ctx context.Context, _ map[string]any) *T
return m.result
}
+type mockPromptMetadataTool struct {
+ mockRegistryTool
+ metadata PromptMetadata
+}
+
+func (m *mockPromptMetadataTool) PromptMetadata() PromptMetadata {
+ return m.metadata
+}
+
type mockAsyncRegistryTool struct {
mockRegistryTool
lastCB AsyncCallback
@@ -375,6 +384,47 @@ func TestToolToSchema(t *testing.T) {
}
}
+func TestToolRegistry_ToProviderDefsAttachesPromptMetadata(t *testing.T) {
+ r := NewToolRegistry()
+ r.Register(newMockTool("native", "native tool"))
+ r.Register(&mockPromptMetadataTool{
+ mockRegistryTool: mockRegistryTool{
+ name: "mcp_demo",
+ desc: "mcp tool",
+ params: map[string]any{"type": "object"},
+ },
+ metadata: PromptMetadata{
+ Layer: ToolPromptLayerCapability,
+ Slot: ToolPromptSlotMCP,
+ Source: "mcp:demo",
+ },
+ })
+
+ defs := r.ToProviderDefs()
+ if len(defs) != 2 {
+ t.Fatalf("ToProviderDefs() len = %d, want 2", len(defs))
+ }
+
+ byName := make(map[string]providers.ToolDefinition, len(defs))
+ for _, def := range defs {
+ byName[def.Function.Name] = def
+ }
+
+ native := byName["native"]
+ if native.PromptLayer != ToolPromptLayerCapability ||
+ native.PromptSlot != ToolPromptSlotTooling ||
+ native.PromptSource != ToolPromptSourceRegistry {
+ t.Fatalf("native prompt metadata = %#v, want default tooling source", native)
+ }
+
+ mcp := byName["mcp_demo"]
+ if mcp.PromptLayer != ToolPromptLayerCapability ||
+ mcp.PromptSlot != ToolPromptSlotMCP ||
+ mcp.PromptSource != "mcp:demo" {
+ t.Fatalf("mcp prompt metadata = %#v, want mcp source", mcp)
+ }
+}
+
func TestToolRegistry_Clone(t *testing.T) {
r := NewToolRegistry()
r.Register(newMockTool("read_file", "reads files"))
diff --git a/pkg/tools/search_tool.go b/pkg/tools/search_tool.go
index f41c80d90..c5884c9de 100644
--- a/pkg/tools/search_tool.go
+++ b/pkg/tools/search_tool.go
@@ -34,6 +34,14 @@ 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) PromptMetadata() PromptMetadata {
+ return PromptMetadata{
+ Layer: ToolPromptLayerCapability,
+ Slot: ToolPromptSlotTooling,
+ Source: ToolPromptSourceDiscovery,
+ }
+}
+
func (t *RegexSearchTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
@@ -95,6 +103,14 @@ 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) PromptMetadata() PromptMetadata {
+ return PromptMetadata{
+ Layer: ToolPromptLayerCapability,
+ Slot: ToolPromptSlotTooling,
+ Source: ToolPromptSourceDiscovery,
+ }
+}
+
func (t *BM25SearchTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
diff --git a/pkg/tools/shared/base.go b/pkg/tools/shared/base.go
index 5498d24ab..298e1b478 100644
--- a/pkg/tools/shared/base.go
+++ b/pkg/tools/shared/base.go
@@ -14,6 +14,24 @@ type Tool interface {
Execute(ctx context.Context, args map[string]any) *ToolResult
}
+const (
+ ToolPromptLayerCapability = "capability"
+ ToolPromptSlotTooling = "tooling"
+ ToolPromptSlotMCP = "mcp"
+ ToolPromptSourceRegistry = "tool_registry:native"
+ ToolPromptSourceDiscovery = "tool_registry:discovery"
+)
+
+type PromptMetadata struct {
+ Layer string
+ Slot string
+ Source string
+}
+
+type PromptMetadataProvider interface {
+ PromptMetadata() PromptMetadata
+}
+
// --- Request-scoped tool context (channel / chatID) ---
//
// Carried via context.Value so that concurrent tool calls each receive
diff --git a/pkg/tools/shared_facade.go b/pkg/tools/shared_facade.go
index 6e40e4e3a..8409ea060 100644
--- a/pkg/tools/shared_facade.go
+++ b/pkg/tools/shared_facade.go
@@ -22,12 +22,20 @@ type (
Tool = toolshared.Tool
AsyncCallback = toolshared.AsyncCallback
AsyncExecutor = toolshared.AsyncExecutor
+ PromptMetadata = toolshared.PromptMetadata
+ PromptMetadataProvider = toolshared.PromptMetadataProvider
ToolResult = toolshared.ToolResult
)
const (
handledToolLLMNote = toolshared.HandledToolLLMNote
artifactPathsLLMNote = toolshared.ArtifactPathsLLMNote
+
+ ToolPromptLayerCapability = toolshared.ToolPromptLayerCapability
+ ToolPromptSlotTooling = toolshared.ToolPromptSlotTooling
+ ToolPromptSlotMCP = toolshared.ToolPromptSlotMCP
+ ToolPromptSourceRegistry = toolshared.ToolPromptSourceRegistry
+ ToolPromptSourceDiscovery = toolshared.ToolPromptSourceDiscovery
)
func WithToolContext(ctx context.Context, channel, chatID string) context.Context {
diff --git a/pkg/utils/tool_feedback.go b/pkg/utils/tool_feedback.go
index 1a8b6c747..de7cb467e 100644
--- a/pkg/utils/tool_feedback.go
+++ b/pkg/utils/tool_feedback.go
@@ -7,21 +7,31 @@ import (
const ToolFeedbackContinuationHint = "Continuing the current task."
-// FormatToolFeedbackMessage renders the model-provided explanation for why a
-// tool is being executed. When the model does not provide one, it keeps only
-// the tool line and does not expose raw arguments or fallback text.
-func FormatToolFeedbackMessage(toolName, explanation string) string {
+// FormatToolFeedbackMessage renders a tool feedback message for chat channels.
+// It keeps the tool name on the first line for animation and can include both
+// a human explanation and the serialized tool arguments in the body.
+func FormatToolFeedbackMessage(toolName, explanation, argsPreview string) string {
toolName = strings.TrimSpace(toolName)
explanation = strings.TrimSpace(explanation)
+ argsPreview = strings.TrimSpace(argsPreview)
+
+ bodyLines := make([]string, 0, 2)
+ if explanation != "" {
+ bodyLines = append(bodyLines, explanation)
+ }
+ if argsPreview != "" {
+ bodyLines = append(bodyLines, "```json\n"+argsPreview+"\n```")
+ }
+ body := strings.Join(bodyLines, "\n")
if toolName == "" {
- return explanation
+ return body
}
- if explanation == "" {
+ if body == "" {
return fmt.Sprintf("\U0001f527 `%s`", toolName)
}
- return fmt.Sprintf("\U0001f527 `%s`\n%s", toolName, explanation)
+ return fmt.Sprintf("\U0001f527 `%s`\n%s", toolName, body)
}
// FitToolFeedbackMessage keeps tool feedback within a single outbound message.
diff --git a/pkg/utils/tool_feedback_test.go b/pkg/utils/tool_feedback_test.go
index 316ce2408..c30f53827 100644
--- a/pkg/utils/tool_feedback_test.go
+++ b/pkg/utils/tool_feedback_test.go
@@ -6,29 +6,38 @@ func TestFormatToolFeedbackMessage(t *testing.T) {
got := FormatToolFeedbackMessage(
"read_file",
"I will read README.md first to confirm the current project structure.",
+ "{\n \"path\": \"README.md\"\n}",
)
- want := "\U0001f527 `read_file`\nI will read README.md first to confirm the current project structure."
+ want := "\U0001f527 `read_file`\nI will read README.md first to confirm the current project structure.\n```json\n{\n \"path\": \"README.md\"\n}\n```"
if got != want {
t.Fatalf("FormatToolFeedbackMessage() = %q, want %q", got, want)
}
}
-func TestFormatToolFeedbackMessage_EmptyExplanationKeepsOnlyToolLine(t *testing.T) {
- got := FormatToolFeedbackMessage("read_file", "")
- want := "\U0001f527 `read_file`"
+func TestFormatToolFeedbackMessage_EmptyExplanationShowsArgs(t *testing.T) {
+ got := FormatToolFeedbackMessage("read_file", "", "{\n \"path\": \"README.md\"\n}")
+ want := "\U0001f527 `read_file`\n```json\n{\n \"path\": \"README.md\"\n}\n```"
if got != want {
t.Fatalf("FormatToolFeedbackMessage() = %q, want %q", got, want)
}
}
func TestFormatToolFeedbackMessage_EmptyToolNameOmitsToolLine(t *testing.T) {
- got := FormatToolFeedbackMessage("", "Continue drafting the final response.")
+ got := FormatToolFeedbackMessage("", "Continue drafting the final response.", "")
want := "Continue drafting the final response."
if got != want {
t.Fatalf("FormatToolFeedbackMessage() = %q, want %q", got, want)
}
}
+func TestFormatToolFeedbackMessage_EmptyExplanationAndArgsKeepsOnlyToolLine(t *testing.T) {
+ got := FormatToolFeedbackMessage("read_file", "", "")
+ want := "\U0001f527 `read_file`"
+ if got != want {
+ t.Fatalf("FormatToolFeedbackMessage() = %q, want %q", got, want)
+ }
+}
+
func TestFitToolFeedbackMessage_TruncatesBodyWithinSingleMessage(t *testing.T) {
got := FitToolFeedbackMessage(
"\U0001f527 `read_file`\nRead README.md first to confirm the current project structure.",
diff --git a/scripts/copydir.go b/scripts/copydir.go
new file mode 100644
index 000000000..6e2777612
--- /dev/null
+++ b/scripts/copydir.go
@@ -0,0 +1,186 @@
+package main
+
+import (
+ "fmt"
+ "io"
+ "os"
+ "path/filepath"
+ "runtime"
+ "strings"
+)
+
+func main() {
+ if len(os.Args) != 3 {
+ fmt.Fprintf(os.Stderr, "usage: go run scripts/copydir.go