Skip to main content

ntex_tls/openssl/
alloc.rs

1/// Route OpenSSL memory allocations to the Rust global allocator.
2///
3/// OpenSSL allocates with the C runtime `malloc` by default. After this call
4/// it allocates through [`std::alloc`], so an application that registers
5/// a `#[global_allocator]` (for example mimalloc) uses it for tls buffers
6/// and handshake state as well.
7///
8/// The setting is process wide and OpenSSL accepts it only before its first
9/// allocation, call this at the start of `main` before any other OpenSSL use.
10/// Returns `false` if OpenSSL has already allocated memory or the linked library
11/// does not support custom allocators (OpenSSL older than 1.1.0, `LibreSSL`,
12/// `BoringSSL`, `aws-lc`), OpenSSL keeps using `malloc` in that case.
13pub fn use_global_allocator() -> bool {
14    #[cfg(ntex_openssl_mem)]
15    {
16        shim::install()
17    }
18    #[cfg(not(ntex_openssl_mem))]
19    {
20        false
21    }
22}
23
24#[cfg(ntex_openssl_mem)]
25mod shim {
26    use std::alloc::{self, Layout};
27    use std::ffi::{c_char, c_int, c_void};
28    use std::ptr;
29
30    // Every block starts with a header that keeps the requested size, `free`
31    // has no size argument but `dealloc` needs the layout. 16 bytes keep the
32    // returned pointer aligned like `malloc` does (`max_align_t`).
33    const ALIGN: usize = 16;
34    const HEADER: usize = 16;
35
36    type MallocFn = unsafe extern "C" fn(usize, *const c_char, c_int) -> *mut c_void;
37    type ReallocFn = unsafe extern "C" fn(*mut c_void, usize, *const c_char, c_int) -> *mut c_void;
38    type FreeFn = unsafe extern "C" fn(*mut c_void, *const c_char, c_int);
39
40    unsafe extern "C" {
41        fn CRYPTO_set_mem_functions(m: MallocFn, r: ReallocFn, f: FreeFn) -> c_int;
42    }
43
44    pub(super) fn install() -> bool {
45        unsafe { CRYPTO_set_mem_functions(malloc, realloc, free) == 1 }
46    }
47
48    fn layout(size: usize) -> Option<Layout> {
49        Layout::from_size_align(size.checked_add(HEADER)?, ALIGN).ok()
50    }
51
52    /// Start and layout of a block returned by `malloc` or `realloc`
53    #[allow(clippy::cast_ptr_alignment)] // blocks are ALIGN aligned
54    unsafe fn block(p: *mut c_void) -> (*mut u8, Layout) {
55        unsafe {
56            let base = p.cast::<u8>().sub(HEADER);
57            let size = base.cast::<usize>().read();
58            (
59                base,
60                Layout::from_size_align_unchecked(size + HEADER, ALIGN),
61            )
62        }
63    }
64
65    #[allow(clippy::cast_ptr_alignment)] // blocks are ALIGN aligned
66    unsafe fn finish(base: *mut u8, size: usize) -> *mut c_void {
67        if base.is_null() {
68            return ptr::null_mut();
69        }
70        unsafe {
71            base.cast::<usize>().write(size);
72            base.add(HEADER).cast()
73        }
74    }
75
76    pub(super) unsafe extern "C" fn malloc(size: usize, _: *const c_char, _: c_int) -> *mut c_void {
77        // OpenSSL's own malloc returns null for empty allocations
78        if size == 0 {
79            return ptr::null_mut();
80        }
81        match layout(size) {
82            Some(layout) => unsafe { finish(alloc::alloc(layout), size) },
83            None => ptr::null_mut(),
84        }
85    }
86
87    pub(super) unsafe extern "C" fn realloc(
88        p: *mut c_void,
89        size: usize,
90        file: *const c_char,
91        line: c_int,
92    ) -> *mut c_void {
93        // OpenSSL forwards these cases to a custom realloc unchanged
94        if p.is_null() {
95            return unsafe { malloc(size, file, line) };
96        }
97        if size == 0 {
98            unsafe { free(p, file, line) };
99            return ptr::null_mut();
100        }
101        let Some(new) = layout(size) else {
102            return ptr::null_mut();
103        };
104        unsafe {
105            let (base, old) = block(p);
106            // on failure the original block stays valid, as with C realloc
107            finish(alloc::realloc(base, old, new.size()), size)
108        }
109    }
110
111    pub(super) unsafe extern "C" fn free(p: *mut c_void, _: *const c_char, _: c_int) {
112        if !p.is_null() {
113            unsafe {
114                let (base, layout) = block(p);
115                alloc::dealloc(base, layout);
116            }
117        }
118    }
119
120    #[cfg(test)]
121    mod tests {
122        use super::*;
123
124        #[test]
125        fn malloc_realloc_free() {
126            let (file, line) = (ptr::null(), 0);
127            unsafe {
128                assert!(malloc(0, file, line).is_null());
129                assert!(malloc(usize::MAX - 8, file, line).is_null());
130                free(ptr::null_mut(), file, line);
131
132                let p = malloc(3, file, line).cast::<u8>();
133                assert!(!p.is_null());
134                assert_eq!(p as usize % ALIGN, 0);
135                p.copy_from(b"abc".as_ptr(), 3);
136
137                // grow keeps the data
138                let p = realloc(p.cast(), 64 * 1024, file, line).cast::<u8>();
139                assert_eq!(p as usize % ALIGN, 0);
140                assert_eq!(std::slice::from_raw_parts(p, 3), b"abc");
141                p.add(64 * 1024 - 1).write(1);
142
143                // shrink keeps the prefix
144                let p = realloc(p.cast(), 2, file, line).cast::<u8>();
145                assert_eq!(std::slice::from_raw_parts(p, 2), b"ab");
146
147                // too large fails and keeps the block
148                assert!(realloc(p.cast(), usize::MAX - 8, file, line).is_null());
149                assert_eq!(std::slice::from_raw_parts(p, 2), b"ab");
150
151                // zero size frees
152                assert!(realloc(p.cast(), 0, file, line).is_null());
153
154                // null acts as malloc
155                let p = realloc(ptr::null_mut(), 5, file, line);
156                assert!(!p.is_null());
157                free(p, file, line);
158            }
159        }
160    }
161}
162
163#[cfg(test)]
164mod tests {
165    use tls_openssl::ssl::{SslConnector, SslMethod};
166
167    #[test]
168    fn rejected_after_openssl_allocated() {
169        drop(SslConnector::builder(SslMethod::tls()).unwrap());
170        assert!(!super::use_global_allocator());
171    }
172}