diff --git a/arkruntime/lib/environments/workdir.go b/arkruntime/lib/environments/workdir.go deleted file mode 100644 index 22ff076..0000000 --- a/arkruntime/lib/environments/workdir.go +++ /dev/null @@ -1,52 +0,0 @@ -// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. -// SPDX-License-Identifier: Apache-2.0 -package environments - -import ( - "crypto/sha256" - "encoding/hex" - "errors" - "path/filepath" - "strings" -) - -func sessionWorkdir(root, sessionID string) (string, error) { - if root == "" { - return "", errors.New("workdir root must not be empty") - } - name := sessionWorkdirName(sessionID) - target := filepath.Join(root, name) - rel, err := filepath.Rel(root, target) - if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) || filepath.IsAbs(rel) { - return "", errors.New("session workdir escapes root") - } - return target, nil -} - -func sessionWorkdirName(sessionID string) string { - if isSafeWorkdirName(sessionID) { - return sessionID - } - sum := sha256.Sum256([]byte(sessionID)) - return "session-" + hex.EncodeToString(sum[:]) -} - -func isSafeWorkdirName(name string) bool { - if name == "" || name == "." || name == ".." { - return false - } - clean := filepath.Clean(name) - if clean != name || strings.Contains(name, "/") || strings.Contains(name, `\`) { - return false - } - for _, c := range name { - if c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' { - continue - } - if c == '.' || c == '_' || c == '-' { - continue - } - return false - } - return true -} diff --git a/arkruntime/lib/environments/worker.go b/arkruntime/lib/environments/worker.go index 4070cdc..da74ef7 100644 --- a/arkruntime/lib/environments/worker.go +++ b/arkruntime/lib/environments/worker.go @@ -27,7 +27,7 @@ type EnvironmentWorkerOptions struct { EnvironmentID string // WorkerID 是上报给控制面的 worker 标识,空值时自动生成。 WorkerID string - // Workdir 是每个 session 工作目录的根目录,空值时使用进程当前目录。 + // Workdir 是 session 工具使用的工作目录,空值时使用进程当前目录。 Workdir string // UnrestrictedPaths 控制文件工具是否允许访问 Workdir 之外的路径。 UnrestrictedPaths bool @@ -115,7 +115,7 @@ func (w *EnvironmentWorker) runClaimedWork(ctx context.Context, item selfhosted. logger.Warn("handle work failed", "work_id", item.ID, "err", err) return } - if err := w.handleItem(ctx, work, false); err != nil && + if err := w.handleItem(ctx, work); err != nil && !isBenignWorkerExit(err) { logger.Warn("handle work failed", "work_id", work.ID, "session_id", work.SessionID, "err", err) } @@ -138,18 +138,18 @@ func (w *EnvironmentWorker) HandleItem(ctx context.Context, opts HandleItemOptio if err != nil { return err } - if err := w.handleItem(ctx, work, true); err != nil && + if err := w.handleItem(ctx, work); err != nil && !isBenignWorkerExit(err) { return err } return nil } -func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork, useWorkdirAsSession bool) (err error) { +func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork) (err error) { if w.api == nil { return errors.New("environments: API is required") } - workdir, err := w.workdirFor(work.SessionID, useWorkdirAsSession) + workdir, err := w.workdir() if err != nil { return err } @@ -196,7 +196,13 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork, us initOpts := toolEnv.InitOptions() initOpts.Workdir = workdir initOpts.Logger = w.opts.Logger - if err := envinit.New(api, initOpts).Setup(workCtx, session); err != nil { + initializer := envinit.New(api, initOpts) + defer func() { + if cleanupErr := initializer.Cleanup(); cleanupErr != nil { + logger.Warn("cleanup session skills failed", "err", cleanupErr) + } + }() + if err := initializer.Setup(workCtx, session); err != nil { return fmt.Errorf("setup session environment: %w", err) } tools, owned, err := w.toolsFor(toolEnv) @@ -206,7 +212,7 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork, us if owned { defer func() { _ = tools.Close() }() } - store, err := selfhosted.NewFileToolResultStore(workdir) + store, err := selfhosted.NewSessionFileToolResultStore(workdir, work.SessionID) if err != nil { return fmt.Errorf("create tool result store: %w", err) } @@ -287,19 +293,12 @@ func (w *EnvironmentWorker) effectiveToolTimeout() time.Duration { return 0 } -func (w *EnvironmentWorker) workdirFor(sessionID string, useWorkdirAsSession bool) (string, error) { +func (w *EnvironmentWorker) workdir() (string, error) { root := w.opts.Workdir if root == "" { root = "." } - absRoot, err := filepath.Abs(root) - if err != nil { - return "", err - } - if useWorkdirAsSession { - return absRoot, nil - } - return sessionWorkdir(absRoot, sessionID) + return filepath.Abs(root) } func (w *EnvironmentWorker) stopItem(api selfhosted.API, work claimedWork) error { diff --git a/arkruntime/lib/environments/worker_test.go b/arkruntime/lib/environments/worker_test.go index adfcee4..e77bdfb 100644 --- a/arkruntime/lib/environments/worker_test.go +++ b/arkruntime/lib/environments/worker_test.go @@ -4,6 +4,19 @@ package environments import "testing" +func TestEnvironmentWorkerUsesConfiguredWorkdir(t *testing.T) { + root := t.TempDir() + worker := NewEnvironmentWorker(nil, EnvironmentWorkerOptions{Workdir: root}) + + got, err := worker.workdir() + if err != nil { + t.Fatal(err) + } + if got != root { + t.Fatalf("workdir = %q, want %q", got, root) + } +} + func TestHandleItemOptionsBuildsClaimedWork(t *testing.T) { got, err := claimedWorkFromOptions(HandleItemOptions{ WorkID: testWorkID, diff --git a/arkruntime/selfhosted/envinit/initializer.go b/arkruntime/selfhosted/envinit/initializer.go index 898ee7d..90142a7 100644 --- a/arkruntime/selfhosted/envinit/initializer.go +++ b/arkruntime/selfhosted/envinit/initializer.go @@ -39,9 +39,10 @@ type Options struct { // Initializer 执行 workdir 与 skills 初始化。 type Initializer struct { - api selfhosted.API - opts Options - logger *selfhostedlog.Logger + api selfhosted.API + opts Options + logger *selfhostedlog.Logger + installedSkillDirs []string } // New 创建环境初始化器。 @@ -83,6 +84,18 @@ func (i *Initializer) Setup(ctx context.Context, session *selfhosted.Session) er return nil } +// Cleanup 删除本次初始化成功安装的 skill 目录。 +func (i *Initializer) Cleanup() error { + var errs []error + for _, dir := range i.installedSkillDirs { + if err := os.RemoveAll(dir); err != nil { + errs = append(errs, fmt.Errorf("remove installed skill %s: %w", dir, err)) + } + } + i.installedSkillDirs = nil + return errors.Join(errs...) +} + func (i *Initializer) installSkill(ctx context.Context, sessionID string, skill selfhosted.SkillRef) error { if resolver, ok := i.api.(selfhosted.SkillResolver); ok && strings.TrimSpace(skill.IDValue()) != "" { resolved, err := resolver.ResolveSkill(ctx, skill) @@ -133,6 +146,7 @@ func (i *Initializer) installSkill(ctx context.Context, sessionID string, skill if err != nil { return fmt.Errorf("commit skill %s: %w", name, err) } + i.installedSkillDirs = append(i.installedSkillDirs, target) committed = true if source != tmp { _ = os.RemoveAll(tmp) diff --git a/arkruntime/selfhosted/envinit/initializer_test.go b/arkruntime/selfhosted/envinit/initializer_test.go index 546bf81..eb88569 100644 --- a/arkruntime/selfhosted/envinit/initializer_test.go +++ b/arkruntime/selfhosted/envinit/initializer_test.go @@ -101,6 +101,19 @@ func TestSetupInstallsZipSkill(t *testing.T) { if string(got) != "hello" { t.Fatalf("skill = %q", got) } + retained := filepath.Join(root, "skills", "retained") + if err := os.MkdirAll(retained, 0o755); err != nil { + t.Fatal(err) + } + if err := init.Cleanup(); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "demo")); !os.IsNotExist(err) { + t.Fatalf("installed skill should be removed, err=%v", err) + } + if _, err := os.Stat(retained); err != nil { + t.Fatalf("unmanaged skills entry should remain: %v", err) + } } func TestReplaceSkillDirRollsBackWhenCommitFails(t *testing.T) { diff --git a/arkruntime/selfhosted/tool_result_store.go b/arkruntime/selfhosted/tool_result_store.go index df2b1d1..f5d21bb 100644 --- a/arkruntime/selfhosted/tool_result_store.go +++ b/arkruntime/selfhosted/tool_result_store.go @@ -35,16 +35,50 @@ type fileToolResultRecord struct { // NewFileToolResultStore 在 workdir 下创建 tool result 持久化 store。 func NewFileToolResultStore(workdir string) (*FileToolResultStore, error) { + return newFileToolResultStore(workdir, "") +} + +// NewSessionFileToolResultStore 在 workdir 下创建按 session 隔离的 tool result 持久化 store。 +func NewSessionFileToolResultStore(workdir, sessionID string) (*FileToolResultStore, error) { + if sessionID == "" { + return nil, errors.New("session id must not be empty") + } + return newFileToolResultStore(workdir, toolResultStoreSessionDir(sessionID)) +} + +func newFileToolResultStore(workdir, sessionDir string) (*FileToolResultStore, error) { if workdir == "" { return nil, errors.New("workdir must not be empty") } - dir := filepath.Join(workdir, ".ma_self_host_worker", "tool_ledger") + dir := filepath.Join(workdir, ".ma_self_hosted_worker", "tool_ledger") + if sessionDir != "" { + dir = filepath.Join(dir, sessionDir) + } if err := os.MkdirAll(dir, 0o700); err != nil { return nil, fmt.Errorf("create tool result store: %w", err) } return &FileToolResultStore{dir: dir}, nil } +func toolResultStoreSessionDir(sessionID string) string { + if sessionID != "." && sessionID != ".." { + safe := true + for _, c := range sessionID { + if c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' || + c == '.' || c == '_' || c == '-' { + continue + } + safe = false + break + } + if safe { + return sessionID + } + } + sum := sha256.Sum256([]byte(sessionID)) + return "session-" + hex.EncodeToString(sum[:]) +} + // Recover 恢复未回写的 tool result 和已经完成回写的 tool_use。 func (s *FileToolResultStore) Recover() (map[string]Event, map[string]bool, error) { pending := map[string]Event{} diff --git a/arkruntime/selfhosted/tool_result_store_test.go b/arkruntime/selfhosted/tool_result_store_test.go index 10b5e2a..1e97dcb 100644 --- a/arkruntime/selfhosted/tool_result_store_test.go +++ b/arkruntime/selfhosted/tool_result_store_test.go @@ -9,10 +9,15 @@ import ( ) func TestFileToolResultStoreRecoverIgnoresInterruptedTempFile(t *testing.T) { - store, err := NewFileToolResultStore(t.TempDir()) + workdir := t.TempDir() + store, err := NewFileToolResultStore(workdir) if err != nil { t.Fatal(err) } + wantDir := filepath.Join(workdir, ".ma_self_hosted_worker", "tool_ledger") + if store.dir != wantDir { + t.Fatalf("store dir = %q, want %q", store.dir, wantDir) + } tempPath := filepath.Join(store.dir, ".tool-result-interrupted") if err := os.WriteFile(tempPath, []byte("{"), 0o600); err != nil { t.Fatal(err) @@ -57,3 +62,41 @@ func TestFileToolResultStoreRecoverUsesPersistedCallID(t *testing.T) { t.Fatalf("tool_use_id=%q", result.ToolUseID) } } + +func TestSessionFileToolResultStoreIsolatesSessions(t *testing.T) { + workdir := t.TempDir() + first, err := NewSessionFileToolResultStore(workdir, "session-a") + if err != nil { + t.Fatal(err) + } + second, err := NewSessionFileToolResultStore(workdir, "session-b") + if err != nil { + t.Fatal(err) + } + wantDir := filepath.Join(workdir, ".ma_self_hosted_worker", "tool_ledger", "session-a") + if first.dir != wantDir { + t.Fatalf("store dir = %q, want %q", first.dir, wantDir) + } + if _, err := first.Begin("call-id", Event{ID: "event-id", Type: EventTypeAgentToolUse}); err != nil { + t.Fatal(err) + } + pending, processed, err := second.Recover() + if err != nil { + t.Fatal(err) + } + if len(pending) != 0 || len(processed) != 0 { + t.Fatalf("pending=%v processed=%v", pending, processed) + } +} + +func TestSessionFileToolResultStoreSanitizesSessionID(t *testing.T) { + workdir := t.TempDir() + store, err := NewSessionFileToolResultStore(workdir, "../../outside") + if err != nil { + t.Fatal(err) + } + base := filepath.Join(workdir, ".ma_self_hosted_worker", "tool_ledger") + if filepath.Dir(store.dir) != base { + t.Fatalf("store dir escaped base: %q", store.dir) + } +}