mirror of
https://github.com/wavetermdev/backup.git
synced 2026-08-05 13:57:07 -07:00
PE-41 remote file api (#1)
* new remote file streaming API packets. implemented 'stat' for remote files * introduce filedata packets. allow streaming RPCs. fix RPC bug with combined packet parsers. implement file streaming for filestream RPC. * checkpoint on adding write-file * completely untested write-file impl -- writefilecontext, condition var for signaling new data packets, cleanup goroutine, ready/done states. * better error messages, also unlock MServer before calling done on wfcs * fix bug with perm json tag. change constant name
This commit is contained in:
+1
-1
@@ -159,7 +159,7 @@ func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType
|
||||
}
|
||||
|
||||
func handleSingle(fromServer bool) {
|
||||
packetParser := packet.MakePacketParser(os.Stdin)
|
||||
packetParser := packet.MakePacketParser(os.Stdin, false)
|
||||
sender := packet.MakePacketSender(os.Stdout, nil)
|
||||
defer func() {
|
||||
sender.Close()
|
||||
|
||||
+193
-21
@@ -26,33 +26,42 @@ import (
|
||||
// server : <init, >run, >cmddata, >cmddone, <cmdstart, <>data, <>dataack, <cmddone
|
||||
// >cd, >getcmd, >untailcmd, >input, <resp
|
||||
// all : <>error, <>message, <>ping, <raw
|
||||
//
|
||||
// >streamfile, <streamfileresp, <filedata*
|
||||
// >writefile, <writefileready, >filedata*, <writefiledone
|
||||
|
||||
const MaxCompGenValues = 100
|
||||
|
||||
var GlobalDebug = false
|
||||
|
||||
const (
|
||||
RunPacketStr = "run" // rpc
|
||||
PingPacketStr = "ping"
|
||||
InitPacketStr = "init"
|
||||
DataPacketStr = "data" // command
|
||||
DataAckPacketStr = "dataack" // command
|
||||
CmdStartPacketStr = "cmdstart" // rpc-response
|
||||
CmdDonePacketStr = "cmddone" // command
|
||||
DataEndPacketStr = "dataend"
|
||||
ResponsePacketStr = "resp" // rpc-response
|
||||
DonePacketStr = "done"
|
||||
CmdErrorPacketStr = "cmderror" // command
|
||||
MessagePacketStr = "message"
|
||||
GetCmdPacketStr = "getcmd" // rpc
|
||||
UntailCmdPacketStr = "untailcmd" // rpc
|
||||
CdPacketStr = "cd" // rpc
|
||||
CmdDataPacketStr = "cmddata" // rpc-response
|
||||
RawPacketStr = "raw"
|
||||
SpecialInputPacketStr = "sinput" // command
|
||||
CompGenPacketStr = "compgen" // rpc
|
||||
ReInitPacketStr = "reinit" // rpc
|
||||
CmdFinalPacketStr = "cmdfinal" // command, pushed at the "end" of a command (fail-safe for no cmddone)
|
||||
RunPacketStr = "run" // rpc
|
||||
PingPacketStr = "ping"
|
||||
InitPacketStr = "init"
|
||||
DataPacketStr = "data" // command
|
||||
DataAckPacketStr = "dataack" // command
|
||||
CmdStartPacketStr = "cmdstart" // rpc-response
|
||||
CmdDonePacketStr = "cmddone" // command
|
||||
DataEndPacketStr = "dataend"
|
||||
ResponsePacketStr = "resp" // rpc-response
|
||||
DonePacketStr = "done"
|
||||
CmdErrorPacketStr = "cmderror" // command
|
||||
MessagePacketStr = "message"
|
||||
GetCmdPacketStr = "getcmd" // rpc
|
||||
UntailCmdPacketStr = "untailcmd" // rpc
|
||||
CdPacketStr = "cd" // rpc
|
||||
CmdDataPacketStr = "cmddata" // rpc-response
|
||||
RawPacketStr = "raw"
|
||||
SpecialInputPacketStr = "sinput" // command
|
||||
CompGenPacketStr = "compgen" // rpc
|
||||
ReInitPacketStr = "reinit" // rpc
|
||||
CmdFinalPacketStr = "cmdfinal" // command, pushed at the "end" of a command (fail-safe for no cmddone)
|
||||
StreamFilePacketStr = "streamfile" // rpc
|
||||
StreamFileResponseStr = "streamfileresp" // rpc-response
|
||||
WriteFilePacketStr = "writefile" // rpc
|
||||
WriteFileReadyPacketStr = "writefileready" // rpc-response
|
||||
WriteFileDonePacketStr = "writefiledone" // rpc-response
|
||||
FileDataPacketStr = "filedata"
|
||||
|
||||
OpenAIPacketStr = "openai" // other
|
||||
)
|
||||
@@ -84,7 +93,13 @@ func init() {
|
||||
TypeStrToFactory[CompGenPacketStr] = reflect.TypeOf(CompGenPacketType{})
|
||||
TypeStrToFactory[ReInitPacketStr] = reflect.TypeOf(ReInitPacketType{})
|
||||
TypeStrToFactory[CmdFinalPacketStr] = reflect.TypeOf(CmdFinalPacketType{})
|
||||
TypeStrToFactory[StreamFilePacketStr] = reflect.TypeOf(StreamFilePacketType{})
|
||||
TypeStrToFactory[StreamFileResponseStr] = reflect.TypeOf(StreamFileResponseType{})
|
||||
TypeStrToFactory[OpenAIPacketStr] = reflect.TypeOf(OpenAIPacketType{})
|
||||
TypeStrToFactory[FileDataPacketStr] = reflect.TypeOf(FileDataPacketType{})
|
||||
TypeStrToFactory[WriteFilePacketStr] = reflect.TypeOf(WriteFilePacketType{})
|
||||
TypeStrToFactory[WriteFileReadyPacketStr] = reflect.TypeOf(WriteFileReadyPacketType{})
|
||||
TypeStrToFactory[WriteFileDonePacketStr] = reflect.TypeOf(WriteFileDonePacketType{})
|
||||
|
||||
var _ RpcPacketType = (*RunPacketType)(nil)
|
||||
var _ RpcPacketType = (*GetCmdPacketType)(nil)
|
||||
@@ -92,10 +107,16 @@ func init() {
|
||||
var _ RpcPacketType = (*CdPacketType)(nil)
|
||||
var _ RpcPacketType = (*CompGenPacketType)(nil)
|
||||
var _ RpcPacketType = (*ReInitPacketType)(nil)
|
||||
var _ RpcPacketType = (*StreamFilePacketType)(nil)
|
||||
var _ RpcPacketType = (*WriteFilePacketType)(nil)
|
||||
|
||||
var _ RpcResponsePacketType = (*CmdStartPacketType)(nil)
|
||||
var _ RpcResponsePacketType = (*ResponsePacketType)(nil)
|
||||
var _ RpcResponsePacketType = (*CmdDataPacketType)(nil)
|
||||
var _ RpcResponsePacketType = (*StreamFileResponseType)(nil)
|
||||
var _ RpcResponsePacketType = (*FileDataPacketType)(nil)
|
||||
var _ RpcResponsePacketType = (*WriteFileReadyPacketType)(nil)
|
||||
var _ RpcResponsePacketType = (*WriteFileDonePacketType)(nil)
|
||||
|
||||
var _ CommandPacketType = (*DataPacketType)(nil)
|
||||
var _ CommandPacketType = (*DataAckPacketType)(nil)
|
||||
@@ -159,6 +180,33 @@ func MakePingPacket() *PingPacketType {
|
||||
return &PingPacketType{Type: PingPacketStr}
|
||||
}
|
||||
|
||||
type FileDataPacketType struct {
|
||||
Type string `json:"type"`
|
||||
RespId string `json:"respid"`
|
||||
Data []byte `json:"data"`
|
||||
Eof bool `json:"eof,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func (*FileDataPacketType) GetType() string {
|
||||
return FileDataPacketStr
|
||||
}
|
||||
|
||||
func MakeFileDataPacket(reqId string) *FileDataPacketType {
|
||||
return &FileDataPacketType{
|
||||
Type: FileDataPacketStr,
|
||||
RespId: reqId,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *FileDataPacketType) GetResponseId() string {
|
||||
return p.RespId
|
||||
}
|
||||
|
||||
func (p *FileDataPacketType) GetResponseDone() bool {
|
||||
return p.Eof || p.Error != ""
|
||||
}
|
||||
|
||||
type DataPacketType struct {
|
||||
Type string `json:"type"`
|
||||
CK base.CommandKey `json:"ck"`
|
||||
@@ -348,6 +396,61 @@ func MakeReInitPacket() *ReInitPacketType {
|
||||
return &ReInitPacketType{Type: ReInitPacketStr}
|
||||
}
|
||||
|
||||
type StreamFilePacketType struct {
|
||||
Type string `json:"type"`
|
||||
ReqId string `json:"reqid"`
|
||||
Path string `json:"path"`
|
||||
ByteRange []int64 `json:"byterange"` // works like the http "Range" header (multiple ranges are not allowed)
|
||||
StatOnly bool `json:"statonly,omitempty"` // set if you just want the stat response (no data returned)
|
||||
}
|
||||
|
||||
func (*StreamFilePacketType) GetType() string {
|
||||
return StreamFilePacketStr
|
||||
}
|
||||
|
||||
func (p *StreamFilePacketType) GetReqId() string {
|
||||
return p.ReqId
|
||||
}
|
||||
|
||||
func MakeStreamFilePacket() *StreamFilePacketType {
|
||||
return &StreamFilePacketType{Type: StreamFilePacketStr}
|
||||
}
|
||||
|
||||
type FileInfo struct {
|
||||
Name string `json:"name"`
|
||||
Size int64 `json:"size"`
|
||||
ModTs int64 `json:"modts"`
|
||||
IsDir bool `json:"isdir,omitempty"`
|
||||
Perm int `json:"perm"`
|
||||
}
|
||||
|
||||
type StreamFileResponseType struct {
|
||||
Type string `json:"type"`
|
||||
RespId string `json:"respid"`
|
||||
Done bool `json:"done,omitempty"`
|
||||
Info *FileInfo `json:"info,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func (*StreamFileResponseType) GetType() string {
|
||||
return StreamFileResponseStr
|
||||
}
|
||||
|
||||
func (p *StreamFileResponseType) GetResponseId() string {
|
||||
return p.RespId
|
||||
}
|
||||
|
||||
func (p *StreamFileResponseType) GetResponseDone() bool {
|
||||
return p.Done
|
||||
}
|
||||
|
||||
func MakeStreamFileResponse(respId string) *StreamFileResponseType {
|
||||
return &StreamFileResponseType{
|
||||
Type: StreamFileResponseStr,
|
||||
RespId: respId,
|
||||
}
|
||||
}
|
||||
|
||||
type CompGenPacketType struct {
|
||||
Type string `json:"type"`
|
||||
ReqId string `json:"reqid"`
|
||||
@@ -668,6 +771,75 @@ func MakeCmdErrorPacket(ck base.CommandKey, err error) *CmdErrorPacketType {
|
||||
return &CmdErrorPacketType{Type: CmdErrorPacketStr, CK: ck, Error: err.Error()}
|
||||
}
|
||||
|
||||
type WriteFilePacketType struct {
|
||||
Type string `json:"type"`
|
||||
ReqId string `json:"reqid"`
|
||||
UseTemp bool `json:"usetemp,omitempty"`
|
||||
Path string `json:"path"`
|
||||
}
|
||||
|
||||
func (*WriteFilePacketType) GetType() string {
|
||||
return WriteFilePacketStr
|
||||
}
|
||||
|
||||
func (p *WriteFilePacketType) GetReqId() string {
|
||||
return p.ReqId
|
||||
}
|
||||
|
||||
func MakeWriteFilePacket() *WriteFilePacketType {
|
||||
return &WriteFilePacketType{Type: WriteFilePacketStr}
|
||||
}
|
||||
|
||||
type WriteFileReadyPacketType struct {
|
||||
Type string `json:"type"`
|
||||
RespId string `json:"reqid"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func (*WriteFileReadyPacketType) GetType() string {
|
||||
return WriteFileReadyPacketStr
|
||||
}
|
||||
|
||||
func (p *WriteFileReadyPacketType) GetResponseId() string {
|
||||
return p.RespId
|
||||
}
|
||||
|
||||
func (p *WriteFileReadyPacketType) GetResponseDone() bool {
|
||||
return p.Error != ""
|
||||
}
|
||||
|
||||
func MakeWriteFileReadyPacket(reqId string) *WriteFileReadyPacketType {
|
||||
return &WriteFileReadyPacketType{
|
||||
Type: WriteFileReadyPacketStr,
|
||||
RespId: reqId,
|
||||
}
|
||||
}
|
||||
|
||||
type WriteFileDonePacketType struct {
|
||||
Type string `json:"type"`
|
||||
RespId string `json:"reqid"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func (*WriteFileDonePacketType) GetType() string {
|
||||
return WriteFileDonePacketStr
|
||||
}
|
||||
|
||||
func (p *WriteFileDonePacketType) GetResponseId() string {
|
||||
return p.RespId
|
||||
}
|
||||
|
||||
func (p *WriteFileDonePacketType) GetResponseDone() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func MakeWriteFileDonePacket(reqId string) *WriteFileDonePacketType {
|
||||
return &WriteFileDonePacketType{
|
||||
Type: WriteFileDonePacketStr,
|
||||
RespId: reqId,
|
||||
}
|
||||
}
|
||||
|
||||
type PacketType interface {
|
||||
GetType() string
|
||||
}
|
||||
|
||||
+64
-25
@@ -16,10 +16,11 @@ import (
|
||||
)
|
||||
|
||||
type PacketParser struct {
|
||||
Lock *sync.Mutex
|
||||
MainCh chan PacketType
|
||||
RpcMap map[string]*RpcEntry
|
||||
Err error
|
||||
Lock *sync.Mutex
|
||||
MainCh chan PacketType
|
||||
RpcMap map[string]*RpcEntry
|
||||
RpcHandler bool
|
||||
Err error
|
||||
}
|
||||
|
||||
type RpcEntry struct {
|
||||
@@ -27,20 +28,37 @@ type RpcEntry struct {
|
||||
RespCh chan RpcResponsePacketType
|
||||
}
|
||||
|
||||
func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser {
|
||||
type RpcResponseIter struct {
|
||||
ReqId string
|
||||
Parser *PacketParser
|
||||
}
|
||||
|
||||
func (iter *RpcResponseIter) Next(ctx context.Context) (RpcResponsePacketType, error) {
|
||||
// will unregister the rpc on ResponseDone
|
||||
return iter.Parser.GetNextResponse(ctx, iter.ReqId)
|
||||
}
|
||||
|
||||
func (iter *RpcResponseIter) Close() {
|
||||
iter.Parser.UnRegisterRpc(iter.ReqId)
|
||||
}
|
||||
|
||||
func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser, rpcHandler bool) *PacketParser {
|
||||
rtnParser := &PacketParser{
|
||||
Lock: &sync.Mutex{},
|
||||
MainCh: make(chan PacketType),
|
||||
RpcMap: make(map[string]*RpcEntry),
|
||||
Lock: &sync.Mutex{},
|
||||
MainCh: make(chan PacketType),
|
||||
RpcMap: make(map[string]*RpcEntry),
|
||||
RpcHandler: rpcHandler,
|
||||
}
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for pk := range p1.MainCh {
|
||||
sent := rtnParser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
if rtnParser.RpcHandler {
|
||||
sent := rtnParser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
}
|
||||
}
|
||||
rtnParser.MainCh <- pk
|
||||
}
|
||||
@@ -48,9 +66,11 @@ func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for pk := range p2.MainCh {
|
||||
sent := rtnParser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
if rtnParser.RpcHandler {
|
||||
sent := rtnParser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
}
|
||||
}
|
||||
rtnParser.MainCh <- pk
|
||||
}
|
||||
@@ -77,6 +97,26 @@ func (p *PacketParser) WaitForResponse(ctx context.Context, reqId string) RpcRes
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PacketParser) GetResponseIter(reqId string) *RpcResponseIter {
|
||||
return &RpcResponseIter{Parser: p, ReqId: reqId}
|
||||
}
|
||||
|
||||
func (p *PacketParser) GetNextResponse(ctx context.Context, reqId string) (RpcResponsePacketType, error) {
|
||||
entry := p.getRpcEntry(reqId)
|
||||
if entry == nil {
|
||||
return nil, nil
|
||||
}
|
||||
select {
|
||||
case resp := <-entry.RespCh:
|
||||
if resp.GetResponseDone() {
|
||||
p.UnRegisterRpc(reqId)
|
||||
}
|
||||
return resp, nil
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PacketParser) UnRegisterRpc(reqId string) {
|
||||
p.Lock.Lock()
|
||||
defer p.Lock.Unlock()
|
||||
@@ -123,10 +163,6 @@ func (p *PacketParser) trySendRpcResponse(pk PacketType) bool {
|
||||
case entry.RespCh <- respPk:
|
||||
default:
|
||||
}
|
||||
if respPk.GetResponseDone() {
|
||||
delete(p.RpcMap, respPk.GetResponseId())
|
||||
close(entry.RespCh)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -144,11 +180,12 @@ func (p *PacketParser) SetErr(err error) {
|
||||
}
|
||||
}
|
||||
|
||||
func MakePacketParser(input io.Reader) *PacketParser {
|
||||
func MakePacketParser(input io.Reader, rpcHandler bool) *PacketParser {
|
||||
parser := &PacketParser{
|
||||
Lock: &sync.Mutex{},
|
||||
MainCh: make(chan PacketType),
|
||||
RpcMap: make(map[string]*RpcEntry),
|
||||
Lock: &sync.Mutex{},
|
||||
MainCh: make(chan PacketType),
|
||||
RpcMap: make(map[string]*RpcEntry),
|
||||
RpcHandler: rpcHandler,
|
||||
}
|
||||
bufReader := bufio.NewReader(input)
|
||||
go func() {
|
||||
@@ -194,9 +231,11 @@ func MakePacketParser(input io.Reader) *PacketParser {
|
||||
if pk.GetType() == PingPacketStr {
|
||||
continue
|
||||
}
|
||||
sent := parser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
if parser.RpcHandler {
|
||||
sent := parser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
}
|
||||
}
|
||||
parser.MainCh <- pk
|
||||
}
|
||||
|
||||
+347
-16
@@ -8,9 +8,12 @@ package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -22,22 +25,109 @@ import (
|
||||
"github.com/commandlinedev/apishell/pkg/shexec"
|
||||
)
|
||||
|
||||
const MaxFileDataPacketSize = 16 * 1024
|
||||
const WriteFileContextTimeout = 30 * time.Second
|
||||
const cleanLoopTime = 5 * time.Second
|
||||
const MaxWriteFileContextData = 100
|
||||
|
||||
// TODO create unblockable packet-sender (backed by an array) for clientproc
|
||||
type MServer struct {
|
||||
Lock *sync.Mutex
|
||||
MainInput *packet.PacketParser
|
||||
Sender *packet.PacketSender
|
||||
ClientMap map[base.CommandKey]*shexec.ClientProc
|
||||
Debug bool
|
||||
StateMap map[string]*packet.ShellState // sha1->state
|
||||
CurrentState string // sha1
|
||||
WriteErrorCh chan bool // closed if there is a I/O write error
|
||||
WriteErrorChOnce *sync.Once
|
||||
Lock *sync.Mutex
|
||||
MainInput *packet.PacketParser
|
||||
Sender *packet.PacketSender
|
||||
ClientMap map[base.CommandKey]*shexec.ClientProc
|
||||
Debug bool
|
||||
StateMap map[string]*packet.ShellState // sha1->state
|
||||
CurrentState string // sha1
|
||||
WriteErrorCh chan bool // closed if there is a I/O write error
|
||||
WriteErrorChOnce *sync.Once
|
||||
WriteFileContextMap map[string]*WriteFileContext
|
||||
Done bool
|
||||
}
|
||||
|
||||
type WriteFileContext struct {
|
||||
CVar *sync.Cond
|
||||
Data []*packet.FileDataPacketType
|
||||
LastActive time.Time
|
||||
Err error
|
||||
Done bool
|
||||
}
|
||||
|
||||
func (m *MServer) Close() {
|
||||
m.Sender.Close()
|
||||
m.Sender.WaitForDone()
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
m.Done = true
|
||||
}
|
||||
|
||||
func (m *MServer) checkDone() bool {
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
return m.Done
|
||||
}
|
||||
|
||||
func (m *MServer) getWriteFileContext(reqId string) *WriteFileContext {
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
wfc := m.WriteFileContextMap[reqId]
|
||||
if wfc == nil {
|
||||
wfc = &WriteFileContext{
|
||||
CVar: sync.NewCond(&sync.Mutex{}),
|
||||
LastActive: time.Now(),
|
||||
}
|
||||
m.WriteFileContextMap[reqId] = wfc
|
||||
}
|
||||
return wfc
|
||||
}
|
||||
|
||||
func (m *MServer) addFileDataPacket(pk *packet.FileDataPacketType) {
|
||||
m.Lock.Lock()
|
||||
wfc := m.WriteFileContextMap[pk.RespId]
|
||||
m.Lock.Unlock()
|
||||
if wfc == nil {
|
||||
return
|
||||
}
|
||||
wfc.CVar.L.Lock()
|
||||
defer wfc.CVar.L.Unlock()
|
||||
if wfc.Done || wfc.Err != nil {
|
||||
return
|
||||
}
|
||||
if len(wfc.Data) > MaxWriteFileContextData {
|
||||
wfc.Err = errors.New("write-file buffer length exceeded")
|
||||
wfc.Data = nil
|
||||
wfc.CVar.Broadcast()
|
||||
return
|
||||
}
|
||||
wfc.LastActive = time.Now()
|
||||
wfc.Data = append(wfc.Data, pk)
|
||||
wfc.CVar.Signal()
|
||||
}
|
||||
|
||||
func (wfc *WriteFileContext) setDone() {
|
||||
wfc.CVar.L.Lock()
|
||||
defer wfc.CVar.L.Unlock()
|
||||
wfc.Done = true
|
||||
wfc.Data = nil
|
||||
wfc.CVar.Broadcast()
|
||||
}
|
||||
|
||||
func (m *MServer) cleanWriteFileContexts() {
|
||||
now := time.Now()
|
||||
var staleWfcs []*WriteFileContext
|
||||
m.Lock.Lock()
|
||||
for reqId, wfc := range m.WriteFileContextMap {
|
||||
if now.Sub(wfc.LastActive) > WriteFileContextTimeout {
|
||||
staleWfcs = append(staleWfcs, wfc)
|
||||
delete(m.WriteFileContextMap, reqId)
|
||||
}
|
||||
}
|
||||
m.Lock.Unlock()
|
||||
|
||||
// we do this outside of m.Lock just in case there is some lock contention (end of WriteFile could theoretically be slow)
|
||||
for _, wfc := range staleWfcs {
|
||||
wfc.setDone()
|
||||
}
|
||||
}
|
||||
|
||||
func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) {
|
||||
@@ -164,6 +254,224 @@ func (m *MServer) reinit(reqId string) {
|
||||
m.Sender.SendPacket(initPk)
|
||||
}
|
||||
|
||||
func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContext) {
|
||||
defer wfc.setDone()
|
||||
if pk.Path == "" {
|
||||
resp := packet.MakeWriteFileReadyPacket(pk.ReqId)
|
||||
resp.Error = "invalid write-file request, no path specified"
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
finfo, err := os.Stat(pk.Path)
|
||||
if err == nil && finfo.IsDir() {
|
||||
err = fmt.Errorf("invalid path, cannot write a directory")
|
||||
}
|
||||
if err == nil {
|
||||
writePerm := (finfo.Mode().Perm() & 0o222)
|
||||
if writePerm == 0 {
|
||||
err = fmt.Errorf("file is not writable, perms: %v", finfo.Mode().Perm())
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
resp := packet.MakeWriteFileReadyPacket(pk.ReqId)
|
||||
resp.Error = err.Error()
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
|
||||
var writeFd *os.File
|
||||
if pk.UseTemp {
|
||||
dirName := filepath.Dir(pk.Path)
|
||||
dirFInfo, err := os.Stat(dirName)
|
||||
if err == nil {
|
||||
writePerm := (dirFInfo.Mode().Perm() & 0o222)
|
||||
if writePerm == 0 {
|
||||
err = fmt.Errorf("file-write tempmode is set, but parent directory is not writeable, perms: %v", dirFInfo.Mode().Perm())
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
resp := packet.MakeWriteFileReadyPacket(pk.ReqId)
|
||||
resp.Error = err.Error()
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
baseName := filepath.Base(pk.Path)
|
||||
writeFd, err = os.CreateTemp(dirName, baseName+".tmp.")
|
||||
if err != nil {
|
||||
resp := packet.MakeWriteFileReadyPacket(pk.ReqId)
|
||||
resp.Error = fmt.Sprintf("write-file could not open tempfile: %v", err)
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
writeFd, err = os.OpenFile(pk.Path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o777) // use 777 because OpenFile respects umask
|
||||
if err != nil {
|
||||
resp := packet.MakeWriteFileReadyPacket(pk.ReqId)
|
||||
resp.Error = fmt.Sprintf("write-file could not open file: %v", err)
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// ok, so now writeFd is valid, send the "ready" response
|
||||
resp := packet.MakeWriteFileReadyPacket(pk.ReqId)
|
||||
m.Sender.SendPacket(resp)
|
||||
|
||||
// now we wait for data (cond var)
|
||||
// this Unlock() runs first (because it is a later defer) so we can still run wfc.setDone() safely
|
||||
wfc.CVar.L.Lock()
|
||||
defer wfc.CVar.L.Unlock()
|
||||
var doneErr error
|
||||
for {
|
||||
if wfc.Done {
|
||||
break
|
||||
}
|
||||
if wfc.Err != nil {
|
||||
doneErr = wfc.Err
|
||||
break
|
||||
}
|
||||
if len(wfc.Data) == 0 {
|
||||
wfc.CVar.Wait()
|
||||
continue
|
||||
}
|
||||
dataPk := wfc.Data[0]
|
||||
wfc.Data = wfc.Data[1:]
|
||||
if dataPk.Error != "" {
|
||||
doneErr = fmt.Errorf("error received from client: %v", errors.New(dataPk.Error))
|
||||
break
|
||||
}
|
||||
if len(dataPk.Data) > 0 {
|
||||
_, err := writeFd.Write(dataPk.Data)
|
||||
if err != nil {
|
||||
doneErr = fmt.Errorf("error writing data to file: %v", err)
|
||||
break
|
||||
}
|
||||
}
|
||||
if dataPk.Eof {
|
||||
break
|
||||
}
|
||||
}
|
||||
closeErr := writeFd.Close()
|
||||
if doneErr == nil && closeErr != nil {
|
||||
doneErr = fmt.Errorf("error closing file: %v", closeErr)
|
||||
}
|
||||
if pk.UseTemp {
|
||||
if doneErr != nil {
|
||||
os.Remove(writeFd.Name())
|
||||
} else {
|
||||
renameErr := os.Rename(writeFd.Name(), pk.Path)
|
||||
if renameErr != nil {
|
||||
doneErr = fmt.Errorf("error renaming temp file: %v", renameErr)
|
||||
// rename failed, try to remove temp file still
|
||||
os.Remove(writeFd.Name())
|
||||
}
|
||||
}
|
||||
}
|
||||
donePk := packet.MakeWriteFileDonePacket(pk.ReqId)
|
||||
if doneErr != nil {
|
||||
donePk.Error = doneErr.Error()
|
||||
}
|
||||
m.Sender.SendPacket(donePk)
|
||||
}
|
||||
|
||||
func (m *MServer) streamFile(pk *packet.StreamFilePacketType) {
|
||||
resp := packet.MakeStreamFileResponse(pk.ReqId)
|
||||
finfo, err := os.Stat(pk.Path)
|
||||
if err != nil {
|
||||
resp.Error = fmt.Sprintf("cannot stat file %q: %v", pk.Path, err)
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
resp.Info = &packet.FileInfo{
|
||||
Name: pk.Path,
|
||||
Size: finfo.Size(),
|
||||
ModTs: finfo.ModTime().UnixMilli(),
|
||||
IsDir: finfo.IsDir(),
|
||||
Perm: int(finfo.Mode().Perm()),
|
||||
}
|
||||
if pk.StatOnly {
|
||||
resp.Done = true
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
// like the http Range header. range header is end inclusive. for us, endByte is non-inclusive (so we add 1)
|
||||
var startByte, endByte int64
|
||||
if len(pk.ByteRange) == 0 {
|
||||
endByte = finfo.Size()
|
||||
} else if len(pk.ByteRange) == 1 && pk.ByteRange[0] >= 0 {
|
||||
startByte = pk.ByteRange[0]
|
||||
endByte = finfo.Size()
|
||||
} else if len(pk.ByteRange) == 1 && pk.ByteRange[0] < 0 {
|
||||
startByte = finfo.Size() + pk.ByteRange[0] // "+" since ByteRange[0] is less than 0
|
||||
endByte = finfo.Size()
|
||||
} else if len(pk.ByteRange) == 2 {
|
||||
startByte = pk.ByteRange[0]
|
||||
endByte = pk.ByteRange[1] + 1
|
||||
} else {
|
||||
resp.Error = fmt.Sprintf("invalid byte range (%d entries)", len(pk.ByteRange))
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
if startByte < 0 {
|
||||
startByte = 0
|
||||
}
|
||||
if endByte > finfo.Size() {
|
||||
endByte = finfo.Size()
|
||||
}
|
||||
if startByte >= endByte {
|
||||
resp.Done = true
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
fd, err := os.Open(pk.Path)
|
||||
if err != nil {
|
||||
resp.Error = fmt.Sprintf("opening file: %v", err)
|
||||
m.Sender.SendPacket(resp)
|
||||
return
|
||||
}
|
||||
defer fd.Close()
|
||||
m.Sender.SendPacket(resp)
|
||||
var buffer [MaxFileDataPacketSize]byte
|
||||
var sentDone bool
|
||||
first := true
|
||||
for ; startByte < endByte; startByte += MaxFileDataPacketSize {
|
||||
if !first {
|
||||
// throttle packet sending @ 1000 packets/s, or 16M/s
|
||||
time.Sleep(1 * time.Millisecond)
|
||||
}
|
||||
first = false
|
||||
readLen := int64Min(MaxFileDataPacketSize, endByte-startByte)
|
||||
bufSlice := buffer[0:readLen]
|
||||
nr, err := fd.ReadAt(bufSlice, startByte)
|
||||
dataPk := packet.MakeFileDataPacket(pk.ReqId)
|
||||
dataPk.Data = make([]byte, nr)
|
||||
copy(dataPk.Data, bufSlice)
|
||||
if err == io.EOF {
|
||||
dataPk.Eof = true
|
||||
} else if err != nil {
|
||||
dataPk.Error = err.Error()
|
||||
}
|
||||
m.Sender.SendPacket(dataPk)
|
||||
if dataPk.GetResponseDone() {
|
||||
sentDone = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !sentDone {
|
||||
dataPk := packet.MakeFileDataPacket(pk.ReqId)
|
||||
dataPk.Eof = true
|
||||
m.Sender.SendPacket(dataPk)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func int64Min(v1 int64, v2 int64) int64 {
|
||||
if v1 < v2 {
|
||||
return v1
|
||||
}
|
||||
return v2
|
||||
}
|
||||
|
||||
func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) {
|
||||
reqId := pk.GetReqId()
|
||||
if cdPk, ok := pk.(*packet.CdPacketType); ok {
|
||||
@@ -183,6 +491,15 @@ func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) {
|
||||
go m.reinit(reqId)
|
||||
return
|
||||
}
|
||||
if streamPk, ok := pk.(*packet.StreamFilePacketType); ok {
|
||||
go m.streamFile(streamPk)
|
||||
return
|
||||
}
|
||||
if writePk, ok := pk.(*packet.WriteFilePacketType); ok {
|
||||
wfc := m.getWriteFileContext(writePk.ReqId)
|
||||
go m.writeFile(writePk, wfc)
|
||||
return
|
||||
}
|
||||
m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid rpc type '%s'", pk.GetType()))
|
||||
return
|
||||
}
|
||||
@@ -288,6 +605,10 @@ func (server *MServer) runReadLoop() {
|
||||
server.ProcessRpcPacket(rpcPk)
|
||||
continue
|
||||
}
|
||||
if fileDataPk, ok := pk.(*packet.FileDataPacketType); ok {
|
||||
server.addFileDataPacket(fileDataPk)
|
||||
continue
|
||||
}
|
||||
server.Sender.SendMessageFmt("invalid packet '%s' sent to mshell server", packet.AsString(pk))
|
||||
continue
|
||||
}
|
||||
@@ -299,17 +620,27 @@ func RunServer() (int, error) {
|
||||
debug = true
|
||||
}
|
||||
server := &MServer{
|
||||
Lock: &sync.Mutex{},
|
||||
ClientMap: make(map[base.CommandKey]*shexec.ClientProc),
|
||||
StateMap: make(map[string]*packet.ShellState),
|
||||
Debug: debug,
|
||||
WriteErrorCh: make(chan bool),
|
||||
WriteErrorChOnce: &sync.Once{},
|
||||
Lock: &sync.Mutex{},
|
||||
ClientMap: make(map[base.CommandKey]*shexec.ClientProc),
|
||||
StateMap: make(map[string]*packet.ShellState),
|
||||
Debug: debug,
|
||||
WriteErrorCh: make(chan bool),
|
||||
WriteErrorChOnce: &sync.Once{},
|
||||
WriteFileContextMap: make(map[string]*WriteFileContext),
|
||||
}
|
||||
go func() {
|
||||
for {
|
||||
if server.checkDone() {
|
||||
return
|
||||
}
|
||||
time.Sleep(cleanLoopTime)
|
||||
server.cleanWriteFileContexts()
|
||||
}
|
||||
}()
|
||||
if debug {
|
||||
packet.GlobalDebug = true
|
||||
}
|
||||
server.MainInput = packet.MakePacketParser(os.Stdin)
|
||||
server.MainInput = packet.MakePacketParser(os.Stdin, false)
|
||||
server.Sender = packet.MakePacketSender(os.Stdout, server.packetSenderErrorHandler)
|
||||
defer server.Close()
|
||||
var err error
|
||||
|
||||
@@ -47,9 +47,9 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.I
|
||||
return nil, nil, fmt.Errorf("running local client: %w", err)
|
||||
}
|
||||
sender := packet.MakePacketSender(inputWriter, nil)
|
||||
stdoutPacketParser := packet.MakePacketParser(stdoutReader)
|
||||
stderrPacketParser := packet.MakePacketParser(stderrReader)
|
||||
packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser)
|
||||
stdoutPacketParser := packet.MakePacketParser(stdoutReader, false)
|
||||
stderrPacketParser := packet.MakePacketParser(stderrReader, false)
|
||||
packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser, true)
|
||||
cproc := &ClientProc{
|
||||
Cmd: ecmd,
|
||||
StartTs: startTs,
|
||||
|
||||
@@ -727,7 +727,7 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, mshe
|
||||
if mshellStream != nil {
|
||||
sendMShellBinary(inputWriter, mshellStream)
|
||||
}
|
||||
packetParser := packet.MakePacketParser(stdoutReader)
|
||||
packetParser := packet.MakePacketParser(stdoutReader, false)
|
||||
err = ecmd.Start()
|
||||
if err != nil {
|
||||
return fmt.Errorf("running ssh command: %w", err)
|
||||
@@ -860,9 +860,9 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon
|
||||
return nil, fmt.Errorf("running ssh command: %w", err)
|
||||
}
|
||||
defer cmd.Close()
|
||||
stdoutPacketParser := packet.MakePacketParser(stdoutReader)
|
||||
stderrPacketParser := packet.MakePacketParser(stderrReader)
|
||||
packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser)
|
||||
stdoutPacketParser := packet.MakePacketParser(stdoutReader, false)
|
||||
stderrPacketParser := packet.MakePacketParser(stderrReader, false)
|
||||
packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser, false)
|
||||
sender := packet.MakePacketSender(inputWriter, nil)
|
||||
versionOk := false
|
||||
for pk := range packetParser.MainCh {
|
||||
|
||||
Reference in New Issue
Block a user