diff --git a/.github/workflows/real-e2e.yml b/.github/workflows/real-e2e.yml index 3e1c85d3b..9d098b186 100644 --- a/.github/workflows/real-e2e.yml +++ b/.github/workflows/real-e2e.yml @@ -35,6 +35,8 @@ jobs: pip install uv - name: Run tests + env: + OPENSANDBOX_SANDBOX_DEFAULT_IMAGE: opensandbox/code-interpreter:latest run: | set -e diff --git a/components/execd/pkg/runtime/bash_session.go b/components/execd/pkg/runtime/bash_session.go new file mode 100644 index 000000000..bde8b7119 --- /dev/null +++ b/components/execd/pkg/runtime/bash_session.go @@ -0,0 +1,282 @@ +// Copyright 2026 Alibaba Group Holding Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build !windows +// +build !windows + +package runtime + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "os" + "os/exec" + "strconv" + "strings" + "time" + + "github.com/google/uuid" + + "github.com/alibaba/opensandbox/execd/pkg/log" +) + +func (c *Controller) createBashSession(_ *CreateContextRequest) (string, error) { + session := newBashSession(nil) + if err := session.start(); err != nil { + return "", fmt.Errorf("failed to start bash session: %w", err) + } + + c.bashSessionClientMap.Store(session.config.Session, session) + log.Info("created bash session %s", session.config.Session) + return session.config.Session, nil +} + +func (c *Controller) runBashSession(_ context.Context, request *ExecuteCodeRequest) error { + if request.Context == "" { + if c.getDefaultLanguageSession(request.Language) == "" { + if err := c.createDefaultBashSession(); err != nil { + return err + } + } + } + + targetSessionID := request.Context + if targetSessionID == "" { + targetSessionID = c.getDefaultLanguageSession(request.Language) + } + + session := c.getBashSession(targetSessionID) + if session == nil { + return ErrContextNotFound + } + + return session.run(request.Code, request.Timeout, &request.Hooks) +} + +func (c *Controller) createDefaultBashSession() error { + session, err := c.createBashSession(&CreateContextRequest{}) + if err != nil { + return err + } + + c.setDefaultLanguageSession(Bash, session) + return nil +} + +func (c *Controller) getBashSession(sessionId string) *bashSession { + if v, ok := c.bashSessionClientMap.Load(sessionId); ok { + if s, ok := v.(*bashSession); ok { + return s + } + } + return nil +} + +func (c *Controller) closeBashSession(sessionId string) error { + session := c.getBashSession(sessionId) + if session == nil { + return ErrContextNotFound + } + + err := session.close() + if err != nil { + return err + } + + c.bashSessionClientMap.Delete(sessionId) + return nil +} + +// nolint:unused +func (c *Controller) listBashSessions() []string { + sessions := make([]string, 0) + c.bashSessionClientMap.Range(func(key, _ any) bool { + sessionID, _ := key.(string) + sessions = append(sessions, sessionID) + return true + }) + + return sessions +} + +// Session implementation (pipe-based, no PTY) +func newBashSession(config *bashSessionConfig) *bashSession { + if config == nil { + config = &bashSessionConfig{ + Session: uuidString(), + StartupTimeout: 5 * time.Second, + } + } + return &bashSession{ + config: config, + stdoutLines: make(chan string, 256), + stdoutErr: make(chan error, 1), + } +} + +func (s *bashSession) start() error { + s.mu.Lock() + defer s.mu.Unlock() + + if s.started { + return errors.New("session already started") + } + + cmd := exec.Command("bash", "--noprofile", "--norc", "-s") + cmd.Env = os.Environ() + + stdin, err := cmd.StdinPipe() + if err != nil { + return fmt.Errorf("stdin pipe: %w", err) + } + stdout, err := cmd.StdoutPipe() + if err != nil { + return fmt.Errorf("stdout pipe: %w", err) + } + stderr, err := cmd.StderrPipe() + if err != nil { + return fmt.Errorf("stderr pipe: %w", err) + } + + if err := cmd.Start(); err != nil { + return fmt.Errorf("start bash: %w", err) + } + + s.cmd = cmd + s.stdin = stdin + s.stdout = stdout + s.stderr = stderr + s.started = true + + // drain stdout/stderr into channel + go s.readStdout(stdout) + go s.discardStderr(stderr) + return nil +} + +func (s *bashSession) readStdout(r io.Reader) { + reader := bufio.NewReader(r) + for { + line, err := reader.ReadString('\n') + if len(line) > 0 { + s.stdoutLines <- strings.TrimRight(line, "\r\n") + } + if err != nil { + if !errors.Is(err, io.EOF) { + s.stdoutErr <- err + } + close(s.stdoutLines) + return + } + } +} + +func (s *bashSession) discardStderr(r io.Reader) { + _, _ = io.Copy(io.Discard, r) +} + +func (s *bashSession) run(command string, timeout time.Duration, hooks *ExecuteResultHook) error { + s.mu.Lock() + defer s.mu.Unlock() + + if !s.started { + return errors.New("session not started") + } + + startAt := time.Now() + + if hooks != nil && hooks.OnExecuteInit != nil { + hooks.OnExecuteInit(s.config.Session) + } + + waitSeconds := timeout + if waitSeconds <= 0 { + waitSeconds = 30 * time.Second + } + + cleanCmd := strings.ReplaceAll(command, "\n", " ; ") + + // send command + marker + cmdText := fmt.Sprintf("%s\nprintf \"%s$?%s\\n\"\n", cleanCmd, exitCodePrefix, exitCodeSuffix) + if _, err := fmt.Fprint(s.stdin, cmdText); err != nil { + return fmt.Errorf("write command: %w", err) + } + + // collect output until marker + timer := time.NewTimer(waitSeconds) + defer timer.Stop() + + for { + select { + case <-timer.C: + return fmt.Errorf("timeout after %s while running command %q", waitSeconds, command) + case err := <-s.stdoutErr: + if err != nil { + return err + } + case line, ok := <-s.stdoutLines: + if !ok { + return errors.New("stdout closed unexpectedly") + } + if _, ok := parseExitCodeLine(line); ok { + if hooks != nil && hooks.OnExecuteComplete != nil { + hooks.OnExecuteComplete(time.Since(startAt)) + } + return nil + } + if hooks != nil && hooks.OnExecuteStdout != nil { + hooks.OnExecuteStdout(line) + } + } + } +} + +func parseExitCodeLine(line string) (int, bool) { + p := strings.Index(line, exitCodePrefix) + q := strings.Index(line, exitCodeSuffix) + if p < 0 || q <= p { + return 0, false + } + text := strings.TrimSpace(line[p+len(exitCodePrefix) : q]) + code, err := strconv.Atoi(text) + if err != nil { + return 0, false + } + return code, true +} + +func (s *bashSession) close() error { + s.mu.Lock() + defer s.mu.Unlock() + + if !s.started { + return nil + } + s.started = false + + if s.stdin != nil { + _ = s.stdin.Close() + } + if s.cmd != nil && s.cmd.Process != nil { + _ = s.cmd.Process.Kill() + } + return nil +} + +func uuidString() string { + return uuid.New().String() +} diff --git a/components/execd/pkg/runtime/bash_session_test.go b/components/execd/pkg/runtime/bash_session_test.go new file mode 100644 index 000000000..ac66ca0a3 --- /dev/null +++ b/components/execd/pkg/runtime/bash_session_test.go @@ -0,0 +1,183 @@ +// Copyright 2026 Alibaba Group Holding Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build !windows +// +build !windows + +package runtime + +import ( + "strings" + "testing" + "time" +) + +func TestBashSessionEnvAndExitCode(t *testing.T) { + session := newBashSession(nil) + t.Cleanup(func() { _ = session.close() }) + + if err := session.start(); err != nil { + t.Fatalf("Start() error = %v", err) + } + + var ( + initCalls int + completeCalls int + stdoutLines []string + ) + + hooks := ExecuteResultHook{ + OnExecuteInit: func(ctx string) { + if ctx != session.config.Session { + t.Fatalf("unexpected session in OnExecuteInit: %s", ctx) + } + initCalls++ + }, + OnExecuteStdout: func(text string) { + t.Log(text) + stdoutLines = append(stdoutLines, text) + }, + OnExecuteComplete: func(_ time.Duration) { + completeCalls++ + }, + } + + // 1) export an env var + if err := session.run("export FOO=hello", 3*time.Second, &hooks); err != nil { + t.Fatalf("runCommand(export) error = %v", err) + } + exportStdoutCount := len(stdoutLines) + + // 2) verify env is persisted + if err := session.run("echo $FOO", 3*time.Second, &hooks); err != nil { + t.Fatalf("runCommand(echo) error = %v", err) + } + echoLines := stdoutLines[exportStdoutCount:] + foundHello := false + for _, line := range echoLines { + if strings.TrimSpace(line) == "hello" { + foundHello = true + break + } + } + if !foundHello { + t.Fatalf("expected echo $FOO to output 'hello', got %v", echoLines) + } + + // 3) ensure exit code of previous command is reflected in shell state + prevCount := len(stdoutLines) + if err := session.run("false; echo EXIT:$?", 3*time.Second, &hooks); err != nil { + t.Fatalf("runCommand(exitcode) error = %v", err) + } + exitLines := stdoutLines[prevCount:] + foundExit := false + for _, line := range exitLines { + if strings.Contains(line, "EXIT:1") { + foundExit = true + break + } + } + if !foundExit { + t.Fatalf("expected exit code output 'EXIT:1', got %v", exitLines) + } + + if initCalls != 3 { + t.Fatalf("OnExecuteInit expected 3 calls, got %d", initCalls) + } + if completeCalls != 3 { + t.Fatalf("OnExecuteComplete expected 3 calls, got %d", completeCalls) + } +} + +func TestBashSessionEnvLargeOutputChained(t *testing.T) { + session := newBashSession(nil) + t.Cleanup(func() { _ = session.close() }) + + if err := session.start(); err != nil { + t.Fatalf("Start() error = %v", err) + } + + var ( + initCalls int + completeCalls int + stdoutLines []string + ) + + hooks := ExecuteResultHook{ + OnExecuteInit: func(ctx string) { + if ctx != session.config.Session { + t.Fatalf("unexpected session in OnExecuteInit: %s", ctx) + } + initCalls++ + }, + OnExecuteStdout: func(text string) { + t.Log(text) + stdoutLines = append(stdoutLines, text) + }, + OnExecuteComplete: func(_ time.Duration) { + completeCalls++ + }, + } + + runAndCollect := func(cmd string) []string { + start := len(stdoutLines) + if err := session.run(cmd, 10*time.Second, &hooks); err != nil { + t.Fatalf("runCommand(%q) error = %v", cmd, err) + } + return append([]string(nil), stdoutLines[start:]...) + } + + lines1 := runAndCollect("export FOO=hello1; for i in $(seq 1 60); do echo A${i}:$FOO; done") + if len(lines1) < 60 { + t.Fatalf("expected >=60 lines for cmd1, got %d", len(lines1)) + } + if !containsLine(lines1, "A1:hello1") || !containsLine(lines1, "A60:hello1") { + t.Fatalf("env not reflected in cmd1 output, got %v", lines1[:3]) + } + + lines2 := runAndCollect("export FOO=${FOO}_next; export BAR=bar1; for i in $(seq 1 60); do echo B${i}:$FOO:$BAR; done") + if len(lines2) < 60 { + t.Fatalf("expected >=60 lines for cmd2, got %d", len(lines2)) + } + if !containsLine(lines2, "B1:hello1_next:bar1") || !containsLine(lines2, "B60:hello1_next:bar1") { + t.Fatalf("env not propagated to cmd2 output, sample %v", lines2[:3]) + } + + lines3 := runAndCollect("export BAR=${BAR}_last; for i in $(seq 1 60); do echo C${i}:$FOO:$BAR; done; echo FINAL_FOO=$FOO; echo FINAL_BAR=$BAR") + if len(lines3) < 62 { // 60 lines + 2 finals + t.Fatalf("expected >=62 lines for cmd3, got %d", len(lines3)) + } + if !containsLine(lines3, "C1:hello1_next:bar1_last") || !containsLine(lines3, "C60:hello1_next:bar1_last") { + t.Fatalf("env not propagated to cmd3 output, sample %v", lines3[:3]) + } + if !containsLine(lines3, "FINAL_FOO=hello1_next") || !containsLine(lines3, "FINAL_BAR=bar1_last") { + t.Fatalf("final env lines missing, got %v", lines3[len(lines3)-5:]) + } + + if initCalls != 3 { + t.Fatalf("OnExecuteInit expected 3 calls, got %d", initCalls) + } + if completeCalls != 3 { + t.Fatalf("OnExecuteComplete expected 3 calls, got %d", completeCalls) + } +} + +func containsLine(lines []string, target string) bool { + for _, l := range lines { + if strings.TrimSpace(l) == target { + return true + } + } + return false +} diff --git a/components/execd/pkg/runtime/bash_session_windows.go b/components/execd/pkg/runtime/bash_session_windows.go new file mode 100644 index 000000000..8b65db812 --- /dev/null +++ b/components/execd/pkg/runtime/bash_session_windows.go @@ -0,0 +1,67 @@ +// Copyright 2026 Alibaba Group Holding Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build windows +// +build windows + +package runtime + +import ( + "context" + "errors" + "time" +) + +var errBashSessionNotSupported = errors.New("bash session is not supported on windows") + +func (c *Controller) createBashSession(_ *CreateContextRequest) (string, error) { + return "", errBashSessionNotSupported +} + +func (c *Controller) runBashSession(_ context.Context, _ *ExecuteCodeRequest) error { //nolint:revive + return errBashSessionNotSupported +} + +func (c *Controller) createDefaultBashSession() error { //nolint:revive + return errBashSessionNotSupported +} + +func (c *Controller) getBashSession(_ string) (*bashSession, error) { //nolint:revive + return nil, errBashSessionNotSupported +} + +func (c *Controller) closeBashSession(_ string) error { //nolint:revive + return errBashSessionNotSupported +} + +func (c *Controller) listBashSessions() []string { //nolint:revive + return nil +} + +// Stub methods on bashSession to satisfy interfaces on non-Linux platforms. +func newBashSession(config *bashSessionConfig) *bashSession { + return &bashSession{config: config} +} + +func (s *bashSession) start() (string, error) { + return "", errBashSessionNotSupported +} + +func (s *bashSession) run(_ string, _ time.Duration, _ *ExecuteResultHook) error { + return errBashSessionNotSupported +} + +func (s *bashSession) close() error { + return nil +} diff --git a/components/execd/pkg/runtime/command_common.go b/components/execd/pkg/runtime/command_common.go index 633efa35b..4f49ebbcb 100644 --- a/components/execd/pkg/runtime/command_common.go +++ b/components/execd/pkg/runtime/command_common.go @@ -45,18 +45,17 @@ func (c *Controller) tailStdPipe(file string, onExecute func(text string), done // getCommandKernel retrieves a command execution context. func (c *Controller) getCommandKernel(sessionID string) *commandKernel { - c.mu.RLock() - defer c.mu.RUnlock() - - return c.commandClientMap[sessionID] + if v, ok := c.commandClientMap.Load(sessionID); ok { + if kernel, ok := v.(*commandKernel); ok { + return kernel + } + } + return nil } // storeCommandKernel registers a command execution context. func (c *Controller) storeCommandKernel(sessionID string, kernel *commandKernel) { - c.mu.Lock() - defer c.mu.Unlock() - - c.commandClientMap[sessionID] = kernel + c.commandClientMap.Store(sessionID, kernel) } // stdLogDescriptor creates temporary files for capturing command output. diff --git a/components/execd/pkg/runtime/command_status.go b/components/execd/pkg/runtime/command_status.go index 97f112b1c..6dbc6d4f2 100644 --- a/components/execd/pkg/runtime/command_status.go +++ b/components/execd/pkg/runtime/command_status.go @@ -40,11 +40,11 @@ type CommandOutput struct { } func (c *Controller) commandSnapshot(session string) *commandKernel { - c.mu.RLock() - defer c.mu.RUnlock() - - kernel, ok := c.commandClientMap[session] - if !ok || kernel == nil { + var kernel *commandKernel + if v, ok := c.commandClientMap.Load(session); ok { + kernel, _ = v.(*commandKernel) + } + if kernel == nil { return nil } @@ -116,8 +116,11 @@ func (c *Controller) markCommandFinished(session string, exitCode int, errMsg st c.mu.Lock() defer c.mu.Unlock() - kernel, ok := c.commandClientMap[session] - if !ok || kernel == nil { + var kernel *commandKernel + if v, ok := c.commandClientMap.Load(session); ok { + kernel, _ = v.(*commandKernel) + } + if kernel == nil { return } diff --git a/components/execd/pkg/runtime/context.go b/components/execd/pkg/runtime/context.go index a11355072..6e7ea8701 100644 --- a/components/execd/pkg/runtime/context.go +++ b/components/execd/pkg/runtime/context.go @@ -32,6 +32,11 @@ import ( // CreateContext provisions a kernel-backed session and returns its ID. func (c *Controller) CreateContext(req *CreateContextRequest) (string, error) { + if req.Language == Bash { + return c.createBashSession(req) + } + + // Create a new Jupyter session. var ( client *jupyter.Client session *jupytersession.Session @@ -42,7 +47,7 @@ func (c *Controller) CreateContext(req *CreateContextRequest) (string, error) { log.Error("failed to create session, retrying: %v", err) return err != nil }, func() error { - client, session, err = c.createContext(*req) + client, session, err = c.createJupyterContext(*req) return err }) if err != nil { @@ -116,15 +121,8 @@ func (c *Controller) deleteSessionAndCleanup(session string) error { return err } - c.mu.Lock() - defer c.mu.Unlock() - - delete(c.jupyterClientMap, session) - for lang, id := range c.defaultLanguageJupyterSessions { - if id == session { - delete(c.defaultLanguageJupyterSessions, lang) - } - } + c.jupyterClientMap.Delete(session) + c.deleteDefaultSessionByID(session) return nil } @@ -143,8 +141,12 @@ func (c *Controller) newIpynbPath(sessionID, cwd string) (string, error) { return filepath.Join(cwd, fmt.Sprintf("%s.ipynb", sessionID)), nil } -// createDefaultLanguageContext prewarms a session for stateless execution. -func (c *Controller) createDefaultLanguageContext(language Language) error { +// createDefaultLanguageJupyterContext prewarms a session for stateless execution. +func (c *Controller) createDefaultLanguageJupyterContext(language Language) error { + if c.getDefaultLanguageSession(language) != "" { + return nil + } + var ( client *jupyter.Client session *jupytersession.Session @@ -154,7 +156,7 @@ func (c *Controller) createDefaultLanguageContext(language Language) error { log.Error("failed to create context, retrying: %v", err) return err != nil }, func() error { - client, session, err = c.createContext(CreateContextRequest{ + client, session, err = c.createJupyterContext(CreateContextRequest{ Language: language, Cwd: "", }) @@ -164,20 +166,17 @@ func (c *Controller) createDefaultLanguageContext(language Language) error { return err } - c.mu.Lock() - defer c.mu.Unlock() - - c.defaultLanguageJupyterSessions[language] = session.ID - c.jupyterClientMap[session.ID] = &jupyterKernel{ + c.setDefaultLanguageSession(language, session.ID) + c.jupyterClientMap.Store(session.ID, &jupyterKernel{ kernelID: session.Kernel.ID, client: client, language: language, - } + }) return nil } -// createContext performs the actual context creation workflow. -func (c *Controller) createContext(request CreateContextRequest) (*jupyter.Client, *jupytersession.Session, error) { +// createJupyterContext performs the actual context creation workflow. +func (c *Controller) createJupyterContext(request CreateContextRequest) (*jupyter.Client, *jupytersession.Session, error) { client := c.jupyterClient() kernel, err := c.searchKernel(client, request.Language) @@ -217,10 +216,7 @@ func (c *Controller) createContext(request CreateContextRequest) (*jupyter.Clien // storeJupyterKernel caches a session -> kernel mapping. func (c *Controller) storeJupyterKernel(sessionID string, kernel *jupyterKernel) { - c.mu.Lock() - defer c.mu.Unlock() - - c.jupyterClientMap[sessionID] = kernel + c.jupyterClientMap.Store(sessionID, kernel) } func (c *Controller) jupyterClient() *jupyter.Client { @@ -236,49 +232,63 @@ func (c *Controller) jupyterClient() *jupyter.Client { jupyter.WithHTTPClient(httpClient)) } -func (c *Controller) listAllContexts() ([]CodeContext, error) { - c.mu.RLock() - defer c.mu.RUnlock() +func (c *Controller) getDefaultLanguageSession(language Language) string { + if v, ok := c.defaultLanguageSessions.Load(language); ok { + if session, ok := v.(string); ok { + return session + } + } + return "" +} + +func (c *Controller) setDefaultLanguageSession(language Language, sessionID string) { + c.defaultLanguageSessions.Store(language, sessionID) +} +func (c *Controller) deleteDefaultSessionByID(sessionID string) { + c.defaultLanguageSessions.Range(func(key, value any) bool { + if s, ok := value.(string); ok && s == sessionID { + c.defaultLanguageSessions.Delete(key) + } + return true + }) +} + +func (c *Controller) listAllContexts() ([]CodeContext, error) { contexts := make([]CodeContext, 0) - for session, kernel := range c.jupyterClientMap { - if kernel != nil { - contexts = append(contexts, CodeContext{ - ID: session, - Language: kernel.language, - }) + c.jupyterClientMap.Range(func(key, value any) bool { + session, _ := key.(string) + if kernel, ok := value.(*jupyterKernel); ok && kernel != nil { + contexts = append(contexts, CodeContext{ID: session, Language: kernel.language}) } - } + return true + }) - for language, defaultContext := range c.defaultLanguageJupyterSessions { - contexts = append(contexts, CodeContext{ - ID: defaultContext, - Language: language, - }) - } + c.defaultLanguageSessions.Range(func(key, value any) bool { + lang, _ := key.(Language) + session, _ := value.(string) + if session == "" { + return true + } + contexts = append(contexts, CodeContext{ID: session, Language: lang}) + return true + }) return contexts, nil } func (c *Controller) listLanguageContexts(language Language) ([]CodeContext, error) { - c.mu.RLock() - defer c.mu.RUnlock() - contexts := make([]CodeContext, 0) - for session, kernel := range c.jupyterClientMap { - if kernel != nil && kernel.language == language { - contexts = append(contexts, CodeContext{ - ID: session, - Language: language, - }) + c.jupyterClientMap.Range(func(key, value any) bool { + session, _ := key.(string) + if kernel, ok := value.(*jupyterKernel); ok && kernel != nil && kernel.language == language { + contexts = append(contexts, CodeContext{ID: session, Language: language}) } - } + return true + }) - if defaultContext := c.defaultLanguageJupyterSessions[language]; defaultContext != "" { - contexts = append(contexts, CodeContext{ - ID: defaultContext, - Language: language, - }) + if defaultContext := c.getDefaultLanguageSession(language); defaultContext != "" { + contexts = append(contexts, CodeContext{ID: defaultContext, Language: language}) } return contexts, nil diff --git a/components/execd/pkg/runtime/context_test.go b/components/execd/pkg/runtime/context_test.go index 6a27ad18b..43efe81c0 100644 --- a/components/execd/pkg/runtime/context_test.go +++ b/components/execd/pkg/runtime/context_test.go @@ -26,8 +26,9 @@ import ( func TestListContextsAndNewIpynbPath(t *testing.T) { c := NewController("http://example", "token") - c.jupyterClientMap["session-python"] = &jupyterKernel{language: Python} - c.defaultLanguageJupyterSessions[Go] = "session-go-default" + + c.jupyterClientMap.Store("session-python", &jupyterKernel{language: Python}) + c.setDefaultLanguageSession(Go, "session-go-default") pyContexts, err := c.listLanguageContexts(Python) if err != nil { @@ -128,8 +129,8 @@ func TestDeleteContext_RemovesCacheOnSuccess(t *testing.T) { defer server.Close() c := NewController(server.URL, "token") - c.jupyterClientMap[sessionID] = &jupyterKernel{language: Python} - c.defaultLanguageJupyterSessions[Python] = sessionID + c.jupyterClientMap.Store(sessionID, &jupyterKernel{language: Python}) + c.setDefaultLanguageSession(Python, sessionID) if err := c.DeleteContext(sessionID); err != nil { t.Fatalf("DeleteContext returned error: %v", err) @@ -138,7 +139,7 @@ func TestDeleteContext_RemovesCacheOnSuccess(t *testing.T) { if kernel := c.getJupyterKernel(sessionID); kernel != nil { t.Fatalf("expected cache to be cleared, found: %+v", kernel) } - if _, ok := c.defaultLanguageJupyterSessions[Python]; ok { + if c.getDefaultLanguageSession(Python) != "" { t.Fatalf("expected default session entry to be removed") } } @@ -166,21 +167,21 @@ func TestDeleteLanguageContext_RemovesCacheOnSuccess(t *testing.T) { defer server.Close() c := NewController(server.URL, "token") - c.jupyterClientMap[session1] = &jupyterKernel{language: lang} - c.jupyterClientMap[session2] = &jupyterKernel{language: lang} - c.defaultLanguageJupyterSessions[lang] = session2 + c.jupyterClientMap.Store(session1, &jupyterKernel{language: lang}) + c.jupyterClientMap.Store(session2, &jupyterKernel{language: lang}) + c.setDefaultLanguageSession(lang, session2) if err := c.DeleteLanguageContext(lang); err != nil { t.Fatalf("DeleteLanguageContext returned error: %v", err) } - if _, ok := c.jupyterClientMap[session1]; ok { + if v, ok := c.jupyterClientMap.Load(session1); ok && v != nil { t.Fatalf("expected session1 removed from cache") } - if _, ok := c.jupyterClientMap[session2]; ok { + if v, ok := c.jupyterClientMap.Load(session2); ok && v != nil { t.Fatalf("expected session2 removed from cache") } - if _, ok := c.defaultLanguageJupyterSessions[lang]; ok { + if c.getDefaultLanguageSession(lang) != "" { t.Fatalf("expected default entry removed") } if deleteCalls[session1] != 1 || deleteCalls[session2] != 1 { diff --git a/components/execd/pkg/runtime/ctrl.go b/components/execd/pkg/runtime/ctrl.go index 20bbecc62..2bb1967be 100644 --- a/components/execd/pkg/runtime/ctrl.go +++ b/components/execd/pkg/runtime/ctrl.go @@ -35,14 +35,15 @@ var kernelWaitingBackoff = wait.Backoff{ // Controller manages code execution across runtimes. type Controller struct { - baseURL string - token string - mu sync.RWMutex - jupyterClientMap map[string]*jupyterKernel - defaultLanguageJupyterSessions map[Language]string - commandClientMap map[string]*commandKernel - db *sql.DB - dbOnce sync.Once + baseURL string + token string + mu sync.RWMutex + jupyterClientMap sync.Map // sessionID -> *jupyterKernel + defaultLanguageSessions sync.Map // Language -> sessionID + commandClientMap sync.Map // sessionID -> *commandKernel + bashSessionClientMap sync.Map // sessionID -> *bashSession + db *sql.DB + dbOnce sync.Once } type jupyterKernel struct { @@ -71,9 +72,10 @@ func NewController(baseURL, token string) *Controller { baseURL: baseURL, token: token, - jupyterClientMap: make(map[string]*jupyterKernel), - defaultLanguageJupyterSessions: make(map[Language]string), - commandClientMap: make(map[string]*commandKernel), + jupyterClientMap: sync.Map{}, + defaultLanguageSessions: sync.Map{}, + commandClientMap: sync.Map{}, + bashSessionClientMap: sync.Map{}, } } @@ -93,10 +95,12 @@ func (c *Controller) Execute(request *ExecuteCodeRequest) error { return c.runCommand(ctx, request) case BackgroundCommand: return c.runBackgroundCommand(ctx, request) - case Bash, Python, Java, JavaScript, TypeScript, Go: + case Python, Java, JavaScript, TypeScript, Go: return c.runJupyter(ctx, request) case SQL: return c.runSQL(ctx, request) + case Bash: + return c.runBashSession(ctx, request) default: return fmt.Errorf("unknown language: %s", request.Language) } diff --git a/components/execd/pkg/runtime/interrupt.go b/components/execd/pkg/runtime/interrupt.go index 1a9515fa1..67902a3d6 100644 --- a/components/execd/pkg/runtime/interrupt.go +++ b/components/execd/pkg/runtime/interrupt.go @@ -38,6 +38,8 @@ func (c *Controller) Interrupt(sessionID string) error { case c.getCommandKernel(sessionID) != nil: kernel := c.getCommandKernel(sessionID) return c.killPid(kernel.pid) + case c.getBashSession(sessionID) != nil: + return c.closeBashSession(sessionID) default: return errors.New("no such session") } diff --git a/components/execd/pkg/runtime/jupyter.go b/components/execd/pkg/runtime/jupyter.go index cdc0a6cc5..9ea33b13b 100644 --- a/components/execd/pkg/runtime/jupyter.go +++ b/components/execd/pkg/runtime/jupyter.go @@ -29,9 +29,8 @@ func (c *Controller) runJupyter(ctx context.Context, request *ExecuteCodeRequest return errors.New("language runtime server not configured, please check your image runtime") } if request.Context == "" { - if _, exists := c.defaultLanguageJupyterSessions[request.Language]; !exists { - err := c.createDefaultLanguageContext(request.Language) - if err != nil { + if c.getDefaultLanguageSession(request.Language) == "" { + if err := c.createDefaultLanguageJupyterContext(request.Language); err != nil { return err } } @@ -39,7 +38,7 @@ func (c *Controller) runJupyter(ctx context.Context, request *ExecuteCodeRequest var targetSessionID string if request.Context == "" { - targetSessionID = c.defaultLanguageJupyterSessions[request.Language] + targetSessionID = c.getDefaultLanguageSession(request.Language) } else { targetSessionID = request.Context } @@ -135,10 +134,12 @@ func (c *Controller) setWorkingDir(_ *jupyterKernel, _ *CreateContextRequest) er // getJupyterKernel retrieves a kernel connection from the session map. func (c *Controller) getJupyterKernel(sessionID string) *jupyterKernel { - c.mu.RLock() - defer c.mu.RUnlock() - - return c.jupyterClientMap[sessionID] + if v, ok := c.jupyterClientMap.Load(sessionID); ok { + if kernel, ok := v.(*jupyterKernel); ok { + return kernel + } + } + return nil } // searchKernel finds a kernel spec name for the given language. diff --git a/components/execd/pkg/runtime/types.go b/components/execd/pkg/runtime/types.go index cb82a11bc..5cd5addae 100644 --- a/components/execd/pkg/runtime/types.go +++ b/components/execd/pkg/runtime/types.go @@ -16,6 +16,9 @@ package runtime import ( "fmt" + "io" + "os/exec" + "sync" "time" "github.com/alibaba/opensandbox/execd/pkg/jupyter/execute" @@ -80,3 +83,33 @@ type CodeContext struct { ID string `json:"id,omitempty"` Language Language `json:"language"` } + +// bashSessionConfig holds bash session configuration. +type bashSessionConfig struct { + // StartupSource is a list of scripts sourced on startup. + StartupSource []string + // Session is the session identifier. + Session string + // StartupTimeout is the startup timeout. + StartupTimeout time.Duration +} + +const ( + // exitCodePrefix marks the beginning of exit code output. + exitCodePrefix = "EXITCODESTART" + // exitCodeSuffix marks the end of exit code output. + exitCodeSuffix = "EXITCODEEND" +) + +// bashSession represents a bash session. +type bashSession struct { + config *bashSessionConfig + cmd *exec.Cmd + stdin io.WriteCloser + stdout io.ReadCloser + stderr io.ReadCloser + stdoutLines chan string + stdoutErr chan error + mu sync.Mutex + started bool +} diff --git a/tests/javascript/tests/test_code_interpreter_e2e.test.ts b/tests/javascript/tests/test_code_interpreter_e2e.test.ts index 686fc4076..6ae00a41d 100644 --- a/tests/javascript/tests/test_code_interpreter_e2e.test.ts +++ b/tests/javascript/tests/test_code_interpreter_e2e.test.ts @@ -267,3 +267,58 @@ test("07 interrupt code execution + fake id", async () => { await expect(ci0.codes.interrupt(`fake-${Date.now()}`)).rejects.toBeTruthy(); }); + +test("08 bash env propagation across sequential executions", async () => { + if (!ci) throw new Error("not initialized"); + + const stdout: string[] = []; + const stderr: string[] = []; + const errors: string[] = []; + + const handlers: ExecutionHandlers = { + onStdout: (m) => { + if (m.text) stdout.push(m.text.trim()); + }, + onStderr: (m) => { + if (m.text) stderr.push(m.text.trim()); + }, + onError: (e) => { + errors.push(e.name); + }, + }; + + const code1 = "export FOO=hello\nexport BAR=world\n"; + const code2 = 'printf "step1:$FOO:$BAR\\n"\n'; + const code3 = + "export FOO=${FOO}_next\n" + + 'printf "step2:$FOO:$BAR\\n"\n' + + "export BAR=${BAR}_next\n" + + 'printf "step3:$FOO:$BAR\\n"\n'; + + const r1 = await ci.codes.run(code1, { + language: SupportedLanguages.BASH, + handlers, + }); + expect(r1.id).toBeTruthy(); + expect(r1.error).toBeUndefined(); + + const r2 = await ci.codes.run(code2, { + language: SupportedLanguages.BASH, + handlers, + }); + expect(r2.id).toBeTruthy(); + expect(r2.error).toBeUndefined(); + + const r3 = await ci.codes.run(code3, { + language: SupportedLanguages.BASH, + handlers, + }); + expect(r3.id).toBeTruthy(); + expect(r3.error).toBeUndefined(); + + expect(stdout).toContain("step1:hello:world"); + expect(stdout).toContain("step2:hello_next:world"); + expect(stdout).toContain("step3:hello_next:world_next"); + expect(errors).toHaveLength(0); + expect(stderr.filter((s) => s.length > 0)).toHaveLength(0); +}); diff --git a/tests/python/tests/test_code_interpreter_e2e_sync.py b/tests/python/tests/test_code_interpreter_e2e_sync.py index 95d8f9b7e..9bae9f2ea 100644 --- a/tests/python/tests/test_code_interpreter_e2e_sync.py +++ b/tests/python/tests/test_code_interpreter_e2e_sync.py @@ -893,3 +893,95 @@ def test_09_context_management_endpoints(self): assert len(final_contexts) == 0 logger.info("✓ delete_contexts removed all bash contexts") + @pytest.mark.timeout(300) + @pytest.mark.order(10) + def test_10_bash_env_propagation(self): + """Ensure bash commands share env/vars across sequential executions.""" + TestCodeInterpreterE2ESync._ensure_code_interpreter_created() + code_interpreter = TestCodeInterpreterE2ESync.code_interpreter + assert code_interpreter is not None + + stdout_messages: list[OutputMessage] = [] + stderr_messages: list[OutputMessage] = [] + errors: list[ExecutionError] = [] + completed_events: list[ExecutionComplete] = [] + init_events: list[ExecutionInit] = [] + + def on_stdout(msg: OutputMessage): + stdout_messages.append(msg) + + def on_stderr(msg: OutputMessage): + stderr_messages.append(msg) + + def on_error(err: ExecutionError): + errors.append(err) + + def on_complete(evt: ExecutionComplete): + completed_events.append(evt) + + def on_init(evt: ExecutionInit): + init_events.append(evt) + + handlers = ExecutionHandlersSync( + on_stdout=on_stdout, + on_stderr=on_stderr, + on_result=None, + on_error=on_error, + on_execution_complete=on_complete, + on_init=on_init, + ) + + # Send three sequential commands in the same session, validating env propagation. + code1 = ( + "export FOO=hello\n" + "export BAR=world\n" + ) + code2 = ( + "printf \"step1:$FOO:$BAR\\n\"\n" + ) + code3 = ( + "export FOO=${FOO}_next\n" + "printf \"step2:$FOO:$BAR\\n\"\n" + "export BAR=${BAR}_next\n" + "printf \"step3:$FOO:$BAR\\n\"\n" + ) + + # export envs + result1 = code_interpreter.codes.run( + code1, + language=SupportedLanguage.BASH, + handlers=handlers, + ) + + assert result1 is not None + assert result1.id is not None and str(result1.id).strip() + assert result1.error is None + + # print env + result2 = code_interpreter.codes.run( + code2, + language=SupportedLanguage.BASH, + handlers=handlers, + ) + + assert result2 is not None + assert result2.id is not None and str(result2.id).strip() + assert result2.error is None + + # print env + result3 = code_interpreter.codes.run( + code3, + language=SupportedLanguage.BASH, + handlers=handlers, + ) + assert result3 is not None + assert result3.id is not None and str(result3.id).strip() + assert result3.error is None + + # Expect at least three stdout lines with propagated env values. + stdout_texts = [m.text.strip() for m in stdout_messages if m.text] + assert "step1:hello:world" in stdout_texts + assert "step2:hello_next:world" in stdout_texts + assert "step3:hello_next:world_next" in stdout_texts + for m in stdout_messages[:3]: + _assert_recent_timestamp_ms(m.timestamp)