1use alloc::collections::BTreeMap;
2#[allow(unused_imports)] use alloc::vec::Vec;
4use core::fmt;
5
6use ax_memory_addr::{AddrRange, MemoryAddr};
7
8use crate::{MappingBackend, MappingError, MappingResult, MemoryArea};
9
10pub struct MemorySet<B: MappingBackend> {
12 areas: BTreeMap<B::Addr, MemoryArea<B>>,
13}
14
15impl<B: MappingBackend> MemorySet<B> {
16 pub const fn new() -> Self {
18 Self {
19 areas: BTreeMap::new(),
20 }
21 }
22
23 pub fn len(&self) -> usize {
25 self.areas.len()
26 }
27
28 pub fn is_empty(&self) -> bool {
30 self.areas.is_empty()
31 }
32
33 pub fn iter(&self) -> impl Iterator<Item = &MemoryArea<B>> {
35 self.areas.values()
36 }
37
38 pub fn overlaps(&self, range: AddrRange<B::Addr>) -> bool {
40 if let Some((_, before)) = self.areas.range(..range.start).last()
41 && before.va_range().overlaps(range)
42 {
43 return true;
44 }
45 if let Some((_, after)) = self.areas.range(range.start..).next()
46 && after.va_range().overlaps(range)
47 {
48 return true;
49 }
50 false
51 }
52
53 pub fn find(&self, addr: B::Addr) -> Option<&MemoryArea<B>> {
55 let candidate = self.areas.range(..=addr).last().map(|(_, a)| a);
56 candidate.filter(|a| a.va_range().contains(addr))
57 }
58
59 pub fn find_free_area(
74 &self,
75 hint: B::Addr,
76 size: usize,
77 limit: AddrRange<B::Addr>,
78 align: usize,
79 ) -> Option<B::Addr> {
80 if !size.is_multiple_of(align) {
81 return None;
83 }
84 let mut last_end: <B as MappingBackend>::Addr = hint.max(limit.start).align_up(align);
86 if let Some((_, area)) = self.areas.range(..last_end).last() {
87 last_end = last_end.max(area.end()).align_up(align);
88 }
89 for (&addr, area) in self.areas.range(last_end..) {
90 if last_end.checked_add(size).is_some_and(|end| end <= addr) {
91 return Some(last_end);
92 }
93 last_end = area.end().align_up(align);
94 }
95 if last_end
96 .checked_add(size)
97 .is_some_and(|end| end <= limit.end)
98 {
99 Some(last_end)
100 } else {
101 None
102 }
103 }
104
105 pub fn extend_area(
107 &mut self,
108 addr: B::Addr,
109 additional_size: usize,
110 page_table: &mut B::PageTable,
111 ) -> MappingResult {
112 if additional_size == 0 {
113 return Ok(());
114 }
115
116 let area_start = self
118 .areas
119 .range(..=addr)
120 .last()
121 .filter(|(_, a)| a.va_range().contains(addr))
122 .map(|(&start, _)| start)
123 .ok_or(MappingError::InvalidParam)?;
124
125 let area_end = self.areas[&area_start].end();
127 let new_end = area_end
128 .checked_add(additional_size)
129 .ok_or(MappingError::InvalidParam)?;
130 if let Some((_, next)) = self.areas.range(area_end..).next()
131 && new_end > next.start()
132 {
133 return Err(MappingError::AlreadyExists);
134 }
135
136 self.areas
137 .get_mut(&area_start)
138 .unwrap()
139 .grow_right(additional_size, page_table)?;
140 Ok(())
141 }
142
143 pub fn map(
152 &mut self,
153 area: MemoryArea<B>,
154 page_table: &mut B::PageTable,
155 unmap_overlap: bool,
156 ) -> MappingResult {
157 if area.va_range().is_empty() {
158 return Err(MappingError::InvalidParam);
159 }
160
161 if self.overlaps(area.va_range()) {
162 if unmap_overlap {
163 self.unmap(area.start(), area.size(), page_table)?;
164 } else {
165 return Err(MappingError::AlreadyExists);
166 }
167 }
168
169 area.map_area(page_table)?;
170 assert!(self.areas.insert(area.start(), area).is_none());
171 Ok(())
172 }
173
174 pub fn unmap(
181 &mut self,
182 start: B::Addr,
183 size: usize,
184 page_table: &mut B::PageTable,
185 ) -> MappingResult {
186 let range =
187 AddrRange::try_from_start_size(start, size).ok_or(MappingError::InvalidParam)?;
188 if range.is_empty() {
189 return Ok(());
190 }
191
192 let end = range.end;
193
194 self.areas.retain(|_, area| {
196 if area.va_range().contained_in(range) {
197 area.unmap_area(page_table).unwrap();
198 false
199 } else {
200 true
201 }
202 });
203
204 if let Some((&before_start, before)) = self.areas.range_mut(..start).last() {
206 let before_end = before.end();
207 if before_end > start {
208 if before_end <= end {
209 before.shrink_right(start.sub_addr(before_start), page_table)?;
211 } else {
212 let right_part = before.split(end).unwrap();
214 before.shrink_right(start.sub_addr(before_start), page_table)?;
215 assert_eq!(right_part.start().into(), Into::<usize>::into(end));
216 self.areas.insert(end, right_part);
217 }
218 }
219 }
220
221 if let Some((&after_start, after)) = self.areas.range_mut(start..).next() {
223 let after_end = after.end();
224 if after_start < end {
225 let mut new_area = self.areas.remove(&after_start).unwrap();
227 new_area.shrink_left(after_end.sub_addr(end), page_table)?;
228 assert_eq!(new_area.start().into(), Into::<usize>::into(end));
229 self.areas.insert(end, new_area);
230 }
231 }
232
233 Ok(())
234 }
235
236 pub fn unmap_metadata(&mut self, start: B::Addr, size: usize) -> MappingResult {
241 let range =
242 AddrRange::try_from_start_size(start, size).ok_or(MappingError::InvalidParam)?;
243 if range.is_empty() {
244 return Ok(());
245 }
246
247 let end = range.end;
248
249 self.areas
250 .retain(|_, area| !area.va_range().contained_in(range));
251
252 if let Some((&before_start, before)) = self.areas.range_mut(..start).last() {
253 let before_end = before.end();
254 if before_end > start {
255 if before_end <= end {
256 before.shrink_right_metadata(start.sub_addr(before_start));
257 } else {
258 let right_part = before.split(end).unwrap();
259 before.shrink_right_metadata(start.sub_addr(before_start));
260 assert_eq!(right_part.start().into(), Into::<usize>::into(end));
261 self.areas.insert(end, right_part);
262 }
263 }
264 }
265
266 if let Some((&after_start, _)) = self.areas.range(start..).next()
267 && after_start < end
268 {
269 let mut new_area = self.areas.remove(&after_start).unwrap();
270 let after_end = new_area.end();
271 new_area.shrink_left_metadata(after_end.sub_addr(end));
272 assert_eq!(new_area.start().into(), Into::<usize>::into(end));
273 self.areas.insert(end, new_area);
274 }
275
276 Ok(())
277 }
278
279 pub fn replace_area_metadata(&mut self, area: MemoryArea<B>) -> MappingResult {
281 if area.va_range().is_empty() {
282 return Err(MappingError::InvalidParam);
283 }
284
285 let start = area.start();
286 let end = area.end();
287
288 let old_start = self
289 .areas
290 .range(..=start)
291 .last()
292 .filter(|(_, old)| old.start() <= start && end <= old.end())
293 .map(|(&old_start, _)| old_start)
294 .ok_or(MappingError::InvalidParam)?;
295
296 let mut old_area = self.areas.remove(&old_start).unwrap();
297 if old_start < start {
298 let right_part = old_area.split(start).unwrap();
299 self.areas.insert(old_start, old_area);
300 old_area = right_part;
301 }
302 if old_area.end() > end {
303 let right_part = old_area.split(end).unwrap();
304 self.areas.insert(right_part.start(), right_part);
305 }
306 assert!(self.areas.insert(start, area).is_none());
307 Ok(())
308 }
309
310 pub fn clear(&mut self, page_table: &mut B::PageTable) -> MappingResult {
312 for area in self.areas.values() {
313 area.unmap_area(page_table)?;
314 }
315 self.areas.clear();
316 Ok(())
317 }
318
319 pub fn protect(
329 &mut self,
330 start: B::Addr,
331 size: usize,
332 update_flags: impl Fn(B::Flags) -> Option<B::Flags>,
333 page_table: &mut B::PageTable,
334 ) -> MappingResult {
335 self.protect_with_reported_flags(
336 start,
337 size,
338 |flags, _reported_flags| update_flags(flags).map(|new_flags| (new_flags, new_flags)),
339 page_table,
340 )
341 }
342
343 pub fn protect_with_reported_flags(
345 &mut self,
346 start: B::Addr,
347 size: usize,
348 update_flags: impl Fn(B::Flags, B::Flags) -> Option<(B::Flags, B::Flags)>,
349 page_table: &mut B::PageTable,
350 ) -> MappingResult {
351 let end = start.checked_add(size).ok_or(MappingError::InvalidParam)?;
352 let mut to_insert = Vec::new();
353 for (&area_start, area) in self.areas.iter_mut() {
354 let area_end = area.end();
355
356 if let Some((new_flags, new_reported_flags)) =
357 update_flags(area.flags(), area.reported_flags())
358 {
359 if area_start >= end {
360 break;
363 } else if area_end <= start {
364 } else if area_start >= start && area_end <= end {
368 area.protect_area(new_flags, page_table)?;
371 area.set_flags_with_reported_flags(new_flags, new_reported_flags);
372 } else if area_start < start && area_end > end {
373 let mut middle_part = area.split(start).unwrap();
376 let right_part = middle_part.split(end).unwrap();
377
378 middle_part.protect_area(new_flags, page_table)?;
379 middle_part.set_flags_with_reported_flags(new_flags, new_reported_flags);
380
381 to_insert.push((right_part.start(), right_part));
382 to_insert.push((middle_part.start(), middle_part));
383 } else if area_end > end {
384 let right_part = area.split(end).unwrap();
387 area.protect_area(new_flags, page_table)?;
388 area.set_flags_with_reported_flags(new_flags, new_reported_flags);
389
390 to_insert.push((right_part.start(), right_part));
391 } else {
392 let mut right_part = area.split(start).unwrap();
395 right_part.protect_area(new_flags, page_table)?;
396 right_part.set_flags_with_reported_flags(new_flags, new_reported_flags);
397
398 to_insert.push((right_part.start(), right_part));
399 }
400 }
401 }
402 self.areas.extend(to_insert);
403 Ok(())
404 }
405}
406
407impl<B: MappingBackend> Default for MemorySet<B> {
408 fn default() -> Self {
409 Self::new()
410 }
411}
412
413impl<B: MappingBackend> fmt::Debug for MemorySet<B>
414where
415 B::Addr: fmt::Debug,
416 B::Flags: fmt::Debug,
417{
418 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
419 f.debug_list().entries(self.areas.values()).finish()
420 }
421}