From eeaeac8dc80e14bb866ff1f84716064be5d44fab Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 10 Jun 2022 00:35:24 -0700 Subject: [PATCH 001/149] initial runner commit --- LICENSE | 373 +++++++++++++++++++++++++++++++++++++++++++ NOTICE.md | 5 + go.mod | 8 + go.sum | 4 + main-runner.go | 60 +++++++ pkg/base/base.go | 156 ++++++++++++++++++ pkg/packet/packet.go | 187 ++++++++++++++++++++++ pkg/shexec/shexec.go | 226 ++++++++++++++++++++++++++ 8 files changed, 1019 insertions(+) create mode 100644 LICENSE create mode 100644 NOTICE.md create mode 100644 go.mod create mode 100644 go.sum create mode 100644 main-runner.go create mode 100644 pkg/base/base.go create mode 100644 pkg/packet/packet.go create mode 100644 pkg/shexec/shexec.go diff --git a/LICENSE b/LICENSE new file mode 100644 index 00000000..ee6256cd --- /dev/null +++ b/LICENSE @@ -0,0 +1,373 @@ +Mozilla Public License Version 2.0 +================================== + +1. Definitions +-------------- + +1.1. "Contributor" + means each individual or legal entity that creates, contributes to + the creation of, or owns Covered Software. + +1.2. "Contributor Version" + means the combination of the Contributions of others (if any) used + by a Contributor and that particular Contributor's Contribution. + +1.3. "Contribution" + means Covered Software of a particular Contributor. + +1.4. "Covered Software" + means Source Code Form to which the initial Contributor has attached + the notice in Exhibit A, the Executable Form of such Source Code + Form, and Modifications of such Source Code Form, in each case + including portions thereof. + +1.5. "Incompatible With Secondary Licenses" + means + + (a) that the initial Contributor has attached the notice described + in Exhibit B to the Covered Software; or + + (b) that the Covered Software was made available under the terms of + version 1.1 or earlier of the License, but not also under the + terms of a Secondary License. + +1.6. "Executable Form" + means any form of the work other than Source Code Form. + +1.7. "Larger Work" + means a work that combines Covered Software with other material, in + a separate file or files, that is not Covered Software. + +1.8. "License" + means this document. + +1.9. "Licensable" + means having the right to grant, to the maximum extent possible, + whether at the time of the initial grant or subsequently, any and + all of the rights conveyed by this License. + +1.10. "Modifications" + means any of the following: + + (a) any file in Source Code Form that results from an addition to, + deletion from, or modification of the contents of Covered + Software; or + + (b) any new file in Source Code Form that contains any Covered + Software. + +1.11. "Patent Claims" of a Contributor + means any patent claim(s), including without limitation, method, + process, and apparatus claims, in any patent Licensable by such + Contributor that would be infringed, but for the grant of the + License, by the making, using, selling, offering for sale, having + made, import, or transfer of either its Contributions or its + Contributor Version. + +1.12. "Secondary License" + means either the GNU General Public License, Version 2.0, the GNU + Lesser General Public License, Version 2.1, the GNU Affero General + Public License, Version 3.0, or any later versions of those + licenses. + +1.13. "Source Code Form" + means the form of the work preferred for making modifications. + +1.14. "You" (or "Your") + means an individual or a legal entity exercising rights under this + License. For legal entities, "You" includes any entity that + controls, is controlled by, or is under common control with You. For + purposes of this definition, "control" means (a) the power, direct + or indirect, to cause the direction or management of such entity, + whether by contract or otherwise, or (b) ownership of more than + fifty percent (50%) of the outstanding shares or beneficial + ownership of such entity. + +2. License Grants and Conditions +-------------------------------- + +2.1. Grants + +Each Contributor hereby grants You a world-wide, royalty-free, +non-exclusive license: + +(a) under intellectual property rights (other than patent or trademark) + Licensable by such Contributor to use, reproduce, make available, + modify, display, perform, distribute, and otherwise exploit its + Contributions, either on an unmodified basis, with Modifications, or + as part of a Larger Work; and + +(b) under Patent Claims of such Contributor to make, use, sell, offer + for sale, have made, import, and otherwise transfer either its + Contributions or its Contributor Version. + +2.2. Effective Date + +The licenses granted in Section 2.1 with respect to any Contribution +become effective for each Contribution on the date the Contributor first +distributes such Contribution. + +2.3. Limitations on Grant Scope + +The licenses granted in this Section 2 are the only rights granted under +this License. No additional rights or licenses will be implied from the +distribution or licensing of Covered Software under this License. +Notwithstanding Section 2.1(b) above, no patent license is granted by a +Contributor: + +(a) for any code that a Contributor has removed from Covered Software; + or + +(b) for infringements caused by: (i) Your and any other third party's + modifications of Covered Software, or (ii) the combination of its + Contributions with other software (except as part of its Contributor + Version); or + +(c) under Patent Claims infringed by Covered Software in the absence of + its Contributions. + +This License does not grant any rights in the trademarks, service marks, +or logos of any Contributor (except as may be necessary to comply with +the notice requirements in Section 3.4). + +2.4. Subsequent Licenses + +No Contributor makes additional grants as a result of Your choice to +distribute the Covered Software under a subsequent version of this +License (see Section 10.2) or under the terms of a Secondary License (if +permitted under the terms of Section 3.3). + +2.5. Representation + +Each Contributor represents that the Contributor believes its +Contributions are its original creation(s) or it has sufficient rights +to grant the rights to its Contributions conveyed by this License. + +2.6. Fair Use + +This License is not intended to limit any rights You have under +applicable copyright doctrines of fair use, fair dealing, or other +equivalents. + +2.7. Conditions + +Sections 3.1, 3.2, 3.3, and 3.4 are conditions of the licenses granted +in Section 2.1. + +3. Responsibilities +------------------- + +3.1. Distribution of Source Form + +All distribution of Covered Software in Source Code Form, including any +Modifications that You create or to which You contribute, must be under +the terms of this License. You must inform recipients that the Source +Code Form of the Covered Software is governed by the terms of this +License, and how they can obtain a copy of this License. You may not +attempt to alter or restrict the recipients' rights in the Source Code +Form. + +3.2. Distribution of Executable Form + +If You distribute Covered Software in Executable Form then: + +(a) such Covered Software must also be made available in Source Code + Form, as described in Section 3.1, and You must inform recipients of + the Executable Form how they can obtain a copy of such Source Code + Form by reasonable means in a timely manner, at a charge no more + than the cost of distribution to the recipient; and + +(b) You may distribute such Executable Form under the terms of this + License, or sublicense it under different terms, provided that the + license for the Executable Form does not attempt to limit or alter + the recipients' rights in the Source Code Form under this License. + +3.3. Distribution of a Larger Work + +You may create and distribute a Larger Work under terms of Your choice, +provided that You also comply with the requirements of this License for +the Covered Software. If the Larger Work is a combination of Covered +Software with a work governed by one or more Secondary Licenses, and the +Covered Software is not Incompatible With Secondary Licenses, this +License permits You to additionally distribute such Covered Software +under the terms of such Secondary License(s), so that the recipient of +the Larger Work may, at their option, further distribute the Covered +Software under the terms of either this License or such Secondary +License(s). + +3.4. Notices + +You may not remove or alter the substance of any license notices +(including copyright notices, patent notices, disclaimers of warranty, +or limitations of liability) contained within the Source Code Form of +the Covered Software, except that You may alter any license notices to +the extent required to remedy known factual inaccuracies. + +3.5. Application of Additional Terms + +You may choose to offer, and to charge a fee for, warranty, support, +indemnity or liability obligations to one or more recipients of Covered +Software. However, You may do so only on Your own behalf, and not on +behalf of any Contributor. You must make it absolutely clear that any +such warranty, support, indemnity, or liability obligation is offered by +You alone, and You hereby agree to indemnify every Contributor for any +liability incurred by such Contributor as a result of warranty, support, +indemnity or liability terms You offer. You may include additional +disclaimers of warranty and limitations of liability specific to any +jurisdiction. + +4. Inability to Comply Due to Statute or Regulation +--------------------------------------------------- + +If it is impossible for You to comply with any of the terms of this +License with respect to some or all of the Covered Software due to +statute, judicial order, or regulation then You must: (a) comply with +the terms of this License to the maximum extent possible; and (b) +describe the limitations and the code they affect. Such description must +be placed in a text file included with all distributions of the Covered +Software under this License. Except to the extent prohibited by statute +or regulation, such description must be sufficiently detailed for a +recipient of ordinary skill to be able to understand it. + +5. Termination +-------------- + +5.1. The rights granted under this License will terminate automatically +if You fail to comply with any of its terms. However, if You become +compliant, then the rights granted under this License from a particular +Contributor are reinstated (a) provisionally, unless and until such +Contributor explicitly and finally terminates Your grants, and (b) on an +ongoing basis, if such Contributor fails to notify You of the +non-compliance by some reasonable means prior to 60 days after You have +come back into compliance. Moreover, Your grants from a particular +Contributor are reinstated on an ongoing basis if such Contributor +notifies You of the non-compliance by some reasonable means, this is the +first time You have received notice of non-compliance with this License +from such Contributor, and You become compliant prior to 30 days after +Your receipt of the notice. + +5.2. If You initiate litigation against any entity by asserting a patent +infringement claim (excluding declaratory judgment actions, +counter-claims, and cross-claims) alleging that a Contributor Version +directly or indirectly infringes any patent, then the rights granted to +You by any and all Contributors for the Covered Software under Section +2.1 of this License shall terminate. + +5.3. In the event of termination under Sections 5.1 or 5.2 above, all +end user license agreements (excluding distributors and resellers) which +have been validly granted by You or Your distributors under this License +prior to termination shall survive termination. + +************************************************************************ +* * +* 6. Disclaimer of Warranty * +* ------------------------- * +* * +* Covered Software is provided under this License on an "as is" * +* basis, without warranty of any kind, either expressed, implied, or * +* statutory, including, without limitation, warranties that the * +* Covered Software is free of defects, merchantable, fit for a * +* particular purpose or non-infringing. The entire risk as to the * +* quality and performance of the Covered Software is with You. * +* Should any Covered Software prove defective in any respect, You * +* (not any Contributor) assume the cost of any necessary servicing, * +* repair, or correction. This disclaimer of warranty constitutes an * +* essential part of this License. No use of any Covered Software is * +* authorized under this License except under this disclaimer. * +* * +************************************************************************ + +************************************************************************ +* * +* 7. Limitation of Liability * +* -------------------------- * +* * +* Under no circumstances and under no legal theory, whether tort * +* (including negligence), contract, or otherwise, shall any * +* Contributor, or anyone who distributes Covered Software as * +* permitted above, be liable to You for any direct, indirect, * +* special, incidental, or consequential damages of any character * +* including, without limitation, damages for lost profits, loss of * +* goodwill, work stoppage, computer failure or malfunction, or any * +* and all other commercial damages or losses, even if such party * +* shall have been informed of the possibility of such damages. This * +* limitation of liability shall not apply to liability for death or * +* personal injury resulting from such party's negligence to the * +* extent applicable law prohibits such limitation. Some * +* jurisdictions do not allow the exclusion or limitation of * +* incidental or consequential damages, so this exclusion and * +* limitation may not apply to You. * +* * +************************************************************************ + +8. Litigation +------------- + +Any litigation relating to this License may be brought only in the +courts of a jurisdiction where the defendant maintains its principal +place of business and such litigation shall be governed by laws of that +jurisdiction, without reference to its conflict-of-law provisions. +Nothing in this Section shall prevent a party's ability to bring +cross-claims or counter-claims. + +9. Miscellaneous +---------------- + +This License represents the complete agreement concerning the subject +matter hereof. If any provision of this License is held to be +unenforceable, such provision shall be reformed only to the extent +necessary to make it enforceable. Any law or regulation which provides +that the language of a contract shall be construed against the drafter +shall not be used to construe this License against a Contributor. + +10. Versions of the License +--------------------------- + +10.1. New Versions + +Mozilla Foundation is the license steward. Except as provided in Section +10.3, no one other than the license steward has the right to modify or +publish new versions of this License. Each version will be given a +distinguishing version number. + +10.2. Effect of New Versions + +You may distribute the Covered Software under the terms of the version +of the License under which You originally received the Covered Software, +or under the terms of any subsequent version published by the license +steward. + +10.3. Modified Versions + +If you create software not governed by this License, and you want to +create a new license for such software, you may create and use a +modified version of this License if you rename the license and remove +any references to the name of the license steward (except to note that +such modified license differs from this License). + +10.4. Distributing Source Code Form that is Incompatible With Secondary +Licenses + +If You choose to distribute Source Code Form that is Incompatible With +Secondary Licenses under the terms of this version of the License, the +notice described in Exhibit B of this License must be attached. + +Exhibit A - Source Code Form License Notice +------------------------------------------- + + This Source Code Form is subject to the terms of the Mozilla Public + License, v. 2.0. If a copy of the MPL was not distributed with this + file, You can obtain one at https://mozilla.org/MPL/2.0/. + +If it is not possible or desirable to put the notice in a particular +file, then You may include the notice in a location (such as a LICENSE +file in a relevant directory) where a recipient would be likely to look +for such a notice. + +You may add additional accurate notices of copyright ownership. + +Exhibit B - "Incompatible With Secondary Licenses" Notice +--------------------------------------------------------- + + This Source Code Form is "Incompatible With Secondary Licenses", as + defined by the Mozilla Public License, v. 2.0. diff --git a/NOTICE.md b/NOTICE.md new file mode 100644 index 00000000..5a9ef9f4 --- /dev/null +++ b/NOTICE.md @@ -0,0 +1,5 @@ +Copyright (c) 2021-2022 Dashborg Inc + +This Source Code Form is subject to the terms of the Mozilla Public +License, v. 2.0. If a copy of the MPL was not distributed with this +file, You can obtain one at https://mozilla.org/MPL/2.0/. diff --git a/go.mod b/go.mod new file mode 100644 index 00000000..a794172c --- /dev/null +++ b/go.mod @@ -0,0 +1,8 @@ +module github.com/scripthaus-dev/sh2-runner + +go 1.17 + +require ( + github.com/creack/pty v1.1.18 // indirect + github.com/google/uuid v1.3.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 00000000..12d5eadc --- /dev/null +++ b/go.sum @@ -0,0 +1,4 @@ +github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= +github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= +github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= +github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= diff --git a/main-runner.go b/main-runner.go new file mode 100644 index 00000000..04fb9f4a --- /dev/null +++ b/main-runner.go @@ -0,0 +1,60 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package main + +import ( + "fmt" + "os" + "os/signal" + "syscall" + + "github.com/scripthaus-dev/sh2-runner/pkg/packet" + "github.com/scripthaus-dev/sh2-runner/pkg/shexec" +) + +func setupSignals(cmd *shexec.ShExecType) { + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + go func() { + for sig := range sigCh { + cmd.Cmd.Process.Signal(sig) + } + }() +} + +func main() { + packetCh := packet.PacketParser(os.Stdin) + var runPacket *packet.RunPacketType + for pk := range packetCh { + if pk.GetType() == packet.PingPacketStr { + continue + } + if pk.GetType() == packet.RunPacketStr { + runPacket, _ = pk.(*packet.RunPacketType) + break + } + if pk.GetType() == packet.ErrorPacketStr { + packet.SendPacket(os.Stdout, pk) + return + } + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + return + } + if runPacket == nil { + packet.SendErrorPacket(os.Stdout, "did not receive a 'run' packet") + return + } + cmd, err := shexec.RunCommand(runPacket) + if err != nil { + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("error running command: %v", err)) + return + } + setupSignals(cmd) + packet.SendPacket(os.Stdout, packet.MakeOkCmdPacket(fmt.Sprintf("running command %s/%s", runPacket.SessionId, runPacket.CmdId), runPacket.CmdId, cmd.Cmd.Process.Pid)) + cmd.WaitForCommand() + packet.SendPacket(os.Stdout, packet.MakeDonePacket()) +} diff --git a/pkg/base/base.go b/pkg/base/base.go new file mode 100644 index 00000000..ae6774ce --- /dev/null +++ b/pkg/base/base.go @@ -0,0 +1,156 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package base + +import ( + "errors" + "fmt" + "io/fs" + "os" + "path" + "path/filepath" +) + +const ScRunnerVarName = "SCRIPTHAUS_RUNNER" +const ScHomeVarName = "SCRIPTHAUS_HOME" +const HomeVarName = "HOME" +const ScShell = "bash" +const SessionsDirBaseName = ".sessions" +const RunnerBaseName = "runner" +const SessionDBName = "session.db" +const ScReadyString = "scripthaus runner ready" + +const OSCEscError = "error" + +type CommandFileNames struct { + PtyOutFile string + StdinFifo string + DoneFile string +} + +func GetScHomeDir() (string, error) { + scHome := os.Getenv(ScHomeVarName) + if scHome == "" { + homeVar := os.Getenv(HomeVarName) + if homeVar == "" { + return "", fmt.Errorf("Cannot resolve scripthaus home directory (SCRIPTHAUS_HOME and HOME not set)") + } + scHome = path.Join(homeVar, "scripthaus") + } + return scHome, nil +} + +func GetCommandFileNames(sessionId string, cmdId string) (*CommandFileNames, error) { + if sessionId == "" || cmdId == "" { + return nil, fmt.Errorf("cannot get command-files when sessionid or cmdid is empty") + } + sdir, err := EnsureSessionDir(sessionId) + if err != nil { + return nil, err + } + base := path.Join(sdir, cmdId) + return &CommandFileNames{ + PtyOutFile: base + ".ptyout", + StdinFifo: base + ".stdin", + DoneFile: base + ".done", + }, nil +} + +func CleanUpCmdFiles(sessionId string, cmdId string) error { + if cmdId == "" { + return fmt.Errorf("bad cmdid, cannot clean up") + } + sdir, err := EnsureSessionDir(sessionId) + if err != nil { + return err + } + cmdFileGlob := path.Join(sdir, cmdId+".*") + matches, err := filepath.Glob(cmdFileGlob) + if err != nil { + return err + } + for _, file := range matches { + rmErr := os.Remove(file) + if err == nil && rmErr != nil { + err = rmErr + } + } + return err +} + +func EnsureSessionDir(sessionId string) (string, error) { + if sessionId == "" { + return "", fmt.Errorf("Bad sessionid, cannot be empty") + } + shhome, err := GetScHomeDir() + if err != nil { + return "", err + } + sdir := path.Join(shhome, ".sessions", sessionId) + info, err := os.Stat(sdir) + if errors.Is(err, fs.ErrNotExist) { + err = os.MkdirAll(sdir, 0777) + if err != nil { + return "", err + } + info, err = os.Stat(sdir) + } + if err != nil { + return "", err + } + if !info.IsDir() { + return "", fmt.Errorf("session dir '%s' must be a directory", sdir) + } + return sdir, nil +} + +func GetScRunnerPath() string { + runnerPath := os.Getenv(ScRunnerVarName) + if runnerPath != "" { + return runnerPath + } + scHome, err := GetScHomeDir() + if err != nil { + panic(err) + } + return path.Join(scHome, RunnerBaseName) +} + +func GetScSessionsDir() string { + scHome, err := GetScHomeDir() + if err != nil { + panic(err) + } + return path.Join(scHome, SessionsDirBaseName) +} + +func GetSessionDBName(sessionId string) string { + scHome, err := GetScHomeDir() + if err != nil { + panic(err) + } + return path.Join(scHome, SessionDBName) +} + +// SH OSC Escapes (code 198, S=19, H=8) +// \e]198;cmdid;(cmd-id)BEL - return command-id to server +// \e]198;remote;0BEL - runner program not available +// \e]198;remote;1BEL - runner program is available +// \e]198;error;(error-str)BEL - communicate an internal error +func MakeSHOSCEsc(escName string, data string) string { + return fmt.Sprintf("\033]198;%s;%s\007", escName, data) +} + +func WriteErrorMsg(fileName string, errVal string) error { + fd, err := os.OpenFile(fileName, os.O_APPEND|os.O_WRONLY, 0600) + if err != nil { + return err + } + oscEsc := MakeSHOSCEsc(OSCEscError, errVal) + _, writeErr := fd.Write([]byte(oscEsc)) + return writeErr +} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go new file mode 100644 index 00000000..52e3a0da --- /dev/null +++ b/pkg/packet/packet.go @@ -0,0 +1,187 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package packet + +import ( + "bufio" + "encoding/json" + "fmt" + "io" +) + +const RunPacketStr = "run" +const PingPacketStr = "ping" +const DonePacketStr = "done" +const ErrorPacketStr = "error" +const OkCmdPacketStr = "okcmd" + +type PingPacketType struct { + Type string `json:"type"` +} + +func (*PingPacketType) GetType() string { + return PingPacketStr +} + +func MakePingPacket() *PingPacketType { + return &PingPacketType{Type: PingPacketStr} +} + +type DonePacketType struct { + Type string `json:"type"` +} + +func (*DonePacketType) GetType() string { + return DonePacketStr +} + +func MakeDonePacket() *DonePacketType { + return &DonePacketType{Type: DonePacketStr} +} + +type OkCmdPacketType struct { + Type string `json:"type"` + Message string `json:"message"` + CmdId string `json:"cmdid"` + Pid int `json:"pid"` +} + +func (*OkCmdPacketType) GetType() string { + return OkCmdPacketStr +} + +func MakeOkCmdPacket(message string, cmdId string, pid int) *OkCmdPacketType { + return &OkCmdPacketType{Type: OkCmdPacketStr, Message: message, CmdId: cmdId, Pid: pid} +} + +type RunPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + ChDir string `json:"chdir"` + Env map[string]string `json:"env"` + Command string `json:"command"` +} + +func (ct *RunPacketType) GetType() string { + return RunPacketStr +} + +type BarePacketType struct { + Type string `json:"type"` +} + +type ErrorPacketType struct { + Type string `json:"type"` + Error string `json:"error"` +} + +func (et *ErrorPacketType) GetType() string { + return ErrorPacketStr +} + +func MakeErrorPacket(errorStr string) *ErrorPacketType { + return &ErrorPacketType{Type: ErrorPacketStr, Error: errorStr} +} + +type PacketType interface { + GetType() string +} + +func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { + var bareCmd BarePacketType + err := json.Unmarshal(jsonBuf, &bareCmd) + if err != nil { + return nil, err + } + if bareCmd.Type == "" { + return nil, fmt.Errorf("received packet with no type") + } + if bareCmd.Type == RunPacketStr { + var runPacket RunPacketType + err = json.Unmarshal(jsonBuf, &runPacket) + if err != nil { + return nil, err + } + return &runPacket, nil + } + if bareCmd.Type == PingPacketStr { + return MakePingPacket(), nil + } + if bareCmd.Type == DonePacketStr { + return MakeDonePacket(), nil + } + if bareCmd.Type == ErrorPacketStr { + var errorPacket ErrorPacketType + err = json.Unmarshal(jsonBuf, &errorPacket) + if err != nil { + return nil, err + } + return &errorPacket, nil + } + if bareCmd.Type == OkCmdPacketStr { + var okPacket OkCmdPacketType + err = json.Unmarshal(jsonBuf, &okPacket) + if err != nil { + return nil, err + } + return &okPacket, nil + } + return nil, fmt.Errorf("invalid packet-type '%s'", bareCmd.Type) +} + +func SendPacket(w io.Writer, packet PacketType) error { + if packet == nil { + return nil + } + barr, err := json.Marshal(packet) + if err != nil { + return fmt.Errorf("marshaling '%s' packet: %w", packet.GetType(), err) + } + barr = append(barr, '\n') + _, err = w.Write(barr) + if err != nil { + return err + } + return nil +} + +func SendErrorPacket(w io.Writer, errorStr string) error { + return SendPacket(w, MakeErrorPacket(errorStr)) +} + +func PacketParser(input io.Reader) chan PacketType { + bufReader := bufio.NewReader(input) + rtnCh := make(chan PacketType) + go func() { + defer func() { + close(rtnCh) + }() + for { + line, err := bufReader.ReadString('\n') + if err == io.EOF { + return + } + if err != nil { + errPacket := MakeErrorPacket(fmt.Sprintf("reading packets from input: %v", err)) + rtnCh <- errPacket + return + } + pk, err := ParseJsonPacket([]byte(line)) + if err != nil { + errPk := MakeErrorPacket(fmt.Sprintf("parsing packet json from input: %v", err)) + rtnCh <- errPk + return + } + if pk.GetType() == DonePacketStr { + return + } + rtnCh <- pk + } + }() + return rtnCh +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go new file mode 100644 index 00000000..4ef5f660 --- /dev/null +++ b/pkg/shexec/shexec.go @@ -0,0 +1,226 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package shexec + +import ( + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "os" + "os/exec" + "strings" + "syscall" + "time" + + "github.com/creack/pty" + "github.com/google/uuid" + "github.com/scripthaus-dev/sh2-runner/pkg/base" + "github.com/scripthaus-dev/sh2-runner/pkg/packet" +) + +type DoneData struct { + DurationMs int64 `json:"durationms"` + ExitCode int `json:"exitcode"` +} + +type ShExecType struct { + FileNames *base.CommandFileNames + Cmd *exec.Cmd + CmdPty *os.File + StartTs time.Time +} + +func (c *ShExecType) Close() { + c.CmdPty.Close() +} + +func getEnvStrKey(envStr string) string { + eqIdx := strings.Index(envStr, "=") + if eqIdx == -1 { + return envStr + } + return envStr[0:eqIdx] +} + +func UpdateCmdEnv(cmd *exec.Cmd, envVars map[string]string) { + if len(envVars) == 0 { + return + } + if cmd.Env != nil { + cmd.Env = os.Environ() + } + found := make(map[string]bool) + var newEnv []string + for _, envStr := range cmd.Env { + envKey := getEnvStrKey(envStr) + newEnvVal, ok := envVars[envKey] + if ok { + if newEnvVal == "" { + continue + } + newEnv = append(newEnv, envKey+"="+newEnvVal) + found[envKey] = true + } else { + newEnv = append(newEnv, envStr) + } + } + for envKey, envVal := range envVars { + if found[envKey] { + continue + } + newEnv = append(newEnv, envKey+"="+envVal) + } + cmd.Env = newEnv +} + +func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { + ecmd := exec.Command("bash", "-c", pk.Command) + UpdateCmdEnv(ecmd, pk.Env) + if pk.ChDir != "" { + ecmd.Dir = pk.ChDir + } + ecmd.Stdin = cmdTty + ecmd.Stdout = cmdTty + ecmd.Stderr = cmdTty + ecmd.SysProcAttr = &syscall.SysProcAttr{ + Setsid: true, + Setctty: true, + } + return ecmd +} + +// this will never return (unless there is an error creating/opening the file), as fifoFile will never EOF +func MakeAndCopyStdinFifo(dst *os.File, fifoName string) error { + os.Remove(fifoName) + err := syscall.Mkfifo(fifoName, 0600) // only read/write from user for security + if err != nil { + return fmt.Errorf("cannot make stdin-fifo '%s': %v", fifoName, err) + } + // rw is non-blocking, will keep the fifo "open" for the blocking reader + rwfd, err := os.OpenFile(fifoName, os.O_RDWR, 0600) + if err != nil { + return fmt.Errorf("cannot open stdin-fifo(1) '%s': %v", fifoName, err) + } + defer rwfd.Close() + fifoReader, err := os.Open(fifoName) // blocking open/reads (open won't block because of rwfd) + if err != nil { + return fmt.Errorf("cannot open stdin-fifo(2) '%s': %w", fifoName, err) + } + defer fifoReader.Close() + io.Copy(dst, fifoReader) + return nil +} + +func ValidateRunPacket(pk *packet.RunPacketType) error { + if pk.Type != packet.RunPacketStr { + return fmt.Errorf("run packet has wrong type: %s", pk.Type) + } + if pk.SessionId == "" { + return fmt.Errorf("run packet does not have sessionid") + } + _, err := uuid.Parse(pk.SessionId) + if err != nil { + return fmt.Errorf("invalid sessionid '%s' for command", pk.SessionId) + } + if pk.CmdId == "" { + return fmt.Errorf("run packet does not have cmdid") + } + _, err = uuid.Parse(pk.CmdId) + if err != nil { + return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) + } + if pk.ChDir != "" { + dirInfo, err := os.Stat(pk.ChDir) + if err != nil { + return fmt.Errorf("invalid cwd '%s' for command: %v", pk.ChDir, err) + } + if !dirInfo.IsDir() { + return fmt.Errorf("invalid cwd '%s' for command, not a directory", pk.ChDir) + } + } + return nil +} + +// returning nil error means the process has successfully been kicked-off +func RunCommand(pk *packet.RunPacketType) (*ShExecType, error) { + if pk.CmdId == "" { + pk.CmdId = uuid.New().String() + } + err := ValidateRunPacket(pk) + if err != nil { + return nil, err + } + fileNames, err := base.GetCommandFileNames(pk.SessionId, pk.CmdId) + if err != nil { + return nil, err + } + if _, err = os.Stat(fileNames.PtyOutFile); !errors.Is(err, fs.ErrNotExist) { + return nil, fmt.Errorf("cmdid '%s' was already used", pk.CmdId) + } + cmdPty, cmdTty, err := pty.Open() + if err != nil { + return nil, fmt.Errorf("opening new pty: %w", err) + } + defer func() { + cmdTty.Close() + }() + startTs := time.Now() + ecmd := MakeExecCmd(pk, cmdTty) + err = ecmd.Start() + if err != nil { + return nil, fmt.Errorf("starting command: %w", err) + } + ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY, 0600) + if err != nil { + return nil, fmt.Errorf("cannot open ptyout file '%s': %w", fileNames.PtyOutFile, err) + } + go func() { + // copy pty output to .ptyout file + _, copyErr := io.Copy(ptyOutFd, cmdPty) + if copyErr != nil { + base.WriteErrorMsg(fileNames.PtyOutFile, fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) + } + }() + go func() { + // copy .stdin fifo contents to pty input + copyFifoErr := MakeAndCopyStdinFifo(cmdPty, fileNames.StdinFifo) + if copyFifoErr != nil { + base.WriteErrorMsg(fileNames.PtyOutFile, fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) + } + }() + return &ShExecType{ + FileNames: fileNames, + Cmd: ecmd, + CmdPty: cmdPty, + StartTs: startTs, + }, nil +} + +func (c *ShExecType) WaitForCommand() { + err := c.Cmd.Wait() + cmdDuration := time.Since(c.StartTs) + exitCode := 0 + if err != nil { + exitErr, ok := err.(*exec.ExitError) + if ok { + exitCode = exitErr.ExitCode() + } + } + doneData := DoneData{ + DurationMs: int64(cmdDuration / time.Millisecond), + ExitCode: exitCode, + } + doneDataBytes, _ := json.Marshal(doneData) + doneDataBytes = append(doneDataBytes, '\n') + err = os.WriteFile(c.FileNames.DoneFile, doneDataBytes, 0600) + if err != nil { + base.WriteErrorMsg(c.FileNames.PtyOutFile, fmt.Sprintf("reading from stdin fifo: %v", err)) + } + return +} From 1a3886c437736a995ae99773aef89d992c907543 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 10 Jun 2022 21:37:21 -0700 Subject: [PATCH 002/149] go runner / runner-single fork flow working --- main-runner.go | 151 ++++++++++++++++++++++++++++++++++++++----- pkg/base/base.go | 52 ++++++++++----- pkg/packet/packet.go | 145 ++++++++++++++++++++++++++++++++++++----- pkg/shexec/shexec.go | 46 ++++++------- 4 files changed, 321 insertions(+), 73 deletions(-) diff --git a/main-runner.go b/main-runner.go index 04fb9f4a..19b0ba3c 100644 --- a/main-runner.go +++ b/main-runner.go @@ -11,23 +11,30 @@ import ( "os" "os/signal" "syscall" + "time" + "github.com/google/uuid" + "github.com/scripthaus-dev/sh2-runner/pkg/base" "github.com/scripthaus-dev/sh2-runner/pkg/packet" "github.com/scripthaus-dev/sh2-runner/pkg/shexec" ) -func setupSignals(cmd *shexec.ShExecType) { +// in single run mode, we don't want the runner to die from signals +// since we want the single runner to persist even if session / main runner +// is terminated. +func setupSingleSignals(cmd *shexec.ShExecType) { sigCh := make(chan os.Signal, 1) - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) go func() { - for sig := range sigCh { - cmd.Cmd.Process.Signal(sig) + for range sigCh { + // do nothing } }() } -func main() { +func doSingle(cmdId string) { packetCh := packet.PacketParser(os.Stdin) + sender := packet.MakePacketSender(os.Stdout) var runPacket *packet.RunPacketType for pk := range packetCh { if pk.GetType() == packet.PingPacketStr { @@ -37,24 +44,134 @@ func main() { runPacket, _ = pk.(*packet.RunPacketType) break } - if pk.GetType() == packet.ErrorPacketStr { - packet.SendPacket(os.Stdout, pk) - return - } - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) return } if runPacket == nil { - packet.SendErrorPacket(os.Stdout, "did not receive a 'run' packet") + sender.SendErrorPacket("did not receive a 'run' packet") return } - cmd, err := shexec.RunCommand(runPacket) + if runPacket.CmdId == "" { + runPacket.CmdId = cmdId + } + if runPacket.CmdId != cmdId { + sender.SendErrorPacket(fmt.Sprintf("run packet cmdid[%s] did not match arg[%s]", runPacket.CmdId, cmdId)) + return + } + cmd, err := shexec.RunCommand(runPacket, sender) if err != nil { - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("error running command: %v", err)) + sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) return } - setupSignals(cmd) - packet.SendPacket(os.Stdout, packet.MakeOkCmdPacket(fmt.Sprintf("running command %s/%s", runPacket.SessionId, runPacket.CmdId), runPacket.CmdId, cmd.Cmd.Process.Pid)) - cmd.WaitForCommand() - packet.SendPacket(os.Stdout, packet.MakeDonePacket()) + setupSingleSignals(cmd) + startPacket := packet.MakeCmdStartPacket() + startPacket.Ts = time.Now().UnixMilli() + startPacket.CmdId = runPacket.CmdId + startPacket.Pid = cmd.Cmd.Process.Pid + startPacket.RunnerPid = os.Getpid() + sender.SendPacket(startPacket) + donePacket := cmd.WaitForCommand(runPacket.CmdId) + sender.SendPacket(donePacket) + sender.CloseSendCh() + sender.WaitForDone() +} + +func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { + if pk.CmdId == "" { + pk.CmdId = uuid.New().String() + } + err := shexec.ValidateRunPacket(pk) + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("invalid run packet: %v", err))) + return + } + fileNames, err := base.GetCommandFileNames(pk.SessionId, pk.CmdId) + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot get command file names: %v", err))) + return + } + cmd, err := shexec.MakeRunnerExec(pk.CmdId) + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot make runner command: %v", err))) + return + } + cmdStdin, err := cmd.StdinPipe() + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot pipe stdin to command: %v", err))) + return + } + runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err))) + return + } + cmd.Stdout = runnerOutFd + cmd.Stderr = runnerOutFd + err = cmd.Start() + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("error starting command: %v", err))) + return + } + go func() { + err = packet.SendPacket(cmdStdin, pk) + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("error sending forked runner command: %v", err))) + return + } + cmdStdin.Close() + + // clean up zombies + cmd.Wait() + }() +} + +func doMain() { + homeDir, err := base.GetScHomeDir() + if err != nil { + packet.SendErrorPacket(os.Stdout, err.Error()) + return + } + err = os.Chdir(homeDir) + if err != nil { + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to scripthaus home '%s': %v", homeDir, err)) + return + } + err = base.EnsureRunnerPath() + if err != nil { + packet.SendErrorPacket(os.Stdout, err.Error()) + return + } + packetCh := packet.PacketParser(os.Stdin) + sender := packet.MakePacketSender(os.Stdout) + sender.SendPacket(packet.MakeMessagePacket(fmt.Sprintf("starting scripthaus runner @ %s", homeDir))) + for pk := range packetCh { + if pk.GetType() == packet.PingPacketStr { + continue + } + if pk.GetType() == packet.RunPacketStr { + doMainRun(pk.(*packet.RunPacketType), sender) + continue + } + if pk.GetType() == packet.ErrorPacketStr { + errPk := pk.(*packet.ErrorPacketType) + errPk.Error = "invalid packet sent to runner: " + errPk.Error + sender.SendPacket(errPk) + continue + } + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + } +} + +func main() { + if len(os.Args) >= 2 { + cmdId, err := uuid.Parse(os.Args[1]) + if err != nil { + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid non-cmdid passed to runner", err)) + return + } + doSingle(cmdId.String()) + return + } else { + doMain() + } } diff --git a/pkg/base/base.go b/pkg/base/base.go index ae6774ce..fd2d64ed 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -27,9 +27,9 @@ const ScReadyString = "scripthaus runner ready" const OSCEscError = "error" type CommandFileNames struct { - PtyOutFile string - StdinFifo string - DoneFile string + PtyOutFile string + StdinFifo string + RunnerOutFile string } func GetScHomeDir() (string, error) { @@ -54,9 +54,9 @@ func GetCommandFileNames(sessionId string, cmdId string) (*CommandFileNames, err } base := path.Join(sdir, cmdId) return &CommandFileNames{ - PtyOutFile: base + ".ptyout", - StdinFifo: base + ".stdin", - DoneFile: base + ".done", + PtyOutFile: base + ".ptyout", + StdinFifo: base + ".stdin", + RunnerOutFile: base + ".runout", }, nil } @@ -108,32 +108,50 @@ func EnsureSessionDir(sessionId string) (string, error) { return sdir, nil } -func GetScRunnerPath() string { +func GetScRunnerPath() (string, error) { runnerPath := os.Getenv(ScRunnerVarName) if runnerPath != "" { - return runnerPath + return runnerPath, nil } scHome, err := GetScHomeDir() if err != nil { - panic(err) + return "", err } - return path.Join(scHome, RunnerBaseName) + return path.Join(scHome, RunnerBaseName), nil } -func GetScSessionsDir() string { - scHome, err := GetScHomeDir() +func EnsureRunnerPath() error { + runnerPath, err := GetScRunnerPath() if err != nil { - panic(err) + return err } - return path.Join(scHome, SessionsDirBaseName) + info, err := os.Stat(runnerPath) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return fmt.Errorf("cannot find scripthaus runner at path '%s'", runnerPath) + } + return fmt.Errorf("error stating scripthaus runner at path '%s'", runnerPath) + } + if info.Mode()&0100 == 0 { + return fmt.Errorf("scripthaus runner at path '%s' is not executable mode=%#o", runnerPath, info.Mode()) + } + return nil } -func GetSessionDBName(sessionId string) string { +func GetScSessionsDir() (string, error) { scHome, err := GetScHomeDir() if err != nil { - panic(err) + return "", err } - return path.Join(scHome, SessionDBName) + return path.Join(scHome, SessionsDirBaseName), nil +} + +func GetSessionDBName(sessionId string) (string, error) { + scHome, err := GetScHomeDir() + if err != nil { + return "", err + } + return path.Join(scHome, SessionDBName), nil } // SH OSC Escapes (code 198, S=19, H=8) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 52e3a0da..62328467 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -11,13 +11,16 @@ import ( "encoding/json" "fmt" "io" + "sync" ) const RunPacketStr = "run" const PingPacketStr = "ping" const DonePacketStr = "done" const ErrorPacketStr = "error" -const OkCmdPacketStr = "okcmd" +const MessagePacketStr = "message" +const CmdStartPacketStr = "cmdstart" +const CmdDonePacketStr = "cmddone" type PingPacketType struct { Type string `json:"type"` @@ -31,6 +34,19 @@ func MakePingPacket() *PingPacketType { return &PingPacketType{Type: PingPacketStr} } +type MessagePacketType struct { + Type string `json:"type"` + Message string `json:"message"` +} + +func (*MessagePacketType) GetType() string { + return MessagePacketStr +} + +func MakeMessagePacket(message string) *MessagePacketType { + return &MessagePacketType{Type: MessagePacketStr, Message: message} +} + type DonePacketType struct { Type string `json:"type"` } @@ -43,27 +59,44 @@ func MakeDonePacket() *DonePacketType { return &DonePacketType{Type: DonePacketStr} } -type OkCmdPacketType struct { - Type string `json:"type"` - Message string `json:"message"` - CmdId string `json:"cmdid"` - Pid int `json:"pid"` +type CmdDonePacketType struct { + Type string `json:"type"` + Ts int64 `json:"ts"` + CmdId string `json:"cmdid"` + ExitCode int `json:"exitcode"` + DurationMs int64 `json:"durationms"` } -func (*OkCmdPacketType) GetType() string { - return OkCmdPacketStr +func (*CmdDonePacketType) GetType() string { + return CmdDonePacketStr } -func MakeOkCmdPacket(message string, cmdId string, pid int) *OkCmdPacketType { - return &OkCmdPacketType{Type: OkCmdPacketStr, Message: message, CmdId: cmdId, Pid: pid} +func MakeCmdDonePacket() *CmdDonePacketType { + return &CmdDonePacketType{Type: CmdDonePacketStr} +} + +type CmdStartPacketType struct { + Type string `json:"type"` + Ts int64 `json:"ts"` + CmdId string `json:"cmdid"` + Pid int `json:"pid"` + RunnerPid int `json:"runnerpid"` +} + +func (*CmdStartPacketType) GetType() string { + return CmdStartPacketStr +} + +func MakeCmdStartPacket() *CmdStartPacketType { + return &CmdStartPacketType{Type: CmdStartPacketStr} } type RunPacketType struct { Type string `json:"type"` SessionId string `json:"sessionid"` CmdId string `json:"cmdid"` - ChDir string `json:"chdir"` - Env map[string]string `json:"env"` + ChDir string `json:"chdir,omitempty"` + Env map[string]string `json:"env,omitempty"` Command string `json:"command"` } @@ -76,6 +109,7 @@ type BarePacketType struct { } type ErrorPacketType struct { + Id string `json:"id,omitempty"` Type string `json:"type"` Error string `json:"error"` } @@ -88,6 +122,10 @@ func MakeErrorPacket(errorStr string) *ErrorPacketType { return &ErrorPacketType{Type: ErrorPacketStr, Error: errorStr} } +func MakeIdErrorPacket(id string, errorStr string) *ErrorPacketType { + return &ErrorPacketType{Type: ErrorPacketStr, Id: id, Error: errorStr} +} + type PacketType interface { GetType() string } @@ -123,13 +161,21 @@ func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { } return &errorPacket, nil } - if bareCmd.Type == OkCmdPacketStr { - var okPacket OkCmdPacketType - err = json.Unmarshal(jsonBuf, &okPacket) + if bareCmd.Type == CmdStartPacketStr { + var startPacket CmdStartPacketType + err = json.Unmarshal(jsonBuf, &startPacket) if err != nil { return nil, err } - return &okPacket, nil + return &startPacket, nil + } + if bareCmd.Type == CmdDonePacketStr { + var donePacket CmdDonePacketType + err = json.Unmarshal(jsonBuf, &donePacket) + if err != nil { + return nil, err + } + return &donePacket, nil } return nil, fmt.Errorf("invalid packet-type '%s'", bareCmd.Type) } @@ -154,6 +200,73 @@ func SendErrorPacket(w io.Writer, errorStr string) error { return SendPacket(w, MakeErrorPacket(errorStr)) } +type PacketSender struct { + Lock *sync.Mutex + SendCh chan PacketType + Err error + Done bool + DoneCh chan bool +} + +func MakePacketSender(output io.Writer) *PacketSender { + sender := &PacketSender{ + Lock: &sync.Mutex{}, + SendCh: make(chan PacketType), + DoneCh: make(chan bool), + } + go func() { + defer func() { + sender.Lock.Lock() + sender.Done = true + sender.Lock.Unlock() + close(sender.DoneCh) + }() + for pk := range sender.SendCh { + err := SendPacket(output, pk) + if err != nil { + sender.Lock.Lock() + sender.Err = err + sender.Lock.Unlock() + return + } + } + }() + return sender +} + +func (sender *PacketSender) CloseSendCh() { + close(sender.SendCh) +} + +func (sender *PacketSender) WaitForDone() { + <-sender.DoneCh +} + +func (sender *PacketSender) checkStatus() error { + sender.Lock.Lock() + defer sender.Lock.Unlock() + if sender.Done { + return fmt.Errorf("cannot send packet, sender write loop is closed") + } + if sender.Err != nil { + return fmt.Errorf("cannot send packet, sender had error: %w", sender.Err) + } + return nil +} + +func (sender *PacketSender) SendPacket(pk PacketType) error { + err := sender.checkStatus() + if err != nil { + return err + } + sender.SendCh <- pk + return nil +} + +func (sender *PacketSender) SendErrorPacket(errVal string) error { + return sender.SendPacket(MakeErrorPacket(errVal)) +} + func PacketParser(input io.Reader) chan PacketType { bufReader := bufio.NewReader(input) rtnCh := make(chan PacketType) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 4ef5f660..70e75f1d 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -7,7 +7,6 @@ package shexec import ( - "encoding/json" "errors" "fmt" "io" @@ -24,11 +23,6 @@ import ( "github.com/scripthaus-dev/sh2-runner/pkg/packet" ) -type DoneData struct { - DurationMs int64 `json:"durationms"` - ExitCode int `json:"exitcode"` -} - type ShExecType struct { FileNames *base.CommandFileNames Cmd *exec.Cmd @@ -95,6 +89,15 @@ func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { return ecmd } +func MakeRunnerExec(cmdId string) (*exec.Cmd, error) { + runnerPath, err := base.GetScRunnerPath() + if err != nil { + return nil, err + } + ecmd := exec.Command(runnerPath, cmdId) + return ecmd, nil +} + // this will never return (unless there is an error creating/opening the file), as fifoFile will never EOF func MakeAndCopyStdinFifo(dst *os.File, fifoName string) error { os.Remove(fifoName) @@ -147,8 +150,8 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { return nil } -// returning nil error means the process has successfully been kicked-off -func RunCommand(pk *packet.RunPacketType) (*ShExecType, error) { +// when err is nil, the command will have already been started +func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { if pk.CmdId == "" { pk.CmdId = uuid.New().String() } @@ -184,14 +187,14 @@ func RunCommand(pk *packet.RunPacketType) (*ShExecType, error) { // copy pty output to .ptyout file _, copyErr := io.Copy(ptyOutFd, cmdPty) if copyErr != nil { - base.WriteErrorMsg(fileNames.PtyOutFile, fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) + sender.SendErrorPacket(fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) } }() go func() { // copy .stdin fifo contents to pty input copyFifoErr := MakeAndCopyStdinFifo(cmdPty, fileNames.StdinFifo) if copyFifoErr != nil { - base.WriteErrorMsg(fileNames.PtyOutFile, fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) + sender.SendErrorPacket(fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) } }() return &ShExecType{ @@ -202,9 +205,10 @@ func RunCommand(pk *packet.RunPacketType) (*ShExecType, error) { }, nil } -func (c *ShExecType) WaitForCommand() { +func (c *ShExecType) WaitForCommand(cmdId string) *packet.CmdDonePacketType { err := c.Cmd.Wait() - cmdDuration := time.Since(c.StartTs) + endTs := time.Now() + cmdDuration := endTs.Sub(c.StartTs) exitCode := 0 if err != nil { exitErr, ok := err.(*exec.ExitError) @@ -212,15 +216,11 @@ func (c *ShExecType) WaitForCommand() { exitCode = exitErr.ExitCode() } } - doneData := DoneData{ - DurationMs: int64(cmdDuration / time.Millisecond), - ExitCode: exitCode, - } - doneDataBytes, _ := json.Marshal(doneData) - doneDataBytes = append(doneDataBytes, '\n') - err = os.WriteFile(c.FileNames.DoneFile, doneDataBytes, 0600) - if err != nil { - base.WriteErrorMsg(c.FileNames.PtyOutFile, fmt.Sprintf("reading from stdin fifo: %v", err)) - } - return + donePacket := packet.MakeCmdDonePacket() + donePacket.Ts = endTs.UnixMilli() + donePacket.CmdId = cmdId + donePacket.ExitCode = exitCode + donePacket.DurationMs = int64(cmdDuration / time.Millisecond) + os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) + return donePacket } From 65094ad0eca136434699e00a00e343f7903817f8 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 14 Jun 2022 14:17:36 -0700 Subject: [PATCH 003/149] updates, groundwork for tail, better parsing logic --- pkg/cmdtail/cmdtail.go | 51 ++++++++++++++++++++ pkg/packet/packet.go | 104 ++++++++++++++++++++++++++--------------- 2 files changed, 118 insertions(+), 37 deletions(-) create mode 100644 pkg/cmdtail/cmdtail.go diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go new file mode 100644 index 00000000..f95c5d82 --- /dev/null +++ b/pkg/cmdtail/cmdtail.go @@ -0,0 +1,51 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package cmdtail + +import ( + "sync" + + "github.com/fsnotify/fsnotify" + "github.com/scripthaus-dev/sh2-runner/pkg/packet" +) + +type TailPos struct { + CmdKey CmdKey + Pos int + RunOut bool + RunOutPos int +} + +type CmdKey struct { + SessionId string + CmdId string +} + +type Tailer struct { + Lock *sync.Mutex + WatchList map[CmdKey]TailPos + Sessions map[string]bool + Watcher *fsnotify.Watcher +} + +func MakeTailer() (*Tailer, error) { + rtn := &Tailer{ + Lock: &sync.Mutex{}, + WatchList: make(map[CmdKey]TailPos), + Sessions: make(map[string]bool), + } + var err error + rtn.Watcher, err = fsnotify.NewWatcher() + if err != nil { + return nil, err + } + return rtn, nil +} + +func AddWatch(getPacket *packet.GetCmdPacketType) error { + return nil +} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 62328467..0a5f2406 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -11,6 +11,7 @@ import ( "encoding/json" "fmt" "io" + "reflect" "sync" ) @@ -21,6 +22,32 @@ const ErrorPacketStr = "error" const MessagePacketStr = "message" const CmdStartPacketStr = "cmdstart" const CmdDonePacketStr = "cmddone" +const ListCmdPacketStr = "lscmd" +const GetCmdPacketStr = "getcmd" + +var TypeStrToFactory map[string]reflect.Type + +func init() { + TypeStrToFactory = make(map[string]reflect.Type) + TypeStrToFactory[RunPacketStr] = reflect.TypeOf(RunPacketType{}) + TypeStrToFactory[PingPacketStr] = reflect.TypeOf(PingPacketType{}) + TypeStrToFactory[DonePacketStr] = reflect.TypeOf(DonePacketType{}) + TypeStrToFactory[ErrorPacketStr] = reflect.TypeOf(ErrorPacketType{}) + TypeStrToFactory[MessagePacketStr] = reflect.TypeOf(MessagePacketType{}) + TypeStrToFactory[CmdStartPacketStr] = reflect.TypeOf(CmdStartPacketType{}) + TypeStrToFactory[CmdDonePacketStr] = reflect.TypeOf(CmdDonePacketType{}) + TypeStrToFactory[ListCmdPacketStr] = reflect.TypeOf(ListCmdPacketType{}) + TypeStrToFactory[GetCmdPacketStr] = reflect.TypeOf(GetCmdPacketType{}) +} + +func MakePacket(packetType string) (PacketType, error) { + rtype := TypeStrToFactory[packetType] + if rtype == nil { + return nil, fmt.Errorf("invalid packet type '%s'", packetType) + } + rtn := reflect.New(rtype) + return rtn.Interface().(PacketType), nil +} type PingPacketType struct { Type string `json:"type"` @@ -34,6 +61,35 @@ func MakePingPacket() *PingPacketType { return &PingPacketType{Type: PingPacketStr} } +type GetCmdPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + Tail bool `json:"tail,omitempty"` + RunOut bool `json:"runout,omitempty"` +} + +func (*GetCmdPacketType) GetType() string { + return GetCmdPacketStr +} + +func MakeGetCmdPacket() *GetCmdPacketType { + return &GetCmdPacketType{Type: GetCmdPacketStr} +} + +type ListCmdPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` +} + +func (*ListCmdPacketType) GetType() string { + return ListCmdPacketStr +} + +func MakeListCmdPacket(sessionId string) *ListCmdPacketType { + return &ListCmdPacketType{Type: ListCmdPacketStr, SessionId: sessionId} +} + type MessagePacketType struct { Type string `json:"type"` Message string `json:"message"` @@ -104,6 +160,10 @@ func (ct *RunPacketType) GetType() string { return RunPacketStr } +func MakeRunPacket() *RunPacketType { + return &RunPacketType{Type: RunPacketStr} +} + type BarePacketType struct { Type string `json:"type"` } @@ -139,45 +199,15 @@ func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { if bareCmd.Type == "" { return nil, fmt.Errorf("received packet with no type") } - if bareCmd.Type == RunPacketStr { - var runPacket RunPacketType - err = json.Unmarshal(jsonBuf, &runPacket) - if err != nil { - return nil, err - } - return &runPacket, nil + pk, err := MakePacket(bareCmd.Type) + if err != nil { + return nil, err } - if bareCmd.Type == PingPacketStr { - return MakePingPacket(), nil + err = json.Unmarshal(jsonBuf, pk) + if err != nil { + return nil, err } - if bareCmd.Type == DonePacketStr { - return MakeDonePacket(), nil - } - if bareCmd.Type == ErrorPacketStr { - var errorPacket ErrorPacketType - err = json.Unmarshal(jsonBuf, &errorPacket) - if err != nil { - return nil, err - } - return &errorPacket, nil - } - if bareCmd.Type == CmdStartPacketStr { - var startPacket CmdStartPacketType - err = json.Unmarshal(jsonBuf, &startPacket) - if err != nil { - return nil, err - } - return &startPacket, nil - } - if bareCmd.Type == CmdDonePacketStr { - var donePacket CmdDonePacketType - err = json.Unmarshal(jsonBuf, &donePacket) - if err != nil { - return nil, err - } - return &donePacket, nil - } - return nil, fmt.Errorf("invalid packet-type '%s'", bareCmd.Type) + return pk, nil } func SendPacket(w io.Writer, packet PacketType) error { From ecceb67f2063353ee94ba7d90be13b731d5d6b40 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 14 Jun 2022 22:16:58 -0700 Subject: [PATCH 004/149] got session/command tailing working. server can send getcmd packets, and client responds with cmddata packets --- go.mod | 2 + go.sum | 4 + main-runner.go | 37 ++++++- pkg/base/base.go | 19 +++- pkg/cmdtail/cmdtail.go | 220 +++++++++++++++++++++++++++++++++++++++-- pkg/packet/packet.go | 80 ++++++++++++++- 6 files changed, 350 insertions(+), 12 deletions(-) diff --git a/go.mod b/go.mod index a794172c..e64ed540 100644 --- a/go.mod +++ b/go.mod @@ -4,5 +4,7 @@ go 1.17 require ( github.com/creack/pty v1.1.18 // indirect + github.com/fsnotify/fsnotify v1.5.4 // indirect github.com/google/uuid v1.3.0 // indirect + golang.org/x/sys v0.0.0-20220412211240-33da011f77ad // indirect ) diff --git a/go.sum b/go.sum index 12d5eadc..fd7b5ca7 100644 --- a/go.sum +++ b/go.sum @@ -1,4 +1,8 @@ github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= +github.com/fsnotify/fsnotify v1.5.4 h1:jRbGcIw6P2Meqdwuo0H1p6JVLbL5DHKAKlYndzMwVZI= +github.com/fsnotify/fsnotify v1.5.4/go.mod h1:OVB6XrOHzAwXMpEM7uPOzcehqUV2UqJxmVXmkdnm1bU= github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +golang.org/x/sys v0.0.0-20220412211240-33da011f77ad h1:ntjMns5wyP/fN65tdBD4g8J5w8n015+iIIs9rtjXkY0= +golang.org/x/sys v0.0.0-20220412211240-33da011f77ad/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= diff --git a/main-runner.go b/main-runner.go index 19b0ba3c..43a32904 100644 --- a/main-runner.go +++ b/main-runner.go @@ -15,6 +15,7 @@ import ( "github.com/google/uuid" "github.com/scripthaus-dev/sh2-runner/pkg/base" + "github.com/scripthaus-dev/sh2-runner/pkg/cmdtail" "github.com/scripthaus-dev/sh2-runner/pkg/packet" "github.com/scripthaus-dev/sh2-runner/pkg/shexec" ) @@ -125,15 +126,26 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { }() } +func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { + // non-tail packets? + sender.SendPacket(packet.MakeMessagePacket(fmt.Sprintf("getcmd %s", pk.CmdId))) + err := tailer.AddWatch(pk) + if err != nil { + return err + } + return nil +} + func doMain() { - homeDir, err := base.GetScHomeDir() + scHomeDir, err := base.GetScHomeDir() if err != nil { packet.SendErrorPacket(os.Stdout, err.Error()) return } + homeDir := base.GetHomeDir() err = os.Chdir(homeDir) if err != nil { - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to scripthaus home '%s': %v", homeDir, err)) + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) return } err = base.EnsureRunnerPath() @@ -143,7 +155,18 @@ func doMain() { } packetCh := packet.PacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) - sender.SendPacket(packet.MakeMessagePacket(fmt.Sprintf("starting scripthaus runner @ %s", homeDir))) + tailer, err := cmdtail.MakeTailer(sender) + if err != nil { + packet.SendErrorPacket(os.Stdout, err.Error()) + return + } + go tailer.Run() + sender.SendPacket(packet.MakeMessagePacket(fmt.Sprintf("starting scripthaus runner @ %s", scHomeDir))) + initPacket := packet.MakeRunnerInitPacket() + initPacket.Env = os.Environ() + initPacket.HomeDir = homeDir + initPacket.ScHomeDir = scHomeDir + sender.SendPacket(initPacket) for pk := range packetCh { if pk.GetType() == packet.PingPacketStr { continue @@ -152,6 +175,14 @@ func doMain() { doMainRun(pk.(*packet.RunPacketType), sender) continue } + if pk.GetType() == packet.GetCmdPacketStr { + err = doGetCmd(tailer, pk.(*packet.GetCmdPacketType), sender) + if err != nil { + errPk := packet.MakeErrorPacket(err.Error()) + sender.SendPacket(errPk) + } + continue + } if pk.GetType() == packet.ErrorPacketStr { errPk := pk.(*packet.ErrorPacketType) errPk.Error = "invalid packet sent to runner: " + errPk.Error diff --git a/pkg/base/base.go b/pkg/base/base.go index fd2d64ed..5da63c01 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -32,6 +32,14 @@ type CommandFileNames struct { RunnerOutFile string } +func GetHomeDir() string { + homeVar := os.Getenv(HomeVarName) + if homeVar == "" { + return "/" + } + return homeVar +} + func GetScHomeDir() (string, error) { scHome := os.Getenv(ScHomeVarName) if scHome == "" { @@ -60,6 +68,15 @@ func GetCommandFileNames(sessionId string, cmdId string) (*CommandFileNames, err }, nil } +func MakeCommandFileNamesWithHome(scHome string, sessionId string, cmdId string) *CommandFileNames { + base := path.Join(scHome, SessionsDirBaseName, sessionId, cmdId) + return &CommandFileNames{ + PtyOutFile: base + ".ptyout", + StdinFifo: base + ".stdin", + RunnerOutFile: base + ".runout", + } +} + func CleanUpCmdFiles(sessionId string, cmdId string) error { if cmdId == "" { return fmt.Errorf("bad cmdid, cannot clean up") @@ -90,7 +107,7 @@ func EnsureSessionDir(sessionId string) (string, error) { if err != nil { return "", err } - sdir := path.Join(shhome, ".sessions", sessionId) + sdir := path.Join(shhome, SessionsDirBaseName, sessionId) info, err := os.Stat(sdir) if errors.Is(err, fs.ErrNotExist) { err = os.MkdirAll(sdir, 0777) diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index f95c5d82..219ec236 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -7,17 +7,30 @@ package cmdtail import ( + "fmt" + "io" + "os" + "path" + "regexp" "sync" + "time" "github.com/fsnotify/fsnotify" + "github.com/google/uuid" + "github.com/scripthaus-dev/sh2-runner/pkg/base" "github.com/scripthaus-dev/sh2-runner/pkg/packet" ) +const MaxDataBytes = 4096 + type TailPos struct { - CmdKey CmdKey - Pos int - RunOut bool - RunOutPos int + CmdKey CmdKey + Running bool // an active tailer sending data + Version int + FilePtyLen int64 + FileRunLen int64 + TailPtyPos int64 + TailRunPos int64 } type CmdKey struct { @@ -30,15 +43,22 @@ type Tailer struct { WatchList map[CmdKey]TailPos Sessions map[string]bool Watcher *fsnotify.Watcher + ScHomeDir string + Sender *packet.PacketSender } -func MakeTailer() (*Tailer, error) { +func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { + scHomeDir, err := base.GetScHomeDir() + if err != nil { + return nil, err + } rtn := &Tailer{ Lock: &sync.Mutex{}, WatchList: make(map[CmdKey]TailPos), Sessions: make(map[string]bool), + ScHomeDir: scHomeDir, + Sender: sender, } - var err error rtn.Watcher, err = fsnotify.NewWatcher() if err != nil { return nil, err @@ -46,6 +66,192 @@ func MakeTailer() (*Tailer, error) { return rtn, nil } -func AddWatch(getPacket *packet.GetCmdPacketType) error { +func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]byte, error) { + fd, err := os.Open(fileName) + defer fd.Close() + if err != nil { + return nil, err + } + buf := make([]byte, maxBytes) + nr, err := fd.ReadAt(buf, pos) + if err != nil && err != io.EOF { // ignore EOF error + return nil, err + } + return buf[0:nr], nil +} + +func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, pos TailPos) *packet.CmdDataPacketType { + dataPacket := packet.MakeCmdDataPacket() + dataPacket.SessionId = pos.CmdKey.SessionId + dataPacket.CmdId = pos.CmdKey.CmdId + dataPacket.PtyPos = pos.TailPtyPos + dataPacket.RunPos = pos.TailRunPos + if pos.FilePtyLen > pos.TailPtyPos { + ptyData, err := t.readDataFromFile(fileNames.PtyOutFile, pos.TailPtyPos, MaxDataBytes) + if err != nil { + dataPacket.Error = err.Error() + return dataPacket + } + dataPacket.PtyData = string(ptyData) + } + if pos.FileRunLen > pos.TailRunPos { + runData, err := t.readDataFromFile(fileNames.RunnerOutFile, pos.TailRunPos, MaxDataBytes) + if err != nil { + dataPacket.Error = err.Error() + return dataPacket + } + dataPacket.RunData = string(runData) + } + return dataPacket +} + +var updateFileRe = regexp.MustCompile("/([a-z0-9-]+)/([a-z0-9-]+)\\.(ptyout|runout)$") + +// returns (data-packet, keepRunning) +func (t *Tailer) runSingleDataTransfer(key CmdKey) (*packet.CmdDataPacketType, bool) { + t.Lock.Lock() + pos, foundPos := t.WatchList[key] + t.Lock.Unlock() + if !foundPos { + return nil, false + } + fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, key.SessionId, key.CmdId) + dataPacket := t.makeCmdDataPacket(fileNames, pos) + + t.Lock.Lock() + defer t.Lock.Unlock() + pos, foundPos = t.WatchList[key] + if !foundPos { + return nil, false + } + // pos was updated between first and second get, throw out data-packet and re-run + if pos.TailPtyPos != dataPacket.PtyPos || pos.TailRunPos != dataPacket.RunPos { + return nil, true + } + if dataPacket.Error != "" { + // error, so return error packet, and stop running + pos.Running = false + t.WatchList[key] = pos + return dataPacket, false + } + pos.TailPtyPos += int64(len(dataPacket.PtyData)) + pos.TailRunPos += int64(len(dataPacket.RunData)) + if pos.TailPtyPos > pos.FilePtyLen { + pos.FilePtyLen = pos.TailPtyPos + } + if pos.TailRunPos > pos.FileRunLen { + pos.FileRunLen = pos.TailRunPos + } + if pos.TailPtyPos >= pos.FilePtyLen && pos.TailRunPos >= pos.FileRunLen { + // we caught up, tail position equals file length + pos.Running = false + } + t.WatchList[key] = pos + return dataPacket, pos.Running +} + +func (t *Tailer) RunDataTransfer(key CmdKey) { + for { + dataPacket, keepRunning := t.runSingleDataTransfer(key) + if dataPacket != nil { + t.Sender.SendPacket(dataPacket) + } + if !keepRunning { + break + } + time.Sleep(10 * time.Millisecond) + } +} + +func (t *Tailer) UpdateFile(relFileName string) { + m := updateFileRe.FindStringSubmatch(relFileName) + if m == nil { + return + } + finfo, err := os.Stat(relFileName) + if err != nil { + t.Sender.SendMessage("error stating file '%s': %w", relFileName, err) + return + } + isPtyFile := m[3] == "ptyout" + cmdKey := CmdKey{m[1], m[2]} + fileSize := finfo.Size() + t.Lock.Lock() + defer t.Lock.Unlock() + pos, foundPos := t.WatchList[cmdKey] + if !foundPos { + return + } + if isPtyFile { + pos.FilePtyLen = fileSize + } else { + pos.FileRunLen = fileSize + } + t.WatchList[cmdKey] = pos + if !pos.Running && (pos.FilePtyLen > pos.TailPtyPos || pos.FileRunLen > pos.TailRunPos) { + go t.RunDataTransfer(cmdKey) + } +} + +func (t *Tailer) Run() { + for { + select { + case event, ok := <-t.Watcher.Events: + if !ok { + return + } + if (event.Op&fsnotify.Write == fsnotify.Write) || (event.Op&fsnotify.Create == fsnotify.Create) { + t.UpdateFile(event.Name) + } + + case err, ok := <-t.Watcher.Errors: + if !ok { + return + } + // what to do with watcher error? + t.Sender.SendMessage("error in tailer '%v'", err) + } + } +} + +func (tp *TailPos) fillFilePos(scHomeDir string) { + fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, tp.CmdKey.SessionId, tp.CmdKey.CmdId) + ptyInfo, _ := os.Stat(fileNames.PtyOutFile) + if ptyInfo != nil { + tp.FilePtyLen = ptyInfo.Size() + } + runoutInfo, _ := os.Stat(fileNames.RunnerOutFile) + if runoutInfo != nil { + tp.FileRunLen = runoutInfo.Size() + } +} + +func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { + if !getPacket.Tail { + return fmt.Errorf("cannot add a watch for non-tail packet") + } + _, err := uuid.Parse(getPacket.SessionId) + if err != nil { + return fmt.Errorf("getcmd, bad sessionid '%s': %w", getPacket.SessionId, err) + } + _, err = uuid.Parse(getPacket.CmdId) + if err != nil { + return fmt.Errorf("getcmd, bad cmdid '%s': %w", getPacket.CmdId, err) + } + t.Lock.Lock() + defer t.Lock.Unlock() + key := CmdKey{getPacket.SessionId, getPacket.CmdId} + if !t.Sessions[getPacket.SessionId] { + sessionDir := path.Join(t.ScHomeDir, base.SessionsDirBaseName, getPacket.SessionId) + err = t.Watcher.Add(sessionDir) + if err != nil { + return fmt.Errorf("error adding watcher for session dir '%s': %v", sessionDir, err) + } + t.Sessions[getPacket.SessionId] = true + } + oldPos := t.WatchList[key] + pos := TailPos{CmdKey: key, TailPtyPos: getPacket.PtyPos, TailRunPos: getPacket.RunPos, Version: oldPos.Version + 1} + pos.fillFilePos(t.ScHomeDir) + t.WatchList[key] = pos return nil } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 0a5f2406..987650eb 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -24,6 +24,10 @@ const CmdStartPacketStr = "cmdstart" const CmdDonePacketStr = "cmddone" const ListCmdPacketStr = "lscmd" const GetCmdPacketStr = "getcmd" +const RunnerInitPacketStr = "runnerinit" +const CdPacketStr = "cd" +const CdResponseStr = "cdresp" +const CmdDataPacketStr = "cmddata" var TypeStrToFactory map[string]reflect.Type @@ -38,6 +42,11 @@ func init() { TypeStrToFactory[CmdDonePacketStr] = reflect.TypeOf(CmdDonePacketType{}) TypeStrToFactory[ListCmdPacketStr] = reflect.TypeOf(ListCmdPacketType{}) TypeStrToFactory[GetCmdPacketStr] = reflect.TypeOf(GetCmdPacketType{}) + TypeStrToFactory[RunnerInitPacketStr] = reflect.TypeOf(RunnerInitPacketType{}) + TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) + TypeStrToFactory[CdResponseStr] = reflect.TypeOf(CdResponseType{}) + TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) + } func MakePacket(packetType string) (PacketType, error) { @@ -49,6 +58,26 @@ func MakePacket(packetType string) (PacketType, error) { return rtn.Interface().(PacketType), nil } +type CmdDataPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + PtyPos int64 `json:"ptypos"` + RunPos int64 `json:"runpos"` + PtyData string `json:"ptydata"` + RunData string `json:"rundata"` + Done bool `json:"done"` + Error string `json:"error"` +} + +func (*CmdDataPacketType) GetType() string { + return CmdDataPacketStr +} + +func MakeCmdDataPacket() *CmdDataPacketType { + return &CmdDataPacketType{Type: CmdDataPacketStr} +} + type PingPacketType struct { Type string `json:"type"` } @@ -65,8 +94,9 @@ type GetCmdPacketType struct { Type string `json:"type"` SessionId string `json:"sessionid"` CmdId string `json:"cmdid"` + PtyPos int64 `json:"ptypos"` + RunPos int64 `json:"runpos"` Tail bool `json:"tail,omitempty"` - RunOut bool `json:"runout,omitempty"` } func (*GetCmdPacketType) GetType() string { @@ -90,6 +120,35 @@ func MakeListCmdPacket(sessionId string) *ListCmdPacketType { return &ListCmdPacketType{Type: ListCmdPacketStr, SessionId: sessionId} } +type CdPacketType struct { + Type string `json:"type"` + PacketId string `json:"packetid"` + Dir string `json:"dir"` +} + +func (*CdPacketType) GetType() string { + return CdPacketStr +} + +func MakeCdPacket() *CdPacketType { + return &CdPacketType{Type: CdPacketStr} +} + +type CdResponseType struct { + Type string `json:"type"` + PacketId string `json:"packetid"` + Success bool `json:"success"` + Error string `json:"error"` +} + +func (*CdResponseType) GetType() string { + return CdResponseStr +} + +func MakeCdResponse() *CdResponseType { + return &CdResponseType{Type: CdResponseStr} +} + type MessagePacketType struct { Type string `json:"type"` Message string `json:"message"` @@ -103,6 +162,21 @@ func MakeMessagePacket(message string) *MessagePacketType { return &MessagePacketType{Type: MessagePacketStr, Message: message} } +type RunnerInitPacketType struct { + Type string `json:"type"` + ScHomeDir string `json:"schomedir"` + HomeDir string `json:"homedir"` + Env []string `json:"env"` +} + +func (*RunnerInitPacketType) GetType() string { + return RunnerInitPacketStr +} + +func MakeRunnerInitPacket() *RunnerInitPacketType { + return &RunnerInitPacketType{Type: RunnerInitPacketStr} +} + type DonePacketType struct { Type string `json:"type"` } @@ -297,6 +371,10 @@ func (sender *PacketSender) SendErrorPacket(errVal string) error { return sender.SendPacket(MakeErrorPacket(errVal)) } +func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) error { + return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) +} + func PacketParser(input io.Reader) chan PacketType { bufReader := bufio.NewReader(input) rtnCh := make(chan PacketType) From 5176128346e167727e36796ba45bb900afd01a23 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 15 Jun 2022 16:29:39 -0700 Subject: [PATCH 005/149] updates to allow parsing non-packets into raw packets (to see stderr output), and new sessionwatcher --- main-runner.go | 2 - pkg/cmdtail/cmdtail.go | 116 +++++++++++++----------- pkg/cmdtail/sessionwatcher.go | 161 ++++++++++++++++++++++++++++++++++ pkg/packet/packet.go | 50 +++++++++-- pkg/shexec/shexec.go | 21 +++-- 5 files changed, 281 insertions(+), 69 deletions(-) create mode 100644 pkg/cmdtail/sessionwatcher.go diff --git a/main-runner.go b/main-runner.go index 43a32904..3c74d49a 100644 --- a/main-runner.go +++ b/main-runner.go @@ -128,7 +128,6 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { // non-tail packets? - sender.SendPacket(packet.MakeMessagePacket(fmt.Sprintf("getcmd %s", pk.CmdId))) err := tailer.AddWatch(pk) if err != nil { return err @@ -161,7 +160,6 @@ func doMain() { return } go tailer.Run() - sender.SendPacket(packet.MakeMessagePacket(fmt.Sprintf("starting scripthaus runner @ %s", scHomeDir))) initPacket := packet.MakeRunnerInitPacket() initPacket.Env = os.Environ() initPacket.HomeDir = homeDir diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 219ec236..f82c41e6 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -10,12 +10,9 @@ import ( "fmt" "io" "os" - "path" - "regexp" "sync" "time" - "github.com/fsnotify/fsnotify" "github.com/google/uuid" "github.com/scripthaus-dev/sh2-runner/pkg/base" "github.com/scripthaus-dev/sh2-runner/pkg/packet" @@ -26,11 +23,15 @@ const MaxDataBytes = 4096 type TailPos struct { CmdKey CmdKey Running bool // an active tailer sending data - Version int FilePtyLen int64 FileRunLen int64 TailPtyPos int64 TailRunPos int64 + Follow bool +} + +func (pos TailPos) IsCurrent() bool { + return pos.TailPtyPos >= pos.FilePtyLen && pos.TailRunPos >= pos.FileRunLen } type CmdKey struct { @@ -41,9 +42,8 @@ type CmdKey struct { type Tailer struct { Lock *sync.Mutex WatchList map[CmdKey]TailPos - Sessions map[string]bool - Watcher *fsnotify.Watcher ScHomeDir string + Watcher *SessionWatcher Sender *packet.PacketSender } @@ -55,11 +55,10 @@ func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { rtn := &Tailer{ Lock: &sync.Mutex{}, WatchList: make(map[CmdKey]TailPos), - Sessions: make(map[string]bool), ScHomeDir: scHomeDir, Sender: sender, } - rtn.Watcher, err = fsnotify.NewWatcher() + rtn.Watcher, err = MakeSessionWatcher() if err != nil { return nil, err } @@ -105,8 +104,6 @@ func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, pos TailPos return dataPacket } -var updateFileRe = regexp.MustCompile("/([a-z0-9-]+)/([a-z0-9-]+)\\.(ptyout|runout)$") - // returns (data-packet, keepRunning) func (t *Tailer) runSingleDataTransfer(key CmdKey) (*packet.CmdDataPacketType, bool) { t.Lock.Lock() @@ -150,6 +147,18 @@ func (t *Tailer) runSingleDataTransfer(key CmdKey) (*packet.CmdDataPacketType, b return dataPacket, pos.Running } +func (t *Tailer) checkRemoveNoFollow(cmdKey CmdKey) { + t.Lock.Lock() + defer t.Lock.Unlock() + pos, foundPos := t.WatchList[cmdKey] + if !foundPos { + return + } + if !pos.Follow { + delete(t.WatchList, cmdKey) + } +} + func (t *Tailer) RunDataTransfer(key CmdKey) { for { dataPacket, keepRunning := t.runSingleDataTransfer(key) @@ -157,73 +166,78 @@ func (t *Tailer) RunDataTransfer(key CmdKey) { t.Sender.SendPacket(dataPacket) } if !keepRunning { + t.checkRemoveNoFollow(key) break } time.Sleep(10 * time.Millisecond) } } -func (t *Tailer) UpdateFile(relFileName string) { - m := updateFileRe.FindStringSubmatch(relFileName) - if m == nil { +// should already hold t.Lock +func (t *Tailer) tryStartRun_nolock(pos TailPos) { + if pos.Running || pos.IsCurrent() { return } - finfo, err := os.Stat(relFileName) - if err != nil { - t.Sender.SendMessage("error stating file '%s': %w", relFileName, err) + pos.Running = true + t.WatchList[pos.CmdKey] = pos + go t.RunDataTransfer(pos.CmdKey) +} + +func (t *Tailer) updateFile(event FileUpdateEvent) { + if event.Err != nil { + t.Sender.SendMessage("error in FileUpdateEvent %s/%s: %v", event.SessionId, event.CmdId, event.Err) return } - isPtyFile := m[3] == "ptyout" - cmdKey := CmdKey{m[1], m[2]} - fileSize := finfo.Size() + cmdKey := CmdKey{SessionId: event.SessionId, CmdId: event.CmdId} t.Lock.Lock() defer t.Lock.Unlock() pos, foundPos := t.WatchList[cmdKey] if !foundPos { return } - if isPtyFile { - pos.FilePtyLen = fileSize - } else { - pos.FileRunLen = fileSize + if event.FileType == FileTypePty { + pos.FilePtyLen = event.Size + } else if event.FileType == FileTypeRun { + pos.FileRunLen = event.Size } t.WatchList[cmdKey] = pos - if !pos.Running && (pos.FilePtyLen > pos.TailPtyPos || pos.FileRunLen > pos.TailRunPos) { - go t.RunDataTransfer(cmdKey) - } + t.tryStartRun_nolock(pos) } -func (t *Tailer) Run() { - for { - select { - case event, ok := <-t.Watcher.Events: - if !ok { - return - } - if (event.Op&fsnotify.Write == fsnotify.Write) || (event.Op&fsnotify.Create == fsnotify.Create) { - t.UpdateFile(event.Name) - } - - case err, ok := <-t.Watcher.Errors: - if !ok { - return - } - // what to do with watcher error? - t.Sender.SendMessage("error in tailer '%v'", err) +func (t *Tailer) Run() error { + go func() { + for event := range t.Watcher.EventCh { + t.updateFile(event) } - } + }() + err := t.Watcher.Run(nil) + return err } +func max(v1 int64, v2 int64) int64 { + if v1 > v2 { + return v1 + } + return v2 +} + +// also converts negative positions to positive positions func (tp *TailPos) fillFilePos(scHomeDir string) { fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, tp.CmdKey.SessionId, tp.CmdKey.CmdId) ptyInfo, _ := os.Stat(fileNames.PtyOutFile) if ptyInfo != nil { tp.FilePtyLen = ptyInfo.Size() } + if tp.TailPtyPos < 0 { + tp.TailPtyPos = max(0, tp.FilePtyLen-tp.TailPtyPos) + } runoutInfo, _ := os.Stat(fileNames.RunnerOutFile) if runoutInfo != nil { tp.FileRunLen = runoutInfo.Size() } + if tp.TailRunPos < 0 { + tp.TailRunPos = max(0, tp.FileRunLen-tp.TailRunPos) + } } func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { @@ -241,17 +255,13 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { t.Lock.Lock() defer t.Lock.Unlock() key := CmdKey{getPacket.SessionId, getPacket.CmdId} - if !t.Sessions[getPacket.SessionId] { - sessionDir := path.Join(t.ScHomeDir, base.SessionsDirBaseName, getPacket.SessionId) - err = t.Watcher.Add(sessionDir) - if err != nil { - return fmt.Errorf("error adding watcher for session dir '%s': %v", sessionDir, err) - } - t.Sessions[getPacket.SessionId] = true + err = t.Watcher.WatchSession(getPacket.SessionId) + if err != nil { + return fmt.Errorf("error trying to watch sesion '%s': %v", getPacket.SessionId, err) } - oldPos := t.WatchList[key] - pos := TailPos{CmdKey: key, TailPtyPos: getPacket.PtyPos, TailRunPos: getPacket.RunPos, Version: oldPos.Version + 1} + pos := TailPos{CmdKey: key, TailPtyPos: getPacket.PtyPos, TailRunPos: getPacket.RunPos, Follow: getPacket.Tail} pos.fillFilePos(t.ScHomeDir) t.WatchList[key] = pos + t.tryStartRun_nolock(pos) return nil } diff --git a/pkg/cmdtail/sessionwatcher.go b/pkg/cmdtail/sessionwatcher.go new file mode 100644 index 00000000..9f3bcf81 --- /dev/null +++ b/pkg/cmdtail/sessionwatcher.go @@ -0,0 +1,161 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package cmdtail + +import ( + "fmt" + "os" + "path" + "regexp" + "sync" + + "github.com/fsnotify/fsnotify" + "github.com/google/uuid" + "github.com/scripthaus-dev/sh2-runner/pkg/base" +) + +const FileTypePty = "ptyout" +const FileTypeRun = "runout" +const eventChSize = 10 + +type FileUpdateEvent struct { + SessionId string + CmdId string + FileType string + Size int64 + Err error +} + +type SessionWatcher struct { + Lock *sync.Mutex + Sessions map[string]bool + ScHomeDir string + Watcher *fsnotify.Watcher + EventCh chan FileUpdateEvent + Err error + Running bool +} + +func MakeSessionWatcher() (*SessionWatcher, error) { + scHomeDir, err := base.GetScHomeDir() + if err != nil { + return nil, err + } + rtn := &SessionWatcher{ + Lock: &sync.Mutex{}, + Sessions: make(map[string]bool), + ScHomeDir: scHomeDir, + EventCh: make(chan FileUpdateEvent, eventChSize), + } + rtn.Watcher, err = fsnotify.NewWatcher() + if err != nil { + return nil, err + } + return rtn, nil +} + +func (w *SessionWatcher) UnWatchSession(sessionId string) error { + _, err := uuid.Parse(sessionId) + if err != nil { + return fmt.Errorf("WatchSession, bad sessionid '%s': %w", sessionId, err) + } + w.Lock.Lock() + defer w.Lock.Unlock() + if !w.Sessions[sessionId] { + return nil + } + sessionDir := path.Join(w.ScHomeDir, base.SessionsDirBaseName, sessionId) + err = w.Watcher.Remove(sessionDir) + if err != nil { + return err + } + w.Sessions[sessionId] = false + return nil +} + +func (w *SessionWatcher) WatchSession(sessionId string) error { + _, err := uuid.Parse(sessionId) + if err != nil { + return fmt.Errorf("WatchSession, bad sessionid '%s': %w", sessionId, err) + } + + w.Lock.Lock() + defer w.Lock.Unlock() + if w.Sessions[sessionId] { + return nil + } + sessionDir := path.Join(w.ScHomeDir, base.SessionsDirBaseName, sessionId) + err = w.Watcher.Add(sessionDir) + if err != nil { + return err + } + w.Sessions[sessionId] = true + return nil +} + +func (w *SessionWatcher) setRunning() bool { + w.Lock.Lock() + defer w.Lock.Unlock() + if w.Running { + return false + } + w.Running = true + return true +} + +var swUpdateFileRe = regexp.MustCompile("/([a-z0-9-]+)/([a-z0-9-]+)\\.(ptyout|runout)$") + +func (w *SessionWatcher) updateFile(relFileName string) { + m := swUpdateFileRe.FindStringSubmatch(relFileName) + if m == nil { + return + } + event := FileUpdateEvent{SessionId: m[1], CmdId: m[2], FileType: m[3]} + finfo, err := os.Stat(relFileName) + if err != nil { + event.Err = err + w.EventCh <- event + return + } + event.Size = finfo.Size() + w.EventCh <- event + return +} + +func (w *SessionWatcher) Run(stopCh chan bool) error { + ok := w.setRunning() + if !ok { + return fmt.Errorf("Cannot run SessionWatcher (alreaady running)") + } + defer func() { + w.Lock.Lock() + defer w.Lock.Unlock() + w.Running = false + close(w.EventCh) + }() + for { + select { + case event, ok := <-w.Watcher.Events: + if !ok { + return nil + } + if (event.Op&fsnotify.Write == fsnotify.Write) || (event.Op&fsnotify.Create == fsnotify.Create) { + w.updateFile(event.Name) + } + + case err, ok := <-w.Watcher.Errors: + if !ok { + return nil + } + return fmt.Errorf("Got error in SessionWatcher: %w", err) + + case <-stopCh: + return nil + } + } + return nil +} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 987650eb..ff56c4ac 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -8,10 +8,13 @@ package packet import ( "bufio" + "bytes" "encoding/json" "fmt" "io" "reflect" + "strconv" + "strings" "sync" ) @@ -28,6 +31,7 @@ const RunnerInitPacketStr = "runnerinit" const CdPacketStr = "cd" const CdResponseStr = "cdresp" const CmdDataPacketStr = "cmddata" +const RawPacketStr = "raw" var TypeStrToFactory map[string]reflect.Type @@ -46,7 +50,7 @@ func init() { TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) TypeStrToFactory[CdResponseStr] = reflect.TypeOf(CdResponseType{}) TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) - + TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) } func MakePacket(packetType string) (PacketType, error) { @@ -63,11 +67,13 @@ type CmdDataPacketType struct { SessionId string `json:"sessionid"` CmdId string `json:"cmdid"` PtyPos int64 `json:"ptypos"` + PtyLen int64 `json:"ptylen"` RunPos int64 `json:"runpos"` + RunLen int64 `json:"runlen"` PtyData string `json:"ptydata"` RunData string `json:"rundata"` - Done bool `json:"done"` Error string `json:"error"` + NotFound bool `json:"notfound,omitempty"` } func (*CmdDataPacketType) GetType() string { @@ -149,6 +155,19 @@ func MakeCdResponse() *CdResponseType { return &CdResponseType{Type: CdResponseStr} } +type RawPacketType struct { + Type string `json:"type"` + Data string `json:"data"` +} + +func (*RawPacketType) GetType() string { + return RawPacketStr +} + +func MakeRawPacket(val string) *RawPacketType { + return &RawPacketType{Type: RawPacketStr, Data: val} +} + type MessagePacketType struct { Type string `json:"type"` Message string `json:"message"` @@ -288,12 +307,16 @@ func SendPacket(w io.Writer, packet PacketType) error { if packet == nil { return nil } - barr, err := json.Marshal(packet) + jsonBytes, err := json.Marshal(packet) if err != nil { return fmt.Errorf("marshaling '%s' packet: %w", packet.GetType(), err) } - barr = append(barr, '\n') - _, err = w.Write(barr) + var outBuf bytes.Buffer + outBuf.WriteByte('\n') + outBuf.WriteString(fmt.Sprintf("##%d", len(jsonBytes))) + outBuf.Write(jsonBytes) + outBuf.WriteByte('\n') + _, err = w.Write(outBuf.Bytes()) if err != nil { return err } @@ -392,7 +415,22 @@ func PacketParser(input io.Reader) chan PacketType { rtnCh <- errPacket return } - pk, err := ParseJsonPacket([]byte(line)) + if line == "\n" { + continue + } + // ##[len][json]\n + // ##14{"hello":true}\n + bracePos := strings.Index(line, "{") + if !strings.HasPrefix(line, "##") || bracePos == -1 { + rtnCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + packetLen, err := strconv.Atoi(line[2:bracePos]) + if err != nil || packetLen != len(line)-bracePos-1 { + rtnCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + pk, err := ParseJsonPacket([]byte(line[bracePos:])) if err != nil { errPk := MakeErrorPacket(fmt.Sprintf("parsing packet json from input: %v", err)) rtnCh <- errPk diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 70e75f1d..3b0e5c20 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -205,17 +205,22 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT }, nil } +func GetExitCode(err error) int { + if err == nil { + return 0 + } + if exitErr, ok := err.(*exec.ExitError); ok { + return exitErr.ExitCode() + } else { + return -1 + } +} + func (c *ShExecType) WaitForCommand(cmdId string) *packet.CmdDonePacketType { - err := c.Cmd.Wait() + exitErr := c.Cmd.Wait() endTs := time.Now() cmdDuration := endTs.Sub(c.StartTs) - exitCode := 0 - if err != nil { - exitErr, ok := err.(*exec.ExitError) - if ok { - exitCode = exitErr.ExitCode() - } - } + exitCode := GetExitCode(exitErr) donePacket := packet.MakeCmdDonePacket() donePacket.Ts = endTs.UnixMilli() donePacket.CmdId = cmdId From c6165f15f450f13ae795892c41555a5bfbd19770 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 16 Jun 2022 01:10:56 -0700 Subject: [PATCH 006/149] update tailer to send packets to a channel instead of hard coding to PacketSender --- main-runner.go | 2 +- pkg/cmdtail/cmdtail.go | 10 +++++----- pkg/packet/packet.go | 32 ++++++++++++++++++++++++++++++++ 3 files changed, 38 insertions(+), 6 deletions(-) diff --git a/main-runner.go b/main-runner.go index 3c74d49a..c9b9fbe7 100644 --- a/main-runner.go +++ b/main-runner.go @@ -154,7 +154,7 @@ func doMain() { } packetCh := packet.PacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) - tailer, err := cmdtail.MakeTailer(sender) + tailer, err := cmdtail.MakeTailer(sender.SendCh) if err != nil { packet.SendErrorPacket(os.Stdout, err.Error()) return diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index f82c41e6..e0940869 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -44,10 +44,10 @@ type Tailer struct { WatchList map[CmdKey]TailPos ScHomeDir string Watcher *SessionWatcher - Sender *packet.PacketSender + SendCh chan packet.PacketType } -func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { +func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { scHomeDir, err := base.GetScHomeDir() if err != nil { return nil, err @@ -56,7 +56,7 @@ func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { Lock: &sync.Mutex{}, WatchList: make(map[CmdKey]TailPos), ScHomeDir: scHomeDir, - Sender: sender, + SendCh: sendCh, } rtn.Watcher, err = MakeSessionWatcher() if err != nil { @@ -163,7 +163,7 @@ func (t *Tailer) RunDataTransfer(key CmdKey) { for { dataPacket, keepRunning := t.runSingleDataTransfer(key) if dataPacket != nil { - t.Sender.SendPacket(dataPacket) + t.SendCh <- dataPacket } if !keepRunning { t.checkRemoveNoFollow(key) @@ -185,7 +185,7 @@ func (t *Tailer) tryStartRun_nolock(pos TailPos) { func (t *Tailer) updateFile(event FileUpdateEvent) { if event.Err != nil { - t.Sender.SendMessage("error in FileUpdateEvent %s/%s: %v", event.SessionId, event.CmdId, event.Err) + t.SendCh <- packet.FmtMessagePacket("error in FileUpdateEvent %s/%s: %v", event.SessionId, event.CmdId, event.Err) return } cmdKey := CmdKey{SessionId: event.SessionId, CmdId: event.CmdId} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index ff56c4ac..78f77cb0 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -181,6 +181,11 @@ func MakeMessagePacket(message string) *MessagePacketType { return &MessagePacketType{Type: MessagePacketStr, Message: message} } +func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { + message := fmt.Sprintf(fmtStr, args...) + return &MessagePacketType{Type: MessagePacketStr, Message: message} +} + type RunnerInitPacketType struct { Type string `json:"type"` ScHomeDir string `json:"schomedir"` @@ -444,3 +449,30 @@ func PacketParser(input io.Reader) chan PacketType { }() return rtnCh } + +type ErrorReporter interface { + ReportError(err error) +} + +func PacketToByteArrBridge(pkCh chan PacketType, byteCh chan []byte, errorReporter ErrorReporter, closeOnDone bool) { + go func() { + defer func() { + if closeOnDone { + close(byteCh) + } + }() + for pk := range pkCh { + if pk == nil { + continue + } + jsonBytes, err := json.Marshal(pk) + if err != nil { + if errorReporter != nil { + errorReporter.ReportError(fmt.Errorf("error marshaling packet: %w", err)) + } + continue + } + byteCh <- jsonBytes + } + }() +} From dc4baaea27c00f022340e7e547895d40fac42b4f Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 16 Jun 2022 17:24:29 -0700 Subject: [PATCH 007/149] cmdtail now can support multiple tails on the same cmd with independent tail positions --- main-runner.go | 1 - pkg/cmdtail/cmdtail.go | 190 +++++++++++++++++++++++++++++------------ pkg/packet/packet.go | 43 +++++++--- 3 files changed, 166 insertions(+), 68 deletions(-) diff --git a/main-runner.go b/main-runner.go index c9b9fbe7..a74bcf03 100644 --- a/main-runner.go +++ b/main-runner.go @@ -127,7 +127,6 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { } func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { - // non-tail packets? err := tailer.AddWatch(pk) if err != nil { return err diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index e0940869..b43e044d 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -21,17 +21,52 @@ import ( const MaxDataBytes = 4096 type TailPos struct { - CmdKey CmdKey + ReqId string Running bool // an active tailer sending data - FilePtyLen int64 - FileRunLen int64 TailPtyPos int64 TailRunPos int64 Follow bool } -func (pos TailPos) IsCurrent() bool { - return pos.TailPtyPos >= pos.FilePtyLen && pos.TailRunPos >= pos.FileRunLen +type CmdWatchEntry struct { + CmdKey CmdKey + FilePtyLen int64 + FileRunLen int64 + Tails []TailPos +} + +func (w CmdWatchEntry) getTailPos(reqId string) (TailPos, bool) { + for _, pos := range w.Tails { + if pos.ReqId == reqId { + return pos, true + } + } + return TailPos{}, false +} + +func (w *CmdWatchEntry) updateTailPos(reqId string, pos TailPos) { + for idx, pos := range w.Tails { + if pos.ReqId == reqId { + w.Tails[idx] = pos + return + } + } + w.Tails = append(w.Tails, pos) +} + +func (w *CmdWatchEntry) removeTailPos(reqId string) { + var newTails []TailPos + for _, pos := range w.Tails { + if pos.ReqId == reqId { + continue + } + newTails = append(newTails, pos) + } + w.Tails = newTails +} + +func (pos TailPos) IsCurrent(entry CmdWatchEntry) bool { + return pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen } type CmdKey struct { @@ -41,12 +76,43 @@ type CmdKey struct { type Tailer struct { Lock *sync.Mutex - WatchList map[CmdKey]TailPos + WatchList map[CmdKey]CmdWatchEntry ScHomeDir string Watcher *SessionWatcher SendCh chan packet.PacketType } +func (t *Tailer) updateTailPos_nolock(cmdKey CmdKey, reqId string, pos TailPos) { + entry, found := t.WatchList[cmdKey] + if !found { + return + } + entry.updateTailPos(reqId, pos) + t.WatchList[cmdKey] = entry +} + +func (t *Tailer) updateEntrySizes_nolock(cmdKey CmdKey, ptyLen int64, runLen int64) { + entry, found := t.WatchList[cmdKey] + if !found { + return + } + entry.FilePtyLen = ptyLen + entry.FileRunLen = runLen + t.WatchList[cmdKey] = entry +} + +func (t *Tailer) getEntryAndPos_nolock(cmdKey CmdKey, reqId string) (CmdWatchEntry, TailPos, bool) { + entry, found := t.WatchList[cmdKey] + if !found { + return CmdWatchEntry{}, TailPos{}, false + } + pos, found := entry.getTailPos(reqId) + if !found { + return CmdWatchEntry{}, TailPos{}, false + } + return entry, pos, true +} + func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { scHomeDir, err := base.GetScHomeDir() if err != nil { @@ -54,7 +120,7 @@ func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { } rtn := &Tailer{ Lock: &sync.Mutex{}, - WatchList: make(map[CmdKey]TailPos), + WatchList: make(map[CmdKey]CmdWatchEntry), ScHomeDir: scHomeDir, SendCh: sendCh, } @@ -79,45 +145,48 @@ func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]b return buf[0:nr], nil } -func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, pos TailPos) *packet.CmdDataPacketType { +func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWatchEntry, pos TailPos) *packet.CmdDataPacketType { dataPacket := packet.MakeCmdDataPacket() - dataPacket.SessionId = pos.CmdKey.SessionId - dataPacket.CmdId = pos.CmdKey.CmdId + dataPacket.ReqId = pos.ReqId + dataPacket.SessionId = entry.CmdKey.SessionId + dataPacket.CmdId = entry.CmdKey.CmdId dataPacket.PtyPos = pos.TailPtyPos dataPacket.RunPos = pos.TailRunPos - if pos.FilePtyLen > pos.TailPtyPos { + if entry.FilePtyLen > pos.TailPtyPos { ptyData, err := t.readDataFromFile(fileNames.PtyOutFile, pos.TailPtyPos, MaxDataBytes) if err != nil { dataPacket.Error = err.Error() return dataPacket } dataPacket.PtyData = string(ptyData) + dataPacket.PtyDataLen = len(ptyData) } - if pos.FileRunLen > pos.TailRunPos { + if entry.FileRunLen > pos.TailRunPos { runData, err := t.readDataFromFile(fileNames.RunnerOutFile, pos.TailRunPos, MaxDataBytes) if err != nil { dataPacket.Error = err.Error() return dataPacket } dataPacket.RunData = string(runData) + dataPacket.RunDataLen = len(runData) } return dataPacket } // returns (data-packet, keepRunning) -func (t *Tailer) runSingleDataTransfer(key CmdKey) (*packet.CmdDataPacketType, bool) { +func (t *Tailer) runSingleDataTransfer(key CmdKey, reqId string) (*packet.CmdDataPacketType, bool) { t.Lock.Lock() - pos, foundPos := t.WatchList[key] + entry, pos, foundPos := t.getEntryAndPos_nolock(key, reqId) t.Lock.Unlock() if !foundPos { return nil, false } fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, key.SessionId, key.CmdId) - dataPacket := t.makeCmdDataPacket(fileNames, pos) + dataPacket := t.makeCmdDataPacket(fileNames, entry, pos) t.Lock.Lock() defer t.Lock.Unlock() - pos, foundPos = t.WatchList[key] + entry, pos, foundPos = t.getEntryAndPos_nolock(key, reqId) if !foundPos { return nil, false } @@ -128,45 +197,44 @@ func (t *Tailer) runSingleDataTransfer(key CmdKey) (*packet.CmdDataPacketType, b if dataPacket.Error != "" { // error, so return error packet, and stop running pos.Running = false - t.WatchList[key] = pos + t.updateTailPos_nolock(key, reqId, pos) return dataPacket, false } pos.TailPtyPos += int64(len(dataPacket.PtyData)) pos.TailRunPos += int64(len(dataPacket.RunData)) - if pos.TailPtyPos > pos.FilePtyLen { - pos.FilePtyLen = pos.TailPtyPos - } - if pos.TailRunPos > pos.FileRunLen { - pos.FileRunLen = pos.TailRunPos - } - if pos.TailPtyPos >= pos.FilePtyLen && pos.TailRunPos >= pos.FileRunLen { + if pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen { // we caught up, tail position equals file length pos.Running = false } - t.WatchList[key] = pos + t.updateTailPos_nolock(key, reqId, pos) return dataPacket, pos.Running } -func (t *Tailer) checkRemoveNoFollow(cmdKey CmdKey) { +func (t *Tailer) checkRemoveNoFollow(cmdKey CmdKey, reqId string) { t.Lock.Lock() defer t.Lock.Unlock() - pos, foundPos := t.WatchList[cmdKey] + entry, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) if !foundPos { return } if !pos.Follow { - delete(t.WatchList, cmdKey) + entry.removeTailPos(reqId) + if len(entry.Tails) == 0 { + delete(t.WatchList, cmdKey) + } else { + t.WatchList[cmdKey] = entry + } } } -func (t *Tailer) RunDataTransfer(key CmdKey) { +func (t *Tailer) RunDataTransfer(key CmdKey, reqId string) { for { - dataPacket, keepRunning := t.runSingleDataTransfer(key) + dataPacket, keepRunning := t.runSingleDataTransfer(key, reqId) if dataPacket != nil { t.SendCh <- dataPacket } if !keepRunning { - t.checkRemoveNoFollow(key) + t.checkRemoveNoFollow(key, reqId) break } time.Sleep(10 * time.Millisecond) @@ -174,13 +242,13 @@ func (t *Tailer) RunDataTransfer(key CmdKey) { } // should already hold t.Lock -func (t *Tailer) tryStartRun_nolock(pos TailPos) { - if pos.Running || pos.IsCurrent() { +func (t *Tailer) tryStartRun_nolock(entry CmdWatchEntry, pos TailPos) { + if pos.Running || pos.IsCurrent(entry) { return } pos.Running = true - t.WatchList[pos.CmdKey] = pos - go t.RunDataTransfer(pos.CmdKey) + t.updateTailPos_nolock(entry.CmdKey, pos.ReqId, pos) + go t.RunDataTransfer(entry.CmdKey, pos.ReqId) } func (t *Tailer) updateFile(event FileUpdateEvent) { @@ -191,17 +259,19 @@ func (t *Tailer) updateFile(event FileUpdateEvent) { cmdKey := CmdKey{SessionId: event.SessionId, CmdId: event.CmdId} t.Lock.Lock() defer t.Lock.Unlock() - pos, foundPos := t.WatchList[cmdKey] - if !foundPos { + entry, foundEntry := t.WatchList[cmdKey] + if !foundEntry { return } if event.FileType == FileTypePty { - pos.FilePtyLen = event.Size + entry.FilePtyLen = event.Size } else if event.FileType == FileTypeRun { - pos.FileRunLen = event.Size + entry.FileRunLen = event.Size + } + t.WatchList[cmdKey] = entry + for _, pos := range entry.Tails { + t.tryStartRun_nolock(entry, pos) } - t.WatchList[cmdKey] = pos - t.tryStartRun_nolock(pos) } func (t *Tailer) Run() error { @@ -221,22 +291,15 @@ func max(v1 int64, v2 int64) int64 { return v2 } -// also converts negative positions to positive positions -func (tp *TailPos) fillFilePos(scHomeDir string) { - fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, tp.CmdKey.SessionId, tp.CmdKey.CmdId) +func (entry *CmdWatchEntry) fillFilePos(scHomeDir string) { + fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, entry.CmdKey.SessionId, entry.CmdKey.CmdId) ptyInfo, _ := os.Stat(fileNames.PtyOutFile) if ptyInfo != nil { - tp.FilePtyLen = ptyInfo.Size() - } - if tp.TailPtyPos < 0 { - tp.TailPtyPos = max(0, tp.FilePtyLen-tp.TailPtyPos) + entry.FilePtyLen = ptyInfo.Size() } runoutInfo, _ := os.Stat(fileNames.RunnerOutFile) if runoutInfo != nil { - tp.FileRunLen = runoutInfo.Size() - } - if tp.TailRunPos < 0 { - tp.TailRunPos = max(0, tp.FileRunLen-tp.TailRunPos) + entry.FileRunLen = runoutInfo.Size() } } @@ -252,6 +315,9 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { if err != nil { return fmt.Errorf("getcmd, bad cmdid '%s': %w", getPacket.CmdId, err) } + if getPacket.ReqId == "" { + return fmt.Errorf("getcmd, no reqid specified") + } t.Lock.Lock() defer t.Lock.Unlock() key := CmdKey{getPacket.SessionId, getPacket.CmdId} @@ -259,9 +325,21 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { if err != nil { return fmt.Errorf("error trying to watch sesion '%s': %v", getPacket.SessionId, err) } - pos := TailPos{CmdKey: key, TailPtyPos: getPacket.PtyPos, TailRunPos: getPacket.RunPos, Follow: getPacket.Tail} - pos.fillFilePos(t.ScHomeDir) - t.WatchList[key] = pos - t.tryStartRun_nolock(pos) + entry, foundEntry := t.WatchList[key] + if !foundEntry { + entry = CmdWatchEntry{CmdKey: key} + entry.fillFilePos(t.ScHomeDir) + } + pos := TailPos{ReqId: getPacket.ReqId, TailPtyPos: getPacket.PtyPos, TailRunPos: getPacket.RunPos, Follow: getPacket.Tail} + // convert negative pos to positive + if pos.TailPtyPos < 0 { + pos.TailPtyPos = max(0, entry.FilePtyLen+pos.TailPtyPos) // + because negative + } + if pos.TailRunPos < 0 { + pos.TailRunPos = max(0, entry.FileRunLen+pos.TailRunPos) // + because negative + } + entry.updateTailPos(pos.ReqId, pos) + t.WatchList[key] = entry + t.tryStartRun_nolock(entry, pos) return nil } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 78f77cb0..dec40066 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -27,6 +27,7 @@ const CmdStartPacketStr = "cmdstart" const CmdDonePacketStr = "cmddone" const ListCmdPacketStr = "lscmd" const GetCmdPacketStr = "getcmd" +const UntailCmdPacketStr = "untailcmd" const RunnerInitPacketStr = "runnerinit" const CdPacketStr = "cd" const CdResponseStr = "cdresp" @@ -46,6 +47,7 @@ func init() { TypeStrToFactory[CmdDonePacketStr] = reflect.TypeOf(CmdDonePacketType{}) TypeStrToFactory[ListCmdPacketStr] = reflect.TypeOf(ListCmdPacketType{}) TypeStrToFactory[GetCmdPacketStr] = reflect.TypeOf(GetCmdPacketType{}) + TypeStrToFactory[UntailCmdPacketStr] = reflect.TypeOf(UntailCmdPacketType{}) TypeStrToFactory[RunnerInitPacketStr] = reflect.TypeOf(RunnerInitPacketType{}) TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) TypeStrToFactory[CdResponseStr] = reflect.TypeOf(CdResponseType{}) @@ -63,17 +65,20 @@ func MakePacket(packetType string) (PacketType, error) { } type CmdDataPacketType struct { - Type string `json:"type"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` - PtyPos int64 `json:"ptypos"` - PtyLen int64 `json:"ptylen"` - RunPos int64 `json:"runpos"` - RunLen int64 `json:"runlen"` - PtyData string `json:"ptydata"` - RunData string `json:"rundata"` - Error string `json:"error"` - NotFound bool `json:"notfound,omitempty"` + Type string `json:"type"` + ReqId string `json:"reqid"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + PtyPos int64 `json:"ptypos"` + PtyLen int64 `json:"ptylen"` + RunPos int64 `json:"runpos"` + RunLen int64 `json:"runlen"` + PtyData string `json:"ptydata"` + PtyDataLen int `json:"ptydatalen"` + RunData string `json:"rundata"` + RunDataLen int `json:"rundatalen"` + Error string `json:"error"` + NotFound bool `json:"notfound,omitempty"` } func (*CmdDataPacketType) GetType() string { @@ -96,8 +101,24 @@ func MakePingPacket() *PingPacketType { return &PingPacketType{Type: PingPacketStr} } +type UntailCmdPacketType struct { + Type string `json:"type"` + ReqId string `json:"reqid"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` +} + +func (*UntailCmdPacketType) GetType() string { + return UntailCmdPacketStr +} + +func MakeUntailCmdPacket() *UntailCmdPacketType { + return &UntailCmdPacketType{Type: UntailCmdPacketStr} +} + type GetCmdPacketType struct { Type string `json:"type"` + ReqId string `json:"reqid"` SessionId string `json:"sessionid"` CmdId string `json:"cmdid"` PtyPos int64 `json:"ptypos"` From 3497f607cec9f3f2e1a84cd98b25e2f4816aa26b Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 16 Jun 2022 22:23:29 -0700 Subject: [PATCH 008/149] fix bug in updateTailPos, add Close() for tailers --- pkg/cmdtail/cmdtail.go | 13 +++++++------ pkg/cmdtail/sessionwatcher.go | 4 ++++ 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index b43e044d..9bf1a6f5 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -44,14 +44,14 @@ func (w CmdWatchEntry) getTailPos(reqId string) (TailPos, bool) { return TailPos{}, false } -func (w *CmdWatchEntry) updateTailPos(reqId string, pos TailPos) { +func (w *CmdWatchEntry) updateTailPos(reqId string, newPos TailPos) { for idx, pos := range w.Tails { if pos.ReqId == reqId { - w.Tails[idx] = pos + w.Tails[idx] = newPos return } } - w.Tails = append(w.Tails, pos) + w.Tails = append(w.Tails, newPos) } func (w *CmdWatchEntry) removeTailPos(reqId string) { @@ -284,6 +284,10 @@ func (t *Tailer) Run() error { return err } +func (t *Tailer) Close() error { + return t.Watcher.Close() +} + func max(v1 int64, v2 int64) int64 { if v1 > v2 { return v1 @@ -304,9 +308,6 @@ func (entry *CmdWatchEntry) fillFilePos(scHomeDir string) { } func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { - if !getPacket.Tail { - return fmt.Errorf("cannot add a watch for non-tail packet") - } _, err := uuid.Parse(getPacket.SessionId) if err != nil { return fmt.Errorf("getcmd, bad sessionid '%s': %w", getPacket.SessionId, err) diff --git a/pkg/cmdtail/sessionwatcher.go b/pkg/cmdtail/sessionwatcher.go index 9f3bcf81..1e736cb8 100644 --- a/pkg/cmdtail/sessionwatcher.go +++ b/pkg/cmdtail/sessionwatcher.go @@ -58,6 +58,10 @@ func MakeSessionWatcher() (*SessionWatcher, error) { return rtn, nil } +func (w *SessionWatcher) Close() error { + return w.Watcher.Close() +} + func (w *SessionWatcher) UnWatchSession(sessionId string) error { _, err := uuid.Parse(sessionId) if err != nil { From c8f9022db4e2f084924659c5af0f8c7f7a6be94d Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 17 Jun 2022 11:29:33 -0700 Subject: [PATCH 009/149] close runout file in main runner --- main-runner.go | 1 + 1 file changed, 1 insertion(+) diff --git a/main-runner.go b/main-runner.go index a74bcf03..12a1e5aa 100644 --- a/main-runner.go +++ b/main-runner.go @@ -106,6 +106,7 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err))) return } + defer runnerOutFd.Close() cmd.Stdout = runnerOutFd cmd.Stderr = runnerOutFd err = cmd.Start() From b6a8550ab8f7dff6a6f642c3adc12b6d585017fc Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 17 Jun 2022 12:27:29 -0700 Subject: [PATCH 010/149] revert cmdtail to watch individual files (inefficient directory watching on osx). touch ptyout file before running (because of file watching) --- main-runner.go | 8 ++ pkg/cmdtail/cmdtail.go | 124 ++++++++++++++++++------- pkg/cmdtail/sessionwatcher.go | 165 ---------------------------------- pkg/shexec/shexec.go | 10 ++- 4 files changed, 107 insertions(+), 200 deletions(-) delete mode 100644 pkg/cmdtail/sessionwatcher.go diff --git a/main-runner.go b/main-runner.go index 12a1e5aa..f8b8dc39 100644 --- a/main-runner.go +++ b/main-runner.go @@ -101,6 +101,13 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot pipe stdin to command: %v", err))) return } + // touch ptyout file (should exist for tailer to work correctly) + ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) + if err != nil { + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err))) + return + } + ptyOutFd.Close() // just opened to create the file, can close right after runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) if err != nil { sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err))) @@ -199,6 +206,7 @@ func main() { return } doSingle(cmdId.String()) + time.Sleep(100 * time.Millisecond) return } else { doMain() diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 9bf1a6f5..44b44cc5 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -10,15 +10,19 @@ import ( "fmt" "io" "os" + "regexp" "sync" "time" + "github.com/fsnotify/fsnotify" "github.com/google/uuid" "github.com/scripthaus-dev/sh2-runner/pkg/base" "github.com/scripthaus-dev/sh2-runner/pkg/packet" ) const MaxDataBytes = 4096 +const FileTypePty = "ptyout" +const FileTypeRun = "runout" type TailPos struct { ReqId string @@ -78,7 +82,7 @@ type Tailer struct { Lock *sync.Mutex WatchList map[CmdKey]CmdWatchEntry ScHomeDir string - Watcher *SessionWatcher + Watcher *fsnotify.Watcher SendCh chan packet.PacketType } @@ -91,6 +95,24 @@ func (t *Tailer) updateTailPos_nolock(cmdKey CmdKey, reqId string, pos TailPos) t.WatchList[cmdKey] = entry } +func (t *Tailer) removeTailPos_nolock(cmdKey CmdKey, reqId string) { + entry, found := t.WatchList[cmdKey] + if !found { + return + } + entry.removeTailPos(reqId) + if len(entry.Tails) > 0 { + t.WatchList[cmdKey] = entry + return + } + + // delete from watchlist, remove watches + fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, cmdKey.SessionId, cmdKey.CmdId) + delete(t.WatchList, cmdKey) + t.Watcher.Remove(fileNames.PtyOutFile) + t.Watcher.Remove(fileNames.RunnerOutFile) +} + func (t *Tailer) updateEntrySizes_nolock(cmdKey CmdKey, ptyLen int64, runLen int64) { entry, found := t.WatchList[cmdKey] if !found { @@ -124,7 +146,7 @@ func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { ScHomeDir: scHomeDir, SendCh: sendCh, } - rtn.Watcher, err = MakeSessionWatcher() + rtn.Watcher, err = fsnotify.NewWatcher() if err != nil { return nil, err } @@ -213,17 +235,12 @@ func (t *Tailer) runSingleDataTransfer(key CmdKey, reqId string) (*packet.CmdDat func (t *Tailer) checkRemoveNoFollow(cmdKey CmdKey, reqId string) { t.Lock.Lock() defer t.Lock.Unlock() - entry, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) + _, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) if !foundPos { return } if !pos.Follow { - entry.removeTailPos(reqId) - if len(entry.Tails) == 0 { - delete(t.WatchList, cmdKey) - } else { - t.WatchList[cmdKey] = entry - } + t.removeTailPos_nolock(cmdKey, reqId) } } @@ -241,9 +258,12 @@ func (t *Tailer) RunDataTransfer(key CmdKey, reqId string) { } } -// should already hold t.Lock func (t *Tailer) tryStartRun_nolock(entry CmdWatchEntry, pos TailPos) { - if pos.Running || pos.IsCurrent(entry) { + if pos.Running { + return + } + if pos.IsCurrent(entry) { + return } pos.Running = true @@ -251,22 +271,30 @@ func (t *Tailer) tryStartRun_nolock(entry CmdWatchEntry, pos TailPos) { go t.RunDataTransfer(entry.CmdKey, pos.ReqId) } -func (t *Tailer) updateFile(event FileUpdateEvent) { - if event.Err != nil { - t.SendCh <- packet.FmtMessagePacket("error in FileUpdateEvent %s/%s: %v", event.SessionId, event.CmdId, event.Err) +var updateFileRe = regexp.MustCompile("/([a-z0-9-]+)/([a-z0-9-]+)\\.(ptyout|runout)$") + +func (t *Tailer) updateFile(relFileName string) { + m := updateFileRe.FindStringSubmatch(relFileName) + if m == nil { return } - cmdKey := CmdKey{SessionId: event.SessionId, CmdId: event.CmdId} + finfo, err := os.Stat(relFileName) + if err != nil { + t.SendCh <- packet.FmtMessagePacket("error trying to stat file '%s': %v", relFileName, err) + return + } + cmdKey := CmdKey{SessionId: m[1], CmdId: m[2]} t.Lock.Lock() defer t.Lock.Unlock() entry, foundEntry := t.WatchList[cmdKey] if !foundEntry { return } - if event.FileType == FileTypePty { - entry.FilePtyLen = event.Size - } else if event.FileType == FileTypeRun { - entry.FileRunLen = event.Size + fileType := m[3] + if fileType == FileTypePty { + entry.FilePtyLen = finfo.Size() + } else if fileType == FileTypeRun { + entry.FileRunLen = finfo.Size() } t.WatchList[cmdKey] = entry for _, pos := range entry.Tails { @@ -274,14 +302,26 @@ func (t *Tailer) updateFile(event FileUpdateEvent) { } } -func (t *Tailer) Run() error { - go func() { - for event := range t.Watcher.EventCh { - t.updateFile(event) +func (t *Tailer) Run() { + for { + select { + case event, ok := <-t.Watcher.Events: + if !ok { + return + } + if event.Op&fsnotify.Write == fsnotify.Write { + t.updateFile(event.Name) + } + + case err, ok := <-t.Watcher.Errors: + if !ok { + return + } + // what to do with this error? just send a message + t.SendCh <- packet.FmtMessagePacket("error in tailer: %v", err) } - }() - err := t.Watcher.Run(nil) - return err + } + return } func (t *Tailer) Close() error { @@ -307,6 +347,13 @@ func (entry *CmdWatchEntry) fillFilePos(scHomeDir string) { } } +func (t *Tailer) RemoveWatch(pk *packet.UntailCmdPacketType) { + t.Lock.Lock() + defer t.Lock.Unlock() + key := CmdKey{pk.SessionId, pk.CmdId} + t.removeTailPos_nolock(key, pk.ReqId) +} + func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { _, err := uuid.Parse(getPacket.SessionId) if err != nil { @@ -319,19 +366,34 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { if getPacket.ReqId == "" { return fmt.Errorf("getcmd, no reqid specified") } + fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, getPacket.SessionId, getPacket.CmdId) t.Lock.Lock() defer t.Lock.Unlock() key := CmdKey{getPacket.SessionId, getPacket.CmdId} - err = t.Watcher.WatchSession(getPacket.SessionId) - if err != nil { - return fmt.Errorf("error trying to watch sesion '%s': %v", getPacket.SessionId, err) - } entry, foundEntry := t.WatchList[key] if !foundEntry { + // add watches, initialize entry + err = t.Watcher.Add(fileNames.PtyOutFile) + if err != nil { + return err + } + err = t.Watcher.Add(fileNames.RunnerOutFile) + if err != nil { + t.Watcher.Remove(fileNames.PtyOutFile) // best effort clean up + return err + } entry = CmdWatchEntry{CmdKey: key} entry.fillFilePos(t.ScHomeDir) } - pos := TailPos{ReqId: getPacket.ReqId, TailPtyPos: getPacket.PtyPos, TailRunPos: getPacket.RunPos, Follow: getPacket.Tail} + pos, foundPos := entry.getTailPos(getPacket.ReqId) + if !foundPos { + // initialize a new tailpos + pos = TailPos{ReqId: getPacket.ReqId} + } + // update tailpos with new values from getpacket + pos.TailPtyPos = getPacket.PtyPos + pos.TailRunPos = getPacket.RunPos + pos.Follow = getPacket.Tail // convert negative pos to positive if pos.TailPtyPos < 0 { pos.TailPtyPos = max(0, entry.FilePtyLen+pos.TailPtyPos) // + because negative diff --git a/pkg/cmdtail/sessionwatcher.go b/pkg/cmdtail/sessionwatcher.go deleted file mode 100644 index 1e736cb8..00000000 --- a/pkg/cmdtail/sessionwatcher.go +++ /dev/null @@ -1,165 +0,0 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - -package cmdtail - -import ( - "fmt" - "os" - "path" - "regexp" - "sync" - - "github.com/fsnotify/fsnotify" - "github.com/google/uuid" - "github.com/scripthaus-dev/sh2-runner/pkg/base" -) - -const FileTypePty = "ptyout" -const FileTypeRun = "runout" -const eventChSize = 10 - -type FileUpdateEvent struct { - SessionId string - CmdId string - FileType string - Size int64 - Err error -} - -type SessionWatcher struct { - Lock *sync.Mutex - Sessions map[string]bool - ScHomeDir string - Watcher *fsnotify.Watcher - EventCh chan FileUpdateEvent - Err error - Running bool -} - -func MakeSessionWatcher() (*SessionWatcher, error) { - scHomeDir, err := base.GetScHomeDir() - if err != nil { - return nil, err - } - rtn := &SessionWatcher{ - Lock: &sync.Mutex{}, - Sessions: make(map[string]bool), - ScHomeDir: scHomeDir, - EventCh: make(chan FileUpdateEvent, eventChSize), - } - rtn.Watcher, err = fsnotify.NewWatcher() - if err != nil { - return nil, err - } - return rtn, nil -} - -func (w *SessionWatcher) Close() error { - return w.Watcher.Close() -} - -func (w *SessionWatcher) UnWatchSession(sessionId string) error { - _, err := uuid.Parse(sessionId) - if err != nil { - return fmt.Errorf("WatchSession, bad sessionid '%s': %w", sessionId, err) - } - w.Lock.Lock() - defer w.Lock.Unlock() - if !w.Sessions[sessionId] { - return nil - } - sessionDir := path.Join(w.ScHomeDir, base.SessionsDirBaseName, sessionId) - err = w.Watcher.Remove(sessionDir) - if err != nil { - return err - } - w.Sessions[sessionId] = false - return nil -} - -func (w *SessionWatcher) WatchSession(sessionId string) error { - _, err := uuid.Parse(sessionId) - if err != nil { - return fmt.Errorf("WatchSession, bad sessionid '%s': %w", sessionId, err) - } - - w.Lock.Lock() - defer w.Lock.Unlock() - if w.Sessions[sessionId] { - return nil - } - sessionDir := path.Join(w.ScHomeDir, base.SessionsDirBaseName, sessionId) - err = w.Watcher.Add(sessionDir) - if err != nil { - return err - } - w.Sessions[sessionId] = true - return nil -} - -func (w *SessionWatcher) setRunning() bool { - w.Lock.Lock() - defer w.Lock.Unlock() - if w.Running { - return false - } - w.Running = true - return true -} - -var swUpdateFileRe = regexp.MustCompile("/([a-z0-9-]+)/([a-z0-9-]+)\\.(ptyout|runout)$") - -func (w *SessionWatcher) updateFile(relFileName string) { - m := swUpdateFileRe.FindStringSubmatch(relFileName) - if m == nil { - return - } - event := FileUpdateEvent{SessionId: m[1], CmdId: m[2], FileType: m[3]} - finfo, err := os.Stat(relFileName) - if err != nil { - event.Err = err - w.EventCh <- event - return - } - event.Size = finfo.Size() - w.EventCh <- event - return -} - -func (w *SessionWatcher) Run(stopCh chan bool) error { - ok := w.setRunning() - if !ok { - return fmt.Errorf("Cannot run SessionWatcher (alreaady running)") - } - defer func() { - w.Lock.Lock() - defer w.Lock.Unlock() - w.Running = false - close(w.EventCh) - }() - for { - select { - case event, ok := <-w.Watcher.Events: - if !ok { - return nil - } - if (event.Op&fsnotify.Write == fsnotify.Write) || (event.Op&fsnotify.Create == fsnotify.Create) { - w.updateFile(event.Name) - } - - case err, ok := <-w.Watcher.Errors: - if !ok { - return nil - } - return fmt.Errorf("Got error in SessionWatcher: %w", err) - - case <-stopCh: - return nil - } - } - return nil -} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 3b0e5c20..e872e34b 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -7,10 +7,8 @@ package shexec import ( - "errors" "fmt" "io" - "io/fs" "os" "os/exec" "strings" @@ -163,8 +161,12 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT if err != nil { return nil, err } - if _, err = os.Stat(fileNames.PtyOutFile); !errors.Is(err, fs.ErrNotExist) { - return nil, fmt.Errorf("cmdid '%s' was already used", pk.CmdId) + ptyOutInfo, err := os.Stat(fileNames.PtyOutFile) + if err == nil { // non-nil error will be caught by regular OpenFile below + // must have size 0 + if ptyOutInfo.Size() != 0 { + return nil, fmt.Errorf("cmdid '%s' was already used (ptyout len=%d)", pk.CmdId, ptyOutInfo.Size()) + } } cmdPty, cmdTty, err := pty.Open() if err != nil { From 0a6d8b8e9f2a3ea8eab1c44320a74b18a1917e83 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 17 Jun 2022 15:31:07 -0700 Subject: [PATCH 011/149] input packet type --- pkg/packet/packet.go | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index dec40066..bd47abf8 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -33,6 +33,7 @@ const CdPacketStr = "cd" const CdResponseStr = "cdresp" const CmdDataPacketStr = "cmddata" const RawPacketStr = "raw" +const InputPacketStr = "input" var TypeStrToFactory map[string]reflect.Type @@ -53,6 +54,7 @@ func init() { TypeStrToFactory[CdResponseStr] = reflect.TypeOf(CdResponseType{}) TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) + TypeStrToFactory[InputPacketStr] = reflect.TypeOf(InputPacketType{}) } func MakePacket(packetType string) (PacketType, error) { @@ -101,6 +103,27 @@ func MakePingPacket() *PingPacketType { return &PingPacketType{Type: PingPacketStr} } +// InputData gets written to PTY directly +// SigNum gets sent to process via a signal +// WinSize, if set, will run TIOCSWINSZ to set size, and then send SIGWINCH +type InputPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + InputData string `json:"inputdata"` + SigNum int `json:"signum,omitempty"` + WinSizeRows int `json:"winsizerows"` + WinSizeCols int `json:"winsizecols"` +} + +func (*InputPacketType) GetType() string { + return InputPacketStr +} + +func MakeInputPacket() *InputPacketType { + return &InputPacketType{Type: InputPacketStr} +} + type UntailCmdPacketType struct { Type string `json:"type"` ReqId string `json:"reqid"` From 315a048f49a156d6c1c6ece6ab610936de298bd9 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 17 Jun 2022 18:11:49 -0700 Subject: [PATCH 012/149] return username with init packet --- main-runner.go | 4 ++++ pkg/packet/packet.go | 1 + 2 files changed, 5 insertions(+) diff --git a/main-runner.go b/main-runner.go index f8b8dc39..fb7cd6e6 100644 --- a/main-runner.go +++ b/main-runner.go @@ -10,6 +10,7 @@ import ( "fmt" "os" "os/signal" + "os/user" "syscall" "time" @@ -171,6 +172,9 @@ func doMain() { initPacket.Env = os.Environ() initPacket.HomeDir = homeDir initPacket.ScHomeDir = scHomeDir + if user, _ := user.Current(); user != nil { + initPacket.User = user.Username + } sender.SendPacket(initPacket) for pk := range packetCh { if pk.GetType() == packet.PingPacketStr { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index bd47abf8..ea791e48 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -235,6 +235,7 @@ type RunnerInitPacketType struct { ScHomeDir string `json:"schomedir"` HomeDir string `json:"homedir"` Env []string `json:"env"` + User string `json:"user"` } func (*RunnerInitPacketType) GetType() string { From 2c628909124575c1c19a834ac680729dc4ab9ec5 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 20 Jun 2022 17:51:28 -0700 Subject: [PATCH 013/149] call SetWinsize to set terminal size always for pty --- pkg/packet/packet.go | 4 +++- pkg/shexec/shexec.go | 18 ++++++++++++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index ea791e48..2e145f3b 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -297,9 +297,11 @@ type RunPacketType struct { ChDir string `json:"chdir,omitempty"` Env map[string]string `json:"env,omitempty"` Command string `json:"command"` + Rows int `json:"rows"` + Cols int `json:'cols"` } -func (ct *RunPacketType) GetType() string { +func (*RunPacketType) GetType() string { return RunPacketStr } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index e872e34b..cd155702 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -21,6 +21,11 @@ import ( "github.com/scripthaus-dev/sh2-runner/pkg/packet" ) +const DefaultRows = 25 +const DefaultCols = 80 +const MaxRows = 1024 +const MaxCols = 1024 + type ShExecType struct { FileNames *base.CommandFileNames Cmd *exec.Cmd @@ -148,6 +153,18 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { return nil } +func GetWinsize(p *packet.RunPacketType) *pty.Winsize { + rows := DefaultRows + cols := DefaultCols + if p.Rows > 0 && p.Rows <= MaxRows { + rows = p.Rows + } + if p.Cols > 0 && p.Cols <= MaxCols { + cols = p.Cols + } + return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} +} + // when err is nil, the command will have already been started func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { if pk.CmdId == "" { @@ -172,6 +189,7 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT if err != nil { return nil, fmt.Errorf("opening new pty: %w", err) } + pty.Setsize(cmdPty, GetWinsize(pk)) defer func() { cmdTty.Close() }() From 0b172cd689e4254986257c9a381f98331fac0f9e Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Jun 2022 09:56:54 -0700 Subject: [PATCH 014/149] add cdpacket --- main-runner.go | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/main-runner.go b/main-runner.go index fb7cd6e6..64a54abd 100644 --- a/main-runner.go +++ b/main-runner.go @@ -189,9 +189,22 @@ func doMain() { if err != nil { errPk := packet.MakeErrorPacket(err.Error()) sender.SendPacket(errPk) + continue } continue } + if pk.GetType() == packet.CdPacketStr { + cdPacket := pk.(*packet.CdPacketType) + err := os.Chdir(cdPacket.Dir) + resp := packet.MakeResponsePacket(cdPacket.PacketId) + if err != nil { + resp.Error = err.Error() + } else { + resp.Success = true + } + sender.SendPacket(resp) + continue + } if pk.GetType() == packet.ErrorPacketStr { errPk := pk.(*packet.ErrorPacketType) errPk.Error = "invalid packet sent to runner: " + errPk.Error From 766d19f1bc0652437894ffce4403dca5d5fd489f Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Jun 2022 10:16:54 -0700 Subject: [PATCH 015/149] checkpoint, switch name from sh2-runner to mshell --- go.mod | 2 +- main-runner.go => main-mshell.go | 10 +-- pkg/base/base.go | 34 +++----- pkg/cmdtail/cmdtail.go | 4 +- pkg/packet/combined.go | 34 ++++++++ pkg/packet/packet.go | 132 ++++++++++++++++++++----------- pkg/shexec/shexec.go | 31 ++++---- 7 files changed, 152 insertions(+), 95 deletions(-) rename main-runner.go => main-mshell.go (96%) create mode 100644 pkg/packet/combined.go diff --git a/go.mod b/go.mod index e64ed540..e1a2464c 100644 --- a/go.mod +++ b/go.mod @@ -1,4 +1,4 @@ -module github.com/scripthaus-dev/sh2-runner +module github.com/scripthaus-dev/mshell go 1.17 diff --git a/main-runner.go b/main-mshell.go similarity index 96% rename from main-runner.go rename to main-mshell.go index 64a54abd..dd551372 100644 --- a/main-runner.go +++ b/main-mshell.go @@ -15,10 +15,10 @@ import ( "time" "github.com/google/uuid" - "github.com/scripthaus-dev/sh2-runner/pkg/base" - "github.com/scripthaus-dev/sh2-runner/pkg/cmdtail" - "github.com/scripthaus-dev/sh2-runner/pkg/packet" - "github.com/scripthaus-dev/sh2-runner/pkg/shexec" + "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/cmdtail" + "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/shexec" ) // in single run mode, we don't want the runner to die from signals @@ -155,7 +155,7 @@ func doMain() { packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) return } - err = base.EnsureRunnerPath() + err = base.EnsureMShellPath() if err != nil { packet.SendErrorPacket(os.Stdout, err.Error()) return diff --git a/pkg/base/base.go b/pkg/base/base.go index 5da63c01..91f2959b 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -11,11 +11,14 @@ import ( "fmt" "io/fs" "os" + "os/exec" "path" "path/filepath" ) -const ScRunnerVarName = "SCRIPTHAUS_RUNNER" +const DefaultMShellPath = "mshell" +const MShellPathVarName = "MSHELL_PATH" +const SSHCommandVarName = "SSH_COMMAND" const ScHomeVarName = "SCRIPTHAUS_HOME" const HomeVarName = "HOME" const ScShell = "bash" @@ -125,33 +128,20 @@ func EnsureSessionDir(sessionId string) (string, error) { return sdir, nil } -func GetScRunnerPath() (string, error) { - runnerPath := os.Getenv(ScRunnerVarName) - if runnerPath != "" { - return runnerPath, nil +func GetMShellPath() string { + msPath := os.Getenv(MShellPathVarName) + if msPath != "" { + return msPath } - scHome, err := GetScHomeDir() - if err != nil { - return "", err - } - return path.Join(scHome, RunnerBaseName), nil + return DefaultMShellPath } -func EnsureRunnerPath() error { - runnerPath, err := GetScRunnerPath() +func EnsureMShellPath() error { + msPath := GetMShellPath() + _, err := exec.LookPath(msPath) if err != nil { return err } - info, err := os.Stat(runnerPath) - if err != nil { - if errors.Is(err, fs.ErrNotExist) { - return fmt.Errorf("cannot find scripthaus runner at path '%s'", runnerPath) - } - return fmt.Errorf("error stating scripthaus runner at path '%s'", runnerPath) - } - if info.Mode()&0100 == 0 { - return fmt.Errorf("scripthaus runner at path '%s' is not executable mode=%#o", runnerPath, info.Mode()) - } return nil } diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 44b44cc5..51cdbfd8 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -16,8 +16,8 @@ import ( "github.com/fsnotify/fsnotify" "github.com/google/uuid" - "github.com/scripthaus-dev/sh2-runner/pkg/base" - "github.com/scripthaus-dev/sh2-runner/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/packet" ) const MaxDataBytes = 4096 diff --git a/pkg/packet/combined.go b/pkg/packet/combined.go new file mode 100644 index 00000000..0f15fd93 --- /dev/null +++ b/pkg/packet/combined.go @@ -0,0 +1,34 @@ +package packet + +type CombinedPacket struct { + Type string `json:"type"` + Success bool `json:"success"` + Ts int64 `json:"ts"` + Id string `json:"id,omitempty"` + + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + + PtyPos int64 `json:"ptypos"` + PtyLen int64 `json:"ptylen"` + RunPos int64 `json:"runpos"` + RunLen int64 `json:"runlen"` + + Error string `json:"error"` + NotFound bool `json:"notfound,omitempty"` + Tail bool `json:"tail,omitempty"` + Dir string `json:"dir"` + ChDir string `json:"chdir,omitempty"` + + Data string `json:"data"` + PtyData string `json:"ptydata"` + RunData string `json:"rundata"` + Message string `json:"message"` + Command string `json:"command"` + + ScHomeDir string `json:"schomedir"` + HomeDir string `json:"homedir"` + Env []string `json:"env"` + ExitCode int `json:"exitcode"` + RunnerPid int `json:"runnerpid"` +} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 2e145f3b..863fc418 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -18,22 +18,29 @@ import ( "sync" ) -const RunPacketStr = "run" -const PingPacketStr = "ping" -const DonePacketStr = "done" -const ErrorPacketStr = "error" -const MessagePacketStr = "message" -const CmdStartPacketStr = "cmdstart" -const CmdDonePacketStr = "cmddone" -const ListCmdPacketStr = "lscmd" -const GetCmdPacketStr = "getcmd" -const UntailCmdPacketStr = "untailcmd" -const RunnerInitPacketStr = "runnerinit" -const CdPacketStr = "cd" -const CdResponseStr = "cdresp" -const CmdDataPacketStr = "cmddata" -const RawPacketStr = "raw" -const InputPacketStr = "input" +// remote: runnerinit, run, ping, data, cmdstart, cmddone +// remote(detached): runnerinit, run, cmdstart +// server: runnerinit, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] +// all: error, message + +const ( + RunPacketStr = "run" + PingPacketStr = "ping" + RunnerInitPacketStr = "runnerinit" + DataPacketStr = "data" + CmdStartPacketStr = "cmdstart" + CmdDonePacketStr = "cmddone" + ResponsePacketStr = "resp" + DonePacketStr = "done" + ErrorPacketStr = "error" + MessagePacketStr = "message" + GetCmdPacketStr = "getcmd" + UntailCmdPacketStr = "untailcmd" + CdPacketStr = "cd" + CmdDataPacketStr = "cmddata" + RawPacketStr = "raw" + InputPacketStr = "input" +) var TypeStrToFactory map[string]reflect.Type @@ -41,20 +48,20 @@ func init() { TypeStrToFactory = make(map[string]reflect.Type) TypeStrToFactory[RunPacketStr] = reflect.TypeOf(RunPacketType{}) TypeStrToFactory[PingPacketStr] = reflect.TypeOf(PingPacketType{}) + TypeStrToFactory[ResponsePacketStr] = reflect.TypeOf(ResponsePacketType{}) TypeStrToFactory[DonePacketStr] = reflect.TypeOf(DonePacketType{}) TypeStrToFactory[ErrorPacketStr] = reflect.TypeOf(ErrorPacketType{}) TypeStrToFactory[MessagePacketStr] = reflect.TypeOf(MessagePacketType{}) TypeStrToFactory[CmdStartPacketStr] = reflect.TypeOf(CmdStartPacketType{}) TypeStrToFactory[CmdDonePacketStr] = reflect.TypeOf(CmdDonePacketType{}) - TypeStrToFactory[ListCmdPacketStr] = reflect.TypeOf(ListCmdPacketType{}) TypeStrToFactory[GetCmdPacketStr] = reflect.TypeOf(GetCmdPacketType{}) TypeStrToFactory[UntailCmdPacketStr] = reflect.TypeOf(UntailCmdPacketType{}) TypeStrToFactory[RunnerInitPacketStr] = reflect.TypeOf(RunnerInitPacketType{}) TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) - TypeStrToFactory[CdResponseStr] = reflect.TypeOf(CdResponseType{}) TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) TypeStrToFactory[InputPacketStr] = reflect.TypeOf(InputPacketType{}) + TypeStrToFactory[DataPacketStr] = reflect.TypeOf(DataPacketType{}) } func MakePacket(packetType string) (PacketType, error) { @@ -103,6 +110,22 @@ func MakePingPacket() *PingPacketType { return &PingPacketType{Type: PingPacketStr} } +type DataPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + FdNum int `json:"fdnum"` + Data string `json:"data"` +} + +func (*DataPacketType) GetType() string { + return DataPacketStr +} + +func MakeDataPacket(fdNum int, data string) *DataPacketType { + return &DataPacketType{Type: DataPacketStr, FdNum: fdNum, Data: data} +} + // InputData gets written to PTY directly // SigNum gets sent to process via a signal // WinSize, if set, will run TIOCSWINSZ to set size, and then send SIGWINCH @@ -157,19 +180,6 @@ func MakeGetCmdPacket() *GetCmdPacketType { return &GetCmdPacketType{Type: GetCmdPacketStr} } -type ListCmdPacketType struct { - Type string `json:"type"` - SessionId string `json:"sessionid"` -} - -func (*ListCmdPacketType) GetType() string { - return ListCmdPacketStr -} - -func MakeListCmdPacket(sessionId string) *ListCmdPacketType { - return &ListCmdPacketType{Type: ListCmdPacketStr, SessionId: sessionId} -} - type CdPacketType struct { Type string `json:"type"` PacketId string `json:"packetid"` @@ -180,23 +190,32 @@ func (*CdPacketType) GetType() string { return CdPacketStr } +func (p *CdPacketType) GetPacketId() string { + return p.PacketId +} + func MakeCdPacket() *CdPacketType { return &CdPacketType{Type: CdPacketStr} } -type CdResponseType struct { - Type string `json:"type"` - PacketId string `json:"packetid"` - Success bool `json:"success"` - Error string `json:"error"` +type ResponsePacketType struct { + Type string `json:"type"` + PacketId string `json:"packetid"` + Success bool `json:"success"` + Error string `json:"error"` + Data interface{} `json:"data"` } -func (*CdResponseType) GetType() string { - return CdResponseStr +func (*ResponsePacketType) GetType() string { + return ResponsePacketStr } -func MakeCdResponse() *CdResponseType { - return &CdResponseType{Type: CdResponseStr} +func (p *ResponsePacketType) GetPacketId() string { + return p.PacketId +} + +func MakeResponsePacket(packetId string) *ResponsePacketType { + return &ResponsePacketType{Type: ResponsePacketStr, PacketId: packetId} } type RawPacketType struct { @@ -232,10 +251,11 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { type RunnerInitPacketType struct { Type string `json:"type"` - ScHomeDir string `json:"schomedir"` - HomeDir string `json:"homedir"` - Env []string `json:"env"` - User string `json:"user"` + Version string `json:"version"` + ScHomeDir string `json:"schomedir,omitempty"` + HomeDir string `json:"homedir,omitempty"` + Env []string `json:"env,omitempty"` + User string `json:"user,omitempty"` } func (*RunnerInitPacketType) GetType() string { @@ -290,15 +310,26 @@ func MakeCmdStartPacket() *CmdStartPacketType { return &CmdStartPacketType{Type: CmdStartPacketStr} } +type TermSize struct { + Rows int `json:"rows"` + Cols int `json:"cols"` +} + +type RemoteFd struct { + FdNum int `json:"fdnum"` + Read bool `json:"read"` + Write bool `json:"write"` +} + type RunPacketType struct { Type string `json:"type"` SessionId string `json:"sessionid"` CmdId string `json:"cmdid"` - ChDir string `json:"chdir,omitempty"` - Env map[string]string `json:"env,omitempty"` Command string `json:"command"` - Rows int `json:"rows"` - Cols int `json:'cols"` + Cwd string `json:"cwd,omitempty"` + Env map[string]string `json:"env,omitempty"` + TermSize TermSize `json:"termsize"` + Fds []RemoteFd `json:"fds"` } func (*RunPacketType) GetType() string { @@ -335,6 +366,11 @@ type PacketType interface { GetType() string } +type RpcPacketType interface { + GetType() string + GetPacketId() string +} + func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { var bareCmd BarePacketType err := json.Unmarshal(jsonBuf, &bareCmd) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index cd155702..600dd379 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -17,8 +17,8 @@ import ( "github.com/creack/pty" "github.com/google/uuid" - "github.com/scripthaus-dev/sh2-runner/pkg/base" - "github.com/scripthaus-dev/sh2-runner/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/packet" ) const DefaultRows = 25 @@ -79,8 +79,8 @@ func UpdateCmdEnv(cmd *exec.Cmd, envVars map[string]string) { func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { ecmd := exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(ecmd, pk.Env) - if pk.ChDir != "" { - ecmd.Dir = pk.ChDir + if pk.Cwd != "" { + ecmd.Dir = pk.Cwd } ecmd.Stdin = cmdTty ecmd.Stdout = cmdTty @@ -93,11 +93,8 @@ func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { } func MakeRunnerExec(cmdId string) (*exec.Cmd, error) { - runnerPath, err := base.GetScRunnerPath() - if err != nil { - return nil, err - } - ecmd := exec.Command(runnerPath, cmdId) + msPath := base.GetMShellPath() + ecmd := exec.Command(msPath, cmdId) return ecmd, nil } @@ -141,13 +138,13 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { if err != nil { return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) } - if pk.ChDir != "" { - dirInfo, err := os.Stat(pk.ChDir) + if pk.Cwd != "" { + dirInfo, err := os.Stat(pk.Cwd) if err != nil { - return fmt.Errorf("invalid cwd '%s' for command: %v", pk.ChDir, err) + return fmt.Errorf("invalid cwd '%s' for command: %v", pk.Cwd, err) } if !dirInfo.IsDir() { - return fmt.Errorf("invalid cwd '%s' for command, not a directory", pk.ChDir) + return fmt.Errorf("invalid cwd '%s' for command, not a directory", pk.Cwd) } } return nil @@ -156,11 +153,11 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { func GetWinsize(p *packet.RunPacketType) *pty.Winsize { rows := DefaultRows cols := DefaultCols - if p.Rows > 0 && p.Rows <= MaxRows { - rows = p.Rows + if p.TermSize.Rows > 0 && p.TermSize.Rows <= MaxRows { + rows = p.TermSize.Rows } - if p.Cols > 0 && p.Cols <= MaxCols { - cols = p.Cols + if p.TermSize.Cols > 0 && p.TermSize.Cols <= MaxCols { + cols = p.TermSize.Cols } return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} } From c43d3ecc85bedf08dd9c7ff8dc6cf91edf800348 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Jun 2022 12:48:45 -0700 Subject: [PATCH 016/149] checkpoint got stdout/stderr data packets working with new remote handler --- main-mshell.go | 116 ++++++++++++++++++++--- pkg/base/base.go | 17 ++-- pkg/packet/packet.go | 77 ++++++++-------- pkg/shexec/shexec.go | 212 ++++++++++++++++++++++++++++++++++++------- 4 files changed, 329 insertions(+), 93 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index dd551372..0bdb4e77 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -11,6 +11,7 @@ import ( "os" "os/signal" "os/user" + "strings" "syscall" "time" @@ -21,8 +22,10 @@ import ( "github.com/scripthaus-dev/mshell/pkg/shexec" ) -// in single run mode, we don't want the runner to die from signals -// since we want the single runner to persist even if session / main runner +const MShellVersion = "0.1.0" + +// in single run mode, we don't want mshell to die from signals +// since we want the single mshell to persist even if session / main mshell // is terminated. func setupSingleSignals(cmd *shexec.ShExecType) { sigCh := make(chan os.Signal, 1) @@ -46,7 +49,7 @@ func doSingle(cmdId string) { runPacket, _ = pk.(*packet.RunPacketType) break } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) return } if runPacket == nil { @@ -66,13 +69,9 @@ func doSingle(cmdId string) { return } setupSingleSignals(cmd) - startPacket := packet.MakeCmdStartPacket() - startPacket.Ts = time.Now().UnixMilli() - startPacket.CmdId = runPacket.CmdId - startPacket.Pid = cmd.Cmd.Process.Pid - startPacket.RunnerPid = os.Getpid() + startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) - donePacket := cmd.WaitForCommand(runPacket.CmdId) + donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) sender.CloseSendCh() sender.WaitForDone() @@ -94,7 +93,7 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { } cmd, err := shexec.MakeRunnerExec(pk.CmdId) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot make runner command: %v", err))) + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot make mshell command: %v", err))) return } cmdStdin, err := cmd.StdinPipe() @@ -155,7 +154,7 @@ func doMain() { packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) return } - err = base.EnsureMShellPath() + _, err = base.GetMShellPath() if err != nil { packet.SendErrorPacket(os.Stdout, err.Error()) return @@ -168,7 +167,7 @@ func doMain() { return } go tailer.Run() - initPacket := packet.MakeRunnerInitPacket() + initPacket := packet.MakeInitPacket() initPacket.Env = os.Environ() initPacket.HomeDir = homeDir initPacket.ScHomeDir = scHomeDir @@ -207,19 +206,106 @@ func doMain() { } if pk.GetType() == packet.ErrorPacketStr { errPk := pk.(*packet.ErrorPacketType) - errPk.Error = "invalid packet sent to runner: " + errPk.Error + errPk.Error = "invalid packet sent to mshell: " + errPk.Error sender.SendPacket(errPk) continue } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) } } +func handleRemote() { + packetCh := packet.PacketParser(os.Stdin) + sender := packet.MakePacketSender(os.Stdout) + defer func() { + // wait for sender to complete + close(sender.SendCh) + <-sender.DoneCh + }() + initPacket := packet.MakeInitPacket() + initPacket.Version = MShellVersion + sender.SendPacket(initPacket) + var runPacket *packet.RunPacketType + for pk := range packetCh { + if pk.GetType() == packet.PingPacketStr { + continue + } + if pk.GetType() == packet.RunPacketStr { + runPacket, _ = pk.(*packet.RunPacketType) + break + } + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) + return + } + cmd, err := shexec.RunCommand(runPacket, sender) + if err != nil { + sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) + return + } + defer cmd.Close() + startPacket := cmd.MakeCmdStartPacket() + sender.SendPacket(startPacket) + cmd.RunIOAndWait(sender) +} + +func handleServer() { +} + +func handleClient() { + fmt.Printf("mshell client\n") +} + +func handleUsage(extended bool) { + usage := ` +Client Usage: mshell [mshell-opts] [ssh-opts] user@host [command] + +mshell multiplexes input and output streams to a remote command over ssh. + +Options: + --env 'X=Y,A=B' - set remote environment variables for command, comma or newline separated + --env-file [file] - load environment variables from [file] (.env format) + --env-copy [glob] - copy local environment variables to remote using [glob] pattern + --cwd [dir] - execute remote command in [dir] + --no-auto-fds - do not auto-detect additional fds + --fds [fdspec] - open fds based off [fdspec], comma separated (implies --no-auto-fds) + <[num] opens for reading + >[num] opens for writing + <>[num] opens for read/write + e.g. --fds '<5,>6,<>7' + +mshell is licensed under the MPLv2 +Please see https://github.com/scripthaus-dev/mshell for extended usage modes, source code, bugs, and feature requests +` + fmt.Printf("%s\n\n", strings.TrimSpace(usage)) +} + func main() { + if len(os.Args) == 1 { + handleUsage(false) + return + } + firstArg := os.Args[1] + if firstArg == "--help" { + handleUsage(true) + return + } else if firstArg == "--version" { + fmt.Printf("mshell v%s\n", MShellVersion) + return + } else if firstArg == "--remote" { + handleRemote() + return + } else if firstArg == "--server" { + handleServer() + return + } else { + handleClient() + return + } + if len(os.Args) >= 2 { cmdId, err := uuid.Parse(os.Args[1]) if err != nil { - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid non-cmdid passed to runner", err)) + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid non-cmdid passed to mshell", err)) return } doSingle(cmdId.String()) diff --git a/pkg/base/base.go b/pkg/base/base.go index 91f2959b..dc13157f 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -17,6 +17,7 @@ import ( ) const DefaultMShellPath = "mshell" +const DefaultUserMShellPath = ".mshell/mshell" const MShellPathVarName = "MSHELL_PATH" const SSHCommandVarName = "SSH_COMMAND" const ScHomeVarName = "SCRIPTHAUS_HOME" @@ -128,21 +129,17 @@ func EnsureSessionDir(sessionId string) (string, error) { return sdir, nil } -func GetMShellPath() string { +func GetMShellPath() (string, error) { msPath := os.Getenv(MShellPathVarName) if msPath != "" { - return msPath + return exec.LookPath(msPath) } - return DefaultMShellPath -} - -func EnsureMShellPath() error { - msPath := GetMShellPath() - _, err := exec.LookPath(msPath) + userMShellPath := path.Join(GetHomeDir(), DefaultUserMShellPath) + msPath, err := exec.LookPath(userMShellPath) if err != nil { - return err + return msPath, nil } - return nil + return exec.LookPath(DefaultMShellPath) } func GetScSessionsDir() (string, error) { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 863fc418..33989978 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -18,28 +18,28 @@ import ( "sync" ) -// remote: runnerinit, run, ping, data, cmdstart, cmddone -// remote(detached): runnerinit, run, cmdstart -// server: runnerinit, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] +// remote: init, run, ping, data, cmdstart, cmddone +// remote(detached): init, run, cmdstart +// server: init, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] // all: error, message const ( - RunPacketStr = "run" - PingPacketStr = "ping" - RunnerInitPacketStr = "runnerinit" - DataPacketStr = "data" - CmdStartPacketStr = "cmdstart" - CmdDonePacketStr = "cmddone" - ResponsePacketStr = "resp" - DonePacketStr = "done" - ErrorPacketStr = "error" - MessagePacketStr = "message" - GetCmdPacketStr = "getcmd" - UntailCmdPacketStr = "untailcmd" - CdPacketStr = "cd" - CmdDataPacketStr = "cmddata" - RawPacketStr = "raw" - InputPacketStr = "input" + RunPacketStr = "run" + PingPacketStr = "ping" + InitPacketStr = "init" + DataPacketStr = "data" + CmdStartPacketStr = "cmdstart" + CmdDonePacketStr = "cmddone" + ResponsePacketStr = "resp" + DonePacketStr = "done" + ErrorPacketStr = "error" + MessagePacketStr = "message" + GetCmdPacketStr = "getcmd" + UntailCmdPacketStr = "untailcmd" + CdPacketStr = "cd" + CmdDataPacketStr = "cmddata" + RawPacketStr = "raw" + InputPacketStr = "input" ) var TypeStrToFactory map[string]reflect.Type @@ -56,7 +56,7 @@ func init() { TypeStrToFactory[CmdDonePacketStr] = reflect.TypeOf(CmdDonePacketType{}) TypeStrToFactory[GetCmdPacketStr] = reflect.TypeOf(GetCmdPacketType{}) TypeStrToFactory[UntailCmdPacketStr] = reflect.TypeOf(UntailCmdPacketType{}) - TypeStrToFactory[RunnerInitPacketStr] = reflect.TypeOf(RunnerInitPacketType{}) + TypeStrToFactory[InitPacketStr] = reflect.TypeOf(InitPacketType{}) TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) @@ -112,18 +112,20 @@ func MakePingPacket() *PingPacketType { type DataPacketType struct { Type string `json:"type"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` FdNum int `json:"fdnum"` Data string `json:"data"` + Eof bool `json:"eof,omitempty"` + Error string `json:"error,omitempty"` } func (*DataPacketType) GetType() string { return DataPacketStr } -func MakeDataPacket(fdNum int, data string) *DataPacketType { - return &DataPacketType{Type: DataPacketStr, FdNum: fdNum, Data: data} +func MakeDataPacket() *DataPacketType { + return &DataPacketType{Type: DataPacketStr} } // InputData gets written to PTY directly @@ -249,7 +251,7 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { return &MessagePacketType{Type: MessagePacketStr, Message: message} } -type RunnerInitPacketType struct { +type InitPacketType struct { Type string `json:"type"` Version string `json:"version"` ScHomeDir string `json:"schomedir,omitempty"` @@ -258,12 +260,12 @@ type RunnerInitPacketType struct { User string `json:"user,omitempty"` } -func (*RunnerInitPacketType) GetType() string { - return RunnerInitPacketStr +func (*InitPacketType) GetType() string { + return InitPacketStr } -func MakeRunnerInitPacket() *RunnerInitPacketType { - return &RunnerInitPacketType{Type: RunnerInitPacketStr} +func MakeInitPacket() *InitPacketType { + return &InitPacketType{Type: InitPacketStr} } type DonePacketType struct { @@ -281,7 +283,8 @@ func MakeDonePacket() *DonePacketType { type CmdDonePacketType struct { Type string `json:"type"` Ts int64 `json:"ts"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` ExitCode int `json:"exitcode"` DurationMs int64 `json:"durationms"` } @@ -297,9 +300,10 @@ func MakeCmdDonePacket() *CmdDonePacketType { type CmdStartPacketType struct { Type string `json:"type"` Ts int64 `json:"ts"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` Pid int `json:"pid"` - RunnerPid int `json:"runnerpid"` + MShellPid int `json:"mshellpid"` } func (*CmdStartPacketType) GetType() string { @@ -323,13 +327,14 @@ type RemoteFd struct { type RunPacketType struct { Type string `json:"type"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` Command string `json:"command"` Cwd string `json:"cwd,omitempty"` Env map[string]string `json:"env,omitempty"` - TermSize TermSize `json:"termsize"` - Fds []RemoteFd `json:"fds"` + TermSize TermSize `json:"termsize,omitempty"` + Fds []RemoteFd `json:"fds,omitempty"` + Detached bool `json:"detached,omitempty"` } func (*RunPacketType) GetType() string { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 600dd379..e2cd6a8e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -12,6 +12,7 @@ import ( "os" "os/exec" "strings" + "sync" "syscall" "time" @@ -27,14 +28,48 @@ const MaxRows = 1024 const MaxCols = 1024 type ShExecType struct { - FileNames *base.CommandFileNames - Cmd *exec.Cmd - CmdPty *os.File - StartTs time.Time + StartTs time.Time + RunPacket *packet.RunPacketType + FileNames *base.CommandFileNames + Cmd *exec.Cmd + CmdPty *os.File + FdReaders map[int]*os.File + FdWriters map[int]*os.File + CloseAfterStart []*os.File +} + +func MakeShExec(pk *packet.RunPacketType) *ShExecType { + return &ShExecType{ + StartTs: time.Now(), + RunPacket: pk, + FdReaders: make(map[int]*os.File), + FdWriters: make(map[int]*os.File), + } } func (c *ShExecType) Close() { - c.CmdPty.Close() + if c.CmdPty != nil { + c.CmdPty.Close() + } + for _, fd := range c.FdReaders { + fd.Close() + } + for _, fd := range c.FdWriters { + fd.Close() + } + for _, fd := range c.CloseAfterStart { + fd.Close() + } +} + +func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { + startPacket := packet.MakeCmdStartPacket() + startPacket.Ts = time.Now().UnixMilli() + startPacket.SessionId = c.RunPacket.SessionId + startPacket.CmdId = c.RunPacket.CmdId + startPacket.Pid = c.Cmd.Process.Pid + startPacket.MShellPid = os.Getpid() + return startPacket } func getEnvStrKey(envStr string) string { @@ -93,7 +128,10 @@ func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { } func MakeRunnerExec(cmdId string) (*exec.Cmd, error) { - msPath := base.GetMShellPath() + msPath, err := base.GetMShellPath() + if err != nil { + return nil, err + } ecmd := exec.Command(msPath, cmdId) return ecmd, nil } @@ -124,19 +162,21 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { if pk.Type != packet.RunPacketStr { return fmt.Errorf("run packet has wrong type: %s", pk.Type) } - if pk.SessionId == "" { - return fmt.Errorf("run packet does not have sessionid") - } - _, err := uuid.Parse(pk.SessionId) - if err != nil { - return fmt.Errorf("invalid sessionid '%s' for command", pk.SessionId) - } - if pk.CmdId == "" { - return fmt.Errorf("run packet does not have cmdid") - } - _, err = uuid.Parse(pk.CmdId) - if err != nil { - return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) + if pk.Detached { + if pk.SessionId == "" { + return fmt.Errorf("run packet does not have sessionid") + } + _, err := uuid.Parse(pk.SessionId) + if err != nil { + return fmt.Errorf("invalid sessionid '%s' for command", pk.SessionId) + } + if pk.CmdId == "" { + return fmt.Errorf("run packet does not have cmdid") + } + _, err = uuid.Parse(pk.CmdId) + if err != nil { + return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) + } } if pk.Cwd != "" { dirInfo, err := os.Stat(pk.Cwd) @@ -164,13 +204,120 @@ func GetWinsize(p *packet.RunPacketType) *pty.Winsize { // when err is nil, the command will have already been started func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - if pk.CmdId == "" { - pk.CmdId = uuid.New().String() - } err := ValidateRunPacket(pk) if err != nil { return nil, err } + if !pk.Detached { + return runCommandSimple(pk, sender) + } else { + return runCommandDetached(pk, sender) + } +} + +// returns the *writer* to connect to process, reader is put in FdReaders +func (cmd *ShExecType) makeReaderPipe(fdNum int) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + cmd.FdReaders[fdNum] = pr + cmd.CloseAfterStart = append(cmd.CloseAfterStart, pw) + return pw, nil +} + +// returns the *reader* to connect to process, writer is put in FdWriters +func (cmd *ShExecType) makeWriterPipe(fdNum int) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + cmd.FdWriters[fdNum] = pw + cmd.CloseAfterStart = append(cmd.CloseAfterStart, pr) + return pr, nil +} + +func (cmd *ShExecType) MakeDataPacket(fdNum int, data []byte) *packet.DataPacketType { + pk := packet.MakeDataPacket() + pk.SessionId = cmd.RunPacket.SessionId + pk.CmdId = cmd.RunPacket.CmdId + pk.FdNum = fdNum + pk.Data = string(data) + return pk +} + +func (cmd *ShExecType) runReadLoop(wg *sync.WaitGroup, fdNum int, fd *os.File, sender *packet.PacketSender) { + go func() { + defer fd.Close() + defer wg.Done() + buf := make([]byte, 4096) + for { + nr, err := fd.Read(buf) + pk := cmd.MakeDataPacket(fdNum, buf[0:nr]) + if err == io.EOF { + pk.Eof = true + sender.SendPacket(pk) + break + } else if err != nil { + pk.Error = err.Error() + sender.SendPacket(pk) + break + } else { + sender.SendPacket(pk) + } + } + }() +} + +func (cmd *ShExecType) RunIOAndWait(sender *packet.PacketSender) { + var wg sync.WaitGroup + wg.Add(len(cmd.FdReaders)) + go func() { + for fdNum, fd := range cmd.FdReaders { + cmd.runReadLoop(&wg, fdNum, fd, sender) + } + }() + donePacket := cmd.WaitForCommand() + wg.Wait() + sender.SendPacket(donePacket) +} + +func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { + cmd := MakeShExec(pk) + cmd.Cmd = exec.Command("bash", "-c", pk.Command) + UpdateCmdEnv(cmd.Cmd, pk.Env) + if pk.Cwd != "" { + cmd.Cmd.Dir = pk.Cwd + } + var err error + cmd.Cmd.Stdin, err = cmd.makeWriterPipe(0) + if err != nil { + cmd.Close() + return nil, err + } + cmd.Cmd.Stdout, err = cmd.makeReaderPipe(1) + if err != nil { + cmd.Close() + return nil, err + } + cmd.Cmd.Stderr, err = cmd.makeReaderPipe(2) + if err != nil { + cmd.Close() + return nil, err + } + err = cmd.Cmd.Start() + if err != nil { + cmd.Close() + return nil, err + } + for _, fd := range cmd.CloseAfterStart { + fd.Close() + } + cmd.CloseAfterStart = nil + return cmd, nil +} + +func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { fileNames, err := base.GetCommandFileNames(pk.SessionId, pk.CmdId) if err != nil { return nil, err @@ -190,7 +337,7 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT defer func() { cmdTty.Close() }() - startTs := time.Now() + rtn := MakeShExec(pk) ecmd := MakeExecCmd(pk, cmdTty) err = ecmd.Start() if err != nil { @@ -214,12 +361,10 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT sender.SendErrorPacket(fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) } }() - return &ShExecType{ - FileNames: fileNames, - Cmd: ecmd, - CmdPty: cmdPty, - StartTs: startTs, - }, nil + rtn.FileNames = fileNames + rtn.Cmd = ecmd + rtn.CmdPty = cmdPty + return rtn, nil } func GetExitCode(err error) int { @@ -233,16 +378,19 @@ func GetExitCode(err error) int { } } -func (c *ShExecType) WaitForCommand(cmdId string) *packet.CmdDonePacketType { +func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { exitErr := c.Cmd.Wait() endTs := time.Now() cmdDuration := endTs.Sub(c.StartTs) exitCode := GetExitCode(exitErr) donePacket := packet.MakeCmdDonePacket() donePacket.Ts = endTs.UnixMilli() - donePacket.CmdId = cmdId + donePacket.SessionId = c.RunPacket.SessionId + donePacket.CmdId = c.RunPacket.CmdId donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) - os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) + if c.FileNames != nil { + os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) + } return donePacket } From 29372be4efa48b94e32ce2d83387c38991a3f2a5 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Jun 2022 17:37:05 -0700 Subject: [PATCH 017/149] working with fdreaders and fdwriters to properly buffer output and not exceed buffer size without acks --- main-mshell.go | 2 +- pkg/packet/packet.go | 19 ++++++ pkg/shexec/bufreader.go | 125 ++++++++++++++++++++++++++++++++++++++++ pkg/shexec/bufwriter.go | 115 ++++++++++++++++++++++++++++++++++++ pkg/shexec/shexec.go | 125 ++++++++++++++++++++++++++-------------- 5 files changed, 342 insertions(+), 44 deletions(-) create mode 100644 pkg/shexec/bufreader.go create mode 100644 pkg/shexec/bufwriter.go diff --git a/main-mshell.go b/main-mshell.go index 0bdb4e77..e5120fd8 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -245,7 +245,7 @@ func handleRemote() { defer cmd.Close() startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) - cmd.RunIOAndWait(sender) + cmd.RunIOAndWait(packetCh, sender) } func handleServer() { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 33989978..fc9fcaa4 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -28,6 +28,7 @@ const ( PingPacketStr = "ping" InitPacketStr = "init" DataPacketStr = "data" + DataAckPacketStr = "dataack" CmdStartPacketStr = "cmdstart" CmdDonePacketStr = "cmddone" ResponsePacketStr = "resp" @@ -62,6 +63,7 @@ func init() { TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) TypeStrToFactory[InputPacketStr] = reflect.TypeOf(InputPacketType{}) TypeStrToFactory[DataPacketStr] = reflect.TypeOf(DataPacketType{}) + TypeStrToFactory[DataAckPacketStr] = reflect.TypeOf(DataAckPacketType{}) } func MakePacket(packetType string) (PacketType, error) { @@ -128,6 +130,23 @@ func MakeDataPacket() *DataPacketType { return &DataPacketType{Type: DataPacketStr} } +type DataAckPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` + FdNum int `json:"fdnum"` + AckLen int `json:"acklen"` + Error string `json:"error"` +} + +func (*DataAckPacketType) GetType() string { + return DataAckPacketStr +} + +func MakeDataAckPacket() *DataAckPacketType { + return &DataAckPacketType{Type: DataAckPacketStr} +} + // InputData gets written to PTY directly // SigNum gets sent to process via a signal // WinSize, if set, will run TIOCSWINSZ to set size, and then send SIGWINCH diff --git a/pkg/shexec/bufreader.go b/pkg/shexec/bufreader.go new file mode 100644 index 00000000..73aefc7b --- /dev/null +++ b/pkg/shexec/bufreader.go @@ -0,0 +1,125 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package shexec + +import ( + "io" + "os" + "sync" + + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +type FdReader struct { + CVar *sync.Cond + SessionId string + CmdId string + FdNum int + Fd *os.File + BufSize int + Closed bool +} + +func MakeFdReader(c *ShExecType, fd *os.File, fdNum int) *FdReader { + return &FdReader{ + CVar: sync.NewCond(&sync.Mutex{}), + SessionId: c.RunPacket.SessionId, + CmdId: c.RunPacket.CmdId, + FdNum: fdNum, + Fd: fd, + BufSize: 0, + } +} + +func (r *FdReader) Close() { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + if r.Closed { + return + } + if r.Fd != nil { + r.Fd.Close() + } + r.CVar.Broadcast() +} + +func (r *FdReader) NotifyAck(ackLen int) { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + r.BufSize -= ackLen + if r.BufSize < 0 { + r.BufSize = 0 + } + r.CVar.Broadcast() +} + +// returns (success) +func (r *FdReader) WriteWait(sender *packet.PacketSender, data []byte, isEof bool) bool { + if len(data) == 0 { + return true + } + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + for { + bufAvail := ReadBufSize - r.BufSize + if r.Closed { + return false + } + if bufAvail == 0 { + r.CVar.Wait() + continue + } + writeLen := min(bufAvail, len(data)) + pk := r.MakeDataPacket(data[0:writeLen], nil) + sender.SendPacket(pk) + r.BufSize += writeLen + data = data[writeLen:] + if len(data) == 0 { + return true + } + r.CVar.Wait() + } +} + +func min(v1 int, v2 int) int { + if v1 <= v2 { + return v1 + } + return v2 +} + +func (r *FdReader) MakeDataPacket(data []byte, err error) *packet.DataPacketType { + pk := packet.MakeDataPacket() + pk.SessionId = r.SessionId + pk.CmdId = r.CmdId + pk.FdNum = r.FdNum + pk.Data = string(data) + if err != nil { + pk.Error = err.Error() + } + return pk +} + +func (r *FdReader) ReadLoop(wg *sync.WaitGroup, sender *packet.PacketSender) { + defer r.Close() + defer wg.Done() + buf := make([]byte, 4096) + for { + nr, err := r.Fd.Read(buf) + if nr > 0 || err == io.EOF { + isOpen := r.WriteWait(sender, buf[0:nr], (err == io.EOF)) + if !isOpen { + return + } + } + if err != nil { + errPk := r.MakeDataPacket(nil, err) + sender.SendPacket(errPk) + return + } + } +} diff --git a/pkg/shexec/bufwriter.go b/pkg/shexec/bufwriter.go new file mode 100644 index 00000000..f51a0604 --- /dev/null +++ b/pkg/shexec/bufwriter.go @@ -0,0 +1,115 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package shexec + +import ( + "fmt" + "os" + "sync" + + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +type FdWriter struct { + CVar *sync.Cond + SessionId string + CmdId string + FdNum int + Buffer []byte + Fd *os.File + Eof bool + Closed bool +} + +func MakeFdWriter(c *ShExecType, fd *os.File, fdNum int) *FdWriter { + return &FdWriter{ + CVar: sync.NewCond(&sync.Mutex{}), + Fd: fd, + SessionId: c.RunPacket.SessionId, + CmdId: c.RunPacket.CmdId, + FdNum: fdNum, + } +} + +func (w *FdWriter) Close() { + w.CVar.L.Lock() + defer w.CVar.L.Unlock() + if w.Closed { + return + } + w.Closed = true + if w.Fd != nil { + w.Fd.Close() + } + w.Buffer = nil + w.CVar.Broadcast() +} + +func (w *FdWriter) WaitForData() ([]byte, bool) { + w.CVar.L.Lock() + defer w.CVar.L.Unlock() + for { + if len(w.Buffer) > 0 || w.Eof || w.Closed { + toWrite := w.Buffer + w.Buffer = nil + return toWrite, w.Eof + } + w.CVar.Wait() + } +} + +func (w *FdWriter) MakeDataAckPacket(ackLen int, err error) *packet.DataAckPacketType { + ack := packet.MakeDataAckPacket() + ack.SessionId = w.SessionId + ack.CmdId = w.CmdId + ack.FdNum = w.FdNum + ack.AckLen = ackLen + if err != nil { + ack.Error = err.Error() + } + return ack +} + +func (w *FdWriter) AddData(data []byte, eof bool) error { + w.CVar.L.Lock() + defer w.CVar.L.Unlock() + if w.Closed { + return fmt.Errorf("write to closed file") + } + if len(data) > 0 { + if len(data)+len(w.Buffer) > WriteBufSize { + return fmt.Errorf("write exceeds buffer size") + } + w.Buffer = append(w.Buffer, data...) + } + if eof { + w.Eof = true + } + w.CVar.Broadcast() + return nil +} + +func (w *FdWriter) WriteLoop(sender *packet.PacketSender) { + defer w.Close() + for { + data, isEof := w.WaitForData() + if w.Closed { + return + } + if len(data) > 0 { + nw, err := w.Fd.Write(data) + ack := w.MakeDataAckPacket(nw, err) + sender.SendPacket(ack) + if err != nil { + return + } + } + if isEof { + return + } + } +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index e2cd6a8e..3afedc29 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -26,36 +26,43 @@ const DefaultRows = 25 const DefaultCols = 80 const MaxRows = 1024 const MaxCols = 1024 +const ReadBufSize = 128 * 1024 +const WriteBufSize = 128 * 1024 type ShExecType struct { + Lock *sync.Mutex StartTs time.Time RunPacket *packet.RunPacketType FileNames *base.CommandFileNames Cmd *exec.Cmd CmdPty *os.File - FdReaders map[int]*os.File - FdWriters map[int]*os.File - CloseAfterStart []*os.File + FdReaders map[int]*FdReader // synchronized + FdWriters map[int]*FdWriter // synchronized + CloseAfterStart []*os.File // synchronized } func MakeShExec(pk *packet.RunPacketType) *ShExecType { return &ShExecType{ + Lock: &sync.Mutex{}, StartTs: time.Now(), RunPacket: pk, - FdReaders: make(map[int]*os.File), - FdWriters: make(map[int]*os.File), + FdReaders: make(map[int]*FdReader), + FdWriters: make(map[int]*FdWriter), } } func (c *ShExecType) Close() { + c.Lock.Lock() + defer c.Lock.Unlock() + if c.CmdPty != nil { c.CmdPty.Close() } for _, fd := range c.FdReaders { fd.Close() } - for _, fd := range c.FdWriters { - fd.Close() + for _, fw := range c.FdWriters { + fw.Close() } for _, fd := range c.CloseAfterStart { fd.Close() @@ -221,7 +228,9 @@ func (cmd *ShExecType) makeReaderPipe(fdNum int) (*os.File, error) { if err != nil { return nil, err } - cmd.FdReaders[fdNum] = pr + cmd.Lock.Lock() + defer cmd.Lock.Unlock() + cmd.FdReaders[fdNum] = MakeFdReader(cmd, pr, fdNum) cmd.CloseAfterStart = append(cmd.CloseAfterStart, pw) return pw, nil } @@ -232,51 +241,81 @@ func (cmd *ShExecType) makeWriterPipe(fdNum int) (*os.File, error) { if err != nil { return nil, err } - cmd.FdWriters[fdNum] = pw + cmd.Lock.Lock() + defer cmd.Lock.Unlock() + cmd.FdWriters[fdNum] = MakeFdWriter(cmd, pw, fdNum) cmd.CloseAfterStart = append(cmd.CloseAfterStart, pr) return pr, nil } -func (cmd *ShExecType) MakeDataPacket(fdNum int, data []byte) *packet.DataPacketType { - pk := packet.MakeDataPacket() - pk.SessionId = cmd.RunPacket.SessionId - pk.CmdId = cmd.RunPacket.CmdId - pk.FdNum = fdNum - pk.Data = string(data) - return pk +func (cmd *ShExecType) MakeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { + ack := packet.MakeDataAckPacket() + ack.SessionId = cmd.RunPacket.SessionId + ack.CmdId = cmd.RunPacket.CmdId + ack.FdNum = fdNum + ack.AckLen = ackLen + if err != nil { + ack.Error = err.Error() + } + return ack } -func (cmd *ShExecType) runReadLoop(wg *sync.WaitGroup, fdNum int, fd *os.File, sender *packet.PacketSender) { - go func() { - defer fd.Close() - defer wg.Done() - buf := make([]byte, 4096) - for { - nr, err := fd.Read(buf) - pk := cmd.MakeDataPacket(fdNum, buf[0:nr]) - if err == io.EOF { - pk.Eof = true - sender.SendPacket(pk) - break - } else if err != nil { - pk.Error = err.Error() - sender.SendPacket(pk) - break - } else { - sender.SendPacket(pk) - } +func (cmd *ShExecType) launchWriters(sender *packet.PacketSender) { + cmd.Lock.Lock() + defer cmd.Lock.Unlock() + for _, fw := range cmd.FdWriters { + go fw.WriteLoop(sender) + } +} + +func (cmd *ShExecType) writeDataPacket(dataPacket *packet.DataPacketType) error { + cmd.Lock.Lock() + defer cmd.Lock.Unlock() + fw := cmd.FdWriters[dataPacket.FdNum] + if fw == nil { + // add a closed FdWriter as a placeholder so we only send one error + fw := MakeFdWriter(cmd, nil, dataPacket.FdNum) + fw.Close() + cmd.FdWriters[dataPacket.FdNum] = fw + return fmt.Errorf("write to closed file") + } + err := fw.AddData([]byte(dataPacket.Data), dataPacket.Eof) + if err != nil { + fw.Close() + return err + } + return nil +} + +func (cmd *ShExecType) runMainWriteLoop(packetCh chan packet.PacketType, sender *packet.PacketSender) { + for pk := range packetCh { + if pk.GetType() != packet.DataPacketStr { + // other packets are ignored + continue } - }() + dataPacket := pk.(*packet.DataPacketType) + err := cmd.writeDataPacket(dataPacket) + if err != nil { + errPacket := cmd.MakeDataAckPacket(dataPacket.FdNum, 0, err) + sender.SendPacket(errPacket) + } + } } -func (cmd *ShExecType) RunIOAndWait(sender *packet.PacketSender) { - var wg sync.WaitGroup +func (cmd *ShExecType) launchReaders(wg *sync.WaitGroup, sender *packet.PacketSender) { + cmd.Lock.Lock() + defer cmd.Lock.Unlock() wg.Add(len(cmd.FdReaders)) - go func() { - for fdNum, fd := range cmd.FdReaders { - cmd.runReadLoop(&wg, fdNum, fd, sender) - } - }() + for _, fr := range cmd.FdReaders { + go fr.ReadLoop(wg, sender) + } +} + +func (cmd *ShExecType) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { + var wg sync.WaitGroup + cmd.launchReaders(&wg, sender) + cmd.launchWriters(sender) + go cmd.runMainWriteLoop(packetCh, sender) donePacket := cmd.WaitForCommand() wg.Wait() sender.SendPacket(donePacket) From 52831dc7231c5f98ce8e904e8efac656edc00e06 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Jun 2022 18:23:30 -0700 Subject: [PATCH 018/149] setup extrafiles using run packet's fds field --- main-mshell.go | 3 +-- pkg/shexec/shexec.go | 52 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 53 insertions(+), 2 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index e5120fd8..be0233b2 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -270,8 +270,7 @@ Options: --fds [fdspec] - open fds based off [fdspec], comma separated (implies --no-auto-fds) <[num] opens for reading >[num] opens for writing - <>[num] opens for read/write - e.g. --fds '<5,>6,<>7' + e.g. --fds '<5,>6,>7' mshell is licensed under the MPLv2 Please see https://github.com/scripthaus-dev/mshell for extended usage modes, source code, bugs, and feature requests diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 3afedc29..9fd73d03 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -28,6 +28,8 @@ const MaxRows = 1024 const MaxCols = 1024 const ReadBufSize = 128 * 1024 const WriteBufSize = 128 * 1024 +const MaxFdNum = 1023 +const FirstExtraFilesFdNum = 3 type ShExecType struct { Lock *sync.Mutex @@ -344,6 +346,56 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmd.Close() return nil, err } + extraFiles := make([]*os.File, 0, MaxFdNum+1) + for _, rfd := range pk.Fds { + if rfd.FdNum < 0 { + cmd.Close() + return nil, fmt.Errorf("mshell negative fd numbers fd=%d", rfd.FdNum) + } + if rfd.FdNum < FirstExtraFilesFdNum { + cmd.Close() + return nil, fmt.Errorf("mshell does not support re-opening fd=%d (0, 1, and 2, are always open)", rfd.FdNum) + } + if rfd.FdNum > MaxFdNum { + cmd.Close() + return nil, fmt.Errorf("mshell does not support opening fd numbers above %d", MaxFdNum) + } + if rfd.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:rfd.FdNum+1] + } + if extraFiles[rfd.FdNum] != nil { + cmd.Close() + return nil, fmt.Errorf("mshell got duplicate entries for fd=%d", rfd.FdNum) + } + if rfd.Read && rfd.Write { + cmd.Close() + return nil, fmt.Errorf("mshell does not support opening fd numbers for reading and writing, fd=%d", rfd.FdNum) + } + if !rfd.Read && !rfd.Write { + cmd.Close() + return nil, fmt.Errorf("invalid fd=%d, neither reading or writing mode specified", rfd.FdNum) + } + if rfd.Read { + // client file is open for reading, so we make a writer pipe + extraFiles[rfd.FdNum], err = cmd.makeWriterPipe(rfd.FdNum) + if err != nil { + cmd.Close() + return nil, err + } + } + if rfd.Write { + // client file is open for writing, so we make a reader pipe + extraFiles[rfd.FdNum], err = cmd.makeReaderPipe(rfd.FdNum) + if err != nil { + cmd.Close() + return nil, err + } + } + } + if len(extraFiles) > FirstExtraFilesFdNum { + cmd.Cmd.ExtraFiles = extraFiles[FirstExtraFilesFdNum:] + } + err = cmd.Cmd.Start() if err != nil { cmd.Close() From 4256ff523109c47aad2c758b8baaebe6535d3348 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 24 Jun 2022 00:02:18 -0700 Subject: [PATCH 019/149] checkpoint -- cleanup and sync optimizations for remote client (basically working). beginning work on local client --- main-mshell.go | 128 +++++++++++++++++++++++++++++++++++++--- pkg/base/optsiter.go | 39 ++++++++++++ pkg/packet/packet.go | 4 +- pkg/shexec/bufreader.go | 30 ++++++++-- pkg/shexec/bufwriter.go | 16 +++-- pkg/shexec/shexec.go | 34 +++++++---- 6 files changed, 222 insertions(+), 29 deletions(-) create mode 100644 pkg/base/optsiter.go diff --git a/main-mshell.go b/main-mshell.go index be0233b2..d9608ccc 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -9,6 +9,7 @@ package main import ( "fmt" "os" + "os/exec" "os/signal" "os/user" "strings" @@ -234,6 +235,11 @@ func handleRemote() { runPacket, _ = pk.(*packet.RunPacketType) break } + if pk.GetType() == packet.RawPacketStr { + rawPk := pk.(*packet.RawPacketType) + sender.SendMessage("got raw packet '%s'", rawPk.Data) + continue + } sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) return } @@ -251,26 +257,128 @@ func handleRemote() { func handleServer() { } -func handleClient() { - fmt.Printf("mshell client\n") +func detectOpenFds() { + } -func handleUsage(extended bool) { +type ClientOpts struct { + IsSSH bool + SSHOptsTerm bool + SSHOpts []string + Command string + Fds []packet.RemoteFd + Cwd string +} + +func parseClientOpts() (*ClientOpts, error) { + opts := &ClientOpts{} + iter := base.MakeOptsIter(os.Args[1:]) + for iter.HasNext() { + argStr := iter.Next() + if argStr == "--ssh" { + if opts.IsSSH { + return nil, fmt.Errorf("duplicate '--ssh' option") + } + opts.IsSSH = true + break + } + } + if opts.IsSSH { + // parse SSH opts + for iter.HasNext() { + argStr := iter.Next() + if argStr == "--" { + opts.SSHOptsTerm = true + break + } + if argStr == "--cwd" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--cwd [dir]' missing directory") + } + } + opts.SSHOpts = append(opts.SSHOpts, argStr) + } + if !opts.SSHOptsTerm { + return nil, fmt.Errorf("ssh options must be terminated with '--' followed by [command]") + } + if !iter.HasNext() { + return nil, fmt.Errorf("no command specified") + } + opts.Command = strings.Join(iter.Rest(), " ") + if strings.TrimSpace(opts.Command) == "" { + return nil, fmt.Errorf("no command or empty command specified") + } + } + return opts, nil +} + +func handleClient() (int, error) { + fmt.Printf("mshell client\n") + opts, err := parseClientOpts() + if err != nil { + return 1, fmt.Errorf("parsing opts: %w", err) + } + if !opts.IsSSH { + return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") + } + fmt.Printf("opts: %v\n", opts) + sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` + sshOpts := append(opts.SSHOpts, sshRemoteCommand) + ecmd := exec.Command("ssh", sshOpts...) + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return 1, fmt.Errorf("creating stdin pipe: %v", err) + } + outputReader, err := ecmd.StdoutPipe() + if err != nil { + return 1, fmt.Errorf("creating stdout pipe: %v", err) + } + ecmd.Stderr = ecmd.Stdout + err = ecmd.Start() + if err != nil { + return 1, fmt.Errorf("running ssh command: %w", err) + } + parser := packet.PacketParser(outputReader) + go func() { + fmt.Printf("%v %v\n", parser, inputWriter) + }() + exitErr := ecmd.Wait() + return shexec.GetExitCode(exitErr), nil +} + +func handleUsage() { usage := ` -Client Usage: mshell [mshell-opts] [ssh-opts] user@host [command] +Client Usage: mshell [mshell-opts] --ssh [ssh-opts] user@host -- [command] mshell multiplexes input and output streams to a remote command over ssh. Options: - --env 'X=Y,A=B' - set remote environment variables for command, comma or newline separated + --env 'X=Y;A=B' - set remote environment variables for command, semicolon separated --env-file [file] - load environment variables from [file] (.env format) --env-copy [glob] - copy local environment variables to remote using [glob] pattern --cwd [dir] - execute remote command in [dir] --no-auto-fds - do not auto-detect additional fds + --sudo - execute "sudo [command]" --fds [fdspec] - open fds based off [fdspec], comma separated (implies --no-auto-fds) <[num] opens for reading >[num] opens for writing e.g. --fds '<5,>6,>7' + [command] - a single argument (should be quoted) + +Examples: + # execute a python script remotely, with stdin still hooked up correctly + mshell --cwd "~/work" --ssh -i key.pem ubuntu@somehost -- "python /dev/fd/4" 4< myscript.py + + # capture multiple outputs + mshell --ssh ubuntu@test -- "cat file1.txt > /dev/fd/3; cat file2.txt > /dev/fd/4" 3> file1.txt 4> file2.txt + + # environment variable copying, setting working directory + # note the single quotes on command (otherwise the local shell will expand the variables) + TEST1=hello TEST2=world mshell --cwd "~/work" --env-copy "TEST*" --ssh user@host -- 'echo $(pwd) $TEST1 $TEST2' + + # execute a script, catpure stdout/stderr in fd-3 and fd-4 + # useful if you need to see stdout for interacting with ssh (password or host auth) + mshell --ssh user@host -- "test.sh > /dev/fd/3 2> /dev/fd/4" 3> test.stdout 4> test.stderr mshell is licensed under the MPLv2 Please see https://github.com/scripthaus-dev/mshell for extended usage modes, source code, bugs, and feature requests @@ -280,12 +388,12 @@ Please see https://github.com/scripthaus-dev/mshell for extended usage modes, so func main() { if len(os.Args) == 1 { - handleUsage(false) + handleUsage() return } firstArg := os.Args[1] if firstArg == "--help" { - handleUsage(true) + handleUsage() return } else if firstArg == "--version" { fmt.Printf("mshell v%s\n", MShellVersion) @@ -297,7 +405,11 @@ func main() { handleServer() return } else { - handleClient() + rtnCode, err := handleClient() + if err != nil { + fmt.Printf("[error] %v\n", err) + } + os.Exit(rtnCode) return } diff --git a/pkg/base/optsiter.go b/pkg/base/optsiter.go new file mode 100644 index 00000000..329455ba --- /dev/null +++ b/pkg/base/optsiter.go @@ -0,0 +1,39 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package base + +import "strings" + +type OptsIter struct { + Pos int + Opts []string +} + +func MakeOptsIter(opts []string) *OptsIter { + return &OptsIter{Opts: opts} +} + +func IsOption(argStr string) bool { + return strings.HasPrefix(argStr, "-") && argStr != "-" && !strings.HasPrefix(argStr, "-/") +} + +func (iter *OptsIter) HasNext() bool { + return iter.Pos <= len(iter.Opts)-1 +} + +func (iter *OptsIter) Next() string { + if iter.Pos >= len(iter.Opts) { + return "" + } + rtn := iter.Opts[iter.Pos] + iter.Pos++ + return rtn +} + +func (iter *OptsIter) Rest() []string { + return iter.Opts[iter.Pos:] +} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index fc9fcaa4..5f6f3bd2 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -43,6 +43,8 @@ const ( InputPacketStr = "input" ) +const PacketSenderQueueSize = 20 + var TypeStrToFactory map[string]reflect.Type func init() { @@ -450,7 +452,7 @@ type PacketSender struct { func MakePacketSender(output io.Writer) *PacketSender { sender := &PacketSender{ Lock: &sync.Mutex{}, - SendCh: make(chan PacketType), + SendCh: make(chan PacketType, PacketSenderQueueSize), DoneCh: make(chan bool), } go func() { diff --git a/pkg/shexec/bufreader.go b/pkg/shexec/bufreader.go index 73aefc7b..e666c495 100644 --- a/pkg/shexec/bufreader.go +++ b/pkg/shexec/bufreader.go @@ -50,6 +50,9 @@ func (r *FdReader) Close() { func (r *FdReader) NotifyAck(ackLen int) { r.CVar.L.Lock() defer r.CVar.L.Unlock() + if r.Closed { + return + } r.BufSize -= ackLen if r.BufSize < 0 { r.BufSize = 0 @@ -57,11 +60,17 @@ func (r *FdReader) NotifyAck(ackLen int) { r.CVar.Broadcast() } +// !! inverse locking. must already hold the lock when you call this method. +// will *unlock*, send the packet, and then *relock* once it is done. +// this can prevent an unlikely deadlock where we are holding r.CVar.L and stuck on sender.SendCh +func (r *FdReader) sendPacket_unlock(sender *packet.PacketSender, pk packet.PacketType) { + r.CVar.L.Unlock() + defer r.CVar.L.Lock() + sender.SendPacket(pk) +} + // returns (success) func (r *FdReader) WriteWait(sender *packet.PacketSender, data []byte, isEof bool) bool { - if len(data) == 0 { - return true - } r.CVar.L.Lock() defer r.CVar.L.Unlock() for { @@ -75,13 +84,15 @@ func (r *FdReader) WriteWait(sender *packet.PacketSender, data []byte, isEof boo } writeLen := min(bufAvail, len(data)) pk := r.MakeDataPacket(data[0:writeLen], nil) - sender.SendPacket(pk) + pk.Eof = isEof && (writeLen == len(data)) r.BufSize += writeLen data = data[writeLen:] + r.sendPacket_unlock(sender, pk) if len(data) == 0 { return true } - r.CVar.Wait() + // do *not* do a CVar.Wait() here -- because we *unlocked* to send the packet, we should + // recheck the condition before waiting to avoid deadlock. } } @@ -104,12 +115,21 @@ func (r *FdReader) MakeDataPacket(data []byte, err error) *packet.DataPacketType return pk } +func (r *FdReader) isClosed() bool { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + return r.Closed +} + func (r *FdReader) ReadLoop(wg *sync.WaitGroup, sender *packet.PacketSender) { defer r.Close() defer wg.Done() buf := make([]byte, 4096) for { nr, err := r.Fd.Read(buf) + if r.isClosed() { + return // should not send data or error if we already closed the fd + } if nr > 0 || err == io.EOF { isOpen := r.WriteWait(sender, buf[0:nr], (err == io.EOF)) if !isOpen { diff --git a/pkg/shexec/bufwriter.go b/pkg/shexec/bufwriter.go index f51a0604..552e4c74 100644 --- a/pkg/shexec/bufwriter.go +++ b/pkg/shexec/bufwriter.go @@ -14,6 +14,8 @@ import ( "github.com/scripthaus-dev/mshell/pkg/packet" ) +const MaxSingleWriteSize = 4 * 1024 + type FdWriter struct { CVar *sync.Cond SessionId string @@ -97,16 +99,20 @@ func (w *FdWriter) WriteLoop(sender *packet.PacketSender) { defer w.Close() for { data, isEof := w.WaitForData() - if w.Closed { - return - } - if len(data) > 0 { - nw, err := w.Fd.Write(data) + // chunk the writes to make sure we send ample ack packets + for len(data) > 0 { + if w.Closed { + return + } + chunkSize := min(len(data), MaxSingleWriteSize) + chunk := data[0:chunkSize] + nw, err := w.Fd.Write(chunk) ack := w.MakeDataAckPacket(nw, err) sender.SendPacket(ack) if err != nil { return } + data = data[chunkSize:] } if isEof { return diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 9fd73d03..52663dcd 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -270,7 +270,7 @@ func (cmd *ShExecType) launchWriters(sender *packet.PacketSender) { } } -func (cmd *ShExecType) writeDataPacket(dataPacket *packet.DataPacketType) error { +func (cmd *ShExecType) processDataPacket(dataPacket *packet.DataPacketType) error { cmd.Lock.Lock() defer cmd.Lock.Unlock() fw := cmd.FdWriters[dataPacket.FdNum] @@ -289,18 +289,32 @@ func (cmd *ShExecType) writeDataPacket(dataPacket *packet.DataPacketType) error return nil } -func (cmd *ShExecType) runMainWriteLoop(packetCh chan packet.PacketType, sender *packet.PacketSender) { +func (cmd *ShExecType) processAckPacket(ackPacket *packet.DataAckPacketType) { + cmd.Lock.Lock() + defer cmd.Lock.Unlock() + fr := cmd.FdReaders[ackPacket.FdNum] + if fr == nil { + return + } + fr.NotifyAck(ackPacket.AckLen) +} + +func (cmd *ShExecType) runPacketInputLoop(packetCh chan packet.PacketType, sender *packet.PacketSender) { for pk := range packetCh { - if pk.GetType() != packet.DataPacketStr { - // other packets are ignored + if pk.GetType() == packet.DataPacketStr { + dataPacket := pk.(*packet.DataPacketType) + err := cmd.processDataPacket(dataPacket) + if err != nil { + errPacket := cmd.MakeDataAckPacket(dataPacket.FdNum, 0, err) + sender.SendPacket(errPacket) + } continue } - dataPacket := pk.(*packet.DataPacketType) - err := cmd.writeDataPacket(dataPacket) - if err != nil { - errPacket := cmd.MakeDataAckPacket(dataPacket.FdNum, 0, err) - sender.SendPacket(errPacket) + if pk.GetType() == packet.DataAckPacketStr { + ackPacket := pk.(*packet.DataAckPacketType) + cmd.processAckPacket(ackPacket) } + // other packet types are ignored } } @@ -317,7 +331,7 @@ func (cmd *ShExecType) RunIOAndWait(packetCh chan packet.PacketType, sender *pac var wg sync.WaitGroup cmd.launchReaders(&wg, sender) cmd.launchWriters(sender) - go cmd.runMainWriteLoop(packetCh, sender) + go cmd.runPacketInputLoop(packetCh, sender) donePacket := cmd.WaitForCommand() wg.Wait() sender.SendPacket(donePacket) From 0267836376173c1bb0759e6fb50d0140846a394b Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 24 Jun 2022 10:24:02 -0700 Subject: [PATCH 020/149] move multiplexed IO to its own package independent of SHExecType (to use in mshell client) --- pkg/{shexec => mpio}/bufreader.go | 64 ++++------ pkg/{shexec => mpio}/bufwriter.go | 53 +++----- pkg/mpio/mpio.go | 206 ++++++++++++++++++++++++++++++ pkg/shexec/shexec.go | 164 +++--------------------- 4 files changed, 270 insertions(+), 217 deletions(-) rename pkg/{shexec => mpio}/bufreader.go (63%) rename pkg/{shexec => mpio}/bufwriter.go (63%) create mode 100644 pkg/mpio/mpio.go diff --git a/pkg/shexec/bufreader.go b/pkg/mpio/bufreader.go similarity index 63% rename from pkg/shexec/bufreader.go rename to pkg/mpio/bufreader.go index e666c495..f78ba4b1 100644 --- a/pkg/shexec/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -4,7 +4,7 @@ // License, v. 2.0. If a copy of the MPL was not distributed with this // file, You can obtain one at https://mozilla.org/MPL/2.0/. -package shexec +package mpio import ( "io" @@ -15,24 +15,23 @@ import ( ) type FdReader struct { - CVar *sync.Cond - SessionId string - CmdId string - FdNum int - Fd *os.File - BufSize int - Closed bool + CVar *sync.Cond + M *Multiplexer + FdNum int + Fd *os.File + BufSize int + Closed bool } -func MakeFdReader(c *ShExecType, fd *os.File, fdNum int) *FdReader { - return &FdReader{ - CVar: sync.NewCond(&sync.Mutex{}), - SessionId: c.RunPacket.SessionId, - CmdId: c.RunPacket.CmdId, - FdNum: fdNum, - Fd: fd, - BufSize: 0, +func MakeFdReader(m *Multiplexer, fd *os.File, fdNum int) *FdReader { + fr := &FdReader{ + CVar: sync.NewCond(&sync.Mutex{}), + M: m, + FdNum: fdNum, + Fd: fd, + BufSize: 0, } + return fr } func (r *FdReader) Close() { @@ -63,14 +62,14 @@ func (r *FdReader) NotifyAck(ackLen int) { // !! inverse locking. must already hold the lock when you call this method. // will *unlock*, send the packet, and then *relock* once it is done. // this can prevent an unlikely deadlock where we are holding r.CVar.L and stuck on sender.SendCh -func (r *FdReader) sendPacket_unlock(sender *packet.PacketSender, pk packet.PacketType) { +func (r *FdReader) sendPacket_unlock(pk packet.PacketType) { r.CVar.L.Unlock() defer r.CVar.L.Lock() - sender.SendPacket(pk) + r.M.sendPacket(pk) } // returns (success) -func (r *FdReader) WriteWait(sender *packet.PacketSender, data []byte, isEof bool) bool { +func (r *FdReader) WriteWait(data []byte, isEof bool) bool { r.CVar.L.Lock() defer r.CVar.L.Unlock() for { @@ -83,11 +82,11 @@ func (r *FdReader) WriteWait(sender *packet.PacketSender, data []byte, isEof boo continue } writeLen := min(bufAvail, len(data)) - pk := r.MakeDataPacket(data[0:writeLen], nil) + pk := r.M.makeDataPacket(r.FdNum, data[0:writeLen], nil) pk.Eof = isEof && (writeLen == len(data)) r.BufSize += writeLen data = data[writeLen:] - r.sendPacket_unlock(sender, pk) + r.sendPacket_unlock(pk) if len(data) == 0 { return true } @@ -103,25 +102,13 @@ func min(v1 int, v2 int) int { return v2 } -func (r *FdReader) MakeDataPacket(data []byte, err error) *packet.DataPacketType { - pk := packet.MakeDataPacket() - pk.SessionId = r.SessionId - pk.CmdId = r.CmdId - pk.FdNum = r.FdNum - pk.Data = string(data) - if err != nil { - pk.Error = err.Error() - } - return pk -} - func (r *FdReader) isClosed() bool { r.CVar.L.Lock() defer r.CVar.L.Unlock() return r.Closed } -func (r *FdReader) ReadLoop(wg *sync.WaitGroup, sender *packet.PacketSender) { +func (r *FdReader) ReadLoop(wg *sync.WaitGroup) { defer r.Close() defer wg.Done() buf := make([]byte, 4096) @@ -131,14 +118,17 @@ func (r *FdReader) ReadLoop(wg *sync.WaitGroup, sender *packet.PacketSender) { return // should not send data or error if we already closed the fd } if nr > 0 || err == io.EOF { - isOpen := r.WriteWait(sender, buf[0:nr], (err == io.EOF)) + isOpen := r.WriteWait(buf[0:nr], (err == io.EOF)) if !isOpen { return } + if err == io.EOF { + return + } } if err != nil { - errPk := r.MakeDataPacket(nil, err) - sender.SendPacket(errPk) + errPk := r.M.makeDataPacket(r.FdNum, nil, err) + r.M.sendPacket(errPk) return } } diff --git a/pkg/shexec/bufwriter.go b/pkg/mpio/bufwriter.go similarity index 63% rename from pkg/shexec/bufwriter.go rename to pkg/mpio/bufwriter.go index 552e4c74..16e139ca 100644 --- a/pkg/shexec/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -4,37 +4,32 @@ // License, v. 2.0. If a copy of the MPL was not distributed with this // file, You can obtain one at https://mozilla.org/MPL/2.0/. -package shexec +package mpio import ( "fmt" "os" "sync" - - "github.com/scripthaus-dev/mshell/pkg/packet" ) -const MaxSingleWriteSize = 4 * 1024 - type FdWriter struct { - CVar *sync.Cond - SessionId string - CmdId string - FdNum int - Buffer []byte - Fd *os.File - Eof bool - Closed bool + CVar *sync.Cond + M *Multiplexer + FdNum int + Buffer []byte + Fd *os.File + Eof bool + Closed bool } -func MakeFdWriter(c *ShExecType, fd *os.File, fdNum int) *FdWriter { - return &FdWriter{ - CVar: sync.NewCond(&sync.Mutex{}), - Fd: fd, - SessionId: c.RunPacket.SessionId, - CmdId: c.RunPacket.CmdId, - FdNum: fdNum, +func MakeFdWriter(m *Multiplexer, fd *os.File, fdNum int) *FdWriter { + fw := &FdWriter{ + CVar: sync.NewCond(&sync.Mutex{}), + Fd: fd, + M: m, + FdNum: fdNum, } + return fw } func (w *FdWriter) Close() { @@ -64,18 +59,6 @@ func (w *FdWriter) WaitForData() ([]byte, bool) { } } -func (w *FdWriter) MakeDataAckPacket(ackLen int, err error) *packet.DataAckPacketType { - ack := packet.MakeDataAckPacket() - ack.SessionId = w.SessionId - ack.CmdId = w.CmdId - ack.FdNum = w.FdNum - ack.AckLen = ackLen - if err != nil { - ack.Error = err.Error() - } - return ack -} - func (w *FdWriter) AddData(data []byte, eof bool) error { w.CVar.L.Lock() defer w.CVar.L.Unlock() @@ -95,7 +78,7 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { return nil } -func (w *FdWriter) WriteLoop(sender *packet.PacketSender) { +func (w *FdWriter) WriteLoop() { defer w.Close() for { data, isEof := w.WaitForData() @@ -107,8 +90,8 @@ func (w *FdWriter) WriteLoop(sender *packet.PacketSender) { chunkSize := min(len(data), MaxSingleWriteSize) chunk := data[0:chunkSize] nw, err := w.Fd.Write(chunk) - ack := w.MakeDataAckPacket(nw, err) - sender.SendPacket(ack) + ack := w.M.makeDataAckPacket(w.FdNum, nw, err) + w.M.sendPacket(ack) if err != nil { return } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go new file mode 100644 index 00000000..c3f82d61 --- /dev/null +++ b/pkg/mpio/mpio.go @@ -0,0 +1,206 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package mpio + +import ( + "fmt" + "os" + "sync" + + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +const ReadBufSize = 128 * 1024 +const WriteBufSize = 128 * 1024 +const MaxSingleWriteSize = 4 * 1024 + +type Multiplexer struct { + Lock *sync.Mutex + SessionId string + CmdId string + FdReaders map[int]*FdReader // synchronized + FdWriters map[int]*FdWriter // synchronized + CloseAfterStart []*os.File // synchronized + + Sender *packet.PacketSender + Input chan packet.PacketType + Started bool +} + +func MakeMultiplexer(sessionId string, cmdId string) *Multiplexer { + return &Multiplexer{ + Lock: &sync.Mutex{}, + SessionId: sessionId, + CmdId: cmdId, + FdReaders: make(map[int]*FdReader), + FdWriters: make(map[int]*FdWriter), + } +} + +func (m *Multiplexer) Close() { + m.Lock.Lock() + defer m.Lock.Unlock() + + for _, fd := range m.FdReaders { + fd.Close() + } + for _, fd := range m.FdWriters { + fd.Close() + } + for _, fd := range m.CloseAfterStart { + fd.Close() + } +} + +// returns the *writer* to connect to process, reader is put in FdReaders +func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum) + m.CloseAfterStart = append(m.CloseAfterStart, pw) + return pw, nil +} + +// returns the *reader* to connect to process, writer is put in FdWriters +func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum) + m.CloseAfterStart = append(m.CloseAfterStart, pr) + return pr, nil +} + +func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { + ack := packet.MakeDataAckPacket() + ack.SessionId = m.SessionId + ack.CmdId = m.CmdId + ack.FdNum = fdNum + ack.AckLen = ackLen + if err != nil { + ack.Error = err.Error() + } + return ack +} + +func (m *Multiplexer) makeDataPacket(fdNum int, data []byte, err error) *packet.DataPacketType { + pk := packet.MakeDataPacket() + pk.SessionId = m.SessionId + pk.CmdId = m.CmdId + pk.FdNum = fdNum + pk.Data = string(data) + if err != nil { + pk.Error = err.Error() + } + return pk +} + +func (m *Multiplexer) sendPacket(p packet.PacketType) { + m.Sender.SendPacket(p) +} + +func (m *Multiplexer) launchWriters() { + m.Lock.Lock() + defer m.Lock.Unlock() + for _, fw := range m.FdWriters { + go fw.WriteLoop() + } +} + +func (m *Multiplexer) launchReaders(wg *sync.WaitGroup) { + m.Lock.Lock() + defer m.Lock.Unlock() + wg.Add(len(m.FdReaders)) + for _, fr := range m.FdReaders { + go fr.ReadLoop(wg) + } +} + +func (m *Multiplexer) startIO(packetCh chan packet.PacketType, sender *packet.PacketSender) { + m.Lock.Lock() + defer m.Lock.Unlock() + if m.Started { + panic("Multiplexer is already running, cannot start again") + } + m.Input = packetCh + m.Sender = sender + m.Started = true +} + +func (m *Multiplexer) runPacketInputLoop() { + for pk := range m.Input { + if pk.GetType() == packet.DataPacketStr { + dataPacket := pk.(*packet.DataPacketType) + err := m.processDataPacket(dataPacket) + if err != nil { + errPacket := m.makeDataAckPacket(dataPacket.FdNum, 0, err) + m.sendPacket(errPacket) + } + continue + } + if pk.GetType() == packet.DataAckPacketStr { + ackPacket := pk.(*packet.DataAckPacketType) + m.processAckPacket(ackPacket) + } + // other packet types are ignored + } +} + +func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { + m.Lock.Lock() + defer m.Lock.Unlock() + fw := m.FdWriters[dataPacket.FdNum] + if fw == nil { + // add a closed FdWriter as a placeholder so we only send one error + fw := MakeFdWriter(m, nil, dataPacket.FdNum) + fw.Close() + m.FdWriters[dataPacket.FdNum] = fw + return fmt.Errorf("write to closed file") + } + err := fw.AddData([]byte(dataPacket.Data), dataPacket.Eof) + if err != nil { + fw.Close() + return err + } + return nil +} + +func (m *Multiplexer) processAckPacket(ackPacket *packet.DataAckPacketType) { + m.Lock.Lock() + defer m.Lock.Unlock() + fr := m.FdReaders[ackPacket.FdNum] + if fr == nil { + return + } + fr.NotifyAck(ackPacket.AckLen) +} + +func (m *Multiplexer) closeTempStartFds() { + m.Lock.Lock() + defer m.Lock.Unlock() + for _, fd := range m.CloseAfterStart { + fd.Close() + } + m.CloseAfterStart = nil +} + +func (m *Multiplexer) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { + m.startIO(packetCh, sender) + m.closeTempStartFds() + var wg sync.WaitGroup + m.launchReaders(&wg) + m.launchWriters() + go m.runPacketInputLoop() + wg.Wait() +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 52663dcd..e8fc52ad 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -19,6 +19,7 @@ import ( "github.com/creack/pty" "github.com/google/uuid" "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" ) @@ -26,49 +27,33 @@ const DefaultRows = 25 const DefaultCols = 80 const MaxRows = 1024 const MaxCols = 1024 -const ReadBufSize = 128 * 1024 -const WriteBufSize = 128 * 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 type ShExecType struct { - Lock *sync.Mutex - StartTs time.Time - RunPacket *packet.RunPacketType - FileNames *base.CommandFileNames - Cmd *exec.Cmd - CmdPty *os.File - FdReaders map[int]*FdReader // synchronized - FdWriters map[int]*FdWriter // synchronized - CloseAfterStart []*os.File // synchronized + Lock *sync.Mutex + StartTs time.Time + RunPacket *packet.RunPacketType + FileNames *base.CommandFileNames + Cmd *exec.Cmd + CmdPty *os.File + Multiplexer *mpio.Multiplexer } func MakeShExec(pk *packet.RunPacketType) *ShExecType { return &ShExecType{ - Lock: &sync.Mutex{}, - StartTs: time.Now(), - RunPacket: pk, - FdReaders: make(map[int]*FdReader), - FdWriters: make(map[int]*FdWriter), + Lock: &sync.Mutex{}, + StartTs: time.Now(), + RunPacket: pk, + Multiplexer: mpio.MakeMultiplexer(pk.SessionId, pk.CmdId), } } func (c *ShExecType) Close() { - c.Lock.Lock() - defer c.Lock.Unlock() - if c.CmdPty != nil { c.CmdPty.Close() } - for _, fd := range c.FdReaders { - fd.Close() - } - for _, fw := range c.FdWriters { - fw.Close() - } - for _, fd := range c.CloseAfterStart { - fd.Close() - } + c.Multiplexer.Close() } func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { @@ -224,116 +209,9 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } } -// returns the *writer* to connect to process, reader is put in FdReaders -func (cmd *ShExecType) makeReaderPipe(fdNum int) (*os.File, error) { - pr, pw, err := os.Pipe() - if err != nil { - return nil, err - } - cmd.Lock.Lock() - defer cmd.Lock.Unlock() - cmd.FdReaders[fdNum] = MakeFdReader(cmd, pr, fdNum) - cmd.CloseAfterStart = append(cmd.CloseAfterStart, pw) - return pw, nil -} - -// returns the *reader* to connect to process, writer is put in FdWriters -func (cmd *ShExecType) makeWriterPipe(fdNum int) (*os.File, error) { - pr, pw, err := os.Pipe() - if err != nil { - return nil, err - } - cmd.Lock.Lock() - defer cmd.Lock.Unlock() - cmd.FdWriters[fdNum] = MakeFdWriter(cmd, pw, fdNum) - cmd.CloseAfterStart = append(cmd.CloseAfterStart, pr) - return pr, nil -} - -func (cmd *ShExecType) MakeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { - ack := packet.MakeDataAckPacket() - ack.SessionId = cmd.RunPacket.SessionId - ack.CmdId = cmd.RunPacket.CmdId - ack.FdNum = fdNum - ack.AckLen = ackLen - if err != nil { - ack.Error = err.Error() - } - return ack -} - -func (cmd *ShExecType) launchWriters(sender *packet.PacketSender) { - cmd.Lock.Lock() - defer cmd.Lock.Unlock() - for _, fw := range cmd.FdWriters { - go fw.WriteLoop(sender) - } -} - -func (cmd *ShExecType) processDataPacket(dataPacket *packet.DataPacketType) error { - cmd.Lock.Lock() - defer cmd.Lock.Unlock() - fw := cmd.FdWriters[dataPacket.FdNum] - if fw == nil { - // add a closed FdWriter as a placeholder so we only send one error - fw := MakeFdWriter(cmd, nil, dataPacket.FdNum) - fw.Close() - cmd.FdWriters[dataPacket.FdNum] = fw - return fmt.Errorf("write to closed file") - } - err := fw.AddData([]byte(dataPacket.Data), dataPacket.Eof) - if err != nil { - fw.Close() - return err - } - return nil -} - -func (cmd *ShExecType) processAckPacket(ackPacket *packet.DataAckPacketType) { - cmd.Lock.Lock() - defer cmd.Lock.Unlock() - fr := cmd.FdReaders[ackPacket.FdNum] - if fr == nil { - return - } - fr.NotifyAck(ackPacket.AckLen) -} - -func (cmd *ShExecType) runPacketInputLoop(packetCh chan packet.PacketType, sender *packet.PacketSender) { - for pk := range packetCh { - if pk.GetType() == packet.DataPacketStr { - dataPacket := pk.(*packet.DataPacketType) - err := cmd.processDataPacket(dataPacket) - if err != nil { - errPacket := cmd.MakeDataAckPacket(dataPacket.FdNum, 0, err) - sender.SendPacket(errPacket) - } - continue - } - if pk.GetType() == packet.DataAckPacketStr { - ackPacket := pk.(*packet.DataAckPacketType) - cmd.processAckPacket(ackPacket) - } - // other packet types are ignored - } -} - -func (cmd *ShExecType) launchReaders(wg *sync.WaitGroup, sender *packet.PacketSender) { - cmd.Lock.Lock() - defer cmd.Lock.Unlock() - wg.Add(len(cmd.FdReaders)) - for _, fr := range cmd.FdReaders { - go fr.ReadLoop(wg, sender) - } -} - func (cmd *ShExecType) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { - var wg sync.WaitGroup - cmd.launchReaders(&wg, sender) - cmd.launchWriters(sender) - go cmd.runPacketInputLoop(packetCh, sender) + cmd.Multiplexer.RunIOAndWait(packetCh, sender) donePacket := cmd.WaitForCommand() - wg.Wait() sender.SendPacket(donePacket) } @@ -345,17 +223,17 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmd.Cmd.Dir = pk.Cwd } var err error - cmd.Cmd.Stdin, err = cmd.makeWriterPipe(0) + cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) if err != nil { cmd.Close() return nil, err } - cmd.Cmd.Stdout, err = cmd.makeReaderPipe(1) + cmd.Cmd.Stdout, err = cmd.Multiplexer.MakeReaderPipe(1) if err != nil { cmd.Close() return nil, err } - cmd.Cmd.Stderr, err = cmd.makeReaderPipe(2) + cmd.Cmd.Stderr, err = cmd.Multiplexer.MakeReaderPipe(2) if err != nil { cmd.Close() return nil, err @@ -391,7 +269,7 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S } if rfd.Read { // client file is open for reading, so we make a writer pipe - extraFiles[rfd.FdNum], err = cmd.makeWriterPipe(rfd.FdNum) + extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum) if err != nil { cmd.Close() return nil, err @@ -399,7 +277,7 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S } if rfd.Write { // client file is open for writing, so we make a reader pipe - extraFiles[rfd.FdNum], err = cmd.makeReaderPipe(rfd.FdNum) + extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeReaderPipe(rfd.FdNum) if err != nil { cmd.Close() return nil, err @@ -415,10 +293,6 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmd.Close() return nil, err } - for _, fd := range cmd.CloseAfterStart { - fd.Close() - } - cmd.CloseAfterStart = nil return cmd, nil } From 5223760a7687e694f72929b4f43c56803300a9b0 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 24 Jun 2022 13:25:09 -0700 Subject: [PATCH 021/149] got basic mshell client working -- still need detectfds and extra files support --- main-mshell.go | 41 +++-------------- pkg/mpio/bufreader.go | 32 +++++++------ pkg/mpio/bufwriter.go | 42 +++++++++++------- pkg/mpio/mpio.go | 95 ++++++++++++++++++++++++++++++++------- pkg/packet/packet.go | 8 +++- pkg/shexec/shexec.go | 101 +++++++++++++++++++++++++++++++++++++----- 6 files changed, 226 insertions(+), 93 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index d9608ccc..418769f2 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -9,7 +9,6 @@ package main import ( "fmt" "os" - "os/exec" "os/signal" "os/user" "strings" @@ -251,7 +250,7 @@ func handleRemote() { defer cmd.Close() startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) - cmd.RunIOAndWait(packetCh, sender) + cmd.RunRemoteIOAndWait(packetCh, sender) } func handleServer() { @@ -261,17 +260,8 @@ func detectOpenFds() { } -type ClientOpts struct { - IsSSH bool - SSHOptsTerm bool - SSHOpts []string - Command string - Fds []packet.RemoteFd - Cwd string -} - -func parseClientOpts() (*ClientOpts, error) { - opts := &ClientOpts{} +func parseClientOpts() (*shexec.ClientOpts, error) { + opts := &shexec.ClientOpts{} iter := base.MakeOptsIter(os.Args[1:]) for iter.HasNext() { argStr := iter.Next() @@ -313,7 +303,6 @@ func parseClientOpts() (*ClientOpts, error) { } func handleClient() (int, error) { - fmt.Printf("mshell client\n") opts, err := parseClientOpts() if err != nil { return 1, fmt.Errorf("parsing opts: %w", err) @@ -321,29 +310,11 @@ func handleClient() (int, error) { if !opts.IsSSH { return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") } - fmt.Printf("opts: %v\n", opts) - sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` - sshOpts := append(opts.SSHOpts, sshRemoteCommand) - ecmd := exec.Command("ssh", sshOpts...) - inputWriter, err := ecmd.StdinPipe() + donePacket, err := shexec.RunClientSSHCommandAndWait(opts) if err != nil { - return 1, fmt.Errorf("creating stdin pipe: %v", err) + return 1, err } - outputReader, err := ecmd.StdoutPipe() - if err != nil { - return 1, fmt.Errorf("creating stdout pipe: %v", err) - } - ecmd.Stderr = ecmd.Stdout - err = ecmd.Start() - if err != nil { - return 1, fmt.Errorf("running ssh command: %w", err) - } - parser := packet.PacketParser(outputReader) - go func() { - fmt.Printf("%v %v\n", parser, inputWriter) - }() - exitErr := ecmd.Wait() - return shexec.GetExitCode(exitErr), nil + return donePacket.ExitCode, nil } func handleUsage() { diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index f78ba4b1..756f9e83 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -15,21 +15,23 @@ import ( ) type FdReader struct { - CVar *sync.Cond - M *Multiplexer - FdNum int - Fd *os.File - BufSize int - Closed bool + CVar *sync.Cond + M *Multiplexer + FdNum int + Fd *os.File + BufSize int + Closed bool + ShouldCloseFd bool } -func MakeFdReader(m *Multiplexer, fd *os.File, fdNum int) *FdReader { +func MakeFdReader(m *Multiplexer, fd *os.File, fdNum int, shouldCloseFd bool) *FdReader { fr := &FdReader{ - CVar: sync.NewCond(&sync.Mutex{}), - M: m, - FdNum: fdNum, - Fd: fd, - BufSize: 0, + CVar: sync.NewCond(&sync.Mutex{}), + M: m, + FdNum: fdNum, + Fd: fd, + BufSize: 0, + ShouldCloseFd: shouldCloseFd, } return fr } @@ -40,7 +42,7 @@ func (r *FdReader) Close() { if r.Closed { return } - if r.Fd != nil { + if r.Fd != nil && r.ShouldCloseFd { r.Fd.Close() } r.CVar.Broadcast() @@ -110,7 +112,9 @@ func (r *FdReader) isClosed() bool { func (r *FdReader) ReadLoop(wg *sync.WaitGroup) { defer r.Close() - defer wg.Done() + if wg != nil { + defer wg.Done() + } buf := make([]byte, 4096) for { nr, err := r.Fd.Read(buf) diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index 16e139ca..9b389678 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -13,21 +13,23 @@ import ( ) type FdWriter struct { - CVar *sync.Cond - M *Multiplexer - FdNum int - Buffer []byte - Fd *os.File - Eof bool - Closed bool + CVar *sync.Cond + M *Multiplexer + FdNum int + Buffer []byte + Fd *os.File + Eof bool + Closed bool + ShouldCloseFd bool } -func MakeFdWriter(m *Multiplexer, fd *os.File, fdNum int) *FdWriter { +func MakeFdWriter(m *Multiplexer, fd *os.File, fdNum int, shouldCloseFd bool) *FdWriter { fw := &FdWriter{ - CVar: sync.NewCond(&sync.Mutex{}), - Fd: fd, - M: m, - FdNum: fdNum, + CVar: sync.NewCond(&sync.Mutex{}), + Fd: fd, + M: m, + FdNum: fdNum, + ShouldCloseFd: shouldCloseFd, } return fw } @@ -39,7 +41,7 @@ func (w *FdWriter) Close() { return } w.Closed = true - if w.Fd != nil { + if w.Fd != nil && w.ShouldCloseFd { w.Fd.Close() } w.Buffer = nil @@ -65,6 +67,9 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { if w.Closed { return fmt.Errorf("write to closed file") } + if w.Eof { + return fmt.Errorf("write to closed file (eof)") + } if len(data) > 0 { if len(data)+len(w.Buffer) > WriteBufSize { return fmt.Errorf("write exceeds buffer size") @@ -78,8 +83,11 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { return nil } -func (w *FdWriter) WriteLoop() { +func (w *FdWriter) WriteLoop(wg *sync.WaitGroup) { defer w.Close() + if wg != nil { + defer wg.Done() + } for { data, isEof := w.WaitForData() // chunk the writes to make sure we send ample ack packets @@ -90,8 +98,10 @@ func (w *FdWriter) WriteLoop() { chunkSize := min(len(data), MaxSingleWriteSize) chunk := data[0:chunkSize] nw, err := w.Fd.Write(chunk) - ack := w.M.makeDataAckPacket(w.FdNum, nw, err) - w.M.sendPacket(ack) + if nw > 0 || err != nil { + ack := w.M.makeDataAckPacket(w.FdNum, nw, err) + w.M.sendPacket(ack) + } if err != nil { return } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index c3f82d61..31a7c577 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -45,17 +45,32 @@ func (m *Multiplexer) Close() { m.Lock.Lock() defer m.Lock.Unlock() - for _, fd := range m.FdReaders { - fd.Close() + for _, fr := range m.FdReaders { + fr.Close() } - for _, fd := range m.FdWriters { - fd.Close() + for _, fw := range m.FdWriters { + fw.Close() } for _, fd := range m.CloseAfterStart { fd.Close() } } +func (m *Multiplexer) HandleInputDone() { + m.Lock.Lock() + defer m.Lock.Unlock() + + // close readers (obviously the done command needs no more input) + for _, fr := range m.FdReaders { + fr.Close() + } + + // ensure EOF on all writers (ignore error) + for _, fw := range m.FdWriters { + fw.AddData(nil, true) + } +} + // returns the *writer* to connect to process, reader is put in FdReaders func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { pr, pw, err := os.Pipe() @@ -64,7 +79,7 @@ func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { } m.Lock.Lock() defer m.Lock.Unlock() - m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum) + m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum, true) m.CloseAfterStart = append(m.CloseAfterStart, pw) return pw, nil } @@ -77,11 +92,23 @@ func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { } m.Lock.Lock() defer m.Lock.Unlock() - m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum) + m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum, true) m.CloseAfterStart = append(m.CloseAfterStart, pr) return pr, nil } +func (m *Multiplexer) MakeRawFdReader(fdNum int, fd *os.File) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, false) +} + +func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd *os.File) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, false) +} + func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { ack := packet.MakeDataAckPacket() ack.SessionId = m.SessionId @@ -110,18 +137,23 @@ func (m *Multiplexer) sendPacket(p packet.PacketType) { m.Sender.SendPacket(p) } -func (m *Multiplexer) launchWriters() { +func (m *Multiplexer) launchWriters(wg *sync.WaitGroup) { m.Lock.Lock() defer m.Lock.Unlock() + if wg != nil { + wg.Add(len(m.FdWriters)) + } for _, fw := range m.FdWriters { - go fw.WriteLoop() + go fw.WriteLoop(wg) } } func (m *Multiplexer) launchReaders(wg *sync.WaitGroup) { m.Lock.Lock() defer m.Lock.Unlock() - wg.Add(len(m.FdReaders)) + if wg != nil { + wg.Add(len(m.FdReaders)) + } for _, fr := range m.FdReaders { go fr.ReadLoop(wg) } @@ -138,7 +170,8 @@ func (m *Multiplexer) startIO(packetCh chan packet.PacketType, sender *packet.Pa m.Started = true } -func (m *Multiplexer) runPacketInputLoop() { +func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { + defer m.HandleInputDone() for pk := range m.Input { if pk.GetType() == packet.DataPacketStr { dataPacket := pk.(*packet.DataPacketType) @@ -152,9 +185,15 @@ func (m *Multiplexer) runPacketInputLoop() { if pk.GetType() == packet.DataAckPacketStr { ackPacket := pk.(*packet.DataAckPacketType) m.processAckPacket(ackPacket) + continue + } + if pk.GetType() == packet.CmdDonePacketStr { + donePacket := pk.(*packet.CmdDonePacketType) + return donePacket } // other packet types are ignored } + return nil } func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { @@ -163,7 +202,7 @@ func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error fw := m.FdWriters[dataPacket.FdNum] if fw == nil { // add a closed FdWriter as a placeholder so we only send one error - fw := MakeFdWriter(m, nil, dataPacket.FdNum) + fw := MakeFdWriter(m, nil, dataPacket.FdNum, false) fw.Close() m.FdWriters[dataPacket.FdNum] = fw return fmt.Errorf("write to closed file") @@ -195,12 +234,38 @@ func (m *Multiplexer) closeTempStartFds() { m.CloseAfterStart = nil } -func (m *Multiplexer) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { +func (m *Multiplexer) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender, waitOnReaders bool, waitOnWriters bool, waitForInputLoop bool) *packet.CmdDonePacketType { m.startIO(packetCh, sender) m.closeTempStartFds() var wg sync.WaitGroup - m.launchReaders(&wg) - m.launchWriters() - go m.runPacketInputLoop() + if waitOnReaders { + m.launchReaders(&wg) + } else { + m.launchReaders(nil) + } + if waitOnWriters { + m.launchWriters(&wg) + } else { + m.launchWriters(nil) + } + var donePacket *packet.CmdDonePacketType + if waitForInputLoop { + wg.Add(1) + } + go func() { + if waitForInputLoop { + defer wg.Done() + } + pkRtn := m.runPacketInputLoop() + if pkRtn != nil { + m.Lock.Lock() + donePacket = pkRtn + m.Lock.Unlock() + } + }() wg.Wait() + + m.Lock.Lock() + defer m.Lock.Unlock() + return donePacket } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 5f6f3bd2..9217b17f 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -23,6 +23,8 @@ import ( // server: init, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] // all: error, message +var GlobalDebug = false + const ( RunPacketStr = "run" PingPacketStr = "ping" @@ -353,7 +355,7 @@ type RunPacketType struct { Command string `json:"command"` Cwd string `json:"cwd,omitempty"` Env map[string]string `json:"env,omitempty"` - TermSize TermSize `json:"termsize,omitempty"` + TermSize *TermSize `json:"termsize,omitempty"` Fds []RemoteFd `json:"fds,omitempty"` Detached bool `json:"detached,omitempty"` } @@ -430,6 +432,10 @@ func SendPacket(w io.Writer, packet PacketType) error { outBuf.WriteString(fmt.Sprintf("##%d", len(jsonBytes))) outBuf.Write(jsonBytes) outBuf.WriteByte('\n') + if GlobalDebug { + outBytes := outBuf.Bytes() + fmt.Printf("SEND>%s", string(outBytes[1:])) + } _, err = w.Write(outBuf.Bytes()) if err != nil { return err diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index e8fc52ad..cc372d92 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -33,19 +33,21 @@ const FirstExtraFilesFdNum = 3 type ShExecType struct { Lock *sync.Mutex StartTs time.Time - RunPacket *packet.RunPacketType + SessionId string + CmdId string FileNames *base.CommandFileNames Cmd *exec.Cmd CmdPty *os.File Multiplexer *mpio.Multiplexer } -func MakeShExec(pk *packet.RunPacketType) *ShExecType { +func MakeShExec(sessionId string, cmdId string) *ShExecType { return &ShExecType{ Lock: &sync.Mutex{}, StartTs: time.Now(), - RunPacket: pk, - Multiplexer: mpio.MakeMultiplexer(pk.SessionId, pk.CmdId), + SessionId: sessionId, + CmdId: cmdId, + Multiplexer: mpio.MakeMultiplexer(sessionId, cmdId), } } @@ -59,8 +61,8 @@ func (c *ShExecType) Close() { func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { startPacket := packet.MakeCmdStartPacket() startPacket.Ts = time.Now().UnixMilli() - startPacket.SessionId = c.RunPacket.SessionId - startPacket.CmdId = c.RunPacket.CmdId + startPacket.SessionId = c.SessionId + startPacket.CmdId = c.CmdId startPacket.Pid = c.Cmd.Process.Pid startPacket.MShellPid = os.Getpid() return startPacket @@ -209,14 +211,89 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } } -func (cmd *ShExecType) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { - cmd.Multiplexer.RunIOAndWait(packetCh, sender) +type ClientOpts struct { + IsSSH bool + SSHOptsTerm bool + SSHOpts []string + Command string + Fds []packet.RemoteFd + Cwd string +} + +func (opts *ClientOpts) MakeRunPacket() *packet.RunPacketType { + runPacket := packet.MakeRunPacket() + runPacket.Command = opts.Command + runPacket.Cwd = opts.Cwd + runPacket.Fds = opts.Fds + return runPacket +} + +func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, error) { + // packet.GlobalDebug = true + cmd := MakeShExec("", "") + sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` + var fullSshOpts []string + fullSshOpts = append(fullSshOpts, opts.SSHOpts...) + fullSshOpts = append(fullSshOpts, sshRemoteCommand) + ecmd := exec.Command("ssh", fullSshOpts...) + cmd.Cmd = ecmd + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return nil, fmt.Errorf("creating stdin pipe: %v", err) + } + stdoutReader, err := ecmd.StdoutPipe() + if err != nil { + return nil, fmt.Errorf("creating stdout pipe: %v", err) + } + stderrReader, err := ecmd.StderrPipe() + if err != nil { + return nil, fmt.Errorf("creating stderr pipe: %v", err) + } + err = ecmd.Start() + if err != nil { + return nil, fmt.Errorf("running ssh command: %w", err) + } + defer cmd.Close() + packetCh := packet.PacketParser(stdoutReader) + go func() { + io.Copy(os.Stderr, stderrReader) + }() + sender := packet.MakePacketSender(inputWriter) + for pk := range packetCh { + if pk.GetType() == packet.RawPacketStr { + rawPk := pk.(*packet.RawPacketType) + fmt.Printf("%s\n", rawPk.Data) + continue + } + if pk.GetType() == packet.InitPacketStr { + initPk := pk.(*packet.InitPacketType) + if initPk.Version != "0.1.0" { + return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) + } + break + } + } + runPacket := opts.MakeRunPacket() + sender.SendPacket(runPacket) + cmd.Multiplexer.MakeRawFdReader(0, os.Stdin) + cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout) + cmd.Multiplexer.MakeRawFdWriter(2, os.Stderr) + remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetCh, sender, false, true, true) + donePacket := cmd.WaitForCommand() + if remoteDonePacket != nil { + donePacket = remoteDonePacket + } + return donePacket, nil +} + +func (cmd *ShExecType) RunRemoteIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { + cmd.Multiplexer.RunIOAndWait(packetCh, sender, true, false, false) donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) } func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - cmd := MakeShExec(pk) + cmd := MakeShExec(pk.SessionId, pk.CmdId) cmd.Cmd = exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(cmd.Cmd, pk.Env) if pk.Cwd != "" { @@ -316,7 +393,7 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( defer func() { cmdTty.Close() }() - rtn := MakeShExec(pk) + rtn := MakeShExec(pk.SessionId, pk.CmdId) ecmd := MakeExecCmd(pk, cmdTty) err = ecmd.Start() if err != nil { @@ -364,8 +441,8 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { exitCode := GetExitCode(exitErr) donePacket := packet.MakeCmdDonePacket() donePacket.Ts = endTs.UnixMilli() - donePacket.SessionId = c.RunPacket.SessionId - donePacket.CmdId = c.RunPacket.CmdId + donePacket.SessionId = c.SessionId + donePacket.CmdId = c.CmdId donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) if c.FileNames != nil { From e6776bd97493b0320f98c478de7760b8744b5292 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 24 Jun 2022 23:42:00 -0700 Subject: [PATCH 022/149] checkpoint. transfer binary data as base64. handle cwd. detect open fds. working to transfer data in non-error cases. --- main-mshell.go | 40 ++++++++++++++---- pkg/base/base.go | 12 ++++++ pkg/mpio/bufreader.go | 6 +++ pkg/mpio/bufwriter.go | 12 +++--- pkg/mpio/mpio.go | 40 +++++++++++++----- pkg/packet/packet.go | 67 +++++++++++++++++++++++++----- pkg/shexec/shexec.go | 94 ++++++++++++++++++++++++++----------------- 7 files changed, 202 insertions(+), 69 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 418769f2..226c6452 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -20,6 +20,7 @@ import ( "github.com/scripthaus-dev/mshell/pkg/cmdtail" "github.com/scripthaus-dev/mshell/pkg/packet" "github.com/scripthaus-dev/mshell/pkg/shexec" + "golang.org/x/sys/unix" ) const MShellVersion = "0.1.0" @@ -256,8 +257,26 @@ func handleRemote() { func handleServer() { } -func detectOpenFds() { - +func detectOpenFds() ([]packet.RemoteFd, error) { + var fds []packet.RemoteFd + for fdNum := 3; fdNum <= 64; fdNum++ { + flags, err := unix.FcntlInt(uintptr(fdNum), unix.F_GETFL, 0) + if err != nil { + continue + } + flags = flags & 3 + rfd := packet.RemoteFd{FdNum: fdNum} + if flags&2 == 2 { + return nil, fmt.Errorf("invalid fd=%d, mshell does not support fds open for reading and writing", fdNum) + } + if flags&1 == 1 { + rfd.Write = true + } else { + rfd.Read = true + } + fds = append(fds, rfd) + } + return fds, nil } func parseClientOpts() (*shexec.ClientOpts, error) { @@ -272,6 +291,13 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.IsSSH = true break } + if argStr == "--cwd" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--cwd [dir]' missing directory") + } + opts.Cwd = iter.Next() + continue + } } if opts.IsSSH { // parse SSH opts @@ -281,11 +307,6 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.SSHOptsTerm = true break } - if argStr == "--cwd" { - if !iter.HasNext() { - return nil, fmt.Errorf("'--cwd [dir]' missing directory") - } - } opts.SSHOpts = append(opts.SSHOpts, argStr) } if !opts.SSHOptsTerm { @@ -310,6 +331,11 @@ func handleClient() (int, error) { if !opts.IsSSH { return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") } + fds, err := detectOpenFds() + if err != nil { + return 1, err + } + opts.Fds = fds donePacket, err := shexec.RunClientSSHCommandAndWait(opts) if err != nil { return 1, err diff --git a/pkg/base/base.go b/pkg/base/base.go index dc13157f..eae9a8cc 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -14,6 +14,7 @@ import ( "os/exec" "path" "path/filepath" + "strings" ) const DefaultMShellPath = "mshell" @@ -176,3 +177,14 @@ func WriteErrorMsg(fileName string, errVal string) error { _, writeErr := fd.Write([]byte(oscEsc)) return writeErr } + +func ExpandHomeDir(pathStr string) string { + if pathStr != "~" && !strings.HasPrefix(pathStr, "~/") { + return pathStr + } + homeDir := GetHomeDir() + if pathStr == "~" { + return homeDir + } + return path.Join(homeDir, pathStr[2:]) +} diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index 756f9e83..654d8d68 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -48,6 +48,12 @@ func (r *FdReader) Close() { r.CVar.Broadcast() } +func (r *FdReader) GetBufSize() int { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + return r.BufSize +} + func (r *FdReader) NotifyAck(ackLen int) { r.CVar.L.Lock() defer r.CVar.L.Unlock() diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index 9b389678..c1b36e2b 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -64,15 +64,15 @@ func (w *FdWriter) WaitForData() ([]byte, bool) { func (w *FdWriter) AddData(data []byte, eof bool) error { w.CVar.L.Lock() defer w.CVar.L.Unlock() - if w.Closed { - return fmt.Errorf("write to closed file") - } - if w.Eof { - return fmt.Errorf("write to closed file (eof)") + if w.Closed || w.Eof { + if len(data) == 0 { + return nil + } + return fmt.Errorf("write to closed file eof[%v]", w.Eof) } if len(data) > 0 { if len(data)+len(w.Buffer) > WriteBufSize { - return fmt.Errorf("write exceeds buffer size") + return fmt.Errorf("write exceeds buffer size bufsize=%d (max=%d)", len(data)+len(w.Buffer), WriteBufSize) } w.Buffer = append(w.Buffer, data...) } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 31a7c577..f911e762 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -7,6 +7,7 @@ package mpio import ( + "encoding/base64" "fmt" "os" "sync" @@ -14,8 +15,8 @@ import ( "github.com/scripthaus-dev/mshell/pkg/packet" ) -const ReadBufSize = 128 * 1024 -const WriteBufSize = 128 * 1024 +const ReadBufSize = 32 * 1024 +const WriteBufSize = 32 * 1024 const MaxSingleWriteSize = 4 * 1024 type Multiplexer struct { @@ -29,6 +30,8 @@ type Multiplexer struct { Sender *packet.PacketSender Input chan packet.PacketType Started bool + + Debug bool } func MakeMultiplexer(sessionId string, cmdId string) *Multiplexer { @@ -97,16 +100,16 @@ func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { return pr, nil } -func (m *Multiplexer) MakeRawFdReader(fdNum int, fd *os.File) { +func (m *Multiplexer) MakeRawFdReader(fdNum int, fd *os.File, shouldClose bool) { m.Lock.Lock() defer m.Lock.Unlock() - m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, false) + m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose) } -func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd *os.File) { +func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd *os.File, shouldClose bool) { m.Lock.Lock() defer m.Lock.Unlock() - m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, false) + m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, shouldClose) } func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { @@ -126,7 +129,7 @@ func (m *Multiplexer) makeDataPacket(fdNum int, data []byte, err error) *packet. pk.SessionId = m.SessionId pk.CmdId = m.CmdId pk.FdNum = fdNum - pk.Data = string(data) + pk.Data64 = base64.StdEncoding.EncodeToString(data) if err != nil { pk.Error = err.Error() } @@ -173,6 +176,9 @@ func (m *Multiplexer) startIO(packetCh chan packet.PacketType, sender *packet.Pa func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { defer m.HandleInputDone() for pk := range m.Input { + if m.Debug { + fmt.Printf("PK> %s\n", packet.AsString(pk)) + } if pk.GetType() == packet.DataPacketStr { dataPacket := pk.(*packet.DataPacketType) err := m.processDataPacket(dataPacket) @@ -191,12 +197,26 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { donePacket := pk.(*packet.CmdDonePacketType) return donePacket } - // other packet types are ignored + if pk.GetType() == packet.ErrorPacketStr { + errPacket := pk.(*packet.ErrorPacketType) + // at this point, just send the error packet to stderr rather than try to do something special + fmt.Fprintf(os.Stderr, "%s\n", errPacket.Error) + return nil + } + if pk.GetType() == packet.RawPacketStr { + rawPacket := pk.(*packet.RawPacketType) + fmt.Fprintf(os.Stderr, "%s\n", rawPacket.Data) + continue + } } return nil } func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { + realData, err := base64.StdEncoding.DecodeString(dataPacket.Data64) + if err != nil { + return fmt.Errorf("decoding base64 data: %w", err) + } m.Lock.Lock() defer m.Lock.Unlock() fw := m.FdWriters[dataPacket.FdNum] @@ -205,9 +225,9 @@ func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error fw := MakeFdWriter(m, nil, dataPacket.FdNum, false) fw.Close() m.FdWriters[dataPacket.FdNum] = fw - return fmt.Errorf("write to closed file") + return fmt.Errorf("write to closed file (no fd)") } - err := fw.AddData([]byte(dataPacket.Data), dataPacket.Eof) + err = fw.AddData(realData, dataPacket.Eof) if err != nil { fw.Close() return err diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 9217b17f..a7674518 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -121,7 +121,7 @@ type DataPacketType struct { SessionId string `json:"sessionid,omitempty"` CmdId string `json:"cmdid,omitempty"` FdNum int `json:"fdnum"` - Data string `json:"data"` + Data64 string `json:"data64"` // base64 encoded Eof bool `json:"eof,omitempty"` Error string `json:"error,omitempty"` } @@ -130,6 +130,32 @@ func (*DataPacketType) GetType() string { return DataPacketStr } +func B64DecodedLen(b64 string) int { + if len(b64) < 4 { + return 0 // we use padded strings, so < 4 is always 0 + } + realLen := 3 * (len(b64) / 4) + if b64[len(b64)-1] == '=' { + realLen-- + } + if b64[len(b64)-2] == '=' { + realLen-- + } + return realLen +} + +func (p *DataPacketType) String() string { + eofStr := "" + if p.Eof { + eofStr = ", eof" + } + errStr := "" + if p.Error != "" { + errStr = fmt.Sprintf(", err=%s", p.Error) + } + return fmt.Sprintf("data[fd=%d, len=%d%s%s]", p.FdNum, B64DecodedLen(p.Data64), eofStr, errStr) +} + func MakeDataPacket() *DataPacketType { return &DataPacketType{Type: DataPacketStr} } @@ -140,13 +166,21 @@ type DataAckPacketType struct { CmdId string `json:"cmdid,omitempty"` FdNum int `json:"fdnum"` AckLen int `json:"acklen"` - Error string `json:"error"` + Error string `json:"error,omitempty"` } func (*DataAckPacketType) GetType() string { return DataAckPacketStr } +func (p *DataAckPacketType) String() string { + errStr := "" + if p.Error != "" { + errStr = fmt.Sprintf(" err=%s", p.Error) + } + return fmt.Sprintf("ack[fd=%d, acklen=%d%s]", p.FdNum, p.AckLen, errStr) +} + func MakeDataAckPacket() *DataAckPacketType { return &DataAckPacketType{Type: DataAckPacketStr} } @@ -252,6 +286,10 @@ func (*RawPacketType) GetType() string { return RawPacketStr } +func (p *RawPacketType) String() string { + return fmt.Sprintf("raw[%s]", p.Data) +} + func MakeRawPacket(val string) *RawPacketType { return &RawPacketType{Type: RawPacketStr, Data: val} } @@ -265,6 +303,10 @@ func (*MessagePacketType) GetType() string { return MessagePacketStr } +func (p *MessagePacketType) String() string { + return fmt.Sprintf("messsage[%s]", p.Message) +} + func MakeMessagePacket(message string) *MessagePacketType { return &MessagePacketType{Type: MessagePacketStr, Message: message} } @@ -394,6 +436,13 @@ type PacketType interface { GetType() string } +func AsString(pk PacketType) string { + if s, ok := pk.(fmt.Stringer); ok { + return s.String() + } + return fmt.Sprintf("%s[]", pk.GetType()) +} + type RpcPacketType interface { GetType() string GetPacketId() string @@ -433,8 +482,7 @@ func SendPacket(w io.Writer, packet PacketType) error { outBuf.Write(jsonBytes) outBuf.WriteByte('\n') if GlobalDebug { - outBytes := outBuf.Bytes() - fmt.Printf("SEND>%s", string(outBytes[1:])) + fmt.Printf("SEND> %s\n", AsString(packet)) } _, err = w.Write(outBuf.Bytes()) if err != nil { @@ -519,12 +567,14 @@ func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) erro } func PacketParser(input io.Reader) chan PacketType { - bufReader := bufio.NewReader(input) rtnCh := make(chan PacketType) + PacketParserAttach(input, rtnCh) + return rtnCh +} + +func PacketParserAttach(input io.Reader, rtnCh chan PacketType) { + bufReader := bufio.NewReader(input) go func() { - defer func() { - close(rtnCh) - }() for { line, err := bufReader.ReadString('\n') if err == io.EOF { @@ -562,7 +612,6 @@ func PacketParser(input io.Reader) chan PacketType { rtnCh <- pk } }() - return rtnCh } type ErrorReporter interface { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index cc372d92..82aa31b4 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -111,7 +111,7 @@ func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { ecmd := exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(ecmd, pk.Env) if pk.Cwd != "" { - ecmd.Dir = pk.Cwd + ecmd.Dir = base.ExpandHomeDir(pk.Cwd) } ecmd.Stdin = cmdTty ecmd.Stdout = cmdTty @@ -175,12 +175,13 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { } } if pk.Cwd != "" { - dirInfo, err := os.Stat(pk.Cwd) + realCwd := base.ExpandHomeDir(pk.Cwd) + dirInfo, err := os.Stat(realCwd) if err != nil { - return fmt.Errorf("invalid cwd '%s' for command: %v", pk.Cwd, err) + return fmt.Errorf("invalid cwd '%s' for command: %v", realCwd, err) } if !dirInfo.IsDir() { - return fmt.Errorf("invalid cwd '%s' for command, not a directory", pk.Cwd) + return fmt.Errorf("invalid cwd '%s' for command, not a directory", realCwd) } } return nil @@ -228,8 +229,37 @@ func (opts *ClientOpts) MakeRunPacket() *packet.RunPacketType { return runPacket } +func ValidateRemoteFds(rfds []packet.RemoteFd) error { + dupMap := make(map[int]bool) + for _, rfd := range rfds { + if rfd.FdNum < 0 { + return fmt.Errorf("mshell negative fd numbers fd=%d", rfd.FdNum) + } + if rfd.FdNum < FirstExtraFilesFdNum { + return fmt.Errorf("mshell does not support re-opening fd=%d (0, 1, and 2, are always open)", rfd.FdNum) + } + if rfd.FdNum > MaxFdNum { + return fmt.Errorf("mshell does not support opening fd numbers above %d", MaxFdNum) + } + if dupMap[rfd.FdNum] { + return fmt.Errorf("mshell got duplicate entries for fd=%d", rfd.FdNum) + } + if rfd.Read && rfd.Write { + return fmt.Errorf("mshell does not support opening fd numbers for reading and writing, fd=%d", rfd.FdNum) + } + if !rfd.Read && !rfd.Write { + return fmt.Errorf("invalid fd=%d, neither reading or writing mode specified", rfd.FdNum) + } + dupMap[rfd.FdNum] = true + } + return nil +} + func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, error) { - // packet.GlobalDebug = true + err := ValidateRemoteFds(opts.Fds) + if err != nil { + return nil, err + } cmd := MakeShExec("", "") sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` var fullSshOpts []string @@ -249,15 +279,27 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if err != nil { return nil, fmt.Errorf("creating stderr pipe: %v", err) } + cmd.Multiplexer.MakeRawFdReader(0, os.Stdin, false) + cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout, false) + cmd.Multiplexer.MakeRawFdWriter(2, os.Stderr, false) + for _, rfd := range opts.Fds { + fd := os.NewFile(uintptr(rfd.FdNum), fmt.Sprintf("/dev/fd/%d", rfd.FdNum)) + if fd == nil { + return nil, fmt.Errorf("cannot open fd %d", rfd.FdNum) + } + if rfd.Read { + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, true) + } else if rfd.Write { + cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true) + } + } err = ecmd.Start() if err != nil { return nil, fmt.Errorf("running ssh command: %w", err) } defer cmd.Close() packetCh := packet.PacketParser(stdoutReader) - go func() { - io.Copy(os.Stderr, stderrReader) - }() + packet.PacketParserAttach(stderrReader, packetCh) sender := packet.MakePacketSender(inputWriter) for pk := range packetCh { if pk.GetType() == packet.RawPacketStr { @@ -275,9 +317,6 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er } runPacket := opts.MakeRunPacket() sender.SendPacket(runPacket) - cmd.Multiplexer.MakeRawFdReader(0, os.Stdin) - cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout) - cmd.Multiplexer.MakeRawFdWriter(2, os.Stderr) remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetCh, sender, false, true, true) donePacket := cmd.WaitForCommand() if remoteDonePacket != nil { @@ -287,6 +326,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er } func (cmd *ShExecType) RunRemoteIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { + defer cmd.Close() cmd.Multiplexer.RunIOAndWait(packetCh, sender, true, false, false) donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) @@ -297,9 +337,13 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmd.Cmd = exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(cmd.Cmd, pk.Env) if pk.Cwd != "" { - cmd.Cmd.Dir = pk.Cwd + cmd.Cmd.Dir = base.ExpandHomeDir(pk.Cwd) + } + err := ValidateRemoteFds(pk.Fds) + if err != nil { + cmd.Close() + return nil, err } - var err error cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) if err != nil { cmd.Close() @@ -317,33 +361,9 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S } extraFiles := make([]*os.File, 0, MaxFdNum+1) for _, rfd := range pk.Fds { - if rfd.FdNum < 0 { - cmd.Close() - return nil, fmt.Errorf("mshell negative fd numbers fd=%d", rfd.FdNum) - } - if rfd.FdNum < FirstExtraFilesFdNum { - cmd.Close() - return nil, fmt.Errorf("mshell does not support re-opening fd=%d (0, 1, and 2, are always open)", rfd.FdNum) - } - if rfd.FdNum > MaxFdNum { - cmd.Close() - return nil, fmt.Errorf("mshell does not support opening fd numbers above %d", MaxFdNum) - } if rfd.FdNum >= len(extraFiles) { extraFiles = extraFiles[:rfd.FdNum+1] } - if extraFiles[rfd.FdNum] != nil { - cmd.Close() - return nil, fmt.Errorf("mshell got duplicate entries for fd=%d", rfd.FdNum) - } - if rfd.Read && rfd.Write { - cmd.Close() - return nil, fmt.Errorf("mshell does not support opening fd numbers for reading and writing, fd=%d", rfd.FdNum) - } - if !rfd.Read && !rfd.Write { - cmd.Close() - return nil, fmt.Errorf("invalid fd=%d, neither reading or writing mode specified", rfd.FdNum) - } if rfd.Read { // client file is open for reading, so we make a writer pipe extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum) From 43ed95f8fc40831c1fd8ee0c492d90bec6357274 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 25 Jun 2022 00:05:37 -0700 Subject: [PATCH 023/149] clean up combining stdout and stderr into packet parsers, combine the channels and close appropriately --- pkg/packet/packet.go | 28 +++++++++++++++++++++++++--- pkg/shexec/shexec.go | 5 +++-- 2 files changed, 28 insertions(+), 5 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index a7674518..5a374a25 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -566,15 +566,36 @@ func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) erro return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) } -func PacketParser(input io.Reader) chan PacketType { +func CombinePacketParsers(p1 chan PacketType, p2 chan PacketType) chan PacketType { rtnCh := make(chan PacketType) - PacketParserAttach(input, rtnCh) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for v := range p1 { + rtnCh <- v + } + }() + go func() { + defer wg.Done() + for v := range p2 { + rtnCh <- v + } + }() + go func() { + wg.Wait() + close(rtnCh) + }() return rtnCh } -func PacketParserAttach(input io.Reader, rtnCh chan PacketType) { +func PacketParser(input io.Reader) chan PacketType { + rtnCh := make(chan PacketType) bufReader := bufio.NewReader(input) go func() { + defer func() { + close(rtnCh) + }() for { line, err := bufReader.ReadString('\n') if err == io.EOF { @@ -612,6 +633,7 @@ func PacketParserAttach(input io.Reader, rtnCh chan PacketType) { rtnCh <- pk } }() + return rtnCh } type ErrorReporter interface { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 82aa31b4..9bf98a56 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -298,8 +298,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, fmt.Errorf("running ssh command: %w", err) } defer cmd.Close() - packetCh := packet.PacketParser(stdoutReader) - packet.PacketParserAttach(stderrReader, packetCh) + stdoutPacketCh := packet.PacketParser(stdoutReader) + stderrPacketCh := packet.PacketParser(stderrReader) + packetCh := packet.CombinePacketParsers(stdoutPacketCh, stderrPacketCh) sender := packet.MakePacketSender(inputWriter) for pk := range packetCh { if pk.GetType() == packet.RawPacketStr { From 935500f1f1d5218d3908758f626daab9bc944fd1 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 25 Jun 2022 00:22:03 -0700 Subject: [PATCH 024/149] packet debugging with --debug --- main-mshell.go | 4 ++++ pkg/shexec/shexec.go | 4 ++++ 2 files changed, 8 insertions(+) diff --git a/main-mshell.go b/main-mshell.go index 226c6452..7cb4836b 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -298,6 +298,10 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.Cwd = iter.Next() continue } + if argStr == "--debug" { + opts.Debug = true + continue + } } if opts.IsSSH { // parse SSH opts diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 9bf98a56..74b98aa2 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -219,6 +219,7 @@ type ClientOpts struct { Command string Fds []packet.RemoteFd Cwd string + Debug bool } func (opts *ClientOpts) MakeRunPacket() *packet.RunPacketType { @@ -318,6 +319,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er } runPacket := opts.MakeRunPacket() sender.SendPacket(runPacket) + if opts.Debug { + cmd.Multiplexer.Debug = true + } remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetCh, sender, false, true, true) donePacket := cmd.WaitForCommand() if remoteDonePacket != nil { From fec7721e32e1e737a2aee975ce935ddc13c49202 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 25 Jun 2022 00:30:41 -0700 Subject: [PATCH 025/149] refuse to run with ssh -t or -tt, detect nil runPacket --- main-mshell.go | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/main-mshell.go b/main-mshell.go index 7cb4836b..6eebbb5f 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -243,6 +243,10 @@ func handleRemote() { sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) return } + if runPacket == nil { + sender.SendErrorPacket(fmt.Sprintf("no run packet received")) + return + } cmd, err := shexec.RunCommand(runPacket, sender) if err != nil { sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) @@ -311,6 +315,9 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.SSHOptsTerm = true break } + if argStr == "-t" || argStr == "-tt" { + return nil, fmt.Errorf("mshell cannot run over ssh -t") + } opts.SSHOpts = append(opts.SSHOpts, argStr) } if !opts.SSHOptsTerm { From e8ae01efaea4552e94fa843403376977cba3cdf7 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 25 Jun 2022 00:33:18 -0700 Subject: [PATCH 026/149] check for input termination before init packet --- pkg/shexec/shexec.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 74b98aa2..25b096bc 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -303,6 +303,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er stderrPacketCh := packet.PacketParser(stderrReader) packetCh := packet.CombinePacketParsers(stdoutPacketCh, stderrPacketCh) sender := packet.MakePacketSender(inputWriter) + versionOk := false for pk := range packetCh { if pk.GetType() == packet.RawPacketStr { rawPk := pk.(*packet.RawPacketType) @@ -314,9 +315,13 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if initPk.Version != "0.1.0" { return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) } + versionOk = true break } } + if !versionOk { + return nil, fmt.Errorf("did not receive version from remote mshell") + } runPacket := opts.MakeRunPacket() sender.SendPacket(runPacket) if opts.Debug { From 222deff0db7bf1de45af7f53eccfdfc6a6a2658a Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 26 Jun 2022 01:41:58 -0700 Subject: [PATCH 027/149] implement sudo dance allowing passing the sudo password on stdin with sudo -S, and passing a different stdin fd to the command --- main-mshell.go | 30 +++++++++++ pkg/mpio/mpio.go | 12 +++++ pkg/packet/packet.go | 8 +-- pkg/shexec/shexec.go | 117 +++++++++++++++++++++++++++++++++++++------ 4 files changed, 149 insertions(+), 18 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 6eebbb5f..7a4447b4 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -306,6 +306,33 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.Debug = true continue } + if argStr == "--sudo" { + opts.Sudo = true + continue + } + if argStr == "--sudo-with-password" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--sudo-with-password [pw]', missing password") + } + opts.Sudo = true + opts.SudoWithPass = true + opts.SudoPw = iter.Next() + continue + } + if argStr == "--sudo-with-passfile" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--sudo-with-passfile [file]', missing file") + } + opts.Sudo = true + opts.SudoWithPass = true + fileName := iter.Next() + contents, err := os.ReadFile(fileName) + if err != nil { + return nil, fmt.Errorf("cannot read --sudo-with-passfile file '%s': %w", fileName, err) + } + opts.SudoPw = string(contents) + continue + } } if opts.IsSSH { // parse SSH opts @@ -339,6 +366,9 @@ func handleClient() (int, error) { if err != nil { return 1, fmt.Errorf("parsing opts: %w", err) } + if opts.Debug { + packet.GlobalDebug = true + } if !opts.IsSSH { return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index f911e762..a1b93bd4 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -100,6 +100,18 @@ func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { return pr, nil } +func (m *Multiplexer) MakeStringFdReader(fdNum int, contents string) error { + pw, err := m.MakeReaderPipe(fdNum) + if err != nil { + return err + } + go func() { + pw.Write([]byte(contents)) + pw.Close() + }() + return nil +} + func (m *Multiplexer) MakeRawFdReader(fdNum int, fd *os.File, shouldClose bool) { m.Lock.Lock() defer m.Lock.Unlock() diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 5a374a25..3c6d9dce 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -385,9 +385,11 @@ type TermSize struct { } type RemoteFd struct { - FdNum int `json:"fdnum"` - Read bool `json:"read"` - Write bool `json:"write"` + FdNum int `json:"fdnum"` + Read bool `json:"read"` + Write bool `json:"write"` + Content string `json:"-"` + DupStdin bool `json:"-"` } type RunPacketType struct { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 25b096bc..73afa925 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -30,6 +30,12 @@ const MaxCols = 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 +const SSHRemoteCommand = `PATH=$PATH:~/.mshell; mshell --remote` + +const RemoteCommandFmt = `%s` +const RemoteSudoCommandFmt = `sudo -C %d bash /dev/fd/%d` +const RemoteSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -S -C %d bash -c "echo '[from-mshell]'; bash /dev/fd/%d < /dev/fd/%d"` + type ShExecType struct { Lock *sync.Mutex StartTs time.Time @@ -213,21 +219,87 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } type ClientOpts struct { - IsSSH bool - SSHOptsTerm bool - SSHOpts []string - Command string - Fds []packet.RemoteFd - Cwd string - Debug bool + IsSSH bool + SSHOptsTerm bool + SSHOpts []string + Command string + Fds []packet.RemoteFd + Cwd string + Debug bool + Sudo bool + SudoWithPass bool + SudoPw string + CommandStdinFdNum int } -func (opts *ClientOpts) MakeRunPacket() *packet.RunPacketType { +func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket := packet.MakeRunPacket() - runPacket.Command = opts.Command runPacket.Cwd = opts.Cwd runPacket.Fds = opts.Fds - return runPacket + if !opts.Sudo { + // normal, non-sudo command + runPacket.Command = opts.Command + return runPacket, nil + } + if opts.SudoWithPass { + pwFdNum, err := opts.NextFreeFdNum() + if err != nil { + return nil, err + } + pwRfd := packet.RemoteFd{FdNum: pwFdNum, Read: true, Content: opts.SudoPw} + opts.Fds = append(opts.Fds, pwRfd) + commandFdNum, err := opts.NextFreeFdNum() + if err != nil { + return nil, err + } + commandRfd := packet.RemoteFd{FdNum: commandFdNum, Read: true, Content: opts.Command} + opts.Fds = append(opts.Fds, commandRfd) + commandStdinFdNum, err := opts.NextFreeFdNum() + if err != nil { + return nil, err + } + commandStdinRfd := packet.RemoteFd{FdNum: commandStdinFdNum, Read: true, DupStdin: true} + opts.Fds = append(opts.Fds, commandStdinRfd) + opts.CommandStdinFdNum = commandStdinFdNum + maxFdNum := opts.MaxFdNum() + runPacket.Command = fmt.Sprintf(RemoteSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, commandFdNum, commandStdinFdNum) + runPacket.Fds = opts.Fds + return runPacket, nil + } else { + commandFdNum, err := opts.NextFreeFdNum() + if err != nil { + return nil, err + } + rfd := packet.RemoteFd{FdNum: commandFdNum, Read: true, Content: opts.Command} + opts.Fds = append(opts.Fds, rfd) + maxFdNum := opts.MaxFdNum() + runPacket.Command = fmt.Sprintf(RemoteSudoCommandFmt, maxFdNum+1, commandFdNum) + runPacket.Fds = opts.Fds + return runPacket, nil + } +} + +func (opts *ClientOpts) NextFreeFdNum() (int, error) { + fdMap := make(map[int]bool) + for _, fd := range opts.Fds { + fdMap[fd.FdNum] = true + } + for i := 3; i <= MaxFdNum; i++ { + if !fdMap[i] { + return i, nil + } + } + return 0, fmt.Errorf("reached maximum number of fds, all fds between 3-%d are in use", MaxFdNum) +} + +func (opts *ClientOpts) MaxFdNum() int { + maxFdNum := 3 + for _, fd := range opts.Fds { + if fd.FdNum > maxFdNum { + maxFdNum = fd.FdNum + } + } + return maxFdNum } func ValidateRemoteFds(rfds []packet.RemoteFd) error { @@ -261,11 +333,14 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if err != nil { return nil, err } + runPacket, err := opts.MakeRunPacket() // modifies opts + if err != nil { + return nil, err + } cmd := MakeShExec("", "") - sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` var fullSshOpts []string fullSshOpts = append(fullSshOpts, opts.SSHOpts...) - fullSshOpts = append(fullSshOpts, sshRemoteCommand) + fullSshOpts = append(fullSshOpts, SSHRemoteCommand) ecmd := exec.Command("ssh", fullSshOpts...) cmd.Cmd = ecmd inputWriter, err := ecmd.StdinPipe() @@ -280,10 +355,23 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if err != nil { return nil, fmt.Errorf("creating stderr pipe: %v", err) } - cmd.Multiplexer.MakeRawFdReader(0, os.Stdin, false) + if !opts.SudoWithPass { + cmd.Multiplexer.MakeRawFdReader(0, os.Stdin, false) + } cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout, false) cmd.Multiplexer.MakeRawFdWriter(2, os.Stderr, false) - for _, rfd := range opts.Fds { + for _, rfd := range runPacket.Fds { + if rfd.Read && rfd.Content != "" { + err = cmd.Multiplexer.MakeStringFdReader(rfd.FdNum, rfd.Content) + if err != nil { + return nil, fmt.Errorf("creating content fd %d", rfd.FdNum) + } + continue + } + if rfd.Read && rfd.DupStdin { + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, os.Stdin, false) + continue + } fd := os.NewFile(uintptr(rfd.FdNum), fmt.Sprintf("/dev/fd/%d", rfd.FdNum)) if fd == nil { return nil, fmt.Errorf("cannot open fd %d", rfd.FdNum) @@ -322,7 +410,6 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if !versionOk { return nil, fmt.Errorf("did not receive version from remote mshell") } - runPacket := opts.MakeRunPacket() sender.SendPacket(runPacket) if opts.Debug { cmd.Multiplexer.Debug = true From 1ea839384480c7cb71305e3a0418295c35f2e326 Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 26 Jun 2022 01:53:07 -0700 Subject: [PATCH 028/149] update docs with sudo example --- main-mshell.go | 27 +++++++++++++-------------- 1 file changed, 13 insertions(+), 14 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 7a4447b4..fd9f03fd 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -391,33 +391,32 @@ Client Usage: mshell [mshell-opts] --ssh [ssh-opts] user@host -- [command] mshell multiplexes input and output streams to a remote command over ssh. Options: - --env 'X=Y;A=B' - set remote environment variables for command, semicolon separated - --env-file [file] - load environment variables from [file] (.env format) - --env-copy [glob] - copy local environment variables to remote using [glob] pattern --cwd [dir] - execute remote command in [dir] - --no-auto-fds - do not auto-detect additional fds - --sudo - execute "sudo [command]" - --fds [fdspec] - open fds based off [fdspec], comma separated (implies --no-auto-fds) - <[num] opens for reading - >[num] opens for writing - e.g. --fds '<5,>6,>7' [command] - a single argument (should be quoted) +Sudo Options: + --sudo + --sudo-with-password [pw] (not recommended, use --sudo-with-passfile if possible) + --sudo-with-passfile [file] + +Sudo options allow you to run the given command using "sudo". The first +option only works when you can sudo without a password. Your password will be passed +securely through a high numbered fd to "sudo -S". See full documentation for more details. + Examples: # execute a python script remotely, with stdin still hooked up correctly - mshell --cwd "~/work" --ssh -i key.pem ubuntu@somehost -- "python /dev/fd/4" 4< myscript.py + mshell --cwd "~/work" --ssh -i key.pem ubuntu@somehost -- "python3 /dev/fd/4" 4< myscript.py # capture multiple outputs mshell --ssh ubuntu@test -- "cat file1.txt > /dev/fd/3; cat file2.txt > /dev/fd/4" 3> file1.txt 4> file2.txt - # environment variable copying, setting working directory - # note the single quotes on command (otherwise the local shell will expand the variables) - TEST1=hello TEST2=world mshell --cwd "~/work" --env-copy "TEST*" --ssh user@host -- 'echo $(pwd) $TEST1 $TEST2' - # execute a script, catpure stdout/stderr in fd-3 and fd-4 # useful if you need to see stdout for interacting with ssh (password or host auth) mshell --ssh user@host -- "test.sh > /dev/fd/3 2> /dev/fd/4" 3> test.stdout 4> test.stderr + # run a script as root (via sudo), capture output + mshell --sudo-with-passfile pw.txt --ssh ubuntu@somehost -- "python3 /dev/fd/3 > /dev/fd/4" 3< myscript.py 4> script-output.txt < script-input.txt + mshell is licensed under the MPLv2 Please see https://github.com/scripthaus-dev/mshell for extended usage modes, source code, bugs, and feature requests ` From 2a6791bcd620367bbf734acdf6ada183002b6eae Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 12:03:47 -0700 Subject: [PATCH 029/149] combine sessionid and cmdid into one field ck (commandkey) --- main-mshell.go | 42 +++++----- pkg/base/base.go | 75 ++++++++++++++++-- pkg/cmdtail/cmdtail.go | 55 +++++-------- pkg/mpio/mpio.go | 15 ++-- pkg/packet/packet.go | 176 ++++++++++++++++++++++++----------------- pkg/shexec/shexec.go | 45 ++++------- 6 files changed, 236 insertions(+), 172 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index fd9f03fd..64ff0948 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -15,7 +15,6 @@ import ( "syscall" "time" - "github.com/google/uuid" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/cmdtail" "github.com/scripthaus-dev/mshell/pkg/packet" @@ -38,7 +37,7 @@ func setupSingleSignals(cmd *shexec.ShExecType) { }() } -func doSingle(cmdId string) { +func doSingle(ck base.CommandKey) { packetCh := packet.PacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) var runPacket *packet.RunPacketType @@ -57,11 +56,11 @@ func doSingle(cmdId string) { sender.SendErrorPacket("did not receive a 'run' packet") return } - if runPacket.CmdId == "" { - runPacket.CmdId = cmdId + if runPacket.CK.IsEmpty() { + runPacket.CK = ck } - if runPacket.CmdId != cmdId { - sender.SendErrorPacket(fmt.Sprintf("run packet cmdid[%s] did not match arg[%s]", runPacket.CmdId, cmdId)) + if runPacket.CK != ck { + sender.SendErrorPacket(fmt.Sprintf("run packet cmdid[%s] did not match arg[%s]", runPacket.CK, ck)) return } cmd, err := shexec.RunCommand(runPacket, sender) @@ -79,39 +78,36 @@ func doSingle(cmdId string) { } func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { - if pk.CmdId == "" { - pk.CmdId = uuid.New().String() - } err := shexec.ValidateRunPacket(pk) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("invalid run packet: %v", err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("invalid run packet: %v", err))) return } - fileNames, err := base.GetCommandFileNames(pk.SessionId, pk.CmdId) + fileNames, err := base.GetCommandFileNames(pk.CK) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot get command file names: %v", err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot get command file names: %v", err))) return } - cmd, err := shexec.MakeRunnerExec(pk.CmdId) + cmd, err := shexec.MakeRunnerExec(pk.CK) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot make mshell command: %v", err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot make mshell command: %v", err))) return } cmdStdin, err := cmd.StdinPipe() if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot pipe stdin to command: %v", err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot pipe stdin to command: %v", err))) return } // touch ptyout file (should exist for tailer to work correctly) ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err))) return } ptyOutFd.Close() // just opened to create the file, can close right after runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err))) return } defer runnerOutFd.Close() @@ -119,13 +115,13 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { cmd.Stderr = runnerOutFd err = cmd.Start() if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("error starting command: %v", err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("error starting command: %v", err))) return } go func() { err = packet.SendPacket(cmdStdin, pk) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("error sending forked runner command: %v", err))) + sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("error sending forked runner command: %v", err))) return } cmdStdin.Close() @@ -451,12 +447,12 @@ func main() { } if len(os.Args) >= 2 { - cmdId, err := uuid.Parse(os.Args[1]) - if err != nil { - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid non-cmdid passed to mshell", err)) + ck := base.CommandKey(os.Args[1]) + if err := ck.Validate("mshell arg"); err != nil { + packet.SendErrorPacket(os.Stdout, err.Error()) return } - doSingle(cmdId.String()) + doSingle(ck) time.Sleep(100 * time.Millisecond) return } else { diff --git a/pkg/base/base.go b/pkg/base/base.go index eae9a8cc..39398e01 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -15,6 +15,8 @@ import ( "path" "path/filepath" "strings" + + "github.com/google/uuid" ) const DefaultMShellPath = "mshell" @@ -37,6 +39,68 @@ type CommandFileNames struct { RunnerOutFile string } +type CommandKey string + +func MakeCommandKey(sessionId string, cmdId string) CommandKey { + if sessionId == "" && cmdId == "" { + return CommandKey("") + } + return CommandKey(fmt.Sprintf("%s/%s", sessionId, cmdId)) +} + +func (ckey CommandKey) IsEmpty() bool { + return string(ckey) == "" +} + +func (ckey CommandKey) GetSessionId() string { + slashIdx := strings.Index(string(ckey), "/") + if slashIdx == -1 { + return "" + } + return string(ckey[0:slashIdx]) +} + +func (ckey CommandKey) GetCmdId() string { + slashIdx := strings.Index(string(ckey), "/") + if slashIdx == -1 { + return "" + } + return string(ckey[slashIdx+1:]) +} + +func (ckey CommandKey) Split() (string, string) { + fields := strings.SplitN(string(ckey), "/", 2) + if len(fields) < 2 { + return "", "" + } + return fields[0], fields[1] +} + +func (ckey CommandKey) Validate(typeStr string) error { + if typeStr == "" { + typeStr = "ck" + } + if ckey == "" { + return fmt.Errorf("%s has empty commandkey", typeStr) + } + sessionId, cmdId := ckey.Split() + if sessionId == "" { + return fmt.Errorf("%s does not have sessionid", typeStr) + } + _, err := uuid.Parse(sessionId) + if err != nil { + return fmt.Errorf("%s has invalid sessionid '%s'", typeStr, sessionId) + } + if cmdId == "" { + return fmt.Errorf("%s does not have cmdid", typeStr) + } + _, err = uuid.Parse(cmdId) + if err != nil { + return fmt.Errorf("%s has invalid cmdid '%s'", typeStr, cmdId) + } + return nil +} + func GetHomeDir() string { homeVar := os.Getenv(HomeVarName) if homeVar == "" { @@ -57,10 +121,11 @@ func GetScHomeDir() (string, error) { return scHome, nil } -func GetCommandFileNames(sessionId string, cmdId string) (*CommandFileNames, error) { - if sessionId == "" || cmdId == "" { - return nil, fmt.Errorf("cannot get command-files when sessionid or cmdid is empty") +func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { + if err := ck.Validate("ck"); err != nil { + return nil, fmt.Errorf("cannot get command files: %w", err) } + sessionId, cmdId := ck.Split() sdir, err := EnsureSessionDir(sessionId) if err != nil { return nil, err @@ -73,8 +138,8 @@ func GetCommandFileNames(sessionId string, cmdId string) (*CommandFileNames, err }, nil } -func MakeCommandFileNamesWithHome(scHome string, sessionId string, cmdId string) *CommandFileNames { - base := path.Join(scHome, SessionsDirBaseName, sessionId, cmdId) +func MakeCommandFileNamesWithHome(scHome string, ck CommandKey) *CommandFileNames { + base := path.Join(scHome, SessionsDirBaseName, ck.GetSessionId(), ck.GetCmdId()) return &CommandFileNames{ PtyOutFile: base + ".ptyout", StdinFifo: base + ".stdin", diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 51cdbfd8..517b35f8 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -15,7 +15,6 @@ import ( "time" "github.com/fsnotify/fsnotify" - "github.com/google/uuid" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" ) @@ -33,7 +32,7 @@ type TailPos struct { } type CmdWatchEntry struct { - CmdKey CmdKey + CmdKey base.CommandKey FilePtyLen int64 FileRunLen int64 Tails []TailPos @@ -73,20 +72,15 @@ func (pos TailPos) IsCurrent(entry CmdWatchEntry) bool { return pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen } -type CmdKey struct { - SessionId string - CmdId string -} - type Tailer struct { Lock *sync.Mutex - WatchList map[CmdKey]CmdWatchEntry + WatchList map[base.CommandKey]CmdWatchEntry ScHomeDir string Watcher *fsnotify.Watcher SendCh chan packet.PacketType } -func (t *Tailer) updateTailPos_nolock(cmdKey CmdKey, reqId string, pos TailPos) { +func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos TailPos) { entry, found := t.WatchList[cmdKey] if !found { return @@ -95,7 +89,7 @@ func (t *Tailer) updateTailPos_nolock(cmdKey CmdKey, reqId string, pos TailPos) t.WatchList[cmdKey] = entry } -func (t *Tailer) removeTailPos_nolock(cmdKey CmdKey, reqId string) { +func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { entry, found := t.WatchList[cmdKey] if !found { return @@ -107,13 +101,13 @@ func (t *Tailer) removeTailPos_nolock(cmdKey CmdKey, reqId string) { } // delete from watchlist, remove watches - fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, cmdKey.SessionId, cmdKey.CmdId) + fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, cmdKey) delete(t.WatchList, cmdKey) t.Watcher.Remove(fileNames.PtyOutFile) t.Watcher.Remove(fileNames.RunnerOutFile) } -func (t *Tailer) updateEntrySizes_nolock(cmdKey CmdKey, ptyLen int64, runLen int64) { +func (t *Tailer) updateEntrySizes_nolock(cmdKey base.CommandKey, ptyLen int64, runLen int64) { entry, found := t.WatchList[cmdKey] if !found { return @@ -123,7 +117,7 @@ func (t *Tailer) updateEntrySizes_nolock(cmdKey CmdKey, ptyLen int64, runLen int t.WatchList[cmdKey] = entry } -func (t *Tailer) getEntryAndPos_nolock(cmdKey CmdKey, reqId string) (CmdWatchEntry, TailPos, bool) { +func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (CmdWatchEntry, TailPos, bool) { entry, found := t.WatchList[cmdKey] if !found { return CmdWatchEntry{}, TailPos{}, false @@ -142,7 +136,7 @@ func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { } rtn := &Tailer{ Lock: &sync.Mutex{}, - WatchList: make(map[CmdKey]CmdWatchEntry), + WatchList: make(map[base.CommandKey]CmdWatchEntry), ScHomeDir: scHomeDir, SendCh: sendCh, } @@ -170,8 +164,7 @@ func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]b func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWatchEntry, pos TailPos) *packet.CmdDataPacketType { dataPacket := packet.MakeCmdDataPacket() dataPacket.ReqId = pos.ReqId - dataPacket.SessionId = entry.CmdKey.SessionId - dataPacket.CmdId = entry.CmdKey.CmdId + dataPacket.CK = entry.CmdKey dataPacket.PtyPos = pos.TailPtyPos dataPacket.RunPos = pos.TailRunPos if entry.FilePtyLen > pos.TailPtyPos { @@ -196,14 +189,14 @@ func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWa } // returns (data-packet, keepRunning) -func (t *Tailer) runSingleDataTransfer(key CmdKey, reqId string) (*packet.CmdDataPacketType, bool) { +func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*packet.CmdDataPacketType, bool) { t.Lock.Lock() entry, pos, foundPos := t.getEntryAndPos_nolock(key, reqId) t.Lock.Unlock() if !foundPos { return nil, false } - fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, key.SessionId, key.CmdId) + fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, key) dataPacket := t.makeCmdDataPacket(fileNames, entry, pos) t.Lock.Lock() @@ -232,7 +225,7 @@ func (t *Tailer) runSingleDataTransfer(key CmdKey, reqId string) (*packet.CmdDat return dataPacket, pos.Running } -func (t *Tailer) checkRemoveNoFollow(cmdKey CmdKey, reqId string) { +func (t *Tailer) checkRemoveNoFollow(cmdKey base.CommandKey, reqId string) { t.Lock.Lock() defer t.Lock.Unlock() _, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) @@ -244,7 +237,7 @@ func (t *Tailer) checkRemoveNoFollow(cmdKey CmdKey, reqId string) { } } -func (t *Tailer) RunDataTransfer(key CmdKey, reqId string) { +func (t *Tailer) RunDataTransfer(key base.CommandKey, reqId string) { for { dataPacket, keepRunning := t.runSingleDataTransfer(key, reqId) if dataPacket != nil { @@ -283,7 +276,7 @@ func (t *Tailer) updateFile(relFileName string) { t.SendCh <- packet.FmtMessagePacket("error trying to stat file '%s': %v", relFileName, err) return } - cmdKey := CmdKey{SessionId: m[1], CmdId: m[2]} + cmdKey := base.MakeCommandKey(m[1], m[2]) t.Lock.Lock() defer t.Lock.Unlock() entry, foundEntry := t.WatchList[cmdKey] @@ -336,7 +329,7 @@ func max(v1 int64, v2 int64) int64 { } func (entry *CmdWatchEntry) fillFilePos(scHomeDir string) { - fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, entry.CmdKey.SessionId, entry.CmdKey.CmdId) + fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, entry.CmdKey) ptyInfo, _ := os.Stat(fileNames.PtyOutFile) if ptyInfo != nil { entry.FilePtyLen = ptyInfo.Size() @@ -350,30 +343,24 @@ func (entry *CmdWatchEntry) fillFilePos(scHomeDir string) { func (t *Tailer) RemoveWatch(pk *packet.UntailCmdPacketType) { t.Lock.Lock() defer t.Lock.Unlock() - key := CmdKey{pk.SessionId, pk.CmdId} - t.removeTailPos_nolock(key, pk.ReqId) + t.removeTailPos_nolock(pk.CK, pk.ReqId) } func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { - _, err := uuid.Parse(getPacket.SessionId) - if err != nil { - return fmt.Errorf("getcmd, bad sessionid '%s': %w", getPacket.SessionId, err) - } - _, err = uuid.Parse(getPacket.CmdId) - if err != nil { - return fmt.Errorf("getcmd, bad cmdid '%s': %w", getPacket.CmdId, err) + if err := getPacket.CK.Validate("getcmd"); err != nil { + return err } if getPacket.ReqId == "" { return fmt.Errorf("getcmd, no reqid specified") } - fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, getPacket.SessionId, getPacket.CmdId) + fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, getPacket.CK) t.Lock.Lock() defer t.Lock.Unlock() - key := CmdKey{getPacket.SessionId, getPacket.CmdId} + key := getPacket.CK entry, foundEntry := t.WatchList[key] if !foundEntry { // add watches, initialize entry - err = t.Watcher.Add(fileNames.PtyOutFile) + err := t.Watcher.Add(fileNames.PtyOutFile) if err != nil { return err } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index a1b93bd4..ef56355f 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -12,6 +12,7 @@ import ( "os" "sync" + "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" ) @@ -21,8 +22,7 @@ const MaxSingleWriteSize = 4 * 1024 type Multiplexer struct { Lock *sync.Mutex - SessionId string - CmdId string + CK base.CommandKey FdReaders map[int]*FdReader // synchronized FdWriters map[int]*FdWriter // synchronized CloseAfterStart []*os.File // synchronized @@ -34,11 +34,10 @@ type Multiplexer struct { Debug bool } -func MakeMultiplexer(sessionId string, cmdId string) *Multiplexer { +func MakeMultiplexer(ck base.CommandKey) *Multiplexer { return &Multiplexer{ Lock: &sync.Mutex{}, - SessionId: sessionId, - CmdId: cmdId, + CK: ck, FdReaders: make(map[int]*FdReader), FdWriters: make(map[int]*FdWriter), } @@ -126,8 +125,7 @@ func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd *os.File, shouldClose bool) func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { ack := packet.MakeDataAckPacket() - ack.SessionId = m.SessionId - ack.CmdId = m.CmdId + ack.CK = m.CK ack.FdNum = fdNum ack.AckLen = ackLen if err != nil { @@ -138,8 +136,7 @@ func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packe func (m *Multiplexer) makeDataPacket(fdNum int, data []byte, err error) *packet.DataPacketType { pk := packet.MakeDataPacket() - pk.SessionId = m.SessionId - pk.CmdId = m.CmdId + pk.CK = m.CK pk.FdNum = fdNum pk.Data64 = base64.StdEncoding.EncodeToString(data) if err != nil { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 3c6d9dce..d996ae9a 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -16,6 +16,8 @@ import ( "strconv" "strings" "sync" + + "github.com/scripthaus-dev/mshell/pkg/base" ) // remote: init, run, ping, data, cmdstart, cmddone @@ -80,26 +82,29 @@ func MakePacket(packetType string) (PacketType, error) { } type CmdDataPacketType struct { - Type string `json:"type"` - ReqId string `json:"reqid"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` - PtyPos int64 `json:"ptypos"` - PtyLen int64 `json:"ptylen"` - RunPos int64 `json:"runpos"` - RunLen int64 `json:"runlen"` - PtyData string `json:"ptydata"` - PtyDataLen int `json:"ptydatalen"` - RunData string `json:"rundata"` - RunDataLen int `json:"rundatalen"` - Error string `json:"error"` - NotFound bool `json:"notfound,omitempty"` + Type string `json:"type"` + ReqId string `json:"reqid"` + CK base.CommandKey `json:"ck"` + PtyPos int64 `json:"ptypos"` + PtyLen int64 `json:"ptylen"` + RunPos int64 `json:"runpos"` + RunLen int64 `json:"runlen"` + PtyData string `json:"ptydata"` + PtyDataLen int `json:"ptydatalen"` + RunData string `json:"rundata"` + RunDataLen int `json:"rundatalen"` + Error string `json:"error"` + NotFound bool `json:"notfound,omitempty"` } func (*CmdDataPacketType) GetType() string { return CmdDataPacketStr } +func (p *CmdDataPacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeCmdDataPacket() *CmdDataPacketType { return &CmdDataPacketType{Type: CmdDataPacketStr} } @@ -117,19 +122,22 @@ func MakePingPacket() *PingPacketType { } type DataPacketType struct { - Type string `json:"type"` - SessionId string `json:"sessionid,omitempty"` - CmdId string `json:"cmdid,omitempty"` - FdNum int `json:"fdnum"` - Data64 string `json:"data64"` // base64 encoded - Eof bool `json:"eof,omitempty"` - Error string `json:"error,omitempty"` + Type string `json:"type"` + CK base.CommandKey `json:"ck"` + FdNum int `json:"fdnum"` + Data64 string `json:"data64"` // base64 encoded + Eof bool `json:"eof,omitempty"` + Error string `json:"error,omitempty"` } func (*DataPacketType) GetType() string { return DataPacketStr } +func (p *DataPacketType) GetCK() base.CommandKey { + return p.CK +} + func B64DecodedLen(b64 string) int { if len(b64) < 4 { return 0 // we use padded strings, so < 4 is always 0 @@ -161,18 +169,21 @@ func MakeDataPacket() *DataPacketType { } type DataAckPacketType struct { - Type string `json:"type"` - SessionId string `json:"sessionid,omitempty"` - CmdId string `json:"cmdid,omitempty"` - FdNum int `json:"fdnum"` - AckLen int `json:"acklen"` - Error string `json:"error,omitempty"` + Type string `json:"type"` + CK base.CommandKey `json:"ck"` + FdNum int `json:"fdnum"` + AckLen int `json:"acklen"` + Error string `json:"error,omitempty"` } func (*DataAckPacketType) GetType() string { return DataAckPacketStr } +func (p *DataAckPacketType) GetCK() base.CommandKey { + return p.CK +} + func (p *DataAckPacketType) String() string { errStr := "" if p.Error != "" { @@ -189,52 +200,61 @@ func MakeDataAckPacket() *DataAckPacketType { // SigNum gets sent to process via a signal // WinSize, if set, will run TIOCSWINSZ to set size, and then send SIGWINCH type InputPacketType struct { - Type string `json:"type"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` - InputData string `json:"inputdata"` - SigNum int `json:"signum,omitempty"` - WinSizeRows int `json:"winsizerows"` - WinSizeCols int `json:"winsizecols"` + Type string `json:"type"` + CK base.CommandKey `json:"ck"` + InputData string `json:"inputdata"` + SigNum int `json:"signum,omitempty"` + WinSizeRows int `json:"winsizerows"` + WinSizeCols int `json:"winsizecols"` } func (*InputPacketType) GetType() string { return InputPacketStr } +func (p *InputPacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeInputPacket() *InputPacketType { return &InputPacketType{Type: InputPacketStr} } type UntailCmdPacketType struct { - Type string `json:"type"` - ReqId string `json:"reqid"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` + Type string `json:"type"` + ReqId string `json:"reqid"` + CK base.CommandKey `json:"ck"` } func (*UntailCmdPacketType) GetType() string { return UntailCmdPacketStr } +func (p *UntailCmdPacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeUntailCmdPacket() *UntailCmdPacketType { return &UntailCmdPacketType{Type: UntailCmdPacketStr} } type GetCmdPacketType struct { - Type string `json:"type"` - ReqId string `json:"reqid"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` - PtyPos int64 `json:"ptypos"` - RunPos int64 `json:"runpos"` - Tail bool `json:"tail,omitempty"` + Type string `json:"type"` + ReqId string `json:"reqid"` + CK base.CommandKey `json:"ck"` + PtyPos int64 `json:"ptypos"` + RunPos int64 `json:"runpos"` + Tail bool `json:"tail,omitempty"` } func (*GetCmdPacketType) GetType() string { return GetCmdPacketStr } +func (p *GetCmdPacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeGetCmdPacket() *GetCmdPacketType { return &GetCmdPacketType{Type: GetCmdPacketStr} } @@ -346,35 +366,41 @@ func MakeDonePacket() *DonePacketType { } type CmdDonePacketType struct { - Type string `json:"type"` - Ts int64 `json:"ts"` - SessionId string `json:"sessionid,omitempty"` - CmdId string `json:"cmdid,omitempty"` - ExitCode int `json:"exitcode"` - DurationMs int64 `json:"durationms"` + Type string `json:"type"` + Ts int64 `json:"ts"` + CK base.CommandKey `json:"ck"` + ExitCode int `json:"exitcode"` + DurationMs int64 `json:"durationms"` } func (*CmdDonePacketType) GetType() string { return CmdDonePacketStr } +func (p *CmdDonePacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeCmdDonePacket() *CmdDonePacketType { return &CmdDonePacketType{Type: CmdDonePacketStr} } type CmdStartPacketType struct { - Type string `json:"type"` - Ts int64 `json:"ts"` - SessionId string `json:"sessionid,omitempty"` - CmdId string `json:"cmdid,omitempty"` - Pid int `json:"pid"` - MShellPid int `json:"mshellpid"` + Type string `json:"type"` + Ts int64 `json:"ts"` + CK base.CommandKey `json:"ck"` + Pid int `json:"pid"` + MShellPid int `json:"mshellpid"` } func (*CmdStartPacketType) GetType() string { return CmdStartPacketStr } +func (p *CmdStartPacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeCmdStartPacket() *CmdStartPacketType { return &CmdStartPacketType{Type: CmdStartPacketStr} } @@ -393,21 +419,24 @@ type RemoteFd struct { } type RunPacketType struct { - Type string `json:"type"` - SessionId string `json:"sessionid,omitempty"` - CmdId string `json:"cmdid,omitempty"` - Command string `json:"command"` - Cwd string `json:"cwd,omitempty"` - Env map[string]string `json:"env,omitempty"` - TermSize *TermSize `json:"termsize,omitempty"` - Fds []RemoteFd `json:"fds,omitempty"` - Detached bool `json:"detached,omitempty"` + Type string `json:"type"` + CK base.CommandKey `json:"ck"` + Command string `json:"command"` + Cwd string `json:"cwd,omitempty"` + Env map[string]string `json:"env,omitempty"` + TermSize *TermSize `json:"termsize,omitempty"` + Fds []RemoteFd `json:"fds,omitempty"` + Detached bool `json:"detached,omitempty"` } func (*RunPacketType) GetType() string { return RunPacketStr } +func (p *RunPacketType) GetCK() base.CommandKey { + return p.CK +} + func MakeRunPacket() *RunPacketType { return &RunPacketType{Type: RunPacketStr} } @@ -417,9 +446,9 @@ type BarePacketType struct { } type ErrorPacketType struct { - Id string `json:"id,omitempty"` - Type string `json:"type"` - Error string `json:"error"` + CK base.CommandKey `json:"ck,omitempty"` + Type string `json:"type"` + Error string `json:"error"` } func (et *ErrorPacketType) GetType() string { @@ -430,8 +459,8 @@ func MakeErrorPacket(errorStr string) *ErrorPacketType { return &ErrorPacketType{Type: ErrorPacketStr, Error: errorStr} } -func MakeIdErrorPacket(id string, errorStr string) *ErrorPacketType { - return &ErrorPacketType{Type: ErrorPacketStr, Id: id, Error: errorStr} +func MakeCKErrorPacket(ck base.CommandKey, errorStr string) *ErrorPacketType { + return &ErrorPacketType{Type: ErrorPacketStr, CK: ck, Error: errorStr} } type PacketType interface { @@ -450,6 +479,11 @@ type RpcPacketType interface { GetPacketId() string } +type CommandPacketType interface { + GetType() string + GetCK() base.CommandKey +} + func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { var bareCmd BarePacketType err := json.Unmarshal(jsonBuf, &bareCmd) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 73afa925..5ea68d2b 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -17,7 +17,6 @@ import ( "time" "github.com/creack/pty" - "github.com/google/uuid" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" @@ -39,21 +38,19 @@ const RemoteSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -S -C %d bash -c "ec type ShExecType struct { Lock *sync.Mutex StartTs time.Time - SessionId string - CmdId string + CK base.CommandKey FileNames *base.CommandFileNames Cmd *exec.Cmd CmdPty *os.File Multiplexer *mpio.Multiplexer } -func MakeShExec(sessionId string, cmdId string) *ShExecType { +func MakeShExec(ck base.CommandKey) *ShExecType { return &ShExecType{ Lock: &sync.Mutex{}, StartTs: time.Now(), - SessionId: sessionId, - CmdId: cmdId, - Multiplexer: mpio.MakeMultiplexer(sessionId, cmdId), + CK: ck, + Multiplexer: mpio.MakeMultiplexer(ck), } } @@ -67,8 +64,7 @@ func (c *ShExecType) Close() { func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { startPacket := packet.MakeCmdStartPacket() startPacket.Ts = time.Now().UnixMilli() - startPacket.SessionId = c.SessionId - startPacket.CmdId = c.CmdId + startPacket.CK = c.CK startPacket.Pid = c.Cmd.Process.Pid startPacket.MShellPid = os.Getpid() return startPacket @@ -129,12 +125,12 @@ func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { return ecmd } -func MakeRunnerExec(cmdId string) (*exec.Cmd, error) { +func MakeRunnerExec(ck base.CommandKey) (*exec.Cmd, error) { msPath, err := base.GetMShellPath() if err != nil { return nil, err } - ecmd := exec.Command(msPath, cmdId) + ecmd := exec.Command(msPath, string(ck)) return ecmd, nil } @@ -165,19 +161,9 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { return fmt.Errorf("run packet has wrong type: %s", pk.Type) } if pk.Detached { - if pk.SessionId == "" { - return fmt.Errorf("run packet does not have sessionid") - } - _, err := uuid.Parse(pk.SessionId) + err := pk.CK.Validate("run packet") if err != nil { - return fmt.Errorf("invalid sessionid '%s' for command", pk.SessionId) - } - if pk.CmdId == "" { - return fmt.Errorf("run packet does not have cmdid") - } - _, err = uuid.Parse(pk.CmdId) - if err != nil { - return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) + return err } } if pk.Cwd != "" { @@ -337,7 +323,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if err != nil { return nil, err } - cmd := MakeShExec("", "") + cmd := MakeShExec("") var fullSshOpts []string fullSshOpts = append(fullSshOpts, opts.SSHOpts...) fullSshOpts = append(fullSshOpts, SSHRemoteCommand) @@ -430,7 +416,7 @@ func (cmd *ShExecType) RunRemoteIOAndWait(packetCh chan packet.PacketType, sende } func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - cmd := MakeShExec(pk.SessionId, pk.CmdId) + cmd := MakeShExec(pk.CK) cmd.Cmd = exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(cmd.Cmd, pk.Env) if pk.Cwd != "" { @@ -491,7 +477,7 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S } func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - fileNames, err := base.GetCommandFileNames(pk.SessionId, pk.CmdId) + fileNames, err := base.GetCommandFileNames(pk.CK) if err != nil { return nil, err } @@ -499,7 +485,7 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( if err == nil { // non-nil error will be caught by regular OpenFile below // must have size 0 if ptyOutInfo.Size() != 0 { - return nil, fmt.Errorf("cmdid '%s' was already used (ptyout len=%d)", pk.CmdId, ptyOutInfo.Size()) + return nil, fmt.Errorf("cmdkey '%s' was already used (ptyout len=%d)", pk.CK, ptyOutInfo.Size()) } } cmdPty, cmdTty, err := pty.Open() @@ -510,7 +496,7 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( defer func() { cmdTty.Close() }() - rtn := MakeShExec(pk.SessionId, pk.CmdId) + rtn := MakeShExec(pk.CK) ecmd := MakeExecCmd(pk, cmdTty) err = ecmd.Start() if err != nil { @@ -558,8 +544,7 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { exitCode := GetExitCode(exitErr) donePacket := packet.MakeCmdDonePacket() donePacket.Ts = endTs.UnixMilli() - donePacket.SessionId = c.SessionId - donePacket.CmdId = c.CmdId + donePacket.CK = c.CK donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) if c.FileNames != nil { From 657440269141d01d9a741e5e83a107f9ab17cc04 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 12:14:07 -0700 Subject: [PATCH 030/149] create packetparser type, refactor to use --- main-mshell.go | 14 +++---- pkg/mpio/mpio.go | 12 +++--- pkg/packet/packet.go | 73 --------------------------------- pkg/packet/parser.go | 97 ++++++++++++++++++++++++++++++++++++++++++++ pkg/shexec/shexec.go | 14 +++---- 5 files changed, 117 insertions(+), 93 deletions(-) create mode 100644 pkg/packet/parser.go diff --git a/main-mshell.go b/main-mshell.go index 64ff0948..ea8df615 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -38,10 +38,10 @@ func setupSingleSignals(cmd *shexec.ShExecType) { } func doSingle(ck base.CommandKey) { - packetCh := packet.PacketParser(os.Stdin) + packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) var runPacket *packet.RunPacketType - for pk := range packetCh { + for pk := range packetParser.MainCh { if pk.GetType() == packet.PingPacketStr { continue } @@ -156,7 +156,7 @@ func doMain() { packet.SendErrorPacket(os.Stdout, err.Error()) return } - packetCh := packet.PacketParser(os.Stdin) + packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) tailer, err := cmdtail.MakeTailer(sender.SendCh) if err != nil { @@ -172,7 +172,7 @@ func doMain() { initPacket.User = user.Username } sender.SendPacket(initPacket) - for pk := range packetCh { + for pk := range packetParser.MainCh { if pk.GetType() == packet.PingPacketStr { continue } @@ -212,7 +212,7 @@ func doMain() { } func handleRemote() { - packetCh := packet.PacketParser(os.Stdin) + packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) defer func() { // wait for sender to complete @@ -223,7 +223,7 @@ func handleRemote() { initPacket.Version = MShellVersion sender.SendPacket(initPacket) var runPacket *packet.RunPacketType - for pk := range packetCh { + for pk := range packetParser.MainCh { if pk.GetType() == packet.PingPacketStr { continue } @@ -251,7 +251,7 @@ func handleRemote() { defer cmd.Close() startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) - cmd.RunRemoteIOAndWait(packetCh, sender) + cmd.RunRemoteIOAndWait(packetParser, sender) } func handleServer() { diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index ef56355f..227625a1 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -28,7 +28,7 @@ type Multiplexer struct { CloseAfterStart []*os.File // synchronized Sender *packet.PacketSender - Input chan packet.PacketType + Input *packet.PacketParser Started bool Debug bool @@ -171,20 +171,20 @@ func (m *Multiplexer) launchReaders(wg *sync.WaitGroup) { } } -func (m *Multiplexer) startIO(packetCh chan packet.PacketType, sender *packet.PacketSender) { +func (m *Multiplexer) startIO(packetParser *packet.PacketParser, sender *packet.PacketSender) { m.Lock.Lock() defer m.Lock.Unlock() if m.Started { panic("Multiplexer is already running, cannot start again") } - m.Input = packetCh + m.Input = packetParser m.Sender = sender m.Started = true } func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { defer m.HandleInputDone() - for pk := range m.Input { + for pk := range m.Input.MainCh { if m.Debug { fmt.Printf("PK> %s\n", packet.AsString(pk)) } @@ -263,8 +263,8 @@ func (m *Multiplexer) closeTempStartFds() { m.CloseAfterStart = nil } -func (m *Multiplexer) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender, waitOnReaders bool, waitOnWriters bool, waitForInputLoop bool) *packet.CmdDonePacketType { - m.startIO(packetCh, sender) +func (m *Multiplexer) RunIOAndWait(packetParser *packet.PacketParser, sender *packet.PacketSender, waitOnReaders bool, waitOnWriters bool, waitForInputLoop bool) *packet.CmdDonePacketType { + m.startIO(packetParser, sender) m.closeTempStartFds() var wg sync.WaitGroup if waitOnReaders { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index d996ae9a..87848fce 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -7,14 +7,11 @@ package packet import ( - "bufio" "bytes" "encoding/json" "fmt" "io" "reflect" - "strconv" - "strings" "sync" "github.com/scripthaus-dev/mshell/pkg/base" @@ -602,76 +599,6 @@ func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) erro return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) } -func CombinePacketParsers(p1 chan PacketType, p2 chan PacketType) chan PacketType { - rtnCh := make(chan PacketType) - var wg sync.WaitGroup - wg.Add(2) - go func() { - defer wg.Done() - for v := range p1 { - rtnCh <- v - } - }() - go func() { - defer wg.Done() - for v := range p2 { - rtnCh <- v - } - }() - go func() { - wg.Wait() - close(rtnCh) - }() - return rtnCh -} - -func PacketParser(input io.Reader) chan PacketType { - rtnCh := make(chan PacketType) - bufReader := bufio.NewReader(input) - go func() { - defer func() { - close(rtnCh) - }() - for { - line, err := bufReader.ReadString('\n') - if err == io.EOF { - return - } - if err != nil { - errPacket := MakeErrorPacket(fmt.Sprintf("reading packets from input: %v", err)) - rtnCh <- errPacket - return - } - if line == "\n" { - continue - } - // ##[len][json]\n - // ##14{"hello":true}\n - bracePos := strings.Index(line, "{") - if !strings.HasPrefix(line, "##") || bracePos == -1 { - rtnCh <- MakeRawPacket(line[:len(line)-1]) - continue - } - packetLen, err := strconv.Atoi(line[2:bracePos]) - if err != nil || packetLen != len(line)-bracePos-1 { - rtnCh <- MakeRawPacket(line[:len(line)-1]) - continue - } - pk, err := ParseJsonPacket([]byte(line[bracePos:])) - if err != nil { - errPk := MakeErrorPacket(fmt.Sprintf("parsing packet json from input: %v", err)) - rtnCh <- errPk - return - } - if pk.GetType() == DonePacketStr { - return - } - rtnCh <- pk - } - }() - return rtnCh -} - type ErrorReporter interface { ReportError(err error) } diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go new file mode 100644 index 00000000..907ae83f --- /dev/null +++ b/pkg/packet/parser.go @@ -0,0 +1,97 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package packet + +import ( + "bufio" + "fmt" + "io" + "strconv" + "strings" + "sync" +) + +type PacketParser struct { + Lock *sync.Mutex + MainCh chan PacketType +} + +func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser { + rtnParser := &PacketParser{ + Lock: &sync.Mutex{}, + MainCh: make(chan PacketType), + } + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for v := range p1.MainCh { + rtnParser.MainCh <- v + } + }() + go func() { + defer wg.Done() + for v := range p2.MainCh { + rtnParser.MainCh <- v + } + }() + go func() { + wg.Wait() + close(rtnParser.MainCh) + }() + return rtnParser +} + +func MakePacketParser(input io.Reader) *PacketParser { + parser := &PacketParser{ + Lock: &sync.Mutex{}, + MainCh: make(chan PacketType), + } + bufReader := bufio.NewReader(input) + go func() { + defer func() { + close(parser.MainCh) + }() + for { + line, err := bufReader.ReadString('\n') + if err == io.EOF { + return + } + if err != nil { + errPacket := MakeErrorPacket(fmt.Sprintf("reading packets from input: %v", err)) + parser.MainCh <- errPacket + return + } + if line == "\n" { + continue + } + // ##[len][json]\n + // ##14{"hello":true}\n + bracePos := strings.Index(line, "{") + if !strings.HasPrefix(line, "##") || bracePos == -1 { + parser.MainCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + packetLen, err := strconv.Atoi(line[2:bracePos]) + if err != nil || packetLen != len(line)-bracePos-1 { + parser.MainCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + pk, err := ParseJsonPacket([]byte(line[bracePos:])) + if err != nil { + errPk := MakeErrorPacket(fmt.Sprintf("parsing packet json from input: %v", err)) + parser.MainCh <- errPk + return + } + if pk.GetType() == DonePacketStr { + return + } + parser.MainCh <- pk + } + }() + return parser +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 5ea68d2b..fad15ca7 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -373,12 +373,12 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, fmt.Errorf("running ssh command: %w", err) } defer cmd.Close() - stdoutPacketCh := packet.PacketParser(stdoutReader) - stderrPacketCh := packet.PacketParser(stderrReader) - packetCh := packet.CombinePacketParsers(stdoutPacketCh, stderrPacketCh) + stdoutPacketParser := packet.MakePacketParser(stdoutReader) + stderrPacketParser := packet.MakePacketParser(stderrReader) + packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) sender := packet.MakePacketSender(inputWriter) versionOk := false - for pk := range packetCh { + for pk := range packetParser.MainCh { if pk.GetType() == packet.RawPacketStr { rawPk := pk.(*packet.RawPacketType) fmt.Printf("%s\n", rawPk.Data) @@ -400,7 +400,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if opts.Debug { cmd.Multiplexer.Debug = true } - remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetCh, sender, false, true, true) + remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetParser, sender, false, true, true) donePacket := cmd.WaitForCommand() if remoteDonePacket != nil { donePacket = remoteDonePacket @@ -408,9 +408,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return donePacket, nil } -func (cmd *ShExecType) RunRemoteIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { +func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sender *packet.PacketSender) { defer cmd.Close() - cmd.Multiplexer.RunIOAndWait(packetCh, sender, true, false, false) + cmd.Multiplexer.RunIOAndWait(packetParser, sender, true, false, false) donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) } From dafe2b5a575c3c3c6049937b970d8d26adb1daa2 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 14:57:01 -0700 Subject: [PATCH 031/149] simplify argument parsing, hard code common ssh options --- go.mod | 1 + go.sum | 2 ++ main-mshell.go | 86 ++++++++++++++++++++++++-------------------- pkg/base/optsiter.go | 7 ++++ pkg/packet/packet.go | 1 + pkg/shexec/shexec.go | 56 +++++++++++++++++++++-------- 6 files changed, 101 insertions(+), 52 deletions(-) diff --git a/go.mod b/go.mod index e1a2464c..735e32fd 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module github.com/scripthaus-dev/mshell go 1.17 require ( + github.com/alessio/shellescape v1.4.1 // indirect github.com/creack/pty v1.1.18 // indirect github.com/fsnotify/fsnotify v1.5.4 // indirect github.com/google/uuid v1.3.0 // indirect diff --git a/go.sum b/go.sum index fd7b5ca7..d802966a 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +github.com/alessio/shellescape v1.4.1 h1:V7yhSDDn8LP4lc4jS8pFkt0zCnzVJlG5JXy9BVKJUX0= +github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPpyFjDCtHLS30= github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= github.com/fsnotify/fsnotify v1.5.4 h1:jRbGcIw6P2Meqdwuo0H1p6JVLbL5DHKAKlYndzMwVZI= diff --git a/main-mshell.go b/main-mshell.go index ea8df615..2b3e4116 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -211,7 +211,7 @@ func doMain() { } } -func handleRemote() { +func handleSingle() { packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) defer func() { @@ -285,14 +285,35 @@ func parseClientOpts() (*shexec.ClientOpts, error) { for iter.HasNext() { argStr := iter.Next() if argStr == "--ssh" { - if opts.IsSSH { - return nil, fmt.Errorf("duplicate '--ssh' option") + if !iter.IsNextPlain() { + return nil, fmt.Errorf("'--ssh [user@host]' missing host") } - opts.IsSSH = true - break + opts.SSHHost = iter.Next() + continue + } + if argStr == "--ssh-opts" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--ssh-opts [options]' missing options") + } + opts.SSHOptsStr = iter.Next() + continue + } + if argStr == "-i" { + if !iter.IsNextPlain() { + return nil, fmt.Errorf("-i [identity-file]' missing file") + } + opts.SSHIdentity = iter.Next() + continue + } + if argStr == "-l" { + if !iter.IsNextPlain() { + return nil, fmt.Errorf("-l [user]' missing user") + } + opts.SSHUser = iter.Next() + continue } if argStr == "--cwd" { - if !iter.HasNext() { + if !iter.IsNextPlain() { return nil, fmt.Errorf("'--cwd [dir]' missing directory") } opts.Cwd = iter.Next() @@ -316,7 +337,7 @@ func parseClientOpts() (*shexec.ClientOpts, error) { continue } if argStr == "--sudo-with-passfile" { - if !iter.HasNext() { + if !iter.IsNextPlain() { return nil, fmt.Errorf("'--sudo-with-passfile [file]', missing file") } opts.Sudo = true @@ -329,29 +350,12 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.SudoPw = string(contents) continue } - } - if opts.IsSSH { - // parse SSH opts - for iter.HasNext() { - argStr := iter.Next() - if argStr == "--" { - opts.SSHOptsTerm = true - break + if argStr == "--" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--' should be followed by command") } - if argStr == "-t" || argStr == "-tt" { - return nil, fmt.Errorf("mshell cannot run over ssh -t") - } - opts.SSHOpts = append(opts.SSHOpts, argStr) - } - if !opts.SSHOptsTerm { - return nil, fmt.Errorf("ssh options must be terminated with '--' followed by [command]") - } - if !iter.HasNext() { - return nil, fmt.Errorf("no command specified") - } - opts.Command = strings.Join(iter.Rest(), " ") - if strings.TrimSpace(opts.Command) == "" { - return nil, fmt.Errorf("no command or empty command specified") + opts.Command = strings.Join(iter.Rest(), " ") + break } } return opts, nil @@ -365,9 +369,12 @@ func handleClient() (int, error) { if opts.Debug { packet.GlobalDebug = true } - if !opts.IsSSH { + if opts.SSHHost == "" { return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") } + if opts.Command == "" { + return 1, fmt.Errorf("no [command] specified. [command] follows '--' option (see usage)") + } fds, err := detectOpenFds() if err != nil { return 1, err @@ -382,17 +389,20 @@ func handleClient() (int, error) { func handleUsage() { usage := ` -Client Usage: mshell [mshell-opts] --ssh [ssh-opts] user@host -- [command] +Client Usage: mshell [opts] --ssh user@host -- [command] mshell multiplexes input and output streams to a remote command over ssh. Options: - --cwd [dir] - execute remote command in [dir] - [command] - a single argument (should be quoted) + -i [identity-file] - used to set '-i' option for ssh command + -l [user] - used to set '-l' option for ssh command + --cwd [dir] - execute remote command in [dir] + --ssh-opts [opts] - addition options to pass to ssh command + [command] - the remote command to execute Sudo Options: - --sudo - --sudo-with-password [pw] (not recommended, use --sudo-with-passfile if possible) + --sudo - use only if sudo never requires a password + --sudo-with-password [pw] - not recommended, use --sudo-with-passfile if possible --sudo-with-passfile [file] Sudo options allow you to run the given command using "sudo". The first @@ -401,7 +411,7 @@ securely through a high numbered fd to "sudo -S". See full documentation for mo Examples: # execute a python script remotely, with stdin still hooked up correctly - mshell --cwd "~/work" --ssh -i key.pem ubuntu@somehost -- "python3 /dev/fd/4" 4< myscript.py + mshell --cwd "~/work" -i key.pem --ssh ubuntu@somehost -- "python3 /dev/fd/4" 4< myscript.py # capture multiple outputs mshell --ssh ubuntu@test -- "cat file1.txt > /dev/fd/3; cat file2.txt > /dev/fd/4" 3> file1.txt 4> file2.txt @@ -431,8 +441,8 @@ func main() { } else if firstArg == "--version" { fmt.Printf("mshell v%s\n", MShellVersion) return - } else if firstArg == "--remote" { - handleRemote() + } else if firstArg == "--single" { + handleSingle() return } else if firstArg == "--server" { handleServer() diff --git a/pkg/base/optsiter.go b/pkg/base/optsiter.go index 329455ba..923d86d2 100644 --- a/pkg/base/optsiter.go +++ b/pkg/base/optsiter.go @@ -25,6 +25,13 @@ func (iter *OptsIter) HasNext() bool { return iter.Pos <= len(iter.Opts)-1 } +func (iter *OptsIter) IsNextPlain() bool { + if !iter.HasNext() { + return false + } + return !IsOption(iter.Opts[iter.Pos]) +} + func (iter *OptsIter) Next() string { if iter.Pos >= len(iter.Opts) { return "" diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 87848fce..3638a04f 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -340,6 +340,7 @@ type InitPacketType struct { HomeDir string `json:"homedir,omitempty"` Env []string `json:"env,omitempty"` User string `json:"user,omitempty"` + NotFound bool `json:"notfound,omitempty"` } func (*InitPacketType) GetType() string { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index fad15ca7..b81cae1e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -16,6 +16,7 @@ import ( "syscall" "time" + "github.com/alessio/shellescape" "github.com/creack/pty" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/mpio" @@ -29,11 +30,20 @@ const MaxCols = 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 -const SSHRemoteCommand = `PATH=$PATH:~/.mshell; mshell --remote` +const SSHRemoteCommand = ` +PATH=$PATH:~/.mshell; +which mshell > /dev/null; +if [[ "$?" -ne 0 ]] +then + printf "\n##34{\"type\": \"init\", \"notfound\": true}\n" +else + mshell --single +fi +` -const RemoteCommandFmt = `%s` -const RemoteSudoCommandFmt = `sudo -C %d bash /dev/fd/%d` -const RemoteSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -S -C %d bash -c "echo '[from-mshell]'; bash /dev/fd/%d < /dev/fd/%d"` +const RunCommandFmt = `%s` +const RunSudoCommandFmt = `sudo -C %d bash /dev/fd/%d` +const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -S -C %d bash -c "echo '[from-mshell]'; bash /dev/fd/%d < /dev/fd/%d"` type ShExecType struct { Lock *sync.Mutex @@ -205,9 +215,10 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } type ClientOpts struct { - IsSSH bool - SSHOptsTerm bool - SSHOpts []string + SSHHost string + SSHOptsStr string + SSHIdentity string + SSHUser string Command string Fds []packet.RemoteFd Cwd string @@ -218,13 +229,29 @@ type ClientOpts struct { CommandStdinFdNum int } +func (opts *ClientOpts) MakeSSHCommandString() string { + var moreSSHOpts []string + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + moreSSHOpts = append(moreSSHOpts, identityOpt) + } + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + moreSSHOpts = append(moreSSHOpts, userOpt) + } + remoteCommand := strings.TrimSpace(SSHRemoteCommand) + // note that SSHOptsStr is *not* escaped + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) + return sshCmd +} + func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket := packet.MakeRunPacket() runPacket.Cwd = opts.Cwd runPacket.Fds = opts.Fds if !opts.Sudo { // normal, non-sudo command - runPacket.Command = opts.Command + runPacket.Command = fmt.Sprintf(RunCommandFmt, opts.Command) return runPacket, nil } if opts.SudoWithPass { @@ -248,7 +275,7 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { opts.Fds = append(opts.Fds, commandStdinRfd) opts.CommandStdinFdNum = commandStdinFdNum maxFdNum := opts.MaxFdNum() - runPacket.Command = fmt.Sprintf(RemoteSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, commandFdNum, commandStdinFdNum) + runPacket.Command = fmt.Sprintf(RunSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, commandFdNum, commandStdinFdNum) runPacket.Fds = opts.Fds return runPacket, nil } else { @@ -259,7 +286,7 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { rfd := packet.RemoteFd{FdNum: commandFdNum, Read: true, Content: opts.Command} opts.Fds = append(opts.Fds, rfd) maxFdNum := opts.MaxFdNum() - runPacket.Command = fmt.Sprintf(RemoteSudoCommandFmt, maxFdNum+1, commandFdNum) + runPacket.Command = fmt.Sprintf(RunSudoCommandFmt, maxFdNum+1, commandFdNum) runPacket.Fds = opts.Fds return runPacket, nil } @@ -324,10 +351,8 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, err } cmd := MakeShExec("") - var fullSshOpts []string - fullSshOpts = append(fullSshOpts, opts.SSHOpts...) - fullSshOpts = append(fullSshOpts, SSHRemoteCommand) - ecmd := exec.Command("ssh", fullSshOpts...) + sshCmdStr := opts.MakeSSHCommandString() + ecmd := exec.Command("bash", "-c", sshCmdStr) cmd.Cmd = ecmd inputWriter, err := ecmd.StdinPipe() if err != nil { @@ -386,6 +411,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er } if pk.GetType() == packet.InitPacketStr { initPk := pk.(*packet.InitPacketType) + if initPk.NotFound { + return nil, fmt.Errorf("mshell command not found on remote server, can install with 'mshell --install'") + } if initPk.Version != "0.1.0" { return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) } From 0f5ee87a768c05cbb6ef426ca9148ae1a9af2920 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 15:10:17 -0700 Subject: [PATCH 032/149] allow mshell to execute local commands --- main-mshell.go | 3 --- pkg/shexec/shexec.go | 37 +++++++++++++++++++++---------------- 2 files changed, 21 insertions(+), 19 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 2b3e4116..9a942492 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -369,9 +369,6 @@ func handleClient() (int, error) { if opts.Debug { packet.GlobalDebug = true } - if opts.SSHHost == "" { - return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") - } if opts.Command == "" { return 1, fmt.Errorf("no [command] specified. [command] follows '--' option (see usage)") } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index b81cae1e..09fcf167 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -30,7 +30,7 @@ const MaxCols = 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 -const SSHRemoteCommand = ` +const ClientCommand = ` PATH=$PATH:~/.mshell; which mshell > /dev/null; if [[ "$?" -ne 0 ]] @@ -229,20 +229,26 @@ type ClientOpts struct { CommandStdinFdNum int } -func (opts *ClientOpts) MakeSSHCommandString() string { - var moreSSHOpts []string - if opts.SSHIdentity != "" { - identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) - moreSSHOpts = append(moreSSHOpts, identityOpt) +func (opts *ClientOpts) MakeExecCmd() *exec.Cmd { + if opts.SSHHost == "" { + ecmd := exec.Command("bash", "-c", strings.TrimSpace(ClientCommand)) + return ecmd + } else { + var moreSSHOpts []string + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + moreSSHOpts = append(moreSSHOpts, identityOpt) + } + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + moreSSHOpts = append(moreSSHOpts, userOpt) + } + remoteCommand := strings.TrimSpace(ClientCommand) + // note that SSHOptsStr is *not* escaped + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) + ecmd := exec.Command("bash", "-c", sshCmd) + return ecmd } - if opts.SSHUser != "" { - userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) - moreSSHOpts = append(moreSSHOpts, userOpt) - } - remoteCommand := strings.TrimSpace(SSHRemoteCommand) - // note that SSHOptsStr is *not* escaped - sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) - return sshCmd } func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { @@ -351,8 +357,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, err } cmd := MakeShExec("") - sshCmdStr := opts.MakeSSHCommandString() - ecmd := exec.Command("bash", "-c", sshCmdStr) + ecmd := opts.MakeExecCmd() cmd.Cmd = ecmd inputWriter, err := ecmd.StdinPipe() if err != nil { From ec4bd5eaa1969f8f7834604508869cf0cce061d1 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 15:59:14 -0700 Subject: [PATCH 033/149] only send 1 line from pw file, explicitly close pw file descriptor before running command --- main-mshell.go | 6 +++++- pkg/shexec/shexec.go | 9 ++++++--- 2 files changed, 11 insertions(+), 4 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 9a942492..6137ab72 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -7,6 +7,7 @@ package main import ( + "bytes" "fmt" "os" "os/signal" @@ -347,7 +348,10 @@ func parseClientOpts() (*shexec.ClientOpts, error) { if err != nil { return nil, fmt.Errorf("cannot read --sudo-with-passfile file '%s': %w", fileName, err) } - opts.SudoPw = string(contents) + if newlineIdx := bytes.Index(contents, []byte{'\n'}); newlineIdx != -1 { + contents = contents[0:newlineIdx] + } + opts.SudoPw = string(contents) + "\n" continue } if argStr == "--" { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 09fcf167..bd6d79b8 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -42,8 +42,8 @@ fi ` const RunCommandFmt = `%s` -const RunSudoCommandFmt = `sudo -C %d bash /dev/fd/%d` -const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -S -C %d bash -c "echo '[from-mshell]'; bash /dev/fd/%d < /dev/fd/%d"` +const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` +const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` type ShExecType struct { Lock *sync.Mutex @@ -281,7 +281,7 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { opts.Fds = append(opts.Fds, commandStdinRfd) opts.CommandStdinFdNum = commandStdinFdNum maxFdNum := opts.MaxFdNum() - runPacket.Command = fmt.Sprintf(RunSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, commandFdNum, commandStdinFdNum) + runPacket.Command = fmt.Sprintf(RunSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, pwFdNum, commandFdNum, commandStdinFdNum) runPacket.Fds = opts.Fds return runPacket, nil } else { @@ -423,6 +423,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) } versionOk = true + if opts.Debug { + fmt.Printf("VERSION> %s\n", initPk.Version) + } break } } From 26479f59c0d6e0dd0cf03d77d58fc71f4da94b1b Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 18:42:56 -0700 Subject: [PATCH 034/149] pass uname back when mshell isn't found, parse, and give install command --- main-mshell.go | 4 +++ pkg/mpio/mpio.go | 4 +-- pkg/packet/packet.go | 1 + pkg/packet/parser.go | 11 ++++--- pkg/shexec/shexec.go | 74 ++++++++++++++++++++++++++++++++++++++++++-- 5 files changed, 85 insertions(+), 9 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 6137ab72..d83aee3c 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -320,6 +320,10 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.Cwd = iter.Next() continue } + if argStr == "--detach" { + opts.Detach = true + continue + } if argStr == "--debug" { opts.Debug = true continue diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 227625a1..6fce2061 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -16,8 +16,8 @@ import ( "github.com/scripthaus-dev/mshell/pkg/packet" ) -const ReadBufSize = 32 * 1024 -const WriteBufSize = 32 * 1024 +const ReadBufSize = 128 * 1024 +const WriteBufSize = 128 * 1024 const MaxSingleWriteSize = 4 * 1024 type Multiplexer struct { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 3638a04f..01191ed3 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -341,6 +341,7 @@ type InitPacketType struct { Env []string `json:"env,omitempty"` User string `json:"user,omitempty"` NotFound bool `json:"notfound,omitempty"` + UName string `json:"uname,omitempty"` } func (*InitPacketType) GetType() string { diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index 907ae83f..40e59b24 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -76,10 +76,13 @@ func MakePacketParser(input io.Reader) *PacketParser { parser.MainCh <- MakeRawPacket(line[:len(line)-1]) continue } - packetLen, err := strconv.Atoi(line[2:bracePos]) - if err != nil || packetLen != len(line)-bracePos-1 { - parser.MainCh <- MakeRawPacket(line[:len(line)-1]) - continue + packetLen := -1 + if line[2:bracePos] != "N" { + packetLen, err = strconv.Atoi(line[2:bracePos]) + if err != nil || packetLen != len(line)-bracePos-1 { + parser.MainCh <- MakeRawPacket(line[:len(line)-1]) + continue + } } pk, err := ParseJsonPacket([]byte(line[bracePos:])) if err != nil { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index bd6d79b8..78316a2d 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -32,10 +32,10 @@ const FirstExtraFilesFdNum = 3 const ClientCommand = ` PATH=$PATH:~/.mshell; -which mshell > /dev/null; +which mshell2 > /dev/null; if [[ "$?" -ne 0 ]] then - printf "\n##34{\"type\": \"init\", \"notfound\": true}\n" + printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s | %s\"}\n" "$(uname -s)" "$(uname -m)" else mshell --single fi @@ -175,6 +175,19 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { if err != nil { return err } + for _, rfd := range pk.Fds { + if rfd.Write { + return fmt.Errorf("cannot detach command with writable remote files fd=%d", rfd.FdNum) + } + if rfd.Read { + if rfd.Content == "" { + return fmt.Errorf("cannot detach command with readable remote files fd=%d", rfd.FdNum) + } + if len(rfd.Content) > mpio.ReadBufSize { + return fmt.Errorf("cannot detach command, constant readable input too large fd=%d, len=%d, max=%d", rfd.FdNum, len(rfd.Content), mpio.ReadBufSize) + } + } + } } if pk.Cwd != "" { realCwd := base.ExpandHomeDir(pk.Cwd) @@ -227,6 +240,7 @@ type ClientOpts struct { SudoWithPass bool SudoPw string CommandStdinFdNum int + Detach bool } func (opts *ClientOpts) MakeExecCmd() *exec.Cmd { @@ -251,8 +265,30 @@ func (opts *ClientOpts) MakeExecCmd() *exec.Cmd { } } +func (opts *ClientOpts) MakeInstallCommandString(goos string, goarch string) string { + var moreSSHOpts []string + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + moreSSHOpts = append(moreSSHOpts, identityOpt) + } + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + moreSSHOpts = append(moreSSHOpts, userOpt) + } + if opts.SSHOptsStr != "" { + optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOptsStr)) + moreSSHOpts = append(moreSSHOpts, optsOpt) + } + if opts.SSHHost != "" { + sshArg := fmt.Sprintf("--ssh %s", shellescape.Quote(opts.SSHHost)) + moreSSHOpts = append(moreSSHOpts, sshArg) + } + return fmt.Sprintf("mshell --install %s %s_%s", strings.Join(moreSSHOpts, " "), goos, goarch) +} + func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket := packet.MakeRunPacket() + runPacket.Detached = opts.Detach runPacket.Cwd = opts.Cwd runPacket.Fds = opts.Fds if !opts.Sudo { @@ -417,7 +453,16 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if pk.GetType() == packet.InitPacketStr { initPk := pk.(*packet.InitPacketType) if initPk.NotFound { - return nil, fmt.Errorf("mshell command not found on remote server, can install with 'mshell --install'") + fmt.Printf("UNAME> %s\n", initPk.UName) + if initPk.UName == "" { + return nil, fmt.Errorf("mshell command not found on remote server, no uname detected") + } + goos, goarch, err := UNameStringToGoArch(initPk.UName) + if err != nil { + return nil, fmt.Errorf("mshell command not found on remote server, architecture cannot be detected (might be incompatible with mshell): %w", err) + } + installCmd := opts.MakeInstallCommandString(goos, goarch) + return nil, fmt.Errorf("mshell command not found on remote server, can install with '%s'", installCmd) } if initPk.Version != "0.1.0" { return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) @@ -444,6 +489,29 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return donePacket, nil } +func UNameStringToGoArch(uname string) (string, string, error) { + fields := strings.SplitN(uname, "|", 2) + if len(fields) != 2 { + return "", "", fmt.Errorf("invalid uname string returned") + } + osVal := strings.TrimSpace(strings.ToLower(fields[0])) + archVal := strings.TrimSpace(strings.ToLower(fields[1])) + if osVal != "darwin" && osVal != "linux" { + return "", "", fmt.Errorf("invalid uname OS '%s', mshell only supports OS X (darwin) and linux", osVal) + } + goos := osVal + goarch := "" + if archVal == "x86_64" || archVal == "i686" || archVal == "amd64" { + goarch = "amd64" + } else if archVal == "aarch64" || archVal == "amd64" { + goarch = "arm64" + } + if goarch == "" { + return "", "", fmt.Errorf("invalid uname machine type '%s', mshell only supports aarch64 (amd64) and x86_64 (amd64)", archVal) + } + return goos, goarch, nil +} + func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sender *packet.PacketSender) { defer cmd.Close() cmd.Multiplexer.RunIOAndWait(packetParser, sender, true, false, false) From afd3bdb315aa14b33faca1695156ea39f499385f Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 22:39:16 -0700 Subject: [PATCH 035/149] implement install command --- main-mshell.go | 132 ++++++++++++++++++++++++++++++++---------- pkg/base/base.go | 9 +++ pkg/base/optsiter.go | 7 +++ pkg/shexec/shexec.go | 133 +++++++++++++++++++++++++++++++++++-------- 4 files changed, 229 insertions(+), 52 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index d83aee3c..db6e34e5 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -23,8 +23,6 @@ import ( "golang.org/x/sys/unix" ) -const MShellVersion = "0.1.0" - // in single run mode, we don't want mshell to die from signals // since we want the single mshell to persist even if session / main mshell // is terminated. @@ -220,8 +218,14 @@ func handleSingle() { close(sender.SendCh) <-sender.DoneCh }() + if len(os.Args) >= 3 && os.Args[2] == "--version" { + initPacket := packet.MakeInitPacket() + initPacket.Version = base.MShellVersion + sender.SendPacket(initPacket) + return + } initPacket := packet.MakeInitPacket() - initPacket.Version = MShellVersion + initPacket.Version = base.MShellVersion sender.SendPacket(initPacket) var runPacket *packet.RunPacketType for pk := range packetParser.MainCh { @@ -280,37 +284,70 @@ func detectOpenFds() ([]packet.RemoteFd, error) { return fds, nil } +func parseInstallOpts() (*shexec.InstallOpts, error) { + opts := &shexec.InstallOpts{} + iter := base.MakeOptsIter(os.Args[2:]) // first arg is --install + for iter.HasNext() { + argStr := iter.Next() + found, err := tryParseSSHOpt(iter, &opts.SSHOpts) + if err != nil { + return nil, err + } + if found { + continue + } + if base.IsOption(argStr) { + return nil, fmt.Errorf("invalid option '%s' passed to mshell --install", argStr) + } + opts.ArchStr = argStr + break + } + return opts, nil +} + +func tryParseSSHOpt(iter *base.OptsIter, sshOpts *shexec.SharedSSHOpts) (bool, error) { + argStr := iter.Current() + if argStr == "--ssh" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("'--ssh [user@host]' missing host") + } + sshOpts.SSHHost = iter.Next() + return true, nil + } + if argStr == "--ssh-opts" { + if !iter.HasNext() { + return false, fmt.Errorf("'--ssh-opts [options]' missing options") + } + sshOpts.SSHOptsStr = iter.Next() + return true, nil + } + if argStr == "-i" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("-i [identity-file]' missing file") + } + sshOpts.SSHIdentity = iter.Next() + return true, nil + } + if argStr == "-l" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("-l [user]' missing user") + } + sshOpts.SSHUser = iter.Next() + return true, nil + } + return false, nil +} + func parseClientOpts() (*shexec.ClientOpts, error) { opts := &shexec.ClientOpts{} iter := base.MakeOptsIter(os.Args[1:]) for iter.HasNext() { argStr := iter.Next() - if argStr == "--ssh" { - if !iter.IsNextPlain() { - return nil, fmt.Errorf("'--ssh [user@host]' missing host") - } - opts.SSHHost = iter.Next() - continue + found, err := tryParseSSHOpt(iter, &opts.SSHOpts) + if err != nil { + return nil, err } - if argStr == "--ssh-opts" { - if !iter.HasNext() { - return nil, fmt.Errorf("'--ssh-opts [options]' missing options") - } - opts.SSHOptsStr = iter.Next() - continue - } - if argStr == "-i" { - if !iter.IsNextPlain() { - return nil, fmt.Errorf("-i [identity-file]' missing file") - } - opts.SSHIdentity = iter.Next() - continue - } - if argStr == "-l" { - if !iter.IsNextPlain() { - return nil, fmt.Errorf("-l [user]' missing user") - } - opts.SSHUser = iter.Next() + if found { continue } if argStr == "--cwd" { @@ -392,6 +429,36 @@ func handleClient() (int, error) { return donePacket.ExitCode, nil } +func handleInstall() (int, error) { + opts, err := parseInstallOpts() + if err != nil { + return 1, fmt.Errorf("parsing opts: %w", err) + } + if opts.SSHOpts.SSHHost == "" { + return 1, fmt.Errorf("cannot install without '--ssh user@host' option") + } + fullArch := opts.ArchStr + fields := strings.SplitN(fullArch, ".", 2) + if len(fields) != 2 { + return 1, fmt.Errorf("invalid arch format '%s' passed to mshell --install", fullArch) + } + goos, goarch := fields[0], fields[1] + if !base.ValidGoArch(goos, goarch) { + return 1, fmt.Errorf("invalid arch '%s' passed to mshell --install", fullArch) + } + optName := base.GoArchOptFile(goos, goarch) + _, err = os.Stat(optName) + if err != nil { + return 1, fmt.Errorf("cannot install mshell to remote host, cannot read '%s': %w", optName, err) + } + opts.OptName = optName + err = shexec.RunInstallSSHCommand(opts) + if err != nil { + return 1, err + } + return 0, nil +} + func handleUsage() { usage := ` Client Usage: mshell [opts] --ssh user@host -- [command] @@ -444,7 +511,7 @@ func main() { handleUsage() return } else if firstArg == "--version" { - fmt.Printf("mshell v%s\n", MShellVersion) + fmt.Printf("mshell v%s\n", base.MShellVersion) return } else if firstArg == "--single" { handleSingle() @@ -452,6 +519,13 @@ func main() { } else if firstArg == "--server" { handleServer() return + } else if firstArg == "--install" { + rtnCode, err := handleInstall() + if err != nil { + fmt.Printf("[error] %v\n", err) + } + os.Exit(rtnCode) + return } else { rtnCode, err := handleClient() if err != nil { diff --git a/pkg/base/base.go b/pkg/base/base.go index 39398e01..b2736c4c 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -30,6 +30,7 @@ const SessionsDirBaseName = ".sessions" const RunnerBaseName = "runner" const SessionDBName = "session.db" const ScReadyString = "scripthaus runner ready" +const MShellVersion = "0.1.0" const OSCEscError = "error" @@ -253,3 +254,11 @@ func ExpandHomeDir(pathStr string) string { } return path.Join(homeDir, pathStr[2:]) } + +func ValidGoArch(goos string, goarch string) bool { + return (goos == "darwin" || goos == "linux") && (goarch == "amd64" || goarch == "arm64") +} + +func GoArchOptFile(goos string, goarch string) string { + return fmt.Sprintf("/opt/mshell/bin/mshell.%s.%s", goos, goarch) +} diff --git a/pkg/base/optsiter.go b/pkg/base/optsiter.go index 923d86d2..aee544e3 100644 --- a/pkg/base/optsiter.go +++ b/pkg/base/optsiter.go @@ -41,6 +41,13 @@ func (iter *OptsIter) Next() string { return rtn } +func (iter *OptsIter) Current() string { + if iter.Pos == 0 { + return "" + } + return iter.Opts[iter.Pos-1] +} + func (iter *OptsIter) Rest() []string { return iter.Opts[iter.Pos:] } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 78316a2d..814cfac4 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -32,7 +32,7 @@ const FirstExtraFilesFdNum = 3 const ClientCommand = ` PATH=$PATH:~/.mshell; -which mshell2 > /dev/null; +which mshell > /dev/null; if [[ "$?" -ne 0 ]] then printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s | %s\"}\n" "$(uname -s)" "$(uname -m)" @@ -41,6 +41,14 @@ else fi ` +const InstallCommand = ` +mkdir -p ~/.mshell/; +cat > ~/.mshell/mshell.temp; +mv ~/.mshell/mshell.temp ~/.mshell/mshell; +chmod a+x ~/.mshell/mshell; +~/.mshell/mshell --single --version +` + const RunCommandFmt = `%s` const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` @@ -227,11 +235,21 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } } +type SharedSSHOpts struct { + SSHHost string + SSHOptsStr string + SSHIdentity string + SSHUser string +} + +type InstallOpts struct { + SSHOpts SharedSSHOpts + ArchStr string + OptName string +} + type ClientOpts struct { - SSHHost string - SSHOptsStr string - SSHIdentity string - SSHUser string + SSHOpts SharedSSHOpts Command string Fds []packet.RemoteFd Cwd string @@ -244,46 +262,63 @@ type ClientOpts struct { } func (opts *ClientOpts) MakeExecCmd() *exec.Cmd { - if opts.SSHHost == "" { + if opts.SSHOpts.SSHHost == "" { ecmd := exec.Command("bash", "-c", strings.TrimSpace(ClientCommand)) return ecmd } else { var moreSSHOpts []string - if opts.SSHIdentity != "" { - identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + if opts.SSHOpts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHOpts.SSHIdentity)) moreSSHOpts = append(moreSSHOpts, identityOpt) } - if opts.SSHUser != "" { - userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + if opts.SSHOpts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHOpts.SSHUser)) moreSSHOpts = append(moreSSHOpts, userOpt) } remoteCommand := strings.TrimSpace(ClientCommand) // note that SSHOptsStr is *not* escaped - sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOpts.SSHOptsStr, shellescape.Quote(opts.SSHOpts.SSHHost), shellescape.Quote(remoteCommand)) ecmd := exec.Command("bash", "-c", sshCmd) return ecmd } } -func (opts *ClientOpts) MakeInstallCommandString(goos string, goarch string) string { +func (opts *InstallOpts) MakeExecCmd() *exec.Cmd { var moreSSHOpts []string - if opts.SSHIdentity != "" { - identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + if opts.SSHOpts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHOpts.SSHIdentity)) moreSSHOpts = append(moreSSHOpts, identityOpt) } - if opts.SSHUser != "" { - userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + if opts.SSHOpts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHOpts.SSHUser)) moreSSHOpts = append(moreSSHOpts, userOpt) } - if opts.SSHOptsStr != "" { - optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOptsStr)) + // note that SSHOptsStr is *not* escaped + installCommand := strings.TrimSpace(InstallCommand) + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOpts.SSHOptsStr, shellescape.Quote(opts.SSHOpts.SSHHost), shellescape.Quote(installCommand)) + ecmd := exec.Command("bash", "-c", sshCmd) + return ecmd +} + +func (opts *ClientOpts) MakeInstallCommandString(goos string, goarch string) string { + var moreSSHOpts []string + if opts.SSHOpts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHOpts.SSHIdentity)) + moreSSHOpts = append(moreSSHOpts, identityOpt) + } + if opts.SSHOpts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHOpts.SSHUser)) + moreSSHOpts = append(moreSSHOpts, userOpt) + } + if opts.SSHOpts.SSHOptsStr != "" { + optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOpts.SSHOptsStr)) moreSSHOpts = append(moreSSHOpts, optsOpt) } - if opts.SSHHost != "" { - sshArg := fmt.Sprintf("--ssh %s", shellescape.Quote(opts.SSHHost)) + if opts.SSHOpts.SSHHost != "" { + sshArg := fmt.Sprintf("--ssh %s", shellescape.Quote(opts.SSHOpts.SSHHost)) moreSSHOpts = append(moreSSHOpts, sshArg) } - return fmt.Sprintf("mshell --install %s %s_%s", strings.Join(moreSSHOpts, " "), goos, goarch) + return fmt.Sprintf("mshell --install %s %s.%s", strings.Join(moreSSHOpts, " "), goos, goarch) } func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { @@ -383,6 +418,55 @@ func ValidateRemoteFds(rfds []packet.RemoteFd) error { return nil } +func RunInstallSSHCommand(opts *InstallOpts) error { + ecmd := opts.MakeExecCmd() + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return fmt.Errorf("creating stdin pipe: %v", err) + } + stdoutReader, err := ecmd.StdoutPipe() + if err != nil { + return fmt.Errorf("creating stdout pipe: %v", err) + } + stderrReader, err := ecmd.StderrPipe() + if err != nil { + return fmt.Errorf("creating stderr pipe: %v", err) + } + go func() { + io.Copy(os.Stderr, stderrReader) + }() + fd, err := os.Open(opts.OptName) + if err != nil { + return fmt.Errorf("cannot open '%s': %w", opts.OptName, err) + } + go func() { + defer inputWriter.Close() + io.Copy(inputWriter, fd) + }() + packetParser := packet.MakePacketParser(stdoutReader) + err = ecmd.Start() + if err != nil { + return fmt.Errorf("running ssh command: %w", err) + } + for pk := range packetParser.MainCh { + if pk.GetType() == packet.InitPacketStr { + initPacket := pk.(*packet.InitPacketType) + if initPacket.Version == base.MShellVersion { + fmt.Printf("mshell %s, installed successfully at %s:~/.mshell/mshell\n", initPacket.Version, opts.SSHOpts.SSHHost) + return nil + } + return fmt.Errorf("invalid version '%s' received from client, expecting '%s'", initPacket.Version, base.MShellVersion) + } + if pk.GetType() == packet.RawPacketStr { + rawPk := pk.(*packet.RawPacketType) + fmt.Printf("%s\n", rawPk.Data) + continue + } + return fmt.Errorf("invalid response packet '%s' received from client", pk.GetType()) + } + return fmt.Errorf("did not receive version string from client, install not successful") +} + func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, error) { err := ValidateRemoteFds(opts.Fds) if err != nil { @@ -464,8 +548,8 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er installCmd := opts.MakeInstallCommandString(goos, goarch) return nil, fmt.Errorf("mshell command not found on remote server, can install with '%s'", installCmd) } - if initPk.Version != "0.1.0" { - return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) + if initPk.Version != base.MShellVersion { + return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) } versionOk = true if opts.Debug { @@ -509,6 +593,9 @@ func UNameStringToGoArch(uname string) (string, string, error) { if goarch == "" { return "", "", fmt.Errorf("invalid uname machine type '%s', mshell only supports aarch64 (amd64) and x86_64 (amd64)", archVal) } + if !base.ValidGoArch(goos, goarch) { + return "", "", fmt.Errorf("invalid arch detected %s.%s", goos, goarch) + } return goos, goarch, nil } From 9377619e4cd9399200f9340757023d2e72aa7b49 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 27 Jun 2022 23:14:53 -0700 Subject: [PATCH 036/149] write auto-detect logic for arch from uname --- main-mshell.go | 37 ++++++++++++++++++---------- pkg/shexec/shexec.go | 57 +++++++++++++++++++++++++++++++++----------- 2 files changed, 68 insertions(+), 26 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index db6e34e5..103bf451 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -296,6 +296,10 @@ func parseInstallOpts() (*shexec.InstallOpts, error) { if found { continue } + if argStr == "--detect" { + opts.Detect = true + continue + } if base.IsOption(argStr) { return nil, fmt.Errorf("invalid option '%s' passed to mshell --install", argStr) } @@ -402,6 +406,7 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.Command = strings.Join(iter.Rest(), " ") break } + return nil, fmt.Errorf("invalid option '%s' passed to mshell", argStr) } return opts, nil } @@ -437,21 +442,29 @@ func handleInstall() (int, error) { if opts.SSHOpts.SSHHost == "" { return 1, fmt.Errorf("cannot install without '--ssh user@host' option") } - fullArch := opts.ArchStr - fields := strings.SplitN(fullArch, ".", 2) - if len(fields) != 2 { - return 1, fmt.Errorf("invalid arch format '%s' passed to mshell --install", fullArch) + if opts.Detect && opts.ArchStr != "" { + return 1, fmt.Errorf("cannot supply both --detect and arch '%s'", opts.ArchStr) } - goos, goarch := fields[0], fields[1] - if !base.ValidGoArch(goos, goarch) { - return 1, fmt.Errorf("invalid arch '%s' passed to mshell --install", fullArch) + if opts.ArchStr == "" && !opts.Detect { + return 1, fmt.Errorf("must supply an arch string or '--detect' to auto detect") } - optName := base.GoArchOptFile(goos, goarch) - _, err = os.Stat(optName) - if err != nil { - return 1, fmt.Errorf("cannot install mshell to remote host, cannot read '%s': %w", optName, err) + if opts.ArchStr != "" { + fullArch := opts.ArchStr + fields := strings.SplitN(fullArch, ".", 2) + if len(fields) != 2 { + return 1, fmt.Errorf("invalid arch format '%s' passed to mshell --install", fullArch) + } + goos, goarch := fields[0], fields[1] + if !base.ValidGoArch(goos, goarch) { + return 1, fmt.Errorf("invalid arch '%s' passed to mshell --install", fullArch) + } + optName := base.GoArchOptFile(goos, goarch) + _, err = os.Stat(optName) + if err != nil { + return 1, fmt.Errorf("cannot install mshell to remote host, cannot read '%s': %w", optName, err) + } + opts.OptName = optName } - opts.OptName = optName err = shexec.RunInstallSSHCommand(opts) if err != nil { return 1, err diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 814cfac4..9a45077b 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -42,6 +42,7 @@ fi ` const InstallCommand = ` +printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s | %s\"}\n" "$(uname -s)" "$(uname -m)"; mkdir -p ~/.mshell/; cat > ~/.mshell/mshell.temp; mv ~/.mshell/mshell.temp ~/.mshell/mshell; @@ -246,6 +247,7 @@ type InstallOpts struct { SSHOpts SharedSSHOpts ArchStr string OptName string + Detect bool } type ClientOpts struct { @@ -294,8 +296,8 @@ func (opts *InstallOpts) MakeExecCmd() *exec.Cmd { moreSSHOpts = append(moreSSHOpts, userOpt) } // note that SSHOptsStr is *not* escaped - installCommand := strings.TrimSpace(InstallCommand) - sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOpts.SSHOptsStr, shellescape.Quote(opts.SSHOpts.SSHHost), shellescape.Quote(installCommand)) + command := strings.TrimSpace(InstallCommand) + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOpts.SSHOptsStr, shellescape.Quote(opts.SSHOpts.SSHHost), shellescape.Quote(command)) ecmd := exec.Command("bash", "-c", sshCmd) return ecmd } @@ -418,7 +420,20 @@ func ValidateRemoteFds(rfds []packet.RemoteFd) error { return nil } +func sendOptFile(input io.WriteCloser, optName string) error { + fd, err := os.Open(optName) + if err != nil { + return fmt.Errorf("cannot open '%s': %w", optName, err) + } + go func() { + defer input.Close() + io.Copy(input, fd) + }() + return nil +} + func RunInstallSSHCommand(opts *InstallOpts) error { + tryDetect := opts.Detect ecmd := opts.MakeExecCmd() inputWriter, err := ecmd.StdinPipe() if err != nil { @@ -435,21 +450,36 @@ func RunInstallSSHCommand(opts *InstallOpts) error { go func() { io.Copy(os.Stderr, stderrReader) }() - fd, err := os.Open(opts.OptName) - if err != nil { - return fmt.Errorf("cannot open '%s': %w", opts.OptName, err) + if opts.OptName != "" { + sendOptFile(inputWriter, opts.OptName) } - go func() { - defer inputWriter.Close() - io.Copy(inputWriter, fd) - }() packetParser := packet.MakePacketParser(stdoutReader) err = ecmd.Start() if err != nil { return fmt.Errorf("running ssh command: %w", err) } + firstInit := true for pk := range packetParser.MainCh { - if pk.GetType() == packet.InitPacketStr { + if pk.GetType() == packet.InitPacketStr && firstInit { + firstInit = false + initPacket := pk.(*packet.InitPacketType) + if !tryDetect { + continue // ignore + } + tryDetect = false + if initPacket.UName == "" { + return fmt.Errorf("cannot detect arch, no uname received from remote server") + } + goos, goarch, err := DetectGoArch(initPacket.UName) + if err != nil { + return fmt.Errorf("arch cannot be detected (might be incompatible with mshell): %w", err) + } + fmt.Printf("mshell detected remote architecture as '%s.%s'\n", goos, goarch) + optName := base.GoArchOptFile(goos, goarch) + sendOptFile(inputWriter, optName) + continue + } + if pk.GetType() == packet.InitPacketStr && !firstInit { initPacket := pk.(*packet.InitPacketType) if initPacket.Version == base.MShellVersion { fmt.Printf("mshell %s, installed successfully at %s:~/.mshell/mshell\n", initPacket.Version, opts.SSHOpts.SSHHost) @@ -537,16 +567,15 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if pk.GetType() == packet.InitPacketStr { initPk := pk.(*packet.InitPacketType) if initPk.NotFound { - fmt.Printf("UNAME> %s\n", initPk.UName) if initPk.UName == "" { return nil, fmt.Errorf("mshell command not found on remote server, no uname detected") } - goos, goarch, err := UNameStringToGoArch(initPk.UName) + goos, goarch, err := DetectGoArch(initPk.UName) if err != nil { return nil, fmt.Errorf("mshell command not found on remote server, architecture cannot be detected (might be incompatible with mshell): %w", err) } installCmd := opts.MakeInstallCommandString(goos, goarch) - return nil, fmt.Errorf("mshell command not found on remote server, can install with '%s'", installCmd) + return nil, fmt.Errorf("mshell command not found on remote server, can install with '%s' (or --auto-install)", installCmd) } if initPk.Version != base.MShellVersion { return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) @@ -573,7 +602,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return donePacket, nil } -func UNameStringToGoArch(uname string) (string, string, error) { +func DetectGoArch(uname string) (string, string, error) { fields := strings.SplitN(uname, "|", 2) if len(fields) != 2 { return "", "", fmt.Errorf("invalid uname string returned") From d7eb2526f07d9a7518798d7610eb4c6936116451 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 28 Jun 2022 15:04:08 -0700 Subject: [PATCH 037/149] refactor RunClientSSHCommandAndWait for server code --- main-mshell.go | 44 +++++++++++------- pkg/cmdtail/cmdtail.go | 12 ++--- pkg/packet/packet.go | 26 ++++++++--- pkg/server/server.go | 53 ++++++++++++++++++++++ pkg/shexec/shexec.go | 100 ++++++++++++++++------------------------- 5 files changed, 146 insertions(+), 89 deletions(-) create mode 100644 pkg/server/server.go diff --git a/main-mshell.go b/main-mshell.go index 103bf451..d43aafd5 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -19,6 +19,7 @@ import ( "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/cmdtail" "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/server" "github.com/scripthaus-dev/mshell/pkg/shexec" "golang.org/x/sys/unix" ) @@ -72,7 +73,7 @@ func doSingle(ck base.CommandKey) { sender.SendPacket(startPacket) donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) - sender.CloseSendCh() + sender.Close() sender.WaitForDone() } @@ -157,7 +158,7 @@ func doMain() { } packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) - tailer, err := cmdtail.MakeTailer(sender.SendCh) + tailer, err := cmdtail.MakeTailer(sender) if err != nil { packet.SendErrorPacket(os.Stdout, err.Error()) return @@ -215,8 +216,8 @@ func handleSingle() { sender := packet.MakePacketSender(os.Stdout) defer func() { // wait for sender to complete - close(sender.SendCh) - <-sender.DoneCh + sender.Close() + sender.WaitForDone() }() if len(os.Args) >= 3 && os.Args[2] == "--version" { initPacket := packet.MakeInitPacket() @@ -259,9 +260,6 @@ func handleSingle() { cmd.RunRemoteIOAndWait(packetParser, sender) } -func handleServer() { -} - func detectOpenFds() ([]packet.RemoteFd, error) { var fds []packet.RemoteFd for fdNum := 3; fdNum <= 64; fdNum++ { @@ -309,7 +307,7 @@ func parseInstallOpts() (*shexec.InstallOpts, error) { return opts, nil } -func tryParseSSHOpt(iter *base.OptsIter, sshOpts *shexec.SharedSSHOpts) (bool, error) { +func tryParseSSHOpt(iter *base.OptsIter, sshOpts *shexec.SSHOpts) (bool, error) { argStr := iter.Current() if argStr == "--ssh" { if !iter.IsNextPlain() { @@ -378,7 +376,7 @@ func parseClientOpts() (*shexec.ClientOpts, error) { return nil, fmt.Errorf("'--sudo-with-password [pw]', missing password") } opts.Sudo = true - opts.SudoWithPass = true + opts.SSHOpts.SudoWithPass = true opts.SudoPw = iter.Next() continue } @@ -387,7 +385,7 @@ func parseClientOpts() (*shexec.ClientOpts, error) { return nil, fmt.Errorf("'--sudo-with-passfile [file]', missing file") } opts.Sudo = true - opts.SudoWithPass = true + opts.SSHOpts.SudoWithPass = true fileName := iter.Next() contents, err := os.ReadFile(fileName) if err != nil { @@ -427,7 +425,15 @@ func handleClient() (int, error) { return 1, err } opts.Fds = fds - donePacket, err := shexec.RunClientSSHCommandAndWait(opts) + err = shexec.ValidateRemoteFds(opts.Fds) + if err != nil { + return 1, err + } + runPacket, err := opts.MakeRunPacket() // modifies opts + if err != nil { + return 1, err + } + donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, opts.SSHOpts, opts.Debug) if err != nil { return 1, err } @@ -530,21 +536,29 @@ func main() { handleSingle() return } else if firstArg == "--server" { - handleServer() + rtnCode, err := server.RunServer() + if err != nil { + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + } + if rtnCode != 0 { + os.Exit(rtnCode) + } return } else if firstArg == "--install" { rtnCode, err := handleInstall() if err != nil { - fmt.Printf("[error] %v\n", err) + fmt.Fprintf(os.Stderr, "[error] %v\n", err) } os.Exit(rtnCode) return } else { rtnCode, err := handleClient() if err != nil { - fmt.Printf("[error] %v\n", err) + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + } + if rtnCode != 0 { + os.Exit(rtnCode) } - os.Exit(rtnCode) return } diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 517b35f8..663f8065 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -77,7 +77,7 @@ type Tailer struct { WatchList map[base.CommandKey]CmdWatchEntry ScHomeDir string Watcher *fsnotify.Watcher - SendCh chan packet.PacketType + Sender *packet.PacketSender } func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos TailPos) { @@ -129,7 +129,7 @@ func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (Cm return entry, pos, true } -func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { +func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { scHomeDir, err := base.GetScHomeDir() if err != nil { return nil, err @@ -138,7 +138,7 @@ func MakeTailer(sendCh chan packet.PacketType) (*Tailer, error) { Lock: &sync.Mutex{}, WatchList: make(map[base.CommandKey]CmdWatchEntry), ScHomeDir: scHomeDir, - SendCh: sendCh, + Sender: sender, } rtn.Watcher, err = fsnotify.NewWatcher() if err != nil { @@ -241,7 +241,7 @@ func (t *Tailer) RunDataTransfer(key base.CommandKey, reqId string) { for { dataPacket, keepRunning := t.runSingleDataTransfer(key, reqId) if dataPacket != nil { - t.SendCh <- dataPacket + t.Sender.SendPacket(dataPacket) } if !keepRunning { t.checkRemoveNoFollow(key, reqId) @@ -273,7 +273,7 @@ func (t *Tailer) updateFile(relFileName string) { } finfo, err := os.Stat(relFileName) if err != nil { - t.SendCh <- packet.FmtMessagePacket("error trying to stat file '%s': %v", relFileName, err) + t.Sender.SendPacket(packet.FmtMessagePacket("error trying to stat file '%s': %v", relFileName, err)) return } cmdKey := base.MakeCommandKey(m[1], m[2]) @@ -311,7 +311,7 @@ func (t *Tailer) Run() { return } // what to do with this error? just send a message - t.SendCh <- packet.FmtMessagePacket("error in tailer: %v", err) + t.Sender.SendPacket(packet.FmtMessagePacket("error in tailer: %v", err)) } } return diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 01191ed3..07f58c19 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -483,6 +483,16 @@ type CommandPacketType interface { GetCK() base.CommandKey } +func AsExtType(pk PacketType) string { + if rpcPacket, ok := pk.(RpcPacketType); ok { + return fmt.Sprintf("%s[%s]", rpcPacket.GetType(), rpcPacket.GetPacketId()) + } else if cmdPacket, ok := pk.(CommandPacketType); ok { + return fmt.Sprintf("%s[%s]", cmdPacket.GetType(), cmdPacket.GetCK()) + } else { + return pk.GetType() + } +} + func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { var bareCmd BarePacketType err := json.Unmarshal(jsonBuf, &bareCmd) @@ -545,12 +555,8 @@ func MakePacketSender(output io.Writer) *PacketSender { DoneCh: make(chan bool), } go func() { - defer func() { - sender.Lock.Lock() - sender.Done = true - sender.Lock.Unlock() - close(sender.DoneCh) - }() + defer close(sender.DoneCh) + defer sender.Close() for pk := range sender.SendCh { err := SendPacket(output, pk) if err != nil { @@ -564,7 +570,13 @@ func MakePacketSender(output io.Writer) *PacketSender { return sender } -func (sender *PacketSender) CloseSendCh() { +func (sender *PacketSender) Close() { + sender.Lock.Lock() + defer sender.Lock.Unlock() + if sender.Done { + return + } + sender.Done = true close(sender.SendCh) } diff --git a/pkg/server/server.go b/pkg/server/server.go new file mode 100644 index 00000000..61bd1760 --- /dev/null +++ b/pkg/server/server.go @@ -0,0 +1,53 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package server + +import ( + "fmt" + "os" + "sync" + + "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +type MServer struct { + Lock *sync.Mutex + MainInput *packet.PacketParser + Sender *packet.PacketSender +} + +func (m *MServer) Close() { + m.Sender.Close() + m.Sender.WaitForDone() +} + +func RunServer() (int, error) { + server := &MServer{ + Lock: &sync.Mutex{}, + } + server.MainInput = packet.MakePacketParser(os.Stdin) + server.Sender = packet.MakePacketSender(os.Stdout) + defer server.Close() + initPacket := packet.MakeInitPacket() + initPacket.Version = base.MShellVersion + server.Sender.SendPacket(initPacket) + for pk := range server.MainInput.MainCh { + fmt.Printf("PK> %s\n", packet.AsString(pk)) + if pk.GetType() == packet.PingPacketStr { + continue + } + if pk.GetType() == packet.RunPacketStr { + runPacket := pk.(*packet.RunPacketType) + fmt.Printf("RUN> %s\n", runPacket) + continue + } + server.Sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", packet.AsExtType(pk))) + continue + } + return 0, nil +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 9a45077b..f0334e3f 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -236,91 +236,74 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } } -type SharedSSHOpts struct { - SSHHost string - SSHOptsStr string - SSHIdentity string - SSHUser string +type SSHOpts struct { + SSHHost string + SSHOptsStr string + SSHIdentity string + SSHUser string + SudoWithPass bool } type InstallOpts struct { - SSHOpts SharedSSHOpts + SSHOpts SSHOpts ArchStr string OptName string Detect bool } type ClientOpts struct { - SSHOpts SharedSSHOpts + SSHOpts SSHOpts Command string Fds []packet.RemoteFd Cwd string Debug bool Sudo bool - SudoWithPass bool SudoPw string CommandStdinFdNum int Detach bool } -func (opts *ClientOpts) MakeExecCmd() *exec.Cmd { - if opts.SSHOpts.SSHHost == "" { - ecmd := exec.Command("bash", "-c", strings.TrimSpace(ClientCommand)) +func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { + remoteCommand = strings.TrimSpace(remoteCommand) + if opts.SSHHost == "" { + ecmd := exec.Command("bash", "-c", remoteCommand) return ecmd } else { var moreSSHOpts []string - if opts.SSHOpts.SSHIdentity != "" { - identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHOpts.SSHIdentity)) + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) moreSSHOpts = append(moreSSHOpts, identityOpt) } - if opts.SSHOpts.SSHUser != "" { - userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHOpts.SSHUser)) + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) moreSSHOpts = append(moreSSHOpts, userOpt) } - remoteCommand := strings.TrimSpace(ClientCommand) // note that SSHOptsStr is *not* escaped - sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOpts.SSHOptsStr, shellescape.Quote(opts.SSHOpts.SSHHost), shellescape.Quote(remoteCommand)) + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) ecmd := exec.Command("bash", "-c", sshCmd) return ecmd } } -func (opts *InstallOpts) MakeExecCmd() *exec.Cmd { +func (opts SSHOpts) MakeMShellSSHOpts() string { var moreSSHOpts []string - if opts.SSHOpts.SSHIdentity != "" { - identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHOpts.SSHIdentity)) + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) moreSSHOpts = append(moreSSHOpts, identityOpt) } - if opts.SSHOpts.SSHUser != "" { - userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHOpts.SSHUser)) + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) moreSSHOpts = append(moreSSHOpts, userOpt) } - // note that SSHOptsStr is *not* escaped - command := strings.TrimSpace(InstallCommand) - sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOpts.SSHOptsStr, shellescape.Quote(opts.SSHOpts.SSHHost), shellescape.Quote(command)) - ecmd := exec.Command("bash", "-c", sshCmd) - return ecmd -} - -func (opts *ClientOpts) MakeInstallCommandString(goos string, goarch string) string { - var moreSSHOpts []string - if opts.SSHOpts.SSHIdentity != "" { - identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHOpts.SSHIdentity)) - moreSSHOpts = append(moreSSHOpts, identityOpt) - } - if opts.SSHOpts.SSHUser != "" { - userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHOpts.SSHUser)) - moreSSHOpts = append(moreSSHOpts, userOpt) - } - if opts.SSHOpts.SSHOptsStr != "" { - optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOpts.SSHOptsStr)) + if opts.SSHOptsStr != "" { + optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOptsStr)) moreSSHOpts = append(moreSSHOpts, optsOpt) } - if opts.SSHOpts.SSHHost != "" { - sshArg := fmt.Sprintf("--ssh %s", shellescape.Quote(opts.SSHOpts.SSHHost)) + if opts.SSHHost != "" { + sshArg := fmt.Sprintf("--ssh %s", shellescape.Quote(opts.SSHHost)) moreSSHOpts = append(moreSSHOpts, sshArg) } - return fmt.Sprintf("mshell --install %s %s.%s", strings.Join(moreSSHOpts, " "), goos, goarch) + return strings.Join(moreSSHOpts, " ") } func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { @@ -333,7 +316,7 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket.Command = fmt.Sprintf(RunCommandFmt, opts.Command) return runPacket, nil } - if opts.SudoWithPass { + if opts.SSHOpts.SudoWithPass { pwFdNum, err := opts.NextFreeFdNum() if err != nil { return nil, err @@ -434,7 +417,7 @@ func sendOptFile(input io.WriteCloser, optName string) error { func RunInstallSSHCommand(opts *InstallOpts) error { tryDetect := opts.Detect - ecmd := opts.MakeExecCmd() + ecmd := opts.SSHOpts.MakeSSHExecCmd(InstallCommand) inputWriter, err := ecmd.StdinPipe() if err != nil { return fmt.Errorf("creating stdin pipe: %v", err) @@ -497,17 +480,9 @@ func RunInstallSSHCommand(opts *InstallOpts) error { return fmt.Errorf("did not receive version string from client, install not successful") } -func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, error) { - err := ValidateRemoteFds(opts.Fds) - if err != nil { - return nil, err - } - runPacket, err := opts.MakeRunPacket() // modifies opts - if err != nil { - return nil, err - } +func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, sshOpts SSHOpts, debug bool) (*packet.CmdDonePacketType, error) { cmd := MakeShExec("") - ecmd := opts.MakeExecCmd() + ecmd := sshOpts.MakeSSHExecCmd(ClientCommand) cmd.Cmd = ecmd inputWriter, err := ecmd.StdinPipe() if err != nil { @@ -521,7 +496,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if err != nil { return nil, fmt.Errorf("creating stderr pipe: %v", err) } - if !opts.SudoWithPass { + if !sshOpts.SudoWithPass { cmd.Multiplexer.MakeRawFdReader(0, os.Stdin, false) } cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout, false) @@ -567,6 +542,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if pk.GetType() == packet.InitPacketStr { initPk := pk.(*packet.InitPacketType) if initPk.NotFound { + if sshOpts.SSHHost == "" { + return nil, fmt.Errorf("mshell command not found on local server") + } if initPk.UName == "" { return nil, fmt.Errorf("mshell command not found on remote server, no uname detected") } @@ -574,14 +552,14 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er if err != nil { return nil, fmt.Errorf("mshell command not found on remote server, architecture cannot be detected (might be incompatible with mshell): %w", err) } - installCmd := opts.MakeInstallCommandString(goos, goarch) - return nil, fmt.Errorf("mshell command not found on remote server, can install with '%s' (or --auto-install)", installCmd) + sshOptsStr := sshOpts.MakeMShellSSHOpts() + return nil, fmt.Errorf("mshell command not found on remote server, can install with 'mshell --install %s %s.%s'", sshOptsStr, goos, goarch) } if initPk.Version != base.MShellVersion { return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) } versionOk = true - if opts.Debug { + if debug { fmt.Printf("VERSION> %s\n", initPk.Version) } break @@ -591,7 +569,7 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, fmt.Errorf("did not receive version from remote mshell") } sender.SendPacket(runPacket) - if opts.Debug { + if debug { cmd.Multiplexer.Debug = true } remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetParser, sender, false, true, true) From 1d44afc10e5bd533f1b84013c24627d771cf7328 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 28 Jun 2022 17:20:01 -0700 Subject: [PATCH 038/149] working on server mode. extract fdcontext as interface. create packet writer/reader for mpio. hook up to serverFdContext. --- main-mshell.go | 6 +-- pkg/mpio/bufreader.go | 5 +-- pkg/mpio/bufwriter.go | 6 +-- pkg/mpio/mpio.go | 5 ++- pkg/mpio/packetreader.go | 96 ++++++++++++++++++++++++++++++++++++++++ pkg/mpio/packetwriter.go | 40 +++++++++++++++++ pkg/server/server.go | 71 ++++++++++++++++++++++++++++- pkg/shexec/shexec.go | 76 ++++++++++++++++++++++++------- 8 files changed, 276 insertions(+), 29 deletions(-) create mode 100644 pkg/mpio/packetreader.go create mode 100644 pkg/mpio/packetwriter.go diff --git a/main-mshell.go b/main-mshell.go index d43aafd5..b8b2b85f 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -376,7 +376,7 @@ func parseClientOpts() (*shexec.ClientOpts, error) { return nil, fmt.Errorf("'--sudo-with-password [pw]', missing password") } opts.Sudo = true - opts.SSHOpts.SudoWithPass = true + opts.SudoWithPass = true opts.SudoPw = iter.Next() continue } @@ -385,7 +385,7 @@ func parseClientOpts() (*shexec.ClientOpts, error) { return nil, fmt.Errorf("'--sudo-with-passfile [file]', missing file") } opts.Sudo = true - opts.SSHOpts.SudoWithPass = true + opts.SudoWithPass = true fileName := iter.Next() contents, err := os.ReadFile(fileName) if err != nil { @@ -433,7 +433,7 @@ func handleClient() (int, error) { if err != nil { return 1, err } - donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, opts.SSHOpts, opts.Debug) + donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, shexec.StdContext{}, opts.SSHOpts, opts.Debug) if err != nil { return 1, err } diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index 654d8d68..bcda317e 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -8,7 +8,6 @@ package mpio import ( "io" - "os" "sync" "github.com/scripthaus-dev/mshell/pkg/packet" @@ -18,13 +17,13 @@ type FdReader struct { CVar *sync.Cond M *Multiplexer FdNum int - Fd *os.File + Fd io.ReadCloser BufSize int Closed bool ShouldCloseFd bool } -func MakeFdReader(m *Multiplexer, fd *os.File, fdNum int, shouldCloseFd bool) *FdReader { +func MakeFdReader(m *Multiplexer, fd io.ReadCloser, fdNum int, shouldCloseFd bool) *FdReader { fr := &FdReader{ CVar: sync.NewCond(&sync.Mutex{}), M: m, diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index c1b36e2b..977a2493 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -8,7 +8,7 @@ package mpio import ( "fmt" - "os" + "io" "sync" ) @@ -17,13 +17,13 @@ type FdWriter struct { M *Multiplexer FdNum int Buffer []byte - Fd *os.File + Fd io.WriteCloser Eof bool Closed bool ShouldCloseFd bool } -func MakeFdWriter(m *Multiplexer, fd *os.File, fdNum int, shouldCloseFd bool) *FdWriter { +func MakeFdWriter(m *Multiplexer, fd io.WriteCloser, fdNum int, shouldCloseFd bool) *FdWriter { fw := &FdWriter{ CVar: sync.NewCond(&sync.Mutex{}), Fd: fd, diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 6fce2061..ca0a0ffa 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -9,6 +9,7 @@ package mpio import ( "encoding/base64" "fmt" + "io" "os" "sync" @@ -111,13 +112,13 @@ func (m *Multiplexer) MakeStringFdReader(fdNum int, contents string) error { return nil } -func (m *Multiplexer) MakeRawFdReader(fdNum int, fd *os.File, shouldClose bool) { +func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose bool) { m.Lock.Lock() defer m.Lock.Unlock() m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose) } -func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd *os.File, shouldClose bool) { +func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd io.WriteCloser, shouldClose bool) { m.Lock.Lock() defer m.Lock.Unlock() m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, shouldClose) diff --git a/pkg/mpio/packetreader.go b/pkg/mpio/packetreader.go new file mode 100644 index 00000000..8dfb3024 --- /dev/null +++ b/pkg/mpio/packetreader.go @@ -0,0 +1,96 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package mpio + +import ( + "encoding/base64" + "errors" + "io" + "sync" + + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +type PacketReader struct { + CVar *sync.Cond + FdNum int + Buf []byte + Eof bool + Err error +} + +func MakePacketReader(fdNum int) *PacketReader { + return &PacketReader{ + CVar: sync.NewCond(&sync.Mutex{}), + FdNum: fdNum, + } +} + +func (pr *PacketReader) AddData(pk *packet.DataPacketType) { + pr.CVar.L.Lock() + defer pr.CVar.L.Unlock() + defer pr.CVar.Broadcast() + if pr.Eof || pr.Err != nil { + return + } + if pk.Data64 != "" { + realData, err := base64.StdEncoding.DecodeString(pk.Data64) + if err != nil { + pr.Err = err + return + } + pr.Buf = append(pr.Buf, realData...) + } + pr.Eof = pk.Eof + if pk.Error != "" { + pr.Err = errors.New(pk.Error) + } + return +} + +func (pr *PacketReader) Read(buf []byte) (int, error) { + pr.CVar.L.Lock() + defer pr.CVar.L.Unlock() + for { + if pr.Err != nil { + return 0, pr.Err + } + if pr.Eof { + return 0, io.EOF + } + if len(pr.Buf) == 0 { + pr.CVar.Wait() + continue + } + nr := copy(buf, pr.Buf) + pr.Buf = pr.Buf[nr:] + if len(pr.Buf) == 0 { + pr.Buf = nil + } + return nr, nil + } +} + +func (pr *PacketReader) Close() error { + pr.CVar.L.Lock() + defer pr.CVar.L.Unlock() + defer pr.CVar.Broadcast() + if pr.Err == nil { + pr.Err = io.ErrClosedPipe + } + return nil +} + +type NullReader struct{} + +func (NullReader) Read(buf []byte) (int, error) { + return 0, io.EOF +} + +func (NullReader) Close() error { + return nil +} diff --git a/pkg/mpio/packetwriter.go b/pkg/mpio/packetwriter.go new file mode 100644 index 00000000..0665f044 --- /dev/null +++ b/pkg/mpio/packetwriter.go @@ -0,0 +1,40 @@ +// Copyright 2022 Dashborg Inc +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at https://mozilla.org/MPL/2.0/. + +package mpio + +import ( + "encoding/base64" + + "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +type PacketWriter struct { + FdNum int + Sender *packet.PacketSender + CK base.CommandKey +} + +func MakePacketWriter(fdNum int, sender *packet.PacketSender, ck base.CommandKey) *PacketWriter { + return &PacketWriter{FdNum: fdNum, Sender: sender, CK: ck} +} + +func (pw *PacketWriter) Write(data []byte) (int, error) { + pk := packet.MakeDataPacket() + pk.CK = pw.CK + pk.FdNum = pw.FdNum + pk.Data64 = base64.StdEncoding.EncodeToString(data) + return len(data), pw.Sender.SendPacket(pk) +} + +func (pw *PacketWriter) Close() error { + pk := packet.MakeDataPacket() + pk.CK = pw.CK + pk.FdNum = pw.FdNum + pk.Eof = true + return pw.Sender.SendPacket(pk) +} diff --git a/pkg/server/server.go b/pkg/server/server.go index 61bd1760..0400d132 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -8,17 +8,21 @@ package server import ( "fmt" + "io" "os" "sync" "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/shexec" ) type MServer struct { Lock *sync.Mutex MainInput *packet.PacketParser Sender *packet.PacketSender + FdContext *serverFdContext } func (m *MServer) Close() { @@ -26,13 +30,73 @@ func (m *MServer) Close() { m.Sender.WaitForDone() } +type serverFdContext struct { + M *MServer + Lock *sync.Mutex + Sender *packet.PacketSender + CK base.CommandKey + Readers map[int]*mpio.PacketReader +} + +func (m *MServer) MakeServerFdContext(ck base.CommandKey) *serverFdContext { + rtn := &serverFdContext{ + M: m, + Lock: &sync.Mutex{}, + Sender: m.Sender, + CK: ck, + Readers: make(map[int]*mpio.PacketReader), + } + return rtn +} + +func (c *serverFdContext) processDataPacket(pk *packet.DataPacketType) { + c.Lock.Lock() + reader := c.Readers[pk.FdNum] + c.Lock.Unlock() + if reader == nil { + ackPacket := packet.MakeDataAckPacket() + ackPacket.CK = c.CK + ackPacket.FdNum = pk.FdNum + ackPacket.Error = "write to closed file (no fd)" + c.M.Sender.SendPacket(ackPacket) + return + } + reader.AddData(pk) + return +} + +func (c *serverFdContext) GetWriter(fdNum int) io.WriteCloser { + return mpio.MakePacketWriter(fdNum, c.Sender, c.CK) +} + +func (c *serverFdContext) GetReader(fdNum int) io.ReadCloser { + c.Lock.Lock() + defer c.Lock.Unlock() + reader := mpio.MakePacketReader(fdNum) + c.Readers[fdNum] = reader + return reader +} + +func (m *MServer) runCommand(runPacket *packet.RunPacketType) { + fdContext := m.MakeServerFdContext(runPacket.CK) + m.Lock.Lock() + m.FdContext = fdContext + m.Lock.Unlock() + go func() { + donePk, err := shexec.RunClientSSHCommandAndWait(runPacket, fdContext, shexec.SSHOpts{}, true) + fmt.Printf("done: err:%v, %v\n", err, donePk) + }() +} + func RunServer() (int, error) { server := &MServer{ Lock: &sync.Mutex{}, } + packet.GlobalDebug = true server.MainInput = packet.MakePacketParser(os.Stdin) server.Sender = packet.MakePacketSender(os.Stdout) defer server.Close() + defer fmt.Printf("runserver done\n") initPacket := packet.MakeInitPacket() initPacket.Version = base.MShellVersion server.Sender.SendPacket(initPacket) @@ -43,7 +107,12 @@ func RunServer() (int, error) { } if pk.GetType() == packet.RunPacketStr { runPacket := pk.(*packet.RunPacketType) - fmt.Printf("RUN> %s\n", runPacket) + server.runCommand(runPacket) + continue + } + if pk.GetType() == packet.DataPacketStr { + dataPacket := pk.(*packet.DataPacketType) + server.FdContext.processDataPacket(dataPacket) continue } server.Sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", packet.AsExtType(pk))) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index f0334e3f..4bc6dadc 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -64,6 +64,41 @@ type ShExecType struct { Multiplexer *mpio.Multiplexer } +type StdContext struct{} + +func (StdContext) GetWriter(fdNum int) io.WriteCloser { + if fdNum == 0 { + return os.Stdin + } + if fdNum == 1 { + return os.Stdout + } + if fdNum == 2 { + return os.Stderr + } + fd := os.NewFile(uintptr(fdNum), fmt.Sprintf("/dev/fd/%d", fdNum)) + return fd +} + +func (StdContext) GetReader(fdNum int) io.ReadCloser { + if fdNum == 0 { + return os.Stdin + } + if fdNum == 1 { + return os.Stdout + } + if fdNum == 2 { + return os.Stdout + } + fd := os.NewFile(uintptr(fdNum), fmt.Sprintf("/dev/fd/%d", fdNum)) + return fd +} + +type FdContext interface { + GetWriter(fdNum int) io.WriteCloser + GetReader(fdNum int) io.ReadCloser +} + func MakeShExec(ck base.CommandKey) *ShExecType { return &ShExecType{ Lock: &sync.Mutex{}, @@ -237,11 +272,10 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } type SSHOpts struct { - SSHHost string - SSHOptsStr string - SSHIdentity string - SSHUser string - SudoWithPass bool + SSHHost string + SSHOptsStr string + SSHIdentity string + SSHUser string } type InstallOpts struct { @@ -258,6 +292,7 @@ type ClientOpts struct { Cwd string Debug bool Sudo bool + SudoWithPass bool SudoPw string CommandStdinFdNum int Detach bool @@ -316,7 +351,7 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket.Command = fmt.Sprintf(RunCommandFmt, opts.Command) return runPacket, nil } - if opts.SSHOpts.SudoWithPass { + if opts.SudoWithPass { pwFdNum, err := opts.NextFreeFdNum() if err != nil { return nil, err @@ -480,7 +515,16 @@ func RunInstallSSHCommand(opts *InstallOpts) error { return fmt.Errorf("did not receive version string from client, install not successful") } -func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, sshOpts SSHOpts, debug bool) (*packet.CmdDonePacketType, error) { +func HasDupStdin(fds []packet.RemoteFd) bool { + for _, rfd := range fds { + if rfd.Read && rfd.DupStdin { + return true + } + } + return false +} + +func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdContext, sshOpts SSHOpts, debug bool) (*packet.CmdDonePacketType, error) { cmd := MakeShExec("") ecmd := sshOpts.MakeSSHExecCmd(ClientCommand) cmd.Cmd = ecmd @@ -496,11 +540,11 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, sshOpts SSHOpts if err != nil { return nil, fmt.Errorf("creating stderr pipe: %v", err) } - if !sshOpts.SudoWithPass { - cmd.Multiplexer.MakeRawFdReader(0, os.Stdin, false) + if !HasDupStdin(runPacket.Fds) { + cmd.Multiplexer.MakeRawFdReader(0, fdContext.GetReader(0), false) } - cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout, false) - cmd.Multiplexer.MakeRawFdWriter(2, os.Stderr, false) + cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false) + cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), false) for _, rfd := range runPacket.Fds { if rfd.Read && rfd.Content != "" { err = cmd.Multiplexer.MakeStringFdReader(rfd.FdNum, rfd.Content) @@ -510,16 +554,14 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, sshOpts SSHOpts continue } if rfd.Read && rfd.DupStdin { - cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, os.Stdin, false) + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fdContext.GetReader(0), false) continue } - fd := os.NewFile(uintptr(rfd.FdNum), fmt.Sprintf("/dev/fd/%d", rfd.FdNum)) - if fd == nil { - return nil, fmt.Errorf("cannot open fd %d", rfd.FdNum) - } if rfd.Read { - cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, true) + fd := fdContext.GetReader(rfd.FdNum) + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, false) } else if rfd.Write { + fd := fdContext.GetWriter(rfd.FdNum) cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true) } } From 9054c3cdccce33969820a67c84814fe3ff2f6782 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 28 Jun 2022 19:01:33 -0700 Subject: [PATCH 039/149] got basic mshell --server functionality working to dispatch multiple commands --- main-mshell.go | 25 +++++------ pkg/mpio/mpio.go | 19 ++++----- pkg/packet/packet.go | 44 ++++++++++---------- pkg/server/server.go | 98 +++++++++++++++++++++++++++++++++----------- pkg/shexec/shexec.go | 16 ++++---- 5 files changed, 119 insertions(+), 83 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index b8b2b85f..04257aa6 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -80,34 +80,34 @@ func doSingle(ck base.CommandKey) { func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { err := shexec.ValidateRunPacket(pk) if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("invalid run packet: %v", err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("invalid run packet: %v", err)) return } fileNames, err := base.GetCommandFileNames(pk.CK) if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot get command file names: %v", err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot get command file names: %v", err)) return } cmd, err := shexec.MakeRunnerExec(pk.CK) if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot make mshell command: %v", err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot make mshell command: %v", err)) return } cmdStdin, err := cmd.StdinPipe() if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot pipe stdin to command: %v", err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot pipe stdin to command: %v", err)) return } // touch ptyout file (should exist for tailer to work correctly) ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err)) return } ptyOutFd.Close() // just opened to create the file, can close right after runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err)) return } defer runnerOutFd.Close() @@ -115,13 +115,13 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { cmd.Stderr = runnerOutFd err = cmd.Start() if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("error starting command: %v", err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error starting command: %v", err)) return } go func() { err = packet.SendPacket(cmdStdin, pk) if err != nil { - sender.SendPacket(packet.MakeCKErrorPacket(pk.CK, fmt.Sprintf("error sending forked runner command: %v", err))) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error sending forked runner command: %v", err)) return } cmdStdin.Close() @@ -237,11 +237,6 @@ func handleSingle() { runPacket, _ = pk.(*packet.RunPacketType) break } - if pk.GetType() == packet.RawPacketStr { - rawPk := pk.(*packet.RawPacketType) - sender.SendMessage("got raw packet '%s'", rawPk.Data) - continue - } sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) return } @@ -251,7 +246,7 @@ func handleSingle() { } cmd, err := shexec.RunCommand(runPacket, sender) if err != nil { - sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) + sender.SendCKErrorPacket(runPacket.CK, fmt.Sprintf("error running command: %v", err)) return } defer cmd.Close() @@ -433,7 +428,7 @@ func handleClient() (int, error) { if err != nil { return 1, err } - donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, shexec.StdContext{}, opts.SSHOpts, opts.Debug) + donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, shexec.StdContext{}, opts.SSHOpts, nil, opts.Debug) if err != nil { return 1, err } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index ca0a0ffa..9528c47a 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -31,16 +31,21 @@ type Multiplexer struct { Sender *packet.PacketSender Input *packet.PacketParser Started bool + UPR packet.UnknownPacketReporter Debug bool } -func MakeMultiplexer(ck base.CommandKey) *Multiplexer { +func MakeMultiplexer(ck base.CommandKey, upr packet.UnknownPacketReporter) *Multiplexer { + if upr == nil { + upr = packet.DefaultUPR{} + } return &Multiplexer{ Lock: &sync.Mutex{}, CK: ck, FdReaders: make(map[int]*FdReader), FdWriters: make(map[int]*FdWriter), + UPR: upr, } } @@ -207,17 +212,7 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { donePacket := pk.(*packet.CmdDonePacketType) return donePacket } - if pk.GetType() == packet.ErrorPacketStr { - errPacket := pk.(*packet.ErrorPacketType) - // at this point, just send the error packet to stderr rather than try to do something special - fmt.Fprintf(os.Stderr, "%s\n", errPacket.Error) - return nil - } - if pk.GetType() == packet.RawPacketStr { - rawPacket := pk.(*packet.RawPacketType) - fmt.Fprintf(os.Stderr, "%s\n", rawPacket.Data) - continue - } + m.UPR.UnknownPacket(pk) } return nil } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 07f58c19..44246347 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -11,6 +11,7 @@ import ( "encoding/json" "fmt" "io" + "os" "reflect" "sync" @@ -609,33 +610,30 @@ func (sender *PacketSender) SendErrorPacket(errVal string) error { return sender.SendPacket(MakeErrorPacket(errVal)) } +func (sender *PacketSender) SendCKErrorPacket(ck base.CommandKey, errVal string) error { + return sender.SendPacket(MakeCKErrorPacket(ck, errVal)) +} + func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) error { return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) } -type ErrorReporter interface { - ReportError(err error) +type UnknownPacketReporter interface { + UnknownPacket(pk PacketType) } -func PacketToByteArrBridge(pkCh chan PacketType, byteCh chan []byte, errorReporter ErrorReporter, closeOnDone bool) { - go func() { - defer func() { - if closeOnDone { - close(byteCh) - } - }() - for pk := range pkCh { - if pk == nil { - continue - } - jsonBytes, err := json.Marshal(pk) - if err != nil { - if errorReporter != nil { - errorReporter.ReportError(fmt.Errorf("error marshaling packet: %w", err)) - } - continue - } - byteCh <- jsonBytes - } - }() +type DefaultUPR struct{} + +func (DefaultUPR) UnknownPacket(pk PacketType) { + if pk.GetType() == ErrorPacketStr { + errPacket := pk.(*ErrorPacketType) + // at this point, just send the error packet to stderr rather than try to do something special + fmt.Fprintf(os.Stderr, "[error] %s\n", errPacket.Error) + } else if pk.GetType() == RawPacketStr { + rawPacket := pk.(*RawPacketType) + fmt.Fprintf(os.Stderr, "%s\n", rawPacket.Data) + } else { + fmt.Fprintf(os.Stderr, "[error] invalid packet received '%s'", AsExtType(pk)) + } + } diff --git a/pkg/server/server.go b/pkg/server/server.go index 0400d132..9406fc66 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -19,10 +19,11 @@ import ( ) type MServer struct { - Lock *sync.Mutex - MainInput *packet.PacketParser - Sender *packet.PacketSender - FdContext *serverFdContext + Lock *sync.Mutex + MainInput *packet.PacketParser + Sender *packet.PacketSender + FdContextMap map[base.CommandKey]*serverFdContext + Debug bool } func (m *MServer) Close() { @@ -38,17 +39,6 @@ type serverFdContext struct { Readers map[int]*mpio.PacketReader } -func (m *MServer) MakeServerFdContext(ck base.CommandKey) *serverFdContext { - rtn := &serverFdContext{ - M: m, - Lock: &sync.Mutex{}, - Sender: m.Sender, - CK: ck, - Readers: make(map[int]*mpio.PacketReader), - } - return rtn -} - func (c *serverFdContext) processDataPacket(pk *packet.DataPacketType) { c.Lock.Lock() reader := c.Readers[pk.FdNum] @@ -62,7 +52,43 @@ func (c *serverFdContext) processDataPacket(pk *packet.DataPacketType) { return } reader.AddData(pk) - return +} + +func (m *MServer) MakeServerFdContext(ck base.CommandKey) *serverFdContext { + rtn := &serverFdContext{ + M: m, + Lock: &sync.Mutex{}, + Sender: m.Sender, + CK: ck, + Readers: make(map[int]*mpio.PacketReader), + } + return rtn +} + +func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { + ck := pk.GetCK() + if ck == "" { + m.Sender.SendErrorPacket(fmt.Sprintf("received '%s' packet without ck", pk.GetType())) + return + } + m.Lock.Lock() + fdContext := m.FdContextMap[ck] + m.Lock.Unlock() + if fdContext == nil { + m.Sender.SendCKErrorPacket(ck, fmt.Sprintf("no server context for ck '%s'", ck)) + return + } + if pk.GetType() == packet.DataPacketStr { + dataPacket := pk.(*packet.DataPacketType) + fdContext.processDataPacket(dataPacket) + return + } else if pk.GetType() == packet.DataAckPacketStr { + m.Sender.SendPacket(pk) + return + } else { + m.Sender.SendCKErrorPacket(ck, fmt.Sprintf("invalid packet '%s' received", packet.AsExtType(pk))) + return + } } func (c *serverFdContext) GetWriter(fdNum int) io.WriteCloser { @@ -78,21 +104,42 @@ func (c *serverFdContext) GetReader(fdNum int) io.ReadCloser { } func (m *MServer) runCommand(runPacket *packet.RunPacketType) { + if err := runPacket.CK.Validate("packet"); err != nil { + m.Sender.SendErrorPacket(fmt.Sprintf("server run packets require valid ck: %s", err)) + return + } fdContext := m.MakeServerFdContext(runPacket.CK) m.Lock.Lock() - m.FdContext = fdContext + m.FdContextMap[runPacket.CK] = fdContext m.Lock.Unlock() go func() { - donePk, err := shexec.RunClientSSHCommandAndWait(runPacket, fdContext, shexec.SSHOpts{}, true) - fmt.Printf("done: err:%v, %v\n", err, donePk) + donePk, err := shexec.RunClientSSHCommandAndWait(runPacket, fdContext, shexec.SSHOpts{}, m, m.Debug) + if donePk != nil { + m.Sender.SendPacket(donePk) + } + if err != nil { + m.Sender.SendCKErrorPacket(runPacket.CK, err.Error()) + } }() } +func (m *MServer) UnknownPacket(pk packet.PacketType) { + m.Sender.SendPacket(pk) +} + func RunServer() (int, error) { + debug := false + if len(os.Args) >= 3 && os.Args[2] == "--debug" { + debug = true + } server := &MServer{ - Lock: &sync.Mutex{}, + Lock: &sync.Mutex{}, + FdContextMap: make(map[base.CommandKey]*serverFdContext), + Debug: debug, + } + if debug { + packet.GlobalDebug = true } - packet.GlobalDebug = true server.MainInput = packet.MakePacketParser(os.Stdin) server.Sender = packet.MakePacketSender(os.Stdout) defer server.Close() @@ -101,7 +148,9 @@ func RunServer() (int, error) { initPacket.Version = base.MShellVersion server.Sender.SendPacket(initPacket) for pk := range server.MainInput.MainCh { - fmt.Printf("PK> %s\n", packet.AsString(pk)) + if server.Debug { + fmt.Printf("PK> %s\n", packet.AsString(pk)) + } if pk.GetType() == packet.PingPacketStr { continue } @@ -110,9 +159,8 @@ func RunServer() (int, error) { server.runCommand(runPacket) continue } - if pk.GetType() == packet.DataPacketStr { - dataPacket := pk.(*packet.DataPacketType) - server.FdContext.processDataPacket(dataPacket) + if cmdPk, ok := pk.(packet.CommandPacketType); ok { + server.ProcessCommandPacket(cmdPk) continue } server.Sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", packet.AsExtType(pk))) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 4bc6dadc..e40c5fd2 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -99,12 +99,12 @@ type FdContext interface { GetReader(fdNum int) io.ReadCloser } -func MakeShExec(ck base.CommandKey) *ShExecType { +func MakeShExec(ck base.CommandKey, upr packet.UnknownPacketReporter) *ShExecType { return &ShExecType{ Lock: &sync.Mutex{}, StartTs: time.Now(), CK: ck, - Multiplexer: mpio.MakeMultiplexer(ck), + Multiplexer: mpio.MakeMultiplexer(ck, upr), } } @@ -524,8 +524,8 @@ func HasDupStdin(fds []packet.RemoteFd) bool { return false } -func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdContext, sshOpts SSHOpts, debug bool) (*packet.CmdDonePacketType, error) { - cmd := MakeShExec("") +func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdContext, sshOpts SSHOpts, upr packet.UnknownPacketReporter, debug bool) (*packet.CmdDonePacketType, error) { + cmd := MakeShExec(runPacket.CK, upr) ecmd := sshOpts.MakeSSHExecCmd(ClientCommand) cmd.Cmd = ecmd inputWriter, err := ecmd.StdinPipe() @@ -656,7 +656,7 @@ func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sen } func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - cmd := MakeShExec(pk.CK) + cmd := MakeShExec(pk.CK, nil) cmd.Cmd = exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(cmd.Cmd, pk.Env) if pk.Cwd != "" { @@ -736,7 +736,7 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( defer func() { cmdTty.Close() }() - rtn := MakeShExec(pk.CK) + rtn := MakeShExec(pk.CK, nil) ecmd := MakeExecCmd(pk, cmdTty) err = ecmd.Start() if err != nil { @@ -750,14 +750,14 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( // copy pty output to .ptyout file _, copyErr := io.Copy(ptyOutFd, cmdPty) if copyErr != nil { - sender.SendErrorPacket(fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) } }() go func() { // copy .stdin fifo contents to pty input copyFifoErr := MakeAndCopyStdinFifo(cmdPty, fileNames.StdinFifo) if copyFifoErr != nil { - sender.SendErrorPacket(fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) + sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) } }() rtn.FileNames = fileNames From c73691ac24eb6bcf4fec3e725bb24ccf2e847408 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 28 Jun 2022 21:57:30 -0700 Subject: [PATCH 040/149] move static files from remotefd content to 'rundata'. send all rundata before command start. parse rundata before command start. compatible with detached commands --- main-mshell.go | 39 ++++++------ pkg/mpio/mpio.go | 28 ++++++--- pkg/packet/packet.go | 80 ++++++++++++++++++++++-- pkg/packet/parser.go | 3 + pkg/server/server.go | 3 - pkg/shexec/shexec.go | 141 +++++++++++++++++++++++++++++-------------- 6 files changed, 214 insertions(+), 80 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 04257aa6..cfb9fca9 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -42,9 +42,6 @@ func doSingle(ck base.CommandKey) { sender := packet.MakePacketSender(os.Stdout) var runPacket *packet.RunPacketType for pk := range packetParser.MainCh { - if pk.GetType() == packet.PingPacketStr { - continue - } if pk.GetType() == packet.RunPacketStr { runPacket, _ = pk.(*packet.RunPacketType) break @@ -173,9 +170,6 @@ func doMain() { } sender.SendPacket(initPacket) for pk := range packetParser.MainCh { - if pk.GetType() == packet.PingPacketStr { - continue - } if pk.GetType() == packet.RunPacketStr { doMainRun(pk.(*packet.RunPacketType), sender) continue @@ -211,6 +205,20 @@ func doMain() { } } +func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType, error) { + rpb := packet.MakeRunPacketBuilder() + for pk := range packetParser.MainCh { + ok, runPacket := rpb.ProcessPacket(pk) + if runPacket != nil { + return runPacket, nil + } + if !ok { + return nil, fmt.Errorf("invalid packet '%s' sent to mshell", pk.GetType()) + } + } + return nil, fmt.Errorf("no run packet received") +} + func handleSingle() { packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) @@ -228,20 +236,13 @@ func handleSingle() { initPacket := packet.MakeInitPacket() initPacket.Version = base.MShellVersion sender.SendPacket(initPacket) - var runPacket *packet.RunPacketType - for pk := range packetParser.MainCh { - if pk.GetType() == packet.PingPacketStr { - continue + runPacket, err := readFullRunPacket(packetParser) + if err != nil { + ck := base.CommandKey("") + if runPacket != nil { + ck = runPacket.CK } - if pk.GetType() == packet.RunPacketStr { - runPacket, _ = pk.(*packet.RunPacketType) - break - } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) - return - } - if runPacket == nil { - sender.SendErrorPacket(fmt.Sprintf("no run packet received")) + sender.SendCKErrorPacket(ck, err.Error()) return } cmd, err := shexec.RunCommand(runPacket, sender) diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 9528c47a..6cff436a 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -20,12 +20,14 @@ import ( const ReadBufSize = 128 * 1024 const WriteBufSize = 128 * 1024 const MaxSingleWriteSize = 4 * 1024 +const MaxTotalRunDataSize = 10 * ReadBufSize type Multiplexer struct { Lock *sync.Mutex CK base.CommandKey FdReaders map[int]*FdReader // synchronized FdWriters map[int]*FdWriter // synchronized + RunData map[int]*FdReader // synchronized CloseAfterStart []*os.File // synchronized Sender *packet.PacketSender @@ -105,16 +107,22 @@ func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { return pr, nil } -func (m *Multiplexer) MakeStringFdReader(fdNum int, contents string) error { - pw, err := m.MakeReaderPipe(fdNum) +// returns the *reader* to connect to process, writer is put in FdWriters +func (m *Multiplexer) MakeStaticWriterPipe(fdNum int, data []byte) (*os.File, error) { + pr, pw, err := os.Pipe() if err != nil { - return err + return nil, err } - go func() { - pw.Write([]byte(contents)) - pw.Close() - }() - return nil + m.Lock.Lock() + defer m.Lock.Unlock() + fdWriter := MakeFdWriter(m, pw, fdNum, true) + err = fdWriter.AddData(data, true) + if err != nil { + return nil, err + } + m.FdWriters[fdNum] = fdWriter + m.CloseAfterStart = append(m.CloseAfterStart, pr) + return pr, nil } func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose bool) { @@ -212,6 +220,10 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { donePacket := pk.(*packet.CmdDonePacketType) return donePacket } + if pk.GetType() == packet.CmdStartPacketStr { + // nothing + continue + } m.UPR.UnknownPacket(pk) } return nil diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 44246347..074b7175 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -8,6 +8,7 @@ package packet import ( "bytes" + "encoding/base64" "encoding/json" "fmt" "io" @@ -33,6 +34,7 @@ const ( DataAckPacketStr = "dataack" CmdStartPacketStr = "cmdstart" CmdDonePacketStr = "cmddone" + DataEndPacketStr = "dataend" ResponsePacketStr = "resp" DonePacketStr = "done" ErrorPacketStr = "error" @@ -68,6 +70,7 @@ func init() { TypeStrToFactory[InputPacketStr] = reflect.TypeOf(InputPacketType{}) TypeStrToFactory[DataPacketStr] = reflect.TypeOf(DataPacketType{}) TypeStrToFactory[DataAckPacketStr] = reflect.TypeOf(DataAckPacketType{}) + TypeStrToFactory[DataEndPacketStr] = reflect.TypeOf(DataEndPacketType{}) } func MakePacket(packetType string) (PacketType, error) { @@ -166,6 +169,23 @@ func MakeDataPacket() *DataPacketType { return &DataPacketType{Type: DataPacketStr} } +type DataEndPacketType struct { + Type string `json:"type"` + CK base.CommandKey `json:"ck"` +} + +func MakeDataEndPacket(ck base.CommandKey) *DataEndPacketType { + return &DataEndPacketType{Type: DataEndPacketStr, CK: ck} +} + +func (*DataEndPacketType) GetType() string { + return DataEndPacketStr +} + +func (p *DataEndPacketType) GetCK() base.CommandKey { + return p.CK +} + type DataAckPacketType struct { Type string `json:"type"` CK base.CommandKey `json:"ck"` @@ -411,11 +431,16 @@ type TermSize struct { } type RemoteFd struct { - FdNum int `json:"fdnum"` - Read bool `json:"read"` - Write bool `json:"write"` - Content string `json:"-"` - DupStdin bool `json:"-"` + FdNum int `json:"fdnum"` + Read bool `json:"read"` + Write bool `json:"write"` + DupStdin bool `json:"-"` +} + +type RunDataType struct { + FdNum int `json:"fdnum"` + DataLen int `json:"datalen"` + Data []byte `json:"-"` } type RunPacketType struct { @@ -426,6 +451,7 @@ type RunPacketType struct { Env map[string]string `json:"env,omitempty"` TermSize *TermSize `json:"termsize,omitempty"` Fds []RemoteFd `json:"fds,omitempty"` + RunData []RunDataType `json:"rundata,omitempty"` Detached bool `json:"detached,omitempty"` } @@ -637,3 +663,47 @@ func (DefaultUPR) UnknownPacket(pk PacketType) { } } + +// todo: clean hanging entries in RunMap when in server mode +type RunPacketBuilder struct { + RunMap map[base.CommandKey]*RunPacketType +} + +func MakeRunPacketBuilder() *RunPacketBuilder { + return &RunPacketBuilder{ + RunMap: make(map[base.CommandKey]*RunPacketType), + } +} + +// returns (consumed, fullRunPacket) +func (b *RunPacketBuilder) ProcessPacket(pk PacketType) (bool, *RunPacketType) { + if pk.GetType() == RunPacketStr { + runPacket := pk.(*RunPacketType) + b.RunMap[runPacket.CK] = runPacket + return true, nil + } + if pk.GetType() == DataEndPacketStr { + endPacket := pk.(*DataEndPacketType) + runPacket := b.RunMap[endPacket.CK] // might be nil + delete(b.RunMap, endPacket.CK) + return true, runPacket + } + if pk.GetType() == DataPacketStr { + dataPacket := pk.(*DataPacketType) + runPacket := b.RunMap[dataPacket.CK] + if runPacket == nil { + return false, nil + } + for idx, runData := range runPacket.RunData { + if runData.FdNum == dataPacket.FdNum { + // can ignore error, will get caught later with RunData.DataLen check + realData, _ := base64.StdEncoding.DecodeString(dataPacket.Data64) + runData.Data = append(runData.Data, realData...) + runPacket.RunData[idx] = runData + break + } + } + return true, nil + } + return false, nil +} diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index 40e59b24..e8b2a96d 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -93,6 +93,9 @@ func MakePacketParser(input io.Reader) *PacketParser { if pk.GetType() == DonePacketStr { return } + if pk.GetType() == PingPacketStr { + continue + } parser.MainCh <- pk } }() diff --git a/pkg/server/server.go b/pkg/server/server.go index 9406fc66..b7996b05 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -151,9 +151,6 @@ func RunServer() (int, error) { if server.Debug { fmt.Printf("PK> %s\n", packet.AsString(pk)) } - if pk.GetType() == packet.PingPacketStr { - continue - } if pk.GetType() == packet.RunPacketStr { runPacket := pk.(*packet.RunPacketType) server.runCommand(runPacket) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index e40c5fd2..144f52a4 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -7,6 +7,7 @@ package shexec import ( + "encoding/base64" "fmt" "io" "os" @@ -223,15 +224,20 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { if rfd.Write { return fmt.Errorf("cannot detach command with writable remote files fd=%d", rfd.FdNum) } - if rfd.Read { - if rfd.Content == "" { - return fmt.Errorf("cannot detach command with readable remote files fd=%d", rfd.FdNum) - } - if len(rfd.Content) > mpio.ReadBufSize { - return fmt.Errorf("cannot detach command, constant readable input too large fd=%d, len=%d, max=%d", rfd.FdNum, len(rfd.Content), mpio.ReadBufSize) - } + if rfd.Read && !rfd.DupStdin { + return fmt.Errorf("cannot detach command with readable remote files fd=%d", rfd.FdNum) } } + totalRunData := 0 + for _, rd := range pk.RunData { + if rd.DataLen > mpio.ReadBufSize { + return fmt.Errorf("cannot detach command, constant rundata input too large fd=%d, len=%d, max=%d", rd.FdNum, rd.DataLen, mpio.ReadBufSize) + } + totalRunData += rd.DataLen + } + if totalRunData > mpio.MaxTotalRunDataSize { + return fmt.Errorf("cannot detach command, constant rundata input too large len=%d, max=%d", totalRunData, mpio.MaxTotalRunDataSize) + } } if pk.Cwd != "" { realCwd := base.ExpandHomeDir(pk.Cwd) @@ -243,6 +249,11 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { return fmt.Errorf("invalid cwd '%s' for command, not a directory", realCwd) } } + for _, runData := range pk.RunData { + if runData.DataLen != len(runData.Data) { + return fmt.Errorf("rundata length mismatch, fd=%d, datalen=%d, expected=%d", runData.FdNum, len(runData.Data), runData.DataLen) + } + } return nil } @@ -286,16 +297,15 @@ type InstallOpts struct { } type ClientOpts struct { - SSHOpts SSHOpts - Command string - Fds []packet.RemoteFd - Cwd string - Debug bool - Sudo bool - SudoWithPass bool - SudoPw string - CommandStdinFdNum int - Detach bool + SSHOpts SSHOpts + Command string + Fds []packet.RemoteFd + Cwd string + Debug bool + Sudo bool + SudoWithPass bool + SudoPw string + Detach bool } func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { @@ -352,48 +362,55 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { return runPacket, nil } if opts.SudoWithPass { - pwFdNum, err := opts.NextFreeFdNum() + pwFdNum, err := AddRunData(runPacket, opts.SudoPw, "sudo pw") if err != nil { return nil, err } - pwRfd := packet.RemoteFd{FdNum: pwFdNum, Read: true, Content: opts.SudoPw} - opts.Fds = append(opts.Fds, pwRfd) - commandFdNum, err := opts.NextFreeFdNum() + commandFdNum, err := AddRunData(runPacket, opts.Command, "command") if err != nil { return nil, err } - commandRfd := packet.RemoteFd{FdNum: commandFdNum, Read: true, Content: opts.Command} - opts.Fds = append(opts.Fds, commandRfd) - commandStdinFdNum, err := opts.NextFreeFdNum() + commandStdinFdNum, err := NextFreeFdNum(runPacket) if err != nil { return nil, err } commandStdinRfd := packet.RemoteFd{FdNum: commandStdinFdNum, Read: true, DupStdin: true} - opts.Fds = append(opts.Fds, commandStdinRfd) - opts.CommandStdinFdNum = commandStdinFdNum - maxFdNum := opts.MaxFdNum() + runPacket.Fds = append(runPacket.Fds, commandStdinRfd) + maxFdNum := MaxFdNumInPacket(runPacket) runPacket.Command = fmt.Sprintf(RunSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, pwFdNum, commandFdNum, commandStdinFdNum) - runPacket.Fds = opts.Fds return runPacket, nil } else { - commandFdNum, err := opts.NextFreeFdNum() + commandFdNum, err := AddRunData(runPacket, opts.Command, "command") if err != nil { return nil, err } - rfd := packet.RemoteFd{FdNum: commandFdNum, Read: true, Content: opts.Command} - opts.Fds = append(opts.Fds, rfd) - maxFdNum := opts.MaxFdNum() + maxFdNum := MaxFdNumInPacket(runPacket) runPacket.Command = fmt.Sprintf(RunSudoCommandFmt, maxFdNum+1, commandFdNum) - runPacket.Fds = opts.Fds return runPacket, nil } } -func (opts *ClientOpts) NextFreeFdNum() (int, error) { +func AddRunData(pk *packet.RunPacketType, data string, dataType string) (int, error) { + if len(data) > mpio.ReadBufSize { + return 0, fmt.Errorf("%s too large, exceeds read buffer size", dataType) + } + fdNum, err := NextFreeFdNum(pk) + if err != nil { + return 0, err + } + runData := packet.RunDataType{FdNum: fdNum, DataLen: len(data), Data: []byte(data)} + pk.RunData = append(pk.RunData, runData) + return fdNum, nil +} + +func NextFreeFdNum(pk *packet.RunPacketType) (int, error) { fdMap := make(map[int]bool) - for _, fd := range opts.Fds { + for _, fd := range pk.Fds { fdMap[fd.FdNum] = true } + for _, rd := range pk.RunData { + fdMap[rd.FdNum] = true + } for i := 3; i <= MaxFdNum; i++ { if !fdMap[i] { return i, nil @@ -402,13 +419,18 @@ func (opts *ClientOpts) NextFreeFdNum() (int, error) { return 0, fmt.Errorf("reached maximum number of fds, all fds between 3-%d are in use", MaxFdNum) } -func (opts *ClientOpts) MaxFdNum() int { +func MaxFdNumInPacket(pk *packet.RunPacketType) int { maxFdNum := 3 - for _, fd := range opts.Fds { + for _, fd := range pk.Fds { if fd.FdNum > maxFdNum { maxFdNum = fd.FdNum } } + for _, rd := range pk.RunData { + if rd.FdNum > maxFdNum { + maxFdNum = rd.FdNum + } + } return maxFdNum } @@ -546,13 +568,6 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false) cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), false) for _, rfd := range runPacket.Fds { - if rfd.Read && rfd.Content != "" { - err = cmd.Multiplexer.MakeStringFdReader(rfd.FdNum, rfd.Content) - if err != nil { - return nil, fmt.Errorf("creating content fd %d", rfd.FdNum) - } - continue - } if rfd.Read && rfd.DupStdin { cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fdContext.GetReader(0), false) continue @@ -610,7 +625,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon if !versionOk { return nil, fmt.Errorf("did not receive version from remote mshell") } - sender.SendPacket(runPacket) + SendRunPacketAndRunData(sender, runPacket) if debug { cmd.Multiplexer.Debug = true } @@ -622,6 +637,32 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon return donePacket, nil } +func min(v1 int, v2 int) int { + if v1 <= v2 { + return v1 + } + return v2 +} + +func SendRunPacketAndRunData(sender *packet.PacketSender, runPacket *packet.RunPacketType) { + sender.SendPacket(runPacket) + for _, runData := range runPacket.RunData { + sendBuf := runData.Data + for len(sendBuf) > 0 { + chunkSize := min(len(sendBuf), mpio.MaxSingleWriteSize) + chunk := sendBuf[0:chunkSize] + dataPk := packet.MakeDataPacket() + dataPk.CK = runPacket.CK + dataPk.FdNum = runData.FdNum + dataPk.Data64 = base64.StdEncoding.EncodeToString(chunk) + dataPk.Eof = (len(chunk) == len(sendBuf)) + sendBuf = sendBuf[chunkSize:] + sender.SendPacket(dataPk) + } + } + sender.SendPacket(packet.MakeDataEndPacket(runPacket.CK)) +} + func DetectGoArch(uname string) (string, string, error) { fields := strings.SplitN(uname, "|", 2) if len(fields) != 2 { @@ -683,6 +724,16 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S return nil, err } extraFiles := make([]*os.File, 0, MaxFdNum+1) + for _, runData := range pk.RunData { + if runData.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:runData.FdNum+1] + } + extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data) + if err != nil { + cmd.Close() + return nil, err + } + } for _, rfd := range pk.Fds { if rfd.FdNum >= len(extraFiles) { extraFiles = extraFiles[:rfd.FdNum+1] From 4d8841a4590248498ab709270813033c25e1540c Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 28 Jun 2022 22:05:47 -0700 Subject: [PATCH 041/149] use RunPacketBuilder in server mode --- pkg/packet/packet.go | 10 +++++++++- pkg/server/server.go | 12 ++++++++---- 2 files changed, 17 insertions(+), 5 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 074b7175..17a5638e 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -477,10 +477,18 @@ type ErrorPacketType struct { Error string `json:"error"` } -func (et *ErrorPacketType) GetType() string { +func (*ErrorPacketType) GetType() string { return ErrorPacketStr } +func (p *ErrorPacketType) String() string { + ckStr := "" + if p.CK != "" { + ckStr = fmt.Sprintf(", ck=%s", p.CK) + } + return fmt.Sprintf("error[%s%s]", p.Error, ckStr) +} + func MakeErrorPacket(errorStr string) *ErrorPacketType { return &ErrorPacketType{Type: ErrorPacketStr, Error: errorStr} } diff --git a/pkg/server/server.go b/pkg/server/server.go index b7996b05..cc0c84f3 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -147,20 +147,24 @@ func RunServer() (int, error) { initPacket := packet.MakeInitPacket() initPacket.Version = base.MShellVersion server.Sender.SendPacket(initPacket) + builder := packet.MakeRunPacketBuilder() for pk := range server.MainInput.MainCh { if server.Debug { fmt.Printf("PK> %s\n", packet.AsString(pk)) } - if pk.GetType() == packet.RunPacketStr { - runPacket := pk.(*packet.RunPacketType) - server.runCommand(runPacket) + ok, runPacket := builder.ProcessPacket(pk) + if ok { + if runPacket != nil { + server.runCommand(runPacket) + continue + } continue } if cmdPk, ok := pk.(packet.CommandPacketType); ok { server.ProcessCommandPacket(cmdPk) continue } - server.Sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", packet.AsExtType(pk))) + server.Sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", packet.AsString(pk))) continue } return 0, nil From b6711e7428405fdbe92736b9c8ffc565b635fabb Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 29 Jun 2022 14:29:38 -0700 Subject: [PATCH 042/149] sanitize packets to be 7-bit ascii without control chars. dont send data/dataend when no rundata present. use os.Executable to locate mshell if running locally. more work on detached mode --- main-mshell.go | 51 ++++++++++--------- pkg/packet/packet.go | 15 +++++- pkg/server/server.go | 14 +++-- pkg/shexec/shexec.go | 118 +++++++++++++++++++++++++++++++++++-------- 4 files changed, 148 insertions(+), 50 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index cfb9fca9..a7487660 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -10,10 +10,8 @@ import ( "bytes" "fmt" "os" - "os/signal" "os/user" "strings" - "syscall" "time" "github.com/scripthaus-dev/mshell/pkg/base" @@ -24,19 +22,6 @@ import ( "golang.org/x/sys/unix" ) -// in single run mode, we don't want mshell to die from signals -// since we want the single mshell to persist even if session / main mshell -// is terminated. -func setupSingleSignals(cmd *shexec.ShExecType) { - sigCh := make(chan os.Signal, 1) - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) - go func() { - for range sigCh { - // do nothing - } - }() -} - func doSingle(ck base.CommandKey) { packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) @@ -60,12 +45,12 @@ func doSingle(ck base.CommandKey) { sender.SendErrorPacket(fmt.Sprintf("run packet cmdid[%s] did not match arg[%s]", runPacket.CK, ck)) return } - cmd, err := shexec.RunCommand(runPacket, sender) + cmd, err := shexec.RunCommandDetached(runPacket, sender) if err != nil { sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) return } - setupSingleSignals(cmd) + shexec.SetupSignalsForDetach() startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) donePacket := cmd.WaitForCommand() @@ -245,15 +230,30 @@ func handleSingle() { sender.SendCKErrorPacket(ck, err.Error()) return } - cmd, err := shexec.RunCommand(runPacket, sender) + err = shexec.ValidateRunPacket(runPacket) if err != nil { - sender.SendCKErrorPacket(runPacket.CK, fmt.Sprintf("error running command: %v", err)) + sender.SendCKErrorPacket(runPacket.CK, err.Error()) + return + } + if runPacket.Detached { + cmd, err := shexec.RunCommandDetached(runPacket, sender) + if err != nil { + sender.SendCKErrorPacket(runPacket.CK, err.Error()) + return + } + cmd.WaitForCommand() + } else { + cmd, err := shexec.RunCommandSimple(runPacket, sender) + if err != nil { + sender.SendCKErrorPacket(runPacket.CK, fmt.Sprintf("error running command: %v", err)) + return + } + defer cmd.Close() + startPacket := cmd.MakeCmdStartPacket() + sender.SendPacket(startPacket) + cmd.RunRemoteIOAndWait(packetParser, sender) return } - defer cmd.Close() - startPacket := cmd.MakeCmdStartPacket() - sender.SendPacket(startPacket) - cmd.RunRemoteIOAndWait(packetParser, sender) } func detectOpenFds() ([]packet.RemoteFd, error) { @@ -494,7 +494,10 @@ Sudo Options: Sudo options allow you to run the given command using "sudo". The first option only works when you can sudo without a password. Your password will be passed -securely through a high numbered fd to "sudo -S". See full documentation for more details. +securely through a high numbered fd to "sudo -S". Note that to use high numbered +file descriptors with sudo, you will need to add this line to your /etc/sudoers file: + Defaults closefrom_override +See full documentation for more details. Examples: # execute a python script remotely, with stdin still hooked up correctly diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 17a5638e..f49db0c4 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -548,6 +548,14 @@ func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { return pk, nil } +func sanitizeBytes(buf []byte) { + for idx, b := range buf { + if b >= 127 || (b < 32 && b != 10 && b != 13) { + buf[idx] = '?' + } + } +} + func SendPacket(w io.Writer, packet PacketType) error { if packet == nil { return nil @@ -564,7 +572,9 @@ func SendPacket(w io.Writer, packet PacketType) error { if GlobalDebug { fmt.Printf("SEND> %s\n", AsString(packet)) } - _, err = w.Write(outBuf.Bytes()) + outBytes := outBuf.Bytes() + sanitizeBytes(outBytes) + _, err = w.Write(outBytes) if err != nil { return err } @@ -687,6 +697,9 @@ func MakeRunPacketBuilder() *RunPacketBuilder { func (b *RunPacketBuilder) ProcessPacket(pk PacketType) (bool, *RunPacketType) { if pk.GetType() == RunPacketStr { runPacket := pk.(*RunPacketType) + if len(runPacket.RunData) == 0 { + return true, runPacket + } b.RunMap[runPacket.CK] = runPacket return true, nil } diff --git a/pkg/server/server.go b/pkg/server/server.go index cc0c84f3..a04f5f98 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -55,6 +55,8 @@ func (c *serverFdContext) processDataPacket(pk *packet.DataPacketType) { } func (m *MServer) MakeServerFdContext(ck base.CommandKey) *serverFdContext { + m.Lock.Lock() + defer m.Lock.Unlock() rtn := &serverFdContext{ M: m, Lock: &sync.Mutex{}, @@ -62,6 +64,7 @@ func (m *MServer) MakeServerFdContext(ck base.CommandKey) *serverFdContext { CK: ck, Readers: make(map[int]*mpio.PacketReader), } + m.FdContextMap[ck] = rtn return rtn } @@ -103,16 +106,20 @@ func (c *serverFdContext) GetReader(fdNum int) io.ReadCloser { return reader } +func (m *MServer) RemoveFdContext(ck base.CommandKey) { + m.Lock.Lock() + defer m.Lock.Unlock() + delete(m.FdContextMap, ck) +} + func (m *MServer) runCommand(runPacket *packet.RunPacketType) { if err := runPacket.CK.Validate("packet"); err != nil { m.Sender.SendErrorPacket(fmt.Sprintf("server run packets require valid ck: %s", err)) return } fdContext := m.MakeServerFdContext(runPacket.CK) - m.Lock.Lock() - m.FdContextMap[runPacket.CK] = fdContext - m.Lock.Unlock() go func() { + defer m.RemoveFdContext(runPacket.CK) donePk, err := shexec.RunClientSSHCommandAndWait(runPacket, fdContext, shexec.SSHOpts{}, m, m.Debug) if donePk != nil { m.Sender.SendPacket(donePk) @@ -143,7 +150,6 @@ func RunServer() (int, error) { server.MainInput = packet.MakePacketParser(os.Stdin) server.Sender = packet.MakePacketSender(os.Stdout) defer server.Close() - defer fmt.Printf("runserver done\n") initPacket := packet.MakeInitPacket() initPacket.Version = base.MShellVersion server.Sender.SendPacket(initPacket) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 144f52a4..788de2cb 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -12,6 +12,7 @@ import ( "io" "os" "os/exec" + "os/signal" "strings" "sync" "syscall" @@ -164,20 +165,59 @@ func UpdateCmdEnv(cmd *exec.Cmd, envVars map[string]string) { cmd.Env = newEnv } -func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { +// returns (pr, err) +func MakeSimpleStaticWriterPipe(data []byte) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + go func() { + defer pw.Close() + pw.Write(data) + }() + return pr, err +} + +func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, error) { ecmd := exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(ecmd, pk.Env) if pk.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(pk.Cwd) } - ecmd.Stdin = cmdTty + if !HasDupStdin(pk.Fds) { + ecmd.Stdin = cmdTty + } ecmd.Stdout = cmdTty ecmd.Stderr = cmdTty ecmd.SysProcAttr = &syscall.SysProcAttr{ Setsid: true, Setctty: true, } - return ecmd + extraFiles := make([]*os.File, 0, MaxFdNum+1) + for _, rfd := range pk.Fds { + if rfd.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:rfd.FdNum+1] + } + if rfd.Read && rfd.DupStdin { + extraFiles[rfd.FdNum] = cmdTty + continue + } + return nil, fmt.Errorf("invalid fd %d passed to detached command", rfd.FdNum) + } + for _, runData := range pk.RunData { + if runData.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:runData.FdNum+1] + } + var err error + extraFiles[runData.FdNum], err = MakeSimpleStaticWriterPipe(runData.Data) + if err != nil { + return nil, err + } + } + if len(extraFiles) > FirstExtraFilesFdNum { + ecmd.ExtraFiles = extraFiles[FirstExtraFilesFdNum:] + } + return ecmd, nil } func MakeRunnerExec(ck base.CommandKey) (*exec.Cmd, error) { @@ -269,19 +309,6 @@ func GetWinsize(p *packet.RunPacketType) *pty.Winsize { return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} } -// when err is nil, the command will have already been started -func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - err := ValidateRunPacket(pk) - if err != nil { - return nil, err - } - if !pk.Detached { - return runCommandSimple(pk, sender) - } else { - return runCommandDetached(pk, sender) - } -} - type SSHOpts struct { SSHHost string SSHOptsStr string @@ -308,6 +335,25 @@ type ClientOpts struct { Detach bool } +func (opts SSHOpts) MakeSSHInstallCmd() (*exec.Cmd, error) { + if opts.SSHHost == "" { + return nil, fmt.Errorf("no ssh host provided, can only install to a remote host") + } + return opts.MakeSSHExecCmd(InstallCommand), nil +} + +func (opts SSHOpts) MakeMShellSingleCmd() (*exec.Cmd, error) { + if opts.SSHHost == "" { + execFile, err := os.Executable() + if err != nil { + return nil, fmt.Errorf("cannot find local mshell executable: %w", err) + } + ecmd := exec.Command(execFile, "--single") + return ecmd, nil + } + return opts.MakeSSHExecCmd(ClientCommand), nil +} + func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { remoteCommand = strings.TrimSpace(remoteCommand) if opts.SSHHost == "" { @@ -474,7 +520,10 @@ func sendOptFile(input io.WriteCloser, optName string) error { func RunInstallSSHCommand(opts *InstallOpts) error { tryDetect := opts.Detect - ecmd := opts.SSHOpts.MakeSSHExecCmd(InstallCommand) + ecmd, err := opts.SSHOpts.MakeSSHInstallCmd() + if err != nil { + return err + } inputWriter, err := ecmd.StdinPipe() if err != nil { return fmt.Errorf("creating stdin pipe: %v", err) @@ -548,7 +597,10 @@ func HasDupStdin(fds []packet.RemoteFd) bool { func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdContext, sshOpts SSHOpts, upr packet.UnknownPacketReporter, debug bool) (*packet.CmdDonePacketType, error) { cmd := MakeShExec(runPacket.CK, upr) - ecmd := sshOpts.MakeSSHExecCmd(ClientCommand) + ecmd, err := sshOpts.MakeMShellSingleCmd() + if err != nil { + return nil, err + } cmd.Cmd = ecmd inputWriter, err := ecmd.StdinPipe() if err != nil { @@ -646,6 +698,9 @@ func min(v1 int, v2 int) int { func SendRunPacketAndRunData(sender *packet.PacketSender, runPacket *packet.RunPacketType) { sender.SendPacket(runPacket) + if len(runPacket.RunData) == 0 { + return + } for _, runData := range runPacket.RunData { sendBuf := runData.Data for len(sendBuf) > 0 { @@ -696,7 +751,7 @@ func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sen sender.SendPacket(donePacket) } -func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { +func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { cmd := MakeShExec(pk.CK, nil) cmd.Cmd = exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(cmd.Cmd, pk.Env) @@ -767,7 +822,19 @@ func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S return cmd, nil } -func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { +// in detached run mode, we don't want mshell to die from signals +// since we want mshell to persist even if the mshell --server is terminated +func SetupSignalsForDetach() { + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) + go func() { + for range sigCh { + // do nothing + } + }() +} + +func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { fileNames, err := base.GetCommandFileNames(pk.CK) if err != nil { return nil, err @@ -788,11 +855,20 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( cmdTty.Close() }() rtn := MakeShExec(pk.CK, nil) - ecmd := MakeExecCmd(pk, cmdTty) + ecmd, err := MakeDetachedExecCmd(pk, cmdTty) + if err != nil { + return nil, err + } + SetupSignalsForDetach() err = ecmd.Start() if err != nil { return nil, fmt.Errorf("starting command: %w", err) } + for _, fd := range ecmd.ExtraFiles { + if fd != cmdTty { + fd.Close() + } + } ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY, 0600) if err != nil { return nil, fmt.Errorf("cannot open ptyout file '%s': %w", fileNames.PtyOutFile, err) From 0a828b718480c27d5b4f9b5c943f782549decd5a Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 1 Jul 2022 17:37:37 -0700 Subject: [PATCH 043/149] tightening up server mode, fix bugs, refactor, etc. --- main-mshell.go | 79 ++------------------------ pkg/base/base.go | 119 ++++++++++++++++++--------------------- pkg/cmdtail/cmdtail.go | 18 +++--- pkg/mpio/mpio.go | 6 +- pkg/packet/packet.go | 39 ++++++++++--- pkg/server/server.go | 17 +++++- pkg/shexec/shexec.go | 125 ++++++++++++++++++++++++++++++++--------- 7 files changed, 212 insertions(+), 191 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index a7487660..03cf3400 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -10,9 +10,7 @@ import ( "bytes" "fmt" "os" - "os/user" "strings" - "time" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/cmdtail" @@ -22,43 +20,6 @@ import ( "golang.org/x/sys/unix" ) -func doSingle(ck base.CommandKey) { - packetParser := packet.MakePacketParser(os.Stdin) - sender := packet.MakePacketSender(os.Stdout) - var runPacket *packet.RunPacketType - for pk := range packetParser.MainCh { - if pk.GetType() == packet.RunPacketStr { - runPacket, _ = pk.(*packet.RunPacketType) - break - } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) - return - } - if runPacket == nil { - sender.SendErrorPacket("did not receive a 'run' packet") - return - } - if runPacket.CK.IsEmpty() { - runPacket.CK = ck - } - if runPacket.CK != ck { - sender.SendErrorPacket(fmt.Sprintf("run packet cmdid[%s] did not match arg[%s]", runPacket.CK, ck)) - return - } - cmd, err := shexec.RunCommandDetached(runPacket, sender) - if err != nil { - sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) - return - } - shexec.SetupSignalsForDetach() - startPacket := cmd.MakeCmdStartPacket() - sender.SendPacket(startPacket) - donePacket := cmd.WaitForCommand() - sender.SendPacket(donePacket) - sender.Close() - sender.WaitForDone() -} - func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { err := shexec.ValidateRunPacket(pk) if err != nil { @@ -122,13 +83,8 @@ func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packe } func doMain() { - scHomeDir, err := base.GetScHomeDir() - if err != nil { - packet.SendErrorPacket(os.Stdout, err.Error()) - return - } homeDir := base.GetHomeDir() - err = os.Chdir(homeDir) + err := os.Chdir(homeDir) if err != nil { packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) return @@ -146,13 +102,7 @@ func doMain() { return } go tailer.Run() - initPacket := packet.MakeInitPacket() - initPacket.Env = os.Environ() - initPacket.HomeDir = homeDir - initPacket.ScHomeDir = scHomeDir - if user, _ := user.Current(); user != nil { - initPacket.User = user.Username - } + initPacket := shexec.MakeInitPacket() sender.SendPacket(initPacket) for pk := range packetParser.MainCh { if pk.GetType() == packet.RunPacketStr { @@ -208,19 +158,14 @@ func handleSingle() { packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) defer func() { - // wait for sender to complete sender.Close() sender.WaitForDone() }() + initPacket := shexec.MakeInitPacket() + sender.SendPacket(initPacket) if len(os.Args) >= 3 && os.Args[2] == "--version" { - initPacket := packet.MakeInitPacket() - initPacket.Version = base.MShellVersion - sender.SendPacket(initPacket) return } - initPacket := packet.MakeInitPacket() - initPacket.Version = base.MShellVersion - sender.SendPacket(initPacket) runPacket, err := readFullRunPacket(packetParser) if err != nil { ck := base.CommandKey("") @@ -236,12 +181,11 @@ func handleSingle() { return } if runPacket.Detached { - cmd, err := shexec.RunCommandDetached(runPacket, sender) + err := shexec.RunCommandDetached(runPacket, sender) if err != nil { sender.SendCKErrorPacket(runPacket.CK, err.Error()) return } - cmd.WaitForCommand() } else { cmd, err := shexec.RunCommandSimple(runPacket, sender) if err != nil { @@ -560,17 +504,4 @@ func main() { } return } - - if len(os.Args) >= 2 { - ck := base.CommandKey(os.Args[1]) - if err := ck.Validate("mshell arg"); err != nil { - packet.SendErrorPacket(os.Stdout, err.Error()) - return - } - doSingle(ck) - time.Sleep(100 * time.Millisecond) - return - } else { - doMain() - } } diff --git a/pkg/base/base.go b/pkg/base/base.go index b2736c4c..4933273c 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -9,6 +9,7 @@ package base import ( "errors" "fmt" + "io" "io/fs" "os" "os/exec" @@ -19,20 +20,15 @@ import ( "github.com/google/uuid" ) -const DefaultMShellPath = "mshell" -const DefaultUserMShellPath = ".mshell/mshell" -const MShellPathVarName = "MSHELL_PATH" -const SSHCommandVarName = "SSH_COMMAND" -const ScHomeVarName = "SCRIPTHAUS_HOME" const HomeVarName = "HOME" -const ScShell = "bash" -const SessionsDirBaseName = ".sessions" -const RunnerBaseName = "runner" -const SessionDBName = "session.db" -const ScReadyString = "scripthaus runner ready" +const DefaultMShellHome = "~/.mshell" +const DefaultMShellName = "mshell" +const MShellPathVarName = "MSHELL_PATH" +const MShellHomeVarName = "MSHELL_HOME" +const SSHCommandVarName = "SSH_COMMAND" +const SessionsDirBaseName = "sessions" const MShellVersion = "0.1.0" - -const OSCEscError = "error" +const RemoteIdFile = "remoteid" type CommandFileNames struct { PtyOutFile string @@ -110,16 +106,12 @@ func GetHomeDir() string { return homeVar } -func GetScHomeDir() (string, error) { - scHome := os.Getenv(ScHomeVarName) - if scHome == "" { - homeVar := os.Getenv(HomeVarName) - if homeVar == "" { - return "", fmt.Errorf("Cannot resolve scripthaus home directory (SCRIPTHAUS_HOME and HOME not set)") - } - scHome = path.Join(homeVar, "scripthaus") +func GetMShellHomeDir() string { + homeVar := os.Getenv(MShellHomeVarName) + if homeVar != "" { + return homeVar } - return scHome, nil + return ExpandHomeDir(DefaultMShellHome) } func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { @@ -139,8 +131,8 @@ func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { }, nil } -func MakeCommandFileNamesWithHome(scHome string, ck CommandKey) *CommandFileNames { - base := path.Join(scHome, SessionsDirBaseName, ck.GetSessionId(), ck.GetCmdId()) +func MakeCommandFileNamesWithHome(mhome string, ck CommandKey) *CommandFileNames { + base := path.Join(mhome, SessionsDirBaseName, ck.GetSessionId(), ck.GetCmdId()) return &CommandFileNames{ PtyOutFile: base + ".ptyout", StdinFifo: base + ".stdin", @@ -174,11 +166,8 @@ func EnsureSessionDir(sessionId string) (string, error) { if sessionId == "" { return "", fmt.Errorf("Bad sessionid, cannot be empty") } - shhome, err := GetScHomeDir() - if err != nil { - return "", err - } - sdir := path.Join(shhome, SessionsDirBaseName, sessionId) + mhome := GetMShellHomeDir() + sdir := path.Join(mhome, SessionsDirBaseName, sessionId) info, err := os.Stat(sdir) if errors.Is(err, fs.ErrNotExist) { err = os.MkdirAll(sdir, 0777) @@ -197,51 +186,22 @@ func EnsureSessionDir(sessionId string) (string, error) { } func GetMShellPath() (string, error) { - msPath := os.Getenv(MShellPathVarName) + msPath := os.Getenv(MShellPathVarName) // use MSHELL_PATH if msPath != "" { return exec.LookPath(msPath) } - userMShellPath := path.Join(GetHomeDir(), DefaultUserMShellPath) + mhome := GetMShellHomeDir() + userMShellPath := path.Join(mhome, DefaultMShellName) // look in ~/.mshell msPath, err := exec.LookPath(userMShellPath) - if err != nil { + if err == nil { return msPath, nil } - return exec.LookPath(DefaultMShellPath) + return exec.LookPath(DefaultMShellName) // standard path lookup for 'mshell' } -func GetScSessionsDir() (string, error) { - scHome, err := GetScHomeDir() - if err != nil { - return "", err - } - return path.Join(scHome, SessionsDirBaseName), nil -} - -func GetSessionDBName(sessionId string) (string, error) { - scHome, err := GetScHomeDir() - if err != nil { - return "", err - } - return path.Join(scHome, SessionDBName), nil -} - -// SH OSC Escapes (code 198, S=19, H=8) -// \e]198;cmdid;(cmd-id)BEL - return command-id to server -// \e]198;remote;0BEL - runner program not available -// \e]198;remote;1BEL - runner program is available -// \e]198;error;(error-str)BEL - communicate an internal error -func MakeSHOSCEsc(escName string, data string) string { - return fmt.Sprintf("\033]198;%s;%s\007", escName, data) -} - -func WriteErrorMsg(fileName string, errVal string) error { - fd, err := os.OpenFile(fileName, os.O_APPEND|os.O_WRONLY, 0600) - if err != nil { - return err - } - oscEsc := MakeSHOSCEsc(OSCEscError, errVal) - _, writeErr := fd.Write([]byte(oscEsc)) - return writeErr +func GetMShellSessionsDir() (string, error) { + mhome := GetMShellHomeDir() + return path.Join(mhome, SessionsDirBaseName), nil } func ExpandHomeDir(pathStr string) string { @@ -262,3 +222,32 @@ func ValidGoArch(goos string, goarch string) bool { func GoArchOptFile(goos string, goarch string) string { return fmt.Sprintf("/opt/mshell/bin/mshell.%s.%s", goos, goarch) } + +func GetRemoteId() (string, error) { + mhome := GetMShellHomeDir() + remoteIdFile := path.Join(mhome, RemoteIdFile) + fd, err := os.Open(remoteIdFile) + if errors.Is(err, fs.ErrNotExist) { + // write the file + remoteId := uuid.New().String() + err = os.WriteFile(remoteIdFile, []byte(remoteId), 0644) + if err != nil { + return "", fmt.Errorf("cannot write remoteid to '%s': %w", remoteIdFile, err) + } + return remoteId, nil + } else if err != nil { + return "", fmt.Errorf("cannot read remoteid file '%s': %w", remoteIdFile, err) + } else { + defer fd.Close() + contents, err := io.ReadAll(fd) + if err != nil { + return "", fmt.Errorf("cannot read remoteid file '%s': %w", remoteIdFile, err) + } + uuidStr := string(contents) + _, err = uuid.Parse(uuidStr) + if err != nil { + return "", fmt.Errorf("invalid uuid read from '%s': %w", remoteIdFile, err) + } + return uuidStr, nil + } +} diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 663f8065..23de26c4 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -75,7 +75,7 @@ func (pos TailPos) IsCurrent(entry CmdWatchEntry) bool { type Tailer struct { Lock *sync.Mutex WatchList map[base.CommandKey]CmdWatchEntry - ScHomeDir string + MHomeDir string Watcher *fsnotify.Watcher Sender *packet.PacketSender } @@ -101,7 +101,7 @@ func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { } // delete from watchlist, remove watches - fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, cmdKey) + fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, cmdKey) delete(t.WatchList, cmdKey) t.Watcher.Remove(fileNames.PtyOutFile) t.Watcher.Remove(fileNames.RunnerOutFile) @@ -130,16 +130,14 @@ func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (Cm } func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { - scHomeDir, err := base.GetScHomeDir() - if err != nil { - return nil, err - } + mhomeDir := base.GetMShellHomeDir() rtn := &Tailer{ Lock: &sync.Mutex{}, WatchList: make(map[base.CommandKey]CmdWatchEntry), - ScHomeDir: scHomeDir, + MHomeDir: mhomeDir, Sender: sender, } + var err error rtn.Watcher, err = fsnotify.NewWatcher() if err != nil { return nil, err @@ -196,7 +194,7 @@ func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*pack if !foundPos { return nil, false } - fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, key) + fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, key) dataPacket := t.makeCmdDataPacket(fileNames, entry, pos) t.Lock.Lock() @@ -353,7 +351,7 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { if getPacket.ReqId == "" { return fmt.Errorf("getcmd, no reqid specified") } - fileNames := base.MakeCommandFileNamesWithHome(t.ScHomeDir, getPacket.CK) + fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, getPacket.CK) t.Lock.Lock() defer t.Lock.Unlock() key := getPacket.CK @@ -370,7 +368,7 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { return err } entry = CmdWatchEntry{CmdKey: key} - entry.fillFilePos(t.ScHomeDir) + entry.fillFilePos(t.MHomeDir) } pos, foundPos := entry.getTailPos(getPacket.ReqId) if !foundPos { diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 6cff436a..4fcea20c 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -200,7 +200,7 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { defer m.HandleInputDone() for pk := range m.Input.MainCh { if m.Debug { - fmt.Printf("PK> %s\n", packet.AsString(pk)) + fmt.Printf("PK-M> %s\n", packet.AsString(pk)) } if pk.GetType() == packet.DataPacketStr { dataPacket := pk.(*packet.DataPacketType) @@ -220,10 +220,6 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { donePacket := pk.(*packet.CmdDonePacketType) return donePacket } - if pk.GetType() == packet.CmdStartPacketStr { - // nothing - continue - } m.UPR.UnknownPacket(pk) } return nil diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index f49db0c4..942a023f 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -73,6 +73,10 @@ func init() { TypeStrToFactory[DataEndPacketStr] = reflect.TypeOf(DataEndPacketType{}) } +func RegisterPacketType(typeStr string, rtype reflect.Type) { + TypeStrToFactory[typeStr] = rtype +} + func MakePacket(packetType string) (PacketType, error) { rtype := TypeStrToFactory[packetType] if rtype == nil { @@ -355,14 +359,15 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { } type InitPacketType struct { - Type string `json:"type"` - Version string `json:"version"` - ScHomeDir string `json:"schomedir,omitempty"` - HomeDir string `json:"homedir,omitempty"` - Env []string `json:"env,omitempty"` - User string `json:"user,omitempty"` - NotFound bool `json:"notfound,omitempty"` - UName string `json:"uname,omitempty"` + Type string `json:"type"` + Version string `json:"version"` + MShellHomeDir string `json:"mshellhomedir,omitempty"` + HomeDir string `json:"homedir,omitempty"` + Env []string `json:"env,omitempty"` + User string `json:"user,omitempty"` + NotFound bool `json:"notfound,omitempty"` + UName string `json:"uname,omitempty"` + RemoteId string `json:"remoteid,omitempty"` } func (*InitPacketType) GetType() string { @@ -615,6 +620,22 @@ func MakePacketSender(output io.Writer) *PacketSender { return sender } +func MakeChannelPacketSender(packetCh chan PacketType) *PacketSender { + sender := &PacketSender{ + Lock: &sync.Mutex{}, + SendCh: make(chan PacketType, PacketSenderQueueSize), + DoneCh: make(chan bool), + } + go func() { + defer close(sender.DoneCh) + defer sender.Close() + for pk := range sender.SendCh { + packetCh <- pk + } + }() + return sender +} + func (sender *PacketSender) Close() { sender.Lock.Lock() defer sender.Lock.Unlock() @@ -676,6 +697,8 @@ func (DefaultUPR) UnknownPacket(pk PacketType) { } else if pk.GetType() == RawPacketStr { rawPacket := pk.(*RawPacketType) fmt.Fprintf(os.Stderr, "%s\n", rawPacket.Data) + } else if pk.GetType() == CmdStartPacketStr { + return // do nothing } else { fmt.Fprintf(os.Stderr, "[error] invalid packet received '%s'", AsExtType(pk)) } diff --git a/pkg/server/server.go b/pkg/server/server.go index a04f5f98..2b28a19a 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -150,8 +150,11 @@ func RunServer() (int, error) { server.MainInput = packet.MakePacketParser(os.Stdin) server.Sender = packet.MakePacketSender(os.Stdout) defer server.Close() - initPacket := packet.MakeInitPacket() - initPacket.Version = base.MShellVersion + var err error + initPacket, err := shexec.MakeServerInitPacket() + if err != nil { + return 1, err + } server.Sender.SendPacket(initPacket) builder := packet.MakeRunPacketBuilder() for pk := range server.MainInput.MainCh { @@ -159,6 +162,9 @@ func RunServer() (int, error) { fmt.Printf("PK> %s\n", packet.AsString(pk)) } ok, runPacket := builder.ProcessPacket(pk) + if server.Debug { + fmt.Printf("PP> %s | %v\n", pk.GetType(), ok) + } if ok { if runPacket != nil { server.runCommand(runPacket) @@ -166,6 +172,13 @@ func RunServer() (int, error) { } continue } + if startPk, ok := pk.(*packet.CmdStartPacketType); ok { + if server.Debug { + fmt.Printf("START> %v", startPk) + } + server.Sender.SendPacket(startPk) + continue + } if cmdPk, ok := pk.(packet.CommandPacketType); ok { server.ProcessCommandPacket(cmdPk) continue diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 788de2cb..a19edafe 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -13,6 +13,7 @@ import ( "os" "os/exec" "os/signal" + "os/user" "strings" "sync" "syscall" @@ -23,6 +24,7 @@ import ( "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" + "golang.org/x/sys/unix" ) const DefaultRows = 25 @@ -57,13 +59,16 @@ const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` type ShExecType struct { - Lock *sync.Mutex - StartTs time.Time - CK base.CommandKey - FileNames *base.CommandFileNames - Cmd *exec.Cmd - CmdPty *os.File - Multiplexer *mpio.Multiplexer + Lock *sync.Mutex + StartTs time.Time + CK base.CommandKey + FileNames *base.CommandFileNames + Cmd *exec.Cmd + CmdPty *os.File + Multiplexer *mpio.Multiplexer + Detached bool + DetachedOutput *packet.PacketSender + RunnerOutFd *os.File } type StdContext struct{} @@ -115,6 +120,13 @@ func (c *ShExecType) Close() { c.CmdPty.Close() } c.Multiplexer.Close() + if c.DetachedOutput != nil { + c.DetachedOutput.Close() + c.DetachedOutput.WaitForDone() + } + if c.RunnerOutFd != nil { + c.RunnerOutFd.Close() + } } func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { @@ -300,11 +312,13 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { func GetWinsize(p *packet.RunPacketType) *pty.Winsize { rows := DefaultRows cols := DefaultCols - if p.TermSize.Rows > 0 && p.TermSize.Rows <= MaxRows { - rows = p.TermSize.Rows - } - if p.TermSize.Cols > 0 && p.TermSize.Cols <= MaxCols { - cols = p.TermSize.Cols + if p.TermSize != nil { + if p.TermSize.Rows > 0 && p.TermSize.Rows <= MaxRows { + rows = p.TermSize.Rows + } + if p.TermSize.Cols > 0 && p.TermSize.Cols <= MaxCols { + cols = p.TermSize.Cols + } } return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} } @@ -834,63 +848,98 @@ func SetupSignalsForDetach() { }() } -func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { +func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) error { fileNames, err := base.GetCommandFileNames(pk.CK) if err != nil { - return nil, err + return err } ptyOutInfo, err := os.Stat(fileNames.PtyOutFile) if err == nil { // non-nil error will be caught by regular OpenFile below // must have size 0 if ptyOutInfo.Size() != 0 { - return nil, fmt.Errorf("cmdkey '%s' was already used (ptyout len=%d)", pk.CK, ptyOutInfo.Size()) + return fmt.Errorf("cmdkey '%s' was already used (ptyout len=%d)", pk.CK, ptyOutInfo.Size()) } } cmdPty, cmdTty, err := pty.Open() if err != nil { - return nil, fmt.Errorf("opening new pty: %w", err) + return fmt.Errorf("opening new pty: %w", err) } pty.Setsize(cmdPty, GetWinsize(pk)) defer func() { cmdTty.Close() }() - rtn := MakeShExec(pk.CK, nil) + cmd := MakeShExec(pk.CK, nil) + cmd.FileNames = fileNames + cmd.CmdPty = cmdPty + cmd.Detached = true + cmd.RunnerOutFd, err = os.OpenFile(fileNames.RunnerOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) + if err != nil { + return fmt.Errorf("cannot open runout file '%s': %w", fileNames.RunnerOutFile, err) + } + nullFd, err := os.OpenFile("/dev/null", os.O_RDWR, 0) + if err != nil { + return fmt.Errorf("cannot open /dev/null: %w", err) + } + cmd.DetachedOutput = packet.MakePacketSender(cmd.RunnerOutFd) ecmd, err := MakeDetachedExecCmd(pk, cmdTty) if err != nil { - return nil, err + return err } + cmd.Cmd = ecmd SetupSignalsForDetach() err = ecmd.Start() if err != nil { - return nil, fmt.Errorf("starting command: %w", err) + return fmt.Errorf("starting command: %w", err) } for _, fd := range ecmd.ExtraFiles { if fd != cmdTty { fd.Close() } } - ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY, 0600) + // after Start(), any errors must go to DetachedOutput + // close stdin/stdout/stderr, but wait for cmdstart packet to get sent + startPacket := cmd.MakeCmdStartPacket() + go func() { + sender.SendPacket(startPacket) + sender.Close() + sender.WaitForDone() + fmt.Printf("sender done! start: %v\n", startPacket) + err = unix.Dup2(int(nullFd.Fd()), int(os.Stdin.Fd())) + if err != nil { + cmd.DetachedOutput.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot dup2 stdin to /dev/null: %w", err)) + } + err = unix.Dup2(int(nullFd.Fd()), int(os.Stdout.Fd())) + if err != nil { + cmd.DetachedOutput.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot dup2 stdin to /dev/null: %w", err)) + } + err = unix.Dup2(int(nullFd.Fd()), int(os.Stderr.Fd())) + if err != nil { + cmd.DetachedOutput.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot dup2 stdin to /dev/null: %w", err)) + } + cmd.DetachedOutput.SendPacket(startPacket) + }() + ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err != nil { - return nil, fmt.Errorf("cannot open ptyout file '%s': %w", fileNames.PtyOutFile, err) + cmd.DetachedOutput.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open ptyout file '%s': %v", fileNames.PtyOutFile, err)) + // don't return (command is already running) } go func() { // copy pty output to .ptyout file _, copyErr := io.Copy(ptyOutFd, cmdPty) if copyErr != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) + cmd.DetachedOutput.SendCKErrorPacket(pk.CK, fmt.Sprintf("copying pty output to ptyout file: %v", copyErr)) } }() go func() { // copy .stdin fifo contents to pty input copyFifoErr := MakeAndCopyStdinFifo(cmdPty, fileNames.StdinFifo) if copyFifoErr != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) + cmd.DetachedOutput.SendCKErrorPacket(pk.CK, fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) } }() - rtn.FileNames = fileNames - rtn.Cmd = ecmd - rtn.CmdPty = cmdPty - return rtn, nil + donePacket := cmd.WaitForCommand() + cmd.DetachedOutput.SendPacket(donePacket) + return nil } func GetExitCode(err error) int { @@ -919,3 +968,25 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { } return donePacket } + +func MakeInitPacket() *packet.InitPacketType { + initPacket := packet.MakeInitPacket() + initPacket.Version = base.MShellVersion + initPacket.HomeDir = base.GetHomeDir() + initPacket.MShellHomeDir = base.GetMShellHomeDir() + if user, _ := user.Current(); user != nil { + initPacket.User = user.Username + } + return initPacket +} + +func MakeServerInitPacket() (*packet.InitPacketType, error) { + var err error + initPacket := MakeInitPacket() + initPacket.Env = os.Environ() + initPacket.RemoteId, err = base.GetRemoteId() + if err != nil { + return nil, err + } + return initPacket, nil +} From ef362e5ee96a9f44c7d5110124a401ca9949daf9 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 5 Jul 2022 16:53:31 -0700 Subject: [PATCH 044/149] tighten up packet interfaces, RpcPacketType, RpcResponsePacketType, and CommandPacketType --- main-mshell.go | 2 +- pkg/cmdtail/cmdtail.go | 2 +- pkg/packet/packet.go | 118 ++++++++++++++++++++++------------------- 3 files changed, 65 insertions(+), 57 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 03cf3400..d7f0a0ad 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -121,7 +121,7 @@ func doMain() { if pk.GetType() == packet.CdPacketStr { cdPacket := pk.(*packet.CdPacketType) err := os.Chdir(cdPacket.Dir) - resp := packet.MakeResponsePacket(cdPacket.PacketId) + resp := packet.MakeResponsePacket(cdPacket.ReqId) if err != nil { resp.Error = err.Error() } else { diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 23de26c4..c301c0cf 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -161,7 +161,7 @@ func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]b func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWatchEntry, pos TailPos) *packet.CmdDataPacketType { dataPacket := packet.MakeCmdDataPacket() - dataPacket.ReqId = pos.ReqId + dataPacket.RespId = pos.ReqId dataPacket.CK = entry.CmdKey dataPacket.PtyPos = pos.TailPtyPos dataPacket.RunPos = pos.TailRunPos diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 942a023f..aab905e7 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -27,24 +27,24 @@ import ( var GlobalDebug = false const ( - RunPacketStr = "run" + RunPacketStr = "run" // rpc PingPacketStr = "ping" InitPacketStr = "init" - DataPacketStr = "data" - DataAckPacketStr = "dataack" - CmdStartPacketStr = "cmdstart" - CmdDonePacketStr = "cmddone" + DataPacketStr = "data" // command + DataAckPacketStr = "dataack" // command + CmdStartPacketStr = "cmdstart" // rpc-response + CmdDonePacketStr = "cmddone" // command DataEndPacketStr = "dataend" - ResponsePacketStr = "resp" + ResponsePacketStr = "resp" // rpc-response DonePacketStr = "done" ErrorPacketStr = "error" MessagePacketStr = "message" - GetCmdPacketStr = "getcmd" - UntailCmdPacketStr = "untailcmd" - CdPacketStr = "cd" - CmdDataPacketStr = "cmddata" + GetCmdPacketStr = "getcmd" // rpc + UntailCmdPacketStr = "untailcmd" // rpc + CdPacketStr = "cd" // rpc + CmdDataPacketStr = "cmddata" // rpc-response RawPacketStr = "raw" - InputPacketStr = "input" + InputPacketStr = "input" // command ) const PacketSenderQueueSize = 20 @@ -71,6 +71,20 @@ func init() { TypeStrToFactory[DataPacketStr] = reflect.TypeOf(DataPacketType{}) TypeStrToFactory[DataAckPacketStr] = reflect.TypeOf(DataAckPacketType{}) TypeStrToFactory[DataEndPacketStr] = reflect.TypeOf(DataEndPacketType{}) + + var _ RpcPacketType = (*RunPacketType)(nil) + var _ RpcPacketType = (*GetCmdPacketType)(nil) + var _ RpcPacketType = (*UntailCmdPacketType)(nil) + var _ RpcPacketType = (*CdPacketType)(nil) + + var _ RpcResponsePacketType = (*CmdStartPacketType)(nil) + var _ RpcResponsePacketType = (*ResponsePacketType)(nil) + var _ RpcResponsePacketType = (*CmdDataPacketType)(nil) + + var _ CommandPacketType = (*DataPacketType)(nil) + var _ CommandPacketType = (*DataAckPacketType)(nil) + var _ CommandPacketType = (*CmdDonePacketType)(nil) + var _ CommandPacketType = (*InputPacketType)(nil) } func RegisterPacketType(typeStr string, rtype reflect.Type) { @@ -88,7 +102,7 @@ func MakePacket(packetType string) (PacketType, error) { type CmdDataPacketType struct { Type string `json:"type"` - ReqId string `json:"reqid"` + RespId string `json:"respid"` CK base.CommandKey `json:"ck"` PtyPos int64 `json:"ptypos"` PtyLen int64 `json:"ptylen"` @@ -106,8 +120,8 @@ func (*CmdDataPacketType) GetType() string { return CmdDataPacketStr } -func (p *CmdDataPacketType) GetCK() base.CommandKey { - return p.CK +func (p *CmdDataPacketType) GetResponseId() string { + return p.RespId } func MakeCmdDataPacket() *CmdDataPacketType { @@ -186,10 +200,6 @@ func (*DataEndPacketType) GetType() string { return DataEndPacketStr } -func (p *DataEndPacketType) GetCK() base.CommandKey { - return p.CK -} - type DataAckPacketType struct { Type string `json:"type"` CK base.CommandKey `json:"ck"` @@ -252,8 +262,8 @@ func (*UntailCmdPacketType) GetType() string { return UntailCmdPacketStr } -func (p *UntailCmdPacketType) GetCK() base.CommandKey { - return p.CK +func (p *UntailCmdPacketType) GetReqId() string { + return p.ReqId } func MakeUntailCmdPacket() *UntailCmdPacketType { @@ -273,8 +283,8 @@ func (*GetCmdPacketType) GetType() string { return GetCmdPacketStr } -func (p *GetCmdPacketType) GetCK() base.CommandKey { - return p.CK +func (p *GetCmdPacketType) GetReqId() string { + return p.ReqId } func MakeGetCmdPacket() *GetCmdPacketType { @@ -282,17 +292,17 @@ func MakeGetCmdPacket() *GetCmdPacketType { } type CdPacketType struct { - Type string `json:"type"` - PacketId string `json:"packetid"` - Dir string `json:"dir"` + Type string `json:"type"` + ReqId string `json:"reqid"` + Dir string `json:"dir"` } func (*CdPacketType) GetType() string { return CdPacketStr } -func (p *CdPacketType) GetPacketId() string { - return p.PacketId +func (p *CdPacketType) GetReqId() string { + return p.ReqId } func MakeCdPacket() *CdPacketType { @@ -300,23 +310,23 @@ func MakeCdPacket() *CdPacketType { } type ResponsePacketType struct { - Type string `json:"type"` - PacketId string `json:"packetid"` - Success bool `json:"success"` - Error string `json:"error"` - Data interface{} `json:"data"` + Type string `json:"type"` + RespId string `json:"respid"` + Success bool `json:"success"` + Error string `json:"error"` + Data interface{} `json:"data"` } func (*ResponsePacketType) GetType() string { return ResponsePacketStr } -func (p *ResponsePacketType) GetPacketId() string { - return p.PacketId +func (p *ResponsePacketType) GetResponseId() string { + return p.RespId } -func MakeResponsePacket(packetId string) *ResponsePacketType { - return &ResponsePacketType{Type: ResponsePacketStr, PacketId: packetId} +func MakeResponsePacket(reqId string) *ResponsePacketType { + return &ResponsePacketType{Type: ResponsePacketStr, RespId: reqId} } type RawPacketType struct { @@ -412,6 +422,7 @@ func MakeCmdDonePacket() *CmdDonePacketType { type CmdStartPacketType struct { Type string `json:"type"` + RespId string `json:"respid"` Ts int64 `json:"ts"` CK base.CommandKey `json:"ck"` Pid int `json:"pid"` @@ -422,8 +433,8 @@ func (*CmdStartPacketType) GetType() string { return CmdStartPacketStr } -func (p *CmdStartPacketType) GetCK() base.CommandKey { - return p.CK +func (p *CmdStartPacketType) GetResponseId() string { + return p.RespId } func MakeCmdStartPacket() *CmdStartPacketType { @@ -450,6 +461,7 @@ type RunDataType struct { type RunPacketType struct { Type string `json:"type"` + ReqId string `json:"packetid"` CK base.CommandKey `json:"ck"` Command string `json:"command"` Cwd string `json:"cwd,omitempty"` @@ -464,8 +476,8 @@ func (*RunPacketType) GetType() string { return RunPacketStr } -func (p *RunPacketType) GetCK() base.CommandKey { - return p.CK +func (p *RunPacketType) GetReqId() string { + return p.ReqId } func MakeRunPacket() *RunPacketType { @@ -477,9 +489,8 @@ type BarePacketType struct { } type ErrorPacketType struct { - CK base.CommandKey `json:"ck,omitempty"` - Type string `json:"type"` - Error string `json:"error"` + Type string `json:"type"` + Error string `json:"error"` } func (*ErrorPacketType) GetType() string { @@ -487,21 +498,13 @@ func (*ErrorPacketType) GetType() string { } func (p *ErrorPacketType) String() string { - ckStr := "" - if p.CK != "" { - ckStr = fmt.Sprintf(", ck=%s", p.CK) - } - return fmt.Sprintf("error[%s%s]", p.Error, ckStr) + return fmt.Sprintf("error[%s]", p.Error) } func MakeErrorPacket(errorStr string) *ErrorPacketType { return &ErrorPacketType{Type: ErrorPacketStr, Error: errorStr} } -func MakeCKErrorPacket(ck base.CommandKey, errorStr string) *ErrorPacketType { - return &ErrorPacketType{Type: ErrorPacketStr, CK: ck, Error: errorStr} -} - type PacketType interface { GetType() string } @@ -515,7 +518,12 @@ func AsString(pk PacketType) string { type RpcPacketType interface { GetType() string - GetPacketId() string + GetReqId() string +} + +type RpcResponsePacketType interface { + GetType() string + GetResponseId() string } type CommandPacketType interface { @@ -525,7 +533,7 @@ type CommandPacketType interface { func AsExtType(pk PacketType) string { if rpcPacket, ok := pk.(RpcPacketType); ok { - return fmt.Sprintf("%s[%s]", rpcPacket.GetType(), rpcPacket.GetPacketId()) + return fmt.Sprintf("%s[%s]", rpcPacket.GetType(), rpcPacket.GetReqId()) } else if cmdPacket, ok := pk.(CommandPacketType); ok { return fmt.Sprintf("%s[%s]", cmdPacket.GetType(), cmdPacket.GetCK()) } else { @@ -676,7 +684,7 @@ func (sender *PacketSender) SendErrorPacket(errVal string) error { } func (sender *PacketSender) SendCKErrorPacket(ck base.CommandKey, errVal string) error { - return sender.SendPacket(MakeCKErrorPacket(ck, errVal)) + return sender.SendPacket(MakeErrorPacket(errVal)) } func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) error { From 0c204e8b2b833d82dbee5f857666179294275f88 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 5 Jul 2022 17:45:46 -0700 Subject: [PATCH 045/149] standardize error reporting, rpc gets resp, command get cmderr, other errors are just sent as messages --- main-mshell.go | 245 +++++++++++++++++++++---------------------- pkg/packet/packet.go | 64 ++++++----- pkg/packet/parser.go | 24 +++-- pkg/server/server.go | 12 +-- pkg/shexec/shexec.go | 12 +-- 5 files changed, 190 insertions(+), 167 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index d7f0a0ad..4e58eba2 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -13,132 +13,131 @@ import ( "strings" "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/cmdtail" "github.com/scripthaus-dev/mshell/pkg/packet" "github.com/scripthaus-dev/mshell/pkg/server" "github.com/scripthaus-dev/mshell/pkg/shexec" "golang.org/x/sys/unix" ) -func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { - err := shexec.ValidateRunPacket(pk) - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("invalid run packet: %v", err)) - return - } - fileNames, err := base.GetCommandFileNames(pk.CK) - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot get command file names: %v", err)) - return - } - cmd, err := shexec.MakeRunnerExec(pk.CK) - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot make mshell command: %v", err)) - return - } - cmdStdin, err := cmd.StdinPipe() - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot pipe stdin to command: %v", err)) - return - } - // touch ptyout file (should exist for tailer to work correctly) - ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err)) - return - } - ptyOutFd.Close() // just opened to create the file, can close right after - runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err)) - return - } - defer runnerOutFd.Close() - cmd.Stdout = runnerOutFd - cmd.Stderr = runnerOutFd - err = cmd.Start() - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error starting command: %v", err)) - return - } - go func() { - err = packet.SendPacket(cmdStdin, pk) - if err != nil { - sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error sending forked runner command: %v", err)) - return - } - cmdStdin.Close() +// func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { +// err := shexec.ValidateRunPacket(pk) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("invalid run packet: %v", err)) +// return +// } +// fileNames, err := base.GetCommandFileNames(pk.CK) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot get command file names: %v", err)) +// return +// } +// cmd, err := shexec.MakeRunnerExec(pk.CK) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot make mshell command: %v", err)) +// return +// } +// cmdStdin, err := cmd.StdinPipe() +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot pipe stdin to command: %v", err)) +// return +// } +// // touch ptyout file (should exist for tailer to work correctly) +// ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err)) +// return +// } +// ptyOutFd.Close() // just opened to create the file, can close right after +// runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err)) +// return +// } +// defer runnerOutFd.Close() +// cmd.Stdout = runnerOutFd +// cmd.Stderr = runnerOutFd +// err = cmd.Start() +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error starting command: %v", err)) +// return +// } +// go func() { +// err = packet.SendPacket(cmdStdin, pk) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error sending forked runner command: %v", err)) +// return +// } +// cmdStdin.Close() - // clean up zombies - cmd.Wait() - }() -} +// // clean up zombies +// cmd.Wait() +// }() +// } -func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { - err := tailer.AddWatch(pk) - if err != nil { - return err - } - return nil -} +// func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { +// err := tailer.AddWatch(pk) +// if err != nil { +// return err +// } +// return nil +// } -func doMain() { - homeDir := base.GetHomeDir() - err := os.Chdir(homeDir) - if err != nil { - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) - return - } - _, err = base.GetMShellPath() - if err != nil { - packet.SendErrorPacket(os.Stdout, err.Error()) - return - } - packetParser := packet.MakePacketParser(os.Stdin) - sender := packet.MakePacketSender(os.Stdout) - tailer, err := cmdtail.MakeTailer(sender) - if err != nil { - packet.SendErrorPacket(os.Stdout, err.Error()) - return - } - go tailer.Run() - initPacket := shexec.MakeInitPacket() - sender.SendPacket(initPacket) - for pk := range packetParser.MainCh { - if pk.GetType() == packet.RunPacketStr { - doMainRun(pk.(*packet.RunPacketType), sender) - continue - } - if pk.GetType() == packet.GetCmdPacketStr { - err = doGetCmd(tailer, pk.(*packet.GetCmdPacketType), sender) - if err != nil { - errPk := packet.MakeErrorPacket(err.Error()) - sender.SendPacket(errPk) - continue - } - continue - } - if pk.GetType() == packet.CdPacketStr { - cdPacket := pk.(*packet.CdPacketType) - err := os.Chdir(cdPacket.Dir) - resp := packet.MakeResponsePacket(cdPacket.ReqId) - if err != nil { - resp.Error = err.Error() - } else { - resp.Success = true - } - sender.SendPacket(resp) - continue - } - if pk.GetType() == packet.ErrorPacketStr { - errPk := pk.(*packet.ErrorPacketType) - errPk.Error = "invalid packet sent to mshell: " + errPk.Error - sender.SendPacket(errPk) - continue - } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) - } -} +// func doMain() { +// homeDir := base.GetHomeDir() +// err := os.Chdir(homeDir) +// if err != nil { +// packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) +// return +// } +// _, err = base.GetMShellPath() +// if err != nil { +// packet.SendErrorPacket(os.Stdout, err.Error()) +// return +// } +// packetParser := packet.MakePacketParser(os.Stdin) +// sender := packet.MakePacketSender(os.Stdout) +// tailer, err := cmdtail.MakeTailer(sender) +// if err != nil { +// packet.SendErrorPacket(os.Stdout, err.Error()) +// return +// } +// go tailer.Run() +// initPacket := shexec.MakeInitPacket() +// sender.SendPacket(initPacket) +// for pk := range packetParser.MainCh { +// if pk.GetType() == packet.RunPacketStr { +// doMainRun(pk.(*packet.RunPacketType), sender) +// continue +// } +// if pk.GetType() == packet.GetCmdPacketStr { +// err = doGetCmd(tailer, pk.(*packet.GetCmdPacketType), sender) +// if err != nil { +// errPk := packet.MakeErrorPacket(err.Error()) +// sender.SendPacket(errPk) +// continue +// } +// continue +// } +// if pk.GetType() == packet.CdPacketStr { +// cdPacket := pk.(*packet.CdPacketType) +// err := os.Chdir(cdPacket.Dir) +// resp := packet.MakeResponsePacket(cdPacket.ReqId) +// if err != nil { +// resp.Error = err.Error() +// } else { +// resp.Success = true +// } +// sender.SendPacket(resp) +// continue +// } +// if pk.GetType() == packet.ErrorPacketStr { +// errPk := pk.(*packet.ErrorPacketType) +// errPk.Error = "invalid packet sent to mshell: " + errPk.Error +// sender.SendPacket(errPk) +// continue +// } +// sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) +// } +// } func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType, error) { rpb := packet.MakeRunPacketBuilder() @@ -168,28 +167,24 @@ func handleSingle() { } runPacket, err := readFullRunPacket(packetParser) if err != nil { - ck := base.CommandKey("") - if runPacket != nil { - ck = runPacket.CK - } - sender.SendCKErrorPacket(ck, err.Error()) + sender.SendErrorResponse(runPacket.ReqId, err) return } err = shexec.ValidateRunPacket(runPacket) if err != nil { - sender.SendCKErrorPacket(runPacket.CK, err.Error()) + sender.SendErrorResponse(runPacket.ReqId, err) return } if runPacket.Detached { err := shexec.RunCommandDetached(runPacket, sender) if err != nil { - sender.SendCKErrorPacket(runPacket.CK, err.Error()) + sender.SendErrorResponse(runPacket.ReqId, err) return } } else { cmd, err := shexec.RunCommandSimple(runPacket, sender) if err != nil { - sender.SendCKErrorPacket(runPacket.CK, fmt.Sprintf("error running command: %v", err)) + sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("error running command: %w", err)) return } defer cmd.Close() diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index aab905e7..dcdd5f09 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -19,10 +19,11 @@ import ( "github.com/scripthaus-dev/mshell/pkg/base" ) -// remote: init, run, ping, data, cmdstart, cmddone -// remote(detached): init, run, cmdstart -// server: init, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] -// all: error, message +// single : run, >cmddata, >cmddone, data, <>dataack, run, >cmddata, >cmddone, run, >cmddata, >cmddone, data, <>dataack, cd, >getcmd, >untailcmd, >input, error, <>message, <>ping, Date: Tue, 5 Jul 2022 23:14:14 -0700 Subject: [PATCH 046/149] checkpoint for tightened runtime semantics for calls -- always send response packets, make sure correct response ids are set, etc. --- main-mshell.go | 24 +++++---- pkg/cmdtail/cmdtail.go | 106 +++++++++++++++++++++---------------- pkg/packet/packet.go | 31 +++++++---- pkg/packet/parser.go | 77 +++++++++++++++++++++++++++ pkg/server/server.go | 4 ++ pkg/shexec/shexec.go | 117 ++++++++++++++++++++--------------------- 6 files changed, 236 insertions(+), 123 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 4e58eba2..351130a0 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -13,6 +13,7 @@ import ( "strings" "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/cmdtail" "github.com/scripthaus-dev/mshell/pkg/packet" "github.com/scripthaus-dev/mshell/pkg/server" "github.com/scripthaus-dev/mshell/pkg/shexec" @@ -73,13 +74,13 @@ import ( // }() // } -// func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { -// err := tailer.AddWatch(pk) -// if err != nil { -// return err -// } -// return nil -// } +func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { + err := tailer.AddWatch(pk) + if err != nil { + return err + } + return nil +} // func doMain() { // homeDir := base.GetHomeDir() @@ -176,11 +177,16 @@ func handleSingle() { return } if runPacket.Detached { - err := shexec.RunCommandDetached(runPacket, sender) + cmd, startPk, err := shexec.RunCommandDetached(runPacket, sender) if err != nil { sender.SendErrorResponse(runPacket.ReqId, err) return } + sender.SendPacket(startPk) + sender.Close() + sender.WaitForDone() + cmd.DetachedWait(startPk) + return } else { cmd, err := shexec.RunCommandSimple(runPacket, sender) if err != nil { @@ -188,7 +194,7 @@ func handleSingle() { return } defer cmd.Close() - startPacket := cmd.MakeCmdStartPacket() + startPacket := cmd.MakeCmdStartPacket(runPacket.ReqId) sender.SendPacket(startPacket) cmd.RunRemoteIOAndWait(packetParser, sender) return diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index c301c0cf..d39c56e1 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -7,6 +7,7 @@ package cmdtail import ( + "encoding/base64" "fmt" "io" "os" @@ -89,6 +90,12 @@ func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos t.WatchList[cmdKey] = entry } +func (t *Tailer) removeTailPos(cmdKey base.CommandKey, reqId string) { + t.Lock.Lock() + defer t.Lock.Unlock() + t.removeTailPos_nolock(cmdKey, reqId) +} + func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { entry, found := t.WatchList[cmdKey] if !found { @@ -107,16 +114,6 @@ func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { t.Watcher.Remove(fileNames.RunnerOutFile) } -func (t *Tailer) updateEntrySizes_nolock(cmdKey base.CommandKey, ptyLen int64, runLen int64) { - entry, found := t.WatchList[cmdKey] - if !found { - return - } - entry.FilePtyLen = ptyLen - entry.FileRunLen = runLen - t.WatchList[cmdKey] = entry -} - func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (CmdWatchEntry, TailPos, bool) { entry, found := t.WatchList[cmdKey] if !found { @@ -159,90 +156,98 @@ func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]b return buf[0:nr], nil } -func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWatchEntry, pos TailPos) *packet.CmdDataPacketType { - dataPacket := packet.MakeCmdDataPacket() - dataPacket.RespId = pos.ReqId +func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWatchEntry, pos TailPos) (*packet.CmdDataPacketType, error) { + dataPacket := packet.MakeCmdDataPacket(pos.ReqId) dataPacket.CK = entry.CmdKey dataPacket.PtyPos = pos.TailPtyPos dataPacket.RunPos = pos.TailRunPos if entry.FilePtyLen > pos.TailPtyPos { ptyData, err := t.readDataFromFile(fileNames.PtyOutFile, pos.TailPtyPos, MaxDataBytes) if err != nil { - dataPacket.Error = err.Error() - return dataPacket + return nil, err } - dataPacket.PtyData = string(ptyData) + dataPacket.PtyData64 = base64.StdEncoding.EncodeToString(ptyData) dataPacket.PtyDataLen = len(ptyData) } if entry.FileRunLen > pos.TailRunPos { runData, err := t.readDataFromFile(fileNames.RunnerOutFile, pos.TailRunPos, MaxDataBytes) if err != nil { - dataPacket.Error = err.Error() - return dataPacket + return nil, err } - dataPacket.RunData = string(runData) + dataPacket.RunData64 = base64.StdEncoding.EncodeToString(runData) dataPacket.RunDataLen = len(runData) } - return dataPacket + return dataPacket, nil } // returns (data-packet, keepRunning) -func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*packet.CmdDataPacketType, bool) { +func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*packet.CmdDataPacketType, bool, error) { t.Lock.Lock() entry, pos, foundPos := t.getEntryAndPos_nolock(key, reqId) t.Lock.Unlock() if !foundPos { - return nil, false + return nil, false, nil } fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, key) - dataPacket := t.makeCmdDataPacket(fileNames, entry, pos) + dataPacket, dataErr := t.makeCmdDataPacket(fileNames, entry, pos) t.Lock.Lock() defer t.Lock.Unlock() entry, pos, foundPos = t.getEntryAndPos_nolock(key, reqId) if !foundPos { - return nil, false + return nil, false, nil } // pos was updated between first and second get, throw out data-packet and re-run if pos.TailPtyPos != dataPacket.PtyPos || pos.TailRunPos != dataPacket.RunPos { - return nil, true + return nil, true, nil } - if dataPacket.Error != "" { + if dataErr != nil { // error, so return error packet, and stop running pos.Running = false t.updateTailPos_nolock(key, reqId, pos) - return dataPacket, false + return nil, false, dataErr } - pos.TailPtyPos += int64(len(dataPacket.PtyData)) - pos.TailRunPos += int64(len(dataPacket.RunData)) + pos.TailPtyPos += int64(dataPacket.PtyDataLen) + pos.TailRunPos += int64(dataPacket.RunDataLen) if pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen { // we caught up, tail position equals file length pos.Running = false } t.updateTailPos_nolock(key, reqId, pos) - return dataPacket, pos.Running + return dataPacket, pos.Running, nil } -func (t *Tailer) checkRemoveNoFollow(cmdKey base.CommandKey, reqId string) { +// returns (removed) +func (t *Tailer) checkRemoveNoFollow(cmdKey base.CommandKey, reqId string) bool { t.Lock.Lock() defer t.Lock.Unlock() _, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) if !foundPos { - return + return false } if !pos.Follow { t.removeTailPos_nolock(cmdKey, reqId) + return true } + return false } func (t *Tailer) RunDataTransfer(key base.CommandKey, reqId string) { for { - dataPacket, keepRunning := t.runSingleDataTransfer(key, reqId) + dataPacket, keepRunning, err := t.runSingleDataTransfer(key, reqId) if dataPacket != nil { t.Sender.SendPacket(dataPacket) } + if err != nil { + t.removeTailPos(key, reqId) + t.Sender.SendErrorResponse(reqId, err) + break + } if !keepRunning { - t.checkRemoveNoFollow(key, reqId) + removed := t.checkRemoveNoFollow(key, reqId) + if removed { + t.Sender.SendResponse(reqId, true) + } break } time.Sleep(10 * time.Millisecond) @@ -254,7 +259,6 @@ func (t *Tailer) tryStartRun_nolock(entry CmdWatchEntry, pos TailPos) { return } if pos.IsCurrent(entry) { - return } pos.Running = true @@ -344,6 +348,19 @@ func (t *Tailer) RemoveWatch(pk *packet.UntailCmdPacketType) { t.removeTailPos_nolock(pk.CK, pk.ReqId) } +func (t *Tailer) AddFileWatches_nolock(fileNames *base.CommandFileNames) error { + err := t.Watcher.Add(fileNames.PtyOutFile) + if err != nil { + return err + } + err = t.Watcher.Add(fileNames.RunnerOutFile) + if err != nil { + t.Watcher.Remove(fileNames.PtyOutFile) // best effort clean up + return err + } + return nil +} + func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { if err := getPacket.CK.Validate("getcmd"); err != nil { return err @@ -357,16 +374,7 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { key := getPacket.CK entry, foundEntry := t.WatchList[key] if !foundEntry { - // add watches, initialize entry - err := t.Watcher.Add(fileNames.PtyOutFile) - if err != nil { - return err - } - err = t.Watcher.Add(fileNames.RunnerOutFile) - if err != nil { - t.Watcher.Remove(fileNames.PtyOutFile) // best effort clean up - return err - } + // initialize entry, add watches entry = CmdWatchEntry{CmdKey: key} entry.fillFilePos(t.MHomeDir) } @@ -387,6 +395,14 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { pos.TailRunPos = max(0, entry.FileRunLen+pos.TailRunPos) // + because negative } entry.updateTailPos(pos.ReqId, pos) + if !pos.Follow && pos.IsCurrent(entry) { + // don't add to t.WatchList, don't t.AddFileWatches_nolock, send rpc response + go func() { t.Sender.SendResponse(getPacket.ReqId, true) }() + return nil + } + if !foundEntry { + t.AddFileWatches_nolock(fileNames) + } t.WatchList[key] = entry t.tryStartRun_nolock(entry, pos) return nil diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index dcdd5f09..b6142406 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -109,12 +109,10 @@ type CmdDataPacketType struct { PtyLen int64 `json:"ptylen"` RunPos int64 `json:"runpos"` RunLen int64 `json:"runlen"` - PtyData string `json:"ptydata"` + PtyData64 string `json:"ptydata64"` PtyDataLen int `json:"ptydatalen"` - RunData string `json:"rundata"` + RunData64 string `json:"rundata64"` RunDataLen int `json:"rundatalen"` - Error string `json:"error"` - NotFound bool `json:"notfound,omitempty"` } func (*CmdDataPacketType) GetType() string { @@ -125,8 +123,12 @@ func (p *CmdDataPacketType) GetResponseId() string { return p.RespId } -func MakeCmdDataPacket() *CmdDataPacketType { - return &CmdDataPacketType{Type: CmdDataPacketStr} +func (*CmdDataPacketType) GetResponseDone() bool { + return false +} + +func MakeCmdDataPacket(reqId string) *CmdDataPacketType { + return &CmdDataPacketType{Type: CmdDataPacketStr, RespId: reqId} } type PingPacketType struct { @@ -326,6 +328,10 @@ func (p *ResponsePacketType) GetResponseId() string { return p.RespId } +func (*ResponsePacketType) GetResponseDone() bool { + return true +} + func MakeErrorResponsePacket(reqId string, err error) *ResponsePacketType { return &ResponsePacketType{Type: ResponsePacketStr, RespId: reqId, Error: err.Error()} } @@ -421,8 +427,8 @@ func (p *CmdDonePacketType) GetCK() base.CommandKey { return p.CK } -func MakeCmdDonePacket() *CmdDonePacketType { - return &CmdDonePacketType{Type: CmdDonePacketStr} +func MakeCmdDonePacket(ck base.CommandKey) *CmdDonePacketType { + return &CmdDonePacketType{Type: CmdDonePacketStr, CK: ck} } type CmdStartPacketType struct { @@ -442,8 +448,12 @@ func (p *CmdStartPacketType) GetResponseId() string { return p.RespId } -func MakeCmdStartPacket() *CmdStartPacketType { - return &CmdStartPacketType{Type: CmdStartPacketStr} +func (*CmdStartPacketType) GetResponseDone() bool { + return true +} + +func MakeCmdStartPacket(reqId string) *CmdStartPacketType { + return &CmdStartPacketType{Type: CmdStartPacketStr, RespId: reqId} } type TermSize struct { @@ -534,6 +544,7 @@ type RpcPacketType interface { type RpcResponsePacketType interface { GetType() string GetResponseId() string + GetResponseDone() bool } type CommandPacketType interface { diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index 5d2c91e3..d09dac47 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -8,6 +8,7 @@ package packet import ( "bufio" + "context" "io" "strconv" "strings" @@ -17,9 +18,15 @@ import ( type PacketParser struct { Lock *sync.Mutex MainCh chan PacketType + RpcMap map[string]*RpcEntry Err error } +type RpcEntry struct { + ReqId string + RespCh chan RpcResponsePacketType +} + func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser { rtnParser := &PacketParser{ Lock: &sync.Mutex{}, @@ -46,6 +53,70 @@ func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser { return rtnParser } +// should have already registered rpc +func (p *PacketParser) WaitForResponse(ctx context.Context, reqId string) RpcResponsePacketType { + entry := p.getRpcEntry(reqId, false) + if entry == nil { + return nil + } + defer p.UnRegisterRpc(reqId) + select { + case resp := <-entry.RespCh: + return resp + case <-ctx.Done(): + return nil + } +} + +func (p *PacketParser) UnRegisterRpc(reqId string) { + p.Lock.Lock() + defer p.Lock.Unlock() + entry := p.RpcMap[reqId] + if entry != nil { + close(entry.RespCh) + delete(p.RpcMap, reqId) + } +} + +func (p *PacketParser) RegisterRpc(reqId string, queueSize int) chan RpcResponsePacketType { + p.Lock.Lock() + defer p.Lock.Unlock() + ch := make(chan RpcResponsePacketType, queueSize) + entry := &RpcEntry{ReqId: reqId, RespCh: ch} + p.RpcMap[reqId] = entry + return ch +} + +func (p *PacketParser) getRpcEntry(reqId string, remove bool) *RpcEntry { + p.Lock.Lock() + defer p.Lock.Unlock() + entry := p.RpcMap[reqId] + if entry != nil && remove { + delete(p.RpcMap, reqId) + close(entry.RespCh) + } + return entry +} + +func (p *PacketParser) trySendRpcResponse(respPk RpcResponsePacketType) bool { + p.Lock.Lock() + defer p.Lock.Unlock() + entry := p.RpcMap[respPk.GetResponseId()] + if entry == nil { + return false + } + // nonblocking send + select { + case entry.RespCh <- respPk: + default: + } + if respPk.GetResponseDone() { + delete(p.RpcMap, respPk.GetResponseId()) + close(entry.RespCh) + } + return true +} + func (p *PacketParser) GetErr() error { p.Lock.Lock() defer p.Lock.Unlock() @@ -108,6 +179,12 @@ func MakePacketParser(input io.Reader) *PacketParser { if pk.GetType() == PingPacketStr { continue } + if respPk, ok := pk.(RpcResponsePacketType); ok { + sent := parser.trySendRpcResponse(respPk) + if sent { + continue + } + } parser.MainCh <- pk } }() diff --git a/pkg/server/server.go b/pkg/server/server.go index fa736aa5..ab66c35e 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -161,6 +161,8 @@ func RunServer() (int, error) { if server.Debug { fmt.Printf("PK> %s\n", packet.AsString(pk)) } + + // run-start combo ok, runPacket := builder.ProcessPacket(pk) if server.Debug { fmt.Printf("PP> %s | %v\n", pk.GetType(), ok) @@ -179,6 +181,8 @@ func RunServer() (int, error) { server.Sender.SendPacket(startPk) continue } + + // command packet if cmdPk, ok := pk.(packet.CommandPacketType); ok { server.ProcessCommandPacket(cmdPk) continue diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 5981fdca..8eb2c46f 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -129,8 +129,8 @@ func (c *ShExecType) Close() { } } -func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { - startPacket := packet.MakeCmdStartPacket() +func (c *ShExecType) MakeCmdStartPacket(reqId string) *packet.CmdStartPacketType { + startPacket := packet.MakeCmdStartPacket(reqId) startPacket.Ts = time.Now().UnixMilli() startPacket.CK = c.CK startPacket.Pid = c.Cmd.Process.Pid @@ -848,21 +848,67 @@ func SetupSignalsForDetach() { }() } -func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) error { +func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { + // after Start(), any output/errors must go to DetachedOutput + // close stdin/stdout/stderr, but wait for cmdstart packet to get sent + nullFd, err := os.OpenFile("/dev/null", os.O_RDWR, 0) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot open /dev/null: %w", err)) + } + if nullFd != nil { + err := unix.Dup2(int(nullFd.Fd()), int(os.Stdin.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) + } + err = unix.Dup2(int(nullFd.Fd()), int(os.Stdout.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) + } + err = unix.Dup2(int(nullFd.Fd()), int(os.Stderr.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) + } + } + cmd.DetachedOutput.SendPacket(startPacket) + ptyOutFd, err := os.OpenFile(cmd.FileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot open ptyout file '%s': %w", cmd.FileNames.PtyOutFile, err)) + // don't return (command is already running) + } + go func() { + // copy pty output to .ptyout file + _, copyErr := io.Copy(ptyOutFd, cmd.CmdPty) + if copyErr != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("copying pty output to ptyout file: %w", copyErr)) + } + }() + go func() { + // copy .stdin fifo contents to pty input + copyFifoErr := MakeAndCopyStdinFifo(cmd.CmdPty, cmd.FileNames.StdinFifo) + if copyFifoErr != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("reading from stdin fifo: %w", copyFifoErr)) + } + }() + donePacket := cmd.WaitForCommand() + cmd.DetachedOutput.SendPacket(donePacket) + return +} + +func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, *packet.CmdStartPacketType, error) { fileNames, err := base.GetCommandFileNames(pk.CK) if err != nil { - return err + return nil, nil, err } ptyOutInfo, err := os.Stat(fileNames.PtyOutFile) if err == nil { // non-nil error will be caught by regular OpenFile below // must have size 0 if ptyOutInfo.Size() != 0 { - return fmt.Errorf("cmdkey '%s' was already used (ptyout len=%d)", pk.CK, ptyOutInfo.Size()) + return nil, nil, fmt.Errorf("cmdkey '%s' was already used (ptyout len=%d)", pk.CK, ptyOutInfo.Size()) } } cmdPty, cmdTty, err := pty.Open() if err != nil { - return fmt.Errorf("opening new pty: %w", err) + return nil, nil, fmt.Errorf("opening new pty: %w", err) } pty.Setsize(cmdPty, GetWinsize(pk)) defer func() { @@ -874,72 +920,26 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) e cmd.Detached = true cmd.RunnerOutFd, err = os.OpenFile(fileNames.RunnerOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err != nil { - return fmt.Errorf("cannot open runout file '%s': %w", fileNames.RunnerOutFile, err) - } - nullFd, err := os.OpenFile("/dev/null", os.O_RDWR, 0) - if err != nil { - return fmt.Errorf("cannot open /dev/null: %w", err) + return nil, nil, fmt.Errorf("cannot open runout file '%s': %w", fileNames.RunnerOutFile, err) } cmd.DetachedOutput = packet.MakePacketSender(cmd.RunnerOutFd) ecmd, err := MakeDetachedExecCmd(pk, cmdTty) if err != nil { - return err + return nil, nil, err } cmd.Cmd = ecmd SetupSignalsForDetach() err = ecmd.Start() if err != nil { - return fmt.Errorf("starting command: %w", err) + return nil, nil, fmt.Errorf("starting command: %w", err) } for _, fd := range ecmd.ExtraFiles { if fd != cmdTty { fd.Close() } } - // after Start(), any errors must go to DetachedOutput - // close stdin/stdout/stderr, but wait for cmdstart packet to get sent - startPacket := cmd.MakeCmdStartPacket() - go func() { - sender.SendPacket(startPacket) - sender.Close() - sender.WaitForDone() - fmt.Printf("sender done! start: %v\n", startPacket) - err = unix.Dup2(int(nullFd.Fd()), int(os.Stdin.Fd())) - if err != nil { - cmd.DetachedOutput.SendCmdError(pk.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) - } - err = unix.Dup2(int(nullFd.Fd()), int(os.Stdout.Fd())) - if err != nil { - cmd.DetachedOutput.SendCmdError(pk.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) - } - err = unix.Dup2(int(nullFd.Fd()), int(os.Stderr.Fd())) - if err != nil { - cmd.DetachedOutput.SendCmdError(pk.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) - } - cmd.DetachedOutput.SendPacket(startPacket) - }() - ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) - if err != nil { - cmd.DetachedOutput.SendCmdError(pk.CK, fmt.Errorf("cannot open ptyout file '%s': %w", fileNames.PtyOutFile, err)) - // don't return (command is already running) - } - go func() { - // copy pty output to .ptyout file - _, copyErr := io.Copy(ptyOutFd, cmdPty) - if copyErr != nil { - cmd.DetachedOutput.SendCmdError(pk.CK, fmt.Errorf("copying pty output to ptyout file: %w", copyErr)) - } - }() - go func() { - // copy .stdin fifo contents to pty input - copyFifoErr := MakeAndCopyStdinFifo(cmdPty, fileNames.StdinFifo) - if copyFifoErr != nil { - cmd.DetachedOutput.SendCmdError(pk.CK, fmt.Errorf("reading from stdin fifo: %w", copyFifoErr)) - } - }() - donePacket := cmd.WaitForCommand() - cmd.DetachedOutput.SendPacket(donePacket) - return nil + startPacket := cmd.MakeCmdStartPacket(pk.ReqId) + return cmd, startPacket, nil } func GetExitCode(err error) int { @@ -958,9 +958,8 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { endTs := time.Now() cmdDuration := endTs.Sub(c.StartTs) exitCode := GetExitCode(exitErr) - donePacket := packet.MakeCmdDonePacket() + donePacket := packet.MakeCmdDonePacket(c.CK) donePacket.Ts = endTs.UnixMilli() - donePacket.CK = c.CK donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) if c.FileNames != nil { From 0d585e5959035408c1a4f1eec89e83a9ef43535e Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 00:21:44 -0700 Subject: [PATCH 047/149] clean up --single detached mode --- pkg/packet/packet.go | 6 +++--- pkg/server/server.go | 2 +- pkg/shexec/shexec.go | 37 ++++++++++++++++++------------------- 3 files changed, 22 insertions(+), 23 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index b6142406..0c1527f2 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -316,8 +316,8 @@ type ResponsePacketType struct { Type string `json:"type"` RespId string `json:"respid"` Success bool `json:"success"` - Error string `json:"error"` - Data interface{} `json:"data"` + Error string `json:"error,omitempty"` + Data interface{} `json:"data,omitempty"` } func (*ResponsePacketType) GetType() string { @@ -476,7 +476,7 @@ type RunDataType struct { type RunPacketType struct { Type string `json:"type"` - ReqId string `json:"packetid"` + ReqId string `json:"reqid"` CK base.CommandKey `json:"ck"` Command string `json:"command"` Cwd string `json:"cwd,omitempty"` diff --git a/pkg/server/server.go b/pkg/server/server.go index ab66c35e..668c0172 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -114,7 +114,7 @@ func (m *MServer) RemoveFdContext(ck base.CommandKey) { func (m *MServer) runCommand(runPacket *packet.RunPacketType) { if err := runPacket.CK.Validate("packet"); err != nil { - m.Sender.SendResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } fdContext := m.MakeServerFdContext(runPacket.CK) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 8eb2c46f..44ed76e1 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -850,33 +850,30 @@ func SetupSignalsForDetach() { func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { // after Start(), any output/errors must go to DetachedOutput - // close stdin/stdout/stderr, but wait for cmdstart packet to get sent - nullFd, err := os.OpenFile("/dev/null", os.O_RDWR, 0) - if err != nil { - cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot open /dev/null: %w", err)) - } - if nullFd != nil { - err := unix.Dup2(int(nullFd.Fd()), int(os.Stdin.Fd())) - if err != nil { - cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) - } - err = unix.Dup2(int(nullFd.Fd()), int(os.Stdout.Fd())) - if err != nil { - cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) - } - err = unix.Dup2(int(nullFd.Fd()), int(os.Stderr.Fd())) - if err != nil { - cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to /dev/null: %w", err)) - } - } + // close stdin, redirect stdout/stderr to /dev/null, but wait for cmdstart packet to get sent cmd.DetachedOutput.SendPacket(startPacket) + err := os.Stdin.Close() + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot close stdin: %w", err)) + } + err = unix.Dup2(int(cmd.RunnerOutFd.Fd()), int(os.Stdout.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to runout: %w", err)) + } + err = unix.Dup2(int(cmd.RunnerOutFd.Fd()), int(os.Stderr.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to runout: %w", err)) + } ptyOutFd, err := os.OpenFile(cmd.FileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err != nil { cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot open ptyout file '%s': %w", cmd.FileNames.PtyOutFile, err)) // don't return (command is already running) } + ptyCopyDone := make(chan bool) go func() { // copy pty output to .ptyout file + defer close(ptyCopyDone) + defer ptyOutFd.Close() _, copyErr := io.Copy(ptyOutFd, cmd.CmdPty) if copyErr != nil { cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("copying pty output to ptyout file: %w", copyErr)) @@ -891,6 +888,8 @@ func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { }() donePacket := cmd.WaitForCommand() cmd.DetachedOutput.SendPacket(donePacket) + <-ptyCopyDone + cmd.Close() return } From 9aa684882bc42634ee8ba955199f506f0ed2c969 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 00:25:59 -0700 Subject: [PATCH 048/149] don't send done packet when detached --- pkg/server/server.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 668c0172..a91fb0f5 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -121,7 +121,7 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { go func() { defer m.RemoveFdContext(runPacket.CK) donePk, err := shexec.RunClientSSHCommandAndWait(runPacket, fdContext, shexec.SSHOpts{}, m, m.Debug) - if donePk != nil { + if donePk != nil && !runPacket.Detached { m.Sender.SendPacket(donePk) } if err != nil { From 95f11fb4187118c1c191ee00e4fc19a958320752 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 11:21:15 -0700 Subject: [PATCH 049/149] checkpoint, getting tty output working in non-detached mode --- main-mshell.go | 4 ++ pkg/packet/packet.go | 1 + pkg/shexec/shexec.go | 89 ++++++++++++++++++++++++++++++-------------- 3 files changed, 66 insertions(+), 28 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 351130a0..dac5d69a 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -304,6 +304,10 @@ func parseClientOpts() (*shexec.ClientOpts, error) { opts.Detach = true continue } + if argStr == "--pty" { + opts.UsePty = true + continue + } if argStr == "--debug" { opts.Debug = true continue diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 0c1527f2..924ba6a1 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -481,6 +481,7 @@ type RunPacketType struct { Command string `json:"command"` Cwd string `json:"cwd,omitempty"` Env map[string]string `json:"env,omitempty"` + UsePty bool `json:"usepty,omitempty"` TermSize *TermSize `json:"termsize,omitempty"` Fds []RemoteFd `json:"fds,omitempty"` RunData []RunDataType `json:"rundata,omitempty"` diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 44ed76e1..96a7d3e1 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -196,9 +196,10 @@ func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, if pk.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(pk.Cwd) } - if !HasDupStdin(pk.Fds) { - ecmd.Stdin = cmdTty + if HasDupStdin(pk.Fds) { + return nil, fmt.Errorf("cannot detach command with dup stdin") } + ecmd.Stdin = cmdTty ecmd.Stdout = cmdTty ecmd.Stderr = cmdTty ecmd.SysProcAttr = &syscall.SysProcAttr{ @@ -206,15 +207,8 @@ func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, Setctty: true, } extraFiles := make([]*os.File, 0, MaxFdNum+1) - for _, rfd := range pk.Fds { - if rfd.FdNum >= len(extraFiles) { - extraFiles = extraFiles[:rfd.FdNum+1] - } - if rfd.Read && rfd.DupStdin { - extraFiles[rfd.FdNum] = cmdTty - continue - } - return nil, fmt.Errorf("invalid fd %d passed to detached command", rfd.FdNum) + if len(pk.Fds) > 0 { + return nil, fmt.Errorf("invalid fd %d passed to detached command", pk.Fds[0].FdNum) } for _, runData := range pk.RunData { if runData.FdNum >= len(extraFiles) { @@ -276,7 +270,10 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { if rfd.Write { return fmt.Errorf("cannot detach command with writable remote files fd=%d", rfd.FdNum) } - if rfd.Read && !rfd.DupStdin { + if rfd.Read && rfd.DupStdin { + return fmt.Errorf("cannot detach command with dup stdin fd=%d", rfd.FdNum) + } + if rfd.Read { return fmt.Errorf("cannot detach command with readable remote files fd=%d", rfd.FdNum) } } @@ -306,6 +303,9 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { return fmt.Errorf("rundata length mismatch, fd=%d, datalen=%d, expected=%d", runData.FdNum, len(runData.Data), runData.DataLen) } } + if pk.UsePty && HasDupStdin(pk.Fds) { + return fmt.Errorf("cannot use pty with command that has dup stdin") + } return nil } @@ -347,6 +347,7 @@ type ClientOpts struct { SudoWithPass bool SudoPw string Detach bool + UsePty bool } func (opts SSHOpts) MakeSSHInstallCmd() (*exec.Cmd, error) { @@ -416,6 +417,7 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket.Detached = opts.Detach runPacket.Cwd = opts.Cwd runPacket.Fds = opts.Fds + runPacket.UsePty = opts.UsePty if !opts.Sudo { // normal, non-sudo command runPacket.Command = fmt.Sprintf(RunCommandFmt, opts.Command) @@ -777,20 +779,51 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmd.Close() return nil, err } - cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) - if err != nil { - cmd.Close() - return nil, err + var cmdPty *os.File + var cmdTty *os.File + if pk.UsePty { + cmdPty, cmdTty, err = pty.Open() + if err != nil { + return nil, fmt.Errorf("opening new pty: %w", err) + } + pty.Setsize(cmdPty, GetWinsize(pk)) + defer func() { + cmdTty.Close() + }() + cmd.CmdPty = cmdPty } - cmd.Cmd.Stdout, err = cmd.Multiplexer.MakeReaderPipe(1) - if err != nil { - cmd.Close() - return nil, err - } - cmd.Cmd.Stderr, err = cmd.Multiplexer.MakeReaderPipe(2) - if err != nil { - cmd.Close() - return nil, err + if cmdTty != nil { + cmd.Cmd.Stdin = cmdTty + cmd.Cmd.Stdout = cmdTty + cmd.Cmd.Stderr = cmdTty + cmd.Cmd.SysProcAttr = &syscall.SysProcAttr{ + Setsid: true, + Setctty: true, + } + cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false) + cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false) + nullFd, err := os.Open("/dev/null") + if err != nil { + cmd.Close() + return nil, fmt.Errorf("cannot open /dev/null: %w", err) + } + cmd.Multiplexer.MakeRawFdReader(2, nullFd, true) + } else { + cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) + if err != nil { + cmd.Close() + return nil, err + } + cmd.Cmd.Stdout, err = cmd.Multiplexer.MakeReaderPipe(1) + if err != nil { + cmd.Close() + return nil, err + } + cmd.Cmd.Stderr, err = cmd.Multiplexer.MakeReaderPipe(2) + if err != nil { + cmd.Close() + return nil, err + } } extraFiles := make([]*os.File, 0, MaxFdNum+1) for _, runData := range pk.RunData { @@ -898,11 +931,11 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( if err != nil { return nil, nil, err } - ptyOutInfo, err := os.Stat(fileNames.PtyOutFile) + runOutInfo, err := os.Stat(fileNames.RunnerOutFile) if err == nil { // non-nil error will be caught by regular OpenFile below // must have size 0 - if ptyOutInfo.Size() != 0 { - return nil, nil, fmt.Errorf("cmdkey '%s' was already used (ptyout len=%d)", pk.CK, ptyOutInfo.Size()) + if runOutInfo.Size() != 0 { + return nil, nil, fmt.Errorf("cmdkey '%s' was already used (runout len=%d)", pk.CK, runOutInfo.Size()) } } cmdPty, cmdTty, err := pty.Open() From a1b82349544d9aa7ae6b9025e13e9ecd50f46192 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 12:16:37 -0700 Subject: [PATCH 050/149] fix tty TERM for ssh connections when usepty is set. also ignore pty read errors --- pkg/mpio/bufreader.go | 8 +++++++- pkg/mpio/mpio.go | 6 +++--- pkg/shexec/shexec.go | 13 +++++++------ 3 files changed, 17 insertions(+), 10 deletions(-) diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index bcda317e..68af4cb5 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -21,9 +21,10 @@ type FdReader struct { BufSize int Closed bool ShouldCloseFd bool + IsPty bool } -func MakeFdReader(m *Multiplexer, fd io.ReadCloser, fdNum int, shouldCloseFd bool) *FdReader { +func MakeFdReader(m *Multiplexer, fd io.ReadCloser, fdNum int, shouldCloseFd bool, isPty bool) *FdReader { fr := &FdReader{ CVar: sync.NewCond(&sync.Mutex{}), M: m, @@ -31,6 +32,7 @@ func MakeFdReader(m *Multiplexer, fd io.ReadCloser, fdNum int, shouldCloseFd boo Fd: fd, BufSize: 0, ShouldCloseFd: shouldCloseFd, + IsPty: isPty, } return fr } @@ -136,6 +138,10 @@ func (r *FdReader) ReadLoop(wg *sync.WaitGroup) { } } if err != nil { + if r.IsPty { + r.WriteWait(nil, true) + return + } errPk := r.M.makeDataPacket(r.FdNum, nil, err) r.M.sendPacket(errPk) return diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 4fcea20c..f16a411c 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -89,7 +89,7 @@ func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { } m.Lock.Lock() defer m.Lock.Unlock() - m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum, true) + m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum, true, false) m.CloseAfterStart = append(m.CloseAfterStart, pw) return pw, nil } @@ -125,10 +125,10 @@ func (m *Multiplexer) MakeStaticWriterPipe(fdNum int, data []byte) (*os.File, er return pr, nil } -func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose bool) { +func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose bool, isPty bool) { m.Lock.Lock() defer m.Lock.Unlock() - m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose) + m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose, isPty) } func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd io.WriteCloser, shouldClose bool) { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 96a7d3e1..bf935e7b 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -150,7 +150,7 @@ func UpdateCmdEnv(cmd *exec.Cmd, envVars map[string]string) { if len(envVars) == 0 { return } - if cmd.Env != nil { + if cmd.Env == nil { cmd.Env = os.Environ() } found := make(map[string]bool) @@ -631,18 +631,18 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon return nil, fmt.Errorf("creating stderr pipe: %v", err) } if !HasDupStdin(runPacket.Fds) { - cmd.Multiplexer.MakeRawFdReader(0, fdContext.GetReader(0), false) + cmd.Multiplexer.MakeRawFdReader(0, fdContext.GetReader(0), false, false) } cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false) cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), false) for _, rfd := range runPacket.Fds { if rfd.Read && rfd.DupStdin { - cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fdContext.GetReader(0), false) + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fdContext.GetReader(0), false, false) continue } if rfd.Read { fd := fdContext.GetReader(rfd.FdNum) - cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, false) + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, false, false) } else if rfd.Write { fd := fdContext.GetWriter(rfd.FdNum) cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true) @@ -791,6 +791,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmdTty.Close() }() cmd.CmdPty = cmdPty + UpdateCmdEnv(cmd.Cmd, map[string]string{"TERM": "xterm-256color"}) } if cmdTty != nil { cmd.Cmd.Stdin = cmdTty @@ -801,13 +802,13 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S Setctty: true, } cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false) - cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false) + cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false, true) nullFd, err := os.Open("/dev/null") if err != nil { cmd.Close() return nil, fmt.Errorf("cannot open /dev/null: %w", err) } - cmd.Multiplexer.MakeRawFdReader(2, nullFd, true) + cmd.Multiplexer.MakeRawFdReader(2, nullFd, true, false) } else { cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) if err != nil { From 56e1ddf8e6abb60a8a7906435743529b9b006e4e Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 14:06:58 -0700 Subject: [PATCH 051/149] run packet opt to set term type --- pkg/packet/packet.go | 9 +++++---- pkg/shexec/shexec.go | 45 +++++++++++++++++++++++++++++++++++++------- 2 files changed, 43 insertions(+), 11 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 924ba6a1..a12092b2 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -456,9 +456,10 @@ func MakeCmdStartPacket(reqId string) *CmdStartPacketType { return &CmdStartPacketType{Type: CmdStartPacketStr, RespId: reqId} } -type TermSize struct { - Rows int `json:"rows"` - Cols int `json:"cols"` +type TermOpts struct { + Rows int `json:"rows"` + Cols int `json:"cols"` + Term string `json:"term"` } type RemoteFd struct { @@ -482,7 +483,7 @@ type RunPacketType struct { Cwd string `json:"cwd,omitempty"` Env map[string]string `json:"env,omitempty"` UsePty bool `json:"usepty,omitempty"` - TermSize *TermSize `json:"termsize,omitempty"` + TermOpts *TermOpts `json:"termopts,omitempty"` Fds []RemoteFd `json:"fds,omitempty"` RunData []RunDataType `json:"rundata,omitempty"` Detached bool `json:"detached,omitempty"` diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index bf935e7b..6b549f56 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -33,6 +33,7 @@ const MaxRows = 1024 const MaxCols = 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 +const DefaultTermType = "xterm-256color" const ClientCommand = ` PATH=$PATH:~/.mshell; @@ -193,6 +194,7 @@ func MakeSimpleStaticWriterPipe(data []byte) (*os.File, error) { func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, error) { ecmd := exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(ecmd, pk.Env) + UpdateCmdEnv(ecmd, map[string]string{"TERM": getTermType(pk)}) if pk.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(pk.Cwd) } @@ -312,12 +314,12 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { func GetWinsize(p *packet.RunPacketType) *pty.Winsize { rows := DefaultRows cols := DefaultCols - if p.TermSize != nil { - if p.TermSize.Rows > 0 && p.TermSize.Rows <= MaxRows { - rows = p.TermSize.Rows + if p.TermOpts != nil { + if p.TermOpts.Rows > 0 && p.TermOpts.Rows <= MaxRows { + rows = p.TermOpts.Rows } - if p.TermSize.Cols > 0 && p.TermSize.Cols <= MaxCols { - cols = p.TermSize.Cols + if p.TermOpts.Cols > 0 && p.TermOpts.Cols <= MaxCols { + cols = p.TermOpts.Cols } } return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} @@ -412,12 +414,33 @@ func (opts SSHOpts) MakeMShellSSHOpts() string { return strings.Join(moreSSHOpts, " ") } +func GetTerminalSize() (int, int, error) { + fd, err := os.Open("/dev/tty") + if err != nil { + return 0, 0, err + } + defer fd.Close() + return pty.Getsize(fd) +} + func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket := packet.MakeRunPacket() runPacket.Detached = opts.Detach runPacket.Cwd = opts.Cwd runPacket.Fds = opts.Fds - runPacket.UsePty = opts.UsePty + if opts.UsePty { + runPacket.UsePty = true + runPacket.TermOpts = &packet.TermOpts{} + rows, cols, err := GetTerminalSize() + if err == nil { + runPacket.TermOpts.Rows = rows + runPacket.TermOpts.Cols = cols + } + term := os.Getenv("TERM") + if term != "" { + runPacket.TermOpts.Term = term + } + } if !opts.Sudo { // normal, non-sudo command runPacket.Command = fmt.Sprintf(RunCommandFmt, opts.Command) @@ -767,6 +790,14 @@ func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sen sender.SendPacket(donePacket) } +func getTermType(pk *packet.RunPacketType) string { + termType := DefaultTermType + if pk.TermOpts != nil && pk.TermOpts.Term != "" { + termType = pk.TermOpts.Term + } + return termType +} + func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { cmd := MakeShExec(pk.CK, nil) cmd.Cmd = exec.Command("bash", "-c", pk.Command) @@ -791,7 +822,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S cmdTty.Close() }() cmd.CmdPty = cmdPty - UpdateCmdEnv(cmd.Cmd, map[string]string{"TERM": "xterm-256color"}) + UpdateCmdEnv(cmd.Cmd, map[string]string{"TERM": getTermType(pk)}) } if cmdTty != nil { cmd.Cmd.Stdin = cmdTty From 51df0479ffe19dabe4f7267c7c16587bb92608f0 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 14:35:27 -0700 Subject: [PATCH 052/149] sendpacket with context, initialize rpcmap --- pkg/packet/packet.go | 14 ++++++++++++++ pkg/packet/parser.go | 8 +++++++- 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index a12092b2..fe4c4990 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -8,6 +8,7 @@ package packet import ( "bytes" + "context" "encoding/base64" "encoding/json" "fmt" @@ -693,6 +694,19 @@ func (sender *PacketSender) checkStatus() error { return nil } +func (sender *PacketSender) SendPacketCtx(ctx context.Context, pk PacketType) error { + err := sender.checkStatus() + if err != nil { + return err + } + select { + case sender.SendCh <- pk: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + func (sender *PacketSender) SendPacket(pk PacketType) error { err := sender.checkStatus() if err != nil { diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index d09dac47..d405d3b0 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -31,6 +31,7 @@ func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser { rtnParser := &PacketParser{ Lock: &sync.Mutex{}, MainCh: make(chan PacketType), + RpcMap: make(map[string]*RpcEntry), } var wg sync.WaitGroup wg.Add(2) @@ -78,7 +79,11 @@ func (p *PacketParser) UnRegisterRpc(reqId string) { } } -func (p *PacketParser) RegisterRpc(reqId string, queueSize int) chan RpcResponsePacketType { +func (p *PacketParser) RegisterRpc(reqId string) chan RpcResponsePacketType { + return p.RegisterRpcSz(reqId, 2) +} + +func (p *PacketParser) RegisterRpcSz(reqId string, queueSize int) chan RpcResponsePacketType { p.Lock.Lock() defer p.Lock.Unlock() ch := make(chan RpcResponsePacketType, queueSize) @@ -135,6 +140,7 @@ func MakePacketParser(input io.Reader) *PacketParser { parser := &PacketParser{ Lock: &sync.Mutex{}, MainCh: make(chan PacketType), + RpcMap: make(map[string]*RpcEntry), } bufReader := bufio.NewReader(input) go func() { From 1b69bb0ac8d6d40b30d48086d5836e2f8402451f Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 17:16:45 -0700 Subject: [PATCH 053/149] fix mshell server to just blindly proxy mshell single command input/output, better simpler code --- pkg/server/server.go | 140 ++++++++++--------------------------------- pkg/shexec/client.go | 121 +++++++++++++++++++++++++++++++++++++ 2 files changed, 154 insertions(+), 107 deletions(-) create mode 100644 pkg/shexec/client.go diff --git a/pkg/server/server.go b/pkg/server/server.go index a91fb0f5..91637778 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -8,22 +8,21 @@ package server import ( "fmt" - "io" "os" "sync" "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" "github.com/scripthaus-dev/mshell/pkg/shexec" ) +// TODO create unblockable packet-sender (backed by an array) for clientproc type MServer struct { - Lock *sync.Mutex - MainInput *packet.PacketParser - Sender *packet.PacketSender - FdContextMap map[base.CommandKey]*serverFdContext - Debug bool + Lock *sync.Mutex + MainInput *packet.PacketParser + Sender *packet.PacketSender + ClientMap map[base.CommandKey]*shexec.ClientProc + Debug bool } func (m *MServer) Close() { @@ -31,43 +30,6 @@ func (m *MServer) Close() { m.Sender.WaitForDone() } -type serverFdContext struct { - M *MServer - Lock *sync.Mutex - Sender *packet.PacketSender - CK base.CommandKey - Readers map[int]*mpio.PacketReader -} - -func (c *serverFdContext) processDataPacket(pk *packet.DataPacketType) { - c.Lock.Lock() - reader := c.Readers[pk.FdNum] - c.Lock.Unlock() - if reader == nil { - ackPacket := packet.MakeDataAckPacket() - ackPacket.CK = c.CK - ackPacket.FdNum = pk.FdNum - ackPacket.Error = "write to closed file (no fd)" - c.M.Sender.SendPacket(ackPacket) - return - } - reader.AddData(pk) -} - -func (m *MServer) MakeServerFdContext(ck base.CommandKey) *serverFdContext { - m.Lock.Lock() - defer m.Lock.Unlock() - rtn := &serverFdContext{ - M: m, - Lock: &sync.Mutex{}, - Sender: m.Sender, - CK: ck, - Readers: make(map[int]*mpio.PacketReader), - } - m.FdContextMap[ck] = rtn - return rtn -} - func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { ck := pk.GetCK() if ck == "" { @@ -75,41 +37,14 @@ func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { return } m.Lock.Lock() - fdContext := m.FdContextMap[ck] + cproc := m.ClientMap[ck] m.Lock.Unlock() - if fdContext == nil { - m.Sender.SendCmdError(ck, fmt.Errorf("no server context for ck '%s'", ck)) + if cproc == nil { + m.Sender.SendCmdError(ck, fmt.Errorf("no client proc for ck '%s'", ck)) return } - if pk.GetType() == packet.DataPacketStr { - dataPacket := pk.(*packet.DataPacketType) - fdContext.processDataPacket(dataPacket) - return - } else if pk.GetType() == packet.DataAckPacketStr { - m.Sender.SendPacket(pk) - return - } else { - m.Sender.SendCmdError(ck, fmt.Errorf("invalid packet '%s' received", packet.AsExtType(pk))) - return - } -} - -func (c *serverFdContext) GetWriter(fdNum int) io.WriteCloser { - return mpio.MakePacketWriter(fdNum, c.Sender, c.CK) -} - -func (c *serverFdContext) GetReader(fdNum int) io.ReadCloser { - c.Lock.Lock() - defer c.Lock.Unlock() - reader := mpio.MakePacketReader(fdNum) - c.Readers[fdNum] = reader - return reader -} - -func (m *MServer) RemoveFdContext(ck base.CommandKey) { - m.Lock.Lock() - defer m.Lock.Unlock() - delete(m.FdContextMap, ck) + cproc.Input.SendPacket(pk) + return } func (m *MServer) runCommand(runPacket *packet.RunPacketType) { @@ -117,32 +52,37 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } - fdContext := m.MakeServerFdContext(runPacket.CK) + cproc, err := shexec.MakeClientProc(runPacket.CK) + if err != nil { + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("starting mshell client: %s", err)) + return + } + fmt.Printf("client start: %v\n", runPacket.CK) + m.Lock.Lock() + m.ClientMap[runPacket.CK] = cproc + m.Lock.Unlock() go func() { - defer m.RemoveFdContext(runPacket.CK) - donePk, err := shexec.RunClientSSHCommandAndWait(runPacket, fdContext, shexec.SSHOpts{}, m, m.Debug) - if donePk != nil && !runPacket.Detached { - m.Sender.SendPacket(donePk) - } - if err != nil { - m.Sender.SendErrorResponse(runPacket.ReqId, err) - } + defer func() { + m.Lock.Lock() + delete(m.ClientMap, runPacket.CK) + m.Lock.Unlock() + cproc.Close() + fmt.Printf("client done: %v\n", runPacket.CK) + }() + shexec.SendRunPacketAndRunData(cproc.Input, runPacket) + cproc.ProxyOutput(m.Sender) }() } -func (m *MServer) UnknownPacket(pk packet.PacketType) { - m.Sender.SendPacket(pk) -} - func RunServer() (int, error) { debug := false if len(os.Args) >= 3 && os.Args[2] == "--debug" { debug = true } server := &MServer{ - Lock: &sync.Mutex{}, - FdContextMap: make(map[base.CommandKey]*serverFdContext), - Debug: debug, + Lock: &sync.Mutex{}, + ClientMap: make(map[base.CommandKey]*shexec.ClientProc), + Debug: debug, } if debug { packet.GlobalDebug = true @@ -161,12 +101,7 @@ func RunServer() (int, error) { if server.Debug { fmt.Printf("PK> %s\n", packet.AsString(pk)) } - - // run-start combo ok, runPacket := builder.ProcessPacket(pk) - if server.Debug { - fmt.Printf("PP> %s | %v\n", pk.GetType(), ok) - } if ok { if runPacket != nil { server.runCommand(runPacket) @@ -174,20 +109,11 @@ func RunServer() (int, error) { } continue } - if startPk, ok := pk.(*packet.CmdStartPacketType); ok { - if server.Debug { - fmt.Printf("START> %v", startPk) - } - server.Sender.SendPacket(startPk) - continue - } - - // command packet if cmdPk, ok := pk.(packet.CommandPacketType); ok { server.ProcessCommandPacket(cmdPk) continue } - server.Sender.SendMessage(fmt.Sprintf("invalid packet '%s' sent to mshell", packet.AsString(pk))) + server.Sender.SendMessage(fmt.Sprintf("invalid packet '%s' sent to mshell server", packet.AsString(pk))) continue } return 0, nil diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go new file mode 100644 index 00000000..f322ed38 --- /dev/null +++ b/pkg/shexec/client.go @@ -0,0 +1,121 @@ +package shexec + +import ( + "fmt" + "io" + "os/exec" + "time" + + "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/packet" +) + +type ClientProc struct { + Cmd *exec.Cmd + CK base.CommandKey + StartTs time.Time + StdinWriter io.WriteCloser + StdoutReader io.ReadCloser + StderrReader io.ReadCloser + Input *packet.PacketSender + Output *packet.PacketParser +} + +func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { + ecmd, err := SSHOpts{}.MakeMShellSingleCmd() + if err != nil { + return nil, err + } + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return nil, fmt.Errorf("creating stdin pipe: %v", err) + } + stdoutReader, err := ecmd.StdoutPipe() + if err != nil { + return nil, fmt.Errorf("creating stdout pipe: %v", err) + } + stderrReader, err := ecmd.StderrPipe() + if err != nil { + return nil, fmt.Errorf("creating stderr pipe: %v", err) + } + startTs := time.Now() + err = ecmd.Start() + if err != nil { + return nil, fmt.Errorf("running local client: %w", err) + } + sender := packet.MakePacketSender(inputWriter) + stdoutPacketParser := packet.MakePacketParser(stdoutReader) + stderrPacketParser := packet.MakePacketParser(stderrReader) + packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) + cproc := &ClientProc{ + Cmd: ecmd, + CK: ck, + StartTs: startTs, + StdinWriter: inputWriter, + StdoutReader: stdoutReader, + StderrReader: stderrReader, + Input: sender, + Output: packetParser, + } + versionOk := false + for pk := range packetParser.MainCh { + if pk.GetType() != packet.InitPacketStr { + cproc.Close() + return nil, fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) + } + initPk := pk.(*packet.InitPacketType) + if initPk.NotFound { + cproc.Close() + return nil, fmt.Errorf("mshell command not found on local server") + } + if initPk.Version != base.MShellVersion { + cproc.Close() + return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) + } + versionOk = true + break + } + if !versionOk { + cproc.Close() + return nil, fmt.Errorf("no init packet received from mshell client") + } + return cproc, nil +} + +func (cproc *ClientProc) Close() { + if cproc.Input != nil { + cproc.Input.Close() + } + if cproc.StdinWriter != nil { + cproc.StdinWriter.Close() + } + if cproc.StdoutReader != nil { + cproc.StdoutReader.Close() + } + if cproc.StderrReader != nil { + cproc.StderrReader.Close() + } + if cproc.Cmd != nil { + cproc.Cmd.Process.Kill() + } +} + +func (cproc *ClientProc) ProxyOutput(sender *packet.PacketSender) { + sentDonePk := false + for pk := range cproc.Output.MainCh { + if pk.GetType() == packet.CmdDonePacketStr { + sentDonePk = true + } + sender.SendPacket(pk) + } + exitErr := cproc.Cmd.Wait() + if !sentDonePk { + endTs := time.Now() + cmdDuration := endTs.Sub(cproc.StartTs) + donePacket := packet.MakeCmdDonePacket(cproc.CK) + donePacket.Ts = endTs.UnixMilli() + donePacket.ExitCode = GetExitCode(exitErr) + donePacket.DurationMs = int64(cmdDuration / time.Millisecond) + sender.SendPacket(donePacket) + } +} From 353605f815916cf5ff1fd51c649101218f2b8305 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 18:59:46 -0700 Subject: [PATCH 054/149] bug fixes and updates for running server with scripthaus --- pkg/server/server.go | 14 +++++++++----- pkg/shexec/client.go | 18 ++++++------------ pkg/shexec/shexec.go | 32 ++++++++++++++++++++++++++------ 3 files changed, 41 insertions(+), 23 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 91637778..40ab29cc 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -7,6 +7,7 @@ package server import ( + "context" "fmt" "os" "sync" @@ -52,12 +53,16 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } - cproc, err := shexec.MakeClientProc(runPacket.CK) + ecmd, err := shexec.SSHOpts{}.MakeMShellSingleCmd() + if err != nil { + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) + return + } + cproc, err := shexec.MakeClientProc(ecmd) if err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("starting mshell client: %s", err)) return } - fmt.Printf("client start: %v\n", runPacket.CK) m.Lock.Lock() m.ClientMap[runPacket.CK] = cproc m.Lock.Unlock() @@ -67,10 +72,9 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { delete(m.ClientMap, runPacket.CK) m.Lock.Unlock() cproc.Close() - fmt.Printf("client done: %v\n", runPacket.CK) }() - shexec.SendRunPacketAndRunData(cproc.Input, runPacket) - cproc.ProxyOutput(m.Sender) + shexec.SendRunPacketAndRunData(context.Background(), cproc.Input, runPacket) + cproc.ProxySingleOutput(runPacket.CK, m.Sender) }() } diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index f322ed38..f95b252f 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -12,7 +12,7 @@ import ( type ClientProc struct { Cmd *exec.Cmd - CK base.CommandKey + InitPk *packet.InitPacketType StartTs time.Time StdinWriter io.WriteCloser StdoutReader io.ReadCloser @@ -21,11 +21,7 @@ type ClientProc struct { Output *packet.PacketParser } -func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { - ecmd, err := SSHOpts{}.MakeMShellSingleCmd() - if err != nil { - return nil, err - } +func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, error) { inputWriter, err := ecmd.StdinPipe() if err != nil { return nil, fmt.Errorf("creating stdin pipe: %v", err) @@ -49,7 +45,6 @@ func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) cproc := &ClientProc{ Cmd: ecmd, - CK: ck, StartTs: startTs, StdinWriter: inputWriter, StdoutReader: stdoutReader, @@ -57,7 +52,6 @@ func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { Input: sender, Output: packetParser, } - versionOk := false for pk := range packetParser.MainCh { if pk.GetType() != packet.InitPacketStr { cproc.Close() @@ -72,10 +66,10 @@ func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { cproc.Close() return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) } - versionOk = true + cproc.InitPk = initPk break } - if !versionOk { + if cproc.InitPk == nil { cproc.Close() return nil, fmt.Errorf("no init packet received from mshell client") } @@ -100,7 +94,7 @@ func (cproc *ClientProc) Close() { } } -func (cproc *ClientProc) ProxyOutput(sender *packet.PacketSender) { +func (cproc *ClientProc) ProxySingleOutput(ck base.CommandKey, sender *packet.PacketSender) { sentDonePk := false for pk := range cproc.Output.MainCh { if pk.GetType() == packet.CmdDonePacketStr { @@ -112,7 +106,7 @@ func (cproc *ClientProc) ProxyOutput(sender *packet.PacketSender) { if !sentDonePk { endTs := time.Now() cmdDuration := endTs.Sub(cproc.StartTs) - donePacket := packet.MakeCmdDonePacket(cproc.CK) + donePacket := packet.MakeCmdDonePacket(ck) donePacket.Ts = endTs.UnixMilli() donePacket.ExitCode = GetExitCode(exitErr) donePacket.DurationMs = int64(cmdDuration / time.Millisecond) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 6b549f56..03d0ba6e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -7,6 +7,7 @@ package shexec import ( + "context" "encoding/base64" "fmt" "io" @@ -359,6 +360,15 @@ func (opts SSHOpts) MakeSSHInstallCmd() (*exec.Cmd, error) { return opts.MakeSSHExecCmd(InstallCommand), nil } +func (opts SSHOpts) MakeMShellServerCmd() (*exec.Cmd, error) { + msPath, err := base.GetMShellPath() + if err != nil { + return nil, err + } + ecmd := exec.Command(msPath, "--server") + return ecmd, nil +} + func (opts SSHOpts) MakeMShellSingleCmd() (*exec.Cmd, error) { if opts.SSHHost == "" { execFile, err := os.Executable() @@ -716,7 +726,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon if !versionOk { return nil, fmt.Errorf("did not receive version from remote mshell") } - SendRunPacketAndRunData(sender, runPacket) + SendRunPacketAndRunData(context.Background(), sender, runPacket) if debug { cmd.Multiplexer.Debug = true } @@ -735,10 +745,13 @@ func min(v1 int, v2 int) int { return v2 } -func SendRunPacketAndRunData(sender *packet.PacketSender, runPacket *packet.RunPacketType) { - sender.SendPacket(runPacket) +func SendRunPacketAndRunData(ctx context.Context, sender *packet.PacketSender, runPacket *packet.RunPacketType) error { + err := sender.SendPacketCtx(ctx, runPacket) + if err != nil { + return err + } if len(runPacket.RunData) == 0 { - return + return nil } for _, runData := range runPacket.RunData { sendBuf := runData.Data @@ -751,10 +764,17 @@ func SendRunPacketAndRunData(sender *packet.PacketSender, runPacket *packet.RunP dataPk.Data64 = base64.StdEncoding.EncodeToString(chunk) dataPk.Eof = (len(chunk) == len(sendBuf)) sendBuf = sendBuf[chunkSize:] - sender.SendPacket(dataPk) + err = sender.SendPacketCtx(ctx, dataPk) + if err != nil { + return err + } } } - sender.SendPacket(packet.MakeDataEndPacket(runPacket.CK)) + err = sender.SendPacketCtx(ctx, packet.MakeDataEndPacket(runPacket.CK)) + if err != nil { + return err + } + return nil } func DetectGoArch(uname string) (string, string, error) { From 2652a3509b617ddb0913da03d904dc75c4bb0758 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 22:46:59 -0700 Subject: [PATCH 055/149] add rpc to combined packet parser --- pkg/mpio/packetreader.go | 96 ---------------------------------------- pkg/mpio/packetwriter.go | 40 ----------------- pkg/packet/parser.go | 38 +++++++++------- 3 files changed, 22 insertions(+), 152 deletions(-) delete mode 100644 pkg/mpio/packetreader.go delete mode 100644 pkg/mpio/packetwriter.go diff --git a/pkg/mpio/packetreader.go b/pkg/mpio/packetreader.go deleted file mode 100644 index 8dfb3024..00000000 --- a/pkg/mpio/packetreader.go +++ /dev/null @@ -1,96 +0,0 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - -package mpio - -import ( - "encoding/base64" - "errors" - "io" - "sync" - - "github.com/scripthaus-dev/mshell/pkg/packet" -) - -type PacketReader struct { - CVar *sync.Cond - FdNum int - Buf []byte - Eof bool - Err error -} - -func MakePacketReader(fdNum int) *PacketReader { - return &PacketReader{ - CVar: sync.NewCond(&sync.Mutex{}), - FdNum: fdNum, - } -} - -func (pr *PacketReader) AddData(pk *packet.DataPacketType) { - pr.CVar.L.Lock() - defer pr.CVar.L.Unlock() - defer pr.CVar.Broadcast() - if pr.Eof || pr.Err != nil { - return - } - if pk.Data64 != "" { - realData, err := base64.StdEncoding.DecodeString(pk.Data64) - if err != nil { - pr.Err = err - return - } - pr.Buf = append(pr.Buf, realData...) - } - pr.Eof = pk.Eof - if pk.Error != "" { - pr.Err = errors.New(pk.Error) - } - return -} - -func (pr *PacketReader) Read(buf []byte) (int, error) { - pr.CVar.L.Lock() - defer pr.CVar.L.Unlock() - for { - if pr.Err != nil { - return 0, pr.Err - } - if pr.Eof { - return 0, io.EOF - } - if len(pr.Buf) == 0 { - pr.CVar.Wait() - continue - } - nr := copy(buf, pr.Buf) - pr.Buf = pr.Buf[nr:] - if len(pr.Buf) == 0 { - pr.Buf = nil - } - return nr, nil - } -} - -func (pr *PacketReader) Close() error { - pr.CVar.L.Lock() - defer pr.CVar.L.Unlock() - defer pr.CVar.Broadcast() - if pr.Err == nil { - pr.Err = io.ErrClosedPipe - } - return nil -} - -type NullReader struct{} - -func (NullReader) Read(buf []byte) (int, error) { - return 0, io.EOF -} - -func (NullReader) Close() error { - return nil -} diff --git a/pkg/mpio/packetwriter.go b/pkg/mpio/packetwriter.go deleted file mode 100644 index 0665f044..00000000 --- a/pkg/mpio/packetwriter.go +++ /dev/null @@ -1,40 +0,0 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - -package mpio - -import ( - "encoding/base64" - - "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/packet" -) - -type PacketWriter struct { - FdNum int - Sender *packet.PacketSender - CK base.CommandKey -} - -func MakePacketWriter(fdNum int, sender *packet.PacketSender, ck base.CommandKey) *PacketWriter { - return &PacketWriter{FdNum: fdNum, Sender: sender, CK: ck} -} - -func (pw *PacketWriter) Write(data []byte) (int, error) { - pk := packet.MakeDataPacket() - pk.CK = pw.CK - pk.FdNum = pw.FdNum - pk.Data64 = base64.StdEncoding.EncodeToString(data) - return len(data), pw.Sender.SendPacket(pk) -} - -func (pw *PacketWriter) Close() error { - pk := packet.MakeDataPacket() - pk.CK = pw.CK - pk.FdNum = pw.FdNum - pk.Eof = true - return pw.Sender.SendPacket(pk) -} diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index d405d3b0..2e4199c2 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -37,14 +37,22 @@ func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser { wg.Add(2) go func() { defer wg.Done() - for v := range p1.MainCh { - rtnParser.MainCh <- v + for pk := range p1.MainCh { + sent := rtnParser.trySendRpcResponse(pk) + if sent { + continue + } + rtnParser.MainCh <- pk } }() go func() { defer wg.Done() - for v := range p2.MainCh { - rtnParser.MainCh <- v + for pk := range p2.MainCh { + sent := rtnParser.trySendRpcResponse(pk) + if sent { + continue + } + rtnParser.MainCh <- pk } }() go func() { @@ -56,7 +64,7 @@ func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser) *PacketParser { // should have already registered rpc func (p *PacketParser) WaitForResponse(ctx context.Context, reqId string) RpcResponsePacketType { - entry := p.getRpcEntry(reqId, false) + entry := p.getRpcEntry(reqId) if entry == nil { return nil } @@ -92,18 +100,18 @@ func (p *PacketParser) RegisterRpcSz(reqId string, queueSize int) chan RpcRespon return ch } -func (p *PacketParser) getRpcEntry(reqId string, remove bool) *RpcEntry { +func (p *PacketParser) getRpcEntry(reqId string) *RpcEntry { p.Lock.Lock() defer p.Lock.Unlock() entry := p.RpcMap[reqId] - if entry != nil && remove { - delete(p.RpcMap, reqId) - close(entry.RespCh) - } return entry } -func (p *PacketParser) trySendRpcResponse(respPk RpcResponsePacketType) bool { +func (p *PacketParser) trySendRpcResponse(pk PacketType) bool { + respPk, ok := pk.(RpcResponsePacketType) + if !ok { + return false + } p.Lock.Lock() defer p.Lock.Unlock() entry := p.RpcMap[respPk.GetResponseId()] @@ -185,11 +193,9 @@ func MakePacketParser(input io.Reader) *PacketParser { if pk.GetType() == PingPacketStr { continue } - if respPk, ok := pk.(RpcResponsePacketType); ok { - sent := parser.trySendRpcResponse(respPk) - if sent { - continue - } + sent := parser.trySendRpcResponse(pk) + if sent { + continue } parser.MainCh <- pk } From eb880e024e913faffe69da3cb036880380d2eeb9 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 7 Jul 2022 13:25:42 -0700 Subject: [PATCH 056/149] add hostname to init packet --- pkg/packet/packet.go | 1 + pkg/server/server.go | 2 +- pkg/shexec/shexec.go | 1 + 3 files changed, 3 insertions(+), 1 deletion(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index fe4c4990..0b72c12c 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -387,6 +387,7 @@ type InitPacketType struct { HomeDir string `json:"homedir,omitempty"` Env []string `json:"env,omitempty"` User string `json:"user,omitempty"` + HostName string `json:"hostname,omitempty"` NotFound bool `json:"notfound,omitempty"` UName string `json:"uname,omitempty"` RemoteId string `json:"remoteid,omitempty"` diff --git a/pkg/server/server.go b/pkg/server/server.go index 40ab29cc..6ebd2ad4 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -41,7 +41,7 @@ func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { cproc := m.ClientMap[ck] m.Lock.Unlock() if cproc == nil { - m.Sender.SendCmdError(ck, fmt.Errorf("no client proc for ck '%s'", ck)) + m.Sender.SendCmdError(ck, fmt.Errorf("no client proc for ck '%s', pk=%s", ck, packet.AsString(pk))) return } cproc.Input.SendPacket(pk) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 03d0ba6e..df740a81 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -1060,6 +1060,7 @@ func MakeInitPacket() *packet.InitPacketType { if user, _ := user.Current(); user != nil { initPacket.User = user.Username } + initPacket.HostName, _ = os.Hostname() return initPacket } From 463187221bb89bf65ee0fcc1b56e63c8fc39dfe0 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 7 Jul 2022 21:37:17 -0700 Subject: [PATCH 057/149] update tailer to use filenamegenerator. allow ptyonly option for getcmd --- pkg/cmdtail/cmdtail.go | 67 ++++++++++++++++++++++++------------------ pkg/packet/packet.go | 15 +++++----- 2 files changed, 46 insertions(+), 36 deletions(-) diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index d39c56e1..fbd777bb 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -39,6 +39,11 @@ type CmdWatchEntry struct { Tails []TailPos } +type FileNameGenerator interface { + PtyOutFile(ck base.CommandKey) string + RunOutFile(ck base.CommandKey) string +} + func (w CmdWatchEntry) getTailPos(reqId string) (TailPos, bool) { for _, pos := range w.Tails { if pos.ReqId == reqId { @@ -76,9 +81,9 @@ func (pos TailPos) IsCurrent(entry CmdWatchEntry) bool { type Tailer struct { Lock *sync.Mutex WatchList map[base.CommandKey]CmdWatchEntry - MHomeDir string Watcher *fsnotify.Watcher Sender *packet.PacketSender + Gen FileNameGenerator } func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos TailPos) { @@ -108,10 +113,9 @@ func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { } // delete from watchlist, remove watches - fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, cmdKey) delete(t.WatchList, cmdKey) - t.Watcher.Remove(fileNames.PtyOutFile) - t.Watcher.Remove(fileNames.RunnerOutFile) + t.Watcher.Remove(t.Gen.PtyOutFile(cmdKey)) + t.Watcher.Remove(t.Gen.RunOutFile(cmdKey)) } func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (CmdWatchEntry, TailPos, bool) { @@ -126,13 +130,12 @@ func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (Cm return entry, pos, true } -func MakeTailer(sender *packet.PacketSender) (*Tailer, error) { - mhomeDir := base.GetMShellHomeDir() +func MakeTailer(sender *packet.PacketSender, gen FileNameGenerator) (*Tailer, error) { rtn := &Tailer{ Lock: &sync.Mutex{}, WatchList: make(map[base.CommandKey]CmdWatchEntry), - MHomeDir: mhomeDir, Sender: sender, + Gen: gen, } var err error rtn.Watcher, err = fsnotify.NewWatcher() @@ -156,13 +159,13 @@ func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]b return buf[0:nr], nil } -func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWatchEntry, pos TailPos) (*packet.CmdDataPacketType, error) { +func (t *Tailer) makeCmdDataPacket(entry CmdWatchEntry, pos TailPos) (*packet.CmdDataPacketType, error) { dataPacket := packet.MakeCmdDataPacket(pos.ReqId) dataPacket.CK = entry.CmdKey dataPacket.PtyPos = pos.TailPtyPos dataPacket.RunPos = pos.TailRunPos if entry.FilePtyLen > pos.TailPtyPos { - ptyData, err := t.readDataFromFile(fileNames.PtyOutFile, pos.TailPtyPos, MaxDataBytes) + ptyData, err := t.readDataFromFile(t.Gen.PtyOutFile(entry.CmdKey), pos.TailPtyPos, MaxDataBytes) if err != nil { return nil, err } @@ -170,7 +173,7 @@ func (t *Tailer) makeCmdDataPacket(fileNames *base.CommandFileNames, entry CmdWa dataPacket.PtyDataLen = len(ptyData) } if entry.FileRunLen > pos.TailRunPos { - runData, err := t.readDataFromFile(fileNames.RunnerOutFile, pos.TailRunPos, MaxDataBytes) + runData, err := t.readDataFromFile(t.Gen.RunOutFile(entry.CmdKey), pos.TailRunPos, MaxDataBytes) if err != nil { return nil, err } @@ -188,8 +191,7 @@ func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*pack if !foundPos { return nil, false, nil } - fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, key) - dataPacket, dataErr := t.makeCmdDataPacket(fileNames, entry, pos) + dataPacket, dataErr := t.makeCmdDataPacket(entry, pos) t.Lock.Lock() defer t.Lock.Unlock() @@ -330,13 +332,12 @@ func max(v1 int64, v2 int64) int64 { return v2 } -func (entry *CmdWatchEntry) fillFilePos(scHomeDir string) { - fileNames := base.MakeCommandFileNamesWithHome(scHomeDir, entry.CmdKey) - ptyInfo, _ := os.Stat(fileNames.PtyOutFile) +func (entry *CmdWatchEntry) fillFilePos(gen FileNameGenerator) { + ptyInfo, _ := os.Stat(gen.PtyOutFile(entry.CmdKey)) if ptyInfo != nil { entry.FilePtyLen = ptyInfo.Size() } - runoutInfo, _ := os.Stat(fileNames.RunnerOutFile) + runoutInfo, _ := os.Stat(gen.RunOutFile(entry.CmdKey)) if runoutInfo != nil { entry.FileRunLen = runoutInfo.Size() } @@ -348,27 +349,33 @@ func (t *Tailer) RemoveWatch(pk *packet.UntailCmdPacketType) { t.removeTailPos_nolock(pk.CK, pk.ReqId) } -func (t *Tailer) AddFileWatches_nolock(fileNames *base.CommandFileNames) error { - err := t.Watcher.Add(fileNames.PtyOutFile) +func (t *Tailer) AddFileWatches_nolock(key base.CommandKey, ptyOnly bool) error { + ptyName := t.Gen.PtyOutFile(key) + runName := t.Gen.RunOutFile(key) + fmt.Printf("WATCH> add %s\n", ptyName) + err := t.Watcher.Add(ptyName) if err != nil { return err } - err = t.Watcher.Add(fileNames.RunnerOutFile) + if ptyOnly { + return nil + } + err = t.Watcher.Add(runName) if err != nil { - t.Watcher.Remove(fileNames.PtyOutFile) // best effort clean up + t.Watcher.Remove(ptyName) // best effort clean up return err } return nil } -func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { +// returns (up-to-date/done, error) +func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) (bool, error) { if err := getPacket.CK.Validate("getcmd"); err != nil { - return err + return false, err } if getPacket.ReqId == "" { - return fmt.Errorf("getcmd, no reqid specified") + return false, fmt.Errorf("getcmd, no reqid specified") } - fileNames := base.MakeCommandFileNamesWithHome(t.MHomeDir, getPacket.CK) t.Lock.Lock() defer t.Lock.Unlock() key := getPacket.CK @@ -376,7 +383,7 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { if !foundEntry { // initialize entry, add watches entry = CmdWatchEntry{CmdKey: key} - entry.fillFilePos(t.MHomeDir) + entry.fillFilePos(t.Gen) } pos, foundPos := entry.getTailPos(getPacket.ReqId) if !foundPos { @@ -397,13 +404,15 @@ func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) error { entry.updateTailPos(pos.ReqId, pos) if !pos.Follow && pos.IsCurrent(entry) { // don't add to t.WatchList, don't t.AddFileWatches_nolock, send rpc response - go func() { t.Sender.SendResponse(getPacket.ReqId, true) }() - return nil + return true, nil } if !foundEntry { - t.AddFileWatches_nolock(fileNames) + err := t.AddFileWatches_nolock(key, getPacket.PtyOnly) + if err != nil { + return false, err + } } t.WatchList[key] = entry t.tryStartRun_nolock(entry, pos) - return nil + return false, nil } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 0b72c12c..5bcb273b 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -39,7 +39,7 @@ const ( DataEndPacketStr = "dataend" ResponsePacketStr = "resp" // rpc-response DonePacketStr = "done" - CmdErrorPacketStr = "cmderror" + CmdErrorPacketStr = "cmderror" // command MessagePacketStr = "message" GetCmdPacketStr = "getcmd" // rpc UntailCmdPacketStr = "untailcmd" // rpc @@ -275,12 +275,13 @@ func MakeUntailCmdPacket() *UntailCmdPacketType { } type GetCmdPacketType struct { - Type string `json:"type"` - ReqId string `json:"reqid"` - CK base.CommandKey `json:"ck"` - PtyPos int64 `json:"ptypos"` - RunPos int64 `json:"runpos"` - Tail bool `json:"tail,omitempty"` + Type string `json:"type"` + ReqId string `json:"reqid"` + CK base.CommandKey `json:"ck"` + PtyPos int64 `json:"ptypos"` + RunPos int64 `json:"runpos"` + Tail bool `json:"tail,omitempty"` + PtyOnly bool `json:"ptyonly,omitempty"` } func (*GetCmdPacketType) GetType() string { From 4f4dd67d0668aae197b244d6dd4a72153e0c3dea Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 7 Jul 2022 21:38:05 -0700 Subject: [PATCH 058/149] comment out tailer for now --- main-mshell.go | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index dac5d69a..6029c773 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -13,7 +13,6 @@ import ( "strings" "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/cmdtail" "github.com/scripthaus-dev/mshell/pkg/packet" "github.com/scripthaus-dev/mshell/pkg/server" "github.com/scripthaus-dev/mshell/pkg/shexec" @@ -74,13 +73,13 @@ import ( // }() // } -func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { - err := tailer.AddWatch(pk) - if err != nil { - return err - } - return nil -} +// func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { +// err := tailer.AddWatch(pk) +// if err != nil { +// return err +// } +// return nil +// } // func doMain() { // homeDir := base.GetHomeDir() From 554f8f1b31be0f37335a3a89a238171154e0eb90 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 7 Jul 2022 22:45:14 -0700 Subject: [PATCH 059/149] input packet should use base64, add remoteid. allow untailing a full entry (when command is done) --- pkg/cmdtail/cmdtail.go | 45 ++++++++++++++++++++++++++++++++++-------- pkg/packet/packet.go | 3 ++- pkg/shexec/client.go | 2 ++ 3 files changed, 41 insertions(+), 9 deletions(-) diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index fbd777bb..ae0ece61 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -37,6 +37,7 @@ type CmdWatchEntry struct { FilePtyLen int64 FileRunLen int64 Tails []TailPos + Done bool } type FileNameGenerator interface { @@ -107,11 +108,13 @@ func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { return } entry.removeTailPos(reqId) - if len(entry.Tails) > 0 { - t.WatchList[cmdKey] = entry - return + t.WatchList[cmdKey] = entry + if len(entry.Tails) == 0 { + t.removeWatch_nolock(cmdKey) } +} +func (t *Tailer) removeWatch_nolock(cmdKey base.CommandKey) { // delete from watchlist, remove watches delete(t.WatchList, cmdKey) t.Watcher.Remove(t.Gen.PtyOutFile(cmdKey)) @@ -211,7 +214,7 @@ func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*pack } pos.TailPtyPos += int64(dataPacket.PtyDataLen) pos.TailRunPos += int64(dataPacket.RunDataLen) - if pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen { + if pos.IsCurrent(entry) { // we caught up, tail position equals file length pos.Running = false } @@ -220,14 +223,17 @@ func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*pack } // returns (removed) -func (t *Tailer) checkRemoveNoFollow(cmdKey base.CommandKey, reqId string) bool { +func (t *Tailer) checkRemove(cmdKey base.CommandKey, reqId string) bool { t.Lock.Lock() defer t.Lock.Unlock() - _, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) + entry, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) if !foundPos { return false } - if !pos.Follow { + if !pos.IsCurrent(entry) { + return false + } + if !pos.Follow || entry.Done { t.removeTailPos_nolock(cmdKey, reqId) return true } @@ -246,7 +252,7 @@ func (t *Tailer) RunDataTransfer(key base.CommandKey, reqId string) { break } if !keepRunning { - removed := t.checkRemoveNoFollow(key, reqId) + removed := t.checkRemove(key, reqId) if removed { t.Sender.SendResponse(reqId, true) } @@ -343,6 +349,29 @@ func (entry *CmdWatchEntry) fillFilePos(gen FileNameGenerator) { } } +func (t *Tailer) KeyDone(key base.CommandKey) { + t.Lock.Lock() + defer t.Lock.Unlock() + entry, foundEntry := t.WatchList[key] + if !foundEntry { + return + } + entry.Done = true + var newTails []TailPos + for _, pos := range entry.Tails { + if pos.IsCurrent(entry) { + continue + } + newTails = append(newTails, pos) + } + entry.Tails = newTails + t.WatchList[key] = entry + if len(entry.Tails) == 0 { + t.removeWatch_nolock(key) + } + t.WatchList[key] = entry +} + func (t *Tailer) RemoveWatch(pk *packet.UntailCmdPacketType) { t.Lock.Lock() defer t.Lock.Unlock() diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 5bcb273b..0961bfaa 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -238,7 +238,8 @@ func MakeDataAckPacket() *DataAckPacketType { type InputPacketType struct { Type string `json:"type"` CK base.CommandKey `json:"ck"` - InputData string `json:"inputdata"` + RemoteId string `json:"remoteid"` + InputData64 string `json:"inputdata"` SigNum int `json:"signum,omitempty"` WinSizeRows int `json:"winsizerows"` WinSizeCols int `json:"winsizecols"` diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index f95b252f..8f66a471 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -10,6 +10,8 @@ import ( "github.com/scripthaus-dev/mshell/pkg/packet" ) +// TODO - track buffer sizes for sending input + type ClientProc struct { Cmd *exec.Cmd InitPk *packet.InitPacketType From 74b88185dcdfc35e1bb4b951e6a342c844a3ee57 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 8 Aug 2022 09:52:50 -0700 Subject: [PATCH 060/149] add cache for ensuresessiondir --- pkg/base/base.go | 31 ++++++++++++++++++++++++++++++- pkg/packet/packet.go | 7 ++++--- 2 files changed, 34 insertions(+), 4 deletions(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index 4933273c..1f848e53 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -16,6 +16,7 @@ import ( "path" "path/filepath" "strings" + "sync" "github.com/google/uuid" ) @@ -30,6 +31,9 @@ const SessionsDirBaseName = "sessions" const MShellVersion = "0.1.0" const RemoteIdFile = "remoteid" +var sessionDirCache = make(map[string]string) +var baseLock = &sync.Mutex{} + type CommandFileNames struct { PtyOutFile string StdinFifo string @@ -114,6 +118,22 @@ func GetMShellHomeDir() string { return ExpandHomeDir(DefaultMShellHome) } +func GetPtyOutFile(ck CommandKey, seqNum int) (string, error) { + if err := ck.Validate("ck"); err != nil { + return "", fmt.Errorf("cannot get command files: %w", err) + } + if seqNum < 0 { + return "", fmt.Errorf("invalid seqnum, cannot be negative") + } + sessionId, cmdId := ck.Split() + sdir, err := EnsureSessionDir(sessionId) + if err != nil { + return "", err + } + base := path.Join(sdir, cmdId) + return fmt.Sprintf("%s.%d.ptyout", base, seqNum), nil +} + func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { if err := ck.Validate("ck"); err != nil { return nil, fmt.Errorf("cannot get command files: %w", err) @@ -166,8 +186,14 @@ func EnsureSessionDir(sessionId string) (string, error) { if sessionId == "" { return "", fmt.Errorf("Bad sessionid, cannot be empty") } + baseLock.Lock() + sdir, ok := sessionDirCache[sessionId] + baseLock.Unlock() + if ok { + return sdir, nil + } mhome := GetMShellHomeDir() - sdir := path.Join(mhome, SessionsDirBaseName, sessionId) + sdir = path.Join(mhome, SessionsDirBaseName, sessionId) info, err := os.Stat(sdir) if errors.Is(err, fs.ErrNotExist) { err = os.MkdirAll(sdir, 0777) @@ -182,6 +208,9 @@ func EnsureSessionDir(sessionId string) (string, error) { if !info.IsDir() { return "", fmt.Errorf("session dir '%s' must be a directory", sdir) } + baseLock.Lock() + sessionDirCache[sessionId] = sdir + baseLock.Unlock() return sdir, nil } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 0961bfaa..5fef24a7 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -461,9 +461,10 @@ func MakeCmdStartPacket(reqId string) *CmdStartPacketType { } type TermOpts struct { - Rows int `json:"rows"` - Cols int `json:"cols"` - Term string `json:"term"` + Rows int `json:"rows"` + Cols int `json:"cols"` + Term string `json:"term"` + CmdSize int64 `json:"cmdsize,omitempty"` } type RemoteFd struct { From fbb523aed8d43ef6ba9600c8bf7eaa74077d23a6 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 9 Aug 2022 14:23:59 -0700 Subject: [PATCH 061/149] implement compgen and cd in mshell server --- pkg/packet/packet.go | 31 +++++++++++++++++++++++++++- pkg/server/server.go | 48 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 78 insertions(+), 1 deletion(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 5fef24a7..2def2c90 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -26,6 +26,8 @@ import ( // >cd, >getcmd, >untailcmd, >input, error, <>message, <>ping, 0 && parts[len(parts)-1] == "" { + parts = parts[0 : len(parts)-1] + } + m.Sender.SendResponse(reqId, map[string]interface{}{"comps": parts}) + return +} + +func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) { + reqId := pk.GetReqId() + if cdPk, ok := pk.(*packet.CdPacketType); ok { + err := os.Chdir(cdPk.Dir) + if err != nil { + m.Sender.SendErrorResponse(reqId, fmt.Errorf("cannot change directory: %w", err)) + return + } + m.Sender.SendResponse(reqId, true) + return + } + if compPk, ok := pk.(*packet.CompGenPacketType); ok { + go m.runCompGen(compPk) + return + } + m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid rpc type '%s'", pk.GetType())) + return +} + func (m *MServer) runCommand(runPacket *packet.RunPacketType) { if err := runPacket.CK.Validate("packet"); err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) @@ -117,6 +161,10 @@ func RunServer() (int, error) { server.ProcessCommandPacket(cmdPk) continue } + if rpcPk, ok := pk.(packet.RpcPacketType); ok { + server.ProcessRpcPacket(rpcPk) + continue + } server.Sender.SendMessage(fmt.Sprintf("invalid packet '%s' sent to mshell server", packet.AsString(pk))) continue } From d4528d1c42669b02f77954e0508c5df04362f58c Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 10 Aug 2022 16:07:41 -0700 Subject: [PATCH 062/149] compgen hasmore --- pkg/packet/packet.go | 14 ++++++++++++++ pkg/server/server.go | 9 +++++++-- 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 2def2c90..74c7e36a 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -11,6 +11,7 @@ import ( "context" "encoding/base64" "encoding/json" + "errors" "fmt" "io" "os" @@ -364,6 +365,19 @@ func (*ResponsePacketType) GetResponseDone() bool { return true } +func (p *ResponsePacketType) Err() error { + if p == nil { + return fmt.Errorf("no response received") + } + if !p.Success { + if p.Error != "" { + return errors.New(p.Error) + } + return fmt.Errorf("rpc failed") + } + return nil +} + func MakeErrorResponsePacket(reqId string, err error) *ResponsePacketType { return &ResponsePacketType{Type: ResponsePacketStr, RespId: reqId, Error: err.Error()} } diff --git a/pkg/server/server.go b/pkg/server/server.go index 9cb8713f..4f2f367c 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -57,7 +57,7 @@ func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid compgen type '%s'", compPk.CompType)) return } - compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | head -n %d", shellescape.Quote(compPk.Cwd), shellescape.Quote(compPk.CompType), shellescape.Quote(compPk.Prefix), packet.MaxCompGenValues) + compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | head -n %d", shellescape.Quote(compPk.Cwd), shellescape.Quote(compPk.CompType), shellescape.Quote(compPk.Prefix), packet.MaxCompGenValues+1) ecmd := exec.Command("bash", "-c", compGenCmdStr) outputBytes, err := ecmd.Output() if err != nil { @@ -69,7 +69,12 @@ func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { if len(parts) > 0 && parts[len(parts)-1] == "" { parts = parts[0 : len(parts)-1] } - m.Sender.SendResponse(reqId, map[string]interface{}{"comps": parts}) + hasMore := false + if len(parts) > packet.MaxCompGenValues { + hasMore = true + parts = parts[0:packet.MaxCompGenValues] + } + m.Sender.SendResponse(reqId, map[string]interface{}{"comps": parts, "hasmore": hasMore}) return } From f4515139a1b1f3b2def622d1cbb81ca408e3970c Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 10 Aug 2022 18:33:50 -0700 Subject: [PATCH 063/149] uniq completions --- pkg/server/server.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 4f2f367c..e9ce72ad 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -57,7 +57,7 @@ func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid compgen type '%s'", compPk.CompType)) return } - compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | head -n %d", shellescape.Quote(compPk.Cwd), shellescape.Quote(compPk.CompType), shellescape.Quote(compPk.Prefix), packet.MaxCompGenValues+1) + compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | sort | uniq | head -n %d", shellescape.Quote(compPk.Cwd), shellescape.Quote(compPk.CompType), shellescape.Quote(compPk.Prefix), packet.MaxCompGenValues+1) ecmd := exec.Command("bash", "-c", compGenCmdStr) outputBytes, err := ecmd.Output() if err != nil { From 5a8ffa3544dbc6df5ca4ce8c3f54b92f275eeba3 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 11 Aug 2022 10:21:11 -0700 Subject: [PATCH 064/149] more complicated compgen for directory vs file completion --- pkg/server/server.go | 77 ++++++++++++++++++++++++++++++++++++++------ 1 file changed, 68 insertions(+), 9 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index e9ce72ad..e0534941 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -11,6 +11,7 @@ import ( "fmt" "os" "os/exec" + "sort" "strings" "sync" @@ -51,18 +52,15 @@ func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { return } -func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { - reqId := compPk.GetReqId() - if !packet.IsValidCompGenType(compPk.CompType) { - m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid compgen type '%s'", compPk.CompType)) - return +func runSingleCompGen(cwd string, compType string, prefix string) ([]string, bool, error) { + if !packet.IsValidCompGenType(compType) { + return nil, false, fmt.Errorf("invalid compgen type '%s'", compType) } - compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | sort | uniq | head -n %d", shellescape.Quote(compPk.Cwd), shellescape.Quote(compPk.CompType), shellescape.Quote(compPk.Prefix), packet.MaxCompGenValues+1) + compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | sort | uniq | head -n %d", shellescape.Quote(cwd), shellescape.Quote(compType), shellescape.Quote(prefix), packet.MaxCompGenValues+1) ecmd := exec.Command("bash", "-c", compGenCmdStr) outputBytes, err := ecmd.Output() if err != nil { - m.Sender.SendErrorResponse(reqId, fmt.Errorf("compgen error: %w", err)) - return + return nil, false, fmt.Errorf("compgen error: %w", err) } outputStr := string(outputBytes) parts := strings.Split(outputStr, "\n") @@ -74,7 +72,68 @@ func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { hasMore = true parts = parts[0:packet.MaxCompGenValues] } - m.Sender.SendResponse(reqId, map[string]interface{}{"comps": parts, "hasmore": hasMore}) + return parts, hasMore, nil +} + +func appendSlashes(comps []string) { + for idx, comp := range comps { + comps[idx] = comp + "/" + } +} + +func strArrToMap(strs []string) map[string]bool { + rtn := make(map[string]bool) + for _, s := range strs { + rtn[s] = true + } + return rtn +} + +func (m *MServer) runFileCompGen(compPk *packet.CompGenPacketType) { + // get directories and files, unique them and put slashes on directories for completion + reqId := compPk.GetReqId() + compDirs, hasMoreDirs, err := runSingleCompGen(compPk.Cwd, "directory", compPk.Prefix) + if err != nil { + m.Sender.SendErrorResponse(reqId, err) + return + } + compFiles, hasMoreFiles, err := runSingleCompGen(compPk.Cwd, "file", compPk.Prefix) + if err != nil { + m.Sender.SendErrorResponse(reqId, err) + return + } + + dirMap := strArrToMap(compDirs) + // seed comps with dirs (but append slashes) + comps := compDirs + appendSlashes(comps) + // add files that are not directories (look up in dirMap) + for _, file := range compFiles { + if dirMap[file] { + continue + } + comps = append(comps, file) + } + sort.Strings(comps) // resort + m.Sender.SendResponse(reqId, map[string]interface{}{"comps": comps, "hasmore": (hasMoreFiles || hasMoreDirs)}) + return +} + +func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { + reqId := compPk.GetReqId() + if compPk.CompType == "file" { + m.runFileCompGen(compPk) + return + } + comps, hasMore, err := runSingleCompGen(compPk.Cwd, compPk.CompType, compPk.Prefix) + if err != nil { + m.Sender.SendErrorResponse(reqId, err) + return + } + if compPk.CompType == "directory" { + appendSlashes(comps) + } + m.Sender.SendResponse(reqId, map[string]interface{}{"comps": comps, "hasmore": hasMore}) return } From 9542d14473bf9b0263f000a0b5b5b81df7c40964 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 16 Aug 2022 16:26:06 -0700 Subject: [PATCH 065/149] make mshell home directory during getremoteid --- pkg/base/base.go | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index 1f848e53..c14736d5 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -198,7 +198,7 @@ func EnsureSessionDir(sessionId string) (string, error) { if errors.Is(err, fs.ErrNotExist) { err = os.MkdirAll(sdir, 0777) if err != nil { - return "", err + return "", fmt.Errorf("cannot make mshell session directory[%s]: %w", sdir, err) } info, err = os.Stat(sdir) } @@ -254,6 +254,20 @@ func GoArchOptFile(goos string, goarch string) string { func GetRemoteId() (string, error) { mhome := GetMShellHomeDir() + homeInfo, err := os.Stat(mhome) + if errors.Is(err, fs.ErrNotExist) { + err = os.MkdirAll(mhome, 0777) + if err != nil { + return "", fmt.Errorf("cannot make mshell home directory[%s]: %w", mhome, err) + } + homeInfo, err = os.Stat(mhome) + } + if err != nil { + return "", fmt.Errorf("cannot stat mshell home directory[%s]: %w", mhome, err) + } + if !homeInfo.IsDir() { + return "", fmt.Errorf("mshell home directory[%s] is not a directory", mhome) + } remoteIdFile := path.Join(mhome, RemoteIdFile) fd, err := os.Open(remoteIdFile) if errors.Is(err, fs.ErrNotExist) { From 77bd1fa7bced7cf0fd15a62c0f2d70dedd815b13 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 17 Aug 2022 12:55:12 -0700 Subject: [PATCH 066/149] return uname --- pkg/server/server.go | 2 +- pkg/shexec/client.go | 21 +++++++++++---------- 2 files changed, 12 insertions(+), 11 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index e0534941..c1cd6bb4 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -166,7 +166,7 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } - cproc, err := shexec.MakeClientProc(ecmd) + cproc, _, err := shexec.MakeClientProc(ecmd) if err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("starting mshell client: %s", err)) return diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index 8f66a471..fb6e6e4e 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -23,23 +23,24 @@ type ClientProc struct { Output *packet.PacketParser } -func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, error) { +// returns (clientproc, uname, error) +func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, string, error) { inputWriter, err := ecmd.StdinPipe() if err != nil { - return nil, fmt.Errorf("creating stdin pipe: %v", err) + return nil, "", fmt.Errorf("creating stdin pipe: %v", err) } stdoutReader, err := ecmd.StdoutPipe() if err != nil { - return nil, fmt.Errorf("creating stdout pipe: %v", err) + return nil, "", fmt.Errorf("creating stdout pipe: %v", err) } stderrReader, err := ecmd.StderrPipe() if err != nil { - return nil, fmt.Errorf("creating stderr pipe: %v", err) + return nil, "", fmt.Errorf("creating stderr pipe: %v", err) } startTs := time.Now() err = ecmd.Start() if err != nil { - return nil, fmt.Errorf("running local client: %w", err) + return nil, "", fmt.Errorf("running local client: %w", err) } sender := packet.MakePacketSender(inputWriter) stdoutPacketParser := packet.MakePacketParser(stdoutReader) @@ -57,25 +58,25 @@ func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, error) { for pk := range packetParser.MainCh { if pk.GetType() != packet.InitPacketStr { cproc.Close() - return nil, fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) + return nil, "", fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) } initPk := pk.(*packet.InitPacketType) if initPk.NotFound { cproc.Close() - return nil, fmt.Errorf("mshell command not found on local server") + return nil, initPk.UName, fmt.Errorf("mshell command not found on local server") } if initPk.Version != base.MShellVersion { cproc.Close() - return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) + return nil, initPk.UName, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) } cproc.InitPk = initPk break } if cproc.InitPk == nil { cproc.Close() - return nil, fmt.Errorf("no init packet received from mshell client") + return nil, "", fmt.Errorf("no init packet received from mshell client") } - return cproc, nil + return cproc, cproc.InitPk.UName, nil } func (cproc *ClientProc) Close() { From 38870f9c6ee0d9e8292dbc8a57aab43c326d5d45 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 18 Aug 2022 20:34:20 -0700 Subject: [PATCH 067/149] circular file buffer. metadata in header. uses flock to synchronize access. write metadata before and after writing file data. --- pkg/cirfile/cirfile.go | 412 ++++++++++++++++++++++++++++++++++++ pkg/cirfile/cirfile_test.go | 222 +++++++++++++++++++ 2 files changed, 634 insertions(+) create mode 100644 pkg/cirfile/cirfile.go create mode 100644 pkg/cirfile/cirfile_test.go diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go new file mode 100644 index 00000000..d40fdceb --- /dev/null +++ b/pkg/cirfile/cirfile.go @@ -0,0 +1,412 @@ +package cirfile + +import ( + "context" + "fmt" + "io" + "os" + "syscall" + "time" +) + +// CBUF[version] [maxsize] [fileoffset] [startpos] [endpos] +const HeaderFmt = "CBUF%02d %19d %19d %19d %19d\n" // 87 bytes +const HeaderLen = 256 // set to 256 for future expandability +const FullHeaderFmt = "%-255s\n" // 256 bytes (255 + newline) +const CurrentVersion = 1 +const FilePosEmpty = -1 // sentinel, if startpos is set to -1, file is empty + +const InitialLockDelay = 10 * time.Millisecond +const InitialLockTries = 5 +const LockDelay = 100 * time.Millisecond + +type File struct { + OSFile *os.File + Version byte + MaxSize int64 + FileOffset int64 + StartPos int64 + EndPos int64 + FileDataSize int64 // size of data (does not include header size) + FlockStatus int +} + +func (f *File) flock(ctx context.Context, lockType int) error { + err := syscall.Flock(int(f.OSFile.Fd()), lockType|syscall.LOCK_NB) + if err == nil { + f.FlockStatus = lockType + return nil + } + if err != syscall.EWOULDBLOCK { + return err + } + if ctx == nil { + return syscall.EWOULDBLOCK + } + // busy-wait with context + numWaits := 0 + for { + numWaits++ + var timeout time.Duration + if numWaits <= InitialLockTries { + timeout = InitialLockDelay + } else { + timeout = LockDelay + } + select { + case <-time.After(timeout): + break + case <-ctx.Done(): + return ctx.Err() + } + err = syscall.Flock(int(f.OSFile.Fd()), lockType|syscall.LOCK_NB) + if err == nil { + f.FlockStatus = lockType + return nil + } + if err != syscall.EWOULDBLOCK { + return err + } + } + return fmt.Errorf("could not acquire lock") +} + +func (f *File) unflock() { + syscall.Flock(int(f.OSFile.Fd()), syscall.LOCK_UN) // ignore error (nothing to do about it anyway) + f.FlockStatus = 0 + return +} + +// does not read metadata because locking could block/fail. we want to be able +// to return a valid file struct without blocking. +func OpenCirFile(fileName string) (*File, error) { + fd, err := os.OpenFile(fileName, os.O_RDWR, 0777) + if err != nil { + return nil, err + } + finfo, err := fd.Stat() + if err != nil { + return nil, err + } + if finfo.Size() < HeaderLen { + return nil, fmt.Errorf("invalid cirfile, file length[%d] less than HeaderLen[%d]", finfo.Size(), HeaderLen) + } + rtn := &File{OSFile: fd} + return rtn, nil +} + +// if the file already exists, it is an error. +// there is a race condition if two goroutines try to create the same file between Stat() and Create(), so +// they both might get no error, but only one file will be valid. if this is a concern, this call +// should be externally synchronized. +func CreateCirFile(fileName string, maxSize int64) (*File, error) { + if maxSize <= 0 { + return nil, fmt.Errorf("invalid maxsize[%d]", maxSize) + } + _, err := os.Stat(fileName) + if err == nil { + return nil, fmt.Errorf("file[%s] already exists", fileName) + } + if !os.IsNotExist(err) { + return nil, fmt.Errorf("cannot stat: %w", err) + } + fd, err := os.Create(fileName) + if err != nil { + return nil, err + } + rtn := &File{OSFile: fd, Version: CurrentVersion, MaxSize: maxSize, StartPos: FilePosEmpty} + err = rtn.flock(nil, syscall.LOCK_EX) + if err != nil { + return nil, err + } + defer rtn.unflock() + err = rtn.writeMeta() + if err != nil { + return nil, err + } + return rtn, nil +} + +func (f *File) ReadMeta(ctx context.Context) error { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return err + } + defer f.unflock() + return f.readMeta() +} + +func (f *File) hasShLock() bool { + return f.FlockStatus == syscall.LOCK_EX || f.FlockStatus == syscall.LOCK_SH +} + +func (f *File) hasExLock() bool { + return f.FlockStatus == syscall.LOCK_EX +} + +func (f *File) readMeta() error { + if f.OSFile == nil { + return fmt.Errorf("no *os.File") + } + if !f.hasShLock() { + return fmt.Errorf("writeMeta must hold LOCK_SH") + } + _, err := f.OSFile.Seek(0, 0) + if err != nil { + return fmt.Errorf("cannot seek file: %w", err) + } + finfo, err := f.OSFile.Stat() + if err != nil { + return fmt.Errorf("cannot stat file: %w", err) + } + if finfo.Size() < 256 { + return fmt.Errorf("invalid cbuf file size[%d] < 256", finfo.Size()) + } + f.FileDataSize = finfo.Size() - 256 + buf := make([]byte, 256) + _, err = io.ReadFull(f.OSFile, buf) + if err != nil { + return fmt.Errorf("error reading header: %w", err) + } + // currently only one version, so we don't need to have special logic here yet + _, err = fmt.Sscanf(string(buf), HeaderFmt, &f.Version, &f.MaxSize, &f.FileOffset, &f.StartPos, &f.EndPos) + if err != nil { + return fmt.Errorf("sscanf error: %w", err) + } + if f.Version != CurrentVersion { + return fmt.Errorf("invalid cbuf version[%d]", f.Version) + } + // possible incomplete write, fix start/end pos to be within filesize + if f.FileDataSize == 0 { + f.StartPos = FilePosEmpty + f.EndPos = 0 + } else if f.StartPos >= f.FileDataSize && f.EndPos >= f.FileDataSize { + f.StartPos = FilePosEmpty + f.EndPos = 0 + } else if f.StartPos >= f.FileDataSize { + f.StartPos = 0 + } else if f.EndPos >= f.FileDataSize { + f.EndPos = f.FileDataSize - 1 + } + if f.MaxSize <= 0 || f.FileOffset < 0 || (f.StartPos < 0 && f.StartPos != FilePosEmpty) || f.StartPos >= f.MaxSize || f.EndPos < 0 || f.EndPos >= f.MaxSize { + return fmt.Errorf("invalid cbuf metadata version[%d] filedatasize[%d] maxsize[%d] fileoffset[%d] startpos[%d] endpos[%d]", f.Version, f.FileDataSize, f.MaxSize, f.FileOffset, f.StartPos, f.EndPos) + } + return nil +} + +// no error checking of meta values +func (f *File) writeMeta() error { + if f.OSFile == nil { + return fmt.Errorf("no *os.File") + } + if !f.hasExLock() { + return fmt.Errorf("writeMeta must hold LOCK_EX") + } + _, err := f.OSFile.Seek(0, 0) + if err != nil { + return fmt.Errorf("cannot seek file: %w", err) + } + metaStr := fmt.Sprintf(HeaderFmt, f.Version, f.MaxSize, f.FileOffset, f.StartPos, f.EndPos) + fullMetaStr := fmt.Sprintf(FullHeaderFmt, metaStr) + _, err = f.OSFile.WriteString(fullMetaStr) + if err != nil { + return fmt.Errorf("write error: %w", err) + } + return nil +} + +// returns (fileOffset, datasize, error) +// datasize is the current amount of readable data held in the cirfile +func (f *File) GetStartOffsetAndSize(ctx context.Context) (int64, int64, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, 0, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, 0, err + } + chunks := f.getFileChunks() + return f.FileOffset, totalChunksSize(chunks), nil +} + +type fileChunk struct { + StartPos int64 + Len int64 +} + +func totalChunksSize(chunks []fileChunk) int64 { + var rtn int64 + for _, chunk := range chunks { + rtn += chunk.Len + } + return rtn +} + +func advanceChunks(chunks []fileChunk, offset int64) []fileChunk { + if offset < 0 { + panic(fmt.Sprintf("invalid negative offset: %d", offset)) + } + if offset == 0 { + return chunks + } + var rtn []fileChunk + for _, chunk := range chunks { + if offset >= chunk.Len { + offset = offset - chunk.Len + continue + } + if offset == 0 { + rtn = append(rtn, chunk) + } else { + rtn = append(rtn, fileChunk{chunk.StartPos + offset, chunk.Len - offset}) + offset = 0 + } + } + return rtn +} + +func (f *File) getFileChunks() []fileChunk { + if f.StartPos == FilePosEmpty { + return nil + } + if f.EndPos >= f.StartPos { + return []fileChunk{fileChunk{f.StartPos, f.EndPos - f.StartPos + 1}} + } + return []fileChunk{ + fileChunk{f.StartPos, f.FileDataSize - f.StartPos}, + fileChunk{0, f.EndPos + 1}, + } +} + +func (f *File) getFreeChunks() []fileChunk { + if f.StartPos == FilePosEmpty { + return []fileChunk{fileChunk{0, f.MaxSize}} + } + if (f.EndPos == f.StartPos-1) || (f.StartPos == 0 && f.EndPos == f.MaxSize-1) { + return nil + } + if f.EndPos < f.StartPos { + return []fileChunk{fileChunk{f.EndPos + 1, f.StartPos - f.EndPos - 1}} + } + var rtn []fileChunk + if f.EndPos < f.MaxSize-1 { + rtn = append(rtn, fileChunk{f.EndPos + 1, f.MaxSize - f.EndPos - 1}) + } + if f.StartPos > 0 { + rtn = append(rtn, fileChunk{0, f.StartPos}) + } + return rtn +} + +// returns (realOffset, data, error) +// will only return io.EOF when len(data) == 0, otherwise will just do a short read +func (f *File) ReadNext(ctx context.Context, buf []byte, offset int64) (int64, int, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, 0, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, 0, err + } + if offset < f.FileOffset { + offset = f.FileOffset + } + relativeOffset := offset - f.FileOffset + chunks := f.getFileChunks() + curSize := totalChunksSize(chunks) + if offset >= f.FileOffset+curSize { + return f.FileOffset + curSize, 0, nil + } + chunks = advanceChunks(chunks, relativeOffset) + numRead := 0 + for _, chunk := range chunks { + if numRead >= len(buf) { + break + } + toRead := len(buf) - numRead + if toRead > int(chunk.Len) { + toRead = int(chunk.Len) + } + nr, err := f.OSFile.ReadAt(buf[numRead:numRead+toRead], chunk.StartPos+HeaderLen) + if err != nil { + return offset, 0, err + } + numRead += nr + } + return offset, numRead, nil +} + +func (f *File) ensureFreeSpace(requiredSpace int64) error { + if f.StartPos == FilePosEmpty { + return nil + } + chunks := f.getFileChunks() + curSpace := f.MaxSize - totalChunksSize(chunks) + if curSpace >= requiredSpace { + return nil + } + neededSpace := requiredSpace - curSpace + if requiredSpace >= f.MaxSize { + f.StartPos = FilePosEmpty + f.EndPos = 0 + f.FileOffset += neededSpace + } else { + f.StartPos = (f.StartPos + neededSpace) % f.MaxSize + f.FileOffset += neededSpace + } + return f.writeMeta() +} + +func (f *File) AppendData(ctx context.Context, buf []byte) error { + err := f.flock(ctx, syscall.LOCK_EX) + if err != nil { + return err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return err + } + err = f.ensureFreeSpace(int64(len(buf))) + if err != nil { + return err + } + if len(buf) >= int(f.MaxSize) { + buf = buf[len(buf)-int(f.MaxSize):] + } + chunks := f.getFreeChunks() + numWrite := 0 + for _, chunk := range chunks { + if numWrite >= len(buf) { + break + } + if chunk.Len == 0 { + continue + } + toRead := len(buf) - numWrite + if toRead > int(chunk.Len) { + toRead = int(chunk.Len) + } + nw, err := f.OSFile.WriteAt(buf[numWrite:numWrite+toRead], chunk.StartPos+HeaderLen) + if err != nil { + return err + } + if chunk.StartPos+int64(nw) > f.FileDataSize { + f.FileDataSize = chunk.StartPos + int64(nw) + } + if f.StartPos == FilePosEmpty { + f.StartPos = chunk.StartPos + } + f.EndPos = chunk.StartPos + int64(nw) - 1 + numWrite += nw + } + err = f.writeMeta() + if err != nil { + return err + } + return nil +} diff --git a/pkg/cirfile/cirfile_test.go b/pkg/cirfile/cirfile_test.go new file mode 100644 index 00000000..4e4f555c --- /dev/null +++ b/pkg/cirfile/cirfile_test.go @@ -0,0 +1,222 @@ +package cirfile + +import ( + "context" + "fmt" + "os" + "path" + "syscall" + "testing" + "time" +) + +func validateFileSize(t *testing.T, name string, size int) { + finfo, err := os.Stat(name) + if err != nil { + t.Fatalf("error stating file[%s]: %v", name, err) + } + if int(finfo.Size()) != size { + t.Fatalf("invalid file[%s] expected[%d] got[%d]", name, size, finfo.Size()) + } +} + +func validateMeta(t *testing.T, desc string, f *File, startPos int64, endPos int64, dataSize int64, offset int64) { + if f.StartPos != startPos || f.EndPos != endPos || f.FileDataSize != dataSize || f.FileOffset != offset { + t.Fatalf("metadata error (%s): startpos[%d %d] endpos[%d %d] filedatasize[%d %d] fileoffset[%d %d]", desc, f.StartPos, startPos, f.EndPos, endPos, f.FileDataSize, dataSize, f.FileOffset, offset) + } +} + +func dumpFile(name string) { + barr, _ := os.ReadFile(name) + fmt.Printf("<<<\n%s\n>>>", string(barr)) +} + +func makeData(size int) string { + var rtn string + for { + if len(rtn) >= size { + break + } + needed := size - len(rtn) + if needed < 10 { + rtn += "123456789\n"[0:needed] + break + } + rtn += "123456789\n" + } + return rtn +} + +func TestCreate(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := OpenCirFile(f1Name) + if err == nil || f != nil { + t.Fatalf("OpenCirFile f1.cf should fail (no file)") + } + f, err = CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("CreateCirFile f1.cf failed: %v", err) + } + if f == nil { + t.Fatalf("CreateCirFile f1.cf returned nil") + } + err = f.ReadMeta(context.Background()) + if err != nil { + t.Fatalf("cannot readmeta from f1.cf: %v", err) + } + validateFileSize(t, f1Name, 256) + if f.Version != CurrentVersion || f.MaxSize != 100 || f.FileOffset != 0 || f.StartPos != FilePosEmpty || f.EndPos != 0 || f.FileDataSize != 0 || f.FlockStatus != 0 { + t.Fatalf("error with initial metadata #%v", f) + } + buf := make([]byte, 200) + realOffset, nr, err := f.ReadNext(context.Background(), buf, 0) + if realOffset != 0 || nr != 0 || err != nil { + t.Fatalf("error with empty read: real-offset[%d] nr[%d] err[%v]", realOffset, nr, err) + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 1000) + if realOffset != 0 || nr != 0 || err != nil { + t.Fatalf("error with empty read: real-offset[%d] nr[%d] err[%v]", realOffset, nr, err) + } + f2, err := CreateCirFile(f1Name, 100) + if err == nil || f2 != nil { + t.Fatalf("should be an error to create duplicate CirFile") + } +} + +func TestFile(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("cannot create cirfile: %v", err) + } + err = f.AppendData(context.Background(), nil) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen) + if f.StartPos != FilePosEmpty || f.EndPos != 0 || f.FileDataSize != 0 { + t.Fatalf("metadata error (1): %#v", f) + } + err = f.AppendData(context.Background(), []byte("hello")) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+5) + if f.StartPos != 0 || f.EndPos != 4 || f.FileDataSize != 5 { + t.Fatalf("metadata error (2): %#v", f) + } + err = f.AppendData(context.Background(), []byte(" foo")) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+9) + validateMeta(t, "3", f, 0, 8, 9, 0) + err = f.AppendData(context.Background(), []byte("\n"+makeData(20))) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+30) + validateMeta(t, "4", f, 0, 29, 30, 0) + + data120 := makeData(120) + err = f.AppendData(context.Background(), []byte(data120)) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+100) + validateMeta(t, "5", f, 0, 99, 100, 50) + err = f.AppendData(context.Background(), []byte("foo ")) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+100) + validateMeta(t, "6", f, 4, 3, 100, 54) + + buf := make([]byte, 5) + realOffset, nr, err := f.ReadNext(context.Background(), buf, 0) + if err != nil { + t.Fatalf("cannot ReadNext: %v", err) + } + if realOffset != 54 { + t.Fatalf("wrong realoffset got[%d] expected[%d]", realOffset, 54) + } + if nr != 5 { + t.Fatalf("wrong nr got[%d] expected[%d]", nr, 5) + } + if string(buf[0:nr]) != "56789" { + t.Fatalf("wrong buf return got[%s] expected[%s]", string(buf[0:nr]), "56789") + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 60) + if err != nil { + t.Fatalf("cannot readnext: %v", err) + } + if realOffset != 60 && nr != 5 { + t.Fatalf("invalid rtn realoffset[%d] nr[%d]", realOffset, nr) + } + if string(buf[0:nr]) != "12345" { + t.Fatalf("invalid rtn buf[%s]", string(buf[0:nr])) + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 800) + if err != nil || realOffset != 154 || nr != 0 { + t.Fatalf("invalid past end read: err[%v] realoffset[%d] nr[%d]", err, realOffset, nr) + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 150) + if err != nil || realOffset != 150 || nr != 4 || string(buf[0:nr]) != "foo " { + t.Fatalf("invalid end read: err[%v] realoffset[%d] nr[%d] buf[%s]", err, realOffset, nr, string(buf[0:nr])) + } +} + +func TestFlock(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("cannot create cirfile: %v", err) + } + fd2, err := os.OpenFile(f1Name, os.O_RDWR, 0777) + if err != nil { + t.Fatalf("cannot open file: %v", err) + } + err = syscall.Flock(int(fd2.Fd()), syscall.LOCK_EX) + if err != nil { + t.Fatalf("cannot lock fd: %v", err) + } + err = f.AppendData(nil, []byte("hello")) + if err != syscall.EWOULDBLOCK { + t.Fatalf("append should fail with EWOULDBLOCK") + } + timeoutCtx, _ := context.WithTimeout(context.Background(), 20*time.Millisecond) + startTs := time.Now() + err = f.ReadMeta(timeoutCtx) + if err != context.DeadlineExceeded { + t.Fatalf("readmeta should fail with context.DeadlineExceeded") + } + dur := time.Now().Sub(startTs) + if dur < 20*time.Millisecond { + t.Fatalf("readmeta should take at least 20ms") + } + syscall.Flock(int(fd2.Fd()), syscall.LOCK_UN) + err = f.ReadMeta(timeoutCtx) + if err != nil { + t.Fatalf("readmeta err: %v", err) + } + err = syscall.Flock(int(fd2.Fd()), syscall.LOCK_SH) + if err != nil { + t.Fatalf("cannot flock: %v", err) + } + err = f.AppendData(nil, []byte("hello")) + if err != syscall.EWOULDBLOCK { + t.Fatalf("append should fail with EWOULDBLOCK") + } + err = f.ReadMeta(timeoutCtx) + if err != nil { + t.Fatalf("readmeta err (should work because LOCK_SH): %v", err) + } + fd2.Close() + err = f.AppendData(nil, []byte("hello")) + if err != nil { + t.Fatalf("append error (should work fd2 was closed): %v", err) + } +} From e1eecae6d3a9d247c706a6a7e5da629c7294126f Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 18 Aug 2022 21:54:42 -0700 Subject: [PATCH 068/149] implement WriteAt, refactored and tested append, writeat still needs testing --- pkg/cirfile/cirfile.go | 175 +++++++++++++++++++++++++++--------- pkg/cirfile/cirfile_test.go | 8 +- 2 files changed, 133 insertions(+), 50 deletions(-) diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go index d40fdceb..5992b876 100644 --- a/pkg/cirfile/cirfile.go +++ b/pkg/cirfile/cirfile.go @@ -20,6 +20,7 @@ const InitialLockDelay = 10 * time.Millisecond const InitialLockTries = 5 const LockDelay = 100 * time.Millisecond +// File objects are *not* multithread safe, operations must be externally synchronized type File struct { OSFile *os.File Version byte @@ -72,8 +73,10 @@ func (f *File) flock(ctx context.Context, lockType int) error { } func (f *File) unflock() { - syscall.Flock(int(f.OSFile.Fd()), syscall.LOCK_UN) // ignore error (nothing to do about it anyway) - f.FlockStatus = 0 + if f.FlockStatus != 0 { + syscall.Flock(int(f.OSFile.Fd()), syscall.LOCK_UN) // ignore error (nothing to do about it anyway) + f.FlockStatus = 0 + } return } @@ -127,6 +130,10 @@ func CreateCirFile(fileName string, maxSize int64) (*File, error) { return rtn, nil } +func (f *File) Close() error { + return f.OSFile.Close() +} + func (f *File) ReadMeta(ctx context.Context) error { err := f.flock(ctx, syscall.LOCK_SH) if err != nil { @@ -341,16 +348,13 @@ func (f *File) ReadNext(ctx context.Context, buf []byte, offset int64) (int64, i } func (f *File) ensureFreeSpace(requiredSpace int64) error { - if f.StartPos == FilePosEmpty { - return nil - } chunks := f.getFileChunks() curSpace := f.MaxSize - totalChunksSize(chunks) if curSpace >= requiredSpace { return nil } neededSpace := requiredSpace - curSpace - if requiredSpace >= f.MaxSize { + if requiredSpace >= f.MaxSize || f.StartPos == FilePosEmpty { f.StartPos = FilePosEmpty f.EndPos = 0 f.FileOffset += neededSpace @@ -361,6 +365,126 @@ func (f *File) ensureFreeSpace(requiredSpace int64) error { return f.writeMeta() } +// does not implement io.WriterAt (needs context) +func (f *File) WriteAt(ctx context.Context, buf []byte, writePos int64) error { + if writePos < 0 { + return fmt.Errorf("WriteAt got invalid writePos[%d]", writePos) + } + err := f.flock(ctx, syscall.LOCK_EX) + if err != nil { + return err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return err + } + chunks := f.getFileChunks() + currentSize := totalChunksSize(chunks) + if writePos < f.FileOffset { + negOffset := f.FileOffset - writePos + if negOffset >= int64(len(buf)) { + return nil + } + buf = buf[negOffset:] + writePos = f.FileOffset + } + if writePos > f.FileOffset+currentSize { + // fill gap with zero bytes + posOffset := writePos - (f.FileOffset + currentSize) + err = f.ensureFreeSpace(int64(posOffset)) + if err != nil { + return err + } + var zeroBuf []byte + if posOffset >= f.MaxSize { + zeroBuf = make([]byte, f.MaxSize) + } else { + zeroBuf = make([]byte, posOffset) + } + err = f.internalAppendData(zeroBuf) + if err != nil { + return err + } + // recalc chunks/currentSize + chunks = f.getFileChunks() + currentSize = totalChunksSize(chunks) + // after writing the zero bytes, writePos == f.FileOffset+currentSize (the rest is a straight append) + } + // now writePos >= f.FileOffset && writePos <= f.FileOffset+currentSize (check invariant) + if writePos < f.FileOffset || writePos > f.FileOffset+currentSize { + panic(fmt.Sprintf("invalid writePos, invariant violated writepos[%d] fileoffset[%d] currentsize[%d]", writePos, f.FileOffset, currentSize)) + } + // overwrite existing data (in chunks). advance by writePosOffset + writePosOffset := writePos - f.FileOffset + if writePosOffset < currentSize { + advChunks := advanceChunks(chunks, writePosOffset) + nw, err := f.writeToChunks(buf, advChunks, false) + if err != nil { + return err + } + buf = buf[nw:] + if len(buf) == 0 { + return nil + } + } + // buf contains what was unwritten. this unwritten data is now just a straight append + return f.internalAppendData(buf) +} + +// try writing to chunks, returns (nw, error) +func (f *File) writeToChunks(buf []byte, chunks []fileChunk, updatePos bool) (int64, error) { + var numWrite int64 + for _, chunk := range chunks { + if numWrite >= int64(len(buf)) { + break + } + if chunk.Len == 0 { + continue + } + toWrite := int64(len(buf)) - numWrite + if toWrite > chunk.Len { + toWrite = chunk.Len + } + nw, err := f.OSFile.WriteAt(buf[numWrite:numWrite+toWrite], chunk.StartPos+HeaderLen) + if err != nil { + return 0, err + } + if updatePos { + if chunk.StartPos+int64(nw) > f.FileDataSize { + f.FileDataSize = chunk.StartPos + int64(nw) + } + if f.StartPos == FilePosEmpty { + f.StartPos = chunk.StartPos + } + f.EndPos = chunk.StartPos + int64(nw) - 1 + } + numWrite += int64(nw) + } + return numWrite, nil +} + +func (f *File) internalAppendData(buf []byte) error { + err := f.ensureFreeSpace(int64(len(buf))) + if err != nil { + return err + } + if len(buf) >= int(f.MaxSize) { + buf = buf[len(buf)-int(f.MaxSize):] + } + chunks := f.getFreeChunks() + // don't track nw because we know we have enough free space to write entire buf + _, err = f.writeToChunks(buf, chunks, true) + if err != nil { + return err + } + err = f.writeMeta() + if err != nil { + return err + } + return nil +} + func (f *File) AppendData(ctx context.Context, buf []byte) error { err := f.flock(ctx, syscall.LOCK_EX) if err != nil { @@ -371,42 +495,5 @@ func (f *File) AppendData(ctx context.Context, buf []byte) error { if err != nil { return err } - err = f.ensureFreeSpace(int64(len(buf))) - if err != nil { - return err - } - if len(buf) >= int(f.MaxSize) { - buf = buf[len(buf)-int(f.MaxSize):] - } - chunks := f.getFreeChunks() - numWrite := 0 - for _, chunk := range chunks { - if numWrite >= len(buf) { - break - } - if chunk.Len == 0 { - continue - } - toRead := len(buf) - numWrite - if toRead > int(chunk.Len) { - toRead = int(chunk.Len) - } - nw, err := f.OSFile.WriteAt(buf[numWrite:numWrite+toRead], chunk.StartPos+HeaderLen) - if err != nil { - return err - } - if chunk.StartPos+int64(nw) > f.FileDataSize { - f.FileDataSize = chunk.StartPos + int64(nw) - } - if f.StartPos == FilePosEmpty { - f.StartPos = chunk.StartPos - } - f.EndPos = chunk.StartPos + int64(nw) - 1 - numWrite += nw - } - err = f.writeMeta() - if err != nil { - return err - } - return nil + return f.internalAppendData(buf) } diff --git a/pkg/cirfile/cirfile_test.go b/pkg/cirfile/cirfile_test.go index 4e4f555c..1e655470 100644 --- a/pkg/cirfile/cirfile_test.go +++ b/pkg/cirfile/cirfile_test.go @@ -96,17 +96,13 @@ func TestFile(t *testing.T) { t.Fatalf("cannot append data: %v", err) } validateFileSize(t, f1Name, HeaderLen) - if f.StartPos != FilePosEmpty || f.EndPos != 0 || f.FileDataSize != 0 { - t.Fatalf("metadata error (1): %#v", f) - } + validateMeta(t, "1", f, FilePosEmpty, 0, 0, 0) err = f.AppendData(context.Background(), []byte("hello")) if err != nil { t.Fatalf("cannot append data: %v", err) } validateFileSize(t, f1Name, HeaderLen+5) - if f.StartPos != 0 || f.EndPos != 4 || f.FileDataSize != 5 { - t.Fatalf("metadata error (2): %#v", f) - } + validateMeta(t, "2", f, 0, 4, 5, 0) err = f.AppendData(context.Background(), []byte(" foo")) if err != nil { t.Fatalf("cannot append data: %v", err) From 26bd499facca42185cd9b00215037b25e2fbc340 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 19 Aug 2022 12:51:37 -0700 Subject: [PATCH 069/149] update to allow a WriteAt call for cirfiles --- pkg/cirfile/cirfile.go | 31 ++++++++++++++--- pkg/cirfile/cirfile_test.go | 66 ++++++++++++++++++++++++++++++++++++- 2 files changed, 91 insertions(+), 6 deletions(-) diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go index 5992b876..ce98b807 100644 --- a/pkg/cirfile/cirfile.go +++ b/pkg/cirfile/cirfile.go @@ -307,18 +307,24 @@ func (f *File) getFreeChunks() []fileChunk { return rtn } -// returns (realOffset, data, error) -// will only return io.EOF when len(data) == 0, otherwise will just do a short read -func (f *File) ReadNext(ctx context.Context, buf []byte, offset int64) (int64, int, error) { +func (f *File) ReadAll(ctx context.Context) (int64, []byte, error) { err := f.flock(ctx, syscall.LOCK_SH) if err != nil { - return 0, 0, err + return 0, nil, err } defer f.unflock() err = f.readMeta() if err != nil { - return 0, 0, err + return 0, nil, err } + chunks := f.getFileChunks() + curSize := totalChunksSize(chunks) + buf := make([]byte, curSize) + realOffset, nr, err := f.internalReadNext(buf, 0) + return realOffset, buf[0:nr], err +} + +func (f *File) internalReadNext(buf []byte, offset int64) (int64, int, error) { if offset < f.FileOffset { offset = f.FileOffset } @@ -347,6 +353,21 @@ func (f *File) ReadNext(ctx context.Context, buf []byte, offset int64) (int64, i return offset, numRead, nil } +// returns (realOffset, numread, error) +// will only return io.EOF when len(data) == 0, otherwise will just do a short read +func (f *File) ReadNext(ctx context.Context, buf []byte, offset int64) (int64, int, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, 0, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, 0, err + } + return f.internalReadNext(buf, offset) +} + func (f *File) ensureFreeSpace(requiredSpace int64) error { chunks := f.getFileChunks() curSpace := f.MaxSize - totalChunksSize(chunks) diff --git a/pkg/cirfile/cirfile_test.go b/pkg/cirfile/cirfile_test.go index 1e655470..d38cb09a 100644 --- a/pkg/cirfile/cirfile_test.go +++ b/pkg/cirfile/cirfile_test.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path" + "strings" "syscall" "testing" "time" @@ -28,7 +29,9 @@ func validateMeta(t *testing.T, desc string, f *File, startPos int64, endPos int func dumpFile(name string) { barr, _ := os.ReadFile(name) - fmt.Printf("<<<\n%s\n>>>", string(barr)) + str := string(barr) + str = strings.ReplaceAll(str, "\x00", ".") + fmt.Printf("%s<<<\n%s\n>>>\n", name, str) } func makeData(size int) string { @@ -216,3 +219,64 @@ func TestFlock(t *testing.T) { t.Fatalf("append error (should work fd2 was closed): %v", err) } } + +func TestWriteAt(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("cannot create cirfile: %v", err) + } + err = f.WriteAt(nil, []byte("hello\nmike"), 4) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("t"), 2) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("more"), 30) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("\n"), 19) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + dumpFile(f1Name) + err = f.WriteAt(nil, []byte("hello"), 200) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + buf := make([]byte, 10) + realOffset, nr, err := f.ReadNext(context.Background(), buf, 200) + if err != nil || realOffset != 200 || nr != 5 || string(buf[0:nr]) != "hello" { + t.Fatalf("invalid readnext: err[%v] realoffset[%d] nr[%d] buf[%s]", err, realOffset, nr, string(buf[0:nr])) + } + err = f.WriteAt(nil, []byte("0123456789\n"), 100) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + dumpFile(f1Name) + dataStr := makeData(200) + err = f.WriteAt(nil, []byte(dataStr), 50) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + dumpFile(f1Name) + + dataStr = makeData(1000) + err = f.WriteAt(nil, []byte(dataStr), 1002) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("hello\n"), 2010) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.AppendData(nil, []byte("foo\n")) + if err != nil { + t.Fatalf("appenddata error: %v", err) + } + dumpFile(f1Name) +} From 86a7cd59e628ca652dbba9d618c64b81bde3e49e Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 19 Aug 2022 15:28:32 -0700 Subject: [PATCH 070/149] use cirfile for detached commands --- main-mshell.go | 3 +++ pkg/base/base.go | 25 ------------------------- pkg/packet/packet.go | 15 +++++++++++---- pkg/shexec/shexec.go | 34 +++++++++++++++++++++++++++++++--- 4 files changed, 45 insertions(+), 32 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 6029c773..0cc8c493 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -377,6 +377,9 @@ func handleClient() (int, error) { if err != nil { return 1, err } + if runPacket.Detached { + return 1, fmt.Errorf("cannot run detached command from command line client") + } donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, shexec.StdContext{}, opts.SSHOpts, nil, opts.Debug) if err != nil { return 1, err diff --git a/pkg/base/base.go b/pkg/base/base.go index c14736d5..972d447e 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -118,22 +118,6 @@ func GetMShellHomeDir() string { return ExpandHomeDir(DefaultMShellHome) } -func GetPtyOutFile(ck CommandKey, seqNum int) (string, error) { - if err := ck.Validate("ck"); err != nil { - return "", fmt.Errorf("cannot get command files: %w", err) - } - if seqNum < 0 { - return "", fmt.Errorf("invalid seqnum, cannot be negative") - } - sessionId, cmdId := ck.Split() - sdir, err := EnsureSessionDir(sessionId) - if err != nil { - return "", err - } - base := path.Join(sdir, cmdId) - return fmt.Sprintf("%s.%d.ptyout", base, seqNum), nil -} - func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { if err := ck.Validate("ck"); err != nil { return nil, fmt.Errorf("cannot get command files: %w", err) @@ -151,15 +135,6 @@ func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { }, nil } -func MakeCommandFileNamesWithHome(mhome string, ck CommandKey) *CommandFileNames { - base := path.Join(mhome, SessionsDirBaseName, ck.GetSessionId(), ck.GetCmdId()) - return &CommandFileNames{ - PtyOutFile: base + ".ptyout", - StdinFifo: base + ".stdin", - RunnerOutFile: base + ".runout", - } -} - func CleanUpCmdFiles(sessionId string, cmdId string) error { if cmdId == "" { return fmt.Errorf("bad cmdid, cannot clean up") diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 74c7e36a..35a8d86d 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -378,6 +378,13 @@ func (p *ResponsePacketType) Err() error { return nil } +func (p *ResponsePacketType) String() string { + if p.Success { + return "response[success]" + } + return fmt.Sprintf("response[error:%s]", p.Error) +} + func MakeErrorResponsePacket(reqId string, err error) *ResponsePacketType { return &ResponsePacketType{Type: ResponsePacketStr, RespId: reqId, Error: err.Error()} } @@ -504,10 +511,10 @@ func MakeCmdStartPacket(reqId string) *CmdStartPacketType { } type TermOpts struct { - Rows int `json:"rows"` - Cols int `json:"cols"` - Term string `json:"term"` - CmdSize int64 `json:"cmdsize,omitempty"` + Rows int `json:"rows"` + Cols int `json:"cols"` + Term string `json:"term"` + MaxPtySize int64 `json:"maxptysize,omitempty"` } type RemoteFd struct { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index df740a81..5dae3e40 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -23,6 +23,7 @@ import ( "github.com/alessio/shellescape" "github.com/creack/pty" "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/scripthaus-dev/mshell/pkg/cirfile" "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" "golang.org/x/sys/unix" @@ -35,6 +36,7 @@ const MaxCols = 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 const DefaultTermType = "xterm-256color" +const DefaultMaxPtySize = 1024 * 1024 const ClientCommand = ` PATH=$PATH:~/.mshell; @@ -67,6 +69,7 @@ type ShExecType struct { FileNames *base.CommandFileNames Cmd *exec.Cmd CmdPty *os.File + MaxPtySize int64 Multiplexer *mpio.Multiplexer Detached bool DetachedOutput *packet.PacketSender @@ -933,6 +936,26 @@ func SetupSignalsForDetach() { }() } +func copyToCirFile(dest *cirfile.File, src io.Reader) error { + buf := make([]byte, 64*1024) + for { + var appendErr error + nr, readErr := src.Read(buf) + if nr > 0 { + appendErr = dest.AppendData(context.Background(), buf[0:nr]) + } + if readErr != nil && readErr != io.EOF { + return readErr + } + if appendErr != nil { + return appendErr + } + if readErr == io.EOF { + return nil + } + } +} + func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { // after Start(), any output/errors must go to DetachedOutput // close stdin, redirect stdout/stderr to /dev/null, but wait for cmdstart packet to get sent @@ -949,7 +972,7 @@ func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { if err != nil { cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to runout: %w", err)) } - ptyOutFd, err := os.OpenFile(cmd.FileNames.PtyOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) + ptyOutFile, err := cirfile.CreateCirFile(cmd.FileNames.PtyOutFile, cmd.MaxPtySize) if err != nil { cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot open ptyout file '%s': %w", cmd.FileNames.PtyOutFile, err)) // don't return (command is already running) @@ -958,8 +981,8 @@ func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { go func() { // copy pty output to .ptyout file defer close(ptyCopyDone) - defer ptyOutFd.Close() - _, copyErr := io.Copy(ptyOutFd, cmd.CmdPty) + defer ptyOutFile.Close() + copyErr := copyToCirFile(ptyOutFile, cmd.CmdPty) if copyErr != nil { cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("copying pty output to ptyout file: %w", copyErr)) } @@ -1002,6 +1025,11 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( cmd.FileNames = fileNames cmd.CmdPty = cmdPty cmd.Detached = true + if pk.TermOpts != nil && pk.TermOpts.MaxPtySize != 0 { + cmd.MaxPtySize = pk.TermOpts.MaxPtySize + } else { + cmd.MaxPtySize = DefaultMaxPtySize + } cmd.RunnerOutFd, err = os.OpenFile(fileNames.RunnerOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err != nil { return nil, nil, fmt.Errorf("cannot open runout file '%s': %w", fileNames.RunnerOutFile, err) From e26705623b3170ca3cac464907380d57e608d905 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 22 Aug 2022 15:59:03 -0700 Subject: [PATCH 071/149] add --env mode to mshell to print cwd and environment (for server initpk). needed because OSX does not support 'env -0' --- main-mshell.go | 21 ++++++++++++++ pkg/packet/packet.go | 21 +++++++------- pkg/shexec/shexec.go | 66 +++++++++++++++++++++++++++++++++++++++++++- 3 files changed, 97 insertions(+), 11 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 0cc8c493..65d7a8c9 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -425,6 +425,19 @@ func handleInstall() (int, error) { return 0, nil } +func handleEnv() (int, error) { + cwd, err := os.Getwd() + if err != nil { + return 1, err + } + fmt.Printf("%s\x00", cwd) + fullEnv := os.Environ() + for _, envLine := range fullEnv { + fmt.Printf("%s\x00", envLine) + } + return 0, nil +} + func handleUsage() { usage := ` Client Usage: mshell [opts] --ssh user@host -- [command] @@ -482,6 +495,14 @@ func main() { } else if firstArg == "--version" { fmt.Printf("mshell v%s\n", base.MShellVersion) return + } else if firstArg == "--env" { + rtnCode, err := handleEnv() + if err != nil { + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + } + if rtnCode != 0 { + os.Exit(rtnCode) + } } else if firstArg == "--single" { handleSingle() return diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 35a8d86d..380fce05 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -433,16 +433,17 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { } type InitPacketType struct { - Type string `json:"type"` - Version string `json:"version"` - MShellHomeDir string `json:"mshellhomedir,omitempty"` - HomeDir string `json:"homedir,omitempty"` - Env []string `json:"env,omitempty"` - User string `json:"user,omitempty"` - HostName string `json:"hostname,omitempty"` - NotFound bool `json:"notfound,omitempty"` - UName string `json:"uname,omitempty"` - RemoteId string `json:"remoteid,omitempty"` + Type string `json:"type"` + Version string `json:"version"` + MShellHomeDir string `json:"mshellhomedir,omitempty"` + HomeDir string `json:"homedir,omitempty"` + Cwd string `json:"cwd,omitempty"` + Env []byte `json:"env,omitempty"` // "env -0" format + User string `json:"user,omitempty"` + HostName string `json:"hostname,omitempty"` + NotFound bool `json:"notfound,omitempty"` + UName string `json:"uname,omitempty"` + RemoteId string `json:"remoteid,omitempty"` } func (*InitPacketType) GetType() string { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 5dae3e40..5cb7fb68 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -7,8 +7,10 @@ package shexec import ( + "bytes" "context" "encoding/base64" + "errors" "fmt" "io" "os" @@ -38,6 +40,8 @@ const FirstExtraFilesFdNum = 3 const DefaultTermType = "xterm-256color" const DefaultMaxPtySize = 1024 * 1024 +const GetStateTimeout = 5 * time.Second + const ClientCommand = ` PATH=$PATH:~/.mshell; which mshell > /dev/null; @@ -1095,10 +1099,70 @@ func MakeInitPacket() *packet.InitPacketType { func MakeServerInitPacket() (*packet.InitPacketType, error) { var err error initPacket := MakeInitPacket() - initPacket.Env = os.Environ() + cwd, env, err := GetCurrentState() + if err != nil { + return nil, err + } + initPacket.Cwd = cwd + initPacket.Env = env initPacket.RemoteId, err = base.GetRemoteId() if err != nil { return nil, err } return initPacket, nil } + +func parseEnv(env []byte) map[string]string { + envLines := bytes.Split(env, []byte{0}) + rtn := make(map[string]string) + for _, envLine := range envLines { + if len(envLine) == 0 { + continue + } + eqIdx := bytes.Index(envLine, []byte{'='}) + if eqIdx == -1 { + continue + } + varName := string(envLine[0:eqIdx]) + varVal := string(envLine[eqIdx+1:]) + rtn[varName] = varVal + } + return rtn +} + +func getStderr(err error) string { + exitErr, ok := err.(*exec.ExitError) + if !ok { + return "" + } + if len(exitErr.Stderr) == 0 { + return "" + } + lines := strings.SplitN(string(exitErr.Stderr), "\n", 2) + if len(lines[0]) > 100 { + return lines[0][0:100] + } + return lines[0] +} + +func GetCurrentState() (string, []byte, error) { + execFile, err := os.Executable() + if err != nil { + return "", nil, fmt.Errorf("cannot find local mshell executable: %w", err) + } + ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) + ecmd := exec.CommandContext(ctx, "bash", "-l", "-c", fmt.Sprintf("%s --env", shellescape.Quote(execFile))) + outputBytes, err := ecmd.Output() + if err != nil { + errMsg := getStderr(err) + if errMsg != "" { + return "", nil, errors.New(errMsg) + } + return "", nil, err + } + idx := bytes.Index(outputBytes, []byte{0}) + if idx == -1 { + return "", nil, fmt.Errorf("invalid current state output no NUL byte separator") + } + return string(outputBytes[0:idx]), outputBytes[idx+1:], nil +} From 7d06bc766cb1bb316be260d3af2b9dcab0323bfd Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 22 Aug 2022 16:24:53 -0700 Subject: [PATCH 072/149] rename env to env0. add envcomplete bool --- pkg/packet/packet.go | 25 +++++++++++++------------ pkg/shexec/shexec.go | 17 ++++++++++------- 2 files changed, 23 insertions(+), 19 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 380fce05..985b89b9 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -438,7 +438,7 @@ type InitPacketType struct { MShellHomeDir string `json:"mshellhomedir,omitempty"` HomeDir string `json:"homedir,omitempty"` Cwd string `json:"cwd,omitempty"` - Env []byte `json:"env,omitempty"` // "env -0" format + Env0 []byte `json:"env0,omitempty"` // "env -0" format User string `json:"user,omitempty"` HostName string `json:"hostname,omitempty"` NotFound bool `json:"notfound,omitempty"` @@ -532,17 +532,18 @@ type RunDataType struct { } type RunPacketType struct { - Type string `json:"type"` - ReqId string `json:"reqid"` - CK base.CommandKey `json:"ck"` - Command string `json:"command"` - Cwd string `json:"cwd,omitempty"` - Env map[string]string `json:"env,omitempty"` - UsePty bool `json:"usepty,omitempty"` - TermOpts *TermOpts `json:"termopts,omitempty"` - Fds []RemoteFd `json:"fds,omitempty"` - RunData []RunDataType `json:"rundata,omitempty"` - Detached bool `json:"detached,omitempty"` + Type string `json:"type"` + ReqId string `json:"reqid"` + CK base.CommandKey `json:"ck"` + Command string `json:"command"` + Cwd string `json:"cwd,omitempty"` + Env0 []byte `json:"env0,omitempty"` // in "env -0" format + EnvComplete bool `json:"envcomplete,omitempty"` // set to true if env0 is complete (the default env should not be set) + UsePty bool `json:"usepty,omitempty"` + TermOpts *TermOpts `json:"termopts,omitempty"` + Fds []RemoteFd `json:"fds,omitempty"` + RunData []RunDataType `json:"rundata,omitempty"` + Detached bool `json:"detached,omitempty"` } func (*RunPacketType) GetType() string { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 5cb7fb68..1fa381a0 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -159,9 +159,6 @@ func UpdateCmdEnv(cmd *exec.Cmd, envVars map[string]string) { if len(envVars) == 0 { return } - if cmd.Env == nil { - cmd.Env = os.Environ() - } found := make(map[string]bool) var newEnv []string for _, envStr := range cmd.Env { @@ -201,7 +198,10 @@ func MakeSimpleStaticWriterPipe(data []byte) (*os.File, error) { func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, error) { ecmd := exec.Command("bash", "-c", pk.Command) - UpdateCmdEnv(ecmd, pk.Env) + if !pk.EnvComplete { + ecmd.Env = os.Environ() + } + UpdateCmdEnv(ecmd, parseEnv0(pk.Env0)) UpdateCmdEnv(ecmd, map[string]string{"TERM": getTermType(pk)}) if pk.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(pk.Cwd) @@ -828,7 +828,10 @@ func getTermType(pk *packet.RunPacketType) string { func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { cmd := MakeShExec(pk.CK, nil) cmd.Cmd = exec.Command("bash", "-c", pk.Command) - UpdateCmdEnv(cmd.Cmd, pk.Env) + if !pk.EnvComplete { + cmd.Cmd.Env = os.Environ() + } + UpdateCmdEnv(cmd.Cmd, parseEnv0(pk.Env0)) if pk.Cwd != "" { cmd.Cmd.Dir = base.ExpandHomeDir(pk.Cwd) } @@ -1104,7 +1107,7 @@ func MakeServerInitPacket() (*packet.InitPacketType, error) { return nil, err } initPacket.Cwd = cwd - initPacket.Env = env + initPacket.Env0 = env initPacket.RemoteId, err = base.GetRemoteId() if err != nil { return nil, err @@ -1112,7 +1115,7 @@ func MakeServerInitPacket() (*packet.InitPacketType, error) { return initPacket, nil } -func parseEnv(env []byte) map[string]string { +func parseEnv0(env []byte) map[string]string { envLines := bytes.Split(env, []byte{0}) rtn := make(map[string]string) for _, envLine := range envLines { From f63a851b1a8c8cc5a8488a914f5bd5243423502e Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 22 Aug 2022 17:27:55 -0700 Subject: [PATCH 073/149] export ParseEnv0 and MakeEnv0 --- pkg/shexec/shexec.go | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 1fa381a0..32fc3448 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -201,7 +201,7 @@ func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, if !pk.EnvComplete { ecmd.Env = os.Environ() } - UpdateCmdEnv(ecmd, parseEnv0(pk.Env0)) + UpdateCmdEnv(ecmd, ParseEnv0(pk.Env0)) UpdateCmdEnv(ecmd, map[string]string{"TERM": getTermType(pk)}) if pk.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(pk.Cwd) @@ -831,7 +831,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*S if !pk.EnvComplete { cmd.Cmd.Env = os.Environ() } - UpdateCmdEnv(cmd.Cmd, parseEnv0(pk.Env0)) + UpdateCmdEnv(cmd.Cmd, ParseEnv0(pk.Env0)) if pk.Cwd != "" { cmd.Cmd.Dir = base.ExpandHomeDir(pk.Cwd) } @@ -1115,7 +1115,7 @@ func MakeServerInitPacket() (*packet.InitPacketType, error) { return initPacket, nil } -func parseEnv0(env []byte) map[string]string { +func ParseEnv0(env []byte) map[string]string { envLines := bytes.Split(env, []byte{0}) rtn := make(map[string]string) for _, envLine := range envLines { @@ -1133,6 +1133,17 @@ func parseEnv0(env []byte) map[string]string { return rtn } +func MakeEnv0(envMap map[string]string) []byte { + var buf bytes.Buffer + for envName, envVal := range envMap { + buf.WriteString(envName) + buf.WriteByte('=') + buf.WriteString(envVal) + buf.WriteByte(0) + } + return buf.Bytes() +} + func getStderr(err error) string { exitErr, ok := err.(*exec.ExitError) if !ok { From 09a072f9007a6670024603fea053203abd6004fd Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 24 Aug 2022 02:11:49 -0700 Subject: [PATCH 074/149] run localhost mshell with cwd at HOME, not in current directory --- pkg/shexec/shexec.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 32fc3448..dbfd35e9 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -391,7 +391,12 @@ func (opts SSHOpts) MakeMShellSingleCmd() (*exec.Cmd, error) { func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { remoteCommand = strings.TrimSpace(remoteCommand) if opts.SSHHost == "" { + homeDir, _ := os.UserHomeDir() // ignore error + if homeDir == "" { + homeDir = "/" + } ecmd := exec.Command("bash", "-c", remoteCommand) + ecmd.Dir = homeDir return ecmd } else { var moreSSHOpts []string From 8274d19e0ba370de954b29ed40f419648fef31bb Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 24 Aug 2022 18:57:13 -0700 Subject: [PATCH 075/149] use interactive shell for environment --- pkg/shexec/shexec.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index dbfd35e9..7c463bc9 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -1170,7 +1170,7 @@ func GetCurrentState() (string, []byte, error) { return "", nil, fmt.Errorf("cannot find local mshell executable: %w", err) } ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - ecmd := exec.CommandContext(ctx, "bash", "-l", "-c", fmt.Sprintf("%s --env", shellescape.Quote(execFile))) + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env", shellescape.Quote(execFile))) outputBytes, err := ecmd.Output() if err != nil { errMsg := getStderr(err) From db993cf00fc18fd2295c594d26b2fb535d304b12 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 24 Aug 2022 21:33:50 -0700 Subject: [PATCH 076/149] add scripthaus.md --- scripthaus.md | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) create mode 100644 scripthaus.md diff --git a/scripthaus.md b/scripthaus.md new file mode 100644 index 00000000..6752bda6 --- /dev/null +++ b/scripthaus.md @@ -0,0 +1,16 @@ + +```bash +# @scripthaus command build +go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell main-mshell.go +``` + +```bash +# @scripthaus command fullbuild +go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell main-mshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.linux.amd64 main-mshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.linux.arm64 main-mshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.darwin.amd64 main-mshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.darwin.arm64 main-mshell.go +``` + + From 8e4b02cec4329d54fcf3393342d4468fe565c1f2 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 30 Aug 2022 00:23:03 -0700 Subject: [PATCH 077/149] run mshell env in interactive bash shell with a new pty. this picks up special interactive environment vars from bash startup scripts --- pkg/shexec/shexec.go | 50 +++++++++++++++++++++++++++++++++++++------- 1 file changed, 43 insertions(+), 7 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 7c463bc9..1929baaf 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -10,7 +10,6 @@ import ( "bytes" "context" "encoding/base64" - "errors" "fmt" "io" "os" @@ -832,7 +831,11 @@ func getTermType(pk *packet.RunPacketType) string { func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { cmd := MakeShExec(pk.CK, nil) - cmd.Cmd = exec.Command("bash", "-c", pk.Command) + if pk.UsePty { + cmd.Cmd = exec.Command("bash", "-i", "-c", pk.Command) + } else { + cmd.Cmd = exec.Command("bash", "-c", pk.Command) + } if !pk.EnvComplete { cmd.Cmd.Env = os.Environ() } @@ -1164,6 +1167,43 @@ func getStderr(err error) string { return lines[0] } +func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { + ecmd.Env = os.Environ() + UpdateCmdEnv(ecmd, map[string]string{"TERM": DefaultTermType}) + cmdPty, cmdTty, err := pty.Open() + if err != nil { + return nil, fmt.Errorf("opening new pty: %w", err) + } + pty.Setsize(cmdPty, &pty.Winsize{Rows: DefaultRows, Cols: DefaultCols}) + ecmd.Stdin = cmdTty + ecmd.Stdout = cmdTty + ecmd.Stderr = cmdTty + ecmd.SysProcAttr = &syscall.SysProcAttr{} + ecmd.SysProcAttr.Setsid = true + ecmd.SysProcAttr.Setctty = true + err = ecmd.Start() + if err != nil { + cmdTty.Close() + cmdPty.Close() + return nil, err + } + cmdTty.Close() + defer cmdPty.Close() + ioDone := make(chan bool) + var outputBuf bytes.Buffer + go func() { + // ignore error (/dev/ptmx has read error when process is done) + io.Copy(&outputBuf, cmdPty) + close(ioDone) + }() + exitErr := ecmd.Wait() + if exitErr != nil { + return nil, exitErr + } + <-ioDone + return outputBuf.Bytes(), nil +} + func GetCurrentState() (string, []byte, error) { execFile, err := os.Executable() if err != nil { @@ -1171,12 +1211,8 @@ func GetCurrentState() (string, []byte, error) { } ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env", shellescape.Quote(execFile))) - outputBytes, err := ecmd.Output() + outputBytes, err := runSimpleCmdInPty(ecmd) if err != nil { - errMsg := getStderr(err) - if errMsg != "" { - return "", nil, errors.New(errMsg) - } return "", nil, err } idx := bytes.Index(outputBytes, []byte{0}) From 3936db0429a890a82b9dba17b03c2902c170ddb8 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 31 Aug 2022 12:45:59 -0700 Subject: [PATCH 078/149] allow nil for display --- pkg/packet/packet.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 985b89b9..fc2e0fd7 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -589,6 +589,9 @@ type PacketType interface { } func AsString(pk PacketType) string { + if pk == nil { + return "nil" + } if s, ok := pk.(fmt.Stringer); ok { return s.String() } From 39dacb988a004e1ab957bf882f98bdbf50771a8b Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 3 Sep 2022 23:26:57 -0700 Subject: [PATCH 079/149] default term rows should be 24 (not 25), add min/max values as well --- pkg/base/base.go | 10 ++++++++++ pkg/shexec/shexec.go | 24 +++++++++++------------- 2 files changed, 21 insertions(+), 13 deletions(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index 972d447e..d89864cf 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -269,3 +269,13 @@ func GetRemoteId() (string, error) { return uuidStr, nil } } + +func BoundInt(ival int, minVal int, maxVal int) int { + if ival < minVal { + return minVal + } + if ival > maxVal { + return maxVal + } + return ival +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 1929baaf..a7879ebd 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -30,10 +30,12 @@ import ( "golang.org/x/sys/unix" ) -const DefaultRows = 25 -const DefaultCols = 80 -const MaxRows = 1024 -const MaxCols = 1024 +const DefaultTermRows = 24 +const DefaultTermCols = 80 +const MinTermRows = 2 +const MinTermCols = 10 +const MaxTermRows = 1024 +const MaxTermCols = 1024 const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 const DefaultTermType = "xterm-256color" @@ -319,15 +321,11 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { } func GetWinsize(p *packet.RunPacketType) *pty.Winsize { - rows := DefaultRows - cols := DefaultCols + rows := DefaultTermRows + cols := DefaultTermCols if p.TermOpts != nil { - if p.TermOpts.Rows > 0 && p.TermOpts.Rows <= MaxRows { - rows = p.TermOpts.Rows - } - if p.TermOpts.Cols > 0 && p.TermOpts.Cols <= MaxCols { - cols = p.TermOpts.Cols - } + rows = base.BoundInt(p.TermOpts.Rows, MinTermRows, MaxTermRows) + cols = base.BoundInt(p.TermOpts.Cols, MinTermCols, MaxTermCols) } return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} } @@ -1174,7 +1172,7 @@ func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { if err != nil { return nil, fmt.Errorf("opening new pty: %w", err) } - pty.Setsize(cmdPty, &pty.Winsize{Rows: DefaultRows, Cols: DefaultCols}) + pty.Setsize(cmdPty, &pty.Winsize{Rows: DefaultTermRows, Cols: DefaultTermCols}) ecmd.Stdin = cmdTty ecmd.Stdout = cmdTty ecmd.Stderr = cmdTty From 57b54198e5fb1c8a8b4e9c5ce6600cc9ac830651 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 3 Sep 2022 23:38:35 -0700 Subject: [PATCH 080/149] limit maxptysize --- pkg/base/base.go | 10 ++++++++++ pkg/shexec/shexec.go | 9 +++++---- 2 files changed, 15 insertions(+), 4 deletions(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index d89864cf..1f30d632 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -279,3 +279,13 @@ func BoundInt(ival int, minVal int, maxVal int) int { } return ival } + +func BoundInt64(ival int64, minVal int64, maxVal int64) int64 { + if ival < minVal { + return minVal + } + if ival > maxVal { + return maxVal + } + return ival +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index a7879ebd..5a841a88 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -40,6 +40,8 @@ const MaxFdNum = 1023 const FirstExtraFilesFdNum = 3 const DefaultTermType = "xterm-256color" const DefaultMaxPtySize = 1024 * 1024 +const MinMaxPtySize = 16 * 1024 +const MaxMaxPtySize = 100 * 1024 * 1024 const GetStateTimeout = 5 * time.Second @@ -1038,10 +1040,9 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( cmd.FileNames = fileNames cmd.CmdPty = cmdPty cmd.Detached = true - if pk.TermOpts != nil && pk.TermOpts.MaxPtySize != 0 { - cmd.MaxPtySize = pk.TermOpts.MaxPtySize - } else { - cmd.MaxPtySize = DefaultMaxPtySize + cmd.MaxPtySize = DefaultMaxPtySize + if pk.TermOpts != nil && pk.TermOpts.MaxPtySize > 0 { + cmd.MaxPtySize = base.BoundInt64(pk.TermOpts.MaxPtySize, MinMaxPtySize, MaxMaxPtySize) } cmd.RunnerOutFd, err = os.OpenFile(fileNames.RunnerOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err != nil { From 082fb7a8b4313520bf34549e774713b4b41c9063 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 5 Sep 2022 16:32:08 -0700 Subject: [PATCH 081/149] updates to inputpacket. inputpacket is split between datapacket and specialinputpacket --- pkg/cirfile/cirfile.go | 1 + pkg/packet/packet.go | 71 +++++++++++++++++++++--------------------- 2 files changed, 37 insertions(+), 35 deletions(-) diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go index ce98b807..51c8cbe7 100644 --- a/pkg/cirfile/cirfile.go +++ b/pkg/cirfile/cirfile.go @@ -307,6 +307,7 @@ func (f *File) getFreeChunks() []fileChunk { return rtn } +// returns (offset, data, err) func (f *File) ReadAll(ctx context.Context) (int64, []byte, error) { err := f.flock(ctx, syscall.LOCK_SH) if err != nil { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index fc2e0fd7..3584d5e8 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -32,25 +32,25 @@ 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" - InputPacketStr = "input" // command - CompGenPacketStr = "compgen" // rpc + 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 ) const PacketSenderQueueSize = 20 @@ -73,7 +73,7 @@ func init() { TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) - TypeStrToFactory[InputPacketStr] = reflect.TypeOf(InputPacketType{}) + TypeStrToFactory[SpecialInputPacketStr] = reflect.TypeOf(SpecialInputPacketType{}) TypeStrToFactory[DataPacketStr] = reflect.TypeOf(DataPacketType{}) TypeStrToFactory[DataAckPacketStr] = reflect.TypeOf(DataAckPacketType{}) TypeStrToFactory[DataEndPacketStr] = reflect.TypeOf(DataEndPacketType{}) @@ -92,7 +92,7 @@ func init() { var _ CommandPacketType = (*DataPacketType)(nil) var _ CommandPacketType = (*DataAckPacketType)(nil) var _ CommandPacketType = (*CmdDonePacketType)(nil) - var _ CommandPacketType = (*InputPacketType)(nil) + var _ CommandPacketType = (*SpecialInputPacketType)(nil) } func RegisterPacketType(typeStr string, rtype reflect.Type) { @@ -238,29 +238,30 @@ func MakeDataAckPacket() *DataAckPacketType { return &DataAckPacketType{Type: DataAckPacketStr} } -// InputData gets written to PTY directly +type WinSize struct { + Rows int `json:"rows"` + Cols int `json:"cols"` +} + // SigNum gets sent to process via a signal // WinSize, if set, will run TIOCSWINSZ to set size, and then send SIGWINCH -type InputPacketType struct { - Type string `json:"type"` - CK base.CommandKey `json:"ck"` - RemoteId string `json:"remoteid"` - InputData64 string `json:"inputdata"` - SigNum int `json:"signum,omitempty"` - WinSizeRows int `json:"winsizerows"` - WinSizeCols int `json:"winsizecols"` +type SpecialInputPacketType struct { + Type string `json:"type"` + CK base.CommandKey `json:"ck"` + SigNum int `json:"signum,omitempty"` + WinSize *WinSize `json:"winsize,omitempty"` } -func (*InputPacketType) GetType() string { - return InputPacketStr +func (*SpecialInputPacketType) GetType() string { + return SpecialInputPacketStr } -func (p *InputPacketType) GetCK() base.CommandKey { +func (p *SpecialInputPacketType) GetCK() base.CommandKey { return p.CK } -func MakeInputPacket() *InputPacketType { - return &InputPacketType{Type: InputPacketStr} +func MakeSpecialInputPacket() *SpecialInputPacketType { + return &SpecialInputPacketType{Type: SpecialInputPacketStr} } type UntailCmdPacketType struct { From 4f4e12c00ab4b6c31c29c044a42a1190d6ede3b3 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 6 Sep 2022 12:57:54 -0700 Subject: [PATCH 082/149] add single-from-server option to mshell, send message packets with ck back to server, report unknown packets back to server --- main-mshell.go | 15 ++++++++++++--- pkg/packet/packet.go | 15 +++++++++++++-- pkg/server/server.go | 2 +- pkg/shexec/shexec.go | 19 ++++++++++++++----- 4 files changed, 40 insertions(+), 11 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 65d7a8c9..014b5a76 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -153,7 +153,7 @@ func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType return nil, fmt.Errorf("no run packet received") } -func handleSingle() { +func handleSingle(fromServer bool) { packetParser := packet.MakePacketParser(os.Stdin) sender := packet.MakePacketSender(os.Stdout) defer func() { @@ -175,6 +175,12 @@ func handleSingle() { sender.SendErrorResponse(runPacket.ReqId, err) return } + if fromServer { + err = runPacket.CK.Validate("run packet") + if err != nil { + sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("run packets from server must have a CK: %v", err)) + } + } if runPacket.Detached { cmd, startPk, err := shexec.RunCommandDetached(runPacket, sender) if err != nil { @@ -187,7 +193,7 @@ func handleSingle() { cmd.DetachedWait(startPk) return } else { - cmd, err := shexec.RunCommandSimple(runPacket, sender) + cmd, err := shexec.RunCommandSimple(runPacket, sender, true) if err != nil { sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("error running command: %w", err)) return @@ -504,7 +510,10 @@ func main() { os.Exit(rtnCode) } } else if firstArg == "--single" { - handleSingle() + handleSingle(false) + return + } else if firstArg == "--single-from-server" { + handleSingle(true) return } else if firstArg == "--server" { rtnCode, err := server.RunServer() diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 3584d5e8..3ec863b8 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -412,8 +412,9 @@ func MakeRawPacket(val string) *RawPacketType { } type MessagePacketType struct { - Type string `json:"type"` - Message string `json:"message"` + Type string `json:"type"` + CK base.CommandKey `json:"ck,omitempty"` + Message string `json:"message"` } func (*MessagePacketType) GetType() string { @@ -813,7 +814,17 @@ func (DefaultUPR) UnknownPacket(pk PacketType) { } else { fmt.Fprintf(os.Stderr, "[error] invalid packet received '%s'", AsExtType(pk)) } +} +type MessageUPR struct { + CK base.CommandKey + Sender *PacketSender +} + +func (upr MessageUPR) UnknownPacket(pk PacketType) { + msg := FmtMessagePacket("[error] invalid packet received %s", AsString(pk)) + msg.CK = upr.CK + upr.Sender.SendPacket(msg) } // todo: clean hanging entries in RunMap when in server mode diff --git a/pkg/server/server.go b/pkg/server/server.go index c1cd6bb4..ad4b3367 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -161,7 +161,7 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } - ecmd, err := shexec.SSHOpts{}.MakeMShellSingleCmd() + ecmd, err := shexec.SSHOpts{}.MakeMShellSingleCmd(true) if err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 5a841a88..a7080007 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -375,13 +375,18 @@ func (opts SSHOpts) MakeMShellServerCmd() (*exec.Cmd, error) { return ecmd, nil } -func (opts SSHOpts) MakeMShellSingleCmd() (*exec.Cmd, error) { +func (opts SSHOpts) MakeMShellSingleCmd(fromServer bool) (*exec.Cmd, error) { if opts.SSHHost == "" { execFile, err := os.Executable() if err != nil { return nil, fmt.Errorf("cannot find local mshell executable: %w", err) } - ecmd := exec.Command(execFile, "--single") + var ecmd *exec.Cmd + if fromServer { + ecmd = exec.Command(execFile, "--single-from-server") + } else { + ecmd = exec.Command(execFile, "--single") + } return ecmd, nil } return opts.MakeSSHExecCmd(ClientCommand), nil @@ -657,7 +662,7 @@ func HasDupStdin(fds []packet.RemoteFd) bool { func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdContext, sshOpts SSHOpts, upr packet.UnknownPacketReporter, debug bool) (*packet.CmdDonePacketType, error) { cmd := MakeShExec(runPacket.CK, upr) - ecmd, err := sshOpts.MakeMShellSingleCmd() + ecmd, err := sshOpts.MakeMShellSingleCmd(false) if err != nil { return nil, err } @@ -829,8 +834,12 @@ func getTermType(pk *packet.RunPacketType) string { return termType } -func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - cmd := MakeShExec(pk.CK, nil) +func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (*ShExecType, error) { + var upr packet.UnknownPacketReporter + if fromServer { + upr = packet.MessageUPR{CK: pk.CK, Sender: sender} + } + cmd := MakeShExec(pk.CK, upr) if pk.UsePty { cmd.Cmd = exec.Command("bash", "-i", "-c", pk.Command) } else { From 670f54a5b4e0026e20878846f9accb5ac18b6222 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 6 Sep 2022 13:58:07 -0700 Subject: [PATCH 083/149] checkpoint --- pkg/mpio/mpio.go | 35 +++++++++++++++++++++++++++++++++++ pkg/shexec/shexec.go | 2 ++ 2 files changed, 37 insertions(+) diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index f16a411c..017ce6de 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -12,7 +12,9 @@ import ( "io" "os" "sync" + "syscall" + "github.com/creack/pty" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" ) @@ -29,6 +31,8 @@ type Multiplexer struct { FdWriters map[int]*FdWriter // synchronized RunData map[int]*FdReader // synchronized CloseAfterStart []*os.File // synchronized + PtyFd *os.File + CmdProc *os.Process Sender *packet.PacketSender Input *packet.PacketParser @@ -51,6 +55,12 @@ func MakeMultiplexer(ck base.CommandKey, upr packet.UnknownPacketReporter) *Mult } } +func (m *Multiplexer) SetPtyFd(ptyFd *os.File) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.PtyFd = ptyFd +} + func (m *Multiplexer) Close() { m.Lock.Lock() defer m.Lock.Unlock() @@ -220,11 +230,36 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { donePacket := pk.(*packet.CmdDonePacketType) return donePacket } + if pk.GetType() == packet.SpecialInputPacketStr { + inputPacket := pk.(*packet.SpecialInputPacketType) + m.processSpecialInputPacket(inputPacket) + } m.UPR.UnknownPacket(pk) } return nil } +func (m *Multiplexer) processSpecialInputPacket(pk *packet.SpecialInputPacketType) { + m.Lock.Lock() + ptyFd := m.PtyFd + cmdProc := m.CmdProc + m.Lock.Unlock() + if ptyFd == nil { + // no pty, maybe send a message back to server, but the server always starts with a pty, so this shouldn't be an issue + return + } + if pk.WinSize != nil { + winSize := &pty.Winsize{ + //Rows: base.BoundInt(pk.WinSize.Rows, shexec.MinTermRows, shexec.MaxTermRows), + //Cols: base.BoundInt(pk.Winsize.Cols, shexec.MinTermCols, shexec.MaxTermCols), + } + pty.Setsize(ptyFd, winSize) + if cmdProc != nil { + cmdProc.Signal(syscall.SIGWINCH) + } + } +} + func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { realData, err := base64.StdEncoding.DecodeString(dataPacket.Data64) if err != nil { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index a7080007..378dafa9 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -869,6 +869,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro cmdTty.Close() }() cmd.CmdPty = cmdPty + cmd.Multiplexer.SetPtyFd(cmdPty) UpdateCmdEnv(cmd.Cmd, map[string]string{"TERM": getTermType(pk)}) } if cmdTty != nil { @@ -1050,6 +1051,7 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( cmd.CmdPty = cmdPty cmd.Detached = true cmd.MaxPtySize = DefaultMaxPtySize + cmd.Multiplexer.SetPtyFd(cmdPty) if pk.TermOpts != nil && pk.TermOpts.MaxPtySize > 0 { cmd.MaxPtySize = base.BoundInt64(pk.TermOpts.MaxPtySize, MinMaxPtySize, MaxMaxPtySize) } From ec143de8b486b8f9dc7f13bd1660d377d713e01f Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 6 Sep 2022 16:40:41 -0700 Subject: [PATCH 084/149] handle term resize / SIGWINCH, move from mpio to shexec (used UnknownPacketReporter). change signum to signame for cross-system compatibility --- pkg/mpio/mpio.go | 35 --------------------------- pkg/packet/packet.go | 2 +- pkg/shexec/shexec.go | 56 +++++++++++++++++++++++++++++++++++++------- 3 files changed, 49 insertions(+), 44 deletions(-) diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 017ce6de..f16a411c 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -12,9 +12,7 @@ import ( "io" "os" "sync" - "syscall" - "github.com/creack/pty" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" ) @@ -31,8 +29,6 @@ type Multiplexer struct { FdWriters map[int]*FdWriter // synchronized RunData map[int]*FdReader // synchronized CloseAfterStart []*os.File // synchronized - PtyFd *os.File - CmdProc *os.Process Sender *packet.PacketSender Input *packet.PacketParser @@ -55,12 +51,6 @@ func MakeMultiplexer(ck base.CommandKey, upr packet.UnknownPacketReporter) *Mult } } -func (m *Multiplexer) SetPtyFd(ptyFd *os.File) { - m.Lock.Lock() - defer m.Lock.Unlock() - m.PtyFd = ptyFd -} - func (m *Multiplexer) Close() { m.Lock.Lock() defer m.Lock.Unlock() @@ -230,36 +220,11 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { donePacket := pk.(*packet.CmdDonePacketType) return donePacket } - if pk.GetType() == packet.SpecialInputPacketStr { - inputPacket := pk.(*packet.SpecialInputPacketType) - m.processSpecialInputPacket(inputPacket) - } m.UPR.UnknownPacket(pk) } return nil } -func (m *Multiplexer) processSpecialInputPacket(pk *packet.SpecialInputPacketType) { - m.Lock.Lock() - ptyFd := m.PtyFd - cmdProc := m.CmdProc - m.Lock.Unlock() - if ptyFd == nil { - // no pty, maybe send a message back to server, but the server always starts with a pty, so this shouldn't be an issue - return - } - if pk.WinSize != nil { - winSize := &pty.Winsize{ - //Rows: base.BoundInt(pk.WinSize.Rows, shexec.MinTermRows, shexec.MaxTermRows), - //Cols: base.BoundInt(pk.Winsize.Cols, shexec.MinTermCols, shexec.MaxTermCols), - } - pty.Setsize(ptyFd, winSize) - if cmdProc != nil { - cmdProc.Signal(syscall.SIGWINCH) - } - } -} - func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { realData, err := base64.StdEncoding.DecodeString(dataPacket.Data64) if err != nil { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 3ec863b8..2d13ad5e 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -248,7 +248,7 @@ type WinSize struct { type SpecialInputPacketType struct { Type string `json:"type"` CK base.CommandKey `json:"ck"` - SigNum int `json:"signum,omitempty"` + SigName string `json:"signame,omitempty"` // passed to unix.SignalNum (needs 'SIG' prefix, e.g. "SIGTERM") WinSize *WinSize `json:"winsize,omitempty"` } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 378dafa9..09ca45ec 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -17,7 +17,6 @@ import ( "os/signal" "os/user" "strings" - "sync" "syscall" "time" @@ -70,7 +69,6 @@ const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` type ShExecType struct { - Lock *sync.Mutex StartTs time.Time CK base.CommandKey FileNames *base.CommandFileNames @@ -81,6 +79,7 @@ type ShExecType struct { Detached bool DetachedOutput *packet.PacketSender RunnerOutFd *os.File + MsgSender *packet.PacketSender // where to send out-of-band messages back to calling proceess } type StdContext struct{} @@ -118,9 +117,50 @@ type FdContext interface { GetReader(fdNum int) io.ReadCloser } +type ShExecUPR struct { + ShExec *ShExecType + UPR packet.UnknownPacketReporter +} + +func (s *ShExecType) processSpecialInputPacket(pk *packet.SpecialInputPacketType) error { + if pk.WinSize != nil { + if s.CmdPty == nil { + return fmt.Errorf("cannot change winsize, cmd was not started with a pty") + } + winSize := &pty.Winsize{ + Rows: uint16(base.BoundInt(pk.WinSize.Rows, MinTermRows, MaxTermRows)), + Cols: uint16(base.BoundInt(pk.WinSize.Cols, MinTermCols, MaxTermCols)), + } + pty.Setsize(s.CmdPty, winSize) + s.Cmd.Process.Signal(syscall.SIGWINCH) + } + if pk.SigName != "" { + sigNum := unix.SignalNum(pk.SigName) + if sigNum == 0 { + return fmt.Errorf("error signal %q not found, cannot send", pk.SigName) + } + } + return nil +} + +func (s ShExecUPR) UnknownPacket(pk packet.PacketType) { + if pk.GetType() == packet.SpecialInputPacketStr { + inputPacket := pk.(*packet.SpecialInputPacketType) + err := s.ShExec.processSpecialInputPacket(inputPacket) + if err != nil && s.ShExec.MsgSender != nil { + msg := packet.MakeMessagePacket(err.Error()) + msg.CK = s.ShExec.CK + s.ShExec.MsgSender.SendPacket(msg) + } + return + } + if s.UPR != nil { + s.UPR.UnknownPacket(pk) + } +} + func MakeShExec(ck base.CommandKey, upr packet.UnknownPacketReporter) *ShExecType { return &ShExecType{ - Lock: &sync.Mutex{}, StartTs: time.Now(), CK: ck, Multiplexer: mpio.MakeMultiplexer(ck, upr), @@ -835,11 +875,13 @@ func getTermType(pk *packet.RunPacketType) string { } func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (*ShExecType, error) { - var upr packet.UnknownPacketReporter + cmd := MakeShExec(pk.CK, nil) if fromServer { - upr = packet.MessageUPR{CK: pk.CK, Sender: sender} + msgUpr := packet.MessageUPR{CK: pk.CK, Sender: sender} + upr := ShExecUPR{ShExec: cmd, UPR: msgUpr} + cmd.Multiplexer.UPR = upr + cmd.MsgSender = sender } - cmd := MakeShExec(pk.CK, upr) if pk.UsePty { cmd.Cmd = exec.Command("bash", "-i", "-c", pk.Command) } else { @@ -869,7 +911,6 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro cmdTty.Close() }() cmd.CmdPty = cmdPty - cmd.Multiplexer.SetPtyFd(cmdPty) UpdateCmdEnv(cmd.Cmd, map[string]string{"TERM": getTermType(pk)}) } if cmdTty != nil { @@ -1051,7 +1092,6 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( cmd.CmdPty = cmdPty cmd.Detached = true cmd.MaxPtySize = DefaultMaxPtySize - cmd.Multiplexer.SetPtyFd(cmdPty) if pk.TermOpts != nil && pk.TermOpts.MaxPtySize > 0 { cmd.MaxPtySize = base.BoundInt64(pk.TermOpts.MaxPtySize, MinMaxPtySize, MaxMaxPtySize) } From 53d710f70918e873659bd60e8520a710e744a683 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 13 Sep 2022 17:10:18 -0700 Subject: [PATCH 085/149] option to send ssh errors to tty instead of stderr --- pkg/shexec/shexec.go | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 09ca45ec..98b2e4cf 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -373,10 +373,11 @@ func GetWinsize(p *packet.RunPacketType) *pty.Winsize { } type SSHOpts struct { - SSHHost string - SSHOptsStr string - SSHIdentity string - SSHUser string + SSHHost string + SSHOptsStr string + SSHIdentity string + SSHUser string + SSHErrorsToTty bool } type InstallOpts struct { @@ -453,7 +454,11 @@ func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { moreSSHOpts = append(moreSSHOpts, userOpt) } // note that SSHOptsStr is *not* escaped - sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) + var errFdStr string + if opts.SSHErrorsToTty { + errFdStr = "-E /dev/tty" + } + sshCmd := fmt.Sprintf("ssh %s %s %s %s %s", errFdStr, strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) ecmd := exec.Command("bash", "-c", sshCmd) return ecmd } From df80d2ac08a3fc7ec7b9aa357e978adb5fd59b34 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 16 Sep 2022 12:27:14 -0700 Subject: [PATCH 086/149] allow MakeClientProc to be canceled with a context --- pkg/server/server.go | 2 +- pkg/shexec/client.go | 14 +++++++++++--- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index ad4b3367..d0cc680f 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -166,7 +166,7 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } - cproc, _, err := shexec.MakeClientProc(ecmd) + cproc, _, err := shexec.MakeClientProc(context.Background(), ecmd) if err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("starting mshell client: %s", err)) return diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index fb6e6e4e..c14b17e5 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -1,6 +1,7 @@ package shexec import ( + "context" "fmt" "io" "os/exec" @@ -24,7 +25,7 @@ type ClientProc struct { } // returns (clientproc, uname, error) -func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, string, error) { +func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, string, error) { inputWriter, err := ecmd.StdinPipe() if err != nil { return nil, "", fmt.Errorf("creating stdin pipe: %v", err) @@ -55,7 +56,15 @@ func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, string, error) { Input: sender, Output: packetParser, } - for pk := range packetParser.MainCh { + + var pk packet.PacketType + select { + case pk = <-packetParser.MainCh: + case <-ctx.Done(): + cproc.Close() + return nil, "", ctx.Err() + } + if pk != nil { if pk.GetType() != packet.InitPacketStr { cproc.Close() return nil, "", fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) @@ -70,7 +79,6 @@ func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, string, error) { return nil, initPk.UName, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) } cproc.InitPk = initPk - break } if cproc.InitPk == nil { cproc.Close() From 9702fb648a57e4011f0477c417709b8ff359a781 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 20 Sep 2022 14:15:39 -0700 Subject: [PATCH 087/149] getsessionsdir --- pkg/base/base.go | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/pkg/base/base.go b/pkg/base/base.go index 1f30d632..9a08cdda 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -157,6 +157,12 @@ func CleanUpCmdFiles(sessionId string, cmdId string) error { return err } +func GetSessionsDir() string { + mhome := GetMShellHomeDir() + sdir := path.Join(mhome, SessionsDirBaseName) + return sdir +} + func EnsureSessionDir(sessionId string) (string, error) { if sessionId == "" { return "", fmt.Errorf("Bad sessionid, cannot be empty") From ca61597a190cc789d22af2f02ae9c0ae34505b9a Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 21 Sep 2022 23:26:53 -0700 Subject: [PATCH 088/149] send uname in initpk. remote space between uname parts --- pkg/shexec/shexec.go | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 98b2e4cf..aa9d683e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -16,6 +16,7 @@ import ( "os/exec" "os/signal" "os/user" + "runtime" "strings" "syscall" "time" @@ -49,14 +50,14 @@ PATH=$PATH:~/.mshell; which mshell > /dev/null; if [[ "$?" -ne 0 ]] then - printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s | %s\"}\n" "$(uname -s)" "$(uname -m)" + printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)" else mshell --single fi ` const InstallCommand = ` -printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s | %s\"}\n" "$(uname -s)" "$(uname -m)"; +printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)"; mkdir -p ~/.mshell/; cat > ~/.mshell/mshell.temp; mv ~/.mshell/mshell.temp ~/.mshell/mshell; @@ -852,7 +853,7 @@ func DetectGoArch(uname string) (string, string, error) { goarch := "" if archVal == "x86_64" || archVal == "i686" || archVal == "amd64" { goarch = "amd64" - } else if archVal == "aarch64" || archVal == "amd64" { + } else if archVal == "aarch64" || archVal == "arm64" { goarch = "arm64" } if goarch == "" { @@ -1159,6 +1160,7 @@ func MakeInitPacket() *packet.InitPacketType { initPacket.User = user.Username } initPacket.HostName, _ = os.Hostname() + initPacket.UName = fmt.Sprintf("%s|%s", runtime.GOOS, runtime.GOARCH) return initPacket } From aa1542cfc09bd369b12f2a07f0423ab3e4285aa7 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 21 Sep 2022 23:51:16 -0700 Subject: [PATCH 089/149] don't overwrite mshell if no input received on stdin during install --- pkg/shexec/shexec.go | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index aa9d683e..b8ae79e7 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -59,10 +59,13 @@ fi const InstallCommand = ` printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)"; mkdir -p ~/.mshell/; -cat > ~/.mshell/mshell.temp; -mv ~/.mshell/mshell.temp ~/.mshell/mshell; -chmod a+x ~/.mshell/mshell; -~/.mshell/mshell --single --version +cat > ~/.mshell/mshell.temp; +if [[ -s ~/.mshell/mshell.temp ]] +then + mv ~/.mshell/mshell.temp ~/.mshell/mshell; + chmod a+x ~/.mshell/mshell; + ~/.mshell/mshell --single --version +fi ` const RunCommandFmt = `%s` From 4550e18b6be92eecb3620af2fbff1de1ff4e2372 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 24 Sep 2022 13:53:19 -0700 Subject: [PATCH 090/149] version string will now be a real semantic version --- go.mod | 12 +++++++----- go.sum | 2 ++ main-mshell.go | 2 +- pkg/base/base.go | 2 +- pkg/shexec/client.go | 2 +- pkg/shexec/shexec.go | 2 +- 6 files changed, 13 insertions(+), 9 deletions(-) diff --git a/go.mod b/go.mod index 735e32fd..02a9eb35 100644 --- a/go.mod +++ b/go.mod @@ -3,9 +3,11 @@ module github.com/scripthaus-dev/mshell go 1.17 require ( - github.com/alessio/shellescape v1.4.1 // indirect - github.com/creack/pty v1.1.18 // indirect - github.com/fsnotify/fsnotify v1.5.4 // indirect - github.com/google/uuid v1.3.0 // indirect - golang.org/x/sys v0.0.0-20220412211240-33da011f77ad // indirect + github.com/alessio/shellescape v1.4.1 + github.com/creack/pty v1.1.18 + github.com/fsnotify/fsnotify v1.5.4 + github.com/google/uuid v1.3.0 + golang.org/x/sys v0.0.0-20220412211240-33da011f77ad ) + +require github.com/Masterminds/semver/v3 v3.1.1 // indirect diff --git a/go.sum b/go.sum index d802966a..32656449 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +github.com/Masterminds/semver/v3 v3.1.1 h1:hLg3sBzpNErnxhQtUy/mmLR2I9foDujNK030IGemrRc= +github.com/Masterminds/semver/v3 v3.1.1/go.mod h1:VPu/7SZ7ePZ3QOrcuXROw5FAcLl4a0cBrbBpGY/8hQs= github.com/alessio/shellescape v1.4.1 h1:V7yhSDDn8LP4lc4jS8pFkt0zCnzVJlG5JXy9BVKJUX0= github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPpyFjDCtHLS30= github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= diff --git a/main-mshell.go b/main-mshell.go index 014b5a76..0902de23 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -499,7 +499,7 @@ func main() { handleUsage() return } else if firstArg == "--version" { - fmt.Printf("mshell v%s\n", base.MShellVersion) + fmt.Printf("mshell %s\n", base.MShellVersion) return } else if firstArg == "--env" { rtnCode, err := handleEnv() diff --git a/pkg/base/base.go b/pkg/base/base.go index 9a08cdda..f8a38607 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -28,7 +28,7 @@ const MShellPathVarName = "MSHELL_PATH" const MShellHomeVarName = "MSHELL_HOME" const SSHCommandVarName = "SSH_COMMAND" const SessionsDirBaseName = "sessions" -const MShellVersion = "0.1.0" +const MShellVersion = "v0.1.0" const RemoteIdFile = "remoteid" var sessionDirCache = make(map[string]string) diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index c14b17e5..fadcb644 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -76,7 +76,7 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, string, e } if initPk.Version != base.MShellVersion { cproc.Close() - return nil, initPk.UName, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) + return nil, initPk.UName, fmt.Errorf("invalid remote mshell version '%s', must be %s", initPk.Version, base.MShellVersion) } cproc.InitPk = initPk } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index b8ae79e7..dc52290d 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -779,7 +779,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon return nil, fmt.Errorf("mshell command not found on remote server, can install with 'mshell --install %s %s.%s'", sshOptsStr, goos, goarch) } if initPk.Version != base.MShellVersion { - return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) + return nil, fmt.Errorf("invalid remote mshell version '%s', must be %s", initPk.Version, base.MShellVersion) } versionOk = true if debug { From b5c67b62606e371ad215fc17e7ac5676342f860b Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 26 Sep 2022 13:02:34 -0700 Subject: [PATCH 091/149] refactoring for versioned mshell binaries on remotes --- go.mod | 5 ++- go.sum | 2 ++ main-mshell.go | 4 +-- pkg/base/base.go | 16 +++++++-- pkg/shexec/client.go | 7 ++-- pkg/shexec/shexec.go | 84 ++++++++++++++++++++++++++++++-------------- scripthaus.md | 12 +++---- 7 files changed, 90 insertions(+), 40 deletions(-) diff --git a/go.mod b/go.mod index 02a9eb35..b6c3b3a7 100644 --- a/go.mod +++ b/go.mod @@ -10,4 +10,7 @@ require ( golang.org/x/sys v0.0.0-20220412211240-33da011f77ad ) -require github.com/Masterminds/semver/v3 v3.1.1 // indirect +require ( + github.com/Masterminds/semver/v3 v3.1.1 // indirect + golang.org/x/mod v0.5.1 // indirect +) diff --git a/go.sum b/go.sum index 32656449..0757a9bd 100644 --- a/go.sum +++ b/go.sum @@ -8,5 +8,7 @@ github.com/fsnotify/fsnotify v1.5.4 h1:jRbGcIw6P2Meqdwuo0H1p6JVLbL5DHKAKlYndzMwV github.com/fsnotify/fsnotify v1.5.4/go.mod h1:OVB6XrOHzAwXMpEM7uPOzcehqUV2UqJxmVXmkdnm1bU= github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +golang.org/x/mod v0.5.1 h1:OJxoQ/rynoF0dcCdI7cLPktw/hR2cueqYfjm43oqK38= +golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad h1:ntjMns5wyP/fN65tdBD4g8J5w8n015+iIIs9rtjXkY0= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= diff --git a/main-mshell.go b/main-mshell.go index 0902de23..5010b0de 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -417,14 +417,14 @@ func handleInstall() (int, error) { if !base.ValidGoArch(goos, goarch) { return 1, fmt.Errorf("invalid arch '%s' passed to mshell --install", fullArch) } - optName := base.GoArchOptFile(goos, goarch) + optName := base.GoArchOptFile(base.MShellVersion, goos, goarch) _, err = os.Stat(optName) if err != nil { return 1, fmt.Errorf("cannot install mshell to remote host, cannot read '%s': %w", optName, err) } opts.OptName = optName } - err = shexec.RunInstallSSHCommand(opts) + err = shexec.RunInstallFromOpts(opts) if err != nil { return 1, err } diff --git a/pkg/base/base.go b/pkg/base/base.go index f8a38607..4a43a4ed 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -19,6 +19,7 @@ import ( "sync" "github.com/google/uuid" + "golang.org/x/mod/semver" ) const HomeVarName = "HOME" @@ -26,10 +27,12 @@ const DefaultMShellHome = "~/.mshell" const DefaultMShellName = "mshell" const MShellPathVarName = "MSHELL_PATH" const MShellHomeVarName = "MSHELL_HOME" +const MShellInstallBinVarName = "MSHELL_INSTALLBIN_PATH" const SSHCommandVarName = "SSH_COMMAND" const SessionsDirBaseName = "sessions" const MShellVersion = "v0.1.0" const RemoteIdFile = "remoteid" +const DefaultMShellInstallBinDir = "/opt/mshell/bin" var sessionDirCache = make(map[string]string) var baseLock = &sync.Mutex{} @@ -229,8 +232,17 @@ func ValidGoArch(goos string, goarch string) bool { return (goos == "darwin" || goos == "linux") && (goarch == "amd64" || goarch == "arm64") } -func GoArchOptFile(goos string, goarch string) string { - return fmt.Sprintf("/opt/mshell/bin/mshell.%s.%s", goos, goarch) +func GoArchOptFile(version string, goos string, goarch string) string { + installBinDir := os.Getenv(MShellInstallBinVarName) + if installBinDir == "" { + installBinDir = DefaultMShellInstallBinDir + } + versionStr := semver.MajorMinor(version) + if versionStr == "" { + versionStr = "unknown" + } + binBaseName := fmt.Sprintf("mshell-%s-%s.%s", versionStr, goos, goarch) + return fmt.Sprintf(path.Join(installBinDir, binBaseName)) } func GetRemoteId() (string, error) { diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index fadcb644..30416ae3 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -9,6 +9,7 @@ import ( "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" + "golang.org/x/mod/semver" ) // TODO - track buffer sizes for sending input @@ -72,11 +73,11 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, string, e initPk := pk.(*packet.InitPacketType) if initPk.NotFound { cproc.Close() - return nil, initPk.UName, fmt.Errorf("mshell command not found on local server") + return nil, initPk.UName, fmt.Errorf("mshell-%s command not found on local server", semver.MajorMinor(base.MShellVersion)) } - if initPk.Version != base.MShellVersion { + if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { cproc.Close() - return nil, initPk.UName, fmt.Errorf("invalid remote mshell version '%s', must be %s", initPk.Version, base.MShellVersion) + return nil, initPk.UName, fmt.Errorf("invalid remote mshell version '%s', must be '=%s'", initPk.Version, semver.MajorMinor(base.MShellVersion)) } cproc.InitPk = initPk } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index dc52290d..b272f3d3 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -27,6 +27,7 @@ import ( "github.com/scripthaus-dev/mshell/pkg/cirfile" "github.com/scripthaus-dev/mshell/pkg/mpio" "github.com/scripthaus-dev/mshell/pkg/packet" + "golang.org/x/mod/semver" "golang.org/x/sys/unix" ) @@ -45,29 +46,37 @@ const MaxMaxPtySize = 100 * 1024 * 1024 const GetStateTimeout = 5 * time.Second -const ClientCommand = ` +const ClientCommandFmt = ` PATH=$PATH:~/.mshell; which mshell > /dev/null; if [[ "$?" -ne 0 ]] then printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)" else - mshell --single + mshell-[%VERSION%] --single fi ` -const InstallCommand = ` +func MakeClientCommandStr() string { + return strings.ReplaceAll(ClientCommandFmt, "[%VERSION%]", semver.MajorMinor(base.MShellVersion)) +} + +const InstallCommandFmt = ` printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)"; mkdir -p ~/.mshell/; cat > ~/.mshell/mshell.temp; if [[ -s ~/.mshell/mshell.temp ]] then - mv ~/.mshell/mshell.temp ~/.mshell/mshell; - chmod a+x ~/.mshell/mshell; - ~/.mshell/mshell --single --version + mv ~/.mshell/mshell.temp ~/.mshell/mshell-[%VERSION%]; + chmod a+x ~/.mshell/mshell-[%VERSION%]; + ~/.mshell/mshell-[%VERSION%] --single --version fi ` +func MakeInstallCommandStr() string { + return strings.ReplaceAll(InstallCommandFmt, "[%VERSION%]", semver.MajorMinor(base.MShellVersion)) +} + const RunCommandFmt = `%s` const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` @@ -389,6 +398,7 @@ type InstallOpts struct { ArchStr string OptName string Detect bool + CmdPty *os.File } type ClientOpts struct { @@ -408,7 +418,8 @@ func (opts SSHOpts) MakeSSHInstallCmd() (*exec.Cmd, error) { if opts.SSHHost == "" { return nil, fmt.Errorf("no ssh host provided, can only install to a remote host") } - return opts.MakeSSHExecCmd(InstallCommand), nil + cmdStr := MakeInstallCommandStr() + return opts.MakeSSHExecCmd(cmdStr), nil } func (opts SSHOpts) MakeMShellServerCmd() (*exec.Cmd, error) { @@ -434,7 +445,8 @@ func (opts SSHOpts) MakeMShellSingleCmd(fromServer bool) (*exec.Cmd, error) { } return ecmd, nil } - return opts.MakeSSHExecCmd(ClientCommand), nil + cmdStr := MakeClientCommandStr() + return opts.MakeSSHExecCmd(cmdStr), nil } func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { @@ -632,12 +644,7 @@ func sendOptFile(input io.WriteCloser, optName string) error { return nil } -func RunInstallSSHCommand(opts *InstallOpts) error { - tryDetect := opts.Detect - ecmd, err := opts.SSHOpts.MakeSSHInstallCmd() - if err != nil { - return err - } +func RunInstallFromCmd(ecmd *exec.Cmd, tryDetect bool, optName string, msgFn func(string)) error { inputWriter, err := ecmd.StdinPipe() if err != nil { return fmt.Errorf("creating stdin pipe: %v", err) @@ -653,8 +660,11 @@ func RunInstallSSHCommand(opts *InstallOpts) error { go func() { io.Copy(os.Stderr, stderrReader) }() - if opts.OptName != "" { - sendOptFile(inputWriter, opts.OptName) + if optName != "" { + err = sendOptFile(inputWriter, optName) + if err != nil { + return fmt.Errorf("cannot send mshell binary: %v", err) + } } packetParser := packet.MakePacketParser(stdoutReader) err = ecmd.Start() @@ -677,15 +687,19 @@ func RunInstallSSHCommand(opts *InstallOpts) error { if err != nil { return fmt.Errorf("arch cannot be detected (might be incompatible with mshell): %w", err) } - fmt.Printf("mshell detected remote architecture as '%s.%s'\n", goos, goarch) - optName := base.GoArchOptFile(goos, goarch) - sendOptFile(inputWriter, optName) + msgStr := fmt.Sprintf("mshell detected remote architecture as '%s.%s'\n", goos, goarch) + msgFn(msgStr) + optName := base.GoArchOptFile(base.MShellVersion, goos, goarch) + fmt.Printf("optname %s\n", optName) + err = sendOptFile(inputWriter, optName) + if err != nil { + return fmt.Errorf("cannot send mshell binary: %v", err) + } continue } if pk.GetType() == packet.InitPacketStr && !firstInit { initPacket := pk.(*packet.InitPacketType) if initPacket.Version == base.MShellVersion { - fmt.Printf("mshell %s, installed successfully at %s:~/.mshell/mshell\n", initPacket.Version, opts.SSHOpts.SSHHost) return nil } return fmt.Errorf("invalid version '%s' received from client, expecting '%s'", initPacket.Version, base.MShellVersion) @@ -700,6 +714,23 @@ func RunInstallSSHCommand(opts *InstallOpts) error { return fmt.Errorf("did not receive version string from client, install not successful") } +func RunInstallFromOpts(opts *InstallOpts) error { + ecmd, err := opts.SSHOpts.MakeSSHInstallCmd() + if err != nil { + return err + } + msgFn := func(str string) { + fmt.Printf("%s", str) + } + err = RunInstallFromCmd(ecmd, opts.Detect, opts.OptName, msgFn) + if err != nil { + return err + } + mmVersion := semver.MajorMinor(base.MShellVersion) + fmt.Printf("mshell installed successfully at %s:~/.mshell/mshell%s\n", opts.SSHOpts.SSHHost, mmVersion) + return nil +} + func HasDupStdin(fds []packet.RemoteFd) bool { for _, rfd := range fds { if rfd.Read && rfd.DupStdin { @@ -764,22 +795,23 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon } if pk.GetType() == packet.InitPacketStr { initPk := pk.(*packet.InitPacketType) + mmVersion := semver.MajorMinor(base.MShellVersion) if initPk.NotFound { if sshOpts.SSHHost == "" { - return nil, fmt.Errorf("mshell command not found on local server") + return nil, fmt.Errorf("mshell-%s command not found on local server", mmVersion) } if initPk.UName == "" { - return nil, fmt.Errorf("mshell command not found on remote server, no uname detected") + return nil, fmt.Errorf("mshell-%s command not found on remote server, no uname detected", mmVersion) } goos, goarch, err := DetectGoArch(initPk.UName) if err != nil { - return nil, fmt.Errorf("mshell command not found on remote server, architecture cannot be detected (might be incompatible with mshell): %w", err) + return nil, fmt.Errorf("mshell-%s command not found on remote server, architecture cannot be detected (might be incompatible with mshell): %w", mmVersion, err) } sshOptsStr := sshOpts.MakeMShellSSHOpts() - return nil, fmt.Errorf("mshell command not found on remote server, can install with 'mshell --install %s %s.%s'", sshOptsStr, goos, goarch) + return nil, fmt.Errorf("mshell-%s command not found on remote server, can install with 'mshell --install %s %s.%s'", mmVersion, sshOptsStr, goos, goarch) } - if initPk.Version != base.MShellVersion { - return nil, fmt.Errorf("invalid remote mshell version '%s', must be %s", initPk.Version, base.MShellVersion) + if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { + return nil, fmt.Errorf("invalid remote mshell version '%s', must be '=%s'", initPk.Version, semver.MajorMinor(base.MShellVersion)) } versionOk = true if debug { diff --git a/scripthaus.md b/scripthaus.md index 6752bda6..3d7df182 100644 --- a/scripthaus.md +++ b/scripthaus.md @@ -1,16 +1,16 @@ ```bash # @scripthaus command build -go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell main-mshell.go +go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.1 main-mshell.go ``` ```bash # @scripthaus command fullbuild -go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell main-mshell.go -GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.linux.amd64 main-mshell.go -GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.linux.arm64 main-mshell.go -GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.darwin.amd64 main-mshell.go -GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell.darwin.arm64 main-mshell.go +go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.1 main-mshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-linux.amd64 main-mshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-linux.arm64 main-mshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-darwin.amd64 main-mshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-darwin.arm64 main-mshell.go ``` From ea6b571184738714e34ef76d7cbfd5086c80ff06 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 26 Sep 2022 21:10:08 -0700 Subject: [PATCH 092/149] pass context to runinstall for cancelation --- pkg/shexec/shexec.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index b272f3d3..28333473 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -644,7 +644,7 @@ func sendOptFile(input io.WriteCloser, optName string) error { return nil } -func RunInstallFromCmd(ecmd *exec.Cmd, tryDetect bool, optName string, msgFn func(string)) error { +func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optName string, msgFn func(string)) error { inputWriter, err := ecmd.StdinPipe() if err != nil { return fmt.Errorf("creating stdin pipe: %v", err) @@ -722,7 +722,7 @@ func RunInstallFromOpts(opts *InstallOpts) error { msgFn := func(str string) { fmt.Printf("%s", str) } - err = RunInstallFromCmd(ecmd, opts.Detect, opts.OptName, msgFn) + err = RunInstallFromCmd(context.Background(), ecmd, opts.Detect, opts.OptName, msgFn) if err != nil { return err } From be1e1dfe909eff3773821282c0098db6bb6ac868 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 26 Sep 2022 23:23:32 -0700 Subject: [PATCH 093/149] allow context cancelation of install --- pkg/shexec/shexec.go | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 28333473..b4fb3e7a 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -672,7 +672,13 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optN return fmt.Errorf("running ssh command: %w", err) } firstInit := true - for pk := range packetParser.MainCh { + for { + var pk packet.PacketType + select { + case pk = <-packetParser.MainCh: + case <-ctx.Done(): + return ctx.Err() + } if pk.GetType() == packet.InitPacketStr && firstInit { firstInit = false initPacket := pk.(*packet.InitPacketType) @@ -690,7 +696,6 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optN msgStr := fmt.Sprintf("mshell detected remote architecture as '%s.%s'\n", goos, goarch) msgFn(msgStr) optName := base.GoArchOptFile(base.MShellVersion, goos, goarch) - fmt.Printf("optname %s\n", optName) err = sendOptFile(inputWriter, optName) if err != nil { return fmt.Errorf("cannot send mshell binary: %v", err) @@ -706,7 +711,7 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optN } if pk.GetType() == packet.RawPacketStr { rawPk := pk.(*packet.RawPacketType) - fmt.Printf("%s\n", rawPk.Data) + msgFn(fmt.Sprintf("%s\n", rawPk.Data)) continue } return fmt.Errorf("invalid response packet '%s' received from client", pk.GetType()) From d8e7c915e5b798b7d814d14200d7284b068f0292 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 30 Sep 2022 17:22:57 -0700 Subject: [PATCH 094/149] add sshport and batchmode --- main-mshell.go | 16 ++++++++++++++++ pkg/shexec/shexec.go | 24 +++++++++++++++++++----- 2 files changed, 35 insertions(+), 5 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 5010b0de..8b5d637e 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -10,6 +10,7 @@ import ( "bytes" "fmt" "os" + "strconv" "strings" "github.com/scripthaus-dev/mshell/pkg/base" @@ -283,6 +284,21 @@ func tryParseSSHOpt(iter *base.OptsIter, sshOpts *shexec.SSHOpts) (bool, error) sshOpts.SSHUser = iter.Next() return true, nil } + if argStr == "-p" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("-p [port]' missing port") + } + nextArgStr := iter.Next() + portVal, err := strconv.Atoi(nextArgStr) + if err != nil { + return false, fmt.Errorf("-p [port]' invalid port: %v", err) + } + if portVal <= 0 { + return false, fmt.Errorf("-p [port]' invalid port: %d", portVal) + } + sshOpts.SSHPort = portVal + return true, nil + } return false, nil } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index b4fb3e7a..54f2afaf 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -390,7 +390,9 @@ type SSHOpts struct { SSHOptsStr string SSHIdentity string SSHUser string + SSHPort int SSHErrorsToTty bool + BatchMode bool } type InstallOpts struct { @@ -469,12 +471,20 @@ func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) moreSSHOpts = append(moreSSHOpts, userOpt) } - // note that SSHOptsStr is *not* escaped - var errFdStr string - if opts.SSHErrorsToTty { - errFdStr = "-E /dev/tty" + if opts.SSHPort != 0 { + portOpt := fmt.Sprintf("-p %d", opts.SSHPort) + moreSSHOpts = append(moreSSHOpts, portOpt) } - sshCmd := fmt.Sprintf("ssh %s %s %s %s %s", errFdStr, strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) + if opts.SSHErrorsToTty { + errFdStr := "-E /dev/tty" + moreSSHOpts = append(moreSSHOpts, errFdStr) + } + if opts.BatchMode { + batchOpt := "-o 'BatchMode=yes'" + moreSSHOpts = append(moreSSHOpts, batchOpt) + } + // note that SSHOptsStr is *not* escaped + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) ecmd := exec.Command("bash", "-c", sshCmd) return ecmd } @@ -490,6 +500,10 @@ func (opts SSHOpts) MakeMShellSSHOpts() string { userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) moreSSHOpts = append(moreSSHOpts, userOpt) } + if opts.SSHPort != 0 { + portOpt := fmt.Sprintf("-p %d", opts.SSHPort) + moreSSHOpts = append(moreSSHOpts, portOpt) + } if opts.SSHOptsStr != "" { optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOptsStr)) moreSSHOpts = append(moreSSHOpts, optsOpt) From 1c2c8c2f4d2bc2b29e5ca10c703c9c8af8865534 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 15 Oct 2022 13:45:52 -0700 Subject: [PATCH 095/149] return defined bash aliases in initpk --- main-mshell.go | 13 +++++++++++-- pkg/packet/packet.go | 1 + pkg/shexec/shexec.go | 37 ++++++++++++++++++++++++++----------- 3 files changed, 38 insertions(+), 13 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 8b5d637e..59a826ae 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -452,10 +452,19 @@ func handleEnv() (int, error) { if err != nil { return 1, err } - fmt.Printf("%s\x00", cwd) + fmt.Printf("%s\x00\x00", cwd) fullEnv := os.Environ() + var linePrinted bool for _, envLine := range fullEnv { - fmt.Printf("%s\x00", envLine) + if envLine != "" { + fmt.Printf("%s\x00", envLine) + linePrinted = true + } + } + if linePrinted { + fmt.Printf("\x00") + } else { + fmt.Printf("\x00\x00") } return 0, nil } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 2d13ad5e..db98f8d8 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -441,6 +441,7 @@ type InitPacketType struct { HomeDir string `json:"homedir,omitempty"` Cwd string `json:"cwd,omitempty"` Env0 []byte `json:"env0,omitempty"` // "env -0" format + Aliases string `json:"aliases,omitempty"` User string `json:"user,omitempty"` HostName string `json:"hostname,omitempty"` NotFound bool `json:"notfound,omitempty"` diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 54f2afaf..3e45494a 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -81,6 +81,12 @@ const RunCommandFmt = `%s` const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` +type CurrentState struct { + Cwd string + Env0 []byte + Aliases string +} + type ShExecType struct { StartTs time.Time CK base.CommandKey @@ -1221,12 +1227,13 @@ func MakeInitPacket() *packet.InitPacketType { func MakeServerInitPacket() (*packet.InitPacketType, error) { var err error initPacket := MakeInitPacket() - cwd, env, err := GetCurrentState() + cstate, err := GetCurrentState() if err != nil { return nil, err } - initPacket.Cwd = cwd - initPacket.Env0 = env + initPacket.Cwd = cstate.Cwd + initPacket.Env0 = cstate.Env0 + initPacket.Aliases = cstate.Aliases initPacket.RemoteId, err = base.GetRemoteId() if err != nil { return nil, err @@ -1315,20 +1322,28 @@ func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { return outputBuf.Bytes(), nil } -func GetCurrentState() (string, []byte, error) { +func GetCurrentState() (*CurrentState, error) { execFile, err := os.Executable() if err != nil { - return "", nil, fmt.Errorf("cannot find local mshell executable: %w", err) + return nil, fmt.Errorf("cannot find local mshell executable: %w", err) } ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env", shellescape.Quote(execFile))) + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env; alias -p", shellescape.Quote(execFile))) outputBytes, err := runSimpleCmdInPty(ecmd) if err != nil { - return "", nil, err + return nil, err } - idx := bytes.Index(outputBytes, []byte{0}) - if idx == -1 { - return "", nil, fmt.Errorf("invalid current state output no NUL byte separator") + firstSep := bytes.Index(outputBytes, []byte{0, 0}) + if firstSep == -1 { + return nil, fmt.Errorf("invalid current state output no NUL separator") } - return string(outputBytes[0:idx]), outputBytes[idx+1:], nil + cwd := string(outputBytes[0:firstSep]) + secondSep := bytes.Index(outputBytes[firstSep+2:], []byte{0, 0}) + if secondSep == -1 { + return nil, fmt.Errorf("invalid current state output, no second NUL separator") + } + secondSep += firstSep + 2 + env0 := outputBytes[firstSep+2 : secondSep+1] // grab one of the NUL bytes (end of env0) + aliases := string(outputBytes[secondSep+2:]) + return &CurrentState{Cwd: cwd, Env0: env0, Aliases: aliases}, nil } From b9c3940b9968b70543c7557bbaa089f8983ed183 Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 16 Oct 2022 23:46:59 -0700 Subject: [PATCH 096/149] big change to execution, run command as a script and set aliases/funcs --- pkg/base/base.go | 2 +- pkg/packet/packet.go | 52 ++++++++++++----------- pkg/shexec/shexec.go | 98 ++++++++++++++++++++++++++------------------ scripthaus.md | 12 +++--- 4 files changed, 94 insertions(+), 70 deletions(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index 4a43a4ed..08a23a11 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -30,7 +30,7 @@ const MShellHomeVarName = "MSHELL_HOME" const MShellInstallBinVarName = "MSHELL_INSTALLBIN_PATH" const SSHCommandVarName = "SSH_COMMAND" const SessionsDirBaseName = "sessions" -const MShellVersion = "v0.1.0" +const MShellVersion = "v0.2.0" const RemoteIdFile = "remoteid" const DefaultMShellInstallBinDir = "/opt/mshell/bin" diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index db98f8d8..46f4c29c 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -108,6 +108,13 @@ func MakePacket(packetType string) (PacketType, error) { return rtn.Interface().(PacketType), nil } +type ShellState struct { + Cwd string `json:"cwd,omitempty"` + Env0 []byte `json:"env0,omitempty"` + Aliases string `json:"aliases,omitempty"` + Funcs string `json:"funcs,omitempty"` +} + type CmdDataPacketType struct { Type string `json:"type"` RespId string `json:"respid"` @@ -435,18 +442,16 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { } type InitPacketType struct { - Type string `json:"type"` - Version string `json:"version"` - MShellHomeDir string `json:"mshellhomedir,omitempty"` - HomeDir string `json:"homedir,omitempty"` - Cwd string `json:"cwd,omitempty"` - Env0 []byte `json:"env0,omitempty"` // "env -0" format - Aliases string `json:"aliases,omitempty"` - User string `json:"user,omitempty"` - HostName string `json:"hostname,omitempty"` - NotFound bool `json:"notfound,omitempty"` - UName string `json:"uname,omitempty"` - RemoteId string `json:"remoteid,omitempty"` + Type string `json:"type"` + Version string `json:"version"` + MShellHomeDir string `json:"mshellhomedir,omitempty"` + HomeDir string `json:"homedir,omitempty"` + State *ShellState `json:"state,omitempty"` + User string `json:"user,omitempty"` + HostName string `json:"hostname,omitempty"` + NotFound bool `json:"notfound,omitempty"` + UName string `json:"uname,omitempty"` + RemoteId string `json:"remoteid,omitempty"` } func (*InitPacketType) GetType() string { @@ -535,18 +540,17 @@ type RunDataType struct { } type RunPacketType struct { - Type string `json:"type"` - ReqId string `json:"reqid"` - CK base.CommandKey `json:"ck"` - Command string `json:"command"` - Cwd string `json:"cwd,omitempty"` - Env0 []byte `json:"env0,omitempty"` // in "env -0" format - EnvComplete bool `json:"envcomplete,omitempty"` // set to true if env0 is complete (the default env should not be set) - UsePty bool `json:"usepty,omitempty"` - TermOpts *TermOpts `json:"termopts,omitempty"` - Fds []RemoteFd `json:"fds,omitempty"` - RunData []RunDataType `json:"rundata,omitempty"` - Detached bool `json:"detached,omitempty"` + Type string `json:"type"` + ReqId string `json:"reqid"` + CK base.CommandKey `json:"ck"` + Command string `json:"command"` + State *ShellState `json:"state"` + StateComplete bool `json:"statecomplete,omitempty"` // set to true if state is complete (the default env should not be set) + UsePty bool `json:"usepty,omitempty"` + TermOpts *TermOpts `json:"termopts,omitempty"` + Fds []RemoteFd `json:"fds,omitempty"` + RunData []RunDataType `json:"rundata,omitempty"` + Detached bool `json:"detached,omitempty"` } func (*RunPacketType) GetType() string { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 3e45494a..22f6a385 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -81,12 +81,6 @@ const RunCommandFmt = `%s` const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` -type CurrentState struct { - Cwd string - Env0 []byte - Aliases string -} - type ShExecType struct { StartTs time.Time CK base.CommandKey @@ -259,14 +253,18 @@ func MakeSimpleStaticWriterPipe(data []byte) (*os.File, error) { } func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, error) { + state := pk.State + if state == nil { + state = &packet.ShellState{} + } ecmd := exec.Command("bash", "-c", pk.Command) - if !pk.EnvComplete { + if !pk.StateComplete { ecmd.Env = os.Environ() } - UpdateCmdEnv(ecmd, ParseEnv0(pk.Env0)) + UpdateCmdEnv(ecmd, ParseEnv0(state.Env0)) UpdateCmdEnv(ecmd, map[string]string{"TERM": getTermType(pk)}) - if pk.Cwd != "" { - ecmd.Dir = base.ExpandHomeDir(pk.Cwd) + if state.Cwd != "" { + ecmd.Dir = base.ExpandHomeDir(state.Cwd) } if HasDupStdin(pk.Fds) { return nil, fmt.Errorf("cannot detach command with dup stdin") @@ -360,8 +358,8 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { return fmt.Errorf("cannot detach command, constant rundata input too large len=%d, max=%d", totalRunData, mpio.MaxTotalRunDataSize) } } - if pk.Cwd != "" { - realCwd := base.ExpandHomeDir(pk.Cwd) + if pk.State != nil && pk.State.Cwd != "" { + realCwd := base.ExpandHomeDir(pk.State.Cwd) dirInfo, err := os.Stat(realCwd) if err != nil { return fmt.Errorf("invalid cwd '%s' for command: %v", realCwd, err) @@ -533,7 +531,8 @@ func GetTerminalSize() (int, int, error) { func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { runPacket := packet.MakeRunPacket() runPacket.Detached = opts.Detach - runPacket.Cwd = opts.Cwd + runPacket.State = &packet.ShellState{} + runPacket.State.Cwd = opts.Cwd runPacket.Fds = opts.Fds if opts.UsePty { runPacket.UsePty = true @@ -940,7 +939,26 @@ func getTermType(pk *packet.RunPacketType) string { return termType } +func makeEnvCommandStr(pk *packet.RunPacketType) string { + fmtStr := ` +shopt -q -s expand_aliases +set +m +%s +%s +%s +` + state := pk.State + if state == nil { + state = &packet.ShellState{} + } + return fmt.Sprintf(fmtStr, state.Aliases, state.Funcs, pk.Command) +} + func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (*ShExecType, error) { + state := pk.State + if state == nil { + state = &packet.ShellState{} + } cmd := MakeShExec(pk.CK, nil) if fromServer { msgUpr := packet.MessageUPR{CK: pk.CK, Sender: sender} @@ -948,19 +966,24 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro cmd.Multiplexer.UPR = upr cmd.MsgSender = sender } - if pk.UsePty { - cmd.Cmd = exec.Command("bash", "-i", "-c", pk.Command) - } else { - cmd.Cmd = exec.Command("bash", "-c", pk.Command) + commandStr := makeEnvCommandStr(pk) + commandFdNum, err := AddRunData(pk, commandStr, "command") + if err != nil { + return nil, err } - if !pk.EnvComplete { + if pk.UsePty { + cmd.Cmd = exec.Command("bash", "-i", fmt.Sprintf("/dev/fd/%d", commandFdNum)) + } else { + cmd.Cmd = exec.Command("bash", fmt.Sprintf("/dev/fd/%d", commandFdNum)) + } + if !pk.StateComplete { cmd.Cmd.Env = os.Environ() } - UpdateCmdEnv(cmd.Cmd, ParseEnv0(pk.Env0)) - if pk.Cwd != "" { - cmd.Cmd.Dir = base.ExpandHomeDir(pk.Cwd) + UpdateCmdEnv(cmd.Cmd, ParseEnv0(state.Env0)) + if state.Cwd != "" { + cmd.Cmd.Dir = base.ExpandHomeDir(state.Cwd) } - err := ValidateRemoteFds(pk.Fds) + err = ValidateRemoteFds(pk.Fds) if err != nil { cmd.Close() return nil, err @@ -1227,13 +1250,11 @@ func MakeInitPacket() *packet.InitPacketType { func MakeServerInitPacket() (*packet.InitPacketType, error) { var err error initPacket := MakeInitPacket() - cstate, err := GetCurrentState() + shellState, err := GetShellState() if err != nil { return nil, err } - initPacket.Cwd = cstate.Cwd - initPacket.Env0 = cstate.Env0 - initPacket.Aliases = cstate.Aliases + initPacket.State = shellState initPacket.RemoteId, err = base.GetRemoteId() if err != nil { return nil, err @@ -1322,28 +1343,27 @@ func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { return outputBuf.Bytes(), nil } -func GetCurrentState() (*CurrentState, error) { +func GetShellState() (*packet.ShellState, error) { execFile, err := os.Executable() if err != nil { return nil, fmt.Errorf("cannot find local mshell executable: %w", err) } ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env; alias -p", shellescape.Quote(execFile))) + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env; alias -p; printf \"\\x00\\x00\"; declare -f", shellescape.Quote(execFile))) outputBytes, err := runSimpleCmdInPty(ecmd) if err != nil { return nil, err } - firstSep := bytes.Index(outputBytes, []byte{0, 0}) - if firstSep == -1 { - return nil, fmt.Errorf("invalid current state output no NUL separator") + fields := bytes.Split(outputBytes, []byte{0, 0}) + if len(fields) != 4 { + return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) } - cwd := string(outputBytes[0:firstSep]) - secondSep := bytes.Index(outputBytes[firstSep+2:], []byte{0, 0}) - if secondSep == -1 { - return nil, fmt.Errorf("invalid current state output, no second NUL separator") + rtn := &packet.ShellState{} + rtn.Cwd = string(fields[0]) + if len(fields[1]) > 0 { + rtn.Env0 = append(fields[1], '\x00') } - secondSep += firstSep + 2 - env0 := outputBytes[firstSep+2 : secondSep+1] // grab one of the NUL bytes (end of env0) - aliases := string(outputBytes[secondSep+2:]) - return &CurrentState{Cwd: cwd, Env0: env0, Aliases: aliases}, nil + rtn.Aliases = strings.ReplaceAll(string(fields[2]), "\r\n", "\n") + rtn.Funcs = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") + return rtn, nil } diff --git a/scripthaus.md b/scripthaus.md index 3d7df182..bdcff57c 100644 --- a/scripthaus.md +++ b/scripthaus.md @@ -1,16 +1,16 @@ ```bash # @scripthaus command build -go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.1 main-mshell.go +go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go ``` ```bash # @scripthaus command fullbuild -go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.1 main-mshell.go -GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-linux.amd64 main-mshell.go -GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-linux.arm64 main-mshell.go -GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-darwin.amd64 main-mshell.go -GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.1-darwin.arm64 main-mshell.go +go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-linux.amd64 main-mshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-linux.arm64 main-mshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-darwin.amd64 main-mshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-darwin.arm64 main-mshell.go ``` From d8b5508b776b78ff4666f63c8e27d72ec10f29f4 Mon Sep 17 00:00:00 2001 From: sawka Date: Sat, 22 Oct 2022 14:45:31 -0700 Subject: [PATCH 097/149] returnstate option for runpk (for sourcing files) --- pkg/packet/packet.go | 2 + pkg/shexec/client.go | 26 ++++--- pkg/shexec/shexec.go | 179 +++++++++++++++++++++++++++++++++++-------- 3 files changed, 163 insertions(+), 44 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 46f4c29c..02346f64 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -480,6 +480,7 @@ type CmdDonePacketType struct { CK base.CommandKey `json:"ck"` ExitCode int `json:"exitcode"` DurationMs int64 `json:"durationms"` + FinalState *ShellState `json:"state,omitempty"` } func (*CmdDonePacketType) GetType() string { @@ -551,6 +552,7 @@ type RunPacketType struct { Fds []RemoteFd `json:"fds,omitempty"` RunData []RunDataType `json:"rundata,omitempty"` Detached bool `json:"detached,omitempty"` + ReturnState bool `json:"returnstate,omitempty"` } func (*RunPacketType) GetType() string { diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index 30416ae3..18ea2478 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -14,6 +14,8 @@ import ( // TODO - track buffer sizes for sending input +const NotFoundVersion = "v0.0" + type ClientProc struct { Cmd *exec.Cmd InitPk *packet.InitPacketType @@ -25,24 +27,24 @@ type ClientProc struct { Output *packet.PacketParser } -// returns (clientproc, uname, error) -func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, string, error) { +// returns (clientproc, initpk, error) +func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.InitPacketType, error) { inputWriter, err := ecmd.StdinPipe() if err != nil { - return nil, "", fmt.Errorf("creating stdin pipe: %v", err) + return nil, nil, fmt.Errorf("creating stdin pipe: %v", err) } stdoutReader, err := ecmd.StdoutPipe() if err != nil { - return nil, "", fmt.Errorf("creating stdout pipe: %v", err) + return nil, nil, fmt.Errorf("creating stdout pipe: %v", err) } stderrReader, err := ecmd.StderrPipe() if err != nil { - return nil, "", fmt.Errorf("creating stderr pipe: %v", err) + return nil, nil, fmt.Errorf("creating stderr pipe: %v", err) } startTs := time.Now() err = ecmd.Start() if err != nil { - return nil, "", fmt.Errorf("running local client: %w", err) + return nil, nil, fmt.Errorf("running local client: %w", err) } sender := packet.MakePacketSender(inputWriter) stdoutPacketParser := packet.MakePacketParser(stdoutReader) @@ -63,29 +65,29 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, string, e case pk = <-packetParser.MainCh: case <-ctx.Done(): cproc.Close() - return nil, "", ctx.Err() + return nil, nil, ctx.Err() } if pk != nil { if pk.GetType() != packet.InitPacketStr { cproc.Close() - return nil, "", fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) + return nil, nil, fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) } initPk := pk.(*packet.InitPacketType) if initPk.NotFound { cproc.Close() - return nil, initPk.UName, fmt.Errorf("mshell-%s command not found on local server", semver.MajorMinor(base.MShellVersion)) + return nil, initPk, fmt.Errorf("mshell-%s command not found on local server", semver.MajorMinor(base.MShellVersion)) } if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { cproc.Close() - return nil, initPk.UName, fmt.Errorf("invalid remote mshell version '%s', must be '=%s'", initPk.Version, semver.MajorMinor(base.MShellVersion)) + return nil, initPk, fmt.Errorf("invalid remote mshell version '%s', must be '=%s'", initPk.Version, semver.MajorMinor(base.MShellVersion)) } cproc.InitPk = initPk } if cproc.InitPk == nil { cproc.Close() - return nil, "", fmt.Errorf("no init packet received from mshell client") + return nil, nil, fmt.Errorf("no init packet received from mshell client") } - return cproc, cproc.InitPk.UName, nil + return cproc, cproc.InitPk, nil } func (cproc *ClientProc) Close() { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 22f6a385..9dc59a40 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -18,6 +18,7 @@ import ( "os/user" "runtime" "strings" + "sync" "syscall" "time" @@ -81,6 +82,20 @@ const RunCommandFmt = `%s` const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` +type ReturnStateBuf struct { + Lock *sync.Mutex + Buf []byte + Done bool + Err error + Reader *os.File + FdNum int + DoneCh chan bool +} + +func MakeReturnStateBuf() *ReturnStateBuf { + return &ReturnStateBuf{Lock: &sync.Mutex{}, DoneCh: make(chan bool)} +} + type ShExecType struct { StartTs time.Time CK base.CommandKey @@ -93,6 +108,7 @@ type ShExecType struct { DetachedOutput *packet.PacketSender RunnerOutFd *os.File MsgSender *packet.PacketSender // where to send out-of-band messages back to calling proceess + ReturnState *ReturnStateBuf } type StdContext struct{} @@ -192,6 +208,9 @@ func (c *ShExecType) Close() { if c.RunnerOutFd != nil { c.RunnerOutFd.Close() } + if c.ReturnState != nil { + c.ReturnState.Reader.Close() + } } func (c *ShExecType) MakeCmdStartPacket(reqId string) *packet.CmdStartPacketType { @@ -926,6 +945,9 @@ func DetectGoArch(uname string) (string, string, error) { func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sender *packet.PacketSender) { defer cmd.Close() + if cmd.ReturnState != nil { + go cmd.ReturnState.Run() + } cmd.Multiplexer.RunIOAndWait(packetParser, sender, true, false, false) donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) @@ -939,42 +961,87 @@ func getTermType(pk *packet.RunPacketType) string { return termType } -func makeEnvCommandStr(pk *packet.RunPacketType) string { - fmtStr := ` -shopt -q -s expand_aliases +func makeRcFileStr(pk *packet.RunPacketType) string { + rcFileStr := ` set +m -%s -%s -%s +set +H +shopt -s extglob ` - state := pk.State - if state == nil { - state = &packet.ShellState{} + if pk.State != nil && pk.State.Funcs != "" { + rcFileStr += pk.State.Funcs + "\n" } - return fmt.Sprintf(fmtStr, state.Aliases, state.Funcs, pk.Command) + if pk.State != nil && pk.State.Aliases != "" { + rcFileStr += pk.State.Aliases + "\n" + } + if pk.ReturnState { + rcFileStr += ` +_scripthaus_exittrap () { + %s --env; alias -p; printf \"\\x00\\x00\"; declare -f; +} +trap _scripthaus_exittrap EXIT +` + } + return rcFileStr } -func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (*ShExecType, error) { +func makeExitTrap(fdNum int) (string, error) { + stateCmd, err := GetShellStateRedirectCommandStr(fdNum) + if err != nil { + return "", err + } + fmtStr := ` +_scripthaus_exittrap () { + %s +} +trap _scripthaus_exittrap EXIT +` + return fmt.Sprintf(fmtStr, stateCmd), nil +} + +func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (rtnShExec *ShExecType, rtnErr error) { state := pk.State if state == nil { state = &packet.ShellState{} } cmd := MakeShExec(pk.CK, nil) + defer func() { + // on error, call cmd.Close() + if rtnErr != nil { + cmd.Close() + } + }() if fromServer { msgUpr := packet.MessageUPR{CK: pk.CK, Sender: sender} upr := ShExecUPR{ShExec: cmd, UPR: msgUpr} cmd.Multiplexer.UPR = upr cmd.MsgSender = sender } - commandStr := makeEnvCommandStr(pk) - commandFdNum, err := AddRunData(pk, commandStr, "command") + var rtnStateWriter *os.File + rcFileStr := makeRcFileStr(pk) + if pk.ReturnState { + pr, pw, err := os.Pipe() + if err != nil { + return nil, fmt.Errorf("cannot create returnstate pipe: %v", err) + } + cmd.ReturnState = MakeReturnStateBuf() + cmd.ReturnState.Reader = pr + cmd.ReturnState.FdNum = 20 + rtnStateWriter = pw + defer pw.Close() + trapCmdStr, err := makeExitTrap(cmd.ReturnState.FdNum) + if err != nil { + return nil, err + } + rcFileStr += trapCmdStr + } + rcFileFdNum, err := AddRunData(pk, rcFileStr, "rcfile") if err != nil { return nil, err } if pk.UsePty { - cmd.Cmd = exec.Command("bash", "-i", fmt.Sprintf("/dev/fd/%d", commandFdNum)) + cmd.Cmd = exec.Command("bash", "--rcfile", fmt.Sprintf("/dev/fd/%d", rcFileFdNum), "-i", "-c", pk.Command) } else { - cmd.Cmd = exec.Command("bash", fmt.Sprintf("/dev/fd/%d", commandFdNum)) + cmd.Cmd = exec.Command("bash", "--rcfile", fmt.Sprintf("/dev/fd/%d", rcFileFdNum), "-c", pk.Command) } if !pk.StateComplete { cmd.Cmd.Env = os.Environ() @@ -985,7 +1052,6 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro } err = ValidateRemoteFds(pk.Fds) if err != nil { - cmd.Close() return nil, err } var cmdPty *os.File @@ -1014,24 +1080,20 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false, true) nullFd, err := os.Open("/dev/null") if err != nil { - cmd.Close() return nil, fmt.Errorf("cannot open /dev/null: %w", err) } cmd.Multiplexer.MakeRawFdReader(2, nullFd, true, false) } else { cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) if err != nil { - cmd.Close() return nil, err } cmd.Cmd.Stdout, err = cmd.Multiplexer.MakeReaderPipe(1) if err != nil { - cmd.Close() return nil, err } cmd.Cmd.Stderr, err = cmd.Multiplexer.MakeReaderPipe(2) if err != nil { - cmd.Close() return nil, err } } @@ -1042,7 +1104,6 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro } extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data) if err != nil { - cmd.Close() return nil, err } } @@ -1054,7 +1115,6 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro // client file is open for reading, so we make a writer pipe extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum) if err != nil { - cmd.Close() return nil, err } } @@ -1062,23 +1122,53 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro // client file is open for writing, so we make a reader pipe extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeReaderPipe(rfd.FdNum) if err != nil { - cmd.Close() return nil, err } } } + if cmd.ReturnState != nil { + if cmd.ReturnState.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:cmd.ReturnState.FdNum+1] + } + extraFiles[cmd.ReturnState.FdNum] = rtnStateWriter + } if len(extraFiles) > FirstExtraFilesFdNum { cmd.Cmd.ExtraFiles = extraFiles[FirstExtraFilesFdNum:] } - err = cmd.Cmd.Start() if err != nil { - cmd.Close() return nil, err } return cmd, nil } +// TODO limit size of read state buffer +func (rs *ReturnStateBuf) Run() { + buf := make([]byte, 1024) + defer func() { + rs.Lock.Lock() + defer rs.Lock.Unlock() + rs.Reader.Close() + rs.Done = true + close(rs.DoneCh) + }() + for { + n, readErr := rs.Reader.Read(buf) + if readErr == io.EOF { + break + } + if readErr != nil { + rs.Lock.Lock() + rs.Err = readErr + rs.Lock.Unlock() + break + } + rs.Lock.Lock() + rs.Buf = append(rs.Buf, buf[0:n]...) + rs.Lock.Unlock() + } +} + // in detached run mode, we don't want mshell to die from signals // since we want mshell to persist even if the mshell --server is terminated func SetupSignalsForDetach() { @@ -1220,11 +1310,16 @@ func GetExitCode(err error) int { } func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { + donePacket := packet.MakeCmdDonePacket(c.CK) exitErr := c.Cmd.Wait() + if c.ReturnState != nil { + <-c.ReturnState.DoneCh + state, _ := ParseShellStateOutput(c.ReturnState.Buf) // TODO what to do with error? + donePacket.FinalState = state + } endTs := time.Now() cmdDuration := endTs.Sub(c.StartTs) exitCode := GetExitCode(exitErr) - donePacket := packet.MakeCmdDonePacket(c.CK) donePacket.Ts = endTs.UnixMilli() donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) @@ -1343,17 +1438,23 @@ func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { return outputBuf.Bytes(), nil } -func GetShellState() (*packet.ShellState, error) { +func GetShellStateCommandStr() (string, error) { execFile, err := os.Executable() if err != nil { - return nil, fmt.Errorf("cannot find local mshell executable: %w", err) + return "", fmt.Errorf("cannot find local mshell executable: %w", err) } - ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", fmt.Sprintf("%s --env; alias -p; printf \"\\x00\\x00\"; declare -f", shellescape.Quote(execFile))) - outputBytes, err := runSimpleCmdInPty(ecmd) + return fmt.Sprintf(`%s --env; alias -p; printf \"\\x00\\x00\"; declare -f`, shellescape.Quote(execFile)), nil +} + +func GetShellStateRedirectCommandStr(outputFdNum int) (string, error) { + cmdStr, err := GetShellStateCommandStr() if err != nil { - return nil, err + return "", err } + return fmt.Sprintf("cat <(%s) > /dev/fd/%d", cmdStr, outputFdNum), nil +} + +func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { fields := bytes.Split(outputBytes, []byte{0, 0}) if len(fields) != 4 { return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) @@ -1367,3 +1468,17 @@ func GetShellState() (*packet.ShellState, error) { rtn.Funcs = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") return rtn, nil } + +func GetShellState() (*packet.ShellState, error) { + ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) + cmdStr, err := GetShellStateCommandStr() + if err != nil { + return nil, err + } + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", cmdStr) + outputBytes, err := runSimpleCmdInPty(ecmd) + if err != nil { + return nil, err + } + return ParseShellStateOutput(outputBytes) +} From 5d6c77491fa5bb1aa073c4de8024abde2c08524a Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 24 Oct 2022 15:35:01 -0700 Subject: [PATCH 098/149] add shell parser --- go.mod | 1 + go.sum | 2 ++ 2 files changed, 3 insertions(+) diff --git a/go.mod b/go.mod index b6c3b3a7..4fa77051 100644 --- a/go.mod +++ b/go.mod @@ -13,4 +13,5 @@ require ( require ( github.com/Masterminds/semver/v3 v3.1.1 // indirect golang.org/x/mod v0.5.1 // indirect + mvdan.cc/sh/v3 v3.5.1 // indirect ) diff --git a/go.sum b/go.sum index 0757a9bd..c9f554e1 100644 --- a/go.sum +++ b/go.sum @@ -12,3 +12,5 @@ golang.org/x/mod v0.5.1 h1:OJxoQ/rynoF0dcCdI7cLPktw/hR2cueqYfjm43oqK38= golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad h1:ntjMns5wyP/fN65tdBD4g8J5w8n015+iIIs9rtjXkY0= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +mvdan.cc/sh/v3 v3.5.1 h1:hmP3UOw4f+EYexsJjFxvU38+kn+V/s2CclXHanIBkmQ= +mvdan.cc/sh/v3 v3.5.1/go.mod h1:1JcoyAKm1lZw/2bZje/iYKWicU/KMd0rsyJeKHnsK4E= From 674a6ef11eea7846cebc733a893cb4bc71ba3cc8 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 24 Oct 2022 21:26:39 -0700 Subject: [PATCH 099/149] grab shell vars with export vars --- main-mshell.go | 11 +-- pkg/packet/packet.go | 10 ++- pkg/shexec/parser.go | 200 +++++++++++++++++++++++++++++++++++++++++++ pkg/shexec/shexec.go | 55 +++--------- 4 files changed, 222 insertions(+), 54 deletions(-) create mode 100644 pkg/shexec/parser.go diff --git a/main-mshell.go b/main-mshell.go index 59a826ae..af680c68 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -526,13 +526,14 @@ func main() { } else if firstArg == "--version" { fmt.Printf("mshell %s\n", base.MShellVersion) return - } else if firstArg == "--env" { - rtnCode, err := handleEnv() + } else if firstArg == "--test-env" { + state, err := shexec.GetShellState() + if state != nil { + + } if err != nil { fmt.Fprintf(os.Stderr, "[error] %v\n", err) - } - if rtnCode != 0 { - os.Exit(rtnCode) + os.Exit(1) } } else if firstArg == "--single" { handleSingle(false) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 02346f64..88cd7369 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -109,10 +109,12 @@ func MakePacket(packetType string) (PacketType, error) { } type ShellState struct { - Cwd string `json:"cwd,omitempty"` - Env0 []byte `json:"env0,omitempty"` - Aliases string `json:"aliases,omitempty"` - Funcs string `json:"funcs,omitempty"` + Version string `json:"version,omitempty"` + Cwd string `json:"cwd,omitempty"` + ShellVars string `json:"shellvars,omitempty"` + Env0 []byte `json:"env0,omitempty"` + Aliases string `json:"aliases,omitempty"` + Funcs string `json:"funcs,omitempty"` } type CmdDataPacketType struct { diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go new file mode 100644 index 00000000..6f7eeef7 --- /dev/null +++ b/pkg/shexec/parser.go @@ -0,0 +1,200 @@ +package shexec + +import ( + "bytes" + "fmt" + "io" + "strings" + + "github.com/scripthaus-dev/mshell/pkg/packet" + "mvdan.cc/sh/v3/expand" + "mvdan.cc/sh/v3/syntax" +) + +type ParseEnviron struct { + Env map[string]string +} + +func (e *ParseEnviron) Get(name string) expand.Variable { + val, ok := e.Env[name] + if !ok { + return expand.Variable{} + } + return expand.Variable{ + Exported: true, + Kind: expand.String, + Str: val, + } +} + +func (e *ParseEnviron) Each(fn func(name string, vr expand.Variable) bool) { + for key, _ := range e.Env { + rtn := fn(key, e.Get(key)) + if !rtn { + break + } + } +} + +func doCmdSubst(commandStr string, w io.Writer, word *syntax.CmdSubst) error { + return nil +} + +func doProcSubst(w *syntax.ProcSubst) (string, error) { + return "", nil +} + +func GetParserConfig(envMap map[string]string) *expand.Config { + cfg := &expand.Config{ + Env: &ParseEnviron{Env: envMap}, + GlobStar: false, + NullGlob: false, + NoUnset: false, + CmdSubst: func(w io.Writer, word *syntax.CmdSubst) error { return doCmdSubst("", w, word) }, + ProcSubst: doProcSubst, + ReadDir: nil, + } + return cfg +} + +func QuotedLitToStr(word *syntax.Word) (string, error) { + cfg := GetParserConfig(nil) + return expand.Literal(cfg, word) +} + +// https://wiki.bash-hackers.org/syntax/shellvars +var NoStoreVarNames = map[string]bool{ + "BASH": true, + "BASHOPTS": true, + "BASHPID": true, + "BASH_ALIASES": true, + "BASH_ARGC": true, + "BASH_ARGV": true, + "BASH_ARGV0": true, + "BASH_CMDS": true, + "BASH_COMMAND": true, + "BASH_EXECUTION_STRING": true, + "BASH_LINENO": true, + "BASH_REMATCH": true, + "BASH_SOURCE": true, + "BASH_SUBSHELL": true, + "BASH_VERSINFO": true, + "BASH_VERSION": true, + "COPROC": true, + "DIRSTACK": true, + "EPOCHREALTIME": true, + "EPOCHSECONDS": true, + "FUNCNAME": true, + "HISTCMD": true, + "OLDPWD": true, + "PIPESTATUS": true, + "PPID": true, + "PWD": true, + "RANDOM": true, + "SECONDS": true, + "SHLVL": true, + "HISTFILE": true, + "HISTFILESIZE": true, + "HISTCONTROL": true, + "HISTIGNORE": true, + "HISTSIZE": true, + "HISTTIMEFORMAT": true, + "SRANDOM": true, +} + +func parseDeclareStmt(envBuffer *bytes.Buffer, varsBuffer *bytes.Buffer, stmt *syntax.Stmt, src []byte) error { + cmd := stmt.Cmd + decl, ok := cmd.(*syntax.DeclClause) + if !ok || decl.Variant.Value != "declare" || len(decl.Args) != 2 { + return fmt.Errorf("invalid declare variant") + } + declArgs := decl.Args[0] + if !declArgs.Naked || len(declArgs.Value.Parts) != 1 { + return fmt.Errorf("wrong number of declare args parts") + } + declArgLit, ok := declArgs.Value.Parts[0].(*syntax.Lit) + if !ok { + return fmt.Errorf("declare args is not a literal") + } + declArgStr := declArgLit.Value + if !strings.HasPrefix(declArgStr, "-") { + return fmt.Errorf("declare args not an argument (does not start with '-')") + } + declAssign := decl.Args[1] + if declAssign.Name == nil { + return fmt.Errorf("declare does not have a valid name") + } + varName := declAssign.Name.Value + if NoStoreVarNames[varName] { + return nil + } + if strings.Index(varName, "=") != -1 || strings.Index(varName, "\x00") != -1 { + return fmt.Errorf("invalid varname (cannot contain '=' or 0 byte)") + } + fullDeclBytes := src[decl.Pos().Offset():decl.End().Offset()] + if strings.Index(declArgStr, "x") == -1 { + // non-exported vars get written to vars as decl statements + varsBuffer.Write(fullDeclBytes) + varsBuffer.WriteRune('\n') + return nil + } + if declArgStr != "-x" { + return fmt.Errorf("can only export plain bash variables (no arrays)") + } + // exported vars are parsed into Env0 format + if declAssign.Naked || declAssign.Array != nil || declAssign.Index != nil || declAssign.Append || declAssign.Value == nil { + return fmt.Errorf("invalid variable to export") + } + varValue := declAssign.Value + varValueStr, err := QuotedLitToStr(varValue) + if err != nil { + return fmt.Errorf("parsing declare value: %w", err) + } + if strings.Index(varValueStr, "\x00") != -1 { + return fmt.Errorf("invalid export var value (cannot contain 0 byte)") + } + envBuffer.WriteString(fmt.Sprintf("%s=%s\x00", varName, varValueStr)) + return nil +} + +func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { + r := bytes.NewReader(declareBytes) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(r, "aliases") + if err != nil { + return err + } + var envBuffer, varsBuffer bytes.Buffer + for _, stmt := range file.Stmts { + err = parseDeclareStmt(&envBuffer, &varsBuffer, stmt, declareBytes) + if err != nil { + // TODO where to put parse errors? + continue + } + } + state.Env0 = envBuffer.Bytes() + state.ShellVars = varsBuffer.String() + return nil +} + +func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { + // 5 fields: version, cwd, env/vars, aliases, funcs + fields := bytes.Split(outputBytes, []byte{0, 0}) + if len(fields) != 5 { + return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) + } + rtn := &packet.ShellState{} + rtn.Version = string(fields[0]) + if strings.Index(rtn.Version, "bash") == -1 { + return nil, fmt.Errorf("invalid shell state output, only bash is supported") + } + cwdStr := string(fields[1]) + if strings.HasSuffix(cwdStr, "\r\n") { + cwdStr = cwdStr[0 : len(cwdStr)-2] + } + rtn.Cwd = string(cwdStr) + parseDeclareOutput(rtn, fields[2]) + rtn.Aliases = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") + rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") + return rtn, nil +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 9dc59a40..e7c8818e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -47,6 +47,8 @@ const MaxMaxPtySize = 100 * 1024 * 1024 const GetStateTimeout = 5 * time.Second +const GetShellStateCmd = `echo bash v${BASH_VERSINFO[0]}.${BASH_VERSINFO[1]}.${BASH_VERSINFO[2]}; printf "\x00\x00"; pwd; printf "\x00\x00"; declare -p $(compgen -A variable); printf "\x00\x00"; alias -p; printf "\x00\x00"; declare -f;` + const ClientCommandFmt = ` PATH=$PATH:~/.mshell; which mshell > /dev/null; @@ -976,7 +978,7 @@ shopt -s extglob if pk.ReturnState { rcFileStr += ` _scripthaus_exittrap () { - %s --env; alias -p; printf \"\\x00\\x00\"; declare -f; +` + GetShellStateCmd + ` } trap _scripthaus_exittrap EXIT ` @@ -984,18 +986,15 @@ trap _scripthaus_exittrap EXIT return rcFileStr } -func makeExitTrap(fdNum int) (string, error) { - stateCmd, err := GetShellStateRedirectCommandStr(fdNum) - if err != nil { - return "", err - } +func makeExitTrap(fdNum int) string { + stateCmd := GetShellStateRedirectCommandStr(fdNum) fmtStr := ` _scripthaus_exittrap () { %s } trap _scripthaus_exittrap EXIT ` - return fmt.Sprintf(fmtStr, stateCmd), nil + return fmt.Sprintf(fmtStr, stateCmd) } func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (rtnShExec *ShExecType, rtnErr error) { @@ -1028,10 +1027,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro cmd.ReturnState.FdNum = 20 rtnStateWriter = pw defer pw.Close() - trapCmdStr, err := makeExitTrap(cmd.ReturnState.FdNum) - if err != nil { - return nil, err - } + trapCmdStr := makeExitTrap(cmd.ReturnState.FdNum) rcFileStr += trapCmdStr } rcFileFdNum, err := AddRunData(pk, rcFileStr, "rcfile") @@ -1438,44 +1434,13 @@ func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { return outputBuf.Bytes(), nil } -func GetShellStateCommandStr() (string, error) { - execFile, err := os.Executable() - if err != nil { - return "", fmt.Errorf("cannot find local mshell executable: %w", err) - } - return fmt.Sprintf(`%s --env; alias -p; printf \"\\x00\\x00\"; declare -f`, shellescape.Quote(execFile)), nil -} - -func GetShellStateRedirectCommandStr(outputFdNum int) (string, error) { - cmdStr, err := GetShellStateCommandStr() - if err != nil { - return "", err - } - return fmt.Sprintf("cat <(%s) > /dev/fd/%d", cmdStr, outputFdNum), nil -} - -func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { - fields := bytes.Split(outputBytes, []byte{0, 0}) - if len(fields) != 4 { - return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) - } - rtn := &packet.ShellState{} - rtn.Cwd = string(fields[0]) - if len(fields[1]) > 0 { - rtn.Env0 = append(fields[1], '\x00') - } - rtn.Aliases = strings.ReplaceAll(string(fields[2]), "\r\n", "\n") - rtn.Funcs = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") - return rtn, nil +func GetShellStateRedirectCommandStr(outputFdNum int) string { + return fmt.Sprintf("cat <(%s) > /dev/fd/%d", GetShellStateCmd, outputFdNum) } func GetShellState() (*packet.ShellState, error) { ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - cmdStr, err := GetShellStateCommandStr() - if err != nil { - return nil, err - } - ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", cmdStr) + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", GetShellStateCmd) outputBytes, err := runSimpleCmdInPty(ecmd) if err != nil { return nil, err From 245a1995e20c4420cb64d501f227ec67b9a8cebc Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 25 Oct 2022 12:31:07 -0700 Subject: [PATCH 100/149] parse shellvariables with args --- pkg/packet/packet.go | 14 ++- pkg/shexec/parser.go | 241 ++++++++++++++++++++++++++++++++++--------- pkg/shexec/shexec.go | 30 ++++-- 3 files changed, 228 insertions(+), 57 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 88cd7369..4a33fee1 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -51,6 +51,7 @@ const ( RawPacketStr = "raw" SpecialInputPacketStr = "sinput" // command CompGenPacketStr = "compgen" // rpc + ReInitPacketStr = "reinit" // rpc ) const PacketSenderQueueSize = 20 @@ -111,10 +112,10 @@ func MakePacket(packetType string) (PacketType, error) { type ShellState struct { Version string `json:"version,omitempty"` Cwd string `json:"cwd,omitempty"` - ShellVars string `json:"shellvars,omitempty"` - Env0 []byte `json:"env0,omitempty"` + ShellVars []byte `json:"shellvars,omitempty"` Aliases string `json:"aliases,omitempty"` Funcs string `json:"funcs,omitempty"` + Error string `json:"error,omitempty"` } type CmdDataPacketType struct { @@ -445,6 +446,7 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { type InitPacketType struct { Type string `json:"type"` + RespId string `json:"respid,omitempty"` Version string `json:"version"` MShellHomeDir string `json:"mshellhomedir,omitempty"` HomeDir string `json:"homedir,omitempty"` @@ -460,6 +462,14 @@ func (*InitPacketType) GetType() string { return InitPacketStr } +func (pk *InitPacketType) GetResponseId() string { + return pk.RespId +} + +func (pk *InitPacketType) GetResponseDone() bool { + return true +} + func MakeInitPacket() *InitPacketType { return &InitPacketType{Type: InitPacketStr} } diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 6f7eeef7..aafb39c3 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -4,8 +4,10 @@ import ( "bytes" "fmt" "io" + "regexp" "strings" + "github.com/alessio/shellescape" "github.com/scripthaus-dev/mshell/pkg/packet" "mvdan.cc/sh/v3/expand" "mvdan.cc/sh/v3/syntax" @@ -78,8 +80,6 @@ var NoStoreVarNames = map[string]bool{ "BASH_REMATCH": true, "BASH_SOURCE": true, "BASH_SUBSHELL": true, - "BASH_VERSINFO": true, - "BASH_VERSION": true, "COPROC": true, "DIRSTACK": true, "EPOCHREALTIME": true, @@ -100,61 +100,200 @@ var NoStoreVarNames = map[string]bool{ "HISTSIZE": true, "HISTTIMEFORMAT": true, "SRANDOM": true, + + // we want these in our remote state object + // "EUID": true, + // "SHELLOPTS": true, + // "UID": true, + // "BASH_VERSINFO": true, + // "BASH_VERSION": true, } -func parseDeclareStmt(envBuffer *bytes.Buffer, varsBuffer *bytes.Buffer, stmt *syntax.Stmt, src []byte) error { +type DeclareDeclType struct { + Args string + Name string + Value string +} + +var declareDeclArgsRe = regexp.MustCompile("^[aAxrifx]*$") +var bashValidIdentifierRe = regexp.MustCompile("^[a-zA-Z_][a-zA-Z0-9_]*$") + +func (d *DeclareDeclType) Validate() error { + if len(d.Name) == 0 || !IsValidBashIdentifier(d.Name) { + return fmt.Errorf("invalid shell variable name (invalid bash identifier)") + } + if strings.Index(d.Value, "\x00") >= 0 { + return fmt.Errorf("invalid shell variable value (cannot contain 0 byte)") + } + if !declareDeclArgsRe.MatchString(d.Args) { + return fmt.Errorf("invalid shell variable type %s", shellescape.Quote(d.Args)) + } + return nil +} + +func (d *DeclareDeclType) Serialize() string { + return fmt.Sprintf("%s|%s=%s\x00", d.Args, d.Name, d.Value) +} + +func (d *DeclareDeclType) EnvString() string { + return d.Name + "=" + d.Value +} + +func (d *DeclareDeclType) DeclareStmt() string { + var argsStr string + if d.Args == "" { + argsStr = "--" + } else { + argsStr = "-" + d.Args + } + return fmt.Sprintf("declare %s %s=%s", argsStr, d.Name, shellescape.Quote(d.Value)) +} + +// envline should be valid +func ParseDeclLine(envLine string) *DeclareDeclType { + eqIdx := strings.Index(envLine, "=") + if eqIdx == -1 { + return nil + } + namePart := envLine[0:eqIdx] + valPart := envLine[eqIdx+1:] + pipeIdx := strings.Index(namePart, "|") + if pipeIdx == -1 { + return nil + } + return &DeclareDeclType{ + Args: namePart[0:pipeIdx], + Name: namePart[pipeIdx+1:], + Value: valPart, + } +} + +func DeclMapFromState(state *packet.ShellState) map[string]*DeclareDeclType { + if state == nil { + return nil + } + rtn := make(map[string]*DeclareDeclType) + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil { + rtn[decl.Name] = decl + } + } + return rtn +} + +func SerializeDeclMap(declMap map[string]*DeclareDeclType) []byte { + var rtn bytes.Buffer + for _, decl := range declMap { + rtn.WriteString(decl.Serialize()) + } + return rtn.Bytes() +} + +func EnvMapFromState(state *packet.ShellState) map[string]string { + if state == nil { + return nil + } + rtn := make(map[string]string) + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil && decl.IsExport() { + rtn[decl.Name] = decl.Value + } + } + return rtn +} + +func ShellVarMapFromState(state *packet.ShellState) map[string]string { + if state == nil { + return nil + } + rtn := make(map[string]string) + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil { + rtn[decl.Name] = decl.Value + } + } + return rtn +} + +func VarDeclsFromState(state *packet.ShellState) []*DeclareDeclType { + if state == nil { + return nil + } + var rtn []*DeclareDeclType + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil { + rtn = append(rtn, decl) + } + } + return rtn +} + +func IsValidBashIdentifier(s string) bool { + return bashValidIdentifierRe.MatchString(s) +} + +func (d *DeclareDeclType) IsExport() bool { + return strings.Index(d.Args, "x") >= 0 +} + +func (d *DeclareDeclType) IsReadOnly() bool { + return strings.Index(d.Args, "r") >= 0 +} + +func parseDeclareStmt(stmt *syntax.Stmt, src []byte) (*DeclareDeclType, error) { cmd := stmt.Cmd decl, ok := cmd.(*syntax.DeclClause) if !ok || decl.Variant.Value != "declare" || len(decl.Args) != 2 { - return fmt.Errorf("invalid declare variant") + return nil, fmt.Errorf("invalid declare variant") } + rtn := &DeclareDeclType{} declArgs := decl.Args[0] if !declArgs.Naked || len(declArgs.Value.Parts) != 1 { - return fmt.Errorf("wrong number of declare args parts") + return nil, fmt.Errorf("wrong number of declare args parts") } - declArgLit, ok := declArgs.Value.Parts[0].(*syntax.Lit) + declArgsLit, ok := declArgs.Value.Parts[0].(*syntax.Lit) if !ok { - return fmt.Errorf("declare args is not a literal") + return nil, fmt.Errorf("declare args is not a literal") } - declArgStr := declArgLit.Value - if !strings.HasPrefix(declArgStr, "-") { - return fmt.Errorf("declare args not an argument (does not start with '-')") + if !strings.HasPrefix(declArgsLit.Value, "-") { + return nil, fmt.Errorf("declare args not an argument (does not start with '-')") + } + if declArgsLit.Value == "--" { + rtn.Args = "" + } else { + rtn.Args = declArgsLit.Value[1:] } declAssign := decl.Args[1] if declAssign.Name == nil { - return fmt.Errorf("declare does not have a valid name") + return nil, fmt.Errorf("declare does not have a valid name") } - varName := declAssign.Name.Value - if NoStoreVarNames[varName] { - return nil + rtn.Name = declAssign.Name.Value + if declAssign.Naked || declAssign.Index != nil || declAssign.Append { + return nil, fmt.Errorf("invalid decl format") } - if strings.Index(varName, "=") != -1 || strings.Index(varName, "\x00") != -1 { - return fmt.Errorf("invalid varname (cannot contain '=' or 0 byte)") + if declAssign.Value != nil { + varValueStr, err := QuotedLitToStr(declAssign.Value) + if err != nil { + return nil, fmt.Errorf("parsing declare value: %w", err) + } + rtn.Value = varValueStr + } else if declAssign.Array != nil { + rtn.Value = string(src[declAssign.Array.Pos().Offset():declAssign.Array.End().Offset()]) + } else { + return nil, fmt.Errorf("invalid decl, not plain value or array") } - fullDeclBytes := src[decl.Pos().Offset():decl.End().Offset()] - if strings.Index(declArgStr, "x") == -1 { - // non-exported vars get written to vars as decl statements - varsBuffer.Write(fullDeclBytes) - varsBuffer.WriteRune('\n') - return nil + if err := rtn.Validate(); err != nil { + return nil, err } - if declArgStr != "-x" { - return fmt.Errorf("can only export plain bash variables (no arrays)") - } - // exported vars are parsed into Env0 format - if declAssign.Naked || declAssign.Array != nil || declAssign.Index != nil || declAssign.Append || declAssign.Value == nil { - return fmt.Errorf("invalid variable to export") - } - varValue := declAssign.Value - varValueStr, err := QuotedLitToStr(varValue) - if err != nil { - return fmt.Errorf("parsing declare value: %w", err) - } - if strings.Index(varValueStr, "\x00") != -1 { - return fmt.Errorf("invalid export var value (cannot contain 0 byte)") - } - envBuffer.WriteString(fmt.Sprintf("%s=%s\x00", varName, varValueStr)) - return nil + return rtn, nil } func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { @@ -164,16 +303,23 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { if err != nil { return err } - var envBuffer, varsBuffer bytes.Buffer + var varsBuffer bytes.Buffer + var firstParseErr error for _, stmt := range file.Stmts { - err = parseDeclareStmt(&envBuffer, &varsBuffer, stmt, declareBytes) + decl, err := parseDeclareStmt(stmt, declareBytes) if err != nil { - // TODO where to put parse errors? - continue + if firstParseErr == nil { + firstParseErr = err + } + } + if decl != nil && !NoStoreVarNames[decl.Name] { + varsBuffer.WriteString(decl.Serialize()) } } - state.Env0 = envBuffer.Bytes() - state.ShellVars = varsBuffer.String() + state.ShellVars = varsBuffer.Bytes() + if firstParseErr != nil { + state.Error = firstParseErr.Error() + } return nil } @@ -193,7 +339,10 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { cwdStr = cwdStr[0 : len(cwdStr)-2] } rtn.Cwd = string(cwdStr) - parseDeclareOutput(rtn, fields[2]) + err := parseDeclareOutput(rtn, fields[2]) + if err != nil { + return nil, err + } rtn.Aliases = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") return rtn, nil diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index e7c8818e..7d29ccdc 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -282,7 +282,7 @@ func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, if !pk.StateComplete { ecmd.Env = os.Environ() } - UpdateCmdEnv(ecmd, ParseEnv0(state.Env0)) + UpdateCmdEnv(ecmd, EnvMapFromState(state)) UpdateCmdEnv(ecmd, map[string]string{"TERM": getTermType(pk)}) if state.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(state.Cwd) @@ -964,26 +964,38 @@ func getTermType(pk *packet.RunPacketType) string { } func makeRcFileStr(pk *packet.RunPacketType) string { - rcFileStr := ` + var rcBuf bytes.Buffer + rcBuf.WriteString(` set +m set +H shopt -s extglob -` +`) + + varDecls := VarDeclsFromState(pk.State) + for _, varDecl := range varDecls { + if varDecl.IsExport() || varDecl.IsReadOnly() { + continue + } + rcBuf.WriteString(varDecl.DeclareStmt()) + rcBuf.WriteString("\n") + } if pk.State != nil && pk.State.Funcs != "" { - rcFileStr += pk.State.Funcs + "\n" + rcBuf.WriteString(pk.State.Funcs) + rcBuf.WriteString("\n") } if pk.State != nil && pk.State.Aliases != "" { - rcFileStr += pk.State.Aliases + "\n" + rcBuf.WriteString(pk.State.Aliases) + rcBuf.WriteString("\n") } if pk.ReturnState { - rcFileStr += ` + rcBuf.WriteString(` _scripthaus_exittrap () { ` + GetShellStateCmd + ` } trap _scripthaus_exittrap EXIT -` +`) } - return rcFileStr + return rcBuf.String() } func makeExitTrap(fdNum int) string { @@ -1042,7 +1054,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro if !pk.StateComplete { cmd.Cmd.Env = os.Environ() } - UpdateCmdEnv(cmd.Cmd, ParseEnv0(state.Env0)) + UpdateCmdEnv(cmd.Cmd, EnvMapFromState(state)) if state.Cwd != "" { cmd.Cmd.Dir = base.ExpandHomeDir(state.Cwd) } From e5d2267f278f95857e1241140dab216b581ab7fe Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 27 Oct 2022 00:34:16 -0700 Subject: [PATCH 101/149] minor updates to get state to be consistent --- pkg/shexec/parser.go | 5 ++++- pkg/shexec/shexec.go | 11 ++++------- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index aafb39c3..ecee19d7 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -76,6 +76,7 @@ var NoStoreVarNames = map[string]bool{ "BASH_CMDS": true, "BASH_COMMAND": true, "BASH_EXECUTION_STRING": true, + "LINENO": true, "BASH_LINENO": true, "BASH_REMATCH": true, "BASH_SOURCE": true, @@ -330,13 +331,15 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) } rtn := &packet.ShellState{} - rtn.Version = string(fields[0]) + rtn.Version = strings.TrimSpace(string(fields[0])) if strings.Index(rtn.Version, "bash") == -1 { return nil, fmt.Errorf("invalid shell state output, only bash is supported") } cwdStr := string(fields[1]) if strings.HasSuffix(cwdStr, "\r\n") { cwdStr = cwdStr[0 : len(cwdStr)-2] + } else if strings.HasSuffix(cwdStr, "\n") { + cwdStr = cwdStr[0 : len(cwdStr)-1] } rtn.Cwd = string(cwdStr) err := parseDeclareOutput(rtn, fields[2]) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 7d29ccdc..30b0f537 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -47,6 +47,7 @@ const MaxMaxPtySize = 100 * 1024 * 1024 const GetStateTimeout = 5 * time.Second +const BaseBashOpts = `set +m; set +H; shopt -s extglob` const GetShellStateCmd = `echo bash v${BASH_VERSINFO[0]}.${BASH_VERSINFO[1]}.${BASH_VERSINFO[2]}; printf "\x00\x00"; pwd; printf "\x00\x00"; declare -p $(compgen -A variable); printf "\x00\x00"; alias -p; printf "\x00\x00"; declare -f;` const ClientCommandFmt = ` @@ -965,12 +966,7 @@ func getTermType(pk *packet.RunPacketType) string { func makeRcFileStr(pk *packet.RunPacketType) string { var rcBuf bytes.Buffer - rcBuf.WriteString(` -set +m -set +H -shopt -s extglob -`) - + rcBuf.WriteString(BaseBashOpts + "\n") varDecls := VarDeclsFromState(pk.State) for _, varDecl := range varDecls { if varDecl.IsExport() || varDecl.IsReadOnly() { @@ -1452,7 +1448,8 @@ func GetShellStateRedirectCommandStr(outputFdNum int) string { func GetShellState() (*packet.ShellState, error) { ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", GetShellStateCmd) + cmdStr := BaseBashOpts + "; " + GetShellStateCmd + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", cmdStr) outputBytes, err := runSimpleCmdInPty(ecmd) if err != nil { return nil, err From 1da450e61c9511a10fc8a22892f614339d1cb412 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 27 Oct 2022 21:59:17 -0700 Subject: [PATCH 102/149] implement reinit, also do not store 'columns' var --- pkg/packet/packet.go | 19 +++++++++++++++++++ pkg/packet/parser.go | 1 + pkg/server/server.go | 14 ++++++++++++++ pkg/shexec/parser.go | 1 + 4 files changed, 35 insertions(+) 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, From 8939c57dd81e7ad636fdc09957916ef6d85d657e Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 30 Oct 2022 12:52:58 -0700 Subject: [PATCH 103/149] minor, add isempty func for shellstate --- pkg/packet/packet.go | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index f409ec9d..553c5b2e 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -120,6 +120,10 @@ type ShellState struct { Error string `json:"error,omitempty"` } +func (state ShellState) IsEmpty() bool { + return state.Version == "" && state.Cwd == "" && len(state.ShellVars) == 0 && state.Aliases == "" && state.Funcs == "" && state.Error == "" +} + type CmdDataPacketType struct { Type string `json:"type"` RespId string `json:"respid"` From ee36078082993fd7ee1ea3339974188b8e99cac9 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 1 Nov 2022 21:19:42 -0700 Subject: [PATCH 104/149] parameterize the mshell bin directory (for packaging). inject MSHELL variables for execution --- pkg/base/base.go | 16 +++++++++++++ pkg/shexec/shexec.go | 53 +++++++++++++++++++++++++++----------------- 2 files changed, 49 insertions(+), 20 deletions(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index 08a23a11..3c9840db 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -245,6 +245,22 @@ func GoArchOptFile(version string, goos string, goarch string) string { return fmt.Sprintf(path.Join(installBinDir, binBaseName)) } +func MShellBinaryFromOptDir(version string, goos string, goarch string) (io.ReadCloser, error) { + if !ValidGoArch(goos, goarch) { + return nil, fmt.Errorf("invalid goos/goarch combination: %s/%s", goos, goarch) + } + versionStr := semver.MajorMinor(version) + if versionStr == "" { + return nil, fmt.Errorf("invalid mshell version: %q", version) + } + fileName := GoArchOptFile(version, goos, goarch) + fd, err := os.Open(fileName) + if err != nil { + return nil, fmt.Errorf("cannot open mshell binary %q: %v", fileName, err) + } + return fd, nil +} + func GetRemoteId() (string, error) { mhome := GetMShellHomeDir() homeInfo, err := os.Stat(mhome) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 30b0f537..0b5ecf04 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -85,6 +85,8 @@ const RunCommandFmt = `%s` const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` +type MShellBinaryReaderFn func(version string, goos string, goarch string) (io.ReadCloser, error) + type ReturnStateBuf struct { Lock *sync.Mutex Buf []byte @@ -284,7 +286,7 @@ func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, ecmd.Env = os.Environ() } UpdateCmdEnv(ecmd, EnvMapFromState(state)) - UpdateCmdEnv(ecmd, map[string]string{"TERM": getTermType(pk)}) + UpdateCmdEnv(ecmd, MShellEnvVars(getTermType(pk))) if state.Cwd != "" { ecmd.Dir = base.ExpandHomeDir(state.Cwd) } @@ -673,19 +675,14 @@ func ValidateRemoteFds(rfds []packet.RemoteFd) error { return nil } -func sendOptFile(input io.WriteCloser, optName string) error { - fd, err := os.Open(optName) - if err != nil { - return fmt.Errorf("cannot open '%s': %w", optName, err) - } +func sendMShellBinary(input io.WriteCloser, mshellStream io.Reader) { go func() { defer input.Close() - io.Copy(input, fd) + io.Copy(input, mshellStream) }() - return nil } -func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optName string, msgFn func(string)) error { +func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, mshellStream io.Reader, mshellReaderFn MShellBinaryReaderFn, msgFn func(string)) error { inputWriter, err := ecmd.StdinPipe() if err != nil { return fmt.Errorf("creating stdin pipe: %v", err) @@ -701,11 +698,8 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optN go func() { io.Copy(os.Stderr, stderrReader) }() - if optName != "" { - err = sendOptFile(inputWriter, optName) - if err != nil { - return fmt.Errorf("cannot send mshell binary: %v", err) - } + if mshellStream != nil { + sendMShellBinary(inputWriter, mshellStream) } packetParser := packet.MakePacketParser(stdoutReader) err = ecmd.Start() @@ -736,11 +730,12 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, optN } msgStr := fmt.Sprintf("mshell detected remote architecture as '%s.%s'\n", goos, goarch) msgFn(msgStr) - optName := base.GoArchOptFile(base.MShellVersion, goos, goarch) - err = sendOptFile(inputWriter, optName) + detectedMSS, err := mshellReaderFn(base.MShellVersion, goos, goarch) if err != nil { - return fmt.Errorf("cannot send mshell binary: %v", err) + return err } + defer detectedMSS.Close() + sendMShellBinary(inputWriter, detectedMSS) continue } if pk.GetType() == packet.InitPacketStr && !firstInit { @@ -768,7 +763,15 @@ func RunInstallFromOpts(opts *InstallOpts) error { msgFn := func(str string) { fmt.Printf("%s", str) } - err = RunInstallFromCmd(context.Background(), ecmd, opts.Detect, opts.OptName, msgFn) + var mshellStream *os.File + if opts.OptName != "" { + mshellStream, err = os.Open(opts.OptName) + if err != nil { + return fmt.Errorf("cannot open mshell binary %q: %v", opts.OptName, err) + } + defer mshellStream.Close() + } + err = RunInstallFromCmd(context.Background(), ecmd, opts.Detect, mshellStream, base.MShellBinaryFromOptDir, msgFn) if err != nil { return err } @@ -1070,7 +1073,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro cmdTty.Close() }() cmd.CmdPty = cmdPty - UpdateCmdEnv(cmd.Cmd, map[string]string{"TERM": getTermType(pk)}) + UpdateCmdEnv(cmd.Cmd, MShellEnvVars(getTermType(pk))) } if cmdTty != nil { cmd.Cmd.Stdin = cmdTty @@ -1407,7 +1410,7 @@ func getStderr(err error) string { func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { ecmd.Env = os.Environ() - UpdateCmdEnv(ecmd, map[string]string{"TERM": DefaultTermType}) + UpdateCmdEnv(ecmd, MShellEnvVars(DefaultTermType)) cmdPty, cmdTty, err := pty.Open() if err != nil { return nil, fmt.Errorf("opening new pty: %w", err) @@ -1456,3 +1459,13 @@ func GetShellState() (*packet.ShellState, error) { } return ParseShellStateOutput(outputBytes) } + +func MShellEnvVars(termType string) map[string]string { + rtn := make(map[string]string) + if termType != "" { + rtn["TERM"] = termType + } + rtn["MSHELL"], _ = os.Executable() + rtn["MSHELL_VERSION"] = base.MShellVersion + return rtn +} From 76d3c10748665a82433eeb2c3d2adf10ef72ea76 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 2 Nov 2022 18:41:53 -0700 Subject: [PATCH 105/149] decls equal fn -- must parse associative arrays (order is not consistent in bash output) --- pkg/shexec/parser.go | 91 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 91 insertions(+) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 25a5369e..78269e16 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -13,6 +13,13 @@ import ( "mvdan.cc/sh/v3/syntax" ) +const ( + DeclTypeArray = "array" + DeclTypeAssocArray = "assoc" + DeclTypeInt = "int" + DeclTypeNormal = "normal" +) + type ParseEnviron struct { Env map[string]string } @@ -250,6 +257,19 @@ func (d *DeclareDeclType) IsReadOnly() bool { return strings.Index(d.Args, "r") >= 0 } +func (d *DeclareDeclType) DataType() string { + if strings.Index(d.Args, "a") >= 0 { + return DeclTypeArray + } + if strings.Index(d.Args, "A") >= 0 { + return DeclTypeAssocArray + } + if strings.Index(d.Args, "i") >= 0 { + return DeclTypeInt + } + return DeclTypeNormal +} + func parseDeclareStmt(stmt *syntax.Stmt, src []byte) (*DeclareDeclType, error) { cmd := stmt.Cmd decl, ok := cmd.(*syntax.DeclClause) @@ -351,3 +371,74 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") return rtn, nil } + +func assocArrayVarToMap(d *DeclareDeclType) (map[string]string, error) { + if d.DataType() != DeclTypeAssocArray { + return nil, fmt.Errorf("decl is not an assoc-array") + } + refStr := "X=" + d.Value + r := strings.NewReader(refStr) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(r, "assocdecl") + if err != nil { + return nil, err + } + if len(file.Stmts) != 1 { + return nil, fmt.Errorf("invalid assoc-array parse (multiple stmts)") + } + stmt := file.Stmts[0] + callExpr, ok := stmt.Cmd.(*syntax.CallExpr) + if !ok || len(callExpr.Args) != 0 || len(callExpr.Assigns) != 1 { + return nil, fmt.Errorf("invalid assoc-array parse (bad expr)") + } + assign := callExpr.Assigns[0] + arrayExpr := assign.Array + if arrayExpr == nil { + return nil, fmt.Errorf("invalid assoc-array parse (no array expr)") + } + rtn := make(map[string]string) + for _, elem := range arrayExpr.Elems { + indexStr := refStr[elem.Index.Pos().Offset():elem.Index.End().Offset()] + valStr := refStr[elem.Value.Pos().Offset():elem.Value.End().Offset()] + rtn[indexStr] = valStr + } + return rtn, nil +} + +func strMapsEqual(m1 map[string]string, m2 map[string]string) bool { + if len(m1) != len(m2) { + return false + } + for key, val1 := range m1 { + val2, found := m2[key] + if !found || val1 != val2 { + return false + } + } + for key, _ := range m2 { + _, found := m1[key] + if !found { + return false + } + } + return true +} + +func DeclsEqual(d1 *DeclareDeclType, d2 *DeclareDeclType) bool { + if d1.IsExport() != d2.IsExport() { + return false + } + if d1.DataType() != d2.DataType() { + return false + } + // comparing value will work for all data types *except* for associative arrays (bash does not output them in a consistent order) + if d1.DataType() == DeclTypeAssocArray { + m1, err1 := assocArrayVarToMap(d1) + m2, err2 := assocArrayVarToMap(d2) + if err1 != nil || err2 != nil { + return d1.Value == d2.Value + } + return strMapsEqual(m1, m2) + } + return d1.Value == d2.Value +} From 4392956f998cd46d1aebc2ca9436c6ca0cf9e723 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 3 Nov 2022 23:53:25 -0700 Subject: [PATCH 106/149] replace expand.Literal with safer hand-written version that only expands quoted strings and literals --- pkg/shexec/parser.go | 135 +++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 129 insertions(+), 6 deletions(-) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 78269e16..32f64e4e 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -66,9 +66,131 @@ func GetParserConfig(envMap map[string]string) *expand.Config { return cfg } -func QuotedLitToStr(word *syntax.Word) (string, error) { - cfg := GetParserConfig(nil) - return expand.Literal(cfg, word) +func writeIndent(buf *bytes.Buffer, num int) { + for i := 0; i < num; i++ { + buf.WriteByte(' ') + } +} + +func makeSpaceStr(num int) string { + barr := make([]byte, num) + for i := 0; i < num; i++ { + barr[i] = ' ' + } + return string(barr) +} + +type SimpleExpandContext struct { + HomeDir string +} + +func expandLiteral(buf *bytes.Buffer, litVal string) { + var lastBackSlash bool + for _, ch := range litVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '\n' { + // special case, backslash *and* newline are ignored + continue + } + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + continue + } + buf.WriteRune(ch) + } + if lastBackSlash { + buf.WriteByte('\\') + } +} + +func expandDQLiteral(buf *bytes.Buffer, litVal string) { + var lastBackSlash bool + for _, ch := range litVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { + buf.WriteRune(ch) + continue + } + buf.WriteRune('\\') + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + continue + } + buf.WriteRune(ch) + } + // in a valid parsed DQ string, you cannot have a trailing backslash (because \" would not end the string) + // still putting the case here though in case we ever deal with incomplete strings (e.g. completion) + if lastBackSlash { + buf.WriteByte('\\') + } +} + +func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts []syntax.WordPart, sourceStr string, inDoubleQuote bool, level int) error { + for partIdx, untypedPart := range parts { + switch part := untypedPart.(type) { + case *syntax.Lit: + if !inDoubleQuote && part.Value == "~" && partIdx == 0 && len(parts) == 1 && level == 1 && ectx.HomeDir != "" { + buf.WriteString(ectx.HomeDir) + continue + } + if !inDoubleQuote && strings.HasPrefix(part.Value, "~/") && partIdx == 0 && level == 1 && ectx.HomeDir != "" { + buf.WriteString(ectx.HomeDir) + buf.WriteString(part.Value[1:]) + continue + } + if inDoubleQuote { + expandDQLiteral(buf, part.Value) + } else { + expandLiteral(buf, part.Value) + } + + case *syntax.SglQuoted: + if part.Dollar { + str, _, _ := expand.Format(nil, part.Value, nil) + buf.WriteString(str) + } else { + buf.WriteString(part.Value) + } + + case *syntax.DblQuoted: + simpleExpandWordInternal(buf, ectx, part.Parts, sourceStr, true, level+1) + + default: + rawStr := sourceStr[part.Pos().Offset():part.End().Offset()] + buf.WriteString(rawStr) + } + } + return nil +} + +// simple word expansion +// expands: literals, single-quoted strings, double-quoted strings (recursively) +// does *not* expand: params (variables), command substitution, arithmetic expressions, process substituions, globs +// for the not expands, they will show up as the literal string +// this is different than expand.Literal which will replace variables as empty string if they aren't defined. +// so "a"'foo'${bar}$x => "afoo${bar}$x", but expand.Literal would produce => "afoo" +// note will do ~ expansion (will not do ~user expansion) +func SimpleExpandWord(ectx SimpleExpandContext, word *syntax.Word, sourceStr string) (string, error) { + var buf bytes.Buffer + err := simpleExpandWordInternal(&buf, ectx, word.Parts, sourceStr, false, 1) + if err != nil { + return "", err + } + return buf.String(), nil } // https://wiki.bash-hackers.org/syntax/shellvars @@ -270,7 +392,7 @@ func (d *DeclareDeclType) DataType() string { return DeclTypeNormal } -func parseDeclareStmt(stmt *syntax.Stmt, src []byte) (*DeclareDeclType, error) { +func parseDeclareStmt(stmt *syntax.Stmt, src string) (*DeclareDeclType, error) { cmd := stmt.Cmd decl, ok := cmd.(*syntax.DeclClause) if !ok || decl.Variant.Value != "declare" || len(decl.Args) != 2 { @@ -302,7 +424,7 @@ func parseDeclareStmt(stmt *syntax.Stmt, src []byte) (*DeclareDeclType, error) { return nil, fmt.Errorf("invalid decl format") } if declAssign.Value != nil { - varValueStr, err := QuotedLitToStr(declAssign.Value) + varValueStr, err := SimpleExpandWord(SimpleExpandContext{}, declAssign.Value, src) if err != nil { return nil, fmt.Errorf("parsing declare value: %w", err) } @@ -319,6 +441,7 @@ func parseDeclareStmt(stmt *syntax.Stmt, src []byte) (*DeclareDeclType, error) { } func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { + declareStr := string(declareBytes) r := bytes.NewReader(declareBytes) parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) file, err := parser.Parse(r, "aliases") @@ -328,7 +451,7 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { var varsBuffer bytes.Buffer var firstParseErr error for _, stmt := range file.Stmts { - decl, err := parseDeclareStmt(stmt, declareBytes) + decl, err := parseDeclareStmt(stmt, declareStr) if err != nil { if firstParseErr == nil { firstParseErr = err From d86bee87d8e3931464bf342b75c2b081fb33d7ef Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 4 Nov 2022 12:28:08 -0700 Subject: [PATCH 107/149] partial word expansion --- pkg/shexec/parser.go | 82 ++++++++++++++++++++++++++++++-------------- 1 file changed, 57 insertions(+), 25 deletions(-) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 32f64e4e..906d25a1 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -84,6 +84,19 @@ type SimpleExpandContext struct { HomeDir string } +func expandHomeDir(litVal string, multiPart bool, homeDir string) string { + if homeDir == "" { + return litVal + } + if litVal == "~" && !multiPart { + return homeDir + } + if strings.HasPrefix(litVal, "~/") { + return homeDir + litVal[1:] + } + return litVal +} + func expandLiteral(buf *bytes.Buffer, litVal string) { var lastBackSlash bool for _, ch := range litVal { @@ -110,6 +123,22 @@ func expandLiteral(buf *bytes.Buffer, litVal string) { } } +// also expands ~ +func expandLiteralPlus(buf *bytes.Buffer, litVal string, multiPart bool, ectx SimpleExpandContext) { + litVal = expandHomeDir(litVal, multiPart, ectx.HomeDir) + expandLiteral(buf, litVal) +} + +func expandSQANSILiteral(buf *bytes.Buffer, litVal string) { + str, _, _ := expand.Format(nil, litVal, nil) + buf.WriteString(str) +} + +func expandSQLiteral(buf *bytes.Buffer, litVal string) { + buf.WriteString(litVal) +} + +// will also work for partial double quoted strings func expandDQLiteral(buf *bytes.Buffer, litVal string) { var lastBackSlash bool for _, ch := range litVal { @@ -139,20 +168,13 @@ func expandDQLiteral(buf *bytes.Buffer, litVal string) { } } -func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts []syntax.WordPart, sourceStr string, inDoubleQuote bool, level int) error { +func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts []syntax.WordPart, sourceStr string, inDoubleQuote bool, level int) { for partIdx, untypedPart := range parts { switch part := untypedPart.(type) { case *syntax.Lit: - if !inDoubleQuote && part.Value == "~" && partIdx == 0 && len(parts) == 1 && level == 1 && ectx.HomeDir != "" { - buf.WriteString(ectx.HomeDir) - continue - } - if !inDoubleQuote && strings.HasPrefix(part.Value, "~/") && partIdx == 0 && level == 1 && ectx.HomeDir != "" { - buf.WriteString(ectx.HomeDir) - buf.WriteString(part.Value[1:]) - continue - } - if inDoubleQuote { + if !inDoubleQuote && partIdx == 0 && level == 1 && ectx.HomeDir != "" { + expandLiteralPlus(buf, part.Value, len(parts) > 1, ectx) + } else if inDoubleQuote { expandDQLiteral(buf, part.Value) } else { expandLiteral(buf, part.Value) @@ -160,10 +182,9 @@ func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts case *syntax.SglQuoted: if part.Dollar { - str, _, _ := expand.Format(nil, part.Value, nil) - buf.WriteString(str) + expandSQANSILiteral(buf, part.Value) } else { - buf.WriteString(part.Value) + expandSQLiteral(buf, part.Value) } case *syntax.DblQuoted: @@ -174,7 +195,6 @@ func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts buf.WriteString(rawStr) } } - return nil } // simple word expansion @@ -184,13 +204,29 @@ func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts // this is different than expand.Literal which will replace variables as empty string if they aren't defined. // so "a"'foo'${bar}$x => "afoo${bar}$x", but expand.Literal would produce => "afoo" // note will do ~ expansion (will not do ~user expansion) -func SimpleExpandWord(ectx SimpleExpandContext, word *syntax.Word, sourceStr string) (string, error) { +func SimpleExpandWord(ectx SimpleExpandContext, word *syntax.Word, sourceStr string) string { var buf bytes.Buffer - err := simpleExpandWordInternal(&buf, ectx, word.Parts, sourceStr, false, 1) - if err != nil { - return "", err + simpleExpandWordInternal(&buf, ectx, word.Parts, sourceStr, false, 1) + return buf.String() +} + +func SimpleExpandPartialWord(ectx SimpleExpandContext, partialWord string, multiPart bool) string { + var buf bytes.Buffer + if partialWord == "" { + return "" } - return buf.String(), nil + if strings.HasPrefix(partialWord, "\"") { + expandDQLiteral(&buf, partialWord[1:]) + } else if strings.HasPrefix(partialWord, "$\"") { + expandDQLiteral(&buf, partialWord[2:]) + } else if strings.HasPrefix(partialWord, "'") { + expandSQLiteral(&buf, partialWord[1:]) + } else if strings.HasPrefix(partialWord, "$'") { + expandSQANSILiteral(&buf, partialWord[2:]) + } else { + expandLiteralPlus(&buf, partialWord, multiPart, ectx) + } + return buf.String() } // https://wiki.bash-hackers.org/syntax/shellvars @@ -424,11 +460,7 @@ func parseDeclareStmt(stmt *syntax.Stmt, src string) (*DeclareDeclType, error) { return nil, fmt.Errorf("invalid decl format") } if declAssign.Value != nil { - varValueStr, err := SimpleExpandWord(SimpleExpandContext{}, declAssign.Value, src) - if err != nil { - return nil, fmt.Errorf("parsing declare value: %w", err) - } - rtn.Value = varValueStr + rtn.Value = SimpleExpandWord(SimpleExpandContext{}, declAssign.Value, src) } else if declAssign.Array != nil { rtn.Value = string(src[declAssign.Array.Pos().Offset():declAssign.Array.End().Offset()]) } else { From 01821ca09496ffc43514520826b818d410ad861f Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 9 Nov 2022 20:38:47 -0800 Subject: [PATCH 108/149] allow variable comptype --- pkg/packet/packet.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 553c5b2e..f4a6cac6 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -364,7 +364,7 @@ type CompGenPacketType struct { } func IsValidCompGenType(t string) bool { - return (t == "file" || t == "command" || t == "directory") + return (t == "file" || t == "command" || t == "directory" || t == "variable") } func (*CompGenPacketType) GetType() string { From da2fe25fb8dfea30b27b317c433b03f4df957db3 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 22 Nov 2022 22:59:29 -0800 Subject: [PATCH 109/149] move to simpleexpand --- pkg/packet/packet.go | 2 +- pkg/shexec/parser.go | 152 +--------------------- pkg/simpleexpand/simpleexpand.go | 213 +++++++++++++++++++++++++++++++ 3 files changed, 216 insertions(+), 151 deletions(-) create mode 100644 pkg/simpleexpand/simpleexpand.go diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index f4a6cac6..403c1cbe 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -683,7 +683,7 @@ func ParseJsonPacket(jsonBuf []byte) (PacketType, error) { } err = json.Unmarshal(jsonBuf, pk) if err != nil { - return nil, err + return nil, fmt.Errorf("unmarshaling %q packet: %v", bareCmd.Type, err) } return pk, nil } diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 906d25a1..3980e721 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -9,6 +9,7 @@ import ( "github.com/alessio/shellescape" "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/simpleexpand" "mvdan.cc/sh/v3/expand" "mvdan.cc/sh/v3/syntax" ) @@ -80,155 +81,6 @@ func makeSpaceStr(num int) string { return string(barr) } -type SimpleExpandContext struct { - HomeDir string -} - -func expandHomeDir(litVal string, multiPart bool, homeDir string) string { - if homeDir == "" { - return litVal - } - if litVal == "~" && !multiPart { - return homeDir - } - if strings.HasPrefix(litVal, "~/") { - return homeDir + litVal[1:] - } - return litVal -} - -func expandLiteral(buf *bytes.Buffer, litVal string) { - var lastBackSlash bool - for _, ch := range litVal { - if ch == 0 { - break - } - if lastBackSlash { - lastBackSlash = false - if ch == '\n' { - // special case, backslash *and* newline are ignored - continue - } - buf.WriteRune(ch) - continue - } - if ch == '\\' { - lastBackSlash = true - continue - } - buf.WriteRune(ch) - } - if lastBackSlash { - buf.WriteByte('\\') - } -} - -// also expands ~ -func expandLiteralPlus(buf *bytes.Buffer, litVal string, multiPart bool, ectx SimpleExpandContext) { - litVal = expandHomeDir(litVal, multiPart, ectx.HomeDir) - expandLiteral(buf, litVal) -} - -func expandSQANSILiteral(buf *bytes.Buffer, litVal string) { - str, _, _ := expand.Format(nil, litVal, nil) - buf.WriteString(str) -} - -func expandSQLiteral(buf *bytes.Buffer, litVal string) { - buf.WriteString(litVal) -} - -// will also work for partial double quoted strings -func expandDQLiteral(buf *bytes.Buffer, litVal string) { - var lastBackSlash bool - for _, ch := range litVal { - if ch == 0 { - break - } - if lastBackSlash { - lastBackSlash = false - if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { - buf.WriteRune(ch) - continue - } - buf.WriteRune('\\') - buf.WriteRune(ch) - continue - } - if ch == '\\' { - lastBackSlash = true - continue - } - buf.WriteRune(ch) - } - // in a valid parsed DQ string, you cannot have a trailing backslash (because \" would not end the string) - // still putting the case here though in case we ever deal with incomplete strings (e.g. completion) - if lastBackSlash { - buf.WriteByte('\\') - } -} - -func simpleExpandWordInternal(buf *bytes.Buffer, ectx SimpleExpandContext, parts []syntax.WordPart, sourceStr string, inDoubleQuote bool, level int) { - for partIdx, untypedPart := range parts { - switch part := untypedPart.(type) { - case *syntax.Lit: - if !inDoubleQuote && partIdx == 0 && level == 1 && ectx.HomeDir != "" { - expandLiteralPlus(buf, part.Value, len(parts) > 1, ectx) - } else if inDoubleQuote { - expandDQLiteral(buf, part.Value) - } else { - expandLiteral(buf, part.Value) - } - - case *syntax.SglQuoted: - if part.Dollar { - expandSQANSILiteral(buf, part.Value) - } else { - expandSQLiteral(buf, part.Value) - } - - case *syntax.DblQuoted: - simpleExpandWordInternal(buf, ectx, part.Parts, sourceStr, true, level+1) - - default: - rawStr := sourceStr[part.Pos().Offset():part.End().Offset()] - buf.WriteString(rawStr) - } - } -} - -// simple word expansion -// expands: literals, single-quoted strings, double-quoted strings (recursively) -// does *not* expand: params (variables), command substitution, arithmetic expressions, process substituions, globs -// for the not expands, they will show up as the literal string -// this is different than expand.Literal which will replace variables as empty string if they aren't defined. -// so "a"'foo'${bar}$x => "afoo${bar}$x", but expand.Literal would produce => "afoo" -// note will do ~ expansion (will not do ~user expansion) -func SimpleExpandWord(ectx SimpleExpandContext, word *syntax.Word, sourceStr string) string { - var buf bytes.Buffer - simpleExpandWordInternal(&buf, ectx, word.Parts, sourceStr, false, 1) - return buf.String() -} - -func SimpleExpandPartialWord(ectx SimpleExpandContext, partialWord string, multiPart bool) string { - var buf bytes.Buffer - if partialWord == "" { - return "" - } - if strings.HasPrefix(partialWord, "\"") { - expandDQLiteral(&buf, partialWord[1:]) - } else if strings.HasPrefix(partialWord, "$\"") { - expandDQLiteral(&buf, partialWord[2:]) - } else if strings.HasPrefix(partialWord, "'") { - expandSQLiteral(&buf, partialWord[1:]) - } else if strings.HasPrefix(partialWord, "$'") { - expandSQANSILiteral(&buf, partialWord[2:]) - } else { - expandLiteralPlus(&buf, partialWord, multiPart, ectx) - } - return buf.String() -} - // https://wiki.bash-hackers.org/syntax/shellvars var NoStoreVarNames = map[string]bool{ "BASH": true, @@ -460,7 +312,7 @@ func parseDeclareStmt(stmt *syntax.Stmt, src string) (*DeclareDeclType, error) { return nil, fmt.Errorf("invalid decl format") } if declAssign.Value != nil { - rtn.Value = SimpleExpandWord(SimpleExpandContext{}, declAssign.Value, src) + rtn.Value, _ = simpleexpand.SimpleExpandWord(simpleexpand.SimpleExpandContext{}, declAssign.Value, src) } else if declAssign.Array != nil { rtn.Value = string(src[declAssign.Array.Pos().Offset():declAssign.Array.End().Offset()]) } else { diff --git a/pkg/simpleexpand/simpleexpand.go b/pkg/simpleexpand/simpleexpand.go new file mode 100644 index 00000000..40691319 --- /dev/null +++ b/pkg/simpleexpand/simpleexpand.go @@ -0,0 +1,213 @@ +package simpleexpand + +import ( + "bytes" + "strings" + + "mvdan.cc/sh/v3/expand" + "mvdan.cc/sh/v3/syntax" +) + +type SimpleExpandContext struct { + HomeDir string +} + +type SimpleExpandInfo struct { + HasTilde bool // only ~ as the first character when SimpleExpandContext.HomeDir is set + HasVar bool // $x, $$, ${...} + HasGlob bool // *, ?, [, { + HasExtGlob bool // ?(...) ... ?*+@! + HasHistory bool // ! (anywhere) + HasSpecial bool // subshell, arith +} + +func expandHomeDir(info *SimpleExpandInfo, litVal string, multiPart bool, homeDir string) string { + if homeDir == "" { + return litVal + } + if litVal == "~" && !multiPart { + return homeDir + } + if strings.HasPrefix(litVal, "~/") { + info.HasTilde = true + return homeDir + litVal[1:] + } + return litVal +} + +func expandLiteral(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string) { + var lastBackSlash bool + var lastExtGlob bool + var lastDollar bool + for _, ch := range litVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '\n' { + // special case, backslash *and* newline are ignored + continue + } + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + lastExtGlob = false + lastDollar = false + continue + } + if ch == '*' || ch == '?' || ch == '[' || ch == '{' { + info.HasGlob = true + } + if ch == '`' { + info.HasSpecial = true + } + if ch == '!' { + info.HasHistory = true + } + if lastExtGlob && ch == '(' { + info.HasExtGlob = true + } + if lastDollar && (ch != ' ' && ch != '"' && ch != '\'' && ch != '(' || ch != '[') { + info.HasVar = true + } + if lastDollar && (ch == '(' || ch == '[') { + info.HasSpecial = true + } + lastExtGlob = (ch == '?' || ch == '*' || ch == '+' || ch == '@' || ch == '!') + lastDollar = (ch == '$') + buf.WriteRune(ch) + } + if lastBackSlash { + buf.WriteByte('\\') + } +} + +// also expands ~ +func expandLiteralPlus(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string, multiPart bool, ectx SimpleExpandContext) { + litVal = expandHomeDir(info, litVal, multiPart, ectx.HomeDir) + expandLiteral(buf, info, litVal) +} + +func expandSQANSILiteral(buf *bytes.Buffer, litVal string) { + // no info specials + str, _, _ := expand.Format(nil, litVal, nil) + buf.WriteString(str) +} + +func expandSQLiteral(buf *bytes.Buffer, litVal string) { + // no info specials + buf.WriteString(litVal) +} + +// will also work for partial double quoted strings +func expandDQLiteral(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string) { + var lastBackSlash bool + var lastDollar bool + for _, ch := range litVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { + buf.WriteRune(ch) + continue + } + buf.WriteRune('\\') + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + lastDollar = false + continue + } + + // similar to expandLiteral, but no globbing + if ch == '`' { + info.HasSpecial = true + } + if ch == '!' { + info.HasHistory = true + } + if lastDollar && (ch != ' ' && ch != '"' && ch != '\'' && ch != '(' || ch != '[') { + info.HasVar = true + } + if lastDollar && (ch == '(' || ch == '[') { + info.HasSpecial = true + } + lastDollar = (ch == '$') + buf.WriteRune(ch) + } + // in a valid parsed DQ string, you cannot have a trailing backslash (because \" would not end the string) + // still putting the case here though in case we ever deal with incomplete strings (e.g. completion) + if lastBackSlash { + buf.WriteByte('\\') + } +} + +func simpleExpandWordInternal(buf *bytes.Buffer, info *SimpleExpandInfo, ectx SimpleExpandContext, parts []syntax.WordPart, sourceStr string, inDoubleQuote bool, level int) { + for partIdx, untypedPart := range parts { + switch part := untypedPart.(type) { + case *syntax.Lit: + if !inDoubleQuote && partIdx == 0 && level == 1 && ectx.HomeDir != "" { + expandLiteralPlus(buf, info, part.Value, len(parts) > 1, ectx) + } else if inDoubleQuote { + expandDQLiteral(buf, info, part.Value) + } else { + expandLiteral(buf, info, part.Value) + } + + case *syntax.SglQuoted: + if part.Dollar { + expandSQANSILiteral(buf, part.Value) + } else { + expandSQLiteral(buf, part.Value) + } + + case *syntax.DblQuoted: + simpleExpandWordInternal(buf, info, ectx, part.Parts, sourceStr, true, level+1) + + default: + rawStr := sourceStr[part.Pos().Offset():part.End().Offset()] + buf.WriteString(rawStr) + } + } +} + +// simple word expansion +// expands: literals, single-quoted strings, double-quoted strings (recursively) +// does *not* expand: params (variables), command substitution, arithmetic expressions, process substituions, globs +// for the not expands, they will show up as the literal string +// this is different than expand.Literal which will replace variables as empty string if they aren't defined. +// so "a"'foo'${bar}$x => "afoo${bar}$x", but expand.Literal would produce => "afoo" +// note will do ~ expansion (will not do ~user expansion) +func SimpleExpandWord(ectx SimpleExpandContext, word *syntax.Word, sourceStr string) (string, SimpleExpandInfo) { + var buf bytes.Buffer + var info SimpleExpandInfo + simpleExpandWordInternal(&buf, &info, ectx, word.Parts, sourceStr, false, 1) + return buf.String(), info +} + +func SimpleExpandPartialWord(ectx SimpleExpandContext, partialWord string, multiPart bool) (string, SimpleExpandInfo) { + var buf bytes.Buffer + var info SimpleExpandInfo + if partialWord == "" { + return "", info + } + if strings.HasPrefix(partialWord, "\"") { + expandDQLiteral(&buf, &info, partialWord[1:]) + } else if strings.HasPrefix(partialWord, "$\"") { + expandDQLiteral(&buf, &info, partialWord[2:]) + } else if strings.HasPrefix(partialWord, "'") { + expandSQLiteral(&buf, partialWord[1:]) + } else if strings.HasPrefix(partialWord, "$'") { + expandSQANSILiteral(&buf, partialWord[2:]) + } else { + expandLiteralPlus(&buf, &info, partialWord, multiPart, ectx) + } + return buf.String(), info +} From 717461908a9dfd5fbbd0b172f2cd36fe2cdcf3cd Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 23 Nov 2022 10:51:51 -0800 Subject: [PATCH 110/149] command compgen also needs to detect directories --- pkg/server/server.go | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 61e649e2..6ad8169d 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -89,7 +89,7 @@ func strArrToMap(strs []string) map[string]bool { return rtn } -func (m *MServer) runFileCompGen(compPk *packet.CompGenPacketType) { +func (m *MServer) runMixedCompGen(compPk *packet.CompGenPacketType) { // get directories and files, unique them and put slashes on directories for completion reqId := compPk.GetReqId() compDirs, hasMoreDirs, err := runSingleCompGen(compPk.Cwd, "directory", compPk.Prefix) @@ -97,7 +97,7 @@ func (m *MServer) runFileCompGen(compPk *packet.CompGenPacketType) { m.Sender.SendErrorResponse(reqId, err) return } - compFiles, hasMoreFiles, err := runSingleCompGen(compPk.Cwd, "file", compPk.Prefix) + compFiles, hasMoreFiles, err := runSingleCompGen(compPk.Cwd, compPk.CompType, compPk.Prefix) if err != nil { m.Sender.SendErrorResponse(reqId, err) return @@ -121,8 +121,8 @@ func (m *MServer) runFileCompGen(compPk *packet.CompGenPacketType) { func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { reqId := compPk.GetReqId() - if compPk.CompType == "file" { - m.runFileCompGen(compPk) + if compPk.CompType == "file" || compPk.CompType == "command" { + m.runMixedCompGen(compPk) return } comps, hasMore, err := runSingleCompGen(compPk.Cwd, compPk.CompType, compPk.Prefix) From c94c0b7c366b0b1cc179bda1f3b6ce4a29464b9d Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 25 Nov 2022 14:21:51 -0800 Subject: [PATCH 111/149] normalize assoc arrays on parse so then values can be compared --- pkg/binpack/binpack.go | 53 ++++++++++++++++++++++++++++++++ pkg/shexec/parser.go | 68 ++++++++++++++++++++++++++++++++---------- 2 files changed, 105 insertions(+), 16 deletions(-) create mode 100644 pkg/binpack/binpack.go diff --git a/pkg/binpack/binpack.go b/pkg/binpack/binpack.go new file mode 100644 index 00000000..ba0e7ce3 --- /dev/null +++ b/pkg/binpack/binpack.go @@ -0,0 +1,53 @@ +package binpack + +import ( + "encoding/binary" + "io" +) + +type FullByteReader interface { + io.ByteReader + io.Reader +} + +func PackValue(w io.Writer, barr []byte) error { + viBuf := make([]byte, binary.MaxVarintLen64) + viLen := binary.PutUvarint(viBuf, uint64(len(barr))) + _, err := w.Write(viBuf[0:viLen]) + if err != nil { + return err + } + _, err = w.Write(barr) + if err != nil { + return err + } + return nil +} + +func PackInt(w io.Writer, ival int) error { + viBuf := make([]byte, binary.MaxVarintLen64) + l := binary.PutUvarint(viBuf, uint64(ival)) + _, err := w.Write(viBuf[0:l]) + return err +} + +func UnpackValue(r FullByteReader) ([]byte, error) { + lenVal, err := binary.ReadUvarint(r) + if err != nil { + return nil, err + } + rtnBuf := make([]byte, int(lenVal)) + _, err = io.ReadFull(r, rtnBuf) + if err != nil { + return nil, err + } + return rtnBuf, nil +} + +func UnpackInt(r io.ByteReader) (int, error) { + ival64, err := binary.ReadVarint(r) + if err != nil { + return 0, err + } + return int(ival64), nil +} diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 3980e721..7b96c1a4 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -5,11 +5,11 @@ import ( "fmt" "io" "regexp" + "sort" "strings" "github.com/alessio/shellescape" "github.com/scripthaus-dev/mshell/pkg/packet" - "github.com/scripthaus-dev/mshell/pkg/simpleexpand" "mvdan.cc/sh/v3/expand" "mvdan.cc/sh/v3/syntax" ) @@ -129,8 +129,10 @@ var NoStoreVarNames = map[string]bool{ } type DeclareDeclType struct { - Args string - Name string + Args string + Name string + + // this holds the raw quoted value suitable for bash. this is *not* the real expanded variable value Value string } @@ -165,7 +167,7 @@ func (d *DeclareDeclType) DeclareStmt() string { } else { argsStr = "-" + d.Args } - return fmt.Sprintf("declare %s %s=%s", argsStr, d.Name, shellescape.Quote(d.Value)) + return fmt.Sprintf("declare %s %s=%s", argsStr, d.Name, d.Value) } // envline should be valid @@ -312,13 +314,17 @@ func parseDeclareStmt(stmt *syntax.Stmt, src string) (*DeclareDeclType, error) { return nil, fmt.Errorf("invalid decl format") } if declAssign.Value != nil { - rtn.Value, _ = simpleexpand.SimpleExpandWord(simpleexpand.SimpleExpandContext{}, declAssign.Value, src) + rtn.Value = string(src[declAssign.Value.Pos().Offset():declAssign.Value.End().Offset()]) } else if declAssign.Array != nil { rtn.Value = string(src[declAssign.Array.Pos().Offset():declAssign.Array.End().Offset()]) } else { return nil, fmt.Errorf("invalid decl, not plain value or array") } - if err := rtn.Validate(); err != nil { + err := rtn.normalize() + if err != nil { + return nil, err + } + if err = rtn.Validate(); err != nil { return nil, err } return rtn, nil @@ -379,6 +385,42 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { return rtn, nil } +func (d *DeclareDeclType) normalize() error { + if d.DataType() == DeclTypeAssocArray { + return d.normalizeAssocArrayDecl() + } + return nil +} + +// normalizes order of assoc array keys so value is stable +func (d *DeclareDeclType) normalizeAssocArrayDecl() error { + if d.DataType() != DeclTypeAssocArray { + return fmt.Errorf("invalid decltype passed to assocArrayDeclToStr: %s", d.DataType()) + } + varMap, err := assocArrayVarToMap(d) + if err != nil { + return err + } + keys := make([]string, 0, len(varMap)) + for key, _ := range varMap { + keys = append(keys, key) + } + sort.Strings(keys) + var buf bytes.Buffer + buf.WriteByte('(') + for _, key := range keys { + buf.WriteByte('[') + buf.WriteString(key) + buf.WriteByte(']') + buf.WriteByte('=') + buf.WriteString(varMap[key]) + buf.WriteByte(' ') + } + buf.WriteByte(')') + d.Value = buf.String() + return nil +} + func assocArrayVarToMap(d *DeclareDeclType) (map[string]string, error) { if d.DataType() != DeclTypeAssocArray { return nil, fmt.Errorf("decl is not an assoc-array") @@ -431,21 +473,15 @@ func strMapsEqual(m1 map[string]string, m2 map[string]string) bool { return true } -func DeclsEqual(d1 *DeclareDeclType, d2 *DeclareDeclType) bool { +func DeclsEqual(compareName bool, d1 *DeclareDeclType, d2 *DeclareDeclType) bool { if d1.IsExport() != d2.IsExport() { return false } if d1.DataType() != d2.DataType() { return false } - // comparing value will work for all data types *except* for associative arrays (bash does not output them in a consistent order) - if d1.DataType() == DeclTypeAssocArray { - m1, err1 := assocArrayVarToMap(d1) - m2, err2 := assocArrayVarToMap(d2) - if err1 != nil || err2 != nil { - return d1.Value == d2.Value - } - return strMapsEqual(m1, m2) + if compareName && d1.Name != d2.Name { + return false } - return d1.Value == d2.Value + return d1.Value == d2.Value // this works even for assoc arrays because we normalize them when parsing } From 5a151369cb6b87f5427bdbe091b41647279fc477 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 25 Nov 2022 15:52:29 -0800 Subject: [PATCH 112/149] line/mapdiff code --- pkg/statediff/linediff.go | 182 ++++++++++++++++++++++++++++++++ pkg/statediff/mapdiff.go | 124 ++++++++++++++++++++++ pkg/statediff/statediff_test.go | 99 +++++++++++++++++ 3 files changed, 405 insertions(+) create mode 100644 pkg/statediff/linediff.go create mode 100644 pkg/statediff/mapdiff.go create mode 100644 pkg/statediff/statediff_test.go diff --git a/pkg/statediff/linediff.go b/pkg/statediff/linediff.go new file mode 100644 index 00000000..a2e36491 --- /dev/null +++ b/pkg/statediff/linediff.go @@ -0,0 +1,182 @@ +package statediff + +import ( + "bytes" + "encoding/binary" + "fmt" + "strings" +) + +const LineDiffVersion = 0 + +type SingleLineEntry struct { + LineVal int + Run int +} + +type LineDiffType struct { + Lines []SingleLineEntry + NewData []string +} + +func (diff LineDiffType) dump() { + fmt.Printf("DIFF:\n") + pos := 1 + for _, entry := range diff.Lines { + fmt.Printf(" %d-%d: %d\n", pos, pos+entry.Run, entry.LineVal) + pos += entry.Run + } + for idx, str := range diff.NewData { + fmt.Printf(" n%d: %s\n", idx+1, str) + } +} + +// simple encoding +// a 0 means read a line from NewData +// a non-zero number means read the 1-indexed line from OldData +func (diff LineDiffType) applyDiff(oldData []string) ([]string, error) { + rtn := make([]string, 0, len(diff.Lines)) + newDataPos := 0 + for _, entry := range diff.Lines { + if entry.LineVal == 0 { + for i := 0; i < entry.Run; i++ { + if newDataPos >= len(diff.NewData) { + return nil, fmt.Errorf("not enough newdata for diff") + } + rtn = append(rtn, diff.NewData[newDataPos]) + newDataPos++ + } + } else { + oldDataPos := entry.LineVal - 1 // 1-indexed + for i := 0; i < entry.Run; i++ { + realPos := oldDataPos + i + if realPos < 0 || realPos >= len(oldData) { + return nil, fmt.Errorf("diff index out of bounds %d old-data-len:%d", realPos, len(oldData)) + } + rtn = append(rtn, oldData[realPos]) + } + } + } + return rtn, nil +} + +func putUVarint(buf *bytes.Buffer, viBuf []byte, ival int) { + l := binary.PutUvarint(viBuf, uint64(ival)) + buf.Write(viBuf[0:l]) +} + +// simple encoding +// write varints. first version, then len, then len-number-of-varints, then fill the rest with newdata +// [version] [len-varint] [varint]xlen... newdata (bytes) +func (diff LineDiffType) encode() []byte { + var buf bytes.Buffer + viBuf := make([]byte, binary.MaxVarintLen64) + putUVarint(&buf, viBuf, LineDiffVersion) + putUVarint(&buf, viBuf, len(diff.Lines)) + for _, entry := range diff.Lines { + putUVarint(&buf, viBuf, entry.LineVal) + putUVarint(&buf, viBuf, entry.Run) + } + for idx, str := range diff.NewData { + buf.WriteString(str) + if idx != len(diff.NewData)-1 { + buf.WriteByte('\n') + } + } + return buf.Bytes() +} + +func (rtn *LineDiffType) decode(diffBytes []byte) error { + r := bytes.NewBuffer(diffBytes) + version, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read version: %v", err) + } + if version != LineDiffVersion { + return fmt.Errorf("invalid diff, bad version: %d", version) + } + linesLen64, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read lines length: %v", err) + } + linesLen := int(linesLen64) + rtn.Lines = make([]SingleLineEntry, linesLen) + for idx := 0; idx < linesLen; idx++ { + lineVal, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read line %d: %v", idx, err) + } + lineRun, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read line-run %d: %v", idx, err) + } + rtn.Lines[idx] = SingleLineEntry{LineVal: int(lineVal), Run: int(lineRun)} + } + restOfInput := string(r.Bytes()) + if len(restOfInput) > 0 { + rtn.NewData = strings.Split(restOfInput, "\n") + } + return nil +} + +func makeLineDiff(oldData []string, newData []string) LineDiffType { + var rtn LineDiffType + oldDataMap := make(map[string]int) // 1-indexed + for idx, str := range oldData { + if _, found := oldDataMap[str]; found { + continue + } + oldDataMap[str] = idx + 1 + } + var cur *SingleLineEntry + rtn.Lines = make([]SingleLineEntry, 0) + for _, str := range newData { + oldIdx, found := oldDataMap[str] + if cur != nil && cur.LineVal != 0 { + checkLine := cur.LineVal + cur.Run - 1 + if checkLine < len(oldData) && oldData[checkLine] == str { + cur.Run++ + continue + } + } else if cur != nil && cur.LineVal == 0 && !found { + cur.Run++ + rtn.NewData = append(rtn.NewData, str) + continue + } + if cur != nil { + rtn.Lines = append(rtn.Lines, *cur) + } + cur = &SingleLineEntry{Run: 1} + if found { + cur.LineVal = oldIdx + } else { + cur.LineVal = 0 + rtn.NewData = append(rtn.NewData, str) + } + } + if cur != nil { + rtn.Lines = append(rtn.Lines, *cur) + } + return rtn +} + +func MakeLineDiff(str1 string, str2 string) []byte { + str1Arr := strings.Split(str1, "\n") + str2Arr := strings.Split(str2, "\n") + diff := makeLineDiff(str1Arr, str2Arr) + return diff.encode() +} + +func ApplyLineDiff(str1 string, diffBytes []byte) (string, error) { + var diff LineDiffType + err := diff.decode(diffBytes) + if err != nil { + return "", err + } + str1Arr := strings.Split(str1, "\n") + str2Arr, err := diff.applyDiff(str1Arr) + if err != nil { + return "", err + } + return strings.Join(str2Arr, "\n"), nil +} diff --git a/pkg/statediff/mapdiff.go b/pkg/statediff/mapdiff.go new file mode 100644 index 00000000..a06298b7 --- /dev/null +++ b/pkg/statediff/mapdiff.go @@ -0,0 +1,124 @@ +package statediff + +import ( + "bytes" + "encoding/binary" + "fmt" +) + +const MapDiffVersion = 0 + +// 0-bytes are not allowed in entries or keys (same as bash) + +type MapDiffType struct { + ToAdd map[string]string + ToRemove []string +} + +func (diff MapDiffType) dump() { + fmt.Printf("VAR-DIFF\n") + for name, val := range diff.ToAdd { + fmt.Printf(" add: %s=%s\n", name, val) + } + for _, name := range diff.ToRemove { + fmt.Printf(" rem: %s\n", name) + } +} + +func makeMapDiff(oldMap map[string]string, newMap map[string]string) MapDiffType { + var rtn MapDiffType + rtn.ToAdd = make(map[string]string) + for name, newVal := range newMap { + oldVal, found := oldMap[name] + if !found || oldVal != newVal { + rtn.ToAdd[name] = newVal + continue + } + } + for name, _ := range oldMap { + _, found := newMap[name] + if !found { + rtn.ToRemove = append(rtn.ToRemove, name) + } + } + return rtn +} + +func (diff MapDiffType) apply(oldMap map[string]string) map[string]string { + rtn := make(map[string]string) + for name, val := range oldMap { + rtn[name] = val + } + for name, val := range diff.ToAdd { + rtn[name] = val + } + for _, name := range diff.ToRemove { + delete(rtn, name) + } + return rtn +} + +func (diff MapDiffType) encode() []byte { + var buf bytes.Buffer + viBuf := make([]byte, binary.MaxVarintLen64) + putUVarint(&buf, viBuf, MapDiffVersion) + putUVarint(&buf, viBuf, len(diff.ToAdd)) + for key, val := range diff.ToAdd { + buf.WriteString(key) + buf.WriteByte(0) + buf.WriteString(val) + buf.WriteByte(0) + } + for _, val := range diff.ToRemove { + buf.WriteString(val) + buf.WriteByte(0) + } + return buf.Bytes() +} + +func (diff *MapDiffType) decode(diffBytes []byte) error { + r := bytes.NewBuffer(diffBytes) + version, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read version: %v", err) + } + if version != MapDiffVersion { + return fmt.Errorf("invalid diff, bad version: %d", version) + } + mapLen64, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot map length: %v", err) + } + mapLen := int(mapLen64) + fields := bytes.Split(r.Bytes(), []byte{0}) + if len(fields) < 2*mapLen { + return fmt.Errorf("invalid diff, not enough fields, maplen:%d fields:%d", mapLen, len(fields)) + } + mapFields := fields[0 : 2*mapLen] + removeFields := fields[2*mapLen:] + diff.ToAdd = make(map[string]string) + for i := 0; i < len(mapFields); i += 2 { + diff.ToAdd[string(mapFields[i])] = string(mapFields[i+1]) + } + for _, removeVal := range removeFields { + if len(removeVal) == 0 { + continue + } + diff.ToRemove = append(diff.ToRemove, string(removeVal)) + } + return nil +} + +func MakeMapDiff(m1 map[string]string, m2 map[string]string) []byte { + diff := makeMapDiff(m1, m2) + return diff.encode() +} + +func ApplyMapDiff(oldMap map[string]string, diffBytes []byte) (map[string]string, error) { + var diff MapDiffType + err := diff.decode(diffBytes) + if err != nil { + return nil, err + } + return diff.apply(oldMap), nil +} diff --git a/pkg/statediff/statediff_test.go b/pkg/statediff/statediff_test.go new file mode 100644 index 00000000..3b231c34 --- /dev/null +++ b/pkg/statediff/statediff_test.go @@ -0,0 +1,99 @@ +package statediff + +import ( + "fmt" + "testing" +) + +const Str1 = ` +hello +line #2 +apple +grapes +banana +apple +` + +const Str2 = ` +line #2 +apple +grapes +banana +` + +const Str3 = ` +more +stuff +banana +coconut +` + +const Str4 = ` +more +stuff +banana2 +coconut +` + +func testLineDiff(t *testing.T, str1 string, str2 string) { + diffBytes := MakeLineDiff(str1, str2) + fmt.Printf("diff-len: %d\n", len(diffBytes)) + out, err := ApplyLineDiff(str1, diffBytes) + if err != nil { + t.Errorf("error in diff: %v", err) + return + } + if out != str2 { + t.Errorf("bad diff output") + } + var dt LineDiffType + err = dt.decode(diffBytes) + if err != nil { + t.Errorf("error decoding diff: %v\n", err) + } +} + +func TestLineDiff(t *testing.T) { + testLineDiff(t, Str1, Str2) + testLineDiff(t, Str2, Str3) + testLineDiff(t, Str1, Str3) + testLineDiff(t, Str3, Str1) + testLineDiff(t, Str3, Str4) +} + +func strMapsEqual(m1 map[string]string, m2 map[string]string) bool { + if len(m1) != len(m2) { + return false + } + for key, val := range m1 { + val2, ok := m2[key] + if !ok || val != val2 { + return false + } + } + for key, val := range m2 { + val2, ok := m1[key] + if !ok || val != val2 { + return false + } + } + return true +} + +func TestMapDiff(t *testing.T) { + m1 := map[string]string{"a": "5", "b": "hello", "c": "mike"} + m2 := map[string]string{"a": "5", "b": "goodbye", "d": "more"} + diffBytes := MakeMapDiff(m1, m2) + fmt.Printf("mapdifflen: %d\n", len(diffBytes)) + var diff MapDiffType + diff.decode(diffBytes) + diff.dump() + mcheck, err := ApplyMapDiff(m1, diffBytes) + if err != nil { + t.Fatalf("error applying map diff: %v", err) + } + if !strMapsEqual(m2, mcheck) { + t.Errorf("maps not equal") + } + fmt.Printf("%v\n", mcheck) +} From 4481cddadc0b47c026f9dcb73a267c0dd1a0927c Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 27 Nov 2022 13:47:18 -0800 Subject: [PATCH 113/149] checkpoint on statediff. bug fixes. working on more robust error handling for packetsender --- main-mshell.go | 2 +- pkg/binpack/binpack.go | 47 +++++++++++- pkg/packet/packet.go | 123 ++++++++++++++++++++---------- pkg/packet/shellstate.go | 131 ++++++++++++++++++++++++++++++++ pkg/server/server.go | 79 ++++++++++++++++--- pkg/shexec/client.go | 7 +- pkg/shexec/parser.go | 48 ++++++++++++ pkg/shexec/shexec.go | 12 +-- pkg/statediff/linediff.go | 10 +-- pkg/statediff/mapdiff.go | 10 +-- pkg/statediff/statediff_test.go | 6 +- 11 files changed, 396 insertions(+), 79 deletions(-) create mode 100644 pkg/packet/shellstate.go diff --git a/main-mshell.go b/main-mshell.go index af680c68..73f2a2d9 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -156,7 +156,7 @@ func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType func handleSingle(fromServer bool) { packetParser := packet.MakePacketParser(os.Stdin) - sender := packet.MakePacketSender(os.Stdout) + sender := packet.MakePacketSender(os.Stdout, nil) defer func() { sender.Close() sender.WaitForDone() diff --git a/pkg/binpack/binpack.go b/pkg/binpack/binpack.go index ba0e7ce3..a1145542 100644 --- a/pkg/binpack/binpack.go +++ b/pkg/binpack/binpack.go @@ -2,9 +2,15 @@ package binpack import ( "encoding/binary" + "fmt" "io" ) +type Unpacker struct { + R FullByteReader + Err error +} + type FullByteReader interface { io.ByteReader io.Reader @@ -17,9 +23,11 @@ func PackValue(w io.Writer, barr []byte) error { if err != nil { return err } - _, err = w.Write(barr) - if err != nil { - return err + if len(barr) > 0 { + _, err = w.Write(barr) + if err != nil { + return err + } } return nil } @@ -36,6 +44,9 @@ func UnpackValue(r FullByteReader) ([]byte, error) { if err != nil { return nil, err } + if lenVal == 0 { + return nil, nil + } rtnBuf := make([]byte, int(lenVal)) _, err = io.ReadFull(r, rtnBuf) if err != nil { @@ -51,3 +62,33 @@ func UnpackInt(r io.ByteReader) (int, error) { } return int(ival64), nil } + +func (u *Unpacker) UnpackValue(name string) []byte { + if u.Err != nil { + return nil + } + rtn, err := UnpackValue(u.R) + if err != nil { + u.Err = fmt.Errorf("cannot unpack %s: %v", name, err) + } + return rtn +} + +func (u *Unpacker) UnpackInt(name string) int { + if u.Err != nil { + return 0 + } + rtn, err := UnpackInt(u.R) + if err != nil { + u.Err = fmt.Errorf("cannot unpack %s: %v", name, err) + } + return rtn +} + +func (u *Unpacker) Error() error { + return u.Err +} + +func MakeUnpacker(r FullByteReader) *Unpacker { + return &Unpacker{R: r} +} diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 403c1cbe..22172355 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -49,9 +49,10 @@ const ( CdPacketStr = "cd" // rpc CmdDataPacketStr = "cmddata" // rpc-response RawPacketStr = "raw" - SpecialInputPacketStr = "sinput" // command - CompGenPacketStr = "compgen" // rpc - ReInitPacketStr = "reinit" // rpc + SpecialInputPacketStr = "sinput" // command + CompGenPacketStr = "compgen" // rpc + ReInitPacketStr = "reinit" // rpc + CmdFinalPacketStr = "cmdfinal" // command, pushed at the "end" of a command (fail-safe for no cmddone) ) const PacketSenderQueueSize = 20 @@ -80,6 +81,7 @@ func init() { TypeStrToFactory[DataEndPacketStr] = reflect.TypeOf(DataEndPacketType{}) TypeStrToFactory[CompGenPacketStr] = reflect.TypeOf(CompGenPacketType{}) TypeStrToFactory[ReInitPacketStr] = reflect.TypeOf(ReInitPacketType{}) + TypeStrToFactory[CmdFinalPacketStr] = reflect.TypeOf(CmdFinalPacketType{}) var _ RpcPacketType = (*RunPacketType)(nil) var _ RpcPacketType = (*GetCmdPacketType)(nil) @@ -96,6 +98,7 @@ func init() { var _ CommandPacketType = (*DataAckPacketType)(nil) var _ CommandPacketType = (*CmdDonePacketType)(nil) var _ CommandPacketType = (*SpecialInputPacketType)(nil) + var _ CommandPacketType = (*CmdFinalPacketType)(nil) } func RegisterPacketType(typeStr string, rtype reflect.Type) { @@ -111,19 +114,6 @@ func MakePacket(packetType string) (PacketType, error) { return rtn.Interface().(PacketType), nil } -type ShellState struct { - Version string `json:"version,omitempty"` - Cwd string `json:"cwd,omitempty"` - ShellVars []byte `json:"shellvars,omitempty"` - Aliases string `json:"aliases,omitempty"` - Funcs string `json:"funcs,omitempty"` - Error string `json:"error,omitempty"` -} - -func (state ShellState) IsEmpty() bool { - return state.Version == "" && state.Cwd == "" && len(state.ShellVars) == 0 && state.Aliases == "" && state.Funcs == "" && state.Error == "" -} - type CmdDataPacketType struct { Type string `json:"type"` RespId string `json:"respid"` @@ -509,13 +499,33 @@ func MakeDonePacket() *DonePacketType { return &DonePacketType{Type: DonePacketStr} } +type CmdFinalPacketType struct { + Type string `json:"type"` + Ts int64 `json:"ts"` + CK base.CommandKey `json:"ck"` + Error string `json:"error"` +} + +func (*CmdFinalPacketType) GetType() string { + return CmdFinalPacketStr +} + +func (pk *CmdFinalPacketType) GetCK() base.CommandKey { + return pk.CK +} + +func MakeCmdFinalPacket(ck base.CommandKey) *CmdFinalPacketType { + return &CmdFinalPacketType{Type: CmdFinalPacketStr, CK: ck} +} + type CmdDonePacketType struct { - Type string `json:"type"` - Ts int64 `json:"ts"` - CK base.CommandKey `json:"ck"` - ExitCode int `json:"exitcode"` - DurationMs int64 `json:"durationms"` - FinalState *ShellState `json:"state,omitempty"` + Type string `json:"type"` + Ts int64 `json:"ts"` + CK base.CommandKey `json:"ck"` + ExitCode int `json:"exitcode"` + DurationMs int64 `json:"durationms"` + FinalState *ShellState `json:"finalstate,omitempty"` + FinalStateDiff *ShellStateDiff `json:"finalstatediff,omitempty"` } func (*CmdDonePacketType) GetType() string { @@ -580,7 +590,8 @@ type RunPacketType struct { ReqId string `json:"reqid"` CK base.CommandKey `json:"ck"` Command string `json:"command"` - State *ShellState `json:"state"` + State *ShellState `json:"state,omitempty"` + StateDiff *ShellStateDiff `json:"statediff,omitempty"` StateComplete bool `json:"statecomplete,omitempty"` // set to true if state is complete (the default env should not be set) UsePty bool `json:"usepty,omitempty"` TermOpts *TermOpts `json:"termopts,omitempty"` @@ -696,13 +707,34 @@ func sanitizeBytes(buf []byte) { } } +type SendError struct { + IsWriteError bool // fatal + IsMarshalError bool // not fatal + PacketType string + Err error +} + +func (e *SendError) Unwrap() error { + return e.Err +} + +func (e *SendError) Error() string { + if e.IsMarshalError { + return fmt.Sprintf("SendPacket marshal-error '%s' packet: %v", e.PacketType, e.Err) + } else if e.IsWriteError { + return fmt.Sprintf("SendPacket write-error: %v", e.Err) + } else { + return e.Err.Error() + } +} + func SendPacket(w io.Writer, packet PacketType) error { if packet == nil { return nil } jsonBytes, err := json.Marshal(packet) if err != nil { - return fmt.Errorf("marshaling '%s' packet: %w", packet.GetType(), err) + return &SendError{IsMarshalError: true, PacketType: packet.GetType(), Err: err} } var outBuf bytes.Buffer outBuf.WriteByte('\n') @@ -716,7 +748,7 @@ func SendPacket(w io.Writer, packet PacketType) error { sanitizeBytes(outBytes) _, err = w.Write(outBytes) if err != nil { - return err + return &SendError{IsWriteError: true, PacketType: packet.GetType(), Err: err} } return nil } @@ -726,18 +758,19 @@ func SendCmdError(w io.Writer, ck base.CommandKey, err error) error { } type PacketSender struct { - Lock *sync.Mutex - SendCh chan PacketType - Err error - Done bool - DoneCh chan bool + Lock *sync.Mutex + SendCh chan PacketType + Done bool + DoneCh chan bool + ErrHandler func(*PacketSender, PacketType, error) } -func MakePacketSender(output io.Writer) *PacketSender { +func MakePacketSender(output io.Writer, errHandler func(*PacketSender, PacketType, error)) *PacketSender { sender := &PacketSender{ - Lock: &sync.Mutex{}, - SendCh: make(chan PacketType, PacketSenderQueueSize), - DoneCh: make(chan bool), + Lock: &sync.Mutex{}, + SendCh: make(chan PacketType, PacketSenderQueueSize), + DoneCh: make(chan bool), + ErrHandler: errHandler, } go func() { defer close(sender.DoneCh) @@ -745,9 +778,12 @@ func MakePacketSender(output io.Writer) *PacketSender { for pk := range sender.SendCh { err := SendPacket(output, pk) if err != nil { - sender.Lock.Lock() - sender.Err = err - sender.Lock.Unlock() + sender.goHandleError(pk, err) + if serr, ok := err.(*SendError); ok && serr.IsMarshalError { + // marshaler errors are recoverable + continue + } + // write errors are not recoverable return } } @@ -755,6 +791,14 @@ func MakePacketSender(output io.Writer) *PacketSender { return sender } +func (sender *PacketSender) goHandleError(pk PacketType, err error) { + sender.Lock.Lock() + defer sender.Lock.Unlock() + if sender.ErrHandler != nil { + go sender.ErrHandler(sender, pk, err) + } +} + func MakeChannelPacketSender(packetCh chan PacketType) *PacketSender { sender := &PacketSender{ Lock: &sync.Mutex{}, @@ -791,9 +835,6 @@ func (sender *PacketSender) checkStatus() error { if sender.Done { return fmt.Errorf("cannot send packet, sender write loop is closed") } - if sender.Err != nil { - return fmt.Errorf("cannot send packet, sender had error: %w", sender.Err) - } return nil } @@ -833,7 +874,7 @@ func (sender *PacketSender) SendResponse(reqId string, data interface{}) error { return sender.SendPacket(pk) } -func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) error { +func (sender *PacketSender) SendMessageFmt(fmtStr string, args ...interface{}) error { return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) } diff --git a/pkg/packet/shellstate.go b/pkg/packet/shellstate.go new file mode 100644 index 00000000..18297a0c --- /dev/null +++ b/pkg/packet/shellstate.go @@ -0,0 +1,131 @@ +package packet + +import ( + "bytes" + "crypto/sha1" + "encoding/base64" + "encoding/json" + "fmt" + + "github.com/scripthaus-dev/mshell/pkg/binpack" + "github.com/scripthaus-dev/mshell/pkg/statediff" +) + +const ShellStatePackVersion = 0 +const ShellStateDiffPackVersion = 0 + +type ShellState struct { + Version string `json:"version"` // [type] [semver] + Cwd string `json:"cwd,omitempty"` + ShellVars []byte `json:"shellvars,omitempty"` + Aliases string `json:"aliases,omitempty"` + Funcs string `json:"funcs,omitempty"` + Error string `json:"error,omitempty"` +} + +type ShellStateDiff struct { + Version string `json:"version"` // [type] [semver] + BaseHash string `json:"basehash"` + Cwd string `json:"cwd,omitempty"` + VarsDiff []byte `json:"shellvarsdiff,omitempty"` // vardiff + AliasesDiff []byte `json:"aliasesdiff,omitempty"` // linediff + FuncsDiff []byte `json:"funcsdiff,omitempty"` // linediff + Error string `json:"error,omitempty"` +} + +func (state ShellState) IsEmpty() bool { + return state.Version == "" && state.Cwd == "" && len(state.ShellVars) == 0 && state.Aliases == "" && state.Funcs == "" && state.Error == "" +} + +// returns (SHA1, encoded-state) +func (state ShellState) EncodeAndHash() (string, []byte) { + var buf bytes.Buffer + binpack.PackInt(&buf, ShellStatePackVersion) + binpack.PackValue(&buf, []byte(state.Version)) + binpack.PackValue(&buf, []byte(state.Cwd)) + binpack.PackValue(&buf, state.ShellVars) + binpack.PackValue(&buf, []byte(state.Aliases)) + binpack.PackValue(&buf, []byte(state.Funcs)) + binpack.PackValue(&buf, []byte(state.Error)) + hvalRaw := sha1.Sum(buf.Bytes()) + hval := base64.StdEncoding.EncodeToString(hvalRaw[:]) + return hval, buf.Bytes() +} + +func (state ShellState) MarshalJSON() ([]byte, error) { + _, encodedState := state.EncodeAndHash() + return json.Marshal(encodedState) +} + +func (state *ShellState) UnmarshalJSON(jsonBytes []byte) error { + var barr []byte + err := json.Unmarshal(jsonBytes, &barr) + if err != nil { + return err + } + buf := bytes.NewBuffer(barr) + u := binpack.MakeUnpacker(buf) + version := u.UnpackInt("ShellState pack version") + if version != ShellStatePackVersion { + return fmt.Errorf("invalid ShellState pack version: %d", version) + } + state.Version = string(u.UnpackValue("ShellState.Version")) + state.Cwd = string(u.UnpackValue("ShellState.Cwd")) + state.ShellVars = u.UnpackValue("ShellState.ShellVars") + state.Aliases = string(u.UnpackValue("ShellState.Aliases")) + state.Funcs = string(u.UnpackValue("ShellState.Funcs")) + state.Error = string(u.UnpackValue("ShellState.Error")) + return u.Error() +} + +func (sdiff ShellStateDiff) MarshalJSON() ([]byte, error) { + var buf bytes.Buffer + binpack.PackInt(&buf, ShellStateDiffPackVersion) + binpack.PackValue(&buf, []byte(sdiff.Version)) + binpack.PackValue(&buf, []byte(sdiff.BaseHash)) + binpack.PackValue(&buf, []byte(sdiff.Cwd)) + binpack.PackValue(&buf, sdiff.VarsDiff) + binpack.PackValue(&buf, sdiff.AliasesDiff) + binpack.PackValue(&buf, sdiff.FuncsDiff) + binpack.PackValue(&buf, []byte(sdiff.Error)) + return buf.Bytes(), nil +} + +func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { + var barr []byte + err := json.Unmarshal(jsonBytes, &barr) + if err != nil { + return err + } + buf := bytes.NewBuffer(barr) + u := binpack.MakeUnpacker(buf) + version := u.UnpackInt("ShellState pack version") + if version != ShellStateDiffPackVersion { + return fmt.Errorf("invalid ShellStateDiff pack version: %d", version) + } + sdiff.Version = string(u.UnpackValue("ShellStateDiff.Version")) + sdiff.BaseHash = string(u.UnpackValue("ShellStateDiff.BaseHash")) + sdiff.Cwd = string(u.UnpackValue("ShellStateDiff.Cwd")) + sdiff.VarsDiff = u.UnpackValue("ShellStateDiff.VarsDiff") + sdiff.AliasesDiff = u.UnpackValue("ShellStateDiff.AliasesDiff") + sdiff.FuncsDiff = u.UnpackValue("ShellStateDiff.FuncsDiff") + sdiff.Error = string(u.UnpackValue("ShellStateDiff.Error")) + return u.Error() +} + +func (sdiff ShellStateDiff) Dump() { + fmt.Printf("ShellStateDiff:\n") + fmt.Printf(" version: %s\n", sdiff.Version) + fmt.Printf(" base: %s\n", sdiff.BaseHash) + var mdiff statediff.MapDiffType + err := mdiff.Decode(sdiff.VarsDiff) + if err != nil { + fmt.Printf(" vars: error[%s]\n", err.Error()) + } else { + mdiff.Dump() + } + fmt.Printf(" aliases: %d, funcs: %d\n", len(sdiff.AliasesDiff), len(sdiff.FuncsDiff)) + if sdiff.Error != "" { + fmt.Printf(" error: %s\n", sdiff.Error) + } +} diff --git a/pkg/server/server.go b/pkg/server/server.go index 6ad8169d..9ff06dff 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -14,6 +14,7 @@ import ( "sort" "strings" "sync" + "time" "github.com/alessio/shellescape" "github.com/scripthaus-dev/mshell/pkg/base" @@ -23,11 +24,13 @@ import ( // 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 + 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 } func (m *MServer) Close() { @@ -38,7 +41,7 @@ func (m *MServer) Close() { func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { ck := pk.GetCK() if ck == "" { - m.Sender.SendMessage(fmt.Sprintf("received '%s' packet without ck", pk.GetType())) + m.Sender.SendMessageFmt("received '%s' packet without ck", pk.GetType()) return } m.Lock.Lock() @@ -137,12 +140,24 @@ func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { return } +func (m *MServer) setCurrentState(state *packet.ShellState) { + if state == nil { + return + } + hval, _ := state.EncodeAndHash() + m.Lock.Lock() + defer m.Lock.Unlock() + m.StateMap[hval] = state + m.CurrentState = hval +} + 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 } + m.setCurrentState(initPk.State) initPk.RespId = reqId m.Sender.SendPacket(initPk) } @@ -170,6 +185,32 @@ func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) { return } +func (m *MServer) getCurrentState() (string, *packet.ShellState) { + m.Lock.Lock() + defer m.Lock.Unlock() + return m.CurrentState, m.StateMap[m.CurrentState] +} + +func (m *MServer) clientPacketCallback(pk packet.PacketType) { + if pk.GetType() != packet.CmdDonePacketStr { + return + } + donePk := pk.(*packet.CmdDonePacketType) + if donePk.FinalState == nil { + return + } + stateHash, curState := m.getCurrentState() + if curState == nil { + return + } + diff, err := shexec.MakeShellStateDiff(*curState, stateHash, *donePk.FinalState) + if err != nil { + return + } + donePk.FinalState = nil + donePk.FinalStateDiff = &diff +} + func (m *MServer) runCommand(runPacket *packet.RunPacketType) { if err := runPacket.CK.Validate("packet"); err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) @@ -190,16 +231,34 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Lock.Unlock() go func() { defer func() { + r := recover() + finalPk := packet.MakeCmdFinalPacket(runPacket.CK) + finalPk.Ts = time.Now().UnixMilli() + if r != nil { + finalPk.Error = fmt.Sprintf("%s", r) + } + m.Sender.SendPacket(finalPk) m.Lock.Lock() delete(m.ClientMap, runPacket.CK) m.Lock.Unlock() cproc.Close() }() shexec.SendRunPacketAndRunData(context.Background(), cproc.Input, runPacket) - cproc.ProxySingleOutput(runPacket.CK, m.Sender) + cproc.ProxySingleOutput(runPacket.CK, m.Sender, m.clientPacketCallback) }() } +func (m *MServer) packetSenderErrorHandler(sender *packet.PacketSender, pk packet.PacketType, err error) { + if serr, ok := err.(*packet.SendError); ok && serr.IsMarshalError { + msg := packet.MakeMessagePacket(err.Error()) + if cpk, ok := pk.(packet.CommandPacketType); ok { + msg.CK = cpk.GetCK() + } + sender.SendPacket(msg) + } + // otherwise ignore (we can't output anything for a I/O error) +} + func RunServer() (int, error) { debug := false if len(os.Args) >= 3 && os.Args[2] == "--debug" { @@ -208,19 +267,21 @@ func RunServer() (int, error) { server := &MServer{ Lock: &sync.Mutex{}, ClientMap: make(map[base.CommandKey]*shexec.ClientProc), + StateMap: make(map[string]*packet.ShellState), Debug: debug, } if debug { packet.GlobalDebug = true } server.MainInput = packet.MakePacketParser(os.Stdin) - server.Sender = packet.MakePacketSender(os.Stdout) + server.Sender = packet.MakePacketSender(os.Stdout, server.packetSenderErrorHandler) defer server.Close() var err error initPacket, err := shexec.MakeServerInitPacket() if err != nil { return 1, err } + server.setCurrentState(initPacket.State) server.Sender.SendPacket(initPacket) builder := packet.MakeRunPacketBuilder() for pk := range server.MainInput.MainCh { @@ -243,7 +304,7 @@ func RunServer() (int, error) { server.ProcessRpcPacket(rpcPk) continue } - server.Sender.SendMessage(fmt.Sprintf("invalid packet '%s' sent to mshell server", packet.AsString(pk))) + server.Sender.SendMessageFmt("invalid packet '%s' sent to mshell server", packet.AsString(pk)) continue } return 0, nil diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index 18ea2478..9dd5d974 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -46,7 +46,7 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.I if err != nil { return nil, nil, fmt.Errorf("running local client: %w", err) } - sender := packet.MakePacketSender(inputWriter) + sender := packet.MakePacketSender(inputWriter, nil) stdoutPacketParser := packet.MakePacketParser(stdoutReader) stderrPacketParser := packet.MakePacketParser(stderrReader) packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) @@ -108,9 +108,12 @@ func (cproc *ClientProc) Close() { } } -func (cproc *ClientProc) ProxySingleOutput(ck base.CommandKey, sender *packet.PacketSender) { +func (cproc *ClientProc) ProxySingleOutput(ck base.CommandKey, sender *packet.PacketSender, packetCallback func(packet.PacketType)) { sentDonePk := false for pk := range cproc.Output.MainCh { + if packetCallback != nil { + packetCallback(pk) + } if pk.GetType() == packet.CmdDonePacketStr { sentDonePk = true } diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 7b96c1a4..7b331927 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -10,6 +10,7 @@ import ( "github.com/alessio/shellescape" "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/statediff" "mvdan.cc/sh/v3/expand" "mvdan.cc/sh/v3/syntax" ) @@ -189,6 +190,32 @@ func ParseDeclLine(envLine string) *DeclareDeclType { } } +func parseDeclLineToKV(envLine string) (string, string) { + eqIdx := strings.Index(envLine, "=") + if eqIdx == -1 { + return "", "" + } + namePart := envLine[0:eqIdx] + valPart := envLine[eqIdx+1:] + return namePart, valPart +} + +func shellStateVarsToMap(shellVars []byte) map[string]string { + if len(shellVars) == 0 { + return nil + } + rtn := make(map[string]string) + vars := bytes.Split(shellVars, []byte{0}) + for _, varLine := range vars { + name, val := parseDeclLineToKV(string(varLine)) + if name == "" { + continue + } + rtn[name] = val + } + return rtn +} + func DeclMapFromState(state *packet.ShellState) map[string]*DeclareDeclType { if state == nil { return nil @@ -485,3 +512,24 @@ func DeclsEqual(compareName bool, d1 *DeclareDeclType, d2 *DeclareDeclType) bool } return d1.Value == d2.Value // this works even for assoc arrays because we normalize them when parsing } + +func MakeShellStateDiff(oldState packet.ShellState, oldStateHash string, newState packet.ShellState) (packet.ShellStateDiff, error) { + var rtn packet.ShellStateDiff + rtn.BaseHash = oldStateHash + if oldState.Version != newState.Version { + return rtn, fmt.Errorf("cannot diff, states have different versions") + } + rtn.Version = newState.Version + if oldState.Cwd != newState.Cwd { + rtn.Cwd = newState.Cwd + } + if oldState.Error != newState.Error { + rtn.Error = newState.Error + } + oldVars := shellStateVarsToMap(oldState.ShellVars) + newVars := shellStateVarsToMap(newState.ShellVars) + rtn.VarsDiff = statediff.MakeMapDiff(oldVars, newVars) + rtn.AliasesDiff = statediff.MakeLineDiff(oldState.Aliases, newState.Aliases) + rtn.FuncsDiff = statediff.MakeLineDiff(oldState.Funcs, newState.Funcs) + return rtn, nil +} diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 0b5ecf04..af56f4b8 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -834,7 +834,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon stdoutPacketParser := packet.MakePacketParser(stdoutReader) stderrPacketParser := packet.MakePacketParser(stderrReader) packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) - sender := packet.MakePacketSender(inputWriter) + sender := packet.MakePacketSender(inputWriter, nil) versionOk := false for pk := range packetParser.MainCh { if pk.GetType() == packet.RawPacketStr { @@ -986,14 +986,6 @@ func makeRcFileStr(pk *packet.RunPacketType) string { rcBuf.WriteString(pk.State.Aliases) rcBuf.WriteString("\n") } - if pk.ReturnState { - rcBuf.WriteString(` -_scripthaus_exittrap () { -` + GetShellStateCmd + ` -} -trap _scripthaus_exittrap EXIT -`) - } return rcBuf.String() } @@ -1285,7 +1277,7 @@ func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( if err != nil { return nil, nil, fmt.Errorf("cannot open runout file '%s': %w", fileNames.RunnerOutFile, err) } - cmd.DetachedOutput = packet.MakePacketSender(cmd.RunnerOutFd) + cmd.DetachedOutput = packet.MakePacketSender(cmd.RunnerOutFd, nil) ecmd, err := MakeDetachedExecCmd(pk, cmdTty) if err != nil { return nil, nil, err diff --git a/pkg/statediff/linediff.go b/pkg/statediff/linediff.go index a2e36491..347b2e37 100644 --- a/pkg/statediff/linediff.go +++ b/pkg/statediff/linediff.go @@ -19,7 +19,7 @@ type LineDiffType struct { NewData []string } -func (diff LineDiffType) dump() { +func (diff LineDiffType) Dump() { fmt.Printf("DIFF:\n") pos := 1 for _, entry := range diff.Lines { @@ -68,7 +68,7 @@ func putUVarint(buf *bytes.Buffer, viBuf []byte, ival int) { // simple encoding // write varints. first version, then len, then len-number-of-varints, then fill the rest with newdata // [version] [len-varint] [varint]xlen... newdata (bytes) -func (diff LineDiffType) encode() []byte { +func (diff LineDiffType) Encode() []byte { var buf bytes.Buffer viBuf := make([]byte, binary.MaxVarintLen64) putUVarint(&buf, viBuf, LineDiffVersion) @@ -86,7 +86,7 @@ func (diff LineDiffType) encode() []byte { return buf.Bytes() } -func (rtn *LineDiffType) decode(diffBytes []byte) error { +func (rtn *LineDiffType) Decode(diffBytes []byte) error { r := bytes.NewBuffer(diffBytes) version, err := binary.ReadUvarint(r) if err != nil { @@ -164,12 +164,12 @@ func MakeLineDiff(str1 string, str2 string) []byte { str1Arr := strings.Split(str1, "\n") str2Arr := strings.Split(str2, "\n") diff := makeLineDiff(str1Arr, str2Arr) - return diff.encode() + return diff.Encode() } func ApplyLineDiff(str1 string, diffBytes []byte) (string, error) { var diff LineDiffType - err := diff.decode(diffBytes) + err := diff.Decode(diffBytes) if err != nil { return "", err } diff --git a/pkg/statediff/mapdiff.go b/pkg/statediff/mapdiff.go index a06298b7..da0d14da 100644 --- a/pkg/statediff/mapdiff.go +++ b/pkg/statediff/mapdiff.go @@ -15,7 +15,7 @@ type MapDiffType struct { ToRemove []string } -func (diff MapDiffType) dump() { +func (diff MapDiffType) Dump() { fmt.Printf("VAR-DIFF\n") for name, val := range diff.ToAdd { fmt.Printf(" add: %s=%s\n", name, val) @@ -58,7 +58,7 @@ func (diff MapDiffType) apply(oldMap map[string]string) map[string]string { return rtn } -func (diff MapDiffType) encode() []byte { +func (diff MapDiffType) Encode() []byte { var buf bytes.Buffer viBuf := make([]byte, binary.MaxVarintLen64) putUVarint(&buf, viBuf, MapDiffVersion) @@ -76,7 +76,7 @@ func (diff MapDiffType) encode() []byte { return buf.Bytes() } -func (diff *MapDiffType) decode(diffBytes []byte) error { +func (diff *MapDiffType) Decode(diffBytes []byte) error { r := bytes.NewBuffer(diffBytes) version, err := binary.ReadUvarint(r) if err != nil { @@ -111,12 +111,12 @@ func (diff *MapDiffType) decode(diffBytes []byte) error { func MakeMapDiff(m1 map[string]string, m2 map[string]string) []byte { diff := makeMapDiff(m1, m2) - return diff.encode() + return diff.Encode() } func ApplyMapDiff(oldMap map[string]string, diffBytes []byte) (map[string]string, error) { var diff MapDiffType - err := diff.decode(diffBytes) + err := diff.Decode(diffBytes) if err != nil { return nil, err } diff --git a/pkg/statediff/statediff_test.go b/pkg/statediff/statediff_test.go index 3b231c34..34dc478c 100644 --- a/pkg/statediff/statediff_test.go +++ b/pkg/statediff/statediff_test.go @@ -47,7 +47,7 @@ func testLineDiff(t *testing.T, str1 string, str2 string) { t.Errorf("bad diff output") } var dt LineDiffType - err = dt.decode(diffBytes) + err = dt.Decode(diffBytes) if err != nil { t.Errorf("error decoding diff: %v\n", err) } @@ -86,8 +86,8 @@ func TestMapDiff(t *testing.T) { diffBytes := MakeMapDiff(m1, m2) fmt.Printf("mapdifflen: %d\n", len(diffBytes)) var diff MapDiffType - diff.decode(diffBytes) - diff.dump() + diff.Decode(diffBytes) + diff.Dump() mcheck, err := ApplyMapDiff(m1, diffBytes) if err != nil { t.Fatalf("error applying map diff: %v", err) From eb3cf8032912003100a20d9e5c0a186949111f37 Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 27 Nov 2022 14:16:25 -0800 Subject: [PATCH 114/149] fix json marshaling bug for statediff --- pkg/packet/shellstate.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/packet/shellstate.go b/pkg/packet/shellstate.go index 18297a0c..f4187681 100644 --- a/pkg/packet/shellstate.go +++ b/pkg/packet/shellstate.go @@ -88,7 +88,7 @@ func (sdiff ShellStateDiff) MarshalJSON() ([]byte, error) { binpack.PackValue(&buf, sdiff.AliasesDiff) binpack.PackValue(&buf, sdiff.FuncsDiff) binpack.PackValue(&buf, []byte(sdiff.Error)) - return buf.Bytes(), nil + return json.Marshal(buf.Bytes()) } func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { From 605d0899cfd10fb05a4b4387167824e912be9013 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 28 Nov 2022 00:15:34 -0800 Subject: [PATCH 115/149] checkpoint --- pkg/binpack/binpack.go | 33 ++++++++++++ pkg/packet/shellstate.go | 103 ++++++++++++++++++++++++++------------ pkg/shexec/parser.go | 83 ++++++++++++++++++++++++++---- pkg/statediff/linediff.go | 6 +++ pkg/statediff/mapdiff.go | 6 +++ 5 files changed, 191 insertions(+), 40 deletions(-) diff --git a/pkg/binpack/binpack.go b/pkg/binpack/binpack.go index a1145542..67e25337 100644 --- a/pkg/binpack/binpack.go +++ b/pkg/binpack/binpack.go @@ -2,6 +2,7 @@ package binpack import ( "encoding/binary" + "encoding/json" "fmt" "io" ) @@ -32,6 +33,14 @@ func PackValue(w io.Writer, barr []byte) error { return nil } +func PackStrArr(w io.Writer, strs []string) error { + barr, err := json.Marshal(strs) + if err != nil { + return err + } + return PackValue(w, barr) +} + func PackInt(w io.Writer, ival int) error { viBuf := make([]byte, binary.MaxVarintLen64) l := binary.PutUvarint(viBuf, uint64(ival)) @@ -55,6 +64,19 @@ func UnpackValue(r FullByteReader) ([]byte, error) { return rtnBuf, nil } +func UnpackStrArr(r FullByteReader) ([]string, error) { + barr, err := UnpackValue(r) + if err != nil { + return nil, err + } + var strs []string + err = json.Unmarshal(barr, &strs) + if err != nil { + return nil, err + } + return strs, nil +} + func UnpackInt(r io.ByteReader) (int, error) { ival64, err := binary.ReadVarint(r) if err != nil { @@ -85,6 +107,17 @@ func (u *Unpacker) UnpackInt(name string) int { return rtn } +func (u *Unpacker) UnpackStrArr(name string) []string { + if u.Err != nil { + return nil + } + rtn, err := UnpackStrArr(u.R) + if err != nil { + u.Err = fmt.Errorf("cannot unpack %s: %v", name, err) + } + return rtn +} + func (u *Unpacker) Error() error { return u.Err } diff --git a/pkg/packet/shellstate.go b/pkg/packet/shellstate.go index f4187681..fce4026c 100644 --- a/pkg/packet/shellstate.go +++ b/pkg/packet/shellstate.go @@ -21,22 +21,32 @@ type ShellState struct { Aliases string `json:"aliases,omitempty"` Funcs string `json:"funcs,omitempty"` Error string `json:"error,omitempty"` + HashVal string `json:"-"` } type ShellStateDiff struct { - Version string `json:"version"` // [type] [semver] - BaseHash string `json:"basehash"` - Cwd string `json:"cwd,omitempty"` - VarsDiff []byte `json:"shellvarsdiff,omitempty"` // vardiff - AliasesDiff []byte `json:"aliasesdiff,omitempty"` // linediff - FuncsDiff []byte `json:"funcsdiff,omitempty"` // linediff - Error string `json:"error,omitempty"` + Version string `json:"version"` // [type] [semver] + BaseHash string `json:"basehash"` + DiffHashArr []string `json:"diffhasharr,omitempty"` + Cwd string `json:"cwd,omitempty"` + VarsDiff []byte `json:"shellvarsdiff,omitempty"` // vardiff + AliasesDiff []byte `json:"aliasesdiff,omitempty"` // linediff + FuncsDiff []byte `json:"funcsdiff,omitempty"` // linediff + Error string `json:"error,omitempty"` + HashVal string `json:"-"` } func (state ShellState) IsEmpty() bool { return state.Version == "" && state.Cwd == "" && len(state.ShellVars) == 0 && state.Aliases == "" && state.Funcs == "" && state.Error == "" } +// returns base64 hash of data +func sha1Hash(data []byte) string { + hvalRaw := sha1.Sum(data) + hval := base64.StdEncoding.EncodeToString(hvalRaw[:]) + return hval +} + // returns (SHA1, encoded-state) func (state ShellState) EncodeAndHash() (string, []byte) { var buf bytes.Buffer @@ -47,22 +57,24 @@ func (state ShellState) EncodeAndHash() (string, []byte) { binpack.PackValue(&buf, []byte(state.Aliases)) binpack.PackValue(&buf, []byte(state.Funcs)) binpack.PackValue(&buf, []byte(state.Error)) - hvalRaw := sha1.Sum(buf.Bytes()) - hval := base64.StdEncoding.EncodeToString(hvalRaw[:]) - return hval, buf.Bytes() + return sha1Hash(buf.Bytes()), buf.Bytes() } func (state ShellState) MarshalJSON() ([]byte, error) { - _, encodedState := state.EncodeAndHash() - return json.Marshal(encodedState) + _, encodedBytes := state.EncodeAndHash() + return json.Marshal(encodedBytes) } -func (state *ShellState) UnmarshalJSON(jsonBytes []byte) error { - var barr []byte - err := json.Unmarshal(jsonBytes, &barr) - if err != nil { - return err +// caches HashVal in struct +func (state *ShellState) GetHashVal(force bool) string { + if state.HashVal == "" || force { + state.HashVal, _ = state.EncodeAndHash() } + return state.HashVal +} + +func (state *ShellState) DecodeShellState(barr []byte) error { + state.HashVal = sha1Hash(barr) buf := bytes.NewBuffer(barr) u := binpack.MakeUnpacker(buf) version := u.UnpackInt("ShellState pack version") @@ -78,25 +90,36 @@ func (state *ShellState) UnmarshalJSON(jsonBytes []byte) error { return u.Error() } -func (sdiff ShellStateDiff) MarshalJSON() ([]byte, error) { - var buf bytes.Buffer - binpack.PackInt(&buf, ShellStateDiffPackVersion) - binpack.PackValue(&buf, []byte(sdiff.Version)) - binpack.PackValue(&buf, []byte(sdiff.BaseHash)) - binpack.PackValue(&buf, []byte(sdiff.Cwd)) - binpack.PackValue(&buf, sdiff.VarsDiff) - binpack.PackValue(&buf, sdiff.AliasesDiff) - binpack.PackValue(&buf, sdiff.FuncsDiff) - binpack.PackValue(&buf, []byte(sdiff.Error)) - return json.Marshal(buf.Bytes()) -} - -func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { +func (state *ShellState) UnmarshalJSON(jsonBytes []byte) error { var barr []byte err := json.Unmarshal(jsonBytes, &barr) if err != nil { return err } + return state.DecodeShellState(barr) +} + +func (sdiff ShellStateDiff) EncodeAndHash() (string, []byte) { + var buf bytes.Buffer + binpack.PackInt(&buf, ShellStateDiffPackVersion) + binpack.PackValue(&buf, []byte(sdiff.Version)) + binpack.PackValue(&buf, []byte(sdiff.BaseHash)) + binpack.PackStrArr(&buf, sdiff.DiffHashArr) + binpack.PackValue(&buf, []byte(sdiff.Cwd)) + binpack.PackValue(&buf, sdiff.VarsDiff) + binpack.PackValue(&buf, sdiff.AliasesDiff) + binpack.PackValue(&buf, sdiff.FuncsDiff) + binpack.PackValue(&buf, []byte(sdiff.Error)) + return sha1Hash(buf.Bytes()), buf.Bytes() +} + +func (sdiff ShellStateDiff) MarshalJSON() ([]byte, error) { + _, encodedBytes := sdiff.EncodeAndHash() + return json.Marshal(encodedBytes) +} + +func (sdiff *ShellStateDiff) DecodeShellStateDiff(barr []byte) error { + sdiff.HashVal = sha1Hash(barr) buf := bytes.NewBuffer(barr) u := binpack.MakeUnpacker(buf) version := u.UnpackInt("ShellState pack version") @@ -105,6 +128,7 @@ func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { } sdiff.Version = string(u.UnpackValue("ShellStateDiff.Version")) sdiff.BaseHash = string(u.UnpackValue("ShellStateDiff.BaseHash")) + sdiff.DiffHashArr = u.UnpackStrArr("ShellStateDiff.DiffHashArr") sdiff.Cwd = string(u.UnpackValue("ShellStateDiff.Cwd")) sdiff.VarsDiff = u.UnpackValue("ShellStateDiff.VarsDiff") sdiff.AliasesDiff = u.UnpackValue("ShellStateDiff.AliasesDiff") @@ -113,6 +137,23 @@ func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { return u.Error() } +func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { + var barr []byte + err := json.Unmarshal(jsonBytes, &barr) + if err != nil { + return err + } + return sdiff.DecodeShellStateDiff(barr) +} + +// caches HashVal in struct +func (sdiff *ShellStateDiff) GetHashVal(force bool) string { + if sdiff.HashVal == "" || force { + sdiff.HashVal, _ = sdiff.EncodeAndHash() + } + return sdiff.HashVal +} + func (sdiff ShellStateDiff) Dump() { fmt.Printf("ShellStateDiff:\n") fmt.Printf(" version: %s\n", sdiff.Version) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 7b331927..e75cbfeb 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -190,14 +190,13 @@ func ParseDeclLine(envLine string) *DeclareDeclType { } } +// returns name => full-line func parseDeclLineToKV(envLine string) (string, string) { - eqIdx := strings.Index(envLine, "=") - if eqIdx == -1 { + decl := ParseDeclLine(envLine) + if decl == nil { return "", "" } - namePart := envLine[0:eqIdx] - valPart := envLine[eqIdx+1:] - return namePart, valPart + return decl.Name, envLine } func shellStateVarsToMap(shellVars []byte) map[string]string { @@ -216,6 +215,35 @@ func shellStateVarsToMap(shellVars []byte) map[string]string { return rtn } +func strMapToShellStateVars(varMap map[string]string) []byte { + var buf bytes.Buffer + orderedKeys := getOrderedKeysStrMap(varMap) + for _, key := range orderedKeys { + val := varMap[key] + buf.WriteString(val) + buf.WriteByte(0) + } + return buf.Bytes() +} + +func getOrderedKeysStrMap(m map[string]string) []string { + keys := make([]string, 0, len(m)) + for key, _ := range m { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +func getOrderedKeysDeclMap(m map[string]*DeclareDeclType) []string { + keys := make([]string, 0, len(m)) + for key, _ := range m { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + func DeclMapFromState(state *packet.ShellState) map[string]*DeclareDeclType { if state == nil { return nil @@ -233,7 +261,9 @@ func DeclMapFromState(state *packet.ShellState) map[string]*DeclareDeclType { func SerializeDeclMap(declMap map[string]*DeclareDeclType) []byte { var rtn bytes.Buffer - for _, decl := range declMap { + orderedKeys := getOrderedKeysDeclMap(declMap) + for _, key := range orderedKeys { + decl := declMap[key] rtn.WriteString(decl.Serialize()) } return rtn.Bytes() @@ -269,6 +299,18 @@ func ShellVarMapFromState(state *packet.ShellState) map[string]string { return rtn } +func DumpVarMapFromState(state *packet.ShellState) { + fmt.Printf("DUMP-STATE-VARS:\n") + if state == nil { + fmt.Printf(" nil\n") + return + } + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + fmt.Printf(" %s\n", varLine) + } +} + func VarDeclsFromState(state *packet.ShellState) []*DeclareDeclType { if state == nil { return nil @@ -365,8 +407,8 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { if err != nil { return err } - var varsBuffer bytes.Buffer var firstParseErr error + declMap := make(map[string]*DeclareDeclType) for _, stmt := range file.Stmts { decl, err := parseDeclareStmt(stmt, declareStr) if err != nil { @@ -375,10 +417,10 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { } } if decl != nil && !NoStoreVarNames[decl.Name] { - varsBuffer.WriteString(decl.Serialize()) + declMap[decl.Name] = decl } } - state.ShellVars = varsBuffer.Bytes() + state.ShellVars = SerializeDeclMap(declMap) // this writes out the decls in a canonical order if firstParseErr != nil { state.Error = firstParseErr.Error() } @@ -533,3 +575,26 @@ func MakeShellStateDiff(oldState packet.ShellState, oldStateHash string, newStat rtn.FuncsDiff = statediff.MakeLineDiff(oldState.Funcs, newState.Funcs) return rtn, nil } + +func ApplyShellStateDiff(oldState packet.ShellState, diff packet.ShellStateDiff) (packet.ShellState, error) { + var rtnState packet.ShellState + var err error + rtnState.Version = oldState.Version + rtnState.Cwd = diff.Cwd + rtnState.Error = diff.Error + oldVars := shellStateVarsToMap(oldState.ShellVars) + newVars, err := statediff.ApplyMapDiff(oldVars, diff.VarsDiff) + if err != nil { + return rtnState, fmt.Errorf("applying mapdiff 'vars': %v", err) + } + rtnState.ShellVars = strMapToShellStateVars(newVars) + rtnState.Aliases, err = statediff.ApplyLineDiff(oldState.Aliases, diff.AliasesDiff) + if err != nil { + return rtnState, fmt.Errorf("applying diff 'aliases': %v", err) + } + rtnState.Funcs, err = statediff.ApplyLineDiff(oldState.Funcs, diff.FuncsDiff) + if err != nil { + return rtnState, fmt.Errorf("applying diff 'funcs': %v", err) + } + return rtnState, nil +} diff --git a/pkg/statediff/linediff.go b/pkg/statediff/linediff.go index 347b2e37..fa7ce11b 100644 --- a/pkg/statediff/linediff.go +++ b/pkg/statediff/linediff.go @@ -161,6 +161,9 @@ func makeLineDiff(oldData []string, newData []string) LineDiffType { } func MakeLineDiff(str1 string, str2 string) []byte { + if str1 == str2 { + return nil + } str1Arr := strings.Split(str1, "\n") str2Arr := strings.Split(str2, "\n") diff := makeLineDiff(str1Arr, str2Arr) @@ -168,6 +171,9 @@ func MakeLineDiff(str1 string, str2 string) []byte { } func ApplyLineDiff(str1 string, diffBytes []byte) (string, error) { + if len(diffBytes) == 0 { + return str1, nil + } var diff LineDiffType err := diff.Decode(diffBytes) if err != nil { diff --git a/pkg/statediff/mapdiff.go b/pkg/statediff/mapdiff.go index da0d14da..8e810f7f 100644 --- a/pkg/statediff/mapdiff.go +++ b/pkg/statediff/mapdiff.go @@ -111,10 +111,16 @@ func (diff *MapDiffType) Decode(diffBytes []byte) error { func MakeMapDiff(m1 map[string]string, m2 map[string]string) []byte { diff := makeMapDiff(m1, m2) + if len(diff.ToAdd) == 0 && len(diff.ToRemove) == 0 { + return nil + } return diff.Encode() } func ApplyMapDiff(oldMap map[string]string, diffBytes []byte) (map[string]string, error) { + if len(diffBytes) == 0 { + return oldMap, nil + } var diff MapDiffType err := diff.Decode(diffBytes) if err != nil { From bdd8381b0166904eb4788e7ab56e0dd7cfd46614 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 28 Nov 2022 18:05:54 -0800 Subject: [PATCH 116/149] updates/bugfixes for statediff --- main-mshell.go | 3 ++ pkg/base/base.go | 30 ++++++++++++++++++++ pkg/packet/shellstate.go | 38 +++++++++++++++++++------ pkg/shexec/parser.go | 48 +++++++++++++++++++++++++------- pkg/simpleexpand/simpleexpand.go | 9 ++++++ pkg/statediff/mapdiff.go | 4 +-- 6 files changed, 111 insertions(+), 21 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 73f2a2d9..0ad82482 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -536,12 +536,15 @@ func main() { os.Exit(1) } } else if firstArg == "--single" { + base.InitDebugLog("single") handleSingle(false) return } else if firstArg == "--single-from-server" { + base.InitDebugLog("single") handleSingle(true) return } else if firstArg == "--server" { + base.InitDebugLog("server") rtnCode, err := server.RunServer() if err != nil { fmt.Fprintf(os.Stderr, "[error] %v\n", err) diff --git a/pkg/base/base.go b/pkg/base/base.go index 3c9840db..4b87a534 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -11,6 +11,7 @@ import ( "fmt" "io" "io/fs" + "log" "os" "os/exec" "path" @@ -33,9 +34,13 @@ const SessionsDirBaseName = "sessions" const MShellVersion = "v0.2.0" const RemoteIdFile = "remoteid" const DefaultMShellInstallBinDir = "/opt/mshell/bin" +const LogFileName = "mshell.log" +const ForceDebugLog = false var sessionDirCache = make(map[string]string) var baseLock = &sync.Mutex{} +var DebugLogEnabled = false +var DebugLogger *log.Logger type CommandFileNames struct { PtyOutFile string @@ -56,6 +61,31 @@ func (ckey CommandKey) IsEmpty() bool { return string(ckey) == "" } +func Logf(fmtStr string, args ...interface{}) { + if (!DebugLogEnabled && !ForceDebugLog) || DebugLogger == nil { + return + } + DebugLogger.Printf(fmtStr, args...) +} + +func InitDebugLog(prefix string) { + homeDir := GetMShellHomeDir() + err := os.MkdirAll(homeDir, 0777) + if err != nil { + return + } + logFile := path.Join(homeDir, LogFileName) + fd, err := os.OpenFile(logFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600) + if err != nil { + return + } + DebugLogger = log.New(fd, prefix+" ", log.LstdFlags) +} + +func SetEnableDebugLog(enable bool) { + DebugLogEnabled = enable +} + func (ckey CommandKey) GetSessionId() string { slashIdx := strings.Index(string(ckey), "/") if slashIdx == -1 { diff --git a/pkg/packet/shellstate.go b/pkg/packet/shellstate.go index fce4026c..3bc5957f 100644 --- a/pkg/packet/shellstate.go +++ b/pkg/packet/shellstate.go @@ -154,19 +154,39 @@ func (sdiff *ShellStateDiff) GetHashVal(force bool) string { return sdiff.HashVal } -func (sdiff ShellStateDiff) Dump() { +func (sdiff ShellStateDiff) Dump(vars bool, aliases bool, funcs bool) { fmt.Printf("ShellStateDiff:\n") fmt.Printf(" version: %s\n", sdiff.Version) fmt.Printf(" base: %s\n", sdiff.BaseHash) - var mdiff statediff.MapDiffType - err := mdiff.Decode(sdiff.VarsDiff) - if err != nil { - fmt.Printf(" vars: error[%s]\n", err.Error()) - } else { - mdiff.Dump() - } - fmt.Printf(" aliases: %d, funcs: %d\n", len(sdiff.AliasesDiff), len(sdiff.FuncsDiff)) + fmt.Printf(" vars: %d, aliases: %d, funcs: %d\n", len(sdiff.VarsDiff), len(sdiff.AliasesDiff), len(sdiff.FuncsDiff)) if sdiff.Error != "" { fmt.Printf(" error: %s\n", sdiff.Error) } + if vars { + var mdiff statediff.MapDiffType + err := mdiff.Decode(sdiff.VarsDiff) + if err != nil { + fmt.Printf(" vars: error[%s]\n", err.Error()) + } else { + mdiff.Dump() + } + } + if aliases && len(sdiff.AliasesDiff) > 0 { + var ldiff statediff.LineDiffType + err := ldiff.Decode(sdiff.AliasesDiff) + if err != nil { + fmt.Printf(" aliases: error[%s]\n", err.Error()) + } else { + ldiff.Dump() + } + } + if funcs && len(sdiff.FuncsDiff) > 0 { + var ldiff statediff.LineDiffType + err := ldiff.Decode(sdiff.FuncsDiff) + if err != nil { + fmt.Printf(" funcs: error[%s]\n", err.Error()) + } else { + ldiff.Dump() + } + } } diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index e75cbfeb..6b6295f3 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -9,7 +9,9 @@ import ( "strings" "github.com/alessio/shellescape" + "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/scripthaus-dev/mshell/pkg/simpleexpand" "github.com/scripthaus-dev/mshell/pkg/statediff" "mvdan.cc/sh/v3/expand" "mvdan.cc/sh/v3/syntax" @@ -157,10 +159,6 @@ func (d *DeclareDeclType) Serialize() string { return fmt.Sprintf("%s|%s=%s\x00", d.Args, d.Name, d.Value) } -func (d *DeclareDeclType) EnvString() string { - return d.Name + "=" + d.Value -} - func (d *DeclareDeclType) DeclareStmt() string { var argsStr string if d.Args == "" { @@ -274,11 +272,12 @@ func EnvMapFromState(state *packet.ShellState) map[string]string { return nil } rtn := make(map[string]string) + ectx := simpleexpand.SimpleExpandContext{} vars := bytes.Split(state.ShellVars, []byte{0}) for _, varLine := range vars { decl := ParseDeclLine(string(varLine)) if decl != nil && decl.IsExport() { - rtn[decl.Name] = decl.Value + rtn[decl.Name], _ = simpleexpand.SimpleExpandPartialWord(ectx, decl.Value, false) } } return rtn @@ -289,11 +288,12 @@ func ShellVarMapFromState(state *packet.ShellState) map[string]string { return nil } rtn := make(map[string]string) + ectx := simpleexpand.SimpleExpandContext{} vars := bytes.Split(state.ShellVars, []byte{0}) for _, varLine := range vars { decl := ParseDeclLine(string(varLine)) if decl != nil { - rtn[decl.Name] = decl.Value + rtn[decl.Name], _ = simpleexpand.SimpleExpandPartialWord(ectx, decl.Value, false) } } return rtn @@ -438,6 +438,7 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { if strings.Index(rtn.Version, "bash") == -1 { return nil, fmt.Errorf("invalid shell state output, only bash is supported") } + rtn.Version = rtn.Version cwdStr := string(fields[1]) if strings.HasSuffix(cwdStr, "\r\n") { cwdStr = cwdStr[0 : len(cwdStr)-2] @@ -451,9 +452,35 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { } rtn.Aliases = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") + rtn.Funcs = removeFunc(rtn.Funcs, "_scripthaus_exittrap") + lines := strings.Split(rtn.Funcs, "\n") + for _, line := range lines { + base.Logf("func-line: [%s]\n", line) + } return rtn, nil } +func removeFunc(funcs string, toRemove string) string { + lines := strings.Split(funcs, "\n") + var newLines []string + removeLine := fmt.Sprintf("%s ()", toRemove) + doingRemove := false + for _, line := range lines { + if line == removeLine { + doingRemove = true + continue + } + if doingRemove { + if line == "}" { + doingRemove = false + } + continue + } + newLines = append(newLines, line) + } + return strings.Join(newLines, "\n") +} + func (d *DeclareDeclType) normalize() error { if d.DataType() == DeclTypeAssocArray { return d.normalizeAssocArrayDecl() @@ -565,9 +592,7 @@ func MakeShellStateDiff(oldState packet.ShellState, oldStateHash string, newStat if oldState.Cwd != newState.Cwd { rtn.Cwd = newState.Cwd } - if oldState.Error != newState.Error { - rtn.Error = newState.Error - } + rtn.Error = newState.Error oldVars := shellStateVarsToMap(oldState.ShellVars) newVars := shellStateVarsToMap(newState.ShellVars) rtn.VarsDiff = statediff.MakeMapDiff(oldVars, newVars) @@ -580,7 +605,10 @@ func ApplyShellStateDiff(oldState packet.ShellState, diff packet.ShellStateDiff) var rtnState packet.ShellState var err error rtnState.Version = oldState.Version - rtnState.Cwd = diff.Cwd + rtnState.Cwd = oldState.Cwd + if diff.Cwd != "" { + rtnState.Cwd = diff.Cwd + } rtnState.Error = diff.Error oldVars := shellStateVarsToMap(oldState.ShellVars) newVars, err := statediff.ApplyMapDiff(oldVars, diff.VarsDiff) diff --git a/pkg/simpleexpand/simpleexpand.go b/pkg/simpleexpand/simpleexpand.go index 40691319..d92d35c4 100644 --- a/pkg/simpleexpand/simpleexpand.go +++ b/pkg/simpleexpand/simpleexpand.go @@ -93,12 +93,18 @@ func expandLiteralPlus(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string, func expandSQANSILiteral(buf *bytes.Buffer, litVal string) { // no info specials + if strings.HasSuffix(litVal, "'") { + litVal = litVal[0 : len(litVal)-1] + } str, _, _ := expand.Format(nil, litVal, nil) buf.WriteString(str) } func expandSQLiteral(buf *bytes.Buffer, litVal string) { // no info specials + if strings.HasSuffix(litVal, "'") { + litVal = litVal[0 : len(litVal)-1] + } buf.WriteString(litVal) } @@ -125,6 +131,9 @@ func expandDQLiteral(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string) { lastDollar = false continue } + if ch == '"' { + break + } // similar to expandLiteral, but no globbing if ch == '`' { diff --git a/pkg/statediff/mapdiff.go b/pkg/statediff/mapdiff.go index 8e810f7f..47db9f4b 100644 --- a/pkg/statediff/mapdiff.go +++ b/pkg/statediff/mapdiff.go @@ -18,10 +18,10 @@ type MapDiffType struct { func (diff MapDiffType) Dump() { fmt.Printf("VAR-DIFF\n") for name, val := range diff.ToAdd { - fmt.Printf(" add: %s=%s\n", name, val) + fmt.Printf(" add[%s] %s\n", name, val) } for _, name := range diff.ToRemove { - fmt.Printf(" rem: %s\n", name) + fmt.Printf(" rem[%s]\n", name) } } From ad2cab595d6881df0ec429560fa5b2d14f60ca20 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 5 Dec 2022 15:38:44 -0800 Subject: [PATCH 117/149] kill server on I/O write error, and add a pinger to continually send ping packets to test connection --- pkg/server/server.go | 103 +++++++++++++++++++++++++++++-------------- 1 file changed, 71 insertions(+), 32 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 9ff06dff..daf5e5bf 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -24,13 +24,15 @@ import ( // 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 + 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 } func (m *MServer) Close() { @@ -255,34 +257,16 @@ func (m *MServer) packetSenderErrorHandler(sender *packet.PacketSender, pk packe msg.CK = cpk.GetCK() } sender.SendPacket(msg) + return + } else { + // I/O error: close the WriteErrorCh to signal that we are dead (cannot continue if we can't write output) + m.WriteErrorChOnce.Do(func() { + close(m.WriteErrorCh) + }) } - // otherwise ignore (we can't output anything for a I/O error) } -func RunServer() (int, error) { - debug := false - if len(os.Args) >= 3 && os.Args[2] == "--debug" { - debug = true - } - server := &MServer{ - Lock: &sync.Mutex{}, - ClientMap: make(map[base.CommandKey]*shexec.ClientProc), - StateMap: make(map[string]*packet.ShellState), - Debug: debug, - } - if debug { - packet.GlobalDebug = true - } - server.MainInput = packet.MakePacketParser(os.Stdin) - server.Sender = packet.MakePacketSender(os.Stdout, server.packetSenderErrorHandler) - defer server.Close() - var err error - initPacket, err := shexec.MakeServerInitPacket() - if err != nil { - return 1, err - } - server.setCurrentState(initPacket.State) - server.Sender.SendPacket(initPacket) +func (server *MServer) runReadLoop() { builder := packet.MakeRunPacketBuilder() for pk := range server.MainInput.MainCh { if server.Debug { @@ -307,5 +291,60 @@ func RunServer() (int, error) { server.Sender.SendMessageFmt("invalid packet '%s' sent to mshell server", packet.AsString(pk)) continue } +} + +func RunServer() (int, error) { + debug := false + if len(os.Args) >= 3 && os.Args[2] == "--debug" { + 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{}, + } + if debug { + packet.GlobalDebug = true + } + server.MainInput = packet.MakePacketParser(os.Stdin) + server.Sender = packet.MakePacketSender(os.Stdout, server.packetSenderErrorHandler) + defer server.Close() + var err error + initPacket, err := shexec.MakeServerInitPacket() + if err != nil { + return 1, err + } + server.setCurrentState(initPacket.State) + server.Sender.SendPacket(initPacket) + ticker := time.NewTicker(1 * time.Minute) + go func() { + for range ticker.C { + server.Sender.SendPacket(packet.MakePingPacket()) + } + }() + defer ticker.Stop() + readLoopDoneCh := make(chan bool) + + go func() { + defer close(readLoopDoneCh) + server.runReadLoop() + }() + + go func() { + time.Sleep(5 * time.Second) + respPk := packet.MakeResponsePacket("NA", make(chan bool)) + server.Sender.SendPacket(respPk) + }() + + select { + case <-readLoopDoneCh: + break + + case <-server.WriteErrorCh: + break + } return 0, nil } From 39e5e6c7297ff85783d6ce2ce2eb323320977af1 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 5 Dec 2022 15:45:26 -0800 Subject: [PATCH 118/149] remove test code --- pkg/server/server.go | 8 -------- 1 file changed, 8 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index daf5e5bf..a2f0890f 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -327,18 +327,10 @@ func RunServer() (int, error) { }() defer ticker.Stop() readLoopDoneCh := make(chan bool) - go func() { defer close(readLoopDoneCh) server.runReadLoop() }() - - go func() { - time.Sleep(5 * time.Second) - respPk := packet.MakeResponsePacket("NA", make(chan bool)) - server.Sender.SendPacket(respPk) - }() - select { case <-readLoopDoneCh: break From f010758b36f842398e998eaea7b8d86d2f9d5e0c Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 5 Dec 2022 22:26:13 -0800 Subject: [PATCH 119/149] mshell single writes ping packets to detect when the server has died. sends SIGHUP to children --- main-mshell.go | 18 +++++++++++++++ pkg/base/base.go | 1 + pkg/packet/packet.go | 18 ++++++++++++--- pkg/shexec/parser.go | 5 ---- pkg/shexec/shexec.go | 55 ++++++++++++++++++++++++++++++++++++++++---- 5 files changed, 85 insertions(+), 12 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 0ad82482..2d66c151 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -12,6 +12,7 @@ import ( "os" "strconv" "strings" + "time" "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" @@ -194,6 +195,16 @@ func handleSingle(fromServer bool) { cmd.DetachedWait(startPk) return } else { + shexec.IgnoreSigPipe() + ticker := time.NewTicker(1 * time.Minute) + go func() { + for range ticker.C { + // this will let the command detect when the server has gone away + // that will then trigger cmd.SendHup() to send SIGHUP to the exec'ed process + sender.SendPacket(packet.MakePingPacket()) + } + }() + defer ticker.Stop() cmd, err := shexec.RunCommandSimple(runPacket, sender, true) if err != nil { sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("error running command: %w", err)) @@ -202,6 +213,13 @@ func handleSingle(fromServer bool) { defer cmd.Close() startPacket := cmd.MakeCmdStartPacket(runPacket.ReqId) sender.SendPacket(startPacket) + go func() { + exitErr := sender.WaitForDone() + if exitErr != nil { + base.Logf("I/O error talking to server, sending SIGHUP to children\n") + cmd.SendHup() + } + }() cmd.RunRemoteIOAndWait(packetParser, sender) return } diff --git a/pkg/base/base.go b/pkg/base/base.go index 4b87a534..fe103bef 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -80,6 +80,7 @@ func InitDebugLog(prefix string) { return } DebugLogger = log.New(fd, prefix+" ", log.LstdFlags) + Logf("logger initialized\n") } func SetEnableDebugLog(enable bool) { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 22172355..eee2e712 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -722,7 +722,7 @@ func (e *SendError) Error() string { if e.IsMarshalError { return fmt.Sprintf("SendPacket marshal-error '%s' packet: %v", e.PacketType, e.Err) } else if e.IsWriteError { - return fmt.Sprintf("SendPacket write-error: %v", e.Err) + return fmt.Sprintf("SendPacket write-error packet[%s]: %v", e.PacketType, e.Err) } else { return e.Err.Error() } @@ -742,7 +742,7 @@ func SendPacket(w io.Writer, packet PacketType) error { outBuf.Write(jsonBytes) outBuf.WriteByte('\n') if GlobalDebug { - fmt.Printf("SEND> %s\n", AsString(packet)) + base.Logf("SEND> %s\n", AsString(packet)) } outBytes := outBuf.Bytes() sanitizeBytes(outBytes) @@ -763,6 +763,7 @@ type PacketSender struct { Done bool DoneCh chan bool ErrHandler func(*PacketSender, PacketType, error) + ExitErr error } func MakePacketSender(output io.Writer, errHandler func(*PacketSender, PacketType, error)) *PacketSender { @@ -784,6 +785,9 @@ func MakePacketSender(output io.Writer, errHandler func(*PacketSender, PacketTyp continue } // write errors are not recoverable + sender.Lock.Lock() + sender.ExitErr = err + sender.Lock.Unlock() return } } @@ -825,10 +829,18 @@ func (sender *PacketSender) Close() { close(sender.SendCh) } -func (sender *PacketSender) WaitForDone() { +// returns ExitErr if set +func (sender *PacketSender) WaitForDone() error { <-sender.DoneCh + sender.Lock.Lock() + defer sender.Lock.Unlock() + return sender.ExitErr } +// this is "advisory", as there is a race condition between the loop closing and setting Done. +// that's okay because that's an impossible race condition anyway (you could enqueue the packet +// and then the connection dies, or it dies half way, etc.). this just stops blindly adding +// packets forever when the loop is done. func (sender *PacketSender) checkStatus() error { sender.Lock.Lock() defer sender.Lock.Unlock() diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 6b6295f3..4b4ff300 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -9,7 +9,6 @@ import ( "strings" "github.com/alessio/shellescape" - "github.com/scripthaus-dev/mshell/pkg/base" "github.com/scripthaus-dev/mshell/pkg/packet" "github.com/scripthaus-dev/mshell/pkg/simpleexpand" "github.com/scripthaus-dev/mshell/pkg/statediff" @@ -453,10 +452,6 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { rtn.Aliases = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") rtn.Funcs = removeFunc(rtn.Funcs, "_scripthaus_exittrap") - lines := strings.Split(rtn.Funcs, "\n") - for _, line := range lines { - base.Logf("func-line: [%s]\n", line) - } return rtn, nil } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index af56f4b8..6c41412a 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -102,6 +102,7 @@ func MakeReturnStateBuf() *ReturnStateBuf { } type ShExecType struct { + Lock *sync.Mutex // only locks "Exited" field StartTs time.Time CK base.CommandKey FileNames *base.CommandFileNames @@ -114,6 +115,7 @@ type ShExecType struct { RunnerOutFd *os.File MsgSender *packet.PacketSender // where to send out-of-band messages back to calling proceess ReturnState *ReturnStateBuf + Exited bool // locked via Lock } type StdContext struct{} @@ -195,6 +197,7 @@ func (s ShExecUPR) UnknownPacket(pk packet.PacketType) { func MakeShExec(ck base.CommandKey, upr packet.UnknownPacketReporter) *ShExecType { return &ShExecType{ + Lock: &sync.Mutex{}, StartTs: time.Now(), CK: ck, Multiplexer: mpio.MakeMultiplexer(ck, upr), @@ -1000,6 +1003,25 @@ trap _scripthaus_exittrap EXIT return fmt.Sprintf(fmtStr, stateCmd) } +func (s *ShExecType) SendHup() { + base.Logf("sendhup start\n") + if s.Cmd == nil || s.Cmd.Process == nil || s.IsExited() { + return + } + pgroup := false + if s.Cmd.SysProcAttr != nil && (s.Cmd.SysProcAttr.Setsid || s.Cmd.SysProcAttr.Setpgid) { + pgroup = true + } + pid := s.Cmd.Process.Pid + if pgroup { + base.Logf("sendhup %d (pgroup)\n", -pid) + syscall.Kill(-pid, syscall.SIGHUP) + } else { + base.Logf("sendhup %d (normal)\n", pid) + syscall.Kill(pid, syscall.SIGHUP) + } +} + func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (rtnShExec *ShExecType, rtnErr error) { state := pk.State if state == nil { @@ -1172,7 +1194,7 @@ func (rs *ReturnStateBuf) Run() { // since we want mshell to persist even if the mshell --server is terminated func SetupSignalsForDetach() { sigCh := make(chan os.Signal, 1) - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP, syscall.SIGPIPE) go func() { for range sigCh { // do nothing @@ -1180,6 +1202,18 @@ func SetupSignalsForDetach() { }() } +// in detached run mode, we don't want mshell to die from signals +// since we want mshell to persist even if the mshell --server is terminated +func IgnoreSigPipe() { + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGPIPE) + go func() { + for sig := range sigCh { + base.Logf("ignoring signal %v\n", sig) + } + }() +} + func copyToCirFile(dest *cirfile.File, src io.Reader) error { buf := make([]byte, 64*1024) for { @@ -1308,9 +1342,23 @@ func GetExitCode(err error) int { } } +func (c *ShExecType) ProcWait() error { + exitErr := c.Cmd.Wait() + c.Lock.Lock() + c.Exited = true + c.Lock.Unlock() + return exitErr +} + +func (c *ShExecType) IsExited() bool { + c.Lock.Lock() + defer c.Lock.Unlock() + return c.Exited +} + func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { donePacket := packet.MakeCmdDonePacket(c.CK) - exitErr := c.Cmd.Wait() + exitErr := c.ProcWait() if c.ReturnState != nil { <-c.ReturnState.DoneCh state, _ := ParseShellStateOutput(c.ReturnState.Buf) // TODO what to do with error? @@ -1318,9 +1366,8 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { } endTs := time.Now() cmdDuration := endTs.Sub(c.StartTs) - exitCode := GetExitCode(exitErr) donePacket.Ts = endTs.UnixMilli() - donePacket.ExitCode = exitCode + donePacket.ExitCode = GetExitCode(exitErr) donePacket.DurationMs = int64(cmdDuration / time.Millisecond) if c.FileNames != nil { os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) From 7250bbb1ada46c6965683a903eded8f5b6440fe2 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 19 Dec 2022 17:38:24 -0800 Subject: [PATCH 120/149] rename to mshell_exittrap --- pkg/shexec/parser.go | 2 +- pkg/shexec/shexec.go | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 4b4ff300..299cecc8 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -451,7 +451,7 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { } rtn.Aliases = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") - rtn.Funcs = removeFunc(rtn.Funcs, "_scripthaus_exittrap") + rtn.Funcs = removeFunc(rtn.Funcs, "_mshell_exittrap") return rtn, nil } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 6c41412a..009f3e6c 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -995,10 +995,10 @@ func makeRcFileStr(pk *packet.RunPacketType) string { func makeExitTrap(fdNum int) string { stateCmd := GetShellStateRedirectCommandStr(fdNum) fmtStr := ` -_scripthaus_exittrap () { +_mshell_exittrap () { %s } -trap _scripthaus_exittrap EXIT +trap _mshell_exittrap EXIT ` return fmt.Sprintf(fmtStr, stateCmd) } From 91667e4dec9e1bcc72739696aab89ca6bd715cda Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 20 Dec 2022 21:58:24 -0800 Subject: [PATCH 121/149] process special input signals. also allow numeric signals. --- main-mshell.go | 3 ++- pkg/packet/packet.go | 2 +- pkg/shexec/shexec.go | 26 ++++++++++++++++++-------- 3 files changed, 21 insertions(+), 10 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 2d66c151..bb41760b 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -12,6 +12,7 @@ import ( "os" "strconv" "strings" + "syscall" "time" "github.com/scripthaus-dev/mshell/pkg/base" @@ -217,7 +218,7 @@ func handleSingle(fromServer bool) { exitErr := sender.WaitForDone() if exitErr != nil { base.Logf("I/O error talking to server, sending SIGHUP to children\n") - cmd.SendHup() + cmd.SendSignal(syscall.SIGHUP) } }() cmd.RunRemoteIOAndWait(packetParser, sender) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index eee2e712..6e6f104b 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -254,7 +254,7 @@ type WinSize struct { type SpecialInputPacketType struct { Type string `json:"type"` CK base.CommandKey `json:"ck"` - SigName string `json:"signame,omitempty"` // passed to unix.SignalNum (needs 'SIG' prefix, e.g. "SIGTERM") + SigName string `json:"signame,omitempty"` // passed to unix.SignalNum (needs 'SIG' prefix, e.g. "SIGTERM"), also accepts a number (e.g. "9") WinSize *WinSize `json:"winsize,omitempty"` } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 009f3e6c..bf19c92b 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -17,6 +17,7 @@ import ( "os/signal" "os/user" "runtime" + "strconv" "strings" "sync" "syscall" @@ -159,6 +160,7 @@ type ShExecUPR struct { } func (s *ShExecType) processSpecialInputPacket(pk *packet.SpecialInputPacketType) error { + base.Logf("processSpecialInputPacket: %#v\n", pk) if pk.WinSize != nil { if s.CmdPty == nil { return fmt.Errorf("cannot change winsize, cmd was not started with a pty") @@ -171,10 +173,17 @@ func (s *ShExecType) processSpecialInputPacket(pk *packet.SpecialInputPacketType s.Cmd.Process.Signal(syscall.SIGWINCH) } if pk.SigName != "" { - sigNum := unix.SignalNum(pk.SigName) - if sigNum == 0 { + var signal syscall.Signal + sigNumInt, err := strconv.Atoi(pk.SigName) + if err == nil { + signal = syscall.Signal(sigNumInt) + } else { + signal = unix.SignalNum(pk.SigName) + } + if signal == 0 { return fmt.Errorf("error signal %q not found, cannot send", pk.SigName) } + s.SendSignal(syscall.Signal(signal)) } return nil } @@ -1003,8 +1012,8 @@ trap _mshell_exittrap EXIT return fmt.Sprintf(fmtStr, stateCmd) } -func (s *ShExecType) SendHup() { - base.Logf("sendhup start\n") +func (s *ShExecType) SendSignal(sig syscall.Signal) { + base.Logf("signal start\n") if s.Cmd == nil || s.Cmd.Process == nil || s.IsExited() { return } @@ -1014,11 +1023,11 @@ func (s *ShExecType) SendHup() { } pid := s.Cmd.Process.Pid if pgroup { - base.Logf("sendhup %d (pgroup)\n", -pid) - syscall.Kill(-pid, syscall.SIGHUP) + base.Logf("send signal %s to %d (pgroup)\n", sig, -pid) + syscall.Kill(-pid, sig) } else { - base.Logf("sendhup %d (normal)\n", pid) - syscall.Kill(pid, syscall.SIGHUP) + base.Logf("send signal %s to %d (normal)\n", sig, pid) + syscall.Kill(pid, sig) } } @@ -1344,6 +1353,7 @@ func GetExitCode(err error) int { func (c *ShExecType) ProcWait() error { exitErr := c.Cmd.Wait() + base.Logf("procwait: %v\n", exitErr) c.Lock.Lock() c.Exited = true c.Lock.Unlock() From 7d887bc2d98d47568dfdcb812179eab72681940f Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 28 Dec 2022 23:08:33 -0800 Subject: [PATCH 122/149] update build output paths --- scripthaus.md | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/scripthaus.md b/scripthaus.md index bdcff57c..bc58b302 100644 --- a/scripthaus.md +++ b/scripthaus.md @@ -1,16 +1,16 @@ ```bash # @scripthaus command build -go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go +go build -ldflags="-s -w" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go ``` ```bash # @scripthaus command fullbuild go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go -GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-linux.amd64 main-mshell.go -GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-linux.arm64 main-mshell.go -GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-darwin.amd64 main-mshell.go -GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o /opt/mshell/bin/mshell-v0.2-darwin.arm64 main-mshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o bin/mshell-v0.2-linux.amd64 main-mshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o bin/mshell-v0.2-linux.arm64 main-mshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o bin/mshell-v0.2-darwin.arm64 main-mshell.go ``` From a4b8819948b5c659dca6841dbb92cb2dad65ed19 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 31 Jan 2023 12:14:41 -0800 Subject: [PATCH 123/149] increase rundata size limits --- pkg/shexec/shexec.go | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index bf19c92b..bfa3a541 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -45,6 +45,8 @@ const DefaultTermType = "xterm-256color" const DefaultMaxPtySize = 1024 * 1024 const MinMaxPtySize = 16 * 1024 const MaxMaxPtySize = 100 * 1024 * 1024 +const MaxRunDataSize = 1024 * 1024 +const MaxTotalRunDataSize = 10 * MaxRunDataSize const GetStateTimeout = 5 * time.Second @@ -385,12 +387,12 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { } totalRunData := 0 for _, rd := range pk.RunData { - if rd.DataLen > mpio.ReadBufSize { + if rd.DataLen > MaxRunDataSize { return fmt.Errorf("cannot detach command, constant rundata input too large fd=%d, len=%d, max=%d", rd.FdNum, rd.DataLen, mpio.ReadBufSize) } totalRunData += rd.DataLen } - if totalRunData > mpio.MaxTotalRunDataSize { + if totalRunData > MaxTotalRunDataSize { return fmt.Errorf("cannot detach command, constant rundata input too large len=%d, max=%d", totalRunData, mpio.MaxTotalRunDataSize) } } @@ -618,8 +620,8 @@ func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { } func AddRunData(pk *packet.RunPacketType, data string, dataType string) (int, error) { - if len(data) > mpio.ReadBufSize { - return 0, fmt.Errorf("%s too large, exceeds read buffer size", dataType) + if len(data) > MaxRunDataSize { + return 0, fmt.Errorf("%s too large, exceeds read buffer size size:%d", dataType, len(data)) } fdNum, err := NextFreeFdNum(pk) if err != nil { From 3240b128e1f8c749c98c3f5b78efd019786e0061 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 1 Feb 2023 00:44:31 -0800 Subject: [PATCH 124/149] fix nil ptr exception (because GetType() returns a value even when pk is nil) --- pkg/shexec/shexec.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index bfa3a541..60c6660f 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -728,6 +728,9 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, mshe case <-ctx.Done(): return ctx.Err() } + if pk == nil { + return fmt.Errorf("no response packet received from client") + } if pk.GetType() == packet.InitPacketStr && firstInit { firstInit = false initPacket := pk.(*packet.InitPacketType) From 573ca55c50e2d3294a2aae3e4cc2b834482c9bf2 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 6 Feb 2023 00:29:33 -0800 Subject: [PATCH 125/149] add stat call --- pkg/cirfile/cirfile.go | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go index 51c8cbe7..797b09bd 100644 --- a/pkg/cirfile/cirfile.go +++ b/pkg/cirfile/cirfile.go @@ -32,6 +32,14 @@ type File struct { FlockStatus int } +type Stat struct { + Location string + Version byte + MaxSize int64 + FileOffset int64 + DataSize int64 +} + func (f *File) flock(ctx context.Context, lockType int) error { err := syscall.Flock(int(f.OSFile.Fd()), lockType|syscall.LOCK_NB) if err == nil { @@ -98,6 +106,25 @@ func OpenCirFile(fileName string) (*File, error) { return rtn, nil } +func StatCirFile(ctx context.Context, fileName string) (*Stat, error) { + file, err := OpenCirFile(fileName) + if err != nil { + return nil, err + } + defer file.Close() + fileOffset, dataSize, err := file.GetStartOffsetAndSize(ctx) + if err != nil { + return nil, err + } + return &Stat{ + Location: fileName, + Version: file.Version, + MaxSize: file.MaxSize, + FileOffset: fileOffset, + DataSize: dataSize, + }, nil +} + // if the file already exists, it is an error. // there is a race condition if two goroutines try to create the same file between Stat() and Create(), so // they both might get no error, but only one file will be valid. if this is a concern, this call From c1b609541064e8d3b0aca00ce776c9b1da1168e5 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Feb 2023 14:50:58 -0800 Subject: [PATCH 126/149] add buildtime --- go.mod | 2 +- main-mshell.go | 2 ++ scripthaus.md | 14 ++++++++------ 3 files changed, 11 insertions(+), 7 deletions(-) diff --git a/go.mod b/go.mod index 4fa77051..a24bde24 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/scripthaus-dev/mshell -go 1.17 +go 1.18 require ( github.com/alessio/shellescape v1.4.1 diff --git a/main-mshell.go b/main-mshell.go index bb41760b..b6b8ab33 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -22,6 +22,8 @@ import ( "golang.org/x/sys/unix" ) +var BuildTime = "-" + // func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { // err := shexec.ValidateRunPacket(pk) // if err != nil { diff --git a/scripthaus.md b/scripthaus.md index bc58b302..39b26787 100644 --- a/scripthaus.md +++ b/scripthaus.md @@ -1,16 +1,18 @@ ```bash # @scripthaus command build -go build -ldflags="-s -w" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go +GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" +go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go ``` ```bash # @scripthaus command fullbuild -go build -ldflags="-s -w" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go -GOOS=linux GOARCH=amd64 go build -ldflags="-s -w" -o bin/mshell-v0.2-linux.amd64 main-mshell.go -GOOS=linux GOARCH=arm64 go build -ldflags="-s -w" -o bin/mshell-v0.2-linux.arm64 main-mshell.go -GOOS=darwin GOARCH=amd64 go build -ldflags="-s -w" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go -GOOS=darwin GOARCH=arm64 go build -ldflags="-s -w" -o bin/mshell-v0.2-darwin.arm64 main-mshell.go +GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" +go build -ldflags="$GO_LDFLAGS" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-linux.amd64 main-mshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-linux.arm64 main-mshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.arm64 main-mshell.go ``` From 0ad1d9236ac099552c41d81bf86255f81474765a Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 8 Mar 2023 09:56:38 -0800 Subject: [PATCH 127/149] commit unused code in working version --- .gitignore | 6 ++++++ pkg/cmdtail/cmdtail.go | 46 ++++++++++++++++++++++++++++++++++-------- 2 files changed, 44 insertions(+), 8 deletions(-) create mode 100644 .gitignore diff --git a/.gitignore b/.gitignore new file mode 100644 index 00000000..ca2468ae --- /dev/null +++ b/.gitignore @@ -0,0 +1,6 @@ +*~ +bin/ +*.out +*.log +.DS_Store + diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index ae0ece61..0526c42b 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -24,6 +24,15 @@ const MaxDataBytes = 4096 const FileTypePty = "ptyout" const FileTypeRun = "runout" +type Tailer struct { + Lock *sync.Mutex + WatchList map[base.CommandKey]CmdWatchEntry + Watcher *fsnotify.Watcher + Sender *packet.PacketSender + Gen FileNameGenerator + Sessions map[string]bool +} + type TailPos struct { ReqId string Running bool // an active tailer sending data @@ -43,6 +52,7 @@ type CmdWatchEntry struct { type FileNameGenerator interface { PtyOutFile(ck base.CommandKey) string RunOutFile(ck base.CommandKey) string + SessionDir(sessionId string) string } func (w CmdWatchEntry) getTailPos(reqId string) (TailPos, bool) { @@ -79,14 +89,6 @@ func (pos TailPos) IsCurrent(entry CmdWatchEntry) bool { return pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen } -type Tailer struct { - Lock *sync.Mutex - WatchList map[base.CommandKey]CmdWatchEntry - Watcher *fsnotify.Watcher - Sender *packet.PacketSender - Gen FileNameGenerator -} - func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos TailPos) { entry, found := t.WatchList[cmdKey] if !found { @@ -133,10 +135,38 @@ func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (Cm return entry, pos, true } +func (t *Tailer) addSessionWatcher(sessionId string) error { + t.Lock.Lock() + defer t.Lock.Unlock() + + if t.Sessions[sessionId] { + return + } + sdir := t.Gen.SessionDir(sessionId) + err := t.Watcher.Add(sdir) + if err != nil { + return err + } + t.Sessions[sessionId] = true + return nil +} + +func (t *Tailer) removeSessionWatcher(sessionId string) { + t.Lock.Lock() + defer t.Lock.Unlock() + + if !t.Sessions[sessionId] { + return + } + sdir := t.Gen.SessionDir(sessionId) + t.Watcher.Remove(sdir) +} + func MakeTailer(sender *packet.PacketSender, gen FileNameGenerator) (*Tailer, error) { rtn := &Tailer{ Lock: &sync.Mutex{}, WatchList: make(map[base.CommandKey]CmdWatchEntry), + Sessions: make(map[string]bool), Sender: sender, Gen: gen, } From 963bef842596bfc1b7322530bf993f6525a89c39 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 20 Mar 2023 19:21:23 -0700 Subject: [PATCH 128/149] add groupid, since it is now a screenid not a sessionid --- pkg/base/base.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/pkg/base/base.go b/pkg/base/base.go index fe103bef..6d9b82aa 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -87,7 +87,12 @@ func SetEnableDebugLog(enable bool) { DebugLogEnabled = enable } +// deprecated (use GetGroupId instead) func (ckey CommandKey) GetSessionId() string { + return ckey.GetGroupId() +} + +func (ckey CommandKey) GetGroupId() string { slashIdx := strings.Index(string(ckey), "/") if slashIdx == -1 { return "" From 01a99cba03fbcbffe3a7c38e60c827674b635552 Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 2 Apr 2023 23:00:55 -0700 Subject: [PATCH 129/149] minor updates/bugfixes --- main-mshell.go | 3 --- pkg/cirfile/cirfile.go | 22 ++++++++++++++++++++++ pkg/shexec/client.go | 2 +- 3 files changed, 23 insertions(+), 4 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index b6b8ab33..b9a5e86f 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -528,9 +528,6 @@ Examples: # run a script as root (via sudo), capture output mshell --sudo-with-passfile pw.txt --ssh ubuntu@somehost -- "python3 /dev/fd/3 > /dev/fd/4" 3< myscript.py 4> script-output.txt < script-input.txt - -mshell is licensed under the MPLv2 -Please see https://github.com/scripthaus-dev/mshell for extended usage modes, source code, bugs, and feature requests ` fmt.Printf("%s\n\n", strings.TrimSpace(usage)) } diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go index 797b09bd..59e44651 100644 --- a/pkg/cirfile/cirfile.go +++ b/pkg/cirfile/cirfile.go @@ -352,6 +352,28 @@ func (f *File) ReadAll(ctx context.Context) (int64, []byte, error) { return realOffset, buf[0:nr], err } +func (f *File) ReadAtWithMax(ctx context.Context, offset int64, maxSize int64) (int64, []byte, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, nil, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, nil, err + } + chunks := f.getFileChunks() + curSize := totalChunksSize(chunks) + var buf []byte + if maxSize > curSize { + buf = make([]byte, curSize) + } else { + buf = make([]byte, maxSize) + } + realOffset, nr, err := f.internalReadNext(buf, offset) + return realOffset, buf[0:nr], err +} + func (f *File) internalReadNext(buf []byte, offset int64) (int64, int, error) { if offset < f.FileOffset { offset = f.FileOffset diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index 9dd5d974..7037c240 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -75,7 +75,7 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.I initPk := pk.(*packet.InitPacketType) if initPk.NotFound { cproc.Close() - return nil, initPk, fmt.Errorf("mshell-%s command not found on local server", semver.MajorMinor(base.MShellVersion)) + return nil, initPk, fmt.Errorf("mshell client not found", semver.MajorMinor(base.MShellVersion)) } if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { cproc.Close() From 181010757247a3f47c69e3a5f7b8f3d68f36cd82 Mon Sep 17 00:00:00 2001 From: sawka Date: Sun, 2 Apr 2023 23:11:29 -0700 Subject: [PATCH 130/149] fix error message --- pkg/shexec/client.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index 7037c240..9b3b1947 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -75,7 +75,7 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.I initPk := pk.(*packet.InitPacketType) if initPk.NotFound { cproc.Close() - return nil, initPk, fmt.Errorf("mshell client not found", semver.MajorMinor(base.MShellVersion)) + return nil, initPk, fmt.Errorf("mshell client not found") } if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { cproc.Close() From fdae19610bfe2521908ae822e9c569a78b9d744c Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 4 Apr 2023 09:03:18 -0700 Subject: [PATCH 131/149] fix path --- scripthaus.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripthaus.md b/scripthaus.md index 39b26787..4a16ac8b 100644 --- a/scripthaus.md +++ b/scripthaus.md @@ -8,7 +8,7 @@ go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go ```bash # @scripthaus command fullbuild GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" -go build -ldflags="$GO_LDFLAGS" -o /Users/mike/.mshell/mshell-v0.2 main-mshell.go +go build -ldflags="$GO_LDFLAGS" -o ~/.mshell/mshell-v0.2 main-mshell.go GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-linux.amd64 main-mshell.go GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-linux.arm64 main-mshell.go GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go From a4a4d53eb0dfaf5917da52ada30bd75a6434c311 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 11 Apr 2023 23:52:58 -0700 Subject: [PATCH 132/149] grab git branch --- pkg/shexec/parser.go | 21 ++++++++++++++++++--- pkg/shexec/shexec.go | 18 +++++++++++++++--- 2 files changed, 33 insertions(+), 6 deletions(-) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 299cecc8..4e0dba22 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -398,7 +398,7 @@ func parseDeclareStmt(stmt *syntax.Stmt, src string) (*DeclareDeclType, error) { return rtn, nil } -func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { +func parseDeclareOutput(state *packet.ShellState, declareBytes []byte, pvarBytes []byte) error { declareStr := string(declareBytes) r := bytes.NewReader(declareBytes) parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) @@ -419,6 +419,21 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { declMap[decl.Name] = decl } } + pvars := bytes.Split(pvarBytes, []byte{0}) + for _, pvarBA := range pvars { + pvarStr := string(pvarBA) + pvarFields := strings.SplitN(pvarStr, " ", 2) + if len(pvarFields) != 2 { + continue + } + if pvarFields[0] == "" || pvarFields[1] == "" { + continue + } + decl := &DeclareDeclType{Args: "x"} + decl.Name = "PROMPTVAR_" + pvarFields[0] + decl.Value = shellescape.Quote(pvarFields[1]) + declMap[decl.Name] = decl + } state.ShellVars = SerializeDeclMap(declMap) // this writes out the decls in a canonical order if firstParseErr != nil { state.Error = firstParseErr.Error() @@ -429,7 +444,7 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte) error { func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { // 5 fields: version, cwd, env/vars, aliases, funcs fields := bytes.Split(outputBytes, []byte{0, 0}) - if len(fields) != 5 { + if len(fields) != 6 { return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) } rtn := &packet.ShellState{} @@ -445,7 +460,7 @@ func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { cwdStr = cwdStr[0 : len(cwdStr)-1] } rtn.Cwd = string(cwdStr) - err := parseDeclareOutput(rtn, fields[2]) + err := parseDeclareOutput(rtn, fields[2], fields[5]) if err != nil { return nil, err } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 60c6660f..3f1f7ee0 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -51,7 +51,15 @@ const MaxTotalRunDataSize = 10 * MaxRunDataSize const GetStateTimeout = 5 * time.Second const BaseBashOpts = `set +m; set +H; shopt -s extglob` -const GetShellStateCmd = `echo bash v${BASH_VERSINFO[0]}.${BASH_VERSINFO[1]}.${BASH_VERSINFO[2]}; printf "\x00\x00"; pwd; printf "\x00\x00"; declare -p $(compgen -A variable); printf "\x00\x00"; alias -p; printf "\x00\x00"; declare -f;` + +var GetShellStateCmds = []string{ + `echo bash v${BASH_VERSINFO[0]}.${BASH_VERSINFO[1]}.${BASH_VERSINFO[2]};`, + `pwd;`, + `declare -p $(compgen -A variable);`, + `alias -p;`, + `declare -f;`, + `printf "GITBRANCH %s\x00" "$(git rev-parse --abbrev-ref HEAD 2>/dev/null)"`, +} const ClientCommandFmt = ` PATH=$PATH:~/.mshell; @@ -161,6 +169,10 @@ type ShExecUPR struct { UPR packet.UnknownPacketReporter } +func GetShellStateCmd() string { + return strings.Join(GetShellStateCmds, ` printf "\x00\x00";`) +} + func (s *ShExecType) processSpecialInputPacket(pk *packet.SpecialInputPacketType) error { base.Logf("processSpecialInputPacket: %#v\n", pk) if pk.WinSize != nil { @@ -1500,12 +1512,12 @@ func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { } func GetShellStateRedirectCommandStr(outputFdNum int) string { - return fmt.Sprintf("cat <(%s) > /dev/fd/%d", GetShellStateCmd, outputFdNum) + return fmt.Sprintf("cat <(%s) > /dev/fd/%d", GetShellStateCmd(), outputFdNum) } func GetShellState() (*packet.ShellState, error) { ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) - cmdStr := BaseBashOpts + "; " + GetShellStateCmd + cmdStr := BaseBashOpts + "; " + GetShellStateCmd() ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", cmdStr) outputBytes, err := runSimpleCmdInPty(ecmd) if err != nil { From bab095d70142d93b3f65de094c27131059c91fdc Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 12 Apr 2023 21:45:45 -0700 Subject: [PATCH 133/149] add buildtime to initpacket --- main-mshell.go | 5 +++-- pkg/base/base.go | 5 +++++ pkg/packet/packet.go | 1 + pkg/shexec/shexec.go | 1 + 4 files changed, 10 insertions(+), 2 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index b9a5e86f..bcad09c0 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -22,7 +22,7 @@ import ( "golang.org/x/sys/unix" ) -var BuildTime = "-" +var BuildTime = "0" // func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { // err := shexec.ValidateRunPacket(pk) @@ -533,6 +533,7 @@ Examples: } func main() { + base.SetBuildTime(BuildTime) if len(os.Args) == 1 { handleUsage() return @@ -542,7 +543,7 @@ func main() { handleUsage() return } else if firstArg == "--version" { - fmt.Printf("mshell %s\n", base.MShellVersion) + fmt.Printf("mshell %s+%s\n", base.MShellVersion, base.BuildTime) return } else if firstArg == "--test-env" { state, err := shexec.GetShellState() diff --git a/pkg/base/base.go b/pkg/base/base.go index 6d9b82aa..bff94cf2 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -41,6 +41,7 @@ var sessionDirCache = make(map[string]string) var baseLock = &sync.Mutex{} var DebugLogEnabled = false var DebugLogger *log.Logger +var BuildTime string = "0" type CommandFileNames struct { PtyOutFile string @@ -50,6 +51,10 @@ type CommandFileNames struct { type CommandKey string +func SetBuildTime(build string) { + BuildTime = build +} + func MakeCommandKey(sessionId string, cmdId string) CommandKey { if sessionId == "" && cmdId == "" { return CommandKey("") diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 6e6f104b..abf05a51 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -461,6 +461,7 @@ type InitPacketType struct { Type string `json:"type"` RespId string `json:"respid,omitempty"` Version string `json:"version"` + BuildTime string `json:"buildtime,omitempty"` MShellHomeDir string `json:"mshellhomedir,omitempty"` HomeDir string `json:"homedir,omitempty"` State *ShellState `json:"state,omitempty"` diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 3f1f7ee0..8019e4e7 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -1405,6 +1405,7 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { func MakeInitPacket() *packet.InitPacketType { initPacket := packet.MakeInitPacket() initPacket.Version = base.MShellVersion + initPacket.BuildTime = base.BuildTime initPacket.HomeDir = base.GetHomeDir() initPacket.MShellHomeDir = base.GetMShellHomeDir() if user, _ := user.Current(); user != nil { From 15c09b78205bb62f16aa7e3cf5d0b91a55889018 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 12 Apr 2023 22:57:17 -0700 Subject: [PATCH 134/149] fix, allow empty values for pvar fields --- pkg/shexec/parser.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/shexec/parser.go b/pkg/shexec/parser.go index 4e0dba22..a4f7e0cb 100644 --- a/pkg/shexec/parser.go +++ b/pkg/shexec/parser.go @@ -426,7 +426,7 @@ func parseDeclareOutput(state *packet.ShellState, declareBytes []byte, pvarBytes if len(pvarFields) != 2 { continue } - if pvarFields[0] == "" || pvarFields[1] == "" { + if pvarFields[0] == "" { continue } decl := &DeclareDeclType{Args: "x"} From 5e212caf83d05116b58872d7ce32112e9719e543 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 17 Apr 2023 15:13:03 -0700 Subject: [PATCH 135/149] allow staticdata to set a higher buffer limit (since it is processed before mpio multiplexer takes over). raise max rundata size to 1M from 128k --- pkg/mpio/bufwriter.go | 12 ++++++++---- pkg/mpio/mpio.go | 35 ++++++++++++++++++++--------------- pkg/shexec/shexec.go | 14 +++++++------- 3 files changed, 35 insertions(+), 26 deletions(-) diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index 977a2493..f2f3cbbb 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -17,19 +17,23 @@ type FdWriter struct { M *Multiplexer FdNum int Buffer []byte + BufferLimit int Fd io.WriteCloser Eof bool Closed bool ShouldCloseFd bool + Desc string } -func MakeFdWriter(m *Multiplexer, fd io.WriteCloser, fdNum int, shouldCloseFd bool) *FdWriter { +func MakeFdWriter(m *Multiplexer, fd io.WriteCloser, fdNum int, shouldCloseFd bool, desc string) *FdWriter { fw := &FdWriter{ CVar: sync.NewCond(&sync.Mutex{}), Fd: fd, M: m, FdNum: fdNum, ShouldCloseFd: shouldCloseFd, + Desc: desc, + BufferLimit: WriteBufSize, } return fw } @@ -68,11 +72,11 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { if len(data) == 0 { return nil } - return fmt.Errorf("write to closed file eof[%v]", w.Eof) + return fmt.Errorf("write to closed file %q (fd:%d) eof[%v]", w.Desc, w.FdNum, w.Eof) } if len(data) > 0 { - if len(data)+len(w.Buffer) > WriteBufSize { - return fmt.Errorf("write exceeds buffer size bufsize=%d (max=%d)", len(data)+len(w.Buffer), WriteBufSize) + if len(data)+len(w.Buffer) > w.BufferLimit { + return fmt.Errorf("write exceeds buffer size %q (fd:%d) bufsize=%d (max=%d)", w.Desc, w.FdNum, len(data)+len(w.Buffer), w.BufferLimit) } w.Buffer = append(w.Buffer, data...) } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index f16a411c..e27d9dcd 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -95,27 +95,28 @@ func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { } // returns the *reader* to connect to process, writer is put in FdWriters -func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { +func (m *Multiplexer) MakeWriterPipe(fdNum int, desc string) (*os.File, error) { pr, pw, err := os.Pipe() if err != nil { return nil, err } m.Lock.Lock() defer m.Lock.Unlock() - m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum, true) + m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum, true, desc) m.CloseAfterStart = append(m.CloseAfterStart, pr) return pr, nil } // returns the *reader* to connect to process, writer is put in FdWriters -func (m *Multiplexer) MakeStaticWriterPipe(fdNum int, data []byte) (*os.File, error) { +func (m *Multiplexer) MakeStaticWriterPipe(fdNum int, data []byte, bufferLimit int, desc string) (*os.File, error) { pr, pw, err := os.Pipe() if err != nil { return nil, err } m.Lock.Lock() defer m.Lock.Unlock() - fdWriter := MakeFdWriter(m, pw, fdNum, true) + fdWriter := MakeFdWriter(m, pw, fdNum, true, desc) + fdWriter.BufferLimit = bufferLimit err = fdWriter.AddData(data, true) if err != nil { return nil, err @@ -131,10 +132,10 @@ func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose b m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose, isPty) } -func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd io.WriteCloser, shouldClose bool) { +func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd io.WriteCloser, shouldClose bool, desc string) { m.Lock.Lock() defer m.Lock.Unlock() - m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, shouldClose) + m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, shouldClose, desc) } func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { @@ -225,22 +226,18 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { return nil } -func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { - realData, err := base64.StdEncoding.DecodeString(dataPacket.Data64) - if err != nil { - return fmt.Errorf("decoding base64 data: %w", err) - } +func (m *Multiplexer) WriteDataToFd(fdNum int, data []byte, isEof bool) error { m.Lock.Lock() defer m.Lock.Unlock() - fw := m.FdWriters[dataPacket.FdNum] + fw := m.FdWriters[fdNum] if fw == nil { // add a closed FdWriter as a placeholder so we only send one error - fw := MakeFdWriter(m, nil, dataPacket.FdNum, false) + fw := MakeFdWriter(m, nil, fdNum, false, "invalid-fd") fw.Close() - m.FdWriters[dataPacket.FdNum] = fw + m.FdWriters[fdNum] = fw return fmt.Errorf("write to closed file (no fd)") } - err = fw.AddData(realData, dataPacket.Eof) + err := fw.AddData(data, isEof) if err != nil { fw.Close() return err @@ -248,6 +245,14 @@ func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error return nil } +func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { + realData, err := base64.StdEncoding.DecodeString(dataPacket.Data64) + if err != nil { + return fmt.Errorf("decoding base64 data: %w", err) + } + return m.WriteDataToFd(dataPacket.FdNum, realData, dataPacket.Eof) +} + func (m *Multiplexer) processAckPacket(ackPacket *packet.DataAckPacketType) { m.Lock.Lock() defer m.Lock.Unlock() diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 8019e4e7..3ac000f3 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -840,8 +840,8 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon if !HasDupStdin(runPacket.Fds) { cmd.Multiplexer.MakeRawFdReader(0, fdContext.GetReader(0), false, false) } - cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false) - cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), false) + cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false, "client") + cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), false, "client") for _, rfd := range runPacket.Fds { if rfd.Read && rfd.DupStdin { cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fdContext.GetReader(0), false, false) @@ -852,7 +852,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, false, false) } else if rfd.Write { fd := fdContext.GetWriter(rfd.FdNum) - cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true) + cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true, "client") } } err = ecmd.Start() @@ -1123,7 +1123,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro Setsid: true, Setctty: true, } - cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false) + cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false, "simple") cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false, true) nullFd, err := os.Open("/dev/null") if err != nil { @@ -1131,7 +1131,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro } cmd.Multiplexer.MakeRawFdReader(2, nullFd, true, false) } else { - cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) + cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0, "simple") if err != nil { return nil, err } @@ -1149,7 +1149,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro if runData.FdNum >= len(extraFiles) { extraFiles = extraFiles[:runData.FdNum+1] } - extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data) + extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data, MaxRunDataSize, "simple-rundata") if err != nil { return nil, err } @@ -1160,7 +1160,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro } if rfd.Read { // client file is open for reading, so we make a writer pipe - extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum) + extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum, "simple") if err != nil { return nil, err } From dbd76e2f40c19fcae14804a60f78c2223cb244f0 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 17 Apr 2023 17:25:04 -0700 Subject: [PATCH 136/149] go mod tidy --- go.mod | 8 ++------ go.sum | 8 ++++++-- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/go.mod b/go.mod index a24bde24..825a9b2d 100644 --- a/go.mod +++ b/go.mod @@ -7,11 +7,7 @@ require ( github.com/creack/pty v1.1.18 github.com/fsnotify/fsnotify v1.5.4 github.com/google/uuid v1.3.0 + golang.org/x/mod v0.5.1 golang.org/x/sys v0.0.0-20220412211240-33da011f77ad -) - -require ( - github.com/Masterminds/semver/v3 v3.1.1 // indirect - golang.org/x/mod v0.5.1 // indirect - mvdan.cc/sh/v3 v3.5.1 // indirect + mvdan.cc/sh/v3 v3.5.1 ) diff --git a/go.sum b/go.sum index c9f554e1..cfc0858c 100644 --- a/go.sum +++ b/go.sum @@ -1,16 +1,20 @@ -github.com/Masterminds/semver/v3 v3.1.1 h1:hLg3sBzpNErnxhQtUy/mmLR2I9foDujNK030IGemrRc= -github.com/Masterminds/semver/v3 v3.1.1/go.mod h1:VPu/7SZ7ePZ3QOrcuXROw5FAcLl4a0cBrbBpGY/8hQs= github.com/alessio/shellescape v1.4.1 h1:V7yhSDDn8LP4lc4jS8pFkt0zCnzVJlG5JXy9BVKJUX0= github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPpyFjDCtHLS30= github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= +github.com/frankban/quicktest v1.14.0 h1:+cqqvzZV87b4adx/5ayVOaYZ2CrvM4ejQvUdBzPPUss= github.com/fsnotify/fsnotify v1.5.4 h1:jRbGcIw6P2Meqdwuo0H1p6JVLbL5DHKAKlYndzMwVZI= github.com/fsnotify/fsnotify v1.5.4/go.mod h1:OVB6XrOHzAwXMpEM7uPOzcehqUV2UqJxmVXmkdnm1bU= +github.com/google/go-cmp v0.5.6 h1:BKbKCqvP6I+rmFHt06ZmyQtvB8xAkWdhFyr0ZUNZcxQ= github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/rogpeppe/go-internal v1.8.1 h1:geMPLpDpQOgVyCg5z5GoRwLHepNdb71NXb67XFkP+Eg= golang.org/x/mod v0.5.1 h1:OJxoQ/rynoF0dcCdI7cLPktw/hR2cueqYfjm43oqK38= golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad h1:ntjMns5wyP/fN65tdBD4g8J5w8n015+iIIs9rtjXkY0= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 h1:go1bK/D/BFZV2I8cIQd1NKEZ+0owSTG1fDTci4IqFcE= mvdan.cc/sh/v3 v3.5.1 h1:hmP3UOw4f+EYexsJjFxvU38+kn+V/s2CclXHanIBkmQ= mvdan.cc/sh/v3 v3.5.1/go.mod h1:1JcoyAKm1lZw/2bZje/iYKWicU/KMd0rsyJeKHnsK4E= From 386b5f7a905089b5850faad2dcee7df3ebcdbae3 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 8 May 2023 18:00:46 -0700 Subject: [PATCH 137/149] add openai packet --- pkg/packet/packet.go | 55 +++++++++++++++++++++++++++++++++++++------- 1 file changed, 47 insertions(+), 8 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index abf05a51..fbb81f29 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -53,6 +53,8 @@ const ( CompGenPacketStr = "compgen" // rpc ReInitPacketStr = "reinit" // rpc CmdFinalPacketStr = "cmdfinal" // command, pushed at the "end" of a command (fail-safe for no cmddone) + + OpenAIPacketStr = "openai" // other ) const PacketSenderQueueSize = 20 @@ -82,6 +84,7 @@ func init() { TypeStrToFactory[CompGenPacketStr] = reflect.TypeOf(CompGenPacketType{}) TypeStrToFactory[ReInitPacketStr] = reflect.TypeOf(ReInitPacketType{}) TypeStrToFactory[CmdFinalPacketStr] = reflect.TypeOf(CmdFinalPacketType{}) + TypeStrToFactory[OpenAIPacketStr] = reflect.TypeOf(OpenAIPacketType{}) var _ RpcPacketType = (*RunPacketType)(nil) var _ RpcPacketType = (*GetCmdPacketType)(nil) @@ -543,11 +546,11 @@ func MakeCmdDonePacket(ck base.CommandKey) *CmdDonePacketType { type CmdStartPacketType struct { Type string `json:"type"` - RespId string `json:"respid"` + RespId string `json:"respid,omitempty"` Ts int64 `json:"ts"` CK base.CommandKey `json:"ck"` - Pid int `json:"pid"` - MShellPid int `json:"mshellpid"` + Pid int `json:"pid,omitempty"` + MShellPid int `json:"mshellpid,omitempty"` } func (*CmdStartPacketType) GetType() string { @@ -614,6 +617,31 @@ func MakeRunPacket() *RunPacketType { return &RunPacketType{Type: RunPacketStr} } +type OpenAIUsageType struct { + PromptTokens int `json:"prompt_tokens,omitempty"` + CompletionTokens int `json:"completion_tokens,omitempty"` + TotalTokens int `json:"total_tokens,omitempty"` +} + +type OpenAIPacketType struct { + Type string `json:"type"` + Model string `json:"model,omitempty"` + Created int64 `json:"created,omitempty"` + FinishReason string `json:"finish_reason,omitempty"` + Usage *OpenAIUsageType `json:"usage,omitempty"` + Index int `json:"index,omitempty"` + Text string `json:"text,omitempty"` + Error string `json:"error,omitempty"` +} + +func (*OpenAIPacketType) GetType() string { + return OpenAIPacketStr +} + +func MakeOpenAIPacket() *OpenAIPacketType { + return &OpenAIPacketType{Type: OpenAIPacketStr} +} + type BarePacketType struct { Type string `json:"type"` } @@ -729,24 +757,35 @@ func (e *SendError) Error() string { } } -func SendPacket(w io.Writer, packet PacketType) error { +func MarshalPacket(packet PacketType) ([]byte, error) { if packet == nil { - return nil + return nil, fmt.Errorf("invalid nil packet") } jsonBytes, err := json.Marshal(packet) if err != nil { - return &SendError{IsMarshalError: true, PacketType: packet.GetType(), Err: err} + return nil, &SendError{IsMarshalError: true, PacketType: packet.GetType(), Err: err} } var outBuf bytes.Buffer outBuf.WriteByte('\n') outBuf.WriteString(fmt.Sprintf("##%d", len(jsonBytes))) outBuf.Write(jsonBytes) outBuf.WriteByte('\n') + outBytes := outBuf.Bytes() + sanitizeBytes(outBytes) + return outBytes, nil +} + +func SendPacket(w io.Writer, packet PacketType) error { + if packet == nil { + return nil + } + outBytes, err := MarshalPacket(packet) + if err != nil { + return err + } if GlobalDebug { base.Logf("SEND> %s\n", AsString(packet)) } - outBytes := outBuf.Bytes() - sanitizeBytes(outBytes) _, err = w.Write(outBytes) if err != nil { return &SendError{IsWriteError: true, PacketType: packet.GetType(), Err: err} From 94827683b59967b58624151ca5409146d0e23061 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 26 Jul 2023 13:00:07 -0700 Subject: [PATCH 138/149] big rename mshell to commandlinedev/apishell --- go.mod | 2 +- main-mshell.go | 8 ++++---- pkg/cmdtail/cmdtail.go | 4 ++-- pkg/mpio/bufreader.go | 2 +- pkg/mpio/mpio.go | 4 ++-- pkg/packet/packet.go | 2 +- pkg/packet/shellstate.go | 4 ++-- pkg/server/server.go | 6 +++--- pkg/shexec/client.go | 4 ++-- pkg/shexec/parser.go | 6 +++--- pkg/shexec/shexec.go | 8 ++++---- 11 files changed, 25 insertions(+), 25 deletions(-) diff --git a/go.mod b/go.mod index 825a9b2d..4efba117 100644 --- a/go.mod +++ b/go.mod @@ -1,4 +1,4 @@ -module github.com/scripthaus-dev/mshell +module github.com/commandlinedev/apishell go 1.18 diff --git a/main-mshell.go b/main-mshell.go index bcad09c0..98ed90f4 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -15,10 +15,10 @@ import ( "syscall" "time" - "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/packet" - "github.com/scripthaus-dev/mshell/pkg/server" - "github.com/scripthaus-dev/mshell/pkg/shexec" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/server" + "github.com/commandlinedev/apishell/pkg/shexec" "golang.org/x/sys/unix" ) diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 0526c42b..9822b2b3 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -16,8 +16,8 @@ import ( "time" "github.com/fsnotify/fsnotify" - "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" ) const MaxDataBytes = 4096 diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index 68af4cb5..00e77514 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -10,7 +10,7 @@ import ( "io" "sync" - "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/packet" ) type FdReader struct { diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index e27d9dcd..6dbb1d36 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -13,8 +13,8 @@ import ( "os" "sync" - "github.com/scripthaus-dev/mshell/pkg/base" - "github.com/scripthaus-dev/mshell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" ) const ReadBufSize = 128 * 1024 diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index fbb81f29..8241e69e 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -18,7 +18,7 @@ import ( "reflect" "sync" - "github.com/scripthaus-dev/mshell/pkg/base" + "github.com/commandlinedev/apishell/pkg/base" ) // single : run, >cmddata, >cmddone, data, <>dataack, Date: Thu, 3 Aug 2023 14:58:15 -0700 Subject: [PATCH 139/149] bump dep versions --- go.mod | 6 +++--- go.sum | 12 ++++++++++++ pkg/cirfile/cirfile.go | 10 +++++----- 3 files changed, 20 insertions(+), 8 deletions(-) diff --git a/go.mod b/go.mod index 4efba117..7fd6d4cd 100644 --- a/go.mod +++ b/go.mod @@ -5,9 +5,9 @@ go 1.18 require ( github.com/alessio/shellescape v1.4.1 github.com/creack/pty v1.1.18 - github.com/fsnotify/fsnotify v1.5.4 + github.com/fsnotify/fsnotify v1.6.0 github.com/google/uuid v1.3.0 golang.org/x/mod v0.5.1 - golang.org/x/sys v0.0.0-20220412211240-33da011f77ad - mvdan.cc/sh/v3 v3.5.1 + golang.org/x/sys v0.10.0 + mvdan.cc/sh/v3 v3.7.0 ) diff --git a/go.sum b/go.sum index cfc0858c..82889351 100644 --- a/go.sum +++ b/go.sum @@ -3,18 +3,30 @@ github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPp github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= github.com/frankban/quicktest v1.14.0 h1:+cqqvzZV87b4adx/5ayVOaYZ2CrvM4ejQvUdBzPPUss= +github.com/frankban/quicktest v1.14.5 h1:dfYrrRyLtiqT9GyKXgdh+k4inNeTvmGbuSgZ3lx3GhA= github.com/fsnotify/fsnotify v1.5.4 h1:jRbGcIw6P2Meqdwuo0H1p6JVLbL5DHKAKlYndzMwVZI= github.com/fsnotify/fsnotify v1.5.4/go.mod h1:OVB6XrOHzAwXMpEM7uPOzcehqUV2UqJxmVXmkdnm1bU= +github.com/fsnotify/fsnotify v1.6.0 h1:n+5WquG0fcWoWp6xPWfHdbskMCQaFnG6PfBrh1Ky4HY= +github.com/fsnotify/fsnotify v1.6.0/go.mod h1:sl3t1tCWJFWoRz9R8WJCbQihKKwmorjAbSClcnxKAGw= github.com/google/go-cmp v0.5.6 h1:BKbKCqvP6I+rmFHt06ZmyQtvB8xAkWdhFyr0ZUNZcxQ= +github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/rogpeppe/go-internal v1.8.1 h1:geMPLpDpQOgVyCg5z5GoRwLHepNdb71NXb67XFkP+Eg= +github.com/rogpeppe/go-internal v1.10.1-0.20230524175051-ec119421bb97 h1:3RPlVWzZ/PDqmVuf/FKHARG5EMid/tl7cv54Sw/QRVY= golang.org/x/mod v0.5.1 h1:OJxoQ/rynoF0dcCdI7cLPktw/hR2cueqYfjm43oqK38= golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad h1:ntjMns5wyP/fN65tdBD4g8J5w8n015+iIIs9rtjXkY0= golang.org/x/sys v0.0.0-20220412211240-33da011f77ad/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20220908164124-27713097b956/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.10.0 h1:SqMFp9UcQJZa+pmYuAKjd9xq1f0j5rLcDIk0mj4qAsA= +golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898 h1:/atklqdjdhuosWIl6AIbOeHJjicWYPqR9bpxqxYG2pA= golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 h1:go1bK/D/BFZV2I8cIQd1NKEZ+0owSTG1fDTci4IqFcE= mvdan.cc/sh/v3 v3.5.1 h1:hmP3UOw4f+EYexsJjFxvU38+kn+V/s2CclXHanIBkmQ= mvdan.cc/sh/v3 v3.5.1/go.mod h1:1JcoyAKm1lZw/2bZje/iYKWicU/KMd0rsyJeKHnsK4E= +mvdan.cc/sh/v3 v3.7.0 h1:lSTjdP/1xsddtaKfGg7Myu7DnlHItd3/M2tomOcNNBg= +mvdan.cc/sh/v3 v3.7.0/go.mod h1:K2gwkaesF/D7av7Kxl0HbF5kGOd2ArupNTX3X44+8l8= diff --git a/pkg/cirfile/cirfile.go b/pkg/cirfile/cirfile.go index 59e44651..fef9a764 100644 --- a/pkg/cirfile/cirfile.go +++ b/pkg/cirfile/cirfile.go @@ -10,9 +10,9 @@ import ( ) // CBUF[version] [maxsize] [fileoffset] [startpos] [endpos] -const HeaderFmt = "CBUF%02d %19d %19d %19d %19d\n" // 87 bytes -const HeaderLen = 256 // set to 256 for future expandability -const FullHeaderFmt = "%-255s\n" // 256 bytes (255 + newline) +const HeaderFmt1 = "CBUF%02d %19d %19d %19d %19d\n" // 87 bytes +const HeaderLen = 256 // set to 256 for future expandability +const FullHeaderFmt = "%-255s\n" // 256 bytes (255 + newline) const CurrentVersion = 1 const FilePosEmpty = -1 // sentinel, if startpos is set to -1, file is empty @@ -203,7 +203,7 @@ func (f *File) readMeta() error { return fmt.Errorf("error reading header: %w", err) } // currently only one version, so we don't need to have special logic here yet - _, err = fmt.Sscanf(string(buf), HeaderFmt, &f.Version, &f.MaxSize, &f.FileOffset, &f.StartPos, &f.EndPos) + _, err = fmt.Sscanf(string(buf), HeaderFmt1, &f.Version, &f.MaxSize, &f.FileOffset, &f.StartPos, &f.EndPos) if err != nil { return fmt.Errorf("sscanf error: %w", err) } @@ -240,7 +240,7 @@ func (f *File) writeMeta() error { if err != nil { return fmt.Errorf("cannot seek file: %w", err) } - metaStr := fmt.Sprintf(HeaderFmt, f.Version, f.MaxSize, f.FileOffset, f.StartPos, f.EndPos) + metaStr := fmt.Sprintf(HeaderFmt1, f.Version, f.MaxSize, f.FileOffset, f.StartPos, f.EndPos) fullMetaStr := fmt.Sprintf(FullHeaderFmt, metaStr) _, err = f.OSFile.WriteString(fullMetaStr) if err != nil { From 093e550d50193510cd3054d22ee2b110cd617e61 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 14 Aug 2023 12:23:33 -0700 Subject: [PATCH 140/149] add a debug flag to log the rc file for bash initialization --- pkg/base/base.go | 21 +++++++++++++++++++++ pkg/shexec/shexec.go | 10 +++++++++- 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index bff94cf2..91b22c95 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -30,6 +30,7 @@ const MShellPathVarName = "MSHELL_PATH" const MShellHomeVarName = "MSHELL_HOME" const MShellInstallBinVarName = "MSHELL_INSTALLBIN_PATH" const SSHCommandVarName = "SSH_COMMAND" +const MShellDebugVarName = "MSHELL_DEBUG" const SessionsDirBaseName = "sessions" const MShellVersion = "v0.2.0" const RemoteIdFile = "remoteid" @@ -37,6 +38,9 @@ const DefaultMShellInstallBinDir = "/opt/mshell/bin" const LogFileName = "mshell.log" const ForceDebugLog = false +const DebugFlag_LogRcFile = "logrc" +const LogRcFileName = "debug.rcfile" + var sessionDirCache = make(map[string]string) var baseLock = &sync.Mutex{} var DebugLogEnabled = false @@ -146,6 +150,23 @@ func (ckey CommandKey) Validate(typeStr string) error { return nil } +func HasDebugFlag(envMap map[string]string, flagName string) bool { + msDebug := envMap[MShellDebugVarName] + flags := strings.Split(msDebug, ",") + Logf("hasdebugflag[%s]: %s [%#v]\n", flagName, msDebug, flags) + for _, flag := range flags { + if strings.TrimSpace(flag) == flagName { + return true + } + } + return false +} + +func GetDebugRcFileName() string { + msHome := GetMShellHomeDir() + return path.Join(msHome, LogRcFileName) +} + func GetHomeDir() string { homeVar := os.Getenv(HomeVarName) if homeVar == "" { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 520c333b..7b38e0a7 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -24,11 +24,11 @@ import ( "time" "github.com/alessio/shellescape" - "github.com/creack/pty" "github.com/commandlinedev/apishell/pkg/base" "github.com/commandlinedev/apishell/pkg/cirfile" "github.com/commandlinedev/apishell/pkg/mpio" "github.com/commandlinedev/apishell/pkg/packet" + "github.com/creack/pty" "golang.org/x/mod/semver" "golang.org/x/sys/unix" ) @@ -1081,6 +1081,14 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro trapCmdStr := makeExitTrap(cmd.ReturnState.FdNum) rcFileStr += trapCmdStr } + shellVarMap := ShellVarMapFromState(state) + if base.HasDebugFlag(shellVarMap, base.DebugFlag_LogRcFile) { + debugRcFileName := base.GetDebugRcFileName() + err := os.WriteFile(debugRcFileName, []byte(rcFileStr), 0600) + if err != nil { + base.Logf("error writing %s: %v\n", debugRcFileName, err) + } + } rcFileFdNum, err := AddRunData(pk, rcFileStr, "rcfile") if err != nil { return nil, err From c29c4a9a2dda697f1fc8e99bef6bf261169e0dd1 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 31 Aug 2023 22:03:38 -0700 Subject: [PATCH 141/149] 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 --- main-mshell.go | 2 +- pkg/packet/packet.go | 214 ++++++++++++++++++++++--- pkg/packet/parser.go | 89 ++++++++--- pkg/server/server.go | 363 +++++++++++++++++++++++++++++++++++++++++-- pkg/shexec/client.go | 6 +- pkg/shexec/shexec.go | 8 +- 6 files changed, 612 insertions(+), 70 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 98ed90f4..31a80f21 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -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() diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 8241e69e..c35ec856 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -26,33 +26,42 @@ import ( // server : run, >cmddata, >cmddone, data, <>dataack, cd, >getcmd, >untailcmd, >input, error, <>message, <>ping, streamfile, writefile, filedata*, 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 diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index efb3414a..c6292b82 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -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, diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 7b38e0a7..63337fa8 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -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 { From bc488cf242dc3e6138a8fb1f3817d071b343fa23 Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 5 Sep 2023 21:21:34 -0700 Subject: [PATCH 142/149] fix PE-63, permissions for temp files --- pkg/server/server.go | 35 ++++++++++++++++++++++++++++++++--- 1 file changed, 32 insertions(+), 3 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index b047abe8..4a3b31d6 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -11,6 +11,7 @@ import ( "errors" "fmt" "io" + "io/fs" "os" "os/exec" "path/filepath" @@ -254,6 +255,23 @@ func (m *MServer) reinit(reqId string) { m.Sender.SendPacket(initPk) } +func makeTemp(path string, mode fs.FileMode) (*os.File, error) { + dirName := filepath.Dir(path) + baseName := filepath.Base(path) + baseTempName := baseName + ".tmp." + writeFd, err := os.CreateTemp(dirName, baseTempName) + if err != nil { + return nil, err + } + err = writeFd.Chmod(mode) + if err != nil { + writeFd.Close() + os.Remove(writeFd.Name()) + return nil, fmt.Errorf("error setting tempfile permissions: %w", err) + } + return writeFd, nil +} + func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContext) { defer wfc.setDone() if pk.Path == "" { @@ -262,10 +280,22 @@ func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContex m.Sender.SendPacket(resp) return } - finfo, err := os.Stat(pk.Path) + var finfo fs.FileInfo + var err error + if pk.UseTemp { + finfo, err = os.Lstat(pk.Path) + } else { + finfo, err = os.Stat(pk.Path) + } if err == nil && finfo.IsDir() { err = fmt.Errorf("invalid path, cannot write a directory") } + if err == nil && ((finfo.Mode() & fs.ModeSymlink) != 0) { + err = fmt.Errorf("writefile (with usetemp) does not support symlinks") + } + if err == nil && ((finfo.Mode() & (fs.ModeNamedPipe | fs.ModeSocket | fs.ModeDevice | fs.ModeSetuid | fs.ModeSetgid)) != 0) { + err = fmt.Errorf("writefile does not support special files (named pipes, sockets, devices, setuid, or setgid): mode=%v", finfo.Mode()) + } if err == nil { writePerm := (finfo.Mode().Perm() & 0o222) if writePerm == 0 { @@ -295,8 +325,7 @@ func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContex m.Sender.SendPacket(resp) return } - baseName := filepath.Base(pk.Path) - writeFd, err = os.CreateTemp(dirName, baseName+".tmp.") + writeFd, err = makeTemp(pk.Path, finfo.Mode().Perm()) if err != nil { resp := packet.MakeWriteFileReadyPacket(pk.ReqId) resp.Error = fmt.Sprintf("write-file could not open tempfile: %v", err) From 71ba4b5b46b5ba8cd88a982a8be82dcc7b04dc0f Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Sep 2023 21:45:15 -0700 Subject: [PATCH 143/149] update writefile code. changed the way usetemp works to make sure file permissions/owner/attributes are kept on original file --- pkg/packet/packet.go | 11 ++-- pkg/server/server.go | 139 +++++++++++++++++++++++++++++-------------- 2 files changed, 101 insertions(+), 49 deletions(-) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index c35ec856..a8c0d033 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -417,11 +417,12 @@ func MakeStreamFilePacket() *StreamFilePacketType { } type FileInfo struct { - Name string `json:"name"` - Size int64 `json:"size"` - ModTs int64 `json:"modts"` - IsDir bool `json:"isdir,omitempty"` - Perm int `json:"perm"` + Name string `json:"name"` + Size int64 `json:"size"` + ModTs int64 `json:"modts"` + IsDir bool `json:"isdir,omitempty"` + Perm int `json:"perm"` + NotFound bool `json:"notfound,omitempty"` // when NotFound is set, Perm will be set to permission for directory } type StreamFileResponseType struct { diff --git a/pkg/server/server.go b/pkg/server/server.go index 4a3b31d6..67461682 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -272,6 +272,58 @@ func makeTemp(path string, mode fs.FileMode) (*os.File, error) { return writeFd, nil } +func checkFileWritable(path string) error { + finfo, err := os.Stat(path) // ok to follow symlinks + if errors.Is(err, fs.ErrNotExist) { + dirName := filepath.Dir(path) + dirInfo, err := os.Stat(dirName) + if err != nil { + return fmt.Errorf("file does not exist, error trying to stat parent directory: %w", err) + } + if !dirInfo.IsDir() { + return fmt.Errorf("file does not exist, parent path [%s] is not a directory", dirName) + } + return nil + } else { + if err != nil { + return fmt.Errorf("cannot stat: %w", err) + } + if finfo.IsDir() { + return fmt.Errorf("invalid path, cannot write a directory") + } + if (finfo.Mode() & fs.ModeSymlink) != 0 { + return fmt.Errorf("writefile does not support symlinks") // note this shouldn't happen because we're using Stat (not Lstat) + } + if (finfo.Mode() & (fs.ModeNamedPipe | fs.ModeSocket | fs.ModeDevice)) != 0 { + return fmt.Errorf("writefile does not support special files (named pipes, sockets, devices): mode=%v", finfo.Mode()) + } + writePerm := (finfo.Mode().Perm() & 0o222) + if writePerm == 0 { + return fmt.Errorf("file is not writable, perms: %v", finfo.Mode().Perm()) + } + return nil + } +} + +func copyFile(dstName string, srcName string) error { + srcFd, err := os.Open(srcName) + if err != nil { + return err + } + defer srcFd.Close() + dstFd, err := os.OpenFile(dstName, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o777) // use 777 because OpenFile respects umask + if err != nil { + return err + } + // we don't defer dstFd.Close() so we can return an error if dstFd.Close() returns an error + _, err = io.Copy(dstFd, srcFd) + if err != nil { + dstFd.Close() + return err + } + return dstFd.Close() +} + func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContext) { defer wfc.setDone() if pk.Path == "" { @@ -280,55 +332,19 @@ func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContex m.Sender.SendPacket(resp) return } - var finfo fs.FileInfo - var err error - if pk.UseTemp { - finfo, err = os.Lstat(pk.Path) - } else { - finfo, err = os.Stat(pk.Path) - } - if err == nil && finfo.IsDir() { - err = fmt.Errorf("invalid path, cannot write a directory") - } - if err == nil && ((finfo.Mode() & fs.ModeSymlink) != 0) { - err = fmt.Errorf("writefile (with usetemp) does not support symlinks") - } - if err == nil && ((finfo.Mode() & (fs.ModeNamedPipe | fs.ModeSocket | fs.ModeDevice | fs.ModeSetuid | fs.ModeSetgid)) != 0) { - err = fmt.Errorf("writefile does not support special files (named pipes, sockets, devices, setuid, or setgid): mode=%v", finfo.Mode()) - } - if err == nil { - writePerm := (finfo.Mode().Perm() & 0o222) - if writePerm == 0 { - err = fmt.Errorf("file is not writable, perms: %v", finfo.Mode().Perm()) - } - } + err := checkFileWritable(pk.Path) 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()) - } - } + writeFd, err = os.CreateTemp("", "mshell.writefile.*") // "" means make this file in standard TempDir if err != nil { resp := packet.MakeWriteFileReadyPacket(pk.ReqId) - resp.Error = err.Error() - m.Sender.SendPacket(resp) - return - } - writeFd, err = makeTemp(pk.Path, finfo.Mode().Perm()) - if err != nil { - resp := packet.MakeWriteFileReadyPacket(pk.ReqId) - resp.Error = fmt.Sprintf("write-file could not open tempfile: %v", err) + resp.Error = fmt.Sprintf("cannot create temp file: %v", err) m.Sender.SendPacket(resp) return } @@ -388,12 +404,12 @@ func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContex 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()) + // copy file between writeFd.Name() and pk.Path + copyErr := copyFile(pk.Path, writeFd.Name()) + if err != nil { + doneErr = fmt.Errorf("error writing file: %v", copyErr) } + os.Remove(writeFd.Name()) } } donePk := packet.MakeWriteFileDonePacket(pk.ReqId) @@ -403,9 +419,44 @@ func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContex m.Sender.SendPacket(donePk) } +func (m *MServer) returnStreamFileNewFileResponse(pk *packet.StreamFilePacketType) { + // ok, file doesn't exist, so try to check the directory at least to see if we can write a file here + resp := packet.MakeStreamFileResponse(pk.ReqId) + defer func() { + if resp.Error == "" { + resp.Done = true + } + m.Sender.SendPacket(resp) + }() + dirName := filepath.Dir(pk.Path) + dirInfo, err := os.Stat(dirName) + if err != nil { + resp.Error = fmt.Sprintf("file does not exist, error trying to stat parent directory: %v", err) + return + } + if !dirInfo.IsDir() { + resp.Error = fmt.Sprintf("file does not exist, parent path [%s] is not a directory", dirName) + return + } + resp.Info = &packet.FileInfo{ + Name: pk.Path, + Size: 0, + ModTs: 0, + IsDir: false, + Perm: int(dirInfo.Mode().Perm()), + NotFound: true, + } + return +} + func (m *MServer) streamFile(pk *packet.StreamFilePacketType) { resp := packet.MakeStreamFileResponse(pk.ReqId) finfo, err := os.Stat(pk.Path) + if errors.Is(err, fs.ErrNotExist) { + // special return + m.returnStreamFileNewFileResponse(pk) + return + } if err != nil { resp.Error = fmt.Sprintf("cannot stat file %q: %v", pk.Path, err) m.Sender.SendPacket(resp) From c7d09b469241ff620f5b19f0709913d518b011d3 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Sep 2023 21:50:59 -0700 Subject: [PATCH 144/149] bump mshell version to v0.3 --- pkg/base/base.go | 2 +- scripthaus.md | 10 +++++----- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/pkg/base/base.go b/pkg/base/base.go index 91b22c95..71c1bc71 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -32,7 +32,7 @@ const MShellInstallBinVarName = "MSHELL_INSTALLBIN_PATH" const SSHCommandVarName = "SSH_COMMAND" const MShellDebugVarName = "MSHELL_DEBUG" const SessionsDirBaseName = "sessions" -const MShellVersion = "v0.2.0" +const MShellVersion = "v0.3.0" const RemoteIdFile = "remoteid" const DefaultMShellInstallBinDir = "/opt/mshell/bin" const LogFileName = "mshell.log" diff --git a/scripthaus.md b/scripthaus.md index 4a16ac8b..57539e2b 100644 --- a/scripthaus.md +++ b/scripthaus.md @@ -2,17 +2,17 @@ ```bash # @scripthaus command build GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" -go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go +go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-mshell.go ``` ```bash # @scripthaus command fullbuild GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" go build -ldflags="$GO_LDFLAGS" -o ~/.mshell/mshell-v0.2 main-mshell.go -GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-linux.amd64 main-mshell.go -GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-linux.arm64 main-mshell.go -GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.amd64 main-mshell.go -GOOS=darwin GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.2-darwin.arm64 main-mshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.amd64 main-mshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.arm64 main-mshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-mshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.arm64 main-mshell.go ``` From d97c1862c79667296e3c2dc7ff4b4e75829e7b0d Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Sep 2023 21:57:41 -0700 Subject: [PATCH 145/149] should be 666 not 777 (don't set executable flag by default) --- pkg/server/server.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 67461682..d8690fba 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -311,7 +311,7 @@ func copyFile(dstName string, srcName string) error { return err } defer srcFd.Close() - dstFd, err := os.OpenFile(dstName, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o777) // use 777 because OpenFile respects umask + dstFd, err := os.OpenFile(dstName, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o666) // use 666 because OpenFile respects umask if err != nil { return err } @@ -349,7 +349,7 @@ func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContex return } } else { - writeFd, err = os.OpenFile(pk.Path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o777) // use 777 because OpenFile respects umask + writeFd, err = os.OpenFile(pk.Path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o666) // use 666 because OpenFile respects umask if err != nil { resp := packet.MakeWriteFileReadyPacket(pk.ReqId) resp.Error = fmt.Sprintf("write-file could not open file: %v", err) From b8ef07e06440066e2ec2e1b73b1b0438a5b32f55 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 7 Sep 2023 15:12:50 -0700 Subject: [PATCH 146/149] updated go.sum --- go.sum | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/go.sum b/go.sum index 82889351..9897e263 100644 --- a/go.sum +++ b/go.sum @@ -2,31 +2,19 @@ github.com/alessio/shellescape v1.4.1 h1:V7yhSDDn8LP4lc4jS8pFkt0zCnzVJlG5JXy9BVK github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPpyFjDCtHLS30= github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= -github.com/frankban/quicktest v1.14.0 h1:+cqqvzZV87b4adx/5ayVOaYZ2CrvM4ejQvUdBzPPUss= github.com/frankban/quicktest v1.14.5 h1:dfYrrRyLtiqT9GyKXgdh+k4inNeTvmGbuSgZ3lx3GhA= -github.com/fsnotify/fsnotify v1.5.4 h1:jRbGcIw6P2Meqdwuo0H1p6JVLbL5DHKAKlYndzMwVZI= -github.com/fsnotify/fsnotify v1.5.4/go.mod h1:OVB6XrOHzAwXMpEM7uPOzcehqUV2UqJxmVXmkdnm1bU= github.com/fsnotify/fsnotify v1.6.0 h1:n+5WquG0fcWoWp6xPWfHdbskMCQaFnG6PfBrh1Ky4HY= github.com/fsnotify/fsnotify v1.6.0/go.mod h1:sl3t1tCWJFWoRz9R8WJCbQihKKwmorjAbSClcnxKAGw= -github.com/google/go-cmp v0.5.6 h1:BKbKCqvP6I+rmFHt06ZmyQtvB8xAkWdhFyr0ZUNZcxQ= github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= -github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= -github.com/rogpeppe/go-internal v1.8.1 h1:geMPLpDpQOgVyCg5z5GoRwLHepNdb71NXb67XFkP+Eg= github.com/rogpeppe/go-internal v1.10.1-0.20230524175051-ec119421bb97 h1:3RPlVWzZ/PDqmVuf/FKHARG5EMid/tl7cv54Sw/QRVY= golang.org/x/mod v0.5.1 h1:OJxoQ/rynoF0dcCdI7cLPktw/hR2cueqYfjm43oqK38= golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro= -golang.org/x/sys v0.0.0-20220412211240-33da011f77ad h1:ntjMns5wyP/fN65tdBD4g8J5w8n015+iIIs9rtjXkY0= -golang.org/x/sys v0.0.0-20220412211240-33da011f77ad/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220908164124-27713097b956/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0 h1:SqMFp9UcQJZa+pmYuAKjd9xq1f0j5rLcDIk0mj4qAsA= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898 h1:/atklqdjdhuosWIl6AIbOeHJjicWYPqR9bpxqxYG2pA= -golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 h1:go1bK/D/BFZV2I8cIQd1NKEZ+0owSTG1fDTci4IqFcE= -mvdan.cc/sh/v3 v3.5.1 h1:hmP3UOw4f+EYexsJjFxvU38+kn+V/s2CclXHanIBkmQ= -mvdan.cc/sh/v3 v3.5.1/go.mod h1:1JcoyAKm1lZw/2bZje/iYKWicU/KMd0rsyJeKHnsK4E= mvdan.cc/sh/v3 v3.7.0 h1:lSTjdP/1xsddtaKfGg7Myu7DnlHItd3/M2tomOcNNBg= mvdan.cc/sh/v3 v3.7.0/go.mod h1:K2gwkaesF/D7av7Kxl0HbF5kGOd2ArupNTX3X44+8l8= From adc4948d172c41b79feb4cd5ce00bbd02853bc26 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 5 Oct 2023 09:50:35 -0700 Subject: [PATCH 147/149] add shell to initpk --- pkg/packet/packet.go | 1 + pkg/shexec/shexec.go | 2 ++ 2 files changed, 3 insertions(+) diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index a8c0d033..d566ffb3 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -576,6 +576,7 @@ type InitPacketType struct { HostName string `json:"hostname,omitempty"` NotFound bool `json:"notfound,omitempty"` UName string `json:"uname,omitempty"` + Shell string `json:"shell,omitempty"` RemoteId string `json:"remoteid,omitempty"` } diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 63337fa8..a2e27487 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -47,6 +47,7 @@ const MinMaxPtySize = 16 * 1024 const MaxMaxPtySize = 100 * 1024 * 1024 const MaxRunDataSize = 1024 * 1024 const MaxTotalRunDataSize = 10 * MaxRunDataSize +const ShellVarName = "SHELL" const GetStateTimeout = 5 * time.Second @@ -1432,6 +1433,7 @@ func MakeServerInitPacket() (*packet.InitPacketType, error) { return nil, err } initPacket.State = shellState + initPacket.Shell = os.Getenv(ShellVarName) initPacket.RemoteId, err = base.GetRemoteId() if err != nil { return nil, err From afa7ef3d76eb07b2c26a2ac9f350b94f0b681e8b Mon Sep 17 00:00:00 2001 From: Red J Adaya Date: Tue, 10 Oct 2023 09:19:47 +0800 Subject: [PATCH 148/149] remove license (#2) --- LICENSE | 373 ----------------------------------------- main-mshell.go | 6 - pkg/base/base.go | 6 - pkg/base/optsiter.go | 6 - pkg/cmdtail/cmdtail.go | 6 - pkg/mpio/bufreader.go | 6 - pkg/mpio/bufwriter.go | 6 - pkg/mpio/mpio.go | 6 - pkg/packet/packet.go | 6 - pkg/packet/parser.go | 6 - pkg/server/server.go | 6 - pkg/shexec/shexec.go | 6 - 12 files changed, 439 deletions(-) delete mode 100644 LICENSE diff --git a/LICENSE b/LICENSE deleted file mode 100644 index ee6256cd..00000000 --- a/LICENSE +++ /dev/null @@ -1,373 +0,0 @@ -Mozilla Public License Version 2.0 -================================== - -1. Definitions --------------- - -1.1. "Contributor" - means each individual or legal entity that creates, contributes to - the creation of, or owns Covered Software. - -1.2. "Contributor Version" - means the combination of the Contributions of others (if any) used - by a Contributor and that particular Contributor's Contribution. - -1.3. "Contribution" - means Covered Software of a particular Contributor. - -1.4. "Covered Software" - means Source Code Form to which the initial Contributor has attached - the notice in Exhibit A, the Executable Form of such Source Code - Form, and Modifications of such Source Code Form, in each case - including portions thereof. - -1.5. "Incompatible With Secondary Licenses" - means - - (a) that the initial Contributor has attached the notice described - in Exhibit B to the Covered Software; or - - (b) that the Covered Software was made available under the terms of - version 1.1 or earlier of the License, but not also under the - terms of a Secondary License. - -1.6. "Executable Form" - means any form of the work other than Source Code Form. - -1.7. "Larger Work" - means a work that combines Covered Software with other material, in - a separate file or files, that is not Covered Software. - -1.8. "License" - means this document. - -1.9. "Licensable" - means having the right to grant, to the maximum extent possible, - whether at the time of the initial grant or subsequently, any and - all of the rights conveyed by this License. - -1.10. "Modifications" - means any of the following: - - (a) any file in Source Code Form that results from an addition to, - deletion from, or modification of the contents of Covered - Software; or - - (b) any new file in Source Code Form that contains any Covered - Software. - -1.11. "Patent Claims" of a Contributor - means any patent claim(s), including without limitation, method, - process, and apparatus claims, in any patent Licensable by such - Contributor that would be infringed, but for the grant of the - License, by the making, using, selling, offering for sale, having - made, import, or transfer of either its Contributions or its - Contributor Version. - -1.12. "Secondary License" - means either the GNU General Public License, Version 2.0, the GNU - Lesser General Public License, Version 2.1, the GNU Affero General - Public License, Version 3.0, or any later versions of those - licenses. - -1.13. "Source Code Form" - means the form of the work preferred for making modifications. - -1.14. "You" (or "Your") - means an individual or a legal entity exercising rights under this - License. For legal entities, "You" includes any entity that - controls, is controlled by, or is under common control with You. For - purposes of this definition, "control" means (a) the power, direct - or indirect, to cause the direction or management of such entity, - whether by contract or otherwise, or (b) ownership of more than - fifty percent (50%) of the outstanding shares or beneficial - ownership of such entity. - -2. License Grants and Conditions --------------------------------- - -2.1. Grants - -Each Contributor hereby grants You a world-wide, royalty-free, -non-exclusive license: - -(a) under intellectual property rights (other than patent or trademark) - Licensable by such Contributor to use, reproduce, make available, - modify, display, perform, distribute, and otherwise exploit its - Contributions, either on an unmodified basis, with Modifications, or - as part of a Larger Work; and - -(b) under Patent Claims of such Contributor to make, use, sell, offer - for sale, have made, import, and otherwise transfer either its - Contributions or its Contributor Version. - -2.2. Effective Date - -The licenses granted in Section 2.1 with respect to any Contribution -become effective for each Contribution on the date the Contributor first -distributes such Contribution. - -2.3. Limitations on Grant Scope - -The licenses granted in this Section 2 are the only rights granted under -this License. No additional rights or licenses will be implied from the -distribution or licensing of Covered Software under this License. -Notwithstanding Section 2.1(b) above, no patent license is granted by a -Contributor: - -(a) for any code that a Contributor has removed from Covered Software; - or - -(b) for infringements caused by: (i) Your and any other third party's - modifications of Covered Software, or (ii) the combination of its - Contributions with other software (except as part of its Contributor - Version); or - -(c) under Patent Claims infringed by Covered Software in the absence of - its Contributions. - -This License does not grant any rights in the trademarks, service marks, -or logos of any Contributor (except as may be necessary to comply with -the notice requirements in Section 3.4). - -2.4. Subsequent Licenses - -No Contributor makes additional grants as a result of Your choice to -distribute the Covered Software under a subsequent version of this -License (see Section 10.2) or under the terms of a Secondary License (if -permitted under the terms of Section 3.3). - -2.5. Representation - -Each Contributor represents that the Contributor believes its -Contributions are its original creation(s) or it has sufficient rights -to grant the rights to its Contributions conveyed by this License. - -2.6. Fair Use - -This License is not intended to limit any rights You have under -applicable copyright doctrines of fair use, fair dealing, or other -equivalents. - -2.7. Conditions - -Sections 3.1, 3.2, 3.3, and 3.4 are conditions of the licenses granted -in Section 2.1. - -3. Responsibilities -------------------- - -3.1. Distribution of Source Form - -All distribution of Covered Software in Source Code Form, including any -Modifications that You create or to which You contribute, must be under -the terms of this License. You must inform recipients that the Source -Code Form of the Covered Software is governed by the terms of this -License, and how they can obtain a copy of this License. You may not -attempt to alter or restrict the recipients' rights in the Source Code -Form. - -3.2. Distribution of Executable Form - -If You distribute Covered Software in Executable Form then: - -(a) such Covered Software must also be made available in Source Code - Form, as described in Section 3.1, and You must inform recipients of - the Executable Form how they can obtain a copy of such Source Code - Form by reasonable means in a timely manner, at a charge no more - than the cost of distribution to the recipient; and - -(b) You may distribute such Executable Form under the terms of this - License, or sublicense it under different terms, provided that the - license for the Executable Form does not attempt to limit or alter - the recipients' rights in the Source Code Form under this License. - -3.3. Distribution of a Larger Work - -You may create and distribute a Larger Work under terms of Your choice, -provided that You also comply with the requirements of this License for -the Covered Software. If the Larger Work is a combination of Covered -Software with a work governed by one or more Secondary Licenses, and the -Covered Software is not Incompatible With Secondary Licenses, this -License permits You to additionally distribute such Covered Software -under the terms of such Secondary License(s), so that the recipient of -the Larger Work may, at their option, further distribute the Covered -Software under the terms of either this License or such Secondary -License(s). - -3.4. Notices - -You may not remove or alter the substance of any license notices -(including copyright notices, patent notices, disclaimers of warranty, -or limitations of liability) contained within the Source Code Form of -the Covered Software, except that You may alter any license notices to -the extent required to remedy known factual inaccuracies. - -3.5. Application of Additional Terms - -You may choose to offer, and to charge a fee for, warranty, support, -indemnity or liability obligations to one or more recipients of Covered -Software. However, You may do so only on Your own behalf, and not on -behalf of any Contributor. You must make it absolutely clear that any -such warranty, support, indemnity, or liability obligation is offered by -You alone, and You hereby agree to indemnify every Contributor for any -liability incurred by such Contributor as a result of warranty, support, -indemnity or liability terms You offer. You may include additional -disclaimers of warranty and limitations of liability specific to any -jurisdiction. - -4. Inability to Comply Due to Statute or Regulation ---------------------------------------------------- - -If it is impossible for You to comply with any of the terms of this -License with respect to some or all of the Covered Software due to -statute, judicial order, or regulation then You must: (a) comply with -the terms of this License to the maximum extent possible; and (b) -describe the limitations and the code they affect. Such description must -be placed in a text file included with all distributions of the Covered -Software under this License. Except to the extent prohibited by statute -or regulation, such description must be sufficiently detailed for a -recipient of ordinary skill to be able to understand it. - -5. Termination --------------- - -5.1. The rights granted under this License will terminate automatically -if You fail to comply with any of its terms. However, if You become -compliant, then the rights granted under this License from a particular -Contributor are reinstated (a) provisionally, unless and until such -Contributor explicitly and finally terminates Your grants, and (b) on an -ongoing basis, if such Contributor fails to notify You of the -non-compliance by some reasonable means prior to 60 days after You have -come back into compliance. Moreover, Your grants from a particular -Contributor are reinstated on an ongoing basis if such Contributor -notifies You of the non-compliance by some reasonable means, this is the -first time You have received notice of non-compliance with this License -from such Contributor, and You become compliant prior to 30 days after -Your receipt of the notice. - -5.2. If You initiate litigation against any entity by asserting a patent -infringement claim (excluding declaratory judgment actions, -counter-claims, and cross-claims) alleging that a Contributor Version -directly or indirectly infringes any patent, then the rights granted to -You by any and all Contributors for the Covered Software under Section -2.1 of this License shall terminate. - -5.3. In the event of termination under Sections 5.1 or 5.2 above, all -end user license agreements (excluding distributors and resellers) which -have been validly granted by You or Your distributors under this License -prior to termination shall survive termination. - -************************************************************************ -* * -* 6. Disclaimer of Warranty * -* ------------------------- * -* * -* Covered Software is provided under this License on an "as is" * -* basis, without warranty of any kind, either expressed, implied, or * -* statutory, including, without limitation, warranties that the * -* Covered Software is free of defects, merchantable, fit for a * -* particular purpose or non-infringing. The entire risk as to the * -* quality and performance of the Covered Software is with You. * -* Should any Covered Software prove defective in any respect, You * -* (not any Contributor) assume the cost of any necessary servicing, * -* repair, or correction. This disclaimer of warranty constitutes an * -* essential part of this License. No use of any Covered Software is * -* authorized under this License except under this disclaimer. * -* * -************************************************************************ - -************************************************************************ -* * -* 7. Limitation of Liability * -* -------------------------- * -* * -* Under no circumstances and under no legal theory, whether tort * -* (including negligence), contract, or otherwise, shall any * -* Contributor, or anyone who distributes Covered Software as * -* permitted above, be liable to You for any direct, indirect, * -* special, incidental, or consequential damages of any character * -* including, without limitation, damages for lost profits, loss of * -* goodwill, work stoppage, computer failure or malfunction, or any * -* and all other commercial damages or losses, even if such party * -* shall have been informed of the possibility of such damages. This * -* limitation of liability shall not apply to liability for death or * -* personal injury resulting from such party's negligence to the * -* extent applicable law prohibits such limitation. Some * -* jurisdictions do not allow the exclusion or limitation of * -* incidental or consequential damages, so this exclusion and * -* limitation may not apply to You. * -* * -************************************************************************ - -8. Litigation -------------- - -Any litigation relating to this License may be brought only in the -courts of a jurisdiction where the defendant maintains its principal -place of business and such litigation shall be governed by laws of that -jurisdiction, without reference to its conflict-of-law provisions. -Nothing in this Section shall prevent a party's ability to bring -cross-claims or counter-claims. - -9. Miscellaneous ----------------- - -This License represents the complete agreement concerning the subject -matter hereof. If any provision of this License is held to be -unenforceable, such provision shall be reformed only to the extent -necessary to make it enforceable. Any law or regulation which provides -that the language of a contract shall be construed against the drafter -shall not be used to construe this License against a Contributor. - -10. Versions of the License ---------------------------- - -10.1. New Versions - -Mozilla Foundation is the license steward. Except as provided in Section -10.3, no one other than the license steward has the right to modify or -publish new versions of this License. Each version will be given a -distinguishing version number. - -10.2. Effect of New Versions - -You may distribute the Covered Software under the terms of the version -of the License under which You originally received the Covered Software, -or under the terms of any subsequent version published by the license -steward. - -10.3. Modified Versions - -If you create software not governed by this License, and you want to -create a new license for such software, you may create and use a -modified version of this License if you rename the license and remove -any references to the name of the license steward (except to note that -such modified license differs from this License). - -10.4. Distributing Source Code Form that is Incompatible With Secondary -Licenses - -If You choose to distribute Source Code Form that is Incompatible With -Secondary Licenses under the terms of this version of the License, the -notice described in Exhibit B of this License must be attached. - -Exhibit A - Source Code Form License Notice -------------------------------------------- - - This Source Code Form is subject to the terms of the Mozilla Public - License, v. 2.0. If a copy of the MPL was not distributed with this - file, You can obtain one at https://mozilla.org/MPL/2.0/. - -If it is not possible or desirable to put the notice in a particular -file, then You may include the notice in a location (such as a LICENSE -file in a relevant directory) where a recipient would be likely to look -for such a notice. - -You may add additional accurate notices of copyright ownership. - -Exhibit B - "Incompatible With Secondary Licenses" Notice ---------------------------------------------------------- - - This Source Code Form is "Incompatible With Secondary Licenses", as - defined by the Mozilla Public License, v. 2.0. diff --git a/main-mshell.go b/main-mshell.go index 31a80f21..36b8ce1c 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package main import ( diff --git a/pkg/base/base.go b/pkg/base/base.go index 71c1bc71..d9204f4f 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package base import ( diff --git a/pkg/base/optsiter.go b/pkg/base/optsiter.go index aee544e3..f607c670 100644 --- a/pkg/base/optsiter.go +++ b/pkg/base/optsiter.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package base import "strings" diff --git a/pkg/cmdtail/cmdtail.go b/pkg/cmdtail/cmdtail.go index 9822b2b3..2578c3b5 100644 --- a/pkg/cmdtail/cmdtail.go +++ b/pkg/cmdtail/cmdtail.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package cmdtail import ( diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index 00e77514..ab4f5709 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package mpio import ( diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index f2f3cbbb..86d2efa3 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package mpio import ( diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index 6dbb1d36..fe8fed09 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package mpio import ( diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index d566ffb3..30c36f77 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package packet import ( diff --git a/pkg/packet/parser.go b/pkg/packet/parser.go index 3eacca6f..f84be258 100644 --- a/pkg/packet/parser.go +++ b/pkg/packet/parser.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package packet import ( diff --git a/pkg/server/server.go b/pkg/server/server.go index d8690fba..3ccbb318 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package server import ( diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index a2e27487..d5dc865f 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -1,9 +1,3 @@ -// Copyright 2022 Dashborg Inc -// -// This Source Code Form is subject to the terms of the Mozilla Public -// License, v. 2.0. If a copy of the MPL was not distributed with this -// file, You can obtain one at https://mozilla.org/MPL/2.0/. - package shexec import ( From b864e1fb1d7b74d8839925c55ee1d6c55955c809 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 16 Oct 2023 13:04:18 -0700 Subject: [PATCH 149/149] move apishell code to subdirectory to prepare for move to main repo --- NOTICE.md | 5 ----- go.mod => waveshell/go.mod | 0 go.sum => waveshell/go.sum | 0 main-mshell.go => waveshell/main-waveshell.go | 0 {pkg => waveshell/pkg}/base/base.go | 0 {pkg => waveshell/pkg}/base/optsiter.go | 0 {pkg => waveshell/pkg}/binpack/binpack.go | 0 {pkg => waveshell/pkg}/cirfile/cirfile.go | 0 {pkg => waveshell/pkg}/cirfile/cirfile_test.go | 0 {pkg => waveshell/pkg}/cmdtail/cmdtail.go | 0 {pkg => waveshell/pkg}/mpio/bufreader.go | 0 {pkg => waveshell/pkg}/mpio/bufwriter.go | 0 {pkg => waveshell/pkg}/mpio/mpio.go | 0 {pkg => waveshell/pkg}/packet/combined.go | 0 {pkg => waveshell/pkg}/packet/packet.go | 0 {pkg => waveshell/pkg}/packet/parser.go | 0 {pkg => waveshell/pkg}/packet/shellstate.go | 0 {pkg => waveshell/pkg}/server/server.go | 0 {pkg => waveshell/pkg}/shexec/client.go | 0 {pkg => waveshell/pkg}/shexec/parser.go | 0 {pkg => waveshell/pkg}/shexec/shexec.go | 0 {pkg => waveshell/pkg}/simpleexpand/simpleexpand.go | 0 {pkg => waveshell/pkg}/statediff/linediff.go | 0 {pkg => waveshell/pkg}/statediff/mapdiff.go | 0 {pkg => waveshell/pkg}/statediff/statediff_test.go | 0 scripthaus.md => waveshell/scripthaus.md | 12 ++++++------ 26 files changed, 6 insertions(+), 11 deletions(-) delete mode 100644 NOTICE.md rename go.mod => waveshell/go.mod (100%) rename go.sum => waveshell/go.sum (100%) rename main-mshell.go => waveshell/main-waveshell.go (100%) rename {pkg => waveshell/pkg}/base/base.go (100%) rename {pkg => waveshell/pkg}/base/optsiter.go (100%) rename {pkg => waveshell/pkg}/binpack/binpack.go (100%) rename {pkg => waveshell/pkg}/cirfile/cirfile.go (100%) rename {pkg => waveshell/pkg}/cirfile/cirfile_test.go (100%) rename {pkg => waveshell/pkg}/cmdtail/cmdtail.go (100%) rename {pkg => waveshell/pkg}/mpio/bufreader.go (100%) rename {pkg => waveshell/pkg}/mpio/bufwriter.go (100%) rename {pkg => waveshell/pkg}/mpio/mpio.go (100%) rename {pkg => waveshell/pkg}/packet/combined.go (100%) rename {pkg => waveshell/pkg}/packet/packet.go (100%) rename {pkg => waveshell/pkg}/packet/parser.go (100%) rename {pkg => waveshell/pkg}/packet/shellstate.go (100%) rename {pkg => waveshell/pkg}/server/server.go (100%) rename {pkg => waveshell/pkg}/shexec/client.go (100%) rename {pkg => waveshell/pkg}/shexec/parser.go (100%) rename {pkg => waveshell/pkg}/shexec/shexec.go (100%) rename {pkg => waveshell/pkg}/simpleexpand/simpleexpand.go (100%) rename {pkg => waveshell/pkg}/statediff/linediff.go (100%) rename {pkg => waveshell/pkg}/statediff/mapdiff.go (100%) rename {pkg => waveshell/pkg}/statediff/statediff_test.go (100%) rename scripthaus.md => waveshell/scripthaus.md (66%) diff --git a/NOTICE.md b/NOTICE.md deleted file mode 100644 index 5a9ef9f4..00000000 --- a/NOTICE.md +++ /dev/null @@ -1,5 +0,0 @@ -Copyright (c) 2021-2022 Dashborg Inc - -This Source Code Form is subject to the terms of the Mozilla Public -License, v. 2.0. If a copy of the MPL was not distributed with this -file, You can obtain one at https://mozilla.org/MPL/2.0/. diff --git a/go.mod b/waveshell/go.mod similarity index 100% rename from go.mod rename to waveshell/go.mod diff --git a/go.sum b/waveshell/go.sum similarity index 100% rename from go.sum rename to waveshell/go.sum diff --git a/main-mshell.go b/waveshell/main-waveshell.go similarity index 100% rename from main-mshell.go rename to waveshell/main-waveshell.go diff --git a/pkg/base/base.go b/waveshell/pkg/base/base.go similarity index 100% rename from pkg/base/base.go rename to waveshell/pkg/base/base.go diff --git a/pkg/base/optsiter.go b/waveshell/pkg/base/optsiter.go similarity index 100% rename from pkg/base/optsiter.go rename to waveshell/pkg/base/optsiter.go diff --git a/pkg/binpack/binpack.go b/waveshell/pkg/binpack/binpack.go similarity index 100% rename from pkg/binpack/binpack.go rename to waveshell/pkg/binpack/binpack.go diff --git a/pkg/cirfile/cirfile.go b/waveshell/pkg/cirfile/cirfile.go similarity index 100% rename from pkg/cirfile/cirfile.go rename to waveshell/pkg/cirfile/cirfile.go diff --git a/pkg/cirfile/cirfile_test.go b/waveshell/pkg/cirfile/cirfile_test.go similarity index 100% rename from pkg/cirfile/cirfile_test.go rename to waveshell/pkg/cirfile/cirfile_test.go diff --git a/pkg/cmdtail/cmdtail.go b/waveshell/pkg/cmdtail/cmdtail.go similarity index 100% rename from pkg/cmdtail/cmdtail.go rename to waveshell/pkg/cmdtail/cmdtail.go diff --git a/pkg/mpio/bufreader.go b/waveshell/pkg/mpio/bufreader.go similarity index 100% rename from pkg/mpio/bufreader.go rename to waveshell/pkg/mpio/bufreader.go diff --git a/pkg/mpio/bufwriter.go b/waveshell/pkg/mpio/bufwriter.go similarity index 100% rename from pkg/mpio/bufwriter.go rename to waveshell/pkg/mpio/bufwriter.go diff --git a/pkg/mpio/mpio.go b/waveshell/pkg/mpio/mpio.go similarity index 100% rename from pkg/mpio/mpio.go rename to waveshell/pkg/mpio/mpio.go diff --git a/pkg/packet/combined.go b/waveshell/pkg/packet/combined.go similarity index 100% rename from pkg/packet/combined.go rename to waveshell/pkg/packet/combined.go diff --git a/pkg/packet/packet.go b/waveshell/pkg/packet/packet.go similarity index 100% rename from pkg/packet/packet.go rename to waveshell/pkg/packet/packet.go diff --git a/pkg/packet/parser.go b/waveshell/pkg/packet/parser.go similarity index 100% rename from pkg/packet/parser.go rename to waveshell/pkg/packet/parser.go diff --git a/pkg/packet/shellstate.go b/waveshell/pkg/packet/shellstate.go similarity index 100% rename from pkg/packet/shellstate.go rename to waveshell/pkg/packet/shellstate.go diff --git a/pkg/server/server.go b/waveshell/pkg/server/server.go similarity index 100% rename from pkg/server/server.go rename to waveshell/pkg/server/server.go diff --git a/pkg/shexec/client.go b/waveshell/pkg/shexec/client.go similarity index 100% rename from pkg/shexec/client.go rename to waveshell/pkg/shexec/client.go diff --git a/pkg/shexec/parser.go b/waveshell/pkg/shexec/parser.go similarity index 100% rename from pkg/shexec/parser.go rename to waveshell/pkg/shexec/parser.go diff --git a/pkg/shexec/shexec.go b/waveshell/pkg/shexec/shexec.go similarity index 100% rename from pkg/shexec/shexec.go rename to waveshell/pkg/shexec/shexec.go diff --git a/pkg/simpleexpand/simpleexpand.go b/waveshell/pkg/simpleexpand/simpleexpand.go similarity index 100% rename from pkg/simpleexpand/simpleexpand.go rename to waveshell/pkg/simpleexpand/simpleexpand.go diff --git a/pkg/statediff/linediff.go b/waveshell/pkg/statediff/linediff.go similarity index 100% rename from pkg/statediff/linediff.go rename to waveshell/pkg/statediff/linediff.go diff --git a/pkg/statediff/mapdiff.go b/waveshell/pkg/statediff/mapdiff.go similarity index 100% rename from pkg/statediff/mapdiff.go rename to waveshell/pkg/statediff/mapdiff.go diff --git a/pkg/statediff/statediff_test.go b/waveshell/pkg/statediff/statediff_test.go similarity index 100% rename from pkg/statediff/statediff_test.go rename to waveshell/pkg/statediff/statediff_test.go diff --git a/scripthaus.md b/waveshell/scripthaus.md similarity index 66% rename from scripthaus.md rename to waveshell/scripthaus.md index 57539e2b..d6a26698 100644 --- a/scripthaus.md +++ b/waveshell/scripthaus.md @@ -2,17 +2,17 @@ ```bash # @scripthaus command build GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" -go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-mshell.go +go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-waveshell.go ``` ```bash # @scripthaus command fullbuild GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" -go build -ldflags="$GO_LDFLAGS" -o ~/.mshell/mshell-v0.2 main-mshell.go -GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.amd64 main-mshell.go -GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.arm64 main-mshell.go -GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-mshell.go -GOOS=darwin GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.arm64 main-mshell.go +go build -ldflags="$GO_LDFLAGS" -o ~/.mshell/mshell-v0.2 main-waveshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.amd64 main-waveshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.arm64 main-waveshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-waveshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.arm64 main-waveshell.go ```