Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 0 additions & 52 deletions arkruntime/lib/environments/workdir.go

This file was deleted.

31 changes: 15 additions & 16 deletions arkruntime/lib/environments/worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ type EnvironmentWorkerOptions struct {
EnvironmentID string
// WorkerID 是上报给控制面的 worker 标识,空值时自动生成。
WorkerID string
// Workdir 是每个 session 工作目录的根目录,空值时使用进程当前目录。
// Workdir session 工具使用的工作目录,空值时使用进程当前目录。
Workdir string
// UnrestrictedPaths 控制文件工具是否允许访问 Workdir 之外的路径。
UnrestrictedPaths bool
Expand Down Expand Up @@ -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)
}
Expand All @@ -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
}
Expand Down Expand Up @@ -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)
Expand All @@ -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)
}
Expand Down Expand Up @@ -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 {
Expand Down
13 changes: 13 additions & 0 deletions arkruntime/lib/environments/worker_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
20 changes: 17 additions & 3 deletions arkruntime/selfhosted/envinit/initializer.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 创建环境初始化器。
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
13 changes: 13 additions & 0 deletions arkruntime/selfhosted/envinit/initializer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
36 changes: 35 additions & 1 deletion arkruntime/selfhosted/tool_result_store.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{}
Expand Down
45 changes: 44 additions & 1 deletion arkruntime/selfhosted/tool_result_store_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
}
}
Loading