fix(security): remove exec context race and harden ipv6 SSRF
This commit is contained in:
parent
e009611f88
commit
441824a0db
6 changed files with 43 additions and 46 deletions
|
|
@ -747,16 +747,6 @@ func (al *AgentLoop) updateToolContexts(agent *AgentInstance, channel, chatID st
|
|||
st.SetContext(channel, chatID)
|
||||
}
|
||||
}
|
||||
if tool, ok := al.tools.Get("exec"); ok {
|
||||
if et, ok := tool.(tools.ContextualTool); ok {
|
||||
et.SetContext(channel, chatID)
|
||||
}
|
||||
}
|
||||
if tool, ok := al.tools.Get("cron"); ok {
|
||||
if ct, ok := tool.(tools.ContextualTool); ok {
|
||||
ct.SetContext(channel, chatID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// maybeSummarize triggers summarization if the session history exceeds thresholds.
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ package tools
|
|||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
|
|
@ -24,9 +23,6 @@ type CronTool struct {
|
|||
executor JobExecutor
|
||||
msgBus *bus.MessageBus
|
||||
execTool *ExecTool
|
||||
channel string
|
||||
chatID string
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewCronTool creates a new CronTool
|
||||
|
|
@ -102,14 +98,6 @@ func (t *CronTool) Parameters() map[string]any {
|
|||
}
|
||||
}
|
||||
|
||||
// SetContext sets the current session context for job creation
|
||||
func (t *CronTool) SetContext(channel, chatID string) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
t.channel = channel
|
||||
t.chatID = chatID
|
||||
}
|
||||
|
||||
// Execute runs the tool with the given arguments
|
||||
func (t *CronTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||
action, ok := args["action"].(string)
|
||||
|
|
@ -134,10 +122,8 @@ func (t *CronTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
|||
}
|
||||
|
||||
func (t *CronTool) addJob(args map[string]any) *ToolResult {
|
||||
t.mu.RLock()
|
||||
channel := t.channel
|
||||
chatID := t.chatID
|
||||
t.mu.RUnlock()
|
||||
channel, _ := args["__channel"].(string)
|
||||
chatID, _ := args["__chat_id"].(string)
|
||||
|
||||
if channel == "" || chatID == "" {
|
||||
return ErrorResult("no session context (channel/chat_id not set). Use this tool in an active conversation.")
|
||||
|
|
|
|||
|
|
@ -77,8 +77,19 @@ func (r *ToolRegistry) ExecuteWithContext(
|
|||
})
|
||||
}
|
||||
|
||||
toolArgs := args
|
||||
if channel != "" || chatID != "" {
|
||||
toolArgs = make(map[string]interface{}, len(args)+2)
|
||||
for k, v := range args {
|
||||
toolArgs[k] = v
|
||||
}
|
||||
// Internal runtime context for auth/policy checks in tools.
|
||||
toolArgs["__channel"] = channel
|
||||
toolArgs["__chat_id"] = chatID
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
result := tool.Execute(ctx, args)
|
||||
result := tool.Execute(ctx, toolArgs)
|
||||
duration := time.Since(start)
|
||||
|
||||
// Log based on result type
|
||||
|
|
|
|||
|
|
@ -24,8 +24,6 @@ type ExecTool struct {
|
|||
allowPatterns []*regexp.Regexp
|
||||
restrictToWorkspace bool
|
||||
allowRemote bool
|
||||
channel string
|
||||
chatID string
|
||||
}
|
||||
|
||||
var defaultDenyPatterns = []*regexp.Regexp{
|
||||
|
|
@ -149,7 +147,8 @@ func (t *ExecTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
|||
}
|
||||
|
||||
if !t.allowRemote {
|
||||
channel := strings.TrimSpace(t.channel)
|
||||
channel, _ := args["__channel"].(string)
|
||||
channel = strings.TrimSpace(channel)
|
||||
if channel == "" || !constants.IsInternalChannel(channel) {
|
||||
return ErrorResult("exec is restricted to internal channels")
|
||||
}
|
||||
|
|
@ -345,8 +344,3 @@ func (t *ExecTool) SetAllowPatterns(patterns []string) error {
|
|||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *ExecTool) SetContext(channel, chatID string) {
|
||||
t.channel = channel
|
||||
t.chatID = chatID
|
||||
}
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import (
|
|||
"net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
|
|
@ -498,7 +499,7 @@ type WebFetchTool struct {
|
|||
|
||||
// allowPrivateWebFetchHosts controls whether loopback/private hosts are allowed.
|
||||
// This is false in normal runtime to reduce SSRF exposure, and tests can override it temporarily.
|
||||
var allowPrivateWebFetchHosts = false
|
||||
var allowPrivateWebFetchHosts atomic.Bool
|
||||
|
||||
func NewWebFetchTool(maxChars int) *WebFetchTool {
|
||||
if maxChars <= 0 {
|
||||
|
|
@ -711,7 +712,7 @@ func (t *WebFetchTool) extractText(htmlContent string) string {
|
|||
|
||||
func newSafeDialContext(dialer *net.Dialer) func(context.Context, string, string) (net.Conn, error) {
|
||||
return func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
if allowPrivateWebFetchHosts {
|
||||
if allowPrivateWebFetchHosts.Load() {
|
||||
return dialer.DialContext(ctx, network, address)
|
||||
}
|
||||
|
||||
|
|
@ -760,11 +761,12 @@ func newSafeDialContext(dialer *net.Dialer) func(context.Context, string, string
|
|||
}
|
||||
|
||||
func isPrivateFetchHost(host string) bool {
|
||||
if allowPrivateWebFetchHosts {
|
||||
if allowPrivateWebFetchHosts.Load() {
|
||||
return false
|
||||
}
|
||||
|
||||
canonicalHost := strings.ToLower(strings.TrimSpace(host))
|
||||
canonicalHost = strings.TrimSuffix(canonicalHost, ".")
|
||||
if canonicalHost == "" {
|
||||
return true
|
||||
}
|
||||
|
|
@ -816,6 +818,20 @@ func isPrivateOrRestrictedIP(ip net.IP) bool {
|
|||
return false
|
||||
}
|
||||
|
||||
if len(ip) == net.IPv6len {
|
||||
// IPv6 unique local addresses (fc00::/7)
|
||||
return len(ip) == net.IPv6len && (ip[0]&0xfe) == 0xfc
|
||||
if (ip[0] & 0xfe) == 0xfc {
|
||||
return true
|
||||
}
|
||||
// 6to4 addresses (2002::/16) can embed private IPv4 targets.
|
||||
if ip[0] == 0x20 && ip[1] == 0x02 {
|
||||
return true
|
||||
}
|
||||
// Teredo tunneling addresses (2001:0000::/32) can encapsulate private endpoints.
|
||||
if ip[0] == 0x20 && ip[1] == 0x01 && ip[2] == 0x00 && ip[3] == 0x00 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
|
|
|||
|
|
@ -324,10 +324,10 @@ func TestWebFetchTool_extractText(t *testing.T) {
|
|||
|
||||
func withPrivateWebFetchHostsAllowed(t *testing.T) {
|
||||
t.Helper()
|
||||
previous := allowPrivateWebFetchHosts
|
||||
allowPrivateWebFetchHosts = true
|
||||
previous := allowPrivateWebFetchHosts.Load()
|
||||
allowPrivateWebFetchHosts.Store(true)
|
||||
t.Cleanup(func() {
|
||||
allowPrivateWebFetchHosts = previous
|
||||
allowPrivateWebFetchHosts.Store(previous)
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue