mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +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
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
+1
-1
@@ -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));
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user