Skip to main content

bge_m3_embedding_server/handler/
dense.rs

1// Copyright (c) 2026 J. Patrick Fulton
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! `POST /v1/embeddings` handler — OpenAI-compatible dense embeddings.
16
17use 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/// Handles `POST /v1/embeddings` — returns dense (float32) embeddings.
28///
29/// # Errors
30///
31/// - [`AppError::ServiceUnavailable`] if the model is not ready or no workers are live.
32/// - [`AppError::InvalidRequest`] if the batch is empty, exceeds `max_batch`, or any
33///   text exceeds the per-string character limit.
34/// - [`AppError::Internal`] if the embedding pool returns an inference error.
35///
36/// # Panics
37///
38/// Panics if the request semaphore has been closed — should not occur in normal operation.
39#[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    // Acquire a concurrency permit before dispatching to the worker pool.
71    // This is released on drop when the handler returns (success or error).
72    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    // x_headers (normalized: hyphens → underscores) are emitted at event level so
90    // they appear under $.fields.x_headers in JSON logs and are accessible to
91    // downstream log processors. Each caller-supplied X-* header is included
92    // generically; no header name is special-cased here.
93    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}