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
24 changes: 24 additions & 0 deletions src-tauri/src/commands/remote_control.rs
Original file line number Diff line number Diff line change
Expand Up @@ -205,6 +205,30 @@ pub async fn remote_force_cutoff(
Ok(())
}

/// Cuts off every managed endpoint on sign-out (PRD "Hard cutoff on
/// revocation or expiry": the Supabase session ended). Tears down pooled
/// managed connections and their PTY sessions, emits `remote://cutoff` for
/// each, and refuses new managed connections until the user signs in again.
#[tauri::command]
pub async fn remote_cut_off_managed(
app: tauri::AppHandle,
state: State<'_, RemoteExecState>,
pty_state: State<'_, crate::commands::remote_pty_commands::RemotePtyState>,
) -> Result<(), String> {
let reason = CutoffReasonDto::SessionEnded;
for endpoint_id in state.0.cut_off_managed(reason.into()).await {
pty_state.0.close_all_for_endpoint(&endpoint_id).await;
let _ = app.emit(
REMOTE_CUTOFF_EVENT,
RemoteCutoffEvent {
endpoint_id,
reason,
},
);
}
Ok(())
}

/// Clears a previously forced cutoff after the user reauthenticates and a
/// fresh certificate is issued through the normal registration and issuance
/// flow (PRD "The user regains access only by reauthenticating and obtaining
Expand Down
108 changes: 107 additions & 1 deletion src-tauri/src/core/remote_ssh_transport.rs
Original file line number Diff line number Diff line change
Expand Up @@ -688,6 +688,8 @@ pub struct PoolKey {
host_key_fingerprints: Vec<String>,
/// Relay URL for relayed endpoints; `None` for direct TCP.
relay_url: Option<String>,
/// Whether this is a Treq-managed instance, which sign-out cuts off.
managed: bool,
}

impl PoolKey {
Expand Down Expand Up @@ -715,10 +717,18 @@ impl PoolKey {
SshTransport::Direct => None,
SshTransport::Relay { url } => Some(url.clone()),
},
managed: is_managed(endpoint),
}
}
}

fn is_managed(endpoint: &SshEndpoint) -> bool {
matches!(
endpoint.source,
crate::core::remote_control_plane::SshEndpointSource::Managed { .. }
)
}

