feat: fence command acceptance by session
This commit is contained in:
@@ -229,6 +229,18 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w
|
||||
server.close(connection, websocket.StatusPolicyViolation, "invalid client capacity")
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
if acknowledgement := envelope.GetCommandAccepted(); acknowledgement != nil {
|
||||
issue, parseErr := domain.ParseUUIDv7(acknowledgement.GetIssueUuid())
|
||||
if parseErr != nil {
|
||||
server.close(connection, websocket.StatusPolicyViolation, "invalid command acceptance")
|
||||
return
|
||||
}
|
||||
if _, acceptErr := server.Store.RecordCommandAcceptance(sessionContext, issue, hello.GetClientId(), registration.Generation, acknowledgement.GetCommandRevision(), acknowledgement.GetAccepted(), server.now()); acceptErr != nil {
|
||||
server.close(connection, websocket.StatusPolicyViolation, "invalid command acceptance")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/rvbox/rvbox/internal/domain"
|
||||
)
|
||||
|
||||
var ErrDispatchMismatch = errors.New("command acceptance does not match dispatched generation or revision")
|
||||
|
||||
// RecordCommandAcceptance fences an acknowledgement to the exact dispatch that
|
||||
// produced it. Repeated acknowledgements are harmless; a stale session cannot
|
||||
// advance a newer dispatch or overwrite a terminal decision.
|
||||
func (store *Store) RecordCommandAcceptance(ctx context.Context, issueUUID domain.UUID, clientID string, generation, revision uint64, accepted bool, now time.Time) (bool, error) {
|
||||
if isZeroUUID([16]byte(issueUUID)) || clientID == "" || generation == 0 || revision == 0 || now.IsZero() {
|
||||
return false, ErrDispatchMismatch
|
||||
}
|
||||
store.writeMu.Lock()
|
||||
defer store.writeMu.Unlock()
|
||||
database, err := store.openDatabase()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
targetLifecycle := 3
|
||||
if !accepted {
|
||||
targetLifecycle = 11
|
||||
}
|
||||
query := `UPDATE commands SET lifecycle = ?`
|
||||
args := []any{targetLifecycle}
|
||||
if !accepted {
|
||||
query += `, terminal_time = ?`
|
||||
args = append(args, now.UTC().UnixNano())
|
||||
}
|
||||
query += ` WHERE issue_uuid = ? AND client_id = ? AND lifecycle = 2 AND target_session_generation = ? AND revision = ?`
|
||||
args = append(args, issueUUID[:], clientID, generation, revision)
|
||||
result, err := database.ExecContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
changed, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if changed == 1 {
|
||||
return true, nil
|
||||
}
|
||||
var lifecycle uint32
|
||||
var storedGeneration, storedRevision uint64
|
||||
err = database.QueryRowContext(ctx, `SELECT lifecycle, COALESCE(target_session_generation, 0), revision FROM commands WHERE issue_uuid = ? AND client_id = ?`, issueUUID[:], clientID).Scan(&lifecycle, &storedGeneration, &storedRevision)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, ErrDispatchMismatch
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if storedGeneration == generation && storedRevision == revision && ((accepted && (lifecycle == 3 || lifecycle == 4)) || (!accepted && lifecycle == 11)) {
|
||||
return false, nil
|
||||
}
|
||||
return false, ErrDispatchMismatch
|
||||
}
|
||||
@@ -123,4 +123,36 @@ func TestClaimDispatchExpiresAndFencesRequeue_HP_DISPATCH_02(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordCommandAcceptanceFencesGeneration_HP_DISPATCH_05(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
opened, err := Open(ctx, Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = opened.Close() })
|
||||
if _, err := opened.RegisterClientSession(ctx, ClientRegistration{ClientID: "win-client", Platform: 2, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\`, SupportedShells: []byte{1}, ClientInstanceID: [16]byte{5}, SessionID: [16]byte{6}, ConnectedAt: time.Now()}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
issue, err := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000095")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := opened.QueueCommand(ctx, QueueCommandInput{IssueUUID: issue, ClientID: "win-client", IssueTime: time.Now(), ReceiptTime: time.Now(), ImmutableSHA256: sha256.Sum256([]byte("accept")), ExecutionSpec: []byte("spec")}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := opened.ClaimNextDispatch(ctx, "win-client", 3, time.Now()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := opened.RecordCommandAcceptance(ctx, issue, "win-client", 2, 1, true, time.Now()); !errors.Is(err, ErrDispatchMismatch) {
|
||||
t.Fatalf("stale acceptance error = %v", err)
|
||||
}
|
||||
if changed, err := opened.RecordCommandAcceptance(ctx, issue, "win-client", 3, 1, true, time.Now()); err != nil || !changed {
|
||||
t.Fatalf("acceptance = %t, %v", changed, err)
|
||||
}
|
||||
if changed, err := opened.RecordCommandAcceptance(ctx, issue, "win-client", 3, 1, true, time.Now()); err != nil || changed {
|
||||
t.Fatalf("duplicate acceptance = %t, %v", changed, err)
|
||||
}
|
||||
}
|
||||
|
||||
func timePtr(value time.Time) *time.Time { return &value }
|
||||
|
||||
Reference in New Issue
Block a user