diff --git a/cmd/main-server.go b/cmd/main-server.go index eefea466..7202f904 100644 --- a/cmd/main-server.go +++ b/cmd/main-server.go @@ -2,14 +2,19 @@ package main import ( "context" + "encoding/base64" "encoding/json" "errors" "fmt" + "io" "io/fs" "log" + "mime/multipart" "net/http" "os" "os/signal" + "path/filepath" + "regexp" "runtime/debug" "strconv" "strings" @@ -20,6 +25,8 @@ import ( "github.com/google/uuid" "github.com/gorilla/mux" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/server" "github.com/commandlinedev/prompt-server/pkg/cmdrunner" "github.com/commandlinedev/prompt-server/pkg/pcloud" "github.com/commandlinedev/prompt-server/pkg/remote" @@ -49,11 +56,14 @@ const InitialTelemetryWait = 30 * time.Second const TelemetryTick = 30 * time.Minute const TelemetryInterval = 8 * time.Hour +const MaxWriteFileMemSize = 20 * (1024 * 1024) // 20M + var GlobalLock = &sync.Mutex{} var WSStateMap = make(map[string]*scws.WSState) // clientid -> WsState var GlobalAuthKey string var BuildTime = "0" var shutdownOnce sync.Once +var ContentTypeHeaderValidRe = regexp.MustCompile(`^\w+/[\w.+-]+$`) type ClientActiveState struct { Fg bool `json:"fg"` @@ -312,6 +322,273 @@ func HandleGetPtyOut(w http.ResponseWriter, r *http.Request) { w.Write(data) } +type writeFileParamsType struct { + ScreenId string `json:"screenid"` + LineId string `json:"lineid"` + Path string `json:"path"` + UseTemp bool `json:"usetemp,omitempty"` +} + +func parseWriteFileParams(r *http.Request) (*writeFileParamsType, multipart.File, error) { + err := r.ParseMultipartForm(MaxWriteFileMemSize) + if err != nil { + return nil, nil, fmt.Errorf("cannot parse multipart form data: %v", err) + } + form := r.MultipartForm + if len(form.Value["params"]) == 0 { + return nil, nil, fmt.Errorf("no params found") + } + paramsStr := form.Value["params"][0] + var params writeFileParamsType + err = json.Unmarshal([]byte(paramsStr), ¶ms) + if err != nil { + return nil, nil, fmt.Errorf("bad params json: %v", err) + } + if len(form.File["data"]) == 0 { + return nil, nil, fmt.Errorf("no data found") + } + fileHeader := form.File["data"][0] + file, err := fileHeader.Open() + if err != nil { + return nil, nil, fmt.Errorf("error opening multipart data file: %v", err) + } + return ¶ms, file, nil +} + +func HandleWriteFile(w http.ResponseWriter, r *http.Request) { + defer func() { + r := recover() + if r == nil { + return + } + log.Printf("[error] in write-file: %v\n", r) + debug.PrintStack() + WriteJsonError(w, fmt.Errorf("panic: %v", r)) + return + }() + w.Header().Set("Cache-Control", "no-cache") + params, mpFile, err := parseWriteFileParams(r) + if err != nil { + WriteJsonError(w, fmt.Errorf("error parsing multipart form params: %w", err)) + return + } + if params.ScreenId == "" || params.LineId == "" || params.Path == "" { + WriteJsonError(w, fmt.Errorf("invalid params, must set screenid, lineid, and path")) + return + } + if _, err := uuid.Parse(params.ScreenId); err != nil { + WriteJsonError(w, fmt.Errorf("invalid screenid: %v", err)) + return + } + if _, err := uuid.Parse(params.LineId); err != nil { + WriteJsonError(w, fmt.Errorf("invalid lineid: %v", err)) + return + } + _, cmd, err := sstore.GetLineCmdByLineId(r.Context(), params.ScreenId, params.LineId) + if err != nil { + WriteJsonError(w, fmt.Errorf("cannot retrieve line/cmd: %v", err)) + return + } + if cmd == nil { + WriteJsonError(w, fmt.Errorf("line not found")) + return + } + if cmd.Remote.RemoteId == "" { + WriteJsonError(w, fmt.Errorf("invalid line, no remote")) + return + } + msh := remote.GetRemoteById(cmd.Remote.RemoteId) + if msh == nil { + WriteJsonError(w, fmt.Errorf("invalid line, cannot resolve remote")) + return + } + cwd := cmd.FeState["cwd"] + writePk := packet.MakeWriteFilePacket() + writePk.ReqId = uuid.New().String() + writePk.UseTemp = params.UseTemp + if filepath.IsAbs(params.Path) { + writePk.Path = params.Path + } else { + writePk.Path = filepath.Join(cwd, params.Path) + } + iter, err := msh.PacketRpcIter(r.Context(), writePk) + if err != nil { + WriteJsonError(w, fmt.Errorf("error: %v", err)) + return + } + // first packet should be WriteFileReady + readyIf, err := iter.Next(r.Context()) + if err != nil { + WriteJsonError(w, fmt.Errorf("error while getting ready response: %w", err)) + return + } + readyPk, ok := readyIf.(*packet.WriteFileReadyPacketType) + if !ok { + WriteJsonError(w, fmt.Errorf("bad ready packet received: %T", readyIf)) + return + } + if readyPk.Error != "" { + WriteJsonError(w, fmt.Errorf("ready error: %s", readyPk.Error)) + return + } + var buffer [server.MaxFileDataPacketSize]byte + bufSlice := buffer[:] + for { + dataPk := packet.MakeFileDataPacket(writePk.ReqId) + nr, err := io.ReadFull(mpFile, bufSlice) + if err == io.ErrUnexpectedEOF || err == io.EOF { + dataPk.Eof = true + } else if err != nil { + dataErr := fmt.Errorf("error reading file data: %v", err) + dataPk.Error = dataErr.Error() + msh.SendFileData(dataPk) + WriteJsonError(w, dataErr) + return + } + if nr > 0 { + dataPk.Data = make([]byte, nr) + copy(dataPk.Data, bufSlice[0:nr]) + } + msh.SendFileData(dataPk) + if dataPk.Eof { + break + } + // slight throttle for sending packets + time.Sleep(10 * time.Millisecond) + } + doneIf, err := iter.Next(r.Context()) + if err != nil { + WriteJsonError(w, fmt.Errorf("error while getting done response: %w", err)) + return + } + donePk, ok := doneIf.(*packet.WriteFileDonePacketType) + if !ok { + WriteJsonError(w, fmt.Errorf("bad done packet received: %T", doneIf)) + return + } + if donePk.Error != "" { + WriteJsonError(w, fmt.Errorf("dne error: %s", donePk.Error)) + return + } + WriteJsonSuccess(w, nil) + return +} + +func HandleReadFile(w http.ResponseWriter, r *http.Request) { + qvals := r.URL.Query() + screenId := qvals.Get("screenid") + lineId := qvals.Get("lineid") + path := qvals.Get("path") // validate path? + contentType := qvals.Get("mimetype") + if contentType == "" { + contentType = "application/octet-stream" + } + if screenId == "" || lineId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify sessionid, screenid, and lineid"))) + return + } + if path == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify path"))) + return + } + if _, err := uuid.Parse(screenId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid screenid: %v", err))) + return + } + if _, err := uuid.Parse(lineId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid lineid: %v", err))) + return + } + if !ContentTypeHeaderValidRe.MatchString(contentType) { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid mimetype specified"))) + return + } + _, cmd, err := sstore.GetLineCmdByLineId(r.Context(), screenId, lineId) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid lineid: %v", err))) + return + } + if cmd == nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid line, no cmd"))) + return + } + if cmd.Remote.RemoteId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid line, no remote"))) + return + } + streamPk := packet.MakeStreamFilePacket() + streamPk.ReqId = uuid.New().String() + cwd := cmd.FeState["cwd"] + if filepath.IsAbs(path) { + streamPk.Path = path + } else { + streamPk.Path = filepath.Join(cwd, path) + } + msh := remote.GetRemoteById(cmd.Remote.RemoteId) + if msh == nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid line, cannot resolve remote"))) + return + } + iter, err := msh.StreamFile(r.Context(), streamPk) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error trying to stream file: %v", err))) + return + } + defer iter.Close() + respIf, err := iter.Next(r.Context()) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error getting streamfile response: %v", err))) + return + } + resp, ok := respIf.(*packet.StreamFileResponseType) + if !ok { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("bad response packet type: %T", respIf))) + return + } + if resp.Error != "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error response: %s", resp.Error))) + return + } + infoJson, _ := json.Marshal(resp.Info) + w.Header().Set("X-FileInfo", base64.StdEncoding.EncodeToString(infoJson)) + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusOK) + for { + dataPkIf, err := iter.Next(r.Context()) + if err != nil { + log.Printf("error in read-file while getting data: %v\n", err) + break + } + if dataPkIf == nil { + break + } + dataPk, ok := dataPkIf.(*packet.FileDataPacketType) + if !ok { + log.Printf("error in read-file, invalid data packet type: %T", dataPkIf) + break + } + if dataPk.Error != "" { + log.Printf("in read-file, data packet error: %s", dataPk.Error) + break + } + w.Write(dataPk.Data) + } + return +} + func WriteJsonError(w http.ResponseWriter, errVal error) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(200) @@ -576,6 +853,8 @@ func main() { gr.HandleFunc("/api/get-client-data", AuthKeyWrap(HandleGetClientData)) gr.HandleFunc("/api/set-winsize", AuthKeyWrap(HandleSetWinSize)) gr.HandleFunc("/api/log-active-state", AuthKeyWrap(HandleLogActiveState)) + gr.HandleFunc("/api/read-file", AuthKeyWrap(HandleReadFile)) + gr.HandleFunc("/api/write-file", AuthKeyWrap(HandleWriteFile)).Methods("POST") serverAddr := MainServerAddr if scbase.IsDevMode() { serverAddr = MainServerDevAddr diff --git a/pkg/cmdrunner/cmdrunner.go b/pkg/cmdrunner/cmdrunner.go index 6d89f487..805546c6 100644 --- a/pkg/cmdrunner/cmdrunner.go +++ b/pkg/cmdrunner/cmdrunner.go @@ -6,9 +6,11 @@ import ( "crypto/rand" "encoding/base64" "fmt" + "io/fs" "log" "net/url" "os" + "path/filepath" "regexp" "sort" "strconv" @@ -56,6 +58,8 @@ const MaxEvalDepth = 5 const MaxOpenAIAPITokenLen = 100 const MaxOpenAIModelLen = 100 +const TsFormatStr = "2006-01-02 15:04:05" + const ( KwArgRenderer = "renderer" KwArgView = "view" @@ -205,6 +209,11 @@ func init() { registerCmdFn("_killserver", KillServerCommand) registerCmdFn("set", SetCommand) + + registerCmdFn("view:stat", ViewStatCommand) + registerCmdFn("view:test", ViewTestCommand) + + registerCmdFn("edit:test", EditTestCommand) } func getValidCommands() []string { @@ -2115,7 +2124,7 @@ func SessionShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) ( if session.Archived { buf.WriteString(fmt.Sprintf(" %-15s %s\n", "archived", "true")) ts := time.UnixMilli(session.ArchivedTs) - buf.WriteString(fmt.Sprintf(" %-15s %s\n", "archivedts", ts.Format("2006-01-02 15:04:05"))) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "archivedts", ts.Format(TsFormatStr))) } stats, err := sstore.GetSessionStats(ctx, ids.SessionId) if err != nil { @@ -2936,6 +2945,7 @@ func LineShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sst return nil, fmt.Errorf("line %q not found", lineArg) } var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "screenid", line.ScreenId)) buf.WriteString(fmt.Sprintf(" %-15s %s\n", "lineid", line.LineId)) buf.WriteString(fmt.Sprintf(" %-15s %s\n", "type", line.LineType)) lineNumStr := strconv.FormatInt(line.LineNum, 10) @@ -2944,7 +2954,7 @@ func LineShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sst } buf.WriteString(fmt.Sprintf(" %-15s %s\n", "linenum", lineNumStr)) ts := time.UnixMilli(line.Ts) - buf.WriteString(fmt.Sprintf(" %-15s %s\n", "ts", ts.Format("2006-01-02 15:04:05"))) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "ts", ts.Format(TsFormatStr))) if line.Ephemeral { buf.WriteString(fmt.Sprintf(" %-15s %v\n", "ephemeral", true)) } @@ -2974,6 +2984,12 @@ func LineShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sst buf.WriteString(fmt.Sprintf(" %-15s %s\n", "file", stat.Location)) buf.WriteString(fmt.Sprintf(" %-15s %s\n", "file-data", fileDataStr)) } + if cmd.DoneTs != 0 { + doneTs := time.UnixMilli(cmd.DoneTs) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "donets", doneTs.Format(TsFormatStr))) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "exitcode", cmd.ExitCode)) + buf.WriteString(fmt.Sprintf(" %-15s %dms\n", "duration", cmd.DurationMs)) + } } update := &sstore.ModelUpdate{ Info: &sstore.InfoMsgType{ @@ -3010,6 +3026,205 @@ func SetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.U return nil, nil } +func makeStreamFilePk(ids resolvedIds, pk *scpacket.FeCommandPacketType) (*packet.StreamFilePacketType, error) { + cwd := ids.Remote.FeState["cwd"] + fileArg := pk.Args[0] + if fileArg == "" { + return nil, fmt.Errorf("/view:stat file argument must be set (cannot be empty)") + } + streamPk := packet.MakeStreamFilePacket() + streamPk.ReqId = uuid.New().String() + if filepath.IsAbs(fileArg) { + streamPk.Path = fileArg + } else { + streamPk.Path = filepath.Join(cwd, fileArg) + } + return streamPk, nil +} + +func ViewStatCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/view:stat requires an argument (file name)") + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + streamPk, err := makeStreamFilePk(ids, pk) + if err != nil { + return nil, err + } + streamPk.StatOnly = true + msh := ids.Remote.MShell + iter, err := msh.StreamFile(ctx, streamPk) + if err != nil { + return nil, fmt.Errorf("/view:stat error: %v", err) + } + defer iter.Close() + respIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/view:stat error getting response: %v", err) + } + resp, ok := respIf.(*packet.StreamFileResponseType) + if !ok { + return nil, fmt.Errorf("/view:stat error, bad response packet type: %T", respIf) + } + if resp.Error != "" { + return nil, fmt.Errorf("/view:stat error: %s", resp.Error) + } + if resp.Info == nil { + return nil, fmt.Errorf("/view:stat error, no file info") + } + var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "path", resp.Info.Name)) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "size", resp.Info.Size)) + modTs := time.UnixMilli(resp.Info.ModTs) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "modts", modTs.Format(TsFormatStr))) + buf.WriteString(fmt.Sprintf(" %-15s %v\n", "isdir", resp.Info.IsDir)) + modeStr := fs.FileMode(resp.Info.Perm).String() + if len(modeStr) > 9 { + modeStr = modeStr[len(modeStr)-9:] + } + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "perms", modeStr)) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("view stat %q", streamPk.Path), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func ViewTestCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/view:test requires an argument (file name)") + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + streamPk, err := makeStreamFilePk(ids, pk) + if err != nil { + return nil, err + } + msh := ids.Remote.MShell + iter, err := msh.StreamFile(ctx, streamPk) + if err != nil { + return nil, fmt.Errorf("/view:test error: %v", err) + } + defer iter.Close() + respIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/view:test error getting response: %v", err) + } + resp, ok := respIf.(*packet.StreamFileResponseType) + if !ok { + return nil, fmt.Errorf("/view:test error, bad response packet type: %T", respIf) + } + if resp.Error != "" { + return nil, fmt.Errorf("/view:test error: %s", resp.Error) + } + if resp.Info == nil { + return nil, fmt.Errorf("/view:test error, no file info") + } + var buf bytes.Buffer + var numPackets int + for { + dataPkIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/view:test error while getting data: %w", err) + } + if dataPkIf == nil { + break + } + dataPk, ok := dataPkIf.(*packet.FileDataPacketType) + if !ok { + return nil, fmt.Errorf("/view:test invalid data packet type: %T", dataPkIf) + } + if dataPk.Error != "" { + return nil, fmt.Errorf("/view:test error returned while getting data: %s", dataPk.Error) + } + numPackets++ + buf.Write(dataPk.Data) + } + buf.WriteString(fmt.Sprintf("\n\ntotal packets: %d\n", numPackets)) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("view file %q", streamPk.Path), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func EditTestCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/edit:test requires an argument (file name)") + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + content, ok := pk.Kwargs["content"] + if !ok { + return nil, fmt.Errorf("/edit:test no content for file specified") + } + fileArg := pk.Args[0] + if fileArg == "" { + return nil, fmt.Errorf("/view:stat file argument must be set (cannot be empty)") + } + writePk := packet.MakeWriteFilePacket() + writePk.ReqId = uuid.New().String() + writePk.UseTemp = true + cwd := ids.Remote.FeState["cwd"] + if filepath.IsAbs(fileArg) { + writePk.Path = fileArg + } else { + writePk.Path = filepath.Join(cwd, fileArg) + } + msh := ids.Remote.MShell + iter, err := msh.PacketRpcIter(ctx, writePk) + if err != nil { + return nil, fmt.Errorf("/edit:test error: %v", err) + } + // first packet should be WriteFileReady + readyIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/edit:test error while getting ready response: %w", err) + } + readyPk, ok := readyIf.(*packet.WriteFileReadyPacketType) + if !ok { + return nil, fmt.Errorf("/edit:test bad ready packet received: %T", readyIf) + } + if readyPk.Error != "" { + return nil, fmt.Errorf("/edit:test %s", readyPk.Error) + } + dataPk := packet.MakeFileDataPacket(writePk.ReqId) + dataPk.Data = []byte(content) + dataPk.Eof = true + err = msh.SendFileData(dataPk) + if err != nil { + return nil, fmt.Errorf("/edit:test error sending data packet: %v", err) + } + doneIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/edit:test error while getting done response: %w", err) + } + donePk, ok := doneIf.(*packet.WriteFileDonePacketType) + if !ok { + return nil, fmt.Errorf("/edit:test bad done packet received: %T", doneIf) + } + if donePk.Error != "" { + return nil, fmt.Errorf("/edit:test %s", donePk.Error) + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("edit test, wrote %q", writePk.Path), + }, + } + return update, nil +} + func SignalCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) if err != nil { diff --git a/pkg/cmdrunner/resolver.go b/pkg/cmdrunner/resolver.go index 64c264fd..7142a369 100644 --- a/pkg/cmdrunner/resolver.go +++ b/pkg/cmdrunner/resolver.go @@ -8,10 +8,10 @@ import ( "strconv" "strings" - "github.com/google/uuid" "github.com/commandlinedev/prompt-server/pkg/remote" "github.com/commandlinedev/prompt-server/pkg/scpacket" "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/google/uuid" ) const ( @@ -242,7 +242,7 @@ func resolveUiIds(ctx context.Context, pk *scpacket.FeCommandPacketType, rtype i if err != nil { return rtn, fmt.Errorf("invalid resolved remote: %v", err) } - rr, err := resolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.ScreenId) + rr, err := ResolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.ScreenId) if err != nil { return rtn, err } @@ -263,7 +263,7 @@ func resolveUiIds(ctx context.Context, pk *scpacket.FeCommandPacketType, rtype i if err != nil { return rtn, fmt.Errorf("error trying to auto-connect remote [%s]: %w", rtn.Remote.DisplayName, err) } - rrNew, err := resolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.ScreenId) + rrNew, err := ResolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.ScreenId) if err != nil { return rtn, err } @@ -450,7 +450,7 @@ func parseFullRemoteRef(fullRemoteRef string) (string, string, string, error) { return fields[0], fields[1], fields[2], nil } -func resolveRemoteFromPtr(ctx context.Context, rptr *sstore.RemotePtrType, sessionId string, screenId string) (*ResolvedRemote, error) { +func ResolveRemoteFromPtr(ctx context.Context, rptr *sstore.RemotePtrType, sessionId string, screenId string) (*ResolvedRemote, error) { if rptr == nil || rptr.RemoteId == "" { return nil, nil } diff --git a/pkg/remote/remote.go b/pkg/remote/remote.go index 326b8b91..5fcdb409 100644 --- a/pkg/remote/remote.go +++ b/pkg/remote/remote.go @@ -1119,6 +1119,10 @@ func (msh *MShellProc) ReInit(ctx context.Context) (*packet.InitPacketType, erro return initPk, nil } +func (msh *MShellProc) StreamFile(ctx context.Context, streamPk *packet.StreamFilePacketType) (*packet.RpcResponseIter, error) { + return msh.PacketRpcIter(ctx, streamPk) +} + func addScVarsToState(state *packet.ShellState) *packet.ShellState { if state == nil { return nil @@ -1374,6 +1378,13 @@ func (msh *MShellProc) SendSpecialInput(siPk *packet.SpecialInputPacketType) err return msh.ServerProc.Input.SendPacket(siPk) } +func (msh *MShellProc) SendFileData(dataPk *packet.FileDataPacketType) error { + if !msh.IsConnected() { + return fmt.Errorf("remote is not connected, cannot send input") + } + return msh.ServerProc.Input.SendPacket(dataPk) +} + func makeTermOpts(runPk *packet.RunPacketType) sstore.TermOpts { return sstore.TermOpts{Rows: int64(runPk.TermOpts.Rows), Cols: int64(runPk.TermOpts.Cols), FlexRows: true, MaxPtySize: DefaultMaxPtySize} } @@ -1577,9 +1588,25 @@ func (msh *MShellProc) RemoveRunningCmd(ck base.CommandKey) { } } +func (msh *MShellProc) PacketRpcIter(ctx context.Context, pk packet.RpcPacketType) (*packet.RpcResponseIter, error) { + if !msh.IsConnected() { + return nil, fmt.Errorf("remote is not connected") + } + if pk == nil { + return nil, fmt.Errorf("PacketRpc passed nil packet") + } + reqId := pk.GetReqId() + msh.ServerProc.Output.RegisterRpc(reqId) + err := msh.ServerProc.Input.SendPacketCtx(ctx, pk) + if err != nil { + return nil, err + } + return msh.ServerProc.Output.GetResponseIter(reqId), nil +} + func (msh *MShellProc) PacketRpcRaw(ctx context.Context, pk packet.RpcPacketType) (packet.RpcResponsePacketType, error) { if !msh.IsConnected() { - return nil, fmt.Errorf("runner is not connected") + return nil, fmt.Errorf("remote is not connected") } if pk == nil { return nil, fmt.Errorf("PacketRpc passed nil packet") @@ -1812,6 +1839,7 @@ func (msh *MShellProc) ProcessPackets() { go sendScreenUpdates(screens) } }) + // TODO need to clean dataPosMap dataPosMap := make(map[base.CommandKey]int64) for pk := range msh.ServerProc.Output.MainCh { if pk.GetType() == packet.DataPacketStr { diff --git a/pkg/utilfn/utilfn.go b/pkg/utilfn/utilfn.go index 0decc6eb..5b3e70b1 100644 --- a/pkg/utilfn/utilfn.go +++ b/pkg/utilfn/utilfn.go @@ -193,3 +193,16 @@ func Sha1Hash(data []byte) string { hval := base64.StdEncoding.EncodeToString(hvalRaw[:]) return hval } + +func ChunkSlice[T any](s []T, chunkSize int) [][]T { + var rtn [][]T + for len(rtn) > 0 { + if len(s) <= chunkSize { + rtn = append(rtn, s) + break + } + rtn = append(rtn, s[:chunkSize]) + s = s[chunkSize:] + } + return rtn +}