diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 4a33fee1..f409ec9d 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -79,12 +79,14 @@ func init() { TypeStrToFactory[DataAckPacketStr] = reflect.TypeOf(DataAckPacketType{}) TypeStrToFactory[DataEndPacketStr] = reflect.TypeOf(DataEndPacketType{}) TypeStrToFactory[CompGenPacketStr] = reflect.TypeOf(CompGenPacketType{}) + TypeStrToFactory[ReInitPacketStr] = reflect.TypeOf(ReInitPacketType{}) var _ RpcPacketType = (*RunPacketType)(nil) var _ RpcPacketType = (*GetCmdPacketType)(nil) var _ RpcPacketType = (*UntailCmdPacketType)(nil) var _ RpcPacketType = (*CdPacketType)(nil) var _ RpcPacketType = (*CompGenPacketType)(nil) + var _ RpcPacketType = (*ReInitPacketType)(nil) var _ RpcResponsePacketType = (*CmdStartPacketType)(nil) var _ RpcResponsePacketType = (*ResponsePacketType)(nil) @@ -332,6 +334,23 @@ func MakeCdPacket() *CdPacketType { return &CdPacketType{Type: CdPacketStr} } +type ReInitPacketType struct { + Type string `json:"type"` + ReqId string `json:"reqid"` +} + +func (*ReInitPacketType) GetType() string { + return ReInitPacketStr +} + +func (p *ReInitPacketType) GetReqId() string { + return p.ReqId +} + +func MakeReInitPacket() *ReInitPacketType { + return &ReInitPacketType{Type: ReInitPacketStr} +} + type CompGenPacketType struct { Type string `json:"type"` ReqId string `json:"reqid"` diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index 2e4199c2..84c8b396 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -169,6 +169,7 @@ func MakePacketParser(input io.Reader) *PacketParser { } // ##[len][json]\n // ##14{"hello":true}\n + // ##N{...} bracePos := strings.Index(line, "{") if !strings.HasPrefix(line, "##") || bracePos == -1 { parser.MainCh <- MakeRawPacket(line[:len(line)-1]) diff --git a/pkg/server/server.go b/pkg/server/server.go index d0cc680f..61e649e2 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -137,6 +137,16 @@ func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { return } +func (m *MServer) reinit(reqId string) { + initPk, err := shexec.MakeServerInitPacket() + if err != nil { + m.Sender.SendErrorResponse(reqId, fmt.Errorf("error creating init packet: %w", err)) + return + } + initPk.RespId = reqId + m.Sender.SendPacket(initPk) +} + func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) { reqId := pk.GetReqId() if cdPk, ok := pk.(*packet.CdPacketType); ok { @@ -152,6 +162,10 @@ func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) { go m.runCompGen(compPk) return } + if _, ok := pk.(*packet.ReInitPacketType); ok { + go m.reinit(reqId) + return + } m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid rpc type '%s'", pk.GetType())) return } diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index ecee19d7..25a5369e 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -101,6 +101,7 @@ var NoStoreVarNames = map[string]bool{ "HISTSIZE": true, "HISTTIMEFORMAT": true, "SRANDOM": true, + "COLUMNS": true, // we want these in our remote state object // "EUID": true,