1use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
14use std::sync::{Arc, Weak};
15
16use parking_lot::Mutex;
17use smallvec::SmallVec;
18
19use crate::{BackendKind, Error};
20
21#[derive(Debug, Default, Copy, Clone)]
24pub struct BudgetCaps {
25 pub soft_cap_bytes: Option<u64>,
26 pub hard_cap_bytes: Option<u64>,
27}
28
29#[derive(Debug, Default, Clone)]
34pub struct BudgetCapsSet {
35 pub cpu: BudgetCaps,
36 #[cfg(feature = "wgpu")]
37 pub wgpu: BudgetCaps,
38 #[cfg(feature = "opencl")]
39 pub opencl: BudgetCaps,
40 #[cfg(feature = "cuda")]
41 pub cuda: BudgetCaps,
42}
43
44impl BudgetCapsSet {
45 fn for_backend(&self, backend: BackendKind) -> BudgetCaps {
46 match backend {
47 BackendKind::Cpu => self.cpu,
48 #[cfg(feature = "wgpu")]
49 BackendKind::Wgpu => self.wgpu,
50 #[cfg(feature = "opencl")]
51 BackendKind::OpenCl => self.opencl,
52 #[cfg(feature = "cuda")]
53 BackendKind::Cuda => self.cuda,
54 #[cfg(feature = "wgpu")]
59 BackendKind::Vulkan
60 | BackendKind::D3D12
61 | BackendKind::D3D11
62 | BackendKind::Metal
63 | BackendKind::OpenGL => self.wgpu,
64 #[cfg(not(feature = "wgpu"))]
65 BackendKind::Vulkan
66 | BackendKind::D3D12
67 | BackendKind::D3D11
68 | BackendKind::Metal
69 | BackendKind::OpenGL => self.cpu,
70 #[cfg(all(target_os = "android", feature = "wgpu"))]
75 BackendKind::AHardwareBuffer
76 | BackendKind::AndroidPresentation
77 | BackendKind::AndroidSurfaceControl
78 | BackendKind::AImageWriter
79 | BackendKind::MediaCodec
80 | BackendKind::ForeignGl
81 | BackendKind::ForeignVulkan => self.wgpu,
82 #[cfg(all(target_os = "android", not(feature = "wgpu")))]
88 BackendKind::AHardwareBuffer
89 | BackendKind::AndroidPresentation
90 | BackendKind::AndroidSurfaceControl
91 | BackendKind::AImageWriter
92 | BackendKind::MediaCodec
93 | BackendKind::ForeignGl
94 | BackendKind::ForeignVulkan => self.cpu,
95 #[cfg(all(feature = "web-codecs", feature = "wgpu"))]
99 BackendKind::WebCodecs => self.wgpu,
100 #[cfg(all(feature = "web-codecs", not(feature = "wgpu")))]
101 BackendKind::WebCodecs => self.cpu,
102 }
103 }
104}
105
106#[derive(Default)]
108struct BudgetCountersSet {
109 cpu: AtomicU64,
110 #[cfg(feature = "wgpu")]
111 wgpu: AtomicU64,
112 #[cfg(feature = "opencl")]
113 opencl: AtomicU64,
114 #[cfg(feature = "cuda")]
115 cuda: AtomicU64,
116}
117
118impl BudgetCountersSet {
119 fn get(&self, backend: BackendKind) -> &AtomicU64 {
120 match backend {
121 BackendKind::Cpu => &self.cpu,
122 #[cfg(feature = "wgpu")]
123 BackendKind::Wgpu => &self.wgpu,
124 #[cfg(feature = "opencl")]
125 BackendKind::OpenCl => &self.opencl,
126 #[cfg(feature = "cuda")]
127 BackendKind::Cuda => &self.cuda,
128 #[cfg(feature = "wgpu")]
130 BackendKind::Vulkan
131 | BackendKind::D3D12
132 | BackendKind::D3D11
133 | BackendKind::Metal
134 | BackendKind::OpenGL => &self.wgpu,
135 #[cfg(not(feature = "wgpu"))]
136 BackendKind::Vulkan
137 | BackendKind::D3D12
138 | BackendKind::D3D11
139 | BackendKind::Metal
140 | BackendKind::OpenGL => &self.cpu,
141 #[cfg(all(target_os = "android", feature = "wgpu"))]
144 BackendKind::AHardwareBuffer
145 | BackendKind::AndroidPresentation
146 | BackendKind::AndroidSurfaceControl
147 | BackendKind::AImageWriter
148 | BackendKind::MediaCodec
149 | BackendKind::ForeignGl
150 | BackendKind::ForeignVulkan => &self.wgpu,
151 #[cfg(all(target_os = "android", not(feature = "wgpu")))]
152 BackendKind::AHardwareBuffer
153 | BackendKind::AndroidPresentation
154 | BackendKind::AndroidSurfaceControl
155 | BackendKind::AImageWriter
156 | BackendKind::MediaCodec
157 | BackendKind::ForeignGl
158 | BackendKind::ForeignVulkan => &self.cpu,
159 #[cfg(all(feature = "web-codecs", feature = "wgpu"))]
162 BackendKind::WebCodecs => &self.wgpu,
163 #[cfg(all(feature = "web-codecs", not(feature = "wgpu")))]
164 BackendKind::WebCodecs => &self.cpu,
165 }
166 }
167}
168
169#[derive(Debug, Copy, Clone, PartialEq, Eq)]
170pub enum BudgetPressure {
171 Ok,
172 Soft,
173 Hard,
174}
175
176#[derive(Debug, Copy, Clone)]
177pub struct BudgetPressureEvent {
178 pub backend: BackendKind,
179 pub level: BudgetPressure,
180 pub current_usage_bytes: u64,
181 pub required_bytes: u64,
182 pub cap_bytes: u64,
183}
184
185#[cfg(not(target_family = "wasm"))]
190pub type PressureCallback = Arc<dyn Fn(&BudgetPressureEvent) + Send + Sync>;
191#[cfg(target_family = "wasm")]
192pub type PressureCallback = Arc<dyn Fn(&BudgetPressureEvent)>;
193
194pub trait EvictablePool: crate::MaybeSendSync {
198 fn evict_unused(&self) -> u64;
200 fn bytes_in_use(&self) -> u64;
202}
203
204#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
205pub struct PoolHandle(pub u32);
206
207#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
208pub struct PressureCallbackHandle(pub u32);
209
210type PoolEntry = (PoolHandle, BackendKind, Weak<dyn EvictablePool>);
211type CallbackEntry = (PressureCallbackHandle, PressureCallback);
212
213pub struct MemoryBudget {
214 caps: BudgetCapsSet,
215 usage: BudgetCountersSet,
216 pools: Mutex<SmallVec<[PoolEntry; 4]>>,
217 callbacks: Mutex<SmallVec<[CallbackEntry; 4]>>,
218 next_pool_handle: AtomicU32,
219 next_callback_handle: AtomicU32,
220}
221
222impl core::fmt::Debug for MemoryBudget {
223 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
224 f.debug_struct("MemoryBudget").field("caps", &self.caps).finish_non_exhaustive()
225 }
226}
227
228pub struct BudgetReservation {
231 budget: Arc<MemoryBudget>,
232 backend: BackendKind,
233 bytes: u64,
234}
235
236impl BudgetReservation {
237 pub fn backend(&self) -> BackendKind {
238 self.backend
239 }
240 pub fn bytes(&self) -> u64 {
241 self.bytes
242 }
243}
244
245impl Drop for BudgetReservation {
246 fn drop(&mut self) {
247 self.budget.usage.get(self.backend).fetch_sub(self.bytes, Ordering::Release);
248 }
249}
250
251impl core::fmt::Debug for BudgetReservation {
252 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
253 f.debug_struct("BudgetReservation").field("backend", &self.backend).field("bytes", &self.bytes).finish()
254 }
255}
256
257impl MemoryBudget {
258 pub fn new(caps: BudgetCapsSet) -> Arc<Self> {
259 Arc::new(Self {
260 caps,
261 usage: BudgetCountersSet::default(),
262 pools: Mutex::new(SmallVec::new()),
263 callbacks: Mutex::new(SmallVec::new()),
264 next_pool_handle: AtomicU32::new(1),
265 next_callback_handle: AtomicU32::new(1),
266 })
267 }
268
269 pub fn caps(&self, backend: BackendKind) -> BudgetCaps {
270 self.caps.for_backend(backend)
271 }
272
273 pub fn current_usage(&self, backend: BackendKind) -> u64 {
274 self.usage.get(backend).load(Ordering::Acquire)
275 }
276
277 pub fn available(&self, backend: BackendKind) -> u64 {
278 let cap = self.caps.for_backend(backend).hard_cap_bytes.unwrap_or(u64::MAX);
279 cap.saturating_sub(self.current_usage(backend))
280 }
281
282 pub fn pressure(&self, backend: BackendKind, required: u64) -> BudgetPressure {
283 let caps = self.caps.for_backend(backend);
284 let usage = self.current_usage(backend);
285 let projected = usage.saturating_add(required);
286 if let Some(hard) = caps.hard_cap_bytes
287 && projected > hard
288 {
289 return BudgetPressure::Hard;
290 }
291 if let Some(soft) = caps.soft_cap_bytes
292 && projected > soft
293 {
294 return BudgetPressure::Soft;
295 }
296 BudgetPressure::Ok
297 }
298
299 pub fn try_reserve(self: &Arc<Self>, backend: BackendKind, bytes: u64) -> Result<BudgetReservation, Error> {
300 let caps = self.caps.for_backend(backend);
304 if caps.soft_cap_bytes.is_none() && caps.hard_cap_bytes.is_none() {
305 self.usage.get(backend).fetch_add(bytes, Ordering::Release);
306 return Ok(BudgetReservation { budget: self.clone(), backend, bytes });
307 }
308
309 let level = self.pressure(backend, bytes);
316 if level != BudgetPressure::Ok {
317 let cap = caps.hard_cap_bytes.or(caps.soft_cap_bytes).unwrap_or(u64::MAX);
318 let event = BudgetPressureEvent {
319 backend,
320 level,
321 current_usage_bytes: self.current_usage(backend),
322 required_bytes: bytes,
323 cap_bytes: cap,
324 };
325
326 {
329 let cbs = self.callbacks.lock().clone();
330 for (_, cb) in cbs.iter() {
331 let cb = cb.clone();
332 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| cb(&event)));
333 }
334 }
335
336 {
338 let pools = self.pools.lock().clone();
339 for (_, bk, weak) in pools.iter() {
340 if *bk != backend {
341 continue;
342 }
343 if let Some(pool) = weak.upgrade() {
344 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| pool.evict_unused()));
345 }
346 }
347 }
348 }
349
350 if self.try_commit(backend, bytes) {
354 return Ok(BudgetReservation { budget: self.clone(), backend, bytes });
355 }
356
357 let available = self.available(backend);
359 Err(Error::OutOfGpuMemory { required_bytes: bytes, available_bytes: available, backend })
360 }
361
362 fn try_commit(&self, backend: BackendKind, bytes: u64) -> bool {
366 let counter = self.usage.get(backend);
367 let hard = self.caps.for_backend(backend).hard_cap_bytes;
368 loop {
369 let cur = counter.load(Ordering::Acquire);
370 let next = cur.saturating_add(bytes);
371 if let Some(h) = hard
372 && next > h
373 {
374 return false;
375 }
376 if counter.compare_exchange(cur, next, Ordering::AcqRel, Ordering::Acquire).is_ok() {
377 return true;
378 }
379 }
380 }
381
382 pub fn register_pool(&self, backend: BackendKind, pool: &Arc<dyn EvictablePool>) -> PoolHandle {
383 let h = PoolHandle(self.next_pool_handle.fetch_add(1, Ordering::Relaxed));
384 self.pools.lock().push((h, backend, Arc::downgrade(pool)));
385 h
386 }
387
388 pub fn unregister_pool(&self, handle: PoolHandle) {
389 self.pools.lock().retain(|(h, _, _)| *h != handle);
390 }
391
392 pub fn register_pressure_callback(&self, cb: PressureCallback) -> PressureCallbackHandle {
393 let h = PressureCallbackHandle(self.next_callback_handle.fetch_add(1, Ordering::Relaxed));
394 self.callbacks.lock().push((h, cb));
395 h
396 }
397
398 pub fn unregister_pressure_callback(&self, handle: PressureCallbackHandle) {
399 self.callbacks.lock().retain(|(h, _)| *h != handle);
400 }
401}