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"))]
76 BackendKind::AHardwareBuffer
77 | BackendKind::AndroidPresentation
78 | BackendKind::AndroidSurfaceControl
79 | BackendKind::AImageWriter
80 | BackendKind::MediaCodec
81 | BackendKind::ForeignGl
82 | BackendKind::ForeignVulkan => self.wgpu,
83 #[cfg(all(target_os = "android", not(feature = "wgpu")))]
89 BackendKind::AHardwareBuffer
90 | BackendKind::AndroidPresentation
91 | BackendKind::AndroidSurfaceControl
92 | BackendKind::AImageWriter
93 | BackendKind::MediaCodec
94 | BackendKind::ForeignGl
95 | BackendKind::ForeignVulkan => self.cpu,
96 #[cfg(all(feature = "web-codecs", feature = "wgpu"))]
100 BackendKind::WebCodecs => self.wgpu,
101 #[cfg(all(feature = "web-codecs", not(feature = "wgpu")))]
102 BackendKind::WebCodecs => self.cpu,
103 }
104 }
105}
106
107#[derive(Default)]
109struct BudgetCountersSet {
110 cpu: AtomicU64,
111 #[cfg(feature = "wgpu")]
112 wgpu: AtomicU64,
113 #[cfg(feature = "opencl")]
114 opencl: AtomicU64,
115 #[cfg(feature = "cuda")]
116 cuda: AtomicU64,
117}
118
119impl BudgetCountersSet {
120 fn get(&self, backend: BackendKind) -> &AtomicU64 {
121 match backend {
122 BackendKind::Cpu => &self.cpu,
123 #[cfg(feature = "wgpu")]
124 BackendKind::Wgpu => &self.wgpu,
125 #[cfg(feature = "opencl")]
126 BackendKind::OpenCl => &self.opencl,
127 #[cfg(feature = "cuda")]
128 BackendKind::Cuda => &self.cuda,
129 #[cfg(feature = "wgpu")]
131 BackendKind::Vulkan
132 | BackendKind::D3D12
133 | BackendKind::D3D11
134 | BackendKind::Metal
135 | BackendKind::OpenGL => &self.wgpu,
136 #[cfg(not(feature = "wgpu"))]
137 BackendKind::Vulkan
138 | BackendKind::D3D12
139 | BackendKind::D3D11
140 | BackendKind::Metal
141 | BackendKind::OpenGL => &self.cpu,
142 #[cfg(all(target_os = "android", feature = "wgpu"))]
145 BackendKind::AHardwareBuffer
146 | BackendKind::AndroidPresentation
147 | BackendKind::AndroidSurfaceControl
148 | BackendKind::AImageWriter
149 | BackendKind::MediaCodec
150 | BackendKind::ForeignGl
151 | BackendKind::ForeignVulkan => &self.wgpu,
152 #[cfg(all(target_os = "android", not(feature = "wgpu")))]
153 BackendKind::AHardwareBuffer
154 | BackendKind::AndroidPresentation
155 | BackendKind::AndroidSurfaceControl
156 | BackendKind::AImageWriter
157 | BackendKind::MediaCodec
158 | BackendKind::ForeignGl
159 | BackendKind::ForeignVulkan => &self.cpu,
160 #[cfg(all(feature = "web-codecs", feature = "wgpu"))]
163 BackendKind::WebCodecs => &self.wgpu,
164 #[cfg(all(feature = "web-codecs", not(feature = "wgpu")))]
165 BackendKind::WebCodecs => &self.cpu,
166 }
167 }
168}
169
170#[derive(Debug, Copy, Clone, PartialEq, Eq)]
171pub enum BudgetPressure {
172 Ok,
173 Soft,
174 Hard,
175}
176
177#[derive(Debug, Copy, Clone)]
178pub struct BudgetPressureEvent {
179 pub backend: BackendKind,
180 pub level: BudgetPressure,
181 pub current_usage_bytes: u64,
182 pub required_bytes: u64,
183 pub cap_bytes: u64,
184}
185
186#[cfg(not(target_family = "wasm"))]
191pub type PressureCallback = Arc<dyn Fn(&BudgetPressureEvent) + Send + Sync>;
192#[cfg(target_family = "wasm")]
193pub type PressureCallback = Arc<dyn Fn(&BudgetPressureEvent)>;
194
195pub trait EvictablePool: crate::MaybeSendSync {
199 fn evict_unused(&self) -> u64;
201 fn bytes_in_use(&self) -> u64;
203}
204
205#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
206pub struct PoolHandle(pub u32);
207
208#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
209pub struct PressureCallbackHandle(pub u32);
210
211type PoolEntry = (PoolHandle, BackendKind, Weak<dyn EvictablePool>);
212type CallbackEntry = (PressureCallbackHandle, PressureCallback);
213
214pub struct MemoryBudget {
215 caps: BudgetCapsSet,
216 usage: BudgetCountersSet,
217 pools: Mutex<SmallVec<[PoolEntry; 4]>>,
218 callbacks: Mutex<SmallVec<[CallbackEntry; 4]>>,
219 next_pool_handle: AtomicU32,
220 next_callback_handle: AtomicU32,
221}
222
223impl core::fmt::Debug for MemoryBudget {
224 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
225 f.debug_struct("MemoryBudget").field("caps", &self.caps).finish_non_exhaustive()
226 }
227}
228
229pub struct BudgetReservation {
232 budget: Arc<MemoryBudget>,
233 backend: BackendKind,
234 bytes: u64,
235}
236
237impl BudgetReservation {
238 pub fn backend(&self) -> BackendKind {
239 self.backend
240 }
241 pub fn bytes(&self) -> u64 {
242 self.bytes
243 }
244}
245
246impl Drop for BudgetReservation {
247 fn drop(&mut self) {
248 self.budget.usage.get(self.backend).fetch_sub(self.bytes, Ordering::Release);
249 }
250}
251
252impl core::fmt::Debug for BudgetReservation {
253 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
254 f.debug_struct("BudgetReservation").field("backend", &self.backend).field("bytes", &self.bytes).finish()
255 }
256}
257
258impl MemoryBudget {
259 #[cfg_attr(target_family = "wasm", allow(clippy::arc_with_non_send_sync))]
263 pub fn new(caps: BudgetCapsSet) -> Arc<Self> {
264 Arc::new(Self {
265 caps,
266 usage: BudgetCountersSet::default(),
267 pools: Mutex::new(SmallVec::new()),
268 callbacks: Mutex::new(SmallVec::new()),
269 next_pool_handle: AtomicU32::new(1),
270 next_callback_handle: AtomicU32::new(1),
271 })
272 }
273
274 pub fn caps(&self, backend: BackendKind) -> BudgetCaps {
275 self.caps.for_backend(backend)
276 }
277
278 pub fn current_usage(&self, backend: BackendKind) -> u64 {
279 self.usage.get(backend).load(Ordering::Acquire)
280 }
281
282 pub fn available(&self, backend: BackendKind) -> u64 {
283 let cap = self.caps.for_backend(backend).hard_cap_bytes.unwrap_or(u64::MAX);
284 cap.saturating_sub(self.current_usage(backend))
285 }
286
287 pub fn pressure(&self, backend: BackendKind, required: u64) -> BudgetPressure {
288 let caps = self.caps.for_backend(backend);
289 let usage = self.current_usage(backend);
290 let projected = usage.saturating_add(required);
291 if let Some(hard) = caps.hard_cap_bytes
292 && projected > hard
293 {
294 return BudgetPressure::Hard;
295 }
296 if let Some(soft) = caps.soft_cap_bytes
297 && projected > soft
298 {
299 return BudgetPressure::Soft;
300 }
301 BudgetPressure::Ok
302 }
303
304 pub fn try_reserve(self: &Arc<Self>, backend: BackendKind, bytes: u64) -> Result<BudgetReservation, Error> {
305 let caps = self.caps.for_backend(backend);
309 if caps.soft_cap_bytes.is_none() && caps.hard_cap_bytes.is_none() {
310 self.usage.get(backend).fetch_add(bytes, Ordering::Release);
311 return Ok(BudgetReservation { budget: self.clone(), backend, bytes });
312 }
313
314 let level = self.pressure(backend, bytes);
321 if level != BudgetPressure::Ok {
322 let cap = caps.hard_cap_bytes.or(caps.soft_cap_bytes).unwrap_or(u64::MAX);
323 let event = BudgetPressureEvent {
324 backend,
325 level,
326 current_usage_bytes: self.current_usage(backend),
327 required_bytes: bytes,
328 cap_bytes: cap,
329 };
330
331 {
334 let cbs = self.callbacks.lock().clone();
335 for (_, cb) in cbs.iter() {
336 let cb = cb.clone();
337 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| cb(&event)));
338 }
339 }
340
341 {
343 let pools = self.pools.lock().clone();
344 for (_, bk, weak) in pools.iter() {
345 if *bk != backend {
346 continue;
347 }
348 if let Some(pool) = weak.upgrade() {
349 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| pool.evict_unused()));
350 }
351 }
352 }
353 }
354
355 if self.try_commit(backend, bytes) {
359 return Ok(BudgetReservation { budget: self.clone(), backend, bytes });
360 }
361
362 let available = self.available(backend);
364 Err(Error::OutOfGpuMemory { required_bytes: bytes, available_bytes: available, backend })
365 }
366
367 fn try_commit(&self, backend: BackendKind, bytes: u64) -> bool {
371 let counter = self.usage.get(backend);
372 let hard = self.caps.for_backend(backend).hard_cap_bytes;
373 loop {
374 let cur = counter.load(Ordering::Acquire);
375 let next = cur.saturating_add(bytes);
376 if let Some(h) = hard
377 && next > h
378 {
379 return false;
380 }
381 if counter.compare_exchange(cur, next, Ordering::AcqRel, Ordering::Acquire).is_ok() {
382 return true;
383 }
384 }
385 }
386
387 pub fn register_pool(&self, backend: BackendKind, pool: &Arc<dyn EvictablePool>) -> PoolHandle {
388 let h = PoolHandle(self.next_pool_handle.fetch_add(1, Ordering::Relaxed));
389 self.pools.lock().push((h, backend, Arc::downgrade(pool)));
390 h
391 }
392
393 pub fn unregister_pool(&self, handle: PoolHandle) {
394 self.pools.lock().retain(|(h, _, _)| *h != handle);
395 }
396
397 pub fn register_pressure_callback(&self, cb: PressureCallback) -> PressureCallbackHandle {
398 let h = PressureCallbackHandle(self.next_callback_handle.fetch_add(1, Ordering::Relaxed));
399 self.callbacks.lock().push((h, cb));
400 h
401 }
402
403 pub fn unregister_pressure_callback(&self, handle: PressureCallbackHandle) {
404 self.callbacks.lock().retain(|(h, _)| *h != handle);
405 }
406}