1use crate::mm::{MemoryManager, PAGE_SIZE};
6use bitflags::bitflags;
7use range_map::RangeMap;
8use starnix_logging::track_stub;
9use starnix_sync::{LockDepMutex, UserFaultInner};
10use starnix_uapi::errors::Errno;
11use starnix_uapi::user_address::UserAddress;
12use starnix_uapi::{
13 UFFD_FEATURE_EVENT_FORK, UFFD_FEATURE_EVENT_REMAP, UFFD_FEATURE_EVENT_REMOVE,
14 UFFD_FEATURE_EVENT_UNMAP, UFFD_FEATURE_MINOR_HUGETLBFS, UFFD_FEATURE_MINOR_SHMEM,
15 UFFD_FEATURE_MISSING_HUGETLBFS, UFFD_FEATURE_MISSING_SHMEM, UFFD_FEATURE_SIGBUS,
16 UFFD_FEATURE_THREAD_ID, UFFDIO_CONTINUE_MODE_DONTWAKE, UFFDIO_COPY_MODE_DONTWAKE,
17 UFFDIO_COPY_MODE_WP, UFFDIO_REGISTER_MODE_MINOR, UFFDIO_REGISTER_MODE_MISSING,
18 UFFDIO_REGISTER_MODE_WP, UFFDIO_ZEROPAGE_MODE_DONTWAKE, errno, error,
19};
20use std::ops::Range;
21use std::sync::{Arc, Weak};
22
23#[derive(Debug)]
24pub struct UserFault {
25 mm: Weak<MemoryManager>,
26 state: LockDepMutex<UserFaultState, UserFaultInner>,
27}
28
29#[derive(Debug, Clone)]
30struct UserFaultState {
31 features: Option<UserFaultFeatures>,
33
34 userfault_pages: RangeMap<UserAddress, bool>,
37}
38
39impl UserFault {
40 pub fn new(mm: Weak<MemoryManager>) -> Self {
41 Self { mm, state: LockDepMutex::new(UserFaultState::new()) }
42 }
43
44 pub fn insert_pages(&self, range: Range<UserAddress>, value: bool) {
45 let _ = self.state.lock().userfault_pages.insert(range, value);
47 }
48
49 pub fn remove_pages(&self, range: Range<UserAddress>) -> bool {
50 !self.state.lock().userfault_pages.remove(range).is_empty()
51 }
52
53 pub fn get_registered_pages_overlapping_range(
54 &self,
55 range: Range<UserAddress>,
56 ) -> Vec<Range<UserAddress>> {
57 self.state.lock().userfault_pages.get_keys(range).cloned().collect()
58 }
59
60 pub fn contains_addr(&self, addr: UserAddress) -> bool {
61 self.state.lock().userfault_pages.get(addr).is_some()
62 }
63
64 pub fn get_first_populated_page_after(&self, addr: UserAddress) -> Option<UserAddress> {
65 self.state.lock().userfault_pages.get(addr).map(|(affected_range, is_populated)| {
66 if *is_populated { addr } else { affected_range.end }
67 })
68 }
69
70 pub fn is_initialized(self: &Arc<Self>) -> bool {
71 self.state.lock().features.is_some()
72 }
73
74 pub fn has_features(self: &Arc<Self>, features: UserFaultFeatures) -> bool {
75 self.state.lock().features.map(|f| f.contains(features)).unwrap_or(false)
76 }
77
78 pub fn initialize(self: &Arc<Self>, features: UserFaultFeatures) {
79 self.state.lock().features = Some(features);
80 }
81
82 pub fn op_register(
83 self: &Arc<Self>,
84 start: UserAddress,
85 len: u64,
86 mode: FaultRegisterMode,
87 ) -> Result<SupportedUserFaultIoctls, Errno> {
88 if !self.is_initialized() {
89 return error!(EINVAL);
90 }
91 if !self.has_features(UserFaultFeatures::SIGBUS) {
92 track_stub!(TODO("https://fxbug.dev/391599171"), "userfault without SIGBUS feature");
93 return error!(ENOTSUP);
94 }
95 check_op_range(start, len)?;
96 let mm = self.mm.upgrade().ok_or_else(|| errno!(EINVAL))?;
97
98 mm.register_with_uffd(start, len as usize, self, mode)?;
99 Ok(SupportedUserFaultIoctls::COPY | SupportedUserFaultIoctls::ZERO_PAGE)
100 }
101
102 pub fn op_unregister(self: &Arc<Self>, start: UserAddress, len: u64) -> Result<(), Errno> {
103 if !self.is_initialized() {
104 return error!(EINVAL);
105 }
106 check_op_range(start, len)?;
107 let mm = self.mm.upgrade().ok_or_else(|| errno!(EINVAL))?;
108 mm.unregister_range_from_uffd(self, start, len as usize)
109 }
110
111 pub fn op_copy(
112 self: &Arc<Self>,
113 mm_source: &MemoryManager,
114 source: UserAddress,
115 dest: UserAddress,
116 len: u64,
117 _mode: FaultCopyMode,
118 ) -> Result<usize, Errno> {
119 if !self.is_initialized() {
120 return error!(EINVAL);
121 }
122 check_op_range(source, len)?;
123 check_op_range(dest, len)?;
124 let mm = self.mm.upgrade().ok_or_else(|| errno!(EINVAL))?;
125
126 if Arc::as_ptr(&mm) == mm_source as *const MemoryManager {
129 mm.copy_from_uffd(source, dest, len as usize, self)
130 } else {
131 let mut buf = vec![std::mem::MaybeUninit::uninit(); len as usize];
132 let buf = mm_source.syscall_read_memory(source, &mut buf)?;
133 mm.fill_from_uffd(dest, buf, len as usize, self)
134 }
135 }
136
137 pub fn op_zero(
138 self: &Arc<Self>,
139 start: UserAddress,
140 len: u64,
141 _mode: FaultZeroMode,
142 ) -> Result<usize, Errno> {
143 if !self.is_initialized() {
144 return error!(EINVAL);
145 }
146 check_op_range(start, len)?;
147 let mm = self.mm.upgrade().ok_or_else(|| errno!(EINVAL))?;
148 mm.zero_from_uffd(start, len as usize, self)
149 }
150
151 pub fn cleanup(self: &Arc<Self>) {
152 if let Some(mm) = self.mm.upgrade() {
153 mm.unregister_uffd(self);
154 }
155 }
156}
157
158impl UserFaultState {
159 pub fn new() -> Self {
160 Self { features: None, userfault_pages: RangeMap::default() }
161 }
162}
163
164bitflags! {
165 #[derive(Debug, Clone, Copy, Eq, PartialEq)]
166 pub struct UserFaultFeatures: u32 {
167 const ALL_SUPPORTED = UFFD_FEATURE_SIGBUS;
168 const EVENT_FORK = UFFD_FEATURE_EVENT_FORK;
169 const EVENT_REMAP = UFFD_FEATURE_EVENT_REMAP;
170 const EVENT_REMOVE = UFFD_FEATURE_EVENT_REMOVE;
171 const EVENT_UNMAP = UFFD_FEATURE_EVENT_UNMAP;
172 const MISSING_HUGETLBFS = UFFD_FEATURE_MISSING_HUGETLBFS;
173 const MISSING_SHMEM = UFFD_FEATURE_MISSING_SHMEM;
174 const SIGBUS = UFFD_FEATURE_SIGBUS;
175 const THREAD_ID = UFFD_FEATURE_THREAD_ID;
176 const MINOR_HUGETLBFS = UFFD_FEATURE_MINOR_HUGETLBFS;
177 const MINOR_SHMEM = UFFD_FEATURE_MINOR_SHMEM;
178 }
179
180 #[derive(Debug, Clone, Copy, Eq, PartialEq)]
181 pub struct FaultRegisterMode: u32 {
182 const MINOR = UFFDIO_REGISTER_MODE_MINOR;
183 const MISSING = UFFDIO_REGISTER_MODE_MISSING;
184 const WRITE_PROTECT = UFFDIO_REGISTER_MODE_WP;
185 }
186
187 pub struct FaultCopyMode: u32 {
188 const DONT_WAKE = UFFDIO_COPY_MODE_DONTWAKE;
189 const WRITE_PROTECT = UFFDIO_COPY_MODE_WP;
190 }
191
192 pub struct FaultZeroMode: u32 {
193 const DONT_WAKE = UFFDIO_ZEROPAGE_MODE_DONTWAKE;
194 }
195
196 pub struct FaultContinueMode: u32 {
197 const DONT_WAKE = UFFDIO_CONTINUE_MODE_DONTWAKE;
198 }
199
200
201 pub struct SupportedUserFaultIoctls: u64 {
202 const COPY = 1 << starnix_uapi::_UFFDIO_COPY;
203 const WAKE = 1 << starnix_uapi::_UFFDIO_WAKE;
204 const WRITE_PROTECT = 1 << starnix_uapi::_UFFDIO_WRITEPROTECT;
205 const ZERO_PAGE = 1 << starnix_uapi::_UFFDIO_ZEROPAGE;
206 const CONTINUE = 1 << starnix_uapi::_UFFDIO_CONTINUE;
207 }
208}
209
210fn check_op_range(addr: UserAddress, len: u64) -> Result<(), Errno> {
211 if addr.is_aligned(*PAGE_SIZE) && len % *PAGE_SIZE == 0 && len > 0 {
212 Ok(())
213 } else {
214 error!(EINVAL)
215 }
216}