use std::sync::Arc;
use axum::{Json, Router, http::StatusCode, response::IntoResponse, routing::get};
use serde_json::json;
use tower_http::limit::RequestBodyLimitLayer;
use crate::auth_resolver::{AuthResolver, ExtraAuthValidator, resolver_from_validator};
use crate::error::ServerError;
use crate::lifecycle;
use crate::middleware;
use crate::openapi;
use crate::routers;
use crate::state::AppState;
async fn root() -> impl IntoResponse {
(
StatusCode::OK,
Json(json!({"message": "Hello, World, I am alive!"})),
)
}
pub struct RouterBuilder {
state: AppState,
extra_routers: Vec<(&'static str, Router<AppState>)>,
extra_validator: Option<Arc<dyn ExtraAuthValidator>>,
auth_resolver: Option<Arc<dyn AuthResolver>>,
}
impl RouterBuilder {
pub fn new(state: AppState) -> Self {
Self {
state,
extra_routers: Vec::new(),
extra_validator: None,
auth_resolver: None,
}
}
pub fn with_router(mut self, mount: &'static str, r: Router<AppState>) -> Self {
self.extra_routers.push((mount, r));
self
}
pub fn with_extra_validator(mut self, v: Arc<dyn ExtraAuthValidator>) -> Self {
self.extra_validator = Some(v);
self
}
pub fn with_auth_resolver(mut self, r: Arc<dyn AuthResolver>) -> Self {
self.auth_resolver = Some(r);
self
}
pub async fn build(mut self) -> Result<Router, ServerError> {
if let Some(r) = self.auth_resolver.take() {
self.state.auth_resolver = Some(r);
} else if let Some(v) = self.extra_validator.take() {
self.state.auth_resolver = Some(resolver_from_validator(v));
}
lifecycle::on_startup(&self.state).await?;
let body_limit = self.state.config.body_limit;
let mut app = Router::new()
.route("/", get(root))
.nest("/health", routers::health::router())
.route("/openapi.json", get(openapi::openapi_json))
.nest("/api/v1/add", routers::add::router())
.nest("/api/v1/datasets", routers::datasets::router())
.nest("/api/v1/ontologies", routers::ontologies::router())
.nest("/api/v1/delete", routers::delete::router())
.nest("/api/v1/update", routers::update::router())
.nest("/api/v1/forget", routers::forget::router())
.nest("/api/v1/cognify", routers::cognify::router())
.nest("/api/v1/memify", routers::memify::router())
.nest("/api/v1/remember", routers::remember::router())
.nest("/api/v1/improve", routers::improve::router())
.nest("/api/v1/search", routers::search::router())
.nest("/api/v1/recall", routers::recall::router())
.nest("/api/v1/sessions", routers::sessions::router())
.nest("/api/v1/llm", routers::llm::router())
.nest("/api/v1/visualize", routers::visualize::router())
.nest("/api/v1/settings", routers::settings::router())
.nest("/api/v1/activity", routers::activity::router())
.nest("/api/v1/notebooks", routers::notebooks::router())
.nest("/api/v1/responses", routers::responses::router());
for (mount, r) in self.extra_routers.drain(..) {
app = app.nest(mount, r);
}
let app = app
.layer(RequestBodyLimitLayer::new(body_limit))
.layer(middleware::cors::cors_layer(&self.state.config))
.layer(middleware::tracing::trace_layer())
.with_state(self.state);
Ok(app)
}
}
pub async fn build_router(state: AppState) -> Result<Router, ServerError> {
RouterBuilder::new(state).build().await
}