diff --git a/internal/executor/sandbox/executor.go b/internal/executor/sandbox/executor.go index 97daa8e..39f2461 100644 --- a/internal/executor/sandbox/executor.go +++ b/internal/executor/sandbox/executor.go @@ -16,10 +16,7 @@ import ( "google.golang.org/protobuf/proto" ) -const ( - tarStreamThreshold = 128 * 1024 - streamChunkSize = 32 * 1024 -) +const tarThreshold = 128 * 1024 func (e *Sandbox) Run(cmds []stage.Cmd) ([]stage.ExecutorResult, error) { var err error @@ -30,7 +27,6 @@ func (e *Sandbox) Run(cmds []stage.Cmd) ([]stage.ExecutorResult, error) { return nil, err } } - // cannot use range loop since we need to change the value for i := 0; i < len(cmds); i += 1 { cmd := &cmds[i] if cmd.CopyIn == nil { @@ -42,19 +38,19 @@ func (e *Sandbox) Run(cmds []stage.Cmd) ([]stage.ExecutorResult, error) { } } } - if needStream, tarData := prepareTarStream(cmds); needStream { - return e.runWithTarStream(cmds, tarData) + if needTar, tarData := prepareTar(cmds); needTar { + return e.runWithTar(cmds, tarData) } return e.runUnary(cmds) } -func prepareTarStream(cmds []stage.Cmd) (bool, []byte) { +func prepareTar(cmds []stage.Cmd) (bool, []byte) { if len(cmds) == 0 { return false, nil } for i := range cmds { if cmds[i].CopyInDir != "" && - estimateCopyInSize(&cmds[i]) >= tarStreamThreshold { + estimateCopyInSize(&cmds[i]) >= tarThreshold { tarData, keysInTar := createCopyInTar(&cmds[i]) for j := range cmds { cmds[j].CopyInDir = "" @@ -62,6 +58,9 @@ func prepareTarStream(cmds []stage.Cmd) (bool, []byte) { delete(cmds[j].CopyIn, k) } } + if tarData == nil { + return false, nil + } return true, tarData } } @@ -90,11 +89,20 @@ func (e *Sandbox) runUnary(cmds []stage.Cmd) ([]stage.ExecutorResult, error) { return results, nil } -func (e *Sandbox) runWithTarStream(cmds []stage.Cmd, tarData []byte) ([]stage.ExecutorResult, error) { +func (e *Sandbox) runWithTar(cmds []stage.Cmd, tarData []byte) ([]stage.ExecutorResult, error) { + fc := &pb.FileContent{} + fc.SetContent(tarData) + fileIDResp, err := e.execClient.FileAdd(context.TODO(), fc) + if err != nil { + return nil, fmt.Errorf("file add tar: %w", err) + } + fid := fileIDResp.GetFileID() + slog.Debug("sandbox tar uploaded", "fileID", fid, "tarSize", len(tarData)) + + tarFileName := "/w/__joj3_copyin.tar" tarCmd := stage.Cmd{ - Args: []string{"/bin/tar", "xf", "-", "-C", "/w"}, + Args: []string{"/bin/tar", "xf", tarFileName, "-C", "/w"}, Env: cmds[0].Env, - Stdin: &stage.CmdFile{StreamIn: true}, CPULimit: cmds[0].CPULimit, ClockLimit: cmds[0].ClockLimit, MemoryLimit: cmds[0].MemoryLimit, @@ -102,6 +110,10 @@ func (e *Sandbox) runWithTarStream(cmds []stage.Cmd, tarData []byte) ([]stage.Ex ProcLimit: cmds[0].ProcLimit, CPURateLimit: cmds[0].CPURateLimit, CPUSetLimit: cmds[0].CPUSetLimit, + CopyIn: map[string]stage.CmdFile{ + tarFileName: {FileID: &fid}, + }, + CopyOutCached: []string{tarFileName}, } allCmds := append([]stage.Cmd{tarCmd}, cmds...) @@ -109,94 +121,40 @@ func (e *Sandbox) runWithTarStream(cmds []stage.Cmd, tarData []byte) ([]stage.Ex pbReq := &pb.Request{} pbReq.SetCmd(pbCmds) - slog.Debug("sandbox stream execute", "pbReq size", proto.Size(pbReq), "tarSize", len(tarData)) + slog.Debug("sandbox tar exec", "pbReq size", proto.Size(pbReq), "tarSize", len(tarData)) - stream, err := e.execClient.ExecStream(context.TODO()) + pbRet, err := e.execClient.Exec(context.TODO(), pbReq) if err != nil { - return nil, fmt.Errorf("exec stream: %w", err) + return nil, fmt.Errorf("tar exec: %w", err) + } + if pbRet.GetError() != "" { + return nil, fmt.Errorf("sandbox execute error: %s", pbRet.GetError()) } - sr := &pb.StreamRequest{} - sr.SetExecRequest(pbReq) - if err := stream.Send(sr); err != nil { - return nil, fmt.Errorf("stream send request: %w", err) + results := convertPBResult(pbRet.GetResults()) + if len(results) == 0 { + return nil, fmt.Errorf("tar exec: empty results") + } + tarResult := results[0] + if tarResult.Status != stage.StatusAccepted || tarResult.ExitStatus != 0 { + return nil, fmt.Errorf( + "tar extraction failed: status=%v, exit=%d, error=%s", + tarResult.Status, tarResult.ExitStatus, tarResult.Error, + ) } - errCh := make(chan error, 1) - respCh := make(chan *pb.Response, 1) - go func() { - for { - resp, recvErr := stream.Recv() - if recvErr == io.EOF { - return - } - if recvErr != nil { - errCh <- recvErr - return - } - if resp.HasExecResponse() { - er := resp.GetExecResponse() - if er.GetError() != "" { - errCh <- fmt.Errorf("server error: %s", er.GetError()) - return - } - respCh <- er - } - } - }() - - for offset := 0; offset < len(tarData); offset += streamChunkSize { - end := offset + streamChunkSize - if end > len(tarData) { - end = len(tarData) - } - select { - case e := <-errCh: - return nil, fmt.Errorf("stream send input: %w", e) - default: - } - input := pb.StreamRequest_builder{ - ExecInput: (&pb.StreamRequest_Input_builder{ - Index: 0, - Fd: 0, - Content: tarData[offset:end], - }).Build(), - }.Build() - if err := stream.Send(input); err != nil { - select { - case e := <-errCh: - return nil, fmt.Errorf("stream send input: %w (server: %v)", err, e) - default: - } - return nil, fmt.Errorf("stream send input: %w", err) - } + results = results[1:] + for _, result := range results { + maps.Copy(e.cachedMap, result.FileIDs) } - if err := stream.CloseSend(); err != nil { - return nil, fmt.Errorf("stream close send: %w", err) + deleteReq := &pb.FileID{} + deleteReq.SetFileID(fid) + if _, err := e.execClient.FileDelete(context.TODO(), deleteReq); err != nil { + slog.Warn("sandbox tar file delete", "fileID", fid, "error", err) } - select { - case finalResponse := <-respCh: - results := convertPBResult(finalResponse.GetResults()) - if len(results) == 0 { - return nil, fmt.Errorf("stream: empty results") - } - tarResult := results[0] - if tarResult.Status != stage.StatusAccepted || tarResult.ExitStatus != 0 { - return nil, fmt.Errorf( - "tar extraction failed: status=%v, exit=%d, error=%s", - tarResult.Status, tarResult.ExitStatus, tarResult.Error, - ) - } - results = results[1:] - for _, result := range results { - maps.Copy(e.cachedMap, result.FileIDs) - } - return results, nil - case e := <-errCh: - return nil, fmt.Errorf("stream recv: %w", e) - } + return results, nil } func estimateCopyInSize(cmd *stage.Cmd) int {