diff --git a/pkg/selfupdate/exec_windows.go b/pkg/selfupdate/exec_windows.go index e5098d1af0..3459dcc7ed 100644 --- a/pkg/selfupdate/exec_windows.go +++ b/pkg/selfupdate/exec_windows.go @@ -3,6 +3,7 @@ package selfupdate import ( + "errors" "fmt" "os" "os/exec" @@ -48,7 +49,7 @@ func reExecProcess(path string, args, env []string) error { childArgs = args[1:] } - cmd := exec.Command(path, childArgs...) //nolint:noctx // path is our own freshly installed binary; no context needed for re-exec + cmd := exec.Command(path, childArgs...) //nolint:noctx // re-exec must outlive any request-scoped context cmd.Env = env cmd.Stdin = os.Stdin cmd.Stdout = os.Stdout @@ -56,7 +57,7 @@ func reExecProcess(path string, args, env []string) error { if err := cmd.Run(); err != nil { var exitErr *exec.ExitError - if ok := asExitError(err, &exitErr); ok { + if errors.As(err, &exitErr) { os.Exit(exitErr.ExitCode()) } return fmt.Errorf("running updated binary: %w", err) @@ -65,13 +66,3 @@ func reExecProcess(path string, args, env []string) error { os.Exit(0) return nil } - -// asExitError is a tiny helper kept separate so exec_unix.go does not need to -// import errors solely for this Windows branch. -func asExitError(err error, target **exec.ExitError) bool { - if e, ok := err.(*exec.ExitError); ok { //nolint:errorlint // direct type assertion is intentional here - *target = e - return true - } - return false -} diff --git a/pkg/tools/builtin/backgroundjobs/cmd_windows.go b/pkg/tools/builtin/backgroundjobs/cmd_windows.go index 0196b62a4e..2cbb43fa86 100644 --- a/pkg/tools/builtin/backgroundjobs/cmd_windows.go +++ b/pkg/tools/builtin/backgroundjobs/cmd_windows.go @@ -31,13 +31,13 @@ func createProcessGroup(proc *os.Process) (*processGroup, error) { if _, err := windows.SetInformationJobObject( job, windows.JobObjectExtendedLimitInformation, - uintptr(unsafe.Pointer(&info)), //nolint:gosec // Windows API requires unsafe pointer + uintptr(unsafe.Pointer(&info)), //nolint:gosec // Windows API requires unsafe.Pointer uint32(unsafe.Sizeof(info))); err != nil { _ = windows.CloseHandle(job) return nil, err } - handle, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(proc.Pid)) //nolint:gosec // Pid is safe to convert to uint32 on Windows + handle, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(proc.Pid)) //nolint:gosec // proc.Pid fits in uint32 on Windows if err != nil { _ = windows.CloseHandle(job) return nil, err diff --git a/pkg/tools/builtin/rag/rag.go b/pkg/tools/builtin/rag/rag.go index 8ca720492a..296cfd69c7 100644 --- a/pkg/tools/builtin/rag/rag.go +++ b/pkg/tools/builtin/rag/rag.go @@ -8,6 +8,7 @@ import ( "fmt" "log/slog" "slices" + "sync" "github.com/docker/docker-agent/pkg/config" "github.com/docker/docker-agent/pkg/config/latest" @@ -49,6 +50,8 @@ type ToolSet struct { manager *rag.Manager toolName string eventCallback EventCallback + cancelWatcher context.CancelFunc + wg sync.WaitGroup } // Verify interface compliance. @@ -84,20 +87,29 @@ func (t *ToolSet) Start(ctx context.Context) error { return nil } + // We create a child context so we can explicitly cancel the watcher and event goroutines + // when Stop() is called, preventing goroutine leaks if the parent context outlives this toolset. + watchCtx, cancel := context.WithCancel(ctx) + t.cancelWatcher = cancel + // Forward RAG manager events if a callback is set. if t.eventCallback != nil { - go t.forwardEvents(ctx) + t.wg.Go(func() { + t.forwardEvents(watchCtx) + }) } if err := t.manager.Initialize(ctx); err != nil { + cancel() + t.wg.Wait() return fmt.Errorf("failed to initialize RAG manager %q: %w", t.toolName, err) } - go func() { - if err := t.manager.StartFileWatcher(ctx); err != nil { - slog.ErrorContext(ctx, "Failed to start RAG file watcher", "tool", t.toolName, "error", err) + t.wg.Go(func() { + if err := t.manager.StartFileWatcher(watchCtx); err != nil && !errors.Is(err, context.Canceled) { + slog.ErrorContext(watchCtx, "Failed to start RAG file watcher", "tool", t.toolName, "error", err) } - }() + }) return nil } @@ -106,6 +118,10 @@ func (t *ToolSet) Stop(_ context.Context) error { if t.manager == nil { return nil } + if t.cancelWatcher != nil { + t.cancelWatcher() + } + t.wg.Wait() return t.manager.Close() } diff --git a/pkg/tools/builtin/rag/rag_test.go b/pkg/tools/builtin/rag/rag_test.go index 69e2d2996d..6666f467b5 100644 --- a/pkg/tools/builtin/rag/rag_test.go +++ b/pkg/tools/builtin/rag/rag_test.go @@ -5,6 +5,7 @@ import ( "context" "slices" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -143,3 +144,46 @@ func TestRAGTool_HandleQuery_Telemetry(t *testing.T) { require.NoError(t, err) assert.NotNil(t, res) } + +type failingMockStrategy struct { + mockStrategy +} + +func (m *failingMockStrategy) Initialize(_ context.Context, _ []string, _ strategy.ChunkingConfig) error { + return assert.AnError +} + +func TestStopAfterFailedStart(t *testing.T) { + t.Parallel() + + strategyMock := &failingMockStrategy{} + cfg := rag.Config{ + StrategyConfigs: []strategy.Config{ + {Name: "failingStrategy", Strategy: strategyMock}, + }, + } + + mgr, err := rag.New(t.Context(), "failing-rag", cfg, nil) + require.NoError(t, err) + + tool := &ToolSet{ + manager: mgr, + toolName: "failing-rag", + } + + err = tool.Start(t.Context()) + require.Error(t, err) + + done := make(chan struct{}) + go func() { + _ = tool.Stop(t.Context()) + close(done) + }() + + select { + case <-done: + // Success: Stop returned without deadlocking + case <-time.After(5 * time.Second): + t.Fatal("Stop() deadlocked after a failed Start()") + } +} diff --git a/pkg/tools/builtin/shell/cmd_windows.go b/pkg/tools/builtin/shell/cmd_windows.go index 03908fcc04..85d4145786 100644 --- a/pkg/tools/builtin/shell/cmd_windows.go +++ b/pkg/tools/builtin/shell/cmd_windows.go @@ -31,13 +31,13 @@ func createProcessGroup(proc *os.Process) (*processGroup, error) { if _, err := windows.SetInformationJobObject( job, windows.JobObjectExtendedLimitInformation, - uintptr(unsafe.Pointer(&info)), //nolint:gosec // Windows API requires unsafe pointer + uintptr(unsafe.Pointer(&info)), //nolint:gosec // Windows API requires unsafe.Pointer uint32(unsafe.Sizeof(info))); err != nil { _ = windows.CloseHandle(job) return nil, err } - handle, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(proc.Pid)) //nolint:gosec // Pid is safe to convert to uint32 on Windows + handle, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(proc.Pid)) //nolint:gosec // proc.Pid fits in uint32 on Windows if err != nil { _ = windows.CloseHandle(job) return nil, err diff --git a/pkg/tools/builtin/shell/script_shell.go b/pkg/tools/builtin/shell/script_shell.go index efb3176299..93ef4f6fbf 100644 --- a/pkg/tools/builtin/shell/script_shell.go +++ b/pkg/tools/builtin/shell/script_shell.go @@ -52,6 +52,12 @@ var ( ) func NewScript(shellTools map[string]latest.ScriptShellToolConfig, env []string) (*ScriptToolSet, error) { + for _, e := range env { + if strings.ContainsRune(e, 0) { + return nil, errors.New("toolset environment contains a NUL byte") + } + } + for toolName, tool := range shellTools { if err := validateConfig(toolName, tool); err != nil { return nil, err @@ -248,7 +254,11 @@ func (t *ScriptToolSet) execute(ctx context.Context, rt tools.Runtime, toolConfi // stay literal because env values may legitimately contain $ (issue // #2615). for _, key := range slices.Sorted(maps.Keys(toolConfig.Env)) { - envCopy = append(envCopy, key+"="+path.ExpandEnvRefs(toolConfig.Env[key])) + val := path.ExpandEnvRefs(toolConfig.Env[key]) + if strings.ContainsRune(val, 0) { + return tools.ResultError(fmt.Sprintf("configured environment variable %q contains a NUL byte", key)), nil + } + envCopy = append(envCopy, key+"="+val) } for key, value := range params { if value == nil { @@ -262,9 +272,8 @@ func (t *ScriptToolSet) execute(ctx context.Context, rt tools.Runtime, toolConfi continue } valueStr := fmt.Sprintf("%v", value) - // A NUL byte mid-string silently truncates env entries at the - // execve boundary; refuse rather than spawn a process with a - // surprising env. + // Go's os/exec rejects NUL bytes with generic errors. We check here + // to provide a clearer error message. if strings.ContainsRune(valueStr, 0) { return tools.ResultError(fmt.Sprintf("argument %q contains a NUL byte", key)), nil }