use crate::request::Request;
use crate::websocket::{CloseCode, ErrorCode, Message, MessageMode, MessageType};
use crate::{HttpError, Method, Uri};
use hpack::Decoder;
use log::{info, warn};
use rustls::{ClientConnection, ServerConnection, StreamOwned};
use std::io::{ErrorKind, Read, Write};
use std::net::TcpStream;
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Duration;
const IO_BUF_SIZE: usize = 1024 * 1024;
const IO_SMALL_BUF_SIZE: usize = 1024 * 64;
const WS_READ_TIMEOUT_MS: u64 = 100;
#[derive(Debug, Clone)]
pub enum Scheme {
Http(Arc<Mutex<TcpStream>>),
Https(Arc<Mutex<StreamOwned<ServerConnection, TcpStream>>>),
}
pub struct SchemeReader {
inner: SchemeReaderInner,
pending: Vec<u8>,
}
#[allow(dead_code)]
enum SchemeReaderInner {
Http(TcpStream),
Https(Box<rustls::StreamOwned<ServerConnection, TcpStream>>),
}
pub struct SchemeWriter {
inner: SchemeWriterInner,
}
#[allow(dead_code)]
enum SchemeWriterInner {
Http(TcpStream),
Https(Box<rustls::StreamOwned<ServerConnection, TcpStream>>),
}
impl Scheme {
pub fn split_for_websocket(
scheme: &Arc<Mutex<Scheme>>,
) -> Result<(SchemeReader, SchemeWriter), HttpError> {
let guard = scheme
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?;
match &*guard {
Scheme::Http(stream) => {
let inner_guard = stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?;
let read_stream = inner_guard.try_clone().map_err(|e| {
HttpError::new(500, &format!("clone read stream failed: {}", e))
})?;
let write_stream = inner_guard.try_clone().map_err(|e| {
HttpError::new(500, &format!("clone write stream failed: {}", e))
})?;
Ok((
SchemeReader {
inner: SchemeReaderInner::Http(read_stream),
pending: vec![],
},
SchemeWriter {
inner: SchemeWriterInner::Http(write_stream),
},
))
}
Scheme::Https(_) => Err(HttpError::new(
500,
"HTTPS split not supported, use shared mode",
)),
}
}
}
impl SchemeReader {
pub fn read_ws_data(
&mut self,
deflate: &crate::websocket::DeflateConfig,
) -> Result<Message, HttpError> {
if !self.pending.is_empty() {
let message = Message::parse_message(&mut self.pending, deflate);
match message.message_type {
MessageType::TimeOut => {} _ => return Ok(message),
}
}
let mut buffer = vec![0u8; IO_BUF_SIZE];
let res = match &mut self.inner {
SchemeReaderInner::Http(stream) => {
stream
.set_read_timeout(Some(Duration::from_millis(WS_READ_TIMEOUT_MS)))
.ok();
let result = stream.read(&mut buffer);
stream.set_read_timeout(None).ok();
result
}
SchemeReaderInner::Https(stream) => {
stream
.get_mut()
.set_read_timeout(Some(Duration::from_millis(WS_READ_TIMEOUT_MS)))
.ok();
let result = stream.read(&mut buffer);
stream.get_mut().set_read_timeout(None).ok();
result
}
};
match res {
Ok(0) => Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Close,
payload: vec![],
text: CloseCode::GoingAway.str(),
close: CloseCode::GoingAway,
error: ErrorCode::None,
}),
Ok(n) => {
self.pending.extend_from_slice(&buffer[..n]);
let start = std::time::Instant::now();
let total_timeout = Duration::from_secs(30);
loop {
if start.elapsed() > total_timeout {
log::warn!("等待 WebSocket 完整帧超时 (30s),关闭连接");
return Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Close,
payload: vec![],
text: "等待完整帧超时".to_string(),
close: CloseCode::ProtocolError,
error: ErrorCode::TimeOut,
});
}
let message = Message::parse_message(&mut self.pending, deflate);
match message.message_type {
MessageType::TimeOut => {
let mut more_buffer = vec![0u8; IO_SMALL_BUF_SIZE];
let more_res = match &mut self.inner {
SchemeReaderInner::Http(stream) => {
stream
.set_read_timeout(Some(Duration::from_millis(
WS_READ_TIMEOUT_MS,
)))
.ok();
let result = stream.read(&mut more_buffer);
stream.set_read_timeout(None).ok();
result
}
SchemeReaderInner::Https(stream) => {
stream
.get_mut()
.set_read_timeout(Some(Duration::from_millis(
WS_READ_TIMEOUT_MS,
)))
.ok();
let result = stream.read(&mut more_buffer);
stream.get_mut().set_read_timeout(None).ok();
result
}
};
match more_res {
Ok(0) => {
return Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Close,
payload: vec![],
text: CloseCode::GoingAway.str(),
close: CloseCode::GoingAway,
error: ErrorCode::None,
});
}
Ok(m) => {
self.pending.extend_from_slice(&more_buffer[..m]);
continue;
}
Err(ref e)
if e.kind() == ErrorKind::WouldBlock
|| e.kind() == ErrorKind::TimedOut =>
{
continue;
}
Err(_) => {
continue;
}
}
}
_ => return Ok(message),
}
}
}
Err(ref e) if e.kind() == ErrorKind::WouldBlock => Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::TimeOut,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::TimeOut,
}),
Err(e) => Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: e.to_string(),
close: CloseCode::Other(1011),
error: ErrorCode::Unknown,
}),
}
}
}
impl SchemeWriter {
pub fn write_all(&mut self, data: &[u8]) -> Result<(), HttpError> {
let result = match &mut self.inner {
SchemeWriterInner::Http(stream) => stream.write_all(data),
SchemeWriterInner::Https(stream) => stream.write_all(data),
};
match result {
Ok(()) => {
self.flush()?;
Ok(())
}
Err(e) => Err(HttpError::new(500, format!("write: {}", e).as_str())),
}
}
pub fn flush(&mut self) -> Result<(), HttpError> {
let result = match &mut self.inner {
SchemeWriterInner::Http(stream) => stream.flush(),
SchemeWriterInner::Https(stream) => stream.flush(),
};
match result {
Ok(()) => Ok(()),
Err(e) => Err(HttpError::new(500, format!("flush: {}", e).as_str())),
}
}
}
impl Scheme {
pub fn read(&mut self, data: &mut Vec<u8>) -> Result<(), HttpError> {
let mut buf = vec![0u8; IO_BUF_SIZE];
let mut index = 2;
loop {
let result = match self {
Self::Http(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.read(&mut buf),
Self::Https(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.read(&mut buf),
};
return match result {
Ok(0) => Err(HttpError::new(500, "read: 客户端主动关闭")),
Ok(n) => {
data.extend(&buf[..n]);
return Ok(());
}
Err(ref e) if e.kind() == ErrorKind::Interrupted => {
if !data.is_empty() {
return Ok(());
}
if index > 0 {
index -= 1;
continue;
}
Err(HttpError::new(
500,
format!("read现在没数据可读: {}", e.to_string().as_str()).as_str(),
))
}
Err(e) => Err(HttpError::new(
500,
format!("read: {}", e.to_string().as_str()).as_str(),
)),
};
}
}
fn read_data(&self, init_data: &mut Vec<u8>, length: usize) -> Result<(), HttpError> {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(30);
loop {
if init_data.len() >= length {
return Ok(());
}
let mut buf = vec![0u8; IO_BUF_SIZE];
let result = match self {
Self::Http(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.read(&mut buf),
Self::Https(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.read(&mut buf),
};
return match result {
Ok(0) => Err(HttpError::new(500, "read_data: 客户端主动关闭")),
Ok(n) => {
init_data.extend(&buf[..n]);
Ok(())
}
Err(ref e) if e.kind() == ErrorKind::WouldBlock => {
if std::time::Instant::now() > deadline {
return Err(HttpError::new(408, "read_data: timeout"));
}
thread::sleep(Duration::from_millis(100));
continue;
}
Err(e) => Err(HttpError::new(
500,
format!("read_data: {}", e.to_string().as_str()).as_str(),
)),
};
}
}
pub fn write(&mut self, data: &[u8]) -> Result<(), HttpError> {
let mut off = 0;
loop {
let result = match self {
Self::Http(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.write(&data[off..]),
Self::Https(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.get_mut()
.write(&data[off..]),
};
match result {
Ok(0) => return Err(HttpError::new(500, "write: 客户端主动关闭")),
Ok(e) => {
if e != data.len() {
off = e;
continue;
}
self.flush()?;
return Ok(());
}
Err(ref e)
if e.kind() == ErrorKind::WouldBlock || e.kind() == ErrorKind::Interrupted => {}
Err(e) => {
return Err(HttpError::new(
500,
format!("write: {}", e.to_string().as_str()).as_str(),
))
}
};
}
}
pub fn write_all(&mut self, data: &[u8]) -> Result<(), HttpError> {
let result = match self {
Self::Http(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.write_all(data),
Self::Https(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.write_all(data),
};
match result {
Ok(()) => {
self.flush()?;
Ok(())
}
Err(e) => Err(HttpError::new(
500,
format!("write: {}", e.to_string().as_str()).as_str(),
)),
}
}
pub fn flush(&mut self) -> Result<(), HttpError> {
let result = match self {
Self::Http(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.flush(),
Self::Https(stream) => stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?
.flush(),
};
match result {
Ok(()) => Ok(()),
Err(e) => Err(HttpError::new(
500,
format!("flush: {}", e.to_string().as_str()).as_str(),
)),
}
}
pub fn read_ws_data(
&mut self,
deflate: &crate::websocket::DeflateConfig,
) -> Result<Message, HttpError> {
let mut response = vec![];
let mut buffer = vec![0u8; IO_BUF_SIZE];
let res = match self {
Self::Http(stream) => {
let mut guard = stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?;
guard
.set_read_timeout(Some(std::time::Duration::from_millis(WS_READ_TIMEOUT_MS)))
.ok();
let result = guard.read(&mut buffer);
guard.set_read_timeout(None).ok();
result
}
Self::Https(ref mut stream) => {
let mut guard = stream
.lock()
.map_err(|e| HttpError::new(500, &format!("lock poisoned: {}", e)))?;
guard
.get_mut()
.set_read_timeout(Some(std::time::Duration::from_millis(WS_READ_TIMEOUT_MS)))
.ok();
let result = guard.read(&mut buffer);
guard.get_mut().set_read_timeout(None).ok();
result
}
};
match res {
Ok(0) => Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Close,
payload: vec![],
text: CloseCode::GoingAway.str(),
close: CloseCode::GoingAway,
error: ErrorCode::None,
}),
Ok(n) => {
response.extend(buffer[..n].to_vec());
let start = std::time::Instant::now();
let total_timeout = std::time::Duration::from_secs(30);
loop {
if start.elapsed() > total_timeout {
log::warn!("等待 WebSocket 完整帧超时 (30s),关闭连接");
return Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Close,
payload: vec![],
text: "等待完整帧超时".to_string(),
close: CloseCode::ProtocolError,
error: ErrorCode::TimeOut,
});
}
let message = Message::parse_message(&mut response, deflate);
match message.message_type {
MessageType::TimeOut => {
let mut more_buffer = vec![0u8; IO_SMALL_BUF_SIZE];
let more_res = match self {
Self::Http(stream) => {
let mut guard = stream.lock().map_err(|e| {
HttpError::new(500, &format!("lock poisoned: {}", e))
})?;
guard
.set_read_timeout(Some(std::time::Duration::from_millis(
WS_READ_TIMEOUT_MS,
)))
.ok();
let result = guard.read(&mut more_buffer);
guard.set_read_timeout(None).ok();
result
}
Self::Https(ref mut stream) => {
let mut guard = stream.lock().map_err(|e| {
HttpError::new(500, &format!("lock poisoned: {}", e))
})?;
guard
.get_mut()
.set_read_timeout(Some(std::time::Duration::from_millis(
WS_READ_TIMEOUT_MS,
)))
.ok();
let result = guard.read(&mut more_buffer);
guard.get_mut().set_read_timeout(None).ok();
result
}
};
match more_res {
Ok(0) => {
return Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Close,
payload: vec![],
text: CloseCode::GoingAway.str(),
close: CloseCode::GoingAway,
error: ErrorCode::None,
});
}
Ok(m) => {
response.extend(more_buffer[..m].to_vec());
continue;
}
Err(ref e)
if e.kind() == ErrorKind::WouldBlock
|| e.kind() == ErrorKind::TimedOut =>
{
continue;
}
Err(_) => {
continue;
}
}
}
_ => return Ok(message),
}
}
}
Err(ref e) if e.kind() == ErrorKind::WouldBlock => Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::TimeOut,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::TimeOut,
}),
Err(e) => Ok(Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: e.to_string(),
close: CloseCode::Other(1011),
error: ErrorCode::Unknown,
}),
}
}
pub fn client_ip(&mut self) -> String {
match self {
Self::Http(stream) => match stream.lock() {
Ok(guard) => match guard.peer_addr() {
Ok(e) => e.ip().to_string(),
Err(_) => "unknown".to_string(),
},
Err(_) => "unknown".to_string(),
},
Self::Https(stream) => stream
.lock()
.ok()
.and_then(|mut guard| guard.get_mut().peer_addr().ok())
.map(|a| a.ip().to_string())
.unwrap_or_else(|| "unknown".to_string()),
}
}
pub fn server_ip(&mut self) -> String {
match self {
Self::Http(stream) => stream
.lock()
.ok()
.and_then(|guard| guard.local_addr().ok())
.map(|a| a.ip().to_string())
.unwrap_or_else(|| "unknown".to_string()),
Self::Https(stream) => stream
.lock()
.ok()
.and_then(|mut guard| guard.get_mut().local_addr().ok())
.map(|a| a.ip().to_string())
.unwrap_or_else(|| "unknown".to_string()),
}
}
pub fn http2_packet(
&mut self,
init_data: &mut Vec<u8>,
) -> Result<(Vec<u8>, FrameType, u8, u32), HttpError> {
let bytes = init_data;
self.read_data(bytes, 9)?;
let headers = bytes.drain(..9).collect::<Vec<u8>>();
let length =
((headers[0] as u32) << 16) | (u32::from(headers[1]) << 8) | u32::from(headers[2]);
let frame_type = headers[3];
let flags = headers[4];
let stream_id =
u32::from_be_bytes([headers[5], headers[6], headers[7], headers[8]]) & 0x7FFF_FFFF;
self.read_data(bytes, length as usize)?;
let payload = bytes.drain(..length as usize).collect::<Vec<u8>>();
Ok((payload, FrameType::from(frame_type), flags, stream_id))
}
pub fn http2_handle_header(
&mut self,
data: &mut Vec<u8>,
request: &mut Request,
) -> Result<(), HttpError> {
loop {
let (payload, frame_type, flags, stream_id) = self.http2_packet(data)?;
if request.config.debug {
info!("http2_handle_header: frame_type: {frame_type:?} flags: {flags} stream_id: {stream_id} payload: {}", payload.len());
}
match frame_type {
FrameType::Settings => {
let is_ack = flags & 0x01 != 0;
if !is_ack {
self.http2_settings_ack()?;
}
}
FrameType::WindowUpdate => {
if payload.len() == 4 {
let raw =
u32::from_be_bytes(<[u8; 4]>::try_from(&payload[..4]).map_err(
|_| HttpError::new(400, "invalid WindowUpdate frame data"),
)?);
let increment = raw & 0x7FFF_FFFF; if request.config.debug {
info!("WindowUpdate: increment = {} {:?}", increment, payload);
}
} else {
return Err(HttpError::new(
400,
format!("Invalid WindowUpdate frame length: {}", payload.len())
.as_str(),
));
}
}
FrameType::Headers => {
let mut decoder = Decoder::new();
let headers = decoder.decode(&payload).map_err(|e| {
HttpError::new(400, &format!("HPACK decode error: {:?}", e))
})?;
if request.config.debug {
println!(
"=================请求头 {:?}=================",
thread::current().id()
);
}
for (name, value) in headers {
let header_name = String::from_utf8_lossy(name.as_slice());
let header_value = String::from_utf8_lossy(value.as_slice());
if request.config.debug {
println!("{header_name}: {header_value}");
}
match header_name.as_ref() {
":method" => request.method = Method::from(header_value.as_ref()),
":path" => request.uri = Uri::from(header_value.as_ref()),
":scheme" => request.set_header("scheme", header_value.as_ref())?,
":authority" => request.set_header("host", header_value.as_ref())?,
_ => request.set_header(&header_name, &header_value)?,
}
}
if request.config.debug {
println!("====================================================");
}
return Ok(());
}
_ => {
return Err(HttpError::new(
400,
format!("Invalid {frame_type:?}").as_str(),
))
}
}
}
}
pub fn http2_handle_body(
&mut self,
data: &mut Vec<u8>,
request: Request,
) -> Result<Vec<u8>, HttpError> {
let mut body = vec![];
loop {
let (payload, frame_type, flags, stream_id) = self.http2_packet(data)?;
if request.config.debug {
info!("http2_handle_body: frame_type: {frame_type:?} flags: {flags} stream_id: {stream_id} data: {}",payload.len());
}
match frame_type {
FrameType::Data => {
body.extend(payload);
if body.len() > request.config.max_body_size {
return Err(HttpError::new(413, "Request body too large"));
}
if flags == 1 {
return Ok(body);
}
}
FrameType::Headers => {}
FrameType::RstStream => {}
FrameType::Settings => {
if !payload.is_empty() {
self.http2_send_server_settings()?;
} else {
self.http2_settings_ack()?;
}
}
FrameType::Ping => {}
FrameType::Goaway => {
let text = String::from_utf8_lossy(&payload);
if request.config.debug {
warn!("Goaway: {text}");
}
return Ok(vec![]);
}
FrameType::WindowUpdate => {
if payload.len() == 4 {
let raw =
u32::from_be_bytes(<[u8; 4]>::try_from(&payload[..4]).map_err(
|_| HttpError::new(400, "invalid WindowUpdate frame data"),
)?);
let increment = raw & 0x7FFF_FFFF; if request.config.debug {
info!("WindowUpdate: increment = {} {:?}", increment, payload);
}
} else {
return Err(HttpError::new(
400,
format!("Invalid WindowUpdate frame length: {}", payload.len())
.as_str(),
));
}
}
FrameType::Continuation => {}
FrameType::None => {}
}
}
}
pub fn http2_send_server_settings(&mut self) -> Result<(), HttpError> {
let payload = {
let mut p = Vec::new();
p.extend_from_slice(&2u16.to_be_bytes());
p.extend_from_slice(&0u32.to_be_bytes());
p.extend_from_slice(&4u16.to_be_bytes());
p.extend_from_slice(&65_535u32.to_be_bytes());
p.extend_from_slice(&5u16.to_be_bytes());
p.extend_from_slice(&16_384u32.to_be_bytes());
p
};
let len = payload.len();
let mut f = Vec::with_capacity(9 + len);
f.extend_from_slice(&[(len >> 16) as u8, (len >> 8) as u8, len as u8]); f.push(0x04); f.push(0x00); f.extend_from_slice(&0u32.to_be_bytes()); f.extend_from_slice(&payload);
self.write_all(&f)?;
Ok(())
}
pub fn http2_settings_ack(&mut self) -> Result<(), HttpError> {
let f = [0x00, 0x00, 0x00, 0x04, 0x01, 0x00, 0x00, 0x00, 0x00];
self.write_all(&f)?;
Ok(())
}
pub fn http2_goaway(&mut self, last_stream_id: u32, error_code: u32) -> Result<(), HttpError> {
let mut frame = Vec::new();
frame.extend_from_slice(&[0x00, 0x00, 0x08]); frame.push(0x07); frame.push(0x00); frame.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); frame.extend_from_slice(&last_stream_id.to_be_bytes());
frame.extend_from_slice(&error_code.to_be_bytes());
self.write_all(frame.as_slice())?;
Ok(())
}
}
#[derive(Debug)]
pub enum FrameType {
Data,
Headers,
RstStream,
Settings,
Ping,
Goaway,
WindowUpdate,
Continuation,
None,
}
impl FrameType {
pub fn from(code: u8) -> Self {
match code {
0x00 => Self::Data,
0x01 => Self::Headers,
0x03 => Self::RstStream,
0x04 => Self::Settings,
0x06 => Self::Ping,
0x07 => Self::Goaway,
0x08 => Self::WindowUpdate,
0x09 => Self::Continuation,
_ => Self::None,
}
}
}
pub enum ClientStream {
Http(TcpStream),
Https(Box<StreamOwned<ClientConnection, TcpStream>>),
}
impl ClientStream {
pub fn write_all(&mut self, data: &[u8]) -> std::io::Result<()> {
match self {
ClientStream::Http(e) => e.write_all(data),
ClientStream::Https(e) => e.write_all(data),
}
}
pub fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
match self {
ClientStream::Http(e) => e.read(buf),
ClientStream::Https(e) => e.read(buf),
}
}
pub fn read_data(&mut self, buffer: &mut Vec<u8>) -> Result<(), String> {
let mut tmp = [0u8; 1024];
let n = self.read(&mut tmp).map_err(|e| e.to_string())?;
if n == 0 {
return Err("unexpected EOF while reading chunk data".to_string());
}
buffer.extend_from_slice(&tmp[..n]);
Ok(())
}
pub fn flush(&mut self) -> std::io::Result<()> {
match self {
ClientStream::Http(e) => e.flush(),
ClientStream::Https(e) => e.flush(),
}
}
}