update statediff algorithm for wavesrv / remote instances (#530)

* remote statemap from waveshell server (diff against initial state)

* move ShellStatePtr from sstore to packet so it can be passed over the wire

* add finalstatebaseptr to cmddone

* much improved diff computation code on wavesrv side

* fix displayname -- now using hash

* add comments, change a couple msh.WriteToPtyBuffer calls to log.Printfs
This commit is contained in:
Mike Sawka
2024-03-28 16:56:39 -07:00
committed by GitHub
parent 5c85b2b786
commit a1e4e807cc
12 changed files with 182 additions and 108 deletions
+108 -35
View File
@@ -1886,7 +1886,7 @@ type RunCommandOpts struct {
// optional, if not provided shellstate will look up state from remote instance
// ReturnState cannot be used with StatePtr
// this will also cause this command to bypass the pending state cmd logic
StatePtr *sstore.ShellStatePtr
StatePtr *packet.ShellStatePtr
// set to true to skip creating the pty file (for restarted commands)
NoCreateCmdPtyFile bool
@@ -1920,6 +1920,9 @@ func RunCommand(ctx context.Context, rcOpts RunCommandOpts, runPacket *packet.Ru
if rcOpts.StatePtr != nil && runPacket.ReturnState {
return nil, nil, fmt.Errorf("RunCommand: cannot use ReturnState with StatePtr")
}
if runPacket.StatePtr != nil {
return nil, nil, fmt.Errorf("runPacket.StatePtr should not be set, it is set in RunCommand")
}
// pending state command logic
// if we are currently running a command that can change the state, we need to wait for it to finish
@@ -1950,7 +1953,7 @@ func RunCommand(ctx context.Context, rcOpts RunCommandOpts, runPacket *packet.Ru
}
// get current remote-instance state
var statePtr *sstore.ShellStatePtr
var statePtr *packet.ShellStatePtr
if rcOpts.StatePtr != nil {
statePtr = rcOpts.StatePtr
} else {
@@ -1963,6 +1966,8 @@ func RunCommand(ctx context.Context, rcOpts RunCommandOpts, runPacket *packet.Ru
return nil, nil, fmt.Errorf("cannot run command: no valid shell state found")
}
}
// statePtr will not be nil
runPacket.StatePtr = statePtr
currentState, err := sstore.GetFullState(ctx, *statePtr)
if err != nil || currentState == nil {
return nil, nil, fmt.Errorf("cannot load current remote state: %w", err)
@@ -2216,31 +2221,95 @@ func (msh *MShellProc) notifyHangups_nolock() {
msh.PendingStateCmds = make(map[pendingStateKey]base.CommandKey)
}
// either fullstate or statediff will be set (not both) <- this is so the result is compatible with the sstore.UpdateRemoteState function
// note that this function *does* touch the DB, if FinalStateDiff is set, will ensure that StateBase is written to DB
func (msh *MShellProc) makeStatePtrFromFinalState(ctx context.Context, donePk *packet.CmdDonePacketType) (*sstore.ShellStatePtr, map[string]string, *packet.ShellState, *packet.ShellStateDiff, error) {
func (msh *MShellProc) resolveFinalState(ctx context.Context, origState *packet.ShellState, origStatePtr *packet.ShellStatePtr, donePk *packet.CmdDonePacketType) (*packet.ShellState, error) {
if donePk.FinalState != nil {
if origStatePtr == nil {
return nil, fmt.Errorf("command must have a stateptr to resolve final state")
}
finalState := stripScVarsFromState(donePk.FinalState)
feState := sstore.FeStateFromShellState(finalState)
statePtr := &sstore.ShellStatePtr{BaseHash: finalState.GetHashVal(false)}
return statePtr, feState, finalState, nil, nil
return finalState, nil
}
if donePk.FinalStateDiff != nil {
if donePk.FinalStateBasePtr == nil {
return nil, fmt.Errorf("invalid rtnstate, has diff but no baseptr")
}
stateDiff := stripScVarsFromStateDiff(donePk.FinalStateDiff)
feState, err := msh.getFeStateFromDiff(stateDiff)
if origStatePtr == donePk.FinalStateBasePtr {
// this is the normal case. the stateptr from the run-packet should match the baseptr from the done-packet
// this is also the most efficient, because we don't need to fetch the original state
sapi, err := shellapi.MakeShellApi(origState.GetShellType())
if err != nil {
return nil, fmt.Errorf("cannot make shellapi from initial state: %w", err)
}
fullState, err := sapi.ApplyShellStateDiff(origState, stateDiff)
if err != nil {
return nil, fmt.Errorf("cannot apply shell state diff: %w", err)
}
return fullState, nil
}
// this is strange (why is backend returning non-original stateptr?)
// but here, we fetch the stateptr, and then apply the diff against that
realOrigState, err := sstore.GetFullState(ctx, *donePk.FinalStateBasePtr)
if err != nil {
return nil, nil, nil, nil, err
return nil, fmt.Errorf("cannot get original state for diff: %w", err)
}
fullState := msh.StateMap.GetStateByHash(stateDiff.GetShellType(), stateDiff.BaseHash)
if fullState != nil {
sstore.StoreStateBase(ctx, fullState)
if realOrigState == nil {
return nil, fmt.Errorf("cannot get original state for diff: not found")
}
diffHashArr := append(([]string)(nil), donePk.FinalStateDiff.DiffHashArr...)
diffHashArr = append(diffHashArr, donePk.FinalStateDiff.GetHashVal(false))
statePtr := &sstore.ShellStatePtr{BaseHash: donePk.FinalStateDiff.BaseHash, DiffHashArr: diffHashArr}
return statePtr, feState, nil, stateDiff, nil
sapi, err := shellapi.MakeShellApi(realOrigState.GetShellType())
if err != nil {
return nil, fmt.Errorf("cannot make shellapi from original state: %w", err)
}
fullState, err := sapi.ApplyShellStateDiff(realOrigState, stateDiff)
if err != nil {
return nil, fmt.Errorf("cannot apply shell state diff: %w", err)
}
return fullState, nil
}
return nil, nil, nil, nil, nil
return nil, nil
}
// after this limit we'll switch to persisting the full state
const NewStateDiffSizeThreshold = 30 * 1024
// will update the remote instance with the final state
// this is complicated because we want to be as efficient as possible.
// so we pull the current remote-instance state (just the baseptr). then we compute the diff.
// then we check the size of the diff, and only persist the diff it is under some size threshold
// also we check to see if the diff succeeds (it can fail if the shell or version changed).
// in those cases we also update the RI with the full state
func (msh *MShellProc) updateRIWithFinalState(ctx context.Context, rct *RunCmdType, newState *packet.ShellState) (*sstore.RemoteInstance, error) {
curRIState, err := sstore.GetRemoteStatePtr(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr)
if err != nil {
return nil, fmt.Errorf("error trying to get current screen stateptr: %w", err)
}
feState := sstore.FeStateFromShellState(newState)
if curRIState == nil {
// no current state, so just persist the full state
return sstore.UpdateRemoteState(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, newState, nil)
}
// pull the base (not the diff) state from the RI (right now we don't want to make multi-level diffs)
riBaseState, err := sstore.GetStateBase(ctx, curRIState.BaseHash)
if err != nil {
return nil, fmt.Errorf("error trying to get statebase: %w", err)
}
sapi, err := shellapi.MakeShellApi(riBaseState.GetShellType())
if err != nil {
return nil, fmt.Errorf("error trying to make shellapi: %w", err)
}
newStateDiff, err := sapi.MakeShellStateDiff(riBaseState, curRIState.BaseHash, newState)
if err != nil {
// if we can't make a diff, just persist the full state (this could happen if the shell type changes)
return sstore.UpdateRemoteState(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, newState, nil)
}
// we have a diff, let's check the diff size first
_, encodedDiff := newStateDiff.EncodeAndHash()
if len(encodedDiff) > NewStateDiffSizeThreshold {
// diff is too large, persist the full state
return sstore.UpdateRemoteState(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, newState, nil)
}
// diff is small enough, persist the diff
return sstore.UpdateRemoteState(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, nil, newStateDiff)
}
func (msh *MShellProc) handleCmdDonePacket(rct *RunCmdType, donePk *packet.CmdDonePacketType) {
@@ -2261,12 +2330,12 @@ func (msh *MShellProc) handleCmdDonePacket(rct *RunCmdType, donePk *packet.CmdDo
// only update DB for non-ephemeral commands
err := sstore.UpdateCmdDoneInfo(ctx, update, donePk.CK, donePk, sstore.CmdStatusDone)
if err != nil {
msh.WriteToPtyBuffer("*error updating cmddone: %v\n", err)
log.Printf("error updating cmddone info (in handleCmdDonePacket): %v\n", err)
return
}
screen, err := sstore.UpdateScreenFocusForDoneCmd(ctx, donePk.CK.GetGroupId(), donePk.CK.GetCmdId())
if err != nil {
msh.WriteToPtyBuffer("*error trying to update screen focus type: %v\n", err)
log.Printf("error trying to update screen focus type (in handleCmdDonePacket): %v\n", err)
// fall-through (nothing to do)
}
if screen != nil {
@@ -2274,24 +2343,28 @@ func (msh *MShellProc) handleCmdDonePacket(rct *RunCmdType, donePk *packet.CmdDo
}
}
// ephemeral commands *do* update the remote state
if donePk.FinalState != nil || donePk.FinalStateDiff != nil {
statePtr, feState, finalState, finalStateDiff, err := msh.makeStatePtrFromFinalState(ctx, donePk)
// not all commands get a final state (only RtnState commands have this returned)
// so in those cases finalState will be nil
finalState, err := msh.resolveFinalState(ctx, rct.RunPacket.State, rct.RunPacket.StatePtr, donePk)
if err != nil {
log.Printf("error resolving final state for cmd: %v\n", err)
// fallthrough
}
if finalState != nil {
newRI, err := msh.updateRIWithFinalState(ctx, rct, finalState)
if err != nil {
msh.WriteToPtyBuffer("*error trying to read final command state: %v\n", err)
log.Printf("error updating RI with final state (in handleCmdDonePacket): %v\n", err)
// fallthrough
}
remoteInst, err := sstore.UpdateRemoteState(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, finalState, finalStateDiff)
if err != nil {
msh.WriteToPtyBuffer("*error trying to update remotestate: %v\n", err)
// fall-through (nothing to do)
}
if remoteInst != nil {
update.AddUpdate(sstore.MakeSessionUpdateForRemote(rct.SessionId, remoteInst))
if newRI != nil {
update.AddUpdate(sstore.MakeSessionUpdateForRemote(rct.SessionId, newRI))
}
// ephemeral commands *do not* update cmd state (there is no command)
if statePtr != nil && !rct.Ephemeral {
err = sstore.UpdateCmdRtnState(ctx, donePk.CK, *statePtr)
if newRI != nil && !rct.Ephemeral {
newRIStatePtr := packet.ShellStatePtr{BaseHash: newRI.StateBaseHash, DiffHashArr: newRI.StateDiffHashArr}
err = sstore.UpdateCmdRtnState(ctx, donePk.CK, newRIStatePtr)
if err != nil {
msh.WriteToPtyBuffer("*error trying to update cmd rtnstate: %v\n", err)
log.Printf("error trying to update cmd rtnstate: %v\n", err)
// fall-through (nothing to do)
}
}
@@ -2605,7 +2678,7 @@ func (msh *MShellProc) getFullState(shellType string, stateDiff *packet.ShellSta
}
return newState, nil
} else {
fullState, err := sstore.GetFullState(context.Background(), sstore.ShellStatePtr{BaseHash: stateDiff.BaseHash, DiffHashArr: stateDiff.DiffHashArr})
fullState, err := sstore.GetFullState(context.Background(), packet.ShellStatePtr{BaseHash: stateDiff.BaseHash, DiffHashArr: stateDiff.DiffHashArr})
if err != nil {
return nil, err
}
@@ -2632,7 +2705,7 @@ func (msh *MShellProc) getFeStateFromDiff(stateDiff *packet.ShellStateDiff) (map
}
return sstore.FeStateFromShellState(newState), nil
} else {
fullState, err := sstore.GetFullState(context.Background(), sstore.ShellStatePtr{BaseHash: stateDiff.BaseHash, DiffHashArr: stateDiff.DiffHashArr})
fullState, err := sstore.GetFullState(context.Background(), packet.ShellStatePtr{BaseHash: stateDiff.BaseHash, DiffHashArr: stateDiff.DiffHashArr})
if err != nil {
return nil, err
}