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
3 changes: 2 additions & 1 deletion src-tauri/gen/android/app/src/main/AndroidManifest.xml
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,8 @@
android:roundIcon="@mipmap/ic_launcher_round"
android:label="@string/app_name"
android:theme="@style/Theme.sable"
android:usesCleartextTraffic="${usesCleartextTraffic}">
android:usesCleartextTraffic="${usesCleartextTraffic}"
android:networkSecurityConfig="@xml/network_security_config">

<receiver
android:name="app.tauri.notification.TauriUnifiedPushMessagingService"
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
<?xml version="1.0" encoding="utf-8"?>
<network-security-config>
<domain-config cleartextTrafficPermitted="true">
<domain includeSubdomains="false">127.0.0.1</domain>
</domain-config>
</network-security-config>
2 changes: 2 additions & 0 deletions src-tauri/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -446,6 +446,8 @@ pub fn run() {
network::media_protocol::set_media_session,
network::media_protocol::clear_media_session,
network::media_protocol::set_media_encryption,
#[cfg(target_os = "android")]
network::media_protocol::prepare_loopback_video,
sentry::set_native_sentry_enabled,
share_inbox::share_inbox_drain,
share_inbox::share_inbox_read,
Expand Down
59 changes: 55 additions & 4 deletions src-tauri/src/network/media_protocol.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,15 @@ use tauri::{
AppHandle, Manager, Runtime, State, UriSchemeContext, UriSchemeResponder,
};

#[cfg(target_os = "android")]
mod android_loopback;
mod crypto;
mod lane;
mod response;
mod session;

#[cfg(target_os = "android")]
use android_loopback::LoopbackMediaServer;
use crypto::EncryptionStore;
use lane::{LanePermit, LifoLane};
use response::{
Expand Down Expand Up @@ -63,6 +67,8 @@ pub struct MediaSessionState {
download_lane: LifoLane,
cache_miss_gates: Mutex<HashMap<String, Weak<AsyncMutex<Option<FetchResult>>>>>,
negative_cache: Mutex<HashMap<String, (StatusCode, Instant)>>,
#[cfg(target_os = "android")]
loopback: Option<LoopbackMediaServer>,
}

impl Default for MediaSessionState {
Expand All @@ -75,6 +81,8 @@ impl Default for MediaSessionState {
download_lane: LifoLane::new(MAX_CONCURRENT_DOWNLOAD_REQUESTS),
cache_miss_gates: Mutex::new(HashMap::new()),
negative_cache: Mutex::new(HashMap::new()),
#[cfg(target_os = "android")]
loopback: LoopbackMediaServer::start().ok(),
}
}
}
Expand Down Expand Up @@ -105,12 +113,26 @@ impl MediaSessionState {
}

fn set_session(&self, session: MediaSession) -> Result<(), String> {
self.session_store
.set(session, || self.forget_client_errors())
self.session_store.set(session, || {
self.forget_client_errors();
#[cfg(target_os = "android")]
self.clear_loopback_media();
})
}

fn clear_session(&self) -> Result<(), String> {
self.session_store.clear(|| self.forget_client_errors())
self.session_store.clear(|| {
self.forget_client_errors();
#[cfg(target_os = "android")]
self.clear_loopback_media();
})
}

#[cfg(target_os = "android")]
fn clear_loopback_media(&self) {
if let Some(loopback) = &self.loopback {
loopback.clear();
}
}

// Shared across requests so the connection pool and TLS sessions stay warm.
Expand Down Expand Up @@ -240,6 +262,24 @@ pub fn set_media_encryption(
.register(&url, &key, &iv, &sha256, &version, mime_type)
}

#[cfg(target_os = "android")]
#[tauri::command]
pub async fn prepare_loopback_video<R: Runtime>(
app: AppHandle<R>,
url: String,
) -> Result<String, String> {
let uri = url.parse::<Uri>().map_err(|err| err.to_string())?;
let response = handle_request(&app, uri, None, true)
.await
.map_err(|status| status.to_string())?;
response
.headers()
.get(header::LOCATION)
.and_then(|value| value.to_str().ok())
.map(str::to_owned)
.ok_or_else(|| "loopback media server unavailable".to_owned())
}

fn media_session_marker(uri: &Uri) -> Result<Option<String>, StatusCode> {
let Some(query) = uri.query() else {
return Ok(None);
Expand Down Expand Up @@ -295,7 +335,7 @@ pub fn respond<R: Runtime>(
let range = header_value(header::RANGE);
let origin = header_value(header::ORIGIN);
tauri::async_runtime::spawn(async move {
let mut response = handle_request(&app, uri, range)
let mut response = handle_request(&app, uri, range, false)
.await
.unwrap_or_else(error_response);
apply_cors_headers(&mut response, origin.as_deref());
Expand All @@ -307,7 +347,11 @@ async fn handle_request<R: Runtime>(
app: &AppHandle<R>,
uri: Uri,
range: Option<String>,
loopback: bool,
) -> Result<Response<Vec<u8>>, StatusCode> {
#[cfg(not(target_os = "android"))]
let _ = loopback;

let target = percent_encoding::percent_decode_str(uri.path().trim_start_matches('/'))
.decode_utf8()
.map_err(|_| StatusCode::BAD_REQUEST)?
Expand Down Expand Up @@ -348,6 +392,13 @@ async fn handle_request<R: Runtime>(
let (content_type, in_memory_body, disk_path) =
ensure_cached(&state, &session, &key, media_url, dir, temp_dir).await?;

#[cfg(target_os = "android")]
if loopback && in_memory_body.is_none() && content_type.starts_with("video/") {
if let Some(loopback) = &state.loopback {
return Ok(loopback.redirect_response(&session, &key, disk_path, &content_type));
}
}

match (range, in_memory_body) {
(Some(range_header), Some(body)) => {
Ok(serve_range_memory(&body, &content_type, &range_header))
Expand Down
239 changes: 239 additions & 0 deletions src-tauri/src/network/media_protocol/android_loopback.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,239 @@
use std::{
collections::HashMap,
fs::File,
io::{BufRead, BufReader, Read, Seek, SeekFrom, Write},
net::{TcpListener, TcpStream},
path::PathBuf,
sync::{Arc, RwLock},
thread,
};

use sha2::{Digest, Sha256};
use tauri::http::{header, Response, StatusCode};

use super::session::MediaSession;

pub(super) struct LoopbackMediaServer {
origin: String,
routes: Arc<RwLock<HashMap<String, CachedMedia>>>,
}

#[derive(Clone)]
struct CachedMedia {
path: PathBuf,
content_type: String,
}

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 })
}

pub(super) fn clear(&self) {
if let Ok(mut routes) = self.routes.write() {
routes.clear();
}
}

pub(super) fn redirect_response(
&self,
session: &MediaSession,
cache_key: &str,
path: PathBuf,
content_type: &str,
) -> Response<Vec<u8>> {
let capability = capability(session, cache_key);
if let Ok(mut routes) = self.routes.write() {
routes.insert(
capability.clone(),
CachedMedia {
path,
content_type: content_type.to_owned(),
},
);
}
Response::builder()
.status(StatusCode::FOUND)
.header(header::LOCATION, format!("{}/{}", self.origin, capability))
.header(header::CACHE_CONTROL, "no-store")
.body(Vec::new())
.unwrap_or_else(|_| Response::new(Vec::new()))
}
}

fn capability(session: &MediaSession, cache_key: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(session.token.as_bytes());
hasher.update([0]);
hasher.update(session.scope.as_bytes());
hasher.update([0]);
hasher.update(cache_key.as_bytes());
hasher
.finalize()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}

fn serve(mut stream: TcpStream, routes: Arc<RwLock<HashMap<String, CachedMedia>>>) {
let Ok((method, capability, range)) = parse_request(&stream) else {
let _ = write_status(&mut stream, 400, "Bad Request", &[]);
return;
};
if method != "GET" && method != "HEAD" {
let _ = write_status(&mut stream, 405, "Method Not Allowed", &[]);
return;
}
let media = routes
.read()
.ok()
.and_then(|routes| routes.get(&capability).cloned());
let Some(media) = media else {
let _ = write_status(&mut stream, 404, "Not Found", &[]);
return;
};
let Ok(mut file) = File::open(media.path) else {
let _ = write_status(&mut stream, 404, "Not Found", &[]);
return;
};
let Ok(total) = file.metadata().map(|metadata| metadata.len()) else {
let _ = write_status(&mut stream, 500, "Internal Server Error", &[]);
return;
};
let selection = range.as_deref().and_then(|value| parse_range(value, total));
if range.is_some() && selection.is_none() {
let _ = write_status(
&mut stream,
416,
"Range Not Satisfiable",
&[("Content-Range", format!("bytes */{total}"))],
);
return;
}
let (start, end, partial) = selection.unwrap_or((0, total.saturating_sub(1), false));
let length = end.saturating_sub(start) + 1;
let mut headers = vec![
("Content-Type", media.content_type),
("Content-Length", length.to_string()),
("Accept-Ranges", "bytes".to_owned()),
(
"Access-Control-Allow-Origin",
"https://tauri.localhost".to_owned(),
),
(
"Cache-Control",
"private, max-age=31536000, immutable".to_owned(),
),
];
if partial {
headers.push(("Content-Range", format!("bytes {start}-{end}/{total}")));
}
if write_status(
&mut stream,
if partial { 206 } else { 200 },
if partial { "Partial Content" } else { "OK" },
&headers,
)
.is_err()
|| method == "HEAD"
{
return;
}
if file.seek(SeekFrom::Start(start)).is_err() {
return;
}
let mut left = length;
let mut buffer = [0_u8; 64 * 1024];
while left > 0 {
let want = left.min(buffer.len() as u64) as usize;
let Ok(read) = file.read(&mut buffer[..want]) else {
return;
};
if read == 0 || stream.write_all(&buffer[..read]).is_err() {
return;
}
left -= read as u64;
}
}

fn parse_request(stream: &TcpStream) -> Result<(String, String, Option<String>), ()> {
let mut reader = BufReader::new(stream);
let mut request_line = String::new();
reader.read_line(&mut request_line).map_err(|_| ())?;
let mut parts = request_line.split_whitespace();
let method = parts.next().ok_or(())?.to_owned();
let path = parts.next().ok_or(())?;
if parts.next().is_none() || !path.starts_with('/') || path[1..].contains('/') {
return Err(());
}
let mut range = None;
let mut line = String::new();
loop {
line.clear();
reader.read_line(&mut line).map_err(|_| ())?;
if line == "\r\n" || line.is_empty() {
break;
}
if let Some(value) = line
.strip_prefix("Range:")
.or_else(|| line.strip_prefix("range:"))
{
range = Some(value.trim().to_owned());
}
if line.len() > 8192 {
return Err(());
}
}
Ok((method, path[1..].to_owned(), range))
}

fn parse_range(value: &str, total: u64) -> Option<(u64, u64, bool)> {
let spec = value.strip_prefix("bytes=")?;
if spec.contains(',') || total == 0 {
return None;
}
let (start, end) = spec.split_once('-')?;
if start.is_empty() {
let length = end.parse::<u64>().ok()?.min(total);
(length > 0).then_some((total - length, total - 1, true))
} else {
let start = start.parse::<u64>().ok()?;
let end = if end.is_empty() {
total.checked_sub(1)?
} else {
end.parse::<u64>().ok()?.min(total.checked_sub(1)?)
};
(start <= end && start < total).then_some((start, end, true))
}
}

fn write_status(
stream: &mut TcpStream,
status: u16,
reason: &str,
headers: &[(&str, String)],
) -> std::io::Result<()> {
write!(
stream,
"HTTP/1.1 {status} {reason}\r\nConnection: close\r\n"
)?;
for (name, value) in headers {
write!(stream, "{name}: {value}\r\n")?;
}
stream.write_all(b"\r\n")
}
Loading
Loading