From a49960b5b1a2077d4d8d272d93bcf1439236d960 Mon Sep 17 00:00:00 2001 From: ivila <390810839@qq.com> Date: Wed, 19 Aug 2026 14:18:17 +0800 Subject: [PATCH] optee-utee: safety of memref shared-memory access --- crates/optee-utee/src/parameter/memref.rs | 171 +++++++++++++++--- crates/optee-utee/src/parameter/mod.rs | 5 +- examples/ta/acipher-rs/src/main.rs | 4 +- examples/ta/aes-rs/src/main.rs | 12 +- examples/ta/authentication-rs/src/main.rs | 31 ++-- examples/ta/big_int-rs/src/main.rs | 4 +- examples/ta/diffie_hellman-rs/src/main.rs | 12 +- examples/ta/digest-rs/src/main.rs | 6 +- examples/ta/hotp-rs/src/main.rs | 2 +- .../message_passing_interface-rs/src/main.rs | 2 +- examples/ta/mnist-rs/inference/src/main.rs | 6 +- examples/ta/mnist-rs/train/src/main.rs | 19 +- examples/ta/random-rs/src/main.rs | 4 +- examples/ta/secure_storage-rs/src/main.rs | 10 +- .../ta/signature_verification-rs/src/main.rs | 16 +- examples/ta/supp_plugin-rs/src/main.rs | 18 +- examples/ta/tcp_client-rs/src/main.rs | 4 +- examples/ta/tls_server-rs/src/main.rs | 6 +- examples/ta/udp_socket-rs/src/main.rs | 2 +- projects/web3/eth_wallet/ta/src/main.rs | 4 +- 20 files changed, 237 insertions(+), 101 deletions(-) diff --git a/crates/optee-utee/src/parameter/memref.rs b/crates/optee-utee/src/parameter/memref.rs index a34cc460..44a2422c 100644 --- a/crates/optee-utee/src/parameter/memref.rs +++ b/crates/optee-utee/src/parameter/memref.rs @@ -37,6 +37,17 @@ //! | `ParameterMemrefInput` | ✓ | ✗ | //! | `ParameterMemrefOutput` | ✗ | ✓ | //! | `ParameterMemrefInout` | ✓ | ✓ | +//! +//! # Shared-memory safety +//! +//! The buffers are mapped from Normal World, which may access them concurrently. +//! Consequently, `get_buffer` and `get_buffer_mut` are unsafe: callers may use +//! them only when their application guarantees that the REE will not access the +//! memory for the returned reference's lifetime. Callers are also responsible +//! for preventing REE-controlled time-of-check-to-time-of-use (TOCTOU) attacks; +//! data must not be validated through a shared slice and then fetched from it +//! again for use. Use `read_to_vec` and validate/use the resulting TA-owned copy, +//! or use `write_at`/`set_output`, when those guarantees are unavailable. use super::{FromRawParameter, ParamType, RawParamType, check_type_is}; use crate::{ErrorKind, Result, raw::TEE_Param}; @@ -51,7 +62,38 @@ pub trait ParameterMemrefRead { /// supplied by the host. For `ParameterMemrefInout` the length is the /// full buffer capacity, not the number of valid bytes (which may have /// been updated by a prior write). - fn get_buffer(&self) -> &[u8]; + /// # Safety + /// + /// This slice points directly into Normal-World shared memory. The caller + /// must ensure the REE cannot mutate that memory for the entire lifetime of + /// the returned slice. The caller must also prevent TOCTOU attacks: do not + /// validate data through this slice and later fetch it again from shared + /// memory for use. If either guarantee cannot be met, use + /// [`Self::read_to_vec`] and perform both validation and use on the same + /// TA-owned copy. + unsafe fn get_buffer(&self) -> &[u8] { + unsafe { core::slice::from_raw_parts(self.buffer_ptr(), self.buffer_len()) } + } + + /// Copies the shared buffer into TA-owned memory. + fn read_to_vec(&self) -> alloc::vec::Vec { + let len = self.buffer_len(); + let mut copy = alloc::vec![0; len]; + if len != 0 { + unsafe { + crate::raw::TEE_MemMove(copy.as_mut_ptr().cast(), self.buffer_ptr().cast(), len); + } + } + copy + } + + /// Returns the start of the shared input buffer. + #[doc(hidden)] + fn buffer_ptr(&self) -> *const u8; + + /// Returns the length of the shared input buffer. + #[doc(hidden)] + fn buffer_len(&self) -> usize; } /// Write access to a memory-reference parameter's buffer. @@ -64,7 +106,20 @@ pub trait ParameterMemrefWrite { /// [`ParameterMemrefWrite::set_updated_size`] to report how many bytes were /// produced. Otherwise the client application may observe an incorrect /// output size. - fn get_buffer_mut(&mut self) -> &mut [u8]; + /// # Safety + /// + /// This slice points directly into Normal-World shared memory. The caller + /// must ensure the REE does not read or write that memory for the entire + /// lifetime of the returned mutable slice. A TA that only writes output + /// need not protect the resulting contents from the REE. However, if the TA + /// also reads, validates, or makes decisions from this slice, it must treat + /// those bytes like input shared memory and prevent TOCTOU attacks. Prefer + /// [`Self::write_at`] for write-only access; copy data into TA-owned memory + /// before validating or otherwise relying on bytes read from this slice. + unsafe fn get_buffer_mut(&mut self) -> &mut [u8] { + let capacity = self.get_capacity(); + unsafe { core::slice::from_raw_parts_mut(self.buffer_ptr(), capacity) } + } /// Returns the maximum allowed buffer size (capacity). fn get_capacity(&self) -> usize; @@ -92,16 +147,29 @@ pub trait ParameterMemrefWrite { /// the buffer capacity. fn write_at>(&mut self, offset: usize, data: T) -> Result<()> { let input = data.as_ref(); - let new_size = offset + input.len(); + let new_size = offset + .checked_add(input.len()) + .ok_or(ErrorKind::ShortBuffer)?; if new_size > self.get_capacity() { return Err(ErrorKind::ShortBuffer.into()); } - let output = self.get_buffer_mut(); - output[offset..new_size].copy_from_slice(input); + if !input.is_empty() { + unsafe { + crate::raw::TEE_MemMove( + self.buffer_ptr().add(offset).cast(), + input.as_ptr().cast(), + input.len(), + ); + } + } unsafe { self.set_updated_size_unchecked(new_size) }; Ok(()) } + /// Returns the start of the shared output buffer. + #[doc(hidden)] + fn buffer_ptr(&mut self) -> *mut u8; + /// Directly sets the updated size without bounds checking. /// /// # Safety @@ -140,12 +208,18 @@ pub struct ParameterMemrefOutput<'a> { impl<'a> FromRawParameter<'a> for ParameterMemrefInput<'a> { unsafe fn from_raw(raw_type: RawParamType, raw_param: &'a mut TEE_Param) -> Result { check_type_is(raw_type, ParamType::MemrefInput)?; + if unsafe { raw_param.memref.buffer }.is_null() { + return Err(ErrorKind::BadParameters.into()); + } Ok(Self(raw_param)) } } impl<'a> FromRawParameter<'a> for ParameterMemrefInout<'a> { unsafe fn from_raw(raw_type: RawParamType, raw_param: &'a mut TEE_Param) -> Result { check_type_is(raw_type, ParamType::MemrefInout)?; + if unsafe { raw_param.memref.buffer }.is_null() { + return Err(ErrorKind::BadParameters.into()); + } Ok(Self { capacity: unsafe { raw_param.memref.size }, raw_param, @@ -155,6 +229,9 @@ impl<'a> FromRawParameter<'a> for ParameterMemrefInout<'a> { impl<'a> FromRawParameter<'a> for ParameterMemrefOutput<'a> { unsafe fn from_raw(raw_type: RawParamType, raw_param: &'a mut TEE_Param) -> Result { check_type_is(raw_type, ParamType::MemrefOutput)?; + if unsafe { raw_param.memref.buffer }.is_null() { + return Err(ErrorKind::BadParameters.into()); + } Ok(Self { capacity: unsafe { raw_param.memref.size }, raw_param, @@ -163,45 +240,95 @@ impl<'a> FromRawParameter<'a> for ParameterMemrefOutput<'a> { } impl<'a> ParameterMemrefWrite for ParameterMemrefInout<'a> { - fn get_buffer_mut(&mut self) -> &mut [u8] { - unsafe { - core::slice::from_raw_parts_mut(self.raw_param.memref.buffer as *mut u8, self.capacity) - } - } fn get_capacity(&self) -> usize { self.capacity } + fn buffer_ptr(&mut self) -> *mut u8 { + unsafe { self.raw_param.memref.buffer as *mut u8 } + } unsafe fn set_updated_size_unchecked(&mut self, size: usize) { self.raw_param.memref.size = size; } } impl<'a> ParameterMemrefWrite for ParameterMemrefOutput<'a> { - fn get_buffer_mut(&mut self) -> &mut [u8] { - unsafe { - core::slice::from_raw_parts_mut(self.raw_param.memref.buffer as *mut u8, self.capacity) - } - } fn get_capacity(&self) -> usize { self.capacity } + fn buffer_ptr(&mut self) -> *mut u8 { + unsafe { self.raw_param.memref.buffer as *mut u8 } + } unsafe fn set_updated_size_unchecked(&mut self, size: usize) { self.raw_param.memref.size = size; } } impl<'a> ParameterMemrefRead for ParameterMemrefInout<'a> { - fn get_buffer(&self) -> &[u8] { - unsafe { - core::slice::from_raw_parts(self.raw_param.memref.buffer as *const u8, self.capacity) - } + fn buffer_ptr(&self) -> *const u8 { + unsafe { self.raw_param.memref.buffer as *const u8 } + } + fn buffer_len(&self) -> usize { + self.capacity } } impl<'a> ParameterMemrefRead for ParameterMemrefInput<'a> { - fn get_buffer(&self) -> &[u8] { - unsafe { - core::slice::from_raw_parts(self.0.memref.buffer as *const u8, self.0.memref.size) + fn buffer_ptr(&self) -> *const u8 { + unsafe { self.0.memref.buffer as *const u8 } + } + fn buffer_len(&self) -> usize { + unsafe { self.0.memref.size } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::raw; + + fn memref(buffer: *mut u8, size: usize) -> TEE_Param { + TEE_Param { + memref: raw::Memref { + buffer: buffer.cast(), + size, + }, + } + } + + #[test] + fn typed_wrappers_reject_null_buffers() { + let cases = [ + (raw::TEE_PARAM_TYPE_MEMREF_INPUT, ParamType::MemrefInput), + (raw::TEE_PARAM_TYPE_MEMREF_INOUT, ParamType::MemrefInout), + (raw::TEE_PARAM_TYPE_MEMREF_OUTPUT, ParamType::MemrefOutput), + ]; + + for (raw_type, param_type) in cases { + let mut raw_param = memref(core::ptr::null_mut(), 0); + let error = match param_type { + ParamType::MemrefInput => unsafe { + ParameterMemrefInput::from_raw(raw_type, &mut raw_param).map(|_| ()) + }, + ParamType::MemrefInout => unsafe { + ParameterMemrefInout::from_raw(raw_type, &mut raw_param).map(|_| ()) + }, + ParamType::MemrefOutput => unsafe { + ParameterMemrefOutput::from_raw(raw_type, &mut raw_param).map(|_| ()) + }, + _ => unreachable!(), + } + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::BadParameters); + } + } + + #[test] + fn zero_length_non_null_buffer_is_valid() { + let mut raw_param = memref(core::ptr::NonNull::::dangling().as_ptr(), 0); + let input = unsafe { + ParameterMemrefInput::from_raw(raw::TEE_PARAM_TYPE_MEMREF_INPUT, &mut raw_param) } + .unwrap(); + assert!(input.read_to_vec().is_empty()); } } diff --git a/crates/optee-utee/src/parameter/mod.rs b/crates/optee-utee/src/parameter/mod.rs index 79422f0b..ffce1cc4 100644 --- a/crates/optee-utee/src/parameter/mod.rs +++ b/crates/optee-utee/src/parameter/mod.rs @@ -176,7 +176,10 @@ fn check_type_is(raw_type: RawParamType, exp_type: ParamType) -> Result<()> { /// match param { /// ParameterAny::None => { /* no parameter */ } /// ParameterAny::MemrefInput(p) => { -/// let data: &[u8] = p.get_buffer(); +/// // SAFETY: this application guarantees the REE will not mutate the +/// // buffer while `data` is in use and does not fetch it again after +/// // validation, preventing a TOCTOU mismatch. +/// let data: &[u8] = unsafe { p.get_buffer() }; /// // process data ... /// } /// ParameterAny::ValueInput(p) => { diff --git a/examples/ta/acipher-rs/src/main.rs b/examples/ta/acipher-rs/src/main.rs index a4c5972a..5eb37689 100644 --- a/examples/ta/acipher-rs/src/main.rs +++ b/examples/ta/acipher-rs/src/main.rs @@ -83,7 +83,7 @@ fn encrypt(rsa: &mut RsaCipher, (p0, p1, _, _): &mut ParametersAny<'_>) -> Resul key_info.object_size(), )?; cipher.set_key(&rsa.key)?; - let cipher_text = cipher.encrypt(&[], p0.get_buffer())?; + let cipher_text = cipher.encrypt(&[], unsafe { p0.get_buffer() })?; p1.set_output(cipher_text)?; Ok(()) } @@ -97,7 +97,7 @@ fn decrypt(rsa: &mut RsaCipher, (p0, p1, _, _): &mut ParametersAny<'_>) -> Resul key_info.object_size(), )?; cipher.set_key(&rsa.key)?; - let plain_text = cipher.decrypt(&[], p0.get_buffer())?; + let plain_text = cipher.decrypt(&[], unsafe { p0.get_buffer() })?; p1.set_output(plain_text) } diff --git a/examples/ta/aes-rs/src/main.rs b/examples/ta/aes-rs/src/main.rs index 926c0390..7aa4c017 100644 --- a/examples/ta/aes-rs/src/main.rs +++ b/examples/ta/aes-rs/src/main.rs @@ -135,7 +135,7 @@ pub fn alloc_resources(aes: &mut AesCipher, (p0, p1, p2, _): &mut ParametersAny< } pub fn set_aes_key(aes: &mut AesCipher, (p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { - let key = p0.as_memref_input()?.get_buffer(); + let key = unsafe { p0.as_memref_input()?.get_buffer() }; if key.len() != aes.key_size { trace_println!("[+] Get wrong key size !\n"); @@ -152,7 +152,7 @@ pub fn set_aes_key(aes: &mut AesCipher, (p0, _, _, _): &mut ParametersAny<'_>) - } pub fn reset_aes_iv(aes: &mut AesCipher, (p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { - let iv = p0.as_memref_input()?.get_buffer(); + let iv = unsafe { p0.as_memref_input()?.get_buffer() }; aes.cipher.init(iv); @@ -163,15 +163,15 @@ pub fn reset_aes_iv(aes: &mut AesCipher, (p0, _, _, _): &mut ParametersAny<'_>) pub fn cipher_buffer(aes: &mut AesCipher, (p0, p1, _, _): &mut ParametersAny<'_>) -> Result<()> { let (input, output) = (p0.as_memref_input()?, p1.as_memref_output()?); - if output.get_capacity() < input.get_buffer().len() { + if output.get_capacity() < unsafe { input.get_buffer() }.len() { return Err(ErrorKind::BadParameters.into()); } trace_println!("[+] TA tries to update ciphers!"); - let tmp_size = aes - .cipher - .update(input.get_buffer(), output.get_buffer_mut())?; + let tmp_size = aes.cipher.update(unsafe { input.get_buffer() }, unsafe { + output.get_buffer_mut() + })?; output.set_updated_size(tmp_size)?; Ok(()) } diff --git a/examples/ta/authentication-rs/src/main.rs b/examples/ta/authentication-rs/src/main.rs index 5a4e093c..dac60b9a 100644 --- a/examples/ta/authentication-rs/src/main.rs +++ b/examples/ta/authentication-rs/src/main.rs @@ -21,10 +21,10 @@ extern crate alloc; use optee_utee::prelude::*; -use optee_utee::{AlgorithmId, OperationMode, AE}; +use optee_utee::{AE, AlgorithmId, OperationMode}; use optee_utee::{AttributeId, AttributeMemref, TransientObject, TransientObjectType}; use optee_utee::{ErrorKind, Result}; -use proto::authentication::{Command, Mode, AAD_LEN, BUFFER_SIZE, KEY_SIZE, TAG_LEN}; +use proto::authentication::{AAD_LEN, BUFFER_SIZE, Command, KEY_SIZE, Mode, TAG_LEN}; pub const PAYLOAD_NUMBER: usize = 2; @@ -90,9 +90,9 @@ pub fn prepare(ae: &mut AEOp, (p0, p1, p2, p3): &mut ParametersAny<'_>) -> Resul Mode::Decrypt => OperationMode::Decrypt, _ => OperationMode::IllegalValue, }; - let nonce = p1.as_memref_input()?.get_buffer(); - let key = p2.as_memref_input()?.get_buffer(); - let aad = p3.as_memref_input()?.get_buffer(); + let nonce = unsafe { p1.as_memref_input()?.get_buffer() }; + let key = unsafe { p2.as_memref_input()?.get_buffer() }; + let aad = unsafe { p3.as_memref_input()?.get_buffer() }; ae.op = AE::allocate(AlgorithmId::AesCcm, mode, KEY_SIZE * 8)?; @@ -108,7 +108,9 @@ pub fn prepare(ae: &mut AEOp, (p0, p1, p2, p3): &mut ParametersAny<'_>) -> Resul pub fn update(digest: &mut AEOp, (p0, p1, _, _): &mut ParametersAny<'_>) -> Result<()> { let (p0, p1) = (p0.as_memref_input()?, p1.as_memref_output()?); - let size = digest.op.update(p0.get_buffer(), p1.get_buffer_mut())?; + let size = digest + .op + .update(unsafe { p0.get_buffer() }, unsafe { p1.get_buffer_mut() })?; p1.set_updated_size(size)?; Ok(()) } @@ -120,10 +122,11 @@ pub fn encrypt_final(digest: &mut AEOp, (p0, p1, p2, _): &mut ParametersAny<'_>) p2.as_memref_output()?, ); - let (ciph_len, tag_len) = - digest - .op - .encrypt_final(p0.get_buffer(), p1.get_buffer_mut(), p2.get_buffer_mut())?; + let (ciph_len, tag_len) = digest.op.encrypt_final( + unsafe { p0.get_buffer() }, + unsafe { p1.get_buffer_mut() }, + unsafe { p2.get_buffer_mut() }, + )?; p1.set_updated_size(ciph_len)?; p2.set_updated_size(tag_len)?; Ok(()) @@ -136,9 +139,11 @@ pub fn decrypt_final(digest: &mut AEOp, (p0, p1, p2, _): &mut ParametersAny<'_>) p2.as_memref_input()?, ); - let len = digest - .op - .decrypt_final(p0.get_buffer(), p1.get_buffer_mut(), p2.get_buffer())?; + let len = digest.op.decrypt_final( + unsafe { p0.get_buffer() }, + unsafe { p1.get_buffer_mut() }, + unsafe { p2.get_buffer() }, + )?; p1.set_updated_size(len)?; Ok(()) } diff --git a/examples/ta/big_int-rs/src/main.rs b/examples/ta/big_int-rs/src/main.rs index 8c46a23d..940815cf 100644 --- a/examples/ta/big_int-rs/src/main.rs +++ b/examples/ta/big_int-rs/src/main.rs @@ -18,8 +18,8 @@ #![cfg_attr(not(feature = "std"), no_std)] #![no_main] -use optee_utee::prelude::*; use optee_utee::BigInt; +use optee_utee::prelude::*; use optee_utee::{ErrorKind, Result}; use proto::big_int::Command; @@ -103,7 +103,7 @@ fn invoke_command(cmd_id: u32, (p0, p1, _, _): &mut ParametersAny<'_>) -> Result let mut n0 = BigInt::new(64); let mut n1 = BigInt::new(2); - n0.convert_from_octet_string(n0_buffer.get_buffer(), 0)?; + n0.convert_from_octet_string(unsafe { n0_buffer.get_buffer() }, 0)?; n1.convert_from_s32(n1_value.get_a() as i32); match Command::from(cmd_id) { diff --git a/examples/ta/diffie_hellman-rs/src/main.rs b/examples/ta/diffie_hellman-rs/src/main.rs index fda3fe1a..bbb453ff 100644 --- a/examples/ta/diffie_hellman-rs/src/main.rs +++ b/examples/ta/diffie_hellman-rs/src/main.rs @@ -70,7 +70,7 @@ fn generate_key(dh: &mut DiffieHellman, (p0, p1, p2, p3): &mut ParametersAny<'_> p3.as_memref_output()?, ); // Extract prime and base from parameters - let prime_base_vec = p0.get_buffer(); + let prime_base_vec = unsafe { p0.get_buffer() }; let prime_slice = &prime_base_vec[..KEY_SIZE / 8]; let base_slice = &prime_base_vec[KEY_SIZE / 8..]; @@ -85,7 +85,7 @@ fn generate_key(dh: &mut DiffieHellman, (p0, p1, p2, p3): &mut ParametersAny<'_> { let key_size = dh .key - .ref_attribute(AttributeId::DhPublicValue, p2.get_buffer_mut())?; + .ref_attribute(AttributeId::DhPublicValue, unsafe { p2.get_buffer_mut() })?; p2.set_updated_size(key_size)?; p1.set_a(key_size as u32); } @@ -93,7 +93,7 @@ fn generate_key(dh: &mut DiffieHellman, (p0, p1, p2, p3): &mut ParametersAny<'_> { let key_size = dh .key - .ref_attribute(AttributeId::DhPrivateValue, p3.get_buffer_mut())?; + .ref_attribute(AttributeId::DhPrivateValue, unsafe { p3.get_buffer_mut() })?; p3.set_updated_size(key_size)?; p1.set_b(key_size as u32); } @@ -106,13 +106,15 @@ fn derive_key(dh: &mut DiffieHellman, (p0, p1, p2, _): &mut ParametersAny<'_>) - p1.as_memref_output()?, p2.as_value_output()?, ); - let received_public = AttributeMemref::from_ref(AttributeId::DhPublicValue, p0.get_buffer()); + let received_public = + AttributeMemref::from_ref(AttributeId::DhPublicValue, unsafe { p0.get_buffer() }); let mut operation = DeriveKey::allocate(AlgorithmId::DhDeriveSharedSecret, KEY_SIZE)?; operation.set_key(&dh.key)?; let mut derived_key = TransientObject::allocate(TransientObjectType::GenericSecret, KEY_SIZE)?; operation.derive(&[received_public.into()], &mut derived_key); - let key_size = derived_key.ref_attribute(AttributeId::SecretValue, p1.get_buffer_mut())?; + let key_size = + derived_key.ref_attribute(AttributeId::SecretValue, unsafe { p1.get_buffer_mut() })?; p1.set_updated_size(key_size)?; p2.set_a(key_size as u32); Ok(()) diff --git a/examples/ta/digest-rs/src/main.rs b/examples/ta/digest-rs/src/main.rs index ca55bc1a..064e114f 100644 --- a/examples/ta/digest-rs/src/main.rs +++ b/examples/ta/digest-rs/src/main.rs @@ -77,7 +77,7 @@ fn invoke_command( } pub fn update(digest: &mut DigestOp, (p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { - let buffer = p0.as_memref_input()?.get_buffer(); + let buffer = unsafe { p0.as_memref_input()?.get_buffer() }; digest.op.update(buffer); Ok(()) } @@ -88,8 +88,8 @@ pub fn do_final(digest: &mut DigestOp, (p0, p1, p2, _): &mut ParametersAny<'_>) p1.as_memref_output()?, p2.as_value_output()?, ); - let input = p0.get_buffer(); - let length = digest.op.do_final(input, p1.get_buffer_mut())?; + let input = unsafe { p0.get_buffer() }; + let length = digest.op.do_final(input, unsafe { p1.get_buffer_mut() })?; p2.set_a(length as u32); p1.set_updated_size(length)?; Ok(()) diff --git a/examples/ta/hotp-rs/src/main.rs b/examples/ta/hotp-rs/src/main.rs index 78a5a865..5a49e196 100644 --- a/examples/ta/hotp-rs/src/main.rs +++ b/examples/ta/hotp-rs/src/main.rs @@ -87,7 +87,7 @@ pub fn register_shared_key( hotp: &mut HmacOtp, (p0, _, _, _): &mut ParametersAny<'_>, ) -> Result<()> { - let buffer = p0.as_memref_input()?.get_buffer(); + let buffer = unsafe { p0.as_memref_input()?.get_buffer() }; hotp.key_len = buffer.len(); hotp.key[..hotp.key_len].clone_from_slice(buffer); Ok(()) diff --git a/examples/ta/message_passing_interface-rs/src/main.rs b/examples/ta/message_passing_interface-rs/src/main.rs index 9da52cf4..a9a6cb42 100644 --- a/examples/ta/message_passing_interface-rs/src/main.rs +++ b/examples/ta/message_passing_interface-rs/src/main.rs @@ -80,7 +80,7 @@ fn invoke_command( ) -> Result<()> { trace_println!("[+] TA invoke command"); let input: proto::message_passing_interface::EnclaveInput = - serde_json::from_slice(p0.get_buffer()).map_err(|e| { + serde_json::from_slice(unsafe { p0.get_buffer() }).map_err(|e| { trace_println!("Failed to deserialize input: {}", e); ErrorKind::BadFormat })?; diff --git a/examples/ta/mnist-rs/inference/src/main.rs b/examples/ta/mnist-rs/inference/src/main.rs index 15537592..92a248d7 100644 --- a/examples/ta/mnist-rs/inference/src/main.rs +++ b/examples/ta/mnist-rs/inference/src/main.rs @@ -20,7 +20,7 @@ extern crate alloc; use burn::{ - backend::{ndarray::NdArrayDevice, NdArray}, + backend::{NdArray, ndarray::NdArrayDevice}, tensor::cast::ToElement, }; @@ -51,7 +51,7 @@ fn open_session( ) -> Result<()> { let mut model = MODEL.lock(); model.replace( - Model::import(&DEVICE, p0.get_buffer().to_vec()).map_err(|err| { + Model::import(&DEVICE, unsafe { p0.get_buffer() }.to_vec()).map_err(|err| { trace_println!("import failed: {:?}", err); ErrorKind::BadParameters })?, @@ -81,7 +81,7 @@ fn invoke_command( ), ) -> Result<()> { trace_println!("[+] TA invoke command"); - let images: &[Image] = bytemuck::cast_slice(p0.get_buffer()); + let images: &[Image] = bytemuck::cast_slice(unsafe { p0.get_buffer() }); let input = NoStdModel::images_to_tensors(&DEVICE, images); let output = MODEL diff --git a/examples/ta/mnist-rs/train/src/main.rs b/examples/ta/mnist-rs/train/src/main.rs index 7039c6ad..89738243 100644 --- a/examples/ta/mnist-rs/train/src/main.rs +++ b/examples/ta/mnist-rs/train/src/main.rs @@ -19,7 +19,7 @@ #![no_main] extern crate alloc; -use burn::backend::{ndarray::NdArrayDevice, Autodiff, NdArray}; +use burn::backend::{Autodiff, NdArray, ndarray::NdArrayDevice}; use optee_utee::prelude::*; use optee_utee::{ErrorKind, Result}; use proto::mnist::train::Command; @@ -47,10 +47,11 @@ fn open_session( ParameterNone, ), ) -> Result<()> { - let learning_rate = f64::from_le_bytes(p0.get_buffer().try_into().map_err(|err| { - trace_println!("bad parameter {:?}", err); - ErrorKind::BadParameters - })?); + let learning_rate = + f64::from_le_bytes(unsafe { p0.get_buffer() }.try_into().map_err(|err| { + trace_println!("bad parameter {:?}", err); + ErrorKind::BadParameters + })?); trace_println!("Initialize with learning_rate: {}", learning_rate); let mut trainer = TRAINER.lock(); @@ -73,8 +74,8 @@ fn destroy() { fn invoke_command(cmd_id: u32, (p0, p1, p2, _): &mut ParametersAny<'_>) -> Result<()> { match Command::try_from(cmd_id) { Ok(Command::Train) => { - let images = p0.as_memref_input()?.get_buffer(); - let labels = p1.as_memref_input()?.get_buffer(); + let images = unsafe { p0.as_memref_input()?.get_buffer() }; + let labels = unsafe { p1.as_memref_input()?.get_buffer() }; let mut trainer = TRAINER.lock(); let result = trainer @@ -88,8 +89,8 @@ fn invoke_command(cmd_id: u32, (p0, p1, p2, _): &mut ParametersAny<'_>) -> Resul p2.as_memref_output()?.set_output(bytes) } Ok(Command::Valid) => { - let images = p0.as_memref_input()?.get_buffer(); - let labels = p1.as_memref_input()?.get_buffer(); + let images = unsafe { p0.as_memref_input()?.get_buffer() }; + let labels = unsafe { p1.as_memref_input()?.get_buffer() }; let trainer = TRAINER.lock(); let result = trainer diff --git a/examples/ta/random-rs/src/main.rs b/examples/ta/random-rs/src/main.rs index 06b06c14..2768ee3b 100644 --- a/examples/ta/random-rs/src/main.rs +++ b/examples/ta/random-rs/src/main.rs @@ -20,8 +20,8 @@ extern crate alloc; -use optee_utee::prelude::*; use optee_utee::Random; +use optee_utee::prelude::*; use optee_utee::{ErrorKind, Result}; use proto::random::Command; @@ -51,7 +51,7 @@ fn destroy() { pub fn random_number_generate((p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { let p0 = p0.as_memref_output()?; - Random::generate(p0.get_buffer_mut()); + Random::generate(unsafe { p0.get_buffer_mut() }); p0.set_updated_size(p0.get_capacity())?; Ok(()) diff --git a/examples/ta/secure_storage-rs/src/main.rs b/examples/ta/secure_storage-rs/src/main.rs index 409bfb54..4a23eccb 100644 --- a/examples/ta/secure_storage-rs/src/main.rs +++ b/examples/ta/secure_storage-rs/src/main.rs @@ -60,7 +60,7 @@ fn invoke_command(cmd_id: u32, params: &mut ParametersAny<'_>) -> Result<()> { pub fn delete_object((p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { // use to_vec to copy into tee memory - let obj_id = p0.as_memref_input()?.get_buffer().to_vec(); + let obj_id = unsafe { p0.as_memref_input()?.get_buffer() }.to_vec(); match PersistentObject::open( ObjectStorageConstants::Private, @@ -78,8 +78,8 @@ pub fn delete_object((p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { pub fn create_raw_object((p0, p1, _, _): &mut ParametersAny<'_>) -> Result<()> { // use to_vec to copy into tee memory - let obj_id = p0.as_memref_input()?.get_buffer().to_vec(); - let data_buffer = p1.as_memref_input()?.get_buffer().to_vec(); + let obj_id = unsafe { p0.as_memref_input()?.get_buffer() }.to_vec(); + let data_buffer = unsafe { p1.as_memref_input()?.get_buffer() }.to_vec(); let obj_data_flag = DataFlag::ACCESS_READ | DataFlag::ACCESS_WRITE @@ -106,7 +106,7 @@ pub fn create_raw_object((p0, p1, _, _): &mut ParametersAny<'_>) -> Result<()> { pub fn read_raw_object((p0, p1, _, _): &mut ParametersAny<'_>) -> Result<()> { // use to_vec to copy into tee memory - let obj_id = p0.as_memref_input()?.get_buffer().to_vec(); + let obj_id = unsafe { p0.as_memref_input()?.get_buffer() }.to_vec(); let p1 = p1.as_memref_output()?; let mut object = PersistentObject::open( @@ -116,7 +116,7 @@ pub fn read_raw_object((p0, p1, _, _): &mut ParametersAny<'_>) -> Result<()> { )?; let obj_info = object.info()?; - let read_bytes = object.read(p1.get_buffer_mut())?; + let read_bytes = object.read(unsafe { p1.get_buffer_mut() })?; if read_bytes != obj_info.data_size() as u32 { return Err(ErrorKind::ExcessData.into()); } diff --git a/examples/ta/signature_verification-rs/src/main.rs b/examples/ta/signature_verification-rs/src/main.rs index a086ae05..d48250ee 100644 --- a/examples/ta/signature_verification-rs/src/main.rs +++ b/examples/ta/signature_verification-rs/src/main.rs @@ -65,7 +65,7 @@ fn sign((p0, p1, p2, _): &mut ParametersAny<'_>) -> Result<()> { let p0 = p0.as_memref_input()?; let p1 = p1.as_memref_output()?; let p2 = p2.as_memref_output()?; - let message = p0.get_buffer(); + let message = unsafe { p0.get_buffer() }; trace_println!("[+] message: {:?}", message); let rsa_key = TransientObject::allocate(TransientObjectType::RsaKeypair, 2048_usize)?; @@ -73,7 +73,7 @@ fn sign((p0, p1, p2, _): &mut ParametersAny<'_>) -> Result<()> { rsa_key.generate_key(2048_usize, &[])?; { - let buffer = p1.get_buffer_mut(); + let buffer = unsafe { p1.get_buffer_mut() }; let modulus_len = rsa_key.ref_attribute(AttributeId::RsaModulus, buffer)?; let exp_len = rsa_key.ref_attribute(AttributeId::RsaPublicExponent, &mut buffer[modulus_len..])?; @@ -94,7 +94,7 @@ fn sign((p0, p1, p2, _): &mut ParametersAny<'_>) -> Result<()> { )?; rsa.set_key(&rsa_key)?; - let len = rsa.sign_digest(&[], &hash, p2.get_buffer_mut())?; + let len = rsa.sign_digest(&[], &hash, unsafe { p2.get_buffer_mut() })?; p2.set_updated_size(len)?; Ok(()) } @@ -104,13 +104,13 @@ fn verify((p0, p1, p2, _): &mut ParametersAny<'_>) -> Result<()> { let p1 = p1.as_memref_input()?; let p2 = p2.as_memref_input()?; - let message = p0.get_buffer(); + let message = unsafe { p0.get_buffer() }; let mut pub_key_mod = vec![0u8; 256]; - let mut pub_key_exp = vec![0u8; p1.get_buffer().len() - 256]; - let signature = p2.get_buffer(); + let mut pub_key_exp = vec![0u8; unsafe { p1.get_buffer() }.len() - 256]; + let signature = unsafe { p2.get_buffer() }; - pub_key_mod.copy_from_slice(&p1.get_buffer()[..256]); - pub_key_exp.copy_from_slice(&p1.get_buffer()[256..]); + pub_key_mod.copy_from_slice(&unsafe { p1.get_buffer() }[..256]); + pub_key_exp.copy_from_slice(&unsafe { p1.get_buffer() }[256..]); trace_println!("[+] message: {:?}", &message); trace_println!("[+] public_key_mod: {:?}", &pub_key_mod); diff --git a/examples/ta/supp_plugin-rs/src/main.rs b/examples/ta/supp_plugin-rs/src/main.rs index 49936d1f..4fd88a21 100644 --- a/examples/ta/supp_plugin-rs/src/main.rs +++ b/examples/ta/supp_plugin-rs/src/main.rs @@ -20,10 +20,10 @@ extern crate alloc; -use optee_utee::prelude::*; use optee_utee::LoadablePlugin; +use optee_utee::prelude::*; use optee_utee::{ErrorKind, Result, Uuid}; -use proto::supp_plugin::{Command, PluginCommand, PLUGIN_SUBCMD_NULL, PLUGIN_UUID}; +use proto::supp_plugin::{Command, PLUGIN_SUBCMD_NULL, PLUGIN_UUID, PluginCommand}; #[ta_create] fn create() -> Result<()> { @@ -51,20 +51,18 @@ fn destroy() { fn invoke_command(cmd_id: u32, (p0, _, _, _): &mut ParametersAny<'_>) -> Result<()> { trace_println!("[+] TA invoke command"); let p0 = p0.as_memref_input()?; - trace_println!( - "[+] TA received value {:?} then send to plugin", + trace_println!("[+] TA received value {:?} then send to plugin", unsafe { p0.get_buffer() - ); + }); let uuid = Uuid::parse_str(PLUGIN_UUID)?; match Command::from(cmd_id) { Command::Ping => { let plugin = LoadablePlugin::new(&uuid); - let outbuf = plugin.invoke( - PluginCommand::Print as u32, - PLUGIN_SUBCMD_NULL, - p0.get_buffer(), - )?; + let outbuf = + plugin.invoke(PluginCommand::Print as u32, PLUGIN_SUBCMD_NULL, unsafe { + p0.get_buffer() + })?; trace_println!( "[+] TA received out value {:?} outlen {:?}", diff --git a/examples/ta/tcp_client-rs/src/main.rs b/examples/ta/tcp_client-rs/src/main.rs index 996271ae..7992de86 100644 --- a/examples/ta/tcp_client-rs/src/main.rs +++ b/examples/ta/tcp_client-rs/src/main.rs @@ -75,7 +75,7 @@ fn invoke_command(cmd_id: u32, (p0, p1, p2, _): &mut ParametersAny<'_>) -> Resul let p1 = p1.as_value_input()?; let p2 = p2.as_memref_input()?; - let address = core::str::from_utf8(p0.get_buffer()).map_err(|e| { + let address = core::str::from_utf8(unsafe { p0.get_buffer() }).map_err(|e| { trace_println!("Failed to parse address from UTF-8: {}", e); ErrorKind::BadParameters })?; @@ -84,7 +84,7 @@ fn invoke_command(cmd_id: u32, (p0, p1, p2, _): &mut ParametersAny<'_>) -> Resul trace_println!("Invalid IP version parameter"); ErrorKind::BadParameters })?; - let http_data = p2.get_buffer(); + let http_data = unsafe { p2.get_buffer() }; tcp_client(address, port, ip_version, http_data) } diff --git a/examples/ta/tls_server-rs/src/main.rs b/examples/ta/tls_server-rs/src/main.rs index 232012a4..48408aa9 100644 --- a/examples/ta/tls_server-rs/src/main.rs +++ b/examples/ta/tls_server-rs/src/main.rs @@ -22,7 +22,7 @@ use lazy_static::lazy_static; use optee_utee::prelude::*; use optee_utee::{ErrorKind, Result}; use proto::tls_server::Command; -use rustls::pki_types::{pem::PemObject, CertificateDer, PrivateKeyDer}; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject}; use std::collections::HashMap; use std::io::{Cursor, Read, Write}; use std::sync::{Arc, Mutex, RwLock}; @@ -79,7 +79,7 @@ fn invoke_command(cmd_id: u32, params: &mut ParametersAny<'_>) -> Result<()> { } Command::DoTlsRead => { let p1 = params.1.as_memref_input()?; - let buffer = p1.get_buffer(); + let buffer = unsafe { p1.get_buffer() }; trace_println!("[+] do_tls_read"); do_tls_read(session_id, buffer).map_err(|e| { trace_println!("[-] Failed to read TLS data: {:?}", e); @@ -89,7 +89,7 @@ fn invoke_command(cmd_id: u32, params: &mut ParametersAny<'_>) -> Result<()> { Command::DoTlsWrite => { trace_println!("[+] do_tls_write"); let p1 = params.1.as_memref_output()?; - let lens = do_tls_write(session_id, p1.get_buffer_mut()).map_err(|e| { + let lens = do_tls_write(session_id, unsafe { p1.get_buffer_mut() }).map_err(|e| { trace_println!("[-] Failed to write TLS data: {:?}", e); ErrorKind::Generic })?; diff --git a/examples/ta/udp_socket-rs/src/main.rs b/examples/ta/udp_socket-rs/src/main.rs index 6674ea25..05ea5c8c 100644 --- a/examples/ta/udp_socket-rs/src/main.rs +++ b/examples/ta/udp_socket-rs/src/main.rs @@ -73,7 +73,7 @@ fn invoke_command(cmd_id: u32, params: &mut ParametersAny<'_>) -> Result<()> { let param0 = params.0.as_memref_input()?; let param1 = params.1.as_value_input()?; - let address = core::str::from_utf8(param0.get_buffer()).map_err(|e| { + let address = core::str::from_utf8(unsafe { param0.get_buffer() }).map_err(|e| { trace_println!("Failed to parse address from UTF-8: {}", e); ErrorKind::BadParameters })?; diff --git a/projects/web3/eth_wallet/ta/src/main.rs b/projects/web3/eth_wallet/ta/src/main.rs index 69e3a966..16a74383 100644 --- a/projects/web3/eth_wallet/ta/src/main.rs +++ b/projects/web3/eth_wallet/ta/src/main.rs @@ -25,7 +25,7 @@ use optee_utee::{Error, ErrorKind}; use proto::Command; use secure_db::SecureStorageClient; -use anyhow::{anyhow, bail, Result}; +use anyhow::{Result, anyhow, bail}; use wallet::Wallet; const DB_NAME: &str = "eth_wallet_db"; @@ -151,7 +151,7 @@ fn invoke_command( dbg_println!("[+] TA invoke command"); p1.set_updated_size(0)?; - let output_vec = match handle_invoke(Command::from(cmd_id), p0.get_buffer()) { + let output_vec = match handle_invoke(Command::from(cmd_id), unsafe { p0.get_buffer() }) { Ok(output) => output, Err(e) => { let err_message = format!("{:?}", e);