pub mod api;
mod auth;
pub mod client_ip;
pub mod error;
pub mod export;
#[cfg(feature = "graphql")]
pub mod gql;
pub(crate) mod headers;
pub mod health;
pub mod import;
mod input;
pub mod key;
pub mod ml;
pub(crate) mod output;
mod params;
pub mod rpc;
mod signals;
pub mod signin;
pub mod signup;
pub mod sql;
pub mod sync;
mod tracer;
pub mod version;
use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use anyhow::Result;
use axum::response::Redirect;
use axum::routing::get;
use axum::{Router, middleware};
use axum_server::Handle;
use axum_server::tls_rustls::RustlsConfig;
use http::header;
use surrealdb::headers::{AUTH_DB, AUTH_NS, DB, ID, NS};
use surrealdb_core::CommunityComposer;
use surrealdb_core::kvs::Datastore;
use tokio_util::sync::CancellationToken;
use tower::ServiceBuilder;
use tower_http::ServiceBuilderExt;
use tower_http::add_extension::AddExtensionLayer;
use tower_http::auth::AsyncRequireAuthorizationLayer;
use tower_http::compression::CompressionLayer;
use tower_http::compression::predicate::{NotForContentType, Predicate, SizeAbove};
use tower_http::cors::{AllowOrigin, Any, CorsLayer};
use tower_http::request_id::MakeRequestUuid;
use tower_http::sensitive_headers::{
SetSensitiveRequestHeadersLayer, SetSensitiveResponseHeadersLayer,
};
use tower_http::trace::TraceLayer;
use crate::cli::Config;
use crate::cnf;
use crate::ntw::signals::graceful_shutdown;
use crate::rpc::{RpcState, notifications};
use crate::telemetry::metrics::HttpMetricsLayer;
const LOG: &str = "surrealdb::net";
pub trait RouterFactory {
fn configure_router() -> Router<Arc<RpcState>>;
}
impl RouterFactory for CommunityComposer {
fn configure_router() -> Router<Arc<RpcState>> {
let router = Router::<Arc<RpcState>>::new()
.route("/", get(|| async { Redirect::temporary(cnf::APP_ENDPOINT) }))
.route("/status", get(|| async {}))
.merge(health::router())
.merge(export::router())
.merge(import::router())
.merge(rpc::router())
.merge(version::router())
.merge(sync::router())
.merge(sql::router())
.merge(signin::router())
.merge(signup::router())
.merge(key::router())
.merge(ml::router())
.merge(api::router());
#[cfg(feature = "graphql")]
let router = router.merge(gql::router());
router
}
}
#[derive(Clone)]
pub struct AppState {
pub client_ip: client_ip::ClientIp,
pub datastore: Arc<Datastore>,
}
pub async fn init<F: RouterFactory>(
opt: &Config,
ds: Arc<Datastore>,
ct: CancellationToken,
) -> Result<()> {
let app_state = AppState {
client_ip: opt.client_ip,
datastore: ds.clone(),
};
let headers: Arc<[_]> = Arc::new([
header::AUTHORIZATION,
header::PROXY_AUTHORIZATION,
header::COOKIE,
header::SET_COOKIE,
]);
let service = ServiceBuilder::new()
.catch_panic()
.set_x_request_id(MakeRequestUuid)
.propagate_x_request_id()
.concurrency_limit(*cnf::NET_MAX_CONCURRENT_REQUESTS);
let service = service.layer(
CompressionLayer::new().compress_when(
SizeAbove::new(512)
.and(NotForContentType::GRPC)
.and(NotForContentType::IMAGES),
),
);
let allow_origin: AllowOrigin = if opt.allow_origin.is_empty() {
Any.into()
} else {
let origins: Vec<http::HeaderValue> = opt
.allow_origin
.iter()
.map(|o| o.parse().expect("CORS origins are validated at startup"))
.collect();
AllowOrigin::list(origins)
};
let allow_header = [
http::header::ACCEPT,
http::header::ACCEPT_ENCODING,
http::header::AUTHORIZATION,
http::header::CONTENT_TYPE,
http::header::ORIGIN,
NS.clone(),
DB.clone(),
ID.clone(),
AUTH_NS.clone(),
AUTH_DB.clone(),
];
let service = service
.layer(AddExtensionLayer::new(app_state))
.layer(middleware::from_fn(client_ip::client_ip_middleware))
.layer(SetSensitiveRequestHeadersLayer::from_shared(Arc::clone(&headers)))
.layer(
TraceLayer::new_for_http()
.make_span_with(tracer::HttpTraceLayerHooks)
.on_request(tracer::HttpTraceLayerHooks)
.on_response(tracer::HttpTraceLayerHooks)
.on_failure(tracer::HttpTraceLayerHooks),
)
.layer(HttpMetricsLayer)
.layer(SetSensitiveResponseHeadersLayer::from_shared(headers))
.layer(AsyncRequireAuthorizationLayer::new(auth::SurrealAuth))
.layer(headers::add_server_header(!opt.no_identification_headers)?)
.layer(headers::add_version_header(!opt.no_identification_headers)?)
.layer(
CorsLayer::new()
.allow_methods([
http::Method::GET,
http::Method::PUT,
http::Method::POST,
http::Method::PATCH,
http::Method::DELETE,
http::Method::OPTIONS,
])
.allow_headers(allow_header)
.allow_origin(allow_origin)
.max_age(Duration::from_secs(86400)),
);
let axum_app = F::configure_router();
let axum_app = axum_app.layer(service);
let handle = Handle::new();
let rpc_state = Arc::new(RpcState::new(ds.clone(), surrealdb_core::dbs::Session::default()));
let shutdown_handler = graceful_shutdown(rpc_state.clone(), ct.clone(), handle.clone());
let axum_app = axum_app.with_state(rpc_state.clone());
tokio::spawn(async move { notifications(ds, rpc_state, ct.clone()).await });
let res = if let (Some(cert), Some(key)) = (&opt.crt, &opt.key) {
let tls = RustlsConfig::from_pem_file(cert, key).await?;
let server = axum_server::bind_rustls(opt.bind, tls);
info!(target: LOG, "Started web server on {}", &opt.bind);
server
.handle(handle)
.serve(axum_app.into_make_service_with_connect_info::<SocketAddr>())
.await
} else {
let server = axum_server::bind(opt.bind);
info!(target: LOG, "Started web server on {}", &opt.bind);
server
.handle(handle)
.serve(axum_app.into_make_service_with_connect_info::<SocketAddr>())
.await
};
if let Err(e) = res {
if opt.bind.port() < 1024
&& let io::ErrorKind::PermissionDenied = e.kind()
{
error!(target: LOG, "Binding to ports below 1024 requires privileged access or special permissions.");
}
return Err(e.into());
}
let _ = shutdown_handler.await;
info!(target: LOG, "Web server stopped. Bye!");
Ok(())
}