pg-api 0.3.4

A high-performance PostgreSQL REST API driver with rate limiting, connection pooling, and observability
mod models;
mod pool;
mod queue;
mod workers;
mod handlers;
mod auth;
mod database;
mod error;
mod introspection;
mod setup;
mod config;
mod rate_limit;
mod connection_limit;
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};

/// Build the CORS layer from operator configuration instead of a blanket
/// permissive policy (v0.3). `origins = ["*"]` keeps the old behavior
/// explicitly; any other list is enforced. Disabled CORS emits no
/// `Access-Control-*` headers at all.
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<()> {
    // Load server configuration
    let server_config = load_server_config().await?;
    
    // Initialize tracing
    tracing_subscriber::fmt()
        .with_env_filter(format!("pg_api={}", server_config.log_level))
        .json()
        .init();

    // Check if running in setup mode
    let args: Vec<String> = env::args().collect();
    if args.len() > 1 && args[1] == "setup" {
        return setup::run_setup().await;
    }

    // Initialize app state (fail-closed: missing/invalid account
    // configuration aborts startup with a non-zero exit, never with
    // default credentials).
    let state = match AppState::new().await {
        Ok(state) => state,
        Err(e) => {
            eprintln!("pg-api: refusing to start: {e:#}");
            std::process::exit(1);
        }
    };
    
    // Initialize observability
    let observability_config = observability::ObservabilityConfig::from_env();
    let observability_client = std::sync::Arc::new(
        observability::ObservabilityClient::new(observability_config.clone())
    );
    
    // Start observability flush task if enabled
    if observability_config.enabled {
        observability_client.clone().start_flush_task();
        info!("Observability enabled, sending metrics to: {:?}", 
              observability_config.opensearch_url);
    } else {
        info!("Observability disabled");
    }

    // Build router
    let app = Router::new()
        // Documentation
        .route("/", get(handlers::serve_docs))
        .route("/docs", get(handlers::serve_docs))
        .route("/openapi.json", get(handlers::serve_openapi))
        
        // Public endpoints
        .route("/health", get(handlers::health_check))
        .route("/v1/status", get(handlers::status_handler))
        
        // Query endpoints
        .route("/v1/query", post(handlers::query_handler))
        .route("/v1/batch", post(handlers::batch_query_handler))
        .route("/v1/transaction", post(handlers::transaction_handler))
        
        // Database management
        .route("/v1/databases", get(handlers::list_databases))
        .route("/v1/databases", post(handlers::create_database))
        .route("/v1/databases/{name}", delete(handlers::drop_database))
        
        // Schema operations
        .route("/v1/databases/{db}/tables", get(handlers::list_tables))
        .route("/v1/databases/{db}/schema", get(handlers::get_schema))
        
        // Account management
        .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))
        // v0.3: operator-configured CORS + enforced body/timeout budgets
        // (these LimitsConfig values were previously parsed but never applied).
        .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);

    // Start server
    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(())
}