Files
teaching-feedback-assistant/server/src/main.rs

105 lines
3.4 KiB
Rust

mod auth;
mod config;
mod error;
mod recording_routes;
mod routes;
mod speech;
mod summary;
mod transcription_worker;
use std::{path::PathBuf, sync::Arc};
use axum::{Json, Router, extract::DefaultBodyLimit, http::HeaderValue, routing::get};
use config::Config;
use sqlx::{PgPool, postgres::PgPoolOptions};
use tower_http::{cors::CorsLayer, trace::TraceLayer};
use tracing_subscriber::EnvFilter;
use utoipa_scalar::{Scalar, Servable};
#[derive(Clone)]
pub struct AppState {
pub pool: Option<PgPool>,
pub audio_storage_dir: PathBuf,
pub speech: speech::SpeechService,
pub summary: summary::SummaryService,
pub auth: auth::AuthService,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
dotenvy::dotenv().ok();
tracing_subscriber::fmt()
.with_env_filter(EnvFilter::try_from_default_env().unwrap_or_else(|_| "info".into()))
.init();
let config = Config::from_env().map_err(std::io::Error::other)?;
let pool = connect_database(&config).await?;
let speech = speech::SpeechService::new(config.speech.clone());
let auth = auth::AuthService::new(config.auth.clone());
let state = AppState {
pool,
audio_storage_dir: config.audio_storage_dir.clone(),
speech: speech.clone(),
summary: summary::SummaryService::new(config.summary.clone()),
auth,
};
if let Some(pool) = state.pool.clone() {
if speech.configured() {
transcription_worker::spawn(pool, speech);
}
}
let app = build_app(state, config.cors_allowed_origin.as_deref())?;
let address = std::net::SocketAddr::from((config.host, config.port));
let listener = tokio::net::TcpListener::bind(address).await?;
tracing::info!(%address, scalar = %format!("http://{address}/scalar"), "API server is listening");
axum::serve(listener, app)
.with_graceful_shutdown(shutdown_signal())
.await?;
Ok(())
}
pub fn build_app(state: AppState, cors_allowed_origin: Option<&str>) -> Result<Router, String> {
let (api_router, openapi) = routes::router();
let openapi_json = openapi.clone();
let mut app = api_router
.route(
"/openapi.json",
get(move || {
let openapi = openapi_json.clone();
async move { Json(openapi) }
}),
)
.merge(Scalar::with_url("/scalar", openapi))
.with_state(Arc::new(state))
.layer(DefaultBodyLimit::max(20 * 1024 * 1024))
.layer(TraceLayer::new_for_http());
if let Some(origin) = cors_allowed_origin {
let allowed_origin = HeaderValue::from_str(origin)
.map_err(|_| "CORS_ALLOWED_ORIGIN must be a valid HTTP header value".to_owned())?;
app = app.layer(CorsLayer::new().allow_origin(allowed_origin));
}
Ok(app)
}
async fn connect_database(config: &Config) -> Result<Option<PgPool>, sqlx::Error> {
let Some(database_url) = &config.database_url else {
tracing::warn!("DATABASE_URL is not set; data endpoints will return HTTP 503");
return Ok(None);
};
let pool = PgPoolOptions::new()
.max_connections(config.database_max_connections)
.connect(database_url)
.await?;
tracing::info!("connected to PostgreSQL");
Ok(Some(pool))
}
async fn shutdown_signal() {
let _ = tokio::signal::ctrl_c().await;
tracing::info!("shutdown signal received");
}