pub mod artifacts;
pub mod clock;
pub mod concurrency;
pub mod config;
pub(crate) mod delivery;
pub mod http;
pub mod npm;
pub mod osv;
pub mod policy;
pub mod pypi;
pub mod store;
pub mod tasks;
pub mod upstream;
use std::fmt;
use std::net::SocketAddr;
use std::sync::Arc;
use arc_swap::ArcSwapOption;
use tokio::net::TcpListener;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use crate::artifacts::content::ContentStore;
use crate::artifacts::download::DownloadCoordinator;
use crate::clock::Clock;
use crate::config::Config;
use crate::http::limits::Limits;
use crate::policy::BlocklistSnapshot;
use crate::store::StoreHandle;
use crate::store::cache::MemoryCaches;
use crate::tasks::Tasks;
use crate::tasks::blocklist_poller::{self, Watcher};
use crate::upstream::{OriginSet, Transport};
pub struct AppDeps {
pub config: Config,
pub clock: Arc<dyn Clock>,
pub transport: Arc<dyn Transport>,
pub origins: OriginSet,
pub osv_client: reqwest::Client,
pub osv_base_url: Option<url::Url>,
}
pub struct App {
pub config: Config,
pub clock: Arc<dyn Clock>,
pub transport: Arc<dyn Transport>,
pub origins: OriginSet,
policy: ArcSwapOption<BlocklistSnapshot>,
store: StoreHandle,
pub content: ContentStore,
pub downloads: DownloadCoordinator,
pub limits: Limits,
pub(crate) delivery: delivery::Sinks,
osv: osv::OsvClient,
}
impl App {
pub fn blocklist(&self) -> Option<Arc<BlocklistSnapshot>> {
self.policy.load_full()
}
pub fn publish_blocklist(&self, snapshot: Arc<BlocklistSnapshot>) {
self.policy.store(Some(snapshot));
}
pub fn blocklist_revision(&self) -> Option<u64> {
self.policy
.load()
.as_ref()
.map(|snapshot| snapshot.revision)
}
pub fn store(&self) -> &StoreHandle {
&self.store
}
#[cfg(feature = "test-support")]
pub fn delivery_is_empty(&self) -> bool {
self.delivery.is_empty()
}
#[cfg(feature = "test-support")]
pub fn delivery_lost_total(&self) -> u64 {
self.delivery.lost_total()
}
pub async fn start(deps: AppDeps) -> Result<Running, StartupError> {
let opened = store::startup::open_and_recover(&deps.config.data_dir)
.await
.map_err(StartupError::DataDir)?;
let caches = Arc::new(MemoryCaches::new(deps.config.memory_cache_max_bytes.get()));
let (store, store_task) = match opened.connection {
Ok(connection) => {
let (store, task) = store::spawn(connection, opened.lock, Arc::clone(&caches));
(store, Some(task))
}
Err(err) => (
StoreHandle::unusable(err.to_string(), opened.lock, caches),
None,
),
};
let content = ContentStore::new(&deps.config.data_dir, deps.config.cache_max_bytes.get());
content.remove_temp_files();
let limits = Limits::new(&deps.config);
let drain = CancellationToken::new();
let (delivery, mut delivery_tasks) = delivery::build(&deps.config, drain.clone())?;
let osv_cache_ttl = std::time::Duration::from_secs(deps.config.osv_cache_ttl_seconds.get());
let osv_request_timeout =
std::time::Duration::from_millis(deps.config.osv_request_timeout_ms.get());
let (osv, osv_task) = match deps.osv_base_url {
Some(url) => osv::OsvClient::spawn_with(
deps.osv_client,
url,
osv_cache_ttl,
osv_request_timeout,
deps.config.osv_mode,
drain.clone(),
),
None => osv::OsvClient::new(
deps.osv_client,
osv_cache_ttl,
osv_request_timeout,
deps.config.osv_mode,
drain.clone(),
),
};
delivery_tasks.push(osv_task);
let app = Arc::new(App {
config: deps.config,
clock: deps.clock,
transport: deps.transport,
origins: deps.origins,
policy: ArcSwapOption::empty(),
store,
content,
downloads: DownloadCoordinator::new(),
limits,
delivery,
osv,
});
restore_blocklist(&app).await;
let mut watcher = Watcher::new();
blocklist_poller::poll_once(&app, &mut watcher).await;
let addr = app.config.listen;
let listener = TcpListener::bind(addr)
.await
.map_err(|source| StartupError::Bind { addr, source })?;
let local_addr = listener.local_addr().map_err(StartupError::Serve)?;
let shutdown = CancellationToken::new();
let tasks = tasks::spawn(Arc::clone(&app), shutdown.clone(), watcher, delivery_tasks);
let signal = shutdown.clone();
let server = tokio::spawn({
let app = Arc::clone(&app);
async move {
let consumer_identification = app.config.log_consumer_identification;
let router = http::router(app);
if consumer_identification {
axum::serve(
listener,
router.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(async move { signal.cancelled().await })
.await
} else {
axum::serve(listener, router)
.with_graceful_shutdown(async move { signal.cancelled().await })
.await
}
}
});
Ok(Running {
local_addr,
app,
shutdown,
drain,
server,
tasks,
store_task,
})
}
}
async fn restore_blocklist(app: &App) {
let row = match app.store().load_blocklist().await {
Ok(Some(row)) => row,
Ok(None) => return,
Err(err) => {
tracing::error!(
error = %err,
"the persisted blocklist could not be read; readiness stays false until a \
current blocklist is loaded"
);
return;
}
};
let now = app.clock.now_utc_micros();
match BlocklistSnapshot::parse_and_validate(&row.snapshot, now) {
Ok(snapshot) => {
tracing::info!(
revision = snapshot.revision,
entries = snapshot.entry_count(),
"the persisted blocklist is still valid and is back in force"
);
app.publish_blocklist(Arc::new(snapshot));
}
Err(err) => tracing::warn!(
revision = row.revision,
error = %err,
"the persisted blocklist is no longer usable; readiness stays false until a \
current blocklist is loaded"
),
}
}
pub struct Running {
pub local_addr: SocketAddr,
app: Arc<App>,
shutdown: CancellationToken,
drain: CancellationToken,
server: JoinHandle<std::io::Result<()>>,
tasks: Tasks,
store_task: Option<JoinHandle<()>>,
}
impl Running {
pub fn app(&self) -> &Arc<App> {
&self.app
}
#[cfg(feature = "test-support")]
pub fn background_task_count(&self) -> usize {
self.tasks.count()
}
pub async fn shutdown(self) -> Result<(), StartupError> {
self.shutdown.cancel();
let served = match self.server.await {
Ok(result) => result.map_err(StartupError::Serve),
Err(_) => Ok(()),
};
http::logging::flush_summary(&self.app.delivery, self.app.clock.as_ref());
self.drain.cancel();
self.tasks.join().await;
http::logging::flush_drop_tail(&self.app.delivery);
drop(self.app);
if let Some(task) = self.store_task {
let _ = task.await;
}
served
}
}
#[derive(Debug)]
pub enum StartupError {
DataDir(store::startup::StartupError),
Bind {
addr: SocketAddr,
source: std::io::Error,
},
Serve(std::io::Error),
Delivery(&'static str),
}
impl fmt::Display for StartupError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
StartupError::DataDir(err) => write!(f, "{err}"),
StartupError::Bind { addr, source } => write!(f, "cannot bind {addr}: {source}"),
StartupError::Serve(source) => write!(f, "server stopped: {source}"),
StartupError::Delivery(reason) => write!(f, "{reason}"),
}
}
}
impl std::error::Error for StartupError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
StartupError::DataDir(err) => Some(err),
StartupError::Bind { source, .. } | StartupError::Serve(source) => Some(source),
StartupError::Delivery(_) => None,
}
}
}