diff --git a/README.md b/README.md index b986a30..b4c208a 100644 --- a/README.md +++ b/README.md @@ -344,6 +344,7 @@ doubletake-ctl status doubletake-ctl connect 192.168.1.77 doubletake-ctl connect 192.168.1.133 doubletake-ctl disconnect 192.168.1.77 +doubletake-ctl reset-restore-token 192.168.1.133 doubletake-ctl disconnect ``` @@ -399,6 +400,7 @@ doubletake-ctl devices doubletake-ctl connect [target] [PIN-or-password] doubletake-ctl pin doubletake-ctl disconnect [target] +doubletake-ctl reset-restore-token doubletake-ctl mute [target] doubletake-ctl unmute [target] ``` @@ -406,6 +408,11 @@ doubletake-ctl unmute [target] - `disconnect` without a target stops all active streams. - `disconnect ` stops only that receiver. - `mute`/`unmute` can operate globally or per target. +- `reset-restore-token ` stops one fully streaming receiver, clears only + its saved Wayland portal restore token, and reconnects it on the same IP and + port. Pairing credentials and other streams are unchanged. The command rejects + targets sharing a capture group; disconnect those peers first so the old portal + source can be stopped before a replacement is authorized. - `pin` retains its historical command name, but submits whichever credential the daemon requests: an on-screen PIN or a configured password. It is targetless and therefore requires exactly one waiting receiver; use diff --git a/cmd/doubletake-ctl/main.go b/cmd/doubletake-ctl/main.go index f60f236..b132424 100644 --- a/cmd/doubletake-ctl/main.go +++ b/cmd/doubletake-ctl/main.go @@ -61,6 +61,12 @@ func main() { } else { resp, err = client.Disconnect() } + case "reset-restore-token": + if len(args) != 2 { + fmt.Fprintln(os.Stderr, "Usage: doubletake-ctl reset-restore-token ") + os.Exit(1) + } + resp, err = client.ResetRestoreToken(args[1]) case "mute": if len(args) >= 2 { resp, err = client.MuteTarget(args[1]) @@ -94,5 +100,5 @@ func main() { } func usage() { - fmt.Fprintf(os.Stderr, "Usage: doubletake-ctl [-socket path] [args]\n\nCommands:\n status Show daemon state and all active streams\n discover Discover AirPlay devices on the network\n devices List cached discovered devices\n connect [target] [PIN-or-password] Start mirroring (to target IP, or first free device)\n pin Submit pairing credentials for a waiting device\n disconnect [target] Stop mirroring (all streams, or only the given IP)\n mute [target] Mute mirrored audio (all streams, or only the given IP)\n unmute [target] Unmute mirrored audio (all streams, or only the given IP)\n\nFlags:\n -socket path Override daemon socket path (default: %s)\n", daemon.DefaultSocketPath()) + fmt.Fprintf(os.Stderr, "Usage: doubletake-ctl [-socket path] [args]\n\nCommands:\n status Show daemon state and all active streams\n discover Discover AirPlay devices on the network\n devices List cached discovered devices\n connect [target] [PIN-or-password] Start mirroring (to target IP, or first free device)\n pin Submit pairing credentials for a waiting device\n disconnect [target] Stop mirroring (all streams, or only the given IP)\n reset-restore-token Clear one Wayland restore token and reconnect that target\n mute [target] Mute mirrored audio (all streams, or only the given IP)\n unmute [target] Unmute mirrored audio (all streams, or only the given IP)\n\nFlags:\n -socket path Override daemon socket path (default: %s)\n", daemon.DefaultSocketPath()) } diff --git a/cmd/doubletake-ctl/main_test.go b/cmd/doubletake-ctl/main_test.go new file mode 100644 index 0000000..d288e2f --- /dev/null +++ b/cmd/doubletake-ctl/main_test.go @@ -0,0 +1,183 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "net" + "os" + "os/exec" + "path/filepath" + "testing" + "time" +) + +const cliHelperEnvironment = "DOUBLETAKE_CTL_TEST_HELPER" + +type controlFixtureResult struct { + request map[string]json.RawMessage + err error +} + +func TestResetRestoreTokenCLI(t *testing.T) { + listener, results := startControlFixture(t) + socketPath := listener.Addr().String() + + code, stdout, stderr := runCLIProcess(t, + "-socket", socketPath, + "reset-restore-token", "192.0.2.10", + ) + if code != 0 { + t.Fatalf("reset-restore-token exit code = %d, want 0; stdout=%q stderr=%q", code, stdout, stderr) + } + + result := awaitControlFixture(t, results) + if result.err != nil { + t.Fatalf("control fixture: %v", result.err) + } + if len(result.request) != 2 { + t.Fatalf("request fields = %v, want exactly cmd and target", result.request) + } + assertJSONStringField(t, result.request, "cmd", "reset-restore-token") + assertJSONStringField(t, result.request, "target", "192.0.2.10") + + if err := listener.Close(); err != nil { + t.Fatalf("close control fixture: %v", err) + } +} + +func TestResetRestoreTokenCLIRejectsInvalidArityWithoutContactingDaemon(t *testing.T) { + for _, test := range []struct { + name string + args []string + }{ + {name: "missing target", args: []string{"reset-restore-token"}}, + {name: "extra target", args: []string{"reset-restore-token", "192.0.2.10", "192.0.2.11"}}, + } { + t.Run(test.name, func(t *testing.T) { + listener, results := startControlFixture(t) + args := append([]string{"-socket", listener.Addr().String()}, test.args...) + + code, stdout, stderr := runCLIProcess(t, args...) + if code == 0 { + t.Fatalf("invalid reset exit code = 0, want nonzero; stdout=%q stderr=%q", stdout, stderr) + } + if err := listener.Close(); err != nil { + t.Fatalf("close control fixture: %v", err) + } + + result := awaitControlFixture(t, results) + if result.err == nil { + t.Fatalf("invalid reset contacted daemon with request %v", result.request) + } + if !errors.Is(result.err, net.ErrClosed) { + t.Fatalf("control fixture stopped with %v, want listener close", result.err) + } + }) + } +} + +func TestDoubletakeCtlHelperProcess(t *testing.T) { + if os.Getenv(cliHelperEnvironment) != "1" { + return + } + separator := -1 + for i, arg := range os.Args { + if arg == "--" { + separator = i + break + } + } + if separator < 0 { + os.Exit(2) + } + os.Args = append([]string{"doubletake-ctl"}, os.Args[separator+1:]...) + main() +} + +func startControlFixture(t *testing.T) (net.Listener, <-chan controlFixtureResult) { + t.Helper() + listener, err := net.Listen("unix", filepath.Join(t.TempDir(), "doubletake.sock")) + if err != nil { + t.Fatalf("listen on control socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + results := make(chan controlFixtureResult, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + results <- controlFixtureResult{err: err} + return + } + defer conn.Close() + + var request map[string]json.RawMessage + if err := json.NewDecoder(conn).Decode(&request); err != nil { + results <- controlFixtureResult{err: err} + return + } + if err := json.NewEncoder(conn).Encode(map[string]any{"ok": true, "state": "connecting"}); err != nil { + results <- controlFixtureResult{err: err} + return + } + results <- controlFixtureResult{request: request} + }() + return listener, results +} + +func runCLIProcess(t *testing.T, args ...string) (int, string, string) { + t.Helper() + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + + commandArgs := []string{"-test.run=^TestDoubletakeCtlHelperProcess$", "--"} + commandArgs = append(commandArgs, args...) + cmd := exec.CommandContext(ctx, os.Args[0], commandArgs...) + cmd.Env = append(os.Environ(), cliHelperEnvironment+"=1") + var stdout bytes.Buffer + var stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + err := cmd.Run() + if ctx.Err() != nil { + t.Fatalf("CLI process did not complete: %v", ctx.Err()) + } + if err == nil { + return 0, stdout.String(), stderr.String() + } + var exitErr *exec.ExitError + if !errors.As(err, &exitErr) { + t.Fatalf("run CLI process: %v", err) + } + return exitErr.ExitCode(), stdout.String(), stderr.String() +} + +func awaitControlFixture(t *testing.T, results <-chan controlFixtureResult) controlFixtureResult { + t.Helper() + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + select { + case result := <-results: + return result + case <-ctx.Done(): + t.Fatalf("control fixture did not finish: %v", ctx.Err()) + return controlFixtureResult{} + } +} + +func assertJSONStringField(t *testing.T, request map[string]json.RawMessage, field, want string) { + t.Helper() + raw, ok := request[field] + if !ok { + t.Fatalf("request is missing %q: %v", field, request) + } + var got string + if err := json.Unmarshal(raw, &got); err != nil { + t.Fatalf("decode request field %q: %v", field, err) + } + if got != want { + t.Fatalf("request field %q = %q, want %q", field, got, want) + } +} diff --git a/cmd/doubletake/main.go b/cmd/doubletake/main.go index b49aab6..772398e 100644 --- a/cmd/doubletake/main.go +++ b/cmd/doubletake/main.go @@ -371,6 +371,7 @@ func main() { } streamCfg.AutomaticHEVCAvailable = capturePreparation.AutomaticHEVCAvailable() streamCfg.MeasuredVideoLatency = capturePreparation.MeasuredVideoLatency() + streamCfg.MinimumVideoLead = capturePreparation.MinimumVideoLead() var capture *airplay.ScreenCapture var broadcast *airplay.BroadcastCapture diff --git a/internal/airplay/capture.go b/internal/airplay/capture.go index 5aaf46b..f36eab5 100644 --- a/internal/airplay/capture.go +++ b/internal/airplay/capture.go @@ -61,19 +61,25 @@ const ( // uses a fixed resolution. testCaptureWidth = 1920 testCaptureHeight = 1080 + + // The isolated Wayland encoder receives copied raw frames over a pipe. Its + // bounded copy/encode interval exceeds Apple's nominal 75 ms screen lead on + // the supported integrated-GPU path. + waylandRawRelayMinimumVideoLead = 250 * time.Millisecond ) // ScreenCapture manages screen capture via GStreamer. type ScreenCapture struct { - cmd *exec.Cmd // gst-launch-1.0 process - stdout io.ReadCloser - frames videoAccessUnitReader - cancel context.CancelFunc - pwNodeID uint32 - dbusConn *dbus.Conn // portal session D-Bus connection (must stay open for Wayland) - waitCh chan struct{} // closed when process exits - waitErr error // set before waitCh is closed - stopped bool + cmd *exec.Cmd // gst-launch-1.0 encoder process + sourceCmd *exec.Cmd // optional Wayland capture/serialization process + stdout io.ReadCloser + frames videoAccessUnitReader + cancel context.CancelFunc + pwNodeID uint32 + dbusConn *dbus.Conn // portal session D-Bus connection (must stay open for Wayland) + waitCh chan struct{} // closed when process exits + waitErr error // set before waitCh is closed + stopped bool } type capturePreparationKind uint8 @@ -84,6 +90,21 @@ const ( capturePreparationTest ) +func captureMinimumVideoLead(kind capturePreparationKind, measured time.Duration) time.Duration { + if kind == capturePreparationWayland && measured < waylandRawRelayMinimumVideoLead { + return waylandRawRelayMinimumVideoLead + } + return measured +} + +func waylandRawVideoSize(receiverWidth, receiverHeight int, streamSize [2]int) (int, int) { + if receiverWidth <= 0 || receiverHeight <= 0 { + receiverWidth, receiverHeight = streamSize[0], streamSize[1] + } + receiverWidth, receiverHeight = fitVideoSize(receiverWidth, receiverHeight, 3840, 2160) + return receiverWidth &^ 1, receiverHeight &^ 1 +} + // CapturePreparation performs the potentially interactive part of screen // capture before the receiver session starts. In particular, a Wayland // preparation completes the screencast portal request and retains its PipeWire @@ -101,9 +122,10 @@ type CapturePreparation struct { timestampedOutput bool automaticHEVCAvail bool - // measuredVideoLatency is the minimum screen lead measured by the local 4K - // HEVC preflight. It is zero for H.264 and unmeasured software fallback. + // measuredVideoLatency is the minimum screen lead required by local capture: + // the 4K HEVC preflight. minimumVideoLead also includes transport overhead. measuredVideoLatency time.Duration + minimumVideoLead time.Duration pwNodeID uint32 pwFd *os.File @@ -204,6 +226,7 @@ func PrepareCapture(ctx context.Context, cfg CaptureConfig) (*CapturePreparation preparation.pwNodeID = nodeID preparation.pwFd = pwFd preparation.dbusConn = dbusConn + preparation.minimumVideoLead = captureMinimumVideoLead(kind, preparation.measuredVideoLatency) return preparation, nil } @@ -391,6 +414,14 @@ func (p *CapturePreparation) MeasuredVideoLatency() time.Duration { return p.measuredVideoLatency } +// MinimumVideoLead returns the full local capture/transport scheduling floor. +func (p *CapturePreparation) MinimumVideoLead() time.Duration { + if p == nil { + return 0 + } + return p.minimumVideoLead +} + // Close releases an unconsumed portal preparation. Once Start has taken // ownership, ScreenCapture.Stop owns the corresponding resources. func (p *CapturePreparation) Close() { @@ -725,21 +756,37 @@ func frameIntervalMillis(fps int) int { return max(1, 1000/fps) } -func pipeWireVideoSourceStage(fd int, nodeID uint32, fps int) gstStage { - return gstStage{ +func pipeWireVideoSourceStage(fd int, nodeID uint32, fps int, copyPortalBuffers bool) gstStage { + stage := gstStage{ "pipewiresrc", fmt.Sprintf("fd=%d", fd), fmt.Sprintf("path=%d", nodeID), "do-timestamp=true", fmt.Sprintf("keepalive-time=%d", frameIntervalMillis(fps)), - // The compositor and pipewiresrc's keepalive path both retain the latest - // GstBuffer. With a small portal pool that can keep every PipeWire buffer - // checked out and freeze screencopy. Copying here returns the portal buffer - // as soon as pipewiresrc pulls it while downstream retains only the copy. - "always-copy=true", + } + if copyPortalBuffers { + // The software path cannot import a portal DMA-BUF through VA-API. Copy + // immediately so downstream never retains a PipeWire-owned buffer. + stage = append(stage, "always-copy=true") + } + return stage +} + +func vaapiVideoImportStages() []gstStage { + return []gstStage{ + {"vapostproc", "disable-passthrough=true"}, + {"video/x-raw,format=NV12"}, } } +func waylandVideoInputStages(fd int, nodeID uint32, fps int, useVAAPI bool) (gstStage, []gstStage) { + source := pipeWireVideoSourceStage(fd, nodeID, fps, !useVAAPI) + if !useVAAPI { + return source, nil + } + return source, vaapiVideoImportStages() +} + func lowLatencyVideoQueueStage() gstStage { // Drop stale raw frames before encoding. Encoded P-frames may reference // earlier frames, so dropping them downstream would corrupt the codec chain. @@ -801,6 +848,10 @@ func buildGstVideoPipeline(source gstStage, beforeConvert, afterScale []gstStage for _, stage := range afterScale { args = appendGstStage(args, stage) } + return appendGstEncoderPipeline(args, encoder, timestampedOutput) +} + +func appendGstEncoderPipeline(args []string, encoder encoderResult, timestampedOutput bool) []string { if encoder.needsVulkan { args = appendGstStage(args, gstStage{"vulkanupload"}) } @@ -827,6 +878,105 @@ func buildGstVideoPipeline(source gstStage, beforeConvert, afterScale []gstStage return appendGstStage(args, gstStage{"fdsink", "fd=1", "sync=false", "async=false"}) } +func buildSplitGstVideoPipeline(source gstStage, beforeConvert, afterScale []gstStage, encoder encoderResult, maxWidth, maxHeight, fps int, timestampedOutput bool) (producer, consumer []string) { + producer = append([]string{"--quiet"}, source...) + for _, stage := range beforeConvert { + producer = appendGstStage(producer, stage) + } + producer = appendGstStage(producer, gstStage{"videoconvert"}) + producer = appendGstStage(producer, gstStage{fmt.Sprintf("video/x-raw,format=%s", encoder.rawFormat)}) + for _, stage := range receiverScaleStages(maxWidth, maxHeight) { + producer = appendGstStage(producer, stage) + } + for _, stage := range afterScale { + producer = appendGstStage(producer, stage) + } + producer = appendGstStage(producer, gstStage{"fdsink", "fd=1", "sync=false", "async=false"}) + + maxWidth &^= 1 + maxHeight &^= 1 + consumer = []string{"--quiet", "fdsrc", "fd=0", "do-timestamp=true"} + consumer = appendGstStage(consumer, gstStage{ + fmt.Sprintf( + "video/x-raw,format=%s,width=%d,height=%d,framerate=%d/1", + encoder.rawFormat, maxWidth, maxHeight, fps, + ), + }) + consumer = appendGstStage(consumer, gstStage{"rawvideoparse", "use-sink-caps=true"}) + consumer = appendGstEncoderPipeline(consumer, encoder, timestampedOutput) + return producer, consumer +} + +type waylandSplitPipes struct { + rawFrames *os.File + sourceOutput *os.File + sourceStderr *os.File + sourceErrOut *os.File + stdout *os.File + encoderOutput *os.File + encoderStderr *os.File + encoderErrOut *os.File +} + +func openWaylandSplitPipes(sourceCmd, encoderCmd *exec.Cmd) (*waylandSplitPipes, error) { + pipes := &waylandSplitPipes{} + var err error + pipes.rawFrames, pipes.sourceOutput, err = os.Pipe() + if err != nil { + return nil, fmt.Errorf("capture serialization pipe: %w", err) + } + sourceCmd.Stdout = pipes.sourceOutput + pipes.sourceStderr, pipes.sourceErrOut, err = os.Pipe() + if err != nil { + pipes.close() + return nil, fmt.Errorf("capture stderr pipe: %w", err) + } + sourceCmd.Stderr = pipes.sourceErrOut + encoderCmd.Stdin = pipes.rawFrames + pipes.stdout, pipes.encoderOutput, err = os.Pipe() + if err != nil { + pipes.close() + return nil, fmt.Errorf("encoder stdout pipe: %w", err) + } + encoderCmd.Stdout = pipes.encoderOutput + pipes.encoderStderr, pipes.encoderErrOut, err = os.Pipe() + if err != nil { + pipes.close() + return nil, fmt.Errorf("encoder stderr pipe: %w", err) + } + encoderCmd.Stderr = pipes.encoderErrOut + return pipes, nil +} + +func (pipes *waylandSplitPipes) close() { + if pipes == nil { + return + } + for _, closer := range []io.Closer{ + pipes.rawFrames, + pipes.sourceOutput, + pipes.sourceStderr, + pipes.sourceErrOut, + pipes.stdout, + pipes.encoderOutput, + pipes.encoderStderr, + pipes.encoderErrOut, + } { + if closer != nil { + _ = closer.Close() + } + } +} + +func startWaylandEncoder(cmd *exec.Cmd, pipes *waylandSplitPipes) (<-chan error, error) { + wait, err := startGStreamerCommand(cmd) + if err != nil { + pipes.close() + return nil, err + } + return wait, nil +} + func startPreparedWaylandCapture(ctx context.Context, cfg CaptureConfig, encoderParts encoderResult, nodeID uint32, pwFd *os.File, dbusConn *dbus.Conn, streamSize [2]int, timestampedOutput bool) (*ScreenCapture, error) { if pwFd == nil || dbusConn == nil { if pwFd != nil { @@ -853,14 +1003,14 @@ func startPreparedWaylandCapture(ctx context.Context, cfg CaptureConfig, encoder // The encoded dimensions are capped to the receiver's advertised display size // when available. The actual result is read back from the codec SPS downstream. const pwFdNum = 3 - source := pipeWireVideoSourceStage(pwFdNum, nodeID, fps) - hasCompositor := streamSize[0] > 0 && streamSize[1] > 0 && hasGstElement("compositor") - - var beforeConvert []gstStage - if hasGstElement("vapostproc") { - beforeConvert = append(beforeConvert, gstStage{"vapostproc"}) - } else { + hasVAAPIPostproc := hasGstElement("vapostproc") + // VA-API needs the portal's original DMA-BUF. pipewiresrc's always-copy path + // can turn DMA-BUF map failures into black fallback frames before vapostproc + // gets a chance to import them. The software path still copies immediately + // so a forced-live compositor cannot exhaust a small portal buffer pool. + source, beforeConvert := waylandVideoInputStages(pwFdNum, nodeID, fps, hasVAAPIPostproc) + if !hasVAAPIPostproc { log.Printf("[CAPTURE] vapostproc unavailable, using software conversion") } @@ -869,58 +1019,102 @@ func startPreparedWaylandCapture(ctx context.Context, cfg CaptureConfig, encoder beforeConvert = append(beforeConvert, gstStage{"compositor", "force-live=true", "ignore-inactive-pads=true", "background=black"}, gstStage{fmt.Sprintf("video/x-raw,width=%d,height=%d,framerate=%d/1", streamSize[0], streamSize[1], fps)}, + lowLatencyVideoQueueStage(), ) } else { log.Printf("[CAPTURE] idle-frame compositor unavailable; using portal frame timing") } - if hasCompositor { - afterScale = append(afterScale, lowLatencyVideoQueueStage()) - } else { + if !hasCompositor { afterScale = append(afterScale, gstStage{"videorate", "drop-only=true", "skip-to-first=true"}, frameRateStage(fps), lowLatencyVideoQueueStage(), ) } - gstArgs := buildGstVideoPipeline(source, beforeConvert, afterScale, encoderParts, cfg.MaxWidth, cfg.MaxHeight, timestampedOutput) - - dbg("[CAPTURE] gst-launch-1.0 (wayland) %s", strings.Join(gstArgs, " ")) - cmd := exec.CommandContext(captureCtx, "gst-launch-1.0", gstArgs...) - cmd.ExtraFiles = []*os.File{pwFd} - - stdout, err := cmd.StdoutPipe() + rawWidth, rawHeight := waylandRawVideoSize(cfg.MaxWidth, cfg.MaxHeight, streamSize) + if rawWidth <= 0 || rawHeight <= 0 { + cancel() + _ = pwFd.Close() + _ = dbusConn.Close() + return nil, fmt.Errorf("Wayland capture is missing both receiver and portal dimensions") + } + sourceArgs, encoderArgs := buildSplitGstVideoPipeline( + source, beforeConvert, afterScale, encoderParts, + rawWidth, rawHeight, fps, timestampedOutput, + ) + dbg("[CAPTURE] gst-launch-1.0 (wayland source) %s", strings.Join(sourceArgs, " ")) + dbg("[CAPTURE] gst-launch-1.0 (wayland encoder) %s", strings.Join(encoderArgs, " ")) + + sourceCmd := exec.CommandContext(captureCtx, "gst-launch-1.0", sourceArgs...) + sourceCmd.ExtraFiles = []*os.File{pwFd} + cmd := exec.CommandContext(captureCtx, "gst-launch-1.0", encoderArgs...) + pipes, err := openWaylandSplitPipes(sourceCmd, cmd) if err != nil { cancel() - pwFd.Close() - dbusConn.Close() - return nil, fmt.Errorf("gst stdout pipe: %w", err) + _ = pwFd.Close() + _ = dbusConn.Close() + return nil, err } - stderr, _ := cmd.StderrPipe() - waitResult, err := startGStreamerCommand(cmd) + encoderWait, err := startWaylandEncoder(cmd, pipes) if err != nil { cancel() - pwFd.Close() - dbusConn.Close() - return nil, fmt.Errorf("start gst-launch: %w", err) + _ = pwFd.Close() + _ = dbusConn.Close() + return nil, fmt.Errorf("start encoder gst-launch: %w", err) + } + _ = pipes.rawFrames.Close() + _ = pipes.encoderOutput.Close() + _ = pipes.encoderErrOut.Close() + sourceWait, err := startGStreamerCommand(sourceCmd) + if err != nil { + cancel() + pipes.close() + _ = pwFd.Close() + _ = dbusConn.Close() + <-encoderWait + return nil, fmt.Errorf("start capture gst-launch: %w", err) } - pwFd.Close() // child inherited it + _ = pipes.sourceOutput.Close() + _ = pipes.sourceErrOut.Close() + _ = pwFd.Close() // source child inherited it - go logStderr("GST", stderr) + go func() { + defer pipes.sourceStderr.Close() + logStderr("GST-SOURCE", pipes.sourceStderr) + }() + go func() { + defer pipes.encoderStderr.Close() + logStderr("GST-ENCODER", pipes.encoderStderr) + }() capture := &ScreenCapture{ - cmd: cmd, - stdout: stdout, - cancel: cancel, - pwNodeID: nodeID, - dbusConn: dbusConn, - waitCh: make(chan struct{}), + cmd: cmd, + sourceCmd: sourceCmd, + stdout: pipes.stdout, + cancel: cancel, + pwNodeID: nodeID, + dbusConn: dbusConn, + waitCh: make(chan struct{}), } if timestampedOutput { - capture.frames = newRTPVideoAccessUnitReader(stdout, encoderParts.codec) + capture.frames = newRTPVideoAccessUnitReader(pipes.stdout, encoderParts.codec) } go func() { - capture.waitErr = <-waitResult + type processResult struct { + name string + err error + } + results := make(chan processResult, 2) + go func() { results <- processResult{name: "capture", err: <-sourceWait} }() + go func() { results <- processResult{name: "encoder", err: <-encoderWait} }() + first := <-results + dbg("[CAPTURE] %s pipeline exited: %v", first.name, first.err) + cancel() + <-results + if first.err != nil { + capture.waitErr = fmt.Errorf("%s pipeline: %w", first.name, first.err) + } close(capture.waitCh) }() @@ -1056,6 +1250,9 @@ func (sc *ScreenCapture) Stop() { if sc.cmd != nil && sc.cmd.Process != nil { _ = sc.cmd.Process.Signal(os.Interrupt) } + if sc.sourceCmd != nil && sc.sourceCmd.Process != nil { + _ = sc.sourceCmd.Process.Signal(os.Interrupt) + } select { case <-sc.waitCh: @@ -1063,6 +1260,9 @@ func (sc *ScreenCapture) Stop() { if sc.cmd != nil && sc.cmd.Process != nil { _ = sc.cmd.Process.Kill() } + if sc.sourceCmd != nil && sc.sourceCmd.Process != nil { + _ = sc.sourceCmd.Process.Kill() + } <-sc.waitCh } } diff --git a/internal/airplay/capture_broadcast.go b/internal/airplay/capture_broadcast.go index 2177cab..c02f5e9 100644 --- a/internal/airplay/capture_broadcast.go +++ b/internal/airplay/capture_broadcast.go @@ -19,13 +19,21 @@ const ( // metadata; it is deliberately generous for normal H.264 buffer cadence. broadcastSinkQueueChunks = 4096 + // Encoded frames are detached from the portal-owned raw-buffer pool. Keep + // enough of them to absorb a receiver's ordinary two-second socket-write + // stall without terminating the session; the independent byte limit still + // bounds memory and a persistently stalled peer is removed. + broadcastSinkQueueDuration = 2 * time.Second + // Source shutdown normally races with consumers draining their last few // buffers. Do not let a receiver that stopped reading keep Run alive forever. broadcastSinkDrainTimeout = 2 * time.Second ) -var errBroadcastSinkBacklog = errors.New("broadcast sink backlog limit exceeded") -var errBroadcastSinkMode = errors.New("backpressured broadcast sink requires an otherwise unused capture") +var ( + errBroadcastSinkBacklog = errors.New("broadcast sink backlog limit exceeded") + errBroadcastSinkMode = errors.New("backpressured broadcast sink requires an otherwise unused capture") +) // BroadcastCapture reads from a single ScreenCapture and fans the raw byte // stream out to multiple registered sinks. Each sink has an independent, @@ -55,7 +63,6 @@ type BroadcastCapture struct { // following sequence, which gives attachment an exact cutover even when a // source read has completed but has not yet been fanned out. sequence uint64 - drainTimeout time.Duration sinks []*BroadcastSink @@ -81,10 +88,9 @@ type BroadcastSink struct { maxQueuedBytes int maxQueuedChunks int - // Apple's ordinary virtual-display source bounds its upstream frame queue to - // 67 ms and drops an incoming source frame at that limit. Doubletake derives - // a downstream encoded-relay ceiling from that value and counts configured - // sample durations. The byte and chunk limits remain independent safeguards. + // This encoded system-memory queue is independent of Apple's upstream raw + // frame queue. It counts nominal sample durations so PTS gaps caused by raw + // frame dropping do not spuriously detach a healthy receiver. maxFrameQueueDuration time.Duration backpressure bool blockedProducers int // number waiting for queue handoff; guarded by mu @@ -105,7 +111,7 @@ func newBroadcastSinkWithPolicy(owner *BroadcastCapture, backpressure bool) *Bro owner: owner, maxQueuedBytes: broadcastSinkQueueBytes, maxQueuedChunks: broadcastSinkQueueChunks, - maxFrameQueueDuration: ordinaryScreenFrameQueueDuration, + maxFrameQueueDuration: broadcastSinkQueueDuration, backpressure: backpressure, frameDuration: frameDuration, done: make(chan struct{}), diff --git a/internal/airplay/capture_broadcast_test.go b/internal/airplay/capture_broadcast_test.go index 616b3f5..4028b0f 100644 --- a/internal/airplay/capture_broadcast_test.go +++ b/internal/airplay/capture_broadcast_test.go @@ -192,15 +192,16 @@ func TestBroadcastSinkNonblockingFrameQueueUsesNominalDuration(t *testing.T) { base := time.Now() // A large or backward PTS gap can be caused by the upstream leaky queue; it // must not turn one queued picture into an artificial duration overflow. - for i, offset := range []time.Duration{0, time.Second} { + for i := 0; i < 60; i++ { + offset := time.Duration(i%2) * time.Second frame := VideoAccessUnit{AnnexB: []byte{byte(i + 1)}, PTS: base.Add(offset)} if err := sink.enqueueFrame(frame); err != nil { t.Fatalf("enqueue frame %d at %v: %v", i, offset, err) } } - third := VideoAccessUnit{AnnexB: []byte{3}, PTS: base.Add(-time.Second)} - if err := sink.enqueueFrame(third); !errors.Is(err, errBroadcastSinkBacklog) { - t.Fatalf("enqueue third nominal 30fps frame = %v, want backlog error", err) + overflow := VideoAccessUnit{AnnexB: []byte{0xff}, PTS: base.Add(-time.Second)} + if err := sink.enqueueFrame(overflow); !errors.Is(err, errBroadcastSinkBacklog) { + t.Fatalf("enqueue frame after 2s nominal 30fps queue = %v, want backlog error", err) } } @@ -210,8 +211,8 @@ func TestBroadcastSinkNominalDurationUsesConfiguredFrameRate(t *testing.T) { acceptedFrames int rejectedOrdinal int }{ - {fps: 20, acceptedFrames: 1, rejectedOrdinal: 2}, - {fps: 60, acceptedFrames: 4, rejectedOrdinal: 5}, + {fps: 20, acceptedFrames: 40, rejectedOrdinal: 41}, + {fps: 60, acceptedFrames: 120, rejectedOrdinal: 121}, } { t.Run(fmt.Sprintf("%dfps", test.fps), func(t *testing.T) { broadcast := NewBroadcastCaptureWithFrameRate(nil, test.fps) @@ -363,7 +364,7 @@ func TestBackpressuredByteBroadcastHandoff(t *testing.T) { func TestLoneSharedTimestampedSinkDoesNotBackpressureCapture(t *testing.T) { base := time.Now() - frames := make([]VideoAccessUnit, 4) + frames := make([]VideoAccessUnit, 62) for i := range frames { frames[i] = VideoAccessUnit{ AnnexB: []byte{byte(i + 1)}, @@ -396,6 +397,20 @@ func TestLoneSharedTimestampedSinkDoesNotBackpressureCapture(t *testing.T) { } } +func TestSharedTimestampedSinkToleratesTransientNetworkStall(t *testing.T) { + sink := newBroadcastSinkWithPolicy(nil, false) + base := time.Now() + for i := 0; i < 15; i++ { + frame := VideoAccessUnit{ + AnnexB: []byte{byte(i + 1)}, + PTS: base.Add(time.Duration(i) * time.Second / 30), + } + if err := sink.enqueueFrame(frame); err != nil { + t.Fatalf("enqueue frame %d during 500ms network stall: %v", i, err) + } + } +} + func TestTimestampedSlowSinkDoesNotStallHealthyPeer(t *testing.T) { frames := make(chan VideoAccessUnit) capture := &ScreenCapture{ @@ -424,7 +439,7 @@ func TestTimestampedSlowSinkDoesNotStallHealthyPeer(t *testing.T) { }() base := time.Now() - for i := 0; i < 6; i++ { + for i := 0; i < 62; i++ { want := VideoAccessUnit{ AnnexB: []byte{byte(i + 1)}, PTS: base.Add(time.Duration(i) * time.Second / 30), diff --git a/internal/airplay/capture_test.go b/internal/airplay/capture_test.go index 7f3ac73..b27218f 100644 --- a/internal/airplay/capture_test.go +++ b/internal/airplay/capture_test.go @@ -1,6 +1,7 @@ package airplay import ( + "bytes" "context" "errors" "os" @@ -14,6 +15,71 @@ import ( "github.com/godbus/dbus/v5" ) +func TestRawVideoRelayCapsNegotiateWithInstalledParser(t *testing.T) { + if _, err := exec.LookPath("gst-launch-1.0"); err != nil { + t.Skip("gst-launch-1.0 is unavailable") + } + const width, height = 4, 2 + args := []string{ + "--quiet", + "fdsrc", "fd=0", + "!", "video/x-raw,format=NV12,width=4,height=2,framerate=30/1", + "!", "rawvideoparse", "use-sink-caps=true", + "!", "fakesink", + } + cmd := exec.Command("gst-launch-1.0", args...) + cmd.Stdin = bytes.NewReader(make([]byte, width*height*3/2)) + if output, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("raw relay caps failed real GStreamer negotiation: %v\n%s", err, output) + } +} + +func TestWaylandEncoderStartFailureClosesEveryPipeDescriptor(t *testing.T) { + before := openDescriptorCount(t) + for i := 0; i < 20; i++ { + source := exec.Command("true") + encoder := exec.Command("/definitely/not/a/doubletake-test-command") + pipes, err := openWaylandSplitPipes(source, encoder) + if err != nil { + t.Fatalf("openWaylandSplitPipes: %v", err) + } + if _, err := startWaylandEncoder(encoder, pipes); err == nil { + t.Fatal("startWaylandEncoder unexpectedly succeeded") + } + } + after := openDescriptorCount(t) + if after > before { + t.Fatalf("stderr setup failures leaked descriptors: before=%d after=%d", before, after) + } +} + +func TestOpenWaylandSplitPipesCloseReleasesEveryDescriptor(t *testing.T) { + before := openDescriptorCount(t) + source := exec.Command("true") + encoder := exec.Command("true") + pipes, err := openWaylandSplitPipes(source, encoder) + if err != nil { + t.Fatalf("openWaylandSplitPipes: %v", err) + } + pipes.close() + after := openDescriptorCount(t) + if after > before { + t.Fatalf("pipe close leaked descriptors: before=%d after=%d", before, after) + } +} + +func openDescriptorCount(t *testing.T) int { + t.Helper() + entries, err := os.ReadDir("/proc/self/fd") + if errors.Is(err, os.ErrNotExist) { + t.Skip("/proc/self/fd is unavailable") + } + if err != nil { + t.Fatalf("read /proc/self/fd: %v", err) + } + return len(entries) +} + func TestCapturePreparationCloseReleasesUnstartedResources(t *testing.T) { portalFD, peerFD, err := os.Pipe() if err != nil { @@ -138,6 +204,30 @@ func TestRecommendedAutomaticVideoLatencyUsesP95AndDeliveryMargin(t *testing.T) } } +func TestCaptureMinimumVideoLeadIncludesWaylandRawRelay(t *testing.T) { + const measured = 300 * time.Millisecond + if got := captureMinimumVideoLead(capturePreparationWayland, 0); got != 250*time.Millisecond { + t.Fatalf("unmeasured Wayland split lead = %v, want 250ms", got) + } + if got := captureMinimumVideoLead(capturePreparationWayland, measured); got != measured { + t.Fatalf("larger measured Wayland lead = %v, want %v", got, measured) + } + if got := captureMinimumVideoLead(capturePreparationX11, 0); got != 0 { + t.Fatalf("X11 minimum lead = %v, want no split-pipeline override", got) + } +} + +func TestWaylandRawVideoSizeBoundsReceiverCanvas(t *testing.T) { + width, height := waylandRawVideoSize(1<<30, 1<<30, [2]int{2880, 1800}) + if width != 2160 || height != 2160 { + t.Fatalf("hostile square receiver canvas = %dx%d, want bounded 2160x2160", width, height) + } + width, height = waylandRawVideoSize(0, 0, [2]int{2881, 1801}) + if width != 2880 || height != 1800 { + t.Fatalf("portal fallback canvas = %dx%d, want even 2880x1800", width, height) + } +} + func TestLiveVideoProbeTimeoutTracksConfiguredFrameRate(t *testing.T) { if got := liveVideoProbeTimeout(30); got != minimumLiveVideoProbeTimeout { t.Fatalf("30fps live probe timeout = %v, want %v", got, minimumLiveVideoProbeTimeout) @@ -361,8 +451,22 @@ func TestFrameIntervalMillis(t *testing.T) { } } -func TestPipeWireVideoSourceCopiesPortalBuffers(t *testing.T) { - got := pipeWireVideoSourceStage(3, 42, 30) +func TestPipeWireVideoSourcePreservesDMABuffersForVAAPI(t *testing.T) { + got := pipeWireVideoSourceStage(3, 42, 30, false) + want := gstStage{ + "pipewiresrc", + "fd=3", + "path=42", + "do-timestamp=true", + "keepalive-time=33", + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("PipeWire VA-API source stage = %v, want %v", got, want) + } +} + +func TestPipeWireVideoSourceCopiesPortalBuffersForSoftwareConversion(t *testing.T) { + got := pipeWireVideoSourceStage(3, 42, 30, true) want := gstStage{ "pipewiresrc", "fd=3", @@ -372,7 +476,52 @@ func TestPipeWireVideoSourceCopiesPortalBuffers(t *testing.T) { "always-copy=true", } if !reflect.DeepEqual(got, want) { - t.Fatalf("PipeWire source stage = %v, want %v", got, want) + t.Fatalf("PipeWire software source stage = %v, want %v", got, want) + } +} + +func TestVAAPIPostprocReceivesOriginalPortalDMABuffer(t *testing.T) { + got := vaapiVideoImportStages() + want := []gstStage{ + {"vapostproc", "disable-passthrough=true"}, + {"video/x-raw,format=NV12"}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("VA-API import stages = %v, want %v", got, want) + } +} + +func TestWaylandVideoInputStagesPreservePortalBufferOwnership(t *testing.T) { + for _, tt := range []struct { + name string + useVAAPI bool + wantSource gstStage + wantImports []gstStage + }{ + { + name: "VA-API imports original portal buffer", + useVAAPI: true, + wantSource: gstStage{"pipewiresrc", "fd=3", "path=42", "do-timestamp=true", "keepalive-time=33"}, + wantImports: []gstStage{ + {"vapostproc", "disable-passthrough=true"}, + {"video/x-raw,format=NV12"}, + }, + }, + { + name: "software conversion copies portal buffer", + useVAAPI: false, + wantSource: gstStage{"pipewiresrc", "fd=3", "path=42", "do-timestamp=true", "keepalive-time=33", "always-copy=true"}, + }, + } { + t.Run(tt.name, func(t *testing.T) { + gotSource, gotImports := waylandVideoInputStages(3, 42, 30, tt.useVAAPI) + if !reflect.DeepEqual(gotSource, tt.wantSource) { + t.Fatalf("source stage = %v, want %v", gotSource, tt.wantSource) + } + if !reflect.DeepEqual(gotImports, tt.wantImports) { + t.Fatalf("import stages = %v, want %v", gotImports, tt.wantImports) + } + }) } } @@ -451,6 +600,53 @@ func TestBuildGstVideoPipeline(t *testing.T) { } } +func TestBuildSplitWaylandVideoPipelineSerializesRawFramesBeforeEncoding(t *testing.T) { + encoder := encoderResult{ + parts: gstStage{"testh264enc", "bitrate=2500"}, + rawFormat: "NV12", + } + producer, consumer := buildSplitGstVideoPipeline( + gstStage{"pipewiresrc", "fd=3", "path=42"}, + []gstStage{ + {"vapostproc"}, + {"compositor", "force-live=true"}, + {"video/x-raw,width=2880,height=1800,framerate=30/1"}, + lowLatencyVideoQueueStage(), + }, + nil, + encoder, + 1280, + 720, + 30, + true, + ) + wantProducerSuffix := []string{ + "!", "video/x-raw,width=1280,height=720,pixel-aspect-ratio=1/1", + "!", "fdsink", "fd=1", "sync=false", "async=false", + } + if got := producer[len(producer)-len(wantProducerSuffix):]; !reflect.DeepEqual(got, wantProducerSuffix) { + t.Fatalf("raw producer suffix = %v, want %v", got, wantProducerSuffix) + } + wantConsumerPrefix := []string{ + "--quiet", "fdsrc", "fd=0", "do-timestamp=true", + "!", "video/x-raw,format=NV12,width=1280,height=720,framerate=30/1", + "!", "rawvideoparse", "use-sink-caps=true", + "!", "testh264enc", "bitrate=2500", + } + if got := consumer[:len(wantConsumerPrefix)]; !reflect.DeepEqual(got, wantConsumerPrefix) { + t.Fatalf("encoder consumer prefix = %v, want %v", got, wantConsumerPrefix) + } + wantConsumerSuffix := []string{ + "!", "rtph264pay", "pt=96", "mtu=60000", "aggregate-mode=none", "timestamp-offset=0", "seqnum-offset=0", + "!", "rtponviftimestamp", "ntp-offset=-1", "set-e-bit=false", "set-t-bit=false", + "!", "rtpstreampay", + "!", "fdsink", "fd=1", "sync=false", "async=false", + } + if got := consumer[len(consumer)-len(wantConsumerSuffix):]; !reflect.DeepEqual(got, wantConsumerSuffix) { + t.Fatalf("timestamped consumer suffix = %v, want %v", got, wantConsumerSuffix) + } +} + func TestBuildGstVideoPipelineTimestampedOutput(t *testing.T) { encoder := encoderResult{parts: gstStage{"testh264enc"}, rawFormat: "I420"} pipeline := buildGstVideoPipeline( diff --git a/internal/airplay/client.go b/internal/airplay/client.go index cd35c0e..0a294c7 100644 --- a/internal/airplay/client.go +++ b/internal/airplay/client.go @@ -1203,6 +1203,7 @@ type StreamConfig struct { VideoCodec VideoCodec // empty/h264, auto, or capability-gated hevc AutomaticHEVCAvailable bool // capture preflight found the hardware HEVC-4K path MeasuredVideoLatency time.Duration // measured minimum lead for the local HEVC capture path + MinimumVideoLead time.Duration // known capture/transport lead before codec selection NoEncrypt bool // Disable encryption for debugging DirectKey bool // Use shk/shiv directly without SHA-512 derivation NoAudio bool // Disable audio streaming diff --git a/internal/airplay/credentials.go b/internal/airplay/credentials.go index c07a8ce..1fe685f 100644 --- a/internal/airplay/credentials.go +++ b/internal/airplay/credentials.go @@ -103,6 +103,16 @@ type CredentialStore struct { backend CredentialBackend } +// RestoreTokenReset is a one-shot restore-token deletion which can be rolled +// back if the daemon loses its stream reservation before reconnecting. +type RestoreTokenReset struct { + store *CredentialStore + deviceID string + previousToken string + changed bool + done bool +} + // NewCredentialStore creates a credential store backed by a JSON file at path. func NewCredentialStore(path string) (*CredentialStore, error) { fb, err := newFileBackend(path) @@ -193,6 +203,81 @@ func (cs *CredentialStore) SaveRestoreToken(deviceID, restoreToken string) error return cs.backend.Save(deviceID, creds) } +// ClearRestoreToken removes only the Wayland screencast restore token for a +// device. Pairing credentials and all other device entries are preserved. +func (cs *CredentialStore) ClearRestoreToken(deviceID string) error { + reset, err := cs.BeginRestoreTokenReset(deviceID) + if err != nil { + return err + } + reset.Commit() + return nil +} + +// BeginRestoreTokenReset clears the restore token while retaining enough state +// to restore it if the caller cannot commit the corresponding reconnect. +func (cs *CredentialStore) BeginRestoreTokenReset(deviceID string) (*RestoreTokenReset, error) { + cs.mu.Lock() + defer cs.mu.Unlock() + + reset := &RestoreTokenReset{store: cs, deviceID: deviceID} + creds, err := cs.backend.Lookup(deviceID) + if err != nil { + return nil, err + } + if creds == nil || creds.RestoreToken == "" { + return reset, nil + } + reset.previousToken = creds.RestoreToken + reset.changed = true + updated := *creds + updated.RestoreToken = "" + if err := cs.backend.Save(deviceID, &updated); err != nil { + return nil, err + } + return reset, nil +} + +// Commit makes a successful deletion permanent. +func (r *RestoreTokenReset) Commit() { + if r == nil { + return + } + r.store.mu.Lock() + r.done = true + r.store.mu.Unlock() +} + +// Rollback restores only the previous token onto a fresh credential snapshot. +// A concurrently saved non-empty token wins, and all other fields are retained. +func (r *RestoreTokenReset) Rollback() error { + if r == nil { + return nil + } + r.store.mu.Lock() + defer r.store.mu.Unlock() + if r.done || !r.changed { + r.done = true + return nil + } + creds, err := r.store.backend.Lookup(r.deviceID) + if err != nil { + return err + } + if creds == nil { + creds = &SavedCredentials{} + } + if creds.RestoreToken == "" { + updated := *creds + updated.RestoreToken = r.previousToken + if err := r.store.backend.Save(r.deviceID, &updated); err != nil { + return err + } + } + r.done = true + return nil +} + // fileBackend stores credentials as a JSON file on disk. type fileBackend struct { path string @@ -222,8 +307,17 @@ func (fb *fileBackend) Lookup(deviceID string) (*SavedCredentials, error) { } func (fb *fileBackend) Save(deviceID string, creds *SavedCredentials) error { + previous, existed := fb.devices[deviceID] fb.devices[deviceID] = creds - return fb.persist() + if err := fb.persist(); err != nil { + if existed { + fb.devices[deviceID] = previous + } else { + delete(fb.devices, deviceID) + } + return err + } + return nil } func (fb *fileBackend) persist() error { diff --git a/internal/airplay/credentials_clear_test.go b/internal/airplay/credentials_clear_test.go new file mode 100644 index 0000000..53d1ceb --- /dev/null +++ b/internal/airplay/credentials_clear_test.go @@ -0,0 +1,157 @@ +package airplay + +import ( + "crypto/ed25519" + "crypto/rand" + "os" + "path/filepath" + "testing" +) + +func TestCredentialStoreClearRestoreTokenPreservesPairingAndOtherDevices(t *testing.T) { + path := filepath.Join(t.TempDir(), "credentials.json") + store, err := NewCredentialStore(path) + if err != nil { + t.Fatalf("NewCredentialStore: %v", err) + } + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatalf("GenerateKey: %v", err) + } + if err := store.SavePairing("device-1", "pair-1", pub, priv, PairingProtocolHAP); err != nil { + t.Fatalf("SavePairing: %v", err) + } + if err := store.SaveRestoreToken("device-1", "restore-1"); err != nil { + t.Fatalf("SaveRestoreToken device-1: %v", err) + } + if err := store.SaveRestoreToken("device-2", "restore-2"); err != nil { + t.Fatalf("SaveRestoreToken device-2: %v", err) + } + + if err := store.ClearRestoreToken("device-1"); err != nil { + t.Fatalf("ClearRestoreToken: %v", err) + } + + reloaded, err := NewCredentialStore(path) + if err != nil { + t.Fatalf("reload credential store: %v", err) + } + cleared := reloaded.Lookup("device-1") + if cleared == nil || !cleared.HasPairingCredentials() { + t.Fatal("clearing the restore token removed pairing credentials") + } + if cleared.PairingID != "pair-1" || cleared.PairingProtocol != PairingProtocolHAP { + t.Fatalf("pairing metadata changed: %+v", cleared) + } + if cleared.RestoreToken != "" { + t.Fatalf("restore token = %q, want empty", cleared.RestoreToken) + } + other := reloaded.Lookup("device-2") + if other == nil || other.RestoreToken != "restore-2" { + t.Fatalf("other device changed: %+v", other) + } +} + +type recordingCredentialBackend struct { + devices map[string]*SavedCredentials + saves []string +} + +func (b *recordingCredentialBackend) Lookup(deviceID string) (*SavedCredentials, error) { + return b.devices[deviceID], nil +} + +func (b *recordingCredentialBackend) Save(deviceID string, creds *SavedCredentials) error { + b.devices[deviceID] = creds + b.saves = append(b.saves, deviceID) + return nil +} + +func TestCredentialStoreClearRestoreTokenRollsBackFileBackendAfterPersistFailure(t *testing.T) { + path := filepath.Join(t.TempDir(), "credentials.json") + store, err := NewCredentialStore(path) + if err != nil { + t.Fatalf("NewCredentialStore: %v", err) + } + if err := store.SaveRestoreToken("device-1", "restore-1"); err != nil { + t.Fatalf("SaveRestoreToken: %v", err) + } + + blockedParent := filepath.Join(t.TempDir(), "not-a-directory") + if err := os.WriteFile(blockedParent, []byte("block mkdir"), 0600); err != nil { + t.Fatalf("create blocked parent: %v", err) + } + backend := store.backend.(*fileBackend) + backend.path = filepath.Join(blockedParent, "credentials.json") + + if err := store.ClearRestoreToken("device-1"); err == nil { + t.Fatal("ClearRestoreToken unexpectedly succeeded with an invalid credential path") + } + creds := store.Lookup("device-1") + if creds == nil || creds.RestoreToken != "restore-1" { + t.Fatalf("failed persistence changed in-memory credentials: %+v", creds) + } +} + +func TestCredentialStoreClearRestoreTokenUsesBackendWithoutDeletingEntry(t *testing.T) { + backend := &recordingCredentialBackend{devices: map[string]*SavedCredentials{ + "device-1": {PairingID: "pair-1", RestoreToken: "restore-1"}, + }} + store := NewCredentialStoreWithBackend(backend) + + if err := store.ClearRestoreToken("device-1"); err != nil { + t.Fatalf("ClearRestoreToken: %v", err) + } + + if len(backend.saves) != 1 || backend.saves[0] != "device-1" { + t.Fatalf("backend saves = %v, want [device-1]", backend.saves) + } + creds := backend.devices["device-1"] + if creds == nil || creds.PairingID != "pair-1" || creds.RestoreToken != "" { + t.Fatalf("backend credentials after clear = %+v", creds) + } +} + +func TestRestoreTokenResetRollbackMergesConcurrentCredentialChanges(t *testing.T) { + backend := &recordingCredentialBackend{devices: map[string]*SavedCredentials{ + "device-1": {PairingID: "pair-1", RestoreToken: "restore-1"}, + }} + store := NewCredentialStoreWithBackend(backend) + + reset, err := store.BeginRestoreTokenReset("device-1") + if err != nil { + t.Fatalf("BeginRestoreTokenReset: %v", err) + } + backend.devices["device-1"] = &SavedCredentials{ + PairingID: "pair-2", + RestoreToken: "", + } + + if err := reset.Rollback(); err != nil { + t.Fatalf("Rollback: %v", err) + } + creds := backend.devices["device-1"] + if creds.PairingID != "pair-2" || creds.RestoreToken != "restore-1" { + t.Fatalf("credentials after rollback = %+v", creds) + } +} + +func TestRestoreTokenResetRollbackDoesNotOverwriteConcurrentToken(t *testing.T) { + backend := &recordingCredentialBackend{devices: map[string]*SavedCredentials{ + "device-1": {RestoreToken: "restore-1"}, + }} + store := NewCredentialStoreWithBackend(backend) + + reset, err := store.BeginRestoreTokenReset("device-1") + if err != nil { + t.Fatalf("BeginRestoreTokenReset: %v", err) + } + backend.devices["device-1"] = &SavedCredentials{RestoreToken: "restore-2"} + + if err := reset.Rollback(); err != nil { + t.Fatalf("Rollback: %v", err) + } + if got := backend.devices["device-1"].RestoreToken; got != "restore-2" { + t.Fatalf("restore token after rollback = %q, want concurrent token", got) + } +} diff --git a/internal/airplay/mirror.go b/internal/airplay/mirror.go index 1784951..d5b981b 100644 --- a/internal/airplay/mirror.go +++ b/internal/airplay/mirror.go @@ -504,6 +504,9 @@ func (c *AirPlayClient) setupMirrorSession(ctx context.Context, cfg StreamConfig var receiverEventPort int var attemptedEventPort int audioControlLPort := audioCtrlConn.LocalAddr().(*net.UDPAddr).Port + if !targetLatencyIsExplicit() { + latencies = latencies.withMinimumVideoLead(cfg.MinimumVideoLead) + } audioLatencySamples := samplesFor44k1(latencies.audio) measuredLatencyApplied := false audioSetupCommitted := false diff --git a/internal/airplay/receiver_server_test.go b/internal/airplay/receiver_server_test.go index c3c28f1..b253850 100644 --- a/internal/airplay/receiver_server_test.go +++ b/internal/airplay/receiver_server_test.go @@ -579,6 +579,37 @@ func TestMediaFirstAutoFallsBackBeforeCreatingLatencyMismatch(t *testing.T) { } } +func TestMediaFirstRelayLeadPreservesAutomaticHEVC(t *testing.T) { + SetTargetLatency(0) + t.Cleanup(func() { SetTargetLatency(0) }) + _, client, ctx := newReceiverServerTestPair(t, ReceiverConfig{ + Profile: ReceiverProfileRoku, DisplayWidth: 3840, DisplayHeight: 2160, + }) + if err := client.Pair(ctx, ""); err != nil { + t.Fatalf("pair: %v", err) + } + session, err := client.SetupMirrorWithVideoCodecPreparation(ctx, StreamConfig{ + VideoCodec: VideoCodecAuto, + AutomaticHEVCAvailable: true, + MinimumVideoLead: 250 * time.Millisecond, + }, func(_, _ int, codec VideoCodec) error { + if codec != VideoCodecHEVC { + return fmt.Errorf("media-first relay codec = %s, want HEVC", codec) + } + return nil + }) + if err != nil { + t.Fatalf("setup mirror: %v", err) + } + defer session.Close() + if session.videoCodec != VideoCodecHEVC || session.timestampBias != 250*time.Millisecond { + t.Fatalf("media-first relay session = codec %s lead %v, want HEVC/250ms", session.videoCodec, session.timestampBias) + } + if session.audioStream == nil || session.audioStream.latencySamples != samplesFor44k1(260*time.Millisecond) { + t.Fatalf("media-first relay audio lead = %#v, want %d samples", session.audioStream, samplesFor44k1(260*time.Millisecond)) + } +} + func TestMediaFirstCalibratedAutoFallsBackBeforeLiveMeasurement(t *testing.T) { SetTargetLatency(0) t.Cleanup(func() { SetTargetLatency(0) }) diff --git a/internal/daemon/daemon.go b/internal/daemon/daemon.go index 3d9d669..3350404 100644 --- a/internal/daemon/daemon.go +++ b/internal/daemon/daemon.go @@ -204,6 +204,7 @@ type activeStream struct { device string // friendly name deviceIP string deviceID string + port int state State audioMuted bool session *airplay.MirrorSession @@ -233,8 +234,11 @@ type videoCaptureGroup struct { capture *airplay.ScreenCapture minimumVideoLead time.Duration cancel context.CancelFunc + resetReservedBy *activeStream } +var errCaptureGroupResetReserved = errors.New("capture group is reserved for restore-token reset") + // daemonCleanup owns resources detached from the daemon's state maps. Building // a cleanup plan while holding d.mu makes the state change atomic; running it // afterwards keeps cancellation, pipe closure, RTSP teardown, socket closure, @@ -264,6 +268,9 @@ func (cleanup *daemonCleanup) addStream(entry *activeStream) { cleanup.clients = append(cleanup.clients, entry.client) entry.client = nil } + if entry.captureGroup != nil && entry.captureGroup.resetReservedBy == entry { + entry.captureGroup.resetReservedBy = nil + } entry.captureGroup = nil } @@ -277,6 +284,7 @@ func (cleanup *daemonCleanup) addCaptureGroup(group *videoCaptureGroup) { group.capture = nil } group.broadcast = nil + group.resetReservedBy = nil } // run may block and therefore must never be called with d.mu held. Cancel all @@ -548,6 +556,8 @@ func (d *Daemon) handleRequest(req Request) Response { return d.handleConnect(req) case "disconnect": return d.handleDisconnect(req) + case "reset-restore-token": + return d.handleResetRestoreToken(req) case "mute": return d.handleSetMute(req, true) case "unmute": @@ -783,6 +793,7 @@ func (d *Daemon) handleConnect(req Request) Response { connCtx, cancel := context.WithCancel(context.Background()) entry := &activeStream{ deviceIP: target, + port: port, state: StateConnecting, cancelFn: cancel, credentialCh: make(chan string, 1), @@ -1124,6 +1135,7 @@ func (d *Daemon) connectAndStream(ctx context.Context, entry *activeStream, targ streamCfg := d.mirrorStreamConfig() streamCfg.AutomaticHEVCAvailable = capturePreparation.AutomaticHEVCAvailable() streamCfg.MeasuredVideoLatency = capturePreparation.MeasuredVideoLatency() + streamCfg.MinimumVideoLead = capturePreparation.MinimumVideoLead() var broadcast *airplay.BroadcastCapture selectedCaptureKey := videoCaptureKey{maxWidth: -1, maxHeight: -1} prepareVideo := func(width, height int, codec airplay.VideoCodec) (airplay.VideoPreparationResult, error) { @@ -1342,20 +1354,38 @@ func (d *Daemon) getOrStartPreparedCaptureGroup(ctx context.Context, entry *acti if d.captureGroups == nil { d.captureGroups = make(map[videoCaptureKey]*videoCaptureGroup) } - if group := d.captureGroups[key]; group != nil { - entry.captureGroup = group - broadcast := group.broadcast + if _, err := d.migrateRestoreTokenResetReservationLocked(entry, key); err != nil { d.mu.Unlock() - preparation.Close() - if broadcast == nil { - return nil, 0, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight) + return nil, 0, err + } + group := d.captureGroups[key] + var captureCtx context.Context + var captureCancel context.CancelFunc + if group != nil { + if group.resetReservedBy != nil && group.resetReservedBy != entry { + d.mu.Unlock() + return nil, 0, fmt.Errorf("%w: %dx%d", errCaptureGroupResetReserved, key.maxWidth, key.maxHeight) + } + if group.resetReservedBy == entry && group.broadcast == nil && group.capture == nil && group.cancel == nil { + captureCtx, captureCancel = context.WithCancel(context.Background()) + group.cancel = captureCancel + group.resetReservedBy = nil + } else { + entry.captureGroup = group + broadcast := group.broadcast + d.mu.Unlock() + preparation.Close() + if broadcast == nil { + return nil, 0, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight) + } + return broadcast, group.minimumVideoLead, nil } - return broadcast, group.minimumVideoLead, nil + } else { + captureCtx, captureCancel = context.WithCancel(context.Background()) + group = &videoCaptureGroup{key: key, cancel: captureCancel} + d.captureGroups[key] = group + entry.captureGroup = group } - - captureCtx, captureCancel := context.WithCancel(context.Background()) - group := &videoCaptureGroup{key: key, cancel: captureCancel} - d.captureGroups[key] = group entry.captureGroup = group d.mu.Unlock() @@ -1448,22 +1478,37 @@ func (d *Daemon) getOrStartCaptureGroup(entry *activeStream, restoreToken, devic if d.captureGroups == nil { d.captureGroups = make(map[videoCaptureKey]*videoCaptureGroup) } - if group := d.captureGroups[key]; group != nil { - entry.captureGroup = group - broadcast := group.broadcast + if _, err := d.migrateRestoreTokenResetReservationLocked(entry, key); err != nil { d.mu.Unlock() - if broadcast == nil { - return nil, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight) + return nil, err + } + group := d.captureGroups[key] + var captureCtx context.Context + var captureCancel context.CancelFunc + if group != nil { + if group.resetReservedBy != nil && group.resetReservedBy != entry { + d.mu.Unlock() + return nil, fmt.Errorf("%w: %dx%d", errCaptureGroupResetReserved, key.maxWidth, key.maxHeight) + } + if group.resetReservedBy == entry && group.broadcast == nil && group.capture == nil && group.cancel == nil { + captureCtx, captureCancel = context.WithCancel(context.Background()) + group.cancel = captureCancel + group.resetReservedBy = nil + } else { + entry.captureGroup = group + broadcast := group.broadcast + d.mu.Unlock() + if broadcast == nil { + return nil, fmt.Errorf("capture group %dx%d has no broadcast", key.maxWidth, key.maxHeight) + } + return broadcast, nil } - return broadcast, nil + } else { + captureCtx, captureCancel = context.WithCancel(context.Background()) + group = &videoCaptureGroup{key: key, cancel: captureCancel} + d.captureGroups[key] = group + entry.captureGroup = group } - - // Publish the group and cancellation hook before entering the display portal - // or launching GStreamer. A targeted disconnect can then cancel an orphaned - // startup without affecting captures used by other canvas groups. - captureCtx, captureCancel := context.WithCancel(context.Background()) - group := &videoCaptureGroup{key: key, cancel: captureCancel} - d.captureGroups[key] = group entry.captureGroup = group d.mu.Unlock() @@ -1528,6 +1573,30 @@ func (d *Daemon) getOrStartCaptureGroup(entry *activeStream, restoreToken, devic return newBC, nil } +// migrateRestoreTokenResetReservationLocked moves an owner-only reset claim +// when the receiver negotiates a different canvas after reconnecting. It never +// joins an occupied destination: reset replacement remains exclusive. +func (d *Daemon) migrateRestoreTokenResetReservationLocked(entry *activeStream, key videoCaptureKey) (*videoCaptureGroup, error) { + reservation := entry.captureGroup + if reservation == nil || reservation.resetReservedBy != entry || + reservation.broadcast != nil || reservation.capture != nil || reservation.cancel != nil { + return nil, nil + } + if reservation.key == key { + return reservation, nil + } + if d.captureGroups[reservation.key] != reservation { + return nil, context.Canceled + } + if d.captureGroups[key] != nil { + return nil, fmt.Errorf("%w: replacement canvas %dx%d is already active", errCaptureGroupResetReserved, key.maxWidth, key.maxHeight) + } + delete(d.captureGroups, reservation.key) + reservation.key = key + d.captureGroups[key] = reservation + return reservation, nil +} + // detachStreamLocked removes a single stream and transfers ownership of its // resources, plus an unused final capture group, to a cleanup plan. The caller // must unlock d.mu before running the plan. diff --git a/internal/daemon/daemonclient/client.go b/internal/daemon/daemonclient/client.go index 2ce7d23..47f8bf5 100644 --- a/internal/daemon/daemonclient/client.go +++ b/internal/daemon/daemonclient/client.go @@ -54,6 +54,12 @@ func (c *Client) DisconnectTarget(target string) (*daemon.Response, error) { return c.send(daemon.Request{Cmd: "disconnect", Target: target}) } +// ResetRestoreToken clears one active receiver's Wayland restore token and +// reconnects that same target. +func (c *Client) ResetRestoreToken(target string) (*daemon.Response, error) { + return c.send(daemon.Request{Cmd: "reset-restore-token", Target: target}) +} + // Mute mutes mirrored audio on all active sessions. func (c *Client) Mute() (*daemon.Response, error) { return c.send(daemon.Request{Cmd: "mute"}) diff --git a/internal/daemon/daemonclient/client_test.go b/internal/daemon/daemonclient/client_test.go new file mode 100644 index 0000000..cfa8da9 --- /dev/null +++ b/internal/daemon/daemonclient/client_test.go @@ -0,0 +1,52 @@ +package daemonclient + +import ( + "encoding/json" + "net" + "path/filepath" + "testing" + "time" + + "doubletake/internal/daemon" +) + +func TestClientResetRestoreTokenSendsTargetedCommand(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "doubletake.sock") + listener, err := net.Listen("unix", socketPath) + if err != nil { + t.Fatalf("listen: %v", err) + } + defer listener.Close() + requestCh := make(chan daemon.Request, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + defer conn.Close() + var request daemon.Request + if json.NewDecoder(conn).Decode(&request) != nil { + return + } + requestCh <- request + _ = json.NewEncoder(conn).Encode(daemon.Response{OK: true, State: daemon.StateConnecting}) + }() + client := New(socketPath) + + response, err := client.ResetRestoreToken("192.0.2.10") + + if err != nil { + t.Fatalf("ResetRestoreToken: %v", err) + } + if response == nil || !response.OK { + t.Fatalf("response = %+v", response) + } + select { + case request := <-requestCh: + if request.Cmd != "reset-restore-token" || request.Target != "192.0.2.10" { + t.Fatalf("request = %+v", request) + } + case <-time.After(time.Second): + t.Fatal("client did not send reset request") + } +} diff --git a/internal/daemon/reset_restore_token.go b/internal/daemon/reset_restore_token.go new file mode 100644 index 0000000..1ffd3cd --- /dev/null +++ b/internal/daemon/reset_restore_token.go @@ -0,0 +1,184 @@ +package daemon + +import ( + "context" + + "doubletake/internal/airplay" +) + +func (d *Daemon) handleResetRestoreToken(req Request) Response { + d.mu.Lock() + if req.Target == "" { + response := Response{OK: false, State: d.overallStateLocked(), Error: "reset-restore-token requires a target IP"} + d.mu.Unlock() + return response + } + if d.shuttingDown { + d.mu.Unlock() + return Response{OK: false, State: StateIdle, Error: "daemon is shutting down"} + } + entry, ok := d.streams[req.Target] + if !ok { + response := Response{OK: false, State: d.overallStateLocked(), Error: "no active stream to " + req.Target} + d.mu.Unlock() + return response + } + if entry.state != StateStreaming { + response := Response{OK: false, State: d.overallStateLocked(), Error: "restore token can be reset only for a fully streaming target: " + req.Target} + d.mu.Unlock() + return response + } + group := entry.captureGroup + if group == nil || d.captureGroups[group.key] != group { + response := Response{OK: false, State: d.overallStateLocked(), Error: "active target has no owned capture group: " + req.Target} + d.mu.Unlock() + return response + } + if group.resetReservedBy != nil { + response := Response{OK: false, State: d.overallStateLocked(), Error: "restore token reset is already in progress for " + req.Target} + d.mu.Unlock() + return response + } + for target, other := range d.streams { + if target != req.Target && other.captureGroup == group { + response := Response{OK: false, State: d.overallStateLocked(), Error: "target uses a shared capture group; disconnect its peers before resetting the restore token"} + d.mu.Unlock() + return response + } + } + if entry.deviceIP != req.Target || entry.deviceID == "" || entry.port == 0 { + response := Response{OK: false, State: d.overallStateLocked(), Error: "active target is missing receiver identity or port: " + req.Target} + d.mu.Unlock() + return response + } + + target := entry.deviceIP + deviceID := entry.deviceID + port := entry.port + group.resetReservedBy = entry + // Register the reset before releasing d.mu. Shutdown may detach the original + // stream while credential I/O is in progress, but it cannot finish Wait until + // this request either abandons the reservation or starts the replacement. + d.streamWorkers.Add(1) + d.mu.Unlock() + + credentialReset, clearErr := d.credStore.BeginRestoreTokenReset(deviceID) + + d.mu.Lock() + reservationCurrent := d.restoreTokenResetReservationCurrentLocked(target, deviceID, port, entry, group) + state := d.overallStateLocked() + if d.shuttingDown { + if group.resetReservedBy == entry { + group.resetReservedBy = nil + } + d.mu.Unlock() + rollbackErr := rollbackRestoreTokenReset(credentialReset, clearErr) + d.streamWorkers.Done() + return Response{OK: false, State: state, Error: resetFailure("daemon is shutting down", rollbackErr)} + } + if !reservationCurrent { + if group.resetReservedBy == entry { + group.resetReservedBy = nil + } + d.mu.Unlock() + rollbackErr := rollbackRestoreTokenReset(credentialReset, clearErr) + d.streamWorkers.Done() + return Response{OK: false, State: state, Error: resetFailure("restore token reset was canceled for "+target, rollbackErr)} + } + if clearErr != nil { + group.resetReservedBy = nil + d.mu.Unlock() + d.streamWorkers.Done() + return Response{OK: false, State: state, Error: "clear restore token: " + clearErr.Error()} + } + + connCtx, cancel := context.WithCancel(context.Background()) + replacement := &activeStream{ + deviceIP: target, + deviceID: deviceID, + port: port, + state: StateConnecting, + cancelFn: cancel, + credentialCh: make(chan string, 1), + } + // Replace the old generation with an owner-only claim before releasing the + // daemon lock. Peers cannot create or join this key while physical cleanup + // runs; the designated replacement later converts the claim into a capture. + group.resetReservedBy = nil + cleanup := d.detachStreamLocked(target) + reservation := &videoCaptureGroup{ + key: group.key, + resetReservedBy: replacement, + } + replacement.captureGroup = reservation + d.captureGroups[group.key] = reservation + d.streams[target] = replacement + d.mu.Unlock() + + cleanup.run() + + // Cleanup can block, so disconnect or shutdown may remove the replacement + // reservation before its worker starts. Only detach the exact generation we + // published and report the daemon's actual aggregate state. + d.mu.Lock() + replacementCurrent := d.streams[target] == replacement + shuttingDown := d.shuttingDown + if shuttingDown || !replacementCurrent { + abandoned := daemonCleanup{} + if replacementCurrent { + abandoned = d.detachStreamLocked(target) + } + state = d.overallStateLocked() + d.mu.Unlock() + abandoned.run() + cancel() + rollbackErr := credentialReset.Rollback() + d.streamWorkers.Done() + if shuttingDown { + return Response{OK: false, State: state, Error: resetFailure("daemon is shutting down", rollbackErr)} + } + return Response{OK: false, State: state, Error: resetFailure("restore token reset was canceled for "+target, rollbackErr)} + } + d.clearLastErrorForTargetLocked(target) + state = d.overallStateLocked() + d.mu.Unlock() + + credentialReset.Commit() + go func() { + defer d.streamWorkers.Done() + d.connectAndStream(connCtx, replacement, target, port, "") + }() + return Response{OK: true, State: state, Device: target, DeviceIP: target} +} + +func rollbackRestoreTokenReset(reset *airplay.RestoreTokenReset, clearErr error) error { + if clearErr != nil { + return nil + } + return reset.Rollback() +} + +func resetFailure(message string, rollbackErr error) string { + if rollbackErr == nil { + return message + } + return message + "; restore token rollback failed: " + rollbackErr.Error() +} + +// restoreTokenResetReservationCurrentLocked reports whether the exact stream +// and exclusive capture generation reserved before credential I/O are still +// current. Must be called with d.mu held. +func (d *Daemon) restoreTokenResetReservationCurrentLocked(target, deviceID string, port int, entry *activeStream, group *videoCaptureGroup) bool { + if d.streams[target] != entry || entry.state != StateStreaming || + entry.deviceIP != target || entry.deviceID != deviceID || entry.port != port || + entry.captureGroup != group || d.captureGroups[group.key] != group || + group.resetReservedBy != entry { + return false + } + for otherTarget, other := range d.streams { + if otherTarget != target && other.captureGroup == group { + return false + } + } + return true +} diff --git a/internal/daemon/reset_restore_token_test.go b/internal/daemon/reset_restore_token_test.go new file mode 100644 index 0000000..2c9420f --- /dev/null +++ b/internal/daemon/reset_restore_token_test.go @@ -0,0 +1,671 @@ +package daemon + +import ( + "context" + "errors" + "net" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "doubletake/internal/airplay" +) + +func TestResetRestoreTokenRejectsMissingUnknownAndNonStreamingTargetsWithoutMutation(t *testing.T) { + for _, test := range []struct { + name string + target string + state State + }{ + {name: "missing target"}, + {name: "unknown target", target: "192.0.2.99"}, + {name: "connecting target", target: "192.0.2.10", state: StateConnecting}, + {name: "credential-waiting target", target: "192.0.2.10", state: StatePINRequired}, + } { + t.Run(test.name, func(t *testing.T) { + d, err := New(Config{CredFile: filepath.Join(t.TempDir(), "credentials.json")}) + if err != nil { + t.Fatalf("New: %v", err) + } + if err := d.credStore.SaveRestoreToken("device-1", "restore-1"); err != nil { + t.Fatalf("SaveRestoreToken: %v", err) + } + entry := &activeStream{deviceIP: "192.0.2.10", deviceID: "device-1", state: test.state} + if test.state != "" { + d.streams[entry.deviceIP] = entry + } + + response := d.handleRequest(Request{Cmd: "reset-restore-token", Target: test.target}) + + if response.OK { + t.Fatalf("reset unexpectedly succeeded: %+v", response) + } + if d.streams[entry.deviceIP] != entry && test.state != "" { + t.Fatal("rejected reset detached the target") + } + creds := d.credStore.Lookup("device-1") + if creds == nil || creds.RestoreToken != "restore-1" { + t.Fatalf("rejected reset mutated credentials: %+v", creds) + } + }) + } +} + +func TestResetRestoreTokenRejectsSharedCaptureGroupWithoutMutation(t *testing.T) { + d, err := New(Config{CredFile: filepath.Join(t.TempDir(), "credentials.json")}) + if err != nil { + t.Fatalf("New: %v", err) + } + if err := d.credStore.SaveRestoreToken("device-1", "restore-1"); err != nil { + t.Fatalf("SaveRestoreToken: %v", err) + } + group := &videoCaptureGroup{key: normalizedVideoCaptureKey(1920, 1080)} + target := &activeStream{deviceIP: "192.0.2.10", deviceID: "device-1", state: StateStreaming, captureGroup: group} + peer := &activeStream{deviceIP: "192.0.2.11", deviceID: "device-2", state: StateStreaming, captureGroup: group} + d.streams[target.deviceIP] = target + d.streams[peer.deviceIP] = peer + d.captureGroups[group.key] = group + + response := d.handleRequest(Request{Cmd: "reset-restore-token", Target: target.deviceIP}) + + if response.OK || !strings.Contains(response.Error, "shared capture group") { + t.Fatalf("shared-group reset response = %+v", response) + } + if d.streams[target.deviceIP] != target || d.streams[peer.deviceIP] != peer || d.captureGroups[group.key] != group { + t.Fatal("shared-group rejection changed active stream state") + } + creds := d.credStore.Lookup("device-1") + if creds == nil || creds.RestoreToken != "restore-1" { + t.Fatalf("shared-group rejection mutated credentials: %+v", creds) + } +} + +func TestResetRestoreTokenKeepsDaemonResponsiveDuringCredentialSave(t *testing.T) { + saveStarted := make(chan struct{}) + saveRelease := make(chan struct{}) + var releaseOnce sync.Once + releaseSave := func() { releaseOnce.Do(func() { close(saveRelease) }) } + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{RestoreToken: "restore-1"}, + saveStarted: saveStarted, + saveRelease: saveRelease, + saveErr: errors.New("save failed"), + } + d, _, _, _ := newResetTestDaemon(t, backend) + defer d.Shutdown() + defer releaseSave() + + resetResponse := make(chan Response, 1) + go func() { + resetResponse <- d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + }() + waitForResetSignal(t, saveStarted, "credential save did not start") + + statusResponse := make(chan Response, 1) + go func() { + statusResponse <- d.handleStatus() + }() + select { + case response := <-statusResponse: + if response.State != StateStreaming { + t.Fatalf("status during credential save = %+v, want original stream", response) + } + case <-time.After(time.Second): + t.Fatal("status blocked behind credential backend I/O") + } + + releaseSave() + response := waitForResetResponse(t, resetResponse) + if response.OK || !strings.Contains(response.Error, "save failed") { + t.Fatalf("reset response = %+v, want credential save failure", response) + } +} + +func TestResetRestoreTokenCredentialFailureIsAtomic(t *testing.T) { + for _, test := range []struct { + name string + lookupErr error + saveErr error + }{ + {name: "lookup failure", lookupErr: errors.New("lookup failed")}, + {name: "save failure", saveErr: errors.New("save failed")}, + } { + t.Run(test.name, func(t *testing.T) { + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{PairingID: "pair-1", RestoreToken: "restore-1"}, + lookupErr: test.lookupErr, + saveErr: test.saveErr, + } + d, entry, group, streamCtx := newResetTestDaemon(t, backend) + defer d.Shutdown() + d.lastError = "existing stream error" + d.lastErrorTarget = resetTestTarget + + response := d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + + if response.OK { + t.Fatalf("reset unexpectedly succeeded: %+v", response) + } + d.mu.Lock() + streamPreserved := d.streams[resetTestTarget] == entry + groupPreserved := d.captureGroups[group.key] == group && entry.captureGroup == group + errorPreserved := d.lastError == "existing stream error" && d.lastErrorTarget == resetTestTarget + d.mu.Unlock() + if !streamPreserved || !groupPreserved { + t.Fatalf("credential failure changed stream state: stream=%t group=%t", streamPreserved, groupPreserved) + } + if !errorPreserved { + t.Fatalf("credential failure changed last error: %q for %q", d.lastError, d.lastErrorTarget) + } + select { + case <-streamCtx.Done(): + t.Fatal("credential failure canceled the original stream") + default: + } + creds := d.credStore.Lookup("device-1") + if creds == nil || creds.RestoreToken != "restore-1" { + t.Fatalf("credential failure changed restore token: %+v", creds) + } + + peer := &activeStream{deviceIP: "192.0.2.11", state: StateConnecting} + d.mu.Lock() + d.streams[peer.deviceIP] = peer + d.mu.Unlock() + broadcast, _, err := d.getOrStartPreparedCaptureGroup(context.Background(), peer, nil, 1920, 1080, airplay.VideoCodecH264) + if err != nil || broadcast != group.broadcast { + t.Fatalf("capture reservation remained after failed reset: broadcast=%p err=%v", broadcast, err) + } + d.mu.Lock() + delete(d.streams, peer.deviceIP) + peer.captureGroup = nil + d.mu.Unlock() + }) + } +} + +func TestResetRestoreTokenReservationExcludesCaptureGroupJoin(t *testing.T) { + saveStarted := make(chan struct{}) + saveRelease := make(chan struct{}) + var releaseOnce sync.Once + releaseSave := func() { releaseOnce.Do(func() { close(saveRelease) }) } + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{RestoreToken: "restore-1"}, + saveStarted: saveStarted, + saveRelease: saveRelease, + saveErr: errors.New("save failed"), + } + d, _, _, _ := newResetTestDaemon(t, backend) + defer d.Shutdown() + defer releaseSave() + peer := &activeStream{deviceIP: "192.0.2.11", state: StateConnecting} + d.streams[peer.deviceIP] = peer + + resetResponse := make(chan Response, 1) + go func() { + resetResponse <- d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + }() + waitForResetSignal(t, saveStarted, "credential save did not start") + + type joinResult struct { + err error + } + joinResults := make(chan joinResult, 1) + go func() { + _, _, err := d.getOrStartPreparedCaptureGroup(context.Background(), peer, nil, 1920, 1080, airplay.VideoCodecH264) + joinResults <- joinResult{err: err} + }() + var result joinResult + select { + case result = <-joinResults: + case <-time.After(time.Second): + t.Fatal("peer capture join blocked behind credential backend I/O") + } + if result.err == nil || !strings.Contains(result.err.Error(), "reserved for restore-token reset") { + t.Fatalf("peer capture join error = %v, want reset reservation rejection", result.err) + } + d.mu.Lock() + joined := peer.captureGroup != nil + delete(d.streams, peer.deviceIP) + d.mu.Unlock() + if joined { + t.Fatal("peer joined capture group while restore-token reset was reserved") + } + + releaseSave() + response := waitForResetResponse(t, resetResponse) + if response.OK || !strings.Contains(response.Error, "save failed") { + t.Fatalf("reset response = %+v, want credential save failure", response) + } +} + +func TestResetRestoreTokenReservationSurvivesPhysicalCaptureCleanup(t *testing.T) { + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{RestoreToken: "restore-1"}, + } + d, entry, oldGroup, _ := newResetTestDaemon(t, backend) + cleanupStarted := make(chan struct{}) + cleanupRelease := make(chan struct{}) + var releaseOnce sync.Once + releaseCleanup := func() { releaseOnce.Do(func() { close(cleanupRelease) }) } + originalCancel := entry.cancelFn + entry.cancelFn = func() { + close(cleanupStarted) + <-cleanupRelease + originalCancel() + } + peer := &activeStream{deviceIP: "192.0.2.11", state: StateConnecting} + d.streams[peer.deviceIP] = peer + defer d.Shutdown() + defer releaseCleanup() + + resetResponse := make(chan Response, 1) + go func() { + resetResponse <- d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + }() + waitForResetSignal(t, cleanupStarted, "old capture cleanup did not start") + + d.mu.Lock() + replacement := d.streams[resetTestTarget] + reservation := d.captureGroups[oldGroup.key] + reserved := reservation != nil && reservation != oldGroup && + reservation.resetReservedBy == replacement && + replacement.captureGroup == reservation + d.mu.Unlock() + if !reserved { + t.Fatal("replacement did not retain an exclusive capture-key reservation during cleanup") + } + + _, _, joinErr := d.getOrStartPreparedCaptureGroup( + context.Background(), peer, nil, 1920, 1080, airplay.VideoCodecH264, + ) + if joinErr == nil || !strings.Contains(joinErr.Error(), "reserved for restore-token reset") { + t.Fatalf("capture join during cleanup = %v, want reset reservation rejection", joinErr) + } + d.mu.Lock() + delete(d.streams, peer.deviceIP) + d.mu.Unlock() + + releaseCleanup() + if response := waitForResetResponse(t, resetResponse); !response.OK { + t.Fatalf("reset response = %+v", response) + } +} + +func TestResetRestoreTokenReservationMakesShutdownWait(t *testing.T) { + saveStarted := make(chan struct{}) + saveRelease := make(chan struct{}) + var releaseOnce sync.Once + releaseSave := func() { releaseOnce.Do(func() { close(saveRelease) }) } + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{RestoreToken: "restore-1"}, + saveStarted: saveStarted, + saveRelease: saveRelease, + } + d, entry, _, _ := newResetTestDaemon(t, backend) + defer releaseSave() + + cancelEvents := make(chan struct{}, 2) + originalCancel := entry.cancelFn + entry.cancelFn = func() { + originalCancel() + cancelEvents <- struct{}{} + } + resetResponse := make(chan Response, 1) + resetDone := make(chan struct{}) + go func() { + resetResponse <- d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + close(resetDone) + }() + defer func() { + releaseSave() + waitForResetSignal(t, resetDone, "reset did not finish during cleanup") + }() + waitForResetSignal(t, saveStarted, "credential save did not start") + select { + case <-cancelEvents: + // The old implementation canceled before credential I/O. Drain that event + // so only shutdown detachment can satisfy the wait below. + default: + } + + shutdownDone := make(chan struct{}) + go func() { + d.Shutdown() + close(shutdownDone) + }() + waitForResetSignal(t, cancelEvents, "shutdown did not detach the reserved stream") + select { + case <-shutdownDone: + t.Fatal("Shutdown returned while credential I/O was still blocked") + case <-time.After(100 * time.Millisecond): + } + + releaseSave() + response := waitForResetResponse(t, resetResponse) + if response.OK || !strings.Contains(response.Error, "shutting down") || response.State != StateIdle { + t.Fatalf("reset response during shutdown = %+v", response) + } + if credentials := d.credStore.Lookup("device-1"); credentials == nil || credentials.RestoreToken != "restore-1" { + t.Fatalf("shutdown cancellation lost restore token: %+v", credentials) + } + waitForResetSignal(t, shutdownDone, "Shutdown did not finish after reset released its reservation") +} + +func TestResetRestoreTokenConcurrentDisconnectReportsCancellationAndOverallState(t *testing.T) { + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{RestoreToken: "restore-1"}, + } + d, entry, _, _ := newResetTestDaemon(t, backend) + cleanupStarted := make(chan struct{}) + cleanupRelease := make(chan struct{}) + var releaseOnce sync.Once + releaseCleanup := func() { releaseOnce.Do(func() { close(cleanupRelease) }) } + originalCancel := entry.cancelFn + entry.cancelFn = func() { + close(cleanupStarted) + <-cleanupRelease + originalCancel() + } + independentGroup := &videoCaptureGroup{key: normalizedVideoCaptureKey(1280, 720)} + independent := &activeStream{ + deviceIP: "192.0.2.20", + deviceID: "device-2", + port: 7000, + state: StateStreaming, + captureGroup: independentGroup, + } + d.streams[independent.deviceIP] = independent + d.captureGroups[independentGroup.key] = independentGroup + defer d.Shutdown() + defer releaseCleanup() + + resetResponse := make(chan Response, 1) + go func() { + resetResponse <- d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + }() + waitForResetSignal(t, cleanupStarted, "old stream cleanup did not start") + + disconnect := d.handleDisconnect(Request{Cmd: "disconnect", Target: resetTestTarget}) + if !disconnect.OK || disconnect.State != StateStreaming { + t.Fatalf("concurrent disconnect response = %+v, want independent stream preserved", disconnect) + } + releaseCleanup() + response := waitForResetResponse(t, resetResponse) + if response.OK || !strings.Contains(response.Error, "canceled") { + t.Fatalf("reset response = %+v, want concurrent cancellation", response) + } + if strings.Contains(response.Error, "shutting down") || response.State != StateStreaming { + t.Fatalf("reset misreported concurrent disconnect: %+v", response) + } + if credentials := d.credStore.Lookup("device-1"); credentials == nil || credentials.RestoreToken != "restore-1" { + t.Fatalf("disconnect cancellation lost restore token: %+v", credentials) + } +} + +func TestCanceledResetReportsRollbackPersistenceFailures(t *testing.T) { + for _, test := range []struct { + name string + lookupErrAt int + saveErrAt int + }{ + {name: "rollback lookup", lookupErrAt: 2}, + {name: "rollback save", saveErrAt: 2}, + } { + t.Run(test.name, func(t *testing.T) { + backend := &controlledCredentialBackend{ + credentials: &airplay.SavedCredentials{RestoreToken: "restore-1"}, + lookupErrAt: test.lookupErrAt, + saveErrAt: test.saveErrAt, + } + d, entry, _, _ := newResetTestDaemon(t, backend) + cleanupStarted := make(chan struct{}) + cleanupRelease := make(chan struct{}) + originalCancel := entry.cancelFn + entry.cancelFn = func() { + close(cleanupStarted) + <-cleanupRelease + originalCancel() + } + defer d.Shutdown() + + resetResponse := make(chan Response, 1) + go func() { + resetResponse <- d.handleResetRestoreToken(Request{Cmd: "reset-restore-token", Target: resetTestTarget}) + }() + waitForResetSignal(t, cleanupStarted, "cleanup did not start") + if response := d.handleDisconnect(Request{Cmd: "disconnect", Target: resetTestTarget}); !response.OK { + t.Fatalf("disconnect response = %+v", response) + } + close(cleanupRelease) + + response := waitForResetResponse(t, resetResponse) + if response.OK || !strings.Contains(response.Error, "rollback failed") { + t.Fatalf("reset response = %+v, want explicit rollback failure", response) + } + }) + } +} + +func TestRestoreTokenResetReservationMigratesChangedCanvasExclusively(t *testing.T) { + d, entry, reservation, _ := newResetTestDaemon(t, &controlledCredentialBackend{}) + oldKey := reservation.key + newKey := normalizedVideoCaptureKey(1280, 720, airplay.VideoCodecH264) + reservation.broadcast = nil + reservation.resetReservedBy = entry + + d.mu.Lock() + migrated, err := d.migrateRestoreTokenResetReservationLocked(entry, newKey) + oldReleased := d.captureGroups[oldKey] == nil + newOwned := d.captureGroups[newKey] == reservation + d.mu.Unlock() + if err != nil || migrated != reservation || !oldReleased || !newOwned { + t.Fatalf("migration = %p, %v, oldReleased=%t newOwned=%t", migrated, err, oldReleased, newOwned) + } + + occupiedKey := normalizedVideoCaptureKey(3840, 2160, airplay.VideoCodecH264) + occupied := &videoCaptureGroup{key: occupiedKey} + d.mu.Lock() + d.captureGroups[occupiedKey] = occupied + _, err = d.migrateRestoreTokenResetReservationLocked(entry, occupiedKey) + stillOwned := d.captureGroups[newKey] == reservation && entry.captureGroup == reservation + d.mu.Unlock() + if err == nil || !errors.Is(err, errCaptureGroupResetReserved) || !stillOwned { + t.Fatalf("occupied migration = %v, stillOwned=%t", err, stillOwned) + } + + d.mu.Lock() + cleanup := d.detachStreamLocked(resetTestTarget) + reservationReleased := d.captureGroups[newKey] == nil + d.mu.Unlock() + cleanup.run() + if !reservationReleased { + t.Fatal("disconnect left migrated reservation behind") + } +} + +func TestResetRestoreTokenClearsExclusiveTargetAndReconnectsActualPort(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer listener.Close() + accepted := make(chan net.Conn, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr == nil { + accepted <- conn + } + }() + + d, err := New(Config{CredFile: filepath.Join(t.TempDir(), "credentials.json")}) + if err != nil { + t.Fatalf("New: %v", err) + } + defer d.Shutdown() + if err := d.credStore.SaveRestoreToken("device-1", "restore-1"); err != nil { + t.Fatalf("SaveRestoreToken: %v", err) + } + address := listener.Addr().(*net.TCPAddr) + group := &videoCaptureGroup{key: normalizedVideoCaptureKey(1920, 1080)} + oldContext, cancelOld := context.WithCancel(context.Background()) + old := &activeStream{ + deviceIP: address.IP.String(), + deviceID: "device-1", + state: StateStreaming, + port: address.Port, + captureGroup: group, + cancelFn: cancelOld, + } + independentGroup := &videoCaptureGroup{key: normalizedVideoCaptureKey(1280, 720)} + independent := &activeStream{ + deviceIP: "192.0.2.20", + deviceID: "device-2", + state: StateStreaming, + captureGroup: independentGroup, + } + d.streams[old.deviceIP] = old + d.streams[independent.deviceIP] = independent + d.captureGroups[group.key] = group + d.captureGroups[independentGroup.key] = independentGroup + + response := d.handleRequest(Request{Cmd: "reset-restore-token", Target: old.deviceIP}) + + if !response.OK { + t.Fatalf("reset response = %+v", response) + } + creds := d.credStore.Lookup("device-1") + if creds == nil || creds.RestoreToken != "" { + t.Fatalf("restore token was not cleared: %+v", creds) + } + select { + case <-oldContext.Done(): + default: + t.Fatal("old stream cleanup did not finish before reset returned") + } + d.mu.Lock() + replacement := d.streams[old.deviceIP] + replacementGroup := d.captureGroups[group.key] + oldGroupStillPresent := replacementGroup == group + independentPreserved := d.streams[independent.deviceIP] == independent && + d.captureGroups[independentGroup.key] == independentGroup + d.mu.Unlock() + if replacement == nil || replacement == old || replacement.port != address.Port { + t.Fatalf("replacement stream = %+v, want new entry on port %d", replacement, address.Port) + } + if oldGroupStillPresent { + t.Fatal("exclusive old capture group remained active") + } + if replacementGroup != nil && replacement.captureGroup != replacementGroup { + t.Fatal("replacement did not own the reserved capture generation") + } + if !independentPreserved { + t.Fatal("reset changed an independent stream or capture group") + } + select { + case conn := <-accepted: + defer conn.Close() + case <-time.After(3 * time.Second): + t.Fatal("reset did not reconnect to the target's actual port") + } +} + +const resetTestTarget = "192.0.2.10" + +type controlledCredentialBackend struct { + credentials *airplay.SavedCredentials + lookupErr error + saveErr error + saveStarted chan struct{} + saveRelease <-chan struct{} + saveOnce sync.Once + lookupCalls int + saveCalls int + lookupErrAt int + saveErrAt int +} + +func (b *controlledCredentialBackend) Lookup(string) (*airplay.SavedCredentials, error) { + b.lookupCalls++ + if b.lookupErrAt > 0 && b.lookupCalls == b.lookupErrAt { + return nil, errors.New("injected lookup failure") + } + if b.lookupErr != nil { + err := b.lookupErr + b.lookupErr = nil + return nil, err + } + if b.credentials == nil { + return nil, nil + } + credentials := *b.credentials + return &credentials, nil +} + +func (b *controlledCredentialBackend) Save(_ string, credentials *airplay.SavedCredentials) error { + if b.saveStarted != nil { + b.saveOnce.Do(func() { close(b.saveStarted) }) + } + if b.saveRelease != nil { + <-b.saveRelease + } + b.saveCalls++ + if b.saveErrAt > 0 && b.saveCalls == b.saveErrAt { + return errors.New("injected save failure") + } + if b.saveErr != nil { + return b.saveErr + } + updated := *credentials + b.credentials = &updated + return nil +} + +func newResetTestDaemon(t *testing.T, backend airplay.CredentialBackend) (*Daemon, *activeStream, *videoCaptureGroup, context.Context) { + t.Helper() + d, err := New(Config{ + SocketPath: filepath.Join(t.TempDir(), "doubletake.sock"), + CredFile: filepath.Join(t.TempDir(), "credentials.json"), + }) + if err != nil { + t.Fatalf("New: %v", err) + } + d.credStore = airplay.NewCredentialStoreWithBackend(backend) + group := &videoCaptureGroup{ + key: normalizedVideoCaptureKey(1920, 1080), + broadcast: airplay.NewBroadcastCapture(nil), + } + streamCtx, cancel := context.WithCancel(context.Background()) + entry := &activeStream{ + deviceIP: resetTestTarget, + deviceID: "device-1", + port: 7000, + state: StateStreaming, + captureGroup: group, + cancelFn: cancel, + } + d.streams[resetTestTarget] = entry + d.captureGroups[group.key] = group + return d, entry, group, streamCtx +} + +func waitForResetSignal(t *testing.T, signal <-chan struct{}, failure string) { + t.Helper() + select { + case <-signal: + case <-time.After(time.Second): + t.Fatal(failure) + } +} + +func waitForResetResponse(t *testing.T, responses <-chan Response) Response { + t.Helper() + select { + case response := <-responses: + return response + case <-time.After(time.Second): + t.Fatal("reset did not return") + return Response{} + } +} diff --git a/man/man1/doubletake-ctl.1 b/man/man1/doubletake-ctl.1 index 99839c9..8bdb270 100644 --- a/man/man1/doubletake-ctl.1 +++ b/man/man1/doubletake-ctl.1 @@ -29,6 +29,12 @@ to select one explicitly. .TP .B disconnect [\fITARGET\fR] Disconnect from an Apple TV. If no target is specified, disconnects all active connections +.TP +.B reset-restore-token \fITARGET\fR +Stop one fully streaming target, clear only its saved Wayland portal restore +token, and reconnect the same IP and port. Pairing credentials and other streams +are preserved. The command refuses a target whose capture group is shared; first +disconnect the peers using that group so the old portal source can be stopped. .SH EXAMPLES .TP Check daemon status: @@ -46,6 +52,9 @@ Disconnect from a specific Apple TV: .TP Disconnect all connections: .B doubletake-ctl disconnect +.TP +Clear one receiver's Wayland restore token and reconnect it: +.B doubletake-ctl reset-restore-token 192.168.1.77 .SH USAGE First, start the doubletake daemon in a separate terminal or background: .IP