Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/fix-ios-loopback-media-after-background.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
default: patch
---

Fix media failing to load after the app has been in the background, which previously needed an app restart to recover.
1 change: 1 addition & 0 deletions src-tauri/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -458,6 +458,7 @@ pub fn run() {
network::media_protocol::clear_media_session,
network::media_protocol::set_media_encryption,
network::media_protocol::prepare_loopback_media,
network::media_protocol::ensure_loopback_media,
sentry::set_native_sentry_enabled,
share_inbox::share_inbox_drain,
share_inbox::share_inbox_read,
Expand Down
25 changes: 24 additions & 1 deletion src-tauri/src/network/media_protocol.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,12 +10,14 @@ use std::{
use sha2::{Digest, Sha256};
use tauri::{
http::{header, Request, Response, StatusCode, Uri},
AppHandle, Manager, Runtime, State, UriSchemeContext, UriSchemeResponder,
AppHandle, Emitter, Manager, Runtime, State, UriSchemeContext, UriSchemeResponder,
};

mod crypto;
mod lane;
mod loopback;

pub const LOOPBACK_REBOUND_EVENT: &str = "sable-media://loopback-rebound";
mod response;
mod session;

Expand Down Expand Up @@ -132,6 +134,13 @@ impl MediaSessionState {
}
}

fn ensure_loopback_media_live(&self) -> Result<bool, String> {
let Some(loopback) = &self.loopback else {
return Ok(false);
};
loopback.ensure_live().map_err(|err| err.to_string())
}

// Shared across requests so the connection pool and TLS sessions stay warm.
fn client(&self) -> Client {
self.client
Expand Down Expand Up @@ -259,6 +268,17 @@ pub fn set_media_encryption(
.register(&url, &key, &iv, &sha256, &version, mime_type)
}

#[tauri::command]
pub fn ensure_loopback_media<R: Runtime>(
app: AppHandle<R>,
state: tauri::State<'_, MediaSessionState>,
) -> Result<(), String> {
if state.ensure_loopback_media_live()? {
let _ = app.emit(LOOPBACK_REBOUND_EVENT, ());
}
Ok(())
}

#[tauri::command]
pub async fn prepare_loopback_media<R: Runtime>(
app: AppHandle<R>,
Expand Down Expand Up @@ -372,6 +392,9 @@ async fn handle_request<R: Runtime>(
if loopback && in_memory_body.is_none() {
let state = app.state::<MediaSessionState>();
if let Some(loopback) = &state.loopback {
if loopback.ensure_live().unwrap_or(false) {
let _ = app.emit(LOOPBACK_REBOUND_EVENT, ());
}
return Ok(loopback.redirect_response(&session, &cache_key, disk_path, &content_type));
}
}
Expand Down
185 changes: 167 additions & 18 deletions src-tauri/src/network/media_protocol/loopback.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,10 @@ use std::{
io::{BufRead, BufReader, Read, Seek, SeekFrom, Write},
net::{TcpListener, TcpStream},
path::PathBuf,
sync::{Arc, RwLock},
sync::{
atomic::{AtomicBool, Ordering},
Arc, RwLock,
},
thread,
};

Expand All @@ -14,9 +17,19 @@ use tauri::http::{header, Response, StatusCode};
use super::response::is_webview_origin;
use super::session::MediaSession;

type Routes = Arc<RwLock<HashMap<String, CachedMedia>>>;

pub(super) struct LoopbackMediaServer {
routes: Routes,
active: RwLock<ActiveListener>,
}

struct ActiveListener {
origin: String,
routes: Arc<RwLock<HashMap<String, CachedMedia>>>,
port: u16,
/// Deliberate retirement, as opposed to `dead`.
shutdown: Arc<AtomicBool>,
dead: Arc<AtomicBool>,
}

#[derive(Clone)]
Expand All @@ -25,24 +38,66 @@ struct CachedMedia {
content_type: String,
}

// iOS reclaims a listening socket while suspended: the descriptor still looks valid but every
// accept() fails with ECONNABORTED (Apple TN2277). `ensure_live` rebinds, so failing fast here
// beats guessing whether an error was transient.
fn accept_loop(
listener: TcpListener,
routes: Routes,
shutdown: Arc<AtomicBool>,
dead: Arc<AtomicBool>,
) {
for stream in listener.incoming() {
if shutdown.load(Ordering::Acquire) {
return;
}
let Ok(stream) = stream else {
dead.store(true, Ordering::Release);
return;
};
let routes = Arc::clone(&routes);
let _ = thread::Builder::new()
.name("sable-media-request".into())
.spawn(move || serve(stream, routes));
}
dead.store(true, Ordering::Release);
}

fn spawn_listener(routes: &Routes) -> std::io::Result<ActiveListener> {
let listener = TcpListener::bind(("127.0.0.1", 0))?;
let port = listener.local_addr()?.port();
let shutdown = Arc::new(AtomicBool::new(false));
let dead = Arc::new(AtomicBool::new(false));
let thread_routes = Arc::clone(routes);
let thread_shutdown = Arc::clone(&shutdown);
let thread_dead = Arc::clone(&dead);
thread::Builder::new()
.name("sable-media-loopback".into())
.spawn(move || accept_loop(listener, thread_routes, thread_shutdown, thread_dead))?;

Ok(ActiveListener {
origin: format!("http://127.0.0.1:{port}"),
port,
shutdown,
dead,
})
}

// A retired thread is parked in accept() and needs a connection before it can see the flag.
fn retire(listener: &ActiveListener) {
listener.shutdown.store(true, Ordering::Release);
let _ = TcpStream::connect(("127.0.0.1", listener.port));
}

impl LoopbackMediaServer {
pub(super) fn start() -> std::io::Result<Self> {
let listener = TcpListener::bind(("127.0.0.1", 0))?;
let origin = format!("http://127.0.0.1:{}", listener.local_addr()?.port());
let routes = Arc::new(RwLock::new(HashMap::<String, CachedMedia>::new()));
let server_routes = Arc::clone(&routes);
thread::Builder::new()
.name("sable-media-loopback".into())
.spawn(move || {
for stream in listener.incoming().flatten() {
let routes = Arc::clone(&server_routes);
let _ = thread::Builder::new()
.name("sable-media-request".into())
.spawn(move || serve(stream, routes));
}
})?;

Ok(Self { origin, routes })
let active = spawn_listener(&routes)?;

Ok(Self {
routes,
active: RwLock::new(active),
})
}

pub(super) fn clear(&self) {
Expand All @@ -51,6 +106,43 @@ impl LoopbackMediaServer {
}
}

fn origin(&self) -> Option<String> {
self.active.read().ok().map(|active| active.origin.clone())
}

/// `true` means the origin moved, so a caller must drop urls held elsewhere.
pub(super) fn ensure_live(&self) -> std::io::Result<bool> {
{
let Ok(active) = self.active.read() else {
return Err(std::io::Error::other("loopback listener lock poisoned"));
};
if !active.dead.load(Ordering::Acquire) {
return Ok(false);
}
}

let Ok(mut active) = self.active.write() else {
return Err(std::io::Error::other("loopback listener lock poisoned"));
};
// Rebinding twice would orphan the origin a concurrent caller just handed out.
if !active.dead.load(Ordering::Acquire) {
return Ok(false);
}

let previous = std::mem::replace(&mut *active, spawn_listener(&self.routes)?);
drop(active);
retire(&previous);

Ok(true)
}

#[cfg(test)]
fn mark_dead(&self) {
if let Ok(active) = self.active.read() {
active.dead.store(true, Ordering::Release);
}
}

pub(super) fn redirect_response(
&self,
session: &MediaSession,
Expand All @@ -59,6 +151,9 @@ impl LoopbackMediaServer {
content_type: &str,
) -> Response<Vec<u8>> {
let capability = capability(session, cache_key);
let Some(origin) = self.origin() else {
return Response::new(Vec::new());
};
if let Ok(mut routes) = self.routes.write() {
routes.insert(
capability.clone(),
Expand All @@ -70,7 +165,7 @@ impl LoopbackMediaServer {
}
Response::builder()
.status(StatusCode::FOUND)
.header(header::LOCATION, format!("{}/{}", self.origin, capability))
.header(header::LOCATION, format!("{origin}/{capability}"))
.header(header::CACHE_CONTROL, "no-store")
.body(Vec::new())
.unwrap_or_else(|_| Response::new(Vec::new()))
Expand Down Expand Up @@ -264,6 +359,60 @@ mod tests {
parsed
}

#[test]
fn ensure_live_is_a_no_op_while_the_listener_is_healthy() {
let server = LoopbackMediaServer::start().expect("loopback test server");
let origin = server.origin().expect("origin");

assert!(!server.ensure_live().expect("ensure_live"));
assert_eq!(server.origin().as_deref(), Some(origin.as_str()));
}

#[test]
fn ensure_live_rebinds_a_dead_listener_and_keeps_the_route_table() {
let server = LoopbackMediaServer::start().expect("loopback test server");
let session = MediaSession {
origin: "https://matrix.example.org".into(),
token: "token".into(),
scope: "@alice:example.org".into(),
generation: 0,
};
let file = std::env::temp_dir().join("sable-loopback-rebind-test.bin");
std::fs::write(&file, b"media").expect("loopback test fixture");

let before = server.redirect_response(&session, "key", file.clone(), "image/png");
let first_location = before
.headers()
.get(header::LOCATION)
.and_then(|value| value.to_str().ok())
.expect("first location")
.to_owned();

server.mark_dead();
assert!(server.ensure_live().expect("ensure_live"));
let next_origin = server.origin().expect("origin after rebind");
assert!(!server.ensure_live().expect("ensure_live"));

let after = server.redirect_response(&session, "key", file, "image/png");
let second_location = after
.headers()
.get(header::LOCATION)
.and_then(|value| value.to_str().ok())
.expect("second location")
.to_owned();

assert_ne!(first_location, second_location);
assert!(second_location.starts_with(&next_origin));
let capability = capability(&session, "key");
assert!(first_location.ends_with(&capability));
assert!(second_location.ends_with(&capability));
assert!(server
.routes
.read()
.expect("routes")
.contains_key(&capability));
}

#[test]
fn extracts_range_and_origin_regardless_of_header_case() {
let (method, capability, range, origin) = parse(
Expand Down
36 changes: 36 additions & 0 deletions src/app/components/message/content/ImageContent.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,18 @@ import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { describe, expect, it, vi } from 'vitest';
import { ImageContent } from './ImageContent';
import { downloadEncryptedMedia, mxcUrlToHttp } from '$utils/matrix';
import type * as PlatformModule from '$utils/platform';

const screenMocks = vi.hoisted(() => ({
isMobile: true,
tauri: false,
loopbackUrl: undefined as string | undefined,
stripsCache: true,
}));

vi.mock('$utils/platform', async (importOriginal) => ({
...(await importOriginal<typeof PlatformModule>()),
webviewStripsCustomProtocolCache: () => screenMocks.stripsCache,
}));
vi.mock('$hooks/useScreenSize', () => ({
ScreenSize: { Desktop: 'Desktop', Tablet: 'Tablet', Mobile: 'Mobile' },
Expand Down Expand Up @@ -236,6 +243,35 @@ describe('ImageContent', () => {
}
});

it('loads a Tauri image from the custom protocol where its cache headers survive', async () => {
screenMocks.tauri = true;
screenMocks.stripsCache = false;
screenMocks.loopbackUrl = 'http://127.0.0.1:45678/capability';
try {
const srcs: string[] = [];
render(
<ImageContent
url="mxc://example.org/abc123"
renderImage={(props) => {
srcs.push(props.src);
return <img alt="preview" src={props.src} onError={props.onError} />;
}}
renderViewer={() => <div>viewer</div>}
/>
);

touchTap(screen.getByRole('button', { name: 'View' }));
await screen.findByAltText('preview');

await waitFor(() => expect(srcs.length).toBeGreaterThan(0));
expect(Array.from(new Set(srcs))).toEqual([SABLE_MEDIA_URL]);
} finally {
screenMocks.tauri = false;
screenMocks.stripsCache = true;
screenMocks.loopbackUrl = undefined;
}
});

it('passes the Tauri media URL straight to the encrypted download', async () => {
screenMocks.tauri = true;
const renderViewer = vi.fn<(props: { getDownloadBlob?: () => Promise<Blob> }) => ReactNode>(
Expand Down
6 changes: 4 additions & 2 deletions src/app/components/message/content/VideoContent.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -62,9 +62,11 @@ describe('VideoContent', () => {
fireEvent.click(screen.getByRole('button', { name: 'Watch' }));

const video = await screen.findByTestId('video');
expect(srcs[srcs.length - 1]).toMatch(/^http:\/\/127\.0\.0\.1:45678\//);
await waitFor(() => {
expect(srcs[srcs.length - 1]).toMatch(/^http:\/\/127\.0\.0\.1:45678\//);
expect(mocks.loopbackTargets[mocks.loopbackTargets.length - 1]).toBe(SABLE_MEDIA_URL);
});
const initialSrc = mocks.loopbackTargets[mocks.loopbackTargets.length - 1];
expect(initialSrc).toBe(SABLE_MEDIA_URL);

fireEvent.error(video);
fireEvent.click(await screen.findByRole('button', { name: 'Retry' }));
Expand Down
3 changes: 1 addition & 2 deletions src/app/components/room-avatar/RoomAvatar.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -25,8 +25,7 @@ type RoomAvatarProps = {
};

export function RoomAvatar({ roomId, src, alt, renderFallback, uniformIcons }: RoomAvatarProps) {
// Mirrors the crossOrigin AvatarImage puts on the rendered element.
const { mediaSrc, error, onError } = useAvatarMediaSource(src, { crossOrigin: 'anonymous' });
const { mediaSrc, error, onError } = useAvatarMediaSource(src);

if (!mediaSrc || error) {
return (
Expand Down
Loading
Loading