dialtone_axum 0.1.0

Dialtone Axum Back-end
Documentation
use std::fmt::Display;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;

use axum::error_handling::HandleErrorLayer;
use axum::extract::{FromRequest, RequestParts};
use axum::http::Method;
use axum::{async_trait, extract::TypedHeader, http::StatusCode, Router};
use chrono::Utc;
use dialtone_common::authz::user_authz_info::UserAuthzInfo;
use headers::{authorization::Bearer, Authorization};
use jsonwebtoken::{decode, DecodingKey, EncodingKey, Validation};
use lazy_static::lazy_static;
use serde::{Deserialize, Serialize};
use sqlx::{Pool, Postgres};
use tokio_cron_scheduler::{Job, JobScheduler};
use tower::{BoxError, ServiceBuilder};
use tower_http::cors::{Any, CorsLayer};
use tower_http::{add_extension::AddExtensionLayer, trace::TraceLayer};

use dialtone_common::rest::users::web_user::UserStatus;
use dialtone_common::utils::version::DT_VERSION;
use dialtone_sqlx::db::persistent_queue::listener::listener;
use dialtone_sqlx::db::persistent_queue::notifier::notify;
use dialtone_sqlx::db::persistent_queue::pop_queue::pop_queue;
use dialtone_sqlx::db::persistent_queue::JobDbType;
use dialtone_sqlx::db::user_principal::fetch_auth_and_mark_seen::fetch_auth_and_mark_seen;

use crate::api_v1::routes::api_v1_routes;
use crate::forwarded_for::get_forwarded;
use crate::response::response_error::ResponseError;

lazy_static! {
    pub static ref KEYS: Keys = {
        let secret = std::env::var("JWT_SECRET").expect("JWT_SECRET must be set");
        Keys::new(secret.as_bytes())
    };
}

pub async fn start_server(pg_pool: Pool<Postgres>) -> anyhow::Result<()> {
    tracing_subscriber::fmt::init();
    tracing::info!("dialtone version {}", DT_VERSION);

    start_queues(&pg_pool)?;
    start_jobs(&pg_pool)?;

    let app = app_router(pg_pool);

    // run it
    let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
    tracing::debug!("listening on {}", addr);
    axum::Server::bind(&addr)
        .serve(app.into_make_service())
        .await?;
    Ok(())
}

fn start_queues(pg_pool: &Pool<Postgres>) -> anyhow::Result<()> {
    let pg_pool = pg_pool.clone();
    tokio::spawn(async move {
        listener(
            &pg_pool,
            &JobDbType::NoOperation,
            move |pool, notification| async move {
                tracing::info!("queue notification {:?}", notification);
                pop_queue(&pool, &JobDbType::NoOperation).await.unwrap();
            },
            true,
        )
        .await
    });
    Ok(())
}

fn start_jobs(pg_pool: &Pool<Postgres>) -> anyhow::Result<()> {
    let pg_pool = pg_pool.clone();
    let sched = JobScheduler::new().expect("cron schedulaer error");
    let job = Job::new_repeated(Duration::from_secs(60), move |_uuid, _l| {
        let pg_pool = pg_pool.clone();
        tokio::spawn(async move { notify(&pg_pool, &JobDbType::NoOperation).await.unwrap() });
    })
    .unwrap();
    sched.add(job).expect("adding job");
    sched.start()?;
    Ok(())
}

pub fn app_router(pg_pool: Pool<Postgres>) -> Router {
    let api_routes = api_v1_routes();
    Router::new().nest(&api_routes.0, api_routes.1).layer(
        ServiceBuilder::new()
            .layer(HandleErrorLayer::new(|error: BoxError| async move {
                if error.is::<tower::timeout::error::Elapsed>() {
                    Ok(StatusCode::REQUEST_TIMEOUT)
                } else {
                    Err((
                        StatusCode::INTERNAL_SERVER_ERROR,
                        format!("Unhandled internal error: {}", error),
                    ))
                }
            }))
            .timeout(Duration::from_secs(10))
            .layer(TraceLayer::new_for_http())
            .layer(
                CorsLayer::new()
                    .allow_origin(Any)
                    .allow_methods(vec![
                        Method::GET,
                        Method::POST,
                        Method::OPTIONS,
                        Method::DELETE,
                        Method::PUT,
                        Method::PATCH,
                    ])
                    .allow_headers(Any),
            )
            .layer(AddExtensionLayer::new(pg_pool))
            .layer(AddExtensionLayer::new(SharedState::default()))
            .into_inner(),
    )
}

pub type SharedState = Arc<tokio::sync::RwLock<State>>;

#[derive(Default)]
pub struct State {
    pub user_authz: UserAuthzInfo,
}

pub fn internal_error<E>(err: E) -> (StatusCode, String)
where
    E: std::error::Error,
{
    (StatusCode::INTERNAL_SERVER_ERROR, err.to_string())
}

impl Display for Claims {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(f, "User Account: {}", self.sub)
    }
}

#[async_trait]
impl<B> FromRequest<B> for Claims
where
    B: Send,
{
    type Rejection = ResponseError;

    async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
        // Extract the token from the authorization header
        let TypedHeader(Authorization(bearer)) =
            TypedHeader::<Authorization<Bearer>>::from_request(req)
                .await
                .map_err(|_| ResponseError::MissingCredentials(None))?;
        // Decode the users data
        let token_data = decode::<Claims>(bearer.token(), &KEYS.decoding, &Validation::default())
            .map_err(|err| {
            tracing::warn!("error with token: {}", err.to_string());
            ResponseError::TokenDecodingError(None)
        })?;

        // get forwarded for
        let headers = req.headers();
        let forwarded_for = get_forwarded(headers);

        // get the pg pool
        let extensions = req.extensions_mut();
        let pool = extensions.get::<Pool<Postgres>>().unwrap();
        let user_authz_info =
            fetch_auth_and_mark_seen(pool, &token_data.claims.sub, &forwarded_for, Utc::now())
                .await;
        match user_authz_info {
            Ok(info) => match &info {
                None => Err(ResponseError::NoAuthorizationsAvailable(None)),
                Some(us) => match us.status {
                    UserStatus::Active => {
                        let shared_state = extensions.get::<SharedState>().unwrap();
                        shared_state.write().await.user_authz = us.clone();
                        Ok(token_data.claims)
                    }
                    UserStatus::Suspended => Err(ResponseError::UnableToAuthorize(None)),
                    UserStatus::PendingApproval => Err(ResponseError::UnableToAuthorize(None)),
                },
            },
            Err(_) => Err(ResponseError::UnableToAuthenticate(None)),
        }
    }
}

pub struct Keys {
    pub encoding: EncodingKey,
    pub decoding: DecodingKey,
}

impl Keys {
    fn new(secret: &[u8]) -> Self {
        Self {
            encoding: EncodingKey::from_secret(secret),
            decoding: DecodingKey::from_secret(secret),
        }
    }
}

#[derive(Debug, Serialize, Deserialize)]
pub struct Claims {
    pub sub: String,
    pub exp: u64,
}