mirror of
https://github.com/lightninglabs/lightning-terminal.git
synced 2026-08-13 12:33:36 +02:00
session: implement DeleteReservedSession
This commit is contained in:
parent
e12c88bc16
commit
a18c0b656f
4 changed files with 171 additions and 62 deletions
|
|
@ -325,6 +325,10 @@ type Store interface {
|
|||
// StateReserved state.
|
||||
DeleteReservedSessions(ctx context.Context) error
|
||||
|
||||
// DeleteReservedSession deletes the session with the given ID if it is
|
||||
// in the StateReserved state.
|
||||
DeleteReservedSession(ctx context.Context, id ID) error
|
||||
|
||||
// ShiftState updates the state of the session with the given ID to the
|
||||
// "dest" state.
|
||||
ShiftState(ctx context.Context, id ID, dest State) error
|
||||
|
|
|
|||
|
|
@ -442,7 +442,10 @@ func (db *BoltStore) DeleteReservedSessions(_ context.Context) error {
|
|||
return err
|
||||
}
|
||||
|
||||
return sessionBucket.ForEach(func(k, v []byte) error {
|
||||
// We create a copy of the sessions to delete so that we are
|
||||
// not iterating and modifying the bucket at the same time.
|
||||
var sessionsToDelete []*Session
|
||||
err = sessionBucket.ForEach(func(k, v []byte) error {
|
||||
// We'll also get buckets here, skip those (identified
|
||||
// by nil value).
|
||||
if v == nil {
|
||||
|
|
@ -458,69 +461,120 @@ func (db *BoltStore) DeleteReservedSessions(_ context.Context) error {
|
|||
return nil
|
||||
}
|
||||
|
||||
err = sessionBucket.Delete(k)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sessionsToDelete = append(sessionsToDelete, session)
|
||||
|
||||
idIndexBkt := sessionBucket.Bucket(idIndexKey)
|
||||
if idIndexBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
// Delete the entire session ID bucket.
|
||||
err = idIndexBkt.DeleteBucket(session.ID[:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
groupIdIndexBkt := sessionBucket.Bucket(groupIDIndexKey)
|
||||
if groupIdIndexBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
groupBkt := groupIdIndexBkt.Bucket(session.GroupID[:])
|
||||
if groupBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
sessionIDsBkt := groupBkt.Bucket(sessionIDKey)
|
||||
if sessionIDsBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
var (
|
||||
seqKey []byte
|
||||
numSessions int
|
||||
)
|
||||
err = sessionIDsBkt.ForEach(func(k, v []byte) error {
|
||||
numSessions++
|
||||
|
||||
if !bytes.Equal(v, session.ID[:]) {
|
||||
return nil
|
||||
}
|
||||
|
||||
seqKey = k
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if numSessions == 0 {
|
||||
return fmt.Errorf("no sessions found for "+
|
||||
"group ID %x", session.GroupID)
|
||||
}
|
||||
|
||||
if numSessions == 1 {
|
||||
// Delete the whole group bucket.
|
||||
return groupBkt.DeleteBucket(sessionIDKey)
|
||||
}
|
||||
|
||||
// Else, delete just the session ID entry.
|
||||
return sessionIDsBkt.Delete(seqKey)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, session := range sessionsToDelete {
|
||||
if err := deleteSession(sessionBucket,
|
||||
session); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// deleteSession deletes all the parts of a session from the database. This
|
||||
// assumes that the session has already been fetched from the db.
|
||||
func deleteSession(sessionBucket *bbolt.Bucket, session *Session) error {
|
||||
sessionKey := getSessionKey(session)
|
||||
err := sessionBucket.Delete(sessionKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
idIndexBkt := sessionBucket.Bucket(idIndexKey)
|
||||
if idIndexBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
// Delete the entire session ID bucket.
|
||||
err = idIndexBkt.DeleteBucket(session.ID[:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
groupIdIndexBkt := sessionBucket.Bucket(groupIDIndexKey)
|
||||
if groupIdIndexBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
groupBkt := groupIdIndexBkt.Bucket(session.GroupID[:])
|
||||
if groupBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
sessionIDsBkt := groupBkt.Bucket(sessionIDKey)
|
||||
if sessionIDsBkt == nil {
|
||||
return ErrDBInitErr
|
||||
}
|
||||
|
||||
var (
|
||||
seqKey []byte
|
||||
numSessions int
|
||||
)
|
||||
err = sessionIDsBkt.ForEach(func(k, v []byte) error {
|
||||
numSessions++
|
||||
|
||||
if !bytes.Equal(v, session.ID[:]) {
|
||||
return nil
|
||||
}
|
||||
|
||||
seqKey = k
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if numSessions == 0 {
|
||||
return fmt.Errorf("no sessions found for "+
|
||||
"group ID %x", session.GroupID)
|
||||
}
|
||||
|
||||
if numSessions == 1 {
|
||||
// If this is the last session in the group, we can delete the
|
||||
// whole group bucket.
|
||||
return groupIdIndexBkt.DeleteBucket(session.GroupID[:])
|
||||
}
|
||||
|
||||
// Else, delete just the session ID entry from the group.
|
||||
return sessionIDsBkt.Delete(seqKey)
|
||||
}
|
||||
|
||||
// DeleteReservedSession removes a given session that is in the reserved state
|
||||
// from the database.
|
||||
//
|
||||
// NOTE: This is part of the Store interface.
|
||||
func (db *BoltStore) DeleteReservedSession(_ context.Context, id ID) error {
|
||||
return db.Update(func(tx *bbolt.Tx) error {
|
||||
sessionBucket, err := getBucket(tx, sessionBucketKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// We'll first get the session to make sure it's actually in the
|
||||
// reserved state before deleting. This gives us a slightly
|
||||
// better error message than just trying to delete and getting a
|
||||
// "not found" if the session was in another state.
|
||||
session, err := getSessionByID(sessionBucket, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if session.State != StateReserved {
|
||||
return fmt.Errorf("session not in reserved state, is "+
|
||||
"%v", session.State)
|
||||
}
|
||||
|
||||
return deleteSession(sessionBucket, session)
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ type SQLQueries interface {
|
|||
SetSessionGroupID(ctx context.Context, arg sqlc.SetSessionGroupIDParams) error
|
||||
UpdateSessionState(ctx context.Context, arg sqlc.UpdateSessionStateParams) error
|
||||
DeleteSessionsWithState(ctx context.Context, state int16) error
|
||||
DeleteSession(ctx context.Context, id int64) error
|
||||
GetAccountIDByAlias(ctx context.Context, alias int64) (int64, error)
|
||||
GetAccount(ctx context.Context, id int64) (sqlc.Account, error)
|
||||
}
|
||||
|
|
@ -431,6 +432,30 @@ func (s *SQLStore) DeleteReservedSessions(ctx context.Context) error {
|
|||
})
|
||||
}
|
||||
|
||||
// DeleteReservedSession removes a given session that is in the reserved state
|
||||
// from the database.
|
||||
//
|
||||
// NOTE: This is part of the Store interface.
|
||||
func (s *SQLStore) DeleteReservedSession(ctx context.Context, id ID) error {
|
||||
var writeTxOpts db.QueriesTxOptions
|
||||
return s.db.ExecTx(ctx, &writeTxOpts, func(db SQLQueries) error {
|
||||
session, err := db.GetSessionByAlias(ctx, id[:])
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return fmt.Errorf("%w: unable to get session: %w",
|
||||
ErrSessionNotFound, err)
|
||||
} else if err != nil {
|
||||
return fmt.Errorf("unable to get session: %w", err)
|
||||
}
|
||||
|
||||
if State(session.State) != StateReserved {
|
||||
return fmt.Errorf("session not in reserved state, is "+
|
||||
"%v", State(session.State))
|
||||
}
|
||||
|
||||
return db.DeleteSession(ctx, session.ID)
|
||||
})
|
||||
}
|
||||
|
||||
// GetSessionByLocalPub fetches the session with the given local pub key.
|
||||
//
|
||||
// NOTE: This is part of the Store interface.
|
||||
|
|
|
|||
|
|
@ -156,6 +156,10 @@ func TestBasicSessionStore(t *testing.T) {
|
|||
// of the sessions are reserved.
|
||||
require.NoError(t, db.DeleteReservedSessions(ctx))
|
||||
|
||||
// Explicitly trying to delete session 1 should fail as it's not
|
||||
// reserved.
|
||||
require.Error(t, db.DeleteReservedSession(ctx, s1.ID))
|
||||
|
||||
sessions, err = db.ListSessionsByState(ctx, StateReserved)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, sessions)
|
||||
|
|
@ -192,6 +196,28 @@ func TestBasicSessionStore(t *testing.T) {
|
|||
_, err = db.GetGroupID(ctx, s4.ID)
|
||||
require.ErrorIs(t, err, ErrSessionNotFound)
|
||||
|
||||
// Reserve a new session and link it to session 1.
|
||||
s5, err := reserveSession(
|
||||
db, "session 5", withLinkedGroupID(&session1.GroupID),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
sessions, err = db.ListSessionsByState(ctx, StateReserved)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, len(sessions))
|
||||
assertEqualSessions(t, s5, sessions[0])
|
||||
|
||||
// Now delete the reserved session by its ID and show that it is no
|
||||
// longer in the database and no longer in the group ID/session ID
|
||||
// index.
|
||||
require.NoError(t, db.DeleteReservedSession(ctx, s5.ID))
|
||||
|
||||
sessions, err = db.ListSessionsByState(ctx, StateReserved)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, sessions)
|
||||
|
||||
_, err = db.GetGroupID(ctx, s5.ID)
|
||||
require.ErrorIs(t, err, ErrSessionNotFound)
|
||||
|
||||
// Only session 1 should remain in this group.
|
||||
sessIDs, err = db.GetSessionIDs(ctx, s4.GroupID)
|
||||
require.NoError(t, err)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue