mirror of
https://github.com/rustfs/rustfs.git
synced 2026-08-10 23:26:53 +00:00
perf(storage): optimize internode RPC transfer path (#2262)
Co-authored-by: momoda693 <momoda693@gmail.com>
This commit is contained in:
+303
-149
@@ -13,13 +13,14 @@
|
||||
// limitations under the License.
|
||||
|
||||
use crate::{EtagResolvable, HashReaderDetector, HashReaderMut};
|
||||
use bytes::Bytes;
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use futures::{Stream, TryStreamExt as _};
|
||||
use http::HeaderMap;
|
||||
use pin_project_lite::pin_project;
|
||||
use reqwest::{Certificate, Client, Identity, Method, RequestBuilder};
|
||||
use rustfs_common::internode_metrics::global_internode_metrics;
|
||||
use rustfs_utils::get_env_opt_str;
|
||||
use std::error::Error as _;
|
||||
use std::io::IoSlice;
|
||||
use std::io::{self, Error};
|
||||
use std::ops::Not as _;
|
||||
use std::pin::Pin;
|
||||
@@ -28,6 +29,7 @@ use std::task::{Context, Poll};
|
||||
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_util::io::StreamReader;
|
||||
use tokio_util::sync::PollSender;
|
||||
use tracing::error;
|
||||
|
||||
/// Get the TLS path from the RUSTFS_TLS_PATH environment variable.
|
||||
@@ -111,24 +113,12 @@ fn get_http_client() -> Client {
|
||||
CLIENT.clone()
|
||||
}
|
||||
|
||||
static HTTP_DEBUG_LOG: bool = false;
|
||||
#[inline(always)]
|
||||
fn http_debug_log(args: std::fmt::Arguments) {
|
||||
if HTTP_DEBUG_LOG {
|
||||
println!("{args}");
|
||||
}
|
||||
}
|
||||
macro_rules! http_log {
|
||||
($($arg:tt)*) => {
|
||||
http_debug_log(format_args!($($arg)*));
|
||||
};
|
||||
}
|
||||
|
||||
pin_project! {
|
||||
pub struct HttpReader {
|
||||
url:String,
|
||||
method: Method,
|
||||
headers: HeaderMap,
|
||||
track_internode_metrics: bool,
|
||||
#[pin]
|
||||
inner: StreamReader<Pin<Box<dyn Stream<Item=std::io::Result<Bytes>>+Send+Sync>>, Bytes>,
|
||||
}
|
||||
@@ -147,53 +137,47 @@ impl HttpReader {
|
||||
body: Option<Vec<u8>>,
|
||||
_read_buf_size: usize,
|
||||
) -> io::Result<Self> {
|
||||
// http_log!(
|
||||
// "[HttpReader::with_capacity] url: {url}, method: {method:?}, headers: {headers:?}, buf_size: {}",
|
||||
// _read_buf_size
|
||||
// );
|
||||
// First, check if the connection is available (HEAD)
|
||||
let client = get_http_client();
|
||||
let head_resp = client.head(&url).headers(headers.clone()).send().await;
|
||||
match head_resp {
|
||||
Ok(resp) => {
|
||||
http_log!("[HttpReader::new] HEAD status: {}", resp.status());
|
||||
if !resp.status().is_success() {
|
||||
return Err(Error::other(format!("HEAD failed: url: {}, status {}", url, resp.status())));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
http_log!("[HttpReader::new] HEAD error: {e}");
|
||||
return Err(Error::other(e.source().map(|s| s.to_string()).unwrap_or_else(|| e.to_string())));
|
||||
}
|
||||
}
|
||||
|
||||
let track_internode_metrics = is_internode_rpc_url(&url);
|
||||
let client = get_http_client();
|
||||
let mut request: RequestBuilder = client.request(method.clone(), url.clone()).headers(headers.clone());
|
||||
if let Some(body) = body {
|
||||
request = request.body(body);
|
||||
}
|
||||
|
||||
let resp = request
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| Error::other(format!("HttpReader HTTP request error: {e}")))?;
|
||||
let resp = request.send().await.map_err(|e| {
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_error();
|
||||
}
|
||||
Error::other(format!("HttpReader HTTP request error: {e}"))
|
||||
})?;
|
||||
|
||||
if resp.status().is_success().not() {
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_error();
|
||||
}
|
||||
return Err(Error::other(format!(
|
||||
"HttpReader HTTP request failed with non-200 status {}",
|
||||
resp.status()
|
||||
)));
|
||||
}
|
||||
|
||||
let stream = resp
|
||||
.bytes_stream()
|
||||
.map_err(|e| Error::other(format!("HttpReader stream error: {e}")));
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_outgoing_request();
|
||||
}
|
||||
|
||||
let stream = resp.bytes_stream().map_err(move |e| {
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_error();
|
||||
}
|
||||
Error::other(format!("HttpReader stream error: {e}"))
|
||||
});
|
||||
|
||||
Ok(Self {
|
||||
inner: StreamReader::new(Box::pin(stream)),
|
||||
url,
|
||||
method,
|
||||
headers,
|
||||
track_internode_metrics,
|
||||
})
|
||||
}
|
||||
pub fn url(&self) -> &str {
|
||||
@@ -209,14 +193,17 @@ impl HttpReader {
|
||||
|
||||
impl AsyncRead for HttpReader {
|
||||
fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<std::io::Result<()>> {
|
||||
// http_log!(
|
||||
// "[HttpReader::poll_read] url: {}, method: {:?}, buf.remaining: {}",
|
||||
// self.url,
|
||||
// self.method,
|
||||
// buf.remaining()
|
||||
// );
|
||||
// Read from the inner stream
|
||||
Pin::new(&mut self.inner).poll_read(cx, buf)
|
||||
let filled_before = buf.filled().len();
|
||||
match Pin::new(&mut self.inner).poll_read(cx, buf) {
|
||||
Poll::Ready(Ok(())) => {
|
||||
let bytes_read = buf.filled().len().saturating_sub(filled_before);
|
||||
if self.track_internode_metrics && bytes_read > 0 {
|
||||
global_internode_metrics().record_recv_bytes(bytes_read);
|
||||
}
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
other => other,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -241,6 +228,7 @@ impl HashReaderDetector for HttpReader {
|
||||
|
||||
struct ReceiverStream {
|
||||
receiver: mpsc::Receiver<Option<Bytes>>,
|
||||
track_internode_metrics: bool,
|
||||
}
|
||||
|
||||
impl Stream for ReceiverStream {
|
||||
@@ -262,7 +250,12 @@ impl Stream for ReceiverStream {
|
||||
// }
|
||||
// }
|
||||
match poll {
|
||||
Poll::Ready(Some(Some(bytes))) => Poll::Ready(Some(Ok(bytes))),
|
||||
Poll::Ready(Some(Some(bytes))) => {
|
||||
if self.track_internode_metrics {
|
||||
global_internode_metrics().record_sent_bytes(bytes.len());
|
||||
}
|
||||
Poll::Ready(Some(Ok(bytes)))
|
||||
}
|
||||
Poll::Ready(Some(None)) => Poll::Ready(None), // Sender shutdown
|
||||
Poll::Ready(None) => Poll::Ready(None),
|
||||
Poll::Pending => Poll::Pending,
|
||||
@@ -276,13 +269,17 @@ pin_project! {
|
||||
method: Method,
|
||||
headers: HeaderMap,
|
||||
err_rx: tokio::sync::oneshot::Receiver<std::io::Error>,
|
||||
sender: tokio::sync::mpsc::Sender<Option<Bytes>>,
|
||||
sender: PollSender<Option<Bytes>>,
|
||||
handle: tokio::task::JoinHandle<std::io::Result<()>>,
|
||||
pending_chunk: BytesMut,
|
||||
finish:bool,
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
const HTTP_WRITER_CHANNEL_CAPACITY: usize = 8;
|
||||
const HTTP_WRITER_BUFFER_SIZE: usize = 1024 * 1024;
|
||||
|
||||
impl HttpWriter {
|
||||
/// Create a new HttpWriter for the given URL. The HTTP request is performed in the background.
|
||||
pub async fn new(url: String, method: Method, headers: HeaderMap) -> io::Result<Self> {
|
||||
@@ -290,28 +287,16 @@ impl HttpWriter {
|
||||
let url_clone = url.clone();
|
||||
let method_clone = method.clone();
|
||||
let headers_clone = headers.clone();
|
||||
let track_internode_metrics = is_internode_rpc_url(&url);
|
||||
|
||||
// First, try to write empty data to check if writable
|
||||
let client = get_http_client();
|
||||
let resp = client.put(&url).headers(headers.clone()).body(Vec::new()).send().await;
|
||||
match resp {
|
||||
Ok(resp) => {
|
||||
// http_log!("[HttpWriter::new] empty PUT status: {}", resp.status());
|
||||
if !resp.status().is_success() {
|
||||
return Err(Error::other(format!("Empty PUT failed: status {}", resp.status())));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
// http_log!("[HttpWriter::new] empty PUT error: {e}");
|
||||
return Err(Error::other(format!("Empty PUT failed: {e}")));
|
||||
}
|
||||
}
|
||||
|
||||
let (sender, receiver) = tokio::sync::mpsc::channel::<Option<Bytes>>(8);
|
||||
let (sender, receiver) = tokio::sync::mpsc::channel::<Option<Bytes>>(HTTP_WRITER_CHANNEL_CAPACITY);
|
||||
let (err_tx, err_rx) = tokio::sync::oneshot::channel::<io::Error>();
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
let stream = ReceiverStream { receiver };
|
||||
let stream = ReceiverStream {
|
||||
receiver,
|
||||
track_internode_metrics,
|
||||
};
|
||||
let body = reqwest::Body::wrap_stream(stream);
|
||||
// http_log!(
|
||||
// "[HttpWriter::spawn] sending HTTP request: url={url_clone}, method={method_clone:?}, headers={headers_clone:?}"
|
||||
@@ -330,6 +315,9 @@ impl HttpWriter {
|
||||
Ok(resp) => {
|
||||
// http_log!("[HttpWriter::spawn] got response: status={}", resp.status());
|
||||
if !resp.status().is_success() {
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_error();
|
||||
}
|
||||
let _ = err_tx.send(Error::other(format!(
|
||||
"HttpWriter HTTP request failed with non-200 status {}",
|
||||
resp.status()
|
||||
@@ -338,6 +326,9 @@ impl HttpWriter {
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_error();
|
||||
}
|
||||
// http_log!("[HttpWriter::spawn] HTTP request error: {e}");
|
||||
let _ = err_tx.send(Error::other(format!("HTTP request failed: {e}")));
|
||||
return Err(Error::other(format!("HTTP request failed: {e}")));
|
||||
@@ -349,13 +340,17 @@ impl HttpWriter {
|
||||
});
|
||||
|
||||
// http_log!("[HttpWriter::new] connection established successfully");
|
||||
if track_internode_metrics {
|
||||
global_internode_metrics().record_outgoing_request();
|
||||
}
|
||||
Ok(Self {
|
||||
url,
|
||||
method,
|
||||
headers,
|
||||
err_rx,
|
||||
sender,
|
||||
sender: PollSender::new(sender),
|
||||
handle,
|
||||
pending_chunk: BytesMut::with_capacity(HTTP_WRITER_BUFFER_SIZE),
|
||||
finish: false,
|
||||
})
|
||||
}
|
||||
@@ -373,8 +368,40 @@ impl HttpWriter {
|
||||
}
|
||||
}
|
||||
|
||||
fn is_internode_rpc_url(url: &str) -> bool {
|
||||
url.contains("/rustfs/rpc/")
|
||||
}
|
||||
|
||||
fn poll_send_error_to_io<T>(err: tokio_util::sync::PollSendError<T>, context: &str) -> io::Error {
|
||||
Error::other(format!("{context}: {err}"))
|
||||
}
|
||||
|
||||
fn send_error_to_io<T>(err: tokio_util::sync::PollSendError<T>, context: &str) -> io::Error {
|
||||
Error::other(format!("{context}: {err}"))
|
||||
}
|
||||
|
||||
impl HttpWriter {
|
||||
fn poll_send_pending_chunk(&mut self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
|
||||
if self.pending_chunk.is_empty() {
|
||||
return Poll::Ready(Ok(()));
|
||||
}
|
||||
|
||||
match self.sender.poll_reserve(cx) {
|
||||
Poll::Ready(Ok(())) => {
|
||||
let chunk = self.pending_chunk.split().freeze();
|
||||
self.sender
|
||||
.send_item(Some(chunk))
|
||||
.map_err(|e| send_error_to_io(e, "HttpWriter send error"))?;
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
Poll::Ready(Err(e)) => Poll::Ready(Err(poll_send_error_to_io(e, "HttpWriter send error"))),
|
||||
Poll::Pending => Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for HttpWriter {
|
||||
fn poll_write(mut self: Pin<&mut Self>, _cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
|
||||
fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
|
||||
// http_log!(
|
||||
// "[HttpWriter::poll_write] url: {}, method: {:?}, buf.len: {}",
|
||||
// self.url,
|
||||
@@ -385,26 +412,104 @@ impl AsyncWrite for HttpWriter {
|
||||
return Poll::Ready(Err(e));
|
||||
}
|
||||
|
||||
self.sender
|
||||
.try_send(Some(Bytes::copy_from_slice(buf)))
|
||||
.map_err(|e| Error::other(format!("HttpWriter send error: {e}")))?;
|
||||
let this = self.as_mut().get_mut();
|
||||
|
||||
if this.pending_chunk.len() >= HTTP_WRITER_BUFFER_SIZE {
|
||||
match this.poll_send_pending_chunk(cx) {
|
||||
Poll::Ready(Ok(())) => {}
|
||||
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
}
|
||||
|
||||
if buf.len() >= HTTP_WRITER_BUFFER_SIZE && this.pending_chunk.is_empty() {
|
||||
match this.sender.poll_reserve(cx) {
|
||||
Poll::Ready(Ok(())) => {
|
||||
this.sender
|
||||
.send_item(Some(Bytes::copy_from_slice(buf)))
|
||||
.map_err(|e| send_error_to_io(e, "HttpWriter send error"))?;
|
||||
return Poll::Ready(Ok(buf.len()));
|
||||
}
|
||||
Poll::Ready(Err(err)) => return Poll::Ready(Err(poll_send_error_to_io(err, "HttpWriter send error"))),
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
}
|
||||
|
||||
this.pending_chunk.extend_from_slice(buf);
|
||||
|
||||
Poll::Ready(Ok(buf.len()))
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
Poll::Ready(Ok(()))
|
||||
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
self.as_mut().get_mut().poll_send_pending_chunk(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
fn poll_write_vectored(mut self: Pin<&mut Self>, cx: &mut Context<'_>, bufs: &[IoSlice<'_>]) -> Poll<io::Result<usize>> {
|
||||
if let Ok(e) = Pin::new(&mut self.err_rx).try_recv() {
|
||||
return Poll::Ready(Err(e));
|
||||
}
|
||||
|
||||
let this = self.as_mut().get_mut();
|
||||
|
||||
if this.pending_chunk.len() >= HTTP_WRITER_BUFFER_SIZE {
|
||||
match this.poll_send_pending_chunk(cx) {
|
||||
Poll::Ready(Ok(())) => {}
|
||||
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
}
|
||||
|
||||
let total_len = bufs.iter().map(|buf| buf.len()).sum::<usize>();
|
||||
if total_len == 0 {
|
||||
return Poll::Ready(Ok(0));
|
||||
}
|
||||
|
||||
if bufs.len() == 1 && this.pending_chunk.is_empty() && total_len >= HTTP_WRITER_BUFFER_SIZE {
|
||||
match this.sender.poll_reserve(cx) {
|
||||
Poll::Ready(Ok(())) => {
|
||||
this.sender
|
||||
.send_item(Some(Bytes::copy_from_slice(bufs[0].as_ref())))
|
||||
.map_err(|e| send_error_to_io(e, "HttpWriter send error"))?;
|
||||
return Poll::Ready(Ok(total_len));
|
||||
}
|
||||
Poll::Ready(Err(err)) => return Poll::Ready(Err(poll_send_error_to_io(err, "HttpWriter send error"))),
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
}
|
||||
|
||||
for buf in bufs {
|
||||
this.pending_chunk.extend_from_slice(buf);
|
||||
}
|
||||
|
||||
Poll::Ready(Ok(total_len))
|
||||
}
|
||||
|
||||
fn is_write_vectored(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
// let url = self.url.clone();
|
||||
// let method = self.method.clone();
|
||||
|
||||
match self.as_mut().get_mut().poll_send_pending_chunk(cx) {
|
||||
Poll::Ready(Ok(())) => {}
|
||||
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
|
||||
if !self.finish {
|
||||
// http_log!("[HttpWriter::poll_shutdown] url: {}, method: {:?}", url, method);
|
||||
self.sender
|
||||
.try_send(None)
|
||||
.map_err(|e| Error::other(format!("HttpWriter shutdown error: {e}")))?;
|
||||
let this = self.as_mut().get_mut();
|
||||
match this.sender.poll_reserve(cx) {
|
||||
Poll::Ready(Ok(())) => {
|
||||
this.sender
|
||||
.send_item(None)
|
||||
.map_err(|e| send_error_to_io(e, "HttpWriter shutdown error"))?;
|
||||
}
|
||||
Poll::Ready(Err(err)) => return Poll::Ready(Err(poll_send_error_to_io(err, "HttpWriter shutdown error"))),
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
// http_log!(
|
||||
// "[HttpWriter::poll_shutdown] sent shutdown signal to HTTP request, url: {}, method: {:?}",
|
||||
// url,
|
||||
@@ -415,7 +520,7 @@ impl AsyncWrite for HttpWriter {
|
||||
}
|
||||
// Wait for the HTTP request to complete
|
||||
use futures::FutureExt;
|
||||
match Pin::new(&mut self.get_mut().handle).poll_unpin(_cx) {
|
||||
match Pin::new(&mut self.get_mut().handle).poll_unpin(cx) {
|
||||
Poll::Ready(Ok(_)) => {
|
||||
// http_log!(
|
||||
// "[HttpWriter::poll_shutdown] HTTP request finished successfully, url: {}, method: {:?}",
|
||||
@@ -437,77 +542,126 @@ impl AsyncWrite for HttpWriter {
|
||||
}
|
||||
}
|
||||
|
||||
// #[cfg(test)]
|
||||
// mod tests {
|
||||
// use super::*;
|
||||
// use reqwest::Method;
|
||||
// use std::vec;
|
||||
// use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use axum::{Router, body::Body, extract::State, http::StatusCode, response::IntoResponse, routing::get};
|
||||
use http_body_util::BodyExt as _;
|
||||
use std::io::IoSlice;
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
};
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::TcpListener,
|
||||
sync::Mutex,
|
||||
};
|
||||
|
||||
// #[tokio::test]
|
||||
// async fn test_http_writer_err() {
|
||||
// // Use a real local server for integration, or mockito for unit test
|
||||
// // Here, we use the Go test server at 127.0.0.1:8081 (scripts/testfile.go)
|
||||
// let url = "http://127.0.0.1:8081/testfile".to_string();
|
||||
// let data = vec![42u8; 8];
|
||||
#[derive(Clone, Default)]
|
||||
struct TestState {
|
||||
head_count: Arc<AtomicUsize>,
|
||||
get_count: Arc<AtomicUsize>,
|
||||
put_count: Arc<AtomicUsize>,
|
||||
put_bodies: Arc<Mutex<Vec<Vec<u8>>>>,
|
||||
}
|
||||
|
||||
// // Write
|
||||
// // Add header X-Deny-Write = 1 to simulate non-writable situation
|
||||
// let mut headers = HeaderMap::new();
|
||||
// headers.insert("X-Deny-Write", "1".parse().unwrap());
|
||||
// // Here we use PUT method
|
||||
// let writer_result = HttpWriter::new(url.clone(), Method::PUT, headers).await;
|
||||
// match writer_result {
|
||||
// Ok(mut writer) => {
|
||||
// // If creation succeeds, write should fail
|
||||
// let write_result = writer.write_all(&data).await;
|
||||
// assert!(write_result.is_err(), "write_all should fail when server denies write");
|
||||
// if let Err(e) = write_result {
|
||||
// println!("write_all error: {e}");
|
||||
// }
|
||||
// let shutdown_result = writer.shutdown().await;
|
||||
// if let Err(e) = shutdown_result {
|
||||
// println!("shutdown error: {e}");
|
||||
// }
|
||||
// }
|
||||
// Err(e) => {
|
||||
// // Direct construction failure is also acceptable
|
||||
// println!("HttpWriter::new error: {e}");
|
||||
// assert!(
|
||||
// e.to_string().contains("Empty PUT failed") || e.to_string().contains("Forbidden"),
|
||||
// "unexpected error: {e}"
|
||||
// );
|
||||
// return;
|
||||
// }
|
||||
// }
|
||||
// // Should not reach here
|
||||
// panic!("HttpWriter should not allow writing when server denies write");
|
||||
// }
|
||||
async fn get_stream(State(state): State<TestState>) -> impl IntoResponse {
|
||||
state.get_count.fetch_add(1, Ordering::SeqCst);
|
||||
(StatusCode::OK, Body::from("hello"))
|
||||
}
|
||||
|
||||
// #[tokio::test]
|
||||
// async fn test_http_writer_and_reader_ok() {
|
||||
// // Use local Go test server
|
||||
// let url = "http://127.0.0.1:8081/testfile".to_string();
|
||||
// let data = vec![99u8; 512 * 1024]; // 512KB of data
|
||||
async fn reject_head(State(state): State<TestState>) -> impl IntoResponse {
|
||||
state.head_count.fetch_add(1, Ordering::SeqCst);
|
||||
StatusCode::METHOD_NOT_ALLOWED
|
||||
}
|
||||
|
||||
// // Write (without X-Deny-Write)
|
||||
// let headers = HeaderMap::new();
|
||||
// let mut writer = HttpWriter::new(url.clone(), Method::PUT, headers).await.unwrap();
|
||||
// writer.write_all(&data).await.unwrap();
|
||||
// writer.shutdown().await.unwrap();
|
||||
async fn accept_put(State(state): State<TestState>, body: Body) -> impl IntoResponse {
|
||||
state.put_count.fetch_add(1, Ordering::SeqCst);
|
||||
let bytes = body.collect().await.unwrap().to_bytes();
|
||||
state.put_bodies.lock().await.push(bytes.to_vec());
|
||||
StatusCode::OK
|
||||
}
|
||||
|
||||
// http_log!("Wrote {} bytes to {} (ok case)", data.len(), url);
|
||||
async fn start_test_server(state: TestState) -> (String, tokio::task::JoinHandle<()>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
let app = Router::new()
|
||||
.route("/stream", get(get_stream).head(reject_head).put(accept_put))
|
||||
.with_state(state);
|
||||
|
||||
// // Read back
|
||||
// let mut reader = HttpReader::with_capacity(url.clone(), Method::GET, HeaderMap::new(), 8192)
|
||||
// .await
|
||||
// .unwrap();
|
||||
// let mut buf = Vec::new();
|
||||
// reader.read_to_end(&mut buf).await.unwrap();
|
||||
// assert_eq!(buf, data);
|
||||
let handle = tokio::spawn(async move {
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
// // println!("Read {} bytes from {} (ok case)", buf.len(), url);
|
||||
// // tokio::time::sleep(std::time::Duration::from_secs(2)).await; // Wait for server to process
|
||||
// // println!("[test_http_writer_and_reader_ok] completed successfully");
|
||||
// }
|
||||
// }
|
||||
(format!("http://{addr}/stream"), handle)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_reader_does_not_send_preflight_head() {
|
||||
let state = TestState::default();
|
||||
let (url, handle) = start_test_server(state.clone()).await;
|
||||
|
||||
let mut reader = HttpReader::new(url, Method::GET, HeaderMap::new(), None).await.unwrap();
|
||||
let mut buf = Vec::new();
|
||||
reader.read_to_end(&mut buf).await.unwrap();
|
||||
|
||||
assert_eq!(buf, b"hello");
|
||||
assert_eq!(state.head_count.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(state.get_count.load(Ordering::SeqCst), 1);
|
||||
|
||||
handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_writer_does_not_send_empty_preflight_put() {
|
||||
let state = TestState::default();
|
||||
let (url, handle) = start_test_server(state.clone()).await;
|
||||
|
||||
let mut writer = HttpWriter::new(url, Method::PUT, HeaderMap::new()).await.unwrap();
|
||||
writer.write_all(b"payload").await.unwrap();
|
||||
writer.shutdown().await.unwrap();
|
||||
|
||||
assert_eq!(state.put_count.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(state.put_bodies.lock().await.as_slice(), &[b"payload".to_vec()]);
|
||||
|
||||
handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_writer_handles_many_small_writes() {
|
||||
let state = TestState::default();
|
||||
let (url, handle) = start_test_server(state.clone()).await;
|
||||
|
||||
let mut writer = HttpWriter::new(url, Method::PUT, HeaderMap::new()).await.unwrap();
|
||||
let chunk = b"0123456789abcdef";
|
||||
let mut expected = Vec::new();
|
||||
for _ in 0..256 {
|
||||
writer.write_all(chunk).await.unwrap();
|
||||
expected.extend_from_slice(chunk);
|
||||
}
|
||||
writer.shutdown().await.unwrap();
|
||||
|
||||
assert_eq!(state.put_count.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(state.put_bodies.lock().await.as_slice(), &[expected]);
|
||||
|
||||
handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_writer_supports_vectored_writes() {
|
||||
let state = TestState::default();
|
||||
let (url, handle) = start_test_server(state.clone()).await;
|
||||
|
||||
let mut writer = HttpWriter::new(url, Method::PUT, HeaderMap::new()).await.unwrap();
|
||||
let bufs = [IoSlice::new(b"hello "), IoSlice::new(b"world")];
|
||||
let written = writer.write_vectored(&bufs).await.unwrap();
|
||||
assert_eq!(written, 11);
|
||||
writer.shutdown().await.unwrap();
|
||||
|
||||
assert_eq!(state.put_count.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(state.put_bodies.lock().await.as_slice(), &[b"hello world".to_vec()]);
|
||||
|
||||
handle.abort();
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user