mirror of
https://github.com/wavetermdev/backup.git
synced 2026-08-05 13:57:07 -07:00
merge waveshell into waveterm repo
This commit is contained in:
+5
-1
@@ -3,13 +3,17 @@ dist-dev/
|
||||
node_modules/
|
||||
*~
|
||||
*.log
|
||||
*.out
|
||||
out/
|
||||
.DS_Store
|
||||
bin/
|
||||
waveshell/bin/
|
||||
wavesrv/bin/
|
||||
dev-bin
|
||||
local-server-bin
|
||||
*.pw
|
||||
build/
|
||||
*.dmg
|
||||
webshare/dist/
|
||||
webshare/dist-dev/
|
||||
webshare/dist-dev/
|
||||
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
module github.com/commandlinedev/apishell
|
||||
|
||||
go 1.18
|
||||
|
||||
require (
|
||||
github.com/alessio/shellescape v1.4.1
|
||||
github.com/creack/pty v1.1.18
|
||||
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.10.0
|
||||
mvdan.cc/sh/v3 v3.7.0
|
||||
)
|
||||
@@ -0,0 +1,20 @@
|
||||
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.5 h1:dfYrrRyLtiqT9GyKXgdh+k4inNeTvmGbuSgZ3lx3GhA=
|
||||
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.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.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||
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-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=
|
||||
mvdan.cc/sh/v3 v3.7.0 h1:lSTjdP/1xsddtaKfGg7Myu7DnlHItd3/M2tomOcNNBg=
|
||||
mvdan.cc/sh/v3 v3.7.0/go.mod h1:K2gwkaesF/D7av7Kxl0HbF5kGOd2ArupNTX3X44+8l8=
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,381 @@
|
||||
package base
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"log"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/mod/semver"
|
||||
)
|
||||
|
||||
const HomeVarName = "HOME"
|
||||
const DefaultMShellHome = "~/.mshell"
|
||||
const DefaultMShellName = "mshell"
|
||||
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.3.0"
|
||||
const RemoteIdFile = "remoteid"
|
||||
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
|
||||
var DebugLogger *log.Logger
|
||||
var BuildTime string = "0"
|
||||
|
||||
type CommandFileNames struct {
|
||||
PtyOutFile string
|
||||
StdinFifo string
|
||||
RunnerOutFile string
|
||||
}
|
||||
|
||||
type CommandKey string
|
||||
|
||||
func SetBuildTime(build string) {
|
||||
BuildTime = build
|
||||
}
|
||||
|
||||
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 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)
|
||||
Logf("logger initialized\n")
|
||||
}
|
||||
|
||||
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 ""
|
||||
}
|
||||
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 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 == "" {
|
||||
return "/"
|
||||
}
|
||||
return homeVar
|
||||
}
|
||||
|
||||
func GetMShellHomeDir() string {
|
||||
homeVar := os.Getenv(MShellHomeVarName)
|
||||
if homeVar != "" {
|
||||
return homeVar
|
||||
}
|
||||
return ExpandHomeDir(DefaultMShellHome)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
base := path.Join(sdir, cmdId)
|
||||
return &CommandFileNames{
|
||||
PtyOutFile: base + ".ptyout",
|
||||
StdinFifo: base + ".stdin",
|
||||
RunnerOutFile: base + ".runout",
|
||||
}, 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 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")
|
||||
}
|
||||
baseLock.Lock()
|
||||
sdir, ok := sessionDirCache[sessionId]
|
||||
baseLock.Unlock()
|
||||
if ok {
|
||||
return sdir, nil
|
||||
}
|
||||
mhome := GetMShellHomeDir()
|
||||
sdir = path.Join(mhome, SessionsDirBaseName, sessionId)
|
||||
info, err := os.Stat(sdir)
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
err = os.MkdirAll(sdir, 0777)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot make mshell session directory[%s]: %w", sdir, 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)
|
||||
}
|
||||
baseLock.Lock()
|
||||
sessionDirCache[sessionId] = sdir
|
||||
baseLock.Unlock()
|
||||
return sdir, nil
|
||||
}
|
||||
|
||||
func GetMShellPath() (string, error) {
|
||||
msPath := os.Getenv(MShellPathVarName) // use MSHELL_PATH
|
||||
if msPath != "" {
|
||||
return exec.LookPath(msPath)
|
||||
}
|
||||
mhome := GetMShellHomeDir()
|
||||
userMShellPath := path.Join(mhome, DefaultMShellName) // look in ~/.mshell
|
||||
msPath, err := exec.LookPath(userMShellPath)
|
||||
if err == nil {
|
||||
return msPath, nil
|
||||
}
|
||||
return exec.LookPath(DefaultMShellName) // standard path lookup for 'mshell'
|
||||
}
|
||||
|
||||
func GetMShellSessionsDir() (string, error) {
|
||||
mhome := GetMShellHomeDir()
|
||||
return path.Join(mhome, SessionsDirBaseName), nil
|
||||
}
|
||||
|
||||
func ExpandHomeDir(pathStr string) string {
|
||||
if pathStr != "~" && !strings.HasPrefix(pathStr, "~/") {
|
||||
return pathStr
|
||||
}
|
||||
homeDir := GetHomeDir()
|
||||
if pathStr == "~" {
|
||||
return homeDir
|
||||
}
|
||||
return path.Join(homeDir, pathStr[2:])
|
||||
}
|
||||
|
||||
func ValidGoArch(goos string, goarch string) bool {
|
||||
return (goos == "darwin" || goos == "linux") && (goarch == "amd64" || goarch == "arm64")
|
||||
}
|
||||
|
||||
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 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)
|
||||
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) {
|
||||
// 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
|
||||
}
|
||||
}
|
||||
|
||||
func BoundInt(ival int, minVal int, maxVal int) int {
|
||||
if ival < minVal {
|
||||
return minVal
|
||||
}
|
||||
if ival > maxVal {
|
||||
return maxVal
|
||||
}
|
||||
return ival
|
||||
}
|
||||
|
||||
func BoundInt64(ival int64, minVal int64, maxVal int64) int64 {
|
||||
if ival < minVal {
|
||||
return minVal
|
||||
}
|
||||
if ival > maxVal {
|
||||
return maxVal
|
||||
}
|
||||
return ival
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
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) 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 ""
|
||||
}
|
||||
rtn := iter.Opts[iter.Pos]
|
||||
iter.Pos++
|
||||
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:]
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package binpack
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
type Unpacker struct {
|
||||
R FullByteReader
|
||||
Err error
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
if len(barr) > 0 {
|
||||
_, err = w.Write(barr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
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))
|
||||
_, 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
|
||||
}
|
||||
if lenVal == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
rtnBuf := make([]byte, int(lenVal))
|
||||
_, err = io.ReadFull(r, rtnBuf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
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 {
|
||||
return 0, err
|
||||
}
|
||||
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) 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
|
||||
}
|
||||
|
||||
func MakeUnpacker(r FullByteReader) *Unpacker {
|
||||
return &Unpacker{R: r}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,282 @@
|
||||
package cirfile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path"
|
||||
"strings"
|
||||
"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)
|
||||
str := string(barr)
|
||||
str = strings.ReplaceAll(str, "\x00", ".")
|
||||
fmt.Printf("%s<<<\n%s\n>>>\n", name, str)
|
||||
}
|
||||
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,471 @@
|
||||
package cmdtail
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"regexp"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"github.com/commandlinedev/apishell/pkg/base"
|
||||
"github.com/commandlinedev/apishell/pkg/packet"
|
||||
)
|
||||
|
||||
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
|
||||
TailPtyPos int64
|
||||
TailRunPos int64
|
||||
Follow bool
|
||||
}
|
||||
|
||||
type CmdWatchEntry struct {
|
||||
CmdKey base.CommandKey
|
||||
FilePtyLen int64
|
||||
FileRunLen int64
|
||||
Tails []TailPos
|
||||
Done bool
|
||||
}
|
||||
|
||||
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) {
|
||||
for _, pos := range w.Tails {
|
||||
if pos.ReqId == reqId {
|
||||
return pos, true
|
||||
}
|
||||
}
|
||||
return TailPos{}, false
|
||||
}
|
||||
|
||||
func (w *CmdWatchEntry) updateTailPos(reqId string, newPos TailPos) {
|
||||
for idx, pos := range w.Tails {
|
||||
if pos.ReqId == reqId {
|
||||
w.Tails[idx] = newPos
|
||||
return
|
||||
}
|
||||
}
|
||||
w.Tails = append(w.Tails, newPos)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos TailPos) {
|
||||
entry, found := t.WatchList[cmdKey]
|
||||
if !found {
|
||||
return
|
||||
}
|
||||
entry.updateTailPos(reqId, 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 {
|
||||
return
|
||||
}
|
||||
entry.removeTailPos(reqId)
|
||||
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))
|
||||
t.Watcher.Remove(t.Gen.RunOutFile(cmdKey))
|
||||
}
|
||||
|
||||
func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, 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 (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,
|
||||
}
|
||||
var err error
|
||||
rtn.Watcher, err = fsnotify.NewWatcher()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rtn, nil
|
||||
}
|
||||
|
||||
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(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(t.Gen.PtyOutFile(entry.CmdKey), pos.TailPtyPos, MaxDataBytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dataPacket.PtyData64 = base64.StdEncoding.EncodeToString(ptyData)
|
||||
dataPacket.PtyDataLen = len(ptyData)
|
||||
}
|
||||
if entry.FileRunLen > pos.TailRunPos {
|
||||
runData, err := t.readDataFromFile(t.Gen.RunOutFile(entry.CmdKey), pos.TailRunPos, MaxDataBytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dataPacket.RunData64 = base64.StdEncoding.EncodeToString(runData)
|
||||
dataPacket.RunDataLen = len(runData)
|
||||
}
|
||||
return dataPacket, nil
|
||||
}
|
||||
|
||||
// returns (data-packet, keepRunning)
|
||||
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, nil
|
||||
}
|
||||
dataPacket, dataErr := t.makeCmdDataPacket(entry, pos)
|
||||
|
||||
t.Lock.Lock()
|
||||
defer t.Lock.Unlock()
|
||||
entry, pos, foundPos = t.getEntryAndPos_nolock(key, reqId)
|
||||
if !foundPos {
|
||||
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, nil
|
||||
}
|
||||
if dataErr != nil {
|
||||
// error, so return error packet, and stop running
|
||||
pos.Running = false
|
||||
t.updateTailPos_nolock(key, reqId, pos)
|
||||
return nil, false, dataErr
|
||||
}
|
||||
pos.TailPtyPos += int64(dataPacket.PtyDataLen)
|
||||
pos.TailRunPos += int64(dataPacket.RunDataLen)
|
||||
if pos.IsCurrent(entry) {
|
||||
// we caught up, tail position equals file length
|
||||
pos.Running = false
|
||||
}
|
||||
t.updateTailPos_nolock(key, reqId, pos)
|
||||
return dataPacket, pos.Running, nil
|
||||
}
|
||||
|
||||
// returns (removed)
|
||||
func (t *Tailer) checkRemove(cmdKey base.CommandKey, reqId string) bool {
|
||||
t.Lock.Lock()
|
||||
defer t.Lock.Unlock()
|
||||
entry, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId)
|
||||
if !foundPos {
|
||||
return false
|
||||
}
|
||||
if !pos.IsCurrent(entry) {
|
||||
return false
|
||||
}
|
||||
if !pos.Follow || entry.Done {
|
||||
t.removeTailPos_nolock(cmdKey, reqId)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *Tailer) RunDataTransfer(key base.CommandKey, reqId string) {
|
||||
for {
|
||||
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 {
|
||||
removed := t.checkRemove(key, reqId)
|
||||
if removed {
|
||||
t.Sender.SendResponse(reqId, true)
|
||||
}
|
||||
break
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Tailer) tryStartRun_nolock(entry CmdWatchEntry, pos TailPos) {
|
||||
if pos.Running {
|
||||
return
|
||||
}
|
||||
if pos.IsCurrent(entry) {
|
||||
return
|
||||
}
|
||||
pos.Running = true
|
||||
t.updateTailPos_nolock(entry.CmdKey, pos.ReqId, pos)
|
||||
go t.RunDataTransfer(entry.CmdKey, pos.ReqId)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
finfo, err := os.Stat(relFileName)
|
||||
if err != nil {
|
||||
t.Sender.SendPacket(packet.FmtMessagePacket("error trying to stat file '%s': %v", relFileName, err))
|
||||
return
|
||||
}
|
||||
cmdKey := base.MakeCommandKey(m[1], m[2])
|
||||
t.Lock.Lock()
|
||||
defer t.Lock.Unlock()
|
||||
entry, foundEntry := t.WatchList[cmdKey]
|
||||
if !foundEntry {
|
||||
return
|
||||
}
|
||||
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 {
|
||||
t.tryStartRun_nolock(entry, pos)
|
||||
}
|
||||
}
|
||||
|
||||
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.Sender.SendPacket(packet.FmtMessagePacket("error in tailer: %v", err))
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (t *Tailer) Close() error {
|
||||
return t.Watcher.Close()
|
||||
}
|
||||
|
||||
func max(v1 int64, v2 int64) int64 {
|
||||
if v1 > v2 {
|
||||
return v1
|
||||
}
|
||||
return v2
|
||||
}
|
||||
|
||||
func (entry *CmdWatchEntry) fillFilePos(gen FileNameGenerator) {
|
||||
ptyInfo, _ := os.Stat(gen.PtyOutFile(entry.CmdKey))
|
||||
if ptyInfo != nil {
|
||||
entry.FilePtyLen = ptyInfo.Size()
|
||||
}
|
||||
runoutInfo, _ := os.Stat(gen.RunOutFile(entry.CmdKey))
|
||||
if runoutInfo != nil {
|
||||
entry.FileRunLen = runoutInfo.Size()
|
||||
}
|
||||
}
|
||||
|
||||
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()
|
||||
t.removeTailPos_nolock(pk.CK, pk.ReqId)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
if ptyOnly {
|
||||
return nil
|
||||
}
|
||||
err = t.Watcher.Add(runName)
|
||||
if err != nil {
|
||||
t.Watcher.Remove(ptyName) // best effort clean up
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// returns (up-to-date/done, error)
|
||||
func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) (bool, error) {
|
||||
if err := getPacket.CK.Validate("getcmd"); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if getPacket.ReqId == "" {
|
||||
return false, fmt.Errorf("getcmd, no reqid specified")
|
||||
}
|
||||
t.Lock.Lock()
|
||||
defer t.Lock.Unlock()
|
||||
key := getPacket.CK
|
||||
entry, foundEntry := t.WatchList[key]
|
||||
if !foundEntry {
|
||||
// initialize entry, add watches
|
||||
entry = CmdWatchEntry{CmdKey: key}
|
||||
entry.fillFilePos(t.Gen)
|
||||
}
|
||||
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
|
||||
}
|
||||
if pos.TailRunPos < 0 {
|
||||
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
|
||||
return true, nil
|
||||
}
|
||||
if !foundEntry {
|
||||
err := t.AddFileWatches_nolock(key, getPacket.PtyOnly)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
}
|
||||
t.WatchList[key] = entry
|
||||
t.tryStartRun_nolock(entry, pos)
|
||||
return false, nil
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package mpio
|
||||
|
||||
import (
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/commandlinedev/apishell/pkg/packet"
|
||||
)
|
||||
|
||||
type FdReader struct {
|
||||
CVar *sync.Cond
|
||||
M *Multiplexer
|
||||
FdNum int
|
||||
Fd io.ReadCloser
|
||||
BufSize int
|
||||
Closed bool
|
||||
ShouldCloseFd bool
|
||||
IsPty bool
|
||||
}
|
||||
|
||||
func MakeFdReader(m *Multiplexer, fd io.ReadCloser, fdNum int, shouldCloseFd bool, isPty bool) *FdReader {
|
||||
fr := &FdReader{
|
||||
CVar: sync.NewCond(&sync.Mutex{}),
|
||||
M: m,
|
||||
FdNum: fdNum,
|
||||
Fd: fd,
|
||||
BufSize: 0,
|
||||
ShouldCloseFd: shouldCloseFd,
|
||||
IsPty: isPty,
|
||||
}
|
||||
return fr
|
||||
}
|
||||
|
||||
func (r *FdReader) Close() {
|
||||
r.CVar.L.Lock()
|
||||
defer r.CVar.L.Unlock()
|
||||
if r.Closed {
|
||||
return
|
||||
}
|
||||
if r.Fd != nil && r.ShouldCloseFd {
|
||||
r.Fd.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()
|
||||
if r.Closed {
|
||||
return
|
||||
}
|
||||
r.BufSize -= ackLen
|
||||
if r.BufSize < 0 {
|
||||
r.BufSize = 0
|
||||
}
|
||||
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(pk packet.PacketType) {
|
||||
r.CVar.L.Unlock()
|
||||
defer r.CVar.L.Lock()
|
||||
r.M.sendPacket(pk)
|
||||
}
|
||||
|
||||
// returns (success)
|
||||
func (r *FdReader) WriteWait(data []byte, isEof bool) bool {
|
||||
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.M.makeDataPacket(r.FdNum, data[0:writeLen], nil)
|
||||
pk.Eof = isEof && (writeLen == len(data))
|
||||
r.BufSize += writeLen
|
||||
data = data[writeLen:]
|
||||
r.sendPacket_unlock(pk)
|
||||
if len(data) == 0 {
|
||||
return true
|
||||
}
|
||||
// do *not* do a CVar.Wait() here -- because we *unlocked* to send the packet, we should
|
||||
// recheck the condition before waiting to avoid deadlock.
|
||||
}
|
||||
}
|
||||
|
||||
func min(v1 int, v2 int) int {
|
||||
if v1 <= v2 {
|
||||
return v1
|
||||
}
|
||||
return v2
|
||||
}
|
||||
|
||||
func (r *FdReader) isClosed() bool {
|
||||
r.CVar.L.Lock()
|
||||
defer r.CVar.L.Unlock()
|
||||
return r.Closed
|
||||
}
|
||||
|
||||
func (r *FdReader) ReadLoop(wg *sync.WaitGroup) {
|
||||
defer r.Close()
|
||||
if wg != nil {
|
||||
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(buf[0:nr], (err == io.EOF))
|
||||
if !isOpen {
|
||||
return
|
||||
}
|
||||
if err == io.EOF {
|
||||
return
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if r.IsPty {
|
||||
r.WriteWait(nil, true)
|
||||
return
|
||||
}
|
||||
errPk := r.M.makeDataPacket(r.FdNum, nil, err)
|
||||
r.M.sendPacket(errPk)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package mpio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type FdWriter struct {
|
||||
CVar *sync.Cond
|
||||
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, desc string) *FdWriter {
|
||||
fw := &FdWriter{
|
||||
CVar: sync.NewCond(&sync.Mutex{}),
|
||||
Fd: fd,
|
||||
M: m,
|
||||
FdNum: fdNum,
|
||||
ShouldCloseFd: shouldCloseFd,
|
||||
Desc: desc,
|
||||
BufferLimit: WriteBufSize,
|
||||
}
|
||||
return fw
|
||||
}
|
||||
|
||||
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.ShouldCloseFd {
|
||||
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) AddData(data []byte, eof bool) error {
|
||||
w.CVar.L.Lock()
|
||||
defer w.CVar.L.Unlock()
|
||||
if w.Closed || w.Eof {
|
||||
if len(data) == 0 {
|
||||
return nil
|
||||
}
|
||||
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) > 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...)
|
||||
}
|
||||
if eof {
|
||||
w.Eof = true
|
||||
}
|
||||
w.CVar.Broadcast()
|
||||
return nil
|
||||
}
|
||||
|
||||
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
|
||||
for len(data) > 0 {
|
||||
if w.Closed {
|
||||
return
|
||||
}
|
||||
chunkSize := min(len(data), MaxSingleWriteSize)
|
||||
chunk := data[0:chunkSize]
|
||||
nw, err := w.Fd.Write(chunk)
|
||||
if nw > 0 || err != nil {
|
||||
ack := w.M.makeDataAckPacket(w.FdNum, nw, err)
|
||||
w.M.sendPacket(ack)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
data = data[chunkSize:]
|
||||
}
|
||||
if isEof {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,303 @@
|
||||
package mpio
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
|
||||
"github.com/commandlinedev/apishell/pkg/base"
|
||||
"github.com/commandlinedev/apishell/pkg/packet"
|
||||
)
|
||||
|
||||
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
|
||||
Input *packet.PacketParser
|
||||
Started bool
|
||||
UPR packet.UnknownPacketReporter
|
||||
|
||||
Debug bool
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Multiplexer) Close() {
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
|
||||
for _, fr := range m.FdReaders {
|
||||
fr.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()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum, true, false)
|
||||
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, 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, 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, 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, desc)
|
||||
fdWriter.BufferLimit = bufferLimit
|
||||
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, isPty bool) {
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose, isPty)
|
||||
}
|
||||
|
||||
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, desc)
|
||||
}
|
||||
|
||||
func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType {
|
||||
ack := packet.MakeDataAckPacket()
|
||||
ack.CK = m.CK
|
||||
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.CK = m.CK
|
||||
pk.FdNum = fdNum
|
||||
pk.Data64 = base64.StdEncoding.EncodeToString(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(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(wg)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Multiplexer) launchReaders(wg *sync.WaitGroup) {
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
if wg != nil {
|
||||
wg.Add(len(m.FdReaders))
|
||||
}
|
||||
for _, fr := range m.FdReaders {
|
||||
go fr.ReadLoop(wg)
|
||||
}
|
||||
}
|
||||
|
||||
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 = packetParser
|
||||
m.Sender = sender
|
||||
m.Started = true
|
||||
}
|
||||
|
||||
func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType {
|
||||
defer m.HandleInputDone()
|
||||
for pk := range m.Input.MainCh {
|
||||
if m.Debug {
|
||||
fmt.Printf("PK-M> %s\n", packet.AsString(pk))
|
||||
}
|
||||
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)
|
||||
continue
|
||||
}
|
||||
if pk.GetType() == packet.CmdDonePacketStr {
|
||||
donePacket := pk.(*packet.CmdDonePacketType)
|
||||
return donePacket
|
||||
}
|
||||
m.UPR.UnknownPacket(pk)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Multiplexer) WriteDataToFd(fdNum int, data []byte, isEof bool) error {
|
||||
m.Lock.Lock()
|
||||
defer m.Lock.Unlock()
|
||||
fw := m.FdWriters[fdNum]
|
||||
if fw == nil {
|
||||
// add a closed FdWriter as a placeholder so we only send one error
|
||||
fw := MakeFdWriter(m, nil, fdNum, false, "invalid-fd")
|
||||
fw.Close()
|
||||
m.FdWriters[fdNum] = fw
|
||||
return fmt.Errorf("write to closed file (no fd)")
|
||||
}
|
||||
err := fw.AddData(data, isEof)
|
||||
if err != nil {
|
||||
fw.Close()
|
||||
return err
|
||||
}
|
||||
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()
|
||||
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(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 {
|
||||
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
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,238 @@
|
||||
package packet
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type PacketParser struct {
|
||||
Lock *sync.Mutex
|
||||
MainCh chan PacketType
|
||||
RpcMap map[string]*RpcEntry
|
||||
RpcHandler bool
|
||||
Err error
|
||||
}
|
||||
|
||||
type RpcEntry struct {
|
||||
ReqId string
|
||||
RespCh chan RpcResponsePacketType
|
||||
}
|
||||
|
||||
type RpcResponseIter struct {
|
||||
ReqId string
|
||||
Parser *PacketParser
|
||||
}
|
||||
|
||||
func (iter *RpcResponseIter) Next(ctx context.Context) (RpcResponsePacketType, error) {
|
||||
// will unregister the rpc on ResponseDone
|
||||
return iter.Parser.GetNextResponse(ctx, iter.ReqId)
|
||||
}
|
||||
|
||||
func (iter *RpcResponseIter) Close() {
|
||||
iter.Parser.UnRegisterRpc(iter.ReqId)
|
||||
}
|
||||
|
||||
func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser, rpcHandler bool) *PacketParser {
|
||||
rtnParser := &PacketParser{
|
||||
Lock: &sync.Mutex{},
|
||||
MainCh: make(chan PacketType),
|
||||
RpcMap: make(map[string]*RpcEntry),
|
||||
RpcHandler: rpcHandler,
|
||||
}
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for pk := range p1.MainCh {
|
||||
if rtnParser.RpcHandler {
|
||||
sent := rtnParser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
}
|
||||
}
|
||||
rtnParser.MainCh <- pk
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for pk := range p2.MainCh {
|
||||
if rtnParser.RpcHandler {
|
||||
sent := rtnParser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
}
|
||||
}
|
||||
rtnParser.MainCh <- pk
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
wg.Wait()
|
||||
close(rtnParser.MainCh)
|
||||
}()
|
||||
return rtnParser
|
||||
}
|
||||
|
||||
// should have already registered rpc
|
||||
func (p *PacketParser) WaitForResponse(ctx context.Context, reqId string) RpcResponsePacketType {
|
||||
entry := p.getRpcEntry(reqId)
|
||||
if entry == nil {
|
||||
return nil
|
||||
}
|
||||
defer p.UnRegisterRpc(reqId)
|
||||
select {
|
||||
case resp := <-entry.RespCh:
|
||||
return resp
|
||||
case <-ctx.Done():
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PacketParser) GetResponseIter(reqId string) *RpcResponseIter {
|
||||
return &RpcResponseIter{Parser: p, ReqId: reqId}
|
||||
}
|
||||
|
||||
func (p *PacketParser) GetNextResponse(ctx context.Context, reqId string) (RpcResponsePacketType, error) {
|
||||
entry := p.getRpcEntry(reqId)
|
||||
if entry == nil {
|
||||
return nil, nil
|
||||
}
|
||||
select {
|
||||
case resp := <-entry.RespCh:
|
||||
if resp.GetResponseDone() {
|
||||
p.UnRegisterRpc(reqId)
|
||||
}
|
||||
return resp, nil
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PacketParser) UnRegisterRpc(reqId string) {
|
||||
p.Lock.Lock()
|
||||
defer p.Lock.Unlock()
|
||||
entry := p.RpcMap[reqId]
|
||||
if entry != nil {
|
||||
close(entry.RespCh)
|
||||
delete(p.RpcMap, reqId)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
entry := &RpcEntry{ReqId: reqId, RespCh: ch}
|
||||
p.RpcMap[reqId] = entry
|
||||
return ch
|
||||
}
|
||||
|
||||
func (p *PacketParser) getRpcEntry(reqId string) *RpcEntry {
|
||||
p.Lock.Lock()
|
||||
defer p.Lock.Unlock()
|
||||
entry := p.RpcMap[reqId]
|
||||
return entry
|
||||
}
|
||||
|
||||
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()]
|
||||
if entry == nil {
|
||||
return false
|
||||
}
|
||||
// nonblocking send
|
||||
select {
|
||||
case entry.RespCh <- respPk:
|
||||
default:
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (p *PacketParser) GetErr() error {
|
||||
p.Lock.Lock()
|
||||
defer p.Lock.Unlock()
|
||||
return p.Err
|
||||
}
|
||||
|
||||
func (p *PacketParser) SetErr(err error) {
|
||||
p.Lock.Lock()
|
||||
defer p.Lock.Unlock()
|
||||
if p.Err == nil {
|
||||
p.Err = err
|
||||
}
|
||||
}
|
||||
|
||||
func MakePacketParser(input io.Reader, rpcHandler bool) *PacketParser {
|
||||
parser := &PacketParser{
|
||||
Lock: &sync.Mutex{},
|
||||
MainCh: make(chan PacketType),
|
||||
RpcMap: make(map[string]*RpcEntry),
|
||||
RpcHandler: rpcHandler,
|
||||
}
|
||||
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 {
|
||||
parser.SetErr(err)
|
||||
return
|
||||
}
|
||||
if line == "\n" {
|
||||
continue
|
||||
}
|
||||
// ##[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])
|
||||
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 {
|
||||
parser.MainCh <- MakeRawPacket(line[:len(line)-1])
|
||||
continue
|
||||
}
|
||||
if pk.GetType() == DonePacketStr {
|
||||
return
|
||||
}
|
||||
if pk.GetType() == PingPacketStr {
|
||||
continue
|
||||
}
|
||||
if parser.RpcHandler {
|
||||
sent := parser.trySendRpcResponse(pk)
|
||||
if sent {
|
||||
continue
|
||||
}
|
||||
}
|
||||
parser.MainCh <- pk
|
||||
}
|
||||
}()
|
||||
return parser
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
package packet
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha1"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/commandlinedev/apishell/pkg/binpack"
|
||||
"github.com/commandlinedev/apishell/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"`
|
||||
HashVal string `json:"-"`
|
||||
}
|
||||
|
||||
type ShellStateDiff struct {
|
||||
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
|
||||
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))
|
||||
return sha1Hash(buf.Bytes()), buf.Bytes()
|
||||
}
|
||||
|
||||
func (state ShellState) MarshalJSON() ([]byte, error) {
|
||||
_, encodedBytes := state.EncodeAndHash()
|
||||
return json.Marshal(encodedBytes)
|
||||
}
|
||||
|
||||
// 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")
|
||||
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 (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")
|
||||
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.DiffHashArr = u.UnpackStrArr("ShellStateDiff.DiffHashArr")
|
||||
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) 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(vars bool, aliases bool, funcs bool) {
|
||||
fmt.Printf("ShellStateDiff:\n")
|
||||
fmt.Printf(" version: %s\n", sdiff.Version)
|
||||
fmt.Printf(" base: %s\n", sdiff.BaseHash)
|
||||
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()
|
||||
}
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,132 @@
|
||||
package shexec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os/exec"
|
||||
"time"
|
||||
|
||||
"github.com/commandlinedev/apishell/pkg/base"
|
||||
"github.com/commandlinedev/apishell/pkg/packet"
|
||||
"golang.org/x/mod/semver"
|
||||
)
|
||||
|
||||
// TODO - track buffer sizes for sending input
|
||||
|
||||
const NotFoundVersion = "v0.0"
|
||||
|
||||
type ClientProc struct {
|
||||
Cmd *exec.Cmd
|
||||
InitPk *packet.InitPacketType
|
||||
StartTs time.Time
|
||||
StdinWriter io.WriteCloser
|
||||
StdoutReader io.ReadCloser
|
||||
StderrReader io.ReadCloser
|
||||
Input *packet.PacketSender
|
||||
Output *packet.PacketParser
|
||||
}
|
||||
|
||||
// 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, nil, fmt.Errorf("creating stdin pipe: %v", err)
|
||||
}
|
||||
stdoutReader, err := ecmd.StdoutPipe()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("creating stdout pipe: %v", err)
|
||||
}
|
||||
stderrReader, err := ecmd.StderrPipe()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("creating stderr pipe: %v", err)
|
||||
}
|
||||
startTs := time.Now()
|
||||
err = ecmd.Start()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("running local client: %w", err)
|
||||
}
|
||||
sender := packet.MakePacketSender(inputWriter, nil)
|
||||
stdoutPacketParser := packet.MakePacketParser(stdoutReader, false)
|
||||
stderrPacketParser := packet.MakePacketParser(stderrReader, false)
|
||||
packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser, true)
|
||||
cproc := &ClientProc{
|
||||
Cmd: ecmd,
|
||||
StartTs: startTs,
|
||||
StdinWriter: inputWriter,
|
||||
StdoutReader: stdoutReader,
|
||||
StderrReader: stderrReader,
|
||||
Input: sender,
|
||||
Output: packetParser,
|
||||
}
|
||||
|
||||
var pk packet.PacketType
|
||||
select {
|
||||
case pk = <-packetParser.MainCh:
|
||||
case <-ctx.Done():
|
||||
cproc.Close()
|
||||
return nil, nil, ctx.Err()
|
||||
}
|
||||
if pk != nil {
|
||||
if pk.GetType() != packet.InitPacketStr {
|
||||
cproc.Close()
|
||||
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, fmt.Errorf("mshell client not found")
|
||||
}
|
||||
if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) {
|
||||
cproc.Close()
|
||||
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, nil, fmt.Errorf("no init packet received from mshell client")
|
||||
}
|
||||
return cproc, cproc.InitPk, 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) 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
|
||||
}
|
||||
sender.SendPacket(pk)
|
||||
}
|
||||
exitErr := cproc.Cmd.Wait()
|
||||
if !sentDonePk {
|
||||
endTs := time.Now()
|
||||
cmdDuration := endTs.Sub(cproc.StartTs)
|
||||
donePacket := packet.MakeCmdDonePacket(ck)
|
||||
donePacket.Ts = endTs.UnixMilli()
|
||||
donePacket.ExitCode = GetExitCode(exitErr)
|
||||
donePacket.DurationMs = int64(cmdDuration / time.Millisecond)
|
||||
sender.SendPacket(donePacket)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user