use std::sync::Arc;
use bytes::Bytes;
use futures_util::StreamExt;
use http::{Request, Response, StatusCode};
use http_body::Body;
use http_body_util::Full;
use engineioxide_core::Sid;
use crate::{
DisconnectReason,
body::ResponseBody,
engine::EngineIo,
errors::Error,
handler::EngineIoHandler,
packet::{OpenPacket, Packet},
service::{ProtocolVersion, TransportType},
transport::polling::payload::Payload,
};
mod payload;
fn http_response<B, D>(
code: StatusCode,
data: D,
is_binary: bool,
) -> Result<Response<ResponseBody<B>>, http::Error>
where
D: Into<Bytes>,
{
use http::header::*;
let body: Bytes = data.into();
let res = Response::builder()
.status(code)
.header(CONTENT_LENGTH, body.len());
if is_binary {
res.header(CONTENT_TYPE, "application/octet-stream")
} else {
res.header(CONTENT_TYPE, "text/plain; charset=UTF-8")
}
.body(ResponseBody::custom_response(Full::new(body)))
}
pub fn open_req<H, B, R>(
engine: Arc<EngineIo<H>>,
protocol: ProtocolVersion,
req: Request<R>,
#[cfg(feature = "v3")] supports_binary: bool,
) -> Result<Response<ResponseBody<B>>, Error>
where
H: EngineIoHandler,
B: Send + 'static,
{
let socket = engine.create_session(
protocol,
TransportType::Polling,
req.into_parts().0,
#[cfg(feature = "v3")]
supports_binary,
);
let packet = OpenPacket::new(TransportType::Polling, socket.id, &engine.config);
socket.spawn_heartbeat(engine.config.ping_interval, engine.config.ping_timeout);
let packet: String = Packet::Open(packet).into();
let packet = {
#[cfg(feature = "v3")]
{
let mut packet = packet;
if protocol == ProtocolVersion::V3 {
packet = format!("{}:{}", packet.chars().count(), packet);
}
packet
}
#[cfg(not(feature = "v3"))]
packet
};
http_response(StatusCode::OK, packet, false).map_err(Error::Http)
}
pub async fn polling_req<B, H>(
engine: Arc<EngineIo<H>>,
protocol: ProtocolVersion,
sid: Sid,
) -> Result<Response<ResponseBody<B>>, Error>
where
B: Send + 'static,
H: EngineIoHandler,
{
let socket = engine.get_socket(sid).ok_or(Error::UnknownSessionID(sid))?;
if !socket.is_http() {
return Err(Error::TransportMismatch);
}
if socket.is_upgrading() {
#[cfg(feature = "tracing")]
tracing::debug!(?sid, "socket is upgrading, sending NOOP packet");
#[cfg(feature = "v3")]
let data = payload::packet_encoder(Packet::Noop, socket.protocol, socket.supports_binary);
#[cfg(not(feature = "v3"))]
let data = payload::packet_encoder(Packet::Noop, socket.protocol);
let is_binary = false; return Ok(http_response(StatusCode::OK, data, is_binary)?);
}
let rx = match socket.internal_rx.try_lock() {
Ok(s) => s,
Err(_) => {
socket.close(DisconnectReason::MultipleHttpPollingError);
return Err(Error::HttpErrorResponse(StatusCode::BAD_REQUEST));
}
};
#[cfg(feature = "tracing")]
tracing::debug!("[sid={sid}] polling request");
let max_payload = engine.config.max_payload;
#[cfg(feature = "v3")]
let Payload { data, has_binary } =
payload::encoder(rx, protocol, socket.supports_binary, max_payload).await?;
#[cfg(not(feature = "v3"))]
let Payload { data, has_binary } = payload::encoder(rx, protocol, max_payload).await?;
#[cfg(feature = "tracing")]
tracing::debug!("[sid={sid}] sending data: {:?}", data);
Ok(http_response(StatusCode::OK, data, has_binary)?)
}
pub async fn post_req<R, B, H>(
engine: Arc<EngineIo<H>>,
protocol: ProtocolVersion,
sid: Sid,
body: Request<R>,
) -> Result<Response<ResponseBody<B>>, Error>
where
H: EngineIoHandler,
R: Body + Send + Unpin + 'static,
<R as Body>::Error: std::fmt::Debug,
<R as Body>::Data: Send,
B: Send + 'static,
{
let socket = engine.get_socket(sid).ok_or(Error::UnknownSessionID(sid))?;
if !socket.is_http() {
return Err(Error::TransportMismatch);
}
let packets = payload::decoder(body, protocol, engine.config.max_payload);
futures_util::pin_mut!(packets);
while let Some(packet) = packets.next().await {
match packet {
Ok(Packet::Close) => {
#[cfg(feature = "tracing")]
tracing::debug!("[sid={sid}] closing session");
socket.send(Packet::Noop)?;
engine.close_session(sid, DisconnectReason::TransportClose);
break;
}
Ok(Packet::Pong | Packet::Ping) => socket
.heartbeat_tx
.try_send(())
.map_err(|_| Error::HeartbeatTimeout),
Ok(Packet::Message(msg)) => {
engine.handler.on_message(msg, socket.clone());
Ok(())
}
Ok(Packet::Binary(bin) | Packet::BinaryV3(bin)) => {
engine.handler.on_binary(bin, socket.clone());
Ok(())
}
Ok(p) => {
#[cfg(feature = "tracing")]
tracing::debug!("[sid={sid}] bad packet received: {:?}", &p);
Err(Error::BadPacket(p))
}
Err(e) => {
#[cfg(feature = "tracing")]
tracing::debug!("[sid={sid}] error parsing packet: {:?}", e);
engine.close_session(sid, DisconnectReason::PacketParsingError);
return Err(e);
}
}?;
}
Ok(http_response(StatusCode::OK, "ok", false)?)
}