1use std::{
46 fmt,
47 ops::RangeInclusive,
48 sync::{
49 Arc, Mutex, OnceLock,
50 atomic::{AtomicBool, Ordering},
51 },
52};
53
54use wayland_backend::{
55 client::{Backend, InvalidId, ObjectData, ObjectId, WaylandError},
56 protocol::{Interface, OwnedMessage},
57};
58
59use crate::{
60 Connection, Dispatch, Proxy, QueueHandle,
61 protocol::{wl_display, wl_fixes, wl_registry},
62};
63
64pub trait GlobalListHandler: Sized {
67 fn runtime_add_global(
71 &mut self,
72 _globals: &GlobalList,
73 _conn: &Connection,
74 _qh: &QueueHandle<Self>,
75 _global: &Global,
76 ) {
77 }
78
79 fn runtime_remove_global(
83 &mut self,
84 _globals: &GlobalList,
85 _conn: &Connection,
86 _qh: &QueueHandle<Self>,
87 _global: &Global,
88 ) {
89 }
90}
91
92#[derive(Clone, Debug)]
96pub struct GlobalList {
97 registry: wl_registry::WlRegistry,
98}
99
100impl GlobalList {
101 pub fn init<State>(
105 conn: &Connection,
106 qh: &QueueHandle<State>,
107 ) -> Result<GlobalList, GlobalError>
108 where
109 State: GlobalListHandler + 'static,
110 {
111 let display = conn.display();
112 let fixes = OnceLock::<wl_fixes::WlFixes>::new();
113
114 let data = Arc::new(RegistryState {
115 data: GlobalListData { contents: Default::default(), fixes },
116 handle: qh.clone(),
117 initial_roundtrip_done: AtomicBool::new(false),
118 });
119 let registry =
120 display.send_constructor(wl_display::Request::GetRegistry {}, data.clone())?;
121 conn.roundtrip()?;
123 data.initial_roundtrip_done.store(true, Ordering::Relaxed);
124 Ok(GlobalList { registry })
125 }
126
127 fn data(&self) -> &GlobalListData {
128 self.registry.data::<GlobalListData>().unwrap()
129 }
130
131 pub fn clone_list(&self) -> Vec<Global> {
133 self.data().contents.lock().unwrap().clone()
134 }
135
136 pub fn bind_singleton<I, State, U>(
157 &self,
158 version: RangeInclusive<u32>,
159 qh: &QueueHandle<State>,
160 udata: U,
161 ) -> Result<I, BindError>
162 where
163 I: Proxy + 'static,
164 State: 'static,
165 U: Dispatch<I, State> + Send + Sync + 'static,
166 {
167 let interface = I::interface();
168 assert_valid_interface_version(&version, interface);
169
170 let guard = self.data().contents.lock().unwrap();
171 let global = guard
172 .iter()
173 .find(|Global { interface: interface_name, .. }| interface.name == interface_name)
175 .ok_or(BindError::NotPresent(interface.name))?;
176
177 self.bind_inner(global, version, qh, udata)
178 }
179
180 pub fn bind_all<I, State, U>(
186 &self,
187 version: std::ops::RangeInclusive<u32>,
188 qh: &QueueHandle<State>,
189 mut make_udata: impl FnMut(&Global) -> U,
190 ) -> Result<Vec<I>, BindError>
191 where
192 I: Proxy + 'static,
193 State: 'static,
194 U: Dispatch<I, State> + Send + Sync + 'static,
195 {
196 let interface = I::interface();
197 assert_valid_interface_version(&version, interface);
198
199 let guard = self.data().contents.lock().unwrap();
200 guard
201 .iter()
202 .filter(|global| global.interface == interface.name)
203 .map(|global| self.bind_inner(global, version.clone(), qh, make_udata(global)))
204 .collect()
205 }
206
207 pub fn bind_specific<I, State, U>(
214 &self,
215 name: u32,
216 version: std::ops::RangeInclusive<u32>,
217 qh: &QueueHandle<State>,
218 udata: U,
219 ) -> Result<I, BindError>
220 where
221 I: Proxy + 'static,
222 State: 'static,
223 U: Dispatch<I, State> + Send + Sync + 'static,
224 {
225 let interface = I::interface();
226 assert_valid_interface_version(&version, interface);
227
228 let guard = self.data().contents.lock().unwrap();
229 let global = guard
230 .iter()
231 .rev()
233 .find(|global| global.name == name && global.interface == interface.name)
235 .ok_or(BindError::NotPresent(interface.name))?;
237
238 self.bind_inner(global, version, qh, udata)
239 }
240
241 fn bind_inner<I, State, U>(
242 &self,
243 global: &Global,
244 version: RangeInclusive<u32>,
245 qh: &QueueHandle<State>,
246 udata: U,
247 ) -> Result<I, BindError>
248 where
249 I: Proxy + 'static,
250 State: 'static,
251 U: Dispatch<I, State> + Send + Sync + 'static,
252 {
253 if *version.start() > global.version {
255 return Err(BindError::UnsupportedVersion {
256 interface: I::interface().name,
257 requested: *version.start(),
258 available: global.version,
259 });
260 }
261
262 let negotiated_version = global.version.min(*version.end());
265
266 Ok(self.registry.bind(global.name, negotiated_version, qh, udata))
267 }
268
269 pub fn registry(&self) -> &wl_registry::WlRegistry {
273 &self.registry
274 }
275
276 pub fn destroy(self) {
284 if let Some(fixes) = self.data().fixes.get() {
285 let id = self.registry.id();
286 fixes.destroy_registry(&self.registry);
287 if let Some(backend) = fixes.backend().upgrade() {
288 backend.destroy_object(id).unwrap();
289 }
290 fixes.destroy();
291 }
292 }
293}
294
295#[derive(Debug)]
297pub enum GlobalError {
298 Backend(WaylandError),
300
301 InvalidId(InvalidId),
303}
304
305impl std::error::Error for GlobalError {
306 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
307 match self {
308 GlobalError::Backend(source) => Some(source),
309 GlobalError::InvalidId(source) => std::error::Error::source(source),
310 }
311 }
312}
313
314impl std::fmt::Display for GlobalError {
315 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
316 match self {
317 GlobalError::Backend(source) => {
318 write!(f, "Backend error: {source}")
319 }
320 GlobalError::InvalidId(source) => write!(f, "{source}"),
321 }
322 }
323}
324
325impl From<WaylandError> for GlobalError {
326 fn from(source: WaylandError) -> Self {
327 GlobalError::Backend(source)
328 }
329}
330
331impl From<InvalidId> for GlobalError {
332 fn from(source: InvalidId) -> Self {
333 GlobalError::InvalidId(source)
334 }
335}
336
337#[derive(Debug)]
339pub enum BindError {
340 UnsupportedVersion {
342 interface: &'static str,
344 requested: u32,
346 available: u32,
348 },
349
350 NotPresent(&'static str),
352}
353
354impl std::error::Error for BindError {}
355
356impl fmt::Display for BindError {
357 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
358 match self {
359 BindError::UnsupportedVersion { interface, requested, available } => {
360 write!(
361 f,
362 "the requested version `{requested}` of the global `{interface}` is not supported, only `{available}` is available"
363 )
364 }
365 BindError::NotPresent(name) => {
366 write!(f, "the requested global `{name}` was not found in the registry")
367 }
368 }
369 }
370}
371
372#[derive(Debug, Clone, PartialEq, Eq)]
374pub struct Global {
375 pub name: u32,
379 pub interface: String,
383 pub version: u32,
388}
389
390#[derive(Debug)]
391struct GlobalListData {
392 contents: Mutex<Vec<Global>>,
393 fixes: OnceLock<wl_fixes::WlFixes>,
394}
395
396impl GlobalListData {
397 fn add(&self, global: Global) {
398 self.contents.lock().unwrap().push(global);
399 }
400
401 fn remove(&self, name: u32) -> Option<Global> {
402 let mut guard = self.contents.lock().unwrap();
403 let idx = guard.iter().position(|i| i.name == name)?;
404 Some(guard.remove(idx))
405 }
406}
407
408impl<D> Dispatch<wl_registry::WlRegistry, D> for GlobalListData
409where
410 D: GlobalListHandler,
411{
412 fn event(
413 &self,
414 state: &mut D,
415 registry: &wl_registry::WlRegistry,
416 event: wl_registry::Event,
417 conn: &Connection,
418 qh: &QueueHandle<D>,
419 ) {
420 let globals = GlobalList { registry: registry.clone() };
421 match event {
422 wl_registry::Event::Global { name, interface, version } => {
423 let global = Global { name, interface, version };
424 self.add(global.clone());
425 state.runtime_add_global(&globals, conn, qh, &global);
426 }
427 wl_registry::Event::GlobalRemove { name } => {
428 if let Some(global) = self.remove(name) {
429 state.runtime_remove_global(&globals, conn, qh, &global);
430 }
431 }
432 }
433 }
434}
435
436struct RegistryState<State> {
437 data: GlobalListData,
438 handle: QueueHandle<State>,
439 initial_roundtrip_done: AtomicBool,
440}
441
442impl<State> ObjectData for RegistryState<State>
443where
444 State: GlobalListHandler + 'static,
445{
446 fn event(
447 self: Arc<Self>,
448 backend: &Backend,
449 msg: OwnedMessage<ObjectId>,
450 ) -> Option<Arc<dyn ObjectData>> {
451 if !self.initial_roundtrip_done.load(Ordering::Relaxed) {
455 let conn = Connection::from_backend(backend.clone());
456 if let Ok((registry, event)) = wl_registry::WlRegistry::parse_event(&conn, msg) {
458 match event {
459 wl_registry::Event::Global { name, interface, version } => {
460 let wl_fixes_ver = 1u32..=1;
461 if interface == "wl_fixes" && version >= *wl_fixes_ver.start() {
462 let _ = self.data.fixes.set(registry.bind(
463 name,
464 version.min(*wl_fixes_ver.end()),
465 &self.handle,
466 crate::Noop,
467 ));
468 }
469
470 self.data.add(Global { name, interface, version });
471 }
472
473 wl_registry::Event::GlobalRemove { name: remove } => {
474 self.data.remove(remove);
475 }
476 }
477 };
478 } else {
479 self.handle
481 .inner
482 .lock()
483 .unwrap()
484 .enqueue_event::<wl_registry::WlRegistry, GlobalListData>(msg, self.clone())
485 }
486
487 None
489 }
490
491 fn destroyed(&self, _id: &ObjectId) {}
492
493 fn data_as_any(&self) -> &dyn std::any::Any {
494 &self.data
495 }
496}
497
498fn assert_valid_interface_version(version: &RangeInclusive<u32>, interface: &'static Interface) {
499 if *version.end() > interface.version {
500 panic!(
502 "Maximum version ({}) of {} was higher than the proxy's maximum version ({}); outdated wayland XML files?",
503 version.end(),
504 interface.name,
505 interface.version
506 );
507 }
508}