use serde_json::Value;
use crate::{unsupported_feature, ApiError};
pub(crate) fn refuse_logit_bias(value: Option<&Value>, route: &str) -> Result<(), ApiError> {
let Some(value) = value else {
return Ok(());
};
if value.is_null() || value.as_object().is_some_and(|m| m.is_empty()) {
return Ok(());
}
Err(unsupported_feature(&format!(
"`logit_bias` is not implemented on {route} (see docs/API.md). \
It is refused rather than ignored: a dropped bias is \
indistinguishable from an honoured one."
)))
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
#[test]
fn a_real_bias_is_refused_by_name() {
let bias = serde_json::json!({"50256": -100.0});
let (status, body) =
refuse_logit_bias(Some(&bias), "/v1/chat/completions").expect_err("refused");
assert_eq!(status, StatusCode::NOT_IMPLEMENTED);
let message = body["error"]["message"].as_str().expect("message");
assert!(message.contains("logit_bias"), "{message}");
assert!(message.contains("/v1/chat/completions"), "{message}");
}
#[test]
fn an_empty_or_absent_bias_is_served() {
assert!(refuse_logit_bias(None, "/v1/completions").is_ok());
assert!(refuse_logit_bias(Some(&serde_json::json!({})), "/v1/completions").is_ok());
}
#[test]
fn a_malformed_bias_is_refused_rather_than_read_as_empty() {
assert!(refuse_logit_bias(Some(&serde_json::json!([])), "/v1/completions").is_err());
assert!(refuse_logit_bias(Some(&serde_json::json!("none")), "/v1/completions").is_err());
}
#[test]
fn both_routes_answer_a_logit_bias_the_same_way() {
for (bias, expected_refusal) in [
(serde_json::json!({"50256": -100.0}), true),
(serde_json::json!({}), false),
] {
let chat: crate::ChatCompletionRequest = serde_json::from_value(serde_json::json!({
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"logit_bias": bias,
}))
.expect("chat request");
let completion: crate::openai_extra::CompletionsRequest =
serde_json::from_value(serde_json::json!({
"prompt": "hi",
"logit_bias": bias,
}))
.expect("completions request");
let chat_status = chat.validate_supported_fields().err().map(|(s, _)| s);
let completion_status = completion.validate().err().map(|(s, _)| s);
assert_eq!(
chat_status, completion_status,
"the two routes disagree about logit_bias {bias}"
);
assert_eq!(
chat_status.is_some(),
expected_refusal,
"wrong verdict for logit_bias {bias}"
);
}
}
}