use std::sync::Arc;
use std::time::Duration;
use anyhow::{Context, Result};
use clap::Parser;
use tracing_subscriber::EnvFilter;
use git_cache_proxy::config::{Config, LogFormat};
use git_cache_proxy::{git, metrics, server};
#[tokio::main(flavor = "multi_thread")]
async fn main() -> Result<()> {
let cfg = Config::parse();
let _log_guard = init_tracing(&cfg.log, cfg.log_format);
tokio::fs::create_dir_all(&cfg.cache_root)
.await
.with_context(|| format!("create cache root {}", cfg.cache_root.display()))?;
let upstream_auth_header = cfg
.upstream_auth_header
.as_deref()
.map(str::trim)
.filter(|h| !h.is_empty())
.map(str::to_string);
if upstream_auth_header.is_none() {
tracing::warn!(
"no upstream auth header set; contacting upstream anonymously - \
private repos will fail with an upstream 401 \
(set --upstream-auth-header / GITCACHEPROXY_UPSTREAM_AUTH_HEADER)"
);
}
let git_cfg = git::GitConfig {
git_binary: cfg.git_binary.clone(),
upstream_auth_header,
fetch_ttl: Duration::from_secs(cfg.fetch_ttl_seconds),
};
let metrics = Arc::new(metrics::Metrics::new());
let state = server::AppState {
cache: Arc::new(git::GitCache::new(git_cfg, metrics.clone())),
upstream_base: cfg.upstream.trim_end_matches('/').to_string(),
cache_root: cfg.cache_root.clone(),
serve_token: cfg.serve_token.clone(),
max_decoded_body: (cfg.max_decoded_body_mb as usize).saturating_mul(1024 * 1024),
max_concurrent: cfg.max_concurrent_requests,
metrics,
};
let listener = tokio::net::TcpListener::bind(&cfg.bind)
.await
.with_context(|| format!("bind {}", cfg.bind))?;
tracing::info!(
"git-cache-proxy listening on {} (upstream {}, cache {})",
cfg.bind,
cfg.upstream,
cfg.cache_root.display(),
);
axum::serve(listener, server::router(state))
.with_graceful_shutdown(shutdown_signal())
.await
.context("http server")?;
tracing::info!("shutdown complete");
Ok(())
}
fn init_tracing(filter: &str, format: LogFormat) -> tracing_appender::non_blocking::WorkerGuard {
let env_filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new(filter));
let (writer, guard) = tracing_appender::non_blocking(std::io::stdout());
let base = tracing_subscriber::fmt()
.with_env_filter(env_filter)
.with_target(false)
.with_writer(writer);
match format {
LogFormat::Json => base.json().flatten_event(true).init(),
LogFormat::Text => base.init(),
}
guard
}
async fn shutdown_signal() {
let ctrl_c = async {
let _ = tokio::signal::ctrl_c().await;
};
#[cfg(unix)]
let term = async {
use tokio::signal::unix::{SignalKind, signal};
match signal(SignalKind::terminate()) {
Ok(mut s) => {
s.recv().await;
}
Err(_) => std::future::pending().await,
}
};
#[cfg(not(unix))]
let term = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {}
_ = term => {}
}
tracing::info!("shutdown signal received; draining");
}