use std::time::Duration;
use lgwks_std::http::{self, BodyPolicy, Options};
use crate::cap::{Auth, Cap};
use crate::error::{BotError, DispatchCertainty};
use crate::verb;
pub const BODY_PREVIEW: usize = 4096;
const BODY_PREVIEW_BYTES: usize = BODY_PREVIEW.saturating_mul(4);
const POLL_TIMEOUT_SECS: u64 = 10;
#[derive(Debug)]
pub struct Endpoint {
url: String,
caps: Vec<Cap>,
}
#[derive(PartialEq, Debug, Clone)]
#[non_exhaustive]
pub struct NetState {
status_code: u16,
reachable: bool,
pub(crate) body: String,
}
impl NetState {
#[must_use]
pub fn new(status_code: u16, reachable: bool, body: impl Into<String>) -> Self {
Self {
status_code,
reachable,
body: body.into(),
}
}
#[must_use]
pub const fn status_code(&self) -> u16 {
self.status_code
}
#[must_use]
pub const fn reachable(&self) -> bool {
self.reachable
}
#[must_use]
pub fn body(&self) -> &str {
&self.body
}
}
impl Endpoint {
pub fn new(url: impl Into<String>) -> Self {
Self {
url: url.into(),
caps: vec![Cap::net()],
}
}
}
impl verb::Observe for Endpoint {
type Output = NetState;
fn required_caps(&self) -> &[Cap] {
&self.caps
}
async fn poll(&self, call: (Auth, ())) -> Result<NetState, BotError> {
call.0.check(self.required_caps())?;
let url = self.url.clone();
let options = Options::default()
.timeout(Duration::from_secs(POLL_TIMEOUT_SECS))
.deadline(Duration::from_secs(POLL_TIMEOUT_SECS))
.max_body_bytes(BODY_PREVIEW_BYTES)
.body_policy(BodyPolicy::Preview);
let request = lgwks_std::task::try_spawn_blocking(move || http::get_with(&url, &options))
.map_err(|error| {
lgwks_std::trace::debug!(%error, "poll: the blocking pool refused the request");
BotError::DomainError {
domain: self.domain_id().into(),
certainty: DispatchCertainty::Refused,
cause: format!("no thread was available for the request: {error}"),
}
})?;
let exchange = request.await;
match exchange {
Ok(response) => Ok(NetState::new(
response.status,
true,
response
.text_lossy()
.chars()
.take(BODY_PREVIEW)
.collect::<String>(),
)),
Err(http::Error::InvalidUrl) => Err(BotError::DomainError {
domain: self.domain_id().into(),
certainty: DispatchCertainty::Refused,
cause: "invalid endpoint URL (absolute http(s) URI required)".into(),
}),
Err(_) => Ok(NetState::new(0, false, String::new())),
}
}
fn domain_id(&self) -> &str {
"net::endpoint"
}
}
impl verb::Query for Endpoint {
type Input = ();
type Output = NetState;
fn required_caps(&self) -> &[Cap] {
&self.caps
}
async fn query(&self, call: (Auth, &())) -> Result<NetState, BotError> {
let (auth, _) = call;
verb::Observe::poll(self, (auth, ())).await
}
fn domain_id(&self) -> &str {
"net::endpoint"
}
}
pub fn endpoint(url: impl Into<String>) -> Endpoint {
Endpoint::new(url)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::gate::GrantSet;
use std::io::{Read, Write};
use std::net::TcpListener;
type TestResult = Result<(), Box<dyn std::error::Error>>;
type PollResult = Result<NetState, Box<dyn std::error::Error>>;
fn net_auth() -> Result<Auth, BotError> {
GrantSet::empty().grant(Cap::net()).issue(&[Cap::net()])
}
fn serve_once(
body: impl Into<String>,
) -> std::io::Result<(u16, lgwks_std::task::JoinHandle<std::io::Result<()>>)> {
let body = body.into();
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let handle = lgwks_std::task::spawn_blocking(move || -> std::io::Result<()> {
let (mut stream, _) = listener.accept()?;
let mut request = vec![0u8; 1024];
let mut head = Vec::new();
loop {
let read = stream.read(&mut request)?;
head.extend_from_slice(&request[..read]);
if head.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
let reply = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
stream.write_all(reply.as_bytes())
});
Ok((port, handle))
}
fn poll_body(body: impl Into<String>) -> PollResult {
use crate::verb::Observe;
let (port, server) = serve_once(body)?;
let state = lgwks_std::task::block_on(
Endpoint::new(format!("http://127.0.0.1:{port}/")).poll((net_auth()?, ())),
)?;
lgwks_std::task::block_on(server)?;
Ok(state)
}
#[test]
fn poll_reports_reachable_loopback() -> TestResult {
use crate::verb::Observe;
let (port, server) = serve_once("alive")?;
let state = lgwks_std::task::block_on(
Endpoint::new(format!("http://127.0.0.1:{port}/")).poll((net_auth()?, ())),
)?;
assert!(
state.reachable(),
"a loopback server that answered must report reachable"
);
assert_eq!(state.status_code(), 200);
assert_eq!(state.body, "alive");
lgwks_std::task::block_on(server)?;
Ok(())
}
#[test]
fn poll_cuts_a_body_larger_than_the_preview() -> TestResult {
let payload = "x".repeat(BODY_PREVIEW_BYTES.saturating_add(64));
let state = poll_body(payload)?;
assert!(
state.reachable(),
"a server that answered with a large body is reachable"
);
assert_eq!(
state.status_code(),
200,
"the status survives the body being cut"
);
assert_eq!(
state.body.chars().count(),
BODY_PREVIEW,
"the preview stops at BODY_PREVIEW characters"
);
assert!(
state.body.chars().all(|character| character == 'x'),
"the preview is the body's own prefix"
);
Ok(())
}
#[test]
fn poll_stops_reading_at_its_preview_ceiling() -> TestResult {
use crate::verb::Observe;
let payload = "x".repeat(BODY_PREVIEW_BYTES.saturating_mul(256));
let (port, server) = serve_once(payload)?;
let state = lgwks_std::task::block_on(
Endpoint::new(format!("http://127.0.0.1:{port}/")).poll((net_auth()?, ())),
)?;
assert!(
state.reachable(),
"a server that answered with a huge body is reachable"
);
assert_eq!(
state.body.chars().count(),
BODY_PREVIEW,
"the preview stops at BODY_PREVIEW characters"
);
let served = lgwks_std::task::block_on(server);
assert!(
served.is_err(),
"a 4 MiB body can only be written in full to a client that keeps reading, so a \
completed write means the probe drained it: a preview must stop at its ceiling"
);
Ok(())
}
#[test]
fn poll_previews_a_multibyte_body_without_panicking() -> TestResult {
let payload = "€".repeat(BODY_PREVIEW.saturating_add(1));
let state = poll_body(payload)?;
assert!(
state.reachable(),
"a multibyte body is still a reachable server"
);
assert_eq!(state.status_code(), 200);
assert_eq!(
state.body.chars().count(),
BODY_PREVIEW,
"the character cap holds on the decoded preview"
);
assert!(
state.body.chars().all(|character| character == '€'),
"a cut inside a multi-byte character must not disturb the prefix: {:?}",
state.body
);
Ok(())
}
#[test]
fn poll_keeps_a_body_exactly_at_the_preview() -> TestResult {
let payload = "y".repeat(BODY_PREVIEW);
let state = poll_body(payload)?;
assert_eq!(
state.body.chars().count(),
BODY_PREVIEW,
"a body of exactly BODY_PREVIEW characters is kept whole"
);
assert!(
state.body.chars().all(|character| character == 'y'),
"no character of an at-the-limit body is dropped or replaced"
);
Ok(())
}
#[test]
fn poll_reports_unreachable_closed_port() -> TestResult {
use crate::verb::Observe;
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
let state = lgwks_std::task::block_on(
Endpoint::new(format!("http://127.0.0.1:{port}/")).poll((net_auth()?, ())),
)?;
assert!(
!state.reachable(),
"a refused connection must report unreachable, not error"
);
assert_eq!(state.status_code(), 0);
Ok(())
}
#[test]
fn ten_thousand_observers_are_refused_typed_when_the_pool_is_full_and_answered_when_free()
-> TestResult {
use crate::verb::Observe;
use std::sync::{Arc, Condvar, Mutex};
const OBSERVERS: usize = 10_000;
let gate = Arc::new((Mutex::new(false), Condvar::new()));
let mut parked = Vec::new();
loop {
let held = Arc::clone(&gate);
let job = lgwks_std::task::try_spawn_blocking(move || {
let (ref open, ref opened) = *held;
let mut is_open = crate::journal::owner::lock(open);
while !*is_open {
is_open = crate::journal::owner::wait(opened, is_open);
}
});
match job {
Ok(handle) => parked.push(handle),
Err(lgwks_std::task::SpawnError::AtCapacity { .. }) => break,
Err(other) => return Err(other.into()),
}
}
let listener = TcpListener::bind("127.0.0.1:0")?;
listener.set_nonblocking(true)?;
let live = Endpoint::new(format!(
"http://127.0.0.1:{}/",
listener.local_addr()?.port()
));
let auth = net_auth()?;
let full = lgwks_std::task::block_on(lgwks_std::task::join_all(
(0..OBSERVERS).map(|_| live.poll((auth.clone(), ()))),
));
let refused = full
.iter()
.filter(|outcome| {
matches!(outcome, Err(BotError::DomainError {
certainty: DispatchCertainty::Refused,
cause,
..
}) if cause.starts_with("no thread was available"))
})
.count();
assert_eq!(
refused, OBSERVERS,
"every poll against a full pool is refused, typed"
);
assert!(
matches!(listener.accept(), Err(ref error) if error.kind() == std::io::ErrorKind::WouldBlock),
"a refused poll sent nothing: no connection reached the listener"
);
{
let (ref open, ref opened) = *gate;
*crate::journal::owner::lock(open) = true;
opened.notify_all();
}
drop(lgwks_std::task::block_on(lgwks_std::task::join_all(parked)));
let closed = {
let probe = TcpListener::bind("127.0.0.1:0")?;
probe.local_addr()?.port()
};
let free = Endpoint::new(format!("http://127.0.0.1:{closed}/"));
let answered = lgwks_std::task::block_on(lgwks_std::task::join_all(
(0..OBSERVERS).map(|_| free.poll((auth.clone(), ()))),
));
let unreachable = answered
.iter()
.filter(|outcome| matches!(outcome, Ok(state) if !state.reachable()))
.count();
assert_eq!(
unreachable, OBSERVERS,
"with the pool free every poll runs to an answer, none refused"
);
Ok(())
}
#[test]
fn poll_rejects_malformed_url_as_spec_bug() -> TestResult {
use crate::verb::Observe;
let Err(error) =
lgwks_std::task::block_on(Endpoint::new("not a url").poll((net_auth()?, ())))
else {
return Err("a malformed URL is a spec bug and must error".into());
};
assert!(
matches!(error, BotError::DomainError { .. }),
"a malformed URL must be a typed DomainError, got {error:?}"
);
Ok(())
}
#[test]
fn poll_without_net_proof_is_denied() -> TestResult {
use crate::verb::Observe;
let vacuous = GrantSet::empty().issue(&[])?;
let Err(error) =
lgwks_std::task::block_on(Endpoint::new("http://127.0.0.1:9/").poll((vacuous, ())))
else {
return Err("a proof covering no capability must not authorize `bot.net`".into());
};
assert!(
matches!(error, BotError::CapabilityDenied { .. }),
"a capped endpoint must deny a vacuous proof, got {error:?}"
);
Ok(())
}
}