use crate::error::Result;
use crate::identity::DriveId;
use crate::scsi::ScsiTransport;
use std::sync::RwLock;
pub trait Unlocker: Send + Sync {
fn name(&self) -> &str;
fn matches(&self, id: &DriveId) -> bool;
fn unlock_drive(&self, scsi: &mut dyn ScsiTransport, id: &DriveId) -> Result<()>;
fn read_volume_id(
&self,
_scsi: &mut dyn ScsiTransport,
_id: &DriveId,
) -> Result<Option<[u8; 16]>> {
Ok(None)
}
fn set_max_read_speed(&self, _scsi: &mut dyn ScsiTransport, _id: &DriveId) -> Result<()> {
Ok(())
}
}
static REGISTRY: RwLock<Vec<Box<dyn Unlocker>>> = RwLock::new(Vec::new());
pub fn register_unlocker(u: Box<dyn Unlocker>) {
if let Ok(mut reg) = REGISTRY.write() {
reg.push(u);
}
}
pub(crate) fn route_unlock(scsi: &mut dyn ScsiTransport, id: &DriveId) -> Result<Option<String>> {
let reg = match REGISTRY.read() {
Ok(r) => r,
Err(_) => return Ok(None),
};
for u in reg.iter() {
if u.matches(id) {
let name = u.name().to_string();
u.unlock_drive(scsi, id)?;
return Ok(Some(name));
}
}
Ok(None)
}
pub(crate) fn unlocker_read_volume_id(
scsi: &mut dyn ScsiTransport,
id: &DriveId,
) -> Result<Option<[u8; 16]>> {
let reg = match REGISTRY.read() {
Ok(r) => r,
Err(_) => return Ok(None),
};
for u in reg.iter() {
if u.matches(id) {
return u.read_volume_id(scsi, id);
}
}
Ok(None)
}
pub(crate) fn unlocker_set_max_read_speed(
scsi: &mut dyn ScsiTransport,
id: &DriveId,
) -> Result<()> {
let reg = match REGISTRY.read() {
Ok(r) => r,
Err(_) => return Ok(()),
};
for u in reg.iter() {
if u.matches(id) {
return u.set_max_read_speed(scsi, id);
}
}
Ok(())
}
#[doc(hidden)]
pub fn registered_count() -> usize {
REGISTRY.read().map(|r| r.len()).unwrap_or(0)
}
pub(crate) fn matching_name(id: &DriveId) -> Option<String> {
let reg = REGISTRY.read().ok()?;
reg.iter()
.find(|u| u.matches(id))
.map(|u| u.name().to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::scsi::{DataDirection, ScsiResult, ScsiTransport};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
struct NoopTransport;
impl ScsiTransport for NoopTransport {
fn execute(
&mut self,
_cdb: &[u8],
_dir: DataDirection,
_data: &mut [u8],
_timeout_ms: u32,
) -> Result<ScsiResult> {
Ok(ScsiResult {
status: 0,
bytes_transferred: 0,
sense: [0u8; 32],
})
}
}
fn fake_id(vendor: &str) -> DriveId {
let mut inquiry = vec![0u8; 96];
let v = vendor.as_bytes();
inquiry[8..8 + v.len().min(8)].copy_from_slice(&v[..v.len().min(8)]);
DriveId::from_inquiry(&inquiry, "")
}
struct FakeUnlocker {
want_vendor: String,
ran: Arc<AtomicBool>,
vid: Option<[u8; 16]>,
vid_ran: Arc<AtomicBool>,
speed_ran: Arc<AtomicBool>,
}
impl FakeUnlocker {
fn new(vendor: &str, ran: Arc<AtomicBool>) -> Self {
Self {
want_vendor: vendor.into(),
ran,
vid: None,
vid_ran: Arc::new(AtomicBool::new(false)),
speed_ran: Arc::new(AtomicBool::new(false)),
}
}
fn with_vid(mut self, vid: Option<[u8; 16]>, vid_ran: Arc<AtomicBool>) -> Self {
self.vid = vid;
self.vid_ran = vid_ran;
self
}
fn with_speed(mut self, speed_ran: Arc<AtomicBool>) -> Self {
self.speed_ran = speed_ran;
self
}
}
impl Unlocker for FakeUnlocker {
fn name(&self) -> &str {
"fake"
}
fn matches(&self, id: &DriveId) -> bool {
id.vendor_id.trim() == self.want_vendor
}
fn unlock_drive(&self, _scsi: &mut dyn ScsiTransport, _id: &DriveId) -> Result<()> {
self.ran.store(true, Ordering::SeqCst);
Ok(())
}
fn read_volume_id(
&self,
_scsi: &mut dyn ScsiTransport,
_id: &DriveId,
) -> Result<Option<[u8; 16]>> {
self.vid_ran.store(true, Ordering::SeqCst);
Ok(self.vid)
}
fn set_max_read_speed(&self, _scsi: &mut dyn ScsiTransport, _id: &DriveId) -> Result<()> {
self.speed_ran.store(true, Ordering::SeqCst);
Ok(())
}
}
#[test]
fn registry_routes_match_else_oem() {
let ran = Arc::new(AtomicBool::new(false));
register_unlocker(Box::new(FakeUnlocker::new("MATCHVND", ran.clone())));
let mut scsi = NoopTransport;
let matched = route_unlock(&mut scsi, &fake_id("MATCHVND")).unwrap();
assert_eq!(matched.as_deref(), Some("fake"), "matching unlocker runs");
assert!(ran.load(Ordering::SeqCst), "unlock_drive() was invoked");
ran.store(false, Ordering::SeqCst);
let none = route_unlock(&mut scsi, &fake_id("OTHERVND")).unwrap();
assert!(none.is_none(), "no match → OEM/cert fallback");
assert!(
!ran.load(Ordering::SeqCst),
"unlock_drive() not invoked on no-match"
);
}
#[test]
fn unlocker_read_volume_id_routes_match_else_cert() {
let mut scsi = NoopTransport;
let vid = [0x5Au8; 16];
let vid_ran = Arc::new(AtomicBool::new(false));
register_unlocker(Box::new(
FakeUnlocker::new("VIDVNDOR", Arc::new(AtomicBool::new(false)))
.with_vid(Some(vid), vid_ran.clone()),
));
let got = unlocker_read_volume_id(&mut scsi, &fake_id("VIDVNDOR")).unwrap();
assert_eq!(got, Some(vid), "matching unlocker's OEM VID is used");
assert!(
vid_ran.load(Ordering::SeqCst),
"read_volume_id() was consulted"
);
let none_ran = Arc::new(AtomicBool::new(false));
register_unlocker(Box::new(
FakeUnlocker::new("NOVIDVND", Arc::new(AtomicBool::new(false)))
.with_vid(None, none_ran.clone()),
));
let got = unlocker_read_volume_id(&mut scsi, &fake_id("NOVIDVND")).unwrap();
assert!(
got.is_none(),
"unlocker without OEM VID falls through to cert"
);
assert!(
none_ran.load(Ordering::SeqCst),
"read_volume_id() consulted even when it returns None"
);
let got = unlocker_read_volume_id(&mut scsi, &fake_id("UNKNWNVD")).unwrap();
assert!(got.is_none(), "no match → cert fallback");
}
#[test]
fn unlocker_set_max_read_speed_routes_match_else_noop() {
let mut scsi = NoopTransport;
let speed_ran = Arc::new(AtomicBool::new(false));
register_unlocker(Box::new(
FakeUnlocker::new("SPEEDVND", Arc::new(AtomicBool::new(false)))
.with_speed(speed_ran.clone()),
));
unlocker_set_max_read_speed(&mut scsi, &fake_id("SPEEDVND")).unwrap();
assert!(
speed_ran.load(Ordering::SeqCst),
"set_max_read_speed() invoked on match"
);
speed_ran.store(false, Ordering::SeqCst);
unlocker_set_max_read_speed(&mut scsi, &fake_id("NOSPEEDV")).unwrap();
assert!(
!speed_ran.load(Ordering::SeqCst),
"no match → safe no-op, nothing invoked"
);
}
}