diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index b86b3670..c9ee5d46 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -1071,7 +1071,7 @@ jobs: - name: Pull Test Images run: | docker pull yaoapp/tai-sandbox-test:latest || true - docker pull yaoapp/tai:1.2.0 + docker pull yaoapp/tai:latest docker pull alpine:latest - name: Install k3d @@ -1088,7 +1088,7 @@ jobs: docker run -d --name tai-docker \ -v /var/run/docker.sock:/var/run/docker.sock \ -p 19100:19100 -p 8099:8099 -p 12375:12375 -p 16080:16080 \ - yaoapp/tai:1.2.0 server \ + yaoapp/tai:latest server \ -grpc 0.0.0.0:19100 -http 0.0.0.0:8099 -vnc 0.0.0.0:16080 -docker 0.0.0.0:12375 for i in $(seq 1 30); do @@ -1142,7 +1142,7 @@ jobs: -v /tmp/kubeconfig-tai-k8s.yml:/etc/tai/kubeconfig.yml:ro \ -e TAI_K8S_UPSTREAM="tcp://${K3D_IP}:6443" \ -e TAI_KUBECONFIG=/etc/tai/kubeconfig.yml \ - yaoapp/tai:1.2.0 server \ + yaoapp/tai:latest server \ -grpc 0.0.0.0:19100 -http 0.0.0.0:8099 -vnc 0.0.0.0:16080 -k8s 0.0.0.0:16443 for i in $(seq 1 30); do diff --git a/.github/workflows/unit-test.yml b/.github/workflows/unit-test.yml index 661b0dad..c6765878 100644 --- a/.github/workflows/unit-test.yml +++ b/.github/workflows/unit-test.yml @@ -782,7 +782,7 @@ jobs: - name: Pull Test Images run: | docker pull yaoapp/tai-sandbox-test:latest || true - docker pull yaoapp/tai:1.2.0 + docker pull yaoapp/tai:latest docker pull alpine:latest - name: Install k3d @@ -799,7 +799,7 @@ jobs: docker run -d --name tai-docker \ -v /var/run/docker.sock:/var/run/docker.sock \ -p 19100:19100 -p 8099:8099 -p 12375:12375 -p 16080:16080 \ - yaoapp/tai:1.2.0 server \ + yaoapp/tai:latest server \ -grpc 0.0.0.0:19100 -http 0.0.0.0:8099 -vnc 0.0.0.0:16080 -docker 0.0.0.0:12375 for i in $(seq 1 30); do @@ -853,7 +853,7 @@ jobs: -v /tmp/kubeconfig-tai-k8s.yml:/etc/tai/kubeconfig.yml:ro \ -e TAI_K8S_UPSTREAM="tcp://${K3D_IP}:6443" \ -e TAI_KUBECONFIG=/etc/tai/kubeconfig.yml \ - yaoapp/tai:1.2.0 server \ + yaoapp/tai:latest server \ -grpc 0.0.0.0:19100 -http 0.0.0.0:8099 -vnc 0.0.0.0:16080 -k8s 0.0.0.0:16443 for i in $(seq 1 30); do diff --git a/.gitignore b/.gitignore index 24fe2d1c..6f3a811a 100644 --- a/.gitignore +++ b/.gitignore @@ -74,4 +74,5 @@ tg-login tg-send registry/data/ registry/manager/DESIGN*.md -tai/testdata/ \ No newline at end of file +tai/testdata/ +agent/sandbox/docs/*.md \ No newline at end of file diff --git a/Makefile b/Makefile index 7844d513..cd762aa4 100644 --- a/Makefile +++ b/Makefile @@ -10,11 +10,11 @@ NOW := $(shell date +"%FT%T%z") OS := $(shell uname) # ROOT_DIR := $(shell dirname $(realpath $(firstword $(MAKEFILE_LIST)))) -TESTFOLDER := $(shell $(GO) list ./... | grep -vE 'examples|openai|aigc|neo|twilio|share*|registry' | awk '!/\/tests\// || /openapi\/tests/') +TESTFOLDER := $(shell $(GO) list ./... | grep -vE 'examples|openai|aigc|neo|twilio|share*|registry|agent/sandbox/v2' | awk '!/\/tests\// || /openapi\/tests/' | grep -vE 'openapi/tests/(nodes|sandbox|workspace)') # Core tests (exclude AI-related: agent, aigc, openai, KB, sandbox, registry, grpc, and integrations which require external services) -TESTFOLDER_CORE := $(shell $(GO) list ./... | grep -vE 'examples|openai|aigc|neo|twilio|share*|agent|kb|sandbox|integrations|registry|tai|grpc' | awk '!/\/tests\// || /openapi\/tests/') -# Agent tests (agent, aigc) - exclude agent/search/handlers/web (requires external API keys) and robot packages (tested in robot job) -TESTFOLDER_AGENT := $(shell $(GO) list ./agent/... ./aigc/... | grep -vE 'agent/search/handlers/web|agent/robot/') +TESTFOLDER_CORE := $(shell $(GO) list ./... | grep -vE 'examples|openai|aigc|neo|twilio|share*|agent|kb|sandbox|integrations|registry|tai|grpc' | awk '!/\/tests\// || /openapi\/tests/' | grep -vE 'openapi/tests/(nodes|sandbox|workspace)') +# Agent tests (agent, aigc) - exclude agent/search/handlers/web (requires external API keys), robot packages (tested in robot job), and agent/sandbox/v2 (WIP, has its own job) +TESTFOLDER_AGENT := $(shell $(GO) list ./agent/... ./aigc/... | grep -vE 'agent/search/handlers/web|agent/robot/|agent/sandbox/v2') # KB tests (kb) TESTFOLDER_KB := $(shell $(GO) list ./kb/...) # Robot tests (agent/robot/... packages, excluding events/integrations which require Telegram etc.) diff --git a/agent/assistant/agent.go b/agent/assistant/agent.go index fca12ab8..9e49226b 100644 --- a/agent/assistant/agent.go +++ b/agent/assistant/agent.go @@ -14,6 +14,8 @@ import ( "github.com/yaoapp/yao/agent/llm" "github.com/yaoapp/yao/agent/output/message" agentsandbox "github.com/yaoapp/yao/agent/sandbox" + sandboxTypes "github.com/yaoapp/yao/agent/sandbox/v2/types" + infraV2 "github.com/yaoapp/yao/sandbox/v2" ) // Stream stream the agent @@ -163,7 +165,25 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa var sandboxExecutor agentsandbox.Executor var sandboxCleanup func() var sandboxLoadingMsgID string - if ast.HasSandbox() { + + // V2 sandbox state + var v2Runner sandboxTypes.Runner + var v2Computer infraV2.Computer + var v2LoadingMsgID string + + if ast.HasSandboxV2() { + ctx.Logger.Phase("Sandbox V2") + var err error + var v2Cleanup func() + v2Runner, v2Computer, v2Cleanup, v2LoadingMsgID, err = ast.initSandboxV2(ctx, opts) + if err != nil { + ast.traceAgentFail(agentNode, err) + ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err) + return nil, err + } + sandboxCleanup = v2Cleanup + ctx.Logger.PhaseComplete("Sandbox V2") + } else if ast.HasSandbox() { ctx.Logger.Phase("Sandbox") var err error sandboxExecutor, sandboxCleanup, sandboxLoadingMsgID, err = ast.initSandbox(ctx, opts) @@ -289,8 +309,17 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa // Execute the LLM streaming call // Choose between sandbox execution or direct LLM execution - if ast.HasSandbox() { - // Sandbox execution path (Claude CLI, Cursor CLI, etc.) + if ast.HasSandboxV2() && v2Runner != nil && v2Computer != nil && v2Runner.Name() != "yao" { + // V2 Sandbox execution path (non-yao runners replace LLM.Stream) + completionResponse, err = ast.executeSandboxV2Stream(ctx, completionMessages, agentNode, streamHandler, v2Runner, v2Computer, v2LoadingMsgID) + } else if ast.HasSandboxV2() && v2Runner != nil && v2Runner.Name() == "yao" { + // V2 yao runner: Prepare is done, close loading, fall through to LLM + if v2LoadingMsgID != "" { + closeLoadingV2(ctx, v2LoadingMsgID, "") + } + completionResponse, err = ast.executeLLMStream(ctx, completionMessages, completionOptions, agentNode, streamHandler, opts) + } else if ast.HasSandbox() { + // V1 Sandbox execution path (Claude CLI, Cursor CLI, etc.) completionResponse, err = ast.executeSandboxStream(ctx, completionMessages, agentNode, streamHandler, sandboxExecutor, sandboxLoadingMsgID) } else { // Direct LLM execution path diff --git a/agent/assistant/load.go b/agent/assistant/load.go index 0333ac58..75092b47 100644 --- a/agent/assistant/load.go +++ b/agent/assistant/load.go @@ -12,8 +12,10 @@ import ( "github.com/yaoapp/gou/fs" "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/i18n" + sandboxTypes "github.com/yaoapp/yao/agent/sandbox/v2/types" searchTypes "github.com/yaoapp/yao/agent/search/types" store "github.com/yaoapp/yao/agent/store/types" + "github.com/yaoapp/yao/config" "gopkg.in/yaml.v3" ) @@ -378,7 +380,45 @@ func LoadPath(path string) (*Assistant, error) { return nil, err } data["locales"] = locales - return loadMap(data) + + // V2 sandbox: load standalone sandbox.yao if present (Path A). + sandboxFile := filepath.Join(path, "sandbox.yao") + if has, _ := app.Exists(sandboxFile); has { + absFile := filepath.Join(config.Conf.AppSource, sandboxFile) + sbCfg, sbErr := store.LoadSandboxConfig(absFile) + if sbErr != nil { + return nil, fmt.Errorf("load sandbox.yao: %w", sbErr) + } + data["__sandbox_v2"] = sbCfg + } + + ast, err := loadMap(data) + if err != nil { + return nil, err + } + + // If V2 sandbox was loaded via Path A, assign it now. + if sbCfg, ok := data["__sandbox_v2"].(*sandboxTypes.SandboxConfig); ok && sbCfg != nil { + ast.SandboxV2 = sbCfg + } + + // Compute config hash for V2 sandbox. + if ast.SandboxV2 != nil { + var mcpServers []store.MCPServerConfig + if ast.MCP != nil { + mcpServers = ast.MCP.Servers + } + skillsDir := "" + if ast.Path != "" { + dir := filepath.Join(config.Conf.AppSource, ast.Path, "skills") + if info, e := os.Stat(dir); e == nil && info.IsDir() { + skillsDir = dir + } + } + ast.ConfigHash = store.ComputeConfigHash(ast.SandboxV2, mcpServers, skillsDir) + } + + return ast, nil } func loadMap(data map[string]interface{}) (*Assistant, error) { @@ -721,12 +761,25 @@ func loadMap(data map[string]interface{}) (*Assistant, error) { } // sandbox (for coding agents like Claude CLI, Cursor CLI) - if sandbox, has := data["sandbox"]; has { - sb, err := store.ToSandbox(sandbox) - if err != nil { - return nil, err + // V2 sandbox via independent sandbox.yao is loaded in LoadPath (below). + // This block handles the package.yao embedded "sandbox" field with version dispatch. + if assistant.SandboxV2 == nil { + if sandbox, has := data["sandbox"]; has { + version := extractSandboxVersion(sandbox) + if version == sandboxTypes.SandboxVersionV2 { + sb, err := store.ToSandboxV2(sandbox) + if err != nil { + return nil, err + } + assistant.SandboxV2 = sb + } else { + sb, err := store.ToSandbox(sandbox) + if err != nil { + return nil, err + } + assistant.Sandbox = sb + } } - assistant.Sandbox = sb } // dependencies (name -> version constraint, like npm dependencies) @@ -1036,3 +1089,13 @@ func mergeSearchConfig(base, override *searchTypes.Config) *searchTypes.Config { return &result } + +// extractSandboxVersion tries to read the "version" field from a sandbox config value. +func extractSandboxVersion(v any) string { + if m, ok := v.(map[string]any); ok { + if ver, ok := m["version"].(string); ok { + return ver + } + } + return "" +} diff --git a/agent/assistant/load_test.go b/agent/assistant/load_test.go index a15406dd..818ee334 100644 --- a/agent/assistant/load_test.go +++ b/agent/assistant/load_test.go @@ -594,6 +594,162 @@ func TestLoadSystemAgents(t *testing.T) { }) } +// TestLoadPathSandboxV2 tests loading assistants with V2 sandbox configuration (standalone sandbox.yao) +func TestLoadPathSandboxV2(t *testing.T) { + prepare(t) + defer test.Clean() + + t.Run("OneshotCLI", func(t *testing.T) { + ast, err := assistant.LoadPath("/assistants/tests/sandbox-v2/oneshot-cli") + require.NoError(t, err) + require.NotNil(t, ast) + + assert.Equal(t, "Sandbox V2 Oneshot CLI", ast.Name) + assert.Contains(t, ast.Tags, "SandboxV2") + + // V2 sandbox should be loaded from sandbox.yao + require.NotNil(t, ast.SandboxV2, "SandboxV2 should be loaded") + assert.Equal(t, "2.0", ast.SandboxV2.Version) + assert.Equal(t, "yaoapp/tai-sandbox-claude:latest", ast.SandboxV2.Computer.Image) + assert.Equal(t, "2GB", ast.SandboxV2.Computer.Memory) + assert.Equal(t, float64(2), ast.SandboxV2.Computer.CPUs) + assert.Equal(t, "/workspace", ast.SandboxV2.Computer.WorkDir) + assert.Equal(t, "claude", ast.SandboxV2.Runner.Name) + assert.Equal(t, "cli", ast.SandboxV2.Runner.Mode) + assert.Equal(t, "oneshot", ast.SandboxV2.Lifecycle) + + // Runner options + assert.NotNil(t, ast.SandboxV2.Runner.Options) + assert.Equal(t, float64(5), ast.SandboxV2.Runner.Options["max_turns"]) + + // V1 Sandbox should be nil + assert.Nil(t, ast.Sandbox, "V1 Sandbox should be nil when V2 is present") + + // ConfigHash should be computed + assert.NotEmpty(t, ast.ConfigHash, "ConfigHash should be computed for V2 sandbox") + + // HasSandboxV2 helper + assert.True(t, ast.HasSandboxV2()) + }) + + t.Run("SessionCLI", func(t *testing.T) { + ast, err := assistant.LoadPath("/assistants/tests/sandbox-v2/session-cli") + require.NoError(t, err) + require.NotNil(t, ast) + + require.NotNil(t, ast.SandboxV2) + assert.Equal(t, "session", ast.SandboxV2.Lifecycle) + assert.Equal(t, "10m", ast.SandboxV2.IdleTimeout) + + // Prepare steps + require.Len(t, ast.SandboxV2.Prepare, 1) + assert.Equal(t, "exec", ast.SandboxV2.Prepare[0].Action) + assert.True(t, ast.SandboxV2.Prepare[0].Once) + }) + + t.Run("LongrunningCLI", func(t *testing.T) { + ast, err := assistant.LoadPath("/assistants/tests/sandbox-v2/longrunning-cli") + require.NoError(t, err) + require.NotNil(t, ast) + + require.NotNil(t, ast.SandboxV2) + assert.Equal(t, "longrunning", ast.SandboxV2.Lifecycle) + assert.Equal(t, "15m", ast.SandboxV2.IdleTimeout) + assert.Equal(t, "2h", ast.SandboxV2.MaxLifetime) + assert.Equal(t, "5s", ast.SandboxV2.StopTimeout) + assert.Equal(t, "4GB", ast.SandboxV2.Computer.Memory) + assert.Equal(t, "rw", ast.SandboxV2.Computer.MountMode) + + // Environment + assert.Equal(t, "test", ast.SandboxV2.Environment["NODE_ENV"]) + assert.Equal(t, "longrunning", ast.SandboxV2.Environment["V2_TEST_MODE"]) + + // Secrets + assert.Equal(t, "sandbox-v2-longrunning-secret", ast.SandboxV2.Secrets["TEST_SECRET"]) + + // Prepare steps + require.Len(t, ast.SandboxV2.Prepare, 3) + assert.True(t, ast.SandboxV2.Prepare[2].IgnoreError) + + // MCP (from package.yao) + require.NotNil(t, ast.MCP) + require.Len(t, ast.MCP.Servers, 1) + assert.Equal(t, "echo", ast.MCP.Servers[0].ServerID) + + // ConfigHash should include MCP servers + hashWithMCP := ast.ConfigHash + assert.NotEmpty(t, hashWithMCP) + }) + + t.Run("HooksOnly_YaoRunner", func(t *testing.T) { + ast, err := assistant.LoadPath("/assistants/tests/sandbox-v2/hooks-only") + require.NoError(t, err) + require.NotNil(t, ast) + + require.NotNil(t, ast.SandboxV2) + assert.Equal(t, "yao", ast.SandboxV2.Runner.Name) + assert.Equal(t, "oneshot", ast.SandboxV2.Lifecycle) + assert.Equal(t, float64(1), ast.SandboxV2.Computer.CPUs) + + // Runner mode should be empty (yao runner ignores mode) + assert.Empty(t, ast.SandboxV2.Runner.Mode) + }) + + t.Run("FullPrepare", func(t *testing.T) { + ast, err := assistant.LoadPath("/assistants/tests/sandbox-v2/full-prepare") + require.NoError(t, err) + require.NotNil(t, ast) + + require.NotNil(t, ast.SandboxV2) + assert.Equal(t, "session", ast.SandboxV2.Lifecycle) + assert.Equal(t, "15m", ast.SandboxV2.IdleTimeout) + + // Prepare: 5 steps with mixed actions + require.Len(t, ast.SandboxV2.Prepare, 5) + assert.Equal(t, "copy", ast.SandboxV2.Prepare[0].Action) + assert.Equal(t, "skills", ast.SandboxV2.Prepare[0].Src) + assert.Equal(t, "~/.claude/skills", ast.SandboxV2.Prepare[0].Dst) + assert.Equal(t, "exec", ast.SandboxV2.Prepare[1].Action) + assert.True(t, ast.SandboxV2.Prepare[1].Once) + assert.True(t, ast.SandboxV2.Prepare[3].IgnoreError) + + // Environment + Secrets + assert.Equal(t, "full", ast.SandboxV2.Environment["V2_PREPARE_TEST"]) + assert.Equal(t, "v2-full-prepare-key", ast.SandboxV2.Secrets["TEST_API_KEY"]) + + // Runner options + assert.Equal(t, "acceptEdits", ast.SandboxV2.Runner.Options["permission_mode"]) + }) + + t.Run("HostMode", func(t *testing.T) { + ast, err := assistant.LoadPath("/assistants/tests/sandbox-v2/host-mode") + require.NoError(t, err) + require.NotNil(t, ast) + + require.NotNil(t, ast.SandboxV2) + // Host mode: no image + assert.Empty(t, ast.SandboxV2.Computer.Image) + assert.Equal(t, "/tmp/yao-sandbox-v2-host-test", ast.SandboxV2.Computer.WorkDir) + assert.Equal(t, "session", ast.SandboxV2.Lifecycle) + }) + + t.Run("ConfigHashDeterministic", func(t *testing.T) { + ast1, err := assistant.LoadPath("/assistants/tests/sandbox-v2/oneshot-cli") + require.NoError(t, err) + ast2, err := assistant.LoadPath("/assistants/tests/sandbox-v2/oneshot-cli") + require.NoError(t, err) + assert.Equal(t, ast1.ConfigHash, ast2.ConfigHash, "same config should produce same hash") + }) + + t.Run("ConfigHashDiffers", func(t *testing.T) { + ast1, err := assistant.LoadPath("/assistants/tests/sandbox-v2/oneshot-cli") + require.NoError(t, err) + ast2, err := assistant.LoadPath("/assistants/tests/sandbox-v2/longrunning-cli") + require.NoError(t, err) + assert.NotEqual(t, ast1.ConfigHash, ast2.ConfigHash, "different configs should produce different hashes") + }) +} + // TestValidate tests the assistant Validate method func TestValidate(t *testing.T) { tests := []struct { diff --git a/agent/assistant/sandbox_v2.go b/agent/assistant/sandbox_v2.go new file mode 100644 index 00000000..4a58b4cc --- /dev/null +++ b/agent/assistant/sandbox_v2.go @@ -0,0 +1,188 @@ +package assistant + +import ( + "fmt" + "os" + "path/filepath" + + "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/i18n" + "github.com/yaoapp/yao/agent/output/message" + sandboxv2 "github.com/yaoapp/yao/agent/sandbox/v2" + sandboxTypes "github.com/yaoapp/yao/agent/sandbox/v2/types" + "github.com/yaoapp/yao/config" + infraV2 "github.com/yaoapp/yao/sandbox/v2" + traceTypes "github.com/yaoapp/yao/trace/types" +) + +// HasSandboxV2 returns true if the assistant has a V2 sandbox configuration. +func (ast *Assistant) HasSandboxV2() bool { + return ast.SandboxV2 != nil +} + +// initSandboxV2 initializes the V2 sandbox: obtains a Computer, gets a Runner, +// runs Prepare, and returns the runner, computer, cleanup closure, loading +// message ID, and any error. +func (ast *Assistant) initSandboxV2(ctx *context.Context, opts *context.Options) ( + sandboxTypes.Runner, infraV2.Computer, func(), string, error, +) { + cfg := ast.SandboxV2 + manager := infraV2.M() + + loadingMsg := &message.Message{ + Type: message.TypeLoading, + Props: map[string]any{ + "message": i18n.T(ctx.Locale, "sandbox.preparing"), + }, + } + loadingMsgID, _ := ctx.SendStream(loadingMsg) + + stdCtx := ctx.Context + + // 1. Resolve connector (before Computer so proxy env vars can be injected). + conn, _, err := ast.GetConnector(ctx, opts) + if err != nil && cfg.Runner.Name != "yao" { + closeLoadingV2(ctx, loadingMsgID, "sandbox.failed") + return nil, nil, nil, "", fmt.Errorf("get connector: %w", err) + } + + // 2. Obtain Computer (passes connector for OPENAI_PROXY_* env injection). + computer, identifier, err := sandboxv2.GetComputer(ctx, cfg, manager, conn) + if err != nil { + closeLoadingV2(ctx, loadingMsgID, "sandbox.failed") + return nil, nil, nil, "", fmt.Errorf("getComputer failed: %w", err) + } + _ = identifier + + // 3. Get Runner. + runner, err := sandboxv2.Get(cfg.Runner.Name) + if err != nil { + sandboxv2.LifecycleAction(stdCtx, cfg, computer, manager) + closeLoadingV2(ctx, loadingMsgID, "sandbox.failed") + return nil, nil, nil, "", fmt.Errorf("get runner %q: %w", cfg.Runner.Name, err) + } + + // 4. Resolve skills directory. + skillsDir := "" + if ast.Path != "" { + dir := filepath.Join(config.Conf.AppSource, ast.Path, "skills") + if info, e := os.Stat(dir); e == nil && info.IsDir() { + skillsDir = dir + } + } + + // 5. Convert MCP servers. + var mcpServers []sandboxTypes.MCPServer + if ast.MCP != nil { + for _, s := range ast.MCP.Servers { + mcpServers = append(mcpServers, sandboxTypes.MCPServer{ + ServerID: s.ServerID, + Resources: s.Resources, + Tools: s.Tools, + }) + } + } + + // 6. Runner.Prepare (standard context). + err = runner.Prepare(stdCtx, &sandboxTypes.PrepareRequest{ + Computer: computer, + Config: cfg, + Connector: conn, + SkillsDir: skillsDir, + MCPServers: mcpServers, + ConfigHash: ast.ConfigHash, + RunSteps: sandboxv2.RunPrepareSteps, + }) + if err != nil { + runner.Cleanup(stdCtx, computer) + sandboxv2.LifecycleAction(stdCtx, cfg, computer, manager) + closeLoadingV2(ctx, loadingMsgID, "sandbox.failed") + return nil, nil, nil, "", fmt.Errorf("runner.Prepare: %w", err) + } + + // Inject computer + workspace into context so Create/Next hooks + // can access ctx.computer and ctx.workspace. + ctx.SetComputer(computer) + + cleanup := func() { + // Defensive fallback — executeSandboxV2Stream defer handles the + // normal case; this covers paths that never reach execution. + } + + return runner, computer, cleanup, loadingMsgID, nil +} + +// executeSandboxV2Stream calls the V2 Runner.Stream and wraps it in the +// standard completion response. +func (ast *Assistant) executeSandboxV2Stream( + ctx *context.Context, + completionMessages []context.Message, + agentNode traceTypes.Node, + streamHandler message.StreamFunc, + runner sandboxTypes.Runner, + computer infraV2.Computer, + loadingMsgID string, +) (*context.CompletionResponse, error) { + _ = agentNode + + cfg := ast.SandboxV2 + manager := infraV2.M() + + // Close the "preparing" loading on first output. + if loadingMsgID != "" { + closeLoadingV2(ctx, loadingMsgID, "") + } + + // Build system prompt. + var systemPrompt string + if len(ast.Prompts) > 0 { + for _, p := range ast.Prompts { + if p.Role == "system" && p.Content != "" { + systemPrompt = p.Content + break + } + } + } + + // Resolve connector for Stream. + conn, _, _ := ast.GetConnector(ctx) + + streamReq := &sandboxTypes.StreamRequest{ + Computer: computer, + Config: cfg, + Connector: conn, + Messages: completionMessages, + SystemPrompt: systemPrompt, + ChatID: ctx.ChatID, + } + + execReq := &sandboxv2.ExecuteRequest{ + Computer: computer, + Runner: runner, + Config: cfg, + StreamReq: streamReq, + Manager: manager, + } + + return sandboxv2.ExecuteSandboxStream(ctx, execReq, streamHandler) +} + +func closeLoadingV2(ctx *context.Context, loadingMsgID, msgKey string) { + if loadingMsgID == "" || ctx == nil { + return + } + props := map[string]any{"done": true} + if msgKey != "" { + props["message"] = i18n.T(ctx.Locale, msgKey) + } else { + props["message"] = "" + } + doneMsg := &message.Message{ + MessageID: loadingMsgID, + Delta: true, + DeltaAction: message.DeltaReplace, + Type: message.TypeLoading, + Props: props, + } + ctx.Send(doneMsg) +} diff --git a/agent/context/jsapi.go b/agent/context/jsapi.go index 1081ba7d..9aa767fc 100644 --- a/agent/context/jsapi.go +++ b/agent/context/jsapi.go @@ -152,7 +152,7 @@ func (ctx *Context) NewObject(v8ctx *v8go.Context) (*v8go.Value, error) { memoryObj.Release() } - // Sandbox object - only set if sandbox executor is available + // Sandbox object - only set if sandbox executor is available (V1) if ctx.sandboxExecutor != nil { sandboxObj := ctx.createSandboxInstance(v8ctx) if sandboxObj != nil { @@ -161,6 +161,24 @@ func (ctx *Context) NewObject(v8ctx *v8go.Context) (*v8go.Value, error) { } } + // Computer object - only set if V2 computer is available + if ctx.computer != nil { + computerObj := ctx.createComputerInstance(v8ctx) + if computerObj != nil { + obj.Set("computer", computerObj) + computerObj.Release() + } + } + + // Workspace object - only set if V2 workspace is available + if ctx.workspace != nil { + wsObj := ctx.createWorkspaceInstance(v8ctx) + if wsObj != nil { + obj.Set("workspace", wsObj) + wsObj.Release() + } + } + return instance.Value, nil } diff --git a/agent/context/jsapi_computer.go b/agent/context/jsapi_computer.go new file mode 100644 index 00000000..ae3c57c1 --- /dev/null +++ b/agent/context/jsapi_computer.go @@ -0,0 +1,228 @@ +package context + +import ( + "context" + "fmt" + "strings" + + "github.com/yaoapp/gou/runtime/v8/bridge" + infraV2 "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai/workspace" + "rogchap.com/v8go" +) + +// SetComputer sets the V2 computer and its workspace for this context. +// Should be called after Runner.Prepare succeeds in initSandboxV2. +func (ctx *Context) SetComputer(computer infraV2.Computer) { + ctx.computer = computer + if computer != nil { + ctx.workspace = computer.Workplace() + } +} + +// GetComputer returns the V2 computer if available. +func (ctx *Context) GetComputer() infraV2.Computer { + return ctx.computer +} + +// GetWorkspace returns the V2 workspace FS if available. +func (ctx *Context) GetWorkspace() workspace.FS { + return ctx.workspace +} + +// HasComputer returns true if V2 computer is available. +func (ctx *Context) HasComputer() bool { + return ctx.computer != nil +} + +// createComputerInstance creates the ctx.computer JavaScript object. +func (ctx *Context) createComputerInstance(v8ctx *v8go.Context) *v8go.Value { + if ctx.computer == nil { + return nil + } + + iso := v8ctx.Isolate() + objTpl := v8go.NewObjectTemplate(iso) + + info := ctx.computer.ComputerInfo() + id := info.BoxID + if id == "" { + id = info.NodeID + } + objTpl.Set("id", id) + + objTpl.Set("Exec", ctx.computerExecMethod(iso)) + objTpl.Set("VNC", ctx.computerVNCMethod(iso)) + objTpl.Set("Proxy", ctx.computerProxyMethod(iso)) + objTpl.Set("Info", ctx.computerInfoMethod(iso)) + + instance, err := objTpl.NewInstance(v8ctx) + if err != nil { + return nil + } + return instance.Value +} + +// computerExecMethod implements ctx.computer.Exec(cmd) +// cmd can be a string or an array of strings. +// Returns: { stdout, stderr, exit_code } +func (ctx *Context) computerExecMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.computer == nil { + return bridge.JsException(v8ctx, "computer not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "Exec requires a command argument") + } + + cmd, err := parseCommandArg(v8ctx, args[0]) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + + result, err := ctx.computer.Exec(context.Background(), cmd) + if err != nil { + return bridge.JsException(v8ctx, "Exec failed: "+err.Error()) + } + + res := map[string]interface{}{ + "stdout": result.Stdout, + "stderr": result.Stderr, + "exit_code": int32(result.ExitCode), + } + jsVal, err := bridge.JsValue(v8ctx, res) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// computerVNCMethod implements ctx.computer.VNC() +// Returns the VNC URL string. +func (ctx *Context) computerVNCMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + + if ctx.computer == nil { + return bridge.JsException(v8ctx, "computer not available") + } + + url, err := ctx.computer.VNC(context.Background()) + if err != nil { + return bridge.JsException(v8ctx, "VNC failed: "+err.Error()) + } + + jsVal, err := v8go.NewValue(iso, url) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// computerProxyMethod implements ctx.computer.Proxy(port, path?) +// Returns the proxy URL string. +func (ctx *Context) computerProxyMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.computer == nil { + return bridge.JsException(v8ctx, "computer not available") + } + if len(args) < 1 || !args[0].IsNumber() { + return bridge.JsException(v8ctx, "Proxy requires a port number") + } + + port := int(args[0].Integer()) + path := "" + if len(args) >= 2 && args[1].IsString() { + path = args[1].String() + } + + url, err := ctx.computer.Proxy(context.Background(), port, path) + if err != nil { + return bridge.JsException(v8ctx, "Proxy failed: "+err.Error()) + } + + jsVal, err := v8go.NewValue(iso, url) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// computerInfoMethod implements ctx.computer.Info() +// Returns a JS object with computer identity and system information. +func (ctx *Context) computerInfoMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + + if ctx.computer == nil { + return bridge.JsException(v8ctx, "computer not available") + } + + ci := ctx.computer.ComputerInfo() + result := map[string]interface{}{ + "kind": ci.Kind, + "node_id": ci.NodeID, + "tai_id": ci.TaiID, + "status": ci.Status, + "system": map[string]interface{}{ + "os": ci.System.OS, + "arch": ci.System.Arch, + "hostname": ci.System.Hostname, + "num_cpu": int32(ci.System.NumCPU), + "shell": ci.System.Shell, + }, + } + if ci.BoxID != "" { + result["box_id"] = ci.BoxID + result["container_id"] = ci.ContainerID + result["image"] = ci.Image + result["policy"] = string(ci.Policy) + } + + jsVal, err := bridge.JsValue(v8ctx, result) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// parseCommandArg converts a JS value (string or string array) to []string. +func parseCommandArg(v8ctx *v8go.Context, val *v8go.Value) ([]string, error) { + if val.IsString() { + raw := val.String() + return strings.Fields(raw), nil + } + + if val.IsArray() { + obj, err := val.AsObject() + if err != nil { + return nil, err + } + lengthVal, err := obj.Get("length") + if err != nil { + return nil, err + } + length := int(lengthVal.Integer()) + cmd := make([]string, length) + for i := 0; i < length; i++ { + item, err := obj.GetIdx(uint32(i)) + if err != nil { + return nil, err + } + cmd[i] = item.String() + } + return cmd, nil + } + + return nil, fmt.Errorf("command must be a string or array of strings") +} diff --git a/agent/context/jsapi_workspace.go b/agent/context/jsapi_workspace.go new file mode 100644 index 00000000..d02c3168 --- /dev/null +++ b/agent/context/jsapi_workspace.go @@ -0,0 +1,291 @@ +package context + +import ( + "io/fs" + "os" + + "github.com/yaoapp/gou/runtime/v8/bridge" + "rogchap.com/v8go" +) + +// createWorkspaceInstance creates the ctx.workspace JavaScript object. +func (ctx *Context) createWorkspaceInstance(v8ctx *v8go.Context) *v8go.Value { + if ctx.workspace == nil { + return nil + } + + iso := v8ctx.Isolate() + objTpl := v8go.NewObjectTemplate(iso) + + objTpl.Set("ReadFile", ctx.wsReadFileMethod(iso)) + objTpl.Set("WriteFile", ctx.wsWriteFileMethod(iso)) + objTpl.Set("ReadDir", ctx.wsReadDirMethod(iso)) + objTpl.Set("MkdirAll", ctx.wsMkdirAllMethod(iso)) + objTpl.Set("Remove", ctx.wsRemoveMethod(iso)) + objTpl.Set("RemoveAll", ctx.wsRemoveAllMethod(iso)) + objTpl.Set("Rename", ctx.wsRenameMethod(iso)) + objTpl.Set("Copy", ctx.wsCopyMethod(iso)) + objTpl.Set("Stat", ctx.wsStatMethod(iso)) + objTpl.Set("Exists", ctx.wsExistsMethod(iso)) + + instance, err := objTpl.NewInstance(v8ctx) + if err != nil { + return nil + } + return instance.Value +} + +// wsReadFileMethod implements ctx.workspace.ReadFile(path) +func (ctx *Context) wsReadFileMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "ReadFile requires a path argument") + } + + data, err := ctx.workspace.ReadFile(args[0].String()) + if err != nil { + return bridge.JsException(v8ctx, "ReadFile failed: "+err.Error()) + } + + jsVal, err := v8go.NewValue(iso, string(data)) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// wsWriteFileMethod implements ctx.workspace.WriteFile(path, content) +func (ctx *Context) wsWriteFileMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 2 { + return bridge.JsException(v8ctx, "WriteFile requires path and content arguments") + } + + path := args[0].String() + content := args[1].String() + + if err := ctx.workspace.WriteFile(path, []byte(content), 0o644); err != nil { + return bridge.JsException(v8ctx, "WriteFile failed: "+err.Error()) + } + return v8go.Undefined(iso) + }) +} + +// wsReadDirMethod implements ctx.workspace.ReadDir(path) +func (ctx *Context) wsReadDirMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + + path := "." + if len(args) >= 1 && args[0].IsString() { + path = args[0].String() + } + + entries, err := ctx.workspace.ReadDir(path) + if err != nil { + return bridge.JsException(v8ctx, "ReadDir failed: "+err.Error()) + } + + result := make([]map[string]interface{}, 0, len(entries)) + for _, e := range entries { + fi, _ := e.Info() + item := map[string]interface{}{ + "name": e.Name(), + "is_dir": e.IsDir(), + } + if fi != nil { + item["size"] = int32(fi.Size()) + } + result = append(result, item) + } + + jsVal, err := bridge.JsValue(v8ctx, result) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// wsMkdirAllMethod implements ctx.workspace.MkdirAll(path) +func (ctx *Context) wsMkdirAllMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "MkdirAll requires a path argument") + } + + if err := ctx.workspace.MkdirAll(args[0].String(), 0o755); err != nil { + return bridge.JsException(v8ctx, "MkdirAll failed: "+err.Error()) + } + return v8go.Undefined(iso) + }) +} + +// wsRemoveMethod implements ctx.workspace.Remove(path) +func (ctx *Context) wsRemoveMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "Remove requires a path argument") + } + + if err := ctx.workspace.Remove(args[0].String()); err != nil { + return bridge.JsException(v8ctx, "Remove failed: "+err.Error()) + } + return v8go.Undefined(iso) + }) +} + +// wsRemoveAllMethod implements ctx.workspace.RemoveAll(path) +func (ctx *Context) wsRemoveAllMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "RemoveAll requires a path argument") + } + + if err := ctx.workspace.RemoveAll(args[0].String()); err != nil { + return bridge.JsException(v8ctx, "RemoveAll failed: "+err.Error()) + } + return v8go.Undefined(iso) + }) +} + +// wsRenameMethod implements ctx.workspace.Rename(oldName, newName) +func (ctx *Context) wsRenameMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 2 { + return bridge.JsException(v8ctx, "Rename requires oldName and newName arguments") + } + + if err := ctx.workspace.Rename(args[0].String(), args[1].String()); err != nil { + return bridge.JsException(v8ctx, "Rename failed: "+err.Error()) + } + return v8go.Undefined(iso) + }) +} + +// wsCopyMethod implements ctx.workspace.Copy(src, dst) +func (ctx *Context) wsCopyMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 2 { + return bridge.JsException(v8ctx, "Copy requires src and dst arguments") + } + + if _, err := ctx.workspace.Copy(args[0].String(), args[1].String()); err != nil { + return bridge.JsException(v8ctx, "Copy failed: "+err.Error()) + } + return v8go.Undefined(iso) + }) +} + +// wsStatMethod implements ctx.workspace.Stat(path) +func (ctx *Context) wsStatMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "Stat requires a path argument") + } + + fi, err := ctx.workspace.Stat(args[0].String()) + if err != nil { + return bridge.JsException(v8ctx, "Stat failed: "+err.Error()) + } + + result := map[string]interface{}{ + "name": fi.Name(), + "size": int32(fi.Size()), + "is_dir": fi.IsDir(), + "mode": int32(fi.Mode()), + "mtime": fi.ModTime().UnixMilli(), + } + jsVal, err := bridge.JsValue(v8ctx, result) + if err != nil { + return bridge.JsException(v8ctx, err.Error()) + } + return jsVal + }) +} + +// wsExistsMethod implements ctx.workspace.Exists(path) +func (ctx *Context) wsExistsMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + if ctx.workspace == nil { + return bridge.JsException(v8ctx, "workspace not available") + } + if len(args) < 1 { + return bridge.JsException(v8ctx, "Exists requires a path argument") + } + + _, err := ctx.workspace.Stat(args[0].String()) + exists := err == nil || !isNotExist(err) + + jsVal, _ := v8go.NewValue(iso, exists) + return jsVal + }) +} + +func isNotExist(err error) bool { + if os.IsNotExist(err) { + return true + } + pathErr, ok := err.(*fs.PathError) + if ok && os.IsNotExist(pathErr.Err) { + return true + } + return false +} diff --git a/agent/context/types.go b/agent/context/types.go index 5f5f39ac..10f3aa40 100644 --- a/agent/context/types.go +++ b/agent/context/types.go @@ -11,6 +11,8 @@ import ( "github.com/yaoapp/yao/agent/output" "github.com/yaoapp/yao/agent/output/message" "github.com/yaoapp/yao/openapi/oauth/types" + infraV2 "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai/workspace" traceTypes "github.com/yaoapp/yao/trace/types" ) @@ -251,6 +253,8 @@ type Context struct { trace traceTypes.Manager `json:"-"` // Trace manager, lazy initialized on first access messageMetadata *messageMetadataStore `json:"-"` // Thread-safe message metadata store for delta operations sandboxExecutor SandboxExecutor `json:"-"` // Sandbox executor for hooks (set by assistant when sandbox is configured) + computer infraV2.Computer `json:"-"` // V2 sandbox computer (set by assistant when V2 sandbox is configured) + workspace workspace.FS `json:"-"` // V2 workspace FS (derived from computer.Workplace()) // Model capabilities (set by assistant, used by output adapters) Capabilities *llm.Capabilities `json:"-"` // Model capabilities for the current connector diff --git a/agent/sandbox/v2/claude/attachments.go b/agent/sandbox/v2/claude/attachments.go new file mode 100644 index 00000000..5e38438e --- /dev/null +++ b/agent/sandbox/v2/claude/attachments.go @@ -0,0 +1,225 @@ +package claude + +import ( + "context" + "fmt" + "path/filepath" + "strings" + + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/attachment" + workspace "github.com/yaoapp/yao/tai/workspace" +) + +// prepareAttachments resolves __yao.attachment:// URLs in messages, +// copies actual files into the workspace .attachments/{chatID}/ directory via ws.Copy, +// and replaces multimodal content parts with text references. +func prepareAttachments(ctx context.Context, messages []agentContext.Message, chatID string, ws workspace.FS) ([]agentContext.Message, error) { + usedNames := make(map[string]int) + attachDir := ".attachments/" + chatID + + result := make([]agentContext.Message, len(messages)) + copy(result, messages) + + for i, msg := range result { + if msg.Role != "user" { + continue + } + + parts, ok := msg.Content.([]interface{}) + if !ok { + if typedParts, ok := msg.Content.([]agentContext.ContentPart); ok { + iparts := make([]interface{}, len(typedParts)) + for j, p := range typedParts { + m := map[string]interface{}{"type": string(p.Type)} + if p.Text != "" { + m["text"] = p.Text + } + if p.ImageURL != nil { + m["image_url"] = map[string]interface{}{ + "url": p.ImageURL.URL, + "detail": string(p.ImageURL.Detail), + } + } + if p.File != nil { + m["file"] = map[string]interface{}{ + "url": p.File.URL, + "filename": p.File.Filename, + } + } + iparts[j] = m + } + parts = iparts + } else { + continue + } + } + + if len(parts) == 0 { + continue + } + + var textParts []string + + for _, item := range parts { + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + partType, _ := m["type"].(string) + + switch partType { + case "text": + if text, ok := m["text"].(string); ok && text != "" { + textParts = append(textParts, text) + } + + case "image_url": + imgData, _ := m["image_url"].(map[string]interface{}) + if imgData == nil { + continue + } + url, _ := imgData["url"].(string) + if url == "" { + continue + } + uploaderName, fileID, isWrapper := attachment.Parse(url) + if !isWrapper { + textParts = append(textParts, fmt.Sprintf("[Image: %s]", url)) + continue + } + ref, err := resolveAttachment(ctx, uploaderName, fileID, "", attachDir, usedNames, ws) + if err != nil { + textParts = append(textParts, "[Attached image: failed to load]") + continue + } + textParts = append(textParts, ref) + + case "file": + fileData, _ := m["file"].(map[string]interface{}) + if fileData == nil { + continue + } + url, _ := fileData["url"].(string) + hintName, _ := fileData["filename"].(string) + if url == "" { + continue + } + uploaderName, fileID, isWrapper := attachment.Parse(url) + if !isWrapper { + textParts = append(textParts, fmt.Sprintf("[File: %s]", url)) + continue + } + ref, err := resolveAttachment(ctx, uploaderName, fileID, hintName, attachDir, usedNames, ws) + if err != nil { + textParts = append(textParts, "[Attached file: failed to load]") + continue + } + textParts = append(textParts, ref) + } + } + + if len(textParts) > 0 { + newMsg := result[i] + newMsg.Content = strings.Join(textParts, "\n\n") + result[i] = newMsg + } + } + + return result, nil +} + +// resolveAttachment gets the local path of an attachment and copies it into +// the workspace via ws.Copy("local:///abs/path", ".attachments/{chatID}/filename"). +func resolveAttachment( + ctx context.Context, + uploaderName, fileID, hintName, attachDir string, + usedNames map[string]int, + ws workspace.FS, +) (string, error) { + manager, exists := attachment.Managers[uploaderName] + if !exists { + return "", fmt.Errorf("attachment manager not found: %s", uploaderName) + } + + fileInfo, err := manager.Info(ctx, fileID) + if err != nil { + return "", fmt.Errorf("failed to get file info: %w", err) + } + + absPath, _, err := manager.LocalPath(ctx, fileID) + if err != nil { + return "", fmt.Errorf("failed to get local path: %w", err) + } + + filename := fileInfo.Filename + if filename == "" && hintName != "" { + filename = hintName + } + if filename == "" { + ext := extensionFromContentType(fileInfo.ContentType) + filename = fileID + ext + } + + baseName := filename + if count, exists := usedNames[baseName]; exists { + ext := filepath.Ext(filename) + name := strings.TrimSuffix(filename, ext) + filename = fmt.Sprintf("%s_%d%s", name, count+1, ext) + usedNames[baseName] = count + 1 + } else { + usedNames[baseName] = 0 + } + + dstPath := attachDir + "/" + filename + src := "local:///" + absPath + + if _, err := ws.Copy(src, dstPath); err != nil { + return "", fmt.Errorf("failed to copy attachment to workspace: %w", err) + } + + sizeStr := formatFileSize(fileInfo.Bytes) + return fmt.Sprintf("[Attached file: %s (%s, %s)]", dstPath, fileInfo.ContentType, sizeStr), nil +} + +func extensionFromContentType(contentType string) string { + switch contentType { + case "image/png": + return ".png" + case "image/jpeg": + return ".jpg" + case "image/gif": + return ".gif" + case "image/webp": + return ".webp" + case "image/svg+xml": + return ".svg" + case "application/pdf": + return ".pdf" + case "text/plain": + return ".txt" + case "text/html": + return ".html" + case "text/css": + return ".css" + case "text/javascript", "application/javascript": + return ".js" + case "application/json": + return ".json" + case "application/zip": + return ".zip" + default: + return "" + } +} + +func formatFileSize(bytes int) string { + switch { + case bytes >= 1024*1024: + return fmt.Sprintf("%.1fMB", float64(bytes)/(1024*1024)) + case bytes >= 1024: + return fmt.Sprintf("%.1fKB", float64(bytes)/1024) + default: + return fmt.Sprintf("%dB", bytes) + } +} diff --git a/agent/sandbox/v2/claude/parse.go b/agent/sandbox/v2/claude/parse.go new file mode 100644 index 00000000..a960a224 --- /dev/null +++ b/agent/sandbox/v2/claude/parse.go @@ -0,0 +1,239 @@ +package claude + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "log" + "strings" + "time" + + goujson "github.com/yaoapp/gou/json" + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/output/message" +) + +// parseStreamJSON reads stream-json lines from Claude CLI stdout and +// pushes them through handler as standard StreamChunkType events. +func parseStreamJSON(_ context.Context, stdout io.ReadCloser, handler message.StreamFunc) error { + scanner := bufio.NewScanner(stdout) + buf := make([]byte, 0, 64*1024) + scanner.Buffer(buf, 1024*1024) + + messageStarted := false + + type toolState struct { + name string + inputJSON strings.Builder + } + var currentTool *toolState + + for scanner.Scan() { + line := scanner.Text() + if line == "" { + continue + } + + var msg map[string]any + if err := json.Unmarshal([]byte(line), &msg); err != nil { + continue + } + + msgType, _ := msg["type"].(string) + stopped := false + + switch msgType { + case "system": + if handler != nil { + data, _ := json.Marshal(msg) + if handler(message.ChunkMetadata, data) != 0 { + stopped = true + } + } + + case "stream_event": + event, _ := msg["event"].(map[string]any) + if event == nil { + continue + } + eventType, _ := event["type"].(string) + + switch eventType { + case "content_block_start": + if cb, ok := event["content_block"].(map[string]any); ok { + blockType, _ := cb["type"].(string) + if blockType == "tool_use" { + toolName, _ := cb["name"].(string) + currentTool = &toolState{name: toolName} + if handler != nil { + data, _ := json.Marshal(map[string]any{"tool": toolName}) + if handler(message.ChunkToolCall, data) != 0 { + stopped = true + } + } + } + } + + case "content_block_delta": + if delta, ok := event["delta"].(map[string]any); ok { + deltaType, _ := delta["type"].(string) + switch deltaType { + case "text_delta": + if text, ok := delta["text"].(string); ok && text != "" { + if handler != nil { + if !messageStarted { + startData := message.EventMessageStartData{ + MessageID: fmt.Sprintf("sandbox-%d", time.Now().UnixNano()), + Type: "text", + Timestamp: time.Now().UnixMilli(), + } + sd, _ := json.Marshal(startData) + if handler(message.ChunkMessageStart, sd) != 0 { + stopped = true + break + } + messageStarted = true + } + if handler(message.ChunkText, []byte(text)) != 0 { + stopped = true + } + } + } + case "input_json_delta": + if currentTool != nil { + if partial, ok := delta["partial_json"].(string); ok { + currentTool.inputJSON.WriteString(partial) + } + } + } + } + + case "content_block_stop": + currentTool = nil + } + + case "assistant": + if msgData, ok := msg["message"].(map[string]any); ok { + stopReason, _ := msgData["stop_reason"].(string) + if stopReason != "" { + if contentArr, ok := msgData["content"].([]any); ok { + for _, item := range contentArr { + ci, ok := item.(map[string]any) + if !ok { + continue + } + itemType, _ := ci["type"].(string) + if itemType == "text" { + if text, ok := ci["text"].(string); ok && text != "" && handler != nil && !messageStarted { + startData := message.EventMessageStartData{ + MessageID: fmt.Sprintf("sandbox-%d", time.Now().UnixNano()), + Type: "text", + Timestamp: time.Now().UnixMilli(), + } + sd, _ := json.Marshal(startData) + if handler(message.ChunkMessageStart, sd) != 0 { + stopped = true + break + } + if handler(message.ChunkText, []byte(text)) != 0 { + stopped = true + break + } + messageStarted = true + } + } + } + } + } + } + + case "result": + isError, _ := msg["is_error"].(bool) + if isError { + if result, ok := msg["result"].(string); ok { + if handler != nil { + handler(message.ChunkError, []byte(result)) + } + return fmt.Errorf("Claude CLI error: %s", result) + } + } + if handler != nil && messageStarted { + handler(message.ChunkMessageEnd, nil) + } + + case "error": + var errMsg string + switch e := msg["error"].(type) { + case string: + errMsg = e + case map[string]any: + errMsg, _ = e["message"].(string) + } + if errMsg != "" { + if handler != nil { + handler(message.ChunkError, []byte(errMsg)) + } + return fmt.Errorf("Claude CLI error: %s", errMsg) + } + } + + if stopped { + break + } + } + + return scanner.Err() +} + +// buildFirstRequestJSONL builds JSONL with all messages for the first request. +func buildFirstRequestJSONL(messages []agentContext.Message) string { + var lines []string + for _, msg := range messages { + if msg.Role == "system" { + continue + } + content := msg.Content + if content == nil { + content = "" + } + streamMsg := map[string]any{ + "type": string(msg.Role), + "message": map[string]any{ + "role": string(msg.Role), + "content": content, + }, + } + data, _ := json.Marshal(streamMsg) + lines = append(lines, string(data)) + } + return strings.Join(lines, "\n") +} + +// buildLastUserMessageJSONL builds JSONL with only the last user message. +func buildLastUserMessageJSONL(messages []agentContext.Message) string { + for i := len(messages) - 1; i >= 0; i-- { + if messages[i].Role == "user" { + content := messages[i].Content + if content == nil { + content = "" + } + msg := map[string]any{ + "type": "user", + "message": map[string]any{ + "role": "user", + "content": content, + }, + } + data, _ := json.Marshal(msg) + return string(data) + } + } + return "" +} + +// Suppress unused import warnings — goujson.Parse is used for tool description +// parsing in V1 and will be used for detailed tool descriptions in future. +var _ = goujson.Parse +var _ = log.Printf diff --git a/agent/sandbox/v2/claude/runner.go b/agent/sandbox/v2/claude/runner.go new file mode 100644 index 00000000..fac6c522 --- /dev/null +++ b/agent/sandbox/v2/claude/runner.go @@ -0,0 +1,405 @@ +package claude + +import ( + "context" + "encoding/json" + "fmt" + "path" + "strings" + "time" + + "github.com/yaoapp/gou/connector" + "github.com/yaoapp/yao/agent/output/message" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +const ( + defaultWorkDir = "/workspace" + defaultUser = "sandbox" + defaultUserHome = "/home/sandbox" + defaultProxyPort = 3456 +) + +// ClaudeRunner implements the Runner interface for Claude CLI (mode=cli). +type ClaudeRunner struct { + mode string + hasMCP bool + mcpToolPattern string // e.g. "mcp__yao__*,mcp__github__*" + servicePort int + servicePath string + serviceProtocol string +} + +// New creates a new ClaudeRunner. +func New() *ClaudeRunner { + return &ClaudeRunner{mode: "cli"} +} + +func (r *ClaudeRunner) Name() string { return "claude" } + +// Prepare executes user-defined and runner-specific prepare steps. +func (r *ClaudeRunner) Prepare(ctx context.Context, req *types.PrepareRequest) error { + r.mode = req.Config.Runner.Mode + if r.mode == "" { + r.mode = "cli" + } + + workDir := resolveWorkDir(req.Config) + + // Merge user-defined steps with runner-specific steps. + steps := append([]types.PrepareStep{}, req.Config.Prepare...) + + // Runner-specific: ensure .claude directory in workDir. + if req.SkillsDir != "" { + steps = append(steps, types.PrepareStep{ + Action: "exec", + Cmd: fmt.Sprintf("mkdir -p %s/.claude", workDir), + Once: true, + }) + } + + // Runner-specific: write MCP config. + if len(req.MCPServers) > 0 { + r.hasMCP = true + r.mcpToolPattern = buildMCPAllowedTools(req.MCPServers) + mcpJSON := buildMCPConfig(req.MCPServers) + steps = append(steps, types.PrepareStep{ + Action: "file", + Path: path.Join(workDir, ".mcp.json"), + Content: mcpJSON, + }) + } + + // Execute all steps via the injected callback. + if req.RunSteps != nil && len(steps) > 0 { + if err := req.RunSteps(ctx, steps, req.Computer, req.Config.ID, req.ConfigHash); err != nil { + return fmt.Errorf("claude prepare steps: %w", err) + } + } + + return nil +} + +// Stream executes the Claude CLI and streams output to handler. +func (r *ClaudeRunner) Stream(ctx context.Context, req *types.StreamRequest, handler message.StreamFunc) error { + computer := req.Computer + if computer == nil { + return fmt.Errorf("computer is nil") + } + + workDir := resolveWorkDir(req.Config) + + // Prepare attachments: resolve __yao.attachment:// URLs, copy files to workspace. + if req.ChatID != "" { + ws := computer.Workplace() + if ws != nil { + processed, err := prepareAttachments(ctx, req.Messages, req.ChatID, ws) + if err != nil { + return fmt.Errorf("prepareAttachments: %w", err) + } + req.Messages = processed + } + } + + // Detect continuation (existing .claude/projects/ directory). + isContinuation := hasExistingSession(ctx, computer, workDir) + + // Build CLI command and env. + cmd, env := r.buildCLICommand(req, isContinuation) + + // Create stream. + execStream, err := computer.Stream(ctx, cmd, infra.WithWorkDir(workDir), infra.WithEnv(env)) + if err != nil { + return fmt.Errorf("computer.Stream: %w", err) + } + + // Monitor for context cancellation — kill the process. + done := make(chan struct{}) + defer func() { + close(done) + }() + + go func() { + select { + case <-ctx.Done(): + killCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + computer.Exec(killCtx, []string{"pkill", "-f", "claude"}) + execStream.Cancel() + case <-done: + } + }() + + // Parse streaming output. + parseErr := parseStreamJSON(ctx, execStream.Stdout, handler) + + // Wait for process exit. + exitCode, waitErr := execStream.Wait() + if parseErr != nil { + return parseErr + } + if waitErr != nil { + return waitErr + } + if exitCode != 0 { + return fmt.Errorf("claude CLI exited with code %d", exitCode) + } + return nil +} + +// Cleanup kills any remaining claude processes. +// mode=cli: kill all claude CLI processes. +func (r *ClaudeRunner) Cleanup(ctx context.Context, computer infra.Computer) error { + if computer == nil { + return nil + } + + if r.mode != "service" { + computer.Exec(ctx, []string{"sh", "-c", "pkill -f 'claude' || true"}) + } + + return nil +} + +// hasExistingSession checks if a Claude CLI session exists in the workspace. +func hasExistingSession(ctx context.Context, computer infra.Computer, workDir string) bool { + sessionDir := path.Join(workDir, ".claude/projects") + result, err := computer.Exec(ctx, []string{"ls", sessionDir}) + if err != nil || result.ExitCode != 0 { + return false + } + return strings.TrimSpace(result.Stdout) != "" +} + +// buildCLICommand constructs the Claude CLI command and environment variables. +func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, isContinuation bool) ([]string, map[string]string) { + workDir := resolveWorkDir(req.Config) + userHome := resolveUserHome(req.Config) + + env := make(map[string]string) + env["HOME"] = workDir + + // User-specific paths (only set when running as non-root user inside container). + if userHome != "" { + env["XAUTHORITY"] = path.Join(userHome, ".Xauthority") + } + + // Connector environment. + if req.Connector != nil { + setting := req.Connector.Setting() + host, _ := setting["host"].(string) + key, _ := setting["key"].(string) + model, _ := setting["model"].(string) + + if req.Connector.Is(connector.ANTHROPIC) { + env["ANTHROPIC_BASE_URL"] = host + env["ANTHROPIC_API_KEY"] = key + } else { + env["ANTHROPIC_BASE_URL"] = fmt.Sprintf("http://127.0.0.1:%d", defaultProxyPort) + env["ANTHROPIC_API_KEY"] = "dummy" + } + + if model != "" { + env["ANTHROPIC_MODEL"] = model + env["ANTHROPIC_DEFAULT_OPUS_MODEL"] = model + env["ANTHROPIC_DEFAULT_SONNET_MODEL"] = model + env["ANTHROPIC_DEFAULT_HAIKU_MODEL"] = model + env["CLAUDE_CODE_SUBAGENT_MODEL"] = model + } + } + + // Secrets from config. + if req.Config != nil && len(req.Config.Secrets) > 0 { + for k, v := range req.Config.Secrets { + env[k] = v + } + } + + // Build system prompt. + var systemPrompt string + envPrompt := buildSandboxEnvPrompt(workDir) + if !isContinuation && req.SystemPrompt != "" { + systemPrompt = req.SystemPrompt + "\n\n" + envPrompt + } else if !isContinuation { + systemPrompt = envPrompt + } + + // Build input JSONL. + var inputJSONL string + if isContinuation { + inputJSONL = buildLastUserMessageJSONL(req.Messages) + } else { + inputJSONL = buildFirstRequestJSONL(req.Messages) + } + + // CLI args. + var args []string + args = append(args, "--dangerously-skip-permissions") + args = append(args, "--permission-mode", "bypassPermissions") + args = append(args, "--input-format", "stream-json") + args = append(args, "--output-format", "stream-json") + args = append(args, "--include-partial-messages") + args = append(args, "--verbose") + + if isContinuation { + args = append(args, "--continue") + } + + // Runner options pass-through. + if req.Config != nil && req.Config.Runner.Options != nil { + for key, val := range req.Config.Runner.Options { + if flag, ok := claudeArgWhitelist[key]; ok { + args = append(args, flag, fmt.Sprintf("%v", val)) + } + } + } + + // MCP config (set by Prepare if MCPServers were present). + if r.hasMCP { + args = append(args, "--mcp-config", path.Join(workDir, ".mcp.json")) + if r.mcpToolPattern != "" { + args = append(args, "--allowedTools", r.mcpToolPattern) + } + } + + // Build bash command with heredoc. + var bash strings.Builder + if userHome != "" { + bash.WriteString(fmt.Sprintf("touch %s/.Xauthority 2>/dev/null; ", userHome)) + } + bash.WriteString("touch \"$HOME/.Xauthority\" 2>/dev/null\n") + + if systemPrompt != "" { + promptFile := path.Join(workDir, ".yao/.system-prompt.txt") + bash.WriteString(fmt.Sprintf("mkdir -p %s/.yao\n", workDir)) + bash.WriteString(fmt.Sprintf("cat << 'PROMPTEOF' > %s\n", promptFile)) + bash.WriteString(systemPrompt) + bash.WriteString("\nPROMPTEOF\n") + args = append(args, "--append-system-prompt-file", promptFile) + } + + bash.WriteString("cat << 'INPUTEOF' | claude -p") + for _, arg := range args { + bash.WriteString(fmt.Sprintf(" %q", arg)) + } + bash.WriteString(" 2>&1\n") + bash.WriteString(inputJSONL) + bash.WriteString("\nINPUTEOF") + + return []string{"bash", "-c", bash.String()}, env +} + +// buildMCPConfig creates the .mcp.json for Claude CLI based on declared servers. +// Each server delegates to "tai mcp" which implements the standard MCP protocol +// over stdio and bridges to Yao gRPC with authentication. +// Connection is configured via env vars (YAO_GRPC_ADDR, YAO_TOKEN, etc.) +// injected by the sandbox infrastructure at container start. +func buildMCPConfig(servers []types.MCPServer) []byte { + mcpServers := make(map[string]any, len(servers)) + for _, s := range servers { + name := s.ServerID + if name == "" { + continue + } + mcpServers[name] = map[string]any{ + "command": "tai", + "args": []string{"mcp"}, + } + } + if len(mcpServers) == 0 { + mcpServers["yao"] = map[string]any{ + "command": "tai", + "args": []string{"mcp"}, + } + } + config := map[string]any{"mcpServers": mcpServers} + data, _ := json.Marshal(config) + return data +} + +// buildMCPAllowedTools generates the --allowedTools pattern from server IDs. +func buildMCPAllowedTools(servers []types.MCPServer) string { + patterns := make([]string, 0, len(servers)) + for _, s := range servers { + if s.ServerID != "" { + patterns = append(patterns, fmt.Sprintf("mcp__%s__*", s.ServerID)) + } + } + if len(patterns) == 0 { + return "mcp__yao__*" + } + return strings.Join(patterns, ",") +} + +// buildSandboxEnvPrompt generates the sandbox environment prompt with the actual working directory. +func buildSandboxEnvPrompt(workDir string) string { + return fmt.Sprintf(`## Sandbox Environment + +You are running in a sandboxed environment with the following setup: + +- **Working Directory**: %[1]s +- **Project Structure**: If this is a new project, create a dedicated project folder (e.g., %[1]s/my-project/) and work inside it +- **File Access**: You have full read/write access to %[1]s +- **Output Files**: Save all output files to the working directory + +When creating new projects: +1. Create a project directory with a descriptive name +2. Initialize the project structure inside that directory +3. Keep all related files organized within the project folder + +## IMPORTANT: Restricted Tools + +The following tools are NOT available in this environment and you must NOT use them: +- EnterPlanMode, ExitPlanMode (use regular text to explain plans instead) +- Task, TaskOutput, TaskStop (complete tasks directly without delegation) +- AskUserQuestion (make reasonable assumptions instead of asking) +- Skill, ToolSearch (not supported) + +Focus on using the core tools: Bash, Read, Write, Edit, Glob, Grep, WebSearch, WebFetch. + +## User Attachments + +User-uploaded files (images, documents, code files, etc.) are placed in %[1]s/.attachments/{chatID}/ +Each chat session has its own subdirectory to avoid conflicts. +When the user references an attached file, read it from this directory using the Read or Bash tool. +For image files, you can view them directly as Claude supports vision on local files. + +## GitHub CLI (gh) Usage + +When working with GitHub and a token is provided: +1. First authenticate gh CLI using the token: echo "TOKEN" | gh auth login --with-token +2. Then use gh commands normally (gh repo create, gh pr create, etc.) +3. Do NOT use curl to call GitHub API directly - always prefer gh CLI +`, workDir) +} + +// resolveWorkDir returns the configured working directory, falling back to default. +func resolveWorkDir(cfg *types.SandboxConfig) string { + if cfg != nil && cfg.Computer.WorkDir != "" { + return cfg.Computer.WorkDir + } + return defaultWorkDir +} + +// resolveUserHome returns the home directory for the container user. +// Returns empty string if no user is configured (root or unspecified). +func resolveUserHome(cfg *types.SandboxConfig) string { + if cfg == nil { + return defaultUserHome + } + user := cfg.Computer.User + if user == "" { + user = defaultUser + } + if user == "root" { + return "/root" + } + return fmt.Sprintf("/home/%s", user) +} + +var claudeArgWhitelist = map[string]string{ + "max_turns": "--max-turns", + "disallowed_tools": "--disallowed-tools", + "allowed_tools": "--allowedTools", +} diff --git a/agent/sandbox/v2/claude/runner_test.go b/agent/sandbox/v2/claude/runner_test.go new file mode 100644 index 00000000..97afa979 --- /dev/null +++ b/agent/sandbox/v2/claude/runner_test.go @@ -0,0 +1,279 @@ +package claude_test + +import ( + "bytes" + "context" + "fmt" + "mime/multipart" + "os" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/caller" + agentcontext "github.com/yaoapp/yao/agent/context" + sandboxtestutils "github.com/yaoapp/yao/agent/sandbox/v2/testutils" + "github.com/yaoapp/yao/attachment" + oauthtypes "github.com/yaoapp/yao/openapi/oauth/types" +) + +type e2eCase struct { + ID string + Prompt string + Timeout time.Duration +} + +var cases = []e2eCase{ + { + ID: "tests.sandbox-v2.oneshot-cli", + Prompt: "Reply exactly with: hello sandbox v2", + Timeout: 3 * time.Minute, + }, +} + +func TestSandboxV2_Claude_E2E(t *testing.T) { + sandboxtestutils.Prepare(t) + defer sandboxtestutils.Clean(t) + + require.NotNil(t, caller.AgentGetterFunc, "AgentGetterFunc should be registered after Prepare") + + for _, tc := range cases { + tc := tc + t.Run(tc.ID, func(t *testing.T) { + agent, err := caller.AgentGetterFunc(tc.ID) + require.NoError(t, err, "should load assistant %s", tc.ID) + + timeout := tc.Timeout + if timeout == 0 { + timeout = 3 * time.Minute + } + + chatID := fmt.Sprintf("e2e-%s-%d", tc.ID, time.Now().UnixMilli()) + ctx := agentcontext.New( + context.Background(), + &oauthtypes.AuthorizedInfo{ + TeamID: "test-team-e2e", + UserID: "test-user-e2e", + }, + chatID, + ) + + messages := []agentcontext.Message{ + {Role: "user", Content: tc.Prompt}, + } + + done := make(chan struct{}) + var resp *agentcontext.Response + var streamErr error + + go func() { + defer close(done) + resp, streamErr = agent.Stream(ctx, messages) + }() + + select { + case <-done: + case <-time.After(timeout): + t.Fatalf("timeout after %v", timeout) + } + + require.NoError(t, streamErr, "Stream should not return error") + require.NotNil(t, resp, "response should not be nil") + + // ── 1. CompletionResponse should behave like the LLM path ── + require.NotNil(t, resp.Completion, "completion should not be nil") + assert.Equal(t, "assistant", resp.Completion.Role, "role should be assistant") + assert.Equal(t, agentcontext.FinishReasonStop, resp.Completion.FinishReason, "finish_reason should be stop") + assert.NotNil(t, resp.Completion.Content, "Content should be populated (same as LLM path)") + + contentStr, ok := resp.Completion.Content.(string) + require.True(t, ok, "Content should be a string, got %T", resp.Completion.Content) + t.Logf("CompletionResponse.Content (%d chars): %s", len(contentStr), contentStr) + assert.Contains(t, contentStr, "hello sandbox v2", "Content should contain expected text") + + // ── 2. Buffer: frame sequence handled correctly ── + require.NotNil(t, ctx.Buffer, "ctx.Buffer should not be nil") + + msgs := ctx.Buffer.GetMessages() + t.Logf("buffer message count: %d", len(msgs)) + for _, m := range msgs { + t.Logf(" seq=%d role=%s type=%s streaming=%v props_keys=%v", + m.Sequence, m.Role, m.Type, m.IsStreaming, mapKeys(m.Props)) + } + + var userInputCount, assistantTextCount, loadingCount int + var bufferTextContent string + for _, m := range msgs { + switch { + case m.Role == "user" && m.Type == "user_input": + userInputCount++ + case m.Role == "assistant" && m.Type == "loading": + loadingCount++ + case m.Role == "assistant" && m.Type == "text": + assistantTextCount++ + assert.False(t, m.IsStreaming, "text message should not be streaming (handleMessageEnd should have finalized it)") + require.NotNil(t, m.Props, "text message props should not be nil") + if c, ok := m.Props["content"].(string); ok { + bufferTextContent += c + } + } + } + + assert.Equal(t, 1, userInputCount, "should have exactly 1 user_input message") + assert.GreaterOrEqual(t, loadingCount, 1, "should have at least 1 loading message") + assert.Equal(t, 1, assistantTextCount, "should have exactly 1 assistant text message (from handleMessageEnd)") + assert.Contains(t, bufferTextContent, "hello sandbox v2", "buffer text should contain expected content") + + // ── 3. Buffer content matches CompletionResponse.Content ── + assert.Equal(t, contentStr, bufferTextContent, + "CompletionResponse.Content and Buffer text should match") + }) + } +} + +func TestSandboxV2_Claude_Attachments(t *testing.T) { + sandboxtestutils.Prepare(t) + defer sandboxtestutils.Clean(t) + + require.NotNil(t, caller.AgentGetterFunc, "AgentGetterFunc should be registered after Prepare") + + agent, err := caller.AgentGetterFunc("tests.sandbox-v2.oneshot-cli") + require.NoError(t, err) + + // ── 1. Locate testdata via runtime.Caller ── + _, thisFile, _, ok := runtime.Caller(0) + require.True(t, ok) + testdataDir := filepath.Join(filepath.Dir(thisFile), "testdata") + + // ── 2. Create attachment manager and upload test files ── + const uploaderName = "__yao.attachment" + manager, err := attachment.New(attachment.ManagerOption{ + Driver: "local", + MaxSize: "50M", + AllowedTypes: []string{"image/*", "text/*", "application/*", "video/*", ".ts", ".js", ".tsx", ".jsx"}, + Options: map[string]interface{}{"path": filepath.Join(os.TempDir(), "test_sandbox_v2_attach")}, + }) + require.NoError(t, err) + manager.Name = uploaderName + attachment.Managers[uploaderName] = manager + t.Cleanup(func() { delete(attachment.Managers, uploaderName) }) + + imgFile := uploadTestFile(t, manager, testdataDir, "test-image.png", "image/png") + codeFile := uploadTestFile(t, manager, testdataDir, "code.ts", "text/plain") + + imgWrapper := fmt.Sprintf("%s://%s", uploaderName, imgFile.ID) + codeWrapper := fmt.Sprintf("%s://%s", uploaderName, codeFile.ID) + t.Logf("image wrapper: %s", imgWrapper) + t.Logf("code wrapper: %s", codeWrapper) + + // ── 3. Build multimodal messages (same as CUI InputArea) ── + chatID := fmt.Sprintf("e2e-attach-%d", time.Now().UnixMilli()) + ctx := agentcontext.New( + context.Background(), + &oauthtypes.AuthorizedInfo{TeamID: "test-team-e2e", UserID: "test-user-e2e"}, + chatID, + ) + + messages := []agentcontext.Message{ + { + Role: "user", + Content: []interface{}{ + map[string]interface{}{"type": "text", "text": "Describe the attached image and summarize the attached code file. Reply in English."}, + map[string]interface{}{ + "type": "image_url", + "image_url": map[string]interface{}{"url": imgWrapper, "detail": "auto"}, + }, + map[string]interface{}{ + "type": "file", + "file": map[string]interface{}{"url": codeWrapper, "filename": "code.ts"}, + }, + }, + }, + } + + // ── 4. Run E2E stream ── + done := make(chan struct{}) + var resp *agentcontext.Response + var streamErr error + + go func() { + defer close(done) + resp, streamErr = agent.Stream(ctx, messages) + }() + + select { + case <-done: + case <-time.After(5 * time.Minute): + t.Fatalf("timeout after 5m") + } + + require.NoError(t, streamErr, "Stream should not return error") + require.NotNil(t, resp) + require.NotNil(t, resp.Completion) + + contentStr, ok := resp.Completion.Content.(string) + require.True(t, ok, "Content should be a string, got %T", resp.Completion.Content) + t.Logf("Response (%d chars): %s", len(contentStr), contentStr) + + lower := strings.ToLower(contentStr) + + // ── 5. Verify Claude actually read the image ── + imageKeywords := []string{"hello", "utf", "chinese", "text", "emoji"} + imgHit := false + for _, kw := range imageKeywords { + if strings.Contains(lower, kw) { + imgHit = true + break + } + } + assert.True(t, imgHit, "response should mention image content (tried: %v)", imageKeywords) + + // ── 6. Verify Claude actually read the code ── + codeKeywords := []string{"excel", "typescript", "class", "volcengine"} + codeHit := false + for _, kw := range codeKeywords { + if strings.Contains(lower, kw) { + codeHit = true + break + } + } + assert.True(t, codeHit, "response should mention code content (tried: %v)", codeKeywords) +} + +func uploadTestFile(t *testing.T, manager *attachment.Manager, testdataDir, filename, contentType string) *attachment.File { + t.Helper() + path := filepath.Join(testdataDir, filename) + data, err := os.ReadFile(path) + require.NoError(t, err, "read testdata/%s", filename) + + fh := &attachment.FileHeader{ + FileHeader: &multipart.FileHeader{ + Filename: filename, + Size: int64(len(data)), + Header: make(map[string][]string), + }, + } + fh.Header.Set("Content-Type", contentType) + + file, err := manager.Upload(context.Background(), fh, bytes.NewReader(data), attachment.UploadOption{ + Groups: []string{"e2e-sandbox-v2"}, + }) + require.NoError(t, err, "upload testdata/%s", filename) + t.Logf("uploaded %s => ID=%s, Path=%s", filename, file.ID, file.Path) + return file +} + +func mapKeys(m map[string]interface{}) []string { + if m == nil { + return nil + } + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + return keys +} diff --git a/agent/sandbox/v2/claude/testdata/code.ts b/agent/sandbox/v2/claude/testdata/code.ts new file mode 100644 index 00000000..458e8de1 --- /dev/null +++ b/agent/sandbox/v2/claude/testdata/code.ts @@ -0,0 +1,904 @@ +import { Process } from "@yao/runtime"; + +/** + * Excel class for manipulating Excel files via Yao's Excel Module + */ +export class Excel { + private handle: string | null = null; + + /** + * Creates a new Excel instance + * @param file Path to the Excel file + */ + constructor(private file: string, writable: boolean = false) { + this.file = file; + this.Open(writable); + } + + /** + * Read each sheet top n rows + * @param file Path to the Excel file + * @param n number of rows to read + * @returns Object with sheet names as keys and arrays of row values as values + */ + static Heads( + file: string, + n: number = 5, + filters?: string[] + ): Record { + const excel = new Excel(file); + const heads = excel.Heads(n, filters); + excel.Close(); + return heads; + } + + /** + * Read each sheet top n rows + * @param n number of rows to read + * @returns Object with sheet names as keys and arrays of row values as values + * @throws Error if file not opened + */ + Heads(n: number = 5, filters: string[] = []): Record { + if (!this.handle) throw new Error("Excel file not opened"); + + const sheets = this.Sheets(); + const result: Record = {}; + + for (const sheet of sheets) { + if (filters.length > 0 && !filters.includes(sheet)) { + continue; + } + + // Open row iterator for the sheet + const iterator = this.each.OpenRow(sheet); + const rows: any[][] = []; + + // Read n rows + let row; + let count = 0; + while ( + count < n && + (row = Process(`excel.each.NextRow`, iterator)) !== null + ) { + // Add column headers (A, B, C, ...) for the first row + if (count === 0) { + const headerRow = []; + for (let i = 0; i < row.length; i++) { + headerRow.push(this.convert.ColumnNumberToName(i + 1)); + } + rows.push(headerRow); + } + + // Trim Each cell's value + row = row.map((cell) => cell?.trim?.()); + rows.push(row); + count++; + } + + // Close the row iterator + this.each.CloseRow(iterator); + + // Find the max length of each row, and pad the column headers(A, B, C, ...) to the same length + const maxLength = Math.max(...rows.map((row) => row.length)); + const start = rows[0].length; + const neededLength = maxLength - rows[0].length; + for (let i = 0; i < neededLength; i++) { + rows[0].push(this.convert.ColumnNumberToName(start + i + 1)); + } + + // Add the sheet's rows to the result + result[sheet] = rows; + } + + return result; + } + + /** + * Check if a sheet exists in the Excel file + * @param file Path to the Excel file + * @param sheet Sheet name to check + * @returns boolean - true if sheet exists, false otherwise + */ + static Exists(file: string, sheet: string) { + const excel = new Excel(file); + const exists = excel.sheet.Exists(sheet); + excel.Close(); + return exists; + } + + /** + * Opens an Excel file for reading or writing + * @param writable Whether to open in writable mode (true) or read-only mode (false) + * @returns Handle ID used for subsequent operations + */ + Open(writable: boolean = false) { + this.handle = Process(`excel.Open`, this.file, writable); + return this.handle; + } + + /** + * Closes the Excel file and releases resources + * IMPORTANT: Always call this method when done to prevent memory leaks + */ + Close() { + if (this.handle) { + Process(`excel.Close`, this.handle); + this.handle = null; + } + } + + /** + * Saves changes to the Excel file + * @throws Error if file not opened + */ + Save() { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.Save`, this.handle); + } + + /** + * Gets all sheet names in the workbook + * @returns Array of sheet names + * @throws Error if file not opened + */ + Sheets() { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.Sheets`, this.handle); + } + + // Sheet operations + sheet = { + /** + * Creates a new sheet in the workbook + * @param name Name for the new sheet + * @returns number Index of the new sheet + * @throws Error if file not opened + */ + Create: (name: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.create`, this.handle, name); + }, + + /** + * Lists all sheets in the workbook + * @returns string[] Array of sheet names + * @throws Error if file not opened + */ + List: () => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.list`, this.handle); + }, + + /** + * Checks if a sheet exists in the workbook + * @param name Sheet name to check + * @returns boolean - true if sheet exists, false otherwise + * @throws Error if file not opened + */ + Exists: (name: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.exists`, this.handle, name); + }, + + /** + * Reads all data from a sheet + * @param name Sheet name + * @returns any[][] Two-dimensional array of cell values + * @throws Error if file not opened + */ + Read: (name: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.read`, this.handle, name); + }, + + /** + * Reads all data from a sheet with pagination support + * @param name Sheet name + * @param from Starting row index (0-based) + * @param chunk_size Number of rows to read + * @returns any[][] Two-dimensional array of cell values + * @throws Error if file not opened + */ + Rows: (name: string, from: number, chunk_size: number) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.rows`, this.handle, name, from, chunk_size); + }, + + /** + * Updates data in a sheet. Creates the sheet if it doesn't exist. + * @param name Sheet name + * @param data Two-dimensional array of values to write + * @throws Error if file not opened + */ + Update: (name: string, data: any[][]) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.update`, this.handle, name, data); + }, + + /** + * Copies a sheet with all its content and formatting + * @param source Source sheet name + * @param target Target sheet name (must not exist) + * @throws Error if file not opened + */ + Copy: (source: string, target: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.copy`, this.handle, source, target); + }, + + /** + * Deletes a sheet from the workbook + * @param name Sheet name to delete + * @throws Error if file not opened + */ + Delete: (name: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.delete`, this.handle, name); + }, + + /** + * Gets the dimensions (number of rows and columns) of a sheet + * @param name Sheet name + * @returns {rows: number, cols: number} - Object containing row and column counts + * @throws Error if file not opened + */ + Dimension: (name: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.sheet.dimension`, this.handle, name); + }, + }; + + // Reading operations + read = { + /** + * Reads a cell's value + * @param sheet Sheet name + * @param cell Cell reference (e.g. "A1") + * @returns Cell value + * @throws Error if file not opened + */ + Cell: (sheet: string, cell: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.read.Cell`, this.handle, sheet, cell); + }, + + /** + * Reads all rows in a sheet + * @param sheet Sheet name + * @returns Two-dimensional array of cell values + * @throws Error if file not opened + */ + Row: (sheet: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.read.Row`, this.handle, sheet); + }, + + /** + * Reads all columns in a sheet + * @param sheet Sheet name + * @returns Two-dimensional array of cell values + * @throws Error if file not opened + */ + Column: (sheet: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.read.Column`, this.handle, sheet); + }, + }; + + // Writing operations + write = { + /** + * Writes a value to a cell + * @param sheet Sheet name + * @param cell Cell reference (e.g. "A1") + * @param value Value to write (string, number, boolean, etc.) + * @throws Error if file not opened + */ + Cell: (sheet: string, cell: string, value: any) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.write.Cell`, this.handle, sheet, cell, value); + }, + + /** + * Writes values to a row starting at the specified cell + * @param sheet Sheet name + * @param startCell Starting cell reference (e.g. "A1") + * @param values Array of values to write + * @throws Error if file not opened + */ + Row: (sheet: string, startCell: string, values: any[]) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.write.Row`, this.handle, sheet, startCell, values); + }, + + /** + * Writes values to a column starting at the specified cell + * @param sheet Sheet name + * @param startCell Starting cell reference (e.g. "A1") + * @param values Array of values to write + * @throws Error if file not opened + */ + Column: (sheet: string, startCell: string, values: any[]) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process( + `excel.write.Column`, + this.handle, + sheet, + startCell, + values + ); + }, + + /** + * Writes a two-dimensional array of values starting at the specified cell + * @param sheet Sheet name + * @param startCell Starting cell reference (e.g. "A1") + * @param values Two-dimensional array of values to write + * @throws Error if file not opened + */ + All: (sheet: string, startCell: string, values: any[][]) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.write.All`, this.handle, sheet, startCell, values); + }, + }; + + // Setting properties + set = { + /** + * Applies a style to a cell + * @param sheet Sheet name + * @param cell Cell reference (e.g. "A1") + * @param styleID Style ID to apply + * @throws Error if file not opened + */ + Style: (sheet: string, cell: string, styleID: number) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.set.Style`, this.handle, sheet, cell, styleID); + }, + + /** + * Sets a row's height + * @param sheet Sheet name + * @param row Row number + * @param height Height in points + * @throws Error if file not opened + */ + RowHeight: (sheet: string, row: number, height: number) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.set.RowHeight`, this.handle, sheet, row, height); + }, + + /** + * Sets column width for a range of columns + * @param sheet Sheet name + * @param startCol Starting column letter + * @param endCol Ending column letter + * @param width Width in points + * @throws Error if file not opened + */ + ColumnWidth: ( + sheet: string, + startCol: string, + endCol: string, + width: number + ) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process( + `excel.set.ColumnWidth`, + this.handle, + sheet, + startCol, + endCol, + width + ); + }, + + /** + * Merges cells in a range + * @param sheet Sheet name + * @param startCell Starting cell reference (e.g. "A1") + * @param endCell Ending cell reference (e.g. "B2") + * @throws Error if file not opened + */ + MergeCell: (sheet: string, startCell: string, endCell: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process( + `excel.set.MergeCell`, + this.handle, + sheet, + startCell, + endCell + ); + }, + + /** + * Unmerges previously merged cells + * @param sheet Sheet name + * @param startCell Starting cell reference (e.g. "A1") + * @param endCell Ending cell reference (e.g. "B2") + * @throws Error if file not opened + */ + UnmergeCell: (sheet: string, startCell: string, endCell: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process( + `excel.set.UnmergeCell`, + this.handle, + sheet, + startCell, + endCell + ); + }, + + /** + * Sets a formula in a cell + * @param sheet Sheet name + * @param cell Cell reference (e.g. "C1") + * @param formula Excel formula without the leading equals sign + * @throws Error if file not opened + */ + Formula: (sheet: string, cell: string, formula: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.set.Formula`, this.handle, sheet, cell, formula); + }, + + /** + * Adds a hyperlink to a cell + * @param sheet Sheet name + * @param cell Cell reference (e.g. "A1") + * @param url URL for the hyperlink + * @param text Display text for the hyperlink + * @throws Error if file not opened + */ + Link: (sheet: string, cell: string, url: string, text: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.set.Link`, this.handle, sheet, cell, url, text); + }, + }; + + // Iteration methods + each = { + /** + * Opens a row iterator + * @param sheet Sheet name + * @returns Row iterator ID + * @throws Error if file not opened + */ + OpenRow: (sheet: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.each.OpenRow`, this.handle, sheet); + }, + + /** + * Gets the next row from the iterator + * @param rowID Row iterator ID from excel.each.OpenRow + * @returns Array of cell values or null if no more rows + */ + NextRow: (rowID: string) => { + return Process(`excel.each.NextRow`, rowID); + }, + + /** + * Closes the row iterator + * @param rowID Row iterator ID from excel.each.OpenRow + */ + CloseRow: (rowID: string) => { + return Process(`excel.each.CloseRow`, rowID); + }, + + /** + * Opens a column iterator + * @param sheet Sheet name + * @returns Column iterator ID + * @throws Error if file not opened + */ + OpenColumn: (sheet: string) => { + if (!this.handle) throw new Error("Excel file not opened"); + return Process(`excel.each.OpenColumn`, this.handle, sheet); + }, + + /** + * Gets the next column from the iterator + * @param colID Column iterator ID from excel.each.OpenColumn + * @returns Array of cell values or null if no more columns + */ + NextColumn: (colID: string) => { + return Process(`excel.each.NextColumn`, colID); + }, + + /** + * Closes the column iterator + * @param colID Column iterator ID from excel.each.OpenColumn + */ + CloseColumn: (colID: string) => { + return Process(`excel.each.CloseColumn`, colID); + }, + }; + + // Conversion utilities + convert = { + /** + * Converts a column name to a column number + * @param colName Column name (e.g. "A", "AB") + * @returns Column number (1-based) + */ + ColumnNameToNumber: (colName: string) => { + return Process(`excel.convert.ColumnNameToNumber`, colName); + }, + + /** + * Converts a column number to a column name + * @param colNum Column number (1-based) + * @returns Column name + */ + ColumnNumberToName: (colNum: number) => { + return Process(`excel.convert.ColumnNumberToName`, colNum); + }, + + /** + * Converts a cell reference to coordinates + * @param cell Cell reference (e.g. "A1") + * @returns Array with [columnNumber, rowNumber] (1-based) + */ + CellNameToCoordinates: (cell: string) => { + return Process(`excel.convert.CellNameToCoordinates`, cell); + }, + + /** + * Converts coordinates to a cell reference + * @param col Column number (1-based) + * @param row Row number (1-based) + * @returns Cell reference + */ + CoordinatesToCellName: (col: number, row: number) => { + return Process(`excel.convert.CoordinatesToCellName`, col, row); + }, + }; +} +/** + * Volcengine OpenAPI SDK + */ +import { Exception, http, Process } from "@yao/runtime"; + +export class Volcengine { + private AccessKeyId: string; + private SecretAccessKey: string; + private Region: string; + private Service: string; + private Endpoint: string; + constructor(option: Option) { + this.AccessKeyId = option.AccessKeyId; + this.SecretAccessKey = option.SecretAccessKey; + this.Region = option.Region; + this.Service = option.Service; + this.Endpoint = option.Endpoint + ? `https://${option.Endpoint}` + : `https://${this.Service}.${this.Region}.volcengineapi.com`; + } + + public Get(query: Record) { + const url = this.Endpoint; + const host = url.split("://")[1].split("/")[0]; + const headers = { host: host }; + const request: Request = { + Method: "GET", + URI: "/", + Query: query, + Headers: headers, + Payload: null, + }; + + const auth = this.getAuthorization(request); + + // Add authorization header + headers["Authorization"] = auth; + headers["Content-Type"] = "application/json"; + + const resp = http.Get(url, query, headers); + if (resp.code > 299 || resp.code < 200) { + const { ResponseMetadata } = resp.data || {}; + const { Error } = ResponseMetadata || {}; + const message = + Error?.Message || (resp.code === 0 ? resp.message : "Unknown error"); + throw new Exception(message, resp.code); + } + + return resp.data; + } + + /** + * Post request + * @param query Query parameters + * @param payload Payload + * @returns Response + */ + public Post(query: Record, payload: Record) { + const url = this.Endpoint; + const host = url.split("://")[1].split("/")[0]; + const headers = { host: host }; + const body = JSON.stringify(payload); + const request: Request = { + Method: "POST", + URI: "/", + Query: query, + Headers: headers, + Payload: body, + }; + + const auth = this.getAuthorization(request); + headers["Authorization"] = auth; + headers["Content-Type"] = "application/json"; + + const resp = http.Post(url, body, null, query, headers); + if (resp.code > 299 || resp.code < 200) { + const { ResponseMetadata } = resp.data || {}; + const { Error } = ResponseMetadata || {}; + const message = + Error?.Message || (resp.code === 0 ? resp.message : "Unknown error"); + throw new Exception(message, resp.code); + } + return resp.data; + } + + /** + * Create a canonical request + * @param request Request object + * @returns Canonical request string + */ + private canonicalRequest(request: Request): string { + const xDate = this.formatDate(new Date()); + + // 1. HTTP Method + const method = request.Method; + + // 2. URI (default to '/' if null) + const uri = request.URI || "/"; + + // 3. Query String + let queryString = ""; + if (request.Query) { + if (Array.isArray(request.Query)) { + // Handle array of query parameters + const queryParams = request.Query.reduce((acc: string[], curr) => { + Object.entries(curr).forEach(([key, value]) => { + if (value !== null && value !== undefined && value !== "") { + acc.push( + `${encodeURIComponent(key)}=${encodeURIComponent(value)}` + ); + } + }); + return acc; + }, []); + queryString = queryParams.sort().join("&"); + } else { + // Handle single query object + const queryParams = Object.entries(request.Query) + .filter( + ([_, value]) => + value !== null && value !== undefined && value !== "" + ) + .map( + ([key, value]) => + `${encodeURIComponent(key)}=${encodeURIComponent(value)}` + ) + .sort(); + queryString = queryParams.join("&"); + } + } + + // 4. Headers + // First, collect all headers in a normalized format + const headers: Record = { "x-date": xDate }; + if (request.Headers) { + if (Array.isArray(request.Headers)) { + request.Headers.forEach((headerObj) => { + Object.entries(headerObj).forEach(([key, value]) => { + if (value !== null && value !== undefined && value.trim() !== "") { + headers[key.toLowerCase()] = value.trim(); + } + }); + }); + } else { + Object.entries(request.Headers).forEach(([key, value]) => { + if (value !== null && value !== undefined && value.trim() !== "") { + headers[key.toLowerCase()] = value.trim(); + } + }); + } + } + + // Get required headers if they exist + const signedHeaderKeys: string[] = []; + const requiredHeaders = ["host", "x-date"]; + + // Add required headers first if they exist + requiredHeaders.forEach((key) => { + if (headers[key]) { + signedHeaderKeys.push(key); + } + }); + + // Add any additional headers + // const additionalHeaders = Object.keys(headers) + // .filter((key) => !requiredHeaders.includes(key)) + // .sort(); + // signedHeaderKeys.push(...additionalHeaders); + + // Build canonical headers string + const canonicalHeaders = signedHeaderKeys + .map((key) => `${key}:${headers[key]}`) + .join("\n"); + + // Build signed headers string + const signedHeaders = signedHeaderKeys.join(";"); + + // 5. Payload/Body + let hashedPayload = Process("crypto.Hash", "SHA256", ""); + if (request.Payload !== null && request.Payload !== undefined) { + if (typeof request.Payload === "string") { + if (request.Payload !== "") { + hashedPayload = Process("crypto.Hash", "SHA256", request.Payload); + } + } else { + const payload = JSON.stringify(request.Payload); + if (payload !== "{}" && payload !== "[]") { + hashedPayload = Process("crypto.Hash", "SHA256", payload); + } + } + } + + // Combine all components + const parts = [ + method, + uri, + queryString, + canonicalHeaders, + "", // Empty line after headers + signedHeaders, + hashedPayload, + ]; + + return parts.join("\n"); + } + + /** + * Format date to YYYYMMDDTHHMMSSZ + * @param date Date object + * @returns Formatted date string + */ + private formatDate(date: Date): string { + const year = date.getUTCFullYear(); + const month = String(date.getUTCMonth() + 1).padStart(2, "0"); + const day = String(date.getUTCDate()).padStart(2, "0"); + const hours = String(date.getUTCHours()).padStart(2, "0"); + const minutes = String(date.getUTCMinutes()).padStart(2, "0"); + const seconds = String(date.getUTCSeconds()).padStart(2, "0"); + return `${year}${month}${day}T${hours}${minutes}${seconds}Z`; + } + + /** + * Create string to sign + * @param canonicalRequest Canonical request string + * @returns String to sign + */ + private stringToSign(canonicalRequest: string): string { + const algorithm = "HMAC-SHA256"; + const requestDateTime = this.formatDate(new Date()); + const requestDate = requestDateTime.slice(0, 8); + const credentialScope = `${requestDate}/${this.Region}/${this.Service}/request`; // YYYYMMDD + + const hashedCanonicalRequest = Process( + "crypto.Hash", + "SHA256", + canonicalRequest + ); + + return `${algorithm}\n${requestDateTime}\n${credentialScope}\n${hashedCanonicalRequest}`; + } + + /** + * Derive signing key + * @param date Date in YYYY/MM/DD format + * @returns Signing key + */ + private getSigningKey(date: string): string { + const kDate = Process("crypto.HMAC", "SHA256", date, this.SecretAccessKey); + const kRegion = Process( + "crypto.HMACWith", + { key: "hex" }, + this.Region, + kDate + ); + const kService = Process( + "crypto.HMACWith", + { key: "hex" }, + this.Service, + kRegion + ); + + const kSigning = Process( + "crypto.HMACWith", + { key: "hex" }, + "request", + kService + ); + return kSigning; + } + + /** + * Calculate signature + * @param stringToSign String to sign + * @param signingKey Signing key + * @returns Signature + */ + private signature(stringToSign: string, signingKey: string): string { + return Process("crypto.HMACWith", { key: "hex" }, stringToSign, signingKey); + } + + /** + * Build authorization header + * @param request Request object + * @returns Authorization header value + */ + public getAuthorization(request: Request): string { + const xDate = this.formatDate(new Date()); + if (request.Headers) { + if (typeof request.Headers === "object") { + request.Headers["x-date"] = request.Headers["x-date"] + ? request.Headers["x-date"] + : xDate; + } + } + + // 1. Create canonical request + const canonicalReq = this.canonicalRequest(request); + + // 2. Create string to sign + const stringToSign = this.stringToSign(canonicalReq); + + // 3. Get date from string to sign + const [algorithm, requestDateTime, credentialScope] = + stringToSign.split("\n"); + const date = requestDateTime.slice(0, 8); + + // 4. Derive signing key + const signingKey = this.getSigningKey(date); + // 5. Calculate signature + const signature = this.signature(stringToSign, signingKey); + + // 6. Build authorization header + let signedHeaders = ""; + if (request.Headers) { + const headers: Record = {}; + if (Array.isArray(request.Headers)) { + request.Headers.forEach((headerObj) => { + Object.entries(headerObj).forEach(([key, value]) => { + headers[key.toLowerCase()] = value.trim(); + }); + }); + } else { + Object.entries(request.Headers).forEach(([key, value]) => { + headers[key.toLowerCase()] = value.trim(); + }); + } + signedHeaders = Object.keys(headers).sort().join(";"); + } + + return `${algorithm} Credential=${this.AccessKeyId}/${credentialScope}, SignedHeaders=${signedHeaders}, Signature=${signature}`; + } +} + +export interface Option { + AccessKeyId: string; + SecretAccessKey: string; + Endpoint?: string; + Region: string; + Service: string; +} + +export interface Request { + Method: "GET" | "POST"; + URI: string | null; // Default / + Query: Record | Record[] | null; + Headers: Record | Record[] | null; + Payload: string | Record | any[] | null; +} diff --git a/agent/sandbox/v2/claude/testdata/test-image.png b/agent/sandbox/v2/claude/testdata/test-image.png new file mode 100644 index 00000000..7c354fd1 Binary files /dev/null and b/agent/sandbox/v2/claude/testdata/test-image.png differ diff --git a/agent/sandbox/v2/init.go b/agent/sandbox/v2/init.go new file mode 100644 index 00000000..cd13d530 --- /dev/null +++ b/agent/sandbox/v2/init.go @@ -0,0 +1,13 @@ +package sandboxv2 + +import ( + "github.com/yaoapp/yao/agent/sandbox/v2/claude" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + yaorunner "github.com/yaoapp/yao/agent/sandbox/v2/yao" +) + +func init() { + Register("claude", func() types.Runner { return claude.New() }) + Register("claude/cli", func() types.Runner { return claude.New() }) + Register("yao", func() types.Runner { return yaorunner.New() }) +} diff --git a/agent/sandbox/v2/lifecycle.go b/agent/sandbox/v2/lifecycle.go new file mode 100644 index 00000000..afe8604e --- /dev/null +++ b/agent/sandbox/v2/lifecycle.go @@ -0,0 +1,156 @@ +package sandboxv2 + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + "log" + + "github.com/yaoapp/gou/connector" + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// BuildIdentifier determines the Computer identifier based on lifecycle policy +// and optional metadata override. Returns "" for oneshot (always new). +func BuildIdentifier(cfg *types.SandboxConfig, ownerID, chatID, assistantID string, metadata map[string]any) string { + if cfg.Lifecycle == "oneshot" { + return "" + } + + // Custom identifier from metadata takes precedence. + if metadata != nil { + if cid, ok := metadata["computer_id"].(string); ok && cid != "" { + return fmt.Sprintf("%s-%s", ownerID, cid) + } + } + + switch cfg.Lifecycle { + case "session": + return fmt.Sprintf("%s-%s", ownerID, chatID) + case "longrunning", "persistent": + return fmt.Sprintf("%s-%s", ownerID, assistantID) + default: + return "" + } +} + +// GetComputer obtains or creates a Computer for the current request. +// An optional connector may be passed to inject OPENAI_PROXY_* env vars. +// Returns the Computer, the resolved identifier, and any error. +func GetComputer(ctx *agentContext.Context, cfg *types.SandboxConfig, manager *infra.Manager, conn ...connector.Connector) (infra.Computer, string, error) { + ownerID := resolveOwnerID(ctx) + identifier := BuildIdentifier(cfg, ownerID, ctx.ChatID, ctx.AssistantID, ctx.Metadata) + + // Fill runtime fields. + cfg.Owner = ownerID + cfg.ID = identifier + + workspaceID := "" + if ctx.Metadata != nil { + if ws, ok := ctx.Metadata["workspace_id"].(string); ok && ws != "" { + workspaceID = ws + } + } + if workspaceID == "" { + workspaceID = ownerID + } + cfg.WorkspaceID = workspaceID + + // Host mode: no image → host computer. + if cfg.Computer.Image == "" { + cfg.Kind = "host" + nodeID := cfg.NodeID + if nodeID == "" { + return nil, identifier, fmt.Errorf("host mode requires a nodeID (set in sandbox.yao or workspace)") + } + host, err := manager.Host(context.Background(), nodeID) + if err != nil { + return nil, identifier, fmt.Errorf("get host computer: %w", err) + } + host.BindWorkplace(workspaceID) + return host, identifier, nil + } + + cfg.Kind = "box" + + // Reuse: non-empty identifier → try Get first. + if identifier != "" { + box, err := manager.Get(context.Background(), identifier) + if err == nil && box != nil { + box.BindWorkplace(workspaceID) + return box, identifier, nil + } + } + + // Create new box. + var c connector.Connector + if len(conn) > 0 { + c = conn[0] + } + createOpts, err := BuildCreateOptions(cfg, identifier, ownerID, workspaceID, c) + if err != nil { + return nil, identifier, fmt.Errorf("build create options: %w", err) + } + + // Oneshot with empty identifier: generate a random one. + if createOpts.ID == "" { + createOpts.ID = randomID() + identifier = createOpts.ID + cfg.ID = identifier + } + + box, err := manager.Create(context.Background(), createOpts) + if err != nil { + return nil, identifier, fmt.Errorf("create computer: %w", err) + } + return box, identifier, nil +} + +// LifecycleAction performs the post-request lifecycle operation based on policy. +// Called in defer after executeSandboxStream completes. +func LifecycleAction(ctx context.Context, cfg *types.SandboxConfig, computer infra.Computer, manager *infra.Manager) { + if computer == nil || cfg == nil { + return + } + + info := computer.ComputerInfo() + + switch cfg.Lifecycle { + case "oneshot": + if info.Kind == "box" && manager != nil { + if err := manager.Remove(ctx, cfg.ID); err != nil { + log.Printf("[sandbox/v2] oneshot remove %s: %v", cfg.ID, err) + } + } + + case "session", "longrunning": + if info.Kind == "box" && manager != nil { + manager.Heartbeat(cfg.ID, false, 0) // active=false: request finished, start idle timer + } + + case "persistent": + // No action — persistent boxes survive indefinitely. + } +} + +// resolveOwnerID returns teamID if available, otherwise userID. +func resolveOwnerID(ctx *agentContext.Context) string { + if ctx.Authorized != nil { + if ctx.Authorized.TeamID != "" { + return ctx.Authorized.TeamID + } + if ctx.Authorized.UserID != "" { + return ctx.Authorized.UserID + } + } + return "anonymous" +} + +func randomID() string { + b := make([]byte, 8) + _, _ = rand.Read(b) + return hex.EncodeToString(b) +} diff --git a/agent/sandbox/v2/lifecycle_test.go b/agent/sandbox/v2/lifecycle_test.go new file mode 100644 index 00000000..64cf2a81 --- /dev/null +++ b/agent/sandbox/v2/lifecycle_test.go @@ -0,0 +1,547 @@ +package sandboxv2_test + +import ( + "context" + "fmt" + "strings" + "testing" + "time" + + agentContext "github.com/yaoapp/yao/agent/context" + sandboxv2 "github.com/yaoapp/yao/agent/sandbox/v2" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + oauthTypes "github.com/yaoapp/yao/openapi/oauth/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// =========================================================================== +// BuildIdentifier — pure-function tests (no infra needed) +// =========================================================================== + +func TestBuildIdentifier_Oneshot(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "oneshot"} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", nil) + if id != "" { + t.Errorf("oneshot should return empty, got %q", id) + } +} + +func TestBuildIdentifier_Session(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "session"} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat42", "ast1", nil) + if id != "owner1-chat42" { + t.Errorf("session: got %q, want %q", id, "owner1-chat42") + } +} + +func TestBuildIdentifier_Longrunning(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "longrunning"} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast99", nil) + if id != "owner1-ast99" { + t.Errorf("longrunning: got %q, want %q", id, "owner1-ast99") + } +} + +func TestBuildIdentifier_Persistent(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "persistent"} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast99", nil) + if id != "owner1-ast99" { + t.Errorf("persistent: got %q, want %q", id, "owner1-ast99") + } +} + +func TestBuildIdentifier_MetadataOverride(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "session"} + meta := map[string]any{"computer_id": "custom-box"} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", meta) + if id != "owner1-custom-box" { + t.Errorf("metadata override: got %q, want %q", id, "owner1-custom-box") + } +} + +func TestBuildIdentifier_MetadataEmptyIgnored(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "session"} + meta := map[string]any{"computer_id": ""} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat42", "ast1", meta) + if id != "owner1-chat42" { + t.Errorf("empty metadata should fall through to session, got %q", id) + } +} + +func TestBuildIdentifier_UnknownLifecycle(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "unknown"} + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", nil) + if id != "" { + t.Errorf("unknown lifecycle should return empty, got %q", id) + } +} + +// =========================================================================== +// GetComputer — real container tests +// =========================================================================== + +func makeAgentCtx(teamID, userID, chatID, assistantID string, metadata map[string]any) *agentContext.Context { + var auth *oauthTypes.AuthorizedInfo + if teamID != "" || userID != "" { + auth = &oauthTypes.AuthorizedInfo{TeamID: teamID, UserID: userID} + } + return &agentContext.Context{ + Context: context.Background(), + Authorized: auth, + ChatID: chatID, + AssistantID: assistantID, + Metadata: metadata, + } +} + +func TestGetComputer_BoxCreate(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + wsID := fmt.Sprintf("lc-create-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + meta := map[string]any{"workspace_id": wsID} + ctx := makeAgentCtx("team-t1", "", "chat-1", "ast-1", meta) + + computer, identifier, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + + if identifier == "" { + t.Fatal("oneshot should get a random identifier, got empty") + } + + info := computer.ComputerInfo() + if info.Kind != "box" { + t.Errorf("kind = %q, want %q", info.Kind, "box") + } + if cfg.Owner != "team-t1" { + t.Errorf("cfg.Owner = %q, want %q", cfg.Owner, "team-t1") + } + if cfg.Kind != "box" { + t.Errorf("cfg.Kind = %q, want %q", cfg.Kind, "box") + } + }) + } +} + +func TestGetComputer_BoxReuse(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + wsID := fmt.Sprintf("lc-reuse-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "session", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + meta := map[string]any{"workspace_id": wsID} + ctx := makeAgentCtx("team-reuse", "", "chat-reuse", "ast-1", meta) + + computer1, id1, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("first GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + + cfg2 := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "session", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + computer2, id2, err := sandboxv2.GetComputer(ctx, cfg2, m) + if err != nil { + t.Fatalf("second GetComputer: %v", err) + } + + if id1 != id2 { + t.Errorf("identifiers differ: %q vs %q", id1, id2) + } + + info1 := computer1.ComputerInfo() + info2 := computer2.ComputerInfo() + if info1.ContainerID != info2.ContainerID { + t.Errorf("container IDs differ: %q vs %q (should reuse)", info1.ContainerID, info2.ContainerID) + } + }) + } +} + +func TestGetComputer_WorkspaceBindAlways(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + wsID := fmt.Sprintf("lc-ws-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + meta := map[string]any{"workspace_id": wsID} + ctx := makeAgentCtx("team-ws", "", "chat-ws", "ast-ws", meta) + + computer, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + + if cfg.WorkspaceID != wsID { + t.Errorf("WorkspaceID = %q, want %q", cfg.WorkspaceID, wsID) + } + + ws := computer.Workplace() + if ws == nil { + t.Fatal("Workplace() returned nil, workspace should always be bound") + } + }) + } +} + +func TestGetComputer_WorkspaceFallbackOwner(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + ownerID := fmt.Sprintf("lc-owner-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, ownerID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx(ownerID, "", "chat-fb", "ast-fb", nil) + + computer, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + + if cfg.WorkspaceID != ownerID { + t.Errorf("WorkspaceID = %q, want %q (should fallback to ownerID)", cfg.WorkspaceID, ownerID) + } + + ws := computer.Workplace() + if ws == nil { + t.Fatal("Workplace() returned nil") + } + }) + } +} + +func TestGetComputer_OwnerPriority(t *testing.T) { + skipIfNoDocker(t) + + nc := boxNodes()[0] + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + t.Run("teamID", func(t *testing.T) { + wsID := fmt.Sprintf("lc-ownp-team-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx("my-team", "my-user", "c", "a", map[string]any{"workspace_id": wsID}) + _, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + if cfg.Owner != "my-team" { + t.Errorf("Owner = %q, want %q (teamID takes precedence)", cfg.Owner, "my-team") + } + }) + + t.Run("userID", func(t *testing.T) { + wsID := fmt.Sprintf("lc-ownp-user-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx("", "my-user", "c", "a", map[string]any{"workspace_id": wsID}) + _, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + if cfg.Owner != "my-user" { + t.Errorf("Owner = %q, want %q", cfg.Owner, "my-user") + } + }) + + t.Run("anonymous", func(t *testing.T) { + wsID := fmt.Sprintf("lc-ownp-anon-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx("", "", "c", "a", map[string]any{"workspace_id": wsID}) + _, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + if cfg.Owner != "anonymous" { + t.Errorf("Owner = %q, want %q", cfg.Owner, "anonymous") + } + }) +} + +func TestGetComputer_HostMode(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "session", + Computer: types.ComputerConfig{}, + NodeID: tgt.TaiID, + } + ctx := makeAgentCtx("team-host", "", "chat-host", "ast-host", nil) + + computer, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer host: %v", err) + } + + if cfg.Kind != "host" { + t.Errorf("Kind = %q, want %q", cfg.Kind, "host") + } + + info := computer.ComputerInfo() + if info.Kind != "host" { + t.Errorf("ComputerInfo.Kind = %q, want %q", info.Kind, "host") + } + + ws := computer.Workplace() + if ws == nil { + t.Fatal("Workplace() returned nil on host mode") + } + }) + } +} + +func TestGetComputer_HostMissingNodeID(t *testing.T) { + skipIfNoDocker(t) + + nc := boxNodes()[0] + m := setupManager(t, &nc) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "session", + Computer: types.ComputerConfig{}, + NodeID: "", + } + ctx := makeAgentCtx("team-err", "", "c", "a", nil) + + _, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err == nil { + t.Fatal("expected error for host mode without nodeID") + } + if !strings.Contains(err.Error(), "nodeID") { + t.Errorf("error should mention nodeID, got: %v", err) + } +} + +// =========================================================================== +// LifecycleAction — behavior tests +// =========================================================================== + +func TestLifecycleAction_Oneshot(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + wsID := fmt.Sprintf("lc-oneshot-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "oneshot", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx("team-oneshot", "", "c", "a", map[string]any{"workspace_id": wsID}) + + computer, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + + boxID := cfg.ID + + sandboxv2.LifecycleAction(context.Background(), cfg, computer, m) + + _, getErr := m.Get(context.Background(), boxID) + if getErr == nil { + t.Error("box should be removed after oneshot LifecycleAction") + } + }) + } +} + +func TestLifecycleAction_Session(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + wsID := fmt.Sprintf("lc-sess-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "session", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx("team-sess", "", "chat-sess", "ast-sess", map[string]any{"workspace_id": wsID}) + + computer, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + + sandboxv2.LifecycleAction(context.Background(), cfg, computer, m) + + box, err := m.Get(context.Background(), cfg.ID) + if err != nil { + t.Fatalf("box should still exist after session LifecycleAction: %v", err) + } + if box == nil { + t.Fatal("box is nil after session LifecycleAction") + } + }) + } +} + +func TestLifecycleAction_Persistent(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + ensureImage(t, m, nc) + + wsID := fmt.Sprintf("lc-pers-%d", time.Now().UnixNano()) + createTestWorkspace(t, nc.TaiID, wsID) + + cfg := &types.SandboxConfig{ + Version: "2.0", + Lifecycle: "persistent", + Computer: types.ComputerConfig{Image: testImage()}, + NodeID: nc.TaiID, + } + ctx := makeAgentCtx("team-pers", "", "chat-pers", "ast-pers", map[string]any{"workspace_id": wsID}) + + computer, _, err := sandboxv2.GetComputer(ctx, cfg, m) + if err != nil { + t.Fatalf("GetComputer: %v", err) + } + defer cleanupComputer(t, m, cfg) + + sandboxv2.LifecycleAction(context.Background(), cfg, computer, m) + + box, err := m.Get(context.Background(), cfg.ID) + if err != nil { + t.Fatalf("box should still exist after persistent LifecycleAction: %v", err) + } + if box == nil { + t.Fatal("box is nil after persistent LifecycleAction") + } + }) + } +} + +func TestLifecycleAction_NilSafe(t *testing.T) { + cfg := &types.SandboxConfig{Lifecycle: "oneshot"} + sandboxv2.LifecycleAction(context.Background(), cfg, nil, nil) + sandboxv2.LifecycleAction(context.Background(), nil, nil, nil) +} + +// =========================================================================== +// helpers +// =========================================================================== + +func ensureImage(t *testing.T, m *infra.Manager, nc nodeConfig) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + defer cancel() + if err := m.EnsureImage(ctx, nc.TaiID, testImage(), infra.ImagePullOptions{}); err != nil { + t.Fatalf("EnsureImage: %v", err) + } +} + +func cleanupComputer(t *testing.T, m *infra.Manager, cfg *types.SandboxConfig) { + t.Helper() + if cfg.ID == "" { + return + } + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + if err := m.Remove(ctx, cfg.ID); err != nil { + t.Logf("cleanup Remove(%s): %v", cfg.ID, err) + } +} diff --git a/agent/sandbox/v2/options.go b/agent/sandbox/v2/options.go new file mode 100644 index 00000000..2ceb3392 --- /dev/null +++ b/agent/sandbox/v2/options.go @@ -0,0 +1,231 @@ +package sandboxv2 + +import ( + "encoding/json" + "fmt" + "os" + "strings" + "time" + + "github.com/yaoapp/gou/connector" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// resolveEnvRef resolves $ENV.XXX references to os.Getenv("XXX"). +func resolveEnvRef(value string) string { + if strings.HasPrefix(value, "$ENV.") { + return os.Getenv(value[5:]) + } + return value +} + +// BuildCreateOptions converts a SandboxConfig into the V2 infrastructure +// CreateOptions. An optional connector is used to inject OPENAI_PROXY_* +// environment variables when the connector is OpenAI-compatible (non-Anthropic). +func BuildCreateOptions(cfg *types.SandboxConfig, identifier, ownerID, workspaceID string, conn ...connector.Connector) (infra.CreateOptions, error) { + opts := infra.CreateOptions{ + ID: identifier, + Owner: ownerID, + Image: cfg.Computer.Image, + WorkDir: cfg.Computer.WorkDir, + User: cfg.Computer.User, + MountPath: cfg.Computer.MountPath, + MountMode: cfg.Computer.MountMode, + WorkspaceID: workspaceID, + Labels: cfg.Labels, + } + + if opts.Labels == nil { + opts.Labels = make(map[string]string) + } + + // Lifecycle policy + switch cfg.Lifecycle { + case "oneshot": + opts.Policy = infra.OneShot + case "session": + opts.Policy = infra.Session + case "longrunning": + opts.Policy = infra.LongRunning + case "persistent": + opts.Policy = infra.Persistent + default: + opts.Policy = infra.OneShot + } + + // Timeouts + if cfg.IdleTimeout != "" { + d, err := time.ParseDuration(cfg.IdleTimeout) + if err != nil { + return opts, fmt.Errorf("idle_timeout: %w", err) + } + opts.IdleTimeout = d + } + if cfg.MaxLifetime != "" { + d, err := time.ParseDuration(cfg.MaxLifetime) + if err != nil { + return opts, fmt.Errorf("max_lifetime: %w", err) + } + opts.MaxLifetime = d + } + if cfg.StopTimeout != "" { + d, err := time.ParseDuration(cfg.StopTimeout) + if err != nil { + return opts, fmt.Errorf("stop_timeout: %w", err) + } + opts.StopTimeout = d + } + + // Memory (string like "4g" → bytes) + if cfg.Computer.Memory != "" { + mem, err := parseMemory(cfg.Computer.Memory) + if err != nil { + return opts, fmt.Errorf("memory: %w", err) + } + opts.Memory = mem + } + + opts.CPUs = cfg.Computer.CPUs + + // VNC + opts.VNC = cfg.Computer.VNC.Enabled + + // Ports + for _, p := range cfg.Computer.Ports { + opts.Ports = append(opts.Ports, infra.PortMapping{ + ContainerPort: p.Port, + HostPort: p.HostPort, + Protocol: p.Protocol, + }) + } + + // NodeID (host mode pre-selection) + if cfg.NodeID != "" { + opts.NodeID = cfg.NodeID + } + + // Merge environment + secrets into CreateOptions.Env. + // Secrets override environment for same-name keys. + // $ENV.XXX references are resolved at runtime. + envSize := len(cfg.Environment) + len(cfg.Secrets) + if envSize > 0 { + opts.Env = make(map[string]string, envSize) + for k, v := range cfg.Environment { + opts.Env[k] = resolveEnvRef(v) + } + for k, v := range cfg.Secrets { + opts.Env[k] = resolveEnvRef(v) + } + } + + if opts.Env == nil { + opts.Env = make(map[string]string) + } + + // Inject OPENAI_PROXY_* when connector is OpenAI-compatible (non-Anthropic). + // The a2o proxy inside the container translates Anthropic API → OpenAI API. + if len(conn) > 0 && conn[0] != nil && !conn[0].Is(connector.ANTHROPIC) { + injectProxyEnv(opts.Env, conn[0]) + } + + // Inject VNC_* environment variables from config. + if cfg.Computer.VNC.Enabled { + opts.Env["VNC_ENABLED"] = "true" + if cfg.Computer.VNC.Password != "" { + opts.Env["VNC_PASSWORD"] = resolveEnvRef(cfg.Computer.VNC.Password) + } + if cfg.Computer.VNC.Resolution != "" { + opts.Env["VNC_RESOLUTION"] = cfg.Computer.VNC.Resolution + } + if cfg.Computer.VNC.ViewOnly { + opts.Env["VNC_VIEW_ONLY"] = "true" + } + } + + return opts, nil +} + +// injectProxyEnv extracts backend URL, model, and API key from an +// OpenAI-compatible connector's settings and writes them as OPENAI_PROXY_* +// environment variables into env. +func injectProxyEnv(env map[string]string, conn connector.Connector) { + settings := conn.Setting() + if settings == nil { + return + } + + if host, ok := settings["host"].(string); ok && host != "" { + env["OPENAI_PROXY_BACKEND"] = host + } + if model, ok := settings["model"].(string); ok && model != "" { + env["OPENAI_PROXY_MODEL"] = model + } + if key, ok := settings["key"].(string); ok && key != "" { + env["OPENAI_PROXY_API_KEY"] = key + } + + // Forward extra options as JSON. + extra := make(map[string]interface{}) + for k, v := range settings { + switch k { + case "host", "model", "key", "proxy", "type": + continue + default: + extra[k] = v + } + } + if len(extra) > 0 { + if data, err := json.Marshal(extra); err == nil { + env["OPENAI_PROXY_OPTIONS"] = string(data) + } + } +} + +// parseMemory converts a human-readable memory string to bytes. +// Supported formats: "4GB", "4G", "4g", "512MB", "512M", "512m", "1024KB", "1024K", "1024". +func parseMemory(s string) (int64, error) { + if len(s) == 0 { + return 0, nil + } + + upper := strings.ToUpper(s) + var num string + var multiplier int64 + + switch { + case strings.HasSuffix(upper, "GB"): + num = s[:len(s)-2] + multiplier = 1 << 30 + case strings.HasSuffix(upper, "MB"): + num = s[:len(s)-2] + multiplier = 1 << 20 + case strings.HasSuffix(upper, "KB"): + num = s[:len(s)-2] + multiplier = 1 << 10 + case strings.HasSuffix(upper, "TB"): + num = s[:len(s)-2] + multiplier = 1 << 40 + case strings.HasSuffix(upper, "G"): + num = s[:len(s)-1] + multiplier = 1 << 30 + case strings.HasSuffix(upper, "M"): + num = s[:len(s)-1] + multiplier = 1 << 20 + case strings.HasSuffix(upper, "K"): + num = s[:len(s)-1] + multiplier = 1 << 10 + case strings.HasSuffix(upper, "T"): + num = s[:len(s)-1] + multiplier = 1 << 40 + default: + num = s + multiplier = 1 + } + + var val float64 + if _, err := fmt.Sscanf(num, "%f", &val); err != nil { + return 0, fmt.Errorf("invalid memory value %q", s) + } + return int64(val * float64(multiplier)), nil +} diff --git a/agent/sandbox/v2/prepare.go b/agent/sandbox/v2/prepare.go new file mode 100644 index 00000000..26e50843 --- /dev/null +++ b/agent/sandbox/v2/prepare.go @@ -0,0 +1,171 @@ +package sandboxv2 + +import ( + "context" + "fmt" + "log" + "path" + "strings" + + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai/workspace" +) + +const onceMarkerDir = ".yao/prepare" + +// RunPrepareSteps executes a list of PrepareStep actions on the given Computer. +// file/copy/marker operations use computer.Workplace() (gRPC volume, cross-platform). +// exec operations use shell via Computer.Exec. +func RunPrepareSteps(ctx context.Context, steps []types.PrepareStep, computer infra.Computer, assistantID, configHash string) error { + if len(steps) == 0 { + return nil + } + + var ws workspace.FS + if computer != nil { + ws = computer.Workplace() + } + + markerDir := onceMarkerDir + if assistantID != "" { + markerDir = onceMarkerDir + "/" + assistantID + } + markerPath := markerDir + "/done" + + skipOnce := false + if configHash != "" && ws != nil { + if data, err := ws.ReadFile(markerPath); err == nil { + if strings.TrimSpace(string(data)) == configHash { + skipOnce = true + } + } + } + + for i, step := range steps { + if step.Once && skipOnce { + continue + } + + var err error + switch step.Action { + case "file": + err = runFileStep(ws, step) + case "copy": + err = runCopyStep(ws, step) + case "exec": + err = runExecStep(ctx, computer, step) + case "process": + log.Printf("[sandbox/v2] prepare step %d: action=process (reserved, skipping)", i) + default: + err = fmt.Errorf("unknown prepare action %q", step.Action) + } + + if err != nil { + if step.IgnoreError { + log.Printf("[sandbox/v2] prepare step %d (%s): ignored error: %v", i, step.Action, err) + continue + } + return fmt.Errorf("prepare step %d (%s): %w", i, step.Action, err) + } + } + + if configHash != "" && ws != nil { + ws.MkdirAll(markerDir, 0755) + ws.WriteFile(markerPath, []byte(configHash), 0644) + } + + return nil +} + +// --------------------------------------------------------------------------- +// Step runners +// --------------------------------------------------------------------------- + +func runFileStep(ws workspace.FS, step types.PrepareStep) error { + if step.Path == "" { + return fmt.Errorf("file step requires path") + } + if ws == nil { + return fmt.Errorf("file step requires workspace") + } + + dir := path.Dir(step.Path) + if dir != "." && dir != "/" { + if err := ws.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("mkdir %s: %w", dir, err) + } + } + + if err := ws.WriteFile(step.Path, step.Content, 0644); err != nil { + return fmt.Errorf("write file %s: %w", step.Path, err) + } + return nil +} + +func runCopyStep(ws workspace.FS, step types.PrepareStep) error { + if step.Src == "" || step.Dst == "" { + return fmt.Errorf("copy step requires src and dst") + } + if ws == nil { + return fmt.Errorf("copy step requires workspace") + } + + data, err := ws.ReadFile(step.Src) + if err != nil { + return fmt.Errorf("read src %s: %w", step.Src, err) + } + + dir := path.Dir(step.Dst) + if dir != "." && dir != "/" { + if err := ws.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("mkdir %s: %w", dir, err) + } + } + + if err := ws.WriteFile(step.Dst, data, 0644); err != nil { + return fmt.Errorf("write dst %s: %w", step.Dst, err) + } + return nil +} + +func runExecStep(ctx context.Context, computer infra.Computer, step types.PrepareStep) error { + if step.Cmd == "" { + return fmt.Errorf("exec step requires cmd") + } + + kind := shellFromSystem(computer) + script := step.Cmd + if step.Background { + if kind == shellSh { + script = fmt.Sprintf("nohup %s > /dev/null 2>&1 &", step.Cmd) + } else { + script = fmt.Sprintf("Start-Process -NoNewWindow -FilePath 'cmd.exe' -ArgumentList '/C %s'", step.Cmd) + } + } + + result, err := computer.Exec(ctx, shellWrap(kind, script), infra.WithWorkDir("/")) + if err != nil { + return err + } + label := "exec" + if step.Background { + label = "exec(background)" + } + return checkResult(result, label) +} + +// checkResult inspects ExecResult for errors. +func checkResult(result *infra.ExecResult, label string) error { + if result.Error != "" { + return fmt.Errorf("%s: %s", label, result.Error) + } + if result.ExitCode != 0 { + stderr := result.Stderr + if len(stderr) > 200 { + stderr = stderr[:200] + "..." + } + return fmt.Errorf("%s: exit %d: %s", label, result.ExitCode, stderr) + } + return nil +} diff --git a/agent/sandbox/v2/prepare_test.go b/agent/sandbox/v2/prepare_test.go new file mode 100644 index 00000000..b607a83e --- /dev/null +++ b/agent/sandbox/v2/prepare_test.go @@ -0,0 +1,549 @@ +package sandboxv2_test + +import ( + "context" + "fmt" + "strings" + "testing" + "time" + + sandboxv2 "github.com/yaoapp/yao/agent/sandbox/v2" + "github.com/yaoapp/yao/agent/sandbox/v2/types" +) + +// --------------------------------------------------------------------------- +// Box tests (local + remote) +// --------------------------------------------------------------------------- + +func TestRunPrepareSteps_Exec(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "exec", Cmd: "echo hello > /tmp/prep-test"}, + {Action: "exec", Cmd: "echo world >> /tmp/prep-test"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err != nil { + t.Fatalf("RunPrepareSteps: %v", err) + } + + result, err := box.Exec(ctx, []string{"cat", "/tmp/prep-test"}) + if err != nil { + t.Fatalf("cat: %v", err) + } + got := strings.TrimSpace(result.Stdout) + if got != "hello\nworld" { + t.Errorf("content = %q, want %q", got, "hello\nworld") + } + }) + } +} + +func TestRunPrepareSteps_File(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + wsID := fmt.Sprintf("test-file-%d", time.Now().UnixNano()) + box.BindWorkplace(wsID) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "file", Path: "config/test.txt", Content: []byte("file-content-v2")}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err != nil { + t.Fatalf("RunPrepareSteps: %v", err) + } + + ws := box.Workplace() + data, err := ws.ReadFile("config/test.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "file-content-v2" { + t.Errorf("content = %q, want %q", string(data), "file-content-v2") + } + }) + } +} + +func TestRunPrepareSteps_Copy(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + wsID := fmt.Sprintf("test-copy-%d", time.Now().UnixNano()) + box.BindWorkplace(wsID) + + ws := box.Workplace() + ws.WriteFile("src.txt", []byte("copy-src"), 0644) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "copy", Src: "src.txt", Dst: "dst.txt"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err != nil { + t.Fatalf("RunPrepareSteps: %v", err) + } + + data, err := ws.ReadFile("dst.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "copy-src" { + t.Errorf("content = %q, want %q", string(data), "copy-src") + } + }) + } +} + +func TestRunPrepareSteps_OnceMarker(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + wsID := fmt.Sprintf("test-once-%d", time.Now().UnixNano()) + box.BindWorkplace(wsID) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + counter := "/tmp/once-counter" + steps := []types.PrepareStep{ + {Action: "exec", Cmd: "echo -n x >> " + counter, Once: true}, + } + + hash := "abc123" + assistantID := "test-once" + + if err := sandboxv2.RunPrepareSteps(ctx, steps, box, assistantID, hash); err != nil { + t.Fatalf("first run: %v", err) + } + r1, _ := box.Exec(ctx, []string{"cat", counter}) + if r1.Stdout != "x" { + t.Fatalf("first run: got %q, want %q", r1.Stdout, "x") + } + + if err := sandboxv2.RunPrepareSteps(ctx, steps, box, assistantID, hash); err != nil { + t.Fatalf("second run: %v", err) + } + r2, _ := box.Exec(ctx, []string{"cat", counter}) + if r2.Stdout != "x" { + t.Errorf("second run: got %q, want %q (once step should be skipped)", r2.Stdout, "x") + } + + if err := sandboxv2.RunPrepareSteps(ctx, steps, box, assistantID, "new-hash"); err != nil { + t.Fatalf("third run: %v", err) + } + r3, _ := box.Exec(ctx, []string{"cat", counter}) + if r3.Stdout != "xx" { + t.Errorf("third run: got %q, want %q (hash changed, should re-execute)", r3.Stdout, "xx") + } + }) + } +} + +func TestRunPrepareSteps_OnceIsolation(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + wsID := fmt.Sprintf("test-iso-%d", time.Now().UnixNano()) + box.BindWorkplace(wsID) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + stepsA := []types.PrepareStep{ + {Action: "exec", Cmd: "echo -n A >> /tmp/iso-a", Once: true}, + } + stepsB := []types.PrepareStep{ + {Action: "exec", Cmd: "echo -n B >> /tmp/iso-b", Once: true}, + } + + hash := "same-hash" + + if err := sandboxv2.RunPrepareSteps(ctx, stepsA, box, "assistant-a", hash); err != nil { + t.Fatalf("assistant-a: %v", err) + } + if err := sandboxv2.RunPrepareSteps(ctx, stepsB, box, "assistant-b", hash); err != nil { + t.Fatalf("assistant-b: %v", err) + } + + rA, _ := box.Exec(ctx, []string{"cat", "/tmp/iso-a"}) + rB, _ := box.Exec(ctx, []string{"cat", "/tmp/iso-b"}) + if rA.Stdout != "A" { + t.Errorf("assistant-a: got %q, want %q", rA.Stdout, "A") + } + if rB.Stdout != "B" { + t.Errorf("assistant-b: got %q, want %q", rB.Stdout, "B") + } + + if err := sandboxv2.RunPrepareSteps(ctx, stepsA, box, "assistant-a", hash); err != nil { + t.Fatalf("assistant-a re-run: %v", err) + } + if err := sandboxv2.RunPrepareSteps(ctx, stepsB, box, "assistant-b", hash); err != nil { + t.Fatalf("assistant-b re-run: %v", err) + } + rA2, _ := box.Exec(ctx, []string{"cat", "/tmp/iso-a"}) + rB2, _ := box.Exec(ctx, []string{"cat", "/tmp/iso-b"}) + if rA2.Stdout != "A" { + t.Errorf("assistant-a re-run: got %q, want %q (should be skipped)", rA2.Stdout, "A") + } + if rB2.Stdout != "B" { + t.Errorf("assistant-b re-run: got %q, want %q (should be skipped)", rB2.Stdout, "B") + } + }) + } +} + +func TestRunPrepareSteps_IgnoreError(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "exec", Cmd: "false", IgnoreError: true}, + {Action: "exec", Cmd: "echo survived > /tmp/survived"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err != nil { + t.Fatalf("RunPrepareSteps: %v (ignore_error should have prevented failure)", err) + } + + result, _ := box.Exec(ctx, []string{"cat", "/tmp/survived"}) + if strings.TrimSpace(result.Stdout) != "survived" { + t.Errorf("second step should have executed, got %q", result.Stdout) + } + }) + } +} + +func TestRunPrepareSteps_FailOnError(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "exec", Cmd: "false"}, + {Action: "exec", Cmd: "echo should-not-reach > /tmp/unreachable"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err == nil { + t.Fatal("expected error from failing step without ignore_error") + } + + result, _ := box.Exec(ctx, []string{"cat", "/tmp/unreachable"}) + if result.ExitCode == 0 { + t.Error("second step should not have executed") + } + }) + } +} + +func TestRunPrepareSteps_UnknownAction(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + _ = createBox(t, m, nc) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "unknown_action"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, nil, "test-assistant", "") + if err == nil { + t.Fatal("expected error for unknown action") + } + if !strings.Contains(err.Error(), "unknown_action") { + t.Errorf("error should mention action name, got: %v", err) + } + }) + } +} + +func TestRunPrepareSteps_EmptySteps(t *testing.T) { + err := sandboxv2.RunPrepareSteps(context.Background(), nil, nil, "test-assistant", "hash") + if err != nil { + t.Fatalf("empty steps should succeed: %v", err) + } +} + +func TestRunPrepareSteps_Background(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "exec", Cmd: "sleep 30", Background: true}, + {Action: "exec", Cmd: "echo after-bg > /tmp/after-bg"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err != nil { + t.Fatalf("RunPrepareSteps: %v", err) + } + + result, _ := box.Exec(ctx, []string{"cat", "/tmp/after-bg"}) + if strings.TrimSpace(result.Stdout) != "after-bg" { + t.Errorf("background step blocked execution, got %q", result.Stdout) + } + }) + } +} + +func TestRunPrepareSteps_MixedActions(t *testing.T) { + skipIfNoDocker(t) + + for _, nc := range boxNodes() { + nc := nc + t.Run(nc.Name, func(t *testing.T) { + m := setupManager(t, &nc) + box := createBox(t, m, nc) + wsID := fmt.Sprintf("test-mixed-%d", time.Now().UnixNano()) + box.BindWorkplace(wsID) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "file", Path: "mixed.conf", Content: []byte("key=value")}, + {Action: "exec", Cmd: "echo exec-ok > /tmp/mixed-exec"}, + {Action: "copy", Src: "mixed.conf", Dst: "mixed-copy.conf"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, box, "test-assistant", "") + if err != nil { + t.Fatalf("RunPrepareSteps: %v", err) + } + + ws := box.Workplace() + data, err := ws.ReadFile("mixed-copy.conf") + if err != nil { + t.Fatalf("ReadFile mixed-copy.conf: %v", err) + } + if string(data) != "key=value" { + t.Errorf("copy result: got %q, want %q", string(data), "key=value") + } + + result, _ := box.Exec(ctx, []string{"cat", "/tmp/mixed-exec"}) + if strings.TrimSpace(result.Stdout) != "exec-ok" { + t.Errorf("exec result: got %q, want %q", result.Stdout, "exec-ok") + } + }) + } +} + +// --------------------------------------------------------------------------- +// HostExec tests +// --------------------------------------------------------------------------- + +func TestRunPrepareSteps_HostExec(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + host := createHost(t, m, tgt) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + t.Logf("SystemInfo: OS=%q Shell=%q TempDir=%q", + host.ComputerInfo().System.OS, + host.ComputerInfo().System.Shell, + host.ComputerInfo().System.TempDir) + + isWin := tgt.Name == "win-native" + var cmd string + if isWin { + cmd = `Write-Output 'host-ok'` + } else { + cmd = "echo host-ok" + } + steps := []types.PrepareStep{ + {Action: "exec", Cmd: cmd}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, host, "test-host", "") + if err != nil { + t.Fatalf("RunPrepareSteps on host: %v", err) + } + }) + } +} + +func TestRunPrepareSteps_HostExecFile(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + host := createHost(t, m, tgt) + wsID := fmt.Sprintf("test-hostfile-%d", time.Now().UnixNano()) + host.BindWorkplace(wsID) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "file", Path: "host-test.txt", Content: []byte("host-file-data")}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, host, "test-host", "") + if err != nil { + t.Fatalf("RunPrepareSteps file: %v", err) + } + + ws := host.Workplace() + data, err := ws.ReadFile("host-test.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "host-file-data" { + t.Errorf("content = %q, want %q", string(data), "host-file-data") + } + }) + } +} + +func TestRunPrepareSteps_HostExecCopy(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + host := createHost(t, m, tgt) + wsID := fmt.Sprintf("test-hostcopy-%d", time.Now().UnixNano()) + host.BindWorkplace(wsID) + + ws := host.Workplace() + ws.WriteFile("copy-src.txt", []byte("copy-data"), 0644) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + steps := []types.PrepareStep{ + {Action: "copy", Src: "copy-src.txt", Dst: "copy-dst.txt"}, + } + + err := sandboxv2.RunPrepareSteps(ctx, steps, host, "test-host", "") + if err != nil { + t.Fatalf("RunPrepareSteps copy: %v", err) + } + + data, err := ws.ReadFile("copy-dst.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "copy-data" { + t.Errorf("content = %q, want %q", string(data), "copy-data") + } + }) + } +} + +func TestRunPrepareSteps_HostExecOnce(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + host := createHost(t, m, tgt) + wsID := fmt.Sprintf("test-hostonce-%d", time.Now().UnixNano()) + host.BindWorkplace(wsID) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + isWin := tgt.Name == "win-native" + var cmd string + if isWin { + cmd = `Write-Output 'once-ok'` + } else { + cmd = "echo once-ok" + } + steps := []types.PrepareStep{ + {Action: "exec", Cmd: cmd, Once: true}, + } + hash := "host-once-hash" + aid := "host-once-aid" + + if err := sandboxv2.RunPrepareSteps(ctx, steps, host, aid, hash); err != nil { + t.Fatalf("first run: %v", err) + } + + ws := host.Workplace() + markerData, err := ws.ReadFile(".yao/prepare/" + aid + "/done") + if err != nil { + t.Fatalf("marker not written: %v", err) + } + if string(markerData) != hash { + t.Errorf("marker = %q, want %q", string(markerData), hash) + } + }) + } +} diff --git a/agent/sandbox/v2/runner.go b/agent/sandbox/v2/runner.go new file mode 100644 index 00000000..b2132607 --- /dev/null +++ b/agent/sandbox/v2/runner.go @@ -0,0 +1,32 @@ +package sandboxv2 + +import ( + "fmt" + "sync" + + "github.com/yaoapp/yao/agent/sandbox/v2/types" +) + +var ( + mu sync.RWMutex + runners = map[string]func() types.Runner{} +) + +// Register adds a runner factory to the global registry. +// Typically called from init() in the runner's package. +func Register(name string, factory func() types.Runner) { + mu.Lock() + defer mu.Unlock() + runners[name] = factory +} + +// Get creates a new Runner instance from the registry. +func Get(name string) (types.Runner, error) { + mu.RLock() + defer mu.RUnlock() + factory, ok := runners[name] + if !ok { + return nil, fmt.Errorf("sandbox runner %q not registered", name) + } + return factory(), nil +} diff --git a/agent/sandbox/v2/shell.go b/agent/sandbox/v2/shell.go new file mode 100644 index 00000000..565ca09b --- /dev/null +++ b/agent/sandbox/v2/shell.go @@ -0,0 +1,47 @@ +package sandboxv2 + +import ( + "strings" + + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// shellKind identifies which shell to use for command execution. +type shellKind int + +const ( + shellSh shellKind = iota // Unix: sh -c + shellPwsh // Windows: pwsh -NoProfile -Command + shellPS // Windows: powershell -NoProfile -Command + shellCmd // Windows: cmd.exe /C (last-resort fallback) +) + +// shellWrap returns the Exec command slice to run a script string. +func shellWrap(kind shellKind, script string) []string { + switch kind { + case shellPwsh: + return []string{"pwsh", "-NoProfile", "-Command", script} + case shellPS: + return []string{"powershell", "-NoProfile", "-Command", script} + case shellCmd: + return []string{"cmd.exe", "/C", script} + default: + return []string{"sh", "-c", script} + } +} + +// shellFromSystem resolves shellKind from ComputerInfo().System.Shell +// reported by the Tai node at registration time. +func shellFromSystem(computer infra.Computer) shellKind { + shell := strings.ToLower(computer.ComputerInfo().System.Shell) + switch shell { + case "pwsh": + return shellPwsh + case "powershell": + return shellPS + case "cmd.exe", "cmd": + return shellCmd + default: + return shellSh + } +} diff --git a/agent/sandbox/v2/stream.go b/agent/sandbox/v2/stream.go new file mode 100644 index 00000000..2066cac5 --- /dev/null +++ b/agent/sandbox/v2/stream.go @@ -0,0 +1,138 @@ +package sandboxv2 + +import ( + "context" + "errors" + "fmt" + "log" + "time" + + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/output/message" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// ExecuteRequest consolidates all parameters for ExecuteSandboxStream. +type ExecuteRequest struct { + Computer infra.Computer + Runner types.Runner + Config *types.SandboxConfig + StreamReq *types.StreamRequest + Manager *infra.Manager +} + +// ExecuteSandboxStream is the V2 replacement for executeSandboxStream. +// It calls runner.Stream, handles interrupts, and performs cleanup/lifecycle +// in defer. +func ExecuteSandboxStream( + ctx *agentContext.Context, + req *ExecuteRequest, + handler message.StreamFunc, +) (*agentContext.CompletionResponse, error) { + + if req.Runner == nil || req.Computer == nil { + return nil, fmt.Errorf("runner and computer are required") + } + + stdCtx := ctx.Context + panicked := true // Assume panic; set false on normal exit. + + // Resolve stop timeout from config (default 2s). + stopTimeout := 2 * time.Second + if req.Config != nil && req.Config.StopTimeout != "" { + if d, err := time.ParseDuration(req.Config.StopTimeout); err == nil { + stopTimeout = d + } + } + + // Panic recovery (registered first, executes last in LIFO order). + defer func() { + if r := recover(); r != nil { + log.Printf("[sandbox/v2] panic in stream: %v", r) + cleanCtx, cancel := context.WithTimeout(context.Background(), stopTimeout) + defer cancel() + req.Runner.Cleanup(cleanCtx, req.Computer) + LifecycleAction(cleanCtx, req.Config, req.Computer, req.Manager) + } + }() + + // Lifecycle action (registered second, executes second-to-last). + defer func() { + if !panicked { + LifecycleAction(stdCtx, req.Config, req.Computer, req.Manager) + } + }() + + // Runner cleanup (registered last, executes first). + defer func() { + if !panicked { + cleanCtx, cancel := context.WithTimeout(context.Background(), stopTimeout) + defer cancel() + req.Runner.Cleanup(cleanCtx, req.Computer) + } + }() + + // Build a cancellable runnerCtx that bridges agentContext interrupts. + runnerCtx, cancelRunner := context.WithCancel(stdCtx) + defer cancelRunner() // Prevent goroutine leak. + + done := make(chan struct{}) + defer close(done) + + go func() { + ticker := time.NewTicker(500 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-done: + return + case <-ticker.C: + if ctx.Interrupt != nil { + if sig := ctx.Interrupt.Peek(); sig != nil { + cancelRunner() + return + } + if ctx.Interrupt.IsInterrupted() { + cancelRunner() + return + } + } + case <-stdCtx.Done(): + cancelRunner() + return + } + } + }() + + var textContent []byte + wrappedHandler := func(chunkType message.StreamChunkType, data []byte) int { + if chunkType == message.ChunkText { + textContent = append(textContent, data...) + } + if handler != nil { + return handler(chunkType, data) + } + return 0 + } + + err := req.Runner.Stream(runnerCtx, req.StreamReq, wrappedHandler) + + panicked = false // Normal exit reached. + + if err != nil { + if errors.Is(err, context.Canceled) { + return nil, err + } + return nil, fmt.Errorf("runner.Stream: %w", err) + } + + resp := &agentContext.CompletionResponse{ + Role: "assistant", + FinishReason: agentContext.FinishReasonStop, + } + if len(textContent) > 0 { + resp.Content = string(textContent) + } + return resp, nil +} diff --git a/agent/sandbox/v2/testutils/testutils.go b/agent/sandbox/v2/testutils/testutils.go new file mode 100644 index 00000000..7e51a105 --- /dev/null +++ b/agent/sandbox/v2/testutils/testutils.go @@ -0,0 +1,47 @@ +package testutils + +import ( + "context" + "os" + "path/filepath" + "testing" + + agenttestutils "github.com/yaoapp/yao/agent/testutils" + "github.com/yaoapp/yao/config" + sandboxv2 "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" +) + +// Prepare initializes the full environment required for sandbox V2 E2E tests: +// - agent layer (assistants, LLM, caller) +// - tai registry + local node +// - sandbox V2 manager +func Prepare(t *testing.T) { + t.Helper() + + agenttestutils.Prepare(t) + + if registry.Global() == nil { + registry.Init(nil) + } + + dataDir := filepath.Join(config.Conf.DataRoot, "workspaces") + os.MkdirAll(dataDir, 0755) + tai.RegisterLocal(tai.WithDataDir(dataDir)) + + sandboxv2.Init() + if err := sandboxv2.M().Start(context.Background()); err != nil { + t.Fatalf("sandbox v2 manager start: %v", err) + } + + t.Cleanup(func() { + sandboxv2.M().Close() + }) +} + +// Clean tears down the test environment. +func Clean(t *testing.T) { + t.Helper() + agenttestutils.Clean(t) +} diff --git a/agent/sandbox/v2/testutils_test.go b/agent/sandbox/v2/testutils_test.go new file mode 100644 index 00000000..1ee2c02b --- /dev/null +++ b/agent/sandbox/v2/testutils_test.go @@ -0,0 +1,228 @@ +package sandboxv2_test + +import ( + "context" + "fmt" + "log" + "os" + "strconv" + "strings" + "testing" + "time" + + sandbox "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" + taisandbox "github.com/yaoapp/yao/tai/sandbox" + "github.com/yaoapp/yao/workspace" +) + +// --------------------------------------------------------------------------- +// node configuration — mirrors sandbox/v2 testutils but scoped to prepare tests +// --------------------------------------------------------------------------- + +type nodeConfig struct { + Name string + Addr string + TaiID string + Options []tai.Option +} + +type hostTarget struct { + Name string + Addr string + TaiID string +} + +// --------------------------------------------------------------------------- +// environment helpers (same conventions as sandbox/v2 + env.local.sh) +// --------------------------------------------------------------------------- + +func testLocalAddr() string { + if addr := os.Getenv("SANDBOX_TEST_LOCAL_ADDR"); addr != "" { + return addr + } + return "local" +} + +func testImage() string { + if img := os.Getenv("SANDBOX_TEST_IMAGE"); img != "" { + return img + } + return "alpine:latest" +} + +func envPort(key string, fallback int) int { + if v := os.Getenv(key); v != "" { + if p, err := strconv.Atoi(v); err == nil { + return p + } + } + return fallback +} + +// --------------------------------------------------------------------------- +// node discovery +// --------------------------------------------------------------------------- + +func boxNodes() []nodeConfig { + nodes := []nodeConfig{ + {Name: "local", Addr: testLocalAddr()}, + } + if addr := os.Getenv("SANDBOX_TEST_REMOTE_ADDR"); addr != "" { + nodes = append(nodes, nodeConfig{Name: "remote", Addr: addr}) + } + return nodes +} + +func hostTargets() []hostTarget { + var targets []hostTarget + if addr := os.Getenv("TAI_TEST_WIN_HOSTEXEC_LINUX"); addr != "" { + targets = append(targets, hostTarget{Name: "win-linux", Addr: addr}) + } + if addr := os.Getenv("TAI_TEST_WIN_HOSTEXEC_NATIVE"); addr != "" { + targets = append(targets, hostTarget{Name: "win-native", Addr: addr}) + } + return targets +} + +// --------------------------------------------------------------------------- +// TestMain — purge stale containers from previous runs +// --------------------------------------------------------------------------- + +func TestMain(m *testing.M) { + purgeStale() + os.Exit(m.Run()) +} + +func purgeStale() { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + for _, nc := range boxNodes() { + client, err := tai.New(nc.Addr, nc.Options...) + if err != nil { + continue + } + sb := client.Sandbox() + if sb == nil { + client.Close() + continue + } + containers, _ := sb.List(ctx, taisandbox.ListOptions{All: true}) + for _, c := range containers { + id := c.Name + if id == "" { + id = c.ID + } + if strings.HasPrefix(id, "sb-prep-") || strings.HasPrefix(id, "sb-lc-") { + sb.Remove(ctx, id, true) + log.Printf("[purge] %s: removed %s", nc.Name, id) + } + } + client.Close() + } +} + +// --------------------------------------------------------------------------- +// Manager + Box helpers +// --------------------------------------------------------------------------- + +func setupManager(t *testing.T, nc *nodeConfig) *sandbox.Manager { + t.Helper() + if registry.Global() == nil { + registry.Init(nil) + } + client, err := tai.New(nc.Addr, nc.Options...) + if err != nil { + t.Fatalf("tai.New(%s): %v", nc.Addr, err) + } + nc.TaiID = client.TaiID() + + sandbox.Init() + m := sandbox.M() + t.Cleanup(func() { m.Close() }) + return m +} + +func createBox(t *testing.T, m *sandbox.Manager, nc nodeConfig) *sandbox.Box { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + defer cancel() + + if err := m.EnsureImage(ctx, nc.TaiID, testImage(), sandbox.ImagePullOptions{}); err != nil { + t.Fatalf("EnsureImage: %v", err) + } + + box, err := m.Create(ctx, sandbox.CreateOptions{ + ID: fmt.Sprintf("sb-prep-%d", time.Now().UnixNano()), + Image: testImage(), + Owner: "test-prepare", + NodeID: nc.TaiID, + }) + if err != nil { + t.Fatalf("Create: %v", err) + } + t.Cleanup(func() { + cCtx, cCancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cCancel() + if err := m.Remove(cCtx, box.ID()); err != nil { + t.Logf("cleanup Remove(%s): %v", box.ID(), err) + } + }) + return box +} + +func createHost(t *testing.T, m *sandbox.Manager, tgt hostTarget) *sandbox.Host { + t.Helper() + host, err := m.Host(context.Background(), tgt.TaiID) + if err != nil { + t.Skipf("Host(%s): %v", tgt.Name, err) + } + return host +} + +func setupHostManager(t *testing.T, tgt *hostTarget) *sandbox.Manager { + t.Helper() + nc := nodeConfig{Name: tgt.Name, Addr: fmt.Sprintf("tai://%s", tgt.Addr)} + m := setupManager(t, &nc) + tgt.TaiID = nc.TaiID + return m +} + +// --------------------------------------------------------------------------- +// skip helpers +// --------------------------------------------------------------------------- + +func skipIfNoDocker(t *testing.T) { + t.Helper() + if testLocalAddr() == "" { + t.Skip("SANDBOX_TEST_LOCAL_ADDR not set") + } +} + +func skipIfNoHostExec(t *testing.T) { + t.Helper() + if len(hostTargets()) == 0 { + t.Skip("no HostExec targets configured") + } +} + +func createTestWorkspace(t *testing.T, taiID, wsID string) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + _, err := workspace.M().Create(ctx, workspace.CreateOptions{ + ID: wsID, + Owner: "test", + Node: taiID, + }) + if err != nil && !strings.Contains(err.Error(), "exists") { + t.Fatalf("create workspace %q: %v", wsID, err) + } + t.Cleanup(func() { + cCtx, cCancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cCancel() + workspace.M().Delete(cCtx, wsID, true) + }) +} diff --git a/agent/sandbox/v2/types/config.go b/agent/sandbox/v2/types/config.go new file mode 100644 index 00000000..2f5a57af --- /dev/null +++ b/agent/sandbox/v2/types/config.go @@ -0,0 +1,137 @@ +package types + +import ( + "encoding/json" + "fmt" +) + +const ( + SandboxVersionV1 = "1.0" + SandboxVersionV2 = "2.0" +) + +// SandboxConfig is the V2 sandbox configuration loaded from sandbox.yao or +// the package.yao "sandbox" block when version == "2.0". +type SandboxConfig struct { + Version string `json:"version" yaml:"version"` + Computer ComputerConfig `json:"computer" yaml:"computer"` + Runner RunnerConfig `json:"runner" yaml:"runner"` + Lifecycle string `json:"lifecycle,omitempty" yaml:"lifecycle,omitempty"` + IdleTimeout string `json:"idle_timeout,omitempty" yaml:"idle_timeout,omitempty"` + MaxLifetime string `json:"max_lifetime,omitempty" yaml:"max_lifetime,omitempty"` + StopTimeout string `json:"stop_timeout,omitempty" yaml:"stop_timeout,omitempty"` + Prepare []PrepareStep `json:"prepare,omitempty" yaml:"prepare,omitempty"` + Environment map[string]string `json:"environment,omitempty" yaml:"environment,omitempty"` + Secrets map[string]string `json:"secrets,omitempty" yaml:"secrets,omitempty"` + + // Populated by the framework at runtime (never serialized). + Owner string `json:"-" yaml:"-"` + ID string `json:"-" yaml:"-"` + Labels map[string]string `json:"-" yaml:"-"` + NodeID string `json:"-" yaml:"-"` + Kind string `json:"-" yaml:"-"` + WorkspaceID string `json:"-" yaml:"-"` +} + +// ComputerConfig describes the execution environment (container or host). +type ComputerConfig struct { + Image string `json:"image,omitempty" yaml:"image,omitempty"` + VNC VNCConfig `json:"vnc,omitempty" yaml:"vnc,omitempty"` + Memory string `json:"memory,omitempty" yaml:"memory,omitempty"` + CPUs float64 `json:"cpus,omitempty" yaml:"cpus,omitempty"` + Ports PortList `json:"ports,omitempty" yaml:"ports,omitempty"` + User string `json:"user,omitempty" yaml:"user,omitempty"` + WorkDir string `json:"work_dir,omitempty" yaml:"work_dir,omitempty"` + MountPath string `json:"mount_path,omitempty" yaml:"mount_path,omitempty"` + MountMode string `json:"mount_mode,omitempty" yaml:"mount_mode,omitempty"` +} + +// RunnerConfig identifies which Runner to use and how. +type RunnerConfig struct { + Name string `json:"name" yaml:"name"` + Mode string `json:"mode,omitempty" yaml:"mode,omitempty"` + Options map[string]any `json:"options,omitempty" yaml:"options,omitempty"` +} + +// PrepareStep is a single action executed during Runner.Prepare. +type PrepareStep struct { + Action string `json:"action" yaml:"action"` + Once bool `json:"once,omitempty" yaml:"once,omitempty"` + IgnoreError bool `json:"ignore_error,omitempty" yaml:"ignore_error,omitempty"` + + // action=copy + Src string `json:"src,omitempty" yaml:"src,omitempty"` + Dst string `json:"dst,omitempty" yaml:"dst,omitempty"` + + // action=exec + Cmd string `json:"cmd,omitempty" yaml:"cmd,omitempty"` + Background bool `json:"background,omitempty" yaml:"background,omitempty"` + + // action=file (internal use by Runner.Prepare) + Path string `json:"path,omitempty" yaml:"path,omitempty"` + Content []byte `json:"-" yaml:"-"` + + // action=process (reserved) + Name string `json:"name,omitempty" yaml:"name,omitempty"` + Args []any `json:"args,omitempty" yaml:"args,omitempty"` +} + +// --------------------------------------------------------------------------- +// VNCConfig — supports both bool and object in JSON/YAML: +// true → VNCConfig{Enabled: true} +// {"enabled": true, "password": "xxx"} → full struct +// --------------------------------------------------------------------------- + +type VNCConfig struct { + Enabled bool `json:"enabled,omitempty" yaml:"enabled,omitempty"` + ViewOnly bool `json:"view_only,omitempty" yaml:"view_only,omitempty"` + Password string `json:"password,omitempty" yaml:"password,omitempty"` + Resolution string `json:"resolution,omitempty" yaml:"resolution,omitempty"` +} + +func (v *VNCConfig) UnmarshalJSON(data []byte) error { + var b bool + if err := json.Unmarshal(data, &b); err == nil { + v.Enabled = b + return nil + } + type alias VNCConfig + var a alias + if err := json.Unmarshal(data, &a); err != nil { + return err + } + *v = VNCConfig(a) + return nil +} + +// --------------------------------------------------------------------------- +// PortList — supports both int array and object array in JSON: +// [3000, 8080] → []PortMapping{{Port: 3000}, {Port: 8080}} +// [{"port": 3000, "host_port": 9000}] → full structs +// --------------------------------------------------------------------------- + +type PortList []PortMapping + +type PortMapping struct { + Port int `json:"port" yaml:"port"` + HostPort int `json:"host_port,omitempty" yaml:"host_port,omitempty"` + Protocol string `json:"protocol,omitempty" yaml:"protocol,omitempty"` +} + +func (p *PortList) UnmarshalJSON(data []byte) error { + var ints []int + if err := json.Unmarshal(data, &ints); err == nil { + out := make(PortList, len(ints)) + for i, port := range ints { + out[i] = PortMapping{Port: port} + } + *p = out + return nil + } + var objs []PortMapping + if err := json.Unmarshal(data, &objs); err != nil { + return fmt.Errorf("ports: expected int array or object array: %w", err) + } + *p = objs + return nil +} diff --git a/agent/sandbox/v2/types/runner.go b/agent/sandbox/v2/types/runner.go new file mode 100644 index 00000000..5a90f80f --- /dev/null +++ b/agent/sandbox/v2/types/runner.go @@ -0,0 +1,53 @@ +package types + +import ( + "context" + + "github.com/yaoapp/gou/connector" + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/output/message" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// Runner is the interface that all sandbox runners must implement. +// A Runner replaces the LLM invocation layer (executeLLMStream) when a +// sandbox is configured. +type Runner interface { + Name() string + Prepare(ctx context.Context, req *PrepareRequest) error + Stream(ctx context.Context, req *StreamRequest, handler message.StreamFunc) error + Cleanup(ctx context.Context, computer infra.Computer) error +} + +// MCPServer mirrors store/types.MCPServerConfig to avoid a cyclic import +// between this leaf package and agent/store/types. +type MCPServer struct { + ServerID string `json:"server_id,omitempty"` + Resources []string `json:"resources,omitempty"` + Tools []string `json:"tools,omitempty"` +} + +// RunStepsFunc is the signature of RunPrepareSteps. Workspace is obtained +// internally via computer.Workplace(). +type RunStepsFunc func(ctx context.Context, steps []PrepareStep, computer infra.Computer, assistantID, configHash string) error + +// PrepareRequest carries everything needed by Runner.Prepare. +type PrepareRequest struct { + Computer infra.Computer + Config *SandboxConfig + Connector connector.Connector + SkillsDir string + MCPServers []MCPServer + ConfigHash string + RunSteps RunStepsFunc +} + +// StreamRequest carries everything needed by Runner.Stream. +type StreamRequest struct { + Computer infra.Computer + Config *SandboxConfig + Connector connector.Connector + Messages []agentContext.Message + SystemPrompt string + ChatID string +} diff --git a/agent/sandbox/v2/types/token.go b/agent/sandbox/v2/types/token.go new file mode 100644 index 00000000..7f1cf525 --- /dev/null +++ b/agent/sandbox/v2/types/token.go @@ -0,0 +1,9 @@ +package types + +import "time" + +// SandboxToken is a short-lived JWT issued for a sandbox computer. +type SandboxToken struct { + Token string + ExpiresAt time.Time +} diff --git a/agent/sandbox/v2/yao/runner.go b/agent/sandbox/v2/yao/runner.go new file mode 100644 index 00000000..7bbd7109 --- /dev/null +++ b/agent/sandbox/v2/yao/runner.go @@ -0,0 +1,38 @@ +package yao + +import ( + "context" + + "github.com/yaoapp/yao/agent/output/message" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// YaoRunner is a no-op Runner for pure Hook-driven sandbox interactions. +// When runner.name == "yao", the assistant relies entirely on Create/Next +// hooks for logic; no external CLI is invoked. +type YaoRunner struct{} + +func New() *YaoRunner { return &YaoRunner{} } + +func (r *YaoRunner) Name() string { return "yao" } + +// Prepare runs user-defined prepare steps (copy, exec, file) but adds +// no runner-specific steps. Connector is not required. +func (r *YaoRunner) Prepare(ctx context.Context, req *types.PrepareRequest) error { + if req.RunSteps != nil && len(req.Config.Prepare) > 0 { + return req.RunSteps(ctx, req.Config.Prepare, req.Computer, req.Config.ID, req.ConfigHash) + } + return nil +} + +// Stream is a no-op — hooks handle all interaction. Returns immediately +// so the assistant framework proceeds to the Next hook. +func (r *YaoRunner) Stream(_ context.Context, _ *types.StreamRequest, _ message.StreamFunc) error { + return nil +} + +// Cleanup is a no-op for the yao runner. +func (r *YaoRunner) Cleanup(_ context.Context, _ infra.Computer) error { + return nil +} diff --git a/agent/sandbox/v2/yao/runner_test.go b/agent/sandbox/v2/yao/runner_test.go new file mode 100644 index 00000000..2f7e4201 --- /dev/null +++ b/agent/sandbox/v2/yao/runner_test.go @@ -0,0 +1,137 @@ +package yao_test + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/caller" + agentcontext "github.com/yaoapp/yao/agent/context" + sandboxtestutils "github.com/yaoapp/yao/agent/sandbox/v2/testutils" + oauthtypes "github.com/yaoapp/yao/openapi/oauth/types" +) + +func TestSandboxV2_Yao_JSAPI(t *testing.T) { + sandboxtestutils.Prepare(t) + defer sandboxtestutils.Clean(t) + + require.NotNil(t, caller.AgentGetterFunc, "AgentGetterFunc should be registered after Prepare") + + agent, err := caller.AgentGetterFunc("tests.sandbox-v2.jsapi-v2") + require.NoError(t, err, "should load assistant tests.sandbox-v2.jsapi-v2") + + chatID := fmt.Sprintf("e2e-jsapi-%d", time.Now().UnixMilli()) + ctx := agentcontext.New( + context.Background(), + &oauthtypes.AuthorizedInfo{ + TeamID: "test-team-jsapi", + UserID: "test-user-jsapi", + }, + chatID, + ) + + messages := []agentcontext.Message{ + {Role: "user", Content: "test jsapi"}, + } + + done := make(chan struct{}) + var resp *agentcontext.Response + var streamErr error + + go func() { + defer close(done) + resp, streamErr = agent.Stream(ctx, messages) + }() + + select { + case <-done: + case <-time.After(3 * time.Minute): + t.Fatalf("timeout after 3m") + } + + require.NoError(t, streamErr, "Stream should not return error") + require.NotNil(t, resp, "response should not be nil") + + // runner=yao goes through executeLLMStream, then Next hook returns { data: results } + // The Next hook result should appear in resp.Next + require.NotNil(t, resp.Next, "resp.Next should not be nil (Next hook returned data)") + t.Logf("resp.Next: %+v", resp.Next) + + nextData, ok := resp.Next.(map[string]interface{}) + if !ok { + t.Fatalf("resp.Next should be a map, got %T: %+v", resp.Next, resp.Next) + } + + // The Next hook returns { data: results }, the framework unwraps .data + data, hasData := nextData["data"] + if hasData { + nextData, ok = data.(map[string]interface{}) + require.True(t, ok, "data should be a map") + } + + t.Logf("JSAPI test results: %+v", nextData) + + // ── Verify ctx.computer was available ── + assert.Equal(t, true, nextData["has_computer"], "ctx.computer should be available") + assert.Equal(t, true, nextData["has_workspace"], "ctx.workspace should be available") + + // ── Verify ctx.computer.Info() ── + if infoRaw, ok := nextData["computer_info"]; ok { + info, ok := infoRaw.(map[string]interface{}) + require.True(t, ok, "computer_info should be a map") + assert.NotEmpty(t, info["kind"], "computer_info.kind should not be empty") + t.Logf("computer info: kind=%v os=%v", info["kind"], info["os"]) + } else { + assert.Nil(t, nextData["computer_info_error"], "computer.Info() should not error") + } + + // ── Verify ctx.computer.Exec() ── + assert.Equal(t, "jsapi-v2-test", nextData["exec_stdout"], "Exec should return expected stdout") + assert.Nil(t, nextData["exec_error"], "Exec should not error") + if exitCode, ok := nextData["exec_exit_code"]; ok { + // JS numbers come back as float64 through JSON + switch v := exitCode.(type) { + case float64: + assert.Equal(t, float64(0), v, "exit_code should be 0") + case int: + assert.Equal(t, 0, v, "exit_code should be 0") + } + } + + // ── Verify ctx.workspace write/read ── + assert.Equal(t, true, nextData["write_read_ok"], "workspace WriteFile+ReadFile round-trip should work") + assert.Equal(t, "hello from jsapi v2", nextData["read_content"], "read content should match") + assert.Nil(t, nextData["write_read_error"], "write/read should not error") + + // ── Verify ctx.workspace MkdirAll + Exists ── + assert.Equal(t, true, nextData["mkdir_exists_ok"], "MkdirAll + Exists should work") + assert.Nil(t, nextData["mkdir_exists_error"], "mkdir/exists should not error") + + // ── Verify ctx.workspace ReadDir ── + assert.Nil(t, nextData["readdir_error"], "ReadDir should not error") + if count, ok := nextData["readdir_count"]; ok { + switch v := count.(type) { + case float64: + assert.Greater(t, v, float64(0), "ReadDir should return entries") + } + } + + // ── Verify ctx.workspace Stat ── + assert.Equal(t, true, nextData["stat_ok"], "Stat should return correct info") + assert.Nil(t, nextData["stat_error"], "Stat should not error") + + // ── Verify ctx.workspace Copy ── + assert.Equal(t, true, nextData["copy_ok"], "Copy should work") + assert.Nil(t, nextData["copy_error"], "Copy should not error") + + // ── Verify ctx.workspace Rename ── + assert.Equal(t, true, nextData["rename_ok"], "Rename should work") + assert.Nil(t, nextData["rename_error"], "Rename should not error") + + // ── Verify ctx.workspace Remove ── + assert.Equal(t, true, nextData["remove_ok"], "Remove should work") + assert.Nil(t, nextData["remove_error"], "Remove should not error") +} diff --git a/agent/store/types/sandbox_v2.go b/agent/store/types/sandbox_v2.go new file mode 100644 index 00000000..9968d613 --- /dev/null +++ b/agent/store/types/sandbox_v2.go @@ -0,0 +1,100 @@ +package types + +import ( + "crypto/sha256" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + + jsoniter "github.com/json-iterator/go" + sandboxTypes "github.com/yaoapp/yao/agent/sandbox/v2/types" +) + +// LoadSandboxConfig reads a sandbox.yao file (JSON or YAML) and returns +// the V2 SandboxConfig. Called during Assistant.Load(). +func LoadSandboxConfig(filePath string) (*sandboxTypes.SandboxConfig, error) { + data, err := os.ReadFile(filePath) + if err != nil { + return nil, fmt.Errorf("read sandbox config %s: %w", filePath, err) + } + + ext := strings.ToLower(filepath.Ext(filePath)) + var cfg sandboxTypes.SandboxConfig + + switch ext { + case ".json", ".yao": + if err := jsoniter.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("parse sandbox config (json): %w", err) + } + default: + if err := jsoniter.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("parse sandbox config: %w", err) + } + } + + if cfg.Version != sandboxTypes.SandboxVersionV2 { + return nil, fmt.Errorf("sandbox.yao version must be %q, got %q", sandboxTypes.SandboxVersionV2, cfg.Version) + } + + return &cfg, nil +} + +// ToSandboxV2 converts a generic value (typically map[string]any from DSL +// parsing) into a V2 SandboxConfig. +func ToSandboxV2(v any) (*sandboxTypes.SandboxConfig, error) { + if v == nil { + return nil, nil + } + + switch sb := v.(type) { + case *sandboxTypes.SandboxConfig: + return sb, nil + case sandboxTypes.SandboxConfig: + return &sb, nil + default: + raw, err := jsoniter.Marshal(v) + if err != nil { + return nil, fmt.Errorf("sandbox v2 format error: %w", err) + } + var cfg sandboxTypes.SandboxConfig + if err := jsoniter.Unmarshal(raw, &cfg); err != nil { + return nil, fmt.Errorf("sandbox v2 format error: %w", err) + } + return &cfg, nil + } +} + +// ComputeConfigHash computes a SHA-256 fingerprint of the sandbox configuration, +// MCP servers, and skills directory. Used for hot-reload detection in prepare +// step "once" logic. +func ComputeConfigHash(cfg *sandboxTypes.SandboxConfig, mcpServers []MCPServerConfig, skillsDir string) string { + h := sha256.New() + + raw, _ := json.Marshal(cfg) + h.Write(raw) + + if len(mcpServers) > 0 { + mcpRaw, _ := json.Marshal(mcpServers) + h.Write(mcpRaw) + } + + if skillsDir != "" { + h.Write([]byte(skillsDir)) + entries, err := os.ReadDir(skillsDir) + if err == nil { + names := make([]string, 0, len(entries)) + for _, e := range entries { + names = append(names, e.Name()) + } + sort.Strings(names) + for _, n := range names { + h.Write([]byte(n)) + } + } + } + + return fmt.Sprintf("%x", h.Sum(nil)) +} diff --git a/agent/store/types/types.go b/agent/store/types/types.go index ea4ae8c1..4063b285 100644 --- a/agent/store/types/types.go +++ b/agent/store/types/types.go @@ -9,6 +9,7 @@ import ( "github.com/yaoapp/xun/dbal/query" "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/i18n" + sandboxTypes "github.com/yaoapp/yao/agent/sandbox/v2/types" searchTypes "github.com/yaoapp/yao/agent/search/types" ) @@ -421,42 +422,44 @@ type ConnectorOptions struct { // AssistantModel the assistant database model type AssistantModel struct { - ID string `json:"assistant_id"` // Assistant ID - Type string `json:"type,omitempty"` // Assistant Type, default is assistant - Name string `json:"name,omitempty"` // Assistant Name - Avatar string `json:"avatar,omitempty"` // Assistant Avatar - Connector string `json:"connector"` // AI Connector (default connector) - ConnectorOptions *ConnectorOptions `json:"connector_options,omitempty"` // Connector selection options for user to choose from - Path string `json:"path,omitempty"` // Assistant Path - BuiltIn bool `json:"built_in,omitempty"` // Whether this is a built-in assistant - Sort int `json:"sort,omitempty"` // Assistant Sort - Description string `json:"description,omitempty"` // Assistant Description - Capabilities string `json:"capabilities,omitempty"` // Assistant capabilities description (useful for Robot orchestration) - Tags []string `json:"tags,omitempty"` // Assistant Tags - Modes []string `json:"modes,omitempty"` // Supported modes (e.g., ["task", "chat"]), null means all modes are supported - DefaultMode string `json:"default_mode,omitempty"` // Default mode, can be empty - Readonly bool `json:"readonly,omitempty"` // Whether this assistant is readonly - Public bool `json:"public,omitempty"` // Whether this assistant is shared across all teams in the platform - Share string `json:"share,omitempty"` // Assistant sharing scope (private/team) - Mentionable bool `json:"mentionable,omitempty"` // Whether this assistant is mentionable - Automated bool `json:"automated,omitempty"` // Whether this assistant is automated - Options map[string]interface{} `json:"options,omitempty"` // AI Options - Prompts []Prompt `json:"prompts,omitempty"` // AI Prompts (default prompts) - PromptPresets map[string][]Prompt `json:"prompt_presets,omitempty"` // Prompt presets organized by mode (e.g., "chat", "task", etc.) - DisableGlobalPrompts bool `json:"disable_global_prompts,omitempty"` // Whether to disable global prompts, default is false - KB *KnowledgeBase `json:"kb,omitempty"` // Knowledge base configuration - DB *Database `json:"db,omitempty"` // Database configuration - MCP *MCPServers `json:"mcp,omitempty"` // MCP servers configuration - Workflow *Workflow `json:"workflow,omitempty"` // Workflow configuration - Sandbox *Sandbox `json:"sandbox,omitempty"` // Sandbox configuration for coding agents - Placeholder *Placeholder `json:"placeholder,omitempty"` // Assistant Placeholder - Source string `json:"source,omitempty"` // Hook script source code - Locales i18n.Map `json:"locales,omitempty"` // Assistant Locales - Uses *context.Uses `json:"uses,omitempty"` // Assistant-specific wrapper configurations for vision, audio, etc. If not set, use global settings - Search *searchTypes.Config `json:"search,omitempty"` // Search configuration (web, kb, db, citation, weights, etc.) - Dependencies map[string]string `json:"dependencies,omitempty"` // Dependencies on other MCP Clients (name -> version constraint) - CreatedAt int64 `json:"created_at"` // Creation timestamp - UpdatedAt int64 `json:"updated_at"` // Last update timestamp + ID string `json:"assistant_id"` // Assistant ID + Type string `json:"type,omitempty"` // Assistant Type, default is assistant + Name string `json:"name,omitempty"` // Assistant Name + Avatar string `json:"avatar,omitempty"` // Assistant Avatar + Connector string `json:"connector"` // AI Connector (default connector) + ConnectorOptions *ConnectorOptions `json:"connector_options,omitempty"` // Connector selection options for user to choose from + Path string `json:"path,omitempty"` // Assistant Path + BuiltIn bool `json:"built_in,omitempty"` // Whether this is a built-in assistant + Sort int `json:"sort,omitempty"` // Assistant Sort + Description string `json:"description,omitempty"` // Assistant Description + Capabilities string `json:"capabilities,omitempty"` // Assistant capabilities description (useful for Robot orchestration) + Tags []string `json:"tags,omitempty"` // Assistant Tags + Modes []string `json:"modes,omitempty"` // Supported modes (e.g., ["task", "chat"]), null means all modes are supported + DefaultMode string `json:"default_mode,omitempty"` // Default mode, can be empty + Readonly bool `json:"readonly,omitempty"` // Whether this assistant is readonly + Public bool `json:"public,omitempty"` // Whether this assistant is shared across all teams in the platform + Share string `json:"share,omitempty"` // Assistant sharing scope (private/team) + Mentionable bool `json:"mentionable,omitempty"` // Whether this assistant is mentionable + Automated bool `json:"automated,omitempty"` // Whether this assistant is automated + Options map[string]interface{} `json:"options,omitempty"` // AI Options + Prompts []Prompt `json:"prompts,omitempty"` // AI Prompts (default prompts) + PromptPresets map[string][]Prompt `json:"prompt_presets,omitempty"` // Prompt presets organized by mode (e.g., "chat", "task", etc.) + DisableGlobalPrompts bool `json:"disable_global_prompts,omitempty"` // Whether to disable global prompts, default is false + KB *KnowledgeBase `json:"kb,omitempty"` // Knowledge base configuration + DB *Database `json:"db,omitempty"` // Database configuration + MCP *MCPServers `json:"mcp,omitempty"` // MCP servers configuration + Workflow *Workflow `json:"workflow,omitempty"` // Workflow configuration + Sandbox *Sandbox `json:"sandbox,omitempty"` // Sandbox configuration for coding agents (V1) + SandboxV2 *sandboxTypes.SandboxConfig `json:"-"` // V2 sandbox configuration (runtime only, not persisted in DB) + ConfigHash string `json:"-"` // V2 sandbox config fingerprint for hot-reload + Placeholder *Placeholder `json:"placeholder,omitempty"` // Assistant Placeholder + Source string `json:"source,omitempty"` // Hook script source code + Locales i18n.Map `json:"locales,omitempty"` // Assistant Locales + Uses *context.Uses `json:"uses,omitempty"` // Assistant-specific wrapper configurations for vision, audio, etc. If not set, use global settings + Search *searchTypes.Config `json:"search,omitempty"` // Search configuration (web, kb, db, citation, weights, etc.) + Dependencies map[string]string `json:"dependencies,omitempty"` // Dependencies on other MCP Clients (name -> version constraint) + CreatedAt int64 `json:"created_at"` // Creation timestamp + UpdatedAt int64 `json:"updated_at"` // Last update timestamp // Permission management fields (not exposed in JSON API responses) YaoCreatedBy string `json:"-"` // User who created the assistant (not exposed in JSON) diff --git a/cmd/start.go b/cmd/start.go index 2c11373b..2f1c98e2 100644 --- a/cmd/start.go +++ b/cmd/start.go @@ -26,12 +26,14 @@ import ( "github.com/yaoapp/yao/engine" yaogrpc "github.com/yaoapp/yao/grpc" _ "github.com/yaoapp/yao/grpc/auth" + sandboxhandler "github.com/yaoapp/yao/grpc/sandbox" "github.com/yaoapp/yao/openapi" + sandbox "github.com/yaoapp/yao/sandbox/v2" ischedule "github.com/yaoapp/yao/schedule" "github.com/yaoapp/yao/service" "github.com/yaoapp/yao/setup" "github.com/yaoapp/yao/share" - tairegistry "github.com/yaoapp/yao/tai/registry" + itask "github.com/yaoapp/yao/task" ) @@ -176,10 +178,6 @@ var startCmd = &cobra.Command{ ischedule.Start() defer ischedule.Stop() - // Initialize the global Tai registry for tunnel and direct connections - // (must happen before HTTP/gRPC start so handlers can access it) - tairegistry.Init(nil) - // Pre-flight: detect port conflicts before attempting to start servers. if occupied, proc := portOccupied(config.Conf.Host, config.Conf.Port); occupied { fmt.Println(color.RedString(L("Fatal: HTTP port %d is already in use%s"), config.Conf.Port, proc)) @@ -194,6 +192,12 @@ var startCmd = &cobra.Command{ } } + // Wire gRPC heartbeat → sandbox Manager so container liveness is tracked. + yaogrpc.SetSandboxOnBeat(func(data *sandboxhandler.HeartbeatData) string { + sandbox.M().Heartbeat(data.SandboxID, true, int(data.RunningProcs)) + return "ok" + }) + // Start all servers (gRPC + HTTP) as a single unit. // Start() blocks until HTTP port is bound (READY) or returns error. svc, err := service.Start(config.Conf, service.ServerHooks{ diff --git a/engine/load.go b/engine/load.go index f96f2090..ac02803b 100644 --- a/engine/load.go +++ b/engine/load.go @@ -40,12 +40,14 @@ import ( "github.com/yaoapp/yao/plugin" "github.com/yaoapp/yao/query" "github.com/yaoapp/yao/runtime" + sandbox "github.com/yaoapp/yao/sandbox/v2" "github.com/yaoapp/yao/schedule" "github.com/yaoapp/yao/script" "github.com/yaoapp/yao/share" "github.com/yaoapp/yao/socket" "github.com/yaoapp/yao/store" sui "github.com/yaoapp/yao/sui/api" + tairegistry "github.com/yaoapp/yao/tai/registry" "github.com/yaoapp/yao/task" "github.com/yaoapp/yao/websocket" "github.com/yaoapp/yao/widget" @@ -130,6 +132,22 @@ func Load(cfg config.Config, options LoadOption, progressCallback ...func(string warnings = append(warnings, Warning{Widget: "DB", Error: err}) } + // Initialize the Tai node registry (idempotent, safe to call early). + loadStep("Registry", func() error { + tairegistry.InitWithWriter(config.LogOutput, cfg.LogMode) + return nil + }, callback) + + // Initialize the Sandbox manager and start it (auto-registers local Docker + // node if available, recovers existing containers, starts cleanup loop). + err = loadStep("Sandbox", func() error { + sandbox.Init() + return sandbox.M().Start(context.Background()) + }, callback) + if err != nil { + warnings = append(warnings, Warning{Widget: "Sandbox", Error: err}) + } + // Load Certs err = loadStep("Cert", func() error { return cert.Load(cfg) diff --git a/openapi/nodes/nodes.go b/openapi/nodes/nodes.go new file mode 100644 index 00000000..7d548ae8 --- /dev/null +++ b/openapi/nodes/nodes.go @@ -0,0 +1,138 @@ +package nodes + +import ( + "net/http" + "time" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/oauth/authorized" + "github.com/yaoapp/yao/openapi/oauth/types" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/tai/registry" +) + +// Attach registers Tai node endpoints on the given group. +// - GET / — list nodes (filtered by team/user from token) +// - GET /:id — get single node (owner check) +func Attach(group *gin.RouterGroup, oauth types.OAuth) { + group.Use(oauth.Guard) + group.GET("", handleList) + group.GET("/:id", handleGet) +} + +type nodeResponse struct { + TaiID string `json:"tai_id"` + MachineID string `json:"machine_id,omitempty"` + Version string `json:"version,omitempty"` + DisplayName string `json:"display_name,omitempty"` + Mode string `json:"mode"` + Addr string `json:"addr,omitempty"` + Status string `json:"status"` + System systemResponse `json:"system"` + Capabilities map[string]bool `json:"capabilities,omitempty"` + Ports map[string]int `json:"ports,omitempty"` + ConnectedAt *time.Time `json:"connected_at,omitempty"` + LastPing *time.Time `json:"last_ping,omitempty"` +} + +type systemResponse struct { + OS string `json:"os"` + Arch string `json:"arch"` + Hostname string `json:"hostname"` + NumCPU int `json:"num_cpu"` + TotalMem int64 `json:"total_mem,omitempty"` + Shell string `json:"shell,omitempty"` +} + +func snapToResponse(s registry.NodeSnapshot) nodeResponse { + r := nodeResponse{ + TaiID: s.TaiID, + MachineID: s.MachineID, + Version: s.Version, + DisplayName: s.DisplayName, + Mode: s.Mode, + Addr: s.Addr, + Status: s.Status, + Capabilities: s.Capabilities, + Ports: s.Ports, + System: systemResponse{ + OS: s.System.OS, + Arch: s.System.Arch, + Hostname: s.System.Hostname, + NumCPU: s.System.NumCPU, + TotalMem: s.System.TotalMem, + Shell: s.System.Shell, + }, + } + if !s.ConnectedAt.IsZero() { + r.ConnectedAt = &s.ConnectedAt + } + if !s.LastPing.IsZero() { + r.LastPing = &s.LastPing + } + return r +} + +// nodeOwnedBy checks whether a node belongs to the caller. +// TeamID match → true; no team and UserID match → true. +func nodeOwnedBy(snap *registry.NodeSnapshot, authInfo *types.AuthorizedInfo) bool { + if authInfo == nil { + return true + } + + if authInfo.TeamID != "" { + return snap.Auth.TeamID == authInfo.TeamID + } + if authInfo.UserID != "" { + return snap.Auth.TeamID == "" && snap.Auth.UserID == authInfo.UserID + } + return true +} + +func handleList(c *gin.Context) { + reg := registry.Global() + if reg == nil { + response.RespondWithSuccess(c, http.StatusOK, []nodeResponse{}) + return + } + + authInfo := authorized.GetInfo(c) + + var snaps []registry.NodeSnapshot + if authInfo != nil && authInfo.TeamID != "" { + snaps = reg.ListByTeam(authInfo.TeamID) + } else if authInfo != nil && authInfo.UserID != "" { + snaps = reg.ListByUser(authInfo.UserID) + } else { + snaps = reg.List() + } + + result := make([]nodeResponse, 0, len(snaps)) + for _, s := range snaps { + result = append(result, snapToResponse(s)) + } + response.RespondWithSuccess(c, http.StatusOK, result) +} + +func handleGet(c *gin.Context) { + reg := registry.Global() + if reg == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "node registry not available"}) + return + } + + id := c.Param("id") + snap, ok := reg.Get(id) + if !ok { + c.JSON(http.StatusNotFound, gin.H{"error": "node not found"}) + return + } + + authInfo := authorized.GetInfo(c) + if !nodeOwnedBy(snap, authInfo) { + c.JSON(http.StatusForbidden, gin.H{"error": "no permission to access this node"}) + return + } + + response.RespondWithSuccess(c, http.StatusOK, snapToResponse(*snap)) +} diff --git a/openapi/oauth.go b/openapi/oauth.go index 52f19b1e..94c23cd5 100644 --- a/openapi/oauth.go +++ b/openapi/oauth.go @@ -2,6 +2,7 @@ package openapi import ( "encoding/base64" + "fmt" "net/http" "strings" @@ -585,11 +586,32 @@ func (openapi *OpenAPI) oauthDeviceAuthorize(c *gin.Context) { if extraClaims == nil { extraClaims = make(map[string]interface{}) } - if tokenClaims.TeamID != "" { - extraClaims["team_id"] = tokenClaims.TeamID + + teamID := tokenClaims.TeamID + if teamID == "" { + switch v := extraClaims["team_id"].(type) { + case string: + teamID = v + case float64: + teamID = fmt.Sprintf("%.0f", v) + } } - if tokenClaims.TenantID != "" { - extraClaims["tenant_id"] = tokenClaims.TenantID + if teamID != "" { + extraClaims["team_id"] = teamID + } + + tenantID := tokenClaims.TenantID + if tenantID == "" { + if v, ok := extraClaims["tenant_id"].(string); ok { + tenantID = v + } + } + if tenantID != "" { + extraClaims["tenant_id"] = tenantID + } + + if tokenClaims.ClientID != "" { + extraClaims["authorizer_client_id"] = tokenClaims.ClientID } userCode := c.PostForm("user_code") diff --git a/openapi/oauth/authenticate.go b/openapi/oauth/authenticate.go index 288ff2ad..16937e77 100644 --- a/openapi/oauth/authenticate.go +++ b/openapi/oauth/authenticate.go @@ -189,13 +189,29 @@ func (s *Service) refreshTokenDirect(refreshToken string, expiredClaims *types.T // buildAuthInfo constructs AuthorizedInfo directly from token claims, // equivalent to the SetInfo+GetInfo round-trip through gin.Context. func (s *Service) buildAuthInfo(claims *types.TokenClaims, sessionID string) *types.AuthorizedInfo { + teamID := claims.TeamID + tenantID := claims.TenantID + + if claims.Extra != nil { + if teamID == "" { + if v, ok := claims.Extra["team_id"].(string); ok && v != "" { + teamID = v + } + } + if tenantID == "" { + if v, ok := claims.Extra["tenant_id"].(string); ok && v != "" { + tenantID = v + } + } + } + info := &types.AuthorizedInfo{ Subject: claims.Subject, ClientID: claims.ClientID, Scope: claims.Scope, SessionID: sessionID, - TeamID: claims.TeamID, - TenantID: claims.TenantID, + TeamID: teamID, + TenantID: tenantID, } userID, err := s.UserID(claims.ClientID, claims.Subject) @@ -203,5 +219,14 @@ func (s *Service) buildAuthInfo(claims *types.TokenClaims, sessionID string) *ty info.UserID = userID } + if info.UserID == "" && claims.Extra != nil { + if authorizerClientID, ok := claims.Extra["authorizer_client_id"].(string); ok && authorizerClientID != "" { + if uid, err := s.UserID(authorizerClientID, claims.Subject); err == nil && uid != "" { + info.UserID = uid + s.copyFingerprint(authorizerClientID, claims.ClientID, claims.Subject) + } + } + } + return info } diff --git a/openapi/oauth/core.go b/openapi/oauth/core.go index 9862a5ee..f6e4fc34 100644 --- a/openapi/oauth/core.go +++ b/openapi/oauth/core.go @@ -241,9 +241,11 @@ func (s *Service) RefreshToken(ctx context.Context, refreshToken string, scope . finalScope = requestedScope } + extraClaims := extractExtraClaims(tokenInfo) + // Generate new access token with final scope expiresIn := int(s.config.Token.AccessTokenLifetime.Seconds()) - newAccessToken, err := s.generateAccessTokenWithScope(clientID, finalScope, originalSubject, expiresIn, nil) + newAccessToken, err := s.generateAccessTokenWithScope(clientID, finalScope, originalSubject, expiresIn, extraClaims) if err != nil { return nil, &types.ErrorResponse{ Code: types.ErrorServerError, @@ -340,9 +342,11 @@ func (s *Service) RotateRefreshToken(ctx context.Context, oldToken string, reque finalScope = scope } + extraClaims := extractExtraClaims(tokenInfo) + // Generate new tokens with final scope and original subject expiresIn := int(s.config.Token.AccessTokenLifetime.Seconds()) - newAccessToken, err := s.generateAccessTokenWithScope(clientID, finalScope, originalSubject, expiresIn, nil) + newAccessToken, err := s.generateAccessTokenWithScope(clientID, finalScope, originalSubject, expiresIn, extraClaims) if err != nil { return nil, &types.ErrorResponse{ Code: types.ErrorServerError, @@ -350,7 +354,7 @@ func (s *Service) RotateRefreshToken(ctx context.Context, oldToken string, reque } } - newRefreshToken, err := s.generateRefreshToken(clientID, finalScope, originalSubject, 0, nil) + newRefreshToken, err := s.generateRefreshToken(clientID, finalScope, originalSubject, 0, extraClaims) if err != nil { return nil, &types.ErrorResponse{ Code: types.ErrorServerError, @@ -497,20 +501,12 @@ func (s *Service) handleRefreshTokenGrant(ctx context.Context, client *types.Cli return nil, err } - // Extract scope and subject from refresh token if available - scope := "" - if scopeVal, ok := refreshTokenInfo["scope"].(string); ok { - scope = scopeVal - } + scope, _ := refreshTokenInfo["scope"].(string) + subject, _ := refreshTokenInfo["subject"].(string) + extraClaims := extractExtraClaims(refreshTokenInfo) - subject := "" - if subjectVal, ok := refreshTokenInfo["subject"].(string); ok { - subject = subjectVal - } - - // Generate and store new access token with proper scope and subject expiresIn := int(s.config.Token.AccessTokenLifetime.Seconds()) - accessToken, err := s.generateAccessTokenWithScope(client.ClientID, scope, subject, expiresIn, nil) + accessToken, err := s.generateAccessTokenWithScope(client.ClientID, scope, subject, expiresIn, extraClaims) if err != nil { return nil, &types.ErrorResponse{ Code: types.ErrorServerError, @@ -524,9 +520,8 @@ func (s *Service) handleRefreshTokenGrant(ctx context.Context, client *types.Cli ExpiresIn: expiresIn, } - // Include refresh token if rotation is enabled if s.config.Features.RefreshTokenRotationEnabled { - newRefreshToken, err := s.generateRefreshToken(client.ClientID, scope, subject, 0, nil) + newRefreshToken, err := s.generateRefreshToken(client.ClientID, scope, subject, 0, extraClaims) if err != nil { return nil, &types.ErrorResponse{ Code: types.ErrorServerError, @@ -534,11 +529,8 @@ func (s *Service) handleRefreshTokenGrant(ctx context.Context, client *types.Cli } } token.RefreshToken = newRefreshToken - - // Revoke old refresh token s.revokeRefreshToken(refreshToken) } else { - // Reuse the same refresh token token.RefreshToken = refreshToken } @@ -713,3 +705,23 @@ func (s *Service) handleDeviceCodeGrant(ctx context.Context, client *types.Clien } } } + +// extractExtraClaims pulls non-reserved fields from a token info map so they +// can be propagated into newly generated access/refresh tokens. +func extractExtraClaims(tokenInfo map[string]interface{}) map[string]interface{} { + reserved := map[string]bool{ + "client_id": true, "scope": true, "subject": true, + "type": true, "issued_at": true, "expires_at": true, + } + var extra map[string]interface{} + for k, v := range tokenInfo { + if reserved[k] { + continue + } + if extra == nil { + extra = make(map[string]interface{}) + } + extra[k] = v + } + return extra +} diff --git a/openapi/oauth/token.go b/openapi/oauth/token.go index 5d8b52d3..bb4f3428 100644 --- a/openapi/oauth/token.go +++ b/openapi/oauth/token.go @@ -320,6 +320,22 @@ func (s *Service) UserID(clientID, subject string) (string, error) { return userIDStr, nil } +// copyFingerprint copies the subject→userID fingerprint mapping from one +// clientID to another so that tokens issued under a different clientID +// (e.g. Device Flow) can resolve the same userID. +func (s *Service) copyFingerprint(srcClientID, dstClientID, subject string) { + srcKey := s.userFingerprintKey(srcClientID, subject) + userID, exists := s.store.Get(srcKey) + if !exists { + return + } + dstKey := s.userFingerprintKey(dstClientID, subject) + if _, already := s.store.Get(dstKey); already { + return + } + s.store.Set(dstKey, userID, 0) +} + // MakeAuthorizationCode generates a new authorization code with specific parameters and stores it // ============================================================================ diff --git a/openapi/openapi.go b/openapi/openapi.go index 23e9c108..dd5c6c9d 100644 --- a/openapi/openapi.go +++ b/openapi/openapi.go @@ -19,6 +19,7 @@ import ( "github.com/yaoapp/yao/openapi/llm" "github.com/yaoapp/yao/openapi/mcp" "github.com/yaoapp/yao/openapi/messenger" + "github.com/yaoapp/yao/openapi/nodes" "github.com/yaoapp/yao/openapi/oauth" "github.com/yaoapp/yao/openapi/oauth/acl" "github.com/yaoapp/yao/openapi/oauth/types" @@ -28,6 +29,7 @@ import ( "github.com/yaoapp/yao/openapi/team" openapiTrace "github.com/yaoapp/yao/openapi/trace" "github.com/yaoapp/yao/openapi/user" + openapiWorkspace "github.com/yaoapp/yao/openapi/workspace" taiapi "github.com/yaoapp/yao/tai/api" taitunnel "github.com/yaoapp/yao/tai/tunnel" ) @@ -173,9 +175,17 @@ func (openapi *OpenAPI) Attach(router *gin.Engine) { // OTP handlers (passwordless authentication) otp.Attach(group.Group("/otp"), openapi.OAuth) - // Sandbox handlers (VNC proxy for visual browser automation) + // Sandbox handlers (VNC proxy + management CRUD) sandbox.SetPathPrefix(baseURL) - sandbox.Attach(group.Group("/sandbox"), openapi.OAuth) + sandboxGroup := group.Group("/sandbox") + sandbox.Attach(sandboxGroup, openapi.OAuth) + sandbox.AttachManage(sandboxGroup) + + // Workspace handlers + openapiWorkspace.Attach(group.Group("/workspace"), openapi.OAuth) + + // Tai nodes handlers + nodes.Attach(group.Group("/nodes"), openapi.OAuth) // Tai tunnel WebSocket and reverse proxy routes group.GET("/ws/tai", taitunnel.HandleControl) diff --git a/openapi/sandbox/manage.go b/openapi/sandbox/manage.go new file mode 100644 index 00000000..2e0a4d02 --- /dev/null +++ b/openapi/sandbox/manage.go @@ -0,0 +1,488 @@ +package sandbox + +import ( + "context" + "net/http" + "sort" + "time" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/oauth/authorized" + "github.com/yaoapp/yao/openapi/oauth/types" + "github.com/yaoapp/yao/openapi/response" + sandboxv2 "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" +) + +// AttachManage registers sandbox management CRUD routes on the given group. +// oauth.Guard is already applied by the parent Attach call on the same group. +// - GET / — list sandboxes (filtered by owner) +// - POST / — create sandbox (owner from token) +// - GET /:id — get sandbox (owner check) +// - DELETE /:id — remove sandbox (owner check) +// - POST /:id/exec — execute command (owner check) +// - POST /:id/heartbeat — heartbeat (owner check) +func AttachManage(group *gin.RouterGroup) { + group.GET("", handleList) + group.POST("", handleCreate) + group.GET("/:id", handleGet) + group.DELETE("/:id", handleRemove) + group.POST("/:id/exec", handleExec) + group.POST("/:id/heartbeat", handleHeartbeat) +} + +// resolveOwner returns TeamID if present, otherwise UserID. +func resolveOwner(authInfo *types.AuthorizedInfo) string { + if authInfo != nil && authInfo.TeamID != "" { + return authInfo.TeamID + } + if authInfo != nil { + return authInfo.UserID + } + return "" +} + +// --- request / response types --- + +type createSandboxRequest struct { + ID string `json:"id,omitempty"` + NodeID string `json:"node_id"` + Image string `json:"image"` + WorkDir string `json:"work_dir,omitempty"` + User string `json:"user,omitempty"` + Env map[string]string `json:"env,omitempty"` + Memory int64 `json:"memory,omitempty"` + CPUs float64 `json:"cpus,omitempty"` + VNC bool `json:"vnc,omitempty"` + Policy string `json:"policy,omitempty"` + Labels map[string]string `json:"labels,omitempty"` + WorkspaceID string `json:"workspace_id,omitempty"` + MountMode string `json:"mount_mode,omitempty"` + MountPath string `json:"mount_path,omitempty"` +} + +type execRequest struct { + Cmd []string `json:"cmd" binding:"required"` + WorkDir string `json:"work_dir,omitempty"` + Env map[string]string `json:"env,omitempty"` + Timeout int `json:"timeout,omitempty"` +} + +type heartbeatRequest struct { + Active bool `json:"active"` + ProcessCount int `json:"process_count"` +} + +type sandboxSystemInfo struct { + OS string `json:"os"` + Arch string `json:"arch"` + Hostname string `json:"hostname"` + NumCPU int `json:"num_cpu"` + TotalMem int64 `json:"total_mem,omitempty"` + Shell string `json:"shell,omitempty"` + TempDir string `json:"temp_dir,omitempty"` +} + +type sandboxResponse struct { + Kind string `json:"kind"` + ID string `json:"id"` + DisplayName string `json:"display_name"` + ContainerID string `json:"container_id,omitempty"` + NodeID string `json:"node_id"` + Owner string `json:"owner"` + Status string `json:"status"` + Policy string `json:"policy,omitempty"` + Labels map[string]string `json:"labels,omitempty"` + Image string `json:"image,omitempty"` + Mode string `json:"mode,omitempty"` + Addr string `json:"addr,omitempty"` + VNC bool `json:"vnc"` + CreatedAt time.Time `json:"created_at"` + LastActive time.Time `json:"last_active"` + ProcessCount int `json:"process_count"` + System sandboxSystemInfo `json:"system"` + WorkspaceID string `json:"workspace_id,omitempty"` +} + +func boxToResponse(b *sandboxv2.Box) sandboxResponse { + snap := b.Snapshot() + info := b.ComputerInfo() + + displayName := info.System.Hostname + if displayName == "" { + displayName = snap.ID + } + + var mode, addr string + if ns, ok := tai.GetNodeSnapshot(snap.NodeID); ok { + mode = ns.Mode + addr = ns.Addr + } + if addr == "" && snap.NodeID != "" { + scheme := mode + if scheme == "" { + scheme = "local" + } + addr = scheme + "://" + snap.NodeID + } + + return sandboxResponse{ + Kind: "box", + ID: snap.ID, + DisplayName: displayName, + ContainerID: snap.ContainerID, + NodeID: snap.NodeID, + Owner: snap.Owner, + Status: snap.Status, + Policy: string(snap.Policy), + Labels: snap.Labels, + Image: snap.Image, + Mode: mode, + Addr: addr, + VNC: snap.VNC, + CreatedAt: snap.CreatedAt, + LastActive: snap.LastActive, + ProcessCount: snap.ProcessCount, + WorkspaceID: b.WorkspaceID(), + System: sandboxSystemInfo{ + OS: info.System.OS, + Arch: info.System.Arch, + Hostname: info.System.Hostname, + NumCPU: info.System.NumCPU, + TotalMem: info.System.TotalMem, + Shell: info.System.Shell, + TempDir: info.System.TempDir, + }, + } +} + +func hostToResponse(s registry.NodeSnapshot) sandboxResponse { + displayName := s.DisplayName + if displayName == "" { + displayName = s.System.Hostname + } + if displayName == "" { + displayName = s.TaiID + } + + status := "stopped" + if s.Status == "online" { + status = "running" + } + + owner := s.Auth.TeamID + if owner == "" { + owner = s.Auth.UserID + } + + addr := s.Addr + if addr == "" { + scheme := s.Mode + if scheme == "" { + scheme = "tai" + } + addr = scheme + "://" + s.TaiID + } + + return sandboxResponse{ + Kind: "host", + ID: s.TaiID, + DisplayName: displayName, + NodeID: s.TaiID, + Owner: owner, + Status: status, + Policy: "persistent", + Mode: s.Mode, + Addr: addr, + VNC: false, + CreatedAt: s.ConnectedAt, + LastActive: s.LastPing, + System: sandboxSystemInfo{ + OS: s.System.OS, + Arch: s.System.Arch, + Hostname: s.System.Hostname, + NumCPU: s.System.NumCPU, + TotalMem: s.System.TotalMem, + Shell: s.System.Shell, + }, + } +} + +func nodeOwnedBy(snap *registry.NodeSnapshot, authInfo *types.AuthorizedInfo) bool { + if authInfo == nil { + return true + } + if authInfo.TeamID != "" { + return snap.Auth.TeamID == authInfo.TeamID + } + if authInfo.UserID != "" { + return snap.Auth.TeamID == "" && snap.Auth.UserID == authInfo.UserID + } + return true +} + +func getManager(c *gin.Context) *sandboxv2.Manager { + defer func() { recover() }() + return sandboxv2.M() +} + +// checkBoxOwner verifies the caller owns the sandbox. +func checkBoxOwner(c *gin.Context, box *sandboxv2.Box, owner string) bool { + if owner == "" { + return true + } + info := box.ComputerInfo() + if info.Owner != "" && info.Owner != owner { + c.JSON(http.StatusForbidden, gin.H{"error": "no permission to access this sandbox"}) + return false + } + return true +} + +// --- handlers --- + +func handleList(c *gin.Context) { + authInfo := authorized.GetInfo(c) + owner := resolveOwner(authInfo) + nodeFilter := c.Query("node_id") + + var result []sandboxResponse + + // Host entries: list all nodes, filter by ownership + host_exec + if reg := registry.Global(); reg != nil { + snaps := reg.List() + for i := range snaps { + s := &snaps[i] + if !nodeOwnedBy(s, authInfo) { + continue + } + if !s.Capabilities["host_exec"] { + continue + } + if nodeFilter != "" && s.TaiID != nodeFilter { + continue + } + result = append(result, hostToResponse(*s)) + } + } + + // Box entries: list all, then filter by owner + if mgr := getManager(c); mgr != nil { + boxes, err := mgr.List(context.Background(), sandboxv2.ListOptions{ + NodeID: nodeFilter, + }) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + for _, b := range boxes { + snap := b.Snapshot() + if snap.Owner != owner { + continue + } + result = append(result, boxToResponse(b)) + } + } + + sort.Slice(result, func(i, j int) bool { + return result[i].LastActive.After(result[j].LastActive) + }) + + if result == nil { + result = []sandboxResponse{} + } + response.RespondWithSuccess(c, http.StatusOK, result) +} + +func handleCreate(c *gin.Context) { + mgr := getManager(c) + if mgr == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "sandbox service not available"}) + return + } + + var req createSandboxRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + if req.Image == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "image is required"}) + return + } + + authInfo := authorized.GetInfo(c) + owner := resolveOwner(authInfo) + + opts := sandboxv2.CreateOptions{ + ID: req.ID, + Owner: owner, + NodeID: req.NodeID, + Image: req.Image, + WorkDir: req.WorkDir, + User: req.User, + Env: req.Env, + Memory: req.Memory, + CPUs: req.CPUs, + VNC: req.VNC, + Policy: sandboxv2.LifecyclePolicy(req.Policy), + Labels: req.Labels, + WorkspaceID: req.WorkspaceID, + MountMode: req.MountMode, + MountPath: req.MountPath, + } + + box, err := mgr.Create(context.Background(), opts) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + response.RespondWithSuccess(c, http.StatusCreated, boxToResponse(box)) +} + +func handleGet(c *gin.Context) { + mgr := getManager(c) + if mgr == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "sandbox service not available"}) + return + } + + id := c.Param("id") + box, err := mgr.Get(context.Background(), id) + if err != nil { + if err == sandboxv2.ErrNotFound { + c.JSON(http.StatusNotFound, gin.H{"error": "sandbox not found"}) + return + } + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + authInfo := authorized.GetInfo(c) + if !checkBoxOwner(c, box, resolveOwner(authInfo)) { + return + } + + response.RespondWithSuccess(c, http.StatusOK, boxToResponse(box)) +} + +func handleRemove(c *gin.Context) { + mgr := getManager(c) + if mgr == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "sandbox service not available"}) + return + } + + id := c.Param("id") + box, err := mgr.Get(context.Background(), id) + if err != nil { + if err == sandboxv2.ErrNotFound { + c.JSON(http.StatusNotFound, gin.H{"error": "sandbox not found"}) + return + } + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + authInfo := authorized.GetInfo(c) + if !checkBoxOwner(c, box, resolveOwner(authInfo)) { + return + } + + if err := mgr.Remove(context.Background(), id); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} + +func handleExec(c *gin.Context) { + mgr := getManager(c) + if mgr == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "sandbox service not available"}) + return + } + + id := c.Param("id") + box, err := mgr.Get(context.Background(), id) + if err != nil { + if err == sandboxv2.ErrNotFound { + c.JSON(http.StatusNotFound, gin.H{"error": "sandbox not found"}) + return + } + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + authInfo := authorized.GetInfo(c) + if !checkBoxOwner(c, box, resolveOwner(authInfo)) { + return + } + + var req execRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + var opts []sandboxv2.ExecOption + if req.WorkDir != "" { + opts = append(opts, sandboxv2.WithWorkDir(req.WorkDir)) + } + if len(req.Env) > 0 { + opts = append(opts, sandboxv2.WithEnv(req.Env)) + } + if req.Timeout > 0 { + opts = append(opts, sandboxv2.WithTimeout(time.Duration(req.Timeout)*time.Second)) + } + + result, err := box.Exec(context.Background(), req.Cmd, opts...) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + response.RespondWithSuccess(c, http.StatusOK, result) +} + +func handleHeartbeat(c *gin.Context) { + mgr := getManager(c) + if mgr == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "sandbox service not available"}) + return + } + + id := c.Param("id") + + box, err := mgr.Get(context.Background(), id) + if err != nil { + if err == sandboxv2.ErrNotFound { + c.JSON(http.StatusNotFound, gin.H{"error": "sandbox not found"}) + return + } + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + authInfo := authorized.GetInfo(c) + if !checkBoxOwner(c, box, resolveOwner(authInfo)) { + return + } + + var req heartbeatRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + if err := mgr.Heartbeat(id, req.Active, req.ProcessCount); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} diff --git a/openapi/tests/nodes/nodes_test.go b/openapi/tests/nodes/nodes_test.go new file mode 100644 index 00000000..b1dfc19c --- /dev/null +++ b/openapi/tests/nodes/nodes_test.go @@ -0,0 +1,91 @@ +package openapi_test + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/openapi" + "github.com/yaoapp/yao/openapi/tests/testutils" +) + +func TestNodesListAuthenticated(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Nodes Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("GET", serverURL+baseURL+"/nodes", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) + + var result []map[string]interface{} + err = json.NewDecoder(resp.Body).Decode(&result) + assert.NoError(t, err) + assert.NotNil(t, result) + t.Logf("Nodes list returned %d items", len(result)) + + for _, node := range result { + assert.NotEmpty(t, node["tai_id"], "node should have tai_id") + assert.NotEmpty(t, node["mode"], "node should have mode") + assert.NotEmpty(t, node["status"], "node should have status") + t.Logf("Node: tai_id=%s, mode=%s, status=%s", node["tai_id"], node["mode"], node["status"]) + } +} + +func TestNodesListUnauthorized(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + resp, err := http.Get(serverURL + baseURL + "/nodes") + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +func TestNodesGetNotFound(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Nodes Get Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("GET", serverURL+baseURL+"/nodes/nonexistent-node", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.True(t, resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusServiceUnavailable, + "expected 404 or 503, got %d", resp.StatusCode) +} diff --git a/openapi/tests/sandbox/sandbox_test.go b/openapi/tests/sandbox/sandbox_test.go new file mode 100644 index 00000000..7f3fedca --- /dev/null +++ b/openapi/tests/sandbox/sandbox_test.go @@ -0,0 +1,150 @@ +package openapi_test + +import ( + "encoding/json" + "net/http" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/openapi" + "github.com/yaoapp/yao/openapi/tests/testutils" +) + +func TestSandboxListPublicDenied(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + resp, err := http.Get(serverURL + baseURL + "/sandbox") + assert.NoError(t, err) + defer resp.Body.Close() + + // Without auth, sandbox list returns 200 (scopes.yml allows GET /sandbox/*) + // but since /sandbox (no trailing wildcard match) could be denied or allowed, + // check that a response is returned. + assert.True(t, resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusUnauthorized, + "expected 200 or 401, got %d", resp.StatusCode) +} + +func TestSandboxListAuthenticated(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Sandbox Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("GET", serverURL+baseURL+"/sandbox", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) + + var result []map[string]interface{} + err = json.NewDecoder(resp.Body).Decode(&result) + assert.NoError(t, err) + assert.NotNil(t, result) + t.Logf("Sandbox list returned %d items", len(result)) +} + +func TestSandboxGetNotFound(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Sandbox NotFound Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("GET", serverURL+baseURL+"/sandbox/nonexistent-id", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Either 404 (sandbox not found) or 503 (sandbox service not available) is acceptable + assert.True(t, resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusServiceUnavailable, + "expected 404 or 503, got %d", resp.StatusCode) +} + +func TestSandboxCreateMissingImage(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Sandbox Create Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + body := `{"node_id": "local"}` + req, err := http.NewRequest("POST", serverURL+baseURL+"/sandbox", jsonBody(body)) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + req.Header.Set("Content-Type", "application/json") + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Should return 400 (image required) or 503 (service unavailable) + assert.True(t, resp.StatusCode == http.StatusBadRequest || resp.StatusCode == http.StatusServiceUnavailable, + "expected 400 or 503, got %d", resp.StatusCode) +} + +func TestSandboxDeleteNotFound(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Sandbox Delete Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("DELETE", serverURL+baseURL+"/sandbox/nonexistent-id", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Either 404 or 503 is acceptable + assert.True(t, resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusServiceUnavailable, + "expected 404 or 503, got %d", resp.StatusCode) +} + +func jsonBody(s string) *strings.Reader { + return strings.NewReader(s) +} diff --git a/openapi/tests/workspace/workspace_test.go b/openapi/tests/workspace/workspace_test.go new file mode 100644 index 00000000..d637f90a --- /dev/null +++ b/openapi/tests/workspace/workspace_test.go @@ -0,0 +1,110 @@ +package openapi_test + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/openapi" + "github.com/yaoapp/yao/openapi/tests/testutils" +) + +func TestWorkspaceListAuthenticated(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Workspace Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("GET", serverURL+baseURL+"/workspace", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) + + var result []map[string]interface{} + err = json.NewDecoder(resp.Body).Decode(&result) + assert.NoError(t, err) + assert.NotNil(t, result) + t.Logf("Workspace list returned %d items", len(result)) +} + +func TestWorkspaceListUnauthorized(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + resp, err := http.Get(serverURL + baseURL + "/workspace") + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +func TestWorkspaceGetNotFound(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Workspace Get Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("GET", serverURL+baseURL+"/workspace/nonexistent-ws", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.True(t, resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusServiceUnavailable, + "expected 404 or 503, got %d", resp.StatusCode) +} + +func TestWorkspaceDeleteNotFound(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + client := testutils.RegisterTestClient(t, "Workspace Delete Test Client", []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, client.ClientID) + + tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile") + + req, err := http.NewRequest("DELETE", serverURL+baseURL+"/workspace/nonexistent-ws", nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken) + + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.True(t, resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusServiceUnavailable, + "expected 404 or 503, got %d", resp.StatusCode) +} diff --git a/openapi/workspace/workspace.go b/openapi/workspace/workspace.go new file mode 100644 index 00000000..5a550199 --- /dev/null +++ b/openapi/workspace/workspace.go @@ -0,0 +1,388 @@ +package workspace + +import ( + "context" + "encoding/base64" + "fmt" + "io" + "mime" + "net/http" + "path/filepath" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/oauth/authorized" + "github.com/yaoapp/yao/openapi/oauth/types" + "github.com/yaoapp/yao/openapi/response" + ws "github.com/yaoapp/yao/workspace" +) + +// Attach registers workspace management routes on the given group. +// - GET / — list workspaces (filtered by owner from token) +// - POST / — create workspace (owner from token) +// - GET /:id — get workspace (owner check) +// - PUT /:id — update workspace (owner check) +// - DELETE /:id — delete workspace (owner check) +// - GET /:id/files — list files +// - GET /:id/files/*path — read file +// - PUT /:id/files/*path — write file +// - DELETE /:id/files/*path — delete file +// - POST /:id/mkdir — create directory +// - POST /:id/rename — rename file/directory +func Attach(group *gin.RouterGroup, oauth types.OAuth) { + group.Use(oauth.Guard) + + group.GET("", handleList) + group.POST("", handleCreate) + group.GET("/:id", handleGet) + group.PUT("/:id", handleUpdate) + group.DELETE("/:id", handleDelete) + + group.GET("/:id/files", handleListFiles) + group.GET("/:id/files/*path", handleReadFile) + group.PUT("/:id/files/*path", handleWriteFile) + group.DELETE("/:id/files/*path", handleDeleteFile) + group.POST("/:id/mkdir", handleMkdir) + group.POST("/:id/rename", handleRename) +} + +// resolveOwner returns TeamID if present, otherwise UserID. +func resolveOwner(authInfo *types.AuthorizedInfo) string { + if authInfo != nil && authInfo.TeamID != "" { + return authInfo.TeamID + } + if authInfo != nil { + return authInfo.UserID + } + return "" +} + +// checkWSOwner verifies the caller owns the workspace. +func checkWSOwner(c *gin.Context, w *ws.Workspace, owner string) bool { + if owner == "" { + return true + } + if w.Owner != "" && w.Owner != owner { + c.JSON(http.StatusForbidden, gin.H{"error": "no permission to access this workspace"}) + return false + } + return true +} + +// --- request / response types --- + +type createRequest struct { + ID string `json:"id,omitempty"` + Name string `json:"name" binding:"required"` + Node string `json:"node" binding:"required"` + Labels map[string]string `json:"labels,omitempty"` +} + +type updateRequest struct { + Name *string `json:"name,omitempty"` + Labels map[string]string `json:"labels,omitempty"` +} + +type mkdirRequest struct { + Path string `json:"path" binding:"required"` +} + +type renameRequest struct { + OldPath string `json:"old_path" binding:"required"` + NewPath string `json:"new_path" binding:"required"` +} + +type workspaceResponse struct { + ID string `json:"id"` + Name string `json:"name"` + Owner string `json:"owner"` + Node string `json:"node"` + Labels map[string]string `json:"labels,omitempty"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +func toResponse(w *ws.Workspace) workspaceResponse { + return workspaceResponse{ + ID: w.ID, + Name: w.Name, + Owner: w.Owner, + Node: w.Node, + Labels: w.Labels, + CreatedAt: w.CreatedAt.Format("2006-01-02T15:04:05Z"), + UpdatedAt: w.UpdatedAt.Format("2006-01-02T15:04:05Z"), + } +} + +func mgr() *ws.Manager { + return ws.M() +} + +// resolveAndCheckWS fetches the workspace and verifies owner permission. +func resolveAndCheckWS(c *gin.Context) (*ws.Workspace, bool) { + m := mgr() + if m == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "workspace service not available"}) + return nil, false + } + + w, err := m.Get(context.Background(), c.Param("id")) + if err != nil { + if err == ws.ErrNotFound { + c.JSON(http.StatusNotFound, gin.H{"error": "workspace not found"}) + return nil, false + } + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return nil, false + } + + authInfo := authorized.GetInfo(c) + if !checkWSOwner(c, w, resolveOwner(authInfo)) { + return nil, false + } + return w, true +} + +// --- handlers --- + +func handleList(c *gin.Context) { + m := mgr() + if m == nil { + response.RespondWithSuccess(c, http.StatusOK, []workspaceResponse{}) + return + } + + authInfo := authorized.GetInfo(c) + owner := resolveOwner(authInfo) + + list, err := m.List(context.Background(), ws.ListOptions{ + Owner: owner, + Node: c.Query("node"), + }) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + result := make([]workspaceResponse, 0, len(list)) + for _, w := range list { + result = append(result, toResponse(w)) + } + response.RespondWithSuccess(c, http.StatusOK, result) +} + +func handleCreate(c *gin.Context) { + m := mgr() + if m == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": "workspace service not available"}) + return + } + + var req createRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + authInfo := authorized.GetInfo(c) + owner := resolveOwner(authInfo) + + w, err := m.Create(context.Background(), ws.CreateOptions{ + ID: req.ID, + Name: req.Name, + Owner: owner, + Node: req.Node, + Labels: req.Labels, + }) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + response.RespondWithSuccess(c, http.StatusCreated, toResponse(w)) +} + +func handleGet(c *gin.Context) { + w, ok := resolveAndCheckWS(c) + if !ok { + return + } + response.RespondWithSuccess(c, http.StatusOK, toResponse(w)) +} + +func handleUpdate(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + var req updateRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + w, err := mgr().Update(context.Background(), c.Param("id"), ws.UpdateOptions{ + Name: req.Name, + Labels: req.Labels, + }) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + response.RespondWithSuccess(c, http.StatusOK, toResponse(w)) +} + +func handleDelete(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + force := c.Query("force") == "true" + if err := mgr().Delete(context.Background(), c.Param("id"), force); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} + +func handleListFiles(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + dir := c.DefaultQuery("path", ".") + entries, err := mgr().ListDir(context.Background(), c.Param("id"), dir) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + response.RespondWithSuccess(c, http.StatusOK, entries) +} + +func handleReadFile(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + path := c.Param("path") + if len(path) > 0 && path[0] == '/' { + path = path[1:] + } + + fmt.Printf("[workspace] handleReadFile id=%s path=%q\n", c.Param("id"), path) + + data, err := mgr().ReadFile(context.Background(), c.Param("id"), path) + if err != nil { + fmt.Printf("[workspace] ReadFile error: %v\n", err) + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + fmt.Printf("[workspace] ReadFile ok, size=%d, encoding=%q\n", len(data), c.Query("encoding")) + + if c.Query("encoding") == "base64" { + response.RespondWithSuccess(c, http.StatusOK, gin.H{ + "content": base64.StdEncoding.EncodeToString(data), + "encoding": "base64", + }) + return + } + + ext := filepath.Ext(path) + mimeType := mime.TypeByExtension(ext) + if mimeType == "" { + mimeType = "application/octet-stream" + } + fmt.Printf("[workspace] serving ext=%q mime=%q size=%d\n", ext, mimeType, len(data)) + c.Data(http.StatusOK, mimeType, data) +} + +func handleWriteFile(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + path := c.Param("path") + if len(path) > 0 && path[0] == '/' { + path = path[1:] + } + + data, err := io.ReadAll(c.Request.Body) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "failed to read body"}) + return + } + + if err := mgr().WriteFile(context.Background(), c.Param("id"), path, data, 0644); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} + +func handleDeleteFile(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + path := c.Param("path") + if len(path) > 0 && path[0] == '/' { + path = path[1:] + } + + if err := mgr().Remove(context.Background(), c.Param("id"), path); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} + +func handleMkdir(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + var req mkdirRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + if err := mgr().MkdirAll(context.Background(), c.Param("id"), req.Path); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} + +func handleRename(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + var req renameRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + if err := mgr().Rename(context.Background(), c.Param("id"), req.OldPath, req.NewPath); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.Status(http.StatusNoContent) +} diff --git a/sandbox/docker/browser/Dockerfile b/sandbox/docker/browser/Dockerfile index eace8b69..611e35f0 100644 --- a/sandbox/docker/browser/Dockerfile +++ b/sandbox/docker/browser/Dockerfile @@ -76,7 +76,7 @@ ENV DISPLAY=:99 ENV VNC_PORT=5900 ENV NOVNC_PORT=6080 ENV RESOLUTION=1920x1080x24 -ENV SANDBOX_VNC_ENABLED=true +ENV VNC_ENABLED=true ENV SANDBOX_DESKTOP=fluxbox # Node.js environment - ensure global modules are accessible diff --git a/sandbox/docker/chrome/Dockerfile b/sandbox/docker/chrome/Dockerfile index 7d456a1b..b7ba373d 100644 --- a/sandbox/docker/chrome/Dockerfile +++ b/sandbox/docker/chrome/Dockerfile @@ -127,7 +127,7 @@ ENV DISPLAY=:99 ENV VNC_PORT=5900 ENV NOVNC_PORT=6080 ENV RESOLUTION=1920x1080x24 -ENV SANDBOX_VNC_ENABLED=true +ENV VNC_ENABLED=true ENV SANDBOX_DESKTOP=fluxbox # Node.js environment diff --git a/sandbox/docker/desktop/Dockerfile b/sandbox/docker/desktop/Dockerfile index aaa26c87..64d6aeed 100644 --- a/sandbox/docker/desktop/Dockerfile +++ b/sandbox/docker/desktop/Dockerfile @@ -85,7 +85,7 @@ ENV DISPLAY=:99 ENV VNC_PORT=5900 ENV NOVNC_PORT=6080 ENV RESOLUTION=1920x1080x24 -ENV SANDBOX_VNC_ENABLED=true +ENV VNC_ENABLED=true ENV SANDBOX_DESKTOP=xfce # Set hostname for XFCE panel display ENV HOSTNAME="Yao Sandbox" diff --git a/sandbox/docker/vnc/entrypoint-vnc.sh b/sandbox/docker/vnc/entrypoint-vnc.sh index 62e48f50..a25a2e14 100644 --- a/sandbox/docker/vnc/entrypoint-vnc.sh +++ b/sandbox/docker/vnc/entrypoint-vnc.sh @@ -5,7 +5,7 @@ # ============================================ # VNC Services Startup # ============================================ -if [ "$SANDBOX_VNC_ENABLED" = "true" ]; then +if [ "$VNC_ENABLED" = "true" ]; then echo "[Entrypoint] Starting VNC services..." /usr/local/bin/start-vnc.sh & # Wait for VNC to initialize diff --git a/sandbox/manager.go b/sandbox/manager.go index 83dd6588..133995bb 100644 --- a/sandbox/manager.go +++ b/sandbox/manager.go @@ -350,8 +350,7 @@ func (m *Manager) createContainer(ctx context.Context, opts CreateOptions) (*Con "6080/tcp": struct{}{}, // noVNC websockify "5900/tcp": struct{}{}, // VNC } - // Enable SANDBOX_VNC_ENABLED environment variable - containerConfig.Env = append(containerConfig.Env, "SANDBOX_VNC_ENABLED=true") + containerConfig.Env = append(containerConfig.Env, "VNC_ENABLED=true") // Map to random available ports on 127.0.0.1 hostConfig.PortBindings = nat.PortMap{ diff --git a/sandbox/v2/DESIGN.md b/sandbox/v2/DESIGN.md index f6797824..243589cc 100644 --- a/sandbox/v2/DESIGN.md +++ b/sandbox/v2/DESIGN.md @@ -41,14 +41,17 @@ Sandbox does NOT import or depend on Agent. Agent is one of many consumers. │ ├── EnsureImage / ImageExists / PullImage │ │ └── guard rails (limits, TTL) + Box factory │ │ │ -│ Box (per-instance) │ +│ Computer (unified interface) │ │ ├── Exec(cmd) → ExecResult │ │ ├── Stream(cmd) → ExecStream (real-time I/O) │ -│ ├── Attach(port) → ServiceConn (WS/SSE) │ -│ ├── Workspace() → workspace.FS │ │ ├── VNC() → url │ -│ ├── Proxy(port) → url │ -│ └── Start / Stop / Remove / Info │ +│ ├── Proxy(port, path) → url │ +│ ├── ComputerInfo() → ComputerInfo │ +│ ├── BindWorkplace(id) / Workplace() → FS │ +│ └── [Box-specific: Attach/Start/Stop/Remove] │ +│ │ +│ Box (container) ── implements Computer │ +│ Host (bare metal) ── implements Computer │ └──────────────────┬──────────────────────────────┘ │ ▼ @@ -251,9 +254,70 @@ const ( const DefaultStopTimeout = 2 * time.Second ``` +## Computer Interface + +`Computer` is the unified interface for execution environments. Both `Box` (container) and `Host` (bare metal) implement it, allowing callers to work with any execution environment without knowing the underlying runtime. + +```go +type Computer interface { + ComputerInfo() ComputerInfo + Exec(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecResult, error) + Stream(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecStream, error) + VNC(ctx context.Context) (string, error) + Proxy(ctx context.Context, port int, path string) (string, error) + BindWorkplace(workspaceID string) + Workplace() workspace.FS +} +``` + +### ComputerInfo + +```go +type ComputerInfo struct { + Kind string // "box" | "host" + Pool string + TaiID string + MachineID string + Version string + System SystemInfo + Mode string // "direct" | "tunnel" + Capabilities map[string]bool + Status string + + // Box-specific (zero values for Host) + BoxID string + ContainerID string + Owner string + Image string + Policy LifecyclePolicy + Labels map[string]string +} + +type SystemInfo struct { + OS string + Arch string + Hostname string + NumCPU int + TotalMem int64 +} +``` + +### Workplace Binding + +Workspace is a Node-level resource, decoupled from the Computer. A Computer can bind to a workspace at session time: + +- `BindWorkplace(workspaceID)` — binds a workspace to this Computer (virtual record, rebind to change) +- `Workplace()` — returns the bound workspace FS, or nil if unbound +- Box: automatically bound via `CreateOptions.WorkspaceID`, can rebind with `BindWorkplace()` +- Host: explicitly bound in the session + +### VNC and Proxy on Host + +Host VNC and Proxy use the special `__host__` identifier to route to the Tai server's localhost instead of a container. The Tai server's VNC router and HTTP proxy both handle `__host__` by connecting to `127.0.0.1:{port}` directly, bypassing the container resolver. + ## Box -A `Box` is a single sandbox instance. All operations go through it. +A `Box` is a single sandbox instance backed by a container. It implements the `Computer` interface and adds container-specific methods (Attach, Start, Stop, Remove, Info). ```go type Box struct { @@ -303,7 +367,9 @@ func (b *Box) Remove(ctx context.Context) error func (b *Box) Info(ctx context.Context) (*BoxInfo, error) ``` -### ExecOption / ExecResult / ExecStream +### ExecOption / ExecResult / ExecStream (unified) + +These types are shared between Box and Host via the Computer interface. ```go type ExecOption func(*execConfig) @@ -311,11 +377,16 @@ type ExecOption func(*execConfig) func WithWorkDir(dir string) ExecOption func WithEnv(env map[string]string) ExecOption func WithTimeout(d time.Duration) ExecOption +func WithStdin(data []byte) ExecOption +func WithMaxOutput(bytes int64) ExecOption type ExecResult struct { - ExitCode int - Stdout string - Stderr string + ExitCode int + Stdout string + Stderr string + DurationMs int64 // Host fills; Box = 0 + Error string // Host fills; Box = "" + Truncated bool // Host fills; Box = false } type ExecStream struct { @@ -531,36 +602,36 @@ type Proxy interface { Local: resolves host ports via `Inspect()`. Remote: routes through Tai HTTP proxy which handles WebSocket upgrade and SSE streaming natively. -## gRPC Token Injection +## gRPC Environment Injection ```go -func CreateContainerTokens(sandboxID, owner string, scopes []string) (access, refresh string, err error) -func RevokeContainerTokens(refresh string) error -func BuildGRPCEnv(pool *Pool, sandboxID, access, refresh string, grpcPort int) map[string]string +func BuildGRPCEnv(pool *Pool, sandboxID string, grpcPort int) map[string]string ``` -Environment variables injected into each container: +`BuildGRPCEnv` sets **only** routing variables — token injection is decoupled: ``` -# All modes +# Set by BuildGRPCEnv (always) YAO_SANDBOX_ID= +YAO_GRPC_ADDR=127.0.0.1:9099 # local / tunnel mode +YAO_GRPC_ADDR=:19100 # remote mode (tai://) + +# Set by caller via CreateOptions.Env (OAuth is caller's responsibility) YAO_TOKEN= YAO_REFRESH_TOKEN= -YAO_GRPC_ADDR=127.0.0.1:9099 - -# Remote mode (tai://) -YAO_GRPC_ADDR=:9100 ``` +`CreateOptions.Env` is merged **after** `BuildGRPCEnv`, so the caller can override any variable including `YAO_GRPC_ADDR`. + ## Errors ```go var ( - ErrNotAvailable = errors.New("sandbox: not available (no pools configured)") + ErrNotAvailable = errors.New("sandbox: not available (no nodes registered)") ErrNotFound = errors.New("sandbox: not found") ErrLimitExceeded = errors.New("sandbox: limit exceeded") - ErrPoolNotFound = errors.New("sandbox: pool not found") - ErrPoolInUse = errors.New("sandbox: pool has running boxes") + ErrNodeNotFound = errors.New("sandbox: node not found") + ErrNodeMissing = errors.New("sandbox: node ID missing") ) ``` @@ -570,14 +641,16 @@ var ( sandbox/v2/ ├── sandbox.go // Init, M(), global singleton ├── manager.go // Manager: CRUD, pool management, image ops, cleanup -├── box.go // Box: Exec, Stream, Attach, Workspace, VNC, Proxy, lifecycle -├── types.go // CreateOptions, ExecResult, ExecStream, ServiceConn, BoxInfo, etc. +├── types.go // Computer interface, ComputerInfo, ExecResult, ExecStream, etc. +├── box.go // Box: implements Computer + Attach/Start/Stop/Remove/Info +├── host.go // Host: implements Computer (HostExec gRPC + __host__ VNC/Proxy) ├── config.go // Config struct ├── errors.go // sentinel errors ├── grpc.go // token creation/revocation, gRPC env var injection ├── jsapi/ // (Phase 2) V8 JSAPI sandbox.* namespace -│ ├── jsapi.go // RegisterObject("sandbox"), Create/Get/List/Delete -│ └── box.go // Box JS object: Exec/Attach/VNC/Proxy/Workspace/Info/Start/Stop/Remove +│ ├── jsapi.go // RegisterObject("sandbox"), Create/Get/List/Delete/Host +│ ├── computer.go // Unified Computer JS object (box + host), sbHost() +│ └── node.go // GetNode/Nodes/NodesByTeam JS bindings ├── export_test.go // ResetForTest() for test isolation ├── testutils_test.go // shared test helpers (multi-pool setup) ├── sandbox_test.go // Init/M singleton tests @@ -586,6 +659,7 @@ sandbox/v2/ ├── box_test.go // Box Exec/Workspace/Info tests ├── box_attach_test.go // Attach WS/SSE/VNC tests ├── box_workspace_test.go // Workspace integration tests +├── host_test.go // Host Exec/Stream/VNC/Proxy/ComputerInfo tests ├── box_image_test.go // Image Pull API tests ├── bench_test.go // Performance benchmarks ├── grpc_test.go // Token/env building tests @@ -851,6 +925,10 @@ Static methods: | `sandbox.Get(id)` | `Manager.Get(ctx, id)` | `Box \| null` | | `sandbox.List(filter?)` | `Manager.List(ctx, ListOptions)` → `Box.Info()` | `BoxInfo[]` | | `sandbox.Delete(id)` | `Manager.Remove(ctx, id)` | `void` | +| `sandbox.Host(pool?)` | `Manager.Host(ctx, pool)` | `Computer (Host)` | +| `sandbox.GetNode(taiID)` | `registry.Global().Get(taiID)` | `NodeInfo \| null` | +| `sandbox.Nodes()` | `registry.Global().List()` | `NodeInfo[]` | +| `sandbox.NodesByTeam(teamID)` | `registry.Global().ListByTeam(teamID)` | `NodeInfo[]` | `sandbox.Create(options)` — JS options → Go `CreateOptions`: @@ -866,7 +944,7 @@ Static methods: memory: number → CreateOptions.Memory // bytes (int64) cpus: number → CreateOptions.CPUs // float64 vnc: boolean → CreateOptions.VNC - ports: array → CreateOptions.Ports // [{container, host, host_ip, protocol}] → []PortMapping + ports: array → CreateOptions.Ports // [{container_port, host_port, host_ip, protocol}] → []PortMapping policy: string → CreateOptions.Policy // "oneshot"|"session"|"longrunning"|"persistent" idle_timeout: number → CreateOptions.IdleTimeout // ms → time.Duration stop_timeout: number → CreateOptions.StopTimeout // ms → time.Duration @@ -918,13 +996,23 @@ Read-only properties: Methods: +Computer interface methods: + | JS | Go | Returns | |----|-----|---------| -| `box.Exec(cmd, opts?)` | `Box.Exec(ctx, cmd, ...ExecOption)` | `ExecResult` | -| `box.Stream(cmd, opts?)` | `Box.Stream(ctx, cmd, ...ExecOption)` | `ExecStream` | -| `box.Attach(port, opts?)` | `Box.Attach(ctx, port, ...AttachOption)` | `ServiceConn` | -| `box.VNC()` | `Box.VNC(ctx)` | `string` | -| `box.Proxy(port, path?)` | `Box.Proxy(ctx, port, path)` | `string` | +| `box.Exec(cmd, opts?)` | `Computer.Exec(ctx, cmd []string, ...ExecOption)` | `ExecResult` | +| `box.Stream(cmd, [opts,] cb)` | `Computer.Stream(ctx, cmd []string, ...ExecOption)` | callback(type, data) | +| `box.VNC()` | `Computer.VNC(ctx)` | `string` | +| `box.Proxy(port, path?)` | `Computer.Proxy(ctx, port, path)` | `string` | +| `box.ComputerInfo()` | `Computer.ComputerInfo()` | `ComputerInfo` | +| `box.BindWorkplace(id)` | `Computer.BindWorkplace(id)` | `void` | +| `box.Workplace()` | `Computer.Workplace()` | `WorkspaceFS \| null` | + +Box-specific methods: + +| JS | Go | Returns | +|----|-----|---------| +| `box.Attach(port, opts?)` | `Proxy.URL(ctx, containerID, port, path)` | `string` (URL) | | `box.Workspace()` | `Box.WorkspaceID()` → `NewFSObject` | `WorkspaceFS` | | `box.Info()` | `Box.Info(ctx)` | `BoxInfo` | | `box.Start()` | `Box.Start(ctx)` | `void` | @@ -936,28 +1024,31 @@ Methods: ``` cmd: string[] → cmd []string options: { - workdir: string, → WithWorkDir(dir) - env: object, → WithEnv(map[string]string) - timeout: number → WithTimeout(ms → time.Duration) + workdir: string, → WithWorkDir(dir) + env: object, → WithEnv(map[string]string) + stdin: string, → WithStdin([]byte) + timeout: number, → WithTimeout(ms → time.Duration) + max_output: number → WithMaxOutput(bytes int64) } returns: { - exit_code: number, ← ExecResult.ExitCode - stdout: string, ← ExecResult.Stdout - stderr: string ← ExecResult.Stderr + exit_code: number, ← ExecResult.ExitCode + stdout: string, ← ExecResult.Stdout + stderr: string, ← ExecResult.Stderr + duration_ms: number, ← ExecResult.DurationMs (Host fills; Box = 0) + error: string, ← ExecResult.Error (Host fills; Box = "") + truncated: boolean ← ExecResult.Truncated (Host fills; Box = false) } ``` -`box.Stream(cmd, options?)`: +`box.Stream(cmd, callback)` / `box.Stream(cmd, options, callback)`: ``` -options: same as Exec -returns: { - stdout: ReadableStream, ← ExecStream.Stdout - stderr: ReadableStream, ← ExecStream.Stderr - stdin: WritableStream, ← ExecStream.Stdin - wait: function() → number, ← ExecStream.Wait() (int, error) - cancel: function() → void ← ExecStream.Cancel() -} +Blocks until exit. Last arg must be a JS function. +options: same as Exec (optional) +callback: function(type, data) + type = "stdout" → data is string (chunk) ← ExecStream.Stdout + type = "stderr" → data is string (chunk) ← ExecStream.Stderr + type = "exit" → data is number (exit code) ← ExecStream.Wait() ``` `box.Attach(port, options?)`: @@ -965,20 +1056,107 @@ returns: { ``` port: number → port int options: { - protocol: "ws"|"sse", → WithProtocol(protocol) - path: string, → WithPath(path) - headers: object → WithHeaders(map[string]string) + protocol: "ws"|"sse", → affects URL scheme (ws:// vs http://) + path: string, → URL path suffix +} +returns: string (URL) ← Proxy.URL(ctx, containerID, port, path) +``` + +Caller (frontend, Agent) establishes the actual WS/SSE connection using the returned URL. +Go-side `ServiceConn` (with Read/Write/Events/Close) is available for Go callers only. + +`box.Info()` returns same structure as `BoxInfo[]` element above. + +#### Host object (Computer) + +Host implements the unified Computer interface for Tai host machines. It executes commands via HostExec gRPC and accesses VNC/Proxy via the `__host__` identifier. Available only when the pool's Tai server exposes HostExec gRPC. JS object holds pool name; all methods delegate to `sandbox.M().Host(ctx, pool)`. + +Read-only properties: + +| JS | Go | +|----|----| +| `host.pool` | `Host.Pool()` | + +Methods (same Computer interface as Box): + +| JS | Go | Returns | +|----|-----|---------| +| `host.Exec(cmd, opts?)` | `Computer.Exec(ctx, cmd []string, ...ExecOption)` | `ExecResult` | +| `host.Stream(cmd, [opts,] cb)` | `Computer.Stream(ctx, cmd []string, ...ExecOption)` | callback(type, data) | +| `host.VNC()` | `Computer.VNC(ctx)` | `string` (URL) | +| `host.Proxy(port, path?)` | `Computer.Proxy(ctx, port, path)` | `string` (URL) | +| `host.ComputerInfo()` | `Computer.ComputerInfo()` | `ComputerInfo` | +| `host.BindWorkplace(id)` | `Computer.BindWorkplace(id)` | `void` | +| `host.Workplace()` | `Computer.Workplace()` | `WorkspaceFS \| null` | + +`host.Exec(cmd, options?)`: + +``` +cmd: string[] → cmd []string (unified with Box) +options: { + workdir: string, → WithWorkDir(dir) + env: object, → WithEnv(map[string]string) + stdin: string, → WithStdin([]byte) + timeout: number, → WithTimeout(ms → time.Duration) + max_output: number → WithMaxOutput(bytes int64) } returns: { - url: string, ← ServiceConn.URL - read: function() → Uint8Array, ← ServiceConn.Read() - write: function(data) → void, ← ServiceConn.Write(data) - events: AsyncIterable, ← ServiceConn.Events - close: function() → void ← ServiceConn.Close() + exit_code: number, ← ExecResult.ExitCode + stdout: string, ← ExecResult.Stdout + stderr: string, ← ExecResult.Stderr + duration_ms: number, ← ExecResult.DurationMs + error: string, ← ExecResult.Error + truncated: boolean ← ExecResult.Truncated } ``` -`box.Info()` returns same structure as `BoxInfo[]` element above. +`host.Stream(cmd, callback)` / `host.Stream(cmd, options, callback)`: + +``` +Blocks until exit. Last arg must be a JS function. +options: same as host.Exec (optional) +callback: function(type, data) + type = "stdout" → data is string (chunk) ← ExecStream.Stdout (io.ReadCloser) + type = "stderr" → data is string (chunk) ← ExecStream.Stderr (io.ReadCloser) + type = "exit" → data is number (exit code) ← ExecStream.Wait() +``` + +#### NodeInfo object + +`sandbox.GetNode()`, `sandbox.Nodes()`, `sandbox.NodesByTeam()` return NodeInfo objects mapped from `registry.NodeSnapshot`. Auth and YaoBase fields are excluded for security. + +``` +{ + tai_id: string, ← NodeSnapshot.TaiID + machine_id: string, ← NodeSnapshot.MachineID + version: string, ← NodeSnapshot.Version + mode: string, ← NodeSnapshot.Mode ("direct"|"tunnel") + addr: string, ← NodeSnapshot.Addr + status: string, ← NodeSnapshot.Status ("online"|"offline"|"connecting") + pool: string, ← NodeSnapshot.PoolName + connected_at: string, ← NodeSnapshot.ConnectedAt (ISO 8601) + last_ping: string, ← NodeSnapshot.LastPing (ISO 8601) + ports: { ← NodeSnapshot.Ports + grpc: number, + http: number, + vnc: number, + docker: number, + k8s: number, + }, + capabilities: { ← NodeSnapshot.Capabilities + docker: boolean, + k8s: boolean, + host_exec: boolean, + }, + system: { ← NodeSnapshot.System (SystemInfo) + os: string, + arch: string, + hostname: string, + num_cpu: number, + total_mem: number, + } +} +``` #### workspace namespace (`RegisterObject("workspace")`) diff --git a/sandbox/v2/IMPL.md b/sandbox/v2/IMPL.md index 4514345d..9c200fe9 100644 --- a/sandbox/v2/IMPL.md +++ b/sandbox/v2/IMPL.md @@ -22,12 +22,11 @@ Reference: [DESIGN.md](./DESIGN.md) | File | What | Status | |------|------|--------| | `sandbox.go` | `Init()`, `M()`, global singleton | DONE | -| `manager.go` | Manager: Create/Get/GetOrCreate/List/Remove/Cleanup/Close, Start (container recovery), AddPool/RemovePool/Pools, Heartbeat, SetGRPCPort, SetWorkspaceManager, ImageExists/PullImage/EnsureImage | DONE | +| `manager.go` | Manager: Create/Get/GetOrCreate/List/Remove/Cleanup/Close, Start (container recovery), Nodes, Heartbeat, ImageExists/PullImage/EnsureImage | DONE | | `box.go` | Box: Exec, Stream, Attach, Workspace, VNC, Proxy, Start/Stop/Remove, Info, touch/lastActiveTime/idleTimeout/maxLifetime/stopTimeout | DONE | -| `types.go` | LifecyclePolicy (OneShot/Session/LongRunning/Persistent), Pool, PoolInfo, PortMapping, CreateOptions (with WorkspaceID/MountMode/MountPath), ListOptions, ExecOption/ExecResult/ExecStream, AttachOption/ServiceConn, ImagePullOptions/RegistryAuth, BoxInfo, DefaultStopTimeout | DONE | -| `config.go` | Config struct | DONE | -| `errors.go` | ErrNotAvailable, ErrNotFound, ErrLimitExceeded, ErrPoolNotFound, ErrPoolInUse | DONE | -| `grpc.go` | CreateContainerTokens, RevokeContainerTokens, BuildGRPCEnv | DONE | +| `types.go` | LifecyclePolicy (OneShot/Session/LongRunning/Persistent), NodeID, PortMapping, CreateOptions (with WorkspaceID/MountMode/MountPath), ListOptions, ExecOption/ExecResult/ExecStream, AttachOption/ServiceConn, ImagePullOptions/RegistryAuth, BoxInfo, DefaultStopTimeout | DONE | +| `errors.go` | ErrNotAvailable, ErrNotFound, ErrNodeNotFound, ErrNodeMissing | DONE | +| `grpc.go` | BuildGRPCEnv (sandbox ID + gRPC addr only; token injection is caller's responsibility via Env) | DONE | ### workspace Module — DONE @@ -51,7 +50,7 @@ Reference: [DESIGN.md](./DESIGN.md) | `box_image_test.go` | ImageExists (Docker+K8s), PullImage (progress+K8s no-op), EnsureImage, bad ref | DONE | | `grpc_test.go` | Token creation/revocation, env var building (local vs remote) | DONE | | `bench_test.go` | ContainerLifecycle, Create, Exec, ExecHeavy, Remove, Info, StopStart, WorkspaceReadWrite | DONE | -| `testutils_test.go` | testPools (local/remote/k8s), setupManager, createTestBox, ensureTestImage | DONE | +| `testutils_test.go` | testNodes (local/remote/k8s), setupManager, setupManagerForNode, createTestBox, ensureTestImage | DONE | | `export_test.go` | ResetForTest | DONE | | **workspace** | | | | `workspace_test.go` | Create (auto/explicit ID, labels, invalid node), Get, List (owner/node filter), Update (name/labels), Delete, Nodes, NodeForWorkspace, AddPool/RemovePool, MountPath | DONE | @@ -78,41 +77,102 @@ Reference: [DESIGN.md](./DESIGN.md) --- -## Phase 2: JSAPI + OAuth — PENDING +## Phase 2: JSAPI + Computer Unification — DONE -| Task | Package | Detail | -|------|---------|--------| -| `jsapi/sandbox.go` | `sandbox/v2/jsapi` | V8 JSAPI `Sandbox()` + `Workspace()` constructors (registered in gou runtime) | -| Wire `openapi/oauth` | `sandbox/v2/grpc.go` | `CreateContainerTokens` currently generates random strings; `RevokeContainerTokens` is a no-op. Replace with real `openapi/oauth` issue/revoke calls | -| `cmd/start.go` integration | `yao` | Call `sandbox.Init(config.Conf.Sandbox)` + `sandbox.M().Start(ctx)` in startup sequence | -| Heartbeat bridge | `yao/grpc` | Wire gRPC Heartbeat handler → `sandbox.M().Heartbeat()` | +### Unified Computer Interface — DONE -### JSAPI (planned) +Box and Host now share a single `Computer` interface (`types.go`). Both `sandbox.Create()` and `sandbox.Host()` return the same JS `Computer` object; `kind` property distinguishes them. Box-only methods (`Info`, `Start`, `Stop`, `Remove`) throw at runtime when called on a host. + +| Step | Package | What | Status | +|------|---------|------|--------| +| Computer interface | `sandbox/v2/types.go` | `Computer` interface: Exec, Stream, VNC, Proxy, ComputerInfo, BindWorkplace, Workplace | DONE | +| Host implementation | `sandbox/v2/host.go` | `Host` struct implements `Computer` via tai HostExec + VNC/Proxy | DONE | +| ComputerInfo | `sandbox/v2/types.go` | `ComputerInfo` struct with Kind, NodeID, TaiID, System, Capabilities, box-specific fields | DONE | + +### JSAPI — DONE + +| File | What | Status | +|------|------|--------| +| `jsapi/jsapi.go` | Static methods: `sandbox.Create`, `Get`, `List`, `Delete` | DONE | +| `jsapi/computer.go` | `NewComputerObject` factory (11 methods + 4 properties), `sbHost`, helpers | DONE | +| `jsapi/node.go` | `sandbox.GetNode`, `Nodes`, `NodesByTeam`, `snapshotToJS` | DONE | +| `jsapi/API.md` | Full JavaScript API reference | DONE | + +Design decisions: +- **No Go objects in V8**: closures capture only `kind` (string) and `identifier` (string); `getComputer()` re-fetches from Manager on each call — prevents memory leaks across runtimes. +- **Stream**: blocking with callback `function(type, data)`, goroutines feed a channel, main V8 thread drains it. +- **Workplace()**: delegates to `workspace/jsapi.NewFSObject()` — reuses existing WorkspaceFS JSAPI. ```javascript -// Sandbox -var sb = Sandbox("my-workspace", { - image: "yaoapp/workspace:latest", - owner: "user-123" +// Unified Computer — same API for box and host +const pc = sandbox.Create({ image: "node:20", owner: "user-123" }) +pc.Exec(["node", "-e", "console.log('hello')"]) +pc.Stream(["npm", "run", "dev"], function(type, data) { + if (type === "stdout") console.log(data) + if (type === "exit") console.log("exited:", data) }) -sb.Exec(["go", "build", "./..."]) -sb.ReadFile("src/main.go") -sb.WriteFile("src/main.go", "package main\n...") -sb.Stream(["npm", "run", "dev"], function(chunk) { ... }) -var conn = sb.Attach(3000, { protocol: "ws", path: "/ws" }) -sb.Info() -sb.Stop() -sb.Start() -sb.Remove() +pc.VNC() // → "ws://host:port/vnc/{id}/ws" +pc.Proxy(3000, "/api") // → "http://host:port/{id}:3000/api" +pc.ComputerInfo() // → { kind, pool, system, ... } +pc.BindWorkplace("ws-abc") +pc.Workplace().ReadFile("main.go") +pc.Info() // box-only +pc.Remove() // box-only -// Workspace -var ws = Workspace("my-workspace") -ws.ReadFile("src/main.go") -ws.WriteFile("src/main.go", "package main\n...") -ws.ListDir("src/") -ws.Remove("tmp.txt") +// Host — same interface, no container +const host = sandbox.Host("gpu") +host.Exec(["nvidia-smi"]) +host.VNC() // → "ws://host:port/vnc/__host__/ws" +host.Proxy(8080) // → "http://host:port/__host__:8080/" +host.kind // "host" +host.Info() // throws: "not supported: Info() requires a box computer" + +// Nodes (registry read-only query) +const nodes = sandbox.Nodes() +const node = sandbox.GetNode("tai-abc123") +const team = sandbox.NodesByTeam("team-001") ``` +### JSAPI Tests — DONE + +| Test | Coverage | Status | +|------|----------|--------| +| `TestCreate` | Create box, verify kind/id | DONE | +| `TestGet` | Get existing box | DONE | +| `TestGetNotFound` | Get non-existent → null | DONE | +| `TestDelete` | Delete + verify gone | DONE | +| `TestList` | List with owner filter | DONE | +| `TestExec` | Exec echo, verify stdout | DONE | +| `TestExecWithOptions` | Exec with workdir option | DONE | +| `TestStream` | Stream with callback, verify chunks + exit code | DONE | +| `TestComputerInfo` | Verify kind field | DONE | +| `TestBoxInfo` | Box-only Info() | DONE | +| `TestHostBoxMethodsThrow` | Host.Info() throws "not supported" | DONE | +| `TestComputerKind` | kind property = "box" | DONE | +| `TestNodes` | Nodes() returns array | DONE | +| `TestGetNodeNotFound` | GetNode non-existent → null | DONE | + +All 14 tests pass in both local and remote modes. + +### OAuth Decoupling — DONE + +Token injection (YAO_TOKEN, YAO_REFRESH_TOKEN) has been **removed from sandbox Manager**. +`CreateContainerTokens`, `RevokeContainerTokens`, and the `Box.refreshToken` field have been deleted. +`BuildGRPCEnv` now only sets `YAO_SANDBOX_ID` and `YAO_GRPC_ADDR`. + +Token provisioning is the **caller's responsibility** via `CreateOptions.Env`: +- The caller (e.g. Agent Hook) already holds an OAuth context +- It calls `oauth.OAuth.MakeAccessToken(...)` to issue a scoped token +- Passes it in `CreateOptions.Env["YAO_TOKEN"]` / `Env["YAO_REFRESH_TOKEN"]` +- `opts.Env` takes priority over `BuildGRPCEnv` output (caller can override anything) + +### Remaining (Startup) + +| Task | Package | Status | Detail | +|------|---------|--------|--------| +| `engine/load.go` integration | `yao` | **DONE** | `sandbox.Init()` + `sandbox.M().Start(ctx)` added as a `loadStep("Sandbox", ...)` right after Registry init | +| Heartbeat bridge | `yao/cmd` | **DONE** | `cmd/start.go` calls `yaogrpc.SetSandboxOnBeat(...)` before `service.Start`, forwarding gRPC heartbeats to `sandbox.M().Heartbeat()` | + --- ## Phase 3: Agent Integration — PENDING @@ -163,7 +223,7 @@ Manager injects these labels at creation time: managed-by=yao-sandbox sandbox-id= sandbox-owner= -sandbox-pool= +sandbox-node-id= sandbox-policy= workspace-id= (if WorkspaceID set) ``` @@ -176,14 +236,14 @@ When `CreateOptions.WorkspaceID` is set: ``` 1. NodeForWorkspace(wsID) → node name -2. Force pool = node name +2. Force nodeID = node name 3. MountPath(wsID) → hostDir 4. Bind: hostDir:/workspace:rw ``` ### Multi-Mode Testing -`testPools()` returns all available pool configurations: +`testNodes()` returns all available node configurations: ```go func testPools() []poolConfig { @@ -204,9 +264,10 @@ Every test iterates over all available pools: ```go func TestSomething(t *testing.T) { for _, pc := range testPools() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - // test logic + m := setupManagerForPool(t, &pc) + // test logic — use pc.TaiID as pool identifier }) } } @@ -228,18 +289,20 @@ K8s-specific behavior: ## File Inventory -### sandbox/v2 (7 source + 10 test = 17 files) +### sandbox/v2 (9 source + 10 test = 19 files) | File | Lines | Purpose | |------|-------|---------| | `sandbox.go` | ~25 | Global singleton | | `manager.go` | ~620 | Manager implementation | -| `box.go` | ~230 | Box implementation | -| `types.go` | ~170 | Type definitions | +| `box.go` | ~317 | Box implementation (Computer interface) | +| `host.go` | ~232 | Host implementation (Computer interface) | +| `types.go` | ~247 | Type definitions (Computer, ComputerInfo, ExecOption, etc.) | | `config.go` | ~5 | Config struct | | `errors.go` | ~10 | Error definitions | -| `grpc.go` | ~55 | Token/env injection | -| `testutils_test.go` | ~130 | Test helpers | +| `grpc.go` | ~50 | BuildGRPCEnv (sandbox ID + addr) | +| `export_test.go` | ~6 | ResetForTest | +| `testutils_test.go` | ~364 | Test helpers (multi-pool, host exec targets) | | `sandbox_test.go` | ~30 | Singleton tests | | `manager_test.go` | ~250 | CRUD tests | | `manager_lifecycle_test.go` | ~120 | Lifecycle tests | @@ -247,9 +310,19 @@ K8s-specific behavior: | `box_attach_test.go` | ~260 | Attach/VNC tests | | `box_workspace_test.go` | ~285 | Workspace tests | | `box_image_test.go` | ~120 | Image tests | -| `grpc_test.go` | ~80 | Token tests | +| `grpc_test.go` | ~40 | BuildGRPCEnv tests | | `bench_test.go` | ~230 | Benchmarks | +### sandbox/v2/jsapi (3 source + 1 test + 1 doc = 5 files) + +| File | Lines | Purpose | +|------|-------|---------| +| `jsapi.go` | ~286 | Static methods (Create/Get/List/Delete) + V8 registration | +| `computer.go` | ~472 | NewComputerObject factory, sbHost, helpers | +| `node.go` | ~143 | Node query methods (GetNode/Nodes/NodesByTeam) + snapshotToJS | +| `jsapi_test.go` | ~430 | 14 test cases (local + remote modes) | +| `API.md` | ~604 | JavaScript API reference | + ### workspace (3 source + 4 test = 7 files) | File | Lines | Purpose | @@ -261,3 +334,12 @@ K8s-specific behavior: | `workspace_test.go` | ~325 | CRUD tests | | `fileio_test.go` | ~235 | File I/O tests | | `bench_test.go` | ~150 | Benchmarks | + +### workspace/jsapi (2 source + 1 test + 1 doc = 4 files) + +| File | Lines | Purpose | +|------|-------|---------| +| `jsapi.go` | ~100 | Static methods (Create/Get/List/Delete) + V8 registration | +| `fs.go` | ~630 | NewFSObject factory (WorkspaceFS methods) | +| `jsapi_test.go` | ~460 | JSAPI tests (local + remote modes) | +| `API.md` | ~220 | Workspace JavaScript API reference | diff --git a/sandbox/v2/TEST.md b/sandbox/v2/TEST.md index a5a3459e..6434d830 100644 --- a/sandbox/v2/TEST.md +++ b/sandbox/v2/TEST.md @@ -426,7 +426,7 @@ func TestRemovePool_InUse(t *testing.T) { // manager_test.go package sandbox_test -func TestMultiPool(t *testing.T) { +func TestMultiNode(t *testing.T) { skipIfNoDocker(t) skipIfNoTai(t) cleanup := setupManagerWithRemote(t) diff --git a/sandbox/v2/bench_test.go b/sandbox/v2/bench_test.go index 605982da..eb535961 100644 --- a/sandbox/v2/bench_test.go +++ b/sandbox/v2/bench_test.go @@ -7,14 +7,17 @@ import ( "time" sandbox "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" ) // BenchmarkContainerLifecycle measures the full Create → Exec → Remove cycle. func BenchmarkContainerLifecycle(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) - ensureTestImageBench(b, m, pc.Name) + m := setupManagerForBench(b, &pc) + ensureTestImageBench(b, m, pc.TaiID) b.ResetTimer() for i := 0; i < b.N; i++ { @@ -41,10 +44,11 @@ func BenchmarkContainerLifecycle(b *testing.B) { // BenchmarkCreate measures container creation time only. func BenchmarkCreate(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) - ensureTestImageBench(b, m, pc.Name) + m := setupManagerForBench(b, &pc) + ensureTestImageBench(b, m, pc.TaiID) ids := make([]string, 0, b.N) b.ResetTimer() @@ -69,9 +73,10 @@ func BenchmarkCreate(b *testing.B) { // BenchmarkExec measures command execution latency on a pre-created container. func BenchmarkExec(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) + m := setupManagerForBench(b, &pc) box := createBoxForBench(b, m) b.ResetTimer() @@ -90,9 +95,10 @@ func BenchmarkExec(b *testing.B) { // BenchmarkExecHeavy measures execution of a heavier command (write + read file). func BenchmarkExecHeavy(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) + m := setupManagerForBench(b, &pc) box := createBoxForBench(b, m) b.ResetTimer() @@ -112,10 +118,11 @@ func BenchmarkExecHeavy(b *testing.B) { // BenchmarkRemove measures container removal time. func BenchmarkRemove(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) - ensureTestImageBench(b, m, pc.Name) + m := setupManagerForBench(b, &pc) + ensureTestImageBench(b, m, pc.TaiID) boxes := make([]*sandbox.Box, b.N) for i := 0; i < b.N; i++ { @@ -141,9 +148,10 @@ func BenchmarkRemove(b *testing.B) { // BenchmarkInfo measures Info() latency on a running container. func BenchmarkInfo(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) + m := setupManagerForBench(b, &pc) box := createBoxForBench(b, m) b.ResetTimer() @@ -159,12 +167,13 @@ func BenchmarkInfo(b *testing.B) { // BenchmarkStopStart measures Stop → Start cycle time. func BenchmarkStopStart(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { if pc.Name == "k8s" { b.Skip("K8s Stop deletes Pod; Stop→Start cycle not applicable") } - m := setupManagerForBench(b, pc) + m := setupManagerForBench(b, &pc) box := createBoxForBench(b, m) b.ResetTimer() @@ -182,9 +191,10 @@ func BenchmarkStopStart(b *testing.B) { // BenchmarkWorkspaceReadWrite measures workspace file read/write via container Box. func BenchmarkWorkspaceReadWrite(b *testing.B) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc b.Run(pc.Name, func(b *testing.B) { - m := setupManagerForBench(b, pc) + m := setupManagerForBench(b, &pc) box := createBoxForBench(b, m) ws := box.Workspace() if ws == nil { @@ -213,38 +223,46 @@ func BenchmarkWorkspaceReadWrite(b *testing.B) { // --- helpers --- -func setupManagerForBench(b *testing.B, pc poolConfig) *sandbox.Manager { +func setupManagerForBench(b *testing.B, pc *nodeConfig) *sandbox.Manager { b.Helper() - pool := sandbox.Pool{Name: pc.Name, Addr: pc.Addr, Options: pc.Options} - cfg := sandbox.Config{Pool: []sandbox.Pool{pool}} - if err := sandbox.Init(cfg); err != nil { - b.Fatalf("Init: %v", err) + reg := registry.Global() + if reg == nil { + registry.Init(nil) } + client, err := tai.New(pc.Addr, pc.Options...) + if err != nil { + b.Fatalf("tai.New(%s): %v", pc.Addr, err) + } + pc.TaiID = client.TaiID() + sandbox.Init() m := sandbox.M() b.Cleanup(func() { m.Close() }) return m } -func ensureTestImageBench(b *testing.B, m *sandbox.Manager, pool string) { +func ensureTestImageBench(b *testing.B, m *sandbox.Manager, nodeID string) { b.Helper() ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) defer cancel() - if err := m.EnsureImage(ctx, pool, testImage(), sandbox.ImagePullOptions{}); err != nil { + if err := m.EnsureImage(ctx, nodeID, testImage(), sandbox.ImagePullOptions{}); err != nil { b.Fatalf("EnsureImage: %v", err) } } func createBoxForBench(b *testing.B, m *sandbox.Manager) *sandbox.Box { b.Helper() - pools := m.Pools() - if len(pools) > 0 { - ensureTestImageBench(b, m, pools[0].Name) + nodes := m.Nodes() + var nodeID string + if len(nodes) > 0 { + nodeID = nodes[0].TaiID + ensureTestImageBench(b, m, nodeID) } ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) defer cancel() box, err := m.Create(ctx, sandbox.CreateOptions{ - Image: testImage(), - Owner: "bench", + Image: testImage(), + Owner: "bench", + NodeID: nodeID, }) if err != nil { b.Fatalf("Create: %v", err) diff --git a/sandbox/v2/box.go b/sandbox/v2/box.go index 5cb5bfc0..9b25013f 100644 --- a/sandbox/v2/box.go +++ b/sandbox/v2/box.go @@ -15,7 +15,7 @@ import ( type Box struct { id string containerID string - pool string + nodeID string owner string policy LifecyclePolicy labels map[string]string @@ -23,20 +23,55 @@ type Box struct { lastHeartbeat atomic.Int64 processCount atomic.Int32 idleTimeoutD time.Duration + maxLifetimeD time.Duration stopTimeoutD time.Duration createdAt time.Time - refreshToken string vnc bool image string workspaceID string + system SystemInfo ws workspace.FS manager *Manager } +// Compile-time check: *Box implements Computer. +var _ Computer = (*Box)(nil) + func (b *Box) ID() string { return b.id } func (b *Box) Owner() string { return b.owner } func (b *Box) ContainerID() string { return b.containerID } -func (b *Box) Pool() string { return b.pool } +func (b *Box) NodeID() string { return b.nodeID } + +// ComputerInfo returns identity and registry information for this Box. +func (b *Box) ComputerInfo() ComputerInfo { + return ComputerInfo{ + Kind: "box", + NodeID: b.nodeID, + System: b.system, + Status: "online", + BoxID: b.id, + ContainerID: b.containerID, + Owner: b.owner, + Image: b.image, + Policy: b.policy, + Labels: b.labels, + } +} + +// BindWorkplace binds (or rebinds) a workspace to this Box. Subsequent calls +// to Workplace() return the FS for this workspace. Overrides the workspace +// set during Create. +func (b *Box) BindWorkplace(workspaceID string) { + b.workspaceID = workspaceID + b.ws = nil // clear cache so Workplace() re-resolves +} + +// Workplace returns the workspace FS bound to this Box. +// If a workspace was bound via CreateOptions.WorkspaceID or BindWorkplace(), +// returns that workspace's FS. Otherwise returns nil. +func (b *Box) Workplace() workspace.FS { + return b.Workspace() +} // Exec runs a command and waits for it to finish. func (b *Box) Exec(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecResult, error) { @@ -46,7 +81,7 @@ func (b *Box) Exec(ctx context.Context, cmd []string, opts ...ExecOption) (*Exec o(cfg) } - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return nil, err } @@ -65,10 +100,6 @@ func (b *Box) Exec(ctx context.Context, cmd []string, opts ...ExecOption) (*Exec Stderr: result.Stderr, } - if b.policy == OneShot { - b.manager.Remove(ctx, b.id) - } - return r, nil } @@ -80,7 +111,7 @@ func (b *Box) Stream(ctx context.Context, cmd []string, opts ...ExecOption) (*Ex o(cfg) } - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return nil, err } @@ -110,7 +141,7 @@ func (b *Box) Attach(ctx context.Context, port int, opts ...AttachOption) (*Serv o(cfg) } - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return nil, err } @@ -156,7 +187,7 @@ func (b *Box) Workspace() workspace.FS { if sessionID == "" { sessionID = b.id } - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return nil } @@ -167,10 +198,29 @@ func (b *Box) Workspace() workspace.FS { // WorkspaceID returns the workspace ID mounted to this sandbox, or empty string. func (b *Box) WorkspaceID() string { return b.workspaceID } +// Snapshot returns a local-only BoxInfo snapshot without any remote calls. +// Status is inferred from local state (not from the container runtime). +func (b *Box) Snapshot() BoxInfo { + return BoxInfo{ + ID: b.id, + ContainerID: b.containerID, + NodeID: b.nodeID, + Owner: b.owner, + Status: "running", + Policy: b.policy, + Labels: b.labels, + Image: b.image, + CreatedAt: b.createdAt, + LastActive: b.lastActiveTime(), + ProcessCount: int(b.processCount.Load()), + VNC: b.vnc, + } +} + // VNC returns the VNC WebSocket URL. func (b *Box) VNC(ctx context.Context) (string, error) { b.touch() - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return "", err } @@ -180,7 +230,7 @@ func (b *Box) VNC(ctx context.Context) (string, error) { // Proxy returns the HTTP URL for a service on the given port inside the sandbox. func (b *Box) Proxy(ctx context.Context, port int, path string) (string, error) { b.touch() - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return "", err } @@ -189,7 +239,7 @@ func (b *Box) Proxy(ctx context.Context, port int, path string) (string, error) // Start starts a stopped sandbox. func (b *Box) Start(ctx context.Context) error { - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return err } @@ -198,7 +248,7 @@ func (b *Box) Start(ctx context.Context) error { // Stop stops the sandbox without removing it. func (b *Box) Stop(ctx context.Context) error { - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return err } @@ -212,7 +262,7 @@ func (b *Box) Remove(ctx context.Context) error { // Info returns current sandbox status. func (b *Box) Info(ctx context.Context) (*BoxInfo, error) { - client, err := b.manager.getPool(b.pool) + client, err := b.manager.getNode(b.nodeID) if err != nil { return nil, err } @@ -225,7 +275,7 @@ func (b *Box) Info(ctx context.Context) (*BoxInfo, error) { return &BoxInfo{ ID: b.id, ContainerID: b.containerID, - Pool: b.pool, + NodeID: b.nodeID, Owner: b.owner, Status: info.Status, Policy: b.policy, @@ -253,31 +303,16 @@ func (b *Box) lastActiveTime() time.Time { } func (b *Box) idleTimeout() time.Duration { - if b.idleTimeoutD > 0 { - return b.idleTimeoutD - } - pd := b.manager.findPoolDef(b.pool) - if pd != nil { - return pd.IdleTimeout - } - return 0 + return b.idleTimeoutD } func (b *Box) maxLifetime() time.Duration { - pd := b.manager.findPoolDef(b.pool) - if pd != nil { - return pd.MaxLifetime - } - return 0 + return b.maxLifetimeD } func (b *Box) stopTimeout() time.Duration { if b.stopTimeoutD > 0 { return b.stopTimeoutD } - pd := b.manager.findPoolDef(b.pool) - if pd != nil && pd.StopTimeout > 0 { - return pd.StopTimeout - } return DefaultStopTimeout } diff --git a/sandbox/v2/box_attach_test.go b/sandbox/v2/box_attach_test.go index 7c91a521..78156387 100644 --- a/sandbox/v2/box_attach_test.go +++ b/sandbox/v2/box_attach_test.go @@ -65,10 +65,11 @@ func TestAttachWS(t *testing.T) { t.Skip("WebSocket test requires tai-sandbox-test image with ws-echo service") } - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.Ports = []sandbox.PortMapping{ {ContainerPort: 9800, HostPort: 0, Protocol: "tcp"}, } @@ -113,10 +114,11 @@ func TestAttachSSE(t *testing.T) { t.Skip("SSE test requires tai-sandbox-test image with sse-server service") } - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.Ports = []sandbox.PortMapping{ {ContainerPort: 9801, HostPort: 0, Protocol: "tcp"}, } @@ -162,10 +164,11 @@ func TestVNCURL(t *testing.T) { t.Skip("VNC test requires tai-sandbox-test image with VNC desktop") } - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.VNC = true }) @@ -192,10 +195,11 @@ func TestVNCConnect(t *testing.T) { t.Skip("VNC test requires tai-sandbox-test image with VNC desktop") } - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.VNC = true }) diff --git a/sandbox/v2/box_image_test.go b/sandbox/v2/box_image_test.go index 1f43fedb..6c3927c9 100644 --- a/sandbox/v2/box_image_test.go +++ b/sandbox/v2/box_image_test.go @@ -11,33 +11,33 @@ import ( ) func TestImageExists(t *testing.T) { - for _, pc := range testPools() { + for _, pc := range testNodes() { pc := pc t.Run(pc.Name, func(t *testing.T) { if pc.Name == "k8s" { t.Run("always_true", func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - exists, err := m.ImageExists(ctx, pc.Name, "anything:nonexistent") + exists, err := m.ImageExists(ctx, pc.TaiID, "anything:nonexistent") require.NoError(t, err) assert.True(t, exists, "k8s mode should always return true") }) return } - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() t.Run("existing", func(t *testing.T) { - exists, err := m.ImageExists(ctx, pc.Name, "alpine:latest") + exists, err := m.ImageExists(ctx, pc.TaiID, "alpine:latest") require.NoError(t, err) assert.True(t, exists) }) t.Run("missing", func(t *testing.T) { - exists, err := m.ImageExists(ctx, pc.Name, "nonexistent/image:no-such-tag-ever-12345") + exists, err := m.ImageExists(ctx, pc.TaiID, "nonexistent/image:no-such-tag-ever-12345") require.NoError(t, err) assert.False(t, exists) }) @@ -46,27 +46,27 @@ func TestImageExists(t *testing.T) { } func TestImagePull(t *testing.T) { - for _, pc := range testPools() { + for _, pc := range testNodes() { pc := pc t.Run(pc.Name, func(t *testing.T) { if pc.Name == "k8s" { t.Run("noop", func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - ch, err := m.PullImage(ctx, pc.Name, "alpine:latest", sandbox.ImagePullOptions{}) + ch, err := m.PullImage(ctx, pc.TaiID, "alpine:latest", sandbox.ImagePullOptions{}) require.NoError(t, err) assert.Nil(t, ch, "k8s mode should return nil channel") }) return } - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) defer cancel() t.Run("pull_with_progress", func(t *testing.T) { - ch, err := m.PullImage(ctx, pc.Name, "alpine:latest", sandbox.ImagePullOptions{}) + ch, err := m.PullImage(ctx, pc.TaiID, "alpine:latest", sandbox.ImagePullOptions{}) require.NoError(t, err) require.NotNil(t, ch) @@ -84,18 +84,18 @@ func TestImagePull(t *testing.T) { } func TestEnsureImage(t *testing.T) { - for _, pc := range testPools() { + for _, pc := range testNodes() { pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) defer cancel() - err := m.EnsureImage(ctx, pc.Name, "alpine:latest", sandbox.ImagePullOptions{}) + err := m.EnsureImage(ctx, pc.TaiID, "alpine:latest", sandbox.ImagePullOptions{}) require.NoError(t, err) if pc.Name != "k8s" { - exists, err := m.ImageExists(ctx, pc.Name, "alpine:latest") + exists, err := m.ImageExists(ctx, pc.TaiID, "alpine:latest") require.NoError(t, err) assert.True(t, exists) } @@ -104,17 +104,17 @@ func TestEnsureImage(t *testing.T) { } func TestEnsureImage_BadRef(t *testing.T) { - for _, pc := range testPools() { + for _, pc := range testNodes() { pc := pc if pc.Name == "k8s" { continue } t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - err := m.EnsureImage(ctx, pc.Name, "nonexistent/image:no-such-tag-ever-12345", sandbox.ImagePullOptions{}) + err := m.EnsureImage(ctx, pc.TaiID, "nonexistent/image:no-such-tag-ever-12345", sandbox.ImagePullOptions{}) assert.Error(t, err) }) } diff --git a/sandbox/v2/box_test.go b/sandbox/v2/box_test.go index 83f9677d..6870e164 100644 --- a/sandbox/v2/box_test.go +++ b/sandbox/v2/box_test.go @@ -13,10 +13,11 @@ import ( func TestBoxExec(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -35,10 +36,11 @@ func TestBoxExec(t *testing.T) { func TestBoxExecWithOptions(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ctx := context.Background() result, err := box.Exec(ctx, []string{"pwd"}, @@ -57,10 +59,11 @@ func TestBoxExecWithOptions(t *testing.T) { func TestBoxStream(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ctx := context.Background() stream, err := box.Stream(ctx, []string{"sh", "-c", "echo line1; echo line2"}) @@ -90,10 +93,11 @@ func TestBoxStream(t *testing.T) { func TestBoxWorkspace(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ws := box.Workspace() if ws == nil { @@ -131,10 +135,11 @@ func TestBoxWorkspace(t *testing.T) { func TestBoxInfo(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ctx := context.Background() info, err := box.Info(ctx) @@ -157,10 +162,11 @@ func TestBoxInfo(t *testing.T) { func TestBoxStopStart(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ctx := context.Background() if err := box.Stop(ctx); err != nil { @@ -185,14 +191,16 @@ func TestBoxStopStart(t *testing.T) { func TestBoxGetOrCreate(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) ctx := context.Background() box1, err := m.GetOrCreate(ctx, sandbox.CreateOptions{ - ID: "goc-" + pc.Name, - Image: testImage(), - Owner: "test-user", + ID: "goc-" + pc.Name, + Image: testImage(), + Owner: "test-user", + NodeID: pc.TaiID, }) if err != nil { t.Fatalf("GetOrCreate first: %v", err) @@ -200,9 +208,10 @@ func TestBoxGetOrCreate(t *testing.T) { defer m.Remove(ctx, box1.ID()) box2, err := m.GetOrCreate(ctx, sandbox.CreateOptions{ - ID: "goc-" + pc.Name, - Image: testImage(), - Owner: "test-user", + ID: "goc-" + pc.Name, + Image: testImage(), + Owner: "test-user", + NodeID: pc.TaiID, }) if err != nil { t.Fatalf("GetOrCreate second: %v", err) diff --git a/sandbox/v2/box_workspace_test.go b/sandbox/v2/box_workspace_test.go index 76c4ee96..6f5edcb8 100644 --- a/sandbox/v2/box_workspace_test.go +++ b/sandbox/v2/box_workspace_test.go @@ -15,20 +15,21 @@ import ( func TestWorkspaceID_Set(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - sbm, wsm := setupManagerWithWorkspace(t, pc) + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "test-ws", Owner: "user", Node: pc.Name, + Name: "test-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID }) @@ -40,10 +41,11 @@ func TestWorkspaceID_Set(t *testing.T) { func TestWorkspaceID_Empty(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) assert.Empty(t, box.WorkspaceID()) }) } @@ -52,24 +54,25 @@ func TestWorkspaceID_Empty(t *testing.T) { func TestWorkspace_NodeRouting(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - sbm, wsm := setupManagerWithWorkspace(t, pc) + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "routed-ws", Owner: "user", Node: pc.Name, + Name: "routed-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID }) - assert.Equal(t, pc.Name, box.Pool()) + assert.Equal(t, pc.TaiID, box.NodeID()) }) } } @@ -77,21 +80,30 @@ func TestWorkspace_NodeRouting(t *testing.T) { func TestWorkspace_InvalidID(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - sbm, _ := setupManagerWithWorkspace(t, pc) - ensureTestImage(t, sbm, pc.Name) + sbm, wsm := setupManagerWithWorkspace(t, &pc) + ensureTestImage(t, sbm, pc.TaiID) - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - _, err := sbm.Create(ctx, sandbox.CreateOptions{ + wsID := "nonexistent-workspace" + + box, err := sbm.Create(ctx, sandbox.CreateOptions{ Image: testImage(), Owner: "user", - WorkspaceID: "nonexistent-workspace", + WorkspaceID: wsID, }) - assert.Error(t, err) - assert.Contains(t, err.Error(), "resolve workspace") + + // With online nodes the manager auto-creates the workspace. + require.NoError(t, err) + require.NotNil(t, box) + defer box.Remove(context.Background()) + if wsm != nil { + defer wsm.Delete(context.Background(), wsID, true) + } }) } } @@ -99,21 +111,21 @@ func TestWorkspace_InvalidID(t *testing.T) { func TestWorkspace_BindMountLocal(t *testing.T) { skipIfNoDocker(t) - pc := poolConfig{Name: "local", Addr: testLocalAddr()} - sbm, wsm := setupManagerWithWorkspace(t, pc) + pc := nodeConfig{Name: "local", Addr: testLocalAddr()} + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "mount-ws", Owner: "user", Node: pc.Name, + Name: "mount-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) require.NoError(t, wsm.WriteFile(ctx, ws.ID, "seed.txt", []byte("hello from workspace"), 0644)) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID }) @@ -125,19 +137,19 @@ func TestWorkspace_BindMountLocal(t *testing.T) { func TestWorkspace_ContainerWriteBack(t *testing.T) { skipIfNoDocker(t) - pc := poolConfig{Name: "local", Addr: testLocalAddr()} - sbm, wsm := setupManagerWithWorkspace(t, pc) + pc := nodeConfig{Name: "local", Addr: testLocalAddr()} + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "writeback-ws", Owner: "user", Node: pc.Name, + Name: "writeback-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID }) @@ -152,21 +164,21 @@ func TestWorkspace_ContainerWriteBack(t *testing.T) { func TestWorkspace_ReadOnlyMount(t *testing.T) { skipIfNoDocker(t) - pc := poolConfig{Name: "local", Addr: testLocalAddr()} - sbm, wsm := setupManagerWithWorkspace(t, pc) + pc := nodeConfig{Name: "local", Addr: testLocalAddr()} + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "ro-ws", Owner: "user", Node: pc.Name, + Name: "ro-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) require.NoError(t, wsm.WriteFile(ctx, ws.ID, "readonly.txt", []byte("immutable"), 0644)) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID co.MountMode = "ro" }) @@ -185,21 +197,21 @@ func TestWorkspace_ReadOnlyMount(t *testing.T) { func TestWorkspace_CustomMountPath(t *testing.T) { skipIfNoDocker(t) - pc := poolConfig{Name: "local", Addr: testLocalAddr()} - sbm, wsm := setupManagerWithWorkspace(t, pc) + pc := nodeConfig{Name: "local", Addr: testLocalAddr()} + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "custom-path-ws", Owner: "user", Node: pc.Name, + Name: "custom-path-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) require.NoError(t, wsm.WriteFile(ctx, ws.ID, "data.json", []byte(`{"ok":true}`), 0644)) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID co.MountPath = "/data" }) @@ -212,27 +224,24 @@ func TestWorkspace_CustomMountPath(t *testing.T) { func TestWorkspace_BoxWorkspaceFS(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { if pc.Name == "local" { - // Local mode: sandbox and workspace use separate tai.Clients with - // different dataDirs, so Box.Workspace() writes to the sandbox volume - // while wsm reads from the workspace volume. Bind mount tests cover - // local workspace I/O end-to-end instead. continue } + pc := pc t.Run(pc.Name, func(t *testing.T) { - sbm, wsm := setupManagerWithWorkspace(t, pc) + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "fs-ws", Owner: "user", Node: pc.Name, + Name: "fs-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID }) @@ -253,24 +262,24 @@ func TestWorkspace_BoxWorkspaceFS(t *testing.T) { func TestWorkspace_LabelPersistence(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - sbm, wsm := setupManagerWithWorkspace(t, pc) + sbm, wsm := setupManagerWithWorkspace(t, &pc) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() ws, err := wsm.Create(ctx, workspace.CreateOptions{ - Name: "label-ws", Owner: "user", Node: pc.Name, + Name: "label-ws", Owner: "user", Node: pc.TaiID, }) require.NoError(t, err) defer wsm.Delete(context.Background(), ws.ID, true) - box := createTestBox(t, sbm, func(co *sandbox.CreateOptions) { + box := createTestBox(t, sbm, pc, func(co *sandbox.CreateOptions) { co.WorkspaceID = ws.ID }) - // WorkspaceID getter should reflect what was set assert.Equal(t, ws.ID, box.WorkspaceID()) // Container should also carry the label (verify via exec reading env or diff --git a/sandbox/v2/config.go b/sandbox/v2/config.go deleted file mode 100644 index 8cc9973c..00000000 --- a/sandbox/v2/config.go +++ /dev/null @@ -1,5 +0,0 @@ -package sandbox - -type Config struct { - Pool []Pool -} diff --git a/sandbox/v2/docs/API.md b/sandbox/v2/docs/API.md index 1ce2b865..32f81b73 100644 --- a/sandbox/v2/docs/API.md +++ b/sandbox/v2/docs/API.md @@ -2,7 +2,7 @@ Package: `github.com/yaoapp/yao/sandbox/v2` -Sandbox V2 manages sandboxes through a pool of Tai nodes. Two primary abstractions: +Sandbox V2 manages sandboxes through a set of Tai nodes. Two primary abstractions: - **Box** — a container (Docker or K8s pod). Created via `Manager.Create`. - **Host** — the Tai host machine itself. Obtained via `Manager.Host` (no Create needed). @@ -16,25 +16,14 @@ Supports workspace mounting, VNC, WebSocket proxying, and HostExec. ### Init ```go -func Init(cfg Config) error +func Init() ``` Initializes the global Manager singleton. Must be called once at startup. +No configuration is needed — node discovery is handled by `tai/registry`. ```go -err := sandbox.Init(sandbox.Config{ - Pool: []sandbox.Pool{ - { - Name: "docker", - Addr: "tai://192.168.1.10:19100", - MaxPerUser: 5, - MaxTotal: 20, - IdleTimeout: 30 * time.Minute, - MaxLifetime: 24 * time.Hour, - StopTimeout: 5 * time.Second, - }, - }, -}) +sandbox.Init() ``` ### M @@ -51,28 +40,12 @@ mgr := sandbox.M() --- -## Config +## Node Discovery -```go -type Config struct { - Pool []Pool -} -``` - -### Pool - -```go -type Pool struct { - Name string - Addr string // "tai://host:port", "tunnel://host:port", or Docker socket - Options []tai.Option // tai.Client options - MaxPerUser int // 0 = unlimited - MaxTotal int // 0 = unlimited - IdleTimeout time.Duration // 0 = no idle cleanup - MaxLifetime time.Duration // 0 = no max lifetime - StopTimeout time.Duration // SIGTERM grace period; 0 = DefaultStopTimeout (2s) -} -``` +Sandbox V2 no longer uses static node configuration. Nodes are discovered dynamically +through `tai/registry`. Each Tai node registers itself with a unique **TaiID** (e.g. +`"192.168.1.10-19100"` for direct mode, `"local"` for Docker). The TaiID is used as the +`NodeID` identifier in `CreateOptions`, `ListOptions`, `Host()`, `ImageExists()`, etc. --- @@ -99,7 +72,7 @@ const ( func (m *Manager) Start(ctx context.Context) error ``` -Recovers existing containers from all pools and starts the background cleanup loop (1 min interval). +Recovers existing containers from all nodes and starts the background cleanup loop (1 min interval). ```go ctx := context.Background() @@ -112,7 +85,7 @@ err := sandbox.M().Start(ctx) func (m *Manager) Close() error ``` -Stops the cleanup loop and closes all pool connections. +Stops the cleanup loop and closes all node connections. ### Create @@ -126,7 +99,7 @@ Creates and starts a new sandbox container. Returns a `Box` handle. box, err := sandbox.M().Create(ctx, sandbox.CreateOptions{ Image: "alpine:latest", Owner: "user-123", - Pool: "docker", + NodeID: "192.168.1.10-19100", // TaiID from registry Policy: sandbox.Session, WorkDir: "/workspace", Env: map[string]string{"LANG": "en_US.UTF-8"}, @@ -148,15 +121,16 @@ box, err := sandbox.M().Create(ctx, sandbox.CreateOptions{ ### Host ```go -func (m *Manager) Host(ctx context.Context, pool string) (*Host, error) +func (m *Manager) Host(ctx context.Context, nodeID string) (*Host, error) ``` -Returns a `Host` handle for the given pool. Unlike `Create`, no container is provisioned — -the Host is available as long as the pool's Tai server reports `host_exec` capability. -Returns `ErrPoolNotFound` if the pool does not exist, or an error if the pool has no `host_exec`. +Returns a `Host` handle for the given node (identified by TaiID). Unlike `Create`, no +container is provisioned — the Host is available as long as the Tai server reports +`host_exec` capability. Returns `ErrNodeNotFound` if the TaiID is not registered, +`ErrNodeMissing` if the nodeID argument is empty, or an error if the node has no `host_exec`. ```go -host, err := sandbox.M().Host(ctx, "remote") +host, err := sandbox.M().Host(ctx, "192.168.1.10-19100") ``` ### Get @@ -198,7 +172,7 @@ Returns all sandboxes matching the given filters. Empty fields = no filter. ```go boxes, err := sandbox.M().List(ctx, sandbox.ListOptions{ Owner: "user-123", - Pool: "docker", + NodeID: "192.168.1.10-19100", Labels: map[string]string{"project": "demo"}, }) ``` @@ -209,7 +183,7 @@ boxes, err := sandbox.M().List(ctx, sandbox.ListOptions{ func (m *Manager) Remove(ctx context.Context, id string) error ``` -Force-removes a sandbox (SIGKILL + delete). Revokes container tokens. +Force-removes a sandbox (SIGKILL + delete). ```go err := sandbox.M().Remove(ctx, "sb-12345") @@ -236,90 +210,48 @@ Updates a sandbox's last-active timestamp. Called by the gRPC heartbeat service. err := sandbox.M().Heartbeat("sb-12345", true, 3) ``` -### AddPool +### Nodes ```go -func (m *Manager) AddPool(ctx context.Context, p Pool) error +func (m *Manager) Nodes() []registry.NodeSnapshot ``` -Registers a new pool at runtime. +Returns all registered Tai nodes from the `tai/registry`. ```go -err := sandbox.M().AddPool(ctx, sandbox.Pool{ - Name: "k8s-gpu", - Addr: "tai://10.0.0.5:19100", - MaxTotal: 10, -}) -``` - -### RemovePool - -```go -func (m *Manager) RemovePool(ctx context.Context, name string, force bool) error -``` - -Removes a pool. Returns `ErrPoolInUse` if the pool has running boxes and `force=false`. -With `force=true`, all boxes in the pool are removed first. - -### Pools - -```go -func (m *Manager) Pools() []PoolInfo -``` - -Returns all registered pools and their status. - -```go -for _, p := range sandbox.M().Pools() { - fmt.Printf("pool=%s addr=%s connected=%v boxes=%d\n", - p.Name, p.Addr, p.Connected, p.Boxes) +for _, n := range sandbox.M().Nodes() { + fmt.Printf("tai_id=%s mode=%s addr=%s status=%s\n", + n.TaiID, n.Mode, n.Addr, n.Status) } ``` -### SetGRPCPort - -```go -func (m *Manager) SetGRPCPort(port int) -``` - -Sets the local gRPC port injected into container env vars (`YAO_GRPC_ADDR`). Default: `9099`. - -### SetWorkspaceManager - -```go -func (m *Manager) SetWorkspaceManager(wm *workspace.Manager) -``` - -Links the workspace manager. When `CreateOptions.WorkspaceID` is set, the Manager uses it -to resolve the workspace's bound node and route the container to the correct pool. - ### ImageExists ```go -func (m *Manager) ImageExists(ctx context.Context, pool, ref string) (bool, error) +func (m *Manager) ImageExists(ctx context.Context, nodeID, ref string) (bool, error) ``` -Reports whether the given image ref exists on the target pool node. -Returns `(true, nil)` when the pool has no image service (e.g. K8s — kubelet handles pulls). +Reports whether the given image ref exists on the target node. +Returns `(true, nil)` when the node has no image service (e.g. K8s — kubelet handles pulls). ```go -exists, err := sandbox.M().ImageExists(ctx, "docker", "alpine:latest") +exists, err := sandbox.M().ImageExists(ctx, "192.168.1.10-19100", "alpine:latest") ``` ### PullImage ```go -func (m *Manager) PullImage(ctx context.Context, pool, ref string, opts ImagePullOptions) (<-chan taisandbox.PullProgress, error) +func (m *Manager) PullImage(ctx context.Context, nodeID, ref string, opts ImagePullOptions) (<-chan taisandbox.PullProgress, error) ``` -Pulls an image to the target pool node. Returns a channel of `taisandbox.PullProgress` -(from `github.com/yaoapp/yao/tai/sandbox`). Returns `(nil, nil)` when the pool has no image +Pulls an image to the target node. Returns a channel of `taisandbox.PullProgress` +(from `github.com/yaoapp/yao/tai/sandbox`). Returns `(nil, nil)` when the node has no image service (e.g. K8s). `PullProgress` fields: `Status string`, `Layer string`, `Current int64`, `Total int64`, `Error string`. ```go -ch, err := sandbox.M().PullImage(ctx, "docker", "myapp:v2", sandbox.ImagePullOptions{ +ch, err := sandbox.M().PullImage(ctx, "192.168.1.10-19100", "myapp:v2", sandbox.ImagePullOptions{ Auth: &sandbox.RegistryAuth{ Username: "user", Password: "pass", @@ -334,13 +266,13 @@ for p := range ch { ### EnsureImage ```go -func (m *Manager) EnsureImage(ctx context.Context, pool, ref string, opts ImagePullOptions) error +func (m *Manager) EnsureImage(ctx context.Context, nodeID, ref string, opts ImagePullOptions) error ``` Checks if the image exists; if not, pulls it and blocks until complete. ```go -err := sandbox.M().EnsureImage(ctx, "docker", "alpine:latest", sandbox.ImagePullOptions{}) +err := sandbox.M().EnsureImage(ctx, "192.168.1.10-19100", "alpine:latest", sandbox.ImagePullOptions{}) ``` --- @@ -355,7 +287,7 @@ A `Box` is a handle to a running sandbox container. func (b *Box) ID() string func (b *Box) Owner() string func (b *Box) ContainerID() string -func (b *Box) Pool() string +func (b *Box) NodeID() string func (b *Box) WorkspaceID() string ``` @@ -491,12 +423,12 @@ fmt.Printf("status=%s processes=%d vnc=%v created=%s\n", ## Host A `Host` represents a Tai host machine execution environment, distinct from `Box` (containers). -No `Create` call is needed — a Host is available as long as the pool's Tai server reports `host_exec`. +No `Create` call is needed — a Host is available as long as the node's Tai server reports `host_exec`. ### Accessors ```go -func (h *Host) Pool() string +func (h *Host) NodeID() string ``` ### Exec @@ -508,7 +440,7 @@ func (h *Host) Exec(ctx context.Context, cmd string, args []string, opts ...Host Runs a command directly on the Tai host machine via HostExec gRPC. ```go -host, _ := sandbox.M().Host(ctx, "remote") +host, _ := sandbox.M().Host(ctx, "192.168.1.10-19100") result, err := host.Exec(ctx, "git", []string{"status"}, sandbox.WithHostWorkDir("/data/repos/project"), sandbox.WithHostEnv(map[string]string{"GIT_AUTHOR_NAME": "bot"}), @@ -529,7 +461,7 @@ Runs a command on the Tai host and streams stdout/stderr in real time via HostEx ExecStream. Returns a `HostExecStream` with separate channels for stdout and stderr. ```go -host, _ := sandbox.M().Host(ctx, "remote") +host, _ := sandbox.M().Host(ctx, "192.168.1.10-19100") stream, err := host.Stream(ctx, "tail", []string{"-f", "/var/log/app.log"}, sandbox.WithHostWorkDir("/data"), sandbox.WithHostTimeout(60000), @@ -607,7 +539,7 @@ type CreateOptions struct { ID string Owner string Labels map[string]string - Pool string // empty = default pool + NodeID string // TaiID from registry (required unless WorkspaceID routes to a node) Image string // required WorkDir string // default "/workspace" User string // container user @@ -617,8 +549,9 @@ type CreateOptions struct { VNC bool Ports []PortMapping Policy LifecyclePolicy // default Session - IdleTimeout time.Duration // overrides pool default - StopTimeout time.Duration // overrides pool default + IdleTimeout time.Duration // 0 = no idle cleanup + MaxLifetime time.Duration // 0 = no max lifetime + StopTimeout time.Duration // SIGTERM grace period; 0 = DefaultStopTimeout (2s) WorkspaceID string // workspace to mount; empty = none MountMode string // "rw" (default) or "ro" MountPath string // default "/workspace" @@ -630,7 +563,7 @@ type CreateOptions struct { ```go type ListOptions struct { Owner string - Pool string + NodeID string Labels map[string]string } ``` @@ -686,7 +619,7 @@ type ServiceConn struct { type BoxInfo struct { ID string ContainerID string - Pool string + NodeID string Owner string Status string // "running", "stopped", etc. Policy LifecyclePolicy @@ -699,21 +632,6 @@ type BoxInfo struct { } ``` -### PoolInfo - -```go -type PoolInfo struct { - Name string - Addr string - Connected bool - Boxes int - MaxPerUser int - MaxTotal int - IdleTimeout time.Duration - MaxLifetime time.Duration -} -``` - ### ImagePullOptions / RegistryAuth ```go @@ -758,11 +676,10 @@ type HostExecStream struct { ```go var ( - ErrNotAvailable = errors.New("sandbox: not available (no pools configured)") - ErrNotFound = errors.New("sandbox: not found") - ErrLimitExceeded = errors.New("sandbox: limit exceeded") - ErrPoolNotFound = errors.New("sandbox: pool not found") - ErrPoolInUse = errors.New("sandbox: pool has running boxes") + ErrNotAvailable = errors.New("sandbox: not available (no nodes registered)") + ErrNotFound = errors.New("sandbox: not found") + ErrNodeNotFound = errors.New("sandbox: node not found") + ErrNodeMissing = errors.New("sandbox: node ID is required") ) ``` @@ -770,38 +687,28 @@ var ( ## Helper Functions -### CreateContainerTokens - -```go -func CreateContainerTokens(sandboxID, owner string, scopes []string) (access, refresh string, err error) -``` - -Creates an OAuth token pair for a sandbox container. - -### RevokeContainerTokens - -```go -func RevokeContainerTokens(refresh string) error -``` - -Revokes a container refresh token. - ### BuildGRPCEnv ```go -func BuildGRPCEnv(pool *Pool, sandboxID, access, refresh string, grpcPort int) map[string]string +func BuildGRPCEnv(mode, addr, sandboxID string) map[string]string ``` -Builds environment variables injected into sandbox containers: +Builds environment variables injected into sandbox containers. The gRPC port is read from +`config.Conf.GRPC.Port` (defaults to `9099`). -| Variable | Description | -|--------------------|--------------------------------------| -| `YAO_SANDBOX_ID` | Sandbox identifier | -| `YAO_TOKEN` | Access token for gRPC auth | -| `YAO_REFRESH_TOKEN` | Refresh token for token rotation | -| `YAO_GRPC_ADDR` | gRPC server address (auto-derived) | +- `mode` — the `TaiNode.Mode` (`"local"`, `"direct"`, `"tunnel"`) +- `addr` — the `TaiNode.Addr` (e.g. `"tai://192.168.1.10:19100"` for direct mode) +- `sandboxID` — the container's sandbox identifier + +| Variable | Description | +|------------------|------------------------------------| +| `YAO_SANDBOX_ID` | Sandbox identifier | +| `YAO_GRPC_ADDR` | gRPC server address (auto-derived) | Address derivation logic: -- `tai://host:port` → `host:port` (default port 19100 when omitted) -- `tunnel://...` → `127.0.0.1:` -- Local/default → `127.0.0.1:` +- `local` → `host.docker.internal:` +- `direct` with `tai://host:port` → `host:port` +- `tunnel` → `127.0.0.1:` + +Token injection (`YAO_TOKEN`, `YAO_REFRESH_TOKEN`) is the **caller's responsibility** via +`CreateOptions.Env`. See IMPL.md "OAuth Decoupling" for details. diff --git a/sandbox/v2/errors.go b/sandbox/v2/errors.go index deda6ea7..55c6285f 100644 --- a/sandbox/v2/errors.go +++ b/sandbox/v2/errors.go @@ -3,9 +3,8 @@ package sandbox import "errors" var ( - ErrNotAvailable = errors.New("sandbox: not available (no pools configured)") - ErrNotFound = errors.New("sandbox: not found") - ErrLimitExceeded = errors.New("sandbox: limit exceeded") - ErrPoolNotFound = errors.New("sandbox: pool not found") - ErrPoolInUse = errors.New("sandbox: pool has running boxes") + ErrNotAvailable = errors.New("sandbox: not available (no nodes registered)") + ErrNotFound = errors.New("sandbox: not found") + ErrNodeNotFound = errors.New("sandbox: node not found") + ErrNodeMissing = errors.New("sandbox: node ID is required") ) diff --git a/sandbox/v2/grpc.go b/sandbox/v2/grpc.go index 5f99eab7..e3d85399 100644 --- a/sandbox/v2/grpc.go +++ b/sandbox/v2/grpc.go @@ -1,63 +1,43 @@ package sandbox import ( - "crypto/rand" - "encoding/hex" "fmt" "net/url" "strconv" - "strings" + + "github.com/yaoapp/yao/config" ) -func createToken() (string, error) { - b := make([]byte, 32) - if _, err := rand.Read(b); err != nil { - return "", err +// BuildGRPCEnv builds the gRPC environment variables for a sandbox container +// based on the Tai node's mode and address from the registry. +// +// mode is the TaiNode.Mode ("local", "direct", "tunnel"). +// addr is the TaiNode.Addr (e.g. "tai://host:port" for direct mode). +// sandboxID is the container's sandbox identifier. +// +// The Yao gRPC port is read from config.Conf.GRPC.Port. +func BuildGRPCEnv(mode, addr, sandboxID string) map[string]string { + grpcPort := config.Conf.GRPC.Port + if grpcPort == 0 { + grpcPort = 9099 } - return hex.EncodeToString(b), nil -} - -// CreateContainerTokens creates an OAuth token pair for a sandbox container. -func CreateContainerTokens(sandboxID, owner string, scopes []string) (access, refresh string, err error) { - access, err = createToken() - if err != nil { - return "", "", err - } - refresh, err = createToken() - if err != nil { - return "", "", err - } - return access, refresh, nil -} - -// RevokeContainerTokens revokes a refresh token for a sandbox container. -func RevokeContainerTokens(refresh string) error { - return nil -} - -// BuildGRPCEnv builds the gRPC environment variables for a sandbox container. -// Supports tai:// (direct), tunnel:// (NAT traversal), and local modes. -func BuildGRPCEnv(pool *Pool, sandboxID, access, refresh string, grpcPort int) map[string]string { portStr := strconv.Itoa(grpcPort) + env := map[string]string{ - "YAO_SANDBOX_ID": sandboxID, - "YAO_TOKEN": access, - "YAO_REFRESH_TOKEN": refresh, + "YAO_SANDBOX_ID": sandboxID, } - if pool == nil { + switch mode { + case "local": + env["YAO_GRPC_ADDR"] = fmt.Sprintf("host.docker.internal:%s", portStr) + + case "tunnel": env["YAO_GRPC_ADDR"] = fmt.Sprintf("127.0.0.1:%s", portStr) - return env - } - switch { - case strings.HasPrefix(pool.Addr, "tunnel://"): - env["YAO_GRPC_ADDR"] = fmt.Sprintf("127.0.0.1:%d", grpcPort) - - case strings.HasPrefix(pool.Addr, "tai://"): - u, err := url.Parse(pool.Addr) - if err != nil { - env["YAO_GRPC_ADDR"] = fmt.Sprintf("127.0.0.1:%s", portStr) + case "direct": + u, err := url.Parse(addr) + if err != nil || u.Hostname() == "" { + env["YAO_GRPC_ADDR"] = fmt.Sprintf("host.docker.internal:%s", portStr) return env } taiHost := u.Hostname() @@ -68,7 +48,7 @@ func BuildGRPCEnv(pool *Pool, sandboxID, access, refresh string, grpcPort int) m env["YAO_GRPC_ADDR"] = fmt.Sprintf("%s:%s", taiHost, taiPort) default: - env["YAO_GRPC_ADDR"] = fmt.Sprintf("127.0.0.1:%s", portStr) + env["YAO_GRPC_ADDR"] = fmt.Sprintf("host.docker.internal:%s", portStr) } return env } diff --git a/sandbox/v2/grpc_test.go b/sandbox/v2/grpc_test.go index b9905988..800a2cae 100644 --- a/sandbox/v2/grpc_test.go +++ b/sandbox/v2/grpc_test.go @@ -3,27 +3,28 @@ package sandbox_test import ( "testing" + "github.com/yaoapp/yao/config" sandbox "github.com/yaoapp/yao/sandbox/v2" ) func TestBuildGRPCEnvLocal(t *testing.T) { - pool := &sandbox.Pool{Name: "local", Addr: "local"} - env := sandbox.BuildGRPCEnv(pool, "sb-001", "access-tok", "refresh-tok", 9099) + config.Conf.GRPC.Port = 9099 + env := sandbox.BuildGRPCEnv("local", "", "sb-001") if env["YAO_SANDBOX_ID"] != "sb-001" { t.Errorf("YAO_SANDBOX_ID = %q", env["YAO_SANDBOX_ID"]) } - if env["YAO_TOKEN"] != "access-tok" { - t.Errorf("YAO_TOKEN = %q", env["YAO_TOKEN"]) + if _, ok := env["YAO_TOKEN"]; ok { + t.Error("YAO_TOKEN should not be set by BuildGRPCEnv") } - if env["YAO_GRPC_ADDR"] != "127.0.0.1:9099" { - t.Errorf("YAO_GRPC_ADDR = %q", env["YAO_GRPC_ADDR"]) + if env["YAO_GRPC_ADDR"] != "host.docker.internal:9099" { + t.Errorf("YAO_GRPC_ADDR = %q, want host.docker.internal:9099", env["YAO_GRPC_ADDR"]) } } -func TestBuildGRPCEnvRemote(t *testing.T) { - pool := &sandbox.Pool{Name: "gpu", Addr: "tai://gpu-server"} - env := sandbox.BuildGRPCEnv(pool, "sb-002", "access", "refresh", 9099) +func TestBuildGRPCEnvDirect(t *testing.T) { + config.Conf.GRPC.Port = 9099 + env := sandbox.BuildGRPCEnv("direct", "tai://gpu-server", "sb-002") if env["YAO_GRPC_ADDR"] != "gpu-server:19100" { t.Errorf("YAO_GRPC_ADDR = %q, want gpu-server:19100", env["YAO_GRPC_ADDR"]) @@ -31,26 +32,10 @@ func TestBuildGRPCEnvRemote(t *testing.T) { } func TestBuildGRPCEnvTunnel(t *testing.T) { - pool := &sandbox.Pool{Name: "tunnel", Addr: "tunnel://relay.example.com"} - env := sandbox.BuildGRPCEnv(pool, "sb-003", "access", "refresh", 9099) + config.Conf.GRPC.Port = 9099 + env := sandbox.BuildGRPCEnv("tunnel", "tunnel://relay.example.com", "sb-003") if env["YAO_GRPC_ADDR"] != "127.0.0.1:9099" { t.Errorf("YAO_GRPC_ADDR = %q, want 127.0.0.1:9099", env["YAO_GRPC_ADDR"]) } } - -func TestCreateContainerTokens(t *testing.T) { - access, refresh, err := sandbox.CreateContainerTokens("sb-001", "user1", nil) - if err != nil { - t.Fatalf("CreateContainerTokens: %v", err) - } - if len(access) != 64 { - t.Errorf("access token len = %d, want 64 hex chars", len(access)) - } - if len(refresh) != 64 { - t.Errorf("refresh token len = %d, want 64 hex chars", len(refresh)) - } - if access == refresh { - t.Error("access and refresh tokens should be different") - } -} diff --git a/sandbox/v2/host.go b/sandbox/v2/host.go index 89ccff4d..125388c3 100644 --- a/sandbox/v2/host.go +++ b/sandbox/v2/host.go @@ -1,8 +1,10 @@ package sandbox import ( + "bytes" "context" "fmt" + "io" hepb "github.com/yaoapp/yao/tai/hostexec/pb" "github.com/yaoapp/yao/tai/workspace" @@ -12,54 +14,78 @@ import ( // Unlike Box (which wraps a container), Host executes commands directly on // the Tai server's OS via HostExec gRPC and accesses files via Volume gRPC. // -// A Host is bound to a pool and does not require Create — it is available as -// long as the pool's Tai server reports host_exec capability. +// Host implements the Computer interface. type Host struct { - pool string - manager *Manager + nodeID string + workplaceID string + system SystemInfo + manager *Manager } -// Pool returns the pool name this Host belongs to. -func (h *Host) Pool() string { return h.pool } +// Compile-time check: *Host implements Computer. +var _ Computer = (*Host)(nil) + +// ComputerInfo returns identity and registry information for the host. +// Registry-level details (TaiID, System, etc.) are populated when the node +// is backed by a registered Tai node; otherwise only Kind and NodeID are set. +func (h *Host) ComputerInfo() ComputerInfo { + return ComputerInfo{ + Kind: "host", + NodeID: h.nodeID, + System: h.system, + Status: "online", + } +} // Exec runs a command on the Tai host machine via HostExec gRPC. -func (h *Host) Exec(ctx context.Context, cmd string, args []string, opts ...HostExecOption) (*HostExecResult, error) { - client, err := h.manager.getPool(h.pool) +// cmd[0] is the program, cmd[1:] are arguments. +func (h *Host) Exec(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecResult, error) { + if len(cmd) == 0 { + return nil, fmt.Errorf("sandbox: empty command") + } + + client, err := h.manager.getNode(h.nodeID) if err != nil { return nil, err } he := client.HostExec() if he == nil { - return nil, fmt.Errorf("sandbox: host_exec not available on pool %q", h.pool) + return nil, fmt.Errorf("sandbox: host_exec not available on node %q", h.nodeID) } - cfg := &hostExecConfig{} + cfg := &execConfig{} for _, o := range opts { o(cfg) } req := &hepb.ExecRequest{ - Command: cmd, - Args: args, - WorkingDir: cfg.WorkDir, - Stdin: cfg.Stdin, - TimeoutMs: cfg.TimeoutMs, - MaxOutputBytes: cfg.MaxOutputBytes, + Command: cmd[0], + Args: cmd[1:], + Stdin: cfg.Stdin, + } + if cfg.WorkDir != "" { + req.WorkingDir = cfg.WorkDir } if cfg.Env != nil { req.Env = cfg.Env } + if cfg.Timeout > 0 { + req.TimeoutMs = cfg.Timeout.Milliseconds() + } + if cfg.MaxOutputBytes > 0 { + req.MaxOutputBytes = cfg.MaxOutputBytes + } resp, err := he.Exec(ctx, req) if err != nil { return nil, fmt.Errorf("hostexec rpc: %w", err) } - return &HostExecResult{ + return &ExecResult{ ExitCode: int(resp.ExitCode), - Stdout: resp.Stdout, - Stderr: resp.Stderr, + Stdout: string(resp.Stdout), + Stderr: string(resp.Stderr), DurationMs: resp.DurationMs, Error: resp.Error, Truncated: resp.Truncated, @@ -67,35 +93,45 @@ func (h *Host) Exec(ctx context.Context, cmd string, args []string, opts ...Host } // Stream runs a command on the Tai host and streams stdout/stderr in real time -// via HostExec gRPC ExecStream. Returns a HostExecStream with separate channels -// for stdout and stderr, plus Wait (blocks until exit) and Cancel. -func (h *Host) Stream(ctx context.Context, cmd string, args []string, opts ...HostExecOption) (*HostExecStream, error) { - client, err := h.manager.getPool(h.pool) +// via HostExec gRPC ExecStream. Returns a unified ExecStream with io.ReadCloser +// for stdout/stderr. +func (h *Host) Stream(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecStream, error) { + if len(cmd) == 0 { + return nil, fmt.Errorf("sandbox: empty command") + } + + client, err := h.manager.getNode(h.nodeID) if err != nil { return nil, err } he := client.HostExec() if he == nil { - return nil, fmt.Errorf("sandbox: host_exec not available on pool %q", h.pool) + return nil, fmt.Errorf("sandbox: host_exec not available on node %q", h.nodeID) } - cfg := &hostExecConfig{} + cfg := &execConfig{} for _, o := range opts { o(cfg) } req := &hepb.ExecRequest{ - Command: cmd, - Args: args, - WorkingDir: cfg.WorkDir, - Stdin: cfg.Stdin, - TimeoutMs: cfg.TimeoutMs, - MaxOutputBytes: cfg.MaxOutputBytes, + Command: cmd[0], + Args: cmd[1:], + Stdin: cfg.Stdin, + } + if cfg.WorkDir != "" { + req.WorkingDir = cfg.WorkDir } if cfg.Env != nil { req.Env = cfg.Env } + if cfg.Timeout > 0 { + req.TimeoutMs = cfg.Timeout.Milliseconds() + } + if cfg.MaxOutputBytes > 0 { + req.MaxOutputBytes = cfg.MaxOutputBytes + } streamCtx, cancel := context.WithCancel(ctx) rpcStream, err := he.ExecStream(streamCtx, req) @@ -104,15 +140,15 @@ func (h *Host) Stream(ctx context.Context, cmd string, args []string, opts ...Ho return nil, fmt.Errorf("hostexec stream rpc: %w", err) } - stdoutCh := make(chan []byte, 64) - stderrCh := make(chan []byte, 64) + stdoutR, stdoutW := io.Pipe() + stderrR, stderrW := io.Pipe() doneCh := make(chan struct{}) var exitCode int var exitErr error go func() { - defer close(stdoutCh) - defer close(stderrCh) + defer stdoutW.Close() + defer stderrW.Close() defer close(doneCh) for { msg, err := rpcStream.Recv() @@ -123,9 +159,9 @@ func (h *Host) Stream(ctx context.Context, cmd string, args []string, opts ...Ho if len(msg.Data) > 0 { switch msg.Stream { case hepb.ExecOutput_STDOUT: - stdoutCh <- msg.Data + stdoutW.Write(msg.Data) case hepb.ExecOutput_STDERR: - stderrCh <- msg.Data + stderrW.Write(msg.Data) } } if msg.Done { @@ -138,9 +174,10 @@ func (h *Host) Stream(ctx context.Context, cmd string, args []string, opts ...Ho } }() - return &HostExecStream{ - Stdout: stdoutCh, - Stderr: stderrCh, + return &ExecStream{ + Stdout: stdoutR, + Stderr: stderrR, + Stdin: nopWriteCloser{&bytes.Buffer{}}, Wait: func() (int, error) { <-doneCh return exitCode, exitErr @@ -149,13 +186,48 @@ func (h *Host) Stream(ctx context.Context, cmd string, args []string, opts ...Ho }, nil } -// Workspace returns a filesystem interface for the given session on the host. -// The sessionID typically corresponds to a workspace ID; files are stored -// under dataDir/{sessionID}/ on the Tai host, accessed via Volume gRPC. -func (h *Host) Workspace(sessionID string) workspace.FS { - client, err := h.manager.getPool(h.pool) +// VNC returns the VNC WebSocket URL for the Tai host machine. +// Uses the special __host__ identifier to route to localhost:5900 on the Tai server. +func (h *Host) VNC(ctx context.Context) (string, error) { + client, err := h.manager.getNode(h.nodeID) + if err != nil { + return "", err + } + return client.VNC().URL(ctx, "__host__") +} + +// Proxy returns the HTTP URL for a service running on the Tai host machine. +// Uses the special __host__ identifier to route to localhost:{port} on the Tai server. +func (h *Host) Proxy(ctx context.Context, port int, path string) (string, error) { + client, err := h.manager.getNode(h.nodeID) + if err != nil { + return "", err + } + return client.Proxy().URL(ctx, "__host__", port, path) +} + +// BindWorkplace binds a workspace to this host by ID. Subsequent calls to +// Workplace() will return the FS for this workspace. Call again to rebind. +func (h *Host) BindWorkplace(workspaceID string) { + h.workplaceID = workspaceID +} + +// Workplace returns the workspace FS bound to this host, or nil if unbound. +func (h *Host) Workplace() workspace.FS { + if h.workplaceID == "" { + return nil + } + client, err := h.manager.getNode(h.nodeID) if err != nil { return nil } - return client.Workspace(sessionID) + return client.Workspace(h.workplaceID) } + +// NodeID returns the node ID this Host belongs to. +func (h *Host) NodeID() string { return h.nodeID } + +// nopWriteCloser wraps an io.Writer with a no-op Close. +type nopWriteCloser struct{ io.Writer } + +func (nopWriteCloser) Close() error { return nil } diff --git a/sandbox/v2/host_test.go b/sandbox/v2/host_test.go index de5d090d..303feca6 100644 --- a/sandbox/v2/host_test.go +++ b/sandbox/v2/host_test.go @@ -3,6 +3,7 @@ package sandbox_test import ( "context" "fmt" + "io" "strings" "testing" "time" @@ -11,16 +12,11 @@ import ( "github.com/yaoapp/yao/tai" ) -func setupHostManager(t *testing.T, tgt hostExecTarget) *sandbox.Manager { +func setupHostManager(t *testing.T, tgt *hostExecTarget) *sandbox.Manager { t.Helper() addr := fmt.Sprintf("tai://%s", tgt.Addr) - pool := sandbox.Pool{Name: tgt.Name, Addr: addr} - cfg := sandbox.Config{Pool: []sandbox.Pool{pool}} - if err := sandbox.Init(cfg); err != nil { - t.Fatalf("Init: %v", err) - } - m := sandbox.M() - t.Cleanup(func() { m.Close() }) + m, nodes := setupManager(t, nodeConfig{Name: tgt.Name, Addr: addr}) + tgt.TaiID = nodes[0].TaiID return m } @@ -28,10 +24,11 @@ func TestHost_Exec_Echo(t *testing.T) { skipIfNoHostExec(t) for _, tgt := range hostExecTargets() { + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } @@ -39,8 +36,8 @@ func TestHost_Exec_Echo(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - cmd, args := linuxCmd(tgt, "echo", "hello", "from", "host") - result, err := host.Exec(ctx, cmd, args) + cmd := hostCmd(tgt, "echo", "hello", "from", "host") + result, err := host.Exec(ctx, cmd) if err != nil { t.Fatalf("Exec: %v", err) } @@ -53,7 +50,7 @@ func TestHost_Exec_Echo(t *testing.T) { if result.ExitCode != 0 { t.Errorf("exit_code = %d, want 0", result.ExitCode) } - got := strings.TrimSpace(string(result.Stdout)) + got := strings.TrimSpace(result.Stdout) if !strings.Contains(got, "hello") { t.Errorf("stdout = %q, want contains 'hello'", got) } @@ -65,10 +62,11 @@ func TestHost_Exec_Env(t *testing.T) { skipIfNoHostExec(t) for _, tgt := range hostExecTargets() { + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } @@ -76,17 +74,14 @@ func TestHost_Exec_Env(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - var cmd string - var args []string + var cmd []string if tgt.IsWinNative { - cmd = "cmd.exe" - args = []string{"/c", "echo", "%MY_VAR%"} + cmd = []string{"cmd.exe", "/c", "echo", "%MY_VAR%"} } else { - cmd = "sh" - args = []string{"-c", "echo $MY_VAR"} + cmd = []string{"sh", "-c", "echo $MY_VAR"} } - result, err := host.Exec(ctx, cmd, args, sandbox.WithHostEnv(map[string]string{"MY_VAR": "host_test_value"})) + result, err := host.Exec(ctx, cmd, sandbox.WithEnv(map[string]string{"MY_VAR": "host_test_value"})) if err != nil { t.Fatalf("Exec: %v", err) } @@ -96,7 +91,7 @@ func TestHost_Exec_Env(t *testing.T) { } t.Fatalf("error: %s", result.Error) } - got := strings.TrimSpace(string(result.Stdout)) + got := strings.TrimSpace(result.Stdout) if !strings.Contains(got, "host_test_value") { t.Errorf("stdout = %q, want contains 'host_test_value'", got) } @@ -104,25 +99,27 @@ func TestHost_Exec_Env(t *testing.T) { } } -func TestHost_Workspace(t *testing.T) { +func TestHost_Workplace(t *testing.T) { skipIfNoHostExec(t) for _, tgt := range hostExecTargets() { + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } sessionID := fmt.Sprintf("host-test-%d", time.Now().UnixNano()) - ws := host.Workspace(sessionID) + host.BindWorkplace(sessionID) + ws := host.Workplace() if ws == nil { - t.Fatal("Workspace returned nil") + t.Fatal("Workplace returned nil after BindWorkplace") } - content := []byte("hello from host workspace test") + content := []byte("hello from host workplace test") if err := ws.WriteFile("test.txt", content, 0644); err != nil { t.Fatalf("WriteFile: %v", err) } @@ -164,10 +161,11 @@ func TestHost_Stream_Incremental(t *testing.T) { if tgt.IsWinNative { continue } + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } @@ -175,15 +173,22 @@ func TestHost_Stream_Incremental(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - stream, err := host.Stream(ctx, "sh", []string{"-c", + stream, err := host.Stream(ctx, []string{"sh", "-c", "for i in 1 2 3 4 5; do echo chunk$i; sleep 0.2; done"}) if err != nil { t.Fatalf("Stream: %v", err) } var chunks []string - for chunk := range stream.Stdout { - chunks = append(chunks, string(chunk)) + buf := make([]byte, 4096) + for { + n, err := stream.Stdout.Read(buf) + if n > 0 { + chunks = append(chunks, string(buf[:n])) + } + if err != nil { + break + } } exitCode, err := stream.Wait() @@ -219,10 +224,11 @@ func TestHost_Stream_MultiLine(t *testing.T) { if tgt.IsWinNative { continue } + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } @@ -230,15 +236,12 @@ func TestHost_Stream_MultiLine(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - stream, err := host.Stream(ctx, "sh", []string{"-c", "for i in 1 2 3; do echo line$i; done"}) + stream, err := host.Stream(ctx, []string{"sh", "-c", "for i in 1 2 3; do echo line$i; done"}) if err != nil { t.Fatalf("Stream: %v", err) } - var stdout []byte - for chunk := range stream.Stdout { - stdout = append(stdout, chunk...) - } + stdout, _ := io.ReadAll(stream.Stdout) exitCode, err := stream.Wait() if err != nil && !strings.Contains(err.Error(), "EOF") { @@ -267,10 +270,11 @@ func TestHost_Stream_Stderr(t *testing.T) { if tgt.IsWinNative { continue } + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } @@ -278,22 +282,17 @@ func TestHost_Stream_Stderr(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - stream, err := host.Stream(ctx, "sh", []string{"-c", "echo err-msg >&2"}) + stream, err := host.Stream(ctx, []string{"sh", "-c", "echo err-msg >&2"}) if err != nil { t.Fatalf("Stream: %v", err) } - var stderr []byte done := make(chan struct{}) go func() { - for chunk := range stream.Stdout { - _ = chunk - } + io.ReadAll(stream.Stdout) close(done) }() - for chunk := range stream.Stderr { - stderr = append(stderr, chunk...) - } + stderr, _ := io.ReadAll(stream.Stderr) <-done exitCode, err := stream.Wait() @@ -321,10 +320,11 @@ func TestHost_Stream_Cancel(t *testing.T) { if tgt.IsWinNative { continue } + tgt := tgt t.Run(tgt.Name, func(t *testing.T) { - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - host, err := m.Host(context.Background(), tgt.Name) + host, err := m.Host(context.Background(), tgt.TaiID) if err != nil { t.Skipf("Host(%s): %v", tgt.Name, err) } @@ -332,7 +332,7 @@ func TestHost_Stream_Cancel(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - stream, err := host.Stream(ctx, "sh", []string{"-c", "while true; do echo tick; sleep 0.1; done"}) + stream, err := host.Stream(ctx, []string{"sh", "-c", "while true; do echo tick; sleep 0.1; done"}) if err != nil { if strings.Contains(err.Error(), "not in the allowed list") { t.Skipf("command not allowed on %s", tgt.Name) @@ -341,13 +341,19 @@ func TestHost_Stream_Cancel(t *testing.T) { } received := 0 - for chunk := range stream.Stdout { - _ = chunk - received++ + buf := make([]byte, 4096) + for { + n, err := stream.Stdout.Read(buf) + if n > 0 { + received++ + } if received >= 3 { stream.Cancel() break } + if err != nil { + break + } } _, waitErr := stream.Wait() @@ -361,49 +367,92 @@ func TestHost_Stream_Cancel(t *testing.T) { } } -func TestHost_CreateRejectsNoContainerPool(t *testing.T) { - // Use the Windows native HostExec target which has no Docker. +func TestHost_ComputerInfo(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostExecTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + + host, err := m.Host(context.Background(), tgt.TaiID) + if err != nil { + t.Skipf("Host(%s): %v", tgt.Name, err) + } + + info := host.ComputerInfo() + if info.Kind != "host" { + t.Errorf("Kind = %q, want 'host'", info.Kind) + } + if info.NodeID != tgt.TaiID { + t.Errorf("NodeID = %q, want %q", info.NodeID, tgt.TaiID) + } + }) + } +} + +func TestHost_ComputerInterface(t *testing.T) { + skipIfNoHostExec(t) + + for _, tgt := range hostExecTargets() { + tgt := tgt + t.Run(tgt.Name, func(t *testing.T) { + m := setupHostManager(t, &tgt) + + host, err := m.Host(context.Background(), tgt.TaiID) + if err != nil { + t.Skipf("Host(%s): %v", tgt.Name, err) + } + + // Verify Host satisfies Computer interface at runtime. + var c sandbox.Computer = host + info := c.ComputerInfo() + if info.Kind != "host" { + t.Errorf("Computer.ComputerInfo().Kind = %q, want 'host'", info.Kind) + } + }) + } +} + +func TestHost_CreateRejectsNoContainerNode(t *testing.T) { tgt := findHostExecOnly(t) if tgt == nil { t.Skip("no host-exec-only target available") } - m := setupHostManager(t, *tgt) + m := setupHostManager(t, tgt) ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() _, err := m.Create(ctx, sandbox.CreateOptions{ - Image: "alpine:latest", - Owner: "test", - Pool: tgt.Name, + Image: "alpine:latest", + Owner: "test", + NodeID: tgt.TaiID, }) if err == nil { - t.Fatal("expected error for Create on host-exec-only pool, got nil") + t.Fatal("expected error for Create on host-exec-only node, got nil") } if !strings.Contains(err.Error(), "no container runtime") { t.Errorf("error = %q, want contains 'no container runtime'", err.Error()) } } -func TestHost_PoolNotFound(t *testing.T) { +func TestHost_NodeNotFound(t *testing.T) { skipIfNoHostExec(t) tgt := hostExecTargets()[0] - m := setupHostManager(t, tgt) + m := setupHostManager(t, &tgt) - _, err := m.Host(context.Background(), "nonexistent-pool") + _, err := m.Host(context.Background(), "nonexistent-node") if err == nil { t.Fatal("expected error, got nil") } } -// findHostExecOnly returns a hostExecTarget that is likely host-exec-only -// (Windows native Tai without Docker). func findHostExecOnly(t *testing.T) *hostExecTarget { t.Helper() for _, tgt := range hostExecTargets() { if tgt.IsWinNative { - // Windows native Tai typically has no Docker addr := fmt.Sprintf("tai://%s", tgt.Addr) client, err := tai.New(addr) if err != nil { @@ -418,3 +467,12 @@ func findHostExecOnly(t *testing.T) *hostExecTarget { } return nil } + +// hostCmd builds a []string command, adapting for Windows targets. +func hostCmd(tgt hostExecTarget, prog string, args ...string) []string { + if tgt.IsWinNative { + cmd, wArgs := linuxCmd(tgt, prog, args...) + return append([]string{cmd}, wArgs...) + } + return append([]string{prog}, args...) +} diff --git a/sandbox/v2/jsapi/API.md b/sandbox/v2/jsapi/API.md new file mode 100644 index 00000000..5eb9c6b8 --- /dev/null +++ b/sandbox/v2/jsapi/API.md @@ -0,0 +1,603 @@ +# Sandbox JavaScript API + +All methods are available on the global `sandbox` object. No constructor needed. + +## Quick Start + +```javascript +// Create a container computer +const pc = sandbox.Create({ image: "node:20", owner: "user-123" }) +const result = pc.Exec(["node", "-e", "console.log('hello')"]) +console.log(result.stdout) // "hello\n" +pc.Remove() + +// Or use the host directly (no container) +const host = sandbox.Host() +const info = host.Exec(["uname", "-a"]) +console.log(info.stdout) // same ExecResult as box +``` + +Both `sandbox.Create()` and `sandbox.Host()` return a **Computer** object with the same interface. The `kind` property tells you which type it is. + +--- + +## Static Methods + +### sandbox.Create(options) → Computer + +Create a new sandbox container. Returns a Computer (`kind = "box"`). If `options.id` is set and a sandbox with that ID already exists, returns the existing one (GetOrCreate semantics). + +```javascript +const pc = sandbox.Create({ + image: "node:20", // required — container image + owner: "user-123", // required — owner identifier + node_id: "192.168.1.10-19100", // optional — TaiID from registry (required unless workspace_id routes to a node) + id: "my-sandbox", // optional — if set, uses GetOrCreate + workdir: "/app", // optional — working directory + user: "1000:1000", // optional — UID:GID + env: { NODE_ENV: "dev" },// optional — environment variables + memory: 536870912, // optional — memory limit in bytes (512MB) + cpus: 1.5, // optional — CPU limit + vnc: true, // optional — enable VNC desktop + ports: [ // optional — port mappings + { container_port: 3000, host_port: 3000, host_ip: "", protocol: "tcp" } + ], + policy: "session", // optional — "oneshot"|"session"|"longrunning"|"persistent" + idle_timeout: 600000, // optional — idle timeout in ms (10min) + stop_timeout: 30000, // optional — stop timeout in ms + workspace_id: "ws-abc", // optional — bind a workspace + mount_mode: "rw", // optional — "rw"|"ro" + mount_path: "/workspace", // optional — mount path in container + labels: { team: "backend" } // optional — custom labels +}) +``` + +### sandbox.Get(id) → Computer | null + +Get an existing sandbox by ID. Returns a Computer (`kind = "box"`) or `null` if not found. + +```javascript +const pc = sandbox.Get("my-sandbox") +if (pc) { + console.log(pc.kind, pc.id, pc.owner, pc.node_id) +} +``` + +### sandbox.List(filter?) → BoxInfo[] + +List all sandboxes, optionally filtered. + +```javascript +// All sandboxes +const all = sandbox.List() + +// Filter by owner +const mine = sandbox.List({ owner: "user-123" }) + +// Filter by node_id (TaiID) and labels +const gpu = sandbox.List({ node_id: "10.0.0.5-19100", labels: { team: "ml" } }) +``` + +Each element in the returned array: + +```javascript +{ + id: "sb-xxx", + container_id: "abc123...", + node_id: "192.168.1.10-19100", + owner: "user-123", + status: "running", // "running"|"stopped"|"creating"|... + image: "node:20", + vnc: false, + policy: "session", + labels: { team: "backend" }, + created_at: "2026-03-07T10:00:00Z", + last_active: "2026-03-07T10:05:00Z", + process_count: 2 +} +``` + +### sandbox.Delete(id) → void + +Remove a sandbox and its container. + +```javascript +sandbox.Delete("my-sandbox") +``` + +### sandbox.Host(nodeID?) → Computer + +Get a Computer (`kind = "host"`) for executing commands directly on the Tai host machine (no container). Only available when the node's Tai server has `host_exec` capability. The `nodeID` argument is the TaiID (e.g. `"192.168.1.10-19100"`). + +```javascript +const host = sandbox.Host("192.168.1.10-19100") +``` + +### sandbox.GetNode(taiID) → NodeInfo | null + +Get information about a registered node by its Tai ID. + +```javascript +const node = sandbox.GetNode("tai-abc123") +if (node) { + console.log(node.status, node.system.hostname) +} +``` + +### sandbox.Nodes() → NodeInfo[] + +List all registered nodes. + +```javascript +const nodes = sandbox.Nodes() +nodes.forEach(function(n) { + console.log(n.tai_id, n.status, n.display_name, n.system.os) +}) +``` + +### sandbox.NodesByTeam(teamID) → NodeInfo[] + +List nodes belonging to a specific team. + +```javascript +const nodes = sandbox.NodesByTeam("team-001") +``` + +--- + +## Computer Object + +Returned by `sandbox.Create()`, `sandbox.Get()`, and `sandbox.Host()`. This is the unified interface for all execution environments — containers and bare-metal hosts. + +Use the `kind` property to check the type. Methods marked **box-only** throw an error when called on a host computer. `Proxy()` covers HTTP, WebSocket, and SSE — use it for all protocol access to container/host services. + +### Properties (read-only) + +| Property | Type | Description | +|----------|------|-------------| +| `pc.kind` | string | `"box"` or `"host"` | +| `pc.id` | string | Sandbox ID (box-only; empty for host) | +| `pc.owner` | string | Owner identifier (box-only; empty for host) | +| `pc.node_id` | string | TaiID (e.g. `"192.168.1.10-19100"`, `"local"`) | + +### pc.Exec(cmd, options?) → ExecResult + +Execute a command and wait for it to finish. + +```javascript +const result = pc.Exec(["ls", "-la", "/app"]) +console.log(result.exit_code) // 0 +console.log(result.stdout) // file listing +``` + +Options: + +```javascript +pc.Exec(["python3", "train.py"], { + workdir: "/workspace/ml", + env: { CUDA_VISIBLE_DEVICES: "0" }, + stdin: "input data", + timeout: 300000, // ms + max_output: 10485760 // bytes (10MB) +}) +``` + +Return value: + +```javascript +{ + exit_code: 0, + stdout: "...", // UTF-8 string + stderr: "...", // UTF-8 string + duration_ms: 1234, // execution time in ms + error: "", // error message (empty on success) + truncated: false // true if output was truncated by max_output +} +``` + +### pc.Stream(cmd, callback) / pc.Stream(cmd, options, callback) + +Execute a command with streaming output via callback. The call blocks until the process exits. + +Callback signature: `function(type, data)` +- `type = "stdout"` → `data` is a string chunk from stdout +- `type = "stderr"` → `data` is a string chunk from stderr +- `type = "exit"` → `data` is the exit code (number) + +```javascript +pc.Stream(["npm", "run", "dev"], function(type, data) { + if (type === "stdout") console.log(data) + if (type === "stderr") console.log("[ERR]", data) + if (type === "exit") console.log("exited:", data) +}) + +// With options +pc.Stream(["npm", "test"], { + workdir: "/app", + env: { CI: "true" }, + timeout: 60000 +}, function(type, data) { + console.log(type, data) +}) +``` + +### pc.VNC() → string + +Get the VNC WebSocket URL. + +- **Box**: routes to the container's VNC server (`:5900`) +- **Host**: routes to the Tai host via `__host__` identifier (configurable via `host_vnc_port`) + +```javascript +const url = pc.VNC() +// Box: "ws://tai-host:16080/vnc/container-id/ws" +// Host: "ws://tai-host:16080/vnc/__host__/ws" +``` + +If no VNC server is running, the WebSocket connection will fail — handle this in the caller. + +### pc.Proxy(port, path?) → string + +Get a proxy URL for a service port. Supports HTTP, WebSocket (`ws://`), and SSE — the Tai proxy handles protocol upgrades automatically. + +- **Box**: routes to `container-ip:{port}` +- **Host**: routes to `127.0.0.1:{port}` on the Tai machine via `__host__` + +```javascript +const url = pc.Proxy(3000) +// Box: "http://tai-host:8099/container-id:3000/" +// Host: "http://tai-host:8099/__host__:3000/" + +const url = pc.Proxy(8080, "/api/v1") +// Box: "http://tai-host:8099/container-id:8080/api/v1" +// Host: "http://tai-host:8099/__host__:8080/api/v1" +``` + +### pc.ComputerInfo() → ComputerInfo + +Get identity and registry information. + +```javascript +const info = pc.ComputerInfo() +console.log(info.kind) // "box" or "host" +console.log(info.node_id) // TaiID +console.log(info.system.os) // "linux" | "windows" | "darwin" +console.log(info.status) // "running" | "stopped" | ... +``` + +Returns a [ComputerInfo](#computerinfo-object) object. + +### pc.BindWorkplace(workspaceID) → void + +Bind a workspace to this computer for the current session. For box computers created with a `workspace_id` option, the workspace is already bound at creation time — calling `BindWorkplace` overrides it. + +```javascript +pc.BindWorkplace("ws-project-abc") +``` + +### pc.Workplace() → WorkspaceFS | null + +Access the workspace filesystem bound via `BindWorkplace()`. Returns `null` if no workspace is bound. ("Workplace" is the binding on a Computer; "Workspace" is the filesystem it points to.) + +```javascript +pc.BindWorkplace("ws-project-abc") +const ws = pc.Workplace() +ws.ReadFile("config.yml") +ws.WriteFile("output.json", JSON.stringify(data)) +``` + +See [WorkspaceFS Object](#workspacefs-object) for the full method list. + +### pc.Info() → BoxInfo — box-only + +Get current container runtime status (process count, last active time, etc.). For node-level identity info (OS, CPU, capabilities), use `ComputerInfo()` instead. Throws on host computers. + +```javascript +const info = pc.Info() +console.log(info.status, info.process_count, info.last_active) +``` + +Returns the same structure as elements in `sandbox.List()`. + +### pc.Start() → void — box-only + +Start a stopped container. Throws on host computers. + +```javascript +pc.Start() +``` + +### pc.Stop() → void — box-only + +Stop a running container. Throws on host computers. + +```javascript +pc.Stop() +``` + +### pc.Remove() → void — box-only + +Remove the container. Throws on host computers. + +```javascript +pc.Remove() +``` + +--- + +## ComputerInfo Object + +Returned by `pc.ComputerInfo()`. Read-only snapshot of a Computer's identity and state. + +```javascript +{ + kind: "box", // "box" | "host" + node_id: "192.168.1.10-19100", // TaiID + tai_id: "tai-abc123", + machine_id: "m-xyz", + version: "1.2.3", + mode: "direct", // "direct" | "tunnel" + status: "running", + capabilities: { docker: true, k8s: false, host_exec: true }, + system: { + os: "linux", + arch: "amd64", + hostname: "gpu-server-01", + num_cpu: 16, + total_mem: 68719476736 + }, + + // Box-only fields (empty/zero for host) + box_id: "sb-xxx", + container_id: "abc123...", + owner: "user-123", + image: "node:20", + policy: "session", + labels: { team: "backend" } +} +``` + +--- + +## NodeInfo Object + +Returned by `sandbox.GetNode()`, `sandbox.Nodes()`, `sandbox.NodesByTeam()`. Read-only view of a registered Tai node. + +```javascript +{ + tai_id: "tai-abc123", + machine_id: "m-xyz", + version: "1.2.3", + mode: "direct", // "direct" | "tunnel" + addr: "tai://192.168.1.100:19100", + status: "online", // "online" | "offline" | "connecting" + display_name: "GPU Node", // optional human-readable name for UI + node_id: "gpu", + connected_at: "2026-03-07T08:00:00Z", + last_ping: "2026-03-07T10:05:00Z", + ports: { + grpc: 19100, + http: 8099, + vnc: 16080, + docker: 12375, + k8s: 16443, + host_vnc: 5900 // VNC port on host for __host__ routing + }, + capabilities: { + docker: true, + k8s: false, + host_exec: true + }, + system: { + os: "linux", + arch: "amd64", + hostname: "gpu-server-01", + num_cpu: 16, + total_mem: 68719476736 // bytes (64GB) + } +} +``` + +--- + +## WorkspaceFS Object + +Returned by `pc.Workplace()`, `workspace.Get()`, and `workspace.Create()`. + +### Properties (read-only) + +| Property | Type | Description | +|----------|------|-------------| +| `ws.id` | string | Workspace ID | +| `ws.name` | string | Workspace name | +| `ws.node` | string | Node name | + +### Methods + +| Method | Returns | Description | +|--------|---------|-------------| +| `ws.ReadFile(path)` | `string` | Read file content as UTF-8 string | +| `ws.WriteFile(path, data, perm?)` | `void` | Write string data to file. `perm` defaults to `0644` | +| `ws.ReadDir(path?)` | `DirEntry[]` | List directory contents. Defaults to root | +| `ws.Stat(path)` | `FileInfo` | Get file/directory metadata | +| `ws.MkdirAll(path, perm?)` | `void` | Create directory tree. `perm` defaults to `0755` | +| `ws.Remove(path)` | `void` | Remove a file | +| `ws.RemoveAll(path)` | `void` | Remove a file or directory recursively | +| `ws.Rename(from, to)` | `void` | Rename/move a file or directory | + +Return types: + +```javascript +// DirEntry +{ name: "main.go", is_dir: false, size: 1234 } + +// FileInfo +{ name: "main.go", size: 1234, is_dir: false, mod_time: "2026-03-07T10:00:00Z" } +``` + +--- + +## Examples + +### Run a build and check output + +```javascript +const pc = sandbox.Create({ + image: "golang:1.23", + owner: "ci-bot", + workspace_id: "ws-project-abc" +}) + +const build = pc.Exec(["go", "build", "./..."], { + workdir: "/workspace", + timeout: 120000 +}) + +if (build.exit_code !== 0) { + console.log("Build failed:", build.stderr) + pc.Remove() + throw new Error("build failed") +} + +const test = pc.Exec(["go", "test", "./..."], { + workdir: "/workspace", + env: { CGO_ENABLED: "0" } +}) + +console.log("Tests:", test.exit_code === 0 ? "PASS" : "FAIL") +pc.Remove() +``` + +### Stream a long-running process + +```javascript +const pc = sandbox.Create({ + image: "node:20", + owner: "user-123", + policy: "session" +}) + +pc.Exec(["npm", "install"], { workdir: "/app" }) + +pc.Stream(["npm", "run", "dev"], { workdir: "/app" }, function(type, data) { + if (type === "stdout") console.log(data) + if (type === "stderr") console.log("[ERR]", data) + if (type === "exit") console.log("dev server exited:", data) +}) +``` + +### Host execution for GPU workloads + +```javascript +const host = sandbox.Host("10.0.0.5-19100") + +const result = host.Exec(["nvidia-smi"]) +console.log(result.stdout) + +const train = host.Exec(["python3", "train.py", "--epochs=10"], { + workdir: "/workspace/ml", + env: { CUDA_VISIBLE_DEVICES: "0,1" }, + timeout: 3600000 +}) +if (train.exit_code !== 0) throw new Error("training failed: " + train.stderr) +``` + +### Uniform interface — same code for box and host + +```javascript +function runTask(pc, cmd, opts) { + const result = pc.Exec(cmd, opts) + if (result.exit_code !== 0) { + throw new Error(pc.kind + " exec failed: " + result.stderr) + } + return result.stdout +} + +// Works the same for both +const box = sandbox.Create({ image: "node:20", owner: "u1" }) +const host = sandbox.Host("10.0.0.5-19100") + +runTask(box, ["node", "-e", "console.log('hi')"]) +runTask(host, ["echo", "hello"]) +``` + +### VNC and HTTP proxy + +```javascript +const pc = sandbox.Create({ + image: "kasmweb/chrome:latest", + owner: "user-123", + vnc: true +}) + +// Get VNC desktop URL +const vncURL = pc.VNC() +// "ws://tai-host:16080/vnc/container-id/ws" + +// Get HTTP proxy to a web service inside the container +const appURL = pc.Proxy(3000) +// "http://tai-host:8099/container-id:3000/" + +// Same methods work on host +const host = sandbox.Host("192.168.1.10-19100") +const hostVNC = host.VNC() +// "ws://tai-host:16080/vnc/__host__/ws" +``` + +### Query cluster nodes + +```javascript +const nodes = sandbox.Nodes() + +// Find online GPU nodes +const gpuNodes = nodes.filter(function(n) { + return n.status === "online" && n.display_name === "gpu" // n.display_name is optional label for UI +}) + +console.log("Available GPU nodes:", gpuNodes.length) +gpuNodes.forEach(function(n) { + console.log( + n.tai_id, + n.system.hostname, + n.system.num_cpu + " CPUs", + Math.round(n.system.total_mem / 1073741824) + "GB RAM" + ) +}) +``` + +### Workspace file operations + +```javascript +const pc = sandbox.Create({ + image: "node:20", + owner: "user-123" +}) + +pc.BindWorkplace("ws-my-project") +const ws = pc.Workplace() + +ws.MkdirAll("src/utils") +ws.WriteFile("src/main.go", 'package main\n\nfunc main() {\n\tprintln("hello")\n}\n') +ws.WriteFile("go.mod", "module myproject\n\ngo 1.23\n") + +const entries = ws.ReadDir("src/") +entries.forEach(function(e) { + console.log(e.name, e.is_dir ? "(dir)" : e.size + " bytes") +}) + +const content = ws.ReadFile("src/main.go") +console.log(content) +``` + +### Permission check pattern + +```javascript +const auth = Authorized() +if (!auth) throw new Error("not authenticated") + +const pc = sandbox.Get(id) +if (!pc) throw new Error("sandbox not found") +if (pc.owner !== auth.user_id) throw new Error("permission denied") + +pc.Exec(["ls", "-la"]) +``` diff --git a/sandbox/v2/jsapi/box.go b/sandbox/v2/jsapi/box.go deleted file mode 100644 index e93b5d84..00000000 --- a/sandbox/v2/jsapi/box.go +++ /dev/null @@ -1,129 +0,0 @@ -package jsapi - -import ( - "rogchap.com/v8go" -) - -// NewBoxObject creates a JS Box object backed by a sandbox ID string. -// All methods delegate to the Go sandbox.M() singleton — no Go object is -// passed to V8, no bridge registration, no Release() needed. -// -// # Properties (read-only) -// -// box.id → string // sandbox ID ← Box.ID() -// box.owner → string // owner user ID ← Box.Owner() -// box.pool → string // pool name ← Box.Pool() -// -// # Methods — Go mapping -// -// box.Exec(cmd, options?) → ExecResult -// -// Go: Box.Exec(ctx, cmd []string, opts ...ExecOption) (*ExecResult, error) -// -// JS args: -// cmd: string[] → cmd []string -// options: { → ExecOption functional options -// workdir: string, → WithWorkDir(dir) -// env: object, → WithEnv(map[string]string) -// timeout: number → WithTimeout(ms → time.Duration) -// } -// JS returns: { -// exit_code: number, ← ExecResult.ExitCode -// stdout: string, ← ExecResult.Stdout -// stderr: string ← ExecResult.Stderr -// } -// -// box.Stream(cmd, options?) → ExecStream -// -// Go: Box.Stream(ctx, cmd []string, opts ...ExecOption) (*ExecStream, error) -// -// JS returns: { -// stdout: ReadableStream, ← ExecStream.Stdout -// stderr: ReadableStream, ← ExecStream.Stderr -// stdin: WritableStream, ← ExecStream.Stdin -// wait: function() → number, ← ExecStream.Wait() (int, error) -// cancel: function() → void ← ExecStream.Cancel() -// } -// -// box.Attach(port, options?) → ServiceConn -// -// Go: Box.Attach(ctx, port int, opts ...AttachOption) (*ServiceConn, error) -// -// JS args: -// port: number → port int -// options: { → AttachOption functional options -// protocol: "ws"|"sse", → WithProtocol(protocol) -// path: string, → WithPath(path) -// headers: object → WithHeaders(map[string]string) -// } -// JS returns: { -// url: string, ← ServiceConn.URL -// read: function() → Uint8Array, ← ServiceConn.Read() ([]byte, error) -// write: function(data) → void, ← ServiceConn.Write(data) error -// events: AsyncIterable, ← ServiceConn.Events <-chan []byte -// close: function() → void ← ServiceConn.Close() error -// } -// -// box.VNC() → string -// -// Go: Box.VNC(ctx) (string, error) -// Returns: VNC WebSocket URL -// -// box.Proxy(port, path?) → string -// -// Go: Box.Proxy(ctx, port int, path string) (string, error) -// Returns: HTTP proxy URL -// -// box.Workspace() → WorkspaceFS -// -// Go: Box.Workspace() workspace.FS -// Box.WorkspaceID() string -// Returns: WorkspaceFS object (see workspace/jsapi/fs.go) -// Uses box.WorkspaceID() to create NewFSObject -// -// box.Info() → BoxInfo -// -// Go: Box.Info(ctx) (*BoxInfo, error) -// JS returns: { -// id: string, ← BoxInfo.ID -// container_id: string, ← BoxInfo.ContainerID -// pool: string, ← BoxInfo.Pool -// owner: string, ← BoxInfo.Owner -// status: string, ← BoxInfo.Status -// image: string, ← BoxInfo.Image -// vnc: boolean, ← BoxInfo.VNC -// policy: string, ← BoxInfo.Policy (LifecyclePolicy) -// labels: object, ← BoxInfo.Labels (map[string]string) -// created_at: string, ← BoxInfo.CreatedAt (ISO 8601) -// last_active: string, ← BoxInfo.LastActive (ISO 8601) -// process_count: number ← BoxInfo.ProcessCount -// } -// -// box.Start() → void -// -// Go: Box.Start(ctx) error -// -// box.Stop() → void -// -// Go: Box.Stop(ctx) error -// -// box.Remove() → void -// -// Go: Box.Remove(ctx) error -func NewBoxObject(v8ctx *v8go.Context, boxID string) (*v8go.Value, error) { - // TODO: Phase 2 implementation - // 1. Create JS object via v8go.NewObjectTemplate - // 2. Set read-only properties: id, owner, pool (from sandbox.M().Get(boxID)) - // 3. Bind each method as FunctionTemplate: - // - Exec → sandbox.M().Get(id).Exec(ctx, cmd, opts...) - // - Stream → sandbox.M().Get(id).Stream(ctx, cmd, opts...) - // - Attach → sandbox.M().Get(id).Attach(ctx, port, opts...) - // - VNC → sandbox.M().Get(id).VNC(ctx) - // - Proxy → sandbox.M().Get(id).Proxy(ctx, port, path) - // - Workspace → NewFSObject(v8ctx, sandbox.M().Get(id).WorkspaceID()) - // - Info → sandbox.M().Get(id).Info(ctx) → JS object - // - Start → sandbox.M().Get(id).Start(ctx) - // - Stop → sandbox.M().Get(id).Stop(ctx) - // - Remove → sandbox.M().Get(id).Remove(ctx) - return nil, nil -} diff --git a/sandbox/v2/jsapi/computer.go b/sandbox/v2/jsapi/computer.go new file mode 100644 index 00000000..4180593c --- /dev/null +++ b/sandbox/v2/jsapi/computer.go @@ -0,0 +1,471 @@ +package jsapi + +import ( + "context" + "encoding/json" + "sync" + "time" + + sandbox "github.com/yaoapp/yao/sandbox/v2" + wsjsapi "github.com/yaoapp/yao/workspace/jsapi" + "rogchap.com/v8go" +) + +// --------------------------------------------------------------------------- +// Helpers — shared across jsapi files +// --------------------------------------------------------------------------- + +func throwError(info *v8go.FunctionCallbackInfo, msg string) *v8go.Value { + iso := info.Context().Isolate() + e, _ := v8go.NewValue(iso, msg) + iso.ThrowException(e) + return v8go.Undefined(iso) +} + +func parseStringArray(val *v8go.Value) []string { + obj, err := val.AsObject() + if err != nil { + return nil + } + lenVal, err := obj.Get("length") + if err != nil { + return nil + } + length := int(lenVal.Int32()) + result := make([]string, 0, length) + for i := 0; i < length; i++ { + item, err := obj.GetIdx(uint32(i)) + if err != nil || !item.IsString() { + continue + } + result = append(result, item.String()) + } + return result +} + +func parseStringMap(v8ctx *v8go.Context, val *v8go.Value) map[string]string { + result := make(map[string]string) + if !val.IsObject() { + return result + } + jsonStr, err := v8go.JSONStringify(v8ctx, val) + if err != nil { + return result + } + _ = json.Unmarshal([]byte(jsonStr), &result) + return result +} + +func parseExecOptions(v8ctx *v8go.Context, args []*v8go.Value) ([]string, []sandbox.ExecOption, *v8go.Value) { + if len(args) < 1 || !args[0].IsObject() { + return nil, nil, nil + } + cmd := parseStringArray(args[0]) + if len(cmd) == 0 { + return nil, nil, nil + } + var opts []sandbox.ExecOption + var callback *v8go.Value + for i := 1; i < len(args); i++ { + v := args[i] + if v.IsFunction() { + callback = v + break + } + if v.IsObject() { + optsObj, err := v.AsObject() + if err != nil { + continue + } + if wd, e := optsObj.Get("workdir"); e == nil && wd.IsString() { + opts = append(opts, sandbox.WithWorkDir(wd.String())) + } + if env, e := optsObj.Get("env"); e == nil && env.IsObject() { + envMap := parseStringMap(v8ctx, env) + if len(envMap) > 0 { + opts = append(opts, sandbox.WithEnv(envMap)) + } + } + if stdin, e := optsObj.Get("stdin"); e == nil && stdin.IsString() { + opts = append(opts, sandbox.WithStdin([]byte(stdin.String()))) + } + if t, e := optsObj.Get("timeout"); e == nil && t.IsNumber() { + opts = append(opts, sandbox.WithTimeout(time.Duration(t.Number())*time.Millisecond)) + } + if mo, e := optsObj.Get("max_output"); e == nil && mo.IsNumber() { + opts = append(opts, sandbox.WithMaxOutput(int64(mo.Number()))) + } + } + } + return cmd, opts, callback +} + +func execResultToJS(v8ctx *v8go.Context, r *sandbox.ExecResult) *v8go.Value { + data, _ := json.Marshal(map[string]interface{}{ + "exit_code": r.ExitCode, + "stdout": r.Stdout, + "stderr": r.Stderr, + "duration_ms": r.DurationMs, + "error": r.Error, + "truncated": r.Truncated, + }) + val, _ := v8go.JSONParse(v8ctx, string(data)) + return val +} + +func boxInfoToJS(v8ctx *v8go.Context, b *sandbox.BoxInfo) *v8go.Value { + data, _ := json.Marshal(map[string]interface{}{ + "id": b.ID, + "container_id": b.ContainerID, + "node_id": b.NodeID, + "owner": b.Owner, + "status": b.Status, + "image": b.Image, + "vnc": b.VNC, + "policy": string(b.Policy), + "labels": b.Labels, + "created_at": b.CreatedAt.Format(time.RFC3339), + "last_active": b.LastActive.Format(time.RFC3339), + "process_count": b.ProcessCount, + }) + val, _ := v8go.JSONParse(v8ctx, string(data)) + return val +} + +func computerInfoToJS(v8ctx *v8go.Context, c sandbox.ComputerInfo) *v8go.Value { + data, _ := json.Marshal(map[string]interface{}{ + "kind": c.Kind, + "node_id": c.NodeID, + "tai_id": c.TaiID, + "machine_id": c.MachineID, + "version": c.Version, + "mode": c.Mode, + "status": c.Status, + "capabilities": c.Capabilities, + "system": map[string]interface{}{ + "os": c.System.OS, + "arch": c.System.Arch, + "hostname": c.System.Hostname, + "num_cpu": c.System.NumCPU, + "total_mem": c.System.TotalMem, + }, + "box_id": c.BoxID, + "container_id": c.ContainerID, + "owner": c.Owner, + "image": c.Image, + "policy": string(c.Policy), + "labels": c.Labels, + }) + val, _ := v8go.JSONParse(v8ctx, string(data)) + return val +} + +// getComputer re-fetches a Computer from the Manager by kind + identifier. +// kind="box" → identifier is boxID, kind="host" → identifier is node ID. +func getComputer(ctx context.Context, kind, identifier string) (sandbox.Computer, error) { + m := sandbox.M() + if kind == "box" { + return m.Get(ctx, identifier) + } + return m.Host(ctx, identifier) +} + +// --------------------------------------------------------------------------- +// sbHost — sandbox.Host(nodeID?) +// --------------------------------------------------------------------------- + +func sbHost(info *v8go.FunctionCallbackInfo) *v8go.Value { + ctx := context.Background() + v8ctx := info.Context() + + nodeID := "" + args := info.Args() + if len(args) > 0 && args[0].IsString() { + nodeID = args[0].String() + } + + if _, err := sandbox.M().Host(ctx, nodeID); err != nil { + return throwError(info, err.Error()) + } + + val, err := NewComputerObject(v8ctx, "host", nodeID) + if err != nil { + return throwError(info, err.Error()) + } + return val +} + +// --------------------------------------------------------------------------- +// NewComputerObject — unified JS Computer object factory +// --------------------------------------------------------------------------- + +// NewComputerObject creates a JS Computer object. Closures capture only +// kind (string) and identifier (string) — no Go objects cross into V8. +func NewComputerObject(v8ctx *v8go.Context, kind string, identifier string) (*v8go.Value, error) { + iso := v8ctx.Isolate() + ctx := context.Background() + + // Mutable workplace binding lives in closure, not in V8 heap. + var workplaceID string + + tpl := v8go.NewObjectTemplate(iso) + + // -- Exec -- + tpl.Set("Exec", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + cmd, opts, _ := parseExecOptions(info.Context(), info.Args()) + if len(cmd) == 0 { + return throwError(info, "Exec requires cmd (string[])") + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + result, err := comp.Exec(ctx, cmd, opts...) + if err != nil { + return throwError(info, err.Error()) + } + return execResultToJS(info.Context(), result) + })) + + // -- Stream -- + tpl.Set("Stream", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + cmd, opts, cbVal := parseExecOptions(info.Context(), info.Args()) + if len(cmd) == 0 { + return throwError(info, "Stream requires cmd (string[]) and callback") + } + if cbVal == nil || !cbVal.IsFunction() { + return throwError(info, "Stream requires a callback function as last argument") + } + cbFn, err := cbVal.AsFunction() + if err != nil { + return throwError(info, "Stream callback is not a function") + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + stream, err := comp.Stream(ctx, cmd, opts...) + if err != nil { + return throwError(info, err.Error()) + } + + type chunk struct { + typ string + data interface{} + } + ch := make(chan chunk, 64) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + buf := make([]byte, 4096) + for { + n, err := stream.Stdout.Read(buf) + if n > 0 { + ch <- chunk{"stdout", string(buf[:n])} + } + if err != nil { + break + } + } + }() + go func() { + defer wg.Done() + buf := make([]byte, 4096) + for { + n, err := stream.Stderr.Read(buf) + if n > 0 { + ch <- chunk{"stderr", string(buf[:n])} + } + if err != nil { + break + } + } + }() + go func() { + code, _ := stream.Wait() + wg.Wait() + ch <- chunk{"exit", code} + close(ch) + }() + + v8c := info.Context() + global := v8c.Global() + for c := range ch { + var dataVal *v8go.Value + switch v := c.data.(type) { + case string: + dataVal, _ = v8go.NewValue(iso, v) + case int: + dataVal, _ = v8go.NewValue(iso, int32(v)) + } + typeVal, _ := v8go.NewValue(iso, c.typ) + if typeVal != nil && dataVal != nil { + _, _ = cbFn.Call(global, typeVal, dataVal) + } + } + return v8go.Undefined(iso) + })) + + // -- VNC -- + tpl.Set("VNC", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + url, err := comp.VNC(ctx) + if err != nil { + return throwError(info, err.Error()) + } + val, _ := v8go.NewValue(iso, url) + return val + })) + + // -- Proxy -- + tpl.Set("Proxy", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 || !args[0].IsNumber() { + return throwError(info, "Proxy requires port (number)") + } + port := int(args[0].Int32()) + path := "/" + if len(args) > 1 && args[1].IsString() { + path = args[1].String() + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + url, err := comp.Proxy(ctx, port, path) + if err != nil { + return throwError(info, err.Error()) + } + val, _ := v8go.NewValue(iso, url) + return val + })) + + // -- ComputerInfo -- + tpl.Set("ComputerInfo", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + return computerInfoToJS(info.Context(), comp.ComputerInfo()) + })) + + // -- BindWorkplace -- + tpl.Set("BindWorkplace", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "BindWorkplace requires workspaceID (string)") + } + workplaceID = args[0].String() + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + comp.BindWorkplace(workplaceID) + return v8go.Undefined(iso) + })) + + // -- Workplace → reuse workspace JSAPI NewFSObject -- + tpl.Set("Workplace", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + if workplaceID == "" { + return v8go.Null(iso) + } + val, err := wsjsapi.NewFSObject(info.Context(), workplaceID) + if err != nil { + return throwError(info, err.Error()) + } + return val + })) + + // -- Box-only: Info -- + tpl.Set("Info", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + if kind == "host" { + return throwError(info, "not supported: Info() requires a box computer") + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + box := comp.(*sandbox.Box) + bi, err := box.Info(ctx) + if err != nil { + return throwError(info, err.Error()) + } + return boxInfoToJS(info.Context(), bi) + })) + + // -- Box-only: Start -- + tpl.Set("Start", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + if kind == "host" { + return throwError(info, "not supported: Start() requires a box computer") + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + if err := comp.(*sandbox.Box).Start(ctx); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + // -- Box-only: Stop -- + tpl.Set("Stop", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + if kind == "host" { + return throwError(info, "not supported: Stop() requires a box computer") + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + if err := comp.(*sandbox.Box).Stop(ctx); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + // -- Box-only: Remove -- + tpl.Set("Remove", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + if kind == "host" { + return throwError(info, "not supported: Remove() requires a box computer") + } + comp, err := getComputer(ctx, kind, identifier) + if err != nil { + return throwError(info, err.Error()) + } + if err := comp.(*sandbox.Box).Remove(ctx); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + // Instantiate and set read-only properties + obj, err := tpl.NewInstance(v8ctx) + if err != nil { + return nil, err + } + + obj.Set("kind", kind) + + idStr := "" + ownerStr := "" + nodeIDStr := identifier + if kind == "box" { + if comp, err := getComputer(ctx, kind, identifier); err == nil { + box := comp.(*sandbox.Box) + idStr = box.ID() + ownerStr = box.Owner() + nodeIDStr = box.NodeID() + } else { + idStr = identifier + } + } + obj.Set("id", idStr) + obj.Set("owner", ownerStr) + obj.Set("node_id", nodeIDStr) + + return obj.Value, nil +} diff --git a/sandbox/v2/jsapi/jsapi.go b/sandbox/v2/jsapi/jsapi.go index 84c71337..d2c32999 100644 --- a/sandbox/v2/jsapi/jsapi.go +++ b/sandbox/v2/jsapi/jsapi.go @@ -1,24 +1,30 @@ // Package jsapi registers the sandbox namespace into the Yao V8 runtime. // // All methods are static on the sandbox object — no constructor. +// Both sandbox.Create() and sandbox.Host() return a unified Computer object. // // # JavaScript API // -// const box = sandbox.Create({ image: "node:20", owner: "user1" }) -// const result = box.Exec(["node", "-e", "console.log('hi')"]) -// console.log(result.stdout) -// -// const box = sandbox.Get(id) // → Box +// const pc = sandbox.Create({ image: "node:20", owner: "user1" }) // → Computer (kind="box") +// const pc = sandbox.Get(id) // → Computer (kind="box") | null // const list = sandbox.List({ owner: "u1" }) // → BoxInfo[] // sandbox.Delete(id) // → void +// const host = sandbox.Host("gpu") // → Computer (kind="host") +// const node = sandbox.GetNode("tai-abc123") // → NodeInfo | null +// const all = sandbox.Nodes() // → NodeInfo[] +// const team = sandbox.NodesByTeam("t-001") // → NodeInfo[] // // # Go mapping // -// sandbox.Create(opts) → Manager.Create(ctx, CreateOptions) → Box -// sandbox.Create(opts) → Manager.GetOrCreate(ctx, opts) → Box (when opts.id is set) -// sandbox.Get(id) → Manager.Get(ctx, id) → Box -// sandbox.List(filter?) → Manager.List(ctx, ListOptions) → []*Box → BoxInfo[] -// sandbox.Delete(id) → Manager.Remove(ctx, id) → void +// sandbox.Create(opts) → Manager.Create(ctx, CreateOptions) → Computer (Box) +// sandbox.Create(opts) → Manager.GetOrCreate(ctx, opts) → Computer (Box) (when opts.id is set) +// sandbox.Get(id) → Manager.Get(ctx, id) → Computer (Box) +// sandbox.List(filter?) → Manager.List(ctx, ListOptions) → BoxInfo[] +// sandbox.Delete(id) → Manager.Remove(ctx, id) → void +// sandbox.Host(nodeID?) → Manager.Host(ctx, nodeID) → Computer (Host) +// sandbox.GetNode(id) → registry.Global().Get(id) → NodeInfo | null +// sandbox.Nodes() → registry.Global().List() → NodeInfo[] +// sandbox.NodesByTeam(t)→ registry.Global().ListByTeam(t) → NodeInfo[] // // Registration happens via init() — import with: // @@ -26,7 +32,12 @@ package jsapi import ( + "context" + "encoding/json" + "time" + v8 "github.com/yaoapp/gou/runtime/v8" + sandbox "github.com/yaoapp/yao/sandbox/v2" "rogchap.com/v8go" ) @@ -41,115 +52,234 @@ func ExportObject(iso *v8go.Isolate) *v8go.ObjectTemplate { obj.Set("Get", v8go.NewFunctionTemplate(iso, sbGet)) obj.Set("List", v8go.NewFunctionTemplate(iso, sbList)) obj.Set("Delete", v8go.NewFunctionTemplate(iso, sbDelete)) + obj.Set("Host", v8go.NewFunctionTemplate(iso, sbHost)) + obj.Set("GetNode", v8go.NewFunctionTemplate(iso, sbGetNode)) + obj.Set("Nodes", v8go.NewFunctionTemplate(iso, sbNodes)) + obj.Set("NodesByTeam", v8go.NewFunctionTemplate(iso, sbNodesByTeam)) return obj } -// sbCreate: `sandbox.Create(options)` → Box -// -// Go: Manager.Create(ctx, CreateOptions) (*Box, error) -// -// Manager.GetOrCreate(ctx, CreateOptions) (*Box, error) — when opts.id is set -// -// JS options → Go CreateOptions mapping: -// -// { -// id: string → CreateOptions.ID // optional; triggers GetOrCreate -// owner: string → CreateOptions.Owner // required -// pool: string → CreateOptions.Pool // default: first pool -// image: string → CreateOptions.Image // required -// workdir: string → CreateOptions.WorkDir -// user: string → CreateOptions.User // e.g. "1000:1000" -// env: object → CreateOptions.Env // map[string]string -// memory: number → CreateOptions.Memory // bytes (int64) -// cpus: number → CreateOptions.CPUs // float64 e.g. 1.5 -// vnc: boolean → CreateOptions.VNC -// ports: array → CreateOptions.Ports // [{container, host, host_ip, protocol}] → []PortMapping -// policy: string → CreateOptions.Policy // "oneshot"|"session"|"longrunning"|"persistent" -// idle_timeout: number → CreateOptions.IdleTimeout // ms → time.Duration -// stop_timeout: number → CreateOptions.StopTimeout // ms → time.Duration -// workspace_id: string → CreateOptions.WorkspaceID -// mount_mode: string → CreateOptions.MountMode // "rw"|"ro" -// mount_path: string → CreateOptions.MountPath -// labels: object → CreateOptions.Labels // map[string]string -// } -// -// Returns: Box object (see box.go) +// sbCreate: `sandbox.Create(options)` → Computer (kind="box") func sbCreate(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. Parse options from info.Args()[0] - // 2. Validate required fields (image, owner) - // 3. If opts.id != "" → sandbox.M().GetOrCreate(ctx, opts) - // else → sandbox.M().Create(ctx, opts) - // 4. Return NewBoxObject(v8ctx, box.ID()) - return v8go.Undefined(info.Context().Isolate()) + v8ctx := info.Context() + ctx := context.Background() + args := info.Args() + if len(args) < 1 || !args[0].IsObject() { + return throwError(info, "Create requires options object") + } + + optsVal := args[0] + jsonStr, err := v8go.JSONStringify(v8ctx, optsVal) + if err != nil { + return throwError(info, "Create: invalid options: "+err.Error()) + } + + var raw map[string]interface{} + if err := json.Unmarshal([]byte(jsonStr), &raw); err != nil { + return throwError(info, "Create: invalid options JSON: "+err.Error()) + } + + opts := sandbox.CreateOptions{} + if v, ok := raw["id"].(string); ok { + opts.ID = v + } + if v, ok := raw["owner"].(string); ok { + opts.Owner = v + } + if v, ok := raw["node_id"].(string); ok { + opts.NodeID = v + } + if v, ok := raw["image"].(string); ok { + opts.Image = v + } + if v, ok := raw["workdir"].(string); ok { + opts.WorkDir = v + } + if v, ok := raw["user"].(string); ok { + opts.User = v + } + if v, ok := raw["env"].(map[string]interface{}); ok { + env := make(map[string]string, len(v)) + for k, val := range v { + if s, ok := val.(string); ok { + env[k] = s + } + } + opts.Env = env + } + if v, ok := raw["memory"].(float64); ok { + opts.Memory = int64(v) + } + if v, ok := raw["cpus"].(float64); ok { + opts.CPUs = v + } + if v, ok := raw["vnc"].(bool); ok { + opts.VNC = v + } + if v, ok := raw["policy"].(string); ok { + opts.Policy = sandbox.LifecyclePolicy(v) + } + if v, ok := raw["idle_timeout"].(float64); ok { + opts.IdleTimeout = time.Duration(v) * time.Millisecond + } + if v, ok := raw["stop_timeout"].(float64); ok { + opts.StopTimeout = time.Duration(v) * time.Millisecond + } + if v, ok := raw["workspace_id"].(string); ok { + opts.WorkspaceID = v + } + if v, ok := raw["mount_mode"].(string); ok { + opts.MountMode = v + } + if v, ok := raw["mount_path"].(string); ok { + opts.MountPath = v + } + if v, ok := raw["labels"].(map[string]interface{}); ok { + labels := make(map[string]string, len(v)) + for k, val := range v { + if s, ok := val.(string); ok { + labels[k] = s + } + } + opts.Labels = labels + } + if v, ok := raw["ports"].([]interface{}); ok { + for _, p := range v { + pm, ok := p.(map[string]interface{}) + if !ok { + continue + } + mapping := sandbox.PortMapping{} + if cp, ok := pm["container_port"].(float64); ok { + mapping.ContainerPort = int(cp) + } + if hp, ok := pm["host_port"].(float64); ok { + mapping.HostPort = int(hp) + } + if hi, ok := pm["host_ip"].(string); ok { + mapping.HostIP = hi + } + if pr, ok := pm["protocol"].(string); ok { + mapping.Protocol = pr + } + opts.Ports = append(opts.Ports, mapping) + } + } + + m := sandbox.M() + var box *sandbox.Box + if opts.ID != "" { + box, err = m.GetOrCreate(ctx, opts) + } else { + box, err = m.Create(ctx, opts) + } + if err != nil { + return throwError(info, err.Error()) + } + + val, err := NewComputerObject(v8ctx, "box", box.ID()) + if err != nil { + return throwError(info, err.Error()) + } + return val } -// sbGet: `sandbox.Get(id)` → Box | null -// -// Go: Manager.Get(ctx, id) (*Box, error) -// -// Args: -// -// id: string — sandbox ID -// -// Returns: Box object if found, null if not found +// sbGet: `sandbox.Get(id)` → Computer (kind="box") | null func sbGet(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. id = info.Args()[0].String() - // 2. box, err := sandbox.M().Get(ctx, id) - // 3. Return NewBoxObject(v8ctx, id) or null - return v8go.Undefined(info.Context().Isolate()) + iso := info.Context().Isolate() + v8ctx := info.Context() + ctx := context.Background() + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "Get requires id (string)") + } + id := args[0].String() + + _, err := sandbox.M().Get(ctx, id) + if err != nil { + return v8go.Null(iso) + } + + val, err := NewComputerObject(v8ctx, "box", id) + if err != nil { + return throwError(info, err.Error()) + } + return val } // sbList: `sandbox.List(filter?)` → BoxInfo[] -// -// Go: Manager.List(ctx, ListOptions) ([]*Box, error) -// -// then Box.Info(ctx) for each → BoxInfo -// -// JS filter → Go ListOptions mapping: -// -// { -// owner: string → ListOptions.Owner // filter by owner; empty = all -// pool: string → ListOptions.Pool // filter by pool; empty = all -// labels: object → ListOptions.Labels // filter by labels -// } -// -// Returns: BoxInfo[] — each element: -// -// { -// id: string ← BoxInfo.ID -// container_id: string ← BoxInfo.ContainerID -// pool: string ← BoxInfo.Pool -// owner: string ← BoxInfo.Owner -// status: string ← BoxInfo.Status -// image: string ← BoxInfo.Image -// vnc: boolean ← BoxInfo.VNC -// policy: string ← BoxInfo.Policy -// labels: object ← BoxInfo.Labels -// created_at: string ← BoxInfo.CreatedAt (ISO 8601) -// last_active: string ← BoxInfo.LastActive (ISO 8601) -// process_count: number ← BoxInfo.ProcessCount -// } func sbList(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. Parse optional filter from info.Args()[0] - // 2. boxes := sandbox.M().List(ctx, opts) - // 3. For each box: box.Info(ctx) → BoxInfo → JS object - // 4. Return JS array of BoxInfo objects - return v8go.Undefined(info.Context().Isolate()) + v8ctx := info.Context() + ctx := context.Background() + args := info.Args() + + opts := sandbox.ListOptions{} + if len(args) > 0 && args[0].IsObject() { + jsonStr, _ := v8go.JSONStringify(v8ctx, args[0]) + var raw map[string]interface{} + if json.Unmarshal([]byte(jsonStr), &raw) == nil { + if v, ok := raw["owner"].(string); ok { + opts.Owner = v + } + if v, ok := raw["node_id"].(string); ok { + opts.NodeID = v + } + if v, ok := raw["labels"].(map[string]interface{}); ok { + labels := make(map[string]string, len(v)) + for k, val := range v { + if s, ok := val.(string); ok { + labels[k] = s + } + } + opts.Labels = labels + } + } + } + + boxes, err := sandbox.M().List(ctx, opts) + if err != nil { + return throwError(info, err.Error()) + } + + items := make([]interface{}, 0, len(boxes)) + for _, b := range boxes { + bi, err := b.Info(ctx) + if err != nil { + continue + } + items = append(items, map[string]interface{}{ + "id": bi.ID, + "container_id": bi.ContainerID, + "node_id": bi.NodeID, + "owner": bi.Owner, + "status": bi.Status, + "image": bi.Image, + "vnc": bi.VNC, + "policy": string(bi.Policy), + "labels": bi.Labels, + "created_at": bi.CreatedAt.Format(time.RFC3339), + "last_active": bi.LastActive.Format(time.RFC3339), + "process_count": bi.ProcessCount, + }) + } + + data, _ := json.Marshal(items) + val, _ := v8go.JSONParse(v8ctx, string(data)) + return val } // sbDelete: `sandbox.Delete(id)` → void -// -// Go: Manager.Remove(ctx, id) error -// -// Args: -// -// id: string — sandbox ID to remove func sbDelete(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. id = info.Args()[0].String() - // 2. sandbox.M().Remove(ctx, id) - return v8go.Undefined(info.Context().Isolate()) + iso := info.Context().Isolate() + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "Delete requires id (string)") + } + ctx := context.Background() + id := args[0].String() + + if err := sandbox.M().Remove(ctx, id); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) } diff --git a/sandbox/v2/jsapi/jsapi_test.go b/sandbox/v2/jsapi/jsapi_test.go new file mode 100644 index 00000000..567dfc2e --- /dev/null +++ b/sandbox/v2/jsapi/jsapi_test.go @@ -0,0 +1,418 @@ +package jsapi_test + +import ( + "fmt" + "os" + "strings" + "testing" + "time" + + v8runtime "github.com/yaoapp/gou/runtime/v8" + "github.com/yaoapp/yao/config" + sandbox "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" + "github.com/yaoapp/yao/test" + + _ "github.com/yaoapp/yao/sandbox/v2/jsapi" +) + +type testMode struct { + Name string + Addr string + TaiID string // filled by setupSandbox + Options []tai.Option +} + +func testModes() []testMode { + modes := []testMode{{Name: "local", Addr: "local"}} + if addr := os.Getenv("SANDBOX_TEST_REMOTE_ADDR"); addr != "" { + modes = append(modes, testMode{Name: "remote", Addr: addr}) + } + return modes +} + +func testImage() string { + if img := os.Getenv("SANDBOX_TEST_IMAGE"); img != "" { + return img + } + return "alpine:latest" +} + +func setupSandbox(t *testing.T, m *testMode) { + t.Helper() + test.Prepare(t, config.Conf) + + reg := registry.Global() + if reg == nil { + registry.Init(nil) + } + + client, err := tai.New(m.Addr, m.Options...) + if err != nil { + t.Fatalf("tai.New: %v", err) + } + m.TaiID = client.TaiID() + + sandbox.Init() + mgr := sandbox.M() + t.Cleanup(func() { mgr.Close() }) +} + +func runJS(t *testing.T, source string) interface{} { + t.Helper() + res, err := v8runtime.Call(v8runtime.CallOptions{ + Sid: "test", + Timeout: 60 * time.Second, + }, source) + if err != nil { + t.Fatalf("JS error: %v", err) + } + return res +} + +func runJSExpectError(t *testing.T, source string) string { + t.Helper() + _, err := v8runtime.Call(v8runtime.CallOptions{ + Sid: "test", + Timeout: 30 * time.Second, + }, source) + if err == nil { + t.Fatal("expected JS error, got nil") + } + return err.Error() +} + +func skipIfNoDocker(t *testing.T) { + t.Helper() + addr := os.Getenv("SANDBOX_TEST_LOCAL_ADDR") + if addr == "" { + addr = "local" + } + _ = addr +} + +// --------------------------------------------------------------------------- +// sandbox.Create / sandbox.Get / sandbox.Delete +// --------------------------------------------------------------------------- + +func TestCreate(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestCreate() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + if (pc.kind !== "box") throw new Error("kind=" + pc.kind); + if (!pc.id) throw new Error("no id"); + var id = pc.id; + sandbox.Delete(id); + return id; + }`, img, m.TaiID)) + if res == nil || res == "" { + t.Error("expected box id") + } + }) + } +} + +func TestGet(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestGet() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var id = pc.id; + var got = sandbox.Get(id); + if (!got) throw new Error("Get returned null"); + if (got.kind !== "box") throw new Error("kind=" + got.kind); + sandbox.Delete(id); + return id; + }`, img, m.TaiID)) + if res == nil || res == "" { + t.Error("expected box id") + } + }) + } +} + +func TestGetNotFound(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + res := runJS(t, `function TestGetNotFound() { + var got = sandbox.Get("sb-nonexistent-id"); + return got === null ? "null" : "found"; + }`) + if res != "null" { + t.Errorf("expected null, got %v", res) + } + }) + } +} + +func TestDelete(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestDelete() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var id = pc.id; + sandbox.Delete(id); + var got = sandbox.Get(id); + return got === null ? "deleted" : "still exists"; + }`, img, m.TaiID)) + if res != "deleted" { + t.Errorf("expected deleted, got %v", res) + } + }) + } +} + +// --------------------------------------------------------------------------- +// sandbox.List +// --------------------------------------------------------------------------- + +func TestList(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestList() { + var a = sandbox.Create({ image: "%s", owner: "list-user", node_id: "%s" }); + var b = sandbox.Create({ image: "%s", owner: "list-user", node_id: "%s" }); + var list = sandbox.List({ owner: "list-user" }); + var count = list.length; + sandbox.Delete(a.id); + sandbox.Delete(b.id); + return count; + }`, img, m.TaiID, img, m.TaiID)) + n := toInt(res) + if n < 2 { + t.Errorf("expected >= 2, got %d", n) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Computer.Exec +// --------------------------------------------------------------------------- + +func TestExec(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestExec() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var r = pc.Exec(["echo", "hello-jsapi"]); + sandbox.Delete(pc.id); + return r.stdout; + }`, img, m.TaiID)) + s := fmt.Sprintf("%v", res) + if !strings.Contains(s, "hello-jsapi") { + t.Errorf("stdout = %q, want contain 'hello-jsapi'", s) + } + }) + } +} + +func TestExecWithOptions(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestExecWithOptions() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var r = pc.Exec(["pwd"], { workdir: "/tmp" }); + sandbox.Delete(pc.id); + return r.stdout; + }`, img, m.TaiID)) + s := fmt.Sprintf("%v", res) + if !strings.Contains(s, "/tmp") { + t.Errorf("stdout = %q, want contain '/tmp'", s) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Computer.Stream +// --------------------------------------------------------------------------- + +func TestStream(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestStream() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var chunks = []; + var exitCode = -1; + pc.Stream(["echo", "streaming"], function(type, data) { + if (type === "stdout") chunks.push(data); + if (type === "exit") exitCode = data; + }); + sandbox.Delete(pc.id); + return chunks.join("").trim() + "|" + exitCode; + }`, img, m.TaiID)) + s := fmt.Sprintf("%v", res) + if !strings.Contains(s, "streaming|0") { + t.Errorf("result = %q, want contain 'streaming|0'", s) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Computer.ComputerInfo +// --------------------------------------------------------------------------- + +func TestComputerInfo(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestComputerInfo() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var info = pc.ComputerInfo(); + sandbox.Delete(pc.id); + return info.kind; + }`, img, m.TaiID)) + if res != "box" { + t.Errorf("kind = %q, want 'box'", res) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Computer.Info (box-only) +// --------------------------------------------------------------------------- + +func TestBoxInfo(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestBoxInfo() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var info = pc.Info(); + sandbox.Delete(pc.id); + return info.id ? "ok" : "no-id"; + }`, img, m.TaiID)) + if res != "ok" { + t.Errorf("expected ok, got %v", res) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Box-only method on host → error +// --------------------------------------------------------------------------- + +func TestHostBoxMethodsThrow(t *testing.T) { + if os.Getenv("SANDBOX_TEST_REMOTE_ADDR") == "" { + t.Skip("no remote host configured") + } + for _, m := range testModes() { + if m.Name == "local" { + continue + } + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + errMsg := runJSExpectError(t, fmt.Sprintf(`function TestHostBoxMethodsThrow() { + var host = sandbox.Host("%s"); + host.Info(); + }`, m.TaiID)) + if !strings.Contains(errMsg, "not supported") { + t.Errorf("expected 'not supported' error, got: %s", errMsg) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Computer.kind property +// --------------------------------------------------------------------------- + +func TestComputerKind(t *testing.T) { + skipIfNoDocker(t) + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupSandbox(t, &m) + img := testImage() + res := runJS(t, fmt.Sprintf(`function TestComputerKind() { + var pc = sandbox.Create({ image: "%s", owner: "test-user", node_id: "%s" }); + var k = pc.kind; + sandbox.Delete(pc.id); + return k; + }`, img, m.TaiID)) + if res != "box" { + t.Errorf("kind = %q, want 'box'", res) + } + }) + } +} + +// --------------------------------------------------------------------------- +// sandbox.Nodes (requires registry) +// --------------------------------------------------------------------------- + +func TestNodes(t *testing.T) { + test.Prepare(t, config.Conf) + registry.Init(nil) + res := runJS(t, `function TestNodes() { + var nodes = sandbox.Nodes(); + return Array.isArray(nodes) ? "array" : typeof nodes; + }`) + if res != "array" { + t.Errorf("expected array, got %v", res) + } +} + +func TestGetNodeNotFound(t *testing.T) { + test.Prepare(t, config.Conf) + registry.Init(nil) + res := runJS(t, `function TestGetNodeNotFound() { + var node = sandbox.GetNode("tai-nonexistent"); + return node === null ? "null" : "found"; + }`) + if res != "null" { + t.Errorf("expected null, got %v", res) + } +} + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +func toInt(v interface{}) int { + switch n := v.(type) { + case int: + return n + case int32: + return int(n) + case int64: + return int(n) + case float64: + return int(n) + case float32: + return int(n) + default: + return 0 + } +} diff --git a/sandbox/v2/jsapi/node.go b/sandbox/v2/jsapi/node.go new file mode 100644 index 00000000..37dc6139 --- /dev/null +++ b/sandbox/v2/jsapi/node.go @@ -0,0 +1,143 @@ +package jsapi + +import ( + "encoding/json" + "time" + + "github.com/yaoapp/yao/tai/registry" + "rogchap.com/v8go" +) + +// sbGetNode: `sandbox.GetNode(taiID)` → NodeInfo | null +func sbGetNode(info *v8go.FunctionCallbackInfo) *v8go.Value { + iso := info.Context().Isolate() + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "GetNode requires taiID (string)") + } + + reg := registry.Global() + if reg == nil { + return throwError(info, "registry not initialized") + } + + snap, ok := reg.Get(args[0].String()) + if !ok { + return v8go.Null(iso) + } + + val, err := snapshotToJS(info.Context(), snap) + if err != nil { + return throwError(info, err.Error()) + } + return val +} + +// sbNodes: `sandbox.Nodes()` → NodeInfo[] +func sbNodes(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + + reg := registry.Global() + if reg == nil { + return throwError(info, "registry not initialized") + } + + snaps := reg.List() + return snapshotsToJSArray(v8ctx, snaps) +} + +// sbNodesByTeam: `sandbox.NodesByTeam(teamID)` → NodeInfo[] +func sbNodesByTeam(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "NodesByTeam requires teamID (string)") + } + + reg := registry.Global() + if reg == nil { + return throwError(info, "registry not initialized") + } + + snaps := reg.ListByTeam(args[0].String()) + return snapshotsToJSArray(v8ctx, snaps) +} + +// snapshotToJS converts a NodeSnapshot to a JS NodeInfo object. +// Auth and YaoBase are excluded for security. +func snapshotToJS(v8ctx *v8go.Context, snap *registry.NodeSnapshot) (*v8go.Value, error) { + ports := make(map[string]interface{}, len(snap.Ports)) + for k, v := range snap.Ports { + ports[k] = v + } + + caps := make(map[string]interface{}, len(snap.Capabilities)) + for k, v := range snap.Capabilities { + caps[k] = v + } + + data, err := json.Marshal(map[string]interface{}{ + "tai_id": snap.TaiID, + "machine_id": snap.MachineID, + "version": snap.Version, + "mode": snap.Mode, + "addr": snap.Addr, + "status": snap.Status, + "display_name": snap.DisplayName, + "connected_at": snap.ConnectedAt.Format(time.RFC3339), + "last_ping": snap.LastPing.Format(time.RFC3339), + "ports": ports, + "capabilities": caps, + "system": map[string]interface{}{ + "os": snap.System.OS, + "arch": snap.System.Arch, + "hostname": snap.System.Hostname, + "num_cpu": snap.System.NumCPU, + "total_mem": snap.System.TotalMem, + }, + }) + if err != nil { + return nil, err + } + + return v8go.JSONParse(v8ctx, string(data)) +} + +func snapshotsToJSArray(v8ctx *v8go.Context, snaps []registry.NodeSnapshot) *v8go.Value { + items := make([]interface{}, 0, len(snaps)) + for i := range snaps { + snap := &snaps[i] + ports := make(map[string]interface{}, len(snap.Ports)) + for k, v := range snap.Ports { + ports[k] = v + } + caps := make(map[string]interface{}, len(snap.Capabilities)) + for k, v := range snap.Capabilities { + caps[k] = v + } + items = append(items, map[string]interface{}{ + "tai_id": snap.TaiID, + "node_id": snap.TaiID, + "machine_id": snap.MachineID, + "version": snap.Version, + "mode": snap.Mode, + "addr": snap.Addr, + "status": snap.Status, + "display_name": snap.DisplayName, + "connected_at": snap.ConnectedAt.Format(time.RFC3339), + "last_ping": snap.LastPing.Format(time.RFC3339), + "ports": ports, + "capabilities": caps, + "system": map[string]interface{}{ + "os": snap.System.OS, + "arch": snap.System.Arch, + "hostname": snap.System.Hostname, + "num_cpu": snap.System.NumCPU, + "total_mem": snap.System.TotalMem, + }, + }) + } + data, _ := json.Marshal(items) + val, _ := v8go.JSONParse(v8ctx, string(data)) + return val +} diff --git a/sandbox/v2/manager.go b/sandbox/v2/manager.go index 5bd85760..1c9659eb 100644 --- a/sandbox/v2/manager.go +++ b/sandbox/v2/manager.go @@ -3,53 +3,46 @@ package sandbox import ( "context" "fmt" + "path/filepath" "sync" "time" + "github.com/yaoapp/yao/config" "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" taisandbox "github.com/yaoapp/yao/tai/sandbox" "github.com/yaoapp/yao/workspace" ) -// Manager manages a pool of tai.Client connections and sandbox lifecycle. +// Manager manages sandbox lifecycle. Node connections are delegated to tai/registry. type Manager struct { - pool map[string]*tai.Client - poolDefs []Pool - defaultPool string - config Config - boxes sync.Map - mu sync.Mutex - cancel context.CancelFunc - grpcPort int - wsManager *workspace.Manager + boxes sync.Map + mu sync.Mutex + cancel context.CancelFunc } -func newManager(cfg Config) (*Manager, error) { - m := &Manager{ - pool: make(map[string]*tai.Client), - poolDefs: cfg.Pool, - config: cfg, - grpcPort: 9099, - } - if len(cfg.Pool) > 0 { - m.defaultPool = cfg.Pool[0].Name - } - return m, nil +func newManager() *Manager { + return &Manager{} } -// Start discovers existing containers from all pools, rebuilds the boxes map, -// and starts the cleanup loop. +// Start discovers existing containers from all registered nodes, rebuilds +// the boxes map, and starts the cleanup loop. +// If no "local" node is registered yet, it probes the local Docker environment +// and auto-registers one when available. func (m *Manager) Start(ctx context.Context) error { - if len(m.poolDefs) == 0 { + reg := registry.Global() + if reg == nil { return nil } - for _, pd := range m.poolDefs { - client, err := m.getPool(pd.Name) + m.ensureLocalNode(reg) + + for _, snap := range reg.List() { + client, err := m.getNode(snap.TaiID) if err != nil { continue } - m.recoverBoxes(ctx, &pd, client) + m.recoverBoxes(ctx, snap.TaiID, client) } loopCtx, cancel := context.WithCancel(ctx) @@ -58,96 +51,22 @@ func (m *Manager) Start(ctx context.Context) error { return nil } -// AddPool registers a new pool at runtime. -func (m *Manager) AddPool(_ context.Context, p Pool) error { - m.mu.Lock() - defer m.mu.Unlock() - - for _, pd := range m.poolDefs { - if pd.Name == p.Name { - return fmt.Errorf("sandbox: pool %q already exists", p.Name) - } - } - m.poolDefs = append(m.poolDefs, p) - if m.defaultPool == "" { - m.defaultPool = p.Name - } - return nil +// ensureLocalNode delegates to tai.RegisterLocal() which probes the local +// Docker environment and registers a "local" node in the registry if available. +// The workspace data directory is derived from config.Conf.DataRoot so that +// workspace files persist across restarts. +func (m *Manager) ensureLocalNode(_ *registry.Registry) { + dataDir := filepath.Join(config.Conf.DataRoot, "workspaces") + tai.RegisterLocal(tai.WithDataDir(dataDir)) } -// RemovePool removes a pool by name. -func (m *Manager) RemovePool(ctx context.Context, name string, force bool) error { - m.mu.Lock() - defer m.mu.Unlock() - - idx := -1 - for i, pd := range m.poolDefs { - if pd.Name == name { - idx = i - break - } +// Nodes returns the list of registered Tai nodes from the registry. +func (m *Manager) Nodes() []registry.NodeSnapshot { + reg := registry.Global() + if reg == nil { + return nil } - if idx < 0 { - return ErrPoolNotFound - } - - count := 0 - m.boxes.Range(func(_, value any) bool { - if value.(*Box).pool == name { - count++ - } - return true - }) - - if count > 0 && !force { - return ErrPoolInUse - } - - if count > 0 { - m.boxes.Range(func(key, value any) bool { - b := value.(*Box) - if b.pool == name { - b.Remove(ctx) - } - return true - }) - } - - m.poolDefs = append(m.poolDefs[:idx], m.poolDefs[idx+1:]...) - if client, ok := m.pool[name]; ok { - client.Close() - delete(m.pool, name) - } - return nil -} - -// Pools returns all registered pool names and their status. -func (m *Manager) Pools() []PoolInfo { - m.mu.Lock() - defer m.mu.Unlock() - - result := make([]PoolInfo, 0, len(m.poolDefs)) - for _, pd := range m.poolDefs { - _, connected := m.pool[pd.Name] - count := 0 - m.boxes.Range(func(_, value any) bool { - if value.(*Box).pool == pd.Name { - count++ - } - return true - }) - result = append(result, PoolInfo{ - Name: pd.Name, - Addr: pd.Addr, - Connected: connected, - Boxes: count, - MaxPerUser: pd.MaxPerUser, - MaxTotal: pd.MaxTotal, - IdleTimeout: pd.IdleTimeout, - MaxLifetime: pd.MaxLifetime, - }) - } - return result + return reg.List() } // Heartbeat updates the box's last heartbeat timestamp. @@ -165,62 +84,80 @@ func (m *Manager) Heartbeat(sandboxID string, active bool, processCount int) err } // Host returns a Host handle for executing commands on the Tai host machine. -// The pool must be connected to a Tai server with host_exec capability. -// Unlike Create/Box, Host does not create a container — it is available -// immediately as long as the pool is reachable. -func (m *Manager) Host(_ context.Context, pool string) (*Host, error) { - if pool == "" { - pool = m.defaultPool +func (m *Manager) Host(_ context.Context, nodeID string) (*Host, error) { + if nodeID == "" { + return nil, ErrNodeMissing } - pd := m.findPoolDef(pool) - if pd == nil { - return nil, ErrPoolNotFound - } - - client, err := m.getPool(pool) + client, err := m.getNode(nodeID) if err != nil { - return nil, fmt.Errorf("sandbox: connect pool %q: %w", pool, err) + return nil, fmt.Errorf("sandbox: connect node %q: %w", nodeID, err) } if client.HostExec() == nil { - return nil, fmt.Errorf("sandbox: pool %q has no host_exec capability", pool) + return nil, fmt.Errorf("sandbox: node %q has no host_exec capability", nodeID) } - return &Host{pool: pool, manager: m}, nil + var sys SystemInfo + if snap, ok := tai.GetNodeSnapshot(nodeID); ok { + sys = SystemInfo{ + OS: snap.System.OS, + Arch: snap.System.Arch, + Hostname: snap.System.Hostname, + NumCPU: snap.System.NumCPU, + TotalMem: snap.System.TotalMem, + Shell: snap.System.Shell, + TempDir: snap.System.TempDir, + } + } + + return &Host{nodeID: nodeID, system: sys, manager: m}, nil } // Create creates and starts a new sandbox. func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) { - if len(m.poolDefs) == 0 { - return nil, ErrNotAvailable - } if opts.Image == "" { return nil, fmt.Errorf("sandbox: image is required") } - poolName := opts.Pool - if poolName == "" { - poolName = m.defaultPool - } + nodeID := opts.NodeID - // Workspace node binding: when WorkspaceID is set, resolve the workspace's - // bound node and force the container onto that pool. - if opts.WorkspaceID != "" && m.wsManager != nil { - node, err := m.wsManager.NodeForWorkspace(ctx, opts.WorkspaceID) - if err != nil { - return nil, fmt.Errorf("sandbox: resolve workspace %q: %w", opts.WorkspaceID, err) + if opts.WorkspaceID != "" { + if wsm := workspace.M(); wsm != nil { + node, err := wsm.NodeForWorkspace(ctx, opts.WorkspaceID) + if err != nil { + targetNode := nodeID + if targetNode == "" { + if nodes := wsm.Nodes(); len(nodes) > 0 { + for _, n := range nodes { + if n.Online { + targetNode = n.Name + break + } + } + } + } + if targetNode == "" { + return nil, fmt.Errorf("sandbox: resolve workspace %q: no available node", opts.WorkspaceID) + } + _, err = wsm.Create(ctx, workspace.CreateOptions{ + ID: opts.WorkspaceID, + Name: opts.WorkspaceID, + Owner: opts.Owner, + Node: targetNode, + }) + if err != nil { + return nil, fmt.Errorf("sandbox: auto-create workspace %q: %w", opts.WorkspaceID, err) + } + nodeID = targetNode + } else { + nodeID = node + } } - poolName = node } - pd := m.findPoolDef(poolName) - if pd == nil { - return nil, ErrPoolNotFound - } - - if err := m.checkLimits(pd, opts.Owner); err != nil { - return nil, err + if nodeID == "" { + return nil, ErrNodeMissing } id := opts.ID @@ -228,21 +165,16 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) id = fmt.Sprintf("sb-%d", time.Now().UnixNano()) } - client, err := m.getPool(poolName) + client, err := m.getNode(nodeID) if err != nil { - return nil, fmt.Errorf("sandbox: connect pool %q: %w", poolName, err) + return nil, fmt.Errorf("sandbox: connect node %q: %w", nodeID, err) } if client.Sandbox() == nil { - return nil, fmt.Errorf("sandbox: pool %q has no container runtime", poolName) + return nil, fmt.Errorf("sandbox: node %q has no container runtime", nodeID) } - access, refresh, err := CreateContainerTokens(id, opts.Owner, nil) - if err != nil { - return nil, fmt.Errorf("sandbox: create tokens: %w", err) - } - - taiOpts := m.buildTaiCreateOptions(opts, pd, id, access, refresh) + taiOpts := m.buildTaiCreateOptions(opts, nodeID, id) containerID, err := client.Sandbox().Create(ctx, taiOpts) if err != nil { @@ -259,21 +191,35 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) policy = Session } + var sys SystemInfo + if snap, ok := tai.GetNodeSnapshot(nodeID); ok { + sys = SystemInfo{ + OS: snap.System.OS, + Arch: snap.System.Arch, + Hostname: snap.System.Hostname, + NumCPU: snap.System.NumCPU, + TotalMem: snap.System.TotalMem, + Shell: snap.System.Shell, + TempDir: snap.System.TempDir, + } + } + box := &Box{ id: id, containerID: containerID, - pool: poolName, + nodeID: nodeID, owner: opts.Owner, policy: policy, labels: opts.Labels, idleTimeoutD: opts.IdleTimeout, + maxLifetimeD: opts.MaxLifetime, stopTimeoutD: opts.StopTimeout, createdAt: time.Now(), - refreshToken: refresh, manager: m, vnc: opts.VNC, image: opts.Image, workspaceID: opts.WorkspaceID, + system: sys, } box.lastCall.Store(time.Now().UnixMilli()) @@ -308,7 +254,7 @@ func (m *Manager) List(_ context.Context, opts ListOptions) ([]*Box, error) { if opts.Owner != "" && b.owner != opts.Owner { return true } - if opts.Pool != "" && b.pool != opts.Pool { + if opts.NodeID != "" && b.nodeID != opts.NodeID { return true } if len(opts.Labels) > 0 { @@ -332,15 +278,11 @@ func (m *Manager) Remove(ctx context.Context, id string) error { } b := v.(*Box) - client, err := m.getPool(b.pool) + client, err := m.getNode(b.nodeID) if err == nil && client.Sandbox() != nil { client.Sandbox().Remove(ctx, b.containerID, true) } - if b.refreshToken != "" { - RevokeContainerTokens(b.refreshToken) - } - m.boxes.Delete(id) return nil } @@ -361,7 +303,7 @@ func (m *Manager) Cleanup(ctx context.Context) error { } case LongRunning: if timeout := b.idleTimeout(); timeout > 0 && idle > timeout { - if client, err := m.getPool(b.pool); err == nil && client.Sandbox() != nil { + if client, err := m.getNode(b.nodeID); err == nil && client.Sandbox() != nil { client.Sandbox().Stop(ctx, b.containerID, b.stopTimeout()) } } @@ -376,32 +318,14 @@ func (m *Manager) Cleanup(ctx context.Context) error { return nil } -// Close stops the cleanup loop and releases all pool connections. +// Close stops the cleanup loop. Node connections are managed by the registry. func (m *Manager) Close() error { if m.cancel != nil { m.cancel() } - m.mu.Lock() - defer m.mu.Unlock() - for name, client := range m.pool { - client.Close() - delete(m.pool, name) - } return nil } -// SetGRPCPort sets the local gRPC port for container env injection. -func (m *Manager) SetGRPCPort(port int) { - m.grpcPort = port -} - -// SetWorkspaceManager links the workspace manager for workspace-aware container creation. -// When CreateOptions.WorkspaceID is set, the sandbox Manager uses the workspace Manager -// to resolve the workspace's bound node and force container routing. -func (m *Manager) SetWorkspaceManager(wm *workspace.Manager) { - m.wsManager = wm -} - func (m *Manager) cleanupLoop(ctx context.Context) { ticker := time.NewTicker(1 * time.Minute) defer ticker.Stop() @@ -415,88 +339,37 @@ func (m *Manager) cleanupLoop(ctx context.Context) { } } -func (m *Manager) getPool(name string) (*tai.Client, error) { - m.mu.Lock() - defer m.mu.Unlock() - - if client, ok := m.pool[name]; ok { - return client, nil +func (m *Manager) getNode(name string) (*tai.Client, error) { + client, ok := tai.GetClient(name) + if !ok { + return nil, ErrNodeNotFound } - - pd := m.findPoolDefLocked(name) - if pd == nil { - return nil, ErrPoolNotFound - } - - client, err := tai.New(pd.Addr, pd.Options...) - if err != nil { - return nil, err - } - m.pool[name] = client return client, nil } -func (m *Manager) findPoolDef(name string) *Pool { - m.mu.Lock() - defer m.mu.Unlock() - return m.findPoolDefLocked(name) -} - -func (m *Manager) findPoolDefLocked(name string) *Pool { - for i := range m.poolDefs { - if m.poolDefs[i].Name == name { - return &m.poolDefs[i] - } - } - return nil -} - -func (m *Manager) checkLimits(pd *Pool, owner string) error { - if pd.MaxTotal > 0 { - count := 0 - m.boxes.Range(func(_, value any) bool { - if value.(*Box).pool == pd.Name { - count++ - } - return true - }) - if count >= pd.MaxTotal { - return ErrLimitExceeded - } - } - - if pd.MaxPerUser > 0 && owner != "" { - count := 0 - m.boxes.Range(func(_, value any) bool { - b := value.(*Box) - if b.pool == pd.Name && b.owner == owner { - count++ - } - return true - }) - if count >= pd.MaxPerUser { - return ErrLimitExceeded - } - } - return nil -} - -func (m *Manager) buildTaiCreateOptions(opts CreateOptions, pd *Pool, sandboxID, access, refresh string) taisandbox.CreateOptions { +func (m *Manager) buildTaiCreateOptions(opts CreateOptions, nodeID, sandboxID string) taisandbox.CreateOptions { env := make(map[string]string) - for k, v := range opts.Env { - env[k] = v + + reg := registry.Global() + if reg != nil { + if snap, ok := reg.Get(nodeID); ok { + grpcEnv := BuildGRPCEnv(snap.Mode, snap.Addr, sandboxID) + for k, v := range grpcEnv { + env[k] = v + } + } } - grpcEnv := BuildGRPCEnv(pd, sandboxID, access, refresh, m.grpcPort) - for k, v := range grpcEnv { + + for k, v := range opts.Env { env[k] = v } labels := map[string]string{ - "managed-by": "yao-sandbox", - "sandbox-id": sandboxID, - "sandbox-owner": opts.Owner, - "sandbox-pool": pd.Name, - "sandbox-policy": string(opts.Policy), + "managed-by": "yao-sandbox", + "sandbox-id": sandboxID, + "sandbox-owner": opts.Owner, + "sandbox-node-id": nodeID, + "sandbox-policy": string(opts.Policy), } if opts.WorkspaceID != "" { labels["workspace-id"] = opts.WorkspaceID @@ -522,20 +395,21 @@ func (m *Manager) buildTaiCreateOptions(opts CreateOptions, pd *Pool, sandboxID, }) } - // Workspace bind mount var binds []string - if opts.WorkspaceID != "" && m.wsManager != nil { - mountPath := opts.MountPath - if mountPath == "" { - mountPath = "/workspace" - } - mode := opts.MountMode - if mode == "" { - mode = "rw" - } - hostPath, _ := m.wsManager.MountPath(context.Background(), opts.WorkspaceID) - if hostPath != "" { - binds = append(binds, fmt.Sprintf("%s:%s:%s", hostPath, mountPath, mode)) + if opts.WorkspaceID != "" { + if wsm := workspace.M(); wsm != nil { + mountPath := opts.MountPath + if mountPath == "" { + mountPath = "/workspace" + } + mode := opts.MountMode + if mode == "" { + mode = "rw" + } + hostPath, _ := wsm.MountPath(context.Background(), opts.WorkspaceID) + if hostPath != "" { + binds = append(binds, fmt.Sprintf("%s:%s:%s", hostPath, mountPath, mode)) + } } } @@ -555,7 +429,7 @@ func (m *Manager) buildTaiCreateOptions(opts CreateOptions, pd *Pool, sandboxID, } } -func (m *Manager) recoverBoxes(ctx context.Context, pd *Pool, client *tai.Client) { +func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, client *tai.Client) { if client.Sandbox() == nil { return } @@ -583,7 +457,7 @@ func (m *Manager) recoverBoxes(ctx context.Context, pd *Pool, client *tai.Client box := &Box{ id: sandboxID, containerID: cid, - pool: c.Labels["sandbox-pool"], + nodeID: c.Labels["sandbox-node-id"], owner: c.Labels["sandbox-owner"], policy: LifecyclePolicy(c.Labels["sandbox-policy"]), labels: c.Labels, @@ -597,11 +471,9 @@ func (m *Manager) recoverBoxes(ctx context.Context, pd *Pool, client *tai.Client } } -// ImageExists reports whether the given image ref exists on the target pool node. -// Returns (true, nil) when the pool has no image service (e.g. K8s — kubelet -// handles image pulls transparently). -func (m *Manager) ImageExists(ctx context.Context, pool, ref string) (bool, error) { - client, err := m.getPool(pool) +// ImageExists reports whether the given image ref exists on the target node. +func (m *Manager) ImageExists(ctx context.Context, nodeID, ref string) (bool, error) { + client, err := m.getNode(nodeID) if err != nil { return false, err } @@ -612,10 +484,10 @@ func (m *Manager) ImageExists(ctx context.Context, pool, ref string) (bool, erro return img.Exists(ctx, ref) } -// PullImage pulls an image to the target pool node, returning a channel of -// real-time progress events. The channel is nil when no pull is needed (e.g. K8s mode). -func (m *Manager) PullImage(ctx context.Context, pool, ref string, opts ImagePullOptions) (<-chan taisandbox.PullProgress, error) { - client, err := m.getPool(pool) +// PullImage pulls an image to the target node, returning a channel of +// real-time progress events. +func (m *Manager) PullImage(ctx context.Context, nodeID, ref string, opts ImagePullOptions) (<-chan taisandbox.PullProgress, error) { + client, err := m.getNode(nodeID) if err != nil { return nil, err } @@ -634,11 +506,10 @@ func (m *Manager) PullImage(ctx context.Context, pool, ref string, opts ImagePul return img.Pull(ctx, ref, pullOpts) } -// EnsureImage checks whether the image exists on the pool node; if not, it -// pulls the image and blocks until the pull completes. Returns the first -// error encountered during pull. For K8s pools this is a no-op. -func (m *Manager) EnsureImage(ctx context.Context, pool, ref string, opts ImagePullOptions) error { - exists, err := m.ImageExists(ctx, pool, ref) +// EnsureImage checks whether the image exists on the node; if not, it +// pulls the image and blocks until the pull completes. +func (m *Manager) EnsureImage(ctx context.Context, nodeID, ref string, opts ImagePullOptions) error { + exists, err := m.ImageExists(ctx, nodeID, ref) if err != nil { return fmt.Errorf("image exists check: %w", err) } @@ -646,7 +517,7 @@ func (m *Manager) EnsureImage(ctx context.Context, pool, ref string, opts ImageP return nil } - ch, err := m.PullImage(ctx, pool, ref, opts) + ch, err := m.PullImage(ctx, nodeID, ref, opts) if err != nil { return fmt.Errorf("image pull: %w", err) } diff --git a/sandbox/v2/manager_lifecycle_test.go b/sandbox/v2/manager_lifecycle_test.go index 06233edf..67412792 100644 --- a/sandbox/v2/manager_lifecycle_test.go +++ b/sandbox/v2/manager_lifecycle_test.go @@ -11,10 +11,11 @@ import ( func TestHeartbeatUpdates(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) err := m.Heartbeat(box.ID(), true, 5) if err != nil { @@ -33,9 +34,10 @@ func TestHeartbeatUpdates(t *testing.T) { } func TestHeartbeatUnknownBox(t *testing.T) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) err := m.Heartbeat("nonexistent", true, 1) if err != sandbox.ErrNotFound { t.Errorf("err = %v, want ErrNotFound", err) @@ -47,18 +49,19 @@ func TestHeartbeatUnknownBox(t *testing.T) { func TestIdleCleanup(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc, func(p *sandbox.Pool) { - p.IdleTimeout = 1 * time.Second - }) - ensureTestImage(t, m, pc.Name) + m := setupManagerForNode(t, &pc) + ensureTestImage(t, m, pc.TaiID) ctx := context.Background() box, err := m.Create(ctx, sandbox.CreateOptions{ - Image: testImage(), - Owner: "test-user", - Policy: sandbox.Session, + Image: testImage(), + Owner: "test-user", + NodeID: pc.TaiID, + Policy: sandbox.Session, + IdleTimeout: 1 * time.Second, }) if err != nil { t.Fatalf("Create: %v", err) @@ -82,18 +85,14 @@ func TestIdleCleanup(t *testing.T) { func TestStartRecovery(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - pool := sandbox.Pool{Name: pc.Name, Addr: pc.Addr, Options: pc.Options} - - m1 := setupManager(t, pool) - box := createTestBox(t, m1) + m1 := setupManagerForNode(t, &pc) + box := createTestBox(t, m1, pc) boxID := box.ID() - cfg := sandbox.Config{Pool: []sandbox.Pool{pool}} - if err := sandbox.Init(cfg); err != nil { - t.Fatalf("Init2: %v", err) - } + sandbox.Init() m2 := sandbox.M() defer m2.Close() @@ -118,14 +117,14 @@ func TestStartRecovery(t *testing.T) { func TestPersistentNotCleaned(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc, func(p *sandbox.Pool) { - p.IdleTimeout = 1 * time.Second - }) + m := setupManagerForNode(t, &pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.Policy = sandbox.Persistent + co.IdleTimeout = 1 * time.Second }) time.Sleep(2 * time.Second) diff --git a/sandbox/v2/manager_test.go b/sandbox/v2/manager_test.go index a6e6be5f..5c0be967 100644 --- a/sandbox/v2/manager_test.go +++ b/sandbox/v2/manager_test.go @@ -11,10 +11,11 @@ import ( func TestCreateAndExec(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -36,10 +37,11 @@ func TestCreateAndExec(t *testing.T) { func TestCreateWithLabels(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.Labels = map[string]string{"app": "test-app"} }) @@ -58,10 +60,11 @@ func TestCreateWithLabels(t *testing.T) { func TestGet(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m) + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc) got, err := m.Get(context.Background(), box.ID()) if err != nil { @@ -75,9 +78,10 @@ func TestGet(t *testing.T) { } func TestGetNotFound(t *testing.T) { - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) + m := setupManagerForNode(t, &pc) _, err := m.Get(context.Background(), "nonexistent") if err != sandbox.ErrNotFound { t.Errorf("err = %v, want ErrNotFound", err) @@ -89,10 +93,11 @@ func TestGetNotFound(t *testing.T) { func TestList(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - box := createTestBox(t, m, func(co *sandbox.CreateOptions) { + m := setupManagerForNode(t, &pc) + box := createTestBox(t, m, pc, func(co *sandbox.CreateOptions) { co.Owner = "user-list" }) @@ -124,14 +129,16 @@ func TestList(t *testing.T) { func TestRemove(t *testing.T) { skipIfNoDocker(t) - for _, pc := range testPools() { + for _, pc := range testNodes() { + pc := pc t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - ensureTestImage(t, m, pc.Name) + m := setupManagerForNode(t, &pc) + ensureTestImage(t, m, pc.TaiID) ctx := context.Background() box, err := m.Create(ctx, sandbox.CreateOptions{ - Image: testImage(), - Owner: "test-user", + Image: testImage(), + Owner: "test-user", + NodeID: pc.TaiID, }) if err != nil { t.Fatalf("Create: %v", err) @@ -149,110 +156,51 @@ func TestRemove(t *testing.T) { } } -func TestPoolLimits_MaxTotal(t *testing.T) { - skipIfNoDocker(t) - - for _, pc := range testPools() { - t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc, func(p *sandbox.Pool) { - p.MaxTotal = 1 - }) - ensureTestImage(t, m, pc.Name) - - box1 := createTestBox(t, m) - _ = box1 - - ctx := context.Background() - _, err := m.Create(ctx, sandbox.CreateOptions{ - Image: testImage(), - Owner: "test-user", - }) - if err != sandbox.ErrLimitExceeded { - t.Errorf("second Create err = %v, want ErrLimitExceeded", err) - } - }) - } -} - -func TestAddPool(t *testing.T) { - m := setupManager(t, sandbox.Pool{ - Name: "default", - Addr: testLocalAddr(), - }) - - err := m.AddPool(context.Background(), sandbox.Pool{ - Name: "extra", - Addr: testLocalAddr(), - }) - if err != nil { - t.Fatalf("AddPool: %v", err) - } - - pools := m.Pools() - if len(pools) != 2 { - t.Fatalf("Pools() = %d, want 2", len(pools)) - } - - err = m.AddPool(context.Background(), sandbox.Pool{ - Name: "extra", - Addr: testLocalAddr(), - }) - if err == nil { - t.Error("expected error for duplicate pool name") - } -} - func TestCreateNoImage(t *testing.T) { - m := setupManager(t, sandbox.Pool{ - Name: "local", - Addr: testLocalAddr(), - }) + m, nodes := setupManager(t, nodeConfig{Name: "local", Addr: testLocalAddr()}) _, err := m.Create(context.Background(), sandbox.CreateOptions{ - Owner: "test", + Owner: "test", + NodeID: nodes[0].TaiID, }) if err == nil { t.Error("expected error for missing image") } } -func TestCreateNoPools(t *testing.T) { - m := setupManager(t) +func TestCreateNoNodeID(t *testing.T) { + m, _ := setupManager(t, nodeConfig{Name: "local", Addr: testLocalAddr()}) _, err := m.Create(context.Background(), sandbox.CreateOptions{ Image: testImage(), }) - if err != sandbox.ErrNotAvailable { - t.Errorf("err = %v, want ErrNotAvailable", err) + if err != sandbox.ErrNodeMissing { + t.Errorf("err = %v, want ErrNodeMissing", err) } } -func TestMultiPool(t *testing.T) { +func TestMultiNode(t *testing.T) { skipIfNoDocker(t) skipIfNoTai(t) - pools := testPools() - if len(pools) < 2 { - t.Skip("need at least 2 pools (local + remote) for multi-pool test") + nodes := testNodes() + if len(nodes) < 2 { + t.Skip("need at least 2 nodes (local + remote) for multi-node test") } - var sps []sandbox.Pool - for _, pc := range pools { - sps = append(sps, sandbox.Pool{Name: pc.Name, Addr: pc.Addr, Options: pc.Options}) - } - m := setupManager(t, sps...) + m, registered := setupManager(t, nodes...) - for _, pc := range pools { - ensureTestImage(t, m, pc.Name) + for _, pc := range registered { + ensureTestImage(t, m, pc.TaiID) } ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) defer cancel() localBox, err := m.Create(ctx, sandbox.CreateOptions{ - Image: testImage(), - Owner: "test-user", - Pool: "local", + Image: testImage(), + Owner: "test-user", + NodeID: registered[0].TaiID, }) if err != nil { t.Fatalf("Create on local: %v", err) @@ -260,9 +208,9 @@ func TestMultiPool(t *testing.T) { defer m.Remove(ctx, localBox.ID()) remoteBox, err := m.Create(ctx, sandbox.CreateOptions{ - Image: testImage(), - Owner: "test-user", - Pool: "remote", + Image: testImage(), + Owner: "test-user", + NodeID: registered[1].TaiID, }) if err != nil { t.Fatalf("Create on remote: %v", err) diff --git a/sandbox/v2/sandbox.go b/sandbox/v2/sandbox.go index 46be58ea..4c321bb7 100644 --- a/sandbox/v2/sandbox.go +++ b/sandbox/v2/sandbox.go @@ -3,15 +3,9 @@ package sandbox var mgr *Manager // Init initializes the global sandbox Manager. -// Config contains pool definitions. At least one Pool entry is required. -// Pass empty Pool list to disable sandbox (methods return ErrNotAvailable). -func Init(cfg Config) error { - m, err := newManager(cfg) - if err != nil { - return err - } - mgr = m - return nil +// Node discovery is handled by the tai/registry; no configuration is needed. +func Init() { + mgr = newManager() } // M returns the global Manager. Panics if Init was not called. diff --git a/sandbox/v2/sandbox_test.go b/sandbox/v2/sandbox_test.go index 7ef408e5..a4ea649d 100644 --- a/sandbox/v2/sandbox_test.go +++ b/sandbox/v2/sandbox_test.go @@ -7,14 +7,7 @@ import ( ) func TestInit(t *testing.T) { - cfg := sandbox.Config{ - Pool: []sandbox.Pool{ - {Name: "test", Addr: "local"}, - }, - } - if err := sandbox.Init(cfg); err != nil { - t.Fatalf("Init: %v", err) - } + sandbox.Init() m := sandbox.M() if m == nil { t.Fatal("M() returned nil") @@ -22,14 +15,6 @@ func TestInit(t *testing.T) { m.Close() } -func TestInitEmpty(t *testing.T) { - cfg := sandbox.Config{} - if err := sandbox.Init(cfg); err != nil { - t.Fatalf("Init with empty config: %v", err) - } - sandbox.M().Close() -} - func TestMPanicWithoutInit(t *testing.T) { sandbox.ResetForTest() defer func() { diff --git a/sandbox/v2/testutils_test.go b/sandbox/v2/testutils_test.go index 4234542f..9f34526b 100644 --- a/sandbox/v2/testutils_test.go +++ b/sandbox/v2/testutils_test.go @@ -13,8 +13,8 @@ import ( sandbox "github.com/yaoapp/yao/sandbox/v2" "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" taisandbox "github.com/yaoapp/yao/tai/sandbox" - "github.com/yaoapp/yao/tai/volume" "github.com/yaoapp/yao/workspace" ) @@ -31,7 +31,7 @@ func TestMain(m *testing.M) { } // purgeStaleContainers removes leftover sb-* containers/pods from previous -// test runs across all configured pools (Docker + K8s). +// test runs across all configured nodes (Docker + K8s). func purgeStaleContainers() { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() @@ -98,35 +98,30 @@ func purgeStaleContainers() { } } -type poolConfig struct { - Name string +type nodeConfig struct { + Name string // human-readable label for t.Run (e.g. "remote", "k8s") Addr string + TaiID string // actual registry key, filled after tai.New Options []tai.Option } -// testPools returns all available pool configurations for multi-mode testing. -// - local: always present (direct Docker daemon) -// - remote: when SANDBOX_TEST_REMOTE_ADDR is set (Tai on host → Docker) -// - containerized: when TAI_TEST_CONTAINERIZED_HOST is set (Tai in container → Docker) -// - k8s: when TAI_TEST_K8S_HOST + TAI_TEST_KUBECONFIG are set (Tai → K8s) -func testPools() []poolConfig { - pools := []poolConfig{ +// testNodes returns all available node configurations for multi-mode testing. +func testNodes() []nodeConfig { + nodes := []nodeConfig{ {Name: "local", Addr: testLocalAddr()}, } if addr := os.Getenv("SANDBOX_TEST_REMOTE_ADDR"); addr != "" { - pools = append(pools, poolConfig{Name: "remote", Addr: addr}) + nodes = append(nodes, nodeConfig{Name: "remote", Addr: addr}) } if host := os.Getenv("TAI_TEST_CONTAINERIZED_HOST"); host != "" { grpcPort := envPort("TAI_TEST_CONTAINERIZED_GRPC_PORT", 9200) addr := fmt.Sprintf("tai://%s:%d", host, grpcPort) - // No WithPorts for HTTP/VNC — Tai self-inspects its container - // and returns host-mapped ports via ServerInfo automatically. - pools = append(pools, poolConfig{Name: "containerized", Addr: addr}) + nodes = append(nodes, nodeConfig{Name: "containerized", Addr: addr}) } if host := os.Getenv("TAI_TEST_K8S_HOST"); host != "" { kubeconfig := os.Getenv("TAI_TEST_KUBECONFIG") if kubeconfig == "" { - return pools + return nodes } grpcPort := envPort("TAI_TEST_K8S_GRPC_PORT", envPort("TAI_TEST_GRPC_PORT", 19100)) addr := fmt.Sprintf("tai://%s:%d", host, grpcPort) @@ -141,9 +136,9 @@ func testPools() []poolConfig { if ns := os.Getenv("TAI_TEST_K8S_NAMESPACE"); ns != "" { opts = append(opts, tai.WithNamespace(ns)) } - pools = append(pools, poolConfig{Name: "k8s", Addr: addr, Options: opts}) + nodes = append(nodes, nodeConfig{Name: "k8s", Addr: addr, Options: opts}) } - return pools + return nodes } func skipIfNoDocker(t *testing.T) { @@ -164,11 +159,10 @@ func skipIfNoTai(t *testing.T) { type hostExecTarget struct { Name string Addr string // host:port (without tai:// prefix) + TaiID string // filled after registration IsWinNative bool } -// hostExecTargets returns all Tai instances that support HostExec gRPC. -// No container creation needed — these are direct gRPC connections. func hostExecTargets() []hostExecTarget { var targets []hostExecTarget if addr := os.Getenv("SANDBOX_TEST_REMOTE_ADDR"); addr != "" { @@ -195,8 +189,6 @@ func skipIfNoHostExec(t *testing.T) { } } -// linuxCmd adapts a Linux command to the equivalent Windows command for -// Windows native Tai targets. func linuxCmd(tgt hostExecTarget, cmd string, args ...string) (string, []string) { if tgt.IsWinNative { switch cmd { @@ -245,83 +237,94 @@ func envPort(key string, fallback int) int { return fallback } -func setupManager(t *testing.T, pools ...sandbox.Pool) *sandbox.Manager { +// registerNode creates a tai.Client and registers it in the global registry. +// It fills pc.TaiID with the actual registry key returned by tai.New. +func registerNode(t *testing.T, pc *nodeConfig) { t.Helper() - cfg := sandbox.Config{Pool: pools} - if err := sandbox.Init(cfg); err != nil { - t.Fatalf("Init: %v", err) + + reg := registry.Global() + if reg == nil { + registry.Init(nil) } + + client, err := tai.New(pc.Addr, pc.Options...) + if err != nil { + t.Fatalf("tai.New(%s): %v", pc.Addr, err) + } + pc.TaiID = client.TaiID() + t.Cleanup(func() { client.Close() }) +} + +func setupManager(t *testing.T, nodes ...nodeConfig) (*sandbox.Manager, []nodeConfig) { + t.Helper() + + reg := registry.Global() + if reg == nil { + registry.Init(nil) + } + _ = reg + + out := make([]nodeConfig, len(nodes)) + copy(out, nodes) + for i := range out { + client, err := tai.New(out[i].Addr, out[i].Options...) + if err != nil { + t.Fatalf("tai.New(%s): %v", out[i].Addr, err) + } + out[i].TaiID = client.TaiID() + } + + sandbox.Init() m := sandbox.M() - t.Cleanup(func() { - m.Close() - }) + t.Cleanup(func() { m.Close() }) + return m, out +} + +func setupManagerForNode(t *testing.T, pc *nodeConfig) *sandbox.Manager { + t.Helper() + m, registered := setupManager(t, *pc) + *pc = registered[0] return m } -func setupManagerForPool(t *testing.T, pc poolConfig, mutators ...func(*sandbox.Pool)) *sandbox.Manager { +// setupManagerWithWorkspace creates a sandbox Manager and returns +// the global workspace.Manager (which uses the registry for client lookups). +func setupManagerWithWorkspace(t *testing.T, pc *nodeConfig) (*sandbox.Manager, *workspace.Manager) { t.Helper() - pool := sandbox.Pool{Name: pc.Name, Addr: pc.Addr, Options: pc.Options} - for _, fn := range mutators { - fn(&pool) - } - return setupManager(t, pool) + sbm := setupManagerForNode(t, pc) + return sbm, workspace.M() } -// setupManagerWithWorkspace creates a sandbox Manager with a linked workspace Manager. -// Returns both managers and a helper to create workspaces on the given pool's node. -func setupManagerWithWorkspace(t *testing.T, pc poolConfig) (*sandbox.Manager, *workspace.Manager) { - t.Helper() - sbm := setupManagerForPool(t, pc) - - var wsClient *tai.Client - var err error - if pc.Addr == "local" || pc.Addr == "" { - dataDir := t.TempDir() - vol := volume.NewLocal(dataDir) - wsClient, err = tai.New("local", tai.WithVolume(vol), tai.WithDataDir(dataDir)) - } else { - wsClient, err = tai.New(pc.Addr, pc.Options...) - } - if err != nil { - t.Fatalf("tai.New for workspace: %v", err) - } - t.Cleanup(func() { wsClient.Close() }) - - wsm := workspace.NewManager(map[string]*tai.Client{pc.Name: wsClient}) - sbm.SetWorkspaceManager(wsm) - return sbm, wsm -} - -// ensureTestImage guarantees testImage() is available on the given pool before -// container creation. Safe for all modes (Docker pull; K8s no-op). -func ensureTestImage(t *testing.T, m *sandbox.Manager, pool string) { +func ensureTestImage(t *testing.T, m *sandbox.Manager, nodeID string) { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) defer cancel() - if err := m.EnsureImage(ctx, pool, testImage(), sandbox.ImagePullOptions{}); err != nil { - t.Fatalf("EnsureImage(%s, %s): %v", pool, testImage(), err) + if err := m.EnsureImage(ctx, nodeID, testImage(), sandbox.ImagePullOptions{}); err != nil { + t.Fatalf("EnsureImage(%s, %s): %v", nodeID, testImage(), err) } } -func createTestBox(t *testing.T, m *sandbox.Manager, opts ...func(*sandbox.CreateOptions)) *sandbox.Box { +func createTestBox(t *testing.T, m *sandbox.Manager, pc nodeConfig, opts ...func(*sandbox.CreateOptions)) *sandbox.Box { t.Helper() co := sandbox.CreateOptions{ - Image: testImage(), - Owner: "test-user", + Image: testImage(), + Owner: "test-user", + NodeID: pc.TaiID, } for _, fn := range opts { fn(&co) } - pool := co.Pool - if pool == "" { - pools := m.Pools() - if len(pools) > 0 { - pool = pools[0].Name + nodeID := co.NodeID + if nodeID == "" { + nodes := m.Nodes() + if len(nodes) > 0 { + nodeID = nodes[0].TaiID + co.NodeID = nodeID } } - isK8s := pool == "k8s" + isK8s := pc.Name == "k8s" if isK8s { k8sSem <- struct{}{} } @@ -329,12 +332,12 @@ func createTestBox(t *testing.T, m *sandbox.Manager, opts ...func(*sandbox.Creat ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - if pool != "" { - if err := m.EnsureImage(ctx, pool, co.Image, sandbox.ImagePullOptions{}); err != nil { + if nodeID != "" { + if err := m.EnsureImage(ctx, nodeID, co.Image, sandbox.ImagePullOptions{}); err != nil { if isK8s { <-k8sSem } - t.Fatalf("EnsureImage(%s, %s): %v", pool, co.Image, err) + t.Fatalf("EnsureImage(%s, %s): %v", nodeID, co.Image, err) } } diff --git a/sandbox/v2/types.go b/sandbox/v2/types.go index d6263300..1d4a26d4 100644 --- a/sandbox/v2/types.go +++ b/sandbox/v2/types.go @@ -1,12 +1,65 @@ package sandbox import ( + "context" "io" "time" - "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/workspace" ) +// --------------------------------------------------------------------------- +// Computer — unified interface for execution environments +// --------------------------------------------------------------------------- + +// Computer is the unified interface for remote execution environments. +// Both Box (container) and Host (bare metal) implement it. +type Computer interface { + ComputerInfo() ComputerInfo + Exec(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecResult, error) + Stream(ctx context.Context, cmd []string, opts ...ExecOption) (*ExecStream, error) + VNC(ctx context.Context) (string, error) + Proxy(ctx context.Context, port int, path string) (string, error) + BindWorkplace(workspaceID string) + Workplace() workspace.FS +} + +// ComputerInfo holds identity and registry information for a Computer. +type ComputerInfo struct { + Kind string // "box" | "host" + NodeID string + TaiID string + MachineID string + Version string + System SystemInfo + Mode string // "direct" | "tunnel" + Capabilities map[string]bool + Status string + + // Box-specific fields (zero values for Host) + BoxID string + ContainerID string + Owner string + Image string + Policy LifecyclePolicy + Labels map[string]string +} + +// SystemInfo describes the hardware and environment of a Tai node. +type SystemInfo struct { + OS string + Arch string + Hostname string + NumCPU int + TotalMem int64 + Shell string // preferred shell: "sh", "pwsh", "powershell", "cmd.exe" + TempDir string // system temp directory +} + +// --------------------------------------------------------------------------- +// Lifecycle +// --------------------------------------------------------------------------- + type LifecyclePolicy string const ( @@ -18,27 +71,9 @@ const ( const DefaultStopTimeout = 2 * time.Second -type Pool struct { - Name string - Addr string - Options []tai.Option - MaxPerUser int - MaxTotal int - IdleTimeout time.Duration - MaxLifetime time.Duration - StopTimeout time.Duration // SIGTERM grace period before SIGKILL; 0 = DefaultStopTimeout -} - -type PoolInfo struct { - Name string - Addr string - Connected bool - Boxes int - MaxPerUser int - MaxTotal int - IdleTimeout time.Duration - MaxLifetime time.Duration -} +// --------------------------------------------------------------------------- +// Create / List options +// --------------------------------------------------------------------------- type PortMapping struct { ContainerPort int @@ -51,7 +86,7 @@ type CreateOptions struct { ID string Owner string Labels map[string]string - Pool string + NodeID string Image string WorkDir string User string @@ -62,52 +97,66 @@ type CreateOptions struct { Ports []PortMapping Policy LifecyclePolicy IdleTimeout time.Duration + MaxLifetime time.Duration + StopTimeout time.Duration - StopTimeout time.Duration // SIGTERM grace period; 0 = pool default or DefaultStopTimeout - - WorkspaceID string // workspace to mount; empty = no workspace - MountMode string // "rw" (default) or "ro" - MountPath string // container path; default "/workspace" + WorkspaceID string + MountMode string + MountPath string } type ListOptions struct { Owner string - Pool string + NodeID string Labels map[string]string } +// --------------------------------------------------------------------------- +// Unified ExecOption / ExecResult / ExecStream +// --------------------------------------------------------------------------- + type execConfig struct { - WorkDir string - Env map[string]string - Timeout time.Duration + WorkDir string + Env map[string]string + Timeout time.Duration + Stdin []byte + MaxOutputBytes int64 } +// ExecOption configures an Exec or Stream call on any Computer. type ExecOption func(*execConfig) func WithWorkDir(dir string) ExecOption { - return func(c *execConfig) { - c.WorkDir = dir - } + return func(c *execConfig) { c.WorkDir = dir } } func WithEnv(env map[string]string) ExecOption { - return func(c *execConfig) { - c.Env = env - } + return func(c *execConfig) { c.Env = env } } func WithTimeout(timeout time.Duration) ExecOption { - return func(c *execConfig) { - c.Timeout = timeout - } + return func(c *execConfig) { c.Timeout = timeout } } +func WithStdin(data []byte) ExecOption { + return func(c *execConfig) { c.Stdin = data } +} + +func WithMaxOutput(bytes int64) ExecOption { + return func(c *execConfig) { c.MaxOutputBytes = bytes } +} + +// ExecResult holds the outcome of a command executed on any Computer. type ExecResult struct { - ExitCode int - Stdout string - Stderr string + ExitCode int + Stdout string + Stderr string + DurationMs int64 + Error string + Truncated bool } +// ExecStream provides real-time streaming I/O for a running command. type ExecStream struct { Stdout io.ReadCloser Stderr io.ReadCloser @@ -116,6 +165,10 @@ type ExecStream struct { Cancel func() } +// --------------------------------------------------------------------------- +// Attach (Box-specific, not part of Computer interface) +// --------------------------------------------------------------------------- + type attachConfig struct { Protocol string Path string @@ -125,26 +178,20 @@ type attachConfig struct { type AttachOption func(*attachConfig) func WithProtocol(protocol string) AttachOption { - return func(c *attachConfig) { - c.Protocol = protocol - } + return func(c *attachConfig) { c.Protocol = protocol } } func WithPath(path string) AttachOption { - return func(c *attachConfig) { - c.Path = path - } + return func(c *attachConfig) { c.Path = path } } func WithHeaders(headers map[string]string) AttachOption { - return func(c *attachConfig) { - c.Headers = headers - } + return func(c *attachConfig) { c.Headers = headers } } // ImagePullOptions configures an image pull operation. type ImagePullOptions struct { - Auth *RegistryAuth // nil = anonymous / public + Auth *RegistryAuth } // RegistryAuth holds credentials for a private container registry. @@ -162,10 +209,11 @@ type ServiceConn struct { Close func() error } +// BoxInfo is a snapshot of a Box's runtime state (used by Manager.List). type BoxInfo struct { ID string ContainerID string - Pool string + NodeID string Owner string Status string Policy LifecyclePolicy @@ -176,53 +224,3 @@ type BoxInfo struct { ProcessCount int VNC bool } - -// HostExecResult holds the outcome of a command executed on the Tai host. -type HostExecResult struct { - ExitCode int - Stdout []byte - Stderr []byte - DurationMs int64 - Error string - Truncated bool -} - -// HostExecStream provides real-time streaming output from a command running -// on the Tai host machine via HostExec gRPC ExecStream. -type HostExecStream struct { - Stdout <-chan []byte - Stderr <-chan []byte - Wait func() (int, error) // blocks until exit; returns exit code - Cancel func() // cancels the stream context -} - -type hostExecConfig struct { - WorkDir string - Env map[string]string - Stdin []byte - TimeoutMs int64 - MaxOutputBytes int64 -} - -// HostExecOption configures an ExecOnHost call. -type HostExecOption func(*hostExecConfig) - -func WithHostWorkDir(dir string) HostExecOption { - return func(c *hostExecConfig) { c.WorkDir = dir } -} - -func WithHostEnv(env map[string]string) HostExecOption { - return func(c *hostExecConfig) { c.Env = env } -} - -func WithHostStdin(data []byte) HostExecOption { - return func(c *hostExecConfig) { c.Stdin = data } -} - -func WithHostTimeout(ms int64) HostExecOption { - return func(c *hostExecConfig) { c.TimeoutMs = ms } -} - -func WithHostMaxOutput(bytes int64) HostExecOption { - return func(c *hostExecConfig) { c.MaxOutputBytes = bytes } -} diff --git a/sandbox/vncproxy/proxy.go b/sandbox/vncproxy/proxy.go index f968f1d9..fc77c53c 100644 --- a/sandbox/vncproxy/proxy.go +++ b/sandbox/vncproxy/proxy.go @@ -323,10 +323,8 @@ func (p *Proxy) checkVNCEnabled(ctx context.Context, containerName string) bool return false } - // Check environment variables for VNC_ENABLED or SANDBOX_VNC_ENABLED for _, env := range info.Config.Env { - if strings.HasPrefix(env, "SANDBOX_VNC_ENABLED=true") || - strings.HasPrefix(env, "VNC_ENABLED=true") { + if strings.HasPrefix(env, "VNC_ENABLED=true") { return true } } diff --git a/tai/api/register.go b/tai/api/register.go index 9c3640cb..c5779de8 100644 --- a/tai/api/register.go +++ b/tai/api/register.go @@ -8,7 +8,9 @@ import ( "github.com/gin-gonic/gin" "github.com/yaoapp/yao/openapi/oauth" + tai "github.com/yaoapp/yao/tai" "github.com/yaoapp/yao/tai/registry" + "github.com/yaoapp/yao/tai/taiid" ) // authenticateBearer validates a Bearer token and returns the caller's identity. @@ -33,6 +35,44 @@ func authenticateBearerDefault(token string) (registry.AuthInfo, error) { info.TeamID = result.Info.TeamID info.TenantID = result.Info.TenantID } + + slog.Info("[auth] buildAuthInfo result", + "subject", info.Subject, "user_id", info.UserID, + "client_id", info.ClientID, "team_id", info.TeamID, + "scope", info.Scope) + + if result.Claims != nil { + slog.Info("[auth] claims", + "claims.TeamID", result.Claims.TeamID, + "claims.TenantID", result.Claims.TenantID, + "claims.ClientID", result.Claims.ClientID, + "claims.Subject", result.Claims.Subject) + if result.Claims.Extra != nil { + slog.Info("[auth] claims.Extra", "extra", fmt.Sprintf("%+v", result.Claims.Extra)) + } else { + slog.Info("[auth] claims.Extra is nil") + } + + if info.TeamID == "" { + switch v := result.Claims.Extra["team_id"].(type) { + case string: + info.TeamID = v + slog.Info("[auth] team_id from Extra (string)", "team_id", v) + case float64: + info.TeamID = fmt.Sprintf("%.0f", v) + slog.Info("[auth] team_id from Extra (float64)", "team_id", info.TeamID) + default: + slog.Info("[auth] team_id not found in Extra or unknown type", + "type", fmt.Sprintf("%T", result.Claims.Extra["team_id"]), + "value", fmt.Sprintf("%v", result.Claims.Extra["team_id"])) + } + } + if info.TenantID == "" { + if v, ok := result.Claims.Extra["tenant_id"].(string); ok { + info.TenantID = v + } + } + } return info, nil } @@ -46,8 +86,10 @@ func extractBearer(r *http.Request) string { // registerRequest is the JSON body for POST /tai-nodes/register. type registerRequest struct { - TaiID string `json:"tai_id"` + NodeID string `json:"node_id,omitempty"` + ClientID string `json:"client_id,omitempty"` MachineID string `json:"machine_id"` + DisplayName string `json:"display_name,omitempty"` Version string `json:"version"` Addr string `json:"addr"` Ports map[string]int `json:"ports"` @@ -87,31 +129,63 @@ func HandleRegister(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": "invalid request body"}) return } - if req.TaiID == "" { - c.JSON(http.StatusBadRequest, gin.H{"error": "tai_id is required"}) + + if req.NodeID == "" || req.MachineID == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "node_id and machine_id are required"}) return } + resolvedTaiID, err := taiid.Generate(req.MachineID, req.NodeID) + if err != nil { + slog.Warn("taiid generation failed", "err", err) + c.JSON(http.StatusBadRequest, gin.H{"error": "failed to generate tai_id"}) + return + } + + remoteIP := c.ClientIP() + addr := req.Addr + if addr == "" && remoteIP != "" { + grpcPort := req.Ports["grpc"] + if grpcPort > 0 { + addr = fmt.Sprintf("tai://%s:%d", remoteIP, grpcPort) + } else { + addr = remoteIP + } + } + node := ®istry.TaiNode{ - TaiID: req.TaiID, + TaiID: resolvedTaiID, MachineID: req.MachineID, Version: req.Version, + DisplayName: req.DisplayName, Auth: authInfo, System: req.System, Mode: "direct", - Addr: req.Addr, + Addr: addr, Ports: req.Ports, Capabilities: req.Capabilities, } reg.Register(node) + slog.Info("[register] node registered via API", + "tai_id", resolvedTaiID, "addr", addr, "remote_ip", remoteIP, + "user_id", authInfo.UserID, "team_id", authInfo.TeamID) - remoteIP := c.ClientIP() - slog.Info("tai node registered via API", - "tai_id", req.TaiID, "remote_ip", remoteIP, "user_id", authInfo.UserID) + allBefore := reg.List() + slog.Info("[register] registry snapshot after Register", + "total", len(allBefore)) + for _, s := range allBefore { + slog.Info("[register] node", "tai_id", s.TaiID, "mode", s.Mode, "addr", s.Addr) + } + + if strings.HasPrefix(addr, "tai://") { + slog.Info("[register] launching connectRegisteredNode goroutine", + "tai_id", resolvedTaiID, "addr", addr) + go connectRegisteredNode(resolvedTaiID, addr, reg) + } c.JSON(http.StatusOK, gin.H{ "status": "registered", - "tai_id": req.TaiID, + "tai_id": resolvedTaiID, "remote_ip": remoteIP, }) } @@ -203,3 +277,46 @@ func HandleUnregister(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"status": "unregistered"}) } + +// connectRegisteredNode dials the self-registered Tai node via gRPC, +// creates a tai.Client, and binds it to the node's TaiID in the registry. +// initRemote internally registers a redundant "host-port" entry; we remove +// it so that the registry contains only the canonical taiID. +func connectRegisteredNode(taiID, addr string, reg *registry.Registry) { + slog.Info("[connect] start", "tai_id", taiID, "addr", addr) + + client, err := tai.New(addr) + if err != nil { + slog.Warn("[connect] tai.New FAILED", + "tai_id", taiID, "addr", addr, "err", err) + + allAfterFail := reg.List() + slog.Info("[connect] registry after tai.New failure", "total", len(allAfterFail)) + for _, s := range allAfterFail { + slog.Info("[connect] node", "tai_id", s.TaiID, "mode", s.Mode, "addr", s.Addr) + } + return + } + + autoID := client.TaiID() + slog.Info("[connect] tai.New OK", "tai_id", taiID, "autoID", autoID) + + allAfterNew := reg.List() + slog.Info("[connect] registry after tai.New", "total", len(allAfterNew)) + for _, s := range allAfterNew { + slog.Info("[connect] node", "tai_id", s.TaiID, "mode", s.Mode, "addr", s.Addr) + } + + if autoID != "" && autoID != taiID { + slog.Info("[connect] removing redundant autoID", "autoID", autoID) + reg.Unregister(autoID) + } + reg.SetClient(taiID, client) + + allFinal := reg.List() + slog.Info("[connect] registry FINAL", "total", len(allFinal)) + for _, s := range allFinal { + slog.Info("[connect] node", "tai_id", s.TaiID, "mode", s.Mode, "addr", s.Addr) + } + slog.Info("[connect] done", "tai_id", taiID) +} diff --git a/tai/api/register_test.go b/tai/api/register_test.go index d027469a..3957735d 100644 --- a/tai/api/register_test.go +++ b/tai/api/register_test.go @@ -46,9 +46,10 @@ func TestHandleRegister_Success(t *testing.T) { defer teardown() body := registerRequest{ - TaiID: "tai-abc123", + NodeID: "9100", MachineID: "m-001", Version: "0.2.0", + DisplayName: "My Dev Machine", Addr: "192.168.1.100", Ports: map[string]int{"grpc": 19100, "http": 8099}, Capabilities: map[string]bool{"docker": true, "host_exec": false}, @@ -74,14 +75,15 @@ func TestHandleRegister_Success(t *testing.T) { if resp["status"] != "registered" { t.Errorf("status = %v, want registered", resp["status"]) } - if resp["tai_id"] != "tai-abc123" { - t.Errorf("tai_id = %v, want tai-abc123", resp["tai_id"]) + taiID, _ := resp["tai_id"].(string) + if taiID == "" || len(taiID) < 5 || taiID[:4] != "tai-" { + t.Errorf("tai_id = %v, want server-generated tai-xxx", resp["tai_id"]) } if _, ok := resp["remote_ip"]; !ok { t.Error("response missing remote_ip") } - snap, ok := registry.Global().Get("tai-abc123") + snap, ok := registry.Global().Get(taiID) if !ok { t.Fatal("node not found in registry after register") } @@ -94,6 +96,74 @@ func TestHandleRegister_Success(t *testing.T) { if snap.Auth.UserID != "user-alice" { t.Errorf("Auth.UserID = %q, want user-alice", snap.Auth.UserID) } + if snap.DisplayName != "My Dev Machine" { + t.Errorf("DisplayName = %q, want %q", snap.DisplayName, "My Dev Machine") + } +} + +func TestHandleRegister_ServerGeneratedTaiID(t *testing.T) { + teardown := setupTest() + defer teardown() + + body := registerRequest{ + NodeID: "19100", + ClientID: "local-uuid-001", + MachineID: "m-001", + Version: "0.2.0", + DisplayName: "Generated ID Node", + Addr: "192.168.1.200", + Ports: map[string]int{"grpc": 19100}, + Capabilities: map[string]bool{"docker": true}, + System: registry.SystemInfo{OS: "darwin", Arch: "arm64", Hostname: "mac-01", NumCPU: 12}, + } + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest("POST", "/tai-nodes/register", jsonBody(body)) + c.Request.Header.Set("Authorization", "Bearer test-token") + c.Request.Header.Set("Content-Type", "application/json") + + HandleRegister(c) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d, want %d; body = %s", w.Code, http.StatusOK, w.Body.String()) + } + + var resp map[string]interface{} + json.Unmarshal(w.Body.Bytes(), &resp) + + generatedID, ok := resp["tai_id"].(string) + if !ok || generatedID == "" { + t.Fatal("response missing tai_id") + } + if generatedID == "19100" { + t.Error("tai_id should be server-generated, not the raw node_id") + } + if len(generatedID) != 26 { + t.Errorf("tai_id length = %d, want 26 (tai- + 22 base62); got %q", len(generatedID), generatedID) + } + + snap, ok2 := registry.Global().Get(generatedID) + if !ok2 { + t.Fatalf("node %q not found in registry", generatedID) + } + if snap.DisplayName != "Generated ID Node" { + t.Errorf("DisplayName = %q, want %q", snap.DisplayName, "Generated ID Node") + } + + // Deterministic: same inputs produce same ID + w2 := httptest.NewRecorder() + c2, _ := gin.CreateTestContext(w2) + c2.Request = httptest.NewRequest("POST", "/tai-nodes/register", jsonBody(body)) + c2.Request.Header.Set("Authorization", "Bearer test-token") + c2.Request.Header.Set("Content-Type", "application/json") + HandleRegister(c2) + + var resp2 map[string]interface{} + json.Unmarshal(w2.Body.Bytes(), &resp2) + if resp2["tai_id"] != generatedID { + t.Errorf("not deterministic: %v != %v", resp2["tai_id"], generatedID) + } } func TestHandleRegister_MissingAuth(t *testing.T) { @@ -102,7 +172,7 @@ func TestHandleRegister_MissingAuth(t *testing.T) { w := httptest.NewRecorder() c, _ := gin.CreateTestContext(w) - c.Request = httptest.NewRequest("POST", "/tai-nodes/register", jsonBody(registerRequest{TaiID: "x"})) + c.Request = httptest.NewRequest("POST", "/tai-nodes/register", jsonBody(registerRequest{NodeID: "x", MachineID: "m1"})) c.Request.Header.Set("Content-Type", "application/json") HandleRegister(c) @@ -112,7 +182,7 @@ func TestHandleRegister_MissingAuth(t *testing.T) { } } -func TestHandleRegister_MissingTaiID(t *testing.T) { +func TestHandleRegister_MissingTaiIDAndClientID(t *testing.T) { teardown := setupTest() defer teardown() diff --git a/tai/docs/registry.md b/tai/docs/registry.md index a8f46834..3d48d248 100644 --- a/tai/docs/registry.md +++ b/tai/docs/registry.md @@ -30,7 +30,7 @@ type TaiNode struct { Status string // "online" | "offline" | "connecting" ConnectedAt time.Time LastPing time.Time - PoolName string + DisplayName string } ``` @@ -48,7 +48,7 @@ type NodeSnapshot struct { Capabilities map[string]bool Status string ConnectedAt, LastPing time.Time - PoolName string + DisplayName string } ``` diff --git a/tai/registry/registry.go b/tai/registry/registry.go index 7386b720..a4683299 100644 --- a/tai/registry/registry.go +++ b/tai/registry/registry.go @@ -7,6 +7,8 @@ import ( "io" "log/slog" "net" + "os" + "strings" "sync" "time" @@ -20,6 +22,8 @@ type SystemInfo struct { Hostname string `json:"hostname"` NumCPU int `json:"num_cpu"` TotalMem int64 `json:"total_mem,omitempty"` + Shell string `json:"shell,omitempty"` + TempDir string `json:"temp_dir,omitempty"` } // TaiNode represents a registered Tai instance (direct or tunnel). @@ -42,7 +46,9 @@ type TaiNode struct { Status string // "online" | "offline" | "connecting" ConnectedAt time.Time LastPing time.Time - PoolName string + DisplayName string // optional human-readable name for UI + + client any // *tai.Client; stored as any to avoid import cycle localListeners map[int]*tunnelListener } @@ -62,7 +68,8 @@ type NodeSnapshot struct { Status string ConnectedAt time.Time LastPing time.Time - PoolName string + DisplayName string + client any } func (n *TaiNode) snapshot() NodeSnapshot { @@ -80,10 +87,15 @@ func (n *TaiNode) snapshot() NodeSnapshot { Mode: n.Mode, Addr: n.Addr, YaoBase: n.YaoBase, Ports: ports, Capabilities: caps, Status: n.Status, ConnectedAt: n.ConnectedAt, LastPing: n.LastPing, - PoolName: n.PoolName, + DisplayName: n.DisplayName, + client: n.client, } } +// Client returns the associated *tai.Client (as any to avoid import cycle). +// Callers should type-assert: snap.Client().(*tai.Client). +func (s *NodeSnapshot) Client() any { return s.client } + // AuthInfo holds Yao user authorization extracted from OAuth token. type AuthInfo struct { Subject string @@ -137,6 +149,23 @@ func Init(logger *slog.Logger) { }) } +// InitWithWriter initializes the global registry using the given io.Writer +// and log format ("JSON" or "TEXT"). If w is nil it falls back to stderr. +// This is the preferred way to integrate with the application log system. +func InitWithWriter(w io.Writer, logMode string) { + if w == nil { + w = os.Stderr + } + opts := &slog.HandlerOptions{Level: slog.LevelInfo} + var handler slog.Handler + if strings.EqualFold(logMode, "JSON") { + handler = slog.NewJSONHandler(w, opts) + } else { + handler = slog.NewTextHandler(w, opts) + } + Init(slog.New(handler)) +} + // Global returns the global registry instance. func Global() *Registry { return global @@ -234,6 +263,31 @@ func (r *Registry) UpdatePing(taiID string) { } } +// SetClient associates a *tai.Client with a registered node. +// Called by tai.New() after successful initialization. +func (r *Registry) SetClient(taiID string, c any) { + r.mu.Lock() + defer r.mu.Unlock() + if n, ok := r.nodes[taiID]; ok { + n.client = c + } +} + +// FindTaiIDByAuthClient returns the TaiID of the first node whose +// Auth.ClientID matches the given OAuth client ID. Returns "" if not found. +// This is needed because Tai's data channel authenticates with its OAuth +// ClientID, which may differ from the server-assigned TaiID. +func (r *Registry) FindTaiIDByAuthClient(clientID string) string { + r.mu.RLock() + defer r.mu.RUnlock() + for _, n := range r.nodes { + if n.Auth.ClientID == clientID { + return n.TaiID + } + } + return "" +} + // ListByTeam returns snapshots of all nodes belonging to the given team. func (r *Registry) ListByTeam(teamID string) []NodeSnapshot { r.mu.RLock() @@ -247,6 +301,20 @@ func (r *Registry) ListByTeam(teamID string) []NodeSnapshot { return result } +// ListByUser returns snapshots of all nodes registered by the given user +// that are NOT associated with any team. +func (r *Registry) ListByUser(userID string) []NodeSnapshot { + r.mu.RLock() + defer r.mu.RUnlock() + var result []NodeSnapshot + for _, n := range r.nodes { + if n.Auth.TeamID == "" && n.Auth.UserID == userID { + result = append(result, n.snapshot()) + } + } + return result +} + // StartHealthCheck runs a background goroutine that periodically checks // direct-mode nodes for heartbeat timeout. Nodes whose LastPing exceeds // timeout are marked offline. Nodes that remain offline longer than diff --git a/tai/sandbox/docker_core.go b/tai/sandbox/docker_core.go index a0e77581..16f950cc 100644 --- a/tai/sandbox/docker_core.go +++ b/tai/sandbox/docker_core.go @@ -59,7 +59,7 @@ func (d *dockerCore) create(ctx context.Context, opts CreateOptions, addVNCPorts shmSize = 256 * 1024 * 1024 } hostCfg.ShmSize = shmSize - cfg.Env = append(cfg.Env, "SANDBOX_VNC_ENABLED=true") + cfg.Env = append(cfg.Env, "VNC_ENABLED=true") if addVNCPorts { for _, p := range []int{6080, 5900} { diff --git a/tai/serverinfo/pb/serverinfo.pb.go b/tai/serverinfo/pb/serverinfo.pb.go index 5d85f88a..ac804da4 100644 --- a/tai/serverinfo/pb/serverinfo.pb.go +++ b/tai/serverinfo/pb/serverinfo.pb.go @@ -2,7 +2,7 @@ // versions: // protoc-gen-go v1.36.11 // protoc v4.25.0 -// source: tai/serverinfo/pb/serverinfo.proto +// source: serverinfo.proto package pb @@ -29,7 +29,7 @@ type GetInfoRequest struct { func (x *GetInfoRequest) Reset() { *x = GetInfoRequest{} - mi := &file_tai_serverinfo_pb_serverinfo_proto_msgTypes[0] + mi := &file_serverinfo_proto_msgTypes[0] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -41,7 +41,7 @@ func (x *GetInfoRequest) String() string { func (*GetInfoRequest) ProtoMessage() {} func (x *GetInfoRequest) ProtoReflect() protoreflect.Message { - mi := &file_tai_serverinfo_pb_serverinfo_proto_msgTypes[0] + mi := &file_serverinfo_proto_msgTypes[0] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -54,7 +54,99 @@ func (x *GetInfoRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use GetInfoRequest.ProtoReflect.Descriptor instead. func (*GetInfoRequest) Descriptor() ([]byte, []int) { - return file_tai_serverinfo_pb_serverinfo_proto_rawDescGZIP(), []int{0} + return file_serverinfo_proto_rawDescGZIP(), []int{0} +} + +type SystemInfo struct { + state protoimpl.MessageState `protogen:"open.v1"` + Os string `protobuf:"bytes,1,opt,name=os,proto3" json:"os,omitempty"` + Arch string `protobuf:"bytes,2,opt,name=arch,proto3" json:"arch,omitempty"` + Hostname string `protobuf:"bytes,3,opt,name=hostname,proto3" json:"hostname,omitempty"` + NumCpu int32 `protobuf:"varint,4,opt,name=num_cpu,json=numCpu,proto3" json:"num_cpu,omitempty"` + TotalMem int64 `protobuf:"varint,5,opt,name=total_mem,json=totalMem,proto3" json:"total_mem,omitempty"` + Shell string `protobuf:"bytes,6,opt,name=shell,proto3" json:"shell,omitempty"` // preferred shell: "sh", "pwsh", "powershell", "cmd.exe" + TempDir string `protobuf:"bytes,7,opt,name=temp_dir,json=tempDir,proto3" json:"temp_dir,omitempty"` // system temp directory + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *SystemInfo) Reset() { + *x = SystemInfo{} + mi := &file_serverinfo_proto_msgTypes[1] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *SystemInfo) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*SystemInfo) ProtoMessage() {} + +func (x *SystemInfo) ProtoReflect() protoreflect.Message { + mi := &file_serverinfo_proto_msgTypes[1] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use SystemInfo.ProtoReflect.Descriptor instead. +func (*SystemInfo) Descriptor() ([]byte, []int) { + return file_serverinfo_proto_rawDescGZIP(), []int{1} +} + +func (x *SystemInfo) GetOs() string { + if x != nil { + return x.Os + } + return "" +} + +func (x *SystemInfo) GetArch() string { + if x != nil { + return x.Arch + } + return "" +} + +func (x *SystemInfo) GetHostname() string { + if x != nil { + return x.Hostname + } + return "" +} + +func (x *SystemInfo) GetNumCpu() int32 { + if x != nil { + return x.NumCpu + } + return 0 +} + +func (x *SystemInfo) GetTotalMem() int64 { + if x != nil { + return x.TotalMem + } + return 0 +} + +func (x *SystemInfo) GetShell() string { + if x != nil { + return x.Shell + } + return "" +} + +func (x *SystemInfo) GetTempDir() string { + if x != nil { + return x.TempDir + } + return "" } type GetInfoResponse struct { @@ -62,13 +154,14 @@ type GetInfoResponse struct { Version string `protobuf:"bytes,1,opt,name=version,proto3" json:"version,omitempty"` Ports map[string]int32 `protobuf:"bytes,2,rep,name=ports,proto3" json:"ports,omitempty" protobuf_key:"bytes,1,opt,name=key" protobuf_val:"varint,2,opt,name=value"` // "grpc", "http", "vnc", "docker", "k8s" Capabilities map[string]bool `protobuf:"bytes,3,rep,name=capabilities,proto3" json:"capabilities,omitempty" protobuf_key:"bytes,1,opt,name=key" protobuf_val:"varint,2,opt,name=value"` // "docker", "k8s" + System *SystemInfo `protobuf:"bytes,4,opt,name=system,proto3" json:"system,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } func (x *GetInfoResponse) Reset() { *x = GetInfoResponse{} - mi := &file_tai_serverinfo_pb_serverinfo_proto_msgTypes[1] + mi := &file_serverinfo_proto_msgTypes[2] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -80,7 +173,7 @@ func (x *GetInfoResponse) String() string { func (*GetInfoResponse) ProtoMessage() {} func (x *GetInfoResponse) ProtoReflect() protoreflect.Message { - mi := &file_tai_serverinfo_pb_serverinfo_proto_msgTypes[1] + mi := &file_serverinfo_proto_msgTypes[2] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -93,7 +186,7 @@ func (x *GetInfoResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use GetInfoResponse.ProtoReflect.Descriptor instead. func (*GetInfoResponse) Descriptor() ([]byte, []int) { - return file_tai_serverinfo_pb_serverinfo_proto_rawDescGZIP(), []int{1} + return file_serverinfo_proto_rawDescGZIP(), []int{2} } func (x *GetInfoResponse) GetVersion() string { @@ -117,17 +210,34 @@ func (x *GetInfoResponse) GetCapabilities() map[string]bool { return nil } -var File_tai_serverinfo_pb_serverinfo_proto protoreflect.FileDescriptor +func (x *GetInfoResponse) GetSystem() *SystemInfo { + if x != nil { + return x.System + } + return nil +} -const file_tai_serverinfo_pb_serverinfo_proto_rawDesc = "" + +var File_serverinfo_proto protoreflect.FileDescriptor + +const file_serverinfo_proto_rawDesc = "" + "\n" + - "\"tai/serverinfo/pb/serverinfo.proto\x12\n" + + "\x10serverinfo.proto\x12\n" + "serverinfo\"\x10\n" + - "\x0eGetInfoRequest\"\xb7\x02\n" + + "\x0eGetInfoRequest\"\xb3\x01\n" + + "\n" + + "SystemInfo\x12\x0e\n" + + "\x02os\x18\x01 \x01(\tR\x02os\x12\x12\n" + + "\x04arch\x18\x02 \x01(\tR\x04arch\x12\x1a\n" + + "\bhostname\x18\x03 \x01(\tR\bhostname\x12\x17\n" + + "\anum_cpu\x18\x04 \x01(\x05R\x06numCpu\x12\x1b\n" + + "\ttotal_mem\x18\x05 \x01(\x03R\btotalMem\x12\x14\n" + + "\x05shell\x18\x06 \x01(\tR\x05shell\x12\x19\n" + + "\btemp_dir\x18\a \x01(\tR\atempDir\"\xe7\x02\n" + "\x0fGetInfoResponse\x12\x18\n" + "\aversion\x18\x01 \x01(\tR\aversion\x12<\n" + "\x05ports\x18\x02 \x03(\v2&.serverinfo.GetInfoResponse.PortsEntryR\x05ports\x12Q\n" + - "\fcapabilities\x18\x03 \x03(\v2-.serverinfo.GetInfoResponse.CapabilitiesEntryR\fcapabilities\x1a8\n" + + "\fcapabilities\x18\x03 \x03(\v2-.serverinfo.GetInfoResponse.CapabilitiesEntryR\fcapabilities\x12.\n" + + "\x06system\x18\x04 \x01(\v2\x16.serverinfo.SystemInfoR\x06system\x1a8\n" + "\n" + "PortsEntry\x12\x10\n" + "\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" + @@ -140,56 +250,58 @@ const file_tai_serverinfo_pb_serverinfo_proto_rawDesc = "" + "\aGetInfo\x12\x1a.serverinfo.GetInfoRequest\x1a\x1b.serverinfo.GetInfoResponseB%Z#github.com/yaoapp/tai/serverinfo/pbb\x06proto3" var ( - file_tai_serverinfo_pb_serverinfo_proto_rawDescOnce sync.Once - file_tai_serverinfo_pb_serverinfo_proto_rawDescData []byte + file_serverinfo_proto_rawDescOnce sync.Once + file_serverinfo_proto_rawDescData []byte ) -func file_tai_serverinfo_pb_serverinfo_proto_rawDescGZIP() []byte { - file_tai_serverinfo_pb_serverinfo_proto_rawDescOnce.Do(func() { - file_tai_serverinfo_pb_serverinfo_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_tai_serverinfo_pb_serverinfo_proto_rawDesc), len(file_tai_serverinfo_pb_serverinfo_proto_rawDesc))) +func file_serverinfo_proto_rawDescGZIP() []byte { + file_serverinfo_proto_rawDescOnce.Do(func() { + file_serverinfo_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_serverinfo_proto_rawDesc), len(file_serverinfo_proto_rawDesc))) }) - return file_tai_serverinfo_pb_serverinfo_proto_rawDescData + return file_serverinfo_proto_rawDescData } -var file_tai_serverinfo_pb_serverinfo_proto_msgTypes = make([]protoimpl.MessageInfo, 4) -var file_tai_serverinfo_pb_serverinfo_proto_goTypes = []any{ +var file_serverinfo_proto_msgTypes = make([]protoimpl.MessageInfo, 5) +var file_serverinfo_proto_goTypes = []any{ (*GetInfoRequest)(nil), // 0: serverinfo.GetInfoRequest - (*GetInfoResponse)(nil), // 1: serverinfo.GetInfoResponse - nil, // 2: serverinfo.GetInfoResponse.PortsEntry - nil, // 3: serverinfo.GetInfoResponse.CapabilitiesEntry + (*SystemInfo)(nil), // 1: serverinfo.SystemInfo + (*GetInfoResponse)(nil), // 2: serverinfo.GetInfoResponse + nil, // 3: serverinfo.GetInfoResponse.PortsEntry + nil, // 4: serverinfo.GetInfoResponse.CapabilitiesEntry } -var file_tai_serverinfo_pb_serverinfo_proto_depIdxs = []int32{ - 2, // 0: serverinfo.GetInfoResponse.ports:type_name -> serverinfo.GetInfoResponse.PortsEntry - 3, // 1: serverinfo.GetInfoResponse.capabilities:type_name -> serverinfo.GetInfoResponse.CapabilitiesEntry - 0, // 2: serverinfo.ServerInfo.GetInfo:input_type -> serverinfo.GetInfoRequest - 1, // 3: serverinfo.ServerInfo.GetInfo:output_type -> serverinfo.GetInfoResponse - 3, // [3:4] is the sub-list for method output_type - 2, // [2:3] is the sub-list for method input_type - 2, // [2:2] is the sub-list for extension type_name - 2, // [2:2] is the sub-list for extension extendee - 0, // [0:2] is the sub-list for field type_name +var file_serverinfo_proto_depIdxs = []int32{ + 3, // 0: serverinfo.GetInfoResponse.ports:type_name -> serverinfo.GetInfoResponse.PortsEntry + 4, // 1: serverinfo.GetInfoResponse.capabilities:type_name -> serverinfo.GetInfoResponse.CapabilitiesEntry + 1, // 2: serverinfo.GetInfoResponse.system:type_name -> serverinfo.SystemInfo + 0, // 3: serverinfo.ServerInfo.GetInfo:input_type -> serverinfo.GetInfoRequest + 2, // 4: serverinfo.ServerInfo.GetInfo:output_type -> serverinfo.GetInfoResponse + 4, // [4:5] is the sub-list for method output_type + 3, // [3:4] is the sub-list for method input_type + 3, // [3:3] is the sub-list for extension type_name + 3, // [3:3] is the sub-list for extension extendee + 0, // [0:3] is the sub-list for field type_name } -func init() { file_tai_serverinfo_pb_serverinfo_proto_init() } -func file_tai_serverinfo_pb_serverinfo_proto_init() { - if File_tai_serverinfo_pb_serverinfo_proto != nil { +func init() { file_serverinfo_proto_init() } +func file_serverinfo_proto_init() { + if File_serverinfo_proto != nil { return } type x struct{} out := protoimpl.TypeBuilder{ File: protoimpl.DescBuilder{ GoPackagePath: reflect.TypeOf(x{}).PkgPath(), - RawDescriptor: unsafe.Slice(unsafe.StringData(file_tai_serverinfo_pb_serverinfo_proto_rawDesc), len(file_tai_serverinfo_pb_serverinfo_proto_rawDesc)), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_serverinfo_proto_rawDesc), len(file_serverinfo_proto_rawDesc)), NumEnums: 0, - NumMessages: 4, + NumMessages: 5, NumExtensions: 0, NumServices: 1, }, - GoTypes: file_tai_serverinfo_pb_serverinfo_proto_goTypes, - DependencyIndexes: file_tai_serverinfo_pb_serverinfo_proto_depIdxs, - MessageInfos: file_tai_serverinfo_pb_serverinfo_proto_msgTypes, + GoTypes: file_serverinfo_proto_goTypes, + DependencyIndexes: file_serverinfo_proto_depIdxs, + MessageInfos: file_serverinfo_proto_msgTypes, }.Build() - File_tai_serverinfo_pb_serverinfo_proto = out.File - file_tai_serverinfo_pb_serverinfo_proto_goTypes = nil - file_tai_serverinfo_pb_serverinfo_proto_depIdxs = nil + File_serverinfo_proto = out.File + file_serverinfo_proto_goTypes = nil + file_serverinfo_proto_depIdxs = nil } diff --git a/tai/serverinfo/pb/serverinfo.proto b/tai/serverinfo/pb/serverinfo.proto index 4eece2ee..844f1881 100644 --- a/tai/serverinfo/pb/serverinfo.proto +++ b/tai/serverinfo/pb/serverinfo.proto @@ -8,8 +8,19 @@ service ServerInfo { message GetInfoRequest {} +message SystemInfo { + string os = 1; + string arch = 2; + string hostname = 3; + int32 num_cpu = 4; + int64 total_mem = 5; + string shell = 6; // preferred shell: "sh", "pwsh", "powershell", "cmd.exe" + string temp_dir = 7; // system temp directory +} + message GetInfoResponse { string version = 1; - map ports = 2; // "grpc", "http", "vnc", "docker", "k8s" + map ports = 2; // "grpc", "http", "vnc", "docker", "k8s" map capabilities = 3; // "docker", "k8s" + SystemInfo system = 4; } diff --git a/tai/serverinfo/pb/serverinfo_grpc.pb.go b/tai/serverinfo/pb/serverinfo_grpc.pb.go index 07e690ec..173e1fa0 100644 --- a/tai/serverinfo/pb/serverinfo_grpc.pb.go +++ b/tai/serverinfo/pb/serverinfo_grpc.pb.go @@ -2,7 +2,7 @@ // versions: // - protoc-gen-go-grpc v1.6.1 // - protoc v4.25.0 -// source: tai/serverinfo/pb/serverinfo.proto +// source: serverinfo.proto package pb @@ -117,5 +117,5 @@ var ServerInfo_ServiceDesc = grpc.ServiceDesc{ }, }, Streams: []grpc.StreamDesc{}, - Metadata: "tai/serverinfo/pb/serverinfo.proto", + Metadata: "serverinfo.proto", } diff --git a/tai/tai.go b/tai/tai.go index 39d2a13d..e1170b85 100644 --- a/tai/tai.go +++ b/tai/tai.go @@ -130,6 +130,7 @@ type Client struct { scheme string // "tai", "docker", or "tunnel" host string addr string + taiID string // registry key — set by initLocal/initRemote/initTunnel ports Ports dataDir string // host-side data directory for local volume vol volume.Volume @@ -209,6 +210,23 @@ func (c *Client) initLocal(cfg *config) (*Client, error) { c.dataDir = dataDir c.vol = volume.NewLocal(dataDir) } + + if reg := registry.Global(); reg != nil { + id := c.host + if id == "" { + id = c.addr + } + if id == "" { + id = "local" + } + c.taiID = id + reg.Register(®istry.TaiNode{ + TaiID: id, + Mode: "local", + Addr: c.addr, + }) + reg.SetClient(id, c) + } return c, nil } @@ -221,15 +239,14 @@ func (c *Client) initRemote(cfg *config) (*Client, error) { c.grpcConn = conn c.he = hepb.NewHostExecClient(conn) - caps, err := c.discoverServerInfo(conn, cfg) + info, err := c.discoverServerInfo(conn, cfg) if err != nil { - // Old Tai without ServerInfo — fall back to legacy behaviour (try Docker). - caps = map[string]bool{"docker": true} + info = &discoveredInfo{Capabilities: map[string]bool{"docker": true}} } - hasDocker := caps["docker"] - hasK8s := caps["k8s"] - hasHostExec := caps["host_exec"] + hasDocker := info.Capabilities["docker"] + hasK8s := info.Capabilities["k8s"] + hasHostExec := info.Capabilities["host_exec"] if !hasDocker && !hasK8s && !hasHostExec { conn.Close() @@ -276,10 +293,15 @@ func (c *Client) initRemote(cfg *config) (*Client, error) { } if reg := registry.Global(); reg != nil { + id := fmt.Sprintf("%s-%d", c.host, c.ports.GRPC) + c.taiID = id reg.Register(®istry.TaiNode{ - TaiID: c.host, - Mode: "direct", - Addr: c.host, + TaiID: id, + Mode: "direct", + Version: info.Version, + System: info.System, + Capabilities: info.Capabilities, + Addr: fmt.Sprintf("tai://%s:%d", c.host, c.ports.GRPC), Ports: map[string]int{ "grpc": c.ports.GRPC, "http": c.ports.HTTP, @@ -288,6 +310,7 @@ func (c *Client) initRemote(cfg *config) (*Client, error) { "k8s": c.ports.K8s, }, }) + reg.SetClient(id, c) } return c, nil @@ -300,6 +323,7 @@ func (c *Client) initTunnel(cfg *config) (*Client, error) { } taiID := c.host // for tunnel:// scheme, host stores the taiID + c.taiID = taiID node, ok := reg.Get(taiID) if !ok || node.Status != "online" { return nil, fmt.Errorf("tai node %s not online", taiID) @@ -310,6 +334,7 @@ func (c *Client) initTunnel(cfg *config) (*Client, error) { HTTP: nodePort(node.Ports, "http", 8099), VNC: nodePort(node.Ports, "vnc", 16080), Docker: nodePort(node.Ports, "docker", 12375), + K8s: nodePort(node.Ports, "k8s", 16443), } grpcLn, err := reg.OpenLocalListener(taiID, c.ports.GRPC) @@ -329,21 +354,36 @@ func (c *Client) initTunnel(cfg *config) (*Client, error) { c.he = hepb.NewHostExecClient(conn) c.vol = volume.NewRemote(conn) - caps, err := c.discoverServerInfo(conn, cfg) + info, err := c.discoverServerInfo(conn, cfg) if err != nil { - caps = map[string]bool{"docker": true} + info = &discoveredInfo{Capabilities: map[string]bool{"docker": true}} } - hasDocker := caps["docker"] - hasHostExec := caps["host_exec"] + hasDocker := info.Capabilities["docker"] + hasK8s := info.Capabilities["k8s"] + hasHostExec := info.Capabilities["host_exec"] - if !hasDocker && !hasHostExec { + if !hasDocker && !hasK8s && !hasHostExec { c.closeTunnelListeners() conn.Close() - return nil, fmt.Errorf("tai %s: no capabilities available via tunnel", taiID) + return nil, fmt.Errorf("tai %s: no capabilities available via tunnel (docker/k8s/host_exec all false)", taiID) } - if hasDocker && c.ports.Docker > 0 { + if cfg.runtime == K8s || (!hasDocker && hasK8s) { + k8sLn, err := reg.OpenLocalListener(taiID, c.ports.K8s) + if err == nil { + c.tunnelListeners = append(c.tunnelListeners, k8sLn) + sbAddr := k8sLn.Addr().String() + sb, err := sandbox.NewK8s(sbAddr, sandbox.K8sOption{ + Namespace: cfg.namespace, + KubeConfig: cfg.kubeConfig, + }) + if err == nil { + c.sb = sb + c.img = sandbox.NewK8sImage() + } + } + } else if hasDocker && c.ports.Docker > 0 { dockerLn, err := reg.OpenLocalListener(taiID, c.ports.Docker) if err == nil { c.tunnelListeners = append(c.tunnelListeners, dockerLn) @@ -360,6 +400,7 @@ func (c *Client) initTunnel(cfg *config) (*Client, error) { c.prx = proxy.NewTunnel(taiID, node.YaoBase) c.vc = vnc.NewTunnel(taiID, node.YaoBase) } + reg.SetClient(taiID, c) return c, nil } @@ -396,9 +437,9 @@ func (c *Client) Close() error { } } c.closeTunnelListeners() - if c.scheme == "tai" { + if c.taiID != "" { if reg := registry.Global(); reg != nil { - reg.Unregister(c.host) + reg.Unregister(c.taiID) } } if len(errs) > 0 { @@ -414,6 +455,12 @@ func (c *Client) Volume() volume.Volume { return c.vol } // Empty for remote (Tai gRPC) connections — the Tai server manages paths. func (c *Client) DataDir() string { return c.dataDir } +// Host returns the raw host parsed from the address (IP or hostname). +func (c *Client) Host() string { return c.host } + +// TaiID returns the registry key for this client. +func (c *Client) TaiID() string { return c.taiID } + // Workspace returns an fs.FS-compatible filesystem for the given session. func (c *Client) Workspace(sessionID string) workspace.FS { return workspace.New(c.vol, sessionID) @@ -515,10 +562,16 @@ func isLocalHost(h string) bool { return h == "127.0.0.1" || h == "localhost" || h == "::1" } +type discoveredInfo struct { + Capabilities map[string]bool + System registry.SystemInfo + Version string +} + // discoverServerInfo calls ServerInfo.GetInfo on the remote Tai server, merges -// discovered ports into c.ports, and returns the server's capabilities map. +// discovered ports into c.ports, and returns capabilities + system info. // Ports explicitly set via WithPorts take precedence over server-reported values. -func (c *Client) discoverServerInfo(conn *grpc.ClientConn, cfg *config) (map[string]bool, error) { +func (c *Client) discoverServerInfo(conn *grpc.ClientConn, cfg *config) (*discoveredInfo, error) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() @@ -547,5 +600,71 @@ func (c *Client) discoverServerInfo(conn *grpc.ClientConn, cfg *config) (map[str if caps == nil { caps = make(map[string]bool) } - return caps, nil + + var sys registry.SystemInfo + if s := resp.System; s != nil { + sys = registry.SystemInfo{ + OS: s.Os, + Arch: s.Arch, + Hostname: s.Hostname, + NumCPU: int(s.NumCpu), + TotalMem: s.TotalMem, + Shell: s.Shell, + TempDir: s.TempDir, + } + } + + return &discoveredInfo{ + Capabilities: caps, + System: sys, + Version: resp.Version, + }, nil +} + +// RegisterLocal probes the local Docker environment and, if reachable, +// creates a Client and registers it as the "local" node in the registry. +// Returns true if a local node was successfully registered. +// Silently returns false if Docker is not available — this is not an error. +func RegisterLocal(opts ...Option) bool { + reg := registry.Global() + if reg == nil { + return false + } + if _, ok := reg.Get("local"); ok { + return true + } + + c, err := New("local", opts...) + if err != nil { + return false + } + _ = c // registered by initLocal → reg.Register + reg.SetClient + return true +} + +// GetClient returns a registered *Client by taiID from the global registry. +func GetClient(taiID string) (*Client, bool) { + reg := registry.Global() + if reg == nil { + return nil, false + } + snap, ok := reg.Get(taiID) + if !ok { + return nil, false + } + c, ok := snap.Client().(*Client) + if !ok || c == nil { + return nil, false + } + return c, true +} + +// GetNodeSnapshot returns the registry snapshot for a Tai node by ID. +// Callers can inspect System, Capabilities, Mode and other registry-level fields. +func GetNodeSnapshot(taiID string) (*registry.NodeSnapshot, bool) { + reg := registry.Global() + if reg == nil { + return nil, false + } + return reg.Get(taiID) } diff --git a/tai/tai_test.go b/tai/tai_test.go index 3695492c..4d2148e6 100644 --- a/tai/tai_test.go +++ b/tai/tai_test.go @@ -5,6 +5,8 @@ import ( "os" "strconv" "testing" + + "github.com/yaoapp/yao/tai/registry" ) func taiTestHost() string { @@ -342,3 +344,77 @@ func TestDiscoverPortsWithUserOverride(t *testing.T) { t.Logf("ports: GRPC=%d HTTP=%d(user) VNC=%d Docker=%d", c.ports.GRPC, c.ports.HTTP, c.ports.VNC, c.ports.Docker) } + +func TestRegisterLocal(t *testing.T) { + registry.Init(nil) + reg := registry.Global() + + dir := t.TempDir() + ok := RegisterLocal(WithDataDir(dir)) + if !ok { + t.Skip("Docker not available, skipping RegisterLocal test") + } + + snap, found := reg.Get("local") + if !found { + t.Fatal("expected 'local' node in registry after RegisterLocal") + } + if snap.Mode != "local" { + t.Errorf("mode = %q, want 'local'", snap.Mode) + } + if snap.Status != "online" { + t.Errorf("status = %q, want 'online'", snap.Status) + } + + c, got := GetClient("local") + if !got { + t.Fatal("GetClient('local') returned false after RegisterLocal") + } + if c.DataDir() != dir { + t.Errorf("DataDir = %q, want %q", c.DataDir(), dir) + } + if c.Sandbox() == nil { + t.Error("local client Sandbox should not be nil") + } + + // Idempotent: second call should return true without error + ok2 := RegisterLocal(WithDataDir(dir)) + if !ok2 { + t.Error("second RegisterLocal should return true (idempotent)") + } + + c.Close() +} + +func TestRegisterLocal_NoRegistry(t *testing.T) { + // RegisterLocal without a registry should return false, not panic + origReg := registry.Global() + defer func() { + if origReg != nil { + registry.Init(nil) + } + }() + + // registry.Global() returns the singleton; we can't un-init it, + // but we can verify RegisterLocal returns true (registry exists from + // other tests) or false gracefully. + ok := RegisterLocal() + // Just verify it doesn't panic; result depends on Docker availability + _ = ok +} + +func TestRegisterLocal_NoDocker(t *testing.T) { + registry.Init(nil) + + // Use an unreachable Docker socket to ensure failure + ok := RegisterLocal(WithDataDir(t.TempDir())) + if !ok { + // Expected when Docker is not available — just ensure no panic + return + } + // If Docker happens to be available, that's also fine + c, _ := GetClient("local") + if c != nil { + c.Close() + } +} diff --git a/tai/taiid/taiid.go b/tai/taiid/taiid.go new file mode 100644 index 00000000..c23f6180 --- /dev/null +++ b/tai/taiid/taiid.go @@ -0,0 +1,37 @@ +package taiid + +import ( + "crypto/sha256" + "fmt" + "math/big" +) + +const base62Chars = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + +// Generate produces a deterministic tai_id from a machine ID and a node ID. +// The result is "tai-" followed by a Base62-encoded truncated SHA-256 hash. +// Both machineID and nodeID must be non-empty. +func Generate(machineID, nodeID string) (string, error) { + if machineID == "" || nodeID == "" { + return "", fmt.Errorf("machineID and nodeID are required") + } + h := sha256.Sum256([]byte(machineID + ":" + nodeID)) + return "tai-" + base62Encode(h[:16]), nil +} + +func base62Encode(data []byte) string { + num := new(big.Int).SetBytes(data) + base := big.NewInt(62) + zero := big.NewInt(0) + mod := new(big.Int) + + var encoded []byte + for num.Cmp(zero) > 0 { + num.DivMod(num, base, mod) + encoded = append([]byte{base62Chars[mod.Int64()]}, encoded...) + } + if len(encoded) == 0 { + return "0" + } + return string(encoded) +} diff --git a/tai/taiid/taiid_test.go b/tai/taiid/taiid_test.go new file mode 100644 index 00000000..66d6d2af --- /dev/null +++ b/tai/taiid/taiid_test.go @@ -0,0 +1,47 @@ +package taiid + +import ( + "testing" +) + +func TestGenerate_Deterministic(t *testing.T) { + id1, err := Generate("machine-abc", "9100") + if err != nil { + t.Fatalf("Generate: %v", err) + } + id2, err := Generate("machine-abc", "9100") + if err != nil { + t.Fatalf("Generate: %v", err) + } + if id1 != id2 { + t.Errorf("same inputs produced different results: %q vs %q", id1, id2) + } + if len(id1) < 5 || id1[:4] != "tai-" { + t.Errorf("result should start with 'tai-', got %q", id1) + } +} + +func TestGenerate_DifferentInputs(t *testing.T) { + id1, _ := Generate("machine-abc", "9100") + id2, _ := Generate("machine-abc", "9200") + id3, _ := Generate("machine-xyz", "9100") + + if id1 == id2 { + t.Errorf("different nodeID should produce different results: %q == %q", id1, id2) + } + if id1 == id3 { + t.Errorf("different machineID should produce different results: %q == %q", id1, id3) + } +} + +func TestGenerate_EmptyInputs(t *testing.T) { + if _, err := Generate("", "9100"); err == nil { + t.Error("empty machineID should return error") + } + if _, err := Generate("machine-abc", ""); err == nil { + t.Error("empty nodeID should return error") + } + if _, err := Generate("", ""); err == nil { + t.Error("both empty should return error") + } +} diff --git a/tai/tunnel/server.go b/tai/tunnel/server.go index 0429e3b8..3f0bb002 100644 --- a/tai/tunnel/server.go +++ b/tai/tunnel/server.go @@ -13,7 +13,9 @@ import ( "github.com/gin-gonic/gin" "github.com/gorilla/websocket" oauth "github.com/yaoapp/yao/openapi/oauth" + tai "github.com/yaoapp/yao/tai" "github.com/yaoapp/yao/tai/registry" + "github.com/yaoapp/yao/tai/taiid" ) var upgrader = websocket.Upgrader{ @@ -62,19 +64,32 @@ func HandleControl(c *gin.Context) { conn.Close() return } - if regMsg.TaiID == "" { - logger.Error("register message missing tai_id") + if regMsg.NodeID == "" || regMsg.MachineID == "" { + logger.Error("register message missing node_id or machine_id") + conn.Close() + return + } + resolvedTaiID, err := taiid.Generate(regMsg.MachineID, regMsg.NodeID) + if err != nil { + logger.Error("taiid generation failed", "err", err) conn.Close() return } + addr := "" + if host, _, err := net.SplitHostPort(c.Request.RemoteAddr); err == nil { + addr = "tunnel://" + host + } + node := ®istry.TaiNode{ - TaiID: regMsg.TaiID, + TaiID: resolvedTaiID, MachineID: regMsg.MachineID, Version: regMsg.Version, + DisplayName: regMsg.DisplayName, Auth: authInfo, System: regMsg.System, Mode: "tunnel", + Addr: addr, YaoBase: regMsg.Server, Ports: regMsg.Ports, Capabilities: regMsg.Capabilities, @@ -82,16 +97,18 @@ func HandleControl(c *gin.Context) { } reg.Register(node) defer func() { - reg.Unregister(regMsg.TaiID) - logger.Info("tai tunnel disconnected", "tai_id", regMsg.TaiID) + reg.Unregister(resolvedTaiID) + logger.Info("tai tunnel disconnected", "tai_id", resolvedTaiID) }() - if err := reg.WriteControlJSON(regMsg.TaiID, map[string]string{"type": "registered", "tai_id": regMsg.TaiID}); err != nil { + if err := reg.WriteControlJSON(resolvedTaiID, map[string]string{"type": "registered", "tai_id": resolvedTaiID}); err != nil { logger.Error("write registered response", "err", err) return } - logger.Info("tai tunnel connected", "tai_id", regMsg.TaiID, "version", regMsg.Version) + logger.Info("tai tunnel connected", "tai_id", resolvedTaiID, "version", regMsg.Version) + + go connectTunnelNode(resolvedTaiID, reg, logger) for { var msg controlMsg @@ -104,8 +121,8 @@ func HandleControl(c *gin.Context) { switch msg.Type { case "ping": - reg.UpdatePing(regMsg.TaiID) - if err := reg.WriteControlJSON(regMsg.TaiID, map[string]string{"type": "pong"}); err != nil { + reg.UpdatePing(resolvedTaiID) + if err := reg.WriteControlJSON(resolvedTaiID, map[string]string{"type": "pong"}); err != nil { logger.Debug("pong write failed", "err", err) return } @@ -150,9 +167,15 @@ func HandleData(c *gin.Context) { return } + resolvedTaiID := reg.FindTaiIDByAuthClient(authInfo.ClientID) + if resolvedTaiID == "" { + resolvedTaiID = authInfo.ClientID + } + wsConn := newWSConn(conn) - if err := reg.AcceptDataChannel(channelID, authInfo.ClientID, wsConn); err != nil { - logger.Debug("accept data channel failed", "channel_id", channelID, "err", err) + if err := reg.AcceptDataChannel(channelID, resolvedTaiID, wsConn); err != nil { + logger.Debug("accept data channel failed", "channel_id", channelID, "err", err, + "auth_client_id", authInfo.ClientID, "resolved_tai_id", resolvedTaiID) conn.Close() return } @@ -161,8 +184,10 @@ func HandleData(c *gin.Context) { // registerMessage is the JSON structure for Tai's register message. type registerMessage struct { Type string `json:"type"` - TaiID string `json:"tai_id"` + NodeID string `json:"node_id,omitempty"` + ClientID string `json:"client_id,omitempty"` MachineID string `json:"machine_id"` + DisplayName string `json:"display_name,omitempty"` Version string `json:"version"` Server string `json:"server"` Ports map[string]int `json:"ports"` @@ -207,6 +232,37 @@ func authenticateBearerDefault(token string) (registry.AuthInfo, error) { info.TeamID = result.Info.TeamID info.TenantID = result.Info.TenantID } + + slog.Info("[tunnel-auth] info from token", + "subject", info.Subject, "user_id", info.UserID, + "client_id", info.ClientID, "team_id", info.TeamID, + "scope", info.Scope) + + if result.Claims != nil { + slog.Info("[tunnel-auth] claims", + "claims.TeamID", result.Claims.TeamID, + "claims.ClientID", result.Claims.ClientID, + "extra", fmt.Sprintf("%+v", result.Claims.Extra)) + + if info.TeamID == "" && result.Claims.TeamID != "" { + info.TeamID = result.Claims.TeamID + } + if info.TeamID == "" { + switch v := result.Claims.Extra["team_id"].(type) { + case string: + info.TeamID = v + case float64: + info.TeamID = fmt.Sprintf("%.0f", v) + } + } + if info.TenantID == "" { + if v, ok := result.Claims.Extra["tenant_id"].(string); ok { + info.TenantID = v + } + } + } + + slog.Info("[tunnel-auth] final", "team_id", info.TeamID, "client_id", info.ClientID) return info, nil } @@ -267,3 +323,15 @@ func (c *wsConn) SetDeadline(t time.Time) error { func (c *wsConn) SetReadDeadline(t time.Time) error { return c.ws.SetReadDeadline(t) } func (c *wsConn) SetWriteDeadline(t time.Time) error { return c.ws.SetWriteDeadline(t) } + +// connectTunnelNode creates a tai.Client through the tunnel and binds it to the taiID. +func connectTunnelNode(taiID string, reg *registry.Registry, logger *slog.Logger) { + client, err := tai.New("tunnel://" + taiID) + if err != nil { + logger.Warn("failed to connect tunnel node", + "tai_id", taiID, "err", err) + return + } + _ = client // initTunnel already calls reg.SetClient(taiID, c) + logger.Info("tai client created for tunnel node", "tai_id", taiID) +} diff --git a/tai/tunnel/server_test.go b/tai/tunnel/server_test.go index 186e13b5..dfab4fde 100644 --- a/tai/tunnel/server_test.go +++ b/tai/tunnel/server_test.go @@ -281,7 +281,7 @@ func TestHandleControl_RegisterAndPing(t *testing.T) { regMsg := registerMessage{ Type: "register", - TaiID: "tai-001", + NodeID: "9100", MachineID: "m-test", Version: "2.0", Ports: map[string]int{"grpc": 9100}, @@ -297,11 +297,12 @@ func TestHandleControl_RegisterAndPing(t *testing.T) { if registered["type"] != "registered" { t.Errorf("response type = %q, want registered", registered["type"]) } - if registered["tai_id"] != "tai-001" { - t.Errorf("response tai_id = %q, want tai-001", registered["tai_id"]) + gotTaiID := registered["tai_id"] + if gotTaiID == "" || len(gotTaiID) < 5 || gotTaiID[:4] != "tai-" { + t.Errorf("response tai_id = %q, want server-generated tai-xxx", gotTaiID) } - snap, ok := reg.Get("tai-001") + snap, ok := reg.Get(gotTaiID) if !ok { t.Fatal("node not found in registry after register") } @@ -332,15 +333,24 @@ func TestHandleControl_RegisterAndPing(t *testing.T) { t.Fatalf("write ping: %v", err) } - var pong map[string]string - if err := conn.ReadJSON(&pong); err != nil { - t.Fatalf("read pong: %v", err) + // Read messages until we get the pong; connectTunnelNode may inject + // "open" messages (with numeric fields) before our pong arrives. + var gotPong bool + for i := 0; i < 10; i++ { + var msg map[string]interface{} + if err := conn.ReadJSON(&msg); err != nil { + t.Fatalf("read message: %v", err) + } + if msg["type"] == "pong" { + gotPong = true + break + } } - if pong["type"] != "pong" { - t.Errorf("pong type = %q, want pong", pong["type"]) + if !gotPong { + t.Error("did not receive pong after ping") } - snap2, _ := reg.Get("tai-001") + snap2, _ := reg.Get(gotTaiID) if !snap2.LastPing.After(snap.LastPing) { t.Error("LastPing should be updated after ping") } @@ -529,9 +539,10 @@ func TestHandleControl_OpenChannelAndBridge(t *testing.T) { defer ctrlConn.Close() ctrlConn.WriteJSON(registerMessage{ - Type: "register", - TaiID: "tai-001", - Ports: map[string]int{"grpc": 9100}, + Type: "register", + NodeID: "9100", + MachineID: "m-test", + Ports: map[string]int{"grpc": 9100}, }) var registered map[string]string if err := ctrlConn.ReadJSON(®istered); err != nil { @@ -540,6 +551,7 @@ func TestHandleControl_OpenChannelAndBridge(t *testing.T) { if registered["type"] != "registered" { t.Fatalf("expected registered, got %v", registered) } + taiID := registered["tai_id"] var wg sync.WaitGroup wg.Add(1) @@ -547,7 +559,7 @@ func TestHandleControl_OpenChannelAndBridge(t *testing.T) { var channelConn net.Conn go func() { defer wg.Done() - _, resultCh, err := reg.RequestChannel("tai-001", 9100) + _, resultCh, err := reg.RequestChannel(taiID, 9100) if err != nil { requestErr = err return diff --git a/tai/volume/local.go b/tai/volume/local.go index 49462fb0..11ca1e11 100644 --- a/tai/volume/local.go +++ b/tai/volume/local.go @@ -1,7 +1,12 @@ package volume import ( + "archive/tar" + "archive/zip" + "compress/gzip" "context" + "fmt" + "io" "io/fs" "os" "path/filepath" @@ -133,12 +138,126 @@ func (l *localStorage) MkdirAll(_ context.Context, sessionID, path string) error return os.MkdirAll(abs, 0o755) } +// Copy duplicates src to dst within the same workspace session. +// Supports single files and directories (recursive). Uses excludes from SyncOption +// and forceFull to overwrite even when mtime+size match. +func (l *localStorage) Copy(_ context.Context, sessionID, src, dst string, opts ...SyncOption) (*SyncResult, error) { + start := time.Now() + cfg := ApplySyncOpts(opts) + + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + + srcInfo, err := os.Stat(srcAbs) + if err != nil { + return nil, err + } + + if !srcInfo.IsDir() { + n, err := l.copyFile(srcAbs, dstAbs, srcInfo, cfg.ForceFull) + if err != nil { + return nil, err + } + synced := 0 + if n > 0 { + synced = 1 + } + return &SyncResult{ + FilesSynced: synced, + BytesTransferred: n, + Duration: time.Since(start), + }, nil + } + + var synced int + var transferred int64 + err = filepath.WalkDir(srcAbs, func(abs string, d fs.DirEntry, walkErr error) error { + if walkErr != nil { + if os.IsNotExist(walkErr) { + return nil + } + return walkErr + } + rel, _ := filepath.Rel(srcAbs, abs) + if rel == "." { + return os.MkdirAll(dstAbs, 0o755) + } + + if isExcluded(rel, d.IsDir(), cfg.Excludes) { + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + + target := filepath.Join(dstAbs, rel) + if d.IsDir() { + return os.MkdirAll(target, 0o755) + } + + info, err := d.Info() + if err != nil { + return nil + } + n, err := l.copyFile(abs, target, info, cfg.ForceFull) + if err != nil { + return err + } + if n > 0 { + synced++ + transferred += n + } + return nil + }) + + return &SyncResult{ + FilesSynced: synced, + BytesTransferred: transferred, + Duration: time.Since(start), + }, err +} + +func (l *localStorage) copyFile(srcAbs, dstAbs string, srcInfo os.FileInfo, force bool) (int64, error) { + if !force { + if dstInfo, e := os.Stat(dstAbs); e == nil { + if dstInfo.Size() == srcInfo.Size() && dstInfo.ModTime().Equal(srcInfo.ModTime()) { + return 0, nil + } + } + } + + data, err := os.ReadFile(srcAbs) + if err != nil { + if os.IsNotExist(err) { + return 0, nil + } + return 0, err + } + if err := os.MkdirAll(filepath.Dir(dstAbs), 0o755); err != nil { + return 0, err + } + if err := os.WriteFile(dstAbs, data, srcInfo.Mode()); err != nil { + return 0, err + } + _ = os.Chtimes(dstAbs, srcInfo.ModTime(), srcInfo.ModTime()) + return int64(len(data)), nil +} + // SyncPush copies changed files from localDir to dataDir/{sessionID}/. // Uses mtime+size to detect changes. Files that vanish during sync are skipped. func (l *localStorage) SyncPush(_ context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) { start := time.Now() - cfg := applySyncOpts(opts) + cfg := ApplySyncOpts(opts) dst := l.root(sessionID) + if cfg.RemotePath != "" { + dst = filepath.Join(dst, filepath.Clean(cfg.RemotePath)) + } if err := os.MkdirAll(dst, 0o755); err != nil { return nil, err } @@ -159,7 +278,7 @@ func (l *localStorage) SyncPush(_ context.Context, sessionID, localDir string, o } rel = filepath.ToSlash(rel) - if isExcluded(rel, d.IsDir(), cfg.excludes) { + if isExcluded(rel, d.IsDir(), cfg.Excludes) { if d.IsDir() { return filepath.SkipDir } @@ -176,7 +295,7 @@ func (l *localStorage) SyncPush(_ context.Context, sessionID, localDir string, o return nil // file vanished between readdir and stat; skip } - if !cfg.forceFull { + if !cfg.ForceFull { if dstInfo, e := os.Stat(target); e == nil { if dstInfo.Size() == srcInfo.Size() && dstInfo.ModTime().Equal(srcInfo.ModTime()) { return nil @@ -214,8 +333,11 @@ func (l *localStorage) SyncPush(_ context.Context, sessionID, localDir string, o // Files that vanish during sync are skipped. func (l *localStorage) SyncPull(_ context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) { start := time.Now() - cfg := applySyncOpts(opts) + cfg := ApplySyncOpts(opts) src := l.root(sessionID) + if cfg.RemotePath != "" { + src = filepath.Join(src, filepath.Clean(cfg.RemotePath)) + } if err := os.MkdirAll(localDir, 0o755); err != nil { return nil, err } @@ -236,7 +358,7 @@ func (l *localStorage) SyncPull(_ context.Context, sessionID, localDir string, o } rel = filepath.ToSlash(rel) - if isExcluded(rel, d.IsDir(), cfg.excludes) { + if isExcluded(rel, d.IsDir(), cfg.Excludes) { if d.IsDir() { return filepath.SkipDir } @@ -253,7 +375,7 @@ func (l *localStorage) SyncPull(_ context.Context, sessionID, localDir string, o return nil // file vanished between readdir and stat; skip } - if !cfg.forceFull { + if !cfg.ForceFull { if dstInfo, e := os.Stat(target); e == nil { if dstInfo.Size() == srcInfo.Size() && dstInfo.ModTime().Equal(srcInfo.ModTime()) { return nil @@ -287,6 +409,365 @@ func (l *localStorage) SyncPull(_ context.Context, sessionID, localDir string, o }, err } +func (l *localStorage) Zip(_ context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) { + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + if err := os.MkdirAll(filepath.Dir(dstAbs), 0o755); err != nil { + return nil, err + } + out, err := os.Create(dstAbs) + if err != nil { + return nil, err + } + defer out.Close() + w := zip.NewWriter(out) + defer w.Close() + var count int + if err := filepath.WalkDir(srcAbs, func(abs string, d fs.DirEntry, err error) error { + if err != nil { + return err + } + rel, _ := filepath.Rel(srcAbs, abs) + if rel == "." { + return nil + } + rel = filepath.ToSlash(rel) + if isExcluded(rel, d.IsDir(), excludes) { + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + if d.IsDir() { + _, e := w.Create(rel + "/") + return e + } + info, err := d.Info() + if err != nil { + return err + } + header, err := zip.FileInfoHeader(info) + if err != nil { + return err + } + header.Name = rel + header.Method = zip.Deflate + writer, err := w.CreateHeader(header) + if err != nil { + return err + } + f, err := os.Open(abs) + if err != nil { + return err + } + defer f.Close() + _, err = io.Copy(writer, f) + if err == nil { + count++ + } + return err + }); err != nil { + return nil, err + } + w.Close() + out.Close() + fi, _ := os.Stat(dstAbs) + return &ArchiveResult{SizeBytes: fi.Size(), FilesCount: count}, nil +} + +func (l *localStorage) Unzip(_ context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + r, err := zip.OpenReader(srcAbs) + if err != nil { + return nil, err + } + defer r.Close() + if err := os.MkdirAll(dstAbs, 0o755); err != nil { + return nil, err + } + var count int + var totalSize int64 + for _, f := range r.File { + target := filepath.Join(dstAbs, filepath.FromSlash(f.Name)) + if !strings.HasPrefix(target, dstAbs+string(filepath.Separator)) && target != dstAbs { + return nil, fmt.Errorf("zip slip: %s", f.Name) + } + if f.FileInfo().IsDir() { + _ = os.MkdirAll(target, 0o755) + continue + } + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return nil, err + } + rc, err := f.Open() + if err != nil { + return nil, err + } + out, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, f.Mode()) + if err != nil { + rc.Close() + return nil, err + } + n, err := io.Copy(out, rc) + out.Close() + rc.Close() + if err != nil { + return nil, err + } + totalSize += n + count++ + } + return &ArchiveResult{SizeBytes: totalSize, FilesCount: count}, nil +} + +func (l *localStorage) Gzip(_ context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + info, err := os.Stat(srcAbs) + if err != nil { + return nil, err + } + if info.IsDir() { + return nil, fmt.Errorf("gzip requires a file, not directory") + } + in, err := os.Open(srcAbs) + if err != nil { + return nil, err + } + defer in.Close() + if err := os.MkdirAll(filepath.Dir(dstAbs), 0o755); err != nil { + return nil, err + } + out, err := os.Create(dstAbs) + if err != nil { + return nil, err + } + defer out.Close() + w := gzip.NewWriter(out) + w.Name = filepath.Base(srcAbs) + if _, err := io.Copy(w, in); err != nil { + w.Close() + return nil, err + } + w.Close() + out.Close() + fi, _ := os.Stat(dstAbs) + return &ArchiveResult{SizeBytes: fi.Size(), FilesCount: 1}, nil +} + +func (l *localStorage) Gunzip(_ context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + in, err := os.Open(srcAbs) + if err != nil { + return nil, err + } + defer in.Close() + r, err := gzip.NewReader(in) + if err != nil { + return nil, err + } + defer r.Close() + if err := os.MkdirAll(filepath.Dir(dstAbs), 0o755); err != nil { + return nil, err + } + out, err := os.Create(dstAbs) + if err != nil { + return nil, err + } + defer out.Close() + n, err := io.Copy(out, r) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: n, FilesCount: 1}, nil +} + +func (l *localStorage) Tar(_ context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) { + return l.tarImpl(sessionID, src, dst, excludes, false) +} + +func (l *localStorage) Tgz(_ context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) { + return l.tarImpl(sessionID, src, dst, excludes, true) +} + +func (l *localStorage) tarImpl(sessionID, src, dst string, excludes []string, useGzip bool) (*ArchiveResult, error) { + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + if err := os.MkdirAll(filepath.Dir(dstAbs), 0o755); err != nil { + return nil, err + } + out, err := os.Create(dstAbs) + if err != nil { + return nil, err + } + defer out.Close() + var tw *tar.Writer + var gw *gzip.Writer + if useGzip { + gw = gzip.NewWriter(out) + defer gw.Close() + tw = tar.NewWriter(gw) + } else { + tw = tar.NewWriter(out) + } + defer tw.Close() + var count int + if err := filepath.WalkDir(srcAbs, func(abs string, d fs.DirEntry, err error) error { + if err != nil { + return err + } + rel, _ := filepath.Rel(srcAbs, abs) + if rel == "." { + return nil + } + rel = filepath.ToSlash(rel) + if isExcluded(rel, d.IsDir(), excludes) { + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + info, err := d.Info() + if err != nil { + return err + } + header, err := tar.FileInfoHeader(info, "") + if err != nil { + return err + } + header.Name = rel + if err := tw.WriteHeader(header); err != nil { + return err + } + if d.IsDir() { + return nil + } + f, err := os.Open(abs) + if err != nil { + return err + } + defer f.Close() + _, err = io.Copy(tw, f) + if err == nil { + count++ + } + return err + }); err != nil { + return nil, err + } + tw.Close() + if gw != nil { + gw.Close() + } + out.Close() + fi, _ := os.Stat(dstAbs) + return &ArchiveResult{SizeBytes: fi.Size(), FilesCount: count}, nil +} + +func (l *localStorage) Untar(_ context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + return l.untarImpl(sessionID, src, dst, false) +} + +func (l *localStorage) Untgz(_ context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + return l.untarImpl(sessionID, src, dst, true) +} + +func (l *localStorage) untarImpl(sessionID, src, dst string, useGzip bool) (*ArchiveResult, error) { + srcAbs, err := l.abs(sessionID, src) + if err != nil { + return nil, err + } + dstAbs, err := l.abs(sessionID, dst) + if err != nil { + return nil, err + } + in, err := os.Open(srcAbs) + if err != nil { + return nil, err + } + defer in.Close() + var reader io.Reader = in + if useGzip { + gr, err := gzip.NewReader(in) + if err != nil { + return nil, err + } + defer gr.Close() + reader = gr + } + tr := tar.NewReader(reader) + if err := os.MkdirAll(dstAbs, 0o755); err != nil { + return nil, err + } + var count int + var totalSize int64 + for { + header, err := tr.Next() + if err == io.EOF { + break + } + if err != nil { + return nil, err + } + target := filepath.Join(dstAbs, filepath.FromSlash(header.Name)) + if !strings.HasPrefix(target, dstAbs+string(filepath.Separator)) && target != dstAbs { + return nil, fmt.Errorf("tar slip: %s", header.Name) + } + switch header.Typeflag { + case tar.TypeDir: + _ = os.MkdirAll(target, 0o755) + case tar.TypeReg: + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return nil, err + } + out, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, os.FileMode(header.Mode)) + if err != nil { + return nil, err + } + n, err := io.Copy(out, tr) + out.Close() + if err != nil { + return nil, err + } + totalSize += n + count++ + } + } + return &ArchiveResult{SizeBytes: totalSize, FilesCount: count}, nil +} + func (l *localStorage) Close() error { return nil } func isExcluded(rel string, isDir bool, patterns []string) bool { diff --git a/tai/volume/mock_test.go b/tai/volume/mock_test.go index 59f0ce06..bdee3352 100644 --- a/tai/volume/mock_test.go +++ b/tai/volume/mock_test.go @@ -197,6 +197,13 @@ func (m *mockVolumeServer) ListDir(_ context.Context, req *pb.FSRequest) (*pb.FS }}, nil } +func (m *mockVolumeServer) Copy(_ context.Context, req *pb.FSCopyRequest) (*pb.SyncResult, error) { + return &pb.SyncResult{ + FilesSynced: 1, + BytesTransferred: 42, + }, nil +} + func startMockServer(t *testing.T, mock *mockVolumeServer) (*grpc.ClientConn, func()) { t.Helper() lis, err := net.Listen("tcp", "127.0.0.1:0") @@ -549,6 +556,10 @@ func (m *errMockVolumeServer) MkdirAll(_ context.Context, _ *pb.FSRequest) (*pb. return nil, fmt.Errorf("injected mkdir error") } +func (m *errMockVolumeServer) Copy(_ context.Context, _ *pb.FSCopyRequest) (*pb.SyncResult, error) { + return nil, fmt.Errorf("injected copy error") +} + func startErrMockServer(t *testing.T) (*grpc.ClientConn, func()) { t.Helper() lis, err := net.Listen("tcp", "127.0.0.1:0") @@ -694,3 +705,46 @@ func TestPbToFileInfo(t *testing.T) { t.Errorf("mode = %v", fi.Mode) } } + +func TestMockRemoteCopy(t *testing.T) { + conn, cleanup := startMockServer(t, &mockVolumeServer{}) + defer cleanup() + + vol := NewRemote(conn) + result, err := vol.Copy(context.Background(), "s1", "src.txt", "dst.txt") + if err != nil { + t.Fatalf("Copy: %v", err) + } + if result.FilesSynced != 1 { + t.Errorf("synced = %d, want 1", result.FilesSynced) + } + if result.BytesTransferred != 42 { + t.Errorf("bytes = %d, want 42", result.BytesTransferred) + } +} + +func TestMockRemoteCopyWithOpts(t *testing.T) { + conn, cleanup := startMockServer(t, &mockVolumeServer{}) + defer cleanup() + + vol := NewRemote(conn) + result, err := vol.Copy(context.Background(), "s1", "src", "dst", + WithExcludes("*.log"), WithForceFull()) + if err != nil { + t.Fatalf("Copy: %v", err) + } + if result.FilesSynced != 1 { + t.Errorf("synced = %d", result.FilesSynced) + } +} + +func TestErrRemoteCopy(t *testing.T) { + conn, cleanup := startErrMockServer(t) + defer cleanup() + + vol := NewRemote(conn) + _, err := vol.Copy(context.Background(), "s1", "a", "b") + if err == nil { + t.Error("expected error") + } +} diff --git a/tai/volume/pb/volume.pb.go b/tai/volume/pb/volume.pb.go index 44460723..96e97bd8 100644 --- a/tai/volume/pb/volume.pb.go +++ b/tai/volume/pb/volume.pb.go @@ -2,7 +2,7 @@ // versions: // protoc-gen-go v1.36.11 // protoc v4.25.0 -// source: volume/pb/volume.proto +// source: tai/volume/pb/volume.proto package pb @@ -57,11 +57,11 @@ func (x FileChunk_ChunkType) String() string { } func (FileChunk_ChunkType) Descriptor() protoreflect.EnumDescriptor { - return file_volume_pb_volume_proto_enumTypes[0].Descriptor() + return file_tai_volume_pb_volume_proto_enumTypes[0].Descriptor() } func (FileChunk_ChunkType) Type() protoreflect.EnumType { - return &file_volume_pb_volume_proto_enumTypes[0] + return &file_tai_volume_pb_volume_proto_enumTypes[0] } func (x FileChunk_ChunkType) Number() protoreflect.EnumNumber { @@ -70,7 +70,7 @@ func (x FileChunk_ChunkType) Number() protoreflect.EnumNumber { // Deprecated: Use FileChunk_ChunkType.Descriptor instead. func (FileChunk_ChunkType) EnumDescriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{4, 0} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{4, 0} } type FileInfo struct { @@ -86,7 +86,7 @@ type FileInfo struct { func (x *FileInfo) Reset() { *x = FileInfo{} - mi := &file_volume_pb_volume_proto_msgTypes[0] + mi := &file_tai_volume_pb_volume_proto_msgTypes[0] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -98,7 +98,7 @@ func (x *FileInfo) String() string { func (*FileInfo) ProtoMessage() {} func (x *FileInfo) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[0] + mi := &file_tai_volume_pb_volume_proto_msgTypes[0] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -111,7 +111,7 @@ func (x *FileInfo) ProtoReflect() protoreflect.Message { // Deprecated: Use FileInfo.ProtoReflect.Descriptor instead. func (*FileInfo) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{0} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{0} } func (x *FileInfo) GetPath() string { @@ -153,14 +153,15 @@ type SyncManifest struct { state protoimpl.MessageState `protogen:"open.v1"` SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` Files []*FileInfo `protobuf:"bytes,2,rep,name=files,proto3" json:"files,omitempty"` - ForceFull bool `protobuf:"varint,3,opt,name=force_full,json=forceFull,proto3" json:"force_full,omitempty"` // skip snapshot cache, diff against actual disk + ForceFull bool `protobuf:"varint,3,opt,name=force_full,json=forceFull,proto3" json:"force_full,omitempty"` // skip snapshot cache, diff against actual disk + RemotePath string `protobuf:"bytes,4,opt,name=remote_path,json=remotePath,proto3" json:"remote_path,omitempty"` // sub-path within workspace root; empty = root unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } func (x *SyncManifest) Reset() { *x = SyncManifest{} - mi := &file_volume_pb_volume_proto_msgTypes[1] + mi := &file_tai_volume_pb_volume_proto_msgTypes[1] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -172,7 +173,7 @@ func (x *SyncManifest) String() string { func (*SyncManifest) ProtoMessage() {} func (x *SyncManifest) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[1] + mi := &file_tai_volume_pb_volume_proto_msgTypes[1] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -185,7 +186,7 @@ func (x *SyncManifest) ProtoReflect() protoreflect.Message { // Deprecated: Use SyncManifest.ProtoReflect.Descriptor instead. func (*SyncManifest) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{1} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{1} } func (x *SyncManifest) GetSessionId() string { @@ -209,6 +210,13 @@ func (x *SyncManifest) GetForceFull() bool { return false } +func (x *SyncManifest) GetRemotePath() string { + if x != nil { + return x.RemotePath + } + return "" +} + type SyncMessage struct { state protoimpl.MessageState `protogen:"open.v1"` // Types that are valid to be assigned to Payload: @@ -224,7 +232,7 @@ type SyncMessage struct { func (x *SyncMessage) Reset() { *x = SyncMessage{} - mi := &file_volume_pb_volume_proto_msgTypes[2] + mi := &file_tai_volume_pb_volume_proto_msgTypes[2] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -236,7 +244,7 @@ func (x *SyncMessage) String() string { func (*SyncMessage) ProtoMessage() {} func (x *SyncMessage) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[2] + mi := &file_tai_volume_pb_volume_proto_msgTypes[2] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -249,7 +257,7 @@ func (x *SyncMessage) ProtoReflect() protoreflect.Message { // Deprecated: Use SyncMessage.ProtoReflect.Descriptor instead. func (*SyncMessage) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{2} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{2} } func (x *SyncMessage) GetPayload() isSyncMessage_Payload { @@ -333,7 +341,7 @@ type SyncDiff struct { func (x *SyncDiff) Reset() { *x = SyncDiff{} - mi := &file_volume_pb_volume_proto_msgTypes[3] + mi := &file_tai_volume_pb_volume_proto_msgTypes[3] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -345,7 +353,7 @@ func (x *SyncDiff) String() string { func (*SyncDiff) ProtoMessage() {} func (x *SyncDiff) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[3] + mi := &file_tai_volume_pb_volume_proto_msgTypes[3] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -358,7 +366,7 @@ func (x *SyncDiff) ProtoReflect() protoreflect.Message { // Deprecated: Use SyncDiff.ProtoReflect.Descriptor instead. func (*SyncDiff) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{3} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{3} } func (x *SyncDiff) GetNeedFiles() []string { @@ -389,7 +397,7 @@ type FileChunk struct { func (x *FileChunk) Reset() { *x = FileChunk{} - mi := &file_volume_pb_volume_proto_msgTypes[4] + mi := &file_tai_volume_pb_volume_proto_msgTypes[4] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -401,7 +409,7 @@ func (x *FileChunk) String() string { func (*FileChunk) ProtoMessage() {} func (x *FileChunk) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[4] + mi := &file_tai_volume_pb_volume_proto_msgTypes[4] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -414,7 +422,7 @@ func (x *FileChunk) ProtoReflect() protoreflect.Message { // Deprecated: Use FileChunk.ProtoReflect.Descriptor instead. func (*FileChunk) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{4} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{4} } func (x *FileChunk) GetPath() string { @@ -470,7 +478,7 @@ type SyncResult struct { func (x *SyncResult) Reset() { *x = SyncResult{} - mi := &file_volume_pb_volume_proto_msgTypes[5] + mi := &file_tai_volume_pb_volume_proto_msgTypes[5] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -482,7 +490,7 @@ func (x *SyncResult) String() string { func (*SyncResult) ProtoMessage() {} func (x *SyncResult) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[5] + mi := &file_tai_volume_pb_volume_proto_msgTypes[5] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -495,7 +503,7 @@ func (x *SyncResult) ProtoReflect() protoreflect.Message { // Deprecated: Use SyncResult.ProtoReflect.Descriptor instead. func (*SyncResult) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{5} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{5} } func (x *SyncResult) GetFilesSynced() int32 { @@ -529,7 +537,7 @@ type FSRequest struct { func (x *FSRequest) Reset() { *x = FSRequest{} - mi := &file_volume_pb_volume_proto_msgTypes[6] + mi := &file_tai_volume_pb_volume_proto_msgTypes[6] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -541,7 +549,7 @@ func (x *FSRequest) String() string { func (*FSRequest) ProtoMessage() {} func (x *FSRequest) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[6] + mi := &file_tai_volume_pb_volume_proto_msgTypes[6] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -554,7 +562,7 @@ func (x *FSRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSRequest.ProtoReflect.Descriptor instead. func (*FSRequest) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{6} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{6} } func (x *FSRequest) GetSessionId() string { @@ -581,7 +589,7 @@ type FSOpResponse struct { func (x *FSOpResponse) Reset() { *x = FSOpResponse{} - mi := &file_volume_pb_volume_proto_msgTypes[7] + mi := &file_tai_volume_pb_volume_proto_msgTypes[7] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -593,7 +601,7 @@ func (x *FSOpResponse) String() string { func (*FSOpResponse) ProtoMessage() {} func (x *FSOpResponse) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[7] + mi := &file_tai_volume_pb_volume_proto_msgTypes[7] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -606,7 +614,7 @@ func (x *FSOpResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use FSOpResponse.ProtoReflect.Descriptor instead. func (*FSOpResponse) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{7} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{7} } func (x *FSOpResponse) GetOk() bool { @@ -633,7 +641,7 @@ type FSReadRequest struct { func (x *FSReadRequest) Reset() { *x = FSReadRequest{} - mi := &file_volume_pb_volume_proto_msgTypes[8] + mi := &file_tai_volume_pb_volume_proto_msgTypes[8] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -645,7 +653,7 @@ func (x *FSReadRequest) String() string { func (*FSReadRequest) ProtoMessage() {} func (x *FSReadRequest) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[8] + mi := &file_tai_volume_pb_volume_proto_msgTypes[8] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -658,7 +666,7 @@ func (x *FSReadRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSReadRequest.ProtoReflect.Descriptor instead. func (*FSReadRequest) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{8} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{8} } func (x *FSReadRequest) GetSessionId() string { @@ -687,7 +695,7 @@ type FSDataChunk struct { func (x *FSDataChunk) Reset() { *x = FSDataChunk{} - mi := &file_volume_pb_volume_proto_msgTypes[9] + mi := &file_tai_volume_pb_volume_proto_msgTypes[9] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -699,7 +707,7 @@ func (x *FSDataChunk) String() string { func (*FSDataChunk) ProtoMessage() {} func (x *FSDataChunk) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[9] + mi := &file_tai_volume_pb_volume_proto_msgTypes[9] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -712,7 +720,7 @@ func (x *FSDataChunk) ProtoReflect() protoreflect.Message { // Deprecated: Use FSDataChunk.ProtoReflect.Descriptor instead. func (*FSDataChunk) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{9} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{9} } func (x *FSDataChunk) GetData() []byte { @@ -756,7 +764,7 @@ type FSWriteChunk struct { func (x *FSWriteChunk) Reset() { *x = FSWriteChunk{} - mi := &file_volume_pb_volume_proto_msgTypes[10] + mi := &file_tai_volume_pb_volume_proto_msgTypes[10] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -768,7 +776,7 @@ func (x *FSWriteChunk) String() string { func (*FSWriteChunk) ProtoMessage() {} func (x *FSWriteChunk) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[10] + mi := &file_tai_volume_pb_volume_proto_msgTypes[10] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -781,7 +789,7 @@ func (x *FSWriteChunk) ProtoReflect() protoreflect.Message { // Deprecated: Use FSWriteChunk.ProtoReflect.Descriptor instead. func (*FSWriteChunk) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{10} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{10} } func (x *FSWriteChunk) GetSessionId() string { @@ -828,7 +836,7 @@ type FSWriteResponse struct { func (x *FSWriteResponse) Reset() { *x = FSWriteResponse{} - mi := &file_volume_pb_volume_proto_msgTypes[11] + mi := &file_tai_volume_pb_volume_proto_msgTypes[11] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -840,7 +848,7 @@ func (x *FSWriteResponse) String() string { func (*FSWriteResponse) ProtoMessage() {} func (x *FSWriteResponse) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[11] + mi := &file_tai_volume_pb_volume_proto_msgTypes[11] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -853,7 +861,7 @@ func (x *FSWriteResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use FSWriteResponse.ProtoReflect.Descriptor instead. func (*FSWriteResponse) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{11} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{11} } func (x *FSWriteResponse) GetSize() int64 { @@ -872,7 +880,7 @@ type FSListResponse struct { func (x *FSListResponse) Reset() { *x = FSListResponse{} - mi := &file_volume_pb_volume_proto_msgTypes[12] + mi := &file_tai_volume_pb_volume_proto_msgTypes[12] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -884,7 +892,7 @@ func (x *FSListResponse) String() string { func (*FSListResponse) ProtoMessage() {} func (x *FSListResponse) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[12] + mi := &file_tai_volume_pb_volume_proto_msgTypes[12] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -897,7 +905,7 @@ func (x *FSListResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use FSListResponse.ProtoReflect.Descriptor instead. func (*FSListResponse) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{12} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{12} } func (x *FSListResponse) GetEntries() []*FileInfo { @@ -918,7 +926,7 @@ type FSRemoveRequest struct { func (x *FSRemoveRequest) Reset() { *x = FSRemoveRequest{} - mi := &file_volume_pb_volume_proto_msgTypes[13] + mi := &file_tai_volume_pb_volume_proto_msgTypes[13] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -930,7 +938,7 @@ func (x *FSRemoveRequest) String() string { func (*FSRemoveRequest) ProtoMessage() {} func (x *FSRemoveRequest) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[13] + mi := &file_tai_volume_pb_volume_proto_msgTypes[13] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -943,7 +951,7 @@ func (x *FSRemoveRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSRemoveRequest.ProtoReflect.Descriptor instead. func (*FSRemoveRequest) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{13} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{13} } func (x *FSRemoveRequest) GetSessionId() string { @@ -978,7 +986,7 @@ type FSRenameRequest struct { func (x *FSRenameRequest) Reset() { *x = FSRenameRequest{} - mi := &file_volume_pb_volume_proto_msgTypes[14] + mi := &file_tai_volume_pb_volume_proto_msgTypes[14] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -990,7 +998,7 @@ func (x *FSRenameRequest) String() string { func (*FSRenameRequest) ProtoMessage() {} func (x *FSRenameRequest) ProtoReflect() protoreflect.Message { - mi := &file_volume_pb_volume_proto_msgTypes[14] + mi := &file_tai_volume_pb_volume_proto_msgTypes[14] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1003,7 +1011,7 @@ func (x *FSRenameRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSRenameRequest.ProtoReflect.Descriptor instead. func (*FSRenameRequest) Descriptor() ([]byte, []int) { - return file_volume_pb_volume_proto_rawDescGZIP(), []int{14} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{14} } func (x *FSRenameRequest) GetSessionId() string { @@ -1027,23 +1035,221 @@ func (x *FSRenameRequest) GetNewPath() string { return "" } -var File_volume_pb_volume_proto protoreflect.FileDescriptor +type FSCopyRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + SrcPath string `protobuf:"bytes,2,opt,name=src_path,json=srcPath,proto3" json:"src_path,omitempty"` + DstPath string `protobuf:"bytes,3,opt,name=dst_path,json=dstPath,proto3" json:"dst_path,omitempty"` + Excludes []string `protobuf:"bytes,4,rep,name=excludes,proto3" json:"excludes,omitempty"` // glob patterns + Force bool `protobuf:"varint,5,opt,name=force,proto3" json:"force,omitempty"` // overwrite even if mtime/size match + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} -const file_volume_pb_volume_proto_rawDesc = "" + +func (x *FSCopyRequest) Reset() { + *x = FSCopyRequest{} + mi := &file_tai_volume_pb_volume_proto_msgTypes[15] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *FSCopyRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*FSCopyRequest) ProtoMessage() {} + +func (x *FSCopyRequest) ProtoReflect() protoreflect.Message { + mi := &file_tai_volume_pb_volume_proto_msgTypes[15] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use FSCopyRequest.ProtoReflect.Descriptor instead. +func (*FSCopyRequest) Descriptor() ([]byte, []int) { + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{15} +} + +func (x *FSCopyRequest) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *FSCopyRequest) GetSrcPath() string { + if x != nil { + return x.SrcPath + } + return "" +} + +func (x *FSCopyRequest) GetDstPath() string { + if x != nil { + return x.DstPath + } + return "" +} + +func (x *FSCopyRequest) GetExcludes() []string { + if x != nil { + return x.Excludes + } + return nil +} + +func (x *FSCopyRequest) GetForce() bool { + if x != nil { + return x.Force + } + return false +} + +type ArchiveRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + SrcPath string `protobuf:"bytes,2,opt,name=src_path,json=srcPath,proto3" json:"src_path,omitempty"` // relative to workspace root + DstPath string `protobuf:"bytes,3,opt,name=dst_path,json=dstPath,proto3" json:"dst_path,omitempty"` // relative to workspace root + Excludes []string `protobuf:"bytes,4,rep,name=excludes,proto3" json:"excludes,omitempty"` // glob patterns (pack ops only) + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *ArchiveRequest) Reset() { + *x = ArchiveRequest{} + mi := &file_tai_volume_pb_volume_proto_msgTypes[16] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *ArchiveRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*ArchiveRequest) ProtoMessage() {} + +func (x *ArchiveRequest) ProtoReflect() protoreflect.Message { + mi := &file_tai_volume_pb_volume_proto_msgTypes[16] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use ArchiveRequest.ProtoReflect.Descriptor instead. +func (*ArchiveRequest) Descriptor() ([]byte, []int) { + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{16} +} + +func (x *ArchiveRequest) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *ArchiveRequest) GetSrcPath() string { + if x != nil { + return x.SrcPath + } + return "" +} + +func (x *ArchiveRequest) GetDstPath() string { + if x != nil { + return x.DstPath + } + return "" +} + +func (x *ArchiveRequest) GetExcludes() []string { + if x != nil { + return x.Excludes + } + return nil +} + +type ArchiveResponse struct { + state protoimpl.MessageState `protogen:"open.v1"` + SizeBytes int64 `protobuf:"varint,1,opt,name=size_bytes,json=sizeBytes,proto3" json:"size_bytes,omitempty"` // output file size (pack) or total extracted size (unpack) + FilesCount int32 `protobuf:"varint,2,opt,name=files_count,json=filesCount,proto3" json:"files_count,omitempty"` // number of files processed + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *ArchiveResponse) Reset() { + *x = ArchiveResponse{} + mi := &file_tai_volume_pb_volume_proto_msgTypes[17] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *ArchiveResponse) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*ArchiveResponse) ProtoMessage() {} + +func (x *ArchiveResponse) ProtoReflect() protoreflect.Message { + mi := &file_tai_volume_pb_volume_proto_msgTypes[17] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use ArchiveResponse.ProtoReflect.Descriptor instead. +func (*ArchiveResponse) Descriptor() ([]byte, []int) { + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{17} +} + +func (x *ArchiveResponse) GetSizeBytes() int64 { + if x != nil { + return x.SizeBytes + } + return 0 +} + +func (x *ArchiveResponse) GetFilesCount() int32 { + if x != nil { + return x.FilesCount + } + return 0 +} + +var File_tai_volume_pb_volume_proto protoreflect.FileDescriptor + +const file_tai_volume_pb_volume_proto_rawDesc = "" + "\n" + - "\x16volume/pb/volume.proto\x12\x06volume\"s\n" + + "\x1atai/volume/pb/volume.proto\x12\x06volume\"s\n" + "\bFileInfo\x12\x12\n" + "\x04path\x18\x01 \x01(\tR\x04path\x12\x12\n" + "\x04size\x18\x02 \x01(\x03R\x04size\x12\x14\n" + "\x05mtime\x18\x03 \x01(\x03R\x05mtime\x12\x12\n" + "\x04mode\x18\x04 \x01(\rR\x04mode\x12\x15\n" + - "\x06is_dir\x18\x05 \x01(\bR\x05isDir\"t\n" + + "\x06is_dir\x18\x05 \x01(\bR\x05isDir\"\x95\x01\n" + "\fSyncManifest\x12\x1d\n" + "\n" + "session_id\x18\x01 \x01(\tR\tsessionId\x12&\n" + "\x05files\x18\x02 \x03(\v2\x10.volume.FileInfoR\x05files\x12\x1d\n" + "\n" + - "force_full\x18\x03 \x01(\bR\tforceFull\"\xcd\x01\n" + + "force_full\x18\x03 \x01(\bR\tforceFull\x12\x1f\n" + + "\vremote_path\x18\x04 \x01(\tR\n" + + "remotePath\"\xcd\x01\n" + "\vSyncMessage\x122\n" + "\bmanifest\x18\x01 \x01(\v2\x14.volume.SyncManifestH\x00R\bmanifest\x12&\n" + "\x04diff\x18\x02 \x01(\v2\x10.volume.SyncDiffH\x00R\x04diff\x12)\n" + @@ -1110,7 +1316,25 @@ const file_volume_pb_volume_proto_rawDesc = "" + "\n" + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x19\n" + "\bold_path\x18\x02 \x01(\tR\aoldPath\x12\x19\n" + - "\bnew_path\x18\x03 \x01(\tR\anewPath2\xfd\x03\n" + + "\bnew_path\x18\x03 \x01(\tR\anewPath\"\x96\x01\n" + + "\rFSCopyRequest\x12\x1d\n" + + "\n" + + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x19\n" + + "\bsrc_path\x18\x02 \x01(\tR\asrcPath\x12\x19\n" + + "\bdst_path\x18\x03 \x01(\tR\adstPath\x12\x1a\n" + + "\bexcludes\x18\x04 \x03(\tR\bexcludes\x12\x14\n" + + "\x05force\x18\x05 \x01(\bR\x05force\"\x81\x01\n" + + "\x0eArchiveRequest\x12\x1d\n" + + "\n" + + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x19\n" + + "\bsrc_path\x18\x02 \x01(\tR\asrcPath\x12\x19\n" + + "\bdst_path\x18\x03 \x01(\tR\adstPath\x12\x1a\n" + + "\bexcludes\x18\x04 \x03(\tR\bexcludes\"Q\n" + + "\x0fArchiveResponse\x12\x1d\n" + + "\n" + + "size_bytes\x18\x01 \x01(\x03R\tsizeBytes\x12\x1f\n" + + "\vfiles_count\x18\x02 \x01(\x05R\n" + + "filesCount2\xfa\a\n" + "\x06Volume\x128\n" + "\bSyncPush\x12\x13.volume.SyncMessage\x1a\x13.volume.SyncMessage(\x010\x01\x127\n" + "\bSyncPull\x12\x14.volume.SyncManifest\x1a\x13.volume.SyncMessage0\x01\x128\n" + @@ -1120,23 +1344,32 @@ const file_volume_pb_volume_proto_rawDesc = "" + "\aListDir\x12\x11.volume.FSRequest\x1a\x16.volume.FSListResponse\x127\n" + "\x06Remove\x12\x17.volume.FSRemoveRequest\x1a\x14.volume.FSOpResponse\x127\n" + "\x06Rename\x12\x17.volume.FSRenameRequest\x1a\x14.volume.FSOpResponse\x123\n" + - "\bMkdirAll\x12\x11.volume.FSRequest\x1a\x14.volume.FSOpResponseB!Z\x1fgithub.com/yaoapp/tai/volume/pbb\x06proto3" + "\bMkdirAll\x12\x11.volume.FSRequest\x1a\x14.volume.FSOpResponse\x121\n" + + "\x04Copy\x12\x15.volume.FSCopyRequest\x1a\x12.volume.SyncResult\x126\n" + + "\x03Zip\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x128\n" + + "\x05Unzip\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x127\n" + + "\x04Gzip\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x129\n" + + "\x06Gunzip\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x126\n" + + "\x03Tar\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x128\n" + + "\x05Untar\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x126\n" + + "\x03Tgz\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponse\x128\n" + + "\x05Untgz\x12\x16.volume.ArchiveRequest\x1a\x17.volume.ArchiveResponseB%Z#github.com/yaoapp/yao/tai/volume/pbb\x06proto3" var ( - file_volume_pb_volume_proto_rawDescOnce sync.Once - file_volume_pb_volume_proto_rawDescData []byte + file_tai_volume_pb_volume_proto_rawDescOnce sync.Once + file_tai_volume_pb_volume_proto_rawDescData []byte ) -func file_volume_pb_volume_proto_rawDescGZIP() []byte { - file_volume_pb_volume_proto_rawDescOnce.Do(func() { - file_volume_pb_volume_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_volume_pb_volume_proto_rawDesc), len(file_volume_pb_volume_proto_rawDesc))) +func file_tai_volume_pb_volume_proto_rawDescGZIP() []byte { + file_tai_volume_pb_volume_proto_rawDescOnce.Do(func() { + file_tai_volume_pb_volume_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_tai_volume_pb_volume_proto_rawDesc), len(file_tai_volume_pb_volume_proto_rawDesc))) }) - return file_volume_pb_volume_proto_rawDescData + return file_tai_volume_pb_volume_proto_rawDescData } -var file_volume_pb_volume_proto_enumTypes = make([]protoimpl.EnumInfo, 1) -var file_volume_pb_volume_proto_msgTypes = make([]protoimpl.MessageInfo, 15) -var file_volume_pb_volume_proto_goTypes = []any{ +var file_tai_volume_pb_volume_proto_enumTypes = make([]protoimpl.EnumInfo, 1) +var file_tai_volume_pb_volume_proto_msgTypes = make([]protoimpl.MessageInfo, 18) +var file_tai_volume_pb_volume_proto_goTypes = []any{ (FileChunk_ChunkType)(0), // 0: volume.FileChunk.ChunkType (*FileInfo)(nil), // 1: volume.FileInfo (*SyncManifest)(nil), // 2: volume.SyncManifest @@ -1153,8 +1386,11 @@ var file_volume_pb_volume_proto_goTypes = []any{ (*FSListResponse)(nil), // 13: volume.FSListResponse (*FSRemoveRequest)(nil), // 14: volume.FSRemoveRequest (*FSRenameRequest)(nil), // 15: volume.FSRenameRequest + (*FSCopyRequest)(nil), // 16: volume.FSCopyRequest + (*ArchiveRequest)(nil), // 17: volume.ArchiveRequest + (*ArchiveResponse)(nil), // 18: volume.ArchiveResponse } -var file_volume_pb_volume_proto_depIdxs = []int32{ +var file_tai_volume_pb_volume_proto_depIdxs = []int32{ 1, // 0: volume.SyncManifest.files:type_name -> volume.FileInfo 2, // 1: volume.SyncMessage.manifest:type_name -> volume.SyncManifest 4, // 2: volume.SyncMessage.diff:type_name -> volume.SyncDiff @@ -1171,28 +1407,46 @@ var file_volume_pb_volume_proto_depIdxs = []int32{ 14, // 13: volume.Volume.Remove:input_type -> volume.FSRemoveRequest 15, // 14: volume.Volume.Rename:input_type -> volume.FSRenameRequest 7, // 15: volume.Volume.MkdirAll:input_type -> volume.FSRequest - 3, // 16: volume.Volume.SyncPush:output_type -> volume.SyncMessage - 3, // 17: volume.Volume.SyncPull:output_type -> volume.SyncMessage - 10, // 18: volume.Volume.ReadFile:output_type -> volume.FSDataChunk - 12, // 19: volume.Volume.WriteFile:output_type -> volume.FSWriteResponse - 1, // 20: volume.Volume.Stat:output_type -> volume.FileInfo - 13, // 21: volume.Volume.ListDir:output_type -> volume.FSListResponse - 8, // 22: volume.Volume.Remove:output_type -> volume.FSOpResponse - 8, // 23: volume.Volume.Rename:output_type -> volume.FSOpResponse - 8, // 24: volume.Volume.MkdirAll:output_type -> volume.FSOpResponse - 16, // [16:25] is the sub-list for method output_type - 7, // [7:16] is the sub-list for method input_type + 16, // 16: volume.Volume.Copy:input_type -> volume.FSCopyRequest + 17, // 17: volume.Volume.Zip:input_type -> volume.ArchiveRequest + 17, // 18: volume.Volume.Unzip:input_type -> volume.ArchiveRequest + 17, // 19: volume.Volume.Gzip:input_type -> volume.ArchiveRequest + 17, // 20: volume.Volume.Gunzip:input_type -> volume.ArchiveRequest + 17, // 21: volume.Volume.Tar:input_type -> volume.ArchiveRequest + 17, // 22: volume.Volume.Untar:input_type -> volume.ArchiveRequest + 17, // 23: volume.Volume.Tgz:input_type -> volume.ArchiveRequest + 17, // 24: volume.Volume.Untgz:input_type -> volume.ArchiveRequest + 3, // 25: volume.Volume.SyncPush:output_type -> volume.SyncMessage + 3, // 26: volume.Volume.SyncPull:output_type -> volume.SyncMessage + 10, // 27: volume.Volume.ReadFile:output_type -> volume.FSDataChunk + 12, // 28: volume.Volume.WriteFile:output_type -> volume.FSWriteResponse + 1, // 29: volume.Volume.Stat:output_type -> volume.FileInfo + 13, // 30: volume.Volume.ListDir:output_type -> volume.FSListResponse + 8, // 31: volume.Volume.Remove:output_type -> volume.FSOpResponse + 8, // 32: volume.Volume.Rename:output_type -> volume.FSOpResponse + 8, // 33: volume.Volume.MkdirAll:output_type -> volume.FSOpResponse + 6, // 34: volume.Volume.Copy:output_type -> volume.SyncResult + 18, // 35: volume.Volume.Zip:output_type -> volume.ArchiveResponse + 18, // 36: volume.Volume.Unzip:output_type -> volume.ArchiveResponse + 18, // 37: volume.Volume.Gzip:output_type -> volume.ArchiveResponse + 18, // 38: volume.Volume.Gunzip:output_type -> volume.ArchiveResponse + 18, // 39: volume.Volume.Tar:output_type -> volume.ArchiveResponse + 18, // 40: volume.Volume.Untar:output_type -> volume.ArchiveResponse + 18, // 41: volume.Volume.Tgz:output_type -> volume.ArchiveResponse + 18, // 42: volume.Volume.Untgz:output_type -> volume.ArchiveResponse + 25, // [25:43] is the sub-list for method output_type + 7, // [7:25] is the sub-list for method input_type 7, // [7:7] is the sub-list for extension type_name 7, // [7:7] is the sub-list for extension extendee 0, // [0:7] is the sub-list for field type_name } -func init() { file_volume_pb_volume_proto_init() } -func file_volume_pb_volume_proto_init() { - if File_volume_pb_volume_proto != nil { +func init() { file_tai_volume_pb_volume_proto_init() } +func file_tai_volume_pb_volume_proto_init() { + if File_tai_volume_pb_volume_proto != nil { return } - file_volume_pb_volume_proto_msgTypes[2].OneofWrappers = []any{ + file_tai_volume_pb_volume_proto_msgTypes[2].OneofWrappers = []any{ (*SyncMessage_Manifest)(nil), (*SyncMessage_Diff)(nil), (*SyncMessage_Chunk)(nil), @@ -1202,18 +1456,18 @@ func file_volume_pb_volume_proto_init() { out := protoimpl.TypeBuilder{ File: protoimpl.DescBuilder{ GoPackagePath: reflect.TypeOf(x{}).PkgPath(), - RawDescriptor: unsafe.Slice(unsafe.StringData(file_volume_pb_volume_proto_rawDesc), len(file_volume_pb_volume_proto_rawDesc)), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_tai_volume_pb_volume_proto_rawDesc), len(file_tai_volume_pb_volume_proto_rawDesc)), NumEnums: 1, - NumMessages: 15, + NumMessages: 18, NumExtensions: 0, NumServices: 1, }, - GoTypes: file_volume_pb_volume_proto_goTypes, - DependencyIndexes: file_volume_pb_volume_proto_depIdxs, - EnumInfos: file_volume_pb_volume_proto_enumTypes, - MessageInfos: file_volume_pb_volume_proto_msgTypes, + GoTypes: file_tai_volume_pb_volume_proto_goTypes, + DependencyIndexes: file_tai_volume_pb_volume_proto_depIdxs, + EnumInfos: file_tai_volume_pb_volume_proto_enumTypes, + MessageInfos: file_tai_volume_pb_volume_proto_msgTypes, }.Build() - File_volume_pb_volume_proto = out.File - file_volume_pb_volume_proto_goTypes = nil - file_volume_pb_volume_proto_depIdxs = nil + File_tai_volume_pb_volume_proto = out.File + file_tai_volume_pb_volume_proto_goTypes = nil + file_tai_volume_pb_volume_proto_depIdxs = nil } diff --git a/tai/volume/pb/volume.proto b/tai/volume/pb/volume.proto index dc39e07a..5fa38279 100644 --- a/tai/volume/pb/volume.proto +++ b/tai/volume/pb/volume.proto @@ -1,8 +1,9 @@ syntax = "proto3"; package volume; -option go_package = "github.com/yaoapp/tai/volume/pb"; +option go_package = "github.com/yaoapp/yao/tai/volume/pb"; -// Volume provides bulk file synchronization and real-time filesystem I/O. +// Volume provides bulk file synchronization, real-time filesystem I/O, +// and archive/compression operations. // Shares gRPC port :19100 with Yao Gateway. service Volume { @@ -29,6 +30,19 @@ service Volume { rpc Remove(FSRemoveRequest) returns (FSOpResponse); rpc Rename(FSRenameRequest) returns (FSOpResponse); rpc MkdirAll(FSRequest) returns (FSOpResponse); + // Copy: copy src to dst within the same workspace (server-side when remote). + rpc Copy(FSCopyRequest) returns (SyncResult); + + // --- Archive / Compression --- + + rpc Zip(ArchiveRequest) returns (ArchiveResponse); + rpc Unzip(ArchiveRequest) returns (ArchiveResponse); + rpc Gzip(ArchiveRequest) returns (ArchiveResponse); + rpc Gunzip(ArchiveRequest) returns (ArchiveResponse); + rpc Tar(ArchiveRequest) returns (ArchiveResponse); + rpc Untar(ArchiveRequest) returns (ArchiveResponse); + rpc Tgz(ArchiveRequest) returns (ArchiveResponse); + rpc Untgz(ArchiveRequest) returns (ArchiveResponse); } // --- File Metadata --- @@ -47,6 +61,7 @@ message SyncManifest { string session_id = 1; repeated FileInfo files = 2; bool force_full = 3; // skip snapshot cache, diff against actual disk + string remote_path = 4; // sub-path within workspace root; empty = root } message SyncMessage { @@ -136,3 +151,25 @@ message FSRenameRequest { string old_path = 2; string new_path = 3; } + +message FSCopyRequest { + string session_id = 1; + string src_path = 2; + string dst_path = 3; + repeated string excludes = 4; // glob patterns + bool force = 5; // overwrite even if mtime/size match +} + +// --- Archive / Compression Messages --- + +message ArchiveRequest { + string session_id = 1; + string src_path = 2; // relative to workspace root + string dst_path = 3; // relative to workspace root + repeated string excludes = 4; // glob patterns (pack ops only) +} + +message ArchiveResponse { + int64 size_bytes = 1; // output file size (pack) or total extracted size (unpack) + int32 files_count = 2; // number of files processed +} diff --git a/tai/volume/pb/volume_grpc.pb.go b/tai/volume/pb/volume_grpc.pb.go index 7924e483..a5370fe3 100644 --- a/tai/volume/pb/volume_grpc.pb.go +++ b/tai/volume/pb/volume_grpc.pb.go @@ -2,7 +2,7 @@ // versions: // - protoc-gen-go-grpc v1.6.1 // - protoc v4.25.0 -// source: volume/pb/volume.proto +// source: tai/volume/pb/volume.proto package pb @@ -28,13 +28,23 @@ const ( Volume_Remove_FullMethodName = "/volume.Volume/Remove" Volume_Rename_FullMethodName = "/volume.Volume/Rename" Volume_MkdirAll_FullMethodName = "/volume.Volume/MkdirAll" + Volume_Copy_FullMethodName = "/volume.Volume/Copy" + Volume_Zip_FullMethodName = "/volume.Volume/Zip" + Volume_Unzip_FullMethodName = "/volume.Volume/Unzip" + Volume_Gzip_FullMethodName = "/volume.Volume/Gzip" + Volume_Gunzip_FullMethodName = "/volume.Volume/Gunzip" + Volume_Tar_FullMethodName = "/volume.Volume/Tar" + Volume_Untar_FullMethodName = "/volume.Volume/Untar" + Volume_Tgz_FullMethodName = "/volume.Volume/Tgz" + Volume_Untgz_FullMethodName = "/volume.Volume/Untgz" ) // VolumeClient is the client API for Volume service. // // For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. // -// Volume provides bulk file synchronization and real-time filesystem I/O. +// Volume provides bulk file synchronization, real-time filesystem I/O, +// and archive/compression operations. // Shares gRPC port :19100 with Yao Gateway. type VolumeClient interface { // SyncPush: Yao sends code to Tai (before container start). @@ -54,6 +64,16 @@ type VolumeClient interface { Remove(ctx context.Context, in *FSRemoveRequest, opts ...grpc.CallOption) (*FSOpResponse, error) Rename(ctx context.Context, in *FSRenameRequest, opts ...grpc.CallOption) (*FSOpResponse, error) MkdirAll(ctx context.Context, in *FSRequest, opts ...grpc.CallOption) (*FSOpResponse, error) + // Copy: copy src to dst within the same workspace (server-side when remote). + Copy(ctx context.Context, in *FSCopyRequest, opts ...grpc.CallOption) (*SyncResult, error) + Zip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Unzip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Gzip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Gunzip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Tar(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Untar(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Tgz(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) + Untgz(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) } type volumeClient struct { @@ -178,11 +198,102 @@ func (c *volumeClient) MkdirAll(ctx context.Context, in *FSRequest, opts ...grpc return out, nil } +func (c *volumeClient) Copy(ctx context.Context, in *FSCopyRequest, opts ...grpc.CallOption) (*SyncResult, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(SyncResult) + err := c.cc.Invoke(ctx, Volume_Copy_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Zip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Zip_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Unzip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Unzip_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Gzip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Gzip_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Gunzip(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Gunzip_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Tar(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Tar_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Untar(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Untar_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Tgz(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Tgz_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *volumeClient) Untgz(ctx context.Context, in *ArchiveRequest, opts ...grpc.CallOption) (*ArchiveResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ArchiveResponse) + err := c.cc.Invoke(ctx, Volume_Untgz_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + // VolumeServer is the server API for Volume service. // All implementations must embed UnimplementedVolumeServer // for forward compatibility. // -// Volume provides bulk file synchronization and real-time filesystem I/O. +// Volume provides bulk file synchronization, real-time filesystem I/O, +// and archive/compression operations. // Shares gRPC port :19100 with Yao Gateway. type VolumeServer interface { // SyncPush: Yao sends code to Tai (before container start). @@ -202,6 +313,16 @@ type VolumeServer interface { Remove(context.Context, *FSRemoveRequest) (*FSOpResponse, error) Rename(context.Context, *FSRenameRequest) (*FSOpResponse, error) MkdirAll(context.Context, *FSRequest) (*FSOpResponse, error) + // Copy: copy src to dst within the same workspace (server-side when remote). + Copy(context.Context, *FSCopyRequest) (*SyncResult, error) + Zip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Unzip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Gzip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Gunzip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Tar(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Untar(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Tgz(context.Context, *ArchiveRequest) (*ArchiveResponse, error) + Untgz(context.Context, *ArchiveRequest) (*ArchiveResponse, error) mustEmbedUnimplementedVolumeServer() } @@ -239,6 +360,33 @@ func (UnimplementedVolumeServer) Rename(context.Context, *FSRenameRequest) (*FSO func (UnimplementedVolumeServer) MkdirAll(context.Context, *FSRequest) (*FSOpResponse, error) { return nil, status.Error(codes.Unimplemented, "method MkdirAll not implemented") } +func (UnimplementedVolumeServer) Copy(context.Context, *FSCopyRequest) (*SyncResult, error) { + return nil, status.Error(codes.Unimplemented, "method Copy not implemented") +} +func (UnimplementedVolumeServer) Zip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Zip not implemented") +} +func (UnimplementedVolumeServer) Unzip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Unzip not implemented") +} +func (UnimplementedVolumeServer) Gzip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Gzip not implemented") +} +func (UnimplementedVolumeServer) Gunzip(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Gunzip not implemented") +} +func (UnimplementedVolumeServer) Tar(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Tar not implemented") +} +func (UnimplementedVolumeServer) Untar(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Untar not implemented") +} +func (UnimplementedVolumeServer) Tgz(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Tgz not implemented") +} +func (UnimplementedVolumeServer) Untgz(context.Context, *ArchiveRequest) (*ArchiveResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Untgz not implemented") +} func (UnimplementedVolumeServer) mustEmbedUnimplementedVolumeServer() {} func (UnimplementedVolumeServer) testEmbeddedByValue() {} @@ -386,6 +534,168 @@ func _Volume_MkdirAll_Handler(srv interface{}, ctx context.Context, dec func(int return interceptor(ctx, in, info, handler) } +func _Volume_Copy_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(FSCopyRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Copy(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Copy_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Copy(ctx, req.(*FSCopyRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Zip_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Zip(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Zip_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Zip(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Unzip_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Unzip(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Unzip_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Unzip(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Gzip_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Gzip(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Gzip_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Gzip(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Gunzip_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Gunzip(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Gunzip_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Gunzip(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Tar_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Tar(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Tar_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Tar(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Untar_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Untar(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Untar_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Untar(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Tgz_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Tgz(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Tgz_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Tgz(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Volume_Untgz_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ArchiveRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Untgz(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Untgz_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Untgz(ctx, req.(*ArchiveRequest)) + } + return interceptor(ctx, in, info, handler) +} + // Volume_ServiceDesc is the grpc.ServiceDesc for Volume service. // It's only intended for direct use with grpc.RegisterService, // and not to be introspected or modified (even as a copy) @@ -413,6 +723,42 @@ var Volume_ServiceDesc = grpc.ServiceDesc{ MethodName: "MkdirAll", Handler: _Volume_MkdirAll_Handler, }, + { + MethodName: "Copy", + Handler: _Volume_Copy_Handler, + }, + { + MethodName: "Zip", + Handler: _Volume_Zip_Handler, + }, + { + MethodName: "Unzip", + Handler: _Volume_Unzip_Handler, + }, + { + MethodName: "Gzip", + Handler: _Volume_Gzip_Handler, + }, + { + MethodName: "Gunzip", + Handler: _Volume_Gunzip_Handler, + }, + { + MethodName: "Tar", + Handler: _Volume_Tar_Handler, + }, + { + MethodName: "Untar", + Handler: _Volume_Untar_Handler, + }, + { + MethodName: "Tgz", + Handler: _Volume_Tgz_Handler, + }, + { + MethodName: "Untgz", + Handler: _Volume_Untgz_Handler, + }, }, Streams: []grpc.StreamDesc{ { @@ -437,5 +783,5 @@ var Volume_ServiceDesc = grpc.ServiceDesc{ ClientStreams: true, }, }, - Metadata: "volume/pb/volume.proto", + Metadata: "tai/volume/pb/volume.proto", } diff --git a/tai/volume/remote.go b/tai/volume/remote.go index 315db6a2..e7019e66 100644 --- a/tai/volume/remote.go +++ b/tai/volume/remote.go @@ -170,7 +170,7 @@ func (r *remoteStorage) MkdirAll(ctx context.Context, sessionID, path string) er // SyncPush sends local files to Tai using the manifest-first bidi streaming protocol. func (r *remoteStorage) SyncPush(ctx context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) { start := time.Now() - cfg := applySyncOpts(opts) + cfg := ApplySyncOpts(opts) // Scan local directory var manifest []*pb.FileInfo @@ -183,7 +183,7 @@ func (r *remoteStorage) SyncPush(ctx context.Context, sessionID, localDir string return nil } rel = filepath.ToSlash(rel) - if isExcluded(rel, d.IsDir(), cfg.excludes) { + if isExcluded(rel, d.IsDir(), cfg.Excludes) { if d.IsDir() { return filepath.SkipDir } @@ -215,9 +215,10 @@ func (r *remoteStorage) SyncPush(ctx context.Context, sessionID, localDir string if err := stream.Send(&pb.SyncMessage{ Payload: &pb.SyncMessage_Manifest{ Manifest: &pb.SyncManifest{ - SessionId: sessionID, - Files: manifest, - ForceFull: cfg.forceFull, + SessionId: sessionID, + Files: manifest, + ForceFull: cfg.ForceFull, + RemotePath: cfg.RemotePath, }, }, }); err != nil { @@ -309,7 +310,7 @@ func (r *remoteStorage) SyncPush(ctx context.Context, sessionID, localDir string // SyncPull receives changed files from Tai. func (r *remoteStorage) SyncPull(ctx context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) { start := time.Now() - cfg := applySyncOpts(opts) + cfg := ApplySyncOpts(opts) // Build local manifest var manifest []*pb.FileInfo @@ -322,7 +323,7 @@ func (r *remoteStorage) SyncPull(ctx context.Context, sessionID, localDir string return nil } rel = filepath.ToSlash(rel) - if isExcluded(rel, d.IsDir(), cfg.excludes) { + if isExcluded(rel, d.IsDir(), cfg.Excludes) { if d.IsDir() { return filepath.SkipDir } @@ -343,9 +344,10 @@ func (r *remoteStorage) SyncPull(ctx context.Context, sessionID, localDir string }) stream, err := r.client.SyncPull(ctx, &pb.SyncManifest{ - SessionId: sessionID, - Files: manifest, - ForceFull: cfg.forceFull, + SessionId: sessionID, + Files: manifest, + ForceFull: cfg.ForceFull, + RemotePath: cfg.RemotePath, }) if err != nil { return nil, err @@ -432,6 +434,107 @@ func (r *remoteStorage) SyncPull(ctx context.Context, sessionID, localDir string }, nil } +func (r *remoteStorage) Zip(ctx context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) { + resp, err := r.client.Zip(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, Excludes: excludes, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Unzip(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + resp, err := r.client.Unzip(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Gzip(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + resp, err := r.client.Gzip(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Gunzip(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + resp, err := r.client.Gunzip(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Tar(ctx context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) { + resp, err := r.client.Tar(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, Excludes: excludes, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Untar(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + resp, err := r.client.Untar(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Tgz(ctx context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) { + resp, err := r.client.Tgz(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, Excludes: excludes, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Untgz(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) { + resp, err := r.client.Untgz(ctx, &pb.ArchiveRequest{ + SessionId: sessionID, SrcPath: src, DstPath: dst, + }) + if err != nil { + return nil, err + } + return &ArchiveResult{SizeBytes: resp.SizeBytes, FilesCount: int(resp.FilesCount)}, nil +} + +func (r *remoteStorage) Copy(ctx context.Context, sessionID, src, dst string, opts ...SyncOption) (*SyncResult, error) { + start := time.Now() + cfg := ApplySyncOpts(opts) + + resp, err := r.client.Copy(ctx, &pb.FSCopyRequest{ + SessionId: sessionID, + SrcPath: src, + DstPath: dst, + Excludes: cfg.Excludes, + Force: cfg.ForceFull, + }) + if err != nil { + return nil, err + } + return &SyncResult{ + FilesSynced: int(resp.FilesSynced), + BytesTransferred: resp.BytesTransferred, + Duration: time.Since(start), + }, nil +} + func (r *remoteStorage) Close() error { return nil } diff --git a/tai/volume/volume.go b/tai/volume/volume.go index 31ac3967..e8498c33 100644 --- a/tai/volume/volume.go +++ b/tai/volume/volume.go @@ -7,7 +7,7 @@ import ( "time" ) -// Volume provides filesystem IO and directory synchronization. +// Volume provides filesystem IO, directory synchronization, and archive operations. // Remote connects to Tai gRPC :19100; Local operates directly on disk. type Volume interface { ReadFile(ctx context.Context, sessionID, path string) ([]byte, os.FileMode, error) @@ -20,6 +20,16 @@ type Volume interface { SyncPush(ctx context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) SyncPull(ctx context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) + Copy(ctx context.Context, sessionID, src, dst string, opts ...SyncOption) (*SyncResult, error) + + Zip(ctx context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) + Unzip(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) + Gzip(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) + Gunzip(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) + Tar(ctx context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) + Untar(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) + Tgz(ctx context.Context, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) + Untgz(ctx context.Context, sessionID, src, dst string) (*ArchiveResult, error) Close() error } @@ -40,26 +50,40 @@ type SyncResult struct { Duration time.Duration } -// SyncOption configures sync behavior. -type SyncOption func(*syncConfig) +// ArchiveResult summarizes an archive/compression operation. +type ArchiveResult struct { + SizeBytes int64 + FilesCount int +} -type syncConfig struct { - forceFull bool - excludes []string +// SyncOption configures sync behavior. +type SyncOption func(*SyncConfig) + +// SyncConfig holds resolved sync options. +type SyncConfig struct { + ForceFull bool + Excludes []string + RemotePath string } // WithForceFull skips snapshot caches and diffs against actual disk. func WithForceFull() SyncOption { - return func(c *syncConfig) { c.forceFull = true } + return func(c *SyncConfig) { c.ForceFull = true } } // WithExcludes adds glob patterns to exclude from sync. func WithExcludes(patterns ...string) SyncOption { - return func(c *syncConfig) { c.excludes = append(c.excludes, patterns...) } + return func(c *SyncConfig) { c.Excludes = append(c.Excludes, patterns...) } } -func applySyncOpts(opts []SyncOption) syncConfig { - var cfg syncConfig +// WithRemotePath sets a sub-path within the workspace root for sync operations. +func WithRemotePath(path string) SyncOption { + return func(c *SyncConfig) { c.RemotePath = path } +} + +// ApplySyncOpts resolves a slice of SyncOption into a SyncConfig. +func ApplySyncOpts(opts []SyncOption) SyncConfig { + var cfg SyncConfig for _, o := range opts { o(&cfg) } diff --git a/tai/volume/volume_test.go b/tai/volume/volume_test.go index 41339325..eae954a8 100644 --- a/tai/volume/volume_test.go +++ b/tai/volume/volume_test.go @@ -383,6 +383,183 @@ func TestRemoteVolume(t *testing.T) { } }) + t.Run("SyncPush with RemotePath", func(t *testing.T) { + srcDir := t.TempDir() + _ = os.WriteFile(filepath.Join(srcDir, "mod.go"), []byte("module test"), 0o644) + result, err := vol.SyncPush(ctx, "rp-test", srcDir, WithForceFull(), WithRemotePath("packages/api")) + if err != nil { + t.Fatalf("SyncPush: %v", err) + } + if result.FilesSynced < 1 { + t.Errorf("synced = %d", result.FilesSynced) + } + _ = vol.Remove(ctx, "rp-test", ".", true) + }) + + t.Run("SyncPull with RemotePath", func(t *testing.T) { + rpSid := "rp-pull-test" + _ = vol.MkdirAll(ctx, rpSid, "sub/deep") + _ = vol.WriteFile(ctx, rpSid, "sub/deep/f.txt", []byte("deep"), 0o644) + _ = vol.WriteFile(ctx, rpSid, "root.txt", []byte("root"), 0o644) + + dstDir := t.TempDir() + result, err := vol.SyncPull(ctx, rpSid, dstDir, WithForceFull(), WithRemotePath("sub/deep")) + if err != nil { + t.Fatalf("SyncPull: %v", err) + } + if result.FilesSynced < 1 { + t.Errorf("synced = %d", result.FilesSynced) + } + data, err := os.ReadFile(filepath.Join(dstDir, "f.txt")) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "deep" { + t.Errorf("content = %q", data) + } + _ = vol.Remove(ctx, rpSid, ".", true) + }) + + t.Run("Zip and Unzip", func(t *testing.T) { + arcSid := "arc-zip-test" + _ = vol.MkdirAll(ctx, arcSid, "src") + _ = vol.WriteFile(ctx, arcSid, "src/a.txt", []byte("zip a"), 0o644) + _ = vol.WriteFile(ctx, arcSid, "src/b.txt", []byte("zip b"), 0o644) + + zr, err := vol.Zip(ctx, arcSid, "src", "out.zip", nil) + if err != nil { + t.Fatalf("Zip: %v", err) + } + if zr.FilesCount != 2 { + t.Errorf("zip files = %d, want 2", zr.FilesCount) + } + + ur, err := vol.Unzip(ctx, arcSid, "out.zip", "extracted") + if err != nil { + t.Fatalf("Unzip: %v", err) + } + if ur.FilesCount != 2 { + t.Errorf("unzip files = %d, want 2", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, arcSid, "extracted/a.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "zip a" { + t.Errorf("content = %q", data) + } + _ = vol.Remove(ctx, arcSid, ".", true) + }) + + t.Run("Zip with excludes", func(t *testing.T) { + arcSid := "arc-zip-excl" + _ = vol.MkdirAll(ctx, arcSid, "src") + _ = vol.WriteFile(ctx, arcSid, "src/keep.txt", []byte("keep"), 0o644) + _ = vol.WriteFile(ctx, arcSid, "src/skip.log", []byte("skip"), 0o644) + + zr, err := vol.Zip(ctx, arcSid, "src", "filtered.zip", []string{"*.log"}) + if err != nil { + t.Fatalf("Zip: %v", err) + } + if zr.FilesCount != 1 { + t.Errorf("zip files = %d, want 1", zr.FilesCount) + } + _ = vol.Remove(ctx, arcSid, ".", true) + }) + + t.Run("Gzip and Gunzip", func(t *testing.T) { + arcSid := "arc-gzip-test" + _ = vol.WriteFile(ctx, arcSid, "data.txt", []byte("gzip remote"), 0o644) + + gr, err := vol.Gzip(ctx, arcSid, "data.txt", "data.txt.gz") + if err != nil { + t.Fatalf("Gzip: %v", err) + } + if gr.FilesCount != 1 { + t.Errorf("gzip files = %d", gr.FilesCount) + } + + ur, err := vol.Gunzip(ctx, arcSid, "data.txt.gz", "restored.txt") + if err != nil { + t.Fatalf("Gunzip: %v", err) + } + if ur.FilesCount != 1 { + t.Errorf("gunzip files = %d", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, arcSid, "restored.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "gzip remote" { + t.Errorf("content = %q", data) + } + _ = vol.Remove(ctx, arcSid, ".", true) + }) + + t.Run("Tar and Untar", func(t *testing.T) { + arcSid := "arc-tar-test" + _ = vol.MkdirAll(ctx, arcSid, "src") + _ = vol.WriteFile(ctx, arcSid, "src/x.txt", []byte("tar remote"), 0o644) + + tr, err := vol.Tar(ctx, arcSid, "src", "out.tar", nil) + if err != nil { + t.Fatalf("Tar: %v", err) + } + if tr.FilesCount != 1 { + t.Errorf("tar files = %d", tr.FilesCount) + } + + ur, err := vol.Untar(ctx, arcSid, "out.tar", "extracted") + if err != nil { + t.Fatalf("Untar: %v", err) + } + if ur.FilesCount != 1 { + t.Errorf("untar files = %d", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, arcSid, "extracted/x.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "tar remote" { + t.Errorf("content = %q", data) + } + _ = vol.Remove(ctx, arcSid, ".", true) + }) + + t.Run("Tgz and Untgz", func(t *testing.T) { + arcSid := "arc-tgz-test" + _ = vol.MkdirAll(ctx, arcSid, "src") + _ = vol.WriteFile(ctx, arcSid, "src/f.txt", []byte("tgz remote"), 0o644) + + tr, err := vol.Tgz(ctx, arcSid, "src", "out.tgz", nil) + if err != nil { + t.Fatalf("Tgz: %v", err) + } + if tr.FilesCount != 1 { + t.Errorf("tgz files = %d", tr.FilesCount) + } + + ur, err := vol.Untgz(ctx, arcSid, "out.tgz", "extracted") + if err != nil { + t.Fatalf("Untgz: %v", err) + } + if ur.FilesCount != 1 { + t.Errorf("untgz files = %d", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, arcSid, "extracted/f.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "tgz remote" { + t.Errorf("content = %q", data) + } + _ = vol.Remove(ctx, arcSid, ".", true) + }) + // Cleanup _ = vol.Remove(ctx, sid, ".", true) } @@ -781,6 +958,290 @@ func TestLocalSyncPullForceFull(t *testing.T) { } } +func TestLocalSyncPushWithRemotePath(t *testing.T) { + dataDir := t.TempDir() + vol := NewLocal(dataDir) + defer vol.Close() + ctx := context.Background() + sid := "remote-path-push" + + srcDir := t.TempDir() + _ = os.WriteFile(filepath.Join(srcDir, "app.js"), []byte("console.log('hi')"), 0o644) + _ = os.MkdirAll(filepath.Join(srcDir, "lib"), 0o755) + _ = os.WriteFile(filepath.Join(srcDir, "lib", "util.js"), []byte("export {}"), 0o644) + + result, err := vol.SyncPush(ctx, sid, srcDir, WithForceFull(), WithRemotePath("packages/frontend")) + if err != nil { + t.Fatalf("SyncPush: %v", err) + } + if result.FilesSynced != 2 { + t.Errorf("synced = %d, want 2", result.FilesSynced) + } + + data, err := os.ReadFile(filepath.Join(dataDir, sid, "packages", "frontend", "app.js")) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "console.log('hi')" { + t.Errorf("content = %q", data) + } + + nested, err := os.ReadFile(filepath.Join(dataDir, sid, "packages", "frontend", "lib", "util.js")) + if err != nil { + t.Fatalf("ReadFile nested: %v", err) + } + if string(nested) != "export {}" { + t.Errorf("nested content = %q", nested) + } +} + +func TestLocalSyncPullWithRemotePath(t *testing.T) { + dataDir := t.TempDir() + vol := NewLocal(dataDir) + defer vol.Close() + ctx := context.Background() + sid := "remote-path-pull" + + sessionDir := filepath.Join(dataDir, sid, "packages", "backend") + _ = os.MkdirAll(sessionDir, 0o755) + _ = os.WriteFile(filepath.Join(sessionDir, "main.go"), []byte("package main"), 0o644) + + dstDir := t.TempDir() + result, err := vol.SyncPull(ctx, sid, dstDir, WithForceFull(), WithRemotePath("packages/backend")) + if err != nil { + t.Fatalf("SyncPull: %v", err) + } + if result.FilesSynced != 1 { + t.Errorf("synced = %d, want 1", result.FilesSynced) + } + + data, err := os.ReadFile(filepath.Join(dstDir, "main.go")) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "package main" { + t.Errorf("content = %q", data) + } +} + +func TestLocalZipUnzip(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "zip-test" + + _ = vol.MkdirAll(ctx, sid, "src") + _ = vol.WriteFile(ctx, sid, "src/a.txt", []byte("aaa"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/b.txt", []byte("bbb"), 0o644) + + zr, err := vol.Zip(ctx, sid, "src", "out.zip", nil) + if err != nil { + t.Fatalf("Zip: %v", err) + } + if zr.FilesCount != 2 { + t.Errorf("zip files = %d, want 2", zr.FilesCount) + } + if zr.SizeBytes <= 0 { + t.Error("zip size should be > 0") + } + + ur, err := vol.Unzip(ctx, sid, "out.zip", "extracted") + if err != nil { + t.Fatalf("Unzip: %v", err) + } + if ur.FilesCount != 2 { + t.Errorf("unzip files = %d, want 2", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, sid, "extracted/a.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "aaa" { + t.Errorf("content = %q", data) + } +} + +func TestLocalZipExcludes(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "zip-excl" + + _ = vol.MkdirAll(ctx, sid, "src") + _ = vol.WriteFile(ctx, sid, "src/keep.txt", []byte("keep"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/skip.log", []byte("skip"), 0o644) + + zr, err := vol.Zip(ctx, sid, "src", "filtered.zip", []string{"*.log"}) + if err != nil { + t.Fatalf("Zip: %v", err) + } + if zr.FilesCount != 1 { + t.Errorf("zip files = %d, want 1", zr.FilesCount) + } + + ur, err := vol.Unzip(ctx, sid, "filtered.zip", "out") + if err != nil { + t.Fatalf("Unzip: %v", err) + } + if ur.FilesCount != 1 { + t.Errorf("unzip files = %d, want 1", ur.FilesCount) + } + + _, err = vol.Stat(ctx, sid, "out/keep.txt") + if err != nil { + t.Error("keep.txt should exist") + } + _, err = vol.Stat(ctx, sid, "out/skip.log") + if !os.IsNotExist(err) { + t.Error("skip.log should not exist") + } +} + +func TestLocalGzipGunzip(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "gzip-test" + + _ = vol.WriteFile(ctx, sid, "data.txt", []byte("gzip test content"), 0o644) + + gr, err := vol.Gzip(ctx, sid, "data.txt", "data.txt.gz") + if err != nil { + t.Fatalf("Gzip: %v", err) + } + if gr.FilesCount != 1 { + t.Errorf("gzip files = %d", gr.FilesCount) + } + if gr.SizeBytes <= 0 { + t.Error("gzip size should be > 0") + } + + ur, err := vol.Gunzip(ctx, sid, "data.txt.gz", "restored.txt") + if err != nil { + t.Fatalf("Gunzip: %v", err) + } + if ur.FilesCount != 1 { + t.Errorf("gunzip files = %d", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, sid, "restored.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "gzip test content" { + t.Errorf("content = %q", data) + } +} + +func TestLocalGzipRejectsDir(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "gzip-dir" + + _ = vol.MkdirAll(ctx, sid, "subdir") + _, err := vol.Gzip(ctx, sid, "subdir", "subdir.gz") + if err == nil { + t.Error("expected error for gzip on directory") + } +} + +func TestLocalTarUntar(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "tar-test" + + _ = vol.MkdirAll(ctx, sid, "src") + _ = vol.WriteFile(ctx, sid, "src/x.txt", []byte("tar x"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/y.txt", []byte("tar y"), 0o644) + + tr, err := vol.Tar(ctx, sid, "src", "out.tar", nil) + if err != nil { + t.Fatalf("Tar: %v", err) + } + if tr.FilesCount != 2 { + t.Errorf("tar files = %d, want 2", tr.FilesCount) + } + + ur, err := vol.Untar(ctx, sid, "out.tar", "extracted") + if err != nil { + t.Fatalf("Untar: %v", err) + } + if ur.FilesCount != 2 { + t.Errorf("untar files = %d, want 2", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, sid, "extracted/x.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "tar x" { + t.Errorf("content = %q", data) + } +} + +func TestLocalTarExcludes(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "tar-excl" + + _ = vol.MkdirAll(ctx, sid, "src") + _ = vol.WriteFile(ctx, sid, "src/keep.txt", []byte("keep"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/skip.log", []byte("skip"), 0o644) + + tr, err := vol.Tar(ctx, sid, "src", "out.tar", []string{"*.log"}) + if err != nil { + t.Fatalf("Tar: %v", err) + } + if tr.FilesCount != 1 { + t.Errorf("tar files = %d, want 1", tr.FilesCount) + } +} + +func TestLocalTgzUntgz(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "tgz-test" + + _ = vol.MkdirAll(ctx, sid, "src") + _ = vol.WriteFile(ctx, sid, "src/f.txt", []byte("tgz content"), 0o644) + + tr, err := vol.Tgz(ctx, sid, "src", "out.tgz", nil) + if err != nil { + t.Fatalf("Tgz: %v", err) + } + if tr.FilesCount != 1 { + t.Errorf("tgz files = %d", tr.FilesCount) + } + + ur, err := vol.Untgz(ctx, sid, "out.tgz", "extracted") + if err != nil { + t.Fatalf("Untgz: %v", err) + } + if ur.FilesCount != 1 { + t.Errorf("untgz files = %d", ur.FilesCount) + } + + data, _, err := vol.ReadFile(ctx, sid, "extracted/f.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "tgz content" { + t.Errorf("content = %q", data) + } +} + func TestLocalSyncPushForceFull(t *testing.T) { dataDir := t.TempDir() vol := NewLocal(dataDir) @@ -825,3 +1286,244 @@ func TestLocalSyncExcludes(t *testing.T) { t.Error("excluded file should not exist") } } + +func TestLocalCopyFile(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "copy-file" + + _ = vol.WriteFile(ctx, sid, "src.txt", []byte("hello copy"), 0o644) + + result, err := vol.Copy(ctx, sid, "src.txt", "dst.txt") + if err != nil { + t.Fatalf("Copy file: %v", err) + } + if result.FilesSynced != 1 { + t.Errorf("synced = %d, want 1", result.FilesSynced) + } + if result.BytesTransferred != 10 { + t.Errorf("bytes = %d, want 10", result.BytesTransferred) + } + + data, _, err := vol.ReadFile(ctx, sid, "dst.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "hello copy" { + t.Errorf("content = %q", data) + } +} + +func TestLocalCopyDir(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "copy-dir" + + _ = vol.MkdirAll(ctx, sid, "src/sub") + _ = vol.WriteFile(ctx, sid, "src/a.txt", []byte("aaa"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/sub/b.txt", []byte("bbb"), 0o644) + + result, err := vol.Copy(ctx, sid, "src", "dst") + if err != nil { + t.Fatalf("Copy dir: %v", err) + } + if result.FilesSynced != 2 { + t.Errorf("synced = %d, want 2", result.FilesSynced) + } + + data, _, err := vol.ReadFile(ctx, sid, "dst/a.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "aaa" { + t.Errorf("content = %q", data) + } + + data, _, err = vol.ReadFile(ctx, sid, "dst/sub/b.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "bbb" { + t.Errorf("content = %q", data) + } +} + +func TestLocalCopyExcludes(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "copy-excl" + + _ = vol.MkdirAll(ctx, sid, "src") + _ = vol.WriteFile(ctx, sid, "src/keep.txt", []byte("keep"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/skip.log", []byte("skip"), 0o644) + + result, err := vol.Copy(ctx, sid, "src", "dst", WithExcludes("*.log")) + if err != nil { + t.Fatalf("Copy: %v", err) + } + if result.FilesSynced != 1 { + t.Errorf("synced = %d, want 1", result.FilesSynced) + } + + _, err = vol.Stat(ctx, sid, "dst/keep.txt") + if err != nil { + t.Error("keep.txt should exist") + } + _, err = vol.Stat(ctx, sid, "dst/skip.log") + if !os.IsNotExist(err) { + t.Error("skip.log should not exist") + } +} + +func TestLocalCopySkipsUnchanged(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "copy-skip" + + _ = vol.WriteFile(ctx, sid, "src.txt", []byte("data"), 0o644) + + result1, err := vol.Copy(ctx, sid, "src.txt", "dst.txt", WithForceFull()) + if err != nil { + t.Fatalf("Copy 1: %v", err) + } + if result1.FilesSynced != 1 { + t.Errorf("first copy synced = %d, want 1", result1.FilesSynced) + } + + result2, err := vol.Copy(ctx, sid, "src.txt", "dst.txt") + if err != nil { + t.Fatalf("Copy 2: %v", err) + } + if result2.FilesSynced != 0 { + t.Errorf("second copy synced = %d, want 0 (unchanged)", result2.FilesSynced) + } +} + +func TestLocalCopyForceFull(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + defer vol.Close() + ctx := context.Background() + sid := "copy-force" + + _ = vol.WriteFile(ctx, sid, "src.txt", []byte("data"), 0o644) + + _, _ = vol.Copy(ctx, sid, "src.txt", "dst.txt", WithForceFull()) + + result, err := vol.Copy(ctx, sid, "src.txt", "dst.txt", WithForceFull()) + if err != nil { + t.Fatalf("Copy: %v", err) + } + if result.FilesSynced != 1 { + t.Errorf("force copy synced = %d, want 1", result.FilesSynced) + } +} + +func TestLocalCopyNotExist(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + ctx := context.Background() + + _, err := vol.Copy(ctx, "test", "nonexistent", "dst") + if err == nil { + t.Error("expected error for copy nonexistent source") + } +} + +func TestLocalCopyPathTraversal(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + ctx := context.Background() + + _, err := vol.Copy(ctx, "test", "../../etc/passwd", "dst") + if err == nil { + t.Error("expected error for path traversal in src") + } + + _, err = vol.Copy(ctx, "test", "src", "../../etc/evil") + if err == nil { + t.Error("expected error for path traversal in dst") + } +} + +func TestRemoteCopy(t *testing.T) { + addr := taiTestGRPC() + conn, err := grpc.NewClient(addr, + grpc.WithTransportCredentials(insecure.NewCredentials()), + ) + if err != nil { + t.Skipf("gRPC dial %s: %v", addr, err) + } + defer conn.Close() + + vol := NewRemote(conn) + defer vol.Close() + ctx := context.Background() + sid := "copy-remote-test" + + _ = vol.WriteFile(ctx, sid, "src.txt", []byte("remote copy"), 0o644) + + result, err := vol.Copy(ctx, sid, "src.txt", "dst.txt", WithForceFull()) + if err != nil { + t.Fatalf("Copy: %v", err) + } + if result.FilesSynced < 1 { + t.Errorf("synced = %d", result.FilesSynced) + } + + data, _, err := vol.ReadFile(ctx, sid, "dst.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "remote copy" { + t.Errorf("content = %q", data) + } + + _ = vol.Remove(ctx, sid, ".", true) +} + +func TestRemoteCopyDir(t *testing.T) { + addr := taiTestGRPC() + conn, err := grpc.NewClient(addr, + grpc.WithTransportCredentials(insecure.NewCredentials()), + ) + if err != nil { + t.Skipf("gRPC dial %s: %v", addr, err) + } + defer conn.Close() + + vol := NewRemote(conn) + defer vol.Close() + ctx := context.Background() + sid := "copy-remote-dir" + + _ = vol.MkdirAll(ctx, sid, "src/sub") + _ = vol.WriteFile(ctx, sid, "src/a.txt", []byte("aaa"), 0o644) + _ = vol.WriteFile(ctx, sid, "src/sub/b.txt", []byte("bbb"), 0o644) + + result, err := vol.Copy(ctx, sid, "src", "dst", WithForceFull()) + if err != nil { + t.Fatalf("Copy: %v", err) + } + if result.FilesSynced < 2 { + t.Errorf("synced = %d", result.FilesSynced) + } + + data, _, err := vol.ReadFile(ctx, sid, "dst/sub/b.txt") + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(data) != "bbb" { + t.Errorf("content = %q", data) + } + + _ = vol.Remove(ctx, sid, ".", true) +} diff --git a/tai/workspace/copy.go b/tai/workspace/copy.go new file mode 100644 index 00000000..516e2cbc --- /dev/null +++ b/tai/workspace/copy.go @@ -0,0 +1,158 @@ +package workspace + +import ( + "context" + "io/fs" + "os" + "path/filepath" + + "github.com/yaoapp/yao/tai/volume" +) + +// Copy implements the FS.Copy method with 4-way dispatch: +// +// ws -> ws : Volume.Copy (server-side for remote, local copy for local) +// host -> ws : Volume.SyncPush +// ws -> host : Volume.SyncPull +// host -> host : os-level recursive copy +func (w *workspaceFS) Copy(src, dst string, opts ...volume.SyncOption) (*volume.SyncResult, error) { + srcURI := parseHostURI(src) + dstURI := parseHostURI(dst) + ctx := context.Background() + + switch { + case !srcURI.IsHost && !dstURI.IsHost: + return w.vol.Copy(ctx, w.session, srcURI.Path, dstURI.Path, opts...) + + case srcURI.IsHost && !dstURI.IsHost: + hostPath, err := resolveAbsHostPath(srcURI) + if err != nil { + return nil, err + } + info, err := os.Stat(hostPath) + if err != nil { + return nil, err + } + if !info.IsDir() { + return nil, w.pushSingleFile(ctx, hostPath, dstURI.Path, info) + } + pushOpts := append(sliceClone(opts), volume.WithRemotePath(dstURI.Path)) + return w.vol.SyncPush(ctx, w.session, hostPath, pushOpts...) + + case !srcURI.IsHost && dstURI.IsHost: + hostPath, err := resolveAbsHostPath(dstURI) + if err != nil { + return nil, err + } + srcInfo, statErr := w.vol.Stat(ctx, w.session, srcURI.Path) + if statErr != nil { + return nil, statErr + } + if !srcInfo.IsDir { + return nil, w.pullSingleFile(ctx, srcURI.Path, hostPath) + } + pullOpts := append(sliceClone(opts), volume.WithRemotePath(srcURI.Path)) + return w.vol.SyncPull(ctx, w.session, hostPath, pullOpts...) + + default: + srcPath, err := resolveAbsHostPath(srcURI) + if err != nil { + return nil, err + } + dstPath, err := resolveAbsHostPath(dstURI) + if err != nil { + return nil, err + } + cfg := volume.ApplySyncOpts(opts) + return nil, copyLocalToLocal(srcPath, dstPath, cfg.Excludes) + } +} + +// pushSingleFile reads a host file and writes it into the workspace at dstPath. +func (w *workspaceFS) pushSingleFile(ctx context.Context, hostPath, dstPath string, info os.FileInfo) error { + data, err := os.ReadFile(hostPath) + if err != nil { + return err + } + dir := filepath.Dir(dstPath) + if dir != "" && dir != "." { + if err := w.vol.MkdirAll(ctx, w.session, dir); err != nil { + return err + } + } + perm := info.Mode() + if perm == 0 { + perm = 0o644 + } + return w.vol.WriteFile(ctx, w.session, dstPath, data, perm) +} + +// pullSingleFile reads a workspace file and writes it to hostPath. +func (w *workspaceFS) pullSingleFile(ctx context.Context, srcPath, hostPath string) error { + data, perm, err := w.vol.ReadFile(ctx, w.session, srcPath) + if err != nil { + return err + } + if perm == 0 { + perm = 0o644 + } + if err := os.MkdirAll(filepath.Dir(hostPath), 0o755); err != nil { + return err + } + return os.WriteFile(hostPath, data, perm) +} + +func copyLocalToLocal(src, dst string, excludes []string) error { + info, err := os.Stat(src) + if err != nil { + return err + } + if !info.IsDir() { + data, err := os.ReadFile(src) + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(dst), 0o755); err != nil { + return err + } + return os.WriteFile(dst, data, info.Mode()) + } + + return filepath.WalkDir(src, func(abs string, d fs.DirEntry, err error) error { + if err != nil { + return err + } + rel, _ := filepath.Rel(src, abs) + if rel == "." { + return os.MkdirAll(dst, 0o755) + } + for _, p := range excludes { + if matched, _ := filepath.Match(p, filepath.Base(rel)); matched { + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + } + target := filepath.Join(dst, rel) + if d.IsDir() { + return os.MkdirAll(target, 0o755) + } + data, err := os.ReadFile(abs) + if err != nil { + return err + } + fi, _ := d.Info() + perm := os.FileMode(0o644) + if fi != nil { + perm = fi.Mode() + } + return os.WriteFile(target, data, perm) + }) +} + +func sliceClone(opts []volume.SyncOption) []volume.SyncOption { + cp := make([]volume.SyncOption, len(opts)) + copy(cp, opts) + return cp +} diff --git a/tai/workspace/uri.go b/tai/workspace/uri.go new file mode 100644 index 00000000..e4ed9056 --- /dev/null +++ b/tai/workspace/uri.go @@ -0,0 +1,49 @@ +package workspace + +import ( + "fmt" + "os" + "path/filepath" + "strings" +) + +// hostURI holds the parsed result of a host URI. +type hostURI struct { + Scheme string // "local" or "tmp"; empty for workspace paths + Path string // resolved absolute path (host) or relative path (workspace) + IsHost bool +} + +// parseHostURI extracts scheme and path from a host URI string. +// +// "local:///abs/path" -> {Scheme:"local", Path:"/abs/path", IsHost:true} +// "tmp:///rel/path" -> {Scheme:"tmp", Path:"rel/path", IsHost:true} +// "some/workspace/path" -> {Scheme:"", Path:"some/workspace/path", IsHost:false} +func parseHostURI(raw string) hostURI { + switch { + case strings.HasPrefix(raw, "local:///"): + return hostURI{Scheme: "local", Path: strings.TrimPrefix(raw, "local:///"), IsHost: true} + case strings.HasPrefix(raw, "tmp:///"): + return hostURI{Scheme: "tmp", Path: strings.TrimPrefix(raw, "tmp:///"), IsHost: true} + default: + return hostURI{Path: raw} + } +} + +// resolveAbsHostPath converts a parsed hostURI into an absolute filesystem path. +// For "local" scheme, Path is already absolute (rooted at /). +// For "tmp" scheme, Path is relative to os.TempDir(). +func resolveAbsHostPath(u hostURI) (string, error) { + switch u.Scheme { + case "local": + abs := filepath.Clean("/" + u.Path) + return abs, nil + case "tmp": + if strings.Contains(u.Path, "..") { + return "", fmt.Errorf("path traversal not allowed in tmp:// URI") + } + return filepath.Join(os.TempDir(), u.Path), nil + default: + return "", fmt.Errorf("not a host URI: %q", u.Path) + } +} diff --git a/tai/workspace/workspace.go b/tai/workspace/workspace.go index 6af4e69d..5ff5b518 100644 --- a/tai/workspace/workspace.go +++ b/tai/workspace/workspace.go @@ -25,6 +25,12 @@ type FS interface { RemoveAll(name string) error Rename(oldname, newname string) error MkdirAll(name string, perm os.FileMode) error + + // Copy copies files between workspace paths and/or host paths. + // Host paths use "local:///" (absolute system path) or "tmp:///" (os.TempDir-relative). + // ws↔ws uses Volume.Copy (server-side for remote volumes, avoiding 2N network round-trips). + // Returns non-nil *SyncResult for host↔workspace and ws↔ws transfers; nil for host↔host. + Copy(src, dst string, opts ...volume.SyncOption) (*volume.SyncResult, error) } // New creates an FS backed by the given Volume for the specified session. diff --git a/workspace/fileio_test.go b/workspace/fileio_test.go index 61a1f86e..38d31190 100644 --- a/workspace/fileio_test.go +++ b/workspace/fileio_test.go @@ -230,3 +230,90 @@ func TestFS_NotFound(t *testing.T) { }) } } + +func TestManagerRename(t *testing.T) { + for _, pc := range testPools() { + t.Run(pc.Name, func(t *testing.T) { + m := setupManagerForPool(t, pc) + ws := createWorkspace(t, m, pc.Name) + ctx := context.Background() + + require.NoError(t, m.WriteFile(ctx, ws.ID, "old.txt", []byte("rename me"), 0644)) + require.NoError(t, m.Rename(ctx, ws.ID, "old.txt", "new.txt")) + + data, err := m.ReadFile(ctx, ws.ID, "new.txt") + require.NoError(t, err) + assert.Equal(t, "rename me", string(data)) + + _, err = m.ReadFile(ctx, ws.ID, "old.txt") + assert.Error(t, err) + }) + } +} + +func TestManagerRename_NotFound(t *testing.T) { + for _, pc := range testPools() { + t.Run(pc.Name, func(t *testing.T) { + m := setupManagerForPool(t, pc) + err := m.Rename(context.Background(), "nonexistent", "a", "b") + assert.ErrorIs(t, err, workspace.ErrNotFound) + }) + } +} + +func TestManagerMkdirAll(t *testing.T) { + for _, pc := range testPools() { + t.Run(pc.Name, func(t *testing.T) { + m := setupManagerForPool(t, pc) + ws := createWorkspace(t, m, pc.Name) + ctx := context.Background() + + require.NoError(t, m.MkdirAll(ctx, ws.ID, "a/b/c")) + require.NoError(t, m.WriteFile(ctx, ws.ID, "a/b/c/test.txt", []byte("deep"), 0644)) + + data, err := m.ReadFile(ctx, ws.ID, "a/b/c/test.txt") + require.NoError(t, err) + assert.Equal(t, "deep", string(data)) + }) + } +} + +func TestManagerMkdirAll_NotFound(t *testing.T) { + for _, pc := range testPools() { + t.Run(pc.Name, func(t *testing.T) { + m := setupManagerForPool(t, pc) + err := m.MkdirAll(context.Background(), "nonexistent", "a/b") + assert.ErrorIs(t, err, workspace.ErrNotFound) + }) + } +} + +func TestManagerVolume(t *testing.T) { + for _, pc := range testPools() { + t.Run(pc.Name, func(t *testing.T) { + m := setupManagerForPool(t, pc) + ws := createWorkspace(t, m, pc.Name) + ctx := context.Background() + + vol, wsID, err := m.Volume(ctx, ws.ID) + require.NoError(t, err) + assert.Equal(t, ws.ID, wsID) + assert.NotNil(t, vol) + + require.NoError(t, vol.WriteFile(ctx, wsID, "via-vol.txt", []byte("volume direct"), 0o644)) + data, _, err := vol.ReadFile(ctx, wsID, "via-vol.txt") + require.NoError(t, err) + assert.Equal(t, "volume direct", string(data)) + }) + } +} + +func TestManagerVolume_NotFound(t *testing.T) { + for _, pc := range testPools() { + t.Run(pc.Name, func(t *testing.T) { + m := setupManagerForPool(t, pc) + _, _, err := m.Volume(context.Background(), "nonexistent") + assert.ErrorIs(t, err, workspace.ErrNotFound) + }) + } +} diff --git a/workspace/jsapi/API.md b/workspace/jsapi/API.md new file mode 100644 index 00000000..b0b605ad --- /dev/null +++ b/workspace/jsapi/API.md @@ -0,0 +1,572 @@ +# Workspace JavaScript API + +All methods are available on the global `workspace` object. No constructor needed. + +## Quick Start + +```javascript +// Create a workspace +const ws = workspace.Create({ name: "my-project", owner: "user-123", node: "default" }) + +// File I/O +ws.WriteFile("src/main.go", 'package main\n\nfunc main() {}\n') +const content = ws.ReadFile("src/main.go") + +// Binary file (Base64) +const b64 = ws.ReadFileBase64("image.png") +ws.WriteFileBase64("copy.png", b64) + +// Clean up +workspace.Delete(ws.id) +``` + +--- + +## Static Methods + +### workspace.Create(options) → WorkspaceFS + +Create a new workspace on a Tai node. + +```javascript +const ws = workspace.Create({ + name: "my-project", // required — human-readable name + owner: "user-123", // required — user ID + node: "default", // required — target Tai node + id: "ws-custom-id", // optional — auto-generated if empty + labels: { team: "backend" } // optional — custom labels +}) +``` + +### workspace.Get(id) → WorkspaceFS | null + +Get an existing workspace by ID. Returns `null` if not found. + +```javascript +const ws = workspace.Get("ws-abc123") +if (ws) { + console.log(ws.id, ws.name, ws.node) +} +``` + +### workspace.List(filter?) → WorkspaceInfo[] + +List all workspaces, optionally filtered. + +```javascript +const all = workspace.List() +const mine = workspace.List({ owner: "user-123" }) +const onNode = workspace.List({ node: "gpu-01" }) +``` + +Each element: + +```javascript +{ + id: "ws-abc123", + name: "my-project", + owner: "user-123", + node: "default", + labels: { team: "backend" }, + created_at: "2026-03-07T10:00:00Z", + updated_at: "2026-03-07T10:05:00Z" +} +``` + +### workspace.Delete(id) → void + +Delete a workspace and its storage. + +```javascript +workspace.Delete("ws-abc123") +``` + +--- + +## WorkspaceFS Object + +Returned by `workspace.Create()`, `workspace.Get()`, `box.Workspace()`, and `host.Workspace()`. + +### Properties (read-only) + +| Property | Type | Description | +|----------|------|-------------| +| `ws.id` | string | Workspace ID | +| `ws.name` | string | Workspace name | +| `ws.node` | string | Tai node name | + +--- + +### File Reading + +#### ws.ReadFile(path) → string + +Read file content as UTF-8 string. + +```javascript +const content = ws.ReadFile("src/main.go") +``` + +Go: `workspace.M().ReadFile(id, name)` → `Volume.ReadFile` → `string(data)` + +#### ws.ReadFileBase64(path) → string + +Read file content as Base64-encoded string. Use for binary files (images, archives, etc.). + +```javascript +const b64 = ws.ReadFileBase64("assets/logo.png") +``` + +Go: `M().ReadFile` → `base64.StdEncoding.EncodeToString(data)` + +#### ws.ReadFileBuffer(path) → string + +Read file content as Base64 string. Alias for `ReadFileBase64` (temporary — Uint8Array support pending v8go upgrade). + +```javascript +const b64 = ws.ReadFileBuffer("data.bin") +``` + +Go: `M().ReadFile` → `base64.StdEncoding.EncodeToString(data)` + +--- + +### File Writing + +#### ws.WriteFile(path, data, perm?) → void + +Write string data to a file. Creates parent directories if needed. + +```javascript +ws.WriteFile("src/main.go", "package main\n...") +ws.WriteFile("config.yml", yamlContent, 0644) +``` + +Go: `workspace.M().WriteFile(id, name, []byte(data), perm)` → `Volume.WriteFile` — perm defaults to `0644` + +#### ws.WriteFileBase64(path, b64, perm?) → void + +Write Base64-encoded data to a file. Use for binary files. + +```javascript +ws.WriteFileBase64("assets/logo.png", b64Data) +``` + +Go: `base64.Decode` → `M().WriteFile` → `Volume.WriteFile` + +#### ws.WriteFileBuffer(path, b64, perm?) → void + +Write Base64-encoded data to a file. Alias for `WriteFileBase64` (temporary — Uint8Array support pending v8go upgrade). + +```javascript +ws.WriteFileBuffer("data.bin", b64Data) +``` + +Go: `base64.Decode` → `M().WriteFile` → `Volume.WriteFile` + +--- + +### Directory Operations + +#### ws.ReadDir(path?, recursive?) → DirEntry[] + +List directory contents. Defaults to root (`"."`), non-recursive. + +```javascript +// One level (default) +const entries = ws.ReadDir("src/") +// → [{ name: "main.go", ... }, { name: "utils", is_dir: true, ... }] + +// Recursive — name becomes relative path +const all = ws.ReadDir("src/", true) +// → [{ name: "main.go", ... }, { name: "utils/helper.go", ... }] + +entries.forEach(function(e) { + console.log(e.name, e.is_dir ? "(dir)" : e.size + " bytes") +}) +``` + +Return type: + +```javascript +{ name: "main.go", is_dir: false, size: 1234 } +// recursive mode: name is relative path, e.g. "utils/helper.go" +``` + +Go: non-recursive → `workspace.M().ListDir` → `Volume.ListDir`; recursive → `M().FS()` → `fs.WalkDir` + +#### ws.MkdirAll(path) → void + +Create a directory tree recursively. Permission is always `0755`. + +```javascript +ws.MkdirAll("src/utils/helpers") +``` + +Go: `workspace.M().MkdirAll(id, name)` → `Volume.MkdirAll` (0755) + +--- + +### File Operations + +#### ws.Remove(path) → void + +Remove a single file or empty directory. + +```javascript +ws.Remove("tmp.txt") +``` + +Go: `FS.Remove(name)` + +#### ws.RemoveAll(path) → void + +Remove a file or directory recursively. + +```javascript +ws.RemoveAll("build/") +``` + +Go: `FS.RemoveAll(name)` + +#### ws.Rename(from, to) → void + +Rename or move a file/directory within the workspace. + +```javascript +ws.Rename("old.txt", "new.txt") +ws.Rename("src/foo.go", "src/bar.go") +``` + +Go: `workspace.M().Rename(id, oldname, newname)` → `Volume.Rename` + +#### ws.Copy(src, dst, options?) → SyncResult | void + +Unified copy method. Supports workspace-internal copy and host ↔ workspace copy via `local://` URI prefix. + +**Path resolution:** + +- No prefix → workspace-internal path (relative to workspace root) +- `local://` prefix → relative to App Root (`config.Conf.AppSource`), e.g. `local:///data/templates` → `{AppRoot}/data/templates` +- `tmp://` prefix → relative to `os.TempDir()`, e.g. `tmp:///workspace-staging` → `/tmp/workspace-staging` + +**Security:** `..` traversal is rejected for both `local://` and `tmp://`. `local://` paths escaping App Root are rejected. + +**Examples:** + +```javascript +// workspace → workspace +ws.Copy("src/main.go", "src/main_backup.go") +ws.Copy("templates/", "projects/new/") + +// host → workspace +ws.Copy("local:///app/templates/nextjs", "prompts/") +ws.Copy("local:///data/assistants/bot-a/config", "config/") + +// workspace → host +ws.Copy("build/dist/", "local:///app/output/dist") + +// host → host (no extra fs needed) +ws.Copy("local:///templates/nextjs", "local:///backup/nextjs-backup") + +// tmp dir → workspace +ws.Copy("tmp:///workspace-staging/data", "imported/") + +// workspace → tmp dir +ws.Copy("build/dist/", "tmp:///export-staging") + +// with options (excludes, force) +ws.Copy("local:///app/templates/nextjs", "project/", { + excludes: ["node_modules", ".git", "*.log"], + force: true +}) +``` + +Options: + +```javascript +{ + excludes: ["node_modules"], // optional — glob patterns to exclude + force: false // optional — skip incremental diff, sync everything +} +``` + +Return value: + +- **Workspace internal**: void +- **Host ↔ workspace** (`local://`/`tmp://` on one side): `SyncResult` +- **Host → host** (both `local://`/`tmp://`): void + +```javascript +{ + files_synced: 42, // number of files transferred + bytes_transferred: 1048576, // total bytes + duration_ms: 1234 // time taken in ms +} +``` + +**Dispatch rules (Go layer):** + +| src | dst | Implementation | +|-----|-----|----------------| +| workspace | workspace | `FS.ReadFile` + `FS.WriteFile` (recursive for dirs) | +| host URI | workspace | `Volume.SyncPush(hostPath, wsPath, opts)` | +| workspace | host URI | `Volume.SyncPull(wsPath, hostPath, opts)` | +| host URI | host URI | `os` package recursive copy | + +Where "host URI" = `local://` (relative to App Root) or `tmp://` (relative to `os.TempDir()`). + +--- + +### Archive & Compression + +All archive methods operate on paths within the workspace. Pack operations support an `excludes` option. + +#### ws.Zip(src, dst, options?) → ArchiveResult + +Create a ZIP archive from `src` directory to `dst` file. + +```javascript +const result = ws.Zip("src/", "dist.zip") +const filtered = ws.Zip("src/", "dist.zip", { excludes: ["*.log", "node_modules"] }) +``` + +#### ws.Unzip(src, dst) → ArchiveResult + +Extract a ZIP archive from `src` file to `dst` directory. + +```javascript +const result = ws.Unzip("dist.zip", "extracted/") +``` + +#### ws.Gzip(src, dst) → ArchiveResult + +Compress a single file with gzip. + +```javascript +ws.Gzip("data.json", "data.json.gz") +``` + +#### ws.Gunzip(src, dst) → ArchiveResult + +Decompress a gzip file. + +```javascript +ws.Gunzip("data.json.gz", "data.json") +``` + +#### ws.Tar(src, dst, options?) → ArchiveResult + +Create a tar archive from `src` directory. + +```javascript +ws.Tar("src/", "archive.tar") +ws.Tar("src/", "archive.tar", { excludes: [".git"] }) +``` + +#### ws.Untar(src, dst) → ArchiveResult + +Extract a tar archive. + +```javascript +ws.Untar("archive.tar", "extracted/") +``` + +#### ws.Tgz(src, dst, options?) → ArchiveResult + +Create a gzip-compressed tar archive (.tar.gz / .tgz). + +```javascript +ws.Tgz("src/", "archive.tgz") +``` + +#### ws.Untgz(src, dst) → ArchiveResult + +Extract a gzip-compressed tar archive. + +```javascript +ws.Untgz("archive.tgz", "extracted/") +``` + +**ArchiveResult:** + +```javascript +{ + size_bytes: 102400, // output file size (pack) or total extracted size (unpack) + files_count: 15 // number of files processed +} +``` + +Go: delegates to `Volume.Zip`, `Volume.Unzip`, `Volume.Gzip`, `Volume.Gunzip`, `Volume.Tar`, `Volume.Untar`, `Volume.Tgz`, `Volume.Untgz`. + +--- + +### File Information + +#### ws.Stat(path) → FileInfo + +Get file or directory metadata. + +```javascript +const info = ws.Stat("src/main.go") +console.log(info.name, info.size, info.is_dir, info.mod_time, info.mode) +``` + +Return type: + +```javascript +{ + name: "main.go", + size: 1234, + is_dir: false, + mod_time: "2026-03-07T10:00:00Z", + mode: 0644 +} +``` + +Go: `FS.Stat(name) → fs.FileInfo` + +#### ws.Exists(path) → boolean + +Check if a file or directory exists. + +```javascript +if (ws.Exists("config.yml")) { + // ... +} +``` + +Go: `FS.Stat(name)` — returns `true` if err == nil + +#### ws.IsDir(path) → boolean + +Check if a path is a directory. Returns `false` if not found. + +```javascript +if (ws.IsDir("src/")) { + // ... +} +``` + +Go: `FS.Stat(name) → info.IsDir()` + +#### ws.IsFile(path) → boolean + +Check if a path is a regular file. Returns `false` if not found. + +```javascript +if (ws.IsFile("main.go")) { + // ... +} +``` + +Go: `FS.Stat(name) → !info.IsDir()` + +--- + +## Go Interface Reference + +WorkspaceFS methods map to `taiworkspace.FS` and `volume.Volume`. Some methods (Stat, Remove, RemoveAll) use `workspace.M().FS()` directly; others (ReadFile, WriteFile, Rename, MkdirAll) go through `workspace.M()` → `Volume`: + +FS interface (defined in `tai/workspace/workspace.go`): + +```go +type FS interface { + fs.FS // Open(name) (fs.File, error) + fs.StatFS // Stat(name) (fs.FileInfo, error) + fs.ReadFileFS // ReadFile(name) ([]byte, error) + fs.ReadDirFS // ReadDir(name) ([]fs.DirEntry, error) + io.Closer + + WriteFile(name string, data []byte, perm os.FileMode) error + Remove(name string) error + RemoveAll(name string) error + Rename(oldname, newname string) error + MkdirAll(name string, perm os.FileMode) error +} +``` + +Archive methods delegate to `volume.Volume` (defined in `tai/volume/volume.go`): + +```go +type Volume interface { + // ... FS methods ... + Zip(ctx, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) + Unzip(ctx, sessionID, src, dst string) (*ArchiveResult, error) + Gzip(ctx, sessionID, src, dst string) (*ArchiveResult, error) + Gunzip(ctx, sessionID, src, dst string) (*ArchiveResult, error) + Tar(ctx, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) + Untar(ctx, sessionID, src, dst string) (*ArchiveResult, error) + Tgz(ctx, sessionID, src, dst string, excludes []string) (*ArchiveResult, error) + Untgz(ctx, sessionID, src, dst string) (*ArchiveResult, error) +} +``` + +`Exists`, `IsDir`, `IsFile` are thin JSAPI wrappers over `FS.Stat`. + +`ReadDir` adds a `recursive` parameter in Go — non-recursive calls `Volume.ListDir`, recursive uses `M().FS()` + `fs.WalkDir`. + +`Copy` is implemented in Go with host URI dispatch: workspace paths use `FS.ReadFile`/`FS.WriteFile`; `local://` (App Root) and `tmp://` (`os.TempDir()`) paths use `Volume.SyncPush`/`SyncPull` with `WithRemotePath`; host-to-host uses `os` package. + +Base64/Buffer variants (`ReadFileBase64`, `ReadFileBuffer`, `WriteFileBase64`, `WriteFileBuffer`) are Go-side encoding wrappers around `M().ReadFile` / `M().WriteFile` (→ `Volume`). Buffer variants currently use Base64 encoding (Uint8Array support pending v8go upgrade). + +--- + +## Method Summary + +### Standard FS — core file operations + +| Method | Returns | Go mapping | +|--------|---------|------------| +| `ws.ReadFile(path)` | string | `M().ReadFile` → `Volume.ReadFile` | +| `ws.WriteFile(path, data, perm?)` | void | `M().WriteFile` → `Volume.WriteFile` | +| `ws.Stat(path)` | FileInfo | `M().FS()` → `FS.Stat` | +| `ws.MkdirAll(path)` | void | `M().MkdirAll` → `Volume.MkdirAll` (0755) | +| `ws.Remove(path)` | void | `M().FS()` → `FS.Remove` | +| `ws.RemoveAll(path)` | void | `M().FS()` → `FS.RemoveAll` | +| `ws.Rename(from, to)` | void | `M().Rename` → `Volume.Rename` | + +### Go wrapper — implemented in Go, called directly from JSAPI + +| Method | Returns | Go implementation | +|--------|---------|-------------------| +| `ws.ReadDir(path?, recursive?)` | DirEntry[] | non-recursive: `Volume.ListDir`; recursive: `FS` + `fs.WalkDir` | +| `ws.ReadFileBase64(path)` | string | `M().ReadFile` → `base64.Encode` | +| `ws.ReadFileBuffer(path)` | string | `M().ReadFile` → `base64.Encode` (temp, Uint8Array pending) | +| `ws.WriteFileBase64(path, b64, perm?)` | void | `base64.Decode` → `M().WriteFile` | +| `ws.WriteFileBuffer(path, b64, perm?)` | void | `base64.Decode` → `M().WriteFile` (temp, Uint8Array pending) | +| `ws.Copy(src, dst, opts?)` | void / SyncResult | dispatch by host URI prefix (see below) | + +`ws.Copy` dispatch (all handled in Go): + +| src | dst | Returns | Go implementation | +|-----|-----|---------|-------------------| +| workspace | workspace | void | `FS.ReadFile` + `FS.WriteFile` (recursive for dirs) | +| host URI | workspace | SyncResult | `Volume.SyncPush` with `WithRemotePath` | +| workspace | host URI | SyncResult | `Volume.SyncPull` with `WithRemotePath` | +| host URI | host URI | void | `os` package recursive copy | + +Where "host URI" = `local://` (App Root) or `tmp://` (`os.TempDir()`). + +### Archive — delegates to `Volume` interface + +| Method | Returns | Go implementation | +|--------|---------|-------------------| +| `ws.Zip(src, dst, opts?)` | ArchiveResult | `Volume.Zip` | +| `ws.Unzip(src, dst)` | ArchiveResult | `Volume.Unzip` | +| `ws.Gzip(src, dst)` | ArchiveResult | `Volume.Gzip` | +| `ws.Gunzip(src, dst)` | ArchiveResult | `Volume.Gunzip` | +| `ws.Tar(src, dst, opts?)` | ArchiveResult | `Volume.Tar` | +| `ws.Untar(src, dst)` | ArchiveResult | `Volume.Untar` | +| `ws.Tgz(src, dst, opts?)` | ArchiveResult | `Volume.Tgz` | +| `ws.Untgz(src, dst)` | ArchiveResult | `Volume.Untgz` | + +### JSAPI composition — thin JS wrappers over standard FS + +| Method | Returns | Composed from | +|--------|---------|---------------| +| `ws.Exists(path)` | boolean | `Stat` → err == nil | +| `ws.IsDir(path)` | boolean | `Stat` → `info.IsDir()` | +| `ws.IsFile(path)` | boolean | `Stat` → `!info.IsDir()` | + +**Total: 24 methods + 3 read-only properties** (7 standard FS + 6 Go wrapper + 8 archive + 3 JSAPI composition) diff --git a/workspace/jsapi/fs.go b/workspace/jsapi/fs.go index 17c0515c..c82444b9 100644 --- a/workspace/jsapi/fs.go +++ b/workspace/jsapi/fs.go @@ -1,130 +1,510 @@ package jsapi import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "io/fs" + "os" + "path/filepath" + "strings" + "time" + + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/tai/volume" + "github.com/yaoapp/yao/workspace" "rogchap.com/v8go" ) // NewFSObject creates a JS WorkspaceFS object backed by a workspace ID string. -// All methods delegate to workspace.M() → FS — no Go object passed to V8. -// -// # Properties (read-only) -// -// ws.id → string // workspace ID ← workspaceID arg -// ws.name → string // workspace name ← Workspace.Name -// ws.node → string // tai node name ← Workspace.Node -// -// # Methods — Go mapping -// -// Each method internally does: fs, _ := workspace.M().FS(ctx, workspaceID) -// then calls the corresponding method on taiworkspace.FS. -// -// ws.ReadFile(path) → string -// -// Go: FS.ReadFile(name string) ([]byte, error) -// — also available via Manager.ReadFile(ctx, id, path) -// JS args: path string -// JS returns: string (UTF-8 content of the file) -// -// ws.WriteFile(path, data, perm?) → void -// -// Go: FS.WriteFile(name string, data []byte, perm os.FileMode) error -// — also available via Manager.WriteFile(ctx, id, path, data, perm) -// JS args: path string, data string|Uint8Array, perm? number (default 0644) -// -// ws.ReadDir(path?) → DirEntry[] -// -// Go: FS.ReadDir(name string) ([]fs.DirEntry, error) -// — also available via Manager.ListDir(ctx, id, path) -// JS args: path string (default ".") -// JS returns: [{ -// name: string, ← DirEntry.Name() -// is_dir: boolean, ← DirEntry.IsDir() -// size: number ← DirEntry.Info().Size() -// }] -// -// ws.Stat(path) → FileInfo -// -// Go: FS.Stat(name string) (fs.FileInfo, error) -// JS args: path string -// JS returns: { -// name: string, ← FileInfo.Name() -// size: number, ← FileInfo.Size() -// is_dir: boolean, ← FileInfo.IsDir() -// mod_time: string ← FileInfo.ModTime() (ISO 8601) -// } -// -// ws.MkdirAll(path, perm?) → void -// -// Go: FS.MkdirAll(name string, perm os.FileMode) error -// JS args: path string, perm? number (default 0755) -// -// ws.Remove(path) → void -// -// Go: FS.Remove(name string) error -// — also available via Manager.Remove(ctx, id, path) -// JS args: path string (single file or empty directory) -// -// ws.RemoveAll(path) → void -// -// Go: FS.RemoveAll(name string) error -// JS args: path string (recursive removal) -// -// ws.Rename(from, to) → void -// -// Go: FS.Rename(oldname, newname string) error -// JS args: from string, to string -// -// # Base64 variants (PLANNED — not yet implemented) -// -// Avoids V8↔Go binary bridge overhead for images, archives, etc. -// -// ws.ReadFileBase64(path) → string -// -// Go: FS.ReadFile(name) → base64.StdEncoding.EncodeToString(data) -// JS args: path string -// JS returns: string (base64-encoded content) -// -// ws.WriteFileBase64(path, b64, perm?) → void -// -// Go: base64.StdEncoding.DecodeString(b64) → FS.WriteFile(name, data, perm) -// JS args: path string, b64 string, perm? number (default 0644) -// -// # Host copy (PLANNED — not yet implemented) -// -// Copy files/dirs from Yao host filesystem into the workspace volume. -// Useful for seeding workspaces with templates, config files, assets, etc. -// -// ws.CopyFromHost(hostPath, destPath?) → void -// -// Copies a single file or directory tree from the Yao host into the workspace. -// Go: read host file(s) → FS.WriteFile / FS.MkdirAll for each entry -// JS args: hostPath string (absolute path on Yao host), -// destPath? string (target path inside workspace, default basename of hostPath) -// -// ws.CopyFromHostArchive(hostPath, destPath?) → void -// -// For large directory trees: zip on host → transfer → unzip on Tai node. -// Requires Tai server-side unarchive support. -// Go: zip hostPath → tai Volume upload → tai unarchive at destPath -// JS args: hostPath string, destPath? string (default ".") +// All methods delegate to workspace.M() — no Go object is passed to V8. func NewFSObject(v8ctx *v8go.Context, workspaceID string) (*v8go.Value, error) { - // TODO: Phase 2 implementation - // 1. Create JS object via v8go.NewObjectTemplate - // 2. Set read-only properties: id, name, node (from workspace.M().Get(workspaceID)) - // 3. Bind each method as FunctionTemplate: - // - ReadFile → workspace.M().FS(ctx, id).ReadFile(path) - // - WriteFile → workspace.M().FS(ctx, id).WriteFile(path, data, perm) - // - ReadDir → workspace.M().FS(ctx, id).ReadDir(path) - // - Stat → workspace.M().FS(ctx, id).Stat(path) - // - MkdirAll → workspace.M().FS(ctx, id).MkdirAll(path, perm) - // - Remove → workspace.M().FS(ctx, id).Remove(path) - // - RemoveAll → workspace.M().FS(ctx, id).RemoveAll(path) - // - Rename → workspace.M().FS(ctx, id).Rename(old, new) - // - // PLANNED (not yet implemented): - // - ReadFileBase64 → ReadFile + base64 encode in Go - // - WriteFileBase64 → base64 decode in Go + WriteFile - // - CopyFromHost → host fs.Read → FS.Write (file-by-file) - // - CopyFromHostArchive → zip on host → tai transfer → unzip (needs Tai support) - return nil, nil + iso := v8ctx.Isolate() + ctx := context.Background() + wsID := workspaceID + + ws, err := workspace.M().Get(ctx, wsID) + if err != nil { + return nil, fmt.Errorf("workspace %s: %w", wsID, err) + } + + tpl := v8go.NewObjectTemplate(iso) + + tpl.Set("ReadFile", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "ReadFile requires a path") + } + data, err := workspace.M().ReadFile(ctx, wsID, args[0].String()) + if err != nil { + return throwError(info, err.Error()) + } + val, _ := v8go.NewValue(iso, string(data)) + return val + })) + + tpl.Set("WriteFile", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 2 { + return throwError(info, "WriteFile requires path and data") + } + perm := os.FileMode(0o644) + if len(args) > 2 && args[2].IsNumber() { + perm = os.FileMode(args[2].Int32()) + } + if err := workspace.M().WriteFile(ctx, wsID, args[0].String(), []byte(args[1].String()), perm); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("ReadDir", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + path := "." + args := info.Args() + if len(args) > 0 && args[0].IsString() { + path = args[0].String() + } + recursive := false + if len(args) > 1 && args[1].IsBoolean() { + recursive = args[1].Boolean() + } + + if recursive { + return readDirRecursive(info, wsID, path) + } + + entries, err := workspace.M().ListDir(ctx, wsID, path) + if err != nil { + return throwError(info, err.Error()) + } + data, _ := json.Marshal(entries) + val, _ := v8go.JSONParse(info.Context(), string(data)) + return val + })) + + tpl.Set("Stat", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "Stat requires a path") + } + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + return throwError(info, err.Error()) + } + fi, err := fsys.Stat(args[0].String()) + if err != nil { + return throwError(info, err.Error()) + } + return fileInfoToJS(info, fi) + })) + + tpl.Set("MkdirAll", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "MkdirAll requires a path") + } + if err := workspace.M().MkdirAll(ctx, wsID, args[0].String()); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("Remove", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "Remove requires a path") + } + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + return throwError(info, err.Error()) + } + if err := fsys.Remove(args[0].String()); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("RemoveAll", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "RemoveAll requires a path") + } + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + return throwError(info, err.Error()) + } + if err := fsys.RemoveAll(args[0].String()); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("Rename", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 2 { + return throwError(info, "Rename requires from and to paths") + } + if err := workspace.M().Rename(ctx, wsID, args[0].String(), args[1].String()); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("ReadFileBase64", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "ReadFileBase64 requires a path") + } + data, err := workspace.M().ReadFile(ctx, wsID, args[0].String()) + if err != nil { + return throwError(info, err.Error()) + } + val, _ := v8go.NewValue(iso, base64.StdEncoding.EncodeToString(data)) + return val + })) + + tpl.Set("WriteFileBase64", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 2 { + return throwError(info, "WriteFileBase64 requires path and b64 data") + } + perm := os.FileMode(0o644) + if len(args) > 2 && args[2].IsNumber() { + perm = os.FileMode(args[2].Int32()) + } + decoded, err := base64.StdEncoding.DecodeString(args[1].String()) + if err != nil { + return throwError(info, "base64 decode: "+err.Error()) + } + if err := workspace.M().WriteFile(ctx, wsID, args[0].String(), decoded, perm); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("ReadFileBuffer", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + return throwError(info, "ReadFileBuffer requires a path") + } + data, err := workspace.M().ReadFile(ctx, wsID, args[0].String()) + if err != nil { + return throwError(info, err.Error()) + } + val, _ := v8go.NewValue(iso, base64.StdEncoding.EncodeToString(data)) + return val + })) + + tpl.Set("WriteFileBuffer", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 2 { + return throwError(info, "WriteFileBuffer requires path and data") + } + perm := os.FileMode(0o644) + if len(args) > 2 && args[2].IsNumber() { + perm = os.FileMode(args[2].Int32()) + } + decoded, err := base64.StdEncoding.DecodeString(args[1].String()) + if err != nil { + return throwError(info, "decode: "+err.Error()) + } + if err := workspace.M().WriteFile(ctx, wsID, args[0].String(), decoded, perm); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) + })) + + tpl.Set("Exists", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + val, _ := v8go.NewValue(iso, false) + return val + } + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + val, _ := v8go.NewValue(iso, false) + return val + } + _, err = fsys.Stat(args[0].String()) + val, _ := v8go.NewValue(iso, err == nil) + return val + })) + + tpl.Set("IsDir", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + val, _ := v8go.NewValue(iso, false) + return val + } + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + val, _ := v8go.NewValue(iso, false) + return val + } + fi, err := fsys.Stat(args[0].String()) + val, _ := v8go.NewValue(iso, err == nil && fi.IsDir()) + return val + })) + + tpl.Set("IsFile", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + args := info.Args() + if len(args) < 1 { + val, _ := v8go.NewValue(iso, false) + return val + } + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + val, _ := v8go.NewValue(iso, false) + return val + } + fi, err := fsys.Stat(args[0].String()) + val, _ := v8go.NewValue(iso, err == nil && !fi.IsDir()) + return val + })) + + tpl.Set("Copy", v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + return copyHandler(info, wsID) + })) + + tpl.Set("Zip", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "zip"))) + tpl.Set("Unzip", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "unzip"))) + tpl.Set("Gzip", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "gzip"))) + tpl.Set("Gunzip", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "gunzip"))) + tpl.Set("Tar", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "tar"))) + tpl.Set("Untar", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "untar"))) + tpl.Set("Tgz", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "tgz"))) + tpl.Set("Untgz", v8go.NewFunctionTemplate(iso, archiveHandler(iso, wsID, "untgz"))) + + obj, err := tpl.NewInstance(v8ctx) + if err != nil { + return nil, err + } + + idVal, _ := v8go.NewValue(iso, ws.ID) + obj.Set("id", idVal) + nameVal, _ := v8go.NewValue(iso, ws.Name) + obj.Set("name", nameVal) + nodeVal, _ := v8go.NewValue(iso, ws.Node) + obj.Set("node", nodeVal) + + return obj.Value, nil +} + +func archiveHandler(iso *v8go.Isolate, wsID, op string) v8go.FunctionCallback { + return func(info *v8go.FunctionCallbackInfo) *v8go.Value { + ctx := context.Background() + args := info.Args() + if len(args) < 2 { + return throwError(info, op+": requires src and dst paths") + } + src := args[0].String() + dst := args[1].String() + + var excludes []string + if len(args) > 2 && args[2].IsObject() { + excObj, _ := args[2].AsObject() + if excObj != nil { + if v, e := excObj.Get("excludes"); e == nil && v.IsObject() { + excludes = parseStringArray(v) + } + } + } + + vol, sid, err := workspace.M().Volume(ctx, wsID) + if err != nil { + return throwError(info, err.Error()) + } + + var result *volume.ArchiveResult + switch op { + case "zip": + result, err = vol.Zip(ctx, sid, src, dst, excludes) + case "unzip": + result, err = vol.Unzip(ctx, sid, src, dst) + case "gzip": + result, err = vol.Gzip(ctx, sid, src, dst) + case "gunzip": + result, err = vol.Gunzip(ctx, sid, src, dst) + case "tar": + result, err = vol.Tar(ctx, sid, src, dst, excludes) + case "untar": + result, err = vol.Untar(ctx, sid, src, dst) + case "tgz": + result, err = vol.Tgz(ctx, sid, src, dst, excludes) + case "untgz": + result, err = vol.Untgz(ctx, sid, src, dst) + } + if err != nil { + return throwError(info, err.Error()) + } + + data, _ := json.Marshal(map[string]interface{}{ + "size_bytes": result.SizeBytes, + "files_count": result.FilesCount, + }) + val, _ := v8go.JSONParse(info.Context(), string(data)) + return val + } +} + +func readDirRecursive(info *v8go.FunctionCallbackInfo, wsID, path string) *v8go.Value { + ctx := context.Background() + fsys, err := workspace.M().FS(ctx, wsID) + if err != nil { + return throwError(info, err.Error()) + } + + type entry struct { + Name string `json:"name"` + IsDir bool `json:"is_dir"` + Size int64 `json:"size"` + } + var entries []entry + + _ = fs.WalkDir(fsys, path, func(p string, d fs.DirEntry, err error) error { + if err != nil || p == path { + return err + } + rel, _ := filepath.Rel(path, p) + var size int64 + if fi, e := d.Info(); e == nil { + size = fi.Size() + } + entries = append(entries, entry{Name: filepath.ToSlash(rel), IsDir: d.IsDir(), Size: size}) + return nil + }) + + data, _ := json.Marshal(entries) + val, _ := v8go.JSONParse(info.Context(), string(data)) + return val +} + +func fileInfoToJS(info *v8go.FunctionCallbackInfo, fi fs.FileInfo) *v8go.Value { + data, _ := json.Marshal(map[string]interface{}{ + "name": fi.Name(), + "size": fi.Size(), + "is_dir": fi.IsDir(), + "mod_time": fi.ModTime().Format(time.RFC3339), + "mode": uint32(fi.Mode()), + }) + val, _ := v8go.JSONParse(info.Context(), string(data)) + return val +} + +func copyHandler(info *v8go.FunctionCallbackInfo, wsID string) *v8go.Value { + iso := info.Context().Isolate() + ctx := context.Background() + args := info.Args() + if len(args) < 2 { + return throwError(info, "Copy requires src and dst paths") + } + + src := args[0].String() + dst := args[1].String() + + var excludes []string + force := false + if len(args) > 2 && args[2].IsObject() { + optsObj, _ := args[2].AsObject() + if optsObj != nil { + if v, e := optsObj.Get("excludes"); e == nil && v.IsObject() { + excludes = parseStringArray(v) + } + if v, e := optsObj.Get("force"); e == nil && v.IsBoolean() { + force = v.Boolean() + } + } + } + + var opts []volume.SyncOption + if len(excludes) > 0 { + opts = append(opts, volume.WithExcludes(excludes...)) + } + if force { + opts = append(opts, volume.WithForceFull()) + } + + // Map JSAPI local:// (AppRoot-relative) to Go-layer local:/// (absolute) + src = mapHostURI(src) + dst = mapHostURI(dst) + + wsFS, err := workspace.M().FS(ctx, wsID) + if err != nil { + return throwError(info, err.Error()) + } + + result, err := wsFS.Copy(src, dst, opts...) + if err != nil { + return throwError(info, err.Error()) + } + if result != nil { + return syncResultToJS(info, result) + } + return v8go.Undefined(iso) +} + +func syncResultToJS(info *v8go.FunctionCallbackInfo, r *volume.SyncResult) *v8go.Value { + data, _ := json.Marshal(map[string]interface{}{ + "files_synced": r.FilesSynced, + "bytes_transferred": r.BytesTransferred, + "duration_ms": r.Duration.Milliseconds(), + }) + val, _ := v8go.JSONParse(info.Context(), string(data)) + return val +} + +// mapHostURI converts JSAPI host URIs to Go-layer absolute URIs. +// local://relative -> local:///{AppSource}/relative (with security checks) +// tmp://relative -> tmp:///relative (Go layer resolves os.TempDir) +// other -> unchanged (workspace-relative path) +func mapHostURI(raw string) string { + switch { + case strings.HasPrefix(raw, "local://"): + rel := strings.TrimPrefix(raw, "local://") + if strings.Contains(rel, "..") { + return raw + } + appRoot := config.Conf.AppSource + abs := filepath.Join(appRoot, rel) + resolved, err := filepath.EvalSymlinks(abs) + if err != nil { + resolved = abs + } + if !strings.HasPrefix(resolved, appRoot) { + return raw + } + return "local:///" + resolved + + case strings.HasPrefix(raw, "tmp://"): + rel := strings.TrimPrefix(raw, "tmp://") + if strings.Contains(rel, "..") { + return raw + } + return "tmp:///" + rel + + default: + return raw + } +} + +func parseStringArray(val *v8go.Value) []string { + obj, err := val.AsObject() + if err != nil { + return nil + } + lenVal, err := obj.Get("length") + if err != nil { + return nil + } + length := int(lenVal.Int32()) + result := make([]string, 0, length) + for i := 0; i < length; i++ { + item, err := obj.GetIdx(uint32(i)) + if err != nil || !item.IsString() { + continue + } + result = append(result, item.String()) + } + return result } diff --git a/workspace/jsapi/jsapi.go b/workspace/jsapi/jsapi.go index 624818da..2944a6c6 100644 --- a/workspace/jsapi/jsapi.go +++ b/workspace/jsapi/jsapi.go @@ -1,30 +1,12 @@ // Package jsapi registers the workspace namespace into the Yao V8 runtime. -// -// All methods are static on the workspace object — no constructor. -// -// # JavaScript API -// -// const ws = workspace.Create({ name: "proj", owner: "user1", node: "default" }) -// const ws = workspace.Get(id) -// ws.ReadFile("main.go") → string -// ws.WriteFile("out.txt", data) → void -// ws.ReadDir("src/") → [{ name, is_dir, size }] -// workspace.Delete(id) → void -// -// # Go mapping -// -// workspace.Create(opts) → Manager.Create(ctx, CreateOptions) → *Workspace → WorkspaceFS -// workspace.Get(id) → Manager.Get(ctx, id) → *Workspace → WorkspaceFS -// workspace.List(filter?) → Manager.List(ctx, ListOptions) → []*Workspace → WorkspaceInfo[] -// workspace.Delete(id) → Manager.Delete(ctx, id, false) → void -// -// Registration happens via init() — import with: -// -// _ "github.com/yaoapp/yao/workspace/jsapi" package jsapi import ( + "context" + "encoding/json" + v8 "github.com/yaoapp/gou/runtime/v8" + "github.com/yaoapp/yao/workspace" "rogchap.com/v8go" ) @@ -32,7 +14,6 @@ func init() { v8.RegisterObject("workspace", ExportObject) } -// ExportObject exports the workspace namespace object to V8. func ExportObject(iso *v8go.Isolate) *v8go.ObjectTemplate { obj := v8go.NewObjectTemplate(iso) obj.Set("Create", v8go.NewFunctionTemplate(iso, wsCreate)) @@ -42,88 +23,152 @@ func ExportObject(iso *v8go.Isolate) *v8go.ObjectTemplate { return obj } -// wsCreate: `workspace.Create(options)` → WorkspaceFS -// -// Go: Manager.Create(ctx, CreateOptions) (*Workspace, error) -// -// JS options → Go CreateOptions mapping: -// -// { -// id: string → CreateOptions.ID // optional; auto-generated if empty -// name: string → CreateOptions.Name // required, human-readable name -// owner: string → CreateOptions.Owner // required, user ID -// node: string → CreateOptions.Node // required, target Tai node -// labels: object → CreateOptions.Labels // optional, map[string]string -// } -// -// Returns: WorkspaceFS object (see fs.go) func wsCreate(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. Parse options from info.Args()[0] - // 2. Validate required fields (name, owner, node) - // 3. ws := workspace.M().Create(ctx, opts) - // 4. Return NewFSObject(v8ctx, ws.ID) - return v8go.Undefined(info.Context().Isolate()) + ctx := context.Background() + + args := info.Args() + if len(args) < 1 || !args[0].IsObject() { + return throwError(info, "workspace.Create requires an options object") + } + + optsObj, err := args[0].AsObject() + if err != nil { + return throwError(info, "invalid options: "+err.Error()) + } + + opts := workspace.CreateOptions{} + if v, e := optsObj.Get("id"); e == nil && v.IsString() { + opts.ID = v.String() + } + if v, e := optsObj.Get("name"); e == nil && v.IsString() { + opts.Name = v.String() + } + if v, e := optsObj.Get("owner"); e == nil && v.IsString() { + opts.Owner = v.String() + } + if v, e := optsObj.Get("node"); e == nil && v.IsString() { + opts.Node = v.String() + } + if v, e := optsObj.Get("labels"); e == nil && v.IsObject() { + opts.Labels = parseStringMapFromValue(info.Context(), v) + } + + if opts.Name == "" || opts.Owner == "" || opts.Node == "" { + return throwError(info, "workspace.Create: name, owner, and node are required") + } + + ws, err := workspace.M().Create(ctx, opts) + if err != nil { + return throwError(info, err.Error()) + } + + val, err := NewFSObject(info.Context(), ws.ID) + if err != nil { + return throwError(info, err.Error()) + } + return val } -// wsGet: `workspace.Get(id)` → WorkspaceFS | null -// -// Go: Manager.Get(ctx, id) (*Workspace, error) -// -// Args: -// -// id: string — workspace ID -// -// Returns: WorkspaceFS object if found, null if not found func wsGet(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. id = info.Args()[0].String() - // 2. ws, err := workspace.M().Get(ctx, id) - // 3. Return NewFSObject(v8ctx, id) or null - return v8go.Undefined(info.Context().Isolate()) + iso := info.Context().Isolate() + ctx := context.Background() + + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "workspace.Get requires a string ID") + } + + id := args[0].String() + _, err := workspace.M().Get(ctx, id) + if err != nil { + return v8go.Null(iso) + } + + val, err := NewFSObject(info.Context(), id) + if err != nil { + return v8go.Null(iso) + } + return val } -// wsList: `workspace.List(filter?)` → WorkspaceInfo[] -// -// Go: Manager.List(ctx, ListOptions) ([]*Workspace, error) -// -// JS filter → Go ListOptions mapping: -// -// { -// owner: string → ListOptions.Owner // filter by owner; empty = all -// node: string → ListOptions.Node // filter by node; empty = all -// } -// -// Returns: WorkspaceInfo[] — each element: -// -// { -// id: string ← Workspace.ID -// name: string ← Workspace.Name -// owner: string ← Workspace.Owner -// node: string ← Workspace.Node -// labels: object ← Workspace.Labels -// created_at: string ← Workspace.CreatedAt (ISO 8601) -// updated_at: string ← Workspace.UpdatedAt (ISO 8601) -// } func wsList(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. Parse optional filter from info.Args()[0] - // 2. list := workspace.M().List(ctx, opts) - // 3. Convert each *Workspace → JS object - // 4. Return JS array - return v8go.Undefined(info.Context().Isolate()) + ctx := context.Background() + v8ctx := info.Context() + + opts := workspace.ListOptions{} + args := info.Args() + if len(args) > 0 && args[0].IsObject() { + filterObj, _ := args[0].AsObject() + if filterObj != nil { + if v, e := filterObj.Get("owner"); e == nil && v.IsString() { + opts.Owner = v.String() + } + if v, e := filterObj.Get("node"); e == nil && v.IsString() { + opts.Node = v.String() + } + } + } + + list, err := workspace.M().List(ctx, opts) + if err != nil { + return throwError(info, err.Error()) + } + + type wsInfo struct { + ID string `json:"id"` + Name string `json:"name"` + Owner string `json:"owner"` + Node string `json:"node"` + Labels map[string]string `json:"labels,omitempty"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` + } + + items := make([]wsInfo, len(list)) + for i, ws := range list { + items[i] = wsInfo{ + ID: ws.ID, Name: ws.Name, Owner: ws.Owner, Node: ws.Node, + Labels: ws.Labels, CreatedAt: ws.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), + UpdatedAt: ws.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), + } + } + + data, _ := json.Marshal(items) + val, err := v8go.JSONParse(v8ctx, string(data)) + if err != nil { + return throwError(info, err.Error()) + } + return val } -// wsDelete: `workspace.Delete(id)` → void -// -// Go: Manager.Delete(ctx, id string, force bool) error -// -// Args: -// -// id: string — workspace ID to remove (force = false) func wsDelete(info *v8go.FunctionCallbackInfo) *v8go.Value { - // TODO: Phase 2 - // 1. id = info.Args()[0].String() - // 2. workspace.M().Delete(ctx, id, false) - return v8go.Undefined(info.Context().Isolate()) + iso := info.Context().Isolate() + ctx := context.Background() + + args := info.Args() + if len(args) < 1 || !args[0].IsString() { + return throwError(info, "workspace.Delete requires a string ID") + } + + if err := workspace.M().Delete(ctx, args[0].String(), false); err != nil { + return throwError(info, err.Error()) + } + return v8go.Undefined(iso) +} + +func throwError(info *v8go.FunctionCallbackInfo, msg string) *v8go.Value { + iso := info.Context().Isolate() + e, _ := v8go.NewValue(iso, msg) + iso.ThrowException(e) + return v8go.Undefined(iso) +} + +func parseStringMapFromValue(v8ctx *v8go.Context, val *v8go.Value) map[string]string { + result := make(map[string]string) + jsonStr, err := v8go.JSONStringify(v8ctx, val) + if err != nil { + return result + } + _ = json.Unmarshal([]byte(jsonStr), &result) + return result } diff --git a/workspace/jsapi/jsapi_test.go b/workspace/jsapi/jsapi_test.go new file mode 100644 index 00000000..63bd786f --- /dev/null +++ b/workspace/jsapi/jsapi_test.go @@ -0,0 +1,460 @@ +package jsapi_test + +import ( + "os" + "strings" + "testing" + "time" + + v8runtime "github.com/yaoapp/gou/runtime/v8" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" + "github.com/yaoapp/yao/tai/volume" + "github.com/yaoapp/yao/test" + _ "github.com/yaoapp/yao/workspace/jsapi" +) + +type testMode struct { + Name string + Addr string // "local" or gRPC address +} + +func testModes() []testMode { + modes := []testMode{{Name: "local", Addr: "local"}} + if addr := os.Getenv("SANDBOX_TEST_REMOTE_ADDR"); addr != "" { + grpc := strings.TrimPrefix(addr, "tai://") + modes = append(modes, testMode{Name: "remote", Addr: grpc}) + } + return modes +} + +func setupForMode(t *testing.T, m testMode) { + t.Helper() + test.Prepare(t, config.Conf) + registry.Init(nil) + + var client *tai.Client + var err error + if m.Addr == "local" { + dataDir := t.TempDir() + vol := volume.NewLocal(dataDir) + client, err = tai.New("local", tai.WithVolume(vol), tai.WithDataDir(dataDir)) + } else { + client, err = tai.New(m.Addr) + } + if err != nil { + t.Fatalf("tai.New(%s): %v", m.Addr, err) + } + t.Cleanup(func() { client.Close() }) +} + +func setupGlobal(t *testing.T) { + t.Helper() + setupForMode(t, testMode{Name: "local", Addr: "local"}) +} + +func runJS(t *testing.T, source string) interface{} { + t.Helper() + opts := v8runtime.CallOptions{ + Sid: "test", + Timeout: 30 * time.Second, + } + res, err := v8runtime.Call(opts, source) + if err != nil { + t.Fatalf("JS error: %v", err) + } + return res +} + +func toInt(v interface{}) int { + switch n := v.(type) { + case int: + return n + case int32: + return int(n) + case int64: + return int(n) + case float64: + return int(n) + case float32: + return int(n) + default: + return 0 + } +} + +func TestWSCreateAndDelete(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSCreateAndDelete() { + var ws = workspace.Create({ name: "test-proj", owner: "u1", node: "local" }); + var id = ws.id; + workspace.Delete(id); + return id; + }`) + if res == nil || res == "" { + t.Error("expected workspace ID") + } + }) + } +} + +func TestWSGet(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSGet() { + var ws = workspace.Create({ name: "get-test", owner: "u1", node: "local" }); + var got = workspace.Get(ws.id); + var result = got ? got.id : "null"; + workspace.Delete(ws.id); + return result; + }`) + if res == "null" { + t.Error("Get returned null") + } + }) + } +} + +func TestWSGetNotFound(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSGetNotFound() { + var got = workspace.Get("ws-nonexistent"); + return got === null ? "null" : "found"; + }`) + if res != "null" { + t.Errorf("expected null, got %v", res) + } + }) + } +} + +func TestWSList(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSList() { + var ws1 = workspace.Create({ name: "list-a", owner: "u1", node: "local" }); + var ws2 = workspace.Create({ name: "list-b", owner: "u1", node: "local" }); + var list = workspace.List({ owner: "u1" }); + var count = list.length; + workspace.Delete(ws1.id); + workspace.Delete(ws2.id); + return count; + }`) + if toInt(res) < 2 { + t.Errorf("expected >= 2, got %v", res) + } + }) + } +} + +func TestWSReadWriteFile(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSReadWriteFile() { + var ws = workspace.Create({ name: "rw-test", owner: "u1", node: "local" }); + ws.WriteFile("hello.txt", "Hello, World!"); + var content = ws.ReadFile("hello.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "Hello, World!" { + t.Errorf("content = %v", res) + } + }) + } +} + +func TestWSReadDir(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSReadDir() { + var ws = workspace.Create({ name: "readdir", owner: "u1", node: "local" }); + ws.WriteFile("a.txt", "aaa"); + ws.MkdirAll("sub"); + ws.WriteFile("sub/b.txt", "bbb"); + var entries = ws.ReadDir("."); + workspace.Delete(ws.id); + return entries.length; + }`) + if toInt(res) < 2 { + t.Errorf("expected >= 2 entries, got %v", res) + } + }) + } +} + +func TestWSReadDirRecursive(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSReadDirRecursive() { + var ws = workspace.Create({ name: "readdir-r", owner: "u1", node: "local" }); + ws.WriteFile("a.txt", "aaa"); + ws.MkdirAll("sub/deep"); + ws.WriteFile("sub/b.txt", "bbb"); + ws.WriteFile("sub/deep/c.txt", "ccc"); + var entries = ws.ReadDir(".", true); + workspace.Delete(ws.id); + return entries.length; + }`) + if toInt(res) < 4 { + t.Errorf("expected >= 4 recursive, got %v", res) + } + }) + } +} + +func TestWSStat(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSStat() { + var ws = workspace.Create({ name: "stat-test", owner: "u1", node: "local" }); + ws.WriteFile("file.txt", "12345"); + var info = ws.Stat("file.txt"); + workspace.Delete(ws.id); + return info.size; + }`) + if toInt(res) != 5 { + t.Errorf("expected size 5, got %v", res) + } + }) + } +} + +func TestWSExistsIsDirIsFile(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSExistsIsDirIsFile() { + var ws = workspace.Create({ name: "checks", owner: "u1", node: "local" }); + ws.WriteFile("f.txt", "data"); + ws.MkdirAll("d"); + var r = [ + ws.Exists("f.txt"), + ws.Exists("nope"), + ws.IsFile("f.txt"), + ws.IsFile("d"), + ws.IsDir("d"), + ws.IsDir("f.txt") + ]; + workspace.Delete(ws.id); + return JSON.stringify(r); + }`) + if res != "[true,false,true,false,true,false]" { + t.Errorf("checks = %v", res) + } + }) + } +} + +func TestWSRemoveAndRename(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSRemoveAndRename() { + var ws = workspace.Create({ name: "ops", owner: "u1", node: "local" }); + ws.WriteFile("del.txt", "x"); + ws.Remove("del.txt"); + var a = ws.Exists("del.txt"); + + ws.MkdirAll("rmdir/sub"); + ws.WriteFile("rmdir/sub/f.txt", "x"); + ws.RemoveAll("rmdir"); + var b = ws.Exists("rmdir"); + + ws.WriteFile("old.txt", "x"); + ws.Rename("old.txt", "new.txt"); + var c = ws.Exists("old.txt"); + var d = ws.Exists("new.txt"); + + workspace.Delete(ws.id); + return JSON.stringify([a, b, c, d]); + }`) + if res != "[false,false,false,true]" { + t.Errorf("ops = %v", res) + } + }) + } +} + +func TestWSBase64(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSBase64() { + var ws = workspace.Create({ name: "b64", owner: "u1", node: "local" }); + ws.WriteFile("src.txt", "base64 test"); + var b64 = ws.ReadFileBase64("src.txt"); + ws.WriteFileBase64("dst.txt", b64); + var content = ws.ReadFile("dst.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "base64 test" { + t.Errorf("base64 roundtrip = %v", res) + } + }) + } +} + +func TestWSCopyInternal(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSCopyInternal() { + var ws = workspace.Create({ name: "copy", owner: "u1", node: "local" }); + ws.WriteFile("src.txt", "copy me"); + ws.Copy("src.txt", "dst.txt"); + var content = ws.ReadFile("dst.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "copy me" { + t.Errorf("copy = %v", res) + } + }) + } +} + +func TestWSCopyLocalToLocal(t *testing.T) { + setupGlobal(t) + + srcDir := t.TempDir() + dstDir := t.TempDir() + os.WriteFile(srcDir+"/test.txt", []byte("local-to-local"), 0o644) + + srcRel := srcDir[len(os.TempDir()):] + dstRel := dstDir[len(os.TempDir()):] + + runJS(t, `function TestWSCopyLocalToLocal() { + var ws = workspace.Create({ name: "l2l", owner: "u1", node: "local" }); + ws.Copy("tmp://`+srcRel+`", "tmp://`+dstRel+`"); + workspace.Delete(ws.id); + return "ok"; + }`) + + data, err := os.ReadFile(dstDir + "/test.txt") + if err != nil { + t.Fatalf("read dst: %v", err) + } + if string(data) != "local-to-local" { + t.Errorf("content = %s", data) + } +} + +func TestWSZipUnzip(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSZipUnzip() { + var ws = workspace.Create({ name: "zip", owner: "u1", node: "local" }); + ws.MkdirAll("src"); + ws.WriteFile("src/a.txt", "zip content"); + ws.WriteFile("src/b.txt", "more"); + var zr = ws.Zip("src", "out.zip"); + var ur = ws.Unzip("out.zip", "extracted"); + var content = ws.ReadFile("extracted/a.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "zip content" { + t.Errorf("zip/unzip = %v", res) + } + }) + } +} + +func TestWSGzipGunzip(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSGzipGunzip() { + var ws = workspace.Create({ name: "gzip", owner: "u1", node: "local" }); + ws.WriteFile("data.txt", "gzip test"); + ws.Gzip("data.txt", "data.txt.gz"); + ws.Gunzip("data.txt.gz", "restored.txt"); + var content = ws.ReadFile("restored.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "gzip test" { + t.Errorf("gzip/gunzip = %v", res) + } + }) + } +} + +func TestWSTarUntar(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSTarUntar() { + var ws = workspace.Create({ name: "tar", owner: "u1", node: "local" }); + ws.MkdirAll("src"); + ws.WriteFile("src/a.txt", "tar a"); + ws.Tar("src", "out.tar"); + ws.Untar("out.tar", "extracted"); + var content = ws.ReadFile("extracted/a.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "tar a" { + t.Errorf("tar/untar = %v", res) + } + }) + } +} + +func TestWSTgzUntgz(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSTgzUntgz() { + var ws = workspace.Create({ name: "tgz", owner: "u1", node: "local" }); + ws.MkdirAll("src"); + ws.WriteFile("src/x.txt", "tgz x"); + ws.Tgz("src", "out.tgz"); + ws.Untgz("out.tgz", "extracted"); + var content = ws.ReadFile("extracted/x.txt"); + workspace.Delete(ws.id); + return content; + }`) + if res != "tgz x" { + t.Errorf("tgz/untgz = %v", res) + } + }) + } +} + +func TestWSZipExcludes(t *testing.T) { + for _, m := range testModes() { + t.Run(m.Name, func(t *testing.T) { + setupForMode(t, m) + res := runJS(t, `function TestWSZipExcludes() { + var ws = workspace.Create({ name: "zip-exc", owner: "u1", node: "local" }); + ws.MkdirAll("src"); + ws.WriteFile("src/keep.txt", "keep"); + ws.WriteFile("src/skip.log", "skip"); + ws.Zip("src", "filtered.zip", { excludes: ["*.log"] }); + ws.Unzip("filtered.zip", "out"); + var hasKeep = ws.Exists("out/keep.txt"); + var hasSkip = ws.Exists("out/skip.log"); + workspace.Delete(ws.id); + return JSON.stringify([hasKeep, hasSkip]); + }`) + if res != "[true,false]" { + t.Errorf("zip excludes = %v", res) + } + }) + } +} diff --git a/workspace/manager.go b/workspace/manager.go index 8c4bb3dd..da56a902 100644 --- a/workspace/manager.go +++ b/workspace/manager.go @@ -4,26 +4,28 @@ import ( "context" "fmt" "os" - "sync" "time" "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" + "github.com/yaoapp/yao/tai/volume" taiworkspace "github.com/yaoapp/yao/tai/workspace" ) -// Manager owns workspace CRUD, file I/O, and node management. -// Pools are shared with sandbox.Manager — both reference the same tai.Client instances. -type Manager struct { - pools map[string]*tai.Client - mu sync.RWMutex +var mgr = NewManager() + +// M returns the global Manager. +func M() *Manager { + return mgr } -// NewManager creates a workspace manager with the given pools. -func NewManager(pools map[string]*tai.Client) *Manager { - if pools == nil { - pools = make(map[string]*tai.Client) - } - return &Manager{pools: pools} +// Manager owns workspace CRUD, file I/O, and node management. +// All node/client lookups go through tai.GetClient → registry. +type Manager struct{} + +// NewManager creates a workspace manager. +func NewManager() *Manager { + return &Manager{} } // Create allocates storage on the target node and persists metadata. @@ -32,9 +34,9 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Workspace, e return nil, ErrNodeMissing } - client, err := m.getClient(opts.Node) - if err != nil { - return nil, err + client, ok := tai.GetClient(opts.Node) + if !ok { + return nil, ErrNodeOffline } id := opts.ID @@ -71,18 +73,19 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Workspace, e } // Get returns a workspace by ID. -// If the node is unknown, scans all pools. +// Scans all registered nodes. func (m *Manager) Get(ctx context.Context, id string) (*Workspace, error) { - m.mu.RLock() - defer m.mu.RUnlock() - - for nodeName, client := range m.pools { - ws, err := m.readMeta(ctx, client, id) + for _, snap := range listNodes() { + client, ok := tai.GetClient(snap.TaiID) + if !ok { + continue + } + ws, err := readMeta(ctx, client, id) if err != nil { continue } if ws.Node == "" { - ws.Node = nodeName + ws.Node = snap.TaiID } return ws, nil } @@ -91,12 +94,13 @@ func (m *Manager) Get(ctx context.Context, id string) (*Workspace, error) { // List returns workspaces, optionally filtered by owner and/or node. func (m *Manager) List(ctx context.Context, opts ListOptions) ([]*Workspace, error) { - m.mu.RLock() - defer m.mu.RUnlock() - var result []*Workspace - for nodeName, client := range m.pools { - if opts.Node != "" && nodeName != opts.Node { + for _, snap := range listNodes() { + if opts.Node != "" && snap.TaiID != opts.Node { + continue + } + client, ok := tai.GetClient(snap.TaiID) + if !ok { continue } entries, err := client.Volume().ListDir(ctx, "", ".") @@ -107,12 +111,12 @@ func (m *Manager) List(ctx context.Context, opts ListOptions) ([]*Workspace, err if !e.IsDir { continue } - ws, err := m.readMeta(ctx, client, e.Path) + ws, err := readMeta(ctx, client, e.Path) if err != nil { continue } if ws.Node == "" { - ws.Node = nodeName + ws.Node = snap.TaiID } if opts.Owner != "" && ws.Owner != opts.Owner { continue @@ -163,28 +167,25 @@ func (m *Manager) Delete(ctx context.Context, id string, force bool) error { return nil } -// Nodes returns all configured Tai nodes with their online status. +// Nodes returns all registered Tai nodes with their online status. func (m *Manager) Nodes() []NodeInfo { - m.mu.RLock() - defer m.mu.RUnlock() - - nodes := make([]NodeInfo, 0, len(m.pools)) - for name := range m.pools { - nodes = append(nodes, NodeInfo{ - Name: name, - Online: true, + nodes := listNodes() + result := make([]NodeInfo, 0, len(nodes)) + for _, snap := range nodes { + result = append(result, NodeInfo{ + Name: snap.TaiID, + Online: snap.Status == "online" || snap.Status == "", }) } - return nodes + return result } // FS returns an fs.FS-compatible filesystem for the given workspace. func (m *Manager) FS(ctx context.Context, id string) (taiworkspace.FS, error) { - ws, client, err := m.resolve(ctx, id) + _, client, err := m.resolve(ctx, id) if err != nil { return nil, err } - _ = ws return client.Workspace(id), nil } @@ -237,18 +238,31 @@ func (m *Manager) Remove(ctx context.Context, id string, path string) error { return client.Volume().Remove(ctx, id, path, true) } -// AddPool registers a new Tai node. -func (m *Manager) AddPool(name string, client *tai.Client) { - m.mu.Lock() - defer m.mu.Unlock() - m.pools[name] = client +// Rename renames a file or directory within the workspace. +func (m *Manager) Rename(ctx context.Context, id string, oldPath, newPath string) error { + _, client, err := m.resolve(ctx, id) + if err != nil { + return err + } + return client.Volume().Rename(ctx, id, oldPath, newPath) } -// RemovePool unregisters a Tai node. -func (m *Manager) RemovePool(name string) { - m.mu.Lock() - defer m.mu.Unlock() - delete(m.pools, name) +// MkdirAll creates a directory (and parents) in the workspace. +func (m *Manager) MkdirAll(ctx context.Context, id string, path string) error { + _, client, err := m.resolve(ctx, id) + if err != nil { + return err + } + return client.Volume().MkdirAll(ctx, id, path) +} + +// Volume returns the Volume interface for the node hosting the given workspace. +func (m *Manager) Volume(ctx context.Context, id string) (volume.Volume, string, error) { + _, client, err := m.resolve(ctx, id) + if err != nil { + return nil, "", err + } + return client.Volume(), id, nil } // NodeForWorkspace returns the node name for a given workspace ID. @@ -263,7 +277,6 @@ func (m *Manager) NodeForWorkspace(ctx context.Context, id string) (string, erro // MountPath returns the host-side directory path for a workspace, // suitable for use as a Docker bind mount source. -// For local volumes this is dataDir/{id}; for remote (Tai) the server handles mounts. func (m *Manager) MountPath(ctx context.Context, id string) (string, error) { _, client, err := m.resolve(ctx, id) if err != nil { @@ -278,23 +291,14 @@ func (m *Manager) MountPath(ctx context.Context, id string) (string, error) { // --- internal --- -func (m *Manager) getClient(node string) (*tai.Client, error) { - m.mu.RLock() - defer m.mu.RUnlock() - client, ok := m.pools[node] - if !ok { - return nil, ErrNodeOffline - } - return client, nil -} - -// resolve finds the workspace and its tai.Client by scanning pools. +// resolve finds the workspace and its tai.Client by scanning all registered nodes. func (m *Manager) resolve(ctx context.Context, id string) (*Workspace, *tai.Client, error) { - m.mu.RLock() - defer m.mu.RUnlock() - - for _, client := range m.pools { - ws, err := m.readMeta(ctx, client, id) + for _, snap := range listNodes() { + client, ok := tai.GetClient(snap.TaiID) + if !ok { + continue + } + ws, err := readMeta(ctx, client, id) if err != nil { continue } @@ -303,7 +307,7 @@ func (m *Manager) resolve(ctx context.Context, id string) (*Workspace, *tai.Clie return nil, nil, ErrNotFound } -func (m *Manager) readMeta(ctx context.Context, client *tai.Client, id string) (*Workspace, error) { +func readMeta(ctx context.Context, client *tai.Client, id string) (*Workspace, error) { data, _, err := client.Volume().ReadFile(ctx, id, metadataFile) if err != nil { return nil, err @@ -311,6 +315,14 @@ func (m *Manager) readMeta(ctx context.Context, client *tai.Client, id string) ( return unmarshalMeta(data) } +func listNodes() []registry.NodeSnapshot { + reg := registry.Global() + if reg == nil { + return nil + } + return reg.List() +} + // DirEntry represents a file or directory entry in a workspace listing. type DirEntry struct { Name string `json:"name"` diff --git a/workspace/testutils_test.go b/workspace/testutils_test.go index 99b37a4f..591d3d44 100644 --- a/workspace/testutils_test.go +++ b/workspace/testutils_test.go @@ -2,11 +2,14 @@ package workspace_test import ( "context" + "net/url" "os" + "strings" "testing" "time" "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/tai/registry" "github.com/yaoapp/yao/tai/volume" "github.com/yaoapp/yao/workspace" ) @@ -21,19 +24,47 @@ func testPools() []poolConfig { {Name: "local", Addr: "local"}, } if addr := os.Getenv("SANDBOX_TEST_REMOTE_ADDR"); addr != "" { - pools = append(pools, poolConfig{Name: "remote", Addr: addr}) + name := taiIDFromAddr(addr) + pools = append(pools, poolConfig{Name: name, Addr: addr}) } return pools } -func setupManagerForPool(tb testing.TB, pc poolConfig) *workspace.Manager { - tb.Helper() - client := clientForPool(tb, pc) - pools := map[string]*tai.Client{pc.Name: client} - return workspace.NewManager(pools) +func taiIDFromAddr(addr string) string { + addr = strings.TrimSpace(addr) + if addr == "local" || addr == "" { + return "local" + } + if !strings.Contains(addr, "://") { + addr = "tai://" + addr + } + u, err := url.Parse(addr) + if err != nil { + return addr + } + h := u.Hostname() + if h == "" { + return addr + } + if p := u.Port(); p != "" { + return h + "-" + p + } + return h } -func clientForPool(tb testing.TB, pc poolConfig) *tai.Client { +func ensureRegistry(tb testing.TB) { + tb.Helper() + registry.Init(nil) +} + +func setupManagerForPool(tb testing.TB, pc poolConfig) *workspace.Manager { + tb.Helper() + ensureRegistry(tb) + registerClient(tb, pc) + return workspace.NewManager() +} + +func registerClient(tb testing.TB, pc poolConfig) *tai.Client { tb.Helper() if pc.Addr == "local" { return localClient(tb, tb.TempDir()) @@ -57,13 +88,25 @@ func localClient(tb testing.TB, dataDir string) *tai.Client { return client } -func setupManagerMultiNode(t *testing.T) *workspace.Manager { +func setupManagerMultiNode(t *testing.T) (*workspace.Manager, string, string) { t.Helper() - pools := map[string]*tai.Client{ - "node-a": localClient(t, t.TempDir()), - "node-b": localClient(t, t.TempDir()), + ensureRegistry(t) + + dir1 := t.TempDir() + vol1 := volume.NewLocal(dir1) + _, err := tai.New("docker://node-a", tai.WithVolume(vol1), tai.WithDataDir(dir1)) + if err != nil { + t.Fatalf("tai.New node-a: %v", err) } - return workspace.NewManager(pools) + + dir2 := t.TempDir() + vol2 := volume.NewLocal(dir2) + _, err = tai.New("docker://node-b", tai.WithVolume(vol2), tai.WithDataDir(dir2)) + if err != nil { + t.Fatalf("tai.New node-b: %v", err) + } + + return workspace.NewManager(), "docker://node-a", "docker://node-b" } func createWorkspace(tb testing.TB, m *workspace.Manager, node string, opts ...func(*workspace.CreateOptions)) *workspace.Workspace { diff --git a/workspace/workspace_test.go b/workspace/workspace_test.go index f5949cfd..f09e4451 100644 --- a/workspace/workspace_test.go +++ b/workspace/workspace_test.go @@ -7,8 +7,6 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/yaoapp/yao/tai" - "github.com/yaoapp/yao/tai/volume" "github.com/yaoapp/yao/workspace" ) @@ -157,18 +155,18 @@ func TestList_FilterOwner(t *testing.T) { } func TestList_FilterNode(t *testing.T) { - m := setupManagerMultiNode(t) + m, nodeA, nodeB := setupManagerMultiNode(t) ctx := context.Background() - _, err := m.Create(ctx, workspace.CreateOptions{Name: "a", Owner: "u", Node: "node-a"}) + _, err := m.Create(ctx, workspace.CreateOptions{Name: "a", Owner: "u", Node: nodeA}) require.NoError(t, err) - _, err = m.Create(ctx, workspace.CreateOptions{Name: "b", Owner: "u", Node: "node-b"}) + _, err = m.Create(ctx, workspace.CreateOptions{Name: "b", Owner: "u", Node: nodeB}) require.NoError(t, err) - list, err := m.List(ctx, workspace.ListOptions{Node: "node-a"}) + list, err := m.List(ctx, workspace.ListOptions{Node: nodeA}) require.NoError(t, err) assert.Len(t, list, 1) - assert.Equal(t, "node-a", list[0].Node) + assert.Equal(t, nodeA, list[0].Node) } func TestUpdate_Name(t *testing.T) { @@ -246,17 +244,16 @@ func TestDelete_NotFound(t *testing.T) { } func TestNodes(t *testing.T) { - m := setupManagerMultiNode(t) + m, nodeA, nodeB := setupManagerMultiNode(t) nodes := m.Nodes() - assert.Len(t, nodes, 2) + assert.GreaterOrEqual(t, len(nodes), 2) names := make(map[string]bool) for _, n := range nodes { names[n.Name] = true - assert.True(t, n.Online) } - assert.True(t, names["node-a"]) - assert.True(t, names["node-b"]) + assert.True(t, names[nodeA]) + assert.True(t, names[nodeB]) } func TestNodeForWorkspace(t *testing.T) { @@ -282,29 +279,17 @@ func TestNodeForWorkspace_NotFound(t *testing.T) { } } -func TestAddPool(t *testing.T) { - for _, pc := range testPools() { - t.Run(pc.Name, func(t *testing.T) { - m := setupManagerForPool(t, pc) - assert.Len(t, m.Nodes(), 1) +func TestRegistryDrivenNodes(t *testing.T) { + m, nodeA, nodeB := setupManagerMultiNode(t) + nodes := m.Nodes() + assert.GreaterOrEqual(t, len(nodes), 2) - vol := volume.NewLocal(t.TempDir()) - client, err := tai.New("local", tai.WithVolume(vol)) - require.NoError(t, err) - defer client.Close() - - m.AddPool("new-node", client) - assert.Len(t, m.Nodes(), 2) - }) + names := make(map[string]bool) + for _, n := range nodes { + names[n.Name] = true } -} - -func TestRemovePool(t *testing.T) { - m := setupManagerMultiNode(t) - assert.Len(t, m.Nodes(), 2) - - m.RemovePool("node-b") - assert.Len(t, m.Nodes(), 1) + assert.True(t, names[nodeA]) + assert.True(t, names[nodeB]) } func TestMountPath(t *testing.T) {