use bytes::Bytes;
use http_body_util::{BodyExt, Limited};
use hyper::body::Incoming;
use hyper::{Request, StatusCode};
use zygo_core::supervisor::WorkspaceRequest;
use super::reply::HttpError;
const MAX_BODY_BYTES: usize = 16 * 1024 * 1024;
pub(super) const DEFAULT_TIMEOUT_MS: u64 = 60_000;
pub(super) const MAX_TIMEOUT_MS: u64 = 24 * 3_600_000;
pub(super) const REQUEST_KEY_HEADER: &str = "x-zygo-request-key";
pub(super) const REQUEST_ID_HEADER: &str = "x-zygo-request-id";
pub(super) fn request_key(req: &Request<Incoming>) -> Result<Option<String>, HttpError> {
let Some(value) = req.headers().get(REQUEST_KEY_HEADER) else {
return Ok(None);
};
let key = value
.to_str()
.map(str::trim)
.ok()
.filter(|k| !k.is_empty())
.filter(|k| k.len() <= 128)
.filter(|k| k.bytes().all(|b| b.is_ascii_graphic()))
.ok_or_else(|| {
HttpError::closing(
StatusCode::BAD_REQUEST,
"X-Zygo-Request-Key must be 1 to 128 printable ASCII characters",
)
})?;
Ok(Some(key.to_string()))
}
pub(super) fn timeout_header(req: &Request<Incoming>) -> Result<u64, HttpError> {
timeout_header_from(req.headers())
}
fn timeout_header_from(headers: &hyper::HeaderMap) -> Result<u64, HttpError> {
match headers.get("x-zygo-timeout-ms") {
None => Ok(DEFAULT_TIMEOUT_MS),
Some(v) => {
let ms = v
.to_str()
.ok()
.and_then(|s| s.trim().parse::<u64>().ok())
.filter(|ms| *ms > 0)
.ok_or_else(|| {
HttpError::closing(
StatusCode::BAD_REQUEST,
"X-Zygo-Timeout-Ms must be a positive integer",
)
})?;
if ms > MAX_TIMEOUT_MS {
return Err(HttpError::closing(
StatusCode::BAD_REQUEST,
format!(
"X-Zygo-Timeout-Ms is {ms}, over this API's ceiling of \
{MAX_TIMEOUT_MS}\n → the function's own `timeout` is the \
real limit; this header only says how long you will wait"
),
));
}
Ok(ms)
}
}
}
pub(super) async fn read_body(req: Request<Incoming>) -> Result<Bytes, HttpError> {
Limited::new(req.into_body(), MAX_BODY_BYTES)
.collect()
.await
.map(|c| c.to_bytes())
.map_err(|e| {
HttpError::new(
StatusCode::PAYLOAD_TOO_LARGE,
format!("body over {MAX_BODY_BYTES} bytes or unreadable: {e}"),
)
})
}
pub(super) fn parse_event(body: &[u8]) -> Result<serde_json::Value, HttpError> {
if body.iter().all(u8::is_ascii_whitespace) {
return Ok(serde_json::Value::Null);
}
serde_json::from_slice(body)
.map_err(|e| HttpError::new(StatusCode::BAD_REQUEST, format!("body is not JSON: {e}")))
}
pub(super) struct Query<'a>(Vec<(&'a str, &'a str)>);
impl<'a> Query<'a> {
pub(super) fn parse(query: &'a str) -> Query<'a> {
Query(
query
.split('&')
.filter(|p| !p.is_empty())
.map(|pair| match pair.split_once('=') {
Some((k, v)) => (k, v),
None => (pair, ""),
})
.collect(),
)
}
fn get(&self, key: &str) -> Option<&'a str> {
self.0.iter().find(|(k, _)| *k == key).map(|(_, v)| *v)
}
pub(super) fn number(&self, key: &str) -> Result<Option<u64>, HttpError> {
match self.get(key) {
None => Ok(None),
Some(raw) => raw.parse().map(Some).map_err(|_| {
HttpError::new(
StatusCode::BAD_REQUEST,
format!("`{key}` must be a non-negative integer, not `{raw}`"),
)
}),
}
}
pub(super) fn blob(&self, key: &str) -> Result<Option<WorkspaceRequest>, HttpError> {
let Some(digest) = self.get(key).filter(|d| !d.is_empty()) else {
return Ok(None);
};
zygo_core::scripts::ScriptDigest::parse(digest)
.map_err(|e| HttpError::new(StatusCode::BAD_REQUEST, e.to_string()))?;
Ok(Some(WorkspaceRequest {
blob: Some(digest.to_string()),
..WorkspaceRequest::default()
}))
}
pub(super) fn flag(&self, key: &str) -> Result<bool, HttpError> {
match self.get(key) {
None => Ok(false),
Some("") | Some("1") | Some("true") => Ok(true),
Some("0") | Some("false") => Ok(false),
Some(raw) => Err(HttpError::new(
StatusCode::BAD_REQUEST,
format!("`{key}` must be true or false, not `{raw}`"),
)),
}
}
}
pub(super) struct CallParams {
pub(super) timeout_ms: u64,
pub(super) tenant: Option<String>,
pub(super) key: Option<String>,
pub(super) streaming: bool,
pub(super) out: bool,
}
impl CallParams {
pub(super) fn parse(
req: &Request<Incoming>,
tenant: Option<String>,
) -> Result<CallParams, HttpError> {
let timeout_ms = timeout_header(req)?;
let key = request_key(req)?;
let query = Query::parse(req.uri().query().unwrap_or(""));
let streaming = query.flag("stream")?;
let out = query.flag("out")?;
Ok(CallParams {
timeout_ms,
tenant,
key,
streaming,
out,
})
}
}
#[cfg(test)]
mod tests {
use super::super::MAX_IDLE_CLIENTS;
use super::super::routes_fn::{BATCH_IN_FLIGHT, MAX_BATCH};
use super::*;
#[test]
fn an_empty_body_is_a_null_event() {
assert_eq!(parse_event(b"").unwrap(), serde_json::Value::Null);
assert_eq!(parse_event(b" \n").unwrap(), serde_json::Value::Null);
assert_eq!(parse_event(br#"{"n":1}"#).unwrap()["n"], 1);
assert!(parse_event(b"{nope").is_err());
}
#[test]
fn log_query_parameters_are_parsed_or_refused() {
let query = Query::parse("after=12&limit=5&failed=true");
assert_eq!(query.number("after").unwrap(), Some(12));
assert_eq!(query.number("limit").unwrap(), Some(5));
assert!(query.flag("failed").unwrap());
let empty = Query::parse("");
assert_eq!(empty.number("after").unwrap(), None);
assert!(!empty.flag("failed").unwrap());
assert!(Query::parse("failed").flag("failed").unwrap());
assert!(!Query::parse("failed=false").flag("failed").unwrap());
assert!(Query::parse("after=soon").number("after").is_err());
assert!(Query::parse("failed=yes").flag("failed").is_err());
}
#[test]
fn a_caller_cannot_ask_for_an_unbounded_amount_of_work() {
let over = Request::builder()
.header("x-zygo-timeout-ms", (MAX_TIMEOUT_MS + 1).to_string())
.body(())
.expect("a request");
let refused = timeout_header_from(over.headers()).expect_err("over the ceiling");
assert_eq!(refused.status, StatusCode::BAD_REQUEST);
assert!(
refused.body["error"]
.as_str()
.expect("a message")
.contains(&MAX_TIMEOUT_MS.to_string()),
"the refusal should name the ceiling: {:?}",
refused.body
);
let at_the_line = Request::builder()
.header("x-zygo-timeout-ms", MAX_TIMEOUT_MS.to_string())
.body(())
.expect("a request");
assert_eq!(
timeout_header_from(at_the_line.headers()).expect("the ceiling itself is allowed"),
MAX_TIMEOUT_MS
);
let ordinary = Request::builder()
.header("x-zygo-timeout-ms", "2500")
.body(())
.expect("a request");
assert_eq!(
timeout_header_from(ordinary.headers()).expect("an ordinary header"),
2500
);
const _: () = assert!(
BATCH_IN_FLIGHT <= MAX_IDLE_CLIENTS,
"one batch may take every pooled connection"
);
const _: () = assert!(MAX_BATCH >= BATCH_IN_FLIGHT, "the cap is below the gate");
}
}