1use alloc::{boxed::Box, collections::BTreeMap, string::String, sync::Arc, vec::Vec};
2use core::{
3 any::Any,
4 fmt::{Debug, Display},
5};
6
7use ax_sync::SpinLock;
8use usb_if::{
9 descriptor::{
10 ConfigurationDescriptor, DescriptorType, DeviceDescriptor, InterfaceDescriptor, LanguageId,
11 decode_string_descriptor,
12 },
13 err::{TransferError, USBError},
14 host::ControlSetup,
15};
16
17use crate::backend::ty::{DeviceInfoOp, DeviceOp, ep::EndpointHandle};
18
19pub struct DeviceInfo {
20 pub(crate) inner: Box<dyn DeviceInfoOp>,
21}
22
23pub struct HubDeviceInfo {
24 pub(crate) inner: Box<dyn DeviceInfoOp>,
25}
26
27pub enum ProbedDevice {
28 Device(DeviceInfo),
29 Hub(HubDeviceInfo),
30}
31
32pub struct ProbeChanges {
33 pub connected: Vec<ProbedDevice>,
34 pub disconnected: Vec<usize>,
35}
36
37impl ProbedDevice {
38 pub fn id(&self) -> usize {
39 match self {
40 Self::Device(info) => info.id(),
41 Self::Hub(info) => info.id(),
42 }
43 }
44
45 pub fn descriptor(&self) -> &DeviceDescriptor {
46 match self {
47 Self::Device(info) => info.descriptor(),
48 Self::Hub(info) => info.descriptor(),
49 }
50 }
51
52 pub fn configurations(&self) -> &[ConfigurationDescriptor] {
53 match self {
54 Self::Device(info) => info.configurations(),
55 Self::Hub(info) => info.configurations(),
56 }
57 }
58
59 pub fn product_id(&self) -> u16 {
60 self.descriptor().product_id
61 }
62
63 pub fn vendor_id(&self) -> u16 {
64 self.descriptor().vendor_id
65 }
66
67 pub fn as_device_info(&self) -> Option<&DeviceInfo> {
68 match self {
69 Self::Device(info) => Some(info),
70 Self::Hub(_) => None,
71 }
72 }
73
74 pub fn into_device_info(self) -> Option<DeviceInfo> {
75 match self {
76 Self::Device(info) => Some(info),
77 Self::Hub(_) => None,
78 }
79 }
80}
81
82impl Debug for ProbedDevice {
83 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
84 match self {
85 Self::Device(info) => f.debug_tuple("ProbedDevice::Device").field(info).finish(),
86 Self::Hub(info) => f.debug_tuple("ProbedDevice::Hub").field(info).finish(),
87 }
88 }
89}
90
91impl Display for ProbedDevice {
92 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
93 match self {
94 Self::Device(info) => Display::fmt(info, f),
95 Self::Hub(info) => Display::fmt(info, f),
96 }
97 }
98}
99
100impl DeviceInfo {
101 pub fn id(&self) -> usize {
102 self.inner.id()
103 }
104
105 pub fn descriptor(&self) -> &DeviceDescriptor {
106 self.inner.descriptor()
107 }
108
109 pub fn configurations(&self) -> &[ConfigurationDescriptor] {
110 self.inner.configuration_descriptors()
111 }
112
113 pub fn interface_descriptors<'a>(
114 &'a self,
115 ) -> impl Iterator<Item = &'a InterfaceDescriptor> + 'a {
116 self.configurations().iter().flat_map(|config| {
117 config
118 .interfaces
119 .iter()
120 .flat_map(|interface| interface.alt_settings.first())
121 })
122 }
123
124 pub fn product_id(&self) -> u16 {
125 self.descriptor().product_id
126 }
127
128 pub fn vendor_id(&self) -> u16 {
129 self.descriptor().vendor_id
130 }
131}
132
133impl HubDeviceInfo {
134 pub fn id(&self) -> usize {
135 self.inner.id()
136 }
137
138 pub fn descriptor(&self) -> &DeviceDescriptor {
139 self.inner.descriptor()
140 }
141
142 pub fn configurations(&self) -> &[ConfigurationDescriptor] {
143 self.inner.configuration_descriptors()
144 }
145
146 pub fn product_id(&self) -> u16 {
147 self.descriptor().product_id
148 }
149
150 pub fn vendor_id(&self) -> u16 {
151 self.descriptor().vendor_id
152 }
153}
154
155impl Debug for DeviceInfo {
156 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
157 f.debug_struct("DeviceInfo")
158 .field("backend", &self.inner.backend_name())
159 .field("vender_id", &self.inner.descriptor().vendor_id)
160 .field("product_id", &self.inner.descriptor().product_id)
161 .finish()
162 }
163}
164
165impl Debug for HubDeviceInfo {
166 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
167 f.debug_struct("HubDeviceInfo")
168 .field("backend", &self.inner.backend_name())
169 .field("vender_id", &self.inner.descriptor().vendor_id)
170 .field("product_id", &self.inner.descriptor().product_id)
171 .finish()
172 }
173}
174
175impl Display for DeviceInfo {
176 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
177 write!(
178 f,
179 "{:04x}:{:04x}",
180 self.inner.descriptor().vendor_id,
181 self.inner.descriptor().product_id
182 )
183 }
184}
185
186impl Display for HubDeviceInfo {
187 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
188 write!(
189 f,
190 "{:04x}:{:04x}",
191 self.inner.descriptor().vendor_id,
192 self.inner.descriptor().product_id
193 )
194 }
195}
196
197pub struct Device {
198 pub(crate) inner: Box<dyn DeviceOp>,
199 lang_id: LanguageId,
200 manufacturer: Option<String>,
201 claimed_interfaces: BTreeMap<u8, InterfaceRegistration>,
202 lifecycle: DeviceLifecycle,
203}
204
205#[derive(Clone, Copy, Debug, Eq, PartialEq)]
206enum InterfaceSessionState {
207 Active,
208 Released,
209 Broken,
210 Disconnected,
211}
212
213#[derive(Clone, Copy, Debug, Eq, PartialEq)]
214enum DeviceLifecycle {
215 Active,
216 Broken,
217 Disconnected,
218}
219
220struct InterfaceRegistration {
221 alternate: u8,
222 endpoints: BTreeMap<u8, EndpointHandle>,
223 state: Arc<SpinLock<InterfaceSessionState>>,
224}
225
226pub struct InterfaceSession {
228 interface: u8,
229 alternate: u8,
230 endpoints: BTreeMap<u8, EndpointHandle>,
231 state: Arc<SpinLock<InterfaceSessionState>>,
232}
233
234impl InterfaceSession {
235 pub fn interface_number(&self) -> u8 {
236 self.interface
237 }
238
239 pub fn alternate_setting(&self) -> u8 {
240 self.alternate
241 }
242
243 pub fn endpoint(&self, address: u8) -> Result<EndpointHandle, USBError> {
245 match *self.state.lock() {
246 InterfaceSessionState::Active => {}
247 InterfaceSessionState::Released => return Err(USBError::NotFound),
248 InterfaceSessionState::Broken => return Err(USBError::InterfaceBroken),
249 InterfaceSessionState::Disconnected => {
250 return Err(TransferError::Disconnected.into());
251 }
252 }
253 self.endpoints
254 .get(&address)
255 .cloned()
256 .ok_or(USBError::NotFound)
257 }
258
259 pub async fn set_alternate(
261 &mut self,
262 device: &mut Device,
263 alternate: u8,
264 ) -> Result<(), USBError> {
265 self.ensure_active()?;
266 device.ensure_active()?;
267 self.ensure_owned_by(device)?;
268 if self.alternate == alternate {
269 return Ok(());
270 }
271
272 let old_endpoints = self.endpoints.clone();
273 for endpoint in old_endpoints.values() {
274 endpoint.revoke();
275 }
276 match device
277 .inner
278 .claim_interface(self.interface, alternate)
279 .await
280 {
281 Ok(endpoints) => {
282 let Some(registration) = device.claimed_interfaces.get_mut(&self.interface) else {
283 for endpoint in endpoints.values() {
284 endpoint.revoke();
285 }
286 *self.state.lock() = InterfaceSessionState::Broken;
287 return Err(USBError::InterfaceBroken);
288 };
289 if !Arc::ptr_eq(®istration.state, &self.state) {
290 for endpoint in endpoints.values() {
291 endpoint.revoke();
292 }
293 *self.state.lock() = InterfaceSessionState::Broken;
294 return Err(USBError::InterfaceBroken);
295 }
296 registration.alternate = alternate;
297 registration.endpoints = endpoints.clone();
298 self.alternate = alternate;
299 self.endpoints = endpoints;
300 Ok(())
301 }
302 Err(USBError::TransferError(TransferError::Disconnected)) => {
303 for endpoint in old_endpoints.values() {
304 endpoint.disconnect();
305 }
306 *self.state.lock() = InterfaceSessionState::Disconnected;
307 device.lifecycle = DeviceLifecycle::Disconnected;
308 Err(TransferError::Disconnected.into())
309 }
310 Err(err) => {
311 if matches!(err, USBError::InterfaceBroken) {
312 *self.state.lock() = InterfaceSessionState::Broken;
313 device.lifecycle = DeviceLifecycle::Broken;
314 } else {
315 for endpoint in old_endpoints.values() {
316 endpoint.reactivate();
317 }
318 }
319 Err(err)
320 }
321 }
322 }
323
324 pub async fn release(&mut self, device: &mut Device) -> Result<(), USBError> {
326 self.ensure_active()?;
327 device.ensure_active()?;
328 self.ensure_owned_by(device)?;
329 for endpoint in self.endpoints.values() {
330 endpoint.revoke();
331 }
332 match device.inner.release_interface(self.interface).await {
333 Ok(()) => {
334 let owns_registration = device
335 .claimed_interfaces
336 .get(&self.interface)
337 .is_some_and(|entry| Arc::ptr_eq(&entry.state, &self.state));
338 if owns_registration {
339 device.claimed_interfaces.remove(&self.interface);
340 }
341 self.endpoints.clear();
342 *self.state.lock() = InterfaceSessionState::Released;
343 Ok(())
344 }
345 Err(USBError::TransferError(TransferError::Disconnected)) => {
346 for endpoint in self.endpoints.values() {
347 endpoint.disconnect();
348 }
349 *self.state.lock() = InterfaceSessionState::Disconnected;
350 device.lifecycle = DeviceLifecycle::Disconnected;
351 Err(TransferError::Disconnected.into())
352 }
353 Err(err) => {
354 if matches!(err, USBError::InterfaceBroken) {
355 *self.state.lock() = InterfaceSessionState::Broken;
356 device.lifecycle = DeviceLifecycle::Broken;
357 } else {
358 for endpoint in self.endpoints.values() {
359 endpoint.reactivate();
360 }
361 }
362 Err(err)
363 }
364 }
365 }
366
367 fn ensure_active(&self) -> Result<(), USBError> {
368 match *self.state.lock() {
369 InterfaceSessionState::Active => Ok(()),
370 InterfaceSessionState::Released => Err(USBError::InvalidParameter),
371 InterfaceSessionState::Broken => Err(USBError::InterfaceBroken),
372 InterfaceSessionState::Disconnected => Err(TransferError::Disconnected.into()),
373 }
374 }
375
376 fn ensure_owned_by(&self, device: &Device) -> Result<(), USBError> {
377 let Some(registration) = device.claimed_interfaces.get(&self.interface) else {
378 return Err(USBError::InvalidParameter);
379 };
380 if Arc::ptr_eq(®istration.state, &self.state) && registration.alternate == self.alternate
381 {
382 Ok(())
383 } else {
384 Err(USBError::InvalidParameter)
385 }
386 }
387}
388
389impl Debug for Device {
390 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
391 f.debug_struct("Device")
392 .field("backend", &self.inner.backend_name())
393 .field("vender_id", &self.inner.descriptor().vendor_id)
394 .field("product_id", &self.inner.descriptor().product_id)
395 .finish()
396 }
397}
398
399impl<T: DeviceOp> From<T> for Device {
400 fn from(inner: T) -> Self {
401 Self {
402 inner: Box::new(inner),
403 claimed_interfaces: BTreeMap::new(),
404 lifecycle: DeviceLifecycle::Active,
405 lang_id: LanguageId::default(),
406 manufacturer: None,
407 }
408 }
409}
410
411impl From<Box<dyn DeviceOp>> for Device {
412 fn from(inner: Box<dyn DeviceOp>) -> Self {
413 Self {
414 inner,
415 claimed_interfaces: BTreeMap::new(),
416 lifecycle: DeviceLifecycle::Active,
417 lang_id: LanguageId::default(),
418 manufacturer: None,
419 }
420 }
421}
422
423impl Device {
424 pub(crate) async fn init(&mut self) -> Result<(), USBError> {
425 self.manufacturer = self.read_manufacturer().await;
426 Ok(())
427 }
428
429 pub fn product_id(&self) -> u16 {
430 self.descriptor().product_id
431 }
432
433 pub fn vendor_id(&self) -> u16 {
434 self.descriptor().vendor_id
435 }
436
437 pub fn slot_id(&self) -> u8 {
438 self.inner.id() as _
439 }
440
441 pub async fn claim_interface(
442 &mut self,
443 interface: u8,
444 alternate: u8,
445 ) -> Result<InterfaceSession, USBError> {
446 self.ensure_active()?;
447 trace!("Claiming interface {interface}, alternate {alternate}");
448 if self.claimed_interfaces.contains_key(&interface) {
449 return Err(USBError::InvalidParameter);
450 }
451 let endpoints = match self.inner.claim_interface(interface, alternate).await {
452 Ok(endpoints) => endpoints,
453 Err(USBError::InterfaceBroken) => {
454 self.lifecycle = DeviceLifecycle::Broken;
455 return Err(USBError::InterfaceBroken);
456 }
457 Err(USBError::TransferError(TransferError::Disconnected)) => {
458 self.ctrl_ep_ref().disconnect();
459 self.lifecycle = DeviceLifecycle::Disconnected;
460 return Err(TransferError::Disconnected.into());
461 }
462 Err(err) => return Err(err),
463 };
464 let state = Arc::new(SpinLock::new(InterfaceSessionState::Active));
465 self.claimed_interfaces.insert(
466 interface,
467 InterfaceRegistration {
468 alternate,
469 endpoints: endpoints.clone(),
470 state: state.clone(),
471 },
472 );
473 Ok(InterfaceSession {
474 interface,
475 alternate,
476 endpoints,
477 state,
478 })
479 }
480
481 pub fn descriptor(&self) -> &DeviceDescriptor {
482 self.inner.descriptor()
483 }
484
485 pub fn configurations(&self) -> &[ConfigurationDescriptor] {
486 self.inner.configuration_descriptors()
487 }
488
489 pub fn manufacturer(&self) -> Option<&str> {
490 self.manufacturer.as_deref()
491 }
492
493 pub async fn set_configuration(&mut self, configuration_value: u8) -> crate::err::Result {
494 self.ensure_active()?;
495 let registrations = self
496 .claimed_interfaces
497 .values()
498 .map(|entry| (entry.state.clone(), entry.endpoints.clone()))
499 .collect::<Vec<_>>();
500 for (_, endpoints) in ®istrations {
501 for endpoint in endpoints.values() {
502 endpoint.revoke();
503 }
504 }
505 let result = self.inner.set_configuration(configuration_value).await;
506 match &result {
507 Ok(()) => {
508 for (state, _) in registrations {
509 *state.lock() = InterfaceSessionState::Released;
510 }
511 self.claimed_interfaces.clear();
512 }
513 Err(USBError::InterfaceBroken) => {
514 for (state, _) in registrations {
515 *state.lock() = InterfaceSessionState::Broken;
516 }
517 self.lifecycle = DeviceLifecycle::Broken;
518 }
519 Err(USBError::TransferError(TransferError::Disconnected)) => {
520 for (state, endpoints) in registrations {
521 for endpoint in endpoints.values() {
522 endpoint.disconnect();
523 }
524 *state.lock() = InterfaceSessionState::Disconnected;
525 }
526 self.ctrl_ep_ref().disconnect();
527 self.lifecycle = DeviceLifecycle::Disconnected;
528 }
529 Err(_) => {
530 for (_, endpoints) in registrations {
531 for endpoint in endpoints.values() {
532 endpoint.reactivate();
533 }
534 }
535 }
536 }
537 result
538 }
539
540 pub async fn disconnect(&mut self) -> crate::err::Result {
542 match self.lifecycle {
543 DeviceLifecycle::Disconnected => return Ok(()),
544 DeviceLifecycle::Broken => return Err(USBError::InterfaceBroken),
545 DeviceLifecycle::Active => {}
546 }
547
548 let registrations = self
549 .claimed_interfaces
550 .values()
551 .map(|entry| (entry.state.clone(), entry.endpoints.clone()))
552 .collect::<Vec<_>>();
553 for (_, endpoints) in ®istrations {
554 for endpoint in endpoints.values() {
555 endpoint.revoke();
556 }
557 }
558 self.ctrl_ep_ref().revoke();
559
560 match self.inner.disconnect().await {
561 Ok(()) => {
562 for (state, endpoints) in registrations {
563 for endpoint in endpoints.values() {
564 endpoint.disconnect();
565 }
566 *state.lock() = InterfaceSessionState::Disconnected;
567 }
568 self.ctrl_ep_ref().disconnect();
569 self.claimed_interfaces.clear();
570 self.lifecycle = DeviceLifecycle::Disconnected;
571 Ok(())
572 }
573 Err(err) => {
574 for (state, _) in registrations {
575 *state.lock() = InterfaceSessionState::Broken;
576 }
577 self.lifecycle = DeviceLifecycle::Broken;
578 Err(err)
579 }
580 }
581 }
582
583 pub fn ctrl_ep_ref(&self) -> &EndpointHandle {
584 self.inner.ctrl_ep_ref()
585 }
586
587 pub fn ctrl_ep_mut(&mut self) -> &mut EndpointHandle {
588 self.inner.ctrl_ep_mut()
589 }
590
591 fn ensure_active(&self) -> Result<(), USBError> {
592 match self.lifecycle {
593 DeviceLifecycle::Active => Ok(()),
594 DeviceLifecycle::Broken => Err(USBError::InterfaceBroken),
595 DeviceLifecycle::Disconnected => Err(TransferError::Disconnected.into()),
596 }
597 }
598
599 async fn read_manufacturer(&mut self) -> Option<String> {
600 let idx = self.descriptor().manufacturer_string_index?;
601 self.string_descriptor(idx.get()).await.ok()
602 }
603
604 pub fn lang_id(&self) -> LanguageId {
605 self.lang_id
606 }
607
608 pub fn set_lang_id(&mut self, lang_id: LanguageId) {
609 self.lang_id = lang_id;
610 }
611
612 pub async fn string_descriptor(&mut self, index: u8) -> Result<String, USBError> {
613 let mut data = alloc::vec![0u8; 256];
614 let lang_id = self.lang_id();
615 let len = self
616 .ctrl_ep_mut()
617 .get_descriptor(DescriptorType::STRING, index, lang_id.into(), &mut data)
618 .await?;
619 let descriptor_len = data
620 .first()
621 .copied()
622 .map(usize::from)
623 .unwrap_or(0)
624 .min(len)
625 .min(data.len());
626 decode_string_descriptor(&data[..descriptor_len]).map_err(USBError::from)
627 }
628
629 pub async fn control_in(
630 &mut self,
631 param: ControlSetup,
632 buff: &mut [u8],
633 ) -> Result<usize, TransferError> {
634 self.ctrl_ep_mut().control_in(param, buff).await
635 }
636
637 pub async fn control_out(
638 &mut self,
639 param: ControlSetup,
640 buff: &[u8],
641 ) -> Result<usize, TransferError> {
642 self.ctrl_ep_mut().control_out(param, buff).await
643 }
644
645 pub async fn update_hub(
646 &mut self,
647 params: crate::backend::ty::HubParams,
648 ) -> Result<(), USBError> {
649 self.inner.update_hub(params).await
650 }
651
652 pub async fn current_configuration_descriptor(
653 &mut self,
654 ) -> Result<ConfigurationDescriptor, USBError> {
655 let value = self.ctrl_ep_mut().get_configuration().await?;
656 if value == 0 {
657 return Err(USBError::NotFound);
658 }
659 for config in self.configurations() {
660 if config.configuration_value == value {
661 return Ok(config.clone());
662 }
663 }
664 Err(USBError::NotFound)
665 }
666
667 #[allow(unused)]
668 pub(crate) fn as_raw<T: DeviceOp>(&self) -> &T {
669 (self.inner.as_ref() as &dyn Any)
670 .downcast_ref::<T>()
671 .unwrap()
672 }
673
674 #[allow(unused)]
675 pub(crate) fn as_raw_mut<T: DeviceOp>(&mut self) -> &mut T {
676 (self.inner.as_mut() as &mut dyn Any)
677 .downcast_mut::<T>()
678 .unwrap()
679 }
680}
681
682impl Display for Device {
683 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
684 write!(
685 f,
686 "{:04x}:{:04x}",
687 self.inner.descriptor().vendor_id,
688 self.inner.descriptor().product_id
689 )
690 }
691}
692
693#[cfg(test)]
694mod tests {
695 use alloc::{collections::BTreeMap, vec::Vec};
696 use core::{
697 future::Future,
698 pin::Pin,
699 ptr,
700 task::{Context, Poll, RawWaker, RawWakerVTable, Waker},
701 };
702
703 use futures::{FutureExt, future::BoxFuture};
704 use usb_if::{
705 descriptor::{ConfigurationDescriptor, DeviceDescriptor, EndpointType},
706 endpoint::{EndpointAddress, EndpointInfo, RequestId, TransferCompletion, TransferRequest},
707 transfer::Direction,
708 };
709
710 use super::{Device, InterfaceSession};
711 use crate::{
712 backend::ty::{
713 DeviceOp, HubParams,
714 ep::{EndpointHandle, EndpointOp},
715 },
716 err::{TransferError, USBError},
717 };
718
719 struct TestEndpoint;
720
721 impl EndpointOp for TestEndpoint {
722 fn submit_request(
723 &mut self,
724 _request: TransferRequest,
725 ) -> Result<RequestId, TransferError> {
726 Ok(RequestId::new(1))
727 }
728
729 fn reclaim_request(
730 &mut self,
731 _id: RequestId,
732 ) -> Option<Result<TransferCompletion, TransferError>> {
733 None
734 }
735
736 fn register_waker(&self, _id: RequestId, _cx: &mut Context<'_>) {}
737 }
738
739 struct TestDevice {
740 descriptor: DeviceDescriptor,
741 configurations: Vec<ConfigurationDescriptor>,
742 control: EndpointHandle,
743 rejected_alternate: Option<u8>,
744 }
745
746 impl TestDevice {
747 fn new(rejected_alternate: Option<u8>) -> Self {
748 Self {
749 descriptor: DeviceDescriptor {
750 usb_version: 0x0200,
751 class: 0,
752 subclass: 0,
753 protocol: 0,
754 max_packet_size_0: 64,
755 vendor_id: 1,
756 product_id: 1,
757 device_version: 0x0100,
758 manufacturer_string_index: None,
759 product_string_index: None,
760 serial_number_string_index: None,
761 num_configurations: 0,
762 },
763 configurations: Vec::new(),
764 control: EndpointHandle::new(EndpointInfo::control(), TestEndpoint),
765 rejected_alternate,
766 }
767 }
768
769 fn data_endpoint() -> EndpointHandle {
770 EndpointHandle::new(
771 EndpointInfo {
772 address: EndpointAddress::new(1),
773 transfer_type: EndpointType::Bulk,
774 direction: Direction::Out,
775 max_packet_size: 512,
776 packets_per_microframe: 1,
777 interval: 0,
778 },
779 TestEndpoint,
780 )
781 }
782 }
783
784 impl DeviceOp for TestDevice {
785 fn id(&self) -> usize {
786 1
787 }
788
789 fn backend_name(&self) -> &str {
790 "test"
791 }
792
793 fn descriptor(&self) -> &DeviceDescriptor {
794 &self.descriptor
795 }
796
797 fn configuration_descriptors(&self) -> &[ConfigurationDescriptor] {
798 &self.configurations
799 }
800
801 fn ctrl_ep_ref(&self) -> &EndpointHandle {
802 &self.control
803 }
804
805 fn ctrl_ep_mut(&mut self) -> &mut EndpointHandle {
806 &mut self.control
807 }
808
809 fn claim_interface<'a>(
810 &'a mut self,
811 _interface: u8,
812 alternate: u8,
813 ) -> BoxFuture<'a, Result<BTreeMap<u8, EndpointHandle>, USBError>> {
814 let rejected = self.rejected_alternate == Some(alternate);
815 async move {
816 if rejected {
817 return Err(USBError::InvalidParameter);
818 }
819 Ok(BTreeMap::from([(1, Self::data_endpoint())]))
820 }
821 .boxed()
822 }
823
824 fn release_interface<'a>(
825 &'a mut self,
826 _interface: u8,
827 ) -> BoxFuture<'a, Result<(), USBError>> {
828 async { Ok(()) }.boxed()
829 }
830
831 fn set_configuration<'a>(
832 &'a mut self,
833 _configuration_value: u8,
834 ) -> BoxFuture<'a, Result<(), USBError>> {
835 async { Ok(()) }.boxed()
836 }
837
838 fn disconnect(&mut self) -> BoxFuture<'_, Result<(), USBError>> {
839 async { Ok(()) }.boxed()
840 }
841
842 fn update_hub(&mut self, _params: HubParams) -> BoxFuture<'_, Result<(), USBError>> {
843 async { Ok(()) }.boxed()
844 }
845 }
846
847 #[test]
848 fn alternate_commit_revokes_old_handle_and_publishes_new_endpoint() {
849 let (mut device, mut session) = claimed_session(None);
850 let old_endpoint = session.endpoint(1).unwrap();
851
852 block_on_ready(session.set_alternate(&mut device, 1)).unwrap();
853
854 assert!(matches!(
855 old_endpoint.submit(TransferRequest::bulk_out(&[])),
856 Err(TransferError::EndpointRevoked)
857 ));
858 assert!(
859 session
860 .endpoint(1)
861 .unwrap()
862 .submit(TransferRequest::bulk_out(&[]))
863 .is_ok()
864 );
865 }
866
867 #[test]
868 fn alternate_failure_reactivates_old_handle() {
869 let (mut device, mut session) = claimed_session(Some(2));
870 let old_endpoint = session.endpoint(1).unwrap();
871
872 assert!(matches!(
873 block_on_ready(session.set_alternate(&mut device, 2)),
874 Err(USBError::InvalidParameter)
875 ));
876
877 assert!(old_endpoint.submit(TransferRequest::bulk_out(&[])).is_ok());
878 assert_eq!(session.alternate_setting(), 0);
879 }
880
881 #[test]
882 fn session_rejects_a_different_device_before_freezing_endpoints() {
883 let (_owner, mut session) = claimed_session(None);
884 let endpoint = session.endpoint(1).unwrap();
885 let mut different_device = Device::from(TestDevice::new(None));
886
887 assert!(matches!(
888 block_on_ready(session.set_alternate(&mut different_device, 1)),
889 Err(USBError::InvalidParameter)
890 ));
891 assert!(endpoint.submit(TransferRequest::bulk_out(&[])).is_ok());
892 }
893
894 #[test]
895 fn disconnect_marks_session_and_old_handle_disconnected() {
896 let (mut device, session) = claimed_session(None);
897 let old_endpoint = session.endpoint(1).unwrap();
898
899 block_on_ready(device.disconnect()).unwrap();
900
901 assert!(matches!(
902 old_endpoint.submit(TransferRequest::bulk_out(&[])),
903 Err(TransferError::Disconnected)
904 ));
905 assert!(matches!(
906 session.endpoint(1),
907 Err(USBError::TransferError(TransferError::Disconnected))
908 ));
909 }
910
911 fn claimed_session(rejected_alternate: Option<u8>) -> (Device, InterfaceSession) {
912 let mut device = Device::from(TestDevice::new(rejected_alternate));
913 let session = block_on_ready(device.claim_interface(0, 0)).unwrap();
914 (device, session)
915 }
916
917 fn block_on_ready<F: Future>(mut future: F) -> F::Output {
918 let waker = noop_waker();
919 let mut context = Context::from_waker(&waker);
920 match unsafe { Pin::new_unchecked(&mut future) }.poll(&mut context) {
921 Poll::Ready(output) => output,
922 Poll::Pending => panic!("test future unexpectedly pending"),
923 }
924 }
925
926 fn noop_waker() -> Waker {
927 unsafe fn clone(_: *const ()) -> RawWaker {
928 RawWaker::new(ptr::null(), &VTABLE)
929 }
930 unsafe fn wake(_: *const ()) {}
931 unsafe fn wake_by_ref(_: *const ()) {}
932 unsafe fn drop(_: *const ()) {}
933
934 static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, wake, wake_by_ref, drop);
935
936 unsafe { Waker::from_raw(RawWaker::new(ptr::null(), &VTABLE)) }
937 }
938}