From 644c9bfda38be2339b9b9e5dc360ac098ad6d320 Mon Sep 17 00:00:00 2001 From: Hammed Abass Date: Wed, 12 Aug 2026 15:06:01 +0200 Subject: [PATCH] fix(core): isolate targeted transfers and persist peer names --- crates/vnidrop/src/api.rs | 1 + crates/vnidrop/src/approval.rs | 13 +- .../src/device_relationship/service.rs | 16 +- .../vnidrop/src/device_relationship/store.rs | 57 ++- crates/vnidrop/src/handshake.rs | 1 + crates/vnidrop/src/invitation/mod.rs | 87 +++-- crates/vnidrop/src/pairing_eligibility/mod.rs | 26 +- .../vnidrop/src/pairing_eligibility/store.rs | 35 +- crates/vnidrop/src/runtime/facade.rs | 25 ++ crates/vnidrop/src/runtime/mod.rs | 5 + crates/vnidrop/src/runtime/receive.rs | 234 ++++++++++-- crates/vnidrop/src/runtime/share.rs | 14 + crates/vnidrop/src/runtime/targeted.rs | 342 ++++++++++++++---- crates/vnidrop/src/targeted_transfer/store.rs | 40 ++ .../vnidrop/src/tests/device_relationship.rs | 123 ++++++- .../vnidrop/src/tests/pairing_eligibility.rs | 38 ++ crates/vnidrop/src/tests/persistence.rs | 43 +++ .../src/tests/platform_contract_android.rs | 9 + crates/vnidrop/src/tests/repository.rs | 4 + crates/vnidrop/src/tests/runtime.rs | 3 + crates/vnidrop/src/tests/targeted_transfer.rs | 123 +++++++ crates/vnidrop/src/ticket.rs | 7 + 22 files changed, 1091 insertions(+), 155 deletions(-) diff --git a/crates/vnidrop/src/api.rs b/crates/vnidrop/src/api.rs index 6153677..8020a9a 100644 --- a/crates/vnidrop/src/api.rs +++ b/crates/vnidrop/src/api.rs @@ -33,6 +33,7 @@ pub fn experimental_saved_device_capabilities() -> ExperimentalSavedDeviceCapabi #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, uniffi::Record)] pub struct PairingEligibilitySummary { pub peer_endpoint_id: String, + pub remote_display_name: Option, pub session_id: String, pub protocol_version: u16, pub created_at: i64, diff --git a/crates/vnidrop/src/approval.rs b/crates/vnidrop/src/approval.rs index 82c7357..7371a12 100644 --- a/crates/vnidrop/src/approval.rs +++ b/crates/vnidrop/src/approval.rs @@ -58,11 +58,12 @@ impl ApprovalService { ) .await { - Ok(()) => { + Ok(remote_display_name) => { if let Some(eligibility) = &self.pairing_eligibility { if let Err(error) = eligibility .activate_after_completed_transfer( &remote_endpoint_id, + remote_display_name.as_deref(), &receipt.request_id, &receipt.token, ) @@ -313,6 +314,11 @@ impl ApprovalService { request_id, token, expires_at, + sender_name: self + .repository + .send_sender_name(request.transfer_id, &request.transfer_hash) + .await + .unwrap_or(None), } } @@ -408,6 +414,11 @@ impl ApprovalService { request_id: decision.request_id, token, expires_at, + sender_name: self + .repository + .send_sender_name(request.transfer_id, &request.transfer_hash) + .await + .unwrap_or(None), } } Ok(Ok(decision)) => { diff --git a/crates/vnidrop/src/device_relationship/service.rs b/crates/vnidrop/src/device_relationship/service.rs index 71b41f0..54d84b1 100644 --- a/crates/vnidrop/src/device_relationship/service.rs +++ b/crates/vnidrop/src/device_relationship/service.rs @@ -242,6 +242,8 @@ impl DeviceRelationshipService { generation, minimum_protocol_version: taken.protocol_version, session_id: Some(&taken.session_id), + remote_display_name: taken.remote_display_name.as_deref(), + last_authenticated_at: Some(taken.authenticated_at), issued_grant_handle: None, held_grant_handle: None, issued_grant_id: None, @@ -271,6 +273,8 @@ impl DeviceRelationshipService { generation, minimum_protocol_version: taken.protocol_version, session_id: Some(&taken.session_id), + remote_display_name: taken.remote_display_name.as_deref(), + last_authenticated_at: Some(taken.authenticated_at), issued_grant_handle: None, held_grant_handle: None, issued_grant_id: None, @@ -501,17 +505,17 @@ impl DeviceRelationshipService { if request.protocol_version != local_protocol { return PairingRequestResponse::Rejected; } - let accepted = match self + let local_observation = match self .eligibility .validate_presented_capability(&remote_endpoint_id, &request.session_id, &capability) .await { - Ok(Some(_)) => true, - Ok(None) | Err(_) => false, + Ok(Some(entry)) => Some(entry.remote_display_name), + Ok(None) | Err(_) => None, }; - if !accepted { + let Some(remote_display_name) = local_observation else { return PairingRequestResponse::Rejected; - } + }; match self.can_create_new_relationship(&remote_endpoint_id).await { Ok(true) => {} @@ -527,6 +531,8 @@ impl DeviceRelationshipService { generation: request.generation, minimum_protocol_version: request.protocol_version, session_id: Some(&request.session_id), + remote_display_name: remote_display_name.as_deref(), + last_authenticated_at: Some(now), issued_grant_handle: None, held_grant_handle: None, issued_grant_id: None, diff --git a/crates/vnidrop/src/device_relationship/store.rs b/crates/vnidrop/src/device_relationship/store.rs index bd2eccb..e7fc53f 100644 --- a/crates/vnidrop/src/device_relationship/store.rs +++ b/crates/vnidrop/src/device_relationship/store.rs @@ -54,6 +54,8 @@ pub(super) struct RelationshipUpsert<'a> { pub(super) generation: u64, pub(super) minimum_protocol_version: u16, pub(super) session_id: Option<&'a str>, + pub(super) remote_display_name: Option<&'a str>, + pub(super) last_authenticated_at: Option, pub(super) issued_grant_handle: Option<&'a str>, pub(super) held_grant_handle: Option<&'a str>, pub(super) issued_grant_id: Option<&'a str>, @@ -84,6 +86,8 @@ impl DeviceRelationshipStore { generation INTEGER NOT NULL, minimum_protocol_version INTEGER NOT NULL, session_id TEXT, + remote_display_name TEXT, + last_authenticated_at INTEGER, issued_grant_handle TEXT, held_grant_handle TEXT, issued_grant_id TEXT, @@ -116,6 +120,18 @@ impl DeviceRelationshipStore { .execute(pool) .await?; } + if !has("remote_display_name") { + sqlx::query("ALTER TABLE device_relationships ADD COLUMN remote_display_name TEXT") + .execute(pool) + .await?; + } + if !has("last_authenticated_at") { + sqlx::query( + "ALTER TABLE device_relationships ADD COLUMN last_authenticated_at INTEGER", + ) + .execute(pool) + .await?; + } sqlx::query( r#" CREATE TABLE IF NOT EXISTS relationship_generation_tombstones ( @@ -182,10 +198,11 @@ impl DeviceRelationshipStore { rows.into_iter().map(row_to_relationship).collect() } - pub(super) async fn list_saved_devices(&self) -> Result, VnidropError> { + pub(crate) async fn list_saved_devices(&self) -> Result, VnidropError> { let rows = sqlx::query( r#" - SELECT remote_endpoint_id, local_label, created_at, updated_at + SELECT remote_endpoint_id, local_label, remote_display_name, created_at, + last_authenticated_at FROM device_relationships WHERE state = 'saved' ORDER BY updated_at DESC @@ -199,13 +216,37 @@ impl DeviceRelationshipStore { .map(|row| SavedDevice { endpoint_id: row.get("remote_endpoint_id"), local_label: row.get("local_label"), - remote_display_name: None, + remote_display_name: row.get("remote_display_name"), created_at: row.get("created_at"), - last_authenticated_at: Some(row.get("updated_at")), + last_authenticated_at: row.get("last_authenticated_at"), }) .collect()) } + pub(crate) async fn refresh_authenticated_peer( + &self, + peer_endpoint_id: &str, + remote_display_name: Option<&str>, + authenticated_at: i64, + ) -> Result { + let result = sqlx::query( + r#" + UPDATE device_relationships + SET remote_display_name = COALESCE(?2, remote_display_name), + last_authenticated_at = ?3, + updated_at = ?3 + WHERE remote_endpoint_id = ?1 AND state = 'saved' + "#, + ) + .bind(peer_endpoint_id) + .bind(remote_display_name) + .bind(authenticated_at) + .execute(&self.pool) + .await + .map_err(VnidropError::repository)?; + Ok(result.rows_affected() > 0) + } + pub(super) async fn set_saved_device_label( &self, peer_endpoint_id: &str, @@ -377,14 +418,18 @@ impl DeviceRelationshipStore { r#" INSERT INTO device_relationships ( remote_endpoint_id, state, generation, minimum_protocol_version, session_id, + remote_display_name, + last_authenticated_at, issued_grant_handle, held_grant_handle, issued_grant_id, held_grant_id, peer_ack, local_ack, created_at, updated_at - ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13) + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15) ON CONFLICT(remote_endpoint_id) DO UPDATE SET state = excluded.state, generation = excluded.generation, minimum_protocol_version = excluded.minimum_protocol_version, session_id = excluded.session_id, + remote_display_name = COALESCE(excluded.remote_display_name, device_relationships.remote_display_name), + last_authenticated_at = COALESCE(excluded.last_authenticated_at, device_relationships.last_authenticated_at), issued_grant_handle = COALESCE(excluded.issued_grant_handle, device_relationships.issued_grant_handle), held_grant_handle = COALESCE(excluded.held_grant_handle, device_relationships.held_grant_handle), issued_grant_id = COALESCE(excluded.issued_grant_id, device_relationships.issued_grant_id), @@ -399,6 +444,8 @@ impl DeviceRelationshipStore { .bind(entry.generation as i64) .bind(i64::from(entry.minimum_protocol_version)) .bind(entry.session_id) + .bind(entry.remote_display_name) + .bind(entry.last_authenticated_at) .bind(entry.issued_grant_handle) .bind(entry.held_grant_handle) .bind(entry.issued_grant_id) diff --git a/crates/vnidrop/src/handshake.rs b/crates/vnidrop/src/handshake.rs index d48f4a3..eff6075 100644 --- a/crates/vnidrop/src/handshake.rs +++ b/crates/vnidrop/src/handshake.rs @@ -142,6 +142,7 @@ pub(crate) enum HandshakeResponse { request_id: String, token: String, expires_at: i64, + sender_name: Option, }, Denied { reason: String, diff --git a/crates/vnidrop/src/invitation/mod.rs b/crates/vnidrop/src/invitation/mod.rs index ff0e2f4..b31c361 100644 --- a/crates/vnidrop/src/invitation/mod.rs +++ b/crates/vnidrop/src/invitation/mod.rs @@ -44,6 +44,7 @@ pub(crate) struct TransferUpsert<'a> { pub(crate) direction: TransferDirection, pub(crate) status: TransferStatus, pub(crate) transfer_name: Option<&'a str>, + pub(crate) sender_name: Option<&'a str>, pub(crate) content_hash: Option<&'a str>, pub(crate) ticket: Option<&'a str>, pub(crate) file_count: u64, @@ -136,6 +137,7 @@ impl Repository { direction TEXT NOT NULL, status TEXT NOT NULL, transfer_name TEXT, + sender_name TEXT, content_hash TEXT, ticket TEXT, file_count INTEGER NOT NULL DEFAULT 0, @@ -155,6 +157,14 @@ impl Repository { let has_access_mode = columns .iter() .any(|row| row.get::(1) == "access_mode"); + if !columns + .iter() + .any(|row| row.get::(1) == "sender_name") + { + sqlx::query("ALTER TABLE transfers ADD COLUMN sender_name TEXT") + .execute(&self.pool) + .await?; + } if !has_access_mode { sqlx::query( "ALTER TABLE transfers ADD COLUMN access_mode TEXT NOT NULL DEFAULT 'approval_required'", @@ -501,19 +511,21 @@ impl Repository { UPDATE transfers SET status = ?1, transfer_name = ?2, - content_hash = ?3, - ticket = ?4, - file_count = ?5, - total_size = ?6, - access_mode = ?7, - updated_at = ?8 - WHERE transfer_id = ?9 + sender_name = ?3, + content_hash = ?4, + ticket = ?5, + file_count = ?6, + total_size = ?7, + access_mode = ?8, + updated_at = ?9 + WHERE transfer_id = ?10 AND direction = 'send' AND status = 'importing' "#, ) .bind(transfer.status.as_str()) .bind(transfer.transfer_name) + .bind(transfer.sender_name) .bind(transfer.content_hash) .bind(transfer.ticket) .bind(to_db_id(transfer.file_count)?) @@ -910,7 +922,7 @@ impl Repository { transfer_id: u64, remote_endpoint_id: &str, token_hash: &str, - ) -> Result<()> { + ) -> Result> { let result = sqlx::query( r#" UPDATE receiver_requests @@ -926,31 +938,45 @@ impl Repository { .bind(token_hash) .execute(&self.pool) .await?; - if result.rows_affected() == 1 { - return Ok(()); - } - let already_recorded = sqlx::query( - r#" + if result.rows_affected() != 1 { + let already_recorded = sqlx::query( + r#" SELECT EXISTS( SELECT 1 FROM receiver_requests WHERE id = ?1 AND transfer_id = ?2 AND remote_endpoint_id = ?3 AND receipt_token_hash = ?4 AND status = 'completed' ) "#, + ) + .bind(id) + .bind(to_db_id(transfer_id)?) + .bind(remote_endpoint_id) + .bind(token_hash) + .fetch_one(&self.pool) + .await? + .get::(0) + != 0; + if !already_recorded { + anyhow::bail!("delivery receipt did not match an accepted receiver request"); + } + } + let row = sqlx::query( + r#" + SELECT receiver_name, receiver_device_name + FROM receiver_requests + WHERE id = ?1 AND transfer_id = ?2 AND remote_endpoint_id = ?3 + AND receipt_token_hash = ?4 AND status = 'completed' + "#, ) .bind(id) .bind(to_db_id(transfer_id)?) .bind(remote_endpoint_id) .bind(token_hash) .fetch_one(&self.pool) - .await? - .get::(0) - != 0; - if already_recorded { - Ok(()) - } else { - anyhow::bail!("delivery receipt did not match an accepted receiver request") - } + .await?; + Ok(row + .get::, _>("receiver_device_name") + .or_else(|| row.get("receiver_name"))) } pub(crate) async fn fail_receiver_delivery( @@ -1057,6 +1083,25 @@ impl Repository { Ok(row.get::(0) != 0) } + pub(crate) async fn send_sender_name( + &self, + transfer_id: u64, + content_hash: &str, + ) -> Result> { + let row = sqlx::query( + r#" + SELECT sender_name FROM transfers + WHERE transfer_id = ?1 AND content_hash = ?2 + AND direction = 'send' AND status = 'sharing' + "#, + ) + .bind(to_db_id(transfer_id)?) + .bind(content_hash) + .fetch_optional(&self.pool) + .await?; + Ok(row.and_then(|row| row.get("sender_name"))) + } + pub(crate) async fn list_transfers(&self) -> Result> { let rows = sqlx::query( r#" diff --git a/crates/vnidrop/src/pairing_eligibility/mod.rs b/crates/vnidrop/src/pairing_eligibility/mod.rs index 5c7f2f4..1a012d6 100644 --- a/crates/vnidrop/src/pairing_eligibility/mod.rs +++ b/crates/vnidrop/src/pairing_eligibility/mod.rs @@ -15,6 +15,7 @@ pub(crate) use store::PairingEligibilityStore; use crate::{ api::{experimental_saved_device_capabilities, PairingEligibilitySummary}, + device_relationship::DeviceRelationshipStore, error::VnidropError, event_hub::EventHub, secure_secret::{SecretCustody, SecretHandle, SecretKind, SecretMaterial}, @@ -27,6 +28,7 @@ const CAPABILITY_CONTEXT: &str = "vnidrop-pairing-eligibility-v1"; #[derive(Clone)] pub(crate) struct PairingEligibilityService { store: PairingEligibilityStore, + relationships: DeviceRelationshipStore, custody: Option>, event_hub: Arc, local_endpoint_id: String, @@ -35,12 +37,14 @@ pub(crate) struct PairingEligibilityService { impl PairingEligibilityService { pub(crate) fn new( store: PairingEligibilityStore, + relationships: DeviceRelationshipStore, custody: Option>, event_hub: Arc, local_endpoint_id: String, ) -> Self { Self { store, + relationships, custody, event_hub, local_endpoint_id, @@ -83,6 +87,7 @@ impl PairingEligibilityService { pub(crate) async fn activate_after_completed_transfer( &self, peer_endpoint_id: &str, + remote_display_name: Option<&str>, session_id: &str, approval_token: &str, ) -> Result<(), VnidropError> { @@ -92,6 +97,18 @@ impl PairingEligibilityService { if peer_endpoint_id.is_empty() || session_id.is_empty() || approval_token.is_empty() { return Ok(()); } + let remote_display_name = remote_display_name + .map(str::trim) + .filter(|name| !name.is_empty()); + let authenticated_at = now_ms(); + if self + .relationships + .refresh_authenticated_peer(peer_endpoint_id, remote_display_name, authenticated_at) + .await? + { + self.remove_for_peer(peer_endpoint_id).await?; + return Ok(()); + } if self.store.find_by_session(session_id).await?.is_some() { return Ok(()); } @@ -111,12 +128,13 @@ impl PairingEligibilityService { let handle = custody .protect(SecretKind::PairingEligibility, capability, None) .await?; - let created_at = now_ms(); + let created_at = authenticated_at; let expires_at = created_at + ELIGIBILITY_TTL_MS; if let Err(error) = self .store .insert(PairingEligibilityInsert { peer_endpoint_id, + remote_display_name, session_id, protocol_version, secret_handle: handle.as_str(), @@ -189,6 +207,8 @@ impl PairingEligibilityService { Ok(Some(TakenEligibility { session_id: entry.session_id, protocol_version: entry.protocol_version, + remote_display_name: entry.remote_display_name, + authenticated_at: entry.created_at, capability, })) } @@ -327,6 +347,7 @@ impl PairingEligibilityService { pub(crate) struct PairingEligibilityInsert<'a> { pub(crate) peer_endpoint_id: &'a str, + pub(crate) remote_display_name: Option<&'a str>, pub(crate) session_id: &'a str, pub(crate) protocol_version: u16, pub(crate) secret_handle: &'a str, @@ -338,12 +359,15 @@ pub(crate) struct PairingEligibilityInsert<'a> { pub(crate) struct TakenEligibility { pub(crate) session_id: String, pub(crate) protocol_version: u16, + pub(crate) remote_display_name: Option, + pub(crate) authenticated_at: i64, pub(crate) capability: SecretMaterial, } #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct PairingEligibilityRecord { pub(crate) peer_endpoint_id: String, + pub(crate) remote_display_name: Option, pub(crate) session_id: String, pub(crate) protocol_version: u16, pub(crate) secret_handle: String, diff --git a/crates/vnidrop/src/pairing_eligibility/store.rs b/crates/vnidrop/src/pairing_eligibility/store.rs index 2bcac15..19c4d41 100644 --- a/crates/vnidrop/src/pairing_eligibility/store.rs +++ b/crates/vnidrop/src/pairing_eligibility/store.rs @@ -23,6 +23,7 @@ impl PairingEligibilityStore { CREATE TABLE IF NOT EXISTS pairing_eligibilities ( session_id TEXT PRIMARY KEY, peer_endpoint_id TEXT NOT NULL, + remote_display_name TEXT, protocol_version INTEGER NOT NULL, secret_handle TEXT NOT NULL UNIQUE, created_at INTEGER NOT NULL, @@ -32,6 +33,17 @@ impl PairingEligibilityStore { ) .execute(pool) .await?; + let columns = sqlx::query("PRAGMA table_info(pairing_eligibilities)") + .fetch_all(pool) + .await?; + if !columns + .iter() + .any(|row| row.get::(1) == "remote_display_name") + { + sqlx::query("ALTER TABLE pairing_eligibilities ADD COLUMN remote_display_name TEXT") + .execute(pool) + .await?; + } sqlx::query( r#" CREATE INDEX IF NOT EXISTS pairing_eligibilities_peer @@ -50,12 +62,14 @@ impl PairingEligibilityStore { sqlx::query( r#" INSERT INTO pairing_eligibilities ( - session_id, peer_endpoint_id, protocol_version, secret_handle, created_at, expires_at - ) VALUES (?1, ?2, ?3, ?4, ?5, ?6) + session_id, peer_endpoint_id, remote_display_name, protocol_version, + secret_handle, created_at, expires_at + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7) "#, ) .bind(entry.session_id) .bind(entry.peer_endpoint_id) + .bind(entry.remote_display_name) .bind(i64::from(entry.protocol_version)) .bind(entry.secret_handle) .bind(entry.created_at) @@ -71,7 +85,8 @@ impl PairingEligibilityStore { ) -> Result, VnidropError> { let rows = sqlx::query( r#" - SELECT peer_endpoint_id, session_id, protocol_version, created_at, expires_at + SELECT peer_endpoint_id, remote_display_name, session_id, protocol_version, + created_at, expires_at FROM pairing_eligibilities ORDER BY created_at DESC "#, @@ -83,6 +98,7 @@ impl PairingEligibilityStore { .into_iter() .map(|row| PairingEligibilitySummary { peer_endpoint_id: row.get("peer_endpoint_id"), + remote_display_name: row.get("remote_display_name"), session_id: row.get("session_id"), protocol_version: row.get::("protocol_version") as u16, created_at: row.get("created_at"), @@ -94,7 +110,8 @@ impl PairingEligibilityStore { pub(crate) async fn list_records(&self) -> Result, VnidropError> { let rows = sqlx::query( r#" - SELECT peer_endpoint_id, session_id, protocol_version, secret_handle, created_at, expires_at + SELECT peer_endpoint_id, remote_display_name, session_id, protocol_version, + secret_handle, created_at, expires_at FROM pairing_eligibilities "#, ) @@ -110,7 +127,8 @@ impl PairingEligibilityStore { ) -> Result, VnidropError> { let rows = sqlx::query( r#" - SELECT peer_endpoint_id, session_id, protocol_version, secret_handle, created_at, expires_at + SELECT peer_endpoint_id, remote_display_name, session_id, protocol_version, + secret_handle, created_at, expires_at FROM pairing_eligibilities WHERE peer_endpoint_id = ?1 "#, @@ -128,7 +146,8 @@ impl PairingEligibilityStore { ) -> Result, VnidropError> { let rows = sqlx::query( r#" - SELECT peer_endpoint_id, session_id, protocol_version, secret_handle, created_at, expires_at + SELECT peer_endpoint_id, remote_display_name, session_id, protocol_version, + secret_handle, created_at, expires_at FROM pairing_eligibilities WHERE expires_at <= ?1 "#, @@ -146,7 +165,8 @@ impl PairingEligibilityStore { ) -> Result, VnidropError> { let row = sqlx::query( r#" - SELECT peer_endpoint_id, session_id, protocol_version, secret_handle, created_at, expires_at + SELECT peer_endpoint_id, remote_display_name, session_id, protocol_version, + secret_handle, created_at, expires_at FROM pairing_eligibilities WHERE session_id = ?1 "#, @@ -186,6 +206,7 @@ impl PairingEligibilityStore { fn row_to_record(row: sqlx::sqlite::SqliteRow) -> PairingEligibilityRecord { PairingEligibilityRecord { peer_endpoint_id: row.get("peer_endpoint_id"), + remote_display_name: row.get("remote_display_name"), session_id: row.get("session_id"), protocol_version: row.get::("protocol_version") as u16, secret_handle: row.get("secret_handle"), diff --git a/crates/vnidrop/src/runtime/facade.rs b/crates/vnidrop/src/runtime/facade.rs index 6f9b4ba..abbd733 100644 --- a/crates/vnidrop/src/runtime/facade.rs +++ b/crates/vnidrop/src/runtime/facade.rs @@ -99,6 +99,31 @@ impl VnidropCore { ) } + pub(crate) fn targeted_blob_ticket_for_test( + &self, + id: String, + ) -> Result<(u64, String), VnidropError> { + self.block_on(async { + let row = self + .inner + .targeted_store() + .get_row(&id) + .await? + .ok_or_else(|| VnidropError::invalid_input(anyhow::anyhow!("unknown transfer")))?; + let encoded = self + .inner + .load_stored_authorization(&row) + .await? + .ok_or_else(|| { + VnidropError::invalid_input(anyhow::anyhow!("authorization unavailable")) + })?; + Ok(( + row.protocol_transfer_id, + crate::targeted_transfer::TargetedAuthorization::decode(&encoded)?.blob_ticket, + )) + }) + } + pub(crate) fn initialize_with_test_secret_store_limits_and_network( app_data_dir: String, event_sink: Arc, diff --git a/crates/vnidrop/src/runtime/mod.rs b/crates/vnidrop/src/runtime/mod.rs index b940b6a..fcc9914 100644 --- a/crates/vnidrop/src/runtime/mod.rs +++ b/crates/vnidrop/src/runtime/mod.rs @@ -359,6 +359,7 @@ impl CoreInner { } let pairing_eligibility = PairingEligibilityService::new( stores.eligibility.clone(), + stores.relationships.clone(), secret_custody.clone(), event_hub.clone(), endpoint.id().to_string(), @@ -575,3 +576,7 @@ fn relay_mode_label(relay_mode: CoreRelayMode) -> &'static str { pub(super) fn share_tag_name(local_id: &str) -> String { format!("vnidrop/share/{local_id}") } + +pub(super) fn targeted_tag_name(transfer_id: &str) -> String { + format!("vnidrop/targeted/{transfer_id}") +} diff --git a/crates/vnidrop/src/runtime/receive.rs b/crates/vnidrop/src/runtime/receive.rs index bf986de..d492154 100644 --- a/crates/vnidrop/src/runtime/receive.rs +++ b/crates/vnidrop/src/runtime/receive.rs @@ -43,7 +43,7 @@ pub(super) enum ReceiveTarget { #[derive(Clone, Copy)] struct ReceivedTransfer<'a> { protocol_id: u64, - local_id: &'a str, + local_id: Option<&'a str>, } pub(super) struct OutputSinkFile<'a> { @@ -182,6 +182,117 @@ impl CoreInner { .await } + pub(super) async fn receive_targeted_payload( + self: &Arc, + transfer_id: u64, + mut blob_ticket: BlobTicket, + target: ReceiveTarget, + ) -> Result<()> { + let _permit = self + .transfer_slots + .acquire() + .await + .context("transfer limiter is closed") + .map_err(VnidropError::internal)?; + let sender_addr = filter_peer_addr_for_relay_mode( + blob_ticket.addr(), + self.relay_mode, + &self.custom_relay_urls, + ) + .map_err(VnidropError::network)?; + blob_ticket = BlobTicket::new(sender_addr, blob_ticket.hash(), blob_ticket.format()); + let (cancel, mut cancelled) = oneshot::channel(); + self.active_transfers + .lock() + .expect("active_transfers") + .insert( + transfer_id, + ActiveTransfer { + direction: TransferDirection::Receive, + cancel, + }, + ); + let result = tokio::select! { + result = self.download_targeted_payload(transfer_id, blob_ticket, target) => result, + _ = &mut cancelled => Err(VnidropError::cancelled("transfer cancelled").into()), + }; + self.active_transfers + .lock() + .expect("active_transfers") + .remove(&transfer_id); + result + } + + async fn download_targeted_payload( + &self, + transfer_id: u64, + blob_ticket: BlobTicket, + target: ReceiveTarget, + ) -> Result<()> { + if let ReceiveTarget::Directory(output_dir) = &target { + tokio::fs::create_dir_all(output_dir) + .await + .map_err(VnidropError::filesystem)?; + } + let connection = self + .endpoint + .connect(blob_ticket.addr().clone(), iroh_blobs::ALPN) + .await + .map_err(VnidropError::network)?; + let hash_and_format = blob_ticket.hash_and_format(); + let (_hash_seq, sizes) = + get_hash_seq_and_sizes(&connection, &hash_and_format.hash, 1024 * 1024 * 32, None) + .await + .context("failed to get targeted payload sizes") + .map_err(VnidropError::network)?; + let total_size = sizes + .iter() + .try_fold(0u64, |total, size| total.checked_add(*size)) + .context("remote collection size overflow")?; + let total_files = sizes.len().saturating_sub(1) as u64; + if total_files > self.limits.max_collection_files { + anyhow::bail!( + "remote collection has {total_files} files, limit is {}", + self.limits.max_collection_files + ); + } + if total_size > self.limits.max_total_bytes { + anyhow::bail!( + "remote collection size {total_size} exceeds limit {}", + self.limits.max_total_bytes + ); + } + let download_tag = self.store.tags().temp_tag(hash_and_format).await?; + let get = self.store.remote().fetch(connection, hash_and_format); + let mut stream = get.stream(); + loop { + let Some(item) = stream.next().await else { + anyhow::bail!("targeted download ended without completion"); + }; + match item { + GetProgressItem::Progress(downloaded) => self.emit_transfer( + transfer_id, + "receive", + "download", + "progress", + json!({ "downloaded": downloaded, "total_size": total_size }), + ), + GetProgressItem::Done(_) => break, + GetProgressItem::Error(error) => { + return Err(VnidropError::network(anyhow::anyhow!( + "targeted download failed: {error}" + )) + .into()); + } + } + } + let collection = Collection::load(hash_and_format.hash, self.store.as_ref()).await?; + self.export_collection_untracked(transfer_id, total_files, target, collection) + .await?; + drop(download_tag); + Ok(()) + } + pub(super) async fn receive_to_target( self: &Arc, ticket: String, @@ -328,7 +439,7 @@ impl CoreInner { self.emit_transfer(transfer_id, "receive", "network", "connecting", json!({})); // Every VniDrop ticket carries metadata and must complete the handshake. - let delivery_receipt = self + let (delivery_receipt, authenticated_sender_name) = self .request_transfer_approval( transfer_id, sender_addr.clone(), @@ -424,6 +535,7 @@ impl CoreInner { .pairing_eligibility .activate_after_completed_transfer( &peer_endpoint_id, + authenticated_sender_name.as_deref(), &delivery_receipt.request_id, &delivery_receipt.token, ) @@ -455,6 +567,7 @@ impl CoreInner { direction: TransferDirection::Receive, status: TransferStatus::Receiving, transfer_name: Some(parsed.metadata.transfer_name.as_str()), + sender_name: parsed.metadata.sender_name.as_deref(), content_hash: Some(parsed.metadata.content_hash.as_str()), ticket: None, file_count: parsed.metadata.file_count, @@ -484,7 +597,7 @@ impl CoreInner { addr: iroh::EndpointAddr, metadata: &TransferMetadata, receiver_name: Option<&str>, - ) -> Result { + ) -> Result<(DeliveryReceipt, Option)> { self.emit_transfer( local_transfer_id, "receive", @@ -507,6 +620,7 @@ impl CoreInner { request_id, token, expires_at, + sender_name, } => { self.emit_transfer( local_transfer_id, @@ -518,11 +632,14 @@ impl CoreInner { "expires_at": expires_at, }), ); - Ok(DeliveryReceipt { - request_id, - transfer_id: metadata.transfer_id, - token, - }) + Ok(( + DeliveryReceipt { + request_id, + transfer_id: metadata.transfer_id, + token, + }, + sender_name, + )) } HandshakeResponse::Denied { reason } => Err(VnidropError::permission(anyhow::anyhow!( "transfer request was denied by sender: {reason}" @@ -545,7 +662,58 @@ impl CoreInner { .map_err(VnidropError::repository)?; let received_transfer = ReceivedTransfer { protocol_id: transfer_id, - local_id: &transfer_local_id, + local_id: Some(&transfer_local_id), + }; + for (i, (name, hash)) in collection.iter().enumerate() { + match &target { + ReceiveTarget::Directory(output_dir) => { + self.export_blob_to_directory( + received_transfer, + total_files, + i as u64, + output_dir, + name.as_ref(), + *hash, + ) + .await?; + } + ReceiveTarget::OutputSink(output_sink) => { + self.export_blob_to_sink( + transfer_id, + total_files, + i as u64, + output_sink.as_ref(), + name.as_ref(), + *hash, + ) + .await?; + } + ReceiveTarget::OutputSinkV2(output_sink) => { + self.export_blob_to_sink_v2( + received_transfer, + total_files, + i as u64, + output_sink.as_ref(), + name.as_ref(), + *hash, + ) + .await?; + } + } + } + Ok(()) + } + + async fn export_collection_untracked( + &self, + transfer_id: u64, + total_files: u64, + target: ReceiveTarget, + collection: Collection, + ) -> Result<()> { + let received_transfer = ReceivedTransfer { + protocol_id: transfer_id, + local_id: None, }; for (i, (name, hash)) in collection.iter().enumerate() { match &target { @@ -656,17 +824,19 @@ impl CoreInner { .map_err(VnidropError::filesystem)?; let locator = pending_file.target().to_string_lossy().to_string(); pending_file.commit().map_err(VnidropError::filesystem)?; - self.repository - .record_received_artifact(ReceivedArtifactInsert { - transfer_local_id: transfer.local_id, - protocol_transfer_id: transfer.protocol_id, - relative_path, - locator_kind: ReceivedLocatorKind::FilesystemPath, - locator: &locator, - logical_size: exported, - }) - .await - .map_err(VnidropError::repository)?; + if let Some(local_id) = transfer.local_id { + self.repository + .record_received_artifact(ReceivedArtifactInsert { + transfer_local_id: local_id, + protocol_transfer_id: transfer.protocol_id, + relative_path, + locator_kind: ReceivedLocatorKind::FilesystemPath, + locator: &locator, + logical_size: exported, + }) + .await + .map_err(VnidropError::repository)?; + } Ok(()) } @@ -784,17 +954,19 @@ impl CoreInner { } let published = output_file.finish()?; - self.repository - .record_received_artifact(ReceivedArtifactInsert { - transfer_local_id: transfer.local_id, - protocol_transfer_id: transfer.protocol_id, - relative_path: &relative_path, - locator_kind: published.locator_kind, - locator: &published.locator, - logical_size: exported, - }) - .await - .map_err(VnidropError::repository)?; + if let Some(local_id) = transfer.local_id { + self.repository + .record_received_artifact(ReceivedArtifactInsert { + transfer_local_id: local_id, + protocol_transfer_id: transfer.protocol_id, + relative_path: &relative_path, + locator_kind: published.locator_kind, + locator: &published.locator, + logical_size: exported, + }) + .await + .map_err(VnidropError::repository)?; + } Ok(()) } } diff --git a/crates/vnidrop/src/runtime/share.rs b/crates/vnidrop/src/runtime/share.rs index 91181cb..873bf17 100644 --- a/crates/vnidrop/src/runtime/share.rs +++ b/crates/vnidrop/src/runtime/share.rs @@ -61,6 +61,17 @@ impl CoreInner { self.limits .validate_metadata_text("sender name", metadata.sender_name.as_deref()) .map_err(VnidropError::invalid_input)?; + if self + .targeted_store() + .contains_protocol_id(metadata.transfer_id) + .await + .map_err(anyhow::Error::new)? + { + return Err(VnidropError::invalid_input(anyhow::anyhow!( + "transfer id is already reserved by a targeted transfer" + )) + .into()); + } self.repository .insert_transfer(TransferUpsert { transfer_id, @@ -68,6 +79,7 @@ impl CoreInner { direction: TransferDirection::Send, status: TransferStatus::Importing, transfer_name: metadata.transfer_name.as_deref(), + sender_name: metadata.sender_name.as_deref(), content_hash: None, ticket: None, file_count: 0, @@ -149,6 +161,7 @@ impl CoreInner { import.file_count, import.total_size, ); + let sender_name = ticket_metadata.sender_name.clone(); let ticket = VnidropTicket::new_with_relay_urls( blob_ticket, ticket_metadata, @@ -182,6 +195,7 @@ impl CoreInner { direction: TransferDirection::Send, status: TransferStatus::Sharing, transfer_name: Some(&transfer_name), + sender_name: sender_name.as_deref(), content_hash: Some(&content_hash), ticket: Some(&ticket), file_count: import.file_count, diff --git a/crates/vnidrop/src/runtime/targeted.rs b/crates/vnidrop/src/runtime/targeted.rs index 7426572..dff8154 100644 --- a/crates/vnidrop/src/runtime/targeted.rs +++ b/crates/vnidrop/src/runtime/targeted.rs @@ -6,12 +6,11 @@ use anyhow::{Context, Result}; use iroh_blobs::{ticket::BlobTicket, BlobFormat}; use uuid::Uuid; -use super::{receive::ReceiveTarget, CoreInner}; +use super::{receive::ReceiveTarget, targeted_tag_name, CoreInner}; use crate::{ api::{ - experimental_saved_device_capabilities, PendingTargetedOffer, ShareMetadataInput, - ShareSource, TargetedOfferResponse, TargetedTransfer, TargetedTransferState, - TransferAccessMode, TransferMetadata, + experimental_saved_device_capabilities, PendingTargetedOffer, ShareSource, + TargetedOfferResponse, TargetedTransfer, TargetedTransferState, TransferAccessMode, }, error::VnidropError, secure_secret::{SecretHandle, SecretKind}, @@ -24,7 +23,6 @@ use crate::{ reconstruct_authorization, TargetedAuthorization, TargetedAuthorizationDraft, TargetedTransferRole, TargetedTransferRow, }, - ticket::VnidropTicket, util::{non_empty, now_ms}, }; @@ -59,10 +57,55 @@ impl CoreInner { } pub(crate) async fn restore_targeted_transfer_access(&self) -> Result<(), VnidropError> { + let mut active_tags = std::collections::HashSet::new(); for row in self.targeted_store().list_resumable_sender_rows().await? { + let root_hash = row + .content_hash + .parse::() + .map_err(|error| VnidropError::transfer(anyhow::anyhow!(error)))?; + let collection = + iroh_blobs::format::collection::Collection::load(root_hash, self.store.as_ref()) + .await + .map_err(VnidropError::transfer)?; + let tag_name = targeted_tag_name(&row.id); + self.store + .tags() + .set(&tag_name, (root_hash, iroh_blobs::BlobFormat::HashSeq)) + .await + .map_err(VnidropError::transfer)?; + self.register_share_hashes( + row.protocol_transfer_id, + std::iter::once(root_hash).chain(collection.iter().map(|(_, hash)| *hash)), + ) + .await; + self.access_policy + .set_mode( + row.protocol_transfer_id, + TransferAccessMode::ApprovalRequired, + ) + .await; self.access_policy .approve_endpoint_until(row.protocol_transfer_id, row.receiver_endpoint_id, None) .await; + active_tags.insert(tag_name); + } + use futures_lite::StreamExt as _; + let mut tags = self + .store + .tags() + .list_prefix("vnidrop/targeted/") + .await + .map_err(VnidropError::transfer)?; + while let Some(tag) = tags.next().await { + let tag = tag.map_err(VnidropError::transfer)?; + let name = String::from_utf8_lossy(tag.name.as_ref()).to_string(); + if !active_tags.contains(&name) { + self.store + .tags() + .delete(name) + .await + .map_err(VnidropError::transfer)?; + } } Ok(()) } @@ -84,12 +127,21 @@ impl CoreInner { .targeted_store() .protocol_ids_for_peer(peer_endpoint_id) .await?; + let sender_payloads = self + .targeted_store() + .sender_payloads_for_peer(peer_endpoint_id) + .await?; // Signal active transfers synchronously before awaiting share teardown. for protocol_transfer_id in &protocol_ids { let _ = self.take_active_transfer(*protocol_transfer_id); } for protocol_transfer_id in &protocol_ids { - let _ = self.cancel_idle_or_share(*protocol_transfer_id).await; + self.teardown_targeted_payload(*protocol_transfer_id, None) + .await; + } + for (id, protocol_transfer_id) in sender_payloads { + self.teardown_targeted_payload(protocol_transfer_id, Some(&id)) + .await; } self.targeted_store().cancel_by_peer(peer_endpoint_id).await } @@ -111,7 +163,8 @@ impl CoreInner { self.access_policy .remove_transfer(row.protocol_transfer_id) .await; - let _ = self.cancel_idle_or_share(row.protocol_transfer_id).await; + self.teardown_targeted_payload(row.protocol_transfer_id, Some(&row.id)) + .await; if !matches!( row.state, TargetedTransferState::Completed @@ -164,7 +217,8 @@ impl CoreInner { self.access_policy .remove_transfer(row.protocol_transfer_id) .await; - let _ = self.cancel_idle_or_share(row.protocol_transfer_id).await; + self.teardown_targeted_payload(row.protocol_transfer_id, Some(&row.id)) + .await; if let Some(handle) = &row.authorization_secret_handle { if let Some(custody) = &self.secret_custody { let _ = custody @@ -251,23 +305,69 @@ impl CoreInner { .require_saved(&receiver_endpoint_id) .await?; - let transfer_uuid = Uuid::new_v4().to_string(); - let protocol_transfer_id = allocate_protocol_transfer_id(&transfer_uuid); + let (transfer_uuid, protocol_transfer_id) = loop { + let transfer_uuid = Uuid::new_v4().to_string(); + let protocol_transfer_id = allocate_protocol_transfer_id(&transfer_uuid); + let invitation_collision = self + .repository + .list_transfers() + .await + .map_err(VnidropError::repository)? + .into_iter() + .any(|transfer| transfer.transfer_id == protocol_transfer_id); + if !invitation_collision + && !self + .targeted_store() + .contains_protocol_id(protocol_transfer_id) + .await? + { + break (transfer_uuid, protocol_transfer_id); + } + }; let sender_endpoint_id = self.endpoint.id().to_string(); let now = now_ms(); - let share = self - .share_files( - sources, - ShareMetadataInput { - transfer_id: protocol_transfer_id, - transfer_name: transfer_name.clone(), - sender_name: None, - access_mode: TransferAccessMode::ApprovalRequired, - }, + if sources.is_empty() { + return Err(VnidropError::invalid_input(anyhow::anyhow!( + "at least one source is required" + ))); + } + if sources.len() as u64 > self.limits.max_sources { + return Err(VnidropError::invalid_input(anyhow::anyhow!( + "source count {} exceeds limit {}", + sources.len(), + self.limits.max_sources + ))); + } + self.limits + .validate_metadata_text("transfer name", transfer_name.as_deref()) + .map_err(VnidropError::invalid_input)?; + let import = self + .import_sources(protocol_transfer_id, sources) + .await + .map_err(VnidropError::transfer)?; + let payload_name = transfer_name + .and_then(non_empty) + .unwrap_or_else(|| import.default_name.clone()); + let blob_ticket = + BlobTicket::new(self.endpoint.addr(), import.root_hash, BlobFormat::HashSeq); + self.store + .tags() + .set( + targeted_tag_name(&transfer_uuid), + (import.root_hash, BlobFormat::HashSeq), ) .await .map_err(VnidropError::transfer)?; + self.register_share_hashes( + protocol_transfer_id, + std::iter::once(import.root_hash).chain(import.member_hashes.iter().copied()), + ) + .await; + self.access_policy + .set_mode(protocol_transfer_id, TransferAccessMode::ApprovalRequired) + .await; + drop(import.tag); let store = self.targeted_store(); let row = TargetedTransferRow { @@ -275,11 +375,11 @@ impl CoreInner { protocol_transfer_id, sender_endpoint_id: sender_endpoint_id.clone(), receiver_endpoint_id: receiver_endpoint_id.clone(), - manifest_id: share.hash.clone(), - content_hash: share.hash.clone(), - transfer_name: share.transfer_name.clone(), - file_count: share.file_count, - total_size: share.total_size, + manifest_id: blob_ticket.hash().to_string(), + content_hash: blob_ticket.hash().to_string(), + transfer_name: payload_name.clone(), + file_count: import.file_count, + total_size: import.total_size, verified_bytes: 0, blob_ticket: None, authorization_secret_handle: None, @@ -288,19 +388,43 @@ impl CoreInner { created_at: now, updated_at: now, }; - store.insert(&row).await?; - store + if let Err(error) = store.insert(&row).await { + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + if let Err(error) = store .set_state( &transfer_uuid, TargetedTransferState::Preparing, TargetedTransferState::Offering, ) - .await?; + .await + { + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } - let addr = self + let addr = match self .device_relationships .peer_addr(&receiver_endpoint_id) - .await?; + .await + { + Ok(addr) => addr, + Err(error) => { + let _ = store + .set_state( + &transfer_uuid, + TargetedTransferState::Offering, + TargetedTransferState::Failed, + ) + .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + }; let client = TargetedTransferProtocol::client(self.endpoint.clone(), addr); let challenge = match tokio::time::timeout(self.connection_timeout(), client.request_challenge()).await @@ -314,7 +438,8 @@ impl CoreInner { TargetedTransferState::Failed, ) .await; - let _ = self.cancel_idle_or_share(protocol_transfer_id).await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(map_connect_failure(error)); } Err(_) => { @@ -325,27 +450,48 @@ impl CoreInner { TargetedTransferState::Failed, ) .await; - let _ = self.cancel_idle_or_share(protocol_transfer_id).await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(VnidropError::device_unavailable(anyhow::anyhow!( "device did not answer in time" ))); } }; - let (proof, generation, relationship_protocol_version) = self + let (proof, generation, relationship_protocol_version) = match self .device_relationships .prove_saved_possession(&receiver_endpoint_id, &challenge) - .await?; + .await + { + Ok(proof) => proof, + Err(error) => { + let _ = store + .set_state( + &transfer_uuid, + TargetedTransferState::Offering, + TargetedTransferState::Failed, + ) + .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + }; let protocol_version = experimental_saved_device_capabilities().targeted_transfer_protocol_version; - store + if let Err(error) = store .set_state( &transfer_uuid, TargetedTransferState::Offering, TargetedTransferState::AwaitingApproval, ) - .await?; + .await + { + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } let response = match tokio::time::timeout( self.connection_timeout() + self.offer_wait_timeout(), @@ -357,11 +503,11 @@ impl CoreInner { transfer_id: transfer_uuid.clone(), sender_endpoint_id: sender_endpoint_id.clone(), receiver_endpoint_id: receiver_endpoint_id.clone(), - manifest_id: share.hash.clone(), - content_hash: share.hash.clone(), - transfer_name: share.transfer_name.clone(), - file_count: share.file_count, - total_size: share.total_size, + manifest_id: blob_ticket.hash().to_string(), + content_hash: blob_ticket.hash().to_string(), + transfer_name: payload_name.clone(), + file_count: import.file_count, + total_size: import.total_size, relay_mode: self.relay_mode, relay_urls: self .custom_relay_urls @@ -381,7 +527,8 @@ impl CoreInner { TargetedTransferState::Failed, ) .await; - let _ = self.cancel_idle_or_share(protocol_transfer_id).await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(map_connect_failure(error)); } Err(_) => { @@ -392,7 +539,8 @@ impl CoreInner { TargetedTransferState::Failed, ) .await; - let _ = self.cancel_idle_or_share(protocol_transfer_id).await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(VnidropError::offer_timeout(anyhow::anyhow!( "offer timed out" ))); @@ -409,7 +557,8 @@ impl CoreInner { TargetedTransferState::Declined, ) .await; - let _ = self.cancel_idle_or_share(protocol_transfer_id).await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(VnidropError::permission(anyhow::anyhow!( "targeted offer declined: {reason}" ))); @@ -422,7 +571,8 @@ impl CoreInner { TargetedTransferState::Failed, ) .await; - let _ = self.cancel_idle_or_share(protocol_transfer_id).await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(map_offer_refuse_reason(&reason)); } } @@ -432,38 +582,71 @@ impl CoreInner { .approve_endpoint_until(protocol_transfer_id, receiver_endpoint_id.clone(), None) .await; - let parsed = crate::ticket::parse_transfer_ticket_with_limits(&share.ticket, &self.limits) - .map_err(VnidropError::ticket)?; - let blob_ticket = BlobTicket::new( - parsed.blob_ticket.addr().clone(), - parsed.blob_ticket.hash(), - BlobFormat::HashSeq, - ); - let authorization = TargetedAuthorization::issue(TargetedAuthorizationDraft { + let authorization = match TargetedAuthorization::issue(TargetedAuthorizationDraft { transfer_id: transfer_uuid.clone(), protocol_transfer_id, sender_endpoint_id, receiver_endpoint_id, - manifest_id: share.hash.clone(), - content_hash: share.hash.clone(), - file_count: share.file_count, - total_size: share.total_size, + manifest_id: blob_ticket.hash().to_string(), + content_hash: blob_ticket.hash().to_string(), + file_count: import.file_count, + total_size: import.total_size, protocol_version, - transfer_name: share.transfer_name.clone(), + transfer_name: payload_name, blob_ticket: blob_ticket.to_string(), - })?; - self.persist_authorization_secret(&transfer_uuid, &authorization) - .await?; - let encoded = authorization.encode()?; + }) { + Ok(authorization) => authorization, + Err(error) => { + let _ = store + .set_state_from_any(&transfer_uuid, TargetedTransferState::Failed) + .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + }; + if let Err(error) = self + .persist_authorization_secret(&transfer_uuid, &authorization) + .await + { + let _ = store + .set_state_from_any(&transfer_uuid, TargetedTransferState::Failed) + .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + let encoded = match authorization.encode() { + Ok(encoded) => encoded, + Err(error) => { + let _ = store + .set_state_from_any(&transfer_uuid, TargetedTransferState::Failed) + .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + }; - let deliver = client + let deliver = match client .deliver_authorization(DeliverTargetedAuthorization { transfer_id: transfer_uuid.clone(), authorization: encoded, }) .await .context("failed to deliver targeted authorization") - .map_err(VnidropError::network)?; + .map_err(VnidropError::network) + { + Ok(deliver) => deliver, + Err(error) => { + let _ = store + .set_state_from_any(&transfer_uuid, TargetedTransferState::Failed) + .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; + return Err(error); + } + }; if deliver != crate::targeted_transfer::protocol::DeliverAuthorizationResponse::Stored { let _ = store .set_state( @@ -472,6 +655,8 @@ impl CoreInner { TargetedTransferState::Failed, ) .await; + self.teardown_targeted_payload(protocol_transfer_id, Some(&transfer_uuid)) + .await; return Err(VnidropError::network(anyhow::anyhow!( "receiver rejected authorization delivery" ))); @@ -491,6 +676,18 @@ impl CoreInner { .ok_or_else(|| VnidropError::internal(anyhow::anyhow!("targeted transfer missing"))) } + async fn teardown_targeted_payload(&self, protocol_transfer_id: u64, id: Option<&str>) { + self.unregister_transfer_hashes(protocol_transfer_id).await; + self.access_policy + .remove_transfer(protocol_transfer_id) + .await; + if let Some(id) = id { + if let Err(error) = self.store.tags().delete(targeted_tag_name(id)).await { + tracing::warn!(%error, transfer_id = id, "failed to release targeted payload tag"); + } + } + } + pub(super) async fn receive_targeted_transfer( self: &Arc, transfer_id: String, @@ -652,20 +849,9 @@ impl CoreInner { let blob_ticket = BlobTicket::from_str_compat(&auth.blob_ticket) .map_err(|error| VnidropError::ticket(anyhow::anyhow!(error)))?; - let metadata = TransferMetadata::new( - auth.protocol_transfer_id, - non_empty(auth.transfer_name.clone()).unwrap_or_else(|| "transfer".to_string()), - None, - blob_ticket.hash(), - auth.file_count, - auth.total_size, - ); - let ticket = - VnidropTicket::new_with_relay_urls(blob_ticket, metadata, &self.custom_relay_urls) - .encode() - .map_err(VnidropError::ticket)?; - - let receive_result = self.receive_to_target(ticket, target, None).await; + let receive_result = self + .receive_targeted_payload(auth.protocol_transfer_id, blob_ticket, target) + .await; match receive_result { Ok(()) => { @@ -752,7 +938,7 @@ impl CoreInner { .await } - async fn load_stored_authorization( + pub(crate) async fn load_stored_authorization( &self, row: &TargetedTransferRow, ) -> Result, VnidropError> { diff --git a/crates/vnidrop/src/targeted_transfer/store.rs b/crates/vnidrop/src/targeted_transfer/store.rs index 976e66f..2eb74ee 100644 --- a/crates/vnidrop/src/targeted_transfer/store.rs +++ b/crates/vnidrop/src/targeted_transfer/store.rs @@ -77,6 +77,20 @@ impl TargetedTransferStore { Self { pool } } + pub(crate) async fn contains_protocol_id( + &self, + protocol_transfer_id: u64, + ) -> Result { + let row = sqlx::query( + "SELECT EXISTS(SELECT 1 FROM targeted_transfers WHERE protocol_transfer_id = ?1)", + ) + .bind(protocol_transfer_id as i64) + .fetch_one(&self.pool) + .await + .map_err(VnidropError::repository)?; + Ok(row.get::(0) != 0) + } + pub(crate) async fn insert(&self, transfer: &TargetedTransferRow) -> Result<(), VnidropError> { sqlx::query( r#" @@ -349,6 +363,32 @@ impl TargetedTransferStore { .collect()) } + pub(crate) async fn sender_payloads_for_peer( + &self, + peer_endpoint_id: &str, + ) -> Result, VnidropError> { + let rows = sqlx::query( + r#" + SELECT id, protocol_transfer_id FROM targeted_transfers + WHERE role = 'sender' + AND (sender_endpoint_id = ?1 OR receiver_endpoint_id = ?1) + "#, + ) + .bind(peer_endpoint_id) + .fetch_all(&self.pool) + .await + .map_err(VnidropError::repository)?; + Ok(rows + .into_iter() + .map(|row| { + ( + row.get("id"), + row.get::("protocol_transfer_id") as u64, + ) + }) + .collect()) + } + pub(crate) async fn mark_interrupted_in_flight(&self) -> Result { let now = now_ms(); let result = sqlx::query( diff --git a/crates/vnidrop/src/tests/device_relationship.rs b/crates/vnidrop/src/tests/device_relationship.rs index 6c91c6f..3847530 100644 --- a/crates/vnidrop/src/tests/device_relationship.rs +++ b/crates/vnidrop/src/tests/device_relationship.rs @@ -22,6 +22,7 @@ impl CoreEventSink for RecordingSink { struct ProtectedNode { _data_dir: tempfile::TempDir, core: Arc, + store: Arc, } impl ProtectedNode { @@ -34,12 +35,13 @@ impl ProtectedNode { let core = VnidropCore::initialize_with_test_secret_store( data_dir.path().to_string_lossy().into_owned(), sink, - store, + store.clone(), ) .expect("protected test core"); Self { _data_dir: data_dir, core, + store, } } } @@ -51,6 +53,15 @@ impl Drop for ProtectedNode { } fn share_path(core: &VnidropCore, source: &Path, transfer_id: u64) -> crate::ShareResult { + share_path_named(core, source, transfer_id, "sender") +} + +fn share_path_named( + core: &VnidropCore, + source: &Path, + transfer_id: u64, + sender_name: &str, +) -> crate::ShareResult { core.share_files( vec![ShareSource { kind: SourceKind::Path, @@ -61,7 +72,7 @@ fn share_path(core: &VnidropCore, source: &Path, transfer_id: u64) -> crate::Sha ShareMetadataInput { transfer_id, transfer_name: Some("hello.txt".to_string()), - sender_name: Some("sender".to_string()), + sender_name: Some(sender_name.to_string()), access_mode: TransferAccessMode::ApprovalRequired, }, ) @@ -88,17 +99,27 @@ fn wait_for_receiver_request(sender: &VnidropCore, transfer_id: u64) -> crate::R } fn complete_transfer(sender: &ProtectedNode, receiver: &ProtectedNode, transfer_id: u64) { + complete_transfer_named(sender, receiver, transfer_id, "sender", "receiver") +} + +fn complete_transfer_named( + sender: &ProtectedNode, + receiver: &ProtectedNode, + transfer_id: u64, + sender_name: &str, + receiver_name: &str, +) { let source_dir = tempfile::tempdir().unwrap(); let output_dir = tempfile::tempdir().unwrap(); let source_path = source_dir.path().join("hello.txt"); std::fs::write(&source_path, b"mutual consent").unwrap(); - let share = share_path(&sender.core, &source_path, transfer_id); + let share = share_path_named(&sender.core, &source_path, transfer_id, sender_name); let output_dir = output_dir.path().to_string_lossy().to_string(); let receiver_core = receiver.core.clone(); let ticket = share.ticket.clone(); - let handle = std::thread::spawn(move || { - receiver_core.receive(ticket, output_dir, Some("receiver".to_string())) - }); + let receiver_name = receiver_name.to_string(); + let handle = + std::thread::spawn(move || receiver_core.receive(ticket, output_dir, Some(receiver_name))); let request = wait_for_receiver_request(&sender.core, share.transfer_id); sender .core @@ -106,6 +127,15 @@ fn complete_transfer(sender: &ProtectedNode, receiver: &ProtectedNode, transfer_ .unwrap(); handle.join().unwrap().unwrap(); + if sender + .core + .list_saved_devices() + .unwrap() + .iter() + .any(|device| device.endpoint_id == receiver.core.status().endpoint_id) + { + return; + } let started = Instant::now(); let peer = receiver.core.status().endpoint_id.clone(); loop { @@ -194,8 +224,13 @@ fn mutual_consent_reaches_saved_after_both_grants_and_acknowledgement() { let bob_saved = bob.core.list_saved_devices().unwrap(); assert_eq!(alice_saved.len(), 1); assert_eq!(alice_saved[0].endpoint_id, bob_id); + assert_eq!( + alice_saved[0].remote_display_name.as_deref(), + Some("receiver") + ); assert_eq!(bob_saved.len(), 1); assert_eq!(bob_saved[0].endpoint_id, alice_id); + assert_eq!(bob_saved[0].remote_display_name.as_deref(), Some("sender")); } #[test] @@ -571,6 +606,82 @@ fn saved_device_local_label_survives_listing_and_rejects_non_saved_peers() { assert!(matches!(err, crate::VnidropError::InvalidInput { .. })); } +#[test] +fn later_authenticated_invitation_refreshes_saved_name_without_new_eligibility() { + let alice = ProtectedNode::new(); + let bob = ProtectedNode::new(); + let alice_id = alice.core.status().endpoint_id.clone(); + let bob_id = bob.core.status().endpoint_id.clone(); + reach_saved(&alice, &bob, 90_060); + + let authenticated_before_label = + alice.core.list_saved_devices().unwrap()[0].last_authenticated_at; + alice + .core + .set_saved_device_label(bob_id.clone(), Some("My tablet".to_string())) + .unwrap(); + let before = alice.core.list_saved_devices().unwrap().remove(0); + assert_eq!(before.last_authenticated_at, authenticated_before_label); + std::thread::sleep(Duration::from_millis(2)); + complete_transfer_named(&alice, &bob, 90_061, "Alice refreshed", "Bob refreshed"); + + let started = Instant::now(); + loop { + let refreshed = alice.core.list_saved_devices().unwrap()[0] + .remote_display_name + .as_deref() + == Some("Bob refreshed"); + if refreshed { + break; + } + assert!( + started.elapsed() < Duration::from_secs(10), + "saved name never refreshed" + ); + std::thread::sleep(Duration::from_millis(25)); + } + let alice_saved = alice.core.list_saved_devices().unwrap().remove(0); + let bob_saved = bob.core.list_saved_devices().unwrap().remove(0); + assert_eq!(alice_saved.local_label.as_deref(), Some("My tablet")); + assert_eq!( + alice_saved.remote_display_name.as_deref(), + Some("Bob refreshed") + ); + assert_eq!( + bob_saved.remote_display_name.as_deref(), + Some("Alice refreshed") + ); + assert!(alice_saved.last_authenticated_at > before.last_authenticated_at); + assert!(alice.core.list_pairing_eligibilities().unwrap().is_empty()); + assert!(bob.core.list_pairing_eligibilities().unwrap().is_empty()); + assert_eq!(alice_id, bob_saved.endpoint_id); +} + +#[test] +fn saved_remote_name_and_local_label_survive_restart() { + let alice = ProtectedNode::new(); + let bob = ProtectedNode::new(); + let bob_id = bob.core.status().endpoint_id.clone(); + reach_saved(&alice, &bob, 90_062); + alice + .core + .set_saved_device_label(bob_id.clone(), Some("My tablet".to_string())) + .unwrap(); + let expected = alice.core.list_saved_devices().unwrap(); + + alice.core.shutdown(); + let restarted = VnidropCore::initialize_with_test_secret_store( + alice._data_dir.path().to_string_lossy().into_owned(), + Arc::new(RecordingSink { + events: Mutex::new(Vec::new()), + }), + alice.store.clone(), + ) + .unwrap(); + assert_eq!(restarted.list_saved_devices().unwrap(), expected); + restarted.shutdown(); +} + #[test] fn events_carry_stable_ids_and_monotonic_revisions() { let alice = ProtectedNode::new(); diff --git a/crates/vnidrop/src/tests/pairing_eligibility.rs b/crates/vnidrop/src/tests/pairing_eligibility.rs index 02b7aa6..799ab04 100644 --- a/crates/vnidrop/src/tests/pairing_eligibility.rs +++ b/crates/vnidrop/src/tests/pairing_eligibility.rs @@ -181,6 +181,14 @@ fn completed_authenticated_transfer_creates_pairing_eligibility_on_both_sides() .unwrap(); assert_eq!(sender_entry.session_id, receiver_entry.session_id); + assert_eq!( + sender_entry.remote_display_name.as_deref(), + Some("receiver") + ); + assert_eq!( + receiver_entry.remote_display_name.as_deref(), + Some("sender") + ); assert_eq!(sender_entry.protocol_version, protocol); assert_eq!(receiver_entry.protocol_version, protocol); assert!(sender_entry.expires_at > sender_entry.created_at); @@ -204,6 +212,36 @@ fn completed_authenticated_transfer_creates_pairing_eligibility_on_both_sides() ); } +#[test] +fn receiver_eligibility_uses_sender_authenticated_name_not_edited_ticket_metadata() { + let source_dir = tempfile::tempdir().unwrap(); + let output_dir = tempfile::tempdir().unwrap(); + let source_path = source_dir.path().join("hello.txt"); + std::fs::write(&source_path, b"canonical sender name").unwrap(); + let sender = ProtectedNode::new(); + let receiver = ProtectedNode::new(); + let sender_id = sender.core.status().endpoint_id.clone(); + let share = share_path(&sender.core, &source_path, 70_002); + let edited_ticket = + crate::ticket::rewrite_sender_name_for_test(&share.ticket, "Attacker").unwrap(); + + receive_with_response( + &sender.core, + share.transfer_id, + receiver.core.clone(), + edited_ticket, + output_dir.path(), + true, + ) + .unwrap(); + wait_for_eligibility(&receiver.core, &sender_id); + let eligibility = receiver.core.list_pairing_eligibilities().unwrap(); + assert_eq!( + eligibility[0].remote_display_name.as_deref(), + Some("sender") + ); +} + #[test] fn declined_cancelled_and_failed_transfers_create_no_eligibility() { let source_dir = tempfile::tempdir().unwrap(); diff --git a/crates/vnidrop/src/tests/persistence.rs b/crates/vnidrop/src/tests/persistence.rs index cb63c95..767f862 100644 --- a/crates/vnidrop/src/tests/persistence.rs +++ b/crates/vnidrop/src/tests/persistence.rs @@ -50,3 +50,46 @@ async fn open_all_returns_all_domain_stores_and_schemas() { ); } } + +#[tokio::test] +async fn open_all_migrates_name_columns_without_losing_existing_local_labels() { + let temp = tempfile::tempdir().unwrap(); + let db = temp.path().join("vnidrop.sqlite3"); + let options = sqlx::sqlite::SqliteConnectOptions::new() + .filename(&db) + .create_if_missing(true); + let pool = sqlx::SqlitePool::connect_with(options).await.unwrap(); + sqlx::query( + r#" + CREATE TABLE device_relationships ( + remote_endpoint_id TEXT PRIMARY KEY, state TEXT NOT NULL, + generation INTEGER NOT NULL, minimum_protocol_version INTEGER NOT NULL, + session_id TEXT, issued_grant_handle TEXT, held_grant_handle TEXT, + issued_grant_id TEXT, held_grant_id TEXT, peer_ack INTEGER NOT NULL, + local_ack INTEGER NOT NULL, created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, local_label TEXT + ) + "#, + ) + .execute(&pool) + .await + .unwrap(); + sqlx::query( + r#" + INSERT INTO device_relationships ( + remote_endpoint_id, state, generation, minimum_protocol_version, + peer_ack, local_ack, created_at, updated_at, local_label + ) VALUES ('peer', 'saved', 1, 1, 1, 1, 10, 20, 'My tablet') + "#, + ) + .execute(&pool) + .await + .unwrap(); + drop(pool); + + let stores = persistence::open_all(temp.path()).await.unwrap(); + let saved = stores.relationships.list_saved_devices().await.unwrap(); + assert_eq!(saved[0].local_label.as_deref(), Some("My tablet")); + assert_eq!(saved[0].remote_display_name, None); + assert_eq!(saved[0].last_authenticated_at, None); +} diff --git a/crates/vnidrop/src/tests/platform_contract_android.rs b/crates/vnidrop/src/tests/platform_contract_android.rs index b3c6623..5778fe8 100644 --- a/crates/vnidrop/src/tests/platform_contract_android.rs +++ b/crates/vnidrop/src/tests/platform_contract_android.rs @@ -311,6 +311,15 @@ fn complete_invitation_transfer( .unwrap(); handle.join().unwrap().unwrap(); + if sender + .core() + .list_saved_devices() + .unwrap() + .iter() + .any(|device| device.endpoint_id == receiver.core().status().endpoint_id) + { + return; + } let started = Instant::now(); let peer = receiver.core().status().endpoint_id.clone(); loop { diff --git a/crates/vnidrop/src/tests/repository.rs b/crates/vnidrop/src/tests/repository.rs index 90bdde2..df019ea 100644 --- a/crates/vnidrop/src/tests/repository.rs +++ b/crates/vnidrop/src/tests/repository.rs @@ -18,6 +18,7 @@ fn transfer( direction, status, transfer_name: Some("demo"), + sender_name: None, content_hash: Some("hash"), ticket: Some("ticket"), file_count: 1, @@ -516,6 +517,7 @@ async fn share_completion_is_conditional_and_atomic() { direction: TransferDirection::Send, status: TransferStatus::Importing, transfer_name: Some("pending"), + sender_name: None, content_hash: None, ticket: None, file_count: 0, @@ -532,6 +534,7 @@ async fn share_completion_is_conditional_and_atomic() { direction: TransferDirection::Send, status: TransferStatus::Sharing, transfer_name: Some("complete"), + sender_name: None, content_hash: Some("final-hash"), ticket: Some("final-ticket"), file_count: 2, @@ -570,6 +573,7 @@ async fn injected_write_failure_preserves_previous_state() { direction: TransferDirection::Send, status: TransferStatus::Importing, transfer_name: Some("pending"), + sender_name: None, content_hash: None, ticket: None, file_count: 0, diff --git a/crates/vnidrop/src/tests/runtime.rs b/crates/vnidrop/src/tests/runtime.rs index b3e2285..8bd6d09 100644 --- a/crates/vnidrop/src/tests/runtime.rs +++ b/crates/vnidrop/src/tests/runtime.rs @@ -144,6 +144,7 @@ fn startup_recovers_interrupted_transfer_and_persists_event() { direction: TransferDirection::Receive, status: TransferStatus::Receiving, transfer_name: Some("interrupted"), + sender_name: None, content_hash: Some("hash"), ticket: None, file_count: 1, @@ -190,6 +191,7 @@ fn startup_processes_persisted_delivery_receipts() { direction: TransferDirection::Receive, status: TransferStatus::Receiving, transfer_name: Some("completed receive"), + sender_name: None, content_hash: Some("hash"), ticket: None, file_count: 1, @@ -249,6 +251,7 @@ fn startup_fails_persisted_share_when_root_blob_is_missing() { direction: TransferDirection::Send, status: TransferStatus::Sharing, transfer_name: Some("missing blob"), + sender_name: None, content_hash: Some(&missing_hash), ticket: Some("ticket"), file_count: 1, diff --git a/crates/vnidrop/src/tests/targeted_transfer.rs b/crates/vnidrop/src/tests/targeted_transfer.rs index 5065270..87e9e38 100644 --- a/crates/vnidrop/src/tests/targeted_transfer.rs +++ b/crates/vnidrop/src/tests/targeted_transfer.rs @@ -1,9 +1,13 @@ use std::{ path::Path, + str::FromStr, sync::{Arc, Mutex}, time::{Duration, Instant}, }; +use iroh::{endpoint::presets, Endpoint}; +use iroh_blobs::{get::request::get_hash_seq_and_sizes, ticket::BlobTicket}; + use crate::{ secure_secret::FaultInjectingSecretStore, CoreEvent, CoreEventSink, CoreNetworkConfig, CoreRelayMode, DeviceRelationshipState, PendingTargetedOffer, PublishedOutput, @@ -242,6 +246,13 @@ fn create_targeted_transfer_is_immutable_and_saved_only() { let bob_id = bob.core().status().endpoint_id.clone(); let stranger_id = stranger.core().status().endpoint_id.clone(); establish_saved(&alice, &bob, 10_001); + let invitation_history_before = alice + .core() + .list_transfers() + .unwrap() + .into_iter() + .map(|transfer| (transfer.local_id, transfer.ticket)) + .collect::>(); let source_dir = tempfile::tempdir().unwrap(); let source_path = source_dir.path().join("payload.txt"); @@ -279,6 +290,17 @@ fn create_targeted_transfer_is_immutable_and_saved_only() { crate::TargetedOfferResponse::Approved { transfer_id } if transfer_id == transfer.id )); + assert_eq!( + alice + .core() + .list_transfers() + .unwrap() + .into_iter() + .map(|transfer| (transfer.local_id, transfer.ticket)) + .collect::>(), + invitation_history_before, + "targeted payloads must not become invitation transfers" + ); assert_eq!( transfer.sender_endpoint_id, @@ -386,6 +408,10 @@ fn explicit_approval_gates_content_and_binds_authorization_to_receiver() { let charlie = ProtectedNode::new(); let bob_id = bob.core().status().endpoint_id.clone(); establish_saved(&alice, &bob, 10_020); + let alice_invitation_count = alice.core().list_transfers().unwrap().len(); + let bob_invitation_count = bob.core().list_transfers().unwrap().len(); + let alice_eligibility_count = alice.core().list_pairing_eligibilities().unwrap().len(); + let bob_eligibility_count = bob.core().list_pairing_eligibilities().unwrap().len(); let source_dir = tempfile::tempdir().unwrap(); let source_path = source_dir.path().join("payload.txt"); @@ -466,6 +492,22 @@ fn explicit_approval_gates_content_and_binds_authorization_to_receiver() { std::fs::read(output.path().join("payload.txt")).unwrap(), payload ); + assert_eq!( + alice.core().list_transfers().unwrap().len(), + alice_invitation_count + ); + assert_eq!( + bob.core().list_transfers().unwrap().len(), + bob_invitation_count + ); + assert_eq!( + alice.core().list_pairing_eligibilities().unwrap().len(), + alice_eligibility_count + ); + assert_eq!( + bob.core().list_pairing_eligibilities().unwrap().len(), + bob_eligibility_count + ); let charlie_output = tempfile::tempdir().unwrap(); let leaked = charlie.core().receive_targeted_transfer( @@ -478,6 +520,59 @@ fn explicit_approval_gates_content_and_binds_authorization_to_receiver() { ); } +#[test] +fn unrelated_endpoint_cannot_fetch_a_leaked_targeted_blob_ticket() { + let alice = ProtectedNode::new(); + let bob = ProtectedNode::new(); + establish_saved(&alice, &bob, 10_025); + let transfer = approve_one(&alice, &bob, b"private payload", "payload.txt"); + let (protocol_transfer_id, leaked) = alice + .core() + .targeted_blob_ticket_for_test(transfer.id) + .unwrap(); + let collision_source_dir = tempfile::tempdir().unwrap(); + let collision_source = collision_source_dir.path().join("public.txt"); + std::fs::write(&collision_source, b"unrelated public payload").unwrap(); + let collision = alice.core().share_files( + vec![targeted_source(&collision_source)], + ShareMetadataInput { + transfer_id: protocol_transfer_id, + transfer_name: Some("public.txt".to_string()), + sender_name: Some("sender".to_string()), + access_mode: TransferAccessMode::Public, + }, + ); + assert!( + collision.is_err(), + "an invitation transfer must not reuse a targeted ACL identity" + ); + let ticket = BlobTicket::from_str(&leaked).unwrap(); + + let runtime = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + .unwrap(); + runtime.block_on(async move { + let attacker = Endpoint::builder(presets::Minimal).bind().await.unwrap(); + let connection = attacker + .connect(ticket.addr().clone(), iroh_blobs::ALPN) + .await + .unwrap(); + let result = get_hash_seq_and_sizes( + &connection, + &ticket.hash_and_format().hash, + 1024 * 1024 * 32, + None, + ) + .await; + assert!( + result.is_err(), + "a leaked targeted blob ticket must not authorize another endpoint" + ); + attacker.close().await; + }); +} + #[test] fn invitation_multi_receiver_shares_remain_independently_authorized() { let source_dir = tempfile::tempdir().unwrap(); @@ -787,6 +882,34 @@ fn delete_removes_authorization_and_resumable_state() { assert_eq!(sender_deleted.state, TargetedTransferState::Deleted); } +#[test] +fn sender_delete_revokes_receiver_bound_payload_access() { + let alice = ProtectedNode::new(); + let bob = ProtectedNode::new(); + establish_saved(&alice, &bob, 11_055); + let transfer = approve_one(&alice, &bob, b"delete at sender", "payload.txt"); + + alice + .core() + .delete_targeted_transfer(transfer.id.clone()) + .unwrap(); + let deleted = alice + .core() + .get_targeted_transfer(transfer.id.clone()) + .unwrap() + .unwrap(); + assert_eq!(deleted.state, TargetedTransferState::Deleted); + + let output = tempfile::tempdir().unwrap(); + let receive = bob + .core() + .receive_targeted_transfer(transfer.id, output.path().to_string_lossy().into_owned()); + assert!( + receive.is_err(), + "sender deletion must revoke the receiver's provider access" + ); +} + #[test] fn concurrent_independent_transfers_between_same_devices_are_isolated() { let alice = ProtectedNode::new(); diff --git a/crates/vnidrop/src/ticket.rs b/crates/vnidrop/src/ticket.rs index 3d25ea8..b381762 100644 --- a/crates/vnidrop/src/ticket.rs +++ b/crates/vnidrop/src/ticket.rs @@ -57,6 +57,13 @@ impl VnidropTicket { } } +#[cfg(test)] +pub(crate) fn rewrite_sender_name_for_test(value: &str, sender_name: &str) -> Result { + let mut ticket = VnidropTicket::decode(value)?; + ticket.metadata.sender_name = Some(sender_name.to_string()); + ticket.encode() +} + #[derive(Debug, Clone)] pub(crate) struct ParsedTransferTicket { pub(crate) blob_ticket: BlobTicket,