From f5b9ea07a19c2871eff77d6deb0751b58f997ded Mon Sep 17 00:00:00 2001 From: sawka Date: Tue, 4 Oct 2022 11:45:24 -0700 Subject: [PATCH] local flag on remote, ensure 1 local remote. on archive, change to local remote --- cmd/main-server.go | 10 ----- db/migrations/000001_init.up.sql | 1 + db/schema.sql | 5 ++- pkg/cmdrunner/cmdrunner.go | 8 ++++ pkg/remote/remote.go | 49 +++++++++++++++++++-- pkg/sstore/dbops.go | 18 +++++++- pkg/sstore/sstore.go | 75 +++----------------------------- 7 files changed, 81 insertions(+), 85 deletions(-) diff --git a/cmd/main-server.go b/cmd/main-server.go index e3a6e4c3..508fea72 100644 --- a/cmd/main-server.go +++ b/cmd/main-server.go @@ -373,16 +373,6 @@ func main() { fmt.Printf("[error] ensuring local remote: %v\n", err) return } - err = sstore.AddTest01Remote(context.Background()) - if err != nil { - fmt.Printf("[error] ensuring test01 remote: %v\n", err) - return - } - //err = sstore.AddTest02Remote(context.Background()) - //if err != nil { - // fmt.Printf("[error] ensuring test02 remote: %v\n", err) - // return - //} _, err = sstore.EnsureDefaultSession(context.Background()) if err != nil { fmt.Printf("[error] ensuring default session: %v\n", err) diff --git a/db/migrations/000001_init.up.sql b/db/migrations/000001_init.up.sql index 40c75e55..08029a1d 100644 --- a/db/migrations/000001_init.up.sql +++ b/db/migrations/000001_init.up.sql @@ -94,6 +94,7 @@ CREATE TABLE remote ( sshopts json NOT NULL, remoteopts json NOT NULL, lastconnectts bigint NOT NULL, + local boolean NOT NULL, archived boolean NOT NULL, remoteidx int NOT NULL ); diff --git a/db/schema.sql b/db/schema.sql index c049c239..0b998b39 100644 --- a/db/schema.sql +++ b/db/schema.sql @@ -5,7 +5,8 @@ CREATE TABLE client ( userid varchar(36) NOT NULL, activesessionid varchar(36) NOT NULL, userpublickeybytes blob NOT NULL, - userprivatekeybytes blob NOT NULL + userprivatekeybytes blob NOT NULL, + winsize json NOT NULL ); CREATE TABLE session ( sessionid varchar(36) PRIMARY KEY, @@ -83,10 +84,12 @@ CREATE TABLE remote ( remoteuser varchar(50) NOT NULL, remotehost varchar(200) NOT NULL, connectmode varchar(20) NOT NULL, + autoinstall boolean NOT NULL, initpk json NOT NULL, sshopts json NOT NULL, remoteopts json NOT NULL, lastconnectts bigint NOT NULL, + local boolean NOT NULL, archived boolean NOT NULL, remoteidx int NOT NULL ); diff --git a/pkg/cmdrunner/cmdrunner.go b/pkg/cmdrunner/cmdrunner.go index c18bfc82..60c04964 100644 --- a/pkg/cmdrunner/cmdrunner.go +++ b/pkg/cmdrunner/cmdrunner.go @@ -805,6 +805,14 @@ func RemoteArchiveCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) return nil, fmt.Errorf("archiving remote: %v", err) } update := sstore.InfoMsgUpdate("remote [%s] archived", ids.Remote.DisplayName) + localRemote := remote.GetLocalRemote() + if localRemote != nil { + update.Window = &sstore.WindowType{ + SessionId: ids.SessionId, + WindowId: ids.WindowId, + CurRemote: sstore.RemotePtrType{RemoteId: localRemote.GetRemoteId()}, + } + } return update, nil } diff --git a/pkg/remote/remote.go b/pkg/remote/remote.go index 10b2fa62..ba7721e9 100644 --- a/pkg/remote/remote.go +++ b/pkg/remote/remote.go @@ -117,6 +117,7 @@ type RemoteRuntimeState struct { UName string `json:"uname"` MShellVersion string `json:"mshellversion"` WaitingForPassword bool `json:"waitingforpassword,omitempty"` + Local bool `json:"local,omitempty"` } func (state RemoteRuntimeState) IsConnected() bool { @@ -129,6 +130,12 @@ func (msh *MShellProc) GetStatus() string { return msh.Status } +func (msh *MShellProc) GetRemoteId() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Remote.RemoteId +} + func (msh *MShellProc) GetInstallStatus() string { msh.Lock.Lock() defer msh.Lock.Unlock() @@ -166,12 +173,22 @@ func LoadRemotes(ctx context.Context) error { if err != nil { return err } + var numLocal int for _, remote := range allRemotes { msh := MakeMShell(remote) GlobalStore.Map[remote.RemoteId] = msh if remote.ConnectMode == sstore.ConnectModeStartup { go msh.Launch() } + if remote.Local { + numLocal++ + } + } + if numLocal == 0 { + return fmt.Errorf("no local remote found") + } + if numLocal > 1 { + return fmt.Errorf("multiple local remotes found") } return nil } @@ -224,6 +241,10 @@ func AddRemote(ctx context.Context, r *sstore.RemoteType) error { } r.RemoteId = erCopy.RemoteId } + if r.Local { + return fmt.Errorf("cannot create another local remote (there can be only one)") + } + err := sstore.UpsertRemote(ctx, r) if err != nil { return fmt.Errorf("cannot create remote %q: %v", r.RemoteCanonicalName, err) @@ -247,6 +268,9 @@ func ArchiveRemote(ctx context.Context, remoteId string) error { if msh.Status == StatusConnected { return fmt.Errorf("cannot archive connected remote") } + if msh.Remote.Local { + return fmt.Errorf("cannot archive local remote") + } rcopy := msh.GetRemoteCopy() archivedRemote := &sstore.RemoteType{ RemoteId: rcopy.RemoteId, @@ -303,6 +327,17 @@ func GetRemoteById(remoteId string) *MShellProc { return GlobalStore.Map[remoteId] } +func GetLocalRemote() *MShellProc { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + for _, msh := range GlobalStore.Map { + if msh.IsLocal() { + return msh + } + } + return nil +} + func ResolveRemoteRef(remoteRef string) *RemoteRuntimeState { GlobalStore.Lock.Lock() defer GlobalStore.Lock.Unlock() @@ -369,6 +404,12 @@ func makeShortHost(host string) string { return host[0:dotIdx] } +func (msh *MShellProc) IsLocal() bool { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Remote.Local +} + func (msh *MShellProc) GetRemoteRuntimeState() RemoteRuntimeState { msh.Lock.Lock() defer msh.Lock.Unlock() @@ -386,6 +427,7 @@ func (msh *MShellProc) GetRemoteRuntimeState() RemoteRuntimeState { UName: msh.UName, InstallStatus: msh.InstallStatus, NeedsMShellUpgrade: msh.NeedsMShellUpgrade, + Local: msh.Remote.Local, } if msh.Err != nil { state.ErrorStr = msh.Err.Error() @@ -396,7 +438,6 @@ func (msh *MShellProc) GetRemoteRuntimeState() RemoteRuntimeState { if msh.Status == StatusConnecting { state.WaitingForPassword = msh.isWaitingForPassword_nolock() } - local := (msh.Remote.SSHOpts == nil || msh.Remote.SSHOpts.Local) vars := make(map[string]string) vars["user"] = msh.Remote.RemoteUser vars["bestuser"] = vars["user"] @@ -411,7 +452,7 @@ func (msh *MShellProc) GetRemoteRuntimeState() RemoteRuntimeState { if msh.Remote.RemoteSudo { vars["sudo"] = "1" } - if local { + if msh.Remote.Local { vars["local"] = "1" } vars["port"] = "22" @@ -437,12 +478,12 @@ func (msh *MShellProc) GetRemoteRuntimeState() RemoteRuntimeState { vars["besthost"] = vars["remotehost"] vars["bestshorthost"] = vars["remoteshorthost"] } - if local && msh.Remote.RemoteSudo { + if msh.Remote.Local && msh.Remote.RemoteSudo { vars["bestuser"] = "sudo" } else if msh.Remote.RemoteSudo { vars["bestuser"] = "sudo@" + vars["bestuser"] } - if local { + if msh.Remote.Local { vars["bestname"] = vars["bestuser"] + "@local" vars["bestshortname"] = vars["bestuser"] + "@local" } else { diff --git a/pkg/sstore/dbops.go b/pkg/sstore/dbops.go index 494eed77..b1f8bf76 100644 --- a/pkg/sstore/dbops.go +++ b/pkg/sstore/dbops.go @@ -73,6 +73,20 @@ func GetRemoteById(ctx context.Context, remoteId string) (*RemoteType, error) { return remote, nil } +func GetLocalRemote(ctx context.Context) (*RemoteType, error) { + var remote *RemoteType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote WHERE local` + m := tx.GetMap(query) + remote = RemoteFromMap(m) + return nil + }) + if err != nil { + return nil, err + } + return remote, nil +} + func GetRemoteByCanonicalName(ctx context.Context, cname string) (*RemoteType, error) { var remote *RemoteType err := WithTx(ctx, func(tx *TxWrap) error { @@ -131,8 +145,8 @@ func UpsertRemote(ctx context.Context, r *RemoteType) error { maxRemoteIdx := tx.GetInt(query) r.RemoteIdx = int64(maxRemoteIdx + 1) query = `INSERT INTO remote - ( remoteid, physicalid, remotetype, remotealias, remotecanonicalname, remotesudo, remoteuser, remotehost, connectmode, autoinstall, initpk, sshopts, remoteopts, lastconnectts, archived, remoteidx) VALUES - (:remoteid,:physicalid,:remotetype,:remotealias,:remotecanonicalname,:remotesudo,:remoteuser,:remotehost,:connectmode,:autoinstall,:initpk,:sshopts,:remoteopts,:lastconnectts,:archived,:remoteidx)` + ( remoteid, physicalid, remotetype, remotealias, remotecanonicalname, remotesudo, remoteuser, remotehost, connectmode, autoinstall, initpk, sshopts, remoteopts, lastconnectts, archived, remoteidx, local) VALUES + (:remoteid,:physicalid,:remotetype,:remotealias,:remotecanonicalname,:remotesudo,:remoteuser,:remotehost,:connectmode,:autoinstall,:initpk,:sshopts,:remoteopts,:lastconnectts,:archived,:remoteidx,:local)` tx.NamedExecWrap(query, r.ToMap()) return nil }) diff --git a/pkg/sstore/sstore.go b/pkg/sstore/sstore.go index d6b6f94f..2225ccdc 100644 --- a/pkg/sstore/sstore.go +++ b/pkg/sstore/sstore.go @@ -508,6 +508,7 @@ type RemoteType struct { LastConnectTs int64 `json:"lastconnectts"` Archived bool `json:"archived"` RemoteIdx int64 `json:"remoteidx"` + Local bool `json:"local"` } func (r *RemoteType) GetName() string { @@ -551,6 +552,7 @@ func (r *RemoteType) ToMap() map[string]interface{} { rtn["lastconnectts"] = r.LastConnectTs rtn["archived"] = r.Archived rtn["remoteidx"] = r.RemoteIdx + rtn["local"] = r.Local return rtn } @@ -575,6 +577,7 @@ func RemoteFromMap(m map[string]interface{}) *RemoteType { quickSetInt64(&r.LastConnectTs, m, "lastconnectts") quickSetBool(&r.Archived, m, "archived") quickSetInt64(&r.RemoteIdx, m, "remoteidx") + quickSetBool(&r.Local, m, "local") return &r } @@ -668,9 +671,9 @@ func EnsureLocalRemote(ctx context.Context) error { if err != nil { return fmt.Errorf("getting local physical remoteid: %w", err) } - remote, err := GetRemoteByPhysicalId(ctx, physicalId) + remote, err := GetLocalRemote(ctx) if err != nil { - return fmt.Errorf("getting remote[%s] from db: %w", physicalId, err) + return fmt.Errorf("getting local remote from db: %w", err) } if remote != nil { return nil @@ -696,77 +699,13 @@ func EnsureLocalRemote(ctx context.Context) error { ConnectMode: ConnectModeStartup, AutoInstall: true, SSHOpts: &SSHOpts{Local: true}, + Local: true, } err = UpsertRemote(ctx, localRemote) if err != nil { return err } - log.Printf("[db] added remote '%s', id=%s\n", localRemote.GetName(), localRemote.RemoteId) - return nil -} - -func AddTest01Remote(ctx context.Context) error { - remote, err := GetRemoteByAlias(ctx, "test01") - if err != nil { - return fmt.Errorf("getting remote[test01] from db: %w", err) - } - if remote != nil { - return nil - } - testRemote := &RemoteType{ - RemoteId: scbase.GenSCUUID(), - RemoteType: RemoteTypeSsh, - RemoteAlias: "test01", - RemoteCanonicalName: "ubuntu@test01.ec2", - RemoteSudo: false, - RemoteUser: "ubuntu", - RemoteHost: "test01.ec2", - SSHOpts: &SSHOpts{ - Local: false, - SSHHost: "test01.ec2", - SSHUser: "ubuntu", - SSHIdentity: "/Users/mike/aws/mfmt.pem", - }, - ConnectMode: ConnectModeStartup, - AutoInstall: true, - } - err = UpsertRemote(ctx, testRemote) - if err != nil { - return err - } - log.Printf("[db] added remote '%s', id=%s\n", testRemote.GetName(), testRemote.RemoteId) - return nil -} - -func AddTest02Remote(ctx context.Context) error { - remote, err := GetRemoteByAlias(ctx, "test2") - if err != nil { - return fmt.Errorf("getting remote[test01] from db: %w", err) - } - if remote != nil { - return nil - } - testRemote := &RemoteType{ - RemoteId: scbase.GenSCUUID(), - RemoteType: RemoteTypeSsh, - RemoteAlias: "test2", - RemoteCanonicalName: "test2@test01.ec2", - RemoteSudo: false, - RemoteUser: "test2", - RemoteHost: "test01.ec2", - SSHOpts: &SSHOpts{ - Local: false, - SSHHost: "test01.ec2", - SSHUser: "test2", - }, - ConnectMode: ConnectModeStartup, - AutoInstall: true, - } - err = UpsertRemote(ctx, testRemote) - if err != nil { - return err - } - log.Printf("[db] added remote '%s', id=%s\n", testRemote.GetName(), testRemote.RemoteId) + log.Printf("[db] added local remote '%s', id=%s\n", localRemote.RemoteCanonicalName, localRemote.RemoteId) return nil }