From e70e440ea5bcc360e4c3a229564aac6e623b02e2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Arnaud=20He=CC=81ritier?= Date: Fri, 31 Jul 2026 07:55:41 +0200 Subject: [PATCH 1/3] feat(#3996): persist absolute WorkingDir provenance at every local creation path Every session created against a local workspace now records the absolute workspace root that owns it, so generated files can later be resolved against the owning workspace instead of whatever cwd the viewer happens to run from: - delegated/nested sub-sessions inherit parent.WorkingDir - serve mcp / a2a / chatserver capture the server workspace once at startup instead of reading os.Getwd per request (or not at all) - embeddedchat conversations and `docker-agent new` capture the configured root or the creation-time cwd - --working-dir writes the absolutized path back into RuntimeConfig so a relative flag value cannot leak into persisted provenance - compaction runs inherit the compacted session's workspace Remote/headless surfaces (remote backend template, API templates without working_dir, ACP without client cwd, evaluation container transcripts) intentionally keep an empty WorkingDir and now document that choice. session.ResolveWorkingDir adds a bounded, cycle-safe, validated parent-chain fallback for old persisted sub-sessions with an empty WorkingDir; it rejects relative/malformed roots and never falls back to the process cwd. It is the primitive the generated-media resolver will build on. --- cmd/root/backend.go | 3 + cmd/root/flags.go | 14 +- cmd/root/flags_test.go | 28 +++ cmd/root/new.go | 7 + cmd/root/run.go | 8 + e2e/mcp_test.go | 4 +- pkg/a2a/adapter.go | 8 +- pkg/a2a/adapter_run_test.go | 18 +- pkg/a2a/adapter_test.go | 4 +- pkg/a2a/server.go | 11 +- pkg/a2a/server_invoke_test.go | 2 +- pkg/a2a/server_test.go | 2 +- pkg/acp/agent.go | 2 + pkg/chatserver/agent.go | 3 +- pkg/chatserver/server.go | 18 +- pkg/chatserver/server_test.go | 19 +- pkg/embeddedchat/embeddedchat.go | 13 +- pkg/embeddedchat/embeddedchat_test.go | 31 +++ pkg/evaluation/save.go | 3 +- pkg/mcp/server.go | 40 ++-- pkg/mcp/server_test.go | 16 ++ pkg/runtime/agent_delegation.go | 4 + pkg/runtime/agent_delegation_test.go | 38 ++++ pkg/runtime/compactor/compactor.go | 3 + pkg/runtime/compactor/compactor_test.go | 31 +++ pkg/server/session_manager.go | 3 + pkg/session/working_dir.go | 112 ++++++++++ pkg/session/working_dir_test.go | 282 ++++++++++++++++++++++++ 28 files changed, 675 insertions(+), 52 deletions(-) create mode 100644 pkg/session/working_dir.go create mode 100644 pkg/session/working_dir_test.go diff --git a/cmd/root/backend.go b/cmd/root/backend.go index be0aa0c85c..f63a180e19 100644 --- a/cmd/root/backend.go +++ b/cmd/root/backend.go @@ -213,6 +213,9 @@ func (b *remoteBackend) CreateSession(ctx context.Context, _ *teamloader.LoadRes sessTemplate := session.New( session.WithToolsApproved(req.ToolsApproved), session.WithSafetyPolicy(req.SafetyPolicy), + // WorkingDir is intentionally not sent: the client checkout is not the + // remote server's workspace. The server establishes its own workspace + // provenance for the session it creates (see server.SessionManager). ) sess, err := client.CreateSession(ctx, sessTemplate) diff --git a/cmd/root/flags.go b/cmd/root/flags.go index e08b5d857e..7cbb714270 100644 --- a/cmd/root/flags.go +++ b/cmd/root/flags.go @@ -59,12 +59,17 @@ func sessionDBPath(flagValue string) string { return filepath.Join(paths.GetDataDir(), "session.db") } -func setupWorkingDirectory(workingDir string) error { - if workingDir == "" { +// setupWorkingDirectory applies a --working-dir override: it chdirs into the +// directory and stores the absolute path back into runConfig.WorkingDir so +// downstream consumers (session workspace provenance, toolsets, servers) see +// an absolute root — after the chdir a relative flag value would otherwise +// resolve against itself. +func setupWorkingDirectory(runConfig *config.RuntimeConfig) error { + if runConfig.WorkingDir == "" { return nil } - absWd, err := filepath.Abs(workingDir) + absWd, err := filepath.Abs(runConfig.WorkingDir) if err != nil { return fmt.Errorf("invalid working directory: %w", err) } @@ -79,6 +84,7 @@ func setupWorkingDirectory(workingDir string) error { } _ = os.Setenv("PWD", absWd) + runConfig.WorkingDir = absWd slog.Debug("Working directory set", "path", absWd) return nil @@ -164,7 +170,7 @@ func addGatewayFlags(cmd *cobra.Command, runConfig *config.RuntimeConfig, loadUs runConfig.Providers = userCfg.GetProviders() } - return setupWorkingDirectory(runConfig.WorkingDir) + return setupWorkingDirectory(runConfig) } } diff --git a/cmd/root/flags_test.go b/cmd/root/flags_test.go index e13ae25030..22bd2f4aae 100644 --- a/cmd/root/flags_test.go +++ b/cmd/root/flags_test.go @@ -409,3 +409,31 @@ func TestEnvFromFileErrorsAbortPreRun(t *testing.T) { assert.Contains(t, err.Error(), "bad.env") }) } + +// TestSetupWorkingDirectory_StoresAbsolutePath pins that a relative +// --working-dir value is absolutized back into the runtime config: after the +// chdir, downstream consumers (session workspace provenance, servers) must +// see the absolute root, not a path relative to itself. +func TestSetupWorkingDirectory_StoresAbsolutePath(t *testing.T) { + base := t.TempDir() + base, err := filepath.EvalSymlinks(base) + require.NoError(t, err) + require.NoError(t, os.Mkdir(filepath.Join(base, "sub"), 0o755)) + t.Chdir(base) + + runConfig := &config.RuntimeConfig{Config: config.Config{WorkingDir: "sub"}} + require.NoError(t, setupWorkingDirectory(runConfig)) + + assert.Equal(t, filepath.Join(base, "sub"), runConfig.WorkingDir) + cwd, err := os.Getwd() + require.NoError(t, err) + cwd, err = filepath.EvalSymlinks(cwd) + require.NoError(t, err) + assert.Equal(t, filepath.Join(base, "sub"), cwd) +} + +func TestSetupWorkingDirectory_EmptyIsNoop(t *testing.T) { + runConfig := &config.RuntimeConfig{} + require.NoError(t, setupWorkingDirectory(runConfig)) + assert.Empty(t, runConfig.WorkingDir) +} diff --git a/cmd/root/new.go b/cmd/root/new.go index f11e9795a2..b5ab138ed6 100644 --- a/cmd/root/new.go +++ b/cmd/root/new.go @@ -78,10 +78,17 @@ func (f *newFlags) runNewCommand(cmd *cobra.Command, args []string) (commandErr } var appOpts []app.Opt + // The creator runs in the user's checkout and writes the generated agent + // YAML there; capture that workspace as the session's provenance. + workingDir, err := session.CaptureLocalWorkingDir(f.runConfig.WorkingDir) + if err != nil { + return err + } sessOpts := []session.Opt{ session.WithTitle("New agent"), session.WithMaxIterations(f.maxIterationsParam), session.WithToolsApproved(true), + session.WithWorkingDir(workingDir), } if len(args) > 0 { arg := strings.Join(args, " ") diff --git a/cmd/root/run.go b/cmd/root/run.go index 306b0b22bc..2830af750f 100644 --- a/cmd/root/run.go +++ b/cmd/root/run.go @@ -1272,6 +1272,14 @@ func (f *runExecFlags) scopedSafetyDefault(safety latestcfg.SafetyMode, legacyYo // createSessionSpawner creates a function that can spawn new sessions with different working directories. func (f *runExecFlags) createSessionSpawner(agentSource config.Source, sessStore session.Store) tui.SessionSpawner { return func(spawnCtx context.Context, workingDir string) (*app.App, *session.Session, func(), error) { + // The spawn dialog may hand us a relative or empty path; pin the + // spawned session's workspace provenance to an absolute root now, + // before anything below captures it. + workingDir, err := session.CaptureLocalWorkingDir(workingDir) + if err != nil { + return nil, nil, nil, err + } + // Create a copy of the runtime config with the new working directory runConfigCopy := f.runConfig.Clone() runConfigCopy.WorkingDir = workingDir diff --git a/e2e/mcp_test.go b/e2e/mcp_test.go index c742ed0781..dee9f87293 100644 --- a/e2e/mcp_test.go +++ b/e2e/mcp_test.go @@ -27,7 +27,7 @@ func TestMCP_SingleAgent(t *testing.T) { require.NoError(t, team.StopToolSets(ctx)) }) - handler := mcp.CreateToolHandler(team, "root", session.SafetyPolicyAutonomous) + handler := mcp.CreateToolHandler(team, "root", session.SafetyPolicyAutonomous, t.TempDir()) _, output, err := handler(ctx, nil, mcp.ToolInput{ Message: "What is 2+2? Answer in one sentence.", }) @@ -50,7 +50,7 @@ func TestMCP_MultiAgent(t *testing.T) { require.NoError(t, team.StopToolSets(ctx)) }) - handler := mcp.CreateToolHandler(team, "web", session.SafetyPolicyAutonomous) + handler := mcp.CreateToolHandler(team, "web", session.SafetyPolicyAutonomous, t.TempDir()) _, output, err := handler(ctx, nil, mcp.ToolInput{ Message: "Say hello in one sentence.", }) diff --git a/pkg/a2a/adapter.go b/pkg/a2a/adapter.go index 8f0663b80c..ab39c87603 100644 --- a/pkg/a2a/adapter.go +++ b/pkg/a2a/adapter.go @@ -6,7 +6,6 @@ import ( "fmt" "iter" "log/slog" - "os" "strings" "go.opentelemetry.io/otel" @@ -29,7 +28,7 @@ import ( // newDockerAgentAdapter creates a new ADK agent adapter from a docker agent team and agent name. // When agentName is empty, the team's default agent (one explicitly named "root" if it // exists, otherwise the first agent declared) is used. -func newDockerAgentAdapter(t *team.Team, agentName string, sessStore session.Store, safety servesafety.Resolved) (agent.Agent, error) { +func newDockerAgentAdapter(t *team.Team, agentName string, sessStore session.Store, safety servesafety.Resolved, workingDir string) (agent.Agent, error) { a, err := t.AgentOrDefault(agentName) if err != nil { return nil, fmt.Errorf("failed to get agent %s: %w", agentName, err) @@ -42,13 +41,13 @@ func newDockerAgentAdapter(t *team.Team, agentName string, sessStore session.Sto Name: agentName, Description: desc, Run: func(ctx agent.InvocationContext) iter.Seq2[*adksession.Event, error] { - return runDockerAgent(ctx, t, agentName, a, sessStore, safety) + return runDockerAgent(ctx, t, agentName, a, sessStore, safety, workingDir) }, }) } // runDockerAgent executes a docker agent and returns ADK session events -func runDockerAgent(ctx agent.InvocationContext, t *team.Team, agentName string, a *dagent.Agent, sessStore session.Store, safety servesafety.Resolved) iter.Seq2[*adksession.Event, error] { +func runDockerAgent(ctx agent.InvocationContext, t *team.Team, agentName string, a *dagent.Agent, sessStore session.Store, safety servesafety.Resolved, workingDir string) iter.Seq2[*adksession.Event, error] { return func(yield func(*adksession.Event, error) bool) { // Decorate the inbound `a2a.message` SERVER span (created by // otelhttp.NewHandler in server.go) with the GenAI semconv @@ -94,7 +93,6 @@ func runDockerAgent(ctx agent.InvocationContext, t *team.Team, agentName string, yield(nil, fmt.Errorf("check A2A context ID: %w", err)) return default: - workingDir, _ := os.Getwd() sess = session.New( session.WithID(sessionID), session.WithOrigin("a2a"), diff --git a/pkg/a2a/adapter_run_test.go b/pkg/a2a/adapter_run_test.go index a8f3600d60..e5aefcbede 100644 --- a/pkg/a2a/adapter_run_test.go +++ b/pkg/a2a/adapter_run_test.go @@ -4,7 +4,6 @@ import ( "context" "errors" "io" - "os" "slices" "sync" "sync/atomic" @@ -175,6 +174,8 @@ func (s *recordingStore) updatedSessions() []*session.Session { return slices.Clone(s.updated) } +const testWorkspaceRoot = "/srv/a2a-workspace" + type yieldedEvent struct { event *adksession.Event err error @@ -182,7 +183,7 @@ type yieldedEvent struct { func collectRunEvents(ctx agent.InvocationContext, tm *team.Team, a *dagent.Agent, store session.Store, policy session.SafetyPolicy) []yieldedEvent { var out []yieldedEvent - for ev, err := range runDockerAgent(ctx, tm, a.Name(), a, store, servesafety.Resolved{Policy: policy}) { + for ev, err := range runDockerAgent(ctx, tm, a.Name(), a, store, servesafety.Resolved{Policy: policy}, testWorkspaceRoot) { out = append(out, yieldedEvent{event: ev, err: err}) } return out @@ -272,7 +273,7 @@ func TestRunDockerAgent_ConsumerStopsEarly(t *testing.T) { ctx := newFakeInvocationContext(t.Context(), "a2a-ctx-early-stop", "Hi") var events []*adksession.Event - for ev, err := range runDockerAgent(ctx, tm, root.Name(), root, store, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}) { + for ev, err := range runDockerAgent(ctx, tm, root.Name(), root, store, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}, testWorkspaceRoot) { require.NoError(t, err) events = append(events, ev) break @@ -291,7 +292,7 @@ func TestRunDockerAgent_EndedInvocationStopsIteration(t *testing.T) { ctx := newFakeInvocationContext(t.Context(), "a2a-ctx-ended", "Hi") var events []*adksession.Event - for ev, err := range runDockerAgent(ctx, tm, root.Name(), root, store, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}) { + for ev, err := range runDockerAgent(ctx, tm, root.Name(), root, store, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}, testWorkspaceRoot) { require.NoError(t, err) events = append(events, ev) // Ending the invocation after the first chunk must stop the @@ -324,12 +325,9 @@ func TestRunDockerAgent_NewSessionUsesA2ASettings(t *testing.T) { assert.False(t, sess.ToolsApproved) assert.True(t, sess.NonInteractive) - // runDockerAgent stamps new sessions with the process working directory - // via os.Getwd, so this assertion resolves the same value and relies on - // nothing in the test process changing directories. - workingDir, err := os.Getwd() - require.NoError(t, err) - assert.Equal(t, workingDir, sess.WorkingDir) + // runDockerAgent receives the server workspace at startup, so tests use a + // fixed value rather than reading the process working directory. + assert.Equal(t, testWorkspaceRoot, sess.WorkingDir) msgs := sess.GetAllMessages() require.NotEmpty(t, msgs) diff --git a/pkg/a2a/adapter_test.go b/pkg/a2a/adapter_test.go index 88298b6b78..ac477189a2 100644 --- a/pkg/a2a/adapter_test.go +++ b/pkg/a2a/adapter_test.go @@ -26,7 +26,7 @@ func TestNewDockerAgentAdapter(t *testing.T) { require.NoError(t, team.StopToolSets(t.Context())) }() - adapter, err := newDockerAgentAdapter(team, "root", nil, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}) + adapter, err := newDockerAgentAdapter(team, "root", nil, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}, "/srv/a2a-workspace") require.NoError(t, err) assert.Equal(t, "root", adapter.Name()) @@ -45,7 +45,7 @@ func TestNewCAgentAdapter_NonExistent(t *testing.T) { require.NoError(t, team.StopToolSets(t.Context())) }() - _, err = newDockerAgentAdapter(team, "nonexistent", nil, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}) + _, err = newDockerAgentAdapter(team, "nonexistent", nil, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}, "/srv/a2a-workspace") assert.Contains(t, err.Error(), "failed to get agent") } diff --git a/pkg/a2a/server.go b/pkg/a2a/server.go index 4282987ff6..08e924479b 100644 --- a/pkg/a2a/server.go +++ b/pkg/a2a/server.go @@ -98,7 +98,12 @@ func Run(ctx context.Context, agentFilename, agentName, sessionDB string, runCon baseURL := &url.URL{Scheme: "http", Host: routableAddr(ln.Addr().String())} slog.DebugContext(ctx, "A2A server listening", "url", baseURL.String()) - e, err := newServer(t, agentFilename, agentName, sessStore, resolvedSafety, ln.Addr().String(), options) + workingDir, err := session.CaptureLocalWorkingDir(runConfig.WorkingDir) + if err != nil { + return err + } + + e, err := newServer(t, agentFilename, agentName, sessStore, resolvedSafety, workingDir, ln.Addr().String(), options) if err != nil { return fmt.Errorf("failed to create A2A server: %w", err) } @@ -118,8 +123,8 @@ func Run(ctx context.Context, agentFilename, agentName, sessionDB string, runCon return nil } -func newServer(t *team.Team, agentFilename, agentName string, sessStore session.Store, safety servesafety.Resolved, listenAddr string, options RunOptions) (*echo.Echo, error) { - adkAgent, err := newDockerAgentAdapter(t, agentName, sessStore, safety) +func newServer(t *team.Team, agentFilename, agentName string, sessStore session.Store, safety servesafety.Resolved, workingDir, listenAddr string, options RunOptions) (*echo.Echo, error) { + adkAgent, err := newDockerAgentAdapter(t, agentName, sessStore, safety, workingDir) if err != nil { return nil, err } diff --git a/pkg/a2a/server_invoke_test.go b/pkg/a2a/server_invoke_test.go index 9e24e43a32..303299c439 100644 --- a/pkg/a2a/server_invoke_test.go +++ b/pkg/a2a/server_invoke_test.go @@ -164,7 +164,7 @@ func startInvokeServer(t *testing.T, tm *team.Team, store session.Store, safety ln, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", "127.0.0.1:0") require.NoError(t, err) - e, err := newServer(tm, "test.yaml", "root", store, safety, ln.Addr().String(), RunOptions{}) + e, err := newServer(tm, "test.yaml", "root", store, safety, testWorkspaceRoot, ln.Addr().String(), RunOptions{}) require.NoError(t, err) go func() { _ = e.Server.Serve(ln) }() t.Cleanup(func() { require.NoError(t, e.Server.Close()) }) diff --git a/pkg/a2a/server_test.go b/pkg/a2a/server_test.go index 6aa0b36221..13a9d863fb 100644 --- a/pkg/a2a/server_test.go +++ b/pkg/a2a/server_test.go @@ -109,7 +109,7 @@ func TestServerSecurity(t *testing.T) { tm := team.New(team.WithAgents(agent.New("root", "test"))) store := session.NewInMemorySessionStore() - server, err := newServer(tm, "test.yaml", "root", store, servesafety.Resolved{}, "127.0.0.1:0", RunOptions{AuthToken: "secret", CORSOrigin: "https://app.example.com"}) + server, err := newServer(tm, "test.yaml", "root", store, servesafety.Resolved{}, testWorkspaceRoot, "127.0.0.1:0", RunOptions{AuthToken: "secret", CORSOrigin: "https://app.example.com"}) require.NoError(t, err) request := func(method, path string, headers map[string]string) *httptest.ResponseRecorder { diff --git a/pkg/acp/agent.go b/pkg/acp/agent.go index f23468ba99..8d3df8304f 100644 --- a/pkg/acp/agent.go +++ b/pkg/acp/agent.go @@ -193,6 +193,8 @@ func (a *Agent) NewSession(ctx context.Context, params acp.NewSessionRequest) (a // An empty cwd is allowed: clients (e.g. zed) may not always supply a // working directory at session creation. We persist it as empty and // later prompts/tools fall back to the agent's default working dir. + // The persisted WorkingDir stays empty too: workspace provenance must + // come from the client, never be inferred from the server's process cwd. if err := validateWorkingDir(workingDir); err != nil { return acp.NewSessionResponse{}, err } diff --git a/pkg/chatserver/agent.go b/pkg/chatserver/agent.go index 1d33542c10..2172abccf9 100644 --- a/pkg/chatserver/agent.go +++ b/pkg/chatserver/agent.go @@ -64,9 +64,10 @@ func (p agentPolicy) pick(model string) string { // // Returns nil when the history contains no usable user message, in which // case the caller should reject the request. -func buildSession(messages []ChatCompletionMessage) *session.Session { +func buildSession(messages []ChatCompletionMessage, workingDir string) *session.Session { sess := session.New( session.WithNonInteractive(true), + session.WithWorkingDir(workingDir), ) hasUser := false diff --git a/pkg/chatserver/server.go b/pkg/chatserver/server.go index 6482dc002e..c97d2a348c 100644 --- a/pkg/chatserver/server.go +++ b/pkg/chatserver/server.go @@ -143,6 +143,18 @@ func Run(ctx context.Context, agentFilename string, opts Options, ln net.Listene opts.OnSafetyPolicy(resolvedSafety) } + // Clients are remote, but tools run against the server's own local + // workspace; capture it once so per-request sessions don't depend on a + // cwd that may change while the server runs. + var configuredWd string + if opts.RunConfig != nil { + configuredWd = opts.RunConfig.WorkingDir + } + workingDir, err := session.CaptureLocalWorkingDir(configuredWd) + if err != nil { + return err + } + // Wrap with otelhttp so incoming /v1/chat/completions requests // (including SSE streams) extract the caller's trace context. // otelhttp ends the span when the response body is closed, so @@ -156,6 +168,7 @@ func Run(ctx context.Context, agentFilename string, opts Options, ln net.Listene conversations: newConversationStore(opts.ConversationsMaxSessions, conversationTTL(opts)), conversationLocks: newConversationLockSet(), runtimes: newRuntimePool(ctx, t, opts.MaxIdleRuntimes), + workingDir: workingDir, }, opts), "chatserver", ) @@ -217,6 +230,9 @@ type server struct { conversations *conversationStore conversationLocks *conversationLockSet runtimes *runtimePool + // workingDir is the server's absolute workspace root, captured once at + // startup; every request session persists it as workspace provenance. + workingDir string } func newRouter(s *server, opts Options) http.Handler { @@ -413,7 +429,7 @@ func (s *server) resolveSession(id string, msgs []ChatCompletionMessage) (*sessi return working, nil } } - sess := buildSession(msgs) + sess := buildSession(msgs, s.workingDir) if sess == nil { return nil, errors.New("no user message provided") } diff --git a/pkg/chatserver/server_test.go b/pkg/chatserver/server_test.go index 13f347b3f6..ca90ebf3be 100644 --- a/pkg/chatserver/server_test.go +++ b/pkg/chatserver/server_test.go @@ -66,7 +66,7 @@ func TestBuildSession_RequiresUserMessage(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - sess := buildSession(tc.messages) + sess := buildSession(tc.messages, "") if tc.wantNil { assert.Nil(t, sess) return @@ -77,6 +77,15 @@ func TestBuildSession_RequiresUserMessage(t *testing.T) { } } +func TestBuildSession_StampsWorkspaceProvenance(t *testing.T) { + t.Parallel() + sess := buildSession([]ChatCompletionMessage{ + {Role: "user", Content: "hello"}, + }, "/srv/workspace") + require.NotNil(t, sess) + assert.Equal(t, "/srv/workspace", sess.WorkingDir) +} + func TestBuildSession_PreservesHistory(t *testing.T) { t.Parallel() sess := buildSession([]ChatCompletionMessage{ @@ -84,7 +93,7 @@ func TestBuildSession_PreservesHistory(t *testing.T) { {Role: "user", Content: "hello"}, {Role: "assistant", Content: "hi there"}, {Role: "user", Content: "how are you?"}, - }) + }, "") require.NotNil(t, sess) // GetAllMessages omits system messages. @@ -111,7 +120,7 @@ func TestBuildSession_PreservesToolMessage(t *testing.T) { {Role: "user", Content: "compute 2+2"}, {Role: "assistant", Content: ""}, // dropped: empty content {Role: "tool", Content: "4", ToolCallID: "call_1"}, - }) + }, "") require.NotNil(t, sess) all := sess.GetAllMessages() @@ -127,7 +136,7 @@ func TestBuildSession_UnknownRoleTreatedAsUser(t *testing.T) { t.Parallel() sess := buildSession([]ChatCompletionMessage{ {Role: "developer", Content: "do this"}, - }) + }, "") require.NotNil(t, sess) all := sess.GetAllMessages() @@ -436,7 +445,7 @@ func TestBuildSession_AcceptsImageParts(t *testing.T) { {Type: "text", Text: "What is this?"}, {Type: "image_url", ImageURL: &ContentImageURL{URL: "https://example.com/x.png"}}, }, - }}) + }}, "") require.NotNil(t, sess) all := sess.GetAllMessages() diff --git a/pkg/embeddedchat/embeddedchat.go b/pkg/embeddedchat/embeddedchat.go index 08041c90d7..87ca9f2724 100644 --- a/pkg/embeddedchat/embeddedchat.go +++ b/pkg/embeddedchat/embeddedchat.go @@ -119,6 +119,9 @@ type Session struct { rt runtimeRunner session *session.Session welcome string + // workingDir is the absolute workspace root captured once at New; every + // conversation session persists it as its workspace provenance. + workingDir string mu sync.Mutex activeCancel context.CancelFunc @@ -180,6 +183,12 @@ func New(ctx context.Context, cfg Config) (*Session, error) { } s := &Session{cfg: cfg, rt: rt} + // Capture the embedder's workspace root once, at initialization: a later + // process chdir must not change which workspace owns the conversations. + s.workingDir, err = session.CaptureLocalWorkingDir(runConfig.WorkingDir) + if err != nil { + return nil, fmt.Errorf("embeddedchat: capture working dir: %w", err) + } if root, err := tm.DefaultAgent(); err == nil { s.welcome = root.WelcomeMessage() } @@ -244,7 +253,9 @@ func (s *Session) Close() error { } func (s *Session) resetConversationLocked() { - opts := append([]session.Opt(nil), s.cfg.SessionOptions...) + // The captured root goes first so an embedder-supplied WithWorkingDir in + // SessionOptions still wins. + opts := append([]session.Opt{session.WithWorkingDir(s.workingDir)}, s.cfg.SessionOptions...) s.session = session.New(opts...) } diff --git a/pkg/embeddedchat/embeddedchat_test.go b/pkg/embeddedchat/embeddedchat_test.go index 19f52e1680..bacc6c5871 100644 --- a/pkg/embeddedchat/embeddedchat_test.go +++ b/pkg/embeddedchat/embeddedchat_test.go @@ -85,6 +85,37 @@ func TestNewFromCodeBuiltTeam(t *testing.T) { require.NotNil(t, s.Conversation()) } +func TestConversationsCarryWorkspaceProvenance(t *testing.T) { + t.Parallel() + root := t.TempDir() + + s, err := New(t.Context(), Config{ + Team: newCodeBuiltTeam(), + RuntimeConfig: &dagentcfg.RuntimeConfig{Config: dagentcfg.Config{WorkingDir: root}}, + }) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, s.Close()) }) + require.Equal(t, root, s.Conversation().WorkingDir) + + require.NoError(t, s.Restart()) + require.Equal(t, root, s.Conversation().WorkingDir, "restarted conversations must keep the workspace root") +} + +func TestSessionOptionsOverrideCapturedWorkingDir(t *testing.T) { + t.Parallel() + configured := t.TempDir() + override := t.TempDir() + + s, err := New(t.Context(), Config{ + Team: newCodeBuiltTeam(), + RuntimeConfig: &dagentcfg.RuntimeConfig{Config: dagentcfg.Config{WorkingDir: configured}}, + SessionOptions: []session.Opt{session.WithWorkingDir(override)}, + }) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, s.Close()) }) + require.Equal(t, override, s.Conversation().WorkingDir) +} + func TestInitialSessionResumesConversation(t *testing.T) { t.Parallel() restored := session.New() diff --git a/pkg/evaluation/save.go b/pkg/evaluation/save.go index 9407b5f38f..b74c5c785b 100644 --- a/pkg/evaluation/save.go +++ b/pkg/evaluation/save.go @@ -55,7 +55,8 @@ func SaveRunSessions(ctx context.Context, run *EvalRun, outputDir string) (strin // SessionFromEvents reconstructs a session from raw container output events. // This parses the JSON events emitted by docker agent run --exec --json and builds a session -// with the conversation history. +// with the conversation history. WorkingDir is intentionally left empty: the +// transcript ran inside a container, so no host workspace owns it. func SessionFromEvents(events []map[string]any, title string, questions []string) *session.Session { sess := session.New( session.WithTitle(title), diff --git a/pkg/mcp/server.go b/pkg/mcp/server.go index 7d0a548abd..e1676a48e9 100644 --- a/pkg/mcp/server.go +++ b/pkg/mcp/server.go @@ -181,6 +181,11 @@ func createMCPServerForTeam(ctx context.Context, t *team.Team, agentFilename, ag return nil, errors.New("--tool-name can only be used when exactly one agent is exposed") } + workingDir, err := session.CaptureLocalWorkingDir(runConfig.WorkingDir) + if err != nil { + return nil, err + } + slog.DebugContext(ctx, "Adding MCP tools for agents", "count", len(agentNames)) for _, agentName := range agentNames { @@ -208,17 +213,17 @@ func createMCPServerForTeam(ctx context.Context, t *team.Team, agentFilename, ag OutputSchema: tools.MustSchemaFor[ToolOutput](), } - mcp.AddTool(server, toolDef, createToolHandler(t, agentName, safety)) + mcp.AddTool(server, toolDef, createToolHandler(t, agentName, safety, workingDir)) } return server, nil } -func CreateToolHandler(t *team.Team, agentName string, safety session.SafetyPolicy) func(context.Context, *mcp.CallToolRequest, ToolInput) (*mcp.CallToolResult, ToolOutput, error) { - return createToolHandler(t, agentName, safety) +func CreateToolHandler(t *team.Team, agentName string, safety session.SafetyPolicy, workingDir string) func(context.Context, *mcp.CallToolRequest, ToolInput) (*mcp.CallToolResult, ToolOutput, error) { + return createToolHandler(t, agentName, safety, workingDir) } -func createToolHandler(t *team.Team, agentName string, safety session.SafetyPolicy) func(context.Context, *mcp.CallToolRequest, ToolInput) (*mcp.CallToolResult, ToolOutput, error) { +func createToolHandler(t *team.Team, agentName string, safety session.SafetyPolicy, workingDir string) func(context.Context, *mcp.CallToolRequest, ToolInput) (*mcp.CallToolResult, ToolOutput, error) { return func(ctx context.Context, req *mcp.CallToolRequest, input ToolInput) (result *mcp.CallToolResult, output ToolOutput, err error) { // Extract W3C trace context from `params._meta` (per the OTel // MCP semconv) so the SERVER span chains onto the calling @@ -245,17 +250,7 @@ func createToolHandler(t *team.Team, agentName string, safety session.SafetyPoli return nil, ToolOutput{}, fmt.Errorf("failed to get agent: %w", err) } - sessionOptions := []session.Opt{ - session.WithTitle("MCP tool call"), - session.WithMaxIterations(ag.MaxIterations()), - session.WithMaxConsecutiveToolCalls(ag.MaxConsecutiveToolCalls()), - session.WithMaxOldToolCallTokens(ag.MaxOldToolCallTokens()), - session.WithMaxToolResultTokens(ag.MaxToolResultTokens()), - session.WithUserMessage(input.Message), - session.WithNonInteractive(true), - session.WithSafetyPolicy(safety), - } - sess := session.New(sessionOptions...) + sess := newToolCallSession(ag, input.Message, safety, workingDir) rt, err := runtime.New(ctx, t, runtime.WithCurrentAgent(agentName), @@ -334,6 +329,21 @@ func agentToolAnnotations(ctx context.Context, ag *agent.Agent) (*mcp.ToolAnnota return annotations, nil } +func newToolCallSession(ag *agent.Agent, message string, safety session.SafetyPolicy, workingDir string) *session.Session { + return session.New( + session.WithTitle("MCP tool call"), + session.WithMaxIterations(ag.MaxIterations()), + session.WithMaxConsecutiveToolCalls(ag.MaxConsecutiveToolCalls()), + session.WithMaxOldToolCallTokens(ag.MaxOldToolCallTokens()), + session.WithMaxToolResultTokens(ag.MaxToolResultTokens()), + session.WithUserMessage(message), + session.WithToolsApproved(true), + session.WithNonInteractive(true), + session.WithSafetyPolicy(safety), + session.WithWorkingDir(workingDir), + ) +} + // optionalBool returns the value of p, or fallback when p is nil. func optionalBool(p *bool, fallback bool) bool { if p != nil { diff --git a/pkg/mcp/server_test.go b/pkg/mcp/server_test.go index 543ce42efe..707ec17604 100644 --- a/pkg/mcp/server_test.go +++ b/pkg/mcp/server_test.go @@ -12,6 +12,7 @@ import ( "github.com/docker/docker-agent/pkg/agent" "github.com/docker/docker-agent/pkg/config" "github.com/docker/docker-agent/pkg/httpsec" + "github.com/docker/docker-agent/pkg/session" "github.com/docker/docker-agent/pkg/tools" ) @@ -180,3 +181,18 @@ func TestCreateMCPServer_ToolNameRejectsMultipleAgents(t *testing.T) { assert.Contains(t, err.Error(), "--tool-name") assert.Contains(t, err.Error(), "exactly one agent") } + +func TestNewToolCallSession(t *testing.T) { + t.Parallel() + + ag := agent.New("root", "test agent", agent.WithMaxIterations(7)) + sess := newToolCallSession(ag, "hello", session.SafetyPolicyAutonomous, "/srv/workspace") + + assert.Equal(t, "MCP tool call", sess.Title) + assert.Equal(t, 7, sess.MaxIterations) + assert.True(t, sess.ToolsApproved) + assert.True(t, sess.NonInteractive) + assert.Equal(t, session.SafetyPolicyAutonomous, sess.SafetyPolicy) + assert.Equal(t, "hello", sess.GetLastUserMessageContent()) + assert.Equal(t, "/srv/workspace", sess.WorkingDir) +} diff --git a/pkg/runtime/agent_delegation.go b/pkg/runtime/agent_delegation.go index 92dd017824..d051b10a8c 100644 --- a/pkg/runtime/agent_delegation.go +++ b/pkg/runtime/agent_delegation.go @@ -248,6 +248,10 @@ func newSubSession(parent *session.Session, cfg SubSessionConfig, childAgent *ag session.WithSendUserMessage(false), session.WithStructuredOutputDisabled(cfg.DisableStructuredOutput), session.WithParentID(parent.ID), + // Delegated children run in the parent's workspace: the persisted + // WorkingDir is the workspace-root provenance later used to resolve + // files the child produced. Empty stays empty (headless parents). + session.WithWorkingDir(parent.WorkingDir), session.WithAttachedFiles(attachedFiles), session.WithAttributes(parent.AttributesSnapshot()), } diff --git a/pkg/runtime/agent_delegation_test.go b/pkg/runtime/agent_delegation_test.go index 399ca7313f..acca5e2d55 100644 --- a/pkg/runtime/agent_delegation_test.go +++ b/pkg/runtime/agent_delegation_test.go @@ -187,6 +187,44 @@ func TestNewSubSession(t *testing.T) { }) } +func TestNewSubSession_WorkingDirInheritance(t *testing.T) { + t.Parallel() + + childAgent := agent.New("worker", "a worker agent") + cfg := SubSessionConfig{Task: "do work", AgentName: "worker", Title: "Task"} + + t.Run("child inherits parent workspace", func(t *testing.T) { + t.Parallel() + parent := session.New(session.WithWorkingDir("/work/project")) + + child := newSubSession(parent, cfg, childAgent) + + assert.Equal(t, "/work/project", child.WorkingDir) + assert.Equal(t, parent.AllowedDirectories(), child.AllowedDirectories()) + }) + + t.Run("nested children keep the original workspace", func(t *testing.T) { + t.Parallel() + parent := session.New(session.WithWorkingDir("/work/project")) + + child := newSubSession(parent, cfg, childAgent) + grandchild := newSubSession(child, cfg, childAgent) + + assert.Equal(t, "/work/project", grandchild.WorkingDir) + assert.Equal(t, child.ID, grandchild.ParentID) + }) + + t.Run("workspace-less parent stays empty", func(t *testing.T) { + t.Parallel() + parent := session.New() + + child := newSubSession(parent, cfg, childAgent) + + assert.Empty(t, child.WorkingDir, "a headless parent must not make the child pick up a cwd") + assert.Nil(t, child.AllowedDirectories()) + }) +} + func TestSubSessionConfig_DefaultValues(t *testing.T) { t.Parallel() diff --git a/pkg/runtime/compactor/compactor.go b/pkg/runtime/compactor/compactor.go index ce4d7d7ed9..14bf45fe8a 100644 --- a/pkg/runtime/compactor/compactor.go +++ b/pkg/runtime/compactor/compactor.go @@ -223,6 +223,9 @@ func RunLLM(ctx context.Context, args LLMArgs) (result *Result, err error) { compactionSession := session.New( session.WithTitle("Generating summary"), session.WithMessages(toItems(messages)), + // The summarization run happens in the same workspace as the session + // being compacted; empty stays empty for workspace-less sessions. + session.WithWorkingDir(args.Session.WorkingDir), ) seedLen := len(compactionSession.Messages) diff --git a/pkg/runtime/compactor/compactor_test.go b/pkg/runtime/compactor/compactor_test.go index d1e5be9184..9547dceb89 100644 --- a/pkg/runtime/compactor/compactor_test.go +++ b/pkg/runtime/compactor/compactor_test.go @@ -571,6 +571,37 @@ func TestRunLLM_ReportsModelAndUsage(t *testing.T) { assert.Equal(t, int64(45), result.Usage.OutputTokens) } +// TestRunLLM_CompactionSessionInheritsWorkspace verifies the ephemeral +// summarization session runs with the compacted session's workspace root. +func TestRunLLM_CompactionSessionInheritsWorkspace(t *testing.T) { + t.Parallel() + + sess := session.New( + session.WithUserMessage("please do the task"), + session.WithWorkingDir("/work/project"), + ) + sess.AddMessage(session.NewAgentMessage("root", &chat.Message{ + Role: chat.MessageRoleAssistant, + Content: "done", + })) + a := agent.New("root", "instr", agent.WithModel(fakeProvider{id: modelsdev.NewID("fake", "model")})) + + _, err := RunLLM(t.Context(), LLMArgs{ + Session: sess, + Agent: a, + ContextLimit: 8_192, + RunAgent: func(_ context.Context, _ *agent.Agent, cs *session.Session) error { + assert.Equal(t, "/work/project", cs.WorkingDir) + cs.AddMessage(session.NewAgentMessage("root", &chat.Message{ + Role: chat.MessageRoleAssistant, + Content: "the summary", + })) + return nil + }, + }) + require.NoError(t, err) +} + // TestRunLLM_NoConversationFits_NoOps pins the safety net behind the // scaled budgets: when not a single conversation message fits the // summarization budget (e.g. one giant tool result), RunLLM must no-op diff --git a/pkg/server/session_manager.go b/pkg/server/session_manager.go index 605deba102..08d3a1e10f 100644 --- a/pkg/server/session_manager.go +++ b/pkg/server/session_manager.go @@ -584,6 +584,9 @@ func (sm *SessionManager) CreateSession(ctx context.Context, sessionTemplate *se opts = append(opts, session.WithTitle(title)) } + // A template without a working_dir creates an intentionally + // workspace-less session: the API caller is remote and the server must + // not guess its own process cwd as the session's workspace provenance. if wd := strings.TrimSpace(sessionTemplate.WorkingDir); wd != "" { // Refuse any raw ".." here in CreateSession, before filepath.Abs // cleans it away: the traversal rejection is auditable on the raw diff --git a/pkg/session/working_dir.go b/pkg/session/working_dir.go new file mode 100644 index 0000000000..a6ae79bddb --- /dev/null +++ b/pkg/session/working_dir.go @@ -0,0 +1,112 @@ +package session + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "strings" +) + +// ErrWorkingDirUnavailable is returned by ResolveWorkingDir when neither the +// session nor any ancestor carries a usable workspace root. Callers must +// treat it as "no workspace" — never substitute the viewer's process cwd, +// which could belong to a different workspace than the one that owned the +// session's files. +var ErrWorkingDirUnavailable = errors.New("session working directory unavailable") + +// maxWorkingDirAncestry bounds how many sessions (including the starting one) +// ResolveWorkingDir inspects while walking parent links, protecting against +// corrupt stores with very deep or cyclic parent chains. +const maxWorkingDirAncestry = 32 + +// Lookup loads a session by ID. Store.GetSession satisfies it; tests can +// supply a map-backed stub. +type Lookup func(ctx context.Context, id string) (*Session, error) + +// CaptureLocalWorkingDir returns the absolute effective workspace root for a +// session being created on a surface with a LOCAL workspace. configured is +// the surface's configured working directory (e.g. --working-dir); when it is +// empty the process cwd — captured once, at creation/startup time — is the +// workspace. Remote or headless surfaces without a local workspace must not +// call this: their sessions intentionally keep an empty WorkingDir. +func CaptureLocalWorkingDir(configured string) (string, error) { + configured = strings.TrimSpace(configured) + if configured == "" { + cwd, err := os.Getwd() + if err != nil { + return "", fmt.Errorf("capturing workspace root: %w", err) + } + return filepath.Clean(cwd), nil + } + abs, err := filepath.Abs(configured) + if err != nil { + return "", fmt.Errorf("capturing workspace root: %w", err) + } + return filepath.Clean(abs), nil +} + +// ResolveWorkingDir returns the workspace root that owns sess. It prefers the +// session's own persisted WorkingDir and, for old persisted sub-sessions that +// predate creation-time provenance, walks the parent chain (bounded, cycle +// safe) until an ancestor provides one. Stored values must be absolute and +// well-formed; anything else fails with ErrWorkingDirUnavailable rather than +// being resolved against the current process cwd. The workspace directory is +// not required to still exist — a deleted workspace is still valid provenance. +func ResolveWorkingDir(ctx context.Context, sess *Session, lookup Lookup) (string, error) { + if sess == nil { + return "", fmt.Errorf("%w: nil session", ErrWorkingDirUnavailable) + } + + visited := make(map[string]struct{}, 4) + for range maxWorkingDirAncestry { + if sess.WorkingDir != "" { + if err := validateStoredWorkingDir(sess.WorkingDir); err != nil { + return "", fmt.Errorf("%w: session %s: %w", ErrWorkingDirUnavailable, sess.ID, err) + } + return sess.WorkingDir, nil + } + if sess.ParentID == "" { + return "", fmt.Errorf("%w: session %s has no workspace root and no parent", ErrWorkingDirUnavailable, sess.ID) + } + if sess.ID != "" { + visited[sess.ID] = struct{}{} + } + if _, seen := visited[sess.ParentID]; seen { + return "", fmt.Errorf("%w: parent chain cycle at session %s", ErrWorkingDirUnavailable, sess.ParentID) + } + if lookup == nil { + return "", fmt.Errorf("%w: session %s needs parent lookup but none was provided", ErrWorkingDirUnavailable, sess.ID) + } + parent, err := lookup(ctx, sess.ParentID) + if err != nil { + return "", fmt.Errorf("%w: loading parent %s: %w", ErrWorkingDirUnavailable, sess.ParentID, err) + } + if parent == nil || parent.ID != sess.ParentID { + return "", fmt.Errorf("%w: lookup for parent %s returned a different session", ErrWorkingDirUnavailable, sess.ParentID) + } + sess = parent + } + return "", fmt.Errorf("%w: parent chain exceeds %d sessions", ErrWorkingDirUnavailable, maxWorkingDirAncestry) +} + +// validateStoredWorkingDir vets a persisted WorkingDir before it is trusted as +// a workspace root. It is strict on purpose: a malformed value must make +// resolution fail closed instead of being repaired relative to whatever +// directory the viewing process happens to run in. +func validateStoredWorkingDir(dir string) error { + if dir != strings.TrimSpace(dir) { + return errors.New("working directory has surrounding whitespace") + } + if strings.ContainsRune(dir, 0) { + return errors.New("working directory contains a NUL byte") + } + if !filepath.IsAbs(dir) { + return fmt.Errorf("working directory %q is not absolute", dir) + } + if cleaned := filepath.Clean(dir); !filepath.IsAbs(cleaned) { + return fmt.Errorf("working directory %q does not clean to an absolute path", dir) + } + return nil +} diff --git a/pkg/session/working_dir_test.go b/pkg/session/working_dir_test.go new file mode 100644 index 0000000000..7f66301e25 --- /dev/null +++ b/pkg/session/working_dir_test.go @@ -0,0 +1,282 @@ +package session + +import ( + "context" + "fmt" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCaptureLocalWorkingDir(t *testing.T) { + base := t.TempDir() + // Resolve symlinks (macOS /var → /private/var) so cwd-derived values + // compare equal to the TempDir path. + base, err := filepath.EvalSymlinks(base) + require.NoError(t, err) + t.Chdir(base) + + t.Run("empty captures process cwd", func(t *testing.T) { + got, err := CaptureLocalWorkingDir("") + require.NoError(t, err) + assert.Equal(t, base, got) + }) + + t.Run("whitespace-only captures process cwd", func(t *testing.T) { + got, err := CaptureLocalWorkingDir(" ") + require.NoError(t, err) + assert.Equal(t, base, got) + }) + + t.Run("relative is absolutized against cwd", func(t *testing.T) { + require.NoError(t, os.Mkdir(filepath.Join(base, "sub"), 0o755)) + got, err := CaptureLocalWorkingDir("sub") + require.NoError(t, err) + assert.Equal(t, filepath.Join(base, "sub"), got) + }) + + t.Run("absolute is cleaned", func(t *testing.T) { + got, err := CaptureLocalWorkingDir(base + string(filepath.Separator) + "a" + string(filepath.Separator) + ".." + string(filepath.Separator) + "b") + require.NoError(t, err) + assert.Equal(t, filepath.Join(base, "b"), got) + }) +} + +// mapLookup builds a Lookup over a fixed set of sessions. +func mapLookup(sessions ...*Session) Lookup { + byID := make(map[string]*Session, len(sessions)) + for _, s := range sessions { + byID[s.ID] = s + } + return func(_ context.Context, id string) (*Session, error) { + s, ok := byID[id] + if !ok { + return nil, ErrNotFound + } + return s, nil + } +} + +func TestResolveWorkingDir_OwnRoot(t *testing.T) { + t.Parallel() + sess := &Session{ID: "s1", WorkingDir: "/work/project"} + + got, err := ResolveWorkingDir(t.Context(), sess, nil) + require.NoError(t, err) + assert.Equal(t, "/work/project", got) +} + +func TestResolveWorkingDir_RootNeedNotExist(t *testing.T) { + t.Parallel() + // A deleted workspace is still valid provenance; existence checks belong + // to whoever opens files under the root. + sess := &Session{ID: "s1", WorkingDir: "/definitely/gone/workspace"} + + got, err := ResolveWorkingDir(t.Context(), sess, nil) + require.NoError(t, err) + assert.Equal(t, "/definitely/gone/workspace", got) +} + +func TestResolveWorkingDir_InheritsFromParentChain(t *testing.T) { + t.Parallel() + root := &Session{ID: "root", WorkingDir: "/work/a"} + mid := &Session{ID: "mid", ParentID: "root"} + leaf := &Session{ID: "leaf", ParentID: "mid"} + + got, err := ResolveWorkingDir(t.Context(), leaf, mapLookup(root, mid)) + require.NoError(t, err) + assert.Equal(t, "/work/a", got) +} + +func TestResolveWorkingDir_Failures(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + sess *Session + lookup Lookup + }{ + {name: "nil session", sess: nil}, + {name: "empty with no parent", sess: &Session{ID: "s1"}}, + { + name: "missing parent", + sess: &Session{ID: "s1", ParentID: "gone"}, + lookup: mapLookup(), + }, + {name: "relative root rejected", sess: &Session{ID: "s1", WorkingDir: "relative/dir"}}, + {name: "dot root rejected", sess: &Session{ID: "s1", WorkingDir: "."}}, + {name: "whitespace-padded root rejected", sess: &Session{ID: "s1", WorkingDir: " /work "}}, + {name: "NUL root rejected", sess: &Session{ID: "s1", WorkingDir: "/work\x00evil"}}, + { + name: "relative parent root rejected, not repaired", + sess: &Session{ID: "leaf", ParentID: "root"}, + lookup: mapLookup( + &Session{ID: "root", WorkingDir: "relative"}, + ), + }, + { + name: "self parent cycle", + sess: &Session{ID: "s1", ParentID: "s1"}, + lookup: mapLookup( + &Session{ID: "s1", ParentID: "s1"}, + ), + }, + { + name: "two-node cycle", + sess: &Session{ID: "a", ParentID: "b"}, + lookup: mapLookup( + &Session{ID: "a", ParentID: "b"}, + &Session{ID: "b", ParentID: "a"}, + ), + }, + {name: "traversal without lookup", sess: &Session{ID: "s1", ParentID: "p1"}}, + { + name: "lookup ID mismatch", + sess: &Session{ID: "s1", ParentID: "p1"}, + lookup: func(context.Context, string) (*Session, error) { + return &Session{ID: "other"}, nil + }, + }, + { + name: "nil parent from lookup", + sess: &Session{ID: "s1", ParentID: "p1"}, + lookup: func(context.Context, string) (*Session, error) { + return nil, nil + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + got, err := ResolveWorkingDir(t.Context(), tc.sess, tc.lookup) + require.ErrorIs(t, err, ErrWorkingDirUnavailable) + assert.Empty(t, got) + }) + } +} + +// chainOfDepth builds a lookup over a parent chain of n sessions where only +// the topmost ancestor carries a WorkingDir, and returns the leaf. +func chainOfDepth(n int, root string) (*Session, Lookup) { + sessions := make([]*Session, 0, n) + for i := range n { + s := &Session{ID: fmt.Sprintf("s%d", i)} + if i < n-1 { + s.ParentID = fmt.Sprintf("s%d", i+1) + } else { + s.WorkingDir = root + } + sessions = append(sessions, s) + } + return sessions[0], mapLookup(sessions...) +} + +func TestResolveWorkingDir_DepthBound(t *testing.T) { + t.Parallel() + + t.Run("exactly at bound succeeds", func(t *testing.T) { + t.Parallel() + leaf, lookup := chainOfDepth(maxWorkingDirAncestry, "/work/deep") + got, err := ResolveWorkingDir(t.Context(), leaf, lookup) + require.NoError(t, err) + assert.Equal(t, "/work/deep", got) + }) + + t.Run("one beyond bound fails", func(t *testing.T) { + t.Parallel() + leaf, lookup := chainOfDepth(maxWorkingDirAncestry+1, "/work/deep") + _, err := ResolveWorkingDir(t.Context(), leaf, lookup) + require.ErrorIs(t, err, ErrWorkingDirUnavailable) + }) +} + +// TestResolveWorkingDir_CrossWorkspaceInvariant pins the media-safety +// invariant: a session created in workspace A and resolved from a process +// running in workspace B yields A (or fails closed) — never B, even when B +// contains a same-named candidate. Not parallel: it changes the process cwd. +func TestResolveWorkingDir_CrossWorkspaceInvariant(t *testing.T) { + workspaceA := t.TempDir() + workspaceB := t.TempDir() + t.Chdir(workspaceB) + + store := NewInMemorySessionStore() + root := New(WithWorkingDir(workspaceA)) + require.NoError(t, store.AddSession(t.Context(), root)) + // Old-style sub-session persisted before creation-time provenance: + // empty WorkingDir, linked to its parent. + child := New(WithParentID(root.ID)) + require.NoError(t, store.AddSubSession(t.Context(), root.ID, child)) + + reloadedRoot, err := store.GetSession(t.Context(), root.ID) + require.NoError(t, err) + got, err := ResolveWorkingDir(t.Context(), reloadedRoot, store.GetSession) + require.NoError(t, err) + assert.Equal(t, workspaceA, got) + + reloadedChild, err := store.GetSession(t.Context(), child.ID) + require.NoError(t, err) + got, err = ResolveWorkingDir(t.Context(), reloadedChild, store.GetSession) + require.NoError(t, err) + assert.Equal(t, workspaceA, got, "old empty sub-session must inherit its parent's workspace, not the viewer cwd") + + // A workspace-less session fails closed instead of picking up B. + orphan := New() + _, err = ResolveWorkingDir(t.Context(), orphan, store.GetSession) + require.ErrorIs(t, err, ErrWorkingDirUnavailable) + + // A persisted relative root is rejected, never resolved against B. + relative := &Session{ID: "rel", WorkingDir: "."} + _, err = ResolveWorkingDir(t.Context(), relative, store.GetSession) + require.ErrorIs(t, err, ErrWorkingDirUnavailable) +} + +// TestSubSessionWorkingDirRoundTrip pins that a sub-session's WorkingDir and +// parent link survive persistence in both store implementations, and that a +// legacy empty sub-session resolves through its stored parent. +func TestSubSessionWorkingDirRoundTrip(t *testing.T) { + t.Parallel() + + stores := map[string]func(t *testing.T) Store{ + "in-memory": func(*testing.T) Store { return NewInMemorySessionStore() }, + "sqlite": func(t *testing.T) Store { + t.Helper() + store, err := newSQLiteStoreForTest(t, filepath.Join(t.TempDir(), "sessions.db")) + require.NoError(t, err) + return store + }, + } + + for name, newStore := range stores { + t.Run(name, func(t *testing.T) { + t.Parallel() + store := newStore(t) + + parent := New(WithWorkingDir("/work/project")) + require.NoError(t, store.AddSession(t.Context(), parent)) + + child := New(WithWorkingDir(parent.WorkingDir)) + require.NoError(t, store.AddSubSession(t.Context(), parent.ID, child)) + + reloaded, err := store.GetSession(t.Context(), child.ID) + require.NoError(t, err) + assert.Equal(t, parent.ID, reloaded.ParentID) + assert.Equal(t, "/work/project", reloaded.WorkingDir) + + // Legacy sub-session persisted without provenance: stays empty in + // the store and resolves via the parent chain. + legacy := New() + require.NoError(t, store.AddSubSession(t.Context(), parent.ID, legacy)) + reloadedLegacy, err := store.GetSession(t.Context(), legacy.ID) + require.NoError(t, err) + assert.Empty(t, reloadedLegacy.WorkingDir) + + resolved, err := ResolveWorkingDir(t.Context(), reloadedLegacy, store.GetSession) + require.NoError(t, err) + assert.Equal(t, "/work/project", resolved) + }) + } +} From 0b1903a77b48267e2714d676fbd9149709324f71 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Arnaud=20He=CC=81ritier?= Date: Fri, 31 Jul 2026 01:12:52 +0200 Subject: [PATCH 2/3] feat(#3996): add collision-safe workspace media writer MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Pure naming layer for saving generated media as workspace deliverables: sanitizes requested relative paths, derives the extension from the MIME type (reporting a corrected conflicting extension so callers can show a notice), classifies absolute/".."/invalid/reserved/symlink-escaping paths as a typed ErrPathEscape, and never overwrites — the final name is claimed with O_CREATE|O_EXCL, collisions retry with a dash suffix, and bytes are published atomically via a sibling temp file inside os.OpenRoot containment. Materialization wiring comes separately. --- pkg/workspacemedia/writer.go | 372 ++++++++++++++++++++++++++++++ pkg/workspacemedia/writer_test.go | 291 +++++++++++++++++++++++ 2 files changed, 663 insertions(+) create mode 100644 pkg/workspacemedia/writer.go create mode 100644 pkg/workspacemedia/writer_test.go diff --git a/pkg/workspacemedia/writer.go b/pkg/workspacemedia/writer.go new file mode 100644 index 0000000000..acf1164012 --- /dev/null +++ b/pkg/workspacemedia/writer.go @@ -0,0 +1,372 @@ +// Package workspacemedia writes model-generated media into the user's +// workspace as ordinary, visible files — the same way generated code or +// text lands there. It is a pure naming/collision layer: the caller hands +// it a workspace root, a requested relative path (or a generic fallback +// name such as "generated"), the bytes, and the MIME type; it returns the +// exact workspace-relative path it wrote. +// +// Guarantees: +// - Never overwrites. The final name is claimed with O_CREATE|O_EXCL and +// an existing entry of any kind (including a symlink) means "taken": +// the writer retries with a dash suffix (name-1.ext, name-2.ext, ...), +// which also makes concurrent same-name writers safe. +// - Atomic publish. Bytes go to a sibling temp file renamed over the +// claimed name, so a reader observes either the empty claim or the +// full content, never a partial write. A failed write removes the +// empty claim instead of leaving a zero-byte "generated" file behind. +// - Workspace containment. Every directory and file operation goes +// through os.Root, so absolute paths, ".." traversal, invalid or +// Windows-reserved segments, and symlink escapes are rejected with +// [ErrPathEscape] rather than followed. +// - The final filename's extension always agrees with the MIME type when +// the type is known; a conflicting requested extension is corrected +// and reported via [Result] so the caller can show a notice. +// +// Files use ordinary user-file modes (0o755 directories, 0o644 files, both +// umask-masked). On Windows modes are ignored and entries inherit the +// parent directory's ACLs — same convention as pkg/atomicfile. +package workspacemedia + +import ( + "bytes" + "crypto/rand" + "errors" + "fmt" + "io" + "io/fs" + "mime" + "os" + "path" + "slices" + "strings" +) + +// ErrPathEscape classifies a requested path the writer refuses to touch: +// absolute, containing "..", an empty/invalid/Windows-reserved segment, or +// resolving outside the workspace root through a symlinked parent. +// Classification only — the caller decides how to react (e.g. ask the user +// to confirm an out-of-workspace target, or redirect to a safe name). +var ErrPathEscape = errors.New("path escapes the workspace or has invalid segments") + +// Result describes a completed write. +type Result struct { + // RelPath is the exact final path written, relative to the workspace + // root and slash-separated. Persist this verbatim. + RelPath string + + // ExtensionCorrected reports that the requested filename's extension + // conflicted with the MIME-derived one and was replaced. Callers should + // surface a notice (e.g. "saved as sunshine.png — the provider returned + // PNG data"). RequestedExtension holds the original, with leading dot. + ExtensionCorrected bool + RequestedExtension string +} + +// maxNameAttempts bounds the dash-suffix collision retry so a pathological +// directory cannot loop forever; exhaustion surfaces as a visible error. +const maxNameAttempts = 10000 + +// Write stores data under workspaceRoot at requestedPath, sanitized and +// collision-avoided per the package contract, and returns the exact +// workspace-relative path written. Prompt-directed subdirectories in +// requestedPath are created as needed. A rejected path returns an error +// matching [ErrPathEscape]; any other failure (unwritable directory, full +// disk, ...) is returned as-is for the caller to surface. +func Write(workspaceRoot, requestedPath string, data []byte, mimeType string) (Result, error) { + return write(workspaceRoot, requestedPath, bytes.NewReader(data), mimeType) +} + +func write(workspaceRoot, requestedPath string, r io.Reader, mimeType string) (Result, error) { + dir, base, requestedExt, err := splitRequestedPath(requestedPath) + if err != nil { + return Result{}, err + } + ext, corrected := finalExtension(requestedExt, mimeType) + + root, err := os.OpenRoot(workspaceRoot) + if err != nil { + return Result{}, fmt.Errorf("open workspace root: %w", err) + } + defer root.Close() + + if dir != "" { + switch err := root.MkdirAll(dir, 0o755); { + case err == nil: + case isPathEscape(err): + return Result{}, escapeError(requestedPath, err) + case errors.Is(err, fs.ErrExist): + // os.Root.MkdirAll reports an existing symlink as ErrExist instead + // of traversing it. Let the OpenFile claim below decide: it resolves + // symlinks, allows in-root targets, and classifies escapes. + default: + return Result{}, fmt.Errorf("create directory %q: %w", dir, err) + } + } + + rel, err := claimAndPublish(root, dir, base, ext, r) + if err != nil { + return Result{}, err + } + res := Result{RelPath: rel} + if corrected { + res.ExtensionCorrected = true + res.RequestedExtension = requestedExt + } + return res, nil +} + +// claimAndPublish reserves the first free candidate name with +// O_CREATE|O_EXCL — the claim is what serializes concurrent writers and +// keeps an existing file (or symlink) from ever being replaced — then +// atomically publishes r's content into it. +func claimAndPublish(root *os.Root, dir, base, ext string, r io.Reader) (string, error) { + for n := range maxNameAttempts { + name := base + ext + if n > 0 { + name = fmt.Sprintf("%s-%d%s", base, n, ext) + } + rel := path.Join(dir, name) + + claim, err := root.OpenFile(rel, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644) + switch { + case errors.Is(err, fs.ErrExist): + continue + case isPathEscape(err): + return "", escapeError(rel, err) + case err != nil: + return "", fmt.Errorf("claim %q: %w", rel, err) + } + claim.Close() + + if err := publish(root, dir, rel, r); err != nil { + // The claim is still the empty placeholder; leaving it around + // would look like a zero-byte generated file. + _ = root.Remove(rel) + return "", err + } + return rel, nil + } + return "", fmt.Errorf("no free name for %q after %d attempts", path.Join(dir, base+ext), maxNameAttempts) +} + +// publish writes r to a sibling temp file, syncs it, and renames it over +// the claimed name. os.Root.Rename replaces the destination atomically on +// POSIX; on Windows it renames with POSIX semantics (see +// pkg/atomicfile/write_windows.go, which this mirrors inside an os.Root). +func publish(root *os.Root, dir, claimed string, r io.Reader) error { + tmp, f, err := createTemp(root, dir) + if err != nil { + return err + } + // The rename consumes the temp name on success, so this only collects + // the temp file after a failure. + defer func() { _ = root.Remove(tmp) }() + + if _, err := io.Copy(f, r); err != nil { + f.Close() + return fmt.Errorf("write %q: %w", tmp, err) + } + // Sync before the rename so a crash cannot publish a torn file. + if err := f.Sync(); err != nil { + f.Close() + return fmt.Errorf("flush %q: %w", tmp, err) + } + if err := f.Close(); err != nil { + return fmt.Errorf("close %q: %w", tmp, err) + } + if err := root.Rename(tmp, claimed); err != nil { + return fmt.Errorf("publish %q: %w", claimed, err) + } + return nil +} + +// createTemp is os.CreateTemp confined to root (os.Root has no CreateTemp +// as of Go 1.26). The generic random name never collides in practice; the +// small retry is only for completeness. +func createTemp(root *os.Root, dir string) (string, *os.File, error) { + for range 3 { + name := path.Join(dir, ".media-"+rand.Text()+".tmp") + f, err := root.OpenFile(name, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644) + if errors.Is(err, fs.ErrExist) { + continue + } + if err != nil { + return "", nil, fmt.Errorf("create temp file: %w", err) + } + return name, f, nil + } + return "", nil, errors.New("create temp file: name collisions") +} + +// splitRequestedPath validates and sanitizes the requested path, returning +// the slash-joined parent directory (possibly ""), the filename without +// extension, and the requested extension (with leading dot, possibly ""). +func splitRequestedPath(requested string) (dir, base, ext string, err error) { + if isAbsolutePath(requested) { + return "", "", "", escapeError(requested, errors.New("absolute path")) + } + // Accept both separators: model-provided paths may be Windows-style. + segments := strings.FieldsFunc(requested, func(r rune) bool { return r == '/' || r == '\\' }) + + var cleaned []string + for _, seg := range segments { + switch seg { + case ".": + continue + case "..": + return "", "", "", escapeError(requested, errors.New(`".." segment`)) + } + s := sanitizeSegment(seg) + switch { + case s == "": + return "", "", "", escapeError(requested, fmt.Errorf("invalid segment %q", seg)) + case isReservedName(s): + return "", "", "", escapeError(requested, fmt.Errorf("reserved name %q", s)) + } + cleaned = append(cleaned, s) + } + if len(cleaned) == 0 { + return "", "", "", escapeError(requested, errors.New("empty path")) + } + + base = cleaned[len(cleaned)-1] + dir = path.Join(cleaned[:len(cleaned)-1]...) + ext = path.Ext(base) + if ext == base { + // Dotfile-style name (".name"): the whole segment is the base. + ext = "" + } + base = strings.TrimSuffix(base, ext) + // Re-trim: stripping the extension can expose trailing dots/spaces + // ("photo..png" -> "photo."), which Windows silently drops. + if base = strings.TrimRight(base, ". "); base == "" { + return "", "", "", escapeError(requested, errors.New("empty filename")) + } + return dir, base, ext, nil +} + +// isAbsolutePath detects rooted paths without filepath.IsAbs, which only +// recognizes drive-letter paths when compiled for Windows. +func isAbsolutePath(p string) bool { + if strings.HasPrefix(p, "/") || strings.HasPrefix(p, `\`) { + return true + } + if len(p) >= 2 && p[1] == ':' && + (('a' <= p[0] && p[0] <= 'z') || ('A' <= p[0] && p[0] <= 'Z')) { + return true + } + return false +} + +// sanitizeSegment replaces characters that are unrepresentable in Windows +// file names (and control characters) with '-', and trims trailing dots +// and spaces, which Windows silently strips — that would break the "exact +// final path" contract. +func sanitizeSegment(seg string) string { + s := strings.Map(func(r rune) rune { + if r < 0x20 || strings.ContainsRune(`<>:"|?*`, r) { + return '-' + } + return r + }, seg) + s = strings.TrimRight(s, ". ") + return strings.TrimLeft(s, " ") +} + +// windowsReservedNames lists device names Windows refuses as file names, +// with or without an extension (CON.png is as unusable as CON). Rejected +// on every platform so a generated tree stays usable on Windows checkouts. +var windowsReservedNames = map[string]bool{ + "CON": true, "PRN": true, "AUX": true, "NUL": true, + "COM1": true, "COM2": true, "COM3": true, "COM4": true, "COM5": true, + "COM6": true, "COM7": true, "COM8": true, "COM9": true, + "LPT1": true, "LPT2": true, "LPT3": true, "LPT4": true, "LPT5": true, + "LPT6": true, "LPT7": true, "LPT8": true, "LPT9": true, +} + +func isReservedName(segment string) bool { + name, _, _ := strings.Cut(segment, ".") + return windowsReservedNames[strings.ToUpper(name)] +} + +// finalExtension picks the filename extension: the MIME-derived one when +// the type is known, otherwise the requested one (".bin" when neither +// exists). corrected is true only when a requested extension conflicted +// with the MIME type and was replaced — the caller should tell the user. +func finalExtension(requestedExt, mimeType string) (ext string, corrected bool) { + if mt, _, err := mime.ParseMediaType(mimeType); err == nil { + mimeType = mt + } + actual := extensionForMIME(mimeType) + switch { + case actual == "": + if requestedExt == "" { + return ".bin", false + } + return requestedExt, false + case requestedExt == "": + return actual, false + case extensionMatchesMIME(requestedExt, mimeType, actual): + return requestedExt, false + default: + return actual, true + } +} + +// knownExtensions pins the extension for MIME types generated media +// commonly uses. mime.ExtensionsByType returns OS-dependent, sometimes +// surprising results (e.g. "image/jpeg" resolving to ".jfif" ahead of +// ".jpg" on some systems), so common types are pinned and the system mime +// database is only a fallback. +var knownExtensions = map[string]string{ + "image/png": ".png", + "image/jpeg": ".jpg", + "image/webp": ".webp", + "image/gif": ".gif", +} + +// extensionForMIME maps a normalized MIME type to an extension (with +// leading dot), or "" when the type is empty or unknown. +func extensionForMIME(mimeType string) string { + if ext, ok := knownExtensions[mimeType]; ok { + return ext + } + if mimeType == "" { + return "" + } + exts, err := mime.ExtensionsByType(mimeType) + if err != nil || len(exts) == 0 { + return "" + } + return exts[0] +} + +// extensionMatchesMIME reports whether requestedExt is an acceptable +// spelling for mimeType (e.g. ".jpeg" for image/jpeg), so valid variants +// are kept as requested instead of being noisily "corrected". +func extensionMatchesMIME(requestedExt, mimeType, actual string) bool { + req := strings.ToLower(requestedExt) + if req == actual { + return true + } + exts, err := mime.ExtensionsByType(mimeType) + if err != nil { + return false + } + return slices.Contains(exts, req) +} + +func escapeError(requested string, cause error) error { + return fmt.Errorf("%w: %q: %w", ErrPathEscape, requested, cause) +} + +// isPathEscape reports whether err is os.Root rejecting a path that +// resolves outside its root via a symlink — the one escape vector the +// lexical checks in splitRequestedPath cannot see. os.Root exports no +// sentinel for it as of Go 1.26, so this matches the PathError message it +// documents ("path escapes from parent"); prefer errors.Is against a +// stdlib sentinel if a future release adds one. +func isPathEscape(err error) bool { + var pathErr *fs.PathError + return errors.As(err, &pathErr) && pathErr.Err != nil && pathErr.Err.Error() == "path escapes from parent" +} diff --git a/pkg/workspacemedia/writer_test.go b/pkg/workspacemedia/writer_test.go new file mode 100644 index 0000000000..6cfba81bbd --- /dev/null +++ b/pkg/workspacemedia/writer_test.go @@ -0,0 +1,291 @@ +package workspacemedia + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "runtime" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +var pngData = []byte("fake-png-bytes") + +func readWorkspaceFile(t *testing.T, root, rel string) []byte { + t.Helper() + data, err := os.ReadFile(filepath.Join(root, filepath.FromSlash(rel))) + require.NoError(t, err) + return data +} + +func TestWrite_PlainName(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "sunshine.png", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "sunshine.png", res.RelPath) + assert.False(t, res.ExtensionCorrected) + assert.Equal(t, pngData, readWorkspaceFile(t, root, res.RelPath)) +} + +func TestWrite_SanitizesFilename(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, `my : "image?".png`, pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "my -cool-- -image--.png", res.RelPath) + assert.Equal(t, pngData, readWorkspaceFile(t, root, res.RelPath)) +} + +func TestWrite_TrimsWindowsUnsafeTrailingChars(t *testing.T) { + root := t.TempDir() + + // Windows silently strips trailing dots/spaces; the writer trims them + // up front so the returned path is exactly what exists on disk. + res, err := Write(root, "photo...png", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "photo.png", res.RelPath) +} + +func TestWrite_CorrectsConflictingExtension(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "sunshine.jpg", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "sunshine.png", res.RelPath) + assert.True(t, res.ExtensionCorrected) + assert.Equal(t, ".jpg", res.RequestedExtension) + assert.Equal(t, pngData, readWorkspaceFile(t, root, res.RelPath)) +} + +func TestWrite_KeepsMatchingExtensionVariant(t *testing.T) { + root := t.TempDir() + + for _, name := range []string{"a.jpeg", "b.JPG"} { + res, err := Write(root, name, []byte("jpeg"), "image/jpeg") + require.NoError(t, err) + assert.Equal(t, name, res.RelPath) + assert.False(t, res.ExtensionCorrected) + } +} + +func TestWrite_NormalizesMIMEParameters(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "pic", pngData, "image/png; some=param") + require.NoError(t, err) + assert.Equal(t, "pic.png", res.RelPath) +} + +func TestWrite_ExtensionlessName(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "sunshine", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "sunshine.png", res.RelPath) + assert.False(t, res.ExtensionCorrected) +} + +func TestWrite_UnknownMIME(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "blob", []byte("data"), "") + require.NoError(t, err) + assert.Equal(t, "blob.bin", res.RelPath) + assert.False(t, res.ExtensionCorrected) + + // An unknown MIME type cannot contradict a requested extension: keep it. + res, err = Write(root, "notes.dat", []byte("data"), "application/x-mystery") + require.NoError(t, err) + assert.Equal(t, "notes.dat", res.RelPath) + assert.False(t, res.ExtensionCorrected) +} + +func TestWrite_GenericFallbackName(t *testing.T) { + root := t.TempDir() + + for i, want := range []string{"generated.png", "generated-1.png", "generated-2.png"} { + res, err := Write(root, "generated", pngData, "image/png") + require.NoError(t, err, "write %d", i) + assert.Equal(t, want, res.RelPath) + } +} + +func TestWrite_NestedSubdirsCreated(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "images/out/pic.png", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "images/out/pic.png", res.RelPath) + assert.Equal(t, pngData, readWorkspaceFile(t, root, res.RelPath)) + + if runtime.GOOS != "windows" { // modes are ACL-inherited on Windows + dirInfo, err := os.Stat(filepath.Join(root, "images", "out")) + require.NoError(t, err) + assert.Zero(t, dirInfo.Mode().Perm()&^0o755, "dir mode %v exceeds 0755", dirInfo.Mode().Perm()) + + fileInfo, err := os.Stat(filepath.Join(root, "images", "out", "pic.png")) + require.NoError(t, err) + assert.Zero(t, fileInfo.Mode().Perm()&^0o644, "file mode %v exceeds 0644", fileInfo.Mode().Perm()) + } +} + +func TestWrite_CurrentDirSegmentsIgnored(t *testing.T) { + root := t.TempDir() + + res, err := Write(root, "./images/./pic.png", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "images/pic.png", res.RelPath) +} + +func TestWrite_RejectsEscapingAndInvalidPaths(t *testing.T) { + root := t.TempDir() + + for _, requested := range []string{ + "/etc/passwd", + `\evil.png`, + `C:\evil.png`, + "c:/evil.png", + "..", + "../pic.png", + `..\pic.png`, + "a/../pic.png", + "a/../../pic.png", + "", + ".", + "...", + " ", + "a/.../b.png", + "CON", + "con.png", + "images/NUL/pic.png", + "lpt1.txt", + } { + t.Run(fmt.Sprintf("%q", requested), func(t *testing.T) { + _, err := Write(root, requested, pngData, "image/png") + require.ErrorIs(t, err, ErrPathEscape) + }) + } + + entries, err := os.ReadDir(root) + require.NoError(t, err) + assert.Empty(t, entries, "rejected paths must not leave files behind") +} + +func TestWrite_RejectsSymlinkedDirEscape(t *testing.T) { + root := t.TempDir() + outside := t.TempDir() + require.NoError(t, os.Symlink(outside, filepath.Join(root, "link"))) + + for _, requested := range []string{"link/pic.png", "link/sub/pic.png"} { + _, err := Write(root, requested, pngData, "image/png") + require.ErrorIs(t, err, ErrPathEscape, "requested %q", requested) + } + + entries, err := os.ReadDir(outside) + require.NoError(t, err) + assert.Empty(t, entries, "nothing may be written through the symlink") +} + +func TestWrite_NeverWritesThroughSymlinkAtLeaf(t *testing.T) { + root := t.TempDir() + outside := t.TempDir() + target := filepath.Join(outside, "target.png") + require.NoError(t, os.Symlink(target, filepath.Join(root, "pic.png"))) + + // O_EXCL treats the existing symlink as "name taken": the write lands + // on the dash-suffixed sibling and the symlink target is never created. + res, err := Write(root, "pic.png", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "pic-1.png", res.RelPath) + assert.NoFileExists(t, target) +} + +func TestWrite_CollisionGetsDashSuffix(t *testing.T) { + root := t.TempDir() + existing := []byte("do-not-touch") + require.NoError(t, os.WriteFile(filepath.Join(root, "sunshine.png"), existing, 0o644)) + + res, err := Write(root, "sunshine.png", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "sunshine-1.png", res.RelPath) + assert.Equal(t, existing, readWorkspaceFile(t, root, "sunshine.png")) + assert.Equal(t, pngData, readWorkspaceFile(t, root, "sunshine-1.png")) +} + +func TestWrite_CollisionAfterExtensionCorrection(t *testing.T) { + root := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(root, "sunshine.png"), []byte("existing"), 0o644)) + + res, err := Write(root, "sunshine.jpg", pngData, "image/png") + require.NoError(t, err) + assert.Equal(t, "sunshine-1.png", res.RelPath) + assert.True(t, res.ExtensionCorrected) +} + +func TestWrite_ConcurrentSameNameWriters(t *testing.T) { + root := t.TempDir() + const writers = 16 + + results := make([]Result, writers) + errs := make([]error, writers) + var wg sync.WaitGroup + for i := range writers { + wg.Go(func() { + results[i], errs[i] = Write(root, "pic.png", fmt.Appendf(nil, "content-%d", i), "image/png") + }) + } + wg.Wait() + + seen := make(map[string]bool, writers) + for i := range writers { + require.NoError(t, errs[i], "writer %d", i) + assert.False(t, seen[results[i].RelPath], "writer %d got duplicate path %q", i, results[i].RelPath) + seen[results[i].RelPath] = true + assert.Equal(t, fmt.Appendf(nil, "content-%d", i), readWorkspaceFile(t, root, results[i].RelPath)) + } + assert.True(t, seen["pic.png"]) + for i := 1; i < writers; i++ { + assert.True(t, seen[fmt.Sprintf("pic-%d.png", i)]) + } +} + +type failingReader struct{} + +func (failingReader) Read([]byte) (int, error) { return 0, errors.New("stream torn") } + +func TestWrite_FailedWriteRemovesClaimAndKeepsExisting(t *testing.T) { + root := t.TempDir() + existing := []byte("keep-me") + require.NoError(t, os.WriteFile(filepath.Join(root, "pic.png"), existing, 0o644)) + + _, err := write(root, "pic.png", failingReader{}, "image/png") + require.ErrorContains(t, err, "stream torn") + + // The pre-existing file is untouched, the claimed dash-suffixed name is + // removed rather than left as an empty file, and no temp file survives. + assert.Equal(t, existing, readWorkspaceFile(t, root, "pic.png")) + entries, err := os.ReadDir(root) + require.NoError(t, err) + require.Len(t, entries, 1) + assert.Equal(t, "pic.png", entries[0].Name()) +} + +func TestWrite_MissingWorkspaceRootFails(t *testing.T) { + _, err := Write(filepath.Join(t.TempDir(), "gone"), "pic.png", pngData, "image/png") + require.Error(t, err) + assert.NotErrorIs(t, err, ErrPathEscape) +} + +func TestWrite_ParentDirIsFileFails(t *testing.T) { + root := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(root, "images"), []byte("file"), 0o644)) + + _, err := Write(root, "images/pic.png", pngData, "image/png") + require.Error(t, err) + assert.NotErrorIs(t, err, ErrPathEscape) +} From 7d142355f0f9504cc0b0ede865554c542c8ca740 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Arnaud=20He=CC=81ritier?= Date: Mon, 27 Jul 2026 22:13:20 +0200 Subject: [PATCH 3/3] test(#3996): cover mid-write cleanup, file modes, and missing parents partialFailReader writes a genuine partial chunk on its first Read then fails on every subsequent Read, so the mid-write-failure tests reach atomicfile.Write's cleanup path AFTER a real os.CreateTemp'd file already has partial content on disk, rather than relying on a permission trick that prevents the temp file from ever being created. Also cover preserved file modes and creation under missing parent directories. This coverage backs the workspace media writer's reliance on atomicfile.Write for never-overwrite, no-residue semantics. --- pkg/atomicfile/atomicfile_test.go | 93 +++++++++++++++++++++++++++++-- 1 file changed, 87 insertions(+), 6 deletions(-) diff --git a/pkg/atomicfile/atomicfile_test.go b/pkg/atomicfile/atomicfile_test.go index e096a41dac..52698baa6b 100644 --- a/pkg/atomicfile/atomicfile_test.go +++ b/pkg/atomicfile/atomicfile_test.go @@ -1,7 +1,8 @@ -package atomicfile_test +package atomicfile import ( "bytes" + "errors" "os" "path/filepath" "runtime" @@ -9,10 +10,13 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - - "github.com/docker/docker-agent/pkg/atomicfile" ) +// TestWriteCreatesFileWithMode is the plan's restored "new file gets the +// requested mode" regression, inadvertently dropped when this file was +// rewritten to add the mid-write cleanup tests below: Write must create a +// brand-new file with exactly the requested permission bits, not whatever +// natefinch/atomic.WriteFile's own default (umask-derived) mode would be. func TestWriteCreatesFileWithMode(t *testing.T) { t.Parallel() if runtime.GOOS == "windows" { @@ -22,7 +26,7 @@ func TestWriteCreatesFileWithMode(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "secret") - require.NoError(t, atomicfile.Write(path, bytes.NewReader([]byte("hello")), 0o600)) + require.NoError(t, Write(path, bytes.NewReader([]byte("hello")), 0o600)) data, err := os.ReadFile(path) require.NoError(t, err) @@ -33,6 +37,10 @@ func TestWriteCreatesFileWithMode(t *testing.T) { assert.Equal(t, os.FileMode(0o600), info.Mode().Perm()) } +// TestWriteOverwritesAndRetightensMode is the restored counterpart for an +// EXISTING destination with a looser mode (0o644): a replacement Write +// must retighten it to exactly the requested (0o600) bits, not merely +// preserve whatever mode the file already had. func TestWriteOverwritesAndRetightensMode(t *testing.T) { t.Parallel() if runtime.GOOS == "windows" { @@ -43,7 +51,7 @@ func TestWriteOverwritesAndRetightensMode(t *testing.T) { path := filepath.Join(dir, "secret") require.NoError(t, os.WriteFile(path, []byte("old"), 0o644)) - require.NoError(t, atomicfile.Write(path, bytes.NewReader([]byte("new")), 0o600)) + require.NoError(t, Write(path, bytes.NewReader([]byte("new")), 0o600)) data, err := os.ReadFile(path) require.NoError(t, err) @@ -54,11 +62,84 @@ func TestWriteOverwritesAndRetightensMode(t *testing.T) { assert.Equal(t, os.FileMode(0o600), info.Mode().Perm()) } +// TestWriteReturnsErrorForMissingDirectory is the restored regression for +// a missing parent directory: Write must surface the underlying error +// rather than creating the parent itself or panicking. func TestWriteReturnsErrorForMissingDirectory(t *testing.T) { t.Parallel() dir := t.TempDir() path := filepath.Join(dir, "missing", "file") - err := atomicfile.Write(path, bytes.NewReader([]byte("x")), 0o600) + err := Write(path, bytes.NewReader([]byte("x")), 0o600) assert.Error(t, err) } + +// partialFailReader returns chunk on its first Read (so the underlying +// os.CreateTemp'd file genuinely receives that data before anything goes +// wrong) and then fails on every subsequent Read. This reproduces a +// mid-write failure AFTER a real, partially written temporary file +// already exists on disk — the seam the plan requires ("deterministically +// fail only after a controlled temporary/partial output has been +// created"), as opposed to a permission trick that would prevent the +// temp file from ever being created at all. +type partialFailReader struct { + chunk []byte + err error + sent bool +} + +func (r *partialFailReader) Read(p []byte) (int, error) { + if !r.sent { + r.sent = true + n := copy(p, r.chunk) + return n, nil + } + return 0, r.err +} + +// TestWrite_CleansUpTempFileOnMidWriteFailure proves the real production +// cleanup path: after Write's underlying temp file has already received +// partial data, a subsequent read failure must leave no temp file behind +// in the target directory, and the destination path itself must remain +// untouched (never partially written, never created if it didn't already +// exist). +func TestWrite_CleansUpTempFileOnMidWriteFailure(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "artifact.bin") + injectedErr := errors.New("injected mid-write failure") + + err := Write(target, &partialFailReader{chunk: []byte("partial-data"), err: injectedErr}, 0o600) + require.Error(t, err) + assert.Contains(t, err.Error(), injectedErr.Error()) + + _, statErr := os.Stat(target) + assert.True(t, os.IsNotExist(statErr), "the destination file must never be created on a mid-write failure") + + entries, readErr := os.ReadDir(dir) + require.NoError(t, readErr) + assert.Empty(t, entries, "no partial/orphan temp file must remain in the directory after a mid-write failure") +} + +// TestWrite_CleansUpTempFileOnMidWriteFailure_ExistingDestination is the +// same regression when the destination already has prior content: a +// mid-write failure for a REPLACEMENT write must leave the original +// content intact and, again, no temp file residue. +func TestWrite_CleansUpTempFileOnMidWriteFailure_ExistingDestination(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "artifact.bin") + require.NoError(t, os.WriteFile(target, []byte("original"), 0o600)) + injectedErr := errors.New("injected mid-write failure") + + err := Write(target, &partialFailReader{chunk: []byte("partial-data"), err: injectedErr}, 0o600) + require.Error(t, err) + assert.Contains(t, err.Error(), injectedErr.Error()) + + data, readErr := os.ReadFile(target) + require.NoError(t, readErr) + assert.Equal(t, "original", string(data), "the original destination content must survive a failed replacement") + + entries, err := os.ReadDir(dir) + require.NoError(t, err) + require.Len(t, entries, 1, "no partial/orphan temp file must remain alongside the untouched destination") + assert.Equal(t, "artifact.bin", entries[0].Name()) +}