From 59004de3e8e6fda8901620072a939827c19d0c27 Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 13 Mar 2026 12:08:32 +0800 Subject: [PATCH 1/8] feat(grpc, openapi): enhance server keepalive settings and improve sorting in responses - Added keepalive parameters to the gRPC server configuration to manage connection health more effectively. - Implemented sorting logic in the computer and sandbox response handling to ensure consistent ordering by display name and last active timestamp. - Refactored node display name retrieval for better clarity and reuse across different functions. Made-with: Cursor --- grpc/grpc.go | 9 +++++++ openapi/computer/computer.go | 32 ++++++++++++---------- openapi/sandbox/manage.go | 5 ++-- openapi/workspace/workspace.go | 7 +++++ tai/tunnel/forward.go | 18 +++++++++---- tai/tunnel/grpc_handler.go | 49 +++++++++++++++++++++++++++------- 6 files changed, 90 insertions(+), 30 deletions(-) diff --git a/grpc/grpc.go b/grpc/grpc.go index ad088c06..08abe1ba 100644 --- a/grpc/grpc.go +++ b/grpc/grpc.go @@ -10,6 +10,7 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/codes" + "google.golang.org/grpc/keepalive" "google.golang.org/grpc/status" "github.com/yaoapp/kun/log" @@ -152,6 +153,14 @@ func StartServer(cfg config.Config) error { defer mu.Unlock() server = grpc.NewServer( + grpc.KeepaliveParams(keepalive.ServerParameters{ + Time: 30 * time.Second, + Timeout: 10 * time.Second, + }), + grpc.KeepaliveEnforcementPolicy(keepalive.EnforcementPolicy{ + MinTime: 15 * time.Second, + PermitWithoutStream: true, + }), grpc.ChainUnaryInterceptor(auth.UnaryInterceptor), grpc.ChainStreamInterceptor(auth.StreamInterceptor), ) diff --git a/openapi/computer/computer.go b/openapi/computer/computer.go index 4d29f432..fcd6420d 100644 --- a/openapi/computer/computer.go +++ b/openapi/computer/computer.go @@ -3,6 +3,7 @@ package computer import ( "context" "net/http" + "sort" "strconv" "strings" @@ -80,6 +81,9 @@ func handleOptions(c *gin.Context) { } snaps := reg.List() + sort.Slice(snaps, func(i, j int) bool { + return strings.ToLower(nodeDisplayName(snaps[i])) < strings.ToLower(nodeDisplayName(snaps[j])) + }) // Host entries: nodes with host_exec capability if kindFilter == "" || kindFilter == "host" { @@ -164,14 +168,18 @@ func matchNodeFilter(s *taitypes.NodeMeta, osFilter, archFilter string, minCPUs return true } +func nodeDisplayName(s taitypes.NodeMeta) string { + if s.DisplayName != "" { + return s.DisplayName + } + if s.System.Hostname != "" { + return s.System.Hostname + } + return s.TaiID +} + func nodeToHostOption(s taitypes.NodeMeta) computerOption { - displayName := s.DisplayName - if displayName == "" { - displayName = s.System.Hostname - } - if displayName == "" { - displayName = s.TaiID - } + displayName := nodeDisplayName(s) status := "stopped" if s.Status == "online" { @@ -195,6 +203,7 @@ func nodeToHostOption(s taitypes.NodeMeta) computerOption { Status: status, Mode: s.Mode, Addr: addr, + VNC: s.Ports.VNC > 0, System: computerSystemInfo{ OS: s.System.OS, Arch: s.System.Arch, @@ -206,13 +215,7 @@ func nodeToHostOption(s taitypes.NodeMeta) computerOption { } func nodeToNodeOption(s taitypes.NodeMeta) computerOption { - displayName := s.DisplayName - if displayName == "" { - displayName = s.System.Hostname - } - if displayName == "" { - displayName = s.TaiID - } + displayName := nodeDisplayName(s) status := "stopped" if s.Status == "online" { @@ -236,6 +239,7 @@ func nodeToNodeOption(s taitypes.NodeMeta) computerOption { Status: status, Mode: s.Mode, Addr: addr, + VNC: s.Ports.VNC > 0, System: computerSystemInfo{ OS: s.System.OS, Arch: s.System.Arch, diff --git a/openapi/sandbox/manage.go b/openapi/sandbox/manage.go index 84e17b19..71754f94 100644 --- a/openapi/sandbox/manage.go +++ b/openapi/sandbox/manage.go @@ -4,6 +4,7 @@ import ( "context" "net/http" "sort" + "strings" "time" "github.com/gin-gonic/gin" @@ -196,7 +197,7 @@ func hostToResponse(s taitypes.NodeMeta) sandboxResponse { Policy: "persistent", Mode: s.Mode, Addr: addr, - VNC: false, + VNC: s.Ports.VNC > 0, CreatedAt: s.ConnectedAt, LastActive: s.LastPing, System: sandboxSystemInfo{ @@ -287,7 +288,7 @@ func handleList(c *gin.Context) { } sort.Slice(result, func(i, j int) bool { - return result[i].LastActive.After(result[j].LastActive) + return strings.ToLower(result[i].DisplayName) < strings.ToLower(result[j].DisplayName) }) if result == nil { diff --git a/openapi/workspace/workspace.go b/openapi/workspace/workspace.go index d5b5d5dd..634fa7a7 100644 --- a/openapi/workspace/workspace.go +++ b/openapi/workspace/workspace.go @@ -8,6 +8,7 @@ import ( "mime" "net/http" "path/filepath" + "sort" "github.com/gin-gonic/gin" "github.com/yaoapp/yao/openapi/oauth/authorized" @@ -168,6 +169,9 @@ func handleList(c *gin.Context) { for _, w := range list { result = append(result, toResponse(w)) } + sort.Slice(result, func(i, j int) bool { + return result[i].CreatedAt > result[j].CreatedAt + }) response.RespondWithSuccess(c, http.StatusOK, result) } @@ -197,6 +201,9 @@ func handleOptions(c *gin.Context) { for _, w := range list { result = append(result, toResponse(w)) } + sort.Slice(result, func(i, j int) bool { + return result[i].CreatedAt > result[j].CreatedAt + }) response.RespondWithSuccess(c, http.StatusOK, result) } diff --git a/tai/tunnel/forward.go b/tai/tunnel/forward.go index 80b3f937..1713de90 100644 --- a/tai/tunnel/forward.go +++ b/tai/tunnel/forward.go @@ -2,6 +2,7 @@ package tunnel import ( "bytes" + "fmt" "io" "net/http" "strings" @@ -22,6 +23,7 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { reg := h.reg taiID := c.Param("taiID") + node, ok := reg.Get(taiID) if !ok || node.Status != "online" { c.JSON(http.StatusBadGateway, gin.H{"error": "tai node not available"}) @@ -34,6 +36,13 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { return } + rewrittenReq := rewriteRequest(c.Request, taiID) + logger.Debug("[forward] "+node.Mode+" → tai:"+fmt.Sprintf("%d", targetPort), + "tai_id", taiID, + "addr", node.Addr, + "path", rewrittenReq.URL.Path, + ) + hijacker, ok := c.Writer.(http.Hijacker) if !ok { c.JSON(http.StatusInternalServerError, gin.H{"error": "hijack not supported"}) @@ -41,21 +50,19 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { } browserConn, bufrw, err := hijacker.Hijack() if err != nil { - logger.Error("hijack failed", "err", err) + logger.Error("[forward] hijack failed", "tai_id", taiID, "err", err) return } defer browserConn.Close() fwd, err := h.RequestForward(taiID, targetPort) if err != nil { - logger.Error("request forward failed", + logger.Error("[forward] stream failed", "tai_id", taiID, "port", targetPort, "err", err) browserConn.Write([]byte("HTTP/1.1 502 Bad Gateway\r\n\r\n")) return } - rewrittenReq := rewriteRequest(c.Request, taiID) - var reqBuf bytes.Buffer rewrittenReq.Write(&reqBuf) if bufrw.Reader.Buffered() > 0 { @@ -63,7 +70,7 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { reqBuf.Write(buffered) } if err := fwd.Send(&taipb.ForwardData{Data: reqBuf.Bytes()}); err != nil { - logger.Error("send initial request", "err", err) + logger.Error("[forward] send failed", "tai_id", taiID, "err", err) return } @@ -72,6 +79,7 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { &netConnAdapter{ReadWriteCloser: browserConn}, streamConn, ) + logger.Debug("[forward] closed", "tai_id", taiID) } // HandleForwardLazy is a gin.HandlerFunc that resolves the global TunnelHandler diff --git a/tai/tunnel/grpc_handler.go b/tai/tunnel/grpc_handler.go index c7f74a48..cf8445e5 100644 --- a/tai/tunnel/grpc_handler.go +++ b/tai/tunnel/grpc_handler.go @@ -102,20 +102,51 @@ func (h *TunnelHandler) Register(stream taipb.TaiTunnel_RegisterServer) error { go h.connectTunnelNode(resolvedTaiID) + const pingTimeout = 90 * time.Second + recvCh := make(chan *taipb.TunnelControl) + errCh := make(chan error, 1) + go func() { + for { + ctrl, err := stream.Recv() + if err != nil { + errCh <- err + return + } + recvCh <- ctrl + } + }() + + timer := time.NewTimer(pingTimeout) + defer timer.Stop() + for { - ctrl, err := stream.Recv() - if err != nil { + select { + case ctrl := <-recvCh: + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(pingTimeout) + + switch ctrl.Type { + case "ping": + h.reg.UpdatePing(resolvedTaiID) + if err := stream.Send(&taipb.TunnelControl{Type: "pong"}); err != nil { + return err + } + } + + case err := <-errCh: if err == io.EOF { return nil } return err - } - switch ctrl.Type { - case "ping": - h.reg.UpdatePing(resolvedTaiID) - if err := stream.Send(&taipb.TunnelControl{Type: "pong"}); err != nil { - return err - } + + case <-timer.C: + h.logger.Warn("tai ping timeout, closing tunnel", "tai_id", resolvedTaiID, "timeout", pingTimeout) + return fmt.Errorf("tai %s: ping timeout (%s)", resolvedTaiID, pingTimeout) } } } From 82bb44cbdaacaf66b712aad97af02a029b1edf51 Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 13 Mar 2026 18:25:14 +0800 Subject: [PATCH 2/8] feat(workspace): add root directory retrieval for workspaces - Implemented a new endpoint to retrieve the absolute path of a workspace's root directory. - Enhanced the workspace interface with a GetRoot method to facilitate this functionality. - Updated the workspace manager to utilize the new method for improved path resolution. Made-with: Cursor --- agent/assistant/agent.go | 9 ++ agent/assistant/sandbox_v2.go | 22 +++ agent/context/jsapi_computer.go | 11 ++ agent/sandbox/v2/claude/oscompat.go | 164 ++++++++++++++++++++++ agent/sandbox/v2/claude/runner.go | 198 ++++++++++++-------------- agent/sandbox/v2/lifecycle.go | 103 ++++++++++++++ agent/sandbox/v2/prepare.go | 11 +- openapi/workspace/workspace.go | 17 +++ sandbox/v2/box.go | 9 ++ sandbox/v2/host.go | 15 ++ sandbox/v2/manager.go | 7 + sandbox/v2/types.go | 1 + tai/tunnel/grpc_handler.go | 49 +++++-- tai/volume/local.go | 4 + tai/volume/mock_test.go | 47 +++++++ tai/volume/pb/volume.pb.go | 206 +++++++++++++++++----------- tai/volume/pb/volume.proto | 6 + tai/volume/pb/volume_grpc.pb.go | 40 ++++++ tai/volume/remote.go | 11 ++ tai/volume/volume.go | 1 + tai/volume/volume_test.go | 73 ++++++++++ tai/workspace/workspace.go | 7 + tai/workspace/workspace_test.go | 19 +++ workspace/manager.go | 15 +- 24 files changed, 834 insertions(+), 211 deletions(-) create mode 100644 agent/sandbox/v2/claude/oscompat.go diff --git a/agent/assistant/agent.go b/agent/assistant/agent.go index 9e49226b..5e4583a8 100644 --- a/agent/assistant/agent.go +++ b/agent/assistant/agent.go @@ -204,6 +204,15 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa } }() + // ================================================ + // Standalone Workspace Loading (no sandbox required) + // ================================================ + // When no sandbox is configured but the user selected a workspace, + // load the workspace FS into context so hooks can access ctx.workspace. + if !ctx.HasWorkspace() { + ast.initStandaloneWorkspace(ctx) + } + // ================================================ // Execute Create Hook // ================================================ diff --git a/agent/assistant/sandbox_v2.go b/agent/assistant/sandbox_v2.go index 4a58b4cc..a63ba159 100644 --- a/agent/assistant/sandbox_v2.go +++ b/agent/assistant/sandbox_v2.go @@ -2,6 +2,7 @@ package assistant import ( "fmt" + "log" "os" "path/filepath" @@ -13,6 +14,7 @@ import ( "github.com/yaoapp/yao/config" infraV2 "github.com/yaoapp/yao/sandbox/v2" traceTypes "github.com/yaoapp/yao/trace/types" + "github.com/yaoapp/yao/workspace" ) // HasSandboxV2 returns true if the assistant has a V2 sandbox configuration. @@ -167,6 +169,26 @@ func (ast *Assistant) executeSandboxV2Stream( return sandboxv2.ExecuteSandboxStream(ctx, execReq, streamHandler) } +// initStandaloneWorkspace loads the workspace FS into context when no sandbox +// is configured but the user selected a workspace (metadata["workspace_id"]). +func (ast *Assistant) initStandaloneWorkspace(ctx *context.Context) { + if ctx.Metadata == nil { + return + } + wsID, _ := ctx.Metadata["workspace_id"].(string) + if wsID == "" { + return + } + + stdCtx := ctx.Context + wsFS, err := workspace.M().FS(stdCtx, wsID) + if err != nil { + log.Printf("[assistant] initStandaloneWorkspace: failed to load workspace %s: %v", wsID, err) + return + } + ctx.SetWorkspace(wsFS) +} + func closeLoadingV2(ctx *context.Context, loadingMsgID, msgKey string) { if loadingMsgID == "" || ctx == nil { return diff --git a/agent/context/jsapi_computer.go b/agent/context/jsapi_computer.go index ae3c57c1..795fd2db 100644 --- a/agent/context/jsapi_computer.go +++ b/agent/context/jsapi_computer.go @@ -20,6 +20,12 @@ func (ctx *Context) SetComputer(computer infraV2.Computer) { } } +// SetWorkspace sets the workspace FS directly without requiring a Computer. +// Use this when the user selected a workspace but no sandbox is configured. +func (ctx *Context) SetWorkspace(ws workspace.FS) { + ctx.workspace = ws +} + // GetComputer returns the V2 computer if available. func (ctx *Context) GetComputer() infraV2.Computer { return ctx.computer @@ -35,6 +41,11 @@ func (ctx *Context) HasComputer() bool { return ctx.computer != nil } +// HasWorkspace returns true if workspace FS is available. +func (ctx *Context) HasWorkspace() bool { + return ctx.workspace != nil +} + // createComputerInstance creates the ctx.computer JavaScript object. func (ctx *Context) createComputerInstance(v8ctx *v8go.Context) *v8go.Value { if ctx.computer == nil { diff --git a/agent/sandbox/v2/claude/oscompat.go b/agent/sandbox/v2/claude/oscompat.go new file mode 100644 index 00000000..ba0f24d7 --- /dev/null +++ b/agent/sandbox/v2/claude/oscompat.go @@ -0,0 +1,164 @@ +package claude + +import ( + "fmt" + "path" + "strings" + + "github.com/yaoapp/yao/agent/sandbox/v2/types" + infra "github.com/yaoapp/yao/sandbox/v2" +) + +// osEnv captures OS-dependent paths and shell settings derived from the +// Computer's SystemInfo. All runner code should use osEnv instead of +// hardcoded Linux constants. +type osEnv struct { + OS string // "windows", "linux", "darwin", ... + Shell string // preferred shell binary: "bash", "pwsh", "cmd.exe", ... + WorkDir string // working directory on the target machine + UserHome string // user home directory (empty if irrelevant) + TempDir string // system temp directory +} + +func (e *osEnv) isWindows() bool { + return strings.EqualFold(e.OS, "windows") +} + +// resolveOSEnv builds an osEnv from the Computer's reported SystemInfo, +// falling back to SandboxConfig values where available, then to per-OS defaults. +func resolveOSEnv(computer infra.Computer, _ *types.SandboxConfig) *osEnv { + sys := computer.ComputerInfo().System + + env := &osEnv{ + OS: strings.ToLower(sys.OS), + Shell: sys.Shell, + TempDir: sys.TempDir, + WorkDir: computer.GetWorkDir(), + } + + if env.TempDir == "" { + env.TempDir = env.pathJoin(env.WorkDir, ".tmp") + } + + return env +} + +// shellCmd returns the command slice to run a script through the appropriate shell. +func (e *osEnv) shellCmd(script string) []string { + shell := strings.ToLower(e.Shell) + switch shell { + case "pwsh": + return []string{"pwsh", "-NoProfile", "-Command", script} + case "powershell": + return []string{"powershell", "-NoProfile", "-Command", script} + case "cmd.exe", "cmd": + return []string{"cmd.exe", "/C", script} + default: + return []string{"bash", "-c", script} + } +} + +// mkdirCmd returns a shell command string to create a directory (with parents). +func (e *osEnv) mkdirCmd(dir string) string { + if e.isWindows() { + return fmt.Sprintf(`if (!(Test-Path '%s')) { New-Item -ItemType Directory -Path '%s' -Force | Out-Null }`, dir, dir) + } + return fmt.Sprintf("mkdir -p %s", dir) +} + +// listDirCmd returns a command slice to list directory contents. +func (e *osEnv) listDirCmd(dir string) []string { + if e.isWindows() { + return e.shellCmd(fmt.Sprintf("Get-ChildItem -Name '%s'", dir)) + } + return []string{"ls", dir} +} + +// killProcessCmd returns a command slice to kill processes matching a pattern. +func (e *osEnv) killProcessCmd(pattern string) []string { + if e.isWindows() { + script := fmt.Sprintf("Get-Process | Where-Object {$_.ProcessName -like '*%s*'} | Stop-Process -Force -ErrorAction SilentlyContinue", pattern) + return e.shellCmd(script) + } + return []string{"sh", "-c", fmt.Sprintf("pkill -f '%s' || true", pattern)} +} + +// rootDir returns the filesystem root for the target OS. +func (e *osEnv) rootDir() string { + if e.isWindows() { + return `C:\` + } + return "/" +} + +// pathJoin joins path segments using the appropriate separator. +func (e *osEnv) pathJoin(parts ...string) string { + if e.isWindows() { + return strings.Join(parts, `\`) + } + return path.Join(parts...) +} + +// buildCLIScript builds the complete CLI invocation script for the target OS. +// Returns (script, stdin) — on Linux stdin is nil (heredoc handles it), +// on Windows stdin contains inputJSONL bytes to pass via gRPC Stdin. +func (e *osEnv) buildCLIScript(args []string, systemPrompt, inputJSONL string) (string, []byte) { + workDir := e.WorkDir + promptFile := e.pathJoin(workDir, ".yao", ".system-prompt.txt") + + if e.isWindows() { + return e.buildPowerShellScript(args, systemPrompt, inputJSONL, workDir, promptFile) + } + return e.buildBashScript(args, systemPrompt, inputJSONL, workDir, promptFile), nil +} + +func (e *osEnv) buildBashScript(args []string, systemPrompt, inputJSONL, workDir, promptFile string) string { + var b strings.Builder + + if e.UserHome != "" { + b.WriteString(fmt.Sprintf("touch %s/.Xauthority 2>/dev/null; ", e.UserHome)) + } + b.WriteString("touch \"$HOME/.Xauthority\" 2>/dev/null\n") + + if systemPrompt != "" { + b.WriteString(fmt.Sprintf("mkdir -p %s/.yao\n", workDir)) + b.WriteString(fmt.Sprintf("cat << 'PROMPTEOF' > %s\n", promptFile)) + b.WriteString(systemPrompt) + b.WriteString("\nPROMPTEOF\n") + args = append(args, "--append-system-prompt-file", promptFile) + } + + b.WriteString("cat << 'INPUTEOF' | claude -p") + for _, arg := range args { + b.WriteString(fmt.Sprintf(" %q", arg)) + } + b.WriteString(" 2>&1\n") + b.WriteString(inputJSONL) + b.WriteString("\nINPUTEOF") + + return b.String() +} + +// buildPowerShellScript builds a script that writes the system prompt file, +// then launches claude -p. inputJSONL is returned as stdin bytes to be passed +// directly via gRPC, bypassing PowerShell's encoding entirely. +func (e *osEnv) buildPowerShellScript(args []string, systemPrompt, inputJSONL, workDir, promptFile string) (string, []byte) { + var b strings.Builder + noBOM := "(New-Object System.Text.UTF8Encoding $false)" + + yaoDir := e.pathJoin(workDir, ".yao") + b.WriteString(fmt.Sprintf("if (!(Test-Path '%s')) { New-Item -ItemType Directory -Path '%s' -Force | Out-Null }\n", yaoDir, yaoDir)) + + if systemPrompt != "" { + escaped := strings.ReplaceAll(systemPrompt, "'", "''") + b.WriteString(fmt.Sprintf("[IO.File]::WriteAllText('%s', @'\n%s\n'@, %s)\n", promptFile, escaped, noBOM)) + args = append(args, "--append-system-prompt-file", promptFile) + } + + b.WriteString("claude -p") + for _, arg := range args { + b.WriteString(fmt.Sprintf(" '%s'", strings.ReplaceAll(arg, "'", "''"))) + } + + return b.String(), []byte(inputJSONL + "\n") +} diff --git a/agent/sandbox/v2/claude/runner.go b/agent/sandbox/v2/claude/runner.go index fac6c522..886fde38 100644 --- a/agent/sandbox/v2/claude/runner.go +++ b/agent/sandbox/v2/claude/runner.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "io" "path" "strings" "time" @@ -14,12 +15,7 @@ import ( infra "github.com/yaoapp/yao/sandbox/v2" ) -const ( - defaultWorkDir = "/workspace" - defaultUser = "sandbox" - defaultUserHome = "/home/sandbox" - defaultProxyPort = 3456 -) +const defaultProxyPort = 3456 // ClaudeRunner implements the Runner interface for Claude CLI (mode=cli). type ClaudeRunner struct { @@ -45,33 +41,30 @@ func (r *ClaudeRunner) Prepare(ctx context.Context, req *types.PrepareRequest) e r.mode = "cli" } - workDir := resolveWorkDir(req.Config) + env := resolveOSEnv(req.Computer, 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 != "" { + claudeDir := env.pathJoin(env.WorkDir, ".claude") steps = append(steps, types.PrepareStep{ Action: "exec", - Cmd: fmt.Sprintf("mkdir -p %s/.claude", workDir), + Cmd: env.mkdirCmd(claudeDir), 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"), + Path: env.pathJoin(env.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) @@ -88,9 +81,8 @@ func (r *ClaudeRunner) Stream(ctx context.Context, req *types.StreamRequest, han return fmt.Errorf("computer is nil") } - workDir := resolveWorkDir(req.Config) + oe := resolveOSEnv(computer, req.Config) - // Prepare attachments: resolve __yao.attachment:// URLs, copy files to workspace. if req.ChatID != "" { ws := computer.Workplace() if ws != nil { @@ -102,90 +94,118 @@ func (r *ClaudeRunner) Stream(ctx context.Context, req *types.StreamRequest, han } } - // Detect continuation (existing .claude/projects/ directory). - isContinuation := hasExistingSession(ctx, computer, workDir) + isContinuation := hasExistingSession(ctx, computer, oe) - // Build CLI command and env. - cmd, env := r.buildCLICommand(req, isContinuation) + cmd, env, stdin := r.buildCLICommand(req, oe, isContinuation) - // Create stream. - execStream, err := computer.Stream(ctx, cmd, infra.WithWorkDir(workDir), infra.WithEnv(env)) + streamOpts := []infra.ExecOption{infra.WithWorkDir(oe.WorkDir), infra.WithEnv(env)} + if len(stdin) > 0 { + streamOpts = append(streamOpts, infra.WithStdin(stdin)) + } + + execStream, err := computer.Stream(ctx, cmd, streamOpts...) 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) - }() + streamCtx, streamCancel := context.WithCancel(ctx) + defer streamCancel() 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: + <-streamCtx.Done() + killCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + computer.Exec(killCtx, oe.killProcessCmd("claude")) + execStream.Cancel() + }() + + var stderrBuf strings.Builder + go func() { + buf := make([]byte, 4096) + for { + n, err := execStream.Stderr.Read(buf) + if n > 0 { + stderrBuf.Write(buf[:n]) + chunk := string(buf[:n]) + if strings.Contains(strings.ToLower(chunk), "error") { + streamCancel() + io.Copy(&stderrBuf, execStream.Stderr) + return + } + } + if err != nil { + return + } } }() - // Parse streaming output. - parseErr := parseStreamJSON(ctx, execStream.Stdout, handler) + parseErr := parseStreamJSON(streamCtx, execStream.Stdout, handler) - // Wait for process exit. exitCode, waitErr := execStream.Wait() + stderrStr := strings.TrimSpace(stderrBuf.String()) + if parseErr != nil { + if stderrStr != "" { + return fmt.Errorf("%w (stderr: %s)", parseErr, stderrStr) + } return parseErr } if waitErr != nil { + if stderrStr != "" { + return fmt.Errorf("%w (stderr: %s)", waitErr, stderrStr) + } return waitErr } if exitCode != 0 { + if stderrStr != "" { + return fmt.Errorf("claude CLI exited with code %d: %s", exitCode, stderrStr) + } 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"}) + oe := resolveOSEnv(computer, nil) + computer.Exec(ctx, oe.killProcessCmd("claude")) } 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}) +func hasExistingSession(ctx context.Context, computer infra.Computer, oe *osEnv) bool { + sessionDir := oe.pathJoin(oe.WorkDir, ".claude", "projects") + result, err := computer.Exec(ctx, oe.listDirCmd(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) - +// buildCLICommand constructs the Claude CLI command, environment variables, and optional stdin bytes. +func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, oe *osEnv, isContinuation bool) ([]string, map[string]string, []byte) { 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") + if oe.isWindows() { + env["USERPROFILE"] = oe.WorkDir + if len(oe.WorkDir) >= 2 && oe.WorkDir[1] == ':' { + env["HOMEDRIVE"] = oe.WorkDir[:2] + env["HOMEPATH"] = oe.WorkDir[2:] + } + } else { + env["HOME"] = oe.WorkDir + if oe.UserHome != "" { + env["XAUTHORITY"] = path.Join(oe.UserHome, ".Xauthority") + } } - // Connector environment. if req.Connector != nil { setting := req.Connector.Setting() host, _ := setting["host"].(string) @@ -209,23 +229,20 @@ func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, isContinuation } } - // 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) + envPrompt := buildSandboxEnvPrompt(oe.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) @@ -233,10 +250,19 @@ func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, isContinuation inputJSONL = buildFirstRequestJSONL(req.Messages) } - // CLI args. var args []string - args = append(args, "--dangerously-skip-permissions") - args = append(args, "--permission-mode", "bypassPermissions") + + permMode := "" + if req.Config != nil && req.Config.Runner.Options != nil { + if v, ok := req.Config.Runner.Options["permission_mode"]; ok { + permMode = fmt.Sprintf("%v", v) + } + } + if permMode == "bypassPermissions" { + args = append(args, "--dangerously-skip-permissions") + args = append(args, "--permission-mode", permMode) + } + args = append(args, "--input-format", "stream-json") args = append(args, "--output-format", "stream-json") args = append(args, "--include-partial-messages") @@ -246,7 +272,6 @@ func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, 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 { @@ -255,39 +280,16 @@ func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, isContinuation } } - // MCP config (set by Prepare if MCPServers were present). if r.hasMCP { - args = append(args, "--mcp-config", path.Join(workDir, ".mcp.json")) + mcpPath := oe.pathJoin(oe.WorkDir, ".mcp.json") + args = append(args, "--mcp-config", mcpPath) 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 + script, stdin := oe.buildCLIScript(args, systemPrompt, inputJSONL) + return oe.shellCmd(script), env, stdin } // buildMCPConfig creates the .mcp.json for Claude CLI based on declared servers. @@ -374,30 +376,6 @@ When working with GitHub and a token is provided: `, 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", diff --git a/agent/sandbox/v2/lifecycle.go b/agent/sandbox/v2/lifecycle.go index afe8604e..d2f28800 100644 --- a/agent/sandbox/v2/lifecycle.go +++ b/agent/sandbox/v2/lifecycle.go @@ -11,6 +11,8 @@ import ( agentContext "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/sandbox/v2/types" infra "github.com/yaoapp/yao/sandbox/v2" + "github.com/yaoapp/yao/tai" + "github.com/yaoapp/yao/workspace" ) // BuildIdentifier determines the Computer identifier based on lifecycle policy @@ -59,6 +61,98 @@ func GetComputer(ctx *agentContext.Context, cfg *types.SandboxConfig, manager *i } cfg.WorkspaceID = workspaceID + // Resolve computer_id from metadata to determine kind and nodeID. + computerID := "" + if ctx.Metadata != nil { + if cid, ok := ctx.Metadata["computer_id"].(string); ok && cid != "" { + computerID = cid + } + } + + // Workspace-wins rule: when both workspace_id and computer_id are present, + // the workspace's bound node takes precedence over computer_id. + if workspaceID != "" && workspaceID != ownerID { + wsNode, err := workspace.M().NodeForWorkspace(context.Background(), workspaceID) + if err == nil && wsNode != "" { + if computerID != "" && computerID != wsNode { + log.Printf("[sandbox/v2] workspace %s bound to node %s overrides computer_id %s", workspaceID, wsNode, computerID) + } + computerID = wsNode + } + } + + if computerID != "" { + return resolveComputerByID(cfg, manager, computerID, ownerID, identifier, workspaceID, conn...) + } + + // No computer_id: fall back to DSL-based dispatch (original logic). + return resolveComputerByDSL(cfg, manager, ownerID, identifier, workspaceID, conn...) +} + +// resolveComputerByID dispatches based on the runtime computer_id from metadata. +// It queries the registry and sandbox manager to determine the computer kind. +func resolveComputerByID( + cfg *types.SandboxConfig, manager *infra.Manager, + computerID, ownerID, identifier, workspaceID string, + conn ...connector.Connector, +) (infra.Computer, string, error) { + + // 1) Check if computer_id is a known Tai node (host or node kind). + if node, ok := tai.GetNodeMeta(computerID); ok { + cfg.NodeID = computerID + hasContainerRuntime := node.Capabilities.Docker || node.Capabilities.K8s + + if node.Capabilities.HostExec && !hasContainerRuntime { + // Host-only node: must use host mode regardless of DSL image config. + cfg.Kind = "host" + host, err := manager.Host(context.Background(), computerID) + if err != nil { + return nil, identifier, fmt.Errorf("get host computer: %w", err) + } + host.BindWorkplace(workspaceID) + return host, identifier, nil + } + + if node.Capabilities.HostExec && hasContainerRuntime && cfg.Computer.Image == "" { + // Dual-capable node with no image in DSL: prefer host mode. + cfg.Kind = "host" + host, err := manager.Host(context.Background(), computerID) + if err != nil { + return nil, identifier, fmt.Errorf("get host computer: %w", err) + } + host.BindWorkplace(workspaceID) + return host, identifier, nil + } + + if !hasContainerRuntime { + return nil, identifier, fmt.Errorf("node %q has no container runtime and no host_exec capability", computerID) + } + + // Node with container runtime and DSL has image: create/reuse a box. + cfg.Kind = "box" + return resolveBox(cfg, manager, ownerID, identifier, workspaceID, conn...) + } + + // 2) Check if computer_id is an existing box ID. + if manager != nil { + box, err := manager.Get(context.Background(), computerID) + if err == nil && box != nil { + cfg.Kind = "box" + box.BindWorkplace(workspaceID) + return box, computerID, nil + } + } + + return nil, identifier, fmt.Errorf("computer %q not found in registry or sandbox manager", computerID) +} + +// resolveComputerByDSL dispatches based on DSL static configuration (cfg.Computer.Image). +func resolveComputerByDSL( + cfg *types.SandboxConfig, manager *infra.Manager, + ownerID, identifier, workspaceID string, + conn ...connector.Connector, +) (infra.Computer, string, error) { + // Host mode: no image → host computer. if cfg.Computer.Image == "" { cfg.Kind = "host" @@ -75,6 +169,15 @@ func GetComputer(ctx *agentContext.Context, cfg *types.SandboxConfig, manager *i } cfg.Kind = "box" + return resolveBox(cfg, manager, ownerID, identifier, workspaceID, conn...) +} + +// resolveBox reuses or creates a box container. +func resolveBox( + cfg *types.SandboxConfig, manager *infra.Manager, + ownerID, identifier, workspaceID string, + conn ...connector.Connector, +) (infra.Computer, string, error) { // Reuse: non-empty identifier → try Get first. if identifier != "" { diff --git a/agent/sandbox/v2/prepare.go b/agent/sandbox/v2/prepare.go index 26e50843..ae8951b1 100644 --- a/agent/sandbox/v2/prepare.go +++ b/agent/sandbox/v2/prepare.go @@ -144,7 +144,12 @@ func runExecStep(ctx context.Context, computer infra.Computer, step types.Prepar } } - result, err := computer.Exec(ctx, shellWrap(kind, script), infra.WithWorkDir("/")) + rootDir := "/" + if isWindowsComputer(computer) { + rootDir = `C:\` + } + + result, err := computer.Exec(ctx, shellWrap(kind, script), infra.WithWorkDir(rootDir)) if err != nil { return err } @@ -155,6 +160,10 @@ func runExecStep(ctx context.Context, computer infra.Computer, step types.Prepar return checkResult(result, label) } +func isWindowsComputer(computer infra.Computer) bool { + return strings.EqualFold(computer.ComputerInfo().System.OS, "windows") +} + // checkResult inspects ExecResult for errors. func checkResult(result *infra.ExecResult, label string) error { if result.Error != "" { diff --git a/openapi/workspace/workspace.go b/openapi/workspace/workspace.go index 634fa7a7..f3eade52 100644 --- a/openapi/workspace/workspace.go +++ b/openapi/workspace/workspace.go @@ -29,6 +29,7 @@ import ( // - DELETE /:id/files/*path — delete file // - POST /:id/mkdir — create directory // - POST /:id/rename — rename file/directory +// - GET /:id/rootdir — get workspace root directory absolute path func Attach(group *gin.RouterGroup, oauth types.OAuth) { group.Use(oauth.Guard) @@ -39,6 +40,7 @@ func Attach(group *gin.RouterGroup, oauth types.OAuth) { group.PUT("/:id", handleUpdate) group.DELETE("/:id", handleDelete) + group.GET("/:id/rootdir", handleRootDir) group.GET("/:id/files", handleListFiles) group.GET("/:id/files/*path", handleReadFile) group.PUT("/:id/files/*path", handleWriteFile) @@ -285,6 +287,21 @@ func handleDelete(c *gin.Context) { c.Status(http.StatusNoContent) } +func handleRootDir(c *gin.Context) { + _, ok := resolveAndCheckWS(c) + if !ok { + return + } + + rootDir, err := mgr().MountPath(context.Background(), c.Param("id")) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + response.RespondWithSuccess(c, http.StatusOK, gin.H{"root_dir": rootDir}) +} + func handleListFiles(c *gin.Context) { _, ok := resolveAndCheckWS(c) if !ok { diff --git a/sandbox/v2/box.go b/sandbox/v2/box.go index 19ba2fe2..aed45d11 100644 --- a/sandbox/v2/box.go +++ b/sandbox/v2/box.go @@ -30,6 +30,7 @@ type Box struct { image string workspaceID string system SystemInfo + workDir string ws taiworkspace.FS manager *Manager } @@ -195,6 +196,14 @@ func (b *Box) Workspace() taiworkspace.FS { return b.ws } +// GetWorkDir returns the container-internal working directory for command execution. +func (b *Box) GetWorkDir() string { + if b.workDir != "" { + return b.workDir + } + return "/workspace" +} + // WorkspaceID returns the workspace ID mounted to this sandbox, or empty string. func (b *Box) WorkspaceID() string { return b.workspaceID } diff --git a/sandbox/v2/host.go b/sandbox/v2/host.go index 8023a308..1fd77921 100644 --- a/sandbox/v2/host.go +++ b/sandbox/v2/host.go @@ -233,6 +233,21 @@ func (h *Host) Workplace() taiworkspace.FS { return taiworkspace.New(res.Volume, h.workplaceID) } +// GetWorkDir returns the host working directory for command execution. +// Resolves from the bound workspace's root path on disk, falling back to +// the system temp directory if no workspace is bound or root resolution fails. +func (h *Host) GetWorkDir() string { + if ws := h.Workplace(); ws != nil { + if root, err := ws.GetRoot(); err == nil && root != "" { + return root + } + } + if h.system.TempDir != "" { + return h.system.TempDir + } + return "/tmp" +} + // NodeID returns the node ID this Host belongs to. func (h *Host) NodeID() string { return h.nodeID } diff --git a/sandbox/v2/manager.go b/sandbox/v2/manager.go index 11d83c99..7a271257 100644 --- a/sandbox/v2/manager.go +++ b/sandbox/v2/manager.go @@ -199,6 +199,11 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) TempDir: res.System.TempDir, } + boxWorkDir := opts.WorkDir + if boxWorkDir == "" { + boxWorkDir = "/workspace" + } + box := &Box{ id: id, containerID: containerID, @@ -214,6 +219,7 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) vnc: opts.VNC, image: opts.Image, workspaceID: opts.WorkspaceID, + workDir: boxWorkDir, system: sys, } box.lastCall.Store(time.Now().UnixMilli()) @@ -459,6 +465,7 @@ func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, res *tai.Conn createdAt: time.Now(), image: c.Image, workspaceID: c.Labels["workspace-id"], + workDir: "/workspace", manager: m, } box.lastCall.Store(time.Now().UnixMilli()) diff --git a/sandbox/v2/types.go b/sandbox/v2/types.go index 1d4a26d4..bc78c85a 100644 --- a/sandbox/v2/types.go +++ b/sandbox/v2/types.go @@ -22,6 +22,7 @@ type Computer interface { Proxy(ctx context.Context, port int, path string) (string, error) BindWorkplace(workspaceID string) Workplace() workspace.FS + GetWorkDir() string } // ComputerInfo holds identity and registry information for a Computer. diff --git a/tai/tunnel/grpc_handler.go b/tai/tunnel/grpc_handler.go index cf8445e5..57d8b93f 100644 --- a/tai/tunnel/grpc_handler.go +++ b/tai/tunnel/grpc_handler.go @@ -30,6 +30,8 @@ type TunnelHandler struct { reg *registry.Registry pending sync.Map // channel_id → chan taipb.TaiTunnel_ForwardServer logger *slog.Logger + + sendMu sync.Map // taiID → *sync.Mutex – serializes Send on each Register stream } // NewTunnelHandler creates a TunnelHandler backed by the given registry. @@ -84,17 +86,24 @@ func (h *TunnelHandler) Register(stream taipb.TaiTunnel_RegisterServer) error { Capabilities: capsFromProto(msg.Caps), } + var mu sync.Mutex + h.sendMu.Store(resolvedTaiID, &mu) + h.reg.Register(node) h.reg.SetRegisterStream(resolvedTaiID, stream) defer func() { + h.sendMu.Delete(resolvedTaiID) h.reg.Unregister(resolvedTaiID) h.logger.Info("tai gRPC tunnel disconnected", "tai_id", resolvedTaiID) }() - if err := stream.Send(&taipb.TunnelControl{ + mu.Lock() + err = stream.Send(&taipb.TunnelControl{ Type: "registered", TaiId: resolvedTaiID, - }); err != nil { + }) + mu.Unlock() + if err != nil { return fmt.Errorf("send registered: %w", err) } @@ -133,8 +142,11 @@ func (h *TunnelHandler) Register(stream taipb.TaiTunnel_RegisterServer) error { switch ctrl.Type { case "ping": h.reg.UpdatePing(resolvedTaiID) - if err := stream.Send(&taipb.TunnelControl{Type: "pong"}); err != nil { - return err + mu.Lock() + sendErr := stream.Send(&taipb.TunnelControl{Type: "pong"}) + mu.Unlock() + if sendErr != nil { + return sendErr } } @@ -163,9 +175,12 @@ func (h *TunnelHandler) Forward(stream taipb.TaiTunnel_ForwardServer) error { } channelID := vals[0] + h.logger.Debug("[forward] Forward stream arrived", "channel_id", channelID[:16]) + if ch, ok := h.pending.LoadAndDelete(channelID); ok { ch.(chan taipb.TaiTunnel_ForwardServer) <- stream } else { + h.logger.Warn("[forward] no pending channel (expired?)", "channel_id", channelID[:16]) return fmt.Errorf("no pending channel for %s", channelID) } @@ -181,6 +196,12 @@ func (h *TunnelHandler) RequestForward(taiID string, targetPort int) (taipb.TaiT return nil, fmt.Errorf("tai %s: no active register stream", taiID) } + muVal, ok := h.sendMu.Load(taiID) + if !ok { + return nil, fmt.Errorf("tai %s: no send mutex (stream closing?)", taiID) + } + mu := muVal.(*sync.Mutex) + channelID, err := registry.GenerateChannelID() if err != nil { return nil, fmt.Errorf("generate channel_id: %w", err) @@ -194,19 +215,31 @@ func (h *TunnelHandler) RequestForward(taiID string, targetPort int) (taipb.TaiT if !ok { return nil, fmt.Errorf("tai %s: register stream type mismatch", taiID) } - if err := regStream.Send(&taipb.TunnelControl{ + + h.logger.Debug("[forward] sending open command", + "tai_id", taiID, "port", targetPort, "channel_id", channelID[:16]) + + mu.Lock() + sendErr := regStream.Send(&taipb.TunnelControl{ Type: "open", ChannelId: channelID, TargetPort: int32(targetPort), - }); err != nil { - return nil, fmt.Errorf("send open: %w", err) + }) + mu.Unlock() + if sendErr != nil { + return nil, fmt.Errorf("send open: %w", sendErr) } + h.logger.Debug("[forward] open sent, waiting for callback", + "tai_id", taiID, "port", targetPort, "channel_id", channelID[:16]) + select { case fwd := <-waitCh: + h.logger.Debug("[forward] callback received", + "tai_id", taiID, "channel_id", channelID[:16]) return fwd, nil case <-time.After(10 * time.Second): - return nil, fmt.Errorf("tai %s: forward timeout (10s)", taiID) + return nil, fmt.Errorf("tai %s: forward timeout (10s) channel=%s", taiID, channelID[:16]) case <-regStream.Context().Done(): return nil, fmt.Errorf("tai %s: register stream closed while waiting for forward", taiID) } diff --git a/tai/volume/local.go b/tai/volume/local.go index 11ca1e11..1494790c 100644 --- a/tai/volume/local.go +++ b/tai/volume/local.go @@ -138,6 +138,10 @@ func (l *localStorage) MkdirAll(_ context.Context, sessionID, path string) error return os.MkdirAll(abs, 0o755) } +func (l *localStorage) Abs(_ context.Context, sessionID, path string) (string, error) { + return l.abs(sessionID, path) +} + // 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. diff --git a/tai/volume/mock_test.go b/tai/volume/mock_test.go index bdee3352..5c8875b8 100644 --- a/tai/volume/mock_test.go +++ b/tai/volume/mock_test.go @@ -190,6 +190,10 @@ func (m *mockVolumeServer) SyncPull(req *pb.SyncManifest, stream grpc.ServerStre return nil } +func (m *mockVolumeServer) Abs(_ context.Context, req *pb.FSRequest) (*pb.FSAbsResponse, error) { + return &pb.FSAbsResponse{Path: "/data/" + req.SessionId + "/" + req.Path}, nil +} + func (m *mockVolumeServer) ListDir(_ context.Context, req *pb.FSRequest) (*pb.FSListResponse, error) { return &pb.FSListResponse{Entries: []*pb.FileInfo{ {Path: "a.txt", Size: 10}, @@ -560,6 +564,10 @@ func (m *errMockVolumeServer) Copy(_ context.Context, _ *pb.FSCopyRequest) (*pb. return nil, fmt.Errorf("injected copy error") } +func (m *errMockVolumeServer) Abs(_ context.Context, _ *pb.FSRequest) (*pb.FSAbsResponse, error) { + return nil, fmt.Errorf("injected abs error") +} + func startErrMockServer(t *testing.T) (*grpc.ClientConn, func()) { t.Helper() lis, err := net.Listen("tcp", "127.0.0.1:0") @@ -748,3 +756,42 @@ func TestErrRemoteCopy(t *testing.T) { t.Error("expected error") } } + +func TestMockRemoteAbs(t *testing.T) { + conn, cleanup := startMockServer(t, &mockVolumeServer{}) + defer cleanup() + + vol := NewRemote(conn) + got, err := vol.Abs(context.Background(), "s1", ".") + if err != nil { + t.Fatalf("Abs: %v", err) + } + if got != "/data/s1/." { + t.Errorf("Abs = %q, want %q", got, "/data/s1/.") + } +} + +func TestMockRemoteAbsRelative(t *testing.T) { + conn, cleanup := startMockServer(t, &mockVolumeServer{}) + defer cleanup() + + vol := NewRemote(conn) + got, err := vol.Abs(context.Background(), "s1", "sub/file.txt") + if err != nil { + t.Fatalf("Abs: %v", err) + } + if got != "/data/s1/sub/file.txt" { + t.Errorf("Abs = %q, want %q", got, "/data/s1/sub/file.txt") + } +} + +func TestErrRemoteAbs(t *testing.T) { + conn, cleanup := startErrMockServer(t) + defer cleanup() + + vol := NewRemote(conn) + _, err := vol.Abs(context.Background(), "s1", ".") + if err == nil { + t.Error("expected error") + } +} diff --git a/tai/volume/pb/volume.pb.go b/tai/volume/pb/volume.pb.go index 96e97bd8..e7bac607 100644 --- a/tai/volume/pb/volume.pb.go +++ b/tai/volume/pb/volume.pb.go @@ -631,6 +631,50 @@ func (x *FSOpResponse) GetError() string { return "" } +type FSAbsResponse struct { + state protoimpl.MessageState `protogen:"open.v1"` + Path string `protobuf:"bytes,1,opt,name=path,proto3" json:"path,omitempty"` // absolute path on the host filesystem + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *FSAbsResponse) Reset() { + *x = FSAbsResponse{} + mi := &file_tai_volume_pb_volume_proto_msgTypes[8] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *FSAbsResponse) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*FSAbsResponse) ProtoMessage() {} + +func (x *FSAbsResponse) ProtoReflect() protoreflect.Message { + mi := &file_tai_volume_pb_volume_proto_msgTypes[8] + 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 FSAbsResponse.ProtoReflect.Descriptor instead. +func (*FSAbsResponse) Descriptor() ([]byte, []int) { + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{8} +} + +func (x *FSAbsResponse) GetPath() string { + if x != nil { + return x.Path + } + return "" +} + type FSReadRequest struct { state protoimpl.MessageState `protogen:"open.v1"` SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` @@ -641,7 +685,7 @@ type FSReadRequest struct { func (x *FSReadRequest) Reset() { *x = FSReadRequest{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[8] + mi := &file_tai_volume_pb_volume_proto_msgTypes[9] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -653,7 +697,7 @@ func (x *FSReadRequest) String() string { func (*FSReadRequest) ProtoMessage() {} func (x *FSReadRequest) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[8] + mi := &file_tai_volume_pb_volume_proto_msgTypes[9] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -666,7 +710,7 @@ func (x *FSReadRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSReadRequest.ProtoReflect.Descriptor instead. func (*FSReadRequest) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{8} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{9} } func (x *FSReadRequest) GetSessionId() string { @@ -695,7 +739,7 @@ type FSDataChunk struct { func (x *FSDataChunk) Reset() { *x = FSDataChunk{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[9] + mi := &file_tai_volume_pb_volume_proto_msgTypes[10] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -707,7 +751,7 @@ func (x *FSDataChunk) String() string { func (*FSDataChunk) ProtoMessage() {} func (x *FSDataChunk) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[9] + mi := &file_tai_volume_pb_volume_proto_msgTypes[10] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -720,7 +764,7 @@ func (x *FSDataChunk) ProtoReflect() protoreflect.Message { // Deprecated: Use FSDataChunk.ProtoReflect.Descriptor instead. func (*FSDataChunk) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{9} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{10} } func (x *FSDataChunk) GetData() []byte { @@ -764,7 +808,7 @@ type FSWriteChunk struct { func (x *FSWriteChunk) Reset() { *x = FSWriteChunk{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[10] + mi := &file_tai_volume_pb_volume_proto_msgTypes[11] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -776,7 +820,7 @@ func (x *FSWriteChunk) String() string { func (*FSWriteChunk) ProtoMessage() {} func (x *FSWriteChunk) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[10] + mi := &file_tai_volume_pb_volume_proto_msgTypes[11] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -789,7 +833,7 @@ func (x *FSWriteChunk) ProtoReflect() protoreflect.Message { // Deprecated: Use FSWriteChunk.ProtoReflect.Descriptor instead. func (*FSWriteChunk) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{10} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{11} } func (x *FSWriteChunk) GetSessionId() string { @@ -836,7 +880,7 @@ type FSWriteResponse struct { func (x *FSWriteResponse) Reset() { *x = FSWriteResponse{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[11] + mi := &file_tai_volume_pb_volume_proto_msgTypes[12] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -848,7 +892,7 @@ func (x *FSWriteResponse) String() string { func (*FSWriteResponse) ProtoMessage() {} func (x *FSWriteResponse) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[11] + mi := &file_tai_volume_pb_volume_proto_msgTypes[12] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -861,7 +905,7 @@ func (x *FSWriteResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use FSWriteResponse.ProtoReflect.Descriptor instead. func (*FSWriteResponse) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{11} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{12} } func (x *FSWriteResponse) GetSize() int64 { @@ -880,7 +924,7 @@ type FSListResponse struct { func (x *FSListResponse) Reset() { *x = FSListResponse{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[12] + mi := &file_tai_volume_pb_volume_proto_msgTypes[13] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -892,7 +936,7 @@ func (x *FSListResponse) String() string { func (*FSListResponse) ProtoMessage() {} func (x *FSListResponse) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[12] + mi := &file_tai_volume_pb_volume_proto_msgTypes[13] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -905,7 +949,7 @@ func (x *FSListResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use FSListResponse.ProtoReflect.Descriptor instead. func (*FSListResponse) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{12} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{13} } func (x *FSListResponse) GetEntries() []*FileInfo { @@ -926,7 +970,7 @@ type FSRemoveRequest struct { func (x *FSRemoveRequest) Reset() { *x = FSRemoveRequest{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[13] + mi := &file_tai_volume_pb_volume_proto_msgTypes[14] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -938,7 +982,7 @@ func (x *FSRemoveRequest) String() string { func (*FSRemoveRequest) ProtoMessage() {} func (x *FSRemoveRequest) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[13] + mi := &file_tai_volume_pb_volume_proto_msgTypes[14] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -951,7 +995,7 @@ func (x *FSRemoveRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSRemoveRequest.ProtoReflect.Descriptor instead. func (*FSRemoveRequest) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{13} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{14} } func (x *FSRemoveRequest) GetSessionId() string { @@ -986,7 +1030,7 @@ type FSRenameRequest struct { func (x *FSRenameRequest) Reset() { *x = FSRenameRequest{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[14] + mi := &file_tai_volume_pb_volume_proto_msgTypes[15] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -998,7 +1042,7 @@ func (x *FSRenameRequest) String() string { func (*FSRenameRequest) ProtoMessage() {} func (x *FSRenameRequest) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[14] + mi := &file_tai_volume_pb_volume_proto_msgTypes[15] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1011,7 +1055,7 @@ func (x *FSRenameRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSRenameRequest.ProtoReflect.Descriptor instead. func (*FSRenameRequest) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{14} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{15} } func (x *FSRenameRequest) GetSessionId() string { @@ -1048,7 +1092,7 @@ type FSCopyRequest struct { func (x *FSCopyRequest) Reset() { *x = FSCopyRequest{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[15] + mi := &file_tai_volume_pb_volume_proto_msgTypes[16] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -1060,7 +1104,7 @@ func (x *FSCopyRequest) String() string { func (*FSCopyRequest) ProtoMessage() {} func (x *FSCopyRequest) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[15] + mi := &file_tai_volume_pb_volume_proto_msgTypes[16] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1073,7 +1117,7 @@ func (x *FSCopyRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use FSCopyRequest.ProtoReflect.Descriptor instead. func (*FSCopyRequest) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{15} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{16} } func (x *FSCopyRequest) GetSessionId() string { @@ -1123,7 +1167,7 @@ type ArchiveRequest struct { func (x *ArchiveRequest) Reset() { *x = ArchiveRequest{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[16] + mi := &file_tai_volume_pb_volume_proto_msgTypes[17] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -1135,7 +1179,7 @@ func (x *ArchiveRequest) String() string { func (*ArchiveRequest) ProtoMessage() {} func (x *ArchiveRequest) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[16] + mi := &file_tai_volume_pb_volume_proto_msgTypes[17] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1148,7 +1192,7 @@ func (x *ArchiveRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use ArchiveRequest.ProtoReflect.Descriptor instead. func (*ArchiveRequest) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{16} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{17} } func (x *ArchiveRequest) GetSessionId() string { @@ -1189,7 +1233,7 @@ type ArchiveResponse struct { func (x *ArchiveResponse) Reset() { *x = ArchiveResponse{} - mi := &file_tai_volume_pb_volume_proto_msgTypes[17] + mi := &file_tai_volume_pb_volume_proto_msgTypes[18] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -1201,7 +1245,7 @@ func (x *ArchiveResponse) String() string { func (*ArchiveResponse) ProtoMessage() {} func (x *ArchiveResponse) ProtoReflect() protoreflect.Message { - mi := &file_tai_volume_pb_volume_proto_msgTypes[17] + mi := &file_tai_volume_pb_volume_proto_msgTypes[18] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1214,7 +1258,7 @@ func (x *ArchiveResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use ArchiveResponse.ProtoReflect.Descriptor instead. func (*ArchiveResponse) Descriptor() ([]byte, []int) { - return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{17} + return file_tai_volume_pb_volume_proto_rawDescGZIP(), []int{18} } func (x *ArchiveResponse) GetSizeBytes() int64 { @@ -1285,7 +1329,9 @@ const file_tai_volume_pb_volume_proto_rawDesc = "" + "\x04path\x18\x02 \x01(\tR\x04path\"4\n" + "\fFSOpResponse\x12\x0e\n" + "\x02ok\x18\x01 \x01(\bR\x02ok\x12\x14\n" + - "\x05error\x18\x02 \x01(\tR\x05error\"B\n" + + "\x05error\x18\x02 \x01(\tR\x05error\"#\n" + + "\rFSAbsResponse\x12\x12\n" + + "\x04path\x18\x01 \x01(\tR\x04path\"B\n" + "\rFSReadRequest\x12\x1d\n" + "\n" + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x12\n" + @@ -1334,7 +1380,7 @@ const file_tai_volume_pb_volume_proto_rawDesc = "" + "\n" + "size_bytes\x18\x01 \x01(\x03R\tsizeBytes\x12\x1f\n" + "\vfiles_count\x18\x02 \x01(\x05R\n" + - "filesCount2\xfa\a\n" + + "filesCount2\xab\b\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" + @@ -1344,7 +1390,8 @@ const file_tai_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.FSOpResponse\x121\n" + + "\bMkdirAll\x12\x11.volume.FSRequest\x1a\x14.volume.FSOpResponse\x12/\n" + + "\x03Abs\x12\x11.volume.FSRequest\x1a\x15.volume.FSAbsResponse\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" + @@ -1368,7 +1415,7 @@ func file_tai_volume_pb_volume_proto_rawDescGZIP() []byte { } 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_msgTypes = make([]protoimpl.MessageInfo, 19) var file_tai_volume_pb_volume_proto_goTypes = []any{ (FileChunk_ChunkType)(0), // 0: volume.FileChunk.ChunkType (*FileInfo)(nil), // 1: volume.FileInfo @@ -1379,16 +1426,17 @@ var file_tai_volume_pb_volume_proto_goTypes = []any{ (*SyncResult)(nil), // 6: volume.SyncResult (*FSRequest)(nil), // 7: volume.FSRequest (*FSOpResponse)(nil), // 8: volume.FSOpResponse - (*FSReadRequest)(nil), // 9: volume.FSReadRequest - (*FSDataChunk)(nil), // 10: volume.FSDataChunk - (*FSWriteChunk)(nil), // 11: volume.FSWriteChunk - (*FSWriteResponse)(nil), // 12: volume.FSWriteResponse - (*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 + (*FSAbsResponse)(nil), // 9: volume.FSAbsResponse + (*FSReadRequest)(nil), // 10: volume.FSReadRequest + (*FSDataChunk)(nil), // 11: volume.FSDataChunk + (*FSWriteChunk)(nil), // 12: volume.FSWriteChunk + (*FSWriteResponse)(nil), // 13: volume.FSWriteResponse + (*FSListResponse)(nil), // 14: volume.FSListResponse + (*FSRemoveRequest)(nil), // 15: volume.FSRemoveRequest + (*FSRenameRequest)(nil), // 16: volume.FSRenameRequest + (*FSCopyRequest)(nil), // 17: volume.FSCopyRequest + (*ArchiveRequest)(nil), // 18: volume.ArchiveRequest + (*ArchiveResponse)(nil), // 19: volume.ArchiveResponse } var file_tai_volume_pb_volume_proto_depIdxs = []int32{ 1, // 0: volume.SyncManifest.files:type_name -> volume.FileInfo @@ -1400,42 +1448,44 @@ var file_tai_volume_pb_volume_proto_depIdxs = []int32{ 1, // 6: volume.FSListResponse.entries:type_name -> volume.FileInfo 3, // 7: volume.Volume.SyncPush:input_type -> volume.SyncMessage 2, // 8: volume.Volume.SyncPull:input_type -> volume.SyncManifest - 9, // 9: volume.Volume.ReadFile:input_type -> volume.FSReadRequest - 11, // 10: volume.Volume.WriteFile:input_type -> volume.FSWriteChunk + 10, // 9: volume.Volume.ReadFile:input_type -> volume.FSReadRequest + 12, // 10: volume.Volume.WriteFile:input_type -> volume.FSWriteChunk 7, // 11: volume.Volume.Stat:input_type -> volume.FSRequest 7, // 12: volume.Volume.ListDir:input_type -> volume.FSRequest - 14, // 13: volume.Volume.Remove:input_type -> volume.FSRemoveRequest - 15, // 14: volume.Volume.Rename:input_type -> volume.FSRenameRequest + 15, // 13: volume.Volume.Remove:input_type -> volume.FSRemoveRequest + 16, // 14: volume.Volume.Rename:input_type -> volume.FSRenameRequest 7, // 15: volume.Volume.MkdirAll:input_type -> volume.FSRequest - 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, // 16: volume.Volume.Abs:input_type -> volume.FSRequest + 17, // 17: volume.Volume.Copy:input_type -> volume.FSCopyRequest + 18, // 18: volume.Volume.Zip:input_type -> volume.ArchiveRequest + 18, // 19: volume.Volume.Unzip:input_type -> volume.ArchiveRequest + 18, // 20: volume.Volume.Gzip:input_type -> volume.ArchiveRequest + 18, // 21: volume.Volume.Gunzip:input_type -> volume.ArchiveRequest + 18, // 22: volume.Volume.Tar:input_type -> volume.ArchiveRequest + 18, // 23: volume.Volume.Untar:input_type -> volume.ArchiveRequest + 18, // 24: volume.Volume.Tgz:input_type -> volume.ArchiveRequest + 18, // 25: volume.Volume.Untgz:input_type -> volume.ArchiveRequest + 3, // 26: volume.Volume.SyncPush:output_type -> volume.SyncMessage + 3, // 27: volume.Volume.SyncPull:output_type -> volume.SyncMessage + 11, // 28: volume.Volume.ReadFile:output_type -> volume.FSDataChunk + 13, // 29: volume.Volume.WriteFile:output_type -> volume.FSWriteResponse + 1, // 30: volume.Volume.Stat:output_type -> volume.FileInfo + 14, // 31: volume.Volume.ListDir:output_type -> volume.FSListResponse + 8, // 32: volume.Volume.Remove:output_type -> volume.FSOpResponse + 8, // 33: volume.Volume.Rename:output_type -> volume.FSOpResponse + 8, // 34: volume.Volume.MkdirAll:output_type -> volume.FSOpResponse + 9, // 35: volume.Volume.Abs:output_type -> volume.FSAbsResponse + 6, // 36: volume.Volume.Copy:output_type -> volume.SyncResult + 19, // 37: volume.Volume.Zip:output_type -> volume.ArchiveResponse + 19, // 38: volume.Volume.Unzip:output_type -> volume.ArchiveResponse + 19, // 39: volume.Volume.Gzip:output_type -> volume.ArchiveResponse + 19, // 40: volume.Volume.Gunzip:output_type -> volume.ArchiveResponse + 19, // 41: volume.Volume.Tar:output_type -> volume.ArchiveResponse + 19, // 42: volume.Volume.Untar:output_type -> volume.ArchiveResponse + 19, // 43: volume.Volume.Tgz:output_type -> volume.ArchiveResponse + 19, // 44: volume.Volume.Untgz:output_type -> volume.ArchiveResponse + 26, // [26:45] is the sub-list for method output_type + 7, // [7:26] 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 @@ -1458,7 +1508,7 @@ func file_tai_volume_pb_volume_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_tai_volume_pb_volume_proto_rawDesc), len(file_tai_volume_pb_volume_proto_rawDesc)), NumEnums: 1, - NumMessages: 18, + NumMessages: 19, NumExtensions: 0, NumServices: 1, }, diff --git a/tai/volume/pb/volume.proto b/tai/volume/pb/volume.proto index 5fa38279..3ecde25a 100644 --- a/tai/volume/pb/volume.proto +++ b/tai/volume/pb/volume.proto @@ -30,6 +30,8 @@ service Volume { rpc Remove(FSRemoveRequest) returns (FSOpResponse); rpc Rename(FSRenameRequest) returns (FSOpResponse); rpc MkdirAll(FSRequest) returns (FSOpResponse); + // Abs: resolve a session-relative path to its absolute path on the host. + rpc Abs(FSRequest) returns (FSAbsResponse); // Copy: copy src to dst within the same workspace (server-side when remote). rpc Copy(FSCopyRequest) returns (SyncResult); @@ -112,6 +114,10 @@ message FSOpResponse { string error = 2; } +message FSAbsResponse { + string path = 1; // absolute path on the host filesystem +} + message FSReadRequest { string session_id = 1; string path = 2; diff --git a/tai/volume/pb/volume_grpc.pb.go b/tai/volume/pb/volume_grpc.pb.go index a5370fe3..a35792e3 100644 --- a/tai/volume/pb/volume_grpc.pb.go +++ b/tai/volume/pb/volume_grpc.pb.go @@ -28,6 +28,7 @@ const ( Volume_Remove_FullMethodName = "/volume.Volume/Remove" Volume_Rename_FullMethodName = "/volume.Volume/Rename" Volume_MkdirAll_FullMethodName = "/volume.Volume/MkdirAll" + Volume_Abs_FullMethodName = "/volume.Volume/Abs" Volume_Copy_FullMethodName = "/volume.Volume/Copy" Volume_Zip_FullMethodName = "/volume.Volume/Zip" Volume_Unzip_FullMethodName = "/volume.Volume/Unzip" @@ -64,6 +65,8 @@ 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) + // Abs: resolve a session-relative path to its absolute path on the host. + Abs(ctx context.Context, in *FSRequest, opts ...grpc.CallOption) (*FSAbsResponse, 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) @@ -198,6 +201,16 @@ func (c *volumeClient) MkdirAll(ctx context.Context, in *FSRequest, opts ...grpc return out, nil } +func (c *volumeClient) Abs(ctx context.Context, in *FSRequest, opts ...grpc.CallOption) (*FSAbsResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(FSAbsResponse) + err := c.cc.Invoke(ctx, Volume_Abs_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + 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) @@ -313,6 +326,8 @@ type VolumeServer interface { Remove(context.Context, *FSRemoveRequest) (*FSOpResponse, error) Rename(context.Context, *FSRenameRequest) (*FSOpResponse, error) MkdirAll(context.Context, *FSRequest) (*FSOpResponse, error) + // Abs: resolve a session-relative path to its absolute path on the host. + Abs(context.Context, *FSRequest) (*FSAbsResponse, 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) @@ -360,6 +375,9 @@ 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) Abs(context.Context, *FSRequest) (*FSAbsResponse, error) { + return nil, status.Error(codes.Unimplemented, "method Abs not implemented") +} func (UnimplementedVolumeServer) Copy(context.Context, *FSCopyRequest) (*SyncResult, error) { return nil, status.Error(codes.Unimplemented, "method Copy not implemented") } @@ -534,6 +552,24 @@ func _Volume_MkdirAll_Handler(srv interface{}, ctx context.Context, dec func(int return interceptor(ctx, in, info, handler) } +func _Volume_Abs_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(FSRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(VolumeServer).Abs(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Volume_Abs_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(VolumeServer).Abs(ctx, req.(*FSRequest)) + } + 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 { @@ -723,6 +759,10 @@ var Volume_ServiceDesc = grpc.ServiceDesc{ MethodName: "MkdirAll", Handler: _Volume_MkdirAll_Handler, }, + { + MethodName: "Abs", + Handler: _Volume_Abs_Handler, + }, { MethodName: "Copy", Handler: _Volume_Copy_Handler, diff --git a/tai/volume/remote.go b/tai/volume/remote.go index e7019e66..59b2aabf 100644 --- a/tai/volume/remote.go +++ b/tai/volume/remote.go @@ -167,6 +167,17 @@ func (r *remoteStorage) MkdirAll(ctx context.Context, sessionID, path string) er return nil } +func (r *remoteStorage) Abs(ctx context.Context, sessionID, path string) (string, error) { + resp, err := r.client.Abs(ctx, &pb.FSRequest{ + SessionId: sessionID, + Path: path, + }) + if err != nil { + return "", err + } + return resp.Path, nil +} + // 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() diff --git a/tai/volume/volume.go b/tai/volume/volume.go index e8498c33..ebf264e9 100644 --- a/tai/volume/volume.go +++ b/tai/volume/volume.go @@ -17,6 +17,7 @@ type Volume interface { Remove(ctx context.Context, sessionID, path string, recursive bool) error Rename(ctx context.Context, sessionID, oldPath, newPath string) error MkdirAll(ctx context.Context, sessionID, path string) error + Abs(ctx context.Context, sessionID, path string) (string, error) SyncPush(ctx context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) SyncPull(ctx context.Context, sessionID, localDir string, opts ...SyncOption) (*SyncResult, error) diff --git a/tai/volume/volume_test.go b/tai/volume/volume_test.go index eae954a8..8a30619e 100644 --- a/tai/volume/volume_test.go +++ b/tai/volume/volume_test.go @@ -560,10 +560,83 @@ func TestRemoteVolume(t *testing.T) { _ = vol.Remove(ctx, arcSid, ".", true) }) + t.Run("Abs dot", func(t *testing.T) { + absSid := "abs-remote-test" + _ = vol.MkdirAll(ctx, absSid, ".") + got, err := vol.Abs(ctx, absSid, ".") + if err != nil { + t.Fatalf("Abs: %v", err) + } + if got == "" { + t.Error("Abs returned empty") + } + _ = vol.Remove(ctx, absSid, ".", true) + }) + + t.Run("Abs relative", func(t *testing.T) { + got, err := vol.Abs(ctx, sid, "sub/file.txt") + if err != nil { + t.Fatalf("Abs: %v", err) + } + if got == "" { + t.Error("Abs returned empty") + } + }) + + t.Run("Abs path traversal", func(t *testing.T) { + _, err := vol.Abs(ctx, sid, "../../etc/passwd") + if err == nil { + t.Error("expected error for Abs path traversal") + } + }) + // Cleanup _ = vol.Remove(ctx, sid, ".", true) } +func TestLocalAbs_Dot(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + ctx := context.Background() + sid := "abs-test" + + got, err := vol.Abs(ctx, sid, ".") + if err != nil { + t.Fatalf("Abs: %v", err) + } + want := dir + "/" + sid + if got != want { + t.Errorf("Abs(\".\") = %q, want %q", got, want) + } +} + +func TestLocalAbs_RelativePath(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + ctx := context.Background() + sid := "abs-rel" + + got, err := vol.Abs(ctx, sid, "sub/file.txt") + if err != nil { + t.Fatalf("Abs: %v", err) + } + want := dir + "/" + sid + "/sub/file.txt" + if got != want { + t.Errorf("Abs = %q, want %q", got, want) + } +} + +func TestLocalAbs_PathTraversal(t *testing.T) { + dir := t.TempDir() + vol := NewLocal(dir) + ctx := context.Background() + + _, err := vol.Abs(ctx, "test", "../../etc/passwd") + if err == nil { + t.Error("expected error for path traversal in Abs") + } +} + func TestLocalPathTraversal(t *testing.T) { dir := t.TempDir() vol := NewLocal(dir) diff --git a/tai/workspace/workspace.go b/tai/workspace/workspace.go index 5ff5b518..08c518b1 100644 --- a/tai/workspace/workspace.go +++ b/tai/workspace/workspace.go @@ -31,6 +31,9 @@ type FS interface { // 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) + + // GetRoot returns the absolute path of this workspace's root directory on the host filesystem. + GetRoot() (string, error) } // New creates an FS backed by the given Volume for the specified session. @@ -119,6 +122,10 @@ func (w *workspaceFS) MkdirAll(name string, _ os.FileMode) error { return w.vol.MkdirAll(context.Background(), w.session, name) } +func (w *workspaceFS) GetRoot() (string, error) { + return w.vol.Abs(context.Background(), w.session, ".") +} + func (w *workspaceFS) Close() error { return nil } // --- fs.FileInfo adapter --- diff --git a/tai/workspace/workspace_test.go b/tai/workspace/workspace_test.go index 97ba9f1a..2a03770f 100644 --- a/tai/workspace/workspace_test.go +++ b/tai/workspace/workspace_test.go @@ -210,6 +210,25 @@ func TestWorkspaceFS(t *testing.T) { }) } +func TestGetRoot(t *testing.T) { + dir := t.TempDir() + vol := volume.NewLocal(dir) + defer vol.Close() + + sid := "getroot-test" + wfs := New(vol, sid) + defer wfs.Close() + + root, err := wfs.GetRoot() + if err != nil { + t.Fatalf("GetRoot: %v", err) + } + want := dir + "/" + sid + if root != want { + t.Errorf("GetRoot() = %q, want %q", root, want) + } +} + // Compile-time interface checks. var ( _ fs.FS = (*workspaceFS)(nil) diff --git a/workspace/manager.go b/workspace/manager.go index e9be938d..3864dbbd 100644 --- a/workspace/manager.go +++ b/workspace/manager.go @@ -275,20 +275,7 @@ func (m *Manager) MountPath(ctx context.Context, id string) (string, error) { if err != nil { return "", err } - _ = vol - for _, snap := range listNodes() { - res, ok := tai.GetResources(snap.TaiID) - if !ok { - continue - } - if res.Volume == vol { - if res.DataDir == "" { - return "", nil - } - return res.DataDir + "/" + id, nil - } - } - return "", nil + return vol.Abs(ctx, id, ".") } // --- internal --- From b667ec310b0f0dd84e9c0357cdad8bc7c4bd6493 Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 13 Mar 2026 18:41:15 +0800 Subject: [PATCH 3/8] feat(registry, grpc): introduce channel ID constants and logging enhancements - Added constants for channel ID byte size and short length to improve clarity and maintainability. - Implemented ShortChannelID function to truncate channel IDs for log messages, enhancing readability. - Updated gRPC tunnel logging to utilize the new ShortChannelID function, ensuring consistent and concise channel ID representation in logs. Made-with: Cursor --- tai/registry/registry.go | 17 ++++++++++++++++- tai/tunnel/grpc_handler.go | 14 ++++++++------ 2 files changed, 24 insertions(+), 7 deletions(-) diff --git a/tai/registry/registry.go b/tai/registry/registry.go index 4784138c..2e8c5300 100644 --- a/tai/registry/registry.go +++ b/tai/registry/registry.go @@ -401,14 +401,29 @@ func (r *Registry) bridgeTunnelConn(taiID string, targetPort int, localConn net. r.logger.Error("no bridge function configured", "tai_id", taiID, "port", targetPort) } +// ChannelIDBytes is the number of random bytes used to generate a channel ID. +// The resulting hex string is 2× this value (64 characters). +const ChannelIDBytes = 32 + +// ChannelIDShortLen is the max characters shown in log messages. +const ChannelIDShortLen = 16 + func generateChannelID() (string, error) { - b := make([]byte, 32) + b := make([]byte, ChannelIDBytes) if _, err := rand.Read(b); err != nil { return "", err } return hex.EncodeToString(b), nil } +// ShortChannelID truncates a channel ID for log display. +func ShortChannelID(id string) string { + if len(id) <= ChannelIDShortLen { + return id + } + return id[:ChannelIDShortLen] +} + type contextCancel struct { done chan struct{} } diff --git a/tai/tunnel/grpc_handler.go b/tai/tunnel/grpc_handler.go index 57d8b93f..bdf39deb 100644 --- a/tai/tunnel/grpc_handler.go +++ b/tai/tunnel/grpc_handler.go @@ -175,12 +175,13 @@ func (h *TunnelHandler) Forward(stream taipb.TaiTunnel_ForwardServer) error { } channelID := vals[0] - h.logger.Debug("[forward] Forward stream arrived", "channel_id", channelID[:16]) + short := registry.ShortChannelID(channelID) + h.logger.Debug("[forward] Forward stream arrived", "channel_id", short) if ch, ok := h.pending.LoadAndDelete(channelID); ok { ch.(chan taipb.TaiTunnel_ForwardServer) <- stream } else { - h.logger.Warn("[forward] no pending channel (expired?)", "channel_id", channelID[:16]) + h.logger.Warn("[forward] no pending channel (expired?)", "channel_id", short) return fmt.Errorf("no pending channel for %s", channelID) } @@ -216,8 +217,9 @@ func (h *TunnelHandler) RequestForward(taiID string, targetPort int) (taipb.TaiT return nil, fmt.Errorf("tai %s: register stream type mismatch", taiID) } + short := registry.ShortChannelID(channelID) h.logger.Debug("[forward] sending open command", - "tai_id", taiID, "port", targetPort, "channel_id", channelID[:16]) + "tai_id", taiID, "port", targetPort, "channel_id", short) mu.Lock() sendErr := regStream.Send(&taipb.TunnelControl{ @@ -231,15 +233,15 @@ func (h *TunnelHandler) RequestForward(taiID string, targetPort int) (taipb.TaiT } h.logger.Debug("[forward] open sent, waiting for callback", - "tai_id", taiID, "port", targetPort, "channel_id", channelID[:16]) + "tai_id", taiID, "port", targetPort, "channel_id", short) select { case fwd := <-waitCh: h.logger.Debug("[forward] callback received", - "tai_id", taiID, "channel_id", channelID[:16]) + "tai_id", taiID, "channel_id", short) return fwd, nil case <-time.After(10 * time.Second): - return nil, fmt.Errorf("tai %s: forward timeout (10s) channel=%s", taiID, channelID[:16]) + return nil, fmt.Errorf("tai %s: forward timeout (10s) channel=%s", taiID, short) case <-regStream.Context().Done(): return nil, fmt.Errorf("tai %s: register stream closed while waiting for forward", taiID) } From f20f7797a3289d8536ff826b97f1ccb223ced01c Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 13 Mar 2026 20:02:50 +0800 Subject: [PATCH 4/8] refactor(tunnel): improve error handling in gRPC stream tests - Updated test cases to handle potential stream closure more gracefully by checking for send errors before proceeding with receive operations. - Enhanced clarity in test logic to ensure that expected errors are correctly asserted when metadata is missing or channel IDs are non-existent. Made-with: Cursor --- tai/tunnel/grpc_handler_test.go | 24 ++++++++++++++---------- 1 file changed, 14 insertions(+), 10 deletions(-) diff --git a/tai/tunnel/grpc_handler_test.go b/tai/tunnel/grpc_handler_test.go index 4246aa6b..278dd300 100644 --- a/tai/tunnel/grpc_handler_test.go +++ b/tai/tunnel/grpc_handler_test.go @@ -204,13 +204,16 @@ func TestForward_MissingMetadata(t *testing.T) { if err != nil { t.Fatal(err) } - err = stream.Send(&taipb.ForwardData{Data: []byte("hello")}) - if err != nil { - t.Fatal(err) + + // Server may close the stream before or after Send completes (race). + // Either Send or Recv returning an error confirms the server rejected. + sendErr := stream.Send(&taipb.ForwardData{Data: []byte("hello")}) + if sendErr != nil { + return // server already closed stream — pass } - _, err = stream.Recv() - if err == nil { + _, recvErr := stream.Recv() + if recvErr == nil { t.Fatal("expected error for missing channel_id metadata") } } @@ -224,13 +227,14 @@ func TestForward_NoPendingChannel(t *testing.T) { if err != nil { t.Fatal(err) } - err = stream.Send(&taipb.ForwardData{Data: []byte("hello")}) - if err != nil { - t.Fatal(err) + + sendErr := stream.Send(&taipb.ForwardData{Data: []byte("hello")}) + if sendErr != nil { + return // server already closed stream — pass } - _, err = stream.Recv() - if err == nil { + _, recvErr := stream.Recv() + if recvErr == nil { t.Fatal("expected error for non-existent channel_id") } } From 223d02ebfe9085478ab709ce7d4b41ea368db8ae Mon Sep 17 00:00:00 2001 From: Max Date: Sat, 14 Mar 2026 15:50:59 +0800 Subject: [PATCH 5/8] feat(sandbox): implement sandbox token handling in stream execution - Added support for issuing and passing a sandbox token during stream execution in the Assistant. - Updated StreamRequest to include a Token field for managing user authentication. - Enhanced ClaudeRunner to set environment variables for the sandbox token and refresh token. - Refactored SandboxToken structure to clarify its purpose and manage token credentials effectively. - Modified sandbox management routes to enforce OAuth guard for improved security. Made-with: Cursor --- agent/assistant/sandbox_v2.go | 10 ++++ agent/sandbox/v2/claude/runner.go | 9 ++++ agent/sandbox/v2/token.go | 90 +++++++++++++++++++++++++++++++ agent/sandbox/v2/types/runner.go | 1 + agent/sandbox/v2/types/token.go | 9 ++-- openapi/computer/computer.go | 2 + openapi/openapi.go | 2 +- openapi/sandbox/manage.go | 15 +++--- sandbox/v2/manager.go | 13 +++++ tai/dial.go | 2 +- 10 files changed, 138 insertions(+), 15 deletions(-) create mode 100644 agent/sandbox/v2/token.go diff --git a/agent/assistant/sandbox_v2.go b/agent/assistant/sandbox_v2.go index a63ba159..a8026ae9 100644 --- a/agent/assistant/sandbox_v2.go +++ b/agent/assistant/sandbox_v2.go @@ -149,6 +149,15 @@ func (ast *Assistant) executeSandboxV2Stream( // Resolve connector for Stream. conn, _, _ := ast.GetConnector(ctx) + var tok *sandboxTypes.SandboxToken + if ctx.Authorized != nil { + var err error + tok, err = sandboxv2.IssueSandboxToken(ctx.Authorized.TeamID, ctx.Authorized.UserID) + if err != nil { + return nil, fmt.Errorf("issue sandbox token: %w", err) + } + } + streamReq := &sandboxTypes.StreamRequest{ Computer: computer, Config: cfg, @@ -156,6 +165,7 @@ func (ast *Assistant) executeSandboxV2Stream( Messages: completionMessages, SystemPrompt: systemPrompt, ChatID: ctx.ChatID, + Token: tok, } execReq := &sandboxv2.ExecuteRequest{ diff --git a/agent/sandbox/v2/claude/runner.go b/agent/sandbox/v2/claude/runner.go index 886fde38..576933a7 100644 --- a/agent/sandbox/v2/claude/runner.go +++ b/agent/sandbox/v2/claude/runner.go @@ -235,6 +235,15 @@ func (r *ClaudeRunner) buildCLICommand(req *types.StreamRequest, oe *osEnv, isCo } } + if req.Token != nil { + if req.Token.Token != "" { + env["YAO_TOKEN"] = req.Token.Token + } + if req.Token.RefreshToken != "" { + env["YAO_REFRESH_TOKEN"] = req.Token.RefreshToken + } + } + var systemPrompt string envPrompt := buildSandboxEnvPrompt(oe.WorkDir) if !isContinuation && req.SystemPrompt != "" { diff --git a/agent/sandbox/v2/token.go b/agent/sandbox/v2/token.go new file mode 100644 index 00000000..dc47f9d0 --- /dev/null +++ b/agent/sandbox/v2/token.go @@ -0,0 +1,90 @@ +package sandboxv2 + +import ( + "fmt" + "time" + + lrustore "github.com/yaoapp/gou/store/lru" + "github.com/yaoapp/yao/agent/sandbox/v2/types" + "github.com/yaoapp/yao/openapi/oauth" +) + +const ( + accessTokenTTL = 2 * time.Hour + refreshTokenTTL = 30 * 24 * time.Hour // 30 days + tokenCacheSize = 1024 +) + +var tokenCache *lrustore.Cache + +func init() { + c, err := lrustore.New(tokenCacheSize) + if err != nil { + panic("sandbox token cache init failed: " + err.Error()) + } + tokenCache = c +} + +func cacheKey(teamID, userID string) string { + if teamID == "" { + return userID + } + return teamID + "/" + userID +} + +func getToken(teamID, userID string) *types.SandboxToken { + val, ok := tokenCache.Get(cacheKey(teamID, userID)) + if !ok { + return nil + } + tok, _ := val.(*types.SandboxToken) + return tok +} + +func setToken(teamID, userID string, tok *types.SandboxToken, ttl time.Duration) { + tokenCache.Set(cacheKey(teamID, userID), tok, ttl) +} + +// IssueSandboxToken returns a valid identity token for the given user. +// Tokens are cached by (teamID, userID); a new token is only issued on +// cache miss or expiry. Returns nil without error when oauth.OAuth is nil. +func IssueSandboxToken(teamID, userID string) (*types.SandboxToken, error) { + if tok := getToken(teamID, userID); tok != nil { + return tok, nil + } + + svc := oauth.OAuth + if svc == nil { + return nil, nil + } + + subject, err := svc.Subject("__yao.sandbox", userID) + if err != nil { + return nil, fmt.Errorf("sandbox token: derive subject: %w", err) + } + + extraClaims := map[string]interface{}{ + "user_id": userID, + } + if teamID != "" { + extraClaims["team_id"] = teamID + } + + tokenStr, err := svc.MakeAccessToken("__yao.sandbox", "sandbox:mcp", subject, + int(accessTokenTTL.Seconds()), extraClaims) + if err != nil { + return nil, fmt.Errorf("sandbox token: issue access token: %w", err) + } + + tok := &types.SandboxToken{Token: tokenStr} + + refreshStr, err := svc.MakeRefreshToken("__yao.sandbox", "sandbox:mcp", subject, + int(refreshTokenTTL.Seconds()), extraClaims) + if err != nil { + return nil, fmt.Errorf("sandbox token: issue refresh token: %w", err) + } + tok.RefreshToken = refreshStr + + setToken(teamID, userID, tok, accessTokenTTL) + return tok, nil +} diff --git a/agent/sandbox/v2/types/runner.go b/agent/sandbox/v2/types/runner.go index 5a90f80f..8674e2e2 100644 --- a/agent/sandbox/v2/types/runner.go +++ b/agent/sandbox/v2/types/runner.go @@ -50,4 +50,5 @@ type StreamRequest struct { Messages []agentContext.Message SystemPrompt string ChatID string + Token *SandboxToken // current user's sandbox token for MCP callbacks } diff --git a/agent/sandbox/v2/types/token.go b/agent/sandbox/v2/types/token.go index 7f1cf525..4900f6eb 100644 --- a/agent/sandbox/v2/types/token.go +++ b/agent/sandbox/v2/types/token.go @@ -1,9 +1,8 @@ package types -import "time" - -// SandboxToken is a short-lived JWT issued for a sandbox computer. +// SandboxToken holds credentials for a sandbox execution session. +// Expiry is managed by the LRU store TTL, not stored here. type SandboxToken struct { - Token string - ExpiresAt time.Time + Token string // access token → YAO_TOKEN + RefreshToken string // refresh token → YAO_REFRESH_TOKEN } diff --git a/openapi/computer/computer.go b/openapi/computer/computer.go index fcd6420d..e0c80837 100644 --- a/openapi/computer/computer.go +++ b/openapi/computer/computer.go @@ -37,6 +37,7 @@ type computerOption 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"` Status string `json:"status"` Mode string `json:"mode,omitempty"` @@ -276,6 +277,7 @@ func boxToOption(b *sandboxv2.Box) computerOption { Kind: "box", ID: snap.ID, DisplayName: displayName, + ContainerID: snap.ContainerID, NodeID: snap.NodeID, Status: snap.Status, Mode: mode, diff --git a/openapi/openapi.go b/openapi/openapi.go index 4a4a0b7c..4c20ccc4 100644 --- a/openapi/openapi.go +++ b/openapi/openapi.go @@ -180,7 +180,7 @@ func (openapi *OpenAPI) Attach(router *gin.Engine) { sandbox.SetPathPrefix(baseURL) sandboxGroup := group.Group("/sandbox") sandbox.Attach(sandboxGroup, openapi.OAuth) - sandbox.AttachManage(sandboxGroup) + sandbox.AttachManage(sandboxGroup, openapi.OAuth) // Computer option handlers (for InputArea selector) openapiComputer.Attach(group.Group("/computer"), openapi.OAuth) diff --git a/openapi/sandbox/manage.go b/openapi/sandbox/manage.go index 71754f94..270c3d62 100644 --- a/openapi/sandbox/manage.go +++ b/openapi/sandbox/manage.go @@ -18,20 +18,19 @@ import ( ) // 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) +func AttachManage(group *gin.RouterGroup, oauth types.OAuth) { + group.GET("", oauth.Guard, handleList) + group.POST("", oauth.Guard, handleCreate) + group.GET("/:id", oauth.Guard, handleGet) + group.DELETE("/:id", oauth.Guard, handleRemove) + group.POST("/:id/exec", oauth.Guard, handleExec) + group.POST("/:id/heartbeat", oauth.Guard, handleHeartbeat) } // resolveOwner returns TeamID if present, otherwise UserID. diff --git a/sandbox/v2/manager.go b/sandbox/v2/manager.go index 7a271257..3306d95d 100644 --- a/sandbox/v2/manager.go +++ b/sandbox/v2/manager.go @@ -372,6 +372,9 @@ func (m *Manager) buildTaiCreateOptions(opts CreateOptions, nodeID, sandboxID st "sandbox-node-id": nodeID, "sandbox-policy": string(opts.Policy), } + if opts.VNC { + labels["sandbox-vnc"] = "true" + } if opts.WorkspaceID != "" { labels["workspace-id"] = opts.WorkspaceID } @@ -455,6 +458,15 @@ func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, res *tai.Conn if c.Name != "" { cid = c.Name } + hasVNC := c.Labels["sandbox-vnc"] == "true" + if !hasVNC { + for _, p := range c.Ports { + if p.ContainerPort == 5900 || p.ContainerPort == 6080 { + hasVNC = true + break + } + } + } box := &Box{ id: sandboxID, containerID: cid, @@ -465,6 +477,7 @@ func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, res *tai.Conn createdAt: time.Now(), image: c.Image, workspaceID: c.Labels["workspace-id"], + vnc: hasVNC, workDir: "/workspace", manager: m, } diff --git a/tai/dial.go b/tai/dial.go index 5844f25d..0f560cd3 100644 --- a/tai/dial.go +++ b/tai/dial.go @@ -189,8 +189,8 @@ func buildResources(conn *grpc.ClientConn, cfg *dialConfig, env dialEnv) (*ConnR if res.Runtime != nil { res.Proxy = env.newProxy(cfg.ports) - res.VNC = env.newVNC(cfg.ports) } + res.VNC = env.newVNC(cfg.ports) return res, nil } From c4dbcdba2bf2d7cae30bd8ee52bdaf609de71fe3 Mon Sep 17 00:00:00 2001 From: Max Date: Sat, 14 Mar 2026 16:23:24 +0800 Subject: [PATCH 6/8] feat(openapi): integrate Tai forward handlers for proxy and VNC operations - Replaced direct tunnel handling with a dedicated Tai integration for managing proxy and VNC requests. - Updated routing to streamline the attachment of Tai handlers, enhancing modularity and maintainability. Made-with: Cursor --- openapi/openapi.go | 7 ++-- openapi/tai/proxy.go | 60 ++++++++++++++++++++++++++++ openapi/tai/tai.go | 39 ++++++++++++++++++ openapi/tai/util.go | 95 ++++++++++++++++++++++++++++++++++++++++++++ openapi/tai/vnc.go | 57 ++++++++++++++++++++++++++ 5 files changed, 254 insertions(+), 4 deletions(-) create mode 100644 openapi/tai/proxy.go create mode 100644 openapi/tai/tai.go create mode 100644 openapi/tai/util.go create mode 100644 openapi/tai/vnc.go diff --git a/openapi/openapi.go b/openapi/openapi.go index 4c20ccc4..49f2c78f 100644 --- a/openapi/openapi.go +++ b/openapi/openapi.go @@ -27,12 +27,12 @@ import ( "github.com/yaoapp/yao/openapi/otp" "github.com/yaoapp/yao/openapi/response" "github.com/yaoapp/yao/openapi/sandbox" + openapiTai "github.com/yaoapp/yao/openapi/tai" "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" ) // Server is the OpenAPI server @@ -191,9 +191,8 @@ func (openapi *OpenAPI) Attach(router *gin.Engine) { // Tai nodes handlers nodes.Attach(group.Group("/nodes"), openapi.OAuth) - // Tai tunnel: gRPC Forward-based HTTP/VNC transparent proxy - group.Any("/tai/:taiID/proxy/*path", taitunnel.HandleForwardLazy) - group.Any("/tai/:taiID/vnc/*path", taitunnel.HandleForwardLazy) + // Tai forward handlers (proxy + VNC, dispatches tunnel vs local) + openapiTai.Attach(group) // Tai direct registration API (uses /tai-nodes/ prefix to avoid routing conflict with /tai/:taiID/) group.POST("/tai-nodes/register", taiapi.HandleRegister) diff --git a/openapi/tai/proxy.go b/openapi/tai/proxy.go new file mode 100644 index 00000000..dabf9bef --- /dev/null +++ b/openapi/tai/proxy.go @@ -0,0 +1,60 @@ +package tai + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" + yaoTai "github.com/yaoapp/yao/tai" +) + +// handleLocalProxy resolves the container's HTTP address via Docker socket +// and reverse-proxies the request. +func handleLocalProxy(c *gin.Context, taiID string) { + res, ok := yaoTai.GetResources(taiID) + if !ok || res.Proxy == nil { + c.JSON(http.StatusBadGateway, gin.H{"error": "proxy not available for node " + taiID}) + return + } + + // path format: /{containerID}:{port}/{rest...} + raw := strings.TrimPrefix(c.Param("path"), "/") + colonIdx := strings.Index(raw, ":") + if colonIdx < 0 { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid proxy path, expected /{containerID}:{port}/{path}"}) + return + } + + containerID := raw[:colonIdx] + rest := raw[colonIdx+1:] + slashIdx := strings.Index(rest, "/") + var portStr, subPath string + if slashIdx >= 0 { + portStr = rest[:slashIdx] + subPath = rest[slashIdx:] + } else { + portStr = rest + subPath = "/" + } + + var port int + for _, ch := range portStr { + if ch < '0' || ch > '9' { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid port in proxy path"}) + return + } + port = port*10 + int(ch-'0') + } + if port == 0 { + c.JSON(http.StatusBadRequest, gin.H{"error": "missing port in proxy path"}) + return + } + + targetURL, err := res.Proxy.URL(c.Request.Context(), containerID, port, subPath) + if err != nil { + c.JSON(http.StatusBadGateway, gin.H{"error": "resolve proxy target: " + err.Error()}) + return + } + + reverseProxy(c, targetURL) +} diff --git a/openapi/tai/tai.go b/openapi/tai/tai.go new file mode 100644 index 00000000..6ddaeba5 --- /dev/null +++ b/openapi/tai/tai.go @@ -0,0 +1,39 @@ +package tai + +import ( + "github.com/gin-gonic/gin" + yaoTai "github.com/yaoapp/yao/tai" + taitunnel "github.com/yaoapp/yao/tai/tunnel" +) + +// Attach registers Tai forward routes on the given group. +// +// - ANY /tai/:taiID/proxy/*path — HTTP forward (tunnel or local) +// - GET /tai/:taiID/vnc/*path — VNC WebSocket forward (tunnel or local) +func Attach(group *gin.RouterGroup) { + group.Any("/tai/:taiID/proxy/*path", handleProxy) + group.GET("/tai/:taiID/vnc/*path", handleVNC) +} + +func handleProxy(c *gin.Context) { + taiID := c.Param("taiID") + if isLocalNode(taiID) { + handleLocalProxy(c, taiID) + return + } + taitunnel.HandleForwardLazy(c) +} + +func handleVNC(c *gin.Context) { + taiID := c.Param("taiID") + if isLocalNode(taiID) { + handleLocalVNC(c, taiID) + return + } + taitunnel.HandleForwardLazy(c) +} + +func isLocalNode(taiID string) bool { + meta, ok := yaoTai.GetNodeMeta(taiID) + return ok && meta.Mode == "local" +} diff --git a/openapi/tai/util.go b/openapi/tai/util.go new file mode 100644 index 00000000..30ae9a5d --- /dev/null +++ b/openapi/tai/util.go @@ -0,0 +1,95 @@ +package tai + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +// extractContainerID parses container ID from *path param. +// /{containerID}/ws → containerID +func extractContainerID(path string) string { + path = strings.TrimPrefix(path, "/") + path = strings.TrimSuffix(path, "/ws") + path = strings.TrimSuffix(path, "/") + if path == "" || path == "__host__" { + return "__host__" + } + return path +} + +// bridgeWebSocket copies messages bidirectionally between two WebSocket connections. +func bridgeWebSocket(client, target *websocket.Conn) { + done := make(chan struct{}, 2) + + go func() { + defer func() { done <- struct{}{} }() + for { + mt, data, err := client.ReadMessage() + if err != nil { + return + } + if err := target.WriteMessage(mt, data); err != nil { + return + } + } + }() + + go func() { + defer func() { done <- struct{}{} }() + for { + mt, data, err := target.ReadMessage() + if err != nil { + return + } + if err := client.WriteMessage(mt, data); err != nil { + return + } + } + }() + + <-done +} + +// reverseProxy forwards an HTTP request to targetURL and streams the response back. +func reverseProxy(c *gin.Context, targetURL string) { + req, err := http.NewRequestWithContext(c.Request.Context(), c.Request.Method, targetURL, c.Request.Body) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "create proxy request: " + err.Error()}) + return + } + for k, vv := range c.Request.Header { + for _, v := range vv { + req.Header.Add(k, v) + } + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + c.JSON(http.StatusBadGateway, gin.H{"error": "proxy request failed: " + err.Error()}) + return + } + defer resp.Body.Close() + + for k, vv := range resp.Header { + for _, v := range vv { + c.Writer.Header().Add(k, v) + } + } + c.Writer.WriteHeader(resp.StatusCode) + c.Writer.Flush() + + buf := make([]byte, 32*1024) + for { + n, readErr := resp.Body.Read(buf) + if n > 0 { + c.Writer.Write(buf[:n]) + c.Writer.Flush() + } + if readErr != nil { + return + } + } +} diff --git a/openapi/tai/vnc.go b/openapi/tai/vnc.go new file mode 100644 index 00000000..472aee39 --- /dev/null +++ b/openapi/tai/vnc.go @@ -0,0 +1,57 @@ +package tai + +import ( + "net/http" + "time" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" + yaoTai "github.com/yaoapp/yao/tai" +) + +var wsUpgrader = websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, + Subprotocols: []string{"binary"}, +} + +// handleLocalVNC resolves the container's VNC address via Docker socket +// and proxies the WebSocket connection. +func handleLocalVNC(c *gin.Context, taiID string) { + containerID := extractContainerID(c.Param("path")) + if containerID == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "missing container ID in path"}) + return + } + + res, ok := yaoTai.GetResources(taiID) + if !ok || res.VNC == nil { + c.JSON(http.StatusBadGateway, gin.H{"error": "VNC not available for node " + taiID}) + return + } + + targetURL, err := res.VNC.URL(c.Request.Context(), containerID) + if err != nil { + c.JSON(http.StatusBadGateway, gin.H{"error": "resolve VNC target: " + err.Error()}) + return + } + + clientConn, err := wsUpgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + return + } + defer clientConn.Close() + + dialer := websocket.Dialer{ + Subprotocols: []string{"binary"}, + HandshakeTimeout: 5 * time.Second, + } + targetConn, _, err := dialer.Dial(targetURL, nil) + if err != nil { + clientConn.WriteMessage(websocket.CloseMessage, + websocket.FormatCloseMessage(websocket.CloseInternalServerErr, "VNC connection failed")) + return + } + defer targetConn.Close() + + bridgeWebSocket(clientConn, targetConn) +} From d26bbd1e1f143b42b3551df01f2f3b4f6f5708ba Mon Sep 17 00:00:00 2001 From: Max Date: Sat, 14 Mar 2026 18:32:09 +0800 Subject: [PATCH 7/8] feat(sandbox): enhance sandbox configuration with display name and identifier updates - Introduced a human-readable display name for the sandbox, constructed from the agent and workspace names. - Updated the BuildIdentifier function to include workspace ID in identifiers for long-running and persistent lifecycles. - Modified sandbox creation options to incorporate the display name, improving clarity in sandbox management. - Refactored related tests to ensure compatibility with the new identifier structure and display name integration. Made-with: Cursor --- agent/assistant/sandbox_v2.go | 41 +++++++++++-- agent/sandbox/v2/lifecycle.go | 17 +++--- agent/sandbox/v2/lifecycle_test.go | 26 ++++----- agent/sandbox/v2/options.go | 1 + agent/sandbox/v2/types/config.go | 1 + openapi/computer/computer.go | 7 ++- openapi/sandbox/manage.go | 7 ++- sandbox/v2/box.go | 2 + sandbox/v2/manager.go | 92 ++++++++++++++++++++++++++---- sandbox/v2/types.go | 2 + tai/runtime/image.go | 9 +++ tai/runtime/image_docker.go | 39 +++++++++++++ tai/runtime/image_k8s.go | 4 ++ tai/tunnel/forward.go | 19 ++++-- tai/tunnel/forward_test.go | 19 ++++-- 15 files changed, 231 insertions(+), 55 deletions(-) diff --git a/agent/assistant/sandbox_v2.go b/agent/assistant/sandbox_v2.go index a8026ae9..5214ade3 100644 --- a/agent/assistant/sandbox_v2.go +++ b/agent/assistant/sandbox_v2.go @@ -48,7 +48,10 @@ func (ast *Assistant) initSandboxV2(ctx *context.Context, opts *context.Options) return nil, nil, nil, "", fmt.Errorf("get connector: %w", err) } - // 2. Obtain Computer (passes connector for OPENAI_PROXY_* env injection). + // 2. Build human-readable DisplayName from real Agent name + Workspace name. + cfg.DisplayName = buildBoxDisplayName(ctx, ast.ID, ast.Name) + + // 3. 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") @@ -56,7 +59,7 @@ func (ast *Assistant) initSandboxV2(ctx *context.Context, opts *context.Options) } _ = identifier - // 3. Get Runner. + // 4. Get Runner. runner, err := sandboxv2.Get(cfg.Runner.Name) if err != nil { sandboxv2.LifecycleAction(stdCtx, cfg, computer, manager) @@ -64,7 +67,7 @@ func (ast *Assistant) initSandboxV2(ctx *context.Context, opts *context.Options) return nil, nil, nil, "", fmt.Errorf("get runner %q: %w", cfg.Runner.Name, err) } - // 4. Resolve skills directory. + // 5. Resolve skills directory. skillsDir := "" if ast.Path != "" { dir := filepath.Join(config.Conf.AppSource, ast.Path, "skills") @@ -73,7 +76,7 @@ func (ast *Assistant) initSandboxV2(ctx *context.Context, opts *context.Options) } } - // 5. Convert MCP servers. + // 6. Convert MCP servers. var mcpServers []sandboxTypes.MCPServer if ast.MCP != nil { for _, s := range ast.MCP.Servers { @@ -85,7 +88,7 @@ func (ast *Assistant) initSandboxV2(ctx *context.Context, opts *context.Options) } } - // 6. Runner.Prepare (standard context). + // 7. Runner.Prepare (standard context). err = runner.Prepare(stdCtx, &sandboxTypes.PrepareRequest{ Computer: computer, Config: cfg, @@ -199,6 +202,34 @@ func (ast *Assistant) initStandaloneWorkspace(ctx *context.Context) { ctx.SetWorkspace(wsFS) } +// buildBoxDisplayName constructs a human-readable display name for a Box +// using the locale-resolved Agent name and Workspace name (matching the UI list pages). +func buildBoxDisplayName(ctx *context.Context, assistantID, rawName string) string { + agentName := i18n.Tr(assistantID, ctx.Locale, rawName) + + wsName := "" + if ctx.Metadata != nil { + if wsID, ok := ctx.Metadata["workspace_id"].(string); ok && wsID != "" { + if wsm := workspace.M(); wsm != nil { + if ws, err := wsm.Get(ctx.Context, wsID); err == nil && ws != nil { + wsName = ws.Name + } + } + } + } + + if agentName != "" && wsName != "" { + return agentName + " / " + wsName + } + if agentName != "" { + return agentName + } + if wsName != "" { + return wsName + } + return "" +} + func closeLoadingV2(ctx *context.Context, loadingMsgID, msgKey string) { if loadingMsgID == "" || ctx == nil { return diff --git a/agent/sandbox/v2/lifecycle.go b/agent/sandbox/v2/lifecycle.go index d2f28800..bad0b57c 100644 --- a/agent/sandbox/v2/lifecycle.go +++ b/agent/sandbox/v2/lifecycle.go @@ -17,7 +17,7 @@ import ( // 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 { +func BuildIdentifier(cfg *types.SandboxConfig, ownerID, chatID, assistantID, workspaceID string, metadata map[string]any) string { if cfg.Lifecycle == "oneshot" { return "" } @@ -25,7 +25,7 @@ func BuildIdentifier(cfg *types.SandboxConfig, ownerID, chatID, assistantID stri // 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) + return fmt.Sprintf("%s-%s.%s", ownerID, cid, workspaceID) } } @@ -33,7 +33,7 @@ func BuildIdentifier(cfg *types.SandboxConfig, ownerID, chatID, assistantID stri case "session": return fmt.Sprintf("%s-%s", ownerID, chatID) case "longrunning", "persistent": - return fmt.Sprintf("%s-%s", ownerID, assistantID) + return fmt.Sprintf("%s-%s.%s", ownerID, assistantID, workspaceID) default: return "" } @@ -44,11 +44,6 @@ func BuildIdentifier(cfg *types.SandboxConfig, ownerID, chatID, assistantID stri // 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 { @@ -59,6 +54,12 @@ func GetComputer(ctx *agentContext.Context, cfg *types.SandboxConfig, manager *i if workspaceID == "" { workspaceID = ownerID } + + identifier := BuildIdentifier(cfg, ownerID, ctx.ChatID, ctx.AssistantID, workspaceID, ctx.Metadata) + + // Fill runtime fields. + cfg.Owner = ownerID + cfg.ID = identifier cfg.WorkspaceID = workspaceID // Resolve computer_id from metadata to determine kind and nodeID. diff --git a/agent/sandbox/v2/lifecycle_test.go b/agent/sandbox/v2/lifecycle_test.go index 64cf2a81..4d1afe4d 100644 --- a/agent/sandbox/v2/lifecycle_test.go +++ b/agent/sandbox/v2/lifecycle_test.go @@ -20,7 +20,7 @@ import ( func TestBuildIdentifier_Oneshot(t *testing.T) { cfg := &types.SandboxConfig{Lifecycle: "oneshot"} - id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", nil) + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", "ws1", nil) if id != "" { t.Errorf("oneshot should return empty, got %q", id) } @@ -28,7 +28,7 @@ func TestBuildIdentifier_Oneshot(t *testing.T) { func TestBuildIdentifier_Session(t *testing.T) { cfg := &types.SandboxConfig{Lifecycle: "session"} - id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat42", "ast1", nil) + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat42", "ast1", "ws1", nil) if id != "owner1-chat42" { t.Errorf("session: got %q, want %q", id, "owner1-chat42") } @@ -36,33 +36,33 @@ func TestBuildIdentifier_Session(t *testing.T) { 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") + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast99", "ws1", nil) + if id != "owner1-ast99.ws1" { + t.Errorf("longrunning: got %q, want %q", id, "owner1-ast99.ws1") } } 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") + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast99", "ws1", nil) + if id != "owner1-ast99.ws1" { + t.Errorf("persistent: got %q, want %q", id, "owner1-ast99.ws1") } } 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") + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", "ws1", meta) + if id != "owner1-custom-box.ws1" { + t.Errorf("metadata override: got %q, want %q", id, "owner1-custom-box.ws1") } } 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) + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat42", "ast1", "ws1", meta) if id != "owner1-chat42" { t.Errorf("empty metadata should fall through to session, got %q", id) } @@ -70,7 +70,7 @@ func TestBuildIdentifier_MetadataEmptyIgnored(t *testing.T) { func TestBuildIdentifier_UnknownLifecycle(t *testing.T) { cfg := &types.SandboxConfig{Lifecycle: "unknown"} - id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", nil) + id := sandboxv2.BuildIdentifier(cfg, "owner1", "chat1", "ast1", "ws1", nil) if id != "" { t.Errorf("unknown lifecycle should return empty, got %q", id) } diff --git a/agent/sandbox/v2/options.go b/agent/sandbox/v2/options.go index 2ceb3392..2ac6ce92 100644 --- a/agent/sandbox/v2/options.go +++ b/agent/sandbox/v2/options.go @@ -34,6 +34,7 @@ func BuildCreateOptions(cfg *types.SandboxConfig, identifier, ownerID, workspace MountMode: cfg.Computer.MountMode, WorkspaceID: workspaceID, Labels: cfg.Labels, + DisplayName: cfg.DisplayName, } if opts.Labels == nil { diff --git a/agent/sandbox/v2/types/config.go b/agent/sandbox/v2/types/config.go index 9eab224e..d5ec1196 100644 --- a/agent/sandbox/v2/types/config.go +++ b/agent/sandbox/v2/types/config.go @@ -32,6 +32,7 @@ type SandboxConfig struct { NodeID string `json:"-" yaml:"-"` Kind string `json:"-" yaml:"-"` WorkspaceID string `json:"-" yaml:"-"` + DisplayName string `json:"-" yaml:"-"` } // ComputerFilter defines the query parameters for GET /computer/options. diff --git a/openapi/computer/computer.go b/openapi/computer/computer.go index e0c80837..93016184 100644 --- a/openapi/computer/computer.go +++ b/openapi/computer/computer.go @@ -45,7 +45,6 @@ type computerOption struct { Image string `json:"image,omitempty"` Policy string `json:"policy,omitempty"` VNC bool `json:"vnc"` - Labels map[string]string `json:"labels,omitempty"` System computerSystemInfo `json:"system"` } @@ -255,7 +254,10 @@ func boxToOption(b *sandboxv2.Box) computerOption { snap := b.Snapshot() info := b.ComputerInfo() - displayName := info.System.Hostname + displayName := info.DisplayName + if displayName == "" { + displayName = info.System.Hostname + } if displayName == "" { displayName = snap.ID } @@ -285,7 +287,6 @@ func boxToOption(b *sandboxv2.Box) computerOption { Image: snap.Image, Policy: string(snap.Policy), VNC: snap.VNC, - Labels: snap.Labels, System: computerSystemInfo{ OS: info.System.OS, Arch: info.System.Arch, diff --git a/openapi/sandbox/manage.go b/openapi/sandbox/manage.go index 270c3d62..316d1c58 100644 --- a/openapi/sandbox/manage.go +++ b/openapi/sandbox/manage.go @@ -94,7 +94,6 @@ type sandboxResponse struct { 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"` @@ -110,7 +109,10 @@ func boxToResponse(b *sandboxv2.Box) sandboxResponse { snap := b.Snapshot() info := b.ComputerInfo() - displayName := info.System.Hostname + displayName := info.DisplayName + if displayName == "" { + displayName = info.System.Hostname + } if displayName == "" { displayName = snap.ID } @@ -137,7 +139,6 @@ func boxToResponse(b *sandboxv2.Box) sandboxResponse { Owner: snap.Owner, Status: snap.Status, Policy: string(snap.Policy), - Labels: snap.Labels, Image: snap.Image, Mode: mode, Addr: addr, diff --git a/sandbox/v2/box.go b/sandbox/v2/box.go index aed45d11..d8412014 100644 --- a/sandbox/v2/box.go +++ b/sandbox/v2/box.go @@ -30,6 +30,7 @@ type Box struct { image string workspaceID string system SystemInfo + displayName string workDir string ws taiworkspace.FS manager *Manager @@ -56,6 +57,7 @@ func (b *Box) ComputerInfo() ComputerInfo { Image: b.image, Policy: b.policy, Labels: b.labels, + DisplayName: b.displayName, } } diff --git a/sandbox/v2/manager.go b/sandbox/v2/manager.go index 3306d95d..4f0fe5aa 100644 --- a/sandbox/v2/manager.go +++ b/sandbox/v2/manager.go @@ -3,7 +3,10 @@ package sandbox import ( "context" "fmt" + "log" "path/filepath" + goruntime "runtime" + "strconv" "sync" "time" @@ -172,7 +175,9 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) return nil, fmt.Errorf("sandbox: node %q has no container runtime", nodeID) } - taiOpts := m.buildTaiCreateOptions(opts, nodeID, id) + sys := inferSystemInfo(ctx, res, opts.Image) + + taiOpts := m.buildTaiCreateOptions(opts, nodeID, id, sys) containerID, err := res.Runtime.Create(ctx, taiOpts) if err != nil { @@ -189,16 +194,6 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) policy = Session } - sys := SystemInfo{ - OS: res.System.OS, - Arch: res.System.Arch, - Hostname: res.System.Hostname, - NumCPU: res.System.NumCPU, - TotalMem: res.System.TotalMem, - Shell: res.System.Shell, - TempDir: res.System.TempDir, - } - boxWorkDir := opts.WorkDir if boxWorkDir == "" { boxWorkDir = "/workspace" @@ -220,6 +215,7 @@ func (m *Manager) Create(ctx context.Context, opts CreateOptions) (*Box, error) image: opts.Image, workspaceID: opts.WorkspaceID, workDir: boxWorkDir, + displayName: opts.DisplayName, system: sys, } box.lastCall.Store(time.Now().UnixMilli()) @@ -348,7 +344,7 @@ func (m *Manager) getNode(name string) (*tai.ConnResources, error) { return res, nil } -func (m *Manager) buildTaiCreateOptions(opts CreateOptions, nodeID, sandboxID string) tairuntime.CreateOptions { +func (m *Manager) buildTaiCreateOptions(opts CreateOptions, nodeID, sandboxID string, sys SystemInfo) tairuntime.CreateOptions { env := make(map[string]string) reg := registry.Global() @@ -378,6 +374,27 @@ func (m *Manager) buildTaiCreateOptions(opts CreateOptions, nodeID, sandboxID st if opts.WorkspaceID != "" { labels["workspace-id"] = opts.WorkspaceID } + if opts.DisplayName != "" { + labels["sandbox-display-name"] = opts.DisplayName + } + if sys.OS != "" { + labels["sandbox-sys-os"] = sys.OS + } + if sys.Arch != "" { + labels["sandbox-sys-arch"] = sys.Arch + } + if sys.Hostname != "" { + labels["sandbox-sys-hostname"] = sys.Hostname + } + if sys.NumCPU > 0 { + labels["sandbox-sys-numcpu"] = strconv.Itoa(sys.NumCPU) + } + if sys.TotalMem > 0 { + labels["sandbox-sys-totalmem"] = strconv.FormatInt(sys.TotalMem, 10) + } + if sys.Shell != "" { + labels["sandbox-sys-shell"] = sys.Shell + } for k, v := range opts.Labels { labels[k] = v } @@ -467,6 +484,10 @@ func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, res *tai.Conn } } } + sys := systemInfoFromLabels(c.Labels) + if sys.OS == "" { + sys = inferSystemInfo(ctx, res, c.Image) + } box := &Box{ id: sandboxID, containerID: cid, @@ -479,6 +500,8 @@ func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, res *tai.Conn workspaceID: c.Labels["workspace-id"], vnc: hasVNC, workDir: "/workspace", + displayName: c.Labels["sandbox-display-name"], + system: sys, manager: m, } box.lastCall.Store(time.Now().UnixMilli()) @@ -486,6 +509,51 @@ func (m *Manager) recoverBoxes(ctx context.Context, nodeID string, res *tai.Conn } } +// inferSystemInfo derives static SystemInfo for a container from image metadata +// and Tai host resources. OS/Arch/Shell come from the image; Hostname/NumCPU/TotalMem +// come from the Tai host. +func inferSystemInfo(ctx context.Context, res *tai.ConnResources, imageRef string) SystemInfo { + sys := SystemInfo{ + Hostname: res.System.Hostname, + NumCPU: res.System.NumCPU, + TotalMem: res.System.TotalMem, + } + + if res.Image != nil { + meta, err := res.Image.Inspect(ctx, imageRef) + if err != nil { + log.Printf("[sandbox/v2] image inspect %q: %v (using fallback)", imageRef, err) + } + if meta != nil { + sys.OS = meta.OS + sys.Arch = meta.Arch + sys.Shell = meta.Shell + return sys + } + } + + sys.OS = "linux" + sys.Arch = goruntime.GOARCH + sys.Shell = "bash" + return sys +} + +// systemInfoFromLabels restores SystemInfo from Docker container labels that +// were persisted at creation time, so recovery doesn't depend on the Tai node +// being connected. +func systemInfoFromLabels(labels map[string]string) SystemInfo { + numCPU, _ := strconv.Atoi(labels["sandbox-sys-numcpu"]) + totalMem, _ := strconv.ParseInt(labels["sandbox-sys-totalmem"], 10, 64) + return SystemInfo{ + OS: labels["sandbox-sys-os"], + Arch: labels["sandbox-sys-arch"], + Hostname: labels["sandbox-sys-hostname"], + NumCPU: numCPU, + TotalMem: totalMem, + Shell: labels["sandbox-sys-shell"], + } +} + // ImageExists reports whether the given image ref exists on the target node. func (m *Manager) ImageExists(ctx context.Context, nodeID, ref string) (bool, error) { res, err := m.getNode(nodeID) diff --git a/sandbox/v2/types.go b/sandbox/v2/types.go index bc78c85a..828ef799 100644 --- a/sandbox/v2/types.go +++ b/sandbox/v2/types.go @@ -44,6 +44,7 @@ type ComputerInfo struct { Image string Policy LifecyclePolicy Labels map[string]string + DisplayName string } // SystemInfo describes the hardware and environment of a Tai node. @@ -104,6 +105,7 @@ type CreateOptions struct { WorkspaceID string MountMode string MountPath string + DisplayName string } type ListOptions struct { diff --git a/tai/runtime/image.go b/tai/runtime/image.go index 45e32209..e18b0e27 100644 --- a/tai/runtime/image.go +++ b/tai/runtime/image.go @@ -8,11 +8,20 @@ import ( // Image manages container images on a runtime node. type Image interface { Exists(ctx context.Context, ref string) (bool, error) + Inspect(ctx context.Context, ref string) (*ImageMeta, error) Pull(ctx context.Context, ref string, opts PullOptions) (<-chan PullProgress, error) Remove(ctx context.Context, ref string, force bool) error List(ctx context.Context) ([]ImageInfo, error) } +// ImageMeta holds static metadata extracted from a container image. +type ImageMeta struct { + OS string // "linux", "windows" + Arch string // "amd64", "arm64" + Shell string // preferred shell: "bash", "sh", "cmd.exe", "pwsh" + WorkDir string // default working directory from Dockerfile WORKDIR +} + // PullOptions configures an image pull operation. type PullOptions struct { Auth *RegistryAuth // nil = anonymous / public diff --git a/tai/runtime/image_docker.go b/tai/runtime/image_docker.go index 09b3edb9..4014bde1 100644 --- a/tai/runtime/image_docker.go +++ b/tai/runtime/image_docker.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "io" + "strings" "time" "github.com/docker/docker/api/types/image" @@ -35,6 +36,44 @@ func (d *dockerImage) Exists(ctx context.Context, ref string) (bool, error) { return true, nil } +func (d *dockerImage) Inspect(ctx context.Context, ref string) (*ImageMeta, error) { + inspect, _, err := d.cli.ImageInspectWithRaw(ctx, ref) + if err != nil { + return nil, fmt.Errorf("image inspect %q: %w", ref, err) + } + + meta := &ImageMeta{ + OS: inspect.Os, + Arch: inspect.Architecture, + } + + if inspect.Config != nil { + meta.WorkDir = inspect.Config.WorkingDir + + if len(inspect.Config.Shell) > 0 { + meta.Shell = inspect.Config.Shell[0] + } + if meta.Shell == "" { + for _, e := range inspect.Config.Env { + if strings.HasPrefix(e, "SHELL=") { + meta.Shell = e[6:] + break + } + } + } + } + + if meta.Shell == "" { + if strings.EqualFold(meta.OS, "windows") { + meta.Shell = "cmd.exe" + } else { + meta.Shell = "bash" + } + } + + return meta, nil +} + func (d *dockerImage) Pull(ctx context.Context, ref string, opts PullOptions) (<-chan PullProgress, error) { pullOpts := image.PullOptions{} if opts.Auth != nil { diff --git a/tai/runtime/image_k8s.go b/tai/runtime/image_k8s.go index bd81a425..2a22952b 100644 --- a/tai/runtime/image_k8s.go +++ b/tai/runtime/image_k8s.go @@ -12,6 +12,10 @@ func (k *k8sImage) Exists(_ context.Context, _ string) (bool, error) { return true, nil } +func (k *k8sImage) Inspect(_ context.Context, _ string) (*ImageMeta, error) { + return nil, nil +} + func (k *k8sImage) Pull(_ context.Context, _ string, _ PullOptions) (<-chan PullProgress, error) { return nil, nil } diff --git a/tai/tunnel/forward.go b/tai/tunnel/forward.go index 1713de90..81ee5896 100644 --- a/tai/tunnel/forward.go +++ b/tai/tunnel/forward.go @@ -113,17 +113,26 @@ func resolveTargetPort(c *gin.Context, node *types.NodeMeta) int { return 0 } -// rewriteRequest clones the request and strips everything up to and including -// /tai/:taiID from the path, handling any baseURL prefix (e.g. /v1/tai/abc/proxy/x → /proxy/x). +// rewriteRequest clones the request and strips the Yao-side route prefix, +// leaving only what the Tai-side handler expects. +// +// The Tai httpproxy expects /{containerID}:{port}/..., so the /proxy prefix +// is stripped. The Tai VNC router expects /vnc/{containerID}/ws, so the /vnc +// prefix is kept. +// +// /v1/tai/abc/proxy/cid:8080/foo → /cid:8080/foo +// /v1/tai/abc/vnc/cid/ws → /vnc/cid/ws func rewriteRequest(orig *http.Request, taiID string) *http.Request { r := orig.Clone(orig.Context()) marker := "/tai/" + taiID if idx := strings.Index(r.URL.Path, marker); idx >= 0 { - r.URL.Path = r.URL.Path[idx+len(marker):] - if r.URL.Path == "" { - r.URL.Path = "/" + rest := r.URL.Path[idx+len(marker):] + rest = strings.TrimPrefix(rest, "/proxy") + if rest == "" { + rest = "/" } + r.URL.Path = rest } r.RequestURI = r.URL.RequestURI() diff --git a/tai/tunnel/forward_test.go b/tai/tunnel/forward_test.go index 3259cc67..703e5198 100644 --- a/tai/tunnel/forward_test.go +++ b/tai/tunnel/forward_test.go @@ -84,8 +84,8 @@ func TestRewriteRequest(t *testing.T) { "proxy_path", "/tai/abc123/proxy/api/v1/data", "abc123", - "/proxy/api/v1/data", - "/proxy/api/v1/data", + "/api/v1/data", + "/api/v1/data", }, { "vnc_path", @@ -98,8 +98,8 @@ func TestRewriteRequest(t *testing.T) { "with_query", "/tai/node-1/proxy/api?foo=bar", "node-1", - "/proxy/api", - "/proxy/api?foo=bar", + "/api", + "/api?foo=bar", }, { "exact_prefix", @@ -112,8 +112,8 @@ func TestRewriteRequest(t *testing.T) { "with_base_url", "/v1/tai/node-1/proxy/api/v1/data", "node-1", - "/proxy/api/v1/data", - "/proxy/api/v1/data", + "/api/v1/data", + "/api/v1/data", }, { "with_base_url_vnc", @@ -122,6 +122,13 @@ func TestRewriteRequest(t *testing.T) { "/vnc/__host__/ws", "/vnc/__host__/ws", }, + { + "proxy_container_port", + "/v1/tai/abc/proxy/cid123:8080/foo", + "abc", + "/cid123:8080/foo", + "/cid123:8080/foo", + }, { "no_match", "/other/path", From 9c9701ed7e052a147aed3d9ce95b6659b9f6f7e8 Mon Sep 17 00:00:00 2001 From: Max Date: Sat, 14 Mar 2026 21:51:32 +0800 Subject: [PATCH 8/8] feat(tunnel): refactor VNC and proxy handling with structured routing - Updated VNC and proxy handling in the tunnel to utilize a structured routing approach, enhancing clarity and maintainability. - Replaced direct port checks with a new `forwardRoute` struct to encapsulate routing information, including channel type, container ID, and port. - Modified request handling to streamline the forwarding process and improve error handling for unknown routes. - Enhanced tests to validate the new routing logic and ensure consistent behavior across VNC and proxy requests. Made-with: Cursor --- openapi/computer/computer.go | 4 +- openapi/sandbox/manage.go | 2 +- tai/api/register.go | 1 + tai/tunnel/forward.go | 148 ++++++++++++------ tai/tunnel/forward_test.go | 240 +++++++++++++---------------- tai/tunnel/grpc_handler.go | 82 ++++++++-- tai/tunnel/grpc_handler_test.go | 16 +- tai/tunnel/taipb/tunnel.pb.go | 126 +++++++++------ tai/tunnel/taipb/tunnel_grpc.pb.go | 4 +- tai/types/types.go | 1 + 10 files changed, 379 insertions(+), 245 deletions(-) diff --git a/openapi/computer/computer.go b/openapi/computer/computer.go index 93016184..09311000 100644 --- a/openapi/computer/computer.go +++ b/openapi/computer/computer.go @@ -203,7 +203,7 @@ func nodeToHostOption(s taitypes.NodeMeta) computerOption { Status: status, Mode: s.Mode, Addr: addr, - VNC: s.Ports.VNC > 0, + VNC: s.Capabilities.VNC, System: computerSystemInfo{ OS: s.System.OS, Arch: s.System.Arch, @@ -239,7 +239,7 @@ func nodeToNodeOption(s taitypes.NodeMeta) computerOption { Status: status, Mode: s.Mode, Addr: addr, - VNC: s.Ports.VNC > 0, + VNC: s.Capabilities.VNC, System: computerSystemInfo{ OS: s.System.OS, Arch: s.System.Arch, diff --git a/openapi/sandbox/manage.go b/openapi/sandbox/manage.go index 316d1c58..3df56482 100644 --- a/openapi/sandbox/manage.go +++ b/openapi/sandbox/manage.go @@ -197,7 +197,7 @@ func hostToResponse(s taitypes.NodeMeta) sandboxResponse { Policy: "persistent", Mode: s.Mode, Addr: addr, - VNC: s.Ports.VNC > 0, + VNC: s.Capabilities.VNC, CreatedAt: s.ConnectedAt, LastActive: s.LastPing, System: sandboxSystemInfo{ diff --git a/tai/api/register.go b/tai/api/register.go index 27378169..fa6a6d59 100644 --- a/tai/api/register.go +++ b/tai/api/register.go @@ -294,6 +294,7 @@ func capsFromMap(m map[string]bool) types.Capabilities { Docker: m["docker"], K8s: m["k8s"], HostExec: m["host_exec"], + VNC: m["vnc"], } } diff --git a/tai/tunnel/forward.go b/tai/tunnel/forward.go index 81ee5896..4bd46302 100644 --- a/tai/tunnel/forward.go +++ b/tai/tunnel/forward.go @@ -5,6 +5,7 @@ import ( "fmt" "io" "net/http" + "strconv" "strings" "github.com/gin-gonic/gin" @@ -12,11 +13,24 @@ import ( "github.com/yaoapp/yao/tai/types" ) +const defaultVNCPort = 5900 + +// forwardRoute holds the structured routing information extracted from the +// incoming request URL. It is passed to RequestForward so that Yao can +// populate the TunnelControl proto fields and Tai can route directly without +// parsing the first packet. +type forwardRoute struct { + channelType string // "proxy" | "vnc" + containerID string // target container or "__host__" + containerPort int // container-internal port (vnc default 5900) + subpath string // rewritten request path for the container +} + // HandleForward handles HTTP/VNC/any TCP-level forwarding through the gRPC tunnel. // Route: ANY /tai/:taiID/proxy/*path and GET /tai/:taiID/vnc/*path // // It hijacks the browser's raw TCP connection, asks Tai to open a Forward stream -// to the resolved target port, rewrites the request path, and then performs +// with explicit routing information, rewrites the request path, and then performs // bidirectional byte-level bridging. No protocol parsing beyond HTTP hijack. func (h *TunnelHandler) HandleForward(c *gin.Context) { logger := h.logger @@ -30,16 +44,18 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { return } - targetPort := resolveTargetPort(c, node) - if targetPort == 0 { - c.JSON(http.StatusBadRequest, gin.H{"error": "cannot resolve target port"}) + route, err := resolveRoute(c, node) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - rewrittenReq := rewriteRequest(c.Request, taiID) - logger.Debug("[forward] "+node.Mode+" → tai:"+fmt.Sprintf("%d", targetPort), + rewrittenReq := rewriteRequest(c.Request, taiID, route) + logger.Debug("[forward] "+node.Mode+" → tai", "tai_id", taiID, - "addr", node.Addr, + "type", route.channelType, + "container", route.containerID, + "container_port", route.containerPort, "path", rewrittenReq.URL.Path, ) @@ -55,10 +71,10 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { } defer browserConn.Close() - fwd, err := h.RequestForward(taiID, targetPort) + fwd, err := h.RequestForward(taiID, route) if err != nil { logger.Error("[forward] stream failed", - "tai_id", taiID, "port", targetPort, "err", err) + "tai_id", taiID, "type", route.channelType, "err", err) browserConn.Write([]byte("HTTP/1.1 502 Bad Gateway\r\n\r\n")) return } @@ -94,47 +110,91 @@ func HandleForwardLazy(c *gin.Context) { h.HandleForward(c) } -// resolveTargetPort determines the Tai-side port from the route pattern. -func resolveTargetPort(c *gin.Context, node *types.NodeMeta) int { +// resolveRoute extracts structured routing info from the request URL path. +// +// For proxy requests (/tai/:taiID/proxy/{containerID}:{port}/{subpath}): +// +// channelType = "proxy", containerPort from URL, subpath = remaining path. +// +// For VNC requests (/tai/:taiID/vnc/{containerID}/ws): +// +// channelType = "vnc", containerPort = 5900, subpath = /vnc/{containerID}/ws. +func resolveRoute(c *gin.Context, node *types.NodeMeta) (*forwardRoute, error) { path := c.Request.URL.Path - - if strings.Contains(path, "/vnc/") { - if node.Ports.VNC != 0 { - return node.Ports.VNC - } - return 16080 - } - if strings.Contains(path, "/proxy/") { - if node.Ports.HTTP != 0 { - return node.Ports.HTTP - } - return 8099 - } - return 0 -} - -// rewriteRequest clones the request and strips the Yao-side route prefix, -// leaving only what the Tai-side handler expects. -// -// The Tai httpproxy expects /{containerID}:{port}/..., so the /proxy prefix -// is stripped. The Tai VNC router expects /vnc/{containerID}/ws, so the /vnc -// prefix is kept. -// -// /v1/tai/abc/proxy/cid:8080/foo → /cid:8080/foo -// /v1/tai/abc/vnc/cid/ws → /vnc/cid/ws -func rewriteRequest(orig *http.Request, taiID string) *http.Request { - r := orig.Clone(orig.Context()) + taiID := c.Param("taiID") marker := "/tai/" + taiID - if idx := strings.Index(r.URL.Path, marker); idx >= 0 { - rest := r.URL.Path[idx+len(marker):] - rest = strings.TrimPrefix(rest, "/proxy") - if rest == "" { - rest = "/" + idx := strings.Index(path, marker) + if idx < 0 { + return nil, fmt.Errorf("cannot locate /tai/%s in path", taiID) + } + rest := path[idx+len(marker):] + + if strings.HasPrefix(rest, "/vnc/") { + // /vnc/{containerID}/ws → containerID, port=5900 + tail := strings.TrimPrefix(rest, "/vnc/") + containerID := tail + if slashIdx := strings.IndexByte(tail, '/'); slashIdx >= 0 { + containerID = tail[:slashIdx] } - r.URL.Path = rest + if containerID == "" { + return nil, fmt.Errorf("missing container ID in VNC path: %s", path) + } + return &forwardRoute{ + channelType: "vnc", + containerID: containerID, + containerPort: defaultVNCPort, + subpath: rest, // keep /vnc/{containerID}/ws + }, nil } + if strings.HasPrefix(rest, "/proxy/") { + // /proxy/{containerID}:{port}/{subpath} + proxyPath := strings.TrimPrefix(rest, "/proxy") + // proxyPath = /{containerID}:{port}/{subpath} + proxyPath = strings.TrimPrefix(proxyPath, "/") + if proxyPath == "" { + return nil, fmt.Errorf("empty proxy path") + } + + slash := strings.IndexByte(proxyPath, '/') + var head, subpath string + if slash == -1 { + head = proxyPath + subpath = "/" + } else { + head = proxyPath[:slash] + subpath = proxyPath[slash:] + } + + colon := strings.LastIndexByte(head, ':') + if colon < 0 { + return nil, fmt.Errorf("missing port in proxy path: %s", path) + } + containerID := head[:colon] + portStr := head[colon+1:] + port, err := strconv.Atoi(portStr) + if err != nil { + return nil, fmt.Errorf("invalid port %q in proxy path: %w", portStr, err) + } + return &forwardRoute{ + channelType: "proxy", + containerID: containerID, + containerPort: port, + subpath: subpath, + }, nil + } + + return nil, fmt.Errorf("unknown route pattern: %s", rest) +} + +// rewriteRequest clones the request and sets the path to the route's subpath. +// +// For proxy: the path becomes the subpath (e.g. /foo/bar). +// For VNC: the path keeps /vnc/{containerID}/ws as-is. +func rewriteRequest(orig *http.Request, taiID string, route *forwardRoute) *http.Request { + r := orig.Clone(orig.Context()) + r.URL.Path = route.subpath r.RequestURI = r.URL.RequestURI() return r } diff --git a/tai/tunnel/forward_test.go b/tai/tunnel/forward_test.go index 703e5198..834c33ff 100644 --- a/tai/tunnel/forward_test.go +++ b/tai/tunnel/forward_test.go @@ -16,155 +16,131 @@ func init() { gin.SetMode(gin.TestMode) } -func TestResolveTargetPort_VNC(t *testing.T) { +func TestResolveRoute_Proxy(t *testing.T) { tests := []struct { - name string - path string - vncPort int - wantPort int + name string + path string + wantType string + wantContainer string + wantPort int + wantSubpath string }{ - {"default_vnc", "/tai/abc/vnc/websockify", 0, 16080}, - {"custom_vnc", "/tai/abc/vnc/websockify", 5900, 5900}, + { + "basic_proxy", + "/tai/abc/proxy/cid123:8080/foo/bar", + "proxy", "cid123", 8080, "/foo/bar", + }, + { + "proxy_root", + "/tai/abc/proxy/cid:3000", + "proxy", "cid", 3000, "/", + }, + { + "proxy_host", + "/v1/tai/abc/proxy/__host__:9090/api", + "proxy", "__host__", 9090, "/api", + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { c, _ := gin.CreateTestContext(httptest.NewRecorder()) c.Request = &http.Request{URL: &url.URL{Path: tt.path}} - node := &types.NodeMeta{Ports: types.Ports{VNC: tt.vncPort}} - got := resolveTargetPort(c, node) - if got != tt.wantPort { - t.Errorf("resolveTargetPort = %d, want %d", got, tt.wantPort) + c.Params = gin.Params{{Key: "taiID", Value: "abc"}} + node := &types.NodeMeta{} + + r, err := resolveRoute(c, node) + if err != nil { + t.Fatalf("resolveRoute error: %v", err) + } + if r.channelType != tt.wantType { + t.Errorf("channelType = %q, want %q", r.channelType, tt.wantType) + } + if r.containerID != tt.wantContainer { + t.Errorf("containerID = %q, want %q", r.containerID, tt.wantContainer) + } + if r.containerPort != tt.wantPort { + t.Errorf("containerPort = %d, want %d", r.containerPort, tt.wantPort) + } + if r.subpath != tt.wantSubpath { + t.Errorf("subpath = %q, want %q", r.subpath, tt.wantSubpath) } }) } } -func TestResolveTargetPort_Proxy(t *testing.T) { +func TestResolveRoute_VNC(t *testing.T) { tests := []struct { - name string - path string - httpPort int - wantPort int + name string + path string + wantContainer string + wantPort int }{ - {"default_proxy", "/tai/abc/proxy/api/v1/foo", 0, 8099}, - {"custom_proxy", "/tai/abc/proxy/api/v1/foo", 9090, 9090}, + {"vnc_basic", "/tai/abc/vnc/container1/ws", "container1", defaultVNCPort}, + {"vnc_host", "/v1/tai/abc/vnc/__host__/ws", "__host__", defaultVNCPort}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { c, _ := gin.CreateTestContext(httptest.NewRecorder()) c.Request = &http.Request{URL: &url.URL{Path: tt.path}} - node := &types.NodeMeta{Ports: types.Ports{HTTP: tt.httpPort}} - got := resolveTargetPort(c, node) - if got != tt.wantPort { - t.Errorf("resolveTargetPort = %d, want %d", got, tt.wantPort) + c.Params = gin.Params{{Key: "taiID", Value: "abc"}} + node := &types.NodeMeta{} + + r, err := resolveRoute(c, node) + if err != nil { + t.Fatalf("resolveRoute error: %v", err) + } + if r.channelType != "vnc" { + t.Errorf("channelType = %q, want vnc", r.channelType) + } + if r.containerID != tt.wantContainer { + t.Errorf("containerID = %q, want %q", r.containerID, tt.wantContainer) + } + if r.containerPort != tt.wantPort { + t.Errorf("containerPort = %d, want %d", r.containerPort, tt.wantPort) } }) } } -func TestResolveTargetPort_Unknown(t *testing.T) { +func TestResolveRoute_Unknown(t *testing.T) { c, _ := gin.CreateTestContext(httptest.NewRecorder()) c.Request = &http.Request{URL: &url.URL{Path: "/tai/abc/unknown/something"}} + c.Params = gin.Params{{Key: "taiID", Value: "abc"}} node := &types.NodeMeta{} - got := resolveTargetPort(c, node) - if got != 0 { - t.Errorf("resolveTargetPort = %d, want 0", got) + + _, err := resolveRoute(c, node) + if err == nil { + t.Error("expected error for unknown route") } } -func TestRewriteRequest(t *testing.T) { - tests := []struct { - name string - origPath string - taiID string - wantPath string - wantURI string - }{ - { - "proxy_path", - "/tai/abc123/proxy/api/v1/data", - "abc123", - "/api/v1/data", - "/api/v1/data", - }, - { - "vnc_path", - "/tai/node-1/vnc/websockify", - "node-1", - "/vnc/websockify", - "/vnc/websockify", - }, - { - "with_query", - "/tai/node-1/proxy/api?foo=bar", - "node-1", - "/api", - "/api?foo=bar", - }, - { - "exact_prefix", - "/tai/node-1", - "node-1", - "/", - "/", - }, - { - "with_base_url", - "/v1/tai/node-1/proxy/api/v1/data", - "node-1", - "/api/v1/data", - "/api/v1/data", - }, - { - "with_base_url_vnc", - "/v1/tai/abc123/vnc/__host__/ws", - "abc123", - "/vnc/__host__/ws", - "/vnc/__host__/ws", - }, - { - "proxy_container_port", - "/v1/tai/abc/proxy/cid123:8080/foo", - "abc", - "/cid123:8080/foo", - "/cid123:8080/foo", - }, - { - "no_match", - "/other/path", - "node-1", - "/other/path", - "/other/path", - }, +func TestRewriteRequest_Proxy(t *testing.T) { + u, _ := url.Parse("http://localhost/v1/tai/abc/proxy/cid:8080/foo") + orig := &http.Request{ + Method: "GET", + URL: u, + RequestURI: u.RequestURI(), + Host: "localhost", + Header: http.Header{}, + } + route := &forwardRoute{ + channelType: "proxy", + containerID: "cid", + containerPort: 8080, + subpath: "/foo", } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - u, _ := url.Parse("http://localhost" + tt.origPath) - orig := &http.Request{ - Method: "GET", - URL: u, - RequestURI: u.RequestURI(), - Host: "localhost", - Header: http.Header{}, - } - got := rewriteRequest(orig, tt.taiID) - - if got.URL.Path != tt.wantPath { - t.Errorf("path = %q, want %q", got.URL.Path, tt.wantPath) - } - if got.RequestURI != tt.wantURI { - t.Errorf("requestURI = %q, want %q", got.RequestURI, tt.wantURI) - } - if got == orig { - t.Error("rewriteRequest should return a clone, not the original") - } - }) + got := rewriteRequest(orig, "abc", route) + if got.URL.Path != "/foo" { + t.Errorf("path = %q, want /foo", got.URL.Path) + } + if got == orig { + t.Error("rewriteRequest should return a clone") } } -func TestRewriteRequest_PreservesHeaders(t *testing.T) { - u, _ := url.Parse("http://localhost/tai/node-1/vnc/websockify") +func TestRewriteRequest_VNC(t *testing.T) { + u, _ := url.Parse("http://localhost/tai/node-1/vnc/cid/ws") orig := &http.Request{ Method: "GET", URL: u, @@ -175,14 +151,20 @@ func TestRewriteRequest_PreservesHeaders(t *testing.T) { "Upgrade": {"websocket"}, }, } + route := &forwardRoute{ + channelType: "vnc", + containerID: "cid", + containerPort: 5900, + subpath: "/vnc/cid/ws", + } - got := rewriteRequest(orig, "node-1") + got := rewriteRequest(orig, "node-1", route) + if got.URL.Path != "/vnc/cid/ws" { + t.Errorf("path = %q, want /vnc/cid/ws", got.URL.Path) + } if got.Header.Get("Connection") != "Upgrade" { t.Error("expected Connection header preserved") } - if got.Header.Get("Upgrade") != "websocket" { - t.Error("expected Upgrade header preserved") - } } func TestHandleForwardLazy_NilHandler(t *testing.T) { @@ -217,29 +199,24 @@ func TestHandleForward_NodeNotFound(t *testing.T) { } } -func TestHandleForward_NodeOffline(t *testing.T) { +func TestHandleForward_UnknownRoute(t *testing.T) { reg := registry.NewForTest() h := NewTunnelHandler(reg) reg.Register(®istry.TaiNode{ - TaiID: "offline-node", + TaiID: "online-node", Mode: "tunnel", - Ports: types.Ports{HTTP: 8099}, }) - // Manually set status to offline via a Get() — the node is online by default - // after Register, but we need an offline one. We'll use Unregister + re-register - // pattern. Actually, let's just test with a node that doesn't exist: - // the NodeNotFound test above covers that case. Instead, test zero port. w := httptest.NewRecorder() c, _ := gin.CreateTestContext(w) - c.Request = httptest.NewRequest("GET", "/tai/offline-node/unknown/foo", nil) - c.Params = gin.Params{{Key: "taiID", Value: "offline-node"}} + c.Request = httptest.NewRequest("GET", "/tai/online-node/unknown/foo", nil) + c.Params = gin.Params{{Key: "taiID", Value: "online-node"}} h.HandleForward(c) if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 for unresolvable port, got %d", w.Code) + t.Errorf("expected 400 for unresolvable route, got %d", w.Code) } } @@ -268,7 +245,6 @@ func TestHandleForward_ViaRealHTTP(t *testing.T) { reg.Register(®istry.TaiNode{ TaiID: "http-node", Mode: "tunnel", - Ports: types.Ports{HTTP: 8099}, }) router := gin.New() @@ -277,16 +253,12 @@ func TestHandleForward_ViaRealHTTP(t *testing.T) { srv := httptest.NewServer(router) defer srv.Close() - resp, err := http.Get(srv.URL + "/tai/http-node/proxy/api") + resp, err := http.Get(srv.URL + "/tai/http-node/proxy/cid:8080/api") if err != nil { t.Fatal(err) } defer resp.Body.Close() - // RequestForward will fail (no register stream) → hijacked conn gets "502" - // or the response will be a 502 written before hijack. - // Since hijack happens, the actual HTTP status may not be set normally. - // We just verify no panic and the request completes. if resp.StatusCode == 200 { t.Error("expected non-200 response for failed forward") } diff --git a/tai/tunnel/grpc_handler.go b/tai/tunnel/grpc_handler.go index bdf39deb..d7fb41c7 100644 --- a/tai/tunnel/grpc_handler.go +++ b/tai/tunnel/grpc_handler.go @@ -191,7 +191,9 @@ func (h *TunnelHandler) Forward(stream taipb.TaiTunnel_ForwardServer) error { // RequestForward sends an "open" command to Tai via the Register stream and // waits for Tai to call back with a Forward stream. Returns the Forward stream. -func (h *TunnelHandler) RequestForward(taiID string, targetPort int) (taipb.TaiTunnel_ForwardServer, error) { +// +// route may be nil for raw TCP tunnels (gRPC, Docker API, K8s API). +func (h *TunnelHandler) RequestForward(taiID string, route *forwardRoute) (taipb.TaiTunnel_ForwardServer, error) { stream := h.reg.GetRegisterStream(taiID) if stream == nil { return nil, fmt.Errorf("tai %s: no active register stream", taiID) @@ -217,23 +219,30 @@ func (h *TunnelHandler) RequestForward(taiID string, targetPort int) (taipb.TaiT return nil, fmt.Errorf("tai %s: register stream type mismatch", taiID) } + ctrl := &taipb.TunnelControl{ + Type: "open", + ChannelId: channelID, + } + if route != nil { + ctrl.ChannelType = route.channelType + ctrl.ContainerId = route.containerID + ctrl.ContainerPort = int32(route.containerPort) + } + short := registry.ShortChannelID(channelID) h.logger.Debug("[forward] sending open command", - "tai_id", taiID, "port", targetPort, "channel_id", short) + "tai_id", taiID, "channel_type", ctrl.ChannelType, + "container", ctrl.ContainerId, "channel_id", short) mu.Lock() - sendErr := regStream.Send(&taipb.TunnelControl{ - Type: "open", - ChannelId: channelID, - TargetPort: int32(targetPort), - }) + sendErr := regStream.Send(ctrl) mu.Unlock() if sendErr != nil { return nil, fmt.Errorf("send open: %w", sendErr) } h.logger.Debug("[forward] open sent, waiting for callback", - "tai_id", taiID, "port", targetPort, "channel_id", short) + "tai_id", taiID, "channel_id", short) select { case fwd := <-waitCh: @@ -261,8 +270,9 @@ func (h *TunnelHandler) connectTunnelNode(taiID string) { // bridgeConn bridges a local TCP connection to a Tai port via gRPC Forward stream. // Called by registry.OpenLocalListener for each accepted TCP connection. +// Uses raw TCP forwarding (TargetPort only, no container routing). func (h *TunnelHandler) bridgeConn(taiID string, targetPort int, localConn net.Conn) { - fwd, err := h.RequestForward(taiID, targetPort) + fwd, err := h.requestForwardRaw(taiID, targetPort) if err != nil { localConn.Close() h.logger.Error("request forward failed", @@ -274,6 +284,59 @@ func (h *TunnelHandler) bridgeConn(taiID string, targetPort int, localConn net.C bridgeTCP(localConn, streamConn) } +// requestForwardRaw sends an "open" command with only TargetPort (no container +// routing). Used by bridgeConn for raw TCP tunnels (gRPC, Docker API, K8s API). +func (h *TunnelHandler) requestForwardRaw(taiID string, targetPort int) (taipb.TaiTunnel_ForwardServer, error) { + stream := h.reg.GetRegisterStream(taiID) + if stream == nil { + return nil, fmt.Errorf("tai %s: no active register stream", taiID) + } + + muVal, ok := h.sendMu.Load(taiID) + if !ok { + return nil, fmt.Errorf("tai %s: no send mutex (stream closing?)", taiID) + } + mu := muVal.(*sync.Mutex) + + channelID, err := registry.GenerateChannelID() + if err != nil { + return nil, fmt.Errorf("generate channel_id: %w", err) + } + + waitCh := make(chan taipb.TaiTunnel_ForwardServer, 1) + h.pending.Store(channelID, waitCh) + defer h.pending.Delete(channelID) + + regStream, ok := stream.(taipb.TaiTunnel_RegisterServer) + if !ok { + return nil, fmt.Errorf("tai %s: register stream type mismatch", taiID) + } + + short := registry.ShortChannelID(channelID) + h.logger.Debug("[forward] sending open command (raw)", + "tai_id", taiID, "port", targetPort, "channel_id", short) + + mu.Lock() + sendErr := regStream.Send(&taipb.TunnelControl{ + Type: "open", + ChannelId: channelID, + TargetPort: int32(targetPort), + }) + mu.Unlock() + if sendErr != nil { + return nil, fmt.Errorf("send open: %w", sendErr) + } + + select { + case fwd := <-waitCh: + return fwd, nil + case <-time.After(10 * time.Second): + return nil, fmt.Errorf("tai %s: forward timeout (10s) channel=%s", taiID, short) + case <-regStream.Context().Done(): + return nil, fmt.Errorf("tai %s: register stream closed while waiting for forward", taiID) + } +} + // forwardConn wraps a Forward stream as a net.Conn-like reader/writer. type forwardConn struct { stream taipb.TaiTunnel_ForwardServer @@ -364,6 +427,7 @@ func capsFromProto(c *taipb.Capabilities) types.Capabilities { Docker: c.Docker, K8s: c.K8S, HostExec: c.HostExec, + VNC: c.Vnc, } } diff --git a/tai/tunnel/grpc_handler_test.go b/tai/tunnel/grpc_handler_test.go index 278dd300..b5d97558 100644 --- a/tai/tunnel/grpc_handler_test.go +++ b/tai/tunnel/grpc_handler_test.go @@ -245,7 +245,7 @@ func TestRequestForward_NoRegisterStream(t *testing.T) { reg.Register(®istry.TaiNode{TaiID: "no-stream", Mode: "tunnel"}) - _, err := h.RequestForward("no-stream", 8099) + _, err := h.requestForwardRaw("no-stream", 8099) if err == nil { t.Fatal("expected error when no register stream") } @@ -258,7 +258,7 @@ func TestRequestForward_TypeMismatch(t *testing.T) { reg.Register(®istry.TaiNode{TaiID: "bad-type", Mode: "tunnel"}) reg.SetRegisterStream("bad-type", "not-a-stream") - _, err := h.RequestForward("bad-type", 8099) + _, err := h.requestForwardRaw("bad-type", 8099) if err == nil { t.Fatal("expected error for type mismatch") } @@ -348,7 +348,7 @@ drainLoop: requestDone.Add(1) go func() { defer requestDone.Done() - requestResult, requestErr = h.RequestForward(taiID, 8099) + requestResult, requestErr = h.requestForwardRaw(taiID, 8099) }() // Receive the "open" command @@ -748,7 +748,7 @@ func TestRequestForward_Timeout(t *testing.T) { // by never sending Forward). We'll use a short context cancel to avoid waiting. done := make(chan error, 1) go func() { - _, err := h.RequestForward(taiID, 8099) + _, err := h.requestForwardRaw(taiID, 8099) done <- err }() @@ -828,7 +828,7 @@ drained: for i := 0; i < N; i++ { port := 8099 + i go func(port int) { - _, err := h.RequestForward(taiID, port) + _, err := h.requestForwardRaw(taiID, port) results <- err }(port) } @@ -926,7 +926,7 @@ drained2: // Start RequestForward fwdResult := make(chan error, 1) go func() { - _, err := h.RequestForward(taiID, 8099) + _, err := h.requestForwardRaw(taiID, 8099) fwdResult <- err }() @@ -1101,7 +1101,7 @@ proxyDrained: }() // Now do an actual RequestForward + simulate browser side - fwd, err := h.RequestForward(taiID, 8099) + fwd, err := h.requestForwardRaw(taiID, 8099) if err != nil { t.Fatal("RequestForward:", err) } @@ -1267,7 +1267,7 @@ vncDrained: }() // Send WS upgrade request through tunnel - fwd, err := h.RequestForward(taiID, 16080) + fwd, err := h.requestForwardRaw(taiID, 16080) if err != nil { t.Fatal("RequestForward:", err) } diff --git a/tai/tunnel/taipb/tunnel.pb.go b/tai/tunnel/taipb/tunnel.pb.go index 33fe5105..ddc6cd8d 100644 --- a/tai/tunnel/taipb/tunnel.pb.go +++ b/tai/tunnel/taipb/tunnel.pb.go @@ -2,7 +2,7 @@ // versions: // protoc-gen-go v1.36.11 // protoc v4.25.0 -// source: tunnel.proto +// source: tunnel/proto/tunnel.proto package taipb @@ -33,8 +33,11 @@ type TunnelControl struct { Caps *Capabilities `protobuf:"bytes,7,opt,name=caps,proto3" json:"caps,omitempty"` System *SystemInfo `protobuf:"bytes,8,opt,name=system,proto3" json:"system,omitempty"` // Carried on "open" (Yao → Tai) - ChannelId string `protobuf:"bytes,10,opt,name=channel_id,json=channelId,proto3" json:"channel_id,omitempty"` - TargetPort int32 `protobuf:"varint,11,opt,name=target_port,json=targetPort,proto3" json:"target_port,omitempty"` + ChannelId string `protobuf:"bytes,10,opt,name=channel_id,json=channelId,proto3" json:"channel_id,omitempty"` + TargetPort int32 `protobuf:"varint,11,opt,name=target_port,json=targetPort,proto3" json:"target_port,omitempty"` + ChannelType string `protobuf:"bytes,12,opt,name=channel_type,json=channelType,proto3" json:"channel_type,omitempty"` // "proxy" | "vnc" | "" (legacy/raw TCP) + ContainerId string `protobuf:"bytes,13,opt,name=container_id,json=containerId,proto3" json:"container_id,omitempty"` // target container or "__host__" + ContainerPort int32 `protobuf:"varint,14,opt,name=container_port,json=containerPort,proto3" json:"container_port,omitempty"` // container-internal port (vnc default 5900) // Carried on "registered" (Yao → Tai) TaiId string `protobuf:"bytes,20,opt,name=tai_id,json=taiId,proto3" json:"tai_id,omitempty"` unknownFields protoimpl.UnknownFields @@ -43,7 +46,7 @@ type TunnelControl struct { func (x *TunnelControl) Reset() { *x = TunnelControl{} - mi := &file_tunnel_proto_msgTypes[0] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[0] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -55,7 +58,7 @@ func (x *TunnelControl) String() string { func (*TunnelControl) ProtoMessage() {} func (x *TunnelControl) ProtoReflect() protoreflect.Message { - mi := &file_tunnel_proto_msgTypes[0] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[0] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -68,7 +71,7 @@ func (x *TunnelControl) ProtoReflect() protoreflect.Message { // Deprecated: Use TunnelControl.ProtoReflect.Descriptor instead. func (*TunnelControl) Descriptor() ([]byte, []int) { - return file_tunnel_proto_rawDescGZIP(), []int{0} + return file_tunnel_proto_tunnel_proto_rawDescGZIP(), []int{0} } func (x *TunnelControl) GetType() string { @@ -141,6 +144,27 @@ func (x *TunnelControl) GetTargetPort() int32 { return 0 } +func (x *TunnelControl) GetChannelType() string { + if x != nil { + return x.ChannelType + } + return "" +} + +func (x *TunnelControl) GetContainerId() string { + if x != nil { + return x.ContainerId + } + return "" +} + +func (x *TunnelControl) GetContainerPort() int32 { + if x != nil { + return x.ContainerPort + } + return 0 +} + func (x *TunnelControl) GetTaiId() string { if x != nil { return x.TaiId @@ -157,7 +181,7 @@ type ForwardData struct { func (x *ForwardData) Reset() { *x = ForwardData{} - mi := &file_tunnel_proto_msgTypes[1] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[1] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -169,7 +193,7 @@ func (x *ForwardData) String() string { func (*ForwardData) ProtoMessage() {} func (x *ForwardData) ProtoReflect() protoreflect.Message { - mi := &file_tunnel_proto_msgTypes[1] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[1] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -182,7 +206,7 @@ func (x *ForwardData) ProtoReflect() protoreflect.Message { // Deprecated: Use ForwardData.ProtoReflect.Descriptor instead. func (*ForwardData) Descriptor() ([]byte, []int) { - return file_tunnel_proto_rawDescGZIP(), []int{1} + return file_tunnel_proto_tunnel_proto_rawDescGZIP(), []int{1} } func (x *ForwardData) GetData() []byte { @@ -205,7 +229,7 @@ type Ports struct { func (x *Ports) Reset() { *x = Ports{} - mi := &file_tunnel_proto_msgTypes[2] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[2] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -217,7 +241,7 @@ func (x *Ports) String() string { func (*Ports) ProtoMessage() {} func (x *Ports) ProtoReflect() protoreflect.Message { - mi := &file_tunnel_proto_msgTypes[2] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[2] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -230,7 +254,7 @@ func (x *Ports) ProtoReflect() protoreflect.Message { // Deprecated: Use Ports.ProtoReflect.Descriptor instead. func (*Ports) Descriptor() ([]byte, []int) { - return file_tunnel_proto_rawDescGZIP(), []int{2} + return file_tunnel_proto_tunnel_proto_rawDescGZIP(), []int{2} } func (x *Ports) GetGrpc() int32 { @@ -273,13 +297,14 @@ type Capabilities struct { Docker bool `protobuf:"varint,1,opt,name=docker,proto3" json:"docker,omitempty"` K8S bool `protobuf:"varint,2,opt,name=k8s,proto3" json:"k8s,omitempty"` HostExec bool `protobuf:"varint,3,opt,name=host_exec,json=hostExec,proto3" json:"host_exec,omitempty"` + Vnc bool `protobuf:"varint,4,opt,name=vnc,proto3" json:"vnc,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } func (x *Capabilities) Reset() { *x = Capabilities{} - mi := &file_tunnel_proto_msgTypes[3] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[3] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -291,7 +316,7 @@ func (x *Capabilities) String() string { func (*Capabilities) ProtoMessage() {} func (x *Capabilities) ProtoReflect() protoreflect.Message { - mi := &file_tunnel_proto_msgTypes[3] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[3] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -304,7 +329,7 @@ func (x *Capabilities) ProtoReflect() protoreflect.Message { // Deprecated: Use Capabilities.ProtoReflect.Descriptor instead. func (*Capabilities) Descriptor() ([]byte, []int) { - return file_tunnel_proto_rawDescGZIP(), []int{3} + return file_tunnel_proto_tunnel_proto_rawDescGZIP(), []int{3} } func (x *Capabilities) GetDocker() bool { @@ -328,6 +353,13 @@ func (x *Capabilities) GetHostExec() bool { return false } +func (x *Capabilities) GetVnc() bool { + if x != nil { + return x.Vnc + } + return false +} + type SystemInfo struct { state protoimpl.MessageState `protogen:"open.v1"` Os string `protobuf:"bytes,1,opt,name=os,proto3" json:"os,omitempty"` @@ -340,7 +372,7 @@ type SystemInfo struct { func (x *SystemInfo) Reset() { *x = SystemInfo{} - mi := &file_tunnel_proto_msgTypes[4] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[4] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -352,7 +384,7 @@ func (x *SystemInfo) String() string { func (*SystemInfo) ProtoMessage() {} func (x *SystemInfo) ProtoReflect() protoreflect.Message { - mi := &file_tunnel_proto_msgTypes[4] + mi := &file_tunnel_proto_tunnel_proto_msgTypes[4] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -365,7 +397,7 @@ func (x *SystemInfo) ProtoReflect() protoreflect.Message { // Deprecated: Use SystemInfo.ProtoReflect.Descriptor instead. func (*SystemInfo) Descriptor() ([]byte, []int) { - return file_tunnel_proto_rawDescGZIP(), []int{4} + return file_tunnel_proto_tunnel_proto_rawDescGZIP(), []int{4} } func (x *SystemInfo) GetOs() string { @@ -396,12 +428,12 @@ func (x *SystemInfo) GetShell() string { return "" } -var File_tunnel_proto protoreflect.FileDescriptor +var File_tunnel_proto_tunnel_proto protoreflect.FileDescriptor -const file_tunnel_proto_rawDesc = "" + +const file_tunnel_proto_tunnel_proto_rawDesc = "" + "\n" + - "\ftunnel.proto\x12\n" + - "tai.tunnel\"\xf6\x02\n" + + "\x19tunnel/proto/tunnel.proto\x12\n" + + "tai.tunnel\"\xe3\x03\n" + "\rTunnelControl\x12\x12\n" + "\x04type\x18\x01 \x01(\tR\x04type\x12\x17\n" + "\anode_id\x18\x02 \x01(\tR\x06nodeId\x12\x1d\n" + @@ -416,7 +448,10 @@ const file_tunnel_proto_rawDesc = "" + "channel_id\x18\n" + " \x01(\tR\tchannelId\x12\x1f\n" + "\vtarget_port\x18\v \x01(\x05R\n" + - "targetPort\x12\x15\n" + + "targetPort\x12!\n" + + "\fchannel_type\x18\f \x01(\tR\vchannelType\x12!\n" + + "\fcontainer_id\x18\r \x01(\tR\vcontainerId\x12%\n" + + "\x0econtainer_port\x18\x0e \x01(\x05R\rcontainerPort\x12\x15\n" + "\x06tai_id\x18\x14 \x01(\tR\x05taiId\"!\n" + "\vForwardData\x12\x12\n" + "\x04data\x18\x01 \x01(\fR\x04data\"k\n" + @@ -425,11 +460,12 @@ const file_tunnel_proto_rawDesc = "" + "\x04http\x18\x02 \x01(\x05R\x04http\x12\x10\n" + "\x03vnc\x18\x03 \x01(\x05R\x03vnc\x12\x16\n" + "\x06docker\x18\x04 \x01(\x05R\x06docker\x12\x10\n" + - "\x03k8s\x18\x05 \x01(\x05R\x03k8s\"U\n" + + "\x03k8s\x18\x05 \x01(\x05R\x03k8s\"g\n" + "\fCapabilities\x12\x16\n" + "\x06docker\x18\x01 \x01(\bR\x06docker\x12\x10\n" + "\x03k8s\x18\x02 \x01(\bR\x03k8s\x12\x1b\n" + - "\thost_exec\x18\x03 \x01(\bR\bhostExec\"b\n" + + "\thost_exec\x18\x03 \x01(\bR\bhostExec\x12\x10\n" + + "\x03vnc\x18\x04 \x01(\bR\x03vnc\"b\n" + "\n" + "SystemInfo\x12\x0e\n" + "\x02os\x18\x01 \x01(\tR\x02os\x12\x12\n" + @@ -438,29 +474,29 @@ const file_tunnel_proto_rawDesc = "" + "\x05shell\x18\x04 \x01(\tR\x05shell2\x92\x01\n" + "\tTaiTunnel\x12D\n" + "\bRegister\x12\x19.tai.tunnel.TunnelControl\x1a\x19.tai.tunnel.TunnelControl(\x010\x01\x12?\n" + - "\aForward\x12\x17.tai.tunnel.ForwardData\x1a\x17.tai.tunnel.ForwardData(\x010\x01B(Z&github.com/yaoapp/yao/tai/tunnel/taipbb\x06proto3" + "\aForward\x12\x17.tai.tunnel.ForwardData\x1a\x17.tai.tunnel.ForwardData(\x010\x01B$Z\"github.com/yaoapp/tai/tunnel/taipbb\x06proto3" var ( - file_tunnel_proto_rawDescOnce sync.Once - file_tunnel_proto_rawDescData []byte + file_tunnel_proto_tunnel_proto_rawDescOnce sync.Once + file_tunnel_proto_tunnel_proto_rawDescData []byte ) -func file_tunnel_proto_rawDescGZIP() []byte { - file_tunnel_proto_rawDescOnce.Do(func() { - file_tunnel_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_tunnel_proto_rawDesc), len(file_tunnel_proto_rawDesc))) +func file_tunnel_proto_tunnel_proto_rawDescGZIP() []byte { + file_tunnel_proto_tunnel_proto_rawDescOnce.Do(func() { + file_tunnel_proto_tunnel_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_tunnel_proto_tunnel_proto_rawDesc), len(file_tunnel_proto_tunnel_proto_rawDesc))) }) - return file_tunnel_proto_rawDescData + return file_tunnel_proto_tunnel_proto_rawDescData } -var file_tunnel_proto_msgTypes = make([]protoimpl.MessageInfo, 5) -var file_tunnel_proto_goTypes = []any{ +var file_tunnel_proto_tunnel_proto_msgTypes = make([]protoimpl.MessageInfo, 5) +var file_tunnel_proto_tunnel_proto_goTypes = []any{ (*TunnelControl)(nil), // 0: tai.tunnel.TunnelControl (*ForwardData)(nil), // 1: tai.tunnel.ForwardData (*Ports)(nil), // 2: tai.tunnel.Ports (*Capabilities)(nil), // 3: tai.tunnel.Capabilities (*SystemInfo)(nil), // 4: tai.tunnel.SystemInfo } -var file_tunnel_proto_depIdxs = []int32{ +var file_tunnel_proto_tunnel_proto_depIdxs = []int32{ 2, // 0: tai.tunnel.TunnelControl.ports:type_name -> tai.tunnel.Ports 3, // 1: tai.tunnel.TunnelControl.caps:type_name -> tai.tunnel.Capabilities 4, // 2: tai.tunnel.TunnelControl.system:type_name -> tai.tunnel.SystemInfo @@ -475,26 +511,26 @@ var file_tunnel_proto_depIdxs = []int32{ 0, // [0:3] is the sub-list for field type_name } -func init() { file_tunnel_proto_init() } -func file_tunnel_proto_init() { - if File_tunnel_proto != nil { +func init() { file_tunnel_proto_tunnel_proto_init() } +func file_tunnel_proto_tunnel_proto_init() { + if File_tunnel_proto_tunnel_proto != nil { return } type x struct{} out := protoimpl.TypeBuilder{ File: protoimpl.DescBuilder{ GoPackagePath: reflect.TypeOf(x{}).PkgPath(), - RawDescriptor: unsafe.Slice(unsafe.StringData(file_tunnel_proto_rawDesc), len(file_tunnel_proto_rawDesc)), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_tunnel_proto_tunnel_proto_rawDesc), len(file_tunnel_proto_tunnel_proto_rawDesc)), NumEnums: 0, NumMessages: 5, NumExtensions: 0, NumServices: 1, }, - GoTypes: file_tunnel_proto_goTypes, - DependencyIndexes: file_tunnel_proto_depIdxs, - MessageInfos: file_tunnel_proto_msgTypes, + GoTypes: file_tunnel_proto_tunnel_proto_goTypes, + DependencyIndexes: file_tunnel_proto_tunnel_proto_depIdxs, + MessageInfos: file_tunnel_proto_tunnel_proto_msgTypes, }.Build() - File_tunnel_proto = out.File - file_tunnel_proto_goTypes = nil - file_tunnel_proto_depIdxs = nil + File_tunnel_proto_tunnel_proto = out.File + file_tunnel_proto_tunnel_proto_goTypes = nil + file_tunnel_proto_tunnel_proto_depIdxs = nil } diff --git a/tai/tunnel/taipb/tunnel_grpc.pb.go b/tai/tunnel/taipb/tunnel_grpc.pb.go index 973075e6..eb0241b2 100644 --- a/tai/tunnel/taipb/tunnel_grpc.pb.go +++ b/tai/tunnel/taipb/tunnel_grpc.pb.go @@ -2,7 +2,7 @@ // versions: // - protoc-gen-go-grpc v1.6.1 // - protoc v4.25.0 -// source: tunnel.proto +// source: tunnel/proto/tunnel.proto package taipb @@ -147,5 +147,5 @@ var TaiTunnel_ServiceDesc = grpc.ServiceDesc{ ClientStreams: true, }, }, - Metadata: "tunnel.proto", + Metadata: "tunnel/proto/tunnel.proto", } diff --git a/tai/types/types.go b/tai/types/types.go index a6c80be7..d47364c3 100644 --- a/tai/types/types.go +++ b/tai/types/types.go @@ -24,6 +24,7 @@ type Capabilities struct { Docker bool `json:"docker"` K8s bool `json:"k8s"` HostExec bool `json:"host_exec"` + VNC bool `json:"vnc"` } // SystemInfo describes the host machine running Tai.