Menu
AkurAI-Build
publicLatest change 834088c92ff481a8a0cc0f5825c8fb9524396e7a - Build lean Git-native AkurAI CI/CD by Ólafur Búi Ólafsson
use std::{net::SocketAddr, path::PathBuf, sync::Arc, time::Duration};
use anyhow::{Context, Result, ensure};
use axum::{
Json, Router,
body::{Body, Bytes},
extract::{DefaultBodyLimit, Path, State},
http::{HeaderMap, HeaderName, HeaderValue, Request, StatusCode, header},
middleware::{self, Next},
response::{Html, IntoResponse, Response},
routing::{get, post},
};
use hmac::{Hmac, Mac};
use minijinja::{Environment, context};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::Sha256;
use subtle::ConstantTimeEq;
use tokio_util::io::ReaderStream;
use tower::limit::GlobalConcurrencyLimitLayer;
use tower_http::{compression::CompressionLayer, timeout::TimeoutLayer, trace::TraceLayer};
use tracing::{info, warn};
use crate::{db::Database, runner::Runner};
const MAX_BODY: usize = 1024 * 1024;
const MAX_CONNECTIONS: usize = 1024;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const INDEX: &str = include_str!("../app/pages/index.html");
const BASE: &str = include_str!("../app/templates/layouts/base.html");
const CSS: &str = include_str!("../public/app.css");
const JS: &str = include_str!("../public/app.js");
type HmacSha256 = Hmac<Sha256>;
#[derive(Clone)]
pub struct ServerOptions {
pub address: SocketAddr,
pub token: String,
pub webhook_secret: String,
pub public_origin: String,
pub data_root: PathBuf,
pub allow_native: bool,
}
#[derive(Clone)]
struct AppState {
database: Database,
runner: Runner,
token: Arc<str>,
webhook_secret: Arc<str>,
public_origin: Arc<str>,
templates: Arc<Environment<'static>>,
}
#[derive(Serialize)]
struct Envelope<T: Serialize> {
ok: bool,
data: T,
}
#[derive(Serialize)]
struct ErrorEnvelope {
ok: bool,
error: ErrorMessage,
}
#[derive(Serialize)]
struct ErrorMessage {
code: &'static str,
message: String,
}
struct ApiError {
status: StatusCode,
code: &'static str,
message: String,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct TriggerRequest {
git_ref: Option<String>,
commit: Option<String>,
}
#[derive(Serialize)]
struct StateResponse {
repositories: Vec<crate::db::Repository>,
runs: Vec<crate::db::Run>,
}
#[derive(Serialize)]
struct IdResponse {
id: i64,
}
impl ApiError {
fn unauthorized() -> Self {
Self {
status: StatusCode::UNAUTHORIZED,
code: "unauthorized",
message: "valid bearer token required".into(),
}
}
fn bad(error: impl std::fmt::Display) -> Self {
Self {
status: StatusCode::BAD_REQUEST,
code: "invalid_request",
message: error.to_string(),
}
}
fn internal(error: impl std::fmt::Display) -> Self {
tracing::error!(error = %error, "request failed");
Self {
status: StatusCode::INTERNAL_SERVER_ERROR,
code: "internal",
message: "request failed".into(),
}
}
}
impl IntoResponse for ApiError {
fn into_response(self) -> Response {
(
self.status,
Json(ErrorEnvelope {
ok: false,
error: ErrorMessage {
code: self.code,
message: self.message,
},
}),
)
.into_response()
}
}
pub async fn serve(database: Database, options: ServerOptions) -> Result<()> {
ensure!(!options.token.is_empty(), "API token must not be empty");
ensure!(
!options.webhook_secret.is_empty(),
"webhook secret must not be empty"
);
database.validate_schema()?;
let recovered = database.recover_interrupted()?;
if recovered > 0 {
warn!(recovered, "marked interrupted work after restart");
}
let mut templates = Environment::new();
templates.add_template("layouts/base.html", BASE)?;
templates.add_template("index.html", INDEX)?;
let runner = Runner::new(database.clone(), options.data_root, options.allow_native)?;
let state = AppState {
database: database.clone(),
runner: runner.clone(),
token: options.token.into(),
webhook_secret: options.webhook_secret.into(),
public_origin: options.public_origin.into(),
templates: Arc::new(templates),
};
tokio::spawn(worker(database, runner));
let app = Router::new()
.route("/", get(index))
.route("/assets/app.css", get(css))
.route("/assets/app.js", get(js))
.route("/api/health", get(health))
.route("/api/state", get(api_state))
.route("/api/runs/{id}", get(run_detail))
.route("/api/runs/{id}/retry", post(retry))
.route("/api/runs/{id}/promote/{environment}", post(promote))
.route("/api/repos/{repository}/runs", post(trigger))
.route("/api/hooks/{repository}", post(webhook))
.route("/api/artifacts/{id}", get(artifact))
.with_state(state)
.layer(DefaultBodyLimit::max(MAX_BODY))
.layer(CompressionLayer::new())
.layer(TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
REQUEST_TIMEOUT,
))
.layer(TraceLayer::new_for_http())
.layer(middleware::from_fn(security_headers))
.layer(GlobalConcurrencyLimitLayer::new(MAX_CONNECTIONS));
let listener = tokio::net::TcpListener::bind(options.address).await?;
serve_bounded(listener, app).await
}
async fn worker(database: Database, runner: Runner) {
loop {
match database.next_queued_run() {
Ok(Some(id)) => {
let runner = runner.clone();
match tokio::task::spawn_blocking(move || runner.process(id)).await {
Ok(Ok(_)) => {}
Ok(Err(error)) => tracing::error!(run_id = id, ?error, "run failed"),
Err(error) => tracing::error!(run_id = id, ?error, "run worker panicked"),
}
}
Ok(None) => tokio::time::sleep(Duration::from_millis(500)).await,
Err(error) => {
tracing::error!(?error, "queue poll failed");
tokio::time::sleep(Duration::from_secs(2)).await;
}
}
}
}
async fn index(State(state): State<AppState>) -> Result<Html<String>, ApiError> {
let template = state
.templates
.get_template("index.html")
.map_err(ApiError::internal)?;
let page = template
.render(context! { version => env!("CARGO_PKG_VERSION") })
.map_err(ApiError::internal)?;
Ok(Html(page))
}
async fn css() -> Response {
asset(CSS, "text/css; charset=utf-8")
}
async fn js() -> Response {
asset(JS, "text/javascript; charset=utf-8")
}
fn asset(body: &'static str, content_type: &'static str) -> Response {
let mut response = Response::new(Body::from(body));
response
.headers_mut()
.insert(header::CONTENT_TYPE, HeaderValue::from_static(content_type));
response.headers_mut().insert(
header::CACHE_CONTROL,
HeaderValue::from_static("public, max-age=300"),
);
response
}
async fn health(State(state): State<AppState>) -> Result<Json<Envelope<Value>>, ApiError> {
state.database.check_ready().map_err(ApiError::internal)?;
Ok(ok(
serde_json::json!({"service":"akurai-build","status":"ok","version":env!("CARGO_PKG_VERSION")}),
))
}
async fn api_state(
State(state): State<AppState>,
headers: HeaderMap,
) -> Result<Json<Envelope<StateResponse>>, ApiError> {
authorize(&headers, &state)?;
let repositories = state.database.repositories().map_err(ApiError::internal)?;
let runs = state.database.runs(None, 100).map_err(ApiError::internal)?;
Ok(ok(StateResponse { repositories, runs }))
}
async fn run_detail(
State(state): State<AppState>,
headers: HeaderMap,
Path(id): Path<i64>,
) -> Result<Json<Envelope<crate::db::RunDetail>>, ApiError> {
authorize(&headers, &state)?;
Ok(ok(state.database.detail(id).map_err(ApiError::bad)?))
}
async fn trigger(
State(state): State<AppState>,
headers: HeaderMap,
Path(repository): Path<String>,
Json(request): Json<TriggerRequest>,
) -> Result<Json<Envelope<IdResponse>>, ApiError> {
authorize_mutation(&headers, &state)?;
let id = state
.runner
.queue(
&repository,
request.git_ref.as_deref(),
request.commit.as_deref(),
"manual",
)
.map_err(ApiError::bad)?;
Ok(ok(IdResponse { id }))
}
async fn retry(
State(state): State<AppState>,
headers: HeaderMap,
Path(id): Path<i64>,
) -> Result<Json<Envelope<IdResponse>>, ApiError> {
authorize_mutation(&headers, &state)?;
let id = state.database.retry(id).map_err(ApiError::bad)?;
Ok(ok(IdResponse { id }))
}
async fn promote(
State(state): State<AppState>,
headers: HeaderMap,
Path((id, environment)): Path<(i64, String)>,
) -> Result<Json<Envelope<Value>>, ApiError> {
authorize_mutation(&headers, &state)?;
let jobs = state
.database
.approve_environment(id, &environment)
.map_err(ApiError::bad)?;
Ok(ok(
serde_json::json!({"run_id":id,"environment":environment,"jobs":jobs}),
))
}
async fn artifact(
State(state): State<AppState>,
headers: HeaderMap,
Path(id): Path<i64>,
) -> Result<Response, ApiError> {
authorize(&headers, &state)?;
let artifact = state.database.artifact(id).map_err(ApiError::bad)?;
let path = state
.runner
.artifact_path(&artifact)
.map_err(ApiError::internal)?;
let file = tokio::fs::File::open(path)
.await
.map_err(ApiError::internal)?;
let mut response = Response::new(Body::from_stream(ReaderStream::new(file)));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/octet-stream"),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
HeaderValue::from_str(&artifact.bytes.to_string()).map_err(ApiError::internal)?,
);
let filename = artifact
.name
.rsplit('/')
.next()
.unwrap_or("artifact")
.replace(['"', '\r', '\n'], "_");
response.headers_mut().insert(
header::CONTENT_DISPOSITION,
HeaderValue::from_str(&format!("attachment; filename=\"{filename}\""))
.map_err(ApiError::internal)?,
);
Ok(response)
}
async fn webhook(
State(state): State<AppState>,
headers: HeaderMap,
Path(repository): Path<String>,
body: Bytes,
) -> Result<Json<Envelope<IdResponse>>, ApiError> {
authorize_webhook(&headers, &state, &body)?;
let payload: Value = serde_json::from_slice(&body).map_err(ApiError::bad)?;
let git_ref = payload
.get("ref")
.and_then(Value::as_str)
.context("webhook has no ref")
.map_err(ApiError::bad)?;
let commit = payload
.get("after")
.or_else(|| payload.get("checkout_sha"))
.and_then(Value::as_str)
.or_else(|| {
payload
.get("head_commit")
.and_then(|value| value.get("id"))
.and_then(Value::as_str)
});
let commit = commit.filter(|value| value.bytes().any(|byte| byte != b'0'));
let id = state
.runner
.queue(
&repository,
Some(git_ref.strip_prefix("refs/heads/").unwrap_or(git_ref)),
commit,
"webhook",
)
.map_err(ApiError::bad)?;
Ok(ok(IdResponse { id }))
}
fn authorize(headers: &HeaderMap, state: &AppState) -> Result<(), ApiError> {
let supplied = headers
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
.ok_or_else(ApiError::unauthorized)?;
if constant_time_equal(supplied.as_bytes(), state.token.as_bytes()) {
Ok(())
} else {
Err(ApiError::unauthorized())
}
}
fn authorize_mutation(headers: &HeaderMap, state: &AppState) -> Result<(), ApiError> {
authorize(headers, state)?;
if let Some(origin) = headers
.get(header::ORIGIN)
.and_then(|value| value.to_str().ok())
&& origin != state.public_origin.as_ref()
{
return Err(ApiError::bad("origin is not allowed"));
}
Ok(())
}
fn authorize_webhook(headers: &HeaderMap, state: &AppState, body: &[u8]) -> Result<(), ApiError> {
if let Some(token) = headers
.get("x-gitlab-token")
.and_then(|value| value.to_str().ok())
&& constant_time_equal(token.as_bytes(), state.webhook_secret.as_bytes())
{
return Ok(());
}
for name in ["x-hub-signature-256", "x-gitea-signature"] {
if let Some(signature) = headers.get(name).and_then(|value| value.to_str().ok()) {
let signature = signature.strip_prefix("sha256=").unwrap_or(signature);
let supplied = decode_hex(signature).map_err(ApiError::bad)?;
let mut mac = HmacSha256::new_from_slice(state.webhook_secret.as_bytes())
.map_err(ApiError::internal)?;
mac.update(body);
if constant_time_equal(&supplied, &mac.finalize().into_bytes()) {
return Ok(());
}
}
}
Err(ApiError::unauthorized())
}
fn decode_hex(value: &str) -> Result<Vec<u8>> {
ensure!(
value.len().is_multiple_of(2) && value.len() <= 128,
"invalid signature"
);
value
.as_bytes()
.chunks_exact(2)
.map(|pair| {
let text = std::str::from_utf8(pair)?;
Ok(u8::from_str_radix(text, 16)?)
})
.collect()
}
fn constant_time_equal(left: &[u8], right: &[u8]) -> bool {
left.len() == right.len() && left.ct_eq(right).into()
}
fn ok<T: Serialize>(data: T) -> Json<Envelope<T>> {
Json(Envelope { ok: true, data })
}
async fn security_headers(request: Request<Body>, next: Next) -> Response {
let mut response = next.run(request).await;
let is_json = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.contains("json"));
let headers = response.headers_mut();
headers.insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
headers.insert(header::X_FRAME_OPTIONS, HeaderValue::from_static("DENY"));
headers.insert(
header::REFERRER_POLICY,
HeaderValue::from_static("no-referrer"),
);
headers.insert(header::CONTENT_SECURITY_POLICY, HeaderValue::from_static("default-src 'self'; base-uri 'none'; object-src 'none'; frame-ancestors 'none'; form-action 'self'; script-src 'self'; style-src 'self'; connect-src 'self'"));
headers.insert(
HeaderName::from_static("permissions-policy"),
HeaderValue::from_static("camera=(), geolocation=(), microphone=()"),
);
headers.insert(
HeaderName::from_static("cross-origin-opener-policy"),
HeaderValue::from_static("same-origin"),
);
headers.insert(
HeaderName::from_static("cross-origin-resource-policy"),
HeaderValue::from_static("same-origin"),
);
if is_json {
headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
}
response
}
async fn serve_bounded(listener: tokio::net::TcpListener, app: Router) -> Result<()> {
use hyper_util::{
rt::{TokioExecutor, TokioIo, TokioTimer},
server::{conn::auto::Builder, graceful::GracefulShutdown},
service::TowerToHyperService,
};
info!(address = %listener.local_addr()?, "AkurAI Build listening");
let connections = Arc::new(tokio::sync::Semaphore::new(MAX_CONNECTIONS));
let graceful = GracefulShutdown::new();
let mut shutdown = std::pin::pin!(shutdown_signal());
loop {
let permit = tokio::select! { () = &mut shutdown => break, permit = connections.clone().acquire_owned() => permit? };
let (socket, _) = tokio::select! { () = &mut shutdown => break, accepted = listener.accept() => accepted? };
let mut builder = Builder::new(TokioExecutor::new());
builder
.http1()
.timer(TokioTimer::new())
.header_read_timeout(REQUEST_TIMEOUT);
let connection = graceful.watch(
builder
.serve_connection_with_upgrades(
TokioIo::new(socket),
TowerToHyperService::new(app.clone()),
)
.into_owned(),
);
tokio::spawn(async move {
if let Err(error) = connection.await {
tracing::debug!(?error, "connection closed");
}
drop(permit);
});
}
drop(listener);
tokio::select! { () = graceful.shutdown() => {}, () = tokio::time::sleep(REQUEST_TIMEOUT) => warn!("graceful shutdown timed out") }
Ok(())
}
async fn shutdown_signal() {
let ctrl_c = async {
if let Err(error) = tokio::signal::ctrl_c().await {
warn!(?error, "failed to install Ctrl-C handler");
}
};
#[cfg(unix)]
let terminate = async {
match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
Ok(mut signal) => {
signal.recv().await;
}
Err(error) => warn!(?error, "failed to install terminate handler"),
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! { () = ctrl_c => {}, () = terminate => {} }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn templates_compile_and_render_without_inline_code() -> Result<()> {
let mut environment = Environment::new();
environment.add_template("layouts/base.html", BASE)?;
environment.add_template("index.html", INDEX)?;
let page = environment
.get_template("index.html")?
.render(context! { version => "test" })?;
assert!(page.contains("AkurAI Builds"));
assert!(!page.contains("<script>"));
Ok(())
}
#[test]
fn authentication_helpers_fail_closed() {
assert!(constant_time_equal(b"same", b"same"));
assert!(!constant_time_equal(b"same", b"different"));
assert!(decode_hex("not-hex").is_err());
}
}