use std::sync::Arc;
use std::time::Duration;
use futures_util::StreamExt;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use wx_rust_common::http::{ReqwestTransport, TransportBody, TransportMethod, TransportRequest};
use wx_rust_common::pipeline::stream::execute_stream;
async fn spawn_chunk_server(
status_line: &str,
chunks: Vec<Vec<u8>>,
) -> (String, Arc<std::sync::Mutex<String>>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("绑定端口");
let addr = listener.local_addr().expect("获取地址");
let status_line = status_line.to_string();
let request_log = Arc::new(std::sync::Mutex::new(String::new()));
let request_log_clone = request_log.clone();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("接受连接");
let mut raw = Vec::new();
let mut buf = [0u8; 8192];
loop {
let n = socket.read(&mut buf).await.expect("读取请求");
if n == 0 {
break;
}
raw.extend_from_slice(&buf[..n]);
let text = String::from_utf8_lossy(&raw).into_owned();
if let Some((head, _)) = text.split_once("\r\n\r\n") {
let len = head
.lines()
.find(|l| l.to_ascii_lowercase().starts_with("content-length:"))
.and_then(|l| l.split_once(':'))
.and_then(|(_, v)| v.trim().parse::<usize>().ok())
.unwrap_or(0);
if raw.len() >= head.len() + 4 + len {
break;
}
}
}
let text = String::from_utf8_lossy(&raw).into_owned();
*request_log_clone.lock().unwrap() = text;
let total: usize = chunks.iter().map(|c| c.len()).sum();
let head = format!(
"{status_line}\r\nContent-Type: text/plain\r\nContent-Length: {total}\r\nConnection: close\r\n\r\n"
);
socket.write_all(head.as_bytes()).await.expect("写响应头");
for chunk in &chunks {
socket.write_all(chunk).await.expect("写响应块");
socket.flush().await.expect("flush 响应块");
tokio::time::sleep(Duration::from_millis(20)).await;
}
socket.shutdown().await.expect("关闭连接");
});
(format!("http://{addr}"), request_log)
}
#[tokio::test]
async fn stream_chunks_arrive_in_order_and_aggregate_to_full_body() {
let chunk_a = vec![b'A'; 32 * 1024];
let chunk_b = vec![b'B'; 32 * 1024];
let chunk_c = vec![b'C'; 32 * 1024];
let mut full = Vec::new();
for c in [&chunk_a, &chunk_b, &chunk_c] {
full.extend_from_slice(c);
}
let (base, request_log) =
spawn_chunk_server("HTTP/1.1 200 OK", vec![chunk_a, chunk_b, chunk_c]).await;
let transport = ReqwestTransport::new(reqwest::Client::new());
let mut stream = execute_stream(
&transport,
TransportRequest {
method: TransportMethod::Get,
url: format!("{base}/pay/downloadbill"),
headers: vec![],
body: TransportBody::None,
},
)
.await
.expect("流式请求建立");
let mut items: Vec<bytes::Bytes> = Vec::new();
while let Some(item) = stream.next().await {
items.push(item.expect("分块读取成功"));
}
assert!(items.len() > 1, "应观察到多个分块,实际 {} 块", items.len());
let mut aggregated = Vec::new();
for item in &items {
aggregated.extend_from_slice(item);
}
assert_eq!(aggregated, full);
assert_eq!(items[0][0], b'A');
assert_eq!(
items[items.len() - 1][items[items.len() - 1].len() - 1],
b'C'
);
let request = request_log.lock().unwrap().clone();
assert!(request.starts_with("GET /pay/downloadbill "), "{request}");
}
#[tokio::test]
async fn non_success_status_returns_err() {
let (base, _) =
spawn_chunk_server("HTTP/1.1 500 Internal Server Error", vec![b"boom".to_vec()]).await;
let transport = ReqwestTransport::new(reqwest::Client::new());
let err = match execute_stream(
&transport,
TransportRequest {
method: TransportMethod::Get,
url: format!("{base}/pay/downloadbill"),
headers: vec![],
body: TransportBody::None,
},
)
.await
{
Ok(_) => panic!("500 应返回 Err"),
Err(e) => e,
};
assert!(err.to_string().contains("500"), "错误应含状态码:{err}");
}
#[tokio::test]
async fn post_body_reaches_server_via_stream() {
let (base, request_log) =
spawn_chunk_server("HTTP/1.1 200 OK", vec![b"<xml>ok</xml>".to_vec()]).await;
let transport = ReqwestTransport::new(reqwest::Client::new());
let mut stream = execute_stream(
&transport,
TransportRequest {
method: TransportMethod::PostXml("<xml>bill</xml>".to_string()),
url: format!("{base}/pay/downloadbill"),
headers: vec![],
body: TransportBody::None,
},
)
.await
.expect("流式请求建立");
let mut aggregated = Vec::new();
while let Some(item) = stream.next().await {
aggregated.extend_from_slice(&item.expect("分块读取成功"));
}
assert_eq!(aggregated, b"<xml>ok</xml>");
let request = request_log.lock().unwrap().clone();
assert!(request.starts_with("POST /pay/downloadbill "), "{request}");
assert!(request.ends_with("<xml>bill</xml>"), "{request}");
}