Why does async-tungstenite hate me?

This commit is contained in:
Alex Emmet 2025-11-11 21:47:58 +00:00
commit 04356f991a
8 changed files with 674 additions and 374 deletions

View file

@ -1,4 +1,4 @@
use color_eyre::eyre::Ok;
use base64::encode;
use futures::{StreamExt, TryFutureExt};
use http_body_util::Full;
use hyper::body::Bytes;
@ -6,20 +6,25 @@ use hyper::{
Request as HttpRequest, Response as HttpResponse, StatusCode, body::Incoming,
server::conn::http1, upgrade,
};
use std::error::Error;
use std::io::{self};
use std::{future::Future, pin::Pin, time::Duration};
use tokio::net::TcpListener;
use tower::Service;
use warp::filters::log::log;
use crate::gui::log_panel::log_message;
use crate::langu::language_manager::format;
use crate::server::socket::handle;
use hyper_util::rt::tokio::TokioIo;
use hyper_util::service::TowerToHyperService;
use rustls::ServerConfig;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use sha1::{Digest, Sha1};
use std::error::Error;
use std::io::ErrorKind;
use std::io::{self, BufReader};
use std::result::Result::Ok;
use std::sync::Arc;
use std::{future::Future, pin::Pin, time::Duration};
use tokio::net::TcpListener;
use tokio_tungstenite::WebSocketStream;
use tower::Service;
use crate::gui::log_panel::log_message;
use crate::server::socket::handle;
use crate::util::file_util::load_file_buf;
use tokio_rustls::TlsAcceptor;
#[derive(Clone)]
struct HttpService;
@ -50,42 +55,57 @@ impl Service<HttpRequest<Incoming>> for HttpService {
{
log_message("Attempting WebSocket upgrade on /ws");
let response = HttpResponse::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header("Upgrade", "websocket")
.header("Connection", "Upgrade")
.body(Full::new(Bytes::from("")))
.unwrap();
tokio::spawn(async move {
match upgrades.await {
std::result::Result::Ok(upgraded_stream) => {
let raw_stream = TokioIo::new(upgraded_stream);
if let Some(sec_websocket_key) = headers.get("sec-websocket-key") {
let sec_websocket_key = sec_websocket_key.to_str().unwrap_or("").to_string();
let sec_websocket_accept = calculate_accept_key(&sec_websocket_key);
let handshake_result = WebSocketStream::from_raw_socket(
raw_stream,
tungstenite::protocol::Role::Server,
None,
)
.await;
let (writer, reader) = handshake_result.split();
handle(path, writer, reader);
let response = HttpResponse::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header("Upgrade", "websocket")
.header("Connection", "Upgrade")
.header("Sec-WebSocket-Accept", sec_websocket_accept)
.body(Full::new(Bytes::from("")))
.unwrap();
tokio::spawn(async move {
match upgrades.await {
std::result::Result::Ok(upgraded_stream) => {
let raw_stream = TokioIo::new(upgraded_stream);
let handshake_result = WebSocketStream::from_raw_socket(
raw_stream,
tungstenite::protocol::Role::Server,
None,
)
.await;
log_message(format!("Handling WebSocket connection",));
let (writer, reader) = handshake_result.split();
handle(path, writer, reader);
}
Err(e) => {
log_message(format!(
"WebSocket upgrade failed after response: {:?}",
e
));
}
}
Err(e) => {
log_message(format!(
"WebSocket upgrade failed after response: {:?}",
e
));
}
}
});
Ok(response)
});
Ok(response)
} else {
log_message("No Sec-WebSocket-Key found in request headers");
// Handle error: No Sec-WebSocket-Key
let response = HttpResponse::builder()
.status(StatusCode::BAD_REQUEST)
.body(Full::new(Bytes::from("Missing Sec-WebSocket-Key")))
.unwrap();
Ok(response)
}
} else {
let (status, body_text) = match path.as_str() {
"/" => (
StatusCode::OK,
"Barebones Server: Try connecting to WebSocket at ws://<host>:<port>/ws or check /status.",
"Server: Try connecting to WebSocket at ws://<host>:<port>/ws or check /status.",
),
"/status" => (StatusCode::OK, "HTTP Server Status: Barebones Online"),
"/status" => (StatusCode::OK, "HTTP Server Status: Online"),
_ => (StatusCode::NOT_FOUND, "404 Not Found"),
};
let body = Full::new(Bytes::from(body_text.to_string()));
@ -98,7 +118,7 @@ impl Service<HttpRequest<Incoming>> for HttpService {
}
};
Box::pin(fut.map_err(|err| {
Box::pin(fut.map_err(|err: color_eyre::eyre::ErrReport| {
io::Error::new(
io::ErrorKind::Other,
format!("Error in request handling: {}", err),
@ -108,6 +128,16 @@ impl Service<HttpRequest<Incoming>> for HttpService {
}
pub async fn start(port: u16) -> bool {
let tls_config = match load_tls_config() {
Ok(config) => config,
Err(e) => {
log_message(format!("Failed to load TLS configuration: {}", e));
log_message("Server stopped. Ensure 'certs/cert.pem' and 'certs/key.pem' exist.");
return false;
}
};
let acceptor = TlsAcceptor::from(tls_config);
// Bind to the port
let listener = TcpListener::bind(format!("0.0.0.0:{}", port)).await;
if let Err(e) = listener {
@ -116,7 +146,7 @@ pub async fn start(port: u16) -> bool {
}
let listener = listener.unwrap();
log_message(format!(
"Barebones Server listening for HTTP and WS on 0.0.0.0:{}",
"Server listening for HTTP and WS on 0.0.0.0:{}",
port
));
@ -125,8 +155,20 @@ pub async fn start(port: u16) -> bool {
match listener.accept().await {
std::result::Result::Ok((stream, _addr)) => {
let service = HttpService;
let acceptor = acceptor.clone();
tokio::spawn(async move {
let io = TokioIo::new(stream);
let tls_stream = match acceptor.accept(stream).await {
Ok(s) => s,
Err(e) => {
// Ignore non-TLS clients connecting to the TLS port
if e.kind() != io::ErrorKind::Interrupted {
log_message(format!("TLS Handshake error: {:?}", e));
}
return;
}
};
let io = TokioIo::new(tls_stream);
if let Err(err) = http1::Builder::new()
.preserve_header_case(true)
@ -158,3 +200,50 @@ pub async fn start(port: u16) -> bool {
});
true
}
fn calculate_accept_key(key: &str) -> String {
let websocket_guid = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
let mut sha1 = Sha1::new();
sha1.update(key.as_bytes());
sha1.update(websocket_guid.as_bytes());
let result = sha1.finalize();
encode(result) // Base64 encode the result
}
fn load_tls_config() -> Result<Arc<ServerConfig>, Box<dyn Error>> {
// Load certificate file
let mut cert_file = BufReader::new(load_file_buf("certs", "cert.pem")?);
let cert_ders = rustls_pemfile::certs(&mut cert_file)
.collect::<Result<Vec<CertificateDer>, io::Error>>()?;
// PKCS8
let mut key_file = BufReader::new(load_file_buf("certs", "cert.key")?);
let mut key_ders = rustls_pemfile::pkcs8_private_keys(&mut key_file)
.map(|r| r.map(Into::into)) // Explicit conversion
.collect::<Result<Vec<PrivateKeyDer>, io::Error>>()?;
if key_ders.is_empty() {
// RSA
key_file = BufReader::new(load_file_buf("certs", "cert.key")?);
key_ders = rustls_pemfile::rsa_private_keys(&mut key_file)
.map(|r| r.map(Into::into))
.collect::<Result<Vec<PrivateKeyDer>, io::Error>>()?;
}
if key_ders.is_empty() {
// EC
key_file = BufReader::new(load_file_buf("certs", "cert.key")?);
key_ders = rustls_pemfile::ec_private_keys(&mut key_file)
.map(|r| r.map(Into::into))
.collect::<Result<Vec<PrivateKeyDer>, io::Error>>()?;
}
if key_ders.is_empty() {
return Err("No private keys found in key file. (Tried PKCS8, RSA, and EC)".into());
}
let config = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(cert_ders, key_ders.remove(0))
.map_err(|e| io::Error::new(ErrorKind::Other, e.to_string()))?;
Ok(Arc::new(config))
}

View file

@ -1,6 +1,6 @@
use crate::communities::{community_connection::CommunityConnection, community_manager};
use crate::gui::log_panel::log_message;
use async_tungstenite::accept_hdr_async;
use futures::StreamExt;
use futures::stream::SplitSink;
use futures::stream::SplitStream;
@ -40,11 +40,13 @@ pub fn handle(
.handle_message(text.to_string())
.await;
} else if msg.is_close() {
log_message(format!("Closing: {}", msg));
community_conn.handle_close().await;
return;
}
}
Some(Err(_)) => {
Some(Err(e)) => {
log_message(format!("Closing ERR: {}", e));
community_conn.handle_close().await;
return;
}
@ -58,30 +60,3 @@ pub fn handle(
}
});
}
pub async fn start(port: u16) -> bool {
let listener = TcpListener::bind(format!("0.0.0.0:{}", port)).await;
if let Err(_) = listener {
return false;
}
let listener = listener.unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let mut path: String = "/".to_string();
let callback = |req: &Request, response: Response| {
path = format!("{}", &req.uri().path());
Ok(response)
};
let ws_stream = match accept_hdr_async(stream.compat(), callback).await {
Ok(ws) => ws,
Err(_) => {
return;
}
};
let (reader, writer) = ws_stream.split();
//handle(path, reader, writer);
}
});
true
}