Refactors code to made it more idiomatic. Reduces boilerplate

This commit is contained in:
Sergi Delgado Segura 2021-10-11 15:09:29 +02:00
parent a24bf131b8
commit bbd3857c70
No known key found for this signature in database
GPG key ID: 633B3A2298D70DD8
5 changed files with 92 additions and 244 deletions

View file

@ -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(),
}
}
}
}

View file

@ -79,85 +79,82 @@ impl Gatekeeper {
message: &[u8],
signature: &str,
) -> Result<UserId, AuthenticationFailure> {
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<RegistrationReceipt, MaxSlotsReached> {
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<u32, NotEnoughSlots> {
// 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<UUID> {
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<UUID, UserId>) {
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();
}

View file

@ -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<TransactionTracker> {
@ -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!()
}

View file

@ -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

View file

@ -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<std::sync::atomic::AtomicBool>,
}
/// Body of HTTP response messages.
pub enum MessageBody<T: ToString> {
Empty,
Content(T),
ChunkedContent(T),
}
impl HttpServer {
fn responding_with_body<T: ToString>(status: &str, body: MessageBody<T>) -> 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<T: ToString>(body: MessageBody<T>) -> Self {
HttpServer::responding_with_body("HTTP/1.1 200 OK", body)
}
pub fn responding_with_not_found() -> Self {
HttpServer::responding_with_body::<String>("HTTP/1.1 404 Not Found", MessageBody::Empty)
}
pub fn responding_with_server_error<T: ToString>(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,