From bbd3857c70e321f1a563036c11deb184d439fb3e Mon Sep 17 00:00:00 2001 From: Sergi Delgado Segura Date: Mon, 11 Oct 2021 15:09:29 +0200 Subject: [PATCH] Refactors code to made it more idiomatic. Reduces boilerplate --- teos/src/bitcoin_cli.rs | 25 ------ teos/src/gatekeeper.rs | 169 +++++++++++++++++++--------------------- teos/src/responder.rs | 11 +-- teos/src/rpc_errors.rs | 1 + teos/src/test_utils.rs | 130 ++----------------------------- 5 files changed, 92 insertions(+), 244 deletions(-) diff --git a/teos/src/bitcoin_cli.rs b/teos/src/bitcoin_cli.rs index c25beaa..3c70ff9 100644 --- a/teos/src/bitcoin_cli.rs +++ b/teos/src/bitcoin_cli.rs @@ -119,28 +119,3 @@ impl BitcoindClient { .await } } - -#[cfg(test)] -mod tests { - use super::*; - use crate::test_utils::HttpServer; - - /// Credentials encoded in base64. - const CREDENTIALS: &'static str = "dXNlcjpwYXNzd29yZA=="; - - impl BitcoindClient { - pub(crate) fn new_dummy(server: HttpServer) -> Self { - let client = Arc::new(Mutex::new( - RpcClient::new(CREDENTIALS, server.endpoint()).unwrap(), - )); - - BitcoindClient { - bitcoind_rpc_client: client, - host: "host".to_string(), - port: 18443, - rpc_user: "user".to_string(), - rpc_password: "password".to_string(), - } - } - } -} diff --git a/teos/src/gatekeeper.rs b/teos/src/gatekeeper.rs index 38741c4..2373027 100644 --- a/teos/src/gatekeeper.rs +++ b/teos/src/gatekeeper.rs @@ -79,85 +79,82 @@ impl Gatekeeper { message: &[u8], signature: &str, ) -> Result { - match cryptography::recover_pk(message, signature) { - Ok(rpk) => { - let user_id = UserId(rpk); - if self.registered_users.borrow().contains_key(&user_id) { - Ok(user_id) - } else { - Err(AuthenticationFailure("User not found.")) - } - } - Err(_) => Err(AuthenticationFailure("Wrong message or signature.")), + let user_id = UserId( + cryptography::recover_pk(message, signature) + .map_err(|_| AuthenticationFailure("Wrong message or signature."))?, + ); + + if self.registered_users.borrow().contains_key(&user_id) { + Ok(user_id) + } else { + Err(AuthenticationFailure("User not found.")) } } + pub fn add_update_user( &self, user_id: &UserId, ) -> Result { let block_count = self.last_known_block_header.height; - if self.registered_users.borrow().contains_key(user_id) { - // TODO: For now, new calls to register add subscription_slots to the current count and reset the expiry time - let mut borrowed = self.registered_users.borrow_mut(); - match borrowed[user_id] - .available_slots - .checked_add(self.subscription_slots) - { - Some(x) => { - borrowed.get_mut(user_id).unwrap().available_slots = x; - borrowed.get_mut(user_id).unwrap().subscription_expiry = - block_count + self.subscription_duration; + // TODO: For now, new calls to register add subscription_slots to the current count and reset the expiry time + let mut borrowed = self.registered_users.borrow_mut(); + let user_info = match borrowed.get_mut(user_id) { + // User already exists, updating the info + Some(user_info) => { + match user_info + .available_slots + .checked_add(self.subscription_slots) + { + Some(x) => { + user_info.available_slots = x; + user_info.subscription_expiry = block_count + self.subscription_duration; + user_info.clone() + } + None => return Err(MaxSlotsReached), } - None => return Err(MaxSlotsReached), } - } else { - self.registered_users.borrow_mut().insert( - user_id.clone(), - UserInfo::new( + // New user + None => { + let user_info = UserInfo::new( self.subscription_slots, block_count + self.subscription_duration, - ), - ); - } + ); + borrowed.insert(user_id.clone(), user_info.clone()); + user_info + } + }; - let user = self.registered_users.borrow()[user_id].clone(); - let receipt = RegistrationReceipt::new( + Ok(RegistrationReceipt::new( user_id.clone(), - user.available_slots, - user.subscription_expiry, - ); - Ok(receipt) + user_info.available_slots, + user_info.subscription_expiry, + )) } pub fn add_update_appointment( &self, user_id: &UserId, uuid: UUID, - appointment: ExtendedAppointment, + appointment: &ExtendedAppointment, ) -> Result { // For updates, the difference between the existing appointment size and the update is computed. - let used_slots = match self.registered_users.borrow()[user_id] - .appointments - .get(&uuid) - { - Some(x) => x.clone(), - None => 0, - }; + let mut borrowed = self.registered_users.borrow_mut(); + let user_info = borrowed.get_mut(user_id).unwrap(); + let used_slots = user_info.appointments.get(&uuid).map_or(0, |x| x.clone()); + let required_slots = (appointment.inner.encrypted_blob.len() as f32 / ENCRYPTED_BLOB_MAX_SIZE as f32) .ceil() as u32; let diff = required_slots as i64 - used_slots as i64; - if diff <= self.registered_users.borrow()[user_id].available_slots as i64 { + if diff <= user_info.available_slots as i64 { // Filling / freeing slots depending on whether this is an update or not, and if it is bigger or smaller // than the old appointment - let mut borrowed = self.registered_users.borrow_mut(); - let mut user = borrowed.get_mut(user_id).unwrap(); - user.appointments.insert(uuid, required_slots); - user.available_slots = (user.available_slots as i64 - diff) as u32; + user_info.appointments.insert(uuid, required_slots); + user_info.available_slots = (user_info.available_slots as i64 - diff) as u32; - Ok(user.available_slots) + Ok(user_info.available_slots) } else { Err(NotEnoughSlots) } @@ -167,11 +164,13 @@ impl Gatekeeper { &self, user_id: &UserId, ) -> Result<(bool, u32), AuthenticationFailure<'_>> { - if self.registered_users.borrow().contains_key(&user_id) { - let expiry = self.registered_users.borrow()[&user_id].subscription_expiry; - Ok((self.last_known_block_header.height >= expiry, expiry)) - } else { - Err(AuthenticationFailure("User not found.")) + match self.registered_users.borrow().get(&user_id) { + Some(user_info) => Ok(( + self.last_known_block_header.height >= user_info.subscription_expiry, + user_info.subscription_expiry, + )), + // This should never happen as long as calls to this method are guarded by authenticate_user + None => Err(AuthenticationFailure("User not found.")), } } @@ -203,7 +202,7 @@ impl Gatekeeper { pub fn get_outdated_appointments(&self, block_height: &u32) -> HashSet { let mut appointments = HashSet::new(); - for (_, uuids) in self.get_outdated_users(block_height).into_iter() { + for uuids in self.get_outdated_users(block_height).into_values() { appointments.extend(uuids); } @@ -220,20 +219,15 @@ impl Gatekeeper { { outdated_users = self.get_outdated_users(block_height); let mut borrowed = self.outdated_users_cache.borrow_mut(); - borrowed.insert(*block_height, outdated_users.clone()); + borrowed.insert(block_height.clone(), outdated_users.clone()); // Remove the first entry from the cache if it grows beyond the limit size + // TODO: This may be simpler using BTreeMaps once first_entry is not nightly anymore if borrowed.len() > OUTDATED_USERS_CACHE_SIZE_BLOCKS { let mut keys: Vec<&u32> = borrowed.keys().to_owned().collect(); keys.sort(); let first = keys[0].clone(); borrowed.remove(&first); - // TODO: This may be a simpler approach, but we need to make sure data is sanitized so non-existing keys are not computed. - - // Since keys are simply block heights we can get the first key by subtracting - // OUTDATED_USERS_CACHE_SIZE_BLOCKS to the given key when the cache is full - // self.outdated_users_cache - // .remove(&(block_height - OUTDATED_USERS_CACHE_SIZE_BLOCKS as u32)); } } @@ -242,21 +236,16 @@ impl Gatekeeper { pub fn delete_appointments(&self, appointments: &HashMap) { for (uuid, user_id) in appointments { - if self.registered_users.borrow().contains_key(&user_id) - && self.registered_users.borrow()[&user_id] - .appointments - .contains_key(uuid) - { - // Remove the appointment from the appointment list and update the available slots - let mut borrowed = self.registered_users.borrow_mut(); - let freed_slots = borrowed - .get_mut(&user_id) - .unwrap() - .appointments - .remove(uuid) - .unwrap(); - borrowed.get_mut(&user_id).unwrap().available_slots += freed_slots; - } + // Remove the appointment from the appointment list and update the available slots + self.registered_users + .borrow_mut() + .get_mut(&user_id) + .map(|user_info| { + user_info + .appointments + .remove(uuid) + .map(|x| user_info.available_slots += x) + }); } } } @@ -272,6 +261,8 @@ impl chain::Listen for Gatekeeper { } } + // FIXME: To be implemented + #[allow(unused_variables)] fn block_disconnected(&self, header: &bitcoin::BlockHeader, height: u32) { todo!() } @@ -387,7 +378,7 @@ mod tests { let uuid = generate_uuid(); let appointment = generate_dummy_appointment(None); let available_slots = gatekeeper - .add_update_appointment(&user_id, uuid, appointment.clone()) + .add_update_appointment(&user_id, uuid, &appointment) .unwrap(); assert!(gatekeeper.registered_users.borrow()[&user_id] @@ -398,7 +389,7 @@ mod tests { // Adding the exact same appointment should leave the slots count unchanged let update_slot_count = gatekeeper - .add_update_appointment(&user_id, uuid, appointment.clone()) + .add_update_appointment(&user_id, uuid, &appointment) .unwrap(); assert!(gatekeeper.registered_users.borrow()[&user_id] .appointments @@ -409,7 +400,7 @@ mod tests { let mut bigger_appointment = appointment.clone(); bigger_appointment.inner.encrypted_blob = Vec::from([0; ENCRYPTED_BLOB_MAX_SIZE + 1]); let update_slot_count = gatekeeper - .add_update_appointment(&user_id, uuid, bigger_appointment) + .add_update_appointment(&user_id, uuid, &bigger_appointment) .unwrap(); assert!(gatekeeper.registered_users.borrow()[&user_id] .appointments @@ -418,7 +409,7 @@ mod tests { // Adding back a smaller update (modulo ENCRYPTED_BLOB_MAX_SIZE) should reduce the count let update_slot_count = gatekeeper - .add_update_appointment(&user_id, uuid, appointment.clone()) + .add_update_appointment(&user_id, uuid, &appointment) .unwrap(); assert!(gatekeeper.registered_users.borrow()[&user_id] .appointments @@ -428,7 +419,7 @@ mod tests { // Adding an appointment with a different uuid should not count as an update let new_uuid = generate_uuid(); let update_slot_count = gatekeeper - .add_update_appointment(&user_id, new_uuid, appointment.clone()) + .add_update_appointment(&user_id, new_uuid, &appointment) .unwrap(); assert!(gatekeeper.registered_users.borrow()[&user_id] .appointments @@ -443,7 +434,7 @@ mod tests { .unwrap() .available_slots = 0; assert!(matches!( - gatekeeper.add_update_appointment(&user_id, generate_uuid(), appointment), + gatekeeper.add_update_appointment(&user_id, generate_uuid(), &appointment), Err(NotEnoughSlots) )); } @@ -503,7 +494,7 @@ mod tests { let appointment = generate_dummy_appointment(None); let uuid = generate_uuid(); gatekeeper - .add_update_appointment(&user_id, uuid, appointment) + .add_update_appointment(&user_id, uuid, &appointment) .unwrap(); // Check that data is not in the cache before querying @@ -583,10 +574,10 @@ mod tests { let appointment = generate_dummy_appointment(None); gatekeeper - .add_update_appointment(&user1_id, uuid1, appointment.clone()) + .add_update_appointment(&user1_id, uuid1, &appointment) .unwrap(); gatekeeper - .add_update_appointment(&user2_id, uuid2, appointment.clone()) + .add_update_appointment(&user2_id, uuid2, &appointment) .unwrap(); let outdated_appointments = gatekeeper.get_outdated_appointments(&start_height); @@ -682,15 +673,15 @@ mod tests { } // Calling the method with unknown data should work but do nothing - assert_eq!(gatekeeper.registered_users.borrow().len(), 0); + assert!(gatekeeper.registered_users.borrow().is_empty()); gatekeeper.delete_appointments(&all_appointments); - assert_eq!(gatekeeper.registered_users.borrow().len(), 0); + assert!(gatekeeper.registered_users.borrow().is_empty()); // If there's matching data in the gatekeeper it should be deleted for (uuid, user_id) in to_be_deleted.iter() { gatekeeper.add_update_user(&user_id).unwrap(); gatekeeper - .add_update_appointment(&user_id, *uuid, generate_dummy_appointment(None)) + .add_update_appointment(&user_id, *uuid, &generate_dummy_appointment(None)) .unwrap(); } diff --git a/teos/src/responder.rs b/teos/src/responder.rs index e636840..d9c9035 100644 --- a/teos/src/responder.rs +++ b/teos/src/responder.rs @@ -120,8 +120,8 @@ impl<'a> Responder<'a> { // Has tracker should return true as long as the given tracker is hold by the Responder. // If the tracker is partially kept, the function will log and the return will be false. // This may point out that some partial data deletion is happening, which must be fixed. - match self.trackers.borrow().get(uuid) { - Some(tracker) => match self.tx_tracker_map.borrow().get(&tracker.penalty_tx.txid()) { + self.trackers.borrow().get(uuid).map_or(false, |tracker| { + match self.tx_tracker_map.borrow().get(&tracker.penalty_tx.txid()) { Some(_) => true, None => { log::debug!( @@ -129,9 +129,8 @@ impl<'a> Responder<'a> { ); false } - }, - None => false, - } + } + }) } pub fn get_tracker(&self, uuid: &UUID) -> Option { @@ -330,6 +329,8 @@ impl<'a> Listen for Responder<'a> { }; } + // FIXME: To be implemented + #[allow(unused_variables)] fn block_disconnected(&self, header: &BlockHeader, height: u32) { todo!() } diff --git a/teos/src/rpc_errors.rs b/teos/src/rpc_errors.rs index dd40ec7..e33b3a4 100644 --- a/teos/src/rpc_errors.rs +++ b/teos/src/rpc_errors.rs @@ -1,3 +1,4 @@ +#![allow(dead_code)] // Ported from https://github.com/bitcoin/bitcoin/blob/0.18/src/rpc/protocol.h // General application defined errors diff --git a/teos/src/test_utils.rs b/teos/src/test_utils.rs index 6b85c55..8c413cf 100644 --- a/teos/src/test_utils.rs +++ b/teos/src/test_utils.rs @@ -7,14 +7,9 @@ * at your option. */ -use chunked_transfer; use rand::Rng; - -use std::io::BufRead; -use std::io::Write; use std::sync::Arc; use std::thread; -use std::time::Duration; use jsonrpc_http_server::jsonrpc_core::error::ErrorCode as JsonRpcErrorCode; use jsonrpc_http_server::jsonrpc_core::{Error as JsonRpcError, IoHandler, Params, Value}; @@ -36,7 +31,6 @@ use bitcoin::secp256k1::{PublicKey, Secp256k1}; use bitcoin::util::hash::bitcoin_merkle_root; use bitcoin::util::psbt::serialize::Deserialize; use bitcoin::util::uint::Uint256; -use lightning_block_sync::http::HttpEndpoint; use lightning_block_sync::poll::{Validate, ValidatedBlockHeader}; use lightning_block_sync::{ AsyncBlockSourceResult, BlockHeaderData, BlockSource, BlockSourceError, UnboundedCache, @@ -333,7 +327,11 @@ pub(crate) fn get_random_tx() -> Transaction { pub(crate) fn generate_dummy_appointment(dispute_txid: Option<&Txid>) -> ExtendedAppointment { let dispute_txid = match dispute_txid { Some(l) => l.clone(), - None => Txid::from_slice(&[1; 32]).unwrap(), + None => { + let mut rng = rand::thread_rng(); + let prev_txid_bytes = rng.gen::<[u8; 32]>(); + Txid::from_slice(&prev_txid_bytes).unwrap() + } }; let tx_bytes = Vec::from_hex(TX_HEX).unwrap(); @@ -379,124 +377,6 @@ pub fn create_carrier(query: MockedServerQuery) -> Carrier { Carrier::new(bitcoin_cli) } -/// Timeout for operations on TCP streams. -const TCP_STREAM_TIMEOUT: Duration = Duration::from_secs(5); - -/// Server for handling HTTP client requests with a stock response. -pub struct HttpServer { - address: std::net::SocketAddr, - handler: std::thread::JoinHandle<()>, - shutdown: std::sync::Arc, -} - -/// Body of HTTP response messages. -pub enum MessageBody { - Empty, - Content(T), - ChunkedContent(T), -} - -impl HttpServer { - fn responding_with_body(status: &str, body: MessageBody) -> Self { - let response = match body { - MessageBody::Empty => format!("{}\r\n\r\n", status), - MessageBody::Content(body) => { - let body = body.to_string(); - format!( - "{}\r\n\ - Content-Length: {}\r\n\ - \r\n\ - {}", - status, - body.len(), - body - ) - } - MessageBody::ChunkedContent(body) => { - let mut chuncked_body = Vec::new(); - { - use chunked_transfer::Encoder; - let mut encoder = Encoder::with_chunks_size(&mut chuncked_body, 8); - encoder.write_all(body.to_string().as_bytes()).unwrap(); - } - format!( - "{}\r\n\ - Transfer-Encoding: chunked\r\n\ - \r\n\ - {}", - status, - String::from_utf8(chuncked_body).unwrap() - ) - } - }; - HttpServer::responding_with(response) - } - - pub fn responding_with_ok(body: MessageBody) -> Self { - HttpServer::responding_with_body("HTTP/1.1 200 OK", body) - } - - pub fn responding_with_not_found() -> Self { - HttpServer::responding_with_body::("HTTP/1.1 404 Not Found", MessageBody::Empty) - } - - pub fn responding_with_server_error(content: T) -> Self { - let body = MessageBody::Content(content); - HttpServer::responding_with_body("HTTP/1.1 500 Internal Server Error", body) - } - - fn responding_with(response: String) -> Self { - let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); - let address = listener.local_addr().unwrap(); - - let shutdown = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); - let shutdown_signaled = std::sync::Arc::clone(&shutdown); - let handler = std::thread::spawn(move || { - for stream in listener.incoming() { - let mut stream = stream.unwrap(); - stream.set_write_timeout(Some(TCP_STREAM_TIMEOUT)).unwrap(); - - let lines_read = std::io::BufReader::new(&stream) - .lines() - .take_while(|line| !line.as_ref().unwrap().is_empty()) - .count(); - if lines_read == 0 { - continue; - } - - for chunk in response.as_bytes().chunks(16) { - if shutdown_signaled.load(std::sync::atomic::Ordering::SeqCst) { - return; - } else { - if let Err(_) = stream.write(chunk) { - break; - } - if let Err(_) = stream.flush() { - break; - } - } - } - } - }); - - Self { - address, - handler, - shutdown, - } - } - - fn shutdown(self) { - self.shutdown - .store(true, std::sync::atomic::Ordering::SeqCst); - self.handler.join().unwrap(); - } - - pub fn endpoint(&self) -> HttpEndpoint { - HttpEndpoint::for_host(self.address.ip().to_string()).with_port(self.address.port()) - } -} - pub struct BitcoindMock { pub url: String, pub server: Server,