use std::collections::{HashMap, HashSet};
use std::fmt;
use std::sync::{Arc, Mutex, PoisonError};
use std::time::{Duration, Instant};
use bevy_app::AppExit;
use bevy_ecs::message::{MessageReader, MessageWriter};
use bevy_ecs::resource::Resource;
use bevy_ecs::system::{Commands, Res, ResMut};
use bevy_time::{Real, Time};
use http::header::HeaderName;
use http::Method;
use crate::client::{HttpClient, Queued, Route};
use crate::config::HttpConfig;
use crate::credentials::BackendCredentials;
use crate::request::{build_uri, OutgoingRequest, PreparedRequest, RequestId};
use crate::response::{BackendError, HttpResponse};
use crate::transport::{HttpTransportRes, HttpTransportResult};
pub const DEADLINE_GRACE: Duration = Duration::from_secs(5);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum RequestKind {
Http,
WebSocket,
Ssh,
Sftp,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct RequestInfo {
pub kind: RequestKind,
pub method: Option<Method>,
pub target: String,
}
struct Entry {
route: Route,
info: RequestInfo,
deadline: Duration,
allowed: Duration,
generation: u64,
}
type Answer = (RequestId, Route, HttpTransportResult);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub(crate) enum Protocol {
#[cfg(feature = "ws")]
WebSocket,
#[cfg(feature = "ssh")]
Ssh,
}
#[derive(Clone, Default)]
pub(crate) struct CancelList(Arc<Mutex<Vec<(RequestId, u8)>>>);
impl CancelList {
fn lock(&self) -> std::sync::MutexGuard<'_, Vec<(RequestId, u8)>> {
self.0.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn push(&self, id: RequestId) {
self.lock().push((id, 0));
}
pub(crate) fn claim(&self, mut owns: impl FnMut(RequestId) -> bool) -> Vec<RequestId> {
let mut list = self.lock();
let mut claimed = Vec::new();
list.retain(|(id, _)| {
if owns(*id) {
claimed.push(*id);
false
} else {
true
}
});
claimed
}
pub(crate) fn age(&self) {
self.lock().retain_mut(|(_, passes)| {
*passes = passes.saturating_add(1);
*passes < 3
});
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.lock().len()
}
}
#[derive(Resource)]
pub struct InFlight {
entries: HashMap<RequestId, Entry>,
rows: HashMap<Protocol, HashMap<RequestId, RequestInfo>>,
ready: Vec<Answer>,
cancels: CancelList,
epoch: Instant,
}
impl Default for InFlight {
fn default() -> Self {
Self { entries: HashMap::new(), rows: HashMap::new(), ready: Vec::new(), cancels: CancelList::default(), epoch: Instant::now() }
}
}
impl fmt::Debug for InFlight {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("InFlight").field("pending", &self.len()).finish_non_exhaustive()
}
}
impl InFlight {
pub fn contains(&self, id: RequestId) -> bool {
self.entries.contains_key(&id) || self.rows.values().any(|rows| rows.contains_key(&id))
}
pub fn len(&self) -> usize {
self.rows.values().fold(self.entries.len(), |n, rows| n.saturating_add(rows.len()))
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty() && self.rows.values().all(HashMap::is_empty)
}
pub fn ids(&self) -> Vec<RequestId> {
let mut ids: Vec<RequestId> = self.entries.keys().chain(self.rows.values().flat_map(HashMap::keys)).copied().collect();
ids.sort_unstable();
ids
}
pub fn describe(&self, id: RequestId) -> Option<&RequestInfo> {
self.entries.get(&id).map(|e| &e.info).or_else(|| self.rows.values().find_map(|rows| rows.get(&id)))
}
fn now(&self, time: Option<&Time<Real>>) -> Duration {
time.map_or_else(|| self.epoch.elapsed(), Time::elapsed)
}
pub(crate) fn cancel_list(&self) -> CancelList {
self.cancels.clone()
}
#[cfg_attr(not(any(feature = "ws", feature = "ssh")), allow(dead_code))]
pub(crate) fn claim_cancels(&self, owns: impl FnMut(RequestId) -> bool) -> Vec<RequestId> {
self.cancels.claim(owns)
}
#[cfg_attr(not(any(feature = "ws", feature = "ssh")), allow(dead_code))]
pub(crate) fn set_rows(&mut self, protocol: Protocol, rows: impl IntoIterator<Item = (RequestId, RequestInfo)>) {
let map = self.rows.entry(protocol).or_default();
map.clear();
map.extend(rows);
}
}
pub(crate) fn prepare(request: OutgoingRequest, config: &HttpConfig, credentials: Option<&BackendCredentials>) -> Result<PreparedRequest, BackendError> {
let mut request = request;
if let Some(error) = request.take_error() {
return Err(error);
}
config.validate().map_err(|e| BackendError::InvalidRequest(e.to_string()))?;
check_method(request.method(), request.body().is_some())?;
let own: HashSet<HeaderName> = request.headers().keys().cloned().collect();
for (name, value) in config.headers() {
if !own.contains(name) {
request.headers_mut().append(name.clone(), value.clone());
}
}
if request.uses_credentials() {
if let Some(credentials) = credentials {
credentials.apply(&mut request);
}
}
if let Some(error) = request.take_error() {
return Err(error);
}
let uri = build_uri(config.base_url(), request.path(), request.query(), config.insecure_http_allowed())?;
let (method, headers, body, timeout, purpose) = request.into_parts();
Ok(PreparedRequest { method, uri, headers, body, timeout: timeout.unwrap_or(config.timeout()), max_body_bytes: config.max_body_bytes(), purpose })
}
fn check_method(method: &Method, has_body: bool) -> Result<(), BackendError> {
let standard = [Method::GET, Method::POST, Method::PUT, Method::PATCH, Method::DELETE, Method::HEAD, Method::OPTIONS, Method::TRACE];
if !standard.contains(method) {
return Err(BackendError::InvalidRequest(format!("method `{method}` is not supported (standard methods only, no CONNECT)")));
}
if has_body && *method == Method::HEAD {
return Err(BackendError::InvalidRequest("a HEAD request cannot have a body".into()));
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn send_requests(
client: Res<HttpClient>,
mut inflight: ResMut<InFlight>,
mut transport: Option<ResMut<HttpTransportRes>>,
config: Res<HttpConfig>,
credentials: Option<Res<BackendCredentials>>,
time: Option<Res<Time<Real>>>,
mut exit: MessageReader<AppExit>,
) {
inflight.cancels.age();
if exit.read().count() > 0 {
return;
}
let queued = client.drain();
let queued_ids: HashSet<RequestId> = queued.iter().map(|Queued::Send { id, .. }| *id).collect();
let cancels = {
let entries = &inflight.entries;
inflight.cancels.claim(|id| queued_ids.contains(&id) || entries.contains_key(&id))
};
if queued.is_empty() && cancels.is_empty() {
return;
}
let now = inflight.now(time.as_deref());
let cancelled: HashSet<RequestId> = cancels.iter().copied().collect();
let mut claimed: HashSet<RequestId> = HashSet::new();
for item in queued {
match item {
Queued::Send { id, route, .. } if cancelled.contains(&id) => {
claimed.insert(id);
inflight.ready.push((id, route, Err(BackendError::Cancelled)));
}
Queued::Send { id, request, route } => {
let info = RequestInfo { kind: RequestKind::Http, method: Some(request.method().clone()), target: request.path().to_string() };
match prepare(*request, &config, credentials.as_deref()) {
Err(error) => inflight.ready.push((id, route, Err(error))),
Ok(prepared) => match transport.as_mut() {
None => inflight.ready.push((id, route, Err(BackendError::NoTransport))),
Some(transport) => {
let allowed = prepared.timeout.saturating_add(DEADLINE_GRACE);
if let Some(method) = &info.method {
tracing::debug!(">>> NET-BACKEND: {id} {method} {}", info.target);
}
inflight
.entries
.insert(id, Entry { route, info, deadline: now.saturating_add(allowed), allowed, generation: transport.generation() });
transport.get_mut().submit(id, prepared);
}
},
}
}
}
}
for id in cancels {
if claimed.contains(&id) {
continue;
}
if let Some(entry) = inflight.entries.remove(&id) {
if let Some(transport) = transport.as_mut() {
transport.get_mut().cancel(id);
}
inflight.ready.push((id, entry.route, Err(BackendError::Cancelled)));
}
}
}
pub(crate) fn receive_answers(
mut inflight: ResMut<InFlight>,
mut transport: Option<ResMut<HttpTransportRes>>,
config: Res<HttpConfig>,
time: Option<Res<Time<Real>>>,
mut raw: MessageWriter<HttpResponse>,
mut commands: Commands,
) {
let mut answers = std::mem::take(&mut inflight.ready);
let generation = transport.as_ref().map(|t| t.generation());
if let Some(transport) = transport.as_mut() {
for (id, result) in transport.get_mut().poll() {
match inflight.entries.get(&id) {
Some(entry) if Some(entry.generation) == generation => {
if let Some(entry) = inflight.entries.remove(&id) {
answers.push((id, entry.route, result));
}
}
_ => tracing::debug!(">>> NET-BACKEND: {id}: a late result was discarded (already answered)"),
}
}
}
if !inflight.entries.is_empty() {
let now = inflight.now(time.as_deref());
let gone: Vec<(RequestId, Entry)> = inflight.entries.extract_if(|_, entry| Some(entry.generation) != generation).collect();
answers.extend(gone.into_iter().map(|(id, entry)| (id, entry.route, Err(BackendError::NoTransport))));
let late: Vec<(RequestId, Entry)> = inflight.entries.extract_if(|_, entry| entry.deadline <= now).collect();
for (id, entry) in late {
if let Some(transport) = transport.as_mut() {
transport.get_mut().cancel(id);
}
answers.push((id, entry.route, Err(BackendError::Timeout(format!("no answer from the transport within {:?}", entry.allowed)))));
}
}
deliver(answers, config.max_body_bytes(), &mut raw, &mut commands);
}
pub(crate) fn shutdown_on_exit(
client: Res<HttpClient>,
mut inflight: ResMut<InFlight>,
mut transport: Option<ResMut<HttpTransportRes>>,
config: Res<HttpConfig>,
mut raw: MessageWriter<HttpResponse>,
mut commands: Commands,
) {
let mut answers = std::mem::take(&mut inflight.ready);
let queued = client.drain();
let queued_ids: HashSet<RequestId> = queued.iter().map(|Queued::Send { id, .. }| *id).collect();
let cancelled: HashSet<RequestId> = {
let entries = &inflight.entries;
inflight.cancels.claim(|id| queued_ids.contains(&id) || entries.contains_key(&id)).into_iter().collect()
};
for Queued::Send { id, route, .. } in queued {
let error = if cancelled.contains(&id) { BackendError::Cancelled } else { BackendError::Shutdown };
answers.push((id, route, Err(error)));
}
for id in &cancelled {
if let Some(entry) = inflight.entries.remove(id) {
answers.push((*id, entry.route, Err(BackendError::Cancelled)));
}
}
let generation = transport.as_ref().map(|t| t.generation());
if let Some(transport) = transport.as_mut() {
for (id, result) in transport.get_mut().poll() {
if inflight.entries.get(&id).is_some_and(|e| Some(e.generation) == generation) {
if let Some(entry) = inflight.entries.remove(&id) {
answers.push((id, entry.route, result));
}
}
}
}
let open = inflight.entries.len();
answers.extend(inflight.entries.drain().map(|(id, entry)| (id, entry.route, Err(BackendError::Shutdown))));
if let Some(transport) = transport.as_mut() {
transport.get_mut().shutdown();
}
if open > 0 {
tracing::info!(">>> NET-BACKEND: app exit: {open} open request(s) answered with Shutdown");
}
deliver(answers, config.max_body_bytes(), &mut raw, &mut commands);
}
fn deliver(mut answers: Vec<Answer>, limit: u64, raw: &mut MessageWriter<HttpResponse>, commands: &mut Commands) {
answers.sort_by_key(|(id, _, _)| *id);
for (id, route, result) in answers {
let result = result.and_then(|response| {
if u64::try_from(response.body.len()).unwrap_or(u64::MAX) > limit {
Err(BackendError::BodyTooLarge { limit })
} else if response.status.is_success() {
Ok(response)
} else {
Err(BackendError::Status(Box::new(response)))
}
});
match &result {
Ok(response) => tracing::debug!(">>> NET-BACKEND: {id} -> {}", response.status),
Err(
error @ (BackendError::InvalidRequest(_)
| BackendError::InsecureHttp { .. }
| BackendError::Encode(_)
| BackendError::RequestTooLarge { .. }
| BackendError::NoTransport),
) => {
tracing::warn!(">>> NET-BACKEND: {id} not sent: {error}")
}
Err(error) => tracing::debug!(">>> NET-BACKEND: {id} -> {error}"),
}
match route {
Route::Raw => {
raw.write(HttpResponse { id, result });
}
#[cfg(feature = "json")]
Route::Json(json) => json.deliver(id, result, commands),
}
}
#[cfg(not(feature = "json"))]
let _ = commands;
}