use std::{
net::SocketAddr,
sync::{Arc, Mutex},
};
use axum::{
Router,
body::{Body, to_bytes},
extract::State,
http::{HeaderMap, Method, Request, StatusCode, header},
response::{IntoResponse, Response},
};
use saddle_core::{ApplicationId, ComponentLifecycle, ErrorKind, LifecycleFuture, SaddleError};
use saddle_observability::Observer;
use saddle_service::{
ExternalDispatcher, ExternalRequest, ExternalResponse, ExternalStatus,
MAX_EXTERNAL_REQUEST_BYTES,
};
use serde_json::json;
use tokio::{sync::oneshot, task::JoinHandle};
const TRACE_HEADER: &str = "x-saddle-trace-id";
struct RunningServer {
shutdown: oneshot::Sender<()>,
worker: JoinHandle<std::io::Result<()>>,
}
pub(crate) struct HttpServer {
listen: SocketAddr,
application: ApplicationId,
dispatcher: Arc<ExternalDispatcher>,
observer: Observer,
running: Mutex<Option<RunningServer>>,
}
impl HttpServer {
pub(crate) fn new(
listen: SocketAddr,
application: ApplicationId,
dispatcher: ExternalDispatcher,
observer: Observer,
) -> Self {
Self {
listen,
application,
dispatcher: Arc::new(dispatcher),
observer,
running: Mutex::new(None),
}
}
}
impl ComponentLifecycle for HttpServer {
fn name(&self) -> &'static str {
"http"
}
fn start(&self) -> LifecycleFuture<'_> {
Box::pin(async move {
if self.running.lock().unwrap().is_some() {
return Err(http_error(
"http.already_started",
"HTTP server has already started",
));
}
let listener = tokio::net::TcpListener::bind(self.listen)
.await
.map_err(|_| http_error("http.bind_failed", "HTTP listener could not bind"))?;
let state = HttpState {
application: self.application.clone(),
dispatcher: Arc::clone(&self.dispatcher),
observer: self.observer.clone(),
};
let router = Router::new().fallback(dispatch).with_state(state);
let (shutdown, shutdown_requested) = oneshot::channel();
let worker = tokio::spawn(async move {
axum::serve(listener, router)
.with_graceful_shutdown(async move {
let _ = shutdown_requested.await;
})
.await
});
*self.running.lock().unwrap() = Some(RunningServer { shutdown, worker });
Ok(())
})
}
fn shutdown(&self) -> LifecycleFuture<'_> {
Box::pin(async move {
let running = self.running.lock().unwrap().take();
let Some(running) = running else {
return Err(http_error(
"http.not_started",
"HTTP server was not running during shutdown",
));
};
let _ = running.shutdown.send(());
running
.worker
.await
.map_err(|_| http_error("http.worker_failed", "HTTP server task failed"))?
.map_err(|_| http_error("http.serve_failed", "HTTP server stopped with an error"))
})
}
}
#[derive(Clone)]
struct HttpState {
application: ApplicationId,
dispatcher: Arc<ExternalDispatcher>,
observer: Observer,
}
async fn dispatch(State(state): State<HttpState>, request: Request<Body>) -> Response {
if request.method() != Method::POST {
return reject_method(&state, request.headers());
}
if !is_json(request.headers()) {
return reject_content_type(&state, request.headers());
}
let route = request.uri().path().to_owned();
let inbound_trace = trace_header(request.headers());
let body = match to_bytes(request.into_body(), MAX_EXTERNAL_REQUEST_BYTES).await {
Ok(body) => body,
Err(_) => return reject_body(&state, inbound_trace.as_deref()),
};
let response = state
.dispatcher
.dispatch(ExternalRequest {
route: &route,
inbound_trace_id: inbound_trace.as_deref(),
body: &body,
})
.await;
external_response(response)
}
fn reject_content_type(state: &HttpState, headers: &HeaderMap) -> Response {
let (call, _) = state.observer.start_external_call(
state.application.clone(),
"saddle",
"http",
"content_type_rejected",
trace_header(headers).as_deref(),
);
let trace = call.context().trace_id();
let error = SaddleError::new(
ErrorKind::InvalidArgument,
"http.unsupported_media_type",
"Saddle HTTP endpoints require application/json",
);
call.fail(&error);
json_error(
StatusCode::UNSUPPORTED_MEDIA_TYPE,
error.code(),
Some(trace.to_string()),
)
}
fn reject_method(state: &HttpState, headers: &HeaderMap) -> Response {
let (call, _) = state.observer.start_external_call(
state.application.clone(),
"saddle",
"http",
"method_rejected",
trace_header(headers).as_deref(),
);
let trace = call.context().trace_id();
let error = SaddleError::new(
ErrorKind::InvalidArgument,
"http.method_not_allowed",
"Saddle HTTP endpoints accept POST only",
);
call.fail(&error);
json_error(
StatusCode::METHOD_NOT_ALLOWED,
error.code(),
Some(trace.to_string()),
)
}
fn reject_body(state: &HttpState, inbound_trace: Option<&str>) -> Response {
let (call, _) = state.observer.start_external_call(
state.application.clone(),
"saddle",
"http",
"body_rejected",
inbound_trace,
);
let trace = call.context().trace_id();
let error = SaddleError::new(
ErrorKind::InvalidArgument,
"service.request_too_large",
"the external request exceeds the request limit",
);
call.fail(&error);
json_error(
StatusCode::PAYLOAD_TOO_LARGE,
error.code(),
Some(trace.to_string()),
)
}
fn external_response(response: ExternalResponse) -> Response {
render_response(
response.status(),
response.body(),
response.error_code(),
response.trace_id().map(|trace| trace.to_string()),
)
}
fn render_response(
external_status: ExternalStatus,
body: Option<&[u8]>,
error_code: Option<&'static str>,
trace: Option<String>,
) -> Response {
let status = match external_status {
ExternalStatus::Ok => StatusCode::OK,
ExternalStatus::InvalidArgument => StatusCode::BAD_REQUEST,
ExternalStatus::NotFound => StatusCode::NOT_FOUND,
ExternalStatus::Conflict => StatusCode::CONFLICT,
ExternalStatus::Business => StatusCode::UNPROCESSABLE_ENTITY,
ExternalStatus::Unavailable => StatusCode::SERVICE_UNAVAILABLE,
ExternalStatus::Internal => StatusCode::INTERNAL_SERVER_ERROR,
};
if let Some(body) = body {
let mut output = (
status,
[(header::CONTENT_TYPE, "application/json")],
body.to_vec(),
)
.into_response();
insert_trace(&mut output, trace.as_deref());
output
} else {
json_error(status, error_code.unwrap_or("saddle.internal"), trace)
}
}
fn json_error(status: StatusCode, code: &'static str, trace: Option<String>) -> Response {
let body = serde_json::to_vec(&json!({
"error": { "code": code },
"trace_id": trace,
}))
.expect("fixed error response must serialize");
let mut response = (status, [(header::CONTENT_TYPE, "application/json")], body).into_response();
insert_trace(&mut response, trace.as_deref());
response
}
fn insert_trace(response: &mut Response, trace: Option<&str>) {
if let Some(trace) = trace
&& let Ok(value) = trace.parse()
{
response.headers_mut().insert(TRACE_HEADER, value);
}
}
fn trace_header(headers: &HeaderMap) -> Option<String> {
headers.get(TRACE_HEADER).map(|value| {
value.to_str().map_or_else(|_| String::new(), str::to_owned)
})
}
fn is_json(headers: &HeaderMap) -> bool {
headers
.get(header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.split(';').next())
.is_some_and(|media_type| media_type.trim().eq_ignore_ascii_case("application/json"))
}
fn http_error(code: &'static str, message: &'static str) -> SaddleError {
SaddleError::new(ErrorKind::Infrastructure, code, message)
}
#[cfg(test)]
mod tests {
use std::{io, sync::Arc};
use axum::body::to_bytes;
use saddle_observability::ObserverConfig;
use saddle_runtime::Application;
use saddle_service::{ExternalDispatcherBuilder, ServiceRegistryBuilder};
use super::*;
fn state() -> HttpState {
let observer = Observer::with_writer(ObserverConfig::default(), io::sink()).unwrap();
let registry = ServiceRegistryBuilder::new()
.build(observer.clone())
.unwrap();
let application = Application::new();
let dispatcher =
ExternalDispatcherBuilder::new("test", registry, application.request_lifecycle())
.build();
HttpState {
application: "test".into(),
dispatcher: Arc::new(dispatcher),
observer,
}
}
fn server(listen: SocketAddr) -> HttpServer {
let state = state();
HttpServer {
listen,
application: state.application,
dispatcher: state.dispatcher,
observer: state.observer,
running: Mutex::new(None),
}
}
async fn error_code(response: Response) -> String {
let body = to_bytes(response.into_body(), 4096).await.unwrap();
serde_json::from_slice::<serde_json::Value>(&body).unwrap()["error"]["code"]
.as_str()
.unwrap()
.to_owned()
}
#[tokio::test]
async fn transport_rejects_non_post_and_missing_json_content_type() {
let state = state();
let response = dispatch(
State(state.clone()),
Request::builder()
.method(Method::GET)
.body(Body::empty())
.unwrap(),
)
.await;
assert_eq!(response.status(), StatusCode::METHOD_NOT_ALLOWED);
assert_eq!(error_code(response).await, "http.method_not_allowed");
let response = dispatch(
State(state),
Request::builder()
.method(Method::POST)
.body(Body::empty())
.unwrap(),
)
.await;
assert_eq!(response.status(), StatusCode::UNSUPPORTED_MEDIA_TYPE);
assert_eq!(error_code(response).await, "http.unsupported_media_type");
}
#[tokio::test]
async fn transport_bounds_request_before_dispatch() {
let state = state();
let response = dispatch(
State(state),
Request::builder()
.method(Method::POST)
.header(header::CONTENT_TYPE, "application/json; charset=utf-8")
.body(Body::from(vec![b'x'; MAX_EXTERNAL_REQUEST_BYTES + 1]))
.unwrap(),
)
.await;
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
assert_eq!(error_code(response).await, "service.request_too_large");
}
#[tokio::test]
async fn managed_server_binds_and_closes_through_component_lifecycle() {
let server = server("127.0.0.1:0".parse().unwrap());
ComponentLifecycle::start(&server).await.unwrap();
ComponentLifecycle::shutdown(&server).await.unwrap();
assert_eq!(
ComponentLifecycle::shutdown(&server)
.await
.unwrap_err()
.code(),
"http.not_started"
);
}
#[tokio::test]
async fn transport_maps_success_and_failure_without_internal_messages() {
let response = render_response(
ExternalStatus::Ok,
Some(br#"{"value":"ok"}"#),
None,
Some("00000000000000000000000000000001".to_owned()),
);
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()[header::CONTENT_TYPE], "application/json");
assert_eq!(
response.headers()[TRACE_HEADER],
"00000000000000000000000000000001"
);
let response = render_response(
ExternalStatus::Business,
None,
Some("orders.rejected"),
Some("00000000000000000000000000000002".to_owned()),
);
assert_eq!(response.status(), StatusCode::UNPROCESSABLE_ENTITY);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let body: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(body["error"]["code"], "orders.rejected");
assert!(body.get("message").is_none());
}
}