AkurAI Build
Menu

AkurAI-Build

public

Latest 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());
    }
}