feat: support per-request model override in /v1/chat/completions (#103)

* feat: support per-request model override for /v1/chat/completions

- add optional model override to completion request types\n- forward request model through gateway, worker, and orchestrator proxy paths\n- use request model in NEAR AI providers with fallback to active model\n- replace model-mismatch integration test with override propagation checks\n- update FEATURE_PARITY.md note for OpenAI-compatible API behavior\n\nRefs #49

* Wire gateway OpenAI-compatible routes to active LLM provider

* Validate OpenAI model name length before streaming

* Address PR103 review feedback on model override and validation

* Report effective model in OpenAI-compatible responses

* Use async mutexes in OpenAI compatibility integration tests

* fix tests for per-request model field in response cache

* fix formatting and clippy lint after main merge

* Fix model override reporting and cache correctness

---------

Co-authored-by: Illia Polosukhin <[email protected]>
This commit is contained in:
Raahim Salman
2026-02-19 21:45:37 +00:00
committed by GitHub
co-authored by Illia Polosukhin
parent 89fdd81420
commit ccf60055f4
13 changed files with 519 additions and 62 deletions
+51 -24
View File
@@ -24,6 +24,8 @@ use crate::llm::{
use super::server::GatewayState;
const MAX_MODEL_NAME_BYTES: usize = 256;
// ---------------------------------------------------------------------------
// OpenAI request types
// ---------------------------------------------------------------------------
@@ -380,6 +382,27 @@ fn unix_timestamp() -> u64 {
.as_secs()
}
fn validate_model_name(model: &str) -> Result<(), String> {
let trimmed = model.trim();
if trimmed.is_empty() {
return Err("model must not be empty".to_string());
}
if trimmed != model {
return Err("model must not have leading or trailing whitespace".to_string());
}
if model.len() > MAX_MODEL_NAME_BYTES {
return Err(format!(
"model must be at most {} bytes",
MAX_MODEL_NAME_BYTES
));
}
if model.chars().any(char::is_control) {
return Err("model contains control characters".to_string());
}
Ok(())
}
/// Extract stop sequences from the flexible `stop` field.
fn parse_stop(val: &serde_json::Value) -> Option<Vec<String>> {
match val {
@@ -426,29 +449,17 @@ pub async fn chat_completions_handler(
"invalid_request_error",
));
}
// Validate the requested model matches the active model.
// Per-request model switching is not yet supported (see GH issue).
let active_model = llm.active_model_name();
if req.model != active_model {
return Err((
StatusCode::NOT_FOUND,
Json(OpenAiErrorResponse {
error: OpenAiErrorDetail {
message: format!(
"Model '{}' not found. The active model is '{}'.",
req.model, active_model
),
error_type: "invalid_request_error".to_string(),
param: Some("model".to_string()),
code: Some("model_not_found".to_string()),
},
}),
if let Err(e) = validate_model_name(&req.model) {
return Err(openai_error(
StatusCode::BAD_REQUEST,
e,
"invalid_request_error",
));
}
let has_tools = req.tools.as_ref().is_some_and(|t| !t.is_empty());
let stream = req.stream.unwrap_or(false);
let requested_model = req.model.clone();
if stream {
return handle_streaming(llm.clone(), req, has_tools)
@@ -460,13 +471,12 @@ pub async fn chat_completions_handler(
let messages = convert_messages(&req.messages)
.map_err(|e| openai_error(StatusCode::BAD_REQUEST, e, "invalid_request_error"))?;
let model_name = llm.active_model_name();
let id = chat_completion_id();
let created = unix_timestamp();
if has_tools {
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
let mut tool_req = ToolCompletionRequest::new(messages, tools);
let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model);
if let Some(t) = req.temperature {
tool_req = tool_req.with_temperature(t);
}
@@ -483,6 +493,7 @@ pub async fn chat_completions_handler(
.complete_with_tools(tool_req)
.await
.map_err(map_llm_error)?;
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
let tool_calls_openai = if resp.tool_calls.is_empty() {
None
@@ -515,7 +526,7 @@ pub async fn chat_completions_handler(
Ok(Json(response).into_response())
} else {
let mut comp_req = CompletionRequest::new(messages);
let mut comp_req = CompletionRequest::new(messages).with_model(req.model);
if let Some(t) = req.temperature {
comp_req = comp_req.with_temperature(t);
}
@@ -527,6 +538,7 @@ pub async fn chat_completions_handler(
}
let resp = llm.complete(comp_req).await.map_err(map_llm_error)?;
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
let response = OpenAiChatResponse {
id,
@@ -570,7 +582,7 @@ async fn handle_streaming(
let messages = convert_messages(&req.messages)
.map_err(|e| openai_error(StatusCode::BAD_REQUEST, e, "invalid_request_error"))?;
let model_name = llm.active_model_name();
let requested_model = req.model.clone();
let id = chat_completion_id();
let created = unix_timestamp();
@@ -584,7 +596,7 @@ async fn handle_streaming(
let llm_result = if has_tools {
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
let mut tool_req = ToolCompletionRequest::new(messages, tools);
let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model);
if let Some(t) = req.temperature {
tool_req = tool_req.with_temperature(t);
}
@@ -602,7 +614,7 @@ async fn handle_streaming(
.map_err(map_llm_error)?,
)
} else {
let mut comp_req = CompletionRequest::new(messages);
let mut comp_req = CompletionRequest::new(messages).with_model(req.model);
if let Some(t) = req.temperature {
comp_req = comp_req.with_temperature(t);
}
@@ -614,6 +626,7 @@ async fn handle_streaming(
}
LlmResult::Simple(llm.complete(comp_req).await.map_err(map_llm_error)?)
};
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
// LLM succeeded — emit the response as SSE chunks
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Event, std::convert::Infallible>>(64);
@@ -1091,4 +1104,18 @@ mod tests {
let v = serde_json::Value::Null;
assert_eq!(parse_stop(&v), None);
}
#[test]
fn test_validate_model_name_rejects_leading_or_trailing_whitespace() {
let err = validate_model_name(" gpt-4").unwrap_err();
assert!(err.contains("leading or trailing whitespace"));
let err = validate_model_name("gpt-4 ").unwrap_err();
assert!(err.contains("leading or trailing whitespace"));
}
#[test]
fn test_validate_model_name_accepts_normal_name() {
assert!(validate_model_name("gpt-4").is_ok());
}
}
+4
View File
@@ -273,6 +273,10 @@ impl LlmProvider for CircuitBreakerProvider {
self.inner.model_metadata().await
}
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
self.inner.effective_model_name(requested_model)
}
fn active_model_name(&self) -> String {
self.inner.active_model_name()
}
+102 -12
View File
@@ -7,8 +7,10 @@
//! so subsequent requests skip them, reducing latency when a provider
//! is known to be down. Cooldown state is lock-free (atomics only).
use std::collections::HashMap;
use std::future::Future;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
@@ -139,6 +141,12 @@ pub struct FailoverProvider {
epoch: Instant,
/// Cooldown configuration.
cooldown_config: CooldownConfig,
/// Request-scoped provider index keyed by Tokio task ID.
///
/// This allows `effective_model_name()` to report the provider that handled
/// the *current* request, even when other concurrent requests update
/// `last_used`.
provider_for_task: Mutex<HashMap<tokio::task::Id, usize>>,
}
impl FailoverProvider {
@@ -171,6 +179,7 @@ impl FailoverProvider {
cooldowns,
epoch: Instant::now(),
cooldown_config,
provider_for_task: Mutex::new(HashMap::new()),
})
}
@@ -182,12 +191,36 @@ impl FailoverProvider {
self.epoch.elapsed().as_nanos() as u64
}
/// Current Tokio task ID if available.
fn current_task_id() -> Option<tokio::task::Id> {
tokio::task::try_id()
}
/// Bind the selected provider index to the current task.
fn bind_provider_to_current_task(&self, provider_idx: usize) {
let Some(task_id) = Self::current_task_id() else {
return;
};
if let Ok(mut guard) = self.provider_for_task.lock() {
guard.insert(task_id, provider_idx);
}
}
/// Take and remove the provider index bound to the current task.
fn take_bound_provider_for_current_task(&self) -> Option<usize> {
let task_id = Self::current_task_id()?;
self.provider_for_task
.lock()
.ok()
.and_then(|mut guard| guard.remove(&task_id))
}
/// Try each provider in sequence until one succeeds or all fail.
///
/// Providers in cooldown are skipped unless *all* providers are in
/// cooldown, in which case the one with the oldest cooldown timestamp
/// (most likely to have recovered) is tried.
async fn try_providers<T, F, Fut>(&self, mut call: F) -> Result<T, LlmError>
async fn try_providers<T, F, Fut>(&self, mut call: F) -> Result<(usize, T), LlmError>
where
F: FnMut(Arc<dyn LlmProvider>) -> Fut,
Fut: Future<Output = Result<T, LlmError>>,
@@ -236,7 +269,7 @@ impl FailoverProvider {
Ok(response) => {
self.last_used.store(i, Ordering::Relaxed);
self.cooldowns[i].reset();
return Ok(response);
return Ok((i, response));
}
Err(err) => {
if !is_retryable(&err) {
@@ -287,22 +320,28 @@ impl LlmProvider for FailoverProvider {
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
self.try_providers(|provider| {
let req = request.clone();
async move { provider.complete(req).await }
})
.await
let (provider_idx, response) = self
.try_providers(|provider| {
let req = request.clone();
async move { provider.complete(req).await }
})
.await?;
self.bind_provider_to_current_task(provider_idx);
Ok(response)
}
async fn complete_with_tools(
&self,
request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
self.try_providers(|provider| {
let req = request.clone();
async move { provider.complete_with_tools(req).await }
})
.await
let (provider_idx, response) = self
.try_providers(|provider| {
let req = request.clone();
async move { provider.complete_with_tools(req).await }
})
.await?;
self.bind_provider_to_current_task(provider_idx);
Ok(response)
}
fn active_model_name(&self) -> String {
@@ -336,6 +375,14 @@ impl LlmProvider for FailoverProvider {
all_models.dedup();
Ok(all_models)
}
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
if let Some(provider_idx) = self.take_bound_provider_for_current_task() {
return self.providers[provider_idx].effective_model_name(requested_model);
}
self.providers[self.last_used.load(Ordering::Relaxed)].effective_model_name(requested_model)
}
}
#[cfg(test)]
@@ -610,6 +657,49 @@ mod tests {
assert_eq!(failover.cost_per_token(), (fallback_cost, fallback_cost));
}
// Test: model reporting is request-scoped under concurrent requests.
#[tokio::test]
async fn effective_model_name_is_request_scoped_under_concurrency() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(60),
failure_threshold: 3,
};
let primary = Arc::new(MultiCallMockProvider::fail_then_ok("primary", 1));
let fallback = Arc::new(MultiCallMockProvider::always_ok("fallback"));
let failover =
Arc::new(FailoverProvider::with_cooldown(vec![primary, fallback], config).unwrap());
let (first_done_tx, first_done_rx) = tokio::sync::oneshot::channel::<()>();
let (second_done_tx, second_done_rx) = tokio::sync::oneshot::channel::<()>();
let failover_a = Arc::clone(&failover);
let task_a = tokio::spawn(async move {
// First request: primary fails once, fallback serves.
let _ = failover_a.complete(make_request()).await.unwrap();
let _ = first_done_tx.send(());
// Wait until the second request finishes and updates global state.
let _ = second_done_rx.await;
failover_a.effective_model_name(None)
});
let failover_b = Arc::clone(&failover);
let task_b = tokio::spawn(async move {
let _ = first_done_rx.await;
// Second request: primary now succeeds.
let _ = failover_b.complete(make_request()).await.unwrap();
let model = failover_b.effective_model_name(None);
let _ = second_done_tx.send(());
model
});
let model_b = task_b.await.unwrap();
let model_a = task_a.await.unwrap();
assert_eq!(model_a, "fallback");
assert_eq!(model_b, "primary");
}
// Test: list_models aggregates from all providers.
#[tokio::test]
async fn list_models_aggregates_all() {
+5 -3
View File
@@ -462,11 +462,12 @@ fn split_messages(
#[async_trait]
impl LlmProvider for NearAiProvider {
async fn complete(&self, req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let model = req.model.unwrap_or_else(|| self.active_model_name());
let thread_id = req.metadata.get("thread_id").cloned();
let (instructions, input) = split_messages(req.messages, false);
let request = NearAiRequest {
model: self.active_model_name(),
model,
instructions,
input,
previous_response_id: None,
@@ -579,6 +580,7 @@ impl LlmProvider for NearAiProvider {
&self,
req: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
let model = req.model.unwrap_or_else(|| self.active_model_name());
let thread_id = req.metadata.get("thread_id").cloned();
// Look up chaining state for this thread
@@ -619,7 +621,7 @@ impl LlmProvider for NearAiProvider {
.collect();
let request = NearAiRequest {
model: self.active_model_name(),
model: model.clone(),
instructions: if chaining { None } else { instructions.clone() },
input,
previous_response_id: previous_response_id.clone(),
@@ -660,7 +662,7 @@ impl LlmProvider for NearAiProvider {
false,
);
let retry_request = NearAiRequest {
model: self.active_model_name(),
model,
instructions: instructions_full,
input: input_full,
previous_response_id: None,
+4 -2
View File
@@ -227,11 +227,12 @@ struct ApiModelEntry {
#[async_trait]
impl LlmProvider for NearAiChatProvider {
async fn complete(&self, req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let model = req.model.unwrap_or_else(|| self.active_model_name());
let messages: Vec<ChatCompletionMessage> =
req.messages.into_iter().map(|m| m.into()).collect();
let request = ChatCompletionRequest {
model: self.active_model_name(),
model,
messages,
temperature: req.temperature,
max_tokens: req.max_tokens,
@@ -273,6 +274,7 @@ impl LlmProvider for NearAiChatProvider {
&self,
req: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
let model = req.model.unwrap_or_else(|| self.active_model_name());
let messages: Vec<ChatCompletionMessage> =
req.messages.into_iter().map(|m| m.into()).collect();
@@ -296,7 +298,7 @@ impl LlmProvider for NearAiChatProvider {
.collect();
let request = ChatCompletionRequest {
model: self.active_model_name(),
model,
messages,
temperature: req.temperature,
max_tokens: req.max_tokens,
+28
View File
@@ -105,6 +105,8 @@ impl ChatMessage {
#[derive(Debug, Clone)]
pub struct CompletionRequest {
pub messages: Vec<ChatMessage>,
/// Optional per-request model override.
pub model: Option<String>,
pub max_tokens: Option<u32>,
pub temperature: Option<f32>,
pub stop_sequences: Option<Vec<String>>,
@@ -117,6 +119,7 @@ impl CompletionRequest {
pub fn new(messages: Vec<ChatMessage>) -> Self {
Self {
messages,
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
@@ -124,6 +127,12 @@ impl CompletionRequest {
}
}
/// Set model override.
pub fn with_model(mut self, model: impl Into<String>) -> Self {
self.model = Some(model.into());
self
}
/// Set max tokens.
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
self.max_tokens = Some(max_tokens);
@@ -188,6 +197,8 @@ pub struct ToolResult {
pub struct ToolCompletionRequest {
pub messages: Vec<ChatMessage>,
pub tools: Vec<ToolDefinition>,
/// Optional per-request model override.
pub model: Option<String>,
pub max_tokens: Option<u32>,
pub temperature: Option<f32>,
/// How to handle tool use: "auto", "required", or "none".
@@ -202,6 +213,7 @@ impl ToolCompletionRequest {
Self {
messages,
tools,
model: None,
max_tokens: None,
temperature: None,
tool_choice: None,
@@ -209,6 +221,12 @@ impl ToolCompletionRequest {
}
}
/// Set model override.
pub fn with_model(mut self, model: impl Into<String>) -> Self {
self.model = Some(model.into());
self
}
/// Set max tokens.
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
self.max_tokens = Some(max_tokens);
@@ -283,6 +301,16 @@ pub trait LlmProvider: Send + Sync {
})
}
/// Resolve which model should be reported for a given request.
///
/// Providers that ignore per-request model overrides should override this
/// and return `active_model_name()`.
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
requested_model
.map(std::borrow::ToOwned::to_owned)
.unwrap_or_else(|| self.active_model_name())
}
/// Get the currently active model name.
///
/// May differ from `model_name()` if the model was switched at runtime
+27 -1
View File
@@ -144,7 +144,8 @@ impl LlmProvider for CachedProvider {
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let key = cache_key(self.inner.model_name(), &request);
let effective_model = self.inner.effective_model_name(request.model.as_deref());
let key = cache_key(&effective_model, &request);
let now = Instant::now();
// Check cache
@@ -216,6 +217,10 @@ impl LlmProvider for CachedProvider {
self.inner.model_metadata().await
}
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
self.inner.effective_model_name(requested_model)
}
fn active_model_name(&self) -> String {
self.inner.active_model_name()
}
@@ -242,6 +247,7 @@ mod tests {
fn simple_request() -> CompletionRequest {
CompletionRequest {
messages: vec![ChatMessage::user("hello")],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
@@ -252,6 +258,7 @@ mod tests {
fn different_request() -> CompletionRequest {
CompletionRequest {
messages: vec![ChatMessage::user("goodbye")],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
@@ -378,6 +385,7 @@ mod tests {
// Add a third: should evict the oldest
let third = CompletionRequest {
messages: vec![ChatMessage::user("third")],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
@@ -396,6 +404,7 @@ mod tests {
let req = ToolCompletionRequest {
messages: vec![ChatMessage::user("use tool")],
tools: vec![],
model: None,
max_tokens: None,
temperature: None,
tool_choice: None,
@@ -444,6 +453,23 @@ mod tests {
assert!(cached.is_empty().await);
}
#[tokio::test]
async fn model_override_gets_distinct_cache_entries() {
let stub = Arc::new(StubLlm::new("cached response"));
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
let mut req_a = simple_request();
req_a.model = Some("model-a".to_string());
let mut req_b = simple_request();
req_b.model = Some("model-b".to_string());
cached.complete(req_a).await.unwrap();
cached.complete(req_b).await.unwrap();
assert_eq!(stub.calls(), 2);
assert_eq!(cached.len().await, 2);
}
#[test]
fn default_config_is_reasonable() {
let cfg = ResponseCacheConfig::default();
+24
View File
@@ -404,6 +404,16 @@ where
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
if let Some(requested_model) = request.model.as_deref()
&& requested_model != self.model_name.as_str()
{
tracing::warn!(
requested_model = requested_model,
active_model = %self.model_name,
"Per-request model override is not supported for this provider; using configured model"
);
}
let (preamble, history) = convert_messages(&request.messages);
let rig_req = build_rig_request(
@@ -439,6 +449,16 @@ where
&self,
request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
if let Some(requested_model) = request.model.as_deref()
&& requested_model != self.model_name.as_str()
{
tracing::warn!(
requested_model = requested_model,
active_model = %self.model_name,
"Per-request model override is not supported for this provider; using configured model"
);
}
let (preamble, history) = convert_messages(&request.messages);
let tools = convert_tools(&request.tools);
let tool_choice = convert_tool_choice(request.tool_choice.as_deref());
@@ -477,6 +497,10 @@ where
self.model_name.clone()
}
fn effective_model_name(&self, _requested_model: Option<&str>) -> String {
self.active_model_name()
}
fn set_model(&self, _model: &str) -> Result<(), LlmError> {
// rig-core models are baked at construction time.
// Switching requires creating a new adapter.
+1 -1
View File
@@ -1305,7 +1305,7 @@ async fn main() -> anyhow::Result<()> {
// Add web gateway channel if configured
let mut gateway_url: Option<String> = None;
if let Some(ref gw_config) = config.channels.gateway {
let mut gw = GatewayChannel::new(gw_config.clone());
let mut gw = GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&llm));
if let Some(ref ws) = workspace {
gw = gw.with_workspace(Arc::clone(ws));
}
+2
View File
@@ -143,6 +143,7 @@ async fn llm_complete(
) -> Result<Json<ProxyCompletionResponse>, StatusCode> {
let completion_req = CompletionRequest {
messages: req.messages,
model: req.model,
max_tokens: req.max_tokens,
temperature: req.temperature,
stop_sequences: req.stop_sequences,
@@ -170,6 +171,7 @@ async fn llm_complete_with_tools(
let tool_req = ToolCompletionRequest {
messages: req.messages,
tools: req.tools,
model: req.model,
max_tokens: req.max_tokens,
temperature: req.temperature,
tool_choice: req.tool_choice,
+4
View File
@@ -40,6 +40,7 @@ pub struct JobDescription {
#[derive(Debug, Serialize, Deserialize)]
pub struct ProxyCompletionRequest {
pub messages: Vec<ChatMessage>,
pub model: Option<String>,
pub max_tokens: Option<u32>,
pub temperature: Option<f32>,
pub stop_sequences: Option<Vec<String>>,
@@ -57,6 +58,7 @@ pub struct ProxyCompletionResponse {
pub struct ProxyToolCompletionRequest {
pub messages: Vec<ChatMessage>,
pub tools: Vec<ToolDefinition>,
pub model: Option<String>,
pub max_tokens: Option<u32>,
pub temperature: Option<f32>,
pub tool_choice: Option<String>,
@@ -210,6 +212,7 @@ impl WorkerHttpClient {
) -> Result<CompletionResponse, WorkerError> {
let proxy_req = ProxyCompletionRequest {
messages: request.messages.clone(),
model: request.model.clone(),
max_tokens: request.max_tokens,
temperature: request.temperature,
stop_sequences: request.stop_sequences.clone(),
@@ -236,6 +239,7 @@ impl WorkerHttpClient {
let proxy_req = ProxyToolCompletionRequest {
messages: request.messages.clone(),
tools: request.tools.clone(),
model: request.model.clone(),
max_tokens: request.max_tokens,
temperature: request.temperature,
tool_choice: request.tool_choice.clone(),