From bdd8381b0166904eb4788e7ab56e0dd7cfd46614 Mon Sep 17 00:00:00 2001 From: sawka Date: Mon, 28 Nov 2022 18:05:54 -0800 Subject: [PATCH] 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) } }