bge_m3_embedding_server/handler/
dense.rs1use std::sync::Arc;
18use std::time::Instant;
19
20use axum::{Json, extract::State, http::HeaderMap};
21
22use super::common::{check_ready, collect_x_headers, validate_input};
23use crate::error::AppError;
24use crate::models::{DenseEmbeddingData, DenseRequest, DenseResponse, Usage};
25use crate::state::AppState;
26
27#[tracing::instrument(
40 skip(state, req, headers),
41 fields(
42 batch_size,
43 prompt_tokens,
44 chunks,
45 max_chunk_seq,
46 tokenize_ms,
47 inference_ms,
48 queue_wait_ms,
49 total_ms,
50 )
51)]
52pub async fn dense_embeddings(
53 State(state): State<Arc<AppState>>,
54 headers: HeaderMap,
55 Json(req): Json<DenseRequest>,
56) -> Result<Json<DenseResponse>, AppError> {
57 check_ready(&state)?;
58 let x_headers = collect_x_headers(&headers);
59 let texts = req.input.0;
60 drop(req.model);
61 validate_input(&texts, state.max_batch)?;
62 let batch_size = texts.len();
63 tracing::Span::current().record("batch_size", batch_size);
64
65 let prompt_tokens: usize = texts.iter().map(|t| t.chars().count() / 4 + 1).sum();
66 tracing::Span::current().record("prompt_tokens", prompt_tokens);
67
68 let t0 = Instant::now();
69
70 let _permit = Arc::clone(&state.request_permits)
73 .acquire_owned()
74 .await
75 .expect("request semaphore is never closed");
76
77 let queue_wait_ms = u64::try_from(t0.elapsed().as_millis()).unwrap_or(u64::MAX);
78
79 let (embeddings, embed_stats) = state.pool.dense(texts).await?;
80
81 let total_ms = u64::try_from(t0.elapsed().as_millis()).unwrap_or(u64::MAX);
82 tracing::Span::current()
83 .record("chunks", embed_stats.chunks)
84 .record("max_chunk_seq", embed_stats.max_chunk_seq)
85 .record("tokenize_ms", embed_stats.tokenize_ms)
86 .record("inference_ms", embed_stats.inference_ms)
87 .record("queue_wait_ms", queue_wait_ms)
88 .record("total_ms", total_ms);
89 let x_headers_val =
94 (!x_headers.is_empty()).then(|| serde_json::to_string(&x_headers).unwrap_or_default());
95 tracing::info!(
96 route = "dense",
97 batch_size,
98 prompt_tokens,
99 chunks = embed_stats.chunks,
100 max_chunk_seq = embed_stats.max_chunk_seq,
101 total_token_positions = embed_stats.total_token_positions,
102 seq_len_min = embed_stats.seq_len_min,
103 seq_len_max = embed_stats.seq_len_max,
104 seq_len_mean = embed_stats.seq_len_mean,
105 seq_len_p95 = embed_stats.seq_len_p95,
106 tokenize_ms = embed_stats.tokenize_ms,
107 inference_ms = embed_stats.inference_ms,
108 queue_wait_ms,
109 total_ms,
110 x_headers = x_headers_val,
111 "embedding request complete"
112 );
113
114 let data = embeddings
115 .into_iter()
116 .enumerate()
117 .map(|(index, embedding)| DenseEmbeddingData {
118 object: "embedding",
119 index,
120 embedding,
121 })
122 .collect();
123
124 Ok(Json(DenseResponse {
125 object: "list",
126 model: "bge-m3",
127 data,
128 usage: Usage {
129 prompt_tokens,
130 total_tokens: prompt_tokens,
131 },
132 }))
133}