diff --git a/pkg/cmdrunner/resolver.go b/pkg/cmdrunner/resolver.go index 046337d2..cbd95abe 100644 --- a/pkg/cmdrunner/resolver.go +++ b/pkg/cmdrunner/resolver.go @@ -31,9 +31,9 @@ type resolvedIds struct { type ResolvedRemote struct { DisplayName string RemotePtr sstore.RemotePtrType - RemoteState *sstore.RemoteState - RState remote.RemoteRuntimeState MShell *remote.MShellProc + RState remote.RemoteRuntimeState + RemoteState *sstore.RemoteState } type ResolveItem struct { @@ -118,6 +118,22 @@ func resolveByPosition(items []ResolveItem, curId string, posStr string) *Resolv return &items[pos-1] } +func resolveRemoteArg(remoteArg string) (*sstore.RemotePtrType, error) { + rrUser, rrRemote, rrName, err := parseFullRemoteRef(remoteArg) + if err != nil { + return nil, err + } + if rrUser != "" { + return nil, fmt.Errorf("remoteusers not supported") + } + msh := remote.GetRemoteByName(rrRemote) + if msh == nil { + return nil, nil + } + rcopy := msh.GetRemoteCopy() + return &sstore.RemotePtrType{RemoteId: rcopy.RemoteId, Name: rrName}, nil +} + func resolveUiIds(ctx context.Context, pk *scpacket.FeCommandPacketType, rtype int) (resolvedIds, error) { rtn := resolvedIds{} uictx := pk.UIContext @@ -125,17 +141,6 @@ func resolveUiIds(ctx context.Context, pk *scpacket.FeCommandPacketType, rtype i rtn.SessionId = uictx.SessionId rtn.ScreenId = uictx.ScreenId rtn.WindowId = uictx.WindowId - if uictx.Remote != nil && rtn.SessionId != "" && rtn.WindowId != "" { - rr, err := resolveRemoteFromPtr(ctx, uictx.Remote, rtn.SessionId, rtn.WindowId) - if err != nil { - if rtype&R_Remote > 0 || rtype&R_RemoteConnected > 0 { - return rtn, err - } - // otherwise just don't set uictx.Remote - } else { - rtn.Remote = rr - } - } } if pk.Kwargs["window"] != "" { windowId, err := resolveWindowArg(rtn.SessionId, rtn.ScreenId, pk.Kwargs["window"]) @@ -146,6 +151,30 @@ func resolveUiIds(ctx context.Context, pk *scpacket.FeCommandPacketType, rtype i rtn.WindowId = windowId } } + var rptr *sstore.RemotePtrType + var err error + if pk.Kwargs["remote"] != "" { + rptr, err = resolveRemoteArg(pk.Kwargs["remote"]) + if err != nil { + return rtn, err + } + if rptr == nil { + return rtn, fmt.Errorf("invalid remote argument %q passed, remote not found", pk.Kwargs["remote"]) + } + } else if uictx.Remote != nil { + rptr = uictx.Remote + } + if rptr != nil { + err = rptr.Validate() + if err != nil { + return rtn, fmt.Errorf("invalid resolved remote: %v", err) + } + rr, err := resolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.WindowId) + if err != nil { + return rtn, err + } + rtn.Remote = rr + } if rtype&R_Session > 0 && rtn.SessionId == "" { return rtn, fmt.Errorf("no session") } @@ -285,7 +314,7 @@ func parseFullRemoteRef(fullRemoteRef string) (string, string, string, error) { func resolveRemoteFromPtr(ctx context.Context, rptr *sstore.RemotePtrType, sessionId string, windowId string) (*ResolvedRemote, error) { if rptr == nil || rptr.RemoteId == "" { - return nil, fmt.Errorf("no remote to resolve") + return nil, nil } msh := remote.GetRemoteById(rptr.RemoteId) if msh == nil { @@ -293,20 +322,24 @@ func resolveRemoteFromPtr(ctx context.Context, rptr *sstore.RemotePtrType, sessi } rstate := msh.GetRemoteRuntimeState() displayName := rstate.GetDisplayName(rptr) - state, err := sstore.GetRemoteState(ctx, sessionId, windowId, *rptr) - if err != nil { - return nil, fmt.Errorf("cannot resolve remote state '%s': %w", displayName, err) - } - if state == nil { - state = rstate.DefaultState - } - return &ResolvedRemote{ + rtn := &ResolvedRemote{ DisplayName: displayName, RemotePtr: *rptr, - RemoteState: state, + RemoteState: nil, RState: rstate, MShell: msh, - }, nil + } + if sessionId != "" && windowId != "" { + state, err := sstore.GetRemoteState(ctx, sessionId, windowId, *rptr) + if err != nil { + return nil, fmt.Errorf("cannot resolve remote state '%s': %w", displayName, err) + } + if state == nil { + state = rstate.DefaultState + } + rtn.RemoteState = state + } + return rtn, nil } // returns (remoteDisplayName, remoteptr, state, rstate, err) diff --git a/pkg/remote/remote.go b/pkg/remote/remote.go index d32a44c9..9a683e3c 100644 --- a/pkg/remote/remote.go +++ b/pkg/remote/remote.go @@ -173,7 +173,7 @@ func AddRemote(ctx context.Context, r *sstore.RemoteType) error { existingRemote := getRemoteByCanonicalName_nolock(r.RemoteCanonicalName) if existingRemote != nil { - erCopy := existingRemote.getRemoteCopy() + erCopy := existingRemote.GetRemoteCopy() if !erCopy.Archived { return fmt.Errorf("duplicate canonical name %q: cannot create new remote", r.RemoteCanonicalName) } @@ -202,7 +202,7 @@ func ArchiveRemote(ctx context.Context, remoteId string) error { if msh.Status == StatusConnected { return fmt.Errorf("cannot archive connected remote") } - rcopy := msh.getRemoteCopy() + rcopy := msh.GetRemoteCopy() archivedRemote := &sstore.RemoteType{ RemoteId: rcopy.RemoteId, RemoteType: rcopy.RemoteType, @@ -224,7 +224,8 @@ func GetRemoteByName(name string) *MShellProc { GlobalStore.Lock.Lock() defer GlobalStore.Lock.Unlock() for _, msh := range GlobalStore.Map { - if msh.Remote.RemoteAlias == name || msh.Remote.GetName() == name { + rcopy := msh.GetRemoteCopy() + if rcopy.RemoteAlias == name || rcopy.RemoteCanonicalName == name { return msh } } @@ -233,7 +234,7 @@ func GetRemoteByName(name string) *MShellProc { func getRemoteByCanonicalName_nolock(name string) *MShellProc { for _, msh := range GlobalStore.Map { - rcopy := msh.getRemoteCopy() + rcopy := msh.GetRemoteCopy() if rcopy.RemoteCanonicalName == name { return msh } @@ -454,7 +455,7 @@ func (msh *MShellProc) setErrorStatus(err error) { go msh.NotifyRemoteUpdate() } -func (msh *MShellProc) getRemoteCopy() sstore.RemoteType { +func (msh *MShellProc) GetRemoteCopy() sstore.RemoteType { msh.Lock.Lock() defer msh.Lock.Unlock() return *msh.Remote @@ -481,7 +482,7 @@ func (msh *MShellProc) GetRemoteName() string { } func (msh *MShellProc) Launch() { - remoteCopy := msh.getRemoteCopy() + remoteCopy := msh.GetRemoteCopy() remoteName := remoteCopy.GetName() if remoteCopy.Archived { logf(&remoteCopy, "cannot launch archived remote") diff --git a/pkg/sstore/sstore.go b/pkg/sstore/sstore.go index 7241df81..9efcbcae 100644 --- a/pkg/sstore/sstore.go +++ b/pkg/sstore/sstore.go @@ -12,6 +12,7 @@ import ( "os" "os/user" "path" + "regexp" "strings" "sync" "time" @@ -138,6 +139,8 @@ func (opts WindowShareOptsType) Value() (driver.Value, error) { return quickValueJson(opts) } +var RemoteNameRe = regexp.MustCompile("^\\*?[a-zA-Z0-9_-]+$") + type RemotePtrType struct { OwnerId string `json:"ownerid"` RemoteId string `json:"remoteid"` @@ -148,6 +151,26 @@ func (r RemotePtrType) IsSessionScope() bool { return strings.HasPrefix(r.Name, "*") } +func (r RemotePtrType) Validate() error { + if r.OwnerId != "" { + if _, err := uuid.Parse(r.OwnerId); err != nil { + return fmt.Errorf("invalid ownerid format: %v", err) + } + } + if r.RemoteId != "" { + if _, err := uuid.Parse(r.RemoteId); err != nil { + return fmt.Errorf("invalid remoteid format: %v", err) + } + } + if r.Name != "" { + ok := RemoteNameRe.MatchString(r.Name) + if !ok { + return fmt.Errorf("invalid remote name") + } + } + return nil +} + func (r RemotePtrType) MakeFullRemoteRef() string { if r.RemoteId == "" { return ""