use std::sync::Arc;
use axum::{
extract::State,
response::IntoResponse,
routing::{get, post},
Json, Router,
};
use openkind_core::{ResponseContract, SystemRequest};
use openkind_engine::{dispatch, EngineRegistry};
use tower_http::trace::{DefaultMakeSpan, TraceLayer};
use tracing::Level;
use crate::error::ApiError;
use crate::middleware::AuthConfig;
use crate::models::ModelsResponse;
use crate::AppState;
pub const MAX_PAYLOAD_SIZE_BYTES: usize = 16 * 1024 * 1024;
pub fn router_with_state_and_limit(
state: AppState,
auth: AuthConfig,
max_payload_bytes: usize,
) -> Router {
router_full(
state,
auth,
max_payload_bytes,
crate::middleware::RateLimiter::new(crate::middleware::RateLimitConfig::default()),
false,
false,
)
}
pub fn router_with_state_auth_rate_limit(
state: AppState,
auth: AuthConfig,
max_payload_bytes: usize,
rate_limiter: crate::middleware::RateLimiter,
) -> Router {
router_full(state, auth, max_payload_bytes, rate_limiter, false, false)
}
pub fn router_daemon(
state: AppState,
auth: AuthConfig,
max_payload_bytes: usize,
rate_limiter: crate::middleware::RateLimiter,
playground: bool,
) -> Router {
router_full(
state,
auth,
max_payload_bytes,
rate_limiter,
playground,
false,
)
}
pub fn router_daemon_with_arrow(
state: AppState,
auth: AuthConfig,
max_payload_bytes: usize,
rate_limiter: crate::middleware::RateLimiter,
playground: bool,
arrow: bool,
) -> Router {
router_full(
state,
auth,
max_payload_bytes,
rate_limiter,
playground,
arrow,
)
}
fn router_full(
state: AppState,
auth: AuthConfig,
max_payload_bytes: usize,
rate_limiter: crate::middleware::RateLimiter,
playground: bool,
arrow: bool,
) -> Router {
let routes = Router::new()
.route("/v1/systemone", post(systemone))
.route("/v1/system_one", post(systemone))
.route("/v1/models", get(list_models))
.route("/health", get(health))
.route("/metrics", get(prometheus_metrics));
let routes = if playground {
routes
.route("/playground", get(crate::playground::playground_page))
.route(
"/playground/api/models",
get(crate::playground::list_models).post(crate::playground::change_model),
)
} else {
routes
};
let routes = if arrow {
routes.route("/v1/arrow", post(crate::arrow::arrow_batch))
} else {
routes
};
let routes = if rate_limiter.is_enabled() {
routes.layer(axum::middleware::from_fn_with_state(
rate_limiter,
crate::middleware::rate_limit_layer,
))
} else {
routes
};
routes
.layer(axum::middleware::from_fn_with_state(
auth,
crate::middleware::auth_layer,
))
.layer(axum::middleware::from_fn(
crate::middleware::request_id_layer,
))
.layer(axum::extract::DefaultBodyLimit::max(max_payload_bytes))
.layer(
TraceLayer::new_for_http().make_span_with(
DefaultMakeSpan::new()
.level(Level::INFO)
.include_headers(false),
),
)
.with_state(Arc::new(state))
}
pub fn router_with_state(state: AppState, auth: AuthConfig) -> Router {
router_with_state_and_limit(state, auth, MAX_PAYLOAD_SIZE_BYTES)
}
pub fn router(registry: EngineRegistry) -> Router {
router_with_state(AppState::new(registry), AuthConfig::default())
}
pub fn router_with_auth(registry: EngineRegistry, auth: AuthConfig) -> Router {
router_with_state(AppState::new(registry), auth)
}
pub use router_with_state as build_router_with_state;
async fn systemone(
State(state): State<Arc<AppState>>,
headers: axum::http::HeaderMap,
req: Result<Json<SystemRequest>, axum::extract::rejection::JsonRejection>,
) -> Result<axum::response::Response, ApiError> {
let Json(req) = match req {
Ok(j) => j,
Err(rejection) => match rejection {
axum::extract::rejection::JsonRejection::BytesRejection(e) => {
return Err(ApiError::PayloadTooLarge(e.to_string()));
}
axum::extract::rejection::JsonRejection::JsonSyntaxError(e) => {
return Err(ApiError::BadJson(e.to_string()));
}
axum::extract::rejection::JsonRejection::JsonDataError(e) => {
return Err(ApiError::InvalidBody(e.to_string()));
}
other => {
return Err(ApiError::InvalidBody(other.to_string()));
}
},
};
if let Some(proxy) = &state.proxy {
if proxy.wants(&req) {
let contract = ResponseContract::from_request(&req)
.map_err(|error| ApiError::InvalidBody(error.to_string()))?;
let caller_key = if proxy.forwards_caller_credentials() {
bearer_of(&headers)
} else {
None
};
let outcome = proxy.evaluate(req, caller_key).await?;
contract.validate(&outcome.response).map_err(|error| {
ApiError::BadGateway(format!("upstream returned an invalid response: {error}"))
})?;
let mut response = (axum::http::StatusCode::OK, Json(outcome.response)).into_response();
let headers = response.headers_mut();
if let Ok(value) = axum::http::HeaderValue::from_str(outcome.source.as_str()) {
headers.insert(
axum::http::HeaderName::from_static("x-openkind-cache"),
value,
);
}
if let Some(detail) = &outcome.detail {
if let Ok(text) = serde_json::to_string(detail) {
if let Ok(value) = axum::http::HeaderValue::from_str(&text) {
headers.insert(
axum::http::HeaderName::from_static("x-openkind-cache-detail"),
value,
);
}
}
}
return Ok(response);
}
}
let resp = dispatch(req, &state.registry).await?;
Ok((axum::http::StatusCode::OK, Json(resp)).into_response())
}
fn bearer_of(headers: &axum::http::HeaderMap) -> Option<String> {
let value = headers.get(axum::http::header::AUTHORIZATION)?;
let value = value.to_str().ok()?;
let token = value
.strip_prefix("Bearer ")
.or_else(|| value.strip_prefix("bearer "))?;
let token = token.trim();
if token.is_empty() {
None
} else {
Some(token.to_owned())
}
}
async fn list_models(State(state): State<Arc<AppState>>) -> impl IntoResponse {
if let Some(proxy) = &state.proxy {
if let Some(models) = proxy.models().await {
return Json(models).into_response();
}
}
let models = state.registry.list_models();
Json(ModelsResponse::new(models)).into_response()
}
async fn health() -> impl IntoResponse {
static HEALTH_BODY: &str = "{\"status\":\"ok\"}";
(
[(
axum::http::header::CONTENT_TYPE,
axum::http::HeaderValue::from_static("application/json"),
)],
HEALTH_BODY,
)
}
static HANDLE: std::sync::OnceLock<metrics_exporter_prometheus::PrometheusHandle> =
std::sync::OnceLock::new();
async fn prometheus_metrics() -> impl IntoResponse {
if let Some(h) = HANDLE.get() {
(
axum::http::StatusCode::OK,
[("content-type", "text/plain; version=0.0.4")],
h.render(),
)
} else {
(
axum::http::StatusCode::OK,
[("content-type", "text/plain; version=0.0.4")],
"# metrics recorder not installed\n".to_string(),
)
}
}
pub fn install_metrics_recorder() -> anyhow::Result<()> {
if HANDLE.get().is_some() {
return Ok(());
}
use metrics_exporter_prometheus::PrometheusBuilder;
let handle = PrometheusBuilder::new()
.install_recorder()
.map_err(|e| anyhow::anyhow!("install metrics recorder: {e}"))?;
let _ = HANDLE.set(handle);
Ok(())
}
#[cfg(test)]
#[path = "http_tests.rs"]
mod tests;