struct PooledConnection {
handle: Arc<AsyncMutex<Handle<TreqSshClientHandler>>>,
last_used: Instant,
Expand Down Expand Up @@ -747,6 +757,9 @@ pub struct SshConnectionPool {
/// connection, so a refreshed token needs no reconnect, and it is never
/// logged or included in an error.
relay_access_token: std::sync::RwLock<Option<String>>,
/// Set by [`Self::cut_off_managed`] on sign-out: every managed endpoint is
/// refused, pooled or not, until a signed-in user's token arrives.
managed_cutoff: std::sync::RwLock<Option<CutoffReason>>,
}

impl Default for SshConnectionPool {
Expand All @@ -765,13 +778,20 @@ impl SshConnectionPool {
cutoffs: AsyncMutex::new(HashMap::new()),
device_key_provider: None,
relay_access_token: std::sync::RwLock::new(None),
managed_cutoff: std::sync::RwLock::new(None),
}
}

/// Sets (or clears, on sign-out) the token used to open relayed
/// connections. See the field doc on `relay_access_token`.
pub fn set_relay_access_token(&self, token: Option<String>) {
let token = token.filter(|token| !token.is_empty());
if token.is_some() {
*self
.managed_cutoff
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = None;
}
match self.relay_access_token.write() {
Ok(mut guard) => *guard = token,
Err(poisoned) => *poisoned.into_inner() = token,
Expand Down Expand Up @@ -828,6 +848,32 @@ impl SshConnectionPool {
tracing::warn!(endpoint_id = %endpoint_id, reason = %reason, "ssh endpoint credential cut off");
}

/// Cuts off every managed endpoint when the user signs out: tears down
/// each pooled managed connection (see [`Self::force_cutoff`]) and refuses
/// new managed connections until [`Self::set_relay_access_token`] receives
/// a token again. Returns the endpoint ids whose connections were torn
/// down. User-managed endpoints are untouched.
pub async fn cut_off_managed(&self, reason: CutoffReason) -> Vec<String> {
*self
.managed_cutoff
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some(reason);
let mut endpoint_ids: Vec<String> = {
let connections = self.connections.lock().await;
connections
.keys()
.filter(|key| key.managed)
.map(|key| key.endpoint_id.clone())
.collect()
};
endpoint_ids.sort();
endpoint_ids.dedup();
for endpoint_id in &endpoint_ids {
self.force_cutoff(endpoint_id, reason).await;
}
endpoint_ids
}

/// Clears a previously forced cutoff, e.g. after the user reauthenticates
/// and a fresh certificate is issued through the normal registration and
/// issuance flow. Does not itself reconnect; the next call reconnects with
Expand Down Expand Up @@ -857,7 +903,15 @@ impl SshConnectionPool {
&self,
endpoint: &SshEndpoint,
) -> Result<Arc<AsyncMutex<Handle<TreqSshClientHandler>>>, SshTransportError> {
if let Some(reason) = self.cutoff_reason(&endpoint.id).await {
let managed_cutoff = if is_managed(endpoint) {
*self
.managed_cutoff
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
} else {
None
};
if let Some(reason) = managed_cutoff.or(self.cutoff_reason(&endpoint.id).await) {
return Err(SshTransportError::CredentialCutOff {
endpoint_id: endpoint.id.clone(),
reason,
Expand Down Expand Up @@ -2651,6 +2705,58 @@ mod tests {
assert!(output.success());
}

#[tokio::test]
async fn cut_off_managed_tears_down_managed_endpoints_until_sign_in() {
let (addr, host_key) = start_mock_server(0).await;
let temp_dir = tempfile::tempdir().unwrap();
let key_reference = write_client_key(temp_dir.path(), &test_host_key());
let user_endpoint = test_endpoint(addr, &host_key, key_reference);
let managed = SshEndpoint {
id: "managed-endpoint".to_string(),
source: SshEndpointSource::Managed {
provider: "fly_sprites".to_string(),
generation: 1,
},
..user_endpoint.clone()
};
let pool = SshConnectionPool::new();
let cancellation = CancellationToken::new();
let run = |endpoint: &SshEndpoint| {
let endpoint = endpoint.clone();
let (pool, cancellation) = (&pool, &cancellation);
async move {
let args = ["repo".to_string(), "inspect".to_string()];
exec_command(pool, &endpoint, &args, ExecLimits::default(), cancellation).await
}
};
run(&managed).await.unwrap();
run(&user_endpoint).await.unwrap();

let cut = pool.cut_off_managed(CutoffReason::SessionEnded).await;

assert_eq!(cut, vec![managed.id.clone()]);
assert_eq!(pool.pooled_connection_count().await, 1);
assert!(matches!(
run(&managed).await,
Err(SshTransportError::CredentialCutOff {
reason: CutoffReason::SessionEnded,
..
})
));
assert!(run(&user_endpoint).await.is_ok());

// Signing in again lifts the blanket block; the torn-down endpoint
// still waits for its own reauthentication.
let fresh = SshEndpoint {
id: "other-managed".to_string(),
..managed.clone()
};
assert!(run(&fresh).await.is_err());
pool.set_relay_access_token(Some("token".to_string()));
assert!(run(&fresh).await.is_ok());
assert!(run(&managed).await.is_err());
}

// -- Retrying after network loss (PRD "Structured command protocol" >
// "Retrying after network loss") ---------------------------------------

Expand Down
1 change: 1 addition & 0 deletions src-tauri/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -875,6 +875,7 @@ pub fn run() {
commands::remote_transport_metrics,
commands::remote_force_cutoff,
commands::remote_clear_cutoff,
commands::remote_cut_off_managed,
commands::remote_set_relay_access_token,
commands::remote_cutoff_reason,
commands::set_window_repo_path,
Expand Down
8 changes: 8 additions & 0 deletions src/lib/api-extra.ts
Original file line number Diff line number Diff line change
Expand Up @@ -842,6 +842,14 @@ export const remoteForceCutoff = (
reason: string,
): Promise<void> => invoke("remote_force_cutoff", { endpointId, reason });

/**
* Cuts off every managed endpoint on sign-out: tears down pooled managed
* connections and PTYs, emits `remote://cutoff` for each, and refuses new
* managed connections until a signed-in token reaches the native pool.
*/
export const remoteCutOffManaged = (): Promise<void> =>
invoke("remote_cut_off_managed");

/**
* Clears a previously forced cutoff after the user reauthenticates and a
* fresh certificate is issued through the normal registration and issuance
Expand Down
8 changes: 8 additions & 0 deletions src/stores/authStore.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import type { Session, User } from "@supabase/supabase-js";
import { create } from "zustand";
import { getSetting, setSetting } from "../lib/api";
import { remoteCutOffManaged } from "../lib/api-extra";
import {
SUPABASE_ANON_KEY,
SUPABASE_URL,
Expand Down Expand Up @@ -55,6 +56,13 @@ export const useAuthStore = create<AuthState>((set, get) => ({
await openUrl(`${WEB_URL}/sign-in?source=desktop`);
},
signOut: async () => {
// Signing out ends the authorization behind managed SSH access, so
// open managed connections and PTYs are cut off, not left running.
try {
await remoteCutOffManaged();
} catch {
// No native pool (e.g. a web preview); nothing to cut off.
}
await supabase.auth.signOut();
set({ user: null, session: null, subscription: null });
await persistSession(null);
Expand Down
Loading