pub mod middleware_auth;
pub mod middleware_correlation;
pub mod middleware_identity;
pub mod middleware_log;
pub mod module;
pub mod routes;
pub use middleware_auth::{
AuthGuard, Authenticated, RequireAuth, auth_middleware, optional_auth_middleware,
require_auth_middleware,
};
pub use middleware_correlation::{X_REQUEST_ID, correlation_id_middleware};
pub use middleware_identity::{
ARQEN_IDENTITY, POWERED_BY_HEADER, SERVER_HEADER, identity_middleware,
};
pub use middleware_log::{RequestLogConfig, logging_middleware};
pub use module::{HttpModule, merge_module_routes};
pub use routes::{agent, agent_manifest, docs, health, ready};
use axum::extract::FromRef;
use axum::{
http::StatusCode,
routing::{get, post},
};
use std::net::SocketAddr;
use tokio::net::TcpListener;
use tower_http::compression::{CompressionLayer, predicate::SizeAbove};
use tower_http::cors::{Any, CorsLayer};
use tower_http::limit::RequestBodyLimitLayer;
use tower_http::timeout::TimeoutLayer;
use crate::state::AppState;
pub use axum::{Router, body, extract, http, middleware, response, routing};
pub fn create_router() -> Router {
let state = AppState::builder()
.build()
.expect("failed to build default state");
create_router_with_state(state)
}
pub fn builtin_routes<S>(state: &AppState) -> Router<S>
where
S: Clone + Send + Sync + 'static,
AppState: FromRef<S>,
{
let cors = CorsLayer::new()
.allow_origin(Any)
.allow_methods(Any)
.allow_headers(Any);
let timeout = TimeoutLayer::with_status_code(
StatusCode::GATEWAY_TIMEOUT,
state.config.server.request_timeout,
);
let body_limit = RequestBodyLimitLayer::new(state.config.server.max_body_size);
let compression = CompressionLayer::new().compress_when(SizeAbove::new(
state
.config
.server
.compression_threshold
.min(u16::MAX as usize) as u16,
));
let request_log_config = RequestLogConfig {
success_sample_rate: if std::env::var("ARQEN_ENV").as_deref() == Ok("production") {
state.config.server.request_log_sample_rate
} else {
1.0
},
slow_request_threshold: state.config.server.slow_request_threshold,
};
Router::new()
.route("/health", get(routes::health))
.route("/ready", get(routes::ready))
.route("/agent", get(routes::agent))
.route("/agent/manifest", get(routes::agent_manifest))
.route("/agent/tools/:name", post(routes::tool_invoke))
.route("/docs", get(routes::docs))
.layer(body_limit)
.layer(compression)
.layer(timeout)
.layer(cors)
.layer(axum::Extension(request_log_config))
.layer(middleware::from_fn(
middleware_identity::identity_middleware,
))
.layer(middleware::from_fn(
middleware_correlation::correlation_id_middleware,
))
.layer(middleware::from_fn(middleware_log::logging_middleware))
}
pub fn create_router_with_state(state: AppState) -> Router {
builtin_routes::<AppState>(&state).with_state(state)
}
pub fn create_router_with_state_and_routes(state: AppState, app_routes: Router) -> Router {
create_router_with_state(state)
.merge(app_routes)
.layer(middleware::from_fn(
middleware_identity::identity_middleware,
))
}
pub fn nest_routes(state: AppState, prefix: &str, app_routes: Router) -> Router {
create_router_with_state(state)
.nest(prefix, app_routes)
.layer(middleware::from_fn(
middleware_identity::identity_middleware,
))
}
pub async fn start_server(
addr: SocketAddr,
router: Router,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let listener = TcpListener::bind(addr).await?;
tracing::info!("Server listening on {}", addr);
axum::serve(
listener,
router.layer(middleware::from_fn(
middleware_identity::identity_middleware,
)),
)
.with_graceful_shutdown(shutdown_signal())
.await?;
Ok(())
}
async fn shutdown_signal() {
let ctrl_c = async {
if let Err(error) = tokio::signal::ctrl_c().await {
tracing::error!(%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) => tracing::error!(%error, "failed to install SIGTERM handler"),
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
() = ctrl_c => {},
() = terminate => {},
}
tracing::info!("shutdown signal received; draining requests");
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::extract::State;
use axum::http::{Request, StatusCode};
use axum::routing::get;
use tower::ServiceExt;
async fn test_handler() -> &'static str {
"ok"
}
#[tokio::test]
async fn test_create_router() {
let router = create_router();
let request = Request::builder()
.uri("/health")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_health_endpoint() {
let router = create_router();
let request = Request::builder()
.uri("/health")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_ready_endpoint() {
let router = create_router();
let request = Request::builder()
.uri("/ready")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_agent_endpoint() {
let router = create_router();
let request = Request::builder()
.uri("/agent")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_docs_endpoint() {
let router = create_router();
let request = Request::builder().uri("/docs").body(Body::empty()).unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_cors_headers() {
let router = create_router();
let request = Request::builder()
.method("OPTIONS")
.uri("/health")
.header("origin", "http://example.com")
.header("access-control-request-method", "GET")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert!(
response
.headers()
.contains_key("access-control-allow-origin"),
"CORS headers should be present"
);
}
#[tokio::test]
async fn test_response_compression_respects_threshold_and_negotiation() {
let mut config = crate::AppConfig::default();
config.server.compression_threshold = 1;
let state = AppState::builder().with_config(config).build().unwrap();
let router = create_router_with_state(state);
let response = router
.oneshot(
Request::builder()
.uri("/docs")
.header("accept-encoding", "gzip")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers().get("content-encoding").unwrap(), "gzip");
}
#[tokio::test]
async fn test_create_router_with_state_and_routes() {
let state = AppState::builder().build().unwrap();
let app_routes = Router::new().route("/api/hello", get(test_handler));
let router = create_router_with_state_and_routes(state, app_routes);
let built_in = router
.clone()
.oneshot(
Request::builder()
.uri("/health")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(built_in.status(), StatusCode::OK);
let app = router
.oneshot(
Request::builder()
.uri("/api/hello")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(app.status(), StatusCode::OK);
assert_eq!(app.headers()[SERVER_HEADER], ARQEN_IDENTITY);
assert_eq!(app.headers()[POWERED_BY_HEADER], ARQEN_IDENTITY);
}
#[tokio::test]
async fn test_nest_routes() {
let state = AppState::builder().build().unwrap();
let app_routes = Router::new().route("/users", get(test_handler));
let router = nest_routes(state, "/api/v1", app_routes);
let built_in = router
.clone()
.oneshot(
Request::builder()
.uri("/health")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(built_in.status(), StatusCode::OK);
let app = router
.oneshot(
Request::builder()
.uri("/api/v1/users")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(app.status(), StatusCode::OK);
assert_eq!(app.headers()[SERVER_HEADER], ARQEN_IDENTITY);
assert_eq!(app.headers()[POWERED_BY_HEADER], ARQEN_IDENTITY);
}
#[tokio::test]
async fn test_nest_routes_returns_404_for_missing() {
let state = AppState::builder().build().unwrap();
let app_routes = Router::new().route("/users", get(test_handler));
let router = nest_routes(state, "/api/v1", app_routes);
let response = router
.oneshot(
Request::builder()
.uri("/api/v1/nonexistent")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[derive(Clone)]
struct CustomState {
arqen: AppState,
custom: String,
}
impl FromRef<CustomState> for AppState {
fn from_ref(state: &CustomState) -> Self {
state.arqen.clone()
}
}
async fn custom_state_handler(State(state): State<CustomState>) -> String {
state.custom
}
fn custom_state_router() -> Router<CustomState> {
let app_state = AppState::builder().build().unwrap();
builtin_routes::<CustomState>(&app_state).route("/custom", get(custom_state_handler))
}
#[tokio::test]
async fn test_builtin_routes_with_custom_state() {
let app_state = AppState::builder().build().unwrap();
let router = custom_state_router().with_state(CustomState {
arqen: app_state,
custom: "hello".to_string(),
});
for path in ["/health", "/ready", "/agent", "/agent/manifest", "/docs"] {
let response = router
.clone()
.oneshot(Request::builder().uri(path).body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(
response.status(),
StatusCode::OK,
"path {path} should respond"
);
}
let app = router
.oneshot(
Request::builder()
.uri("/custom")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(app.status(), StatusCode::OK);
let body = axum::body::to_bytes(app.into_body(), 1024).await.unwrap();
assert_eq!(&body[..], b"hello");
}
}