mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-09-01 17:19:24 +00:00
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:
co-authored by
Illia Polosukhin
parent
89fdd81420
commit
ccf60055f4
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user