use std::path::Path;
use std::sync::Arc;
use axum::Router;
use axum::http::StatusCode;
use axum::response::IntoResponse;
use axum::routing::get;
use swagger_ui_dist::{ApiDefinition, OpenApiSource};
use tokio::net::TcpListener;
use tokio::signal;
use tower_http::LatencyUnit;
use tower_http::services::ServeDir;
use tower_http::trace::{DefaultMakeSpan, DefaultOnRequest, DefaultOnResponse, TraceLayer};
use tracing::Level;
use unitycatalog_client::UnityCatalogClient;
use unitycatalog_common::services::encryption::EnvelopeEncryptor;
use unitycatalog_common::store::ObjectStoreAdapter;
use unitycatalog_common::{Error, Result};
use unitycatalog_postgres::PgCommitCoordinator;
use unitycatalog_sqlite::SqliteCommitCoordinator;
use crate::api::RequestContext;
use crate::config::PostgresBackendConfig;
use crate::config::{AuthMode, Backend, Config, SqliteBackendConfig, StorageProxyConfig, UiConfig};
use crate::policy::{ConstantPolicy, Policy};
use crate::rest::{
AnonymousAuthenticator, AuthenticationLayer, OnMissingIdentity, ReverseProxyAuthenticator,
create_agent_skills_router, create_agents_router, create_catalogs_router,
create_credentials_router, create_delta_router, create_entity_tag_assignments_router,
create_external_locations_router, create_functions_router, create_model_versions_router,
create_policies_router, create_providers_router, create_recipients_router,
create_registered_models_router, create_schemas_router, create_shares_router,
create_staging_tables_router, create_tables_router, create_tag_policies_router,
create_temporary_credentials_router, create_volumes_router,
};
use crate::services::{LocalStoragePolicy, ServerHandler, location::StorageLocationUrl};
const UI_DIR: &str = "web";
pub(crate) type LocalHandler = (
ServerHandler<RequestContext>,
Arc<dyn Policy<RequestContext>>,
);
pub async fn serve(config: Config) -> Result<()> {
let host = config.resolved_host().to_string();
let port = config.resolved_port();
let encryptor = match config.encryption.as_ref() {
Some(enc) => enc.build_encryptor().map_err(Error::Generic)?,
None if config.backend.is_ephemeral() => crate::config::EncryptionConfig::dev_default()
.build_encryptor()
.map_err(Error::Generic)?,
None => {
return Err(Error::Generic(
"missing `encryption` configuration: an active KEK is required to store secrets \
with a durable backend"
.into(),
));
}
};
let local_storage_policy = LocalStoragePolicy::new(&config.local_storage.allowed_roots)
.map_err(|e| Error::Generic(format!("invalid local_storage config: {e}")))?;
if let Some(root) = config
.managed_storage_root
.as_deref()
.filter(|s| !s.is_empty())
{
let url = StorageLocationUrl::parse(root)
.map_err(|e| Error::Generic(format!("invalid managed_storage_root '{root}': {e}")))?;
local_storage_policy
.check(&url)
.map_err(|e| Error::Generic(format!("invalid managed_storage_root '{root}': {e}")))?;
}
let migrate_on_connect = config.backend.is_ephemeral();
if migrate_on_connect {
tracing::info!("ephemeral in-memory backend: applying migrations at startup");
}
let (handler, policy) = match &config.backend {
Backend::Postgres(pg) => connect_postgres(pg, encryptor, migrate_on_connect).await?,
Backend::Sqlite(cfg) => connect_sqlite(cfg, encryptor, migrate_on_connect).await?,
};
let handler = handler
.with_local_storage_policy(local_storage_policy)
.with_managed_storage_root(config.managed_storage_root.clone());
let proxy_handler = handler.clone();
let api_router = if config.routing.any_upstream() {
let unsupported = config.routing.unsupported_upstream();
if !unsupported.is_empty() {
return Err(Error::Generic(format!(
"upstream routing is not yet implemented for: {}",
unsupported.join(", ")
)));
}
let upstream = config.upstream.as_ref().ok_or_else(|| {
Error::Generic(
"routing marks surfaces as upstream but no `upstream` config is set".to_string(),
)
})?;
let upstream_url = upstream
.url
.parse()
.map_err(|e| Error::Generic(format!("invalid upstream url: {e}")))?;
let client = UnityCatalogClient::new_unauthenticated(upstream_url);
crate::hybrid::build_hybrid_router(handler, policy, client, &config.routing)
} else {
build_rest_router(handler)
};
if let Some(err) = config.storage_proxy.startup_error() {
return Err(Error::Generic(format!("storage-proxy config: {err}")));
}
let api_router = if config.storage_proxy.enabled {
let proxy = build_storage_proxy_router(&config, proxy_handler).await?;
api_router.merge(proxy)
} else {
api_router
};
let app = match config.auth.mode {
AuthMode::Anonymous => api_router.layer(AuthenticationLayer::new(AnonymousAuthenticator)),
AuthMode::ReverseProxy => {
let mut authenticator = ReverseProxyAuthenticator::new();
if let Some(header) = &config.auth.forwarded_user_header {
authenticator = authenticator.with_header(header.clone());
}
if config.auth.allow_missing_identity {
authenticator = authenticator.with_on_missing(OnMissingIdentity::Anonymous);
}
api_router.layer(AuthenticationLayer::new(authenticator))
}
};
run(app, &config.ui, &config.storage_proxy, &host, port).await
}
async fn build_storage_proxy_router(
config: &Config,
handler: ServerHandler<RequestContext>,
) -> Result<Router> {
use crate::config::ProxyArm;
match config.storage_proxy.arm {
ProxyArm::Local => Ok(crate::rest::create_storage_proxy_router(handler)),
ProxyArm::Client => {
let client = config.storage_proxy.client.as_ref().ok_or_else(|| {
Error::Generic("storage-proxy client arm requires a `client` block".into())
})?;
let token = client.token.as_ref().and_then(|t| t.value());
let backend = unitycatalog_storage_proxy::UnityFactoryProxyBackend::connect(
&client.base_url,
token,
)
.await
.map_err(|e| Error::Generic(format!("storage-proxy client connect: {e}")))?;
use crate::policy::Principal;
use std::sync::Arc;
use unitycatalog_storage_proxy::{ContextExtractor, router_with_context};
let extract_cx: ContextExtractor<RequestContext> = Arc::new(|parts| {
let recipient = parts
.extensions
.get::<Principal>()
.cloned()
.unwrap_or_else(Principal::anonymous);
Box::pin(async move { Ok(RequestContext { recipient }) })
});
Ok(router_with_context::<(), _>(Arc::new(backend), extract_cx).with_state(()))
}
}
}
pub async fn migrate(config: Config) -> Result<()> {
match &config.backend {
Backend::Postgres(pg) => {
let db_url = pg.connection_string().ok_or_else(|| {
Error::Generic("incomplete postgres backend configuration".into())
})?;
let pool = unitycatalog_postgres::connect_pool(&db_url)
.await
.map_err(|e| Error::Generic(format!("connecting to database: {e}")))?;
unitycatalog_postgres::unified_migrator()
.run(&pool)
.await
.map_err(|e| Error::Generic(format!("running migrations: {e}")))?;
}
Backend::Sqlite(cfg) => {
let path = cfg
.database_path()
.ok_or_else(|| Error::Generic("incomplete sqlite backend configuration".into()))?;
let pool = unitycatalog_sqlite::connect_pool(&path)
.await
.map_err(|e| Error::Generic(format!("opening sqlite database: {e}")))?;
unitycatalog_sqlite::unified_migrator()
.run(&pool)
.await
.map_err(|e| Error::Generic(format!("running migrations: {e}")))?;
}
}
tracing::info!("migrations applied");
Ok(())
}
pub(crate) fn swagger_api_defs() -> Vec<ApiDefinition<&'static str>> {
vec![
ApiDefinition {
uri_prefix: "/api/2.1/unity-catalog",
api_definition: OpenApiSource::Inline(include_str!("../openapi/openapi.yaml")),
title: Some("Unity Catalog API"),
},
ApiDefinition {
uri_prefix: "/api/2.1/unity-catalog/delta",
api_definition: OpenApiSource::Inline(include_str!("../openapi/delta.yaml")),
title: Some("UC Delta API"),
},
]
}
pub(crate) fn build_rest_router(handler: ServerHandler<RequestContext>) -> Router {
let api_routes = create_catalogs_router(handler.clone())
.merge(create_schemas_router(handler.clone()))
.merge(create_staging_tables_router(handler.clone()))
.merge(create_tables_router(handler.clone()))
.merge(create_volumes_router(handler.clone()))
.merge(create_agent_skills_router(handler.clone()))
.merge(create_agents_router(handler.clone()))
.merge(create_credentials_router(handler.clone()))
.merge(create_external_locations_router(handler.clone()))
.merge(create_temporary_credentials_router(handler.clone()))
.merge(create_functions_router(handler.clone()))
.merge(create_registered_models_router(handler.clone()))
.merge(create_model_versions_router(handler.clone()))
.merge(create_recipients_router(handler.clone()))
.merge(create_providers_router(handler.clone()))
.merge(create_shares_router(handler.clone()))
.merge(create_delta_router(handler.clone()))
.merge(create_entity_tag_assignments_router(handler.clone()))
.merge(create_policies_router(handler.clone()));
Router::new()
.nest("/api/2.1/unity-catalog", api_routes)
.nest("/api/2.1", create_tag_policies_router(handler))
}
fn operational_router(storage_proxy: &StorageProxyConfig) -> Router {
let capabilities = capabilities_body(storage_proxy);
Router::new()
.route("/health", get(|| async { "OK" }))
.route("/version", get(|| async { env!("CARGO_PKG_VERSION") }))
.route(
"/capabilities",
get(move || async move {
(
[(axum::http::header::CONTENT_TYPE, "application/json")],
capabilities,
)
}),
)
}
fn capabilities_body(storage_proxy: &StorageProxyConfig) -> String {
if storage_proxy.enabled {
r#"{"storageAccess":"proxy","storageProxy":{"basePath":"/storage-proxy","conditionalWrites":true}}"#
.to_string()
} else {
r#"{"storageAccess":"direct"}"#.to_string()
}
}
pub(crate) async fn run(
api_router: Router,
ui: &UiConfig,
storage_proxy: &StorageProxyConfig,
host: &str,
port: u16,
) -> Result<()> {
let mut router = swagger_api_defs()
.into_iter()
.fold(api_router, |router, api| {
router.merge(swagger_ui_dist::generate_routes(api))
});
if ui.serve {
router = mount_spa(router);
}
let router = mount_under_base(router, &ui.normalized_base_path());
let router = operational_router(storage_proxy).merge(router);
let router = router.layer(
TraceLayer::new_for_http()
.make_span_with(DefaultMakeSpan::new().include_headers(true))
.on_request(DefaultOnRequest::new().level(Level::INFO))
.on_response(
DefaultOnResponse::new()
.level(Level::INFO)
.latency_unit(LatencyUnit::Micros),
),
);
let listener = TcpListener::bind(format!("{host}:{port}"))
.await
.map_err(|e| Error::Generic(e.to_string()))?;
let addr = listener
.local_addr()
.map_err(|e| Error::Generic(e.to_string()))?;
tracing::info!("Listening on: {addr}");
axum::serve(listener, router)
.with_graceful_shutdown(shutdown_signal())
.await
.map_err(|e| Error::Generic(e.to_string()))?;
Ok(())
}
async fn connect_postgres(
pg: &PostgresBackendConfig,
encryptor: EnvelopeEncryptor,
migrate: bool,
) -> Result<LocalHandler> {
let db_url = pg
.connection_string()
.ok_or_else(|| Error::Generic("incomplete postgres backend configuration".into()))?;
let pool = unitycatalog_postgres::connect_pool(&db_url)
.await
.map_err(|e| Error::Generic(format!("connecting to database: {e}")))?;
if migrate {
unitycatalog_postgres::unified_migrator()
.run(&pool)
.await
.map_err(|e| Error::Generic(format!("running migrations: {e}")))?;
}
let policy: Arc<dyn Policy<RequestContext>> = Arc::new(ConstantPolicy::default());
let graph = unitycatalog_postgres::connect_graph(pool.clone(), encryptor);
let resource_store = Arc::new(ObjectStoreAdapter::new(graph));
let coordinator = Arc::new(PgCommitCoordinator::new(pool));
let handler =
ServerHandler::try_new_tokio_with_coordinator(policy.clone(), resource_store, coordinator)
.map_err(|e| Error::Generic(e.to_string()))?;
Ok((handler, policy))
}
async fn connect_sqlite(
cfg: &SqliteBackendConfig,
encryptor: EnvelopeEncryptor,
migrate: bool,
) -> Result<LocalHandler> {
let path = cfg
.database_path()
.ok_or_else(|| Error::Generic("incomplete sqlite backend configuration".into()))?;
let pool = unitycatalog_sqlite::connect_pool(&path)
.await
.map_err(|e| Error::Generic(format!("opening sqlite database: {e}")))?;
if migrate {
unitycatalog_sqlite::unified_migrator()
.run(&pool)
.await
.map_err(|e| Error::Generic(format!("running migrations: {e}")))?;
}
let policy: Arc<dyn Policy<RequestContext>> = Arc::new(ConstantPolicy::default());
let graph = unitycatalog_sqlite::connect_graph(pool.clone(), encryptor);
let resource_store = Arc::new(ObjectStoreAdapter::new(graph));
let coordinator = Arc::new(SqliteCommitCoordinator::new(pool));
let handler =
ServerHandler::try_new_tokio_with_coordinator(policy.clone(), resource_store, coordinator)
.map_err(|e| Error::Generic(e.to_string()))?;
Ok((handler, policy))
}
fn mount_spa(app: Router) -> Router {
let index_path = Path::new(UI_DIR).join("index.html");
let index_html: Arc<Option<String>> = Arc::new(std::fs::read_to_string(&index_path).ok());
if index_html.is_none() {
tracing::warn!(
"ui.serve is enabled but no bundle found at `{}`; SPA routes will 404",
index_path.display()
);
}
let index_handler = move || {
let index_html = index_html.clone();
get(move || {
let index_html = index_html.clone();
async move {
match index_html.as_ref() {
Some(html) => axum::response::Html(html.clone()).into_response(),
None => StatusCode::NOT_FOUND.into_response(),
}
}
})
};
let serve_assets = ServeDir::new(UI_DIR)
.append_index_html_on_directories(false)
.fallback(index_handler());
app
.route("/", index_handler())
.route("/index.html", index_handler())
.fallback_service(serve_assets)
}
fn mount_under_base(app: Router, base_path: &str) -> Router {
if base_path.is_empty() {
return app;
}
let prefix = base_path.to_string();
let stripped = tower::ServiceBuilder::new()
.layer(axum::middleware::from_fn(
move |mut req: axum::extract::Request, next: axum::middleware::Next| {
let prefix = prefix.clone();
async move {
let path = req.uri().path();
let new_path = match path.strip_prefix(&prefix) {
Some("") => Some("/".to_string()),
Some(rest) if rest.starts_with('/') => Some(rest.to_string()),
_ => None,
};
match new_path {
Some(new_path) => {
rewrite_path(&mut req, &new_path);
next.run(req).await
}
None => StatusCode::NOT_FOUND.into_response(),
}
}
},
))
.service(app);
Router::new().fallback_service(stripped)
}
fn rewrite_path(req: &mut axum::extract::Request, new_path: &str) {
let uri = req.uri();
let path_and_query = match uri.query() {
Some(q) => format!("{new_path}?{q}"),
None => new_path.to_string(),
};
let mut parts = uri.clone().into_parts();
parts.path_and_query = Some(
path_and_query
.parse()
.expect("rewritten path-and-query is valid"),
);
if let Ok(new_uri) = axum::http::Uri::from_parts(parts) {
*req.uri_mut() = new_uri;
}
}
async fn shutdown_signal() {
let ctrl_c = async {
signal::ctrl_c()
.await
.expect("failed to install Ctrl+C handler");
};
#[cfg(unix)]
let terminate = async {
signal::unix::signal(signal::unix::SignalKind::terminate())
.expect("failed to install signal handler")
.recv()
.await;
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {},
_ = terminate => {},
}
}