Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
171 changes: 149 additions & 22 deletions crates/optee-utee/src/parameter/memref.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand All @@ -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<u8> {
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.
Expand All @@ -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;
Expand Down Expand Up @@ -92,16 +147,29 @@ pub trait ParameterMemrefWrite {
/// the buffer capacity.
fn write_at<T: AsRef<[u8]>>(&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
Expand Down Expand Up @@ -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<Self> {
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<Self> {
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,
Expand All @@ -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<Self> {
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,
Expand All @@ -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::<u8>::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());
}
}
5 changes: 4 additions & 1 deletion crates/optee-utee/src/parameter/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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) => {
Expand Down
4 changes: 2 additions & 2 deletions examples/ta/acipher-rs/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(())
}
Expand All @@ -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)
}

Expand Down
12 changes: 6 additions & 6 deletions examples/ta/aes-rs/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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");
Expand All @@ -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);

Expand All @@ -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(())
}
Expand Down
31 changes: 18 additions & 13 deletions examples/ta/authentication-rs/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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)?;

Expand All @@ -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(())
}
Expand All @@ -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(())
Expand All @@ -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(())
}
Expand Down
Loading
Loading