1use std::sync::Arc;
17
18use axum::{
19 extract::State,
20 response::IntoResponse,
21 routing::{get, post},
22 Json, Router,
23};
24use openkind_core::{ResponseContract, SystemRequest};
25use openkind_engine::{dispatch, EngineRegistry};
26use tower_http::trace::{DefaultMakeSpan, TraceLayer};
27use tracing::Level;
28
29use crate::error::ApiError;
30use crate::middleware::AuthConfig;
31use crate::models::ModelsResponse;
32use crate::AppState;
33
34pub const MAX_PAYLOAD_SIZE_BYTES: usize = 16 * 1024 * 1024;
36
37pub fn router_with_state_and_limit(
40 state: AppState,
41 auth: AuthConfig,
42 max_payload_bytes: usize,
43) -> Router {
44 router_full(
45 state,
46 auth,
47 max_payload_bytes,
48 crate::middleware::RequestLimits::default(),
49 false,
50 false,
51 )
52}
53
54pub fn router_with_state_auth_rate_limit(
56 state: AppState,
57 auth: AuthConfig,
58 max_payload_bytes: usize,
59 rate_limiter: crate::middleware::RateLimiter,
60) -> Router {
61 router_full(
62 state,
63 auth,
64 max_payload_bytes,
65 rate_limiter.into(),
66 false,
67 false,
68 )
69}
70
71pub fn router_daemon(
75 state: AppState,
76 auth: AuthConfig,
77 max_payload_bytes: usize,
78 rate_limiter: crate::middleware::RateLimiter,
79 playground: bool,
80) -> Router {
81 router_full(
82 state,
83 auth,
84 max_payload_bytes,
85 rate_limiter.into(),
86 playground,
87 false,
88 )
89}
90
91pub fn router_daemon_with_arrow(
94 state: AppState,
95 auth: AuthConfig,
96 max_payload_bytes: usize,
97 rate_limiter: crate::middleware::RateLimiter,
98 playground: bool,
99 arrow: bool,
100) -> Router {
101 router_full(
102 state,
103 auth,
104 max_payload_bytes,
105 rate_limiter.into(),
106 playground,
107 arrow,
108 )
109}
110
111pub fn router_daemon_with_arrow_and_limits(
113 state: AppState,
114 auth: AuthConfig,
115 max_payload_bytes: usize,
116 limits: crate::middleware::RequestLimits,
117 playground: bool,
118 arrow: bool,
119) -> Router {
120 router_full(state, auth, max_payload_bytes, limits, playground, arrow)
121}
122
123fn router_full(
124 state: AppState,
125 auth: AuthConfig,
126 max_payload_bytes: usize,
127 limits: crate::middleware::RequestLimits,
128 playground: bool,
129 arrow: bool,
130) -> Router {
131 let routes = Router::new()
132 .route("/v1/systemone", post(systemone))
134 .route("/v1/system_one", post(systemone))
136 .route("/v1/models", get(list_models))
138 .route("/health", get(health))
140 .route("/metrics", get(prometheus_metrics));
142 let routes = if playground {
146 routes
147 .route("/playground", get(crate::playground::playground_page))
148 .route(
149 "/playground/api/models",
150 get(crate::playground::list_models).post(crate::playground::change_model),
151 )
152 } else {
153 routes
154 };
155 let routes = if arrow {
159 routes.route("/v1/arrow", post(crate::arrow::arrow_batch))
160 } else {
161 routes
162 };
163 let routes = if limits.evaluation.is_enabled() {
166 routes.layer(axum::middleware::from_fn_with_state(
167 limits.evaluation,
168 crate::middleware::rate_limit_layer,
169 ))
170 } else {
171 routes
172 };
173 routes
174 .layer(axum::middleware::from_fn_with_state(
182 (auth, limits.failed_auth),
183 crate::middleware::auth_layer_with_rate_limit,
184 ))
185 .layer(axum::middleware::from_fn(
186 crate::middleware::request_id_layer,
187 ))
188 .layer(axum::extract::DefaultBodyLimit::max(max_payload_bytes))
189 .layer(
193 TraceLayer::new_for_http().make_span_with(
194 DefaultMakeSpan::new()
195 .level(Level::INFO)
196 .include_headers(false),
197 ),
198 )
199 .with_state(Arc::new(state))
200}
201
202pub fn router_with_state(state: AppState, auth: AuthConfig) -> Router {
204 router_with_state_and_limit(state, auth, MAX_PAYLOAD_SIZE_BYTES)
205}
206
207pub fn router(registry: EngineRegistry) -> Router {
209 router_with_state(AppState::new(registry), AuthConfig::default())
210}
211
212pub fn router_with_auth(registry: EngineRegistry, auth: AuthConfig) -> Router {
214 router_with_state(AppState::new(registry), auth)
215}
216
217pub use router_with_state as build_router_with_state;
219
220async fn systemone(
222 State(state): State<Arc<AppState>>,
223 headers: axum::http::HeaderMap,
224 req: Result<Json<SystemRequest>, axum::extract::rejection::JsonRejection>,
225) -> Result<axum::response::Response, ApiError> {
226 let Json(req) = match req {
227 Ok(j) => j,
228 Err(rejection) => match rejection {
229 axum::extract::rejection::JsonRejection::BytesRejection(e) => {
230 return Err(ApiError::PayloadTooLarge(e.to_string()));
231 }
232 axum::extract::rejection::JsonRejection::JsonSyntaxError(e) => {
233 return Err(ApiError::BadJson(e.to_string()));
234 }
235 axum::extract::rejection::JsonRejection::JsonDataError(e) => {
236 return Err(ApiError::InvalidBody(e.to_string()));
237 }
238 other => {
239 return Err(ApiError::InvalidBody(other.to_string()));
240 }
241 },
242 };
243 if let Some(proxy) = &state.proxy {
247 if proxy.wants(&req) {
248 let contract = ResponseContract::from_request(&req)
249 .map_err(|error| ApiError::InvalidBody(error.to_string()))?;
250 let caller_key = if proxy.forwards_caller_credentials() {
251 bearer_of(&headers)
252 } else {
253 None
254 };
255 let outcome = proxy.evaluate(req, caller_key).await?;
256 contract.validate(&outcome.response).map_err(|error| {
257 ApiError::BadGateway(format!("upstream returned an invalid response: {error}"))
258 })?;
259 let mut response = (axum::http::StatusCode::OK, Json(outcome.response)).into_response();
260 let headers = response.headers_mut();
261 if let Ok(value) = axum::http::HeaderValue::from_str(outcome.source.as_str()) {
262 headers.insert(
263 axum::http::HeaderName::from_static("x-openkind-cache"),
264 value,
265 );
266 }
267 if let Some(detail) = &outcome.detail {
268 if let Ok(text) = serde_json::to_string(detail) {
269 if let Ok(value) = axum::http::HeaderValue::from_str(&text) {
270 headers.insert(
271 axum::http::HeaderName::from_static("x-openkind-cache-detail"),
272 value,
273 );
274 }
275 }
276 }
277 return Ok(response);
278 }
279 }
280 let resp = dispatch(req, &state.registry).await?;
281 Ok((axum::http::StatusCode::OK, Json(resp)).into_response())
282}
283
284fn bearer_of(headers: &axum::http::HeaderMap) -> Option<String> {
287 let value = headers.get(axum::http::header::AUTHORIZATION)?;
288 let value = value.to_str().ok()?;
289 let token = value
290 .strip_prefix("Bearer ")
291 .or_else(|| value.strip_prefix("bearer "))?;
292 let token = token.trim();
293 if token.is_empty() {
294 None
295 } else {
296 Some(token.to_owned())
297 }
298}
299
300async fn list_models(State(state): State<Arc<AppState>>) -> impl IntoResponse {
302 if let Some(proxy) = &state.proxy {
306 if let Some(models) = proxy.models().await {
307 return Json(models).into_response();
308 }
309 }
310 let models = state.registry.list_models();
311 Json(ModelsResponse::new(models)).into_response()
312}
313
314async fn health() -> impl IntoResponse {
320 static HEALTH_BODY: &str = "{\"status\":\"ok\"}";
321 (
322 [(
323 axum::http::header::CONTENT_TYPE,
324 axum::http::HeaderValue::from_static("application/json"),
325 )],
326 HEALTH_BODY,
327 )
328}
329
330static HANDLE: std::sync::OnceLock<metrics_exporter_prometheus::PrometheusHandle> =
331 std::sync::OnceLock::new();
332
333async fn prometheus_metrics() -> impl IntoResponse {
338 if let Some(h) = HANDLE.get() {
339 (
340 axum::http::StatusCode::OK,
341 [("content-type", "text/plain; version=0.0.4")],
342 h.render(),
343 )
344 } else {
345 (
346 axum::http::StatusCode::OK,
347 [("content-type", "text/plain; version=0.0.4")],
348 "# metrics recorder not installed\n".to_string(),
349 )
350 }
351}
352
353pub fn install_metrics_recorder() -> anyhow::Result<()> {
357 if HANDLE.get().is_some() {
358 return Ok(());
359 }
360 use metrics_exporter_prometheus::PrometheusBuilder;
361 let handle = PrometheusBuilder::new()
362 .install_recorder()
363 .map_err(|e| anyhow::anyhow!("install metrics recorder: {e}"))?;
364 let _ = HANDLE.set(handle);
365 Ok(())
366}
367
368#[cfg(test)]
369#[path = "http_tests.rs"]
370mod tests;