use std::{net::SocketAddr, sync::Arc};
use anyhow::{Context, Result};
use async_trait::async_trait;
use log::{debug, error, info, trace};
use serde_json::Value;
use tokio::{
io::{AsyncWriteExt, BufReader},
sync::{mpsc::Sender, Mutex},
};
use tokio_rustls::server::TlsStream;
use x509_parser::prelude::FromDer;
use crate::{
config::Config,
doslimit::{ConnectionLimits, GlobalLimits},
query::Query,
rpc::{daemon::connection::Connection, rpcstats::RpcStats},
signal::NetworkNotifier,
util::Channel,
utilnet::get_global_ip_from_hostname,
};
use super::{
communication::Communcation,
parse_requests::{no_validation, parse_requests},
server::ConnectionId,
ssl_utils::create_tls_acceptor,
Message,
};
use crate::query::scripthash_subscriptions::ScripthashSubscriptions;
pub struct TcpSslCommunication {
reader: Option<tokio::io::ReadHalf<TlsStream<tokio::net::TcpStream>>>,
writer: tokio::io::WriteHalf<TlsStream<tokio::net::TcpStream>>,
}
impl TcpSslCommunication {
pub fn new(stream: TlsStream<tokio::net::TcpStream>) -> Self {
let (reader, writer) = tokio::io::split(stream);
Self {
reader: Some(reader),
writer,
}
}
}
#[async_trait]
impl Communcation for TcpSslCommunication {
async fn send_values(&mut self, values: &[Value]) -> Result<()> {
for value in values {
let line = value.to_string() + "\n";
if let Err(e) = self.writer.write_all(line.as_bytes()).await {
let truncated: String = line.chars().take(80).collect();
return Err(e).context(format!("failed to send {}", truncated));
}
}
Ok(())
}
fn start_request_receiver(&mut self, rpc_queue_sender: Sender<Message>, idle_timeout: u64) {
let reader = self.reader.take();
crate::thread::spawn_task("ssl_peer_reader", async move {
let bufreader = BufReader::new(reader.unwrap());
if let Err(e) =
parse_requests(bufreader, rpc_queue_sender, idle_timeout, no_validation).await
{
trace!("ssl_peer_reader thread: {}", e);
}
Ok(())
});
}
async fn shutdown(&mut self) {
if let Err(e) = self.writer.shutdown().await {
trace!("Failed to shutdown SSL writer: {}", e);
}
drop(self.reader.take());
}
}
pub(crate) async fn extract_hostname_from_cert(
cert_file: &std::path::Path,
) -> Result<Option<String>> {
use std::fs::File;
use std::io::BufReader;
let cert_file = File::open(cert_file).context("Failed to open SSL certificate file")?;
let mut cert_reader = BufReader::new(cert_file);
use rustls::pki_types::{pem::PemObject, CertificateDer};
let cert_data: Vec<CertificateDer> = CertificateDer::pem_reader_iter(&mut cert_reader)
.collect::<Result<Vec<_>, _>>()
.context("Failed to read certificate file")?;
if cert_data.is_empty() {
return Ok(None);
}
let cert_der = &cert_data[0];
let cert = x509_parser::prelude::X509Certificate::from_der(cert_der)
.context("Failed to parse certificate DER")?;
let subject = &cert.1.tbs_certificate.subject;
for name in subject.iter_common_name() {
if let Ok(cn) = name.as_str() {
warn!("Found cert name {}.", cn);
if !cn.is_empty() && get_global_ip_from_hostname(cn).await.is_some() {
warn!("Cert is global routable");
return Ok(Some(cn.to_string()));
} else {
warn!("Cert is not globally routable");
}
}
}
if let Ok(Some(ext)) = cert
.1
.get_extension_unique(&x509_parser::oid_registry::OID_X509_EXT_SUBJECT_ALT_NAME)
{
if let Ok((_, x509_parser::extensions::GeneralName::DNSName(dns_name))) =
x509_parser::extensions::GeneralName::from_der(ext.value)
{
if !dns_name.is_empty() && get_global_ip_from_hostname(dns_name).await.is_some() {
return Ok(Some(dns_name.to_string()));
}
}
}
Ok(None)
}
#[allow(clippy::too_many_arguments)]
pub async fn start_tcp_ssl_server(
conn_acceptor: Channel<Option<(tokio::net::TcpStream, SocketAddr)>>,
cert_file: std::path::PathBuf,
key_file: std::path::PathBuf,
global_limits: Arc<GlobalLimits>,
config: Arc<Config>,
query: Arc<Query>,
stats: Arc<RpcStats>,
connection_limits: ConnectionLimits,
network_notifier: Arc<NetworkNotifier>,
rpc_queue_senders: Arc<Mutex<std::collections::HashMap<ConnectionId, Sender<Message>>>>,
subscription_manager: Arc<ScripthashSubscriptions>,
) -> Result<()> {
let tls_acceptor = match create_tls_acceptor(&cert_file, &key_file).await {
Ok(acceptor) => acceptor,
Err(e) => {
error!("Failed to create TLS acceptor: {}", e);
return Err(e);
}
};
while let Some(conn) = conn_acceptor.receiver().await.recv().await {
let (mut stream, addr) = match conn {
Some(c) => c,
None => break,
};
let global_limits = global_limits.clone();
let tls_acceptor = tls_acceptor.clone();
let mut connections = match global_limits.inc_connection(&addr.ip()).await {
Err(e) => {
trace!("[{}] dropping tcp ssl peer - {}", addr, e);
let _ = stream.shutdown().await;
drop(stream);
continue;
}
Ok(n) => n,
};
let query_cpy = Arc::clone(&query);
let stats_cpy = Arc::clone(&stats);
let network_notifier_cpy = Arc::clone(&network_notifier);
let config = Arc::clone(&config);
let rpc_queue = Channel::bounded(config.rpc_buffer_size);
let conn_id = ConnectionId::new();
let subscription_manager_cpy = Arc::clone(&subscription_manager);
let rpc_queue_senders_cpy = Arc::clone(&rpc_queue_senders);
crate::thread::spawn_task("tcp_ssl_peer", async move {
if connections != (0, 0) {
debug!(
"[{} tcp_ssl] connected peer ({:?} out of {:?} connection slots used)",
addr,
connections,
global_limits.connection_limits(),
);
}
let tls_stream = match tls_acceptor.accept(stream).await {
Ok(stream) => stream,
Err(e) => {
trace!("[{} tcp_ssl] {}", addr, e);
match global_limits.dec_connection(&addr.ip()).await {
Ok(n) => connections = n,
Err(e) => warn!("Failed to decrement connection count: {}", e),
};
debug!(
"[{} tcp_ssl] disconnected peer ({:?} out of {:?} connection slots used)",
addr,
connections,
global_limits.connection_limits(),
);
return Ok(());
}
};
let communication = TcpSslCommunication::new(tls_stream);
rpc_queue_senders_cpy
.lock()
.await
.insert(conn_id, rpc_queue.sender());
let conn = Connection::new(
config,
query_cpy,
addr,
stats_cpy,
connection_limits,
Box::new(communication),
network_notifier_cpy,
conn_id,
subscription_manager_cpy.clone(),
);
let conn_id_for_cleanup = conn_id;
conn.run(rpc_queue).await;
subscription_manager_cpy
.remove_connection(conn_id_for_cleanup)
.await;
rpc_queue_senders_cpy
.lock()
.await
.remove(&conn_id_for_cleanup);
match global_limits.dec_connection(&addr.ip()).await {
Ok(n) => connections = n,
Err(e) => warn!("Failed to decrement connection count: {}", e),
};
debug!(
"[{} tcp_ssl] disconnected peer ({:?} out of {:?} connection slots used)",
addr,
connections,
global_limits.connection_limits(),
);
Ok(())
});
}
info!("TCP SSL RPC connections no longer accepted");
Ok(())
}