Skip to content

Commit d6fb969

Browse files
Make restore-token reset atomic
Ultraworked with [omo](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: sisyphus-dev-ai <sisyphus-dev-ai@users.noreply.github.com>
1 parent 0f38ba6 commit d6fb969

5 files changed

Lines changed: 250 additions & 36 deletions

File tree

internal/airplay/credentials.go

Lines changed: 70 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -103,6 +103,16 @@ type CredentialStore struct {
103103
backend CredentialBackend
104104
}
105105

106+
// RestoreTokenReset is a one-shot restore-token deletion which can be rolled
107+
// back if the daemon loses its stream reservation before reconnecting.
108+
type RestoreTokenReset struct {
109+
store *CredentialStore
110+
deviceID string
111+
previousToken string
112+
changed bool
113+
done bool
114+
}
115+
106116
// NewCredentialStore creates a credential store backed by a JSON file at path.
107117
func NewCredentialStore(path string) (*CredentialStore, error) {
108118
fb, err := newFileBackend(path)
@@ -196,19 +206,76 @@ func (cs *CredentialStore) SaveRestoreToken(deviceID, restoreToken string) error
196206
// ClearRestoreToken removes only the Wayland screencast restore token for a
197207
// device. Pairing credentials and all other device entries are preserved.
198208
func (cs *CredentialStore) ClearRestoreToken(deviceID string) error {
209+
reset, err := cs.BeginRestoreTokenReset(deviceID)
210+
if err != nil {
211+
return err
212+
}
213+
reset.Commit()
214+
return nil
215+
}
216+
217+
// BeginRestoreTokenReset clears the restore token while retaining enough state
218+
// to restore it if the caller cannot commit the corresponding reconnect.
219+
func (cs *CredentialStore) BeginRestoreTokenReset(deviceID string) (*RestoreTokenReset, error) {
199220
cs.mu.Lock()
200221
defer cs.mu.Unlock()
201222

223+
reset := &RestoreTokenReset{store: cs, deviceID: deviceID}
202224
creds, err := cs.backend.Lookup(deviceID)
203225
if err != nil {
204-
return err
226+
return nil, err
205227
}
206228
if creds == nil || creds.RestoreToken == "" {
207-
return nil
229+
return reset, nil
208230
}
231+
reset.previousToken = creds.RestoreToken
232+
reset.changed = true
209233
updated := *creds
210234
updated.RestoreToken = ""
211-
return cs.backend.Save(deviceID, &updated)
235+
if err := cs.backend.Save(deviceID, &updated); err != nil {
236+
return nil, err
237+
}
238+
return reset, nil
239+
}
240+
241+
// Commit makes a successful deletion permanent.
242+
func (r *RestoreTokenReset) Commit() {
243+
if r == nil {
244+
return
245+
}
246+
r.store.mu.Lock()
247+
r.done = true
248+
r.store.mu.Unlock()
249+
}
250+
251+
// Rollback restores only the previous token onto a fresh credential snapshot.
252+
// A concurrently saved non-empty token wins, and all other fields are retained.
253+
func (r *RestoreTokenReset) Rollback() error {
254+
if r == nil {
255+
return nil
256+
}
257+
r.store.mu.Lock()
258+
defer r.store.mu.Unlock()
259+
if r.done || !r.changed {
260+
r.done = true
261+
return nil
262+
}
263+
creds, err := r.store.backend.Lookup(r.deviceID)
264+
if err != nil {
265+
return err
266+
}
267+
if creds == nil {
268+
creds = &SavedCredentials{}
269+
}
270+
if creds.RestoreToken == "" {
271+
updated := *creds
272+
updated.RestoreToken = r.previousToken
273+
if err := r.store.backend.Save(r.deviceID, &updated); err != nil {
274+
return err
275+
}
276+
}
277+
r.done = true
278+
return nil
212279
}
213280

214281
// fileBackend stores credentials as a JSON file on disk.

internal/airplay/credentials_clear_test.go

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -111,3 +111,47 @@ func TestCredentialStoreClearRestoreTokenUsesBackendWithoutDeletingEntry(t *test
111111
t.Fatalf("backend credentials after clear = %+v", creds)
112112
}
113113
}
114+
115+
func TestRestoreTokenResetRollbackMergesConcurrentCredentialChanges(t *testing.T) {
116+
backend := &recordingCredentialBackend{devices: map[string]*SavedCredentials{
117+
"device-1": {PairingID: "pair-1", RestoreToken: "restore-1"},
118+
}}
119+
store := NewCredentialStoreWithBackend(backend)
120+
121+
reset, err := store.BeginRestoreTokenReset("device-1")
122+
if err != nil {
123+
t.Fatalf("BeginRestoreTokenReset: %v", err)
124+
}
125+
backend.devices["device-1"] = &SavedCredentials{
126+
PairingID: "pair-2",
127+
RestoreToken: "",
128+
}
129+
130+
if err := reset.Rollback(); err != nil {
131+
t.Fatalf("Rollback: %v", err)
132+
}
133+
creds := backend.devices["device-1"]
134+
if creds.PairingID != "pair-2" || creds.RestoreToken != "restore-1" {
135+
t.Fatalf("credentials after rollback = %+v", creds)
136+
}
137+
}
138+
139+
func TestRestoreTokenResetRollbackDoesNotOverwriteConcurrentToken(t *testing.T) {
140+
backend := &recordingCredentialBackend{devices: map[string]*SavedCredentials{
141+
"device-1": {RestoreToken: "restore-1"},
142+
}}
143+
store := NewCredentialStoreWithBackend(backend)
144+
145+
reset, err := store.BeginRestoreTokenReset("device-1")
146+
if err != nil {
147+
t.Fatalf("BeginRestoreTokenReset: %v", err)
148+
}
149+
backend.devices["device-1"] = &SavedCredentials{RestoreToken: "restore-2"}
150+
151+
if err := reset.Rollback(); err != nil {
152+
t.Fatalf("Rollback: %v", err)
153+
}
154+
if got := backend.devices["device-1"].RestoreToken; got != "restore-2" {
155+
t.Fatalf("restore token after rollback = %q, want concurrent token", got)
156+
}
157+
}

internal/daemon/daemon.go

Lines changed: 43 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1354,24 +1354,34 @@ func (d *Daemon) getOrStartPreparedCaptureGroup(ctx context.Context, entry *acti
13541354
if d.captureGroups == nil {
13551355
d.captureGroups = make(map[videoCaptureKey]*videoCaptureGroup)
13561356
}
1357-
if group := d.captureGroups[key]; group != nil {
1357+
group := d.captureGroups[key]
1358+
var captureCtx context.Context
1359+
var captureCancel context.CancelFunc
1360+
if group != nil {
13581361
if group.resetReservedBy != nil && group.resetReservedBy != entry {
13591362
d.mu.Unlock()
13601363
return nil, 0, fmt.Errorf("%w: %dx%d", errCaptureGroupResetReserved, key.maxWidth, key.maxHeight)
13611364
}
1362-
entry.captureGroup = group
1363-
broadcast := group.broadcast
1364-
d.mu.Unlock()
1365-
preparation.Close()
1366-
if broadcast == nil {
1367-
return nil, 0, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight)
1365+
if group.resetReservedBy == entry && group.broadcast == nil && group.capture == nil && group.cancel == nil {
1366+
captureCtx, captureCancel = context.WithCancel(context.Background())
1367+
group.cancel = captureCancel
1368+
group.resetReservedBy = nil
1369+
} else {
1370+
entry.captureGroup = group
1371+
broadcast := group.broadcast
1372+
d.mu.Unlock()
1373+
preparation.Close()
1374+
if broadcast == nil {
1375+
return nil, 0, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight)
1376+
}
1377+
return broadcast, group.minimumVideoLead, nil
13681378
}
1369-
return broadcast, group.minimumVideoLead, nil
1379+
} else {
1380+
captureCtx, captureCancel = context.WithCancel(context.Background())
1381+
group = &videoCaptureGroup{key: key, cancel: captureCancel}
1382+
d.captureGroups[key] = group
1383+
entry.captureGroup = group
13701384
}
1371-
1372-
captureCtx, captureCancel := context.WithCancel(context.Background())
1373-
group := &videoCaptureGroup{key: key, cancel: captureCancel}
1374-
d.captureGroups[key] = group
13751385
entry.captureGroup = group
13761386
d.mu.Unlock()
13771387

@@ -1464,26 +1474,33 @@ func (d *Daemon) getOrStartCaptureGroup(entry *activeStream, restoreToken, devic
14641474
if d.captureGroups == nil {
14651475
d.captureGroups = make(map[videoCaptureKey]*videoCaptureGroup)
14661476
}
1467-
if group := d.captureGroups[key]; group != nil {
1477+
group := d.captureGroups[key]
1478+
var captureCtx context.Context
1479+
var captureCancel context.CancelFunc
1480+
if group != nil {
14681481
if group.resetReservedBy != nil && group.resetReservedBy != entry {
14691482
d.mu.Unlock()
14701483
return nil, fmt.Errorf("%w: %dx%d", errCaptureGroupResetReserved, key.maxWidth, key.maxHeight)
14711484
}
1472-
entry.captureGroup = group
1473-
broadcast := group.broadcast
1474-
d.mu.Unlock()
1475-
if broadcast == nil {
1476-
return nil, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight)
1485+
if group.resetReservedBy == entry && group.broadcast == nil && group.capture == nil && group.cancel == nil {
1486+
captureCtx, captureCancel = context.WithCancel(context.Background())
1487+
group.cancel = captureCancel
1488+
group.resetReservedBy = nil
1489+
} else {
1490+
entry.captureGroup = group
1491+
broadcast := group.broadcast
1492+
d.mu.Unlock()
1493+
if broadcast == nil {
1494+
return nil, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight)
1495+
}
1496+
return broadcast, nil
14771497
}
1478-
return broadcast, nil
1498+
} else {
1499+
captureCtx, captureCancel = context.WithCancel(context.Background())
1500+
group = &videoCaptureGroup{key: key, cancel: captureCancel}
1501+
d.captureGroups[key] = group
1502+
entry.captureGroup = group
14791503
}
1480-
1481-
// Publish the group and cancellation hook before entering the display portal
1482-
// or launching GStreamer. A targeted disconnect can then cancel an orphaned
1483-
// startup without affecting captures used by other canvas groups.
1484-
captureCtx, captureCancel := context.WithCancel(context.Background())
1485-
group := &videoCaptureGroup{key: key, cancel: captureCancel}
1486-
d.captureGroups[key] = group
14871504
entry.captureGroup = group
14881505
d.mu.Unlock()
14891506

internal/daemon/reset_restore_token.go

Lines changed: 26 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -58,25 +58,35 @@ func (d *Daemon) handleResetRestoreToken(req Request) Response {
5858
d.streamWorkers.Add(1)
5959
d.mu.Unlock()
6060

61-
clearErr := d.credStore.ClearRestoreToken(deviceID)
61+
credentialReset, clearErr := d.credStore.BeginRestoreTokenReset(deviceID)
6262

6363
d.mu.Lock()
6464
reservationCurrent := d.restoreTokenResetReservationCurrentLocked(target, deviceID, port, entry, group)
65-
if group.resetReservedBy == entry {
66-
group.resetReservedBy = nil
67-
}
6865
state := d.overallStateLocked()
6966
if d.shuttingDown {
67+
if group.resetReservedBy == entry {
68+
group.resetReservedBy = nil
69+
}
7070
d.mu.Unlock()
71+
if clearErr == nil {
72+
_ = credentialReset.Rollback()
73+
}
7174
d.streamWorkers.Done()
7275
return Response{OK: false, State: state, Error: "daemon is shutting down"}
7376
}
7477
if !reservationCurrent {
78+
if group.resetReservedBy == entry {
79+
group.resetReservedBy = nil
80+
}
7581
d.mu.Unlock()
82+
if clearErr == nil {
83+
_ = credentialReset.Rollback()
84+
}
7685
d.streamWorkers.Done()
7786
return Response{OK: false, State: state, Error: "restore token reset was canceled for " + target}
7887
}
7988
if clearErr != nil {
89+
group.resetReservedBy = nil
8090
d.mu.Unlock()
8191
d.streamWorkers.Done()
8292
return Response{OK: false, State: state, Error: "clear restore token: " + clearErr.Error()}
@@ -91,7 +101,17 @@ func (d *Daemon) handleResetRestoreToken(req Request) Response {
91101
cancelFn: cancel,
92102
credentialCh: make(chan string, 1),
93103
}
104+
// Replace the old generation with an owner-only claim before releasing the
105+
// daemon lock. Peers cannot create or join this key while physical cleanup
106+
// runs; the designated replacement later converts the claim into a capture.
107+
group.resetReservedBy = nil
94108
cleanup := d.detachStreamLocked(target)
109+
reservation := &videoCaptureGroup{
110+
key: group.key,
111+
resetReservedBy: replacement,
112+
}
113+
replacement.captureGroup = reservation
114+
d.captureGroups[group.key] = reservation
95115
d.streams[target] = replacement
96116
d.mu.Unlock()
97117

@@ -112,6 +132,7 @@ func (d *Daemon) handleResetRestoreToken(req Request) Response {
112132
d.mu.Unlock()
113133
abandoned.run()
114134
cancel()
135+
_ = credentialReset.Rollback()
115136
d.streamWorkers.Done()
116137
if shuttingDown {
117138
return Response{OK: false, State: state, Error: "daemon is shutting down"}
@@ -122,6 +143,7 @@ func (d *Daemon) handleResetRestoreToken(req Request) Response {
122143
state = d.overallStateLocked()
123144
d.mu.Unlock()
124145

146+
credentialReset.Commit()
125147
go func() {
126148
defer d.streamWorkers.Done()
127149
d.connectAndStream(connCtx, replacement, target, port, "")

0 commit comments

Comments
 (0)