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::RateLimiter::new(crate::middleware::RateLimitConfig::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(state, auth, max_payload_bytes, rate_limiter, false, false)
62}
63
64pub fn router_daemon(
68 state: AppState,
69 auth: AuthConfig,
70 max_payload_bytes: usize,
71 rate_limiter: crate::middleware::RateLimiter,
72 playground: bool,
73) -> Router {
74 router_full(
75 state,
76 auth,
77 max_payload_bytes,
78 rate_limiter,
79 playground,
80 false,
81 )
82}
83
84pub fn router_daemon_with_arrow(
87 state: AppState,
88 auth: AuthConfig,
89 max_payload_bytes: usize,
90 rate_limiter: crate::middleware::RateLimiter,
91 playground: bool,
92 arrow: bool,
93) -> Router {
94 router_full(
95 state,
96 auth,
97 max_payload_bytes,
98 rate_limiter,
99 playground,
100 arrow,
101 )
102}
103
104fn router_full(
105 state: AppState,
106 auth: AuthConfig,
107 max_payload_bytes: usize,
108 rate_limiter: crate::middleware::RateLimiter,
109 playground: bool,
110 arrow: bool,
111) -> Router {
112 let routes = Router::new()
113 .route("/v1/systemone", post(systemone))
115 .route("/v1/system_one", post(systemone))
117 .route("/v1/models", get(list_models))
119 .route("/health", get(health))
121 .route("/metrics", get(prometheus_metrics));
123 let routes = if playground {
127 routes
128 .route("/playground", get(crate::playground::playground_page))
129 .route(
130 "/playground/api/models",
131 get(crate::playground::list_models).post(crate::playground::change_model),
132 )
133 } else {
134 routes
135 };
136 let routes = if arrow {
140 routes.route("/v1/arrow", post(crate::arrow::arrow_batch))
141 } else {
142 routes
143 };
144 let routes = if rate_limiter.is_enabled() {
147 routes.layer(axum::middleware::from_fn_with_state(
148 rate_limiter,
149 crate::middleware::rate_limit_layer,
150 ))
151 } else {
152 routes
153 };
154 routes
155 .layer(axum::middleware::from_fn_with_state(
163 auth,
164 crate::middleware::auth_layer,
165 ))
166 .layer(axum::middleware::from_fn(
167 crate::middleware::request_id_layer,
168 ))
169 .layer(axum::extract::DefaultBodyLimit::max(max_payload_bytes))
170 .layer(
174 TraceLayer::new_for_http().make_span_with(
175 DefaultMakeSpan::new()
176 .level(Level::INFO)
177 .include_headers(false),
178 ),
179 )
180 .with_state(Arc::new(state))
181}
182
183pub fn router_with_state(state: AppState, auth: AuthConfig) -> Router {
185 router_with_state_and_limit(state, auth, MAX_PAYLOAD_SIZE_BYTES)
186}
187
188pub fn router(registry: EngineRegistry) -> Router {
190 router_with_state(AppState::new(registry), AuthConfig::default())
191}
192
193pub fn router_with_auth(registry: EngineRegistry, auth: AuthConfig) -> Router {
195 router_with_state(AppState::new(registry), auth)
196}
197
198pub use router_with_state as build_router_with_state;
200
201async fn systemone(
203 State(state): State<Arc<AppState>>,
204 headers: axum::http::HeaderMap,
205 req: Result<Json<SystemRequest>, axum::extract::rejection::JsonRejection>,
206) -> Result<axum::response::Response, ApiError> {
207 let Json(req) = match req {
208 Ok(j) => j,
209 Err(rejection) => match rejection {
210 axum::extract::rejection::JsonRejection::BytesRejection(e) => {
211 return Err(ApiError::PayloadTooLarge(e.to_string()));
212 }
213 axum::extract::rejection::JsonRejection::JsonSyntaxError(e) => {
214 return Err(ApiError::BadJson(e.to_string()));
215 }
216 axum::extract::rejection::JsonRejection::JsonDataError(e) => {
217 return Err(ApiError::InvalidBody(e.to_string()));
218 }
219 other => {
220 return Err(ApiError::InvalidBody(other.to_string()));
221 }
222 },
223 };
224 if let Some(proxy) = &state.proxy {
228 if proxy.wants(&req) {
229 let contract = ResponseContract::from_request(&req)
230 .map_err(|error| ApiError::InvalidBody(error.to_string()))?;
231 let caller_key = if proxy.forwards_caller_credentials() {
232 bearer_of(&headers)
233 } else {
234 None
235 };
236 let outcome = proxy.evaluate(req, caller_key).await?;
237 contract.validate(&outcome.response).map_err(|error| {
238 ApiError::BadGateway(format!("upstream returned an invalid response: {error}"))
239 })?;
240 let mut response = (axum::http::StatusCode::OK, Json(outcome.response)).into_response();
241 let headers = response.headers_mut();
242 if let Ok(value) = axum::http::HeaderValue::from_str(outcome.source.as_str()) {
243 headers.insert(
244 axum::http::HeaderName::from_static("x-openkind-cache"),
245 value,
246 );
247 }
248 if let Some(detail) = &outcome.detail {
249 if let Ok(text) = serde_json::to_string(detail) {
250 if let Ok(value) = axum::http::HeaderValue::from_str(&text) {
251 headers.insert(
252 axum::http::HeaderName::from_static("x-openkind-cache-detail"),
253 value,
254 );
255 }
256 }
257 }
258 return Ok(response);
259 }
260 }
261 let resp = dispatch(req, &state.registry).await?;
262 Ok((axum::http::StatusCode::OK, Json(resp)).into_response())
263}
264
265fn bearer_of(headers: &axum::http::HeaderMap) -> Option<String> {
268 let value = headers.get(axum::http::header::AUTHORIZATION)?;
269 let value = value.to_str().ok()?;
270 let token = value
271 .strip_prefix("Bearer ")
272 .or_else(|| value.strip_prefix("bearer "))?;
273 let token = token.trim();
274 if token.is_empty() {
275 None
276 } else {
277 Some(token.to_owned())
278 }
279}
280
281async fn list_models(State(state): State<Arc<AppState>>) -> impl IntoResponse {
283 if let Some(proxy) = &state.proxy {
287 if let Some(models) = proxy.models().await {
288 return Json(models).into_response();
289 }
290 }
291 let models = state.registry.list_models();
292 Json(ModelsResponse::new(models)).into_response()
293}
294
295async fn health() -> impl IntoResponse {
301 static HEALTH_BODY: &str = "{\"status\":\"ok\"}";
302 (
303 [(
304 axum::http::header::CONTENT_TYPE,
305 axum::http::HeaderValue::from_static("application/json"),
306 )],
307 HEALTH_BODY,
308 )
309}
310
311static HANDLE: std::sync::OnceLock<metrics_exporter_prometheus::PrometheusHandle> =
312 std::sync::OnceLock::new();
313
314async fn prometheus_metrics() -> impl IntoResponse {
319 if let Some(h) = HANDLE.get() {
320 (
321 axum::http::StatusCode::OK,
322 [("content-type", "text/plain; version=0.0.4")],
323 h.render(),
324 )
325 } else {
326 (
327 axum::http::StatusCode::OK,
328 [("content-type", "text/plain; version=0.0.4")],
329 "# metrics recorder not installed\n".to_string(),
330 )
331 }
332}
333
334pub fn install_metrics_recorder() -> anyhow::Result<()> {
338 if HANDLE.get().is_some() {
339 return Ok(());
340 }
341 use metrics_exporter_prometheus::PrometheusBuilder;
342 let handle = PrometheusBuilder::new()
343 .install_recorder()
344 .map_err(|e| anyhow::anyhow!("install metrics recorder: {e}"))?;
345 let _ = HANDLE.set(handle);
346 Ok(())
347}
348
349#[cfg(test)]
350#[path = "http_tests.rs"]
351mod tests;