use std::collections::BTreeMap;
use serde::{Deserialize, Deserializer, Serialize};
use super::error::ApiError;
pub const MAX_TOP_LOGPROBS: usize = 20;
pub const MAX_STOP: usize = 4;
pub const MAX_N: usize = 8;
fn default_max_tokens() -> usize {
64
}
#[derive(Deserialize)]
#[serde(untagged)]
pub enum StopField {
One(String),
Many(Vec<String>),
}
impl StopField {
pub fn to_vec(&self) -> Vec<String> {
match self {
StopField::One(s) => vec![s.clone()],
StopField::Many(v) => v.clone(),
}
}
}
#[derive(Deserialize)]
#[serde(untagged)]
pub enum PromptField {
Text(String),
Texts(Vec<String>),
}
#[derive(Deserialize, Default)]
pub struct StreamOptions {
#[serde(default)]
pub include_usage: bool,
}
fn de_null_string<'de, D>(d: D) -> Result<String, D::Error>
where
D: Deserializer<'de>,
{
Ok(Option::<String>::deserialize(d)?.unwrap_or_default())
}
#[derive(Deserialize)]
pub struct ToolCallMsg {
#[serde(default)]
pub id: Option<String>,
#[serde(rename = "type", default)]
pub kind: Option<String>,
pub function: FunctionCallMsg,
}
#[derive(Deserialize)]
pub struct FunctionCallMsg {
pub name: String,
#[serde(default, deserialize_with = "de_null_string")]
pub arguments: String,
}
#[derive(Deserialize)]
pub struct ChatMessage {
pub role: String,
#[serde(default, deserialize_with = "de_null_string")]
pub content: String,
#[serde(default)]
pub tool_calls: Vec<ToolCallMsg>,
#[serde(default)]
pub tool_call_id: Option<String>,
#[serde(default)]
pub name: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct CommonParams {
pub max_tokens: usize,
pub n: usize,
pub temperature: f32,
pub top_p: f32,
pub top_k: Option<usize>,
pub min_p: Option<f32>,
pub seed: u64,
pub presence_penalty: f32,
pub frequency_penalty: f32,
pub repetition_penalty: f32,
pub logit_bias: Vec<(u32, f32)>,
pub stop: Vec<String>,
pub stop_token_ids: Vec<u32>,
pub logprobs: Option<usize>,
pub stream: bool,
pub include_usage: bool,
}
#[derive(Deserialize, Default)]
pub struct RawParams {
#[serde(default = "default_max_tokens")]
pub max_tokens: usize,
#[serde(default)]
pub n: Option<usize>,
#[serde(default)]
pub temperature: Option<f32>,
#[serde(default)]
pub top_p: Option<f32>,
#[serde(default)]
pub top_k: Option<usize>,
#[serde(default)]
pub min_p: Option<f32>,
#[serde(default)]
pub seed: Option<u64>,
#[serde(default)]
pub presence_penalty: Option<f32>,
#[serde(default)]
pub frequency_penalty: Option<f32>,
#[serde(default)]
pub repetition_penalty: Option<f32>,
#[serde(default)]
pub logit_bias: Option<BTreeMap<String, f32>>,
#[serde(default)]
pub stop: Option<StopField>,
#[serde(default)]
pub stop_token_ids: Option<Vec<u32>>,
#[serde(default)]
pub stream: bool,
#[serde(default)]
pub stream_options: Option<StreamOptions>,
#[serde(default)]
pub best_of: Option<usize>,
#[serde(default)]
pub suffix: Option<String>,
#[serde(default)]
pub echo: Option<bool>,
#[serde(default)]
pub tools: Option<Vec<serde_json::Value>>,
#[serde(default)]
pub tool_choice: Option<serde_json::Value>,
#[serde(default)]
pub functions: Option<Vec<serde_json::Value>>,
#[serde(default)]
pub function_call: Option<serde_json::Value>,
#[serde(default)]
pub response_format: Option<serde_json::Value>,
}
fn in_range(name: &str, v: f32, lo: f32, hi: f32) -> Result<(), ApiError> {
if v.is_finite() && (lo..=hi).contains(&v) {
Ok(())
} else {
Err(ApiError::invalid_request(format!("{name} must be in [{lo}, {hi}]")).with_param(name))
}
}
impl RawParams {
pub fn validate(&self, logprobs: Option<usize>) -> Result<CommonParams, ApiError> {
if self.best_of.is_some_and(|b| b > 1) {
return Err(
ApiError::invalid_request("best_of > 1 is not implemented").with_param("best_of")
);
}
if self.suffix.is_some() {
return Err(ApiError::invalid_request("suffix is not implemented").with_param("suffix"));
}
if self.echo == Some(true) {
return Err(ApiError::invalid_request("echo is not implemented").with_param("echo"));
}
if self.functions.as_ref().is_some_and(|f| !f.is_empty()) {
return Err(
ApiError::invalid_request("functions are not implemented").with_param("functions")
);
}
if self.function_call.is_some() {
return Err(
ApiError::invalid_request("function_call is not implemented")
.with_param("function_call"),
);
}
if self.response_format.is_some() {
return Err(ApiError::invalid_request(
"response_format (structured output) is not implemented on this endpoint",
)
.with_param("response_format"));
}
let temperature = self.temperature.unwrap_or(0.0);
in_range("temperature", temperature, 0.0, 10.0)?;
let top_p = self.top_p.unwrap_or(1.0);
if !(top_p.is_finite() && top_p > 0.0 && top_p <= 1.0) {
return Err(ApiError::invalid_request("top_p must be in (0, 1]").with_param("top_p"));
}
if let Some(k) = self.top_k
&& k == 0
{
return Err(
ApiError::invalid_request("top_k must be >= 1 (omit for unlimited)")
.with_param("top_k"),
);
}
if let Some(mp) = self.min_p {
in_range("min_p", mp, 0.0, 1.0)?;
}
let presence_penalty = self.presence_penalty.unwrap_or(0.0);
in_range("presence_penalty", presence_penalty, -2.0, 2.0)?;
let frequency_penalty = self.frequency_penalty.unwrap_or(0.0);
in_range("frequency_penalty", frequency_penalty, -2.0, 2.0)?;
let repetition_penalty = self.repetition_penalty.unwrap_or(1.0);
if !(repetition_penalty.is_finite() && repetition_penalty > 0.0) {
return Err(ApiError::invalid_request("repetition_penalty must be > 0")
.with_param("repetition_penalty"));
}
let n = self.n.unwrap_or(1);
if n == 0 || n > MAX_N {
return Err(
ApiError::invalid_request(format!("n must be in 1..={MAX_N}")).with_param("n"),
);
}
if self.max_tokens == 0 {
return Err(
ApiError::invalid_request("max_tokens must be >= 1").with_param("max_tokens")
);
}
let logit_bias = match &self.logit_bias {
None => Vec::new(),
Some(m) => {
let mut out = Vec::with_capacity(m.len());
for (k, &v) in m {
let id: u32 = k.parse().map_err(|_| {
ApiError::invalid_request(format!(
"logit_bias keys must be integer token ids, got `{k}`"
))
.with_param("logit_bias")
})?;
out.push((id, v));
}
out
}
};
let stop = match &self.stop {
None => Vec::new(),
Some(field) => {
let v = field.to_vec();
if v.len() > MAX_STOP {
return Err(ApiError::invalid_request(format!(
"at most {MAX_STOP} stop strings are allowed"
))
.with_param("stop"));
}
v.into_iter().filter(|s| !s.is_empty()).collect()
}
};
if let Some(l) = logprobs
&& l > MAX_TOP_LOGPROBS
{
return Err(ApiError::invalid_request(format!(
"top_logprobs must be in 0..={MAX_TOP_LOGPROBS}"
))
.with_param("top_logprobs"));
}
Ok(CommonParams {
max_tokens: self.max_tokens,
n,
temperature,
top_p,
top_k: self.top_k,
min_p: self.min_p,
seed: self.seed.unwrap_or(0),
presence_penalty,
frequency_penalty,
repetition_penalty,
logit_bias,
stop,
stop_token_ids: self.stop_token_ids.clone().unwrap_or_default(),
logprobs,
stream: self.stream,
include_usage: self
.stream_options
.as_ref()
.map(|o| o.include_usage)
.unwrap_or(false),
})
}
pub fn tool_choice_mode(&self) -> Result<ToolMode, ApiError> {
match &self.tool_choice {
None => Ok(ToolMode::Auto),
Some(v) => match v.as_str() {
Some("auto") => Ok(ToolMode::Auto),
Some("none") => Ok(ToolMode::None),
Some("required") => Err(ApiError::invalid_request(
"tool_choice \"required\" needs constrained decoding, which is not implemented",
)
.with_param("tool_choice")),
_ => Err(ApiError::invalid_request(
"tool_choice must be \"auto\" or \"none\" (forced/named tool choice is not implemented)",
)
.with_param("tool_choice")),
},
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ToolMode {
Auto,
None,
}
#[derive(Deserialize)]
pub struct CompletionReq {
#[serde(default)]
pub model: Option<String>,
pub prompt: PromptField,
#[serde(default)]
pub images: Vec<String>,
#[serde(default)]
pub logprobs: Option<usize>,
#[serde(flatten)]
pub params: RawParams,
}
impl CompletionReq {
pub fn prompt_text(&self) -> Result<&str, ApiError> {
match &self.prompt {
PromptField::Text(s) => Ok(s.as_str()),
PromptField::Texts(v) if v.len() == 1 => Ok(v[0].as_str()),
PromptField::Texts(_) => Err(ApiError::invalid_request(
"multiple prompts in one request are not implemented (send one prompt, or use n)",
)
.with_param("prompt")),
}
}
pub fn common(&self) -> Result<CommonParams, ApiError> {
if self.params.tools.as_ref().is_some_and(|t| !t.is_empty())
|| self.params.tool_choice.is_some()
{
return Err(ApiError::invalid_request(
"tools / function calling is not supported on /v1/completions — use /v1/chat/completions",
)
.with_param("tools"));
}
self.params.validate(self.logprobs)
}
}
#[derive(Deserialize)]
pub struct ChatReq {
#[serde(default)]
pub model: Option<String>,
pub messages: Vec<ChatMessage>,
#[serde(default)]
pub logprobs: Option<bool>,
#[serde(default)]
pub top_logprobs: Option<usize>,
#[serde(default)]
pub max_completion_tokens: Option<usize>,
#[serde(flatten)]
pub params: RawParams,
}
impl ChatReq {
pub fn common(&self) -> Result<CommonParams, ApiError> {
let logprobs = if self.logprobs == Some(true) {
Some(self.top_logprobs.unwrap_or(0))
} else {
None
};
let mut c = self.params.validate(logprobs)?;
if let Some(m) = self.max_completion_tokens {
if m == 0 {
return Err(
ApiError::invalid_request("max_completion_tokens must be >= 1")
.with_param("max_completion_tokens"),
);
}
c.max_tokens = m;
}
Ok(c)
}
pub fn tools_to_render(&self) -> Result<Vec<serde_json::Value>, ApiError> {
let mode = self.params.tool_choice_mode()?;
match &self.params.tools {
Some(tools) if !tools.is_empty() && mode == ToolMode::Auto => Ok(tools.clone()),
_ => Ok(Vec::new()),
}
}
}
#[derive(Serialize, Clone)]
pub struct TopLogprob {
pub token: String,
pub logprob: f32,
}
#[derive(Serialize, Clone, Debug, PartialEq)]
pub struct Usage {
pub prompt_tokens: usize,
pub completion_tokens: usize,
pub total_tokens: usize,
pub prompt_tokens_details: PromptTokensDetails,
}
#[derive(Serialize, Clone, Debug, PartialEq)]
pub struct PromptTokensDetails {
pub cached_tokens: usize,
}
impl Usage {
pub fn new(prompt_tokens: usize, completion_tokens: usize, cached_tokens: usize) -> Self {
Self {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
prompt_tokens_details: PromptTokensDetails { cached_tokens },
}
}
}
pub fn created_epoch() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
fn monotonic() -> u64 {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
COUNTER.fetch_add(1, Ordering::Relaxed)
}
pub fn response_id(prefix: &str) -> String {
format!("{prefix}-{}-{}", std::process::id(), monotonic())
}
#[cfg(test)]
mod tests {
use super::*;
fn raw(json: serde_json::Value) -> RawParams {
serde_json::from_value(json).expect("raw params")
}
#[test]
fn should_default_and_validate_the_happy_path() {
let c = raw(serde_json::json!({})).validate(None).expect("valid");
assert_eq!(c.max_tokens, 64);
assert_eq!(c.n, 1);
assert_eq!(c.temperature, 0.0);
assert_eq!(c.top_p, 1.0);
assert_eq!(c.repetition_penalty, 1.0);
assert!(c.stop.is_empty());
assert!(c.logit_bias.is_empty());
assert_eq!(c.logprobs, None);
}
#[test]
fn should_reject_out_of_range_top_p() {
let e = raw(serde_json::json!({"top_p": 1.5}))
.validate(None)
.expect_err("reject");
assert_eq!(e.param.as_deref(), Some("top_p"));
}
#[test]
fn should_reject_n_outside_one_to_eight() {
assert!(raw(serde_json::json!({"n": 0})).validate(None).is_err());
assert!(raw(serde_json::json!({"n": 9})).validate(None).is_err());
assert_eq!(
raw(serde_json::json!({"n": 8}))
.validate(None)
.expect("ok")
.n,
8
);
}
#[test]
fn should_reject_unimplemented_features_loudly() {
for (field, body) in [
("best_of", serde_json::json!({"best_of": 2})),
("suffix", serde_json::json!({"suffix": "x"})),
("echo", serde_json::json!({"echo": true})),
(
"response_format",
serde_json::json!({"response_format": {"type": "json_object"}}),
),
] {
let e = raw(body).validate(None).expect_err("reject");
assert_eq!(e.param.as_deref(), Some(field), "field {field}");
assert_eq!(e.status, axum::http::StatusCode::BAD_REQUEST);
}
}
#[test]
fn should_allow_default_valued_unimplemented_features() {
let c = raw(serde_json::json!({"best_of": 1, "echo": false, "tools": []})).validate(None);
assert!(c.is_ok(), "defaults must pass: {c:?}");
}
fn chat(json: serde_json::Value) -> ChatReq {
serde_json::from_value(json).expect("chat req")
}
#[test]
fn chat_advertises_tools_only_under_auto() {
let tools = serde_json::json!([{"type": "function", "function": {"name": "f"}}]);
let auto = chat(serde_json::json!({"messages": [], "tools": tools}));
assert_eq!(auto.tools_to_render().expect("ok").len(), 1);
let explicit = chat(serde_json::json!({
"messages": [], "tools": tools, "tool_choice": "auto"
}));
assert_eq!(explicit.tools_to_render().expect("ok").len(), 1);
let none = chat(serde_json::json!({
"messages": [], "tools": tools, "tool_choice": "none"
}));
assert!(none.tools_to_render().expect("ok").is_empty());
let bare = chat(serde_json::json!({"messages": []}));
assert!(bare.tools_to_render().expect("ok").is_empty());
}
#[test]
fn chat_rejects_forced_and_named_tool_choice() {
for tc in [
serde_json::json!("required"),
serde_json::json!({"type": "function", "function": {"name": "f"}}),
] {
let req = chat(serde_json::json!({"messages": [], "tool_choice": tc}));
let e = req.tools_to_render().expect_err("reject");
assert_eq!(e.param.as_deref(), Some("tool_choice"));
}
}
#[test]
fn completions_reject_tools() {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"prompt": "hi", "tools": [{"type": "function", "function": {"name": "f"}}]
}))
.expect("parse");
let e = req.common().expect_err("reject");
assert_eq!(e.param.as_deref(), Some("tools"));
}
#[test]
fn chat_message_parses_tool_calls_and_null_content() {
let req = chat(serde_json::json!({
"messages": [{
"role": "assistant",
"content": null,
"tool_calls": [{
"id": "call_1", "type": "function",
"function": {"name": "get_weather", "arguments": "{\"city\":\"Paris\"}"}
}]
}, {
"role": "tool", "tool_call_id": "call_1", "content": "18C"
}]
}));
let asst = &req.messages[0];
assert_eq!(asst.content, "");
assert_eq!(asst.tool_calls.len(), 1);
assert_eq!(asst.tool_calls[0].function.name, "get_weather");
assert_eq!(
asst.tool_calls[0].function.arguments,
"{\"city\":\"Paris\"}"
);
let tool = &req.messages[1];
assert_eq!(tool.content, "18C");
assert_eq!(tool.tool_call_id.as_deref(), Some("call_1"));
}
#[test]
fn should_normalize_stop_string_and_array() {
let one = raw(serde_json::json!({"stop": "END"}))
.validate(None)
.expect("ok");
assert_eq!(one.stop, vec!["END".to_string()]);
let many = raw(serde_json::json!({"stop": ["A", "B"]}))
.validate(None)
.expect("ok");
assert_eq!(many.stop, vec!["A".to_string(), "B".to_string()]);
assert!(
raw(serde_json::json!({"stop": ["a", "b", "c", "d", "e"]}))
.validate(None)
.is_err()
);
}
#[test]
fn should_parse_logit_bias_keys_as_token_ids() {
let c = raw(serde_json::json!({"logit_bias": {"5": 10.0, "42": -100.0}}))
.validate(None)
.expect("ok");
let mut got = c.logit_bias.clone();
got.sort_by_key(|(id, _)| *id);
assert_eq!(got, vec![(5, 10.0), (42, -100.0)]);
assert!(
raw(serde_json::json!({"logit_bias": {"foo": 1.0}}))
.validate(None)
.is_err()
);
}
#[test]
fn should_cap_top_logprobs_at_twenty() {
assert!(raw(serde_json::json!({})).validate(Some(21)).is_err());
assert_eq!(
raw(serde_json::json!({}))
.validate(Some(20))
.expect("ok")
.logprobs,
Some(20)
);
}
#[test]
fn should_thread_stream_options_include_usage() {
let c = raw(serde_json::json!({"stream": true, "stream_options": {"include_usage": true}}))
.validate(None)
.expect("ok");
assert!(c.stream);
assert!(c.include_usage);
}
#[test]
fn usage_totals_are_the_sum() {
let u = Usage::new(10, 5, 3);
assert_eq!(u.total_tokens, 15);
assert_eq!(u.prompt_tokens_details.cached_tokens, 3);
}
#[test]
fn response_ids_are_unique_and_prefixed() {
let a = response_id("cmpl");
let b = response_id("cmpl");
assert_ne!(a, b, "ids must be distinct: {a} vs {b}");
assert!(a.starts_with("cmpl-"));
}
#[test]
fn chat_max_completion_tokens_overrides_max_tokens() {
let req: ChatReq = serde_json::from_value(serde_json::json!({
"messages": [], "max_tokens": 64, "max_completion_tokens": 7
}))
.expect("parse");
assert_eq!(req.common().expect("ok").max_tokens, 7);
}
#[test]
fn chat_logprobs_bool_plus_top_logprobs_becomes_count() {
let req: ChatReq = serde_json::from_value(serde_json::json!({
"messages": [], "logprobs": true, "top_logprobs": 3
}))
.expect("parse");
assert_eq!(req.common().expect("ok").logprobs, Some(3));
let off: ChatReq = serde_json::from_value(serde_json::json!({
"messages": [], "logprobs": false, "top_logprobs": 3
}))
.expect("parse");
assert_eq!(off.common().expect("ok").logprobs, None);
}
#[test]
fn completion_rejects_multi_prompt_arrays() {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"prompt": ["a", "b"]
}))
.expect("parse");
assert!(req.prompt_text().is_err());
let one: CompletionReq =
serde_json::from_value(serde_json::json!({"prompt": ["solo"]})).expect("parse");
assert_eq!(one.prompt_text().expect("ok"), "solo");
}
}