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);
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> {
let TypedHeader(Authorization(bearer)) =
TypedHeader::<Authorization<Bearer>>::from_request(req)
.await
.map_err(|_| ResponseError::MissingCredentials(None))?;
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)
})?;
let headers = req.headers();
let forwarded_for = get_forwarded(headers);
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,
}