refactor(rust/trezor-thp): add use_std feature
What changed, and why it matters
This commit is a routine code refactor for the Trezor hardware wallet firmware. It adds an optional 'use_std' feature to a Rust library so it can use standard Rust vector-based buffers when running on a full computer (for example, in test/example code), while keeping the no-standard-library mode for the embedded device. There is no indication of a security fix or vulnerability.
No security action required. Treat as normal maintenance/refactoring.
Security signals we found
No strong security signals were identified.
Evidence from the diff
The change introduces a Cargo feature ‘use_std’ to rust/trezor-thp. When enabled, the crate no longer uses #![no_std] and exposes a new buffered::Buffered wrapper that manages send/receive buffers with std::vec::Vec. The host CLI example is updated to use this wrapper, removing manual fixed-size buffers. The underlying ChannelIO API remains unchanged. No security-relevant logic is modified.
Changed components
rust/trezor-thp/Cargo.tomlrust/trezor-thp/src/lib.rsrust/trezor-thp/src/channel/mod.rsrust/trezor-thp/src/channel/buffered.rs (new)rust/trezor-thp/examples/host-cli/client.rsInspect captured patch +161 / −37
diff --git a/rust/trezor-thp/Cargo.toml b/rust/trezor-thp/Cargo.toml
index ddb319ac..99d72e87 100644
--- a/rust/trezor-thp/Cargo.toml
+++ b/rust/trezor-thp/Cargo.toml
@@ -3,6 +3,9 @@ name = "trezor-thp"
version = "0.1.0"
edition = "2024"
+[features]
+use_std = []
+
[dependencies]
log = "0.4.29"
@@ -24,6 +27,10 @@ hex = "0.4.3"
protobuf = "=3.7.2"
protobuf-codegen = "=3.7.2"
+[dev-dependencies.trezor-thp]
+features = ["use_std"]
+path = "."
+
[dev-dependencies.trezor-noise-rust-crypto]
version = "0.6.2"
default-features = false
diff --git a/rust/trezor-thp/examples/host-cli/client.rs b/rust/trezor-thp/examples/host-cli/client.rs
index 263f0f04..b7022c15 100644
--- a/rust/trezor-thp/examples/host-cli/client.rs
+++ b/rust/trezor-thp/examples/host-cli/client.rs
@@ -3,7 +3,8 @@ use std::net::{SocketAddr, UdpSocket};
use std::time::Duration;
use trezor_thp::{
- Backend, ChannelIO, Error, channel::host::ChannelOpen, credential::CredentialStore,
+ Backend, ChannelIO, Error, channel::buffered::Buffered, channel::host::ChannelOpen,
+ credential::CredentialStore,
};
use protobuf::{Enum, Message};
@@ -11,13 +12,12 @@ use protobuf::{Enum, Message};
const ACK_TIMEOUT: Duration = Duration::from_secs(1);
const READ_TIMEOUT: Duration = Duration::from_secs(30);
const PACKET_LEN: usize = 64;
-const BUFFER_LEN: usize = 1024;
const MESSAGE_TYPE_BUTTONREQUEST: u16 = 26;
const MESSAGE_TYPE_BUTTONACK: u16 = 27;
pub struct Client<C> {
- pub channel: C,
+ pub channel: Buffered<C>,
socket: UdpSocket,
emu_addr: SocketAddr,
}
@@ -28,6 +28,8 @@ where
C: CredentialStore,
{
pub fn open(emu_addr: SocketAddr, channel: ChannelOpen<C, B>) -> Self {
+ let mut channel = Buffered::new(channel);
+ channel.set_packet_len(PACKET_LEN);
Client {
channel,
socket: UdpSocket::bind("127.0.0.1:0").unwrap(),
@@ -39,7 +41,7 @@ where
impl<C: ChannelIO> Client<C> {
pub fn map<D>(self, func: impl FnOnce(C) -> D) -> Client<D> {
Client {
- channel: func(self.channel),
+ channel: self.channel.map(|c| Ok(func(c))).unwrap(),
socket: self.socket,
emu_addr: self.emu_addr,
}
@@ -50,14 +52,16 @@ impl<C: ChannelIO> Client<C> {
self.socket.send_to(buf, self.emu_addr).unwrap();
}
- fn recv_from(&mut self, buf: &mut [u8], timeout: Duration) -> Option<usize> {
+ fn recv_from(&mut self, timeout: Duration) -> Option<Vec<u8>> {
self.socket.set_read_timeout(Some(timeout)).unwrap();
- let res = self.socket.recv_from(buf);
+ let mut sockbuf = vec![0u8; PACKET_LEN];
+ let res = self.socket.recv_from(sockbuf.as_mut_slice());
match res {
Ok((reply_len, src_addr)) => {
assert_eq!(src_addr, self.emu_addr);
- log::trace!("< {}", hex::encode(&buf[..reply_len]));
- Some(reply_len)
+ sockbuf.truncate(reply_len);
+ log::trace!("< {}", hex::encode(&sockbuf));
+ Some(sockbuf)
}
Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
log::debug!("UDP receive timeout after {:?}.", timeout);
@@ -70,34 +74,25 @@ impl<C: ChannelIO> Client<C> {
}
}
- fn write_ack(&mut self, packet_buffer: &mut [u8]) {
- self.channel.packet_out(packet_buffer, &[]).unwrap()
+ fn write_ack(&mut self) -> Vec<u8> {
+ self.channel.packet_out().unwrap()
}
- fn read_ack(&mut self, packet_buffer: &[u8]) -> bool {
- let pir = self
- .channel
- .packet_in(packet_buffer, &[])
- .check_failed()
- .unwrap();
+ fn read_ack(&mut self, packet: &[u8]) -> bool {
+ let pir = self.channel.packet_in(packet).check_failed().unwrap();
assert!(!pir.got_transport_error());
pir.got_ack()
}
pub fn write(&mut self, sid: u8, message_type: u16, message: &[u8]) {
- let mut send_buffer = vec![0; message.len() + C::BUFFER_OVERHEAD];
- self.channel
- .message_in_from(sid, message_type, message, send_buffer.as_mut_slice())
- .unwrap();
+ self.channel.message_in(sid, message_type, message).unwrap();
let mut sockbuf = [0u8; PACKET_LEN];
let mut acked = false;
while !acked {
while self.channel.packet_out_ready() {
- self.channel
- .packet_out(&mut sockbuf, send_buffer.as_mut_slice())
- .unwrap();
- self.send_to(&sockbuf);
+ let packet = self.channel.packet_out().unwrap();
+ self.send_to(&packet);
sockbuf.fill(0);
}
// Only true if channel ID is not known, otherwise we need to wait for an ACK.
@@ -105,14 +100,14 @@ impl<C: ChannelIO> Client<C> {
break;
}
while !acked {
- match self.recv_from(&mut sockbuf, ACK_TIMEOUT) {
+ match self.recv_from(ACK_TIMEOUT) {
None => {
// timeout
self.channel.message_retransmit().unwrap();
break;
}
- Some(reply_len) => {
- acked = self.read_ack(&sockbuf[..reply_len]);
+ Some(packet) => {
+ acked = self.read_ack(&packet);
}
}
}
@@ -120,25 +115,23 @@ impl<C: ChannelIO> Client<C> {
}
pub fn read(&mut self) -> (u8, u16, Vec<u8>) {
- let mut recv_buf = vec![0u8; BUFFER_LEN];
- let mut sockbuf = [0u8; PACKET_LEN];
let mut result: Option<(u8, u16, Vec<u8>)> = None;
while result.is_none() {
let mut message_ready = false;
while !message_ready {
- let Some(reply_len) = self.recv_from(&mut sockbuf, READ_TIMEOUT) else {
+ let Some(packet) = self.recv_from(READ_TIMEOUT) else {
log::error!("Timed out waiting for response for {:?}.", READ_TIMEOUT);
panic!();
};
message_ready = self
.channel
- .packet_in(&sockbuf[..reply_len], recv_buf.as_mut_slice())
+ .packet_in(&packet)
.check_failed()
.unwrap()
.got_message();
}
- result = match self.channel.message_out(recv_buf.as_mut_slice()) {
- Ok((sid, message_type, message)) => Some((sid, message_type, Vec::from(message))),
+ result = match self.channel.message_out() {
+ Ok(r) => Some(r),
Err(Error::InvalidChecksum | Error::MalformedData) => {
log::error!("Received bad message, waiting for retransmission.");
continue;
@@ -150,9 +143,8 @@ impl<C: ChannelIO> Client<C> {
}
}
// Send ACK
- sockbuf.fill(0);
- self.write_ack(&mut sockbuf);
- self.send_to(&sockbuf);
+ let packet = self.write_ack();
+ self.send_to(&packet);
result.unwrap()
}
diff --git a/rust/trezor-thp/src/channel/buffered.rs b/rust/trezor-thp/src/channel/buffered.rs
new file mode 100644
index 00000000..6e4e9059
--- /dev/null
+++ b/rust/trezor-thp/src/channel/buffered.rs
@@ -0,0 +1,123 @@
+use crate::{
+ ChannelIO,
+ channel::{APP_HEADER_LEN, PacketInResult},
+ error::Result,
+};
+
+use std::ops::{Deref, DerefMut};
+use std::vec::Vec;
+
+const INITIAL_BUFFER_LEN: usize = 1024;
+const DEFAULT_PACKET_LEN: usize = 64;
+
+/// Wrapper for [`ChannelIO`] that deals with send/receive buffer using `std::vec::Vec`.
+/// It needs to know what size packets to produce - make sure to call [`Buffered::set_packet_len`].
+pub struct Buffered<C> {
+ channel: C,
+ packet_len: usize,
+ send_buffer: Vec<u8>,
+ receive_buffer: Vec<u8>,
+}
+
+impl<C: ChannelIO> Buffered<C> {
+ pub fn new(channel: C) -> Self {
+ Self {
+ channel,
+ packet_len: DEFAULT_PACKET_LEN,
+ send_buffer: Vec::new(),
+ receive_buffer: vec![0u8; INITIAL_BUFFER_LEN],
+ }
+ }
+
+ pub fn set_packet_len(&mut self, packet_len: usize) {
+ self.packet_len = packet_len;
+ }
+
+ pub fn packet_in(&mut self, packet_buffer: &[u8]) -> PacketInResult {
+ let res = self
+ .channel
+ .packet_in(packet_buffer, self.receive_buffer.as_mut_slice());
+ if let PacketInResult::EnlargeBuffer {
+ ack_received,
+ buffer_size,
+ } = res
+ {
+ log::debug!("Resizing receive buffer to {}.", buffer_size);
+ self.receive_buffer.resize(buffer_size.into(), 0u8);
+ return PacketInResult::Accepted {
+ ack_received,
+ message_ready: false,
+ pong: false,
+ };
+ }
+ res
+ }
+
+ pub fn packet_out(&mut self) -> Result<Vec<u8>> {
+ let mut channel_buffer = vec![0u8; self.packet_len];
+ self.channel
+ .packet_out(channel_buffer.as_mut_slice(), &self.send_buffer)?;
+ Ok(channel_buffer)
+ }
+
+ pub fn message_in(&mut self, session_id: u8, message_type: u16, message: &[u8]) -> Result<()> {
+ let mut send_buffer = vec![0; message.len() + C::BUFFER_OVERHEAD];
+ let res = self.channel.message_in_from(
+ session_id,
+ message_type,
+ message,
+ send_buffer.as_mut_slice(),
+ );
+ if res.is_ok() {
+ self.send_buffer = send_buffer;
+ }
+ res
+ }
+
+ pub fn message_out<'a>(&mut self) -> Result<(u8, u16, Vec<u8>)> {
+ let (session_id, message_type, message) = self
+ .channel
+ .message_out(self.receive_buffer.as_mut_slice())?;
+ let len = message.len();
+ self.receive_buffer
+ .copy_within(APP_HEADER_LEN..APP_HEADER_LEN + len, 0);
+ self.receive_buffer.truncate(len);
+ let res = core::mem::replace(&mut self.receive_buffer, vec![0u8; INITIAL_BUFFER_LEN]);
+ Ok((session_id, message_type, res))
+ }
+
+ pub fn message_retransmit(&mut self) -> Result<()> {
+ self.channel.message_retransmit()
+ }
+
+ pub fn map<D>(self, func: impl FnOnce(C) -> Result<D>) -> Result<Buffered<D>> {
+ Ok(Buffered {
+ channel: func(self.channel)?,
+ packet_len: self.packet_len,
+ send_buffer: self.send_buffer,
+ receive_buffer: self.receive_buffer,
+ })
+ }
+}
+
+impl<C: ChannelIO> Deref for Buffered<C> {
+ type Target = C;
+
+ fn deref(&self) -> &Self::Target {
+ &self.channel
+ }
+}
+
+impl<C: ChannelIO> DerefMut for Buffered<C> {
+ fn deref_mut(&mut self) -> &mut Self::Target {
+ &mut self.channel
+ }
+}
+
+pub trait ChannelExt: ChannelIO + Sized {
+ fn into_buffered(self) -> Buffered<Self> {
+ Buffered::new(self)
+ }
+}
+
+impl<C: ChannelIO> ChannelExt for C {}
diff --git a/rust/trezor-thp/src/channel/mod.rs b/rust/trezor-thp/src/channel/mod.rs
index abe5d4e2..e37e11f7 100644
--- a/rust/trezor-thp/src/channel/mod.rs
+++ b/rust/trezor-thp/src/channel/mod.rs
@@ -1,3 +1,5 @@
+#[cfg(feature = "use_std")]
+pub mod buffered;
pub mod host;
mod noise;
diff --git a/rust/trezor-thp/src/lib.rs b/rust/trezor-thp/src/lib.rs
index b9460e4c..d1ff02b4 100644
--- a/rust/trezor-thp/src/lib.rs
+++ b/rust/trezor-thp/src/lib.rs
@@ -1,6 +1,6 @@
#![doc = include_str!("../README.md")]
-#![no_std]
#![forbid(unsafe_code)]
+#![cfg_attr(not(feature = "use_std"), no_std)]
mod alternating_bit;
pub mod channel;
Why this scored 15/100
Community notes
Notes can correct, qualify, or add evidence to the AI analysis. Every note shown here has been validated by a human moderator.
The AI analysis stands alone for now. Submit a note if you can add evidence or important context.