use axum::Json;
use axum::extract::State;
use axum::http::HeaderMap;
use axum::response::{IntoResponse, Response};
use serde::Deserialize;
use serde_json::json;
use crate::worker::{CaptureSpec, Cmd, Event};
use crate::{AppState, Envelope, auth, metering, worker};
#[derive(Deserialize)]
pub(crate) struct EmbeddingsReq {
pub model: String,
pub input: EmbeddingsInput,
#[serde(default)]
pub dimensions: Option<usize>,
#[serde(default)]
pub encoding_format: Option<String>,
#[serde(default)]
pub timeout_ms: Option<serde_json::Value>,
}
#[derive(Deserialize)]
#[serde(untagged)]
pub(crate) enum EmbeddingsInput {
One(String),
Many(Vec<String>),
}
#[derive(Deserialize)]
pub(crate) struct RerankReq {
pub model: String,
pub query: String,
pub documents: Vec<String>,
#[serde(default)]
pub top_n: Option<usize>,
#[serde(default)]
pub instruction: Option<String>,
#[serde(default)]
pub return_documents: Option<bool>,
#[serde(default)]
pub timeout_ms: Option<serde_json::Value>,
}
const RERANK_DEFAULT_INSTRUCTION: &str =
"Given a web search query, retrieve relevant passages that answer the query";
fn rerank_prompt(instruction: &str, query: &str, document: &str) -> String {
format!(
"<|im_start|>system\nJudge whether the Document meets the requirements based on \
the Query and the Instruct provided. Note that the answer can only be \"yes\" \
or \"no\".<|im_end|>\n<|im_start|>user\n<Instruct>: {instruction}\n<Query>: \
{query}\n<Document>: {document}<|im_end|>\n<|im_start|>assistant\n<think>\n\n\
</think>\n\n"
)
}
struct CaptureOutcome {
hidden: Option<Vec<f32>>,
logits: Vec<f32>,
n_prompt: usize,
}
#[allow(clippy::too_many_arguments)]
async fn run_capture(
st: &AppState,
headers: &HeaderMap,
parent: &Envelope,
capture_index: usize,
tenant: &auth::TenantCtx,
model: &str,
prompt_text: String,
capture: CaptureSpec,
route: &'static str,
deadline: &crate::RequestDeadline,
body_admission: Option<&crate::BodyAdmissionGuard>,
) -> Result<CaptureOutcome, Response> {
let capture_env = parent.capture_child(capture_index);
let env = &capture_env;
let cache_ns = match crate::tenant_namespace(tenant, &None::<String>) {
Ok(ns) => ns,
Err(msg) => return Err(crate::bad_request(msg, Some("cache_salt"))),
};
crate::lane_for_tenant(headers, tenant)?;
let lane = crate::lanes::Lane::Harvest;
let (tx, rx) = worker::event_channel();
let mut request = worker::Request {
model: model.to_string(),
prompt_ids: Vec::new(),
prompt_text,
chat: false,
chat_turns: Vec::new(),
tools_json: Vec::new(),
tools_struct: Vec::new(),
think: memra_tokenizer::chat::ThinkMode::Default,
reasoning_effort: None,
params: memra_engine::decode::GenParams {
max_new: 0, max_ctx: None,
eos: Vec::new(),
},
sampler_cfg: memra_engine::sampler::SamplerConfig::default(),
stop_strings: Vec::new(),
trace_id: None,
request_id: env.id.clone(),
admit_predict_logged: false,
max_prompt_tokens: None,
cache_ns,
affinity: None,
lane,
grammar: None,
prepared_constraint: None,
constraint_ready: None,
oom_retries: 0,
spec_k_replay: None,
prepared_prompt: None,
images: Vec::new(),
gemma_images: Vec::new(),
glm5_images: Vec::new(),
step_images: Vec::new(),
capture: Some(capture),
vision_memory: None,
wire_deadline: None,
ttft: None,
tx,
};
if let Err((message, param)) = crate::apply_model_request_limits(
&mut request,
st.openrouter_metadata.get(model),
st.caps.get(model),
) {
return Err(crate::bad_request(&message, Some(param)));
}
if crate::draining() {
let receipt = crate::start_request_receipt(
st,
env,
tenant,
model,
route,
lane,
false,
crate::effective_max_tokens(&request),
None,
None,
);
return Err(crate::ledger_rejected(
receipt,
crate::drain_response(),
"draining",
&env.id,
));
}
let budget = match crate::admit_tenant_budget(st, tenant, &mut request) {
Ok(budget) => budget,
Err(rejection) => {
let (response, error_code) = rejection.into_response();
let receipt = crate::start_request_receipt(
st,
env,
tenant,
model,
route,
lane,
false,
crate::effective_max_tokens(&request),
None,
None,
);
return Err(crate::ledger_rejected(
receipt, response, error_code, &env.id,
));
}
};
let receipt = crate::start_request_receipt(
st,
env,
tenant,
model,
route,
lane,
false,
crate::effective_max_tokens(&request),
budget.reserved_ctx,
budget.permit,
);
let (guard, rl) = match crate::acquire_request_slot(st, lane, tenant, env) {
Ok(slot) => slot,
Err(resp) => {
return Err(crate::ledger_rejected(
receipt,
resp,
"rate_limit_exceeded",
&env.id,
));
}
};
if let Some(admission) = body_admission {
admission.release();
}
let pending_admit = match crate::reserve_pending_admit(st, lane, &rl, *deadline) {
Ok(guard) => guard,
Err((resp, outcome)) => {
return Err(crate::ledger_unbilled(
receipt,
rl.attach(resp),
outcome,
outcome,
&env.id,
));
}
};
crate::meter_admit(env, tenant, model, lane);
if st.cmd_tx.send(Cmd::Generate(Box::new(request))).is_err() {
drop(pending_admit);
return Err(crate::ledger_rejected(
receipt,
rl.attach(crate::worker_unavailable_response()),
"worker_unavailable",
&env.id,
));
}
pending_admit.commit();
let rx = match tokio::time::timeout_at(deadline.at, crate::peek_admission(rx)).await {
Ok(Ok(rx)) => rx,
Ok(Err((resp, error_code))) => {
return Err(crate::ledger_rejected(
receipt,
rl.attach(resp),
error_code,
&env.id,
));
}
Err(_) => {
return Err(crate::ledger_unbilled(
receipt,
rl.attach(crate::deadline_exceeded_response(deadline.ms, false)),
"deadline_exceeded",
"deadline_exceeded",
&env.id,
));
}
};
collect_capture(rx, receipt, guard, rl, env, deadline).await
}
async fn collect_capture(
mut rx: worker::EventReceiver,
mut receipt: Option<Box<dyn crate::metering::Receipt>>,
_guard: crate::InflightGuard,
rl: crate::RateLimit,
env: &Envelope,
deadline: &crate::RequestDeadline,
) -> Result<CaptureOutcome, Response> {
let mut hidden: Option<Vec<f32>> = None;
let mut logits: Vec<f32> = Vec::new();
let mut got_capture = false;
let n_prompt;
loop {
let ev = match tokio::time::timeout_at(deadline.at, rx.recv()).await {
Ok(Some(ev)) => ev,
Ok(None) => {
if let Some(mut receipt) = receipt.take() {
let _ = receipt.reject(503, "worker_unavailable");
}
return Err(rl.attach(crate::worker_unavailable_response()));
}
Err(_) => {
if let Some(mut receipt) = receipt.take() {
let _ = receipt.reject(408, "deadline_exceeded");
}
return Err(rl.attach(crate::deadline_exceeded_response(deadline.ms, false)));
}
};
match ev {
Event::PromptUsage {
n_prompt: np,
n_cached,
} => {
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.record_prompt_usage(np as u64, n_cached as u64)
{
eprintln!(
"[ledger] ERROR: request {} partial prompt receipt failed: {err}",
env.id
);
let _ = receipt.reject(500, "request_ledger_unavailable");
return Err(rl.attach(crate::request_ledger_error_response()));
}
}
Event::PromptCapture {
hidden: h,
logits: l,
} => {
hidden = h;
logits = l;
got_capture = true;
}
Event::Token { .. } | Event::TokenSnapshot(_) => {}
Event::Done {
n_prompt: np,
n_cached,
elapsed_s,
..
} => {
n_prompt = np;
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.complete(
metering::UsageCounts {
prompt_tokens: np as u64,
cached_prompt_tokens: n_cached as u64,
completion_tokens: 0,
},
elapsed_s,
)
{
eprintln!(
"[ledger] ERROR: request {} completion receipt failed: {err}",
env.id
);
let _ = receipt.reject(500, "request_ledger_unavailable");
return Err(rl.attach(crate::request_ledger_error_response()));
}
if !got_capture {
return Err(rl.attach(crate::engine_error_response(
&worker::EngineError::engine(
"capture request finished without a PromptCapture event",
),
)));
}
return Ok(CaptureOutcome {
hidden,
logits,
n_prompt,
});
}
Event::Error(err) => {
if let Some(receipt) = receipt.as_mut() {
let _ = receipt.reject(
crate::class_http(err.class).0.as_u16(),
crate::engine_error_code(err.class),
);
}
return Err(rl.attach(crate::engine_error_response(&err)));
}
}
}
}
fn l2_normalize(v: &mut [f32]) {
let norm = v
.iter()
.map(|x| (*x as f64) * (*x as f64))
.sum::<f64>()
.sqrt();
if norm > 0.0 {
for x in v.iter_mut() {
*x = (*x as f64 / norm) as f32;
}
}
}
pub(crate) async fn embeddings_admitted(
state: State<AppState>,
headers: HeaderMap,
crate::AdmittedJson(req, admission): crate::AdmittedJson<EmbeddingsReq>,
) -> Response {
embeddings_with_admission(state, headers, Json(req), Some(admission)).await
}
async fn embeddings_with_admission(
State(st): State<AppState>,
headers: HeaderMap,
Json(mut req): Json<EmbeddingsReq>,
body_admission: Option<crate::BodyAdmissionLease>,
) -> Response {
let env = Envelope::new(false);
match crate::canonical_model_id(&st.models, &req.model) {
Some(canonical) => req.model = canonical,
None => {
return crate::with_request_id(
&env.id,
crate::model_not_found_response(&st.models, &req.model),
);
}
}
if req.encoding_format.as_deref().is_some_and(|f| f != "float") {
return crate::with_request_id(
&env.id,
crate::bad_request(
"only encoding_format=\"float\" is served",
Some("encoding_format"),
),
);
}
let inputs: Vec<String> = match req.input {
EmbeddingsInput::One(s) => vec![s],
EmbeddingsInput::Many(v) => v,
};
const MAX_INPUTS: usize = 32;
if inputs.is_empty() || inputs.len() > MAX_INPUTS {
return crate::with_request_id(
&env.id,
crate::bad_request(
&format!("input must carry 1..={MAX_INPUTS} strings"),
Some("input"),
),
);
}
if req.dimensions.is_some_and(|d| d == 0) {
return crate::with_request_id(
&env.id,
crate::bad_request("dimensions must be >= 1", Some("dimensions")),
);
}
let tenant = match crate::authenticate(&st.api_auth, &headers) {
Ok(t) => t,
Err(resp) => return crate::with_request_id(&env.id, resp),
};
let deadline = match crate::parse_timeout_ms(req.timeout_ms.as_ref()) {
Ok(ms) => crate::RequestDeadline::starting_now(ms),
Err(msg) => {
return crate::with_request_id(&env.id, crate::bad_request(&msg, Some("timeout_ms")));
}
};
let mut data = Vec::with_capacity(inputs.len());
let mut prompt_tokens = 0usize;
for (index, input) in inputs.into_iter().enumerate() {
let outcome = match run_capture(
&st,
&headers,
&env,
index,
&tenant,
&req.model,
input,
CaptureSpec {
hidden: true,
logit_pieces: Vec::new(),
},
"/v1/embeddings",
&deadline,
body_admission
.as_ref()
.and_then(|admission| admission.guard()),
)
.await
{
Ok(o) => o,
Err(resp) => return crate::with_request_id(&env.id, resp),
};
let Some(mut vector) = outcome.hidden else {
return crate::with_request_id(
&env.id,
crate::bad_request(
"this model cannot serve embeddings (no prime-path hidden state)",
Some("model"),
),
);
};
if let Some(d) = req.dimensions {
if d > vector.len() {
return crate::with_request_id(
&env.id,
crate::bad_request(
&format!("dimensions exceeds the model width {}", vector.len()),
Some("dimensions"),
),
);
}
vector.truncate(d);
}
l2_normalize(&mut vector);
prompt_tokens += outcome.n_prompt;
data.push(json!({
"object": "embedding",
"index": index,
"embedding": vector,
}));
}
let body = json!({
"object": "list",
"data": data,
"model": req.model,
"usage": { "prompt_tokens": prompt_tokens, "total_tokens": prompt_tokens },
});
crate::with_request_id(&env.id, Json(body).into_response())
}
pub(crate) async fn rerank_admitted(
state: State<AppState>,
headers: HeaderMap,
crate::AdmittedJson(req, admission): crate::AdmittedJson<RerankReq>,
) -> Response {
rerank_with_admission(state, headers, Json(req), Some(admission)).await
}
async fn rerank_with_admission(
State(st): State<AppState>,
headers: HeaderMap,
Json(mut req): Json<RerankReq>,
body_admission: Option<crate::BodyAdmissionLease>,
) -> Response {
let env = Envelope::new(false);
match crate::canonical_model_id(&st.models, &req.model) {
Some(canonical) => req.model = canonical,
None => {
return crate::with_request_id(
&env.id,
crate::model_not_found_response(&st.models, &req.model),
);
}
}
const MAX_DOCS: usize = 64;
if req.documents.is_empty() || req.documents.len() > MAX_DOCS {
return crate::with_request_id(
&env.id,
crate::bad_request(
&format!("documents must carry 1..={MAX_DOCS} strings"),
Some("documents"),
),
);
}
let tenant = match crate::authenticate(&st.api_auth, &headers) {
Ok(t) => t,
Err(resp) => return crate::with_request_id(&env.id, resp),
};
let deadline = match crate::parse_timeout_ms(req.timeout_ms.as_ref()) {
Ok(ms) => crate::RequestDeadline::starting_now(ms),
Err(msg) => {
return crate::with_request_id(&env.id, crate::bad_request(&msg, Some("timeout_ms")));
}
};
let instruction = req
.instruction
.as_deref()
.unwrap_or(RERANK_DEFAULT_INSTRUCTION)
.to_string();
let mut scored: Vec<(usize, f64)> = Vec::with_capacity(req.documents.len());
let mut total_tokens = 0usize;
for (index, document) in req.documents.iter().enumerate() {
let outcome = match run_capture(
&st,
&headers,
&env,
index,
&tenant,
&req.model,
rerank_prompt(&instruction, &req.query, document),
CaptureSpec {
hidden: false,
logit_pieces: vec!["yes".to_string(), "no".to_string()],
},
"/v1/rerank",
&deadline,
body_admission
.as_ref()
.and_then(|admission| admission.guard()),
)
.await
{
Ok(o) => o,
Err(resp) => return crate::with_request_id(&env.id, resp),
};
let (yes, no) = match outcome.logits.as_slice() {
[y, n] if *y > f32::MIN && *n > f32::MIN => (*y as f64, *n as f64),
_ => {
return crate::with_request_id(
&env.id,
crate::bad_request(
"this model cannot serve rerank (\"yes\"/\"no\" are not single \
vocabulary tokens)",
Some("model"),
),
);
}
};
let m = yes.max(no);
let score = ((yes - m).exp()) / ((yes - m).exp() + (no - m).exp());
scored.push((index, score));
total_tokens += outcome.n_prompt;
}
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let top_n = req.top_n.unwrap_or(scored.len()).min(scored.len());
let return_documents = req.return_documents.unwrap_or(false);
let results: Vec<serde_json::Value> = scored[..top_n]
.iter()
.map(|(index, score)| {
let mut row = json!({ "index": index, "relevance_score": score });
if return_documents {
row["document"] = json!({ "text": req.documents[*index] });
}
row
})
.collect();
let body = json!({
"model": req.model,
"results": results,
"usage": { "total_tokens": total_tokens },
});
crate::with_request_id(&env.id, Json(body).into_response())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rerank_prompt_renders_the_vendor_judge_format() {
let p = rerank_prompt("inst", "q", "d");
assert!(p.starts_with("<|im_start|>system\nJudge whether the Document"));
assert!(p.contains("<Instruct>: inst\n<Query>: q\n<Document>: d<|im_end|>"));
assert!(p.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"));
}
#[test]
fn l2_normalize_produces_a_unit_vector_and_survives_zero() {
let mut v = vec![3.0f32, 4.0];
l2_normalize(&mut v);
assert!((v[0] - 0.6).abs() < 1e-6 && (v[1] - 0.8).abs() < 1e-6);
let mut z = vec![0.0f32, 0.0];
l2_normalize(&mut z);
assert_eq!(z, vec![0.0, 0.0]);
}
#[test]
fn mrl_truncation_renormalizes_the_prefix() {
let mut v = vec![1.0f32, 1.0, 1.0, 1.0];
v.truncate(2);
l2_normalize(&mut v);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-6);
}
#[test]
fn rerank_score_is_p_yes_over_the_pair() {
let (yes, no) = (2.0f64, 0.0f64);
let m = yes.max(no);
let score = ((yes - m).exp()) / ((yes - m).exp() + (no - m).exp());
assert!((score - 1.0 / (1.0 + (-2.0f64).exp())).abs() < 1e-12);
}
}