fix(zmodem): resolve transfer races and task lifecycle

This commit is contained in:
胡飞
2026-08-28 09:11:20 +08:00
parent f610aff434
commit f2d1301ac8
22 changed files with 550 additions and 113 deletions
+6
View File
@@ -166,6 +166,12 @@ jobs:
- name: Install system dependencies
if: ${{ matrix.platform != 'windows' }}
run: script/bootstrap
- name: Install lrzsz on Linux
if: ${{ matrix.platform == 'linux' }}
run: sudo apt-get update && sudo apt-get install -y lrzsz
- name: Install lrzsz on macOS
if: ${{ matrix.platform == 'macos' }}
run: brew list lrzsz >/dev/null 2>&1 || brew install lrzsz
- name: Install wasm-tools
run: cargo install wasm-tools --locked --version 1.251.0
- name: Test Linux
+73 -2
View File
@@ -254,11 +254,13 @@ impl BackgroundTaskFilter {
}
type CancelCallback = Arc<dyn Fn() + Send + Sync>;
type CancelResultCallback = Arc<dyn Fn() -> bool + Send + Sync>;
#[derive(Clone, Default)]
pub struct BackgroundTaskCancellation {
token: Option<CancellationToken>,
callback: Option<CancelCallback>,
result_callback: Option<CancelResultCallback>,
}
impl BackgroundTaskCancellation {
@@ -266,6 +268,7 @@ impl BackgroundTaskCancellation {
Self {
token: Some(token),
callback: None,
result_callback: None,
}
}
@@ -273,6 +276,7 @@ impl BackgroundTaskCancellation {
Self {
token: None,
callback: Some(Arc::new(callback)),
result_callback: None,
}
}
@@ -283,6 +287,15 @@ impl BackgroundTaskCancellation {
Self {
token: Some(token),
callback: Some(Arc::new(callback)),
result_callback: None,
}
}
pub fn callback_with_result(callback: impl Fn() -> bool + Send + Sync + 'static) -> Self {
Self {
token: None,
callback: None,
result_callback: Some(Arc::new(callback)),
}
}
@@ -295,11 +308,12 @@ impl BackgroundTaskCancellation {
cb();
true
});
token_cancelled || callback_called
let result_callback_called = self.result_callback.as_ref().is_some_and(|cb| cb());
token_cancelled || callback_called || result_callback_called
}
fn is_configured(&self) -> bool {
self.token.is_some() || self.callback.is_some()
self.token.is_some() || self.callback.is_some() || self.result_callback.is_some()
}
}
@@ -383,6 +397,14 @@ impl BackgroundTaskManager {
.map(|task| task.id)
}
pub fn find_latest_by_key(&self, key: &str) -> Option<BackgroundTaskId> {
self.tasks
.iter()
.rev()
.find(|task| task.key.as_deref() == Some(key))
.map(|task| task.id)
}
pub fn ensure_by_key(
&mut self,
spec: BackgroundTaskSpec,
@@ -605,6 +627,12 @@ impl BackgroundTaskManager {
// 更新的百分比(如下载 98/100 即成功时仍显示 98%)。
if let Some(total) = progress.total {
progress.current = total;
} else if progress.current > 0 {
// 总量未知的任务也要在成功时显示 100%,而不是停在 0。
progress.total = Some(progress.current);
} else {
progress.current = 1;
progress.total = Some(1);
}
progress.message = None;
}
@@ -1098,6 +1126,7 @@ mod tests {
callback: Some(Arc::new(move || {
called_for_cb.store(true, Ordering::SeqCst);
})),
result_callback: None,
},
cx,
)
@@ -1227,6 +1256,24 @@ mod tests {
assert_eq!(1, calls.load(Ordering::SeqCst));
}
#[gpui::test]
fn rejected_cancellation_callback_keeps_task_running(cx: &mut gpui::TestAppContext) {
let manager = new_manager(cx);
let id = manager.update(cx, |m, cx| {
let id = m.register(BackgroundTaskSpec::new("kind", "task"), cx);
m.set_cancellation(
id,
BackgroundTaskCancellation::callback_with_result(|| false),
cx,
);
m.mark_running(id, cx);
id
});
manager.update(cx, |m, cx| assert!(!m.request_cancel(id, cx)));
assert_eq!(BackgroundTaskStatus::Running, task(&manager, id, cx).status);
}
#[gpui::test]
fn cancellation_wins_over_late_success_or_failure(cx: &mut gpui::TestAppContext) {
let manager = new_manager(cx);
@@ -1304,6 +1351,30 @@ mod tests {
assert_eq!(Some(progress.current), progress.total);
}
#[gpui::test]
fn succeeded_task_with_unknown_total_pins_to_one_hundred(cx: &mut gpui::TestAppContext) {
let manager = new_manager(cx);
let id = manager.update(cx, |m, cx| {
let id = m.register(
BackgroundTaskSpec::new("kind", "download")
.progress_unit(BackgroundTaskProgressUnit::Bytes),
cx,
);
m.mark_running(id, cx);
// 下载开始时总大小未知(total = None),成功时也应显示 100%。
m.update_progress(id, 5_897, None, None, None, cx);
m.succeed(id, None, cx);
id
});
let task = task(&manager, id, cx);
assert_eq!(BackgroundTaskStatus::Succeeded, task.status);
let progress = task.progress.expect("progress should be kept");
assert_eq!(100, progress.percent());
assert_eq!(Some(progress.current), progress.total);
assert_eq!(5_897, progress.current);
}
#[gpui::test]
fn terminal_task_rejects_progress_and_repeated_finish(cx: &mut gpui::TestAppContext) {
let manager = new_manager(cx);
+2
View File
@@ -84,6 +84,8 @@ pub enum Event<'a> {
FileStarted(FileInfo<'a>),
/// The current file completed.
FileCompleted,
/// The receiver declined the current file with `ZSKIP`.
FileSkipped,
/// The session completed successfully.
SessionCompleted,
/// The session was aborted.
+15 -4
View File
@@ -48,6 +48,7 @@ pub struct Receiver {
manual_accept: bool,
zrpos_retries: u8,
file_active: bool,
session_complete_pending: bool,
}
/// Consecutive corrupt data subpackets, without any forward progress in
@@ -121,6 +122,7 @@ impl Receiver {
manual_accept: false,
zrpos_retries: 0,
file_active: false,
session_complete_pending: false,
};
receiver.queue_zrinit()?;
Ok(receiver)
@@ -262,6 +264,10 @@ impl Receiver {
if self.outgoing_offset >= self.outgoing.len() {
self.outgoing.clear();
self.outgoing_offset = 0;
if self.state == ReceiverPhase::SessionFinishWriting {
self.state = ReceiverPhase::SessionEnd;
self.session_complete_pending = true;
}
}
}
@@ -340,7 +346,6 @@ impl Receiver {
Some(Position::new(self.file_size)),
)),
ReceiverEvent::FileComplete => Event::FileCompleted,
ReceiverEvent::SessionComplete => Event::SessionCompleted,
ReceiverEvent::Aborted => Event::Aborted,
});
}
@@ -349,6 +354,11 @@ impl Receiver {
return Action::WriteWire(self.drain_outgoing());
}
if self.session_complete_pending {
self.session_complete_pending = false;
return Action::Event(Event::SessionCompleted);
}
if !self.drain_file().is_empty() {
return Action::WriteFile(self.drain_file());
}
@@ -514,12 +524,13 @@ impl Receiver {
Frame::ZFIN
if matches!(
self.state,
ReceiverPhase::FileWaitingSubpacket | ReceiverPhase::FileBegin
ReceiverPhase::SessionBegin
| ReceiverPhase::FileWaitingSubpacket
| ReceiverPhase::FileBegin
) =>
{
self.queue_zfin()?;
self.state = ReceiverPhase::SessionEnd;
self.push_event(ReceiverEvent::SessionComplete)?;
self.state = ReceiverPhase::SessionFinishWriting;
}
_ => {}
}
+29 -14
View File
@@ -274,6 +274,10 @@ impl Sender {
if self.outgoing_offset >= self.outgoing.len() {
self.outgoing.clear();
self.outgoing_offset = 0;
if self.state == SenderPhase::FinishWriting {
self.state = SenderPhase::Done;
self.pending_event = Some(SenderEvent::SessionComplete);
}
}
}
@@ -291,6 +295,8 @@ impl Sender {
match self.state {
SenderPhase::WaitReceiverInit => self.queue_zrqinit()?,
SenderPhase::WaitFilePos => self.queue_zfile()?,
SenderPhase::WaitFileDone => self.queue_zeof(self.file_size)?,
SenderPhase::WaitFinish => self.queue_zfin()?,
_ => {}
}
Ok(())
@@ -312,6 +318,7 @@ impl Sender {
if let Some(event) = self.pending_event.take() {
return Action::Event(match event {
SenderEvent::FileComplete => Event::FileCompleted,
SenderEvent::FileSkipped => Event::FileSkipped,
SenderEvent::SessionComplete => Event::SessionCompleted,
SenderEvent::Aborted => Event::Aborted,
});
@@ -427,14 +434,21 @@ impl Sender {
match header.frame() {
Frame::ZRINIT => self.on_zrinit(header),
Frame::ZRPOS | Frame::ZACK => self.on_zrpos(header.count()),
Frame::ZSKIP => {
self.on_zskip();
Ok(())
}
Frame::ZSKIP => self.on_zskip(),
Frame::ZABORT | Frame::ZCAN => {
self.on_abort();
Ok(())
}
Frame::ZFERR => {
self.on_abort();
Ok(())
}
Frame::ZNAK => match self.state {
SenderPhase::WaitFilePos => self.queue_zfile(),
SenderPhase::WaitFileDone => self.queue_zeof(self.file_size),
SenderPhase::WaitFinish => self.queue_zfin(),
_ => Ok(()),
},
Frame::ZFIN => self.on_zfin(),
_ => {
if self.state == SenderPhase::WaitReceiverInit {
@@ -470,11 +484,7 @@ impl Sender {
self.state = SenderPhase::ReadyForFile;
}
}
SenderPhase::WaitFinish => {
self.queue_oo()?;
self.state = SenderPhase::Done;
self.pending_event = Some(SenderEvent::SessionComplete);
}
SenderPhase::WaitFinish => self.queue_zfin()?,
_ => {}
}
Ok(())
@@ -534,7 +544,7 @@ impl Sender {
Ok(())
}
fn on_zskip(&mut self) {
fn on_zskip(&mut self) -> Result<(), Error> {
if matches!(
self.state,
SenderPhase::WaitFilePos
@@ -545,9 +555,15 @@ impl Sender {
self.has_file = false;
self.pending_request = None;
self.frame_remaining = 0;
self.pending_event = Some(SenderEvent::FileComplete);
self.state = SenderPhase::ReadyForFile;
self.pending_event = Some(SenderEvent::FileSkipped);
if self.finish_requested {
self.queue_zfin()?;
self.state = SenderPhase::WaitFinish;
} else {
self.state = SenderPhase::ReadyForFile;
}
}
Ok(())
}
fn on_abort(&mut self) {
@@ -559,8 +575,7 @@ impl Sender {
fn on_zfin(&mut self) -> Result<(), Error> {
if self.state == SenderPhase::WaitFinish {
self.queue_oo()?;
self.state = SenderPhase::Done;
self.pending_event = Some(SenderEvent::SessionComplete);
self.state = SenderPhase::FinishWriting;
}
Ok(())
}
+3 -1
View File
@@ -16,6 +16,7 @@ pub(crate) struct FileRequest {
#[derive(Clone, Copy, Debug, PartialEq)]
pub(crate) enum SenderEvent {
FileComplete,
FileSkipped,
SessionComplete,
Aborted,
}
@@ -24,7 +25,6 @@ pub(crate) enum SenderEvent {
pub(crate) enum ReceiverEvent {
FileStart,
FileComplete,
SessionComplete,
Aborted,
}
@@ -49,6 +49,7 @@ pub(crate) enum SenderPhase {
WaitFileAck,
WaitFileDone,
WaitFinish,
FinishWriting,
Done,
}
@@ -61,5 +62,6 @@ pub(crate) enum ReceiverPhase {
FileAcceptPending,
FileReadingSubpacket,
FileWaitingSubpacket,
SessionFinishWriting,
SessionEnd,
}
+95 -1
View File
@@ -5,12 +5,15 @@
//! Protocol-level unit tests exercising crate-internal framing and the
//! public poll/submit API.
extern crate std;
use crate::buffer::Buffer;
use crate::header::{Encoding, EscapeMode, Frame, Header, Zrinit, write_slice_escaped};
use crate::receiver::MAX_ZRPOS_RETRIES;
use crate::wire::{BufferWriter, HeaderReader, SliceReader, SubpacketType};
use crate::{Action, Error, Event, FileInfo, Position, Receiver, Sender, ZDLE, ZPAD};
use rstest::rstest;
use std::{vec, vec::Vec};
fn write_header(header: Header) -> Vec<u8> {
let mut buf = Buffer::<64>::new();
@@ -532,7 +535,7 @@ fn test_sender_timeout_retries_zfile_while_waiting_for_zrpos() {
}
#[test]
fn test_sender_skips_file_on_zskip() {
fn test_sender_reports_skipped_file_on_zskip() {
let mut sender = Sender::new().unwrap();
drain_wire_sender(&mut sender);
sender
@@ -546,7 +549,98 @@ fn test_sender_skips_file_on_zskip() {
let zskip = write_header(Header::new(Encoding::ZHEX, Frame::ZSKIP, [0; 4]));
sender.submit_wire(&zskip).unwrap();
assert_eq!(sender.poll(), Action::Event(Event::FileSkipped));
}
#[test]
fn skipped_last_file_still_finishes_the_session() {
let mut sender = Sender::new().unwrap();
drain_wire_sender(&mut sender);
sender
.start_file(FileInfo::new(b"skip.bin", Some(Position::new(16))))
.unwrap();
sender.finish().unwrap();
let zrinit = write_header(Header::new(Encoding::ZHEX, Frame::ZRINIT, [0; 4]));
sender.submit_wire(&zrinit).unwrap();
drain_wire_sender(&mut sender);
let zskip = write_header(Header::new(Encoding::ZHEX, Frame::ZSKIP, [0; 4]));
sender.submit_wire(&zskip).unwrap();
assert_eq!(sender.poll(), Action::Event(Event::FileSkipped));
match sender.poll() {
Action::WriteWire(bytes) => assert_eq!(parse_first_header(bytes).frame(), Frame::ZFIN),
action => panic!("skipped final file should send ZFIN, got {action:?}"),
}
}
#[test]
fn receiver_accepts_empty_session_zfin() {
let mut receiver = Receiver::new().unwrap();
drain_wire_receiver(&mut receiver);
let zfin = write_header(Header::new(Encoding::ZHEX, Frame::ZFIN, [0; 4]));
receiver.submit_wire(&zfin).unwrap();
let Action::WriteWire(bytes) = receiver.poll() else {
panic!("receiver should acknowledge empty session ZFIN");
};
let len = bytes.len();
receiver.wire_written(len);
assert_eq!(receiver.poll(), Action::Event(Event::SessionCompleted));
}
#[test]
fn sender_waits_for_zfin_and_flushes_oo_before_session_complete() {
let mut sender = Sender::new().unwrap();
drain_wire_sender(&mut sender);
sender
.start_file(FileInfo::new(b"done.bin", Some(Position::new(1))))
.unwrap();
sender.finish().unwrap();
let zrinit = write_header(Header::new(Encoding::ZHEX, Frame::ZRINIT, [0; 4]));
sender.submit_wire(&zrinit).unwrap();
drain_wire_sender(&mut sender);
let zrpos0 = write_header(Header::new(Encoding::ZHEX, Frame::ZRPOS, [0; 4]));
sender.submit_wire(&zrpos0).unwrap();
let Action::ReadFile { .. } = sender.poll() else {
panic!("sender should request file data");
};
sender.submit_file(&[1]).unwrap();
drain_wire_sender(&mut sender);
let zrpos1 = write_header(Header::new(
Encoding::ZHEX,
Frame::ZRPOS,
1_u32.to_le_bytes(),
));
sender.submit_wire(&zrpos1).unwrap();
drain_wire_sender(&mut sender);
sender.submit_wire(&zrinit).unwrap();
assert_eq!(sender.poll(), Action::Event(Event::FileCompleted));
drain_wire_sender(&mut sender);
sender.submit_wire(&zrinit).unwrap();
match sender.poll() {
Action::WriteWire(bytes) => {
assert_eq!(parse_first_header(bytes).frame(), Frame::ZFIN);
let len = bytes.len();
sender.wire_written(len);
}
action => panic!("duplicate ZRINIT should retry ZFIN, got {action:?}"),
}
assert_eq!(sender.poll(), Action::Idle);
let zfin = write_header(Header::new(Encoding::ZHEX, Frame::ZFIN, [0; 4]));
sender.submit_wire(&zfin).unwrap();
let Action::WriteWire(bytes) = sender.poll() else {
panic!("sender should flush OO before completion");
};
assert_eq!(bytes, b"OO");
let len = bytes.len();
sender.wire_written(len);
assert_eq!(sender.poll(), Action::Event(Event::SessionCompleted));
}
#[test]
+1 -1
View File
@@ -2,7 +2,7 @@
// Copyright (c) 2017-2020 Alexey Arbuzov
// Copyright (c) 2023-2026 Jarkko Sakkinen
#![cfg(has_lrzsz)]
#![cfg(all(has_lrzsz, unix))]
use nix::fcntl::{self, OFlag};
use std::cmp::{max, min};
+3 -2
View File
@@ -22,7 +22,7 @@ use crate::exec_supervisor::{ExecEffect, ExecPhase, ExecSupervisor, TerminalInpu
use crate::osc::extract_osc_events;
use crate::osc::{OscEvent, OscStreamParser};
use crate::recording::RecordingTap;
use crate::zmodem::{ZmodemTransferId, ZmodemTransferOutcome};
use crate::zmodem::{ZmodemTransferId, ZmodemTransferOutcome, ZmodemTransferProgress};
use crate::{
TerminalBackend, TerminalControlError, TerminalControlHandle, TerminalControlOutput,
TerminalControlRequest, TerminalExecError, TerminalExecHandle, TerminalExecOutput,
@@ -40,11 +40,12 @@ pub enum TerminalEvent {
/// SSH ZMODEM 文件选择请求状态变化
ZmodemRequestChanged,
/// SSH ZMODEM 文件传输进度变化
ZmodemProgressChanged(ZmodemTransferId),
ZmodemProgressChanged(ZmodemTransferProgress),
/// SSH ZMODEM 文件传输结束
ZmodemTransferFinished {
transfer_id: ZmodemTransferId,
outcome: ZmodemTransferOutcome,
progress: Option<ZmodemTransferProgress>,
},
/// shell 开始渲染新的 promptOSC 133;A
PromptStart,
+79 -1
View File
@@ -349,6 +349,62 @@ enum DeferredSshActorInput {
TerminalResponse(Vec<u8>),
}
fn defer_zmodem_actor_command(
deferred_inputs: &mut VecDeque<DeferredSshActorInput>,
command: SshCommand,
) {
match command {
SshCommand::CancelExec { id } => {
if let Some(result) = take_deferred_exec(deferred_inputs, id) {
let _ = result.send(Err(TerminalExecError::CancelledBeforeSubmit));
} else {
deferred_inputs.push_back(DeferredSshActorInput::Command(SshCommand::CancelExec {
id,
}));
}
}
SshCommand::ExecTimeout { id, phase } => {
if let Some(result) = take_deferred_exec(deferred_inputs, id) {
let error = match phase {
ExecPhase::WaitingForReady | ExecPhase::Observing => {
TerminalExecError::ReadyTimeout
}
ExecPhase::ClearingInput => TerminalExecError::ClearInputTimeout,
};
let _ = result.send(Err(error));
} else {
deferred_inputs.push_back(DeferredSshActorInput::Command(
SshCommand::ExecTimeout { id, phase },
));
}
}
command => {
deferred_inputs.push_back(DeferredSshActorInput::Command(command));
}
}
}
fn take_deferred_exec(
deferred_inputs: &mut VecDeque<DeferredSshActorInput>,
id: u64,
) -> Option<ExecResultSender> {
let index = deferred_inputs.iter().position(|input| {
matches!(
input,
DeferredSshActorInput::Command(SshCommand::StartExec {
id: deferred_id,
..
}) if *deferred_id == id
)
})?;
let DeferredSshActorInput::Command(SshCommand::StartExec { result, .. }) =
deferred_inputs.remove(index)?
else {
return None;
};
Some(result)
}
const SSH_TERMINAL_INPUT_CHUNK_BYTES: usize = 4 * 1024;
#[derive(Default)]
@@ -481,7 +537,7 @@ async fn run_zmodem_transfer_while_servicing_actor<C: SshChannel>(
// ZMODEM owns the SSH channel until its protocol session
// has ended. Keep the actor responsive without injecting
// terminal input, resize, or exec bytes into that session.
deferred_inputs.push_back(DeferredSshActorInput::Command(command));
defer_zmodem_actor_command(deferred_inputs, command);
}
}
Some(data) = terminal_response_rx.recv() => {
@@ -1713,6 +1769,28 @@ mod tests {
);
}
#[tokio::test]
async fn cancelled_exec_is_removed_from_zmodem_deferred_inputs() {
let mut deferred_inputs = VecDeque::new();
let (result_tx, result_rx) = oneshot::channel();
defer_zmodem_actor_command(
&mut deferred_inputs,
SshCommand::StartExec {
id: 42,
request: request("must-not-run"),
result: result_tx,
},
);
defer_zmodem_actor_command(&mut deferred_inputs, SshCommand::CancelExec { id: 42 });
assert!(deferred_inputs.is_empty());
assert_eq!(
Err(TerminalExecError::CancelledBeforeSubmit),
result_rx.await.expect("deferred result")
);
}
#[test]
fn ssh_backend_records_direct_and_handle_input_without_double_counting() {
let (command_tx, mut command_rx) = unbounded_channel();
+6 -3
View File
@@ -109,11 +109,12 @@ pub enum TerminalModelEvent {
/// SSH ZMODEM 文件选择请求状态变化
ZmodemRequestChanged,
/// SSH ZMODEM 文件传输进度变化
ZmodemProgressChanged(ZmodemTransferId),
ZmodemProgressChanged(ZmodemTransferProgress),
/// SSH ZMODEM 文件传输结束
ZmodemTransferFinished {
transfer_id: ZmodemTransferId,
outcome: ZmodemTransferOutcome,
progress: Option<ZmodemTransferProgress>,
},
/// shell 开始渲染新的 promptOSC 133;A
PromptStart,
@@ -3144,16 +3145,18 @@ impl Terminal {
TerminalEvent::ZmodemRequestChanged => {
cx.emit(TerminalModelEvent::ZmodemRequestChanged);
}
TerminalEvent::ZmodemProgressChanged(transfer_id) => {
cx.emit(TerminalModelEvent::ZmodemProgressChanged(transfer_id));
TerminalEvent::ZmodemProgressChanged(progress) => {
cx.emit(TerminalModelEvent::ZmodemProgressChanged(progress));
}
TerminalEvent::ZmodemTransferFinished {
transfer_id,
outcome,
progress,
} => {
cx.emit(TerminalModelEvent::ZmodemTransferFinished {
transfer_id,
outcome,
progress,
});
}
TerminalEvent::PromptStart => {
+40 -5
View File
@@ -249,11 +249,13 @@ impl ZmodemResponder {
}
fn release_picker_claim(&self, request_id: u64) {
let released = self
.state
.lock()
.ok()
.is_some_and(|mut state| state.picker_claim.take() == Some(request_id));
let released = self.state.lock().ok().is_some_and(|mut state| {
if state.picker_claim != Some(request_id) {
return false;
}
state.picker_claim = None;
true
});
if released {
self.notify_changed();
}
@@ -474,4 +476,37 @@ mod tests {
assert!(claim.submit(ZmodemPickerResponse::Cancel));
assert!(second.await.unwrap().is_ok());
}
#[tokio::test]
async fn dropping_stale_picker_claim_does_not_release_new_claim() {
let (event_tx, mut event_rx) = unbounded_channel();
let responder = ZmodemResponder::new(event_tx);
let first_responder = responder.clone();
let first =
tokio::spawn(
async move { first_responder.request(ZmodemPickerKind::UploadFiles).await },
);
event_rx.recv().await.expect("first request event");
let first_id = responder.pending_request().unwrap().id();
let stale_claim = responder.try_claim_picker(first_id).unwrap();
assert!(responder.cancel());
assert!(first.await.unwrap().is_err());
event_rx.recv().await.expect("first clear event");
let second_responder = responder.clone();
let second = tokio::spawn(async move {
second_responder
.request(ZmodemPickerKind::DownloadDirectory)
.await
});
event_rx.recv().await.expect("second request event");
let second_id = responder.pending_request().unwrap().id();
let second_claim = responder.try_claim_picker(second_id).unwrap();
drop(stale_claim);
assert!(second_claim.submit(ZmodemPickerResponse::Cancel));
assert!(second.await.unwrap().is_ok());
}
}
+19 -2
View File
@@ -1,7 +1,10 @@
use super::{
ZmodemResponder, ZmodemTransferDirection, ZmodemTransferId, ZmodemTransferProgress,
download_path,
transfer::{MAX_PROTOCOL_TIMEOUTS, consume_hex_header_terminator, receive_wire, send_wire},
transfer::{
MAX_PROTOCOL_TIMEOUTS, consume_hex_header_terminator, receive_finish_wire, receive_wire,
send_wire,
},
};
use anyhow::{Context as _, Result, bail};
use ssh::SshChannel;
@@ -45,6 +48,20 @@ pub(super) async fn run_download(
) -> Result<Vec<u8>> {
validate_directory(&directory).await?;
let mut receiver = Receiver::new().context("create ZMODEM receiver")?;
if !initial_wire.is_empty() {
// The receiver pre-queues a ZRINIT before it has seen the remote
// ZRQINIT. In the transfer flow the initial wire already contains
// that request, so drop the premature handshake and let the
// ZRQINIT handler emit the response. This prevents lrzsz `sz -e`
// from restarting its handshake when it receives a duplicate ZRINIT.
let initial_wire_len = match receiver.poll() {
zmodem2::Action::WriteWire(bytes) => Some(bytes.len()),
_ => None,
};
if let Some(initial_wire_len) = initial_wire_len {
receiver.wire_written(initial_wire_len);
}
}
let mut current = None;
let mut progress = DownloadProgressTracker {
responder,
@@ -132,7 +149,7 @@ async fn finish_download(
if pending.len() >= 2 || pending.first().is_some_and(|byte| *byte != b'O') {
return Ok(pending);
}
let Some(data) = receive_wire(channel, cancellation).await? else {
let Some(data) = receive_finish_wire(channel, cancellation).await? else {
return Ok(pending);
};
pending.extend(data);
@@ -79,6 +79,16 @@ async fn cancelled_picker_sends_zcan() {
let response_task = tokio::spawn(async move {
event_rx.recv().await.expect("picker request event");
assert!(task_responder.submit(ZmodemPickerResponse::Cancel));
loop {
match event_rx.recv().await.expect("transfer outcome event") {
TerminalEvent::ZmodemTransferFinished {
outcome: super::ZmodemTransferOutcome::Cancelled,
..
} => break,
TerminalEvent::ZmodemRequestChanged | TerminalEvent::ZmodemProgressChanged(_) => {}
event => panic!("unexpected terminal event: {event:?}"),
}
}
});
let mut channel = MockChannel::default();
let sent = channel.sent.clone();
@@ -114,6 +124,7 @@ async fn cancellation_during_picker_reports_cancelled_outcome() {
TerminalEvent::ZmodemTransferFinished {
transfer_id: _,
outcome: super::ZmodemTransferOutcome::Cancelled,
..
} => {
break;
}
+70 -48
View File
@@ -111,10 +111,7 @@ struct ProgressState {
snapshot: Option<ZmodemTransferProgress>,
}
pub(crate) struct TransferProgressGuard {
progress: ZmodemProgressState,
transfer_id: ZmodemTransferId,
}
pub(crate) struct TransferProgressGuard;
impl ZmodemProgressState {
pub(crate) fn new(event_tx: UnboundedSender<TerminalEvent>) -> Self {
@@ -133,10 +130,18 @@ impl ZmodemProgressState {
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let displaced = state
.active_id
.take()
.map(|transfer_id| (transfer_id, state.snapshot.take()));
state.next_id = state.next_id.wrapping_add(1).max(1);
let transfer_id = ZmodemTransferId(state.next_id);
state.active_id = Some(transfer_id);
state.snapshot = None;
drop(state);
if let Some((displaced_id, progress)) = displaced {
self.notify_finished(displaced_id, ZmodemTransferOutcome::Cancelled, progress);
}
transfer_id
}
@@ -146,10 +151,7 @@ impl ZmodemProgressState {
progress: ZmodemTransferProgress,
) -> TransferProgressGuard {
self.set(transfer_id, progress, true);
TransferProgressGuard {
progress: self.clone(),
transfer_id,
}
TransferProgressGuard
}
pub(crate) fn start(&self, transfer_id: ZmodemTransferId, progress: ZmodemTransferProgress) {
@@ -172,68 +174,57 @@ impl ZmodemProgressState {
mut progress: ZmodemTransferProgress,
force: bool,
) {
progress.transfer_id = transfer_id;
let notification = progress.clone();
let changed = self.state.lock().ok().is_some_and(|mut state| {
if state.active_id != Some(transfer_id) {
return false;
}
progress.transfer_id = transfer_id;
let changed = force || state.snapshot.as_ref() != Some(&progress);
state.snapshot = Some(progress);
changed
});
if changed {
self.notify(transfer_id);
}
}
fn clear(&self, transfer_id: ZmodemTransferId) {
let cleared = self.state.lock().ok().is_some_and(|mut state| {
if state.active_id != Some(transfer_id) {
return false;
}
state.snapshot.take().is_some()
});
if cleared {
self.notify(transfer_id);
self.notify(notification);
}
}
fn finish_inner(&self, transfer_id: ZmodemTransferId, outcome: ZmodemTransferOutcome) {
let should_notify = self.state.lock().ok().is_some_and(|mut state| {
let Some(progress) = self.state.lock().ok().and_then(|mut state| {
if state.active_id != Some(transfer_id) {
return false;
return None;
}
state.snapshot.take();
let progress = state.snapshot.take();
state.active_id = None;
true
});
if should_notify {
self.notify_finished(transfer_id, outcome);
}
Some(progress)
}) else {
return;
};
self.notify_finished(transfer_id, outcome, progress);
}
fn notify(&self, transfer_id: ZmodemTransferId) {
fn notify(&self, progress: ZmodemTransferProgress) {
if let Some(event_tx) = &self.event_tx {
let _ = event_tx.send(TerminalEvent::ZmodemProgressChanged(transfer_id));
let _ = event_tx.send(TerminalEvent::ZmodemProgressChanged(progress));
}
}
fn notify_finished(&self, transfer_id: ZmodemTransferId, outcome: ZmodemTransferOutcome) {
fn notify_finished(
&self,
transfer_id: ZmodemTransferId,
outcome: ZmodemTransferOutcome,
progress: Option<ZmodemTransferProgress>,
) {
if let Some(event_tx) = &self.event_tx {
let _ = event_tx.send(TerminalEvent::ZmodemTransferFinished {
transfer_id,
outcome,
progress,
});
}
}
}
impl Drop for TransferProgressGuard {
fn drop(&mut self) {
self.progress.clear(self.transfer_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -253,7 +244,7 @@ mod tests {
}
#[tokio::test]
async fn upload_progress_notifies_on_transferred_change_and_clear() {
async fn upload_progress_notifies_on_transferred_change() {
let (event_tx, mut event_rx) = unbounded_channel();
let progress_state = ZmodemProgressState::new(event_tx);
let transfer_id = progress_state.begin_transfer(ZmodemTransferDirection::Upload);
@@ -269,8 +260,8 @@ mod tests {
assert_eq!(10, progress_state.snapshot().unwrap().transferred());
drop(guard);
assert!(progress_state.snapshot().is_none());
assert_progress_event_for(&mut event_rx, transfer_id).await;
assert_eq!(10, progress_state.snapshot().unwrap().transferred());
assert!(event_rx.try_recv().is_err());
}
#[tokio::test]
@@ -282,6 +273,7 @@ mod tests {
assert_progress_event_for(&mut event_rx, stale_id).await;
let current_id = progress_state.begin_transfer(ZmodemTransferDirection::Upload);
assert_cancelled_event_for(&mut event_rx, stale_id).await;
let _current_guard = progress_state.begin(current_id, upload_progress(2, 1_000));
assert_progress_event_for(&mut event_rx, current_id).await;
@@ -305,6 +297,7 @@ mod tests {
assert_progress_event_for(&mut event_rx, stale_id).await;
let current_id = progress_state.begin_transfer(ZmodemTransferDirection::Download);
assert_cancelled_event_for(&mut event_rx, stale_id).await;
progress_state.start(current_id, download_progress(2));
assert_progress_event_for(&mut event_rx, current_id).await;
@@ -333,6 +326,7 @@ mod tests {
Some(TerminalEvent::ZmodemTransferFinished {
transfer_id: id,
outcome: ZmodemTransferOutcome::Succeeded,
..
}) if id == transfer_id
));
drop(guard);
@@ -340,14 +334,15 @@ mod tests {
}
#[tokio::test]
async fn explicit_finish_still_notifies_after_guard_clears_snapshot() {
async fn explicit_finish_still_notifies_after_guard_drop() {
let (event_tx, mut event_rx) = unbounded_channel();
let progress_state = ZmodemProgressState::new(event_tx);
let transfer_id = progress_state.begin_transfer(ZmodemTransferDirection::Upload);
let guard = progress_state.begin(transfer_id, upload_progress(1, 1_000));
assert_progress_event_for(&mut event_rx, transfer_id).await;
drop(guard);
assert_progress_event_for(&mut event_rx, transfer_id).await;
assert!(progress_state.snapshot().is_some());
assert!(event_rx.try_recv().is_err());
progress_state.finish(transfer_id, ZmodemTransferOutcome::Succeeded);
assert!(matches!(
@@ -355,6 +350,7 @@ mod tests {
Some(TerminalEvent::ZmodemTransferFinished {
transfer_id: id,
outcome: ZmodemTransferOutcome::Succeeded,
..
}) if id == transfer_id
));
}
@@ -373,6 +369,7 @@ mod tests {
Some(TerminalEvent::ZmodemTransferFinished {
transfer_id: id,
outcome: ZmodemTransferOutcome::Cancelled,
..
}) if id == transfer_id
));
drop(guard);
@@ -396,6 +393,7 @@ mod tests {
Some(TerminalEvent::ZmodemTransferFinished {
transfer_id: id,
outcome: ZmodemTransferOutcome::Failed(error),
..
}) if id == transfer_id && error == "boom"
));
drop(guard);
@@ -415,7 +413,8 @@ mod tests {
);
assert!(matches!(
event_rx.try_recv(),
Ok(TerminalEvent::ZmodemProgressChanged(id)) if id == transfer_id
Ok(TerminalEvent::ZmodemProgressChanged(progress))
if progress.transfer_id() == transfer_id
));
assert!(event_rx.try_recv().is_err());
}
@@ -442,13 +441,14 @@ mod tests {
Some(TerminalEvent::ZmodemTransferFinished {
transfer_id: id,
outcome: ZmodemTransferOutcome::Succeeded,
..
}) if id == transfer_id
));
assert!(progress_state.snapshot().is_none());
}
#[test]
fn begin_transfer_resets_the_previous_snapshot_without_notifying() {
fn begin_transfer_cancels_the_previous_snapshot() {
let (event_tx, mut event_rx) = unbounded_channel();
let progress_state = ZmodemProgressState::new(event_tx);
let upload_id = progress_state.begin_transfer(ZmodemTransferDirection::Upload);
@@ -458,7 +458,14 @@ mod tests {
let download_id = progress_state.begin_transfer(ZmodemTransferDirection::Download);
assert_ne!(upload_id, download_id);
assert!(progress_state.snapshot().is_none());
assert!(event_rx.try_recv().is_err());
assert!(matches!(
event_rx.try_recv(),
Ok(TerminalEvent::ZmodemTransferFinished {
transfer_id,
outcome: ZmodemTransferOutcome::Cancelled,
..
}) if transfer_id == upload_id
));
drop(guard);
assert!(event_rx.try_recv().is_err());
}
@@ -469,7 +476,22 @@ mod tests {
) {
assert!(matches!(
event_rx.recv().await,
Some(TerminalEvent::ZmodemProgressChanged(id)) if id == transfer_id
Some(TerminalEvent::ZmodemProgressChanged(progress))
if progress.transfer_id() == transfer_id
));
}
async fn assert_cancelled_event_for(
event_rx: &mut tokio::sync::mpsc::UnboundedReceiver<TerminalEvent>,
transfer_id: ZmodemTransferId,
) {
assert!(matches!(
event_rx.recv().await,
Some(TerminalEvent::ZmodemTransferFinished {
transfer_id: id,
outcome: ZmodemTransferOutcome::Cancelled,
..
}) if id == transfer_id
));
}
+32 -5
View File
@@ -12,12 +12,16 @@ pub(crate) const ZCAN: &[u8] = b"\x18\x18\x18\x18\x18\x18\x18\x18\x08\x08\x08\x0
const SEND_TIMEOUT: Duration = Duration::from_secs(30);
const CANCEL_SEND_TIMEOUT: Duration = Duration::from_secs(1);
const RECEIVE_TIMEOUT: Duration = Duration::from_secs(10);
const FINISH_RECEIVE_TIMEOUT: Duration = Duration::from_secs(1);
const MAX_PICKER_WIRE_BYTES: usize = 1024 * 1024;
pub(super) const MAX_PROTOCOL_TIMEOUTS: usize = 6;
#[derive(Debug)]
struct ChannelClosed;
#[derive(Debug)]
struct TransferCancelled;
impl fmt::Display for ChannelClosed {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("SSH channel closed before ZMODEM transfer completed")
@@ -26,6 +30,14 @@ impl fmt::Display for ChannelClosed {
impl StdError for ChannelClosed {}
impl fmt::Display for TransferCancelled {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("ZMODEM transfer was cancelled")
}
}
impl StdError for TransferCancelled {}
pub(crate) async fn run_transfer(
channel: &mut dyn SshChannel,
detected: DetectedZmodem,
@@ -48,7 +60,7 @@ pub(crate) async fn run_transfer(
match result {
Ok(_) => ZmodemTransferOutcome::Succeeded,
Err(ref error) => {
if was_cancelled {
if was_cancelled || error.downcast_ref::<TransferCancelled>().is_some() {
ZmodemTransferOutcome::Cancelled
} else {
ZmodemTransferOutcome::Failed(format!("{error:#}"))
@@ -78,7 +90,7 @@ async fn run_selected_transfer(
match direction {
ZmodemDirection::Upload => {
let ZmodemPickerResponse::UploadFiles(paths) = response else {
bail!("ZMODEM upload was cancelled");
return Err(TransferCancelled.into());
};
let request = super::upload::UploadRequest {
initial_wire: wire,
@@ -90,7 +102,7 @@ async fn run_selected_transfer(
}
ZmodemDirection::Download => {
let ZmodemPickerResponse::DownloadDirectory(directory) = response else {
bail!("ZMODEM download was cancelled");
return Err(TransferCancelled.into());
};
super::download::run_download(
channel,
@@ -134,7 +146,7 @@ pub(super) async fn consume_hex_header_terminator(
return Ok(());
}
while pending.len() < 2 {
let Some(data) = receive_wire(channel, cancellation).await? else {
let Some(data) = receive_finish_wire(channel, cancellation).await? else {
return Ok(());
};
pending.extend(data);
@@ -199,11 +211,26 @@ pub(super) async fn send_wire(
pub(super) async fn receive_wire(
channel: &mut dyn SshChannel,
cancellation: &CancellationToken,
) -> Result<Option<Vec<u8>>> {
receive_wire_with_timeout(channel, cancellation, RECEIVE_TIMEOUT).await
}
pub(super) async fn receive_finish_wire(
channel: &mut dyn SshChannel,
cancellation: &CancellationToken,
) -> Result<Option<Vec<u8>>> {
receive_wire_with_timeout(channel, cancellation, FINISH_RECEIVE_TIMEOUT).await
}
async fn receive_wire_with_timeout(
channel: &mut dyn SshChannel,
cancellation: &CancellationToken,
receive_timeout: Duration,
) -> Result<Option<Vec<u8>>> {
loop {
let event = tokio::select! {
_ = cancellation.cancelled() => bail!("ZMODEM transfer was cancelled"),
result = timeout(RECEIVE_TIMEOUT, channel.recv()) => {
result = timeout(receive_timeout, channel.recv()) => {
match result {
Ok(event) => event,
Err(_) => return Ok(None),
+5
View File
@@ -187,6 +187,11 @@ async fn drive_sender(
progress.start_file(current.as_ref());
Ok(SenderStep::Progress)
}
Action::Event(Event::FileSkipped) => {
*current = start_next(sender, queue, progress.file_count).await?;
progress.start_file(current.as_ref());
Ok(SenderStep::Progress)
}
Action::Event(Event::SessionCompleted) => Ok(SenderStep::Complete),
Action::Event(Event::Aborted) => bail!("remote aborted ZMODEM upload"),
Action::Event(Event::FileStarted(_)) | Action::WriteFile(_) => {
@@ -8,22 +8,15 @@ impl TerminalView {
/// 将 ZMODEM 传输进度同步到全局后台任务面板。
pub(super) fn sync_zmodem_background_task(
&mut self,
expected_transfer_id: Option<ZmodemTransferId>,
progress: Option<terminal::zmodem::ZmodemTransferProgress>,
cx: &mut Context<Self>,
) {
let Some(progress) = self.terminal.read(cx).zmodem_transfer_progress() else {
let Some(progress) = progress.or_else(|| self.terminal.read(cx).zmodem_transfer_progress())
else {
return;
};
if expected_transfer_id.is_some_and(|id| id != progress.transfer_id()) {
return;
}
let direction = progress.direction();
let entity_id = self.terminal.entity_id().as_u64();
let key = SharedString::from(format!(
"zmodem-{}:{entity_id}:{}",
direction.as_str(),
progress.transfer_id().as_u64()
));
let key = self.zmodem_background_task_key(&progress);
let title = progress.file_name().to_string();
let file_number = progress.file_index().saturating_add(1);
let file_count = progress.file_count();
@@ -42,15 +35,17 @@ impl TerminalView {
format_zmodem_bytes(total)
)
};
let detail = format!("{file_progress} · {byte_progress}");
let detail = format!("{title} · {file_progress} · {byte_progress}");
let manager = background_tasks::global(cx);
let active_id = manager.read(cx).find_by_key(&key);
let existing_id = self
.zmodem_background_tasks
.get(&progress.transfer_id())
.copied()
.filter(|id| manager.read(cx).find_by_key(&key) == Some(*id));
.filter(|id| active_id == Some(*id))
.or(active_id);
let id = if let Some(id) = existing_id {
id
} else {
@@ -69,8 +64,8 @@ impl TerminalView {
if let Some(cancellation) = cancellation {
manager.set_cancellation(
id,
BackgroundTaskCancellation::callback(move || {
cancellation.cancel();
BackgroundTaskCancellation::callback_with_result(move || {
cancellation.cancel()
}),
cx,
);
@@ -100,18 +95,58 @@ impl TerminalView {
&mut self,
transfer_id: ZmodemTransferId,
outcome: &ZmodemTransferOutcome,
progress: Option<terminal::zmodem::ZmodemTransferProgress>,
cx: &mut Context<Self>,
) {
let Some(id) = self.zmodem_background_tasks.remove(&transfer_id) else {
let manager = background_tasks::global(cx);
let mut id = self.zmodem_background_tasks.remove(&transfer_id);
if id.is_none() {
if let Some(progress) = progress.as_ref() {
let key = self.zmodem_background_task_key(progress);
id = manager.read(cx).find_latest_by_key(&key);
}
}
if id.is_none() {
if let Some(progress) = progress {
self.sync_zmodem_background_task(Some(progress), cx);
id = self.zmodem_background_tasks.remove(&transfer_id);
}
}
let Some(id) = id else {
return;
};
let manager = background_tasks::global(cx);
manager.update(cx, |manager, cx| match outcome {
ZmodemTransferOutcome::Succeeded => manager.succeed(id, None, cx),
ZmodemTransferOutcome::Cancelled => manager.cancel_confirmed(id, None, cx),
ZmodemTransferOutcome::Failed(error) => manager.fail(id, error.clone(), cx),
});
}
fn zmodem_background_task_key(
&self,
progress: &terminal::zmodem::ZmodemTransferProgress,
) -> SharedString {
let entity_id = self.terminal.entity_id().as_u64();
SharedString::from(format!(
"zmodem-{}:{entity_id}:{}",
progress.direction().as_str(),
progress.transfer_id().as_u64()
))
}
pub(super) fn cancel_zmodem_background_tasks(&mut self, cx: &mut App) {
self.terminal.read(cx).cancel_zmodem_transfer();
let transfer_tasks = std::mem::take(&mut self.zmodem_background_tasks);
if transfer_tasks.is_empty() {
return;
}
let manager = background_tasks::global(cx);
manager.update(cx, |manager, cx| {
for task_id in transfer_tasks.into_values() {
manager.cancel_confirmed(task_id, None, cx);
}
});
}
}
fn format_zmodem_bytes(bytes: u64) -> String {
+1
View File
@@ -16,6 +16,7 @@ impl TerminalView {
self.unregister_broadcast_input(cx);
self.unregister_public_mcp_session(cx);
self.release_active_connection(cx);
self.cancel_zmodem_background_tasks(cx);
self.terminal.read(cx).shutdown();
}
@@ -296,6 +296,7 @@ impl TerminalView {
this.sync_ssh_mfa_inputs(window, cx);
this.register_broadcast_input(cx);
this.start_performance_diagnostics(connection_id, connection_kind, cx);
cx.on_release(|this, cx| this.cancel_zmodem_background_tasks(cx));
this
}
}
@@ -71,15 +71,16 @@ impl TerminalView {
TerminalModelEvent::ZmodemRequestChanged => {
self.sync_zmodem_picker(cx);
}
TerminalModelEvent::ZmodemProgressChanged(transfer_id) => {
self.sync_zmodem_background_task(Some(*transfer_id), cx);
TerminalModelEvent::ZmodemProgressChanged(progress) => {
self.sync_zmodem_background_task(Some(progress.clone()), cx);
cx.notify();
}
TerminalModelEvent::ZmodemTransferFinished {
transfer_id,
outcome,
progress,
} => {
self.finish_zmodem_background_task(*transfer_id, outcome, cx);
self.finish_zmodem_background_task(*transfer_id, outcome, progress.clone(), cx);
cx.notify();
}
TerminalModelEvent::PromptStart
@@ -483,10 +483,9 @@ fn zmodem_progress_uses_only_the_global_background_task_panel() {
assert!(!view_source.contains("mod zmodem_progress;"));
assert!(!render_source.contains("render_zmodem_progress"));
assert!(event_source.contains("self.sync_zmodem_background_task(None, cx);"));
assert!(event_source.contains("self.sync_zmodem_background_task(Some(*transfer_id), cx);"));
assert!(
event_source.contains("self.finish_zmodem_background_task(*transfer_id, outcome, cx);")
);
assert!(event_source.contains("self.sync_zmodem_background_task(Some(progress.clone()), cx);"));
assert!(event_source.contains("self.finish_zmodem_background_task("));
assert!(event_source.contains("progress.clone(),"));
assert!(background_task_source.contains(r#"BackgroundTaskSpec::new("zmodem-transfer""#));
}