Skip to main content

gwr_models/memory/
memory_map.rs

1// Copyright (c) 2023 Graphcore Ltd. All rights reserved.
2
3use std::collections::BTreeMap;
4
5use gwr_engine::sim_error;
6use gwr_engine::types::SimError;
7
8#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
9pub struct DeviceId(pub u64);
10
11#[derive(Clone, Debug)]
12pub struct MemoryRegion {
13    pub start: u64,
14    pub end: u64,
15    pub device: DeviceId,
16}
17
18pub struct MemoryMap {
19    // key = start address of region
20    regions: BTreeMap<u64, MemoryRegion>,
21}
22
23impl Default for MemoryMap {
24    fn default() -> Self {
25        Self::new()
26    }
27}
28
29impl MemoryMap {
30    #[must_use]
31    pub fn new() -> Self {
32        Self {
33            regions: BTreeMap::new(),
34        }
35    }
36
37    /// Map a [start, start+size-1] region to a device.
38    pub fn insert(&mut self, start: u64, size: u64, device: DeviceId) -> Result<(), SimError> {
39        let end = if size > 0 {
40            start + size - 1
41        } else {
42            return sim_error!("Invalid region size {size}");
43        };
44
45        // Check overlap with previous region (if any)
46        if let Some((_, prev)) = self.regions.range(..=start).next_back()
47            && prev.end >= start
48        {
49            return sim_error!("Region overlap at {start}");
50        }
51
52        // Check overlap with next region (if any)
53        if let Some((_, next)) = self.regions.range(start..).next()
54            && next.start <= end
55        {
56            return sim_error!("Region overlap at {end}");
57        }
58
59        let region = MemoryRegion { start, end, device };
60        self.regions.insert(start, region);
61        Ok(())
62    }
63
64    /// Remove a region by its exact start address.
65    #[must_use]
66    pub fn unmap(&mut self, start: u64) -> Option<MemoryRegion> {
67        self.regions.remove(&start)
68    }
69
70    /// Resolve an address to (device_id, offset_in_region).
71    #[must_use]
72    pub fn lookup(&self, addr: u64) -> Option<(DeviceId, u64)> {
73        // Find region with greatest start <= addr
74        let (_, region) = self.regions.range(..=addr).next_back()?;
75        if addr <= region.end {
76            let offset = addr - region.start;
77            Some((region.device, offset))
78        } else {
79            None
80        }
81    }
82
83    #[must_use]
84    pub fn num_regions(&self) -> usize {
85        self.regions.len()
86    }
87
88    /// Iterate all mapped ranges.
89    pub fn regions(&self) -> impl Iterator<Item = &MemoryRegion> {
90        self.regions.values()
91    }
92}
93
94#[cfg(test)]
95mod tests {
96    use crate::memory::memory_map::{DeviceId, MemoryMap};
97
98    fn setup_map() -> MemoryMap {
99        let mut memory_map = MemoryMap::new();
100        memory_map.insert(0x0000_0000, 0x1000, DeviceId(1)).unwrap();
101        memory_map.insert(0x0000_2000, 0x1000, DeviceId(2)).unwrap();
102        memory_map.insert(0x0000_4000, 0x1000, DeviceId(3)).unwrap();
103        memory_map
104    }
105
106    #[test]
107    fn insert_successfully() {
108        let memory_map = setup_map();
109        assert_eq!(memory_map.num_regions(), 3);
110    }
111
112    #[test]
113    fn insert_in_between() {
114        let mut memory_map = setup_map();
115        memory_map.insert(0x0000_3000, 0x1000, DeviceId(4)).unwrap();
116        assert_eq!(memory_map.num_regions(), 4);
117    }
118
119    #[test]
120    #[should_panic(expected = "Region overlap")]
121    fn insert_overlap() {
122        let mut memory_map = setup_map();
123        memory_map.insert(0x0000_0F00, 0x200, DeviceId(3)).unwrap();
124    }
125
126    #[test]
127    #[should_panic(expected = "Region overlap")]
128    fn insert_overlap_inserted() {
129        let mut memory_map = setup_map();
130        memory_map.insert(0x0000_3000, 0x2000, DeviceId(3)).unwrap();
131    }
132
133    #[test]
134    fn address_lookup() {
135        let memmory_map = setup_map();
136        let (dev, offset) = memmory_map.lookup(0x0000_2004).unwrap();
137
138        assert_eq!(dev, DeviceId(2));
139        assert_eq!(offset, 0x4);
140    }
141
142    #[test]
143    fn address_lookup_begin() {
144        let memmory_map = setup_map();
145        let (dev, offset) = memmory_map.lookup(0x0000_4000).unwrap();
146
147        assert_eq!(dev, DeviceId(3));
148        assert_eq!(offset, 0x0);
149    }
150
151    #[test]
152    fn address_lookup_end() {
153        let memmory_map = setup_map();
154        let (dev, offset) = memmory_map.lookup(0x0000_4fff).unwrap();
155
156        assert_eq!(dev, DeviceId(3));
157        assert_eq!(offset, 0xfff);
158    }
159
160    #[test]
161    fn address_lookup_after() {
162        let memmory_map = setup_map();
163        assert!(memmory_map.lookup(0x0000_5000).is_none());
164    }
165
166    #[test]
167    #[should_panic(expected = "Invalid region size 0")]
168    fn insert_zero_sized() {
169        let mut memory_map = setup_map();
170        memory_map.insert(0x0000_8000, 0x0, DeviceId(4)).unwrap();
171    }
172}