use std::collections::HashMap;
use std::convert::Infallible;
use std::sync::Arc;
use axum::response::{IntoResponse, Sse};
use axum::{Json, extract::State, response::sse::Event};
use serde::{Deserialize, Serialize};
use tokio_stream::wrappers::ReceiverStream;
use super::AppState;
use super::CancelOnDrop;
use super::infer::{run_embeddings, run_text_inference_with_config, run_text_inference_with_logprobs};
#[derive(Serialize)]
pub(super) struct ModelsList {
object: &'static str,
data: Vec<ModelEntry>,
}
#[derive(Serialize)]
pub(super) struct ModelEntry {
id: String,
object: &'static str,
created: u64,
owned_by: &'static str,
}
pub(super) async fn list_models(State(state): State<Arc<AppState>>) -> Json<ModelsList> {
Json(ModelsList {
object: "list",
data: vec![ModelEntry {
id: state.name.clone(),
object: "model",
created: 0,
owned_by: "modelc",
}],
})
}
pub(super) async fn retrieve_model(
State(state): State<Arc<AppState>>,
axum::extract::Path(id): axum::extract::Path<String>,
) -> axum::response::Response {
if id != state.name {
return (
axum::http::StatusCode::NOT_FOUND,
Json(serde_json::json!({
"error": {
"message": format!("The model '{}' does not exist", id),
"type": "invalid_request_error",
"code": "model_not_found"
}
})),
)
.into_response();
}
Json(ModelEntry {
id: state.name.clone(),
object: "model",
created: 0,
owned_by: "modelc",
})
.into_response()
}
#[derive(Deserialize)]
pub(super) struct ChatCompletionRequest {
#[serde(default)]
model: String,
messages: Vec<ChatMessage>,
#[serde(default)]
stream: bool,
#[serde(default)]
response_format: Option<ResponseFormat>,
#[serde(default)]
tools: Vec<Tool>,
#[serde(default)]
tool_choice: Option<serde_json::Value>,
#[serde(default)]
max_tokens: Option<usize>,
#[serde(default)]
max_completion_tokens: Option<usize>,
#[serde(default)]
temperature: Option<f32>,
#[serde(default)]
top_p: Option<f32>,
#[serde(default)]
min_p: Option<f32>,
#[serde(default)]
logprobs: Option<bool>,
#[serde(default)]
top_logprobs: Option<u8>,
#[serde(default)]
grammar: Option<String>,
#[serde(default)]
json_schema: Option<serde_json::Value>,
#[serde(default)]
stop: Vec<String>,
#[serde(default)]
seed: Option<u64>,
#[serde(default)]
repetition_penalty: Option<f32>,
#[serde(default)]
presence_penalty: Option<f32>,
#[serde(default)]
frequency_penalty: Option<f32>,
#[serde(default)]
logit_bias: Option<HashMap<u32, f32>>,
#[serde(default)]
n: Option<u32>,
#[serde(default)]
stream_options: Option<StreamOptions>,
#[serde(default)]
#[allow(dead_code)]
user: Option<String>,
}
#[derive(Deserialize)]
pub(super) struct StreamOptions {
#[serde(default)]
include_usage: bool,
}
#[derive(Deserialize)]
pub(super) struct ResponseFormat {
#[serde(rename = "type")]
format_type: String,
#[serde(default)]
json_schema: Option<JsonSchemaDef>,
}
#[derive(Deserialize)]
pub(super) struct JsonSchemaDef {
#[serde(default)]
#[allow(dead_code)] name: Option<String>,
#[serde(default)]
schema: Option<serde_json::Value>,
#[serde(default)]
#[allow(dead_code)] strict: Option<bool>,
}
#[derive(Deserialize, Clone)]
pub(super) struct Tool {
function: ToolFunction,
}
#[derive(Deserialize, Clone)]
pub(super) struct ToolFunction {
name: String,
#[serde(default)]
description: String,
}
#[derive(Deserialize, Serialize, Clone)]
pub(super) struct ChatMessage {
role: String,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<ToolCall>>,
}
#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)]
pub(super) struct ToolCall {
id: String,
#[serde(rename = "type")]
call_type: String,
function: ToolCallFunction,
}
#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)]
pub(super) struct ToolCallFunction {
name: String,
arguments: String,
}
#[derive(Serialize)]
pub(super) struct ChatCompletionResponse {
id: String,
object: &'static str,
created: u64,
model: String,
choices: Vec<Choice>,
usage: Usage,
}
#[derive(Serialize)]
pub(super) struct Choice {
index: usize,
message: ChatMessage,
finish_reason: &'static str,
logprobs: Option<Logprobs>,
}
#[derive(Serialize)]
pub(super) struct Logprobs {
content: Vec<ContentLogprob>,
}
#[derive(Serialize)]
pub(super) struct ContentLogprob {
token: String,
logprob: f64,
bytes: Vec<u8>,
top_logprobs: Vec<TopLogprob>,
}
#[derive(Serialize)]
pub(super) struct TopLogprob {
token: String,
logprob: f64,
bytes: Vec<u8>,
}
#[derive(Serialize)]
pub(super) struct Usage {
prompt_tokens: usize,
completion_tokens: usize,
total_tokens: usize,
}
#[derive(Serialize)]
pub(super) struct ChatCompletionChunk {
id: String,
object: &'static str,
created: u64,
model: String,
choices: Vec<ChunkChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
usage: Option<Usage>,
}
#[derive(Serialize)]
pub(super) struct ChunkChoice {
index: usize,
delta: ChunkDelta,
finish_reason: Option<&'static str>,
}
#[derive(Serialize, Default)]
pub(super) struct ChunkDelta {
#[serde(skip_serializing_if = "Option::is_none")]
role: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<String>,
}
pub(super) async fn chat_completion(
State(state): State<Arc<AppState>>,
Json(req): Json<ChatCompletionRequest>,
) -> axum::response::Response {
let _guard = super::metrics::ActiveRequestGuard::new(&state.metrics);
let _timer = super::metrics::InferenceTimer::new(&state.metrics);
let is_json_mode = req
.response_format
.as_ref()
.is_some_and(|rf| rf.format_type == "json_object");
let is_json_schema_mode = req
.response_format
.as_ref()
.is_some_and(|rf| rf.format_type == "json_schema");
let response_schema = req
.response_format
.as_ref()
.and_then(|rf| rf.json_schema.as_ref())
.and_then(|js| js.schema.clone());
let has_tools = !req.tools.is_empty();
let tool_choice_none = req
.tool_choice
.as_ref()
.and_then(|v| v.as_str())
.is_some_and(|s| s == "none");
let tool_choice_required = req
.tool_choice
.as_ref()
.and_then(|v| v.as_str())
.is_some_and(|s| s == "required");
let tool_choice_specific = req
.tool_choice
.as_ref()
.and_then(|v| v.get("function"))
.and_then(|f| f.get("name"))
.and_then(|n| n.as_str())
.map(|s| s.to_string());
let effective_has_tools = has_tools && !tool_choice_none;
let stream = req.stream;
let mut messages: Vec<crate::chat_template::ChatMessage> = req
.messages
.iter()
.map(|m| crate::chat_template::ChatMessage {
role: m.role.clone(),
content: m.content.clone().unwrap_or_default(),
})
.collect();
if is_json_mode {
messages.push(crate::chat_template::ChatMessage {
role: "system".to_string(),
content: "Respond with valid JSON only. Do not include any explanatory text before or after the JSON.".to_string(),
});
}
if is_json_schema_mode {
let schema_str = response_schema
.as_ref()
.map(|s| serde_json::to_string_pretty(s).unwrap_or_default())
.unwrap_or_default();
messages.push(crate::chat_template::ChatMessage {
role: "system".to_string(),
content: format!(
"Respond with valid JSON that conforms to the following JSON Schema. Do not include any explanatory text before or after the JSON.\n\n{}",
schema_str
),
});
}
if effective_has_tools {
let tool_desc = build_tool_description(&req.tools);
let mut content = tool_desc;
if tool_choice_required {
content.push_str("\nYou must call one of the available tools. Do not respond with text only.");
}
if let Some(ref name) = tool_choice_specific {
content.push_str(&format!(
"\nYou must call the tool named \"{}\". Do not call any other tool.",
name
));
}
messages.push(crate::chat_template::ChatMessage {
role: "system".to_string(),
content,
});
}
let prompt =
crate::chat_template::apply_chat_template(state.chat_template.as_deref(), &messages);
let constraint = req
.grammar
.and_then(|pat| {
crate::constraint::RegexConstraint::new(&pat).map(|c| {
std::sync::Arc::new(c) as std::sync::Arc<dyn crate::constraint::Constraint>
})
})
.or_else(|| state.generation.constraint.clone());
let gen_cfg = crate::generate::GenerationConfig {
max_tokens: req.max_completion_tokens.or(req.max_tokens).unwrap_or(state.generation.max_tokens),
temperature: req.temperature.unwrap_or(state.generation.temperature),
top_p: req.top_p.unwrap_or(state.generation.top_p),
min_p: req.min_p.unwrap_or(state.generation.min_p),
gamma: state.generation.gamma,
use_int8_kv: state.generation.use_int8_kv,
use_mixed_kv: state.generation.use_mixed_kv,
constraint,
max_context: state.generation.max_context,
anchor_tokens: state.generation.anchor_tokens,
stop: if req.stop.is_empty() {
state.generation.stop.clone()
} else {
req.stop
},
seed: req.seed.or(state.generation.seed),
repetition_penalty: req
.repetition_penalty
.unwrap_or(state.generation.repetition_penalty),
presence_penalty: req
.presence_penalty
.unwrap_or(state.generation.presence_penalty),
frequency_penalty: req
.frequency_penalty
.unwrap_or(state.generation.frequency_penalty),
logit_bias: req
.logit_bias
.unwrap_or_else(|| state.generation.logit_bias.clone()),
cancel: None,
};
let model_name = if req.model.is_empty() {
state.name.clone()
} else {
req.model
};
let id = format!("chatcmpl-{}-0", state.name);
let include_usage = req
.stream_options
.as_ref()
.is_some_and(|so| so.include_usage);
if stream {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Event, Infallible>>(4);
let state_clone = Arc::clone(&state);
let cancel = Arc::new(std::sync::atomic::AtomicBool::new(false));
let mut gen_cfg = gen_cfg;
gen_cfg.cancel = Some(cancel.clone());
tokio::spawn(async move {
let token_ids =
super::infer::run_text_inference_token_ids(&state_clone, &prompt, &gen_cfg);
let tokenizer = crate::tokenizer::BpeTokenizer::byte_fallback();
let prompt_ids = tokenizer.encode(&prompt);
let mut prev_text = String::new();
let mut first = true;
for (idx, &_token_id) in token_ids.iter().enumerate() {
let cumulative = [prompt_ids.as_slice(), &token_ids[..=idx]].concat();
let text = tokenizer.decode(&cumulative);
if let Some(delta) = text.strip_prefix(&prev_text) {
if !delta.is_empty() {
let chunk = ChatCompletionChunk {
id: id.clone(),
object: "chat.completion.chunk",
created: 0,
model: model_name.clone(),
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta {
role: if first {
Some("assistant".to_string())
} else {
None
},
content: Some(delta.to_string()),
},
finish_reason: None,
}],
usage: None,
};
if tx
.send(Ok(
Event::default().data(serde_json::to_string(&chunk).unwrap())
))
.await
.is_err()
{
break;
}
first = false;
}
prev_text = text;
}
}
let final_chunk = ChatCompletionChunk {
id: id.clone(),
object: "chat.completion.chunk",
created: 0,
model: model_name.clone(),
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta::default(),
finish_reason: Some("stop"),
}],
usage: None,
};
let _ = tx
.send(Ok(
Event::default().data(serde_json::to_string(&final_chunk).unwrap())
))
.await;
if include_usage {
let usage_chunk = ChatCompletionChunk {
id: id.clone(),
object: "chat.completion.chunk",
created: 0,
model: model_name.clone(),
choices: vec![],
usage: Some(Usage {
prompt_tokens: prompt_ids.len(),
completion_tokens: token_ids.len(),
total_tokens: prompt_ids.len() + token_ids.len(),
}),
};
let _ = tx
.send(Ok(
Event::default().data(serde_json::to_string(&usage_chunk).unwrap())
))
.await;
}
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
});
return Sse::new(CancelOnDrop::new(ReceiverStream::new(rx), cancel)).into_response();
}
let want_logprobs = req.logprobs.unwrap_or(false);
let top_n = req.top_logprobs.unwrap_or(0).min(20) as usize;
let n = req.n.unwrap_or(1).clamp(1, 20) as usize;
let tokenizer = crate::tokenizer::BpeTokenizer::byte_fallback();
let prompt_tokens = tokenizer.encode(&prompt).len();
let mut choices: Vec<Choice> = Vec::with_capacity(n);
let mut total_completion_tokens = 0usize;
for i in 0..n {
let (raw_output, logprobs_field): (String, Option<Logprobs>) = if want_logprobs {
match run_text_inference_with_logprobs(&state, &prompt, &gen_cfg, top_n) {
Some((ids, lps)) => {
let text = tokenizer.decode(&ids);
let content = build_logprobs_content(&tokenizer, &lps);
(text, Some(Logprobs { content }))
}
None => {
let text = run_text_inference_with_config(&state, &prompt, &gen_cfg);
(text, Some(Logprobs { content: Vec::new() }))
}
}
} else {
let text = if let Some(ref schema) = req.json_schema {
crate::json_schema::generate_with_schema(
|cfg| run_text_inference_with_config(&state, &prompt, cfg),
schema,
&gen_cfg,
3,
)
} else if is_json_schema_mode {
if let Some(ref schema) = response_schema {
crate::json_schema::generate_with_schema(
|cfg| run_text_inference_with_config(&state, &prompt, cfg),
schema,
&gen_cfg,
3,
)
} else {
run_text_inference_with_config(&state, &prompt, &gen_cfg)
}
} else {
run_text_inference_with_config(&state, &prompt, &gen_cfg)
};
(text, None)
};
let (message, finish_reason) = if effective_has_tools {
if let Some(tool_calls) = parse_tool_calls(&raw_output) {
(
ChatMessage {
role: "assistant".to_string(),
content: None,
name: None,
tool_calls: Some(tool_calls),
},
"tool_calls",
)
} else {
let content = extract_content(&raw_output, is_json_mode);
(
ChatMessage {
role: "assistant".to_string(),
content: Some(content),
name: None,
tool_calls: None,
},
"stop",
)
}
} else {
let content = extract_content(&raw_output, is_json_mode);
(
ChatMessage {
role: "assistant".to_string(),
content: Some(content),
name: None,
tool_calls: None,
},
"stop",
)
};
total_completion_tokens += tokenizer.encode(&raw_output).len();
choices.push(Choice {
index: i,
message,
finish_reason,
logprobs: logprobs_field,
});
}
Json(ChatCompletionResponse {
id,
object: "chat.completion",
created: 0,
model: model_name,
choices,
usage: Usage {
prompt_tokens,
completion_tokens: total_completion_tokens,
total_tokens: prompt_tokens + total_completion_tokens,
},
})
.into_response()
}
fn extract_content(raw: &str, is_json_mode: bool) -> String {
if is_json_mode {
extract_json_object(raw).unwrap_or_else(|| raw.to_string())
} else {
raw.to_string()
}
}
fn build_logprobs_content(
tokenizer: &crate::tokenizer::BpeTokenizer,
lps: &[crate::generate::TokenLogprob],
) -> Vec<ContentLogprob> {
lps.iter()
.map(|lp| {
let bytes = tokenizer
.token_bytes(lp.token)
.map(|b| b.to_vec())
.unwrap_or_default();
let token_str = String::from_utf8_lossy(&bytes).into_owned();
let top_logprobs = lp
.top_logprobs
.iter()
.map(|(id, logprob)| {
let tb = tokenizer
.token_bytes(*id)
.map(|b| b.to_vec())
.unwrap_or_default();
TopLogprob {
token: String::from_utf8_lossy(&tb).into_owned(),
logprob: *logprob as f64,
bytes: tb,
}
})
.collect();
ContentLogprob {
token: token_str,
logprob: lp.logprob as f64,
bytes,
top_logprobs,
}
})
.collect()
}
#[derive(Deserialize)]
#[serde(untagged)]
pub(super) enum PromptInput {
Single(String),
Batch(Vec<String>),
}
impl PromptInput {
fn into_vec(self) -> Vec<String> {
match self {
PromptInput::Single(s) => vec![s],
PromptInput::Batch(v) => v,
}
}
}
impl Default for PromptInput {
fn default() -> Self {
PromptInput::Single(String::new())
}
}
#[derive(Deserialize)]
pub(super) struct CompletionRequest {
#[serde(default)]
model: String,
#[serde(default)]
prompt: PromptInput,
#[serde(default)]
max_tokens: Option<usize>,
#[serde(default)]
max_completion_tokens: Option<usize>,
#[serde(default)]
temperature: Option<f32>,
#[serde(default)]
top_p: Option<f32>,
#[serde(default)]
min_p: Option<f32>,
#[serde(default)]
logprobs: Option<bool>,
#[serde(default)]
top_logprobs: Option<u8>,
#[serde(default)]
grammar: Option<String>,
#[serde(default)]
stop: Vec<String>,
#[serde(default)]
seed: Option<u64>,
#[serde(default)]
repetition_penalty: Option<f32>,
#[serde(default)]
presence_penalty: Option<f32>,
#[serde(default)]
frequency_penalty: Option<f32>,
#[serde(default)]
logit_bias: Option<HashMap<u32, f32>>,
#[serde(default)]
echo: Option<bool>,
#[serde(default)]
n: Option<u32>,
#[serde(default)]
best_of: Option<u32>,
#[serde(default)]
suffix: Option<String>,
#[serde(default)]
stream_options: Option<StreamOptions>,
#[serde(default)]
#[allow(dead_code)]
user: Option<String>,
#[serde(default)]
stream: bool,
}
#[derive(Serialize)]
pub(super) struct CompletionResponse {
id: String,
object: &'static str,
created: u64,
model: String,
choices: Vec<CompletionChoice>,
usage: Usage,
}
#[derive(Serialize)]
pub(super) struct CompletionChoice {
index: usize,
text: String,
finish_reason: &'static str,
logprobs: Option<Logprobs>,
}
#[derive(Serialize)]
pub(super) struct CompletionChunk {
id: String,
object: &'static str,
created: u64,
model: String,
choices: Vec<CompletionChunkChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
usage: Option<Usage>,
}
#[derive(Serialize)]
pub(super) struct CompletionChunkChoice {
index: usize,
text: String,
finish_reason: Option<&'static str>,
}
pub(super) async fn completions(
State(state): State<Arc<AppState>>,
Json(req): Json<CompletionRequest>,
) -> axum::response::Response {
let _guard = super::metrics::ActiveRequestGuard::new(&state.metrics);
let _timer = super::metrics::InferenceTimer::new(&state.metrics);
let stream = req.stream;
let constraint = req
.grammar
.and_then(|pat| {
crate::constraint::RegexConstraint::new(&pat).map(|c| {
std::sync::Arc::new(c) as std::sync::Arc<dyn crate::constraint::Constraint>
})
})
.or_else(|| state.generation.constraint.clone());
let gen_cfg = crate::generate::GenerationConfig {
max_tokens: req.max_completion_tokens.or(req.max_tokens).unwrap_or(state.generation.max_tokens),
temperature: req.temperature.unwrap_or(state.generation.temperature),
top_p: req.top_p.unwrap_or(state.generation.top_p),
min_p: req.min_p.unwrap_or(state.generation.min_p),
gamma: state.generation.gamma,
use_int8_kv: state.generation.use_int8_kv,
use_mixed_kv: state.generation.use_mixed_kv,
constraint,
max_context: state.generation.max_context,
anchor_tokens: state.generation.anchor_tokens,
stop: if req.stop.is_empty() {
state.generation.stop.clone()
} else {
req.stop
},
seed: req.seed.or(state.generation.seed),
repetition_penalty: req
.repetition_penalty
.unwrap_or(state.generation.repetition_penalty),
presence_penalty: req
.presence_penalty
.unwrap_or(state.generation.presence_penalty),
frequency_penalty: req
.frequency_penalty
.unwrap_or(state.generation.frequency_penalty),
logit_bias: req
.logit_bias
.unwrap_or_else(|| state.generation.logit_bias.clone()),
cancel: None,
};
let model_name = if req.model.is_empty() {
state.name.clone()
} else {
req.model
};
let id = format!("cmpl-{}-0", state.name);
let prompts = req.prompt.into_vec();
let prompt = prompts[0].clone();
let include_usage = req
.stream_options
.as_ref()
.is_some_and(|so| so.include_usage);
if stream {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Event, Infallible>>(4);
let state_clone = Arc::clone(&state);
let prompt_clone = prompt.clone();
let cancel = Arc::new(std::sync::atomic::AtomicBool::new(false));
let mut gen_cfg = gen_cfg;
gen_cfg.cancel = Some(cancel.clone());
let echo_prompt = req.echo.unwrap_or(false);
tokio::spawn(async move {
if echo_prompt && !prompt_clone.is_empty() {
let chunk = CompletionChunk {
id: id.clone(),
object: "text_completion",
created: 0,
model: model_name.clone(),
choices: vec![CompletionChunkChoice {
index: 0,
text: prompt_clone.clone(),
finish_reason: None,
}],
usage: None,
};
if tx
.send(Ok(
Event::default().data(serde_json::to_string(&chunk).unwrap())
))
.await
.is_err()
{
return;
}
}
let token_ids =
super::infer::run_text_inference_token_ids(&state_clone, &prompt_clone, &gen_cfg);
let tokenizer = crate::tokenizer::BpeTokenizer::byte_fallback();
let prompt_ids = tokenizer.encode(&prompt_clone);
let mut prev_text = if echo_prompt {
prompt_clone.clone()
} else {
String::new()
};
for (idx, &_token_id) in token_ids.iter().enumerate() {
let cumulative = [prompt_ids.as_slice(), &token_ids[..=idx]].concat();
let text = tokenizer.decode(&cumulative);
if let Some(delta) = text.strip_prefix(&prev_text) {
if !delta.is_empty() {
let chunk = CompletionChunk {
id: id.clone(),
object: "text_completion",
created: 0,
model: model_name.clone(),
choices: vec![CompletionChunkChoice {
index: 0,
text: delta.to_string(),
finish_reason: None,
}],
usage: None,
};
if tx
.send(Ok(
Event::default().data(serde_json::to_string(&chunk).unwrap())
))
.await
.is_err()
{
break;
}
}
prev_text = text;
}
}
let final_chunk = CompletionChunk {
id: id.clone(),
object: "text_completion",
created: 0,
model: model_name.clone(),
choices: vec![CompletionChunkChoice {
index: 0,
text: String::new(),
finish_reason: Some("stop"),
}],
usage: None,
};
let _ = tx
.send(Ok(
Event::default().data(serde_json::to_string(&final_chunk).unwrap())
))
.await;
if include_usage {
let usage_chunk = CompletionChunk {
id: id.clone(),
object: "text_completion",
created: 0,
model: model_name.clone(),
choices: vec![],
usage: Some(Usage {
prompt_tokens: prompt_ids.len(),
completion_tokens: token_ids.len(),
total_tokens: prompt_ids.len() + token_ids.len(),
}),
};
let _ = tx
.send(Ok(
Event::default().data(serde_json::to_string(&usage_chunk).unwrap())
))
.await;
}
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
});
return Sse::new(CancelOnDrop::new(ReceiverStream::new(rx), cancel)).into_response();
}
let want_logprobs = req.logprobs.unwrap_or(false);
let top_n = req.top_logprobs.unwrap_or(0).min(20) as usize;
let n = req.n.unwrap_or(1).clamp(1, 20) as usize;
let best_of = req.best_of.unwrap_or(n as u32).max(n as u32).min(20) as usize;
let echo_prompt = req.echo.unwrap_or(false);
let suffix = req.suffix.as_deref().unwrap_or("");
let tokenizer = crate::tokenizer::BpeTokenizer::byte_fallback();
let mut total_prompt_tokens = 0usize;
let mut choices: Vec<CompletionChoice> = Vec::with_capacity(n * prompts.len());
let mut total_completion_tokens = 0usize;
let mut choice_index = 0usize;
for p in &prompts {
total_prompt_tokens += tokenizer.encode(p).len();
let mut all: Vec<(String, Option<Logprobs>, usize)> = Vec::with_capacity(best_of);
for _ in 0..best_of {
let (raw_output, logprobs_field): (String, Option<Logprobs>) = if want_logprobs {
match run_text_inference_with_logprobs(&state, p, &gen_cfg, top_n) {
Some((ids, lps)) => {
let text = tokenizer.decode(&ids);
let content = build_logprobs_content(&tokenizer, &lps);
(text, Some(Logprobs { content }))
}
None => {
let text = run_text_inference_with_config(&state, p, &gen_cfg);
(text, Some(Logprobs { content: Vec::new() }))
}
}
} else {
let text = run_text_inference_with_config(&state, p, &gen_cfg);
(text, None)
};
let tok_count = tokenizer.encode(&raw_output).len();
all.push((raw_output, logprobs_field, tok_count));
}
all.sort_by_key(|b| std::cmp::Reverse(b.2));
all.truncate(n);
for (raw_output, logprobs_field, _) in all.into_iter() {
let display_text = if echo_prompt {
format!("{}{}{}", p, raw_output, suffix)
} else {
format!("{}{}", raw_output, suffix)
};
total_completion_tokens += tokenizer.encode(&display_text).len();
choices.push(CompletionChoice {
index: choice_index,
text: display_text,
finish_reason: "stop",
logprobs: logprobs_field,
});
choice_index += 1;
}
}
Json(CompletionResponse {
id,
object: "text_completion",
created: 0,
model: model_name,
choices,
usage: Usage {
prompt_tokens: total_prompt_tokens,
completion_tokens: total_completion_tokens,
total_tokens: total_prompt_tokens + total_completion_tokens,
},
})
.into_response()
}
fn build_tool_description(tools: &[Tool]) -> String {
let mut lines = vec![
"You have access to the following tools. Respond with a JSON object containing 'tool_calls', an array of objects with 'name' and 'arguments' (a JSON object).".to_string(),
];
for tool in tools {
let desc = if tool.function.description.is_empty() {
format!("- {}: no description", tool.function.name)
} else {
format!("- {}: {}", tool.function.name, tool.function.description)
};
lines.push(desc);
}
lines.join("\n")
}
fn parse_tool_calls(text: &str) -> Option<Vec<ToolCall>> {
let json_str = extract_json_object(text)?;
let val: serde_json::Value = serde_json::from_str(&json_str).ok()?;
let arr = val.get("tool_calls")?.as_array()?;
let mut calls = Vec::new();
for (idx, item) in arr.iter().enumerate() {
let name = item.get("name")?.as_str()?;
let args = item.get("arguments")?;
let args_str = if args.is_string() {
args.as_str()?.to_string()
} else {
serde_json::to_string(args).ok()?
};
calls.push(ToolCall {
id: format!("call-{}", idx),
call_type: "function".to_string(),
function: ToolCallFunction {
name: name.to_string(),
arguments: args_str,
},
});
}
if calls.is_empty() { None } else { Some(calls) }
}
fn extract_json_object(text: &str) -> Option<String> {
let trimmed = text.trim();
if (trimmed.starts_with('{') || trimmed.starts_with('['))
&& serde_json::from_str::<serde_json::Value>(trimmed).is_ok()
{
return Some(trimmed.to_string());
}
for (start, ch) in trimmed.char_indices() {
if ch == '{' || ch == '[' {
for (end, end_ch) in trimmed[start..].char_indices().rev() {
let abs_end = start + end;
if end_ch == '}' || end_ch == ']' {
let candidate = &trimmed[start..=abs_end];
if serde_json::from_str::<serde_json::Value>(candidate).is_ok() {
return Some(candidate.to_string());
}
}
}
}
}
None
}
#[derive(Deserialize)]
pub(super) struct V1EmbeddingsRequest {
input: serde_json::Value,
#[serde(default)]
model: String,
#[serde(default = "default_encoding_format")]
encoding_format: String,
}
fn default_encoding_format() -> String {
"float".to_string()
}
fn base64_encode(bytes: &[u8]) -> String {
const ALPHABET: &[u8; 64] =
b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
let mut out = String::with_capacity(bytes.len().div_ceil(3) * 4);
for chunk in bytes.chunks(3) {
let b0 = chunk[0];
let b1 = chunk.get(1).copied().unwrap_or(0);
let b2 = chunk.get(2).copied().unwrap_or(0);
let n = ((b0 as u32) << 16) | ((b1 as u32) << 8) | (b2 as u32);
out.push(ALPHABET[((n >> 18) & 0x3F) as usize] as char);
out.push(ALPHABET[((n >> 12) & 0x3F) as usize] as char);
if chunk.len() > 1 {
out.push(ALPHABET[((n >> 6) & 0x3F) as usize] as char);
} else {
out.push('=');
}
if chunk.len() > 2 {
out.push(ALPHABET[(n & 0x3F) as usize] as char);
} else {
out.push('=');
}
}
out
}
#[derive(Serialize)]
pub(super) struct V1EmbeddingsResponse {
object: &'static str,
data: Vec<V1Embedding>,
model: String,
usage: Usage,
}
#[derive(Serialize)]
pub(super) struct V1Embedding {
object: &'static str,
embedding: serde_json::Value,
index: usize,
dimensions: usize,
}
pub(super) async fn v1_embeddings(
State(state): State<Arc<AppState>>,
Json(req): Json<V1EmbeddingsRequest>,
) -> Json<V1EmbeddingsResponse> {
let _guard = super::metrics::ActiveRequestGuard::new(&state.metrics);
let _timer = super::metrics::InferenceTimer::new(&state.metrics);
let inputs = extract_embedding_inputs(&req.input);
let use_base64 = req.encoding_format == "base64";
let mut entries = Vec::with_capacity(inputs.len());
let mut total_tokens = 0usize;
for (idx, text) in inputs.iter().enumerate() {
let input_f32: Vec<f32> = text.bytes().map(|b| b as f32 / 255.0).collect();
total_tokens += text.len();
let embedding = run_embeddings(&state, &input_f32).unwrap_or_default();
let dimensions = embedding.len();
let embedding_value = if use_base64 {
let bytes: Vec<u8> = embedding
.iter()
.flat_map(|v| v.to_le_bytes())
.collect();
serde_json::Value::String(base64_encode(&bytes))
} else {
serde_json::to_value(embedding).unwrap_or(serde_json::Value::Null)
};
entries.push(V1Embedding {
object: "embedding",
embedding: embedding_value,
index: idx,
dimensions,
});
}
Json(V1EmbeddingsResponse {
object: "list",
data: entries,
model: if req.model.is_empty() {
state.name.clone()
} else {
req.model
},
usage: Usage {
prompt_tokens: total_tokens,
completion_tokens: 0,
total_tokens,
},
})
}
fn extract_embedding_inputs(input: &serde_json::Value) -> Vec<String> {
if let Some(arr) = input.as_array() {
arr.iter()
.filter_map(|v| v.as_str().map(String::from))
.collect()
} else if let Some(s) = input.as_str() {
vec![s.to_string()]
} else {
Vec::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extract_json_parses_object() {
let s = "Here is the result: {\"answer\": 42} thanks!";
assert_eq!(
extract_json_object(s),
Some(r#"{"answer": 42}"#.to_string())
);
}
#[test]
fn extract_json_parses_array() {
let s = "Some text [1, 2, 3] more text";
assert_eq!(extract_json_object(s), Some("[1, 2, 3]".to_string()));
}
#[test]
fn extract_json_returns_none_for_no_json() {
assert_eq!(extract_json_object("no json here"), None);
}
#[test]
fn extract_json_uses_whole_string_when_valid() {
let s = r#"{"key": "value"}"#;
assert_eq!(
extract_json_object(s),
Some(r#"{"key": "value"}"#.to_string())
);
}
#[test]
fn build_tool_description_lists_tools() {
let tools = vec![Tool {
function: ToolFunction {
name: "get_weather".to_string(),
description: "Get current weather".to_string(),
},
}];
let desc = build_tool_description(&tools);
assert!(desc.contains("get_weather"));
assert!(desc.contains("Get current weather"));
}
#[test]
fn parse_tool_calls_extracts_name_and_args() {
let text = r#"{"tool_calls": [{"name": "get_weather", "arguments": {"city": "NYC"}}]}"#;
let calls = parse_tool_calls(text).expect("parses");
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].function.name, "get_weather");
assert_eq!(calls[0].function.arguments, r#"{"city":"NYC"}"#);
assert_eq!(calls[0].call_type, "function");
}
#[test]
fn parse_tool_calls_returns_none_for_plain_json() {
assert_eq!(parse_tool_calls(r#"{"answer": 42}"#), None);
}
#[test]
fn parse_tool_calls_handles_arguments_as_string() {
let text = r#"{"tool_calls": [{"name": "foo", "arguments": "{\"x\":1}"}]}"#;
let calls = parse_tool_calls(text).expect("parses");
assert_eq!(calls[0].function.arguments, "{\"x\":1}");
}
#[test]
fn base64_encode_matches_known_vectors() {
assert_eq!(base64_encode(b""), "");
assert_eq!(base64_encode(b"f"), "Zg==");
assert_eq!(base64_encode(b"fo"), "Zm8=");
assert_eq!(base64_encode(b"foo"), "Zm9v");
assert_eq!(base64_encode(b"foob"), "Zm9vYg==");
assert_eq!(base64_encode(b"fooba"), "Zm9vYmE=");
assert_eq!(base64_encode(b"foobar"), "Zm9vYmFy");
}
#[test]
fn base64_encode_handles_arbitrary_bytes() {
let bytes: Vec<u8> = (0u8..=255).collect();
let encoded = base64_encode(&bytes);
assert_eq!(encoded.len() % 4, 0);
let mut decoded = Vec::with_capacity(bytes.len());
let chars: Vec<u8> = encoded.bytes().collect();
for chunk in chars.chunks(4) {
let mut vals = [0u32; 4];
for (i, &c) in chunk.iter().enumerate() {
vals[i] = match c {
b'A'..=b'Z' => (c - b'A') as u32,
b'a'..=b'z' => (c - b'a' + 26) as u32,
b'0'..=b'9' => (c - b'0' + 52) as u32,
b'+' => 62,
b'/' => 63,
b'=' => 0,
_ => panic!("unexpected base64 char: {}", c as char),
};
}
let n = (vals[0] << 18) | (vals[1] << 12) | (vals[2] << 6) | vals[3];
decoded.push((n >> 16) as u8);
if chunk[2] != b'=' {
decoded.push((n >> 8) as u8);
}
if chunk[3] != b'=' {
decoded.push(n as u8);
}
}
assert_eq!(decoded, bytes);
}
}