use apimock_config::Config;
use apimock_config::config::constant::{
SERVICE_DEFAULT_MAX_REQUEST_BODY_BYTES, SERVICE_DEFAULT_MIDDLEWARE_MAX_OPERATIONS,
TLS_DEFAULT_HANDSHAKE_TIMEOUT_SECONDS, TLS_DEFAULT_MAX_CONNECTIONS,
};
use apimock_routing::ParsedRequest;
use console::style;
use http_body_util::{BodyExt, Empty};
use hyper::{
HeaderMap, Response, body,
header::{CONTENT_LENGTH, HeaderValue},
service::service_fn,
};
use hyper_util::{
rt::{TokioExecutor, TokioIo},
server::conn::auto::Builder,
};
use tokio::net::TcpListener;
use tokio::sync::Semaphore;
use tokio_rustls::TlsAcceptor;
use std::net::{SocketAddr, ToSocketAddrs};
use std::path::PathBuf;
use std::sync::Arc;
use crate::{
dyn_route::dyn_route_content_traced,
error::{ServerError, ServerResult},
middleware::LoadedMiddlewares,
parsed_request::{ParsedRequestError, capture_in_log_with_trace_config, parsed_request_from},
respond_response::respond_response,
response::{
confine::canonical_dir,
error_response::{internal_server_error_response, payload_too_large_response},
},
response_handler::default_response_headers,
tls::{build_server_config_reloadable, load_certs, load_private_key},
types::BoxBody,
};
pub use crate::control::{ReloadHint, ServerControl, ServerHandle, ServerState};
use crate::trace::{Outcome, RequestSummary, TraceEmitter};
#[non_exhaustive]
pub struct AppState {
pub config: Config,
pub middlewares: LoadedMiddlewares,
pub tracer: TraceEmitter,
canonical_fallback_respond_dir: Option<PathBuf>,
canonical_rule_set_respond_dirs: Vec<Option<PathBuf>>,
}
impl AppState {
pub fn new(config: Config, middlewares: LoadedMiddlewares, tracer: TraceEmitter) -> Self {
let canonical_fallback_respond_dir =
canonical_dir(config.service.fallback_respond_dir.as_str());
let canonical_rule_set_respond_dirs = config
.service
.rule_sets
.iter()
.map(|rule_set| canonical_dir(rule_set.dir_prefix().as_str()))
.collect();
Self {
config,
middlewares,
tracer,
canonical_fallback_respond_dir,
canonical_rule_set_respond_dirs,
}
}
}
#[non_exhaustive]
pub struct Server {
pub app_state: Arc<AppState>,
pub http_addr: Option<SocketAddr>,
pub https_addr: Option<SocketAddr>,
https_tls: Option<HttpsTls>,
}
#[derive(Clone)]
struct HttpsTls {
acceptor: TlsAcceptor,
handshake_timeout: std::time::Duration,
max_connections: usize,
}
impl Server {
pub async fn new(config: Config) -> ServerResult<Self> {
let http_addr = resolve_listener(config.listener_http_addr().as_deref())?;
let https_addr = resolve_listener(config.listener_https_addr().as_deref())?;
let https_tls = match https_addr {
Some(addr) => Some(build_https_tls(&config, addr)?),
None => None,
};
let relative_dir_path = config
.current_dir_to_parent_dir_relative_path()
.map_err(ServerError::Config)?;
let middleware_max_operations = config
.service
.middleware_max_operations
.unwrap_or(SERVICE_DEFAULT_MIDDLEWARE_MAX_OPERATIONS);
let middlewares = LoadedMiddlewares::compile(
config
.service
.middlewares_file_paths
.as_deref()
.unwrap_or(&[]),
relative_dir_path.as_str(),
middleware_max_operations,
)?;
if !middlewares.is_empty() {
log::info!("middleware is activated: {} file(s)", middlewares.len());
}
Ok(Server {
http_addr,
https_addr,
https_tls,
app_state: Arc::new(AppState::new(config, middlewares, TraceEmitter::new())),
})
}
pub async fn start(&self) {
let http = self.http_start();
let https = self.https_start();
tokio::join!(http, https);
}
pub async fn bind_http(&self) -> ServerResult<Option<TcpListener>> {
let Some(addr) = self.http_addr else {
return Ok(None);
};
let listener =
TcpListener::bind(addr)
.await
.map_err(|err| ServerError::ListenerAddress {
addr: addr.to_string(),
reason: err.to_string(),
})?;
Ok(Some(listener))
}
pub async fn serve_http(&self, listener: TcpListener) {
if let Ok(addr) = listener.local_addr() {
log::info!(
"Greetings from apimock-rs (API Mock) !!\nListening on {} ...\n",
style(format!("http://{}", addr)).cyan()
);
}
let app_state = Arc::clone(&self.app_state);
loop {
let (stream, _) = match listener.accept().await {
Ok(pair) => pair,
Err(err) => {
log::error!("HTTP accept failed: {}", err);
continue;
}
};
let io = TokioIo::new(stream);
let app_state = Arc::clone(&app_state);
tokio::task::spawn(async move {
if let Err(err) = Builder::new(TokioExecutor::new())
.serve_connection(
io,
service_fn(move |request: hyper::Request<body::Incoming>| {
service(request, app_state.clone())
}),
)
.await
{
log::error!("{} to build connection: {:?}", style("failed").red(), err);
}
});
}
}
async fn http_start(&self) {
match self.bind_http().await {
Ok(Some(listener)) => self.serve_http(listener).await,
Ok(None) => (),
Err(err) => log::error!("{}", err),
}
}
pub async fn bind_https(&self) -> ServerResult<Option<(TcpListener, TlsAcceptor)>> {
let (Some(addr), Some(https_tls)) = (self.https_addr, self.https_tls.as_ref()) else {
return Ok(None);
};
let listener =
TcpListener::bind(addr)
.await
.map_err(|err| ServerError::ListenerAddress {
addr: addr.to_string(),
reason: err.to_string(),
})?;
Ok(Some((listener, https_tls.acceptor.clone())))
}
pub async fn serve_https(&self, listener: TcpListener, acceptor: TlsAcceptor) {
if let Ok(addr) = listener.local_addr() {
log::info!(
"Greetings from apimock-rs (API Mock) !!\nListening on {} ...\n",
style(format!("https://{}", addr)).cyan()
);
}
let (handshake_timeout, max_connections) = self
.https_tls
.as_ref()
.map(|t| (t.handshake_timeout, t.max_connections))
.unwrap_or((
std::time::Duration::from_secs(TLS_DEFAULT_HANDSHAKE_TIMEOUT_SECONDS),
TLS_DEFAULT_MAX_CONNECTIONS,
));
let connection_slots = Arc::new(Semaphore::new(max_connections));
let app_state = Arc::clone(&self.app_state);
loop {
let (stream, _) = match listener.accept().await {
Ok(pair) => pair,
Err(err) => {
log::error!("HTTPS accept failed: {}", err);
continue;
}
};
let Ok(permit) = Arc::clone(&connection_slots).acquire_owned().await else {
continue;
};
let acceptor = acceptor.clone();
let app_state = Arc::clone(&app_state);
tokio::spawn(async move {
let _permit = permit; let tls_stream =
match tokio::time::timeout(handshake_timeout, acceptor.accept(stream)).await {
Ok(Ok(s)) => s,
Ok(Err(e)) => {
log::error!("TLS handshake failed: {:?}", e);
return;
}
Err(_elapsed) => {
log::error!(
"TLS handshake timed out after {:?}; dropping connection",
handshake_timeout
);
return;
}
};
let io = TokioIo::new(tls_stream);
let app_state = app_state.clone();
tokio::task::spawn(async move {
if let Err(err) = Builder::new(TokioExecutor::new())
.serve_connection(
io,
service_fn(move |request: hyper::Request<body::Incoming>| {
service(request, app_state.clone())
}),
)
.await
{
log::error!("{} to build connection: {:?}", style("failed").red(), err);
}
});
});
}
}
async fn https_start(&self) {
match self.bind_https().await {
Ok(Some((listener, acceptor))) => self.serve_https(listener, acceptor).await,
Ok(None) => (),
Err(err) => log::error!("{}", err),
}
}
}
fn build_https_tls(config: &Config, addr: SocketAddr) -> ServerResult<HttpsTls> {
let tls = config
.listener
.as_ref()
.and_then(|l| l.tls.as_ref())
.cloned()
.ok_or_else(|| ServerError::ListenerAddress {
addr: addr.to_string(),
reason: "internal: HTTPS listener scheduled without TLS config".to_owned(),
})?;
let certs = load_certs(tls.cert.as_str())?;
let key = load_private_key(tls.key.as_str())?;
let (tls_config, resolver) =
build_server_config_reloadable(certs, key).map_err(|err| ServerError::ListenerAddress {
addr: addr.to_string(),
reason: format!("failed to build TLS config: {}", err),
})?;
let acceptor = TlsAcceptor::from(Arc::new(tls_config));
drop(resolver);
let handshake_timeout = std::time::Duration::from_secs(
tls.handshake_timeout_seconds
.unwrap_or(TLS_DEFAULT_HANDSHAKE_TIMEOUT_SECONDS),
);
let max_connections = tls.max_connections.unwrap_or(TLS_DEFAULT_MAX_CONNECTIONS);
Ok(HttpsTls {
acceptor,
handshake_timeout,
max_connections,
})
}
fn resolve_listener(addr_str: Option<&str>) -> ServerResult<Option<SocketAddr>> {
let Some(addr_str) = addr_str else {
return Ok(None);
};
let mut addrs = addr_str
.to_socket_addrs()
.map_err(|e| ServerError::ListenerAddress {
addr: addr_str.to_owned(),
reason: e.to_string(),
})?;
addrs
.next()
.map(Some)
.ok_or_else(|| ServerError::ListenerAddress {
addr: addr_str.to_owned(),
reason: "address resolved to no socket addresses".to_owned(),
})
}
pub async fn service(
request: hyper::Request<body::Incoming>,
app_state: Arc<AppState>,
) -> Result<hyper::Response<BoxBody>, hyper::http::Error> {
let request_headers = request.headers().clone();
let config = &app_state.config;
let middlewares = &app_state.middlewares;
let tracer = &app_state.tracer;
let cors_allow_credentials_origins = config
.service
.cors_allow_credentials_origins
.clone()
.unwrap_or_default();
if request.method() == hyper::Method::OPTIONS {
return handle_options(&request_headers, &cors_allow_credentials_origins);
}
let max_request_body_bytes = config
.service
.max_request_body_bytes
.unwrap_or(SERVICE_DEFAULT_MAX_REQUEST_BODY_BYTES);
let max_request_body_bytes = usize::try_from(max_request_body_bytes).unwrap_or(usize::MAX);
let parsed_request = match parsed_request_from(request, max_request_body_bytes).await {
Ok(x) => x,
Err(ParsedRequestError::BodyTooLarge) => {
return payload_too_large_response(
&format!(
"request body exceeds the configured limit ({} bytes)",
max_request_body_bytes
),
&request_headers,
&cors_allow_credentials_origins,
);
}
Err(ParsedRequestError::Other(err)) => {
return internal_server_error_response(
err.as_str(),
&request_headers,
&cors_allow_credentials_origins,
);
}
};
let received_at_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64;
let start = std::time::Instant::now();
capture_in_log_with_trace_config(
&parsed_request,
config.log.clone().unwrap_or_default().verbose,
&tracer.config,
);
let trace_summary = if tracer.has_subscribers() {
let headers = parsed_request
.component_parts
.headers
.iter()
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.to_string(), v.to_owned())))
.collect();
let mut summary = RequestSummary::new(
parsed_request.component_parts.method.to_string(),
parsed_request.url_path.clone(),
headers,
parsed_request.body_len,
&tracer.config,
);
tracer.enrich_with_body(&mut summary, parsed_request.body_json.as_ref());
Some(summary)
} else {
None
};
if let Some((middleware_file_path, response)) = middleware_response(
middlewares,
&parsed_request,
&cors_allow_credentials_origins,
)
.await
{
let status = response_status_or(&response, 500);
emit_trace_event(
tracer,
received_at_ms,
start,
trace_summary,
Outcome::Middleware {
file_path: middleware_file_path,
status,
},
);
return response;
}
if let Some((rule_set_idx, rule_idx, response)) = rule_set_response(
config,
&parsed_request,
app_state.canonical_rule_set_respond_dirs.as_slice(),
&cors_allow_credentials_origins,
)
.await
{
emit_trace_event(
tracer,
received_at_ms,
start,
trace_summary,
Outcome::Matched {
rule_set_index: rule_set_idx,
rule_index: rule_idx,
},
);
return response;
}
let (dyn_route_response, resolved_file_path) = dyn_route_content_traced(
parsed_request.url_path.as_str(),
config.service.fallback_respond_dir.as_str(),
&request_headers,
app_state.canonical_fallback_respond_dir.as_deref(),
&cors_allow_credentials_origins,
)
.await;
let status = response_status_or(&dyn_route_response, 500);
let outcome = match resolved_file_path {
Some(file_path) if status != 404 => Outcome::Fallback { file_path, status },
_ => Outcome::Miss { status },
};
emit_trace_event(tracer, received_at_ms, start, trace_summary, outcome);
dyn_route_response
}
fn response_status_or(
response: &Result<hyper::Response<BoxBody>, hyper::http::Error>,
fallback: u16,
) -> u16 {
response
.as_ref()
.map(|r| r.status().as_u16())
.unwrap_or(fallback)
}
fn emit_trace_event(
tracer: &TraceEmitter,
received_at_ms: u64,
start: std::time::Instant,
trace_summary: Option<RequestSummary>,
outcome: Outcome,
) {
if let Some(summary) = trace_summary {
tracer.emit(
received_at_ms,
start.elapsed().as_millis() as u32,
summary,
outcome,
);
}
}
async fn middleware_response(
middlewares: &LoadedMiddlewares,
parsed_request: &ParsedRequest,
cors_allow_credentials_origins: &[String],
) -> Option<(String, Result<hyper::Response<BoxBody>, hyper::http::Error>)> {
for handler in middlewares.iter() {
match handler
.handle(
parsed_request.url_path.as_str(),
parsed_request.body_json.as_ref(),
&parsed_request.component_parts.headers,
cors_allow_credentials_origins,
)
.await
{
Some(x) => return Some((handler.file_path.clone(), x)),
None => continue,
}
}
None
}
async fn rule_set_response(
config: &Config,
parsed_request: &ParsedRequest,
canonical_rule_set_respond_dirs: &[Option<PathBuf>],
cors_allow_credentials_origins: &[String],
) -> Option<(
usize,
usize,
Result<hyper::Response<BoxBody>, hyper::http::Error>,
)> {
for (rule_set_idx, rule_set) in config.service.rule_sets.iter().enumerate() {
if let Some((rule_idx, respond)) = rule_set.find_matched(
parsed_request,
config.service.strategy.as_ref(),
rule_set_idx,
) {
let dir_prefix = rule_set.dir_prefix();
let rule_set_default_delay_ms = rule_set
.default
.as_ref()
.and_then(|default| default.delay_response_milliseconds);
let confine_to = canonical_rule_set_respond_dirs
.get(rule_set_idx)
.and_then(|dir| dir.as_deref());
return Some((
rule_set_idx,
rule_idx,
respond_response(
&respond,
dir_prefix.as_str(),
parsed_request,
rule_set_default_delay_ms,
confine_to,
cors_allow_credentials_origins,
)
.await,
));
}
}
None
}
pub fn handle_options(
request_headers: &HeaderMap,
cors_allow_credentials_origins: &[String],
) -> Result<hyper::Response<BoxBody>, hyper::http::Error> {
let mut response = Response::new(Empty::new().boxed());
*response.status_mut() = hyper::StatusCode::NO_CONTENT;
response
.headers_mut()
.insert(CONTENT_LENGTH, HeaderValue::from_static("0"));
for (header_key, header_value) in
default_response_headers(request_headers, cors_allow_credentials_origins).into_iter()
{
if let Some(header_key) = header_key {
response.headers_mut().insert(header_key, header_value);
}
}
Ok(response)
}