371 lines
8.8 KiB
Go
371 lines
8.8 KiB
Go
package sandbox
|
|
|
|
import (
|
|
"archive/tar"
|
|
"bytes"
|
|
"context"
|
|
"fmt"
|
|
"io"
|
|
"log/slog"
|
|
"maps"
|
|
"os"
|
|
"path/filepath"
|
|
|
|
"github.com/criyle/go-judge/pb"
|
|
"github.com/joint-online-judge/JOJ3/internal/stage"
|
|
"google.golang.org/protobuf/proto"
|
|
)
|
|
|
|
const (
|
|
tarStreamThreshold = 128 * 1024
|
|
streamChunkSize = 32 * 1024
|
|
)
|
|
|
|
func (e *Sandbox) Run(cmds []stage.Cmd) ([]stage.ExecutorResult, error) {
|
|
var err error
|
|
if e.execClient == nil {
|
|
slog.Debug("create exec client", "server", e.execServer)
|
|
e.execClient, err = createExecClient(e.execServer, e.token)
|
|
if err != nil {
|
|
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 {
|
|
cmd.CopyIn = make(map[string]stage.CmdFile)
|
|
}
|
|
for k, v := range cmd.CopyInCached {
|
|
if fileID, ok := e.cachedMap[v]; ok {
|
|
cmd.CopyIn[k] = stage.CmdFile{FileID: &fileID}
|
|
}
|
|
}
|
|
}
|
|
if needStream, tarData := prepareTarStream(cmds); needStream {
|
|
return e.runWithTarStream(cmds, tarData)
|
|
}
|
|
return e.runUnary(cmds)
|
|
}
|
|
|
|
func prepareTarStream(cmds []stage.Cmd) (bool, []byte) {
|
|
if len(cmds) == 0 {
|
|
return false, nil
|
|
}
|
|
for i := range cmds {
|
|
if cmds[i].CopyInDir != "" &&
|
|
estimateCopyInSize(&cmds[i]) >= tarStreamThreshold {
|
|
tarData, keysInTar := createCopyInTar(&cmds[i])
|
|
for j := range cmds {
|
|
cmds[j].CopyInDir = ""
|
|
for _, k := range keysInTar {
|
|
delete(cmds[j].CopyIn, k)
|
|
}
|
|
}
|
|
return true, tarData
|
|
}
|
|
}
|
|
return false, nil
|
|
}
|
|
|
|
func (e *Sandbox) runUnary(cmds []stage.Cmd) ([]stage.ExecutorResult, error) {
|
|
pbCmds := convertPBCmd(cmds)
|
|
for i, pbCmd := range pbCmds {
|
|
slog.Debug("sandbox execute", "i", i, "pbCmd size", proto.Size(pbCmd))
|
|
}
|
|
pbReq := &pb.Request{}
|
|
pbReq.SetCmd(pbCmds)
|
|
slog.Debug("sandbox execute", "pbReq size", proto.Size(pbReq))
|
|
pbRet, err := e.execClient.Exec(context.TODO(), pbReq)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if pbRet.GetError() != "" {
|
|
return nil, fmt.Errorf("sandbox execute error: %s", pbRet.GetError())
|
|
}
|
|
results := convertPBResult(pbRet.GetResults())
|
|
for _, result := range results {
|
|
maps.Copy(e.cachedMap, result.FileIDs)
|
|
}
|
|
return results, nil
|
|
}
|
|
|
|
func (e *Sandbox) runWithTarStream(cmds []stage.Cmd, tarData []byte) ([]stage.ExecutorResult, error) {
|
|
tarCmd := stage.Cmd{
|
|
Args: []string{"/bin/tar", "xf", "-", "-C", "/w"},
|
|
Env: cmds[0].Env,
|
|
Stdin: &stage.CmdFile{StreamIn: true},
|
|
CPULimit: cmds[0].CPULimit,
|
|
ClockLimit: cmds[0].ClockLimit,
|
|
MemoryLimit: cmds[0].MemoryLimit,
|
|
StackLimit: cmds[0].StackLimit,
|
|
ProcLimit: cmds[0].ProcLimit,
|
|
CPURateLimit: cmds[0].CPURateLimit,
|
|
CPUSetLimit: cmds[0].CPUSetLimit,
|
|
}
|
|
|
|
allCmds := append([]stage.Cmd{tarCmd}, cmds...)
|
|
pbCmds := convertPBCmd(allCmds)
|
|
|
|
pbReq := &pb.Request{}
|
|
pbReq.SetCmd(pbCmds)
|
|
slog.Debug("sandbox stream execute", "pbReq size", proto.Size(pbReq), "tarSize", len(tarData))
|
|
|
|
stream, err := e.execClient.ExecStream(context.TODO())
|
|
if err != nil {
|
|
return nil, fmt.Errorf("exec stream: %w", err)
|
|
}
|
|
|
|
sr := &pb.StreamRequest{}
|
|
sr.SetExecRequest(pbReq)
|
|
if err := stream.Send(sr); err != nil {
|
|
return nil, fmt.Errorf("stream send request: %w", err)
|
|
}
|
|
|
|
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)
|
|
}
|
|
}
|
|
|
|
if err := stream.CloseSend(); err != nil {
|
|
return nil, fmt.Errorf("stream close send: %w", 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)
|
|
}
|
|
}
|
|
|
|
func estimateCopyInSize(cmd *stage.Cmd) int {
|
|
total := 0
|
|
if cmd.CopyInDir != "" {
|
|
_ = filepath.Walk(cmd.CopyInDir,
|
|
func(path string, info os.FileInfo, err error) error {
|
|
if err != nil || info.IsDir() {
|
|
return nil
|
|
}
|
|
relPath, err := filepath.Rel(cmd.CopyInDir, path)
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
if _, exists := cmd.CopyIn[relPath]; !exists {
|
|
total += int(info.Size())
|
|
}
|
|
return nil
|
|
})
|
|
}
|
|
for _, f := range cmd.CopyIn {
|
|
if f.Symlink != nil {
|
|
continue
|
|
}
|
|
if f.Src != nil {
|
|
if fi, err := os.Stat(*f.Src); err == nil {
|
|
total += int(fi.Size())
|
|
}
|
|
} else if f.Content != nil {
|
|
total += len(*f.Content)
|
|
}
|
|
}
|
|
return total
|
|
}
|
|
|
|
func createCopyInTar(cmd *stage.Cmd) ([]byte, []string) {
|
|
var buf bytes.Buffer
|
|
tw := tar.NewWriter(&buf)
|
|
tarKeys := make([]string, 0, len(cmd.CopyIn))
|
|
|
|
if cmd.CopyInDir != "" {
|
|
err := filepath.Walk(cmd.CopyInDir,
|
|
func(path string, info os.FileInfo, err error) error {
|
|
if err != nil {
|
|
return err
|
|
}
|
|
relPath, err := filepath.Rel(cmd.CopyInDir, path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if relPath == "." {
|
|
return nil
|
|
}
|
|
if _, exists := cmd.CopyIn[relPath]; exists {
|
|
if info.IsDir() {
|
|
return filepath.SkipDir
|
|
}
|
|
return nil
|
|
}
|
|
if info.Mode()&os.ModeSymlink != 0 {
|
|
link, err := os.Readlink(path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
hdr, err := tar.FileInfoHeader(info, link)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
hdr.Name = relPath
|
|
return tw.WriteHeader(hdr)
|
|
}
|
|
hdr, err := tar.FileInfoHeader(info, "")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
hdr.Name = relPath
|
|
if info.IsDir() {
|
|
hdr.Name += "/"
|
|
}
|
|
if err := tw.WriteHeader(hdr); err != nil {
|
|
return err
|
|
}
|
|
if info.IsDir() {
|
|
return nil
|
|
}
|
|
f, err := os.Open(path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
_, err = io.Copy(tw, f)
|
|
f.Close()
|
|
return err
|
|
})
|
|
if err != nil {
|
|
slog.Error("create copyIn tar walk", "error", err)
|
|
return nil, nil
|
|
}
|
|
}
|
|
|
|
for k, f := range cmd.CopyIn {
|
|
if f.FileID != nil || f.Symlink != nil || f.StreamIn || f.StreamOut || f.Pipe {
|
|
continue
|
|
}
|
|
if f.Content != nil {
|
|
hdr := &tar.Header{
|
|
Name: k,
|
|
Mode: 0o644,
|
|
Size: int64(len(*f.Content)),
|
|
}
|
|
if err := tw.WriteHeader(hdr); err != nil {
|
|
slog.Error("create copyIn tar write header", "key", k, "error", err)
|
|
continue
|
|
}
|
|
if _, err := tw.Write([]byte(*f.Content)); err != nil {
|
|
slog.Error("create copyIn tar write content", "key", k, "error", err)
|
|
continue
|
|
}
|
|
tarKeys = append(tarKeys, k)
|
|
} else if f.Src != nil {
|
|
fi, err := os.Stat(*f.Src)
|
|
if err != nil {
|
|
slog.Error("create copyIn tar stat", "key", k, "src", *f.Src, "error", err)
|
|
continue
|
|
}
|
|
if fi.IsDir() {
|
|
continue
|
|
}
|
|
hdr, err := tar.FileInfoHeader(fi, "")
|
|
if err != nil {
|
|
slog.Error("create copyIn tar file info header", "key", k, "error", err)
|
|
continue
|
|
}
|
|
hdr.Name = k
|
|
if err := tw.WriteHeader(hdr); err != nil {
|
|
slog.Error("create copyIn tar write header", "key", k, "error", err)
|
|
continue
|
|
}
|
|
srcFile, err := os.Open(*f.Src)
|
|
if err != nil {
|
|
slog.Error("create copyIn tar open src", "key", k, "src", *f.Src, "error", err)
|
|
continue
|
|
}
|
|
_, err = io.Copy(tw, srcFile)
|
|
srcFile.Close()
|
|
if err != nil {
|
|
slog.Error("create copyIn tar copy", "key", k, "error", err)
|
|
continue
|
|
}
|
|
tarKeys = append(tarKeys, k)
|
|
}
|
|
}
|
|
|
|
if err := tw.Close(); err != nil {
|
|
slog.Error("create copyIn tar close", "error", err)
|
|
return nil, nil
|
|
}
|
|
return buf.Bytes(), tarKeys
|
|
}
|
|
|
|
func (e *Sandbox) Cleanup() error {
|
|
for k, fileID := range e.cachedMap {
|
|
req := &pb.FileID{}
|
|
req.SetFileID(fileID)
|
|
_, err := e.execClient.FileDelete(context.TODO(), req)
|
|
if err != nil {
|
|
slog.Error("sandbox cleanup", "error", err)
|
|
}
|
|
delete(e.cachedMap, k)
|
|
}
|
|
return nil
|
|
}
|