#[allow(dead_code)]
mod models;
#[allow(dead_code)]
mod pool;
#[allow(dead_code)]
mod queue;
#[allow(dead_code)]
mod workers;
mod handlers;
mod auth;
mod database;
#[allow(dead_code)]
mod error;
mod introspection;
mod setup;
mod config;
mod rate_limit;
mod connection_limit;
#[allow(dead_code)]
mod observability;
use axum::{
error_handling::HandleErrorLayer,
extract::DefaultBodyLimit,
http::StatusCode,
middleware,
routing::{delete, get, post},
Router,
};
use std::net::SocketAddr;
use std::env;
use std::time::Duration;
use tower::timeout::TimeoutLayer;
use tower::ServiceBuilder;
use tower_http::cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer};
use tower_http::set_header::SetResponseHeaderLayer;
use tracing::info;
use crate::models::AppState;
use crate::config::{load_server_config, CorsConfig};
fn cors_layer(cfg: &CorsConfig) -> CorsLayer {
if !cfg.enabled {
return CorsLayer::new();
}
let mut layer = CorsLayer::new()
.allow_methods(AllowMethods::list(vec![
axum::http::Method::GET,
axum::http::Method::POST,
axum::http::Method::DELETE,
axum::http::Method::OPTIONS,
]))
.allow_headers(AllowHeaders::list(vec![
axum::http::header::CONTENT_TYPE,
axum::http::header::AUTHORIZATION,
"x-api-key".parse().expect("static header name"),
"x-request-id".parse().expect("static header name"),
]));
if cfg.origins.iter().any(|o| o == "*") {
layer = layer.allow_origin(AllowOrigin::any());
} else {
let origins: Vec<axum::http::HeaderValue> = cfg
.origins
.iter()
.filter_map(|o| o.parse().ok())
.collect();
layer = layer.allow_origin(AllowOrigin::list(origins));
}
layer
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
let server_config = load_server_config().await?;
tracing_subscriber::fmt()
.with_env_filter(format!("pg_api={}", server_config.log_level))
.json()
.init();
let args: Vec<String> = env::args().collect();
if args.len() > 1 && args[1] == "setup" {
return setup::run_setup().await;
}
let state = match AppState::new().await {
Ok(state) => state,
Err(e) => {
eprintln!("pg-api: refusing to start: {e:#}");
std::process::exit(1);
}
};
let observability_config = observability::ObservabilityConfig::from_env();
let observability_client = std::sync::Arc::new(
observability::ObservabilityClient::new(observability_config.clone())
);
if observability_config.enabled {
observability_client.clone().start_flush_task();
info!("Observability enabled, sending metrics to: {:?}",
observability_config.opensearch_url);
} else {
info!("Observability disabled");
}
let app = Router::new()
.route("/", get(handlers::serve_docs))
.route("/docs", get(handlers::serve_docs))
.route("/openapi.json", get(handlers::serve_openapi))
.route("/health", get(handlers::health_check))
.route("/v1/status", get(handlers::status_handler))
.route("/v1/query", post(handlers::query_handler))
.route("/v1/batch", post(handlers::batch_query_handler))
.route("/v1/transaction", post(handlers::transaction_handler))
.route("/v1/databases", get(handlers::list_databases))
.route("/v1/databases", post(handlers::create_database))
.route("/v1/databases/{name}", delete(handlers::drop_database))
.route("/v1/databases/{db}/tables", get(handlers::list_tables))
.route("/v1/databases/{db}/schema", get(handlers::get_schema))
.route("/v1/account", get(handlers::get_account_info))
.route("/v1/account/usage", get(handlers::get_usage_stats))
.layer(middleware::from_fn_with_state(observability_client.clone(), observability::metrics_middleware))
.layer(middleware::from_fn_with_state(state.clone(), connection_limit::connection_limit_middleware))
.layer(middleware::from_fn_with_state(state.clone(), rate_limit::rate_limit_middleware))
.layer(middleware::from_fn_with_state(state.clone(), auth::auth_middleware))
.layer(middleware::from_fn(auth::request_id_middleware))
.layer(cors_layer(&server_config.cors))
.layer(DefaultBodyLimit::max(
server_config.limits.max_request_size_mb * 1024 * 1024,
))
.layer(
ServiceBuilder::new()
.layer(HandleErrorLayer::new(
|_: Box<dyn std::error::Error + Send + Sync>| async {
StatusCode::REQUEST_TIMEOUT
},
))
.layer(TimeoutLayer::new(Duration::from_secs(
server_config.limits.request_timeout_seconds,
))),
)
.layer(SetResponseHeaderLayer::if_not_present(
axum::http::header::X_CONTENT_TYPE_OPTIONS,
axum::http::HeaderValue::from_static("nosniff"),
))
.with_state(state);
let addr = SocketAddr::new(server_config.host, server_config.port);
let listener = tokio::net::TcpListener::bind(addr).await?;
info!("pg-api running on {}", addr);
info!("Documentation available at http://{}/docs", addr);
axum::serve(listener, app).await?;
Ok(())
}