use std::ops::Deref;
use anyhow::Result;
use axum::body::Body;
use axum::response::{IntoResponse, Response};
use axum::routing::options;
use axum::{Extension, Router};
use axum_extra::TypedHeader;
use bytes::Bytes;
use http::StatusCode;
use surrealdb_core::dbs::Session;
use surrealdb_core::dbs::capabilities::RouteTarget;
use surrealdb_core::iam::Action::View;
use surrealdb_core::iam::ResourceKind::Any;
use surrealdb_core::iam::check::check_ns_db;
use surrealdb_core::kvs::export;
use surrealdb_core::rpc::format::Format;
use surrealdb_types::SurrealValue;
use super::AppState;
use super::error::ResponseError;
use super::headers::ContentType;
use crate::ntw::error::Error as NetError;
pub fn router<S>() -> Router<S>
where
S: Clone + Send + Sync + 'static,
{
Router::new().route("/export", options(|| async {}).get(get_handler).post(post_handler))
}
async fn get_handler(
Extension(state): Extension<AppState>,
Extension(session): Extension<Session>,
) -> Result<impl IntoResponse, ResponseError> {
let cfg = export::Config::default();
handle_inner(state, session, cfg).await
}
async fn post_handler(
Extension(state): Extension<AppState>,
Extension(session): Extension<Session>,
content_type: TypedHeader<ContentType>,
body: Bytes,
) -> Result<impl IntoResponse, ResponseError> {
let fmt = content_type.deref();
let fmt: Format = fmt.into();
let val = match fmt {
Format::Json => surrealdb_core::rpc::format::json::decode(&body)
.map_err(anyhow::Error::msg)
.map_err(ResponseError)?,
Format::Cbor => surrealdb_core::rpc::format::cbor::decode(&body)
.map_err(anyhow::Error::msg)
.map_err(ResponseError)?,
Format::Flatbuffers => surrealdb_core::rpc::format::flatbuffers::decode(&body)
.map_err(anyhow::Error::msg)
.map_err(ResponseError)?,
Format::Unsupported => {
return Err(ResponseError(anyhow::Error::msg("unsupported body format")));
}
};
let cfg =
export::Config::from_value(val).map_err(|e| ResponseError(anyhow::anyhow!("{}", e)))?;
handle_inner(state, session, cfg).await
}
async fn handle_inner(
state: AppState,
session: Session,
cfg: export::Config,
) -> Result<impl IntoResponse, ResponseError> {
let db = &state.datastore;
if !db.allows_http_route(&RouteTarget::Export) {
warn!("Capabilities denied HTTP route request attempt, target: '{}'", &RouteTarget::Export);
return Err(NetError::ForbiddenRoute(RouteTarget::Export.to_string()).into());
}
let (chn, body_stream) = surrealdb::channel::bounded::<Result<Bytes>>(1);
let body = Body::from_stream(body_stream);
let (nsv, dbv) = check_ns_db(&session).map_err(ResponseError)?;
db.check(&session, View, Any.on_db(&nsv, &dbv)).map_err(ResponseError)?;
let (snd, rcv) = surrealdb::channel::bounded(1);
let task = db.export_with_config(&session, snd, cfg).await.map_err(ResponseError)?;
tokio::spawn(task);
tokio::spawn(async move {
while let Ok(v) = rcv.recv().await {
if let Err(err) = chn.send(Ok(Bytes::from(v))).await {
tracing::warn!("Error sending bytes: {:?}", err);
}
}
});
Ok(Response::builder().status(StatusCode::OK).body(body)?)
}