bge_m3_embedding_server/bootstrap/
router.rs1use std::sync::Arc;
19
20use axum::extract::DefaultBodyLimit;
21use axum::http::HeaderValue;
22use axum::{Router, routing::get, routing::post};
23use tower_http::request_id::{
24 MakeRequestId, PropagateRequestIdLayer, RequestId, SetRequestIdLayer,
25};
26use tower_http::trace::{DefaultOnFailure, DefaultOnResponse, MakeSpan, TraceLayer};
27use tracing::Level;
28
29use crate::handler;
30use crate::state::AppState;
31
32#[derive(Clone, Default)]
35pub(super) struct UuidRequestId;
36
37impl MakeRequestId for UuidRequestId {
38 fn make_request_id<B>(&mut self, _request: &axum::http::Request<B>) -> Option<RequestId> {
39 let id = uuid::Uuid::new_v4().to_string();
40 HeaderValue::from_str(&id).ok().map(RequestId::new)
41 }
42}
43
44#[derive(Clone)]
50pub(super) struct RouteAwareSpan;
51
52impl<B> MakeSpan<B> for RouteAwareSpan {
53 fn make_span(&mut self, request: &axum::http::Request<B>) -> tracing::Span {
54 let path = request.uri().path();
55 let is_noisy = matches!(path, "/health" | "/health/deep" | "/v1/models");
56 let method = request.method().as_str();
57 if is_noisy {
58 tracing::debug_span!(
59 "http_request",
60 method = method,
61 uri = %request.uri(),
62 version = ?request.version(),
63 )
64 } else {
65 tracing::info_span!(
66 "http_request",
67 method = method,
68 uri = %request.uri(),
69 version = ?request.version(),
70 )
71 }
72 }
73}
74
75pub fn build_router(state: Arc<AppState>, max_body_bytes: usize) -> Router {
79 Router::new()
80 .route("/v1/embeddings", post(handler::dense_embeddings))
81 .route("/v1/sparse-embeddings", post(handler::sparse_embeddings))
82 .route("/v1/embeddings:both", post(handler::both_embeddings))
90 .route("/v1/embeddings%3Aboth", post(handler::both_embeddings))
91 .route("/v1/embeddings%3aboth", post(handler::both_embeddings))
92 .route("/v1/models", get(handler::models))
93 .route("/health", get(handler::health))
94 .route("/health/deep", get(handler::health_deep))
95 .layer(DefaultBodyLimit::max(max_body_bytes))
96 .layer(PropagateRequestIdLayer::x_request_id())
97 .layer(
98 TraceLayer::new_for_http()
99 .make_span_with(RouteAwareSpan)
100 .on_response(
101 DefaultOnResponse::new()
102 .level(Level::INFO)
103 .latency_unit(tower_http::LatencyUnit::Millis),
104 )
105 .on_failure(DefaultOnFailure::new().level(Level::ERROR)),
106 )
107 .layer(SetRequestIdLayer::x_request_id(UuidRequestId))
108 .with_state(state)
109}