fix: remap max_completion_tokens to max_tokens for OpenAI-compatible providers (#7765)
Signed-off-by: fre <anonwurcod@proton.me>
This commit is contained in:
@@ -251,6 +251,34 @@ impl OpenAiProvider {
|
|||||||
Self::is_responses_model(model_name)
|
Self::is_responses_model(model_name)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Providers known to reject `max_completion_tokens` and require
|
||||||
|
/// the legacy `max_tokens` field instead.
|
||||||
|
const PROVIDERS_NEEDING_MAX_TOKENS_REMAP: &[&str] = &[
|
||||||
|
"cerebras",
|
||||||
|
"custom_deepseek",
|
||||||
|
"groq",
|
||||||
|
"inception",
|
||||||
|
"kimi",
|
||||||
|
"lmstudio",
|
||||||
|
"mistral",
|
||||||
|
"moonshot",
|
||||||
|
"ovhcloud",
|
||||||
|
];
|
||||||
|
|
||||||
|
fn sanitize_request_for_compat(&self, mut payload: serde_json::Value) -> serde_json::Value {
|
||||||
|
if !Self::PROVIDERS_NEEDING_MAX_TOKENS_REMAP.contains(&self.name.as_str()) {
|
||||||
|
return payload;
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(obj) = payload.as_object_mut() {
|
||||||
|
if let Some(value) = obj.remove("max_completion_tokens") {
|
||||||
|
obj.entry("max_tokens").or_insert(value);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
payload
|
||||||
|
}
|
||||||
|
|
||||||
fn map_base_path(base_path: &str, target: &str, fallback: &str) -> String {
|
fn map_base_path(base_path: &str, target: &str, fallback: &str) -> String {
|
||||||
let normalized = Self::normalize_base_path(base_path);
|
let normalized = Self::normalize_base_path(base_path);
|
||||||
if normalized.ends_with(target) || normalized.contains(&format!("/{target}")) {
|
if normalized.ends_with(target) || normalized.contains(&format!("/{target}")) {
|
||||||
@@ -457,6 +485,7 @@ impl Provider for OpenAiProvider {
|
|||||||
&ImageFormat::OpenAi,
|
&ImageFormat::OpenAi,
|
||||||
self.supports_streaming,
|
self.supports_streaming,
|
||||||
)?;
|
)?;
|
||||||
|
let payload = self.sanitize_request_for_compat(payload);
|
||||||
let mut log = RequestLog::start(model_config, &payload)?;
|
let mut log = RequestLog::start(model_config, &payload)?;
|
||||||
|
|
||||||
let response = self
|
let response = self
|
||||||
@@ -568,7 +597,98 @@ impl EmbeddingCapable for OpenAiProvider {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::OpenAiProvider;
|
use super::*;
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
fn make_provider(name: &str) -> OpenAiProvider {
|
||||||
|
OpenAiProvider {
|
||||||
|
api_client: ApiClient::new("http://localhost".to_string(), AuthMethod::NoAuth).unwrap(),
|
||||||
|
base_path: "v1/chat/completions".to_string(),
|
||||||
|
organization: None,
|
||||||
|
project: None,
|
||||||
|
model: ModelConfig::new_or_fail("test-model"),
|
||||||
|
custom_headers: None,
|
||||||
|
supports_streaming: true,
|
||||||
|
name: name.to_string(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sanitize_remaps_max_completion_tokens_for_compat_provider() {
|
||||||
|
let provider = make_provider("mistral");
|
||||||
|
let payload = json!({
|
||||||
|
"model": "mistral-medium-latest",
|
||||||
|
"messages": [],
|
||||||
|
"max_completion_tokens": 16384
|
||||||
|
});
|
||||||
|
|
||||||
|
let result = provider.sanitize_request_for_compat(payload);
|
||||||
|
let obj = result.as_object().unwrap();
|
||||||
|
|
||||||
|
assert!(!obj.contains_key("max_completion_tokens"));
|
||||||
|
assert_eq!(obj.get("max_tokens").unwrap(), &json!(16384));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sanitize_preserves_existing_max_tokens_for_compat_provider() {
|
||||||
|
let provider = make_provider("mistral");
|
||||||
|
let payload = json!({
|
||||||
|
"model": "mistral-medium-latest",
|
||||||
|
"messages": [],
|
||||||
|
"max_tokens": 4096,
|
||||||
|
"max_completion_tokens": 16384
|
||||||
|
});
|
||||||
|
|
||||||
|
let result = provider.sanitize_request_for_compat(payload);
|
||||||
|
let obj = result.as_object().unwrap();
|
||||||
|
|
||||||
|
assert!(!obj.contains_key("max_completion_tokens"));
|
||||||
|
assert_eq!(obj.get("max_tokens").unwrap(), &json!(4096));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sanitize_noop_for_native_openai_provider() {
|
||||||
|
let provider = make_provider("openai");
|
||||||
|
let payload = json!({
|
||||||
|
"model": "o3",
|
||||||
|
"messages": [],
|
||||||
|
"max_completion_tokens": 16384
|
||||||
|
});
|
||||||
|
|
||||||
|
let result = provider.sanitize_request_for_compat(payload);
|
||||||
|
let obj = result.as_object().unwrap();
|
||||||
|
|
||||||
|
assert!(obj.contains_key("max_completion_tokens"));
|
||||||
|
assert!(!obj.contains_key("max_tokens"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sanitize_noop_for_unknown_provider() {
|
||||||
|
let provider = make_provider("some_future_provider");
|
||||||
|
let payload = json!({
|
||||||
|
"model": "future-model",
|
||||||
|
"messages": [],
|
||||||
|
"max_completion_tokens": 16384
|
||||||
|
});
|
||||||
|
|
||||||
|
let result = provider.sanitize_request_for_compat(payload);
|
||||||
|
let obj = result.as_object().unwrap();
|
||||||
|
|
||||||
|
assert!(obj.contains_key("max_completion_tokens"));
|
||||||
|
assert!(!obj.contains_key("max_tokens"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sanitize_no_token_params() {
|
||||||
|
let provider = make_provider("groq");
|
||||||
|
let payload = json!({
|
||||||
|
"model": "llama-3.3-70b-versatile",
|
||||||
|
"messages": []
|
||||||
|
});
|
||||||
|
|
||||||
|
let result = provider.sanitize_request_for_compat(payload.clone());
|
||||||
|
assert_eq!(result, payload);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn gpt_5_2_codex_uses_responses_when_base_path_is_default() {
|
fn gpt_5_2_codex_uses_responses_when_base_path_is_default() {
|
||||||
|
|||||||
Reference in New Issue
Block a user