use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use anyhow::{anyhow, Result};
use super::protocol::*;
use super::types::{DfuPhase, DfuProgress};
const REPORT_ID_OUTPUT: u8 = crate::kernel::protocol_hid::REPORT_ID_OUTPUT;
pub(crate) const VID: u16 = crate::kernel::protocol_hid::VID;
pub(crate) const NORMAL_PID: u16 = 0xED20;
pub type ProgressCallback = Arc<dyn Fn(DfuProgress) + Send + Sync>;
pub struct DfuClient {
cancel_flag: Arc<AtomicBool>,
on_progress: ProgressCallback,
}
impl DfuClient {
pub fn new(cancel_flag: Arc<AtomicBool>, on_progress: ProgressCallback) -> Self {
Self {
cancel_flag,
on_progress,
}
}
pub fn upgrade<F>(&self, firmware: &[u8], enter_dfu: F) -> Result<()>
where
F: FnOnce() -> Result<()>,
{
let total_len = firmware.len() as u32;
if total_len == 0 {
return Err(anyhow!("固件文件为空"));
}
log::info!(target: "board", "[dfu] 开始升级 ({} bytes)", total_len);
self.emit(DfuPhase::EnteringDfu, 0, total_len, None);
log::info!(target: "board", "[dfu] 进入 DFU 模式...");
if let Err(e) = enter_dfu() {
let msg = format!("进入 DFU 失败: {e}");
self.emit_failed(0, total_len, &msg);
return Err(anyhow!(msg));
}
self.emit(
DfuPhase::EnteringDfu,
0,
total_len,
Some("等待 DFU 设备...".into()),
);
let dfu_device = match self.wait_for_dfu_device() {
Ok(d) => d,
Err(e) => {
let msg = format!("{e}");
self.emit_failed(0, total_len, &msg);
return Err(anyhow!(msg));
}
};
log::info!(target: "board", "[dfu] DFU 设备已连接 (PID 0xFF06)");
self.emit(DfuPhase::Preparing, 0, total_len, None);
if let Err(e) = self.run_prepare(&dfu_device, total_len) {
self.fail_with_reset(&dfu_device, 0, total_len, "PREPARE", e);
return Err(anyhow!("PREPARE 失败"));
}
log::debug!(target: "board", "[dfu] PREPARE 成功");
if let Err(e) = self.run_start(&dfu_device) {
self.fail_with_reset(&dfu_device, 0, total_len, "START", e);
return Err(anyhow!("START 失败"));
}
log::debug!(target: "board", "[dfu] START 成功");
self.emit(DfuPhase::Transferring, 0, total_len, None);
self.run_data_loop(&dfu_device, firmware, total_len)?;
log::info!(target: "board", "[dfu] DATA 传输完成 ({} bytes)", total_len);
self.emit(
DfuPhase::Verifying,
total_len,
total_len,
Some("设备正在验证固件...".into()),
);
match self.send_and_recv(&dfu_device, &DfuPacketEncoder::end(), DFU_END_TIMEOUT_MS) {
Ok(resp) if resp.is_success() => {
log::info!(target: "board", "[dfu] END 验证成功,设备即将重启");
}
Ok(_) => {
log::warn!(target: "board", "[dfu] END 验证失败,设备可能重启回旧固件");
}
Err(e) => {
log::warn!(target: "board", "[dfu] END 通信失败: {e}(设备可能需要物理重插)");
}
}
self.emit(
DfuPhase::Rebooting,
total_len,
total_len,
Some("设备重启中...".into()),
);
self.wait_for_normal_device();
self.emit(DfuPhase::Completed, total_len, total_len, None);
log::info!(target: "board", "[dfu] 固件升级完成");
Ok(())
}
fn run_prepare(&self, device: &hidapi::HidDevice, total_len: u32) -> Result<()> {
let resp = self.send_and_recv_with_retry(
device,
&DfuPacketEncoder::prepare(total_len),
DFU_PREPARE_TIMEOUT_MS,
2,
"PREPARE",
)?;
if !resp.is_success() {
return Err(anyhow!("设备返回错误 result=0x{:02X}", resp.result));
}
Ok(())
}
fn run_start(&self, device: &hidapi::HidDevice) -> Result<()> {
let resp = self.send_and_recv_with_retry(
device,
&DfuPacketEncoder::start(),
DFU_PREPARE_TIMEOUT_MS,
2,
"START",
)?;
if !resp.is_success() {
return Err(anyhow!("设备返回错误 result=0x{:02X}", resp.result));
}
Ok(())
}
fn run_data_loop(
&self,
device: &hidapi::HidDevice,
firmware: &[u8],
total_len: u32,
) -> Result<()> {
let mut offset: usize = 0;
let mut retry_count: u8 = 0;
const MAX_RETRY: u8 = 3;
while offset < firmware.len() {
if self.cancel_flag.load(Ordering::SeqCst) {
log::warn!(target: "board", "[dfu] 用户取消 (offset={})", offset);
self.try_reset_device(device);
let msg = "用户取消升级".to_string();
self.emit_failed(offset as u32, total_len, &msg);
return Err(anyhow!(msg));
}
let end = (offset + DFU_DATA_PAYLOAD_MAX).min(firmware.len());
let chunk = &firmware[offset..end];
let packet = DfuPacketEncoder::data(chunk)?;
match self.send_and_recv(device, &packet, DFU_RW_TIMEOUT_MS) {
Ok(resp) if resp.is_success() => {
offset = end;
retry_count = 0;
self.emit(DfuPhase::Transferring, offset as u32, total_len, None);
}
Ok(resp) => {
retry_count += 1;
log::warn!(
target: "board",
"[dfu] DATA 失败 (offset={}, retry={}/{}): result=0x{:02X}",
offset, retry_count, MAX_RETRY, resp.result
);
if retry_count >= MAX_RETRY {
self.try_reset_device(device);
let msg =
format!("DATA 传输失败,重试 {MAX_RETRY} 次后放弃 (offset={offset})");
self.emit_failed(offset as u32, total_len, &msg);
return Err(anyhow!(msg));
}
std::thread::sleep(Duration::from_millis(100));
}
Err(e) => {
retry_count += 1;
log::warn!(
target: "board",
"[dfu] DATA 通信错误 (offset={}, retry={}/{}): {}",
offset, retry_count, MAX_RETRY, e
);
if retry_count >= MAX_RETRY {
self.try_reset_device(device);
let msg = format!("DATA 通信失败: {e} (offset={offset})");
self.emit_failed(offset as u32, total_len, &msg);
return Err(anyhow!(msg));
}
std::thread::sleep(Duration::from_millis(100));
}
}
}
Ok(())
}
fn try_reset_device(&self, device: &hidapi::HidDevice) {
if super::recover::send_recovery_sequence_on(device) {
self.wait_for_normal_device();
} else {
log::warn!(target: "board", "[dfu] 复位包未能送达设备,跳过等待(设备可能需要物理重插)");
}
}
fn send_and_recv(
&self,
device: &hidapi::HidDevice,
packet: &[u8],
timeout_ms: i32,
) -> Result<DfuResponse> {
send_and_recv(device, packet, timeout_ms)
}
fn send_and_recv_with_retry(
&self,
device: &hidapi::HidDevice,
packet: &[u8],
timeout_ms: i32,
max_attempts: u8,
label: &str,
) -> Result<DfuResponse> {
let mut last_err = String::new();
for attempt in 1..=max_attempts {
if attempt > 1 {
log::warn!(target: "board", "[dfu] {label} 超时,重试 {attempt}/{max_attempts}");
std::thread::sleep(Duration::from_millis(500));
}
match self.send_and_recv(device, packet, timeout_ms) {
Ok(resp) => return Ok(resp),
Err(e) => last_err = e.to_string(),
}
}
Err(anyhow!("{label} 失败: {last_err}"))
}
fn wait_for_dfu_device(&self) -> Result<hidapi::HidDevice> {
let poll = Duration::from_millis(500);
let timeout = Duration::from_secs(WAIT_DFU_DEVICE_TIMEOUT_SECS);
let start = std::time::Instant::now();
while start.elapsed() < timeout {
if self.cancel_flag.load(Ordering::SeqCst) {
return Err(anyhow!("用户取消"));
}
if let Ok(dev) = open_dfu_device() {
return Ok(dev);
}
std::thread::sleep(poll);
}
Err(anyhow!(
"等待 DFU 设备超时 ({}s)",
WAIT_DFU_DEVICE_TIMEOUT_SECS
))
}
fn wait_for_normal_device(&self) {
let poll = Duration::from_millis(500);
let timeout = Duration::from_secs(WAIT_NORMAL_DEVICE_TIMEOUT_SECS);
let start = std::time::Instant::now();
while start.elapsed() < timeout {
if let Ok(api) = hidapi::HidApi::new() {
let found = api
.device_list()
.any(|d| d.vendor_id() == VID && d.product_id() == NORMAL_PID);
if found {
log::info!(target: "board", "[dfu] 正常设备已重新枚举");
return;
}
}
std::thread::sleep(poll);
}
log::warn!(
target: "board",
"[dfu] 等待正常设备超时 ({}s),hotplug 将自动重连",
WAIT_NORMAL_DEVICE_TIMEOUT_SECS
);
}
fn emit(&self, phase: DfuPhase, bytes_written: u32, total_bytes: u32, message: Option<String>) {
let p = DfuProgress::new(phase, bytes_written, total_bytes, message);
log::debug!(target: "board", "[dfu] progress {:?} {}% ({}/{})", p.phase, p.percent, p.bytes_written, p.total_bytes);
(self.on_progress)(p);
}
fn emit_failed(&self, bytes_written: u32, total_bytes: u32, msg: &str) {
log::warn!(target: "board", "[dfu] 升级失败: {msg}");
let p = DfuProgress::new(
DfuPhase::Failed,
bytes_written,
total_bytes,
Some(msg.to_string()),
);
(self.on_progress)(p);
}
fn fail_with_reset(
&self,
device: &hidapi::HidDevice,
bytes_written: u32,
total_bytes: u32,
label: &str,
e: anyhow::Error,
) {
log::warn!(target: "board", "[dfu] {label} 失败: {e}");
self.try_reset_device(device);
self.emit_failed(bytes_written, total_bytes, &format!("{label} 失败: {e}"));
}
}
pub(crate) fn open_dfu_device() -> Result<hidapi::HidDevice> {
let api = hidapi::HidApi::new().map_err(|e| anyhow!("HidApi 创建失败: {e}"))?;
let dev_info = api
.device_list()
.find(|d| d.vendor_id() == VID && d.product_id() == DFU_PID)
.ok_or_else(|| anyhow!("DFU 设备未找到"))?;
let device = api
.open_path(dev_info.path())
.map_err(|e| anyhow!("打开 DFU 设备失败: {e}"))?;
device
.set_blocking_mode(true)
.map_err(|e| anyhow!("设置阻塞模式失败: {e}"))?;
Ok(device)
}
pub(crate) fn send_and_recv(
device: &hidapi::HidDevice,
packet: &[u8],
timeout_ms: i32,
) -> Result<DfuResponse> {
device
.write(packet)
.map_err(|e| anyhow!("DFU write 失败: {e}"))?;
read_response(device, timeout_ms)
}
pub(crate) fn read_response(device: &hidapi::HidDevice, timeout_ms: i32) -> Result<DfuResponse> {
let mut buf = [0u8; DFU_INPUT_MAX_SIZE];
let len = device
.read_timeout(&mut buf, timeout_ms)
.map_err(|e| anyhow!("DFU read 失败: {e}"))?;
if len == 0 {
return Err(anyhow!("DFU 响应超时 ({timeout_ms}ms)"));
}
let data_start = if len > 1 && buf[0] == DFU_REPORT_ID_INPUT {
1
} else {
0
};
DfuResponse::parse(&buf[data_start..len])
}
pub const CMD_ENTER_HID_DFU_MODE: u8 = 0xEF;
pub fn build_enter_dfu_hid_command() -> [u8; 64] {
let mut cmd = [0u8; 64];
cmd[0] = REPORT_ID_OUTPUT;
cmd[1] = CMD_ENTER_HID_DFU_MODE;
cmd
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn enter_dfu_command_layout() {
let cmd = build_enter_dfu_hid_command();
assert_eq!(cmd[0], REPORT_ID_OUTPUT);
assert_eq!(cmd[1], CMD_ENTER_HID_DFU_MODE);
assert!(cmd[2..].iter().all(|&b| b == 0));
assert_eq!(cmd.len(), 64);
}
#[test]
fn cancel_flag_is_checked() {
let cancel = Arc::new(AtomicBool::new(false));
let _client = DfuClient::new(cancel.clone(), Arc::new(|_| {}));
assert!(!cancel.load(Ordering::SeqCst));
}
}