Skip to main content

emerald_hwkey/ledger/connect/
shared.rs

1use std::thread;
2use std::sync::mpsc::{Receiver, Sender, channel};
3use std::sync::{Arc, Mutex, OnceLock, RwLock};
4use crate::errors::HWKeyError;
5use crate::ledger::apdu::APDU;
6use crate::ledger::comm;
7use crate::ledger::comm::{sendrecv_timeout, LedgerTransport};
8use crate::ledger::connect::direct::{AppDetails, LedgerHidKey};
9use std::convert::TryFrom;
10use std::marker::PhantomData;
11use log;
12use crate::ledger::connect::LedgerKey;
13#[cfg(feature = "speculos")]
14use crate::ledger::connect::speculos::LedgerSpeculosKey;
15
16static INSTANCE: OnceLock<LedgerKeyShared<LedgerHidKey>> = OnceLock::new();
17
18#[cfg(feature = "speculos")]
19static INSTANCE_SPECULOS: OnceLock<LedgerKeyShared<LedgerSpeculosKey>> = OnceLock::new();
20
21pub struct LedgerKeyShared<LK: LedgerKey> {
22    t: PhantomData<LK>,
23    channel: Sender<Command>,
24    state: Arc<RwLock<State>>,
25}
26
27impl<LK: LedgerKey> Clone for LedgerKeyShared<LK> {
28    fn clone(&self) -> Self {
29        Self { t: self.t, channel: self.channel.clone(), state: self.state.clone() }
30    }
31}
32
33enum Command {
34    Connect(Sender<Result<(), HWKeyError>>),
35    HaveDevice(Sender<bool>),
36    Write(Vec<u8>, Sender<Result<usize, HWKeyError>>),
37    Read(i32, Sender<Result<Vec<u8>, HWKeyError>>),
38    SetDisconnected,
39}
40
41#[derive(Clone, Eq, PartialEq)]
42enum State {
43    Init,
44    Working,
45    Disconnected,
46    Stopped,
47}
48
49impl LedgerKeyShared<LedgerHidKey> {
50    ///
51    /// Get (or create on the first call) a shared instance to access LedgerKey
52    pub fn instance() -> Result<LedgerKeyShared<LedgerHidKey>, HWKeyError> {
53        let instance = INSTANCE.get_or_init(|| {
54            let mut value = Self::new();
55            let connected = value.connect();
56            if connected.is_ok() {
57                let ping = comm::ping(&value);
58                if ping.is_err() {
59                    log::warn!("No immediate response from Ledger: {:?}. Will try to reconnect", ping);
60                    value.set_disconnected();
61                }
62            }
63            value
64        });
65
66        Ok(instance.clone())
67    }
68}
69
70#[cfg(feature = "speculos")]
71impl LedgerKeyShared<LedgerSpeculosKey> {
72    ///
73    /// Get (or create on the first call) a shared instance to access LedgerKey
74    pub fn instance() -> Result<LedgerKeyShared<LedgerSpeculosKey>, HWKeyError> {
75        let instance = INSTANCE_SPECULOS.get_or_init(|| {
76            let mut value = Self::new();
77            let connected = value.connect();
78            if connected.is_ok() {
79                let _ = comm::ping(&value);
80            }
81            value
82        });
83
84        Ok(instance.clone())
85    }
86}
87
88///
89/// An instance of LedgerKey that can be safely used between different threads because it opens the device and executes all commands in its own thread.
90/// This is important on macOS where the HID device is not thread safe (and sometimes it even needs to "sleep" between uses from different threads,
91/// so a Mutex doesn't always help), and when a HID device is used improperly it could panic with SIGILL.
92///
93impl<LK: LedgerKey> LedgerKeyShared<LK> {
94
95    fn new() -> Self {
96        let (tx, rx) = channel();
97        let state = Arc::new(RwLock::new(State::Init));
98        Self::run(rx, state.clone());
99        Self { t: PhantomData, channel: tx, state }
100    }
101
102    fn is_working(&self) -> bool {
103        let r = self.state.read().unwrap();
104        *r != State::Init
105    }
106
107    fn ensure_working(&self) -> Result<(), HWKeyError> {
108        if !self.is_working() {
109            return Err(HWKeyError::Unavailable)
110        }
111        Ok(())
112    }
113
114    pub fn is_connected(&self) -> bool {
115        if self.ensure_working().is_err() {
116            return false
117        }
118        let (tx, rx) = channel();
119        if let Err(e) = self.channel.send(Command::HaveDevice(tx)) {
120            log::error!("Error sending command: {:?}", e);
121            return false
122        }
123        rx.recv().unwrap_or_default()
124    }
125
126    fn set_disconnected(&self) {
127        let _ = self.channel.send(Command::SetDisconnected);
128    }
129
130    fn run(channel: Receiver<Command>, state: Arc<RwLock<State>>) {
131        thread::spawn( move || {
132            let ledger = LK::create();
133            if let Err(e) = ledger {
134                {
135                    let mut w = state.write().unwrap();
136                    *w = State::Stopped;
137                }
138                log::error!("Error creating a Ledger instance: {:?}", e);
139                return
140            }
141            let mut ledger = ledger.unwrap();
142            log::debug!("Ledger is working");
143            let mut device: Option<Arc<Mutex<LK::Transport>>> = None;
144            loop {
145                match channel.recv() {
146                    Ok(command) => {
147                        match command {
148                            Command::Connect(resp) => {
149                                log::info!("Connecting to Ledger...");
150                                if let Err(e) = ledger.connect() {
151                                    {
152                                        let mut w = state.write().unwrap();
153                                        *w = State::Disconnected;
154                                    }
155                                    log::error!("Error connecting to Ledger: {:?}", e);
156                                    let _ = resp.send(Err(e));
157                                } else {
158                                    match ledger.open_exclusive() {
159                                        Ok(transport) => {
160                                            device = Some(transport);
161                                            let mut w = state.write().unwrap();
162                                            *w = State::Working;
163                                            let _ = resp.send(Ok(()));
164                                        }
165                                        Err(e) => {
166                                            {
167                                                let mut w = state.write().unwrap();
168                                                *w = State::Disconnected;
169                                            }
170                                            log::error!("Error opening transport: {:?}", e);
171                                            let _ = resp.send(Err(e));
172                                        }
173                                    }
174                                }
175                            }
176                            Command::SetDisconnected => {
177                                log::info!("Setting Ledger as disconnected");
178                                let mut w = state.write().unwrap();
179                                *w = State::Disconnected;
180                            }
181                            Command::HaveDevice(resp) => {
182                                log::trace!("Check if device is connected");
183                                let is_connected = {
184                                    let state = state.read().unwrap().clone();
185                                    matches!(state, State::Working)
186                                };
187                                let _ = resp.send(device.is_some() && is_connected);
188                            }
189                            Command::Write(data, resp) => {
190                                match &device {
191                                    None => {
192                                        log::debug!("Ledger device not connected");
193                                        {
194                                            let mut w = state.write().unwrap();
195                                            *w = State::Disconnected;
196                                        }
197                                        let _ = resp.send(Err(HWKeyError::Unavailable));
198                                    }
199                                    Some(transport) => {
200                                        log::trace!("Writing: {}", hex::encode(&data));
201                                        let result = {
202                                            let transport_guard = transport.lock().unwrap();
203                                            transport_guard.write(data.as_slice())
204                                        };
205                                        if result.is_err() {
206                                            {
207                                                let mut w = state.write().unwrap();
208                                                *w = State::Disconnected;
209                                            }
210                                        }
211                                        let _ = resp.send(result);
212                                    }
213                                }
214                            }
215                            Command::Read(timeout, resp) => {
216                                match &device {
217                                    None => {
218                                        log::debug!("Ledger device not connected");
219                                        {
220                                            let mut w = state.write().unwrap();
221                                            *w = State::Disconnected;
222                                        }
223                                        let _ = resp.send(Err(HWKeyError::Unavailable));
224                                    }
225                                    Some(transport) => {
226                                        let mut data = [0u8; comm::HID_RPT_SIZE];
227                                        let len = {
228                                            let transport_guard = transport.lock().unwrap();
229                                            transport_guard.read_timeout(&mut data, timeout)
230                                        };
231                                        match len {
232                                            Err(e) => {
233                                                {
234                                                    let mut w = state.write().unwrap();
235                                                    *w = State::Disconnected;
236                                                }
237                                                log::trace!("Disconnected due to error: {:?}", e);
238                                                let _ = resp.send(Err(e));
239                                            }
240                                            Ok(len) => {
241                                                let mut result = Vec::with_capacity(len);
242                                                result.extend_from_slice(&data[..len]);
243                                                log::trace!("Received data: {}", hex::encode(&result));
244                                                let _ = resp.send(Ok(result));
245                                            }
246                                        }
247                                    }
248                                }
249                            }
250                        }
251                    }
252                    Err(e) => {
253                        log::error!("Stopping the manager: {:?}", e);
254                        break
255                    }
256                }
257            }
258            log::info!("Ledger is stopped");
259            {
260                let mut w = state.write().unwrap();
261                *w = State::Stopped;
262            }
263        });
264    }
265}
266
267impl<LK: LedgerKey> LedgerKey for LedgerKeyShared<LK> {
268    type Transport = Self;
269
270    fn create() -> Result<Self, HWKeyError> {
271        panic!("Not implemented. use LedgerKeyShared::instance() instead")
272    }
273
274    fn connect(&mut self) -> Result<(), HWKeyError> {
275        if self.is_connected() {
276            return Ok(())
277        }
278        let (tx, rx) = channel();
279        if let Err(e) = self.channel.send(Command::Connect(tx)) {
280            log::error!("Error sending command: {:?}", e);
281            return Err(HWKeyError::Unavailable)
282        }
283        rx.recv().unwrap_or_else(|e| {
284            log::error!("Error receiving response: {:?}", e);
285            Err(HWKeyError::Unavailable)
286        })
287    }
288
289    fn get_app_details(&self) -> Result<AppDetails, HWKeyError> {
290        self.ensure_working()?;
291
292        let apdu = APDU {
293            cla: 0xb0,
294            ins: 0x01,
295            ..APDU::default()
296        };
297        match sendrecv_timeout(self, &apdu, 5000) {
298            Err(e) => {
299                debug!("Error sending APDU: {:?}", e);
300                match e {
301                    HWKeyError::EmptyResponse => Ok(AppDetails::default()),
302                    _ => Err(e),
303                }
304            }
305            Ok(resp) => AppDetails::try_from(resp)
306        }
307    }
308
309    fn open_exclusive(&self) -> Result<Arc<Mutex<Self::Transport>>, HWKeyError> {
310        //TODO have a single instance as a field
311        Ok(Arc::new(Mutex::new(self.clone())))
312    }
313}
314
315impl<LK: LedgerKey> LedgerTransport for LedgerKeyShared<LK> {
316    fn write(&self, data: &[u8]) -> Result<usize, HWKeyError> {
317        self.ensure_working()?;
318        let (tx, rx) = channel();
319        if let Err(e) = self.channel.send(Command::Write(data.to_vec(), tx)) {
320            log::warn!("Error sending command: {:?}", e);
321            return Err(HWKeyError::Unavailable)
322        }
323        rx.recv().unwrap_or_else(|e| {
324            log::warn!("Error write receiving response: {:?}", e);
325            Err(HWKeyError::Unavailable)
326        })
327    }
328    fn read(&self, buf: &mut [u8]) -> Result<usize, HWKeyError> {
329        self.read_timeout(buf, -1)
330    }
331    fn read_timeout(&self, buf: &mut [u8], timeout_ms: i32) -> Result<usize, HWKeyError> {
332        self.ensure_working()?;
333        if buf.len() < comm::HID_RPT_SIZE {
334            log::error!("Buffer is too small");
335            return Err(HWKeyError::CommError("Buffer is too small".to_string()))
336        }
337        let (tx, rx) = channel();
338        if let Err(e) = self.channel.send(Command::Read(timeout_ms, tx)) {
339            log::warn!("Error sending command: {:?}", e);
340            return Err(HWKeyError::Unavailable)
341        }
342        let data = rx.recv().unwrap_or_else(|e| {
343            log::warn!("Error read receiving response: {:?}", e);
344            Err(HWKeyError::Unavailable)
345        })?;
346        let len = data.len();
347        buf[..len].copy_from_slice(&data);
348        Ok(len)
349    }
350}
351
352#[cfg(test)]
353mod tests {
354    use super::*;
355    use crate::ledger::connect::direct::LedgerHidKey;
356    use crate::ledger::comm;
357
358    /// Give the background thread a moment to initialize
359    fn wait_for_thread_init() {
360        thread::sleep(std::time::Duration::from_millis(10));
361    }
362
363    fn set_state<T: LedgerKey>(key: &LedgerKeyShared<T>, new_state: State) {
364        let mut w = key.state.write().unwrap();
365        *w = new_state;
366    }
367
368    #[test]
369    fn test_write_when_device_not_connected() {
370        let shared = LedgerKeyShared::<LedgerHidKey>::new();
371        set_state(&shared, State::Working);
372        wait_for_thread_init();
373
374        // Test write operation - should fail because device is None
375        let test_data = vec![0x01, 0x02, 0x03];
376        let result = shared.write(&test_data);
377        
378        assert!(result.is_err());
379        assert!(matches!(result.unwrap_err(), HWKeyError::Unavailable));
380    }
381
382    #[test]
383    fn test_read_when_device_not_connected() {
384        let shared = LedgerKeyShared::<LedgerHidKey>::new();
385        set_state(&shared, State::Working);
386        wait_for_thread_init();
387
388        // Test read operation - should fail because device is None
389        let mut buffer = [0u8; comm::HID_RPT_SIZE];
390        let result = shared.read_timeout(&mut buffer, 1000);
391        
392        assert!(result.is_err());
393        assert!(matches!(result.unwrap_err(), HWKeyError::Unavailable));
394    }
395
396    #[test]
397    fn test_is_connected_returns_false_when_not_connected() {
398        let shared = LedgerKeyShared::<LedgerHidKey>::new();
399        set_state(&shared, State::Working);
400        wait_for_thread_init();
401
402        // Test is_connected - should return false because device is None
403        let result = shared.is_connected();
404        assert_eq!(result, false);
405    }
406
407    #[test]
408    fn test_send_receive_data() {
409        use crate::ledger::connect::mock::MockLedgerKey;
410        let shared = LedgerKeyShared::<MockLedgerKey>::new();
411        wait_for_thread_init();
412        
413        let mut shared_for_connect = shared.clone();
414        let connect_result = shared_for_connect.connect();
415        assert!(connect_result.is_ok());
416        
417        assert!(shared.is_connected());
418        
419        // Test write operation
420        let test_data = vec![0x01, 0x02, 0x03];
421        let write_result = shared.write(&test_data);
422        assert!(write_result.is_ok());
423        assert_eq!(write_result.unwrap(), 3);
424        
425        // Test read operation
426        let mut buffer = [0u8; comm::HID_RPT_SIZE];
427        let read_result = shared.read_timeout(&mut buffer, 1000);
428        assert!(read_result.is_ok());
429        let bytes_read = read_result.unwrap();
430        assert_eq!(bytes_read, 2); // Mock returns success response [0x90, 0x00]
431        assert_eq!(&buffer[..bytes_read], &[0x90, 0x00]);
432    }
433}